我猜你也遇到过这种场景:模型调参调得好好的,训练到凌晨三点,突然蹦出一行 CUDA out of memory。上 nvidia-smi 一看,显存明明还有十几 G,PyTorch 却报分配失败。这可能是你第一次意识到,PyTorch 的 GPU 内存管理远不是“装得下就装、装不下就爆”那么简单。
这篇文章围绕 PyTorch GPU 内存优化,把我这几年踩过的坑、用过的策略和完整的排查思路整理了一遍。既适合刚入门、准备跑通第一个训练脚本的同学,也适合已经在微调大模型、天天跟 OOM 打交道的工程师。内容不会只停留在“调小 batch size、用 AMP”这种层面,而是会解释为什么这些方法有效、什么时候该用哪一种、以及组合使用时怎么排序最划算。
1. 先搞清楚显存到底被谁吃了:PyTorch缓存分配器与OOM的三种形态
1.1 为什么 nvidia-smi 显示的显存和你的 tensor 对不上
很多人第一次排查 OOM,都会打开 nvidia-smi 看当前显存占用,发现明明没有多少进程在跑,可用显存却比预期少一大截。这是正常的,因为 PyTorch 的 CUDA 后端默认使用 缓存分配器(Caching Allocator),它的工作方式很像操作系统的内存池:不是每次 torch.Tensor 分配都直接调用 cudaMalloc,而是先向驱动申请一大块显存,然后在内部按块分配。
这样做的目的是省掉频繁 cudaMalloc / cudaFree 的系统调用开销,因为这两个操作在训练循环里如果每步都做,会严重拖慢速度。但也带来了一个副作用:即使你的 Python 代码里删掉了一些 tensor,分配器也不会立刻把显存还给驱动,而是留在自己的缓存池里,等待下一个分配请求复用。
所以 nvidia-smi 里显示的是“进程从驱动那里申请到的显存总量”,这个数字通常会大于“当前实际活着的数据占用的显存”。如果只看 nvidia-smi 就来判断程序到底用掉多少显存,很容易被误导。
1.2 OOM 并不都是“显存真的不够”
我整理了平时遇到最多的三种 OOM 形态,它们的表象都是程序报错,但根因和解决手段完全不同:
| 形态 | 本质 | 典型表现 |
|---|---|---|
| 硬性不足 | 峰值显存超过 GPU 物理容量 | 任何优化都失效,只能用更小的模型或更小的 batch |
| 碎片化 | 总空闲量足够,但没有连续块 | 训练到某一步突然 OOM,重启后可能又能跑一会 |
| 缓存膨胀 | 分配器持有过多空闲缓存,挤压可用空间 | 多进程共享 GPU 时,一个进程的缓存影响另一个进程 |
碎片化是最容易让人困惑的。比如一块 24G 的卡,nvidia-smi 显示还有 14G 空闲,但你的程序尝试分配一个 8G 的连续空间时却失败了。原因是缓存池里的小块散落各处,没有一个连续区域能满足请求。PyTorch 2.x 提供了 PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True 来缓解这种问题,后面我会具体讲。
1.3 训练时显存大头到底在哪
很多人以为模型参数量最大,所以显存主要被模型占掉。但你跑一次正常训练就会发现,中间激活值(activation)往往是最大的显存消耗者,尤其是在大 batch、长序列、高分辨率输入的场景下。
一个典型的自动微分训练过程,显存占用大体分为四块:
- 模型权重(weights)
- 梯度(gradients)
- 优化器状态(optimizer states)
- 前向传播过程中暂存的中间激活值(用于反向传播计算梯度)
以 1B 参数的模型为例,如果用 FP32 训练:
- 权重:4GB
- 梯度:4GB
- Adam 优化器状态:8GB(Adam 需要保存一阶动量和二阶动量,每个参数量要额外占两份)
- 中间激活值:取决于输入形状和网络结构,可能轻松超过 10GB
所以你会发现,光是没有优化器状态的推理场景,1B 模型 FP32 只需要 4GB 左右;但一旦进入训练,20GB 打底很常见。理解了这张账本,后面讲优化策略才有意义——你至少要知道自己在砍哪块开销,以及砍完会不会有副作用。
需要模型API调用? 免费领10W Token,多模型网关一键接入 Claude、DeepSeek 等主流模型。
2. 动手优化前先量化:三套显存观测手段与基线记录方法
2.1 用 torch.cuda 自带 API 定位峰值
优化显存的第一步不是改代码,而是先记录基线。我建议在任何项目启动时,都先写几行显存监测代码,把关键阶段的占用打印出来。PyTorch 自带的 API 足够完成大部分工作:
python复制import torch
def print_gpu_memory(tag=""):
allocated = torch.cuda.memory_allocated() / 1024**3
reserved = torch.cuda.memory_reserved() / 1024**3
max_allocated = torch.cuda.max_memory_allocated() / 1024**3
print(f"{tag}: allocated={allocated:.2f}GB, reserved={reserved:.2f}GB, max={max_allocated:.2f}GB")
print_gpu_memory("init")
model = MyModel().cuda()
print_gpu_memory("after model")
for step, batch in enumerate(dataloader):
batch = {k: v.cuda() for k, v in batch.items()}
loss = model(**batch).loss
loss.backward()
print_gpu_memory(f"after backward step {step}")
torch.cuda.memory_allocated() 返回当前被 tensor 实际占用的显存,torch.cuda.memory_reserved() 返回缓存分配器从驱动手里拿到的总量,torch.cuda.max_memory_allocated() 返回从程序开始到当前时刻的峰值占用。
这三个值一起看,能帮你快速判断一个关键问题:你的程序是“真的需要这么多显存”,还是“分配器预留太多但实际用到的很少”。如果 max_allocated 跟硬件容量很接近,那就说明峰值的硬需求偏高,你要从算法层面去压;如果 allocated 不高但 reserved 很高,说明缓存池里躺着大量空闲块,可以考虑 empty_cache() 或在分配参数上做文章。
2.2 用系统级工具看清多进程和驱动层面的真实占用
PyTorch 的 API 只能看到自己进程内部的显存情况。如果一张卡上同时跑着多个进程,或者你想确认是否有残留进程占着显存,还是得靠系统级工具。
nvidia-smi 本身就能查看进程级显存占用:
bash复制nvidia-smi --query-compute-apps=pid,used_memory,process_name --format=csv
这个命令会把每张 GPU 上正在运行的进程和显存占用列出来,是排查看有没有“僵尸进程”占显存的最好方式。
如果你希望实时监控一整段时间的变化,用 nvidia-smi dmon:
bash复制nvidia-smi dmon -i 0 -d 1
它会每秒刷新一次 GPU 利用率、显存读写、温度等指标。我自己习惯在训练到峰值阶段时开一个 dmon 窗口,看显存占用曲线有没有突然暴涨。
另外,gpustat 是社区里常用的封装,提供更友好的终端显示。但它本质还是基于 nvidia-smi,所以如果机器上没有权限装,也不影响排查。
2.3 逐层记录激活值,找到真正的“大头”
API 和系统工具能告诉你显存占用在哪个阶段飙升,但没法告诉你具体是哪一层网络结构导致的。想精准定位,我一般会给模型注册 forward hook,打印每一层的输出 tensor 大小和累积显存变化:
python复制def register_memory_hooks(model):
hooks = []
def make_hook(name):
def hook_fn(module, input, output):
if isinstance(output, torch.Tensor):
print(f"{name} output shape: {tuple(output.shape)}, "
f"dtype: {output.dtype}, bytes: {output.numel() * output.element_size() / 1024**3:.3f}GB")
return hook_fn
for name, module in model.named_modules():
hooks.append(module.register_forward_hook(make_hook(name)))
return hooks
这不是一个高精度工具,但用来定性分析很有价值。你经常会发现,某几个 Transformer Block 的中间输出占了显存的一半以上,或者某个多头注意力层因为序列长度变长导致激活爆炸。找到大头之后,再去选择用激活检查点还是降精度,思路就清晰了。
我在实际项目中固定下来的流程是:新拿到一个模型,先把上面的 hook 打一遍,把每层的输出形状、显存占用、梯度是否保存储进一个日志文件,形成“显存基线表”。后续每次改模型结构,都拿新快照和基线对比,而不是靠感觉猜。
3. 训练场景的核心三连:梯度累积、混合精度与激活检查点
3.1 梯度累积:不降 batch 也能等效增大 batch
当显存不够时,很多人第一反应是把 batch size 调小。这确实立竿见影,因为它直接减少了单次前向传播需要保存的激活值。但 batch size 调得太小会导致梯度噪声变大,收敛不稳定,某些情况下还会影响 BatchNorm 的统计量。
梯度累积(Gradient Accumulation)是典型的“换时间换空间”方案:前向传播仍然用小 batch,但每跑几个小 batch 才更新一次参数,让梯度在反向传播后累积起来,等效于一个更大的 batch。
python复制accumulation_steps = 4
optimizer.zero_grad(set_to_none=True)
for step, batch in enumerate(dataloader):
loss = model(batch) / accumulation_steps
loss.backward()
if (step + 1) % accumulation_steps == 0:
optimizer.step()
optimizer.zero_grad(set_to_none=True)
注意两点:
第一,每次 loss 一定要除以 accumulation_steps,否则梯度会是原来大 batch 的 N 倍,学习率等效被放大了。很多人踩过这个坑,训练直接发散。
第二,如果你用了 BatchNorm,梯度累积和真正的大 batch 并不完全等价。BatchNorm 在小 batch 上统计的均值和方差会更不稳定。如果模型依赖 BatchNorm,可以考虑用 SyncBN 或者在评估时用累计的统计量来缓解。
梯度累积解决的是“batch 太大”的问题,对激活值占用的压降非常直接。但它不会减少模型参数、梯度和优化器状态的占用,所以当模型本身都放不下时,靠它没用。
3.2 AMP 混合精度:显存减半收益下隐藏的 loss scale 坑
混合精度(AMP)是 PyTorch 里性价比最高的一招,代码改动量小,收益却非常明显。核心思路是:前向传播用 FP16 或 BF16 计算,反向传播得到的梯度也用半精度,但优化器仍然维护一份 FP32 的 master weight,更新参数时用 FP32。
一句话概括:计算半精度,更新全精度。
python复制scaler = torch.cuda.amp.GradScaler()
for batch in dataloader:
optimizer.zero_grad(set_to_none=True)
with torch.autocast(device_type="cuda", dtype=torch.float16):
loss = model(batch)
scaler.scale(loss).backward()
scaler.step(optimizer)
scaler.update()
FP16 能把绝大多数 tensor 的内存占用砍半,同时因为半精度张量在计算时带宽压力更小,训练速度往往也会提升。但 FP16 的坑在于动态范围小,梯度在反向传播中可能溢出为 inf 或者下溢为 0,所以需要 GradScaler 来做 loss scaling,在反向传播前把 loss 放大,算完梯度后再缩回去。
如果你用的是 Ampere 及之后架构的 GPU(比如 A100、RTX 30/40 系),我更推荐直接用 BF16,因为它的指数范围和 FP32 一样,基本不会出现溢出问题,虽然尾数精度低一些,但绝大多数训练场景下训练稳定性比 FP16 好。代码只改动一个参数:
python复制with torch.autocast(device_type="cuda", dtype=torch.bfloat16):
loss = model(batch)
AMP 的主要收益在激活值和梯度这两块,而优化器状态依然占着 FP32 的内存。所以别指望它能把显存降到原来的四分之一,降一半左右是合理的预期。
3.3 Activation Checkpointing:用算力换显存
如果 AMP 之后显存还是不够,或者你想继续加大 batch,那就要考虑激活检查点(Activation Checkpointing)了,PyTorch 里一般直接叫 torch.utils.checkpoint。
它的原理非常朴素:默认情况下,PyTorch 会在前向传播时把所有中间激活值都保存下来,供反向传播使用。激活检查点则是“选择性放弃保存”,让某些层的前向传播结果不缓存,在反向传播需要梯度的那一刻,再重新计算一次前向传播,拿到中间结果。
torch.utils.checkpoint.checkpoint 的用法是在模型前向代码里包一层:
python复制from torch.utils.checkpoint import checkpoint
def forward(self, x):
x = checkpoint(self.layer1, x)
x = checkpoint(self.layer2, x)
return x
这种方式能把激活值占用从“所有层的中间结果”降到“只保存每个检查点之间的边”,省出的显存相当可观。代价是反向传播时要多算一遍前向,训练总时间可能增加 20%-50%,具体取决于你的网络层计算量。
使用时有几个限制要注意:
- 被包裹的函数不能改变输入(不能 in-place 修改)。
- 如果层里有 BatchNorm,因为前向被重复执行,BatchNorm 的运行统计量更新会被影响,需要小心。通常把 BatchNorm 层留在检查点外面更安全。
- Dropout 被重复计算会有不同的随机 mask,这个问题可以通过传入相同的 seed 或者把随机状态管理好来规避,但初学者最好先避开这种层。
激活值占用大头一般在 Transformer 的 attention 和 FFN 层,所以对大模型来说,把整个 decoder layer 包成 checkpoint 是很常见的做法。
3.4 三招组合使用的顺序建议
梯度累积、AMP、激活检查点这三招不是互斥的,但组合使用时有个推荐顺序,能让你每一步都尽量少改代码、少损失训练质量:
- 先开 AMP / BF16。改动最小,收益最大,对训练速度通常还有正面影响。
- 如果显存还紧,把 batch size 降下来,用梯度累积补足等效 batch size。
- 最后再用 activation checkpointing 去压激活值,同时清楚它会让训练变慢。
这个顺序背后的逻辑是:先做不影响模型收敛质量的改动,把更“伤”的优化手段留到最后。AMP 只要处理好 loss scale,对收敛影响很小;梯度累积对优化动态的扰动相对可控;而 checkpoint 直接增加计算量,且可能会影响某些层的运行方式。
我见过一些人一上来就把 checkpoint 全包了,结果训练速度慢得没法接受,最后排查发现其实开个 AMP 就能解决。先量化,再动刀,能少走很多弯路。
4. 零敲碎打但不积少成多:容易被忽略的显存细节
4.1 梯度置空:zero_grad(set_to_none=True) 到底省在哪
训练循环里那句 optimizer.zero_grad() 是很多人不会多看一眼的代码,但它其实有一个很实用的优化开关:set_to_none=True。
python复制optimizer.zero_grad(set_to_none=True)
默认情况下,zero_grad() 是把梯度张量清零(写入 0 值),这需要一次遍历和写操作。而 set_to_none=True 是直接把梯度张量置为 None,不保留原来的数据缓冲区。这样一来,原本存放梯度的那段显存可以被分配器回收复用,给下一个 batch 的激活值或其他临时张量使用。
实际效果有两层:
- 少了清空梯度的写操作,训练循环会快一点点。
- 梯度缓冲区不再被长期按住,显存使用局部峰值可能会下降。
这个改动基本没有副作用,唯一要注意的是如果代码里手动使用了 param.grad,要记得在需要时先初始化,因为置 None 后 param.grad 是 None 而不是全零张量。
4.2 del 显式释放 + empty_cache 的正确使用姿势
很多人会把 torch.cuda.empty_cache() 当成“清理显存”的万能药,在训练循环里每隔几步就调一次。这其实是不推荐的。
正确理解是:empty_cache() 只会释放缓存分配器当前持有但未被使用的空闲块,不会释放还在被 tensor 引用的显存。如果你在训练循环里频繁调用它,反而会让缓存池形同虚设,每次释放后再分配都要向驱动重新申请,拖慢训练速度。
它适合用在这些场景:
- 训练结束后、启动验证或推理前。
- 捕获到 OOM 异常后,释放掉已经不再需要的缓存,再尝试降低 batch 重跑。
- 一段长期训练中间需要切换数据集或模型结构时。
在代码里,配合 del 使用:
python复制output = model(batch)
loss = loss_fn(output, target)
# 不再需要的中间变量
del output, batch
torch.cuda.empty_cache()
注意,del 是删除 Python 引用,如果其他地方还持有这个 tensor 的引用,del 并不会真正释放底层内存。最稳妥的办法是用完临时中间结果就 del,再加一个 empty_cache(),而不是平时滥调。
4.3 pin_memory 与 non_blocking:隐藏的显存吞吐优化
严格来说,pin_memory 不直接减少显存占用,但它能明显减少数据从 CPU 拷贝到 GPU 的时间。如果你把数据加载的瓶颈降下来,就能在同样的训练时间里跑更多的迭代,间接提高了显存使用效率。
做法是在 DataLoader 里开启 pin_memory=True,然后把数据搬到 GPU 时用 non_blocking=True:
python复制dataloader = DataLoader(dataset, batch_size=32, num_workers=4, pin_memory=True)
for batch in dataloader:
batch = {k: v.cuda(non_blocking=True) for k, v in batch.items()}
pin_memory 会把 CPU 数据放到页锁定内存里,这样 GPU 可以直接通过 DMA 访问,不用通过可分页内存中转。non_blocking=True 则让拷贝操作异步化,可以和计算重叠。
代价是页锁定内存会占用一部分系统物理内存,所以如果机器 CPU 内存本身很紧张,pin_memory 的开销需要评估一下。我一般在 CPU 内存充足时打开,如果机器只有 16G 内存,又跑着大模型,反而会关掉。
4.4 优化器状态:从 Adam 到 8-bit 优化器的取舍
前面那张显存账本里已经提到,Adam 优化器会在模型参数和梯度之外,额外保存两份优化器状态。对于大模型来说,这部分开销非常可观。所以优化器状态是“零碎优化”里最大的一个可压缩项。
如果模型参数是 FP32,Adam 的一阶动量和二阶动量也会是 FP32,等于每个参数要多占 8 字节。换成 SGD + momentum,每个参数只多占 4 字节。如果你用 SGD 能达到差不多的收敛效果,显存压力会小一截。
如果非用 Adam 不可,可以试试低精度优化器状态方案。比如 bitsandbytes 的 8-bit Adam:
python复制import bitsandbytes as bnb
optimizer = bnb.optim.Adam8bit(model.parameters(), lr=1e-3)
它把优化器状态压缩到 8-bit,Adam 状态的显存占用能从 8GB 降到 2GB 左右(以 1B 参数为例),训练效果通常接近标准 Adam。代价是这个库需要额外安装,而且部分环境可能存在算子兼容性问题,建议先在单卡小规模上验证。
还有一个思路是换用 Adafactor,它不保存完整的二阶动量,而是用一个近似值,额外显存占用更低。对 Transformer 类模型收敛效果整体不错,但并非所有场景都适用,最好做对比实验再决定。
4.5 推理阶段:torch.no_grad 和 inference_mode 不要混为一谈
训练显存优化和推理显存优化是两个方向,但很多人会把它们混在一起。推理阶段没有反向传播,不需要保存中间激活值,也不需要梯度,所以显存占用低得多。
在推理代码里,最基础的是 torch.no_grad()。更深一层是 torch.inference_mode(),它比 no_grad 做了更激进的优化,完全禁用 autograd 的部分跟踪机制,在推理循环里通常更快、更省内存:
python复制@torch.inference_mode()
def predict(model, x):
return model(x)
另外,如果模型已经训练完,建议把参数 requires_grad 设成 False:
python复制for param in model.parameters():
param.requires_grad = False
这样即使某些推理代码不小心走到了 autograd 分支,也不会因为“需要计算梯度”而保留不必要的中间节点。
5. 分布式与大模型场景下的显存账本:从DDP到FSDP的选择逻辑
5.1 DDP 为什么没有减少单卡显存
单卡显存不够时,很多人第一反应是上多卡,用 DistributedDataParallel(DDP)。但 DDP 的并行方式是数据并行:每张卡都保存一份完整的模型副本、梯度缓冲区和优化器状态,只是在反向传播时对梯度做一次 All-Reduce 同步。
所以 DDP 并不会让单卡显存需求下降。比如一个模型单卡训练需要 20GB,换 4 卡 DDP,每张卡依然需要 20GB 才能跑起来。它带来的是等效 batch size 变大、训练吞吐提升,而不是单卡内存压力的缓解。
使用 DDP 时有个容易忽略的点:梯度 All-Reduce 会产生额外的通信缓冲区,NCCL 可能会在大 batch 下额外申请一部分显存。所以 DDP 多卡训练时,单卡显存占用反而可能比单卡多出 1-2GB。
5.2 FSDP 和 ZeRO:把显存分摊到卡上
如果模型本身的参数量太大,单卡放不下一个完整副本,那就要考虑模型状态的切分了。PyTorch 原生提供了 torch.distributed.fsdp.FullyShardedDataParallel(FSDP),核心技术来源于 DeepSpeed 的 ZeRO 系列思想。
ZeRO 分几个阶段:
- ZeRO-1:把优化器状态分片到各卡
- ZeRO-2:优化器状态 + 梯度分片
- ZeRO-3:优化器状态 + 梯度 + 模型参数分片
FSDP 相当于把 ZeRO-3 的能力集成进了 PyTorch。在 FSDP 下,每张卡只保存一部分模型分片,需要用到某一层参数时,再从其他卡拉取。因此单卡显存占用会随卡数近似线性下降。
python复制from torch.distributed.fsdp import FullyShardedDataParallel as FSDP
model = FSDP(model)
但天下没有免费的午餐:参数分片意味着在 forward/backward 时需要频繁通信来获取参数,通信开销远高于 DDP。所以 FSDP 通常适用在“模型大到单卡放不下”的场景,而不是小模型为了提速。
使用 FSDP 时还有一个非常实用的开关:CPUOffload。它可以把优化器状态或参数卸载到 CPU 内存,进一步压降 GPU 显存占用,代价是 CPU 与 GPU 之间的传输会拖慢速度。
python复制from torch.distributed.fsdp import CPUOffload
model = FSDP(model, cpu_offload=CPUOffload(offload_params=True))
这个方案很适合“GPU 显存不够但 CPU 内存比较充裕”的开发机。
5.3 加载大模型的 CPU 内存峰值陷阱
大模型加载到 GPU 时,有一个很隐蔽的内存峰值问题。很多人会写:
python复制model = ModelClass.from_pretrained("some-large-model")
model = model.cuda()
这行代码在 from_pretrained 阶段会把完整的模型权重加载到 CPU 内存里,然后 .cuda() 再拷贝到 GPU。如果模型是 70B 参数,光是 FP32 权重就有 280GB,绝大多数机器 CPU 内存直接爆掉。
HuggingFace Transformers 的 from_pretrained 提供了 low_cpu_mem_usage=True 参数,能在加载时尽量减少 CPU 内存占用,直接分阶段放到 GPU:
python复制model = ModelClass.from_pretrained("some-large-model", low_cpu_mem_usage=True, torch_dtype=torch.float16)
model = model.to("cuda")
更细的控制是用 accelerate 库,加载到 meta device,再按需把参数分配到 GPU 或 CPU:
python复制with torch.device("meta"):
model = ModelClass.from_config(config)
后面再通过 from_pretrained 配合 device_map="auto" 把不同层自动分配到可用的 GPU/CPU/磁盘上。这些工具本质上都是在优化“加载过程”的内存峰值,避免在真正开始训练或推理之前就 OOM。
5.4 自回归推理的 KV Cache 不容忽视
如果你跑的是语言模型推理,还有一个显存消耗随着生成长度不断增长的东西:KV Cache。它缓存了注意力计算中的 Key 和 Value,避免每一步都重新计算历史 token 的注意力值。
KV Cache 的大小大约等于:
2(K 和 V) × 层数 × 头数 × 每头维度 × 序列长度 × 精度字节数
假设一个 7B 模型,batch size 为 8,生成长度为 2048,KV Cache 可能轻松占掉 8-12GB 显存。所以在大模型推理场景里,显存优化不只是模型权重的压缩问题。
常用的优化方向包括:
- 降低精度,比如把 KV Cache 从 FP16 量化到 INT8。
- 使用 PagedAttention 之类的显存管理方案,把 KV Cache 分页管理,减少碎片和预留浪费。vLLM 的核心优化就在这一点。
- 控制最大生成长度和 batch size,避免 KV Cache 无限增长。
6. 一次OOM排除的完整链路:从报错信息到根因定位
6.1 先看懂 CUDA out of memory 报错里的信息
PyTorch 的 OOM 报错比很多人想象中更有信息量。完整报错一般长这样:
code复制RuntimeError: CUDA out of memory. Tried to allocate 512.00 MiB (GPU 0; 23.69 GiB total capacity; 19.82 GiB already allocated; 3.30 GiB free; 18.42 GiB reserved in total by PyTorch)
这段信息里最关键的是:
already allocated:当前被 tensor 实际占用的显存。free:分配器视角里的空闲显存。reserved in total by PyTorch:分配器从驱动手里申请的总量。
你会发现 reserved 通常比 already allocated 大,这中间的差值就是缓存池里的空闲块。如果报错时 free 很小,说明硬性峰值快到了;如果 free 看着还行但分配失败,那大概率是碎片化。
拿到报错信息后,我一般先记录当时的 max_memory_allocated,再结合前面提到的逐层 hook 看哪个阶段涨得最猛。
6.2 多进程、残留进程和 WSL 环境带来的假象
很多 OOM 不是你的代码有问题,而是环境里还有其他进程占着显存。最典型的是上次训练没退出,残留的 Python 进程还在持卡。
排查命令:
bash复制nvidia-smi --query-compute-apps=pid,used_memory,process_name --format=csv
看到不需要的进程,直接 kill -9 <pid>。
如果你在 WSL2 里跑 PyTorch,有时会看到这样的系统错误:
code复制failed to initialize nvml: gpu access blocked by the operating system
这通常不是显存优化问题,而是 WSL 里的 GPU 访问没有正确建立。WSL2 本身不需要在 Linux 内单独安装 NVIDIA 驱动,它依赖 Windows 侧安装的支持 WSL 的 GPU 驱动。如果 Windows 驱动版本太老或没装好,WSL 里就无法初始化和 GPU 通信。解决办法是去 Windows 侧更新 NVIDIA 驱动,然后重启 WSL,而不是在 Linux 环境里折腾驱动。
另外,如果机器有多张 GPU,一定要关注 CUDA_VISIBLE_DEVICES 环境变量,否则 PyTorch 默认只看到 cuda:0,你可能以为自己在用某张卡,实际却跑在占用最高的那张卡上。
6.3 动态 shape 导致的碎片化:越跑越容易 OOM
如果输入数据的序列长度或图像尺寸不固定,你会发现训练刚开始时挺正常,越往后越容易出现 OOM。这不是模型变大了,而是缓存分配器里的块被切得越来越碎。
最典型的场景是 NLP 里不等长 batch 直接 pad 到当前 batch 的最大长度。每次最大长度都在变,分配器不断申请不同大小的块,旧块释放后无法被新请求复用,碎片越积累越多。
缓解手段有几个:
- 把序列长度归到固定的 bucket,比如 128、256、512、1024,让每次分配的形状更稳定。
- 在 PyTorch 2.x 里设置环境变量:
bash复制export PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True
这个选项会启用可扩展段分配,减少碎片化。它适用于局部显存不足的场景,但对某些算子可能带来少量额外开销,建议实际跑一轮对比。
- 训练循环里捕获 OOM 后,做一次
torch.cuda.empty_cache()并重置动态 shape 状态,再继续。
6.4 遇到 GPU Crash Dump 时的处理思路
还有一种崩溃和 OOM 不同,日志里会出现类似 GPU crash dump triggered 的信息。
这类问题通常是 CUDA context 已经损坏,比如某个 kernel 访问了非法显存地址、显存超限后没有正常恢复、驱动进入了异常状态。在 Windows 的 WDDM 模式下,还可能出现驱动重置。
遇到这种崩溃,我的处理原则是:
- 不要再尝试在当前进程里捕获异常后继续跑,CUDA context 可能已经不可信。
- 把训练时的错误信息、显存快照和退出点记下来,杀掉进程重跑。
- 检查是不是最近改动了某些自定义 CUDA 算子,或者把某个变量释放之后又被使用了。
很多时候 GPU crash dump 的诱因并不是 OOM,而是代码里某个越界写。这类 bug 在普通 CPU 上不容易发现,在 GPU 上会直接搞坏显存。比较有效的排查方式是先把自定义算子和新改动回退,用最小化脚本复现,确认不是框架层面的问题。
6.5 一条可以参考的显存排查行动路径
如果你现在正好被 OOM 折磨得焦头烂额,可以按下面这条路径走一遍:
- 打开
nvidia-smi dmon -i <卡号> -d 1,确认没有其他进程抢占显存。 - 在训练脚本关键位置打印
torch.cuda.memory_summary(),定位显存飙升的阶段。 - 如果峰值出现在前向传播,优先考虑激活值占用,开 AMP,再看要不要加激活检查点。
- 如果峰值出现在反向传播,考虑梯度累积 +
zero_grad(set_to_none=True),检查是否有不必要的中间变量。 - 如果模型本身大到单卡放不下,直接用 FSDP 并把优化器状态 offload 到 CPU。
- 每次修改后都记录
max_memory_allocated和训练速度,避免优化完显存却发现训练时间翻倍。
我在实际项目里的习惯是,把显存基线、训练速度和模型结构变更记录在同一个地方,这样每次 OOM 都不是从零开始排查,而是直接对比上一次的基线数据,很快就能锁定变更点。
最后再分享一个小技巧:在训练脚本初始化阶段,用 torch.cuda.set_per_process_memory_fraction(0.9) 给进程设置一个显存使用上限,比如 90%。这样即使代码里出现了意外的大分配,进程也是优先 OOM 报错而不是把整张卡撑满后拖垮其他任务。这个上限不会优化你的显存占用,但能让你的程序在共享 GPU 的环境里表现得更有边界感,排查问题时也更干净。
