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


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


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
/approveor/lgtm- Commenting
/approveimplies 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. 👍


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


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


mamba3_mimo_bwd_fwd Triton-Ascend 实践
背景
Mamba-3 MIMO 的 combined backward 沿用“前向重算 + 反向扫描”的两阶段结构。
mamba3_mimo_bwd_fwd 虽然属于 backward 流程,但计算方向仍是从序列头到序列尾:它重新建立
每个 chunk 的进入状态,计算输出投影与门控梯度,并为第二阶段保存 states、qk_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 的 PsiV 为 v * mimo_v。raw_y 由三部分组成:
- 当前 q 与 chunk 进入状态的交互;
- chunk 内严格因果的 q-k 与 PsiV 交互;
- 同位置 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_rev 和 segsum。这种做法避免在每个点积上引入尾块分支,同时由 tail UT
验证补齐区不会影响有效 token。
数值与验证方法
所有状态、点积和梯度归约使用 fp32;output_dtype 只控制最终写回。验证分三层:
- PyTorch 小算子以 fp64 重算
raw_y/states/qk_dot并通过 autograd 求投影、门控梯度; - stage0 用相同输入输出契约验证迁移前的 Triton 计算;
- 生产实现通过 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 基线。


mamba3_mimo_bwd_fwd 算子迁移说明
算子概述
mamba3_mimo_bwd_fwd_kernel 是 Mamba-3 MIMO combined backward 的第一阶段。算子按
chunk 重算前向中间量,计算输出投影与可选门控分支的梯度,并生成第二阶段反向扫描所需的
states 和 qk_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_z 和 dz 仅在 mimo_z、z 均参与门控计算时返回张量。
在 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_o、dmimo_z 和 dz 只依赖本次重算得到的 raw_y 与 dout,可以在前向方向
完成;状态递推的反向依赖则留给 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 实现分为预处理、状态扫描和可选后处理:
- 预处理 kernel 按
(B, H, chunk)并行,完成 bias-add、rotate-half,并计算或准备
qk_dot。小网格会把 rotary 与qk_dot合在一次 launch 中。 - scan kernel 以
(B, H)为任务粒度顺序遍历 chunk。chunk 内先计算
gamma = dt * sigmoid(trap)及因果衰减,再组合上一状态、chunk 内 q-k 交互、Psi 投影与
D-skip,得到各 rank 的raw_y。 - 状态更新使用当前 chunk 的 key 与 Psi-value 投影生成下一状态。所有点积和归约使用 fp32
累加,写回时转换为output_dtype。 - 对较长序列,非 rank-fold 路径可将 Phi/Zeta 收缩和
dz移到并行后处理 kernel,降低串行
scan 中的向量工作量;短序列保留内联计算,避免额外 launch 和 scratch 往返。
aux 实现把计算进一步拆为三段:合并的 pre/aux kernel 并行生成旋转结果、qk_dot、chunk 内
输出 INTRA 和状态增量 KV;scan_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_rev 与 segsum;kernel 执行结束后再把
带序列维的输出裁回原长度。
支持范围
- 硬件:Ascend arch32。
- dtype:
float16、bfloat16、float32。 - 布局:dense、非 varlen,
H % G == 0。 - rank:支持通用
R;已覆盖R = 1, 2, 4, 8。 N必须为偶数,且旋转维度不得超过N / 2。mimo_o为必需输入;当前公开入口只提供 reduce-O 形式的dout。mimo_z与z应同时提供。只提供其中一个不构成有效门控配置。- 支持 reduce-O、可选 Z/SiLU 门控、可选 D-skip 和非整 chunk 的序列长度。
- 当前公开 API 不包含
cu_seqlens、fuse_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,
)
states 和 qk_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..256、P=64..128 |
| official full | B=4, S=2048, H=16, G=1 的 11 组完整网格,包含 R8;由环境变量开启 |
| tail | S={33,40,48,50} 与 C={16,32},校验右填、离散量重建和输出裁剪 |
逐项检查 states、qk_dot、dmimo_o,启用 Z 时额外检查 dmimo_z 和 dz。
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 通过。


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.py、mamba3_mimo_bwd_fwd_aux_impl.py 和
mamba3_mimo_bwd_fwd_dispatch_impl.py。优化没有改变公开 API、输出布局和 fp32 累加方式。
基线与生产版均保留以下计算约束,性能比较不包含语义裁剪:
- 相同的 bias-add 与 rotate-half 位置编码;
- 相同的 chunk 状态递推、Psi/Phi 投影和可选 D-skip;
- 相同的
states、qk_dot、dmimo_o、dmimo_z、dz输出; - 相同的 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=256、P=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_P64、N128_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 <= 4 且 R != 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_P64与N128_P128需要更保守的 tile,分块和 HBM scratch 抵消了一部分收益,但仍达到
1.52x和1.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_rev 与 segsum。这保证性能 kernel 仍只处理整 chunk,而任意正序列长度由公开入口完整承接。
性能数据只比较能够保持相同输入、输出和计算语义的实现;生产路径不使用 PyTorch reference 或其他
整算子作为运行时回退。


compile


| 阶段 | 任务名 | 状态 | 详情 |
|---|---|---|---|
| 恶意代码检查 | Antipoison | ✅ | >>> |
| 编码安全与规范检查 | Only_doc_commit | ✅ | >>> |
| pre-commit | ❌ | >>> | |
| 开源片段检查 | SCA | ❌ | >>> |
| 开发者测试 | UT_MindSpeed | 🕚 | >>> |
| 流水线 | PR-pipeline_MindSpeed-Ops | ❌ | >>> |
- compile : 运行流水线
- retry : 重试流水线所有失败子任务
- retry <任务名> : 仅重试指定失败子任务
- stop : 停止流水线


【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 |


compile


| 阶段 | 任务名 | 状态 | 详情 |
|---|---|---|---|
| 恶意代码检查 | Antipoison | ✅ | >>> |
| 编码安全与规范检查 | Only_doc_commit | ✅ | >>> |
| pre-commit | ❌ | >>> | |
| 开源片段检查 | SCA | ✅ | >>> |
| 开发者测试 | UT_MindSpeed | 🕚 | >>> |
| 流水线 | PR-pipeline_MindSpeed-Ops | ❌ | >>> |
- compile : 运行流水线
- retry : 重试流水线所有失败子任务
- retry <任务名> : 仅重试指定失败子任务
- stop : 停止流水线


compile


| 阶段 | 任务名 | 状态 | 详情 |
|---|---|---|---|
| 恶意代码检查 | Antipoison | ✅ | >>> |
| 编码安全与规范检查 | Only_doc_commit | ✅ | >>> |
| pre-commit | ✅ | >>> | |
| 开源片段检查 | SCA | ✅ | >>> |
| 开发者测试 | UT_MindSpeed | ✅ | >>> |
| 流水线 | PR-pipeline_MindSpeed-Ops | ✅ | >>> |
- compile : 运行流水线
- retry : 重试流水线所有失败子任务
- retry <任务名> : 仅重试指定失败子任务
- stop : 停止流水线


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计算),难以继续优化。


@LinShua 已经更新


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/38This pass recomputes the chunk-wise forward intermediates, calculates the Phi/Zeta projection gradients and optional gate-input gradient, and produces
statesandqk_dotcaches for the subsequentmamba3_mimo_bwd_bwd_kernel.The PR includes:
qk_dotgeneration, and rank-loop de-unrolling.chunk_size.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:
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:
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.