GoodWin/Deep-Multi-scale
0
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 