已开启
feat: mamba3-mimo-bwd-fwd-triton-kernel #101
bitszh3271创建于 7月16日
feat: mamba3-mimo-bwd-fwd-triton-kernel #101
已开启
共 17 个文件变更+7256-1
| @@ -39,7 +39,7 @@ repos: | |||
| 39 | - id: codespell | 39 | - 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 | ] |
| @@ -0,0 +1,219 @@ | |||
| 1 | +# mamba3_mimo_bwd_fwd 算子迁移说明 | ||
| 2 | + | ||
| 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 | + | ||
| 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 | + | ||
| 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 | + | ||
| 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 | + | ||
| 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 | + | ||
| 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 | + | ||
| 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 | + | ||
| 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 | + | ||
| 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 | + | ||
| 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 | + | ||
| 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 | ||


【openlibing.ci】检测到当前PR中存在代码检查告警抑制 7 处,详情见下表,请Committer检视合理性。 / Detected 7 code check alert suppression(s) in this PR, see table below. Committers please review.
mamba3_mimo_bwd_fwd_kernel.py
# pylint: disable=duplicate-codetriton/mamba3/mamba3_mimo_bwd_fwd_aux_impl.py
# pylint: disable=duplicate-code,too-many-linestriton/mamba3/mamba3_mimo_bwd_fwd_baseline_impl.py
# pylint: disable=duplicate-codetriton/mamba3/mamba3_mimo_bwd_fwd_dispatch_impl.py
# pylint: disable=duplicate-codetriton/mamba3_mimo_fwd.py
# pylint: disable=possibly-used-before-assignment,too-many-nested-blocksmamba3_mimo_bwd_fwd_kernel/
generate_mamba3_mimo_bwd_fwd_kernel.py
# pylint: disable=unsubscriptable-objectmamba3_mimo_bwd_fwd_kernel/
reference_impl.py
# pylint: disable=duplicate-code