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

PyTorch模型转TorchScript遇输入类型错误,求排查解决

PyTorch转TorchScript报错原因及修复方案

错误原因

  • 输入结构不匹配:torch.jit.trace时传入了3个Tensor作为位置参数,但后续调用traced_model仅传入1个Tensor,导致模型内部处理输入时,字典类型的参数同时出现Tensor和复杂Tuple结构,触发类型不一致错误。
  • 参数传递方式错误:T5模型的forward方法对参数顺序和传递逻辑有特定要求,直接将decoder_input_ids作为位置参数传入,可能与模型期望的参数映射不匹配,引发内部字典输入的类型混乱。

代码存在的问题

  • 调用traced模型时输入数量与trace阶段不一致:trace时传入了input_ids、attention_mask、decoder_input_ids三个参数,但调用时仅传input_ids,破坏了trace记录的输入结构。
  • 未遵循模型参数传递规范:T5模型更适合通过关键字参数传递输入,直接使用位置参数易导致参数错位,触发内部处理逻辑的类型错误。

修复方案

方案1:统一输入格式,使用关键字参数trace和调用

确保trace与调用时的输入结构一致,用字典传递关键字参数避免位置错位:

model = AutoModelForSeq2SeqLM.from_pretrained("Seungjun/t5-small-finetuned-xsum")
model.eval()

# 准备符合模型要求的输入字典
full_inputs = {
    "input_ids": inputs['input_ids'],
    "attention_mask": inputs['attention_mask'],
    "decoder_input_ids": output
}

# 使用关键字参数示例进行trace
traced_model = torch.jit.trace(model, example_kwarg_inputs=full_inputs)

# 调用时传入完整的输入字典(或保持与trace时一致的输入结构)
out = traced_model(**full_inputs)

方案2:改用torch.jit.script处理复杂模型

T5这类含动态控制流的模型,用script比trace更稳定,无需依赖示例输入的执行路径:

model = AutoModelForSeq2SeqLM.from_pretrained("Seungjun/t5-small-finetuned-xsum")
model.eval()

# 直接脚本化模型
scripted_model = torch.jit.script(model)

# 按模型正常调用方式传入参数
out = scripted_model(inputs['input_ids'], inputs['attention_mask'])

方案3:让模型自动处理decoder输入

如果不需要手动传入decoder_input_ids,可仅trace encoder相关参数,让模型内部生成初始decoder输入:

model = AutoModelForSeq2SeqLM.from_pretrained("Seungjun/t5-small-finetuned-xsum")
model.eval()

# 仅传入encoder所需参数进行trace
traced_model = torch.jit.trace(model, (inputs['input_ids'], inputs['attention_mask']))

# 调用时同样传入encoder参数
out = traced_model(inputs['input_ids'], inputs['attention_mask'])

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.03 08:05:23