mindspore-ai / mindspore-ai/hyper-parallel

[RFC]: FSDP支持参数非均匀切分

Open
#196 0 comments 0 reactions 0 assignees View on GitHub

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 CallableNone 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_paramparam_data 的实际本地分片,允许 actual_shard_length 为 0。
self.sharded_size 记录实际形状;self.padded_sharded_param_size 将 dim-0 设置为
dim_shard_size,记录集合通信要求的统一形状。

placements 生成
  • dim-0 均匀切分继续使用现有 ShardStridedShard
  • 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.pypadded_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_sizepadded_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

  1. Read the whole issue, then the project's contributing guide.
  2. Comment on the issue to say you are picking it up — it saves two people doing the same work.
  3. Fork the repository and make your change on a branch.
  4. 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

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.