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

PyTorch中如何正确打印神经网络各层的激活值输出

问题原因
  • nn.ReLU()属于无参数运算层,本身不存在可训练的权重、偏置项,调用list(model.myRelu.parameters())返回空列表是完全正常的现象,parameters()接口仅用于获取层的可训练参数,和激活值获取没有关系。
  • 激活值是模型前向传播过程中产生的动态中间结果,不会默认持久化存储在层对象的属性中,无法通过直接访问层参数、权重字典的方式获取,需要通过钩子捕获或者改写前向逻辑的方式留存。
实现方案

两种常用方案都可以拿到各层激活值,按需选择即可。

方案1:注册前向钩子(无需修改原有模型结构)

PyTorch提供了register_forward_hook接口,可以给指定层注册回调函数,前向传播经过该层时会自动触发回调,把层的输入、输出传到回调函数里,用这个机制可以无侵入地捕获任意层的输出,代码示例:

import torch
import torch.nn as nn
from torchvision import datasets, transforms

# 修正原有模型缺失的展平维度计算(原代码out_features未定义,MNIST输入为28*28单通道图,按结构计算conv+pool后展平维度为256*4*4=4096)
class Net(nn.Module):
    def __init__(self):
        super().__init__()
        self.conv1 = nn.Conv2d(1, 128, 5)
        self.pool = nn.MaxPool2d(2, 2)
        self.conv2 = nn.Conv2d(128, 256, 5)
        self.fc1 = nn.Linear(256*4*4, 120)
        self.fc2 = nn.Linear(120, 84)
        self.fc3 = nn.Linear(84, 10)
        self.myRelu = nn.ReLU()

    def forward(self, x):
        x = self.pool(self.myRelu(self.conv1(x)))
        x = self.pool(self.myRelu(self.conv2(x)))
        x = torch.flatten(x, 1)
        x = self.myRelu(self.fc1(x))
        x = self.myRelu(self.fc2(x))
        x = self.fc3(x)
        return x

model = Net()
activation_map = {}

# 定义钩子生成函数,传入层名作为key存储激活值
def get_activation(layer_name):
    def hook_fn(module, input_tensor, output_tensor):
        # 用detach切断梯度,避免占用多余显存
        activation_map[layer_name] = output_tensor.detach()
    return hook_fn

# 给需要捕获输出的层注册钩子
model.conv1.register_forward_hook(get_activation("conv1输出"))
model.conv2.register_forward_hook(get_activation("conv2输出"))
model.fc1.register_forward_hook(get_activation("fc1输出"))
model.fc2.register_forward_hook(get_activation("fc2输出"))
model.fc3.register_forward_hook(get_activation("fc3输出"))

# 执行一次前向传播触发钩子
test_transform = transforms.Compose([transforms.ToTensor()])
test_set = datasets.MNIST(root="./mnist", train=False, download=True, transform=test_transform)
test_loader = torch.utils.data.DataLoader(test_set, batch_size=1, shuffle=True)
test_img, _ = next(iter(test_loader))
_ = model(test_img)

# 打印各层激活值
for name, act_val in activation_map.items():
    print(f"===== {name} =====")
    print(f"激活值形状:{act_val.shape}")
    print(f"激活值内容:{act_val}\n")

注意:你当前代码里所有位置的ReLU运算都复用了同一个self.myRelu实例,给这个实例注册钩子只会捕获它最后一次被调用的输出(也就是fc2层之后的ReLU结果)。如果要单独捕获每一层ReLU的输出,要么在__init__里为每个激活位置单独定义ReLU实例(比如self.relu1、self.relu2...),要么直接给卷积、全连接层注册钩子,拿到层输出后手动计算ReLU结果即可。

方案2:改写前向传播方法(逻辑直白适合调试)

如果不想用钩子,可以直接在forward函数里把需要留存的中间激活值存为模型的属性,每次前向传播后直接访问属性即可拿到结果,代码示例:

class Net(nn.Module):
    def __init__(self):
        super().__init__()
        self.conv1 = nn.Conv2d(1, 128, 5)
        self.pool = nn.MaxPool2d(2, 2)
        self.conv2 = nn.Conv2d(128, 256, 5)
        self.fc1 = nn.Linear(256*4*4, 120)
        self.fc2 = nn.Linear(120, 84)
        self.fc3 = nn.Linear(84, 10)
        self.relu = nn.ReLU()
        # 初始化存储激活值的字典
        self.activations = {}

    def forward(self, x):
        x = self.conv1(x)
        self.activations["conv1输出"] = x.detach()
        x = self.relu(x)
        self.activations["relu1输出"] = x.detach()
        x = self.pool(x)
        self.activations["pool1输出"] = x.detach()

        x = self.conv2(x)
        self.activations["conv2输出"] = x.detach()
        x = self.relu(x)
        self.activations["relu2输出"] = x.detach()
        x = self.pool(x)
        self.activations["pool2输出"] = x.detach()

        x = torch.flatten(x, 1)
        x = self.fc1(x)
        self.activations["fc1输出"] = x.detach()
        x = self.relu(x)
        self.activations["relu3输出"] = x.detach()

        x = self.fc2(x)
        self.activations["fc2输出"] = x.detach()
        x = self.relu(x)
        self.activations["relu4输出"] = x.detach()

        x = self.fc3(x)
        self.activations["fc3输出"] = x.detach()
        return x

使用时跑完前向传播,直接访问model.activations就能拿到所有层的激活值,不需要额外注册接口,适合快速调试场景。

避坑提醒
  • 不要尝试通过parameters()、state_dict()接口获取激活值:这两个接口仅存储层的可训练参数、持久化缓存,不会保存前向传播产生的动态计算结果。
  • 捕获激活值时记得调用detach()切断梯度关联,避免留存的激活值占用计算图资源,增加不必要的显存开销。
  • 不要复用同一个层实例处理网络中不同位置的运算,否则不管是钩子还是属性存储,都会被后续调用的结果覆盖,无法拿到每个位置的独立输出。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.30 09:24:21