PyTorch中CNN每层神经元输出打印及数量与输出获取方法咨询
PyTorch中CNN神经元输出与层信息获取方案
一、打印每个神经元的输出
PyTorch没有专门的内置函数直接打印所有神经元输出,但可以通过**注册前向钩子(Forward Hook)**来捕获任意层的输出。钩子能在模型前向传播时自动记录指定层的输出张量,灵活实现神经元输出的查看。
示例代码:
import torch import torch.nn as nn # 匹配你提供的4层卷积CNN架构 class CNN(nn.Module): def __init__(self): super().__init__() self.conv1 = nn.Conv2d(3, 6, kernel_size=5) self.pool1 = nn.MaxPool2d(2, 2) self.conv2 = nn.Conv2d(6, 16, kernel_size=5) self.pool2 = nn.MaxPool2d(2, 2) self.conv3 = nn.Conv2d(16, 32, kernel_size=3) self.pool3 = nn.MaxPool2d(2, 2) self.conv4 = nn.Conv2d(32, 64, kernel_size=3) self.pool4 = nn.MaxPool2d(2, 2) def forward(self, x): x = self.pool1(torch.relu(self.conv1(x))) x = self.pool2(torch.relu(self.conv2(x))) x = self.pool3(torch.relu(self.conv3(x))) x = self.pool4(torch.relu(self.conv4(x))) return x model = CNN() layer_outputs = {} # 定义钩子函数,捕获层输出 def output_hook(layer_name): def hook(module, input, output): layer_outputs[layer_name] = output.detach() return hook # 为所有卷积层注册钩子 for name, layer in model.named_modules(): if isinstance(layer, nn.Conv2d): layer.register_forward_hook(output_hook(name)) # 输入示例张量(假设为3通道224x224图像) input_tensor = torch.randn(1, 3, 224, 224) _ = model(input_tensor) # 查看输出:如需打印具体神经元值,可直接输出张量(建议取部分避免刷屏) for layer_name, output in layer_outputs.items(): print(f"层 {layer_name} 输出形状: {output.shape}") # 打印前1个样本、前1个通道的5x5区域神经元值 # print(output[:1, :1, :5, :5])
二、获取每层的神经元数量
神经元总数对应层输出张量的总元素数,有两种计算方式:
通过输出张量计算
利用钩子捕获的输出张量,调用output.numel()即可得到该层神经元总数。比如conv1输出形状为(1,6,220,220)时,神经元数为6*220*220=290400。通过层参数推导
卷积层输出特征图尺寸公式:输出尺寸 = (输入尺寸 - 核尺寸 + 2*填充) / 步长 + 1
神经元总数 = 输出通道数 × 输出高度 × 输出宽度结合你提供的架构(假设输入224x224图像,无填充、步长1,池化步长2):
- Conv1:输出通道6,特征图尺寸220×220 → 神经元数=6×220×220=290400
- 池化后尺寸110×110
- Conv2:输出通道16,特征图尺寸106×106 → 神经元数=16×106×106=179776
- 池化后尺寸53×53
- Conv3:输出通道32,特征图尺寸51×51 → 神经元数=32×51×51=83232
- 池化后尺寸25×25
- Conv4:输出通道64,特征图尺寸23×23 → 神经元数=64×23×23=33856
三、获取各层输出结果
除钩子外,还有两种实用方式:
- 修改模型forward函数:直接在forward中返回各层输出,示例:
调用时直接接收所有返回值即可。def forward(self, x): conv1_out = self.pool1(torch.relu(self.conv1(x))) conv2_out = self.pool2(torch.relu(self.conv2(conv1_out))) conv3_out = self.pool3(torch.relu(self.conv3(conv2_out))) conv4_out = self.pool4(torch.relu(self.conv4(conv3_out))) return conv4_out, conv1_out, conv2_out, conv3_out, conv4_out - 使用
torchsummary工具:安装后通过summary(model, input_size=(3,224,224))快速查看各层输出形状、参数数量等信息。
内容的提问来源于stack exchange,提问作者Abhiram T Babu
相关产品推荐
相关产品推荐

