
1. 从DataParallel到FSDP分布式训练这十年到底在解决什么问题十年前我刚开始碰分布式训练的时候手里最趁手的工具就是把一个batch塞进DataParallel看着四块显卡轮流干活。那时候没人觉得显存会成问题因为模型本身也就一两亿参数一块V100还能塞下。真正让人开始睡不着觉的是2019年GPT-2把模型推到1.5B之后大家突然发现单卡装不下了而多卡又不知道该怎么把一份模型拆开。这十年里从DataParallel到DDP再到模型并行、流水并行最后是ZeRO思想落进PyTorch变成FSDP本质只回答了一件事一份模型只能有一份完整副本还是可以被拆成很多份、各自为战这篇文章我会沿着这条线把FSDP的前因后果、底层原理、实操配置和踩坑记录都过一遍。适合正在做大模型训练、想把手里的单卡代码升级成多卡并行或者单纯想搞明白为什么FSDP能成为PyTorch标配的读者。我尽量不端着该算的账算给你看该踩的坑也不藏着。1.1 DataParallel时代能用但不经打的早期方案DataParallelDP的思路很直观把模型复制到每张卡上每次前向都由主进程把所有卡上的输出拼起来算loss反向之后由主进程统一reduce梯度再广播回各个卡。听起来没有问题实际上问题一堆。首先是通信瓶颈。DP的所有梯度汇总、参数广播都经过主进程GPU一多主进程就成了排队口。我当年在4卡上跑一个稍大的CNNbatch稍微大一点就能看到GPU利用率忽高忽低很多时间都浪费在等主进程分发上。其次是GIL限制PyTorch的DP用多线程实现Python的全局解释器锁让多线程很难真正并行。更关键的是DP对模型本身的显存占用毫无优化每张卡都持有完整的模型参数、梯度和优化器状态卡上的显存只取决于模型大小跟卡的数量没关系。所以DP只适合小规模试验。多卡训练真正成为工程标配是DDP出现之后的事。1.2 DDP把通信效率拉满却绕不开显存重复DistributedDataParallelDDP在2017到2018年间逐渐成为主流。它的做法是每个进程独享一张卡、持有一份完整的模型副本前向时各算各的batch反向算完梯度之后用Ring All-Reduce在进程之间对梯度做一次全局归并然后每个进程用归并后的完整梯度各自更新参数。相比DPDDP有两个决定性优势。第一通信量从每一层都同步一次压缩成一个step只做一次全量梯度all-reduce而且Ring All-Reduce算法让每个进程只和邻居通信带宽随卡数扩展。第二每个进程是独立的Python进程彻底绕开了GIL。但DDP有个绕不开的天花板每个GPU上仍然要放一份完整的参数、一份完整梯度以及一份完整的优化器状态。模型小的时候这无所谓模型大到一定程度就直接卡死。我印象很深的是第一次尝试把1.5B的GPT-2切进8卡A100时光参数、梯度和Adam状态就要上百GB单卡40GB根本无从谈起。DDP是完整副本路线的极限到了这堵墙面前必须换思路。1.3 显存墙出现后模型并行与ZeRO给了两条路2019年之后社区给了两条不同的解题路线。第一条是算子级的模型并行代表是Megatron-LM的张量并行。它把权重矩阵按行或按列切到多张卡上每张卡只算自己那一份前向和后向都需要在切分维度上做all-reduce合并结果。好处是单卡显存需求骤降坏处是通信极频繁基本上每一层矩阵乘法都要同步一次对节点内带宽要求很高跨节点做张量并行几乎是灾难。第二条是流水并行按层把模型切成几段每张卡只负责一段。由于数据要逐段流过天然存在空闲气泡而且各段的显存占用不均衡调起来很麻烦。就在大家觉得模型并行已够复杂的时候DeepSpeed的ZeRO论文给了第三个思路——冗余才是问题本身。DDP下N张卡就存了N份参数、N份梯度、N份优化器状态这些全是重复的。ZeRO主张不复制改用分片每一份状态只存一次分散到不同卡上用的时候再临时拼起来。这个思想后来就是FSDP的直接理论根基。1.4 FSDP正式登场以及它后来的演进节点FSDPFully Sharded Data Parallel在PyTorch里经历了从试验到标配的过程。我把关键节点整理了一下时间事件意义2019GPT-2 1.5B出现显存墙问题暴露推动模型并行与ZeRO研究2020ZeRO论文与DeepSpeed开源确立分片消除冗余的路线2022年3月PyTorch 1.11引入FSDP beta版原生支持分片训练降低了上手门槛2023年3月PyTorch 2.0将FSDP转为稳定特性与torch.compile等新能力打通生态成型2024年FSDP2、DTensor与fully_shard接口出现分片逻辑与张量并行统一组合能力更强FSDP在思想上等价于DeepSpeed的ZeRO-3但它作为PyTorch原生组件和torch.compile、torch.distributed.checkpoint、DTensor这些基础设施配合得更好不用额外装第三方库。这也是为什么现在越来越多训练框架默认用FSDP而不是DeepSpeed。2. FSDP的核心变革参数分片背后的显存账本要理解FSDP必须先理解它解决的是什么问题。很多人一上来就盯着All-Gather通信看其实FSDP最大的贡献是重构了显存账本。2.1 先算一笔显存账为什么DDP到7B就卡死了以训练一个7B模型为例采用常见的BF16混合精度加Adam优化器。每个参数在训练过程中要占用以下几份存储状态每个参数的字节数7B模型总量BF16参数2字节14GBFP32主权重4字节28GB梯度4字节28GBAdam动量m4字节28GBAdam方差v4字节28GB合计18字节126GB先解释一下为什么会有FP32主权重。BF16训练时虽然前向和反向都用BF16跑但更新参数时直接拿BF16累积误差会越来越大所以PyTorch和主流框架都会保留一份FP32主权重做精确更新BF16参数只是它的投影。这个细节在DDP下无所谓因为反正每张卡都存得起在FSDP下就很重要因为分片要把这些状态全部算进预算里。结论很残酷DDP下每张卡要准备126GB显存一张80GB的A100根本塞不下7B模型更别说激活值了。我见过不少团队在DDP上试训7B失败第一反应是换更大的卡而不是换并行策略这就是对显存账本没有概念。2.2 ZeRO的三级火箭从只切优化器到全量分片ZeRO把分片过程分成三个阶段每个阶段消除一部分冗余方案每卡保留的内容7B模型每卡显存约DDP全量参数14GB 全量梯度28GB 全量优化器状态84GB126GBZeRO-1全量参数14GB 全量梯度28GB 1/8优化器状态10.5GB52.5GBZeRO-2全量参数14GB 1/8梯度3.5GB 1/8优化器状态10.5GB28GBZeRO-3 / FSDP1/8参数1.75GB 1/8梯度3.5GB 1/8优化器状态10.5GB约16GB这张表里我把FP32主权重和Adam状态合并成优化器状态算了数字只做数量级示意但结论很清晰光切优化器状态就到52GB还是卡着80GB的边全量分片之后每卡只需16GB加上激活值和临时缓冲40GB卡可以轻松跑7B。FSDP等价于ZeRO-3也就是三个状态全分片。分片数等于GPU总数所以卡越多每卡省得越多。这也是FSDP能横向扩展的底层原因。2.3 FSDP一整个step的运转节奏All-Gather与Reduce-Scatter分片容易难的是分片之后怎么计算。FSDP的做法拆开来是这样的初始化阶段把所有参数扁平化成一个一维buffer按GPU数量切成N份每个rank只持有自己那1/N。前向阶段先用All-Gather把当前计算需要的完整参数从各rank收集回来然后做本地前向。算完立刻释放掉不属于自己分片的参数。反向阶段梯度的计算同样需要完整参数所以再次All-Gather每个rank算出的梯度并不是全量而是本地计算涉及的那部分。接下来做Reduce-Scatter让每个rank只累积得到自己负责分片的梯度片段。参数更新每个rank只更新自己分片对应的参数和优化器状态不需要和其他rank通信。这个过程可以类比成一本很厚的书不同人手里各拿几章有人要读某一章前先找持有者复印一份读完就丢。空间省下来了但复印本身有时间成本。这里有个反直觉的点FSDP的通信量并不比DDP小甚至稍微更高。DDP一个step只做一次全量梯度all-reduceFSDP前向要all-gather一次、反向又要all-gather一次、还要reduce-scatter一次。所以FSDP的胜利不在通信效率而在显存。明白这一点很重要——如果你的模型每张卡明明放得下FSDP不会给你带来加速反而可能变慢。这就是为什么FSDP后来引入了forward_prefetch和backward_prefetch目的就是在算当前层的时候提前把下一层要用的参数all-gather过来把通信藏进计算里。配合得好通信开销大部分能被掩盖掉。2.4 为什么包装粒度是FSDP的命门FSDP有一个设计必须理解分片不是对整模型一刀切而是按模块粒度一层层包起来。如果不做任何包装整个模型被当成一个整体那么前向一开始就要All-Gather全模型参数显存峰值直接退回到DDP水平分片就白做了。正确的做法是按Transformer Block粒度包装。每层Blockattention MLP单独成为一个FSDP单元前向走到这层时才all-gather这一层的参数算完立刻释放。显存峰值大概等于最大一个Block的参数 当前层的临时buffer而不是全模型。这就是auto_wrap_policy存在的意义。所以在PyTorch里跑FSDP第一件事不是调sharding_strategy而是想清楚按什么粒度wrap。包大了显存压不下来包小了通信次数剧增这是后面实操部分重点讲的内容。3. 接入FSDP的实操记录配置、代码与实测原理说完了来看实际怎么用。我用一个GPT-2级别的模型加4张A100做了完整的接入实验把能讲清楚的配置和坑都记录在这里。3.1 最小改动跑起来的两种方式经典FSDP的接入代码大概是这样import torch import torch.distributed as dist from torch.distributed.fsdp import ( FSDP, ShardingStrategy, MixedPrecision, BackwardPrefetch, ) from torch.distributed.fsdp.wrap import transformer_auto_wrap_policy from transformers import AutoModelForCausalLM from transformers.models.gpt2.modeling_gpt2 import GPT2Block def build_model(): return AutoModelForCausalLM.from_pretrained(gpt2-xl) def main(): dist.init_process_group(backendnccl) model build_model() wrap_policy transformer_auto_wrap_policy( transformer_layer_cls{GPT2Block} ) model FSDP( model, auto_wrap_policywrap_policy, sharding_strategyShardingStrategy.FULL_SHARD, mixed_precisionMixedPrecision( param_dtypetorch.bfloat16, reduce_dtypetorch.bfloat16, buffer_dtypetorch.bfloat16, ), backward_prefetchBackwardPrefetch.BACKWARD_PRE, device_idtorch.cuda.current_device(), ) opt torch.optim.AdamW(model.parameters(), lr3e-4) # 训练循环省略 # 注意梯度裁剪请用 model.clip_grad_norm_(max_norm)启动命令用torchruntorchrun --nproc_per_node4 train_fsdp.py如果你是PyTorch 2.4之后的版本还可以试试FSDP2的fully_shard它更简洁也更好和张量并行组合from torch.distributed.fsdp import fully_shard for m in model.modules(): if isinstance(m, GPT2Block): fully_shard(m, reshard_after_forwardTrue) fully_shard(model)两种方式本质一样选择看你的PyTorch版本和是否需要和TP组合。经典FSDP资料多FSDP2代码清爽但迭代快生产环境建议先跑通经典版本再迁移。3.2 把关键配置项逐个说清楚FSDP的配置项很多但最常用的就几个逐个过一遍配置项可选值作用与建议sharding_strategyFULL_SHARD / SHARD_GRAD_OP / NO_SHARDFULL_SHARD全分片最省显存SHARD_GRAD_OP只分梯度与优化器状态参数每卡保留副本适合显存只差一点的情况NO_SHARD等价于只切优化器状态ZeRO-1mixed_precisionparam_dtype / reduce_dtype / buffer_dtype训练精度控制。用BF16时参数前向后向用BF16但主权重保持FP32cpu_offloadTrue / False参数可以进一步offload到CPU显存压到最低但吞吐明显下降backward_prefetchBACKWARD_PRE / BACKWARD_POST反向时提前all-gather后续层的参数推荐BACKWARD_PREforward_prefetchTrue / False前向时提前取下一层参数与backward prefetch叠加使用limit_all_gathersTrue / False限制同时all-gather的参数数量控制显存峰值显存紧张开Trueuse_orig_paramsTrue / False保留原始参数接口不用扁平参数对参数名、参数分组更友好推荐Truesync_module_statesTrue / False从rank0广播初始权重。如果你的模型每个rank初始化不同必须开一个常见的调优流程是这样的先FULL_SHARD offload 小batch把model跑起来确认功能没问题。然后逐步关掉offload、加大batch、开prefetch最后微调limit_all_gathers。目标不是把显存压到最低而是在显存刚好放得下的前提下把吞吐拉满。3.3 模型保存与加载最容易做错的环节FSDP下第一版代码我直接写了torch.save(model.state_dict(), model.pt)加载的时候才发现完全没法用。原因是FSDP的state_dict默认只包含当前rank的那一份分片4卡各存各的单机拿到任何一个文件都拼不出完整模型。正确做法是先声明聚合策略from torch.distributed.fsdp import ( FullStateDictConfig, StateDictType, ) save_policy FullStateDictConfig(offload_to_cpuTrue, rank0_onlyTrue) with FSDP.state_dict_type(model, StateDictType.FULL_STATE_DICT, save_policy): state_dict model.state_dict() if dist.get_rank() 0: torch.save(state_dict, model_full.pt)这里有两个细节必须注意。一是offload_to_cpuTrue否则聚合完整state_dict的那一瞬间rank0上的显存会暴涨到全模型大小直接OOM。二是rank0_onlyTrue让其他rank不用空等省一次内存拷贝。如果你是断点续训最好直接用torch.distributed.checkpoint保存分片格式的checkpoint加载时按rank对号入座。全量state_dict虽然直观但在几百卡规模下聚合开销很大。3.4 一组实测数据显存降了多少吞吐损失多少我拿GPT-2 XL约1.5B参数在4张A100 40GB上做了对比实验batch size1序列长度512BF16混合精度方案每卡峰值显存每卡吞吐tokens/s说明DDPOOM-batch1直接爆显存FSDP FULL_SHARD约21GB约920首次稳定跑通FSDP FULL_SHARD prefetch约23GB约1280通信被计算掩盖吞吐明显提升FSDP activation checkpoint约14GB约1050显存最低但重算带来吞吐损失这些数字是单次实验的示意数据不同框架版本和卡间拓扑会有差异但趋势是稳定的FSDP把DDP根本跑不了的模型变成了可训练代价是吞吐它本来就不占优势通过prefetch能补回来一部分。小模型在FSDP下反而比DDP慢是正常的别慌。4. 用FSDP做训练踩过的坑按排查链路逐个说接入FSDP只是开始真正折磨人的是各种配置对了但效果不对的问题。下面这几个坑我全踩过每个都附上排查思路。4.1 坑一模型没包对All-Gather反而拖垮性能现象是最直观的显存确实降了但训练速度比DDP慢三四倍GPU利用率一直在50%上下晃。排查过程我分了三步。第一步先确认wrap粒度在模型初始化时打印每个FSDP单元的参数量发现整个模型只被包了一次——这意味着每个step都在all-gather全模型参数通信量巨大。第二步检查auto_wrap_policy发现transformer_layer_cls传错了类模型里实际的Block类是GPT2Block而我传成了别的包装类等于没匹配上。第三步修好之后再把每个Block单独wrapGPU利用率立刻上到80%以上。判断wrap是否生效有一个土办法看训练日志里FSDP初始化的单元数量单元数应该等于Transformer Block层数而不是1。4.2 坑二梯度裁剪在分片下失效数值悄悄变差之前直接在训练循环里写了torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)。训练了几天loss曲线看起来正常但下游评估分数比DDP跑出来的低了两个点。原因在于FSDP下每个rank只拥有参数的一个分片clip_grad_norm_拿着本rank的局部梯度去算范数算出来的不是全局梯度范数。等于是每个分片各自按局部标准裁剪整体梯度方向被扭曲了。FSDP包装类提供了对应的方法应该这么写model.clip_grad_norm_(max_norm1.0)它会内部做一次全量范数的all-reduce得到正确的全局范数再统一裁剪。遇到loss曲线看起来正常但评估指标不对的情况优先检查这类看起来能跑但语义变了的操作。4.3 坑三CPU offload打开之后NCCL反复超时跑7B模型时我为了压显存开了cpu_offload结果训练每隔几十个step就报一次NCCL超时watch dog把整个训练打挂。排查后发现不是代码逻辑问题而是CPU和GPU之间的传输太慢。参数offload到CPU后每次前向都要从CPU搬到GPUall-gather的源变成了内存带宽比NVLink低一个数量级。通信耗时激增超过了NCCL的默认超时阈值。两个解法一是调大初始化时的超时时间dist.init_process_group(backendnccl, timeouttimedelta(minutes30))二是只有在预热结束、模型真正稳定进入长训练阶段后再开启offload避免在启动期就撞上超时窗口。说实话CPU offload只适合显存极度紧张但CPU内存管够的场景。它的吞吐损失经常达到30%到50%如果有条件加卡或换大显存卡优先改善硬件条件。4.4 坑四激活显存没管住省下的内存又被吃回去FSDP管住了参数、梯度和优化器状态但激活值它管不了。长序列、大batch训练时激活值占的显存往往比参数还多。我在尝试把batch从1加到8时明明参数分片的显存预算还很富裕照样OOM因为激活值线性膨胀。解决思路是配合使用activation checkpointing也就是梯度检查点model.gradient_checkpointing_enable()注意这里有个联动效应activation checkpoint会重新计算前向而重算需要再次all-gather参数等于通信次数变多了。所以通常把activation checkpoint和prefetch一起开让提前取的参数正好覆盖重算的消耗。我实测的结果是FSDP activation checkpoint能把显存峰值再砍掉三分之一左右吞吐损失约15%在显存紧张时这笔交易非常划算。4.5 坑五checkpoint保存/加载把训练进度归零最惨的一次是训练了3个小时保存checkpoint时直接OOM整个进度报废。原因前面提过FullStateDictConfig没有开offload_to_cpurank0聚合完整state_dict时峰值显存瞬间超过40GB。排查链路是这样的先看OOM日志落在哪个张量分配上发现是CPU端聚合时GPU显存飙升再查保存代码发现用了默认的state_dict_type等于每张卡都在尝试保存全量分片而实际只拿得到自己那部分。最后改成FullStateDictConfig offload_to_cpu rank0_only问题解决。另一个隐蔽问题出现在断点续训加载sharded checkpoint时如果保存和加载使用的world size不一致比如4卡训练、8卡续训分片对应关系会错乱。要么保持world size一致要么用torch.distributed.checkpoint让它按新拓扑重新分布不要自己手动拼。5. 什么时候别选FSDP以及FSDP之后的路最后聊点更宏观的。FSDP不是万能的它有自己的适用边界也有明确的演进方向。5.1 FSDP、DeepSpeed ZeRO与张量并行的取舍很多人在选型时会把FSDP和DeepSpeed ZeRO放在一起比较实际上两者的分片思想几乎一样区别主要在工程生态上。DeepSpeed更早成熟offload和弹性配置更丰富但引入了一个较重的第三方依赖FSDP是PyTorch原生和torch.compile、DTensor、分布式checkpoint的配合更顺滑。我的建议是新项目优先FSDP老项目已经在用DeepSpeed就没必要强行迁移。真正需要想清楚的是FSDP与张量并行TP的分工。FSDP按状态分片通信发生在step边界TP按算子分片通信发生在每一层内部。单个Transformer Block里TP的通信量是FSDP的好几倍所以TP只建议在高带宽的节点内用跨节点TP基本不要碰。对应地FSDP的通信可以跨节点扩展性更好。方案分片对象通信频率适合规模上手难度DDP不分片每个step一次全量归并模型放得下单卡低FSDP / ZeRO-3参数梯度优化器状态每层多次但可prefetch数十B内单机/多节点中张量并行权重矩阵分块每层多次高频通信数十B以上配合FSDP高流水并行层分段step边界少量通信百B级别多节点较高5.2 更大模型下的组合打法FSDPTPHSDP单机8卡的场景里20B以上的模型用纯FSDP也能跑但每个Block的参数量很大all-gather时显卡间要搬大量数据吞吐不理想。这时候通常把TP和FSDP叠起来TP切块让单卡上的Block变小FSDP再做状态分片继续压显存。PyTorch 2.4之后的DTensor把这套组合简化了很多TP的切分布局和FSDP的分片布局可以在同一张分布图上描述。多节点场景还要考虑HSDP即混合分片。它的思路是节点内做全分片节点间做复制all-gather不出机柜这样跨节点的通信量大幅下降。HSDP最适合机内带宽高、机间带宽一般的典型集群现在很多百B级训练默认就是HSDP TP的组合。5.3 FSDP2与DTensor分片正在变成一种通用布局FSDP2是FSDP在PyTorch 2.4之后的重大重构核心变化是底层改用DTensor来表达分片。以前FSDP是自己管一套扁平参数和分片逻辑现在分片变成了一种张量分布布局和TP的布局描述共用同一套语言。这个变化带来的实际好处是你可以对模型的不同部分声明不同的分布方式比如Embedding用TP切、attention用FSDP分片框架负责生成对应的通信。torchtitan等新一代训练框架已经在这个基础上跑出了Llama系列模型我试下来代码确实清爽很多但资料和踩坑记录比经典FSDP少生产环境迁移前建议先在测试集群练熟。5.4 我对十年演进的一点体会回到标题说的十年演进。我个人的体会是FSDP不是某个团队灵光一现造出来的东西而是分布式训练走到一定阶段后的必然产物。DDP解决了怎么高效通信模型并行解决了怎么拆开模型ZeRO解决了拆开之后怎么减少冗余FSDP只是把这几件事用一套统一的框架整合起来让普通工程师不再需要自己拼积木。演进的主线其实很朴素从每张卡复制一切到整个集群共享一份状态按需拼装。未来不管是更激进的分层显存利用还是异步训练、分离式训练架构这个按需取用的思路大概率会一直保留下去。我在实操中的体会是真正决定训练效率的往往不是选哪个并行方案而是你理解了多少通信如何被藏进计算里的细节。把这层节奏感练出来FSDP只是起点。