从零手写大模型 · 推理篇 09 - FFN(前馈神经网络)
推理篇 · 查看系列猫吃鱼走到这一步:Attention 之后,还差一步
前几期我们把"猫吃鱼"这句话,一步步从三个 token 变成了一个融合了上下文信息的 3×4 矩阵:
- Embedding:三个词各自变成 4 维向量
- Attention:每个词看了看其他词,把信息"借"了过来
- 残差连接 + LayerNorm:借来的信息叠加回原始向量,再做一次归一化整理
到这一步,X(3×4 矩阵)里的每一行,已经不是孤立的"猫""吃""鱼"了,而是三个都掺了点彼此味道的向量。这一期要讲的 FFN(Feed-Forward Network,前馈神经网络),做的是一件方向完全不同的事:不再看别的词,而是让每个词自己"消化"一下刚刚借来的信息。
FFN 的公式,先看骨架
FFN(x) = W2 · Activation(W1 · x + b1) + b2
三步:升维、激活、降维。用形状表示会更直观(以"猫吃鱼"为例,d_model=4,GPT-2 的惯例是把中间维度升到 4 倍,也就是 d_ff=16):
X : [3, 4] ← 猫吃鱼,3个token,每个4维
W1 : [4, 16] ← 升维矩阵
X·W1+b1 : [3, 16] ← 升到16维
激活函数 : [3, 16] ← 形状不变,数值被"筛选"
W2 : [16, 4] ← 降维矩阵
输出 : [3, 4] ← 降回4维,形状和输入一致
一个很重要的细节:FFN 是逐行(逐 token)独立计算的,W1、W2 对"猫""吃""鱼"这三行用的是完全相同的一套参数,词与词之间在这一步不发生任何交互——这正好跟 Attention 相反。Attention 是"跨 token 换信息",FFN 是"关起门来自己整理"。两者交替堆叠,才是一层 Transformer Block 的完整能力。
拿"鱼"这个 token 手算一遍
为了让升维降维不只是一个抽象公式,我们挑"鱼"这一行,代入具体数字走一遍。假设经过 Attention + 残差 + LayerNorm 之后,"猫吃鱼"对应的矩阵 X 是(数值为教学简化,非真实训练结果):
第1维 第2维 第3维 第4维
猫 0.9 -0.3 0.1 -0.7
吃 -0.4 1.0 -0.6 0.2
鱼 0.6 -0.5 0.8 -0.4
取"鱼"这一行:x_鱼 = [0.6, -0.5, 0.8, -0.4]
W1 是 4×16 的矩阵,也就是有 16 个"探测方向"(16 列)。为了方便手算,这里只展示其中 4 列(完整的 16 列在代码里跑,逻辑是一样的):
列1 列2 列3 列4
[ 0.5 -1.0 0.3 -0.6 ]
[ 0.5 1.0 -0.2 -0.6 ]
[ 0.5 -1.0 0.4 0.6 ]
[ 0.5 1.0 -0.1 0.6 ]
x_鱼 分别和这 4 列做点积(先忽略 b1,设为 0):
h1 = 0.6×0.5 + (-0.5)×0.5 + 0.8×0.5 + (-0.4)×0.5 = 0.25
h2 = 0.6×(-1) + (-0.5)×1 + 0.8×(-1) + (-0.4)×1 = -2.30
h3 = 0.6×0.3 + (-0.5)×(-0.2) + 0.8×0.4 + (-0.4)×(-0.1) = 0.64
h4 = 0.6×(-0.6) + (-0.5)×(-0.6) + 0.8×0.6 + (-0.4)×0.6 = 0.18
升维之后,"鱼"这个词在这 4 个探测方向上的原始得分是 [0.25, -2.30, 0.64, 0.18](其余 12 维省略,逻辑相同)。
激活函数登场:GELU 怎么处理这几个数
GPT-2 用的激活函数是 GELU:GELU(x) = x · Φ(x),Φ 是标准正态分布的累积概率,直觉上是一个"按概率放行"的软开关,负数不会被一刀切掉,而是保留一点点、平滑地压小。代入刚才这 4 个数:
原始值 GELU(软开关) ReLU(硬开关,作对比)
h1 0.25 0.150 0.25
h2 -2.30 -0.025 0.00
h3 0.64 0.473 0.64
h4 0.18 0.103 0.18
重点看 h2:原始值是 -2.30,说明"鱼"在这个探测方向上是明显"不匹配"的。如果用 ReLU,这一项直接被清零,后面完全学不到"到底负了多少";GELU 给出的是 -0.025,虽然接近 0,但保留了一点负向的痕迹——这也是我们将在第 10 期展开讲解的"硬开关 vs 软开关"在真实数值上的样子。
降维:把 16 维的判断结果收拢回 4 维
激活之后的 16 维向量,再乘以 W2(16×4),加上 b2,收拢回原来的 4 维空间。这一步可以理解成:16 个探测器分别给出"匹配程度"之后,按各自的权重加权汇总,写回原始维度。汇总完的结果,会在下一步和 x_鱼 做残差相加(x2 = x1 + FFN(LN(x1))),跟这一期没被激活、也没被丢弃的原始信息合并到一起。
一个类比:升维是造探测器,激活函数是决定谁点亮
如果只把 FFN 理解成"矩阵乘法升维再降维",会漏掉最关键的部分——没有激活函数,升维降维在数学上等价于什么都没做(两次线性变换可以合并成一次,参考文末的推导)。真正在起作用的,是升维之后那一大排"探测器"(W1 的每一列),加上激活函数决定每个探测器"点亮"到什么程度。可以类比成:
W1的每一列:一个探测器,负责检测当前词在某个方向上是否有分量- 激活函数:决定这个探测器的输出强度——GELU 是按概率连续调光,ReLU 是非0即1的开关
W2的每一行:探测器点亮之后,把对应的信息按强度加权写回原始空间
升维让探测器数量变多(更细的判断粒度),激活函数决定这些判断怎么被筛选和组合——这两者配合,才是 FFN 真正的非线性来源。
PyTorch 代码:还原"猫吃鱼"的完整 FFN
import torch
import torch.nn as nn
d_model = 4
d_ff = 4 * d_model # GPT-2 惯例:中间维度是 d_model 的4倍
class FeedForward(nn.Module):
def __init__(self, d_model, d_ff):
super().__init__()
self.W1 = nn.Linear(d_model, d_ff) # 升维
self.act = nn.GELU() # GPT-2 用 GELU,Transformer原论文用 ReLU
self.W2 = nn.Linear(d_ff, d_model) # 降维
def forward(self, x):
x = self.W1(x) # [3, 4] -> [3, 16]
x = self.act(x) # 形状不变,数值被筛选
x = self.W2(x) # [3, 16] -> [3, 4]
return x
ffn = FeedForward(d_model, d_ff)
# 猫吃鱼,LayerNorm之后的输出,3个token,每个4维
x = torch.tensor([
[0.9, -0.3, 0.1, -0.7], # 猫
[-0.4, 1.0, -0.6, 0.2], # 吃
[0.6, -0.5, 0.8, -0.4], # 鱼
])
out = ffn(x)
print(out.shape) # torch.Size([3, 4]),形状和输入一致,方便后面做残差相加
一个容易漏掉的推导:为什么没有激活函数就白算了
W2(W1·x + b1) + b2 = W2W1·x + (W2b1 + b2) = W'·x + b'
如果 self.act 换成什么都不做(恒等函数),上面这段代码在数学上会退化成一个单独的线性层,d_ff=16 这个中间维度形同虚设。这也是为什么激活函数不是"锦上添花",而是 FFN 存在的必要条件。
老规矩:这些数字目前还没有意义
这一期代码里 W1、W2 都是随机初始化的,我们手算的那几个数字,只是用来展示"升维、激活、降维"这套流程本身是怎么运作的,跟"鱼"这个词真正的语义完全无关——真正有意义的权重,要等到后面加载 GPT-2 预训练参数那一期才会出现。
下一期,我们会进一步拆解激活函数,对比 ReLU 和 GELU。之后再把 Attention、残差连接、LayerNorm、FFN 这四块正式拼装成一个完整的 Transformer Block,让"猫吃鱼"第一次完整地走完一层 Transformer 的全部路径。
从零手写大模型 系列 项目地址:llm-from-scratch