(dataset_len, train_size, val_size, test_size, seed, filename=None, splits=None)
| 55 | |
| 56 | |
| 57 | def make_splits(dataset_len, train_size, val_size, test_size, seed, filename=None, splits=None): |
| 58 | if splits is not None: |
| 59 | splits = np.load(splits) |
| 60 | idx_train = splits["idx_train"] |
| 61 | idx_val = splits["idx_val"] |
| 62 | idx_test = splits["idx_test"] |
| 63 | else: |
| 64 | idx_train, idx_val, idx_test = train_val_test_split(dataset_len, train_size, val_size, test_size, seed) |
| 65 | |
| 66 | if filename is not None: |
| 67 | np.savez(filename, idx_train=idx_train, idx_val=idx_val, idx_test=idx_test) |
| 68 | |
| 69 | return torch.from_numpy(idx_train), torch.from_numpy(idx_val), torch.from_numpy(idx_test) |
| 70 | |
| 71 | |
| 72 | class LoadFromFile(argparse.Action): |
no test coverage detected