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

torchsummary输出重复问题求助:附CNN模型复现代码

Torchsummary输出重复问题排查

我正在复现一篇基于CNN的分类方案论文,写出了如下简化版代码,但运行torchsummary时出现输出重复的情况,查过其GitHub问答区没找到相关问题记录。

复现代码

import torch
import torch.nn as nn
from torchsummary import summary

class CNN_Pred2D(nn.Module):
    def __init__(self, n_filters=[8,8,8], debug=True):
        super().__init__()
        self.debug = debug
        
        self.model = nn.Sequential(
            nn.Conv2d(1, n_filters[0], kernel_size=(1,82)),
            nn.ReLU(),
            nn.Conv2d(n_filters[0], n_filters[0], kernel_size=(3,1)),
            nn.ReLU(),
            nn.MaxPool2d(kernel_size=(2,1)),
            
            nn.Conv2d(n_filters[0], n_filters[1], kernel_size=(3,1)),
            nn.ReLU(),
            nn.MaxPool2d(kernel_size=(2,1)),
            
            nn.Flatten(),
            nn.Linear(104,1),
            nn.Sigmoid()
        )

        
    def forward(self, X):
        out = self.model(X)
#         print(out.shape)
        return out

device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
model = CNN_Pred2D().to(device)

summary(model, [(1, 60,82)])

异常输出截图

torchsummary重复输出

解决方法

  • 升级/替换torchsummary:旧版torchsummary对嵌套nn.Sequential的模型结构识别存在bug,建议升级到最新稳定版,或改用torchinfo(torchsummary的官方维护分支),它对复杂模型结构的兼容性更好。
  • 取消嵌套Sequential:将self.model中的层直接定义在模型类的__init__方法中,不通过嵌套的nn.Sequential封装,让torchsummary能正确遍历每一层结构。修改示例如下:
    class CNN_Pred2D(nn.Module):
        def __init__(self, n_filters=[8,8,8], debug=True):
            super().__init__()
            self.debug = debug
            
            # 直接定义每层,取消嵌套Sequential
            self.conv1 = nn.Conv2d(1, n_filters[0], kernel_size=(1,82))
            self.relu1 = nn.ReLU()
            self.conv2 = nn.Conv2d(n_filters[0], n_filters[0], kernel_size=(3,1))
            self.relu2 = nn.ReLU()
            self.pool1 = nn.MaxPool2d(kernel_size=(2,1))
            
            self.conv3 = nn.Conv2d(n_filters[0], n_filters[1], kernel_size=(3,1))
            self.relu3 = nn.ReLU()
            self.pool2 = nn.MaxPool2d(kernel_size=(2,1))
            
            self.flatten = nn.Flatten()
            self.fc = nn.Linear(104,1)
            self.sigmoid = nn.Sigmoid()
    
        def forward(self, X):
            x = self.relu1(self.conv1(X))
            x = self.relu2(self.conv2(x))
            x = self.pool1(x)
            x = self.relu3(self.conv3(x))
            x = self.pool2(x)
            x = self.flatten(x)
            x = self.sigmoid(self.fc(x))
            return x
    
  • 验证模型结构:通过print(list(model.named_modules()))打印模型所有模块,确认是否存在重复定义的层,排除代码逻辑错误。
  • 清理环境缓存:重启Python内核,或执行torch.cuda.empty_cache()清理缓存,避免因缓存导致的结构识别异常。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.08 08:35:23