BatchLoRA合批为什么会让四张80G显卡同时OOM
前言
一次 BatchLoRA 任务把两个各 32 样本的事件合并后,四张 80G 训练卡同时 OOM:
1 | 64 samples |
看到四张卡都接近 80G,很容易先想到显存碎片、缓存没有释放或者模型本身太大。但对比同一 Trainer 上其他任务后发现,单个 3~4 万 token Batch 只占 21~25 GB,问题只在合并后的大事件出现。
最终根因不是 BatchLoRA 不能处理两个 adapter,而是纯 forward 事件没有复用已有的 DDP 样本分片能力,导致 20 万 token 被完整复制到每一张卡。
合批后发生了什么
调度器在同一轮 fetch 到两个 forward 事件:
1 | event A:32 samples |
BatchLoRA 将它们合并为一个 packed batch:
1 | 64 samples |
padding 很低,说明 OOM 不是传统的“按最大长度补齐浪费”。问题在于 DDP 执行方式。
在数据并行中,模型副本会存在于每张 GPU,但样本应该按 replica 切分:
1 | 4张卡、64个样本 |
forward-backward 已经有分片路径,但是纯 forward 却没有适配,合并后的 64 个样本全部送到了每张卡:
1 | 期望:每卡约5万tokens |
所以四张卡会在同一个算子附近一起接近显存上限。
为什么堆栈落在FLA算子
最终 OOM 出现在 Qwen3.5 Gated Delta Rule 的:
1 | o = torch.empty_like(v) |
这个算子只是最后申请额外 1.58 GiB 时失败,并不代表它泄漏显存。前面的大 Batch 已经把每卡推到 78 GB 左右,任何需要一块较大连续输出的后续算子都可能成为最后一根稻草。
错误日志中还有约 15 GiB reserved but unallocated,确实可能存在碎片,但即使开启 expandable_segments,也只是缓解分配形状,不能修复“每卡处理了四倍数据”这个根因。
修复
让纯forward也按DDP切样本
修复后,每个 replica 只处理自己的本地样本。
流程是:
1 | callback rank接收完整Batch |
这里不仅要返回一个总 loss。用户还可能依赖:
- 每个样本的 loss 输出
- per-adapter metrics
- per-slice metrics
- active token 数
- 原始提交顺序
所以聚合时必须带原始样本下标,不能按各 rank 返回顺序简单拼接。
最后
这次 OOM 的日志里,最显眼的是“再申请 1.58 GiB 失败”和“15 GiB reserved”,但真正有决定性的证据是:64 样本、20.6 万 token 被合并后,四张卡同时接近相同显存上限。
排查此类问题时应先确认每张设备实际拿到了什么数据
相关记录:
BatchLoRA合批为什么会让四张80G显卡同时OOM
https://blog.novashen.top/2026/07/29/tech/AI Infra/BatchLoRA合批为什么会让四张80G显卡同时OOM/