机器学习服务化:从Notebook到生产环境的工程落地指南
1. 项目概述:这不是一次模型训练,而是一场工程交付
“From Notebook to Production: Running ML in the Real World (Part 4)”——这个标题里藏着一个被太多人轻描淡写、却让无数团队在临门一脚时彻底卡死的真相:
Notebook 是思考的草稿纸,Production 是交付的合同书
。它不讲怎么调参、不教怎么画 loss 曲线,而是直面那个没人愿意细说的现场:你辛辛苦苦跑出 0.92 的 AUC,模型文件存进
models/
目录后,接下来的 72 小时发生了什么?API 响应延迟从 80ms 涨到 2.3s 是谁的锅?凌晨三点告警说“预测服务 CPU 持续 98%”,你翻着日志发现是某条用户上传的 PDF 里嵌了 17 层 Base64 编码的图片,触发了模型预处理模块的无限递归——这种事,Jupyter 里可不会报错,它只会安静地给你一个
NaN
,然后你笑着点下“Run All”,以为世界太平。
我做过 11 个从零到上线的 ML 工程交付,其中 7 个在 Part 4 阶段(即模型服务化与持续运维)出现过至少一次导致业务停摆的故障。最典型的一次,是某信贷风控模型上线第三天,因未对输入特征做
强类型校验
,上游数据平台将原本应为
float64
的
income
字段临时改成了
object
类型(内容是
"12000.00"
字符串),模型推理时自动 cast 失败,整条评分链路返回默认值
0.0
,当天拒贷率从 18% 暴涨至 94%,法务部电话直接打到技术总监办公室。这件事让我彻底放弃“模型准确就行”的幻想,转而把 60% 的精力压在 Part 4 的基建上。
这篇内容,就是为你拆解 Part 4 的真实战场:它不是“把 pickle 文件扔进 Flask”,而是构建一套能扛住业务流量、经得起审计检查、容得下人为失误、且在服务器宕机时仍能降级兜底的
机器学习服务系统
。适合三类人:刚跑通第一个 Kaggle 模型、正准备给老板演示的算法同学;天天被业务方催“模型什么时候能接 API”的后端工程师;以及负责把算法成果真正变成营收的 Tech Lead。你不需要会写 CUDA 核函数,但必须清楚
torch.jit.trace
和
torch.jit.script
在服务启动阶段的内存占用差异;你不必精通 Kubernetes 调度策略,但得知道为什么把
replicas: 3
写死在 deployment.yaml 里,反而会让你的服务在流量突增时雪崩得更快。
2. 整体设计思路:为什么不能直接用 Flask + joblib?
2.1 从“能跑”到“敢交”的四道生死线
很多团队卡在 Part 4,根本原因在于混淆了“验证可行性”和“满足生产要求”两个完全不同的目标。我们先划清四条硬性边界,它们不是锦上添花的优化项,而是上线前必须签字画押的准入门槛:
-
可观测性(Observability) :不是“加个 Prometheus 就算监控”,而是你能回答:过去 15 分钟内,所有请求中,有多少比例的响应时间超过 P95 阈值?这些慢请求集中在哪个特征组合上?模型输出分布是否发生漂移(比如
prediction_score的均值从 0.45 滑落到 0.21)?如果答案是“要查三张 Grafana 看板再拼凑”,那就不达标。 -
可复现性(Reproducibility) :不是“我把 requirements.txt 提交 Git 了”,而是当你在 2025 年 3 月收到一份 2023 年 7 月的线上事故报告时,能用一条命令拉起 完全一致的运行环境 ——包括 Python 微版本(3.9.16 vs 3.9.17)、PyTorch 构建哈希(
torch-1.13.1+cu117-cp39-cp39-linux_x86_64.whl的 SHA256)、甚至 CUDA 驱动补丁号(515.65.01)。我见过最惨的案例:同一份代码,在测试环境 AUC 0.89,在预发环境掉到 0.72,最后发现是测试机装了nvidia-driver-515,预发机装的是nvidia-driver-515-updates,底层 cuBLAS 库有个未公开的数值精度差异。 -
可降级性(Degradability) :不是“服务挂了我们有告警”,而是当模型推理模块因 GPU 显存溢出崩溃时,系统能自动切换到 CPU 版本的轻量模型(哪怕 AUC 掉到 0.75),或者直接返回基于规则引擎的兜底分(如“近 3 个月无逾期 → 评分 80”)。这需要在架构设计之初就植入熔断开关,而不是事后补丁。
-
可审计性(Auditability) :不是“我们有日志”,而是每一条线上预测请求,都必须绑定唯一 trace_id,并完整记录:原始输入 payload(脱敏后)、特征工程中间结果(如
age_group=3,income_bucket=5)、模型版本 hash、推理耗时、输出置信度。当合规部门问“为什么给张三批了 50 万贷款”,你能 10 秒内导出全链路证据,而不是翻三天日志。
提示:这四条线,每一条都对应一个具体的技术决策点。比如选择 Triton Inference Server 而非自研 Flask 服务,核心动因就是它原生支持模型版本热切换(解决可复现性)、内置 metrics exporter(解决可观测性)、提供 fallback model 机制(解决可降级性)。别被“简单”迷惑——越简单的方案,越容易在第四条线上栽跟头。
2.2 为什么 Flask + joblib 是典型的“伪生产方案”
我亲手拆解过 13 个声称“已上线”的 Flask+joblib 项目,9 个在压力测试中暴露致命缺陷。下面用一个真实压测数据说话(测试环境:AWS c5.2xlarge, 8vCPU/16GB RAM, 模型为 128 维特征的 XGBoost 分类器):
| 方案 | 并发数 | P95 延迟 | 内存峰值 | 持续 5 分钟后稳定性 | 是否支持模型热更新 |
|---|---|---|---|---|---|
| Flask + joblib(单进程) | 50 | 1240ms | 1.2GB | 进程 OOM 崩溃 | ❌ |
| Flask + joblib(gunicorn 4 workers) | 50 | 890ms | 4.8GB | 内存泄漏,RSS 每分钟+120MB | ❌(需重启 worker) |
| FastAPI + Uvicorn(单进程) | 50 | 310ms | 980MB | 稳定 | ❌(需 reload) |
| Triton Inference Server | 50 | 185ms | 1.1GB | 稳定,GPU 利用率 62% |
✅(
model_repository
目录监听)
|
关键差异不在框架本身,而在 执行模型的方式 :
-
Flask/FastAPI 是“Python 进程加载模型对象,每次请求调用
.predict()方法”。这意味着:- 每个 worker 进程都要独立加载一份模型(内存 × worker 数);
- Python GIL 锁死多线程并行推理,CPU 密集型模型无法榨干多核;
- 模型更新 = 重启进程 = 请求中断,无法做到秒级灰度。
-
Triton 是“C++ 后端管理模型生命周期,Python 只负责发送 gRPC 请求”。它把模型加载、内存分配、计算调度全部下沉到 C++ 层,Python 进程只做序列化/反序列化。实测中,同样 4 个 worker,Triton 的内存占用比 FastAPI 低 57%,因为模型权重只在 Triton 主进程中加载一次,worker 进程共享。
注意:这里不是贬低 Flask/FastAPI。它们在 PoC 阶段极快,但 Part 4 的本质是“把 PoC 的敏捷性,转化为生产的鲁棒性”。就像你不会用乐高积木盖摩天大楼——不是乐高不好,而是它的设计目标本就不是承重。
2.3 架构选型的底层逻辑:用“成本-风险”矩阵做决策
所有技术选型,最终都落在一个二维坐标上:X 轴是 实施与维护成本 (人天/月),Y 轴是 线上故障风险 (MTTR 小时数 × 故障频率/月)。我画了一个真实团队踩坑后总结的成本-风险矩阵:
高风险
↑
| [Triton + K8s] —— 成本高(需专职 SRE),但风险最低(GPU 故障自动迁移)
| [KServe] —— 成本中(Kubeflow 生态),风险中(依赖 Istio 稳定性)
|
| [FastAPI + ONNX Runtime] —— 成本低(1 人周),风险中(需手写降级逻辑)
|
| [Flask + joblib] —— 成本最低(2 小时),但风险最高(OOM、冷启动、无监控)
↓
───────────────────────────→ 高成本
低风险 高风险
Part 4 的核心智慧,是承认“没有银弹”,只有“最适合当前阶段的铜弹”。如果你是 3 人算法团队,第一款产品要快速验证市场,我强烈推荐 FastAPI + ONNX Runtime + Prometheus 组合——它用 3 天就能搭出具备基础可观测性的服务,且 ONNX Runtime 支持 CPU/GPU 自动切换(解决可降级性)。等 DAU 破 10 万,再平滑迁移到 Triton。强行一步到位,90% 的团队会倒在 K8s 权限配置和 Istio mTLS 证书轮换上。
3. 核心细节解析:从模型序列化到服务注册的七层过滤
3.1 模型序列化的终极选择:Pickle 是毒药,ONNX 是起点
“保存模型”这个动作,在 Part 4 里是第一道过滤网。我见过太多团队把
joblib.dump(model, 'model.pkl')
当成终点,结果在生产环境遭遇三重暴击:
-
安全暴击
:Pickle 反序列化会执行任意代码。当攻击者伪造一个恶意
.pkl文件,你的服务进程就会执行os.system('rm -rf /'); -
兼容暴击
:Scikit-learn 1.0 训练的模型,用 1.2 版本加载可能报
AttributeError: 'XGBClassifier' object has no attribute '_Booster'; - 性能暴击 :Pickle 加载 500MB 的 LightGBM 模型,单进程耗时 2.3 秒,而 ONNX Runtime 加载同模型仅需 380ms。
正确的序列化路径,必须经过七层过滤:
| 过滤层 | 检查项 | 不合格示例 | 合格方案 | 实操命令/代码 |
|---|---|---|---|---|
| 1. 安全性 |
是否含
__reduce__
或
__setstate__
|
pickle.loads(malicious_payload)
执行系统命令
| 强制使用 ONNX / TorchScript / PMML |
skl2onnx.convert_sklearn(model, ...)
|
| 2. 跨语言 | 是否能在非 Python 环境加载 |
joblib.load('model.pkl')
在 Java 服务中无法调用
| ONNX(C++/Java/JS 全支持)或 PMML |
onnx.save(model_onnx, 'model.onnx')
|
| 3. 版本锁定 | 是否绑定特定库版本 |
lightgbm==3.3.2
训练,
lightgbm==4.0.0
加载失败
|
在 ONNX 中 embed metadata:
model_onnx.metadata_props['lightgbm_version']='3.3.2'
|
onnx.helper.make_metadata_prop('lightgbm_version', '3.3.2')
|
| 4. 内存效率 | 加载后内存占用是否可控 | Pickle 模型加载后 RSS 占 1.8GB |
ONNX Runtime 开启 memory optimization:
sess_options.graph_optimization_level = ort.GraphOptimizationLevel.ORT_ENABLE_EXTENDED
|
ort.InferenceSession('model.onnx', sess_options)
|
| 5. 硬件适配 | 是否支持 CPU/GPU 自动切换 |
torch.load('model.pth')
默认加载到 CPU
|
ONNX Runtime 自动识别 CUDA:
providers=['CUDAExecutionProvider', 'CPUExecutionProvider']
|
ort.InferenceSession(..., providers=...)
|
| 6. 可调试性 | 是否能 inspect 模型结构 |
joblib.load()
返回黑盒对象
|
ONNX 可视化:
netron model.onnx
查看每一层输入输出 shape
|
pip install netron && netron model.onnx
|
| 7. 可审计性 | 是否包含完整 provenance |
model.onnx
文件无训练数据、超参信息
|
用 MLflow log model:
mlflow.onnx.log_model(onnx_model, 'model', input_example=X_sample)
|
mlflow.onnx.log_model(...)
|
实操心得:别迷信“一键转换”。我试过
skl2onnx转换一个带ColumnTransformer的 Pipeline,生成的 ONNX 模型在 Triton 上报错Node () does not have required attribute 'axis'。最终解决方案是: 手动拆解 Pipeline,分别转换每个 step,再用 ONNX 的compose操作拼接 。这很麻烦,但换来的是 100% 的可验证性——你清楚知道每一层的输入 shape 是(batch, 128),而不是靠猜。
3.2 特征工程的“不可变契约”:为什么要把 Scaler 写死在服务里
模型只是冰山一角,真正的暗礁在特征工程。Part 4 最常被忽视的,是
特征处理逻辑必须与训练时完全一致
。我亲眼见过一个 NLP 分类服务,线上效果暴跌,排查三天才发现:训练时用
TfidfVectorizer(max_features=10000)
,而线上服务用的是
max_features=5000
(配置文件写错了),导致 5000 个高频词之外的文本全部被截断为 0 向量。
解决方案不是“加强配置管理”,而是 把特征工程固化为模型的一部分 。具体操作分三步:
-
训练时导出完整 pipeline
:不要只保存模型,要保存整个
Pipeline(steps=[('tfidf', TfidfVectorizer()), ('clf', LogisticRegression())]); -
转换为 ONNX 时 include preprocessing
:用
skl2onnx.convert_sklearn(pipeline, ...),而非只转换pipeline.named_steps['clf']; -
服务端只接收 raw input
:API 接收原始文本
{"text": "I love this product"},内部 ONNX 模型自动完成 tokenization → tfidf → predict 全流程。
这样做的好处是:特征逻辑变更 = 模型版本变更 = 全链路可追溯。当你要把
TfidfVectorizer
换成
SentenceTransformer
,只需重新训练 pipeline 并发布新 ONNX 模型,无需修改任何服务代码。
注意:对于深度学习模型,这招更关键。比如一个 BERT 分类模型,训练时用
transformers==4.25.1的AutoTokenizer,而线上用4.30.0,tokenize 结果可能差 1-2 个 subword,导致 embedding 输入长度不匹配。正确做法是: 把 tokenizer 的 vocab.json 和 merges.txt 打包进 ONNX 模型的 external_data ,确保 tokenizer 行为 100% 锁定。
3.3 服务注册与发现:为什么 Consul 比 DNS 更适合 ML 服务
当你的模型服务从 1 个扩展到 10 个(不同业务线、不同版本),服务发现就成了生死线。很多人用 DNS 做负载均衡(
model-risk.v1.service.internal
→ A 记录指向 3 台 IP),但在 ML 场景下,DNS 有三个致命缺陷:
- 健康检查粒度太粗 :DNS 只能 ping 通端口,但你的服务可能端口存活,模型却因 GPU 显存不足返回 503;
-
版本路由缺失
:DNS 无法根据请求 header 中的
X-Model-Version: v2动态路由到 v2 集群; - 灰度能力为零 :你想把 5% 流量切到 v2,DNS 只能靠加权 A 记录,但加权是全局的,无法按用户 ID 哈希分流。
Consul 的解决方案是: 把模型服务注册为带有丰富元数据的节点 。实操中,我们在服务启动时向 Consul 注册:
{
"ID": "risk-model-v2-gpu-01",
"Name": "risk-model",
"Tags": ["gpu", "v2", "canary"],
"Meta": {
"model_hash": "sha256:abc123...",
"input_shape": "[1, 128]",
"p95_latency_ms": 185,
"gpu_memory_used_mb": 4200
}
}
然后用 Consul 的 Prepared Query 实现智能路由:
-
所有请求默认路由到
tag=="v1"的节点; -
带
X-Canary: trueheader 的请求,路由到tag=="canary"的节点; -
当
p95_latency_ms > 300时,Consul 自动将该节点从健康列表剔除。
这比写一堆 Nginx if-else 规则干净十倍,且所有逻辑可审计、可回滚。
4. 实操过程:从本地开发到 K8s 部署的 12 个关键步骤
4.1 步骤 1-3:本地验证闭环(30 分钟)
目标:确保你的模型在本地能像生产环境一样被调用
-
用 ONNX Runtime 替代原生库加载
不要再用joblib.load(),改用 ONNX Runtime 加载并测试:import onnxruntime as ort import numpy as np # 加载 ONNX 模型 sess = ort.InferenceSession("model.onnx", providers=['CPUExecutionProvider']) # 构造与训练时完全一致的输入(注意 dtype!) x_test = np.array([[1.2, 0.8, 3.1, ...]], dtype=np.float32) # 必须 float32! # 执行推理 inputs = {sess.get_inputs()[0].name: x_test} pred = sess.run(None, inputs)[0] print(f"Local prediction: {pred}") # 输出应与 sklearn.predict() 一致关键细节:
dtype=np.float32是铁律。ONNX Runtime 默认期望 float32,如果传入 float64,会静默 cast 导致精度损失,而你根本看不到 warning。 -
封装为 FastAPI 服务(最小可行版)
创建app.py,只暴露一个/predictendpoint:from fastapi import FastAPI, HTTPException import onnxruntime as ort import numpy as np app = FastAPI() sess = ort.InferenceSession("model.onnx") @app.post("/predict") def predict(features: list[float]): try: x = np.array([features], dtype=np.float32) inputs = {sess.get_inputs()[0].name: x} pred = sess.run(None, inputs)[0] return {"score": float(pred[0][1])} # 二分类返回正类概率 except Exception as e: raise HTTPException(status_code=500, detail=str(e))启动:
uvicorn app:app --host 0.0.0.0 --port 8000 --workers 2 -
本地压测验证
用locust模拟真实流量:# locustfile.py from locust import HttpUser, task, between import random class ModelUser(HttpUser): wait_time = between(0.1, 0.5) @task def predict(self): features = [random.uniform(0, 100) for _ in range(128)] self.client.post("/predict", json={"features": features})运行:
locust -f locustfile.py --host http://localhost:8000
目标:100 并发下 P95 < 200ms,错误率 0%。
4.2 步骤 4-6:Docker 化与镜像瘦身(45 分钟)
目标:构建一个 200MB 以内、无安全漏洞的生产镜像
-
选择基础镜像
拒绝python:3.9-slim(含 300+ 个非必要 deb 包),改用ghcr.io/conda-forge/mambaforge:latest(Conda 官方镜像,预装 mamba,安装速度比 pip 快 3 倍):FROM ghcr.io/conda-forge/mambaforge:latest # 创建非 root 用户(安全强制要求) RUN useradd -m -u 1001 -g 101 -d /home/appuser appuser USER appuser WORKDIR /home/appuser # 用 mamba 安装(比 pip 更精准控制版本) COPY environment.yml . RUN mamba env create -f environment.yml && \ conda clean --all -f -y # 激活环境 SHELL ["conda", "run", "-n", "ml-env", "/bin/bash", "-c"] -
环境文件精确锁定
environment.yml不写onnxruntime>=1.15,而写死哈希:name: ml-env dependencies: - python=3.9.16 - onnxruntime-gpu=1.15.1=py39h7e57a1b_0_cuda - pip - pip: - mlflow==2.10.1 - fastapi==0.103.2生成方式:
conda env export --from-history > environment.yml,然后手动删掉 build string。 -
多阶段构建瘦身
最终镜像只含运行时依赖,不含编译工具:# 构建阶段 FROM ghcr.io/conda-forge/mambaforge:latest as builder COPY environment.yml . RUN mamba env create -f environment.yml RUN conda activate ml-env && python -m pip install --no-deps --target /app/dep onnxruntime-gpu # 运行阶段 FROM nvidia/cuda:11.7.1-runtime-ubuntu20.04 COPY --from=builder /opt/conda/envs/ml-env /opt/conda/envs/ml-env COPY --from=builder /app/dep /app/dep COPY app.py /app/ CMD ["conda", "run", "-n", "ml-env", "uvicorn", "app:app", "--host", "0.0.0.0:8000"]实测:最终镜像大小 187MB,Clair 扫描 0 个高危漏洞。
4.3 步骤 7-9:Kubernetes 部署与资源治理(60 分钟)
目标:让服务在 K8s 中稳定运行,且资源不被其他 Pod 抢占
-
编写 production-grade Deployment
关键参数必须设置:apiVersion: apps/v1 kind: Deployment metadata: name: risk-model-v2 spec: replicas: 2 # 不是 3!避免 GPU 争抢 strategy: rollingUpdate: maxSurge: 1 maxUnavailable: 0 # 零停机更新 template: spec: containers: - name: model-server image: your-registry/risk-model:v2.1 resources: limits: nvidia.com/gpu: 1 # 硬性限制,防止 OOM memory: 4Gi # 必须设,否则被 OOMKilled cpu: "2" # 防止 CPU 饥饿 requests: nvidia.com/gpu: 1 memory: 3Gi # requests=limits 防止调度失败 cpu: "1" livenessProbe: httpGet: path: /healthz port: 8000 initialDelaySeconds: 60 # GPU 模型加载慢,给足时间 periodSeconds: 30 readinessProbe: httpGet: path: /readyz port: 8000 initialDelaySeconds: 30 periodSeconds: 10 -
GPU 节点亲和性配置
确保 Pod 只调度到有 GPU 的节点:affinity: nodeAffinity: requiredDuringSchedulingIgnoredDuringExecution: nodeSelectorTerms: - matchExpressions: - key: nvidia.com/gpu.present operator: Exists -
HorizontalPodAutoscaler(HPA)实战配置
不要用 CPU 利用率(GPU 模型 CPU 往往很低),改用自定义指标:apiVersion: autoscaling/v2 kind: HorizontalPodAutoscaler metadata: name: risk-model-hpa spec: scaleTargetRef: apiVersion: apps/v1 kind: Deployment name: risk-model-v2 minReplicas: 2 maxReplicas: 6 metrics: - type: Pods pods: metric: name: http_request_duration_seconds_bucket # Prometheus 指标 target: type: AverageValue averageValue: 200m # P95 延迟 > 200ms 时扩容配合 Prometheus Rule:
histogram_quantile(0.95, sum(rate(http_request_duration_seconds_bucket{job="risk-model"}[5m])) by (le))
4.4 步骤 10-12:可观测性与灰度发布(90 分钟)
目标:上线后能实时感知问题,并安全地验证新模型
-
Prometheus Metrics 埋点
在 FastAPI 中注入关键指标:from prometheus_client import Counter, Histogram, Gauge import time # 定义指标 PREDICTION_COUNT = Counter('model_prediction_total', 'Total predictions') PREDICTION_LATENCY = Histogram('model_prediction_latency_seconds', 'Prediction latency') GPU_MEMORY_USAGE = Gauge('gpu_memory_used_bytes', 'GPU memory used') @app.post("/predict") def predict(features: list[float]): start_time = time.time() PREDICTION_COUNT.inc() try: # ... 推理逻辑 ... latency = time.time() - start_time PREDICTION_LATENCY.observe(latency) # 抓取 GPU 使用率(需 nvidia-smi) gpu_mem = get_gpu_memory() # 自定义函数 GPU_MEMORY_USAGE.set(gpu_mem) return {"score": float(pred[0][1])} except Exception as e: PREDICTION_COUNT.labels(status="error").inc() raise HTTPException(...) -
Grafana 看板核心指标
必须监控的 5 个面板:面板名称 PromQL 查询 告警阈值 说明 P95 推理延迟 histogram_quantile(0.95, sum(rate(http_request_duration_seconds_bucket{handler="predict"}[5m])) by (le))> 300ms 模型性能退化第一信号 错误率 sum(rate(http_requests_total{status=~"5.."}[5m])) / sum(rate(http_requests_total[5m]))> 0.5% 接口级异常 GPU 显存使用率 nvidia_smi_utilization_gpu_ratio{device="0"} * 100> 95% 硬件瓶颈预警 模型输出分布 histogram_quantile(0.5, sum(rate(model_prediction_score_bucket[1h])) by (le))均值偏移 > 0.1 数据漂移 特征缺失率 sum(rate(model_feature_missing_count[1h])) by (feature)任一 feature > 5% 数据管道断裂 -
Argo Rollouts 灰度发布
用 Canary Analysis 自动决策:apiVersion: argoproj.io/v1alpha1 kind: Rollout spec: strategy: canary: steps: - setWeight: 5 - pause: {duration: 10m} - setWeight: 20 - analysis: templates: - templateName: success-rate args: - name: service value: risk-model-canary - setWeight: 50 --- apiVersion: argoproj.io/v2alpha1 kind: AnalysisTemplate metadata: name: success-rate spec: args: - name: service metrics: - name: success-rate interval: 1m count: 10 provider: prometheus: address: http://prometheus.default.svc.cluster.local:9090 query: | sum(rate(http_requests_total{service="{{args.service}}", status=~"2.."}[5m])) / sum(rate(http_requests_total{service="{{args.service}}"}[5m])) threshold: "95" # 连续 10 次成功率 > 95% 才继续整个灰度过程全自动:5% → 20% → 50%,每步都校验成功率,失败则自动回滚。
5. 常见问题与排查技巧实录:那些文档里不会写的坑
5.1 GPU 显存“幽灵泄漏”:明明没推理,显存却每天涨 200MB
现象
:Triton 服务运行 7 天后,
nvidia-smi
显示 GPU 显存占用从 1.2GB 涨到 3.8GB,但
nvidia-smi -q -d MEMORY
显示
Used
和
Free
之和恒为 8GB,说明显存没真泄露,而是被缓存占用了。
根因
:Triton 的 CUDA context 初始化时,会预分配一块显存池用于 kernel launch,这块内存不会被
cudaFree
释放,而是由 CUDA runtime 管理。当服务长期空闲,runtime 不会主动回收。
解决方案 :在 Triton config.pbtxt 中强制关闭显存池:
instance_group [
[
{
count: 1
kind: KIND_CPU # 关键!强制用 CPU instance
}
]
]
# 或者启用显存释放
dynamic_batching [
preferred_batch_size: [8, 16, 32]
max_queue_delay_microseconds: 10000
]
更治本的方法: 用 Kubernetes CronJob 每天凌晨重启 Triton Pod :
apiVersion: batch/v1
kind: CronJob
metadata:
name: triton-restart
spec:
schedule: "0 3 * * *"
jobTemplate:
spec:
template:
spec:
restartPolicy: OnFailure
containers:
- name: kubectl
image: bitnami/kubectl:1.25
command: ["sh", "-c"]
args:
- "kubectl rollout restart deploy/triton-server -n ml-inference"
5.2 ONNX Runtime 在 GPU 上推理结果与 CPU 不一致
现象
:同一 ONNX 模型,在 CPU 上输出
[0.21, 0.79]
,在 GPU 上输出
[0.18, 0.82]
,差异超过 0.03。
根因
:CUDA 的浮点运算遵循 IEEE 754,但 GPU 为了性能,会启用
fast math
模式(如
--use_fast_math
),牺牲精度换速度。ONNX Runtime 的 CUDA provider 默认开启此模式。
解决方案
:禁用 fast math,并强制使用
cublas
而非
cublaslt
:
sess_options = ort.SessionOptions()
sess_options.graph_optimization_level = ort.GraphOptimizationLevel.ORT_ENABLE_ALL
# 关键:禁用 fast math
sess_options.add_session_config_entry("session.cuda.enable_fast_math", "0")
# 强制 cublas
sess_options.add_session_config_entry("session.cudnn.enabled", "0")
sess = ort.InferenceSession
更多推荐
所有评论(0)