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

PyTorch训练SRGAN报二次反向传播计算图错误如何调试

报错根因

该错误由训练逻辑错误导致,对应代码问题如下:

  • 训练生成器时复用了判别器阶段前向输出的disc_gen张量,该张量关联的计算图在执行loss_disc.backward()后就被PyTorch自动释放,再次基于它计算损失反向传播就会触发二次遍历计算图的报错。同时判别器阶段对gen_image做了detach()操作,截断了和生成器参数的关联,即使计算图存在也无法把梯度回传给生成器。
  • 优化器执行顺序错误:opt_disc.step()被放到了生成器反向传播步骤之后,会导致生成器反向传播产生的梯度错误更新判别器参数。
  • 运算符使用错误:Python中^是按位异或运算符,不是幂运算,代码里10^-3的计算结果为-9,会导致对抗损失权重完全异常。
  • 显存泄漏问题:损失累加时直接保存带计算图的张量,会导致计算图无法持续释放,训练过程中显存占用会持续上涨。
  • 冗余与格式错误:每个batch内重复实例化BCEWithLogitsLoss浪费算力;打印损失时使用整数格式符%d,无法正确显示浮点型损失值。
  • 变量命名混淆:判别器阶段真实样本、生成样本的损失变量名写反,提升后续调试成本。
修复逻辑

GAN训练需要严格按顺序执行判别器、生成器的更新流程:

  1. 训练判别器时,先对真实样本计算损失,再将生成器输出的图片做detach()截断梯度后喂给判别器,计算生成样本损失,组合后反向传播,立刻执行判别器参数更新,完成判别器训练步骤。
  2. 训练生成器时,不能复用判别器阶段的前向结果,需要将未做detach()的生成图片重新喂给判别器做前向传播,再依次计算内容损失、对抗损失,组合为生成器总损失后反向传播,执行生成器参数更新。
修复后可运行代码
gen_model = Generator().to(device, non_blocking=True)
disc_model  = Discriminator().to(device, non_blocking=True)
opt_gen = optim.Adam(gen_model.parameters(), lr=0.01)
opt_disc = optim.Adam(disc_model.parameters(), lr=0.01)
# 损失函数实例化放到循环外,避免重复创建
bce_loss = nn.BCEWithLogitsLoss()

def train_model(gen, disc):
  for epoch in range(20):
    run_loss_disc = 0.0
    run_loss_gen = 0.0
    for data in train:
      low_res, high_res = (
          data[0].to(device, non_blocking=True, dtype=torch.float).permute(0, 3, 1, 2),
          data[1].to(device, non_blocking=True, dtype=torch.float).permute(0, 3, 1, 2)
      )
      # ---------------- 训练判别器 ----------------
      opt_disc.zero_grad()
      # 真实样本判别损失
      disc_real = disc(high_res)
      loss_real = bce_loss(disc_real, torch.ones_like(disc_real))
      # 生成样本判别损失,生成图detach截断生成器梯度
      gen_image = gen(low_res)
      disc_fake = disc(gen_image.detach())
      loss_fake = bce_loss(disc_fake, torch.zeros_like(disc_fake))
      
      loss_disc = loss_real + loss_fake
      loss_disc.backward()
      opt_disc.step() # 判别器反向传播后立刻更新参数
      run_loss_disc += loss_disc.item()

      # ---------------- 训练生成器 ----------------
      opt_gen.zero_grad()
      # 重新前向计算判别器结果,不复用之前的计算值,gen_image不做detach保留计算图
      disc_gen_output = disc(gen_image)
      # 注意vgg_loss输入顺序要和你定义的函数匹配,一般是预测值在前、真实值在后
      cont_loss = vgg_loss(gen_image, high_res)
      adv_loss = 1e-3 * bce_loss(disc_gen_output, torch.ones_like(disc_gen_output))
      gen_loss = cont_loss + adv_loss
      gen_loss.backward()
      opt_gen.step()
      run_loss_gen += gen_loss.item()

    # 打印epoch平均损失,使用浮点格式
    avg_disc_loss = run_loss_disc / len(train)
    avg_gen_loss = run_loss_gen / len(train)
    print(f"Epoch {epoch+1} | 判别器平均损失: {avg_disc_loss:.4f} | 生成器平均损失: {avg_gen_loss:.4f}")

train_model(gen_model, disc_model)

注意:如果你的vgg_loss定义时要求输入顺序为真实值在前、生成值在后,调整vgg_loss的入参顺序即可。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.29 17:01:04