| 127 | |
| 128 | |
| 129 | def _setup_data_reader(inputs: npt.NDArray, workdir: str, training: bool): |
| 130 | # Save inputs |
| 131 | if len(inputs.shape) == 1: # Minibatch dimension must exist |
| 132 | inputs = inputs.reshape(1, inputs.shape[0]) |
| 133 | elif len(inputs.shape) > 2: |
| 134 | warnings.warn( |
| 135 | f'{len(inputs.shape)}-dimensional tensor given, all dimensions ' |
| 136 | 'beyond the second one will be flattened') |
| 137 | inputs = inputs.reshape(inputs.shape[0], -1) |
| 138 | |
| 139 | np.save(os.path.join(workdir, 'data.npy'), inputs) |
| 140 | |
| 141 | # Construct protobuf message for data reader |
| 142 | file_name = os.path.realpath(single_tensor_data_reader.__file__) |
| 143 | dir_name = os.path.dirname(file_name) |
| 144 | module_name = os.path.splitext(os.path.basename(file_name))[0] |
| 145 | reader = lbann.reader_pb2.Reader() |
| 146 | reader.name = 'python' |
| 147 | reader.role = 'test' |
| 148 | reader.shuffle = False |
| 149 | reader.fraction_of_data_to_use = 1.0 |
| 150 | reader.python.module = module_name |
| 151 | reader.python.module_dir = dir_name |
| 152 | reader.python.sample_function = 'get_sample' |
| 153 | reader.python.num_samples_function = 'num_samples' |
| 154 | reader.python.sample_dims_function = 'sample_dims' |
| 155 | |
| 156 | data_reader = lbann.reader_pb2.DataReader() |
| 157 | |
| 158 | if training: |
| 159 | reader.role = 'train' |
| 160 | data_reader.reader.extend([reader]) |
| 161 | else: |
| 162 | train_reader = copy.deepcopy(reader) |
| 163 | train_reader.role = 'train' |
| 164 | data_reader.reader.extend([train_reader, reader]) |
| 165 | |
| 166 | return data_reader |
| 167 | |
| 168 | |
| 169 | def _collect_outputs(output_names: List[str], workdir: str, dtype: np.dtype, |