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
相关产品推荐
相关产品推荐

