自注意力将输入序列中的每个位置与所有位置关联:
Attention(Q,K,V)=Softmax(dkQKT)V
其中:
- Q=XWQ:查询矩阵
- K=XWK:键矩阵
- V=XWV:值矩阵
- dk:键向量维度
除以 dk 的原因:
当 dk 较大时,点积结果方差增大,Softmax进入饱和区,梯度极小:
Var(q⋅k)=dk⋅Var(qi)⋅Var(ki)
缩放后:Var(dkq⋅k)=Var(qi)⋅Var(ki)
| 操作 | 复杂度 |
|---|
| QKV投影 | O(n⋅d⋅dk) |
| 注意力矩阵 | O(n2⋅dk) |
| 加权求和 | O(n2⋅dv) |
| 总计 | O(n2d+nd2) |
将Q、K、V投影到 h 个子空间,分别计算注意力后拼接:
MultiHead(Q,K,V)=Concat(head1,…,headh)WO
headi=Attention(QWiQ,KWiK,VWiV)
其中 WiQ∈Rd×dk,dk=d/h。
- 每个头关注不同的子空间信息
- 类似于CNN中多个卷积核提取不同特征
- 增强模型的表达能力
由于自注意力是置换不变的,需要位置编码注入位置信息:
PE(pos,2i)=sin(100002i/dpos)
PE(pos,2i+1)=cos(100002i/dpos)
性质:
- 每个维度对应不同频率的正弦波
- 相对位置关系可通过线性变换表达
- 可外推到更长序列
通过旋转矩阵编码相对位置:
qm=RΘ,mWqxm
kn=RΘ,nWkxn
qmTkn=xmTWqTRΘ,n−mWkxn
内积只依赖相对位置 n−m。
直接在注意力分数上添加线性偏置:
Attentionij=dkqiTkj+m⋅∣i−j∣
输入 → [Embedding + Positional Encoding]
→ [Multi-Head Attention] → [Add & Norm]
→ [Feed-Forward Network] → [Add & Norm]
→ ... (×N层)
子层结构:
- 多头自注意力 + 残差连接 + LayerNorm
- 前馈网络(FFN)+ 残差连接 + LayerNorm
FFN:
FFN(x)=max(0,xW1+b1)W2+b2
扩展比:dff=4dmodel
目标 → [Embedding + Positional Encoding]
→ [Masked Multi-Head Attention] → [Add & Norm]
→ [Cross-Attention] → [Add & Norm]
→ [Feed-Forward Network] → [Add & Norm]
→ [Linear + Softmax]
→ ... (×N层)
Masked Attention:防止看到未来信息
Maskij={0−∞i≥ji<j
| 方式 | 公式 | 训练稳定性 |
|---|
| Post-Norm | LayerNorm(x+Sublayer(x)) | 较差 |
| Pre-Norm | x+Sublayer(LayerNorm(x)) | 较好 |
现代Transformer普遍采用Pre-Norm。
| 方法 | 复杂度 | 思路 |
|---|
| Longformer | O(n) | 局部窗口+全局token |
| BigBird | O(n) | 随机+窗口+全局 |
| Sparse Transformer | O(nn) | 固定稀疏模式 |
Attention(Q,K,V)=ϕ(Q)(ϕ(K)T1)ϕ(Q)(ϕ(K)TV)
先计算 ϕ(K)TV(d×d),复杂度降为 O(nd2)。
通过分块计算和IO感知优化,减少HBM访问次数:
- 数学等价:结果与标准注意力完全一致
- 内存优化:O(n) 而非 O(n2)
- 速度提升:2~4x