SECTOR 02 / KNOWLEDGE EXPEDITION

Transformer 架构详解

深入理解 Transformer 架构:自注意力机制、多头注意力、位置编码等核心组件的数学原理与工程实现

01

Transformer 整体架构

EXPLORE +

为什么是 Transformer?

2017 年 Google 在《Attention Is All You Need》中提出 Transformer,它用自注意力机制替代了传统的 RNN/LSTM,解决了长序列建模和并行化训练两大难题。

架构概览

Transformer 采用 Encoder-Decoder 结构(现代 LLM 多只用 Decoder):

  • Encoder:将输入序列编码为上下文表示
  • Decoder:基于编码表示自回归生成输出

核心组件

组件作用
Self-Attention让每个位置关注序列中所有位置,捕获全局依赖
Multi-Head Attention从多个角度计算注意力,捕获不同类型的关系
Position Encoding为模型提供位置信息(Attention 本身不感知顺序)
Feed Forward Network非线性变换增强表达能力(两层线性 + ReLU)
Layer Normalization稳定训练,加速收敛(在序列维度做归一化)
Residual Connection缓解深层网络梯度消失(x + F(x))

数据流

输入序列 -→ Embedding + Position Encoding
-→ [Multi-Head Attention + Residual + LayerNorm]
-→ [FFN + Residual + LayerNorm]
-→ × N 层
-→ 输出表示
02

自注意力机制(Self-Attention)

EXPLORE +

核心思想

Self-Attention 让输入序列中的每个位置都与其他所有位置建立直接联系,捕获全局依赖关系。与 RNN 的逐步传递不同,Self-Attention 一步到位建立所有位置间的连接。

QKV 机制详解

# 每个 Token 产生三个向量
Q (Query) = X × W_Q # 查询:我在找什么?
K (Key) = X × W_K # 键:我有什么可匹配的?
V (Value) = X × W_V # 值:我的内容是什么?

# 注意力计算(核心公式)
Attention(Q, K, V) = softmax(QK^T / √d_k) × V

# 其中 d_k 是 Key 的维度,除以 √d_k 是为了缩放
# 防止点积结果过大导致 softmax 梯度消失

计算步骤(矩阵形式)

  1. 通过线性变换生成 Q, K, V 三个矩阵(形状: n×d_k, n×d_k, n×d_v)
  2. 计算 Q 与 K 的点积相似度,得到 n×n 的注意力分数矩阵
  3. 除以 √d_k 缩放(防止方差过大导致 softmax 饱和)
  4. 对每行做 Softmax 归一化,得到注意力权重(每行和为 1)
  5. 用权重对 V 加权求和,得到输出(n×d_v)

直观理解

类比搜索引擎:Query = 搜索词,Key = 网页标题,Value = 网页内容。系统先匹配 Q-K 相似度,再返回最相关的 V(加权求和)。

时间复杂度

Self-Attention 的时间复杂度为 O(n²·d),其中 n 是序列长度,d 是维度。这解释了为什么长上下文计算的成本会平方级增长。

03

注意力公式推导与数学原理

EXPLORE +

为什么需要缩放?

假设 q 和 k 是 d_k 维的独立随机向量,每个分量的均值为 0、方差为 1,则点积 q·k 的均值为 0、方差为 d_k。方差越大,softmax 的输入值就越大,导致梯度趋近于 0。

# 证明(简化):
q · k = Σ_{i=1}^{d_k} q_i · k_i
Var(q_i · k_i) = 1(假设 q_i, k_i ~ N(0,1) 且独立)
Var(q · k) = Σ Var(q_i · k_i) = d_k
Std(q · k) = √d_k

# 因此除以 √d_k 后方差变为 1

Softmax 的梯度性质

当输入很大时,softmax 输出接近 one-hot,梯度趋近于 0(饱和区)。缩放保证了输入值落入 softmax 的非饱和区,维持梯度流动。

掩码(Masking)

Decoder 中的因果掩码(Causal Mask)确保位置 i 只能关注位置 1 到 i(不能「看到」未来):

# 因果掩码矩阵(上三角设为 -∞)
Mask = [[0, -∞, -∞, -∞],
[0, 0, -∞, -∞],
[0, 0, 0, -∞],
[0, 0, 0, 0]]

# 掩码加到注意力分数上:
Attention = softmax(QK^T/√d_k + Mask) × V
04

多头注意力与位置编码(RoPE)

EXPLORE +

多头注意力(Multi-Head Attention)

不是只做一次注意力计算,而是并行做多次,每次关注不同的方面:

MultiHead(Q, K, V) = Concat(head_1, ..., head_h) × W_O
head_i = Attention(Q × W_{Qi}, K × W_{Ki}, V × W_{Vi})

# 其中每个 head_i 的维度为 d_k/h
# 总计算量与单头注意力基本相同(分而治之)

不同「头」可能学会关注:语法结构、指代关系、语义相似性、长距离依赖、实体关系等。

位置编码

Attention 本身不感知顺序,需要额外注入位置信息:

  • 正弦位置编码(原论文):PE(pos,2i) = sin(pos/10000^{2i/d}), PE(pos,2i+1) = cos(pos/10000^{2i/d})
  • 可学习位置编码:让模型自己学位置表示(BERT 使用)
  • RoPE(旋转位置编码):通过旋转矩阵编码相对位置,当前主流方案(Claude/GPT-4/LLaMA)

RoPE 原理

RoPE 的核心思想是对 Q 和 K 向量施加与位置相关的旋转:

# RoPE 在二维子空间上的旋转
f(q, m) = (q_1 cos mθ - q_2 sin mθ, q_1 sin mθ + q_2 cos mθ)
# m 是位置,θ 是预定义的旋转角
# 内积会自然地编码相对位置信息:
= · cos((m-n)θ)
05

KV-Cache 与推理优化

EXPLORE +

什么是 KV-Cache?

在自回归生成中,每步生成一个新 token,但前面的 token 已经计算过的 K 和 V 矩阵可以缓存复用,避免重复计算。

# 无 KV-Cache(每一步重算所有)
step 1: 计算 token_1 的 K,V
step 2: 计算 token_1, token_2 的 K,V ← token_1 重复计算!
step 3: 计算 token_1, token_2, token_3 的 K,V ← 越来越慢

# 有 KV-Cache
step 1: 计算并缓存 token_1 的 K1,V1
step 2: 只计算 token_2 的 K2,V2,拼接缓存 [K1|K2, V1|V2]
step 3: 只计算 token_3 的 K3,V3,拼接缓存
# 每步计算量恒定 O(1),而非 O(n)

KV-Cache 的内存消耗

KV-Cache 的内存消耗非常大:

# 估算:一个 70B 模型,h=80 层,d=8192,b=1
KV_cache_size = 2(K和V)× h × n × d × bytes_per_element
# 对于 4096 tokens:约 80 × 4096 × 8192 × 2 × 2 bytes ≈ 10 GB

推理优化技术

  • Grouped Query Attention(GQA):多个 Query 头共享一组 Key/Value 头,减少 KV-Cache 占用(LLaMA 2/3 使用)
  • Multi-Query Attention(MQA):所有 Query 头只用一个 Key/Value 头,极致减少缓存(PaLM 使用)
  • PagedAttention:vLLM 的核心技术,像虚拟内存一样管理 KV-Cache,减少碎片
  • 推测解码(Speculative Decoding):用小模型先生成草稿,大模型验证,加速 2-3 倍
06

FlashAttention 原理

EXPLORE +

FlashAttention 的核心问题

标准注意力计算在 GPU 上存在显存带宽瓶颈:

# 标准实现
S = QK^T # (n×d) × (d×n) = n×n
P = softmax(S) # n×n ← 写入 HBM(高带宽显存)
O = P × V # n×n × n×d = n×d
# 两次 HBM 读写,S 和 P 都是 n×n 矩阵

FlashAttention 的解决思路

FlashAttention 通过分块计算(Tiling)内核融合避免显式实例化 n×n 的注意力矩阵:

  • 将 Q, K, V 分块,每个块完全在 SRAM(片上高速缓存)中计算
  • 重新计算部分注意力(用计算换带宽,在 SRAM 上算比从 HBM 读更快)
  • 从 O(n²) 显存需求降低到 O(n)

性能提升

指标标准 AttentionFlashAttention
显存占用O(n²)O(n)
速度(GPT-3 训练)基线快 2-4 倍
HBM 读写大量减少到 1/10
精度标准完全等价(数学等值)

FlashAttention-2

  • 减少非矩阵运算(重计算 mask 等)
  • 更好的块调度策略
  • 在 A100/H100 上达到理论峰值 FLOPS 的 70-80%