1. 为什么大模型训练绕不开分布式
在接触大模型之前,我训练最大的模型也就是一两亿参数的CV模型,单张V100能跑,顶多两张卡做一下DataParallel。直到开始接手真正的大语言模型训练,才发现情况完全不一样:参数规模从1亿跳到70亿、130亿、甚至更大,单单把模型权重加载到显存里就已经“卡脖子”了,更别说还要做前向、反向、保存梯度、更新优化器状态。分布式训练这个课题,从“可选项”变成了“必选项”。
对大模型从业者来说,分布式训练不只是一个工具,它是一套完整的基础设施思维。不了解它,你甚至看不懂别人训练脚本里的init_process_group、DistributedDataParallel、FSDP这些词在干什么;不了解它,你在排查“为什么训练速度上不去”“为什么偶尔卡死”时会完全无从下手。这篇文章是我自己把分布式训练的基础理论整理成的一篇笔记,适合两类人看:一类是刚入门大模型、想看明白训练代码在做什么的同学;另一类是已经能跑通单机训练、想进一步搞清楚系统设计和通信原理的工程师。
1.1 显存、算力与时间的“三座大山”
为什么单卡放不下大模型?先算一笔账。一个70亿参数的模型,如果以FP32精度存储,每个参数4字节,光权重就是28GB。训练过程中还需要保存梯度,又是28GB。优化器状态如果用AdamW,每个参数还要额外保存一阶动量(fp32)和二阶动量(fp32),又是56GB。全部加起来,单卡显存需求高达112GB,这还没算中间激活值、临时缓冲区、通信缓冲区。目前民用卡和大多数数据中心常用的卡,单卡显存也就16GB到80GB这个量级,根本塞不下。
就算模型小一点、勉强塞进了单卡,还有算力和时间的限制。假设单卡每秒能做10万亿次浮点运算(10 TFLOPS),训练一个1万亿token的小模型,理论计算量大约是模型的参数乘以token数乘以6,也就是6×10^23次浮点运算量级。用单卡跑到猴年马月都跑不完。所以分布式训练的本质就两个词:拆和合——把数据拆开、把模型拆开、把流水线拆开,分散到多张卡上并行计算,再通过通信把结果汇集起来。
这也是我建议每个准备做大模型的同学,先花一个晚上把分布式训练基础理论过一遍的原因。你不需要立刻精通集群运维,但至少要理解:数据并行是在哪个维度上切、模型并行在哪个维度上切、通信开销从哪里来、显存和速度的权衡在哪里。搞懂这些,后面看任何训练框架都会轻松很多。
2. 核心并行策略:数据、模型、流水与混合
2.1 首先理解四种并行维度
大模型的分布式策略可以分成几个维度,我不按教科书顺序讲,而是按“从最容易理解到最复杂”的顺序来讲,更符合实际学习的思路。
第一种是数据并行(Data Parallelism,DP)。这是最直觉的思路:把一批训练数据切成多份,每张卡分一份,每张卡上放一份完整的模型副本。每张卡独立做前向和反向,算出各自的梯度,然后通过集合通信把梯度同步成一样的,再各自用这个同步后的梯度更新自己的模型副本。好处是简单,坏处是模型必须能塞进单卡。我最早接触多卡训练时用的就是DP,代码上几乎不需要改模型本身,torch.nn.DataParallel一行包装就好了。
第二种是模型并行(Model Parallelism,MP),更准确地说是张量并行(Tensor Parallelism,TP)。当模型太大放不进单卡时,把某层内部的矩阵运算拆开,分到多张卡上,每张卡只算矩阵乘法的一部分。比如一个线性层Y = XW,可以把权重矩阵W按列切成两块,两张卡分别算Y1 = XW1、Y2 = XW2,最后拼起来。这样每张卡的显存负载降下来了,但每次前向和反向都要做额外的通信。
第三种是流水线并行(Pipeline Parallelism,PP)。它把网络按层切成若干段,每个设备负责其中一段。比如一个12层的模型,切成4段,每张卡负责3层。数据像流水线一样,先经过第一段再传到第二段。这里最大的问题是气泡(bubble)——当第一段算完第一批数据发给第二段时,第一段要等后面的数据送进来,或者等第二段把梯度传回来,中间会出现空闲时间。
第四种是混合并行。现实中训练大模型,基本不会单一使用某一种策略,而是按模型规模、集群拓扑、通信带宽综合设计,常见的有3D并行和4D并行。这个不用急,先把单种策略的原理和边界条件搞清楚,混合并行就是它们的排列组合。
2.2 四种并行策略的核心权衡
下面这张表是我整理出来的对比,直接看就能快速理解各种策略的定位和代价:
| 策略 | 切分维度 | 单卡显存占用 | 通信开销 | 主要瓶颈 | 适用场景 |
|---|---|---|---|---|---|
| 数据并行(DDP) | 数据 | 高(需完整模型副本) | 中等,梯度同步 | 显存容量 | 模型能塞进单卡时提升吞吐 |
| 张量并行(TP) | 层内矩阵 | 低(每卡只存分片) | 高,每层前向/反向都要通信 | 通信带宽与延迟 | 超大单层、无法放入单卡 |
| 流水线并行(PP) | 层间分段 | 低(每卡只存若干层) | 中等,段间传递激活值 | 流水线气泡 | 层数多、串行依赖强的模型 |
| 混合并行(3D) | 数据+模型+流水线 | 可调 | 综合 | 拓扑与带宽 | 千亿级大模型训练 |
我特别想强调数据并行和张量并行的一个本质区别:数据并行是算完再同步,通信发生在反向传播后,次数少但每次通信的数据量大;张量并行的通信发生在计算过程中,每一次矩阵乘法前后都要通信,次数非常多。所以张量并行对卡间通信带宽的要求极高,多机跨节点的场景一般不太适合做TP,除非你的网络带宽非常充裕。
数据并行则是“先拆数据、内容各算各的”,通信频率低,对带宽要求相对宽松,多机扩展也更友好。这也是为什么DDP成为最常见入门方案的原因。实际训练千亿模型时,常用的是在节点内用TP和PP,节点之间用DP,这就能同时利用NVLink的高带宽和跨节点集群的可扩展性。
2.3 数据并行的同步机制值得细讲
数据并行虽然看起来简单,但同步梯度这个环节里面有细节坑。最朴素的方案是“All-Reduce”:所有卡算出梯度之后,把各自的梯度发送出去并求和,最终每张卡都拿到全量梯度的平均值。这个操作可以用一个生活类比来理解:一群学生各自做了同一张卷子的一部分题目,最后要把答案汇总,每个人都要得到完整的标准答案,那就需要把所有人的答案都广播一遍再合并。
实际实现上,PyTorch DDP会对梯度进行桶(bucket)划分:把反向传播过程中产生的梯度按参数顺序装进一个桶里,当一个桶的梯度全部计算完成后就开始通信。这样梯度计算和通信可以重叠一部分,避免了“先全部算完再一起通信”的等待。理解这个机制对排查性能问题很重要——如果你发现训练速度呈“锯齿状”、一步快一步慢,很可能就是桶的划分和通信重叠没有做好。
另外要注意梯度累积(gradient accumulation)在数据并行下的语义。梯度累积是模拟更大的batch size,但当你有N张卡时,同步后的梯度默认是所有卡梯度的平均值,如果配合累积使用,需要仔细计算学习率缩放、BatchNorm等变动,否则效果会飘。我在实操中见到最多的问题就是:加了大batch,忘了调学习率,或者累积踩了两轮但更新时机不对,导致损失曲线变得很奇怪。
2.4 张量并行与流水线并行的细节差异
张量并行在Megatron-LM中得到了经典实现。以Transformer中的MLP层为例,标准实现是先过一个线性层A,再过激活函数,最后过线性层B。Megatron的做法是:把第一个线性层的权重按列切分,输入分别算;第二个线性层的权重按行切分,把A的结果拼起来算。这样逐层交错切分,避免了中途的重复All-Reduce。列切分和行切分不能乱用,必须按矩阵乘法的维度规则来,否则形状对不上。
流水线并行在实践中最常用的是GPipe和PipeDream两种调度方式,它们的主要区别在于对气泡的处理。GPipe比较朴素,一批数据切成多个micro-batch,一个接一个灌入流水线,前向走完再统一反向;PipeDream则尝试让不同设备交替执行前向和反向任务,减少空闲。需要注意的是,流水线并行中每一层设备上的显存压力并不均匀,最重的往往是第一层和最后一层附近的激活存储,很多框架会建议给边界设备适当分配更小的batch。
这两种策略组合起来的经典模式,基本就是Megatron-Turing、DeepSpeed等框架的底层设计了。所以说,理论不是空中楼阁,所有开源框架的代码就是这些基础策略的工程实现。
3. 分布式训练的基石:通信库与集合通信
3.1 NCCL、GLOO与集合通信原语
通信效率是分布式训练的生命线。PyTorch分布式训练常见的后端有两个:NCCL(NVIDIA Collective Communications Library)和GLOO。NCCL是英伟达官方推出的集合通信库,专为GPU和GPU之间高带宽通信设计,支持NVLink、PCIe、InfiniBand等,是目前GPU训练的事实标准。GLOO是PyTorch自带的通用通信库,CPU和GPU都能用,但性能一般,通常只在调试或CPU环境下用。
集合通信原语不是只有All-Reduce,还包括:
- All-Reduce:所有设备的张量归约为一个值,再广播回所有设备,梯度同步常用。
- Broadcast:把一个设备上的张量广播到所有设备,初始化权重时常用。
- Gather:把所有设备的张量收集到一个设备上。
- Reduce-Scatter:把所有设备的张量归约后,按设备切分,每个设备只保留对应分片。ZeRO、FSDP中用的是这个。
- All-Gather:把所有设备的张量分片收集拼接成完整张量,再分发给所有设备,FSDP反向前需要用到。
我给学生讲这些原语时,会用食堂打饭来比喻:Broadcast就像食堂阿姨把一盘菜端到每个人面前;Gather就是每个人把菜端到阿姨那里汇总;All-Reduce就是每个人炒一个菜,最后把所有人炒的菜混匀,再给每人盛一份混合菜——每个人拿到的都是完整混合后的结果。这个类比虽然粗糙,但用来理解数据流向足够了。
3.2 通信量估算:为什么梯度同步耗时很重要
DDP的通信量与模型大小直接相关。每一轮反向传播结束,需要同步的梯度数据的量大约是模型参数量的两倍(每个梯度fp32是4字节,比fp16的模型参数多不少)。一个70亿参数的模型,假设梯度用fp32传输,那么一轮All-Reduce就有14GB的数据在集群里流动。如果是在单机8卡、NVLink带宽约600GB/s的环境中,理论上需要约0.024秒;但如果跨越节点走千兆以太网,整个代价就要翻几十上百倍,训练效率会被通信直接拖垮。
所以做分布式训练时,有一个铁律:能用NVLink不走PCIe,能用InfiniBand不走以太网,这也是为什么训练集群的造价远高于普通服务器集群。理解了通信量,就能明白FSDP和ZeRO为什么能赢——它们把“所有人同步完整梯度”变成了“每个人只同步自己负责的梯度分片”,通信总量从两倍模型大小降到了跟单设备负责的分片相当,代价是在前向反向过程中额外插入几次All-Gather。
3.3 通信拓扑对训练效率的真实影响
曾经我在一个只有廉价千兆网卡的两机集群上尝试跑8卡DDP,70亿模型训练基本卡死在通信IO里。同一批实验搬到单机8卡NVLink的机器上,同样代码,速度提升了近一个数量级。这个对比给我的冲击很大:分布式训练的性能瓶颈,往往不是GPU算力不够,而是网络带宽和延迟不够。
通信库的配置也要注意。NCCL的环状算法(Ring All-Reduce)在带宽高的时候表现好,树状算法(Tree All-Reduce)在延迟敏感的场景下更稳。实际使用中,可以设置NCCL_DEBUG=INFO来观察通信使用的算法、传输类型和耗时分布。此外,多机训练如果走的是TCP/IP,建议配置好网卡绑定和内核参数,避免NCCL探测到错误的网卡;如果卡的数量和节点数不匹配,也容易让NCCL自动选择效率较低的拓扑。
4. 实操篇:从单机多卡到多机多卡的落地过程
4.1 单机多卡环境与PyTorch DDP最小实例
在动手之前,先把环境理清。单机多卡最常用的框架就是PyTorch的torch.distributed模块,封装了DDP。很多人分不清DataParallel和DistributedDataParallel的区别,我直接说结论:新代码一律用DDP,旧的DataParallel尽量别碰——它把整个模型复制到每张卡的显存里,每次前向要同步所有输出,通信效率低,而且多线程模型调试起来很痛苦。
DDP的基本运行流程分四步:
init_process_group初始化进程组,指定后端(NCCL)、init_method(通常用env://或tcp://)、rank和world_size。- 用
torch.utils.data.distributed.DistributedSampler包装数据集,让每个进程只取自己对应的数据分片。 - 把模型放到对应GPU卡上,再用
DistributedDataParallel包装。 - 训练循环内部需要调用
loss.backward()后由DDP自动同步梯度,注意同步前要调用model.zero_grad()或optimizer.zero_grad()。
下面是我调试过的一个最小可用实例框架,可以直接参考:
import os import torch import torch.distributed as dist import torch.nn as nn import torch.optim as optim from torch.utils.data import DataLoader, Dataset from torch.utils.data.distributed import DistributedSampler from torch.nn.parallel import DistributedDataParallel as DDP def main(): # 1. 进程组初始化 dist.init_process_group(backend="nccl", init_method="env://") local_rank = int(os.environ["LOCAL_RANK"]) world_size = int(os.environ["WORLD_SIZE"]) torch.cuda.set_device(local_rank) # 2. 构造一个示意数据集 class DummyDataset(Dataset): def __len__(self): return 1024 def __getitem__(self, idx): return torch.randn(128), torch.randn(1) dataset = DummyDataset() sampler = DistributedSampler(dataset, num_replicas=world_size, rank=dist.get_rank(), shuffle=True) dataloader = DataLoader(dataset, batch_size=32, sampler=sampler, num_workers=4) # 3. 定义模型并用DDP包装 model = nn.Sequential( nn.Linear(128, 256), nn.ReLU(), nn.Linear(256, 1) ).cuda() model = DDP(model, device_ids=[local_rank], output_device=local_rank) optimizer = optim.Adam(model.parameters(), lr=1e-3) loss_fn = nn.MSELoss() # 4. 训练循环 for epoch in range(3): sampler.set_epoch(epoch) # 重要:重排数据,否则每轮epoch shuffle无效 for x, y in dataloader: x = x.cuda() y = y.cuda() optimizer.zero_grad() pred = model(x) loss = loss_fn(pred, y) loss.backward() optimizer.step() if dist.get_rank() == 0: print(f"epoch {epoch}, loss {loss.item():.4f}") if __name__ == "__main__": main()启动方式有两种。如果你用torchrun,它会自动帮你注入环境变量并生成多进程:
torchrun --nproc_per_node=4 --nnodes=1 train_ddp.py如果手动用mp.spawn启动,则需要在代码里自己设置LOCAL_RANK等参数,但我个人更推荐torchrun,因为它处理了异常恢复、环境变量注入等一堆麻烦事,省心很多。
4.2 多机多卡:网络、共享存储与启动流程
从单机扩展到多机,增加的复杂度主要在三个方面:通信网络、文件系统、启动编排。
多机训练必须保证所有节点使用同一种init方式。常用的是环境变量初始化:指定一个主节点的IP和端口,其他节点通过MASTER_ADDR和MASTER_PORT连接到主节点。需要注意的是,MASTER_ADDR填的一定是主节点的IP,所有节点都要能互相访问该端口,防火墙开好,云主机的安全组也要放行。
文件系统方面,所有节点必须能读到同一份代码和数据。如果各节点本地磁盘不同步,模型初始化权重都会不一致,训练就会“起飞即崩溃”。实际工程上最稳妥的做法是:数据和代码放在共享存储(NFS、GPFS、Lustre等)上,每轮epoch的数据读取走共享存储;实在没有共享存储,至少保证启动时把代码同步到各节点并保持MD5一致。
多机启动的示例命令如下:
# 节点0 torchrun --nnodes=2 --nproc_per_node=8 \ --rdzv_endpoint=192.168.1.10:29500 \ --rdzv_backend=c10d \ train_ddp.py # 节点1 torchrun --nnodes=2 --nproc_per_node=8 \ --rdzv_endpoint=192.168.1.10:29500 \ --rdzv_backend=c10d \ train_ddp.py4.3 显存不够时的FSDP与ZeRO方案
DDP尽管扩展性好,但每张卡上还是要维护一份完整的模型参数、梯度和优化器状态。当模型大到一定程度,DDP就显得无能为力了。这时候可以采用ZeRO的思想——把优化器状态、梯度、模型参数分片到不同设备上。DeepSpeed的ZeRO-Stage或者PyTorch原生的FSDP(Fully Sharded Data Parallel)就是这种思路的实现。
FSDP的思路很直接:把每个模型的参数分片,前向计算前通过All-Gather临时把完整权重聚齐,计算完后马上释放其他设备的参数分片,只保留自己这部分。反向计算同理。这样显存从“每个人带着完整背包”变成了“每个人只背一份行李,用的时候互相借”,适合超大模型训练。
FSDP的显存收益和通信代价需要权衡:同样规模的模型,如果通信带宽不够,FSDP反而比DDP慢。我个人的经验是:7B以内的模型,单机DDP完全够用;13B以上单卡塞不下时,优先考虑FSDP;如果是在多机超大集群训练100B级别,交给Megatron-DeepSpeed这类框架做精细的3D并行更合适。
5. 训练过程中的典型问题与排查实录
5.1 卡死与NCCL超时
分布式训练最常见的坑就是“卡死”。表现是训练跑到某一个步骤后,GPU利用率掉到0,日志停住不动,喝个水回来它还是老样子。大多数情况是集合通信死锁——某个进程在等待一个永远不会到达的消息。
比如,不同进程的数据长度不一致,导致DDP的某些rank提前退出或者没进入同步点;或者代码里自己写了gather但忘了在所有rank上同步调用;又或者网络闪断导致NCCL通信中断。排查方法,一是在启动命令里开启NCCL_DEBUG=INFO,二是看训练日志的长尾位置,三是利用torch.distributed.barrier()手动制造同步点来定位是哪个阶段卡住。正经的训练框架还会配置通信超时时间,比如torchrun默认的--timeout、或初始化进程组时的timeout参数,超时后自动报错而不是无限等待,这是救命的设置。
5.2 OOM(显存溢出)与batch size的关系
分布式训练中OOM比单卡情况更微妙一点。你以为减小全局batch size就行,但实际上由于每张卡各自跑各自的batch,OOM可能只发生在一张卡上。常见诱因有多机间数据长度不均、BatchNorm中同步统计量引入额外显存、激活值过大、自动混合精度时临时缓冲区占用过高等。
我的排查顺序:先把batch size减半测试不稳定范围;再检查是否开了torch.cuda.empty_cache()(这个未必有用但能排除缓存碎片问题);接着看模型结构和激活内存分配,如果用的是激活重计算(activation checkpointing),计算图里会多存一层,显存占用会显著下降;最后考虑梯度累积配合更小的微批。OOM之后还要小心一个问题:有些进程已经挂了但其他进程还在跑,此时如果继续训练就会形成死锁,所以最好在训练脚本里加保护逻辑,任何一个rank OOM就全局退出。
5.3 负载不均衡与效率瓶颈
训练速度上不去,GPU利用率只有百分之三四十,最好的观测工具是nvidia-smi每50毫秒采样一次,看各卡的算力利用率和显存占用是否平均。如果明显有几张卡的利用率高、另几张卡长期闲着,很可能就是数据拆分不均匀,或者并行策略与集群拓扑不匹配。
另一个我踩过的坑是DataLoader的num_workers太小。分布式训练时,每张卡要独立消费数据,如果数据加载速度跟不上GPU的消费速度,GPU训练计算会频繁等待数据。解决办法通常是把num_workers提高到4到8,并配合prefetch_factor和持久化worker。但num_workers也不是越大越好,太大会让每个进程的内存耗尽,还会产生IO瓶颈。
5.4 常见问题速查表
为方便回查,我把高频问题和对应的处理思路整理成了表格,稳定复用的概率不小:
| 表现 | 可能原因 | 排查与解决 |
|---|---|---|
| 启动后立刻报rank初始化失败 | MASTER_ADDR/MASTER_PORT配置错误、防火墙没放行 | 确认主节点IP、端口可互通,检查网络策略 |
| 训练中途某个rank挂掉 | 某张卡OOM、数据长度不齐、进程异常退出 | 减小batch、用DistributedSampler确保长度一致,设置超时 |
| 通信超时(NCCL timeout) | 网络抖动、集合通信死锁、IB/网卡选错 | 开启NCCL_DEBUG,检查网卡绑定,增大timeout初值但也要找根因 |
| 多卡利用率不均匀 | 数据并行sampler没生效、模型并行切分不均 | 检查DistributedSampler和切分规则,观测各rank的idle时间 |
| 训练速度低于预期 | 带宽受限、DataLoader瓶颈、梯度通信未重叠 | 使用Profiler分析step时间构成,优先优化最大的部分 |
| 保存checkpoint不一致 | 只在一个rank上保存或多个rank同时写同一路径 | 只允许rank0保存,或引入分布式barrier后统一存储 |
| 日志满天飞无法定位 | 每个rank都在打印 | 只在主rank打印,辅以rank字段标记,或用日志聚合 |
还有一个很值得提的经验:分布式训练出现问题时,先把“多卡”简化成“单卡”复现一遍。如果单卡能跑通,问题基本就在通信和并行逻辑上;如果单卡也跑不通,那是模型和数据的问题,跟分布式没太大关系。这个分诊思路帮我省掉过大量无意义排查时间。
5.5 大模型训练中的日志与检查点技巧
训练大模型时,日志不合规会让人彻底崩溃。我常用的策略是:只在主rank(dist.get_rank()==0)打印训练信息,其他rank出错时用单独的日志文件记录错误堆栈。多机时,每台机器的日志按节点号分文件。这样排查问题时,可以按rank和节点快速定位。
检查点保存尤其要注意:如果所有rank同时往同一个路径写,会导致文件冲突甚至模型文件损坏。标准做法是只让主rank保存,或者每个rank保存到自己的路径。但要注意,FSDP和ZeRO分片模式下,模型参数不是完整存在的,必须靠框架提供的save_state_dict和load_state_dictAPI配合保存完整参数。不要自己手动序列化模型。此外,我在保存优化器状态时也会把学习率调度器的状态一起保存,否则恢复训练时学习率曲线会“跳崖”。
6. 一些心得体会
这套分布式训练基础理论我前前后后整理了三轮才形成清晰脉络。第一轮是死记硬背概念,看什么都懂,一写代码就懵;第二轮跟着教程跑通了DDP最小实例,才真正理解进程组、rank、world size这套概念是干什么用的;第三轮是在真实多机集群上反复踩坑,把通信超时、负载不均、显存爆炸这些问题全部碰过一轮,才算把这些理论内化成了自己的经验体系。
如果让我给准备入坑大模型训练的同学一个建议,那就是:先用手头能用的机器,把DDP最小实例彻底跑明白,中间最好故意制造几个错误——比如故意让rank数量不一致、故意让某张卡OOM、故意关掉一个进程——亲眼看看会发生什么,然后再去读FSDP和Megatron的实现。这种“主动制造故障”的训练方式,比多看十篇博客都有用。
还有一个隐藏技巧:训练脚本里把torch.cuda.set_device(local_rank)和device_ids配对好,以及正确设置环境变量,能避免大量莫名奇妙的报错。千万不要依赖torch.device("cuda")这种省事写法。代码里显式指定设备,是分布式训练的基本素养。
理论虽然基础,但它决定了你日后排查问题的上限。把分布式训练的这套思维建立起来,后面看到任何新框架、新策略,你都能很快识别出它属于并行策略里的哪一类、解决了什么问题、牺牲了什么资源。这比追着热点工具跑要值钱得多。