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 ParallelismDP。这是最直觉的思路把一批训练数据切成多份每张卡分一份每张卡上放一份完整的模型副本。每张卡独立做前向和反向算出各自的梯度然后通过集合通信把梯度同步成一样的再各自用这个同步后的梯度更新自己的模型副本。好处是简单坏处是模型必须能塞进单卡。我最早接触多卡训练时用的就是DP代码上几乎不需要改模型本身torch.nn.DataParallel一行包装就好了。第二种是模型并行Model ParallelismMP更准确地说是张量并行Tensor ParallelismTP。当模型太大放不进单卡时把某层内部的矩阵运算拆开分到多张卡上每张卡只算矩阵乘法的一部分。比如一个线性层Y XW可以把权重矩阵W按列切成两块两张卡分别算Y1 XW1、Y2 XW2最后拼起来。这样每张卡的显存负载降下来了但每次前向和反向都要做额外的通信。第三种是流水线并行Pipeline ParallelismPP。它把网络按层切成若干段每个设备负责其中一段。比如一个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分布式训练常见的后端有两个NCCLNVIDIA 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卡DDP70亿模型训练基本卡死在通信IO里。同一批实验搬到单机8卡NVLink的机器上同样代码速度提升了近一个数量级。这个对比给我的冲击很大分布式训练的性能瓶颈往往不是GPU算力不够而是网络带宽和延迟不够。通信库的配置也要注意。NCCL的环状算法Ring All-Reduce在带宽高的时候表现好树状算法Tree All-Reduce在延迟敏感的场景下更稳。实际使用中可以设置NCCL_DEBUGINFO来观察通信使用的算法、传输类型和耗时分布。此外多机训练如果走的是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(backendnccl, init_methodenv://) 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_replicasworld_size, rankdist.get_rank(), shuffleTrue) dataloader DataLoader(dataset, batch_size32, samplersampler, num_workers4) # 3. 定义模型并用DDP包装 model nn.Sequential( nn.Linear(128, 256), nn.ReLU(), nn.Linear(256, 1) ).cuda() model DDP(model, device_ids[local_rank], output_devicelocal_rank) optimizer optim.Adam(model.parameters(), lr1e-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(fepoch {epoch}, loss {loss.item():.4f}) if __name__ __main__: main()启动方式有两种。如果你用torchrun它会自动帮你注入环境变量并生成多进程torchrun --nproc_per_node4 --nnodes1 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 --nnodes2 --nproc_per_node8 \ --rdzv_endpoint192.168.1.10:29500 \ --rdzv_backendc10d \ train_ddp.py # 节点1 torchrun --nnodes2 --nproc_per_node8 \ --rdzv_endpoint192.168.1.10:29500 \ --rdzv_backendc10d \ train_ddp.py4.3 显存不够时的FSDP与ZeRO方案DDP尽管扩展性好但每张卡上还是要维护一份完整的模型参数、梯度和优化器状态。当模型大到一定程度DDP就显得无能为力了。这时候可以采用ZeRO的思想——把优化器状态、梯度、模型参数分片到不同设备上。DeepSpeed的ZeRO-Stage或者PyTorch原生的FSDPFully 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_DEBUGINFO二是看训练日志的长尾位置三是利用torch.distributed.barrier()手动制造同步点来定位是哪个阶段卡住。正经的训练框架还会配置通信超时时间比如torchrun默认的--timeout、或初始化进程组时的timeout参数超时后自动报错而不是无限等待这是救命的设置。5.2 OOM显存溢出与batch size的关系分布式训练中OOM比单卡情况更微妙一点。你以为减小全局batch size就行但实际上由于每张卡各自跑各自的batchOOM可能只发生在一张卡上。常见诱因有多机间数据长度不均、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 大模型训练中的日志与检查点技巧训练大模型时日志不合规会让人彻底崩溃。我常用的策略是只在主rankdist.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)这种省事写法。代码里显式指定设备是分布式训练的基本素养。理论虽然基础但它决定了你日后排查问题的上限。把分布式训练的这套思维建立起来后面看到任何新框架、新策略你都能很快识别出它属于并行策略里的哪一类、解决了什么问题、牺牲了什么资源。这比追着热点工具跑要值钱得多。