这两年只要聊到大模型,分布式训练就是个绕不开的话题。很多刚入门的同学都问过我同一个问题:为什么非得把训练摊到多张卡上,单卡硬塞不行吗?答案其实很现实:大模型训练不仅卡在显存,更卡在时间。这篇大模型基础理论笔记,我准备把分布式训练从底层逻辑到实操经验完整捋一遍,包括并行策略怎么选、通信机制怎么工作、代码怎么改、踩过的坑怎么排,尽量一次说透。适合正在入门大模型训练、或者已经跑过单卡训练准备上多卡的同学参考。
1. 为什么大模型训练必须走分布式:显存和算力两道坎
1.1 显存装不下:参数只是最基础的账
先说显存。大模型最直观的问题就是“模型太大,一张卡放不下”。很多人以为显存占用就等于参数量乘以4字节,比如一个7B模型,FP32权重就是28GB。这个算法没错,但只算了静态权重,漏掉了训练时的动态开销。
训练一个模型,显存里除了权重,还要放梯度、优化器状态、激活值。以Adam优化器为例,每个参数要额外保存一阶动量、二阶动量,通常各占4字节。也就是说,仅优化器状态就是8字节每个参数,比模型自身权重还多一倍。再加上梯度、中间激活值和通信缓冲区,实际显存占用通常是参数量的3到6倍。7B模型在训练场景下随便就是几十GB显存,一张卡根本兜不住。
这就是分布式训练要跨过的第一道坎:让多个设备共同承担模型、梯度和优化器状态的存储,而不是把所有东西都压在单个设备上。理解这一点非常重要,因为很多并行策略的设计思路,本质就是在考虑“什么东西放在哪”和“什么数据需要在设备之间流动”。
1.2 算力不够用:训练时间是指数级增长的
显存问题解决了,算力又是另一道坎。即使模型能被塞进一张卡,这张卡训练一个大规模模型也可能要跑几个月,这在工程上是不可接受的。算力不像显存那样“换个设备就完事”,它必须靠数量来堆。多卡并行后,总计算吞吐量能线性提升,训练时间才能从“几个月”压缩到“几周”甚至“几天”。
不过要注意,分布式训练的加速比并不是简单的线性增长。多卡之间要同步信息,存在通信开销;数据加载、日志、checkpoint也需要额外管理。所以分布式训练的本质,是在多加计算资源的同时,尽量降低通信和同步带来的损耗。谁能在通信和计算之间找到平衡,谁就能真正把集群吃满。
2. 并行策略全景拆解:从数据并行到混合并行
2.1 数据并行:最朴素也是最常用的方案
数据并行是大多数人上手的第一个分布式方案。核心思路很直白:每张卡都保存一份完整的模型副本,训练数据切成多份分发给不同设备,各设备独立做前向和反向计算,最后把梯度同步并汇总,再用汇总后的梯度更新参数。
这个过程有一个关键点:梯度同步。每张卡算出的梯度只是“自己看到的那批数据”的梯度,如果不进行同步,模型参数就会各走各的,训练就崩了。所以数据并行每次迭代都要让所有设备把自己的梯度交换一遍,最常见的方式就是下文会细说的AllReduce。PyTorch里封装好的DDP,做的就是这件事。
数据并行最大的优点是简单,模型代码几乎不用改,训练逻辑和单卡非常接近。缺点是每张卡都要存一份完整模型,显存开销大。模型小一点(比如几十亿参数)用数据并行很舒服;但模型大到单卡根本塞不下完整副本时,数据并行就无能为力了,这时候得靠下面两种思路。
2.2 张量并行:把每一层切开分给多张卡
张量并行解决的是“单层塞不进一张卡”的问题。它的思路是把一个算子或者一个权重矩阵切分成多个分片,分别放到不同设备上执行,最后汇总结果。
以Transformer里的矩阵乘法为例,一个线性层Y=XW,可以把权重矩阵W按列切分成W1、W2,分别放到两个设备上,然后在设备1算XW1、设备2算XW2,最后把两个结果拼接起来得到Y。这个过程看起来只是切分了一次矩阵乘法,但实际落地时,算子的切分、结果聚合、通信时机都需要精确设计,否则每个矩阵乘法都产生通信,训练效率会非常差。
张量并行的问题在于通信极其频繁。每做一次前向计算,都要进行AllReduce或者AllGather,通信量跟模型的宽度、层数强相关。所以张量并行通常会用在卡间通信带宽很高的场景,比如通过NVLink把多张卡紧密绑定。跨多台机器做张量并行,网络延迟往往扛不住。这也是为什么超大模型训练里,张量并行一般只在单机内部使用。
2.3 流水线并行:按网络层切分,让数据流动起来
流水线并行是完全不同的一种切法。它不切单个权重矩阵,而是把模型按网络层划分成若干个阶段,每个阶段放到不同的设备上。第1层到第10层放在设备A,第11层到第20层放在设备B,数据按顺序流过每个设备,像工厂流水线一样。
流水线并行看着很合理,但有一个经典问题:bubble,也就是流水线空闲率。如果只是简单地把一个batch的数据从头流到尾,每个设备在等待上下游计算时都会空闲,GPU利用率很低。解决办法是把batch划分成多个micro-batch,让不同micro-batch在不同设备上交错执行。前一个micro-batch还没算完,下一个micro-batch已经追上来了,设备之间的空闲时间就被填上了。
这里的核心体会是:流水线并行要调好micro-batch数量和阶段划分。micro-batch太少,bubble太大;micro-batch太多,通信调度又复杂。合理划分流水线阶段,尽量让每个阶段的计算量均衡,否则最慢的那一段会拖住整个训练速度。
2.4 混合并行与显存卸载:现代大模型训练的标配
现实中的大规模训练几乎不会只用一种并行策略,基本都是混合并行。数据并行负责扩大整体吞吐量,张量并行解决单层过大,流水线并行按阶段切分网络结构。三者在不同层级叠加起来,构成一个三维并行空间。在此基础上,还会配合专家并行、序列并行等策略,进一步优化特定模型结构的效率。
还有一个技术方向非常适合大模型微调和预训练:显存分级,代表方案是DeepSpeed的ZeRO和PyTorch的FSDP。它们的核心思路不再是把模型复制到每张卡,而是把模型的参数、梯度、优化器状态进行分片,各设备只保存一部分。计算时需要哪部分,就从对应设备取哪部分。这样既享受了数据并行“大家算不同数据”的吞吐优势,又避免了每张卡都保存完整模型副本的记忆体浪费。
我个人的建议是:如果你只是想让一个单卡放不下的模型能跑起来,优先考虑FSDP或DeepSpeed ZeRO,它们是当前性价比最高的选择。如果你要预训练一个真正的大规模基础模型,那就要在训练集群设计时直接考虑三维并行,把张量并行、流水线并行和数据并行一起规划进去。
3. 通信机制:分布式训练里最容易翻车的环节
3.1 梯度同步与AllReduce:分布式训练的“神经系统”
分布式训练的很多坑,本质上不是出在模型结构上,而是出在通信上。拿数据并行举例,每个设备反向计算出梯度后,必须把梯度汇总起来得到一个“全局一致”的梯度。这个过程牵扯点对点通信、同步屏障、数据合流,都要靠集合通信来完成。
以Ring-AllReduce为例,它的做法很巧妙:把参与通信的设备排成一个环,数据逐步传递和累加,最后每个设备都拿到完整的梯度平均值。这个算法的好处是通信量不随设备数量线性增长,每个设备只需要和相邻设备通信,整体负载都比较均衡。
实际工程中,梯度同步是分批进行的,不是等所有层反向计算完才一次性同步。PyTorch DDP会把梯度分成一个个bucket,每算完一组梯度就开始通信,等一个bucket满了就启动AllReduce。这样通信计算能重叠起来,训练更快。理解了这一点,你就知道为什么训练时日志里偶尔看到某个阶段耗时忽高忽低,往往就是通信和计算重叠不理想导致的。
3.2 通信拓扑对比:参数服务器架构与环型All-Reduce
早期分布式训练常用参数服务器架构:中心节点负责聚合梯度、更新参数,其他节点只做计算。这种架构逻辑简单,但中心节点的通信带宽很容易成为瓶颈。几十张卡同时往参数服务器发梯度,网络很容易被打爆。
环型AllReduce把压力分散到了所有设备上,不再有单点瓶颈,也因此成为现代分布式训练的主流通信方式。不过在设备数量非常大、跨机房部署时,AllReduce的时延和链路波动仍然不可小视。
通信拓扑的选择不是一个纯技术问题,还要考虑硬件条件。如果所有卡都在同一台机器上,经过高速卡间互联通信,延迟就低;如果跨多台机器走普通千兆万兆以太网,通信延迟和带宽就会急剧上升,这时候能用梯度累积减少通信次数,能合并通信分组,能用异步日志减少同步阻塞,每一点通信优化都直接影响训练效率。
3.3 通信量估算:怎么判断训练速度是不是被通信拖死
判断你的分布式训练瓶颈到底在计算还是通信,最简单的方法就是估算通信量和实际吞吐量。以数据并行为例,每个参数对应一个梯度,梯度通常是4字节。一个7B模型所有梯度合计28GB,Ring-AllReduce一次梯度同步的通信传输量大约是梯度的两倍,也就是约56GB。如果卡间实际带宽是50GB/s,那么一次同步至少要1秒多。如果一个batch的计算时间不到1秒,通信占比就超过50%,训练怎么都快不起来。
我通常会在开始大训练前,先做一个小的通信压测,只测数据并行的AllReduce时延,看看吞吐量是否符合预期。如果实测带宽远低于理论值,就要检查网络配置、TCP/NCCL参数、多卡通信绑核是否正常。通信出问题往往不是代码逻辑错误,而是环境问题,排查起来比模型代码还费时间,这也是分布式训练“重灾区”之一。
4. 常用分布式框架与选型建议
4.1 按场景选工具:DDP、FSDP、DeepSpeed、Megatron
选择框架不能只看名气,得看场景,配得上你现有的模型规模、硬件数量和团队经验。下面是几个常用方案的对比和适用场景:
| 方案 | 核心思路 | 适用场景 | 上手难度 |
|---|---|---|---|
| PyTorch DDP | 数据并行,复制完整模型,同步梯度 | 模型能放进单卡,想快速提升吞吐 | 低 |
| FSDP | 将参数、梯度、优化器状态分片,按需重组 | 模型超单卡显存,单机多卡或跨机微调 | 中 |
| DeepSpeed ZeRO | 分级显存优化,类似FSDP但更成熟 | 大模型微调、大规模训练,兼容性强 | 中 |
| Megatron-LM | 三维并行,张量+流水线+数据并行 | 超大模型预训练,需要精细控制并行度 | 高 |
我给自己选型的原则很简单:模型不超过单卡显存,直接用DDP;模型放不下但不想折腾并行细节,上FSDP;做预训练、模型特别大、需要精细调并行策略,才考虑DeepSpeed ZeRO加Megatron的组合。
4.2 硬件通信设施怎么配:不只看卡,还要看互联和网络
分布式训练性能受两个硬件因素影响特别大:一个是卡间互联带宽,一个是节点间网络。同机内部卡间通信能走高速互联,延迟低、带宽高;跨节点数据交换则依赖网络质量,万兆以太网和更高速的RDMA网络效果差距非常明显。
如果条件有限,只能走普通以太网,那就要尽量减小跨节点通信频率。你可以优先把张量并行的维度限制在单机内,跨节点只跑数据并行,这样每个迭代只有一次梯度同步,通信开销受限于网络但不会频繁触发。这种“单机张量并行加跨机数据并行”的布置,是用常规设备做大模型训练时非常实用的折中方案。
反过来,如果硬件条件很好,通信带宽充足,那就可以更自由地设计并行维度。但好的硬件同样需要好的配置,例如NCCL相关的环境变量、网络协议栈设置、绑核方式,都会影响实际能达到的带宽。我在实际项目中见过很多“明明设备很好,跑分却很低”的案例,大多数都是网络参数没调好,不是设备出了问题。
5. 实操笔记:把单卡训练改造为分布式训练
5.1 环境准备与启动方式
以PyTorch为例,最常见的分布式启动方式是torchrun,它负责初始化进程组,并为每个进程分配一个rank(全局编号)和local_rank(本机编号)。每个进程的序号是理解分布式代码的第一把钥匙:模型代码在每个进程里都会跑一遍,所以每个进程要使用自己对应的GPU设备。
启动命令通常长这样:
torchrun --nproc_per_node=8 train.py如果是多机训练,还要加上--nnodes和--node_rank参数。nnodes表示参与训练的机器数量,node_rank表示当前机器在训练任务中的编号。每个节点都要跑同样的启动命令,只不过参数里的node_rank对应各自机器。
在代码里,第一件事就是初始化进程组:
import torch.distributed as dist dist.init_process_group(backend="nccl")这里的backend="nccl"指定了通信后端。PyTorch支持多种后端,但GPU分布式训练几乎都用NCCL,因为它是针对GPU通信高度优化过的。初始化完成后,还要设置当前进程使用的GPU设备,通常用local_rank来指定,避免多个进程争抢同一块卡。
5.2 数据并行改造:从单卡到DDP
把一个普通单卡训练代码改成DDP,改动量其实非常小。核心是几件事:初始化进程组、把模型包装成DDP模型、使用DistributedSampler让不同进程拿不同数据,以及调整保存checkpoint的方式。
下面是一个非常简化的DDP训练流程:
from torch.nn.parallel import DistributedDataParallel as DDP from torch.utils.data import DataLoader, DistributedSampler def train(local_rank, world_size): dist.init_process_group(backend="nccl") torch.cuda.set_device(local_rank) model = MyModel().to(local_rank) ddp_model = DDP(model, device_ids=[local_rank]) dataset = MyDataset() sampler = DistributedSampler(dataset, num_replicas=world_size, rank=local_rank) loader = DataLoader(dataset, batch_size=32, sampler=sampler) for epoch in range(epochs): sampler.set_epoch(epoch) # 保证每个epoch数据顺序重新打乱 for x, y in loader: loss = ddp_model(x) loss.backward() optimizer.step() optimizer.zero_grad()这里最常被忽略的是sampler.set_epoch(epoch)。没有这一步,每个epoch的数据划分顺序会完全一样,模型在每个epoch看到的样本组合是固定的,训练效果和泛化能力都会受干扰。
还有一点要特别说明:DDP包装后的模型,真正参与参数存取时要用ddp_model.module.state_dict(),而不是直接ddp_model.state_dict()。很多人在保存checkpoint时踩过这个坑,加载时发现键名对不上,或者参数根本没同步好。
5.3 显存优化:用FSDP/DeepSpeed把模型跑起来
当模型超过单卡显存,或者你想用更小的显存跑更大的模型,FSDP和DeepSpeed ZeRO就是利器。它们的共同思路是把模型参数、梯度、优化器状态切分到多张卡上,每张卡只持有其中一部分。
以FSDP为例,代码改造很简洁:
from torch.distributed.fsdp import FullyShardedDataParallel as FSDP model = FSDP(model)但FSDP背后做的事情并不简单。前向计算时,如果某个权重分片不在当前设备上,就要触发一次AllGather把它从其他设备聚合过来;反向计算时还要再次聚合和释放参数。这带来的直接代价就是额外的通信量,所以FSDP并不是无脑提升速度,而是用通信换显存。
使用FSDP时,有几个配置需要关注:sharding_strategy控制分片策略,可以只分片优化器状态,也可以把参数、梯度、优化器状态全部打散;cpu_offload可以把部分状态暂存到CPU内存,进一步压低显存占用。根据我的经验,单机多卡微调大模型时,FSDP的分片策略从“全分片”开始调,遇到通信太频繁导致速度下降,再改用混分片,是一个比较稳妥的路径。
DeepSpeed ZeRO的用法也类似,但它在配置文件中暴露的参数更多,比如optimizer、scheduler、zero_allow_untested_optimizer、reduce_bucket_size等。这些参数单独看都不复杂,组合起来却非常影响性能。建议先按官方文档默认值跑通,再针对自身模型逐步优化,不要一上来就追求终极配置。
5.4 checkpoint:分布式训练的隐藏雷区
checkpoint在分布式训练里特别容易出问题,因为模型参数、优化器状态被分散在各卡上。如果直接在某一张卡上保存模型,保存的内容多半不完整,恢复训练时会直接报错或得到错误结果。
在DDP模式下,主进程(rank为0)约定俗成地负责保存模型和优化器状态,保存时要取module.state_dict()。在FSDP或DeepSpeed场景下,不同并行度下模型参数是分片的,直接保存每个进程的分片之后,可能并不好恢复。更推荐的做法是先用FullStateDictConfig把完整参数聚回来,再在CPU上保存,或者直接使用框架提供的统一checkpoint接口。
保存的动作只用rank 0来做,但每个rank都要调用torch.distributed.barrier()等待所有进程都到达同步点,否则其他进程可能在保存尚未完成时就继续往下跑,导致数据不一致。这个同步等待在断点续训里尤其重要,很多“恢复训练后loss异常”的问题,追溯到最后都是因为checkpoint只保存了部分rank的状态,加载时恢复不完整。
6. 常见问题与排查技巧实录
6.1 显存OOM
即使上了分布式,OOM还是最常见的报错。出现OOM先别急着加卡,先按这个顺序自查:当前用的是不是FP32,改成BF16或FP16能省一半显存;是否打开了梯度检查点来用算力换显存;全局batch size是不是设得过大,可以把batch拆小,用梯度累积替代。每一条都试过,再考虑是不是真的要加卡。
用DeepSpeed ZeRO或FSDP后仍然OOM,还有一个容易被忽略的点:激活值占的显存可能比模型本身还大。这时要重点优化的是激活值管理和重计算,而不是继续切分模型参数。激活值重计算可以让训练过程反复重新计算前向结果,避免把整条链路的中间结果都保留在显存里,效果非常明显。
6.2 训练速度上不去:通信和负载失衡
训练速度慢,先看是通信慢还是计算慢。我的排查办法是分别压测:先跑一个纯计算脚本,不看通信;再用NCCL测试工具测AllReduce带宽;最后再跑真实训练,对比时间分布。如果真实训练时间明显高于前两者之和,通常是通信计算没有重叠,或者同步点太多。
还有一种典型情况是负载不均衡:某个rank切分到的数据或计算特别多,导致其他rank都在等它。比如张量并行切分权重时,没有按维度均匀切分;流水线并行时,各阶段计算量差异过大。这时候需要重新设计并行切分方式。
6.3 数据加载拖后腿
多卡训练中,数据读取往往成为隐藏瓶颈。如果DataLoader的num_workers设置太小,GPU会长期等待数据;如果每个rank都加载同一份数据,又会造成重复训练。用DistributedSampler后还要注意,如果batch_size是单卡的批量大小,总batch size等于单卡batch乘卡数,学习率也要相应调整,否则收敛曲线会异常。
经验上,一个batch的读取时间明显大于前向和反向的时间,就要加大num_workers、开启pin_memory=True,并考虑用内存文件系统做数据缓存。在超大模型训练场景下,数据读取管道的优化效果有时候比换更大带宽的网络还要明显。
6.4 训练过程中偶发卡住或崩溃
分布式训练经常出现“跑着跑着某个rank卡住”的现象,最常见原因是通信等待超时,某个进程没进入通信调用,其他进程卡在AllReduce上。遇到这种情况,第一步是加NCCL_TIMEOUT和日志输出,定位卡在哪个通信原语上。如果是多机训练,还要检查端口开放情况和防火墙规则。
另外,随机性和复现问题在分布式环境下更加复杂。设置相同的随机种子只能保证参数初始化一致,无法完全保证数据加载顺序和通信顺序的一致。要在固定随机种子的同时,处理好DistributedSampler的shuffle和GPU随机种子,才能真正做到可复现。调试阶段,我习惯先在小规模环境、单进程模拟场景下验证逻辑,再上多卡,这样能大幅减少排查时间。
7. 我的分布式训练检查清单与经验
启动任何分布式训练任务前,我都会过一遍自己的检查清单,这里直接分享出来:
- 确认模型复制或分片策略适合当前显存和卡数;
- 第一次启动前先用一两百个batch做冒烟测试,确认loss曲线和通信都正常;
- 压测AllReduce带宽,确认通信不是性能瓶颈;
- 检查checkpoint保存和恢复流程,确保掉电断点可以恢复;
- 数据加载用DistributedSampler,并用
set_epoch保证每轮shuffle; - 保存模型时只让rank 0动手,并加上
barrier()同步; - 多机环境下确认节点间网络通畅,NCCL相关环境变量无遗漏。
踩过几次坑之后,我个人最大的体会是:分布式训练更像一个系统工程,模型代码只是其中一环。硬件网络、通信配置、显存管理、数据管道和checkpoint策略,一个环节出问题,整体训练就会被拖垮。所以不要只看模型结构本身,花时间把基础设施调稳,收益远大于换更复杂的并行策略。