MCPcopy Create free account
hub / github.com/LBANN/lbann / _setup_data_reader

Function _setup_data_reader

python/lbann/core/evaluate.py:129–166  ·  view source on GitHub ↗
(inputs: npt.NDArray, workdir: str, training: bool)

Source from the content-addressed store, hash-verified

127
128
129def _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
169def _collect_outputs(output_names: List[str], workdir: str, dtype: np.dtype,

Callers 1

evaluateFunction · 0.85

Calls 1

saveMethod · 0.45

Tested by

no test coverage detected