distributed
checkpoint
optimizer
trainer
全分片数据并行(Fully Sharded Data Parallel,FSDP)当前要求参数在切分轴上能够被 shard_world_size 均匀切分。模型参数的 dim-0 长度无法整除时,FSDP 不能生成正确的实际分片、 通信存储和 DTensor placements,训练因此无法启动。
shard_world_size
FSDP 的全收集(AllGather)和归约散射(ReduceScatter)要求每个进程(rank)提供等长输入; 分布式检查点(Distributed Checkpoint,DCP)则需要根据 sharded_param.placements 识别每个 rank 真正持有的参数范围。通信 padding 和实际参数分片必须分别表达。
sharded_param.placements
fully_shard()
本 RFC 要解决的问题:FSDP 无法训练 dim-0 长度不能被 shard_world_size 整除的参数。 完成后的成功标准:FSDP 生成正确的实际分片和 placements,通信 padding 正确,端到端精度与非分布式基线一致。
1. 支持 FSDP 和混合分片数据并行(Hybrid Sharded Data Parallel,HSDP)对参数执行 dim-0 非均匀切分。 2. 使用 DTensor 的 RaggedShard 或 RaggedStridedShard 表达 sharded_param.placements,不再用 Shard(0) 表达非均匀分片。 3. 区分 sharded_size 表示的实际分片和 padded_sharded_param_size 表示的通信补齐形状。 4. 支持 FSDP/HSDP 与基于 DTensor 的张量并行(Tensor Parallel,TP)、序列并行(Sequence Parallel,SP)组合训练。 5. 保持 dim-0 均匀切分的接口和执行路径不变。
1. 本期不支持 dim-0 以外轴的非均匀切分;shard_dim != 0 时继续要求参数能够被 shard_world_size 整除。 2. 本期不新增 fully_shard() 对外参数,非均匀切分由参数形状自动触发。 3. 本期不修改 DTensor 或分布式检查点的 RaggedShard 实现,这些能力由 DTensor 非均匀切分需求提供。 4. 本需求是功能完善,不以提升性能或降低显存为目标。
mesh = init_device_mesh( device_type="npu", mesh_shape=(4,), mesh_dim_names=("dp",), ) model = fully_shard(model, mesh=mesh)
module
Module
mesh
DeviceMesh
None
shard_placement_fn
Callable
NotImplementedError
mesh = init_device_mesh("npu", (4,), mesh_dim_names=("dp",)) # model 包含 shape 为 (7, 8) 的参数 weight model = fully_shard(model, mesh=mesh)
参数 weight 的 dim-0 长度为 7,不能被 4 均匀切分。用户不需要增加配置,FSDP 自动进入 dim-0 非均匀切分路径。
weight
为什么这样设计:非均匀切分是参数形状与 shard_world_size 的结果,不需要用户开关。 和已有接口是否一致:一致,继续使用 fully_shard() 和 shard_placement_fn。
flowchart TD A["读取 param_data 和 shard_dim"] --> B{"shard_dim 是否为 0"} B -- "否" --> C["保持现有均匀切分校验和显式通信转换"] B -- "是" --> D["计算 actual_shard_offset 和 actual_shard_length"] D --> E["生成 RaggedShard 或 RaggedStridedShard placements"] E --> F["构造 sharded_param 实际分片"] F --> G["按 padded_sharded_param_size 构造通信存储"] G --> H["AllGather 恢复逻辑完整参数"] H --> I["forward 和 backward"] I --> J["ReduceScatter 输入尾部补 0"] J --> K["按 sharded_size 回填实际梯度"]
flowchart LR subgraph FSDP["FSDP 参数管理"] Param["sharded_param 实际分片"] Padded["_sharded_param_data 通信存储"] Grad["实际分片梯度"] end subgraph DTensor["DTensor 非均匀切分能力"] Placement["RaggedShard / RaggedStridedShard"] Logical["DTensor.shape"] end subgraph Communication["集合通信"] AG["AllGather"] RS["ReduceScatter"] end subgraph Checkpoint["分布式检查点"] DCP["按 placements 保存和加载实际分片"] end Param --> Placement Placement --> Logical Param --> DCP Placement --> DCP Padded --> AG RS --> Grad
sequenceDiagram participant S as HSDPState participant P as HSDPParamV2 participant D as DTensor participant C as Collective participant O as Optimizer S->>P: 初始化参数 P->>P: 计算 sharded_size 和 padded_sharded_param_size P->>D: 用 ragged placements 构造 sharded_param P->>C: 使用 _sharded_param_data 发起 AllGather C-->>P: 返回补齐后的完整 buffer P->>S: 暴露逻辑完整 unsharded_param S->>P: backward 后归约梯度 P->>C: 发起补 0 后的 ReduceScatter C-->>P: 返回补齐后的分片区域 P->>O: 仅提交 sharded_size 对应的实际梯度
PT/MS 采用一致的固定大小分块语义计算实际范围:
dim_shard_size = ( param_data.shape[0] + self.shard_world_size - 1 ) // self.shard_world_size actual_shard_offset = min( self.shard_rank * dim_shard_size, param_data.shape[0], ) actual_shard_length = min( dim_shard_size, param_data.shape[0] - actual_shard_offset, )
sharded_param 是 param_data 的实际本地分片,允许 actual_shard_length 为 0。 self.sharded_size 记录实际形状;self.padded_sharded_param_size 将 dim-0 设置为 dim_shard_size,记录集合通信要求的统一形状。
sharded_param
param_data
actual_shard_length
self.sharded_size
self.padded_sharded_param_size
dim_shard_size
Shard
StridedShard
local_units
RaggedShard
RaggedStridedShard
split_factor
_spmd_placements
self.sharded_param.placements
DCP 根据 ragged placements 识别各 rank 的实际分片范围。 本 RFC 不在 FSDP 内重复实现分片 offset 推导。
self._sharded_param_data
sharded_param.view(-1)
padded_sharded_param
self.sharded_param
reset_sharded_param()
all_gather_inputs
reduce_scatter_grad()
shard_dim != 0
reduce_scatter_copy_in()
padded_sharded_param_size.numel()
model
platform/{torch,mindspore}/fully_shard/param.py
fully_shard/param_group.py
padded_sharded_param_size
sharded_size
Shard(0)
完整能力需要:DTensor 提供 RaggedShard、RaggedStridedShard、logical global shape 及 DCP 所需的实际分片表达。 本期最小可交付能力:PT/MS 的 FSDP/HSDP 均可训练 dim-0 非均匀参数,且 sharded_param.placements 正确。
单元测试(Unit Test,UT)验证 placements 和 buffer 逻辑;系统测试(System Test,ST)按 Level0 和 Level1 覆盖基础训练闭环及多组件组合。
UT 需要至少覆盖以下边界:
sharded_size[0]
1. 基本信息
distributed/checkpoint/optimizer/trainer2. 背景
全分片数据并行(Fully Sharded Data Parallel,FSDP)当前要求参数在切分轴上能够被
shard_world_size均匀切分。模型参数的 dim-0 长度无法整除时,FSDP 不能生成正确的实际分片、通信存储和 DTensor placements,训练因此无法启动。
FSDP 的全收集(AllGather)和归约散射(ReduceScatter)要求每个进程(rank)提供等长输入;
分布式检查点(Distributed Checkpoint,DCP)则需要根据
sharded_param.placements识别每个 rank真正持有的参数范围。通信 padding 和实际参数分片必须分别表达。
fully_shard()接口完成分布式训练3. 目标和非目标
3.1 目标
3.2 非目标
5. 对外接口
5.1 接口定义
mesh = init_device_mesh( device_type="npu", mesh_shape=(4,), mesh_dim_names=("dp",), ) model = fully_shard(model, mesh=mesh)moduleModulemeshDeviceMeshNoneshard_placement_fnCallable或NoneNoneNone表示 dim-0NotImplementedError5.2 使用示例
mesh = init_device_mesh("npu", (4,), mesh_dim_names=("dp",)) # model 包含 shape 为 (7, 8) 的参数 weight model = fully_shard(model, mesh=mesh)参数
weight的 dim-0 长度为 7,不能被 4 均匀切分。用户不需要增加配置,FSDP 自动进入 dim-0非均匀切分路径。
5.3 接口说明
6. 方案设计
6.1 总体流程
flowchart TD A["读取 param_data 和 shard_dim"] --> B{"shard_dim 是否为 0"} B -- "否" --> C["保持现有均匀切分校验和显式通信转换"] B -- "是" --> D["计算 actual_shard_offset 和 actual_shard_length"] D --> E["生成 RaggedShard 或 RaggedStridedShard placements"] E --> F["构造 sharded_param 实际分片"] F --> G["按 padded_sharded_param_size 构造通信存储"] G --> H["AllGather 恢复逻辑完整参数"] H --> I["forward 和 backward"] I --> J["ReduceScatter 输入尾部补 0"] J --> K["按 sharded_size 回填实际梯度"]6.2 架构参考
flowchart LR subgraph FSDP["FSDP 参数管理"] Param["sharded_param 实际分片"] Padded["_sharded_param_data 通信存储"] Grad["实际分片梯度"] end subgraph DTensor["DTensor 非均匀切分能力"] Placement["RaggedShard / RaggedStridedShard"] Logical["DTensor.shape"] end subgraph Communication["集合通信"] AG["AllGather"] RS["ReduceScatter"] end subgraph Checkpoint["分布式检查点"] DCP["按 placements 保存和加载实际分片"] end Param --> Placement Placement --> Logical Param --> DCP Placement --> DCP Padded --> AG RS --> Grad6.3 时序参考
sequenceDiagram participant S as HSDPState participant P as HSDPParamV2 participant D as DTensor participant C as Collective participant O as Optimizer S->>P: 初始化参数 P->>P: 计算 sharded_size 和 padded_sharded_param_size P->>D: 用 ragged placements 构造 sharded_param P->>C: 使用 _sharded_param_data 发起 AllGather C-->>P: 返回补齐后的完整 buffer P->>S: 暴露逻辑完整 unsharded_param S->>P: backward 后归约梯度 P->>C: 发起补 0 后的 ReduceScatter C-->>P: 返回补齐后的分片区域 P->>O: 仅提交 sharded_size 对应的实际梯度6.4 关键逻辑
PT/MS 采用一致的固定大小分块语义计算实际范围:
dim_shard_size = ( param_data.shape[0] + self.shard_world_size - 1 ) // self.shard_world_size actual_shard_offset = min( self.shard_rank * dim_shard_size, param_data.shape[0], ) actual_shard_length = min( dim_shard_size, param_data.shape[0] - actual_shard_offset, )sharded_param是param_data的实际本地分片,允许actual_shard_length为 0。self.sharded_size记录实际形状;self.padded_sharded_param_size将 dim-0 设置为dim_shard_size,记录集合通信要求的统一形状。placements 生成
Shard或StridedShard。actual_shard_length生成local_units。RaggedShard;原参数已有 TP/SP placements 且需要保留切分顺序时,使用RaggedStridedShard,其split_factor沿用现有_spmd_placements的计算规则。self.sharded_param.placements只描述实际分片。padding 不进入 placements,也不进入 DTensor 的逻辑全局 shape。
DCP 根据 ragged placements 识别各 rank 的实际分片范围。
本 RFC 不在 FSDP 内重复实现分片 offset 推导。
参数和通信存储
self._sharded_param_data直接指向sharded_param.view(-1)。padded_sharded_param,把实际分片复制到前缀;self._sharded_param_data指向完整通信存储,self.sharded_param只引用实际范围。reset_sharded_param()在延迟初始化、参数 dtype 转换和加载参数后重建上述关系,不能把 padding暴露给 optimizer 或 DCP。
AllGather 和 ReduceScatter
all_gather_inputs继续读取self._sharded_param_data,保证每个 rank 输入元素数一致。param_data形状,尾部 padding 不进入模块计算。reduce_scatter_grad()对 dim-0 非均匀梯度创建 padded input,尾部必须补 0;shard_dim != 0保持现有显式 chunk/cat 处理。
reduce_scatter_copy_in()直接把每个梯度写入最终融合输入区域。dim-0 由后端分块算子完成分块和补 0,输入 offset 按
padded_sharded_param_size.numel()前进。self.sharded_size,optimizer 不持有 padding 对应的梯度。6.5 代码改动点
modelplatform/{torch,mindspore}/fully_shard/param.py生成实际分片、ragged placements 和补齐通信存储distributedfully_shard/param_group.py按padded_sharded_param_size组织通信,按sharded_size回填梯度checkpointsharded_param.placements;DCP 逻辑由 DTensor 非均匀切分需求提供optimizerself.sharded_param,只能看到实际分片和实际梯度trainerreset_sharded_param()重建 padding storage6.6 方案取舍
Shard语义Shard(0),在 Layout 中额外保存 logical shapeShard(0)不能表达非均匀分片,DCP 无法仅根据 placements 得到实际范围sharded_param7. 组件依赖
Shard(0);本需求不能完整交付8. 约束与兼容性
Shard(0)兼容旁路9. 验证设计
9.1 用例分层
单元测试(Unit Test,UT)验证 placements 和 buffer 逻辑;系统测试(System Test,ST)按 Level0 和
Level1 覆盖基础训练闭环及多组件组合。
shard_world_size整除的参数9.2 交互验证(举例)
sharded_size、padded_sharded_param_size正确;端到端精度与非分布式基线一致reset_sharded_param()正确重建实际分片视图和补齐通信存储sharded_param.placements能表达实际分片;保存和加载由 DTensor 非均匀切分需求验收UT 需要至少覆盖以下边界:
shard_world_size。shard_world_size,后部 rank 的sharded_size[0]为 0。padded_sharded_param_size对齐,输出只暴露逻辑参数范围。9.3 性能 / 显存验证
10. CheckList