(config, is_train = True)
| 29 | random.seed(123) |
| 30 | |
| 31 | def get_dataloader(config, is_train = True): |
| 32 | if is_train: |
| 33 | if config.training.get('dpo', False): |
| 34 | datasets = {} |
| 35 | datasets['dpo'] = DpoDataset(cfg=config, train=True) |
| 36 | elif config.training.get('kto', False): |
| 37 | datasets = {} |
| 38 | datasets['kto'] = KTODataset(cfg=config, train=True) |
| 39 | else: |
| 40 | datasets = {'mix': MixDataset(cfg=config, train=True)} |
| 41 | else: |
| 42 | dataset1 = h36m( |
| 43 | cfg=config, |
| 44 | ann_file=config.dataset.set_list[0].test_set, |
| 45 | root = config.dataset.set_list[0].root, |
| 46 | train=False) |
| 47 | dataset2 = pw3d( |
| 48 | cfg=config, |
| 49 | ann_file=config.dataset.set_list[3].test_set, |
| 50 | root = config.dataset.set_list[3].root, |
| 51 | train=False) |
| 52 | datasets = {'h36m': dataset1, '3dpw':dataset2, } |
| 53 | dataloaders = {} |
| 54 | samplers = {} |
| 55 | for key, dataset in datasets.items(): |
| 56 | sampler = torch.utils.data.distributed.DistributedSampler(dataset, shuffle=True) # [DEBUG] |
| 57 | shuffle = False |
| 58 | batch_size = config.sampling.batch_size |
| 59 | if is_train: |
| 60 | sampler = torch.utils.data.distributed.DistributedSampler(dataset) # [DEBUG] |
| 61 | shuffle = (sampler is None) |
| 62 | batch_size = config.training.batch_size |
| 63 | dataloader = torch.utils.data.DataLoader( |
| 64 | dataset, |
| 65 | batch_size=batch_size, |
| 66 | shuffle=shuffle, |
| 67 | num_workers=config.dataset.workers, |
| 68 | sampler=sampler, |
| 69 | pin_memory=True, |
| 70 | drop_last=False, |
| 71 | worker_init_fn=_init_fn, |
| 72 | ) |
| 73 | dataloaders[key] = dataloader |
| 74 | samplers[key] = sampler |
| 75 | logging.info(f"dataset [{key}] length is {len(dataset)}") |
| 76 | return dataloaders, datasets, samplers |
| 77 | |
| 78 | |
| 79 | def get_optimizer(config, parameters,lr): |
no test coverage detected