PyTorch多分支模型中如何阻止Backbone通过proxy_target更新梯度?
PyTorch多分支模型梯度控制问题解答
你的修改思路核心方向是对的,但存在两处小问题需要修正,同时还有更简洁的实现方式:
现有修改的问题点
- 外层的
with torch.set_grad_enabled(True)完全冗余,训练模式下默认就是开启梯度计算的,除非你主动全局关闭过。 - 代码里写的
loss.backward()是笔误,你的总损失变量是total_loss,应该调用total_loss.backward()。
为什么修改能达到预期效果?
with torch.set_grad_enabled(False):包裹transform_to_targets的计算后,这段操作不会被纳入Autograd的计算图,生成的proxy_target会变成无梯度张量(requires_grad=False)。
反向传播时:
proxy_loss的梯度只会从proxy_output流向proxyModule和backbone(因为proxy_output是proxyModule基于backbone_output计算的,所以backbone会通过这个路径获得梯度)。backbone_loss的梯度直接流向backbone。- 最终总损失的梯度会同时更新
backbone和proxyModule,且backbone不会通过transform_to_targets的路径产生任何梯度,完全符合你的需求。
更简洁稳妥的实现方式
用detach()方法可以更直观地切断梯度流,效果和torch.set_grad_enabled(False)完全一致,代码更简洁:
class Net(nn.Module): def __init__(self): super(Net, self).__init__() self.backbone = Backbone() self.proxyModule = ProxyModule() def forward(self, x): backbone_output = self.backbone(x) # 用detach切断backbone_output到proxy_target的梯度流 proxy_target = transform_to_targets(backbone_output).detach() proxy_output = self.proxyModule(backbone_output) return backbone_output, proxy_target, proxy_output net = Net() x,y = get_some_data() optimizer.zero_grad() backbone_output, proxy_target, proxy_output = net(x) backbone_loss = Loss(backbone_output, y) proxy_loss = Loss(proxy_output, proxy_target) total_loss = backbone_loss + proxy_loss total_loss.backward() optimizer.step()
验证梯度是否符合预期的方法
可以在反向传播后检查参数梯度,确认逻辑正确:
total_loss.backward() # 检查backbone的参数是否有梯度(应该返回True) print(next(net.backbone.parameters()).grad is not None) # 可以单独注释掉proxy_loss,只计算backbone_loss的梯度,对比和之前的梯度是否一致,确认proxy_target的生成没有额外贡献梯度
内容的提问来源于stack exchange,提问作者Ufuk Can Bicici
相关产品推荐
相关产品推荐

