FFN、SwiGLU 与 MoE
1. FFN
FFN(Feed-Forward Network)是 Transformer Block 中除 Attention 外的另一核心模块,对每个 token 独立进行非线性特征变换。
经典 FFN:
FFN(x)=W2σ(W1x) \mathrm{FFN}(x)=W_2\sigma(W_1x)FFN(x)=W2σ(W1x)
通常维度变化为:
dmodel→dff→dmodel d_{\text{model}} \rightarrow d_{\text{ff}} \rightarrow d_{\text{model}}dmodel→dff→dmodel
一般有:
dff>dmodel d_{\text{ff}}>d_{\text{model}}dff>dmodel
例如:
4096→16384→4096 4096\rightarrow16384\rightarrow40964096→16384→4096
即:
先升维扩展特征,再降回 hidden size。
不过这并不是强制要求,在一些现代模型,尤其是 MoE Expert 中,也可能有:
dff<dmodel d_{\text{ff}}<d_{\text{model}}dff<dmodel
PyTorch 实现
importtorchimporttorch.nnasnnclassFFN(nn.Module):def__init__(self,d_model,d_ff):super().__init__()self.up=nn.Linear(d_model,d_ff)self.activation=nn.GELU()self.down=nn.Linear(d_ff,d_model)defforward(self,x):x=self.up(x)x=self.activation(x)x=self.down(x)returnx2. SwiGLU
SwiGLU 可以看作一种带门控机制的 FFN 变体,广泛应用于 LLaMA、Qwen 等现代大模型。
核心公式:
SwiGLU(x)=Wdown[SiLU(Wgatex)⊙(Wupx)] SwiGLU(x)= W_{\text{down}} \left[ \mathrm{SiLU}(W_{\text{gate}}x) \odot (W_{\text{up}}x) \right]SwiGLU(x)=Wdown[SiLU(Wgatex)⊙(Wupx)]
相比普通 FFN,SwiGLU 增加了一条 Gate 分支:
g=SiLU(Wgatex) g=\mathrm{SiLU}(W_{\text{gate}}x)g=SiLU(Wgatex)
u=Wupx u=W_{\text{up}}xu=Wupx
然后进行逐元素相乘:
h=g⊙u h=g\odot uh=g⊙u
最后再映射回模型维度:
y=Wdownh y=W_{\text{down}}hy=Wdownh
其中包含三个线性投影:
Wgate,Wup,Wdown W_{\text{gate}}, \quad W_{\text{up}}, \quad W_{\text{down}}Wgate,Wup,Wdown
其中:
SiLU(Wgatex) \mathrm{SiLU}(W_{\text{gate}}x)SiLU(Wgatex)
可以理解为一个Gate,动态控制 (W_{\text{up}}x) 中哪些特征被保留。
因此可以简单记为:
SwiGLU = 带 SiLU 门控机制的 FFN。
PyTorch 实现
importtorchimporttorch.nnasnnimporttorch.nn.functionalasFclassSwiGLU(nn.Module):def__init__(self,d_model,d_ff):super().__init__()self.gate_proj=nn.Linear(d_model,d_ff)self.up_proj=nn.Linear(d_model,d_ff)self.down_proj=nn.Linear(d_ff,d_model)defforward(self,x):gate=F.silu(self.gate_proj(x))up=self.up_proj(x)x=gate*up x=self.down_proj(x)returnx3. MoE
MoE(Mixture of Experts,混合专家)可以理解为:
将 Transformer 中的一个 FFN 替换成多个 FFN Expert,再通过 Router 为每个 token 选择少量 Expert。
假设共有 (N) 个 Expert:
E1,E2,…,EN E_1,E_2,\ldots,E_NE1,E2,…,EN
Router 首先根据 token hidden state (x) 计算各个 Expert 的路由分数:
r=Wrx r=W_rxr=Wrx
经过 Softmax:
p=softmax(r) p=\operatorname{softmax}(r)p=softmax(r)
其中:
pi p_ipi
表示当前 token 被分配到第 (i) 个 Expert 的概率。
随后只选择概率最高的 Top-(K) 个 Expert:
E(x)=TopK(p) \mathcal E(x)=\operatorname{TopK}(p)E(x)=TopK(p)
最终输出:
y=∑i∈E(x)piEi(x) y= \sum_{i\in\mathcal E(x)} p_iE_i(x)y=i∈E(x)∑piEi(x)
例如共有 8 个 Expert,使用 Top-2 Routing,则每个 token 只经过其中两个 Expert,而不是计算全部 8 个 Expert。
MoE 的核心优势在于:
大量总参数+少量激活参数 \boxed{ \text{大量总参数} + \text{少量激活参数} }大量总参数+少量激活参数
即模型可以拥有大量不同的 Expert 参数,但一个 token 只需要计算其中少数几个 Expert。
MoE 中的 Expert 本质上通常就是 FFN,现代模型中也经常直接使用SwiGLU FFN作为 Expert。
简单 PyTorch 实现
下面以Top-2 MoE为例:
importtorchimporttorch.nnasnnimporttorch.nn.functionalasFclassExpert(nn.Module):def__init__(self,d_model,d_ff):super().__init__()self.gate_proj=nn.Linear(d_model,d_ff)self.up_proj=nn.Linear(d_model,d_ff)self.down_proj=nn.Linear(d_ff,d_model)defforward(self,x):gate=F.silu(self.gate_proj(x))up=self.up_proj(x)returnself.down_proj(gate*up)classMoE(nn.Module):def__init__(self,d_model,d_ff,num_experts=8,top_k=2):super().__init__()self.num_experts=num_experts self.top_k=top_k# Routerself.router=nn.Linear(d_model,num_experts)# Expertsself.experts=nn.ModuleList([Expert(d_model,d_ff)for_inrange(num_experts)])defforward(self,x):# x: [batch, seq_len, d_model]batch_size,seq_len,d_model=x.shape# 将 token 展平x=x.reshape(-1,d_model)# [num_tokens, d_model]# Router 计算每个 Expert 的概率router_logits=self.router(x)router_probs=F.softmax(router_logits,dim=-1)# 选择 Top-K Experttopk_probs,topk_indices=torch.topk(router_probs,k=self.top_k,dim=-1)# Top-K 权重重新归一化topk_probs=topk_probs/topk_probs.sum(dim=-1,keepdim=True)output=torch.zeros_like(x)# 将 token 分配给对应 Expertforexpert_id,expertinenumerate(self.experts):forkinrange(self.top_k):mask=(topk_indices[:,k]==expert_id)ifmask.any():expert_input=x[mask]expert_output=expert(expert_input)weight=topk_probs[mask,k].unsqueeze(-1)output[mask]+=(weight*expert_output)returnoutput.reshape(batch_size,seq_len,d_model)例如:
moe=MoE(d_model=4096,d_ff=1024,num_experts=8,top_k=2)- Hidden Size:4096
- 每个 Expert 的中间维度:1024
- 总 Expert 数:8
- 每个 token 激活:2 个 Expert