尧图网站设计 尧图网站设计YAOTU DESIGN
ARTICLE DETAIL

资讯详情

深耕网站设计与一线实操的经验洞察。

分布式训练核心指南:并行策略、AllReduce与容错实践

分布式训练核心指南:并行策略、AllReduce与容错实践 1. 从显存不够用说起分布式训练解决的不是一个问题是三个大概两年前我第一次尝试在单卡上训练一个十亿参数规模的模型。当时用的是旗舰级数据中心卡显存48GB听着已经很唬人了。结果模型参数一加载优化器状态一分配再塞进一个batch的数据显存直接爆掉OOM报错红成一片。我当时的反应是换更大的卡。但很快发现这条路走不通——你换到80GB的卡模型能塞进去了可训练速度慢到让人怀疑人生。一个batch的前向反向要算十几秒照这个速度跑完一个完整的训练周期得按年计算。那是我第一次认真审视分布式训练这件事。很多人提到分布式训练脑子里只有一句话多卡并行加速训练。但真正动手之后你会发现分布式训练解决的是三个完全不同的痛点显存装不下、算力跟不上、单点扛不住。这三个痛点对应着完全不同的技术路线混为一谈是后面所有踩坑的根源。先说显存。深度学习模型训练时的显存消耗不是参数大小这么简单它由四部分构成模型参数本身、优化器状态Adam一阶二阶动量几乎是参数量的两倍起步、前向计算过程中保存的激活值、以及每个batch送入的数据和临时缓冲区。以GPT-3这种1750亿参数的稠密模型为例仅模型参数就占约350GB内存按半精度存储也要350GB优化器状态再用FP32存一份轻轻松松超过700GB。单张GPU卡哪怕堆到下一代短期也不可能物理装下。这就逼出了模型并行与流水线并行这套把模型拆开的思路。再说算力。哪怕你有一张理论算力几十TFLOPS的卡甚至能把模型勉强装进去训练时间依旧不现实。举个具体的数字一个百亿参数模型训练大约需要10^17次浮点运算量。单卡算力按100 TFLOPS还要打折扣算理想情况持续全速跑也要快半个月真实场景要算上前向反向的重复计算、数据加载瓶颈、通信等待往往放大三到五倍。要在这个时间周期上做迭代实验任何算法团队都会崩溃。于是数据并行登场目标就是把同一份计算复制到多张卡上按数据分片把吞吐量线性放大。还有第三个稳定性。单机训练时一个节点某个CUDA报错、掉驱动、断电最多你自己重跑。可分布式训练中几十台机器一起跑任何一台出问题如果不做容错整个任务跟着陪葬。这个问题比前两个更隐蔽也更容易在项目中期爆发。所以你在看任何分布式训练框架的设计时脑子里要有这根弦它究竟在优化显存还是加速计算还是保证任务不死。很多方案看似复杂一旦先搞清楚它服务的目标底层逻辑就顺了。我还想泼一盆冷水不是所有训练任务都该上分布式。如果你的模型单卡能装下训练时间在几个小时以内多卡引入的通信开销和运维成本可能比收益还大。尤其是一个batch都跑不满一张卡的小项目分布式纯属自找麻烦。分布式训练是锦上添花不是银弹。2. 三种主流并行范式数据并行、张量并行、流水线并行分布式训练的世界里并行策略基本可以归成三大类数据并行、张量并行也叫模型并行/算子内并行、流水线并行。它们从不同维度对训练过程做切分。2.1 数据并行最简单、最常用、最适合起步数据并行的核心思想特别朴素把训练数据集切成多份每张卡上放一份完整的模型副本各自拿不同的数据同时做前向反向算完梯度后把梯度汇总、求平均再更新每一份模型参数。这里有个关键点每张卡的模型参数初始值必须完全一致否则各卡各自迭代模型早就发散到姥姥家了。所以数据并行的每一次迭代末尾都离不开一次全局梯度同步。在PyTorch DDP里这一步是通过进程间通信把每张卡的梯度做AllReduce完成的。数据并行最大的优点是实现简单几乎不需要改动模型结构。你写好的单卡训练代码包一个DDP改一下数据采样器基本就能跑。缺点是它只能解决算力不够的问题解决不了显存装不下的问题——模型本身还是完整地放在每张卡上。2.2 张量并行把一个算子的计算拆到多张卡当单卡显存放不下整个模型时就得考虑把模型切开。张量并行是切得最细的一种方式把某一层的权重矩阵拆成几块分别放在不同卡上。比如一个线性层Y XW权重W是4096×4096的矩阵我可以把它按列切成四块每块4096×1024放在四张卡上让每张卡只算Y的一个片段。前向时每张卡拿着完整输入X和自己的权重分块算出一个部分结果最后通过AllGather把输出拼回完整矩阵。这种并行方式能省显存但通信量也相当大。因为它每一层前向都要做一次全量特征拼接而且计算过程中每张卡都需要拿到完整的输入X这本身就是一份显存开销。张量并行通常只在模型层内维度非常大的时候才值得用常见的比如Transformer的注意力多头、MLP中间层都是天然的拆分点。做张量并行要付出什么代价最典型的是你的模型代码不能再是写一份到处运行的朴素PyTorch得考虑分片逻辑、通信原语插入、序列化并行区域。这也是为什么大家通常不手写而直接用Megatron-LM或DeepSpeed这类框架的原因。2.3 流水线并行按层切分让接力棒跑起来流水线并行比张量并行更粗粒度它按神经网络层来切分算子把网络前几层放在GPU 0中间层放在GPU 1后几层放在GPU 2数据像流水线一样依次流过所有卡。最原始的按层切分有一个致命问题在某一个时刻只有一张卡在算其他卡全部闲着。GPU利用率直接打一折谁用谁亏。于是有了micro-batch切分和流水线调度算法。经典的GPipe把一个小batch再切成更小的micro-batch让前一个micro-batch计算完第一层后第二层卡立即开始处理它同时第一层卡就能处理下一个micro-batch。这才形成流水线重叠。后续的PipeDream、Interleaved调度进一步减少气泡比例。流水线并行省显存、通信量相对小只需要层间传激活值和梯度但调度复杂还伴有一个独特的麻烦梯度更新延迟。因为一个完整batch的样本要经过整条流水线才完成一次前向反向不同micro-batch的梯度产生时间不一致这会影响BatchNorm这类依赖全局统计量的层。2.4 混合并行现实世界里的唯一答案现实中的大模型训练很少只用一种并行。通常是数据并行×张量并行×流水线并行一起上。比如一个模型用张量并行把单层算力扩展到4张卡用流水线并行把层切到8组每组内部再套一个4路数据并行一共32张卡协作。刚接触这套概念的人很容易被绕晕我的理解方式是把三个维度当成切蛋糕的三种方式数据并行是蛋糕不变多复制几个蛋糕师傅同时切不同块张量并行是把蛋糕切成小块分给几个人拼着切流水线并行是按工序分工每个人只做自己这一段。混合并行则是三种切法叠起来。你怎么组合取决于模型结构、显存预算、集群拓扑和可用卡数没有绝对最优全靠试。3. 通信是分布式训练的中枢神经AllReduce到底在做什么很多做算法的同学第一次接触框架时最迷惑的不是模型怎么改而是为什么代码里会有那么多看似无关的通信操作。你可以在DDP里只写几行代码但底层每迭代一次梯度都要经历一场完整的数据全省大集合。这个集合动作就是AllReduce。3.1 梯度同步的本质AllReduceAllReduce是一个分布式计算通信原语含义是所有节点参与把每个节点的数据做某种归约操作最常见的是求和再把结果广播给所有节点。在数据并行训练里每张卡根据自己的数据子集算出本地梯度这只是一个局部信息。要让所有卡保持模型一致就需要把每个参数位置的梯度跨卡求和再除以卡数取平均梯度平均后的结果发给每一张卡。如果不用AllReduce换一种天真的做法让0号卡把所有人的梯度收上来算完再广播回去。这是AllGather/Reduce的串行版本通信量一样但0号卡会成为瓶颈和单点故障。AllReduce的价值在于它通过巧妙的算法让所有卡都参与数据转发没有单一热点每张卡只负责自己应传的那份数据总通信量可扩展。3.2 Ring AllReduce把数据绕圈传Ring AllReduce是NCCL和Horovod中非常经典的实现。基本思想是把N张卡想象成一个环形每张卡只和相邻的两张卡通信。完整的Ring AllReduce分成两步。第一步叫Reduce-Scatter每张卡把自己本地的梯度数据切成N份在第k轮通信时把第k份发给下一张卡同时从上一张卡接收第k份然后在本地做加法。经过N-1轮后每张卡汇总了某个特定分片的全局和。第二步叫AllGather把已经汇总好的每个分片沿着环广播出去同样N-1轮后每张卡都拥有了完整的全局梯度。Ring AllReduce的好处是通信量不随卡数增加而爆炸理论上扩展性很好坏处是延迟会随卡数的增加线性增长并且环形上的每一条边带宽都必须充足否则拖慢全环。所以它适合卡数适中、数据量大的场景。3.3 树状AllReduce和物理拓扑感知另一种主流实现是树状的NCCL在大规模多节点场景也常会用树结构。树状AllReduce把节点组织成树树叶先向上归约根节点算完再向下广播。好处是延迟对数级增长适合卡数特别多的跨节点场景但对根节点带宽要求高且要求网络拓扑确实存在层级关系。这里就引出一个分布式训练的经典问题通信带宽的物理上限。NCCL默认用的是GPU Direct RDMA多卡之间通过NVLink或InfiniBand高速互联。但如果你买的是云主机每台机器上多张卡共享同一个网卡带宽多节点之间通信就非常容易撞车。我第一次跑跨节点训练时单机内8卡AllReduce只花几十毫秒加了4台机器后一轮迭代的通信时间直接从50ms飙到500ms原因就是节点间走的是千兆以太网共享带宽。后来调整NCCL环境变量设置NCCL_P2P_DISABLE1配合NCCL_SOCKET_IFNAME指定高速网卡才把通信时间降下来。3.4 通信隐藏在反向传播里PyTorch DDP的巧妙之处PyTorch DDP最精妙的设计之一是它把梯度同步嵌进了反向传播的过程。它会在反向传播时注册hook当一个参数的梯度算完之后不等整个模型反向算完就立即启动该参数的AllReduce。这样梯度通信和下一层反向计算可以重叠GPU在等数据的同时也在算通信开销被大量掩盖。很多不熟悉DDP的人会以为它只是简单地在每轮迭代末尾做一次同步其实那是Horovod更早期的做法。DDP的梯度通信粒度是参数级重叠这也是为什么它能做到不错的扩展性。理解这个机制之后你就会明白为什么DDP里有时显存占用比单卡高——它需要额外的通信缓冲区加上每个进程持有完整的模型副本显存开销自然上去了。4. 同步更新与异步更新不是非黑即白的选择题说完通信接下来是训练策略层面的经典争论同步更新还是异步更新。这两者的取舍直接影响收敛效果和训练速度。4.1 同步训练稳定但逃不过木桶效应数据并行最常见的模式是同步训练所有worker用各自的数据分片算完梯度做一次AllReduce求平均然后统一更新模型参数进入下一轮迭代。这种方式的优点是梯度的全局一致性非常好每一步direction都代表所有数据子集的共识优化过程稳定收敛曲线可预测。绝大多数学术基准测试和大模型预训练都在同步模式下完成。代价是同步训练有木桶效应每一轮迭代要等最慢的那张卡完成计算。如果集群里有几张卡因为散热、邻居的虚拟机抢占、硬件老化而变慢整体训练速度就会被拖到和它们一样慢。而且同步频率越高每步都同步通信占比越高。一个很实用的缓解方式是梯度累积gradient accumulation。比如你想用1024的batch size但每张卡只能塞下16条样本那就让每张卡连续算32个micro-batch把梯度累加起来做完32次本地反向后才做一次AllReduce。这样能把通信频率降低32倍同时保持足够大的有效batch size。很多大模型训练实际就是这么干的。但要注意梯度累积后BatchNorm的统计量会受影响如果是CNN类模型需要额外小心。4.2 异步更新快递员各送各的和大家一起核对账本异步训练的口号是让每张卡自己跑自己的。每个worker算完梯度后直接更新全局参数不用等别人。这在参数服务器架构中很常见。优点不言而喻没有等待单卡吞吐量最高一台卡慢了不影响其他人。缺点更明显——梯度stale。某个worker算梯度时读到的模型参数是T时刻的等它算完想更新时全局参数可能已经被其他worker更新到T1000了。用一份过时的梯度去更新最新参数轻则收敛变慢重则Loss震荡、模型发散。所以异步训练在深度学习中远没有在大数据领域那么受欢迎。它只适合模型更新不频繁、容忍噪声的场景。一般工程上的做法是在使用异步时降低学习率、增加梯度审查或用半异步折中一部分worker同步、一部分异步。4.3 通信重叠与梯度压缩不换框架也能压掉通信成本除了选择同步或异步还有两个从工程层面削减通信开销的经典手段。第一个是通信计算重叠。前面提到DDP已经做了参数级重叠你还可以从数据加载、前向计算和跨层传输上继续挖掘重叠空间。最典型的做法是预取下一批数据、提前压缩激活值、反向中先通信再计算。用CUDA Graph或torch.cuda.graphs可以把一堆小操作合并成一个大图减少kernel launch开销。第二个是梯度压缩。通信的数据量如果能缩小AllReduce自然就快。常见方法有梯度量化把FP32压成FP16或INT8传输、梯度稀疏化只传超过阈值的梯度其他本地做动量补偿、低秩分解。这类技术以误差反馈为代价换取带宽在大规模跨广域网训练时特别有用。但要记住压缩比越高优化收敛性质越容易被破坏一定要实验验证。我自己的经验是不要一上来就搞异步、搞压缩。先跑一个同步版本把通信开销用profiler测出来。如果通信占比不到20%那说明你的计算已经很饱和不值得为那点收益引入复杂机制。真实生产中简单可靠的同步数据并行往往已经能解决大部分问题。5. 节点故障与容错设计分布式训练最容易翻车的地方如果你觉得把训练代码跑起来就万事大吉那一定是还没经过大规模训练的毒打。分布式训练面对的是一群随时可能出问题的物理设备和系统进程网卡松了、温度过高、电源波动、邻居家虚拟机跑了个吃满CPU的进程……任何一个硬件故障都可能让整个训练任务中断。而中断一次的代价是你前面几天甚至几周跑出的进度全部归零。5.1 别把故障当异常它是分布式系统的默认状态在单机时代蓝屏死机是偶发事件。在分布式集群里故障是常态。几百块GPU长时间高负载运行每周至少有一次卡要报错或掉线。如果你没有任何保护机制任务会直接崩溃退出。尤其是训练大型模型时一个迭代动辄几十分钟甚至几小时重启一次的代价不是简单恢复而是可能要从最近的checkpoint重新热身后再继续。我们曾经跑一个中规模预训练任务三天内遇到两次NCCL超时导致的任务失败。起初以为是自己代码有bug排查后才知道同一批机器上另一位同事的任务占了带宽把通信挤挂了。从那以后我再也不迷信云上机器可靠这种说法。5.2 Checkpoint分布式训练的生命线最简单的保障就是定期保存checkpoint。但分布式训练里的checkpoint不是把模型权重存个文件而已有四个细节要特别留意。第一保存频率要和训练代价匹配。如果你一个epoch要跑8小时每5分钟存一次是合理的如果一个epoch才20分钟存太频繁反而浪费IO。通常按step间隔×存一次需要的大概时间来核算保证最多损失不超过半小时进度。第二不只是模型权重要存数据加载状态也要存。很多训练过程中断后能接上但数据分布对不上比如shuffle状态没保存导致某些样本被重复训练、另一些样本从没出现过。正确做法是保存sampler或dataloader的迭代位置包括随机种子。第三优化器状态必须一起存。Adam里的动量信息决定了后续更新方向如果不存从权重恢复继续训练等于换了一个优化器初值Loss曲线会跳变。第四分布式场景checkpoint目录要能原子提交。多个进程同时写同一个文件是非常典型的崩溃现场。最好每个rank把自身状态写到独立目录全部写完后再改写一个标记文件。恢复时检查标记文件否则读到写了一半的文件恢复出来就是乱参数。5.3 从断点续跑到弹性训练断点续跑是挂了然后手动重启这已经是事故后的补救。更进一步的做法是弹性训练elastic training让集群在节点增减时自动调整参与训练的worker数量不中断任务。Ray Train、PyTorch 2.0的elastic DDP都支持类似能力。弹性训练的原理并不神秘动态监听节点集合变化发生变化时触发一次全局barrier把变化后的workers重新组织从最近一部checkpoint恢复。难点在于怎么让正在执行的梯度计算安全地停止并重新分配这对通信组、数据切分、学习率调度都有连锁影响。如果你还在成长阶段我建议先不做弹性把checkpoint做好才是性价比最高的方案。等任务规模大到人为重启都会耽误大量人力时再上弹性也不迟。6. 一次多机多卡训练实录环境准备、关键参数与踩坑清单纸上谈兵结束分享一次我实际跑多机多卡训练的过程。场景是8卡单机训练调通后扩展到4台机器共32卡训练一个大Transformer模型。这里不说具体模型把通用经验和坑位讲明白。6.1 准备阶段最容易翻车的三件事主机名、密钥、初始同步多机训练和单机最大区别是环境一致性。每台机器得能通过SSH免密互相连接PyTorch的init_process_group需要知道所有rank的地址和端口。我建议用共享文件系统如NFS或云盘作为初始化后端把rank0的地址写在一个共享文件里其他rank去读省去手动传入MASTER_ADDR的麻烦。真正耗时的坑往往在环境依赖上。比如某台机器上的CUDA驱动和另一台不一致或者显卡驱动版本和PyTorch编译版本不匹配跑起来时静默crash。我的习惯是先写一个健康检查脚本在所有机器上统一检查GPU型号、驱动版本、NCCL版本、PyTorch版本、Python包版本逐项对比。这一步看起来琐碎但能避免你在一堆报错信息里找共同点浪费一下午。6.2 PYTHON脚本和DDP初始化代码级别的关键点主流程参考PyTorch DDP有几个细节import torch import torch.distributed as dist from torch.nn.parallel import DistributedDataParallel as DDP def setup(rank, world_size): dist.init_process_group( backendnccl, init_methodfile:///shared/train_store, rankrank, world_sizeworld_size, ) torch.cuda.set_device(rank) # 每个进程指定的rank必须和物理GPU对应否则会发生隐形的算力错位 model DDP(model.to(rank), device_ids[rank])要注意world_size是总进程数不是单机卡数。如果你的每台机器8张卡world_size就是32。进程数和GPU必须一一对应不能出现一个进程管多卡这种自杀式行为除非你想做单节点内的多进程控制。数据加载也要特别调整。DistributedSampler会自动按rank切分数据但每轮epoch开始时要调用sampler.set_epoch(epoch)否则每个epoch都是相同的数据shuffle结果模型会见过不公平的训练分布。这是很多人忽视的小坑。6.3 通信超时、OOM和Load Imbalance的排障思路第一次跨节点跑的时候撞上了典型的NCCL超时报错信息类似NCCL error: timeout. 一开始以为是网卡问题后来用nccl-tests做一次allreduce基准测试发现单机内8卡快跨4机后延迟翻了好几倍。顺着网络诊断才发现四台机器里有两台节点间走的是慢速网络另一对走的是高速网络节点间带宽不对等导致整体排队。另一次OOM是载入阶段显存分配问题。虽然模型理论上能装下但PyTorch默认的显存缓存策略会在一次大张量申请失败时报错而不是寻找可回收的碎片。我通常设置PYTORCH_CUDA_ALLOC_CONFmax_split_size_mb:128并在代码里尽量把初始化搬到大显存操作前。切忌在加载大权重的同一瞬间并行allocate其他大tensor。还有负载不均四台机器算力相同但一台机器总是比其他三台慢10%。一开始以为是网络带宽后来发现是那台机器上还有另一个后台任务在跑CPU密集操作导致数据预处理跟不上。把数据管线改成独立进程并关闭线程竞争后速度就齐了。6.4 分布式训练的Checklist以下这张清单是我现在每次跑分布式训练前都会过一遍的送给需要的朋友类别检查项说明环境所有机器GPU型号一致混用不同代际卡可能造成AllReduce卡在慢卡上环境NCCL、CUDA、PyTorch版本一致不一致会出现莫名的初始化失败网络节点间使用高带宽内网千兆以太网跑大模型训练是自杀启动MASTER_ADDR/rank/world_size正确用共享文件init_method更省心数据DistributedSampler每epoch重新set_epoch保证shuffle有效且均衡模型BatchNorm换成SyncBN或谨慎使用单卡统计量在多卡下不再准确存储checkpoint保存到共享存储所有rank可见才能恢复验证前几步打印loss和模型参数hash确认多卡同步后初始状态一致写在最后的一点个人体会分布式训练在我眼里本质上是一个用通信换计算、用冗余换稳定的系统工程。刚开始接触时你可能会被各种并行模式、通信原语、容错方案劝退但只要亲手把一个模型从单卡推到多机多卡看着训练吞吐量按预期上涨、Loss稳定下降那种成就感是很实在的。这些年我最大的感受是别追求最炫的方案先追求最稳的方案。数据并行能解决80%的需求模型并行用于突破单卡显存上限通信优化用来填最后那部分效率缺口。每一层都建立在前一层正确的基础之上。至于怎么判断哪一层该做到多深只能依靠一次次实测、profiling和复盘。希望这篇内容能帮你少走一些我走过的弯路。如果你也在跑分布式训练或者正准备上多卡欢迎在评论区交流你遇到的报错或者心得很多坑计算机书上是不会写的但现实中它就在那儿等着你。
返回列表