Transformer架构

8 minAdvanced2026/6/14

自注意力机制、多头注意力、位置编码与Transformer架构详解。

1. 自注意力机制

1.1 核心公式

自注意力将输入序列中的每个位置与所有位置关联:

Attention(Q,K,V)=Softmax(QKTdk)V\text{Attention}(\mathbf{Q}, \mathbf{K}, \mathbf{V}) = \text{Softmax}\left(\frac{\mathbf{Q}\mathbf{K}^T}{\sqrt{d_k}}\right)\mathbf{V}

其中:

  • Q=XWQ\mathbf{Q} = \mathbf{X}\mathbf{W}_Q:查询矩阵
  • K=XWK\mathbf{K} = \mathbf{X}\mathbf{W}_K:键矩阵
  • V=XWV\mathbf{V} = \mathbf{X}\mathbf{W}_V:值矩阵
  • dkd_k:键向量维度

1.2 缩放因子

除以 dk\sqrt{d_k} 的原因:

dkd_k 较大时,点积结果方差增大,Softmax进入饱和区,梯度极小:

Var(qk)=dkVar(qi)Var(ki)\text{Var}(q \cdot k) = d_k \cdot \text{Var}(q_i) \cdot \text{Var}(k_i)

缩放后:Var(qkdk)=Var(qi)Var(ki)\text{Var}\left(\frac{q \cdot k}{\sqrt{d_k}}\right) = \text{Var}(q_i) \cdot \text{Var}(k_i)

1.3 计算复杂度

操作复杂度
QKV投影O(nddk)O(n \cdot d \cdot d_k)
注意力矩阵O(n2dk)O(n^2 \cdot d_k)
加权求和O(n2dv)O(n^2 \cdot d_v)
总计O(n2d+nd2)O(n^2 d + nd^2)

2. 多头注意力

2.1 多头机制

将Q、K、V投影到 hh 个子空间,分别计算注意力后拼接:

MultiHead(Q,K,V)=Concat(head1,,headh)WO\text{MultiHead}(\mathbf{Q}, \mathbf{K}, \mathbf{V}) = \text{Concat}(\text{head}_1, \ldots, \text{head}_h)\mathbf{W}^O

headi=Attention(QWiQ,KWiK,VWiV)\text{head}_i = \text{Attention}(\mathbf{Q}\mathbf{W}_i^Q, \mathbf{K}\mathbf{W}_i^K, \mathbf{V}\mathbf{W}_i^V)

其中 WiQRd×dk\mathbf{W}_i^Q \in \mathbb{R}^{d \times d_k}dk=d/hd_k = d/h

2.2 多头的意义

  • 每个头关注不同的子空间信息
  • 似于CNN中多个卷积核提取不同特征
  • 增强模型的表达能力

3. 位置编码

3.1 正弦位置编码

由于自注意力是置换不变的,需要位置编码注入位置信息:

PE(pos,2i)=sin(pos100002i/d)PE_{(pos, 2i)} = \sin\left(\frac{pos}{10000^{2i/d}}\right) PE(pos,2i+1)=cos(pos100002i/d)PE_{(pos, 2i+1)} = \cos\left(\frac{pos}{10000^{2i/d}}\right)

性质

  • 每个维度对应不同频率的正弦波
  • 相对位置关系可通过线性变换表达
  • 可外推到更长序列

3.2 旋转位置编码(RoPE)

通过旋转矩阵编码相对位置

qm=RΘ,mWqxm\mathbf{q}_m = \mathbf{R}_{\Theta,m}\mathbf{W}_q\mathbf{x}_m kn=RΘ,nWkxn\mathbf{k}_n = \mathbf{R}_{\Theta,n}\mathbf{W}_k\mathbf{x}_n

qmTkn=xmTWqTRΘ,nmWkxn\mathbf{q}_m^T\mathbf{k}_n = \mathbf{x}_m^T\mathbf{W}_q^T\mathbf{R}_{\Theta,n-m}\mathbf{W}_k\mathbf{x}_n

内积只依赖相对位置 nmn - m

3.3 ALiBi位置编码

直接在注意力分数上添加线性偏置:

Attentionij=qiTkjdk+mij\text{Attention}_{ij} = \frac{\mathbf{q}_i^T\mathbf{k}_j}{\sqrt{d_k}} + m \cdot |i - j|

  • 无需位置嵌入
  • 支持长度外推

4. Transformer架构

4.1 编码器

输入 → [Embedding + Positional Encoding]
  → [Multi-Head Attention] → [Add & Norm]
  → [Feed-Forward Network] → [Add & Norm]
  → ... (×N层)

子层结构

  1. 多头自注意力 + 残差连接 + LayerNorm
  2. 前馈网络(FFN)+ 残差连接 + LayerNorm

FFN

FFN(x)=max(0,xW1+b1)W2+b2\text{FFN}(\mathbf{x}) = \max(0, \mathbf{x}\mathbf{W}_1 + \mathbf{b}_1)\mathbf{W}_2 + \mathbf{b}_2

扩展比:dff=4dmodeld_{ff} = 4d_{model}

4.2 解码器

目标 → [Embedding + Positional Encoding]
  → [Masked Multi-Head Attention] → [Add & Norm]
  → [Cross-Attention] → [Add & Norm]
  → [Feed-Forward Network] → [Add & Norm]
  → [Linear + Softmax]
  → ... (×N层)

Masked Attention:防止看到未来信息

Maskij={0iji<j\text{Mask}_{ij} = \begin{cases} 0 & i \geq j \\ -\infty & i < j \end{cases}

4.3 Pre-Norm vs Post-Norm

方式公式训练稳定性
Post-NormLayerNorm(x+Sublayer(x))\text{LayerNorm}(\mathbf{x} + \text{Sublayer}(\mathbf{x}))较差
Pre-Normx+Sublayer(LayerNorm(x))\mathbf{x} + \text{Sublayer}(\text{LayerNorm}(\mathbf{x}))较好

现代Transformer普遍采用Pre-Norm。

5. 高效注意力

5.1 稀疏注意力

方法复杂度思路
LongformerO(n)O(n)局部窗口+全局token
BigBirdO(n)O(n)随机+窗口+全局
Sparse TransformerO(nn)O(n\sqrt{n})固定稀疏模式

5.2 线性注意力

Attention(Q,K,V)=ϕ(Q)(ϕ(K)TV)ϕ(Q)(ϕ(K)T1)\text{Attention}(\mathbf{Q}, \mathbf{K}, \mathbf{V}) = \frac{\phi(\mathbf{Q})(\phi(\mathbf{K})^T\mathbf{V})}{\phi(\mathbf{Q})(\phi(\mathbf{K})^T\mathbf{1})}

先计算 ϕ(K)TV\phi(\mathbf{K})^T\mathbf{V}d×dd \times d),复杂度降为 O(nd2)O(nd^2)

5.3 Flash Attention

通过分块计算IO感知优化,减少HBM访问次数:

  • 数学等价:结果与标准注意力完全一致
  • 内存优化:O(n)O(n) 而非 O(n2)O(n^2)
  • 速度提升:2~4x