已开启
feat: mamba3-mimo-bwd-fwd-triton-kernel #101
feat: mamba3-mimo-bwd-fwd-triton-kernel #101
已开启
bitszh3271创建于 7月16日
bitszh3271成员
7月16日

What this PR does / why we need it?

This PR adds the Triton-Ascend implementation of the first pass of the Mamba-3 MIMO combined backward computation (mamba3_mimo_bwd_fwd_kernel). Related issue: #38. https://gitcode.com/Ascend/MindSpeed-Ops/issues/38

This pass recomputes the chunk-wise forward intermediates, calculates the Phi/Zeta projection gradients and optional gate-input gradient, and produces states and qk_dot caches for the subsequent
mamba3_mimo_bwd_bwd_kernel.

The PR includes:

  • A PyTorch small-op reference implementation used for correctness verification.
  • The original stage0 Triton implementation as the optimization baseline.
  • An optimized Triton-Ascend implementation for arch32.
  • Shape-adaptive dispatch between the staged scan and batch/chunk-parallel auxiliary implementation.
  • Rank-fold scan, large-tile N-blocking, fused qk_dot generation, and rank-loop de-unrolling.
  • Host-side padding and output slicing for sequence lengths not divisible by chunk_size.
  • Unit tests and ATK test configuration.
  • Migration, optimization, and Triton implementation documentation.

On the upstream 11-shape grid, the optimized implementation achieves a geometric mean speedup of 1.94x over the stage0 Triton baseline.

Does this PR introduce any user-facing change?

Yes. It introduces the following public API:

from mindspeed_ops.api.triton.mamba3_mimo_bwd_fwd_kernel import (
    mamba3_mimo_bwd_fwd_kernel,
)

The API returns:

(states, qk_dot, dmimo_o, dmimo_z, dz)

It supports fp16, bf16, and fp32 inputs on Ascend arch32, including GQA, optional Z/SiLU gating, optional D-skip, general MIMO rank, and non-aligned sequence tails.

The current public API does not support arch35, varlen input through cu_seqlens, fused pregate RMSNorm, or returning the final state.

Related documentation:

  • docs/triton/mamba3_mimo_bwd_fwd.md
  • docs/triton/mamba3_mimo_bwd_fwd_optimization.md
  • docs/triton/mamba3_mimo_bwd_fwd_practice.md

How was this patch tested?

The implementation was verified on Ascend 910B with:

pytest -q tests/unit_tests/triton/test_mamba3_mimo_bwd_fwd_kernel.py

Result:

19 passed, 11 skipped

The skipped cases are the environment-gated full B=4, S=2048, H=16 shape grid. The default test set includes:

  • fp16, bf16, and fp32.
  • Z gate enabled/disabled.
  • D-skip enabled/disabled.
  • GQA configurations.
  • Five representative upstream Mamba-3 shapes.
  • Ranks 1, 2, and 4.
  • Large N/P tile paths.
  • Sequence lengths not divisible by chunk_size.

The five official quick-grid cases passed 5/5. ATK accuracy testing passed 24/24 cases.

The tests compare the Triton result with both an fp32 PyTorch reference and an fp64 golden implementation. Near-zero qk_dot and dmimo_z outputs use a global error relative to the fp64 golden result
to avoid unstable element-wise relative-error ratios.

likedislike
合并受阻
Bbitszh3271成员
7月16日 关联了issue:【Feature】triton算子(_layer_norm_bwd_kernel)迁移mindspeed-ops仓
atomgit-bot
atomgit-bot
7月16日 评论:

🤖 正在生成合并请求摘要,请稍候…

likedislike
不准确?
atomgit-bot
atomgit-bot
7月16日 评论:

🤖 AI 代码检视正在进行中,请稍候…

likedislike
不准确?
ascend-robotascend-robot成员
7月16日 添加了label:stat/needs-squash
ascend-robotascend-robot成员
7月16日 添加了label:ascend-cla/yes
ascend-robot
ascend-robot成员
7月16日 评论:

Thanks for your pull-request.
The full list of commands accepted by me can be found at here.
You can get sig-info at here.
You can self-configure the PR merge rules for this repository. For more details, please refer to Here.


PR Approval Progress

⚠️ This PR does not yet meet the following requirements:lgtm (requires ≥ 2 person(s) per module)、approve (requires ≥ 1 person(s) per module)

Module Approval Details

module lgtm status approve status
docs LinShua (1/2)(You can also ask: 郑加利, bigdog1206, zhizaidicengshehua, feng0w0, LinMingZhe) ❌ (0/1)(You can also ask: 周蓓蓉, 孙银磊, 王晓歆, 郑加利, 刘哲续)
repo-Ascend/MindSpeed-Ops LinShua (1/2)(You can also ask: iansheng, 华郁秀, 丁子霖, bigdog1206, 朱彦儒) ❌ (0/1)(You can also ask: 周蓓蓉, 朱彦儒, 华郁秀, 刘哲续, 孙银磊)

💡 Tip:

  • Committer can comment /approve or /lgtm
  • Commenting /approve implies both code review (lgtm) and intent to merge (approve)

CLA Signature Pass

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

likedislike
atomgit-bot
atomgit-bot
7月16日 评论:

⚠️ 本次变更过大(15 个文件、5166 行),已超出 AI 代码评审的处理范围,本次跳过。建议拆分为更小的 PR 以获得有效评审。

likedislike
不准确?
atomgit-bot
atomgit-bot
7月16日 评论:

⚠️ 本次变更过大(15 个文件、5166 行),已超出 AI 代码评审的处理范围,本次跳过。建议拆分为更小的 PR 以获得有效评审。

likedislike
不准确?
bitszh3271成员
7月16日 评论:

mamba3_mimo_bwd_fwd Triton-Ascend 实践

背景

Mamba-3 MIMO 的 combined backward 沿用“前向重算 + 反向扫描”的两阶段结构。
mamba3_mimo_bwd_fwd 虽然属于 backward 流程,但计算方向仍是从序列头到序列尾:它重新建立
每个 chunk 的进入状态,计算输出投影与门控梯度,并为第二阶段保存 statesqk_dot

上游实现以 GPU TileLang kernel 为基础。迁移到 Triton-Ascend 时,核心工作不是逐行替换语法,而是
重新确定任务粒度、UB 中的常驻数据和 Vector/Cube 的职责边界。本实践文档记录这部分实现选择;
稳定接口见 mamba3_mimo_bwd_fwd.md,性能演进见 mamba3_mimo_bwd_fwd_optimization.md

计算语义拆解

GQA 与旋转位置编码

输入 q/k 的 head 维为 G,状态和 v 的 head 维为 H。每个 attention head 使用:

g = h // (H / G)
q_h = q[..., g, :] + q_bias[h, ...]
k_h = k[..., g, :] + k_bias[h, ...]

随后对前 RD = N / rotary_dim_divisor 维执行 rotate-half。第 n 维与第 N/2+n 维配对:

q_rot[n]       = cos(angle[n]) * q[n] - sin(angle[n]) * q[N/2+n]
q_rot[N/2+n]   = sin(angle[n]) * q[n] + cos(angle[n]) * q[N/2+n]

k 使用相同变换。旋转以 fp32 计算,并写入 [B,H,S,R,N] scratch,后续 scan 直接按 head 连续读取。

chunk 内输出与状态更新

设当前 chunk 的长度为 C,将 (time, rank) 展平为 C*R。离散化链的两个主要量为:

gamma[t]      = dt[t] * sigmoid(trap[t])
trap_scale[t] = gamma[t] + dt[t+1] * sigmoid(-trap[t+1])

序列末尾没有后一项,第二项取零。每个 rank 的 PsiVv * mimo_vraw_y 由三部分组成:

  1. 当前 q 与 chunk 进入状态的交互;
  2. chunk 内严格因果的 q-k 与 PsiV 交互;
  3. 同位置 q-k 对角项,以及可选 D-skip。

chunk 结束时,key 按 trap_scale 与反向累计衰减缩放,再与 PsiV 收缩得到状态增量:

state_next = state_in * exp(da_cs_sum) + K_state.T @ PsiV

state_in 在计算前写入 states[:, :, chunk]。Z 门控不参与状态更新,只作用于 raw_y 到最终输出
的投影,因此 Phi、Zeta 和 z 的梯度可以在本阶段完成。

从基线到生产实现

stage0 切分

优化前实现保留在 mamba3_mimo_bwd_fwd_baseline_impl.py,由三个步骤组成:

步骤 作用
rotary bias-add 与 rotate-half,生成 Qr/Kr
qkdot 计算每个 token 的 R x R q-k 点积
scan (B,H) 顺序扫描 chunk,生成状态、投影梯度和门控梯度

这种切分先保证数学路径完整,并使 PyTorch reference、stage0 和生产实现共享相同的输入输出契约。

staged 生产路径

生产实现仍保留 (B,H) 串行 scan,但根据 shape 选择预处理与 scan 变体:

kernel 使用场景
_mamba3_mimo_rotary_qkdot_kernel 小预处理网格,合并 rotary 与 qk_dot
_mamba3_mimo_rotary_kernel + _mamba3_mimo_qkdot_kernel 大网格,避免融合后重复读 key
_mamba3_mimo_bwd_fwd_scan_kernel 常规 R1/R2/R3 路径
_mamba3_mimo_bwd_fwd_scan_rf_kernel 可单 tile 驻留的 R4/R8 rank-fold 路径
_mamba3_mimo_bwd_fwd_scan_kernel_bt 大 N/P 的 N-blocked 路径
*_post_kernel 长序列的 Phi/Zeta/dz 并行收缩

aux 生产路径

B*H 不能占满 AI Core 时,仅增加 scan 内部优化无法获得足够并行度。aux 路径将每个 chunk
可独立计算的部分放到 (B,H,nchunks) 网格:

pre_aux ──► INTRA, KV, qk_dot, 离散化 scratch
                      │
                      ▼
               scan_inter ──► states, dmimo_o

INTRA 保存不依赖前一 chunk 状态的输出,KV 保存状态增量。scan_inter 只组合前一状态与这两项。
这条路径不处理 Z 门控和非整 chunk,因此 dispatcher 只在参数契约完整匹配时启用。

Ascend 侧实现要点

UB 预算

scan 同时需要旧状态和新状态,单份状态占 N * P * 4 字节。若直接让两份 [N,P] fp32 tile
常驻,N256_P128 仅状态就需要 256 KiB,尚未计入 q/k、PsiV 和输出 tile。生产实现以
N * P > 16384 作为常规大 tile 边界,并在 rank-fold 的 host 计划中进一步估算折叠矩阵占用。

大 tile 路径把 running state 放到 HBM scratch,按 N block 流式更新。该做法增加 HBM 往返,但把
UB 峰值限制在 BN * P,是覆盖大 shape 所需的容量交换。

rank 维与对齐

典型 R=2 的 fp32 rank 尾轴只有 8 字节,不适合作为独立的一维搬运单位。实现不把 rank 作为
最内层窄向量,而是使用 [C,N][C,P] 二维 tile,并在 rank 循环中读写。R4 以上的矩阵收缩
则将 rank 与 chunk 折为 R*C,扩大点积规模并减少地址生成。

编译期展开

static_range(R) 对小 rank 有利,但 R6/R8 的 R x R 展开会显著增加 IR 与 cbuf 需求。host 对
R >= 6 选择运行时 rank 循环;aux 预处理在大 N 或窄 P 的 R4 shape 也使用 de-unroll 版本,以控制
UB 和编译规模。该选择不改变循环内浮点运算顺序。

qk_dot 的计算位置

同一位置对 q/k 同时应用正交旋转不会改变点积,因此 qk_dot 可以从 bias-add 后的 q/k 直接计算。
低精度输入使用向量 reduce,fp32 输入使用 dot,以匹配 reference 的累加误差。rank-fold scan 已经
持有相关 q/k tile 时直接写 qk_dot,避免额外 kernel。

非整 chunk

kernel 内部只处理完整 chunk。公开入口在 host 侧补齐 q/k/v/dout/angles/z,使用末值延展
da_cs,并重建 da_cs_revsegsum。这种做法避免在每个点积上引入尾块分支,同时由 tail UT
验证补齐区不会影响有效 token。

数值与验证方法

所有状态、点积和梯度归约使用 fp32;output_dtype 只控制最终写回。验证分三层:

  1. PyTorch 小算子以 fp64 重算 raw_y/states/qk_dot 并通过 autograd 求投影、门控梯度;
  2. stage0 用相同输入输出契约验证迁移前的 Triton 计算;
  3. 生产实现通过 pairwise、上游 shape 网格和 tail 用例与 fp32/fp64 双标杆比较。

qk_dot 在随机输入下可能接近零,使用全局最大绝对误差相对 golden 最大值的指标;其它输出沿用
仓库的 dual-benchmark 判据。dmimo_z 也会出现少量近零元素,因此同样直接约束候选对 fp64 golden
的全局最大相对误差,避免逐元素相对误差掩盖整体 RMSE。生产 API 不导入测试 reference,也不在
不支持的架构上静默切换实现。

扩展实现时的检查项

  • 新增输入分支时,同时确认 dispatcher 的 aux 条件是否仍完整覆盖其语义。
  • 修改状态更新前,检查 states 保存的是 chunk 进入态而不是离开态。
  • 调整 tile 后重新核算双状态、折叠 q/k、PsiV 和归约 scratch 的总 UB 占用。
  • 修改 rank 循环后至少覆盖 R1/R2/R4/R8,并分别检查编译时间和数值结果。
  • 修改尾块填充值时同步更新 da_cs_rev/segsum 重建逻辑与 tail reference。
  • 性能对比使用仓内 stage0,不使用 PyTorch eager 代替优化前 Triton 基线。
likedislike
bitszh3271成员
7月16日 评论:

mamba3_mimo_bwd_fwd 算子迁移说明

算子概述

mamba3_mimo_bwd_fwd_kernel 是 Mamba-3 MIMO combined backward 的第一阶段。算子按
chunk 重算前向中间量,计算输出投影与可选门控分支的梯度,并生成第二阶段反向扫描所需的
statesqk_dot 缓存。

当前实现面向 Ascend arch32,公开入口位于
mindspeed_ops.api.triton.mamba3_mimo_bwd_fwd_kernel。arch35 暂未实现,调用时会明确抛出
NotImplementedError

函数签名

def mamba3_mimo_bwd_fwd_kernel(
    dout,
    q,
    k,
    v,
    q_bias,
    k_bias,
    mimo_v,
    mimo_o,
    angles,
    da_cs,
    da_cs_rev,
    dt,
    trap,
    segsum,
    mimo_z=None,
    d=None,
    z=None,
    chunk_size=16,
    rotary_dim_divisor=4,
    output_dtype=torch.float32,
):
    ...

输入与输出

设 batch size 为 B,序列长度为 S,MIMO rank 为 R,KV head 数为 G,query/key
维度为 N,attention head 数为 H,value 维度为 P,chunk 大小为 C

参数 形状 说明
dout [B, S, H, P] reduce-O 路径的上游梯度
q, k [B, S, R, G, N] query、key
v [B, S, H, P] value
q_bias, k_bias [H, R, N] 旋转位置编码前的偏置
mimo_v [H, R, P] value 投影参数 Psi
mimo_o [H, R, P] output 投影参数 Phi
angles [B, S, H, N / rotary_dim_divisor] 旋转位置编码角度
da_cs, da_cs_rev, dt, trap [B, H, S] 状态离散化中间量
segsum [B, H, ceil(S/C), C, C] chunk 内段和
mimo_z [H, R, P]None 可选门控投影参数 Zeta
d [H]None 可选 D-skip 参数
z [B, S, H, P]None 可选 SiLU 门控输入

公开 API 返回五元组:

返回值 形状 说明
states [B, H, ceil(S/C), N, P] 每个 chunk 的递推起始状态
qk_dot [B, H, S, R, R] 同位置 query-key 点积缓存
dmimo_o [B, H, R, P] Phi 梯度;调用方按 batch 维求和
dmimo_z [B, H, R, P]None Zeta 梯度
dz [B, S, H, P]None 门控输入梯度

dmimo_zdz 仅在 mimo_zz 均参与门控计算时返回张量。

在 combined backward 中的位置

Mamba-3 MIMO 的反向计算分为两次扫描:

上游梯度 dout
      │
      ▼
bwd_fwd:前向重算 + Phi/Zeta 梯度
      │
      ├── states ──┐
      └── qk_dot ──┼──► bwd_bwd:倒序状态扫描 ──► q/k/v 等其余梯度
                   │
                   └── 与离散化中间量共同描述前向状态

第一阶段不对状态递推求导,而是无梯度重算 raw_y。设 r 表示 MIMO rank,Phi 为
mimo_o,Zeta 为 mimo_z,则 reduce-O 路径的输出可写为:

u[b,s,h,r,p]    = z[b,s,h,p] * Zeta[h,r,p]
gate            = SiLU(u)                         # 未启用 Z 分支时为 1
out[b,s,h,p]    = sum_r Phi[h,r,p] * gate * raw_y[b,s,h,r,p]

因此 dmimo_odmimo_zdz 只依赖本次重算得到的 raw_ydout,可以在前向方向
完成;状态递推的反向依赖则留给 bwd_bwd。这种拆分与上游 combined backward 的缓存契约一致,
也避免在同一个 kernel 中同时维护正向和反向两条串行依赖链。

每个 chunk 开始前的状态写入 states[:, :, chunk]qk_dot 保存同一 token 上各 rank 的
query-key 点积:

qk_dot[b,h,s,r_out,r_in] = dot(q_bias_rot[b,h,s,r_out],
                               k_bias_rot[b,h,s,r_in])

rotate-half 是正交变换,同一位置同时旋转 q 和 k 不改变点积;实现可以复用旋转前 bias-add
结果计算该缓存,同时保持与后续反向公式相同的数学含义。

实现说明

staged 实现分为预处理、状态扫描和可选后处理:

  1. 预处理 kernel 按 (B, H, chunk) 并行,完成 bias-add、rotate-half,并计算或准备
    qk_dot。小网格会把 rotary 与 qk_dot 合在一次 launch 中。
  2. scan kernel 以 (B, H) 为任务粒度顺序遍历 chunk。chunk 内先计算
    gamma = dt * sigmoid(trap) 及因果衰减,再组合上一状态、chunk 内 q-k 交互、Psi 投影与
    D-skip,得到各 rank 的 raw_y
  3. 状态更新使用当前 chunk 的 key 与 Psi-value 投影生成下一状态。所有点积和归约使用 fp32
    累加,写回时转换为 output_dtype
  4. 对较长序列,非 rank-fold 路径可将 Phi/Zeta 收缩和 dz 移到并行后处理 kernel,降低串行
    scan 中的向量工作量;短序列保留内联计算,避免额外 launch 和 scratch 往返。

aux 实现把计算进一步拆为三段:合并的 pre/aux kernel 并行生成旋转结果、qk_dot、chunk 内
输出 INTRA 和状态增量 KVscan_inter 只处理跨 chunk 状态递推与 dmimo_o;dispatcher
仅在这条路径完整覆盖调用参数时选择它。

生产入口根据 shape 选择两条实现路径:

  • 当未启用 Z 门控、S 能被 chunk_size 整除、单个 (B, H) 网格不能占满 AI Core 且
    N * P <= 16384 时,使用 batch/chunk 并行的 aux 实现补足并行度。
  • 其余情况使用功能完整的 staged scan 实现。该路径支持 Z 门控、D-skip 和尾块处理。

S % chunk_size != 0 时,host wrapper 将序列右填到完整 chunk。普通序列输入补零,
da_cs 延用最后一个有效值,并重建尾块的 da_cs_revsegsum;kernel 执行结束后再把
带序列维的输出裁回原长度。

支持范围

  • 硬件:Ascend arch32。
  • dtype:float16bfloat16float32
  • 布局:dense、非 varlen,H % G == 0
  • rank:支持通用 R;已覆盖 R = 1, 2, 4, 8
  • N 必须为偶数,且旋转维度不得超过 N / 2
  • mimo_o 为必需输入;当前公开入口只提供 reduce-O 形式的 dout
  • mimo_zz 应同时提供。只提供其中一个不构成有效门控配置。
  • 支持 reduce-O、可选 Z/SiLU 门控、可选 D-skip 和非整 chunk 的序列长度。
  • 当前公开 API 不包含 cu_seqlensfuse_pregate_headwise_rms_norm
    return_final_state

调用示例

from mindspeed_ops.api.triton.mamba3_mimo_bwd_fwd_kernel import (
    mamba3_mimo_bwd_fwd_kernel,
)

states, qk_dot, dmimo_o, dmimo_z, dz = mamba3_mimo_bwd_fwd_kernel(
    dout,
    q,
    k,
    v,
    q_bias,
    k_bias,
    mimo_v,
    mimo_o,
    angles,
    da_cs,
    da_cs_rev,
    dt,
    trap,
    segsum,
    mimo_z=mimo_z,
    d=d,
    z=z,
    chunk_size=16,
    output_dtype=torch.float32,
)

statesqk_dot 作为 mamba3_mimo_bwd_bwd_kernel 的输入继续完成第二阶段反向扫描。

代码与测试

文件 内容
mindspeed_ops/api/triton/mamba3_mimo_bwd_fwd_kernel.py 公开 API 与架构检查
mindspeed_ops/arch32/triton/mamba3/mamba3_mimo_bwd_fwd_baseline_impl.py 优化前 Triton 基线
mindspeed_ops/arch32/triton/mamba3/mamba3_mimo_bwd_fwd_impl.py staged 生产实现
mindspeed_ops/arch32/triton/mamba3/mamba3_mimo_bwd_fwd_aux_impl.py 高并行度 aux 实现
mindspeed_ops/arch32/triton/mamba3/mamba3_mimo_bwd_fwd_dispatch_impl.py shape 自适应调度
tests/atk_tests/triton/mamba3_mimo_bwd_fwd_kernel/reference_impl.py PyTorch 小算子参考实现
tests/unit_tests/triton/test_mamba3_mimo_bwd_fwd_kernel.py 双标杆、官方 shape 网格和尾块测试
tests/atk_tests/triton/mamba3_mimo_bwd_fwd_kernel/ ATK 用例配置与适配代码

单元测试使用同一份 PyTorch 参考分别计算 fp32 reference 和 fp64 golden,并比较 Triton 输出与两套
标杆的误差。参考实现通过 PyTorch 小算子重建 chunk 状态,再用 autograd 独立计算 Phi、Zeta 和 z
的梯度,未复用生产 kernel 的梯度公式。

用例组 覆盖内容
pairwise 常规用例 B={1,2}S={32,64}H={2,4}、fp16/bf16/fp32、Z、D、GQA
official quick 从 11 组上游 shape 中选取 5 组,覆盖 R={1,2,4}N=16..256P=64..128
official full B=4, S=2048, H=16, G=1 的 11 组完整网格,包含 R8;由环境变量开启
tail S={33,40,48,50}C={16,32},校验右填、离散量重建和输出裁剪

逐项检查 statesqk_dotdmimo_o,启用 Z 时额外检查 dmimo_zdz
qk_dot 的真值接近零,专项网格采用全局最大绝对误差相对 golden 最大值的指标,避免逐元素相对误差
在零点附近失真。dmimo_z 也包含接近零的投影梯度,统一使用候选对 fp64 golden 的全局相对误差;
其余输出使用 dual-benchmark。完整 shape 网格通过 MIMO22_RUN_OFFICIAL_FULL_GRID=1 开启。

pytest -q tests/unit_tests/triton/test_mamba3_mimo_bwd_fwd_kernel.py

ATK 精度任务覆盖 24 个用例,当前记录为 24/24 通过。

likedislike
bitszh3271成员
7月16日 评论:

mamba3_mimo_bwd_fwd 优化记录

基线与目标

优化基线保存在
mindspeed_ops/arch32/triton/mamba3/mamba3_mimo_bwd_fwd_baseline_impl.py。该版本按上游
TileLang 计算顺序拆出 rotary、qk_dot 和串行 scan,主要用于精度对照和量化后续优化收益。

生产实现保存在 mamba3_mimo_bwd_fwd_impl.pymamba3_mimo_bwd_fwd_aux_impl.py
mamba3_mimo_bwd_fwd_dispatch_impl.py。优化没有改变公开 API、输出布局和 fp32 累加方式。

基线与生产版均保留以下计算约束,性能比较不包含语义裁剪:

  • 相同的 bias-add 与 rotate-half 位置编码;
  • 相同的 chunk 状态递推、Psi/Phi 投影和可选 D-skip;
  • 相同的 statesqk_dotdmimo_odmimo_zdz 输出;
  • 相同的 fp32 中间累加和 output_dtype 转换位置。

瓶颈分析

B=4, S=2048, H=16, G=1 的上游 shape 网格中,stage0 的 scan 占主要执行时间。
scan 必须沿 chunk 保持状态依赖,但每个 chunk 内又包含大量 R x R 小矩阵运算。profile 显示该路径
主要受标量地址生成和 tiny-dot 启动开销限制,Cube 与 Vector 的有效计算占比偏低。仅减少 HBM load
不能消除这一开销。

另一个问题出现在较小的 B * H:按 (batch, head) 启动的串行 scan 无法占满 AI Core,设备上
仍有可用于 chunk-local 计算的并行度。

基线还存在两个 shape 扩展问题:一是 N * P 较大时,递推状态与新状态同时驻留 UB 会超出容量;
二是 R6/R8 下静态展开 R x R rank 循环会导致编译期 IR 与 cbuf 需求快速增长。这两个问题虽不总是
表现为运行时热点,但决定了优化实现能否覆盖完整 shape 网格。

优化项与收益

1. shape 自适应的 aux 路径

对不含 Z 门控、序列长度为完整 chunk、B * H 不足以占满 AI Core 且 N * P <= 16384
shape,将 chunk-local 计算拆到 (B, H, nchunks) 网格,再由轻量 scan 完成跨 chunk 状态递推。
该路径提高了小 batch、小 head 场景的核利用率;其余 shape 自动走 staged 实现,保持完整功能覆盖。

dispatcher 的生产选择条件如下:

条件 选择
z is None、整 chunk、B * H <= AI Core 数N * P <= 16384 aux + inter scan
启用 Z、存在尾块、并行度已足够或 tile 超预算 staged scan

aux 路径将 INTRA[B,H,S,R,P]KV[B,H,nchunks,N,P] 作为并行阶段与串行阶段之间的中间量,
用额外 scratch 换取 chunk 维并行度。只有在 AI Core 欠占用时这笔交换才有收益,因此不作为无条件路径。

2. rank-fold scan

R >= 4 且未启用融合后处理时,将 rank 输入维与 chunk 维折叠为 R * C

  • 原先逐 rank 发起的 R x R 组小矩阵点积合并为较大的矩阵运算;
  • 状态更新、qk 收缩和 mimo_v 投影复用折叠后的 tile;
  • 保留 r_out 循环以控制 UB 使用量。

该优化直接减少 tiny-dot 数量与地址生成次数,是 R4/R8 shape 的主要收益来源。

折叠前,一个 chunk 内需要为多个 (r_out, r_in) 组合分别建立 block pointer 并启动点积;折叠后,
q-k、状态和 Psi-value 的主要收缩由少量 [C, R*C][R*C, N][R*C, P] 矩阵完成。
浮点乘加的结合顺序可能与 stage0 不同,因此正确性以 fp64 golden 而非逐位一致作为判断依据。

3. 大 tile 分块

N * P 超出单 tile 预算时,使用 _bt 路径沿 N 分块。当前分块优先将 N 划分为两个
block,并在 N-block 之间复用每个 chunk 的离散化量与投影中间量,避免因 P-blocking 重复执行标量逻辑。
该路径使 N=256P=128 等 shape 在 UB 约束内继续使用 rank-fold。

是否能使用单块 rank-fold 由 host 端按 tile 字节数估算。估算包含双份状态、折叠后的 q/k、输出投影、
因果系数和更新状态;若完整 P tile 超预算,先缩小 P block,最终路由到 N-blocked _bt。对于
非 rank-fold shape,仅 N * P > 16384 才进入大 tile 路径,避免让本可驻留 UB 的
N256_P64N128_P128 承担不必要的 HBM 往返。

4. 融合 qk_dot

rank-fold scan 已经加载同一位置的 query 和 key,因此直接在 scan 内写出 qk_dot,省去独立
qk_dot kernel。对 R=8,独立 kernel 中的 rank 组合最多,这项融合的收益最明显。

小预处理网格仍可选择 rotary + qk_dot 合并 kernel;大网格若不走 rank-fold,则保留独立
qk_dot kernel。这样既减少 launch,也避免在大网格上因融合而重复读取 key。

5. rank 循环 de-unroll

大 rank 的 static_range 会显著放大编译期 IR 和片上资源需求。实现对 R >= 6 使用运行时 rank
循环,在不改变浮点运算顺序的前提下避免 R8 编译失败;较小 rank 仍保留静态展开。

这项改动本身用于解除编译上限,不把“原始 R8 stage0 无法编译”计作无限加速。性能表中的 R8
基线应用同样的 de-unroll,只比较后续 rank-fold 与 qk_dot 融合带来的运行时收益。

6. 串行 scan 的局部优化

未进入 rank-fold 的 staged 路径仍按 shape 使用两项低风险优化:R <= 4R != 2 时把
mimo_v 提到 chunk 循环外复用;C <= 16 时直接加载离散化量,减少 scan 内重建工作。
这些改动主要改善 R1/R3 和小 chunk,不改变 rank-fold 的调度范围。

性能结果

以下结果在 Ascend 910B 单卡、bf16、上游 B=4, S=2048, H=16, G=1 shape 网格上测得。
候选实现与 stage0 基线在同卡交错执行,每个 shape 取 3 轮结果。R8 的 stage0 使用仅将 rank 循环
de-unroll 的可编译等价版本作为基线。

Shape (N_P_R_C) 相对 stage0
N16_P64_R4_C8 2.13x
N32_P64_R4_C16 2.05x
N64_P64_R4_C16 2.12x
N128_P64_R4_C16 1.98x
N256_P64_R4_C8 1.52x
N64_P128_R4_C16 1.99x
N128_P32_R4_C16 2.13x
N128_P128_R4_C8 1.78x
N128_P64_R8_C8 7.12x
N128_P64_R2_C32 1.03x
N128_P64_R1_C64 0.95x

11 组 shape 的几何平均加速比为 1.94x。R1/R2 不满足 rank-fold 调度条件,作为未使用该优化的
对照组保留在统计中。

结果可以分为三类理解:

  • R4 常规 shape 稳定在 1.98x–2.13x,说明 rank-fold 对常见矩阵尺寸的收益稳定。
  • N256_P64N128_P128 需要更保守的 tile,分块和 HBM scratch 抵消了一部分收益,但仍达到
    1.52x1.78x
  • R8 同时减少大量 rank 组合的 tiny-dot 与独立 qk_dot 工作,因此达到 7.12x。R1/R2 不使用
    rank-fold,结果接近基线,符合调度预期。

除完整 11-shape 网格外,四组常用生产 shape 的几何平均加速比记录为 1.224x。该组包含一个
B=8, H=16、AI Core 已饱和的控制 shape,因此 aux 并行化不会被错误地计为普遍收益。

精度与边界验证

  • stage0 与生产实现使用同一套 PyTorch 小算子参考进行 fp32/fp64 双标杆校验。
  • 官方 quick shape、常规 dtype/特性组合以及非整 chunk 尾块均由 UT 覆盖。
  • qk_dot 的真值接近零,测试使用 max(abs(diff)) / max(abs(golden)) 的全局相对误差,避免逐元素
    相对误差放大近零噪声。
  • dmimo_z 同样含有近零投影梯度,使用候选对 fp64 golden 的全局最大相对误差;其它梯度仍使用
    fp32/fp64 dual-benchmark。
  • ATK 精度任务为 24/24 通过。

尾块测试单独验证 host padding 契约:q/k/v/dout/angles/z 补零,da_cs 使用末值延展,随后重建
da_cs_revsegsum。这保证性能 kernel 仍只处理整 chunk,而任意正序列长度由公开入口完整承接。

性能数据只比较能够保持相同输入、输出和计算语义的实现;生产路径不使用 PyTorch reference 或其他
整算子作为运行时回退。

likedislike
Bbitszh3271成员
7月16日 修改了pull request 的描述
bitszh3271成员
7月16日 评论:

compile

likedislike
ascend-robotascend-robot成员
7月16日 添加了label:ci-pipeline-running
ascend-robotascend-robot成员
7月16日 删除了label:ci-pipeline-running
ascend-robotascend-robot成员
7月16日 添加了label:ci-pipeline-failed
ascend-robot
ascend-robot成员
7月16日 评论:
流水线 PR-pipeline_MindSpeed-Ops#495 [ commitID:3b4283c1 ] 运行失败
>>>代码风格自动修复执行失败,具体请查看日志,不影响流水线执行及PR合入
阶段 任务名 状态 详情
恶意代码检查 Antipoison >>>
编码安全与规范检查 Only_doc_commit >>>
pre-commit >>>
开源片段检查 SCA >>>
开发者测试 UT_MindSpeed 🕚 >>>
流水线 PR-pipeline_MindSpeed-Ops >>>
此流水线已支持下列评论快捷指令,仅PR创建者和白名单成员[wujinyuan1, liuzhexu]评论有效
  • compile : 运行流水线
  • retry : 重试流水线所有失败子任务
  • retry <任务名> : 仅重试指定失败子任务
  • stop : 停止流水线
likedislike
Bbitszh3271成员
7月16日 推送  1 个提交:0b521e69-fix: address Triton compliance and static checks
ascend-robotascend-robot成员
7月16日 删除了label:ci-pipeline-failed
ascend-robot
ascend-robot成员7月16日进行代码检视1
docs/triton/mamba3_mimo_bwd_fwd.md
ascend-robot
ascend-robot7月16日评论:

【openlibing.ci】检测到当前PR中存在代码检查告警抑制 7 处,详情见下表,请Committer检视合理性。 / Detected 7 code check alert suppression(s) in this PR, see table below. Committers please review.

文件路径/File 行号/Line 代码片段/Snippet 工具/Tool
mindspeed_ops/api/triton/
mamba3_mimo_bwd_fwd_kernel.py
3 # pylint: disable=duplicate-code pylint
mindspeed_ops/arch32/
triton/mamba3/mamba3_mimo_bwd_fwd_aux_impl.py
3 # pylint: disable=duplicate-code,too-many-lines pylint
mindspeed_ops/arch32/
triton/mamba3/mamba3_mimo_bwd_fwd_baseline_impl.py
3 # pylint: disable=duplicate-code pylint
mindspeed_ops/arch32/
triton/mamba3/mamba3_mimo_bwd_fwd_dispatch_impl.py
3 # pylint: disable=duplicate-code pylint
mindspeed_ops/arch32/
triton/mamba3_mimo_fwd.py
3 # pylint: disable=possibly-used-before-assignment,too-many-nested-blocks pylint
tests/atk_tests/triton/
mamba3_mimo_bwd_fwd_kernel/
generate_mamba3_mimo_bwd_fwd_kernel.py
2 # pylint: disable=unsubscriptable-object pylint
tests/atk_tests/triton/
mamba3_mimo_bwd_fwd_kernel/
reference_impl.py
3 # pylint: disable=duplicate-code pylint
likedislike
bitszh3271成员
7月16日 评论:

compile

likedislike
ascend-robotascend-robot成员
7月16日 添加了label:ci-pipeline-running
ascend-robotascend-robot成员
7月16日 删除了label:ci-pipeline-running
ascend-robotascend-robot成员
7月16日 添加了label:ci-pipeline-failed
ascend-robot
ascend-robot成员
7月16日 评论:
流水线 PR-pipeline_MindSpeed-Ops#496 [ commitID:0b521e69 ] 运行失败
>>>代码风格自动修复执行成功(无修复内容)
阶段 任务名 状态 详情
恶意代码检查 Antipoison >>>
编码安全与规范检查 Only_doc_commit >>>
pre-commit >>>
开源片段检查 SCA >>>
开发者测试 UT_MindSpeed 🕚 >>>
流水线 PR-pipeline_MindSpeed-Ops >>>
此流水线已支持下列评论快捷指令,仅PR创建者和白名单成员[wujinyuan1, liuzhexu]评论有效
  • compile : 运行流水线
  • retry : 重试流水线所有失败子任务
  • retry <任务名> : 仅重试指定失败子任务
  • stop : 停止流水线
likedislike
Bbitszh3271成员
7月16日 推送  1 个提交:a3f935d0-chore: allow Triton terms in lint checks
ascend-robotascend-robot成员
7月16日 删除了label:ci-pipeline-failed
bitszh3271成员
7月16日 评论:

compile

likedislike
ascend-robotascend-robot成员
7月16日 添加了label:ci-pipeline-running
ascend-robotascend-robot成员
7月16日 删除了label:ci-pipeline-running
ascend-robotascend-robot成员
7月16日 添加了label:ci-pipeline-passed
ascend-robot
ascend-robot成员
7月16日 评论:
流水线 PR-pipeline_MindSpeed-Ops#501 [ commitID:a3f935d0 ] 已完成
>>>代码风格自动修复执行成功(无修复内容)
阶段 任务名 状态 详情
恶意代码检查 Antipoison >>>
编码安全与规范检查 Only_doc_commit >>>
pre-commit >>>
开源片段检查 SCA >>>
开发者测试 UT_MindSpeed >>>
流水线 PR-pipeline_MindSpeed-Ops >>>
此流水线已支持下列评论快捷指令,仅PR创建者和白名单成员[wujinyuan1, liuzhexu]评论有效
  • compile : 运行流水线
  • retry : 重试流水线所有失败子任务
  • retry <任务名> : 仅重试指定失败子任务
  • stop : 停止流水线
likedislike
LinShua成员16 天前进行代码检视1
tests/atk_tests/triton/mamba3_mimo_bwd_fwd_kernel/generate_mamba3_mimo_bwd_fwd_kernel.py
@@ -0,0 +1,0 @@
1+# Copyright (c) 2026, Huawei Technologies Co., Ltd. All rights reserved.
LinShua16 天前评论:

PR描述补充ATK精度和性能通过截图

likedislike
LinShua成员15 天前进行代码检视1
docs/triton/mamba3_mimo_bwd_fwd.md
@@ -0,0 +1,1 @@
1+# mamba3_mimo_bwd_fwd 算子迁移说明
2+ 
LinShua15 天前评论:

请在PR描述中补充当前优化后的triton算子与开源triton算子在GPU上的性能对比和精度对比数据(或小算子与GPU上triton算子运行精度对比数据)

likedislike
bitszh3271成员
14 天前 评论:

AscendC Official 11 逐 Shape 三轮中位数统计(单位:ms)

Ratio 定义为 A100耗时 / 910B3耗时

Official Shape mimo bwdfwd A100 mimo bwdfwd AscendC mimo bwdfwd Ratio mimo bwdbwd A100 mimo bwdbwd AscendC v119c mimo bwdbwd Ratio
N16_P64_R4_C8_BB128 1.538 7.842 0.1962x 1.940 15.498 0.1252x
N32_P64_R4_C16_BB256 1.479 5.536 0.2671x 2.259 10.252 0.2203x
N64_P64_R4_C16_BB256 1.682 6.306 0.2668x 2.669 11.132 0.2398x
N128_P64_R4_C16_BB256 2.323 7.981 0.2910x 3.935 14.013 0.2808x
N256_P64_R4_C8_BB256 5.181 14.565 0.3557x 6.950 26.396 0.2633x
N64_P128_R4_C16_BB256 2.328 7.987 0.2914x 3.468 13.123 0.2643x
N128_P32_R4_C16_BB256 1.740 7.046 0.2469x 3.771 12.630 0.2986x
N128_P128_R4_C8_BB256 5.009 13.247 0.3781x 5.913 22.522 0.2625x
N128_P64_R8_C8_BB256 4.579 16.257 0.2817x 8.229 37.347 0.2203x
N128_P64_R2_C32_BB256 1.269 4.193 0.3025x 1.994 8.910 0.2238x
N128_P64_R1_C64_BB256 0.773 2.330 0.3320x 1.294 4.122 0.3139x
11-shape GM 2.135 7.423 0.2876x 3.311 13.761 0.2406x

Triton Official 11 逐 Shape 三轮中位数统计(单位:ms)

Ratio 定义为 A100耗时 / 910B3耗时

Official Shape mimo bwdfwd A100 mimo bwdfwd AscendC mimo bwdfwd Ratio mimo bwdbwd A100 mimo bwdbwd AscendC v119c mimo bwdbwd Ratio
N16_P64_R4_C8_BB128 1.538 30.424 0.0506x 1.940 30.662 0.0633x
N32_P64_R4_C16_BB256 1.479 15.124 0.0978x 2.259 18.876 0.1197x
N64_P64_R4_C16_BB256 1.682 14.772 0.1139x 2.669 20.209 0.1321x
N128_P64_R4_C16_BB256 2.323 15.461 0.1502x 3.935 24.499 0.1606x
N256_P64_R4_C8_BB256 5.181 42.057 0.1232x 6.950 45.009 0.1544x
N64_P128_R4_C16_BB256 2.328 18.564 0.1254x 3.468 23.432 0.1480x
N128_P32_R4_C16_BB256 1.740 14.540 0.1197x 3.771 22.523 0.1674x
N128_P128_R4_C8_BB256 5.009 41.127 0.1218x 5.913 38.051 0.1554x
N128_P64_R8_C8_BB256 4.579 45.983 0.0996x 8.229 50.917 0.1616x
N128_P64_R2_C32_BB256 1.269 7.345 0.1727x 1.994 16.912 0.1179x
N128_P64_R1_C64_BB256 0.773 2.569 0.3011x 1.294 4.231 0.3058x
11-shape GM 2.135 17.338 0.1231x 3.311 22.980 0.1441x

反向传播算子复杂,输入输出多,主要受限UB大小,HBM搬运效率,以及算法中包含很多大batch小矩阵乘法(MIMO计算),难以继续优化。

likedislike
bitszh3271成员
14 天前 评论:
likedislike
bitszh3271成员
14 天前 评论:

@LinShua 已经更新

likedislike
LinShua成员
9 天前 评论:

/lgtm

likedislike