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

