MCPcopy Create free account
hub / github.com/MotrixLab/ADHMR / get_dataloader

Function get_dataloader

ADHMR/lib/utils/function.py:31–76  ·  view source on GitHub ↗
(config, is_train = True)

Source from the content-addressed store, hash-verified

29 random.seed(123)
30
31def 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
79def get_optimizer(config, parameters,lr):

Callers 4

trainMethod · 0.90
validateMethod · 0.90
trainMethod · 0.90
validateMethod · 0.90

Calls 6

DpoDatasetClass · 0.90
KTODatasetClass · 0.90
MixDatasetClass · 0.90
infoMethod · 0.80
getMethod · 0.45
itemsMethod · 0.45

Tested by

no test coverage detected