【AscendC】MoeInitRoutingGroupedMatmulGrad 算子设计
作者:昇腾实战派
知识地图:https://blog.csdn.net/Lumos_Lovegood/article/details/161601003
1. 需求分析
1.1 背景
在MoE模型的训练过程中,进行首次反向(最后一层 layer)存在显存溢出风险。针对 moe_permute + moe_group_matmul 需要做反向融合算子,以节省显存,同时保证确定性运算。
在 MoE(Mixture of Experts)模型中,前向计算流程为:
- MoeInitRouting(语义等价于 Permute):根据路由结果将 token 分发到对应 expert,输出
expanded_x(按 expert 分组排列的 token)和路由索引expandedRowIdx - GroupedMatmul:对每组 expert 的 token 分别执行矩阵乘法
expanded_x @ weight → y
反向传播需要计算 d_x(对输入的梯度)和 d_weight(对权重的梯度)。非融合方案需要先计算 d_expanded_x = d_y @ weight^T,产生一个与 expanded_x 同形的中间张量([M, K],在大 token 数场景下可达数 GB),再通过 MoeInitRoutingGrad(scatter-add)归约到 d_x。这个中间张量是显存瓶颈。
本融合算子将两步反向合并在一个 kernel 中完成:Matmul 结果分块写入 workspace ring buffer,由 Vector 引擎立即消费并 scatter-add 到 d_x,不产生完整的 d_expanded_x 中间张量。
1.2 功能分析
- 正向链路:
MoeInitRouting(x, rowIdx, expertIdx)→expanded_x→GroupedMatmul(expanded_x, weight)→y - 反向链路(本算子):
d_y→GroupedMatmulGrad(d_y, weight^T)→d_expanded_x→MoeInitRoutingGrad(scatter_add)→d_x - 不计算
d_weight(本算子仅输出d_x)
为何不计算 d_weight 可规避显存峰值
d_weight 通过另外的反向算子单独生成,不在本算子中计算。这样设计可规避显存峰值,原因如下:
1. 避免 d_expanded_x 全量物化
非融合方案中,d_expanded_x = d_y @ weight^T 产出完整形状为 [M, K] 的中间张量。以 ZMoE 典型配置(M = 108K tokens, K = 7168)为例:
- FP16 下:
108K × 7168 × 2B ≈ 1.5 GB - FP32 下:
108K × 7168 × 4B ≈ 3.0 GB
本融合算子使用 workspace ring buffer 替代全量中间张量,ring buffer 大小仅为 parallNum × baseM × baseN × coreNum × sizeof(float)。以 parallNum=3, baseM=128, baseN=128, coreNum=32 为例:
3 × 128 × 128 × 32 × 4B ≈ 6 MB
节省约 1.5~3.0 GB 显存。
2. 避免 d_weight 与 d_expanded_x 同时驻留
如果在同一算子中同时计算 d_x 和 d_weight,需要同时持有:
d_y [M, N]— 上游梯度weight [G, K, N]— 权重矩阵d_expanded_x [M, K]— 中间结果(或 ring buffer)d_weight [G, K, N]— 权重梯度输出d_x [batch, K]— 输入梯度输出
此时 workspace 中同时存在 d_expanded_x 和 d_weight 两份大张量,显存峰值 = M×K + G×K×N。
以 G=60, K=7168, N=2048 为例:d_weight 占 60×7168×2048×2B ≈ 1.7 GB(FP16),加上 d_expanded_x 的 1.5 GB,仅这两个中间结果即达 3.2 GB。
3. 分离策略
将 d_weight 计算拆分到独立算子(通过 npu_grouped_matmul 的 split_item=3 路径,即 expanded_x^T @ d_y),该算子的输入 expanded_x 由前向保存(autograd 已持有),输出 d_weight 直接累加到权重梯度。分离后:
| 阶段 | 主要显存占用 | 峰值 |
|---|---|---|
| 本融合算子 (d_x) | d_y + weight + ring buffer(~6MB) + d_x | ~weight 大小 + 数 MB |
| 单独 d_weight 算子 | expanded_x + d_y + d_weight | ~expanded_x + d_y + d_weight |
两个算子的峰值在不同时间点出现,总峰值由各自的最大值决定,而非叠加,从而规避了 d_expanded_x + d_weight 同时驻留的显存尖峰。
数学定义和公式
记号约定:
| 符号 | 含义 | 形状 |
|---|---|---|
M |
expanded token 总数 =batch × topK |
标量 |
K |
特征维度(weight.shape[1]) | 标量 |
N |
隐藏维度(weight.shape[2]) | 标量 |
G |
expert 数量(group_num) | 标量 |
g(i) |
第 i 个 expanded token 所属的 expert ID | 标量 |
前向公式:
- expanded_x = Permute(x, routing_indices) # 形状: [M, K]
y[i, n] = Σ_k expanded_x[i, k] × weight[g(i), k, n] # 形状: [M, N]
反向核心公式:
给定上游梯度 d_y = ∂L/∂y,形状 [M, N]:
**Step 1 — Grouped Matmul 反向(对 expanded_x 的梯度)**:
- d_expanded_x[i, k] = Σ_n d_y[i, n] × weight[g(i), k, n]
即 d_expanded_x = d_y @ weight^T,其中 weight^T 对每个 expert 在最后两维做转置:weight^T[g, n, k] = weight[g, k, n]。
**Step 2 — MoeInitRouting 反向(Scatter-Add 归约)**:
- d_x[b, k] = Σ_{i ∈ expanded_tokens_of_batch(b)} d_expanded_x[i, k]
其中 expanded_tokens_of_batch(b) = { i | oriExpandedRowIdx[i] / topK == b },即 original_pos = oriExpandedRowIdx[i],batch_row = original_pos / topK。
逐元素梯度公式(合并 Step 1 和 Step 2):
- d_x[b, k] = Σ_{i: batch(i)=b} Σ_n d_y[i, n] × weight[g(i), k, n]
展开为标量形式:
- ∂L/∂x[b, k] = Σ_{i=0}^{M-1} Σ_{n=0}^{N-1} 1_{batch(i)=b} × ∂L/∂y[i, n] × weight[g(i), k, n]
其中 1_{batch(i)=b} 是指示函数,当 expanded 位置 i 映射回原始 batch 位置 b 时取 1,否则为 0。
输入输出
| 参数 | 角色 | 形状 | 数据类型 |
|---|---|---|---|
inputGradY |
上游梯度 d_y | [M, N]其中M = batch × topK |
FLOAT / FLOAT16 / BF16 |
inputWeight |
权重矩阵 | [G, K, N]其中G = group_num |
FLOAT / FLOAT16 / BF16 |
expandedRowIdx |
展开路由索引 | [L],batch ≤ L ≤ M |
INT64 |
groupList |
每组 token 数(非 cumsum) | [G] |
INT64 |
outputGradX |
输入梯度 d_x | [batch, K] |
FLOAT(固定 FP32) |
注:
expandedRowIdx的语义由rowIdxType控制:0表示 gather 类型(expanded→original),1表示 scatter 类型(original→expanded)。反向需要expanded→original映射。outputGradX固定为 FP32 输出,框架/模型侧再 cast 为 FP16/BF16,以保证 scatter-add 累加精度。- 本算子不输出
d_weight,d_weight由独立算子计算(expanded_x^T @ d_y)。
约束
| 约束项 | 约束条件 |
|---|---|
| gradY 维度 | 必须为 2D |
| weight 维度 | 必须为 3D |
| expandedRowIdx 维度 | 必须为 1D,长度满足batch ≤ L ≤ batch × topK,且L ≥ M |
| groupList 维度 | 必须为 1D,长度等于G(weight.shape[0]) |
| 维度一致性 | gradY.shape[1]必须等于weight.shape[2](即 N 维度一致) |
| groupList 语义 | Σ groupList[i]必须等于M(gradY.shape[0]),即每组 token 数之和等于总 token 数 |
| topK | 必须 > 0 |
| batch | 必须 > 0 |
| rowIdxType | 必须为 0 或 1 |
| groupList 单组 | groupList[i]可为 0(空 expert),kernel 会跳过空组 |
| 数据类型 | gradY 和 weight 的 dtype 必须相同(同属 {FLOAT, FLOAT16, BF16}) |
| 芯片支持 | ascend910b, ascend910_93 |
2 特性实现方案
2.1 算子整体计算流程图

2.2 aclnn 设计
2.2.1 执行场景
本算子仅支持 AICORE 路径,不存在 AICPU 与 AICORE 的场景区分。
PrecisionReduceFlag(true):开启精度还原,支持 FP16/BF16 输入下的高精度累加。DynamicCompileStaticFlag(true):开启动静态编译。DynamicShapeSupportFlag(true)/DynamicRankSupportFlag(true):支持动态 Shape 和动态 Rank。NeedCheckSupportFlag(false):不进行额外支持性检查。
2.2.2 入图支持
通过 ExtendCfgInfo("opFile.value", "moe_init_routing_grouped_matmul_grad") 指定 kernel 入口文件。算子通过 OP_ADD 注册到 GE 图,支持图编译。
aclnn 接口:
aclnnMoeInitRoutingGroupedMatmulGradGetWorkspaceSize(...)— 获取 workspace 大小aclnnMoeInitRoutingGroupedMatmulGrad(...)— 执行算子
算子入图,支持动态 Shape / 动态 Rank。
2.3 infershape 设计
2.3.1 输入 Shape 和 Dtype 校验
| 输入 | 维度数 | 各维约束 |
|---|---|---|
| gradY | 2 | [M, N],其中M = batch × topK,N 任意 |
| weight | 3 | [G, K, N],其中 N 必须与gradY.shape[1]一致 |
| expandedRowIdx | 1 | 长度L满足batch ≤ L ≤ M,且L ≥ M |
| groupList | 1 | 长度等于G(weight.shape[0]),且Σ groupList[i] = M |
2.3.2 输出 Shape 和 Dtype 推导
- outputGradX shape:
[batch, K],其中batch来自 attr(必须由用户显式传入),K = weight.shape[1] - outputGradX dtype:恒为
DT_FLOAT(FP32),与输入精度解耦
2.3.3 关键推导参数
从各输入 shape 中提取:
dm = gradY.shape[0]— M 维度(= batch × topK)dn = gradY.shape[1]— N 维度(= weight.shape[2])dk = weight.shape[1]— K 维度(= weight.shape[2] 对应 gradY.shape[1] 的匹配维度的”另一端”)gmmGroupNum = weight.shape[0]— 分组数(expert 数)
校验逻辑:
gradY.shape[1] == weight.shape[2](N 维度一致)groupList.shape[0] == weight.shape[0](G 维度一致)batch ≤ expandedRowIdx.shape[0] ≤ batch × topKexpandedRowIdx.shape[0] ≥ dmdm % batch == 0(用于反推 topK)
2.3.4 对 -1/-2 场景的适配
- **-1(UnknownShape)**:支持。
DynamicShapeSupportFlag(true)已启用,shape 在运行期确定。 - **-2(UnknownRank)**:支持。
DynamicRankSupportFlag(true)已启用,rank 在运行期确定。
infershape 在 compile 期(InferShapeMoeInitRoutingGroupedMatmulGrad)仅做连接校验,不做具体 shape 推导(输出 shape 依赖 attr batch,在 tiling 阶段完成)。
2.4 tiling 设计
2.4.1 分核策略
采用多核流水并行 + 确定性核间同步的分核策略:
核心思路:
- 将所有 Group 的所有 Block(按
baseM × baseN切分)展平为全局 Block 序列 - Block 按
coreNum轮转分配给各 AICore(curr_abs_block_num初始 =coreIdx,每次递增coreNum) - 每个 AICore 维护独立的 Cube 计算流水(
cubeCount计数),Cube 结果写入 workspace 的环形 slot(共parallNum × coreNum个 slot) - AIV 与 AIC 协同:AIC 写完 slot 后通知 AIV 搬运,AIV 消费完通知 AIC 复用 slot
2.4.2 维度泛化限制
| 维度 | 是否可泛化至 2^31 | 说明 |
|---|---|---|
| M (dm) | 是 | uint32_t 存储 |
| N (dn) | 是 | uint32_t 存储 |
| K (dk) | 是 | uint32_t 存储 |
| G (gmmGroupNum) | 是 | uint32_t 存储 |
| batch | 是 | uint32_t 存储,但M = batch × topK需 ≤ 2^31 |
| topK | 是 | uint32_t 存储 |
| expandedRowIdx 长度 | 是 | 通过totalInGroup引用,uint32_t 存储 |
注意:mid 的中间计算(如
cubeCount * coreNum)使用uint64_t防止溢出,最终截断到uint32_t与totalWids取 min。
2.4.3 TilingData 内存对齐
PostTiling 中强制校验 tilingData->GetDataSize() % sizeof(uint64_t) != 0,确保 tiling data 8 字节对齐。
Workspace 各区域按 512 字节对齐(tmp_offset = ((tmp_offset + 511) / 512) * 512)。
2.4.4 Tiling 变量表
| 变量名 | 语义 | 设计上下界 |
|---|---|---|
matmulTiling(TCubeTiling) |
Cube matmul 的精细化 tiling 参数(baseM/baseN/baseK/stepKa/stepKb/depthA1/depthB1 等),由MultiCoreMatmulTiling::GetTiling计算填充 |
由 matmul 库决定 |
irTilingData.batch |
原始 batch 大小,来自 attr | [1, ∞) |
irTilingData.topK |
每个 token 选择的 expert 数,来自 attr | [1, ∞) |
irTilingData.rowIdxType |
路由索引类型:0 = gather(expanded→original),1 = scatter(original→expanded) | {0, 1} |
coreNum |
参与计算的 AICore 总数,来自硬件平台信息 | 硬件决定(典型 32) |
vBaseM |
向量计算基 M 大小 =UBCALSIZE / baseN |
[1, baseM] |
ubCalSize |
UB 计算区大小 =16 × 256 = 4096(与GroupedMatmulFinalizeRouting一致) |
固定值 4096 |
parallNum |
每核并行深度(ring buffer slot 数)=CV_PARALL_NUM = 3 |
固定值 3(必须 ≥ 2) |
k |
Matmul 内维度 =dn(weight 最后一维 N) |
[1, 2^31) |
n |
Matmul 输出维度 =dk(weight 中间维 K) |
[1, 2^31) |
dm |
Matmul M 维度 =batch × topK |
[1, 2^31) |
dn |
gradY 第二维 = weight 第三维 N | [1, 2^31) |
dk |
weight 第二维 K | [1, 2^31) |
gmmGroupNum |
分组/专家数 G | [1, 2^31) |
batch |
原始 batch 大小 | [1, ∞) |
totalInGroup |
所有 group 的 token 数之和 =dm = M |
[1, 2^31) |
topk |
每个 token 选择的 expert 数 | [1, ∞) |
deterministicFlag |
确定性计算标志 = 1(固定) | 固定值 1 |
deterWorkspaceSize |
确定性 workspace 总字节数(含 MM 输出区 + 逆置换表区 + 系统区,均 512B 对齐) | [1, UINT64_MAX) |
deterWorkspaceMMOutOffsetPtr |
MM 输出区在 workspace 中的起始偏移 | 0 |
deterWorkspaceMMOutBytesSize |
MM 输出区字节数 =parallNum × baseM × baseN × coreNum × sizeof(float) |
[1, UINT32_MAX) |
deterWorkspaceOriExpandedRowIdxOffsetPtr |
逆置换表区起始偏移(512B 对齐) | [0, UINT64_MAX) |
deterWorkspaceSystemOffsetPtr |
系统 workspace 区起始偏移(512B 对齐) | [0, UINT64_MAX) |
baseM |
Matmul M 轴分块大小 =min(BEST_BASE_M, m),其中BEST_BASE_M = 128 |
[1, 128] |
baseN |
Matmul N 轴分块大小 =min(BEST_BASE_N, n),其中BEST_BASE_N = 128 |
[1, 128] |
baseK |
Matmul K 轴分块大小 =min(BEST_BASE_K, k),其中BEST_BASE_K = 128 |
[1, 128] |
2.5 片上内存资源设计
2.5.1 UB 空间分配
| 缓冲区 | 用途 | 大小 |
|---|---|---|
| Matmul 内部 L0A Buffer | Cube 左矩阵 A(gradY 块)缓存 | 由 Matmul 库内部管理 |
| Matmul 内部 L0B Buffer | Cube 右矩阵 B(weight^T 块)缓存 | 由 Matmul 库内部管理 |
| Matmul 内部 L0C Buffer | Cube 累加结果 C 缓存 | 由 Matmul 库内部管理 |
queBind(TQueBind) |
确定性同步数据搬运队列(VECIN→VECOUT),用于 copyOneRow 中 workspace→UB→gradX 的数据搬移 | BUFFER_NUM × DETER_UB_SIZE = 2 × 12KB = 24KB |
vecInQueue(TQue) |
Vector 输入队列 | 1 slot(按需分配) |
vecOutQueue(TQue) |
Vector 输出队列 | 1 slot(按需分配) |
expertIDxInQueue(TQue) |
VectorCalcOriExpandedRowIdx 中的 chunk 数据缓冲区 | 1 slot,最大ORI_EXPANDED_ROW_IDX_CHUNK_BYTES = 32KB |
2.5.2 L1 空间分配
| 缓冲区 | 用途 | 大小 |
|---|---|---|
| L1 A Matrix Buffer | 左矩阵 A 从 GM→L1 的 DMA buffer(stepKa × baseK 行块) | 由 Matmul 库管理(与baseM/baseK/stepKa相关) |
| L1 B Matrix Buffer | 右矩阵 B 从 GM→L1 的 DMA buffer(stepKb × baseK 行块) | 由 Matmul 库管理(与baseN/baseK/stepKb相关) |
| L1 C Buffer | Matmul 结果回写 buffer | 由 Matmul 库管理 |
L1 总大小通过 compileInfo->l1Size 获取(910B 上典型为 192KB/核),Matmul 库在 SetBufferSpace(l1Size, l0CSize, ubSize) 时内部分配。
2.5.3 LCM(Local Core Memory)分配
| 缓冲区 | 用途 | 大小 |
|---|---|---|
groupInfoAIVBuf(TBuf) |
AIV 使用的 group 前缀和信息:grouplistAIVLocal、grouplistcumsumAIVLocal、groupWidCumsumAIVLocal、groupBMIDCumsumAIVLocal |
4 × Ceil(G × sizeof(uint64_t), 32) × 32字节(G=60 时约 2KB) |
2.6 kernel 方案
2.6.1 Kernel 计算流程图

2.6.2 精度转换策略
精度策略要点:
- Cube 累加全精度:Matmul 的 C 类型固定为
float(FP32),不受输入精度影响。即 FP16/BF16 输入时,Cube 内部仍使用 FP32 累加,保证矩阵乘的计算精度。 - Scatter-Add 全精度:
gradX固定 FP32,AtomicAdd以 FP32 精度执行。当多个 expanded token 映射回同一个batch行时(topK > 1),多次 FP32 累加不会损失精度。相比 FP16 累加(尾数仅 10 bits),FP32(尾数 23 bits)可避免大数值与小梯度相加时的截断误差。 - 输出端 cast:FP32 输出由框架/模型侧根据需要 cast 为 FP16/BF16,保持与主模型精度一致。这比在 kernel 内部直接输出低精度更灵活且更安全。
- Tiling Key 按输入精度分流:通过
SetTilingKey按grad_y的 dtype 选择编译模板(ELEMENTWISE_TPL_SCH_MODE_FLOAT/FLOAT16/BF16),使不同精度场景使用不同的 Matmul 指令路径,同时保持 C 类型为 FP32。
4. 测试设计
4.1 测试标杆
- 非融合反向标杆:用
torch_npu.grouped_matmul(split_item=2, weight^T)+torch_npu.moe_token_permute_grad拼装的反向组合作为 ground truth。该组合等价于本融合算子的计算语义。 - Naive 标杆:用纯 PyTorch 小算子(
torch.mm+torch.scatter_add)手写反向,确保业务逻辑正确,排除 torch_npu 算子本身的实现差异。
4.2 测试用例设计
4.2.1 白盒用例
| 场景 | 描述 | 关键参数 |
|---|---|---|
| 单 token 单 expert | 最小功能验证:1 个 token 的 topK=1,1 个 expert | batch=1, K=128, N=256, topK=1, G=1, ratio=1.0 |
| 单 token 多 expert | topK>1 累加验证:同一 token 被多个 expert 处理,验证 scatter-add 归并 | batch=1, K=128, N=256, topK=3, G=3, ratio=1.0 |
| 多 token 单 expert | 所有 token 路由到同一 expert | batch=128, K=128, N=256, topK=1, G=1, ratio=1.0 |
| 多 token 多 expert | 常规场景:多 token 分散到多 expert | batch=128, K=128, N=256, topK=3, G=10, ratio=0.8 |
| 空 expert | groupList 中存在 mg=0 的 expert(完全不命中),验证空组跳过逻辑 | batch=10, K=128, N=256, topK=2, G=4, ratio=0.5,强制groupList[1]=0 |
| 尾块对齐 | M/K/N 不能被 baseM/baseN/baseK 整除,验证 curSingleM/curSingleN 尾块处理 | K=128+50, N=256+30, batch=10, topK=3, G=4 |
| rowIdxType=0 (gather) | expandedRowIdx 已是 expanded→original,直接复用 | 同常规用例 +rowIdxType=0 |
| rowIdxType=1 (scatter) | expandedRowIdx 为 original→expanded,先逆置换再使用 | 同常规用例 +rowIdxType=1 |
| 三种输入精度 | FP32 / FP16 / BF16 分别验证 | dtype ∈ {float32, float16, bfloat16} |
| 边界 M=1 | M 维度为 1(最小 batch×topK) | batch=1, topK=1, K=128, N=256, G=1 |
| 边界 N=1 | N 维度为 1 | batch=128, K=128, N=1, topK=3, G=4 |
| 边界 K=1 | K 维度为 1 | batch=128, K=1, N=256, topK=3, G=4 |
| 大 token 数 | 验证大 M 下 uint32 不溢出、chunk 分块正确 | batch=108×1024, K=128, N=256, topK=2, G=60 |
鲲鹏昇腾开发者社区是面向全社会开放的“联接全球计算开发者,聚合华为+生态”的社区,内容涵盖鲲鹏、昇腾资源,帮助开发者快速获取所需的知识、经验、软件、工具、算力,支撑开发者易学、好用、成功,成为核心开发者。
更多推荐



所有评论(0)