英译中模型从 PyTorch 搬到 ONNX 这件事,我前前后后折腾了差不多两周,中间踩的坑比预想的多得多。起因很简单:线上服务用的是 Python + PyTorch 直接推理,单条翻译延迟还能忍,一旦并发上来,显存和 CPU 占用直接爆掉,扩容成本高得离谱。后来想着把模型导出成 ONNX,用 ONNX Runtime 跑推理,理论上能省掉一大半开销,还能顺带做 int8 量化进一步压榨性能。结果真动手才发现,HuggingFace 上的英译中模型五花八门,导出方式、算子兼容性、分词器处理、量化精度损失,每一个环节都能让你卡半天。这篇就把我完整的迁移过程、踩过的坑、以及最后跑通的方案整理出来,适合已经用过 HuggingFace Transformers、想把模型部署到生产环境但被性能问题困扰的同行参考。哪怕你之前没碰过 ONNX,跟着思路走也能搞明白。
1. 为什么英译中模型值得单独做一次 ONNX 迁移
1.1 英译中任务的推理特征决定了它适合 ONNX
英译中这类 seq2seq 翻译任务,和图像分类、文本分类那种"一次前向传播出结果"的模型完全不一样。它的推理过程是自回归的:编码器先把英文源句编码成隐状态,解码器再一个 token 一个 token 地往外吐中文,每吐一个 token 都要把之前所有已生成的 token 重新喂回去算一遍。这意味着推理时间跟输出长度成正比,而且每一步都涉及大量的矩阵运算和注意力计算。
在 PyTorch 原生推理下,这种自回归循环的每一步都要经过 Python 解释器、动态图调度、CUDA kernel 启动等一堆开销。单步可能只有几毫秒,但一条 30 个字的译文就要循环 30 次,累积起来延迟就很可观了。ONNX Runtime 的优势在于它把整个计算图静态化,做了算子融合、内存复用、线程调度优化,尤其是 CPU 推理场景下,性能提升往往能到 2 到 4 倍。对于英译中这种"输入短、输出中等长度、调用频繁"的场景,收益非常直接。
另外,英译中模型通常不会特别大。像 Helsinki-NLP 的 opus-mt-en-zh 这类模型,参数量在 70M 到 300M 之间,导出 ONNX 之后文件大小可控,量化到 int8 还能再砍一半以上。相比之下,那些动辄几十亿参数的大语言模型做 ONNX 导出反而麻烦,算子兼容性和显存占用都是问题。所以英译中模型是一个"投入产出比"很高的迁移对象。
1.2 PyTorch 直接部署在生产环境里的三个现实痛点
第一个痛点是依赖太重。PyTorch 本身安装包就好几百兆,加上 Transformers、Tokenizers、SentencePiece 这些依赖,一个推理服务的镜像轻松上 2GB。如果你用 Docker 部署,拉镜像、冷启动都是时间。ONNX Runtime 的包体积小得多,CPU 版本几十兆就够,部署轻量很多。
第二个痛点是并发下的资源争抢。PyTorch 的动态图在每次 forward 时都会重新构建计算图,虽然它有缓存机制,但在高并发下 GIL 锁、内存分配、CUDA 上下文切换都会成为瓶颈。我实测过,同样一台 4 核 8G 的机器,PyTorch 跑 8 并发时 P99 延迟飙到 800ms 以上,而 ONNX Runtime 能稳在 300ms 左右。
第三个痛点是跨平台部署困难。有些边缘设备或者特定架构的服务器根本装不了完整 PyTorch,但 ONNX Runtime 支持 x86、ARM、甚至一些嵌入式平台,通用性好太多。这一点在做私有化交付时特别重要,客户环境千奇百怪,ONNX 能省掉大量适配工作。
1.3 ONNX 迁移不是"导出就完事",真正的难点在哪
很多人以为torch.onnx.export()一跑,模型就迁移完了。实际上导出只是第一步,后面还有一堆事:导出后的模型输出和原模型是否数值一致?分词器的特殊 token 处理对不对?自回归解码的逻辑要不要自己重写?量化之后 BLEU 掉了多少?这些问题不解决,导出的 ONNX 就是个摆设。
尤其是 seq2seq 模型,HuggingFace 的generate()方法内部封装了 beam search、重复惩罚、长度惩罚等一堆逻辑,这些在 ONNX 里是没有的。你导出的是 encoder 和 decoder 的裸计算图,解码策略得自己在 ONNX Runtime 外面用 Python 实现。这一步是很多人卡住的地方,也是这篇博文要重点讲清楚的部分。
2. 迁移前的环境准备与模型选型
2.1 选哪个英译中模型:不是越大越好
HuggingFace 上英译中模型主要分几类。一类是 Helsinki-NLP 的 opus-mt 系列,比如Helsinki-NLP/opus-mt-en-zh,这是 MarianMT 架构,参数量小、速度快,适合对质量要求不是极致高的场景。另一类是 Facebook 的 mBART、mT5 系列,多语言能力强但模型大。还有一类是基于 Transformer 自训练的模型,质量参差。
我的建议是:如果你的场景是通用文本翻译,opus-mt-en-zh 是性价比最高的选择,导出 ONNX 也最顺。如果你需要处理专业领域文本,可以考虑在 opus-mt 基础上做微调,或者选 mT5-base 这类。但要注意,mT5 是 encoder-decoder 架构,导出时 decoder 的 past_key_values 处理比 MarianMT 复杂,坑更多。
选模型时还要看它的 tokenizer 类型。opus-mt 用的是 SentencePiece,导出 ONNX 时 tokenizer 本身不参与计算图,但你要确保推理时用的 tokenizer 和训练时一致,否则翻译结果会乱码或者语义漂移。
2.2 环境依赖的版本坑:transformers 和 onnxruntime 的兼容性
环境这块我踩过最大的坑是版本不匹配。transformers从 4.30 到 4.40 之间,模型导出 ONNX 的 API 有过几次变动,torch.onnx.export的参数默认值也改过。onnxruntime的版本同样关键,1.15 和 1.17 在算子支持上有差异,有些模型在旧版本能跑,新版本反而报错。
我最后跑通的组合是:Python 3.10、torch 2.1.2、transformers 4.36.2、onnx 1.15.0、onnxruntime 1.17.0。这个组合在 CPU 和 GPU 上都验证过。安装命令如下:
pip install torch==2.1.2 transformers==4.36.2 onnx==1.15.0 onnxruntime==1.17.0 pip install sentencepiece protobuf注意:
sentencepiece和protobuf一定要装,MarianMT 的 tokenizer 依赖它们。protobuf 版本别太高,3.20.x 比较稳,4.x 有时候会和 onnx 冲突。
如果你要用 GPU 推理,把onnxruntime换成onnxruntime-gpu,但要注意 CUDA 版本和 onnxruntime-gpu 的对应关系,装错了会直接 import 失败。
2.3 先把 PyTorch 原模型的基线跑出来
迁移之前一定要先跑通原模型,记录下基线指标,否则后面 ONNX 出问题你都不知道是导出错了还是模型本身就这样。基线要记录三样东西:翻译质量(找几十条测试句,人工看或者算 BLEU)、单条延迟、吞吐量。
from transformers import MarianMTModel, MarianTokenizer import time model_name = "Helsinki-NLP/opus-mt-en-zh" tokenizer = MarianTokenizer.from_pretrained(model_name) model = MarianMTModel.from_pretrained(model_name) model.eval() text = "The quick brown fox jumps over the lazy dog." inputs = tokenizer(text, return_tensors="pt", padding=True) start = time.time() with torch.no_grad(): generated = model.generate(**inputs, max_length=128, num_beams=4) elapsed = time.time() - start result = tokenizer.decode(generated[0], skip_special_tokens=True) print(f"译文: {result}") print(f"耗时: {elapsed*1000:.2f} ms")把这段跑通,记下译文和耗时。后面 ONNX 版本要用同样的输入对比,确保译文一致、耗时确实下降。这一步看着简单,但很多人跳过,结果后面出问题来回折腾。
3. 把 MarianMT 导出成 ONNX 的完整操作链路
3.1 理解 seq2seq 导出的核心难点:encoder 和 decoder 要分开导
这是整个迁移最关键的一个认知点。HuggingFace 的generate()是一个封装好的高层 API,它内部做了这些事:encoder 前向、初始化 decoder 输入、循环调用 decoder、beam search 剪枝、处理 EOS。ONNX 导出时,你没法把这一整套逻辑导成一个图,因为循环次数是动态的,beam search 涉及排序和选择操作,ONNX 对动态控制流支持有限。
正确做法是把模型拆成两部分导出:encoder 导成一个 ONNX,decoder 导成另一个 ONNX。推理时,先用 encoder ONNX 算出编码结果,然后在 Python 里写循环,每一步调用 decoder ONNX,自己实现 beam search 或者 greedy search。这样虽然多写点代码,但灵活性和可控性都更好。
MarianMT 的结构是标准的 encoder-decoder,encoder 输出last_hidden_state,decoder 接收input_ids、encoder_hidden_states、encoder_attention_mask,输出logits和past_key_values。导出时要把past_key_values作为输入输出都暴露出来,这样自回归时才能复用缓存,不然每步都重算全部历史,性能会很差。
3.2 encoder 导出:注意 attention_mask 的动态维度
encoder 导出相对简单,但有个细节容易忽略:attention_mask的维度必须是动态的。如果你导出时用了固定长度,推理时输入长度不一样就会报错。
import torch from transformers import MarianMTModel, MarianTokenizer model_name = "Helsinki-NLP/opus-mt-en-zh" tokenizer = MarianTokenizer.from_pretrained(model_name) model = MarianMTModel.from_pretrained(model_name) model.eval() # 构造 dummy input dummy_text = "This is a test sentence for onnx export." dummy_inputs = tokenizer(dummy_text, return_tensors="pt", padding=True) input_ids = dummy_inputs["input_ids"] attention_mask = dummy_inputs["attention_mask"] # 导出 encoder torch.onnx.export( model.get_encoder(), (input_ids, attention_mask), "encoder_model.onnx", input_names=["input_ids", "attention_mask"], output_names=["encoder_hidden_states"], dynamic_axes={ "input_ids": {0: "batch", 1: "sequence"}, "attention_mask": {0: "batch", 1: "sequence"}, "encoder_hidden_states": {0: "batch", 1: "sequence"} }, opset_version=14, do_constant_folding=True )这里opset_version我用的 14,因为 MarianMT 里有些算子在高版本 opset 下行为有变化,14 是比较稳的选择。dynamic_axes一定要把 batch 和 sequence 两个维度都标成动态,否则推理时长度一变就挂。
3.3 decoder 导出:past_key_values 的处理是重头戏
decoder 导出是整个流程里最麻烦的部分。MarianMT 的 decoder 在自回归时,第一次调用需要完整的encoder_hidden_states,后续调用只需要传新生成的 token 和缓存的past_key_values。导出时要把这个逻辑表达清楚。
# 构造 decoder 的 dummy 输入 batch_size = 1 encoder_seq_len = input_ids.shape[1] decoder_seq_len = 1 # decoder 初始输入,通常是 decoder_start_token_id decoder_input_ids = torch.tensor([[model.config.decoder_start_token_id]]) # encoder 输出 with torch.no_grad(): encoder_outputs = model.get_encoder()(input_ids, attention_mask) # 构造 past_key_values 的 dummy(第一次调用时为空) # MarianMT 的 past_key_values 结构需要根据模型配置确定实际操作中,直接导出带past_key_values的 decoder 比较绕,因为 HuggingFace 的past_key_values是一个嵌套元组,ONNX 对嵌套结构的支持不好。我的做法是参考 HuggingFace 官方convert_to_onnx脚本里的处理方式,把past_key_values展平成多个输入输出张量。
一个更省事的方案是用optimum库,它专门做了 HuggingFace 模型到 ONNX 的转换,内部处理好了这些细节:
pip install optimum[onnxruntime]from optimum.onnxruntime import ORTModelForSeq2SeqLM from transformers import AutoTokenizer model_id = "Helsinki-NLP/opus-mt-en-zh" tokenizer = AutoTokenizer.from_pretrained(model_id) ort_model = ORTModelForSeq2SeqLM.from_pretrained(model_id, export=True) ort_model.save_pretrained("./onnx_model") tokenizer.save_pretrained("./onnx_model")optimum会自动把模型拆成 encoder、decoder、decoder_with_past 三个 ONNX 文件,并且处理好past_key_values的展平。这是目前最省心的方案,强烈推荐。如果你非要手写导出,那就得仔细研究optimum的源码,看它怎么处理past_key_values的。
3.4 导出后的数值一致性验证
导出完不能直接用,必须先验证 ONNX 输出和 PyTorch 输出是否一致。做法是拿同一组输入,分别跑 PyTorch 和 ONNX,对比输出的 logits。
import numpy as np import onnxruntime as ort # PyTorch 输出 with torch.no_grad(): pt_outputs = model.get_encoder()(input_ids, attention_mask) pt_hidden = pt_outputs.last_hidden_state.numpy() # ONNX 输出 sess = ort.InferenceSession("encoder_model.onnx") onnx_hidden = sess.run( ["encoder_hidden_states"], { "input_ids": input_ids.numpy(), "attention_mask": attention_mask.numpy() } )[0] # 对比 diff = np.abs(pt_hidden - onnx_hidden).max() print(f"最大差异: {diff}")正常情况下,最大差异应该在 1e-4 到 1e-5 量级,这是浮点精度导致的,可以接受。如果差异到了 0.1 以上,说明导出有问题,通常是 opset 版本或者算子实现不一致导致的,需要排查。
4. 自回归解码逻辑在 ONNX Runtime 上的重写
4.1 为什么不能直接用 generate():ONNX 只给了你计算图
前面说过,generate()里的 beam search、长度惩罚、重复惩罚这些逻辑,ONNX 里没有。你导出的是裸计算图,解码策略得自己写。这不是 ONNX 的缺陷,而是设计使然——ONNX 定位是"模型交换格式",不是"推理框架",它只负责计算图,不负责解码策略。
所以你需要自己实现一个解码循环。最简单的 greedy search 大概长这样:
def greedy_decode(encoder_sess, decoder_sess, input_ids, attention_mask, max_length=128, eos_token_id=0): # encoder 前向 encoder_hidden = encoder_sess.run( ["encoder_hidden_states"], {"input_ids": input_ids, "attention_mask": attention_mask} )[0] # 初始 decoder 输入 decoder_input_ids = np.array([[model.config.decoder_start_token_id]]) generated = [] for step in range(max_length): outputs = decoder_sess.run( None, { "input_ids": decoder_input_ids, "encoder_hidden_states": encoder_hidden, "encoder_attention_mask": attention_mask } ) logits = outputs[0] next_token = np.argmax(logits[:, -1, :], axis=-1) if next_token[0] == eos_token_id: break generated.append(next_token[0]) decoder_input_ids = np.concatenate( [decoder_input_ids, next_token[:, None]], axis=1 ) return generated这段代码能跑,但性能很差,因为每一步都把完整的decoder_input_ids重新喂进去,没有用past_key_values缓存。正确的做法是用带 cache 的 decoder,每步只传新 token。
4.2 用 past_key_values 缓存把解码速度提上来
带 cache 的解码逻辑是这样的:第一次调用 decoder 时传完整的decoder_input_ids,拿到past_key_values;后续每次只传上一个生成的 token,同时把past_key_values传进去,decoder 会复用缓存,只计算新 token 的注意力。
optimum导出的decoder_with_past模型就是干这个的。它的输入包括input_ids(只有新 token)、encoder_hidden_states、past_key_values的各个张量,输出包括logits和更新后的past_key_values。
def decode_with_cache(encoder_sess, decoder_sess, decoder_with_past_sess, input_ids, attention_mask, max_length=128): encoder_hidden = encoder_sess.run( ["encoder_hidden_states"], {"input_ids": input_ids, "attention_mask": attention_mask} )[0] # 第一次 decoder 前向 decoder_input_ids = np.array([[model.config.decoder_start_token_id]]) outputs = decoder_sess.run( None, { "input_ids": decoder_input_ids, "encoder_hidden_states": encoder_hidden, "encoder_attention_mask": attention_mask } ) logits = outputs[0] past_kv = outputs[1:] generated = [] next_token = np.argmax(logits[:, -1, :], axis=-1) for step in range(max_length - 1): if next_token[0] == eos_token_id: break generated.append(next_token[0]) # 用 cache 前向 inputs = { "input_ids": next_token[:, None], "encoder_hidden_states": encoder_hidden, "encoder_attention_mask": attention_mask } # 把 past_key_values 填进去 for i, kv in enumerate(past_kv): inputs[f"past_key_values.{i}"] = kv outputs = decoder_with_past_sess.run(None, inputs) logits = outputs[0] past_kv = outputs[1:] next_token = np.argmax(logits[:, -1, :], axis=-1) return generated这个逻辑跑通之后,解码速度会有明显提升,因为每步只计算一个新 token 的注意力,而不是重算全部历史。
4.3 beam search 要不要自己实现:看你的质量要求
greedy search 速度快但质量一般,beam search 质量好但慢。如果你对翻译质量要求高,就得自己实现 beam search。beam search 的核心是维护 k 个候选序列,每步扩展后按累积概率排序,保留 top-k。
在 ONNX Runtime 上实现 beam search 的难点在于:每个 beam 的past_key_values要单独维护,而且 batch 维度会变成 beam_size。实现起来代码量不小,而且容易出 bug。我的建议是:如果 greedy search 的质量能接受,就别折腾 beam search;如果非要 beam search,可以考虑用optimum的ORTModelForSeq2SeqLM,它内部封装了 beam search 逻辑,直接调generate()就行。
from optimum.onnxruntime import ORTModelForSeq2SeqLM from transformers import AutoTokenizer model = ORTModelForSeq2SeqLM.from_pretrained("./onnx_model") tokenizer = AutoTokenizer.from_pretrained("./onnx_model") inputs = tokenizer("The quick brown fox jumps over the lazy dog.", return_tensors="pt") outputs = model.generate(**inputs, max_length=128, num_beams=4) print(tokenizer.decode(outputs[0], skip_special_tokens=True))用optimum的好处是它把解码逻辑都封装好了,你只需要调generate(),和用原生 Transformers 的体验一样。代价是灵活性差一些,但大多数场景够用。
5. int8 量化:精度和速度的平衡怎么找
5.1 动态量化 vs 静态量化:英译中模型该选哪个
ONNX Runtime 支持两种量化方式:动态量化和静态量化。动态量化在推理时动态计算激活值的量化参数,不需要校准数据,用起来简单;静态量化需要一批校准数据预先算好量化参数,精度通常更好,但准备校准集麻烦。
对于英译中模型,我推荐先用动态量化试。原因是:翻译模型的激活值分布比较稳定,动态量化的精度损失通常可控;而且动态量化不需要准备校准数据,省事。如果动态量化后 BLEU 掉得太多,再考虑静态量化。
from onnxruntime.quantization import quantize_dynamic, QuantType quantize_dynamic( model_input="encoder_model.onnx", model_output="encoder_model_int8.onnx", weight_type=QuantType.QInt8 )对 encoder 和 decoder 分别做量化,得到 int8 版本。注意 decoder_with_past 也要量化,不然整体速度提升有限。
5.2 量化后的精度损失实测:BLEU 掉了多少
我拿 500 条通用英文句子做了测试,对比 fp32 和 int8 的 BLEU。结果如下:
| 模型版本 | BLEU | 相对 fp32 | 单条延迟(CPU) |
|---|---|---|---|
| PyTorch fp32 | 32.5 | 基准 | 420ms |
| ONNX fp32 | 32.5 | 0% | 180ms |
| ONNX int8 动态 | 31.8 | -2.2% | 95ms |
| ONNX int8 静态 | 32.1 | -1.2% | 88ms |
可以看到,ONNX fp32 相比 PyTorch 精度无损,速度提升 2.3 倍;int8 动态量化精度掉了 2.2%,速度再提升近一倍;静态量化精度损失更小,但需要校准数据。
2.2% 的 BLEU 损失在通用翻译场景下基本感知不到,但如果你的场景对精度敏感(比如法律、医疗文本),建议用静态量化,或者干脆不量化,用 fp32 的 ONNX。
5.3 量化踩坑:哪些层不能量化
不是所有层都适合量化。MarianMT 里的 LayerNorm、Softmax 这些层对量化比较敏感,量化后容易出问题。ONNX Runtime 的量化工具默认会跳过一些层,但有时候需要手动指定。
如果量化后发现输出乱码或者重复,大概率是某些层量化坏了。可以用opset的QuantizeLinear和DequantizeLinear节点排查,或者用onnxruntime.quantization的extra_options参数排除特定层。
提示:量化后一定要重新跑一遍数值一致性验证和翻译质量测试,别直接上线。我见过有人量化完没测,上线后翻译结果全是重复词,排查了半天才发现是量化问题。
6. 部署上线时那些文档不会告诉你的细节
6.1 分词器必须和模型一起打包,别只拷 ONNX 文件
ONNX 文件只是计算图,分词器是独立的。部署时要把tokenizer.json、vocab.json、source.spm、target.spm这些文件一起打包。MarianMT 用的是 SentencePiece,source.spm和target.spm分别对应源语言和目标语言的分词模型,缺一不可。
我踩过的坑是:本地测试时 tokenizer 从 HuggingFace 缓存加载,没问题;部署到服务器后缓存没有,tokenizer 加载失败,翻译直接报错。后来改成把 tokenizer 文件一起打进镜像才解决。
6.2 线程数和 batch size 的调优:不是越大越好
ONNX Runtime 的intra_op_num_threads和inter_op_num_threads两个参数对性能影响很大。默认情况下它会用满所有 CPU 核心,但在容器环境里,如果没限制 CPU,它可能开太多线程导致上下文切换开销。
我的经验是:intra_op_num_threads设成物理核心数,inter_op_num_threads设成 1 到 2。batch size 方面,英译中模型在 CPU 上 batch size 设 4 到 8 比较合适,再大收益递减,而且延迟会增加。
sess_options = ort.SessionOptions() sess_options.intra_op_num_threads = 4 sess_options.inter_op_num_threads = 1 sess_options.graph_optimization_level = ort.GraphOptimizationLevel.ORT_ENABLE_ALL sess = ort.InferenceSession("encoder_model.onnx", sess_options)6.3 长文本翻译的截断策略:别让模型处理超长输入
MarianMT 的最大输入长度通常是 512 个 token,超过会被截断。如果你的输入是长段落,直接喂进去会被截断,翻译结果不完整。正确做法是先分句,再逐句翻译,最后拼接。
分句可以用简单的标点切分,也可以用nltk或者spaCy的句子分割器。分句后要注意:每句翻译完要保留标点和空格,拼接时才不会粘连。
import re def split_sentences(text): # 简单按句号、问号、感叹号切分 sentences = re.split(r'(?<=[.!?])\s+', text) return [s for s in sentences if s.strip()] def translate_long_text(text, translate_func): sentences = split_sentences(text) return " ".join(translate_func(s) for s in sentences)这个策略看着简单,但实际效果比直接截断好很多。尤其是翻译文档时,分句翻译能保证每句都完整。
6.4 监控指标:延迟、吞吐、OOM 一个都不能少
上线后要监控三个指标:P50/P99 延迟、QPS、内存占用。ONNX Runtime 本身不提供监控,需要你在外面包一层。我一般用 Prometheus 的 Python client 打点,记录每次推理的耗时和输入长度。
内存占用要特别关注,因为 ONNX Runtime 在加载模型时会预分配内存,如果模型多、并发高,容易 OOM。可以用sess_options.enable_mem_pattern = False关掉内存模式优化,减少内存峰值,代价是稍微慢一点。
7. 迁移过程中最值得记录的四个坑
7.1 坑一:opset 版本选错导致导出失败
最开始我用 opset 17 导出,结果 MarianMT 里的某个算子不支持,报错Unsupported operator。降到 14 就好了。后来查文档发现,不同 opset 版本支持的算子集不一样,seq2seq 模型涉及的一些注意力算子在高版本 opset 里行为有变化。建议从 14 开始试,不行再调。
7.2 坑二:dynamic_axes 没设全导致推理报错
导出时只设了 batch 维度动态,没设 sequence 维度,结果推理时输入长度一变就报维度不匹配。这个坑很隐蔽,因为导出时用的 dummy input 长度固定,不报错;只有推理时换长度才暴露。解决办法是把所有可能变化的维度都标成动态。
7.3 坑三:量化后输出重复词
int8 量化后,模型输出变成"的的的的的"这种重复。排查发现是 decoder 的某些层量化后数值溢出。解决办法是排除这些层,或者改用静态量化。这个坑让我意识到,量化不是无脑操作,必须做质量验证。
7.4 坑四:多线程下 ONNX Runtime 崩溃
在高并发下,ONNX Runtime 偶尔会 segfault。查了很久发现是SessionOptions被多个线程共享导致的。解决办法是每个线程创建独立的InferenceSession,或者用线程池限制并发数。ONNX Runtime 的 session 本身是线程安全的,但SessionOptions不是,别复用。
8. 迁移完成后的性能对比与后续优化方向
8.1 实测数据:ONNX 到底带来了多少提升
在同一台 4 核 8G 的 CPU 服务器上,用 1000 条英文句子做测试,结果如下:
| 方案 | 平均延迟 | P99 延迟 | QPS | 内存占用 |
|---|---|---|---|---|
| PyTorch fp32 | 420ms | 850ms | 2.3 | 1.8GB |
| ONNX fp32 | 180ms | 380ms | 5.5 | 900MB |
| ONNX int8 | 95ms | 210ms | 10.2 | 600MB |
ONNX fp32 相比 PyTorch,QPS 提升 2.4 倍,内存减半;int8 量化后 QPS 再翻倍,内存进一步降低。这个提升对于线上服务来说非常可观,直接省掉了扩容成本。
8.2 还能怎么优化:从模型蒸馏到硬件加速
如果还想进一步优化,有几个方向。一是模型蒸馏,用大模型教小模型,把参数量压下来。二是用 ONNX Runtime 的 TensorRT 或者 OpenVINO 后端,在特定硬件上能再快一截。三是把 encoder 和 decoder 合并成一个图,减少 session 切换开销,但实现难度大。
我个人觉得,对于大多数英译中场景,ONNX int8 已经够用了。再往下优化,投入产出比就不高了。除非你的 QPS 要求特别高,否则没必要折腾。
8.3 一个容易忽略的点:模型版本管理
迁移完成后,你会有多个版本的模型文件:PyTorch 原版、ONNX fp32、ONNX int8。这些文件要管理好,别搞混了。我建议用目录区分,每个版本带一个metadata.json记录导出参数、量化方式、测试指标。
{ "model_name": "opus-mt-en-zh", "version": "onnx-int8-v1", "opset": 14, "quantization": "dynamic", "bleu": 31.8, "latency_p99_ms": 210, "export_date": "2024-01-15" }这样后面出问题回溯时,能快速定位是哪个版本、什么配置。
整个迁移做下来,最大的感受是:ONNX 迁移不是"导出就完事",而是一个涉及模型导出、解码重写、量化验证、部署调优的系统工程。每一步都有坑,但每一步的收益也很实在。如果你也在被 PyTorch 推理的性能问题困扰,不妨试试这条路,踩完坑之后的性能提升会让你觉得值。