第29讲:显存分片技术——ZeRO-1/2/3 与 PyTorch FSDP 深度剖析
主讲人:👓 Ringi(大厂 AI Infrastructure 资深性能架构师)
所属专栏:《AI_Infra大话西游之水滴石穿》 ➔ Module 04: 大模型分布式训练系统
篇章范式:🌐 大规模分布式训练系统范式(Distributed Training Systems Paradigm)
源码与实验环境:NVIDIA A100-SXM4-80GB / H100-SXM5-80GB | CUDA 12.4 | Python 3.10 | PyTorch 2.3+ | DeepSpeed 0.14+
知识底账索引:
- 数据并行与 FSDP 详解:4.1 数据并行详解(AIInfraGuide)
- FSDP 实战与策略配置:4.2 PyTorch 数据并行从原理到实战(AIInfraGuide)
- 集合通信原语量化:2.1 集合通信原语详解(AIInfraGuide)
- ZeRO-DP 源码剖析:完全分片数据并行 FSDP 实现(AISystem)

📑 目录导航
- 0. Ringi 开场:生产真实现场与痛点冲突
- 1. 显存状态第一性原理:为什么大模型状态可以分片?
- 2. 通信代价的数学推导:为什么 FSDP 是 ,DDP 是 ?
- 3. PyTorch FSDP 系统架构与内核机制(Under the Hood)
- 4. 工业生产的终极解:
HYBRID_SHARD(混合分片)架构 - 5. 全场景实战与实验代码(Minimal Runnable Code)
- 6. Ringi 避坑指南与生产黄金准则
- 7. Ringi 5 点核心速记口诀、自我检验清单与课后深度思考题
- 8. 📚 参考资料与核心源码/经典论文指引
- 附录:Appendix A — 大厂硬核高频面试题与白板推导(Interview Drill)
0. Ringi 开场:生产真实现场与痛点冲突
0.1 真实工程矛盾:DDP 的“显存死局”与 FSDP 的“通信杠杆”
在第 28 讲中,我们详细解构了经典数据并行(DDP)的优雅之处:去中心化的 Ring-AllReduce 将单卡通信量死死压制在 ,配合异步 Bucket Overlap,在单机 8 卡上几乎能做到完美的线性加速。 但是,当大模型的参数量一路狂飙突进到 7B、13B、70B 时,DDP 的天花板被瞬间撞碎:- DDP 只能摊薄计算时间,绝不能摊薄单卡显存!
- 对于任意一个参数量为 的模型,用混合精度 AdamW 训练时,每张 GPU 必须死死背负 的完整静态显存(权重 + 梯度 + 优化器 );
- 面对一个经典的 7B 模型( ), !哪怕单张卡 Batch Size 设为 1,哪怕调用 10,000 张卡,单张 80GB 的 A100/H100 显卡连静态模型都放不进去,任务在初始化阶段就因 OOM 胎死腹中!
“完全不需要改动算子!我们依然做数据并行,但把单卡冗余的 静态状态彻底切成 份。平时各存 ,算哪一层临时拼哪一层,算完立刻销毁扔掉!”显存从 暴降至 ,但代价是:每步通信量从 增加到了 (增加了整整 50%!)。
如何驾驭这多出来的 50% 通信?如何不让它拖垮千卡集群的 MFU?这就是性能工程的核心战役。
0.2 线上真实事故复盘:某 70B 模型盲开全切分引发的“跨机网络大堵塞”
2024 年秋,国内某头部科技团队在 32 台 8 卡 H800(共 256 张 GPU)集群上预训练 70B 稠密语言模型。 团队初次采用 PyTorch FSDP,算法同学直接按网上的快速入门教程,给最外层的模型套上了一个大的FSDP(model):
- 显存瞬间尖刺爆炸:前向传播刚一启动,由于整模型被作为一个单元,FSDP 在第 0 步就强行发起了一个覆盖全部 70B 参数的超巨型 AllGather!单卡瞬间试图在显存里拼出完整的 140 GB 权重,显存分片完全失效,多台节点当场 OOM 熔断;
- 紧急打补丁后又陷入通信黑洞:架构师介入配置了简单的参数量包裹阈值,虽然显存降下来了,但单步时间(Step Time)却高达 14.8 秒,GPU 算力利用率(MFU)仅有可怜的 12%!
- 深入 Profiler 时间线抓出真凶:
- 该集群节点内是 NVLink(400GB/s),但节点间仅配备了双口 100G RoCE(跨机双向带宽仅约 25GB/s);
- 采用
FULL_SHARD意味着前向 80 层 Transformer Block 的每一层都要跨机发起 AllGather,反向每一层都要跨机发起 AllGather + ReduceScatter; - 跨机 100G 网络被这 的海量数据流彻底打穿,交换机 PFC(Priority Flow Control)拥塞流控频繁触发,GPU 有整整 85% 的时间都在空等跨节点网络发包!
- 将分片策略彻底重构成
HYBRID_SHARD:将高频的参数全切分(FULL_SHARD)严格限制在机内 8 卡(走 400GB/s NVLink 极速通道);跨 32 个节点之间退化为传统数据并行(仅在反向结束时跨机做一次 AllReduce); - 配合配置精细的 TransformerBlock 递归包裹 与 Backward Prefetch;
- 单步迭代时间从 14.8 秒直接压缩至 2.9 秒,吞吐暴涨 5.1 倍,MFU 成功挽救至 52.4%!
0.3 显存分片架构与 FSDP 策略演进速查表
1. 显存状态第一性原理:为什么大模型状态可以分片?
💡 架构全景速览:在深潜源码前,先在白板上建立坚不可摧的 ZeRO 显存分片演进与 PyTorch FSDP 物理底账。
1.1 静态显存的三座大山剖析: 底账的结构性冗余
在第 27 讲中,我们手算了混合精度 AdamW 训练的静态显存公式: 仔细审视这 的构成,一个极度不合理的架构缺陷浮出水面:- 优化器状态独占 (占全卡静态显存的整整 75%!);
- 在传统的 DDP 中,假设我们有 64 张卡,这 64 张卡在更新参数时,维护的优化器动量数值是 100% 完全相同的!
- 为什么 64 张卡要各自存一份一模一样的 FP32 动量和 Master 权重?这是巨大的物理显存浪费!
1.2 冗余消除的三阶演进:ZeRO-1、ZeRO-2 与 ZeRO-3
Samyam Rajbhandari 等人在 ZeRO 论文中,提出了一套阶梯式的“手术刀方案”:1. ZeRO-1(优化器分片):
- 每张卡只保存 的优化器状态(Master 权重、一阶动量、二阶动量);
- 反向传播时,依然做传统的全卡梯度 AllReduce(通信量 );
- 更新时,每张卡只用自己负责的那部分梯度更新自己负责的那 权重;
- 更新完毕后,执行一次轻量的跨卡分片收集(AllGather 权重,或等效广播);
- 结论:消灭了 75% 冗余的大头,单卡直接省下约 优化器显存,且通信量与 DDP 完全相同!
2. ZeRO-2(梯度分片):
- 既然每张卡只更新 的权重,那为什么每张卡要存全量的梯度?
- 在反向传播求导时,直接使用 ReduceScatter(规约分散)替代 AllReduce!
- 每张卡只接收并保存自己负责更新的那 梯度分片;
- 通信量分析:
- 结论:单卡静态显存降至 ,通信量依然是 零增加!这是工业界性价比极高的模式(PyTorch FSDP 中的
SHARD_GRAD_OP)。
3. ZeRO-3(参数全分片):
- 连模型参数(权重)也不保留全量了,每张卡只持久常驻 权重;
- 前向算到某一层,通过 AllGather 临时拼出来,算完立刻销毁;
- 反向算到某一层,再次 AllGather 临时拼出来,求导后通过 ReduceScatter 同步并分发梯度分片;
- 结论:单卡显存暴降到 ,但代价是多了一次前向 AllGather,通信量上升至 。
1.3 通信与显存的黄金权衡曲线(Pareto Frontier)
以一个 7B 模型( )在 8 张 80GB A100 上的表现为例:2. 通信代价的数学推导:为什么 FSDP 是 ,DDP 是 ?
2.1 No Naked Formula 2.0 穿透 FSDP 通信量模型
我们必须走完 No Naked Formula 2.0(公式五步穿透法),将这多出来的 到底花在了哪里手算清楚。
① 为什么需要算它?
很多团队盲目将 DDP 迁移到 FSDP 后,发现集群网络带宽被打爆,训练步时变慢。只有推导出 的精确构成,才能用数学确定当前的网卡物理带宽(GB/s)到底能不能在计算时间内把这笔账盖住。② Mental Model(物理直觉比喻)
想象 4 个人合伙拼装一台复杂的机器(80 个零件模块):- DDP 模式:每个人家里都买了一套完整的 80 个零件。大家各自拼装,最后大家围在一起,把各自调校的数据抄写汇总一遍(AllReduce);
- FSDP 模式:每个人家里只放 20 个零件(显存大省!)。但是拼装第 1 个模块时,你必须先把其他 3 个人的零件借过来拼在一起,拼完测量完数据,立刻把别人的零件还回去(前向 AllGather);等到检查故障反向调校时,你又必须重新把别人的零件再借过来拼一次(反向 AllGather);调校完后,大家把测试结果计算出平均数,各自只带走属于自己那份的数据(反向 ReduceScatter)。
- 借零件是需要时间的,多借了一次,就多付了一次搬运费!
③ Tiny Calculator(极简数字手算)
设:- GPU 卡数 ;
- 某一层的权重参数量为 ;
- 静态时,每张卡只存 。
- 前向计算该层:
- 必须通过 AllGather 拼出完整 ;
- 每张卡把自己持有的 发送给其他 3 张卡,同时接收其他卡各 ;
- 单卡发送量:
- 算完前向后,释放非本地的 ;
- 反向求导该层:
- 必须再次执行 AllGather 拼出完整 (因为前向完已经释放了!);
- 单卡再次发送: ;
- 梯度同步并分片:
- 求导计算出的梯度也是 ;
- 执行 ReduceScatter,把 4 张卡的梯度累加,并切分成 4 份,每卡只收回属于自己的那 聚合梯度;
- 单卡发送量:
④ Formal Model(标准公式与渐进极限)
累加全模型所有层(全模型参数为 ),在包含 张 GPU 的 FSDP(FULL_SHARD)集群中:
单张 GPU 在一个完整的训练迭代(Step)中发送的总数据量严格为:
当 时:
对比 DDP 的通信量公式(基于第 28 讲证明的 ):
⑤ Sanity Check(数量级校验)
对于 70B 模型( 参数,BF16 下为 权重):- DDP 模式单卡每步通信量:
- FSDP 全切分单卡每步通信量: !
- 差额净增:单卡整整多出了 的物理传输负荷!
2.2 单层 Transformer Block 通信时序白板拆解
我们把一个 FSDP Unit(通常为一个 Transformer Block)的前向与反向流水线绘制在时序图上:2.3 为什么增加 50% 通信量在工业界依然“极度划算”?
多出了 50% 通信量,为什么从 Meta 到各个大模型巨头依然把 FSDP 作为标准基础设施? 掏出工程算盘手算收益与代价的收支平衡:- 显存杠杆极大:单卡显存从 暴跌到 (节省了整整 的单卡物理显存!)。这使得原本根本不能跑的模型可以跑了,原本只能设 Batch Size = 1 的任务可以直接拉到 Batch Size = 8;
- 机内高带宽完全能够吸收增量:在具备 NVLink(450~900 GB/s)的单机 8 卡节点内,搬运这额外的 只需要:
3. PyTorch FSDP 系统架构与内核机制(Under the Hood)
3.1 核心抽象:FSDP Unit 分片单元与计算生命周期
在 PyTorch FSDP 中,最小的管理与通信粒度被称为 FSDP Unit(分片单元)。 在底层,FSDP 并没有把每个孤立的线性层单独包成一个 Unit。因为如果每个 Linear 层都单独触发一次 AllGather,数千个微小 Kernel 会让 GPU 陷入灾难级的小包发射排队。工业级标准划分:将一个完整的 Transformer Decoder Block(包含 RMSNorm + Attention + RMSNorm + SwiGLU)包裹为一个独立的 FSDP Unit。
3.2 自动包裹策略(Auto Wrap Policy)的生与死
这是所有使用 PyTorch FSDP 的工程师最容易踩雷的深坑:如果不配置 Auto Wrap Policy,或者包裹方式错误,FSDP 的显存节省会彻底归零!错误范式:全模型单层外壳包裹(Flat Wrap)
- 物理机理:FSDP 把整个 70B 模型当成了一个单一的 Unit;
- 前向第 0 步:直接发起一个囊括全模型所有 80 层参数的巨型 AllGather;
- 显存血崩:为了执行第 1 层计算,系统不得不先把全模型 140 GB 权重在显存中全量拼装出来;
- 惨烈结局:分片形同虚设,单卡峰值显存直接冲破物理上限,原地 OOM 暴毙!
正确范式:基于 Transformer Block 递归包裹(Module Wrap)
- 物理机理:FSDP 递归遍历网络树,为每一个
LlamaDecoderLayer实例化一个独立的 Unit; - 流式组装:算第 1 层时只拼第 1 层的参数(仅几百 MB),算完立刻释放;算第 2 层时再拼第 2 层;
- 显存恒定:全网常驻的始终只有单个 Block 拼出来的微小内存,显存峰值被死死压制!
3.3 预取流水线(Backward & Forward Prefetch)的重叠魔法
既然 FSDP 多出了 50% 的通信量,如何保证 GPU 不干等?PyTorch FSDP 提供了强大的流水线机制:
backward_prefetch 与 forward_prefetch。
1. 前向预取(Forward Prefetch):
- 当 Default Stream 正在执行 Layer 1 的 GEMM 矩阵乘时;
- 专用的 NCCL 通信 Stream 不闲着,提前向其他卡发起 Layer 2 的 AllGather 请求;
- 当 Layer 1 计算完毕时,Layer 2 的完整权重早已经在 HBM 缓冲区拼装完毕,计算核心零等待无缝切入!
2. 反向预取(Backward Prefetch - BACKWARD_PRE vs BACKWARD_POST):
- 在反向传播中,靠后层的梯度算完后需要执行 ReduceScatter;
BACKWARD_PRE:在当前层反向计算开始之前,就提前发起更前一层权重的 AllGather;BACKWARD_POST:在当前层反向求导一结束、发射 ReduceScatter 的同时,立刻链式触发下一层的 AllGather;- 实测结论:在开启
backward_prefetch=BackwardPrefetch.BACKWARD_PRE后,FSDP 的端到端通信暴露时间缩短了 70% 以上!
3.4 FSDP1 (Module Wrapper) 到 FSDP2 (DTensor / Per-Param) 的演进
截至 PyTorch 2.x,官方正在全力推进 FSDP2(API 为torch.distributed.checkpoint 与 fully_shard):
4. 工业生产的终极解:HYBRID_SHARD(混合分片)架构
4.1 物理网络的阶梯现实:NVLink 900GB/s vs 机间 IB 50GB/s
在大规模 AI 集群中,通信网络存在着极其悬殊的“带宽阶梯”:FULL_SHARD,就相当于把前向和反向每层的 AllGather 强行推向了只有 50 GB/s 的机间慢速网络,导致千卡集群性能彻底崩盘!
4.2 混合分片机制:机内 FULL_SHARD + 机间数据并行

HYBRID_SHARD(混合分片策略) 给出了终极工业解:
- 机内 8 卡(Intra-Node):
- 组成一个局部进程组(Local Process Group,大小为 8);
- 在机内执行
FULL_SHARD:将参数、梯度、优化器状态切成 8 份; - 依赖机内 900GB/s 的 NVLink 飞速完成每层的 AllGather 与 ReduceScatter;
- 机间节点(Inter-Node):
- 跨机器之间组成一个全局数据并行组(Replication Group);
- 跨机之间不切分参数,退化为经典的数据并行(DDP);
- 仅在反向传播全部结束时,跨节点网卡执行一次聚合通信。
4.3 多机大模型训练的黄金参数组合
在工业级生产实践中,面对多机分布式预训练,推荐的标准配置矩阵如下:5. 全场景实战与实验代码(Minimal Runnable Code)
5.1 实验一:原生 PyTorch FSDP 多进程分片与前向还原最小实战
本实验通过纯原生 PyTorchtorch.distributed.fsdp 启动 2 个 Worker 进程,演示:
- 构造带 Transformer Block 的网络结构;
- 配置基于类的
transformer_auto_wrap_policy递归分片; - 打印分片前后的参数物理尺寸(验证
FlatParameter切分); - 验证前向传播与反向传播的梯度更新正确性。
5.2 实验二:工业级显存分片决策与显存/通信模拟评估器 fsdp_capacity_simulator.py
面对任意模型和集群架构,架构师绝不靠肉眼去猜该选哪种策略。本脚本实现了工业级显存分片多维评估器,精确输出 DDP、ZeRO-1、ZeRO-2、ZeRO-3 及 HYBRID_SHARD 的静态显存、单步通信量、通信时间及决策建议:
6. Ringi 避坑指南与生产黄金准则
6.1 7 大常见小白认知误区 vs 大厂 AI Infra 正确物理认知
6.2 生产 FSDP / ZeRO 黄金 Checklist
- 1. 【严禁未包装直出】:生产使用 PyTorch FSDP 必须显式传入
auto_wrap_policy,强制按LlamaDecoderLayer等核心 Block 粒度进行单元切分。 - 2. 【跨节点首选 HYBRID_SHARD】:多机分布式训练若跨机网络非 8x800G IB 顶配互联,必须优先采用
ShardingStrategy.HYBRID_SHARD抑制跨机网络风暴。 - 3. 【预取流水线必开】:必须开启
backward_prefetch=BackwardPrefetch.BACKWARD_PRE,利用双 CUDA Stream 将下层参数 AllGather 掩盖在当前层求导时间内。 - 4. 【激活值协同重算】:FSDP 仅切分模型状态,无法解决长序列动态激活值;对于长文本任务,必须与
checkpoint_wrapper(选择性激活重算)协同启用。 - 5. 【CPU Offload 严格评估】:生产分布式预训练严禁开启
cpu_offload=True;仅在轻量微调或单卡调试极限 OOM 边缘时作为保命兜底手段。 - 6. 【限制 AllGather 发射窗口】:配置
limit_all_gathers=True,防止预取流过早拼装多个后续层参数导致前向瞬时显存反向飙升。 - 7. 【混合精度类型统合】:在 FSDP 的
MixedPrecision中,将param_dtype与reduce_dtype严格统一配置为torch.bfloat16,消除运行时隐式类型转换开销。
7. Ringi 5 点核心速记口诀、自我检验清单与课后深度思考题
7.1 5 点押韵核心速记口诀
7.2 10 条白板自我检验清单
- 能否闭卷默写出混合精度 AdamW 训练下,模型参数、梯度与优化器各自占用的字节比例?
- 为什么 ZeRO-1 和 ZeRO-2 可以在大幅削减显存的同时,做到每步通信量与 DDP 严格相同(都是 )?
- 能否在白板上推导为什么 ZeRO-3 / FSDP 的单步通信量是 ?这多出的 发生在哪个阶段?
- 如果对一个 Transformer 模型不设置任何 Auto Wrap Policy,直接整体外包一层 FSDP,底层前向会发生什么?
- 为什么说优化器状态(Optimizer States)是大模型训练静态显存中“性价比最高”的切分目标?
- 能否画出 FSDP 中
forward_prefetch和backward_prefetch如何通过专用 Stream 掩盖通信的时间线图? - 阐明
HYBRID_SHARD的设计哲学:它是如何根据 NVLink 与 InfiniBand 的物理带宽差异进行分层切分的? - 在 FSDP 中,反向传播的梯度同步为什么使用的是
ReduceScatter,而不是 DDP 中的AllReduce? - 为什么说 FSDP 能够解决“单卡装不下”的问题,但对超长文本下的“激活值显存爆炸”却无能为力?
- FSDP1(模块包装器)与 FSDP2(DTensor 参数级切分)在多维混合并行(如 TP+FSDP)时有何根本性优势?
7.3 3 道高阶开放式课后思考题(含极限 Corner Case)
- 【极限动态显存与 FSDP 的踩踏事故】:在长文本(如 32K)训练中,假设我们开启了 FSDP FULL_SHARD,单卡静态显存被压缩到了极低的 10 GB。但是在反向传播阶段,由于当前层需要保留完整的输出激活值、同时正在拼装完整参数分片,且上一层的 ReduceScatter 缓冲区尚未完全释放。这种微观时间片上的“三军汇聚”是如何导致显存瞬间刺穿 OOM 的?工程上如何通过
limit_all_gathers避免流水线过冲? - 【DTensor 抽象下的 FSDP2 革命】:在 PyTorch 2.x 的 FSDP2 中,底层彻底重构为了基于
DTensor的 Sharding。请从张量步长(Strides)、连续性(Contiguous)和内存视图(View)的角度分析:DTensor 是如何做到既能维持单个 Parameter 的独立切片,又能在底层调用 NCCL 时零拷贝拼接成大 Buffer 发射通信的? - 【ZeRO++ 的网络带宽极致压榨】:微软在 ZeRO 的基础上进一步提出了
ZeRO++,利用量化与跨节点辅助通信来压榨带宽。请推导:如果在前向 AllGather 时将参数动态量化为 INT8 或 FP8 传输,通信量能从 压缩到多少?这会对反向传播的梯度精度产生什么连锁影响?
8. 📚 参考资料与核心源码/经典论文指引
权威学术论文:
- ZeRO 奠基之作:Rajbhandari et al., “ZeRO: Memory Optimizations Toward Training Trillion Parameter Models”, SC 2020. arXiv:1910.02054
- PyTorch FSDP 官方系统论文:Zhao et al., “PyTorch FSDP: Experiences on Scaling Fully Sharded Data Parallel”, VLDB 2023. arXiv:2304.11277
- ZeRO-Offload 架构:Ren et al., “ZeRO-Offload: Democratizing Billion-Scale Model Training”, USENIX ATC 2021. arXiv:2101.06840
- ZeRO++ 极致通信优化:Wang et al., “ZeRO++: Extremely Efficient Collective Communication for Giant Model Training”, 2023. arXiv:2306.10209
工业级开源源码指引:
- PyTorch FSDP 官方实现:
torch/distributed/fsdp/fully_sharded_data_parallel.py(包含_auto_wrap与通信挂载) - PyTorch FSDP 预取流水线:
torch/distributed/fsdp/_runtime_utils.py(核心流同步与 prefetch 逻辑) - DeepSpeed ZeRO-3 引擎:
deepspeed/runtime/zero/stage3.py(包含状态机、参数分片与动态 fetch)
本地 AI_BOOK 知识库精准映射:
- 数据并行详解:4.1 数据并行详解.md
- FSDP 实战指南:4.2 PyTorch 数据并行从原理到实战.md
- 集合通信原语:2.1 集合通信原语详解.md
- FSDP 源码剖析:03ZeRODP.md
附录:Appendix A — 大厂硬核高频面试题与白板推导(Interview Drill)
面试真题 1:请白板画图推导 FSDP 训练一个完整的 Transformer Block,前向和反向分别发生了哪几次通信?为什么总通信量是 ?
考察维度:FSDP 通信机制精细化追踪、集合通信量推导、反向 ReduceScatter 动因。
标准推导路径:
- 单层 Block 参数与切片基准:
- 设单层参数量为 ,卡数为 ;
- 静态常驻状态下,每张卡只持有大小为 的参数分片。
- 阶段一:前向传播(Forward Pass):
- 在计算该 Block 前,必须持有完整的 权重矩阵;
- 触发通信:执行 AllGather,收集所有卡的参数切片;
- 单卡发送量: ;
- 计算完毕后,立即执行内存释放(Free),显存回落至 。
- 阶段二:反向传播(Backward Pass):
- 在求导计算时,由于前向权重已被销毁,必须再次获取完整参数;
- 触发通信 1:再次执行 AllGather,重新拼装出 ;
- 单卡发送量: ;
- 执行矩阵求导,计算出该层完整的权重梯度 ;
- 触发通信 2:由于每张卡最终只更新属于自己的 权重分片,因此无需将全量梯度广播回所有卡,而是执行 ReduceScatter(全局规约求和并分散切片);
- 单卡发送量: ;
- 随后释放完整参数,每张卡仅持有自身负责的 聚合梯度。
- 全流程累加:
面试真题 2:为什么 FSDP 必须配置 Auto Wrap Policy?如果不配置或者把整个网络包成一个顶层 FSDP,底层会发生什么物理灾难?
考察维度:PyTorch FSDP 架构抽象、内存流水线与显存峰值控制。
标准参考答案:
- FSDP 的分片单元哲学:
- FSDP 的显存节省依赖于“流式按需加载(On-demand Streaming)”——即只有当计算推进到某一特定子模块时,才将该模块参数拼装到显存,计算完立即卸载;
- 这个加载与卸载的控制边界就是 FSDP Unit。
- 不配置 Wrap Policy 的物理灾难:
- 如果直接对整个根模型
FSDP(model)进行包裹,整个模型(无论是 32 层还是 80 层)被归为唯一的一个巨型 FSDP Unit; - 前向第 0 步:在执行整个网络的前向传播前,FSDP 必须一次性将全模型所有层的参数全部执行 AllGather 拼装出来;
- 显存瞬间雪崩:对于 70B 模型,这意味着单卡必须在显存中强行开辟 140 GB 连续空间容纳全量权重;
- 灾难后果:分片带来的显存节省在第一毫秒就被彻底抹平,单卡显存峰值直接飙升到与未切分状态完全相同,系统当场报
CUDA out of memory崩溃。
- 如果直接对整个根模型
- 正确工程实践:
- 必须通过
transformer_auto_wrap_policy将每一个单独的TransformerBlock(如LlamaDecoderLayer)包裹为独立的叶子 FSDP Unit; - 保证系统在任意时刻,显存中拼装出来的完整参数最多只有当前正在计算的这 1 个 Block(仅占全模型的 ,通常只有几百 MB),实现显存的大幅压缩。
- 必须通过
面试真题 3:在什么硬件网络条件下,ZeRO-2 / SHARD_GRAD_OP 的端到端训练吞吐反而会大幅超越 ZeRO-3 / FULL_SHARD?
考察维度:网络带宽瓶颈诊断、Communication-to-Computation Ratio、架构选型 Trade-off。
标准参考答案:
- 根本原因:通信量的本质差距( vs ):
- ZeRO-2 /
SHARD_GRAD_OP只切分优化器状态和梯度,参数全量常驻,每步单卡通信量严格为 (仅在反向时做一次 ReduceScatter); - ZeRO-3 /
FULL_SHARD参数全切分,每步单卡通信量为 (多了前向和反向两次 AllGather,通信量净增 50%)。
- ZeRO-2 /
- 发生性能反转的硬件网络工况:
- 跨机低带宽网络互联:当训练扩展到多机跨节点,且网络仅配备千兆、万兆网卡,或单口 100G RoCE 时;
- 网络成为全系统绝对瓶颈(Communication-Bound):此时机间网络带宽较窄,计算内核耗时远远小于数据传输耗时,多出来的这 通信量根本无法被前向/反向计算所掩盖(Overlap 彻底失效);
- 显存尚有裕量:如果单卡物理显存(如 80GB)在容纳了 ZeRO-2 的静态显存( )以及动态激活值之后仍有剩余;
- 选型决策结论:
- 在此工况下,硬上 ZeRO-3 会让每张 GPU 花费大量时间在慢速跨机网络上空等 AllGather;
- 而采用 ZeRO-2,直接抹掉了 33.3% 的网络传输负载,使得通信等待时间大幅缩短,因此端到端训练吞吐(Tokens/s)和 MFU 往往能够高出 ZeRO-3 整整 30% ~ 50% 以上!