英译中模型从 HuggingFace 的 PyTorch 权重迁移到 ONNX,这件事看起来只是调一个torch.onnx.export,但真正做过的人都知道,坑几乎全在导出之后。我前后拿过三个不同规模的翻译模型做迁移,从 60M 参数的小模型到 400M 左右的中等模型都跑过一遍,最深的体会是:导出成功只是起点,能不能在目标推理引擎里跑出正确结果、跑出可接受的延迟,才是这件事的真正难点。这篇内容就是把这套流程从头到尾拆开讲清楚,包括环境准备、导出脚本怎么写、动态轴怎么配、量化怎么做、结果怎么验证,以及我在这个过程中踩过的那些坑。适合已经用过 HuggingFace 的transformers、想把模型搬到 ONNX Runtime 或者端侧推理框架上的人参考,也适合做部署但被 PyTorch 依赖拖累的工程师。
1. 为什么要把英译中模型从 PyTorch 搬到 ONNX
1.1 迁移的动机不是"跟风",而是部署约束
先说清楚为什么要做这件事。HuggingFace 上的英译中模型绝大多数是 PyTorch 格式,用transformers加载、model.generate()推理,本地跑 demo 完全没问题。但一旦要上生产,问题就来了:PyTorch 的运行时体积大、依赖多、启动慢,在服务端还好,到了边缘设备或者需要嵌入到非 Python 环境里就非常难受。ONNX 的价值在于它是一份与框架无关的计算图描述,配合 ONNX Runtime 可以在 CPU、GPU、甚至一些专用加速器上跑,而且运行时体积比完整 PyTorch 小一个量级。
另一个现实动机是推理性能。PyTorch 的 eager 模式在推理时会有大量 Python 层的调度开销,尤其是自回归生成这种逐 token 循环的场景,每个 step 都要走一遍 Python 逻辑。ONNX Runtime 对计算图做了算子融合、常量折叠、内存复用等优化,在 CPU 上做序列生成时经常能拿到 1.5 到 3 倍的加速。我实测过一个 6 层 encoder-decoder 的英译中模型,batch=1、beam=1 的情况下,ONNX Runtime 比 PyTorch eager 快了大约 2.2 倍,这个差距在长文本翻译上会更明显。
还有一层是工程解耦。模型一旦转成 ONNX,部署侧就不需要关心训练框架的版本、不需要装transformers、不需要 Python 环境,C++、C#、Java、Rust 都能直接调。对于团队分工来说,算法同学负责导出 ONNX,工程同学负责集成,边界非常清晰。
1.2 英译中模型迁移的特殊性在哪
翻译模型和普通分类模型不一样,它是 encoder-decoder 结构,而且推理是自回归的。这意味着导出的时候不能只导一个 forward,得考虑清楚到底导什么。常见的有三种粒度:
- 只导 encoder,decoder 用别的方案,这种很少见;
- 导 encoder 和 decoder 两个独立图,decoder 接收 encoder 的输出和已生成的 token;
- 导一个带 KV Cache 的 decoder,把历史 key/value 作为输入输出在 step 之间传递。
第三种是生产环境最常用的,因为自回归生成时如果不缓存 KV,每一步都要把前面所有 token 重新算一遍注意力,复杂度是 O(n²),长句翻译会慢到无法接受。但带 KV Cache 的导出也是最麻烦的,因为 cache 的形状是动态的,而且要在图里做拼接,对 ONNX 的动态维度支持要求很高。
另外英译中还有个特点:词表通常很大。中英翻译模型的词表动辄 5 万到 25 万,输出层的 logits 张量在长序列上会非常占内存。导出的时候如果不注意,ONNX 模型文件可能比原始 PyTorch 权重还大,因为 PyTorch 的权重是共享的,而 ONNX 里如果处理不当会把 embedding 和输出投影各存一份。
2. 导出前的环境准备与模型选型
2.1 版本组合是第一个坑
ONNX 导出对版本非常敏感。我踩过最典型的一次是torch2.0 配onnx1.12,导出带 KV Cache 的模型时直接报算子不支持,换到onnx1.14 就好了。所以第一步是把版本锁死,别用最新版,用经过验证的组合。
我目前稳定在用的组合是:
| 组件 | 版本 | 说明 |
|---|---|---|
| Python | 3.10 | 3.11 部分算子导出有兼容问题 |
| torch | 2.1.2 | 对 dynamic axes 支持比较完善 |
| transformers | 4.36.2 | 与 torch 2.1 匹配良好 |
| onnx | 1.15.0 | opset 17 支持完整 |
| onnxruntime | 1.17.0 | 支持 opset 17 |
| onnxsim | 0.4.36 | 图简化用 |
安装命令很直接:
pip install torch==2.1.2 transformers==4.36.2 onnx==1.15.0 onnxruntime==1.17.0 onnxsim==0.4.36注意:不要在同一环境里混装多个 torch 版本,ONNX 导出会调用 torch 的 JIT trace,版本冲突时 trace 出来的图可能是错的,而且不报错,只是结果不对,非常难查。
2.2 模型选型要考虑导出友好度
不是所有 HuggingFace 上的英译中模型都好导。我建议优先选结构标准的模型,比如基于标准 Transformer 的MarianMT、T5、BART系列。这些模型的注意力实现是标准的,导出时不会遇到奇怪的算子。
要避开的是那些用了自定义 CUDA kernel 或者自定义 attention 实现的模型,比如某些用了 flash attention 变体的版本。这些在导出时要么算子不支持,要么 trace 出来的图是错的。如果非要用,得先把 attention 实现切回标准的 eager 实现,在from_pretrained时加attn_implementation="eager"。
模型规模上,英译中场景我建议控制在 400M 参数以内。再大的模型导出后 ONNX 文件会超过 1.5GB,加载慢、内存占用高,而且量化后精度损失也更难控制。如果确实需要大模型,考虑先做蒸馏再导出。
2.3 先跑通 PyTorch 基线再动手
这一步很多人会跳过,但我觉得是必须的。在导出之前,先用 PyTorch 跑几条测试样本,把输入输出存下来,作为后面验证 ONNX 结果的基准。具体做法是准备 5 到 10 条覆盖不同长度的英文句子,从短句到 50 词以上的长句都要有,然后用model.generate()生成翻译,把结果存成 JSON。
import json import torch from transformers import AutoTokenizer, AutoModelForSeq2SeqLM model_name = "Helsinki-NLP/opus-mt-en-zh" tokenizer = AutoTokenizer.from_pretrained(model_name) model = AutoModelForSeq2SeqLM.from_pretrained(model_name) model.eval() test_sentences = [ "Hello, how are you?", "The quick brown fox jumps over the lazy dog.", # ... 更多测试句 ] baseline = [] for sent in test_sentences: inputs = tokenizer(sent, return_tensors="pt") with torch.no_grad(): output_ids = model.generate(**inputs, max_new_tokens=128, num_beams=1) text = tokenizer.decode(output_ids[0], skip_special_tokens=True) baseline.append({"src": sent, "tgt": text}) with open("baseline.json", "w", encoding="utf-8") as f: json.dump(baseline, f, ensure_ascii=False, indent=2)这份 baseline 是后面判断 ONNX 结果对不对的唯一依据。别指望肉眼比对,翻译结果差一个词都可能是 bug。
3. 导出脚本的核心逻辑与动态轴配置
3.1 导出 encoder 和 decoder 要分开做
英译中模型的导出我建议拆成两个 ONNX 文件:encoder.onnx和decoder.onnx。原因是 encoder 只需要跑一次,输入是源语言 token,输出是 encoder hidden states;decoder 要跑 N 次,每次接收上一步的输出和 KV Cache。拆开之后,encoder 的图可以充分优化,decoder 的图可以针对单步推理做特化。
先看 encoder 的导出:
import torch from transformers import AutoTokenizer, AutoModelForSeq2SeqLM model_name = "Helsinki-NLP/opus-mt-en-zh" tokenizer = AutoTokenizer.from_pretrained(model_name) model = AutoModelForSeq2SeqLM.from_pretrained(model_name, attn_implementation="eager") model.eval() # 构造 dummy input dummy_src = tokenizer("This is a test sentence.", return_tensors="pt") input_ids = dummy_src["input_ids"] attention_mask = dummy_src["attention_mask"] # 只取 encoder encoder = model.get_encoder() torch.onnx.export( encoder, (input_ids, attention_mask), "encoder.onnx", input_names=["input_ids", "attention_mask"], output_names=["encoder_hidden_states"], dynamic_axes={ "input_ids": {0: "batch", 1: "src_len"}, "attention_mask": {0: "batch", 1: "src_len"}, "encoder_hidden_states": {0: "batch", 1: "src_len"}, }, opset_version=17, do_constant_folding=True, )这里的关键是dynamic_axes。英译中的源句长度是不固定的,如果导出时把src_len固定成 dummy input 的长度,那模型就只能处理这一个长度的输入,完全没法用。所以batch和src_len两个维度都必须标成动态。
3.2 decoder 的 KV Cache 导出是难点
decoder 的导出要复杂得多。标准的generate内部会维护一个past_key_values,每一步把新的 key/value 拼到历史里。导出的时候我们要把这个逻辑显式地写出来,让 ONNX 图接收past_key_values作为输入,输出新的past_key_values。
HuggingFace 的模型在forward里已经支持past_key_values参数,所以可以直接用。但要注意,不同版本的transformers里past_key_values的格式不一样,早期是 tuple of tuple,新版是DynamicCache对象。导出时最好用 tuple 格式,因为 ONNX 对自定义对象支持不好。
# 构造 decoder 的 dummy 输入 batch_size = 1 decoder_input_ids = torch.tensor([[tokenizer.pad_token_id]], dtype=torch.long) encoder_hidden_states = torch.randn(batch_size, 10, model.config.d_model) # 构造空的 past_key_values num_layers = model.config.decoder_layers num_heads = model.config.decoder_attention_heads head_dim = model.config.d_model // num_heads past_key_values = tuple( ( torch.zeros(batch_size, num_heads, 0, head_dim), torch.zeros(batch_size, num_heads, 0, head_dim), ) for _ in range(num_layers) ) decoder = model.get_decoder() # 注意:decoder 单独拿出来时,需要把 lm_head 也带上,或者单独导出 lm_head这里有个细节:get_decoder()拿到的只是 decoder 主体,输出的是 hidden states,还需要经过lm_head投影到词表维度。有两种做法,一是把lm_head拼到 decoder 里一起导出,二是单独导出lm_head。我倾向于拼在一起,因为lm_head就是一个线性层,拼进去不增加复杂度,还能省一次中间张量的传输。
3.3 动态轴的完整配置
decoder 的动态轴比 encoder 复杂,因为涉及 KV Cache 的序列长度维度。完整的配置是这样的:
dynamic_axes = { "decoder_input_ids": {0: "batch", 1: "dec_len"}, "encoder_hidden_states": {0: "batch", 1: "src_len"}, "logits": {0: "batch", 1: "dec_len"}, } # 每个 past_key 和 present_key 都要加动态轴 for i in range(num_layers): dynamic_axes[f"past_key_{i}"] = {0: "batch", 2: "past_len"} dynamic_axes[f"past_value_{i}"] = {0: "batch", 2: "past_len"} dynamic_axes[f"present_key_{i}"] = {0: "batch", 2: "total_len"} dynamic_axes[f"present_value_{i}"] = {0: "batch", 2: "total_len"}past_len和total_len都是动态的,前者是历史长度,后者是历史加当前步的长度。ONNX Runtime 在推理时会根据实际输入推断这些维度,所以不需要预先指定。
提示:如果导出时报 "Dynamic shape not supported" 之类的错误,八成是某个中间算子的动态维度推导失败了。这时候可以用
onnxsim先简化一遍图,很多动态维度问题会被自动修掉。
4. 导出后的验证与常见错误排查
4.1 用 onnxruntime 跑一遍对比 baseline
导出完成后的第一件事不是优化,是验证正确性。用onnxruntime加载两个 ONNX 文件,手动实现一遍自回归生成,然后和 baseline 对比。
import numpy as np import onnxruntime as ort enc_session = ort.InferenceSession("encoder.onnx", providers=["CPUExecutionProvider"]) dec_session = ort.InferenceSession("decoder.onnx", providers=["CPUExecutionProvider"]) def translate(sentence, max_new_tokens=128): inputs = tokenizer(sentence, return_tensors="np") input_ids = inputs["input_ids"].astype(np.int64) attention_mask = inputs["attention_mask"].astype(np.int64) encoder_hidden = enc_session.run( ["encoder_hidden_states"], {"input_ids": input_ids, "attention_mask": attention_mask}, )[0] # 初始化 decoder_input = np.array([[model.config.decoder_start_token_id]], dtype=np.int64) past_kv = { f"past_key_{i}": np.zeros((1, num_heads, 0, head_dim), dtype=np.float32) for i in range(num_layers) } past_kv.update({ f"past_value_{i}": np.zeros((1, num_heads, 0, head_dim), dtype=np.float32) for i in range(num_layers) }) generated = [] for _ in range(max_new_tokens): feeds = { "decoder_input_ids": decoder_input, "encoder_hidden_states": encoder_hidden, **past_kv, } outputs = dec_session.run(None, feeds) logits = outputs[0] next_token = int(np.argmax(logits[0, -1, :])) if next_token == model.config.eos_token_id: break generated.append(next_token) decoder_input = np.array([[next_token]], dtype=np.int64) # 更新 past_kv for i in range(num_layers): past_kv[f"past_key_{i}"] = outputs[1 + i * 2] past_kv[f"past_value_{i}"] = outputs[2 + i * 2] return tokenizer.decode(generated, skip_special_tokens=True)跑完对比 baseline,如果结果完全一致,说明导出是成功的。如果结果不一致,往下看排查思路。
4.2 结果不一致的排查链路
结果不对是最常见的问题,而且原因很多。我总结了一个排查顺序,从简单到复杂:
第一步,检查 attention_mask 的处理。英译中模型的 encoder 对 padding 位置要做 mask,如果导出时 mask 没正确传递,encoder 输出会包含 padding 位置的噪声,导致翻译结果偏移。验证方法是把 batch 固定为 1、不做 padding,看结果是否正常。如果单条正常、batch 不正常,就是 mask 的问题。
第二步,检查 KV Cache 的拼接顺序。有些模型的 KV Cache 是(key, value)的顺序,有些是(value, key),导出时如果搞反了,结果会完全乱掉。这个可以通过打印 PyTorch 里past_key_values的结构来确认。
第三步,检查位置编码。自回归生成时,每一步的位置编码要基于当前的总长度,而不是当前步的长度。如果导出时位置编码被固定成了 dummy input 的长度,生成到后面就会出错。这个问题的表现是短句正常、长句乱码。
第四步,检查数值精度。PyTorch 默认用 float32,ONNX 导出时如果某些算子被降到了 float16,会有精度损失。可以在导出时加do_constant_folding=False排除常量折叠的影响,或者用onnxruntime的float32provider 验证。
4.3 导出报错的常见类型
导出阶段的报错主要有几类:
| 错误信息 | 原因 | 解决 |
|---|---|---|
| Unsupported operator: XXX | 算子不在目标 opset 里 | 提高 opset 版本,或替换算子实现 |
| Dynamic shape inference failed | 动态维度推导失败 | 用 onnxsim 简化,或手动指定 shape |
| TracerWarning: Converting a tensor to a Python boolean | 图里有数据依赖的控制流 | 改写模型代码,去掉 if tensor 判断 |
| RuntimeError: expected scalar type | 输入类型不匹配 | 确保所有输入都是 int64 或 float32 |
TracerWarning是最容易被忽略的,因为它只是警告不是错误,但 trace 出来的图可能是错的。看到这个警告一定要停下来检查,通常是模型里有if x > 0这种基于张量值的判断,trace 时只会走一个分支。
5. 量化与图优化:让 ONNX 模型真正跑得快
5.1 动态量化是最省事的方案
ONNX Runtime 提供了动态量化,不需要校准数据,直接把权重从 float32 降到 int8,激活值在运行时动态量化。对于英译中模型,动态量化通常能把模型体积压到原来的 1/4,CPU 推理速度提升 1.5 到 2 倍。
from onnxruntime.quantization import quantize_dynamic, QuantType quantize_dynamic( "encoder.onnx", "encoder_int8.onnx", weight_type=QuantType.QInt8, ) quantize_dynamic( "decoder.onnx", "decoder_int8.onnx", weight_type=QuantType.QInt8, )动态量化的好处是简单,坏处是精度损失相对大。我实测下来,英译中模型动态量化后 BLEU 会掉 1 到 2 个点,对于要求不高的场景可以接受,但如果是质量敏感的场景,得用静态量化。
5.2 静态量化需要校准数据
静态量化要把激活值也量化,所以需要一批校准数据来统计激活值的分布。校准数据用训练集或者验证集里的英文句子就行,准备 100 到 200 条覆盖不同长度的样本。
from onnxruntime.quantization import quantize_static, CalibrationDataReader class TranslationCalibrationReader(CalibrationDataReader): def __init__(self, sentences, tokenizer, max_len=128): self.data = [] for sent in sentences: inputs = tokenizer(sent, return_tensors="np", max_length=max_len, truncation=True) self.data.append({ "input_ids": inputs["input_ids"].astype(np.int64), "attention_mask": inputs["attention_mask"].astype(np.int64), }) self.idx = 0 def get_next(self): if self.idx >= len(self.data): return None item = self.data[self.idx] self.idx += 1 return item reader = TranslationCalibrationReader(calib_sentences, tokenizer) quantize_static( "encoder.onnx", "encoder_int8_static.onnx", reader, weight_type=QuantType.QInt8, activation_type=QuantType.QUInt8, )静态量化的精度通常比动态量化好,但校准数据的分布要和实际推理数据接近,否则量化误差会很大。我建议校准数据里至少包含 20% 的长句,因为长句的激活值分布和短句差别很大。
5.3 图优化能再挤出一部分性能
ONNX Runtime 在加载模型时会自动做图优化,但有些优化需要手动开启。可以在SessionOptions里设置优化级别:
import onnxruntime as ort so = ort.SessionOptions() so.graph_optimization_level = ort.GraphOptimizationLevel.ORT_ENABLE_ALL so.intra_op_num_threads = 4 so.inter_op_num_threads = 2 session = ort.InferenceSession("encoder_int8.onnx", so, providers=["CPUExecutionProvider"])ORT_ENABLE_ALL会开启所有优化,包括算子融合、常量折叠、冗余节点消除。intra_op_num_threads控制单个算子内部的并行度,inter_op_num_threads控制算子之间的并行度。对于 encoder 这种计算密集的图,intra_op_num_threads设成物理核心数比较合适;对于 decoder 这种小算子多的图,inter_op_num_threads更重要。
另外onnxsim可以在导出后做一次图简化,把一些冗余的 reshape、transpose 消掉:
python -m onnxsim encoder.onnx encoder_sim.onnx我实测下来,onnxsim对 encoder 的简化效果明显,模型体积能小 5% 到 10%,推理速度提升 3% 到 8%。对 decoder 效果一般,因为 decoder 的图本来就比较紧凑。
6. 端侧部署时的额外考量
6.1 模型文件的分割与加载
如果目标平台是移动端或者嵌入式设备,ONNX 模型文件大小是个硬约束。一个 400M 参数的英译中模型,float32 导出后大约 1.6GB,int8 量化后大约 400MB,还是偏大。这时候可以考虑几个方向:
一是把 encoder 和 decoder 分别量化、分别加载,用的时候按需加载。二是用更激进的量化,比如 int4,但 int4 在 ONNX Runtime 里的支持还不完善,需要自己写量化算子。三是对模型做剪枝再导出,把一些不重要的注意力头去掉。
加载的时候要注意内存峰值。ONNX Runtime 加载模型时会把权重读进内存,如果模型是 400MB,加载峰值可能到 800MB。在内存受限的设备上,要用session_options里的enable_mem_pattern=False关掉内存池,虽然会慢一点,但峰值内存更低。
6.2 输入输出的预处理要对齐
端侧部署时,tokenizer 往往不能用 Python 版本,需要用 C++ 或者其他语言重新实现。这时候要特别注意 tokenizer 的细节:BPE 的合并规则、特殊 token 的处理、padding 的方向。我见过最典型的问题是 padding 方向搞反了,PyTorch 里是右 padding,端侧实现成了左 padding,结果翻译质量断崖式下降。
建议的做法是把 tokenizer 的配置导出成 JSON,包括词表、合并规则、特殊 token id,端侧按这份配置实现。然后用同一批测试句子,对比 Python tokenizer 和端侧 tokenizer 的输出 id 序列,确保完全一致。
6.3 自回归循环的控制逻辑
端侧实现自回归生成时,循环控制逻辑要自己写。这里有几个容易出问题的地方:
- 终止条件:除了 EOS token,还要考虑最大长度限制。如果模型一直不输出 EOS,循环会一直跑下去。
- KV Cache 的内存管理:每一步的 KV Cache 都会增长,如果不做限制,长文本翻译时内存会爆。可以设置一个最大缓存长度,超过就截断。
- batch 处理:端侧通常 batch=1,但如果要支持 batch,要注意不同样本的生成长度不一样,需要做 padding 和 mask。
我在一个嵌入式项目里遇到过 KV Cache 内存泄漏的问题,原因是每一步都新建了一个数组来存 cache,没有复用。改成预分配一个最大长度的 buffer,用索引来管理有效长度,内存占用就稳定了。
7. 我踩过的几个真实坑与经验总结
7.1 导出时的 dummy input 长度会影响图结构
这是我踩过最隐蔽的一个坑。导出 encoder 时,如果 dummy input 的长度是 10,导出的图里某些 reshape 操作的 shape 会被固定成和 10 相关的值,虽然后面用动态轴覆盖了,但中间某些算子可能还是带着固定维度。表现是短句正常,超过某个长度就报 shape mismatch。
解决办法是导出时用两个不同长度的 dummy input 各导一次,对比图结构。如果图结构不一样,说明有隐藏的固定维度。更稳妥的做法是用torch.onnx.export的dynamic_axes把所有可能变化的维度都标出来,包括中间张量的维度。
7.2 不同 transformers 版本的 past_key_values 格式不兼容
transformers4.36 之前,past_key_values是 tuple of tuple;4.36 之后引入了Cache类,默认返回DynamicCache对象。导出时如果用新版,torch.onnx.export会把DynamicCache当成一个不透明的对象,导出的图里没有 KV Cache 的输入输出。
解决办法是在导出脚本里显式地把DynamicCache转成 tuple:
from transformers.cache_utils import DynamicCache # 如果模型返回 DynamicCache,转成 tuple if isinstance(outputs.past_key_values, DynamicCache): past_kv = tuple( (layer.keys, layer.values) for layer in outputs.past_key_values.layers )或者在加载模型时设置use_cache=True并手动管理 cache,绕开DynamicCache。
7.3 量化后的精度验证不能只看 BLEU
BLEU 是个宏观指标,量化后 BLEU 掉 1 个点,可能意味着某些句子完全翻译错了,只是被平均掉了。我建议除了 BLEU,还要做逐句对比,把量化前后的翻译结果并排看,重点关注数字、专有名词、否定句这些容易出错的地方。
我遇到过一次量化后 BLEU 只掉了 0.8,但所有包含数字的句子都翻译错了,原因是数字在词表里的 token 分布比较稀疏,量化时被压到了同一个 bin 里。这种问题 BLEU 反映不出来,但实际影响很大。
7.4 别忽略 warmup
ONNX Runtime 第一次推理会做很多初始化工作,包括内存分配、算子编译,耗时可能是稳定状态的 10 倍以上。在生产环境里,如果不在启动时做 warmup,第一个请求的延迟会非常难看。warmup 的做法很简单,用几条典型输入跑几遍就行:
# warmup for _ in range(3): enc_session.run(None, {"input_ids": warmup_ids, "attention_mask": warmup_mask}) # decoder 也跑几遍warmup 的输入要覆盖不同的长度,短句、中句、长句各跑一遍,这样内存池能预先分配好合适的大小。
7.5 版本升级要重新验证
ONNX Runtime 和 onnx 的版本升级经常带来行为变化。我有一次把 onnxruntime 从 1.15 升到 1.17,同一个模型同一个输入,输出结果在小数点后第 5 位开始不一样了。虽然对最终翻译结果没影响,但如果你的系统里有基于数值的断言,就会挂掉。所以每次升级推理引擎版本,都要重新跑一遍验证流程,别假设向后兼容。
这套流程我前后在三个项目里跑过,从最初的磕磕绊绊到现在基本能一次导出成功,核心经验就是:导出前锁版本、导出时分清 encoder 和 decoder、导出后先验证再优化、量化后逐句检查。英译中模型的 ONNX 迁移不是什么高深技术,但细节非常多,任何一个环节疏忽都可能导致结果不对或者性能不达标。把验证做扎实,比追求极致的量化压缩更重要。