跳转至

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)

参考:

1.1 核心思想(FA-2 三个 trick)

  1. Tiling:把 Q / K / V 切成 block,每个 block 在 SRAM 里完成 score + softmax + V 累加,避免 attention matrix 写入 HBM
  2. Online softmax:流式更新 max + 归一化常数,让 block-wise 计算等价于全序列 softmax
  3. 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 321 MB
2048 13.41 1.43 1183 MB
4096 47.84 6.01 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:

\[ Q_\text{absorbed} = Q W_{uk}^\top \quad\Rightarrow\quad \text{score} = Q_\text{absorbed} \cdot c_{kv}^\top \]

直接在压缩空间算 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 衰减:

\[ S_t = \mathrm{diag}(g_t)\, S_{t-1} + k_t v_t^\top, \qquad o_t = q_t^\top S_t \]

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:

\[ \mathrm{Lightning}(Q, K, V) = \frac{\phi(Q) \cdot (\phi(K)^\top V)}{\mathrm{norm}(\cdots)}, \quad \phi(x) = \mathrm{ELU}(x) + 1 \]

复杂度从 \(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 能"覆盖":

\[ S_t = (I - \beta_t v_t k_t^\top) S_{t-1} + \alpha_t v_t k_t^\top \]

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 离散化:

\[ h_t = \bar A_t h_{t-1} + \bar B_t x_t, \qquad y_t = C_t^\top h_t \]

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):

  1. softmax 唯一线性增长:4K → 64K 时间 0.122 → 1.796 ms(15× 慢
  2. 4 个 Triton linear kernel 全部接近常数:~0.02-0.03 ms 全程,跟 ctx 长度无关
  3. Mamba-2 Triton 最快(0.018 ms 稳定),尽管 D_inner=4096 比其他大 32×;它的 [BLOCK_DI, N_state] tile 设计正好打满 SRAM bandwidth
  4. GLA / KDA / Lightning 量级相同(0.023-0.027 ms),差别都是常数因子,不是 scaling
  5. 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

参考链接


上级 · A2 注意力机制全景