Skip to content

Gated MLPs: SwiGLU, GeGLU

ModuleL7.2 · build · Python · Pass 5 · 2 h, plus your graded tests (rung R5)
You buildpython/tinyllm/modern/mlp.py: GatedMLP and llama_ffn_dim; and your own oracle tests in python/tests/l7-2-mlp/
Contractcourse/contracts/py/tinyllm/modern/mlp.pyi
Testscourse/tests/L7.2/test_mlp.py (what they check: section 4), golden values from transformers 5.19.0 LlamaMLP and GemmaMLP in course/fixtures/L7.2/gated_mlp_hf.npz (course/oracle/L7.2/gated_mlp_hf.py); your tests are graded by mutation, threshold 0.80 with every pitfall fault required
NeedsL0.4 Linear, Module · L0.2 F.silu, F.gelu · L0.1 Tensor · M06.3 PCG32 (default initialization) · reading: M01.3 (silu and the tanh GELU) (or --ref-deps)
Used byL7.8 every expert is a gated MLP · L7.9 every Llama MLP · later: L9.6 the C tl_silu_mul_f32
MilestoneMS-L7 (SmolLM2-135M logits match Hugging Face)
Optional depthShazeer, “GLU Variants Improve Transformer” (2020); Dauphin et al., “Language Modeling with Gated Convolutional Networks” (2017), section 3; Touvron et al., “LLaMA” (2023), section 2.2
  • A gated MLP computes two projections of the input and multiplies one by the activation of the other: down(act(gate(x))⊙up(x))\text{down}(\text{act}(\text{gate}(x)) \odot \text{up}(x)) (test_hand_example).
  • The activation belongs on the gate: a closed gate shuts its unit off whatever the other projection says (test_closed_gate_blocks_the_unit).
  • Three matrices at width 83d\frac{8}{3} d cost what two at 4d4d cost, which is where Llama’s odd widths come from (test_parameter_count_matches_a_4d_mlp, test_llama_ffn_dim).
  • Gemma’s GeGLU uses the tanh approximation of GELU, not the erf form (test_golden_hf_geglu).
Terminal window
ol start L7.2 # stubs mlp.py; prints your test path and rung (R5)
ol tests L7.2 # the course tests
# write your oracle tests in python/tests/l7-2-mlp/, then:
ol check L7.2 # course tests and the mutation grade of your tests
ol diff L7.2 # after passing: your code against the reference

Your 2017 transformer’s MLP is Linear(d, 4d), ReLU or GELU, Linear(4d, d). Open SmolLM2-135M’s checkpoint and the MLP has three matrices per layer, gate_proj, up_proj, and down_proj, of width 1536 for a model width of 576: not 4⋅576=23044 \cdot 576 = 2304. Your MLP cannot load these weights, and if you guess how the three combine you get logits that look plausible and are wrong. Every model in the Llama family, Mistral, Qwen, Gemma, and the experts of every MoE (L7.8) use this gated form. This module builds it and explains the width.

SymbolMeaningType / shape
xxone token’s hidden vectorfloat32[..., d]
dd, ffmodel width, hidden width d_ffint
WgW_g, WuW_ugate and up weightsfloat32[f, d]
WdW_ddown weightfloat32[d, f]
σ(z)\sigma(z)logistic sigmoid 1/(1+e−z)1 / (1 + e^{-z})elementwise
silu⁡(z)\operatorname{silu}(z)zσ(z)z \sigma(z)elementwise
gelu⁡tanh⁡(z)\operatorname{gelu}_{\tanh}(z)z2(1+tanh⁡(2/π (z+0.044715z3)))\frac{z}{2}\left(1 + \tanh\left(\sqrt{2/\pi}\,(z + 0.044715 z^3)\right)\right)elementwise
⊙\odotelementwise product

The plain MLP is W2 ϕ(W1x)W_2\, \phi(W_1 x): one hidden vector, one nonlinearity per unit. A gated linear unit (Dauphin et al.) computes two hidden vectors from the same input and lets one control the other elementwise: (Wgx)⊙(Wux)(W_g x) \odot (W_u x) passed through a nonlinearity on the gate side.

GatedMLP⁡(x)=Wd(act⁡(Wgx)⊙Wux).\operatorname{GatedMLP}(x) = W_d \left( \operatorname{act}(W_g x) \odot W_u x \right).

Unit ii outputs act⁡(gi)⋅ui\operatorname{act}(g_i) \cdot u_i. When gig_i is very negative, silu⁡(gi)≈0\operatorname{silu}(g_i) \approx 0 and the unit is closed regardless of uiu_i; when gig_i is large, silu⁡(gi)≈gi\operatorname{silu}(g_i) \approx g_i and the unit passes giuig_i u_i, a product of two linear functions of xx. The network can therefore form products of input features, which a single ReLU layer cannot. Shazeer compared the activations: SwiGLU (act⁡=silu⁡\operatorname{act} = \operatorname{silu}, Llama and SmolLM2) and GeGLU (act⁡=gelu⁡tanh⁡\operatorname{act} = \operatorname{gelu}_{\tanh}, Gemma) both reach lower perplexity than ReLU or GELU MLPs at equal parameters. Neither has a closed-form reason; it is an empirical result that every model since has kept.

F.silu and F.gelu (L0.2, from M01.3’s derivatives) and the products give the backward through autograd. The gate receives the upstream gradient times ui⋅act⁡′(gi)u_i \cdot \operatorname{act}'(g_i), the up projection receives it times act⁡(gi)\operatorname{act}(g_i): both projections learn, and a closed gate also stops the gradient to its up row.

Three f×df \times d matrices hold 3df3 d f numbers. To compare fairly with the plain MLP’s 2⋅d⋅4d=8d22 \cdot d \cdot 4d = 8 d^2, set f=83df = \frac{8}{3} d. Meta’s code computes h=⌊23⋅4d⌋h = \lfloor \frac{2}{3} \cdot 4d \rfloor, optionally scales it by ffn_dim_multiplier, and rounds up to a multiple of multiple_of (256 or 1024) so the matrix products tile well on hardware.

d=f=1d = f = 1, Wg=1W_g = 1, Wu=2W_u = 2, Wd=3W_d = 3, x=1x = 1.

  1. Gate: g=1g = 1. Up: u=2u = 2.
  2. SwiGLU: σ(1)=1/(1+e−1)=0.731059\sigma(1) = 1 / (1 + e^{-1}) = 0.731059, so silu⁡(1)=0.731059\operatorname{silu}(1) = 0.731059; times uu: 1.4621171.462117; times WdW_d: 4.3863524.386352.
  3. GeGLU: 2/π⋅1.044715=0.833562\sqrt{2/\pi} \cdot 1.044715 = 0.833562, tanh⁡=0.682384\tanh = 0.682384, gelu⁡tanh⁡(1)=0.841192\operatorname{gelu}_{\tanh}(1) = 0.841192; output 3⋅2⋅0.841192=5.0471523 \cdot 2 \cdot 0.841192 = 5.047152.

Sizes: Llama-2 7B has d=4096d = 4096: ⌊2⋅16384/3⌋=10922\lfloor 2 \cdot 16384 / 3 \rfloor = 10922, rounded up to a multiple of 256 is 43⋅256=1100843 \cdot 256 = 11008. SmolLM2: ⌊2⋅2304/3⌋=1536\lfloor 2 \cdot 2304 / 3 \rfloor = 1536, already a multiple of 256.

These are test_hand_example and test_llama_ffn_dim.

class GatedMLP(Module):
def __init__(self, d, d_ff, act: Literal["silu", "gelu_tanh"] = "silu", bias=False, rng=None): ...
def forward(self, x: Tensor) -> Tensor: ... # [..., d] -> [..., d]
def llama_ffn_dim(d, multiple_of=256, ffn_dim_multiplier=None) -> int: ...
TestKINDChecksWhy it matters downstream
test_hand_exampleunitsection 3 for both activationsyou and the test agree on the formula
test_golden_hf_swiglugoldenoutput and gradients against LlamaMLPSmolLM2’s MLPs in L7.9
test_golden_hf_swiglu_biasgoldenmlp_bias checkpointsbiased variants load
test_golden_hf_geglugoldenGemmaMLP with gelu_pytorch_tanhGemma-style checkpoints
test_gradcheck_every_parametergradcheckfloat64, both activations, with biasesevery matrix trains
test_parameter_names_order_and_shapesunitHF names, order, shapes, actthe safetensors keys
test_rng_draw_orderunitgate, up, down drawn from one rng in orderone seed fixes the weights everywhere
test_closed_gate_blocks_the_unitpropertya negative gate shuts the unitthe activation is on the gate
test_parameter_count_matches_a_4d_mlpproperty3d⋅83d=8d23 d \cdot \frac{8}{3} d = 8 d^2equal-parameter comparisons
test_llama_ffn_dimunit11008, 14336, 8192, 1536configs without intermediate_size
test_validationboundaryunknown activation, zero widthconfig typos fail loudly

Your oracle is the formula in numpy float64 with the module’s own weights from state_dict(): compute gate, up, the activation, the product, and down, and compare with forward for both activations, with and without biases. Add the hand example, the parameter names, your own central differences for xx and every parameter, and the published llama_ffn_dim values. Import only tinyllm.modern.mlp, tinyllm.autograd.tensor, and tinyllm.autograd.functional. ol check L7.2 requires 0.80 with every pitfall fault killed.

PitfallSymptomCaught by
1. activation on the up projectionloads fine, every logit wrongtest_golden_hf_swiglu, test_closed_gate_blocks_the_unit (mutant s01)
2. exact GELU for GeGLUGemma outputs off by about 1e-4 per unit, growing with depthtest_golden_hf_geglu (mutant s02)
3. rounding d_ff down10752 instead of 11008; the checkpoint shapes do not matchtest_llama_ffn_dim (mutant s03)
4. activating the product, act⁡(g⊙u)\operatorname{act}(g \odot u)a different function with the same parameterstest_hand_example (mutant s04)
registering up before gatestate_dict order and draws differ from Hugging Facetest_parameter_names_order_and_shapes (mutant s05)
dropping the down biasbiased checkpoints load and are off by the biastest_golden_hf_swiglu_bias (mutant s06)
ignoring ffn_dim_multiplierLlama-3 widths wrongtest_llama_ffn_dim (mutant s07)
the up projection as a constantup never trainstest_gradcheck_every_parameter (mutant s08)
a fresh rng for the last layertwo seeds give the same down weightstest_rng_draw_order (mutant s09)
DirectionModuleHow it uses this
BackL0.4three Linear layers
BackL0.2F.silu, F.gelu(approximate="tanh") with their VJPs
BackL0.1Tensor products
BackM06.3PCG32 initializes the layers when no rng is given
ForwardL7.8each expert of a mixture of experts is a GatedMLP
ForwardL7.9mlp of every Llama decoder layer
ForwardL9.6tl_silu_mul_f32 fuses silu⁡(g)⊙u\operatorname{silu}(g) \odot u in C
Your pieceProduction equivalentWhat it addsWhere to look
GatedMLPHF LlamaMLP, GemmaMLPthe same three matrices; tensor-parallel splits shard gate and up by rows and down by columnstransformers/models/llama/modeling_llama.py
GatedMLPvLLM MergedColumnParallelLinear + SiluAndMulgate and up fused into one [2f,d][2f, d] matrix, the activation fused with the productvLLM model_executor/layers/activation.py
llama_ffn_dimMeta’s FeedForwardthe reference rulellama/model.py in Meta’s llama repository