Causal Self-Attention
Attention 精讲
用图把 mask、多头、GQA、KV Cache 钉死。公式用等宽文本,避免 Markdown 下划线坑。
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 必须加在 softmax 之前。加在之后会破坏概率归一,且 -inf 技巧就是为 softmax 准备的。
3. 多头与 GQA
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/8 → 2 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. 动手
python code/01b_attention_numeric_walkthrough.py— 看 T=3 mask- 改
01_llama2_decoder.py的n_kv_heads,对照上面公式 - 默画图 1 的下三角,口述「为何 mask 在 softmax 前」