ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

英译中模型从PyTorch迁移到ONNX:导出、量化与部署实战

英译中模型从PyTorch迁移到ONNX:导出、量化与部署实战 英译中模型从 HuggingFace 的 PyTorch 权重迁移到 ONNX这件事看起来只是调一个torch.onnx.export但真正做过的人都知道坑几乎全在导出之后。我前后拿过三个不同规模的翻译模型做迁移从 60M 参数的小模型到 400M 左右的中等模型都跑过一遍最深的体会是导出成功只是起点能不能在目标推理引擎里跑出正确结果、跑出可接受的延迟才是这件事的真正难点。这篇内容就是把这套流程从头到尾拆开讲清楚包括环境准备、导出脚本怎么写、动态轴怎么配、量化怎么做、结果怎么验证以及我在这个过程中踩过的那些坑。适合已经用过 HuggingFace 的transformers、想把模型搬到 ONNX Runtime 或者端侧推理框架上的人参考也适合做部署但被 PyTorch 依赖拖累的工程师。1. 为什么要把英译中模型从 PyTorch 搬到 ONNX1.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 的英译中模型batch1、beam1 的情况下ONNX Runtime 比 PyTorch eager 快了大约 2.2 倍这个差距在长文本翻译上会更明显。还有一层是工程解耦。模型一旦转成 ONNX部署侧就不需要关心训练框架的版本、不需要装transformers、不需要 Python 环境C、C#、Java、Rust 都能直接调。对于团队分工来说算法同学负责导出 ONNX工程同学负责集成边界非常清晰。1.2 英译中模型迁移的特殊性在哪翻译模型和普通分类模型不一样它是 encoder-decoder 结构而且推理是自回归的。这意味着导出的时候不能只导一个 forward得考虑清楚到底导什么。常见的有三种粒度只导 encoderdecoder 用别的方案这种很少见导 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 就好了。所以第一步是把版本锁死别用最新版用经过验证的组合。我目前稳定在用的组合是组件版本说明Python3.103.11 部分算子导出有兼容问题torch2.1.2对 dynamic axes 支持比较完善transformers4.36.2与 torch 2.1 匹配良好onnx1.15.0opset 17 支持完整onnxruntime1.17.0支持 opset 17onnxsim0.4.36图简化用安装命令很直接pip install torch2.1.2 transformers4.36.2 onnx1.15.0 onnxruntime1.17.0 onnxsim0.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_implementationeager。模型规模上英译中场景我建议控制在 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_tensorspt) with torch.no_grad(): output_ids model.generate(**inputs, max_new_tokens128, num_beams1) text tokenizer.decode(output_ids[0], skip_special_tokensTrue) baseline.append({src: sent, tgt: text}) with open(baseline.json, w, encodingutf-8) as f: json.dump(baseline, f, ensure_asciiFalse, indent2)这份 baseline 是后面判断 ONNX 结果对不对的唯一依据。别指望肉眼比对翻译结果差一个词都可能是 bug。3. 导出脚本的核心逻辑与动态轴配置3.1 导出 encoder 和 decoder 要分开做英译中模型的导出我建议拆成两个 ONNX 文件encoder.onnx和decoder.onnx。原因是 encoder 只需要跑一次输入是源语言 token输出是 encoder hidden statesdecoder 要跑 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_implementationeager) model.eval() # 构造 dummy input dummy_src tokenizer(This is a test sentence., return_tensorspt) 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_version17, do_constant_foldingTrue, )这里的关键是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]], dtypetorch.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[fpast_key_{i}] {0: batch, 2: past_len} dynamic_axes[fpast_value_{i}] {0: batch, 2: past_len} dynamic_axes[fpresent_key_{i}] {0: batch, 2: total_len} dynamic_axes[fpresent_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_tokens128): inputs tokenizer(sentence, return_tensorsnp) 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]], dtypenp.int64) past_kv { fpast_key_{i}: np.zeros((1, num_heads, 0, head_dim), dtypenp.float32) for i in range(num_layers) } past_kv.update({ fpast_value_{i}: np.zeros((1, num_heads, 0, head_dim), dtypenp.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]], dtypenp.int64) # 更新 past_kv for i in range(num_layers): past_kv[fpast_key_{i}] outputs[1 i * 2] past_kv[fpast_value_{i}] outputs[2 i * 2] return tokenizer.decode(generated, skip_special_tokensTrue)跑完对比 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 默认用 float32ONNX 导出时如果某些算子被降到了 float16会有精度损失。可以在导出时加do_constant_foldingFalse排除常量折叠的影响或者用onnxruntime的float32provider 验证。4.3 导出报错的常见类型导出阶段的报错主要有几类错误信息原因解决Unsupported operator: XXX算子不在目标 opset 里提高 opset 版本或替换算子实现Dynamic shape inference failed动态维度推导失败用 onnxsim 简化或手动指定 shapeTracerWarning: Converting a tensor to a Python boolean图里有数据依赖的控制流改写模型代码去掉 if tensor 判断RuntimeError: expected scalar type输入类型不匹配确保所有输入都是 int64 或 float32TracerWarning是最容易被忽略的因为它只是警告不是错误但 trace 出来的图可能是错的。看到这个警告一定要停下来检查通常是模型里有if x 0这种基于张量值的判断trace 时只会走一个分支。5. 量化与图优化让 ONNX 模型真正跑得快5.1 动态量化是最省事的方案ONNX Runtime 提供了动态量化不需要校准数据直接把权重从 float32 降到 int8激活值在运行时动态量化。对于英译中模型动态量化通常能把模型体积压到原来的 1/4CPU 推理速度提升 1.5 到 2 倍。from onnxruntime.quantization import quantize_dynamic, QuantType quantize_dynamic( encoder.onnx, encoder_int8.onnx, weight_typeQuantType.QInt8, ) quantize_dynamic( decoder.onnx, decoder_int8.onnx, weight_typeQuantType.QInt8, )动态量化的好处是简单坏处是精度损失相对大。我实测下来英译中模型动态量化后 BLEU 会掉 1 到 2 个点对于要求不高的场景可以接受但如果是质量敏感的场景得用静态量化。5.2 静态量化需要校准数据静态量化要把激活值也量化所以需要一批校准数据来统计激活值的分布。校准数据用训练集或者验证集里的英文句子就行准备 100 到 200 条覆盖不同长度的样本。from onnxruntime.quantization import quantize_static, CalibrationDataReader class TranslationCalibrationReader(CalibrationDataReader): def __init__(self, sentences, tokenizer, max_len128): self.data [] for sent in sentences: inputs tokenizer(sent, return_tensorsnp, max_lengthmax_len, truncationTrue) 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_typeQuantType.QInt8, activation_typeQuantType.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.6GBint8 量化后大约 400MB还是偏大。这时候可以考虑几个方向一是把 encoder 和 decoder 分别量化、分别加载用的时候按需加载。二是用更激进的量化比如 int4但 int4 在 ONNX Runtime 里的支持还不完善需要自己写量化算子。三是对模型做剪枝再导出把一些不重要的注意力头去掉。加载的时候要注意内存峰值。ONNX Runtime 加载模型时会把权重读进内存如果模型是 400MB加载峰值可能到 800MB。在内存受限的设备上要用session_options里的enable_mem_patternFalse关掉内存池虽然会慢一点但峰值内存更低。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 处理端侧通常 batch1但如果要支持 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 tuple4.36 之后引入了Cache类默认返回DynamicCache对象。导出时如果用新版torch.onnx.export会把DynamicCache当成一个不透明的对象导出的图里没有 KV Cache 的输入输出。解决办法是在导出脚本里显式地把DynamicCache转成 tuplefrom 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_cacheTrue并手动管理 cache绕开DynamicCache。7.3 量化后的精度验证不能只看 BLEUBLEU 是个宏观指标量化后 BLEU 掉 1 个点可能意味着某些句子完全翻译错了只是被平均掉了。我建议除了 BLEU还要做逐句对比把量化前后的翻译结果并排看重点关注数字、专有名词、否定句这些容易出错的地方。我遇到过一次量化后 BLEU 只掉了 0.8但所有包含数字的句子都翻译错了原因是数字在词表里的 token 分布比较稀疏量化时被压到了同一个 bin 里。这种问题 BLEU 反映不出来但实际影响很大。7.4 别忽略 warmupONNX 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 迁移不是什么高深技术但细节非常多任何一个环节疏忽都可能导致结果不对或者性能不达标。把验证做扎实比追求极致的量化压缩更重要。
返回列表