【PyTorch】显存虚高?一文读懂 torch.cuda.empty_cache() 的原理与正确用法
摘要:你是否遇到过这样的情况:模型实际占用的显存很小(比如 2GB),但 nvidia-smi 却显示占用了 10GB 甚至更多?调用 torch.cuda.empty_cache() 后,显存占用瞬间暴跌。本文将揭示 PyTorch 显存管理的幕后机制,解释“显存碎片化”现象,并给出 empty_cache 的最佳实践建议。
1. 诡异的现象:显存去哪了?
在跑深度学习模型(特别是像 SAM 这种处理不定长输入的网络)时,经常会出现以下现象:
- 代码打印:
torch.cuda.memory_allocated()显示实际只用了 0.3GB。 - 系统监控:
nvidia-smi显示显存占用了 10GB,甚至逼近显卡上限。 - 神奇操作:执行一行
torch.cuda.empty_cache()。 - 结果:
nvidia-smi的占用瞬间从 10GB 掉到了 2GB 多。
这中间的 8GB 到底是什么?是内存泄漏吗?
2. 原理解析:PyTorch 的“囤积癖”
这并非内存泄漏,而是 PyTorch 为了速度而设计的一种缓存机制(Caching Allocator)。
2.1 显存管理的机制
向 GPU 申请(cudaMalloc)和释放(cudaFree)显存是非常耗时的操作。为了不让训练过程卡在申请显存上,PyTorch 采用了一种策略:
- 只借不还:当你删除了一个变量(
del tensor),PyTorch 不会立刻把这块显存还给操作系统(GPU驱动)。 - 放入缓存:它把这就这块空闲内存标记为“缓存(Reserved)”,留给下一个新的 Tensor 直接复用。
2.2 为什么会膨胀到 10GB?(显存碎片化)
如果你的输入数据大小是固定的(比如 ResNet 输入全是 224x224),缓存机制运行得很完美。
但如果输入数据忽大忽小(例如 SAM 处理复杂图像生成 100 个 Mask,下一张简单图生成 5 个 Mask):
- 处理复杂图,申请了一块大显存。
- 释放后,显存里留下了一个大坑。
- 下一张图来了,需要一块中等大小的显存。
- PyTorch 发现那个大坑“形状不匹配”或“切分太麻烦”,于是重新向系统申请一块新的显存。
- 结果:旧的显存没还回去,新的显存又申请了。显存池(Reserved)里充满了各种无法利用的“孔洞”(碎片),导致
nvidia-smi看到的占用越来越高。
2.3 empty_cache() 做了什么?
torch.cuda.empty_cache() 就像是一次强制垃圾回收:
它会检查 PyTorch 缓存池中所有当前未使用的内存块,并将它们真正地通过 cudaFree 归还给操作系统。
这就是为什么你的显存从 10GB 瞬间跌回 2GB(实际占用 + 必要的 Context 开销)。
3. 正确用法与避坑指南
既然 empty_cache() 这么好用,是不是应该一直用?绝对不是!
❌ 错误用法:放在训练循环里
# 千万不要这样做!
for i, (input, target) in enumerate(dataloader):
output = model(input)
loss = criterion(output, target)
loss.backward()
optimizer.step()
# ❌ 严重错误:每一步都清理
torch.cuda.empty_cache()
后果:训练速度会大幅下降。因为这一步会强制 CPU 和 GPU 同步,并且让 PyTorch 频繁地找驱动申请显存,由于显存分配的开销,你的 GPU 利用率会跌到谷底。
✅ 推荐用法 1:处理完重型任务后
在推理(Inference)阶段,特别是处理像 SAM 这种显存波动极大的模型时,可以在处理完一张(或一批)图片后调用。
# 推荐:按图片或 Batch 清理
for img in images:
# 1. 跑完一张可能产生大量碎片的复杂图
process_complex_image(img)
# 2. 手动清理,防止显存无限膨胀导致 OOM
torch.cuda.empty_cache()
✅ 推荐用法 2:OOM(显存溢出)异常捕获
作为一种兜底机制,当显存真的不够用报错时,尝试清理后重试。
try:
output = model(input)
except RuntimeError as e:
if "out of memory" in str(e):
print("检测到OOM,尝试清理缓存...")
torch.cuda.empty_cache()
output = model(input) # 重试
else:
raise e
✅ 推荐用法 3:调试显存泄漏
当你怀疑代码有内存泄漏时,在关键位置插入 empty_cache()。如果清理后显存依然持续上涨(Allocated 变大),那是真泄漏;如果清理后回落,那就是单纯的碎片化问题。
4. 进阶方案:从根源解决碎片化
如果你不想频繁调用 empty_cache() 牺牲性能,可以通过设置环境变量来优化 PyTorch 的显存分配策略,减少碎片产生。
在运行代码前设置:
# Linux
export PYTORCH_CUDA_ALLOC_CONF=max_split_size_mb:128
# Windows PowerShell
$env:PYTORCH_CUDA_ALLOC_CONF = "max_split_size_mb:128"
或者在代码最开头:
import os
os.environ["PYTORCH_CUDA_ALLOC_CONF"] = "max_split_size_mb:128"
import torch
这告诉 PyTorch:“如果内存块小于 128MB 就不拆分了”,这能显著减少细碎孔洞的产生,往往能达到不调用 empty_cache 也能控制显存占用的效果。
5. 总结
- nvidia-smi 高、实际占用低,通常是显存碎片化导致的,不是 BUG。
torch.cuda.empty_cache()可以释放未使用的缓存,让nvidia-smi数据变好看。- 不要滥用:频繁调用会严重拖慢速度。
- 最佳场景:推理阶段处理完复杂数据后、或者 OOM 救急时使用。
更多推荐

所有评论(0)