Skip to content
AI-grafen
EUniversityTransformer architecture· about 60 min· evolving, reviewed regularly· verified 2026-09-20· EN

The MLP block: GELU, SwiGLU

Be able to explain the feed-forward block and the gated variants in the Llama style.

Prerequisites

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: MLP(x)=W2 GELU(W1x)\text{MLP}(x) = W_2\,\text{GELU}(W_1 x) 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: SwiGLU(x)=(Swish(Wgatex)⊙Wupx)Wdown\text{SwiGLU}(x) = \big(\text{Swish}(W_{\text{gate}}x)\odot W_{\text{up}}x\big)W_{\text{down}}

where Swish(z)=z⋅σ(z)\text{Swish}(z) = z\cdot\sigma(z) (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 83≈2.67×\tfrac{8}{3}\approx 2.67\times, so that the parameter count becomes comparable:

VariantMatricesThe hidden dim (d = 4096)Parameters per block
GELU MLP216 3842⋅4096⋅16384≈1342\cdot 4096\cdot 16384 \approx 134 M
SwiGLU311 0083⋅4096⋅11008≈1353\cdot 4096\cdot 11008 \approx 135 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 83⋅4096=10922\tfrac{8}{3}\cdot 4096 = 10922 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

All the sources and licences