CoTyle / models /lakonlab /datasets /builder.py
liuhuijie
update
619344d
Raw
History Blame
1.96 kB
import warnings
from functools import partial
from mmcv.parallel import collate
from mmcv.runner import get_dist_info
from mmcv.utils import TORCH_VERSION, digit_version
from torch.utils.data import DataLoader
from mmgen.datasets.builder import worker_init_fn
from .samplers import DistributedSampler
def build_dataloader(dataset,
samples_per_gpu,
workers_per_gpu,
num_gpus=1,
dist=True,
shuffle=True,
seed=None,
persistent_workers=False,
sampler=None,
**kwargs):
rank, world_size = get_dist_info()
if dist:
assert sampler is None, 'sampler is not supported in distributed mode'
sampler = DistributedSampler(
dataset,
world_size,
rank,
shuffle=shuffle,
samples_per_gpu=samples_per_gpu,
seed=seed)
shuffle = False
batch_size = samples_per_gpu
num_workers = workers_per_gpu
else:
batch_size = num_gpus * samples_per_gpu
num_workers = num_gpus * workers_per_gpu
init_fn = partial(
worker_init_fn, num_workers=num_workers, rank=rank,
seed=seed) if seed is not None else None
if (digit_version(TORCH_VERSION) >= digit_version('1.7.0')
and TORCH_VERSION != 'parrots'):
kwargs['persistent_workers'] = persistent_workers
elif persistent_workers is True:
warnings.warn('persistent_workers is invalid because your pytorch '
'version is lower than 1.7.0')
data_loader = DataLoader(
dataset,
batch_size=batch_size,
sampler=sampler,
num_workers=num_workers,
collate_fn=partial(collate, samples_per_gpu=samples_per_gpu),
shuffle=shuffle,
worker_init_fn=init_fn,
**kwargs)
return data_loader