Understanding LLMs II: Forward Pass & Attention
模型第一件事是把离散的 token id 翻译成连续向量。字符 ‘e’ 的 id=4,但 4 这个数字本身没有”语义”——4 和 5 之间的距离不比 4 和 26 更近。所以用一张可学习查找表(embedding matrix)查表:
# wte: (vocab_size, d_model) = (27, 16),每行一个字符的"含义向量"
# wpe: (block_size, d_model) = (8, 16),每行一个位置的"位置向量"
x = wte[ids] + wpe[positions] # (T, 16)
wte[4] 取出 ‘e’ 的 16 维向量。这两张表都是模型参数(4000 多个数字里的一部分),由训练学到。
one-hot 把 ‘e’ 表示为 27 维向量 。它有两个致命问题:
拼接 concat(wte[4], wpe[0]) 会让维度翻倍(32 维),且模型必须从头学”前一半是字符、后一半是位置”的分工。相加保持 16 维,且让位置和字符信息纠缠在同一组分量里——后续注意力能从这种纠缠里”解码”出两者(线性层 对加法可分)。
更深的原因来自原始 Transformer 的正弦位置编码:相加等价于在频域里调制,类似 AM 广播把信号叠加到载波上。学习式位置嵌入(microgpt 用的)继承了这个加法约定。
注意力本身是集合操作——对输入顺序不敏感(精确地说,self-attention 对输入做置换等变:重排输入序列,输出也跟着重排,但每位置的值不变)。如果没有位置信息,“abc” 和 “cba” 在模型看来完全一样(每个字符的查询结果只取决于”是哪些字符”,不取决于顺序)。位置嵌入打破这种对称性,让模型知道”e 在第 0 位”和”e 在第 3 位”是不同的。
wpe 的最大长度限制了上下文窗口。microgpt 的 block_size=8,wpe 只有 8 行,序列长度超过 8 就越界——这就是为什么 GPT 有”上下文窗口”限制(GPT-4 是 8K/128K)。位置编码必须为该长度预先准备好。wte 和输出层 lm_head 可以共享权重(weight tying)。nanoGPT 可选 tie_word_embeddings:让 lm_head = wte.T,节省参数且效果常更好。microgpt 不共享(独立 lm_head),见 B16。Transformer 由若干相同 block 堆叠(GPT-4 几十层,microgpt 1 层,结构完全一致)。每个 block 内部是对称的两阶段:
x ─→ RMSNorm ─→ Attention ─→(+)─→ RMSNorm ─→ MLP ─→(+)─→ 输出
└──────────────────────────┘ 残差 └──────────────────┘ 残差
def block(x):
x = x + attn(rmsnorm(x)) # 子层 1:注意力 + 残差
x = x + mlp(rmsnorm(x)) # 子层 2:MLP + 残差
return x
这两个子层功能正交,缺一不可:
只堆 attention 没法做复杂特征变换(attention 本身是线性的点积 + softmax,表达力有限);只堆 MLP 没法跨位置通信。两者交替:attention 收集上下文 → MLP 加工成更抽象的特征 → 下一层 attention 再收集……这种”通信-计算”的交替堆叠是 Transformer 取代 RNN/CNN 的核心结构创新。
每层 x + f(x) 的形式让深度网络可训练(详见 B14)。GPT-3 有 96 层,若没有残差,反向传播的梯度穿过 96 层线性 + softmax 早就消失了。残差提供一条”无障碍通道”,梯度能直接传到底层。
每子层前 rmsnorm(x) 把向量幅度归一化,防止深堆叠下数值爆炸/消失。microgpt 用 pre-norm(归一化在残差内部),比原始 Transformer 的 post-norm 更稳定。
不是”必须相同”,而是”相同已足够”。层与层之间参数独立(每层自己的 attn、mlp 权重),但结构一致。这种 uniform 设计有两个好处:
不同层学到的功能不同(浅层学语法、深层学语义),但这是训练自适应的结果,不需要人为设计差异。
注意力本质是一个软查询(soft lookup):每个位置都问”我该从前面哪些位置取多少信息”。三个矩阵把同一个输入向量投到三个不同子空间:
Q = x @ wq # (T, 16) @ (16, 16) → (T, 16) 查询:我要找什么
K = x @ wk # 键:我能提供什么
V = x @ wv # 值:你看上我,我给你什么
三个矩阵 wq/wk/wv 是模型参数,训练学到。
| 数据库 | 注意力 |
|---|---|
Query SELECT ... WHERE key = ? | 和每个 算相似度 |
| 返回完全匹配的行的 value | 返回所有行 value 的加权和 |
| 离散、硬匹配 | 连续、软匹配(按相似度加权) |
软匹配让模型可以”半信半疑”地从多个位置各取一部分——比如预测 “joh” 后的字符时,‘n’ 的位置贡献大,但 ‘o’ 的位置也可能贡献一点(提供”这是元音-辅音交替”的线索)。
直接用 也能跑(叫 “linear attention” 的退化形式),但表达力大减。用三个独立投影矩阵 wq/wk/wv 让模型学到三种不同的”视角”:
wq 把 投成”我要找的查询条件”。wk 把 投成”我对外广告的标签”。wv 把 投成”如果有人选我,我真正交付的内容”。同一个 ‘n’ 字符,作为查询时它的”问题”(“还有别的辅音吗?“)和作为键时它的”标签”(“我是辅音 n”)是不同的投影——分开训练才能学到。
注意 Q、K、V 都从同一个序列 算出来,这是 self-attention(自注意力)。区别于 cross-attention(Q 来自 decoder、K/V 来自 encoder,用于翻译模型)。self-attention 让序列内部任意两位置直接交互,是 Transformer 取代 RNN 的关键——RNN 要逐步传递信息,self-attention 一步到位(任意两位置 跳跃)。
设序列长 、模型维度 :
每个位置 的查询向量 、每个位置 的键 、值 。
wq/wk/wv 在多头里会按头拆开(B13)。microgpt 是 4 头 × 4 维 = 16 维,每个 wq 实际是 16×16 但被 reshape 成 (4 头, 4 维) 分头用。完整的注意力运算:
分三步:① 算 logits ;② 缩放 ;③ softmax 后对 加权求和。
是位置 的查询 与位置 的键 的点积。点积越大,两者越”匹配”(在投影后的子空间里方向越一致)。这个 矩阵就是 attention logits。
scores = Q @ K.transpose(-2,-1) # (T, d) @ (d, T) → (T, T)
点积 。若 的分量是均值 0、方差 1 的独立随机变量,点积的方差正比于 ()。
当 大(microgpt 是 4,GPT-2 是 64),点积数值会很大。softmax 对大数值敏感:最大的那个 logit 经过 会指数膨胀,把概率几乎全压在一个位置,梯度趋近 0(饱和区)——训练停滞。除以 把方差拉回 1 附近,softmax 工作在线性区:
每行 是一个概率分布(和为 1),表示位置 对各位置的”关注权重”。
attn = (scores / sqrt(d_k)).softmax(dim=-1) # (T, T),每行和为 1
out = attn @ V # (T, T) @ (T, d) → (T, d)
GPT 是自回归模型(autoregressive),预测第 个字符时不能偷看第 的真实值。所以在算 attention 前,把 中 的位置(未来)设为 :
tril = torch.tril(torch.ones(T, T)) # 下三角
scores = scores.masked_fill(tril==0, float('-inf'))
softmax 后这些位置的权重变成 ,位置 只能看 。这就是 causal self-attention。
位置 的输出是所有”前文” value 的加权和,权重由 Q-K 相似度决定。这就是”当前位置融合了相关位置信息”的数学实现。
softmax(dim=-1) 沿最后一个维度(即 维度)归一化,不是沿 。意思是”每个 query 对所有 key 的权重和为 1”。若搞错 dim,模型完全错乱。K.shape[-1]**0.5 自动取对。单头注意力只能学到一种”匹配模式”。多头让模型在同一层里同时跑多个独立的注意力,每个在更小的子空间里学不同的关系:
# microgpt: d_model=16, n_heads=4, head_dim=4
Q = x @ wq # (T, 16)
Q = Q.view(T, 4, 4).transpose(0,1) # (4 头, T, 4 维)
# 每头独立做 scaled dot-product attention
heads = [attention(Q[h], K[h], V[h]) for h in range(4)] # 每头输出 (T, 4)
out = torch.cat(heads, dim=-1) @ wo # 拼回 (T, 16) 再投影
16 维拆成 4 头 × 4 维:
总计算量与单头 16 维几乎相同(都是 量级),但拆成 4 个独立的”视角”,每个 4 维子空间里学自己的 Q-K 匹配模式。
单头在一个 16 维空间里只能学一种相似度度量(一个 的”问题-标签”配对)。多头让模型有 4 套独立的 ,每套学不同的关系:
这些模式在不同子空间里互不打架。最后 concat 拼起来再用 wo 线性混合,让模型决定怎么综合多视角信息。
每头参数量 ,4 头共 参数——和单头 完全一样。多头不增加参数量,只是把它们重新组织:单头是一个 16→16 矩阵,多头是 4 个 16→4 矩阵(拼接后等价于一个 block-diagonal 结构)。
数学上,多头注意力等价于把 限制为分块对角形式(每个头的 query/key/value 在自己的 4 维子空间里不与其他头交互),再用 做一次跨头混合。
单纯 concat 只是堆叠 4 头的输出,没有信息交流。wo(16×16 矩阵)做线性变换,让头 1 的信息可以影响头 2 对应的维度——这是”多头协同”的关键。如果省掉 wo(直接用 concat 结果),各头的预测能力被局限在自己的子空间,效果显著下降。
d_model % n_heads 必须整除。microgpt 是 16/4=4 整除;如果 d_model=17, n_heads=4 就无法均分,工程上会报错。GPT-2 (768/12=64)、GPT-3 (12288/96=128) 都满足。scores / math.sqrt(K.size(-1)) 自动取每头维度。每个 Transformer block 的结构(pre-norm 变体,microgpt 用这个):
x = x + attn(rmsnorm(x)) # 残差 + 注意力(输入先归一化)
x = x + mlp(rmsnorm(x)) # 残差 + MLP
两个细节:残差连接 x + f(x) 和 RMSNorm 归一化。
形式 ,“短路”绕过 。两层意义:
残差由 He et al. (2015) 在 ResNet 提出,解决”网络越深反而越差”的悖论。Transformer 几乎所有层都用它。
LayerNorm 对每个 token 的向量做去均值 + 除标准差 + 仿射:
RMSNorm (Zhang & Sennrich 2019) 去掉减均值,只除 RMS(root mean square):
只保留幅度归一化,去掉”中心化”。作者论证:LayerNorm 起作用的主要是”除以方差”(缩放不变性),减均值贡献小但占计算量。RMSNorm 在大模型(LLaMA、microgpt)广泛替代 LayerNorm,速度快 10-50% 且效果相当。
γ 是可学习的逐分量缩放(microgpt 是 16 维权重),初始化为 1。 防止除零。
# Pre-norm (microgpt / GPT-2 / LLaMA)
x = x + attn(rmsnorm(x))
# Post-norm (原始 Transformer 论文)
x = rmsnorm(x + attn(x))
Post-norm 把归一化放在残差外,梯度流要穿过 norm 层,深层容易不稳定(原始论文要 warmup 学习率才能训)。Pre-norm 让残差主路”裸奔”,归一化只作用于 的输入,梯度可以无障碍地沿残差通道传到底——这是为什么现代大模型都改用 pre-norm,能直接训上百层。
rmsnorm 沿最后一个维度(特征维 d_model=16)做,不是沿序列维 T。意思是每个 token 独立归一化自己的 16 个分量,token 之间互不影响。这与 BatchNorm(沿 batch 维归一化)正交——LayerNorm/RMSNorm 不依赖 batch,所以推理(batch=1)和训练行为一致,适合自回归生成。
γ 是逐分量的,不是标量。16 维向量有 16 个独立的 ,让模型对每个特征维度做不同缩放。误写成标量会失去灵活性。x + f(x) 要求 输入输出同维(都是 16)。若中间层改维度(如 16→64→16,见 B15 MLP),残差必须加在 16 维那一层不能跨 64 维。epsilon 的位置在根号内: 不是 。前者在 RMS 接近 0 时主导项是 ,仍能稳定除;后者主导是 本身,量级太小。每个 Transformer block 在注意力之后还有一个 MLP,对每个位置独立做:
# d_model=16, d_ff=64(4× 升维)
h = relu(x @ w1) # (T, 16) @ (16, 64) → (T, 64)
out = h @ w2 # (T, 64) @ (64, 16) → (T, 16)
两层全连接,中间夹 ReLU,升维 4 倍再压回。
16→64→16 看似”先扩张后收缩”的浪费,但这是有意的:
GPT-2/3 也用 4× 比例(d_model=768 → d_ff=3072),是经验上的甜点。LLaMA 等用 SwiGLU 时会调整成约 8/3×(因为有门控分支,参数重平衡)。
没有非线性激活,两层全连接 等价于一个单层线性变换——无论堆多少层,整体仍是线性的,学不到复杂函数。ReLU 的”折线”形状提供分段线性近似能力:足够多的 ReLU 神经元可以逼近任意连续函数(万能逼近定理)。
GPT-2 用 GELU,LLaMA 用 SwiGLU,microgpt 用最简单的 ReLU——核心作用都是引入非线性。
MLP 对每个 token 独立做,token 之间不交互(每个位置的 16→64→16 完全独立于其他位置)。这和注意力正交:
两者交替堆叠:注意力收集上下文信息 → MLP 加工成更抽象的特征 → 下一层注意力再收集……这就是 Transformer block 的设计逻辑。
Geva et al. (2020) 提出:MLP 的第一层 的每行可看作一个”键”(识别某种输入模式),第二层 的对应列是该键激活时输出的”值”。MLP 整体是 的软版本——所以 MLP 在大模型里常被解读为事实记忆存储(如”巴黎是法国首都”存在某些 MLP 神经元里)。这解释了为什么 MLP 占模型参数的 2/3(microgpt 也如此)。
走完 Transformer block,每个位置得到一个 16 维向量 。最后一步用 lm_head(language modeling head)把它映射到词表大小:
logits = h @ lm_head # (T, 16) @ (16, 27) → (T, 27)
每个位置输出 27 个数(对应 26 个字母 + BOS),叫 logits(未归一化的对数概率)。最大值的位置就是模型预测的下一个字符。
logits 可以是任意实数(负无穷到正无穷),不归一化。要变成概率必须再过 softmax:
但 microgpt 推理时不需要算 softmax 就能取 argmax——因为 softmax 是单调函数,。所以推理只做 argmax(logits),省一次 softmax 计算。
lm_head 的形状是 (16, 27),把 16 维”语义向量”投到 27 维”字符空间”。它和 embedding wte (27, 16) 形状互为转置。两者可以共享权重(weight tying):
lm_head = wte.T # (27,16).T = (16,27)
直觉:嵌入学的是”字符 → 语义向量”,输出层是逆向的”语义向量 → 字符”,用同一个映射合理。nanoGPT 默认不共享(独立 lm_head),microgpt 也是。共享的优点是参数少(27×16=432 个参数省下来)、有时泛化更好;缺点是限制了模型对输入和输出用不同表征的能力。
是 与第 个字符的”输出向量”的点积——衡量 这个语义向量”有多像字符 “。所以输出层本质是用点积度量当前隐状态与每个候选字符的相似度,点积最大的就是预测。
这与 B11 QKV 注意力里 的几何完全同构:都是”用点积在向量空间里找最匹配的”。
推理时除 argmax 外还可以从 softmax 分布里采样:
温度 控制分布锐度: 退化为 argmax(贪心), 是原始分布, 让分布更平均(更有创造性、更随机)。microgpt 训练用 argmax 评估,但生成时可以加温度增加多样性。
lm_head 的方向:是 不是 。h @ lm_head 把 16 维投到 27 维;写反成 lm_head @ h 形状不对。模型输出 logits,真实答案是某个字符(one-hot 标签 )。需要一个数字衡量”模型预测与真相差多远”——交叉熵损失:
最后等号因为 是 one-hot(只有真实字符位置为 1,其余 0)。所以交叉熵损失就是真实字符被赋予的概率的负对数。
交叉熵 度量”用分布 编码来自分布 的事件,平均需要多少比特”。模型越接近真实分布(),交叉熵越小;模型把概率给错地方(),,损失爆炸。
模型预测真实字符的概率 p_true → loss = -log(p_true)
p_true = 0.9 → loss ≈ 0.11 (自信且对)
p_true = 0.5 → loss ≈ 0.69 (不确定)
p_true = 0.01 → loss ≈ 4.6 (自信但错)
p_true = 0.001 → loss ≈ 6.9 (大错)
负对数让”低概率给真值”受重罚,符合直觉。
直接 log(softmax(logits)) 在 logits 大时数值溢出( 爆炸)。实际用 log-sum-exp trick:
其中 。减去最大值让所有 参数 ≤ 0,结果 ∈ (0, 1],避免溢出。PyTorch 的 F.cross_entropy 内部就是这个:
loss = F.cross_entropy(logits.view(-1, V), targets.view(-1))
# 等价于:先 log_softmax(带 log-sum-exp trick),再 NLL
microgpt 训练时一次喂 8 个字符的序列,每个位置都预测下一个字符,共 8 个 loss:
取平均让 loss 不依赖序列长度,便于跨 batch 比较。这个标量就是”模型本次考试的分数”,反向传播会从它出发算每个参数的梯度。
分类问题用交叉熵而不是 MSE(mean squared error),有两个理由:
targets 是 token id 不是 one-hot。PyTorch 的 cross_entropy 内部把整数标签当作 one-hot 处理,比显式 one-hot + 计算高效。误传 one-hot 进去维度错乱。<PAD>),这些位置不该贡献 loss——实际实现里用 ignore_index 参数屏蔽。microgpt 固定长度 8,没有 padding 问题。