MiniMind 学习笔记之 04 注意力之外:位置、记忆、省显存、省时间、深加工

5 阅读9分钟

上承〔03 · 自注意力:模型怎么“看懂”上下文?〕。上一篇把自注意力拆完了:Q、K、V 怎么来,注意力权重怎么算,多头怎么并行。但一个完整的 Transformer 块,不是只有注意力。注意力只是“让 token 互相看”的那一半。

03 篇源码走读时,有几处只点了名、没展开。其中四个——Flash Attention、RoPE、KV Cache、GQA——是注意力的配套:位置、记忆、省显存、省时间。它们在源码里出现,但每个都是独立的工程话题,放在源码走读的语境里讲不透。这一篇把这四个逐一展开。

注意力之外,还有三块没讲过:FFN 负责每个 token 的深加工,残差和 RMSNorm 保证深层能训得动,最后这些零件怎么拼成一个完整的块。这一篇一并补上。

六块拼完,一个完整的 Transformer 块就出来了。

一、位置:RoPE

1.1 自注意力不知道谁在前,谁在后

第三篇讲注意力的时候,一直回避了一个问题。

自注意力的计算是并行的。所有 token 的 Q、K、V 同时算出来,同时做匹配,同时加权汇总。没有先后。

这意味着什么?如果打乱“我”、“爱”、“你”的顺序,注意力机制算出来的结果是一样的。“我爱你”和“你爱我”,在模型眼里暂时没有区别。

但语言是有顺序的。“猫追老鼠”和“老鼠追猫”,字一样,意思相反。模型要能区分,就必须给每个 token 一个信号:你在第几个位置。

这个信号,就是位置编码。

1.2 从正弦编码到 RoPE

大模型刚出来时,位置编码用的是正弦编码(Sinusoidal Positional Encoding),这也是2017 年《Attention Is All You Need》那篇论文提出的。

做法很直接:给每个位置的 token,在它的输入向量上加一个固定的向量。这个向量用正弦和余弦函数算出来,每个位置不同,每个维度不同。

但正弦编码有两个问题。

第一,它是“加”上去的。 位置信息和内容信息混在同一个向量里,模型需要自己学会怎么把它们分开。位置信息容易被内容信息淹没。

第二,它外推能力差。 训练时序列长度是 512,测试时来了 1024,后面那些位置的正弦值模型没见过,效果会变差。【后面会解释训练时的序列长度问题】

所以后来换了方案。2021 年,苏剑林等人提出了旋转位置编码(RoPE,Rotary Position Embedding)。目前主流大模型——Llama、Qwen、DeepSeek、MiniMind——用的都是 RoPE。

1.3 RoPE 的核心思路

RoPE 不往向量里“加”位置,而是把向量旋转一个角度。

角度由位置决定。位置越靠后,旋转角度越大。

具体怎么做?把 Q 和 K 向量里的数两两分组,每组两个数 (x1,x2)(x_1, x_2),看作一个二维平面上的点。然后把这个点按一个角度旋转:

(x1′x2′)=(cos⁡θ−sin⁡θsin⁡θcos⁡θ)(x1x2)\begin{pmatrix} x_1' \\ x_2' \end{pmatrix} = \begin{pmatrix} \cos\theta & -\sin\theta \\ \sin\theta & \cos\theta \end{pmatrix} \begin{pmatrix} x_1 \\ x_2 \end{pmatrix}

旋转角度 θ\theta 和 token 的位置 mm 有关:θ=m⋅ω\theta = m \cdot \omega。位置越靠后,旋转越多。

每一对分配一个不同的频率 ω\omega。低频对旋转得慢,高频对旋转得快。

注意力是按头算位置的: 每个头的向量是 96 维(768 ÷ 8 头),拆成 48 对,就分配 48 个不同的频率。8 个头共用同一张频率表。

这个多频率设计,让不同维度捕捉不同尺度的位置关系。低频维度关注长距离,高频维度关注短距离。

1.4 为什么旋转能编码相对位置?

关键在点积。

两个 token 做注意力匹配时,算的是 q⋅kq \cdot k。经过 RoPE 之后,q 和 k 都旋转了各自的角度。

数学上有一个漂亮的结论:两个旋转后的向量做点积,结果只和它们的旋转角度之差有关。

假设“猫”在位置 mm,“坐”在位置 nn。经过 RoPE 后,“猫”的 q 旋转了 mωm\omega,“坐”的 k 旋转了 nωn\omega。点积的结果,只取决于 (m−n)ω(m - n)\omega,也就是两个 token 的距离。

这意味着什么?注意力分数天然地包含了“这两个 token 隔多远”的信息,而且是相对距离,不是绝对位置。

“猫”在位置 2,“坐”在位置 3,距离是 1。“猫”在位置 100,“坐”在位置 101,距离也是 1。RoPE 算出来的注意力分数,对这两个场景是一样的。

这就是 RoPE 的核心优势:相对位置编码,自然融进点积里。

1.5 用一个具体例子走一遍

前面讲了原理,现在拿具体数字走一遍。

第三篇里,我们拿“他把苹果吃了”举过例。这一节沿用同样的分词假定:

他 | 把 | 苹果 | 吃了
位置 0 | 1 | 2 | 3

四个 token。苹果在位置 2,吃了在位置 3。

假设 head_dim = 4(真实是 96,这里只留两对,方便手算)。每一对分配一个不同的频率:

  • 第 1 对:ω1=1.0\omega_1 = 1.0(高频,转得快)
  • 第 2 对:ω2=0.1\omega_2 = 0.1(低频,转得慢)

“苹果”的 q 和“吃了”的 k,初始都是最简单的向量:

苹果的 q = [1, 0, 1, 0]
吃了的 k = [1, 0, 1, 0]

前两个数是第 1 对,后两个数是第 2 对。

现在对 q 和 k 做 RoPE 旋转。旋转角度 = 位置 × 频率。

“苹果”在位置 2:

  • 第 1 对旋转角度:2×1.0=22 \times 1.0 = 2 弧度
  • 第 2 对旋转角度:2×0.1=0.22 \times 0.1 = 0.2 弧度

“吃了”在位置 3:

  • 第 1 对旋转角度:3×1.0=33 \times 1.0 = 3 弧度
  • 第 2 对旋转角度:3×0.1=0.33 \times 0.1 = 0.3 弧度

每一对分别算

第 1 对(高频,ω1=1.0\omega_1 = 1.0):

苹果的 q 第 1 对 (1, 0),旋转 2 弧度:

q₁' = [cos(2), sin(2)] ≈ [-0.416, 0.909]

吃了的 k 第 1 对 (1, 0),旋转 3 弧度:

k₁' = [cos(3), sin(3)] ≈ [-0.990, 0.141]

点积:

q₁' · k₁' = (-0.416)(-0.990) + (0.909)(0.141) ≈ 0.540

第 2 对(低频,ω2=0.1\omega_2 = 0.1):

苹果的 q 第 2 对 (1, 0),旋转 0.2 弧度:

q₂' = [cos(0.2), sin(0.2)] ≈ [0.980, 0.199]

吃了的 k 第 2 对 (1, 0),旋转 0.3 弧度:

k₂' = [cos(0.3), sin(0.3)] ≈ [0.955, 0.296]

点积:

q₂' · k₂' = (0.980)(0.955) + (0.199)(0.296) ≈ 0.995

总点积:

q' · k' = 0.540 + 0.995 = 1.535

这个 1.535,就是“苹果”和“吃了”的注意力分数,包含位置信息。

换个距离,看分数怎么变

现在把“吃了”挪到位置 9。距离从 1 变成 7。

“吃了”在位置 9:

  • 第 1 对旋转角度:9×1.0=99 \times 1.0 = 9 弧度
  • 第 2 对旋转角度:9×0.1=0.99 \times 0.1 = 0.9 弧度

第 1 对(高频):

吃了的 k 第 1 对 (1, 0),旋转 9 弧度:

k₁' = [cos(9), sin(9)] ≈ [-0.911, 0.412]

点积:

q₁' · k₁' = (-0.416)(-0.911) + (0.909)(0.412) ≈ 0.379 + 0.375 ≈ 0.754

等等,这个值反而比距离 1 时还大了?这就是绕圈的问题。点积只取决于角度差:k 转了 9 弧度,q 转了 2 弧度,差是 7 弧度。7 弧度绕了一圈多(一圈 ≈ 6.28 弧度),7 − 6.28 = 0.72 弧度——比距离 1 时的 1 弧度还小。高频对在这种情况下已经分不清“距离 7”和“距离 0.72”了。

第 2 对(低频):

吃了的 k 第 2 对 (1, 0),旋转 0.9 弧度:

k₂' = [cos(0.9), sin(0.9)] ≈ [0.622, 0.783]

点积:

q₂' · k₂' = (0.980)(0.622) + (0.199)(0.783) ≈ 0.610 + 0.156 ≈ 0.766

总点积:

q' · k' = 0.754 + 0.766 ≈ 1.520

距离 1 时总分是 1.535,距离 7 时总分是 1.520。非常接近。

问题出在哪?高频对绕圈了,把“距离 7”误判成了“距离 0.72”。低频对因为转得慢,角度差只从 0.1 变成 0.9,还在合理范围内,但它对总分的贡献被高频对的误判抵消了一部分。

多频率配合的意义

如果只有高频对,距离一长就绕圈,分不清距离 1 和距离 7。

如果只有低频对,相邻位置的角度差太小(距离 1 时只有 0.1 弧度),区分不了相邻位置。

两个频率一起,高频负责近处,低频负责远处。距离 1 的时候,高频对贡献 0.540,低频对贡献 0.995,总分 1.535。距离 7 的时候,高频对贡献 0.754(虽然绕圈了,但值还是不同),低频对贡献 0.766,总分 1.520。

单个频率会误判,但两个频率的“误判方式”不同。 模型可以通过多个频率的组合,反推出真实距离。这就是多频率设计的核心:每一对都可能绕圈,但绕圈的时机不同,组合起来就能覆盖完整的距离范围。

MiniMind 的 head_dim 是 96,分成 48 对。48 个频率从高到低,覆盖了从“相邻位置”到“整个序列长度”的所有尺度。这就是为什么 RoPE 能编码相对位置。

为什么两两分组

最后回答“为什么是两两分组”。

因为二维平面上的旋转有现成的公式:

(x1′x2′)=(cos⁡θ−sin⁡θsin⁡θcos⁡θ)(x1x2)\begin{pmatrix} x_1' \\ x_2' \end{pmatrix} = \begin{pmatrix} \cos\theta & -\sin\theta \\ \sin\theta & \cos\theta \end{pmatrix} \begin{pmatrix} x_1 \\ x_2 \end{pmatrix}

高维空间里的旋转没有这么简洁的公式,需要构造复杂的旋转矩阵。但拆成一对一对的二维平面,每一对独立旋转,就很好算。

以 MiniMind 为例:每个头的 96 维向量,拆成 48 对。每对分配一个不同的频率。48 个二维旋转拼起来,就是这个头的旋转;8 个头用同一张频率表各转各的,合起来就是整个 768 维向量的位置编码。

这就是“两两分组”的由来:不是设计上的玄机,是数学上的简化。

两个 token 都在转,为什么能区分位置?

关键在转的角度不一样。

  • “苹果”在位置 2,转 2 弧度。
  • “吃了”在位置 3,转 3 弧度。
  • 角度差 = 3 - 2 = 1 弧度。

如果“吃了”在位置 6,转 6 弧度,角度差 = 6 - 2 = 4 弧度。 角度差不同,点积结果就不同。 两个都在转,但转的速度一样(频率相同),谁在后面谁就多转一点。多转的这一点,就是位置差。

为什么只旋转 Q 和 K,不旋转 V

回到注意力的完整公式:

Attention(Q,K,V)=Softmax(QKTdk)V\text{Attention}(Q, K, V) = \text{Softmax}\left(\frac{QK^T}{\sqrt{d_k}}\right) V

这个公式分两步:

第一步: QKTQK^T 算出匹配分数,Softmax 变成权重。

第二步: 权重乘以 VV,加权求和,得到输出。

位置信息只需要影响第一步,不需要影响第二步。

为什么?

第一步决定“谁看谁”。 苹果应该多看吃了,还是多看手机,这是匹配问题。匹配需要位置信息——两个 token 隔得近还是远,直接影响它们该不该互相关注。所以 Q 和 K 必须带位置。

第二步决定“取回什么”。 权重算完之后,从每个 token 的 V 里取材料,加权汇总。这一步只关心“取什么内容”,不关心“这两个 token 隔多远”。位置信息在这里没有用处。

如果硬把位置信息也塞进 V,会怎样?

旋转后的输出会变成:

输出=∑jwij⋅(Rn⋅vj)\text{输出} = \sum_j w_{ij} \cdot (R_n \cdot v_j)

每个 token 的 V 被它自己的位置旋转了。这意味着,同一个内容,出现在位置 3 和出现在位置 7,取回来的材料不同。但内容本身和位置无关——“吃了”这个动作,不管它在句子的哪个位置,它提供的“动作信息”应该是一样的。

位置只影响匹配,不影响内容。 所以只旋转 Q 和 K,不旋转 V。

数学上还有一个更直接的观察。RoPE 的核心性质是:旋转后的点积只和位置差有关:

(Rmq)⋅(Rnk)=q⋅Rn−m⋅k(R_m q) \cdot (R_n k) = q \cdot R_{n-m} \cdot k

这个性质只在 Q 和 K 都被旋转、且旋转角度由各自位置决定时成立。V 不参与这个点积,所以旋不旋转都不影响这个性质。

V 的职责是提供内容,内容不随位置变。 这就是为什么 RoPE 只动 Q 和 K。

1.6 MiniMind 的 RoPE 实现

打开 MiniMind 的 model/model_minimind.py,RoPE 分两步。

第一步:预计算频率。

def precompute_freqs_cis(dim, end, rope_base, rope_scaling=None):
    freqs = 1.0 / (rope_base ** (torch.arange(0, dim, 2)[: (dim // 2)].float() / dim))
    t = torch.arange(end, device=freqs.device)
    freqs = torch.outer(t, freqs).float()
    freqs_cos = torch.cat([torch.cos(freqs), torch.cos(freqs)], dim=-1)
    freqs_sin = torch.cat([torch.sin(freqs), torch.sin(freqs)], dim=-1)
    return freqs_cos, freqs_sin

逐一拆解:

  • dim 是 head_dim,MiniMind 里是 96。
  • rope_base 是基频,MiniMind 用的是 1,000,000。LLaMA-1 用的是 10,000。基频越大,频率衰减越慢,长距离依赖效果越好。
  • torch.arange(0, dim, 2) 生成 0, 2, 4, ..., 94,共 48 个数。除以 dim 后,得到 48 个不同的频率。
  • t 是位置序列,从 0 到 end - 1。MiniMind 的 end 是 32768。
  • torch.outer(t, freqs) 得到一个 32768 × 48 的矩阵。第 mm 行第 ii 列,就是位置 mm 在第 ii 个频率上的旋转角度。
  • 最后把 cos 和 sin 各拼一份,变成 32768 × 96。前 48 列和后 48 列相同,是为了和 rotate_half 配合。

第二步:应用旋转。

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

def apply_rotary_pos_emb(q, k, cos, sin):
    q_embed = (q * cos) + (rotate_half(q) * sin)
    k_embed = (k * cos) + (rotate_half(k) * sin)
    return q_embed, k_embed

rotate_half 的作用是:把向量分成前后两半,交换位置,前半取负。这正好对应了旋转矩阵里的 −sin⁡-\sin 那一项。

q_embed = q * cos + rotate_half(q) * sin,就是旋转公式的向量化写法。对 q 和 k 都做一遍,v 不参与——因为 v 不参与匹配度计算。

1.6.1 hidden_size 是什么,为什么叫“hidden”

MiniMind 的 hidden_size 是 768。这个数,是模型内部每个 token 向量的维度。 维度越高,能容纳的信息越丰富。768 个数,就是模型对每个 token 的“内部描述”。这个描述不是人写的,是训练中学出来的。【别嫌烦】

那为什么叫 hidden?

这个名字来自早期神经网络的分层叫法。网络分三层:输入层、隐藏层、输出层。

  • 输入层:直接接收外部数据。
  • 输出层:直接给出最终结果。
  • 隐藏层:夹在中间的那些层,既不直接接触输入,也不直接给出输出。

“hidden”的意思是“不直接面对外部”,不是“藏起来看不见”。隐藏层是模型真正做计算的地方,只是它的内部状态不直接暴露给用户。

hidden_size 就是隐藏层里向量的维度。在 Transformer 里,从 embedding 之后到输出层之前,所有中间表示都是 hidden_size 维。MiniMind 选了 768,所以整条链路上流动的向量,长度都是 768。

这个维度和头数的关系是:

hidden_size=头数×head_dim\text{hidden\_size} = \text{头数} \times \text{head\_dim}

MiniMind 有 8 个 Q 头,每个头分到 768 ÷ 8 = 96 维。这个 96 就是 head_dim。

“8 层”是另一回事。层数是 Transformer 块叠了几个,MiniMind 叠了 8 个。每一层里都有自己的一套注意力头。层数和头数是两个独立的维度。

1.6.2 end = 32768 是什么

end 是预计算频率时,位置序列的最大值。MiniMind 的配置里,它等于 max_position_embeddings,设为 32768。

但这 32768 不是预训练时实际使用的序列长度。

MiniMind 预训练时,每条训练数据实际切成的序列长度是几百个 token(脚本默认 340,官方推荐 380~768 这个量级)。模型一次只看这么多,只见过位置 0 到几百。

那为什么频率表要算到 32768?

因为推理的时候,用户可能输入更长的文本。模型需要有能力处理比训练时更长的序列。max_position_embeddings = 32768 就是给推理预留的空间。

训练时只见过几百个位置,推理时却要处理 8192,这中间的差距靠 YaRN 来填。YaRN 把训练时没见过的那些长距离角度,压缩到模型熟悉的范围内。频率表预先算到 32768,就是为了让 YaRN 有足够的表可以查。

所以两个数字的分工是:

  • 几百(默认 340):训练时实际用的序列长度。
  • 32768:频率表预留的最大位置,推理时通过 YaRN 外推才能用到。

1.6.3 逐行拆代码

第一行:算频率。

freqs = 1.0 / (rope_base ** (torch.arange(0, dim, 2)[: (dim // 2)].float() / dim))

torch.arange(0, dim, 2):

dim = 96,所以生成 [0, 2, 4, 6, ..., 94],共 48 个数。

为什么隔 2 取?因为 RoPE 是两两分组。96 维分 48 对,每一对给一个频率。这一行取的是每一对的编号。

[: (dim // 2)]:

dim // 2 = 48,[:48] 就是取前 48 个。torch.arange(0, 96, 2) 本来就只有 48 个数,这里写 [:48] 是保险写法,防止越界。

.float() / dim:

把 48 个数除以 96,得到 [0, 0.0208, 0.0417, ..., 0.979]。这些是归一化的位置编号。

rope_base ** (...):

把上一步的结果作为指数,底数是 rope_base。MiniMind 的 rope_base 是 1,000,000。

当指数是 0 时,10000000=11000000^0 = 1,频率是 1/1=11/1 = 1。 当指数是 0.979 时,10000000.979≈7500001000000^{0.979} \approx 750000,频率是 1/750000≈1.3×10−61/750000 \approx 1.3 \times 10^{-6}。

1.0 / (...):

取倒数,得到频率。

最终 48 个频率,从高到低:

ω₀ ≈ 1.0
ω₁ ≈ 0.75
ω₂ ≈ 0.56
...
ω₄₇ ≈ 1.3 × 10⁻⁶

第 0 对频率最高(1.0),转得最快。第 47 对频率最低(约 1.3e-6),转得最慢。

为什么基频是 1,000,000?

LLaMA-1 用的是 10,000,MiniMind 用的是 1,000,000。基频越大,低频部分的频率越低。

低频对负责长距离。频率越低,绕圈越慢,能区分的距离越长。1,000,000 比 10,000 大 100 倍,长距离依赖的能力也强得多。这是 MiniMind 支持 32768 上下文的一个基础。

第二行:位置序列。

t = torch.arange(end, device=freqs.device)

end 是 32768,所以 t = [0, 1, 2, ..., 32767]。

这是 token 可能出现的所有位置。位置 0 到位置 32767。

第三行:算角度矩阵。

freqs = torch.outer(t, freqs).float()

outer 是外积。

t 是 32768 个位置,freqs 是 48 个频率。外积得到一个 32768 × 48 的矩阵。

第 mm 行第 ii 列的数,是:

m×ωim \times \omega_i

也就是“位置 mm 在第 ii 对上的旋转角度”。

比如:

  • 第 0 行(位置 0):所有角度都是 0。
  • 第 1 行(位置 1):角度是 ω0,ω1,...,ω47\omega_0, \omega_1, ..., \omega_{47}。
  • 第 100 行(位置 100):角度是 100ω0,100ω1,...,100ω47100\omega_0, 100\omega_1, ..., 100\omega_{47}。

位置越靠后,整行的角度越大。

第四、五行:算 cos 和 sin。

freqs_cos = torch.cat([torch.cos(freqs), torch.cos(freqs)], dim=-1)
freqs_sin = torch.cat([torch.sin(freqs), torch.sin(freqs)], dim=-1)

对上一步的 32768 × 48 矩阵,逐元素取 cos 和 sin。每个都得到一个 32768 × 48 的矩阵。

然后 torch.cat([...], dim=-1) 把两个同样的矩阵拼起来,变成 32768 × 96。

为什么要拼一份重复的?

因为 rotate_half 的写法。

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

rotate_half 把 96 维的向量分成前后两半(各 48 维),交换位置,前半取负。

RoPE 的旋转公式是:

q′=q⊙cos⁡+rotate_half(q)⊙sin⁡q' = q \odot \cos + \text{rotate\_half}(q) \odot \sin

其中 ⊙\odot 是逐元素相乘。

要让这个式子对 96 维向量成立,cos 和 sin 也必须是 96 维的。前半段对应“原向量”,后半段对应“rotate_half 之后的向量”。

如果 cos 只算 48 维,前半段和后半段需要用不同的值。但旋转公式里,每一对的两个数用的 cos 和 sin 是同一个(都是 cos⁡θ\cos\theta 和 sin⁡θ\sin\theta)。所以 cos 拼一份重复的,让前后两半都用同一组值。

这就是 torch.cat 的原因:不是有 96 个不同的角度,而是 48 个角度,每个用了两次。

1.6.4 两张表

跑完这个函数,得到两张 32768 × 96 的表:

  • freqs_cos:每个位置、每一维的 cos 值。
  • freqs_sin:每个位置、每一维的 sin 值。

用的时候,位置 mm 的 token,直接取第 mm 行,和它的 q、k 逐元素相乘,就完成了旋转。

不用每次重新算 cos 和 sin,直接查表。这是典型的“以空间换时间”——25MB 的表,换来训练时每一步的省时。

1.6.5 总结

这段代码就做了一件事:把所有可能的“位置 × 频率”组合的旋转角度,预先算好,存成 cos 和 sin 两张表。

  • dim = 96:每个头 96 维,分 48 对。
  • rope_base = 1000000:基频,决定低频对能覆盖多远的距离。
  • end = 32768:频率表预留的最大位置,不是训练时的实际序列长度。
  • 输出两张 32768 × 96 的表,用的时候按位置查。

1.7 外推:训练没见过那么长怎么办

先解释一下“序列长度”。

模型一次处理的 token 数量,是固定的。比如训练时,每条训练数据都切成几百个 token,模型一次看这么多。这个量级就是训练时的序列长度。

但推理的时候,用户可能输入一段很长的文本,比如 4096 个 token。模型要处理比训练时更长的序列。

这就出问题了。

RoPE 是靠旋转角度标记位置的。位置 0 旋转 0 度,位置 1 旋转 ω\omega 度,位置 2 旋转 2ω2\omega 度……训练时,模型只见过位置 0 到几百的旋转角度。位置 3000 该旋转多少度,公式能算出来,但模型没见过这个角度对应的模式,它不知道该怎么处理。

就像一个人只学过 1 到 100 的数,突然让他算 3000 加 5000,他知道规则,但没见过这么大的数,容易出错。

所以需要长度外推:想办法让训练时只见过几百的模型,能处理 4096 甚至更长的序列。

MiniMind 支持 YaRN(Yet another RoPE extensioN)做长度外推。

YaRN 的思路是:不改模型,改旋转频率。

具体做法:把旋转频率分档。高频部分(对应短距离位置关系)保持原样,低频部分(对应长距离位置关系)做插值——把训练时没见过的那些“远距离角度”,压缩到训练时见过的范围内。

这样,模型遇到长序列时,旋转角度不会超出它熟悉的模式,效果就保住了。

MiniMind 的配置里,max_position_embeddings=32768,通过 YaRN 可以把有效上下文从几百扩展到 32768。控制开关是 inference_rope_scaling 标志。

二、记忆:KV Cache

2.1 推理时的浪费

模型推理的时候,是一个词一个词往外吐的。

第一次:输入“今天”,算出“天气”。 第二次:输入“今天天气”,算出“真”。 第三次:输入“今天天气真”,算出“好”。

每次都要重新算一遍前面所有 token 的 k 和 v。但前面 token 的 k 和 v 在第一次就算过了,不会因为后面新增了 token 而改变。

这个浪费有多大?生成第 n 个词时,要重算前面 n-1 个 token 的 k 和 v。生成一个长度为 N 的序列,总计算量是 O(N²)。序列越长,浪费越惊人。

2.2 缓存的做法

解决办法:把算过的 k 和 v 缓存起来,下次直接用。

第一次算“今天”,得到“今天”的 k 和 v,存起来。 第二次算“天气”,只需要算“天气”的 k 和 v,然后把缓存的“今天”的 k、v 拼在前面。 第三次同理。

每次只需要算当前 token 的 k 和 v,前面所有 token 的 k、v 从缓存里取。

MiniMind 的代码:

if past_key_value is not None:
    xk = torch.cat([past_key_value[0], xk], dim=1)
    xv = torch.cat([past_key_value[1], xv], dim=1)

缓存的 k、v 和新算的拼在一起,只用算当前 token 的 q。

有了 KV Cache,每生成一个 token 的计算量从 O(n) 降到 O(1)。生成长度为 N 的序列,总计算量从 O(N²) 降到 O(N)。

2.3 缓存的代价

缓存不是免费的。

KV Cache 的大小 = 2 × 层数 × KV 头数 × head_dim × 序列长度 × 字节数。

MiniMind 的配置:8 层,4 个 KV 头,head_dim 96,fp16 存储。序列长度 2048 时:

2 × 8 × 4 × 96 × 2048 × 2 ≈ 25MB

一条序列 25MB。如果同时处理 100 条序列,就是 2.5GB。这还只是 64M 的小模型。万亿参数的大模型,KV Cache 是推理显存的主要瓶颈。

所以有了 GQA——让多个 Q 头共享一组 K 和 V,把 KV 头数降到 Q 头数的几分之一。

2.4 为什么只缓存 K 和 V,不缓存 Q

这个问题问到了 KV Cache 的关键。

先说为什么 Q 不需要缓存。

注意力分数的计算方式是:当前 token 的 q,和所有 token 的 k 做点积。

生成第 1 个 token 时,用“今天”的 q₁,和 k₁ 做点积。 生成第 2 个 token 时,用“天气”的 q₂,和 k₁、k₂ 做点积。 生成第 3 个 token 时,用“真”的 q₃,和 k₁、k₂、k₃ 做点积。

每一步,只用当前这个 token 的 q。前面 token 的 q 用完了就不再用了。所以 q 不需要缓存,算了就扔。

K 和 V 不一样。k₁ 和 v₁ 在第 1 步算出来,第 2 步还要用,第 3 步还要用,一直用到序列结束。所以必须缓存。

2.5 K 要参与位置编码,为什么前面算过的 K 不变

这是另一个关键问题。

RoPE 给 k 加位置信息,方式是旋转。旋转角度由 token 的位置决定。

“今天”在位置 0,它的 k 旋转 0 度。 “天气”在位置 1,它的 k 旋转 1 度。 “真”在位置 2,它的 k 旋转 2 度。

位置是绝对的。 “今天”永远在位置 0,不会因为后面新增了“天气”、“真”而变到位置 1 去。

所以“今天”的 k,经过 RoPE 旋转后,永远是同一个向量。第 1 步算出来是什么样,第 2 步还是什么样,第 3 步还是什么样。它不会变。

缓存的是“已经应用了 RoPE 的 k”。 存进去的时候,位置信息已经旋转进去了。取出来直接用,不需要重新加位置编码,因为位置没变。

如果换个场景,把“今天”从位置 0 挪到位置 5,那它的 k 确实会变。但推理的时候,token 的位置是固定的——先生成的在位置 0,后生成的往后排,不会往前挪。

所以:

  • K 的旋转角度由位置决定。
  • 位置是绝对的,前面 token 的位置不变。
  • 因此前面 token 的 K 不变,可以安全缓存。

这就是 KV Cache 能成立的根本原因:位置不变,K 就不变。

三、省显存:GQA

3.1 三种注意力:MHA、MQA、GQA

MHA、MQA、GQA,说的是 Q 头和 KV 头之间的数量关系。

MiniMind 有 8 个 Q 头,每个头 96 维。下面用这组配置,把三种方案各走一遍。

MHA:每个 Q 头都有自己的一套 K 和 V

MHA 是原版 Transformer 的做法。Q 头有几个,KV 头就有几个。

8 个 Q 头,8 个 KV 头。每个 Q 头配一组自己的 K 和 V:

Q 头 1 → K 头 1、V 头 1
Q 头 2 → K 头 2、V 头 2
Q 头 3 → K 头 3、V 头 3
...
Q 头 8 → K 头 8、V 头 8

每个头各看各的,互不干扰。质量最好,但 KV Cache 最大。

KV Cache 的大小是:

2×8×8×96×序列长度×2字节2 \times 8 \times 8 \times 96 \times \text{序列长度} \times 2 \text{字节}

8 个 KV 头,每个 96 维。

MQA:所有 Q 头共用一套 K 和 V

MQA 走另一个极端。8 个 Q 头,只有 1 个 KV 头。

Q 头 1 ┐
Q 头 2 ├→ K 头 1、V 头 1
Q 头 3 │
...   │
Q 头 8 ┘

所有 Q 头都去匹配同一组 K,都从同一组 V 里取材料。

KV Cache 直接降到原来的 1/8:

2×8×1×96×序列长度×2字节2 \times 8 \times 1 \times 96 \times \text{序列长度} \times 2 \text{字节}

省得最多,但质量有损失。因为 8 个 Q 头本来想从不同角度提问,现在只能共用一套“名片”,各自失去了匹配的自由度。

GQA:折中方案

GQA 把 Q 头分组,每组共享一套 K 和 V。

MiniMind 是 8 个 Q 头,4 个 KV 头。两个 Q 头一组:

Q 头 1、Q 头 2 → K 头 1、V 头 1
Q 头 3、Q 头 4 → K 头 2、V 头 2
Q 头 5、Q 头 6 → K 头 3、V 头 3
Q 头 7、Q 头 8 → K 头 4、V 头 4

第 1 组两个 Q 头,共享第 1 组 K 和 V。第 2 组两个 Q 头,共享第 2 组 K 和 V。以此类推。

KV Cache 是 MHA 的一半:

2×8×4×96×序列长度×2字节2 \times 8 \times 4 \times 96 \times \text{序列长度} \times 2 \text{字节}

3.2 K 和 V 被投影降维了

前面讲 GQA 时说 K 和 V 是 384 维,Q 是 768 维。这个“降维”具体发生在哪一步,单独说一下。

一个 token 的向量 x 是 768 维。它进来之后,分三条路走:

x (768 维)
  │
  ├─ q_proj → q (768 维)
  ├─ k_proj → k (384 维)
  └─ v_proj → v (384 维)

Q 没有降维,K 和 V 降了。

q_proj 的形状是 768 × 768,输入 768 维,输出还是 768 维。8 个 Q 头,每头 96 维。

k_proj 和 v_proj 的形状是 768 × 384,输入 768 维,输出只有 384 维。4 个 KV 头,每头 96 维。

降的是总维度,不是每个头的维度。 每个 KV 头还是 96 维,和 Q 头一样。变的是头的数量:Q 有 8 个头,K 和 V 只有 4 个。

这是 GQA 的设计。如果换成 MHA,k_proj 和 v_proj 的输出也是 768 维,8 个 KV 头,不降。如果换成 MQA,k_proj 和 v_proj 的输出只有 96 维,1 个 KV 头,降得更狠。

所以“降维”这件事,不是模型把算出来的 768 维 K、V 压缩了,而是 k_proj 和 v_proj 这两个矩阵从一开始就只输出 384 维。投影矩阵本身就是 768 × 384,不是 768 × 768。

参数上也能对上。第一篇算过:

  • q_proj:768 × 768 = 589,824 个数。
  • k_proj:768 × 384 = 294,912 个数。
  • v_proj:768 × 384 = 294,912 个数。

k_proj 和 v_proj 的参数,正好是 q_proj 的一半。少的这一半,就是 KV 头从 8 个减到 4 个省下来的。

省的不只是参数,还有推理时的 KV Cache。KV Cache 存的是每个 token 的 K 和 V。KV 头数减半,缓存也减半。

一句话:Q 保持 768 维不变,K 和 V 被投影降到 384 维。降的是头的数量,不是每个头的维度。

3.3 具体走一遍:“苹果”的 Q 和“吃了”的 K 怎么匹配

拿“他把苹果吃了”举例。假设“苹果”在位置 2,“吃了”在位置 3。

“苹果”的向量 x 是 768 维。经过 q_proj,输出 768 维,拆成 8 个 Q 头,每头 96 维:

苹果的 Q 头 1: [0.12, -0.45, ..., 0.87]  (96 个数)
苹果的 Q 头 2: [-0.33, 0.61, ..., -0.19] (96 个数)
...
苹果的 Q 头 8: [0.48, 0.22, ..., 0.95]   (96 个数)

“吃了”的向量经过 k_proj,输出 384 维,拆成 4 个 KV 头,每头 96 维:

吃了的 K 头 1: [0.91, -0.12, ..., 0.44]  (96 个数)
吃了的 K 头 2: [0.25, 0.78, ..., -0.33]  (96 个数)
吃了的 K 头 3: [...]
吃了的 K 头 4: [...]

GQA 的匹配方式:

  • 苹果的 Q 头 1 和 Q 头 2,都去和吃了的 K 头 1 匹配。
  • 苹果的 Q 头 3 和 Q 头 4,都去和吃了的 K 头 2 匹配。
  • 苹果的 Q 头 5 和 Q 头 6,都去和吃了的 K 头 3 匹配。
  • 苹果的 Q 头 7 和 Q 头 8,都去和吃了的 K 头 4 匹配。

同一组内的两个 Q 头,看到的是同一套 K 和 V。 它们提问的角度不同(各自的 Q 不同),但被匹配的“名片”是一样的。

为什么这样能省显存

推理时,KV Cache 存的是所有 token 的 K 和 V。

  • MHA 要存 8 组 K 和 V。
  • GQA 只存 4 组。
  • MQA 只存 1 组。

GQA 用 4 组 K 和 V,服务 8 个 Q 头。每两个 Q 头共享一组。这就像 8 个人开会,本来每人配一个秘书(MHA),现在两个共用一个秘书(GQA),8 个人共用一个秘书(MQA)。秘书少了,记录的东西就少了,省地方。但共享的人越多,每个人的个性化需求就越难满足,质量就越容易掉。

一张表对比

方案Q 头数KV 头数KV Cache质量
MHA88最大最好
GQA84MHA 的一半接近 MHA
MQA81MHA 的 1/8明显下降

MiniMind 选 GQA,就是在质量和不显存之间取了个平衡点。Llama 2、Llama 3、Qwen2、Qwen3,也都用 GQA。

3.4 GQA 省了多少

回到 KV Cache 的公式:

KV Cache 大小=2×层数×KV 头数×head_dim×序列长度×字节数\text{KV Cache 大小} = 2 \times \text{层数} \times \text{KV 头数} \times \text{head\_dim} \times \text{序列长度} \times \text{字节数}

如果 KV 头数等于 Q 头数(MHA),MiniMind 的 KV Cache 就是:

2 × 8 × 8 × 96 × 2048 × 2 ≈ 50MB

用 GQA,KV 头数从 8 降到 4,KV Cache 直接减半,变成 25MB。

省下的显存,可以用来放更长的序列,或者并发处理更多的请求。

3.5 为什么 GQA 不怎么掉质量

MQA 把 KV 头数降到 1,省得最多,但质量掉得明显。因为所有 Q 头被迫看同一组 K 和 V,各自失去了一部分“被匹配”的自由度。

GQA 保留了分组,每组内部共享,组间独立。8 个 Q 头分成 4 组,每组 2 个 Q 头共享一组 KV。这样既省了缓存,又保留了一定程度的多样性。

Llama 2 的 70B 版本最早大规模用了 GQA,之后 Llama 3、Qwen2、Qwen3 都跟进。现在 GQA 基本是新模型标配。

四、省时间:Flash Attention

4.1 GPU 的内存墙

Flash Attention 是一种针对注意力机制的计算优化技术,由斯坦福大学 Tri Dao 等人在 2022 年提出。它的目标是在不损失精度的前提下,提升 Transformer 训练和推理的速度,并降低显存占用。

要理解它,得先知道 GPU 的内存是分层的,主要分两级:

  • SRAM(高速缓存):容量很小(约 20MB),但读写速度极快(约 19TB/s)。
  • HBM(高带宽显存):容量很大(数十 GB),但读写速度相对慢得多(约 1.5TB/s)。

标准的注意力计算会生成一个 N×N 的注意力矩阵(N 是序列长度)。当序列很长时,这个矩阵非常庞大。传统实现会把这个大矩阵在慢速的 HBM 和快速的 SRAM 之间来回搬运,绝大部分时间浪费在数据搬运上,而不是实际计算上。

4.2 分块计算:不让大矩阵完整出现

Flash Attention 的核心思路:不要让那个 N×N 的大矩阵完整地出现在 HBM 里。

先看这个矩阵有多大。序列长度 2048 时,N×N 就是 2048 × 2048 = 419 万个数字。fp16 存储,一个数字 2 字节,这个矩阵就是 8MB。

SRAM 只有 20MB,理论上放得下。但 SRAM 还要放 Q、K、V 的中间结果,还要放其他计算数据,留给这个矩阵的空间并不多。而且序列再长一点,比如 4096,矩阵就变成 32MB,SRAM 彻底放不下了。

所以传统实现只能把这个矩阵放到 HBM 里。每次算一部分,就从 HBM 读一部分,算完再写回去。HBM 慢,读写次数一多,时间就耗在搬运上。

Flash Attention 的做法:把这个大矩阵切成小块,一块一块地算。

还是 2048 × 2048 的矩阵,切成 64 × 64 的小块。一块只有 4096 个数字,8KB。SRAM 放这一块绰绰有余。每次只把一块加载到 SRAM 里,算完这块的输出,扔掉,再加载下一块。

这样,N×N 的大矩阵从头到尾没有在 HBM 里完整出现过。HBM 的读写量大幅降低。

但这里有一个问题:Softmax 需要看到一整行的所有分数,才能算归一化。切成小块之后,每一块只看到了一部分分数,怎么算 Softmax?

Flash Attention 用了一个数学技巧:在线 Softmax。它一边遍历小块,一边维护当前的最大值和累加和,最后再统一归一化。这样就不需要一次性看到整行。

这个技巧保证了 Flash Attention 的结果和标准注意力完全一致,不是近似。

4.3 重计算:以计算换内存

反向传播时,需要用到前向算出的注意力矩阵来算梯度。

标准实现会把前向的注意力矩阵存下来,反向时直接读。但 Flash Attention 没有把这个矩阵完整地算出来,也没存,怎么办?

答案是:重新算一遍。

反向传播走到注意力这一步时,Flash Attention 拿着已经存下来的 Q、K、V,重新执行一遍分块计算,把需要的中间结果算出来。

这是典型的“以计算换内存”。代价是训练时多算一遍,总计算量增加;收益是显存占用大幅降低。

4.4 效果

通过分块计算和重计算,Flash Attention 把注意力对 HBM 的访问次数大幅降低,显存占用从随序列长度平方增长(O(N²))降为线性增长(O(N))。

实践中,它带来的速度提升显著。在 GPT-2 上训练速度可提升 3 倍,并能支持长达 64K 的序列长度。目前 Flash Attention 已经成为现代大模型训练的事实标准,主流框架(PyTorch、Hugging Face)都集成了它。

4.5 MiniMind 怎么用 Flash Attention

前面讲了 Flash Attention 的原理:分块计算、在线 Softmax、重计算。这些是它自己就做的事,不需要使用者操心。

但语言模型有一个额外的需求:因果掩码。

第三篇讲过,语言模型是自回归的——每个位置只能看自己和前面的位置,不能看后面的。所以在算注意力分数时,要把每个位置对应未来位置的分数设成负无穷,Softmax 之后这些位置的权重就变成 0。

标准实现的做法是:手动构造一个下三角的掩码矩阵,把上三角填成负无穷。

Flash Attention 把这个需求也接管了。MiniMind 的注意力实现里,核心的注意力计算调用是这样的:

# 手工路径(SDPA 不可用时):上三角加 -inf
scores[:, :, :, -seq_len:] += torch.full((seq_len, seq_len), float("-inf"), device=scores.device).triu(1)
# Flash 路径(默认):因果掩码就是 kernel 的一个开关
attn_output = F.scaled_dot_product_attention(xq, xk, xv, dropout_p=self.dropout if self.training else 0.0, is_causal=self.is_causal)

is_causal 参数告诉 kernel:

  • is_causal=True:kernel 内部自动遮住未来位置,不需要外部的掩码矩阵。
  • is_causal=False:不遮,所有位置互相可见。

所以 MiniMind 不需要自己构造掩码矩阵。因果掩码这件事,从外部的手工操作,变成了 kernel 内部的一个开关。

分工是这样的:

  • 分块计算、在线 Softmax、重计算:Flash Attention 默认就做,使用者不用管。
  • 因果掩码:通过 is_causal=True 告诉它要不要做。

MiniMind 两样都用了:分块计算是自动的,因果掩码是靠这个参数打开的。

五、深加工:FFN

5.1 FFN 是什么,为什么叫“前馈”

FFN 的全称是 Feed-Forward Network,中文叫前馈网络。

“前馈”这个词,意思是信息只向前流动,不回头。数据从输入层进来,穿过一层或多层,直接到达输出层,中间没有循环、没有反馈。

这个名字来自早期神经网络的分类。和它相对的是“循环神经网络”(RNN),RNN 的信息会从后面的步骤回传到前面的步骤。FFN 不循环,数据进去、出来,就完事了。

但在 Transformer 的语境里,FFN 特指 Transformer 块里注意力后面的那个子层。它的工作对象是每个 token 的向量,一个一个独立处理——token 之间不发生任何信息交互。交互的事,已经由注意力在前面做完了。

所以 FFN 和注意力,是一个块里分工明确的两半:

  • 注意力:token 和 token 之间交换信息。横向的。
  • FFN:每个 token 自己消化信息。纵向的。

一句话:注意力负责开会讨论,FFN 负责会后自己消化。

5.2 为什么需要一个“非线性”的加工

注意力算完之后,每个 token 拿到了一个包含上下文信息的新向量。但这个向量还只是“汇总”,还没有被深度加工。

而且注意力有一个根本性的限制:它的输出对 V 是线性加权和,但权重本身的计算(QK^T + Softmax)是非线性的。

但即使权重是非线性的,整个注意力模块——从输出角度看——仍然缺少一个关键的东西:逐元素的非线性变换。

线性变换有一个致命弱点:叠多少层,整体还是线性的。

假设两层线性变换,y=W2(W1x)y = W_2(W_1 x),展开就是 y=(W2W1)xy = (W_2 W_1)x,等价于一个矩阵。叠 100 层也一样,最终还是等价于一个矩阵。

纯线性模型,不管多深,表达能力都只相当于一层。它没法处理“如果这个数大于 0,就往一个方向走;否则往另一个方向走”这种带条件的判断。

而语言里到处是这种判断。“这个词如果是名词,就按名词处理;如果是动词,就按动词处理”——这种逻辑,线性变换做不到。

所以注意力之后,必须有一个地方引入非线性。这就是 FFN 的任务。

5.3 激活函数:非线性的来源

引入非线性的关键,是激活函数。

激活函数是一个作用在单个数字上的函数。它的特点是不是直线——输入和输出之间不是简单的比例关系。

最早的激活函数是 Sigmoid,后来是 ReLU,再后来是 SiLU、GELU 等。不同的激活函数,图像不同,脾气不同。

ReLU:一个折线

ReLU 的公式:

ReLU(x)=max⁡(0,x)\text{ReLU}(x) = \max(0, x)

图像是一条折线:

        |
        |       /
        |      /
        |     /
        |    /
        |   /
        |  /
        | /
        |/
--------+--------
        |

左半边(x < 0)全是 0,一条水平线。右半边(x > 0)是一条 45 度的斜线,y=xy = x。

ReLU 做的事:负数砍成 0,正数原样通过。

它简单、快、有效。整个深度学习的复兴,ReLU 功不可没。但它有一个问题:0 点不可导——左边斜率是 0,右边斜率是 1,中间断开了。训练时遇到这个问题,梯度会不稳定。

而且 ReLU 的“硬切”太粗暴:一个数只要小于 0,直接归零,信息全丢。这在大模型里不是最优的。

SiLU:一条光滑的曲线

SiLU 的公式:

SiLU(x)=x⋅σ(x)\text{SiLU}(x) = x \cdot \sigma(x)

其中 σ(x)\sigma(x) 是 Sigmoid 函数:

σ(x)=11+e−x\sigma(x) = \frac{1}{1 + e^{-x}}

SiLU 的图像是一条光滑的曲线:

        |
        |         /
        |        /
        |       /
        |      /
        |     /
        |    /
        |   /
--------+--/--------
       /|
      / |
     /  |
    /   |

负数区域,它不直接归零,而是缓慢趋近 0(负得越多,越接近 0)。正数区域,它近似线性增长,但有一个平滑的过渡。0 附近,它有一个轻微的下凹,形状像一个小山谷。

SiLU 做的事:负数区域几乎归零,但保留了微弱的信号;正数区域通过,但过渡是光滑的。

比 ReLU 好在哪?

  • 光滑可导:没有断点,梯度处处存在,训练更稳定。
  • 保留微弱信号:负数不直接砍成 0,而是保留一小部分。这在某些情况下有助于模型学习更细腻的模式。
  • 非单调:SiLU 在 0 附近有一个轻微的下凹,所以不是单调递增的。这一点看起来很反直觉,但实验证明它比单调函数效果更好。

激活函数越光滑,训练越稳定。SiLU 是“光滑版”的 ReLU,这也是为什么现代 LLM 更偏爱它。

5.4 FFN 在干什么:升维、非线性、降维

有了激活函数,FFN 的完整流程就能讲了。

最简单的一层 FFN,是两层线性加一个激活:

FFN(x)=ReLU(xW1)W2\text{FFN}(x) = \text{ReLU}(x W_1) W_2

三步:

第一步:升维。 W1W_1 把 768 维的 x 升到更高维,比如 3072 维。

第二步:非线性。 ReLU(或 SiLU)作用在每一个维度上。

第三步:降维。 W2W_2 把 3072 维压回 768 维。

为什么要升维?

打个比方。你在二维平面上画一条线,最多能把平面分成两块。在三维空间里画一个平面,也能分两块。但如果要分开一个不规则的区域,二维平面上的直线就不够了——你需要曲线,需要复杂的形状。

维度越高,能表达的模式越复杂。 768 维升到 3072 维,模型有了一块更大的“画布”,在这个高维空间里,它能刻出更精细的特征。

然后激活函数在这里切割空间:哪些区域激活,哪些区域抑制。这一步之后,再降维压回原空间。降维不是丢信息,而是把高维空间里学到的“模式”压缩成一个新的 768 维表示。

升维 → 非线性切割 → 降维,这就是 FFN 做的全部事情。每个 token 的向量,经过这一圈,被重新“雕刻”了一遍。

5.5 SwiGLU:给 FFN 加一道门

ReLU FFN 有两层。现代 LLM 换成了 SwiGLU,有三层。

SwiGLU 的公式:

FFN(x)=(SiLU(xWgate)⊙xWup)Wdown\text{FFN}(x) = \big(\text{SiLU}(x W_{\text{gate}}) \odot x W_{\text{up}}\big) W_{\text{down}}

比 ReLU FFN 多了一个矩阵,多了一个逐元素相乘。这个多出来的部分,就是一个“门”。

三个矩阵各司其职:

gate_proj:升维,过 SiLU。输出是门控值——一串介于负数和正数之间的数,决定每个维度“放多少信息过去”。

up_proj:升维,不过激活。输出是原始信息——真正要搬运的内容。

逐元素相乘:门控值 × 原始信息。这就是“门”的动作。

举个具体例子。假设某个维度上:

  • gate 输出 2.0,SiLU(2.0) ≈ 1.76
  • up 输出 3.0
  • 相乘:1.76 × 3.0 = 5.28

换一个维度:

  • gate 输出 -2.0,SiLU(-2.0) ≈ -0.24
  • up 输出 3.0
  • 相乘:-0.24 × 3.0 = -0.72

同一个 3.0 的信息,因为 gate 不同,通过的量和符号都变了。 这就是门控的意义:不是“开或关”的二选一,而是按比例缩放,甚至可以翻转符号。

down_proj:把高维压回 768 维。

5.6 中间维度为什么是 π 倍

原版 FFN 的中间维度是 hidden_size 的 4 倍。768 升到 3072。

SwiGLU 多了一个矩阵,参数自然变多。如果还保持 4 倍,参数量就比原版多 50%。为了参数持平,SwiGLU 的中间维度通常取 8/3 倍。

但 MiniMind 用的不是 8/3。它用的是 π。

self.intermediate_size = kwargs.get("intermediate_size", math.ceil(hidden_size * math.pi / 64) * 64)

算一下:

  • 768 × π / 64 = 37.699
  • 向上取整到 38
  • 38 × 64 = 2432

所以 MiniMind 实际用的是 2432。而且 2432 就是默认公式的结果,不需要任何显式配置——MiniMind 的倍率不是 8/3≈2.67,而是 π≈3.14。

为什么取整到 64 的倍数? 为了让矩阵维度对齐硬件的计算友好度。64 是 GPU 上矩阵乘法的常见对齐粒度,2432 是 64 的倍数,计算效率更高。

5.7 MiniMind 的 FFN

class FeedForward(nn.Module):
    def __init__(self, config):
        super().__init__()
        self.intermediate_size = kwargs.get("intermediate_size", math.ceil(hidden_size * math.pi / 64) * 64)
        self.gate_proj = nn.Linear(config.hidden_size, config.intermediate_size, bias=False)
        self.down_proj = nn.Linear(config.intermediate_size, config.hidden_size, bias=False)
        self.up_proj = nn.Linear(config.hidden_size, config.intermediate_size, bias=False)
        self.act_fn = ACT2FN[config.hidden_act]

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

逐段拆。

第一段:算中间维度。

self.intermediate_size = kwargs.get("intermediate_size", math.ceil(hidden_size * math.pi / 64) * 64)

hidden_size = 768,768 × π / 64 = 37.699,向上取整到 38,再乘回 64,得到 2432。

所以 MiniMind 实际用的是 2432。而且 2432 就是默认公式的结果,不需要任何显式配置。

第二段:三个矩阵。

self.gate_proj = nn.Linear(768, 2432, bias=False)
self.up_proj   = nn.Linear(768, 2432, bias=False)
self.down_proj = nn.Linear(2432, 768, bias=False)
  • gate_proj:768 → 2432,负责门控。
  • up_proj:768 → 2432,负责提供信息。
  • down_proj:2432 → 768,负责压回原维度。

bias=False:不加偏置项。现代 LLM 的线性层通常都不加,省参数。

第三段:forward。

return self.down_proj(self.act_fn(self.gate_proj(x)) * self.up_proj(x))

从里往外读:

  1. self.gate_proj(x):把 x 升到 2432 维。
  2. self.act_fn(...):SiLU 激活,得到门控值。
  3. self.up_proj(x):把 x 升到 2432 维,得到原始信息。
  4. 两者逐元素相乘:按门控值筛信息。
  5. self.down_proj(...):压回 768 维。

走完这一圈,每个 token 的向量从 768 维进去,768 维出来,但中间在高维空间里被非线性地“雕刻”了一遍。

5.8 FFN 占了多少参数

把 FFN 和注意力对比一下。

每层注意力:

矩阵形状参数量
q_proj768 × 768589,824
k_proj768 × 384294,912
v_proj768 × 384294,912
o_proj768 × 768589,824
合计1,769,472

每层 FFN:

矩阵形状参数量
gate_proj768 × 24321,867,776
up_proj768 × 24321,867,776
down_proj2432 × 7681,867,776
合计5,603,328

FFN 的参数是注意力的 3 倍多。

这不是 MiniMind 的特例。所有主流 Transformer 都这样。FFN 承载了模型大部分参数,也承载了大部分“知识存储”。

有研究认为,FFN 的高维中间层像一个 key-value 记忆库:每个中间维度对应一种模式,up_proj 把这些模式提出来,gate_proj 决定哪些模式激活,down_proj 把激活的模式重新组合成输出。

所以,注意力和 FFN 的分工是:

  • 注意力:找关系,决定 token 之间怎么互动。
  • FFN:记知识,决定每个 token 自己该被加工成什么样。

两者缺一不可。

六、训得动:残差 + RMSNorm

6.1 深层网络的困境

8 层积木,每层都做一次注意力、一次 FFN。如果每层都把向量彻底改写,会发生什么?

第一层的输入是 embedding 向量,经过注意力、FFN,变成一个新向量。这个新向量进入第二层,又被彻底改写。到第八层出来,原始的 embedding 信息可能已经被磨得差不多了。

更严重的是,训练时梯度要从最后一层反向传到第一层。每经过一层,梯度都可能被放大或缩小。8 层下来,要么梯度爆炸,要么梯度消失。层数越多,越难训。

6.2 残差连接

残差连接的做法极其简单:把输入直接加到输出上。

输出=层(x)+x\text{输出} = \text{层}(x) + x

不是“彻底改写”,而是“在原向量上加一个修正量”。

这样做的效果是:原始信息永远有一条直通路。 不管中间的层把 x 变换成什么样,x 本身始终保留在输出里。每一层只需要学“在这个基础上,还应该修正什么”,而不是“从头重建整个表示”。

梯度也受益。反向传播时,残差连接提供了一个恒等路径,梯度可以沿着这条路直接传回去,不受中间层的影响。这就是为什么深层网络能训得动。

6.3 RMSNorm:把数值拉回可控范围

神经网络逐层计算时,矩阵乘法会不断放大数值尺度。一层一层叠加,数值可能变得极大或极小。太大就溢出,太小就变成零,梯度也跟着消失。

归一化的作用,是在每个计算模块的入口,把数值拉回一个标准范围。

传统的 LayerNorm 做两件事:减去均值,除以标准差。

RMSNorm 只做一件事:除以均方根。

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

不减去均值。就这一个区别。

为什么可以省掉?因为实践发现,减去均值对最终效果影响很小,但计算量少了。大模型里,少一步计算,乘以几百层,就是可观的节省。

6.4 MiniMind 的实现

class RMSNorm(nn.Module):
    def __init__(self, dim, eps=1e-5):
        super().__init__()
        self.eps = eps
        self.weight = nn.Parameter(torch.ones(dim))

    def norm(self, x):
        return x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + self.eps)

    def forward(self, x):
        return self.weight * self.norm(x.float()).type_as(x)
  • x.pow(2).mean(-1) 算的是每个 token 向量的均方值。
  • torch.rsqrt 是平方根的倒数。
  • self.weight 是一个可学习的缩放因子,初始为 1。归一化改变了向量的“长度”,这个权重让模型有机会把长度调回来。

归一化只改变向量的尺度,不改变方向。语义信息主要编码在方向里,所以不会因为归一化而丢失。

6.5 残差和归一化的组合

每个 Transformer 块里,注意力之前和 FFN 之前,各有一处 RMSNorm。每次计算完,都要加回原始的输入。

residual = hidden_states
hidden_states = self.self_attn(self.input_layernorm(hidden_states))
hidden_states += residual   # 第一次残差:注意力之后

residual = hidden_states
hidden_states = self.feed_forward(self.post_attention_layernorm(hidden_states))
hidden_states += residual   # 第二次残差:FFN 之后

两次残差,保证了信息在注意力模块和 FFN 模块之间流动时,原始内容始终有一份保留。

整个模型最后还有一处 final norm,在送入输出层之前做最终校准。

七、组装:一个完整的 Transformer 块

现在把六块拼起来。

一个 MiniMind 的 Transformer 块,按顺序做这些事:

输入 hidden_states
  │
  ├─ 保存 residual = hidden_states
  │
  ├─ hidden_states = RMSNorm(hidden_states)          ← 归一化
  │
  ├─ hidden_states = Attention(hidden_states)         ← 注意力
  │     ├─ Q、K、V 投影
  │     ├─ RoPE 旋转 Q 和 K
  │     ├─ KV Cache 拼接
  │     ├─ 注意力计算
  │     └─ o_proj 输出
  │
  ├─ hidden_states = hidden_states + residual         ← 残差连接
  │
  ├─ 保存 residual = hidden_states
  │
  ├─ hidden_states = RMSNorm(hidden_states)          ← 归一化
  │
  ├─ hidden_states = FeedForward(hidden_states)       ← FFN
  │
  └─ hidden_states = hidden_states + residual         ← 残差连接
  │
输出 hidden_states

这个块重复 8 次,就是 MiniMind 的全部积木。

输入是 (batch, seq_len, 768),输出也是 (batch, seq_len, 768)。维度不变,但每个位置的向量,已经经过了 8 轮“注意力 + FFN”的加工。

具体配置:

组件配置
hidden_size768
层数8
Q 头数8
KV 头数4
head_dim96
FFN 中间维度2432
归一化RMSNorm
激活函数SiLU
位置编码RoPE(支持 YaRN)

对比一下 03 篇里讲过的注意力,这一篇补上了剩下的半边。注意力让 token 互相看,位置告诉模型谁在前谁在后,KV Cache 让推理不用重算,GQA 省显存,Flash Attention 省时间,FFN 让每个 token 自己消化,残差和 RMSNorm 保证深层能训得动。

这六块合起来,才是一个完整的 Transformer 块。

八、动手实验

实验一:看 RoPE 的旋转效果

import torch

def precompute_freqs(dim, end, rope_base=1000000):
    freqs = 1.0 / (rope_base ** (torch.arange(0, dim, 2)[: (dim // 2)].float() / dim))
    t = torch.arange(end)
    freqs = torch.outer(t, freqs).float()
    return freqs

freqs = precompute_freqs(96, 100)
print("位置 0 的前 5 个频率:", freqs[0, :5])
print("位置 1 的前 5 个频率:", freqs[1, :5])
print("位置 50 的前 5 个频率:", freqs[50, :5])

观察不同位置的旋转角度差异。位置越靠后,角度越大。前几个频率(高频)变化快,后面的频率(低频)变化慢。

实验二:对比有无残差连接

import torch

x = torch.randn(1, 10, 768)

# 无残差:每层彻底改写
h = x
for _ in range(8):
    h = torch.randn(1, 10, 768) * 0.1
print("无残差,最终与原始输入的相似度:",
      torch.cosine_similarity(x.flatten(), h.flatten(), dim=0).item())

# 有残差:每层加修正量
h = x
for _ in range(8):
    h = h + torch.randn(1, 10, 768) * 0.1
print("有残差,最终与原始输入的相似度:",
      torch.cosine_similarity(x.flatten(), h.flatten(), dim=0).item())

残差连接让原始信息在 8 层之后仍然高度保留。

实验三:KV Cache 的效果

import time
import torch

# 模拟 KV Cache 的节省
seq_len = 1000
cache_size = 0
no_cache_size = 0

for i in range(seq_len):
    no_cache_size += i   # 无缓存,每次重算前面所有 token

print(f"无缓存的总计算量: {no_cache_size}")
print(f"有缓存的总计算量: {seq_len}")
print(f"节省比例: {1 - seq_len / no_cache_size:.1%}")

序列越长,KV Cache 节省越多。这是 O(N²) 到 O(N) 的差距。

九、本篇概念清单

概念本篇交代到什么程度
RoPE(旋转位置编码)讲透:原理、公式、MiniMind 实现、多频率设计
相对位置编码讲清:为什么旋转后点积只和距离有关
YaRN讲清:外推的思路,MiniMind 的支持
KV Cache讲透:为什么需要、怎么做、代价是什么
MHA / MQA / GQA讲透:三种注意力的关系和取舍
Flash Attention讲清:内存墙、分块计算、重计算
FFN / SwiGLU讲透:注意力负责什么、FFN 负责什么、SwiGLU 的三层结构
残差连接讲透:为什么需要、怎么实现
RMSNorm讲透:和 LayerNorm 的区别、公式、实现
Transformer 块讲透:完整组装顺序

本篇要牢记的只有三个词:RoPE、KV Cache、SwiGLU。

十、回到开头的问题

第三篇结尾说,注意力让 token 互相看。这一篇补上了剩下的部分。

模型怎么知道词序?RoPE 把位置信息变成旋转角度,融进 Q 和 K 的点积里。

怎么记住前文?KV Cache 把算过的 K、V 存起来,推理时不用重算。

显存不够怎么办?GQA 让多个 Q 头共享一组 K 和 V,KV Cache 减半。

算得太慢怎么办?Flash Attention 分块计算,不把 N×N 矩阵物化到 HBM 里。

每个 token 自己怎么消化?FFN 用 SwiGLU 做通道维度的非线性变换。

深层怎么训得动?残差连接保证信息直通,RMSNorm 把数值拉回可控范围。

六块拼起来,就是一个完整的 Transformer 块。8 个块叠起来,就是 MiniMind 的全部积木。

下一篇,我们看最后一个环节:积木的输出怎么变成 6400 个分数,又怎么从分数变成下一个 token。

十一、思考题

  1. 如果把 RoPE 的 rope_base 从 1,000,000 改回 LLaMA-1 的 10,000,会发生什么?提示:想想频率衰减速度。
  2. KV Cache 在训练时用不用?为什么?
  3. 残差连接是 输出 = 层(x) + x。如果改成 输出 = 层(x) + 0.1 * x,会有什么问题?
  4. SwiGLU 的中间维度在 MiniMind 里用的是 π 倍率,不是 8/3 倍。为什么要取整到 64 的倍数?

备注:本篇的源码分析基于 MiniMind 仓库的 model/model_minimind.py。RoPE 的公式推导参考了苏剑林等人的原始论文。YaRN 的细节可以参考其论文《YaRN: Efficient Context Window Extension of Large Language Models》。Flash Attention 的原理参考 Tri Dao 等人的论文《FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness》。