The MLP block: GELU, SwiGLU
Be able to explain the feed-forward block and the gated variants in the Llama style.
Prerequisites
- DActivation functionsrequired
- DTransformers — the architecturerequired
Intuition
Every transformer block has two parts: attention (which moves information between positions) and the MLP (which processes each position on its own).
The classic MLP block: with an expansion of 4× in the middle: 4096 → 16384 → 4096.
The MLP accounts for roughly two thirds of the parameters in a transformer — it is where most of the «knowledge» sits, while the attention handles the routing.
Formal
Gated Linear Units (Shazeer 2020) add a gate:
where (also called SiLU). The gate lets the network multiplicatively damp or let through each channel depending on the input — a form of content-dependent filtering that an ordinary MLP cannot express as easily.
Three matrices instead of two means more parameters at the same width. So the expansion factor is adjusted from 4× to about , so that the parameter count becomes comparable:
| Variant | Matrices | The hidden dim (d = 4096) | Parameters per block |
|---|---|---|---|
| GELU MLP | 2 | 16 384 | M |
| SwiGLU | 3 | 11 008 | M |
At a comparable parameter count SwiGLU consistently performs somewhat better — which is why it is used in Llama, Mistral, Qwen and PaLM.
Llama's odd hidden dimension of 11 008 is precisely rounded up to a multiple of 256 (for efficient matrix multiplication).
Code
import torch, torch.nn as nn, torch.nn.functional as F
class SwiGLU(nn.Module):
def __init__(self, d, hidden=None, multiple=256):
super().__init__()
hidden = hidden or int(((8 / 3) * d + multiple - 1) // multiple * multiple)
self.gate = nn.Linear(d, hidden, bias=False)
self.up = nn.Linear(d, hidden, bias=False)
self.down = nn.Linear(hidden, d, bias=False)
self.hidden = hidden
def forward(self, x):
return self.down(F.silu(self.gate(x)) * self.up(x))
class GeluMLP(nn.Module):
def __init__(self, d, hidden=None):
super().__init__()
hidden = hidden or 4 * d
self.up = nn.Linear(d, hidden, bias=False)
self.down = nn.Linear(hidden, d, bias=False)
def forward(self, x):
return self.down(F.gelu(self.up(x)))
for name, m in (("GELU MLP", GeluMLP(4096)), ("SwiGLU", SwiGLU(4096))):
print(f"{name:9s} parameters {sum(p.numel() for p in m.parameters()):,}")
# GELU MLP parameters 134,217,728
# SwiGLU parameters 135,266,304 ← comparable, thanks to the 8/3 adjustment
Mastery means
- Describes the feed-forward block and its dimensions
- Explains SwiGLU and why the expansion factor is adjusted
- Counts the parameters in the block
Sign in to do the exercises and build your mastery up.
Sources
- arXiv — GLU Variants Improve Transformer — arXiv (open access; licence per article)
- arXiv — LLaMA: Open and Efficient Foundation Language Models — arXiv (open access; licence per article)