之前写从零训练 Transformer时,我更关注完整工程链路:数据怎么准备,Tokenizer 怎么选,模型如何保存,以及低 Loss 为什么不等于真实能力。这次读 Karpathy 的 MicroGPT,我想再往下一层。
如果暂时拿掉 MLX、PyTorch、张量算子和 GPU,GPT 的训练究竟还剩下什么?
我的答案是:一个给下一个 token 分配概率的函数,一张记录计算依赖的图,以及一套根据误差修改参数的规则。 文本生成听起来很复杂,但在这个尺度上,每一次乘法如何影响最终概率,都可以追到具体的数字。
本文以 MicroGPT 源码快照 14fb038 为阅读基线。下面的推导、数值例子和配图是我的拆解;数值例子用于解释计算,不是训练成绩。我也不会把一个名字生成器的行为写成通用语言能力。
先确定它到底在学什么
假设一条训练数据是 emma。模型学习的是一串条件概率:
P(emma, END | START)
= P(e | START)
× P(m | START,e)
× P(m | START,e,m)
× P(a | START,e,m,m)
× P(END | START,e,m,m,a)
这是概率链式法则在序列上的展开。GPT 用同一组参数实现每一个条件分布,输入前缀不同,输出概率不同。
在这个版本里,字符被映射为整数,特殊 token BOS 同时放在序列两端。因此,概念上的 START 和 END 实际对应同一个 ID。它既教模型如何开始,也教模型何时停止。
| 位置 | 当前输入 | 可见前缀 | 预测目标 |
|---|---|---|---|
| 0 | BOS | BOS | e |
| 1 | e | BOS e | m |
| 2 | m | BOS e m | m |
| 3 | m | BOS e m m | a |
| 4 | a | BOS e m m a | BOS |
训练标签向右错开一位,注意力只能使用当前及以前的输入。两件事一起成立,才是在预测未来。 如果标签没有错位,模型可以学复制;如果未来位置可见,模型可以直接偷看答案。
训练时下一步喂入真实字符,这叫 teacher forcing。模型这一刻猜错了 m,下一位置仍然会收到训练数据中的 m。生成时则不同,下一步输入来自模型刚刚采样的结果,错误会改变后续整个前缀。
一个参数,怎样知道自己该往哪边改
我觉得理解 MicroGPT 最好的入口不是 Attention,而是一个乘法节点。
设 a = 2、b = 3,计算 u = a × b,再计算 L = u + a。前向结果是 L = 8,反向要回答的是:稍微改变 a 或 b,L 会变化多少?
∂L/∂u = 1
∂u/∂a = b = 3
∂u/∂b = a = 2
∂L/∂a = (∂L/∂u)(∂u/∂a) + 1 = 4
∂L/∂b = (∂L/∂u)(∂u/∂b) = 2
这里 a 有两条通往 L 的路径,贡献必须相加。把梯度写成赋值而不是累加,会悄悄漏掉一条路径。
一个自动求导标量至少要保存四类信息:当前数值、当前累计梯度、它依赖的输入节点,以及输出相对各输入的局部导数。运算时建立这些依赖,反向时沿依赖传播梯度。
| 局部运算 | 对输入的局部导数 |
|---|---|
z = a + b | 对 a、b 都是 1 |
z = a × b | 对 a 是 b,对 b 是 a |
z = a^k | k × a^(k−1),这里 k 为常数 |
z = log(a) | 1/a |
z = exp(a) | exp(a) |
z = ReLU(a) | a 大于 0 时为 1,小于 0 时为 0;零点采用 0 |
这些规则足以拼出线性层、归一化、Softmax 和损失。矩阵只是标量的容器;在最朴素的实现里,一次矩阵乘法仍然是很多次乘法与加法。
反向传播需要逆拓扑顺序:先处理更靠近损失的节点,等它所有下游贡献到齐,再向输入传播。起点设置 ∂L/∂L = 1。概念上只有一句:
输入节点的梯度 += 输出节点的梯度 × 这条边的局部导数
遍历节点时去重,传播边时却不能把重复输入去掉。例如 a × a 对 a 的导数是 2a,两条输入边都必须贡献一次。
这也解释了为什么共享参数能训练。同一个 embedding 行可能在多个位置被使用,最终只有一个参数对象,但来自不同位置的梯度会汇总到它身上。
把网络尺寸算清楚
源码配置是 1 层、16 维隐藏状态、4 个头、16 个上下文位置;每个头 4 维,MLP 中间宽度 64。字符词表大小记作 V,矩阵按“输出维度 × 输入维度”存放。
| 参数 | 形状 | 参数量 |
|---|---|---|
| Token embedding | V × 16 | 16V |
| Position embedding | 16 × 16 | 256 |
| Q、K、V、O 四个投影 | 各 16 × 16 | 1,024 |
| MLP 第一层 | 64 × 16 | 1,024 |
| MLP 第二层 | 16 × 64 | 1,024 |
| 输出 LM head | V × 16 | 16V |
| 合计 | 无 bias、无可学习归一化增益 | 32V + 3,328 |
若词表是 26 个字母加一个 BOS,V = 27,总参数量就是 4,192。 这是按形状计算出的值,不是所有自定义数据都固定有 4,192 个参数。输入 embedding 和输出 head 在这个版本中是独立矩阵,没有共享权重。
Token ID 本身不是语义坐标。ID 为 20 的字符不会天然比 ID 为 10 的字符“大两倍”;ID 只是用于取出一行可学习向量。
位置向量则为同一个字符加入“出现在第几个位置”的信息。模型接收到的是字符向量与位置向量之和,而不是直接把整数 ID 送进乘法。
一次前向计算,数据经过了哪些地方
为避免张量维度遮住主线,我用单个位置的列向量来描述。记 R 为 RMSNorm,E 为字符表,P 为位置表:
x₀ = R(E[token] + P[position])
u = R(x₀)
q = Wq u, k = Wk u, v = Wv u
h = x₀ + Wo · MultiHeadAttention(q, K≤t, V≤t)
r = R(h)
x₁ = h + W₂ · ReLU(W₁ r)
logits = Wout x₁
probabilities = softmax(logits)
这是对当前一层实现的结构表达。不要在结尾自行补一个 final norm:这个快照从最后的残差输出直接进入 LM head。它也不是原样复刻 GPT-2,而是保留了 decoder-only 主干的一种简化。
RMSNorm:先控制输入尺度
对 d 维向量 x:
mean_square = (x₁² + ... + x_d²) / d
R(xᵢ) = xᵢ / sqrt(mean_square + ε)
例如 x 为 [3, 4],忽略很小的 ε,分母是 sqrt(12.5),结果约为 [0.8485, 1.1314]。它没有减去均值,也不会把每个元素变成同一个值;它主要调整整体尺度,同时保留方向信息。
MicroGPT 使用的版本没有额外可学习的缩放向量,ε 为 1e-5。
开头连续出现两次归一化容易让人想删掉一个。但第一份归一化后的 x 同时进入残差支路,第二次归一化只位于注意力支路。即使它们前向数值可能很接近,改动第一个也会改变残差路径及其梯度,不能仅凭“看起来重复”就认为等价。
残差:让每一块学习修正量
残差结构可以写成 y = x + F(x),因此导数包含一条直接路径:∂y/∂x = I + ∂F/∂x。网络既能保留已有表示,也能加入新的修正。
这有助于梯度传播,但不意味着任何深度、任何学习率都会稳定。残差不是一张免除数值检查的通行证。
Attention 到底在计算什么
我习惯把三个投影理解成三种职责:Q 决定当前位置如何寻找信息,K 决定历史位置如何参与匹配,V 提供被汇总的内容。它们都是从隐藏状态学出的向量,不是程序员手工定义的字段。
对当前位置 t、某一个注意力头:
score(t,j) = dot(q_t, k_j) / sqrt(d_head) j ≤ t
weight(t,j) = exp(score(t,j)) / Σ exp(score(t,r))
output_t = Σ weight(t,j) × v_j
分母与加权求和都只遍历当前可见位置。四个头分别完成这个过程,各输出 4 维,拼成 16 维后再做一次输出投影。这是缩放点积与多头注意力在当前尺寸下的应用。
为什么除以 sqrt(d_head)?在分量独立、零均值、方差适当的理想化假设下,点积方差随维度增长。缩放可以避免仅因为维度变大就让 Softmax 过度尖锐。这里每个头 4 维,因此除数是 2,而不是 4 或 16。
看一个我构造的例子。当前 q 为 [1, 0, 1, 0],三个可见 k 分别为 [1, 0, 0, 0]、[0, 1, 0, 0]、[1, 0, 1, 0]:
点积 = [1, 0, 2]
除以 sqrt(4) = [0.5, 0, 1]
Softmax 权重 ≈ [0.3072, 0.1863, 0.5065]
若三个 v 分别为:
[1, 0, 0, 0]、[0, 2, 0, 0]、[0, 0, 3, 0]
则输出约为:
[0.3072, 0.3726, 1.5194, 0]
它不是只挑中一个位置,而是按权重混合多个位置的信息。注意力权重也不是字符最终被输出的概率;前者在历史位置上归一化,后者在词表上归一化,中间还隔着投影、残差和 MLP。
没有显式 Mask,为什么仍然是因果注意力
批量计算通常先构造所有位置之间的分数矩阵,再把未来位置设为负无穷,使其 Softmax 权重为零。MicroGPT 则按位置前进:算出当前位置的 K、V,追加进当前层列表,然后只访问已有列表。
计算 t 时,列表中只有 0...t。未来信息尚未进入可访问的数据结构,因果约束由执行顺序实现。
训练时,这些 K、V 仍然连接着计算图。后面位置的损失可以通过注意力,回到前面位置的 K/V 投影与 embedding。向前不能偷看未来,不代表反向梯度不能从后面的损失流向前面的计算。
每个新文档、新生成样本都应重新建立这些列表;不同层也要有各自的 K/V。推理时可以把历史 K/V 当作无需梯度的数值缓存,训练时如果直接 detach,会截断原本需要的梯度路径。
另外,这里不是一个自动滑动的无限窗口。位置表只有 16 行,训练取序列前面最多 16 个预测位置,生成也最多循环 16 次。换成长文档后,后面的内容不会自动变成新窗口;需要自己增加切片策略,否则连结尾标记都可能被截掉。
MLP、Softmax 与交叉熵如何接上
Attention 负责跨位置汇总信息,MLP 对每个位置的表示进行非线性变换。在这里,它把 16 维扩到 64 维,经过 ReLU,再投影回 16 维。
没有非线性时,两次线性变换可以合并成一次:W₂(W₁x) = (W₂W₁)x。ReLU 让这个合并不再普遍成立,也让模型能够根据输入激活不同的特征组合。
LM head 随后给每个候选 token 一个 logit。logit 不是概率,可以为负,也不需要加起来等于 1。Softmax 才把它们变成分布:
pᵢ = exp(zᵢ − c) / Σⱼ exp(zⱼ − c)
c = max(z)
减去相同常数不会改变概率,因为分子分母共同的因子会抵消;选择最大值能避免较大的正数进入指数。把 c 作为普通数值使用不破坏这里的正确导数,因为 Softmax 对整体平移不敏感。不过极端小概率仍可能下溢;生产训练一般使用稳定的 log-softmax / 交叉熵组合,而不是先算很小的概率再取 log。
目标字符为 y 时,单位置损失是 L_t = −log(p_y)。正确字符概率从 0.1 升到 0.5,损失就从约 2.3026 降到 0.6931。整个文档取各位置损失的平均值,而不是只训练最后一个字符。
这个组合的导数尤其值得记住。由 L = −z_y + log(Σ exp(z_j)) 可以直接得到:
∂L/∂zᵢ = pᵢ − 1[i = y]
假设预测 [0.2, 0.5, 0.3],正确答案是第一项,则 logit 梯度为 [-0.8, 0.5, 0.3]。梯度下降会直接提高正确项的 logit、降低其他项;再通过链式法则把这个误差信号传回所有相关参数。对 n 个位置取平均后,每个位置的这份贡献还要除以 n。
在均匀猜测的参照下,27 类的损失是 ln(27) ≈ 3.2958,困惑度为 27。这是理论基线,随机初始化并不保证输出严格均匀,也就不保证第一步恰好打印这个值。
Adam 如何把梯度变成参数更新
拿到梯度以后,最直接的做法是 θ ← θ − ηg。Adam 还会维护梯度的一阶、二阶指数移动平均,为不同参数调整更新尺度。下面直接展开它的更新规则。
用 s 表示从 1 开始的优化步数,避免与 token 位置混淆:
m_s = β₁ m_(s−1) + (1−β₁) g_s
v_s = β₂ v_(s−1) + (1−β₂) g_s²
m̂_s = m_s / (1−β₁^s)
v̂_s = v_s / (1−β₂^s)
θ_s = θ_(s−1) − η_s × m̂_s / (sqrt(v̂_s) + ε)
偏差修正是因为 m、v 从零开始,早期估计会被零值拉低。源码使用 β₁ = 0.85、β₂ = 0.99,初始学习率 0.01,并按步数线性衰减;这是 Adam,没有 AdamW 的解耦权重衰减项。
例如第一次梯度为 0.2,则 m 为 0.03、v 为 0.0004,修正后分别为 0.2 和 0.04。忽略极小 ε,第一次更新约为减去 0.01。这个例子说明“梯度大多少,参数就一定多改多少”并不是 Adam 的规则。
训练循环的职责可以独立地写成下面的伪代码;这是阅读用结构,不是复制一份可直接运行的 GPT:
初始化参数 θ,以及 Adam 状态 m、v
每一步:
取一条文档,编码并加入边界标记
为每层建立新的 K/V 列表
对每个有效位置,用真实前缀计算下一个 token 的损失
对这些损失取平均,执行一次反向传播
用 Adam 更新 θ,清零参数梯度
每条文档平均后再更新一次,意味着这里每一步的采样单位是文档。它不自动等同于把整个语料所有 token 放在一起求平均:不同长度文档会获得不同的每 token 相对权重。
默认 1,000 步也不等于 1,000 个 epoch。一次 step 处理一条文档;只有覆盖完整个数据集才构成一次遍历。学习率公式使用从 0 开始的 step,因此最后一次更新仍有一个很小的正学习率,而不是恰好为零。
生成时,学习已经停止
训练后,从 BOS 开始重复以下流程:计算 logits,按温度变成概率,采样一个 token,把它作为下一次输入。采到 BOS 就结束,否则最多生成 16 个字符。
pᵢ(T) = exp(zᵢ / T) / Σⱼ exp(zⱼ / T) T > 0
若两个候选的 logit 差为 1,T 为 1 时,它们的概率比是 e¹ ≈ 2.72;T 为 0.5 时是 e² ≈ 7.39。降低温度使偏好更集中,但不会增加知识,也不保证事实更正确。T 可以大于 1;T 为 0 时不能直接套用除法,通常要另行实现贪心选择。
推理不执行损失反传和 Adam 更新,所以“读过了当前前缀”改变的是激活和缓存,不是长期参数。这个极简实现复用自动求导运算,因此推理仍有建图开销;没有调用 backward,不代表已经具备框架中的 no-grad 执行模式。
生成一个像名字的字符串,能说明模型捕捉到某些字符组合规律;不能仅凭结果新颖就断言没有记忆训练数据。要讨论泛化,还需要保留集、重合率检查,以及与任务对应的评价。
我会怎样验证自己的实现
短代码并不天然正确。我会先用能够手算的输入检查计算,再看训练曲线。
- 共享节点梯度:
L = a×b+a在 a=2、b=3 时,应得到 4 和 2;a×a应得到2a。 - 数值梯度:对选定参数比较自动求导结果与中心差分
(L(θ+h)−L(θ−h))/(2h);检查时避开 ReLU 的零点。 - 因果性:保留相同前缀、改变后缀,前缀位置的 logits 应保持一致;若改写为并行版本,再与逐位置版本对照。
- 损失与更新:核对 Softmax 总和、交叉熵导数、第一次 Adam 更新,以及参数梯度是否在每步后清零。
- 数据边界:检查 BOS 配对、长文档截断和词表覆盖;最后才做小样本过拟合与独立验证。
我为本文准备了一个不依赖第三方库的计算核对脚本,覆盖共享节点反传、中心差分、注意力数值、交叉熵梯度和 Adam 首步更新。它核对的是文中的数学例子,不是完整 GPT 的训练或性能测试。
想直接运行原作者实现,可以在一个独立空目录保存上述固定版本为 microgpt.py,执行 python3 microgpt.py。首次没有 input.txt 时,脚本会联网下载默认名字数据;使用自己的文件时,每行应为一条非空样本,并留意字符词表与 16 位置截断规则。原始 Gist 的当前代码及修订记录也值得与固定版本对照。
读完之后,我怎样理解“完整算法”
MicroGPT 对我最有价值的地方,是把几件经常分开学习的事情连成了一条可追踪路径:字符决定输入向量,注意力混合上下文,输出层给出概率,损失产生梯度,优化器修改参数,下一次预测因此发生变化。
但从这个闭环到一个可用的大模型,中间仍有大量会影响能力与可靠性的决策。数据质量、Tokenizer、训练目标、模型结构、评测和后训练,都不能简单归结为“跑快一点”。GPU、批处理和融合算子主要改变执行效率;换语料、换监督方式,改变的则是模型究竟在学什么。
回到我自己的工程习惯,这篇代码提供了一种很好的调试顺序:预测不对,先查目标是否错位;训练不动,查梯度是否断开;Loss 异常地好,查未来是否泄漏;采样很奇怪,查概率和结束条件。
我希望自己对模型的理解,最终能落到这些可以检查的关系上。知道每个数字从哪里来,才有依据判断下一步应该改哪里。