第 11 章 · Attention 入门

⏱️ 70 分钟🎯 写出 SDPA📂 code/ch11_attention/

学习目标

前置知识

已完成 Ch09–10,能实现 GEMM、数值稳定 Softmax,并理解 row-major 张量布局。

核心概念

11.1 数学复习

Scaled Dot-Product Attention (single head, single batch):

S = Q @ K^T / sqrt(D)     # (T, T)
P = softmax_row(S + mask) # (T, T)
O = P @ V                 # (T, D)

关键代码

11.2 Fused QKV Projection

输入 X: (T, D)。原本要做三次 GEMM:

Q = X @ W_q   (T, D) @ (D, D)
K = X @ W_k
V = X @ W_v

W_qkv = [W_q; W_k; W_v] 水平拼成 (D, 3D),一次 GEMM 出 (T, 3D) 再 split:

// 1) 1 个 GEMM
gemm(X, W_qkv, QKV, T, 3*D, D);

// 2) split (T, 3*D) → 三个 (T, D)
__global__ void split_qkv(...) {
    int t = ..., d = ...;
    q[t*D+d] = qkv[t*3*D + 0*D + d];
    k[t*D+d] = qkv[t*3*D + 1*D + d];
    v[t*D+d] = qkv[t*3*D + 2*D + d];
}

好处:

工业实践: TRT-LLM、vLLM、llama.cpp 都用 fused QKV。Llama 还做了进一步的 GQA(Grouped Query Attention)让 K/V 投影更少,下一章后会顺带提到。

11.3 朴素三阶段 Attention

源码:attention_naive.cu

// Step 1: S = Q @ K^T * scale   (T, T)
__global__ void qkt_scale(const float* Q, const float* K, float* S,
                          int T, int D, float scale, bool causal) {
    int i = ...,  j = ...;
    if (causal && j > i) { S[i*T+j] = -1e30f; return; }
    float s = 0;
    for (int d = 0; d < D; ++d) s += Q[i*D+d] * K[j*D+d];
    S[i*T + j] = s * scale;
}

// Step 2: P = softmax_row(S)
softmax_rows<256><<<T, 256>>>(S, T, T);

// Step 3: O = P @ V             (T, D)
__global__ void pv_kernel(const float* P, const float* V, float* O, int T, int D) {
    int i = ..., d = ...;
    float acc = 0;
    for (int j = 0; j < T; ++j) acc += P[i*T+j] * V[j*D+d];
    O[i*D+d] = acc;
}

11.4 朴素实现的两个致命问题

问题 1:HBM 物化 T×T

S 和 P 都是 T × T,必须先全部写到 HBM 再读。T=2048 时 S 占 16 MiB,T=8192 时 S 占 256 MiB—— 不光(多 3 次全量 HBM 读写),显存还放不下长上下文

问题 2:softmax 是 memory-bound

softmax 算每个元素 O(1) FLOP 但要扫两遍 T。kernel 间没法融合,每次都得回 HBM。

graph LR
    Q["Q in HBM"] --> K1["QKᵀ kernel"]
    K1 --> S1["S in HBM"]
    S1 --> K2["softmax kernel"]
    K2 --> P1["P in HBM"]
    P1 --> K3["PV kernel"]
    K3 --> O1["O in HBM"]
    style S1 fill:#f3f1e8,stroke:#a82c30
    style P1 fill:#f3f1e8,stroke:#a82c30

两次完整的 T×T 显存读写——这就是 FlashAttention 要彻底干掉的地方。

运行结果

11.5 性能与瓶颈

TDS 显存朴素 ms (T4)FlashAttn ms (Ch12)
256640.25 MiBTODO(on GPU)TODO(on GPU)
1024644 MiBTODO(on GPU)TODO(on GPU)
40966464 MiBTODO(on GPU)TODO(on GPU)
163841281 GiBTODO(on GPU)TODO(on GPU)

长上下文场景下 FlashAttention 不是"快几倍"——是"能不能跑"的差别。

自检清单

Q1: 为什么 attention 要除 √D?

D 维度的 dot product 会随 D 增大变大(方差 = D)。除 √D 把 logits 方差归一到 1,避免 softmax 进饱和区梯度消失。

Q2: causal mask 为啥用 -1e30 不用 -inf?

fp32 里 -inf 也行,但 fp16 里 -inf - 大数 可能产生 NaN(uncovered subtraction)。-1e30 在 fp32 里 exp(-1e30) = 0 同效果,且更稳。

Q3: multi-head 怎么处理?

把 Q, K, V 改成 (B, H, T, D) 4D 张量。最简单的实现:每 head 一个 grid block,复用 single-head kernel。fused QKV 时 W_qkv = (D, 3*H*D)。

Q4: GQA 是什么?

Grouped Query Attention:多个 Q head 共享同一组 K/V head。Llama 2/3、Mistral 都用。好处:KV cache 显存减小(H_kv < H_q),decode 阶段更快。

Q5: attention 的 FLOPs 算多少?

QKᵀ: 2·T²·D,PV: 2·T²·D,加 softmax 的少量 FLOPs。总共 ≈ 4·T²·D。所以 T 翻倍 FLOPs 翻 4 倍,T=8K 时 attention 已经超过 GEMM 部分的成本。

练习题

  1. attention_naive 扩成 multi-head:B=1, H=4, T=128, D=64。
  2. fused_qkv 用 Ch9 的 gemm_reg_tile 替换 tiled gemm,看是否更快。
  3. 测一下 T 从 64 到 4096 的耗时曲线,画出 O(T²) 增长曲线。
  4. 实现一个 attention 输出与 PyTorch 的 F.scaled_dot_product_attention 对比(Colab 上)。

11.8 工业实战:MHA/MQA/GQA、prefill vs decode、causal mask 优化

11.8.1 MHA / MQA / GQA — KV cache 的演进

变体Q headsKV heads代表模型KV 显存质量
MHA (Multi-Head Attn)n_qn_q (= Q heads)GPT-2/3/4, Llama 1最大 baseline最好
MQA (Multi-Query Attn)n_q1 (共享)PaLM, Falcon1/n_q小幅下降
GQA (Grouped-Query Attn)n_qn_kv (n_q/n_kv = group size)Llama 2/3, Mistral, Qwenn_kv/n_q接近 MHA

例:Llama 3 8B 用 n_q=32, n_kv=8 (GQA-4),KV cache 立刻缩到 MHA 的 1/4。

实现差异

// MHA: 各 head 完全独立
attention(Q[b, h, :, :], K[b, h, :, :], V[b, h, :, :]) → O[b, h, :, :]

// MQA: 所有 Q head 共享同一对 K/V
attention(Q[b, h, :, :], K[b, 0, :, :], V[b, 0, :, :]) → O[b, h, :, :]

// GQA: head 分组共享
int g = h / group_size;
attention(Q[b, h, :, :], K[b, g, :, :], V[b, g, :, :]) → O[b, h, :, :]

kernel 层面只是 K/V 索引换一下,FlashAttention 默认支持三种。

11.8.2 Prefill vs Decode — 两个完全不同的世界

LLM 推理本质上是两个独立 phase

PrefillDecode
触发用户提交 prompt 后一次之后每生成 1 个 token
序列长度T = prompt 长度 (100-32K)T = 1 (单 token)
attention shape(T, T)(1, T_total)
瓶颈compute-bound (大 GEMM + T² attn)memory-bound (小 GEMV + 全 KV 读)
Tensor Core 利用率高 (大批量)低 (M=1)
典型优化FlashAttention v2, fused QKVFlashAttention v2-decode, paged KV, 量化, spec decoding
耗时采集TODO(on GPU)TODO(on GPU)

为什么要分两个 kernel

同一个 attention kernel 在 T=2048 和 T=1 上的最优 tile 完全不同:

vLLM、TensorRT-LLM 都分别实现两个 kernel:flash_attn_varlen_func(prefill)和 flash_attn_with_kvcache(decode)。

11.8.3 Causal mask 的高效实现

Ch11 朴素实现把 -∞ 写到所有 j>i 的位置,但 FlashAttention 里这是浪费:

// ❌ 朴素:扫全 N×N, 一半位置算了又 mask 掉
if (j > i) S[i, j] = -inf;

// ✅ 优化 1: 跳过完全 mask 掉的 tile (block-level skip)
// FlashAttention 里, Q tile 在 [i_lo, i_hi), K/V tile 在 [j_lo, j_hi):
//   - 如果 j_lo > i_hi: 整个 K/V tile 都被 mask, 直接 skip 不计算
//   - 如果 j_hi <= i_lo: 完全 visible, 不用 mask
//   - 部分重叠: 逐元素 mask
if (j_lo > i_hi) continue;             // skip 这个 K/V tile
if (j_hi <= i_lo) compute_no_mask();   // 不用 mask
else compute_with_mask();              // 部分 mask

对长 context,causal mask 让大量上三角 tile 可以 skip;精确节省取决于 tile 大小和调度方式。FlashAttention v2 的 causal 实现会利用这个结构,但仍要处理负载均衡。

11.8.4 长 context 的工业方案

T 从 2K 增到 128K 的发展:

11.8.5 attention dtype 实战

张量推荐 dtype原因
Q, K, V (in/out)fp16 / bf16权重 dtype 决定
S (QK^T)fp32 累加避免 fp16 范围溢出
softmax expfp32数值稳定
P (softmax 输出)fp16压缩 shared mem
O (PV 累加)fp32累加避免精度损失
最终 O (输出)fp16 / bf16权重 dtype

FlashAttention 严格遵循此模式:fp16 I/O + fp32 累加。不要偷懒全 fp16,长 T 时 P @ V 会丢精度,模型输出乱。

11.8.6 attention 在 LLM 推理总耗时中的占比

场景attention 占比主导算子
短 prompt (T=128), batch=1TODO(on GPU)FFN GEMM 常见
长 prompt (T=8K) prefillTODO(on GPU)attention 常见
decode (任意 T)TODO(on GPU)weight load / KV 读
大 batch (B=64) decodeTODO(on GPU)attention (KV 累加)

所以:短 prompt 优化 FFN,长 prompt 优化 attention(FlashAttention),decode 优化权重加载(量化),大 batch 优化 PagedAttention。

11.9 研究前沿(2025-2026):MLA、线性 attention、SSM 混合、Reasoning

11.9.1 Multi-head Latent Attention (MLA) — DeepSeek 杀器

DeepSeek-V2/V3 引入,核心价值是把完整 K/V cache 压到 latent 表示里;显存节省比例由 latent 维度、head 数和 dtype 决定,应按模型配置计算。

// 标准 MHA: cache K, V shape (T, n_head, d_head), 共 2 × n_head × d_head
// MLA: 把 K/V 投影到低维潜空间 c_kv (T, d_c)
c_kv = X @ W_DKV    # (T, d_c)        d_c=512 远小于 n_head*d_head
K_h  = c_kv @ W_UK[h] + position      # 各 head 上投影回 d_head
V_h  = c_kv @ W_UV[h]

// Key point: cache c_kv instead of full K/V; compute bytes from model config.

工程关键:"吸收"W_UK 进 W_Q,让 Q @ K^T 直接 = Q' @ c_kv^T,不必构造完整 K

Q @ K^T = (Q @ W_UK^T) @ c_kv^T
          ─────────────
          推理时只算一次, Q' 跟着 token 走

DeepSeek 2025 开源 FlashMLA kernel,FA v2 的 MLA 变种。MLA 已成为 2025 中国大模型主流(Qwen 3 也用)。

11.9.2 Linear Attention 重生

O(N²) → O(N·D²):用 kernel function φ + 结合律:

softmax(QK^T) V          O(N²·D)
       ↓ φ(Q) @ φ(K)^T 近似 + 结合律
φ(Q) @ (φ(K)^T @ V)      O(N·D²)
              ↑
        (D,D) 大小, 与 N 无关

2024-2026 主要变体:

评价:精度仍略输 softmax attention,但 1M+ context 唯一可行

11.9.3 Mamba / SSM 混合架构

共识:纯 SSM 仍不及 Transformer,但 SSM + 少量 Attention 块在长 context 推理 KV cache 大幅缩小,是产业新方向。Llama 4 据传用类似混合。

11.9.4 Sparse Attention 新进展

11.9.5 Reasoning 模型对 attention 的新挑战

o1 / o3 / R1 / Claude reasoning:单次回答生成 10K-100K thinking token。

这是 2025-2026 推理优化(FA v3/v4、FlashMLA、MoBA、Lookahead 解码)远比训练优化更受关注的原因。

11.9.6 一图总结 2026 attention 全景

           ┌─ 标准 MHA    GPT-2/3/4 老路, 已过时
softmax    │
attention ─┼─ GQA         Llama 2/3, Mistral — 当前实用主流
 O(N²)     │
           └─ MLA         DeepSeek V2/V3 — KV cache 结构性压缩

           ┌─ Sliding W.  Mistral 7B, 长 context 推理
sparse    ─┼─ NSA / MoBA  DeepSeek / Kimi 2025
attention  │
           └─ InfLLM      检索式 KV 取舍

           ┌─ Lightning   MiniMax — 线性 + softmax 混合
linear /  ─┼─ RWKV-7      纯 RNN-like, 零 KV
SSM        │
           ├─ Mamba-2     selective SSM, Tensor Core 友好
           │
           └─ Jamba/Samba SSM+Attn 混合 (产业方向)

对工程师的影响:单一 attention kernel 不够。生产推理引擎需要 FlashInfer / xFormers / FlexAttention 这种"attention dispatch 层",根据模型自动选 backend。

常见坑

11.11 CUDA 官方手册精讲(CUDA Programming Guide 13.2(核验:2026-07-20))

Memory Traffic 分析、MHA/MQA/GQA layout、Causal Mask 优化

本节定位:把 NVIDIA 官方 CUDA Programming Guide 13.2(核验:2026-07-20) 当前版中和本章直接相关的硬核细节抽出来——概念、API、踩坑点、版本兼容性—— 让你不必通读官方手册也能掌握本章主题的"标准答案"。引用按命名章节回查,避免把旧版编号当作稳定接口。

朴素 attention 的 HBM 流量公式

11.5 节给了"T=4096, D=64, S 占 64 MiB"的结果,但没拆账。这里用 CUDA Programming Guide 13.2 Asynchronous Data Copies 的 transaction-bytes 视角分析朴素三阶段 attention:

设单头 (single-head): Q, K, V 形状均为 (T, D);fp16 = 2 字节。

阶段读 (HBM → SM)写 (SM → HBM)FLOPs算术强度 (FLOP/byte)
① QKᵀ kernel2·T·D·2 (Q+K)T²·4 (S fp32)2·T²·DD / (1 + T/D)
② softmax kernelT²·4 (S)T²·2 (P fp16)~5·T²~0.8
③ PV kernelT²·2 + T·D·2 (P+V)T·D·2 (O)2·T²·DD / (1 + T/D)
总计~6·T² 字节~6·T² 字节~4·T²·D~0.7·D

关键观察:

    graph LR
        Naive["朴素三阶段
HBM: 12·T² + 6·T·D B
算术强度 ~0.7·D
memory-bound"] FA["FlashAttention
HBM: ~8·T·D B
算术强度 ~T
compute-bound (大 T)"] Naive -.->|"省去 S, P 物化"| FA style Naive fill:#f3f1e8,stroke:#a82c30 style FA fill:#f3f1e8,stroke:#2f5d3a

这就是为什么下一章的"省 HBM"常常直接体现在端到端速度上:不是 FLOPs 省了,而是把中间矩阵从 HBM 移回 shared/register,让 roofline 落点从 memory-bound 往 compute-bound 移动。具体加速必须按序列长度、head_dim、dtype 和实现测量。

Causal mask 的 tile-level skip 收益量化

11.8.3 给了 block-level skip 代码,"causal 让一半 tile 可以 skip → 复杂度 O(T²/2)"——这只是 FLOPs 账。完整收益还有 负载均衡,分两个维度看:

tile 数 (Br=Bc=64, T=4096)朴素 (无 skip)tile-level skipelementwise mask
Q tile × K/V tile64 × 64 = 4096~2080 (~一半 + 对角线)同朴素
S 全部 compute FLOPs4·T²·D~2·T²·D4·T²·D
跨 Q tile 的 max FLOPs (调度不均)1× (均匀)1.95× (最后 Q tile 算 64 个, 第 0 Q tile 算 1 个)1× (但全 mask 项)
耗时 vs 无 maskTODO(on GPU)TODO(on GPU)TODO(on GPU)
FA v2 用 fixed work per block + persistent kernel 缓解负载不均: PG §4.12 表格指出"fixed work per block"虽然 load balance 好,但 cluster launch control (Blackwell sm_100+) 才能做 work stealing。FA v2 在 sm_80/sm_90 上用 persistent CTAs(block 数 = SM 数, 内部 grid-stride loop 抢 Q tile)缓解 causal long context 的尾部负载不均;收益要看 SM idle 和各 CTA 工作量分布。

KV cache 的 VMM 视角:为什么 PagedAttention 选 16 token / 页

vLLM 的 PagedAttention block size 是软件层设计选择;CUDA Programming Guide 13.2 的 Virtual Memory Management 章节则给出底层虚拟内存映射约束,两者不能直接等同:

// 查询当前 GPU 的最小可分配粒度
    size_t granularity = 0;
    CUmemAllocationProp prop = {};
    prop.type           = CU_MEM_ALLOCATION_TYPE_PINNED;
    prop.location.type  = CU_MEM_LOCATION_TYPE_DEVICE;
    prop.location.id    = device;
    prop.requestedHandleTypes = CU_MEM_HANDLE_TYPE_POSIX_FILE_DESCRIPTOR;

    cuMemGetAllocationGranularity(&granularity, &prop,
                                  CU_MEM_ALLOC_GRANULARITY_MINIMUM);
    // Print this value; do not assume it is identical across GPUs/drivers.
    

先用 cuMemGetAllocationGranularity 查询目标 GPU/driver 的最小粒度,再设计 KV cache 块大小。一个 KV cache 块要:

block_size每块字节 (Llama-3 8B)显存碎片 (per request)kernel indirection 开销
14 KiB~0极高
16 (vLLM 默认)64 KiB≤63 KiB
2561 MiB≤1 MiB极低
20488 MiB短请求可能浪费接近整块
为什么不用更细 (8 / 4 token):
  • PG §4.16.1.2 多次强调 "allocation must be aligned to granularity"。VMM 强制 2 MiB 对齐,block 太小 → 一次 alloc 多块,block_table 索引变长、cache 不友好。
  • attention kernel 内层每读一个 KV 都做 block_id = block_table[t / B],B 太小 → 每次 iter 都查表,把 register 都吃光。
vLLM 常见 block size 是 16 token;FlashInfer 等库允许调到 32/64 用于长 context 场景。最佳值取决于页表 indirection、内部碎片、L2 命中和 batch 长度分布。
    graph TB
        VMM["cuMemCreate
granularity = 2 MiB"] VMM --> Pool["KV pool:
N 个 2 MiB 物理块"] Pool --> Map["block_table[req][t/16]
→ 物理 block id"] Map --> Attn["attention kernel:
indirection load K, V"] style VMM fill:#f3f1e8,stroke:#a86420 style Attn fill:#f3f1e8,stroke:#2f5d3a

下一章 §12.9.4 给出 attention kernel 怎么消费 block_table

MHA vs MQA vs GQA 的 kernel 代码差异(v13.1 API 完整版)

11.8.1 给了一行伪代码,看着像三种变体要写三个 kernel;事实上生产 attention 库(FlashInfer / vLLM / FA v3)一份代码搞定,只在 K/V 加载 那一步换 indexing:

template <int N_Q_HEAD, int N_KV_HEAD, int D>
    __global__ void unified_attention_kernel(
        const __half* Q,          // (T, N_Q_HEAD, D)
        const __half* K,          // (T, N_KV_HEAD, D)  注意 N_KV_HEAD 维变
        const __half* V,          // (T, N_KV_HEAD, D)
        __half* O,
        int T, float scale)
    {
        static_assert(N_Q_HEAD % N_KV_HEAD == 0, "group_size must divide");
        constexpr int GROUP_SIZE = N_Q_HEAD / N_KV_HEAD;

        int h_q  = blockIdx.x;                        // 0..N_Q_HEAD-1
        int h_kv = h_q / GROUP_SIZE;                  // ★ 唯一区别 ★

        // ---- 1) load Q, K, V tile, swizzle 128B ----
        // Q: 永远用 h_q
        // K/V: 用 h_kv (MHA 时 h_kv == h_q, MQA 时 h_kv = 0, GQA 时 h_q/GROUP_SIZE)

        /* ... 后面 online softmax + GEMM 跟 single-head 完全一样 ... */
    }

    // 启动配置:
    // MHA: N_Q_HEAD = N_KV_HEAD = 32   → GROUP_SIZE = 1
    // MQA: N_Q_HEAD = 32, N_KV_HEAD = 1 → GROUP_SIZE = 32
    // GQA: N_Q_HEAD = 32, N_KV_HEAD = 8 → GROUP_SIZE = 4
    
变体K/V HBM 读取量 (per Q tile)多 Q tile 间 K/V 复用decode 观测重点
MHA (group=1)无 (每 head 独立)baseline
GQA (group=4)1/4×同组 4 个 Q 复用DRAM bytes/token 是否下降,L2 hit 是否上升
MQA (group=32)1/32×所有 Q 共享KV 读字节下降,但质量和模型结构约束更强
为什么生产倾向 GQA 而非 MQA: MQA 的 KV head=1 会显著压缩表示能力;GQA 在减少 KV cache 和保持质量之间更稳。 Llama 3 / Mistral / Qwen 等现代模型大量采用 GQA。具体 group size 是模型训练时确定的,推理侧不能随意改。

kernel 层面统一写法的隐藏好处:用 PG §4.11 的 TMA descriptor 共享(一个 CUtensorMap 描述整个 K tensor),h_kv 仅参与 box_dim 起始坐标计算,硬件 cache 命中率不变——这就是为什么 FlashInfer 11.9.6 提到的"backend 自动 dispatch"可以做到 zero overhead。

下一章导览

第 12 章避免物化完整 attention score 矩阵,用 online softmax 与 tile 复用降低 HBM 流量。