Transformer 在做完「注意力机制」之后,会给每个 token 单独再过一遍一个固定的小网络,把它「加工」一下,再传给下一层。这个固定的小网络就叫前馈神经网络(Feed-Forward Network,简称 FFN)。
它干的事用一句话概括:对每个位置的向量,独立地做两次线性变换,中间夹一个非线性激活函数。听起来抽象,但它要解决的问题其实很朴素。
为什么需要它:没有非线性,再多层也白搭
「神经网络」之所以能学复杂的东西,靠的是「线性变换 + 非线性激活」反复叠加。如果只有线性变换(就是矩阵乘法和加法),那么不管你堆多少层,整个网络在数学上等价于「只乘了一个大矩阵」——表达能力跟一层完全一样。
打个比方:你对一个数字做「先乘 2,再加 3,再乘 4」三步运算,完全可以合成「乘 26 再加 12」一步完成。线性运算怎么叠加都能压缩成一步,所以光靠它,模型根本学不出复杂模式。
FFN 在 Transformer 里的核心作用,就是把这个非线性重新塞回去:在两次线性变换之间,插一个 ReLU(或更现代的 SwiGLU)这种「非线性函数」,让多层叠加真正变得有意义。
图 1:上排三个线性运算能合成一步(叠了等于没叠);下排中间塞一个 ReLU,就再也压不回去了——这就是 FFN 必须插非线性的原因。
它是怎么工作的:一个公式
Transformer 原论文里,FFN 的公式是:
逐个符号拆开看:
:前馈网络对输入 加工后的输出向量,维度回到 。 :输入向量,是一个 token 在当前层的表示(比如 512 维)。可以理解成「这个 token 现在的长相」。 :第一层线性变换的权重矩阵和偏置。 通常把维度从 (如 512)放大到 (如 2048),让信息在更高维的空间里展开。 :就是 ReLU 激活函数。它做的事很简单——负数变 0,正数原样保留。这一步是整个公式里唯一的非线性,没有它,多层网络就退化成一个线性变换。 :第二层线性变换,把维度从 (2048)压回 (512),保证输出能接到下一层去。
一进一出:先升维展开,用 ReLU 切掉负信号,再压回原维度。模型在这个过程中可以挑选「哪些特征该保留、哪些该丢」。
图 2:FFN 的三步——先用 W₁ 把 512 维升到 2048 维展开信息,用 ReLU 切掉负信号,再用 W₂ 压回 512 维接下一层。
一个极简算例
设
第一步,算
第二步,套 ReLU,
第三步,再乘
最终这个 token 经过 FFN 后,输出是
逐位置(position-wise):每个 token 走同一套
FFN 还有一个关键特性:对序列里每个位置(token)都独立、且用同一套参数做变换。
注意「独立」的意思——token A 和 token B 在 FFN 里不交换信息。信息交换是注意力机制的活,FFN 只负责「把每个 token 自己打磨一遍」。而「同一套参数」是说:所有 token 共享同一组
这有点像流水线上的一道「统一加工站」:每个零件单独过一遍同一台机器,机器的设置(参数)对所有零件都一样。这样既保证每个 token 都被深度加工,又不会让参数量随序列长度爆炸。
图 3:三个 token 各自独立过同一台 FFN——参数完全共享,token 之间不交换信息,信息交换是注意力机制的活。
现代大模型:SwiGLU 换掉了 ReLU
原始 Transformer 用 ReLU,但近几年主流大模型(LLaMA、PaLM、Mistral 等)几乎都改用了 SwiGLU。它是 Shazeer 在 2020 年提出的「带门控」的激活函数,大致形式是:
其中
符号解读:
:SwiGLU 前馈网络的输出向量。
直觉上:它把
小结
FFN 是 Transformer 里和注意力机制并列的两大核心组件之一。它做的事就三步:升维、激活(塞进非线性)、降回原维度,并且对每个 token 独立施加同一套变换。注意力负责让 token 之间「互相看见」,FFN 负责让每个 token「自己被深度加工」——二者一外一内,缺一不可。
完整代码
下面是一个可直接运行的 PyTorch 实现,包含经典的 ReLU 版本和现代的 SwiGLU 版本:
import torch
import torch.nn as nn
import torch.nn.functional as F
class PositionWiseFFN(nn.Module):
"""经典版 FFN:两层线性 + ReLU,对应原论文公式"""
def __init__(self, d_model: int, d_ff: int):
super().__init__()
# linear1 对应 W1, b1:把 d_model 维升到 d_ff 维
self.linear1 = nn.Linear(d_model, d_ff)
# linear2 对应 W2, b2:把 d_ff 维压回 d_model 维
self.linear2 = nn.Linear(d_ff, d_model)
def forward(self, x):
# x 形状: (batch, seq_len, d_model)
h = self.linear1(x) # 升维:xW1 + b1
h = F.relu(h) # max(0, ·),引入非线性
return self.linear2(h) # 降维:(...)W2 + b2
class SwiGLUFFN(nn.Module):
"""现代版 FFN:SwiGLU 门控激活(LLaMA 等大模型常用)"""
def __init__(self, d_model: int, d_ff: int):
super().__init__()
# SwiGLU 需要两个并列的线性投影:w_gate 和 w_value
self.w_gate = nn.Linear(d_model, d_ff)
self.w_value = nn.Linear(d_model, d_ff)
# 再把门控后的 d_ff 维压回 d_model
self.w_out = nn.Linear(d_ff, d_model)
def forward(self, x):
# silu(x) = x * sigmoid(x),即 Swish;门控 = silu(gate) * value
gated = F.silu(self.w_gate(x)) * self.w_value(x)
return self.w_out(gated)
# === 跑一下试试 ===
torch.manual_seed(0)
d_model, d_ff = 8, 16
ffn = PositionWiseFFN(d_model, d_ff)
# 假数据:1 个 batch、3 个 token、每个 token 8 维
x = torch.randn(1, 3, d_model)
y = ffn(x) # 前向:每个 token 独立过同一套 FFN
print("输出形状:", y.shape) # (1, 3, 8),维度压回了 d_model
# 训练一步演示(参数会被更新)
target = torch.randn_like(y)
loss = F.mse_loss(y, target) # 算一个简单的损失
loss.backward() # 反向传播,算梯度
for p in ffn.parameters():
p.data -= 0.01 * p.grad # 手动做一步梯度下降
print("loss:", loss.item())输出形状会是 (1, 3, 8)——三个 token 各自被同一套参数加工了一遍,维度回到
参考资料
- Attention Is All You Need — Vaswani et al., 2017 https://arxiv.org/abs/1706.03762
- GLU Variants Improve Transformer — Shazeer, 2020 https://arxiv.org/abs/2002.05202
- Position-wise Feed-Forward Network (FFN) — labml.ai https://nn.labml.ai/transformers/feed_forward.html