mindspore-ai / mindspore-ai/hyper-parallel
【RFC】HyperParallel Multicore HyperMegaMoE 正反向融合算子
Nobody has claimed this yet.
- Dominant language
- Python
- Stars
- 53
- Forks
- 63
- Avg merge
- 23h 45m
- Merged PRs (30d)
- 63
Description
1 背景与目标
MoE 训练中的 Expert Parallel 计算链路同时包含 AllToAll 通信、Grouped MatMul、激活计算和路由归并。若各阶段按独立 kernel 串行执行,会导致通信暴露、AIC/AIV 资源交替空闲,以及中间结果的频繁读写。
本期计划交付面向 Multicore 场景的完整 MoE-FFN 融合算子 HyperMegaMoE,目标是:
- 以单个融合算子完成 Dispatch、GMM1、SwiGLU、GMM2 和 Combine 的正向计算;
- 提供与正向语义对应的融合反向,完成激活梯度、专家权重梯度与反向通信;
- 封装为
DFunction自动求导接口,通过 HyperParallel 自定义算子 API 对外提供; - 用户仅传入 MoE 计算所需的数据、路由结果和专家权重,调度、通信缓冲区和运行时资源由接口内部管理;
- 在保证数值正确性的前提下,提升 AIC/AIV 并行度和通信—计算掩盖效果。
2 交付范围
2.1 正向融合
正向算子的逻辑名称为 HyperMegaMoE,ACLNN 执行接口为 aclnnHyperMegaMoE。
融合计算链路为:
AllToAll-Dispatch (AIV)
↓
GMM1 / up_proj (AIC)
↓
SwiGLU (AIV)
↓
GMM2 / down_proj (AIC)
↓
AllToAll-Combine (AIV)
各阶段按 token/expert tile 拆分为任务,由 AIC worker 和 AIV worker 并发执行。任务之间通过事件计数表达生产者—消费者依赖,满足数据就绪条件后即可放行下游任务,不引入全局屏障。
正向对外输出最终 token 结果 y,shape 与输入 x 一致。中间的 dispatch buffer、GMM 输出、SwiGLU 输出和 combine buffer 不作为公共返回值。
2.2 反向融合
反向算子的逻辑名称为 HyperMegaMoEGrad,ACLNN 执行接口为 aclnnHyperMegaMoEGrad。
反向计算链路为:
AllToAll-Dispatch(dY) (AIV)
↓
act_grad: dHidden = dispatched_dY × W2ᵀ (AIC)
│
├── SwiGLUGrad (AIV) ──→ input_grad GMM (AIC)
│
└── W2Grad GMM (AIC)
AllToAll-Combine(dX) (AIV) ∥ W1Grad GMM (AIC)
反向必须生成:
- 输入 token 梯度
dx; up_proj_weight梯度d_up_proj_weight;down_proj_weight梯度d_down_proj_weight;- 路由权重参与加权归并时的
d_topk_weights; - 其他不可导输入返回
None。
W2Grad 与激活梯度分支可并行执行,W1Grad 与 Combine(dX) 可并行执行,以充分利用通信和计算之间的独立性。
2.3 命名约定
| 层级 | 正向 | 反向 |
|---|---|---|
| 逻辑算子名 | HyperMegaMoE |
HyperMegaMoEGrad |
| Workspace 查询 | aclnnHyperMegaMoEGetWorkspaceSize |
aclnnHyperMegaMoEGradGetWorkspaceSize |
| ACLNN 执行接口 | aclnnHyperMegaMoE |
aclnnHyperMegaMoEGrad |
| DFunction | HyperMegaMoEDFunction |
由 backward 自动调用 HyperMegaMoEGrad |
| Python 公共接口 | hyper_mega_moe |
无需用户显式调用 |
3 用户接口
3.1 Python API
通过 HyperParallel 自定义算子接口导出:
from hyper_parallel.custom_ops.experimental import hyper_mega_moe
y = hyper_mega_moe(
x,
topk_ids,
topk_weights,
up_proj_weight,
down_proj_weight,
)
hyper_mega_moe 保持完整 MoE 语义:路由预处理、offset 构造、unpermute 和 top-k 加权归并均由接口内部完成。基础多核算子的融合边界为 Dispatch → GMM1 → SwiGLU → GMM2 → Combine;Permute/Unpermute 进一步融入核心 kernel 属于可选性能增量,不改变公共签名。
公共接口仅暴露以下必要参数:
| 参数 | 类型 | 建议 shape | 语义 |
|---|---|---|---|
x |
Tensor / DTensor | [T, H] |
本 rank 的未展开 token |
topk_ids |
Tensor / DTensor | [T, K] |
每个 token 的 top-k 专家索引,不参与求导 |
topk_weights |
Tensor / DTensor | [T, K] |
每个 token 的 top-k 路由权重 |
up_proj_weight |
Tensor / DTensor | [E_local, H, 2I] |
本 rank 的 GMM1/SwiGLU 专家权重 |
down_proj_weight |
Tensor / DTensor | [E_local, I, H] |
本 rank 的 GMM2 专家权重 |
返回值:
y:完成 top-k 加权归并后的 token 输出,shape 为[T, H]。
3.2 DFunction 语义
HyperMegaMoEDFunction 封装完整的正反向语义:
forward调用HyperMegaMoE,并仅保存反向必需的激活、路由信息和权重;backward调用HyperMegaMoEGrad,向自动求导系统返回与公共参数一一对应的梯度;- 输入为普通 Tensor 时执行本地路径;
- 输入为 DTensor 时,由 DFunction 的分布式调度机制完成 local tensor 提取、layout 推导和输出封装;
- 输出
y的 DeviceMesh 与 token 维布局与x保持一致; - 专家权重及其梯度沿 Expert Parallel 维分片。
3.3 内部管理的运行时资源
下列内容属于算子实现和运行时资源,不进入 Python 公共签名:
- rank id、EP world size、本地专家数、hidden size 和 sequence/token 数;
- dispatch/combine 的 offset、size 和 group list;
- 对称通信 buffer 及其生命周期;
- 输出与中间张量 buffer;
- event counter 及事件触发阈值;
- RuntimeConfig、TaskDesc 和任务队列;
- GMM/SwiGLU tiling 数据;
- ACLNN workspace 和子算子 workspace。
上述资源根据 Tensor/DTensor 的 shape、dtype、DeviceMesh、Placements 和路由结果生成或复用,并保证多次迭代时的事件状态正确重置。
4 功能语义
4.1 路由与专家并行
topk_ids使用全局专家编号;- 专家按 Expert Parallel mesh 均匀分布到各 rank;
- 每个 token 按
topk_ids展开为 K 个 expert slot; - Dispatch 将 expert slot 发送到专家所在 rank;
- 本地专家执行
GMM1 → SwiGLU → GMM2; - Combine 将 expert 输出发回 token 所在 rank,并按
topk_weights完成加权归并。
4.2 数值与边界约束
- 激活函数固定为 SwiGLU;
topk_ids的取值必须落在全局专家编号范围内;topk_weights的 token 维和 top-k 维必须与topk_ids一致;up_proj_weight的输出维必须等于2I,down_proj_weight的输入维必须等于I;- 各 rank 的专家分片数、权重 shape 和通信域配置必须一致;
- 支持非均匀路由、空专家、热点专家和尾块 token;
- 输入不满足 shape、dtype、layout、容量或分布式约束时,在 host/Python 边界返回明确异常,不进入 device kernel。
5 多核调度与性能目标
5.1 AIC/AIV 并发
- AllToAll 通信任务和 SwiGLU 任务映射到 AIV;
- GMM1、GMM2 及反向 GMM 任务映射到 AIC;
- AIC/AIV worker 独立轮询任务队列,按事件依赖放行;
- 任务粒度同时考虑首块延迟、专家 token 不均衡和 GMM 计算效率;
- 正向和反向均保持单次融合调用语义。
5.2 通信—计算重叠
- Dispatch 与已就绪专家的 GMM1 重叠;
- GMM1 与已就绪 tile 的 SwiGLU 重叠;
- GMM2 完成后尽早触发 Combine;
- 反向 W2Grad 与 act_grad 分支交错执行;
- 反向 W1Grad 与 Combine(dX) 重叠;
- 通信顺序避免多 rank 同时集中访问同一目标 rank。
5.3 运行时要求
- 任务配置与 tiling 可按 shape/dtype/topology 缓存并复用;
- 动态路由信息在每次调用时正确更新;
- 所有异步通信和跨流数据访问具备明确的完成或同步语义;
- 多次迭代、连续正反向和异常恢复后不得复用过期事件状态;
- 缓冲区容量与 workspace 大小由 host 侧计算并校验。
6 可选增量能力
以下能力不作为 HyperMegaMoE 基础交付的阻断项,可根据本期资源与性能收益选择纳入。
6.1 共享专家融合 MoE
支持可选共享专家权重,将共享专家的 GMM1、SwiGLU、GMM2 与路由专家流水协同调度,并在 Combine 阶段合并共享专家输出。公共接口仅在启用该能力时增加 keyword-only 共享专家权重参数。
6.2 Permute/Unpermute 融合扩展
将 top-k token 展开、按专家排序、路由计数、反排序与 top-k 加权归并进一步融入 HyperMegaMoE 执行流程,减少额外 kernel launch 和中间 token 读写。扩展前后保持 hyper_mega_moe 公共签名与数学语义不变。
6.3 GMM 专家计算分组粒度优化
将 GMM 就绪与调度粒度从单一专家级拓展为 expert × block,根据专家实际 token 数选择 block M。小批次或长尾专家使用较小 block 缩短首块等待,大批次使用较大 block 保持 GMM 效率。
6.4 Token Wave 优化
将专家 token 分成多个 wave,为每个 expert × wave 设置独立就绪事件。前缀 wave 到齐后立即启动对应 GMM M-tile,并将局部完成信号逐级传递到 SwiGLU、GMM2 和 Combine。尾 wave 的就绪阈值使用真实 token 数,避免等待不存在的补齐 token。
7 正确性与测试
7.1 功能测试
- 正向结果与
AllToAll + GroupedMatMul + SwiGLU + GroupedMatMul + AllToAll + weighted combine参考实现对齐; - 反向
dx、d_topk_weights、d_up_proj_weight和d_down_proj_weight与参考实现对齐; - 覆盖普通 Tensor 和 DTensor 输入;
- 覆盖 MindSpore 与 PyTorch 的正向、自动反向和多输入 layout 组合;
- 覆盖多次迭代、连续多次 forward 后 backward、梯度累积和参数更新;
- BF16 场景建议使用
rtol=atol=1e-3作为基础精度阈值,并在测试记录中说明 shape 和路由分布。
7.2 路由与边界测试
- 均匀路由与非均匀路由;
- 空专家、单热点专家和多热点专家;
- token 数不整除 tile/wave 大小的尾块;
- 不同 T/H/I/E/K 组合与至少两种 EP 规模;
- 非法专家索引、shape 不匹配、dtype 不支持、非法 layout 和缓冲区越界;
- 可选增量被纳入时,分别增加对应的正反向和边界用例。
7.3 稳定性测试
- 连续迭代不出现死锁、事件残留或跨迭代数据污染;
- 异常输入引起的失败可诊断,不使通信域进入不可恢复状态;
- 中间激活和 workspace 生命周期正确,不产生越界、野指针或显存持续增长;
- 异步通信输出在被 AIC/AIV 读取前已完成所需的流间同步。
8 性能验收
性能验收使用固定的硬件、CANN 版本、dtype、T/H/I/E/K、EP 规模和路由分布,并同时报告正向、反向与端到端训练步耗时。
必须提供以下数据:
- HyperMegaMoE 与等价非融合参考链路的耗时对比;
- Dispatch/Combine 通信暴露时间;
- AIC 与 AIV 的并发区间和利用率;
- 均匀路由、非均匀路由和热点专家场景的结果;
- 首轮与稳态迭代数据,明确是否包含编译、tiling 或缓存构建开销;
- 若纳入可选增量,单独报告增量前后的收益和额外显存成本。
性能结论不跨硬件、精度、shape 或路由分布做直接排名。
9 文档与示例
交付内容包含:
hyper_mega_moeAPI 文档,说明参数、shape、dtype、DTensor layout、返回值和异常;- DFunction 正反向语义及梯度列表;
- 单卡 Tensor 与多卡 DTensor 的最小可运行示例;
- 专家并行 DeviceMesh/Placements 配置示例;
- 精度、性能和边界范围说明;
- 可选增量的开关、额外参数和支持范围(如纳入)。
10 验收标准
- 正向算子命名为
HyperMegaMoE,并提供aclnnHyperMegaMoEGetWorkspaceSize/aclnnHyperMegaMoE; - 反向算子命名为
HyperMegaMoEGrad,并提供aclnnHyperMegaMoEGradGetWorkspaceSize/aclnnHyperMegaMoEGrad; -
HyperMegaMoEDFunction封装正向与反向,用户只需调用hyper_mega_moe; - Python 公共接口仅暴露
x、topk_ids、topk_weights、up_proj_weight和down_proj_weight; - 运行时配置、tiling、workspace、通信 buffer、offset/size 和 event counter 由内部管理;
- 正向覆盖 Dispatch → GMM1 → SwiGLU → GMM2 → Combine;
- 反向正确返回
dx、d_topk_weights、d_up_proj_weight和d_down_proj_weight; - Tensor 与 DTensor 调用语义一致,输出 layout 推导正确;
- MindSpore 与 PyTorch 接口语义一致,正反向用例通过;
- 均匀/非均匀路由、空专家、热点专家、尾块与多迭代用例通过;
- 提供可复现的精度与性能报告;
- API 文档、使用示例和支持边界完整。
11 非本期必须项
以下内容只在选入可选增量时纳入对应验收:
- 共享专家融合 MoE;
- Permute/Unpermute 进一步融入核心算子;
expert × block的动态 GMM 分组粒度;expert × wave的 token wave 就绪与流水。
schema_version: 1
source: gitcode
gitcode_repo: mindspore/hyper-parallel
gitcode_issue: 312
source_url: https://gitcode.com/mindspore/hyper-parallel/issues/312
Contributor guide
No contributing guide indexed for this repository
First steps
- Read the whole issue, then the project's contributing guide.
- Comment on the issue to say you are picking it up — it saves two people doing the same work.
- Fork the repository and make your change on a branch.
- Open a pull request that references the issue number.
Research direction
Start at the documented hyper_parallel.custom_ops.experimental.hyper_mega_moe entry point and trace how Tensor and DTensor inputs, expert-parallel layouts, and automatic differentiation are exposed. Done means the forward and backward interfaces, routing and boundary cases, cross-framework tests, performance measurements, documentation, and examples described in the acceptance checklist are complete.
Written by the indexing model from the issue text.
Assessment
- Tech stack
- python, pytorch
- Domain
- backend-api-design, distributed-systems, machine-learning
- Issue type
- Feature
- Difficulty
- 5/5
- Estimated time
- Over a week
- Activity status
- Active
- Clarity
- Mostly clear
- Newbie friendliness
- 20/100