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
相关产品推荐
相关产品推荐

