一句话总结:Attention 不会在前向计算时直接改写词嵌入表;它先用注意力权重混合 V,生成当前语境下的临时表示。只有训练时经过损失、反向传播和优化器更新,词嵌入与 QKV 权重这些参数才会改变。
12.1 今天这一章非常轻松
今天这一章非常轻松,也是 Attention 机制的收尾章节。
上一章我们已经把多个头拆开,看过每个头怎样独立计算。现在只剩两个问题:
- 每个头算出的结果怎样合并回来?
- Attention 的输出到底改了什么?
第二个问题最容易混淆。包括我第一次学的时候,也会把下面两件事当成同一件事:
- 当前这句话里的 Token 表示发生了变化;
- 模型文件里的词嵌入参数发生了变化。
它们不是一回事。
前者发生在每一次前向计算里;后者只有在训练时,等反向传播算出梯度,再由优化器执行更新之后才会发生。
把这条界线看清楚,QKV、残差连接和后面的训练过程就都能接起来了。
12.2 先给 A 一个更准确的名字
前面为了方便,我们把 Attention 的结果写成 A。但在不同资料里,A 有时指注意力权重,有时又指加权后的输出,很容易看混。
这一章把它们分开:
这里:
- P 是注意力权重;
- M 是 Mask,被遮住的位置在 Softmax 后权重为 0;
- H_heads 才是各个头加权混合 V 之后的输出。
所以,不是 Q×K 一乘就直接得到“百分比”。完整顺序是:
- QKᵀ 得到原始匹配分数;
- 除以 √d_k 做缩放;
- 加上 Mask;
- 经过 Softmax,才得到每行和为 1 的权重 P;
- 用 P 对 V 做加权求和。
这些权重很适合解释成“这一次混合 V 时,各位置占多大比重”,但不要把它理解成模型对语义关系作出的严格概率判断。
12.3 H_heads 的形状:四维矩阵不可怕
继续沿用前面的例子:
H_heads 的形状:[4, 4, 16, 128] │ │ │ └── 每个头的输出维度 │ │ └────── Query 的序列长度 │ └───────── 头的数量 └──────────── 批次大小
四个维度分别代表:
- 4 个训练样本;
- 每个样本有 4 个头;
- 每个头要为 16 个 Query Token 生成输出;
- 每个 Token 在这个头里得到 128 维结果。
如果是普通的自注意力,Query、Key、Value 的序列长度相同,所以 P 的形状是:
[batch, heads, seq_len, seq_len]
但概念上最好保留两个名字:
[batch, heads, query_len, key_len]
这样以后看到交叉注意力,Query 和 Key 长度不同时,就不会觉得奇怪。
12.4 Concatenate:把多个头合并回来
12.4.1 合并操作
Concatenate 就是把各个头最后一维并回去:
合并前:[4, 4, 16, 128] 调整顺序:[4, 16, 4, 128] 合并后:[4, 16, 512]
这里要注意,程序通常不是随手把两个维度直接压平,而是先把 Token 维和 Head 维排到正确位置,再 reshape。
每个 Token 原来在 4 个头里各有 128 维。合并后,它重新拥有一个 512 维向量。
12.4.2 为什么先切开再合并?
这个概念是很多人不清晰的,包括我第一次学的时候也不明白:为什么要先切开再合并?绕一圈有什么意义?
因为每个头有自己对应的投影切片,可以在不同的表示子空间里计算 Attention。某些训练好的头可能更明显地响应位置、句法或特定关系,但这不是人为指定的岗位,也不是每个头一定会形成一种清楚可说的功能。
合并的意义,就是把这些不同子空间的结果重新放到同一个 Token 表示里。
12.5 W_O:合并后还要再学一次怎样组合
合并多头只是把结果排列到一起,还没有真正决定这些结果该怎样互相组合。因此后面还有一个输出投影 W_O:
在我们这个等宽例子里:
Concat 输出:[batch, 16, 512] W_O:[512, 512] H_attn:[batch, 16, 512]
W_O 是可训练参数。它不是简单地把 4 个头做平均,而是学习怎样重新组合各个头送来的特征。
如果 h×d_v 不等于 d_model,W_O 的一般形状是:
[h × d_v, d_model]
所以“W_O 一定是 512×512”只对当前这个等宽例子成立。
12.6 QK 和 V 分别在做什么?
12.6.1 QKᵀ 决定怎样混合
拿一个头来说:
Q:[query_len, d_k] K:[key_len, d_k] QKᵀ:[query_len, key_len]
QKᵀ 的每一个格子,是一个 Query 位置和一个 Key 位置的匹配分数。经过缩放、Mask 和 Softmax 之后,它才变成注意力权重 P。
12.6.2 V 提供被混合的内容
V 不是“原始文字”,也不是原封不动的词嵌入。它是当前这一层输入 X 经过 W_V 投影之后得到的特征:
然后:
对每一个 Query 位置来说,输出都是各个 V 向量的加权和。
一句最短的记法:
- Q 和 K 决定“从哪些位置取多少”;
- V 决定“实际取回什么特征”。
12.7 最容易混淆的地方:隐藏状态不是词嵌入表
12.7.1 词嵌入表是什么?
词嵌入表 E 是模型参数。输入 Token ID 后,它负责查出最初的向量:
同一个 Token ID 在查表时拿到的是同一个基础向量。位置信息加入以后,再经过一层层 Transformer Block,才逐渐变成当前语境下的表示。
12.7.2 隐藏状态是什么?
X_0、X_1、X_2……是这一次前向计算产生的中间结果,通常叫隐藏状态或激活值。
它们更像模型在处理当前这句话时写下的“临时工作笔记”。句子一换,即使其中出现同一个 Token,它在后面层里的隐藏状态也可能不同。
12.7.3 Attention 有没有直接改写输入?
没有。
Attention 分支从 X 计算出一个新张量 H_attn。下一章会讲到,常见做法是通过残差连接把它加回主干:
这里产生的是新的隐藏状态。原来的 X 没有被 Attention 原地改写,模型文件里的词嵌入表 E 也没有在这个前向步骤里被改写。
所以,不能把 PV 说成“把原始 Token 的初始化数值调准了一点”。更准确的说法是:
PV 为当前序列、当前层、当前头生成了一份结合上下文的临时表示。
12.8 训练时,参数到底什么时候改变?
这里把一次训练拆成三步,就不会再混淆。
12.8.1 第一步:前向计算
模型读取词嵌入表 E,以及 W_Q、W_K、W_V、W_O 等权重,算出隐藏状态、预测结果和损失。
这一步会产生很多激活值,但参数本身还没有改变。
12.8.2 第二步:反向传播
从损失出发,计算每个可训练参数的梯度:
反向传播算出来的是“应该往哪个方向改、影响有多大”,还不是参数更新本身。
12.8.3 第三步:优化器更新
最后由优化器根据梯度修改参数。用最简单的梯度下降写法:
这一步之后,词嵌入表和各层权重才真正有了新数值。下一批数据再经过模型时,读取的就是更新后的参数。
推理时通常只做前向计算,没有 backward,也没有 optimizer step,所以正常推理不会一边聊天一边改写模型参数。KV Cache 虽然会保存中间 K、V,但它也是本次推理的缓存,不是训练参数。
12.9 回到原来的“小沈阳”例子
假设训练语料里多次出现“小沈阳”这个名称。
第一次遇到它时:
- 输入中的每个 Token 先从词嵌入表查出基础向量;
- Attention 和后续网络根据整句话生成上下文化的隐藏状态;
- 模型根据目标算出损失;
- 反向传播算出梯度;
- 优化器才对相关词嵌入和其他权重做一次小更新。
下一次又遇到“小沈阳”时,模型确实会使用上一次优化器更新后的参数。但这并不是因为上一次的 Attention 输出被直接存回了词嵌入表,而是因为训练完整走过了:
前向计算 → 损失 → 反向传播 → 优化器更新
训练许多次以后,基础词嵌入会学到跨语料都能复用的信息;而每次句子里的隐藏状态,还会继续根据周围 Token 形成这一处特有的语境信息。
这两个层次要同时记住:
- 词嵌入参数:模型长期学到、保存在模型文件里的基础表示;
- 隐藏状态:当前输入临时算出来、随语境改变的表示。
12.10 哪些是参数,哪些不是?
| 名称 | 是否可训练参数 | 什么时候存在或变化 |
|---|---|---|
| 词嵌入表 E | 是 | 优化器更新后改变 |
| W_Q、W_K、W_V、W_O | 是 | 优化器更新后改变 |
| Q、K、V | 否 | 每次前向计算重新产生 |
| 注意力权重 P | 否 | 每次前向计算重新产生 |
| Attention 输出 | 否 | 每次前向计算重新产生 |
| 隐藏状态 X_l | 否 | 每次前向计算重新产生 |
| KV Cache | 否 | 推理过程中暂存和追加 |
这张表就是本章最重要的边界。
12.11 多头是不是共用同一个 W_Q?
这个问题最好换一种说法。
在常见实现里,程序可以用一个大的线性层一次算出完整 Q,再 reshape 成多个头。这个大矩阵包含了各个头对应的不同列切片:
因此:
- 从存储和计算上看,可以说它们来自同一个大 W_Q;
- 从学习到的参数上看,每个头使用的是其中不同的参数切片,并不是所有头拿同一组数字重复计算。
W_K 和 W_V 也是同样道理。
不同 Transformer Block 通常有各自独立的一套 Attention 参数;也存在跨层共享参数的特殊架构,但那不是我们这里讲的标准情况。
12.12 Attention 参数量怎么算?
在最常见的等宽多头注意力中,忽略 bias:
如果 d_model = 512:
每层 Attention 权重 = 4 × 512 × 512 = 1,048,576
如果有 12 层:
12 × 1,048,576 = 12,582,912
这是 Attention 投影权重的数量,不包含 bias、词嵌入、FFN、LayerNorm 和输出层。
训练循环多少次不会让参数数量越来越多。训练改变的是这些参数的数值,不是模型结构中参数的个数。模型是几百万、几十亿还是更多参数,主要由词表大小、层数、宽度、FFN 维度以及具体架构决定。
12.13 把整条因果链连起来
现在可以把 Multi-Head Attention 的最后一段完整写出来:
当前隐藏状态 X ↓ 线性投影得到 Q、K、V ↓ QKᵀ → 缩放 → Mask → Softmax ↓ 权重 P 与 V 相乘 ↓ 各头输出 Concatenate ↓ 经过 W_O ↓ 得到 Attention 分支输出
在一次普通前向计算中,这条链只生成新的激活值。
如果是在训练,还会继续:
预测 → 损失 → 反向传播 → 优化器更新参数
这时词嵌入表和 W_Q、W_K、W_V、W_O 才真正改变。
本章要点回顾
- QKᵀ 先产生匹配分数,经过缩放、Mask 和 Softmax 后才得到注意力权重
- P 决定怎样混合,V 提供被混合的特征
- Concatenate 把各头结果恢复到一个 Token 表示中,W_O 再学习怎样组合
- Attention 的前向计算生成新的隐藏状态,不会直接改写词嵌入表
- 反向传播计算梯度,优化器才真正更新参数
- Q、K、V、注意力权重和 KV Cache 都是运行时产生的值,不是模型参数
- 一个大投影矩阵可以包含各个头不同的参数切片,不等于所有头共用同一组数字
Part 3 小结
到这里,我们把 Attention 的几何直觉、QKV、多头拆分、合并和输出投影完整连起来了。
如果现在让你看下面这条公式:
你应该不只知道怎样算,也知道其中哪些是参数变换后的激活值,哪些是当前输入临时产生的权重,以及训练时参数到底在什么时候改变。
这才是 QKV 输出的本质。
下一章预告
下一章讲残差连接和 Dropout。
- 残差连接给主干留下一条直接通路,让深层网络更容易优化;
- Dropout 在训练时随机丢弃一部分激活,用来做正则化。
而 LayerNorm 放在残差分支之前还是之后,又会形成 Pre-Norm 和 Post-Norm 的区别。我们下一章接着看。
好了,这一章就到这里,拜拜。