从零手写大模型 · 推理篇 10 - Activation Function(激活函数)
推理篇 · 查看系列上期回顾:FFN 里最关键的那一步,我们还没细讲
EP09 我们拆完了 FFN(Feed-Forward Network)的整体结构:
FFN(x) = W2 · Activation(W1 · x + b1) + b2
升维、激活、降维,三步走。升维降维本质是矩阵乘法,容易理解。但中间那个 Activation,也就是激活函数,才是整个 FFN 真正"干活"的地方——甚至可以说,没有它,前后两次线性变换毫无意义。
这一期,我们把激活函数单独拎出来讲清楚:Transformer 原论文用的 ReLU,和 GPT-2 用的 GELU,到底差在哪。
为什么没有激活函数,升维降维就是白费功夫
矩阵乘法是可结合的。如果 FFN 中间没有非线性函数:
W2 · (W1 · x + b1) + b2 = W2W1 · x + (W2b1 + b2) = W' · x + b'
两次线性变换会直接合并成一次。维度从 d_model 升到 4×d_model 再降回来,参数量多花了,效果却和没升过一样。
所以升维这个动作,只有配合非线性函数才有意义——它不是独立生效的,而是在为激活函数"搭建更大的舞台",让更多可调的非线性判断单元参与进来。
ReLU:一个非黑即白的开关
Transformer 原论文(2017,Vaswani et al.)用的是 ReLU:
ReLU(x) = max(0, x)
规则很直白:
- x > 0,原样通过
- x ≤ 0,直接归零
拿几个数值代入感受一下:
x -3 -2 -1 -0.5 0 0.5 1 2 3
y 0.000 0.000 0.000 0.000 0.000 0.500 1.000 2.000 3.000
在 x=0 这个点,函数图像是一个尖角,不平滑——这是 ReLU 天生的一个特点:负数区间的信息被彻底、干脆地丢弃,没有任何缓冲。
GELU:一个按概率放行的软开关
GPT-2 没有沿用 ReLU,而是换成了 GELU(Gaussian Error Linear Unit):
GELU(x) = x · Φ(x)
Φ(x) 是标准正态分布的累积分布函数(CDF),可以理解成一个"通过概率":
- x 越大,Φ(x) 越接近 1,输入几乎原样通过
- x 越小(负得越多),Φ(x) 越接近 0,输入被压得越小
- x = 0 时,Φ(x) = 0.5,正好放行一半
代入同样的数值,对比会更直观:
x -3 -2 -1 -0.5 0 0.5 1 2 3
GELU -0.004 -0.045 -0.159 -0.154 0.000 0.346 0.841 1.954 2.996
ReLU 0.000 0.000 0.000 0.000 0.000 0.500 1.000 2.000 3.000
关键差异就在负数区间:x=-1 时,ReLU 直接给 0,GELU 却是 -0.159。GELU 不会把负数信息一刀切掉,而是保留一小部分、平滑地过渡到 0。这也是为什么说 GELU 是"软开关",ReLU 是"硬开关"。
一个类比:查表机制里的"探测器"
EP09 我们提过一个理解 FFN 的视角:可以把它看成一次逐 token 的"查表"操作。
W1的每一列,是一个"探测器",负责检测当前 token 的向量是否匹配某种模式(做点积,看数值大小)- 激活函数,决定这个探测器"点亮"到什么程度——ReLU 是非0即1的开关,匹配就完全点亮,不匹配就完全熄灭;GELU 则是按概率连续调节亮度
W2的每一行,是这个探测器对应的"取值",一旦被点亮,就按亮度强弱把对应的信息写回残差流
升维的意义,就是造出更多这样的探测器(d_model 越大,探测器数量越多),让模型有能力捕捉更丰富的模式;而激活函数决定了这些探测器是"生硬地开关"还是"柔和地调光"。
GPT-2 为什么选 GELU 而不是 ReLU
工程上的经验总结,主要是这两点:
- 梯度更连续:ReLU 在 x=0 处不可导(严格来说是次梯度),负数区间梯度恒为 0,一旦某个神经元的输入长期为负,它就会"死掉",不再参与训练(业内称为 dying ReLU 问题)。GELU 处处平滑,负数区间仍有微弱梯度,缓解了这个问题。
- 经验效果更好:在 GPT-2、BERT 等大模型的实践中,GELU 相比 ReLU 有稳定的性能提升,逐渐成为 Transformer 系列模型的标配。
需要说明的是,实际框架里(包括 GPT-2 官方实现)很少直接计算 Φ(x),因为它涉及误差函数没有闭式解,计算较慢。工程上普遍使用一个近似公式(tanh 近似)来加速运算,PyTorch 的 nn.GELU() 默认就是走这个近似版本。近似公式本身不影响我们对"软开关"这个核心直觉的理解,这里不展开推导。
PyTorch 代码:把两种激活函数放进"猫吃鱼"里对比
import torch
import torch.nn as nn
# "猫吃鱼" 3个token,d_model=4,升维到4*d_model=16
x = torch.randn(3, 4) # shape: [seq_len=3, d_model=4]
W1 = nn.Linear(4, 16) # 升维
W2 = nn.Linear(16, 4) # 降维
relu = nn.ReLU()
gelu = nn.GELU()
hidden = W1(x) # shape: [3, 16]
hidden_relu = relu(hidden) # 负数全部归零
hidden_gelu = gelu(hidden) # 负数保留一小部分,平滑过渡
out_relu = W2(hidden_relu) # shape: [3, 4]
out_gelu = W2(hidden_gelu) # shape: [3, 4]
print("ReLU 激活后有多少维度被清零:", (hidden_relu == 0).sum().item(), "/", hidden_relu.numel())
print("GELU 激活后有多少维度接近但不等于0:", ((hidden_gelu.abs() < 0.01) & (hidden_gelu != 0)).sum().item())
跑一下这段代码,你会发现 ReLU 那一路清零的维度是一大片、干净利落;GELU 那一路即便数值很小,也很少真正等于 0——这正是"硬开关"和"软开关"在数字上的直接体现。
同样要提醒一句:这里的 W1、W2 都是随机初始化的权重,此刻算出来的数字本身不携带任何语义信息,我们只是在观察激活函数这个"数学开关"的行为方式,不是在看模型真的学到了什么。
小结
| ReLU | GELU | |
|---|---|---|
| 公式 | max(0, x) | x · Φ(x) |
| 负数区间 | 直接归零 | 平滑衰减,保留一小部分 |
| 图像 | x=0处有尖角,不平滑 | 处处平滑 |
| 使用场景 | Transformer 原论文(2017) | GPT-2 及后续主流大模型 |
| 一句话 | 硬截断的开关 | 概率性的软阀门 |
FFN 这一块,我们至此把"升维降维的矩阵乘法"和"中间的非线性开关"都讲透了。下一期,我们要把 Attention、残差连接、LayerNorm、FFN 这四个部件正式拼装成一个完整的 Transformer Block——这是"猫吃鱼"第一次完整地走完一层 Transformer 的全部路径。
从零手写大模型 系列 项目地址:llm-from-scratch