如何在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()
关键说明
register_hook:给单个张量绑定反向hook函数,反向传播时会自动传入该张量的梯度并执行hook逻辑。retain_grad():强制PyTorch保留非叶子节点的梯度(默认会释放以节省内存),否则hook无法捕获到梯度。- 移除了原代码中绑定在模型上的
register_full_backward_hook,因为它无法捕获中间张量的梯度。
运行修改后的代码,就能在每次反向传播时看到mul_x和sum_x的梯度值了。
内容的提问来源于stack exchange,提问作者Newbie
相关产品推荐
相关产品推荐

