"""Compute configuration — bound every library to the available/allowed CPU cores (CPU-only, no GPU). ``configure_compute(cpu_threads)`` points the thread-pool env vars (BLAS / OpenMP / polars / tokenizers) at the resolved core budget. It MUST run before numpy / polars / onnxruntime are imported (those read the env at import), so the entry points (rank.py, cli, streamlit) call it as their very first action — but reading ``Settings.compute.cpu_threads`` first is safe, since loading config never imports those libraries. ``cpu_threads=0`` (the default, "auto") resolves to every core this PROCESS can actually use — via ``os.sched_getaffinity`` where the platform supports it (Linux; respects a container's cpuset/cgroup limit, not just the host's total core count), falling back to ``os.cpu_count()`` elsewhere (Windows/macOS). Passing a positive ``cpu_threads`` caps the budget below that (e.g. to leave headroom for other processes on a shared host); a value above the detected core count is clamped down, never up. ONNX sessions additionally get explicit ``SessionOptions`` (``onnx_session_options``): graph opt = ALL, and intra-op threads either follow the same resolved cap or (when uncapped) defer to onnxruntime's own auto-tuning, which is usually smarter than "all logical cores" (it tends to land on physical cores, avoiding HT contention). """ from __future__ import annotations import os from typing import Any _THREAD_ENV_VARS = ( "OMP_NUM_THREADS", # OpenMP (onnxruntime, many BLAS backends) "OPENBLAS_NUM_THREADS", # numpy OpenBLAS "MKL_NUM_THREADS", # numpy MKL "NUMEXPR_NUM_THREADS", "POLARS_MAX_THREADS", # polars worker pool ) def _available_cores() -> int: """Cores this process can actually use: cgroup/affinity-aware where supported (Linux), else the raw count.""" getaffinity = getattr(os, "sched_getaffinity", None) # absent on Windows/macOS if getaffinity is not None: try: return len(getaffinity(0)) or 1 except OSError: pass return os.cpu_count() or 1 def resolve_cpu_threads(cpu_threads: int = 0) -> int: """0 (or negative) → every available core (cgroup-aware); N>0 → min(N, available), so a stale/oversized config value can never ask a host for more threads than it actually has.""" available = _available_cores() return available if cpu_threads <= 0 else min(cpu_threads, available) def configure_compute(cpu_threads: int = 0) -> None: """Set every thread-pool env var to the resolved core budget so numpy/polars/onnx/tokenizers agree on it. Idempotent (``setdefault``). ``cpu_threads=0`` (default) auto-detects; pass a positive value to cap it.""" cores = str(resolve_cpu_threads(cpu_threads)) for env_var in _THREAD_ENV_VARS: os.environ.setdefault(env_var, cores) os.environ.setdefault("TOKENIZERS_PARALLELISM", "true") # parallel batch tokenization def onnx_session_options(cpu_threads: int = 0) -> Any: """A SessionOptions tuned for CPU: all graph optimizations. ``cpu_threads=0`` (default) leaves onnxruntime's own thread-count heuristic in charge (typically the physical-core auto-detect, no HT contention); a positive value caps intra-op threads to that count instead.""" import onnxruntime as ort options = ort.SessionOptions() options.graph_optimization_level = ort.GraphOptimizationLevel.ORT_ENABLE_ALL options.intra_op_num_threads = max(0, cpu_threads) # 0 → onnxruntime auto; N>0 → explicit cap options.execution_mode = ort.ExecutionMode.ORT_SEQUENTIAL return options