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

如何在PyTorch中绘制激活函数的梯度?

获取PyTorch激活层输出的梯度(ReLU/tanh)

首先要明确:model[1].grad 无法获取梯度的核心原因是——ReLU这类激活层没有可学习参数,grad 属性仅针对模型的可训练参数(比如Linear层的weight和bias)。你需要的是损失相对于激活层输出的梯度,这需要通过Tensor的钩子(hook)或直接梯度计算来捕获。

以下是两种可行的实现方法:

方法1:为激活层输出注册梯度钩子

通过在激活层的输出Tensor上注册register_hook,反向传播时就能自动保存对应的梯度:

import torch

# 注意:将函数式的torch.tanh()改为nn.Tanh()模块,方便捕获输出
model = torch.nn.Sequential(
    torch.nn.Linear(1, 2),
    torch.nn.ReLU(),
    torch.nn.Linear(2, 1),
    torch.nn.Tanh()
)

# 用于存储激活层的输出和梯度
relu_grad = None
tanh_grad = None

# 定义钩子函数:保存梯度副本
def save_relu_grad(grad):
    global relu_grad
    relu_grad = grad.clone()

def save_tanh_grad(grad):
    global tanh_grad
    tanh_grad = grad.clone()

# 前向传播,分步捕获激活层输出并注册钩子
x = torch.randn(1, 1, requires_grad=True)
linear1_out = model[0](x)
relu_out = model[1](linear1_out)
relu_out.register_hook(save_relu_grad)  # 为ReLU输出挂勾

linear2_out = model[2](relu_out)
tanh_out = model[3](linear2_out)
tanh_out.register_hook(save_tanh_grad)  # 为Tanh输出挂勾

# 计算损失并反向传播
loss = tanh_out.sum()
loss.backward()

# 查看捕获的梯度
print("ReLU输出的梯度:", relu_grad)
print("Tanh输出的梯度:", tanh_grad)

方法2:用torch.autograd.grad直接计算梯度

如果只需要单次计算损失对激活层输出的梯度,也可以跳过钩子,直接用torch.autograd.grad计算:

import torch

model = torch.nn.Sequential(
    torch.nn.Linear(1, 2),
    torch.nn.ReLU(),
    torch.nn.Linear(2, 1),
    torch.nn.Tanh()
)

x = torch.randn(1, 1, requires_grad=True)

# 前向传播,分步记录激活层输出
linear1_out = model[0](x)
relu_out = model[1](linear1_out)
linear2_out = model[2](relu_out)
tanh_out = model[3](linear2_out)

loss = tanh_out.sum()

# 直接计算损失对ReLU输出的梯度(retain_graph=True允许后续计算)
relu_grad = torch.autograd.grad(loss, relu_out, retain_graph=True)[0]
# 计算损失对Tanh输出的梯度
tanh_grad = torch.autograd.grad(loss, tanh_out)[0]

print("ReLU输出的梯度:", relu_grad)
print("Tanh输出的梯度:", tanh_grad)

额外说明

  • 如果你使用的是函数式激活(比如torch.tanh()而非nn.Tanh()),只需在调用函数后捕获输出Tensor,再用上述任一方法处理即可。
  • 拿到梯度后,你可以统计梯度的L2范数、分布情况,以此判断是否存在梯度消失(梯度趋近于0)或爆炸(梯度范数过大)的问题。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.21 10:17:22