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

资讯详情

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

PyTorch FSDP2 全解:fully_shard 逐参数分片的全分片数据并行实现

PyTorch FSDP2 全解:fully_shard 逐参数分片的全分片数据并行实现 PyTorch FSDP2 全解fully_shard 逐参数分片的全分片数据并行实现【免费下载链接】pytorchTensors and Dynamic neural networks in Python with strong GPU acceleration项目地址: https://gitcode.com/GitHub_Trending/py/pytorchPyTorch FSDP2 以torch.distributed.fsdp.fully_shard为前端入口提供基于逐参数per-parameterDTensor分片的全分片数据并行Fully Sharded Data ParallelismFSDP实现面向高性能 eager 模式训练场景。本文基于官方 API 文档结合本仓库torch/distributed/fsdp/_fully_shard/源码与测试系统讲解fully_shard的用户契约、通信分组与调度原理、与 FSDP1 的差异以及FSDPModule提供的全部运行期控制 API帮助读者完成从 FSDP1 到 FSDP2 的迁移并掌握大模型训练中的显存与通信调优手段。一、什么是 FSDP2 与fully_shardFSDP2 是 PyTorch 提出的新一代 FSDP 实现RFC 见 PyTorch 官方 issue #114299核心设计目标是在 eager 模式下保持高性能的同时通过逐参数分片提升可用性。其顶层接口fully_shard(module)位于torch.distributed.fsdp命名空间源码实现在本仓库的 torch/distributed/fsdp/_fully_shard/_fully_shard.py模块边界参数dataclass集中定义在 torch/distributed/fsdp/_fully_shard/_fsdp_api.py。与 FSDP1FullyShardedDataParallel相比FSDP2 放弃了扁平化参数 拼接 分块flat-parameter sharding的表示方式改为在每个数据并行 worker 上沿 dim-0 对单个参数做torch.chunk(dim0)分片并把分片后的参数表示为DTensor。这一设计带来三个直接收益推理更直观每个 worker 上实际持有哪些数据一目了然无需跨参数思考约束更宽松对冻结参数frozen parameters、不同并行度之间重新分片reshard的处理更简单状态字典更轻可支持无需通信的shardedstate dict——在 FSDP1 中这通常需要 all-gather。仓库中的完整测试集位于 test/distributed/_composable/fsdp例如test_fully_shard_comm.py、test_fully_shard_autograd.py、test_fully_shard_dtensor.py、test_fully_shard_frozen.py等分别覆盖通信、自动求导、DTensor 语义、冻结参数等行为是阅读该特性的最佳参照。二、fully_shard(model)的用户契约文档给出的用户契约分为初始化、前后向与优化器三个阶段下面结合源码逐一展开。1. 初始化原地把参数转换为 DTensor调用fully_shard后model.parameters()会从普通torch.Tensor原地in-place变成DTensor并根据meshDeviceMesh移动到对应设备。参见 torch/distributed/fsdp/_fully_shard/_fsdp_init.py 中_init_param_group与_get_device_from_mesh的实现。若用户不显式传入meshfully_shard会调用_init_default_mesh()见 _fsdp_init.py#L198-L211默认取全局进程组构造init_device_mesh——能取到全局 CUDA mesh 就用 CUDA否则用全局 CPU mesh。其设备类型同时决定了通信所用的设备类型。2. 前向与反向钩子负责 all-gather 与参数形态切换前向/反向之前pre-forward / pre-backward 钩子负责把分片参数 all-gather 出来并将model.parameters()从DTensor还原为普通torch.Tensor前向/反向之后post-forward / post-backward 钩子负责释放非分片unsharded参数——该释放无需任何通信——再把model.parameters()从普通torch.Tensor切回DTensor。上述机制在_fsdp_state.py中实现FSDPState._pre_forward、_post_forward、_pre_backward、_post_backward等构成完整的钩子调度循环。fully_shard通过contract(state_clsFSDPState)装饰器把 state 对象与 module 一一绑定可通过fully_shard.state(module)访问。3. 优化器必须基于 DTensor 参数优化器必须用fully_shard之后的model.parameters()即DTensor初始化optimizer.step()也必须在DTensor参数上执行——这也是为什么 FSDP2 天然持有高精度分片参数、无需为 optimizer step 额外保存一份高精度拷贝见下文混合精度。4. 用model(input)而不是model.forward(input)pre-forward 钩子只会在真正调用 forward 方法时触发。文档明确要求调用model(input)以触发 all-gather若确实需要model.forward(input)直接工作必须显式先model.unshard()或使用register_fsdp_forward_method(model, forward)注册 forward 方法以便挂钩子。register_fsdp_forward_method的通用形态是注册任意自定义方法为 forward 方法它会把该方法包装成先执行 state 的_pre_forward再执行原方法最后执行_post_forward的闭包。若传入的 module 不是FSDPModule则该调用为 no-op见 torch/distributed/fsdp/_fully_shard/_fully_shard.py#L927-L962因此可以放心地同时用于启用/未启用 FSDP 的代码路径。5. 自底向上bottom-up应用fully_shard一次fully_shard调用会把这些参数归为一个通信组该模块module.parameters()中、尚未被更早的子模块调用所归属的参数。因此在 Transformer 中应当先对每一层调用fully_shard最后再对根模型调用对根模型调用时各层已有归属的参数会被排除剩余的如 embedding、输出投影会被归并到同一个 all-gather 组。6.type(model)与FSDPModule的运行时联合fully_shard原地改变type(model)例如原来类型为nn.Linear的 model调用后会变成FSDPLinear。FSDPLinear同时是nn.Linear与FSDPModule的实例既保留nn.Linear的全部方法又额外暴露 FSDP2 专属 API如reshard()、unshard()。其实现方式是在 MRO方法解析顺序最左侧插入 FSDP 类_apply_to_module见 _fsdp_init.py#L404。源码注释给出了 MRO 形态[FSDPOrig, FSDPModule, Orig, ...]并借助_orig_cls_mro_index 2与重写后的__new__在索引容器模块等场景下直接构造原类。注意文档与源码均指出 FSDP 不支持deepcopy_unimplemented_deepcopy会直接断言报错序列化请走 state dict。7. 参数 FQN 保持不变由于fully_shard只注册钩子、并不对模块做包装参数的 Fully Qualified NameFQN不会改变对 model 应用fully_shard前后model.state_dict()的键名完全一致。三、FQN 不变与分组通信分组如何决定通信边界每次fully_shard调用都会创建一个通信组组内包含该模块中尚未归属到任何组的全部参数。组的边界直接决定通信边界前向一组参数在一次collective 中完成 all-gather反向这些参数的梯度在一次collective 中完成 reduce-scatter。与 DDP 不同FSDP2没有bucket_cap_mb参数——通信边界完全由你对哪些模块应用fully_shard决定不存在自动 bucketing。场景一只对根模型调用考虑一个含四个子模块m1~m4、参数数分别为a~d的模型model[ m1[a] - m2[b] - m3[c] - m4[d] ]若只调用fully_shard(model)仅根模块则所有参数在同一组整个前向与反向退化为all-gather(abcd) - forward(m1,m2,m3,m4) - backward(m4,m3,m2,m1) - reduce-scatter(abcd)全部通信变成两个巨大的阻塞式操作与计算完全无重叠——文档明确说这几乎从不应该是你想要的。场景二按子模块拆分若按子模块应用例如依次调用fully_shard(m2)、fully_shard(m3)、fully_shard(model)则m2、m3各成一个组剩余参数a、d组成根组。这样多个较小的通信组可以在不同 CUDA stream 上与计算重叠。显存-通信粒度权衡想控制通信组大小就选择要包裹哪些模块包裹更细粒度的模块 → 组更小、更易重叠类似更小的 DDP bucket包裹更少模块 → 组更大。通信边界是显式的完全由模块结构决定。四、前向与反向的通信调度与重叠FSDP2 把 all-gatherAG与 reduce-scatterRS放到独立 CUDA stream 上执行从而与计算流compute stream重叠。前向重叠天然形成、可进一步用 prefetch 强化每个模块的 pre-forward 钩子都会发起自己的 all-gather并在运行模块前等待其完成。因为 CPU 通常比 GPU 跑得快下一个模块的 all-gather 会在当前模块 forward 仍在计算流上执行时就已经在 AG stream 上被发起time ──────────────────────────────────────────────► compute: [wait] [ fwd(m1) | fwd(m2) | fwd(m3,m4) ] AG stream: [AG(a,d)] [AG(b) | AG(c) ]当fwd(m1)在计算流上运行时CPU 已触发m2的 pre-forward 钩子并在 AG stream 上发起AG(b)。若希望这种重叠更稳健例如 CPU 侧开销让 CPU 领先优势缩小时可调用set_modules_to_forward_prefetch让下一个 all-gather在当前模块的 pre-forward 钩子内部就被更早发出而不是等下一个模块钩子触发。反向重叠零配置即插即用反向中 FSDP2无需任何额外配置就会显式预取下一个模块的 all-gather并把 reduce-scatter 放到独立 CUDA streamtime ──────────────────────────────────────────────► compute: [ bwd(m4,m3) | bwd(m2) | bwd(m1) ] AG stream: [AG(c)] [ AG(b) | AG(a,d) ] RS stream: |[RS(c)] [ RS(b)| RS(a,d) ]bwd(m4,m3)在计算流运行时b为m2所需的 all-gather 已在 AG stream 上被预取bwd(m2)运行时AG(a,d)与RS(c)同时与计算重叠。这种流水线正是对每层先自底向上应用fully_shard、再应用到根这一推荐模式的原因。通过模块结构控制分组大小组的粒度由你决定更细的包裹 → 更小、更易重叠的组更粗的包裹 → 更大的组。没有自动 bucketing分组完全显式且由模块结构确定。五、fully_shard的完整参数说明fully_shard同时支持单个模块与模块列表两种输入核心签名源码 torch/distributed/fsdp/_fully_shard/_fully_shard.py#L98-L108如下fully_shard( module, # nn.Module | list[nn.Module] *, mesh: DeviceMesh | None None, reshard_after_forward: bool | int | None None, shard_placement_fn: Callable[[nn.Parameter], ShardPlacementFnResult] | None None, mp_policy: MixedPrecisionPolicy MixedPrecisionPolicy(), offload_policy: OffloadPolicy OffloadPolicy(), ignored_params: set[nn.Parameter] | None None, dp_mesh_dims: DataParallelMeshDims | None None, ) - FSDPModule | list[FSDPModule]各参数要点如下参数含义与取值范围module要分片的模块传列表fully_shard([a, b, ...])时模型前向可能只运行其中一部分模块其余在本轮迭代稍后再被调用如 chunked-loss 训练的fully_shard([norm, head])主前向只跑 normhead 逐 chunk 被调用。mesh数据并行 mesh同时决定分片方式与设备。1D mesh参数沿该 mesh 做全分片FSDP(Shard(0),)placement2D mesh第 1 维分片、第 0 维复制HSDP(Replicate(), Shard(0))placement。mesh 的 device type 决定通信所用设备类型。不传时用默认全局 CUDA/CPU mesh。reshard_after_forward控制 forward 之后的参数行为权衡显存与通信。Trueforward 后立即 reshard反向时重新 all-gatherFalseforward 后保留非分片参数、省掉反向中的一次 all-gather根模块通常设False因为反向开始时根模块几乎立刻就需要参数None默认非根模块为True根模块为Falseintforward 后 reshard 到该 world size须为 mesh 分片维大小的非平凡约数典型选择是节点内大小如torch.cuda.device_count()可让反向 all-gather 在更小 world size 上进行代价是显存高于True。forward 与 backward 之间如需修改参数注册在模块上的参数必须是分片参数——False/int时可用reshard()手动完成。shard_placement_fn逐参数覆写分片 placement 与/或 mesh。返回None用默认Shard(0)返回Shard可指定分片维度返回ShardPlacementResult可同时指定分片与自定义FSDPMeshInfo让不同参数在不同进程组上分片如 MoE 中专家参数与常规参数使用不同 mesh。注意在非零维分片时目前要求均匀分片该维大小须能被 FSDP shard mesh 大小整除。mp_policy混合精度策略见下文MixedPrecisionPolicy。offload_policy卸载策略见下文OffloadPolicy/CPUOffloadPolicy。ignored_params一组被 FSDP 忽略的参数不参与分片、初始化时不搬设备、反向不 reduce 其梯度。dp_mesh_dims提供时mesh被当作完整 SPMD mesh参数应已是该 mesh 上的 DTensor所有 DP 维Replicate()shard字段命名要分片的维多维会被拍平replicate字段命名 HSDP 复制维多维会被拍平。列表式分片的注意点列表分组chunked-loss 场景下源码文档明确了两点 caveat每次独立的逐 chunk 调用都会注册自己的 post_backward autograd 节点因此 N 次 chunk 调用会产生该组N 次 reduce-scattermp_policy.cast_forward_inputs与mp_policy.output_dtype均按组内每个模块分别生效——每次调用含逐 chunk 的独立调用都会把输入 cast 到param_dtype、输出 cast 到output_dtype。异常恢复reset_iter_state文档与源码都强调若forward()/backward()抛异常FSDP 每轮迭代的状态迭代 forward-root 标记、分组模块运行 tracker、在途 collective 状态、各组训练状态会处于未定义状态。要恢复并运行下一轮需在根FSDP 模块上调用FSDPModule.reset_iter_state()失败轮次的梯度会被丢弃包括no_sync/HSDP partial-reduce 累积状态。做梯度累积时应将这段 micro-batch 序列视为失效并重新开始。六、FSDP2 与 FSDP1 的核心差异以 torch/distributed/fsdp/fully_sharded_data_parallel.py 为代表的 FSDP1 与 FSDP2 的差异主要体现在四方面分片表示FSDP2 用基于DTensor的 dim-0 逐参数分片分片表示更简单同时保持相近的吞吐性能。具体来说FSDP2 沿 dim-0 用torch.chunk(dim0)切分每个参数FSDP1 则把一组张量 flatten、concat 后一起切分导致每个 worker 上到底有什么数据、如何 reshard 到其他并行都难以推理。逐参数分片体验更直观、对冻结参数约束更松还能支持免通信的分片state dictFSDP1 中通常需要 all-gather。内存管理FSDP2 以不同方式处理多流使用避免了torch.Tensor.record_stream显存使用确定、可预期也无需像 FSDP1limit_all_gathersTrue那样阻塞 CPU。调度可定制性FSDP2 暴露了手动控制 prefetch 与 collective 调度的 API即下文FSDPModule上的一系列方法让高级用户可以精细定制。API 面简化FSDP2 不直接支持 full state dict。用户可自行用DTensorAPI如DTensor.full_tensor()把含DTensor的分片 state dict 重分片为 full state dict或使用 PyTorch Distributed Checkpoint 这类更高层 API 的分布式 state dict 接口。此外部分历史参数被移除。七、MixedPrecisionPolicy模块级混合精度MixedPrecisionPolicy定义见 _fsdp_api.py#L13-L54与 autocast 不同它在模块级而非算子级应用混合精度为反向保存的是低精度激活高精度→低精度的 cast 只在模块边界发生一次。FSDP 非常适合模块级混合精度因为分片的高精度参数本就常驻内存——不需要为 optimizer step 额外保留一份高精度参数拷贝。字段默认值说明param_dtypeNone指定非分片参数的 dtype即前向/反向计算与参数 all-gather 所用 dtypeNone时非分片参数保持原始 dtype。optimizer step 始终使用原始 dtype 的分片参数。reduce_dtypeNone指定梯度归约reduce-scatter / all-reduce的 dtype。若为None但param_dtype非空则归约使用计算 dtype。可借此低精度计算 全精度梯度归约若同时通过set_requires_gradient_sync关闭梯度归约FSDP 会用reduce_dtype累积梯度。output_dtypeNone浮点前向输出的 cast dtype可用于不同模块不同混合精度策略的场景。cast_forward_inputsTrue是否把 forward 的浮点输入 cast 到param_dtype。对列表分组fully_shard([a, b, ...])cast 按模块逐个生效各模块 forward 之前。八、卸载策略OffloadPolicy与CPUOffloadPolicyOffloadPolicy仅作为不卸载的基类是offload_policy参数的默认值。CPUOffloadPolicy把参数、梯度与优化器状态卸载到 CPU。分片参数在 all-gather 前先 host→device 拷贝all-gather 出的参数按reshard_after_forward释放分片梯度在反向中 device→host 拷贝optimizer step 在 CPU 上用 CPU 优化器状态执行。其唯一字段为pin_memory默认True是否固定分片参数/梯度的内存。固定内存可让 H2D/D2H 拷贝更高效且能与计算重叠但该部分固定内存无法被其他进程使用CPU 内存不足时应设为False。九、FSDPModuleFSDP2 的运行期控制 APIfully_shard返回并在原地改变得到的FSDPModule暴露了一整套手动调度与训练控制方法。源码 torch/distributed/fsdp/_fully_shard/_fully_shard.py#L318-L899 中这些方法均是非递归作用于模块自身reshard/unshard除外或可通过recurse控制是否下推到全部 FSDP 子模块。按用途分类如下参数手动 unshard / reshard方法说明reshard()重分片本模块参数若非分片参数已分配则释放并把分片参数重新注册到模块。非递归。unshard(async_opFalse)分配内存并 all-gather 本模块参数遵循MixedPrecisionPolicy设置param_dtype时按该 dtype all-gather。async_opTrue时返回带wait()的UnshardHandleFalse时函数内等待并返回None。若async_opTrueFSDP 会在模块 pre-forward 中替用户等待挂起的 unshard——只有需要在 pre-forward 之前等待时才需显式wait()。UnshardHandle是一个可等待 unshard 操作的句柄其唯一公开方法wait()确保当前 stream 可以使用已注册到模块的非分片参数。调度与预取prefetch控制方法说明set_modules_to_forward_prefetch(modules)设置本模块在前向中应显式预取 all-gather 的 FSDP 模块预取在本模块 all-gather copy-out 之后运行。传只含下一个 FSDP 模块的单元素列表可获得与默认重叠一致的行为只是从 CPU 更早发出要更激进的重叠代价是更多 reserved 内存需要传至少两个模块。set_modules_to_backward_prefetch(modules)覆盖默认按反向 post-forward 顺序预取下一个 FSDP 模块的实现单元素列表与默认行为一致长度 ≥ 2 用于更激进的重叠。set_is_last_backward(is_last_backward)设置下一次反向是否为最后一次最后一次反向时 FSDP 会等待挂起的梯度归约并清理反向预取相关的内部数据结构对 micro-batching 有用。set_post_optim_event(event)为根 FSDP 模块设置optimizer step 之后的事件让 all-gather stream 等待它。默认根模块在当前 stream 上等待 AG stream以确保 optimizer step 完成后再 all-gather这可能在 optimizer step 后存在无关计算时引入假依赖。调用方需每轮迭代传入新事件。梯度归约与累积控制方法说明set_requires_gradient_sync(requires, *, recurseTrue)设置是否同步梯度可用于实现无通信的梯度累积对应 FSDP1 的no_sync。对 HSDP 同时控制 reduce-scatter 与 all-reduce。set_requires_all_reduce(requires, *, recurseTrue)设置是否 all-reduce 梯度可用于实现 HSDP 下只 reduce-scatter、不 all-reduce的梯度累积。set_reshard_after_forward(bool, recurseTrue)运行期更改reshard_after_forward。例如把 FSDP 根模块的值改为True根模块默认被特殊设为False或在 eval 时设为False、训练时改回True。set_reshard_after_backward(bool, *, recurseTrue)设置反向之后是否 reshard 参数。梯度累积时可用更高显存换更少通信非分片参数下次 forward 无需重新 all-gather。set_gradient_divide_factor(factor)为梯度归约设置自定义除数因子可使用 NCCLPreMulSum在归约前先乘上该因子。set_reduce_scatter_divide_factor为其废弃别名。set_force_sum_reduction_for_comms(enable)是否要求底层 collective 原语只用 sum 类归约哪怕需要额外的 pre/post 缩放步骤。NCCL 目前仅对这类 collective 支持零拷贝传输MTIA 设备恒为隐式开启。若在 FSDP 下使用set_all_reduce_hook调用方需自行保证自定义 all-reduce 也遵循该策略。set_reduce_scatter_unused_params(enable, *, recurseTrue)是否在归约中为未收到梯度的参数补零梯度。用于不同 rank 因条件控制流多模态、MoE 等使用不同参数导致 reduce-scatter 不匹配的场景类似 DDP 的find_unused_parameters。set_all_reduce_hook(hook, *, streamNone)注册自定义 all-reduce 钩子签名hook(reduce_output: Tensor) - None其中reduce_output在纯 FSDP 下是 reduce-scatter 输出、在原生 HSDP 下是 all-reduce 输出。原生 HSDP 下stream不可设置由内部 all-reduce stream 运行钩子。通信实现级定制方法说明set_custom_all_gather(comm)/set_custom_reduce_scatter(comm)覆盖默认 all-gather / reduce-scatter 通信行为。Comm抽象接口AllGather、ReduceScatter均为其子类见 _fsdp_api.py#L57-L128需实现三件事如何分配通信内存可每调用临时 buffer也可为效率复用持久 buffer、在哪里分配如 NCCL mem pool 或常规 caching allocator、通信被调用时做什么。注意二者均不支持多参数组来自shard_placement_fn的逐参数 mesh否则会抛ValueError。set_allocate_memory_from_process_group_for_comm(enable)是否让集体通信收发所用的临时 staging buffer 使用进程组自带的优化分配器若有。例如 NCCL 下可启用经 SHARPNVLink/InfiniBand的零拷贝传输。不能与自定义 all-gather/reduce-scatter 同时使用。set_symm_mem_for_comm(backendNCCL)用对称内存symm_mem后端为 all-gather collective 分配 staging buffer使 NCCL 能走优化实现单节点可能用 Copy Engine All-Gather多节点可能用 Symmetric Kernel All-Gather。启用 Copy Engine All-Gather 需以 zero-CTA 策略创建 NCCL 进程组pg_options中cta_policy NCCL_CTA_POLICY_ZERO或将环境变量NCCL_CTA_POLICY设为2。目前仅支持NCCL后端不能与自定义 comm API 同用。set_separate_reduce_scatter_group(enableTrue, *, recurseTrue)实验性默认 FSDP 在 separate CUDA stream 上跑 all-gather 与 reduce-scatter但走同一个进程组单个 NCCL communicator 同一时刻只处理一个 collective通信上串行。启用后FSDP 会为 shard rank 集创建一个专用进程组dist.new_group(..., use_local_synchronizationTrue)使两类 collective 可在网络允许时并发推进。该调用对每个 shard rank 集是集合性的需在用到该 FSDP mesh 的各 rank 上一致调用。set_reduce_scatter_max_input_buffers(max_input_buffers, *, recurseTrue)实验性设置同一时刻在途的梯度 reduce-scatter 输入 buffer 数量上限copy-inchunk_catbuffer 的 cap-K。默认只保留 1 个在途 buffer因此下一个 copy-in 必须等上一次 reduce-scatter 释放该 buffer——当 reduce-scatter 暴露通信慢于被隐藏的反向计算时这个回收等待会卡住计算流提高上限可让下一次 copy-in 写全新 buffer 从而消除停顿代价是更高的峰值显存。取值必须为 1的 intbool 会被拒绝避免True被误当成 1。高级场景与训练状态方法说明set_unshard_in_backward(unshard_in_backward)设置本 FSDP 模块的参数是否需要反向 unshard。用于明确知道该参数组在反向计算中不需要的专家场景如 embedding。reset_iter_state()前向/反向中途异常后重置 FSDP 每轮迭代状态见上文异常恢复一节。等待在途 all-gather/reduce-scatter 事件、重分片所有参数组、清理迭代 tracker在途梯度归约被丢弃。必须在根 FSDP 模块上调用对非根模块调用抛RuntimeError。_set_unshard_async_op(async_op)设置 pre-forward/pre-backward unshard 是否使用async_opTrue开启后 all-gather 分配发生在默认 stream可避免跨 stream 显存碎片但前向必须使用显式 prefetch如unshard才能保留重叠且 dtype cast、copy-in 等 pre-all-gather 操作不再与计算重叠。十、其他模块级辅助 APIshare_comm_ctx(modules)让多个FSDPModule共享 CUDA streamall-gather、reduce-scatter、all-reduce 的通信上下文。典型场景是流水线并行PP每个模型 chunk 是一个 FSDP root共享 stream 可避免跨 stream 通信造成的显存碎片。示例share_comm_ctx([fsdp_model_1, fsdp_model_2, ...])传入非FSDPModule会抛ValueError。register_fsdp_forward_method(module, method_name)见上文用户契约一节若 module 不是FSDPModule则为 no-op。get_cls_to_fsdp_cls()返回类名到 FSDP 类的映射字典cls_to_fsdp_cls可用于了解当前进程内有哪些类已被 FSDP 化。disable_fsdp_module_new_init()上下文管理器临时关闭 FSDP 化模块的__init__配合FSDPModule.__new__的构造逻辑使用。十一、DataParallelMeshDimsSPMD mesh 下的 DP 维度声明DataParallelMeshDims_fsdp_api.py#L131-L171用于当参数本身已经是某个完整 SPMDDeviceMesh上的 DTensor 时指定fully_shard应对 mesh 的哪些维度做数据并行。shardFSDP 进行参数分片的 mesh 维名称。若为名称元组这些维会被拍平成一个分片维。replicate用于 HSDP / DDP 复制的 mesh 维名称。若为名称元组这些维会被拍平成一个复制维。shard与replicate至少必须设置其一否则__post_init__抛ValueError。此外使用 SPMD meshdp_mesh_dims时目前不支持把reshard_after_forward设为 int源码会抛NotImplementedError。十二、快速上手一个可复现的集成骨架下面给出把上述概念串起来的典型集成骨架以 1D mesh 的纯 FSDP 为例所有 API 均为本仓库torch.distributed.fsdp现有导出import torch import torch.nn as nn from torch.distributed.device_mesh import init_device_mesh from torch.distributed.fsdp import fully_shard, MixedPrecisionPolicy from torch.distributed.tensor import DTensor def build_model(): 自底向上先逐层 fully_shard最后再 fully_shard 根模型。 layer1 nn.Linear(4096, 4096) layer2 nn.Linear(4096, 4096) fully_shard(layer1) # 每个 layer 一个通信组 fully_shard(layer2) root nn.Sequential(layer1, layer2) fully_shard(root) # 根组收纳剩余参数若有 return root model build_model() # mesh 为 None 时会 fallback 到默认全局 CUDA/CPU mesh # 这里显式传入 1D 设备 mesh 以明确语义 fully_shard(model, meshinit_device_mesh(cuda, (world_size,)), mp_policyMixedPrecisionPolicy(param_dtypetorch.bfloat16), reshard_after_forwardFalse) # 根模块保留非分片参数以省去反向 all-gather # 优化器必须基于 DTensor 参数Fully Sharded 的 model.parameters() optimizer torch.optim.AdamW(model.parameters(), lr1e-4) # 训练循环始终使用 model(input) 触发 pre-forward 钩子 for x, y in dataloader: optimizer.zero_grad() out model(x) # 触发逐组 all-gather loss loss_fn(out, y) loss.backward() # 逐组 reduce-scatter 梯度自动与计算重叠 optimizer.step() # 在 DTensor 分片参数上执行 # 手动调度示例前向中显式预取下一组 / 显式 unshard fully_shard.state(model)._get_fsdp_state() # 通过 state 访问内部状态 # model.unshard() # 显式 all-gather配合 model.forward(input) 使用 # model.reshard() # 显式重分片 # model.set_modules_to_forward_prefetch([next_fsdp_module])需要留意几个实践要点触发条件必须调用model(input)而不是model.forward(input)或先unshard()/register_fsdp_forward_method否则参数不会被 all-gatherFQN 一致性应用fully_shard前后state_dict()的键名一致便于无缝接入已有的 checkpoint 逻辑分片 state dict 可通过DTensor.full_tensor()或 Distributed Checkpoint 还原为 full state dict粒度不要只在最顶层根模块调用fully_shard否则通信会退化为两次巨大的阻塞 collective先逐层自底向上调用以获得计算/通信重叠深拷贝FSDP 不支持deepcopy序列化统一走 state dict。十三、进一步阅读官方教程 Getting Started with FSDP2 提供更系统的上手演示其中包含 FSDP1→FSDP2 的迁移指南完整实现与类型声明见本仓库前端 APIfully_shard、FSDPModule、UnshardHandle、register_fsdp_forward_method、share_comm_ctx在 torch/distributed/fsdp/_fully_shard/_fully_shard.py策略 dataclass 与通信原语接口MixedPrecisionPolicy、OffloadPolicy、CPUOffloadPolicy、DataParallelMeshDims、Comm/AllGather/ReduceScatter在 torch/distributed/fsdp/_fully_shard/_fsdp_api.py状态机与钩子调度、初始化与 mesh 解析分别位于 torch/distributed/fsdp/_fully_shard/_fsdp_state.py 与 torch/distributed/fsdp/_fully_shard/_fsdp_init.py行为级验证测试集中在 test/distributed/_composable/fsdp如test_fully_shard_comm.py、test_fully_shard_autograd.py、test_fully_shard_dtensor.py、test_fully_shard_frozen.py。【免费下载链接】pytorchTensors and Dynamic neural networks in Python with strong GPU acceleration项目地址: https://gitcode.com/GitHub_Trending/py/pytorch创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表