CoolFace
Apppublic

GoodWin/Deep-Multi-scale

sourceHugging Faceupdated 5y agoView on Hugging Face
0likes
__init__.py81 linesDownload Raw Back to data
1# -- coding: utf-8 --2import importlib3import torch.utils.data4from data.base_data_loader import BaseDataLoader5from data.base_dataset import BaseDataset6 7 8def find_dataset_using_name(dataset_name):9 10    # Given the option --dataset_mode [datasetname],11    # the file "data/datasetname_dataset.py"12    # will be imported.13    dataset_filename = "data." + dataset_name + "_dataset"14    datasetlib = importlib.import_module(dataset_filename)15 16    # In the file, the class called DatasetNameDataset() will17    # be instantiated. It has to be a subclass of BaseDataset,18    # and it is case-insensitive.19    dataset = None20    target_dataset_name = dataset_name.replace('_', '') + 'dataset'21    for name, cls in datasetlib.__dict__.items():22        if name.lower() == target_dataset_name.lower() \23           and issubclass(cls, BaseDataset):24            dataset = cls25            26    if dataset is None:27        print("In %s.py, there should be a subclass of BaseDataset with class name that matches %s in lowercase." % (dataset_filename, target_dataset_name))28        exit(0)29 30    return dataset31 32 33def get_option_setter(dataset_name):    34    dataset_class = find_dataset_using_name(dataset_name)35    return dataset_class.modify_commandline_options36 37 38def create_dataset(opt):39    dataset = find_dataset_using_name(opt.dataset_mode)40    instance = dataset()41    instance.initialize(opt)42    print("dataset [%s] was created" % (instance.name()))43    return instance44 45 46def CreateDataLoader(opt):47    data_loader = CustomDatasetDataLoader()48    data_loader.initialize(opt)49    return data_loader50 51 52# Wrapper class of Dataset class that performs53# multi-threaded data loading54class CustomDatasetDataLoader(BaseDataLoader):55    def name(self):56        return 'CustomDatasetDataLoader'57 58    def initialize(self, opt):59        BaseDataLoader.initialize(self, opt)60        self.dataset = create_dataset(opt)61        self.dataloader = torch.utils.data.DataLoader(62            self.dataset,63            batch_size=opt.batchSize,64            shuffle=not opt.serial_batches,65            num_workers=int(opt.nThreads))66            # DataLoader(dataset, batch_size=1, shuffle=False, sampler=None, num_workers=0, collate_fn=default_collate, pin_memory=False, drop_last=False)67            # 加载的数据集,DataSet对象,是否打乱,样本抽样,使用多线程加载的进程数,0表示不使用多线程,如何将多样本数据拼接成一个batch,是否将数据保存到pin memory,dataset种数据可数可能不是\68            # 一个batch_size的整数倍,drop_last 为True将多出来不足一个batch的数据丢弃69 70    def load_data(self):71        return self72 73    def __len__(self):74        return min(len(self.dataset), self.opt.max_dataset_size)75 76    def __iter__(self):77        for i, data in enumerate(self.dataloader):78            if i * self.opt.batchSize >= self.opt.max_dataset_size:79                break80            yield data81