File size: 3,622 Bytes
a41b023
66cf40d
a41b023
 
 
 
66cf40d
a41b023
 
 
 
 
 
 
 
 
66cf40d
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
a41b023
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
66cf40d
 
 
 
 
a41b023
 
 
 
66cf40d
 
 
 
a41b023
66cf40d
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
"""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