CoolFace
Apppublic

sczhou/CodeFormer

sourceHugging Faceupdated 4mo agoView on Hugging Face
2.4klikes
__init__.py101 linesDownload Raw Back to data
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