PyTorch FSDP2 全解:fully_shard 逐参数分片的全分片数据并行实现
【免费下载链接】pytorchTensors and Dynamic neural networks in Python with strong GPU acceleration项目地址: https://gitcode.com/GitHub_Trending/py/pytorch
PyTorch FSDP2 以torch.distributed.fsdp.fully_shard为前端入口,提供基于逐参数(per-parameter)DTensor分片的全分片数据并行(Fully Sharded Data Parallelism,FSDP)实现,面向高性能 eager 模式训练场景。本文基于官方 API 文档,结合本仓库torch/distributed/fsdp/_fully_shard/源码与测试,系统讲解fully_shard的用户契约、通信分组与调度原理、与 FSDP1 的差异,以及FSDPModule提供的全部运行期控制 API,帮助读者完成从 FSDP1 到 FSDP2 的迁移并掌握大模型训练中的显存与通信调优手段。
一、什么是 FSDP2 与fully_shard
FSDP2 是 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。
与 FSDP1(FullyShardedDataParallel)相比,FSDP2 放弃了"扁平化参数 + 拼接 + 分块"(flat-parameter sharding)的表示方式,改为在每个数据并行 worker 上沿 dim-0 对单个参数做torch.chunk(dim=0)分片,并把分片后的参数表示为DTensor。这一设计带来三个直接收益:
- 推理更直观:每个 worker 上实际持有哪些数据一目了然,无需跨参数思考;
- 约束更宽松:对冻结参数(frozen parameters)、不同并行度之间重新分片(reshard)的处理更简单;
- 状态字典更轻:可支持无需通信的(sharded)state 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,并根据mesh(DeviceMesh)移动到对应设备。参见 torch/distributed/fsdp/_fully_shard/_fsdp_init.py 中_init_param_group与_get_device_from_mesh的实现。
若用户不显式传入mesh,fully_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_cls=FSDPState)装饰器把 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 形态:[FSDP<Orig>, FSDPModule, Orig, ...],并借助_orig_cls_mro_index = 2与重写后的__new__在"索引容器模块"等场景下直接构造原类。
注意:文档与源码均指出 FSDP 不支持
deepcopy(_unimplemented_deepcopy会直接断言报错),序列化请走 state dict。
7. 参数 FQN 保持不变
由于fully_shard只注册钩子、并不对模块做包装,参数的 Fully Qualified Name(FQN)不会改变:对 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(a+b+c+d) -> forward(m1,m2,m3,m4) -> backward(m4,m3,m2,m1) -> reduce-scatter(a+b+c+d)全部通信变成两个巨大的阻塞式操作,与计算完全无重叠——文档明确说这"几乎从不应该是你想要的"。
场景二:按子模块拆分
若按子模块应用,例如依次调用fully_shard(m2)、fully_shard(m3)、fully_shard(model),则m2、m3各成一个组,剩余参数a、d组成根组。这样多个较小的通信组可以在不同 CUDA stream 上与计算重叠。
显存-通信粒度权衡
想控制通信组大小,就选择要"包裹"哪些模块:包裹更细粒度的模块 → 组更小、更易重叠(类似更小的 DDP bucket);包裹更少模块 → 组更大。通信边界是显式的,完全由模块结构决定。
四、前向与反向的通信调度与重叠
FSDP2 把 all-gather(AG)与 reduce-scatter(RS)放到独立 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 stream:
time ──────────────────────────────────────────────► 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]):主前向只跑 norm,head 逐 chunk 被调用)。 |
mesh | 数据并行 mesh,同时决定分片方式与设备。1D mesh:参数沿该 mesh 做全分片(FSDP),(Shard(0),)placement;2D mesh:第 1 维分片、第 0 维复制(HSDP),(Replicate(), Shard(0))placement。mesh 的 device type 决定通信所用设备类型。不传时用默认全局 CUDA/CPU mesh。 |
reshard_after_forward | 控制 forward 之后的参数行为,权衡显存与通信。True:forward 后立即 reshard,反向时重新 all-gather;False:forward 后保留非分片参数、省掉反向中的一次 all-gather(根模块通常设False,因为反向开始时根模块几乎立刻就需要参数);None(默认):非根模块为True,根模块为False;int:forward 后 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-scatter;
mp_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(dim=0)切分每个参数;FSDP1 则把一组张量 flatten、concat 后一起切分,导致"每个 worker 上到底有什么数据、如何 reshard 到其他并行"都难以推理。逐参数分片体验更直观、对冻结参数约束更松,还能支持免通信的(分片)state dict(FSDP1 中通常需要 all-gather)。 - 内存管理:FSDP2 以不同方式处理多流使用,避免了
torch.Tensor.record_stream;显存使用确定、可预期,也无需像 FSDP1limit_all_gathers=True那样阻塞 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_dtype | None | 指定非分片参数的 dtype,即前向/反向计算与参数 all-gather 所用 dtype;None时非分片参数保持原始 dtype。optimizer step 始终使用原始 dtype 的分片参数。 |
reduce_dtype | None | 指定梯度归约(reduce-scatter / all-reduce)的 dtype。若为None但param_dtype非空,则归约使用计算 dtype。可借此"低精度计算 + 全精度梯度归约";若同时通过set_requires_gradient_sync关闭梯度归约,FSDP 会用reduce_dtype累积梯度。 |
output_dtype | None | 浮点前向输出的 cast dtype,可用于"不同模块不同混合精度策略"的场景。 |
cast_forward_inputs | True | 是否把 forward 的浮点输入 cast 到param_dtype。对列表分组fully_shard([a, b, ...]),cast 按模块逐个生效(各模块 forward 之前)。 |
八、卸载策略:OffloadPolicy与CPUOffloadPolicy
OffloadPolicy:仅作为"不卸载"的基类,是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。
九、FSDPModule:FSDP2 的运行期控制 API
fully_shard返回(并在原地改变得到的)FSDPModule暴露了一整套手动调度与训练控制方法。源码 torch/distributed/fsdp/_fully_shard/_fully_shard.py#L318-L899 中这些方法均是非递归作用于模块自身(reshard/unshard除外),或可通过recurse控制是否下推到全部 FSDP 子模块。按用途分类如下:
参数手动 unshard / reshard
| 方法 | 说明 |
|---|---|
reshard() | 重分片本模块参数:若非分片参数已分配则释放,并把分片参数重新注册到模块。非递归。 |
unshard(async_op=False) | 分配内存并 all-gather 本模块参数,遵循MixedPrecisionPolicy(设置param_dtype时按该 dtype all-gather)。async_op=True时返回带wait()的UnshardHandle,False时函数内等待并返回None。若async_op=True,FSDP 会在模块 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, *, recurse=True) | 设置是否同步梯度,可用于实现无通信的梯度累积(对应 FSDP1 的no_sync)。对 HSDP 同时控制 reduce-scatter 与 all-reduce。 |
set_requires_all_reduce(requires, *, recurse=True) | 设置是否 all-reduce 梯度,可用于实现 HSDP 下"只 reduce-scatter、不 all-reduce"的梯度累积。 |
set_reshard_after_forward(bool, recurse=True) | 运行期更改reshard_after_forward。例如把 FSDP 根模块的值改为True(根模块默认被特殊设为False),或在 eval 时设为False、训练时改回True。 |
set_reshard_after_backward(bool, *, recurse=True) | 设置反向之后是否 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, *, recurse=True) | 是否在归约中为"未收到梯度的参数"补零梯度。用于不同 rank 因条件控制流(多模态、MoE 等)使用不同参数导致 reduce-scatter 不匹配的场景,类似 DDP 的find_unused_parameters。 |
set_all_reduce_hook(hook, *, stream=None) | 注册自定义 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 下可启用经 SHARP(NVLink/InfiniBand)的零拷贝传输。不能与自定义 all-gather/reduce-scatter 同时使用。 |
set_symm_mem_for_comm(backend="NCCL") | 用对称内存(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(enable=True, *, recurse=True)(实验性) | 默认 FSDP 在 separate CUDA stream 上跑 all-gather 与 reduce-scatter,但走同一个进程组(单个 NCCL communicator 同一时刻只处理一个 collective,通信上串行)。启用后,FSDP 会为 shard rank 集创建一个专用进程组(dist.new_group(..., use_local_synchronization=True)),使两类 collective 可在网络允许时并发推进。该调用对每个 shard rank 集是集合性的,需在用到该 FSDP mesh 的各 rank 上一致调用。 |
set_reduce_scatter_max_input_buffers(max_input_buffers, *, recurse=True)(实验性) | 设置同一时刻在途的梯度 reduce-scatter 输入 buffer 数量上限(copy-inchunk_catbuffer 的 cap-K)。默认只保留 1 个在途 buffer,因此下一个 copy-in 必须等上一次 reduce-scatter 释放该 buffer——当 reduce-scatter 暴露(通信慢于被隐藏的反向计算)时,这个回收等待会卡住计算流;提高上限可让下一次 copy-in 写全新 buffer 从而消除停顿,代价是更高的峰值显存。取值必须为>= 1的 int(bool 会被拒绝,避免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_op=True:开启后 all-gather 分配发生在默认 stream,可避免跨 stream 显存碎片,但前向必须使用显式 prefetch(如unshard)才能保留重叠,且 dtype cast、copy-in 等 pre-all-gather 操作不再与计算重叠。 |
十、其他模块级辅助 API
share_comm_ctx(modules):让多个FSDPModule共享 CUDA stream(all-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__的构造逻辑使用)。
十一、DataParallelMeshDims:SPMD mesh 下的 DP 维度声明
DataParallelMeshDims(_fsdp_api.py#L131-L171)用于:当参数本身已经是某个完整 SPMDDeviceMesh上的 DTensor 时,指定fully_shard应对 mesh 的哪些维度做数据并行。
shard:FSDP 进行参数分片的 mesh 维名称。若为名称元组,这些维会被拍平成一个分片维。replicate:用于 HSDP / DDP 复制的 mesh 维名称。若为名称元组,这些维会被拍平成一个复制维。
shard与replicate至少必须设置其一(否则__post_init__抛ValueError)。此外,使用 SPMD mesh(dp_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, mesh=init_device_mesh("cuda", (world_size,)), mp_policy=MixedPrecisionPolicy(param_dtype=torch.bfloat16), reshard_after_forward=False) # 根模块保留非分片参数以省去反向 all-gather # 优化器必须基于 DTensor 参数(Fully Sharded 的 model.parameters()) optimizer = torch.optim.AdamW(model.parameters(), lr=1e-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-gather; - FQN 一致性:应用
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 的迁移指南;
- 完整实现与类型声明见本仓库:
- 前端 API:
fully_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)。
- 前端 API:
【免费下载链接】pytorchTensors and Dynamic neural networks in Python with strong GPU acceleration项目地址: https://gitcode.com/GitHub_Trending/py/pytorch
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考