上一篇文章解释了点积和矩阵乘法,说白了,矩阵乘法就是一种“转换”。这一篇,我们来看看这种转换在 Transformer 里到底是怎么被用上的。

我们先从本质说起。LLM 做的这件事,本质上就是预测下一个 token 是什么。在阶段二,模型用海量的互联网内容做训练,通过自监督学习的方式,不断调整那 1750 亿个参数,直到它能正确地补全文本。
但要注意,阶段二训练结束时,模型其实只会一件事——补全文本。它还不具备问答、对话的能力。
好,现在假设所有参数都已经调好了。我们用一个具体的例子,输入 "The cat sat",来看看模型是怎么一步步预测出下一个词是 "on" 的。
从输入到向量坐标
输入 "The cat sat" 首先会经过 Tokenizer,把这串文字拆成 the、cat、sat 三个 token。然后,每个 token 会去词汇对照表(Vocab Map
于是,我们得到了三个形状为 [1, 4096] 的向量,比如:Vector_The: [0.1, -0.5, ...]、Vector_cat: [0.8, 0.2, ...]、Vector_sat: [-0.1, 0.9, ...]。把它们合并成一个矩阵,就得到了一个形状为 [3, 4096] 的输入。
进入 Layer 一步步加工
接下来,就正式进入 Layer 的加工流程了。之前说过,整个模型有 96 层 Layers,每一层里都包含两个核心模块:MHA(多头注意力机制)和 FFN(前馈神经网络)。那个形状为 [3, 4096] 的矩阵,会完整地经历所有 96 层,最终得到加工后的 [3, 4096]。
每一层 Layer 的完整流程可以用下面这个公式来表示:
x_in ➔ [Norm] ➔ [MHA] ➔ (+ 残差连接) ➔ x_mid ➔ [Norm] ➔ [FFN] ➔ (+ 残差连接) ➔ x_out
而这一层的输出 x_out,又会成为下一层 Layer 的输入 x_in。
先说 Norm 层归一化
Norm 就是 Layer Normalization(层归一化)。矩阵乘法的结果,数值范围可能非常大,有的值能到 50000,有的却只有 0.000003。为了防止计算溢出或者梯度在训练时乱跳,需要把这些值统一处理一下,让它们的均值为 0,方差为 1。
Norm 的计算分四步。我们以向量 cat: [10, 2, 12, 0] 为例来看。
-
求均值 (μ)
(10 + 2 + 12 + 0) ÷ 4 = 6 -
求方差 (σ²)
10 → (10-6)² = 16
2 → (2-6)² = 16
12 → (12-6)² = 36
0 → (0-6)² = 36
方差 = (16 + 16 + 36 + 36) ÷ 4 = 26,标准差 σ = √26 ≈ 5.1 -
归一化 (Normalize)
公式是:(x - 均值) / 标准差。这一步的目的,就是把数据强行拉到“均值为0,方差为1”的标准形态。
10 → (10 - 6) / 5.1 ≈ 0.78
2 → (2 - 6) / 5.1 ≈ -0.78
12 → (12 - 6) / 5.1 ≈ 1.17
0 → (0 - 6) / 5.1 ≈ -1.17
结果向量: [0.78, -0.78, 1.17, -1.17] -
缩放与平移 (Scale & Shift)
如果每次都强行变成0均值,可能会破坏数据本身的含义。所以模型学了两个可变的参数:γ (缩放) 和 β (平移)。可以理解为一个线性函数。假设模型觉得这一层数值需要稍微大一点:
γ = [2, 2, 2, 2],β = [1, 1, 1, 1]
最终输出 = 归一化结果 × γ + β
0.78 × 2 + 1 = 2.56,以此类推。
残差连接
在经过 MHA 和 FFN 处理后,为了防止原始信息在层层传递中丢失,会把处理前的原始值再加回来。
Output = New_Process(x) + x
MHA 多头注意力机制
回到公式。x_in 是形状为 [3, 4096] 的矩阵,经过 Norm 后还是 [3, 4096]。接着就进入了 MHA。
MHA 里有四个在阶段二训练好的矩阵:W_Q、W_K、W_V 和 W_O,它们的形状都是 [d_model, d_model],也就是 [4096, 4096]。
Q 是 Question,K 是 Key,V 是 Value。这三个名字很抽象,但作用其实很直观:它们负责把 "The cat sat" 中三个 token 的向量坐标互相融合。比如,the 这个 token 需要更关注 cat,经过融合后,the 的向量值里就包含了大量 cat 的向量信息。
至于“多头”是什么意思?假设有 32 个头,4096 ÷ 32 = 128。就是把 W_Q、W_K、W_V 这三个大矩阵,分别拆成 32 个形状为 [4096, 128] 的小矩阵,让它们并行计算,得到 32 个结果,最后再合并起来乘以 W_O,得到最终的产物。
这也是为什么 Transformer 如此依赖 GPU 算力——大量的并行矩阵计算,正是 GPU 的强项。
“头”这个概念,可以用一个角度来类比。以 "The cat sat on the mat because it was tired." 这句为例:
- Head 1 (语法眼):专门盯着主谓关系。它发现 "it" 指代的是 "cat"。
- Head 2 (逻辑眼):专门盯着因果关系。它发现 "because" 导致了 "tired"。
- Head 3 (位置眼):专门盯着方位关系。它关注的是 "on the mat"。
整个 MHA 的计算过程,可以用一个公式来概括:
其中,Scores = (X W_Q)(X W_K)^T。
X W_Q 算出 Q 矩阵,X W_K 算出 K 矩阵,然后对 K 矩阵做转置。为什么要转置?很简单,为了能让矩阵乘法顺利进行——前一个矩阵的列数,必须等于后一个矩阵的行数。
除以 √d_model,也是为了控制数值范围,防止向量值之间的差距过大。
接着,softmax 把 Scores 变成一堆总和为 1 的概率。然后再乘以 V 矩阵,得到最终的注意力输出 Z。
这里有一个关键点:输入的 X 是形状为 [3, 4096] 的矩阵,而不是单个 token。在这个过程中,每个 token 之间都会互相融合。但为了确保模型在预测时,前一个 token 看不到后一个 token,需要用 Mask 矩阵来实现。
Mask 矩阵是一个上三角矩阵,右上角全是负无穷。任何矩阵加上它,右上角都会变成负无穷。而 e^(-∞) = 0,这就保证了在预测时,后面的 token 不会偷看到前面的信息。
可能有人会问:既然后一个 token 的向量已经包含了前面所有信息,那为什么还要算所有向量的 Q、K、V?这是因为,后一个向量的计算过程,本身就用到了这些数据。
不妨先用一个极度简化的例子(d_model = 3)来感受一下 MHA 的完整计算过程。
输入 X:
假设阶段二训练好的参数矩阵如下:
经过矩阵乘法后,得到 Q、K、V:
计算 Scores = Q * K^T,再加上 Mask 矩阵,得到:
经过 softmax 后,得到注意力权重 A:
第三行是怎么算出来的?softmax(Scores / √3)。√3 ≈ 1.73。cat 得分 4 / 1.73 ≈ 2.3,The 和 sat 得分 0。e^2.3 ≈ 10,e^0 = 1。概率 = 10 / (1 + 10 + 1) ≈ 0.84。所以是 [0.08, 0.84, 0.08]。
最后,Z = A × V:
可以看到,经过 MHA 后,sat 的向量明显向 cat 倾斜了。
在实际的 32 头注意力中,每个头都会并行得到这样一个 Z 向量。我们只需要取最后一个 token(sat)对应的那一行,把 32 个头的结果拼接起来,就又成了一个形状为 [1, 4096] 的向量。向量的 0~128 位是头 1 表示语法,头 2 的 128~256 位表示位置,以此类推到第 32 个头。最后再乘以 W_O 矩阵,把这些独立的信息融为一体,得到最终的结果。这就是 MHA 的全部过程。
FFN 前馈神经网络
Z 经过残差连接和层归一化后,进入 FFN。
FFN 里有两个在阶段二训练好的矩阵,W₁ 和 W₂。W₁ 是升维矩阵,形状是 [d_model, 4*d_model];W₂ 是降维矩阵,形状是 [4*d_model, d_model]。
假设 Z 已经完成残差连接和层归一化。乘以 W₁ 矩阵:
H_up = Z_sat × W₁
升维之后,模型能捕捉到更细节的信息。比如,一个 token “Apple”,升维前它的向量可能是 [0.8, 0.1, 0.5],分别代表“是水果”、“是公司”、“是红色的”。
W₁ 就像一份包含 6 个问题的问卷:是电子产品吗?是红色的吗?能吃吗?是交通工具吗?有毛吗?是液体吗?
计算 H = x × W₁,结果向量可能是 [10, 8, 9, -5, -10, -2]。然后经过 ReLU 激活函数(f(x) = max(0, x)),把所有负数置为 0,变成 [10, 8, 9, 0, 0, 0]。再通过 W₂ 把结果压缩回去,比如变成 [5.0, 2.0, 8.0]。
最后再做一次残差连接,整个 Layer 层的加工就结束了。这个结果会被作为输入,传给下一层 Layer。
所有 Layers 处理完后
输入的是 "The cat sat" 三个 token 组成的形状为 [3, 4096] 的矩阵,经过 96 层 Layers 处理后,输出的还是这个形状。但我们只需要 sat 对应的那一行 H = [1, 4096],因为它已经包含了 the、cat 和 sat 的混合信息。
接着,H 做完 Norm(处理为均值为 0、方差为 1),形状依然是 [1, 4096]。
还记得最开始那个形状为 [128k, 4096] 的 Embedding Table 吗?这里有一个 unembedding 矩阵,形状是 [4096, 128k],为了节省空间,这两个 table 其实是一样的。
用 H 和 unembedding 矩阵做乘法,得到 [1, 128k] 的结果,正好对应整个词汇表里的每个词,这个结果被称为 Logits。
再用 Logits 做 softmax 转成概率,同时结合 temperature 配置,并降低之前已经出现过词语的 Logits。这时候的概率,就对应着词汇表中,下一个词应该是什么的概率。
最后,再结合 Top-K(保留概率最高的前 K 个)和 Top-P(保留概率累计超过 P 的前 n 个),求它们的交集,从保留的 token 里随机选一个,拼接到整个句子后面。然后,把新的句子从头开始,进入下一轮循环。
这个过程会一直持续,直到达到 output token 的限制,或者遇到了标记为 EOS(End of Sequence)的 token,整个推理过程才算结束。
这才是关键所在。模型从头到尾,其实就是在做这件事:根据已有的上下文,计算出下一个最有可能出现的词是什么。
