摘要:你是否遇到过这样的情况:模型实际占用的显存很小(比如 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):

  1. 处理复杂图,申请了一块大显存。
  2. 释放后,显存里留下了一个大坑。
  3. 下一张图来了,需要一块中等大小的显存。
  4. PyTorch 发现那个大坑“形状不匹配”或“切分太麻烦”,于是重新向系统申请一块新的显存。
  5. 结果:旧的显存没还回去,新的显存又申请了。显存池(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. 总结

  1. nvidia-smi 高、实际占用低,通常是显存碎片化导致的,不是 BUG。
  2. torch.cuda.empty_cache() 可以释放未使用的缓存,让 nvidia-smi 数据变好看。
  3. 不要滥用:频繁调用会严重拖慢速度。
  4. 最佳场景:推理阶段处理完复杂数据后、或者 OOM 救急时使用。
Logo

更多推荐