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

Function setup_func

ci_test/common_python/test_util.py:105–138  ·  view source on GitHub ↗
(lbann, weekly)

Source from the content-addressed store, hash-verified

103 file = inspect.getfile(f)
104
105 def setup_func(lbann, weekly):
106 # Get minibatch size from tensor
107 mini_batch_size = tester.input_tensor.shape[0]
108
109 # Save combined input/reference data to file
110 work_dir = _get_work_dir(file)
111 os.makedirs(work_dir, exist_ok=True)
112 if tester.reference_tensor is not None:
113 flat_inp = tester.input_tensor.reshape(mini_batch_size, -1)
114 flat_ref = tester.reference_tensor.reshape(
115 mini_batch_size, -1)
116 np.save(os.path.join(work_dir, 'data.npy'),
117 np.concatenate((flat_inp, flat_ref), axis=1))
118 else:
119 np.save(os.path.join(work_dir, 'data.npy'),
120 tester.input_tensor.reshape(mini_batch_size, -1))
121
122 # Setup data reader
123 data_reader = lbann.reader_pb2.DataReader()
124 data_reader.reader.extend([
125 tools.create_python_data_reader(
126 lbann, single_tensor_data_reader.__file__,
127 'get_sample', 'num_samples', 'sample_dims', 'train')
128 ])
129 if not train:
130 data_reader.reader.extend([
131 tools.create_python_data_reader(
132 lbann, single_tensor_data_reader.__file__,
133 'get_sample', 'num_samples', 'sample_dims', 'test')
134 ])
135
136 trainer = lbann.Trainer(mini_batch_size)
137 optimizer = lbann.SGD(learn_rate=0)
138 return trainer, model, data_reader, optimizer, None # Don't request any specific number of nodes
139
140 test = tools.create_tests(setup_func, file, **decorator_kwargs)[0]
141 cluster = kwargs.get('cluster', None)

Callers 1

test_funcFunction · 0.85

Calls 2

_get_work_dirFunction · 0.85
saveMethod · 0.45

Tested by

no test coverage detected