如何将PyTorch编解码模型opus-mt-fr-en导出为单个ONNX文件?
将Helsinki-NLP/opus-mt-fr-en导出为单个ONNX文件的解决方案
你当前通过optimum.onnxruntime导出的Seq2Seq模型被拆分为三个ONNX文件,核心问题是共享嵌入层重复存储和decoder两种模式参数冗余,导致体积膨胀至原PyTorch模型的3倍。以下是两种可行的合并方案:
方案1:自定义模型包装器,直接导出单个ONNX文件
通过构建一个整合encoder和decoder的包装模型,强制复用共享嵌入层,跳过optimum的默认拆分逻辑,直接用torch.onnx.export导出单个文件。
代码实现
import torch from transformers import AutoModelForSeq2SeqLM, AutoTokenizer import os hf_model_id = "Helsinki-NLP/opus-mt-fr-en" onnx_save_path = "./opus-mt-fr-en_single.onnx" # 加载预训练模型与分词器 model = AutoModelForSeq2SeqLM.from_pretrained(hf_model_id) tokenizer = AutoTokenizer.from_pretrained(hf_model_id) model.eval() # 定义包装模型,整合encoder与decoder并共享嵌入层 class UnifiedSeq2SeqModel(torch.nn.Module): def __init__(self, original_model): super().__init__() self.model = original_model self.shared_embeddings = self.model.shared # 复用共享嵌入层 def forward(self, input_ids, attention_mask, decoder_input_ids=None, decoder_attention_mask=None): # 自动初始化decoder输入(默认用bos token) if decoder_input_ids is None: decoder_input_ids = torch.tensor([[tokenizer.bos_token_id]], device=input_ids.device) decoder_attention_mask = torch.ones_like(decoder_input_ids) # 编码器前向传播 encoder_outputs = self.model.encoder(input_ids=input_ids, attention_mask=attention_mask) # 解码器前向传播 decoder_outputs = self.model.decoder( input_ids=decoder_input_ids, attention_mask=decoder_attention_mask, encoder_hidden_states=encoder_outputs.last_hidden_state, encoder_attention_mask=attention_mask ) # 生成最终logits logits = self.model.lm_head(decoder_outputs.last_hidden_state) return logits # 初始化包装模型 unified_model = UnifiedSeq2SeqModel(model) # 准备示例输入用于导出 sample_input = tokenizer("je regarde la tele", return_tensors="pt") input_args = (sample_input["input_ids"], sample_input["attention_mask"]) # 导出为单个ONNX文件 torch.onnx.export( unified_model, input_args, onnx_save_path, export_params=True, opset_version=17, do_constant_folding=True, input_names=["input_ids", "attention_mask"], output_names=["logits"], dynamic_axes={ "input_ids": {0: "batch_size", 1: "seq_len"}, "attention_mask": {0: "batch_size", 1: "seq_len"}, "logits": {0: "batch_size", 1: "decoder_seq_len"} } ) print(f"单个ONNX模型已保存至: {onnx_save_path}")
优势
- 彻底避免共享嵌入层的重复存储,体积可压缩至接近原PyTorch模型大小
- 导出流程可控,无需额外处理拆分后的文件
- 支持动态batch和序列长度,适配不同输入场景
方案2:对已导出的ONNX文件进行权重去重合并
如果你已经导出了三个ONNX文件,可以通过ONNX工具手动合并重复权重,适合熟悉ONNX图结构的用户。
核心步骤
- 加载三个ONNX模型,提取所有初始化权重
- 通过哈希值识别重复权重,只保留一份副本
- 重构ONNX图,将encoder和decoder的节点整合到同一个图中,并引用共享权重
- 保存合并后的单个ONNX文件
简化代码示例(仅展示权重去重逻辑)
import onnx from onnx import numpy_helper import hashlib # 加载已导出的三个模型 encoder_model = onnx.load("./onnx_model_fr_en/encoder_model.onnx") decoder_model = onnx.load("./onnx_model_fr_en/decoder_model.onnx") # 权重去重:通过哈希值识别重复权重 weight_map = {} for model in [encoder_model, decoder_model]: for initializer in model.graph.initializer: weight_data = numpy_helper.to_array(initializer) hash_key = hashlib.sha256(weight_data.tobytes()).hexdigest() if hash_key not in weight_map: weight_map[hash_key] = initializer # 构建新的ONNX图(需手动处理节点连接,此步骤需熟悉ONNX结构) merged_graph = onnx.GraphProto() merged_graph.name = "opus-mt-fr-en-unified" merged_graph.initializer.extend(weight_map.values()) # 此处需手动复制encoder和decoder的节点、输入输出定义,过程较繁琐,建议优先使用方案1 # 保存合并后的模型 # onnx.save(onnx.ModelProto(graph=merged_graph), "./merged_opus_mt.onnx")
注意事项
- 手动重构ONNX图容易出错,需验证节点连接的正确性
- 如果需要保留
decoder_with_past模式(用于增量生成优化),需额外处理past key/value的输入输出,复杂度会大幅提升
验证合并结果
导出单个ONNX文件后,可通过以下代码测试翻译功能是否正常:
import onnxruntime as ort import numpy as np from transformers import AutoTokenizer tokenizer = AutoTokenizer.from_pretrained("Helsinki-NLP/opus-mt-fr-en") ort_session = ort.InferenceSession("./opus-mt-fr-en_single.onnx") # 测试输入 french_text = "je regarde la tele" inputs = tokenizer(french_text, return_tensors="np") # 推理 outputs = ort_session.run(None, { "input_ids": inputs["input_ids"], "attention_mask": inputs["attention_mask"] }) # 解码结果 logits = outputs[0] predicted_ids = np.argmax(logits, axis=-1) english_translation = tokenizer.decode(predicted_ids[0], skip_special_tokens=True) print(f"输入(法语): {french_text}") print(f"输出(英语): {english_translation}")
内容的提问来源于stack exchange,提问作者Franck Dernoncourt
相关产品推荐
相关产品推荐

