第一次读 "Attention Is All You Need" 的时候,我只记住了 Q、K、V 三个字母。过了一周连公式都默写不出来。后来下定决心,关了所有参考资料,拿纸笔从 definition 开始手推——这次终于「通」了。

从 RNN 的局限说起

RNN 处理序列的核心问题是无法并行——第 t 个 token 的计算依赖第 t-1 个 token 的 hidden state。而且长序列的梯度传播会指数衰减(所谓的「长期依赖」问题)。LSTM 和 GRU 加了门控机制缓解了梯度消失,但并行的瓶颈依然在。

Self-Attention 换个思路:不再按顺序处理,而是让每个 token 直接「看到」序列中所有其他 token。这句话听起来简单,但实现起来需要一套精巧的机制。

Q、K、V 到底在做什么

我把 Self-Attention 理解成一个软寻址过程。每个 token 的 embedding 向量通过三个不同的线性变换矩阵 WQ、WK、WV 投影成三个角色:

计算流程:Attention(Q,K,V) = softmax(QKᵀ / √dk) · V。Q 和 K 做点积得到「相关性分数」,softmax 归一化成权重,再用权重去加权求和 V。本质上就是「根据 query 和 key 的匹配程度,从各个 value 中抽取信息」。

√dk:一个小小的缩放因子,影响巨大

手推的时候我最疑惑的就是分母那个 √dk。为什么不是 dk?为什么不是 1?

假设 Q 和 K 的每个分量独立同分布,均值为 0,方差为 1。点积 QKᵀ 的方差是 dk——dk 越大,点积的数值范围越大。大的点积值进入 softmax 后会挤压到梯度极小的饱和区,导致训练不动。除以 √dk 把方差拉回 1,让 softmax 的输入保持在梯度合理的范围内。不是一个 tricks,而是数学必然

Multi-Head:切分维度,不是并行计算

我之前一直以为 Multi-Head 是为了并行加速——八个头同时算。错。Multi-Head 的核心是把 d_model 切分成 h 份,每份独立做 Attention。比如 d_model=512,8 个头,每个头只在 64 维子空间里做 Attention。

这样做的好处是让不同头关注不同位置的信息——一个头可能专注邻近 token 的语法关系,另一个头可能关注远距离的语义关联。concat 后再投影回原维度,信息就融合了。推荐阅读 Jay Alammar 的 The Illustrated Transformer,图解非常直观。

Positional Encoding:为什么用正余弦

Self-Attention 本身对位置无感知——"我吃苹果"和"苹果吃我"在它看来是一样的。所以必须在输入端加上位置信息。论文用了正弦和余弦函数的组合——不同频率对应不同位置尺度。我自己推了一下:相邻位置的编码相似,远处位置差异大;而且 sin/cos 能外推到训练时没见过的序列长度。这个设计实在太优雅了。

--- 约 720 字