MCPcopy Create free account
hub / github.com/Zhiyuan-R/Tiger-Diffusion / get_dataloader

Function get_dataloader

train_generation.py:514–543  ·  view source on GitHub ↗
(opt, train_dataset, test_dataset=None)

Source from the content-addressed store, hash-verified

512
513
514def get_dataloader(opt, train_dataset, test_dataset=None):
515
516 if opt.distribution_type == 'multi':
517 train_sampler = torch.utils.data.distributed.DistributedSampler(
518 train_dataset,
519 num_replicas=opt.world_size,
520 rank=opt.rank
521 )
522 if test_dataset is not None:
523 test_sampler = torch.utils.data.distributed.DistributedSampler(
524 test_dataset,
525 num_replicas=opt.world_size,
526 rank=opt.rank
527 )
528 else:
529 test_sampler = None
530 else:
531 train_sampler = None
532 test_sampler = None
533
534 train_dataloader = torch.utils.data.DataLoader(train_dataset, batch_size=opt.bs,sampler=train_sampler,
535 shuffle=train_sampler is None, num_workers=int(opt.workers), drop_last=True)
536
537 if test_dataset is not None:
538 test_dataloader = torch.utils.data.DataLoader(train_dataset, batch_size=opt.bs,sampler=test_sampler,
539 shuffle=False, num_workers=int(opt.workers), drop_last=False)
540 else:
541 test_dataloader = None
542
543 return train_dataloader, test_dataloader, train_sampler, test_sampler
544
545
546def train(gpu, opt, output_dir, noises_init):

Callers 1

trainFunction · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected