OOM 的排查链路
🚨 训练过程中突然 OOM,你应该从哪些地方开始排查?¶
💡 面对突然的 OOM,不要慌,沿着“数据 → 配置 → 代码 → 框架”的脉络逐一排查,通常能在几分钟内定位元凶。 我习惯按以下顺序执行:
- 检查是否某个 batch 的数据异常 突然 OOM 最常见的原因是数据集中出现了超长样本。若未对文本进行截断或 padding 策略不当,个别样本长度可能是平均的 3-5 倍,导致激活值瞬间暴涨。
- 排查方法:在 DataLoader 中打印当前 batch 的最大序列长度,或在
collate_fn中记录统计。如果是语音或图像,检查是否有超高分辨率输入。 -
临时方案:开启动态 padding 或强制截断,确保所有样本不超过预设的最大长度。
-
确认是否开启验证(eval)阶段 训练每隔若干步会进行验证。验证时通常会关闭梯度计算(
torch.no_grad()),但可能因为使用了更大的 batch size 或缓存了所有验证结果而导致显存超出。 - 特征:OOM 发生在第一个 eval batch 或 eval 结束时。
-
解决:减小 eval batch size,或在 eval 中只保留必要的指标而不缓存所有 logits。
-
检查是否开启了不必要的缓存或日志
- 某些调试钩子(如记录每层梯度范数)可能意外保留了张量引用,导致显存泄漏。
- 使用
torch.cuda.memory_summary()或torch.cuda.memory_snapshot()检查是否有大量张量未释放。 -
如果使用了 WandB/TensorBoard 记录直方图或嵌入,可能占用额外显存。
-
查看显存碎片情况 即使总空闲显存足够,但若缺乏连续大块,
cudaMalloc仍然可能失败。这种情况在长时间训练后容易出现。 - 指标:
nvidia-smi显示 Memory-Usage 接近 100%,而torch.cuda.memory_allocated()显示已分配却远小于总显存,说明碎片严重。 -
缓解:设置环境变量
PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True或定期调用torch.cuda.empty_cache()(谨慎使用,可能影响性能)。 -
排查模型/优化器状态是否意外膨胀
- 如果使用了混合精度训练,检查 loss scaler 是否产生了 NaN/Inf,导致梯度更新异常,优化器状态可能积累非法值。
- 某些优化器(如 LAMB)可能在内部维护额外的缓冲区。
-
是否无意中解冻了一些原本冻结的参数(例如在 LoRA 中误把基础模型设为可训练)。
-
审视并行策略与通信缓冲区
- DeepSpeed ZeRO 或 FSDP 在通信时可能临时需要额外缓冲区。如果同时使用多个并行策略,检查是否开启了
overlap_comm等,其缓冲区可能超出预期。 -
在多卡训练中,OOM 可能只发生在某一张卡上,用
pdsh或监控查看各卡显存使用是否均衡。 -
回退最近的代码/配置变更 突然 OOM 往往和最近的修改有关。对比上一次稳定训练的配置(学习率、batch size、序列长度、模型结构、是否开启新优化等)。
✅ 总结排查流程
数据 → 评估阶段 → 缓存/日志泄漏 → 碎片 → 模型状态 → 并行配置。遵循这个顺序,80% 的突然 OOM 能找到直接原因,并能快速给出应对措施。
📊 如何通过 nvidia-smi 观察显存变化趋势?¶
💡 nvidia-smi 是快速诊断 GPU 显存的瑞士军刀,通过持续监控、记录日志和观察模式变化,可以判断训练中的显存是稳定、泄漏、还是突发尖峰。
🔧 常用命令与技巧
- 持续监控模式
-
-s pucm:显示功耗、利用率、时钟、显存信息。 -
-d 2:每 2 秒刷新一次。 输出中mem列表示当前使用的显存百分比,观察其随时间的变化曲线:如果单调递增且不回落,极可能是内存泄漏;如果出现周期性尖峰,可能与数据加载或 checkpoint 保存相关。 -
实时查看进程级显存
观察 Processes 部分,找到自己的训练进程 PID,看 Memory-Usage 慢慢上升还是突然跳变。若能结合 nvtop 工具,效果更直观。
- 记录显存日志以便事后分析
用 Python 读取 CSV 并绘图,可以精确看到 OOM 前的显存爬升曲线。如果是阶梯状上升,通常对应梯度累积或 optimizer step 的分配;如果是陡增,则对应某个大张量分配(如注意力矩阵)。
- 识别 OOM 时的模式
- 突然垂直上升后崩溃:往往因为一个 batch 中出现了异常大的激活张量(如超长序列)。查看 log 中最后几个时间点的显存使用量,通常会在几秒内冲顶。
-
缓慢增长直至溢出:可能为内存泄漏(某些张量未释放),也可能是碎片积累。
nvidia-smi显示的Used包含了 PyTorch 缓存,可能掩盖碎片问题,需结合 PyTorch 自身 API 判断。 -
检查温度与功耗相关性 若显存占用高同时温度飙升,可能触发了硬件保护降频,导致训练变慢但显存不减,这虽不直接 OOM,但会间接影响性能。
nvidia-smi能同时观察温度、功耗和显存。
🔬 高级用法
-
通过
nvidia-smi nvlink -e 0查看 NVLink 传输,辅助判断多卡通信是否导致显存占用。 -
结合
nvidia-smi -i 0 -q -d MEMORY输出详细的显存信息,包括 BAR1 使用等。
✅ 总结:用 nvidia-smi 就像看心电图,把显存使用率绘制成时间序列,能最快速识别 OOM 的宏观类型,为进一步精细化定位提供方向。
🛠️ PyTorch 的 memory_allocated() 和 memory_reserved() 在 OOM 时如何帮助定位?¶
💡 当 OOM 发生时,torch.cuda.memory_allocated() 告诉你 PyTorch 实际在张量上占用了多少显存,memory_reserved() 告诉你 PyTorch 缓存分配器总共向 CUDA 驱动申请了多少显存。两者之差就是分配器内部的“空闲但不归还”的缓存。分析这两个值能区分是真实显存不足还是碎片/缓存导致的 OOM。
🔍 具体诊断方法
-
torch.cuda.memory_allocated(device):返回当前所有活的张量占用的字节数。如果你在 OOM 前调用它,发现它接近 GPU 总显存,说明确实是模型、优化器、激活等太大,需要减少 batch/序列或使用并行/offload。 -
torch.cuda.memory_reserved(device):返回缓存分配器持有的总字节数。如果reserved远大于allocated,比如 allocated=15GB, reserved=23GB(24GB 卡),意味着大量显存被分配器缓存但未被实际张量使用。此时nvidia-smi显示 Used 约 23GB,而实际需要的张量只有 15GB。这时候如果发生 OOM,很可能是因为碎片:缓存池里没有足够大的连续块来满足新分配请求,尽管总空闲(reserved - allocated)有 8GB,但无法利用。
🧪 使用方式
- 在代码中插入监控点
print(f"Allocated: {torch.cuda.memory_allocated()/1024**3:.2f} GB")
print(f"Reserved: {torch.cuda.memory_reserved()/1024**3:.2f} GB")
可以在每个训练步或每个 epoch 后记录,观察趋势。
- 定位碎片化
- 如果
reserved - allocated很大,但偶尔出现CUDA out of memory,可在 OOM 前调用torch.cuda.memory_summary()打印详细统计,它会显示每个 block 的大小分布。如果有很多小的空闲 block 却没有大块,即可确认碎片。 -
缓解办法:设置
PYTORCH_CUDA_ALLOC_CONF=max_split_size_mb:512或启用expandable_segments。 -
OOM 后捕获分析 虽然 OOM 异常无法完全避免,但可以使用
try-except包裹训练步骤,在捕获到RuntimeError: out of memory时打印memory_allocated()和memory_reserved(),然后torch.cuda.empty_cache()释放缓存。通过对比前后值,判断此次分配请求的大小。例如分配前 allocated=20GB,请求一个 2GB 张量报 OOM,reserved=23GB,那表明碎片导致无 2GB 连续块。
📊 实战案例
训练 7B 模型,单卡 24GB,开启梯度检查点后依然 OOM。打印发现 allocated 仅 18GB,但 reserved 达 23.5GB。进一步 memory_summary 显示最大连续空闲块只有 300MB,尽管总空闲 5.5GB。结论:碎片严重。解决方案:设置 PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True,显存碎片消除,训练顺利进行。
✅ 因此,这两个函数是深入 PyTorch 显存管理的探针,OOM 时务检查之,区分真性不足与假性不足。
🤔 如果模型参数量没有变化,但训练后期突然 OOM,可能是什么原因?¶
💡 模型参数量固定,但 OOM 仍可能发生,原因常出在数据、优化器状态、激活值、碎片、日志/缓存这五个方面。即使参数量不变,这些因素可在训练过程中动态膨胀。
🔍 具体可能原因
- 数据集中后期出现超长样本 若数据未做有效截断,训练到后期遇到极长文本,激活值随序列长度线性(注意力 O(L²))增加,瞬间撑爆显存。这是最常见且隐蔽的原因。
- 验证:在 dataloader 中打印每个 batch 的最大 token 数,观察 OOM 时的 batch 长度。
-
解决:对输入进行强制截断,或使用动态序列长度的 batch 策略但设置上限。
-
优化器状态积累或泄漏
- 某些优化器如
AdamW可能在内部分配额外的临时缓冲区,但这些通常在初始化时固定,不会在后期突然增大。 - 更可能的是代码 bug 导致优化器状态中出现 NaNs 或 Infs,框架分配额外空间处理异常?罕见。
- 真实常见的是:无意中累积了计算图。例如每步保存了 loss 或 logits 的引用,导致
backward后计算图没释放,梯度累积,优化器状态虽不变,但 autograd 保留的中间张量不断增加,最终 OOM。 -
排查:使用
torch.cuda.memory_summary()查看是否有大量的autograd节点占用。 -
激活值策略发生变化
- 若在训练过程中动态关闭了梯度检查点(比如基于某个条件),激活值会突然增多。
- 或者序列长度因数据 padding 策略变化而被动增加。
-
使用 FlashAttention 的版本切换也可能改变激活峰值。
-
显存碎片累积 训练初期显存充足,碎片被大块分配掩盖。随着不断分配释放,缓存池碎片化严重。后期即使总空闲足够,但分配一个稍大的连续张量就可能 OOM。这在长时间训练中极其常见。
- 指标:
reserved稳定但allocated波动,nvidia-smi 显示高占用。 -
缓解:设置
expandable_segments:True或定期重启训练。 -
日志、checkpoint 或其他外部操作
- 在训练后期,可能开启额外的验证、保存 checkpoint 或记录大批量 embedding 等,这些操作可能产生额外的显存分配。
- 例如,在保存模型时,某些框架会先在 CPU 上创建副本,但也可能占用 GPU 显存放序列化中间数据。
-
如果使用了 DDP,可能在 checkpoint 时进行额外通信同步,消耗通信缓冲区。
-
混合精度 Loss Scale 异常 若 loss scale 变得极大,梯度值异常放大,虽不直接增大显存,但可能导致优化器状态中的统计量异常,进而引发额外内存分配或内核错误,有时表现为 OOM 而非报数值错误。
🧰 排查步骤
① 在 OOM 步捕获 batch 数据,检查 max_seq_len。
② 打印 memory_allocated() 和 reserved,看是否碎片。
③ 检查代码中是否缓存了任何张量(如历史预测值)。
④ 关闭可能导致额外开销的日志(如 profiling)。
✅ 所以,参数量不变但后期 OOM,重点查数据、碎片和隐式张量泄漏。
🔧 batch size 设置过大导致 OOM,除了减小 batch size 还有什么办法?¶
💡 减小 batch size 是直接减少激活和梯度的有效手段,但若必须保持全局 batch 规模,可以使用梯度累积、混合精度、更大力度的梯度检查点、卸载到 CPU等多种技巧来在不变动有效 batch size 的情况下降低单次显存占用。
🔢 替代方案清单
- 梯度累积 (Gradient Accumulation) 将所需的全局 batch size 拆分成多个 micro-batch,每次前向/反向只计算一个 micro-batch,累积梯度后一次性更新。这样每步的显存占用等于 micro-batch 的大小,但最终更新效果等价于大 batch。
- 例如:目标 batch=64,设
per_device_train_batch_size=8,gradient_accumulation_steps=8。 -
代价:训练速度变慢(因多次前向),但不增加显存。
-
启用混合精度训练 (AMP) 使用 FP16 或 BF16 代替 FP32 进行前向和反向,将激活和梯度的精度减半,显存显著下降。PyTorch 的
torch.cuda.amp自动处理,通常能节省 30-40% 显存,允许使用更大的 batch size。 -
注意:某些操作仍需要 FP32。
-
升级梯度检查点 (Gradient Checkpointing) 若已开启,可尝试使用更激进的检查点策略(如对整个模块而非仅 Block),或用
checkpoint_sequential。也可以结合选择性检查点,丢弃更多激活。 -
代价:速度变慢,但换回空间。
-
使用 FlashAttention 消除注意力矩阵的 O(n²) 显存,通常在长序列下节省显著,能腾出空间给更大的 batch。
-
优化优化器状态
- 使用 8-bit Adam (
bitsandbytes) 将优化器状态量化,节省约一半优化器显存。 -
若用 ZeRO-1/2,可以将优化器状态分片到多卡,但单卡的话可考虑 CPU Offload 优化器状态(例如 DeepSpeed ZeRO-Offload 或
torch.optim.Adam+ 自定义 offload),将动量、方差存在 CPU。 -
序列长度维度调整
- 若任务允许,缩短输入序列长度能极大节省激活。可通过动态截断、或对输入进行摘要压缩。
-
使用随机长度采样,避免极端长样本撑爆。
-
使用参数高效微调 (PEFT)
-
若任务是微调,用 LoRA/QLoRA 冻结主干,只有少量参数有梯度和优化器,大幅度降低训练所需显存,自然能使用更大 batch。
-
开启激活卸载 (Activation Offloading)
-
将部分中间激活迁移到 CPU 内存,仅保留当前计算所需在 GPU。DeepSpeed 的 ZeRO-Infinity 或
torch.utils.checkpoint可与 offload 配合。 -
使用更轻量的模型架构
-
如采用 GQA/MQA 降低 KV Cache(训练时影响较小但仍有激活节省),或缩小 hidden 维度。
-
调整 batch size 为动态
- 采用 adaptive batch size,根据当前序列长度动态调整每个 batch 的样本数,长序列用小 batch,短序列用大 batch,最大化利用率。
📊 综合运用示例
单卡 24GB 训练 7B 模型,目标 batch=32。标准配置 OOM。解决:
-
使用 LoRA + 梯度检查点 + 混合精度 + 梯度累积(micro batch=2, 累计步=16)。
-
开启 FlashAttention。 最终可成功训练。
✅ 因此,面对 OOM,不要只盯着 batch size 的数字,组合使用多种省显存技术才是正解。
📏 序列长度变化如何导致 OOM?举例说明。¶
💡 序列长度 L 是激活显存最重要的因子之一,尤其注意力机制会产生 O(L²) 的中间矩阵。当序列长度突然增大,激活显存可能平方级爆发,瞬间导致 OOM。
🔍 机制与示例
假设训练一个 Transformer 模型,hidden=4096,层数 32,heads=32,训练时没有使用 FlashAttention 或梯度检查点。
-
基准 (L=512): 注意力分数矩阵 S 大小为 [batch, heads, 512, 512],FP16 下占 1×32×512×512×2 ≈ 16 MB。 其它线性层激活约 [batch, L, hidden] 形状,共 32 层,每层 Q、K、V 各约 1×512×4096×2 = 4 MB,总计约 4×3×32 ≈ 384 MB。 整个激活量可能 1 GB 左右,加上模型权重,24GB 显存绰绰有余。
-
情况 1:数据中出现 L=4096 的样本 S 变为 [1, 32, 4096, 4096],大小为 32×4096²×2 ≈ 1 GB!仅注意力矩阵就暴增 64 倍。 其他线性激活 (Q,K,V) 也变为 4 MB→ 32 MB 每层,总计 32×32×3 ≈ 3 GB。激活总量迅速达到 4-5 GB,加上权重、优化器,可能超出 24GB,OOM。
-
情况 2:使用 FlashAttention 但未用梯度检查点,L 从 2048 提升到 8192 FlashAttention 不存注意力矩阵,但 Q、K、V 仍为 O(L),且前向中会保留这些张量用于反向。 L=2048 时,每层 Q/K/V 各约 2048×4096×2=16 MB,32层共约 1.5 GB,激活总和可能 3 GB。 L=8192 时,每层 8192×4096×2=64 MB,32层共约 6 GB,激活总和可能 12 GB。加上模型,可能 OOM。
⚙️ 为何序列长度变化导致显存非线性增长
因为注意力矩阵理论上是 O(L²)(无优化时),而其他激活是 O(L)。即使有 FlashAttention,重计算过程也可能临时产生 O(L²) 的工作缓冲区(但被分块控制在 SRAM 内),显存仍然 O(L) 但基数变大。
🔧 如何避免
-
对训练数据进行严格截断,设定
max_seq_length,超出截断或滑动窗口。 -
使用梯度检查点,即便长序列,激活驻留也受控。
-
结合 FlashAttention,消除 O(L²) 矩阵的显存占用。
-
动态 batch:根据序列长度动态调整 batch size,长序列时自动减小 batch。Hugging Face 的
Trainer支持pad_to_multiple_of等,但更好是用自定义 batch sampler。
📊 实例
微调 Longformer 时,某 batch 中出现了 8192 token 的文档,OOM。后来加入截断(max=4096)并开启 FlashAttention,问题解决。
✅ 因此,序列长度是激活显存的主导因素,控制其峰值是关键。
🔁 激活值重计算开启后仍然 OOM,可能是什么原因?¶
💡 梯度检查点只能减少部分激活,但其他部分如模型权重、优化器状态、梯度、KV Cache(如果适用)仍然占用大量显存,并且重计算本身也会产生临时峰值。若 OOM 依然发生,通常是总体显存仍超限,或者重计算实现没有覆盖最主要的激活源。
🔍 可能原因分析
- 优化器状态和梯度仍然过大 梯度检查点对权重、梯度、优化器状态无影响。若模型很大,优化器状态(如 Adam 的 8 bytes/param)和梯度(FP16 2 bytes/param)仍可能占总显存 80% 以上。开启检查点只是将激活从 50% 降低到 10%,但优化器状态没变,总显存依然超标。
-
解决:使用 ZeRO-2/3 分片优化器状态,或采用 8-bit 优化器。
-
基础模型权重未压缩 即使冻结主模型或使用 LoRA,若基础权重为 FP16,7B 占 14GB,13B 占 26GB。检查点省下激活的几 GB 不足以弥补。此时需量化基础权重(QLoRA)。
-
重计算未涵盖所有大激活 默认的梯度检查点可能只对
torch.utils.checkpoint.checkpoint包裹的模块生效,但某些操作(如自定义注意力)可能遗漏,导致其内部激活仍被保留。另外,若只对一部分层设置检查点,其他层激活照常保留。 -
排查:使用 PyTorch Profiler 查看哪些操作保存了大量激活,确保它们被检查点覆盖。
-
重计算产生的临时峰值超出 开启检查点后,反向传播时会重新运行前向,临时分配该段的所有激活。若检查点分段太大(例如 10 个 Block 一段),重计算时激活峰值可能接近无检查点的水平,这可能导致瞬时 OOM。
-
方案:减小检查点段大小,例如每个 Block 作为一个段,而不是每 4 个 Block。
-
KV Cache 占用被忽视 训练时编码器-解码器模型或某些自回归训练(如 teacher forcing)也会需要存储 K、V 以供交叉注意力,这部分与序列长度成正比,检查点不能减少它。长序列下 KV Cache 可能成为大头。
-
其他占用:通信缓冲区、数据加载 分布式训练中的通信缓冲(如 all-reduce 的临时空间)以及
pin_memory的数据缓冲区也占用显存。检查点对此无效。可尝试关闭overlap_comm或调小batch size。 -
碎片或缓存导致 false OOM 如前述,碎片可能使得分配失败。检查
memory_reserved与allocated。
✅ 解决顺序
若开启检查点仍 OOM:先用 memory_summary 确认是真实占用还是碎片;接着考虑量化权重、分片优化器、减小 batch/序列;再检查重计算的粒度。
⚙️ ZeRO 配置不当导致 OOM,应该如何调整?¶
💡 ZeRO 各阶段对显存的分片程度不同,配置不当主要表现为 stage 选低导致显存不足,或 offload 未开启导致 OOM。调整方法是:逐渐升级 ZeRO stage,适当开启 CPU/NVMe offload,同时调小通信相关缓冲区。
🔧 调整策略
- 先从 ZeRO-1 开始,不够上 ZeRO-2,再不够上 ZeRO-3
- ZeRO-1 仅分片优化器状态,节省约 4× 的优化器显存。若优化器是瓶颈,立竿见影。
- ZeRO-2 额外分片梯度,节省 1× 梯度。
-
ZeRO-3 分片模型参数,节省 1× 权重。根据模型大小选择。例如,单卡 80G 训练 13B,ZeRO-2 可能足够;70B 则必须 ZeRO-3。
-
开启 CPU Offload 在 DeepSpeed 配置中设置
"offload_optimizer": {"device": "cpu"}将优化器状态移至 CPU,甚至"offload_param": {"device": "cpu"}将参数也移至 CPU(ZeRO-3 下)。这会大幅降低 GPU 显存,但训练变慢。如果 ZeRO-3 仍然 OOM,开启 offload 是最后利器。 -
调整 ZeRO-3 的参数获取策略
stage3_prefetch_bucket_size:控制预取的参数分片大小。太小则通信频繁,太大则占用显存。可适当调小以降低峰值。stage3_param_persistence_threshold:小于该阈值的参数将常驻 GPU 而不释放,可避免小参数反复收集的开销。如果显存紧张,调小此值使更多参数释放。-
stage3_max_live_parameters:限制同时存在于 GPU 上的参数数量,减小可缓解峰值。 -
优化通信缓冲
- 关闭
overlap_comm("overlap_comm": false),避免额外的梯度通信缓冲区,可释放约 1 个梯度大小的显存。 -
关闭
contiguous_gradients,同样省缓冲。 -
减少 ZeRO 的中间分配 在 DeepSpeed 配置中,
"reduce_bucket_size"和"allgather_bucket_size"控制通信时的临时缓冲区大小。减小这些值可降低瞬时显存峰值,但可能增加通信次数。 -
使用 ZeRO++ 或优化内核 若网络带宽有限,ZeRO++ 可通过量化通信减少缓冲区大小,间接缓解显存。
📊 实例
8 张 V100 32GB 训练 30B,ZeRO-3 默认配置 OOM。调整后:开启 CPU offload optimizer,关闭 overlap_comm,将 stage3_max_live_parameters 设为 1e8,训练正常运行。
✅ 因此,ZeRO OOM 的调优是“分片深度+Offload+缓冲区调节”的组合拳,需根据模型和硬件逐步试验。
🎛️ 混合精度训练中的 Loss Scale 异常是否与显存有关?为什么?¶
💡 Loss Scale 本身不直接占用显存,但其异常(如变成 inf 或极大值)会导致梯度变为 NaN/Inf,进而触发 AMP 的特定行为(如跳过更新、调整 scale),有时会引发框架分配额外内存或导致优化器状态异常,间接引起 OOM。更常见的是,loss scale 异常是显存不足或其他问题的“症状”,而非原因。
🔍 关联机制
-
Loss Scale 的运作 混合精度训练中,前向和反向用 FP16,但权重更新用 FP32。为避免小梯度在 FP16 下溢为零,AMP 自动将 loss 乘以一个大的 scale(如 65536),反向后再 unscale 恢复梯度。当出现异常(如溢出),AMP 会降低 scale 并跳过本次更新。
-
OOM 与 Loss Scale 的间接关系
- 显存不足导致计算失败:当 GPU 显存耗尽,某些算子可能无法分配所需内存,产生非数 (NaN) 或未定义值,这些值进入 loss 和梯度,使得 loss scale 不断降低(因为检测到 inf/NaN)。因此,观察到 Loss Scale 持续下降或跳变,可能暗示训练过程存在不稳定,而根源可能是逼近显存极限导致的数值错误。
-
内存泄漏或碎片引发异常值:碎片可能导致某些张量分配失败,但程序捕获异常后可能产生未初始化数据,进而污染梯度,影响 scale。
-
Loss Scale 异常可能导致 OOM 吗? 理论上不会直接增加显存,但实践中如果 scale 剧增(虽然 AMP 通常限制最大值),会产生极大梯度,可能导致优化器状态中动量等出现极大值,进而需要更多位数?不太可能。更大的风险是,当动态调整 scale 时,AMP 会保留一些备份状态(如 master weights 的副本)?也没有。所以,Loss Scale 异常通常是结果,而非原因。
🔧 排查建议
-
当看到 “Loss scale dropped below” 或 “Gradient overflow” 伴随 OOM,首先优化显存,而非纠结 scale。
-
检查是否使用了不稳定的操作(如大 reduction)导致溢出。
-
可以设置固定的 loss scale(如 1.0)来测试,但可能因下溢导致收敛慢,非标准操作。
✅ 总结:Loss Scale 异常与显存无直接因果,但它是训练数值不稳定的警铃,常由显存边缘状态引发。解决 OOM 后,scale 往往会恢复正常。
🚀 如果训练开始时能运行,几个 step 后 OOM,检查哪些动态增长的部分?¶
💡 训练开始正常、几 step 后 OOM,说明有某种资源的消耗在随着训练步骤单调递增,很可能未及时释放或不断累积。主要排查这些“动态增长”的部分:
-
📈 数据输入侧的动态形状 如果使用了动态 padding,早期 batch 可能序列较短,越后面遇到长样本,激活显存峰值冲高致 OOM。检查每个 step 的
max_seq_len是否逐步增大。 -
🧠 计算图/autograd 泄漏 每步若不小心保留了对输出或 loss 的引用(例如存入列表),反向传播后计算图无法释放,梯度累积不清理,导致中间激活和梯度占用持续膨胀。查看是否有类似
losses.append(loss)且loss保留图的情况。 -
💾 缓存或日志累积 训练中记录大量指标、存储预测结果或中间嵌入向量,如果这些张量长时间存留(如放在
list中),显存占用会线性增长。尤其是在验证阶段保留所有 logits 用于后续计算。 -
🪣 优化器状态泄漏 某些自适应优化器(如带 momentum 的 SGD)在每个 step 后不会额外增长,但若代码 bug 导致每步新建 optimizer 实例,旧优化器未释放,可能累积。
-
🧩 碎片累积效应 即使没有显式泄漏,频繁分配释放不同大小的张量导致 PyTorch 缓存池碎片化严重,某些分配请求因无连续块而失败。尽管总
allocated未大幅增长,但reserved维持高位且碎片逐步恶化,可能在某个分配较大张量时触发 OOM。 -
🔁 通信缓冲区叠加 分布式训练中,若 overlap_comm 开启,梯度通信缓冲区可能未及时清理或因流水线调度导致多步缓冲叠加。
✅ 排查顺序:
-
打印每步的
torch.cuda.memory_allocated()观察是否单调递增。 -
若递增,检查代码中是否有
list.append(tensor)或保留loss等图引用。 -
若无明显泄漏,检查数据集中最长样本是否出现在后期,并关注碎片。
🧮 使用梯度累积为什么有时反而 OOM?(梯度在累积前已占显存)¶
💡 梯度累积目的是用小 batch 替代大 batch,节省激活显存。但梯度本身在整个累积周期内持续驻留并累加,若模型很大,梯度本身就占大量内存,累积步数过多可能导致梯度显存超过节约的激活量,反而 OOM。
🔍 机理:
-
标准训练中,每个 micro-batch 完成反向传播后,梯度立即被 All-Reduce(分布式)并用于优化器更新,然后释放。
-
梯度累积时,每个 micro-batch 反向产生梯度,并累加到
.grad中。这期间所有累积的梯度都保持活跃,直到累积步数完成后才更新并清零。 -
对于大模型,梯度 FP16 占用等于权重大小(如 7B 为 14 GB)。若累积步数 8,梯度始终占用 14 GB,与单次大 batch 相同,反而因为 micro-batch 小导致激活虽小但需要多次前向,可能某些框架未释放中间缓存?激活因 micro-batch 小确实减少,但梯度没有减少,如模型原来单次 batch=64 需 14 GB 梯度,现在 micro-batch=8 累积 8 次,梯度依然占用 14 GB,而激活峰值降为 1/8。若之前 OOM 是因为激活过大,累积可解决;若 OOM 主因是梯度+优化器+权重太大,梯度累积无帮助,甚至可能因为多步间的临时缓冲区叠加而略微增加显存。
📊 场景:
13B 模型,优化器状态 52 GB,权重 26 GB,梯度 26 GB,激活 15 GB,总 >119 GB。单卡 80 GB 无法运行。若用梯度累积,梯度仍 26 GB,激活可能降至 5 GB,总和仍 109 GB,依然 OOM。此情况下需要 ZeRO 分片梯度/优化器,而非梯度累积。
⚠️ 另外:若累积步数过多,PyTorch 的 autograd 引擎可能保留中间计算图(如果 loss.backward() 调用时未正确释放),导致额外的激活占用。必须确保使用 loss.backward() 后正确清理图(optimizer.zero_grad() 在最后调用)。
✅ 所以,梯度累积只对激活瓶颈有效,对权重、梯度、优化器瓶颈无用,甚至可能因保留梯度时间延长而加剧碎片。
🤖 推理时 OOM,是权重、KV Cache 还是其他?¶
💡 推理时显存占用主要由权重和KV Cache构成,激活极小。OOM 时需快速判断瓶颈是谁:若模型文件大小接近或超过显存,通常是权重过大;若模型能加载但生成长文本时崩,通常是 KV Cache 耗尽显存。
🔍 判断方法:
-
仅加载模型(不推理)就 OOM → 权重显存超限(可能未量化)。
-
短文本推理正常,长文本 OOM → KV Cache 是主因。
-
使用
nvidia-smi:若 Memory-Usage 缓慢上升直至满,常为 KV Cache 增长。
🧮 量化估算:
-
权重:参数量 × 精度。7B FP16 = 14 GB。
-
KV Cache = 2 × 层数 × KV 头数 × 头维度 × 精度 × 当前序列长度。
-
例如 7B 无 GQA,4096 token 时 KV Cache ≈ 1 GB;32768 token ≈ 8 GB。若权重 14 GB,总 22 GB,24GB 卡可能刚好满。
⚙️ 其他可能:
-
临时激活缓冲区:推理前向会产生一些中间结果(尤其长 prompt 编码时),但很快释放。
-
框架预留:vLLM 的块池默认会占满 90% 显存,即使实际使用少,也显示高占用。但 OOM 通常不会因为预留本身,而是实际分配超限。
✅ 结论:推理 OOM 先看权重精度,再看目标序列长度估算 KV Cache,两者一算即知。量化权重 + 限制 max_length 通常解之。
📏 推理长文本时 OOM,如何通过限制最大生成长度缓解?¶
💡 限制 max_new_tokens 或 max_length 是直接控制 KV Cache 最大尺寸的手段,可确保生成阶段不会无限增长致 OOM。此外,配合 KV Cache 量化、滑动窗口注意力或前缀共享也能从不同维度缓解。
🔧 具体措施:
-
设置最大生成长度:在
generate中指定max_new_tokens=512,保证总序列长度 ≤ prompt_len + 512,KV Cache 被硬限制。 -
使用更短的系统提示:如果 system prompt 很长,会占据大量 KV Cache 空间,压缩可用生成长度。压缩提示或使用共享前缀技术(PagedAttention 共享相同前缀的物理块)可减少重复存储。
-
调整 PagedAttention 块数量:vLLM 中通过
--max-model-len设定最大序列长度,超过则截断或拒绝请求。同时gpu_memory_utilization控制预留给 KV Cache 的显存比例,间接限制并发长请求数量。 -
动态限制用户输入:在应用层对输入 prompt 长度进行截断或摘要,保留关键信息,缩短序列。
-
启用 KV Cache 量化:如 FP8 KV Cache,相同显存可存放 2 倍长度 token,间接缓解 OOM。
-
采用 Streaming 输出:虽然不省显存,但可早期感知长生成的风险,如果即将超限可提前终止。
📊 示例:一张 24GB 卡,7B FP16 权重 14GB,剩余 10GB。FP16 KV Cache 每 token 约 0.5 MB,则最大 token 数 ≈ 10GB / 0.5MB ≈ 20,000。若 prompt 已 15K,则 max_new_tokens 必须 ≤ 5K,否则 OOM。据此设限。
✅ 所以,max_new_tokens 是长文本推理的安全阀,务必根据显存容量计算好上限。
🌐 推理服务中,多用户并发导致 OOM,如何动态管理 KV Cache?¶
💡 多用户并发时,OOM 往往因为同时活跃的请求过多,KV Cache 总量超出预留池。解决思路:限制并发数、抢占与交换、前缀共享、动态伸缩池大小。
🔧 vLLM 中的动态管理策略:
-
限制最大并发序列数 (
--max-num-seqs) 直接控制同时处理的请求数。超过的请求排队等待。这从源头上避免了 KV Cache 池耗尽。 -
抢占式调度与交换 (Swapping) 当新请求到来且 GPU 块不足时,可将某些低优先级或等待较久的请求的 KV Cache 块交换到 CPU 内存,释放 GPU 块给新请求,待 GPU 有空闲时再换回。vLLM 支持自动 swap,通过
--swap-space设定 CPU 交换空间大小。这是动态平衡的关键。 -
KV 块共享(前缀缓存) 若多个请求共用相同的 system prompt,PagedAttention 可使它们映射到同一物理块,只需存一份。显著降低并发时前缀重复存储开销。
-
动态调整块池(通过
gpu_memory_utilization) 虽然这个参数是启动时固定,但可以在启动时设置为略小值,留一些显存给波动。运行时,vLLM 的调度器会根据当前活跃请求和空闲块数决定是否接受新请求。 -
请求级超时与回收 设置最大生成 token 数或超时时间,防止个别请求占据块过久。强制终止超长生成,立即回收 KV Cache 块。
-
自适应批次大小 推理引擎动态决定每个 step 执行多少请求,避免一次性接受太多请求造成峰值。
📊 效果:在没有 swap 时,并发上限受总 KV 块数卡死;加入 swap 后,可允许短时超载,通过 CPU 内存缓冲,吞吐提升且不 OOM。
✅ 因此,动态管理 KV Cache 的核心是抢占调度、交换和前缀共享,实现显存的弹性利用。
🎛️ 使用 vLLM 时,为什么 gpu_memory_utilization 参数很重要?¶
💡 gpu_memory_utilization 决定了 vLLM 在初始化时预留多少比例的 GPU 显存用于 KV Cache 块池,直接影响到可并发的请求数、最大序列长度以及是否能容纳模型权重。设太小,剩余显存放空,吞吐低;设太大,可能导致权重加载或临时缓冲区 OOM。
🔍 深入解析:
-
vLLM 启动时,先加载模型权重,剩余显存 ×
gpu_memory_utilization全部分配给 KV 块池。 -
例如 24GB 卡,权重占 14GB,剩余 10GB,若设置 0.9,则块池大小 = 10GB × 0.9 = 9GB。剩下的 1GB 留作激活等开销。
-
过大风险:设 1.0 会尝试占满全部剩余显存,若模型前向需要临时分配一些张量(如长 prompt 编码),可能 OOM。因此默认 0.9 留出安全缓冲。
-
过小浪费:设 0.5 则仅用 5GB 作池,可支持的总 token 数少,并发能力不足,大量显存闲置。
-
分块调度核心:池里全是相同大小的 KV 块,其数量决定了系统能同时维护多少个 token 的 KV Cache。该参数直接换算成最大 token 容量。
📊 调优建议:
-
若观察到 OOM 在第一个请求或 prompt 处理阶段,适当降低该值(如 0.85)。
-
若显存尚有较多空闲且请求排队,提高至 0.95 以增大池。
-
结合
max-model-len一起调,确保池大小能满足模型最大长度的单请求需求。
✅ 所以,gpu_memory_utilization 是 vLLM 显存管理的总阀门,需根据实际权重、流量和显存量精确调整。
📦 模型加载时 OOM,但模型文件大小远小于显存,为什么?¶
💡 模型文件通常是压缩或分片存储格式(如 .bin 或 .safetensors),其磁盘大小不等于显存占用。加载后权重会以更高精度展开,并伴随优化器状态、CUDA context 和框架缓存,往往数倍于文件大小。
🔍 具体原因:
-
精度膨胀:磁盘保存可能是 FP16(2 字节/参数),但加载到 GPU 后可能转为 FP32(4 字节)或框架自动保留副本。即使模型文件 14GB(7B FP16),加载后占用 14GB,尚可。但有些 checkpoint 保存时用了压缩(如分片+序列化),磁盘小,内存大。
-
多份副本:训练时通常会加载模型,并创建 optimizer、梯度等。但仅模型加载阶段,若使用 DeepSpeed ZeRO-3 或 FSDP,会存在参数收集缓冲区。即使推理,某些框架(如 PyTorch)可能保留参数的 FP16 和 FP32 master copy。
-
CUDA 上下文与驱动开销:GPU 初始化会占用约 0.5-1 GB。
-
框架预分配:如 vLLM 的块池、DeepSpeed 的通信缓冲,在模型加载后立即占据大量显存,可能显示高占用。
-
文件损坏或格式不匹配:加载错误文件可能导致异常分配,比如误加载为 FP32 权重却当 FP16 处理,占用翻倍。
📊 举例:一个 7B 模型文件 pytorch_model.bin 大小 13.5GB(FP16),加载到 GPU 后占用 14GB。但若同时加载 tokenizer 或其他组件,并开启 device_map="auto",可能会 CPU/GPU 分散,但仍可能触发额外拷贝。若显存 16GB,看似 13.5<16,但实际 14GB + 1GB context = 15GB,接近极限,可能分配失败。
✅ 排查:检查加载时精度,确认 torch_dtype;用 device_map="cpu" 先加载到 CPU 测大小,再逐步移 GPU。避免框架自动预留过多。
🐳 Docker 容器内 OOM,和宿主机显存有什么关系?¶
💡 Docker 容器本身不提供显存隔离,容器内进程直接使用宿主机 GPU 物理显存。如果宿主机显存被其他容器或进程占满,容器内训练/推理就会 OOM。此外,Docker 的 --memory 限制只能限制 RAM,不能限制 GPU 显存(需额外方案)。
🔗 关系:
-
默认情况下,Docker 容器内的
nvidia-smi看到的是宿主机全部的 GPU 显存。多个容器可能争抢同一 GPU 显存,未协调时导致一个容器 OOM。 -
可通过 NVIDIA Container Toolkit 结合 GPU 共享方案(如 MIG、vGPU 或 Hami)对显存进行限制,但必须显式配置。如果没做限制,容器间相互影响。
-
宿主机上其他非容器化进程(如另一训练任务)同样占用显存,可能导致容器内看似空闲但实际不足。
🔧 解决方案:
-
使用 Kubernetes 的资源限制并配合设备插件(如
nvidia.com/gpumem)约束每个 Pod 的显存。 -
单机 Docker 可利用环境变量
NVIDIA_VISIBLE_DEVICES选择特定 GPU,并搭配nvidia-smi监控。 -
多容器共享同一 GPU 时,应用层自己限制(如设置
gpu_memory_utilization)或使用 MIG 硬件隔离。
✅ 因此,容器 OOM 本质是宿主机显存争抢问题,需在调度层面隔离或限制。
👻 为什么有时 PyTorch 报告 OOM,但 nvidia-smi 显示显存未满?¶
💡 PyTorch 报告 OOM 时,是指缓存分配器无法找到足够大的连续显存块来满足本次分配请求,而 nvidia-smi 显示的是进程占用的总显存(可能包含大量闲置碎片或缓存),所以总占用未满但无连续大块,导致 OOM。
🔍 核心原因——显存碎片:
-
PyTorch 的 CUDA 缓存分配器为减少
cudaMalloc调用,会保留已释放的张量空间在缓存池中。随着训练进行,频繁分配不同大小的张量,池中会出现许多不连续的小空洞。 -
当请求一个大张量(如注意力矩阵、梯度桶)时,分配器在池中找不到连续满足尺寸的空闲块,即使总空闲字节足够,也会抛出 OOM。
-
nvidia-smi 只看到进程
used总量,不清楚连续块可用性,所以显示未满。
🩻 诊断方法:
-
使用
torch.cuda.memory_summary()查看缓存池的空闲块分布,会显示最大连续空闲块大小。若该值远小于所需分配量,证实为碎片。 -
对比
torch.cuda.memory_allocated()和memory_reserved(),若差值大且 allocated 未达上限,则是碎片。
🔧 缓解:
-
设置
PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True允许内存段动态扩展,减少碎片。 -
定期调用
torch.cuda.empty_cache()释放缓存(但会短暂拖慢速度)。 -
使用统一内存分配器或预分配大块池(如 vLLM)。
✅ 所以,nvidia-smi 看的是容量,PyTorch 看的是连续性,两者分歧是碎片惹的祸。
🎭 在多卡训练中,只有一张卡 OOM,其他卡正常,可能是什么原因?¶
💡 这种不对称 OOM 通常源于负载不均衡:不同卡上的数据、模型层、激活或优化器状态分布不均,导致某卡峰值过高。
🔍 常见原因:
-
数据分布差异:如果使用不均匀的数据分割(如某些卡处理显著更长的样本),激活显存差异大。
-
流水线并行(PP)的层分配不均:首卡可能有 Embedding 层,末卡有输出头,参数和激活更多,若切分不调整,易 OOM。
-
张量并行(TP)头数分配:在 TP 中,若某卡负责的注意力头数更多(设计错误),权重和激活不对称。
-
ZeRO 分片不完美:ZeRO-3 虽然分片,但某些小参数(如 bias)可能被复制到所有卡,叠加通信缓冲区导致某卡额外占用。
-
通信任务差异:在部分 Mesh 拓扑中,某卡可能承担更多的 NCCL 集合通信根节点角色,额外分配缓冲。
-
CUDA 上下文或驱动版本差异:极少见,不同卡可能由于硬件微小差异或温度降频影响内存分配行为?
🛠️ 解决:
-
检查每卡
nvidia-smi和torch.cuda.memory_allocated()的差异。 -
对于 PP,手动调整层分配或使用自动平衡工具。
-
确保数据加载中
DistributedSampler正确且不做额外排序导致某卡序列长度集中。 -
禁用或减小通信缓冲(
overlap_comm=False)。 -
尝试关闭 NCCL 的
PXN等特性。
✅ 因此,单卡 OOM 往往是模型或数据切分不均匀的信号,需针对性均衡。
🔍 如何通过 Profiler 找出显存分配的热点?¶
💡 使用 PyTorch Profiler 的 Memory 视图或 Nsight Systems,可以记录每个算子的显存分配/释放时间线、峰值大小和调用堆栈,直指占用最大的操作和代码行。
🔧 步骤:
- PyTorch Profiler + TensorBoard
with torch.profiler.profile(
activities=[torch.profiler.ProfilerActivity.CPU, torch.profiler.ProfilerActivity.CUDA],
record_shapes=True,
profile_memory=True,
with_stack=True,
) as prof:
model(input)
prof.export_chrome_trace("trace.json")
在 TensorBoard 中打开,使用 Memory Viewer 按 Allocation Size 排序,即可看到哪些算子分配了最大块的显存,双击可查看 Python 调用栈。
-
Nsight Systems 捕获 CUDA API 和内存分配事件,提供更底层视图,适合分析碎片和分配时机。
nsys profile -o report python train.py。 -
PyTorch Memory Snapshot
torch.cuda.memory._record_memory_history(max_entries=100000)
# 训练
torch.cuda.memory._dump_snapshot("snapshot.pickle")
- 使用 PyTorch 官方 visualizer 生成交互式图表,显示每个张量的生命周期和累积占用。
🎯 查找热点:按算子类型聚合,通常会发现 aten::linear、aten::scaled_dot_product_attention 和 aten::embedding 占据最多。然后可针对性优化(如 FlashAttention、梯度检查点)。
✅ Profiler 是显存优化的显微镜,不用它就像闭眼调参。
📊 训练中如何监控显存的实时使用情况?用什么工具?¶
💡 可用命令行、Python API 和仪表盘三类工具:
-
命令行轻量监控:
nvidia-smi循环,nvtop图形化。 -
代码内监控:
torch.cuda.memory_allocated()打印趋势。 -
全功能仪表盘:
Weights & Biases、TensorBoard可记录显存曲线;Prometheus + Grafana配合nvidia-dcgm实现集群级监控。
🔧 实操示例:
-
训练循环中每 10 step 记录
torch.cuda.max_memory_allocated()和memory_reserved()到 TensorBoard。 -
使用
py-spy或memray无法直接看 GPU,但可辅助 CPU 内存泄漏。
✅ 实时监控能尽早发现泄漏和趋势,避免突然 OOM。
⚡ 遇到 OOM 后,你如何快速决定是减小模型、序列长度还是 batch size?¶
💡 按“优化代价→任务需求”矩阵快速决策:先尝试无损或微调的方法,最后才考虑牺牲模型容量。
🔢 决策顺序:
-
启用/升级显存优化技术:FlashAttention、梯度检查点、混合精度(若未开),这些几乎无精度损失。
-
减小 batch size + 梯度累积:保持有效 batch 不变,牺牲训练速度,但不影响模型质量。
-
缩短序列长度:如果任务允许截断或滑动窗口,显存节省明显(尤其 O(L²))。代价是可能丢失上下文。
-
使用 PEFT/LoRA:改为冻结主干,只训练少量参数,大幅降低权重和优化器占用,可保持较大 batch 和序列。
-
减小模型(层次/维度):最终手段,影响模型能力。
📊 决策树:
-
OOM 发生在反向传播 → 优先减小 batch 或开启梯度检查点。
-
OOM 发生在前向 → 通常序列太长或模型太大,先查序列长度,再考虑模型。
-
OOM 在优化器 update → 优化器状态过大,启用 ZeRO-1 或 8-bit Adam。
✅ 快速决策依赖对显存分布的理解,平时做好 profiling,临阵不乱。
🚨 推理服务中,如何设置显存告警阈值?¶
💡 利用 Prometheus + GPU Exporter 或 nvidia-dcgm 收集显存指标,在 Grafana 中设置告警规则,当显存使用率超过预设百分比(如 90%)并持续一段时间时触发通知。
🔧 具体方案:
-
DCGM Exporter:采集 GPU 各项指标,暴露给 Prometheus。
-
PromQL 示例:
DCGM_FI_DEV_FB_USED / DCGM_FI_DEV_FB_TOTAL > 0.9。 -
告警动作:发送到 Slack/PagerDuty,或触发自动扩容、请求限流。
-
应用层自监控:在推理框架中嵌入回调,当
num_free_blocks低于阈值时主动返回压力信号。
✅ 设定阈值需结合正常波动幅度,避免频繁误报。
💧 什么是“显存泄漏”?如何检测?¶
💡 显存泄漏指程序在运行过程中,不再需要的张量未被释放,导致已分配显存单调递增,最终耗尽所有可用显存。本质是对象引用未断开。
🔍 检测方法:
-
观察单调增长:每步打印
torch.cuda.memory_allocated(),若持续上升不回落,基本可判定泄漏。 -
定位泄漏源:使用 PyTorch Profiler 的 Memory 视图,结合
with_stack=True找出哪些分配从未释放,查看对应 Python 代码行。 -
使用
gc和weakref:Python 层面检查循环引用。 -
内存快照对比:
torch.cuda.memory._record_memory_history()记录两个时间点的快照,对比哪些张量新增且未死。 -
常见泄漏场景:记录 loss 到列表、缓存了非叶子张量、自定义层保留输入引用、多线程共享张量等。
✅ 检测泄漏 = 监控 + 快照对比,找到“只分配不释放”的代码。
💬 多轮对话的显存泄漏通常出在哪?¶
💡 多轮对话中,最常见的是 KV Cache 管理不当:历史轮次的 K、V 未被正确回收,或者每次新轮次都复制完整历史导致旧版本未释放。
🔍 具体原因:
-
缓存历史轮次的完整编码:有些实现会将每轮拼接后的完整序列重新编码,生成新的 KV Cache,而旧的 Cache 未被释放,不断累积。
-
对话管理器存储了全部历史的张量引用:如保留了每个轮次的 logits 或 hidden states。
-
动态 batch 中会话生命周期过长:推理服务未及时终止并释放已完成对话的 KV 块。
-
PagedAttention 块泄漏:如果请求结束但对应的物理块未被标记为空闲,导致池子慢慢耗尽。
🛠️ 解决:
-
使用增量式 KV Cache(只处理新 token,历史 K、V 直接追加)。
-
设定对话最大轮次或 token 上限,主动清理。
-
确保推理引擎正确处理请求结束时的块回收(vLLM 通常自动)。
-
监控
num_free_blocks趋势,若持续下降说明存在泄漏。
✅ 多轮对话的显存管理核心就是生命周期管理:何时创建、何时回收 KV Cache。