多卡训练里出现单卡 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,否则多进程日志交错后无法对位;只打印不比较,看不出不均。
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 的切分顺序都一样。
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;验证阶段如果也要跑模型,尽量放到单独进程里做,避免验证峰值叠加在训练峰值上。
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 也不是办法,那会引入新的失败,而不是解决分配不均。