Self-Attention

自注意力(Self-Attention)是一种让序列每个位置直接关注序列中所有位置(含自身)的注意力机制:查询(Query)、键(Key)、值(Value)均来自同一输入序列。Vaswani 等(2017)在 Transformer 中将其作为核心构建块,以 Scaled Dot-Product Attention 实现,彻底摆脱 RNN 的时间步串行约束,成为现代大语言模型(Large Language Model,LLM)的基石。

段末注释Self-Attention 的”自”指 Q/K/V 同源;与之相对,Cross-Attention(交叉注意力)的 Q 来自一方序列、K/V 来自另一方(如解码器 attend 编码器)。

读前说明:配图位于同名目录 2105.神经网络-自注意力-Self-Attention/。本文独立成篇;更宽泛的注意力机制入门(含 Nadaraya-Watson 核回归)见 2105.神经网络-注意力机制;Transformer 整体见 2105.神经网络-Transformer


1. 从何处来:演进脉络

图 1 从 RNN+Attention 到 Self-Attention 与 Transformer(科普示意)

1.1 前驱:RNN 与编码器瓶颈

循环神经网络(Recurrent Neural Network,RNN)及其门控变体(LSTMGRU)按时间步递推隐藏状态,存在:

  • 串行计算:$t$ 必须等 $t-1$ 完成;
  • 固定向量瓶颈:seq2seq 编码器将整个源句压入单个 $\mathbf{h}_T$,长句信息损失严重;
  • 长距离路径:位置 $i$ 与 $j$ 的信息需经 $|i-j|$ 步递归传递。

1.2 注意力机制的介入

年份 工作 贡献
2014 Bahdanau 等 加性注意力(Additive Attention):解码每步动态 attend 编码器所有隐状态
2015 Luong 等 乘性/点积注意力(Multiplicative/Dot Attention):更简洁高效
2016 Cheng 等、Lin 等 在 RNN 内部引入 self-attention 思想
2017 Vaswani 等,Attention Is All You Need Scaled Dot-Product Self-Attention + Multi-Head Attention,纯注意力架构 Transformer

Bahdanau 注意力解决的是 Cross-Attention(解码 query attend 编码 key/value);Transformer 将其推广为同一序列内部的 Self-Attention,并用堆叠层替代全部 RNN。

1.3 后续发展

阶段 代表 要点
2018 BERT、GPT 预训练 + Self-Attention 成为 NLP 默认
2020+ ViT、CLIP Self-Attention 扩展至视觉
2022+ LLaMA、GPT-4 等 仅微调架构细节,核心仍为 Multi-Head Self-Attention
2023+ FlashAttention、线性注意力 缓解 $O(T^2)$ 显存与计算

段末注释Multi-Head Attention(多头注意力)并行运行 $h$ 组独立 Self-Attention,再拼接投影,使模型在不同子空间捕获不同关系模式。


2. 要解决什么问题

2.1 核心问题形式化

给定序列 $\mathbf{X} = [\mathbf{x}_1, \ldots, \mathbf{x}_T]^\top \in \mathbb{R}^{T \times d}$,希望每个位置 $i$ 的输出 $\mathbf{o}_i$ 能直接融合全序列信息:

$$
\mathbf{o}i = \sum{j=1}^{T} \alpha_{ij}, \mathbf{v}j, \qquad \sum_j \alpha{ij} = 1,; \alpha_{ij} \ge 0
$$

其中 $\alpha_{ij}$ 为位置 $i$ 对位置 $j$ 的注意力权重。理想性质:

性质 RNN 的问题 Self-Attention
任意距离依赖 路径长 $O(|i-j|)$ 路径长 $O(1)$
训练并行 时间维串行 一次矩阵乘覆盖全序列
动态权重 隐状态固定维度压缩 权重随 $(i,j)$ 内容自适应
可解释性 门控/隐状态难读 $\alpha_{ij}$ 可可视化

2.2 Q/K/V 抽象

Self-Attention 将”谁查、查谁、取什么”参数化为:

  • Query $\mathbf{q}_i$:位置 $i$ 的查询——“我在找什么信息”;
  • Key $\mathbf{k}_j$:位置 $j$ 的——“我提供什么索引”;
  • Value $\mathbf{v}_j$:位置 $j$ 的——“被取走的内容”。

相似度 $\mathrm{sim}(\mathbf{q}_i, \mathbf{k}j)$ 越大,$\alpha{ij}$ 越大。Q/K/V 均由同一 $\mathbf{x}_i$ 线性投影得到,故为 Self-Attention。


3. 模型架构与原理

图 2 Scaled Dot-Product Self-Attention 结构(科普示意)

3.1 符号约定

符号 含义 典型维度
$\mathbf{X}$ 输入序列 $\mathbb{R}^{T \times d_{\mathrm{model}}}$
$\mathbf{W}_Q, \mathbf{W}_K, \mathbf{W}_V$ 投影矩阵 $\mathbb{R}^{d_{\mathrm{model}} \times d_k}$(或 $d_v$)
$\mathbf{Q}, \mathbf{K}, \mathbf{V}$ 投影后的 Q/K/V $\mathbb{R}^{T \times d_k}$、$\mathbb{R}^{T \times d_k}$、$\mathbb{R}^{T \times d_v}$
$d_k$ 键/查询维度 常取 $d_{\mathrm{model}}/h$($h$ 为头数)
$\alpha_{ij}$ 注意力权重 标量

3.2 Scaled Dot-Product Attention

矩阵形式(一次处理全序列):

$$
\mathrm{Attention}(\mathbf{Q}, \mathbf{K}, \mathbf{V}) = \mathrm{softmax}!\left(\frac{\mathbf{Q}\mathbf{K}^\top}{\sqrt{d_k}}\right) \mathbf{V}
$$

逐元素理解

$$
\alpha_{ij} = \frac{\exp(\mathbf{q}i^\top \mathbf{k}j / \sqrt{d_k})}{\sum{j’=1}^{T} \exp(\mathbf{q}i^\top \mathbf{k}{j’} / \sqrt{d_k})}, \qquad
\mathbf{o}i = \sum{j=1}^{T} \alpha
{ij}, \mathbf{v}_j
$$

缩放因子 $\sqrt{d_k}$:当 $d_k$ 较大时,点积方差随 $d_k$ 增大,softmax 易饱和导致梯度消失;除以 $\sqrt{d_k}$ 稳定训练(Vaswani 等,2017)。

投影

$$
\mathbf{Q} = \mathbf{X}\mathbf{W}_Q, \quad \mathbf{K} = \mathbf{X}\mathbf{W}_K, \quad \mathbf{V} = \mathbf{X}\mathbf{W}_V
$$

3.3 Multi-Head Self-Attention

$h$ 个头并行计算,每头维度 $d_k = d_v = d_{\mathrm{model}}/h$:

$$
\mathrm{head}_i = \mathrm{Attention}(\mathbf{X}\mathbf{W}_Q^{(i)}, \mathbf{X}\mathbf{W}_K^{(i)}, \mathbf{X}\mathbf{W}_V^{(i)})
$$

$$
\mathrm{MultiHead}(\mathbf{X}) = \mathrm{Concat}(\mathrm{head}_1, \ldots, \mathrm{head}_h), \mathbf{W}_O
$$

不同头可学习关注语法、共指、局部 n-gram、长程依赖等不同模式。

3.4 掩码变体

类型 作用 典型场景
Padding Mask 忽略填充位 batch 内不等长序列
Causal Mask(下三角) 位置 $i$ 只能看 $j \le i$ GPT 类自回归解码
无掩码 双向全连接 BERT 编码器

因果掩码矩阵 $\mathbf{M}$:$M_{ij} = -\infty$(或极大负数)当 $j > i$,使 $\alpha_{ij} = 0$。

3.5 与 Bahdanau 加性注意力对比

维度 Bahdanau(加性) Scaled Dot-Product
分数 $v^\top \tanh(W_q q + W_k k)$ $\mathbf{q}^\top \mathbf{k} / \sqrt{d_k}$
计算 需额外 MLP 纯矩阵乘,GPU 友好
并行 相对差 极佳

3.6 在 Transformer 块中的位置

标准 Transformer Encoder Layer

$$
\mathbf{X}’ = \mathrm{LayerNorm}(\mathbf{X} + \mathrm{MultiHead}(\mathbf{X}))
$$
$$
\mathbf{O} = \mathrm{LayerNorm}(\mathbf{X}’ + \mathrm{FFN}(\mathbf{X}’))
$$

FFN(前馈网络)为逐位置 MLP;Self-Attention 负责跨位置信息交换,FFN 负责逐位置非线性变换。


4. 训练哪些参数

4.1 单层 Multi-Head Self-Attention 参数

参数 形状 作用
$\mathbf{W}_Q^{(i)}, \mathbf{W}_K^{(i)}, \mathbf{W}_V^{(i)}$ 各 $\mathbb{R}^{d_{\mathrm{model}} \times d_k}$ 第 $i$ 头的 Q/K/V 投影
$\mathbf{W}_O$ $\mathbb{R}^{h d_v \times d_{\mathrm{model}}}$ 多头拼接后的输出投影

等价合并实现:PyTorch 常用单个大矩阵 $\mathbf{W}{QKV} \in \mathbb{R}^{d{\mathrm{model}} \times 3 d_{\mathrm{model}}}$ 一次投影再拆分。

参数量估算(单头简化,$d_k = d_v = d_{\mathrm{model}}$):

$$
|\Theta|{\mathrm{MHA}} \approx 4, d{\mathrm{model}}^2
$$

($\mathbf{W}_Q, \mathbf{W}_K, \mathbf{W}_V, \mathbf{W}O$ 各 $d{\mathrm{model}}^2$)

完整 Transformer 层还需 FFN(通常 $4 d_{\mathrm{model}}^2$ 量级)及 LayerNorm 的 $\gamma, \beta$。

4.2 通常还包括(模型级)

  • 词嵌入 $\mathbf{E}{\mathrm{token}} \in \mathbb{R}^{V \times d{\mathrm{model}}}$;
  • 位置编码(正弦或可学习 $\mathbf{E}_{\mathrm{pos}}$)——Self-Attention 本身置换等变,不含顺序信息;
  • 输出层(语言建模头)$\mathbf{W}{\mathrm{lm}} \in \mathbb{R}^{d{\mathrm{model}} \times V}$。

5. 如何训练

图 3 Self-Attention 训练:并行前向与端到端反向传播(科普示意)

5.1 与 RNN 训练的关键差异

维度 RNN (LSTM/GRU) Self-Attention
展开方式 时间步串行 全序列一次矩阵运算
反向传播 BPTT 标准 autograd(无时间展开瓶颈)
并行度 高(训练吞吐的核心优势)

5.2 损失函数

依任务而定,与序列模型通用:

  • 语言建模:$\mathcal{L} = -\sum_{t} \log P(x_t \mid x_{<t})$(因果掩码);
  • 掩码语言模型(MLM):仅对 masked 位置算交叉熵(BERT);
  • 序列标注:每 token 交叉熵;
  • 回归/分类:在 [CLS] 或池化表示上算损失。

5.3 反向传播要点

注意力权重 $\mathbf{A} = \mathrm{softmax}(\mathbf{Q}\mathbf{K}^\top / \sqrt{d_k})$ 对 $\mathbf{W}_Q, \mathbf{W}_K, \mathbf{W}_V$ 可微;梯度经 $\mathbf{A}$ 与 $\mathbf{V}$ 回传至输入 $\mathbf{X}$。

Softmax 梯度:若某行 $\mathbf{A}$ 接近 one-hot,其余 key 的梯度近零(注意力饱和);缩放 $\sqrt{d_k}$ 与多头机制有助于缓解。

深度堆叠:$L$ 层 Encoder 重复 Self-Attention + FFN,梯度经残差连接 $\mathbf{X} + \mathrm{SubLayer}(\mathbf{X})$ 传播,需 LayerNorm 与适当初始化(如 Xavier)稳定训练。

5.4 优化实践

技巧 说明
AdamW Transformer 默认;weight decay 与 Adam 解耦
Warmup + 衰减 学习率先线性增后余弦/多项式降(原论文)
Label Smoothing 缓解过拟合与过度自信
Dropout 加在注意力权重 $\mathbf{A}$ 或子层输出
Gradient Checkpointing 以算换显存,应对深模型
FlashAttention IO 感知的精确注意力,降显存

5.5 PyTorch 最小示例

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
import math
import torch
import torch.nn as nn
import torch.nn.functional as F

class ScaledDotProductSelfAttention(nn.Module):
"""单头 Scaled Dot-Product Self-Attention"""

def __init__(self, d_model: int, d_k: int):
super().__init__()
self.d_k = d_k
self.W_q = nn.Linear(d_model, d_k, bias=False)
self.W_k = nn.Linear(d_model, d_k, bias=False)
self.W_v = nn.Linear(d_model, d_k, bias=False)

def forward(self, x: torch.Tensor, mask: torch.Tensor | None = None) -> torch.Tensor:
# x: (batch, T, d_model)
q, k, v = self.W_q(x), self.W_k(x), self.W_v(x)
scores = q @ k.transpose(-2, -1) / math.sqrt(self.d_k) # (batch, T, T)
if mask is not None:
scores = scores.masked_fill(mask == 0, float("-inf"))
attn = F.softmax(scores, dim=-1)
return attn @ v # (batch, T, d_k)

nn.MultiheadAttentionnn.TransformerEncoder 提供生产级实现。


6. 应用与局限

图 4 Self-Attention 的典型局限(科普示意)

6.1 典型应用场景

场景 说明
大语言模型 GPT、LLaMA、Qwen 等解码器
预训练编码器 BERT、RoBERTa 双向建模
机器翻译 Transformer 仍是工业标准
视觉 ViT、Swin 将 patch 序列化后 Self-Attention
多模态 CLIP、LLaVA 的跨模态 Cross-Attention
生物序列 蛋白质/ DNA 长序列建模(需线性注意力等优化)

6.2 主要局限性

(1)$O(T^2)$ 时间与显存

注意力矩阵 $\mathbf{A} \in \mathbb{R}^{T \times T}$:序列长度翻倍,内存约 4 倍。$T > 10\text{K}$ 时常不可行;需 FlashAttention稀疏注意力线性注意力(Performers 等)或 分块/滑动窗口

(2)无内置顺序,依赖位置编码

Self-Attention 对输入置换等变(无 mask 时);必须额外注入绝对/相对位置编码。外推更长序列时,位置编码泛化常是瓶颈(ALiBiRoPE 等为此改进)。

(3)长上下文外推弱

训练长度 $L_{\mathrm{train}}$ 有限,推理 $L_{\mathrm{test}} \gg L_{\mathrm{train}}$ 时性能常显著下降;RNN 虽慢但在理论上无固定长度上限(实践中也受 BPTT 截断)。

(4)注意力权重 $\neq$ 因果解释

$\alpha_{ij}$ 可视化直观,但不保证对应人类语义”重要性”;对抗扰动下权重可大幅变化。

(5)小数据易过拟合

参数量 $O(d_{\mathrm{model}}^2)$ 级,小语料需预训练、正则或蒸馏。

(6)归纳偏置弱于 CNN/RNN

局部平移不变性(CNN)、顺序递归偏置(RNN)需靠数据量学出来;小样本或强结构数据上未必最优。

6.3 与 RNN 的定位对比

维度 LSTM/GRU Self-Attention
长依赖路径 $O(T)$ $O(1)$
训练并行
推理(自回归) 每步 $O(d^2)$ 需 KV Cache,仍可行
序列长度扩展 线性内存 二次内存(朴素实现)
归纳偏置 强顺序 弱,需位置编码

结论:Self-Attention 是大规模序列建模的默认选择;在极长流式、极低算力场景,RNN/GRU 或线性注意力变体仍有 niche。


7. 小结

主题 要点
由来 Bahdanau/Luong 注意力 → 2017 Transformer 的 Scaled Dot-Product Self-Attention
解决什么 任意距离依赖、训练并行、动态内容寻址
架构 $\mathrm{softmax}(\mathbf{Q}\mathbf{K}^\top/\sqrt{d_k})\mathbf{V}$;可扩展为 Multi-Head
训练参数 $\mathbf{W}_Q, \mathbf{W}_K, \mathbf{W}_V, \mathbf{W}_O$ 及 FFN、嵌入、位置编码
训练方式 标准 autograd + AdamW + warmup;因果/填充掩码
局限 $O(T^2)$、需位置编码、长上下文外推、权重非因果解释

参考文献

  1. Vaswani, A., et al. (2017). Attention Is All You Need. NeurIPS.
  2. Bahdanau, D., Cho, K., & Bengio, Y. (2014). Neural Machine Translation by Jointly Learning to Align and Translate. ICLR.
  3. Luong, T., Pham, H., & Manning, C. D. (2015). Effective Approaches to Attention-based Neural Machine Translation. EMNLP.
  4. Devlin, J., et al. (2018). BERT: Pre-training of Deep Bidirectional Transformers. NAACL.
  5. Dao, T., et al. (2022). FlashAttention: Fast and Memory-Efficient Exact Attention. NeurIPS.
-------------本文结束感谢您的阅读-------------