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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.30 08:33:10