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

Method setup

src/data_readers/data_reader_python.cpp:178–332  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

176}
177
178void python_reader::setup(int num_io_threads,
179 observer_ptr<thread_pool> io_thread_pool)
180{
181 generic_data_reader::setup(num_io_threads, io_thread_pool);
182
183 // Acquire Python GIL
184 python::global_interpreter_lock gil;
185
186 // Import modules
187 python::object main_module = PyImport_ImportModule("__main__");
188 python::object ctypes_module = PyImport_ImportModule("ctypes");
189 python::object multiprocessing_module =
190 PyImport_ImportModule("multiprocessing");
191
192 // Stop process pool if needed
193 if (m_process_pool != nullptr) {
194 PyObject_CallMethod(m_process_pool, "terminate", nullptr);
195 m_process_pool = nullptr;
196 }
197
198 // Allocate shared memory array
199 /// @todo Figure out more robust way to get max mini-batch size
200 const El::Int sample_size = get_linearized_data_size();
201 const El::Int mini_batch_size = get_trainer().get_max_mini_batch_size();
202 std::string datatype_typecode;
203 switch (sizeof(DataType)) {
204 case 4:
205 datatype_typecode = "f";
206 break;
207 case 8:
208 datatype_typecode = "d";
209 break;
210 default:
211 LBANN_ERROR("invalid data type for Python data reader "
212 "(only float and double are supported)");
213 }
214 m_shared_memory_array = PyObject_CallMethod(multiprocessing_module,
215 "RawArray",
216 "(s, l)",
217 datatype_typecode.c_str(),
218 sample_size * mini_batch_size);
219
220 // Get address of shared memory buffer
221 python::object shared_memory_ptr =
222 PyObject_CallMethod(ctypes_module,
223 "addressof",
224 "(O)",
225 m_shared_memory_array.get());
226 m_shared_memory_array_ptr =
227 reinterpret_cast<DataType*>(PyLong_AsLong(shared_memory_ptr));
228
229 // Create global variables in Python
230 // Note: The static counter makes sure variable names are unique.
231 static El::Int instance_id = 0;
232 instance_id++;
233 const std::string sample_func_name =
234 ("_DATA_READER_PYTHON_CPP_sample_function_wrapper" +
235 std::to_string(instance_id));

Callers

nothing calls this directly

Calls 5

check_errorFunction · 0.85
to_stringFunction · 0.70
setupFunction · 0.50
getMethod · 0.45

Tested by

no test coverage detected