65 lines
2.3 KiB
Python
65 lines
2.3 KiB
Python
# -*- coding: utf-8 -*-
|
||||
|
|
"""推理后端工厂:按配置选择后端实现(可插拔适配层)。
|
|||
|
|
|
|||
|
|
PRD 5.6「配置点:资源配额、推理后端选择、灰度发布策略」——
|
|||
|
|
`backend: gpu|npu` 即推理后端选择;业务代码只调用 `build_backend()`
|
|||
|
|
返回的接口对象,不感知具体硬件。
|
|||
|
|
"""
|
|||
|
|
import os
|
|||
|
|
|
|||
|
|
from inference_backend.base import InferenceBackend
|
|||
|
|
from inference_backend.gpu_backend import NvidiaGpuBackend
|
|||
|
|
from inference_backend.npu_backend import AscendNpuBackend
|
|||
|
|
|
|||
|
|
#: 可插拔后端注册表:配置名 -> 实现类。新增硬件只需注册新类。
|
|||
|
|
BACKEND_REGISTRY = {
|
|||
|
|
"gpu": NvidiaGpuBackend, # NVIDIA 5090(Triton/ONNX)
|
|||
|
|
"npu": AscendNpuBackend, # 华为昇腾(ACL/CANN)
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
|
|||
|
|
def build_backend(config: dict) -> InferenceBackend:
|
|||
|
|
"""依据配置构建推理后端实例(切换后端仅改配置,业务代码零改动)。
|
|||
|
|
|
|||
|
|
Args:
|
|||
|
|
config: 后端配置字典(见 config/backends.template.yaml),
|
|||
|
|
至少包含 ``backend`` 键(gpu | npu)。
|
|||
|
|
|
|||
|
|
Returns:
|
|||
|
|
实现了 :class:`InferenceBackend` 接口的后端实例。
|
|||
|
|
|
|||
|
|
Raises:
|
|||
|
|
ValueError: 配置缺失或指定了未注册的后端。
|
|||
|
|
"""
|
|||
|
|
if not isinstance(config, dict) or not config.get("backend"):
|
|||
|
|
raise ValueError("推理后端配置缺失:需要 backend: gpu|npu")
|
|||
|
|
name = str(config["backend"]).lower()
|
|||
|
|
if name not in BACKEND_REGISTRY:
|
|||
|
|
raise ValueError(
|
|||
|
|
f"未注册的推理后端: {name!r},可用: {sorted(BACKEND_REGISTRY)}"
|
|||
|
|
)
|
|||
|
|
cls = BACKEND_REGISTRY[name]
|
|||
|
|
inf = config.get("inference", {}) or {}
|
|||
|
|
return cls(
|
|||
|
|
endpoint=inf.get("endpoint", ""),
|
|||
|
|
model=inf.get("model", "iaop-default"),
|
|||
|
|
timeout_seconds=float(inf.get("timeout_seconds", 10)),
|
|||
|
|
runtime=inf.get("runtime", ""),
|
|||
|
|
device=inf.get("device", ""),
|
|||
|
|
cann_version=inf.get("cann_version", ""),
|
|||
|
|
)
|
|||
|
|
|
|||
|
|
|
|||
|
|
def load_backend_config(path: str) -> dict:
|
|||
|
|
"""从 YAML 配置资产加载后端配置(模板可覆盖资产)。"""
|
|||
|
|
import yaml
|
|||
|
|
|
|||
|
|
with open(path, "r", encoding="utf-8") as fh:
|
|||
|
|
return yaml.safe_load(fh) or {}
|
|||
|
|
|
|||
|
|
|
|||
|
|
def default_config_path() -> str:
|
|||
|
|
"""返回本模块模板配置资产的默认路径。"""
|
|||
|
|
return os.path.join(os.path.dirname(os.path.abspath(__file__)),
|
|||
|
|
"config", "backends.template.yaml")
|