跳到主要内容

Transformer 与 Self-Attention 原理

🧠 Transformer 用注意力机制替代循环,让每个 token 直接与所有其他 token 交互,从而并行处理序列并捕获长距离依赖。

核心问题

RNN 的两大痛点:① 序列依赖无法并行;② 长距离梯度消失。Transformer 用自注意力同时解决两者。


Self-Attention 机制

QKV 计算

每个 token 的 Embedding 经过三个线性变换生成 Q(Query)、K(Key)、V(Value):

Attention(Q, K, V) = softmax(QK^T / sqrt(d_k)) * V
  • Q:当前 token 想要查询什么
  • K:每个 token 能提供的信息标签
  • V:每个 token 实际携带的信息
  • sqrt(d_k):防止点积过大导致 softmax 梯度消失

直觉

每个 token 向所有其他 token "提问"(Q × K),根据相关程度加权聚合所有 token 的信息(× V)。


Multi-Head Attention

MultiHead(Q, K, V) = Concat(head_1, ..., head_h) * W_O

多个注意力头并行运行,每个头关注不同的语义维度(如句法关系、语义相似、指代关系)。


Transformer 架构

Encoder 层(每层包含)

  1. Multi-Head Self-Attention
  2. Add & LayerNorm(残差连接)
  3. Feed-Forward Network(两层 MLP + ReLU)
  4. Add & LayerNorm

Decoder 层(额外增加)

  • Masked Self-Attention:只看历史 token,防止看到未来
  • Cross-Attention:Q 来自 Decoder,K/V 来自 Encoder 输出

位置编码

Attention 本身无序列顺序感知,需要注入位置信息:

类型方法特点
正弦位置编码固定公式 sin/cos可外推,无需训练
可学习位置嵌入训练参数灵活,不可外推
RoPE旋转矩阵相对位置,LLM 常用
ALiBi注意力偏置线性衰减,外推性好

注意力复杂度

指标复杂度
计算复杂度O(n² · d)
内存复杂度O(n²)
最大路径长O(1)

n 为序列长度——这是长上下文的瓶颈,FlashAttention 通过分块计算优化内存。


现代 LLM 的改进

  • GQA(Grouped Query Attention):多个 Q 共享一组 K/V,降低 KV Cache 显存
  • Flash Attention:分块计算,IO 复杂度从 O(n²) 降至 O(n)
  • Sliding Window Attention:每个 token 只看局部窗口,适合超长文本

常见误区

  • Transformer 不是 "Attention 替代 Embedding",Embedding 仍然存在
  • LayerNorm 在现代 LLM 中通常用 Pre-Norm(在 Attention 前),而非 Post-Norm
  • Decoder-Only 模型(GPT)只有 Masked Self-Attention,没有 Cross-Attention