should be starting from X, and Y X_{}_trainval, X_{}_test, Y_{}_trainval, Y_{}_test
(*args,**kwargs)
| 26 | |
| 27 | # def dataloader(X_train_val_set, X_test_set, Y_train_val_set, Y_test_set): |
| 28 | def dataloader(*args,**kwargs): |
| 29 | """ |
| 30 | should be starting from X, and Y |
| 31 | X_{}_trainval, X_{}_test, Y_{}_trainval, Y_{}_test |
| 32 | """ |
| 33 | temp_loader = {} |
| 34 | for name, input in kwargs.items(): |
| 35 | # First, format |
| 36 | input_name = name.split('_') |
| 37 | if input_name[0].startswith('Y'): |
| 38 | input = input.astype('float32') |
| 39 | input = torch.from_numpy(input) |
| 40 | input = input.unsqueeze(1) |
| 41 | temp_loader[name] = input |
| 42 | |
| 43 | elif input_name[0].startswith('X'): |
| 44 | #input is tabular format |
| 45 | input = input.astype('float32') |
| 46 | input = torch.from_numpy(input) |
| 47 | temp_loader[name] = input |
| 48 | |
| 49 | elif input_name[0].startswith('dummy'): |
| 50 | input = input.astype('float32') |
| 51 | input = torch.from_numpy(input) |
| 52 | temp_loader[name] = input |
| 53 | |
| 54 | # Second, to tensordataset |
| 55 | temp_loader_trainval, temp_loader_test = [], [] |
| 56 | for key, val in temp_loader.items(): |
| 57 | load_name = key.split('_') |
| 58 | # train_val_dataset |
| 59 | if load_name[-1].endswith('trainval'): |
| 60 | temp_loader_trainval.append(val) |
| 61 | elif load_name[-1].endswith('test'): |
| 62 | temp_loader_test.append(val) |
| 63 | |
| 64 | train_val_dataset = torch.utils.data.TensorDataset(*temp_loader_trainval) |
| 65 | test_dataset = torch.utils.data.TensorDataset(*temp_loader_test) |
| 66 | test_loader = DataLoader(test_dataset, batch_size=256,shuffle = False, sampler=sampler.SequentialSampler(test_dataset)) |
| 67 | #list(BatchSampler(SequentialSampler(range(10)), batch_size=3, drop_last=False)) |
| 68 | return train_val_dataset, test_loader |
| 69 | |
| 70 | |
| 71 | def dataloader_graph(*args,**kwargs): |