(dataset, dataset_dir, distributed, T)
| 301 | |
| 302 | |
| 303 | def 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 | |
| 331 | def main(args): |
no test coverage detected