CoolFace
Apppublic

MLBench/ReaLens

sourceHugging Facemitupdated 8mo agoView on Hugging Face
0likes
unaligned_dataset.py72 linesDownload Raw Back to data
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