PyTorch中GAN多轮生成器训练反向传播报错问题求助
问题分析与解决方案
错误原因解析
第一个错误(
RuntimeError: Trying to backward through the graph a second time):
生成器训练循环结束后,你使用最后一次生成的fake张量训练判别器,但该fake的计算图已经在生成器最后一次lossG.backward()调用时被默认释放(retain_graph=False),判别器反向传播时需要追溯该计算图,因此触发错误。第二个错误(
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
相关产品推荐
相关产品推荐

