Skip to content

注意力机制

本页速览 注意力是 Transformer 与一切现代大模型的心脏。本文讲清"软检索"思想与 QKV 缩放点积注意力、多头动机、三类位置编码、自/交叉注意力与因果掩码,并覆盖 O(n²) 复杂度、FlashAttention 与 KV cache 等工程要点。

注意力机制

一句话定义:注意力(attention)是一种"按相关性加权聚合信息"的机制——给定一个查询,从一组候选中软性地挑出最相关的部分来读取。它最初由 Bahdanau 等在 2014 年用于机器翻译(对齐源词),2017 年"Attention Is All You Need"(Vaswani 等)将其推到聚光灯下,成为 Transformer 架构与今天所有大模型的基石。它与传统 RNN 的对比见RNN 与序列建模

一、注意力思想:软检索

把注意力想象成一个软性数据库查询:你有一个"问题"(query),一堆"待查条目"(每个条目有 key 和 value)。传统硬检索是"找到唯一命中的条目,取出它的内容";软检索是"给每个条目按匹配度打分,加权求和所有条目":

Attention = Σᵢ (匹配度(query, keyᵢ) 归一化) × valueᵢ

三条直觉值得记住:

  1. 匹配度用点积/相似度衡量,归一化成权重(和为 1),即"软性地分配注意力预算"。
  2. 输出是 value 的加权和——永远可微,梯度能顺畅流过(见反向传播与自动微分)。
  3. 它不假设顺序依赖:任何 query 可以直接"看"到任何位置的 key,因此天然解决长距离依赖——这正是 RNN 无法并行、难以捕捉长程依赖的痛点(见RNN 与序列建模)。

二、QKV 与缩放点积注意力

每个输入 token 通过三个不同的线性投影得到三组向量:

Q = X·W_Q      K = X·W_K      V = X·W_V

缩放点积注意力(scaled dot-product attention)公式:

Attention(Q, K, V) = softmax(Q·Kᵀ / √d_k) · V

逐项解释:

  • Q·Kᵀ:每对 token 之间的相似度得分,形状 (n, n)n 是序列长度。
  • /√d_k 缩放d_k 是 key 的维度。为什么除?因为点积的方差随 d_k 线性增长,不缩放的 softmax 会进入饱和区,梯度趋近 0。缩放让 logits 方差保持在 O(1)——这和我们讲初始化与归一化时反复强调的"方差守恒"是同一个道理。
  • softmax 沿行归一化:每个 query 的注意力权重和为 1。
  • 最后乘以 V,得到每个位置"按注意力加权融合"后的表示。

PyTorch 手写实现(教学版,不追求最优):

python
import torch, torch.nn.functional as F

def scaled_dot_product_attention(q, k, v, mask=None):
    # q,k,v: (batch, heads, seq, d_k)
    scores = q @ k.transpose(-2, -1) / (k.size(-1) ** 0.5)
    if mask is not None:
        scores = scores.masked_fill(mask == 0, float('-inf'))
    weights = F.softmax(scores, dim=-1)
    return weights @ v

三、多头注意力的动机

单个注意力只算一种"关系"。**多头注意力(MHA)**把 Q/K/V 各自切成 h 份,并行算 h 组注意力,最后拼接再投影:

MultiHead(Q,K,V) = Concat(head₁,…,head_h)·W_O
headᵢ = Attention(Q·W_Qⁱ, K·W_Kⁱ, V·W_Vⁱ)

动机有三层:

  1. 不同的头学到不同的关系模式——有的头关注句法依赖,有的关注指代,有的关注位置邻近(已被大量可视化研究证实,见可解释性与公平性)。
  2. 多组低维子空间比一组高维空间表达力更强——总计算量基本不变(切分后每头维度变小),但参数在多个表示子空间里并行工作。
  3. 提升优化稳定性:多头平均了单头的噪声。

四、位置编码:绝对、相对、RoPE

注意力是"置换等变"的——它对输入顺序不敏感(Q·Kᵀ 打乱顺序结果不变)。要让模型"知道顺序",必须注入位置信息。三类主流方案:

方案做法特点代表
绝对位置编码每个位置一个向量,加到 token 上简单直接;但外推到更长序列能力弱正弦-余弦(原版 Transformer)、可学习位置编码(GPT 系列)
相对位置编码编码"两个位置的相对偏移",在注意力打分时注入显式建模相对距离,外推更好Shaw 2018、Transformer-XL
RoPE(旋转位置编码)用旋转矩阵把位置信息编进 Q/K 的夹角,内积自然含相对位置兼具绝对编码的实现与相对编码的性质,长度外推能力强LLaMA、Qwen、DeepSeek 等当代 LLM

为什么当代大模型几乎都转投 RoPE?核心诉求是长度外推:预训练时序列长 4k,推理时希望支持 8k/32k。RoPE 的位置信息以"相位"形式进入点积,天然只依赖相对位移,配合插值技巧(NTK-aware、YaRN)可以低成本扩展到更长上下文。详见大语言模型(LLM)

五、自注意力 vs 交叉注意力

  • 自注意力(self-attention):Q、K、V 来自同一个序列——每个 token 与同序列其他 token 交互。作用:捕获序列内部的长程依赖。Transformer encoder/decoder 内部都用它。
  • 交叉注意力(cross-attention):Q 来自一个序列(如解码器的当前状态),K、V 来自另一个序列(如编码器的输出)。作用:让解码器"读取"编码器信息——机器翻译、图文多模态对齐(多模态模型)都靠它。

本质区别:自注意力建模"内部关系",交叉注意力建模"两个序列之间的对齐关系"

六、因果掩码

自回归生成(一次预测一个 token)要求模型只能看到"当前位置及之前"的 token,否则就作弊了(未来的 token 泄入预测)。实现上很简单:在 Q·Kᵀ 得分矩阵的上三角填 -inf,softmax 后权重自然为 0:

mask 下三角为 True,上三角为 False
scores = scores.masked_fill(~mask, float('-inf'))

这就是因果掩码(causal mask),也是自回归模型的"注意力只能朝左看"的正式含义。它与生成式模型里自回归生成、"下一个 token 预测"的训练目标天然配套。

七、计算复杂度 O(n²) 与 FlashAttention

自注意力的得分矩阵是 (n, n),时间和显存都随序列长度平方级增长:

复杂度 O(n²·d) —— n=8k 时就有 6.4×10⁷ 个得分

这是 Transformer 的长序列之痛。缓解手段分两路:

  1. 稀疏/线性注意力:让每个 query 只跟部分 key 交互(局部窗口、全局 token、滑动窗口),如 Longformer、BigBird、Linformer;复杂度降到 O(n) 或 O(n log n)。
  2. FlashAttention(Dao 等,2022):不改变注意力数学、只改 IO——把注意力分块计算,中间结果不落显存(tiling + online softmax),让慢速显存与快速 SRAM 之间的搬运最小化。结果:速度 2–4 倍、显存从 O(n²) 降到 O(n),并且数值等价。它已成为 A100/H100 上训练一切大模型的默认实现,Triton 教程里甚至用几十行就能复现。

工程选择口诀:序列短(<1k)用标准注意力,长序列先上 FlashAttention,还不够再上稀疏注意力

八、KV cache:推理加速

生成式推理时每步只新增一个 token,但注意力需要"看历史所有 token"的 K、V。如果每步重算全部历史,复杂度随步数二次方,慢得不可用。

KV cache 的做法:把已生成 token 的 K、V 缓存在显存里,每步只算新 token 的 Q/K/V,与缓存拼接再算注意力:

K_cache = concat(K_cache, K_new)     # 每步追加
scores = Q_new · K_cacheᵀ / √d_k

效果:生成延迟从 O(t²) 降到 O(t)(t 是已生成步数)。代价是显存:KV cache 大小 = 层数 × 头数 × 序列长 × 维度 × 2 × 精度 × batch。这也是为什么大模型有"max context 显存不足"、以及 GQA/MQA(多头共享 K/V)这类显存优化——见大语言模型(LLM)

九、注意力可视化与可解释性

注意力权重可以当热力图直接可视化:某个 query token 对哪些 key token 给了高权重。经典发现(如"it"的指代消解、翻译对齐)看起来很有解释力,但要谨慎

  • 注意力高 ≠ 因果上重要。有研究(Jain & Wallace 2019)发现删掉高注意力头,预测常不变——注意力权重存在冗余。
  • 注意力是多层多头的复杂组合,单层热力图只能说明"局部相关性"。

更可靠的做法是把注意力机制放进机制可解释性(probes、circuits)框架里研究,见可解释性与公平性

十、权衡与取舍

权衡与取舍

表达力 vs 计算复杂度:全注意力表达力最强但 O(n²);稀疏注意力线性复杂度但牺牲部分长程建模(需局部窗口+全局 token 补足)。长序列优先省算力,短序列别折腾。

KV cache 的延迟 vs 显存:缓存让生成快一个数量级,但显存随上下文线性吃紧;GQA/MQA 用共享 KV 换显存,牺牲少量质量换吞吐。

RoPE 的外推 vs 训练内插:RoPE 外推能力最好但并非无限;超长上下文还要配合 NTK/YaRN 插值,插值会引入少量精度损失。

绝对编码的简单 vs 相对/RoPE 的通用:简单场景(定长序列)绝对编码够用;要长度外推必须上相对系。当代大模型选 RoPE 不是没有代价——它的实现与融合算子复杂度更高。

注意力机制把"相关性"变成了第一公民,它既是神经网络基础里"层"的一种,也是表征学习与预训练里"上下文表征"的核心来源。理解了注意力,Transformer 与 LLM 的骨架就清晰了,下一站是Transformer 架构

延伸阅读

参考资料