Attention Kernel 实现与正确性验证¶
理论复杂度只是上限。实际跑多快、能不能 work,完全取决于 kernel 实现。本文给出 6 个 attention kernel 的 Triton / PyTorch reference + RTX 4070 16GB 实测数据。
代码 kernels/ + tests/test_attention_kernels.py,跑 pytest tests/test_attention_kernels.py -xvs(Windows 用 triton-windows)。
一、Triton FlashAttention-2 (GQA-aware)¶
参考:
- 原版 FlashAttention-2(CUDA)
- Triton tutorial
fused-attention.py
1.1 核心思想(FA-2 三个 trick)¶
- Tiling:把 Q / K / V 切成 block,每个 block 在 SRAM 里完成 score + softmax + V 累加,避免 attention matrix 写入 HBM
- Online softmax:流式更新 max + 归一化常数,让 block-wise 计算等价于全序列 softmax
- Backward 重算 attention:训练时不存 \(S = QK^\top\),反向用同一个 kernel 重算
1.2 Forward kernel(GQA-aware)¶
完整 ~80 行 kernel 在 kernels/flash_attn.py。核心结构:
import triton
import triton.language as tl
import torch
@triton.jit
def _fwd_kernel_flash_attn(
Q, K, V, Out, L,
sm_scale,
stride_qz, stride_qh, stride_qm, stride_qd,
stride_kz, stride_kh, stride_kn, stride_kd,
stride_vz, stride_vh, stride_vn, stride_vd,
stride_oz, stride_oh, stride_om, stride_od,
Z, H_q, H_kv, N_CTX,
BLOCK_M: tl.constexpr,
BLOCK_N: tl.constexpr,
D: tl.constexpr,
IS_CAUSAL: tl.constexpr,
):
"""每 program 负责 [B, H_q] × Q block。"""
start_m = tl.program_id(0)
off_hz = tl.program_id(1)
off_z = off_hz // H_q
off_h_q = off_hz % H_q
off_h_kv = off_h_q * H_kv // H_q # GQA: 多个 q head 映射到同 kv head
# ... block_ptr 偏移 + load Q block 到 SRAM ...
q = tl.load(Q_block_ptr, boundary_check=(0, 1))
# online softmax 累积变量
m_i = tl.full((BLOCK_M,), -float("inf"), dtype=tl.float32)
l_i = tl.zeros((BLOCK_M,), dtype=tl.float32)
acc = tl.zeros((BLOCK_M, D), dtype=tl.float32)
hi = (start_m + 1) * BLOCK_M if IS_CAUSAL else N_CTX
for start_n in range(0, hi, BLOCK_N):
k = tl.load(K_block_ptr, boundary_check=(0, 1))
v = tl.load(V_block_ptr, boundary_check=(0, 1))
s = tl.dot(q, k) * sm_scale
if IS_CAUSAL:
offs_m = start_m * BLOCK_M + tl.arange(0, BLOCK_M)
offs_n = start_n + tl.arange(0, BLOCK_N)
s = tl.where(offs_m[:, None] >= offs_n[None, :], s, -float("inf"))
# online softmax 更新
m_ij = tl.maximum(m_i, tl.max(s, axis=1))
alpha = tl.exp(m_i - m_ij)
p = tl.exp(s - m_ij[:, None])
l_i = l_i * alpha + tl.sum(p, axis=1)
acc = acc * alpha[:, None] + tl.dot(p.to(v.dtype), v)
m_i = m_ij
K_block_ptr = tl.advance(K_block_ptr, (0, BLOCK_N))
V_block_ptr = tl.advance(V_block_ptr, (BLOCK_N, 0))
acc = acc / l_i[:, None]
tl.store(O_block_ptr, acc.to(Out.dtype.element_ty), boundary_check=(0, 1))
注意省略了 make_block_ptr 配置 + flash_attn_fwd Python wrapper —— 见 kernels/flash_attn.py。
1.3 数值正确性¶
4 个 shape 全部跟 FP32 reference 对齐(BF16 atol/rtol = 2e-2):
$ python -m pytest tests/test_attention_kernels.py::test_flash_attn_numerical -xvs
test_flash_attn_numerical[2-32-8-1024-128] PASSED # GQA-4 (LLaMA-3 8B-ish)
test_flash_attn_numerical[1-64-8-4096-128] PASSED # GQA-8 (LLaMA-3 70B-ish)
test_flash_attn_numerical[1-32-32-512-128] PASSED # MHA
test_flash_attn_numerical[1-16-1-512-128] PASSED # MQA
1.4 性能对比 vs eager(N sweep)¶
SDPA 不画了 —— 内部直接调 FlashAttention,跟本文 Triton kernel 完全同 backend,数字差不多。直接对照 eager(朴素 einsum + softmax)就够了。
| N | eager (ms) | Triton (ms) | speedup | eager 显存 |
|---|---|---|---|---|
| 1024 | 2.81 | 0.47 | 6× | 321 MB |
| 2048 | 13.41 | 1.43 | 9× | 1183 MB |
| 4096 | 47.84 | 6.01 | 8× | 4530 MB |
| 8192 | 8762.71 | 21.22 | 413× | 17716 MB → system RAM paging |
| 16384 | OOM | 79.09 | ∞ | 34 GB attn 矩阵 |
| 32768 | OOM | 318.99 | ∞ | 137 GB attn 矩阵 |
关键观察:
- N=1K–4K:eager 跟 Triton 都还能跑,speedup 6–9×,主要来自 online softmax + tile IO(FA-2 两个 trick)
- N=8192 是悬崖:eager 的 attn 矩阵 17.7 GB 已经超过 4070 的 16 GB 显存,但 PyTorch 不直接 OOM —— 它走 system RAM paging(unified memory),latency 从 48 ms 飙到 8.7 秒,慢 413 倍
- N≥16384:eager 真 OOM;Triton 在 4070 上稳定跑到 32K(55 TFLOPS)
- Triton TFLOPS 趋势:1K=37 → 8K=52 → 32K=55 —— N 越长 Tensor Core 利用率越高(block 越大越能 amortize launch 开销)
二、MLA absorbed — DeepSeek 吸收 trick¶
参考:
2.1 吸收 trick 数学¶
标准 MLA 推理:从 c_kv 上投影回 K = W_uk c_kv + 算 Q K^T + softmax + 上投影回 V + 乘 V。
吸收:把 W_uk 预合并到 Q:
直接在压缩空间算 score,KV cache 真就是 c_kv (d_c) + k_pe (d_R),不用从 c_kv 重建 K/V。
2.2 PyTorch reference 实现¶
# tests/test_attention_kernels.py 完整版:mla_naive + mla_absorbed 都给了
def mla_absorbed(c_kv, k_pe, q_nope, q_pe, W_uk, W_uv):
"""吸收:先 Q @ W_uk^T 把 Q 拉进压缩空间,再跟 c_kv 算 score。"""
B, T, d_c = c_kv.shape
H = W_uk.shape[0]
# 1. 吸收:Q @ W_uk^T → q_absorbed in compressed space
q_absorbed = torch.einsum("bhsd,hdc->bhsc", q_nope, W_uk) # [B, H, S, d_c]
# 2. score split: nope (压缩空间) + RoPE
s_nope = torch.einsum("bhsc,btc->bhst", q_absorbed, c_kv) # [B, H, S, T]
k_pe_b = k_pe.unsqueeze(1).expand(-1, H, -1, -1)
s_pe = torch.einsum("bhsd,bhtd->bhst", q_pe, k_pe_b)
d_h = q_nope.size(-1) + q_pe.size(-1)
s = (s_nope + s_pe) / math.sqrt(d_h)
p = F.softmax(s, dim=-1)
# 3. 在压缩空间算 out, 然后 W_uv 上投影
out_c = torch.einsum("bhst,btc->bhsc", p, c_kv) # [B, H, S, d_c]
return torch.einsum("bhsc,hvc->bhsv", out_c, W_uv) # [B, H, S, d_v]
def mla_naive(c_kv, k_pe, q_nope, q_pe, W_uk, W_uv):
"""HF DeepseekV3Attention 的默认路径:从 c_kv 重建完整 K/V,再标准 attention。"""
# 上投影回完整 K_nope, V
H = W_uk.shape[0]
k_nope = torch.einsum("btc,hdc->bhtd", c_kv, W_uk) # [B, H, T, d_h_nope]
v = torch.einsum("btc,hvc->bhtv", c_kv, W_uv) # [B, H, T, d_v]
k_pe_b = k_pe.unsqueeze(1).expand(-1, H, -1, -1)
k = torch.cat([k_nope, k_pe_b], dim=-1) # [B, H, T, d_h]
q = torch.cat([q_nope, q_pe], dim=-1)
s = torch.einsum("bhsd,bhtd->bhst", q, k) / math.sqrt(q.size(-1))
return torch.einsum("bhst,bhtv->bhsv", F.softmax(s, dim=-1), v)
数值等价:test_mla_numerical 在 FP32 下 assert_close(absorbed, naive, rtol=1e-4, atol=1e-4) 通过。
2.3 性能对比 vs naive 重建¶
DeepSeek-V3 spec(H=128, d_c=512, d_h_nope=d_v=128, d_R=64),decode 1 token:
| T | naive | absorbed | speedup | mem saved |
|---|---|---|---|---|
| 4096 | 3.66 ms / 472 MB | 0.193 ms / 4.5 MB | 19× | 105× |
| 16384 | 14.87 ms / 1888 MB | 0.383 ms / 17 MB | 38.9× | 110× |
take-away:MLA 吸收在 decode 阶段是数量级量变——长 ctx 越长收益越大。这是 DeepSeek-V3 单卡 H100 serve 128K context 的关键。完整 GPU kernel(DeepSeek 自家 FlashMLA)在 H100 上更快,但 PyTorch reference 已经能体现 trick 本身的价值。
2.4 TileLang FlashMLA decode(v1 + split-T v2,Linux/WSL)¶
TileLang 是 tile-ai 团队开发的 GPU kernel DSL,比 Triton 抽象层更高(pipeline / shared / fragment 直接是一等公民)。Windows 没 wheel(PR #2093 进行中,未 merge),目前要 WSL2 / aliyun99 跑。
实现见 kernels/mla_tilelang.py,两版:
- v1:grid
(B, H/BLOCK_H),每 block 串完整 T 轴,online softmax。简单但低 batch 时 SM 占用率低。 - v2:split-T pattern。grid
(B, H/BLOCK_H, NUM_SPLIT_T)跑 partial kernel,pytorch 收尾做 safe-softmax 合并。RTX 4070 Ti SUPER (16GB) fp16 实测:
| B | torch ref | TL v1 | TL v2 | v1 vs torch | v2 vs torch | v1 max err |
|---|---|---|---|---|---|---|
| 1 | 0.36 ms | 2.15 ms | 0.46 ms | 0.17× | 0.78× | 0.0000 |
| 8 | 2.23 ms | 1.84 ms | 3.68 ms | 1.21× | 0.60× | 0.0000 |
| 32 | 0.78 ms | 0.85 ms | 1.06 ms | 0.92× | 0.73× | 0.0000 |
数值精确(max err 0),但 v2 split-T 在这个 RTX 4070 形状下没赢 v1 —— pytorch reduce 加了 4 个 launch 开销 (~40us),掩盖了 split-T 的并行度收益。在 H100 + 真 FlashMLA 形状下 split-T 会赢,因为:(1) NUM_SPLIT 用 fused TileLang reduce 而不是 pytorch;(2) batch 更大、heads 更多时 v1 的低占用瓶颈才显著。
TileLang vs Triton 体感(同样 FlashMLA decode,假想 Triton 版):
| TileLang v1 | Triton 等价 | |
|---|---|---|
| LOC | 96 行 | ~150-200 行(手写 async copy + double buffer) |
| 流水线 | T.Pipelined(num_stages=2) 一行 |
自己 tl.async_copy + barrier + advance ptr |
| GEMM | T.gemm(A, B, C, transpose_B=True) |
tl.dot(a, tl.trans(b)) + 自选 BLOCK_M / BLOCK_N |
| Layout | 自动推断(偶尔需 shared 中转) | 自己定 BLOCK + warp tiling |
| 调试 | TileLang IR 报错可读 | PTX 报错经常断在 register spill |
| 坑 | fragment SSA 不能跨循环重赋 Python 变量;上下游 gemm fragment layout 推断会冲突,要用 T.copy(P_f, P_s) 走 shared 解耦 |
自己手写 epilogue + 数值不稳时要 fp32 accum |
结论:teaching / 中等复杂度 attention 变体 → TileLang 占优;极致 perf 调到底 → Triton + ncu 还是天花板(因为可以手撸 register usage / async pipeline / warp specialize 每一处)。
三、GLA — Gated Linear Attention recurrent step¶
参考 GLA paper (Yang et al., 2023)。
3.1 核心思想¶
state matrix \(S \in \mathbb{R}^{d \times d}\) 累积所有历史,每步用 sigmoid gate 衰减:
3.2 PyTorch reference 实现¶
def gla_step(q_t, k_t, v_t, g_t, S_prev):
"""One recurrent step.
q_t, k_t, v_t, g_t: [B, H, D]
S_prev: [B, H, D, D] ← 固定大小!
"""
# state 更新:每行 i 按 g_t[i] 衰减,再加 outer(k, v)
S_new = g_t.unsqueeze(-1) * S_prev + torch.einsum("bhi,bhj->bhij", k_t, v_t)
o_t = torch.einsum("bhi,bhij->bhj", q_t, S_new)
return o_t, S_new
PyTorch 版本会拆成 3-4 个独立 kernel(broadcast multiply / outer / add / matmul),每次 HBM round-trip。
3.3 Triton fused kernel — decode step¶
完整代码:kernels/gla_triton.py::gla_decode_step_triton。核心 kernel:
@triton.jit
def _gla_decode_step_kernel(
Q_PTR, K_PTR, V_PTR, G_PTR, S_PREV_PTR, S_NEW_PTR, OUT_PTR,
stride_qbh, stride_qd,
stride_sbh, stride_si, stride_sj,
BH: tl.constexpr, D: tl.constexpr,
):
"""One GLA step, fully fused. Per (batch, head):
- load q, k, v, g [D]
- load s_prev [D, D] into SRAM
- s_new[i,j] = g[i] * s_prev[i,j] + k[i] * v[j]
- o[j] = sum_i q[i] * s_new[i, j]
- write s_new, o
State [D, D] (D=128 → 64 KB BF16) fits easily in 4070 SRAM.
"""
pid = tl.program_id(0)
offs_d = tl.arange(0, D)
q = tl.load(Q_PTR + pid * stride_qbh + offs_d * stride_qd).to(tl.float32)
k = tl.load(K_PTR + pid * stride_qbh + offs_d * stride_qd).to(tl.float32)
v = tl.load(V_PTR + pid * stride_qbh + offs_d * stride_qd).to(tl.float32)
g = tl.load(G_PTR + pid * stride_qbh + offs_d * stride_qd).to(tl.float32)
offs_i = offs_d[:, None]; offs_j = offs_d[None, :]
s_ptr = S_PREV_PTR + pid * stride_sbh + offs_i * stride_si + offs_j * stride_sj
s_prev = tl.load(s_ptr).to(tl.float32)
s_new = g[:, None] * s_prev + k[:, None] * v[None, :] # [D, D]
o = tl.sum(q[:, None] * s_new, axis=0) # [D]
s_new_ptr = S_NEW_PTR + pid * stride_sbh + offs_i * stride_si + offs_j * stride_sj
tl.store(s_new_ptr, s_new)
tl.store(OUT_PTR + pid * stride_qbh + offs_d * stride_qd, o)
关键:整个 [D, D] state 在 SRAM 里完成 update + readout,HBM 只读一次 s_prev、写一次 s_new + o。PyTorch 版本会跑 3-4 趟 HBM。
3.4 Triton chunk-wise kernel — prefill scan¶
production GLA 训练用的两层并行模式(mirrors flash-linear-attention fla):
- chunk 内:T_chunk 个 token 在 SRAM 里串行(state 不出 register),单 program 处理整段
- chunk 间:跨 (batch, head) 并行 dispatch;同一 (batch, head) 顺序 launch chunk
完整代码:kernels/gla_triton.py::gla_chunk_fwd_triton。核心 kernel 对每个 (batch, head) 处理一个 chunk:
@triton.jit
def _gla_chunk_kernel(Q_PTR, K_PTR, V_PTR, G_PTR, S_PTR, OUT_PTR, ..., T, D):
pid = tl.program_id(0)
offs_d = tl.arange(0, D)
s = tl.load(S_PTR + pid * stride_sbh + ...).to(tl.float32) # carrier state from prev chunk
for t in range(T): # ← 串行 in registers,每步不出 SRAM
q = tl.load(Q_PTR + ...)
# ... k, v, g 同理
s = g[:, None] * s + k[:, None] * v[None, :]
o = tl.sum(q[:, None] * s, axis=0)
tl.store(OUT_PTR + ..., o)
tl.store(S_PTR + ..., s) # 写回 carrier state for next chunk
for t in range(T) 在 Triton 里展开成串行 instructions,但 state s 全程在 register/SRAM,没 HBM round-trip。这是 chunk-wise scan 的核心 trick:用串行换 IO。
3.5 性能对比¶
3.5.1 Decode step(softmax / PyTorch GLA / Triton fused)¶
RTX 4070 Ti SUPER, BF16, B=1, H=32, D=128(稳态数字,warmup 之后):
| ctx N | softmax decode | PyTorch GLA step | Triton GLA step | Triton vs PyTorch | Triton vs softmax |
|---|---|---|---|---|---|
| 4096 | 0.125 ms | 0.062 ms | 0.023 ms | 2.7× | 5.5× |
| 16384 | 0.494 ms | 0.046 ms | 0.025 ms | 1.8× | 20× |
| 65536 | 1.893 ms | 0.044 ms | 0.025 ms | 1.8× | 77× |
Triton fused 稳定比 PyTorch 快 ~2×(PyTorch warmup 后其实也不慢,但 Triton fused kernel 把 4 个 op 合成 1 个 + state 留 SRAM 还是省了 HBM round-trip)。真正的胜利是 vs softmax decode:N=64K 时 Triton GLA 比 softmax 快 77×,因为 GLA 不读 N-token KV cache。
3.5.2 Prefill scan(PyTorch loop / Triton chunk-wise)¶
跑整段 T-token 训练 prefill(稳态):
| T | PyTorch step-by-step loop | Triton chunk-wise (chunk=64) | speedup |
|---|---|---|---|
| 1024 | 72.55 ms | 1.68 ms | 43× |
| 4096 | 308.68 ms | 6.60 ms | 47× |
PyTorch loop 的 ~300 ms 大头是 每步 4-5 个 kernel launch × T 步 + Python 循环 overhead;Triton chunk-wise 把每个 chunk(64 token)的 state 留在 SRAM 串行处理,HBM 只读 q/k/v/g 一次 + 写 out 一次,一整段 chunk 摊一次 launch。~45× 加速就是 production GLA / Mamba kernel 的真实样子。
Linear 家族 production 还是 Triton 主导(
flash-linear-attention(fla) 的 GLA / DeltaNet / KDA / RWKV 全是 Triton),TileLang 现阶段没强生态。但 GLA chunk-wise 我也写了 TileLang 版做对照(kernels/gla_tilelang.py),WSL fp16 实测:
| T | Triton chunk-wise | TileLang chunk-wise | TL vs Triton |
|---|---|---|---|
| 256 | 0.50 ms | 0.53 ms | 0.95× |
| 1024 | 1.55 ms | 1.88 ms | 0.82× |
数值对齐(max err 0 vs Triton),TileLang 略慢 5-18%。state 矩阵 [D, D] 走 shared mem + per-token sequential update + readout gemv 三件套,Triton 那边把 state 直接放 register(更激进,但要写 boilerplate)。TileLang 写起来短 ~30 行,调试更线性。
3.6 数值正确性¶
$ python -m pytest tests/test_attention_kernels.py -k "triton" -xvs
test_gla_triton_decode_numerical PASSED # FP32 atol=1e-5, BF16 atol=0.5
test_gla_triton_chunk_numerical PASSED # chunk-wise vs step-by-step loop bit-exact (FP32)
gla_decode_step_triton 跟 PyTorch reference gla_step FP32 下 atol=1e-5 完全对齐,BF16 累加 8-bit mantissa 误差在合理范围(D=128 sum 后 ~0.5 abs diff,~5% rel diff)。gla_chunk_fwd_triton 跟 step-by-step PyTorch loop FP32 atol=1e-4 对齐。
3.7 backward¶
PyTorch reference 走 autograd —— 反向时为了不存所有 [T, D, D] state,autograd 实际是 recompute(重新跑 forward 拿 intermediate state)。production GLA backward kernel 同思路:把 forward 拆成 chunk,反向重算每个 chunk 内的 state。本文不展开(参考 fla repo gla/chunk.py)。
四、Lightning Attention — transnormer 风格¶
参考 MiniMax-M1 paper (2025) + Yang et al. transnormer。
4.1 核心思想¶
不用 softmax,改用 ELU+1 feature map 让 K, Q 非负,然后先算 \(K^\top V\) 拿到固定大小 state,再 \(Q \cdot\) state:
复杂度从 \(O(N^2 d)\) 降到 \(O(N d^2)\),因为 \(K^\top V\) 是 \(d \times d\) 跟 \(N\) 无关。
4.2 PyTorch reference 实现¶
def lightning_step(q_t, k_t, v_t, kv_state, eps=1e-6):
"""Transnormer 风格 Lightning Attention 单步 decode。
kv_state: [B, H, D, D] ← 累加 K^T V,固定大小
"""
q_phi = F.elu(q_t) + 1
k_phi = F.elu(k_t) + 1
kv_state = kv_state + torch.einsum("bhi,bhj->bhij", k_phi, v_t)
o_t = torch.einsum("bhi,bhij->bhj", q_phi, kv_state)
z = torch.einsum("bhi,bhij->bhj", q_phi, kv_state) + eps
return o_t / (z.norm(dim=-1, keepdim=True) + eps), kv_state
PyTorch 拆成 4 个 op:两次 ELU+1 + 外积 einsum + 矩阵-向量 einsum + 除法。
4.3 Triton fused kernel¶
完整代码:kernels/linear_triton.py::lightning_step_triton。把 ELU+1 / state update / readout / L2 norm 全 fuse 进一个 kernel,state 在 SRAM:
@triton.jit
def _lightning_step_kernel(Q_PTR, K_PTR, V_PTR, S_PTR, ..., D):
pid = tl.program_id(0)
offs_d = tl.arange(0, D)
q = tl.load(...).to(tl.float32)
k = tl.load(...).to(tl.float32)
v = tl.load(...).to(tl.float32)
# ELU+1 feature map (transnormer style)
q_phi = tl.where(q >= 0, q + 1.0, tl.exp(q))
k_phi = tl.where(k >= 0, k + 1.0, tl.exp(k))
s_prev = tl.load(S_PTR + ...).to(tl.float32)
s_new = s_prev + k_phi[:, None] * v[None, :] # [D, D] state update
o_pre = tl.sum(q_phi[:, None] * s_new, axis=0) # [D] readout
# L2 norm output
z = o_pre + EPS
z_norm = tl.sqrt(tl.sum(z * z, axis=0)) + EPS
o = o_pre / z_norm
tl.store(...)
4.4 性能对比¶
RTX 4070 Ti SUPER, BF16, B=1, H=32, D=128:
| ctx N | softmax decode | PyTorch step | Triton fused | Triton vs PyTorch | Triton vs softmax |
|---|---|---|---|---|---|
| 4096 | 0.122 ms | 0.131 ms | 0.026 ms | 5.0× | 4.7× |
| 16384 | 0.475 ms | 0.131 ms | 0.027 ms | 4.9× | 17.5× |
| 65536 | 1.796 ms | 0.130 ms | 0.027 ms | 4.9× | 67.7× |
PyTorch 在短 ctx (N=4K) 实际比 softmax 慢(0.131 vs 0.122 ms)—— ELU+1 + 4 个 einsum + L2 norm 加起来的 kernel-launch overhead 比单 softmax 还重。Triton fused 把这些全压成 1 个 kernel + state SRAM resident,5× 加速 vs PyTorch,crossover 消失。
take-away:Lightning 在短 ctx (4K) 反而慢 —— normalization + ELU+1 的常数开销大于 softmax 在小 N 下的成本。只有长 ctx 才划算。这就是 MiniMax-M1 用 7:1 (Lightning : softmax MHA) 而不是纯 Lightning 的原因之一:短段交给 softmax,长段交给 Lightning。
五、KDA — Kimi Delta Attention(delta rule)¶
参考 Kimi Linear paper (2025) + GitHub。
5.1 核心思想¶
GLA 只能"衰减"旧 (k, v),KDA 用 delta rule 让 state 能"覆盖":
per-dim \(\alpha, \beta\) 双向量(fine-grained),转移矩阵是 DPLR(diagonal-plus-low-rank)特殊变体。
5.2 PyTorch reference 实现¶
def kda_step(q_t, k_t, v_t, alpha_t, beta_t, S_prev):
"""KDA delta-rule step.
S_t = (I - β v k^T) S_{t-1} + α v k^T
per-dim α, β; rank-1 修正让 state 精确覆盖旧 (k, v).
"""
B, H, D = q_t.shape
vk = torch.einsum("bhi,bhj->bhij", v_t, k_t) # [B, H, D, D]
beta_b = beta_t.unsqueeze(-1) # [B, H, D, 1]
alpha_b = alpha_t.unsqueeze(-1) # [B, H, D, 1]
# rank-1 校正项 - β·outer(v,k)·S 撤销旧、+ α·outer(v,k) 写新
S_new = S_prev - beta_b * torch.einsum("bhij,bhjk->bhik", vk, S_prev) + alpha_b * vk
o_t = torch.einsum("bhi,bhij->bhj", q_t, S_new)
return o_t, S_new
5.3 Triton fused kernel¶
完整代码:kernels/linear_triton.py::kda_step_triton。关键 trick:把 vk @ S_prev 的矩阵乘展开成两个累加 —— 先 r[j] = Σ_l k[l] · S_prev[l, j](matrix-vector),再 delta[i, j] = v[i] · (α[i]·k[j] - β[i]·r[j]),避免显式构造 [D, D] 的 vk 中间张量:
@triton.jit
def _kda_step_kernel(Q_PTR, K_PTR, V_PTR, A_PTR, B_PTR, S_PTR, ..., D):
pid = tl.program_id(0)
offs_d = tl.arange(0, D)
q = tl.load(...).to(tl.float32)
k = tl.load(...).to(tl.float32)
v = tl.load(...).to(tl.float32)
alpha = tl.load(...).to(tl.float32)
beta = tl.load(...).to(tl.float32)
s_prev = tl.load(...).to(tl.float32)
# r[j] = sum_l k[l] * S_prev[l, j] ← matrix-vector (k^T S_prev)
r = tl.sum(k[:, None] * s_prev, axis=0) # [D]
# delta[i, j] = v[i] * (α[i] · k[j] - β[i] · r[j])
delta = v[:, None] * (alpha[:, None] * k[None, :] - beta[:, None] * r[None, :])
s_new = s_prev + delta
o = tl.sum(q[:, None] * s_new, axis=0)
tl.store(...)
不构造 vk 中间矩阵省一遍 D² 写入;整个 DPLR rank-1 update + readout 在 SRAM 完成。
5.4 性能对比¶
RTX 4070 Ti SUPER, BF16, B=1, H=32, D=128:
| ctx N | softmax decode | PyTorch step | Triton fused | Triton vs PyTorch | Triton vs softmax |
|---|---|---|---|---|---|
| 4096 | 0.122 ms | 0.084 ms | 0.027 ms | 3.2× | 4.6× |
| 16384 | 0.463 ms | 0.082 ms | 0.022 ms | 3.7× | 21× |
| 65536 | 1.794 ms | 0.096 ms | 0.022 ms | 4.4× | 83× |
KDA 比 GLA 多了 vk @ S_prev 的 matrix-vector reduction,PyTorch 时数组重排开销大;Triton fuse 后稳定 3-4× 加速 vs PyTorch。Production Kimi Linear KDA kernel 用 chunk-wise matmul 进一步 fuse training 的整段 prefill。
六、Mamba-2 — 选择性 SSM 单步¶
参考 Mamba-2 paper (Dao & Gu, ICML 2024) + mamba_ssm repo。
6.1 核心思想¶
SSM 离散化:
Mamba (S6) 让 \(\bar A, \bar B, C, \Delta\) 全部输入相关(selective)。Mamba-2 SSD 框架证明等价于 linear attention,可用 chunk-wise matmul 训练(生产级 kernel selective_scan_cuda)。
6.2 PyTorch reference 实现¶
def mamba2_step(x_t, A_log, B_t, C_t, dt_t, h_prev):
"""Mamba-2 single decode step.
x_t: [B, D_inner]
A_log: [D_inner, N_state] (learned, A = -exp(A_log))
B_t/C_t: [B, N_state] (data-dependent)
dt_t: [B, D_inner] (data-dependent timescale)
h_prev: [B, D_inner, N_state]
"""
A = -torch.exp(A_log) # [D_inner, N_state]
# Discretize per-channel: A_bar[b,d,n] = exp(dt[b,d] * A[d,n])
A_bar = torch.exp(dt_t.unsqueeze(-1) * A.unsqueeze(0)) # [B, D_inner, N_state]
B_bar = dt_t.unsqueeze(-1) * B_t.unsqueeze(1) # [B, D_inner, N_state]
h_new = A_bar * h_prev + B_bar * x_t.unsqueeze(-1)
y_t = (h_new * C_t.unsqueeze(1)).sum(-1) # [B, D_inner]
return y_t, h_new
6.3 Triton fused kernel¶
Mamba-2 的状态形状跟 GLA 不同(per-channel [B, D_inner, N_state],D_inner=4096 在 Nemotron-Nano-2 量级),需要沿 D_inner 维 tile。完整代码:kernels/mamba2_triton.py::mamba2_step_triton:
@triton.jit
def _mamba2_step_kernel(X_PTR, A_LOG_PTR, B_PTR, C_PTR, DT_PTR, H_PTR, ...):
pid_b = tl.program_id(0) # batch
pid_di = tl.program_id(1) # D_inner tile index
BLOCK_DI: tl.constexpr = 64
di_offs = pid_di * BLOCK_DI + tl.arange(0, BLOCK_DI)
n_offs = tl.arange(0, N_STATE)
# Load per-batch B, C [N_state] + per-(batch,di) x, dt [BLOCK_DI]
Bv = tl.load(B_PTR + ...).to(tl.float32)
Cv = tl.load(C_PTR + ...).to(tl.float32)
x = tl.load(X_PTR + ...).to(tl.float32)
dt = tl.load(DT_PTR + ...).to(tl.float32)
# Load A_log [BLOCK_DI, N_state] + h_prev [BLOCK_DI, N_state]
A_log = tl.load(A_LOG_PTR + ...).to(tl.float32)
h_prev = tl.load(H_PTR + ...).to(tl.float32)
# Discretize + state update + output
dA = tl.exp(dt[:, None] * (-tl.exp(A_log))) # [BLOCK_DI, N_state]
dB = dt[:, None] * Bv[None, :]
h_new = dA * h_prev + dB * x[:, None] # selective scan step
y = tl.sum(h_new * Cv[None, :], axis=1) # [BLOCK_DI]
tl.store(...)
核心优化:把 5 个 PyTorch op(exp(A_log), exp(dt*A), dt*B*x, A*h+..., (h*C).sum)合成一个 kernel,且 h_prev 整 [BLOCK_DI, N_state] tile 留 SRAM,省 4 趟 HBM round-trip。
6.4 性能对比¶
Mamba-2 跟 softmax 直接对比是 apples to oranges(Mamba 是 per-channel state,softmax 是 KV cache),但都是 1-token decode 延迟,可以横向对照:
RTX 4070 Ti SUPER, BF16, D_inner=4096, N_state=128:
| ctx N (softmax) | softmax decode | Mamba-2 PyTorch step | Mamba-2 Triton fused | Triton vs PyTorch | Triton vs softmax |
|---|---|---|---|---|---|
| 4096 | 0.121 ms | 0.071 ms | 0.016 ms | 4.5× | 7.7× |
| 16384 | 0.455 ms | 0.074 ms | 0.019 ms | 3.9× | 24× |
| 65536 | 1.795 ms | 0.076 ms | 0.018 ms | 4.3× | 101× |
Mamba-2 Triton 在 4070 上 0.018 ms 单步 —— 已经接近官方 selective_scan_cuda 量级(H100 上 ~0.05 ms 是因为 batch 维度 dispatch 不同;本文 B=1 单 batch 4070 反而占便宜)。101× vs softmax @ N=64K 的优势完全来自 state 不依赖 ctx 长度。
七、Linear 路线 4 个 Triton kernel 横向对比¶
GLA / Lightning / KDA / Mamba-2 全部用 Triton fused 实现,跟 softmax decode 画一张图直接看常数延迟:
KV cache vs constant state 是 linear 路线真正的卖点:N=64K 时 softmax 要 1 GB KV cache(∝ N),linear-family 不管多长都 1 MB state。decode 单步显存差 1024×。
几个观察(Triton fused 实测,B=1, H=32, D=128, BF16):
- softmax 唯一线性增长:4K → 64K 时间 0.122 → 1.796 ms(15× 慢)
- 4 个 Triton linear kernel 全部接近常数:~0.02-0.03 ms 全程,跟 ctx 长度无关
- Mamba-2 Triton 最快(0.018 ms 稳定),尽管 D_inner=4096 比其他大 32×;它的 [BLOCK_DI, N_state] tile 设计正好打满 SRAM bandwidth
- GLA / KDA / Lightning 量级相同(0.023-0.027 ms),差别都是常数因子,不是 scaling
- N=64K 时 4 个 linear kernel 全部 vs softmax 60-100×(Mamba-2 最高 101×, Lightning 最低 67×)
这就是为什么 frontier hybrid(Kimi Linear 3:1, MiniMax M1 7:1, Nemotron-H 1:11)都把 Linear 当主力 + 少量 Full attention 补精度 —— decode 速度 + 内存常数,长 ctx 唯一可行选项。
4 个 Triton kernel 都比对应 PyTorch reference 快 3-5×(来自 op fusion + state SRAM resident)。production 版本(fla / Kimi-Linear / mamba_ssm / MiniMax kernel)会再加 chunk-wise prefill scan + autotuning,比本文教学版还快 1-2×。
八、跑法¶
# 跑全部 23 个测试(数值正确性 + 6 个 kernel 的 bench):
python -m pytest tests/test_attention_kernels.py -xvs --capture=no
# 只跑数值正确性(任何 CUDA GPU 都行):
python -m pytest tests/test_attention_kernels.py -k numerical -xvs
# 只跑 bench(数据用来更新本文图表):
python -m pytest tests/test_attention_kernels.py -k bench -xvs --capture=no
# 重新生成图表(用最新的 bench 数据,需要手动更新 scripts/gen_kernel_charts.py 里的常量):
python scripts/gen_kernel_charts.py
依赖:
- Linux / WSL / macOS:
torch >= 2.4+triton+pytest+pyecharts - Windows:
torch >= 2.4+triton-windows+pytest+pyecharts - TileLang: 仅 Linux + macOS arm64(Windows 需源码 build 或 WSL)
九、生产 kernel 选型对照¶
实践中谁也不会从头写 kernel,下面是直接调库的选型:
| 场景 | 选什么 | 理由 |
|---|---|---|
| LLaMA / Mistral / Qwen 推理 | flash-attn (Tri Dao) |
GQA 原生支持,FA-2 / FA-3 自动选 |
| DeepSeek-V3 / V3.1 推理 | FlashMLA (DeepSeek 官方) |
吸收 trick + Hopper TMA |
| GLA / DeltaNet / KDA | flash-linear-attention (fla) |
chunk-wise scan production kernel |
| Mamba-2 / SSM | mamba_ssm (selective_scan_cuda) |
SSD chunk matmul + 官方 |
| Lightning Attention | MiniMax-M1 自家 kernel | 跟 hybrid attention 配合 |
| 自定义 attention 变体(科研) | Triton DSL | 上手快,1 天能写出 baseline |
| 极致性能 + 熟悉 CUDA | 手写 CUDA | 最后 5-10% 性能时才值得 |
| 跨平台(AMD ROCm) | TileLang | 编译目标可换 |
| PyTorch 2.x 自动加速 | torch.compile + F.scaled_dot_product_attention |
默认 FA-2 / FA-3 backend |
参考链接¶
- Full 路线: FlashAttention (FA-2 论文 / FA-3 论文) · Triton tutorial
- MLA: DeepSeek FlashMLA · TileLang FlashMLA 教程
- Linear 系: GLA paper · flash-linear-attention · Kimi Linear · MiniMax-M1
- SSM: Mamba-2 / SSD ·
mamba_ssm - 跨家族: The Anatomy of a Triton Attention Kernel
↑ 上级 · A2 注意力机制全景