跳转到内容

基础原理第 6 篇,共 13 篇

Attention 与因果遮罩

用 QKV、因果遮罩和多头形状解释当前 token 如何读取上下文

更新于 阅读约 5 分钟

Attention 为当前 query 位置计算各 key 位置的权重,再混合它们的 value。Causal mask 限定可以读取哪些位置。它限制信息依赖,不限制所有位置同时计算。

接着读:Block 如何接入 Attention。

Q=XWQ,K=XWK,V=XWVQ=XW_Q,\quad K=XW_K,\quad V=XW_V A=softmax⁡key(QK⊤/d+M),Z=AVA=\operatorname{softmax}_{key}(QK^\top/\sqrt d+M),\quad Z=AV

本例 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 表示最终被加权混合的内容。三者是从表示经过不同可训练投影得到的,不是 token 的三种自然属性。

本例 d=4,所以除以 √4=2。缩放用于控制点积量级,避免 softmax 过于饱和。这个方差解释依赖输入分量的简化统计假设,不是普遍的精确方差结论。

允许读取(行=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。

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-2Llama / Qwen 现代
Q/K/V 投影1 个 c_attn Linear(768→2304),split3 个独立 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。


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。

Backend实现速度显存平台
eager纯 PyTorch 5 步公式慢O(N²)全平台,可调试
sdpaPyTorch 内置融合 kernel中-快O(N²)~O(N)全平台(Mac MPS 默认)
flash_attention_2Tri Dao 的 CUDA kernel最快O(N)(不存全 QK^T)以所用实现和硬件为准
flex_attentionPyTorch 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)。


# 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 对所有历史

完整流程:

hidden_states [B, T, 768]
↓ c_attn (Linear 768→2304)
qkv [B, T, 2304]
↓ split(768) × 3
Q, 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 * scale
scores [B, 12, T_q, T_k]
↓ + causal mask, softmax
attn_weights [B, 12, T_q, T_k]
↓ @ V
out [B, 12, T_q, 64]
↓ transpose(1, 2) (eager 内部已做)
out [B, T_q, 12, 64]
↓ reshape
out [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 没保存)。


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 的关注模式;头的语义分工不是训练目标保证的性质。


src/transformers/models/gpt2/modeling_gpt2.py:75 GPT2Attention
src/transformers/models/llama/modeling_llama.py:225 LlamaAttention
src/transformers/integrations/sdpa_attention.py sdpa_attention_forward
src/transformers/integrations/flash_attention.py flash_attention_forward
src/transformers/integrations/flex_attention.py flex_attention_forward

完整张量示例 | KV Cache →