注意力机制(Attention Mechanism)是深度学习中最关键的技术突破之一。通俗地说,它让模型能够像人类阅读一样——在处理某个词时,不仅关注当前词本身,还能”看到”上下文中所有相关的词,并根据相关程度分配不同的”注意力权重”。

自 2014 年 Bahdanau 等人首次将注意力引入神经机器翻译以来,这一机制经历了从附加组件到核心架构的演变。2017 年 Vaswani 等人提出的 Transformer 模型更以”Attention Is All You Need”为题,将注意力推向了舞台中央。如今,从 BERT 到 GPT-4,几乎所有大规模语言模型都建立在注意力机制之上。

注意力机制的起源与动机

在注意力机制出现之前,序列到序列(Seq2Seq)模型是机器翻译的主流方案。它由编码器(Encoder)和解码器(Decoder)两部分组成:编码器将输入序列逐词读入 RNN/LSTM,将最后一个时刻的隐状态作为”上下文向量”传递给解码器,解码器再根据这个向量逐步生成输出序列。

问题就出在这个上下文向量上。无论输入是 5 个词还是 50 个词,编码器都将其压缩为一个固定长度的向量。短句子还能勉强装下,长句子必然丢信息——这就是著名的”瓶颈问题”。

Bahdanau 等人提出了一个优雅的解决方案:不再把所有信息塞进一个向量,而是让解码器在每一步生成时,都能”查看”编码器所有时刻的隐状态,并学会根据当前需要选择性地关注最相关的部分。具体做法是:解码器在每一步用当前隐状态作为”查询”(Query),与编码器各时刻的隐状态计算相关度,得到一组权重后对所有隐状态加权求和,生成当前步的动态上下文。

注意力机制的本质,是让模型在每一步解码时”查看”所有编码器输出,按相关程度分配权重。 这个看似简单的改动,彻底解决了瓶颈问题,也为后来的自注意力奠定了基础。

注意力的核心原理:Query、Key、Value

理解注意力机制最直观的方式是图书馆检索类比:你去图书馆找书,你的检索词就是 Query(查询),每本书的书名或关键词就是 Key(键),书的内容就是 Value(值)。注意力机制做的事就是——用 Query 和所有 Key 做匹配,匹配度越高的书,你越”关注”,最终拿到的是所有书内容的加权组合。

在 2017 年的 Transformer 论文中,Vaswani 等人将这一过程形式化为缩放点积注意力(Scaled Dot-Product Attention),其公式为:

$$\text{Attention}(Q, K, V) = \text{softmax}\left(\frac{QK^T}{\sqrt{d_k}}\right)V$$

这个公式看起来简洁,但信息密度很高。注意力机制的核心操作可以概括为三步:计算相似度、归一化、加权求和。 让我们逐一拆解:

  • 计算相似度:$QK^T$ 是 Query 与 Key 的点积矩阵,衡量每个 Query 与每个 Key 的相关程度。如果 Q 和 K 的维度都是 $d_k$,点积越大说明两者越”对齐”。
  • 归一化:softmax 将每个 Query 对所有 Key 的分数转化为概率分布(非负、求和为 1),这就是”注意力权重”。
  • 加权求和:用权重对 Value 加权求和,相关度高的 Value 贡献大,低的贡献小。

那么为什么要除以 $\sqrt{d_k}$ 呢?当 $d_k$ 较大时,点积的数值也会变大,导致 softmax 输出趋近于 one-hot 分布(某个位置接近 1,其余趋近 0),梯度几乎为零,训练停滞。除以 $\sqrt{d_k}$ 将点积的方差缩放回合理范围,使 softmax 输出更平滑,梯度更健康。

下图展示了一个简化的中英翻译示例中,注意力权重的分布模式:

可以看到,对角线区域权重最高(如”我”对应”I”的权重为 0.746),说明模型正确地将源词对齐到目标词。同时,非对角线也有非零权重(如”自然”对应”language”为 0.228),说明模型还能捕捉跨词的语义关联。

下面用 NumPy 实现一个最小化的缩放点积注意力:

import numpy as np

def softmax(x, axis=-1):
    """数值稳定的 softmax"""
    x_max = np.max(x, axis=axis, keepdims=True)
    exp_x = np.exp(x - x_max)
    return exp_x / np.sum(exp_x, axis=axis, keepdims=True)

def scaled_dot_product_attention(Q, K, V, d_k=64):
    """
    缩放点积注意力
    Q: (batch, seq_q, d_k)
    K: (batch, seq_k, d_k)
    V: (batch, seq_v, d_v)
    """
    # 1. 计算注意力分数: Q @ K^T / sqrt(d_k)
    scores = np.matmul(Q, K.transpose(0, 2, 1)) / np.sqrt(d_k)
    # scores shape: (batch, seq_q, seq_k)

    # 2. softmax 归一化
    attention_weights = softmax(scores, axis=-1)

    # 3. 加权求和
    output = np.matmul(attention_weights, V)
    # output shape: (batch, seq_q, d_v)

    return output, attention_weights

# 示例: 序列长度 4, 维度 64
batch_size = 1
seq_len = 4
d_k = 64

Q = np.random.randn(batch_size, seq_len, d_k)
K = np.random.randn(batch_size, seq_len, d_k)
V = np.random.randn(batch_size, seq_len, d_k)

output, weights = scaled_dot_product_attention(Q, K, V, d_k)

print(f"输出形状: {output.shape}")         # (1, 4, 64)
print(f"注意力权重形状: {weights.shape}")   # (1, 4, 4)
print(f"第一行权重和: {weights[0, 0].sum():.4f}")  # 1.0000

代码中,Q、K、V 的形状为 (batch, seq_len, d_k),np.matmul 批量计算点积,softmax 沿最后一个维度归一化。运行后,每行注意力权重之和为 1.0,输出形状与 Query 序列长度对齐。需要说明的是,上述示例中 Q、K、V 是独立随机生成的,展示的是注意力的一般计算流程;在自注意力场景下,三者由同一输入序列经不同的线性变换得到(详见“自注意力与多头注意力”一节)。

自注意力与多头注意力

自注意力

自注意力(Self-Attention) 是一种让序列中的每个元素能够直接关注同一序列中其他所有元素的机制。直观理解:对于一句话中的每个词,自注意力会计算它与其他所有词之间的相关性,然后根据相关性大小,从其他词中抽取信息,更新当前词的表示。

例如句子:The animal didn’t cross the street because it was too tired.在这句话中,代词 “it” 指代的是 “animal” 而不是 “street”。自注意力机制可以通过计算 “it” 与句中其他词的相关性,自动学习到 “it” 与 “animal” 的关系更强,从而让 “it” 的表示更多地包含 “animal” 的信息。

为什么需要自注意力?

在自注意力出现之前,序列建模主要依赖 RNN/LSTM/GRU 与 CNN,二者在长距离依赖与并行化方面均存在局限。下表对比了自注意力与传统序列模型的核心差异:

维度 RNN / LSTM / GRU CNN 自注意力
长距离依赖 弱:逐步顺序传递,梯度易消失 较弱:需堆叠多层扩大感受野 强:任意两位置直接交互
并行化 难:依赖时间步递推 可以并行 高度并行:不依赖时间步递推
权重性质 固定参数 固定卷积核 由输入内容动态决定
可解释性 较好:注意力矩阵可可视化

需要指出的是,自注意力的时间与空间复杂度均为 $O(n^2)$——序列中任意两个位置都要两两计算相关性。这一特性既是其全局建模能力的来源,也是后文稀疏注意力、线性注意力等高效变体的优化起点。

自注意力与交叉注意力的区别

Query、Key、Value 的检索框架已在“注意力的核心原理”一节中详细介绍,此处不再重复。自注意力与 Bahdanau 注意力的关键区别在于三者的来源:在 Bahdanau 注意力中,Query 来自解码器,Key 和 Value 来自编码器,这是一种跨序列的交叉注意力(Cross-Attention);而在自注意力中,Query、Key、Value 全部由同一个输入序列经三个不同的线性变换得到,序列中的每个元素既是“查询者”,也是“被查询者”——这正是“自”的含义。

自注意力的位置信息问题

自注意力本身不包含位置信息。因为计算过程对序列顺序不敏感,即打乱输入顺序,注意力计算结果中每个位置的输出会相应打乱,但计算方式不变。

为了让模型知道元素的位置,通常需要显式添加位置编码(Positional Encoding)。

常见方式:

  • 正弦位置编码:Transformer 原始论文使用不同频率的正弦和余弦函数。
  • 可学习位置嵌入:将位置作为可训练向量,与词嵌入相加。
  • 相对位置编码:在注意力分数中加入位置之间的相对关系。
  • 旋转位置编码(RoPE):通过旋转矩阵编码相对位置,广泛用于现代大模型。

掩码自注意力(Masked Self-Attention)

在解码器(如 GPT)中,需要防止当前 token 看到未来的 token,因此使用掩码自注意力

具体做法:

  • 在计算注意力分数后,对未来的位置施加一个非常大的负数(如-∞)。
  • 经过 softmax 后,未来位置的权重趋近于 0。

掩码矩阵示例:

$$\text{Mask} = \begin{bmatrix} 0 & -\infty & -\infty & -\infty \\ 0 & 0 & -\infty & -\infty \\ 0 & 0 & 0 & -\infty \\ 0 & 0 & 0 & 0 \end{bmatrix}$$

计算:

$$A = \text{softmax}\left(\frac{Q K^T}{\sqrt{d_k}} + \text{Mask}\right)$$

这样每个位置只能关注自己和之前的 token。

自注意力代码实现

import torch
import torch.nn as nn
import torch.nn.functional as F

class SelfAttention(nn.Module):
    """单头自注意力(支持因果掩码)"""
    def __init__(self, d_model):
        super().__init__()
        self.d_model = d_model
        self.W_Q = nn.Linear(d_model, d_model)
        self.W_K = nn.Linear(d_model, d_model)
        self.W_V = nn.Linear(d_model, d_model)
        self.W_O = nn.Linear(d_model, d_model)

    def forward(self, x, mask=None):
        # 线性变换:Q、K、V 均来自同一输入 x,故称“自”注意力
        Q = self.W_Q(x)  # (batch, seq_len, d_model)
        K = self.W_K(x)
        V = self.W_V(x)

        # 计算注意力分数
        scores = torch.matmul(Q, K.transpose(-2, -1)) / (self.d_model ** 0.5)

        # 因果掩码:屏蔽未来位置
        if mask is not None:
            scores = scores.masked_fill(mask == 0, float('-inf'))

        # softmax 归一化
        attn_weights = F.softmax(scores, dim=-1)

        # 加权求和
        output = torch.matmul(attn_weights, V)

        # 输出线性变换
        return self.W_O(output), attn_weights

# 测试:带因果掩码的自注意力
d_model = 64
seq_len = 5
x = torch.randn(2, seq_len, d_model)

# 下三角掩码:位置 i 只能看到 0..i
mask = torch.tril(torch.ones(seq_len, seq_len))

model = SelfAttention(d_model)
output, weights = model(x, mask=mask)
print(f"输出形状: {output.shape}")      # torch.Size([2, 5, 64])
print(f"权重形状: {weights.shape}")     # torch.Size([2, 5, 5])
# 验证因果性:每个位置之后的权重应为 0
print(f"位置 2 的权重: {weights[0, 2]}")

多头注意力

多头注意力是 Transformer 架构中的核心模块,它在自注意力(Self-Attention)的基础上,通过并行地使用多组不同的线性投影,让模型在不同的表示子空间中分别计算注意力,从而增强模型的表达能力。

多头注意力 并不是一个完全独立的注意力机制,而是将多个自注意力头并行组合起来的结构。

直观理解:与其只使用一组 Query、Key、Value 来计算注意力,不如使用多组 Query、Key、Value,让不同的头关注不同的关系模式,最后把各个头得到的信息拼接起来。

例如在句子:The animal didn’t cross the street because it was too tired.

  • 一个头可能重点关注 “it” 与 “animal” 的指代关系;
  • 另一个头可能关注 “cross” 与 “street” 的动宾关系;
  • 另一个头可能关注整体语义信息。

多个头从不同角度抽取信息,最终合并成一个更丰富的表示。

为什么需要多头注意力?

单头自注意力虽然也能建模全局依赖,但存在一个问题:单个注意力头只能产生一种注意力分布,即每个位置对所有位置的权重分配是固定的。

这意味着单头自注意力在某一时刻只能关注一种模式,而语言等复杂数据中往往同时存在多种关系:

  • 语法依存
  • 指代关系
  • 局部短语结构
  • 全局语义关联
  • 位置关系

多头注意力通过引入多个子空间,让不同的头学习到不同的注意力模式,从而提升模型的容量和灵活性。此外,多头将高维注意力拆分到多个低维子空间,有助于缓解单个头维度过大带来的优化困难。

多头注意力的计算过程

其公式为:

$$\text{head}_i = \text{Attention}(QW_i^Q, KW_i^K, VW_i^V)$$

$$\text{MultiHead}(Q,K,V) = \text{Concat}(\text{head}_1, \ldots, \text{head}_h)W^O$$

其中 $W_i^Q \in \mathbb{R}^{d_{\text{model}} \times d_k}$、$W_i^K \in \mathbb{R}^{d_{\text{model}} \times d_k}$、$W_i^V \in \mathbb{R}^{d_{\text{model}} \times d_v}$、$W^O \in \mathbb{R}^{hd_v \times d_{\text{model}}}$ 均为可学习的投影矩阵。以 Transformer 基础模型为例,$d_{\text{model}}=512$、$h=8$、$d_k=d_v=64$,8 个头各自在 64 维子空间中计算注意力,拼接后经 $W^O$ 映射回 512 维。

下面用 PyTorch 实现完整的多头注意力:

import torch
import torch.nn as nn
import torch.nn.functional as F
import math

class MultiHeadAttention(nn.Module):
    def __init__(self, d_model=512, n_heads=8):
        super().__init__()
        self.d_model = d_model
        self.n_heads = n_heads
        self.d_k = d_model // n_heads  # 64

        # 线性投影矩阵
        self.W_q = nn.Linear(d_model, d_model)
        self.W_k = nn.Linear(d_model, d_model)
        self.W_v = nn.Linear(d_model, d_model)
        self.W_o = nn.Linear(d_model, d_model)

    def forward(self, x):
        batch_size, seq_len, _ = x.size()

        # 1. 线性投影并分头
        Q = self.W_q(x).view(batch_size, seq_len, self.n_heads, self.d_k)
        K = self.W_k(x).view(batch_size, seq_len, self.n_heads, self.d_k)
        V = self.W_v(x).view(batch_size, seq_len, self.n_heads, self.d_k)

        # 转置: (batch, n_heads, seq_len, d_k)
        Q = Q.transpose(1, 2)
        K = K.transpose(1, 2)
        V = V.transpose(1, 2)

        # 2. 缩放点积注意力
        scores = torch.matmul(Q, K.transpose(-2, -1)) / math.sqrt(self.d_k)
        attention = F.softmax(scores, dim=-1)
        context = torch.matmul(attention, V)

        # 3. 合并多头
        context = context.transpose(1, 2).contiguous()
        context = context.view(batch_size, seq_len, self.d_model)

        # 4. 输出投影
        return self.W_o(context)

# 测试
model = MultiHeadAttention(d_model=512, n_heads=8)
x = torch.randn(2, 10, 512)  # batch=2, seq_len=10
output = model(x)
print(f"输出形状: {output.shape}")  # torch.Size([2, 10, 512])

代码中,d_k = d_model // n_heads 确保各头维度均分。view + transpose 操作将张量从 (batch, seq, d_model) 变形为 (batch, n_heads, seq, d_k),使各头在头维度(n_heads)上并行计算。最终 transpose + view 将多头输出拼接,经 W_o 投影回原维度。

注意力机制的演进与优化

标准注意力的时间复杂度为 $O(n^2 \cdot d)$,空间复杂度同样为 $O(n^2)$。当序列长度从 512 增长到 32K 甚至 128K 时,注意力矩阵的大小爆炸式增长,成为训练大模型的核心瓶颈。为此,研究者提出了多种优化方向。

  • 稀疏注意力(Sparse Attention)的核心思路是:大部分注意力权重接近零,只保留少数高权重的连接。Longformer 采用“滑动窗口 + 全局 token”模式,将复杂度降至 $O(n \cdot w)$($w$ 为窗口大小);BigBird 在此基础上加入随机连接,理论上将复杂度进一步降至 $O(n)$,并在图论意义上证明了表达能力。
  • 线性注意力(Linear Attention)用核函数 $\phi(\cdot)$ 近似 softmax,将 $O(n^2)$ 的注意力矩阵分解为 $O(n \cdot d^2)$ 的矩阵乘法:$\phi(Q)(\phi(K)^T V)$。当 $n \gg d$ 时,这一改写带来显著加速,但核函数近似会损失精度。
  • Flash Attention 是 2022 年 Dao 等人提出的里程碑工作。它不改变数学公式(精确注意力),而是优化 GPU 内存访问模式:将注意力矩阵分块加载到 SRAM(片上高速缓存)中计算,避免 HBM(主显存)与 SRAM 之间的冗余读写。这一 IO 优化的实际效果是在不损失任何精度的前提下,将注意力计算加速 2-4 倍、内存占用降低 5-20 倍。Flash Attention-2 进一步优化了并行度和 warp 级调度,已成为 LLaMA、GPT-4 等大模型的标准配置。
  • 分组查询注意力(Grouped-Query Attention, GQA)针对推理阶段的内存瓶颈。标准多头注意力中,每个头有独立的 K/V 矩阵,在自回归生成时需要缓存所有头的 KV——当头数 $h=96$ 时,KV 缓存占用惊人。GQA 让多个 Query 头共享同一组 K/V,在质量几乎不损的前提下大幅减少 KV 缓存大小。LLaMA-2/3、Mistral 等模型均采用 GQA。

下表汇总了主要注意力变体的对比:

变体 核心思路 复杂度 代表模型
标准注意力 全连接 softmax O(n² · d) Transformer / BERT
稀疏注意力 固定 / 可学习稀疏模式 O(n·w) ~ O(n) Longformer / BigBird
线性注意力 核函数近似 softmax O(n · d²) Linformer / Performer
Flash Attention IO 感知分块计算 O(n² · d) LLaMA / GPT-4
分组查询注意力 共享 K/V 头 O(n² · d) LLaMA-2/3 / Mistral

需要注意的是,Flash Attention 和 GQA 的复杂度仍为 $O(n^2 \cdot d)$——它们不是从数学上降低复杂度,而是从工程上减少常量因子。真正将复杂度降至线性的稀疏注意力和线性注意力,则在不同程度上牺牲了精度或表达能力。在实际应用中,Flash Attention + GQA 的组合已成为当前大模型的事实标准,兼顾了精度与效率。

结论

注意力机制从 2014 年的机器翻译附加组件,演变为今天大语言模型的核心架构,其发展轨迹始终围绕一个主线:在表达能力与计算效率之间寻找更优的平衡点。回顾全文,我们看到了几个关键节点:Bahdanau 注意力解决了 Seq2Seq 的信息瓶颈;缩放点积注意力将这一过程形式化为简洁的 QKV 框架;自注意力让序列内部的位置直接交互,取代了 RNN 的逐步传递;多头注意力让模型同时关注多种语义关系;Flash Attention 和 GQA 则在工程层面将注意力推向了千行级序列的实用化。

注意力机制的本质,是用可微的”软查找”替代离散的”硬检索”,让梯度优化成为可能。 这句话既是理解注意力机制的钥匙,也是展望其未来的起点——未来更高效的注意力变体,无论是稀疏化、线性化还是硬件感知优化,都是在同一个框架下追求更好的”查找效率”。

参考文献 / 扩展阅读

  • Bahdanau, D., Cho, K., & Bengio, Y. (2014).Neural Machine Translation by Jointly Learning to Align and Translate. arXiv:1409.0473.
  • Vaswani, A., et al. (2017).Attention Is All You Need. NeurIPS 2017. arXiv:1706.03762.
  • Dao, T., et al. (2022).FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness. NeurIPS 2022.
  • Ainslie, J., et al. (2023).GQA: Training Generalized Multi-Query Transformer Models from Multi-Head Checkpoints. arXiv:2305.13245.
  • Choromanski, K., et al. (2020).Rethinking Attention with Performers. ICLR 2021. arXiv:2009.14794.
0