|
| 1 | +import functools |
| 2 | +import logging |
| 3 | +import os |
| 4 | + |
| 5 | +logger = logging.getLogger("app.services.onnx_providers") |
| 6 | + |
| 7 | +# 推理后端优先级:CUDA (gpu extra) > OpenVINO (openvino extra) > CPU (cpu extra)。 |
| 8 | +# 通过 ort.get_available_providers() 动态探测,避免硬编码导致未安装的 EP 报警告。 |
| 9 | +_PRIORITY = ["CUDAExecutionProvider", "OpenVINOExecutionProvider", "CPUExecutionProvider"] |
| 10 | + |
| 11 | + |
| 12 | +def _openvino_device() -> str: |
| 13 | + """选择 OpenVINO EP 的 device_type:NPU 优先,否则 CPU。 |
| 14 | +
|
| 15 | + OPENVINO_DEVICE 环境变量可强制覆盖(设为 NPU / GPU / CPU);留空时按 |
| 16 | + openvino.Core().available_devices() 探测——有 NPU 走 NPU,否则回退 CPU, |
| 17 | + 避免在无 NPU 的机器上指定 NPU 导致会话创建失败。 |
| 18 | + """ |
| 19 | + override = os.getenv("OPENVINO_DEVICE", "").strip() |
| 20 | + if override: |
| 21 | + return override |
| 22 | + try: |
| 23 | + from openvino import Core |
| 24 | + devices = set(Core().available_devices) |
| 25 | + if "NPU" in devices: |
| 26 | + return "NPU" |
| 27 | + except Exception as e: |
| 28 | + logger.warning(f"Failed to probe OpenVINO devices, fallback to CPU: {e}") |
| 29 | + return "CPU" |
| 30 | + |
| 31 | + |
| 32 | +@functools.lru_cache(maxsize=1) |
| 33 | +def get_onnx_providers(): |
| 34 | + """返回 (providers, provider_options),按可用性筛选并保持优先级。 |
| 35 | +
|
| 36 | + 结果在进程生命周期内缓存(可用 EP 不会变),避免每次加载模型都重复探测与打日志。 |
| 37 | +
|
| 38 | + provider_options 与 providers 一一对齐: |
| 39 | + - CUDAExecutionProvider -> {"device_id": 0} |
| 40 | + - OpenVINOExecutionProvider -> {"device_type": "NPU" | "CPU"}(NPU 优先) |
| 41 | + - CPUExecutionProvider -> {} |
| 42 | + """ |
| 43 | + try: |
| 44 | + import onnxruntime as ort |
| 45 | + available = set(ort.get_available_providers()) |
| 46 | + except Exception as e: |
| 47 | + logger.warning(f"Failed to probe onnxruntime providers, fallback to CPU: {e}") |
| 48 | + return ["CPUExecutionProvider"], [{}] |
| 49 | + |
| 50 | + providers = [p for p in _PRIORITY if p in available] |
| 51 | + if not providers: |
| 52 | + providers = ["CPUExecutionProvider"] |
| 53 | + |
| 54 | + options = [] |
| 55 | + for p in providers: |
| 56 | + if p == "CUDAExecutionProvider": |
| 57 | + options.append({"device_id": 0}) |
| 58 | + elif p == "OpenVINOExecutionProvider": |
| 59 | + options.append({"device_type": _openvino_device()}) |
| 60 | + else: |
| 61 | + options.append({}) |
| 62 | + |
| 63 | + logger.info(f"ONNX Runtime providers selected: {providers} with options {options}") |
| 64 | + return providers, options |
0 commit comments