模型基础03:Transformer
从 Attention 机制、位置编码、编码器-解码器架构到上下文建模,理解 Transformer 为什么能成为大模型的基础架构。
Transformer 是大模型的核心架构,理解它才能真正明白大模型为什么能处理长文本、理解上下文。这篇文章从 Attention 机制入手,把 Transformer 的设计原理讲清楚。
Attention 机制:Transformer 的核心
Attention 机制让模型能够关注输入序列中不同位置的信息,这是 Transformer 的灵魂。
为什么需要 Attention
传统的 RNN/LSTM 存在两个问题:
- 顺序依赖:必须按顺序处理输入,无法并行计算
- 长距离依赖:梯度消失,难以捕捉长文本中的依赖关系
Attention 机制解决了这两个问题:
- 并行计算:所有位置同时计算 Attention
- 直接依赖:任意两个位置之间可以直接建立联系
Scaled Dot-Product Attention
Transformer 使用的是 Scaled Dot-Product Attention:
Attention(Q, K, V) = softmax(QKᵀ/√dₖ)V
三个关键向量:
| 向量 | 含义 | 计算方式 |
|---|---|---|
| Q(Query) | 查询向量,当前位置要查询的信息 | Q = XW_q |
| K(Key) | 键向量,所有位置的信息标识 | K = XW_k |
| V(Value) | 值向量,所有位置的实际内容 | V = XW_v |
计算过程:
- 计算相似度:Q 和 K 的点积,表示查询与每个键的匹配程度
- 缩放:除以 √dₖ,防止点积值过大导致 softmax 梯度消失
- Softmax:将相似度转换为权重
- 加权求和:用权重对 V 加权求和,得到最终输出
Multi-Head Attention
Multi-Head Attention 将输入投影到多个子空间,学习不同类型的关系:
MultiHead(Q, K, V) = Concat(head₁, ..., headₕ)Wᵒ
where headᵢ = Attention(QWᵠⁱ, KWᵏⁱ, VWᵛⁱ)
多头的好处:
- 每个头学习不同的注意力模式
- 模型可以同时关注不同类型的关系
- 增加模型的表达能力
三种 Attention 模式
Transformer 中有三种 Attention 模式:
| 模式 | 输入 | 适用场景 |
|---|---|---|
| Self-Attention | Q=K=V,同一序列内部的 Attention | 理解文本内部关系 |
| Encoder-Decoder Attention | Q来自 Decoder,K=V来自 Encoder | 翻译、摘要等生成任务 |
| Causal Attention | 只允许关注前面的位置 | 语言模型生成 |
位置编码:给模型感知顺序
Transformer 没有循环结构,需要位置编码来让模型感知序列顺序。
为什么需要位置编码
Transformer 的输入是并行处理的,模型本身不知道输入的顺序。位置编码给每个位置一个唯一的标识,让模型知道哪个词在前、哪个词在后。
正弦余弦位置编码
Transformer 论文中使用的是正弦余弦位置编码:
PE(pos, 2i) = sin(pos / 10000^(2i/d_model))
PE(pos, 2i+1) = cos(pos / 10000^(2i/d_model))
特点:
- 位置编码是固定的,不参与训练
- 相对位置信息可以通过三角函数的性质推导
- 可以处理任意长度的序列
学习型位置编码
后来的模型(如 BERT)使用学习型位置编码:
PE = learnable_parameters
特点:
- 位置编码是训练出来的,更灵活
- 但只能处理训练时见过的长度
RoPE(旋转位置编码)
LLaMA、GPT 等模型使用 RoPE:
将位置信息编码到 Attention 的计算中,通过旋转矩阵实现
特点:
- 保持相对位置信息
- 支持外推,能处理更长的序列
编码器-解码器架构
Transformer 由编码器和解码器两部分组成。
编码器
编码器负责理解输入序列,由多个相同的层堆叠而成:
Encoder = [EncoderLayer] × N
EncoderLayer 结构:
- Multi-Head Self-Attention:输入序列内部的 Attention
- Add & Norm:残差连接 + Layer Normalization
- Feed-Forward Network:两层全连接网络,中间用 ReLU
残差连接的作用:
- 缓解梯度消失,让深层网络能够训练
- 允许信息直接传递,不被多层变换破坏
Layer Normalization:
- 对每个样本的特征进行标准化
- 稳定训练过程,加速收敛
解码器
解码器负责生成输出序列,也由多个相同的层堆叠而成:
Decoder = [DecoderLayer] × N
DecoderLayer 结构:
- Masked Multi-Head Self-Attention:只能关注前面已生成的位置
- Add & Norm
- Encoder-Decoder Attention:关注编码器的输出
- Add & Norm
- Feed-Forward Network
- Add & Norm
掩码的作用:
- Causal Mask:防止模型看到未来的位置
- Padding Mask:忽略填充的位置
Transformer 的变体
基于 Transformer 架构,出现了多种变体:
Encoder-only 模型
只使用编码器,适合理解任务:
| 模型 | 特点 | 适用场景 |
|---|---|---|
| BERT | 双向 Attention,MLM 预训练 | 分类、问答、抽取 |
| RoBERTa | BERT 的改进版本 | 通用理解任务 |
| ALBERT | 参数共享,更小更快 | 资源受限场景 |
Decoder-only 模型
只使用解码器,适合生成任务:
| 模型 | 特点 | 适用场景 |
|---|---|---|
| GPT | 单向 Attention,自回归生成 | 文本生成、对话 |
| GPT-2 | 更大参数,更好的生成能力 | 通用生成任务 |
| GPT-3 | 175B 参数,强大的 Few-shot 能力 | 通用 AI |
| LLaMA | Meta 开源,高性能 | 开源部署 |
Encoder-Decoder 模型
同时使用编码器和解码器,适合转换任务:
| 模型 | 特点 | 适用场景 |
|---|---|---|
| T5 | 统一框架,所有任务都转为文本生成 | 通用任务 |
| BART | 基于 BERT 的生成模型 | 摘要、翻译 |
| Switch Transformer | MoE 架构,稀疏激活 | 超大规模模型 |
上下文建模:Transformer 的核心能力
Transformer 的核心能力是上下文建模,它能根据上下文理解和生成文本。
上下文窗口
上下文窗口是模型能处理的最大文本长度:
| 模型 | 上下文窗口 | 特点 |
|---|---|---|
| GPT-3 | 2048 tokens | 基础版本 |
| GPT-3.5 | 4096/16384 tokens | 更大的窗口 |
| GPT-4 | 8192/32768 tokens | 长文本理解 |
| LLaMA-2 | 4096/8192/32768 tokens | 可配置 |
| Claude 3 | 200K tokens | 超长上下文 |
上下文长度的影响
- 更长的上下文:能处理更长的文档,但计算成本更高
- 上下文窗口限制:模型无法记住窗口外的信息
- 上下文压缩:需要智能选择重要信息
注意力模式
不同的注意力模式适用于不同任务:
| 模式 | 特点 | 适用场景 |
|---|---|---|
| 全局 Attention | 关注所有位置 | 理解任务 |
| 局部 Attention | 只关注附近位置 | 长文本处理 |
| 稀疏 Attention | 只关注关键位置 | 超大上下文 |
| 滑动窗口 Attention | 窗口滑动处理 | 长文档处理 |
Transformer 的优势与挑战
优势
- 并行计算:比 RNN 快很多
- 长距离依赖:能捕捉长文本中的依赖关系
- 灵活架构:可以组合成多种变体
- 可扩展性:容易扩展到更大参数
挑战
- 计算复杂度:O(n²),长文本时计算量很大
- 内存占用:存储 Attention 矩阵需要大量内存
- 推理速度:生成时需要逐个 token 计算
- 上下文窗口:有最大长度限制
优化方向
| 方向 | 方法 | 效果 |
|---|---|---|
| 稀疏 Attention | 只计算部分 Attention | 减少计算量 |
| Flash Attention | 优化内存访问 | 加速计算 |
| 量化 | 降低权重精度 | 减少内存占用 |
| MoE | 稀疏激活 | 增加模型容量 |
| RAG | 检索增强 | 扩展知识范围 |
项目判断清单
- 需要理解任务(分类、问答)→ 用 Encoder-only 模型(BERT)
- 需要生成任务(写文章、对话)→ 用 Decoder-only 模型(GPT、LLaMA)
- 需要转换任务(翻译、摘要)→ 用 Encoder-Decoder 模型(T5)
- 文本超过上下文窗口 → 截断、分块或用 RAG
- 推理速度慢 → 用更小的模型、量化或缓存
- 需要开源部署 → 用 LLaMA、Mistral 等开源模型
- 需要超长上下文 → 用 Claude 3 或支持长窗口的开源模型