Skip to main content

04 - 从单卡到多卡

这一篇把 03 篇的 train.py 改成能在多张卡上跑,产出 train_ddp.py。同时要回答一个问题:多卡到底快了多少,代价是什么。

这也是整个专题里唯一必须真开多卡才能验的一篇。原理和账我会算清楚,但最终那张实测表得在租来的机器上填。

前置:03 篇的单卡训练已经跑通。

零、开始之前:多卡在并行什么

0.1 三种切法

一个模型要放到多张卡上,无非是切三样东西:

叫法切什么直觉
数据并行切数据每张卡放一份完整模型,各自算不同的数据,再把梯度对齐
张量并行切单个矩阵一个大矩阵乘法拆到多张卡上算,算完拼起来
流水线并行切层前 9 层在卡 0,后 9 层在卡 1,像流水线一样传递

真正的大模型训练是三种混着用,叫 3D 并行。

0.2 我们只用数据并行

0.5B 的模型单卡放得下(02 篇算过静态显存 8 GB),所以张量并行和流水线并行完全用不上。它们是为了「一张卡装不下一个模型」而存在的,我们没这个问题。

这一篇只讲数据并行,但会讲透:DDP 怎么工作、通信量多大、什么时候它不够用、不够用了换成 FSDP 又要付什么代价。

0.3 数据并行的基本流程

每张卡各拿一批不同的数据,各自跑前向反向,得到各自的梯度。这些梯度不一样,因为看的数据不一样。

关键在于:参数更新前必须让所有卡的梯度变成同一个值,否则各卡的参数就会越走越远,最后变成四个不同的模型。

对齐的方式是求平均。这个「把所有卡的数据加起来再分回去」的操作,就是下一节的 all-reduce。

一、集合通信

1.1 几个原语

多卡之间怎么交换数据,是一组标准操作,叫集合通信(collective communication)。NVIDIA 的实现叫 NCCL,PyTorch 底层用的就是它。

操作干什么
broadcast把某张卡的数据复制给所有卡
reduce把所有卡的数据求和(或求最大等),结果放在一张卡上
all-reduce求和,结果每张卡都有一份
reduce-scatter求和,但每张卡只拿走结果的一部分
all-gather每张卡的一小块,拼成完整的一份,每张卡都有

要记的关键关系是:all-reduce = reduce-scatter + all-gather。这个拆解决定了它的通信量,也是 FSDP 的实现基础。

1.2 ring all-reduce 的通信量

朴素的做法是所有卡都把数据发给卡 0,卡 0 加完再发回去。这样卡 0 的网卡要承担 N 倍流量,成为瓶颈。

实际用的是 ring all-reduce:把 N 张卡连成一个环,数据切成 N 块。

第一阶段 reduce-scatter,转 N-1 步,每步每张卡收发一块,结束后每张卡持有完整结果的其中一块。第二阶段 all-gather,再转 N-1 步,把每块广播到所有卡。

所以每张卡收发的数据量是:

通信量=2×N1N×数据总量\text{通信量} = 2 \times \frac{N-1}{N} \times \text{数据总量}

N 很大时这个系数趋近 2,跟卡数无关。这是 ring all-reduce 的核心优点,也是它能扩展到几千张卡的原因。

二、DDP

2.1 原理

PyTorch 的 DistributedDataParallel(DDP)做的事:

  1. 启动时把 rank 0 的参数 broadcast 给所有卡,保证起点一致
  2. 每张卡各自前向反向
  3. 反向过程中,梯度一算出来就 all-reduce,跟反向计算重叠
  4. 每张卡拿到平均梯度,各自 optimizer.step()

第 3 点是 DDP 快的关键。它不等整个反向结束才开始通信,而是把参数分成若干桶(bucket),某个桶的梯度齐了就立刻开始传,同时反向还在继续算更前面的层。通信和计算重叠,通信开销被藏起来一大半。

注意每张卡都是完整地跑一遍 optimizer.step(),输入(平均梯度)相同、初值相同,所以结果必然相同。参数不需要额外同步。

2.2 代码改动

03 篇的 train.py 改成多卡,实际只动四个地方。

一是初始化进程组、绑定设备。

def setup_dist():
dist.init_process_group(backend="nccl")
rank = int(os.environ["RANK"])
local_rank = int(os.environ["LOCAL_RANK"])
world_size = int(os.environ["WORLD_SIZE"])
torch.cuda.set_device(local_rank) # 必须在建模型之前
return rank, local_rank, world_size, torch.device(f"cuda:{local_rank}")

RANKLOCAL_RANKWORLD_SIZE 这三个环境变量由 torchrun 注入,不用自己管。rank 是全局编号,local_rank 是这台机器内的编号,单机情况下两者相同。

torch.cuda.set_device(local_rank) 必须在建模型之前调用,否则所有进程会往卡 0 上堆,直接 OOM。

二是包一层 DDP。

model = DDP(model, device_ids=[local_rank])

三是随机种子要错开。

torch.manual_seed(args.seed + rank)
np.random.seed(args.seed + rank)

我们的 get_batch 是随机取起点的,如果所有卡用同一个种子,四张卡会取到完全一样的数据,等于白算三份。加上 rank 就错开了。

四是日志和存盘只在 rank 0 做。

is_main = rank == 0
if is_main and step % args.log_every == 0:
print(...)

不然四个进程会往屏幕上刷四份一样的日志,checkpoint 也会被互相覆盖。

2.3 通信量算一算

code/parallel_math.py 的实际输出:

二、每卡每步通信量
方案 通信量 NVLink PCIe
DDP + no_sync(每步只同步一次) 1.51 GB 2.5ms 60.3ms
DDP 无 no_sync(每次累积都同步) 24.11 GB 40.2ms 964.2ms
FSDP full shard 36.16 GB 60.3ms 1446.3ms

三、通信占一步的比例(一步约 6.42 秒)
方案 NVLink PCIe
DDP + no_sync(每步只同步一次) 0.04% 0.94%
DDP 无 no_sync(每次累积都同步) 0.63% 15.02%
FSDP full shard 0.94% 22.53%

先看第一行。DDP 每步要 all-reduce 一遍全部梯度,bf16 梯度是 1.004 GB,乘上 ring 的系数 2×34=1.52 \times \frac{3}{4} = 1.5,得到 1.51 GB。

在 NVLink 上这只要 2.5 ms,占一步(6.42 秒)的 0.04%,基本免费。在 PCIe 上是 60 ms,占 0.94%,也还好。

结论是:0.5B 模型配 16 的梯度累积,DDP 的通信开销几乎可以忽略。 这个结论有点反直觉,很多人以为多卡一定被通信拖累。原因下一节说。

2.4 no_sync:最重要的一个优化

注意上表第二行,不用 no_sync 的话通信量是 24.11 GB,整整 16 倍

原因在于 DDP 的默认行为是「每次 backward() 都触发 all-reduce」。而我们有 16 次梯度累积,就是 16 次 backward(),于是同步了 16 遍。但前 15 次的梯度还没累加完,同步它们毫无意义。

no_sync() 这个上下文管理器就是用来关掉中间那些同步的:

for micro in range(args.grad_accum):
x, y = get_batch(train_data, args.micro_batch, args.seq_len, device)

last = (micro == args.grad_accum - 1)
sync_ctx = nullcontext() if last else model.no_sync()
with sync_ctx:
with ctx:
_, loss = model(x, y)
loss = loss / args.grad_accum
loss.backward()

只有最后一次累积走正常路径触发 all-reduce,前 15 次都在 no_sync() 里,梯度只在本地累加。

这就解释了 2.3 节那个反直觉的结论:梯度累积把通信摊薄了。累积次数越多,同一次通信服务的计算越多,通信占比越低。反过来说,如果 grad_accum=1,通信占比会直接乘 16。

这行代码少写不会报错,loss 曲线也完全正常,只是白白多花十几倍的通信。在 PCIe 机器上这个差别是 0.94% 对 15.02%,很可观。

2.5 用 torchrun 启动

torchrun --standalone --nproc_per_node=4 train_ddp.py --mode ddp

--nproc_per_node=4 表示起 4 个进程,一个进程一张卡。--standalone 是单机模式,多机要改成指定 master 地址和端口。

不要用 python train_ddp.py,那样只会起一个进程,环境变量也不会注入。

三、显存不够的时候:ZeRO 与 FSDP

3.1 DDP 浪费在哪

DDP 每张卡都存一份完整的参数、梯度、优化器状态。四张卡就是四份一模一样的优化器状态。

02 篇算过每参数 16 字节,其中 12 字节是优化器状态(fp32 master weights 加 Adam 的两个动量)。这 12 字节在四张卡上完全重复,是纯浪费。

ZeRO(Zero Redundancy Optimizer)要解决的就是这个。

3.2 ZeRO 的三个阶段

按切分程度分三级,一级比一级省:

阶段切什么每参数字节(4 卡)
ZeRO-0(就是 DDP)什么都不切16
ZeRO-1切优化器状态7
ZeRO-2再切梯度5.5
ZeRO-3再切参数本身4

ZeRO-3 就是 PyTorch FSDP 的 FULL_SHARD 模式。

ZeRO-3 的思路:平时每张卡只存 1/N 的参数,要用到某一层的时候临时 all-gather 出完整参数,算完立刻扔掉。这样显存里任何时刻只有一层的完整参数,其余都是分片。

3.3 显存账

code/parallel_math.py 的输出:

一、每卡静态显存(不含激活值)
方案 字节/参数 显存
DDP / ZeRO-0(什么都不切) 16.00 8.04 GB
ZeRO-1(切优化器状态) 7.00 3.52 GB
ZeRO-2(再切梯度) 5.50 2.76 GB
ZeRO-3 / FSDP full shard(全切) 4.00 2.01 GB

从 8.04 GB 降到 2.01 GB,省了四倍,正好等于卡数。

但对我们这个模型,这个收益毫无意义,因为 8.04 GB 本来就装得下,省到 2 GB 只是把用不上的显存空出来。

3.4 FSDP 的代价

省显存不是白省的。看 2.3 节那张通信表的第三行:FSDP 是 36.16 GB,比 DDP 加 no_sync24 倍

为什么这么多。FSDP 每次前向要 all-gather 一遍参数,反向还要再 all-gather 一遍(因为前向算完就扔了),梯度还要 reduce-scatter。三趟,而且每个 micro batch 都要走一遍,没法像 DDP 那样靠 no_sync 摊薄。

16×(2×参数×N1N+梯度×N1N)=36.16 GB16 \times (2 \times \text{参数} \times \tfrac{N-1}{N} + \text{梯度} \times \tfrac{N-1}{N}) = 36.16\ \text{GB}

在 NVLink 上还好(0.94%),在 PCIe 机器上是 22.53%,也就是四分之一的时间在等通信

3.5 什么时候该用哪个

情况用什么
模型 + 优化器状态单卡装得下DDP,别想别的
装不下,但差得不多ZeRO-1 或 ZeRO-2,通信代价小
差很多ZeRO-3 / FSDP full shard
单层都装不下得上张量并行,超出这一篇范围

判据很简单:先按每参数 16 字节算一遍(02 篇 8.5 节),能装下就用 DDP。

3.6 FSDP 代码

policy = functools.partial(transformer_auto_wrap_policy,
transformer_layer_cls={Block})
model = FSDP(
model,
auto_wrap_policy=policy,
sharding_strategy=ShardingStrategy.FULL_SHARD,
mixed_precision=MixedPrecision(param_dtype=torch.bfloat16,
reduce_dtype=torch.bfloat16,
buffer_dtype=torch.bfloat16),
device_id=local_rank,
use_orig_params=True,
)

几个要点。

auto_wrap_policy 决定按什么粒度切。按 Block 切是标准做法:切得太细通信次数多,太粗则一次 all-gather 出来的参数太大,省不下显存。

mixed_precision 是 FSDP 自己的混合精度配置,配了它就不要再套 autocast,两套机制会打架。这是 train_ddp.pyctx 只在 DDP 模式下启用的原因。

use_orig_params=True 让参数保持原来的结构,torch.compile 和 03 篇按维度分组的 weight decay 才能正常工作。

梯度裁剪也要换:FSDP 的参数是分片的,普通 clip_grad_norm_ 算出来的范数是错的,得用 model.clip_grad_norm_()

四、动手实测

原理和账都算完了,剩下的必须真跑。

4.1 要测的四组

# 1 单卡基线
python train.py --steps 50

# 2 DDP 四卡
torchrun --standalone --nproc_per_node=4 train_ddp.py --mode ddp --steps 50

# 3 FSDP 四卡
torchrun --standalone --nproc_per_node=4 train_ddp.py --mode fsdp --steps 50

# 4 DDP 四卡但故意去掉 no_sync(改一行代码),验证 2.4 节那 16 倍
torchrun --standalone --nproc_per_node=4 train_ddp.py --mode ddp --steps 50

--steps 50 是只跑 50 步就退出,够测速度了,不用等训练跑完。前几步要跳过(有 CUDA 初始化和编译开销),从第 10 步开始计时。

4.2 待填的实测表

配置tokens/s相对单卡MFU峰值显存
单卡待测1.00×待测待测
DDP 4 卡待测待测待测待测
DDP 4 卡(无 no_sync)待测待测待测待测
FSDP 4 卡待测待测待测待测

按前面的账,预期是这样:DDP 四卡的加速比应该接近 3.7 到 3.9 倍(不会满 4 倍,有通信和同步开销);FSDP 峰值显存明显更低但速度略慢;去掉 no_sync 在 NVLink 机器上差别很小,在 PCIe 机器上会明显掉速。

如果实测和预期差很多,那说明账算错了或者哪里有瓶颈,这比数字本身更有价值。 05 篇会把真实结果和这里的预期做对照。

4.3 故意撑爆单卡

00 篇说过,0.5B 上 FSDP 看不出好处,所以要人为制造一个装不下的场景。parallel_math.py 算了阈值:

四、多大的模型才会撑爆单卡
显卡 DDP 能放下的上限
A100 40G 1.75 B 参数
A100 80G 3.50 B 参数
H100 80G 3.50 B 参数

(按静态 16 字节每参数、留三成显存给激活值估的。)

所以把配置放大到 4B 左右,在 80G 卡上单卡和 DDP 都会 OOM,而 FSDP 能跑起来。比如:

Config(d_model=3072, n_layers=32, n_heads=24, n_kv_heads=8, ffn_dim=8192)

这个实验值得做一次。 亲眼看着 DDP 报 OOM、换成 FSDP 就跑起来,比读十遍 ZeRO 论文管用。跑通即可,不用真训。

五、租卡实操

5.1 开卡之前

01 到 03 篇反复强调过,这里汇总一次:

  • train.bin 已经生成好并传到持久化云盘(01 篇)
  • 单卡训练已经跑通,loss 会降(03 篇)
  • checkpoint 保存和恢复已经验证过(03 篇第十节第 5 条)
  • nvidia-smi topo -m 看一眼卡之间是 NVLink 还是 PCIe,这直接决定 2.3 节那张表看哪一列

5.2 checkpoint 在多卡下的坑

DDP 的 model.state_dict() 会带上 module. 前缀,因为模型被包了一层。存之前要脱掉,否则单卡加载时对不上:

raw = model.module if hasattr(model, "module") else model
torch.save({"model": raw.state_dict(), ...}, path)

FSDP 更麻烦,参数是分片的,直接 state_dict() 拿到的是分片。要用 FullStateDictConfig 把完整参数聚到 rank 0 再存。

还有一点:只让 rank 0 存盘,但所有 rank 都要走到 dist.barrier()。如果只有 rank 0 在存盘、其他 rank 继续往下跑,进程之间会失去同步,后面的 all-reduce 会挂住。

5.3 训练挂了怎么排查

多卡训练卡死是常见现象,通常没有报错,就是不动了。排查顺序:

  1. nvidia-smi 看是不是所有卡的利用率都是 100% 且不变,那多半是某个 collective 在等一个永远不来的对端
  2. 检查是不是有条件分支导致某些 rank 调了 all-reduce 而另一些没调,这是最常见的死因
  3. NCCL_DEBUG=INFO 看 NCCL 的日志
  4. TORCH_NCCL_BLOCKING_WAIT=1,让 NCCL 超时报错而不是无限等

第 2 条特别值得注意。比如「loss 是 NaN 就跳过这一步」这种逻辑,如果只有一张卡的 loss 是 NaN,它跳过了而其他三张卡还在等它 all-reduce,整个训练就死锁了。多卡代码里,所有 rank 必须执行完全相同的 collective 序列。

六、验收

  • 四张卡都在跑,nvidia-smi 看得到四个进程
  • rank 0 的日志正常,没有四份重复输出
  • DDP 四卡的 loss 曲线和单卡基本重合(同样的总 batch)
  • 加速比测出来了,和 4.2 节的预期对照过
  • no_sync 去掉前后的差异测出来了
  • FSDP 能跑,峰值显存明显低于 DDP
  • 放大到 4B,DDP OOM 而 FSDP 能跑
  • checkpoint 存了能被单卡加载(前缀脱干净了)
  • 主动 kill 一次,四卡能从 checkpoint 恢复

第 3 条是正确性的核心检查。多卡只是加速手段,不该改变训练结果。 如果 DDP 四卡的 loss 曲线和单卡差得明显,说明有 bug,最常见的是种子没错开(四张卡在算同样的数据)或者 loss 没除以累积次数。

七、常见问题

现象大概率是什么原因
启动就 OOM,且只有卡 0 满忘了 torch.cuda.set_device(local_rank)
四份重复日志没判断 rank == 0
加速比只有 2 倍左右数据加载成瓶颈;或没开 no_sync(PCIe 机器)
loss 曲线和单卡差很多种子没按 rank 错开,四卡在算同样的数据
训练卡死不动也不报错某些 rank 少调了一次 collective,看 5.3 节
checkpoint 单卡加载报 key 不匹配module. 前缀没脱
FSDP 下梯度裁剪的值明显不对用了 torch.nn.utils.clip_grad_norm_,要换成 model.clip_grad_norm_
FSDP 报混合精度相关的错autocast 和 FSDP 的 MixedPrecision 同时开了
torchrun 起不来,端口被占上一次训练的僵尸进程还在,pkill -f train_ddp
先在两张卡上验,再上四张

--nproc_per_node=2 能暴露绝大多数多卡 bug(前缀、种子、collective 不对称),而且便宜一半。两卡跑通了再加到四卡,基本不会再有新问题。