多卡训练中显存分配不均导致单卡OOM的定位与解决办法

文章导读
多卡训练里出现单卡 OOM,先别急着调小 batch size。全局同时 OOM 才说明模型或 batch 整体吃不下;单卡先爆,通常意味着各 rank 拿到的输入量不一致,或者某个 rank 额外承担了评估、聚合、保存等显存负担。排查方向应该是“为什么只有这一张卡多用了显存”,而不是“怎么把全体显存降下来”。
📋 目录
  1. 先把问题分成“输入不均”和“负担不均”
  2. 定位:逐 rank 记录“压力表”
  3. 常见原因与对应动作
  4. 修改后如何验证
A A

多卡训练里出现单卡 OOM,先别急着调小 batch size。全局同时 OOM 才说明模型或 batch 整体吃不下;单卡先爆,通常意味着各 rank 拿到的输入量不一致,或者某个 rank 额外承担了评估、聚合、保存等显存负担。排查方向应该是“为什么只有这一张卡多用了显存”,而不是“怎么把全体显存降下来”。

单卡 OOM 多数来自两类原因:输入分配不均(含变长样本 padding 后某 rank 批更大),以及某个 rank 承担了评估、指标聚合、checkpoint 保存等额外计算。定位顺序:确认哪张卡在哪个阶段先爆,逐 rank 打印批大小和峰值显存,再用 memory snapshot 看多出来的张量;修复重点是修正 sampler 或把辅助操作移出该 rank。

先把问题分成“输入不均”和“负担不均”

输入不均指每个 rank 拿到的样本数或 token 数不相等;变长输入场景下,即使样本数相同,样本内 padding 也会让某个 rank 的 batch 明显更大。负担不均指代码里有 if rank == 0 之类的分支,只在某一张卡上做验证、梯度范数聚合、checkpoint 序列化等操作,这些操作产生的临时张量正好叠加在训练峰值上。两种情况的定位手段相同,但修复动作不同。

定位:逐 rank 记录“压力表”

在训练循环内每隔固定步数打印当前显存、峰值显存和输入 shape。打印要带 rank 前缀并 flush,否则多进程日志交错后无法对位;只打印不比较,看不出不均。

多卡训练中显存分配不均导致单卡OOM的定位与解决办法
import torch
import torch.distributed as dist

def step_report(step, batch, tag="train"):
    rank = dist.get_rank() if dist.is_initialized() else 0
    cur = torch.cuda.memory_allocated() / 1024 ** 2
    peak = torch.cuda.max_memory_allocated() / 1024 ** 2
    x = batch[0] if isinstance(batch, (list, tuple)) else batch
    print(f"[{tag}] rank={rank} step={step} shape={tuple(x.shape)} cur={cur:.1f}MB peak={peak:.1f}MB", flush=True)

# 在 OOM 前一步追加快照,供离线分析
if rank == oom_rank and step == crash_step - 1:
    torch.cuda.memory._dump_snapshot(f"rank{rank}_step{step}.pickle")

max_memory_allocated 是累计峰值,能看出趋势;torch.cuda.memory._dump_snapshot 是 PyTorch 内部接口,版本间可能有变化,只用于抓现场。分析快照时,重点找该 rank 比其他 rank 多出的连续张量,看它的名字和产生位置,就能判断是输入变大还是额外计算引入。

常见原因与对应动作

输入分配不均

最常见的情况是换了自定义 Dataset,却没有用 DistributedSampler,或者自定义 sampler 没有按 num_replicas 分片。此时改用带 drop_last=True 的 DistributedSampler,让每个 rank 拿到的样本数严格一致;代价是数据集末尾的少量样本会被丢弃,需要结合数据量确认是否可以接受。每个 epoch 开始前还要调用 sampler.set_epoch(epoch),否则所有 epoch 的切分顺序都一样。

多卡训练中显存分配不均导致单卡OOM的定位与解决办法
from torch.utils.data import DataLoader
from torch.utils.data.distributed import DistributedSampler

sampler = DistributedSampler(
    dataset,
    num_replicas=world_size,
    rank=rank,
    shuffle=True,
    drop_last=True,
)
loader = DataLoader(dataset, batch_size=batch_size, sampler=sampler)

变长输入不能只看样本数。建议改用按 token 数打包的 batch_sampler,并按 rank 轮询分发,避免最长的一批集中落在同一张卡;这个方式保证的是各 rank 的 token 量接近,不保证严格相等。

def token_chunks(lengths, num_replicas, rank, max_tokens):
    order = sorted(range(len(lengths)), key=lambda i: lengths[i])
    chunks, cur, cur_sum = [], [], 0
    for i in order:
        if cur and cur_sum + lengths[i] > max_tokens:
            chunks.append(cur)
            cur, cur_sum = [], 0
        cur.append(i)
        cur_sum += lengths[i]
    if cur:
        chunks.append(cur)
    return [c for idx, c in enumerate(chunks) if idx % num_replicas == rank]

某个 rank 额外承担辅助计算

检查代码里所有按 rank 分支的操作。评估指标若把整批预测 all_gather 到 rank0,峰值就堆在 rank0;checkpoint 若直接保存 CUDA 张量,序列化过程也会显著抬高显存。通常建议:聚合前先把张量 detach().cpu(),checkpoint 保存前同样先搬回 CPU;验证阶段如果也要跑模型,尽量放到单独进程里做,避免验证峰值叠加在训练峰值上。

多卡训练中显存分配不均导致单卡OOM的定位与解决办法
if rank == 0:
    state = {k: v.detach().cpu() for k, v in model.state_dict().items()}
    torch.save(state, "epoch.pt")

修改后如何验证

保留前面的 step_report,先跑若干训练步,确认各 rank 的输入 shape 和 token 数一致;再观察完整一个 epoch 的 max_memory_allocated,峰值应趋于接近。验证要覆盖训练和验证两个阶段,因为单卡 OOM 往往只出现在其中一个阶段。

如果上述路径都排除后某张卡仍稳定先爆,需要结合具体环境确认是否模型并行切分不均匀、通信后端 buffer 分配差异或驱动版本差异。临时调小 batch size 或开启梯度累计可以止血,但输入不均导致的不均无法靠降 batch 根治,反而会掩盖问题;调低显存上限或加大 swap 也不是办法,那会引入新的失败,而不是解决分配不均。