Skip to content

Attention 机制详解

Attention(注意力)机制是现代深度学习最重要的架构创新之一,它是 Transformer 的核心组件,也是所有现代大语言模型(LLM)的基础。本页从物理直觉出发,逐步推导 Scaled Dot-Product Attention、Multi-Head Attention,以及现代 LLM 中广泛使用的 RoPE(Rotary Position Embedding)旋转位置编码。

什么是 Attention? 简单说,Attention 是一种”选择性聚焦”机制——就像你读一篇文章时,眼睛不会均匀地看每个字,而是根据当前关心的内容,把注意力集中在某些关键句上。Attention 让神经网络也能这样做:面对一长串输入,模型自动学会”该多关注哪个位置的信息”。

Attention 的数学结构可以从一个日常场景理解:你在搜索引擎里搜索东西。

  • 你在搜索框里输入的查询词(Query),代表”我想要什么”
  • 数据库中每个网页有一个关键词标签(Key),代表”这个网页讲什么”
  • 网页的实际内容(Value)才是你最终想看的

搜索引擎的做法是:计算你的 Query 与每个网页的 Key 之间的相似度,相似度越高的网页排名越靠前,然后把排名靠前的 Value 内容汇总返回给你。Attention 做的完全一样:

三个矩阵的物理含义:

名称含义类比
Query(Q)当前位置”想找什么类型的信息”搜索框里输入的关键词
Key(K)每个位置”能提供什么类型的信息”网页的标题和元数据
Value(V)每个位置的”实际内容”网页正文
Attention WeightQ 和 K 的匹配程度搜索结果的相关性分数
Output按 Weight 对 V 加权求和搜索结果的摘要(最相关的内容被放大)

关键洞察:Q、K、V 都是从同一个输入向量经过不同的线性变换得到的。这意味着”我想找什么”、“我能提供什么”和”我的实际内容”是同一信息的三个不同视角——就像同一段文字,可以分别提取出”问题视角”、“标签视角”和”内容视角”。

给定输入序列 X∈Rn×dX \in \mathbb{R}^{n \times d}(nn 是序列长度,dd 是模型维度),通过三个可学习的权重矩阵 WQW_Q、WKW_K、WVW_V 进行线性变换:

Q=XWQ,K=XWK,V=XWVQ = X W_Q, \quad K = X W_K, \quad V = X W_V

其中 WQ∈Rd×dkW_Q \in \mathbb{R}^{d \times d_k},WK∈Rd×dkW_K \in \mathbb{R}^{d \times d_k},WV∈Rd×dvW_V \in \mathbb{R}^{d \times d_v}。通常 dk=dv=d/hd_k = d_v = d / h(hh 是头数),在单头注意力中 dk=dv=dd_k = d_v = d。

为什么要做线性变换,而不直接用原始输入? 因为”找什么”、“提供什么线索”、“内容是什么”三个视角需要不同的表示。线性变换让模型自己学习如何从同一输入中提取出这三个不同的视角——就像同一个人在面试时,会根据面试官的提问展现不同侧面的自己。

对于 Query 向量 qq 和一组 Key 向量 {k1,k2,…,kn}\{k_1, k_2, \ldots, k_n\},注意力分数衡量 qq 与每个 kik_i 的匹配程度。最自然的度量是点积(Dot Product)——两个向量方向越接近,点积越大:

scorei=q⋅ki=∑j=1dkqj⋅ki,j\text{score}_i = q \cdot k_i = \sum_{j=1}^{d_k} q_j \cdot k_{i,j}

点积相似度之所以有效,可以从几何角度理解:两个向量的点积等于 ∥q∥∥ki∥cos⁡θ\|q\| \|k_i\| \cos\theta,其中 θ\theta 是夹角。方向越接近(cos⁡θ→1\cos\theta \to 1),点积越大,表示”越匹配”。

原始点积存在一个问题:当 dkd_k(Key 维度)很大时,点积的值会非常大,导致 Softmax 函数进入梯度极小的饱和区(尾部),训练停滞。解决方案是除以一个缩放因子 dk\sqrt{d_k}:

Attention(Q,K,V)=Softmax ⁣(QKTdk)V\text{Attention}(Q, K, V) = \text{Softmax}\!\left(\frac{Q K^T}{\sqrt{d_k}}\right) V

为什么除以 dk\sqrt{d_k}? 假设 QQ 和 KK 的每个元素是均值为 0、方差为 1 的独立随机变量,则点积 q⋅k=∑i=1dkqikiq \cdot k = \sum_{i=1}^{d_k} q_i k_i 的方差为 dkd_k(独立变量乘积之和的方差等于方差之和)。当 dk=512d_k = 512 时,点积的标准差约为 512≈22.6\sqrt{512} \approx 22.6,这个量级的输入会让 Softmax 输出接近 one-hot(某个位置接近 1,其余接近 0),梯度几乎为零。除以 dk\sqrt{d_k} 后方差变为 1,保持稳定的梯度流。

完整的 Scaled Dot-Product Attention 流程:

展开成矩阵形式,注意力权重矩阵 AA 的每个元素为:

Aij=exp⁡ ⁣(qi⋅kjdk)∑l=1nexp⁡ ⁣(qi⋅kldk)A_{ij} = \frac{\exp\!\left(\frac{q_i \cdot k_j}{\sqrt{d_k}}\right)}{\sum_{l=1}^{n} \exp\!\left(\frac{q_i \cdot k_l}{\sqrt{d_k}}\right)}

其中 AijA_{ij} 表示位置 ii 对位置 jj 的注意力权重。最终输出 O=AVO = AV,即每个位置的输出是所有位置的 Value 按注意力权重加权求和。

Softmax 函数将任意实数向量转换为概率分布(非负且和为 1):

Softmax(zi)=ezi∑jezj\text{Softmax}(z_i) = \frac{e^{z_i}}{\sum_{j} e^{z_j}}

在 Attention 中,Softmax 的作用是:把原始的相似度分数变成”概率化”的注意力权重——所有位置的权重之和为 1,权重越大的位置贡献越多信息。这也意味着每个位置都会”看到”所有其他位置(只是程度不同),这就是 Attention 全局建模能力的来源。

第四步:Multi-Head Attention(多头注意力)

Section titled “第四步:Multi-Head Attention(多头注意力)”

单头注意力的问题:一组 Q/K/V 只能学到一种”关注模式”。但语言理解需要同时关注多个层面——比如翻译 “The animal didn’t cross the street because it was tired” 时,“it”既要关注到 “animal”(指代消解),也要关注到 “tired”(状态信息)。单头难以同时捕捉多种模式。

Multi-Head Attention 的解法:把 Q/K/V 分成 hh 组(“头”),每组独立做 Attention,最后拼接:

MultiHead(Q,K,V)=Concat(head1,…,headh)WO\text{MultiHead}(Q, K, V) = \text{Concat}(\text{head}_1, \ldots, \text{head}_h) W^O headi=Attention(QWiQ,KWiK,VWiV)\text{head}_i = \text{Attention}(Q W_i^Q, K W_i^K, V W_i^V)

每个头使用独立的投影矩阵 WiQ∈Rd×dkW_i^Q \in \mathbb{R}^{d \times d_k}、WiK∈Rd×dkW_i^K \in \mathbb{R}^{d \times d_k}、WiV∈Rd×dvW_i^V \in \mathbb{R}^{d \times d_v},让模型从不同子空间学习不同的关注模式。

维度分配举例:以 BERT-base 为例,d=768d = 768,h=12h = 12,则每个头的维度 dk=dv=768/12=64d_k = d_v = 768 / 12 = 64。12 个头的输出拼接后恢复到 768 维,再经过线性变换 WOW^O 输出。总参数量与单头注意力(dk=768d_k = 768)相同——多头不是”更多参数”,而是”更多视角”。

在自回归生成(GPT 类模型)中,位置 ii 只能看到位置 ≤i\leq i 的信息,不能”偷看”未来。实现方式是在 Softmax 之前,把未来位置的分数设为 −∞-\infty:

scoreij={qi⋅kjdk,j≤i−∞,j>i\text{score}_{ij} = \begin{cases} \frac{q_i \cdot k_j}{\sqrt{d_k}}, & j \leq i \\ -\infty, & j > i \end{cases}

经过 Softmax 后,e−∞=0e^{-\infty} = 0,未来位置的权重恰好为零。这个上三角掩码(causal mask)让 Attention 矩阵变成下三角形式,是 GPT/LLaMA 等生成模型的核心设计。

下面是完整的 Multi-Head Attention 实现,包含 Scaled Dot-Product Attention 和 Causal Mask 支持:

import torch
import torch.nn as nn
import torch.nn.functional as F
import math
class MultiHeadAttention(nn.Module):
def __init__(self, d_model: int, num_heads: int, dropout: float = 0.0):
super().__init__()
assert d_model % num_heads == 0, "d_model must be divisible by num_heads"
self.d_model = d_model
self.num_heads = num_heads
self.d_k = d_model // num_heads # 每个头的维度
# Q/K/V 投影矩阵(合并成一个矩阵提高效率)
self.W_qkv = nn.Linear(d_model, 3 * d_model, bias=False)
# 输出投影
self.W_o = nn.Linear(d_model, d_model, bias=False)
self.dropout = nn.Dropout(dropout)
def forward(self, x: torch.Tensor, causal: bool = False) -> torch.Tensor:
"""
Args:
x: (batch, seq_len, d_model)
causal: 是否使用因果掩码(自回归生成时为 True)
Returns:
(batch, seq_len, d_model)
"""
B, N, D = x.shape
# 一次性计算 Q/K/V,然后切分
qkv = self.W_qkv(x) # (B, N, 3*D)
qkv = qkv.reshape(B, N, 3, self.num_heads, self.d_k)
qkv = qkv.permute(2, 0, 3, 1, 4) # (3, B, H, N, d_k)
q, k, v = qkv[0], qkv[1], qkv[2] # 各为 (B, H, N, d_k)
# Scaled Dot-Product Attention
scores = torch.matmul(q, k.transpose(-2, -1)) / math.sqrt(self.d_k)
# scores: (B, H, N, N)
if causal:
# 上三角掩码(不包括对角线),阻止"看到未来"
mask = torch.triu(
torch.ones(N, N, device=x.device, dtype=torch.bool), diagonal=1
)
scores = scores.masked_fill(mask, float('-inf'))
attn_weights = F.softmax(scores, dim=-1) # (B, H, N, N)
attn_weights = self.dropout(attn_weights)
# 加权求和
out = torch.matmul(attn_weights, v) # (B, H, N, d_k)
# 拼接所有头
out = out.transpose(1, 2).reshape(B, N, D) # (B, N, D)
out = self.W_o(out) # 输出投影
return out
# --- 使用示例 ---
d_model = 512
num_heads = 8
mha = MultiHeadAttention(d_model, num_heads, dropout=0.1)
# 模拟一个 batch:batch_size=2, seq_len=10
x = torch.randn(2, 10, d_model)
# Encoder 风格(双向,所有位置互相可见)
out_enc = mha(x, causal=False)
# Decoder 风格(因果掩码,只能看过去)
out_dec = mha(x, causal=True)
print(f"Input shape: {x.shape}")
print(f"Output shape: {out_enc.shape}")
print(f"Causal output (未来被屏蔽):\n{mha(x, causal=True).shape}")

实现细节:实际工程中(如 HuggingFace Transformers)会将 Q/K/V 的投影合并为一次矩阵乘法(W_qkv),减少 kernel launch 开销。此外,现代 LLM 推理使用 KV Cache 和 FlashAttention 等优化技术来加速,详见 LLM 推理优化。

原始 Attention 机制本身是顺序无关的(permutation equivariant)——打乱输入顺序,输出只是跟着打乱,不会改变每个 token 的表示。这意味着”猫追狗”和”狗追猫”在模型看来没有区别。为了让模型感知位置,需要注入位置信息。

主流方案有两种:

  • 绝对位置编码(Learned / Sinusoidal,BERT/GPT-2 使用):直接给每个位置一个固定的编码向量
  • 相对位置编码(RoPE/ALiBi,LLaMA/GPT-NeoX 使用):编码两个位置之间的相对距离

RoPE(Rotary Position Embedding,旋转位置编码) 的核心思想:用旋转操作将绝对位置信息编码到相对位置信息中。具体来说,在二维空间中,如果把向量 (q0,q1)(q_0, q_1) 旋转角度 θ\theta,那么两个位置 mm 和 nn 的 Query 和 Key 在做点积时,结果只依赖于它们的相对位置 m−nm - n。

对于位置 mm 的向量 qq,RoPE 将其乘以旋转矩阵 RmR_m:

Rm=(cos⁡mθ0−sin⁡mθ0sin⁡mθ0cos⁡mθ0)R_m = \begin{pmatrix} \cos m\theta_0 & -\sin m\theta_0 \\ \sin m\theta_0 & \cos m\theta_0 \end{pmatrix}

对于 dd 维向量,RoPE 将其视为 d/2d/2 个二维子空间的组合,每个子空间应用不同频率 θi\theta_i 的旋转:

θi=10000−2i/d,i=0,1,…,d/2−1\theta_i = 10000^{-2i/d}, \quad i = 0, 1, \ldots, d/2 - 1

完整的旋转操作(以 Query 为例):

q~m=Rmq,其中 Rm=diag ⁣(R(mθ0),R(mθ1),…,R(mθd/2−1))\tilde{q}_m = R_m q, \quad \text{其中 } R_m = \text{diag}\!\left(R(m\theta_0), R(m\theta_1), \ldots, R(m\theta_{d/2-1})\right)

RoPE 的精妙之处在于:当旋转后的 qmq_m 和 knk_n 做点积时,结果只依赖于 m−nm - n:

q~m⋅k~n=qTRmTRnk=qTRn−mk\tilde{q}_m \cdot \tilde{k}_n = q^T R_m^T R_n k = q^T R_{n-m} k

因为旋转矩阵满足 RmTRn=Rn−mR_m^T R_n = R_{n-m}(旋转的复合性质)。这意味着 Attention 分数自动编码了相对距离 n−mn - m,而无需显式建模。

def precompute_freqs_cis(dim: int, max_seq_len: int, theta: float = 10000.0):
"""预计算 RoPE 的旋转频率(只需计算一次)"""
freqs = 1.0 / (theta ** (torch.arange(0, dim, 2).float() / dim))
t = torch.arange(max_seq_len).float()
freqs = torch.outer(t, freqs) # (max_seq_len, dim/2)
freqs_cis = torch.polar(torch.ones_like(freqs), freqs) # 复数形式 e^{iθ}
return freqs_cis
def apply_rotary_emb(x: torch.Tensor, freqs_cis: torch.Tensor) -> torch.Tensor:
"""
Args:
x: (B, num_heads, seq_len, head_dim)
freqs_cis: (seq_len, head_dim/2) 复数
Returns:
旋转后的 x,形状不变
"""
# 将相邻两维视为复数的实部和虚部
x_complex = torch.view_as_complex(x.float().reshape(*x.shape[:-1], -1, 2))
# 复数乘法 = 旋转
x_rotated = torch.view_as_real(x_complex * freqs_cis).flatten(-2)
return x_rotated.type_as(x)
# 在 Multi-Head Attention 中使用 RoPE:
# 在计算 scores = Q @ K^T 之前,对 Q 和 K 分别应用 apply_rotary_emb

RoPE vs 可学习位置编码:可学习位置编码(如 BERT)给每个位置一个可训练的向量,简单有效但无法外推——训练时最大长度 512,推理时超过 512 就不知道用什么编码了。RoPE 通过数学性质保证了相对位置的编码,配合 Length Extension 技术(如 NTK-aware scaling、YaRN),可以在推理时处理远超训练长度的序列,这也是现代 LLM 普遍采用 RoPE 的原因。

注意力计算的时间与空间复杂度

Section titled “注意力计算的时间与空间复杂度”

标准 Self-Attention 的计算复杂度为:

时间复杂度: O(n2⋅d),空间复杂度: O(n2)\text{时间复杂度: } O(n^2 \cdot d), \quad \text{空间复杂度: } O(n^2)

其中 nn 是序列长度,dd 是模型维度。n2n^2 来源于注意力矩阵 A∈Rn×nA \in \mathbb{R}^{n \times n}——每个位置都要和所有其他位置计算相似度。当 nn 很大时(如长文本),这成为严重的瓶颈。

序列长度 nn注意力矩阵大小显存占用(FP16)
512512 × 512~0.5 MB
4,0964,096 × 4,096~32 MB
32,76832,768 × 32,768~2 GB
131,072131,072 × 131,072~32 GB

这就是为什么处理超长上下文需要 FlashAttention、Sparse Attention、Sliding Window Attention 等优化技术。详见 注意力变体和 LLM 推理优化。

场景推荐方案原因
短序列编码(BERT 类)标准 Self-Attentionnn 小,n2n^2 可接受
自回归生成(GPT/LLaMA 类)Causal Self-Attention + KV Cache因果掩码 + 缓存避免重复计算
超长上下文(>32K tokens)FlashAttention + Sliding Window缓解 O(n2)O(n^2) 瓶颈
位置编码选择RoPE(首选)相对位置 + 可外推
多头数选择h=d/64h = d/64(经验值)每头 64 维是平衡点

FlashAttention(Tri Dao 等)通过 tiling(分块) 和 kernel fusion 技术,在不改变 Attention 数学结果的前提下,将 HBM(显存)读写量从 O(n2)O(n^2) 降到 O(n)O(n),实现了 2-4 倍的速度提升和 5-20 倍的显存节省。FlashAttention-2(2023)进一步优化了并行度和 warp 级效率。FlashAttention-3(2024)针对 H100 GPU 的异步特性进行深度优化,在 H100 上实现了接近理论峰值的吞吐。详见 LLM 推理优化。

注意力变体:Linear Attention 与 Sparse Attention

Section titled “注意力变体:Linear Attention 与 Sparse Attention”

为突破 O(n2)O(n^2) 瓶颈,多种高效注意力变体被提出:

  • Sliding Window Attention(Mistral/GLM 系列):每个位置只关注局部窗口 ww 个 token,复杂度降为 O(n⋅w)O(n \cdot w)。通过堆叠多层,感受野逐层扩大。
  • Multi-Query Attention (MQA) 与 Grouped-Query Attention (GQA):多个 Query 头共享同一组 Key/Value,显著减少 KV Cache 显存。LLaMA-2/3 使用 GQA。
  • Linear Attention:将 Softmax 近似为线性核函数 ϕ(q)⋅ϕ(k)\phi(q) \cdot \phi(k),避免显式构建 n×nn \times n 矩阵,复杂度降为 O(n⋅d)O(n \cdot d)。

详见 注意力变体。

为了让 RoPE 支持超长上下文(如 128K、1M tokens),一系列长度外推方法被提出:

  • Position Interpolation (PI):将目标位置等比例缩放到训练范围内,简单但牺牲远处位置的分辨率。
  • NTK-aware Scaling:调整 RoPE 的 base frequency θ\theta,让低频分量(编码远距离)保持外推能力。Code Llama 和早期 LLaMA-2 long 使用。
  • YaRN(Yet another RoPE extensioN):分段缩放——对高频分量插值、对低频分量外推,是 LLaMA-3、Qwen2 等模型支持超长上下文的标准方案。
  • LongRoPE(2024):微软提出,通过进化搜索自动寻找每个维度最优的缩放因子,支持将 4K 训练的模型外推到 2M 上下文。

DeepSeek 在 2025 年提出 Native Sparse Attention (NSA),这是一种硬件对齐的稀疏注意力机制。NSA 结合了 token-level 压缩和 block-level 选择,在保持接近 Full Attention 性能的同时,将推理速度提升数倍。与传统稀疏注意力不同,NSA 从预训练阶段就端到端地学习稀疏模式,而非事后修改,这使得稀疏模式更加自然高效。

Attention-free / Linear 模型的复兴(2024-2025)

Section titled “Attention-free / Linear 模型的复兴(2024-2025)”

一系列工作试图完全替代 Attention 或将其变为线性复杂度:

  • Mamba / State Space Models (SSM):用状态空间模型替代 Attention,具有 O(n)O(n) 复杂度和类似 RNN 的推理效率。Mamba-2(2024)改进了表达力。
  • Linear RNN / RWKV:将 Attention 替换为线性递归,推理时无需 KV Cache。RWKV-6(2024)在多项基准上接近同规模 Transformer。
  • Hybrid 架构:Jamba(AI21,2024)交替堆叠 Mamba 和 Attention 层,兼顾效率和表达力。这成为 2025 年长上下文模型的热门方向。

展望:Attention 机制从 2017 年提出至今,仍然是 LLM 的核心组件。但 O(n2)O(n^2) 复杂度的天然瓶颈驱动着替代方案的研究。2025 年的趋势是 混合架构(少量 Attention + 大量线性层),以及硬件感知的稀疏注意力设计。详见 注意力变体。

  • Vaswani et al. “Attention Is All You Need” (NeurIPS 2017)
  • Su et al. “RoFormer: Enhanced Transformer with Rotary Position Embedding” (2021)
  • Dao et al. “FlashAttention: Fast and Memory-Efficient Exact Attention” (NeurIPS 2022)
  • Peng et al. “YaRN: Efficient Context Window Extension of Large Language Models” (2023)
  • Liu et al. “Ring Attention with Blockwise Transformers for Near-Infinite Context” (2023)
  • DeepSeek “Native Sparse Attention” (2025)