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

将Hugging Face的Flan-T5-xl转为ONNX时遇InvalidGraph错误求助

解决Flan-T5-xl转ONNX后加载失败的问题

错误原因分析

你遇到的InvalidGraph错误是因为T5模型自注意力层中的Min算子收到了int64类型输入,但该算子要求输入为浮点类型。旧版本ONNX opset(比如你用的opset11)无法自动处理这种类型转换,加上手动导出时的输入设置未贴合模型实际需求,导致导出的ONNX模型结构无效。

解决方案

方案1:使用Hugging Face官方ONNX导出工具(推荐)

Hugging Face的transformers库提供了专门的ONNX导出工具,针对Seq2Seq模型做了适配,能自动处理类型转换和输入输出结构问题。

步骤1:升级依赖

确保安装最新版本的相关库:

pip install --upgrade transformers onnx onnxruntime-gpu

步骤2:导出模型代码

from transformers import AutoTokenizer, AutoModelForSeq2SeqLM
from transformers.onnx import export, OnnxConfig

model_name = "google/flan-t5-xl"
tokenizer = AutoTokenizer.from_pretrained(model_name)
model = AutoModelForSeq2SeqLM.from_pretrained(model_name)

# 定义适配T5的ONNX配置,设置动态轴
class T5OnnxConfig(OnnxConfig):
    @property
    def inputs(self):
        return {
            "input_ids": {0: "batch_size", 1: "sequence_length"},
            "attention_mask": {0: "batch_size", 1: "sequence_length"},
            "decoder_input_ids": {0: "batch_size", 1: "decoder_sequence_length"},
        }

    @property
    def outputs(self):
        return {
            "logits": {0: "batch_size", 1: "decoder_sequence_length", 2: "vocab_size"},
        }

# 执行导出
onnx_path = "flan-t5-xl.onnx"
export(
    preprocessor=tokenizer,
    model=model,
    config=T5OnnxConfig(model.config),
    opset=16,
    output=onnx_path,
)

方案2:手动修正torch.onnx.export代码

如果坚持手动导出,需要调整以下几点:

  1. 升级opset版本到16或更高
  2. 用模型内置方法生成decoder初始输入,替代手动指定<pad>
  3. 让ONNX自动处理类型转换逻辑

修正后的导出代码:

from transformers import AutoTokenizer, AutoModelForSeq2SeqLM
import torch

device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
model = AutoModelForSeq2SeqLM.from_pretrained("google/flan-t5-xl").to(device)
tokenizer = AutoTokenizer.from_pretrained("google/flan-t5-xl")

onnx_path = "flan-t5-xl.onnx"

# 生成符合模型要求的dummy输入
dummy_input_text = "What's the disease name in this text: Example text"
dummy_inputs = tokenizer(dummy_input_text, return_tensors="pt", padding=True).to(device)
# 用模型内置方法生成decoder初始输入
dummy_decoder_input_ids = model._prepare_decoder_input_ids_from_labels(
    labels=tokenizer("<pad>", return_tensors="pt").input_ids.to(device)
)

with torch.no_grad():
    torch.onnx.export(
        model,
        (dummy_inputs["input_ids"], dummy_inputs["attention_mask"], dummy_decoder_input_ids),
        onnx_path,
        opset_version=16,  # 升级到高版本opset
        input_names=["input_ids", "attention_mask", "decoder_input_ids"],
        output_names=["logits"],
        dynamic_axes={
            "input_ids": {0: "batch_size", 1: "sequence_length"},
            "attention_mask": {0: "batch_size", 1: "sequence_length"},
            "decoder_input_ids": {0: "batch_size", 1: "decoder_sequence_length"},
            "logits": {0: "batch_size", 1: "decoder_sequence_length", 2: "vocab_size"},
        },
        do_constant_folding=True,
    )

修正后的推理代码

无论用哪种方案导出,推理时需确保输入类型匹配,调整输出处理逻辑:

import onnxruntime
import torch
from transformers import AutoTokenizer

tokenizer = AutoTokenizer.from_pretrained("google/flan-t5-xl")
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")

# 加载ONNX模型
onnx_model = onnxruntime.InferenceSession(
    "flan-t5-xl.onnx",
    providers=["CUDAExecutionProvider"] if device.type == "cuda" else ["CPUExecutionProvider"]
)

# 执行推理
input_text = input("Enter Disease/Symptom Detail: ")
inputs = tokenizer(input_text, return_tensors="pt", padding=True)
decoder_input_ids = tokenizer("<pad>", return_tensors="pt").input_ids

# 转换为numpy数组,确保类型匹配
onnx_inputs = {
    "input_ids": inputs["input_ids"].numpy(),
    "attention_mask": inputs["attention_mask"].numpy(),
    "decoder_input_ids": decoder_input_ids.numpy(),
}

# 运行推理并解析结果
onnx_output = onnx_model.run(None, onnx_inputs)[0]
predicted_ids = onnx_output.argmax(axis=-1)
decoded_output = tokenizer.decode(predicted_ids[0], skip_special_tokens=True)

print('-' * 100)
print(f"Name of Disease based on Entered Text: {decoded_output}")

内容的提问来源于stack exchange,提问作者Romi

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.23 23:14:58