TorchTitan FSDP2 完全指南DTensor 原生参数分片的分布式预训练实践【免费下载链接】AI-Research-SKILLsComprehensive open-source library of AI research and engineering skills for any AI model. Package the skills and your claude code/codex/gemini agent will be an AI research agent with full horsepower. Maintained by Orchestra Research.项目地址: https://gitcode.com/gh_mirrors/ai/AI-Research-SKILLs本篇技术指南以 TorchTitan 的 FSDP2 参考文档 为骨架系统讲解 PyTorch 新一代全分片数据并行Fully Sharded Data ParallelFSDP2的设计动机、fully_shardAPI 用法、ZeRO 等价分片策略、Meta 设备初始化、混合精度与 HSDP 混合分片并结合本仓库中 TorchTitan 主 SKILL、PyTorch FSDP2 参考材料 与 分布式检查点文档 进行源码级深化。读完本文你将掌握如何用 FSDP2 替换 FSDP1 完成大模型预训练、如何在 TorchTitan 的 TOML 配置中启用 FSDP/HSDP以及 Sharded State Dict、Meta-Device 初始化等关键工作流。为什么需要 FSDP2FSDP2 是 PyTorch 对 Fully Sharded Data ParallelFSDPAPI 的一次重写其核心是移除FlatParameter抽象从而获得更好的可组合性composability与更简单的实现。在 FSDP1 中一个模块的所有参数会被拍平flatten成一个FlatParameter再统一分片这带来两个问题一是参数形状被破坏用户难以直接操纵分片后的参数二是与张量并行TP、torch.compile等其它特性的组合不够自然。FSDP2 则把分片参数直接表示为DTensor让分片这一概念在张量层面显式可见。从本仓库 PyTorch FSDP2 教程参考 可以看到FSDP2 的分片对象是逐参数per-parameter的而非逐模块拍平这使得操纵单个分片参数、构建无通信的分片状态字典都成为可能。相对 FSDP1 的三大核心改进改进维度说明DTensor 分片分片后的参数是沿 dim-0 分片的DTensor可直接读写操纵支持无通信communication-free的分片 state dict更优的内存管理通过避免recordStream引入内存占用确定且更低约降低 7%API 更简洁参数更少无需包装类wrapper class直接用fully_shard函数式接口其中内存管理这一点的意义值得展开FSDP1 依赖 CUDA 流记录recordStream来保证异步通信与计算的安全重叠这既增加了峰值内存开销也让内存行为难以预测。FSDP2 在 TorchTitan 官方 FSDP 文档 有收录中即被描述为面向生产环境的配置选择其确定性内存行为直接带来了可预测的显存规划能力。性能表现文档记载在Llama-7B 8 张 H100的配置下FSDP2 相比 FSDP1 实现了更高的 MFU模型浮点利用率峰值显存降低约7%同时 loss 曲线保持一致——即性能提升不伴随收敛质量损失。此外TorchTitan 主 SKILL 文档给出了更贴近真实预训练场景的 H100 基准数据见 SKILL.md 的 Performance benchmarks 小节模型GPU 数并行方案TPS/GPU说明Llama 8B8FSDP5,762基线Llama 8B8FSDP compile FP88,53248%Llama 70B256FSDP TP AsyncTP8762D 并行Llama 405B512FSDP TP PP1283D 并行需要说明的是这些数字来自仓库内文档记载的基准实际收益会随 GPU 型号、模型规模、通信拓扑和框架版本变化建议以自己环境下的实测为准。fully_shardAPI 参考FSDP2 的核心入口是torch.distributed._composable.fsdp.fully_shard其签名如下完整源码注释见 pytorch_fully_shard_api.mdfrom torch.distributed._composable.fsdp import fully_shard, MixedPrecisionPolicy, OffloadPolicy contract(state_clsFSDPState) def fully_shard( module: nn.Module, *, mesh: Optional[DeviceMesh] None, reshard_after_forward: Union[bool, int] True, mp_policy: MixedPrecisionPolicy MixedPrecisionPolicy(), offload_policy: OffloadPolicy OffloadPolicy(), ) - nn.Module:关键语义来自 PyTorch 官方 API 文档本仓库有收录meshDeviceMesh类型。1D mesh 即经典 FSDP 分片参数 placement 为(Shard(0),)2D mesh 即混合分片 HSDP一个维度复制、一个维度分片placement 为(Replicate(), Shard(0))。reshard_after_forwardTrue前向后释放未分片参数反向时重新 all-gatherFalse前向后保留未分片参数省去反向时的 all-gatherNone默认对非根模块为True、对根模块为False整数前向之后重分片到更小的 world size必须能整除 shard 维度大小对应 ZeRO hpZ 风格的分层分片。mp_policy混合精度策略控制参数、梯度归约、前向输出的 dtype详见下文混合精度小节。offload_policy卸载策略可控制参数/归约/输出的设备CPUOffloadPolicy还可额外指定优化器状态驻留设备。从 API 文档还可以提炼出三个必须遵守的使用约定从底向上bottom-up分片先对子模块调用fully_shard最后再对根模块调用。这直接影响通信分组的形成与通信计算重叠官方文档明确指出用户一般不应只对最顶层的根模块调用fully_shard。必须通过model(input)触发钩子fully_shard依赖前向/反向钩子完成 all-gather 与释放调度直接调用model.forward(input)会绕过钩子导致分片失效。优化器必须在fully_shard之后、基于 DTensor 参数构造优化器 step 必须在 DTensor 上进行先建优化器再分片会导致优化器持有的参数引用与分片后的 DTensor 不一致。分片策略与 ZeRO 等价对照FSDP2 用mesh 维度数 reshard_after_forward取值这一组合即可表达 FSDP1 与 DeepSpeed 中几乎所有主流分片策略FSDP2 配置FSDP1 等价DeepSpeed 等价1D mesh reshard_after_forwardTrueFULL_SHARDZeRO-31D mesh reshard_after_forwardFalseSHARD_GRAD_OPZeRO-22D mesh reshard_after_forwardTrueHYBRID_SHARDMiCS1D/2D mesh reshard_after_forward8int—ZeRO hpZ这张对照表是选择并行策略的核心决策工具ZeRO-3 / FULL_SHARD参数、梯度、优化器状态全部分片显存节省最大但每次前反向都要 all-gather 参数通信开销最高是 FSDP2 的默认行为。ZeRO-2 / SHARD_GRAD_OP只分片梯度与优化器状态参数保留完整副本前向后无需重新 all-gather通信量更小适合参数量适中的模型。HSDP / MiCS跨机复制、机内分片把昂贵的跨节点 all-gather 变成廉价的机内操作是大规模多节点训练的主流选择。ZeRO hpZint 语义前向之后只保留部分世界内的分片进一步降低通信适合对带宽敏感的场景。Meta-Device 初始化先分片、后物化FSDP2 支持在meta 设备上初始化模型不占任何显存完成分片之后再物化materialize到 GPU。这一模式避免了 FSDP1 中先完整建模型再分片导致的瞬时显存峰值是从源码层面支持超大模型训练的基石。import itertools import torch # 1. 在 meta 设备上初始化不占内存 with torch.device(meta): model Transformer() # 2. 应用 FSDP2 分片自底向上先子模块后根模块 for module in model.modules(): if isinstance(module, TransformerBlock): fully_shard(module) fully_shard(model) # 3. 此刻所有参数仍在 meta 设备上 for tensor in itertools.chain(model.parameters(), model.buffers()): assert tensor.device torch.device(meta) # 4. 在 GPU 上分配分片后的参数只分配本 rank 拥有的分片 model.to_empty(devicecuda) # 5. 初始化权重 model.init_weights()要点解读分片发生在物化之前fully_shard在参数还是 meta 张量时就完成 DTensor 分片规划to_empty(devicecuda)随后只为每个 rank 分配其持有的参数分片整体显存峰值即为最终稳态占用。init_weights()必须在to_empty之后此时参数才真正落在 GPU 上权重初始化如 custom-models.md 中要求实现的递归init_weights才有实际意义。这一模式与 TorchTitan 的torch.device(meta)初始化流程完全一致也是其能扩展到 405B 参数规模的关键前提见 SKILL.md 的 4D 并行工作流。State Dict 差异从全量到分片FSDP2 的分片状态字典是 DTensor 设计的直接受益者。由于分片参数本身就是DTensormodel.state_dict()与optim.state_dict()默认返回的即是分片状态字典无需跨 rank 通信操作FSDP1FSDP2model.state_dict()全量 state dict分片 state dict无通信optim.state_dict()本地 state dict分片 state dict无通信summon_full_params()支持改用DTensorAPI如full_tensor()梯度裁剪FSDP.clip_grad_norm_()nn.utils.clip_grad_norm_()实战含义保存分片 state dict 由各 rank 并行写出自己持有的分片天然适配 PyTorch Distributed CheckpointDCP 的多 rank 并行保存与**加载时重分片load-time resharding**能力——允许用一套并行配置保存、另一套并行配置加载。这正是 TorchTitan 检查点体系的基础见 checkpoint.md。全量存取需要与第三方格式如 HuggingFace互操作时通过 DCP 的StateDictOptions(full_state_dictTrue, ...)或DTensor.full_tensor()显式 all-gather 出全量张量建议在 rank 0 上配合 CPU offload 以避免峰值显存详见 pytorch_fsdp2_tutorial.md。加载model.load_state_dict(..., assignTrue)配合distribute_tensor(full_tensor, mesh, placements)可实现全量→分片的直接写入若无需自定义处理官方推荐直接使用set_model_state_dict/get_model_state_dict这类 DCP 辅助函数。梯度裁剪由于不再有FlatParameterFSDP2 直接使用标准的nn.utils.clip_grad_norm_()与单机代码一致这也是简化 API的一个体现。混合精度Mixed PrecisionMixedPrecisionPolicy允许将参数存储 dtype、梯度归约 dtype与前向输出 dtype 解耦from torch.distributed._composable.fsdp import MixedPrecisionPolicy mp_policy MixedPrecisionPolicy( param_dtypetorch.bfloat16, # 前向/反向时未分片参数的 dtype reduce_dtypetorch.float32, # 梯度归约的 dtype output_dtypetorch.bfloat16, # 前向输出的 dtype cast_forward_inputsTrue, # 是否将前向输入 cast 到 param_dtype ) fully_shard(model, mp_policymp_policy)四个字段的语义依据 pytorch_fully_shard_api.mdparam_dtype前向/反向计算时未分片参数使用的 dtype。bfloat16是常见选择可显著降低计算与通信量。reduce_dtype梯度 all-reduce / reduce-scatter 时使用的 dtype。设为float32可提高跨节点梯度归约的数值精度。output_dtype前向输出张量的 dtype。cast_forward_inputs为True时前向输入会自动 cast 到param_dtype保证计算精度一致。此外OffloadPolicy系列OffloadPolicy/CPUOffloadPolicy可配置参数、梯度归约、前向输出CPU 卸载时还有优化器状态的设备对应 FSDP1 的cpu_offload参数迁移映射见 pytorch_fsdp2_tutorial.md。HSDP混合分片数据并行当模型规模跨越多个节点时纯 1D 分片会让跨节点 all-gather 成为瓶颈。HSDPHybrid Sharded Data Parallel通过 2D mesh 实现跨组复制 组内分片from torch.distributed.device_mesh import init_device_mesh # 在 4 组之间复制每组内 8 张 GPU 分片 mesh init_device_mesh(cuda, (4, 8), mesh_dim_names(replicate, shard)) fully_shard(model, meshmesh)第一维replicate复制维度对应跨节点/跨机柜的冗余副本第二维shard分片维度对应节点内部的显存分片效果通信从跨 32 卡的全局 all-gather降级为8 卡内 all-gather大幅削减跨节点带宽压力同时保留分片带来的显存节省。在 TorchTitan 的 TOML 配置中HSDP 通过data_parallel_replicate_degree 1启用见下节这也是 70B/405B 级模型多节点训练如 SKILL.md 中 256 GPU、512 GPU 工作流的默认路径。在 TorchTitan 中配置 FSDP/HSDPTorchTitan 把 FSDP 相关配置收敛在[parallelism]段中只需两个键即可在纯 FSDP 与 HSDP 之间切换[parallelism] # FSDP 分片维度-1 自动使用所有可用 GPU data_parallel_shard_degree -1 # HSDP 复制维度1 纯 FSDP1 HSDP data_parallel_replicate_degree 1参数解读data_parallel_shard_degree -1自动把全部可用 GPU 纳入 FSDP 分片是单节点训练的最简配置。data_parallel_replicate_degree 1复制度为 1 即纯 FSDP设为 4 时等价于上面 HSDP 示例中的(4, 8)2D mesh4 组复制 × 每组 8 卡分片。这两个维度与tensor_parallel_degree、pipeline_parallel_degree、context_parallel_degree共同构成 4D 并行网格。例如 70B 模型在 256 GPU 上的典型配置SKILL.md[parallelism] data_parallel_shard_degree 32 # FSDP 跨 32 个 rank 分片 tensor_parallel_degree 8 # 节点内张量并行 pipeline_parallel_degree 1 # 70B 暂不启用 PP context_parallel_degree 1 # 长序列时可增大启动方式配置写好后通过环境变量CONFIG_FILE传入./run_train.sh会自动读取或显式用torchrun启动CONFIG_FILE./torchtitan/models/llama3/train_configs/llama3_8b.toml ./run_train.sh # 等价于 torchrun --nproc_per_node8 -m torchtitan.train --job.config_file ./llama3_8b.toml需要说明的是运行上述命令需要已安装torchtitanpip install torchtitan依赖 PyTorch ≥ 2.6见 SKILL.md 的 frontmatter 元数据并完成 Llama tokenizer 下载。与检查点体系的衔接FSDP2 的分片 state dict 天然对接 TorchTitan 基于 DCP 的检查点机制checkpoint.md[checkpoint] enable true folder checkpoint interval 500DCP 保存的是分片检查点可在改变并行配置后直接重分片加载这正是 FSDP2 无通信分片 state dict 的工程化收益需要与 HuggingFace 互操作时可用last_save_in_hf true在训练中直接导出或离线运行scripts/checkpoint_conversion/convert_to_hf.py变更并行配置后若加载失败可用torch.distributed.checkpoint.format_utils将分片检查点合并为单文件排查。FSDP1 中被移除的参数由于去掉了FlatParameter并重写了内存管理FSDP1 的以下参数在 FSDP2 中不再需要FSDP1 参数FSDP2 替代方案auto_wrap_policy直接对模块调用fully_shard策略即显式调用本身backward_prefetch始终使用BACKWARD_PRE提前预取策略param_init_fn使用 meta-device 初始化 to_emptyinit_weightsdevice_id自动使用 mesh 对应的设备sync_module_statesDTensor 下不再需要广播语义由 DCPbroadcast_from_rank0承担limit_all_gathers新的内存管理机制天然规避该问题use_orig_params恒为True不再有FlatParameter迁移视角对照 pytorch_fsdp2_tutorial.md 的官方迁移映射还能补上几条sharding_strategy→reshard_after_forward 2D mesh 表达 HYBRIDcpu_offload→offload_policyCPUOffloadPolicyno_sync()→set_requires_gradient_syncsync_module_states的从 rank0 广播初始化职责 → DCP 的broadcast_from_rank0流程。这套精简后的参数面配合 custom-models.md 中并行化在模型外部应用的设计原则parallelize_your_model中按 TP → AC → compile → FSDP 的顺序依次包装使得新模型接入分布式训练的开销降到最低——你只需在parallelize_fn里调用fully_shard即可获得与官方模型同级的 FSDP2 能力。最佳实践速查分片顺序始终自底向上调用fully_shard先 TransformerBlock 层再根模块以获得高效的通信分组与通信计算重叠。优化器构造时机在fully_shard之后、基于 DTensor 参数构造优化器。前向入口用model(input)而非model.forward(input)确保 all-gather 钩子生效。超大模型meta-device 初始化 to_emptyinit_weights三步走避免完整建模型的显存峰值。多节点优先data_parallel_replicate_degree 1走 HSDP压低跨节点通信长序列场景再叠加 TP/CP。检查点利用分片 state dict DCP 的 resharding 能力允许并行配置变更后无缝续训。数值精度跨节点梯度归约对精度敏感时将reduce_dtype设为float32。【免费下载链接】AI-Research-SKILLsComprehensive open-source library of AI research and engineering skills for any AI model. Package the skills and your claude code/codex/gemini agent will be an AI research agent with full horsepower. Maintained by Orchestra Research.项目地址: https://gitcode.com/gh_mirrors/ai/AI-Research-SKILLs创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
