diff --git a/numexpr/interpreter.cpp b/numexpr/interpreter.cpp index f377721..6d1393c 100644 --- a/numexpr/interpreter.cpp +++ b/numexpr/interpreter.cpp @@ -18,6 +18,22 @@ #include "complex_functions.hpp" #include "interpreter.hpp" #include "numexpr_object.hpp" + +class NumExprRunLock { +public: + explicit NumExprRunLock(PyThread_type_lock lock) : lock(lock) { + Py_BEGIN_ALLOW_THREADS; + PyThread_acquire_lock(lock, WAIT_LOCK); + Py_END_ALLOW_THREADS; + } + + ~NumExprRunLock() { + PyThread_release_lock(lock); + } + +private: + PyThread_type_lock lock; +}; #include "bespoke_functions.hpp" #ifdef _MSC_VER @@ -1059,6 +1075,8 @@ NumExpr_run(NumExprObject *self, PyObject *args, PyObject *kwds) int is_reduction = 0; bool reduction_outer_loop = false, need_output_buffering = false, full_reduction = false; + NumExprRunLock run_lock((PyThread_type_lock)self->run_lock); + // To specify axes when doing a reduction int op_axes_values[NE_MAXARGS][NPY_MAXDIMS], op_axes_reduction_values[NE_MAXARGS]; diff --git a/numexpr/numexpr_object.cpp b/numexpr/numexpr_object.cpp index b6e2f9c..6da28d2 100644 --- a/numexpr/numexpr_object.cpp +++ b/numexpr/numexpr_object.cpp @@ -45,6 +45,9 @@ NumExpr_dealloc(NumExprObject *self) PyMem_Del(self->rawmem); PyMem_Del(self->memsteps); PyMem_Del(self->memsizes); + if (self->run_lock != NULL) { + PyThread_free_lock((PyThread_type_lock)self->run_lock); + } Py_TYPE(self)->tp_free((PyObject*)self); } @@ -53,6 +56,7 @@ NumExpr_new(PyTypeObject *type, PyObject *args, PyObject *kwds) { NumExprObject *self = (NumExprObject *)type->tp_alloc(type, 0); if (self != NULL) { + self->run_lock = NULL; #define INIT_WITH(name, object) \ self->name = object; \ if (!self->name) { \ @@ -76,6 +80,11 @@ NumExpr_new(PyTypeObject *type, PyObject *args, PyObject *kwds) self->n_inputs = 0; self->n_constants = 0; self->n_temps = 0; + self->run_lock = (void *)PyThread_allocate_lock(); + if (self->run_lock == NULL) { + Py_DECREF(self); + return PyErr_NoMemory(); + } #undef INIT_WITH } return (PyObject *)self; diff --git a/numexpr/numexpr_object.hpp b/numexpr/numexpr_object.hpp index cd5bb44..4f523c7 100644 --- a/numexpr/numexpr_object.hpp +++ b/numexpr/numexpr_object.hpp @@ -27,6 +27,7 @@ struct NumExprObject int n_inputs; int n_constants; int n_temps; + void *run_lock; }; extern PyTypeObject NumExprType; diff --git a/numexpr/tests/test_numexpr.py b/numexpr/tests/test_numexpr.py index 5fdeebb..3436d1d 100644 --- a/numexpr/tests/test_numexpr.py +++ b/numexpr/tests/test_numexpr.py @@ -1448,6 +1448,42 @@ def work(n): for t in threads: t.join() + @pytest.mark.thread_unsafe + def test_shared_compiled_expression(self): + import threading + + expression = NumExpr("where(norm == 0.0, dummy, signal / norm)") + size = 200_000 + signal = np.random.random(size) + norm = np.random.random(size) + dummy = np.float64(0.0) + expected = signal / norm + barrier = threading.Barrier(4) + errors = [] + + def work(): + try: + for _ in range(10): + barrier.wait() + result = expression(dummy, norm, signal) + assert_allclose(result, expected) + except BaseException as error: + errors.append(error) + barrier.abort() + + old_nthreads = numexpr.set_num_threads(1) + try: + threads = [threading.Thread(target=work) for _ in range(4)] + for thread in threads: + thread.start() + for thread in threads: + thread.join() + finally: + numexpr.set_num_threads(old_nthreads) + + if errors: + raise errors[0] + def test_thread_safety(self): """ Expected output