在 3.1 节中,我们看到 RNN 的串行结构无法高效利用现代算力;在 3.2 节中,我们介绍了 Transformer 的编码器-解码器框架,并指出 Self-Attention 是支撑这一框架的“核心引擎”。现在,我们进入 Transformer 的心脏地带,把 Self-Attention 的数学过程完整拆开。
理解这一节的推导,是你后续掌握多头注意力(3.4 节)、位置编码(3.5 节)以及解码器-only 架构(3.7 节)的必备前提。更重要的是,很多工程决策——比如“为什么上下文加长会爆显存”“KV-Cache 到底在缓存什么”——答案都藏在这几个矩阵乘法里。
一、“自”注意力:同一序列的三重角色
Self-Attention 直译为“自注意力”,关键词是“自”:对于输入序列中的每一个 Token,模型把它同时当作三种角色来使用:
- 查询(Query, Q):当前 Token 在问:“上下文中哪些词与我相关?”
- 键(Key, K):序列中每个 Token 打出的标签:“我包含什么信息,来回答查询?”
- 值(Value, V):序列中每个 Token 实际携带的语义内容。
直觉类比:想象你在图书馆查资料。Query 是你提出的问题;Key 是每本书的索引标签;Value 是书里的实际内容。Self-Attention 就是在做一场“全场检索”——每个词都向全场发问,并根据相关性加权整合所有人的内容。
二、数学推导:六步完整过程
我们假设输入是一个长度为 $n$ 的序列,经过词嵌入(可能已加上位置编码,见 3.5 节)后,得到输入矩阵:
$$\mathbf{X} \in \mathbb{R}^{n \times d_{\text{model}}}$$
其中,$d_{\text{model}}$ 是模型隐藏层维度(如 512、1024、4096 等)。
步骤 1:可学习的线性投影
模型通过三组独立的、可训练的权重矩阵,将 $\mathbf{X}$ 投影到三个不同的子空间:
$$
\mathbf{W}_Q \in \mathbb{R}^{d_{\text{model}} \times d_k}, \quad
\mathbf{W}_K \in \mathbb{R}^{d_{\text{model}} \times d_k}, \quad
\mathbf{W}_V \in \mathbb{R}^{d_{\text{model}} \times d_v}
$$
步骤 2:计算 Q、K、V
$$
\mathbf{Q} = \mathbf{X}\mathbf{W}_Q \in \mathbb{R}^{n \times d_k} \\
\mathbf{K} = \mathbf{X}\mathbf{W}_K \in \mathbb{R}^{n \times d_k} \\
\mathbf{V} = \mathbf{X}\mathbf{W}_V \in \mathbb{R}^{n \times d_v}
$$
注意:$\mathbf{Q}$、$\mathbf{K}$、$\mathbf{V}$ 并非存储在模型中的固定知识,而是每次前向传播时由输入 $\mathbf{X}$ 动态计算得出。这意味着 Attention 是一种“输入自适应”的计算。
步骤 3:计算相似度分数(Scaled Dot-Product)
为了让每个 Token 知道该“关注”谁,我们计算 Query 与 Key 的点积:
$$\text{Scores} = \mathbf{Q}\mathbf{K}^\top \in \mathbb{R}^{n \times n}$$
矩阵中第 $i$ 行第 $j$ 列的元素,表示第 $i$ 个 Token 对第 $j$ 个 Token 的原始注意力分数。
步骤 4:缩放(Scaling)
直接将点积送入 Softmax 会有数值稳定性问题:当 $d_k$ 较大时,点积的数值方差会随之放大,导致 Softmax 进入梯度极小的饱和区。因此需要除以 $\sqrt{d_k}$ 进行缩放:
$$\text{Scores}_{\text{scaled}} = \frac{\mathbf{Q}\mathbf{K}^\top}{\sqrt{d_k}}$$
实用认知:在混合精度训练(FP16/BF16)中,这一步缩放尤为关键。如果不做 Scaling,点积结果可能直接溢出 FP16 的表示范围(~65504),导致训练 NaN。
步骤 5:Softmax 归一化
对每一行应用 Softmax,将分数转化为概率分布(每行之和为 1):
$$\mathbf{A} = \text{softmax}\left(\frac{\mathbf{Q}\mathbf{K}^\top}{\sqrt{d_k}}\right) \in \mathbb{R}^{n \times n}$$
这里的 $\mathbf{A}$ 就是注意力权重矩阵。你可以把它看作一张“关系热力图”:第 $i$ 行描述了第 $i$ 个 Token 对序列中所有 Token(包括自己)的关注程度。
步骤 6:加权聚合输出
用注意力权重对 Value 做加权求和,得到最终输出:
$$\mathbf{Z} = \mathbf{A}\mathbf{V} \in \mathbb{R}^{n \times d_v}$$
每一行 $\mathbf{Z}_i$ 是第 $i$ 个 Token 经过“上下文聚合”后的新向量表示。它既保留了自身的语义(因为对角线元素通常有较高权重),也融入了与之高度相关的其他 Token 的信息。
完整公式
将上述过程浓缩为一个紧凑的表达式,即原始 Transformer 论文中的 Scaled Dot-Product Attention:
$$\text{Attention}(\mathbf{Q}, \mathbf{K}, \mathbf{V}) = \text{softmax}\left(\frac{\mathbf{Q}\mathbf{K}^\top}{\sqrt{d_k}}\right)\mathbf{V}$$
三、因果掩码:自回归生成的“守门人”
在 3.2 节我们提到,解码器中的自注意力层是“带掩码的”。这并非可选配置,而是 LLM(如 GPT 系列)实现自回归生成的数学保障。
核心问题:在训练时,如果解码器能看到未来的 Token,模型就会偷懒——直接复制答案,而不是学习逐步推理。在推理时,未来 Token 根本还不存在。
解决方案:引入一个上三角掩码矩阵 $\mathbf{M}$,其主对角线以上元素为 $-\infty$(或一个极大的负数),以下及对角线为 0。
$$\mathbf{A} = \text{softmax}\left(\frac{\mathbf{Q}\mathbf{K}^\top}{\sqrt{d_k}} + \mathbf{M}\right)$$
经过 Softmax 后,$-\infty$ 对应的位置权重变为 0。这确保了在计算第 $i$ 个位置时,模型只能“看”到位置 $\leq i$ 的信息,绝对无法偷看未来。
实用影响:
- 这个掩码在 Decoder-only 架构(3.7 节)中是标配;
- 它使得 $\mathbf{Q}\mathbf{K}^\top$ 的注意力矩阵虽然仍是 $n \times n$,但实际有效的计算量是 $n(n+1)/2$;
- 这也是推理阶段可以使用 KV-Cache 优化的数学基础(后续工程章节详解)。
四、计算复杂度与工程现实
Self-Attention 的每一步都对应明确的计算与显存开销,这直接决定了你能部署的上下文长度:
| 操作 | 计算复杂度 | 显存占用 |
|------|-----------|----------|
| $\mathbf{Q}, \mathbf{K}, \mathbf{V}$ 投影 | $O(3 \cdot n \cdot d_{\text{model}} \cdot d_k)$ | $O(n \cdot d_k)$ 各存一份 |
| $\mathbf{Q}\mathbf{K}^\top$ | $O(n^2 \cdot d_k)$ | $O(n^2)$(注意力矩阵)|
| Softmax + 乘 $\mathbf{V}$ | $O(n^2 \cdot d_k)$ | 复用 $O(n^2)$ |
| 输出投影(如有) | $O(n \cdot d_v \cdot d_{\text{model}})$ | — |
关键结论:
- 序列长度的平方瓶颈:Self-Attention 的计算量和显存占用与 $n^2$ 成正比。当上下文从 4K 扩展到 128K,注意力矩阵的显存需求会暴涨 1024 倍。这是长上下文(Long Context)技术的核心攻坚点。
- 并行性的来源:与 RNN 的逐步递推不同,$\mathbf{Q}\mathbf{K}^\top$ 是一个巨大的矩阵乘法,可以一次性在 GPU 上并行完成。这正是 Transformer 能够 Scale 到千卡集群训练的根本原因。
五、与 RNN 的本质差异:一步全局 vs. 逐步传递
我们用一张对比表,把 Self-Attention 与 RNN 在信息传递上的差异固化下来:
| 特性 | RNN / LSTM | Self-Attention |
|------|-----------|----------------|
| 长程依赖 | 信息需逐步传递,易衰减 | 任意两 Token 直接交互 |
| 计算方式 | 时间步串行 | 矩阵并行 |
| 训练效率 | 难以大规模并行 | 天然适配 GPU/TPU |
| 位置感知 | 隐含在递推中 | 需外挂位置编码(3.5 节)|
| 路径长度 | $O(n)$ | $O(1)$(任意两点)|
六、实用认知:四个关键真相
1. Q、K、V 不是“三种记忆库”
很多初学者误以为模型里有三个独立的“知识库”。实际上,$\mathbf{W}_Q, \mathbf{W}_K, \mathbf{W}_V$ 只是投影矩阵,它们把同一个输入 $\mathbf{X}$ 映射到三个不同的语义空间,以便计算相关性。知识仍然分布式地存储在所有参数中。
2. Attention 权重是可解释的(有限地)
你可以可视化 $\mathbf{A}$ 矩阵,观察模型在生成某个词时到底“看”了哪些前文。这在调试模型行为、分析幻觉来源时非常有用。但要注意:高权重不等于“模型在复制”,低权重也不等于“完全忽略”。
3. $d_k$ 的选择是速度与质量的权衡
$d_k$ 越大,Q/K 的表达能力越强,但点积计算量和显存压力也越大。在多头注意力(3.4 节)中,通常将 $d_k$ 设为 $d_{\text{model}} / h$($h$ 为头数),以在总参数量不变的前提下,换取多视角关注的能力。
4. Self-Attention 本身没有位置概念
如果你把输入序列打乱,Self-Attention 的输出只会在行/列上同步 permute,其计算逻辑完全感知不到“谁在前、谁在后”。这就是为什么 3.5 节的位置编码(Position Encoding)是不可或缺的——没有它,Transformer 就是一个“词袋”模型。
七、小结
Self-Attention 的数学本质可以概括为:
将输入序列通过可学习的线性投影,分解为查询、键、值三重表示;通过缩放点积计算全局 pairwise 相关性,以 Softmax 权重对值向量做加权聚合,从而让每个 Token 在一次矩阵运算中获得全序列的上下文信息。
它用纯矩阵乘法替代了 RNN 的循环递推,用全局直连打破了长程依赖的瓶颈,也为 Transformer 的规模化训练铺平了道路。
在 3.4 节中,我们将回答一个自然的问题:如果一组 Q/K/V 已经能捕捉上下文关系,为什么要重复做 $h$ 组?多头注意力(Multi-Head Attention)的设计,正是为了把 Self-Attention 的单一视角扩展为语义空间的多维扫描。