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

上承〔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),看作一个二维平面上的点。然后把这个点按一个角度旋转:
( x1′ x2′ ) = ( cos⁡θ −sin⁡θ sin⁡θ cos⁡θ ) ( x1 x2 ) \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} (x1′x2′)=(cosθsinθ−sinθcosθ)(x1x2)

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

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

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

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

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

关键在点积。

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

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

假设"猫"在位置 mm m,"坐"在位置 nn n。经过 RoPE 后,"猫"的 q 旋转了 mωm\omega mω,"坐"的 k 旋转了 nωn\omega nω。点积的结果,只取决于 (m−n)ω(m - n)\omega (m−n)ω,也就是两个 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 ω1=1.0(高频,转得快)
  • 第 2 对: ω2=0.1 \omega_2 = 0.1 ω2=0.1(低频,转得慢)

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

ini 复制代码
苹果的 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×1.0=2 弧度
  • 第 2 对旋转角度: 2×0.1=0.22 \times 0.1 = 0.2 2×0.1=0.2 弧度

"吃了"在位置 3:

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

每一对分别算

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

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

css 复制代码
q₁' = [cos(2), sin(2)] ≈ [-0.416, 0.909]

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

css 复制代码
k₁' = [cos(3), sin(3)] ≈ [-0.990, 0.141]

点积:

scss 复制代码
q₁' · k₁' = (-0.416)(-0.990) + (0.909)(0.141) ≈ 0.540

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

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

css 复制代码
q₂' = [cos(0.2), sin(0.2)] ≈ [0.980, 0.199]

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

css 复制代码
k₂' = [cos(0.3), sin(0.3)] ≈ [0.955, 0.296]

点积:

scss 复制代码
q₂' · k₂' = (0.980)(0.955) + (0.199)(0.296) ≈ 0.995

总点积:

ini 复制代码
q' · k' = 0.540 + 0.995 = 1.535

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

换个距离,看分数怎么变

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

"吃了"在位置 9:

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

第 1 对(高频):

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

css 复制代码
k₁' = [cos(9), sin(9)] ≈ [-0.911, 0.412]

点积:

scss 复制代码
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 弧度:

css 复制代码
k₂' = [cos(0.9), sin(0.9)] ≈ [0.622, 0.783]

点积:

scss 复制代码
q₂' · k₂' = (0.980)(0.622) + (0.199)(0.783) ≈ 0.610 + 0.156 ≈ 0.766

总点积:

ini 复制代码
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⁡θ ) ( x1 x2 ) \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} (x1′x2′)=(cosθsinθ−sinθcosθ)(x1x2)

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

以 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 ( QKT dk ) V\text{Attention}(Q, K, V) = \text{Softmax}\left(\frac{QK^T}{\sqrt{d_k}}\right) V Attention(Q,K,V)=Softmax(dk QKT)V

这个公式分两步:

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

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

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

为什么?

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

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

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

旋转后的输出会变成:
输出=∑j wij ⋅(Rn⋅vj) \text{输出} = \sum_j w_{ij} \cdot (R_n \cdot v_j) 输出=j∑wij⋅(Rn⋅vj)

每个 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 (Rmq)⋅(Rnk)=q⋅Rn−m⋅k

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

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

1.6 MiniMind 的 RoPE 实现

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

第一步:预计算频率。

python 复制代码
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 m 行第 ii i 列,就是位置 mm m 在第 ii i 个频率上的旋转角度。
  • 最后把 cos 和 sin 各拼一份,变成 32768 × 96。前 48 列和后 48 列相同,是为了和 rotate_half 配合。

第二步:应用旋转。

python 复制代码
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 −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} hidden_size=头数×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 逐行拆代码

第一行:算频率。

python 复制代码
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 10000000=1,频率是 1/1=11/1 = 1 1/1=1。 当指数是 0.979 时, 10000000.979≈7500001000000^{0.979} \approx 750000 10000000.979≈750000,频率是 1/750000≈1.3×10−61/750000 \approx 1.3 \times 10^{-6} 1/750000≈1.3×10−6。

1.0 / (...):

取倒数,得到频率。

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

erlang 复制代码
ω₀ ≈ 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 上下文的一个基础。

第二行:位置序列。

python 复制代码
t = torch.arange(end, device=freqs.device)

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

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

第三行:算角度矩阵。

python 复制代码
freqs = torch.outer(t, freqs).float()

outer 是外积。

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

第 mm m 行第 ii i 列的数,是:
m×ωi m \times \omega_i m×ωi

也就是"位置 mm m 在第 ii i 对上的旋转角度"。

比如:

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

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

第四、五行:算 cos 和 sin。

python 复制代码
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 的写法。

python 复制代码
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 q′=q⊙cos+rotate_half(q)⊙sin

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

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

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

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

1.6.4 两张表

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

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

用的时候,位置 mm m 的 token,直接取第 mm m 行,和它的 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 2ω 度......训练时,模型只见过位置 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 的代码:

python 复制代码
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:

css 复制代码
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{字节} 2×8×8×96×序列长度×2字节

8 个 KV 头,每个 96 维。

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

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

css 复制代码
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{字节} 2×8×1×96×序列长度×2字节

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

GQA:折中方案

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

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

css 复制代码
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{字节} 2×8×4×96×序列长度×2字节

3.2 K 和 V 被投影降维了

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

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

perl 复制代码
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 维:

less 复制代码
苹果的 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 维:

ini 复制代码
吃了的 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 质量
MHA 8 8 最大 最好
GQA 8 4 MHA 的一半 接近 MHA
MQA 8 1 MHA 的 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 Cache 大小=2×层数×KV 头数×head_dim×序列长度×字节数

如果 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 的注意力实现里,核心的注意力计算调用是这样的:

python 复制代码
# 手工路径(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=W2(W1x),展开就是 y=(W2W1)x y = (W_2 W_1)x y=(W2W1)x,等价于一个矩阵。叠 100 层也一样,最终还是等价于一个矩阵。

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

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

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

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

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

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

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

ReLU:一个折线

ReLU 的公式:
ReLU(x)=max⁡(0,x)\text{ReLU}(x) = \max(0, x) ReLU(x)=max(0,x)

图像是一条折线:

perl 复制代码
        |
        |       /
        |      /
        |     /
        |    /
        |   /
        |  /
        | /
        |/
--------+--------
        |

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

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

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

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

SiLU:一条光滑的曲线

SiLU 的公式:
SiLU(x)=x⋅σ(x)\text{SiLU}(x) = x \cdot \sigma(x) SiLU(x)=x⋅σ(x)

其中 σ(x)\sigma(x) σ(x) 是 Sigmoid 函数:
σ(x)= 11+e−x \sigma(x) = \frac{1}{1 + e^{-x}} σ(x)=1+e−x1

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

perl 复制代码
        |
        |         /
        |        /
        |       /
        |      /
        |     /
        |    /
        |   /
--------+--/--------
       /|
      / |
     /  |
    /   |

负数区域,它不直接归零,而是缓慢趋近 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 FFN(x)=ReLU(xW1)W2

三步:

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

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

第三步:降维。 W2 W_2 W2 把 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}} FFN(x)=(SiLU(xWgate)⊙xWup)Wdown

比 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。它用的是 π。

python 复制代码
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

python 复制代码
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))

逐段拆。

第一段:算中间维度。

python 复制代码
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 就是默认公式的结果,不需要任何显式配置。

第二段:三个矩阵。

python 复制代码
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。

python 复制代码
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_proj 768 × 768 589,824
k_proj 768 × 384 294,912
v_proj 768 × 384 294,912
o_proj 768 × 768 589,824
合计 1,769,472

每层 FFN:

矩阵 形状 参数量
gate_proj 768 × 2432 1,867,776
up_proj 768 × 2432 1,867,776
down_proj 2432 × 768 1,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

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

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

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

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

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

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

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

RMSNorm 只做一件事:除以均方根。
RMSNorm(x)= x 1d ∑i=1d xi2+ϵ ⋅γ\text{RMSNorm}(x) = \frac{x}{\sqrt{\frac{1}{d}\sum_{i=1}^{d} x_i^2 + \epsilon}} \cdot \gamma RMSNorm(x)=d1∑i=1dxi2+ϵ x⋅γ

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

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

6.4 MiniMind 的实现

python 复制代码
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。每次计算完,都要加回原始的输入。

python 复制代码
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 块,按顺序做这些事:

ini 复制代码
输入 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_size 768
层数 8
Q 头数 8
KV 头数 4
head_dim 96
FFN 中间维度 2432
归一化 RMSNorm
激活函数 SiLU
位置编码 RoPE(支持 YaRN)

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

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

八、动手实验

实验一:看 RoPE 的旋转效果

python 复制代码
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])

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

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

python 复制代码
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 的效果

python 复制代码
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》。

相关推荐
嘎嘎风1 小时前
MiniMind 学习笔记之 07 模型怎么学会猜词(旋钮是怎么被拧动的)
算法·源码阅读
国科安芯2 小时前
商业立方星平台中高集成度抗辐射MCU的功耗优化与批量化应用可行性探讨
单片机·嵌入式硬件·架构·risc-v·抗辐射·as32x601
ESDWAN4 小时前
企业跨境网络合规建设指南:从线路选择到数据出境的全流程方案
网络·架构
用户5372312882465 小时前
从一句话到水密 STL:给科研工具装 LLM Agent 的安全架构实录
架构
嘎嘎风5 小时前
MiniMind 学习笔记之 05 从 768 维到下一个 token
github·源码阅读
沫璃染墨5 小时前
《从零入门Linux系统篇(五十八):线程篇·十一——线程安全与死锁详解:从可重入到多锁管理》
linux·服务器·开发语言·c++·驱动开发·安全·架构
用户5708462574405 小时前
AI 会话该什么时候重开?一份「上下文卫生」的判断清单
架构
励志不掉头发的内向程序员5 小时前
从鼠标点击到画出一条线:CAD 交互层的状态机设计
后端·架构
bullkingluo5 小时前
从零到一搭建企业级智能问答系统:Ch14 · 三层记忆与断点续跑
架构·llm·agent