Transformer中的前馈网络设计从ReLU到SwiGLU的激活函数演进一、前馈网络在Transformer中的功能定位Transformer的每一层由两个核心子层组成多头自注意力Multi-Head Self-Attention和前馈网络Feed-Forward Network, FFN。注意力负责token间的信息交互这个token应该关注哪些tokenFFN负责token内的特征变换给定交互后的信息如何映射到更有用的表示。原始Transformer的FFN设计出奇地简单——两个线性变换夹一个ReLU激活函数$$FFN(x) ReLU(xW_1 b_1)W_2 b_2$$其中$W_1 \in \mathbb{R}^{d_{model} \times d_{ff}}$将维度从$d_{model}$扩展到$d_{ff}$通常$d_{ff}4\times d_{model}$$W_2 \in \mathbb{R}^{d_{ff} \times d_{model}}$再压缩回原始维度。ReLU作为唯一的非线性来源承载了整个token级别特征转换的表示能力。近年来的研究尤其是LLaMA、PaLM等大规模语言模型的实践表明FFN的激活函数选择对模型最终性能有显著影响——在同等参数量和训练数据下激活函数的改进可以带来1-2%的困惑度Perplexity提升这在LLM领域是值得关注的提升幅度。二、ReLU → GELU → Swish → SwiGLU的演进逻辑这一演进可以分解为两个独立的改进方向。方向一从ReLU到更平滑的激活函数。ReLU在x0时梯度为零导致死亡神经元问题某些神经元可能在训练中永久失活。GELUGaussian Error Linear Unit, 用于BERT和GPT-2通过对输入乘以标准正态CDF来平滑激活$$GELU(x) x \cdot \Phi(x) \approx 0.5x(1 \tanh(\sqrt{2/\pi}(x 0.044715x^3)))$$GELU在x0附近保持非零梯度缓解了死亡神经元问题。在BERT的训练中GELU相对于ReLU在GLUE benchmark上带来了约0.5-1%的提升。SwishRamachandran et al., 2017使用sigmoid作为门控函数$$Swish(x) x \cdot \sigma(x) \frac{x}{1e^{-x}}$$Swish在深层网络中的优势在于其非单调性——在x0时不是单调递减的这为网络提供了额外的表示灵活性。方向二从简单激活到门控线性单元GLU。GLUDauphin et al., 2017将激活函数从对线性变换的输出进行非线性映射改为使用门控机制动态控制信息流$$GLU(x) (xW_1 b_1) \otimes \sigma(xW_2 b_2)$$其中$\otimes$是逐元素乘法$\sigma$是sigmoid门控函数。GLU的关键创新在于一半的变换负责生成内容另一半负责生成门控信号来决定内容中哪些部分可以通过。 各种FFN激活函数的PyTorch实现与参数量对比 import torch import torch.nn as nn import torch.nn.functional as F class ReLUFFN(nn.Module): 原始Transformer的FFNReLU激活 def __init__(self, d_model: int, d_ff: int): super().__init__() self.w1 nn.Linear(d_model, d_ff) self.w2 nn.Linear(d_ff, d_model) def forward(self, x: torch.Tensor) - torch.Tensor: return self.w2(F.relu(self.w1(x))) class GELUFFN(nn.Module): BERT/GPT-2风格的FFNGELU激活 def __init__(self, d_model: int, d_ff: int): super().__init__() self.w1 nn.Linear(d_model, d_ff) self.w2 nn.Linear(d_ff, d_model) def forward(self, x: torch.Tensor) - torch.Tensor: return self.w2(F.gelu(self.w1(x), approximatetanh)) class SwiGLUFFN(nn.Module): LLaMA/PaLM风格的FFNSwiGLU激活。 关键架构差异 1. 使用门控机制content * gate 2. d_ff 通常取 8/3 * d_model而非4×以补偿门控引入的额外参数 3. 三个线性层两个上投影 一个下投影 参数计算 - ReLU FFN: 2 * d_model * d_ff 8 * d_model^2当d_ff4*d_model时 - SwiGLU FFN: 3 * d_model * d_ff ≈ 8 * d_model^2当d_ff8/3*d_model时 两者参数量近似相等但SwiGLU的性能更好。 def __init__(self, d_model: int, d_ff: int None): super().__init__() # d_ff通常设为 8/3 * d_model向上取最近的64的倍数 if d_ff is None: d_ff int(8 * d_model / 3) d_ff ((d_ff 63) // 64) * 64 # 向上取最近的64倍数 self.w_gate nn.Linear(d_model, d_ff, biasFalse) # 门控路径 self.w_up nn.Linear(d_model, d_ff, biasFalse) # 内容路径 self.w_down nn.Linear(d_ff, d_model, biasFalse) # 下投影 def forward(self, x: torch.Tensor) - torch.Tensor: SwiGLU前向传播。 SwiGLU(x) (xW_up ⊙ Swish(xW_gate)) W_down 其中 Swish(z) z * sigmoid(z) ⊙ 表示逐元素乘法 Args: x: (batch, seq_len, d_model) Returns: (batch, seq_len, d_model) # 门控信号swish x * sigmoid(x)即SiLU激活 gate F.silu(self.w_gate(x)) # SiLU Swish不同名称相同公式 # 内容信号 up self.w_up(x) # 门控乘法 下投影 return self.w_down(gate * up) def count_parameters(self) - dict: 计算各部分的参数量 return { w_gate: sum(p.numel() for p in self.w_gate.parameters()), w_up: sum(p.numel() for p in self.w_up.parameters()), w_down: sum(p.numel() for p in self.w_down.parameters()), total: sum(p.numel() for p in self.parameters()), }三、门控机制为何有效信息流控制的视角SwiGLU相对于ReLU FFN的性能提升可以从信息流控制的角度进行解释。在ReLU FFN中非线性变换是全或无的——ReLU将负值全部截断为0这意味着在每层FFN中一部分信息被不可逆地丢弃。被截断的信息无法在后续层中恢复。SwiGLU通过门控机制sigmoid提供了软的信息流控制门控值在0-1之间可以精细地控制信息保留的比例。与ReLU的硬截断不同sigmoid的平滑过渡使得模型可以学习部分保留某些特征。从优化角度看门控机制的梯度总是非零的sigmoid的导数在任意点都0避免了ReLU的死亡神经元问题。还有另一个角度SwiGLU中的门控路径和内容路径是两个独立的线性投影这意味着模型可以为是否通过和通过什么学习不同的特征变换。这种分离可能提供了比ReLU中单一的变换截断更丰富的表示能力。四、不同规模下的选择建议激活函数的最优选择依赖于模型规模和训练数据量。基于现有文献和实践经验小模型100M参数GELU和ReLU的差距很小0.5%选择主要取决于是否与其他模型如BERT保持一致以利用其预训练权重。中等模型100M-1B参数SwiGLU开始展示优势但需要调整d_ff以保持参数量可比使用8/3倍而非4倍d_model。大模型1B参数SwiGLU已经成为事实标准——LLaMA、LLaMA 2、Mistral、Qwen等主流LLM均采用此方案。在大规模训练中SwiGLU相对于ReLU在同等计算预算下的困惑度提升约为2-5%。五、总结Transformer前馈网络中的激活函数从ReLU到SwiGLU的演进反映了一个从简单非线性到学习的信息门控的范式迁移。GELU通过平滑激活函数缓解了ReLU的梯度消失问题Swish/SiLU通过自门控sigmoid(x)·x进一步提升了表示能力SwiGLU将门控机制与线性变换解耦为模型提供了独立学习信息内容和信息通过率的能力。这一演进在大规模语言模型的实践中已被充分验证——SwiGLU已成为2023年后发布的几乎所有主流LLM的默认选择。在选择激活函数时应同时调整FFN的中间维度d_ff以保持参数量可比——SwiGLU的d_ff通常设为8/3×d_model而非传统ReLU FFN的4×d_model。