diff --git a/numexpr/interpreter.cpp b/numexpr/interpreter.cpp index f377721..9359e85 100644 --- a/numexpr/interpreter.cpp +++ b/numexpr/interpreter.cpp @@ -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 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, @@ -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); /* @@ -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 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; diff --git a/numexpr/module.cpp b/numexpr/module.cpp index 67629bd..5858f5e 100644 --- a/numexpr/module.cpp +++ b/numexpr/module.cpp @@ -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) @@ -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; @@ -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). */ @@ -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 @@ -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; diff --git a/numexpr/tests/test_numexpr.py b/numexpr/tests/test_numexpr.py index 5fdeebb..9b2b754 100644 --- a/numexpr/tests/test_numexpr.py +++ b/numexpr/tests/test_numexpr.py @@ -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 :-/)