PyTorch模型导出ONNX格式时遇TypeError:forward()缺少必需位置参数'text'的问题排查
解决PyTorch模型导出ONNX时的TypeError: forward() missing 'text'参数问题
这个错误的核心原因很清晰:你的Model类的forward()方法需要两个输入参数(图像和text),但你在调用torch.onnx.export()时只传入了图像对应的dummy张量,导致调用forward时缺少了必填的text参数。
具体排查与解决步骤:
确认模型forward方法的签名
先检查你的Model类的forward函数定义,它应该类似这样:def forward(self, img, text): # 模型前向传播逻辑 ...这意味着模型推理时必须同时传入图像和text两个输入。
创建text对应的dummy输入张量
你需要根据实际业务中text输入的形状、数据类型,创建一个匹配的dummy张量。比如如果text是长度为25的序列,类型是LongTensor(通常用于文本索引),可以这样创建:# 假设词汇表大小为1000,序列长度为25 dummy_text = torch.randint(low=0, high=1000, size=(1, 25))注意:dummy输入的形状、数据类型必须和你实际推理时传入的text完全一致,否则导出的ONNX模型会存在兼容性问题。
修改ONNX导出代码,传入所有输入参数
在torch.onnx.export()中,需要把所有输入参数打包成元组传入第二个参数位置,而不是只传单个张量:# 替换原来的导出代码 torch.onnx.export(model, (dummy_input, dummy_text), "vitstr.onnx", verbose=True)
额外注意事项:
- 如果你后续启用了
DataParallel(你注释掉的model = torch.nn.DataParallel(model).to(device)),导出ONNX时需要使用model.module作为导出的模型对象,否则会因为DataParallel的包装层导致输入参数不匹配。 - 导出前可以先手动调用一次模型的forward方法验证输入是否正确:
model(dummy_input, dummy_text),如果这一步不报错,再执行ONNX导出会更稳妥。
内容的提问来源于stack exchange,提问作者Roua Rouatbi
相关产品推荐
相关产品推荐

