Skip to main content

02 - 模型结构从零实现

这一篇产出一个 model.py。跑完验收脚本,它的参数量要正好是 502,193,664,前向能出 logits,初始 loss 落在 10.37 附近。

不用 from_pretrained,不用 transformers 的任何模型类。整个文件只依赖 torchtorch.nn,四百行以内。

需要的前置:知道矩阵乘法,写过一点 PyTorch(会 nn.Linearforward 就够)。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

关键在于中间 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 Attentiontoken 之间交换信息第四节
SwiGLU每个 token 自己加工第五节
Block / LLM把上面四个拼起来第六节

每一节的套路都一样:先说这个部件解决什么问题、不要它会怎样,再讲它怎么做,最后给代码。

一、为什么照抄 Llama,改了 GPT-2 的哪五处

不做架构创新,结构照抄 Llama。理由很实在:出了问题可以直接跟现成实现对照排查,而且将来想加载别人的权重也方便。

相对于最经典的 GPT-2,Llama 改了五处。这五处正好就是第二到第六节的内容:

位置GPT-2Llama(我们用的)换掉的理由
归一化LayerNormRMSNorm少算一个均值,快 7% 左右,效果不掉
位置编码可学习的绝对位置RoPE能表达相对位置,且能外推到更长序列
注意力MHAGQAKV cache 直接小 3 倍
FFNGELU,中间维 4dSwiGLU,中间维 8/3 d同参数量下效果更好
norm 位置Post-normPre-norm深层训练稳定得多

另外所有线性层都不带 bias。原因在 2.3 节顺带说。

二、RMSNorm

2.1 为什么需要归一化

神经网络堆深了会有个麻烦:每一层的输出分布会漂。第一层输出的数值范围可能是 ±1,传到第十层可能变成 ±100,再往后可能溢出,或者反过来缩到接近 0 梯度消失。

归一化就是在每层入口把数值拉回一个稳定的范围,让后面的层总是面对差不多尺度的输入。

2.2 LayerNorm 在做什么

LayerNorm 对每个 token 的 1536 维向量,做四件事:

LayerNorm(x)=xμσ2+ϵγ+β\text{LayerNorm}(x) = \frac{x - \mu}{\sqrt{\sigma^2 + \epsilon}} \cdot \gamma + \beta

减均值 μ\mu、除标准差 σ\sigma、乘一个可学习的缩放 γ\gamma、加一个可学习的偏置 β\beta

注意这里是对每个 token 自己的 1536 维算均值和方差,不是跨 batch 也不是跨序列。所以它跟 batch size 无关,这点比 BatchNorm 好用。

2.3 RMSNorm 砍掉了什么

RMSNorm 的做法是把「减均值」和「加偏置」都去掉,只留缩放:

RMSNorm(x)=x1di=1dxi2+ϵγ\text{RMSNorm}(x) = \frac{x}{\sqrt{\frac{1}{d}\sum_{i=1}^{d} x_i^2 + \epsilon}} \cdot \gamma

分母那一坨就是均方根(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 之前的数值拉回来。

三、RoPE 旋转位置编码

3.1 Attention 根本看不见位置

这是一个不那么直观但很重要的事实:attention 本身对输入顺序是无感的

原因在于 attention 的计算方式。每个位置的输出是所有位置的 value 的加权和,权重由 query 和 key 的内积决定。这里面没有任何一项跟「第几个位置」有关。把输入序列打乱顺序,输出也只是跟着打乱,内容完全一样。

所以「我打你」和「你打我」在纯 attention 眼里是同一个东西。必须额外把位置信息喂进去。

3.2 早期做法和它的问题

GPT-2 的做法是搞一个可学习的位置向量表,大小是 (max_seq_len, d_model),第 t 个位置就查第 t 行,加到 embedding 上。

两个问题。一是只能处理训练时见过的长度,表只有 1024 行,第 1025 个位置查不到,模型直接不能用。二是它编码的是绝对位置,但语言里真正重要的往往是相对距离,「形容词修饰它后面第一个名词」这种规律跟这个词在全文第几个位置没关系。

3.3 RoPE 的想法:把位置变成旋转角度

RoPE 的做法很巧。它不给 embedding 加东西,而是在 attention 内部,把 query 和 key 按位置旋转一个角度

具体做法:把 128 维的 head 向量两两分成 64 组,每组当成一个二维平面上的向量。位置 mm 处的向量,第 ii 组按角度 mθim\theta_i 旋转:

(xiyi)=(cosmθisinmθisinmθicosmθi)(xiyi)\begin{pmatrix} x_i' \\ y_i' \end{pmatrix} = \begin{pmatrix} \cos m\theta_i & -\sin m\theta_i \\ \sin m\theta_i & \cos m\theta_i \end{pmatrix} \begin{pmatrix} x_i \\ y_i \end{pmatrix}

位置越靠后,转得越多。

3.4 为什么旋转能表达相对位置

这是 RoPE 的精髓,值得推一下。

二维旋转矩阵有个性质:R(α)R(β)=R(βα)R(\alpha)^\top R(\beta) = R(\beta - \alpha)。也就是说一个旋转的转置乘另一个旋转,等于「差角」的旋转。

现在看 attention 里的内积。位置 mm 的 query 和位置 nn 的 key,各自旋转之后再做内积:

(R(mθ)q)(R(nθ)k)=qR(mθ)R(nθ)k=qR((nm)θ)k(R(m\theta)q)^\top (R(n\theta)k) = q^\top R(m\theta)^\top R(n\theta) k = q^\top R((n-m)\theta) k

结果里只剩下 nmn-m绝对位置 mmnn 都消掉了,只留下它们的差。

这正是我们想要的:模型感知到的是「这两个 token 隔多远」,而不是「它们分别在第几个位置」。

原理讲完要验。code/rope_demo.py 是个纯 Python 的验证,固定两个向量,固定相对距离 3,挪动绝对位置:

固定相对距离 n-m=3,挪动绝对位置:
m n n-m 旋转后内积
0 3 3 -1.7349277069
1 4 3 -1.7349277069
2 5 3 -1.7349277069
3 6 3 -1.7349277069
4 7 3 -1.7349277069
5 8 3 -1.7349277069

改变相对距离,内积才跟着变:
n-m 旋转后内积
0 -0.9800000000
1 -1.7064616876
2 -2.0008180700
3 -1.7349277069
4 -1.2924048008
5 -1.2145554743

小数点后 10 位完全一致。相对距离一变,内积立刻跟着变。

3.5 频率怎么定

64 组用同一个角度是不行的,那样只能表达一种尺度的距离。RoPE 给每组配不同的频率:

θi=base2i/d,base=10000\theta_i = \text{base}^{-2i/d}, \quad \text{base} = 10000

第 0 组频率是 1,转得最快,每挪一个位置就转 1 弧度,用来分辨近距离。最后一组频率是 100001=0.000110000^{-1} = 0.0001,转得极慢,几千个位置才转完一圈,用来分辨远距离。

这样 64 组合起来,就能同时表达从「隔 1 个」到「隔几千个」的各种距离。跟时钟的秒针分针时针是一个道理。

base 这个数还有个用处:把它调大能让所有频率变慢,等效于把位置「压缩」,这是长上下文外推最常用的手段之一。我们这篇不展开,但要知道 rope_theta 这个超参是干这个的。

3.6 代码

def build_rope_cache(seq_len, head_dim, theta, device, dtype):
"""预计算 cos / sin 表,形状都是 (seq_len, head_dim)"""
idx = torch.arange(0, head_dim, 2, device=device).float() / head_dim
inv_freq = 1.0 / (theta ** idx) # (head_dim/2,)
pos = torch.arange(seq_len, device=device).float() # (seq_len,)
freqs = torch.outer(pos, inv_freq) # (seq_len, head_dim/2)
emb = torch.cat((freqs, freqs), dim=-1) # (seq_len, head_dim)
return emb.cos().to(dtype), emb.sin().to(dtype)


def rotate_half(x):
half = x.shape[-1] // 2
x1, x2 = x[..., :half], x[..., half:]
return torch.cat((-x2, x1), dim=-1)


def apply_rope(q, k, cos, sin):
cos = cos[None, None, :, :] # 广播到 (B, n_heads, T, head_dim)
sin = sin[None, None, :, :]
return q * cos + rotate_half(q) * sin, k * cos + rotate_half(k) * sin

cossin 只跟位置有关,跟输入内容无关,所以建模型的时候算一次存下来就行,不用每次前向重算。代码里用 register_buffer 挂在模型上,它不是参数、不参与梯度,但会跟着 .to(device) 一起搬到 GPU。

3.7 一个必须知道的坑:两套约定

rotate_half 的写法是把 128 维切成前 64 和后 64,配对方式是「第 0 维配第 64 维」。这是 HuggingFace 的约定。

RoPE 原论文用的是另一套:相邻两维配对,也就是「第 0 维配第 1 维」。

两套都能正常训练,效果也一样,但它们互不兼容。 用 A 约定训出来的权重,拿 B 约定的代码加载,前向输出全是乱的,而且不会报任何错,只会表现为模型胡说八道。

这个坑在 08 篇会真的碰上,因为那时候要把权重喂给自制推理框架。训练端和推理端必须用同一套约定,现在选了 rotate_half,到时候推理框架也得是 rotate_half

四、GQA 分组查询注意力

4.1 先说 MHA

标准多头注意力(MHA)的做法:把 1536 维切成 12 个头,每个头 128 维,各自独立算一遍 attention,最后拼回来。

每个头都有自己的 Q、K、V 三个投影。多个头的意义是让不同的头关注不同类型的关系,有的头管语法,有的头管指代。

4.2 推理时的真正瓶颈是 KV cache

训练的时候一次算整个序列,MHA 没什么问题。问题出在推理。

推理是一个 token 一个 token 往外吐的。生成第 100 个 token 时,需要用到前面 99 个 token 的 K 和 V。如果每次都重算,复杂度是平方级的,慢得没法用。

所以实际做法是把算过的 K 和 V 缓存起来,这就是 KV cache。每生成一个 token,就往缓存里追加一份 K 和 V。

问题是这个缓存很占显存,而且跟并发数成正比。服务 100 个用户就要 100 份 KV cache。在推理服务里,KV cache 经常比模型权重本身还占地方。vLLM 的 PagedAttention 解决的就是这块的碎片问题(见 PagedAttention)。

4.3 MQA 和 GQA

既然 KV cache 是瓶颈,一个自然的想法是:能不能让多个头共享同一份 K 和 V

MQA(Multi-Query Attention)走到极端,所有 Q 头共享唯一一组 K/V。缓存直接小 12 倍,但效果掉得比较明显。

GQA(Grouped-Query Attention)取中间:把 Q 头分组,每组共享一组 K/V。我们的配置是 12 个 Q 头、4 个 KV 头,也就是每 3 个 Q 头共享 1 组 K/V。

4.4 我们这个配置省了多少

每个 token 要缓存的字节数:

2×nlayers×nkv_heads×dhead×bytes2 \times n_{\text{layers}} \times n_{\text{kv\_heads}} \times d_{\text{head}} \times \text{bytes}

开头的 2 是因为 K 和 V 各存一份。代进去(BF16 所以每个数 2 字节):

方案n_kv_heads每 token2048 长的一条序列
MHA12108.0 KB226.5 MB
GQA(我们的)436.0 KB75.5 MB

省了整整 3 倍。 这个数字是 code/count_params.py 算出来的。

同时注意 4.1 节说的:Q 头还是 12 个,所以模型的表达能力基本没损失,掉的只是 K/V 的多样性。这是 GQA 的性价比所在,现在从 Llama 2 70B 到 Qwen 系列基本都用它。

顺带看参数量。因为 K/V 投影变窄了,Attention 的参数也跟着少了:

投影形状参数量
q_proj1536 → 12×128 = 15362,359,296
k_proj1536 → 4×128 = 512786,432
v_proj1536 → 4×128 = 512786,432
o_proj1536 → 15362,359,296
合计6,291,456

如果是 MHA,k/v 各是 2,359,296,Attention 合计会变成 9,437,184。

4.5 代码

Q 头和 KV 头数量对不上,算 attention 之前要把 KV 复制对齐:

def repeat_kv(x, n_rep):
"""(B, n_kv_heads, T, hd) -> (B, n_kv_heads * n_rep, T, hd)"""
b, n_kv, t, hd = x.shape
if n_rep == 1:
return x
return (x[:, :, None, :, :]
.expand(b, n_kv, n_rep, t, hd)
.reshape(b, n_kv * n_rep, t, hd))

这里用 expand 而不是 repeat,因为 expand 不真的复制数据,只是改 stride。真正的 KV cache 里只存 4 份,复制只发生在计算的那一刻。 要是这里写成 repeat,显存收益就没了。

完整的 Attention:

class Attention(nn.Module):
def __init__(self, c: Config):
super().__init__()
self.n_heads, self.n_kv_heads = c.n_heads, c.n_kv_heads
self.head_dim = c.head_dim
self.n_rep = c.n_heads // c.n_kv_heads # 12 // 4 = 3

self.q_proj = nn.Linear(c.d_model, c.n_heads * c.head_dim, bias=False)
self.k_proj = nn.Linear(c.d_model, c.n_kv_heads * c.head_dim, bias=False)
self.v_proj = nn.Linear(c.d_model, c.n_kv_heads * c.head_dim, bias=False)
self.o_proj = nn.Linear(c.n_heads * c.head_dim, c.d_model, bias=False)

def forward(self, x, cos, sin):
B, T, _ = x.shape

q = self.q_proj(x).view(B, T, self.n_heads, self.head_dim).transpose(1, 2)
k = self.k_proj(x).view(B, T, self.n_kv_heads, self.head_dim).transpose(1, 2)
v = self.v_proj(x).view(B, T, self.n_kv_heads, self.head_dim).transpose(1, 2)

q, k = apply_rope(q, k, cos[:T], sin[:T]) # RoPE 只作用于 q 和 k
k, v = repeat_kv(k, self.n_rep), repeat_kv(v, self.n_rep)

out = F.scaled_dot_product_attention(q, k, v, is_causal=True)

out = out.transpose(1, 2).contiguous().view(B, T, -1)
return self.o_proj(out)

两个要点。

RoPE 只加在 q 和 k 上,不加在 v 上。 因为位置信息是通过内积起作用的(3.4 节那个推导),而 v 不参与内积,它只是被加权求和。给 v 加旋转没有意义。

is_causal=True 自带因果掩码。 因果的意思是第 t 个位置只能看到 0 到 t,不能看到未来。F.scaled_dot_product_attention 会自动生成下三角掩码,而且它会走 FlashAttention 的融合实现,比自己建一个 (T, T) 的 mask 矩阵又快又省显存。序列长 2048 时,显式 mask 光自己就要 8 MB。

4.6 顺带说 head_dim 为什么是 128

head_dim = d_model / n_heads = 1536 / 12 = 128

128 是个几乎所有主流模型都在用的数。原因是 FlashAttention 这类融合 kernel 是按 64 / 128 这些尺寸做了专门优化的,取别的值会掉到通用实现上,慢不少。定超参的时候,让 head_dim 落在 64 或 128 是条实用的约束。

五、SwiGLU

5.1 标准 FFN 长什么样

原始 Transformer 的前馈层就两层线性加一个激活:

def ffn(x):
return W2(gelu(W1(x))) # 1536 -> 4096 -> 1536

先升维再降维,中间那层一般取 4 倍宽。它对每个位置独立操作,是模型存储「知识」的主要地方。

5.2 门控的想法

SwiGLU 的改动是加一条门控支路:

SwiGLU(x)=Wdown(SiLU(Wgatex)Wupx)\text{SwiGLU}(x) = W_{\text{down}}\big(\text{SiLU}(W_{\text{gate}}\,x) \odot W_{\text{up}}\,x\big)

\odot 是逐元素相乘。直观理解:up 那一路算出候选值,gate 那一路算出「每一维该放行多少」,两者相乘。SiLU 的输出可以接近 0,等于把某些维度关掉。

相比固定的激活函数,门控让「哪些信息通过」变成数据相关的,模型能学得更细。

这类改动 Noam Shazeer 那篇论文(GLU Variants Improve Transformer)做了一组消融,结论是同参数量下 SwiGLU 稳定地好一点。没有特别深刻的理论解释,论文里那句结论大意是把它归功于运气。工程上照用就是了。

5.3 为什么中间维是 8/3 d,不是 4 d

这是个容易被忽略但很实际的问题。

标准 FFN 两个矩阵,参数量是 2×d×dffn2 \times d \times d_{\text{ffn}}。取 dffn=4dd_{\text{ffn}} = 4d 时是 8d28d^2

SwiGLU 有三个矩阵,参数量是 3×d×dffn3 \times d \times d_{\text{ffn}}。要让参数量保持在 8d28d^2 不变:

3×d×dffn=8d2    dffn=83d3 \times d \times d_{\text{ffn}} = 8d^2 \implies d_{\text{ffn}} = \frac{8}{3}d

所以 83×1536=4096\frac{8}{3} \times 1536 = 4096,正好是整数,也正好是 2 的幂,对齐得很舒服。多一个矩阵就把中间维按比例调窄,总参数量不变,这样跟标准 FFN 比才是公平的。

5.4 代码

class SwiGLU(nn.Module):
def __init__(self, c: Config):
super().__init__()
self.gate_proj = nn.Linear(c.d_model, c.ffn_dim, bias=False)
self.up_proj = nn.Linear(c.d_model, c.ffn_dim, bias=False)
self.down_proj = nn.Linear(c.ffn_dim, c.d_model, bias=False)

def forward(self, x):
return self.down_proj(F.silu(self.gate_proj(x)) * self.up_proj(x))

六、组装

6.1 Block

按 0.3 节那张图,Pre-norm 加两条残差:

class Block(nn.Module):
def __init__(self, c: Config):
super().__init__()
self.attn_norm = RMSNorm(c.d_model, c.norm_eps)
self.attn = Attention(c)
self.mlp_norm = RMSNorm(c.d_model, c.norm_eps)
self.mlp = SwiGLU(c)

def forward(self, x, cos, sin):
x = x + self.attn(self.attn_norm(x), cos, sin)
x = x + self.mlp(self.mlp_norm(x))
return x

注意 x = x + f(norm(x)),不是 x = f(norm(x))。少写那个 x + 就没有残差了,模型能跑但训不深。

6.2 完整模型

class LLM(nn.Module):
def __init__(self, c: Config):
super().__init__()
self.config = c
self.embed = nn.Embedding(c.vocab_size, c.d_model)
self.blocks = nn.ModuleList(Block(c) for _ in range(c.n_layers))
self.norm = RMSNorm(c.d_model, c.norm_eps)
self.lm_head = nn.Linear(c.d_model, c.vocab_size, bias=False)

if c.tie_embeddings:
self.lm_head.weight = self.embed.weight

cos, sin = build_rope_cache(c.max_seq_len, c.head_dim, c.rope_theta,
device="cpu", dtype=torch.float32)
self.register_buffer("cos", cos, persistent=False)
self.register_buffer("sin", sin, persistent=False)

def forward(self, idx, targets=None):
B, T = idx.shape
x = self.embed(idx)
cos, sin = self.cos[:T].to(x.dtype), self.sin[:T].to(x.dtype)
for blk in self.blocks:
x = blk(x, cos, sin)
x = self.norm(x)
logits = self.lm_head(x)

if targets is None:
return logits, None
loss = F.cross_entropy(logits.view(-1, logits.size(-1)), targets.reshape(-1))
return logits, loss

persistent=False 的意思是这两个 buffer 不存进 checkpoint。它们是纯计算出来的,加载时重新算一遍就行,没必要占 checkpoint 的体积。

6.3 权重共享

self.lm_head.weight = self.embed.weight 这一行让输入的 embedding 表和输出的投影矩阵用同一组权重。

两者形状正好都是 (32000, 1536),物理上可以共享。直觉上也说得通:embedding 是「token id 到语义向量」,lm_head 是「语义向量到 token id」,互为逆向,共享一组参数是合理的。

收益是省 49,152,000 个参数,占总量 9.8%。对小模型来说这个比例不小,所以小模型基本都开权重共享,大模型反而常常不开,因为那时候 embedding 占比已经很低了。

七、初始化

7.1 为什么不能随便初始化

初始权重全设成 0,所有神经元输出一样、梯度一样,永远学不出差异。设得太大,前向传几层就爆掉。设得太小,信号一路衰减到 0。

通行做法是从均值 0、标准差 0.02 的正态分布里采样。0.02 这个数来自 GPT-2,后来大家一直沿用。

7.2 残差累积和缩放

有个专门针对残差结构的补充处理,值得说清楚。

Pre-norm 的每一层都是 x = x + f(x)。假设每个 ff 的输出方差是 σ2\sigma^2,那么 18 层加下来,方差会累积到大约 18σ218\sigma^2。层数越多,进 lm_head 之前的数值越大。

GPT-2 的对策是:把每个残差分支最后那个投影的初始化标准差按 1/2L1/\sqrt{2L} 缩小(L 是层数,2 是因为每层有 attention 和 mlp 两条残差分支)。

for name, p in self.named_parameters():
if name.endswith(("o_proj.weight", "down_proj.weight")):
nn.init.normal_(p, mean=0.0, std=0.02 / (2 * c.n_layers) ** 0.5)

o_projdown_proj 正好就是 attention 和 mlp 各自的出口。缩放之后,无论堆多少层,累积方差都维持在同一量级。

不做这一步模型也能训,但初期 loss 会更抖,学习率要设得更保守。这是一行代码就能拿到的稳定性,没理由不做。

八、验收

写完必须验,四项都过了才进 03 篇。验收脚本是 code/verify_model.py

cd code && python verify_model.py

8.1 参数量对账

code/count_params.py 是个纯 Python 的独立实现,不 import torch,按公式逐项算。拿它的结果跟 model.num_params() 对,两边独立算出同一个数才说明没写错。

它的实际输出:

配置: d_model=1536 n_layers=18 n_heads=12 n_kv_heads=4 head_dim=128 ffn=4096
ffn / d_model = 2.6667 (8/3 = 2.6667)

单层 Block 内部
Attention (q/k/v/o) 6,291,456 25.0%
SwiGLU MLP 18,874,368 75.0%
2 × RMSNorm 3,072 0.0%
单层合计 25,168,896

整个模型
Embedding(与输出层共享) 49,152,000 9.8%
18 层 Block 453,040,128 90.2%
最后的 RMSNorm 1,536 0.0%
合计 502,193,664 (502.2 M)

502,193,664,跟 00 篇定的配置对上。

值得注意的是 MLP 占了单层参数的 75%。这是 Transformer 的常态,也是为什么做量化、剪枝、MoE 的时候都优先动 MLP,那里油水最多。

8.2 前向形状

输入 (2, 128) 的整数张量,输出必须是 (2, 128, 32000)

这一项主要抓 view / transpose 写错的问题。GQA 那段的形状变换比较绕,n_headsn_kv_heads 很容易填反,填反了 shape 就对不上。

8.3 初始 loss 必须落在 10.37 附近

这是最便宜也最有用的一项检查。

刚初始化的模型对 32000 个候选是均匀猜的,此时交叉熵就是均匀分布的熵:

loss=ln132000=ln32000=10.3735\text{loss} = -\ln\frac{1}{32000} = \ln 32000 = 10.3735

跑出来的数偏离 10.37 太多,说明有问题:

  • 明显偏高(比如 15、20):初始化标准差太大,或者某处数值爆了
  • 明显偏低(比如 7、8):更危险,说明模型能看到不该看到的信息,八成是因果掩码没生效

第二种情况尤其要警惕,因为它的表现是「loss 降得特别快」,看曲线像是训练特别顺利,实际是模型在抄答案。

8.4 因果性

这一项专门抓 8.3 里说的第二种错误,做法很直接:

改动第 t 个位置之后的输入,第 t 个位置的输出必须一点都不变

base, _ = model(idx)
modified = idx.clone()
modified[:, t + 1:] = torch.randint(0, vocab_size, (B, T - t - 1))
after, _ = model(modified)

# 位置 0..t 的输出必须完全一致
assert (base[:, :t + 1] - after[:, :t + 1]).abs().max() < 1e-5
# 位置 t+1.. 的输出必须变了(否则说明输入压根没改动,测试本身失效)
assert (base[:, t + 1:] - after[:, t + 1:]).abs().max() > 1e-3

第二个断言容易被忽略但必须有。只写第一个断言的话,如果测试代码本身有 bug(比如没真的改到输入),断言会平凡地通过,测了等于没测。

8.5 顺带看 KV cache 的账

count_params.py 还会打印 KV cache 的估算,这个数 08 篇接推理框架时要用:

KV cache(BF16)
GQA n_kv_heads=4 36,864 B/token = 36.0 KB 序列 2048 时 75.5 MB/条
MHA n_kv_heads=12 110,592 B/token = 108.0 KB 序列 2048 时 226.5 MB/条
GQA 省了 3.0 倍

以及静态显存,对应 00 篇 3.3 节的 8 GB:

静态显存(BF16 混合精度 + AdamW)
BF16 权重 1.00 GB
BF16 梯度 1.00 GB
FP32 master weights 2.01 GB
AdamW 动量 m 2.01 GB
AdamW 动量 v 2.01 GB
小计 8.04 GB

九、常见问题

现象大概率是什么原因
初始 loss 远低于 10.37因果掩码没生效,检查 is_causal=True 有没有传
初始 loss 远高于 10.37初始化 std 太大;或者 RMSNorm 忘了转 fp32
参数量比预期多 49Mtie_embeddings 没生效,检查那行赋值在不在
参数量对不上但差得不多多半是某个 nn.Linear 忘了写 bias=False
shape 报错在 attention 里n_headsn_kv_heads 填反了,或者 repeat_kv 漏调用
显存比算的多很多repeat_kv 写成了 repeat,真复制了数据,改回 expand
加载别人的权重后输出是乱码RoPE 约定不一致,见 3.7 节
训到几百步 loss 变 NaN残差缩放初始化没做,或者学习率太高(03 篇细说)
这一篇不需要 GPU

model.py 在 CPU 上就能建起来跑前向,四项验收全都能在本地做完。结构验对了再去租卡,别把调 shape 的时间花在按小时计费的机器上。