You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何将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图结构的用户。

核心步骤

  1. 加载三个ONNX模型,提取所有初始化权重
  2. 通过哈希值识别重复权重,只保留一份副本
  3. 重构ONNX图,将encoder和decoder的节点整合到同一个图中,并引用共享权重
  4. 保存合并后的单个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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.13 10:04:56