跳转至

Activation Functions 与 FFN 设计:从 ReLU 到 SwiGLU

更新日期:2026-04-17

本文目标:深入理解 Transformer 中 FFN 层的角色、激活函数的演进、GLU 变体的数学原理与工程实现。能够为新模型选择合适的 FFN 架构并正确计算参数量与显存。


一、FFN 在 Transformer 中的角色

1.1 Attention 是通信,FFN 是计算

Transformer 的每一层由两个子模块构成:

\(\text{TransformerLayer}(x) = \text{FFN}(\text{Attention}(x) + x) + \text{Attention}(x) + x\)

两者的分工: 以 LLaMA-2 7B 为例(\(d=4096\)):每层 Attention 参数 \(\approx 4 \times 4096^2 = 67M\)(因 GQA 实际更少),FFN 参数 \(= 3 \times 4096 \times 11008 = 135M\),FFN 占每层参数的 ~67%。

1.2 FFN 作为知识存储

Dai et al. (2022) 发现 FFN 中存在"知识神经元"(Knowledge Neurons):特定事实知识(如"东京是日本的首都")会激活 FFN 中间层的少数特定神经元。

  • Key-Value Memory 解释(Geva et al., 2021):将 FFN 视为一个键值存储系统。\(W_1\) 的每一行是一个"键"(pattern detector),\(W_2\) 的对应列是一个"值"(output distribution over vocab)

  • 当输入匹配某个键时,对应的值被激活并加入残差流

\(\text{FFN}(x) = \sum_{i=1}^{d_{ff}} f(x \cdot w_1^{(i)}) \cdot w_2^{(i)}\)

其中 \(f\) 是激活函数,\(w_1^{(i)}\)\(W_1\) 的第 \(i\) 行,\(w_2^{(i)}\)\(W_2\) 的第 \(i\) 列。每个 \(f(x \cdot w_1^{(i)})\) 是一个标量"匹配度",\(w_2^{(i)}\) 是对应的"记忆内容"。

参考:Transformer Feed-Forward Layers Are Key-Value Memories (Geva et al., 2021) arxiv:2012.14913Knowledge Neurons in Pretrained Transformers (Dai et al., 2022) arxiv:2104.08696


二、Activation 函数演进

参考:GLU Variants Improve Transformer (Shazeer, 2020) arxiv:2002.05202Gaussian Error Linear Units (GELUs) (Hendrycks & Gimpel, 2016) arxiv:1606.08415Primer: Searching for Efficient Transformers for Language Modeling (So et al., 2021) arxiv:2109.08668


三、GLU 变体深度解析

3.1 从原始 GLU 到 SwiGLU

原始 GLU(Gated Linear Unit, Dauphin et al. 2017):

\(\text{GLU}(x) = (xW_1 + b_1) \otimes \sigma(xW_2 + b_2)\)

其中 \(\sigma\) 是 sigmoid,\(\otimes\) 是逐元素乘法。右半部分 \(\sigma(xW_2)\) 是一个"门",控制左半部分 \(xW_1\) 的哪些维度被通过。

Shazeer (2020) 将 sigmoid 门替换为其他激活函数,形成一族 GLU 变体:

\(\text{SwiGLU}(x, W_g, W_u, W_d) = (\text{Swish}(xW_g) \otimes xW_u) W_d\)

\(\text{GeGLU}(x, W_g, W_u, W_d) = (\text{GELU}(xW_g) \otimes xW_u) W_d\)

\(\text{ReGLU}(x, W_g, W_u, W_d) = (\text{ReLU}(xW_g) \otimes xW_u) W_d\)

3.2 为什么门控有效

门控的本质是乘法交互(multiplicative interaction):

  • 加法交互:\(f(x) + g(x)\) — 两个信号独立叠加

  • 乘法交互:\(f(x) \otimes g(x)\) — 一个信号控制另一个信号的通过量

乘法交互能产生比加法更"锐利"的特征选择。设 \(f(x)\) 在某维度输出 0.01,则不论 \(g(x)\) 多大,乘积都接近 0——相当于硬关闭该维度。

直觉上:gate 路径 \(\text{Swish}(xW_g)\) 学到的是"特征是否重要",而 up 路径 \(xW_u\) 学到的是"特征的值是多少"。两者分离再相乘,比单一路径同时学"是否"和"多少"更有效。

3.3 三矩阵设计:\(W_g\), \(W_u\), \(W_d\)

标准 ReLU FFN 只需 2 个矩阵:

\(\text{FFN}_{\text{ReLU}}(x) = \text{ReLU}(xW_1)W_2\)

  • \(W_1 \in \mathbb{R}^{d \times d_{ff}}\):上投影(expand)

  • \(W_2 \in \mathbb{R}^{d_{ff} \times d}\):下投影(contract)

  • 参数量:\(2 \times d \times d_{ff} = 2 \times d \times 4d = 8d^2\)

SwiGLU FFN 需要 3 个矩阵:

\(\text{FFN}_{\text{SwiGLU}}(x) = (\text{Swish}(xW_g) \otimes xW_u) W_d\)

  • \(W_g \in \mathbb{R}^{d \times d_{ff}}\):gate 投影

  • \(W_u \in \mathbb{R}^{d \times d_{ff}}\):up 投影

  • \(W_d \in \mathbb{R}^{d_{ff} \times d}\):down 投影

  • 参数量:\(3 \times d \times d_{ff}\)

3.4 \(d_{ff} = \frac{8}{3}d\) 的由来

为保持 SwiGLU 和 ReLU FFN 的总参数量相等:

\(3 \times d \times d_{ff}^{\text{SwiGLU}} = 2 \times d \times d_{ff}^{\text{ReLU}}\)

\(d_{ff}^{\text{SwiGLU}} = \frac{2}{3} \times d_{ff}^{\text{ReLU}} = \frac{2}{3} \times 4d = \frac{8}{3}d \approx 2.667d\)

实际实现中 \(d_{ff}\) 还需对齐到特定倍数(通常是 128 或 256)以确保 Tensor Core 高效运算。

def compute_ffn_dim(d_model, multiplier=8/3, align_to=256):
    raw = int(d_model * multiplier)
    return ((raw + align_to - 1) // align_to) * align_to

3.5 SwiGLU FFN 伪代码

class SwiGLUFFN(nn.Module):
    def __init__(self, d_model, d_ff):
        self.w_gate = nn.Linear(d_model, d_ff, bias=False)
        self.w_up   = nn.Linear(d_model, d_ff, bias=False)
        self.w_down = nn.Linear(d_ff, d_model, bias=False)

    def forward(self, x):
        gate = F.silu(self.w_gate(x))   # Swish/SiLU activation
        up   = self.w_up(x)
        return self.w_down(gate * up)

3.6 实际模型的 \(d_{ff}\) 取值

参考:LLaMA: Open and Efficient Foundation Language Models (Touvron et al., 2023) arxiv:2302.13971Mistral 7B (Jiang et al., 2023) arxiv:2310.06825DeepSeek-V3 Technical Report (DeepSeek-AI, 2024) arxiv:2412.19437


四、GELU vs SwiGLU 实证对比

4.1 PaLM 消融实验

PaLM (Chowdhery et al., 2022) 在 8B 规模上做了激活函数消融实验。控制总参数量相同(SwiGLU 使用 \(\frac{8}{3}\) 比值): 参考:PaLM: Scaling Language Modeling with Pathways (Chowdhery et al., 2022) arxiv:2204.02311

4.2 训练 Loss 曲线特征

SwiGLU 与 GELU 的 loss 曲线差异:

  • 早期收敛(前 10% tokens):SwiGLU loss 下降更快,门控机制让模型更早学会"关闭"无用特征

  • 中期稳定(10%~80%):SwiGLU 维持 ~0.02-0.05 的 loss 优势,这一差距在 scaling law 下对应显著的下游性能差异

  • 后期(>80%):差距趋于稳定,不会继续扩大

4.3 Compute-Quality 权衡

结论:SwiGLU 在理论 FLOPs 相同、参数量相同的条件下,稳定赢 0.5-1%——这在大模型训练中是极显著的免费收益。


五、Squared ReLU 与其他新兴方案

5.1 Squared ReLU

\(f(x) = (\max(0, x))^2\)

梯度:

\(f'(x) = \begin{cases} 2x & \text{if } x > 0 \\ 0 & \text{if } x \leq 0 \end{cases}\)

稀疏性分析

  • ReLU 输出中约 50% 的值为零(负半轴被截断)

  • Squared ReLU 的非零值中,小值被进一步抑制(\(0.1^2 = 0.01\)),产生有效稀疏性——虽然非零但对输出贡献极小

  • 实测 Squared ReLU 的"\(\epsilon\)-稀疏率"(\(|f(x)| < \epsilon\) 的比例)比 ReLU 高 20-30%

class SquaredReLUFFN(nn.Module):
    def __init__(self, d_model, d_ff):
        self.w1 = nn.Linear(d_model, d_ff, bias=False)
        self.w2 = nn.Linear(d_ff, d_model, bias=False)

    def forward(self, x):
        h = F.relu(self.w1(x))
        return self.w2(h * h)

5.2 与 MoE 的联系

稀疏激活与 MoE 的显式专家路由之间存在深层联系: 参考:Primer: Searching for Efficient Transformers for Language Modeling (So et al., 2021) arxiv:2109.08668

5.3 其他探索方向


六、FFN 设计空间

6.1 宽度比 \(d_{ff}/d_{\text{model}}\)

6.2 跨层共享 FFN 参数

Universal Transformer (Dehghani et al., 2019) 提出所有层共享同一套参数(包括 FFN),通过重复执行实现"自适应计算"。 现代 LLM 几乎不共享 FFN:FFN 是模型容量的主要来源,共享等于削减模型容量。但在 MoE 中,shared expert 的思路部分复活了这一概念。

6.3 FFN 作为 Key-Value Memory

Geva et al. (2021) 的"Key-Value Memory"解释已在 1.2 节介绍。这一视角的工程启示:

def ffn_as_memory_lookup(x, W_key, W_value, activation):
    # W_key: [d_model, d_ff] — 每列是一个 "记忆键"
    # W_value: [d_ff, d_model] — 每行是一个 "记忆值"

    match_scores = activation(x @ W_key)   # [B, S, d_ff]
    # match_scores[i] = 输入与第 i 个记忆键的匹配度

    output = match_scores @ W_value        # [B, S, d_model]
    # output = 所有匹配记忆的加权求和
    return output

这一视角解释了几个观察:

  • 为什么 \(d_{ff}\) 要大\(d_{ff}\) 等于记忆条目的数量,更大 = 能存更多知识

  • 为什么 ReLU 的稀疏性有益:每个输入只匹配少量记忆条目,类似稀疏检索

  • 为什么 SwiGLU 更好:门控机制让匹配和值读取解耦,更像一个真正的 key-value 存储


七、Activation Memory 与工程考量

7.1 为什么 FFN 激活值占据显存大头

训练时需要保存前向传播的中间激活值(用于反向传播计算梯度)。每层的激活值显存: 对于 SwiGLU FFN,需要保存的激活值总量约为:

\(\text{FFN activations} = (d + 3 \times d_{ff}) \times B \times S \times \text{bytes\_per\_element}\)

以 LLaMA-7B 为例(\(d=4096, d_{ff}=11008, B=1, S=4096\), BF16):

\(= (4096 + 3 \times 11008) \times 1 \times 4096 \times 2 \approx 300 \text{ MB/layer}\)

32 层总计 ~9.6 GB 仅用于 FFN 激活值。

7.2 激活检查点策略

class CheckpointedSwiGLUFFN(nn.Module):
    def __init__(self, d_model, d_ff):
        self.w_gate = nn.Linear(d_model, d_ff, bias=False)
        self.w_up   = nn.Linear(d_model, d_ff, bias=False)
        self.w_down = nn.Linear(d_ff, d_model, bias=False)

    def _inner(self, x):
        return self.w_down(F.silu(self.w_gate(x)) * self.w_up(x))

    def forward(self, x):
        if self.training:
            return torch.utils.checkpoint.checkpoint(self._inner, x)
        return self._inner(x)

7.3 每层激活显存公式

完整的每层激活显存(BF16 训练,不含 Attention 的 softmax 矩阵):

\(M_{\text{act/layer}} = B \times S \times (10d + 2d_{\text{attn}} + 5d_{ff}) \times 2 \text{ bytes}\)

其中:

  • \(10d\):LayerNorm 输入/输出、残差连接等(约 5 个 \(d\) 维张量,每个需保存输入和输出)

  • \(2d_{\text{attn}}\):QKV 投影输出(GQA 下 KV 头数少于 Q)

  • \(5d_{ff}\):SwiGLU 三个线性层的输入/输出 + 两路中间结果

以 LLaMA-7B(\(d=4096, d_{ff}=11008, S=4096, B=4\))为例:\(M_{\text{act/layer}} \approx 4 \times 4096 \times (40960 + 8192 + 55040) \times 2 \approx 3.4\) GB/layer,32 层总计 ~109 GB。这就是为什么激活检查点是必需的。


八、追问延伸

Q1: 为什么 GPT-4 / Claude 等闭源模型不公开激活函数选择?

答:激活函数选择本身不构成核心壁垒(SwiGLU 已成公开共识),但具体的 \(d_{ff}\) 比值、是否使用混合激活(不同层不同函数)、以及与 MoE 的配合细节属于 recipe 调优的一部分,是竞争优势。

Q2: 能否不同层使用不同激活函数?

理论上可以,且有初步研究表明浅层和深层的最优激活函数可能不同(浅层偏好更稀疏的激活以做粗粒度过滤,深层偏好更平滑的激活以做精细组合)。但工程复杂性增加显著(需要为每层维护不同的 kernel),目前没有主流模型采用。

Q3: SwiGLU 的 Swish 中 \(\beta\) 应该设为多少?

几乎所有实现都固定 \(\beta=1\)(即 SiLU)。可学习 \(\beta\) 的实验(Ramachandran et al., 2017)表明最终学到的 \(\beta\) 值集中在 0.8-1.2 之间,收益极小。LLaMA、PaLM、Mistral 均使用 \(\beta=1\)

Q4: 未来 FFN 设计的方向?

  • 稀疏 FFN:只激活 \(d_{ff}\) 中的一小部分神经元,减少实际计算量(类似 MoE 但在神经元级别而非专家级别)

  • FFN-free 架构:部分研究探索用更大的 Attention 替代 FFN(如 Hyena),但目前效果不如标准 Transformer

  • 动态宽度 FFN:根据输入的"难度"自适应调整 \(d_{ff}\)——简单 token 用更窄的 FFN,复杂 token 用更宽的

Q5: SwiGLU 的梯度流有什么特殊性质?

SwiGLU 的反向传播涉及乘法规则:

\(\frac{\partial}{\partial x}[\text{Swish}(xW_g) \otimes xW_u] = \text{Swish}'(xW_g) \cdot W_g \cdot (xW_u) + \text{Swish}(xW_g) \cdot W_u\)

两条梯度路径——一条经过 gate,一条经过 up——提供了更丰富的梯度信号。即使 gate 路径的梯度接近零(gate 关闭),up 路径仍提供梯度,避免了类似 dying ReLU 的梯度消失问题。


参考文献

  1. Vaswani, A., et al. (2017). Attention Is All You Need. arxiv:1706.03762

  2. Hendrycks, D. & Gimpel, K. (2016). Gaussian Error Linear Units (GELUs). arxiv:1606.08415

  3. Ramachandran, P., Zoph, B., & Le, Q. V. (2017). Searching for Activation Functions. arxiv:1710.05941

  4. Dauphin, Y., et al. (2017). Language Modeling with Gated Convolutional Networks. arxiv:1612.08083

  5. Shazeer, N. (2020). GLU Variants Improve Transformer. arxiv:2002.05202

  6. Geva, M., et al. (2021). Transformer Feed-Forward Layers Are Key-Value Memories. arxiv:2012.14913

  7. Dai, D., et al. (2022). Knowledge Neurons in Pretrained Transformers. arxiv:2104.08696

  8. So, D., et al. (2021). Primer: Searching for Efficient Transformers for Language Modeling. arxiv:2109.08668

  9. Chowdhery, A., et al. (2022). PaLM: Scaling Language Modeling with Pathways. arxiv:2204.02311

  10. Touvron, H., et al. (2023). LLaMA: Open and Efficient Foundation Language Models. arxiv:2302.13971

  11. Touvron, H., et al. (2023). Llama 2: Open Foundation and Fine-Tuned Chat Models. arxiv:2307.09288

  12. Jiang, A. Q., et al. (2023). Mistral 7B. arxiv:2310.06825

  13. DeepSeek-AI. (2024). DeepSeek-V3 Technical Report. arxiv:2412.19437

  14. Dehghani, M., et al. (2019). Universal Transformers. arxiv:1807.03819


上级 · A. 基础理论