| 512 | |
| 513 | |
| 514 | def 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 | |
| 546 | def train(gpu, opt, output_dir, noises_init): |