已合并
feat: add distributed op RotaryPositionEmbedding with UT and MindSpore ST #661
feat: add distributed op RotaryPositionEmbedding with UT and MindSpore ST #661
已合并
hedongdong创建于 5月19日
hedongdong成员
5月19日

What type of PR is this?

/kind feature


What does this PR do / why do we need it:

为 HyperParallel 接入 RotaryPositionEmbedding 算子的分布式调度支持:

  • RotaryPositionEmbedding 是 Transformer 模型中广泛使用的位置编码算子(公式:y = x * cos + x_rotate * sin),在分布式并行训练中需要正确推断 DTensor 的输出分布状态。

HyperParallel 的工作包括:

  1. DTensor 输入解包:通过 preprocess() 的 to_local() 提取本地 tensor 后传给 CANN 内核
  2. 输出 DTensor Layout 推断:通过 infer_layout() 基于输入分片方式推断输出 layout

分片规则

规则 说明
D(最后维度)必须 Replicate 内核在 D 维度内旋转,不可分割
B / N / S 维可自由切分 算子在各位置独立,可并行
cos/sin 广播语义 cos/sin 可比 x 少分片(Replicate 维度广播)
输出 layout = x layout 深拷贝 输出 shape 等于 x shape

测试覆盖

测试类型 文件 内容
UT test_parallel_rotary_position_embedding.py (22 cases) _normalize_rpe_args / YAML 注册 / infer_layout 正例(全 Replicate、B-DP、N-TP、2D/3D mesh、cos/sin 广播)/ 错误处理(D 维切分、cos/sin 切分不一致、Partial 输入)
ST rotary_position_embedding_shard_in_python.py (8 cases) Group1(2 卡):replicated + dp_b(fwd+bwd) + tp_n(fwd+bwd);Group2(4 卡):dp_tp(fwd+bwd) + dp_sp(fwd);Group3(4 卡):tp_sp(fwd) + dp_tp_cos_full(fwd);Group4(8 卡):dp_tp_sp(fwd)

Which issue(s) this PR fixes:

Fixes #155


Test Plan and Test result:What scenarios were tested, and what were the verification results(Function, performance, reliability, etc.):

UT(单元测试)

  • tests/ut/core/shard/ops/test_parallel_rotary_position_embedding.py — 22 cases,覆盖参数归一化、YAML 注册、infer_layout(全 Replicate / B-DP / N-TP / 2D-3D mesh / cos/sin 广播)、错误处理(D 维切分、cos/sin 切分不一致、Partial 输入)

ST(系统测试 - distributed, MindSpore)

  • tests/mindspore/st/shard/ops/test_rotary_position_embedding_shard_in_python.py — 4 组:
    • Group1(3 用例,2 卡):replicated / dp_b(fwd+bwd) / tp_n(fwd+bwd)
    • Group2(2 用例,4 卡):dp_tp(fwd+bwd) / dp_sp(fwd)
    • Group3(2 用例,4 卡):tp_sp(fwd) / dp_tp_cos_full(fwd)
    • Group4(1 用例,8 卡):dp_tp_sp(fwd)

验证结果

  • 所有 22 条 UT 通过(mock 平台,无 GPU/NPU 依赖)
  • 所有 4 组 ST 在 Ascend 910B NPU 上通过(float16 精度,atol/rtol=1e-3)
  • 反向验证通过:dp_b(grad_x Shard(0), grad_cos/sin Partial("sum"))、tp_n(全 Shard(1))、dp_tp(Shard(0),Shard(1))
  • 广播场景通过:cos/sin (1,1,S,D) 在 N-TP 和 S-SP 维度的 Replicate 广播正确推断
  • pylint 通过

Self-checklist:(请自检,在[ ]内打上x,我们将检视你的完成情况,否则会导致pr无法合入)

likedislike
Pull Request已成功合入, 合并人@MindSpore-Bot
(感谢 hedongdong 的贡献)
Hhedongdong成员
5月19日 创建了 pull request,commit aaafe6b1
Hhedongdong成员
5月19日 关联了issue:feat: 接入 RotaryPositionEmbedding 算子分布式调度支持
MindSpore-BotMindSpore-Bot成员
5月19日 添加了label:mindspore-cla/yes
MindSpore-Bot
MindSpore-Bot成员
5月19日 评论:

CLA Signature Pass

david-he91, thanks for your pull request. All authors of the commits have signed the CLA. 👍

likedislike
司小南(机器人)司小南(机器人)成员
5月19日 添加了label:pr-check-pass
MindSpore-BotMindSpore-Bot成员
5月19日 添加了label:no-pass-all-review
MindSpore-Bot
MindSpore-Bot成员
5月19日 评论:

Notice

The PR needs 2 assignees to review. 2 does not review. if all are passed, please comment /check-pr to try merge the PR. 😄

likedislike
hedongdong成员
5月19日 评论:

/retest

likedislike
司小南(机器人)
司小南(机器人)成员
5月19日 评论:

🔵 The pipeline #2570 is running. Please wait a moment... (Link)

likedislike
司小南(机器人)司小南(机器人)成员
5月19日 添加了label:ci-pipeline-running
liuchongming74liuchongming74成员
5月19日 通过审查
MindSpore-Bot
MindSpore-Bot成员
5月19日 评论:

Notice

The PR needs 2 assignees to review. 1 does not review. if all are passed, please comment /check-pr to try merge the PR. 😄

likedislike
司小南(机器人)
司小南(机器人)成员5月19日进行代码检视1
hyper_parallel/core/shard/ops/parallel_rotary_position_embedding.py
@@ -0,0 +96,4 @@
96+ f"For {op}, D (last dim) of {name} must be replicated, "
97+ f"but got tensor_map={tm}"
98+ )
99+ for d in range(len(tm) - 1):
司小南(机器人)
司小南(机器人)5月19日评论:

🚨 [关键问题] cos/sin 分片验证逻辑不完整,可能漏检非法分片

问题陈述:
cos/sin 分片参数的校验逻辑存在漏洞,仅验证了分片参数的基础有效性(如非负或非空),却遗漏了对分片完整性(Coverage)和与输入张量维度一致性(Consistency)的校验,导致非法分片配置可能绕过检查进入计算流程。

证据支撑:

  1. 推断校验逻辑处(如 rotary_embedding_kernel.cpp 或相关校验函数中)仅包含形如 assert(split_size > 0) 的基础检查,未对分片总和进行约束
  2. 缺失关键校验逻辑:assert(accumulate(split_sizes) == input_dim) 以及 assert(cos.shape == expected_split_shape)
  3. 问题影响:若分片总和小于输入维度,部分数据将被静默忽略,导致计算结果错误;若分片总和大于输入维度,将引发内存越界访问;若 cos/sin 形状与分片不匹配,将导致计算结果非预期

修改建议:
完善校验逻辑,增加分片总和与输入维度的对齐检查,以及 cos/sin 张量形状与分片配置的匹配检查:

// 修复建议示例
// 1. 校验分片总和覆盖整个维度
int total_split = std::accumulate(split_sizes.begin(), split_sizes.end(), 0);
CHECK_EQ(total_split, hidden_dim) << "Split sizes must sum to hidden dimension (" << hidden_dim << ")";

// 2. 校验 cos/sin 张量形状与分片配置匹配
for (size_t i = 0; i < split_sizes.size(); ++i) {
    CHECK_EQ(cos_tensors[i].size(0), split_sizes[i]) << "Cos tensor " << i << " size mismatch";
}
likedislike
司小南(机器人)
司小南(机器人)成员
5月19日 评论:

📖 Python注释检查:

根据您提供的代码差异,我仔细审查了其中的英文文档修改,发现以下不符合Python注释格式规范的地方:

1.💡问题描述:Args部分参数格式缺失类型说明,且参数mode有默认值但未按规范标注optional和Default。

  • 📄 文件: hyper_parallel/core/shard/ops/parallel_rotary_position_embedding.py
  • ✍ 修改建议:
    -     Args:
    -         x: Input tensor.
    -         cos: Cosine position encoding tensor.
    -         sin: Sine position encoding tensor.
    -         mode: Rotation mode. 0=rotate_half, 1=rotate_interleaved, 2=quarter,
    -             3=interleave-half. Defaults to 0.
    +     Args:
    +         x (Tensor): Input tensor.
    +         cos (Tensor): Cosine position encoding tensor.
    +         sin (Tensor): Sine position encoding tensor.
    +         mode (int, optional): Rotation mode. 0=rotate_half, 1=rotate_interleaved, 2=quarter,
    +             3=interleave-half. Default: 0.
    

2.💡问题描述:Args部分参数格式缺失类型说明。

  • 📄 文件: hyper_parallel/core/shard/ops/parallel_rotary_position_embedding.py
  • ✍ 修改建议:
    -     Args:
    -         x_layout: Layout of the x tensor.
    -         cos_layout: Layout of the cos tensor.
    -         sin_layout: Layout of the sin tensor.
    +     Args:
    +         x_layout (Layout): Layout of the x tensor.
    +         cos_layout (Layout): Layout of the cos tensor.
    +         sin_layout (Layout): Layout of the sin tensor.
    

3.💡问题描述:Args部分参数格式缺失类型说明。

  • 📄 文件: hyper_parallel/core/shard/ops/parallel_rotary_position_embedding.py
  • ✍ 修改建议:
    -     Args:
    -         args: Positional arguments, may include DTensors.
    -         kwargs: Keyword arguments.
    +     Args:
    +         args (tuple): Positional arguments, may include DTensors.
    +         kwargs (dict): Keyword arguments.
    

4.💡问题描述:Args部分参数格式缺失类型说明。

  • 📄 文件: hyper_parallel/core/shard/ops/parallel_rotary_position_embedding.py
  • ✍ 修改建议:
    -     Args:
    -         cache_values: [x_layout, cos_layout, sin_layout]
    +     Args:
    +         cache_values (list): [x_layout, cos_layout, sin_layout]
    

5.💡问题描述:Args部分参数格式缺失类型说明。

  • 📄 文件: tests/mindspore/st/shard/ops/rotary_position_embedding_shard_in_python.py
  • ✍ 修改建议:
    -     Args:
    -         mode: Rotation mode (0=rotate_half, 1=rotate_interleaved,
    -               2=quarter, 3=interleave-half).
    +     Args:
    +         mode (int): Rotation mode (0=rotate_half, 1=rotate_interleaved,
    +               2=quarter, 3=interleave-half).
    

6.💡问题描述:Args部分参数格式缺失类型说明,且参数mode有默认值但未按规范标注optional和Default。

  • 📄 文件: tests/mindspore/st/shard/ops/rotary_position_embedding_shard_in_python.py
  • ✍ 修改建议:
    -     Args:
    -         x_np: float32 ndarray, shape (B, N, S, D).
    -         cos_np: float32 ndarray, shape (B, N, S, D) or (1, 1, S, D).
    -         sin_np: float32 ndarray, same shape as cos_np.
    -         mode: RPE rotation mode.
    +     Args:
    +         x_np (np.ndarray): float32 ndarray, shape (B, N, S, D).
    +         cos_np (np.ndarray): float32 ndarray, shape (B, N, S, D) or (1, 1, S, D).
    +         sin_np (np.ndarray): float32 ndarray, same shape as cos_np.
    +         mode (int, optional): RPE rotation mode. Default: 0.
    

7.💡问题描述:Args部分参数格式缺失类型说明,且参数mode有默认值但未按规范标注optional和Default。

  • 📄 文件: tests/mindspore/st/shard/ops/rotary_position_embedding_shard_in_python.py
  • ✍ 修改建议:
    -     Args:
    -         x_np: float32 ndarray, shape (B, N, S, D).
    -         cos_np: float32 ndarray, shape (B, N, S, D) or (1, 1, S, D).
    -         sin_np: float32 ndarray, same shape as cos_np.
    -         mode: RPE rotation mode (must be 0 or 1).
    +     Args:
    +         x_np (np.ndarray): float32 ndarray, shape (B, N, S, D).
    +         cos_np (np.ndarray): float32 ndarray, shape (B, N, S, D) or (1, 1, S, D).
    +         sin_np (np.ndarray): float32 ndarray, same shape as cos_np.
    +         mode (int, optional): RPE rotation mode (must be 0 or 1). Default: 0.
    

8.💡问题描述:Args部分参数格式缺失类型说明。

  • 📄 文件: tests/mindspore/st/shard/ops/rotary_position_embedding_shard_in_python.py
  • ✍ 修改建议:
    -     Args:
    -         d_inputs: Tuple of DTensor inputs (dx, dcos, dsin).
    +     Args:
    +         d_inputs (tuple): Tuple of DTensor inputs (dx, dcos, dsin).
    

9.💡问题描述:Args部分参数格式缺失类型说明。

  • 📄 文件: tests/mindspore/st/shard/ops/rotary_position_embedding_shard_in_python.py
  • ✍ 修改建议:
    -     Args:
    -         raw_grads: Tuple of local gradient tensors from ms.grad.
    -         d_inputs: Tuple of DTensor inputs (dx, dcos, dsin).
    -         ref_grads: Tuple of float32 reference gradient ndarrays.
    -         tag: Human-readable label for assertion messages.
    +     Args:
    +         raw_grads (tuple): Tuple of local gradient tensors from ms.grad.
    +         d_inputs (tuple): Tuple of DTensor inputs (dx, dcos, dsin).
    +         ref_grads (tuple): Tuple of float32 reference gradient ndarrays.
    +         tag (str): Human-readable label for assertion messages.
    

🤗 上述内容仅供参考。

likedislike
司小南(机器人)司小南(机器人)成员
5月19日 添加了label:ai-reviewed
司小南(机器人)
司小南(机器人)成员
5月19日 评论:
Project Name Build_Stage Build Result Details(Full/Incremental) Detail
Coverage Line Coverage ✅ SUCCESS 40%/0% Link
likedislike
司小南(机器人)
司小南(机器人)成员
5月19日 评论:
Project Name Build_Stage Build Result Details
Hyper-parallel_Atomgit_Gate - ✅ SUCCESS 2570
likedislike
司小南(机器人)司小南(机器人)成员
5月19日 删除了label:ci-pipeline-running
司小南(机器人)司小南(机器人)成员
5月19日 添加了label:ci-pipeline-passed
Yyangzhenzhang成员
5月19日 通过审查
MindSpore-BotMindSpore-Bot成员
5月19日 删除了label:no-pass-all-review
MindSpore-BotMindSpore-Bot成员
5月19日 合入了pull request,合并节点 SHA:f417a67a046efda75ca836e63070fc14355a71b0