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

PyTorch中GAN多轮生成器训练反向传播报错问题求助

问题分析与解决方案

错误原因解析

  1. 第一个错误(RuntimeError: Trying to backward through the graph a second time):
    生成器训练循环结束后,你使用最后一次生成的fake张量训练判别器,但该fake的计算图已经在生成器最后一次lossG.backward()调用时被默认释放(retain_graph=False),判别器反向传播时需要追溯该计算图,因此触发错误。

  2. 第二个错误(RuntimeError: one of the variables needed for gradient computation has been modified by an inplace operation):
    给lossG.backward()添加retain_graph=True后,生成器的计算图被保留,但opt_gen.step()会原地更新生成器参数,导致计算图中记录的参数版本与当前参数版本不匹配,反向传播时触发版本冲突。

修正后的代码

核心思路:训练判别器时,重新生成不追踪梯度的fake张量,避免依赖生成器训练阶段的旧计算图,同时确保判别器反向传播仅更新自身参数。

for epoch in range(num_epochs):
    for batch_idx, (real, _) in enumerate(loader):
        real = real.view(-1, 784).to(device)
        batch_size = real.shape[0]

        # 训练生成器(gen_advantage次更新)
        for i in range(gen_advantage):
            noise = torch.randn(batch_size, z_dim).to(device)
            fake = gen(noise)
            output = disc(fake).view(-1)
            lossG = criterion(output, torch.ones_like(output))
            # 无需保留计算图,生成器训练后的计算图无需复用
            lossG.backward()
            opt_gen.step()
            gen.zero_grad()

        # 训练判别器(disc_advantage次更新)
        for i in range(disc_advantage):
            # 重新生成fake,用torch.no_grad()避免追踪生成器的梯度
            with torch.no_grad():
                noise = torch.randn(batch_size, z_dim).to(device)
                fake = gen(noise)
            # 计算真实样本损失
            disc_real = disc(real).view(-1)
            lossD_real = criterion(disc_real, torch.ones_like(disc_real))
            # 计算伪造样本损失
            disc_fake = disc(fake).view(-1)
            lossD_fake = criterion(disc_fake, torch.zeros_like(disc_fake))
            # 总损失
            lossD = (lossD_real + lossD_fake) * 0.5
            lossD.backward()
            opt_disc.step()
            disc.zero_grad()

额外说明

  • 训练判别器时重新生成fake是合理的:判别器需要学习区分当前生成器的最新输出,而非生成器训练过程中旧版本的输出。
  • 使用torch.no_grad()包裹生成fake的过程,既能避免不必要的梯度计算,也能彻底切断判别器反向传播与生成器计算图的关联,从根源解决计算图复用和参数版本冲突问题。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.13 01:20:00