GPT-2 架构拆解与理解
TL;DR: GPT-2 的核心目标,是把一段上文压缩成一个高维表示,再让这个表示在词表空间里靠近下一个正确 token 的嵌入。嵌入层提供初始语义和位置,注意力负责跨 token 汇聚上下文,FFN 负责逐 token 加工概念,残差和 LayerNorm 则让 12 层堆叠保持可训练。
0. 一句话总结 Transformer 在做什么
把整段上文压缩成一个高维向量,让这个向量在词表空间中恰好落在”正确答案”那个 token 的向量附近。
所有 124M 参数的梯度更新,都服务于这一个目标。
1. 核心超参数(GPT-2 Small)
V = 50257:词表大小d = 768:每个 token 的表示向量是 768 维L = 12:12 层反复加工h = 12:每层注意力有 12 个并行专家(头)d_k = 64:每个专家看 64 维(768 ÷ 12)d_ff = 3072:FFN 中间层膨胀到 3072 维(4 倍)ctx = 1024:最多一次看 1024 个 token
2. 输入层:两块地基
Token Embedding(W_E)
一张大表,每一行是一个 token 的”身份证向量”。输入整数 → 查表 → 768 维向量。五万多个 token,每人一张 768 分的评分卡。
Position Embedding(W_P)
如果只看 token 向量,“The cat sat” 和 “sat cat The” 在模型眼里是三组完全相同的数字。所以再建一张表,每个位置有一个独立的”座位向量”。位置 0 有一个 768 维向量表示”我是第一个”,位置 1 有另一个。
嵌入求和
Token 向量 + 位置向量 → 逐元素相加。为什么用加法而非拼接?拼接使维度翻倍计算量翻倍,实验证明加法已足够——模型通过训练学会区分哪些维度来自语义、哪些来自位置。
3. LayerNorm:数据标准化车间
每层有两次 LayerNorm(注意力前一次、FFN 前一次)。每个 LayerNorm 有两个可学习的向量:
- γ(gamma):缩放系数——归一化后乘多少
- β(beta):平移系数——归一化后加多少
操作:对一个 token 的 768 个数值,减去均值、除以标准差(强制复位为均值 0 标准差 1),然后用 γ 和 β 调到”下一层需要的尺度”。
为什么每层有独立的 LayerNorm? 浅层处理的是词嵌入(原始),深层处理的是抽象语义——不同工序需要不同的数据尺度。
4. 多头自注意力(MHSA):让 token 互相看
为什么需要注意力?
经过嵌入层后,每个 token 只知道自己的信息——“cat”不知道前面有”The”,“sat”不知道前面有”cat”。注意力让每个 token 从序列中其他 token 提取信息,但要有权重——关系大的 token 权重大,关系小的权重小。
Q、K、V 三个投影
输入经过三个矩阵乘法,变成三个新向量:
- Q(Query):“我是谁,我想找什么?“——当前 token 发出的搜索需求
- K(Key):“我是谁,我能提供什么信息?“——每个 token 提供的索引标签
- V(Value):“如果别人关注我,我应该输出什么?“——被注意时传递的实际内容
拆分成多头
12 个头不是 12 个注意力层串联,而是把 768 维切成 12 块,每块 64 维,12 个头并行计算。
Scaled Dot-Product Attention(核心公式)
对每个头:
- 算分数:
Q 和 K 做点积——“我想找的东西”和”你能提供的东西”是否匹配 - 除以 √d_k:64 维点积的方差 = 64,标准差 = 8。不除 8 → 分数太大 → softmax 极端 → 梯度消失
- 加因果掩码:GPT-2 是自回归模型,token
t不能偷看j > t的 token。未来位置分数设为 -∞ - softmax:分数 → 概率(每行和为 1)
- 加权求和 V:注意力权重 × V = 从各 token 提取信息
W_O:融合 12 个专家的意见
12 个头各自输出 64 维 → 拼成 768 维 → 乘以 W_O(另一个 768×768 矩阵)。
W_O 和 V 的关系
W_O 不直接作用于 V——它作用于注意力加权后的结果。V 是原始信息库,softmax(QK^T) 是从信息库中筛选要取多少,W_O 是把筛选出的摘录整合成报告。
残差连接 1
H_mid = H_in + A(原始输入 + 注意力的输出)
残差为什么重要? 无残差→12 层深网络梯度消失。残差提供”高速公路”——梯度可直接从高层跳到底层,不经过中间计算。
5. Feed-Forward Network(FFN):每个 token 单独消化
注意力让 token 之间交换了信息,但交换来的信息还需要加工。FFN 和注意力完全不同——注意力是跨 token 的操作(会议讨论),FFN 是每个 token 内部的操作(各自回位置消化)。
三拍子:展开 → 筛选 → 回收
768 维 → W₁ 升维到 3072 → GELU 筛选 → W₂ 降维回 768- 升维(W₁:768→3072):把紧凑的表示”展开”——原本混在一起的多个概念,在 3072 维空间中可以各占各的维度不干扰
- 筛选(GELU):对每个概念通道独立做门控——正激活放行(保留),负激活屏蔽(接近 0)
- 降维(W₂:3072→768):把幸存的概念压缩回来,作为对原始表示的补充
残差连接 2
同注意力一样:H_out = H_mid + F。FFN 学的是对当前表示的修正,不是从零重建。
6. 形式化推导:一个 token 向量的一生
以纯数学语言完整推导前向、损失与生成流程。所有步骤显式写出可训练的权重矩阵和偏置,模型的最小单元就是这些可训练参数。
6.1 符号约定
| 符号 | 含义 | GPT-2 Small 取值 |
|---|---|---|
V | 词汇表大小 | 50257 |
T_max | 最大序列长度 | 1024 |
L | 层数 | 12 |
D | 模型维度 | 768 |
D_ff | 前馈中间维度 | 3072 |
H | 注意力头数 | 12 |
d_k | 每头维度 D / H | 64 |
注视一个长度为 T 的输入序列:x = (x_1, x_2, ..., x_T),每个 x_t ∈ {0, 1, ..., V-1}。
主人公是第 t 个 token x_t,它的向量表示将经历完整一生,最终用于预测 x_{t+1}。
6.2 诞生:嵌入
可训练参数:词嵌入 W_E ∈ R^{V×D},位置嵌入 W_P ∈ R^{T_max×D}。
Token x_t 的初始表示(第 0 层输出):
逐元素相加。同一词坐不同位置 → 加上不同的位置偏移 → 不同表示。
6.3 成长:穿越 L 层 Transformer Block
对每一层 l = 1, 2, ..., L,向量 h^{(l-1)}_t 依次经历两个子层,每个子层使用 Pre-LN 残差结构。
2.1 第一子层:多头因果自注意力
2.1.1 Pre-LayerNorm
可训练参数:γ₁^{(l)} ∈ R^D,β₁^{(l)} ∈ R^D。
其中:
2.1.2 Q / K / V 投影
可训练参数:W_Q^{(l)}, W_K^{(l)}, W_V^{(l)} ∈ R^{D×D},偏置 b_Q^{(l)}, b_K^{(l)}, b_V^{(l)} ∈ R^D。
对整个序列所有位置 t' ∈ [1, T] 计算:
2.1.3 多头拆分
将 D 维向量均分成 H 个头,每个头维度为 d_k。对头 h ∈ {1, ..., H},位置 t' 的查询、键、值向量为:
2.1.4 因果掩码注意力(主人公 t 的视角)
主人公位置 t 只能注意到 t' ≤ t 的 token。
计算头 h 下,t 对 t' 的未归一化注意力分数:
掩码确保 t' > t 的分数为 -∞。随后 softmax:
加权求和得到该头的输出:
2.1.5 合并多头与输出投影
可训练参数:W_O^{(l)} ∈ R^{D×D},b_O^{(l)} ∈ R^D。
2.1.6 残差连接
2.2 第二子层:前馈网络
2.2.1 Pre-LayerNorm
可训练参数:γ₂^{(l)}, β₂^{(l)} ∈ R^D。
2.2.2 升维 → GELU 筛选 → 降维
可训练参数:W₁^{(l)} ∈ R^{D×D_ff},b₁^{(l)} ∈ R^{D_ff},W₂^{(l)} ∈ R^{D_ff×D},b₂^{(l)} ∈ R^D。
GELU 近似公式:
2.2.3 残差连接
6.4 使命:最终输出与预测
经过全部 L 层后,得到最终表示。最后做一次 LayerNorm:
可训练参数:γ_f, β_f ∈ R^D。
输出投影(Weight Tying — 通常无独立偏置):
可训练参数:W_{lm} ∈ R^{V×D}(实际实现中 W_{lm} = W_E)。
z_t和W_{lm}的每一行(每个候选 token 的嵌入)做点积。点积越大 → 两个向量越接近 → 该 token 概率越高。
6.5 学习:损失函数
训练时一次性计算所有位置的 logits,第 t 个位置预测 x_{t+1}:
即交叉熵——最大化正确 token 的概率,等价于让 z_t 和 W_E[x_{t+1}] 尽可能靠近。反向传播计算损失对每一个可训练参数的梯度,用 AdamW 等优化器更新。
6.6 应用:自回归生成
- 给定前缀
x_1, ..., x_t,算出logits_t - 温度调节:
logits'_t = logits_t / τ(τ 越小越贪婪,τ 越大越随机) - 采样:
x_{t+1} ∼ softmax(logits'_t) - 将
x_{t+1}拼接到序列末尾,重复直到终止符或最大长度
注意:生成第 t+1 个 token 时,所有 ≤t 的 K/V 可缓存(KV-Cache)避免重复计算,但数学本质不变。
6.7 可训练参数全家福
| 模块 | 参数 | 形状 | 数量 |
|---|---|---|---|
| 嵌入 | W_E | V × D | 1 |
W_P | T_max × D | 1 | |
| 每层注意力 | γ₁, β₁ | D | 2 |
W_Q, W_K, W_V | D × D | 3 | |
b_Q, b_K, b_V | D | 3 | |
W_O | D × D | 1 | |
b_O | D | 1 | |
| 每层 FFN | γ₂, β₂ | D | 2 |
W₁ | D × D_ff | 1 | |
b₁ | D_ff | 1 | |
W₂ | D_ff × D | 1 | |
b₂ | D | 1 | |
| 最终 LN | γ_f, β_f | D | 2 |
| 输出头 | W_{lm} | V × D | 1 |
所有这些参数构成了可训练的 GPT-2。每一个正向传递都沿着上述数学路径一步一步塑造 token 向量的生命轨迹。每一条梯度都沿着完全相同的路径反向传播,逐层修正这些参数,直到模型学会在 768 维空间中让预测向量落在正确答案的嵌入向量附近。
7. 为什么需要 12 层?
- 第 1 层:cat 看到 The(距离 1)
- 第 2 层:sat 看到 The+cat(距离 2)
- …
- 第 12 层:第 1024 个 token 可以看到第 1 个
更深层的另一个好处是抽象层级:浅层看语法(局部短语),中层看语义(实体指代),深层看推理(篇章连贯性、世界知识)。
8. 输出层
12 层处理完 → 最终 LayerNorm → @ W_E^T(输出投影)→ softmax → 取概率最高的 token。
Weight Tying:输出投影不设独立矩阵,直接用输入 Token Embedding 的转置。H_final @ W_E^T 等价于”把最后位置的隐状态,和五万个候选 token 的嵌入向量逐一比较相似度”。最像的那个就是预测的下一个 token。
9. FFN 的可解释性(为什么 FFN 占 70% 参数但少有人提?)
核心发现:FFN 是一个键值记忆网络
FFN 的 3072 个维度各是一个”概念槽位”,每个槽位有两样东西:
- Key(W₁ 的一行):输入长什么样时,激活这个槽位?
- Value(W₂ 的一列):激活后,往输出里加什么信息?
前向传播就是:逐一检查 3072 个概念 → 匹配的用 GELU 决定强度 → value 加权累加。
实验证据
- 把 W₂ 的列投影到词表:某列的 top token 是
Paris, France, Lyon→ 对应”法国”概念;某列是DNA, protein, gene→ 对应”分子生物学” - ROME 实验:修改 FFN 中几个特定神经元 → 模型从认为”埃菲尔铁塔在巴黎”变成”在罗马”。证明事实知识以高度局部化的方式存储在 FFN 中
为什么 FFN 少被提及?
注意力是 2017 年的革命性发明(可以可视化热力图),FFN 是 1986 年的 MLP(一张无法可视化的数字表)。但参数量不说谎——70% 参数在 FFN,因为它存储了模型知道的几乎所有事实、语法、常识和推理模式。注意力是名片,FFN 是肌肉。
10. GPT-2 → 现代大模型架构演进
FFN:GELU → SwiGLU
GPT-2 用 GELU 做隐式筛选。LLaMA 用 SwiGLU——把信息分两路:一路学门控(哪些信息放行),一路学内容(放行什么信息),两路逐元素相乘。门控 + 内容的显式分离比 GELU 的隐式筛选更强。
bias:有 → 无
LayerNorm 的 β 已经能替代 bias 的大部分功能。去掉 bias 简化实现、不减性能,还省去了 weight decay 时对 bias 该不该加正则化的纠结。
位置编码:可学习表 → RoPE
LLaMA 用旋转位置编码(Rotary Position Embedding),天然支持任意长度外推——训练时用 2K 上下文,推理时可扩展到 32K 甚至更长。
11. GQA:Grouped-Query Attention
现代大模型(LLaMA-2/3、Mistral、Qwen)的标配。
演变:MHA(每个头独立 K、V)→ MQA(所有头共享一套 K、V,太激进)→ GQA(折中:分若干组,组内共享 K、V)。
为什么 K、V 可以共享但 Q 不行? Q 是”我在找什么”——每个头在找不同的东西,多样性来自 Q。K 是”我提供什么标签”、V 是”我的内容”——同一个 token 的索引和内容只有一份,不需要重复存储。GQA 让 KV Cache 缩小若干倍,困惑度几乎不变。
12. FlashAttention 与 PagedAttention
两者优化的是完全不同的问题,正交使用。
FlashAttention:计算的 IO 优化
原版 Attention 每一步都把整个 N×N 注意力矩阵在显存(HBM)和计算单元间搬运。FlashAttention 采用切块策略——把 Q、K、V 切成小块,在片上高速缓存(SRAM)内算完就扔,不写回显存。同时用 Online Softmax(不需要看到整行就能算 softmax 的数学 trick)让切块计算可行。
效果:HBM 读写降为原来的 1/N,显存省 90%+,速度提升数倍。
PagedAttention:KV Cache 的内存管理
多请求并发推理时,每个请求的 KV Cache 需要预分配最大长度连续显存 → 碎片严重。PagedAttention 把 KV Cache 切成页(类似操作系统的虚拟内存分页),按需分配不预占,不要求连续,用页表索引。
效果:显存利用率从 <40% 提升到 >95%。
13. 预训练 / SFT / RLHF
我们讨论的所有架构细节,三阶段完全相同,变的是数据和 loss:
- 预训练:万亿 token 互联网文本 →“下一个词预测”的交叉熵 loss → 得到一个”会接话”的基础模型(99.9% 计算量)
- SFT:十万条人工标注问答对 → 同样的交叉熵 loss(user 部分不算)→ 变成”会对话”(0.05% 计算量)
- RLHF:人类偏好对比数据 → 偏好 loss / PPO → 变成”回答让人满意且对齐人类价值观”的最终模型(0.05% 计算量)
14. 完整前向传播(端到端流程)
输入 token IDs → W_E 查身份向量 + W_P 查座位向量 → 相加 = 初始表示
12 次循环: LayerNorm(洗数据)→ QKV投影 → 拆12头 → 每头: QK^T/√64 + 因果掩码 → softmax → ×V → 拼回头 × W_O(融合专家意见) → + 残差(存量 + 上下文增量) → LayerNorm(再洗)→ W₁ 升维 → GELU 筛选 → W₂ 降维 → + 残差(存量 + 消化增量)
最终 LayerNorm → @ W_E^T(和所有 token 比点积)→ softmax → 取概率最高最精妙的悖论:这 300 行代码的全部逻辑,论文中 FFN 部分只给了 4 行描述——一个拥有上万个概念槽位、可定位、可编辑的知识库,在 2017 年只是 “two linear transformations with a ReLU activation”。发明者不一定是理解者。
部分信息可能已经过时








