已关闭
【RFC】DTensor支持非均匀切分的表达及重排 #314
yide12创建于  8月4日关闭于  2 天前
yide12
yide12成员
8月4日 创建

HP DTensor 非均匀切分(RaggedShard)设计文档

0. 基本信息

项目 内容
特性名称 HP DTensor 支持 RaggedShard 非均匀连续切分
开发分支 hp_ragged_shard
适用后端 PyTorch、MindSpore
已验证设备 PyTorch NPU/HCCL 实机;PyTorch CPU/Gloo 路径和 MindSpore 接口已完成 mock UT
目标合入时间 2026/08/10
当前阶段 Phase 1:单个 Ragged mesh 维、其余 mesh 维仅允许 Replicate

1. 背景

普通 Shard(dim) 描述的是沿单个逻辑维度进行规则切分,不能表达“各 rank 按业务指定比例持有不同数据量”。长序列、变长样本、MoE token、分块量化和零拷贝 FSDP 等场景需要一种显式的非均匀 placement:

RaggedShard(
    dims=(0, 1),
    local_units=(2, 3, 0, 5),
)
  • dims 描述参与连续展平切分的逻辑前缀维度。
  • local_units 描述该 mesh 维上各 rank 的相对持有量。
  • local_units 可以包含 0,因此允许空分片 rank。
  • Ragged local tensor 使用一维连续 flat storage,逻辑全局 shape 单独保存在 DTensor 中。

该能力不能只通过增加一个 Placement 类完成,还需要贯通:

Placement -> Layout -> DTensor 元数据/切分 -> 变长通信
          -> redistribute -> op dispatch -> DCP save/load/reshard

相关资料:


2. 本期目标与非目标

2.1 本期目标

  1. 新增公开接口 RaggedShard(dims, local_units),并提供完整校验、相等性、hash 和字符串表达。
  2. Layout 无损保留 Ragged placement,同时为旧 tensor-map 流程提供 RaggedShard -> Replicate 的 normal view。
  3. 支持从全局 tensor 创建 Ragged DTensor,以及从 flat local tensor 构造 Ragged DTensor。
  4. 支持 full_tensor() 和以下重排:
    • normal -> ragged;
    • ragged -> normal;
    • ragged -> ragged,仅 local_units 变化时使用变长 all-to-all;
    • ragged -> ragged,dims 或 ragged mesh 维变化时经 Replicate 中转。
  5. PT/MS 平台层提供可微变长 all-gather 和可微变长 all-to-all。
  6. 白名单 elementwise 算子在 flat local storage 上执行,并继承 Ragged layout/global shape。
  7. DCP 不汇总完整 tensor,直接将 Ragged flat interval 映射为标准 N-D chunks,复用现有 chunk 求交和 reshard 流程。

2.2 本期非目标

  • 一个 layout 中存在多个 RaggedShard。
  • 非 prefix dims,例如 (1,)、(0, 2)。
  • Ragged 所在 layout 的其他 mesh 维使用 Shard 或 Partial。
  • _StridedRaggedShard,以及同一逻辑维度上嵌套 Ragged/Shard 的顺序表达。
  • Ragged DTensor factory,例如 DTensor.empty/full/rand(..., RaggedShard(...))。
  • reduction、view/reshape、全局索引、matmul、attention 等复杂算子的自动 Ragged propagation。
  • 整网性能、显存收益承诺。

3. RaggedShard 语义

3.1 Placement 约束

RaggedShard(dims, local_units) 当前满足:

  • dims 必须是非空 tuple[int, ...]。
  • dims == tuple(range(len(dims))),即必须是连续前缀维度。
  • local_units 必须是非空 tuple[int, ...]。
  • 每个 unit 必须是非负整数,且 sum(local_units) > 0。
  • len(local_units) 必须等于 Ragged 所在 mesh 维的大小。
  • 一个 layout 最多包含一个 RaggedShard。
  • 其他 mesh 维在 Phase 1 中必须是 Replicate()。

其中,Placement 构造阶段校验 tuple、类型、prefix 和 unit 非负性;依赖 global shape/mesh 的几何校验在 DTensor 构造或切分阶段完成。

3.2 Flat interval 计算

设:

global_shape = (s0, s1, ..., sn)
dims         = (0, 1, ..., k - 1)
units        = (u0, u1, ..., up - 1)
rank         = r

计算公式:

prefix_cells      = prod(global_shape[:k])
suffix_numel      = prod(global_shape[k:])
cells_per_unit    = prefix_cells / sum(units)
prefix_start      = sum(units[:r]) * cells_per_unit
local_prefix      = units[r] * cells_per_unit
flat_start        = prefix_start * suffix_numel
flat_end          = (prefix_start + local_prefix) * suffix_numel
local_numel       = flat_end - flat_start

要求 prefix_cells % sum(units) == 0,保证一个 unit 对应整数个 prefix cell,并且不会切穿 dims 之后的后缀块。

3.3 真实示例

global_shape = (6, 4, 8)
placement = RaggedShard(dims=(0, 1), local_units=(1, 2))
prefix_cells   = 6 * 4 = 24
total_units    = 1 + 2 = 3
cells_per_unit = 8
suffix_numel   = 8
rank prefix cell 区间 flat 区间 local flat shape
0 [0, 8) [0, 64) (64,)
1 [8, 24) [64, 192) (128,)

本例的边界刚好与第 0 维行边界对齐,可以概念性理解为 rank 0 持有 (2, 4, 8)、rank 1 持有 (4, 4, 8)。但 DTensor 内部统一保存 (64,) 和 (128,) 的一维连续 tensor,不能依赖 local N-D shape 反推 global shape。


4. 总体设计

4.1 架构与数据流

total_dtensor_ragged.png

用户 API
  RaggedShard / distribute_tensor / DTensor.from_local
       |
       v
Layout 表达
  original placements          normal placements
  (RaggedShard(...),)   <->    (Replicate(),)
       |                              |
       |                              +--> 复用已有 tensor-map/normal redistribute
       v
Ragged 几何
  global_shape + dims + local_units + mesh local rank
       |
       +--> flat_start / flat_end / all-gather splits / all-to-all splits
       |
       +--> distribute/full_tensor/redistribute
       |
       +--> DCP N-D boxes

设计原则是:Ragged 元数据由原始 placements 无损保存;旧流程看到的 normal view 是 Replicate;只有实际切分、通信和 checkpoint 几何进入 Ragged 专用逻辑。

4.2 Placement 与 Layout

代码位置:

  • hyper_parallel/core/dtensor/placement_types.py
  • hyper_parallel/core/dtensor/layout.py

Layout 保存三种视图:

视图 含义
placements 原始 placement,保留完整 RaggedShard(dims, local_units)
ragged_shard RaggedShardInfo(mesh_dim, placement),用于快速识别
normal_placements 将 Ragged 替换为 Replicate(),供旧 tensor-map 和普通重排使用

核心行为:

def set_placements(placements):
    self._placements = placements
    self._ragged_shard = extract_single_ragged(placements)

@property
def normal_placements(self):
    return tuple(
        Replicate() if p.is_ragged_shard() else p
        for p in self._placements
    )

@property
def alias_placements(self):
    if self._ragged_shard is not None:
        return self._placements
    return existing_alias_behavior()
  • placement_to_tensor_map() 基于 normal_placements 工作,因此 Ragged mesh 维不会被错误编码成普通 Shard。
  • tensor_map_to_placement() 完成普通 placement 恢复后,会在保存的 mesh 维重新注入原始 Ragged placement。
  • alias_placements 对 Ragged 返回原始 placements,避免重建 DTensor 时丢失 dims/local_units。
  • RaggedShard.__hash__() 包含 dims/local_units,不同 Ragged 布局不会命中同一个 Layout cache key。

4.3 DTensor 元数据与本地存储

代码位置:

  • hyper_parallel/core/dtensor/dtensor.py
  • hyper_parallel/core/dtensor/_ragged_utils.py

Ragged DTensor 的不变量:

logical shape: DTensor._global_shape
local storage: contiguous 1-D tensor
local numel:   compute_ragged_slice(global_shape, layout).local_numel

DTensor.from_local() 在 Ragged 场景必须显式传入 shape:

dt = DTensor.from_local(
    local_flat,
    mesh,
    (RaggedShard((0, 1), (1, 2)),),
    shape=(6, 4, 8),
)

构造时校验:

  • global shape 由非负整数构成;
  • global shape rank 与 Layout tensor-map rank 一致;
  • local tensor 连续且为一维;
  • local numel 与当前 rank 的 Ragged interval 一致。

普通 DTensor 同样保存 _global_shape,未显式传入时继续由 Layout 和 local shape 推导;Ragged 场景不能通过 local flat shape 推导,因此强制显式提供。

当前不向 Layout 增加 global shape 或 stride 字段,也不保存全局 stride。Ragged Phase 1 仅支持连续 row-major flat storage。

4.4 distribute_tensor 与 full_tensor

4.4.1 本地切片模式

distribute_tensor(global_tensor, mesh, placements, src_data_rank=None)

每个 rank 都持有完整且相同的 global tensor,调用 _slice_ragged_tensor() 计算本 rank 的 [flat_start, flat_end) 并 clone 为独立 local flat tensor。

4.4.2 源 rank 分发模式

distribute_tensor(global_tensor, mesh, placements, src_data_rank=0)

调用链:

distribute_tensor
  -> _scatter_ragged_tensor
  -> 为每个 group rank 计算 flat interval
  -> mesh_scatter_ragged
     -> source rank: 本地 copy + isend
     -> other ranks: irecv
  -> DTensor.from_local_with_layout(shape=global_shape)

src_data_rank 是 Ragged mesh 通信组内的相对 rank。该 scatter 用于创建 local shard;当前不承诺梯度跨 P2P 回传到 source rank 的原始 global input。创建后的 local Ragged DTensor 可以正常参与已支持算子的 autograd。

4.4.3 full_tensor

full_tensor() 构造全 Replicate 目标 Layout,并以逻辑 global rank 建立 tensor map:

replicated_layout.placement_to_tensor_map(len(self._global_shape))

随后进入 ragged_to_normal(),用可微变长 all-gather 按 rank 顺序拼接 flat shards,最后 reshape(global_shape)。

4.5 平台通信原语

统一接口位于 hyper_parallel/platform/platform.py:

differentiable_variable_all_gather(input_tensor, output_splits, group)
differentiable_all_to_all_single(input_tensor, input_splits, output_splits, group)

split 的统一语义是 dim 0 行数;Ragged 重排传入一维 tensor,因此行数等同于元素数。

PyTorch

能力 前向 反向
NPU 变长 all-gather dist.all_gather(),按真实长度预分配 list torch_npu.distributed.reduce_scatter_tensor_uneven()
CPU/Gloo 变长 all-gather pad 到最大长度后 all_gather(),再 trim 对完整梯度 all_reduce(),再截取本 rank 区间
变长 all-to-all torch.distributed.nn.functional.all_to_all_single() PyTorch autograd 执行反向 A2A

MindSpore

能力 前向 反向
变长 all-gather 将 N-D tensor 展平,把行 split 转成元素 split,调用 ops.AllGatherV 依赖 MindSpore AllGatherV 自动微分
变长 all-to-all comm_func.all_to_all_single(),split 为 dim 0 行数 自定义 Function 交换 input/output splits,执行反向 A2A

变长 scatter 没有作为公开可微 collective 增加,而是由 mesh_scatter_ragged() 使用 isend/irecv 服务 distribute_tensor(src_data_rank=...)。

4.6 redistribute

代码位置:hyper_parallel/core/dtensor/tensor_redistribution.py。

外层先判断 source/target 是否包含 Ragged;均不包含时完全复用原有重排流程。Ragged 分支只有四种状态转换:

source target 实现
normal ragged normal 重排到 target normal view,再本地 slice
ragged normal 变长 all-gather 到 source normal view,再走普通重排
ragged ragged,mesh dim/dims 相同 根据 source/target flat interval 交集计算 splits,执行变长 all-to-all
ragged ragged,mesh dim 或 dims 不同 ragged -> Replicate -> target ragged

重排.png

伪代码:

if not src_ragged and not dst_ragged:
    return normal_redistribute(x, dst)

if src_ragged and dst_ragged:
    if same_ragged_axis_and_dims and same_normal_view:
        return ragged_to_ragged_all_to_all(x, dst)
    full = ragged_to_normal(x, src_normal)
    return normal_to_ragged(full, dst)

if src_ragged:
    normal = ragged_to_normal(x, src_normal)
    return normal if normal.layout == dst else normal_redistribute(normal, dst)

normal = x if x.layout == dst_normal else normal_redistribute(x, dst_normal)
return normal_to_ragged(normal, dst)

同 dims 的 Ragged-to-Ragged 不 materialize 完整 tensor。每个 source rank 向每个 target rank 发送两个 flat interval 的交集长度:

input_splits[target_rank]  = overlap(source_local_interval, target_interval)
output_splits[source_rank] = overlap(source_interval, target_local_interval)

Partial -> Ragged 会先通过原有 reduce_partial() 消除 Partial,再进入 normal-to-ragged;Ragged -> Shard 会先恢复 source normal view,再复用普通重排。本期不把 Ragged -> Partial 作为受支持语义,因为不能从一个完整值无条件反推出 pending-reduction 状态。

4.7 Op dispatch

发现任一 Ragged DTensor 输入后,dispatcher 采用 fail-closed 策略:只允许白名单 elementwise,本地执行后使用第一个 Ragged 输入的 Layout 和 global shape 包装输出。

当前白名单:

unary:
abs, absolute, clone, cos, exp, gelu, isinf, isnan, log, neg,
negative, relu, rsqrt, sigmoid, silu, sin, sqrt, square

binary:
add, div, mul, pow, real_div, sub, __rsub__, __rpow__, true_divide

特殊场景:

full_x = [[1,2,3,4], [5,6,7,8],[9,10,11,12],[13,14,15,16]]
x = distribute_tensor(full_x, tp_mesh, (shard(0))
x[2].zero_()
rank0:
x[2]._local_tensor = []
rank1:
x[2]._local_tensor=full_tensor

当前实现不预先校验所有 Ragged 输入的 layout 是否相同,也不重复实现广播/shape 校验;白名单命中后交给底层本地算子执行,输出构造和后续流程在不满足 Ragged 不变量时继续报错。

以下算子仍 fail-closed:

  • reduction,例如 mean/sum;
  • view/reshape/flatten;
  • global indexing 和跨 rank 原地修改;
  • matmul、norm、attention 等需要独立分布式语义的算子;
  • 任意未加入白名单的算子。

4.8 Distributed Checkpoint

代码位置:

  • hyper_parallel/core/distributed_checkpoint/ragged_utils.py
  • hyper_parallel/core/distributed_checkpoint/standard_planner.py
  • hyper_parallel/core/distributed_checkpoint/filesystem_storage.py
  • hyper_parallel/core/distributed_checkpoint/util.py
  • hyper_parallel/core/distributed_checkpoint/async_staging.py

4.8.1 几何适配

DCP 原有 reshard 基于标准 N-D ChunkStorageMetadata(offsets, sizes) 求交。Ragged local storage 虽然是一维连续区间,但该区间可能跨越 N-D 行或平面边界,因此保存前将 flat interval 分解为一组有序 N-D boxes:

compute_ragged_boxes(dtensor)
  -> _compute_ragged_slice(global_shape, layout)
  -> _decompose_flat_interval(shape, flat_start, flat_end)
  -> [(offsets, sizes, local_flat_start, local_flat_end), ...]

每个 box 对应一个普通 WriteItem。真正的数据仍来自:

local_flat[local_flat_start:local_flat_end].reshape(box_sizes)

不需要先将 Ragged tensor 通信成 Replicate。

4.8.2 保存调用链

StandardSavePlanner.build_local_plan
  -> create_ragged_write_items
  -> 每个 N-D box 生成一个 WriteItem/ChunkStorageMetadata

StandardSavePlanner.get_data
  -> get_ragged_box_tensor
  -> detach().cpu()

同一个逻辑 FQN 可能在同一 safetensors 文件内对应多个 box。FileSystemWriter 为其生成唯一物理 key:

逻辑 FQN: model.weight

物理 key:
model.weight.__dcp_chunk_0
model.weight.__dcp_chunk_1

StorageInfo.tensor_key 保存逻辑 chunk 到物理 key 的映射。.metadata 中的 TensorStorageMetadata 和 MetadataIndex.fqn 仍使用原参数名,因此 load planner 和 reshard 逻辑不受物理改名影响。

4.8.3 加载与 reshard

目标 Ragged DTensor 的 create_chunk_list_for_tensor() 同样生成目标 N-D boxes。现有 create_read_items_for_chunk_list() 和 chunk intersection 逻辑直接计算 checkpoint chunks 与目标 boxes 的交集。

StandardLoadPlanner.acquire_tensor
  -> get_ragged_box_tensor(target, dest_index)
  -> narrow_tensor_by_index(...)
  -> reader 将 checkpoint slice 写入目标 local flat view

因此可复用现有 DCP 流程支持:

  • Ragged -> Ragged,相同或不同 local_units;
  • Ragged -> Ragged,不同 prefix dims;
  • Ragged -> Replicate/Shard;
  • Replicate/Shard -> Ragged。

异步 staging 重建 Ragged DTensor 时显式传入 shape=tuple(obj.shape)。第一版只要 state_dict 中存在 Ragged DTensor,就禁用 SavePlan cache,避免不同 global shape、dims 或 local_units 复用错误计划。

更完整的文件结构和 metadata 示例见 dcp流程.md。

4.9 缓存与兼容策略

  • Placement/Layout cache:key 包含 RaggedShard,其 hash 包含 dims/local_units。
  • normal redistribution cache:Ragged 专用流程先转换成确定的 normal Layout,再使用旧缓存。
  • op dispatch:Ragged 白名单路径不复用可能丢失 Ragged 元数据的普通 layout 推导。
  • DCP SavePlan cache:检测到 Ragged state_dict 时禁用。
  • 普通 DTensor:没有 Ragged 时继续走原有 Layout、通信、重排、算子和 DCP 分支。

5. 对外接口

5.1 Placement

from hyper_parallel import RaggedShard

placement = RaggedShard(
    dims=(0, 1),
    local_units=(1, 2),
)

RaggedShard 已从 hyper_parallel 顶层导出。

5.2 从全局 tensor 创建

dt = distribute_tensor(
    global_tensor,
    mesh,
    (RaggedShard((0, 1), (1, 2)),),
    src_data_rank=0,
)
  • src_data_rank=None:各 rank 本地切片,不通信。
  • src_data_rank=int:从 Ragged group 内指定相对 rank 进行变长 P2P 分发。

5.3 从本地 flat tensor 创建

dt = DTensor.from_local(
    local_flat,
    mesh,
    (RaggedShard((0, 1), (1, 2)),),
    shape=(6, 4, 8),
)

Ragged 场景 shape 必填;local_flat 必须连续、一维且 numel 与本 rank 配额一致。

5.4 重排与恢复

full = dt.full_tensor()

changed = dt.redistribute(
    mesh,
    (RaggedShard((0, 1), (2, 1)),),
)

replicated = dt.redistribute(mesh, (Replicate(),))
sharded = dt.redistribute(mesh, (Shard(0),))

6. 当前支持矩阵

能力 PT MS 当前状态/限制
Placement/Layout 表达 支持 支持 最多一个 Ragged
DTensor.from_local 支持 支持 必须显式 global shape;local storage 为连续 1-D
distribute_tensor(src_data_rank=None) 支持 支持 每个 rank 必须持有相同全局输入
distribute_tensor(src_data_rank=int) 支持 支持 P2P 变长 scatter;不承诺跨 rank 输入梯度
full_tensor() 支持 支持 通过可微变长 all-gather
Ragged <-> Replicate 支持 支持 核心路径
Ragged <-> Shard 支持 支持 先转换到 normal view,再复用普通重排
Partial -> Ragged 支持 支持 先执行 reduce_partial();不支持反向转换到 Partial
Ragged -> Ragged,仅 units 变化 支持 支持 同 mesh dim/dims,变长 A2A
Ragged -> Ragged,dims/mesh dim 变化 支持 支持 经 Replicate 中转,通信量更大
其他 mesh 维为 Replicate 支持 支持 Phase 1 唯一 mixed-mesh 形式
其他 mesh 维为 Shard/Partial 不支持 不支持 _compute_ragged_slice() 明确报错
elementwise 白名单 支持 支持 本地计算,继承第一个 Ragged 输入布局
reduction/view/matmul/attention 不支持 不支持 fail-closed
DCP save/load/reshard 支持 支持 flat interval 转 N-D boxes;Ragged 时禁用 SavePlan cache
Ragged DTensor factory 不支持 不支持 明确 NotImplementedError
_StridedRaggedShard 不支持 不支持 后续阶段
CUDA/NCCL 未验证 不涉及 本期不适配

“支持”表示实现路径已补齐;MindSpore 当前 PR 内验证以 CPU mock UT 为主,仍需补充真实多卡 NPU ST 作为合入验收证据。


7. 风险与限制

7.1 Flat storage 对算子的影响

Ragged local tensor 的物理 shape 是 (local_numel,),而不是逻辑 N-D local shape。任何依赖维度语义的算子都不能直接复用普通 DTensor 推导,否则可能按物理一维 shape 推导错误。因此一期仅开放与 shape 语义无关的 elementwise 白名单。

7.2 经 Replicate 中转的通信代价

当 Ragged 的 dims 或 mesh dim 改变时,当前正确性路径会先执行变长 all-gather,形成 Replicate,再本地 slice。每个参与 rank 都会 materialize 完整 normal tensor,通信量和峰值内存高于直接 Ragged-to-Ragged 重排。只有同 dims、同 mesh dim、仅 units 变化的路径使用直接 A2A。

7.3 mixed-mesh 限制

当前几何计算只建立一个全局连续 flat interval,没有表达 Ragged 与其他 Shard 的应用顺序。若允许 (RaggedShard(...), Shard(...)),local interval、全局 offset、A2A overlap 和 DCP boxes 都会依赖另一个 mesh 维的切分结果。Phase 1 因此明确拒绝,后续需要 _StridedRaggedShard 或等价的顺序元数据。

7.4 空分片 rank

local_units 可以为 0,几何层会生成 local_numel == 0 的一维 tensor。collective 和 DCP 路径必须持续覆盖空输入,避免底层后端对 0 长度 buffer 的行为差异造成 hang。当前已有 zero-unit 创建、elementwise、full tensor 和重排验证;MS 真实多卡仍需补测。

7.5 DCP 物理 key 兼容

新增 StorageInfo.tensor_key 是可选字段。旧 checkpoint 没有该字段时,Reader 回退到逻辑 FQN,因此原有单 tensor-key checkpoint 保持兼容。


8. 验证设计与当前结果

8.1 UT 覆盖

当前 UT 覆盖:

  • Placement:构造、非法 dims/units、repr、eq、hash。
  • Layout:原始 placements、normal view、alias 恢复、单 Ragged 限制。
  • 几何:flat slice、zero unit、units 变化的 A2A overlap splits。
  • DTensor:from_local global shape 校验、flat storage、distribute/full tensor、四类重排。
  • Op dispatch:elementwise 前后向、layout/global shape 继承、非白名单 fail-closed。
  • 通信:Torch 变长 all-gather 前后向;MindSpore AllGatherV split 转换和变长 A2A 反向 split 交换。
  • DCP:flat interval 分解、WriteItem、box view、filesystem tensor key、save/load planner、reshard、async staging。

warning 为环境中的 torch_npu TypedStorage 弃用提示,与 Ragged 功能无关。


9. 验收标准

9.1 功能验收

  • RaggedShard 校验、表达和公共导入正确。
  • Ragged local storage、global shape、flat offset 和 rank 顺序一致。
  • src_data_rank=None/int 创建流程均正确。
  • full_tensor() 恢复原始全局 tensor。
  • 四类重排前向正确;涉及 gather/A2A 的路径反向梯度正确。
  • elementwise 白名单输出保持 Ragged placement/global shape;非白名单明确报错。
  • DCP same-layout、changed-units、Ragged/normal reshard 数据一致。
  • zero-unit rank 不错位、不崩溃、不 hang。

9.2 兼容性验收

  • 普通 Placement/Layout/DTensor/redistribute/op/DCP 路径行为不变。
  • 旧 checkpoint 没有 tensor_key 时仍可读取。
  • PT/MS 使用相同的公开接口和 split 语义;后端实现差异仅保留在 platform 层。

9.3 明确报错

以下场景必须 fail-closed,而不是静默按普通 Shard/Replicate 处理:

  • 多个 Ragged placements;
  • 非 prefix dims;
  • prefix cells 不能被总 units 整除;
  • local_units 长度与 mesh dim size 不一致;
  • Ragged 之外 mesh 维不是 Replicate;
  • from_local 缺少 global shape,或 local tensor 非一维/非连续/numel 不匹配;
  • 未支持的算子或 Ragged factory。
likedislike
yide12yide12成员
8月4日 issue类型由 Documentation 改变为 RFC
yide12yide12成员
8月4日 修改了issue 的描述
yide12
yide12成员
8月5日 评论:

RaggedShard DCP 保存加载流程

本文用一个真实的两卡 HCCL 场景说明 RaggedShard 如何保存、加载,以及 metadata 和 checkpoint 文件之间的对应关系。

1. 示例输入

示例使用:

global_shape = (3, 4, 3)
ragged_dims = (0, 1)
local_units = (1, 5)
mesh = init_device_mesh(
    device_type="npu",
    mesh_shape=(2,),
    mesh_dim_names=("ragged",),
)

源数据只在 rank 0 有效,rank 1 传入同 shape 的零 tensor:

global_tensor = make_global_tensor((3, 4, 3))
source_input = (
    global_tensor
    if rank == 0
    else torch.zeros_like(global_tensor)
)

weight = distribute_tensor(
    source_input,
    mesh,
    (RaggedShard((0, 1), (1, 5)),),
    src_data_rank=0,
)

调用链:

distribute_tensor()
  -> _build_layout()
  -> _scatter_ragged_tensor()
  -> mesh_scatter_ragged()
  -> rank 0 P2P send / rank 1 P2P recv
  -> DTensor.from_local_with_layout(
         local_flat_tensor,
         ragged_layout,
         shape=(3, 4, 3),
     )

1.1 Ragged 几何

Ragged 前缀是前两维:

prefix cells = 3 * 4 = 12
total units  = 1 + 5 = 6
cells/unit   = 12 / 6 = 2
suffix numel = 3

两个 rank 持有的 flat interval:

rank unit 数 prefix cell 区间 flat 区间 local numel
0 1 [0, 2) [0, 6) 6
1 5 [2, 12) [6, 36) 30

rank 0 对应一个 N-D box:

offsets = (0, 0, 0)
sizes   = (1, 2, 3)

rank 1 的 flat interval 跨越二维行边界,被拆成两个 N-D box:

box 0:
  offsets = (0, 2, 0)
  sizes   = (1, 2, 3)

box 1:
  offsets = (1, 0, 0)
  sizes   = (2, 4, 3)

这部分由 hyper_parallel/core/distributed_checkpoint/ragged_utils.py 完成:

_compute_ragged_slice()
  -> 得到当前 rank 的 flat_start / flat_end

_decompose_flat_interval()
  -> 将 flat interval 拆成有序 N-D boxes

compute_ragged_boxes()
  -> 记录 offsets、sizes 和 local_flat_start/end

local flat tensor 没有被重新拼接成 global tensor。保存数据仍来自连续的:

local_flat[start:end]

保存前只将这一段 reshape 成对应 N-D box。

2. 保存流程

调用:

metadata = save(
    {"model.weight": weight},
    checkpoint_id="/tmp/hp_ragged_dcp_inspect",
)

完整调用链:

save()
  -> _save_impl()
  -> StandardSavePlanner.configure_planner()
  -> StandardSavePlanner.build_local_plan()
  -> create_ragged_write_items()
  -> FileSystemWriter._collect_tensors()
  -> FileSystemWriter._write_tensors()
  -> FileSystemWriter.finalize_checkpoint()

2.1 Save Planner 生成 WriteItem

StandardSavePlanner.build_local_plan() 检测到 Ragged layout 后,不走普通 DTensor 的单 chunk 分支,而是调用:

items.extend(create_ragged_write_items(fqn, obj))

每一个 box 对应一个 WriteItem:

WriteItem(
    index=MetadataIndex(
        fqn="model.weight",
        offset=box.offsets,
        index=None,
    ),
    type=WriteItemType.TENSOR,
    tensor_data={
        "chunk": ChunkStorageMetadata(
            offsets=box.offsets,
            sizes=box.sizes,
        ),
        "properties": TensorProperties(
            dtype="torch.float32",
        ),
        "size": (3, 4, 3),
    },
)

字段含义:

  • fqn 是逻辑参数名;
  • offset 用于区分同一个参数的不同 box;
  • chunk.offsets/sizes 是 global N-D 坐标;
  • size 是整个 tensor 的 global shape,不是 local flat shape。

每个 rank 的 local plan 生成后,会通过 all_gather_object 汇总,由 planner 形成 global plan 和 TensorStorageMetadata。

2.2 获取具体 box 数据

StandardSavePlanner.get_data() 对 Ragged DTensor 调用:

get_ragged_box_tensor(obj, item.index)

函数根据 MetadataIndex.offset 找到 box,然后执行:

local_flat = tensor.to_local().reshape((-1,))
box_tensor = local_flat[
    box.local_flat_start:box.local_flat_end
].reshape(box.sizes)

最后保存前执行 detach().cpu()。

3. checkpoint 文件结构

本示例实际生成的目录:

/tmp/hp_ragged_dcp_inspect/
├── .metadata
├── _rank0_.safetensors
└── _rank1_.safetensors

默认 use_collectives=True,因此:

  • 每个 rank 保存自己的 tensor 文件;
  • 只有 coordinator rank 0 写全局 .metadata;
  • 不会生成 rank-local metadata 文件。

3.1 rank 0 文件

rank 0 只有一个 box,safetensors 内部 key 仍是原始 FQN:

_rank0_.safetensors
└── model.weight                  shape=(1, 2, 3)

3.2 rank 1 文件

rank 1 有两个 box。同一个 FQN 不能在 safetensors 字典中重复,因此物理 key 被区分:

_rank1_.safetensors
├── model.weight.__dcp_chunk_0    shape=(1, 2, 3)
└── model.weight.__dcp_chunk_1    shape=(2, 4, 3)

FileSystemWriter._collect_tensors() 的核心规则:

if fqn_counts[fqn] > 1:
    tensor_key = f"{fqn}.__dcp_chunk_{chunk_index}"

tensor_key 是物理存储 key,不改变逻辑 FQN。

4. .metadata 内容

.metadata 是 pickle 序列化的 Metadata 对象:

Metadata(
    state_dict_metadata=...,  # 逻辑 tensor 和 global chunks
    planner_data=...,         # planner 扩展信息
    storage_data=...,         # 逻辑 chunk 到物理文件/key 的映射
    version="1.0",
)

本示例反序列化后的核心 state_dict_metadata:

{
    "model.weight": TensorStorageMetadata(
        properties=TensorProperties(
            dtype="torch.float32",
            requires_grad=False,
            memory_format=None,
        ),
        size=(3, 4, 3),
        chunks=[
            ChunkStorageMetadata(
                offsets=(0, 0, 0),
                sizes=(1, 2, 3),
            ),
            ChunkStorageMetadata(
                offsets=(0, 2, 0),
                sizes=(1, 2, 3),
            ),
            ChunkStorageMetadata(
                offsets=(1, 0, 0),
                sizes=(2, 4, 3),
            ),
        ],
    ),
}

state_dict_metadata 只描述逻辑数据:

FQN: model.weight
global shape: (3, 4, 3)
global chunks: 3 个 N-D box

它不直接保存 rank 文件名或 safetensors 物理 key。

4.1 storage_data

storage_data 把逻辑 chunk 映射到实际文件和物理 key:

{
    MetadataIndex(
        fqn="model.weight",
        offset=(0, 0, 0),
        index=0,
    ): StorageInfo(
        relative_path="_rank0_.safetensors",
        offset=0,
        length=-1,
        tensor_key="model.weight",
    ),
    MetadataIndex(
        fqn="model.weight",
        offset=(0, 2, 0),
        index=1,
    ): StorageInfo(
        relative_path="_rank1_.safetensors",
        offset=0,
        length=-1,
        tensor_key="model.weight.__dcp_chunk_0",
    ),
    MetadataIndex(
        fqn="model.weight",
        offset=(1, 0, 0),
        index=2,
    ): StorageInfo(
        relative_path="_rank1_.safetensors",
        offset=0,
        length=-1,
        tensor_key="model.weight.__dcp_chunk_1",
    ),
}

关系是:

MetadataIndex(fqn + global offset)
    -> StorageInfo(relative_path + tensor_key)
    -> safetensors 中的实际 tensor

offset=0、length=-1 是 safetensors 容器的文件级记录。当前实现无法用单一 byte range 表示容器内部 tensor,因此真正的 tensor 定位由 tensor_key 完成。

5. 加载流程

本示例把保存布局 (1, 5) 加载到不同的 Ragged 布局 (5, 1):

target = distribute_tensor(
    torch.zeros_like(global_tensor),
    mesh,
    (RaggedShard((0, 1), (5, 1)),),
    src_data_rank=None,
)

load(
    {"model.weight": target},
    checkpoint_id="/tmp/hp_ragged_dcp_inspect",
)

目标 local shape:

rank 0: (30,)
rank 1: (6,)

真实验证结果:

rank 0: FULL_MATCH=True
rank 1: FULL_MATCH=True

加载调用链:

load()
  -> FileSystemReader.load_metadata()
  -> StandardLoadPlanner.configure_planner()
  -> StandardLoadPlanner.build_local_plan()
  -> create_chunk_list_for_tensor(target)
  -> compute_ragged_boxes(target)
  -> create_read_items_for_chunk_list()
  -> N-D chunk intersection
  -> FileSystemReader._load_tensor_file()
  -> StandardLoadPlanner.acquire_tensor()
  -> get_ragged_box_tensor(target, dest_index)
  -> 写回目标 flat local tensor

5.1 目标 chunk 和 ReadItem

加载 planner 根据目标 Ragged layout 重新计算目标 boxes,不使用保存时的 local flat 区间:

create_chunk_list_for_tensor(target)

之后标准 create_read_items_for_chunk_list() 对保存 box 和目标 box 求交,生成:

ReadItem(
    storage_index=...,   # 要读取的保存 chunk
    storage_offsets=..., # 在保存 box 中的偏移
    dest_index=...,      # 目标 box
    dest_offsets=...,    # 在目标 box 中的偏移
    lengths=...,         # 交集长度
)

所以 changed-units 加载仍然是标准 N-D chunk intersection,不需要 Ragged 专用 DCP reshard 算法。

5.2 读取物理 tensor

FileSystemReader._load_tensor_file() 先通过 storage_index 找到 StorageInfo,再选择真实 safetensors key:

tensor_key = storage_info.tensor_key or req.storage_index.fqn

然后从对应 physical tensor 中按 storage_offsets 和 lengths 读取交集数据。

5.3 写回目标 flat storage

目标是 Ragged DTensor 时,StandardLoadPlanner.acquire_tensor() 不直接用 N-D offset 索引一维 local tensor,而是:

box_tensor = get_ragged_box_tensor(target, read_item.dest_index)
target_slice = narrow_tensor_by_index(
    box_tensor,
    read_item.dest_offsets,
    read_item.lengths,
)

过程是:

目标 flat local tensor
    -> 找到目标 box 对应的 flat 区间
    -> view 成 N-D box
    -> 写入当前 ReadItem 对应的交集

6. 各类信息的职责

信息 保存位置 作用
FQN MetadataIndex.fqn / state_dict_metadata 逻辑参数名,例如 model.weight
global shape TensorStorageMetadata.size 校验保存和目标 tensor 的逻辑 shape
N-D chunk offsets/sizes TensorStorageMetadata.chunks 描述全局 box,参与 reshard 求交
local flat start/end 运行时计算 将 box 映射回当前 Ragged local storage,不写入 metadata
文件路径 StorageInfo.relative_path 定位 rank 文件
safetensors 物理 key StorageInfo.tensor_key 定位容器内具体 tensor
读取交集 ReadItem.storage_offsets/lengths 从保存 box 截取需要的数据
目标写入位置 ReadItem.dest_offsets/lengths 写入目标 Ragged box

local_flat_start/end 不需要持久化,因为它可以由以下信息重新计算:

global shape + Ragged dims + local_units + rank

7. 当前实现边界

当前方案默认:

  • Ragged local storage 是连续的一维 flat tensor;
  • global_shape 保存在 DTensor 中;
  • Ragged box 按 row-major flat interval 拆分;
  • checkpoint 文件使用 safetensors;
  • 同一个 FQN 的多个 box 通过 tensor_key 区分;
  • DCP 主流程、metadata 模型、chunk intersection 和普通 DTensor reshard 流程复用现有实现;
  • Ragged 与普通 Shard/Replicate 的加载依赖保存和目标 chunk 在 N-D global 坐标上的交集。

8. 实际验证

真实运行使用两卡 HCCL,关键命令:

cd /home/wyd/code_gen_hp_ragged_shard/hp_test

HYPER_PARALLEL_PLATFORM=torch \
HCCL_NPU_SOCKET_PORT_RANGE=51600-51700 \
PYTHONPATH=/home/wyd/code_gen_hp_ragged_shard/hp_v1:/home/wyd/code_gen_hp_ragged_shard/hp_test:/tmp \
python -m torch.distributed.run \
  --nproc-per-node=2 \
  --master-addr=127.0.0.1 \
  --master-port=29580 \
  /tmp/hp_dcp_inspect_case.py

实际输出的关键结果:

FILES ['.metadata', '_rank0_.safetensors', '_rank1_.safetensors']
RANK_0_KEYS ['model.weight']
RANK_1_KEYS [
    'model.weight.__dcp_chunk_0',
    'model.weight.__dcp_chunk_1',
]
RANK_0_TARGET_LOCAL_SHAPE (30,) FULL_MATCH=True
RANK_1_TARGET_LOCAL_SHAPE (6,) FULL_MATCH=True
likedislike
yide12yide12成员
8月5日 修改了issue 的描述
yide12yide12成员
8月5日 修改了issue 的描述
chopin_syp
8月11日 评论:

支持 load-time resharding 流程。(优先级低)
优先级低的实际是否支持

likedislike
zhu-jun-an
8月12日 评论:

1.“Layout 无损保留 Ragged placement,同时为旧 tensor-map 流程提供 RaggedShard -> Replicate 的 normal view” 这里具体表现是什么,旧 tensor-map 流程有哪些?
2. "DCP 不汇总完整 tensor,直接映射为标准 N-D chunks",这个如何验证?

likedislike
caoruyue1
caoruyue1
8月12日 评论:

1.未加入白名单的算子,Ragged DTensor 采取 fail-closed 策略,"fail-closed" 在测试验证中的具体表现是什么
2.缺少与其他模块的相关性说明, 比如支持哪些交互, 不支持哪些交互

likedislike
yide12
yide12成员
8月13日 评论:

问题总结:
问题1: 支持 load-time resharding 流程。(优先级低)优先级低的实际是否支持?
答:现在已经支持。
问题2:“Layout 无损保留 Ragged placement,同时为旧 tensor-map 流程提供 RaggedShard -> Replicate 的 normal view” 这里具体表现是什么,旧 tensor-map 流程有哪些?
答:表现是layout额外保存ragged placement,原流程里tensor_map用replicate替换raggedshard,但不起作用。
问题3: "DCP 不汇总完整 tensor,直接映射为标准 N-D chunks",这个如何验证?
答:信息会到保存metadate.json文件里。
问题4:未加入白名单的算子,Ragged DTensor 采取 fail-closed 策略,"fail-closed" 在测试验证中的具体表现是什么?
答:fail-closed表现是直接报错。
问题5:缺少与其他模块的相关性说明, 比如支持哪些交互, 不支持哪些交互
答:只支持重排、白名单算子、DCP模块。

likedislike
yide12yide12成员
8月13日 修改了issue 的描述
yide12yide12成员
8月18日 修改了issue 的描述
yide12
yide12成员
8月27日 评论:

已收敛为指定的 27 个逻辑算子,并使用真实公开接口验证 dispatcher 名称,没有 mock/修改算子名。

逻辑算子 Torch 接口 → 白名单名 MindSpore 接口 → 白名单名
abs torch.abs(x) → abs mint.abs(x) → Abs
absolute torch.absolute(x) → absolute ops.absolute(x) → Abs
clone torch.clone(x) → clone mint.clone(x) → Clone
cos torch.cos(x) → cos mint.cos(x) → Cos
exp torch.exp(x) → exp mint.exp(x) → Exp
gelu torch.nn.functional.gelu(x) → gelu ops.GeLU()(x) → GeLU
isinf torch.isinf(x) → isinf mint.isinf(x) → IsInf
isnan torch.isnan(x) → isnan ops.isnan(x) → IsNan
log torch.log(x) → log mint.log(x) → Log
neg torch.neg(x) → neg mint.neg(x) → Neg
negative torch.negative(x) → negative mint.negative(x) → Neg
relu torch.relu(x) → relu ops.relu(x) → ReLU
rsqrt torch.rsqrt(x) → rsqrt mint.rsqrt(x) → Rsqrt
sigmoid torch.sigmoid(x) → sigmoid mint.sigmoid(x) → Sigmoid
silu torch.nn.functional.silu(x) → silu mint.nn.functional.silu(x) → SiLU
sin torch.sin(x) → sin mint.sin(x) → Sin
sqrt torch.sqrt(x) → sqrt mint.sqrt(x) → Sqrt
square torch.square(x) → square mint.square(x) → Square
add torch.add(x,y) → add ops.add(x,y) → Add; mint.add → AddExt
div torch.div(x,y) → div mint.div(x,y) → Div
mul torch.mul(x,y) → mul mint.mul(x,y) → Mul
pow torch.pow(x,y) → pow mint.pow(x,y) → Pow
real_div 没有公开接口 ops.RealDiv()(x,y) → RealDiv
sub torch.sub(x,y) → sub ops.sub → Sub; mint.sub → SubExt
__rsub__ 2.0 - x → __rsub__ 2.0 - x → Sub
__rpow__ 2.0 ** x → __rpow__ 2.0 ** x → Pow
true_divide torch.true_divide(x,y) → true_divide ops.true_divide(x,y) → Div
likedislike
fangwenyifangwenyi成员
2 天前 issue状态由 TODO 改变为 DONE
fangwenyifangwenyi成员
2 天前 关闭了 issue