Spaces:
Running on Zero
Running on Zero
File size: 1,027 Bytes
2680bd5 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 | from torch.utils.data import DataLoader, Dataset, Subset
from torch.utils.data.distributed import DistributedSampler
from .. import dist as dist_utils
from ..config import DataLoaderConfig
def get_dataloader(dataset:Dataset, cfg:DataLoaderConfig):
is_distributed = dist_utils.is_dist_avail_and_initialized()
sampler = DistributedSampler(dataset, shuffle=cfg.shuffle) if is_distributed else None
if sampler is not None:
shuffle = False
else:
shuffle = cfg.shuffle
if cfg.batch_size % dist_utils.get_world_size() != 0:
raise ValueError(f"Batch size {cfg.batch_size} must be divisible by world size {dist_utils.get_world_size()}.")
batch_size = cfg.batch_size // dist_utils.get_world_size()
loader = DataLoader(
dataset,
batch_size=batch_size,
num_workers=cfg.num_workers,
shuffle=shuffle,
sampler=sampler,
collate_fn=dataset.dataset.collate_fn if isinstance(dataset, Subset) else dataset.collate_fn,
)
return loader
|