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

如何在PyTorch反向传播中打印非叶子节点mul_x、sum_x的梯度?

解决PyTorch中间非叶子节点梯度打印问题

你当前代码使用register_full_backward_hook绑定在模型模块上,只能捕获整个模块的输入和输出梯度,无法获取forward函数内定义的mul_x、sum_x这类中间非叶子节点的梯度。另外PyTorch默认会在反向传播结束后释放非叶子节点的梯度,需要显式保留才能访问。

解决方案:给中间张量注册反向hook并保留梯度

直接给mul_x、sum_x这类中间张量注册register_hook,同时调用retain_grad()保留它们的梯度,具体修改如下:

import torch

class func_NN(torch.nn.Module):
    def __init__(self,) :
        super().__init__()
        self.a = torch.nn.Parameter(torch.rand(1))
        self.b = torch.nn.Parameter(torch.rand(1))

    def forward(self, inp):
        mul_x = torch.cos(self.a.view(-1,1)*inp)
        # 给mul_x注册反向hook并保留梯度
        mul_x.register_hook(lambda grad: print(f"mul_x梯度: {grad}"))
        mul_x.retain_grad()
        
        sum_x = mul_x - self.b
        # 给sum_x注册反向hook并保留梯度
        sum_x.register_hook(lambda grad: print(f"sum_x梯度: {grad}"))
        sum_x.retain_grad()
        
        return sum_x

# Training
# Generate labels
a = torch.Tensor([0.5])
b = torch.Tensor([0.8])
x = torch.linspace(-1, 1, 10)
y = a*x + (0.1**0.5)*torch.randn_like(x)*(0.001) + b
inp = torch.linspace(-1, 1, 10)
foo = func_NN()
loss = torch.nn.MSELoss()
optim = torch.optim.Adam(foo.parameters(),lr=0.001)

t_l = []
for i in range(2):
    optim.zero_grad()
    l = loss(y, foo.forward(inp=inp))
    t_l.append(l.detach())
    print(f"\n第{i+1}次反向传播:")
    l.backward()
    optim.step()

关键说明

  1. register_hook:给单个张量绑定反向hook函数,反向传播时会自动传入该张量的梯度并执行hook逻辑。
  2. retain_grad():强制PyTorch保留非叶子节点的梯度(默认会释放以节省内存),否则hook无法捕获到梯度。
  3. 移除了原代码中绑定在模型上的register_full_backward_hook,因为它无法捕获中间张量的梯度。

运行修改后的代码,就能在每次反向传播时看到mul_x和sum_x的梯度值了。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.04 21:27:38