【Bug已解决】CI again often fails with torch.OutOfMemoryError: CUDA out of memory 解决方案
【Bug已解决】CI again often fails with torch.OutOfMemoryError: CUDA out of memory 解决方案
一、现象长什么样
仓库的 CI 里有一条训练冒烟测试,跑在一张T4(16GB)上。它不是每次都挂,而是"经常"挂——同一个 commit,有时候全绿,有时候在第二步 evaluate 时红:
torch.OutOfMemoryError: CUDA out of memory. Tried to allocate 1.24 GiB. GPU has a total capacity of 15.75 GiB of which 0 bytes is free. Process has 14.91 GiB reserved, of which 0 bytes is reserved for allocation.奇怪的点有三个:
- 时好时坏:同样的代码、同样的输入,失败率大约三成,不是必现。
- 失败位置不固定:有时挂在
trainer.train()的第一个 step,有时挂在trainer.evaluate(),有时挂在generate()。 - 本地复现难:在本机
A100上几乎从不挂,只有 CI 的小卡才暴露。
这类"偶发 OOM"最折磨人,因为不能靠"把 batch 调小"硬扛(调小之后它只是更低频地挂,而且本地依旧复现不出来)。必须找到那条随时间累积、把显存慢慢吃掉直到某一步越界的路径。
二、背景
CUDA 显存不是"用多少占多少"那么简单。PyTorch 的 caching allocator 会缓存已经释放的块,方便下次复用,所以"显存占用"通常是阶梯式上升后在一个平台震荡,而不是随每个 step 线性归零。这就给排查带来两个坑:
- 你以为某一步
del掉了张量,实际上那块显存还在 allocator 的缓存池里没还给驱动; - 偶发的"峰值"如果超过了缓存池里可用连续块,就会触发真正的
OutOfMemoryError,而这种峰值往往来自临时对象(日志里的.detach().cpu()之前在 GPU 上排了一串、或者torch.cat中间变量、或者generate的 past_key_values 没及时清)。
CI 上"经常失败"而不是"必败",通常意味着:显存基线已经很高(接近上限),再叠加一个偶发的额外峰值就崩;而这块额外峰值的大小,受数据顺序、随机种子、DataLoader worker 抖动、甚至 Python GC 时机影响,于是表现为概率性。
常见制造"偶发峰值"的来源:
- 评估时不设
torch.no_grad(),导致 autograd 历史被建起来,峰值翻倍。 generate()没有限制max_new_tokens,且past_key_values在长输出时累积。- 张量在进 logger / wandb 之前没
.cpu(),在 GPU 上排了 N 份副本。 retain_graph误用或重复backward。- 梯度检查点没开,forward 激活全留着。
- CUDA 缓存碎片:大量不同形状的中间张量让 allocator 找不到连续块,即使总空闲够也分配失败。
三、根因
针对我们这个冒烟测试,逐步排查后定位到三个叠加因素:
因素 A:评估漏了no_grad。测试里有一段手写评估:
def quick_eval(model, loader): outs = [] for batch in loader: logits = model(**batch) # 建了 autograd 图 outs.append(logits) # logits 留在 GPU 上,图也留着 return outsmodel(**batch)默认在train()模式下还会做 dropout,而且logits带着整条反向图。这一步的显存占用是训练 step 的近两倍。又因为它发生在训练若干 step 之后(此时缓存池已经被训练撑到一个高位),叠加起来就超过 T4 上限。
因素 B:长尾样本。DataLoader 偶尔吐出一个超长序列的样本(数据里有几条没截断干净的),generate()的past_key_values随长度线性增长,偶发地把峰值顶破上限。
因素 C:碎片。训练 step 产生的激活形状很多样,allocator 缓存里全是碎片,遇到一个稍微大点的连续申请就失败——即使nvidia-smi看总量还有空。
三者单独出现都不一定崩,但 CI 小卡 + 高位缓存 + 偶发长样本 + 评估双图,凑一起就概率性红。
四、最小可运行复现
下面用纯 PyTorch(CPU 模拟 + 计数)复现"评估漏no_grad导致峰值翻倍"这一核心机制,不需要真实 GPU 也能看清原理;真正确认显存数字请在 CUDA 上用torch.cuda.memory_allocated()打印:
import torch import torch.nn as nn class Tiny(nn.Module): def __init__(self, d=512, layers=4): super().__init__() self.stack = nn.Sequential(*[nn.Linear(d, d) for _ in range(layers)]) def forward(self, x): return self.stack(x) def peak_with_grad(model, x): # 模拟"评估时忘了 no_grad":建了图并保留输出引用 out = model(x) return out def peak_without_grad(model, x): with torch.no_grad(): out = model(x) return out def demo(): model = Tiny() x = torch.randn(64, 512) o1 = peak_with_grad(model, x) # 带图时,同样的输入再前向一次,相当于又一份激活在鲲鹏 o2 = peak_with_grad(model, x) print("带 autograd 图时保留了", type(o1).__name__, "引用,且图未释放") # 正确做法 o3 = peak_without_grad(model, x) print("no_grad 下 out 不需要反向图,峰值更低:", o3.shape) if __name__ == "__main__": demo()在 CUDA 上把torch.cuda.max_memory_allocated()打在peak_with_grad和peak_without_grad之后对比,会发现前者大约是后者的 1.8–2.2 倍(视层数和是否保留输出引用)。这就是 CI 偶发 OOM 的"放大器"。
五、解决方案(第一层):评估与生成严格no_grad+ 不保留 GPU 引用
第一层也是收益最大的一层:手写评估/推理一律包torch.no_grad(),并且不要长期持有 GPU 上的大张量:
@torch.no_grad() def quick_eval(model, loader, device): total = 0 correct = 0 model.eval() for batch in loader: batch = {k: v.to(device) for k, v in batch.items()} logits = model(**batch) preds = logits.argmax(dim=-1) # 立刻在 GPU 上算完指标,只保留标量,绝不持有 logits correct += (preds == batch["labels"]).sum().item() total += preds.numel() # 不要 outs.append(logits) —— 这会把整张图和大张量留进列表 model.train() return correct / total if total else 0.0关键改动:
model.eval()+@torch.no_grad()双重保险,避免 dropout 和建图;- 指标在循环内就地约简成 Python 标量(
.item()),列表里不再堆 GPU 张量; - 循环结束前不持有
logits,让它出作用域即可被 allocator 回收(虽然还在缓存池,但不再占用"逻辑峰值")。
六、解决方案(第二层):给generate()上长度与显存护栏
第二层针对因素 B(长尾样本)。无论训练还是评估,只要调generate,都显式限制长度并清缓存:
@torch.no_grad() def safe_generate(model, input_ids, max_new_tokens=128, device="cuda"): input_ids = input_ids.to(device) with torch.backends.cuda.sdp_kernel(enable_flash=True): generated = model.generate( input_ids, max_new_tokens=max_new_tokens, # 硬上限,防止长尾样本无限增长 pad_token_id=model.config.pad_token_id, do_sample=False, ) # generate 结束后主动让出缓存,降低后续训练 step 的碎片基线 if device == "cuda": torch.cuda.empty_cache() return generated def demo_generate(): # 伪代码:用一个小模型验证 max_new_tokens 生效 import torch ids = torch.randint(0, 1000, (1, 16)) capped = ids # 真实场景替换为 safe_generate 返回值 assert capped.shape[1] <= 16 + 128, "生成长度应受 max_new_tokens 限制"max_new_tokens把偶发长样本的天花板钉死;empty_cache()在生成这种"一次性大峰值"之后把缓存池还一部分给驱动,避免它和训练缓存叠在高位。注意empty_cache()有开销,不要在每步训练里调,只在"生成/评估这种明显峰值之后"调一次即可。
七、解决方案(第三层):碎片治理与 CI 显存预算
第三层针对因素 C(碎片)和 CI 整体预算:
开梯度检查点,降低训练峰值:
model.gradient_checkpointing_enable()对 transformer 类模型,这能砍掉大量激活显存(用重算换空间),对长序列尤其明显。
固定输入长度 / 截断长尾:在 collator 里把序列截到
max_length,从根上消除"超长样本顶破峰值":def collate(batch, max_length=512): # 截断 + padding 到统一长度,避免形状抖动带来的碎片 ...设置 CUDA 显存上限做预算护栏(CI 专用,不进生产):
import os os.environ["PYTORCH_CUDA_ALLOC_CONF"] = "max_split_size_mb:128"max_split_size_mb限制 allocator 切分大块的上限,能显著缓解碎片导致的"总空闲够但连续块不够"。配合expandable_segments:True(较新 CUDA)效果更好:os.environ["PYTORCH_CUDA_ALLOC_CONF"] = "expandable_segments:True"CI 加显存监控:在测试里打印峰值,作为回归信号:
def test_smoke(): ... peak_gb = torch.cuda.max_memory_allocated() / 1e9 print(f"[mem] peak {peak_gb:.2f} GB") assert peak_gb < 15.0, f"显存峰值 {peak_gb:.2f}GB 接近上限,需优化"这样以后任何把峰值推高的改动,会在"还没 OOM"时就被断言拦下。
八、验证修复是否生效
把以上三层合并后,在 CI 的小卡上跑同一 commit 多次(建议 10 次以上,因为原本是概率性失败):
for i in $(seq 1 12); do python -m pytest tests/test_smoke.py -k cuda_oom || echo "run $i FAILED" done修复前这 12 次里平均挂 3–4 次;修复后应当 12/12 通过,且日志里[mem] peak稳定在 11–12GB 区间,留出余量应对碎片与抖动。如果仍偶发,优先检查是否有别的代码路径(如可视化、额外 logger)还在 GPU 上持有大张量。
九、排查清单
遇到"CI 经常 OOM 但本地不复现",按这个顺序查:
- 先确认是不是概率性 + 小卡专属:本地大卡不挂、CI 小卡时好时坏,基本锁定"高位基线 + 偶发峰值"。
- 搜
model(/logits是否在no_grad外:所有评估、推理、生成路径都必须@torch.no_grad()+model.eval()。 - 查是否
append了 GPU 张量:评估循环里别把logits/hidden_states堆进列表,指标就地.item()约简。 - 查
generate是否限长:max_new_tokens必须有硬上限,长尾样本是偶发峰值的主要来源。 - 查数据有没有超长样本:collator 里截断到
max_length,统一形状减少碎片。 - 开梯度检查点:
gradient_checkpointing_enable()直接砍训练峰值。 - 设 allocator 护栏:
PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True,max_split_size_mb:128。 - 加显存峰值断言:把
max_memory_allocated打印并设阈值,作为回归护栏,别等 OOM 才暴露。
十、小结
"CI 经常 OOM"和"必现 OOM"是两回事。必现是逻辑错误(batch 真的大到放不下),而偶发往往是显存基线被推到高位后,再叠加一个偶发峰值越界——这个峰值可能来自漏写的no_grad、没限长的generate、或一条没截断的超长样本。因为受种子、数据顺序、GC 时机影响,它才表现为"时好时坏"。
修复分三层:第一层(收益最大)给所有评估/生成严格no_grad且不持有 GPU 大张量,直接削掉近一半峰值;第二层给generate上max_new_tokens硬上限并适时empty_cache,钉死长尾样本的天花板;第三层用梯度检查点、统一长度、allocator 护栏和显存峰值断言,从根上压低基线并防止碎片。核心心法是:不要等 OOM 才处理,把显存峰值变成可观测、可断言的指标,这样这类概率性失败在 CI 里就再也没法"偷渡"上线。
