AI 加速器系列 · 第 1 篇
一台 GPU 装不下模型,训练要跑三个月------这就是大模型训练的真实处境。分布式训练不是在"加速",它是在"救火"。当 LLaMA 405B 用 16,000 张 H100 训练 54 天时,你面对的不是"怎么写训练循环",而是"怎么让 16,000 张卡像一台机器一样工作"。本文从数据并行的本质讲起,到 Ring AllReduce 通信原语,再到 PyTorch DDP 实战,帮你建立分布式训练的完整心智模型。
1. 为什么单卡不够
半精度(FP16/BF16)下 1 个参数约 2 bytes,加上 Adam 优化器状态(fp32 param + momentum + variance = 12 bytes/param),实际显存需求约参数的 14 倍:
scss
模型 参数量 半精度显存 单张 A100(80G)
GPT-2 XL 1.5B 3 GB 轻松
LLaMA 7B 7B 14 GB 可以
Falcon 40B 40B 80 GB 勉勉强强
LLaMA 70B 70B 140 GB 做梦
LLaMA 405B 405B 810 GB 洗洗睡吧
即使奇迹发生,一张卡装下了,时间呢?GPT-3 175B 单卡 A100 也要跑 127 天。1024 张 A100:1.5 天。
分布式训练的本质:用通信换时间,用多卡换显存。"通信"二字,就是整个分布式训练工程的灵魂。
2. 数据并行(Data Parallelism)原理
思路朴素到令人发指------每张 GPU 持有完整模型副本,各自吃不同 mini-batch:
ini
GPU 0 GPU 1 GPU 2 GPU 3
[完整模型副本] [完整模型副本] [完整模型副本] [完整模型副本]
| | | |
Batch[0:8] Batch[8:16] Batch[16:24] Batch[24:32]
| | | |
Forward Forward Forward Forward
| | | |
Backward Backward Backward Backward
| | | |
grad_0 grad_1 grad_2 grad_3
\ | | /
+-------------------+---------------------+-------------------+
|
AllReduce 求平均
avg = (grad_0+grad_1+grad_2+grad_3)/4
|
+------------+-------+-------+------------+
| | | |
optimizer optimizer optimizer optimizer
| | | |
(四张 GPU 上的模型权重现在完全一致)
Forward 和 Backward 完全独立,唯一的同步点是梯度聚合。核心问题:梯度怎么同步?
3. Ring AllReduce:梯度同步的核心算法
3.1 朴素方案为什么不行
最直观的想法:搞一台 Parameter Server 收集所有梯度求平均再广播回去。
lua
GPU 0 --+ GPU 0 <--+
GPU 1 --+--> [PS: 收 N 份, 求平均] --> GPU 1 <--+-- 广播平均梯度
GPU 2 --+ GPU 2 <--+
GPU 3 --+ GPU 3 <--+
PS 的网卡带宽是硬上限。N 越大,PS 死的越惨------这叫通信瓶颈。
3.2 Ring AllReduce 两阶段图解
把 N 张 GPU 排成逻辑环,每节点只跟左右邻居通信,梯度切成 N 等份。
阶段一:Scatter-Reduce(N-1 步)------每步传一个 chunk 给下游,收到后与本地的 reduce:
ini
初始(GPU 0 的梯度 = [a0 a1 a2 a3],切 4 份):
GPU 0: [a0,a1,a2,a3] GPU 1: [b0,b1,b2,b3]
GPU 2: [c0,c1,c2,c3] GPU 3: [d0,d1,d2,d3]
第 1 步(传 chunk0):
GPU0 发 a0 -> GPU1 / GPU1 发 b0 -> GPU2 / GPU2 发 c0 -> GPU3 / GPU3 发 d0 -> GPU0
GPU 0: [a0+d0, a1,a2,a3] GPU 1: [a0+b0, b1,b2,b3]
GPU 2: [b0+c0, c1,c2,c3] GPU 3: [c0+d0, d1,d2,d3]
第 2 步(chunk0 继续绕环,带着累积值):
GPU 0: [a0+c0+d0, a1,a2,a3] GPU 1: [a0+b0+d0, b1,b2,b3]
GPU 2: [a0+b0+c0, c1,c2,c3] GPU 3: [a0+b0+c0+d0, d1,d2,d3] <- 完整!
第 3 步(N-1=3 步完成,同时其他 chunk 各自也在流转):
GPU 0: [a0+b0+c0+d0, a1,a2,a3] <- 完整! GPU 1: [a0+b0+c0+d0, b1,b2,b3] <- 完整!
GPU 2: [a0+b0+c0+d0, c1,c2,c3] <- 完整! GPU 3: [a0+b0+c0+d0, d1,d2,d3] <- 完整!
Scatter-Reduce 结束:每 GPU 持有 1 个 chunk 的全局归约结果(互不相同)。
实际实现中 chunk 是按轮流转的(第 i 步 GPU j 发送 chunk(j-i) mod N),上面只追踪 chunk0 的传播路径以解释"数据沿环逐步累积"的核心思想。
阶段二:AllGather(N-1 步)------各 GPU 把手里完整的那个 chunk 沿环传给邻居,N-1 步后所有人集齐全部 chunk,得到完整平均梯度。
3.3 带宽分析
scss
每个 GPU 发送量:Scatter-Reduce (N-1)/N × G + AllGather (N-1)/N × G = 2(N-1)/N × G
当 N 很大时 ≈ 2G ------ 等价于每卡只发了两份完整梯度!
Ring vs Tree 对比:
| 维度 | Ring AllReduce | Tree AllReduce |
|---|---|---|
| 步数/延迟 | O(N) | O(log N) |
| 带宽利用率 | ~2(N-1)/N,无单点瓶颈 | 叶节点只发不收,根节点瓶颈 |
| 容错 | 任一条链路慢则整体慢 | 根/中间节点崩则全崩 |
| 适用场景 | 梯度 GB 级,带宽为王 | 小数据量,延迟敏感 |
大模型梯度动辄几十 GB,带宽是瓶颈而非延迟,所以 Ring 统治了 DL 训练。
NCCL 还做了 GPU Direct RDMA:GPU 之间显存直传,数据不经过 CPU/内存。这是 NCCL Ring AllReduce 能跑到接近网卡理论带宽的原因。
4. PyTorch DDP 实战
4.1 DDP vs DP:别再写 DataParallel
| 维度 | DataParallel (DP) | DistributedDataParallel (DDP) |
|---|---|---|
| 进程模型 | 单进程多线程 | 每 GPU 独占一进程 |
| Python GIL | 严重瓶颈 | 无竞争 |
| 梯度聚合 | GPU 0 搜集 → 平均 → 广播 | Ring AllReduce 对等参与 |
| 主卡负担 | GPU 0 显存/带宽压爆 | 不存在主卡 |
| 跨机 | 不支持 | 原生支持 |
| 实测速度 | 比单卡还慢 | 接近线性加速 |
DP 慢的根因:Python GIL 让多线程 Forward 变串行,梯度全塞给 GPU 0------这不是优化问题,是架构死刑。
4.2 启动方式与最小脚本
python
import os
import torch
import torch.nn as nn
import torch.distributed as dist
from torch.nn.parallel import DistributedDataParallel as DDP
from torch.utils.data import DataLoader, DistributedSampler
def setup():
dist.init_process_group(backend="nccl")
local_rank = int(os.environ["LOCAL_RANK"])
torch.cuda.set_device(local_rank)
return local_rank
def cleanup():
dist.destroy_process_group()
def main():
local_rank = setup()
model = DDP(nn.Linear(1024, 1024).cuda(), device_ids=[local_rank])
dataset = torch.randn(10000, 1024)
sampler = DistributedSampler(dataset, num_replicas=dist.get_world_size(), rank=dist.get_rank())
loader = DataLoader(dataset, batch_size=32, sampler=sampler, num_workers=4)
optimizer = torch.optim.Adam(model.parameters(), lr=1e-4)
for epoch in range(10):
sampler.set_epoch(epoch) # 关键!忘记写则每 epoch shuffle 相同
for data in loader:
data = data.cuda()
optimizer.zero_grad()
loss = nn.MSELoss()(model(data), torch.randn(32, 1024).cuda())
loss.backward() # DDP 自动触发 AllReduce
optimizer.step()
if local_rank == 0:
print(f"Epoch {epoch} loss: {loss.item():.4f}")
cleanup()
if __name__ == "__main__":
main()
启动:torchrun --nproc_per_node=4 train.py
4.3 三个高频坑
坑 1:sampler.set_epoch(epoch) 忘记写。 不写,每个 epoch shuffle 顺序相同。代码不报错,loss 也降,但最终精度差一截------极隐蔽。
坑 2:find_unused_parameters=True 的隐性开销。 开启后 DDP 在 backward 结束后额外遍历计算图找未参与参数。固定结构网络设 False(默认),少等一个全局 barrier。
坑 3:小 per-GPU batch size 时 BN 不稳定。 torch.nn.SyncBatchNorm.convert_sync_batchnorm(model) 让所有 GPU 的 BN 统计量也走 AllReduce,大幅改善稳定性。
5. 梯度累积:用小 GPU 模拟大 Batch
5.1 原理
论文里 batch_size=4096 在 256 张 A100 上跑。你只有 4 张卡?梯度累积:
rust
正常:micro_batch -> forward -> backward -> AllReduce -> step
累积(steps=4):
micro_0 -> fwd -> bwd (不 sync)
micro_1 -> fwd -> bwd (不 sync)
micro_2 -> fwd -> bwd (不 sync)
micro_3 -> fwd -> bwd -> AllReduce -> step
等效 batch size = micro_batch × accumulation_steps × world_size
loss.backward() 默认累加梯度(不清零),4 次 backward 叠加 = 一个 4 倍大的 batch。
5.2 跟 DDP 配合的关键
python
for i, data in enumerate(loader):
if (i + 1) % accum_steps != 0:
with model.no_sync(): # 禁用本次 backward 的 AllReduce
loss = criterion(model(data.cuda()), target)
loss.backward()
else:
loss = criterion(model(data.cuda()), target)
loss.backward() # 内部自动 AllReduce
optimizer.step()
optimizer.zero_grad()
model.no_sync() 是 DDP 上下文管理器,告诉 DDP 这次 backward 别做 AllReduce。不写的话每一步 backward 都触发 AllReduce------不仅浪费带宽,而且梯度被反复平均(AllReduce 是平均不是求和,多次平均会稀释梯度)。
5.3 代价
| 方案 | 通信频率 | 显存 | 注意事项 |
|---|---|---|---|
| 大 batch | 每步 AllReduce | 高 | BN 统计量大 batch 更准 |
| 梯度累积 | N 步一次 AllReduce | 低 | BN 统计量仍是 micro-batch 级别 |
如果任务对 BN 敏感,配合 SyncBN(坑 3)或换 LayerNorm。
一句话总结
分布式训练的本质是一场通信工程------Ring AllReduce 用"环形接力"替代"中心汇聚",每张 GPU 只跟两个邻居说话就把全局梯度算完,带宽利用率接近 100%,这就是 Data Parallelism 能线性加速的根本原因。
下一篇:ZeRO 优化与混合精度训练------当数据并行不够用时,如何把模型参数、优化器状态和梯度切碎,将显存压到极致。