MCPcopy Create free account
hub / github.com/bic-L/MaxFormer / load_data

Function load_data

event/train.py:303–328  ·  view source on GitHub ↗
(dataset, dataset_dir, distributed, T)

Source from the content-addressed store, hash-verified

301
302
303def load_data(dataset, dataset_dir, distributed, T):
304 # Data loading code
305 print("Loading data")
306
307 st = time.time()
308
309 if dataset == 'cifar10dvs':
310 origin_set = cifar10_dvs.CIFAR10DVS(root=dataset_dir, data_type='frame', frames_number=T, split_by='number')
311 dataset_train, dataset_test = split_to_train_test_set(0.9, origin_set, 10)
312
313 elif dataset == 'dvsgesture':
314 # dataset_train, dataset_test = Dvs128Gesture(root="/home/dataset/DvsGesture", resolution=(128, 128))
315 dataset_train = DVS128Gesture(root=dataset_dir, train=True, data_type='frame', frames_number=T, split_by='number')
316 dataset_test = DVS128Gesture(root=dataset_dir, train=False, data_type='frame', frames_number=T, split_by='number')
317
318 print("Took", time.time() - st)
319
320 print("Creating data loaders")
321 if distributed:
322 train_sampler = torch.utils.data.distributed.DistributedSampler(dataset_train)
323 test_sampler = torch.utils.data.distributed.DistributedSampler(dataset_test)
324 else:
325 train_sampler = torch.utils.data.RandomSampler(dataset_train)
326 test_sampler = torch.utils.data.SequentialSampler(dataset_test)
327
328 return dataset_train, dataset_test, train_sampler, test_sampler
329
330
331def main(args):

Callers 1

mainFunction · 0.85

Calls 2

split_to_train_test_setFunction · 0.85
printFunction · 0.70

Tested by

no test coverage detected