memory_allocated()跳变是因cuda缓存机制,它只统计实际张量占用而不计缓存碎片;oom常由缓存无法拼合大块内存触发,需用memory_summary()查看缓存与碎片,配合empty_cache()、eval()模式、bn处理、梯度检查点及dataloader调优综合解决。

为什么memory_allocated()突然跳变却没真爆显存?
这不是显存真的“炸了”,而是PyTorch的CUDA缓存机制在作怪:memory_allocated()只统计当前张量实际占用,不包含缓存碎片。当新张量需要连续大块内存,而缓存里全是零散小块拼不出时,就直接OOM——哪怕nvidia-smi还显示几GB free。
别只盯着memory_allocated(),改用torch.cuda.memory_summary()看真实情况,它会明确列出cached memory和allocated memory,以及碎片占比。
- 训练前加
torch.cuda.memory_summary()打个底,对比每轮前后变化 - 发现
cached memory持续上涨、碎片率>30%,说明缓存没及时回收 -
torch.cuda.empty_cache()不是万能药,它只释放未被引用的缓存,不能解决batch过大或中间变量堆积问题
batch_size=1还OOM?大概率是nn.BatchNorm2d在捣鬼
BN层在model.train()模式下,每个batch都要更新running_mean和running_var,并为每个channel分配临时缓冲区。哪怕batch_size=1,高分辨率输入(如512×512)也会触发巨大显存开销。
- 推理阶段必须调用
model.eval(),否则BN继续统计更新,显存峰值翻倍 - 训练时若必须小batch,可冻结BN:遍历模型模块,对所有
nn.BatchNorm2d调用m.eval() - 或直接替换为
nn.InstanceNorm2d,它不依赖batch维度统计
loss.backward()才OOM?计算图留得太久
前向没问题、反向崩,说明中间激活值堆得太多。尤其常见于自定义loss、RNN展开、torch.cat()拼接大量张量,或者误写total_loss += loss(这会不断延长计算图)。
- 改用
total_loss = total_loss + loss,并在合适位置用loss.item()提取标量,切断梯度链 - 长序列任务启用梯度检查点:
torch.utils.checkpoint.checkpoint(model, x),显存可降30%–50% - 确认没在
torch.no_grad()外做推理——否则计算图照常构建,白占显存
DataLoader设num_workers>0反而更吃显存?
多进程本身不占GPU显存,但每个worker会预加载batch并常驻内存;若主进程GPU显存已近极限,worker的内存映射(尤其pin_memory=True时)可能触发系统级OOM,报错却是CUDA out of memory。
- 先试
num_workers=0,看是否稳定——这是最干净的baseline - 若必须多进程,关闭
pin_memory,或降低prefetch_factor(默认2,可设为1) - 避免在
DataLoader里做重图像变换(如resize+normalize),移到dataset__getitem__中统一处理,减少worker内存压力
显存问题从来不是单点故障,而是多个环节叠加泄漏的结果。最容易被忽略的是:BN状态、计算图生命周期、worker内存与GPU显存的耦合关系——这三者不动手查,光调batch_size只会反复踩坑。
Python免费学习笔记(深入):立即使用
在学习笔记中,你将探索 Python 的核心概念和高级技巧!











