02 - 模型结构从零实现
这一篇产出一个 model.py。跑完验收脚本,它的参数量要正好是 502,193,664,前向能出 logits,初始 loss 落在 10.37 附近。
不用 from_pretrained,不用 transformers 的任何模型类。整个文件只依赖 torch 和 torch.nn,四百行以内。
需要的前置:知道矩阵乘法,写过一点 PyTorch(会 nn.Linear 和 forward 就够)。01 篇的 train.bin 这一篇用不上,纯写模型。
零、开始之前:Transformer 到底在算什么
0.1 任务还是 01 篇那个任务
再确认一遍目标,因为整个模型结构都是围着它设计的:看着前面的 token,猜下一个 token。
输入是一串 token id,输出是每个位置上「下一个 token 是词表里哪个」的概率分布。就这样。
0.2 数据在模型里走一遍
先不管内部细节,看形状怎么变。设 batch 大小为 B、序列长度为 T:
| 步骤 | 张量形状 | 在干什么 |
|---|---|---|
| 输入 | (B, T) | 整数,每个数是一个 token id |
| Embedding 查表 | (B, T, 1536) | 每个 id 换成一个 1536 维向量 |
| Block × 18 | (B, T, 1536) | 形状不变,内容被反复加工 18 次 |
| 最后的 RMSNorm | (B, T, 1536) | 归一化 |
| lm_head | (B, T, 32000) | 投影到词表大小,得到 logits |
画成图:
(B, T, 1536),这就是 Transformer 能随意堆深的结构性原因。最上和最下的 Embedding 与 lm_head 共享同一组权重(6.3 节)。关键在于中间 18 层形状完全不变。每一层的输入输出都是 (B, T, 1536),所以可以随便堆几层。这是 Transformer 能做深的结构性原因。
最后那个 (B, T, 32000) 叫 logits,每个位置一个长度 32000 的向量,softmax 之后就是概率分布。
0.3 一个 Block 里有什么
每层 Block 干两件事,顺序固定:
分工可以这么理解。Attention 负责 token 之间的信息交换:第 5 个位置想知道第 2 个位置说了什么,靠它。MLP 负责每个 token 自己的加工:拿到信息之后做非线性变换,它对每个位置独立操作,位置之间不通信。
那两个圆圈是残差连接,也就是 x = x + f(x)。它的作用是给梯度留一条直通的路。没有它,18 层的梯度传到第一层基本就没了。这个结构 02 篇不展开推导,记住「残差是恒等通路」就够用。
0.4 这一篇要写的五个部件
| 部件 | 作用 | 在哪一节 |
|---|---|---|
| RMSNorm | 归一化,稳住数值 | 第二节 |
| RoPE | 告诉模型 token 的位置 | 第三节 |
| GQA Attention | token 之间交换信息 | 第四节 |
| SwiGLU | 每个 token 自己加工 | 第五节 |
| Block / LLM | 把上面四个拼起来 | 第六节 |
每一节的套路都一样:先说这个部件解决什么问题、不要它会怎样,再讲它怎么做,最后给代码。
一、为什么照抄 Llama,改了 GPT-2 的哪五处
不做架构创新,结构照抄 Llama。理由很实在:出了问题可以直接跟现成实现对照排查,而且将来想加载别人的权重也方便。
相对于最经典的 GPT-2,Llama 改了五处。这五处正好就是第二到第六节的内容:
| 位置 | GPT-2 | Llama(我们用的) | 换掉的理由 |
|---|---|---|---|
| 归一化 | LayerNorm | RMSNorm | 少算一个均值,快 7% 左右,效果不掉 |
| 位置编码 | 可学习的绝对位置 | RoPE | 能表达相对位置,且能外推到更长序列 |
| 注意力 | MHA | GQA | KV cache 直接小 3 倍 |
| FFN | GELU,中间维 4d | SwiGLU,中间维 8/3 d | 同参数量下效果更好 |
| norm 位置 | Post-norm | Pre-norm | 深层训练稳定得多 |
另外所有线性层都不带 bias。原因在 2.3 节顺带说。
二、RMSNorm
2.1 为什么需要归一化
神经网络堆深了会有个麻烦:每一层的输出分布会漂。第一层输出的数值范围可能是 ±1,传到第十层可能变成 ±100,再往后可能溢出,或者反过来缩到接近 0 梯度消失。
归一化就是在每层入口把数值拉回一个稳定的范围,让后面的层总是面对差不多尺度的输入。
2.2 LayerNorm 在做什么
LayerNorm 对每个 token 的 1536 维向量,做四件事:
减均值 、除标准差 、乘一个可学习的缩放 、加一个可学习的偏置 。
注意这里是对每个 token 自己的 1536 维算均值和方差,不是跨 batch 也不是跨序列。所以它跟 batch size 无关,这点比 BatchNorm 好用。
2.3 RMSNorm 砍掉了什么
RMSNorm 的做法是把「减均值」和「加偏置」都去掉,只留缩放:
分母那一坨就是均方根(root mean square),名字由此而来。
为什么能砍掉减均值这一步,是这一节的关键。RMSNorm 原论文的观察是:LayerNorm 真正起作用的是「把向量缩放到统一尺度」这件事,而不是「把中心移到 0」。做了消融实验,去掉中心化之后效果基本不掉,但省掉了一遍求均值和一遍减法。
砍掉 bias 也是同样的道理,实测加不加差别很小,而少一组参数就少一份显存和一次加法。这也是为什么整个模型的所有 nn.Linear 都写 bias=False。
省下来的计算量不算大,但归一化在每层要做两次、18 层就是 36 次,累积起来推理能快百分之几。免费的收益没理由不要。
2.4 代码
class RMSNorm(nn.Module):
def __init__(self, dim: int, eps: float = 1e-5):
super().__init__()
self.eps = eps
self.weight = nn.Parameter(torch.ones(dim)) # 就是公式里的 gamma
def forward(self, x: torch.Tensor) -> torch.Tensor:
# 统计量始终用 fp32 算,否则 BF16 下 x^2 容易损失精度
dtype = x.dtype
x = x.float()
x = x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + self.eps)
return x.to(dtype) * self.weight
有个容易忽略的细节:中间统计量要转成 fp32 再算。BF16 只有 8 位尾数,x.pow(2) 之后动态范围会被压缩得很厉害,直接在 BF16 上求和容易丢精度。转 fp32 算完再转回来,代价可以忽略,但能避免一类很难查的数值问题。
weight 初始化成全 1,也就是一开始不做任何缩放,让模型自己学。
2.5 Pre-norm 还是 Post-norm
同样是 RMSNorm,放的位置不同,训练难度差很远。
Post-norm(原始 Transformer 论文的做法)是先过子层再归一化:
x = Norm(x + Attention(x))
Pre-norm(现在几乎所有大模型的做法)是先归一化再过子层:
x = x + Attention(Norm(x))
差别在残差通路。Pre-norm 的写法里,x 是一路直接加过去的,从最后一层到第一层存在一条完全没有归一化操作的恒等通路,梯度可以无损地流回去。Post-norm 里每一层的残差都要再过一次 Norm,梯度传递会被反复缩放,层数一深就容易出问题。
代价是 Pre-norm 每层输出的方差会随层数累加,所以最后要额外补一个 RMSNorm(就是 0.2 节表里倒数第二行那个),把进 lm_head 之前的数值拉回来。
Norm(x + f(x)) 和 x + f(Norm(x)) 很难感觉到差别,画成通路就一目了然:前者的主干被 Norm 切成了 18 段,后者的主干是通的。深层网络能不能训,很大程度上就取决于反向传播时有没有这样一条不被打断的路。