已开启
[RFC]: FSDP支持参数非均匀切分 #319
MengXY107创建于  8月6日
MengXY107
MengXY107成员
8月6日 创建

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是内部处理行为,脚本侧不感知。
likedislike
MengXY107MengXY107成员
8月6日 修改了issue 的描述
MengXY107MengXY107成员
8月13日 修改了issue 的描述
MengXY107MengXY107成员
8月13日 修改了issue 的描述
MengXY107MengXY107成员
8月13日 修改了issue 的描述
MengXY107MengXY107成员
8月14日 修改了issue 的描述