MLBench/ReaLens
0
1import os2from data.base_dataset import BaseDataset, get_transform3from data.image_folder import make_dataset4from PIL import Image5import random6 7 8class UnalignedDataset(BaseDataset):9 """10 This dataset class can load unaligned/unpaired datasets.11 12 It requires two directories to host training images from domain A '/path/to/data/trainA'13 and from domain B '/path/to/data/trainB' respectively.14 You can train the model with the dataset flag '--dataroot /path/to/data'.15 Similarly, you need to prepare two directories:16 '/path/to/data/testA' and '/path/to/data/testB' during test time.17 """18 19 def __init__(self, opt):20 """Initialize this dataset class.21 22 Parameters:23 opt (Option class) -- stores all the experiment flags; needs to be a subclass of BaseOptions24 """25 BaseDataset.__init__(self, opt)26 self.dir_A = os.path.join(opt.dataroot, opt.phase + "A") # create a path '/path/to/data/trainA'27 self.dir_B = os.path.join(opt.dataroot, opt.phase + "B") # create a path '/path/to/data/trainB'28 29 self.A_paths = sorted(make_dataset(self.dir_A, opt.max_dataset_size)) # load images from '/path/to/data/trainA'30 self.B_paths = sorted(make_dataset(self.dir_B, opt.max_dataset_size)) # load images from '/path/to/data/trainB'31 self.A_size = len(self.A_paths) # get the size of dataset A32 self.B_size = len(self.B_paths) # get the size of dataset B33 btoA = self.opt.direction == "BtoA"34 input_nc = self.opt.output_nc if btoA else self.opt.input_nc # get the number of channels of input image35 output_nc = self.opt.input_nc if btoA else self.opt.output_nc # get the number of channels of output image36 self.transform_A = get_transform(self.opt, grayscale=(input_nc == 1))37 self.transform_B = get_transform(self.opt, grayscale=(output_nc == 1))38 39 def __getitem__(self, index):40 """Return a data point and its metadata information.41 42 Parameters:43 index (int) -- a random integer for data indexing44 45 Returns a dictionary that contains A, B, A_paths and B_paths46 A (tensor) -- an image in the input domain47 B (tensor) -- its corresponding image in the target domain48 A_paths (str) -- image paths49 B_paths (str) -- image paths50 """51 A_path = self.A_paths[index % self.A_size] # make sure index is within then range52 if self.opt.serial_batches: # make sure index is within then range53 index_B = index % self.B_size54 else: # randomize the index for domain B to avoid fixed pairs.55 index_B = random.randint(0, self.B_size - 1)56 B_path = self.B_paths[index_B]57 A_img = Image.open(A_path).convert("RGB")58 B_img = Image.open(B_path).convert("RGB")59 # apply image transformation60 A = self.transform_A(A_img)61 B = self.transform_B(B_img)62 63 return {"A": A, "B": B, "A_paths": A_path, "B_paths": B_path}64 65 def __len__(self):66 """Return the total number of images in the dataset.67 68 As we have two datasets with potentially different number of images,69 we take a maximum of70 """71 return max(self.A_size, self.B_size)72 