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


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


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


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


问题总结:
问题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模块。


已收敛为指定的 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 |


HP DTensor 非均匀切分(RaggedShard)设计文档
0. 基本信息
RaggedShard非均匀连续切分hp_ragged_shardReplicate1. 背景
普通
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。该能力不能只通过增加一个 Placement 类完成,还需要贯通:
相关资料:
2. 本期目标与非目标
2.1 本期目标
RaggedShard(dims, local_units),并提供完整校验、相等性、hash 和字符串表达。Layout无损保留 Ragged placement,同时为旧 tensor-map 流程提供RaggedShard -> Replicate的 normal view。full_tensor()和以下重排:local_units变化时使用变长 all-to-all;dims或 ragged mesh 维变化时经 Replicate 中转。2.2 本期非目标
RaggedShard。dims,例如(1,)、(0, 2)。Shard或Partial。_StridedRaggedShard,以及同一逻辑维度上嵌套 Ragged/Shard 的顺序表达。DTensor.empty/full/rand(..., RaggedShard(...))。3. RaggedShard 语义
3.1 Placement 约束
RaggedShard(dims, local_units)当前满足:dims必须是非空tuple[int, ...]。dims == tuple(range(len(dims))),即必须是连续前缀维度。local_units必须是非空tuple[int, ...]。sum(local_units) > 0。len(local_units)必须等于 Ragged 所在 mesh 维的大小。RaggedShard。Replicate()。其中,Placement 构造阶段校验 tuple、类型、prefix 和 unit 非负性;依赖 global shape/mesh 的几何校验在 DTensor 构造或切分阶段完成。
3.2 Flat interval 计算
设:
计算公式:
要求
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))[0, 8)[0, 64)(64,)[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 架构与数据流
设计原则是:Ragged 元数据由原始 placements 无损保存;旧流程看到的 normal view 是
Replicate;只有实际切分、通信和 checkpoint 几何进入 Ragged 专用逻辑。4.2 Placement 与 Layout
代码位置:
hyper_parallel/core/dtensor/placement_types.pyhyper_parallel/core/dtensor/layout.pyLayout保存三种视图:placementsRaggedShard(dims, local_units)ragged_shardRaggedShardInfo(mesh_dim, placement),用于快速识别normal_placementsReplicate(),供旧 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.pyhyper_parallel/core/dtensor/_ragged_utils.pyRagged DTensor 的不变量:
DTensor.from_local()在 Ragged 场景必须显式传入shape:dt = DTensor.from_local( local_flat, mesh, (RaggedShard((0, 1), (1, 2)),), shape=(6, 4, 8), )构造时校验:
普通 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)调用链:
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:split 的统一语义是 dim 0 行数;Ragged 重排传入一维 tensor,因此行数等同于元素数。
PyTorch
dist.all_gather(),按真实长度预分配 listtorch_npu.distributed.reduce_scatter_tensor_uneven()all_gather(),再 trimall_reduce(),再截取本 rank 区间torch.distributed.nn.functional.all_to_all_single()MindSpore
ops.AllGatherVAllGatherV自动微分comm_func.all_to_all_single(),split 为 dim 0 行数变长 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 分支只有四种状态转换:
伪代码:
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 的交集长度:
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 包装输出。
当前白名单:
特殊场景:
当前实现不预先校验所有 Ragged 输入的 layout 是否相同,也不重复实现广播/shape 校验;白名单命中后交给底层本地算子执行,输出构造和后续流程在不满足 Ragged 不变量时继续报错。
以下算子仍 fail-closed:
mean/sum;4.8 Distributed Checkpoint
代码位置:
hyper_parallel/core/distributed_checkpoint/ragged_utils.pyhyper_parallel/core/distributed_checkpoint/standard_planner.pyhyper_parallel/core/distributed_checkpoint/filesystem_storage.pyhyper_parallel/core/distributed_checkpoint/util.pyhyper_parallel/core/distributed_checkpoint/async_staging.py4.8.1 几何适配
DCP 原有 reshard 基于标准 N-D
ChunkStorageMetadata(offsets, sizes)求交。Ragged local storage 虽然是一维连续区间,但该区间可能跨越 N-D 行或平面边界,因此保存前将 flat interval 分解为一组有序 N-D boxes:每个 box 对应一个普通
WriteItem。真正的数据仍来自:不需要先将 Ragged tensor 通信成 Replicate。
4.8.2 保存调用链
同一个逻辑 FQN 可能在同一 safetensors 文件内对应多个 box。
FileSystemWriter为其生成唯一物理 key: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 的交集。因此可复用现有 DCP 流程支持:
local_units;dims;异步 staging 重建 Ragged DTensor 时显式传入
shape=tuple(obj.shape)。第一版只要 state_dict 中存在 Ragged DTensor,就禁用 SavePlan cache,避免不同 global shape、dims 或 local_units 复用错误计划。更完整的文件结构和 metadata 示例见 dcp流程.md。
4.9 缓存与兼容策略
RaggedShard,其 hash 包含dims/local_units。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. 当前支持矩阵
DTensor.from_localdistribute_tensor(src_data_rank=None)distribute_tensor(src_data_rank=int)full_tensor()reduce_partial();不支持反向转换到 Partial_compute_ragged_slice()明确报错NotImplementedError_StridedRaggedShard“支持”表示实现路径已补齐;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 覆盖:
warning 为环境中的
torch_npu TypedStorage弃用提示,与 Ragged 功能无关。9. 验收标准
9.1 功能验收
RaggedShard校验、表达和公共导入正确。src_data_rank=None/int创建流程均正确。full_tensor()恢复原始全局 tensor。9.2 兼容性验收
tensor_key时仍可读取。9.3 明确报错
以下场景必须 fail-closed,而不是静默按普通 Shard/Replicate 处理: