sczhou/CodeFormer
2.4k
1import importlib2import numpy as np3import random4import torch5import torch.utils.data6from copy import deepcopy7from functools import partial8from os import path as osp9 10from basicsr.data.prefetch_dataloader import PrefetchDataLoader11from basicsr.utils import get_root_logger, scandir12from basicsr.utils.dist_util import get_dist_info13from basicsr.utils.registry import DATASET_REGISTRY14 15__all__ = ['build_dataset', 'build_dataloader']16 17# automatically scan and import dataset modules for registry18# scan all the files under the data folder with '_dataset' in file names19data_folder = osp.dirname(osp.abspath(__file__))20dataset_filenames = [osp.splitext(osp.basename(v))[0] for v in scandir(data_folder) if v.endswith('_dataset.py')]21# import all the dataset modules22_dataset_modules = [importlib.import_module(f'basicsr.data.{file_name}') for file_name in dataset_filenames]23 24 25def build_dataset(dataset_opt):26 """Build dataset from options.27 28 Args:29 dataset_opt (dict): Configuration for dataset. It must constain:30 name (str): Dataset name.31 type (str): Dataset type.32 """33 dataset_opt = deepcopy(dataset_opt)34 dataset = DATASET_REGISTRY.get(dataset_opt['type'])(dataset_opt)35 logger = get_root_logger()36 logger.info(f'Dataset [{dataset.__class__.__name__}] - {dataset_opt["name"]} ' 'is built.')37 return dataset38 39 40def build_dataloader(dataset, dataset_opt, num_gpu=1, dist=False, sampler=None, seed=None):41 """Build dataloader.42 43 Args:44 dataset (torch.utils.data.Dataset): Dataset.45 dataset_opt (dict): Dataset options. It contains the following keys:46 phase (str): 'train' or 'val'.47 num_worker_per_gpu (int): Number of workers for each GPU.48 batch_size_per_gpu (int): Training batch size for each GPU.49 num_gpu (int): Number of GPUs. Used only in the train phase.50 Default: 1.51 dist (bool): Whether in distributed training. Used only in the train52 phase. Default: False.53 sampler (torch.utils.data.sampler): Data sampler. Default: None.54 seed (int | None): Seed. Default: None55 """56 phase = dataset_opt['phase']57 rank, _ = get_dist_info()58 if phase == 'train':59 if dist: # distributed training60 batch_size = dataset_opt['batch_size_per_gpu']61 num_workers = dataset_opt['num_worker_per_gpu']62 else: # non-distributed training63 multiplier = 1 if num_gpu == 0 else num_gpu64 batch_size = dataset_opt['batch_size_per_gpu'] * multiplier65 num_workers = dataset_opt['num_worker_per_gpu'] * multiplier66 dataloader_args = dict(67 dataset=dataset,68 batch_size=batch_size,69 shuffle=False,70 num_workers=num_workers,71 sampler=sampler,72 drop_last=True)73 if sampler is None:74 dataloader_args['shuffle'] = True75 dataloader_args['worker_init_fn'] = partial(76 worker_init_fn, num_workers=num_workers, rank=rank, seed=seed) if seed is not None else None77 elif phase in ['val', 'test']: # validation78 dataloader_args = dict(dataset=dataset, batch_size=1, shuffle=False, num_workers=0)79 else:80 raise ValueError(f'Wrong dataset phase: {phase}. ' "Supported ones are 'train', 'val' and 'test'.")81 82 dataloader_args['pin_memory'] = dataset_opt.get('pin_memory', False)83 84 prefetch_mode = dataset_opt.get('prefetch_mode')85 if prefetch_mode == 'cpu': # CPUPrefetcher86 num_prefetch_queue = dataset_opt.get('num_prefetch_queue', 1)87 logger = get_root_logger()88 logger.info(f'Use {prefetch_mode} prefetch dataloader: ' f'num_prefetch_queue = {num_prefetch_queue}')89 return PrefetchDataLoader(num_prefetch_queue=num_prefetch_queue, **dataloader_args)90 else:91 # prefetch_mode=None: Normal dataloader92 # prefetch_mode='cuda': dataloader for CUDAPrefetcher93 return torch.utils.data.DataLoader(**dataloader_args)94 95 96def worker_init_fn(worker_id, num_workers, rank, seed):97 # Set the worker seed to num_workers * rank + worker_id + seed98 worker_seed = num_workers * rank + worker_id + seed99 np.random.seed(worker_seed)100 random.seed(worker_seed)101 