【Bug已解决】[Feature Request] Respect OMP_NUM_THREADS (and new ORT_*_NUM_THREADS env vars) when sizing default thread pools 解决方案
一、现象长什么样
在用 ONNX Runtime 做推理时,用户希望通过环境变量控制默认线程池大小(比如容器里限制 CPU 核数、或和 OpenMP 对齐),但发现 ORT完全忽略OMP_NUM_THREADS,也不认自己新增的ORT_INTRA_OP_NUM_THREADS/ORT_INTER_OP_NUM_THREADS之类的环境变量,硬是用“逻辑 CPU 数”把线程池开满,导致:容器里 CPU 配额为 4 核却被开成 64 线程、和同进程的 OpenMP 线程数不一致引发 oversubscription、或在受限环境里因线程过多被限流。现象:
# 现象 A:设了 OMP_NUM_THREADS=4,ORT 仍开满核 # ORT 用 std::thread::hardware_concurrency() 直接开 64 线程, # 无视 OMP_NUM_THREADS=4 # 现象 B:新加的 ORT_*_NUM_THREADS 不生效 # 用户设 ORT_INTRA_OP_NUM_THREADS=2,ORT 没读这个变量,线程数不变 # 现象 C:和 OpenMP 线程 oversubscription # ORT 开 64 + OpenMP 开 64 = 128 线程抢 64 核,性能反而下降、延迟抖动最坑的是现象 A:用户以为设了OMP_NUM_THREADS就控制了所有计算库,结果 ORT 自作主张开满,容器配额形同虚设,排查半天才发现是 ORT 没读这个变量。
二、背景
多线程推理库通常会读环境变量来决定默认线程池大小:OMP_NUM_THREADS是 OpenMP 的事实标准,MKL_NUM_THREADS、OMP_NUM_THREADS等也是常见约定。ONNX Runtime 自己有 intra-op(算子内并行)和 inter-op(算子间并行)两套线程池,理应也尊重这些环境变量,尤其是OMP_NUM_THREADS作为“默认并行度”的通用信号。
问题在于 ORT 在“未显式指定线程数”时,直接调用hardware_concurrency()取全核数,跳过了对环境变量的读取。新增的ORT_*_NUM_THREADS变量要么没在“默认大小”分支里被读,要么读取优先级低于“全核数”。于是用户的所有环境变量控制都失效。
这是线程池/资源配置审查里典型的坑:默认线程数计算直接取硬件核数,忽略了既有的环境变量约定(OMP_NUM_THREADS 等)和新变量。
三、根因
默认大小分支不读环境变量:
SessionOptions在未设intra_op_num_threads时,直接用hardware_concurrency(),没先看OMP_NUM_THREADS/ORT_*_NUM_THREADS。新增变量未被读取:
ORT_INTRA_OP_NUM_THREADS等只在“显式 API 设置”路径生效,没在“默认推导”路径读取环境变量。优先级混乱:即使读了,多个变量(OMP / ORT_*/MKL)之间的优先级没定义,导致行为不确定。
本质:是默认线程池大小的推导跳过环境变量读取,且新增 ORT 变量未被纳入默认推导、优先级未定义。
四、最小可运行复现
下面用 Python 模拟“默认线程数直接取核数,忽略 OMP_NUM_THREADS”:
import os def size_default_pool_buggy(): """buggy: 直接取硬件核数,不读环境变量。""" return os.cpu_count() or 1 # 忽略 OMP_NUM_THREADS def size_default_pool_fixed(): """fixed: 按优先级读 OMP_NUM_THREADS / ORT_INTRA_OP_NUM_THREADS, 都没有才退回硬件核数。""" for var in ("ORT_INTRA_OP_NUM_THREADS", "OMP_NUM_THREADS"): val = os.environ.get(var) if val and val.isdigit(): return int(val) return os.cpu_count() or 1 os.environ["OMP_NUM_THREADS"] = "4" print("buggy:", size_default_pool_buggy()) # 64(忽略 4) print("fixed:", size_default_pool_fixed()) # 4(尊重环境变量)buggy返回 64(核数),fixed返回 4(读到了OMP_NUM_THREADS)。
五、解决方案(第一层:最小直接修复)
最小修复:默认线程池大小推导时,按优先级读取环境变量,都没设才退回硬件核数:
// 修正:默认线程数推导尊重环境变量 int GetDefaultIntraOpThreadCount() { // 优先级:ORT_INTRA_OP_NUM_THREADS > OMP_NUM_THREADS > 硬件核数 if (const char* ort = std::getenv("ORT_INTRA_OP_NUM_THREADS")) { if (int n = ParsePositiveInt(ort)) return n; } if (const char* omp = std::getenv("OMP_NUM_THREADS")) { if (int n = ParsePositiveInt(omp)) return n; } return std::thread::hardware_concurrency(); }这一层改动最小:默认分支先读变量再退回核数,环境变量控制恢复。但依赖“每处默认推导都加这套读取”,下看第二层。
六、解决方案(第二层:结构性改进)
把“默认线程池大小的推导规则(环境变量优先级 + 退回核数)”固化成单一事实来源。下面这个 dataclass 集中管理,C++ 侧和 Python 校验侧共享同一规则:
from dataclasses import dataclass, field from typing import Dict, List, Optional @dataclass class OrtThreadpoolEnvPolicy: """单一事实来源:默认线程池大小的推导契约。""" # 优先级从高到低 env_priority: List[str] = field(default_factory=lambda: [ "ORT_INTRA_OP_NUM_THREADS", "ORT_INTER_OP_NUM_THREADS", "OMP_NUM_THREADS", "MKL_NUM_THREADS"]) def resolve(self, env: Dict[str, str], hardware_concurrency: int) -> int: for var in self.env_priority: val = env.get(var) if val and val.strip().isdigit(): n = int(val) if n > 0: return n return max(1, hardware_concurrency) # 都没设则退回核数 def assert_respects_env(self, env: Dict[str, str], hw: int, expected: int) -> None: got = self.resolve(env, hw) if got != expected: raise AssertionError( f"thread pool size {got} ignores env (expected {expected})")这一层的关键收益:
- 统一优先级:
env_priority定义清晰的变量优先级,行为确定; - 退回核数兜底:都没设才用硬件核数,且
max(1, ...)防 0; - 可校验:
assert_respects_env验证环境变量确实被尊重; - 单一事实来源:所有线程池大小推导收口在
OrtThreadpoolEnvPolicy。
七、解决方案(第三层:断言 / CI 守护)
把第二层钉成 pytest,挂进 CI,确保环境变量被尊重:
import pytest from your_package.ort_threadpool_env import OrtThreadpoolEnvPolicy def test_omp_respected(): # 断言 1:OMP_NUM_THREADS 被尊重 p = OrtThreadpoolEnvPolicy() assert p.resolve({"OMP_NUM_THREADS": "4"}, 64) == 4 def test_ort_var_takes_precedence(): # 断言 2:ORT_* 变量优先级高于 OMP p = OrtThreadpoolEnvPolicy() env = {"ORT_INTRA_OP_NUM_THREADS": "2", "OMP_NUM_THREADS": "8"} assert p.resolve(env, 64) == 2 def test_fallback_to_hw(): # 断言 3:都不设时退回硬件核数 p = OrtThreadpoolEnvPolicy() assert p.resolve({}, 64) == 64 def test_zero_core_safe(): # 断言 4:硬件核数为 0 时至少返回 1 p = OrtThreadpoolEnvPolicy() assert p.resolve({}, 0) == 1四条断言从“OMP 被尊重”“ORT 优先”“退回核数”“零核安全”四面把环境变量回归钉死在 CI。
八、排查清单
ORT 不尊重线程数环境变量时:
- 设了
OMP_NUM_THREADS线程池仍开满?查默认大小推导是否直接取hardware_concurrency()而没读变量(现象 A)。 - 新增的
ORT_INTRA_OP_NUM_THREADS不生效?查它是否在“默认推导”路径被读,而非只在显式 API 路径(现象 B)。 - 多个变量同设谁优先?定义清晰优先级(ORT_* > OMP > MKL > 核数),避免不确定。
- 用第二层
OrtThreadpoolEnvPolicy:优先级统一 + 退回核数 + 可校验。 - 加第三层 pytest,断言“OMP 被尊重、ORT 优先、退回核数、零核安全”。
- 容器/受限环境必须能靠环境变量限制线程数,否则 oversubscription 拖垮性能。
九、小结
ORT 默认线程池不尊重环境变量的 bug 本质是未显式指定线程数时,默认推导直接取硬件核数,跳过了OMP_NUM_THREADS等约定变量,且新增的ORT_*_NUM_THREADS也没纳入默认推导、优先级未定义,导致容器配额形同虚设、与 OpenMP oversubscription。修复分三层——第一层默认推导按优先级读ORT_*/OMP_NUM_THREADS,都没设才退回核数;第二层用OrtThreadpoolEnvPolicy这个 dataclass 把推导规则(优先级 + 兜底 + 校验)收口成单一事实来源;第三层用四条 pytest 把“OMP 被尊重、ORT 优先、退回核数、零核安全”钉死在 CI。核心心法:默认线程池大小必须尊重既有环境变量约定(OMP_NUM_THREADS 等)并定义清晰优先级,仅在全未设置时才退回硬件核数。