Skip to content

Commit 79c9db9

Browse files
committed
Fixed parallelisation of automated ICA.
1 parent d0eb22f commit 79c9db9

1 file changed

Lines changed: 23 additions & 0 deletions

File tree

‎osl_dynamics/meeg/parallel.py‎

Lines changed: 23 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -14,13 +14,36 @@
1414
"MKL_NUM_THREADS",
1515
"OPENBLAS_NUM_THREADS",
1616
"VECLIB_MAXIMUM_THREADS",
17+
"NUMEXPR_MAX_THREADS",
1718
]
1819

1920

21+
def _limit_onnx_threads():
22+
"""Patch ONNX Runtime to use 1 thread per session."""
23+
try:
24+
import onnxruntime as ort
25+
26+
_OriginalSession = ort.InferenceSession
27+
28+
class _SingleThreadSession(_OriginalSession):
29+
def __init__(self, *args, **kwargs):
30+
if "sess_options" not in kwargs or kwargs["sess_options"] is None:
31+
opts = ort.SessionOptions()
32+
opts.intra_op_num_threads = 1
33+
opts.inter_op_num_threads = 1
34+
kwargs["sess_options"] = opts
35+
super().__init__(*args, **kwargs)
36+
37+
ort.InferenceSession = _SingleThreadSession
38+
except ImportError:
39+
pass
40+
41+
2042
def _worker(
2143
args: Tuple[Callable, str, Any, Path, dict],
2244
) -> Tuple[str, bool]:
2345
"""Wrapper that handles logging and error catching for a single item."""
46+
_limit_onnx_threads()
2447
func, id, item, log_dir, kwargs = args
2548
with MEEGSessionLogger(id, log_dir) as logger:
2649
try:

0 commit comments

Comments
 (0)