如何在PyTorch嵌套神经网络中实现局部梯度更新以避免多次反向传播
问题
我在PyTorch中构建了两个嵌套的神经网络,针对输出计算两个不同的损失函数,对应不同参数的梯度更新。简单示例如下:
# 两个神经网络 A = nn.Linear(10,10) B = nn.Linear(10,1) # 虚拟输入 x = torch.rand(1,10, requires_grad=True) # 嵌套计算 y = B(A(x)) # 基于输出计算两个独立损失 Loss1 = f(y) Loss2 = g(y) # 基于两个损失反向传播 (Loss1+Loss2).backward()
我希望Loss1同时更新神经网络A和B的参数,但Loss2仅更新神经网络A的参数。我知道可以通过拆分两次反向传播步骤实现,示例如下:
# 两个神经网络 A = nn.Linear(10,10) B = nn.Linear(10,1) # 虚拟输入 x = torch.rand(1,10, requires_grad=True) # 嵌套计算 y = B(A(x)) # 计算第一个损失 Loss1 = f(y) # 基于第一个损失反向传播 Loss1.backward() # 关闭B的梯度计算 B.requires_grad_(False) # 重新进行嵌套计算 y = B(A(x)) # 计算第二个损失 Loss2 = g(y) # 基于第二个损失反向传播 Loss2.backward()
但我不喜欢这种方法,因为它需要重复执行嵌套神经网络的前向计算。我试过类似g(y).detach()的方式,但这会同时移除针对网络A的梯度。有没有办法标记第二个损失,使其不更新网络B的参数?
解决方案
你可以通过手动计算Loss2对A参数的梯度并累加的方式,避免重复前向传播,同时实现Loss2仅更新A的需求:
import torch import torch.nn as nn def f(y): # 示例Loss1,比如MSE到目标值0 return torch.mean(y**2) def g(y): # 示例Loss2,比如MAE到目标值1 return torch.mean(torch.abs(y - 1)) # 初始化网络 A = nn.Linear(10,10) B = nn.Linear(10,1) x = torch.rand(1,10, requires_grad=True) # 单次前向传播得到中间结果 a = A(x) y = B(a) # 处理Loss1:更新A和B的梯度 Loss1 = f(y) Loss1.backward(retain_graph=True) # retain_graph=True保留计算图,用于后续Loss2的梯度计算 # 处理Loss2:仅更新A的梯度 Loss2 = g(y) # 计算Loss2对A所有参数的梯度 grads_A = torch.autograd.grad(Loss2, A.parameters(), retain_graph=True) # 将梯度手动累加到A的参数.grad上 for param, grad in zip(A.parameters(), grads_A): if param.grad is None: param.grad = grad else: param.grad += grad # 后续即可正常执行优化步骤,比如: # optimizer = torch.optim.SGD(list(A.parameters()) + list(B.parameters()), lr=0.01) # optimizer.step()
另一种更简洁的方式是,在计算Loss2时,阻止梯度流向B的参数,同时保留A的梯度传递:
# 单次前向传播 a = A(x) with torch.no_grad(): # 用no_grad包裹B的计算,避免B的参数被跟踪梯度 y = B(a) # 这里需要将y和a关联起来,让Loss2的梯度能传到a y = y.detach() y.requires_grad = True # 现在计算Loss2并反向传播,梯度只会传到a,进而更新A的参数 Loss2 = g(y) Loss2.backward()
这种方式下,y被detach后重新开启requires_grad,此时y的梯度会反向传播到a,而不会涉及B的参数,因为B的计算是在no_grad环境中进行的,不会被记录到计算图里。
内容的提问来源于stack exchange,提问作者cdmath
相关产品推荐
相关产品推荐

