Causal Self-Attention

Attention 精讲

用图把 mask、多头、GQA、KV Cache 钉死。公式用等宽文本,避免 Markdown 下划线坑。

完整推导亦见 notes/01b_Attention精读.md · 跑 code/01b_*.py

1. 一条链

x → Wq/Wk/Wv → Q,K,V
  → (RoPE)
  → S = (Q @ Kᵀ) / √d
  → causal mask(未来 = -∞)
  → A = softmax(S)      # 每行和为 1
  → Y = A @ V
  → concat heads → Wo

一句话:用 query–key 相似度当权重,对 value 做加权平均;因果约束让位置 t 看不到未来。

2. 因果 Mask 长什么样

因果 mask 矩阵 允许看(保留分数) 屏蔽(置 -∞) t=0 t=1 t=2 t=3 j=0 1 2 3 下三角(含对角)= 可见 上三角 = 未来,softmax 前抹掉 PyTorch: triu(..., diagonal=1) masked_fill(-inf)
图 1 · 因果 mask:行是 query 位置 t,列是 key 位置 j
Mask 必须加在 softmax 之前。加在之后会破坏概率归一,且 -inf 技巧就是为 softmax 准备的。

3. 多头与 GQA

MHA vs GQA MHA nh 个 Q 头 nh 个 K 头 nh 个 V 头 KV Cache ∝ nh GQA nh 个 Q 头 nkv 个 K/V 头 计算时 repeat 对齐 KV Cache ∝ nkv 体积比 = nkv/nh
图 2 · GQA 用更少 KV 头省 cache;Llama2-70B 常见 nh=64, nkv=8 → 约 1/8

4. KV Cache 字节(一层、fp16)

bytes_MHA = 2 * B * T * nh * d * 2
          = 4 * B * T * C          # 因 nh*d = C

bytes_GQA = 4 * B * T * C * (nkv / nh)

例:B=1, T=1024, C=4096 → MHA 一层约 16 MiB;GQA nkv=nh/82 MiB

GPT-3 量级 B=16, T=4096, C=12288 → 一层 MHA 约 3 GiB。8MB SRAM 塞不下 1 层——这就是飞书题要的数量级直觉。

5. Prefill vs Decode

阶段Attention 算力直觉常像
Prefill 长度 T∝ T² · C算力 / IO 双吃
Decode 一步(cache 长 L)∝ L · C读 KV → 带宽瓶颈

这和 NPU/编译器主线直接相关:Prefill 要 tiling / 融合;Decode 要减 KV、布局、量化。

6. 动手

  1. python code/01b_attention_numeric_walkthrough.py — 看 T=3 mask
  2. 01_llama2_decoder.pyn_kv_heads,对照上面公式
  3. 默画图 1 的下三角,口述「为何 mask 在 softmax 前」

← GPT 结构 · 发布与变现 →