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

如何在包含两个优化器的训练循环中正确调用backward反向传播?

问题解决方法

核心问题:让训练流程正常运行的修改方案

你遇到的报错本质是计算图出现了非预期的跨迭代依赖,且RL损失的计算图意外包含了FA网络的节点,按以下两步修改即可解决,不需要加retain_graph=True:

  • 第一步:计算RL损失用到的奖励属于“外部反馈”,不需要梯度回传到FA网络,因此算奖励时对FA的输出做detach处理
  • 第二步:baseline是历史奖励的移动平均,不需要关联任何计算图,更新时对奖励做detach处理

修改后的训练循环代码如下:

inps = torch.tensor([[1.0]])
y = torch.tensor(10.0)

opt_RL = optim.Adam(net_RL.parameters())
opt_FA = optim.Adam(net_FA.parameters()) 

baseline = 0
baseline_lr = 0.1

epochs = 100

for _ in tqdm(range(epochs)):
    for inp in inps:
        with torch.no_grad():
            net_FA(inp)
               
        for layer in range(3):
            out_RL = net_RL(torch.tensor([1.0,2.0,3.0]))
            mu, std = out_RL
            dist = Normal(mu, std)
            update_values = dist.sample() 
            log_p = dist.log_prob(update_values).mean()

            # 修改点1:FA输出加detach,不让RL损失的计算图包含FA节点
            out = net_FA(inp).detach() 
            reward = -torch.square((y - out)) 
            # 修改点2:reward加detach,不让baseline关联计算图,避免跨迭代依赖
            baseline = (1 - baseline_lr) * baseline + baseline_lr * reward.detach()

            loss_RL = - (reward - baseline) * log_p            
            opt_RL.zero_grad()
            loss_RL.backward()
            opt_RL.step()            

            out = net_FA(inp) 
            loss_FA = torch.mean(torch.square(y - out)) 
            opt_FA.zero_grad()
            loss_FA.backward()
            opt_FA.step()

print("Mean: " + str(mu.detach().numpy()) + ", Goal: " + str(y))
print("Standard deviation: " + str(softplus(std).detach().numpy()) + ", Goal: 0ish")    

衍生问题解答

1. 为什么你当前场景会被提示需要加retain_graph=True?

PyTorch默认每次反向传播后释放当前计算图,你遇到提示的核心原因是你的代码无意中让计算图产生了跨迭代的依赖:你没有对reward做detach,导致用来计算下一轮损失的baseline始终关联上一轮的FA网络计算图,第一轮反向传播释放了该计算图后,第二轮反向传播需要用到已经被释放的节点,就会触发该报错。并不是正常场景下需要这个参数,是你的代码逻辑导致计算图没有按预期在每轮迭代独立。

2. 加了retain_graph=True后训练变慢的原因?

retain_graph=True的作用是反向传播后不释放当前计算图,你因为跨迭代依赖的存在,每一轮迭代的计算图都会和上一轮的计算图绑定,不会被释放,随着迭代次数增加,内存中会堆积大量的历史计算图节点,每次反向传播需要遍历的节点数越来越多,自然会出现训练速度越来越慢的情况,同时内存占用也会持续升高。


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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.24 21:54:04