第二章 Transformer 与 MoE:从矩阵计算到稀疏专家¶
Transformer 的核心并不神秘:每个位置先从其他位置汇总信息,再独立通过前馈网络加工;残差连接和归一化保证深层网络可训练。难点在于掩码、形状、数值稳定性和训练/推理差异。本章从一次注意力计算开始,逐步组装完整结构,再解释 MoE 如何用稀疏激活扩大参数容量。

图 2-1 Dense Transformer 每个 token 经过同一 FFN;MoE 用路由器为 token 选择少数专家。
2.1 Token、嵌入与位置¶
先建立直觉。 语言模型不能直接处理字符串,必须先把文本切成 token id,再通过嵌入表把离散 id 映射为连续向量。位置表示则回答相同 token 出现在不同位置时怎样区分。 读这一节时,先不要急着记缩写,而要不断追问:输入是什么,经过了哪些可观察变换,输出又怎样被验证。
把过程拆开看:
- 规范化与分词得到 token。
- 查词表得到整数 id。
- Embedding 查表得到 d_model 维向量。
- 叠加或注入位置信息。
最小例子。 词表大小 V=10,000、维度 D=512 的嵌入层本质是一个 V×D 参数矩阵;一批 B×T 的 id 会得到 B×T×D。 例子越小,越容易手算、打印中间量并判断实现是否偏离定义。
容易踩坑。 token 不等于汉字或单词;不同 tokenizer 不能随意共用权重。padding、BOS、EOS 的 id 和 mask 也必须与训练配置一致。 排错时应保存失败输入和版本,先定位最早出现异常的步骤,再决定是否调整模型或超参数。
动手任务:用任意公开 tokenizer 编码中英文、代码和数字,比较 token 数并解释差异。 完成后不要只保存最终输出,还要保存代码、配置、中间张量或检索结果,以及你对结果的解释。
模型不直接读取汉字或单词,而是读取 tokenizer 产生的 token id。BPE、WordPiece、Unigram 等算法都试图在词表规模、序列长度和未登录词之间折中。tokenizer 是模型的一部分:更换 tokenizer 会改变 id 语义,不能只替换词表文件。
设词表大小为 \(V\)、隐藏维度为 \(d\),嵌入矩阵 \(E\in\mathbb{R}^{V\times d}\)。查表把 (B,T) 的 token id 变为 (B,T,d)。注意力本身对输入顺序置换等变,因此还需注入位置信息。原始 Transformer 使用正弦位置编码,现代解码器常用 RoPE;第二者将在第三章推导。
2.2 Self-Attention:每一步都写出形状¶

图 02-2 注意力实现的第一道正确性门槛,是每次 reshape、transpose 和矩阵乘都能写出形状。 从问题出发。 Self-Attention 让每个位置根据内容选择其他位置的信息。Q 表示当前要找什么,K 表示每个位置可被怎样匹配,V 表示真正被汇总的内容。 先把名词放到一边,沿着输入、状态变化和输出走一遍;只有能指出证据落在哪一步,概念才算真正掌握。
沿数据流逐步检查:
- X 线性投影为 Q/K/V。
- 计算 QK^T 并除以 sqrt(dk)。
- 叠加因果或 padding mask。
- Softmax 后乘 V。
用小数据走一遍。 B=2、T=4、D=8、H=2 时,每头维度为 4;Q/K/V 形状均为 (2,2,4,4),注意力矩阵为 (2,2,4,4)。 先在纸面预测结果,再让程序打印中间状态;预测与运行不一致的地方,正是需要继续追查的知识缺口。
这里最容易出现的误解。 Softmax 轴必须是 key 轴;mask 的布尔含义与填充值要核对。把缩放因子写成 sqrt(d_model) 或在半精度中使用过小负数都可能出错。 出现异常时先冻结样本、配置和环境,从最靠近输入的环节开始验证;不要同时更换模型、数据和超参数。
小实验:在手写注意力中返回权重,验证每行和为 1,并检查未来位置权重为 0。 提交物应包含预期、运行证据和失败记录;只有别人能按同样步骤复现,结果才具有学习价值。
输入 \(X\in\mathbb{R}^{B\times T\times d}\) 经过三组线性投影:
$$
Q=XW_Q,\qquad K=XW_K,\qquad V=XW_V.
$$
对单头注意力,\(QK^\top\) 得到 (B,T,T) 分数矩阵。第 \(i\) 行表示第 \(i\) 个查询对所有键的相似度。除以 \(\sqrt{d_k}\) 是为了避免维度增大时点积方差过大,导致 Softmax 过饱和。掩码 \(M\) 把不可见位置加上极小值;解码器的因果掩码保证位置 \(i\) 只能看到不晚于自己的 token。
多头注意力把隐藏维分为 \(H\) 个头,每个头维度 \(d_h=d/H\)。头并非简单重复:不同投影可学习不同关系。各头输出拼接后再经过 \(W_O\) 混合。
import math
import torch
from torch import nn
class MultiHeadSelfAttention(nn.Module):
def __init__(self, d_model: int, n_heads: int, causal: bool = True):
super().__init__()
if d_model % n_heads != 0:
raise ValueError("d_model 必须能被 n_heads 整除")
self.n_heads = n_heads
self.head_dim = d_model // n_heads
self.causal = causal
self.qkv = nn.Linear(d_model, 3 * d_model, bias=False)
self.out = nn.Linear(d_model, d_model, bias=False)
def forward(self, x, padding_mask=None):
b, t, d = x.shape
qkv = self.qkv(x).view(b, t, 3, self.n_heads, self.head_dim)
q, k, v = qkv.unbind(dim=2)
q, k, v = [z.transpose(1, 2) for z in (q, k, v)] # (B,H,T,Dh)
scores = q @ k.transpose(-2, -1) / math.sqrt(self.head_dim)
if self.causal:
mask = torch.triu(
torch.ones(t, t, dtype=torch.bool, device=x.device), diagonal=1
)
scores = scores.masked_fill(mask, torch.finfo(scores.dtype).min)
if padding_mask is not None:
# padding_mask: (B,T),True 表示有效 token
scores = scores.masked_fill(~padding_mask[:, None, None, :],
torch.finfo(scores.dtype).min)
attn = torch.softmax(scores.float(), dim=-1).to(x.dtype)
context = attn @ v
context = context.transpose(1, 2).contiguous().view(b, t, d)
return self.out(context), attn
代码解读:view 之前要确认内存布局,转置后合并维度使用 contiguous() 更稳妥;Softmax 临时升到 FP32 可减小低精度下的数值风险;padding mask 作用在键的位置,因果 mask 作用在时间关系。若某一行全部被屏蔽,Softmax 可能产生 nan,数据拼接和 mask 逻辑必须避免这种情况。
2.3 Encoder、Decoder 与 Cross-Attention¶
先看它解决什么。 Encoder 让所有位置双向交互,Decoder 的自注意力必须因果遮蔽;Cross-Attention 则让 Decoder 的查询读取 Encoder 产生的键和值。 理解的标准不是会复述定义,而是能画出数据流、预测中间结果,并设计一个让错误暴露出来的检查。
可以把实现分成以下环节:
- Encoder 形成上下文化记忆。
- Decoder 读取已生成前缀。
- Cross-Attention 对齐目标与源。
- 输出层预测下一个 token。
一个可以手算的例子。 翻译中,生成目标词时 Q 来自目标端当前状态,K/V 来自整句源语言表示,因此可以动态关注源句不同位置。 这里故意不用大模型或大数据,因为小输入能把每个轴、分数和状态变化完整暴露出来。
排错时先看这些地方。 BERT 是 Encoder-only,GPT 是 Decoder-only,T5 是 Encoder-Decoder;三者差异不能简化为是否含注意力。 修复不能止于‘这次跑通’,还要把失败样本变成自动测试,防止同类问题在下一版本重新出现。
现在动手:画出三种架构的数据流,并为分类、翻译、开放式生成各选择一种架构说明理由。 把关键断言写进测试,并在 README 说明怎样运行、怎样判断正确、当前实现还不支持什么。
Encoder 使用双向 Self-Attention,每个位置可读取整个输入,适合分类、抽取和表征学习。Decoder 使用因果 Self-Attention,自回归预测下一个 token,GPT、Llama、Qwen 等生成模型采用这种结构。Encoder-Decoder 模型先编码输入,再由 Decoder 通过 Cross-Attention 读取编码结果,适合翻译、摘要和条件生成。
Cross-Attention 中,\(Q\) 来自 Decoder 当前状态,\(K,V\) 来自 Encoder 输出。它回答的是“生成当前 token 时,输入中的哪些位置最相关”。T5 是典型 Encoder-Decoder;BERT 是 Encoder-only;GPT 系列是 Decoder-only。三者的差异不仅是层数排列,还包括预训练目标与注意力可见范围。
2.4 FFN、残差和归一化¶
抓住这一节的主线。 注意力负责 token 之间通信,FFN 负责每个 token 内部的非线性变换;残差通路保留信息与梯度,归一化控制数值尺度。 遇到新术语,先找它在系统中的位置:它读取什么、保存什么、改变什么,以及失败时会留下什么信号。
真正动手时按这个顺序走:
- 子层计算变换 F(x)。
- Dropout 或门控调节输出。
- 与残差 x 相加。
- 按 Pre-Norm 或 Post-Norm 放置归一化。
先做最小实验。 FFN 通常先把 D 扩张到约数倍隐藏维,再压回 D。它对每个位置使用同一组权重,因此不混合序列位置。 把随机性固定并保留中间量,这个例子就能成为后续优化时的回归基线。
别被表面现象带偏。 把 FFN 当卷积、把 LayerNorm 放在 batch 维、忘记残差两端形状必须一致,都会破坏结构。Pre-Norm 稳定性也不代表任何模型都可随意改。 先区分定义错误、实现错误、数据问题和性能瓶颈。四类问题需要的证据不同,不能靠盲目调参混在一起处理。
验证任务:记录一个 Transformer Block 经过 Attention、残差、FFN 后的均值、标准差与梯度范数。 实验结束后写三句话:观察到了什么、这些证据支持什么结论、还有哪种解释尚未排除。
注意力在 token 之间混合信息,FFN 在每个 token 内独立变换通道:
$$ \mathrm{FFN}(x)=W_2\,\phi(W_1x+b_1)+b_2. $$ 中间维度通常大于隐藏维度。原始 Transformer 使用 ReLU,BERT 常用 GELU,Llama 使用 SwiGLU。残差连接写作 \(x+F(x)\),让网络能学习对恒等映射的修正,并给梯度提供短路径。
Post-Norm 先做子层再归一化,Pre-Norm 先归一化再做子层。深层语言模型常采用 Pre-Norm,因为训练通常更稳定:
class TransformerBlock(nn.Module):
def __init__(self, d_model=256, n_heads=8, mlp_ratio=4):
super().__init__()
self.norm1 = nn.LayerNorm(d_model)
self.attn = MultiHeadSelfAttention(d_model, n_heads, causal=True)
self.norm2 = nn.LayerNorm(d_model)
self.ffn = nn.Sequential(
nn.Linear(d_model, mlp_ratio * d_model),
nn.GELU(),
nn.Linear(mlp_ratio * d_model, d_model),
)
def forward(self, x, padding_mask=None):
a, _ = self.attn(self.norm1(x), padding_mask)
x = x + a
x = x + self.ffn(self.norm2(x))
return x
LayerNorm 对每个 token 的隐藏维做标准化;RMSNorm 只按均方根缩放,不减均值。归一化中的 eps 太小会在低精度下不稳定,过大又改变尺度。
2.5 从输入到下一个 token¶
先把概念落到可观察对象上。 自回归模型训练时一次并行预测所有位置的下一个 token,推理时却必须逐 token 生成。这解释了训练吞吐高而生成延迟明显的根本差异。 这一节的重点是建立因果链。先说明问题,再看计算或流程,最后用可重复实验验证结论。
把抽象概念还原成操作:
- 右移标签形成输入与目标。
- 用因果 mask 并行计算 logits。
- 交叉熵只统计有效标签。
- 推理时按解码策略循环。
把它缩小到能逐项检查。 序列 [BOS,我,爱,猫,EOS] 的输入可为前四个 token,标签为后四个;每个位置只学习预测紧邻的下一个 token。 如果这个小例子还不能解释清楚,扩大数据只会让错误更难发现。
需要特别守住的边界。 训练标签未右移、padding 未设为 ignore_index、推理时忘记 EOS、把 temperature 用在训练损失上,都是常见混淆。 保留完整 trace 比猜原因更重要。找到第一次偏离预期的位置,通常比分析最终错误输出更高效。
本节练习:构造长度不同的两条序列,完成 padding、attention mask 和 label mask,并逐位置核对。 除代码外,请保留一份结果说明,标明环境、随机种子、输入规模和你主动检查过的边界。
解码语言模型的训练输入通常是 token 序列 \([x_0,x_1,\ldots,x_{T-1}]\),目标是右移后的 \([x_1,x_2,\ldots,x_T]\)。模型一次并行计算所有位置的 logits,因果掩码防止偷看未来。推理时没有真实的下一个 token,只能生成一个、追加到序列、再生成下一个,因此解码天然串行。KV Cache 通过复用历史键值减少重复计算,第十三章会详细说明。
采样策略决定如何从概率分布选 token。贪心每次取最大概率,稳定但容易僵化;temperature 调整分布锐度;top-k 只保留概率最高的 \(k\) 个;top-p 保留累计概率达到阈值的最小集合。Beam Search 更适合有明确序列得分的任务,不一定适合开放对话。
2.6 MoE:容量扩大不等于计算同比扩大¶
先建立直觉。 MoE 用路由器让每个 token 只经过少数专家,从而扩大总参数容量而控制单 token 计算量。它省的是激活的专家计算,不会自动消除通信、存储和负载不均。 读这一节时,先不要急着记缩写,而要不断追问:输入是什么,经过了哪些可观察变换,输出又怎样被验证。
把过程拆开看:
- 路由器产生专家分数。
- 选 Top-k 专家并归一化权重。
- 按专家重排 token。
- 专家计算后按权重合并。
最小例子。 8 个专家、Top-2 路由时,一个 token 只计算两个 FFN;但所有专家权重仍要分布在设备上,且热门专家可能溢出容量。 例子越小,越容易手算、打印中间量并判断实现是否偏离定义。
容易踩坑。 总参数量不能直接当激活参数量;只写 Top-k 而不处理容量、辅助损失、token 丢弃和专家并行,不能构成可训练 MoE。 排错时应保存失败输入和版本,先定位最早出现异常的步骤,再决定是否调整模型或超参数。
动手任务:模拟 100 个 token 的路由计数,计算每个专家负载和变异系数,再尝试加入均衡惩罚。 完成后不要只保存最终输出,还要保存代码、配置、中间张量或检索结果,以及你对结果的解释。
Mixture of Experts 通常把 Dense FFN 替换为多个专家 FFN。路由器为每个 token 计算专家分数,选择 Top-K 专家:
$$ g(x)=\mathrm{softmax}(W_rx),\qquad y=\sum_{e\in\mathrm{TopK}(g)}\tilde{g}_e(x)E_e(x). $$ 总参数量可以很大,但每个 token 只激活少量专家,因此激活参数量和计算量受控。代价是路由、All-to-All 通信、负载不均和专家容量管理。
class ToyMoE(nn.Module):
def __init__(self, d_model=128, hidden=256, n_experts=4, top_k=2):
super().__init__()
self.top_k = top_k
self.router = nn.Linear(d_model, n_experts, bias=False)
self.experts = nn.ModuleList([
nn.Sequential(nn.Linear(d_model, hidden), nn.GELU(),
nn.Linear(hidden, d_model))
for _ in range(n_experts)
])
def forward(self, x):
shape = x.shape
flat = x.reshape(-1, shape[-1])
probs = torch.softmax(self.router(flat), dim=-1)
weights, indices = probs.topk(self.top_k, dim=-1)
weights = weights / weights.sum(dim=-1, keepdim=True)
out = torch.zeros_like(flat)
# 教学写法:逐专家分派;生产实现会使用分组、融合内核和跨卡通信
for expert_id, expert in enumerate(self.experts):
token_pos, slot = torch.where(indices == expert_id)
if token_pos.numel() == 0:
continue
expert_out = expert(flat[token_pos])
out[token_pos] += weights[token_pos, slot, None] * expert_out
return out.view(shape), probs
仅有 Top-K 还不够。若路由器把多数 token 送给少数专家,热门专家会溢出,冷门专家学不到东西。常见做法包括辅助负载均衡损失、容量因子、路由噪声和无辅助损失的动态偏置策略。负载均衡不是越均匀越好;专家适度分工是 MoE 的价值,目标是避免塌缩和硬件闲置。
2.7 复杂度与常见误区¶
从问题出发。 复杂度分析必须说明变量和瓶颈。标准注意力的分数矩阵随 T² 增长,但短序列时 FFN、投影、kernel 启动和内存访问可能占主导。 先把名词放到一边,沿着输入、状态变化和输出走一遍;只有能指出证据落在哪一步,概念才算真正掌握。
沿数据流逐步检查:
- 分别估算参数量、FLOPs 与激活。
- 区分训练和推理。
- 区分 prefill 与 decode。
- 用 profiler 验证理论判断。
用小数据走一遍。 Attention 分数约需 O(T²D),FFN 约需 O(TD·Dff)。T 增大时前者更快增长,但真实速度还受实现与硬件影响。 先在纸面预测结果,再让程序打印中间状态;预测与运行不一致的地方,正是需要继续追查的知识缺口。
这里最容易出现的误解。 把大 O 当实际耗时、忽略 batch/head/dtype、声称 FlashAttention 改变数学结果,或用参数量推断显存,都是错误。 出现异常时先冻结样本、配置和环境,从最靠近输入的环节开始验证;不要同时更换模型、数据和超参数。
小实验:选择三组 T 和 D,计算注意力分数矩阵元素数与 FFN 乘加量,画出交叉趋势。 提交物应包含预期、运行证据和失败记录;只有别人能按同样步骤复现,结果才具有学习价值。
标准注意力的分数矩阵大小为 \(T\times T\),长序列下时间和显存压力显著;FFN 的计算量则随 \(T\) 线性增长但常占大量 FLOPs。FlashAttention 改变内存访问方式,不改变精确注意力的数学结果;线性注意力或稀疏注意力则会改变计算结构。
- “注意力权重就是解释”:权重可提供线索,但不是充分的因果解释。
- “多头越多越好”:固定隐藏维时,头数增加会减小单头维度,存在表达与效率折中。
- “MoE 的总参数都参与一次推理”:应区分总参数和每 token 激活参数。
- “Decoder 训练也必须逐 token”:训练可并行计算所有位置,推理生成才逐 token。
2.8 练习与资料¶
先看它解决什么。 练习应能证明你会推导、实现和验证,而不是只会复述。一个合格的注意力练习至少包含形状断言、mask 测试、数值稳定性和与框架结果对齐。 理解的标准不是会复述定义,而是能画出数据流、预测中间结果,并设计一个让错误暴露出来的检查。
可以把实现分成以下环节:
- 先手算极小输入。
- 再写无多头版本。
- 扩展为多头与 mask。
- 与 PyTorch 官方实现比较。
一个可以手算的例子。 固定同一组投影权重,将自写模块与 scaled_dot_product_attention 输出比较,最大绝对误差应在 dtype 合理范围内。 这里故意不用大模型或大数据,因为小输入能把每个轴、分数和状态变化完整暴露出来。
排错时先看这些地方。 只看 loss 是否下降不能证明注意力实现正确;错误 mask 有时仍能在小数据上拟合。 修复不能止于‘这次跑通’,还要把失败样本变成自动测试,防止同类问题在下一版本重新出现。
现在动手:完成一个测试文件,至少覆盖因果遮蔽、padding、不同 batch、不同头数和半精度有限值。 把关键断言写进测试,并在 README 说明怎样运行、怎样判断正确、当前实现还不支持什么。
练习:用小矩阵手算一次 \(QK^\top\)、缩放、掩码、Softmax 和加权求和;为注意力代码加入 Dropout 与 KV Cache;统计 ToyMoE 的专家负载并设计一个辅助损失。
- Attention Is All You Need
- PyTorch
scaled_dot_product_attention - Switch Transformers:稀疏专家模型
- B 站:6 分钟理解词嵌入与注意力
本章配套代码¶
下面的脚本与正文使用相同符号。建议先在代码中打印形状和中间量,再运行断言;若依赖尚未安装,至少先阅读入口函数、输入输出和测试部分。
examples/02_attention.py:带 mask 的多头注意力。examples/02_tokenizer_and_mask.py:token、padding、因果 mask 与标签右移。
本章端到端实验:把知识变成可复现证据¶
本实验不是把本章代码重新抄一遍,而是把概念、实现、测试和解释串成一个小型工程。请新建独立目录,保存 README.md、环境文件、源代码、测试、运行日志和结果图。README 至少说明任务、输入输出、运行命令、预期现象、已知限制和复现条件。
实验步骤¶
- 用任意公开 tokenizer 编码中英文、代码和数字,比较 token 数并解释差异。
- 在手写注意力中返回权重,验证每行和为 1,并检查未来位置权重为 0。
- 画出三种架构的数据流,并为分类、翻译、开放式生成各选择一种架构说明理由。
- 记录一个 Transformer Block 经过 Attention、残差、FFN 后的均值、标准差与梯度范数。
- 构造长度不同的两条序列,完成 padding、attention mask 和 label mask,并逐位置核对。
- 模拟 100 个 token 的路由计数,计算每个专家负载和变异系数,再尝试加入均衡惩罚。
每完成一步,先写下预期,再运行代码。若结果与预期不一致,不要覆盖旧日志;建立 failures.md,记录现象、假设、证据、修复和回归测试。这样得到的不是一次性 Demo,而是一份能证明你真正理解本章内容的实验档案。
验收标准¶
- 全新环境能够按照 README 从头运行,依赖和随机种子已记录。
- 关键函数至少有正常、边界和错误输入三类测试;涉及数值计算时检查有限值与合理误差。
- 结果包含一个基线和至少一个受控改动,能够说明变化来自哪里。
- 日志保留输入规模、耗时、内存或显存、软件版本和失败样本。
- 结论区分“实验直接证明的事实”“根据事实做出的推断”和“仍未验证的猜想”。
本章自测¶
- 不看正文,用自己的话解释“语言模型不能直接处理字符串,必须先把文本切成 token id,再通过嵌入表把离散 id 映射为连续向量”,并给出一个可以证伪的测试。
- 不看正文,用自己的话解释“Self-Attention 让每个位置根据内容选择其他位置的信息”,并给出一个可以证伪的测试。
- 不看正文,用自己的话解释“Encoder 让所有位置双向交互,Decoder 的自注意力必须因果遮蔽”,并给出一个可以证伪的测试。
- 不看正文,用自己的话解释“注意力负责 token 之间通信,FFN 负责每个 token 内部的非线性变换”,并给出一个可以证伪的测试。
- 不看正文,用自己的话解释“自回归模型训练时一次并行预测所有位置的下一个 token,推理时却必须逐 token 生成”,并给出一个可以证伪的测试。
- 不看正文,用自己的话解释“MoE 用路由器让每个 token 只经过少数专家,从而扩大总参数容量而控制单 token 计算量”,并给出一个可以证伪的测试。
回答时先画图或写形状,再给结论。若只能说出术语而不能给出最小例子、边界条件和验证方法,说明这一节仍需要回到代码中重做。