mindspore-ai / mindspore-ai/hyper-parallel
[RFC]: FSDP支持参数非均匀切分
Nobody has claimed this yet.
- Dominant language
- Python
- Stars
- 53
- Forks
- 63
- Avg merge
- 23h 45m
- Merged PRs (30d)
- 63
Description
1. 基本信息
| 项目 | 内容 |
|---|---|
| 作者 | MengXY107 |
| 相关模块 | distributed / checkpoint / optimizer / trainer |
| 相关 issue / PR | DTensor 非均匀切分需求;实现 PR 待创建 |
| 适用后端 | PyTorch(PT)+ MindSpore(MS) |
2. 背景
全分片数据并行(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 和实际参数分片必须分别表达。
| 类型 | 需要说明的内容 |
|---|---|
| 功能补全 | 支持参数沿 dim-0 非均匀切分,并允许部分 rank 持有 dim-0 长度为 0 的分片 |
| 用户需求 | 包含非整除参数的模型可以直接使用现有 fully_shard() 接口完成分布式训练 |
本 RFC 要解决的问题:FSDP 无法训练 dim-0 长度不能被 shard_world_size 整除的参数。
完成后的成功标准:FSDP 生成正确的实际分片和 placements,通信 padding 正确,端到端精度与非分布式基线一致。
3. 目标和非目标
3.1 目标
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 均匀切分的接口和执行路径不变。
3.2 非目标
1. 本期不支持 dim-0 以外轴的非均匀切分;shard_dim != 0 时继续要求参数能够被 shard_world_size 整除。
2. 本期不新增 fully_shard() 对外参数,非均匀切分由参数形状自动触发。
3. 本期不修改 DTensor 或分布式检查点的 RaggedShard 实现,这些能力由 DTensor 非均匀切分需求提供。
4. 本需求是功能完善,不以提升性能或降低显存为目标。
5. 对外接口
5.1 接口定义
mesh = init_device_mesh(
device_type="npu",
mesh_shape=(4,),
mesh_dim_names=("dp",),
)
model = fully_shard(model, mesh=mesh)
| 入参 / 配置项 | 类型 | 默认值 | 是否必填 | 含义 | 合法范围 | 错误处理 |
|---|---|---|---|---|---|---|
module |
Module |
无 | 是 | 需要执行 FSDP 的模块 | PT/MS 模块 | 类型不合法时沿用现有报错 |
mesh |
DeviceMesh |
None |
否 | 定义 FSDP/HSDP 通信域 | 现有 FSDP/HSDP mesh | mesh 不合法时沿用现有报错 |
shard_placement_fn |
Callable 或 None |
None |
否 | 指定参数切分轴;None 表示 dim-0 |
非均匀切分仅支持 dim-0 | 非 dim-0 无法整除时抛出 NotImplementedError |
5.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 接口说明
为什么这样设计:非均匀切分是参数形状与 shard_world_size 的结果,不需要用户开关。
和已有接口是否一致:一致,继续使用 fully_shard() 和 shard_placement_fn。
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 --> Grad
6.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 生成
- dim-0 均匀切分继续使用现有
Shard或StridedShard。 - dim-0 非均匀切分根据各 rank 的
actual_shard_length生成local_units。 - FSDP/HSDP 使用
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)。 - 非均匀切分时,创建全 0 的
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 输入元素数一致。- AllGather 完成后只暴露原始
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 代码改动点
| 模块 | 改动内容 | 是否影响已有行为 |
|---|---|---|
model |
platform/{torch,mindspore}/fully_shard/param.py 生成实际分片、ragged placements 和补齐通信存储 |
仅影响 dim-0 非均匀参数 |
distributed |
两个后端的 fully_shard/param_group.py 按 padded_sharded_param_size 组织通信,按 sharded_size 回填梯度 |
均匀路径不变 |
checkpoint |
FSDP 提供正确的 sharded_param.placements;DCP 逻辑由 DTensor 非均匀切分需求提供 |
不修改 DCP 接口 |
optimizer |
optimizer 继续持有 self.sharded_param,只能看到实际分片和实际梯度 |
接口和均匀场景不变 |
trainer |
延迟初始化后的 reset_sharded_param() 重建 padding storage |
现有调用顺序不变 |
6.6 方案取舍
| 方案 | 优点 | 缺点 | 是否选择 | 原因 |
|---|---|---|---|---|
| 使用 DTensor ragged placements 表达实际分片,FSDP 私有 storage 处理通信 padding | placements、optimizer 和 DCP 看到的都是实际分片;职责清晰 | 依赖 DTensor 非均匀切分能力 | 是 | 原生表达非均匀分片,不污染普通 Shard 语义 |
继续使用 Shard(0),在 Layout 中额外保存 logical shape |
改动集中 | Shard(0) 不能表达非均匀分片,DCP 无法仅根据 placements 得到实际范围 |
否 | 属于旁路元数据,无法形成完整语义 |
把 padding 直接放入 sharded_param |
通信输入简单 | optimizer 和 DCP 会看到无效元素 | 否 | 参数语义错误 |
7. 组件依赖
| 依赖组件 | 强依赖 / 弱依赖 | 当前状态 | 未 ready 时本期能力 |
|---|---|---|---|
| FSDP | 强依赖 | 已有均匀切分流程 | 无法训练 dim-0 非均匀参数 |
| DTensor 非均匀切分 | 强依赖 | 独立需求开发中 | 不允许回退到 Shard(0);本需求不能完整交付 |
| TP / SP | 强依赖于组合场景 | 已有基于 DTensor 的流程 | 只能交付纯 FSDP/HSDP 场景 |
| checkpoint | 强依赖于 DCP 场景 | 依赖 DTensor ragged placements | 训练可验证,但不能声明支持非均匀参数 DCP |
| optimizer | 弱依赖 | 复用现有 optimizer | 只要实际参数和梯度形状正确,无需新增适配 |
| PT / MS 后端 | 均为强依赖 | 两个后端均有 FSDP/HSDP 基础流程 | 任一后端未完成时,该后端不具备本特性 |
完整能力需要:DTensor 提供 RaggedShard、RaggedStridedShard、logical global shape 及 DCP 所需的实际分片表达。
本期最小可交付能力:PT/MS 的 FSDP/HSDP 均可训练 dim-0 非均匀参数,且 sharded_param.placements 正确。
8. 约束与兼容性
| 类型 | 内容 |
|---|---|
| 不支持项 | 不支持 dim-0 以外轴的非均匀切分;DTensor ragged 依赖未就绪时不提供 Shard(0) 兼容旁路 |
| 性能收益 | 无。本需求是功能完善,不设置吞吐或 step time 收益目标 |
| 显存收益 | 无。非均匀场景需要保留通信 padding,不承诺降低峰值显存 |
| 性能劣化 | 非均匀参数增加 padding 初始化、拷贝和无效通信元素;开销是功能正确性的必要代价,可以接受 |
| PT / MS 差异 | 两个后端支持范围和分片语义一致;张量、参数和集合通信操作分别使用对应后端实现 |
| 和已有行为不一致 | 仅 dim-0 不能整除时 placements 改为 ragged 类型;均匀参数继续使用现有 placements 和通信快路径,无需迁移配置 |
9. 验证设计
9.1 用例分层
单元测试(Unit Test,UT)验证 placements 和 buffer 逻辑;系统测试(System Test,ST)按 Level0 和
Level1 覆盖基础训练闭环及多组件组合。
| 用例级别 | 数量 | 覆盖内容 | 通过标准 |
|---|---|---|---|
| UT | 每后端不少于 8 个 | FSDP/HSDP/TP+FSDP/TP+HSDP placements;实际分片;padding 初始化、重建、融合通信和梯度回填 | 两个后端的 placements、shape、offset 和 padding 值精确匹配预期 |
| Level0 | 每后端不少于 2 个 | FSDP、HSDP 中包含 dim-0 不能被 shard_world_size 整除的参数 |
两个后端的 loss、梯度和参数更新分别与非分布式基线一致 |
| Level1 | 两个后端按组合矩阵覆盖 | 在 FSDP/HSDP 基础上组合基于 DTensor 的 TP/SP、预取、重计算和延迟初始化 | 训练无异常退出或通信挂起,精度符合现有门槛 |
9.2 交互验证(举例)
| 组合 | 是否验证 | 通过标准 |
|---|---|---|
| FSDP / HSDP + dim-0 非均匀参数 | 是 | sharded_size、padded_sharded_param_size 正确;端到端精度与非分布式基线一致 |
| TP + FSDP / TP + HSDP | 是 | 正确生成并保留 ragged 与 TP placements,梯度分片与基线一致 |
| SP + FSDP / SP + HSDP | 是 | SP 输入输出布局不变,FSDP 参数和梯度归约正确 |
| 本特性 + 预取 | 是 | 预取 AllGather 使用补齐输入,无越界、无通信挂起 |
| 本特性 + 重计算 | 是 | 重计算前后参数 unshard/reshard 正确,loss 与基线一致 |
| 本特性 + 延迟初始化 | 是 | reset_sharded_param() 正确重建实际分片视图和补齐通信存储 |
| 本特性 + DCP | 依赖项验证 | sharded_param.placements 能表达实际分片;保存和加载由 DTensor 非均匀切分需求验收 |
| PT / MS 对齐 | 是 | 支持范围、placements、padding 和梯度语义一致 |
UT 需要至少覆盖以下边界:
- dim-0 长度不能整除
shard_world_size。 - dim-0 长度小于
shard_world_size,后部 rank 的sharded_size[0]为 0。 - padding 初始化为 0,buffer 复用前 padding 仍为 0。
- AllGather 输入按
padded_sharded_param_size对齐,输出只暴露逻辑参数范围。 - ReduceScatter 输入尾部补 0,输出 offset 按补齐后的分片区域前进,最终梯度只使用实际范围。
- 均匀参数不创建额外 padding 分配或拷贝。
9.3 性能 / 显存验证
| 场景 | 基线 | 开启本特性 | 指标 | 通过标准 |
|---|---|---|---|---|
| PT/MS dim-0 均匀参数 | 当前 FSDP/HSDP | 新实现的均匀快路径 | step time / peak memory | 不引入额外 padding 分配或拷贝;不设置性能收益目标 |
| PT/MS dim-0 非均匀参数 | 无可运行基线 | ragged placements + padding communication storage | step time / peak memory | 仅记录数据,不作为性能或显存收益验收项 |
10. CheckList
- FSDP遇到参数不能均匀切分的场景,和torch fully_shard的显存、性能应当基本持平。
- FSDP / HSDP + dim-0 非均匀参数存在, 精度误差符合当前精度阈值标准。
- FSDP / HSDP + dim-0 非均匀参数存在 + 参数预取(fully_shard wrap之后, module.set_forward_prefetch_modules, module.set_backward_prefetch_modules) + meta初始化 精度符合标准
- FSDP / HSDP + dim-0 非均匀参数 + 重计算 精度符合标准
- TP切分后, fully_shard拿到的参数切片如果不能均匀切分的场景,fully_shard初始化后,model.parameters()应该是DTensor,placements信息应该正确生成并保留 ragged FSDP与 TP placements,backward后(MindSpore侧要在enable_mindspore_backward_compact 下写测试用例脚本)梯度的placements应当和参数的placements保持一致。正反向训练流程应当功能无误, 精度符合标准。
- fully_shard 对外接口无变化。padding是内部处理行为,脚本侧不感知。
schema_version: 1
source: gitcode
gitcode_repo: mindspore/hyper-parallel
gitcode_issue: 319
source_url: https://gitcode.com/mindspore/hyper-parallel/issues/319
Contributor guide
No contributing guide indexed for this repository
First steps
- Read the whole issue, then the project's contributing guide.
- Comment on the issue to say you are picking it up — it saves two people doing the same work.
- Fork the repository and make your change on a branch.
- Open a pull request that references the issue number.
Research direction
Start with platform/{torch,mindspore}/fully_shard/param.py and the corresponding fully_shard/param_group.py files, then review the DTensor ragged-sharding dependency before changing FSDP behavior. Use the proposed per-backend unit tests and Level0/Level1 validation matrix to verify actual shard sizes, padding, placements, communication, and end-to-end accuracy.
Written by the indexing model from the issue text.
Assessment
- Tech stack
- python, pytorch
- Domain
- distributed-systems, machine-learning
- Issue type
- Feature
- Difficulty
- 5/5
- Estimated time
- Over a week
- Activity status
- Active
- Clarity
- Mostly clear
- Newbie friendliness
- 35/100