GAN模型训练出现Trying to backward through the graph报错如何解决?
报错原因分析
触发的报错信息如下:
RuntimeError: Trying to backward through the graph a second time (or directly access saved variables after they have already been freed). Saved intermediate values of the graph are freed when you call .backward() or autograd.grad(). Specify retain_graph=True if you need to backward through the graph a second time or if you need to access saved variables after calling backward.
你的代码存在两个核心问题导致报错:
- 更新生成器时复用了判别器之前对
global_output.detach()的输出output_disc计算损失。首先global_output已经做了detach操作,这个损失计算完全不会传导梯度到生成器,其次这个output_disc对应的计算图已经在error_fake.backward()执行时被释放,第二次对基于这个计算图的损失error_gen调用backward就会触发二次反向的报错。 - 额外逻辑问题:计算判别器的真实样本损失时,用的是原始未加mask的batch,但是生成器输入的是加了mask的batch,两者数据逻辑不匹配。
修复方案
修正后的训练代码如下:
for epoch in range(N_EPOCHS): for i, batch in enumerate(dataloader, 0): real_batch = batch.to(device) # 判别器训练 global_disc.zero_grad() # 计算真实样本损失 output_disc_real = global_disc(real_batch) error_target = loss(output_disc_real, torch.ones(output_disc_real.shape, device=device)) error_target.backward() # 生成假样本 masked_batch = apply_mask(real_batch) global_output, local_output = gen(masked_batch) # 计算假样本损失,detach切断生成器梯度,只更新判别器 output_disc_fake = global_disc(global_output.detach()) error_fake = loss(output_disc_fake, torch.zeros(output_disc_fake.shape, device=device)) error_fake.backward() # 更新判别器参数 optimizer_disc.step() # 生成器训练:重新将未detach的生成器输出传入判别器计算损失 gen.zero_grad() # 此处不做detach,梯度可以正常传回生成器 output_disc_gen = global_disc(global_output) error_gen = loss(output_disc_gen, torch.ones(output_disc_gen.shape, device=device)) error_gen.backward() # 更新生成器参数 optimizer_gen.step() break break
核心改动点:
- 训练生成器时,重新将没有detach的
global_output传入判别器得到新的输出,再计算生成器损失,既保证计算图完整未被释放,也能正确传导梯度到生成器 - 拆分了判别器真实/假样本的输出变量,避免变量覆盖带来的计算图混淆
- 统一了张量的device指定方式,避免跨设备错误
内容的提问来源于stack exchange,提问作者user2962197
相关产品推荐
相关产品推荐

