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
53 changes: 38 additions & 15 deletions numexpr/interpreter.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -765,6 +765,32 @@ vm_engine_iter_outer_reduce_task(NpyIter *iter, npy_intp *memsteps,
return 0;
}

/* Single-task version of the VM engine, for when there is no worker pool.
`params` is taken by value: out_buffer is rewritten below. */
static int
vm_engine_iter_serial(NpyIter *iter, vm_params params,
bool need_output_buffering, int *pc_error,
char **errmsg)
{
int r;

// Allocate memory for output buffering if needed
vector<char> out_buffer(need_output_buffering ?
(params.memsizes[0] * BLOCK_SIZE1) : 0);
params.out_buffer = need_output_buffering ? &out_buffer[0] : NULL;
// Reset the iterator to allocate its buffers
if (NpyIter_Reset(iter, NULL) != NPY_SUCCEED) {
return -1;
}
get_temps_space(params, params.mem, BLOCK_SIZE1);
Py_BEGIN_ALLOW_THREADS;
r = vm_engine_iter_task(iter, params.memsteps, params, pc_error, errmsg);
Py_END_ALLOW_THREADS;
free_temps_space(params, params.mem);

return r;
}

/* Parallel iterator version of VM engine */
static int
vm_engine_iter_parallel(NpyIter *iter, const vm_params& params,
Expand All @@ -784,6 +810,16 @@ vm_engine_iter_parallel(NpyIter *iter, const vm_params& params,
pthread_mutex_lock(&gs.parallel_mutex);
Py_END_ALLOW_THREADS;

/* numexpr_set_nthreads() takes this same mutex, so gs.nthreads no longer
moves under our feet -- but it may have dropped to 1 between the
dispatch decision in run_interpreter() and here, and a pool of one has
no worker threads to meet at the barrier below. */
if (gs.nthreads == 1) {
ret = vm_engine_iter_serial(iter, params, need_output_buffering,
pc_error, errmsg);
goto end;
}

/* Populate parameters for worker threads */
NpyIter_GetIterIndexRange(iter, &th_params.start, &th_params.vlen);
/*
Expand Down Expand Up @@ -910,21 +946,8 @@ run_interpreter(NumExprObject *self, NpyIter *iter, NpyIter *reduce_iter,
if ((gs.nthreads == 1) || gs.force_serial) {
// Can do it as one "task"
if (reduce_iter == NULL) {
// Allocate memory for output buffering if needed
vector<char> out_buffer(need_output_buffering ?
(self->memsizes[0] * BLOCK_SIZE1) : 0);
params.out_buffer = need_output_buffering ? &out_buffer[0] : NULL;
// Reset the iterator to allocate its buffers
if(NpyIter_Reset(iter, NULL) != NPY_SUCCEED) {
return -1;
}
get_temps_space(params, params.mem, BLOCK_SIZE1);
Py_BEGIN_ALLOW_THREADS;
r = vm_engine_iter_task(iter, params.memsteps,
params, pc_error, &errmsg);
Py_END_ALLOW_THREADS;
free_temps_space(params, params.mem);
}
r = vm_engine_iter_serial(iter, params, need_output_buffering,
pc_error, &errmsg); }
else {
if (reduction_outer_loop) {
char **dataptr;
Expand Down
104 changes: 78 additions & 26 deletions numexpr/module.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -192,6 +192,36 @@ void *th_worker(void *tidptr)
/* This should never be reached, but anyway */
return(0);
}
/* Create the mutexes and the condition variable.

These outlive any particular thread pool, so this runs exactly once per
process: re-creating a mutex that some thread still holds is undefined
behaviour, and numexpr_set_nthreads() now resizes the pool while holding
gs.parallel_mutex. */
static void init_sync_primitives(void)
{
pthread_mutex_init(&gs.count_mutex, NULL);
pthread_mutex_init(&gs.parallel_mutex, NULL);

/* Barrier initialization */
pthread_mutex_init(&gs.count_threads_mutex, NULL);
pthread_cond_init(&gs.count_threads_cv, NULL);
gs.count_threads = 0; /* Reset threads counter */
gs.barrier_passed = 0;
}

#ifndef _WIN32
/* A fork() gives the child `gs` describing worker threads that did not
survive, and copies of the primitives frozen in whatever state they were
in -- possibly held by a thread that no longer exists. Start the child
from a clean slate; NumExpr_run() rebuilds the pool on the next call. */
static void reinit_after_fork(void)
{
init_sync_primitives();
gs.init_threads_done = 0;
gs.end_threads = 0;
}
#endif

/* Initialize threads */
int init_threads(void)
Expand All @@ -203,13 +233,6 @@ int init_threads(void)
return(0);
}

/* Initialize mutex and condition variable objects */
pthread_mutex_init(&gs.count_mutex, NULL);
pthread_mutex_init(&gs.parallel_mutex, NULL);

/* Barrier initialization */
pthread_mutex_init(&gs.count_threads_mutex, NULL);
pthread_cond_init(&gs.count_threads_cv, NULL);
gs.count_threads = 0; /* Reset threads counter */
gs.barrier_passed = 0;

Expand Down Expand Up @@ -263,30 +286,13 @@ int init_threads(void)
return(0);
}

/* Set the number of threads in numexpr's VM */
int numexpr_set_nthreads(int nthreads_new)
/* Rebuild the thread pool. The caller must hold gs.parallel_mutex. */
static int set_nthreads_locked(int nthreads_new)
{
int nthreads_old = gs.nthreads;
int t, rc;
void *status;

// if (nthreads_new > MAX_THREADS) {
// fprintf(stderr,
// "Error. nthreads cannot be larger than MAX_THREADS (%d)",
// MAX_THREADS);
// return -1;
// }
if (nthreads_new > global_max_threads) {
fprintf(stderr,
"Error. nthreads cannot be larger than environment variable \"NUMEXPR_MAX_THREADS\" (%ld)",
global_max_threads);
return -1;
}
else if (nthreads_new <= 0) {
fprintf(stderr, "Error. nthreads must be a positive integer");
return -1;
}

/* Only join threads if they are not initialized or if our PID is
different from that in pid var (probably means that we are a
subprocess, and thus threads are non-existent). */
Expand Down Expand Up @@ -329,6 +335,48 @@ int numexpr_set_nthreads(int nthreads_new)
return nthreads_old;
}

/* Set the number of threads in numexpr's VM */
int numexpr_set_nthreads(int nthreads_new)
{
int nthreads_old;

// if (nthreads_new > MAX_THREADS) {
// fprintf(stderr,
// "Error. nthreads cannot be larger than MAX_THREADS (%d)",
// MAX_THREADS);
// return -1;
// }
if (nthreads_new > global_max_threads) {
fprintf(stderr,
"Error. nthreads cannot be larger than environment variable \"NUMEXPR_MAX_THREADS\" (%ld)",
global_max_threads);
return -1;
}
else if (nthreads_new <= 0) {
fprintf(stderr, "Error. nthreads must be a positive integer");
return -1;
}

/* Resizing the pool has to exclude both a parallel job in flight -- whose
workers we would otherwise join out from under it, leaving it waiting
on a barrier nobody reaches -- and another resize, which would join the
same workers twice and get ESRCH. gs.parallel_mutex already serializes
parallel jobs, so reuse it here.

The GIL must be dropped while waiting for it: vm_engine_iter_parallel()
re-acquires the GIL before releasing the mutex, so holding on to it
here would deadlock the two against each other. */
Py_BEGIN_ALLOW_THREADS;
pthread_mutex_lock(&gs.parallel_mutex);
Py_END_ALLOW_THREADS;

nthreads_old = set_nthreads_locked(nthreads_new);

pthread_mutex_unlock(&gs.parallel_mutex);

return nthreads_old;
}


#ifdef USE_VML

Expand Down Expand Up @@ -470,6 +518,10 @@ PyInit_interpreter(void) {
gs.tids = (int*)calloc(sizeof(int), global_max_threads);
// TODO: for Py3, deallocate: https://docs.python.org/3/c-api/module.html#c.PyModuleDef.m_free
// For Python 2.7, people have to exit the process to reclaim the memory.
init_sync_primitives();
#ifndef _WIN32
pthread_atfork(NULL, NULL, reinit_after_fork);
#endif

if (PyType_Ready(&NumExprType) < 0)
INITERROR;
Expand Down
91 changes: 91 additions & 0 deletions numexpr/tests/test_numexpr.py
Original file line number Diff line number Diff line change
Expand Up @@ -1519,6 +1519,97 @@ def test_thread_safety_with_numexpr():

test_thread_safety_with_numexpr()

# Both scripts below are run in a subprocess: the failures they expose are a
# deadlock and an `exit(-1)` raised from C, neither of which can be caught in
# the running interpreter -- in-process they would take the whole test session
# down instead of failing a single test.

# `numexpr_set_nthreads()` tears the thread pool down and builds it back up
# again without holding any lock. Two Python threads calling it at the same
# time both pass the `gs.init_threads_done` check and both `pthread_join()`
# the same workers; the loser gets ESRCH and numexpr calls `exit(-1)`.
# Only reachable on a free-threaded build: with the GIL, `Py_set_num_threads`
# never releases it, so the calls are serialized for free.
_CONCURRENT_SET_NUM_THREADS = """
import threading
import numexpr

errors = []

def churn():
try:
for i in range(500):
numexpr.set_num_threads(1 + i % 4)
except BaseException as exc:
errors.append(exc)

threads = [threading.Thread(target=churn) for _ in range(4)]
for thread in threads:
thread.start()
for thread in threads:
thread.join()
assert not errors, errors
"""

# `vm_engine_iter_parallel()` holds `gs.parallel_mutex` and releases the GIL
# while the workers run, but `numexpr_set_nthreads()` takes no lock at all, so
# it can join those workers mid-flight. The evaluating thread is then left
# waiting on `count_threads_cv` for a pool that no longer exists.
# This one does not need free-threading: releasing the GIL is enough.
_SET_NUM_THREADS_DURING_EVALUATE = """
import threading
import numpy as np
import numexpr

numexpr.set_num_threads(4)
a = np.random.random(1_000_000)
b = np.random.random(1_000_000)
stop = threading.Event()
errors = []

def compute():
try:
while not stop.is_set():
numexpr.evaluate("sin(a) + cos(b) * a / (b + 1.0)")
except BaseException as exc:
errors.append(exc)

workers = [threading.Thread(target=compute) for _ in range(3)]
for worker in workers:
worker.start()
try:
for i in range(200):
numexpr.set_num_threads(1 + i % 4)
finally:
stop.set()
for worker in workers:
worker.join()
assert not errors, errors
"""


@pytest.mark.thread_unsafe
class test_set_num_threads_concurrency(TestCase):
"""set_num_threads() resizes a process-global thread pool; doing so
concurrently with itself or with a running evaluation must not corrupt it.
"""

def _run_isolated(self, script, timeout=120):
try:
proc = subprocess.run([sys.executable, '-c', script],
timeout=timeout, capture_output=True,
universal_newlines=True)
except subprocess.TimeoutExpired:
self.fail(f'deadlock: no exit after {timeout}s')
if proc.returncode != 0:
self.fail(f'exited with {proc.returncode}\n{proc.stderr}')

def test_concurrent_set_num_threads(self):
self._run_isolated(_CONCURRENT_SET_NUM_THREADS)

def test_set_num_threads_during_evaluate(self):
self._run_isolated(_SET_NUM_THREADS_DURING_EVALUATE)


# The worker function for the subprocess (needs to be here because Windows
# has problems pickling nested functions with the multiprocess module :-/)
Expand Down