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

资讯详情

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

PyTorch分布式训练:DP、DDP、ZeRO、FSDP原理对比与选型

PyTorch分布式训练:DP、DDP、ZeRO、FSDP原理对比与选型 DP、DDP、ZeRO、FSDP这四个词应该是PyTorch分布式训练里最容易被搞混的一组缩写。我在各种技术交流群里见过太多次类似提问DDP和DP到底差在哪FSDP是不是就是DDP升级版ZeRO为什么要多一个阶段这些问题背后其实是一条很清晰的演进线——从参数服务器到Ring-AllReduce从显存冗余到分片存储。这篇文章就把整条线捋一遍讲清楚每个方案的原理、定位和适用场景也顺手把实际工程里最容易踩的坑指出来。不管你是刚入门分布式训练的新手还是已经用DDP跑过几个模型、但一直没搞懂FSDP到底做了什么这篇都值得看完。看完之后你应该能回答这些问题我的场景该用哪个为什么DDP把DP碾压了FSDP是怎么做到“既能用DDP的代码又能跑更大模型”的1. 为什么“数据并行”是所有并行策略里最该先掌握的1.1 从单卡到多卡我们到底在解决什么问题训练一个模型核心动作是不断做前向、反向、更新参数。当模型在一张GPU上能够完整放下时你想提高训练速度最直接的办法就是同时塞进去更多数据、用多张卡一起算。这就是数据并行的直觉来源每个设备持有完整模型副本但处理不同的batch数据最后把梯度同步一下再各自更新参数。这个思路之所以“最该先掌握”是因为它解决的问题最普遍。绝大多数实际项目单卡放得下模型只是显存不够用或训练太慢。数据并行能在不改变模型结构、不修改训练逻辑的前提下把吞吐量提上去。相比之下模型并行和流水线并行是为了解决“单卡放不下模型”而设计的它们在工程上的复杂度比数据并行高一个量级。先从数据并行入手能让你把分布式训练里最核心的概念——通信、同步、梯度聚合——完整过一遍。1.2 数据并行与模型并行、流水线并行的分界线我把三者的区别打个比方。数据并行像开多个班同时讲同一门课每个班老师都一样、教材都一样只是学生不同最后把各班的考试成绩汇总。模型并行则像把一门课拆成多个章节不同老师各讲一段学生听完A老师再去听B老师。流水线并行更像工厂流水线每个工位负责一道工序前一个工位做完传给下一个工位。数据并行有一个前提条件模型必须完整存在于每张卡上。如果模型太大单卡显存放不下那就得先考虑模型并行、流水线并行或后面的FSDP这类分片方案。但即便如此绝大多数大模型训练框架仍然会先引入数据并行再在数据并行之上叠加模型/流水线并行。原因很简单数据并行通信模式规整、负载均衡、工程实现最成熟。1.3 任何数据并行方案都绕不开的三件事情不管DP、DDP、ZeRO还是FSDP它们共同面对三件事模型参数广播/复制训练开始每张卡都要有一份一致的模型参数。梯度同步每张卡用各自的batch算完梯度后需要把梯度汇总成全局梯度。参数更新一致性所有卡用同一份全局梯度更新参数保证模型不会发散。不同的方案主要区别在于“梯度怎么同步”“参数怎么存储”“显存和通信怎么权衡”。DP选择的是“单机多卡、主卡聚合、主卡更新后广播”DDP选择的是“每张卡All-Reduce后各自更新”ZeRO和FSDP则更进一步把优化器状态、梯度甚至参数本身都分区存放。理解了这三个共性后面看每个方案就觉得不乱了。2. DP一个“经典但有点过时”的起点2.1 DP是怎样工作的PyTorch的torch.nn.DataParallelDP是很多人接触到的第一个多卡训练接口。它做的事情可以概括为把输入batch在batch维度切分成多份分发给多张GPU每张GPU独立做前向和反向得到梯度所有梯度汇总到主卡通常是GPU 0上主卡做参数更新新参数再广播回所有卡。这个过程在代码里很简单只要用model nn.DataParallel(model)包一层甚至不用改训练循环。但它的架构决定了上限很低所有梯度都要经过主卡主卡既是求和节点又是更新节点也是广播源。多卡通信全部是“星型拓扑”GPU 0要同时跟所有其他卡通信。如果你看过训练时的nvidia-smi大概率会发现DP模式下GPU 0的利用率或显存占用比其它卡高一些严重时甚至直接爆显存。2.2 DP的性能短板主卡负载、通信瓶颈、BN统计DP的三个明显问题但凡在多机或大规模卡数场景下用过的人都会有体感主卡通信和计算双重瓶颈。每轮迭代其它卡都要把梯度发到GPU 0GPU 0求和后再把新参数广播回去。通信数据量与模型参数量成正比模型一旦上亿参数每次迭代的通信开销非常大GPU 0还要承担额外的加减和广播操作。这个瓶颈让DP几乎无法扩展超过4卡。不支持多机。DP是单机多卡方案跨机器时它根本没有对应的进程通信设计。很多人以为DP能自动多机但官方文档明确说明DataParallel只支持单机且不推荐在多卡训练中继续使用。BatchNorm统计量不一致。DP默认会在每个设备上独立计算BN的均值和方差再汇总到主卡。这样的统计结果和真正“全局BN”有偏差在batch较小、卡数较多时BN的移动均值会变得不稳定影响收敛。2.3 什么时候用DP并不算坏选择虽然DP被诟病但也不是一无是处。我的经验是单机2卡、模型很小、只想快速验证一个想法时DP仍然能让你一行代码跑多卡。比如写个简单的分类器、跑个小规模消融实验DP的通信开销完全不成问题。如果你追求调试效率和代码简洁DP可以当做一个过渡方案。但只要你打算认真跑实验、卡数超过2张或者模型超过千万参数我都建议直接换成DDP别在DP上浪费时间。3. DDP当前多卡训练的默认选项3.1 DDP与DP的本质差异All-Reduce替代参数服务器DDPtorch.nn.parallel.DistributedDataParallel为什么能成为当前事实标准核心原因是它把DP的“主从参数服务器”模式换成了“All-Reduce集体通信”模式。All-Reduce是一类集体通信操作它让所有参与进程共同执行一个操作最后每个进程都拿到聚合后的结果。最常见的是Ring-AllReduce所有GPU按顺序排成一个逻辑环每个GPU只与相邻GPU通信把本地梯度分成块像流水线一样在环上转一圈最终每个GPU都拥有所有梯度的和。这样就不存在一个中心节点通信负载被均匀摊到所有卡上扩展性远好于星型拓扑。DDP的具体流程是每张卡用自己的数据算出完整的梯度然后所有进程一起执行一次梯度All-Reduce每个进程拿到全局平均梯度最后每个进程各自调用优化器更新本地模型。因为模型初始参数一样、每次更新的梯度也一样所以所有进程里的模型始终保持一致。不需要每轮广播权重通信量就是一次全梯度同步。3.2 为什么DDP比DP快通信方式、冗余计算、同步机制从通信量上看DP每轮需要“梯度上传参数广播”两次通信DDP每轮只需要一次梯度All-Reduce。Ring-AllReduce优化得很好通信量与总数据量成正比且每个GPU只传输自己负责的部分带宽利用率很高。再加上DDP在每张卡上独立更新参数省掉了主卡的计算步骤主卡不再是瓶颈。还有一点容易被忽略DDP默认在反向传播过程中就开始梯度通信而不是等所有梯度都算完再通信。它会为每个参数注册一个梯度hook当某个参数的梯度计算完成后立即启动异步All-Reduce这样通信和计算可以部分重叠进一步缩短每轮迭代时间。3.3 怎么写一个标准的DDP训练脚本一个标准的DDP训练脚本关键步骤并不复杂但顺序不能乱import torch import torch.distributed as dist from torch.nn.parallel import DistributedDataParallel as DDP from torch.utils.data.distributed import DistributedSampler def main(): dist.init_process_group(backendnccl) rank dist.get_rank() local_rank rank % torch.cuda.device_count() torch.cuda.set_device(local_rank) model torch.nn.Linear(128, 10).cuda(local_rank) model DDP(model, device_ids[local_rank]) dataset torch.randn(1000, 128) sampler DistributedSampler(dataset, num_replicasdist.get_world_size(), rankrank, shuffleTrue) loader torch.utils.data.DataLoader(dataset, batch_size64, samplersampler) optimizer torch.optim.SGD(model.parameters(), lr0.1) loss_fn torch.nn.MSELoss() for epoch in range(10): sampler.set_epoch(epoch) for x, y in loader: x, y x.cuda(local_rank), y.cuda(local_rank) optimizer.zero_grad() out model(x) loss loss_fn(out, y) loss.backward() optimizer.step() dist.destroy_process_group() if __name__ __main__: main()启动方式通常是torchrun --nproc_per_node4 train.py它会帮每个进程配置好MASTER_ADDR、MASTER_PORT、RANK和LOCAL_RANK环境变量。注意这里有个关键点DistributedSampler和set_epoch是配套的如果不设置epoch每个epoch的数据shuffle结果会完全一样训练会退化。3.4 DDP实际使用中的几个坑模型保存/加载。保存DDP模型时不能只保存model.state_dict()因为state_dict的key里会带module.前缀。最常见做法是model.module.state_dict()加载时直接load_state_dict即可。如果是checkpoint同时包含优化器和调度器建议把rank信息也存进去方便恢复训练时对齐随机种子。BN同步。DDP本身不会改变BN的统计方式。如果你的模型里有大量BN层、且每张卡的batch很小我强烈建议用SyncBatchNorm包一层model torch.nn.SyncBatchNorm.convert_sync_batchnorm(model)然后再包DDP。这会增加一次通信但能显著提升小batch下的精度稳定性。后端选择。GPU上优先用NCCL。gloo后端虽然能在CPU和GPU之间通用但性能差很多。跨机器时还要注意环境变量NCCL_SOCKET_IFNAME和NCCL_IB_IFNAME很多多机训练卡住都是因为网络设备没选对。学习率与warmup。DDP并不会帮你调整学习率。当全局batch变大后学习率需要相应调整通常可以按线性缩放规则但更稳妥的做法是先做一次小规模lr range test。不要指望DDP能自动解决这个问题。4. ZeRO把优化器状态、梯度、参数全部分区4.1 训练一块大模型显存到底花在哪里很多人以为模型显存占用主要来自参数这是误解。拿一个1B参数、Adam优化器、FP16混合精度训练的模型举例模型参数本身只有2GBFP16但在混合精度下Adam优化器需要保存两份FP32副本和两份动量状态这四份都要FP32总共16GB。再加上反向传播产生的梯度以及激活值显存大头其实是优化器状态和激活值。换句话说一个1B参数的模型用混合精度在单卡上训练光优化器状态就可能吃掉16GB以上显存。这就是为什么模型参数明明只有几GB却很难塞进一张卡。ZeRO的核心洞察是这些显存消耗存在大量冗余尤其是优化器状态人人都存一份没必要。4.2 ZeRO-DP的三个阶段怎么省显存ZeRO-DP是DeepSpeed提出的方案核心是把数据并行中重复存储的状态分片到各个进程上。它分三个递进阶段Stage 1优化器状态分片。每个进程只保存1/N优化器状态负责更新自己那部分参数。更新前需要全量参数吗不需要因为每个进程只更新自己拥有的那部分参数。但反向传播时每个进程需要全量梯度来计算各自的更新。这个阶段就把优化器内存从原来的N份降到1份。Stage 2梯度分片。除了优化器状态梯度也按参数分片存储每个进程只保留自己负责更新那部分参数的梯度。反向传播结束后通过一次通信把整份梯度聚合每个进程只保留自己分到的片段。Stage 3参数分片。这是最激进的阶段连模型参数本身也按层/参数分片存储。前向或反向计算到某个参数时需要通过通信临时收集需要的参数片段计算完再释放。这个过程叫“参数收集”没有它任何层都无法计算。到了Stage 3模型才能真正突破单卡显存限制。4.3 ZeRO不是免通信费的午餐ZeRO经常被误解成“又省显存又通信少”这是错的。以Stage 2为例它和DDP的通信量其实差不多都是同步一次梯度。但Stage 3在每次前向计算每个分区参数时都要做一次全收集通信次数显著增加甚至可能比DDP多出数倍通信量。DeepSpeed官方资料里也明确说过ZeRO用通信换显存是“显存受限时才会用”的方案。所以ZeRO的适用场景很明确模型接近或超过单卡显存上限但又不想写复杂的模型并行代码。如果你单卡显存还很宽裕强行上ZeRO反而可能速度更慢。4.4 ZeRO与混合精度、激活重计算的配合实际使用ZeRO时混合精度基本是标配。混精训练下模型参数和梯度用FP16存储但优化器状态必须用FP32否则Adam更新会不稳定。ZeRO Stage 1混合精度能省掉大量FP32优化器状态显存收益非常明显。激活重计算activation checkpointing和ZeRO是互补关系。ZeRO负责切分参数、梯度、优化器状态激活重计算负责切分前向保存的中间激活值两者加在一起才能让超大模型真正跑起来。但激活重计算会增加约30-40%的计算开销需要根据训练时长和显存压力做取舍。我的经验是显存不够时先重计算再考虑ZeRO Stage升级因为重计算的额外代价是可控的而Stage 3的通信开销有时比计算代价更难看。5. FSDPPyTorch官方给出的ZeRO实现5.1 FSDP说到底做了什么事情FSDPtorch.distributed.fsdp.FullyShardedDataParallel是PyTorch官方实现的“ZeRO风格”分片并行方案。它在概念上对应ZeRO Stage 3参数、梯度、优化器状态都分片存储每个进程只持有自己那部分模型副本计算前通过All-Gather临时收集需要的参数计算后立即释放。和ZeRO一样FSDP也会做统一的梯度同步。反向传播过程中收集梯度做完梯度All-Reduce后丢弃非本进程负责的梯度片段优化器更新只更新本地持有的参数片段。这套设计把“分片”的复杂度隐藏在包装器内部用户写的模型前向代码几乎不用改这是它比DeepSpeed ZeRO更易用的关键原因。5.2 FSDP的分片策略与sharding_strategyFSDP提供三种分片策略对应不同需求FULL_SHARD参数、梯度、优化器状态全部分片对应ZeRO Stage 3省显存最彻底但通信最多。SHARD_GRAD_OP只分片梯度和优化器状态参数不全分片对应ZeRO Stage 2通信更少适合单卡显存还够但想跑更快时使用。NO_SHARD不分片相当于用FSDP的包装器跑DDP逻辑主要是为了平滑切换代码。我个人的建议是刚开始尝试FSDP时用默认的FULL_SHARD因为它是FSDP最核心的价值。如果模型不是特别大或者你感觉通信开销拖慢了训练再换成SHARD_GRAD_OP试试。5.3 FSDP的配置和踩坑FSDP不像DDP那样一行代码搞定它要求你显式地配置哪些层要分片。常见的写法是用auto_wrap_policy指定包装策略from torch.distributed.fsdp import FullyShardedDataParallel as FSDP from torch.distributed.fsdp.wrap import transformer_auto_wrap_policy policy transformer_auto_wrap_policy( transformer_layer_cls{YourTransformerBlock} ) model FSDP(model, auto_wrap_policypolicy, sharding_strategyFSDP.ShardingStrategy.FULL_SHARD)transformer_auto_wrap_policy是给Transformer模型用的它会自动把每个TransformerBlock包成一个FSDP通信单元。如果你用的是CNN或MLP需要自己写一个lambda策略或者更简单一点直接对整个模型做FSDP但那样粒度太粗分片效果差。包装粒度很关键包装太细通信次数爆炸包装太粗显存省不下来。还有一个很容易踩的坑是前向和反向期间不能随意访问模型参数。因为参数只有在计算时才被临时收集如果你在训练循环里直接打印model.parameters()或对参数做某些自定义统计会触发隐式全收集拖慢速度甚至报错。类似地FSDP对torch.no_grad也有影响推理时如果不在FSDP上下文中管理好参数可能不会re-shard长期持有全量参数反而增加显存。5.4 FSDP与DDP如何平滑迁移FSDP的接口设计让迁移非常平滑。最常见路径是先写好一个DDP训练脚本然后把DDP(model, device_ids[local_rank])换成一个可配置的包装函数传入参数决定用FSDP的哪种策略。训练循环、优化器、数据加载器基本不用动。唯一需要特别注意的地方是optimizer.step()之前FSDP要求所有参数已经收集完整。内部机制已经处理了大部分情况但如果你自定义了梯度裁剪推荐用FSDP提供的clip_grad_norm_它会先收集完整梯度再做裁剪。如果你用PyTorch新版2.0还能用torch.compile搭配FSDP性能和显存都有进一步优化但也要注意编译本身耗时适合长训练任务。6. 一张表看清DP/DDP/ZeRO/FSDP以及选型建议6.1 横向对比表方案通信模式模型参数存储梯度同步显存占用工程复杂度DP星型主卡聚合每卡全量副本主卡收集后广播较高主卡更突出极低一行代码DDPRing-AllReduce每卡全量副本每卡All-Reduce中主要冗余在参数/梯度/优化器状态低需要进程组初始化ZeRO Stage 1/2梯度All-Reduce全量参数梯度分片分片梯度同步比DDP低取决于阶段中需要DeepSpeed接入ZeRO Stage 3All-Gather 梯度同步参数分片分片梯度同步很低可训练单卡放不下的模型较高需配置FSDPAll-Gather 梯度同步参数分片按策略分片梯度同步很低与ZeRO Stage 3相当中PyTorch原生接口相对友好从这张表能看出来DP和DDP属于传统数据并行模型每卡都有一份完整副本ZeRO和FSDP属于分片存储型数据并行把重复状态拆到多卡上从而突破显存上限。6.2 选型建议按规模与场景我一般按照下面几条经验来选1卡训练或2卡快速调试直接用普通单卡代码或DP。DP虽然性能不如DDP但代码量最少调试方便。单机4-16卡模型单卡能放下无脑选DDP。DDP性能好、生态成熟、调试资料多是当前性价比最高的方案。模型单卡放不下但不想改模型并行代码优先FSDP。先用FULL_SHARD跑通再用SHARD_GRAD_OP调速度和显存。DeepSpeed ZeRO Stage 3也能做但需要额外引入DeepSpeed的配置和启动方式除非你要用DeepSpeed的其它能力比如Offload、压缩否则我更推荐FSDP。大规模多机超大模型训练通常会把DDP/FSDP与模型并行、流水线并行混合起来用。比如用FSDP做数据并行维度再用Megatron做张量并行。到这一步工程复杂度很高不是单个框架能解决的。6.3 我在实战中的体会最后分享一点个人经验。我见过不少人一上来就听说FSDP省显存不管什么项目都套FSDP结果发现训练速度比DDP慢了不少。FSDP的通信开销是真实存在的尤其是模型层数深、包装粒度细的时候每层都要All-Gather。如果你的模型单卡能放下DDP通常更快如果你的模型放不下但激活值占用了大头先试激活重计算再考虑FSDP。DP这个老方案我是真不推荐在正式实验里用了。它最大的问题是主卡瓶颈和无法多机扩展性太差而且现在DDP的代码成本已经很低没有理由为了省两行代码牺牲整体性能。每次从DP迁移到DDP我都建议顺手把模型保存逻辑也规范成model.module.state_dict()否则后面对接评估或部署时会多踩几个坑。FSDP还有一个容易被忽略的优点是它和DDP可以共存。我在一个多卡任务里曾经把一部分模型用DDP包装另一部分用FSDP包装用来精细控制哪部分参数需要分片、哪部分需要全量同步。这种灵活性是DeepSpeed ZeRO不容易做到的也是PyTorch原生方案比较难得的优势。如果要从这套方案里总结一条最实用的建议那就是先清楚地知道你的显存瓶颈到底在哪里——是模型参数、优化器状态、梯度还是激活值。然后再去选对应方案。很多人把DDP和FSDP的差异记在脑子里却不知道自己的显存到底被谁吃掉了结果换了方案也没解决问题。搞清楚这一步选型基本不会错。
返回列表