From cb860d6d8b9223a763572054a3cf3752379f2766 Mon Sep 17 00:00:00 2001 From: Vinicius Faria Date: Sun, 26 Jul 2026 23:21:50 -0300 Subject: [PATCH] fix: move to thread paralellism to avoid broken jvm fork state --- pyterrier/_evaluation/_grid.py | 31 +++++++++++++++++++++++-------- 1 file changed, 23 insertions(+), 8 deletions(-) diff --git a/pyterrier/_evaluation/_grid.py b/pyterrier/_evaluation/_grid.py index 9d4b3c3b..e8fb17fa 100644 --- a/pyterrier/_evaluation/_grid.py +++ b/pyterrier/_evaluation/_grid.py @@ -248,11 +248,9 @@ def _evaluate_several_settings(inputs : List[Tuple]): else: import itertools import more_itertools - try: - from pyterrier_alpha.parallel import parallel_lambda # type: ignore - except ImportError as ie: - raise ImportError("pyterrier-alpha[parallel] must be installed for jobs>1") from ie - + from concurrent.futures import ThreadPoolExecutor + from jnius import detach + all_inputs = [(keys, values) for values in combinations] # how many jobs to distribute this to @@ -261,9 +259,26 @@ def _evaluate_several_settings(inputs : List[Tuple]): # built the batches to distribute batched_inputs = list(more_itertools.chunked(all_inputs, num_batches)) assert len(batched_inputs) > 0, "No inputs identified for parallel_lambda" - eval_list = parallel_lambda(_evaluate_several_settings, batched_inputs, jobs, backend=backend) - eval_list = list(itertools.chain(*eval_list)) - assert len(eval_list) > 0, "parallel_lambda returned 0 rows" + + if backend == 'ray': + # preserve alpha behavior since ray has its own JVM lifecycle + try: + from pyterrier_alpha.parallel import parallel_lambda # type: ignore + except ImportError as ie: + raise ImportError("pyterrier-alpha[parallel] must be installed for backend='ray'") from ie + eval_list = parallel_lambda(_evaluate_several_settings, batched_inputs, jobs, backend=backend) + eval_list = list(itertools.chain(*eval_list)) + else: + # avoid process parallelism by using threads, which share the parent's JVM + # instead of calling fork() and leaving leave the JVM in a weird state + def _eval_chunk(chunk): + out = [_evaluate_one_setting(k, v) for k, v in chunk] + detach() # release JNI refs + return out + with ThreadPoolExecutor(max_workers=jobs) as ex: + per_chunk = list(ex.map(_eval_chunk, batched_inputs)) + eval_list = list(itertools.chain(*per_chunk)) + assert len(eval_list) > 0, "GridScan produced 0 rows" # resulting eval_list has the form [ # ( [(BR, 'wmodel', 'BM25'), (BR, 'c', 0.2)] , {"map" : 0.2654} )