Attention 为当前 query 位置计算各 key 位置的权重,再混合它们的 value。Causal mask 限定可以读取哪些位置。它限制信息依赖,不限制所有位置同时计算。
接着读:Block 如何接入 Attention。
一、核心数学
Section titled “一、核心数学”本例 X=[1,3,8]。投影并拆头后 Q/K/V 各 [1,2,3,4],scores 和 A 为 [1,2,3,3],Z 为 [1,2,3,4]。拼回 [1,3,8] 后再做输出投影 WO。Softmax 沿 key 轴;每个 query/head 的可见 key 权重和为 1。
二、Q, K, V 的几何直觉
Section titled “二、Q, K, V 的几何直觉”Q 表示此位置用于查询的特征;K 表示被匹配的特征;V 表示最终被加权混合的内容。三者是从表示经过不同可训练投影得到的,不是 token 的三种自然属性。
本例 d=4,所以除以 √4=2。缩放用于控制点积量级,避免 softmax 过于饱和。这个方差解释依赖输入分量的简化统计假设,不是普遍的精确方差结论。
因果遮罩与一个权重例子
Section titled “因果遮罩与一个权重例子”允许读取(行=query,列=key) 小猫 坐在 窗边小猫 1 0 0坐在 1 1 0窗边 1 1 1加法 mask 的允许项为 0,禁止项为 −∞。位置“坐在”的示意 scores [0,0,9] 加 mask 后为 [0,0,−∞],softmax 得 [0.5,0.5,0]。即使“窗边”的原始分数最大,它也不能进入该位置输出。
训练输入虽然已经含“窗边”,位置“小猫”预测“坐在”时也不能读取它。标签只用于计算 loss,不进入当前位置的 Attention。
多头与两种 softmax
Section titled “多头与两种 softmax”H=2 个头各在 4 维空间进行查询,结果拼接回 8 维。多头让模型具有不同投影和权重分布,不保证某个头负责固定语义。注意力 softmax 是“对 key 位置分配权重”;LM Head 后的 softmax 是“对词表分配下一 token 概率”。二者的轴和用途不同。
来源:Attention Is All You Need、3Blue1Brown Attention。
实现细节:模型配置与 Attention 后端
Section titled “实现细节:模型配置与 Attention 后端”下面介绍 GPT-2、Llama、Qwen 的配置差异,以及 Attention 的几种后端、online softmax 和源码位置。backend 可用性、返回值和代码行号来自特定版本,需要对照当前安装的版本;表格不是性能保证。
五、GPT-2 vs Llama 的 attention 差异
Section titled “五、GPT-2 vs Llama 的 attention 差异”| GPT-2 | Llama / Qwen 现代 | |
|---|---|---|
| Q/K/V 投影 | 1 个 c_attn Linear(768→2304),split | 3 个独立 Linear(q_proj, k_proj, v_proj) |
| Q/K/V 头数 | 同(MHA) | Q 多,K/V 少(GQA) |
| 位置编码 | 不在 attention 里(在 input embedding) | RoPE,旋转 Q, K 向量 |
repeat_kv | 不用(头数相同) | 视实现而定:可显式扩展或由 kernel 直接处理 |
| 输出投影 | c_proj (Conv1D) | o_proj (Linear) |
六、GQA(Grouped Query Attention,Llama / Qwen 用)
Section titled “六、GQA(Grouped Query Attention,Llama / Qwen 用)”GQA 让 K, V 头数比 Q 少,省 KV cache 显存:
GQA 真实例子(常见模型):
Qwen2.5-0.5B-Instruct: Q 14 头, KV 2 头 → 7:1 比例 Llama-3-8B: Q 32 头, KV 8 头 → 4:1 比例 Llama-3-70B: Q 64 头, KV 8 头 → 8:1 比例 Mistral-7B-v0.1: Q 32 头, KV 8 头 → 4:1 比例以 Qwen 0.5B 为例:
Q 14 头(每头 64 维)K 2 头(每头 64 维) ★ 头数少!V 2 头(每头 64 维)
每 7 个 Q head 共享 1 组 KV head Q[0..6] ←→ K[0], V[0] Q[7..13] ←→ K[1], V[1]好处:KV cache 显存减少 7 倍(KV 头少)— 比例随模型不同。
源码实例:repeat_kv 将 K/V 逻辑扩展到 Q 的头数;支持 GQA 的 kernel 也可以直接处理,不必实际复制完整 K/V。
七、Attention Dispatch(backend 切换)
Section titled “七、Attention Dispatch(backend 切换)”GPT2Attention.forward / LlamaAttention.forward 里:
attention_interface = ALL_ATTENTION_FUNCTIONS.get_interface( self.config._attn_implementation, # "sdpa" / "eager" / "flash_attention_2" eager_attention_forward # fallback)attn_output, attn_weights = attention_interface(self, Q, K, V, mask, ...)→ 同一份模型代码,通过字符串选 backend。
4 种 attention 实现
Section titled “4 种 attention 实现”| Backend | 实现 | 速度 | 显存 | 平台 |
|---|---|---|---|---|
| eager | 纯 PyTorch 5 步公式 | 慢 | O(N²) | 全平台,可调试 |
| sdpa | PyTorch 内置融合 kernel | 中-快 | O(N²)~O(N) | 全平台(Mac MPS 默认) |
| flash_attention_2 | Tri Dao 的 CUDA kernel | 最快 | O(N)(不存全 QK^T) | 以所用实现和硬件为准 |
| flex_attention | PyTorch 2.5+,可编程 mask | 快 | 优化 | CUDA + 新 PyTorch |
八、F.scaled_dot_product_attention(SDPA)详解
Section titled “八、F.scaled_dot_product_attention(SDPA)详解”PyTorch 2.0+ 内置的优化 attention,自动选 3 种 backend:
F.scaled_dot_product_attention(Q, K, V, ...) │ ├── flash_attention (CUDA, 分块,online softmax,显存 O(N)) ├── memory_efficient (CUDA,Flash 的宽松版) └── math (C++ 朴素实现,所有平台 fallback)Mac 上 SDPA 通常走 Math backend(MPS 上 Flash/MemEff 不可用,有时也走 MPS 专用 fast path)— 等价于”优化版 eager”,具体走哪条随 PyTorch 版本和 dtype 变化,可以用 torch.backends.cuda.sdp_kernel(...) context manager 强制指定。
注意:transformers 的 sdpa_attention_forward 不返回 attention weights(永远 None)。原因是 SDPA 的 API/backend 路径不暴露权重——即使内部走 math backend 实际上算了 weights,但 wrapper 不取出来传出去(因为 Flash 走的话权重根本没材料化,统一起见 sdpa 路径就不返回)。
→ 想看 attention 矩阵(可视化、debug)必须用 attn_implementation="eager"。
九、Flash Attention 核心 trick(online softmax)
Section titled “九、Flash Attention 核心 trick(online softmax)”问题:朴素 softmax 需要先看完所有 logits 才能归一化 → 必须 materialize 完整 N×N 矩阵 → O(N²) 显存。
Flash 解法:边算边更新 running statistics(running max、running denominator、rescaled output),不存全矩阵。
简化伪代码(关键是合并块时重新 rescale,公式比下面写的更精细):
维护 running 状态:m(running max), l(running sum), O(running output)
看到一块 K 块 (with scores S_k 和 values V_k) 时: m_new = max(m_old, max(S_k)) # 更新 max # 用新 max 把"老 O"和"新块"都 rescale 到统一基准: l_new = exp(m_old - m_new) * l_old + sum_j exp(S_k[j] - m_new) O_new = (l_old / l_new) * exp(m_old - m_new) * O_old + sum_j (exp(S_k[j] - m_new) / l_new) * V_k[j]
m_old, l_old, O_old = m_new, l_new, O_new
→ 数学上等价于全 N 个一起 softmax(QK^T)V→ 显存 O(N)(只存 running 状态,不存 N×N 矩阵),速度收益依工作负载而定(GPU SRAM 友好,减少 HBM 读写)。完整公式见 FlashAttention-2 论文 (Dao 2023)。
十、KV cache 在 attention 里怎么用
Section titled “十、KV cache 在 attention 里怎么用”# GPT2Attention.forward 关键 3 行if past_key_values is not None: key_states, value_states = past_key_values.update( key_states, value_states, self.layer_idx, ) # ↑ 把新 K, V append 到 cache,返回完整历史 K, V详见 KV Cache。
Decoding 阶段: Q (新 token) [B, 12, 1, 64] K (完整历史) [B, 12, T_total, 64] V (完整历史) [B, 12, T_total, 64]
attention 算 [B, 12, 1, T_total] ← 1 个 query 对所有历史十一、输入输出张量形状
Section titled “十一、输入输出张量形状”完整流程:
hidden_states [B, T, 768] ↓ c_attn (Linear 768→2304)qkv [B, T, 2304] ↓ split(768) × 3Q, K, V (各) [B, T, 768] ↓ view(B, T, 12, 64).transpose(1, 2)Q, K, V (各) [B, 12, T, 64] ↓ [KV cache update,GQA 用 repeat_kv]Q, K, V (适配) [B, 12, T_q, 64] / [B, 12, T_k, 64] ↓ scores = Q @ K^T * scalescores [B, 12, T_q, T_k] ↓ + causal mask, softmaxattn_weights [B, 12, T_q, T_k] ↓ @ Vout [B, 12, T_q, 64] ↓ transpose(1, 2) (eager 内部已做)out [B, T_q, 12, 64] ↓ reshapeout [B, T_q, 768] ↓ c_proj (Linear 768→768)attn_output [B, T_q, 768]十二、attn_output 和 attn_weights 后续用途
Section titled “十二、attn_output 和 attn_weights 后续用途”| 张量 | 形状 | 后续用途 |
|---|---|---|
attn_output | [B, T, 768] | 加到 residual,进入 MLP 子层。深度参与后续计算 |
attn_weights | [B, 12, T_q, T_k] | 通常丢弃(用 _ 接住)。只在 output_attentions=True 时保存,用于可视化 / debug |
注意:SDPA 返回 attn_weights = None(Flash 没保存)。
十三、长 context 的显存瓶颈
Section titled “十三、长 context 的显存瓶颈”attn_weights shape = [B, 12, T_q, T_k] ∝ T²
seq=1000 时: 12M 数 ≈ 24 MB (bf16),一层seq=8000 时: 768M 数 ≈ 1.5 GB,一层!seq=8000, 32 层: 50 GB!→ Flash Attention 不存这个矩阵,所以能跑 long context。
十四、可视化每个 head 的关注模式(用 eager + output_attentions)
Section titled “十四、可视化每个 head 的关注模式(用 eager + output_attentions)”mdl = AutoModelForCausalLM.from_pretrained("openai-community/gpt2", attn_implementation="eager")
with torch.no_grad(): out = mdl(input_ids, output_attentions=True)
# out.attentions 是 tuple,12 层每层一个# 每个 shape: [B, 12, T, T]
# 看第 0 层 head 0 的 attention 矩阵attn_w = out.attentions[0][0, 0] # [T, T]print(attn_w)# 每行和=1,上三角=0观察具体 head 的关注模式;头的语义分工不是训练目标保证的性质。
十五、相关源码位置
Section titled “十五、相关源码位置”src/transformers/models/gpt2/modeling_gpt2.py:75 GPT2Attentionsrc/transformers/models/llama/modeling_llama.py:225 LlamaAttentionsrc/transformers/integrations/sdpa_attention.py sdpa_attention_forwardsrc/transformers/integrations/flash_attention.py flash_attention_forwardsrc/transformers/integrations/flex_attention.py flex_attention_forward