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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.04 22:27:04