PyTorch实现WGAN-GP时Autograd模块触发CUDA显存溢出问题如何解决?
WGAN-GP训练CUDA显存不足解决方案
本次报错为PyTorch运行时显存耗尽(OOM),可通过以下方法逐一调整解决:
- 优先降低批次大小(batch size)
报错显示当前仅缺64MB显存,可先尝试将batch size下调25%~50%,例如从64调整为32/16,是生效最快的修复方案。 - 修复损失反向传播逻辑,减少计算图冗余
现有代码对判别器损失做了3次独立反向传播(真实样本、生成样本、梯度惩罚各一次),会多存储2份中间计算图占用显存。可调整为合并所有损失后仅做1次反向传播,同时删除梯度惩罚计算中不必要的retain_graph=True参数:
对应调整# 判别器迭代内调整后的代码逻辑 D.zero_grad() images = data.__next__() if images.size()[0] != batch_size: continue images = images.to(device) z = torch.randn(batch_size, 100, 1, 1).to(device) # 计算真实样本损失 d_loss_real = D(images).mean(0).view(1) # 计算生成样本损失 fake_images = G(z) d_loss_fake = D(fake_images).mean(0).view(1) # 计算梯度惩罚 gradient_penalty = calculate_gradient_penalty(images.detach(), fake_images.detach()) # 合并损失后仅做一次反向传播 d_loss = d_loss_fake - d_loss_real + gradient_penalty d_loss.backward() Wasserstein_D = d_loss_real - d_loss_fake d_optimizer.step()calculate_gradient_penalty函数中autograd.grad的参数,删除不必要的retain_graph=True和allow_unused=True:grad = torch.autograd.grad( outputs=disc_interpolates, inputs=interpolates, grad_outputs=torch.ones_like(disc_interpolates), create_graph=True)[0] - 修复损失存储逻辑,避免冗余张量驻留显存
现有代码直接将GPU上带计算图的张量存入进度列表,迭代次数增加后会持续占用显存。修改为仅存储CPU侧的标量数值:# 原错误写法:d_progress.append(d_loss) # 调整后写法 d_progress.append(d_loss.detach().cpu().item()) d_fake_progress.append(d_loss_fake.detach().cpu().item()) d_real_progress.append(d_loss_real.detach().cpu().item()) penalty.append(gradient_penalty.detach().cpu().item()) g_progress.append(g_loss.detach().cpu().item()) - 启用混合精度训练
使用PyTorch自带的torch.cuda.amp模块启用半精度训练,可降低约40%~50%的显存占用,且对WGAN-GP训练效果影响极小。 - 清理显存碎片
每次判别器迭代结束后,删除不需要的中间张量并清空显存缓存:del fake_images, images, z, d_loss_real, d_loss_fake, gradient_penalty, d_loss torch.cuda.empty_cache() - 其他可选优化
- 适当降低判别器的模型大小,减少卷积层通道数或层数
- 用
nvidia-smi检查GPU上是否有其他无关进程占用显存,关闭无用进程释放空间
内容的提问来源于stack exchange,提问作者Tim
相关产品推荐
相关产品推荐

