Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
18 changes: 18 additions & 0 deletions numexpr/interpreter.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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];
Expand Down
9 changes: 9 additions & 0 deletions numexpr/numexpr_object.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -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);
}

Expand All @@ -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) { \
Expand All @@ -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;
Expand Down
1 change: 1 addition & 0 deletions numexpr/numexpr_object.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -27,6 +27,7 @@ struct NumExprObject
int n_inputs;
int n_constants;
int n_temps;
void *run_lock;
};

extern PyTypeObject NumExprType;
Expand Down
36 changes: 36 additions & 0 deletions numexpr/tests/test_numexpr.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
Loading