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

如何获取以字典为输入的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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.16 11:25:59