如何获取以字典为输入的PyTorch模型的结构摘要?
解决PyTorch字典输入模型的torch-summary报错问题
报错原因
torchsummary默认会将input_size列表中的每个元组作为独立位置参数传入模型的forward方法,但你的模型forward仅接受一个字典类型的参数,导致参数数量不匹配,触发TypeError: MyModel.forward() takes 2 positional arguments but 3 were given错误(self是第一个参数,torchsummary额外传入了两个参数,总参数数超出预期)。
解决方案
方案1:修改模型兼容多参数输入
调整forward方法,同时支持字典输入和独立参数输入,这样可以直接沿用原有的summary调用方式:
import torch.nn as nn class MyModel(nn.Module): def __init__(self): super(MyModel, self).__init__() self.imgmodule = ImgModule() self.nummodule = NumModule() # 替换为实际的imgmodule输出维度 + nummodule输出维度 self.predict = nn.Linear(256 + 32, 100) def forward(self, x=None, image=None, number=None): # 兼容两种输入模式 if x is not None: image = x['image'] number = x['number'] xImg = self.imgmodule(image) xNum = self.nummodule(number) x = self.predict(torch.cat([xImg, xNum], dim=1)) return x
调用summary:
from torchsummary import summary summary(model, input_size=[(3, 224, 224), (1, )])
方案2:构造字典格式的虚拟输入(不修改模型)
直接创建与模型输入格式一致的虚拟张量字典,传给summary作为输入参数:
import torch from torchsummary import summary # 构造虚拟输入,batch_size设为1,维度与实际输入一致 dummy_input = { 'image': torch.randn(1, 3, 224, 224), 'number': torch.randn(1, 1) } # 调用summary,传入虚拟输入并指定设备(与模型所在设备一致) summary(model, dummy_input, device='cpu')
注意事项
- 虚拟输入的维度、数据类型需与模型实际输入完全匹配;
- 如果模型部署在CUDA上,需将虚拟输入移至对应设备:
dummy_input[k] = dummy_input[k].cuda()。
内容的提问来源于stack exchange,提问作者arr10
相关产品推荐
相关产品推荐

