自注意力让每个 token 直接关注序列中所有 token,动态计算彼此的关联权重,从而建模长程依赖。具体做法:输入向量分别经三个线性变换得到查询 Q、键 K、值 V;用 Q 与所有 K 的点积衡量相关性,除以 √d_k 缩放后过 softmax 得到注意力权重,再对 V 加权求和得到输出。公式:Attention(Q,K,V) = softmax(QKᵀ/√d_k)V。
直觉理解:Q 是当前 token 提出的问题,K 是各 token 的索引标签,V 是真正被聚合的内容。相比 RNN 的逐步传递,自注意力任意两个位置之间路径长度为 1,且整层计算可完全并行,这是它取代 RNN 的根本原因;代价是计算和显存随序列长度平方增长。
易错点:混淆自注意力与交叉注意力(后者 Q 来自解码器、KV 来自编码器);忽略 mask 的作用。追问方向:注意力复杂度为什么是 O(n²),FlashAttention 如何优化?为什么需要多头而不是单头?