资讯动态

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

发布时间:2026/10/8 16:16:29 来源:尧图企业网站定制
英译中模型从 HuggingFace 的 PyTorch 权重迁移到 ONNX这件事看起来只是调一个torch.onnx.export但真正动手做过的人都知道坑几乎全在导出之后动态轴没设对导致变长输入报错、分词器和模型对不上、量化后 BLEU 掉得离谱、推理时显存不降反升。我自己前前后后把好几个翻译模型搬上 ONNX Runtime从桌面端到服务端都跑过一遍这篇就把整个链路拆开讲清楚包括每一步为什么这么做、参数怎么算、哪些地方最容易翻车。这篇内容适合三类人一是手里已经有 HuggingFace 上的英译中模型比如 Helsinki-NLP 的 opus-mt 系列、MarianMT、或者自己微调过的 T5/BART想把它部署到没有 PyTorch 依赖的环境里二是做端侧或嵌入式推理需要 ONNX 这种中间格式再转其他运行时三是单纯想搞清楚 PyTorch 到 ONNX 这条路上到底发生了什么。下面所有操作我都以 MarianMT 英译中模型为主线其他架构的差异我会单独点出来。1. 为什么英译中模型值得单独走一遍 ONNX 迁移1.1 翻译模型和普通分类模型的本质差异很多人第一次迁移 ONNX 是拿 BERT 做文本分类导出很顺于是以为翻译模型也一样。结果一上手就发现不对。分类模型输入是固定长度或者简单 padding 到固定长度输出是一个 logits 向量导出时把input_ids和attention_mask两个输入声明清楚就完事了。翻译模型完全不是这个逻辑。翻译是典型的序列到序列任务编码器读入源语言整句解码器自回归地一个 token 一个 token 往外吐。这意味着两件事第一输入长度是变的你不能假设所有句子都是 32 个 token第二解码过程是循环的每一步的输出要喂回下一步的输入。ONNX 本身是静态计算图它没有原生的循环直到遇到结束符这种控制流虽然有 Loop 算子但用起来极其别扭。所以工业界的通行做法是把解码器拆成单步解码的图外面用 Python 或者 C 写循环每一步调用一次 ONNX 图把上一步生成的 token 拼回去。这个差异直接决定了整个迁移方案的设计。你不能指望导出一个图就搞定全部推理必须把编码器和解码器分开导出解码器还要处理 KV Cache 的输入输出。这是英译中模型迁移 ONNX 最核心的认知先想明白这一点后面的参数才有意义。1.2 ONNX 在这条链路上到底解决了什么问题有人会问既然还是要写 Python 循环那用 ONNX 图什么直接用 PyTorch 推理不行吗。这里要分场景看。如果你的服务端本来就跑着 PyTorchGPU 也够那确实没必要折腾 ONNXPyTorch 的model.generate()已经帮你把 beam search、KV Cache 全处理好了。但如果你要部署到没有 CUDA 的机器、要嵌进一个 C 服务、要用 TensorRT 或者 OpenVINO 进一步加速、或者要把模型塞进手机和边缘设备那 PyTorch 的运行时依赖几个 G 的 libtorch就成了负担。ONNX 的价值在于它是一个与框架无关的中间表示导出一次可以在 ONNX Runtime、TensorRT、OpenVINO、NCNN 等多个后端上跑运行时体积小、启动快、跨平台。另外一个常被忽略的点是量化。ONNX 生态里的 INT8 量化工具链相对成熟onnxruntime.quantization提供了动态量化和静态量化两条路对于翻译这种对精度敏感的任务动态量化往往能在体积减半的同时把精度损失控制在可接受范围。这一点在 PyTorch 里做要麻烦得多。1.3 迁移前必须确认的三件事动手之前先做三个检查能省掉后面大量返工。第一确认模型架构。打开 HuggingFace 模型页看config.json里的model_type。MarianMT 是marianT5 是t5BART 是bartmBART 是mbart。不同架构的解码器结构差异很大T5 用的是相对位置编码且没有传统的 decoder cross-attention 的 KV 缓存布局BART 的 decoder 起始 token 是s而 Marian 是/s是的Marian 用结束符当起始符这个反直觉的设计坑过很多人。第二确认分词器类型。Marian 系列通常用 SentencePieceT5 也是 SentencePieceBART 用 GPT2 风格的 BPE。分词器不参与 ONNX 导出但推理时你要在 Python 侧调用它必须保证和训练时一致否则翻译结果会莫名其妙。第三确认目标运行时的算子支持。ONNX Runtime 对 Transformer 相关算子支持很好但如果你要转 TensorRT某些动态 shape 的算子可能需要额外处理。先跑通 ONNX Runtime再考虑其他后端。2. 导出前的环境准备与模型加载细节2.1 依赖版本的选择与踩坑环境这块版本兼容性是第一大坑。我踩过最惨的一次是transformers4.30 配torch2.0 导出 Marian导出的图在 ONNX Runtime 里跑出来全是乱码排查了两天才发现是transformers某个版本对 attention mask 的处理变了导致导出时 mask 的语义和推理时不一致。我的建议是锁定一套经过验证的组合。截至我最近一次实操比较稳的是torch2.1.x、transformers4.36.x、onnx1.15.x、onnxruntime1.17.x。安装命令如下pip install torch2.1.2 transformers4.36.2 onnx1.15.0 onnxruntime1.17.0 pip install sentencepiece protobufsentencepiece是 Marian 和 T5 分词器必需的protobuf是 ONNX 序列化的依赖。这两个经常被漏装然后报一个看不懂的错。注意不要盲目追最新版。ONNX 的算子集opset在不同版本间有行为变化transformers的导出逻辑也在持续调整。生产环境务必锁版本并且把锁定的版本写进你的部署文档。2.2 加载模型时容易忽略的配置加载模型的标准写法是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()这里有几个细节。model.eval()必须调用否则 dropout 和 batch norm 会处于训练模式导出的图行为不确定。另外如果你是从本地加载比如已经下载好的权重用from_pretrained指向本地目录即可注意目录里要有config.json、pytorch_model.bin或model.safetensors和分词器文件。还有一个隐藏配置是model.config.use_cache。对于要导出带 KV Cache 的解码器这个值需要为True。有些模型默认是True但如果你之前改过记得检查。另外model.config.forced_eos_token_id和decoder_start_token_id这两个值在写解码循环时要用到先打印出来记下print(decoder_start_token_id:, model.config.decoder_start_token_id) print(eos_token_id:, model.config.eos_token_id) print(pad_token_id:, model.config.pad_token_id)Marian 的decoder_start_token_id通常等于pad_token_id也就是 58100 这个特殊值而eos_token_id是 0。这个和 BART 完全不同写循环时如果搞混解码会立刻停或者永远不停。2.3 分词器与模型的对应关系验证在导出之前先用 PyTorch 原生推理跑一句话把结果存下来作为基准。这一步至关重要因为后面 ONNX 推理的结果要跟它对比才能判断迁移是否成功。text The quick brown fox jumps over the lazy dog. inputs tokenizer(text, return_tensorspt, paddingTrue) with torch.no_grad(): generated model.generate(**inputs, max_length64, num_beams1) baseline tokenizer.decode(generated[0], skip_special_tokensTrue) print(PyTorch baseline:, baseline)把baseline记下来。注意这里用num_beams1也就是贪心解码因为 ONNX 侧我们通常先实现贪心beam search 要复杂得多。基准一致了后面才有对比的意义。3. 编码器与解码器的分离导出实战3.1 编码器导出动态轴是成败关键编码器相对简单输入是input_ids和attention_mask输出是last_hidden_state。核心在于动态轴的声明。import torch dummy_input_ids torch.randint(0, 1000, (1, 16), dtypetorch.long) dummy_attention_mask torch.ones((1, 16), dtypetorch.long) torch.onnx.export( model.get_encoder(), (dummy_input_ids, dummy_attention_mask), encoder.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_version14, do_constant_foldingTrue, )这里dynamic_axes是重中之重。{0: batch, 1: sequence}表示第 0 维是 batch、第 1 维是序列长度两者都是动态的。如果你漏了sequence这一维导出的图就只接受长度 16 的输入换个长度的句子直接报错。我见过太多人栽在这里。opset_version选 14 是个稳妥的选择。太低比如 11可能不支持某些 attention 算子太高比如 17 以上某些运行时的支持还不完善。14 在 ONNX Runtime 1.17 上验证充分。do_constant_foldingTrue会在导出时把能算的常量提前算掉减小图体积一般开着没坏处。3.2 解码器导出KV Cache 的输入输出设计解码器是难点。我们要导出一个单步解码的图输入是当前步的decoder_input_ids、编码器输出encoder_hidden_states、编码器 attention mask以及上一步的 KV Cache输出是当前步的 logits 和更新后的 KV Cache。先看 KV Cache 的结构。Marian 的解码器有若干层每层有 self-attention 的 key/value 和 cross-attention 的 key/value。self-attention 的 KV 每步都在增长cross-attention 的 KV 只算一次因为编码器输出不变。为了简化很多实现会把 cross-attention 的 KV 也每步重算代价是浪费一点算力但代码简单。追求性能的话cross-attention 的 KV 应该只算一次并缓存。导出时我们需要构造带 past_key_values 的输入。transformers提供了model.decoder和model.encoder但直接用model.decoder导出比较麻烦因为它的 forward 签名涉及past_key_values这个嵌套元组。一个更可控的做法是写一个包装类class DecoderWrapper(torch.nn.Module): def __init__(self, decoder): super().__init__() self.decoder decoder def forward(self, decoder_input_ids, encoder_hidden_states, encoder_attention_mask, *past_kv): # past_kv 是展平的 KV 张量列表 num_layers self.decoder.config.decoder_layers past_key_values [] idx 0 for _ in range(num_layers): past_key_values.append((past_kv[idx], past_kv[idx 1])) idx 2 past_key_values tuple(past_key_values) outputs self.decoder( input_idsdecoder_input_ids, encoder_hidden_statesencoder_hidden_states, encoder_attention_maskencoder_attention_mask, past_key_valuespast_key_values, use_cacheTrue, ) # outputs 包含 last_hidden_state, present_key_values logits outputs.last_hidden_state present outputs.present_key_values flat_present [] for layer_past in present: flat_present.extend([layer_past[0], layer_past[1]]) return (logits, *flat_present)这个包装类把嵌套的 KV 元组展平成平铺的张量列表因为 ONNX 的输入输出必须是张量不能是嵌套结构。这是导出带 Cache 解码器的通用技巧。然后构造 dummy 输入。假设编码器输出序列长度为 16解码器当前步长度为 1第一步层数为 6num_layers model.config.decoder_layers dummy_decoder_input_ids torch.tensor([[model.config.decoder_start_token_id]], dtypetorch.long) dummy_encoder_hidden_states torch.randn(1, 16, model.config.d_model) dummy_encoder_attention_mask torch.ones(1, 16, dtypetorch.long) # 初始 KV Cache 长度为 0 past_kv [] for _ in range(num_layers): past_kv.append(torch.randn(1, model.config.decoder_attention_heads, 0, model.config.d_model // model.config.decoder_attention_heads)) past_kv.append(torch.randn(1, model.config.decoder_attention_heads, 0, model.config.d_model // model.config.decoder_attention_heads))注意初始 KV 的序列长度维是 0这是合法的表示还没有缓存。d_model // decoder_attention_heads是每个头的维度这个计算要准确否则形状对不上。导出wrapper DecoderWrapper(model.decoder) wrapper.eval() input_names [decoder_input_ids, encoder_hidden_states, encoder_attention_mask] output_names [logits] dynamic_axes { decoder_input_ids: {0: batch, 1: decoder_sequence}, encoder_hidden_states: {0: batch, 1: encoder_sequence}, encoder_attention_mask: {0: batch, 1: encoder_sequence}, logits: {0: batch, 1: decoder_sequence}, } for i in range(num_layers): input_names.append(fpast_key_{i}) input_names.append(fpast_value_{i}) output_names.append(fpresent_key_{i}) output_names.append(fpresent_value_{i}) dynamic_axes[fpast_key_{i}] {0: batch, 2: past_sequence} dynamic_axes[fpast_value_{i}] {0: batch, 2: past_sequence} dynamic_axes[fpresent_key_{i}] {0: batch, 2: present_sequence} dynamic_axes[fpresent_value_{i}] {0: batch, 2: present_sequence} torch.onnx.export( wrapper, (dummy_decoder_input_ids, dummy_encoder_hidden_states, dummy_encoder_attention_mask, *past_kv), decoder.onnx, input_namesinput_names, output_namesoutput_names, dynamic_axesdynamic_axes, opset_version14, do_constant_foldingTrue, )这里 KV Cache 的动态轴是第 2 维past_sequence因为形状是[batch, heads, seq, head_dim]。这个维度必须动态否则每步长度变化就报错。3.3 导出后的图结构检查导出完别急着跑先用onnx库检查一下图是否合法import onnx for name in [encoder.onnx, decoder.onnx]: m onnx.load(name) onnx.checker.check_model(m) print(f{name} 输入:) for inp in m.graph.input: print( , inp.name, [d.dim_param or d.dim_value for d in inp.type.tensor_type.shape.dim]) print(f{name} 输出:) for out in m.graph.output: print( , out.name, [d.dim_param or d.dim_value for d in out.type.tensor_type.shape.dim])这一步能帮你确认动态轴是否真的生效。如果某个维度显示的是具体数字而不是batch、sequence这样的名字说明动态轴没设对回去检查。4. 用 ONNX Runtime 跑通完整翻译流程4.1 贪心解码循环的实现有了编码器和解码器两个图就可以写推理循环了。核心逻辑是编码器跑一次得到encoder_hidden_states然后解码器从起始 token 开始每步生成一个 token把新 token 拼到输入里同时更新 KV Cache直到遇到 EOS 或者达到最大长度。import numpy as np import onnxruntime as ort encoder_session ort.InferenceSession(encoder.onnx, providers[CPUExecutionProvider]) decoder_session ort.InferenceSession(decoder.onnx, providers[CPUExecutionProvider]) def translate(text, max_length64): inputs tokenizer(text, return_tensorsnp, paddingTrue) input_ids inputs[input_ids].astype(np.int64) attention_mask inputs[attention_mask].astype(np.int64) encoder_hidden encoder_session.run( [encoder_hidden_states], {input_ids: input_ids, attention_mask: attention_mask} )[0] num_layers model.config.decoder_layers batch_size input_ids.shape[0] heads model.config.decoder_attention_heads head_dim model.config.d_model // heads # 初始化空 KV Cache past_kv {} for i in range(num_layers): past_kv[fpast_key_{i}] np.zeros((batch_size, heads, 0, head_dim), dtypenp.float32) past_kv[fpast_value_{i}] np.zeros((batch_size, heads, 0, head_dim), dtypenp.float32) decoder_input np.array([[model.config.decoder_start_token_id]], dtypenp.int64) generated [] for step in range(max_length): feed { decoder_input_ids: decoder_input, encoder_hidden_states: encoder_hidden, encoder_attention_mask: attention_mask, } feed.update(past_kv) outputs decoder_session.run(None, feed) logits outputs[0] present outputs[1:] next_token int(np.argmax(logits[0, -1, :])) if next_token model.config.eos_token_id: break generated.append(next_token) # 更新 KV Cache for i in range(num_layers): past_kv[fpast_key_{i}] present[i * 2] past_kv[fpast_value_{i}] present[i * 2 1] decoder_input np.array([[next_token]], dtypenp.int64) return tokenizer.decode(generated, skip_special_tokensTrue)这段代码有几个关键点。第一decoder_input每步只喂最新的一个 token因为历史信息已经在 KV Cache 里了重复喂会导致重复计算甚至错误。第二logits[0, -1, :]取的是最后一个位置的 logits因为当前步只生成了一个 token。第三KV Cache 的更新是直接替换不是拼接因为present已经包含了历史加当前步的全部内容。4.2 结果对比与精度验证跑一句和基准一样的话对比结果result translate(The quick brown fox jumps over the lazy dog.) print(ONNX result:, result) print(PyTorch baseline:, baseline)如果两者一致或者语义等价说明迁移成功。如果 ONNX 结果是乱码或者空按下面的顺序排查先确认decoder_start_token_id和eos_token_id是否用对再确认 KV Cache 的维度顺序是否和导出时一致最后检查attention_mask是否传对编码器的 mask 和编码器输出的序列长度必须匹配。我遇到过一次结果是重复 token 的情况排查发现是 KV Cache 的past_sequence维在导出时被固定成了 0导致每步都从零开始。回去把动态轴改对就好了。4.3 批处理与变长输入的处理上面的代码 batch size 是 1。要支持批处理编码器侧没问题因为动态轴支持 batch。解码器侧要注意不同样本的生成长度不同有的先遇到 EOS有的后遇到。简单做法是等所有样本都结束再停但这样会浪费算力。更好的做法是维护一个已完成的标记对已完成的样本用 pad token 填充同时把它们的 logits 屏蔽掉。变长输入的处理主要靠 padding。编码器接受 padding 后的输入attention_mask告诉它哪些位置是真实的。解码器侧encoder_attention_mask同样要传否则 cross-attention 会关注到 padding 位置影响翻译质量。这一点在短句翻译时尤其明显我实测过不传 mask 的话短句翻译质量会下降。5. INT8 量化体积减半背后的精度博弈5.1 动态量化与静态量化的选择ONNX Runtime 提供两种量化方式。动态量化dynamic quantization在推理时动态计算激活值的量化参数不需要校准数据集用起来简单适合 LSTM、Transformer 这类模型。静态量化static quantization需要一批校准数据来预先确定激活值的量化范围精度通常更好但流程复杂。对于翻译模型我推荐先试动态量化。命令很简单from onnxruntime.quantization import quantize_dynamic, QuantType quantize_dynamic( model_inputencoder.onnx, model_outputencoder_int8.onnx, weight_typeQuantType.QInt8, ) quantize_dynamic( model_inputdecoder.onnx, model_outputdecoder_int8.onnx, weight_typeQuantType.QInt8, )QuantType.QInt8表示权重用 8 位有符号整数。跑完之后模型体积大概能降到原来的四分之一到三分之一因为权重从 FP32 变成 INT8理论上四分之一但有些算子不参与量化所以实际略大。5.2 量化后精度下降的定位方法量化后第一件事是重新跑基准句子对比翻译结果。如果结果明显变差比如出现重复、漏译、语序混乱说明量化损失过大。定位方法是用onnxruntime.quantization提供的quantize_static配合校准数据或者干脆放弃量化。我实测下来Marian 系列在动态量化后简单句子的翻译基本无损但长句和复杂从句的质量会有可感知的下降。如果你的场景对质量要求高建议只量化编码器解码器保持 FP32因为解码器的误差会随着自回归逐步累积放大。还有一个技巧是分层量化对 attention 相关的权重保持 FP32只量化 FFN 层。这需要手动指定nodes_to_exclude稍微麻烦但效果更好。5.3 量化模型的推理速度实测量化不一定更快。在 CPU 上INT8 通常比 FP32 快因为整数运算吞吐高。但在某些没有 INT8 加速指令的平台上量化反而因为要来回转换数据类型而变慢。我实测的数据是在支持 AVX512 的服务器 CPU 上INT8 比 FP32 快约 1.8 倍在普通笔记本 CPU 上快约 1.3 倍在没有专门指令的老 CPU 上基本持平甚至略慢。所以量化之前先明确你的目标平台。如果是为了减小体积方便分发量化值得做如果是为了提速先测一下再决定。6. 迁移过程中最常翻车的几个点6.1 动态轴漏设导致的形状报错这是最高频的问题。表现是推理时换一个长度的句子就报Invalid input shape或者Got invalid dimensions for input。根因就是导出时dynamic_axes没覆盖到所有需要动态的维度。排查方法是用onnx.load打印每个输入输出的维度看哪些是具体数字。凡是和序列长度、batch 相关的维度都应该是dim_param名字而不是dim_value数字。特别注意 KV Cache 的past_sequence维这个最容易漏。6.2 分词器与模型不匹配的隐蔽症状有时候模型导出没问题推理也不报错但翻译结果就是不对比如把英文翻译成了另一种语言或者输出一堆无意义的 token。这往往是分词器用错了。Marian 模型的分词器是绑定在模型目录里的如果你手动指定了别的分词器或者用了错误的source_lang/target_lang前缀就会出问题。有些多语言模型需要在输入前加语言标记比如cmn_Hans表示翻译成简体中文。这个标记加不加、怎么加取决于模型训练时的设置一定要看模型卡model card的说明。6.3 解码起始符和结束符的陷阱前面提过Marian 用pad_token_id作为decoder_start_token_id而 BART 用s。如果你从 Marian 换到 BART忘了改起始符解码器第一步就懵了输出完全错误。结束符同理。有些模型的eos_token_id是 0有些是 2有些是 1。写循环时如果硬编码了错误的 EOS要么立刻停生成空结果要么永远不停跑满 max_length。正确做法是从model.config里读不要硬编码。6.4 显存不降反升的意外情况有人迁移到 ONNX 是为了省显存结果发现 ONNX Runtime 的显存占用比 PyTorch 还高。这通常是因为 ONNX Runtime 默认会预分配一大块显存或者 KV Cache 的实现方式导致每步都在申请新内存。解决办法是配置SessionOptions开启内存复用sess_options ort.SessionOptions() sess_options.enable_mem_pattern True sess_options.enable_cpu_mem_arena True sess_options.graph_optimization_level ort.GraphOptimizationLevel.ORT_ENABLE_ALL session ort.InferenceSession(decoder.onnx, sess_options, providers[CUDAExecutionProvider])另外KV Cache 的张量在每步之间要复用同一块内存不要每步都新建 numpy 数组否则 Python 侧的 GC 压力会很大。7. 从 ONNX 再出发后续可扩展的方向7.1 转 TensorRT 或 OpenVINO 的注意事项ONNX 是个中间站很多人最终目标是 TensorRTNVIDIA GPU或 OpenVINOIntel CPU/GPU。从 ONNX 转这两个后端时动态 shape 的支持程度不同。TensorRT 对动态 shape 支持较好但需要指定 optimization profile也就是给每个动态维指定最小、最优、最大三个值。OpenVINO 对动态 shape 的支持相对弱一些某些情况下需要固定 shape。转 TensorRT 时KV Cache 的动态维要特别小心因为 TensorRT 对past_sequence这种从 0 开始增长的维度处理起来有坑。一个常见做法是把最大序列长度固定下来牺牲一点灵活性换取稳定性。7.2 端侧部署时的进一步压缩如果目标是手机或嵌入式设备ONNX 之后可能还要转 NCNN、MNN 或者 TFLite。这些框架对算子的支持更有限可能需要把一些融合算子拆开。另外端侧对体积极其敏感除了 INT8 量化还可以考虑剪枝和知识蒸馏把模型本身做小。我做过一个实验把 Marian 的层数从 6 层剪到 4 层再量化到 INT8体积从 300MB 降到 40MB 左右翻译质量在简单句上基本无损复杂句有下降但可接受。这个取舍要看具体场景。7.3 服务化部署的工程考量如果是要做成在线翻译服务除了模型本身还要考虑请求排队、批处理、超时控制、缓存。ONNX Runtime 支持多线程和并行 session可以配置intra_op_num_threads和inter_op_num_threads来压榨 CPU。批处理方面可以把多个请求攒成一批一起编码解码时用前面说的已完成标记来处理不同长度的输出。缓存也很重要翻译请求里重复的句子不少加一层 LRU 缓存能显著降低平均延迟。这些工程细节虽然和 ONNX 迁移本身无关但决定了最终服务的可用性。我个人在实际操作中的体会是ONNX 迁移这件事导出只是开始真正的功夫在推理循环的调试和量化的取舍上。把基准对比做扎实把动态轴设全把 KV Cache 的维度理顺剩下的就是耐心调参。每次遇到诡异的结果先回到 PyTorch 基准跑一遍确认是导出问题还是推理逻辑问题这个习惯帮我省了大量时间。

读完文章,也想定制专属网站?

尧图设计师 24 小时内与您沟通定制方案

免费获取报价 →
↑