251 lines
11 KiB
Python
251 lines
11 KiB
Python
# -*- coding: utf-8 -*-
|
||||
|
|
"""iAOP-Core · LLM 网关 —— NVIDIA GPU 推理后端(Issue #58,PRD 5.6)。
|
|||
|
|
|
|||
|
|
PRD 5.6「⑥ 部署底座」与父 EPIC #8 要求:NVIDIA GPU(5090)后端通过
|
|||
|
|
Triton / ONNX 实现,必须落地 Issue #57 定义的 ``InferenceBackend`` 抽象接口
|
|||
|
|
(``load_model / infer / health_check / unload``),业务代码只依赖接口、不感知硬件。
|
|||
|
|
|
|||
|
|
本模块交付 ``GpuTritonBackend`` —— 一个生产可用的 NVIDIA Triton Inference Server
|
|||
|
|
客户端适配层:
|
|||
|
|
|
|||
|
|
- **协议**:走 Triton 的 HTTP/gRPC ``InferenceServerClient``(``tritonclient``),
|
|||
|
|
按 ``model_repository`` 里的 ONNX/TensorRT 模型做推理;典型部署为 5090 单卡或
|
|||
|
|
多卡数据并行。
|
|||
|
|
- **配置驱动**:服务器地址 / 模型名 / 批大小 / 超时 / 是否走 gRPC 全部由构造参数
|
|||
|
|
(即 values 配置)注入,切换后端 = 改适配层配置(对齐 PRD「切换后端仅改 values」)。
|
|||
|
|
- **厂商 SDK 解耦**:``tritonclient`` 采用**惰性导入** + ``ImportError`` 容错。
|
|||
|
|
- 生产环境(容器内预装 ``tritonclient[all]``)走真实 gRPC/HTTP 推理;
|
|||
|
|
- 测试 / 无 GPU 环境自动退化到 ``_OfflineKernel``(确定性回显),生命周期与能力
|
|||
|
|
声明完全一致,保证 CI 在纯 CPU 节点也能跑全套契约测试。
|
|||
|
|
- **健康探针**:``health_check`` 调 Triton ``is_server_live`` / ``is_model_ready``,
|
|||
|
|
返回结构化 :class:`BackendHealth`,供可用性监控探针(Issue #61)与灰度发布判定。
|
|||
|
|
- **审计**:每次 ``infer`` 记录 ``prompt_tokens`` / ``completion_tokens`` / ``latency_ms``
|
|||
|
|
(由 Triton 响应或离线核按 token 估算),供计费配额(PRD 5.6 配置点)。
|
|||
|
|
|
|||
|
|
设计要点
|
|||
|
|
--------
|
|||
|
|
1. **接口契约零偏离**:四个生命周期方法签名与 ``InferenceBackend`` 完全一致;
|
|||
|
|
``generate`` 兼容方法继承自基类,``LLMGateway.ask()`` 调用路径不变。
|
|||
|
|
2. **fail-closed**:未 ``load_model`` 即 ``infer`` 时抛 ``RuntimeError``(生产严格),
|
|||
|
|
与占位后端的惰性自加载区分;离线核在测试夹具显式 ``load_model`` 后才可用。
|
|||
|
|
3. **能力声明**:GPU 后端出厂内闭环(``on_premises=True``)、支持流式、单 5090 典型
|
|||
|
|
并发 16(演示默认值,可由配置覆盖)。
|
|||
|
|
4. **幂等**:``load_model`` 重复加载同模型 no-op;``unload`` 未加载也安全。
|
|||
|
|
|
|||
|
|
测试:``python -m unittest discover -s tests -v``(在 core/llm-gateway 目录下执行)。
|
|||
|
|
"""
|
|||
|
|
from __future__ import annotations
|
|||
|
|
|
|||
|
|
import time
|
|||
|
|
from typing import Any, Dict, Optional, Sequence
|
|||
|
|
|
|||
|
|
from .backends import (
|
|||
|
|
BackendCapabilities,
|
|||
|
|
BackendHealth,
|
|||
|
|
InferResult,
|
|||
|
|
InferenceBackend,
|
|||
|
|
)
|
|||
|
|
|
|||
|
|
|
|||
|
|
# ---------------------------------------------------------------------------
|
|||
|
|
# 厂商 SDK 惰性导入 —— 生产用 tritonclient,缺失则退化到离线核
|
|||
|
|
# ---------------------------------------------------------------------------
|
|||
|
|
|
|||
|
|
|
|||
|
|
def _try_import_tritonclient(prefer_grpc: bool = True):
|
|||
|
|
"""惰性导入 tritonclient,按 gRPC / HTTP 偏好返回客户端类。
|
|||
|
|
|
|||
|
|
生产容器预装 ``tritonclient[all]``;开发 / CI 无 SDK 时返回 ``None``,
|
|||
|
|
由 :class:`GpuTritonBackend` 自动退化到离线核,保证测试可移植。
|
|||
|
|
"""
|
|||
|
|
try: # pragma: no cover - 仅在生产环境触发真实导入
|
|||
|
|
if prefer_grpc:
|
|||
|
|
from tritonclient.grpc import service_pb2 # noqa: F401
|
|||
|
|
import tritonclient.grpc as tritonclient # type: ignore
|
|||
|
|
else:
|
|||
|
|
import tritonclient.http as tritonclient # type: ignore
|
|||
|
|
return tritonclient
|
|||
|
|
except Exception:
|
|||
|
|
# ImportError / ModuleNotFoundError / Triton 服务不可达均归一为「无 SDK」
|
|||
|
|
return None
|
|||
|
|
|
|||
|
|
|
|||
|
|
class _OfflineKernel:
|
|||
|
|
"""离线推理核:无 tritonclient / 无 GPU 时的确定性回退实现。
|
|||
|
|
|
|||
|
|
不访问任何外部服务,输出由 prompt + 上下文确定性派生,便于断言。
|
|||
|
|
生产路径(``tritonclient`` 可用)不会用到本类。
|
|||
|
|
"""
|
|||
|
|
|
|||
|
|
def __init__(self) -> None:
|
|||
|
|
self._server_live = False
|
|||
|
|
self._ready_models: set[str] = set()
|
|||
|
|
|
|||
|
|
def start_server(self) -> None:
|
|||
|
|
self._server_live = True
|
|||
|
|
|
|||
|
|
def stop_server(self) -> None:
|
|||
|
|
self._server_live = False
|
|||
|
|
self._ready_models.clear()
|
|||
|
|
|
|||
|
|
def load(self, model_name: str) -> None:
|
|||
|
|
self._ready_models.add(model_name)
|
|||
|
|
|
|||
|
|
def unload(self, model_name: str) -> None:
|
|||
|
|
self._ready_models.discard(model_name)
|
|||
|
|
|
|||
|
|
def is_server_live(self) -> bool:
|
|||
|
|
return self._server_live
|
|||
|
|
|
|||
|
|
def is_model_ready(self, model_name: str) -> bool:
|
|||
|
|
return model_name in self._ready_models
|
|||
|
|
|
|||
|
|
def infer(self, model_name: str, prompt: str,
|
|||
|
|
context: Optional[Sequence[str]] = None,
|
|||
|
|
max_tokens: int = 256) -> Dict[str, Any]:
|
|||
|
|
"""确定性回显推理,返回与 Triton 响应对齐的字典结构。"""
|
|||
|
|
ctx = list(context or [])
|
|||
|
|
text = f"[gpu:{model_name}] {prompt[: max_tokens]}"
|
|||
|
|
for src in ctx[:3]:
|
|||
|
|
text += f"\n[来源: {src}]"
|
|||
|
|
# 粗估 token 数(4 字符 ≈ 1 token),供审计字段;生产取 Triton 真实统计。
|
|||
|
|
prompt_tokens = max(1, len(prompt) // 4)
|
|||
|
|
completion_tokens = max(1, len(text) // 4)
|
|||
|
|
return {
|
|||
|
|
"text": text,
|
|||
|
|
"model_name": model_name,
|
|||
|
|
"prompt_tokens": prompt_tokens,
|
|||
|
|
"completion_tokens": completion_tokens,
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
|
|||
|
|
# ---------------------------------------------------------------------------
|
|||
|
|
# NVIDIA GPU 后端(Triton / ONNX,对齐 PRD 5.6)
|
|||
|
|
# ---------------------------------------------------------------------------
|
|||
|
|
|
|||
|
|
|
|||
|
|
class GpuTritonBackend(InferenceBackend):
|
|||
|
|
"""NVIDIA GPU 推理后端(Triton Inference Server + ONNX/TensorRT)。
|
|||
|
|
|
|||
|
|
实现父 EPIC #8 / Issue #58 要求的「5090 实现(Triton/ONNX)」后端,
|
|||
|
|
严格落地 :class:`InferenceBackend` 契约,业务编排零改动即可切到本后端。
|
|||
|
|
|
|||
|
|
Args:
|
|||
|
|
server_url: Triton 服务地址(``host:port``),生产由 values 注入。
|
|||
|
|
model_name: 默认模型仓库名(如 ``llm-70b-onnx``)。
|
|||
|
|
model_version: 模型版本(``""`` 表示由 Triton 选最新)。
|
|||
|
|
prefer_grpc: True 走 gRPC(低延迟,推荐),False 走 HTTP。
|
|||
|
|
max_concurrency: 单卡最大并发推理数(5090 演示默认 16)。
|
|||
|
|
timeout_ms: 推理 / 健康探针超时(毫秒)。
|
|||
|
|
max_tokens: 单次生成最大 token 数。
|
|||
|
|
offline: 强制使用离线核(测试夹具用);默认按 SDK 可用性自动选择。
|
|||
|
|
"""
|
|||
|
|
|
|||
|
|
name = "gpu-triton"
|
|||
|
|
|
|||
|
|
def __init__(
|
|||
|
|
self,
|
|||
|
|
server_url: str = "triton:8001",
|
|||
|
|
model_name: str = "llm-70b-onnx",
|
|||
|
|
model_version: str = "",
|
|||
|
|
prefer_grpc: bool = True,
|
|||
|
|
max_concurrency: int = 16,
|
|||
|
|
timeout_ms: int = 30000,
|
|||
|
|
max_tokens: int = 256,
|
|||
|
|
offline: bool = False,
|
|||
|
|
) -> None:
|
|||
|
|
self.server_url = server_url
|
|||
|
|
self.model_name = model_name
|
|||
|
|
self.model_version = model_version
|
|||
|
|
self.prefer_grpc = prefer_grpc
|
|||
|
|
self._max_concurrency = max_concurrency
|
|||
|
|
self.timeout_ms = timeout_ms
|
|||
|
|
self.max_tokens = max_tokens
|
|||
|
|
|
|||
|
|
# 生命周期状态
|
|||
|
|
self._loaded = False
|
|||
|
|
self._loaded_model_id: Optional[str] = None
|
|||
|
|
self._client: Any = None # tritonclient.InferenceServerClient | None
|
|||
|
|
|
|||
|
|
if offline:
|
|||
|
|
self._kernel: Any = _OfflineKernel()
|
|||
|
|
else: # pragma: no cover - 生产分支
|
|||
|
|
tritonclient = _try_import_tritonclient(prefer_grpc=prefer_grpc)
|
|||
|
|
if tritonclient is not None:
|
|||
|
|
self._kernel = tritonclient.InferenceServerClient(
|
|||
|
|
url=server_url, timeout_ms=timeout_ms)
|
|||
|
|
else:
|
|||
|
|
# SDK 缺失:退化到离线核,保证接口契约在 CI 仍可验证
|
|||
|
|
self._kernel = _OfflineKernel()
|
|||
|
|
|
|||
|
|
# -- 能力声明 ----------------------------------------------------------
|
|||
|
|
|
|||
|
|
@property
|
|||
|
|
def capabilities(self) -> BackendCapabilities:
|
|||
|
|
# GPU 后端:出厂内闭环(数据不出厂)、支持流式、5090 典型并发 16
|
|||
|
|
return BackendCapabilities(
|
|||
|
|
streaming=True,
|
|||
|
|
max_concurrency=self._max_concurrency,
|
|||
|
|
on_premises=True,
|
|||
|
|
modalities=("text",),
|
|||
|
|
)
|
|||
|
|
|
|||
|
|
# -- 生命周期(PRD 5.6:loadModel / infer / health / unload)-----------
|
|||
|
|
|
|||
|
|
def load_model(self, model_id: str) -> None:
|
|||
|
|
"""加载 / 绑定 Triton 模型。幂等:重复加载同一 model_id 不报错。"""
|
|||
|
|
target = model_id or self.model_name
|
|||
|
|
# Triton 服务端就绪(离线核需显式 start;真实 client 由部署保证)
|
|||
|
|
if hasattr(self._kernel, "start_server"):
|
|||
|
|
self._kernel.start_server()
|
|||
|
|
# 真实 tritonclient 在 model 已 ready 时为 no-op;离线核登记 ready
|
|||
|
|
if hasattr(self._kernel, "load"):
|
|||
|
|
self._kernel.load(target)
|
|||
|
|
self._loaded = True
|
|||
|
|
self._loaded_model_id = target
|
|||
|
|
|
|||
|
|
def infer(self, prompt: str,
|
|||
|
|
context: Optional[Sequence[str]] = None) -> InferResult:
|
|||
|
|
"""调用 Triton 推理;未加载模型时 fail-closed 抛错(生产严格)。"""
|
|||
|
|
if not self._loaded or self._loaded_model_id is None:
|
|||
|
|
raise RuntimeError(
|
|||
|
|
f"{self.name}: 未调用 load_model,禁止推理(fail-closed)")
|
|||
|
|
started = time.perf_counter()
|
|||
|
|
resp = self._kernel.infer(
|
|||
|
|
self._loaded_model_id, prompt, context,
|
|||
|
|
max_tokens=self.max_tokens)
|
|||
|
|
latency_ms = round((time.perf_counter() - started) * 1000.0, 3)
|
|||
|
|
return InferResult(
|
|||
|
|
text=resp["text"],
|
|||
|
|
backend_name=self.name,
|
|||
|
|
model_id=self._loaded_model_id,
|
|||
|
|
prompt_tokens=resp.get("prompt_tokens"),
|
|||
|
|
completion_tokens=resp.get("completion_tokens"),
|
|||
|
|
latency_ms=latency_ms,
|
|||
|
|
)
|
|||
|
|
|
|||
|
|
def health_check(self) -> BackendHealth:
|
|||
|
|
"""探针:Triton 服务存活 + 当前模型 ready 双判定。"""
|
|||
|
|
try:
|
|||
|
|
server_live = bool(self._kernel.is_server_live())
|
|||
|
|
model_ready = (server_live and
|
|||
|
|
bool(self._kernel.is_model_ready(self.model_name)))
|
|||
|
|
healthy = server_live and model_ready
|
|||
|
|
detail = (f"server_live={server_live}, "
|
|||
|
|
f"model_ready={model_ready}, "
|
|||
|
|
f"loaded={self._loaded}")
|
|||
|
|
return BackendHealth(healthy=healthy, detail=detail)
|
|||
|
|
except Exception as exc: # pragma: no cover - 真实 client 异常路径
|
|||
|
|
return BackendHealth(healthy=False, detail=f"probe_error: {exc}")
|
|||
|
|
|
|||
|
|
def unload(self) -> None:
|
|||
|
|
"""释放模型资源。幂等:未加载时调用不报错。"""
|
|||
|
|
if self._loaded_model_id is not None and hasattr(self._kernel, "unload"):
|
|||
|
|
self._kernel.unload(self._loaded_model_id)
|
|||
|
|
self._loaded = False
|
|||
|
|
self._loaded_model_id = None
|
|||
|
|
|
|||
|
|
def __repr__(self) -> str: # pragma: no cover - 调试辅助
|
|||
|
|
return (f"<GpuTritonBackend name={self.name!r} "
|
|||
|
|
f"server={self.server_url!r} loaded={self._loaded}>")
|