BatchLoRA合批为什么会让四张80G显卡同时OOM

前言

一次 BatchLoRA 任务把两个各 32 样本的事件合并后,四张 80G 训练卡同时 OOM:

1
2
3
4
5
64 samples
206729 raw tokens
206848 packed tokens
max_seq_len = 8447
每张卡还差约1.58 GiB

看到四张卡都接近 80G,很容易先想到显存碎片、缓存没有释放或者模型本身太大。但对比同一 Trainer 上其他任务后发现,单个 3~4 万 token Batch 只占 21~25 GB,问题只在合并后的大事件出现。

最终根因不是 BatchLoRA 不能处理两个 adapter,而是纯 forward 事件没有复用已有的 DDP 样本分片能力,导致 20 万 token 被完整复制到每一张卡。

合批后发生了什么

调度器在同一轮 fetch 到两个 forward 事件:

1
2
event A:32 samples
event B:32 samples

BatchLoRA 将它们合并为一个 packed batch:

1
2
3
64 samples
206729 tokens
padding只有0.1%

padding 很低,说明 OOM 不是传统的“按最大长度补齐浪费”。问题在于 DDP 执行方式。

在数据并行中,模型副本会存在于每张 GPU,但样本应该按 replica 切分:

1
2
4张卡、64个样本
→ 每卡大约16个样本

forward-backward 已经有分片路径,但是纯 forward 却没有适配,合并后的 64 个样本全部送到了每张卡:

1
2
期望:每卡约5万tokens
实际:每卡约20万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
2
3
4
5
callback rank接收完整Batch
→ 按原始样本下标分片
→ 每个DDP replica执行本地forward
→ 聚合loss与指标
→ callback rank按原始顺序恢复输出

这里不仅要返回一个总 loss。用户还可能依赖:

  • 每个样本的 loss 输出
  • per-adapter metrics
  • per-slice metrics
  • active token 数
  • 原始提交顺序

所以聚合时必须带原始样本下标,不能按各 rank 返回顺序简单拼接。

最后

这次 OOM 的日志里,最显眼的是“再申请 1.58 GiB 失败”和“15 GiB reserved”,但真正有决定性的证据是:64 样本、20.6 万 token 被合并后,四张卡同时接近相同显存上限。

排查此类问题时应先确认每张设备实际拿到了什么数据

相关记录:

作者

Noah Shen

发布于

2026-07-29

更新于

2026-07-29

许可协议

评论