已开启
feat: mamba3-mimo-bwd-fwd-triton-kernel #101
feat: mamba3-mimo-bwd-fwd-triton-kernel #101
已开启
bitszh3271创建于 7月16日
17 个文件变更+7256-1
@@ -39,7 +39,7 @@ repos:
39 - id: codespell39 - id: codespell
40 args: [40 args: [
41 "-L",41 "-L",
42- "CANN,cann,NNAL,nnal,ASCEND,ascend,EnQue,CopyIn,ArchType,AND,ND,tbe,copyin,alog",42+ "CANN,cann,NNAL,nnal,ASCEND,ascend,EnQue,CopyIn,ArchType,AND,ND,tbe,copyin,alog,dout",
43 "--skip",43 "--skip",
44 "*.toml,*.py,*.cpp,*.hpp,*.c,*.h",44 "*.toml,*.py,*.cpp,*.hpp,*.c,*.h",
45 ]45 ]
ascend-robotascend-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
@@ -0,0 +1,219 @@
1+# mamba3_mimo_bwd_fwd 算子迁移说明
2+ 
L
LLinShua19 天前

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

likedislike
3+## 算子概述
4+ 
5+`mamba3_mimo_bwd_fwd_kernel` 是 Mamba-3 MIMO combined backward 的第一阶段。算子按
6+chunk 重算前向中间量,计算输出投影与可选门控分支的梯度,并生成第二阶段反向扫描所需的
7+`states``qk_dot` 缓存。
8+ 
9+当前实现面向 Ascend arch32,公开入口位于
10+`mindspeed_ops.api.triton.mamba3_mimo_bwd_fwd_kernel`。arch35 暂未实现,调用时会明确抛出
11+`NotImplementedError`
12+ 
13+## 函数签名
14+ 
15+```python
16+def mamba3_mimo_bwd_fwd_kernel(
17+ dout,
18+ q,
19+ k,
20+ v,
21+ q_bias,
22+ k_bias,
23+ mimo_v,
24+ mimo_o,
25+ angles,
26+ da_cs,
27+ da_cs_rev,
28+ dt,
29+ trap,
30+ segsum,
31+ mimo_z=None,
32+ d=None,
33+ z=None,
34+ chunk_size=16,
35+ rotary_dim_divisor=4,
36+ output_dtype=torch.float32,
37+):
38+ ...
39+```
40+ 
41+## 输入与输出
42+ 
43+设 batch size 为 `B`,序列长度为 `S`,MIMO rank 为 `R`,KV head 数为 `G`,query/key
44+维度为 `N`,attention head 数为 `H`,value 维度为 `P`,chunk 大小为 `C`
45+ 
46+| 参数 | 形状 | 说明 |
47+| --- | --- | --- |
48+| `dout` | `[B, S, H, P]` | reduce-O 路径的上游梯度 |
49+| `q`, `k` | `[B, S, R, G, N]` | query、key |
50+| `v` | `[B, S, H, P]` | value |
51+| `q_bias`, `k_bias` | `[H, R, N]` | 旋转位置编码前的偏置 |
52+| `mimo_v` | `[H, R, P]` | value 投影参数 Psi |
53+| `mimo_o` | `[H, R, P]` | output 投影参数 Phi |
54+| `angles` | `[B, S, H, N / rotary_dim_divisor]` | 旋转位置编码角度 |
55+| `da_cs`, `da_cs_rev`, `dt`, `trap` | `[B, H, S]` | 状态离散化中间量 |
56+| `segsum` | `[B, H, ceil(S/C), C, C]` | chunk 内段和 |
57+| `mimo_z` | `[H, R, P]``None` | 可选门控投影参数 Zeta |
58+| `d` | `[H]``None` | 可选 D-skip 参数 |
59+| `z` | `[B, S, H, P]``None` | 可选 SiLU 门控输入 |
60+ 
61+公开 API 返回五元组:
62+ 
63+| 返回值 | 形状 | 说明 |
64+| --- | --- | --- |
65+| `states` | `[B, H, ceil(S/C), N, P]` | 每个 chunk 的递推起始状态 |
66+| `qk_dot` | `[B, H, S, R, R]` | 同位置 query-key 点积缓存 |
67+| `dmimo_o` | `[B, H, R, P]` | Phi 梯度;调用方按 batch 维求和 |
68+| `dmimo_z` | `[B, H, R, P]``None` | Zeta 梯度 |
69+| `dz` | `[B, S, H, P]``None` | 门控输入梯度 |
70+ 
71+`dmimo_z``dz` 仅在 `mimo_z``z` 均参与门控计算时返回张量。
72+ 
73+## 在 combined backward 中的位置
74+ 
75+Mamba-3 MIMO 的反向计算分为两次扫描:
76+ 
77+```text
78+上游梯度 dout
79+
80+
81+bwd_fwd:前向重算 + Phi/Zeta 梯度
82+
83+ ├── states ──┐
84+ └── qk_dot ──┼──► bwd_bwd:倒序状态扫描 ──► q/k/v 等其余梯度
85+
86+ └── 与离散化中间量共同描述前向状态
87+```
88+ 
89+第一阶段不对状态递推求导,而是无梯度重算 `raw_y`。设 `r` 表示 MIMO rank,Phi 为
90+`mimo_o`,Zeta 为 `mimo_z`,则 reduce-O 路径的输出可写为:
91+ 
92+```text
93+u[b,s,h,r,p] = z[b,s,h,p] * Zeta[h,r,p]
94+gate = SiLU(u) # 未启用 Z 分支时为 1
95+out[b,s,h,p] = sum_r Phi[h,r,p] * gate * raw_y[b,s,h,r,p]
96+```
97+ 
98+因此 `dmimo_o``dmimo_z``dz` 只依赖本次重算得到的 `raw_y``dout`,可以在前向方向
99+完成;状态递推的反向依赖则留给 `bwd_bwd`。这种拆分与上游 combined backward 的缓存契约一致,
100+也避免在同一个 kernel 中同时维护正向和反向两条串行依赖链。
101+ 
102+每个 chunk 开始前的状态写入 `states[:, :, chunk]``qk_dot` 保存同一 token 上各 rank 的
103+query-key 点积:
104+ 
105+```text
106+qk_dot[b,h,s,r_out,r_in] = dot(q_bias_rot[b,h,s,r_out],
107+ k_bias_rot[b,h,s,r_in])
108+```
109+ 
110+rotate-half 是正交变换,同一位置同时旋转 q 和 k 不改变点积;实现可以复用旋转前 bias-add
111+结果计算该缓存,同时保持与后续反向公式相同的数学含义。
112+ 
113+## 实现说明
114+ 
115+staged 实现分为预处理、状态扫描和可选后处理:
116+ 
117+1. 预处理 kernel 按 `(B, H, chunk)` 并行,完成 bias-add、rotate-half,并计算或准备
118+ `qk_dot`。小网格会把 rotary 与 `qk_dot` 合在一次 launch 中。
119+2. scan kernel 以 `(B, H)` 为任务粒度顺序遍历 chunk。chunk 内先计算
120+ `gamma = dt * sigmoid(trap)` 及因果衰减,再组合上一状态、chunk 内 q-k 交互、Psi 投影与
121+ D-skip,得到各 rank 的 `raw_y`
122+3. 状态更新使用当前 chunk 的 key 与 Psi-value 投影生成下一状态。所有点积和归约使用 fp32
123+ 累加,写回时转换为 `output_dtype`
124+4. 对较长序列,非 rank-fold 路径可将 Phi/Zeta 收缩和 `dz` 移到并行后处理 kernel,降低串行
125+ scan 中的向量工作量;短序列保留内联计算,避免额外 launch 和 scratch 往返。
126+ 
127+aux 实现把计算进一步拆为三段:合并的 pre/aux kernel 并行生成旋转结果、`qk_dot`、chunk 内
128+输出 `INTRA` 和状态增量 `KV``scan_inter` 只处理跨 chunk 状态递推与 `dmimo_o`;dispatcher
129+仅在这条路径完整覆盖调用参数时选择它。
130+ 
131+生产入口根据 shape 选择两条实现路径:
132+ 
133+- 当未启用 Z 门控、`S` 能被 `chunk_size` 整除、单个 `(B, H)` 网格不能占满 AI Core 且
134+ `N * P <= 16384` 时,使用 batch/chunk 并行的 aux 实现补足并行度。
135+- 其余情况使用功能完整的 staged scan 实现。该路径支持 Z 门控、D-skip 和尾块处理。
136+ 
137+`S % chunk_size != 0` 时,host wrapper 将序列右填到完整 chunk。普通序列输入补零,
138+`da_cs` 延用最后一个有效值,并重建尾块的 `da_cs_rev``segsum`;kernel 执行结束后再把
139+带序列维的输出裁回原长度。
140+ 
141+## 支持范围
142+ 
143+- 硬件:Ascend arch32。
144+- dtype:`float16``bfloat16``float32`
145+- 布局:dense、非 varlen,`H % G == 0`
146+- rank:支持通用 `R`;已覆盖 `R = 1, 2, 4, 8`
147+- `N` 必须为偶数,且旋转维度不得超过 `N / 2`
148+- `mimo_o` 为必需输入;当前公开入口只提供 reduce-O 形式的 `dout`
149+- `mimo_z``z` 应同时提供。只提供其中一个不构成有效门控配置。
150+- 支持 reduce-O、可选 Z/SiLU 门控、可选 D-skip 和非整 chunk 的序列长度。
151+- 当前公开 API 不包含 `cu_seqlens``fuse_pregate_headwise_rms_norm`
152+ `return_final_state`
153+ 
154+## 调用示例
155+ 
156+```python
157+from mindspeed_ops.api.triton.mamba3_mimo_bwd_fwd_kernel import (
158+ mamba3_mimo_bwd_fwd_kernel,
159+)
160+ 
161+states, qk_dot, dmimo_o, dmimo_z, dz = mamba3_mimo_bwd_fwd_kernel(
162+ dout,
163+ q,
164+ k,
165+ v,
166+ q_bias,
167+ k_bias,
168+ mimo_v,
169+ mimo_o,
170+ angles,
171+ da_cs,
172+ da_cs_rev,
173+ dt,
174+ trap,
175+ segsum,
176+ mimo_z=mimo_z,
177+ d=d,
178+ z=z,
179+ chunk_size=16,
180+ output_dtype=torch.float32,
181+)
182+```
183+ 
184+`states``qk_dot` 作为 `mamba3_mimo_bwd_bwd_kernel` 的输入继续完成第二阶段反向扫描。
185+ 
186+## 代码与测试
187+ 
188+| 文件 | 内容 |
189+| --- | --- |
190+| `mindspeed_ops/api/triton/mamba3_mimo_bwd_fwd_kernel.py` | 公开 API 与架构检查 |
191+| `mindspeed_ops/arch32/triton/mamba3/mamba3_mimo_bwd_fwd_baseline_impl.py` | 优化前 Triton 基线 |
192+| `mindspeed_ops/arch32/triton/mamba3/mamba3_mimo_bwd_fwd_impl.py` | staged 生产实现 |
193+| `mindspeed_ops/arch32/triton/mamba3/mamba3_mimo_bwd_fwd_aux_impl.py` | 高并行度 aux 实现 |
194+| `mindspeed_ops/arch32/triton/mamba3/mamba3_mimo_bwd_fwd_dispatch_impl.py` | shape 自适应调度 |
195+| `tests/atk_tests/triton/mamba3_mimo_bwd_fwd_kernel/reference_impl.py` | PyTorch 小算子参考实现 |
196+| `tests/unit_tests/triton/test_mamba3_mimo_bwd_fwd_kernel.py` | 双标杆、官方 shape 网格和尾块测试 |
197+| `tests/atk_tests/triton/mamba3_mimo_bwd_fwd_kernel/` | ATK 用例配置与适配代码 |
198+ 
199+单元测试使用同一份 PyTorch 参考分别计算 fp32 reference 和 fp64 golden,并比较 Triton 输出与两套
200+标杆的误差。参考实现通过 PyTorch 小算子重建 chunk 状态,再用 autograd 独立计算 Phi、Zeta 和 z
201+的梯度,未复用生产 kernel 的梯度公式。
202+ 
203+| 用例组 | 覆盖内容 |
204+| --- | --- |
205+| pairwise 常规用例 | `B={1,2}``S={32,64}``H={2,4}`、fp16/bf16/fp32、Z、D、GQA |
206+| official quick | 从 11 组上游 shape 中选取 5 组,覆盖 `R={1,2,4}``N=16..256``P=64..128` |
207+| official full | `B=4, S=2048, H=16, G=1` 的 11 组完整网格,包含 R8;由环境变量开启 |
208+| tail | `S={33,40,48,50}``C={16,32}`,校验右填、离散量重建和输出裁剪 |
209+ 
210+逐项检查 `states``qk_dot``dmimo_o`,启用 Z 时额外检查 `dmimo_z``dz`
211+`qk_dot` 的真值接近零,专项网格采用全局最大绝对误差相对 golden 最大值的指标,避免逐元素相对误差
212+在零点附近失真。`dmimo_z` 也包含接近零的投影梯度,统一使用候选对 fp64 golden 的全局相对误差;
213+其余输出使用 dual-benchmark。完整 shape 网格通过 `MIMO22_RUN_OFFICIAL_FULL_GRID=1` 开启。
214+ 
215+```bash
216+pytest -q tests/unit_tests/triton/test_mamba3_mimo_bwd_fwd_kernel.py
217+```
218+ 
219+ATK 精度任务覆盖 24 个用例,当前记录为 24/24 通过。
@@ -0,0 +1,146 @@
1+# mamba3_mimo_bwd_fwd 优化记录
2+ 
3+## 基线与目标
4+ 
5+优化基线保存在
6+`mindspeed_ops/arch32/triton/mamba3/mamba3_mimo_bwd_fwd_baseline_impl.py`。该版本按上游
7+TileLang 计算顺序拆出 rotary、`qk_dot` 和串行 scan,主要用于精度对照和量化后续优化收益。
8+ 
9+生产实现保存在 `mamba3_mimo_bwd_fwd_impl.py``mamba3_mimo_bwd_fwd_aux_impl.py`
10+`mamba3_mimo_bwd_fwd_dispatch_impl.py`。优化没有改变公开 API、输出布局和 fp32 累加方式。
11+ 
12+基线与生产版均保留以下计算约束,性能比较不包含语义裁剪:
13+ 
14+- 相同的 bias-add 与 rotate-half 位置编码;
15+- 相同的 chunk 状态递推、Psi/Phi 投影和可选 D-skip;
16+- 相同的 `states``qk_dot``dmimo_o``dmimo_z``dz` 输出;
17+- 相同的 fp32 中间累加和 `output_dtype` 转换位置。
18+ 
19+## 瓶颈分析
20+ 
21+`B=4, S=2048, H=16, G=1` 的上游 shape 网格中,stage0 的 scan 占主要执行时间。
22+scan 必须沿 chunk 保持状态依赖,但每个 chunk 内又包含大量 `R x R` 小矩阵运算。profile 显示该路径
23+主要受标量地址生成和 tiny-dot 启动开销限制,Cube 与 Vector 的有效计算占比偏低。仅减少 HBM load
24+不能消除这一开销。
25+ 
26+另一个问题出现在较小的 `B * H`:按 `(batch, head)` 启动的串行 scan 无法占满 AI Core,设备上
27+仍有可用于 chunk-local 计算的并行度。
28+ 
29+基线还存在两个 shape 扩展问题:一是 `N * P` 较大时,递推状态与新状态同时驻留 UB 会超出容量;
30+二是 R6/R8 下静态展开 `R x R` rank 循环会导致编译期 IR 与 cbuf 需求快速增长。这两个问题虽不总是
31+表现为运行时热点,但决定了优化实现能否覆盖完整 shape 网格。
32+ 
33+## 优化项与收益
34+ 
35+### 1. shape 自适应的 aux 路径
36+ 
37+对不含 Z 门控、序列长度为完整 chunk、`B * H` 不足以占满 AI Core 且 `N * P <= 16384`
38+shape,将 chunk-local 计算拆到 `(B, H, nchunks)` 网格,再由轻量 scan 完成跨 chunk 状态递推。
39+该路径提高了小 batch、小 head 场景的核利用率;其余 shape 自动走 staged 实现,保持完整功能覆盖。
40+ 
41+dispatcher 的生产选择条件如下:
42+ 
43+| 条件 | 选择 |
44+| --- | --- |
45+| `z is None`、整 chunk、`B * H <= AI Core 数``N * P <= 16384` | aux + inter scan |
46+| 启用 Z、存在尾块、并行度已足够或 tile 超预算 | staged scan |
47+ 
48+aux 路径将 `INTRA[B,H,S,R,P]``KV[B,H,nchunks,N,P]` 作为并行阶段与串行阶段之间的中间量,
49+用额外 scratch 换取 chunk 维并行度。只有在 AI Core 欠占用时这笔交换才有收益,因此不作为无条件路径。
50+ 
51+### 2. rank-fold scan
52+ 
53+`R >= 4` 且未启用融合后处理时,将 rank 输入维与 chunk 维折叠为 `R * C`
54+ 
55+- 原先逐 rank 发起的 `R x R` 组小矩阵点积合并为较大的矩阵运算;
56+- 状态更新、`qk` 收缩和 `mimo_v` 投影复用折叠后的 tile;
57+- 保留 `r_out` 循环以控制 UB 使用量。
58+ 
59+该优化直接减少 tiny-dot 数量与地址生成次数,是 R4/R8 shape 的主要收益来源。
60+ 
61+折叠前,一个 chunk 内需要为多个 `(r_out, r_in)` 组合分别建立 block pointer 并启动点积;折叠后,
62+q-k、状态和 Psi-value 的主要收缩由少量 `[C, R*C]``[R*C, N]``[R*C, P]` 矩阵完成。
63+浮点乘加的结合顺序可能与 stage0 不同,因此正确性以 fp64 golden 而非逐位一致作为判断依据。
64+ 
65+### 3. 大 tile 分块
66+ 
67+`N * P` 超出单 tile 预算时,使用 `_bt` 路径沿 `N` 分块。当前分块优先将 `N` 划分为两个
68+block,并在 N-block 之间复用每个 chunk 的离散化量与投影中间量,避免因 P-blocking 重复执行标量逻辑。
69+该路径使 `N=256``P=128` 等 shape 在 UB 约束内继续使用 rank-fold。
70+ 
71+是否能使用单块 rank-fold 由 host 端按 tile 字节数估算。估算包含双份状态、折叠后的 q/k、输出投影、
72+因果系数和更新状态;若完整 `P` tile 超预算,先缩小 P block,最终路由到 N-blocked `_bt`。对于
73+非 rank-fold shape,仅 `N * P > 16384` 才进入大 tile 路径,避免让本可驻留 UB 的
74+`N256_P64``N128_P128` 承担不必要的 HBM 往返。
75+ 
76+### 4. 融合 `qk_dot`
77+ 
78+rank-fold scan 已经加载同一位置的 query 和 key,因此直接在 scan 内写出 `qk_dot`,省去独立
79+`qk_dot` kernel。对 `R=8`,独立 kernel 中的 rank 组合最多,这项融合的收益最明显。
80+ 
81+小预处理网格仍可选择 rotary + `qk_dot` 合并 kernel;大网格若不走 rank-fold,则保留独立
82+`qk_dot` kernel。这样既减少 launch,也避免在大网格上因融合而重复读取 key。
83+ 
84+### 5. rank 循环 de-unroll
85+ 
86+大 rank 的 `static_range` 会显著放大编译期 IR 和片上资源需求。实现对 `R >= 6` 使用运行时 rank
87+循环,在不改变浮点运算顺序的前提下避免 R8 编译失败;较小 rank 仍保留静态展开。
88+ 
89+这项改动本身用于解除编译上限,不把“原始 R8 stage0 无法编译”计作无限加速。性能表中的 R8
90+基线应用同样的 de-unroll,只比较后续 rank-fold 与 `qk_dot` 融合带来的运行时收益。
91+ 
92+### 6. 串行 scan 的局部优化
93+ 
94+未进入 rank-fold 的 staged 路径仍按 shape 使用两项低风险优化:`R <= 4``R != 2` 时把
95+`mimo_v` 提到 chunk 循环外复用;`C <= 16` 时直接加载离散化量,减少 scan 内重建工作。
96+这些改动主要改善 R1/R3 和小 chunk,不改变 rank-fold 的调度范围。
97+ 
98+## 性能结果
99+ 
100+以下结果在 Ascend 910B 单卡、bf16、上游 `B=4, S=2048, H=16, G=1` shape 网格上测得。
101+候选实现与 stage0 基线在同卡交错执行,每个 shape 取 3 轮结果。R8 的 stage0 使用仅将 rank 循环
102+de-unroll 的可编译等价版本作为基线。
103+ 
104+| Shape (`N_P_R_C`) | 相对 stage0 |
105+| --- | ---: |
106+| `N16_P64_R4_C8` | 2.13x |
107+| `N32_P64_R4_C16` | 2.05x |
108+| `N64_P64_R4_C16` | 2.12x |
109+| `N128_P64_R4_C16` | 1.98x |
110+| `N256_P64_R4_C8` | 1.52x |
111+| `N64_P128_R4_C16` | 1.99x |
112+| `N128_P32_R4_C16` | 2.13x |
113+| `N128_P128_R4_C8` | 1.78x |
114+| `N128_P64_R8_C8` | 7.12x |
115+| `N128_P64_R2_C32` | 1.03x |
116+| `N128_P64_R1_C64` | 0.95x |
117+ 
118+11 组 shape 的几何平均加速比为 **1.94x**。R1/R2 不满足 rank-fold 调度条件,作为未使用该优化的
119+对照组保留在统计中。
120+ 
121+结果可以分为三类理解:
122+ 
123+- R4 常规 shape 稳定在 `1.98x–2.13x`,说明 rank-fold 对常见矩阵尺寸的收益稳定。
124+- `N256_P64``N128_P128` 需要更保守的 tile,分块和 HBM scratch 抵消了一部分收益,但仍达到
125+ `1.52x``1.78x`
126+- R8 同时减少大量 rank 组合的 tiny-dot 与独立 `qk_dot` 工作,因此达到 `7.12x`。R1/R2 不使用
127+ rank-fold,结果接近基线,符合调度预期。
128+ 
129+除完整 11-shape 网格外,四组常用生产 shape 的几何平均加速比记录为 `1.224x`。该组包含一个
130+`B=8, H=16`、AI Core 已饱和的控制 shape,因此 aux 并行化不会被错误地计为普遍收益。
131+ 
132+## 精度与边界验证
133+ 
134+- stage0 与生产实现使用同一套 PyTorch 小算子参考进行 fp32/fp64 双标杆校验。
135+- 官方 quick shape、常规 dtype/特性组合以及非整 chunk 尾块均由 UT 覆盖。
136+- `qk_dot` 的真值接近零,测试使用 `max(abs(diff)) / max(abs(golden))` 的全局相对误差,避免逐元素
137+ 相对误差放大近零噪声。
138+- `dmimo_z` 同样含有近零投影梯度,使用候选对 fp64 golden 的全局最大相对误差;其它梯度仍使用
139+ fp32/fp64 dual-benchmark。
140+- ATK 精度任务为 24/24 通过。
141+ 
142+尾块测试单独验证 host padding 契约:q/k/v/dout/angles/z 补零,`da_cs` 使用末值延展,随后重建
143+`da_cs_rev``segsum`。这保证性能 kernel 仍只处理整 chunk,而任意正序列长度由公开入口完整承接。
144+ 
145+性能数据只比较能够保持相同输入、输出和计算语义的实现;生产路径不使用 PyTorch reference 或其他
146+整算子作为运行时回退。
@@ -0,0 +1,155 @@
1+# mamba3_mimo_bwd_fwd Triton-Ascend 实践
2+ 
3+## 背景
4+ 
5+Mamba-3 MIMO 的 combined backward 沿用“前向重算 + 反向扫描”的两阶段结构。
6+`mamba3_mimo_bwd_fwd` 虽然属于 backward 流程,但计算方向仍是从序列头到序列尾:它重新建立
7+每个 chunk 的进入状态,计算输出投影与门控梯度,并为第二阶段保存 `states``qk_dot`
8+ 
9+上游实现以 GPU TileLang kernel 为基础。迁移到 Triton-Ascend 时,核心工作不是逐行替换语法,而是
10+重新确定任务粒度、UB 中的常驻数据和 Vector/Cube 的职责边界。本实践文档记录这部分实现选择;
11+稳定接口见 `mamba3_mimo_bwd_fwd.md`,性能演进见 `mamba3_mimo_bwd_fwd_optimization.md`
12+ 
13+## 计算语义拆解
14+ 
15+### GQA 与旋转位置编码
16+ 
17+输入 q/k 的 head 维为 `G`,状态和 v 的 head 维为 `H`。每个 attention head 使用:
18+ 
19+```text
20+g = h // (H / G)
21+q_h = q[..., g, :] + q_bias[h, ...]
22+k_h = k[..., g, :] + k_bias[h, ...]
23+```
24+ 
25+随后对前 `RD = N / rotary_dim_divisor` 维执行 rotate-half。第 `n` 维与第 `N/2+n` 维配对:
26+ 
27+```text
28+q_rot[n] = cos(angle[n]) * q[n] - sin(angle[n]) * q[N/2+n]
29+q_rot[N/2+n] = sin(angle[n]) * q[n] + cos(angle[n]) * q[N/2+n]
30+```
31+ 
32+k 使用相同变换。旋转以 fp32 计算,并写入 `[B,H,S,R,N]` scratch,后续 scan 直接按 head 连续读取。
33+ 
34+### chunk 内输出与状态更新
35+ 
36+设当前 chunk 的长度为 `C`,将 `(time, rank)` 展平为 `C*R`。离散化链的两个主要量为:
37+ 
38+```text
39+gamma[t] = dt[t] * sigmoid(trap[t])
40+trap_scale[t] = gamma[t] + dt[t+1] * sigmoid(-trap[t+1])
41+```
42+ 
43+序列末尾没有后一项,第二项取零。每个 rank 的 `PsiV``v * mimo_v``raw_y` 由三部分组成:
44+ 
45+1. 当前 q 与 chunk 进入状态的交互;
46+2. chunk 内严格因果的 q-k 与 PsiV 交互;
47+3. 同位置 q-k 对角项,以及可选 D-skip。
48+ 
49+chunk 结束时,key 按 `trap_scale` 与反向累计衰减缩放,再与 `PsiV` 收缩得到状态增量:
50+ 
51+```text
52+state_next = state_in * exp(da_cs_sum) + K_state.T @ PsiV
53+```
54+ 
55+`state_in` 在计算前写入 `states[:, :, chunk]`。Z 门控不参与状态更新,只作用于 `raw_y` 到最终输出
56+的投影,因此 Phi、Zeta 和 z 的梯度可以在本阶段完成。
57+ 
58+## 从基线到生产实现
59+ 
60+### stage0 切分
61+ 
62+优化前实现保留在 `mamba3_mimo_bwd_fwd_baseline_impl.py`,由三个步骤组成:
63+ 
64+| 步骤 | 作用 |
65+| --- | --- |
66+| rotary | bias-add 与 rotate-half,生成 Qr/Kr |
67+| qkdot | 计算每个 token 的 `R x R` q-k 点积 |
68+| scan | 按 `(B,H)` 顺序扫描 chunk,生成状态、投影梯度和门控梯度 |
69+ 
70+这种切分先保证数学路径完整,并使 PyTorch reference、stage0 和生产实现共享相同的输入输出契约。
71+ 
72+### staged 生产路径
73+ 
74+生产实现仍保留 `(B,H)` 串行 scan,但根据 shape 选择预处理与 scan 变体:
75+ 
76+| kernel | 使用场景 |
77+| --- | --- |
78+| `_mamba3_mimo_rotary_qkdot_kernel` | 小预处理网格,合并 rotary 与 qk_dot |
79+| `_mamba3_mimo_rotary_kernel` + `_mamba3_mimo_qkdot_kernel` | 大网格,避免融合后重复读 key |
80+| `_mamba3_mimo_bwd_fwd_scan_kernel` | 常规 R1/R2/R3 路径 |
81+| `_mamba3_mimo_bwd_fwd_scan_rf_kernel` | 可单 tile 驻留的 R4/R8 rank-fold 路径 |
82+| `_mamba3_mimo_bwd_fwd_scan_kernel_bt` | 大 N/P 的 N-blocked 路径 |
83+| `*_post_kernel` | 长序列的 Phi/Zeta/dz 并行收缩 |
84+ 
85+### aux 生产路径
86+ 
87+`B*H` 不能占满 AI Core 时,仅增加 scan 内部优化无法获得足够并行度。aux 路径将每个 chunk
88+可独立计算的部分放到 `(B,H,nchunks)` 网格:
89+ 
90+```text
91+pre_aux ──► INTRA, KV, qk_dot, 离散化 scratch
92+
93+
94+ scan_inter ──► states, dmimo_o
95+```
96+ 
97+`INTRA` 保存不依赖前一 chunk 状态的输出,`KV` 保存状态增量。`scan_inter` 只组合前一状态与这两项。
98+这条路径不处理 Z 门控和非整 chunk,因此 dispatcher 只在参数契约完整匹配时启用。
99+ 
100+## Ascend 侧实现要点
101+ 
102+### UB 预算
103+ 
104+scan 同时需要旧状态和新状态,单份状态占 `N * P * 4` 字节。若直接让两份 `[N,P]` fp32 tile
105+常驻,`N256_P128` 仅状态就需要 256 KiB,尚未计入 q/k、PsiV 和输出 tile。生产实现以
106+`N * P > 16384` 作为常规大 tile 边界,并在 rank-fold 的 host 计划中进一步估算折叠矩阵占用。
107+ 
108+大 tile 路径把 running state 放到 HBM scratch,按 N block 流式更新。该做法增加 HBM 往返,但把
109+UB 峰值限制在 `BN * P`,是覆盖大 shape 所需的容量交换。
110+ 
111+### rank 维与对齐
112+ 
113+典型 `R=2` 的 fp32 rank 尾轴只有 8 字节,不适合作为独立的一维搬运单位。实现不把 rank 作为
114+最内层窄向量,而是使用 `[C,N]``[C,P]` 二维 tile,并在 rank 循环中读写。R4 以上的矩阵收缩
115+则将 rank 与 chunk 折为 `R*C`,扩大点积规模并减少地址生成。
116+ 
117+### 编译期展开
118+ 
119+`static_range(R)` 对小 rank 有利,但 R6/R8 的 `R x R` 展开会显著增加 IR 与 cbuf 需求。host 对
120+`R >= 6` 选择运行时 rank 循环;aux 预处理在大 N 或窄 P 的 R4 shape 也使用 de-unroll 版本,以控制
121+UB 和编译规模。该选择不改变循环内浮点运算顺序。
122+ 
123+### `qk_dot` 的计算位置
124+ 
125+同一位置对 q/k 同时应用正交旋转不会改变点积,因此 `qk_dot` 可以从 bias-add 后的 q/k 直接计算。
126+低精度输入使用向量 reduce,fp32 输入使用 dot,以匹配 reference 的累加误差。rank-fold scan 已经
127+持有相关 q/k tile 时直接写 `qk_dot`,避免额外 kernel。
128+ 
129+### 非整 chunk
130+ 
131+kernel 内部只处理完整 chunk。公开入口在 host 侧补齐 q/k/v/dout/angles/z,使用末值延展
132+`da_cs`,并重建 `da_cs_rev``segsum`。这种做法避免在每个点积上引入尾块分支,同时由 tail UT
133+验证补齐区不会影响有效 token。
134+ 
135+## 数值与验证方法
136+ 
137+所有状态、点积和梯度归约使用 fp32;`output_dtype` 只控制最终写回。验证分三层:
138+ 
139+1. PyTorch 小算子以 fp64 重算 `raw_y/states/qk_dot` 并通过 autograd 求投影、门控梯度;
140+2. stage0 用相同输入输出契约验证迁移前的 Triton 计算;
141+3. 生产实现通过 pairwise、上游 shape 网格和 tail 用例与 fp32/fp64 双标杆比较。
142+ 
143+`qk_dot` 在随机输入下可能接近零,使用全局最大绝对误差相对 golden 最大值的指标;其它输出沿用
144+仓库的 dual-benchmark 判据。`dmimo_z` 也会出现少量近零元素,因此同样直接约束候选对 fp64 golden
145+的全局最大相对误差,避免逐元素相对误差掩盖整体 RMSE。生产 API 不导入测试 reference,也不在
146+不支持的架构上静默切换实现。
147+ 
148+## 扩展实现时的检查项
149+ 
150+- 新增输入分支时,同时确认 dispatcher 的 aux 条件是否仍完整覆盖其语义。
151+- 修改状态更新前,检查 `states` 保存的是 chunk 进入态而不是离开态。
152+- 调整 tile 后重新核算双状态、折叠 q/k、PsiV 和归约 scratch 的总 UB 占用。
153+- 修改 rank 循环后至少覆盖 R1/R2/R4/R8,并分别检查编译时间和数值结果。
154+- 修改尾块填充值时同步更新 `da_cs_rev/segsum` 重建逻辑与 tail reference。
155+- 性能对比使用仓内 stage0,不使用 PyTorch eager 代替优化前 Triton 基线。
@@ -0,0 +1,98 @@
1+# Copyright (c) 2026, Huawei Technologies Co., Ltd. All rights reserved.
2+# Copyright (c) 2025, Dao AI Lab, Goombalab
3+# pylint: disable=duplicate-code
4+# The public API and production dispatcher intentionally forward the same operator contract.
5+ 
6+from __future__ import annotations
7+ 
8+import torch
9+ 
10+from mindspeed_ops.api.triton.utils import input_guard
11+from mindspeed_ops.utils import is_arch35
12+ 
13+__all__ = ["mamba3_mimo_bwd_fwd_kernel"]
14+ 
15+ 
16+@input_guard
17+def mamba3_mimo_bwd_fwd_kernel(
18+ dout: torch.Tensor,
19+ q: torch.Tensor,
20+ k: torch.Tensor,
21+ v: torch.Tensor,
22+ q_bias: torch.Tensor,
23+ k_bias: torch.Tensor,
24+ mimo_v: torch.Tensor,
25+ mimo_o: torch.Tensor,
26+ angles: torch.Tensor,
27+ da_cs: torch.Tensor,
28+ da_cs_rev: torch.Tensor,
29+ dt: torch.Tensor,
30+ trap: torch.Tensor,
31+ segsum: torch.Tensor,
32+ mimo_z: torch.Tensor | None = None,
33+ d: torch.Tensor | None = None,
34+ z: torch.Tensor | None = None,
35+ chunk_size: int = 16,
36+ rotary_dim_divisor: int = 4,
37+ output_dtype: torch.dtype = torch.float32,
38+):
39+ """Run the Mamba3 MIMO backward first-pass (bwd_fwd) Triton-Ascend kernel.
40+ 
41+ 迁移自 state-spaces/mamba ``mamba3_mimo_bwd.py`` 的 ``mamba_mimo_bwd_fwd_kernel``
42+ (MindSpeed-Ops #22)。按 chunk 重算 MIMO 前向中间量(逐 rank 下投影前输出 raw_y),据此
43+ 就地累加投影梯度并缓存递推状态 / QK 对角块供第二遍反向 (bwd_bwd, #23) 使用。覆盖路径:
44+ reduceO、通用 R、任意 S(尾块经 host tail_len 包装右填→跑核→切回)、可选 Z(SiLU)门控、
45+ 可选 D-skip、GQA;不含 fuse_pregate_headwise_rms_norm(核内已实现但本 API 未暴露)。
46+ 
47+ Args:
48+ dout (torch.Tensor): [B, S, H, P] 上游梯度 (reduceO 布局).
49+ q, k (torch.Tensor): [B, S, R, G, N].
50+ v (torch.Tensor): [B, S, H, P].
51+ q_bias, k_bias (torch.Tensor): [H, R, N] 旋转前偏置.
52+ mimo_v (torch.Tensor): [H, R, P] X 投影 (Psi).
53+ mimo_o (torch.Tensor): [H, R, P] O 投影 (Phi); 本最小路径必需.
54+ angles (torch.Tensor): [B, S, H, N // rotary_dim_divisor] 旋转角.
55+ da_cs, da_cs_rev, dt, trap (torch.Tensor): [B, H, S] 离散化量.
56+ segsum (torch.Tensor): [B, H, nchunks, chunk_size, chunk_size] chunk 内段和.
57+ mimo_z (torch.Tensor | None): [H, R, P] Z 投影 (Zeta); 配合 z 做 SiLU 门控.
58+ d (torch.Tensor | None): [H] D-skip.
59+ z (torch.Tensor | None): [B, S, H, P] 门控输入.
60+ chunk_size (int): 分块大小, 默认 16 (S 无需整除: 尾块由 host tail_len 包装处理).
61+ rotary_dim_divisor (int): 旋转维度除子, 默认 4.
62+ output_dtype (torch.dtype): 输出 dtype, 默认 float32.
63+ 
64+ Returns:
65+ tuple:
66+ states (torch.Tensor): [B, H, nchunks, N, P] 每 chunk 进入态.
67+ qk_dot (torch.Tensor): [B, H, S, R, R] 逐步 q·k 对角块.
68+ dmimo_o (torch.Tensor): [B, H, R, P] Phi 梯度 (逐 batch, 调用方再 ``sum(dim=0)``).
69+ dmimo_z (torch.Tensor | None): [B, H, R, P] Zeta 梯度 (hasZ).
70+ dz (torch.Tensor | None): [B, S, H, P] 门控输入梯度 (hasZ).
71+ """
72+ if is_arch35():
73+ raise NotImplementedError("mamba3_mimo_bwd_fwd_kernel is currently implemented for arch32 only")
74+ 
75+ from mindspeed_ops.arch32.triton.mamba3.mamba3_mimo_bwd_fwd_dispatch_impl import mamba3_mimo_bwd_fwd_prod
76+ 
77+ return mamba3_mimo_bwd_fwd_prod(
78+ dout=dout,
79+ q=q,
80+ k=k,
81+ v=v,
82+ q_bias=q_bias,
83+ k_bias=k_bias,
84+ mimo_v=mimo_v,
85+ mimo_o=mimo_o,
86+ angles=angles,
87+ da_cs=da_cs,
88+ da_cs_rev=da_cs_rev,
89+ dt=dt,
90+ trap=trap,
91+ segsum=segsum,
92+ mimo_z=mimo_z,
93+ d=d,
94+ z=z,
95+ chunk_size=int(chunk_size),
96+ rotary_dim_divisor=int(rotary_dim_divisor),
97+ output_dtype=output_dtype,
98+ )
@@ -0,0 +1,1203 @@
1+# Copyright (c) 2026, Dao AI Lab, Goombalab
2+# Copyright (c) 2026, Huawei Technologies Co., Ltd.
3+# pylint: disable=duplicate-code,too-many-lines
4+# Auxiliary and optimized kernels intentionally share formula-level Triton building blocks.
5+#
6+# SPDX-License-Identifier: Apache-2.0
7+#
8+# #22 batch-parallel-aux fast path (the #23 four-stage trick applied to #22).
9+#
10+# Measurement (this repo + two probe agents): the serial B*H scan is scalar-
11+# bound and, at underutilized shapes (B*H < ncore), the chunk-LOCAL matmuls
12+# (qk = Q@Kᵀ, intra = (qk*coeff)@psiv, KV = Kᵀ@psiv) sit on the serial critical
13+# path even though they do NOT depend on the recurrent state. This moves ALL of
14+# them to a B*H*nchunks-parallel ``aux`` kernel that fills every core, leaving the
15+# serial scan with only the state-dependent inter = Q@state + recurrence + dphi.
16+# Validated to help golden/big/s2048 (underutilized); b8h16 (saturated) keeps the
17+# staged path via the dispatcher. Numerics match the staged baseline.
18+# ruff: noqa: E501
19+from __future__ import annotations
20+ 
21+import torch
22+import triton
23+import triton.language as tl
24+ 
25+from mindspeed_ops.api.triton.utils import get_vector_num, get_aicore_num, input_guard
26+ 
27+ 
28+# --- inlined from bwd_fwd_fast2 (aux 唯一依赖: pre-rope+qkdot+离散化预算核) ---
29+@triton.jit
30+def _pre_rope_qkdot_disc_kernel(
31+ # rotary + qkdot
32+ Q,
33+ K,
34+ Q_BIAS,
35+ K_BIAS,
36+ ANGLES,
37+ Qr,
38+ Kr,
39+ QK_DOT,
40+ # disc
41+ DT,
42+ TRAP,
43+ DA_CS,
44+ COEFF,
45+ TS_REV,
46+ EXP_DA_CS,
47+ DECAY,
48+ sq_b,
49+ sq_s,
50+ sq_r,
51+ sq_g,
52+ sq_n,
53+ sqb_h,
54+ sqb_r,
55+ sqb_n,
56+ sa_b,
57+ sa_s,
58+ sa_h,
59+ sa_n,
60+ sqr_b,
61+ sqr_h,
62+ sqr_s,
63+ sqr_r,
64+ sqr_n,
65+ sqk_b,
66+ sqk_h,
67+ sqk_s,
68+ sqk_ro,
69+ sqk_ri,
70+ sdisc_b,
71+ sdisc_h,
72+ sdisc_s,
73+ sco_b,
74+ sco_h,
75+ sco_c,
76+ sco_i,
77+ sco_j,
78+ sde_b,
79+ sde_h,
80+ sde_c,
81+ B,
82+ S,
83+ H,
84+ G,
85+ num_chunks,
86+ R: tl.constexpr,
87+ N: tl.constexpr,
88+ RD: tl.constexpr,
89+ CHUNK: tl.constexpr,
90+ HALF: tl.constexpr,
91+ QK_USE_DOT: tl.constexpr = True,
92+):
93+ """Fused parallel (B*H*nchunks) prepass: bias+RoPE -> Qr/Kr, same-token qk_dot
94+ diag cache, AND the discretization scalars (gamma/trap_scale/exp/coeff) hoisted
95+ off the serial scan. One launch does all the scalar-heavy chunk-local work.
96+ """
97+ pid = tl.program_id(0)
98+ npg = tl.num_programs(0)
99+ total = B * H * num_chunks
100+ rep = H // G
101+ offs_n = tl.arange(0, N)
102+ offs_c = tl.arange(0, CHUNK)
103+ causal = offs_c[:, None] > offs_c[None, :]
104+ diagm = offs_c[:, None] == offs_c[None, :]
105+ 
106+ for wi in range(pid, total, npg):
107+ c = wi % num_chunks
108+ tmp = wi // num_chunks
109+ i_h = tmp % H
110+ i_b = tmp // H
111+ i_hqk = i_h // rep
112+ s0 = c * CHUNK
113+ s_c = s0 + offs_c
114+ mask_c = s_c < S
115+ 
116+ p_ang = tl.make_block_ptr(ANGLES + i_b * sa_b + i_h * sa_h, (S, RD), (sa_s, sa_n), (s0, 0), (CHUNK, RD), (1, 0))
117+ ang = tl.load(p_ang, boundary_check=(0,), padding_option="zero").to(tl.float32)
118+ cos = tl.cos(ang)
119+ sin = tl.sin(ang)
120+ 
121+ for r_out in tl.static_range(R):
122+ q_in = Q + i_b * sq_b + i_hqk * sq_g + r_out * sq_r
123+ k_in = K + i_b * sq_b + i_hqk * sq_g + r_out * sq_r
124+ q_out = Qr + i_b * sqr_b + i_h * sqr_h + r_out * sqr_r
125+ k_out = Kr + i_b * sqr_b + i_h * sqr_h + r_out * sqr_r
126+ qb = tl.load(Q_BIAS + i_h * sqb_h + r_out * sqb_r + offs_n * sqb_n).to(tl.float32)
127+ kb = tl.load(K_BIAS + i_h * sqb_h + r_out * sqb_r + offs_n * sqb_n).to(tl.float32)
128+ p_q = tl.make_block_ptr(q_in, (S, N), (sq_s, sq_n), (s0, 0), (CHUNK, N), (1, 0))
129+ p_k = tl.make_block_ptr(k_in, (S, N), (sq_s, sq_n), (s0, 0), (CHUNK, N), (1, 0))
130+ qf = tl.load(p_q, boundary_check=(0,), padding_option="zero").to(tl.float32) + qb[None, :]
131+ kf = tl.load(p_k, boundary_check=(0,), padding_option="zero").to(tl.float32) + kb[None, :]
132+ # RoPE rotate-half (cols [0,RD) with [HALF,HALF+RD))
133+ offs_d = tl.arange(0, RD)
134+ qba = tl.load(Q_BIAS + i_h * sqb_h + r_out * sqb_r + offs_d * sqb_n).to(tl.float32)
135+ qbb = tl.load(Q_BIAS + i_h * sqb_h + r_out * sqb_r + (HALF + offs_d) * sqb_n).to(tl.float32)
136+ kba = tl.load(K_BIAS + i_h * sqb_h + r_out * sqb_r + offs_d * sqb_n).to(tl.float32)
137+ kbb = tl.load(K_BIAS + i_h * sqb_h + r_out * sqb_r + (HALF + offs_d) * sqb_n).to(tl.float32)
138+ p_qa = tl.make_block_ptr(q_in, (S, N), (sq_s, sq_n), (s0, 0), (CHUNK, RD), (1, 0))
139+ p_qb2 = tl.make_block_ptr(q_in, (S, N), (sq_s, sq_n), (s0, HALF), (CHUNK, RD), (1, 0))
140+ p_ka = tl.make_block_ptr(k_in, (S, N), (sq_s, sq_n), (s0, 0), (CHUNK, RD), (1, 0))
141+ p_kb2 = tl.make_block_ptr(k_in, (S, N), (sq_s, sq_n), (s0, HALF), (CHUNK, RD), (1, 0))
142+ qa = tl.load(p_qa, boundary_check=(0,), padding_option="zero").to(tl.float32) + qba[None, :]
143+ qb2 = tl.load(p_qb2, boundary_check=(0,), padding_option="zero").to(tl.float32) + qbb[None, :]
144+ ka = tl.load(p_ka, boundary_check=(0,), padding_option="zero").to(tl.float32) + kba[None, :]
145+ kb2 = tl.load(p_kb2, boundary_check=(0,), padding_option="zero").to(tl.float32) + kbb[None, :]
146+ qar = cos * qa - sin * qb2
147+ qbr = sin * qa + cos * qb2
148+ kar = cos * ka - sin * kb2
149+ kbr = sin * ka + cos * kb2
150+ p_qr = tl.make_block_ptr(q_out, (S, N), (sqr_s, sqr_n), (s0, 0), (CHUNK, N), (1, 0))
151+ p_kr = tl.make_block_ptr(k_out, (S, N), (sqr_s, sqr_n), (s0, 0), (CHUNK, N), (1, 0))
152+ tl.store(p_qr, qf.to(p_qr.dtype.element_ty), boundary_check=(0,))
153+ tl.store(p_kr, kf.to(p_kr.dtype.element_ty), boundary_check=(0,))
154+ p_qr_a = tl.make_block_ptr(q_out, (S, N), (sqr_s, sqr_n), (s0, 0), (CHUNK, RD), (1, 0))
155+ p_qr_b = tl.make_block_ptr(q_out, (S, N), (sqr_s, sqr_n), (s0, HALF), (CHUNK, RD), (1, 0))
156+ p_kr_a = tl.make_block_ptr(k_out, (S, N), (sqr_s, sqr_n), (s0, 0), (CHUNK, RD), (1, 0))
157+ p_kr_b = tl.make_block_ptr(k_out, (S, N), (sqr_s, sqr_n), (s0, HALF), (CHUNK, RD), (1, 0))
158+ tl.store(p_qr_a, qar.to(p_qr_a.dtype.element_ty), boundary_check=(0,))
159+ tl.store(p_qr_b, qbr.to(p_qr_b.dtype.element_ty), boundary_check=(0,))
160+ tl.store(p_kr_a, kar.to(p_kr_a.dtype.element_ty), boundary_check=(0,))
161+ tl.store(p_kr_b, kbr.to(p_kr_b.dtype.element_ty), boundary_check=(0,))
162+ # qk_dot diag (rotation invariant -> use bias-added qf/kfd)
163+ for r_in in tl.static_range(R):
164+ if r_in == r_out:
165+ kfd = kf
166+ else:
167+ kbd = tl.load(K_BIAS + i_h * sqb_h + r_in * sqb_r + offs_n * sqb_n).to(tl.float32)
168+ p_kd = tl.make_block_ptr(
169+ K + i_b * sq_b + i_hqk * sq_g + r_in * sq_r, (S, N), (sq_s, sq_n), (s0, 0), (CHUNK, N), (1, 0)
170+ )
171+ kfd = tl.load(p_kd, boundary_check=(0,), padding_option="zero").to(tl.float32) + kbd[None, :]
172+ qkmat = tl.dot(qf, tl.trans(kfd))
173+ qkdiag = tl.sum(tl.where(diagm, qkmat, 0.0), axis=1)
174+ tl.store(
175+ QK_DOT + i_b * sqk_b + i_h * sqk_h + s_c * sqk_s + r_out * sqk_ro + r_in * sqk_ri,
176+ qkdiag.to(QK_DOT.dtype.element_ty),
177+ mask=mask_c,
178+ )
179+ 
180+ # --- discretization (hoisted off the serial scan) ---
181+ base = i_b * sdisc_b + i_h * sdisc_h
182+ dt = tl.load(DT + base + s_c * sdisc_s).to(tl.float32)
183+ trap = tl.load(TRAP + base + s_c * sdisc_s).to(tl.float32)
184+ gamma = dt * tl.sigmoid(trap)
185+ sh = s_c + 1
186+ shm = sh < S
187+ dt_sh = tl.load(DT + base + sh * sdisc_s, mask=shm, other=0.0).to(tl.float32)
188+ trap_sh = tl.load(TRAP + base + sh * sdisc_s, mask=shm, other=0.0).to(tl.float32)
189+ guard = (s_c < (S - 1)).to(tl.float32)
190+ trap_scale = gamma + dt_sh * tl.sigmoid(-trap_sh) * guard
191+ dacs_v = tl.load(DA_CS + base + s_c * sdisc_s).to(tl.float32)
192+ exp_da_cs = tl.exp(dacs_v)
193+ da_cs_sum = tl.load(DA_CS + base + (s0 + CHUNK - 1) * sdisc_s).to(tl.float32)
194+ exp_rev = tl.exp(da_cs_sum - dacs_v)
195+ seg = dacs_v[:, None] - dacs_v[None, :]
196+ coeff = (
197+ causal.to(tl.float32) * trap_scale[None, :] * tl.exp(tl.where(causal, seg, 0.0))
198+ + diagm.to(tl.float32) * gamma[:, None]
199+ )
200+ ts_rev = trap_scale * exp_rev
201+ p_co = tl.make_block_ptr(
202+ COEFF + i_b * sco_b + i_h * sco_h + c * sco_c,
203+ (CHUNK, CHUNK),
204+ (sco_i, sco_j),
205+ (0, 0),
206+ (CHUNK, CHUNK),
207+ (1, 0),
208+ )
209+ tl.store(p_co, coeff)
210+ tl.store(TS_REV + base + s_c * sdisc_s, ts_rev)
211+ tl.store(EXP_DA_CS + base + s_c * sdisc_s, exp_da_cs)
212+ tl.store(DECAY + i_b * sde_b + i_h * sde_h + c * sde_c, tl.exp(da_cs_sum))
213+ 
214+ 
215+@triton.jit
216+def _aux_kernel(
217+ Qr,
218+ Kr,
219+ V,
220+ MIMO_V,
221+ COEFF,
222+ TS_REV,
223+ INTRA,
224+ KV,
225+ sqr_b,
226+ sqr_h,
227+ sqr_s,
228+ sqr_r,
229+ sqr_n,
230+ sv_b,
231+ sv_s,
232+ sv_h,
233+ sv_p,
234+ smv_h,
235+ smv_r,
236+ smv_p,
237+ sdisc_b,
238+ sdisc_h,
239+ sdisc_s,
240+ sco_b,
241+ sco_h,
242+ sco_c,
243+ sco_i,
244+ sco_j,
245+ sin_b,
246+ sin_h,
247+ sin_s,
248+ sin_r,
249+ sin_p,
250+ skv_b,
251+ skv_h,
252+ skv_c,
253+ skv_n,
254+ skv_p,
255+ B,
256+ S,
257+ H,
258+ num_chunks,
259+ R: tl.constexpr,
260+ N: tl.constexpr,
261+ P: tl.constexpr,
262+ CHUNK: tl.constexpr,
263+):
264+ pid = tl.program_id(0)
265+ npg = tl.num_programs(0)
266+ total = B * H * num_chunks
267+ offs_c = tl.arange(0, CHUNK)
268+ offs_p = tl.arange(0, P)
269+ for wi in range(pid, total, npg):
270+ c = wi % num_chunks
271+ tmp = wi // num_chunks
272+ i_h = tmp % H
273+ i_b = tmp // H
274+ disc_base = i_b * sdisc_b + i_h * sdisc_h
275+ s0 = c * CHUNK
276+ s_c = s0 + offs_c
277+ mask_c = s_c < S
278+ coeff = tl.load(
279+ tl.make_block_ptr(
280+ COEFF + i_b * sco_b + i_h * sco_h + c * sco_c,
281+ (CHUNK, CHUNK),
282+ (sco_i, sco_j),
283+ (0, 0),
284+ (CHUNK, CHUNK),
285+ (1, 0),
286+ )
287+ )
288+ ts_rev = tl.load(TS_REV + disc_base + s_c * sdisc_s)[:, None]
289+ p_v = tl.make_block_ptr(V + i_b * sv_b + i_h * sv_h, (S, P), (sv_s, sv_p), (s0, 0), (CHUNK, P), (1, 0))
290+ v_tile = tl.load(p_v, boundary_check=(0,), padding_option="zero").to(tl.float32)
291+ kv = tl.zeros((N, P), dtype=tl.float32)
292+ for r_out in tl.static_range(R):
293+ p_q = tl.make_block_ptr(
294+ Qr + i_b * sqr_b + i_h * sqr_h + r_out * sqr_r, (S, N), (sqr_s, sqr_n), (s0, 0), (CHUNK, N), (1, 0)
295+ )
296+ q_ro = tl.load(p_q, boundary_check=(0,), padding_option="zero")
297+ o_intra = tl.zeros((CHUNK, P), dtype=tl.float32)
298+ for r_in in tl.static_range(R):
299+ p_k = tl.make_block_ptr(
300+ Kr + i_b * sqr_b + i_h * sqr_h + r_in * sqr_r, (S, N), (sqr_s, sqr_n), (s0, 0), (CHUNK, N), (1, 0)
301+ )
302+ k_ri = tl.load(p_k, boundary_check=(0,), padding_option="zero")
303+ mv_i = tl.load(MIMO_V + i_h * smv_h + r_in * smv_r + offs_p * smv_p).to(tl.float32)
304+ psiv = v_tile * mv_i[None, :]
305+ qk = tl.dot(q_ro, tl.trans(k_ri))
306+ o_intra += tl.dot(qk * coeff, psiv)
307+ tl.store(
308+ INTRA + i_b * sin_b + i_h * sin_h + s_c[:, None] * sin_s + r_out * sin_r + offs_p[None, :] * sin_p,
309+ o_intra,
310+ mask=mask_c[:, None],
311+ )
312+ for r_in in tl.static_range(R):
313+ p_k = tl.make_block_ptr(
314+ Kr + i_b * sqr_b + i_h * sqr_h + r_in * sqr_r, (S, N), (sqr_s, sqr_n), (s0, 0), (CHUNK, N), (1, 0)
315+ )
316+ k_ri = tl.load(p_k, boundary_check=(0,), padding_option="zero").to(tl.float32)
317+ mv_i = tl.load(MIMO_V + i_h * smv_h + r_in * smv_r + offs_p * smv_p).to(tl.float32)
318+ kv += tl.dot(tl.trans(k_ri * ts_rev), v_tile * mv_i[None, :])
319+ p_kv = tl.make_block_ptr(
320+ KV + i_b * skv_b + i_h * skv_h + c * skv_c, (N, P), (skv_n, skv_p), (0, 0), (N, P), (1, 0)
321+ )
322+ tl.store(p_kv, kv)
323+ 
324+ 
325+@triton.jit
326+def _scan_inter_kernel(
327+ Qr,
328+ V,
329+ MIMO_V,
330+ DOUT,
331+ D,
332+ INTRA,
333+ KV,
334+ EXP_DA_CS,
335+ DECAY,
336+ STATES,
337+ DMIMO_O,
338+ sqr_b,
339+ sqr_h,
340+ sqr_s,
341+ sqr_r,
342+ sqr_n,
343+ sv_b,
344+ sv_s,
345+ sv_h,
346+ sv_p,
347+ smv_h,
348+ smv_r,
349+ smv_p,
350+ sd_h,
351+ sdo_b,
352+ sdo_s,
353+ sdo_h,
354+ sdo_p,
355+ sdisc_b,
356+ sdisc_h,
357+ sdisc_s,
358+ sin_b,
359+ sin_h,
360+ sin_s,
361+ sin_r,
362+ sin_p,
363+ skv_b,
364+ skv_h,
365+ skv_c,
366+ skv_n,
367+ skv_p,
368+ sde_b,
369+ sde_h,
370+ sde_c,
371+ sst_b,
372+ sst_h,
373+ sst_c,
374+ sst_n,
375+ sst_p,
376+ sdmo_b,
377+ sdmo_h,
378+ sdmo_r,
379+ sdmo_p,
380+ B,
381+ S,
382+ H,
383+ num_chunks,
384+ R: tl.constexpr,
385+ N: tl.constexpr,
386+ P: tl.constexpr,
387+ CHUNK: tl.constexpr,
388+ HAS_D: tl.constexpr,
389+):
390+ pid = tl.program_id(0)
391+ npg = tl.num_programs(0)
392+ total = B * H
393+ offs_c = tl.arange(0, CHUNK)
394+ offs_p = tl.arange(0, P)
395+ offs_r = tl.arange(0, R)
396+ for wi in range(pid, total, npg):
397+ i_b = wi // H
398+ i_h = wi % H
399+ disc_base = i_b * sdisc_b + i_h * sdisc_h
400+ d_val = tl.load(D + i_h * sd_h).to(tl.float32) if HAS_D else 0.0
401+ states = tl.zeros((N, P), dtype=tl.float32)
402+ dphi = tl.zeros((R, P), dtype=tl.float32)
403+ for c in range(num_chunks):
404+ s0 = c * CHUNK
405+ s_c = s0 + offs_c
406+ mask_c = s_c < S
407+ p_st = tl.make_block_ptr(
408+ STATES + i_b * sst_b + i_h * sst_h + c * sst_c, (N, P), (sst_n, sst_p), (0, 0), (N, P), (1, 0)
409+ )
410+ tl.store(p_st, states.to(p_st.dtype.element_ty))
411+ exp_da_cs = tl.load(EXP_DA_CS + disc_base + s_c * sdisc_s)
412+ decay = tl.load(DECAY + i_b * sde_b + i_h * sde_h + c * sde_c)
413+ v_tile = tl.load(
414+ tl.make_block_ptr(V + i_b * sv_b + i_h * sv_h, (S, P), (sv_s, sv_p), (s0, 0), (CHUNK, P), (1, 0)),
415+ boundary_check=(0,),
416+ padding_option="zero",
417+ ).to(tl.float32)
418+ dout = tl.load(
419+ DOUT + i_b * sdo_b + i_h * sdo_h + s_c[:, None] * sdo_s + offs_p[None, :] * sdo_p,
420+ mask=mask_c[:, None],
421+ other=0.0,
422+ ).to(tl.float32)
423+ for r_out in tl.static_range(R):
424+ p_q = tl.make_block_ptr(
425+ Qr + i_b * sqr_b + i_h * sqr_h + r_out * sqr_r, (S, N), (sqr_s, sqr_n), (s0, 0), (CHUNK, N), (1, 0)
426+ )
427+ q_ro = tl.load(p_q, boundary_check=(0,), padding_option="zero").to(tl.float32)
428+ intra = tl.load(
429+ INTRA + i_b * sin_b + i_h * sin_h + s_c[:, None] * sin_s + r_out * sin_r + offs_p[None, :] * sin_p,
430+ mask=mask_c[:, None],
431+ other=0.0,
432+ )
433+ o_r = tl.dot(q_ro, states) * exp_da_cs[:, None] + intra
434+ if HAS_D:
435+ mv_o = tl.load(MIMO_V + i_h * smv_h + r_out * smv_r + offs_p * smv_p).to(tl.float32)
436+ o_r += d_val * (v_tile * mv_o[None, :])
437+ sel = (offs_r == r_out).to(tl.float32)
438+ dphi += sel[:, None] * tl.sum(o_r * dout, axis=0)[None, :]
439+ kv = tl.load(
440+ tl.make_block_ptr(
441+ KV + i_b * skv_b + i_h * (skv_c * num_chunks) + c * skv_c,
442+ (N, P),
443+ (skv_n, skv_p),
444+ (0, 0),
445+ (N, P),
446+ (1, 0),
447+ )
448+ )
449+ states = states * decay + kv
450+ p_dmo = tl.make_block_ptr(
451+ DMIMO_O + i_b * sdmo_b + i_h * sdmo_h, (R, P), (sdmo_r, sdmo_p), (0, 0), (R, P), (1, 0)
452+ )
453+ tl.store(p_dmo, dphi.to(p_dmo.dtype.element_ty))
454+ 
455+ 
456+@triton.jit
457+def _pre_aux_kernel(
458+ Q,
459+ K,
460+ Q_BIAS,
461+ K_BIAS,
462+ ANGLES,
463+ Qr,
464+ Kr,
465+ QK_DOT,
466+ DT,
467+ TRAP,
468+ DA_CS,
469+ COEFF,
470+ TS_REV,
471+ EXP_DA_CS,
472+ DECAY,
473+ V,
474+ MIMO_V,
475+ INTRA,
476+ KV,
477+ sq_b,
478+ sq_s,
479+ sq_r,
480+ sq_g,
481+ sq_n,
482+ sqb_h,
483+ sqb_r,
484+ sqb_n,
485+ sa_b,
486+ sa_s,
487+ sa_h,
488+ sa_n,
489+ sqr_b,
490+ sqr_h,
491+ sqr_s,
492+ sqr_r,
493+ sqr_n,
494+ sqk_b,
495+ sqk_h,
496+ sqk_s,
497+ sqk_ro,
498+ sqk_ri,
499+ sdisc_b,
500+ sdisc_h,
501+ sdisc_s,
502+ sco_b,
503+ sco_h,
504+ sco_c,
505+ sco_i,
506+ sco_j,
507+ sde_b,
508+ sde_h,
509+ sde_c,
510+ sv_b,
511+ sv_s,
512+ sv_h,
513+ sv_p,
514+ smv_h,
515+ smv_r,
516+ smv_p,
517+ sin_b,
518+ sin_h,
519+ sin_s,
520+ sin_r,
521+ sin_p,
522+ skv_b,
523+ skv_h,
524+ skv_c,
525+ skv_n,
526+ skv_p,
527+ B,
528+ S,
529+ H,
530+ G,
531+ num_chunks,
532+ R: tl.constexpr,
533+ N: tl.constexpr,
534+ P: tl.constexpr,
535+ RD: tl.constexpr,
536+ CHUNK: tl.constexpr,
537+ HALF: tl.constexpr,
538+ QK_USE_DOT: tl.constexpr = True,
539+):
540+ """Merged prepass+aux (1 launch): rotary+qkdot+disc AND the chunk-local
541+ INTRA/KV, all on the B*H*nchunks grid. Saves a launch (golden is host-bound)
542+ and consumes Qr/Kr in-kernel.
543+ """
544+ offs_p = tl.arange(0, P)
545+ pid = tl.program_id(0)
546+ npg = tl.num_programs(0)
547+ total = B * H * num_chunks
548+ rep = H // G
549+ offs_n = tl.arange(0, N)
550+ offs_c = tl.arange(0, CHUNK)
551+ causal = offs_c[:, None] > offs_c[None, :]
552+ diagm = offs_c[:, None] == offs_c[None, :]
553+ 
554+ for wi in range(pid, total, npg):
555+ c = wi % num_chunks
556+ tmp = wi // num_chunks
557+ i_h = tmp % H
558+ i_b = tmp // H
559+ i_hqk = i_h // rep
560+ s0 = c * CHUNK
561+ s_c = s0 + offs_c
562+ mask_c = s_c < S
563+ 
564+ p_ang = tl.make_block_ptr(ANGLES + i_b * sa_b + i_h * sa_h, (S, RD), (sa_s, sa_n), (s0, 0), (CHUNK, RD), (1, 0))
565+ ang = tl.load(p_ang, boundary_check=(0,), padding_option="zero").to(tl.float32)
566+ cos = tl.cos(ang)
567+ sin = tl.sin(ang)
568+ 
569+ for r_out in tl.static_range(R):
570+ q_in = Q + i_b * sq_b + i_hqk * sq_g + r_out * sq_r
571+ k_in = K + i_b * sq_b + i_hqk * sq_g + r_out * sq_r
572+ q_out = Qr + i_b * sqr_b + i_h * sqr_h + r_out * sqr_r
573+ k_out = Kr + i_b * sqr_b + i_h * sqr_h + r_out * sqr_r
574+ qb = tl.load(Q_BIAS + i_h * sqb_h + r_out * sqb_r + offs_n * sqb_n).to(tl.float32)
575+ kb = tl.load(K_BIAS + i_h * sqb_h + r_out * sqb_r + offs_n * sqb_n).to(tl.float32)
576+ p_q = tl.make_block_ptr(q_in, (S, N), (sq_s, sq_n), (s0, 0), (CHUNK, N), (1, 0))
577+ p_k = tl.make_block_ptr(k_in, (S, N), (sq_s, sq_n), (s0, 0), (CHUNK, N), (1, 0))
578+ qf = tl.load(p_q, boundary_check=(0,), padding_option="zero").to(tl.float32) + qb[None, :]
579+ kf = tl.load(p_k, boundary_check=(0,), padding_option="zero").to(tl.float32) + kb[None, :]
580+ # RoPE rotate-half (cols [0,RD) with [HALF,HALF+RD))
581+ offs_d = tl.arange(0, RD)
582+ qba = tl.load(Q_BIAS + i_h * sqb_h + r_out * sqb_r + offs_d * sqb_n).to(tl.float32)
583+ qbb = tl.load(Q_BIAS + i_h * sqb_h + r_out * sqb_r + (HALF + offs_d) * sqb_n).to(tl.float32)
584+ kba = tl.load(K_BIAS + i_h * sqb_h + r_out * sqb_r + offs_d * sqb_n).to(tl.float32)
585+ kbb = tl.load(K_BIAS + i_h * sqb_h + r_out * sqb_r + (HALF + offs_d) * sqb_n).to(tl.float32)
586+ p_qa = tl.make_block_ptr(q_in, (S, N), (sq_s, sq_n), (s0, 0), (CHUNK, RD), (1, 0))
587+ p_qb2 = tl.make_block_ptr(q_in, (S, N), (sq_s, sq_n), (s0, HALF), (CHUNK, RD), (1, 0))
588+ p_ka = tl.make_block_ptr(k_in, (S, N), (sq_s, sq_n), (s0, 0), (CHUNK, RD), (1, 0))
589+ p_kb2 = tl.make_block_ptr(k_in, (S, N), (sq_s, sq_n), (s0, HALF), (CHUNK, RD), (1, 0))
590+ qa = tl.load(p_qa, boundary_check=(0,), padding_option="zero").to(tl.float32) + qba[None, :]
591+ qb2 = tl.load(p_qb2, boundary_check=(0,), padding_option="zero").to(tl.float32) + qbb[None, :]
592+ ka = tl.load(p_ka, boundary_check=(0,), padding_option="zero").to(tl.float32) + kba[None, :]
593+ kb2 = tl.load(p_kb2, boundary_check=(0,), padding_option="zero").to(tl.float32) + kbb[None, :]
594+ qar = cos * qa - sin * qb2
595+ qbr = sin * qa + cos * qb2
596+ kar = cos * ka - sin * kb2
597+ kbr = sin * ka + cos * kb2
598+ p_qr = tl.make_block_ptr(q_out, (S, N), (sqr_s, sqr_n), (s0, 0), (CHUNK, N), (1, 0))
599+ p_kr = tl.make_block_ptr(k_out, (S, N), (sqr_s, sqr_n), (s0, 0), (CHUNK, N), (1, 0))
600+ tl.store(p_qr, qf.to(p_qr.dtype.element_ty), boundary_check=(0,))
601+ tl.store(p_kr, kf.to(p_kr.dtype.element_ty), boundary_check=(0,))
602+ p_qr_a = tl.make_block_ptr(q_out, (S, N), (sqr_s, sqr_n), (s0, 0), (CHUNK, RD), (1, 0))
603+ p_qr_b = tl.make_block_ptr(q_out, (S, N), (sqr_s, sqr_n), (s0, HALF), (CHUNK, RD), (1, 0))
604+ p_kr_a = tl.make_block_ptr(k_out, (S, N), (sqr_s, sqr_n), (s0, 0), (CHUNK, RD), (1, 0))
605+ p_kr_b = tl.make_block_ptr(k_out, (S, N), (sqr_s, sqr_n), (s0, HALF), (CHUNK, RD), (1, 0))
606+ tl.store(p_qr_a, qar.to(p_qr_a.dtype.element_ty), boundary_check=(0,))
607+ tl.store(p_qr_b, qbr.to(p_qr_b.dtype.element_ty), boundary_check=(0,))
608+ tl.store(p_kr_a, kar.to(p_kr_a.dtype.element_ty), boundary_check=(0,))
609+ tl.store(p_kr_b, kbr.to(p_kr_b.dtype.element_ty), boundary_check=(0,))
610+ # qk_dot diag (rotation invariant -> use bias-added qf/kfd). Static
611+ # variant: reuse kf when r_in==r_out (python shortcut, lowers UB so the
612+ # R=4 static unroll fits, e.g. big N128_P128).
613+ for r_in in tl.static_range(R):
614+ if r_in == r_out:
615+ kfd = kf
616+ else:
617+ kbd = tl.load(K_BIAS + i_h * sqb_h + r_in * sqb_r + offs_n * sqb_n).to(tl.float32)
618+ p_kd = tl.make_block_ptr(
619+ K + i_b * sq_b + i_hqk * sq_g + r_in * sq_r, (S, N), (sq_s, sq_n), (s0, 0), (CHUNK, N), (1, 0)
620+ )
621+ kfd = tl.load(p_kd, boundary_check=(0,), padding_option="zero").to(tl.float32) + kbd[None, :]
622+ qkmat = tl.dot(qf, tl.trans(kfd))
623+ qkdiag = tl.sum(tl.where(diagm, qkmat, 0.0), axis=1)
624+ tl.store(
625+ QK_DOT + i_b * sqk_b + i_h * sqk_h + s_c * sqk_s + r_out * sqk_ro + r_in * sqk_ri,
626+ qkdiag.to(QK_DOT.dtype.element_ty),
627+ mask=mask_c,
628+ )
629+ 
630+ # --- discretization (hoisted off the serial scan) ---
631+ base = i_b * sdisc_b + i_h * sdisc_h
632+ dt = tl.load(DT + base + s_c * sdisc_s).to(tl.float32)
633+ trap = tl.load(TRAP + base + s_c * sdisc_s).to(tl.float32)
634+ gamma = dt * tl.sigmoid(trap)
635+ sh = s_c + 1
636+ shm = sh < S
637+ dt_sh = tl.load(DT + base + sh * sdisc_s, mask=shm, other=0.0).to(tl.float32)
638+ trap_sh = tl.load(TRAP + base + sh * sdisc_s, mask=shm, other=0.0).to(tl.float32)
639+ guard = (s_c < (S - 1)).to(tl.float32)
640+ trap_scale = gamma + dt_sh * tl.sigmoid(-trap_sh) * guard
641+ dacs_v = tl.load(DA_CS + base + s_c * sdisc_s).to(tl.float32)
642+ exp_da_cs = tl.exp(dacs_v)
643+ da_cs_sum = tl.load(DA_CS + base + (s0 + CHUNK - 1) * sdisc_s).to(tl.float32)
644+ exp_rev = tl.exp(da_cs_sum - dacs_v)
645+ seg = dacs_v[:, None] - dacs_v[None, :]
646+ coeff = (
647+ causal.to(tl.float32) * trap_scale[None, :] * tl.exp(tl.where(causal, seg, 0.0))
648+ + diagm.to(tl.float32) * gamma[:, None]
649+ )
650+ ts_rev = trap_scale * exp_rev
651+ p_co = tl.make_block_ptr(
652+ COEFF + i_b * sco_b + i_h * sco_h + c * sco_c,
653+ (CHUNK, CHUNK),
654+ (sco_i, sco_j),
655+ (0, 0),
656+ (CHUNK, CHUNK),
657+ (1, 0),
658+ )
659+ tl.store(p_co, coeff)
660+ tl.store(TS_REV + base + s_c * sdisc_s, ts_rev)
661+ tl.store(EXP_DA_CS + base + s_c * sdisc_s, exp_da_cs)
662+ tl.store(DECAY + i_b * sde_b + i_h * sde_h + c * sde_c, tl.exp(da_cs_sum))
663+ disc_base = base
664+ coeff = tl.load(
665+ tl.make_block_ptr(
666+ COEFF + i_b * sco_b + i_h * sco_h + c * sco_c,
667+ (CHUNK, CHUNK),
668+ (sco_i, sco_j),
669+ (0, 0),
670+ (CHUNK, CHUNK),
671+ (1, 0),
672+ )
673+ )
674+ ts_rev = tl.load(TS_REV + disc_base + s_c * sdisc_s)[:, None]
675+ p_v = tl.make_block_ptr(V + i_b * sv_b + i_h * sv_h, (S, P), (sv_s, sv_p), (s0, 0), (CHUNK, P), (1, 0))
676+ v_tile = tl.load(p_v, boundary_check=(0,), padding_option="zero").to(tl.float32)
677+ kv = tl.zeros((N, P), dtype=tl.float32)
678+ for r_out in tl.static_range(R):
679+ p_q = tl.make_block_ptr(
680+ Qr + i_b * sqr_b + i_h * sqr_h + r_out * sqr_r, (S, N), (sqr_s, sqr_n), (s0, 0), (CHUNK, N), (1, 0)
681+ )
682+ q_ro = tl.load(p_q, boundary_check=(0,), padding_option="zero")
683+ o_intra = tl.zeros((CHUNK, P), dtype=tl.float32)
684+ for r_in in tl.static_range(R):
685+ p_k = tl.make_block_ptr(
686+ Kr + i_b * sqr_b + i_h * sqr_h + r_in * sqr_r, (S, N), (sqr_s, sqr_n), (s0, 0), (CHUNK, N), (1, 0)
687+ )
688+ k_ri = tl.load(p_k, boundary_check=(0,), padding_option="zero")
689+ mv_i = tl.load(MIMO_V + i_h * smv_h + r_in * smv_r + offs_p * smv_p).to(tl.float32)
690+ psiv = v_tile * mv_i[None, :]
691+ qk = tl.dot(q_ro, tl.trans(k_ri))
692+ o_intra += tl.dot(qk * coeff, psiv)
693+ tl.store(
694+ INTRA + i_b * sin_b + i_h * sin_h + s_c[:, None] * sin_s + r_out * sin_r + offs_p[None, :] * sin_p,
695+ o_intra,
696+ mask=mask_c[:, None],
697+ )
698+ for r_in in tl.static_range(R):
699+ p_k = tl.make_block_ptr(
700+ Kr + i_b * sqr_b + i_h * sqr_h + r_in * sqr_r, (S, N), (sqr_s, sqr_n), (s0, 0), (CHUNK, N), (1, 0)
701+ )
702+ k_ri = tl.load(p_k, boundary_check=(0,), padding_option="zero").to(tl.float32)
703+ mv_i = tl.load(MIMO_V + i_h * smv_h + r_in * smv_r + offs_p * smv_p).to(tl.float32)
704+ kv += tl.dot(tl.trans(k_ri * ts_rev), v_tile * mv_i[None, :])
705+ p_kv = tl.make_block_ptr(
706+ KV + i_b * skv_b + i_h * skv_h + c * skv_c, (N, P), (skv_n, skv_p), (0, 0), (N, P), (1, 0)
707+ )
708+ tl.store(p_kv, kv)
709+ 
710+ 
711+@triton.jit
712+def _pre_aux_kernel_rl(
713+ Q,
714+ K,
715+ Q_BIAS,
716+ K_BIAS,
717+ ANGLES,
718+ Qr,
719+ Kr,
720+ QK_DOT,
721+ DT,
722+ TRAP,
723+ DA_CS,
724+ COEFF,
725+ TS_REV,
726+ EXP_DA_CS,
727+ DECAY,
728+ V,
729+ MIMO_V,
730+ INTRA,
731+ KV,
732+ sq_b,
733+ sq_s,
734+ sq_r,
735+ sq_g,
736+ sq_n,
737+ sqb_h,
738+ sqb_r,
739+ sqb_n,
740+ sa_b,
741+ sa_s,
742+ sa_h,
743+ sa_n,
744+ sqr_b,
745+ sqr_h,
746+ sqr_s,
747+ sqr_r,
748+ sqr_n,
749+ sqk_b,
750+ sqk_h,
751+ sqk_s,
752+ sqk_ro,
753+ sqk_ri,
754+ sdisc_b,
755+ sdisc_h,
756+ sdisc_s,
757+ sco_b,
758+ sco_h,
759+ sco_c,
760+ sco_i,
761+ sco_j,
762+ sde_b,
763+ sde_h,
764+ sde_c,
765+ sv_b,
766+ sv_s,
767+ sv_h,
768+ sv_p,
769+ smv_h,
770+ smv_r,
771+ smv_p,
772+ sin_b,
773+ sin_h,
774+ sin_s,
775+ sin_r,
776+ sin_p,
777+ skv_b,
778+ skv_h,
779+ skv_c,
780+ skv_n,
781+ skv_p,
782+ B,
783+ S,
784+ H,
785+ G,
786+ num_chunks,
787+ R: tl.constexpr,
788+ N: tl.constexpr,
789+ P: tl.constexpr,
790+ RD: tl.constexpr,
791+ CHUNK: tl.constexpr,
792+ HALF: tl.constexpr,
793+ QK_USE_DOT: tl.constexpr = True,
794+):
795+ """Merged prepass+aux (1 launch): rotary+qkdot+disc AND the chunk-local
796+ INTRA/KV, all on the B*H*nchunks grid. Saves a launch (golden is host-bound)
797+ and consumes Qr/Kr in-kernel.
798+ """
799+ offs_p = tl.arange(0, P)
800+ pid = tl.program_id(0)
801+ npg = tl.num_programs(0)
802+ total = B * H * num_chunks
803+ rep = H // G
804+ offs_n = tl.arange(0, N)
805+ offs_c = tl.arange(0, CHUNK)
806+ causal = offs_c[:, None] > offs_c[None, :]
807+ diagm = offs_c[:, None] == offs_c[None, :]
808+ 
809+ for wi in range(pid, total, npg):
810+ c = wi % num_chunks
811+ tmp = wi // num_chunks
812+ i_h = tmp % H
813+ i_b = tmp // H
814+ i_hqk = i_h // rep
815+ s0 = c * CHUNK
816+ s_c = s0 + offs_c
817+ mask_c = s_c < S
818+ 
819+ p_ang = tl.make_block_ptr(ANGLES + i_b * sa_b + i_h * sa_h, (S, RD), (sa_s, sa_n), (s0, 0), (CHUNK, RD), (1, 0))
820+ ang = tl.load(p_ang, boundary_check=(0,), padding_option="zero").to(tl.float32)
821+ cos = tl.cos(ang)
822+ sin = tl.sin(ang)
823+ 
824+ for r_out in range(R):
825+ q_in = Q + i_b * sq_b + i_hqk * sq_g + r_out * sq_r
826+ k_in = K + i_b * sq_b + i_hqk * sq_g + r_out * sq_r
827+ q_out = Qr + i_b * sqr_b + i_h * sqr_h + r_out * sqr_r
828+ k_out = Kr + i_b * sqr_b + i_h * sqr_h + r_out * sqr_r
829+ qb = tl.load(Q_BIAS + i_h * sqb_h + r_out * sqb_r + offs_n * sqb_n).to(tl.float32)
830+ kb = tl.load(K_BIAS + i_h * sqb_h + r_out * sqb_r + offs_n * sqb_n).to(tl.float32)
831+ p_q = tl.make_block_ptr(q_in, (S, N), (sq_s, sq_n), (s0, 0), (CHUNK, N), (1, 0))
832+ p_k = tl.make_block_ptr(k_in, (S, N), (sq_s, sq_n), (s0, 0), (CHUNK, N), (1, 0))
833+ qf = tl.load(p_q, boundary_check=(0,), padding_option="zero").to(tl.float32) + qb[None, :]
834+ kf = tl.load(p_k, boundary_check=(0,), padding_option="zero").to(tl.float32) + kb[None, :]
835+ # RoPE rotate-half (cols [0,RD) with [HALF,HALF+RD))
836+ offs_d = tl.arange(0, RD)
837+ qba = tl.load(Q_BIAS + i_h * sqb_h + r_out * sqb_r + offs_d * sqb_n).to(tl.float32)
838+ qbb = tl.load(Q_BIAS + i_h * sqb_h + r_out * sqb_r + (HALF + offs_d) * sqb_n).to(tl.float32)
839+ kba = tl.load(K_BIAS + i_h * sqb_h + r_out * sqb_r + offs_d * sqb_n).to(tl.float32)
840+ kbb = tl.load(K_BIAS + i_h * sqb_h + r_out * sqb_r + (HALF + offs_d) * sqb_n).to(tl.float32)
841+ p_qa = tl.make_block_ptr(q_in, (S, N), (sq_s, sq_n), (s0, 0), (CHUNK, RD), (1, 0))
842+ p_qb2 = tl.make_block_ptr(q_in, (S, N), (sq_s, sq_n), (s0, HALF), (CHUNK, RD), (1, 0))
843+ p_ka = tl.make_block_ptr(k_in, (S, N), (sq_s, sq_n), (s0, 0), (CHUNK, RD), (1, 0))
844+ p_kb2 = tl.make_block_ptr(k_in, (S, N), (sq_s, sq_n), (s0, HALF), (CHUNK, RD), (1, 0))
845+ qa = tl.load(p_qa, boundary_check=(0,), padding_option="zero").to(tl.float32) + qba[None, :]
846+ qb2 = tl.load(p_qb2, boundary_check=(0,), padding_option="zero").to(tl.float32) + qbb[None, :]
847+ ka = tl.load(p_ka, boundary_check=(0,), padding_option="zero").to(tl.float32) + kba[None, :]
848+ kb2 = tl.load(p_kb2, boundary_check=(0,), padding_option="zero").to(tl.float32) + kbb[None, :]
849+ qar = cos * qa - sin * qb2
850+ qbr = sin * qa + cos * qb2
851+ kar = cos * ka - sin * kb2
852+ kbr = sin * ka + cos * kb2
853+ p_qr = tl.make_block_ptr(q_out, (S, N), (sqr_s, sqr_n), (s0, 0), (CHUNK, N), (1, 0))
854+ p_kr = tl.make_block_ptr(k_out, (S, N), (sqr_s, sqr_n), (s0, 0), (CHUNK, N), (1, 0))
855+ tl.store(p_qr, qf.to(p_qr.dtype.element_ty), boundary_check=(0,))
856+ tl.store(p_kr, kf.to(p_kr.dtype.element_ty), boundary_check=(0,))
857+ p_qr_a = tl.make_block_ptr(q_out, (S, N), (sqr_s, sqr_n), (s0, 0), (CHUNK, RD), (1, 0))
858+ p_qr_b = tl.make_block_ptr(q_out, (S, N), (sqr_s, sqr_n), (s0, HALF), (CHUNK, RD), (1, 0))
859+ p_kr_a = tl.make_block_ptr(k_out, (S, N), (sqr_s, sqr_n), (s0, 0), (CHUNK, RD), (1, 0))
860+ p_kr_b = tl.make_block_ptr(k_out, (S, N), (sqr_s, sqr_n), (s0, HALF), (CHUNK, RD), (1, 0))
861+ tl.store(p_qr_a, qar.to(p_qr_a.dtype.element_ty), boundary_check=(0,))
862+ tl.store(p_qr_b, qbr.to(p_qr_b.dtype.element_ty), boundary_check=(0,))
863+ tl.store(p_kr_a, kar.to(p_kr_a.dtype.element_ty), boundary_check=(0,))
864+ tl.store(p_kr_b, kbr.to(p_kr_b.dtype.element_ty), boundary_check=(0,))
865+ # qk_dot diag (rotation invariant -> use bias-added qf/kfd).
866+ # runtime range (de-unroll): reuse tile buffers to keep UB bounded at
867+ # R>=4 (static R*R unroll overflows UB at e.g. N128_P32 / N256); a
868+ # runtime index can't take the r_in==r_out python shortcut, so always
869+ # reload kfd (bit-identical: same fp32 value).
870+ for r_in in range(R):
871+ kbd = tl.load(K_BIAS + i_h * sqb_h + r_in * sqb_r + offs_n * sqb_n).to(tl.float32)
872+ p_kd = tl.make_block_ptr(
873+ K + i_b * sq_b + i_hqk * sq_g + r_in * sq_r, (S, N), (sq_s, sq_n), (s0, 0), (CHUNK, N), (1, 0)
874+ )
875+ kfd = tl.load(p_kd, boundary_check=(0,), padding_option="zero").to(tl.float32) + kbd[None, :]
876+ qkmat = tl.dot(qf, tl.trans(kfd))
877+ qkdiag = tl.sum(tl.where(diagm, qkmat, 0.0), axis=1)
878+ tl.store(
879+ QK_DOT + i_b * sqk_b + i_h * sqk_h + s_c * sqk_s + r_out * sqk_ro + r_in * sqk_ri,
880+ qkdiag.to(QK_DOT.dtype.element_ty),
881+ mask=mask_c,
882+ )
883+ 
884+ # --- discretization (hoisted off the serial scan) ---
885+ base = i_b * sdisc_b + i_h * sdisc_h
886+ dt = tl.load(DT + base + s_c * sdisc_s).to(tl.float32)
887+ trap = tl.load(TRAP + base + s_c * sdisc_s).to(tl.float32)
888+ gamma = dt * tl.sigmoid(trap)
889+ sh = s_c + 1
890+ shm = sh < S
891+ dt_sh = tl.load(DT + base + sh * sdisc_s, mask=shm, other=0.0).to(tl.float32)
892+ trap_sh = tl.load(TRAP + base + sh * sdisc_s, mask=shm, other=0.0).to(tl.float32)
893+ guard = (s_c < (S - 1)).to(tl.float32)
894+ trap_scale = gamma + dt_sh * tl.sigmoid(-trap_sh) * guard
895+ dacs_v = tl.load(DA_CS + base + s_c * sdisc_s).to(tl.float32)
896+ exp_da_cs = tl.exp(dacs_v)
897+ da_cs_sum = tl.load(DA_CS + base + (s0 + CHUNK - 1) * sdisc_s).to(tl.float32)
898+ exp_rev = tl.exp(da_cs_sum - dacs_v)
899+ seg = dacs_v[:, None] - dacs_v[None, :]
900+ coeff = (
901+ causal.to(tl.float32) * trap_scale[None, :] * tl.exp(tl.where(causal, seg, 0.0))
902+ + diagm.to(tl.float32) * gamma[:, None]
903+ )
904+ ts_rev = trap_scale * exp_rev
905+ p_co = tl.make_block_ptr(
906+ COEFF + i_b * sco_b + i_h * sco_h + c * sco_c,
907+ (CHUNK, CHUNK),
908+ (sco_i, sco_j),
909+ (0, 0),
910+ (CHUNK, CHUNK),
911+ (1, 0),
912+ )
913+ tl.store(p_co, coeff)
914+ tl.store(TS_REV + base + s_c * sdisc_s, ts_rev)
915+ tl.store(EXP_DA_CS + base + s_c * sdisc_s, exp_da_cs)
916+ tl.store(DECAY + i_b * sde_b + i_h * sde_h + c * sde_c, tl.exp(da_cs_sum))
917+ disc_base = base
918+ coeff = tl.load(
919+ tl.make_block_ptr(
920+ COEFF + i_b * sco_b + i_h * sco_h + c * sco_c,
921+ (CHUNK, CHUNK),
922+ (sco_i, sco_j),
923+ (0, 0),
924+ (CHUNK, CHUNK),
925+ (1, 0),
926+ )
927+ )
928+ ts_rev = tl.load(TS_REV + disc_base + s_c * sdisc_s)[:, None]
929+ p_v = tl.make_block_ptr(V + i_b * sv_b + i_h * sv_h, (S, P), (sv_s, sv_p), (s0, 0), (CHUNK, P), (1, 0))
930+ v_tile = tl.load(p_v, boundary_check=(0,), padding_option="zero").to(tl.float32)
931+ kv = tl.zeros((N, P), dtype=tl.float32)
932+ for r_out in range(R):
933+ p_q = tl.make_block_ptr(
934+ Qr + i_b * sqr_b + i_h * sqr_h + r_out * sqr_r, (S, N), (sqr_s, sqr_n), (s0, 0), (CHUNK, N), (1, 0)
935+ )
936+ q_ro = tl.load(p_q, boundary_check=(0,), padding_option="zero")
937+ o_intra = tl.zeros((CHUNK, P), dtype=tl.float32)
938+ for r_in in range(R):
939+ p_k = tl.make_block_ptr(
940+ Kr + i_b * sqr_b + i_h * sqr_h + r_in * sqr_r, (S, N), (sqr_s, sqr_n), (s0, 0), (CHUNK, N), (1, 0)
941+ )
942+ k_ri = tl.load(p_k, boundary_check=(0,), padding_option="zero")
943+ mv_i = tl.load(MIMO_V + i_h * smv_h + r_in * smv_r + offs_p * smv_p).to(tl.float32)
944+ psiv = v_tile * mv_i[None, :]
945+ qk = tl.dot(q_ro, tl.trans(k_ri))
946+ o_intra += tl.dot(qk * coeff, psiv)
947+ tl.store(
948+ INTRA + i_b * sin_b + i_h * sin_h + s_c[:, None] * sin_s + r_out * sin_r + offs_p[None, :] * sin_p,
949+ o_intra,
950+ mask=mask_c[:, None],
951+ )
952+ for r_in in range(R):
953+ p_k = tl.make_block_ptr(
954+ Kr + i_b * sqr_b + i_h * sqr_h + r_in * sqr_r, (S, N), (sqr_s, sqr_n), (s0, 0), (CHUNK, N), (1, 0)
955+ )
956+ k_ri = tl.load(p_k, boundary_check=(0,), padding_option="zero").to(tl.float32)
957+ mv_i = tl.load(MIMO_V + i_h * smv_h + r_in * smv_r + offs_p * smv_p).to(tl.float32)
958+ kv += tl.dot(tl.trans(k_ri * ts_rev), v_tile * mv_i[None, :])
959+ p_kv = tl.make_block_ptr(
960+ KV + i_b * skv_b + i_h * skv_h + c * skv_c, (N, P), (skv_n, skv_p), (0, 0), (N, P), (1, 0)
961+ )
962+ tl.store(p_kv, kv)
963+ 
964+ 
965+_PLAN = {}
966+ 
967+ 
968+def _plan(q, v, chunk_size, rdd, out_dtype):
969+ B, S, R, G, N = q.shape
970+ H, P = v.shape[-2], v.shape[-1]
971+ key = (B, S, H, G, N, P, R, chunk_size, rdd, str(out_dtype), q.device.index)
972+ pl = _PLAN.get(key)
973+ if pl:
974+ return pl
975+ nch = S // chunk_size
976+ dev = q.device
977+ f32 = torch.float32
978+ # fp32 RoPE'd scratch (bf16 measured net-negative for this aux structure).
979+ scratch = f32
980+ pl = dict(
981+ B=B,
982+ S=S,
983+ H=H,
984+ G=G,
985+ N=N,
986+ P=P,
987+ R=R,
988+ C=chunk_size,
989+ RD=N // rdd,
990+ HALF=N // 2,
991+ nch=nch,
992+ grid_pre=(min(get_vector_num(), B * H * nch),),
993+ grid_aux=(min(get_aicore_num(), B * H * nch),),
994+ grid_scan=(min(get_aicore_num(), B * H),),
995+ Qr=torch.empty((B, H, S, R, N), device=dev, dtype=scratch),
996+ Kr=torch.empty((B, H, S, R, N), device=dev, dtype=scratch),
997+ COEFF=torch.empty((B, H, nch, chunk_size, chunk_size), device=dev, dtype=f32),
998+ TS_REV=torch.empty((B, H, S), device=dev, dtype=f32),
999+ EXP_DA_CS=torch.empty((B, H, S), device=dev, dtype=f32),
1000+ DECAY=torch.empty((B, H, nch), device=dev, dtype=f32),
1001+ INTRA=torch.empty((B, H, S, R, P), device=dev, dtype=f32),
1002+ KV=torch.empty((B, H, nch, N, P), device=dev, dtype=f32),
1003+ dev=dev,
1004+ )
1005+ _PLAN[key] = pl
1006+ return pl
1007+ 
1008+ 
1009+@input_guard
1010+def mamba3_mimo_bwd_fwd_aux(
1011+ dout,
1012+ q,
1013+ k,
1014+ v,
1015+ q_bias,
1016+ k_bias,
1017+ mimo_v,
1018+ mimo_o,
1019+ angles,
1020+ da_cs,
1021+ dt,
1022+ trap,
1023+ d=None,
1024+ chunk_size: int = 16,
1025+ rotary_dim_divisor: int = 4,
1026+ output_dtype: torch.dtype = torch.float32,
1027+):
1028+ """Batch-parallel-aux golden-common-path #22. Returns ``(states, qk_dot, dmimo_o)``."""
1029+ pl = _plan(q, v, chunk_size, rotary_dim_divisor, output_dtype)
1030+ B, S, H, G, N, P, R, C = pl["B"], pl["S"], pl["H"], pl["G"], pl["N"], pl["P"], pl["R"], pl["C"]
1031+ nch, RD, HALF = pl["nch"], pl["RD"], pl["HALF"]
1032+ Qr, Kr = pl["Qr"], pl["Kr"]
1033+ COEFF, TS_REV, EXP_DA_CS, DECAY = pl["COEFF"], pl["TS_REV"], pl["EXP_DA_CS"], pl["DECAY"]
1034+ INTRA, KV = pl["INTRA"], pl["KV"]
1035+ dev = pl["dev"]
1036+ HAS_D = d is not None
1037+ _d = d if HAS_D else mimo_v.new_zeros(H)
1038+ 
1039+ states = torch.empty((B, H, nch, N, P), device=dev, dtype=output_dtype)
1040+ qk_dot = torch.empty((B, H, S, R, R), device=dev, dtype=output_dtype)
1041+ dmimo_o = torch.empty((B, H, R, P), device=dev, dtype=output_dtype)
1042+ 
1043+ # The static-unrolled merged pre_aux is fastest but the R*R static rank unroll
1044+ # (a) overflows UB when the [CHUNK,N]/[N,P] tiles are wide/tall (empirically
1045+ # R>=4 with N>128 or P<64 — e.g. N256, P32), and (b) hits the cbuf compile
1046+ # wall once R>=6 (R*R>=36 unrolled dots, e.g. N128_P64_R8 compiles >25min).
1047+ # Both cases use the de-unrolled (runtime rank-loop) variant, which reuses
1048+ # tile buffers to keep UB and IR bounded (bit-identical fp32 ops). This R>=6
1049+ # gate mirrors the staged host ``rank_loop = 1 if R >= 6`` de-unroll wall fix.
1050+ # All four perf shapes (R in {2,4}, N<=128, P>=64) stay on the static path.
1051+ _pre_kernel = _pre_aux_kernel_rl if (R >= 6 or (R >= 4 and (N > 128 or P < 64))) else _pre_aux_kernel
1052+ _pre_kernel[pl["grid_pre"]](
1053+ q,
1054+ k,
1055+ q_bias,
1056+ k_bias,
1057+ angles,
1058+ Qr,
1059+ Kr,
1060+ qk_dot,
1061+ dt,
1062+ trap,
1063+ da_cs,
1064+ COEFF,
1065+ TS_REV,
1066+ EXP_DA_CS,
1067+ DECAY,
1068+ v,
1069+ mimo_v,
1070+ INTRA,
1071+ KV,
1072+ q.stride(0),
1073+ q.stride(1),
1074+ q.stride(2),
1075+ q.stride(3),
1076+ q.stride(4),
1077+ q_bias.stride(0),
1078+ q_bias.stride(1),
1079+ q_bias.stride(2),
1080+ angles.stride(0),
1081+ angles.stride(1),
1082+ angles.stride(2),
1083+ angles.stride(3),
1084+ Qr.stride(0),
1085+ Qr.stride(1),
1086+ Qr.stride(2),
1087+ Qr.stride(3),
1088+ Qr.stride(4),
1089+ qk_dot.stride(0),
1090+ qk_dot.stride(1),
1091+ qk_dot.stride(2),
1092+ qk_dot.stride(3),
1093+ qk_dot.stride(4),
1094+ dt.stride(0),
1095+ dt.stride(1),
1096+ dt.stride(2),
1097+ COEFF.stride(0),
1098+ COEFF.stride(1),
1099+ COEFF.stride(2),
1100+ COEFF.stride(3),
1101+ COEFF.stride(4),
1102+ DECAY.stride(0),
1103+ DECAY.stride(1),
1104+ DECAY.stride(2),
1105+ v.stride(0),
1106+ v.stride(1),
1107+ v.stride(2),
1108+ v.stride(3),
1109+ mimo_v.stride(0),
1110+ mimo_v.stride(1),
1111+ mimo_v.stride(2),
1112+ INTRA.stride(0),
1113+ INTRA.stride(1),
1114+ INTRA.stride(2),
1115+ INTRA.stride(3),
1116+ INTRA.stride(4),
1117+ KV.stride(0),
1118+ KV.stride(1),
1119+ KV.stride(2),
1120+ KV.stride(3),
1121+ KV.stride(4),
1122+ B,
1123+ S,
1124+ H,
1125+ G,
1126+ nch,
1127+ R=R,
1128+ N=N,
1129+ P=P,
1130+ RD=RD,
1131+ CHUNK=C,
1132+ HALF=HALF,
1133+ QK_USE_DOT=(q.dtype == torch.float32),
1134+ multibuffer=False,
1135+ enable_auto_bind_sub_block=True,
1136+ )
1137+ _scan_inter_kernel[pl["grid_scan"]](
1138+ Qr,
1139+ v,
1140+ mimo_v,
1141+ dout,
1142+ _d,
1143+ INTRA,
1144+ KV,
1145+ EXP_DA_CS,
1146+ DECAY,
1147+ states,
1148+ dmimo_o,
1149+ Qr.stride(0),
1150+ Qr.stride(1),
1151+ Qr.stride(2),
1152+ Qr.stride(3),
1153+ Qr.stride(4),
1154+ v.stride(0),
1155+ v.stride(1),
1156+ v.stride(2),
1157+ v.stride(3),
1158+ mimo_v.stride(0),
1159+ mimo_v.stride(1),
1160+ mimo_v.stride(2),
1161+ _d.stride(0),
1162+ dout.stride(0),
1163+ dout.stride(1),
1164+ dout.stride(2),
1165+ dout.stride(3),
1166+ dt.stride(0),
1167+ dt.stride(1),
1168+ dt.stride(2),
1169+ INTRA.stride(0),
1170+ INTRA.stride(1),
1171+ INTRA.stride(2),
1172+ INTRA.stride(3),
1173+ INTRA.stride(4),
1174+ KV.stride(0),
1175+ KV.stride(1),
1176+ KV.stride(2),
1177+ KV.stride(3),
1178+ KV.stride(4),
1179+ DECAY.stride(0),
1180+ DECAY.stride(1),
1181+ DECAY.stride(2),
1182+ states.stride(0),
1183+ states.stride(1),
1184+ states.stride(2),
1185+ states.stride(3),
1186+ states.stride(4),
1187+ dmimo_o.stride(0),
1188+ dmimo_o.stride(1),
1189+ dmimo_o.stride(2),
1190+ dmimo_o.stride(3),
1191+ B,
1192+ S,
1193+ H,
1194+ nch,
1195+ R=R,
1196+ N=N,
1197+ P=P,
1198+ CHUNK=C,
1199+ HAS_D=HAS_D,
1200+ multibuffer=False,
1201+ enable_auto_bind_sub_block=True,
1202+ )
1203+ return states, qk_dot, dmimo_o
@@ -0,0 +1,581 @@
1+# Copyright (c) 2026, Dao AI Lab, Goombalab
2+# Copyright (c) 2026, Huawei Technologies Co., Ltd.
3+# pylint: disable=duplicate-code
4+# Baseline and optimized kernels intentionally retain formula-level code for comparison.
5+#
6+# Licensed under the Apache License, Version 2.0 (the "License");
7+# you may not use this file except in compliance with the License.
8+#
9+# Triton-Ascend 迁移自 state-spaces/mamba:
10+# mamba_ssm/ops/tilelang/mamba3/mamba3_mimo_bwd.py (mamba_mimo_bwd_fwd_kernel, TileLang DSL)
11+# 对应 MindSpeed-Ops issue #22。
12+#
13+# 迁移范围: Mamba-3 MIMO 反向第一遍 (bwd_fwd) 的 *最小可测路径* —— batched(非 varlen)、
14+# S 可被 chunk_size 整除(tail_len==0)、reduceO(投影输出 [B,S,H,P])、无
15+# fuse_pregate_headwise_rms_norm。支持可选 Z 门控(SiLU)、可选 D-skip、GQA(G<=H)。
16+#
17+# bwd_fwd 的语义: 在反向时按 chunk 重算 MIMO 前向中间量 raw_y(逐 rank 的下投影前输出),
18+# 据此就地累加投影梯度并缓存供第二遍(bwd_bwd)使用。关键观察: Z/Zeta 门控只作用在每个时间
19+# 步的输出上, *不* 回灌进状态递推, 因此 DMIMO_O / DZ / DMIMO_Z 都是纯 "局部输出收缩",
20+# 可在一次前向重算中得到(与 torch autograd 完全等价)。
21+#
22+# 本遍产出:
23+# - STATES [B,H,nchunks,N,P]: 每个 chunk 进入时的递推状态(更新前), 供 bwd_bwd。
24+# - QK_DOT [B,H,S,R,R]: 逐时间步 q·k 的 R×R 对角块(旋转保内积, 同位置点积与旋转前相等)。
25+# - DMIMO_O [B,H,R,P]: Phi(下投影)梯度 dPhi = Σ_cs (gate·raw_y)·dout, host 端再对 batch 求和。
26+# - DZ [B,S,H,P]: 门控输入梯度(hasZ)。
27+# - DMIMO_Z [B,H,R,P]: Zeta(门控投影)梯度(hasZ), host 端再对 batch 求和。
28+#
29+# 复用 mamba3_mimo_fwd 的 rotary kernel(加 bias + 旋转编码, 写 Qr/Kr fp32 scratch)。
30+# scan kernel 在前向重算骨架上叠加上述缓存与梯度收缩。迁移要点详见 docs/triton/mamba3_mimo_bwd_fwd.md。
31+ 
32+import triton
33+import triton.language as tl
34+import torch
35+ 
36+from mindspeed_ops.api.triton.utils import get_vector_num, get_aicore_num, input_guard
37+from mindspeed_ops.arch32.triton.mamba3_mimo_fwd import _mamba3_mimo_rotary_kernel
38+ 
39+ 
40+@triton.jit
41+def _mamba3_mimo_qkdot_kernel(
42+ Q,
43+ K,
44+ Q_BIAS,
45+ K_BIAS,
46+ QK_DOT,
47+ sq_b,
48+ sq_s,
49+ sq_r,
50+ sq_g,
51+ sq_n,
52+ sqb_h,
53+ sqb_r,
54+ sqb_n,
55+ sqk_b,
56+ sqk_h,
57+ sqk_s,
58+ sqk_ro,
59+ sqk_ri,
60+ B,
61+ S,
62+ H,
63+ G,
64+ num_chunks,
65+ R: tl.constexpr,
66+ N: tl.constexpr,
67+ CHUNK: tl.constexpr,
68+ QK_USE_DOT: tl.constexpr = False,
69+):
70+ # 逐 (b,h,chunk) 在向量核上计算 QK_DOT 的 R×R 对角块, 用 *旋转前* 的 bias 化 q,k
71+ # (与源 TileLang 一致: 旋转为正交变换, 同位置点积旋转前后相等; 旋转前计算可避免与
72+ # torch 参考的旋转实现产生 fp32 舍入差, 让缓存对齐更紧)。
73+ pid = tl.program_id(0)
74+ npg = tl.num_programs(0)
75+ total = B * H * num_chunks
76+ rep = H // G
77+ offs_n = tl.arange(0, N)
78+ offs_c = tl.arange(0, CHUNK)
79+ # C1: 对角块 = 同位置 q·k 的逐行点积; 直接用纯向量 reduce tl.sum(qf*kf, axis=1) 取代
80+ # [CHUNK,CHUNK] 对角 matmul + gather, 省 16x cube 过算且 k 只读 R 次。旋转前 bias 化输入
81+ # (qf,kf 来自 Q/K+bias, 旋转前), 与 ref(旋转前 einsum 对角)累加对齐, MERE 不回退。
82+ 
83+ for wi in range(pid, total, npg):
84+ c = wi % num_chunks
85+ tmp = wi // num_chunks
86+ i_h = tmp % H
87+ i_b = tmp // H
88+ i_hqk = i_h // rep
89+ s0 = c * CHUNK
90+ s_c = s0 + offs_c
91+ mask_c = s_c < S
92+ 
93+ for r_out in tl.static_range(R):
94+ qb = tl.load(Q_BIAS + i_h * sqb_h + r_out * sqb_r + offs_n * sqb_n).to(tl.float32)
95+ p_q = tl.make_block_ptr(
96+ Q + i_b * sq_b + i_hqk * sq_g + r_out * sq_r, (S, N), (sq_s, sq_n), (s0, 0), (CHUNK, N), (1, 0)
97+ )
98+ qf = tl.load(p_q, boundary_check=(0,), padding_option="zero").to(tl.float32) + qb[None, :]
99+ for r_in in tl.static_range(R):
100+ kb = tl.load(K_BIAS + i_h * sqb_h + r_in * sqb_r + offs_n * sqb_n).to(tl.float32)
101+ p_k = tl.make_block_ptr(
102+ K + i_b * sq_b + i_hqk * sq_g + r_in * sq_r, (S, N), (sq_s, sq_n), (s0, 0), (CHUNK, N), (1, 0)
103+ )
104+ kf = tl.load(p_k, boundary_check=(0,), padding_option="zero").to(tl.float32) + kb[None, :]
105+ if QK_USE_DOT:
106+ # fp32 输入: 用 cube matmul 取对角, 累加序与 ref einsum 对齐 (MERE~0.7, 稳过 L0)
107+ qkmat = tl.dot(qf, tl.trans(kf)) # [CHUNK, CHUNK]
108+ qkdiag = tl.sum(tl.where(offs_c[:, None] == offs_c[None, :], qkmat, 0.0), axis=1)
109+ else:
110+ qkdiag = tl.sum(qf * kf, axis=1) # [CHUNK] 对角 = 逐行 q·k over N (低精度输入安全且快)
111+ tl.store(
112+ QK_DOT + i_b * sqk_b + i_h * sqk_h + s_c * sqk_s + r_out * sqk_ro + r_in * sqk_ri,
113+ qkdiag.to(QK_DOT.dtype.element_ty),
114+ mask=mask_c,
115+ )
116+ 
117+ 
118+@triton.jit
119+def _mamba3_mimo_bwd_fwd_scan_kernel(
120+ Qr,
121+ Kr,
122+ V,
123+ MIMO_V,
124+ MIMO_O,
125+ MIMO_Z,
126+ D,
127+ Z,
128+ DOUT,
129+ DT,
130+ TRAP,
131+ DA_CS,
132+ DA_CS_REV,
133+ SEGSUM,
134+ STATES,
135+ DMIMO_O,
136+ DMIMO_Z,
137+ DZ,
138+ sqr_b,
139+ sqr_h,
140+ sqr_s,
141+ sqr_r,
142+ sqr_n,
143+ sv_b,
144+ sv_s,
145+ sv_h,
146+ sv_p,
147+ smv_h,
148+ smv_r,
149+ smv_p,
150+ sd_h,
151+ sz_b,
152+ sz_s,
153+ sz_h,
154+ sz_p,
155+ sdo_b,
156+ sdo_s,
157+ sdo_h,
158+ sdo_p,
159+ sdisc_b,
160+ sdisc_h,
161+ sdisc_s,
162+ sseg_b,
163+ sseg_h,
164+ sseg_c,
165+ sseg_i,
166+ sseg_j,
167+ sst_b,
168+ sst_h,
169+ sst_c,
170+ sst_n,
171+ sst_p,
172+ sdmo_b,
173+ sdmo_h,
174+ sdmo_r,
175+ sdmo_p,
176+ sdz_b,
177+ sdz_s,
178+ sdz_h,
179+ sdz_p,
180+ B,
181+ S,
182+ H,
183+ num_chunks,
184+ R: tl.constexpr,
185+ N: tl.constexpr,
186+ P: tl.constexpr,
187+ CHUNK: tl.constexpr,
188+ HAS_D: tl.constexpr,
189+ HAS_Z: tl.constexpr,
190+):
191+ pid = tl.program_id(0)
192+ npg = tl.num_programs(0)
193+ total = B * H
194+ offs_c = tl.arange(0, CHUNK)
195+ offs_p = tl.arange(0, P)
196+ offs_r = tl.arange(0, R)
197+ cs_i = offs_c[:, None]
198+ cs_j = offs_c[None, :]
199+ causal = cs_i > cs_j
200+ diagm = cs_i == cs_j
201+ 
202+ for wi in range(pid, total, npg):
203+ i_b = wi // H
204+ i_h = wi % H
205+ disc_base = i_b * sdisc_b + i_h * sdisc_h
206+ 
207+ if HAS_D:
208+ d_val = tl.load(D + i_h * sd_h).to(tl.float32)
209+ else:
210+ d_val = 0.0
211+ 
212+ states = tl.zeros((N, P), dtype=tl.float32)
213+ dphi = tl.zeros((R, P), dtype=tl.float32)
214+ dzeta = tl.zeros((R, P), dtype=tl.float32)
215+ 
216+ for c in range(num_chunks):
217+ s0 = c * CHUNK
218+ s_c = s0 + offs_c
219+ mask_c = s_c < S
220+ 
221+ # --- cache entering recurrent state (before this chunk's update) ---
222+ p_st = tl.make_block_ptr(
223+ STATES + i_b * sst_b + i_h * sst_h + c * sst_c, (N, P), (sst_n, sst_p), (0, 0), (N, P), (1, 0)
224+ )
225+ tl.store(p_st, states.to(p_st.dtype.element_ty))
226+ 
227+ # --- discretization scalars at CHUNK resolution (identical to fwd) ---
228+ dt = tl.load(DT + disc_base + s_c * sdisc_s).to(tl.float32)
229+ trap = tl.load(TRAP + disc_base + s_c * sdisc_s).to(tl.float32)
230+ gamma = dt * tl.sigmoid(trap)
231+ sh = s_c + 1
232+ shm = sh < S
233+ dt_sh = tl.load(DT + disc_base + sh * sdisc_s, mask=shm, other=0.0).to(tl.float32)
234+ trap_sh = tl.load(TRAP + disc_base + sh * sdisc_s, mask=shm, other=0.0).to(tl.float32)
235+ guard = (s_c < (S - 1)).to(tl.float32)
236+ trap_scale = gamma + dt_sh * tl.sigmoid(-trap_sh) * guard
237+ exp_da_cs = tl.exp(tl.load(DA_CS + disc_base + s_c * sdisc_s).to(tl.float32))
238+ da_cs_rev = tl.load(DA_CS_REV + disc_base + s_c * sdisc_s).to(tl.float32)
239+ da_cs_sum = tl.load(DA_CS + disc_base + (s0 + CHUNK - 1) * sdisc_s).to(tl.float32)
240+ exp_rev = tl.exp(da_cs_rev)
241+ seg = tl.load(
242+ SEGSUM + i_b * sseg_b + i_h * sseg_h + c * sseg_c + offs_c[:, None] * sseg_i + offs_c[None, :] * sseg_j
243+ ).to(tl.float32)
244+ seg_use = tl.where(causal, seg, 0.0)
245+ coeff = (
246+ causal.to(tl.float32) * trap_scale[None, :] * tl.exp(seg_use) + diagm.to(tl.float32) * gamma[:, None]
247+ )
248+ 
249+ p_v = tl.make_block_ptr(V + i_b * sv_b + i_h * sv_h, (S, P), (sv_s, sv_p), (s0, 0), (CHUNK, P), (1, 0))
250+ v_tile = tl.load(p_v, boundary_check=(0,), padding_option="zero").to(tl.float32)
251+ ts_rev = (trap_scale * exp_rev)[:, None]
252+ 
253+ # upstream grad for this (b,h,chunk): DOUT[b, s, h, :] (reduceO layout [B,S,H,P])
254+ dout = tl.load(
255+ DOUT + i_b * sdo_b + i_h * sdo_h + s_c[:, None] * sdo_s + offs_p[None, :] * sdo_p,
256+ mask=mask_c[:, None],
257+ other=0.0,
258+ ).to(tl.float32)
259+ if HAS_Z:
260+ zt = tl.load(
261+ Z + i_b * sz_b + i_h * sz_h + s_c[:, None] * sz_s + offs_p[None, :] * sz_p,
262+ mask=mask_c[:, None],
263+ other=0.0,
264+ ).to(tl.float32)
265+ 
266+ dz_acc = tl.zeros((CHUNK, P), dtype=tl.float32)
267+ 
268+ # --- per rank: recompute raw_y, cache QK_DOT diag, accumulate projection grads ---
269+ for r_out in tl.static_range(R):
270+ p_q = tl.make_block_ptr(
271+ Qr + i_b * sqr_b + i_h * sqr_h + r_out * sqr_r, (S, N), (sqr_s, sqr_n), (s0, 0), (CHUNK, N), (1, 0)
272+ )
273+ q_ro = tl.load(p_q, boundary_check=(0,), padding_option="zero")
274+ o_r = tl.dot(q_ro, states) * exp_da_cs[:, None]
275+ for r_in in tl.static_range(R):
276+ p_k = tl.make_block_ptr(
277+ Kr + i_b * sqr_b + i_h * sqr_h + r_in * sqr_r,
278+ (S, N),
279+ (sqr_s, sqr_n),
280+ (s0, 0),
281+ (CHUNK, N),
282+ (1, 0),
283+ )
284+ k_ri = tl.load(p_k, boundary_check=(0,), padding_option="zero")
285+ mv_i = tl.load(MIMO_V + i_h * smv_h + r_in * smv_r + offs_p * smv_p).to(tl.float32)
286+ psiv_i = v_tile * mv_i[None, :]
287+ qk = tl.dot(q_ro, tl.trans(k_ri))
288+ o_r += tl.dot(qk * coeff, psiv_i)
289+ if HAS_D:
290+ mv_o = tl.load(MIMO_V + i_h * smv_h + r_out * smv_r + offs_p * smv_p).to(tl.float32)
291+ o_r += d_val * (v_tile * mv_o[None, :])
292+ 
293+ # raw_y for rank r_out == o_r
294+ phi = tl.load(MIMO_O + i_h * smv_h + r_out * smv_r + offs_p * smv_p).to(tl.float32)
295+ sel = (offs_r == r_out).to(tl.float32) # [R] one-hot, R leading (P stays inner/aligned)
296+ if HAS_Z:
297+ mz = tl.load(MIMO_Z + i_h * smv_h + r_out * smv_r + offs_p * smv_p).to(tl.float32)
298+ u = zt * mz[None, :]
299+ sig = tl.sigmoid(u)
300+ gate = u * sig
301+ dgate = sig * (1.0 + u * (1.0 - sig)) # C2: silu'(u), sigmoid(-u)=1-sig 省一次超越函数
302+ o_gated = gate * o_r
303+ contrib = tl.sum(o_gated * dout, axis=0) # [P]
304+ dphi += sel[:, None] * contrib[None, :]
305+ dldu = dout * phi[None, :] * o_r * dgate # dL/du (per rank)
306+ dz_acc += dldu * mz[None, :] # Σ_r dL/du · Zeta
307+ contrib_z = tl.sum(dldu * zt, axis=0) # [P] = Σ_cs dL/du · Z
308+ dzeta += sel[:, None] * contrib_z[None, :]
309+ else:
310+ contrib = tl.sum(o_r * dout, axis=0) # [P]
311+ dphi += sel[:, None] * contrib[None, :]
312+ 
313+ if HAS_Z:
314+ tl.store(
315+ DZ + i_b * sdz_b + i_h * sdz_h + s_c[:, None] * sdz_s + offs_p[None, :] * sdz_p,
316+ dz_acc.to(DZ.dtype.element_ty),
317+ mask=mask_c[:, None],
318+ )
319+ 
320+ # --- recurrent state update (identical to fwd) ---
321+ new_state = tl.zeros((N, P), dtype=tl.float32)
322+ for r_in in tl.static_range(R):
323+ p_k = tl.make_block_ptr(
324+ Kr + i_b * sqr_b + i_h * sqr_h + r_in * sqr_r, (S, N), (sqr_s, sqr_n), (s0, 0), (CHUNK, N), (1, 0)
325+ )
326+ k_ri = tl.load(p_k, boundary_check=(0,), padding_option="zero")
327+ mv_i = tl.load(MIMO_V + i_h * smv_h + r_in * smv_r + offs_p * smv_p).to(tl.float32)
328+ k_state = k_ri * ts_rev
329+ new_state += tl.dot(tl.trans(k_state), v_tile * mv_i[None, :])
330+ states = states * tl.exp(da_cs_sum) + new_state
331+ 
332+ # --- write per-(b,h) projection grads ---
333+ p_dmo = tl.make_block_ptr(
334+ DMIMO_O + i_b * sdmo_b + i_h * sdmo_h, (R, P), (sdmo_r, sdmo_p), (0, 0), (R, P), (1, 0)
335+ )
336+ tl.store(p_dmo, dphi.to(p_dmo.dtype.element_ty))
337+ if HAS_Z:
338+ p_dmz = tl.make_block_ptr(
339+ DMIMO_Z + i_b * sdmo_b + i_h * sdmo_h, (R, P), (sdmo_r, sdmo_p), (0, 0), (R, P), (1, 0)
340+ )
341+ tl.store(p_dmz, dzeta.to(p_dmz.dtype.element_ty))
342+ 
343+ 
344+@input_guard
345+def mamba3_mimo_bwd_fwd(
346+ dout: torch.Tensor,
347+ q: torch.Tensor,
348+ k: torch.Tensor,
349+ v: torch.Tensor,
350+ q_bias: torch.Tensor,
351+ k_bias: torch.Tensor,
352+ mimo_v: torch.Tensor,
353+ mimo_o: torch.Tensor,
354+ angles: torch.Tensor,
355+ da_cs: torch.Tensor,
356+ da_cs_rev: torch.Tensor,
357+ dt: torch.Tensor,
358+ trap: torch.Tensor,
359+ segsum: torch.Tensor,
360+ mimo_z: torch.Tensor = None,
361+ d: torch.Tensor = None,
362+ z: torch.Tensor = None,
363+ chunk_size: int = 16,
364+ rotary_dim_divisor: int = 4,
365+ output_dtype: torch.dtype = torch.float32,
366+):
367+ """Mamba-3 MIMO 反向第一遍 (bwd_fwd, 最小可测路径)。
368+ 
369+ 形状:
370+ dout: [B, S, H, P] (reduceO 上游梯度)
371+ q, k: [B, S, R, G, N]
372+ v: [B, S, H, P]
373+ q_bias, k_bias: [H, R, N]
374+ mimo_v (Psi): [H, R, P]
375+ mimo_o (Phi): [H, R, P] (reduceO 下投影; 本最小路径必需)
376+ angles: [B, S, H, N//rotary_dim_divisor]
377+ da_cs, da_cs_rev, dt, trap: [B, H, S]
378+ segsum: [B, H, nchunks, chunk_size, chunk_size]
379+ mimo_z (Zeta): [H, R, P] (可选, 配合 z 的 SiLU 门控)
380+ d: [H] (可选 D-skip)
381+ z: [B, S, H, P] (可选 门控输入)
382+ 返回 (均为 output_dtype):
383+ states: [B, H, nchunks, N, P] 每 chunk 进入时的递推状态
384+ qk_dot: [B, H, S, R, R] 逐步 q·k 对角块缓存
385+ dmimo_o: [B, H, R, P] Phi 梯度(逐 batch, 调用方再 sum(dim=0))
386+ dmimo_z: [B, H, R, P] 或 None Zeta 梯度(hasZ, 逐 batch)
387+ dz: [B, S, H, P] 或 None 门控输入梯度(hasZ)
388+ """
389+ B, S, R, G, N = q.shape
390+ H, P = v.shape[-2], v.shape[-1]
391+ RD = angles.shape[-1]
392+ assert mimo_o is not None, "bwd_fwd 最小路径要求 reduceO (mimo_o 非空)"
393+ assert S % chunk_size == 0, "本最小迁移要求 S 可被 chunk_size 整除 (tail_len==0)"
394+ assert N % 2 == 0 and RD <= N // 2
395+ assert H % G == 0
396+ 
397+ nchunks = S // chunk_size
398+ HALF = N // 2
399+ has_z = z is not None
400+ has_d = d is not None
401+ 
402+ Qr = torch.empty((B, H, S, R, N), device=v.device, dtype=torch.float32)
403+ Kr = torch.empty((B, H, S, R, N), device=v.device, dtype=torch.float32)
404+ 
405+ states = torch.empty((B, H, nchunks, N, P), device=v.device, dtype=output_dtype)
406+ qk_dot = torch.empty((B, H, S, R, R), device=v.device, dtype=output_dtype)
407+ dmimo_o = torch.empty((B, H, R, P), device=v.device, dtype=output_dtype)
408+ dmimo_z = torch.empty((B, H, R, P), device=v.device, dtype=output_dtype) if has_z else None
409+ dz = torch.empty((B, S, H, P), device=v.device, dtype=output_dtype) if has_z else None
410+ 
411+ total_pre = B * H * nchunks
412+ total_scan = B * H
413+ 
414+ def grid_pre(meta):
415+ return (min(get_vector_num(), total_pre),)
416+ 
417+ def grid_scan(meta):
418+ return (min(get_aicore_num(), total_scan),)
419+ 
420+ _mamba3_mimo_rotary_kernel[grid_pre](
421+ q,
422+ k,
423+ q_bias,
424+ k_bias,
425+ angles,
426+ Qr,
427+ Kr,
428+ q.stride(0),
429+ q.stride(1),
430+ q.stride(2),
431+ q.stride(3),
432+ q.stride(4),
433+ q_bias.stride(0),
434+ q_bias.stride(1),
435+ q_bias.stride(2),
436+ angles.stride(0),
437+ angles.stride(1),
438+ angles.stride(2),
439+ angles.stride(3),
440+ Qr.stride(0),
441+ Qr.stride(1),
442+ Qr.stride(2),
443+ Qr.stride(3),
444+ Qr.stride(4),
445+ B,
446+ S,
447+ H,
448+ G,
449+ nchunks,
450+ R=R,
451+ N=N,
452+ RD=RD,
453+ CHUNK=chunk_size,
454+ HALF=HALF,
455+ multibuffer=False,
456+ enable_auto_bind_sub_block=True,
457+ )
458+ 
459+ # C1: 对角块用纯向量 reduce (低精度输入); fp32 输入退回 cube matmul 取对角以对齐 ref 累加序
460+ # (vector reduce 在 fp32 下相对 ref 的 MERE 比值会顶到 ~2.05 略超 L0=2.0, 见 doc 要点3)。
461+ qk_use_dot = q.dtype == torch.float32
462+ _mamba3_mimo_qkdot_kernel[grid_pre](
463+ q,
464+ k,
465+ q_bias,
466+ k_bias,
467+ qk_dot,
468+ q.stride(0),
469+ q.stride(1),
470+ q.stride(2),
471+ q.stride(3),
472+ q.stride(4),
473+ q_bias.stride(0),
474+ q_bias.stride(1),
475+ q_bias.stride(2),
476+ qk_dot.stride(0),
477+ qk_dot.stride(1),
478+ qk_dot.stride(2),
479+ qk_dot.stride(3),
480+ qk_dot.stride(4),
481+ B,
482+ S,
483+ H,
484+ G,
485+ nchunks,
486+ R=R,
487+ N=N,
488+ CHUNK=chunk_size,
489+ QK_USE_DOT=qk_use_dot,
490+ multibuffer=False,
491+ enable_auto_bind_sub_block=True,
492+ )
493+ 
494+ mimo_z_arg = mimo_z if has_z else mimo_v
495+ d_arg = d if has_d else q
496+ z_arg = z if has_z else v
497+ dz_arg = dz if has_z else v
498+ dmz_arg = dmimo_z if has_z else dmimo_o
499+ sd_h = d.stride(0) if has_d else 0
500+ sz = z.stride() if has_z else (0, 0, 0, 0)
501+ sdz = dz.stride() if has_z else (0, 0, 0, 0)
502+ 
503+ _mamba3_mimo_bwd_fwd_scan_kernel[grid_scan](
504+ Qr,
505+ Kr,
506+ v,
507+ mimo_v,
508+ mimo_o,
509+ mimo_z_arg,
510+ d_arg,
511+ z_arg,
512+ dout,
513+ dt,
514+ trap,
515+ da_cs,
516+ da_cs_rev,
517+ segsum,
518+ states,
519+ dmimo_o,
520+ dmz_arg,
521+ dz_arg,
522+ Qr.stride(0),
523+ Qr.stride(1),
524+ Qr.stride(2),
525+ Qr.stride(3),
526+ Qr.stride(4),
527+ v.stride(0),
528+ v.stride(1),
529+ v.stride(2),
530+ v.stride(3),
531+ mimo_v.stride(0),
532+ mimo_v.stride(1),
533+ mimo_v.stride(2),
534+ sd_h,
535+ sz[0],
536+ sz[1],
537+ sz[2],
538+ sz[3],
539+ dout.stride(0),
540+ dout.stride(1),
541+ dout.stride(2),
542+ dout.stride(3),
543+ dt.stride(0),
544+ dt.stride(1),
545+ dt.stride(2),
546+ segsum.stride(0),
547+ segsum.stride(1),
548+ segsum.stride(2),
549+ segsum.stride(3),
550+ segsum.stride(4),
551+ states.stride(0),
552+ states.stride(1),
553+ states.stride(2),
554+ states.stride(3),
555+ states.stride(4),
556+ dmimo_o.stride(0),
557+ dmimo_o.stride(1),
558+ dmimo_o.stride(2),
559+ dmimo_o.stride(3),
560+ sdz[0],
561+ sdz[1],
562+ sdz[2],
563+ sdz[3],
564+ B,
565+ S,
566+ H,
567+ nchunks,
568+ R=R,
569+ N=N,
570+ P=P,
571+ CHUNK=chunk_size,
572+ HAS_D=has_d,
573+ HAS_Z=has_z,
574+ multibuffer=False,
575+ enable_auto_bind_sub_block=True,
576+ )
577+ return states, qk_dot, dmimo_o, dmimo_z, dz
578+ 
579+ 
580+# arch35 re-export 期望的代表 kernel 名
581+mamba3_mimo_bwd_fwd_kernel = _mamba3_mimo_bwd_fwd_scan_kernel
@@ -0,0 +1,108 @@
1+# Copyright (c) 2026, Dao AI Lab, Goombalab
2+# Copyright (c) 2026, Huawei Technologies Co., Ltd.
3+# pylint: disable=duplicate-code
4+# The public API and dispatcher intentionally forward the same operator contract.
5+#
6+# SPDX-License-Identifier: Apache-2.0
7+#
8+# Shape-adaptive production dispatcher for #22 (plan §3.5 candidate 7). Picks the
9+# fastest correct implementation per shape from measured paired A/B:
10+# * batch-parallel aux (chunk-local qk/intra/KV on B*H*nchunks) — wins when
11+# B*H underfills the cores (golden/big/s2048);
12+# * staged scalar scan — wins when B*H saturates the cores (b8h16), where the
13+# aux's INTRA/KV scratch traffic is a net loss.
14+# The common stage0 feature path (reduce_o=True, has_d, has_z=False, no fuse,
15+# S%C==0) routes through the fast variants; anything else falls back to the
16+# staged production ``mamba3_mimo_bwd_fwd`` (full feature coverage).
17+# ruff: noqa: E501
18+from __future__ import annotations
19+ 
20+import torch
21+ 
22+from mindspeed_ops.api.triton.utils import get_aicore_num
23+from mindspeed_ops.arch32.triton.mamba3.mamba3_mimo_bwd_fwd_impl import mamba3_mimo_bwd_fwd as _staged
24+from mindspeed_ops.arch32.triton.mamba3.mamba3_mimo_bwd_fwd_aux_impl import mamba3_mimo_bwd_fwd_aux as _aux
25+ 
26+ 
27+def mamba3_mimo_bwd_fwd_prod(
28+ dout,
29+ q,
30+ k,
31+ v,
32+ q_bias,
33+ k_bias,
34+ mimo_v,
35+ mimo_o,
36+ angles,
37+ da_cs,
38+ da_cs_rev,
39+ dt,
40+ trap,
41+ segsum,
42+ mimo_z=None,
43+ d=None,
44+ z=None,
45+ chunk_size: int = 16,
46+ rotary_dim_divisor: int = 4,
47+ output_dtype: torch.dtype = torch.float32,
48+):
49+ """Production #22: dispatch the common feature path to the fastest variant,
50+ else the full staged path. Returns ``(states, qk_dot, dmimo_o, dmimo_z, dz)``.
51+ """
52+ B, S, R, G, N = q.shape
53+ H = v.shape[-2]
54+ common = (mimo_o is not None) and (z is None) and (S % chunk_size == 0)
55+ if common:
56+ ncore = get_aicore_num()
57+ underutilized = (B * H) <= ncore # aux fills cores only when work < cores
58+ # The aux carries states[N,P] in one UB tile with no N-blocking; the
59+ # merged pre_aux is de-unrolled (runtime rank loops) so its UB stays
60+ # bounded across R (validated N16..256 x P32..128 x R1..8). Very large
61+ # N*P beyond the single-tile UB budget still routes to the staged _bt
62+ # N-blocked scan. (All four perf shapes + the official CASE_GRID fit.)
63+ P_ = v.shape[-1]
64+ aux_safe = (N * P_) <= 16384
65+ if underutilized and aux_safe:
66+ states, qk_dot, dmimo_o = _aux(
67+ dout=dout,
68+ q=q,
69+ k=k,
70+ v=v,
71+ q_bias=q_bias,
72+ k_bias=k_bias,
73+ mimo_v=mimo_v,
74+ mimo_o=mimo_o,
75+ angles=angles,
76+ da_cs=da_cs,
77+ dt=dt,
78+ trap=trap,
79+ d=d,
80+ chunk_size=chunk_size,
81+ rotary_dim_divisor=rotary_dim_divisor,
82+ output_dtype=output_dtype,
83+ )
84+ return states, qk_dot, dmimo_o, None, None
85+ # fallback: full staged production path (also handles z / fuse / non-reduce)
86+ out = _staged(
87+ dout=dout,
88+ q=q,
89+ k=k,
90+ v=v,
91+ q_bias=q_bias,
92+ k_bias=k_bias,
93+ mimo_v=mimo_v,
94+ mimo_o=mimo_o,
95+ angles=angles,
96+ da_cs=da_cs,
97+ da_cs_rev=da_cs_rev,
98+ dt=dt,
99+ trap=trap,
100+ segsum=segsum,
101+ mimo_z=mimo_z,
102+ d=d,
103+ z=z,
104+ chunk_size=chunk_size,
105+ rotary_dim_divisor=rotary_dim_divisor,
106+ output_dtype=output_dtype,
107+ )
108+ return out[0], out[1], out[2], out[3], out[4]
@@ -0,0 +1,33 @@
1+# Copyright (c) 2026, Huawei Technologies Co., Ltd. All rights reserved.
2+ 
3+import torch
4+ 
5+ 
6+def pad_tail_seq(x, dim, npad, mode="zero"):
7+ """Right-pad a tensor along one dimension for a complete final chunk."""
8+ if x is None or npad == 0:
9+ return x
10+ 
11+ shape = list(x.shape)
12+ shape[dim] = npad
13+ if mode == "edge":
14+ index = [slice(None)] * x.ndim
15+ index[dim] = slice(x.shape[dim] - 1, x.shape[dim])
16+ block = x[tuple(index)].expand(shape).contiguous()
17+ else:
18+ block = x.new_zeros(shape)
19+ return torch.cat([x, block], dim=dim)
20+ 
21+ 
22+def rebuild_dacs_tail(da_cs_p, da_cs_rev_p, segsum_p, chunk_size):
23+ """Rebuild the final chunk metadata after flat-tail padding ``da_cs``."""
24+ sequence_length = da_cs_p.shape[2]
25+ chunk_count = sequence_length // chunk_size
26+ chunk_start = (chunk_count - 1) * chunk_size
27+ chunk = da_cs_p[:, :, chunk_start : chunk_start + chunk_size]
28+ 
29+ da_cs_rev_p = da_cs_rev_p.clone()
30+ da_cs_rev_p[:, :, chunk_start : chunk_start + chunk_size] = chunk[:, :, -1:] - chunk
31+ segsum_p = segsum_p.clone()
32+ segsum_p[:, :, chunk_count - 1] = chunk[:, :, :, None] - chunk[:, :, None, :]
33+ return da_cs_rev_p, segsum_p