| 122 | |
| 123 | # continuum iterator ######################################################### |
| 124 | def load_datasets(args): |
| 125 | print("path", args.data_path + '/' + args.data_file) |
| 126 | d_tr, d_te = torch.load(args.data_path + '/' + args.data_file) |
| 127 | n_inputs = d_tr[0][1].size(1) |
| 128 | n_outputs = 0 |
| 129 | for i in range(len(d_tr)): |
| 130 | n_outputs = max(n_outputs, d_tr[i][2].max().item()) |
| 131 | n_outputs = max(n_outputs, d_te[i][2].max().item()) |
| 132 | return d_tr, d_te, n_inputs, n_outputs + 1, len(d_tr) |
| 133 | |
| 134 | |
| 135 | class Continuum: |