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

batch_size=1时GPU显存递增致OOM,如何训练完整数据集?

解决单Batch训练显存递增OOM问题的实用方案
  • 手动触发显存回收:每次Batch训练完成后,在梯度更新步骤之后强制释放未使用的显存并回收垃圾。代码示例:

    import torch
    import gc
    
    # 训练循环内,每个batch的末尾
    loss.backward()
    optimizer.step()
    # 释放显存+垃圾回收
    torch.cuda.empty_cache()
    gc.collect()
    

    这能清理掉训练过程中未被自动释放的临时张量占用的显存。

  • 定位显存泄漏点:用torch.cuda.memory_summary(device=None, abbreviated=False)在每个Batch前后打印显存使用明细,对比找出持续增长的张量来源。重点排查模型中是否有循环创建的持久化张量、未被正确释放的中间变量,或是某些自定义模块的缓存机制。

  • 关闭无关梯度计算:确保只有训练相关的参数开启梯度,对于验证、测试环节或模型中固定参数的分支,用torch.no_grad()包裹前向传播,避免不必要的梯度张量累积。比如:

    with torch.no_grad():
        val_output = model(val_input)
    
  • 检查数据加载流程:排查DataLoader是否存在内存泄漏,比如是否在每个Epoch重复加载数据时没有释放旧的数据集实例,或是预处理函数中创建了未被回收的大张量。可以尝试每个Epoch结束后重新初始化DataLoader,或简化预处理逻辑。

  • 优化CUDA算子缓存:关闭torch.backends.cudnn.benchmark(设为False),避免cuDNN为不同输入尺寸缓存过多卷积算子,减少显存占用。如果模型有大尺寸的中间激活,也可以用torch.utils.checkpoint.checkpoint对部分层做梯度检查点,以计算量换显存空间。

内容的提问来源于stack exchange,提问作者mmx

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.28 15:22:10