| 176 | } |
| 177 | |
| 178 | void 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)); |
nothing calls this directly
no test coverage detected