PyTorch模型生产部署实战:ONNX导出+Triton+K8s高可用架构
前阵子我把一个基于PyTorch训练的语义分割模型真正搬到生产环境最终选型是ONNX导出 Triton Inference Server部署再配合Kubernetes做高可用。整个链路跑通之后收益很明显推理性能提升、服务不再依赖Python解释器、扩容和版本发布也规范了很多。这篇内容适合刚接触模型生产部署的算法工程师也适合需要维护推理服务的后端同学我会把从模型导出到高可用架构的实战过程摊开讲包括那些文档里不会写的坑。1. 为什么生产部署要绕开PyTorch原生服务1.1 直接用PyTorch写推理服务的三个硬伤很多人训练完模型第一反应是写个Flask API里面model.load_state_dict然后model(input)直接返回。这个方案在原型验证时没问题但一旦到了生产问题一个接一个。第一个硬伤是依赖体积和运行环境。PyTorch的Python包动辄几个GBCUDAToolkit和cuDNN的版本还和驱动强绑定。你在一台干净机器上部署光是装环境就能劝退运维。再加上训练机和服务机的PyTorch版本不一致偶尔会出现load_state_dict报错或者算子行为变化排查成本很高。我们团队有一次就是训练环境升级了PyTorch结果线上服务因为某个算子的实现变化输出悄悄漂移了一个点查了大半天才定位到版本差异。第二个硬伤是Python的GIL和并发吞吐。Python线程面对高并发请求时计算密集的推理任务会被锁卡住。虽然可以用多进程或者异步但进程内存、模型副本数量、显存占用都要自己管理。实际压测下来单机吞吐很容易到达瓶颈而GPU利用率却很低。尤其是多路并发推理场景如果不想把模型拷贝到多个进程里就只能靠队列串行处理延迟和吞吐都很难看。第三个硬伤是可运维性。纯手写的推理服务没有内置的模型版本管理、健康检查、动态批处理、指标暴露这些能力。你想监控推理延迟和吞吐要么自己埋点要么再套一层Prometheus exporter。想平滑升级模型也得自己写切换逻辑否则流量直接断掉。到了后期模型数量变多每个模型一个Python服务资源占用和运维压力都会爆表。1.2 ONNX和Triton的组合到底解决了什么ONNXOpen Neural Network Exchange是一个开放模型格式核心作用是把模型从PyTorch的训练框架中“抽离”出来。导出之后模型的推理不再依赖PyTorch环境只需要一个ONNX Runtime、TensorRT或者其它后端就能跑。这意味着服务端可以更轻也可以按需选择特定硬件上最快的执行引擎。如果你后续要换到CPU推理、ARM边缘设备或者想用TensorRT榨干GPU性能ONNX都是一个很好的中间层。Triton Inference Server现在叫NVIDIA Triton本质是一个高性能推理服务框架。它本身不做训练专注解决“模型部署和服务化”这一层问题加载模型、管理模型版本、自动批量请求、调度GPU实例、输出Prometheus指标。把ONNX模型丢给Triton之后客户端只需要通过HTTP或GRPC发送输入Triton内部已经把这些调度都处理好了。我自己的体会是这套组合最大的价值是“解耦”。模型文件和推理服务完全分开算法团队负责导出并验证ONNX平台团队负责Triton和K8s的运维。模型要更新时不用改一行服务代码直接替换模型仓库里的文件再加载新版本就行。这也是后来我把整个部署链路定成这个组合的根本原因。2. ONNX导出从PyTorch到ONNX的完整流程与避坑清单2.1 导出前需要确认的三件事在敲torch.onnx.export之前一定要先确认三件事否则后面大概率会返工。第一模型是否处于eval()模式。如果忘记切模型里的Dropout和BatchNorm行为会不一致导出的ONNX在线上的表现会和训练时完全不同。这个问题我见过不止一次都是线上效果异常后回来查才发现导错了。如果在训练模式下导出Dropout的随机性会被一起trace进去导出结果等于一个随机版本线上推理结果可能每次都不同。第二输入维度和dtype是否跟真实部署一致。torch.onnx.export依赖一个dummy input去trace模型结构如果这个输入只是随便生成的很可能导出后和线上实际数据不匹配。尤其注意batch维度、图片通道顺序以及输入是float32还是int64。这些不一致会导致后续客户端传数据时各种隐性问题比如onnxruntime报shape不匹配或者输出结果异常。第三目标opset版本。ONNX有个概念叫Operator Set不同版本支持的算子数量不同。PyTorch 2.x导出时建议opset 17或18onnxruntime一般兼容这些版本。版本太低容易遇到“算子不支持”的报错版本太高又可能在老的推理引擎上不兼容。我自己的习惯是先在本地用目标推理后端的版本实测一下确认能加载再批量导出。做完这三件事就可以导出第一个模型了。但是要注意dummy_input的shape和数据类型要跟真实部署对齐否则ONNX里记录的shape信息会带偏后续优化器。下面这个例子是常见分类模型的导出方式包含动态batch轴和常规参数import torch import torch.onnx model MyModel().eval() dummy_input torch.randn(1, 3, 224, 224, devicecuda) torch.onnx.export( model, dummy_input, model.onnx, input_names[input], output_names[output], dynamic_axes{input: {0: batch}, output: {0: batch}}, opset_version17, )这里最需要注意的是dynamic_axes。因为我最初导出时想支持任意batch大小就把batch设成动态。后来才发现动态轴虽然方便但会让ONNX里多出很多Reshape、Expand这类算子推理性能会有折扣。如果你的线上请求batch大小相对固定比如Triton会做动态batching那建议导出时直接用固定shape把batch留给Triton去拼。2.2 动态轴和静态轴该怎么选很多教程都会说“设置动态轴可以提升灵活性”但实际生产里这是有代价的。动态轴的好处是模型能接收任意batch大小也可以处理不同尺寸的输入比如检测任务里图片长宽不固定。坏处是ONNX导出器会插入大量动态shape相关的算子有些推理后端比如TensorRT优化起来很吃力延迟会变高。另外一旦目标后端不支持某些动态算子部署就直接失败。比如TensorRT对动态shape的支持虽然这几年好了很多但在一些特殊结构上仍然会有额外构建时间和显存开销。我的建议是分场景处理如果是纯CNN分类、分割模型输入shape通常固定尽量导出静态shape如果是检测、关键点这类需要原生支持动态输入的任务再开动态轴但最好在ONNX导出后跑一次onnxruntime确认shape变化路径正常。Triton配置里可以通过max_batch_size和dynamic_batching来承载动态请求这比把动态轴全部暴露给模型更可控。我自己最终导出的是静态shape然后在Triton上做动态batching性能和灵活性都兼顾了。2.3 算子兼容与降级处理导出ONNX遇到最多的坑就是“算子不支持”。常见报错是Unsupported operator: aten::xxx。我遇到过一个自定义Attention模块里面用了torch.einsum和高阶函数导出时直接失败。解决方案是先用标准算子重写模型比如把einsum换成matmultranspose再用torch.onnx.is_onnx_support(module)检查支持情况。如果某个算子实在无法替换可以在onnxruntime里注册自定义算子但这样需要自己实现CPU/GPU内核维护成本高能避免尽量避免。还有一次我们用了PyTorch 2.0刚出的某个新算子导出报错后来升级opset版本就通过了所以看到报错别急着改模型先查一下是不是版本兼容问题。另一个建议是使用onnx-simplifier。这个工具能识别冗余算子、合并常量把导出的ONNX瘦身。有时候导出结果包含了大量训练框架引入的shape处理逻辑simplify之后不仅文件小了推理速度也更快。命令行如下python -m onnxsim model.onnx model_sim.onnx跑完simplify后最好再重新校验一下精度因为个别结构重写会引入微小差异。我之前遇到过simplify后某个Reshape被错误合并导致输出shape对不上最后是在onnxruntime里跑了一遍才发现。2.4 导出后的精度校验与INT8量化模型导出后第一件事不是急着部署而是做精度对齐。我最常用的方法是拿同一个输入分别跑PyTorch和onnxruntime比较输出差异。import onnxruntime as ort import numpy as np sess ort.InferenceSession(model_sim.onnx, providers[CUDAExecutionProvider]) ort_out sess.run([output], {input: dummy_input.cpu().numpy()})[0] torch_out model(dummy_input).detach().cpu().numpy() print(np.max(np.abs(ort_out - torch_out)))如果差值在1e-4量级甚至更小说明导出没问题如果差值很大就要怀疑算子精度或者模型内部某些操作被错误trace了。最简单的排查方法是二分模型结构注释掉部分模块看哪个模块导出后差异变大。精度对齐之后如果模型太大或线上QPS要求高可以考虑量化。onnxruntime提供动态量化和静态量化。动态量化最简单from onnxruntime.quantization import quantize_dynamic, QuantType quantize_dynamic(model_sim.onnx, model_int8.onnx, weight_typeQuantType.QInt8)动态量化不需要校准数据直接把权重转成int8速度快但精度损失可能稍大。CNN模型如果想压得更狠要用静态量化需要准备一批有代表性的校准数据统计每层激活值的范围再量化。量化后务必要用相同的数据集重新评估模型指标不能只看单样本误差。曾经有个模型量化后单样本误差不大但在测试集上mAP掉了两个点最后只能回退到FP16。3. Triton Inference Server部署实战3.1 模型仓库目录结构Triton加载模型的方式不是直接指定一个文件而是指定一个“模型仓库”model repository。我的模型仓库长这样models/ └── my_model/ ├── 1/ │ └── model.onnx └── config.pbtxt1是版本号。Triton会默认加载每个模型目录里数字最大的版本。这意味着你以后放一个2/目录进去就能实现模型版本更新旧版本还在回滚也方便。实际运维时我用版本号目录做灰度先把新版本传上去观察一段时间确认稳定后再卸载旧版本。config.pbtxt是模型配置文件告诉Triton怎么处理这个模型。以我的分割模型为例name: my_model platform: onnxruntime_onnx max_batch_size: 8 input [ { name: input data_type: TYPE_FP32 dims: [3, 224, 224] } ] output [ { name: output data_type: TYPE_FP32 dims: [1024] } ]这里有个容易踩的坑Triton的max_batch_size不一定会自动把请求拼成batch。要开启自动拼batch还需要在config里配置dynamic_batching。我一开始只设了max_batch_size: 8以为就能自动合并请求结果压测发现每个请求还是单独推理GPU利用率很低。另外要注意dims里不需要写batch维度因为Triton会把batch维度和max_batch_size统一管理。3.2 动态批处理与实例配置动态批处理是Triton的看家本领之一。它允许把多个并发请求攒在一起凑成一个batch然后推理一次。配置很简单dynamic_batching { preferred_batch_size: [2, 4, 8] max_queue_delay_microseconds: 200 }preferred_batch_size告诉Triton尽量凑到2、4、8再推理max_queue_delay_microseconds告诉它最多等200微秒如果等不到足够请求就先发出去。这两个参数需要根据线上真实请求速率调。设置太小batch凑不满浪费设置太大请求排队时间变长p99延迟飙升。我一般先用perf_analyzer压测再结合线上请求到达率微调。另外还要配置GPU实例数量instance_group [ { count: 1 kind: KIND_GPU } ]这里的count表示启动几个模型实例。如果模型很小可以尝试count: 2让Triton在两个GPU stream上并行跑吞吐会明显提升但显存占用也会翻倍需要平衡。对于CPU模型可以把kind换成KIND_CPU并用count控制并发线程数。我通常先用默认值跑perf_analyzer再根据指标调整不会一上来就开满。3.3 客户端调用与压测Triton部署起来后客户端调用很简单。我用Python HTTP客户端时是这样写的import tritonclient.http as httpclient client httpclient.InferenceServerClient(urllocalhost:8000) input_data httpclient.InferInput(input, [1, 3, 224, 224], FP32) input_data.set_data_from_numpy(preprocessed_numpy) result client.infer(my_model, [input_data]) output_data result.as_numpy(output)这里要注意InferInput的shape和数据类型必须与config里的定义匹配batch维度要放在最前面。如果数据是uint8但模型要求FP32需要自己转换。这类问题在联调时很常见我建议客户端和服务端的schema都用同样的protobuf或者json定义管理避免两边手动对齐。压测推荐直接用Triton自带的perf_analyzerperf_analyzer -m my_model -u localhost:8000 --concurrency-range 1:10 --shape input:3,224,224它会自动发请求输出吞吐量、平均延迟、p99延迟。我一般先看GPU利用率再调整dynamic_batching目标是让GPU利用率达到70%以上同时p99延迟满足业务SLA。不要一味追求吞吐如果你在延迟敏感场景batch太大反而会拖垮尾延迟。4. 高可用架构把推理服务变成生产级4.1 容器化部署要点Triton本身提供了官方镜像生产环境建议直接基于它构建Docker镜像不要从零装环境。下面是我常用的DockerfileFROM nvcr.io/nvidia/tritonserver:23.10-py3 COPY ./models /models EXPOSE 8000 8001 8002 CMD [tritonserver, --model-repository/models]8000是HTTP端口8001是GRPC端口8002是Prometheus metrics端口。如果你只需要预测不一定要暴露所有端口但开放metrics对后面的监控很有用。我习惯把metrics端口单独暴露给集群内的Prometheus抓取不开放公网。这里有一个需要注意的坑镜像里的Triton版本和导出ONNX时使用的opset要匹配。如果opset版本过高Triton内置的onnxruntime版本可能不支持启动时会报错。我的做法是先在本地用同样版本镜像起一次加载模型确认日志里没有error再推仓库。另外模型文件比较大的时候不要把ONNX打进镜像而是用共享存储挂载进去这样模型更新不需要重新构建镜像。4.2 Kubernetes部署与自动扩缩容生产环境我一般用Kubernetes管理Triton。Kubernetes能提供副本扩展、滚动更新、健康检查这些能力配合Triton天然适合。Deployment的最小配置如下apiVersion: apps/v1 kind: Deployment metadata: name: triton-inference spec: replicas: 3 selector: matchLabels: app: triton-inference template: metadata: labels: app: triton-inference spec: containers: - name: triton image: myregistry/triton-server:1.0 ports: - containerPort: 8000 - containerPort: 8001 resources: limits: nvidia.com/gpu: 1 memory: 8Gi cpu: 4 readinessProbe: httpGet: path: /v2/health/ready port: 8000 initialDelaySeconds: 30 periodSeconds: 10注意resources这块GPU资源要写nvidia.com/gpu: 1。如果你的集群没有安装NVIDIA Device Plugin这个字段不会被识别。内存和CPU建议也写上防止某一路请求把节点资源打满。扩容方面我倾向于用HPA基于CPU虽然GPU利用率才是更准确指标但CPU如果是瓶颈同样有效。简单HPA配置apiVersion: autoscaling/v2 kind: HorizontalPodAutoscaler metadata: name: triton-hpa spec: scaleTargetRef: apiVersion: apps/v1 kind: Deployment name: triton-inference minReplicas: 2 maxReplicas: 10 metrics: - type: Resource resource: name: cpu target: type: Utilization averageUtilization: 60GPU副本扩容要谨慎。每扩一个副本就意味显存多占一份。我遇到过因为HPA扩容太多直接把集群里所有GPU节点显存占满其他训练任务被挤掉的情况。所以GPU服务更推荐把多个小模型塞进同一个Triton实例用一个副本扛住流量真上来了再考虑扩。4.3 健康检查与优雅下线Triton提供了两个健康检查端点/v2/health/live表示进程活着/v2/health/ready表示模型已加载完毕、可以接流量。Kubernetes的readinessProbe应该用/v2/health/ready否则刚启动时还没加载完模型流量就进来了会直接失败。livenessProbe可以用live端点或者不配也行因为这种服务很少会死锁。另外一个容易被忽略的是优雅下线。Kubernetes滚动更新时会先停掉旧Pod再起新Pod如果没有优雅停机正在处理的请求会被硬生生掐断。我在项目里给Triton容器加了preStop hook让它收到停止信号后先把自身从负载均衡摘除再处理完存量请求最后退出。一个简单的做法是lifecycle: preStop: exec: command: [/bin/sh, -c, sleep 10]或者更规范一点调用Triton的model control API先把模型unload再等Pod里的进程退出。你也可以通过调整terminationGracePeriodSeconds给足时间。我踩过坑某次升级没配这个线上一直有少量5xx查了一个多小时才发现是滚动更新导致的。从那以后所有推理服务的Pod模板里都默认加上preStop hook。5. 常见问题与排查技巧实录5.1 ONNX导出的精度偏差ONNX导出后精度和PyTorch不一致算是最高频的问题。除了前面说的eval模式还有几个原因。Transformer类模型尤其容易遇到输入长度动态时attention_mask可能被trace成常量导致新序列长度一变就出错。这种情况我建议固定序列长度或者把mask的处理放到模型外部ONNX只接收已经处理好的输入。另外PyTorch里的nn.functional.scaled_dot_product_attention在不同opset下展开结果可能不一样需要做精度对比验证。还有一个是opset版本。我遇到过同样一个模型opset 11导出精度正常opset 17导出后特征分布漂移最后查出来是某个算子在不同opset下行为有差异。解决办法是逐个opset版本测试选一个精度对齐最好的。这个工作最好提前做不要在线上出问题了再一个个版本试验。5.2 Triton请求超时和排队表现是客户端偶尔报超时但服务端日志看不出错误。这种情况多半是请求排队太长。我一般会看/v2/metrics关注nv_inference_queue_duration_us这个指标。如果queue duration很高说明动态批处理攒batch攒得太贪或者实例数不够。优化方向有三个调低max_queue_delay_microseconds让请求更快被处理适当增加instance_group.count让更多实例并行或者优化preferred_batch_size让batch更容易凑满。如果还是不行看客户端并发连接数是否打满。有一个细节是GRPC客户端默认有连接池大小限制如果并发超过连接数也会在客户端侧排队服务端完全无感知。5.3 GPU显存OOM显存不够是最头疼的。通常发生在几种情况多个模型同时加载实例数调太大动态批处理的batch太大。我的排查顺序是先用nvidia-smi看显存占用再用tritonserver --log-verbose1看模型加载日志最后调整config。一个比较实用的经验是给每个模型设置合理max_batch_size不要为了峰值吞吐牺牲稳定性。同时不用的模型版本记得清理Triton默认会加载所有版本如果你旧版本忘删显存会被白白占掉。可以通过--model-control-modeexplicit和model control API让模型按需加载。这样能节省不少显存但也增加了调用复杂度需要自己管理模型生命周期。5.4 模型热更新与版本管理Triton的模型版本目录天然支持热更新。你只要把新版本的模型文件放到新的数字目录下Triton会重新加载旧版本仍然保留出问题可以手动回滚。但我建议更新时不要直接覆盖正在服务的目录而是上传新版本目录观察指标稳定后再通过model control API把旧版本unload。这里还有一个细节如果模型文件很大加载过程会占住GPU显存和CPU。如果同时更新多个模型可能瞬间造成资源峰值。我把上线流程拆成“先上传、再加载、然后摘流量、最后卸载旧版”每一步都加监控基本没再出过幺蛾子。最后把常见问题整理成速查表方便遇到问题时快速对照现象可能原因解决办法导出ONNX后精度明显下降模型不是eval模式、动态轴trace异常、opset不兼容切eval、固定输入shape、逐个opset测试onnxruntime推理报shape错误客户端输入shape与config不匹配统一schema用InferInput显式设置shape请求超时但服务端无error排队延迟高、客户端连接池不够调整dynamic_batching参数、增加实例数GPU显存OOM实例数过多、max_batch_size过大、旧版本未清理调小批处理、清理旧版本、按需加载模型滚动更新时频繁5xx没有优雅停机、readiness探针配置不当加preStop hook、用ready端点做readinessProbe最后再分享一个小技巧Triton部署好后一定要把metrics接进Prometheus和Grafana。延迟、吞吐、排队时长、显存占用这些指标不怕多就怕没有。模型上线的第一周我基本每天都在看这些曲线很多隐患都是在曲线异常时提前发现的。等你把ONNX导出、Triton服务和K8s高可用都跑顺之后后续再上新模型就只是替换模型仓库版本的问题了。