MiniTransformer-91M / model /feed_forward.py
Vivid86's picture
Upload folder using huggingface_hub
a2ec932 verified
Raw History Blame Contribute Delete
2.05 kB
"""
SwiGLU Feed-Forward Network
────────────────────────────
Replaces the original GELU/ReLU FFN with SwiGLU β€” the activation function
used in Llama, Mistral, PaLM, and most modern open-weight LLMs.
WHY SWIGLU OVER GELU/RELU:
Standard FFN: FFN(x) = max(0, xW1 + b1) W2 + b2 (ReLU)
SwiGLU FFN: FFN(x) = (SiLU(xW1) βŠ— xW3) W2 (SwiGLU)
The key difference is the *gating* term (xW3). Instead of a fixed
activation, SwiGLU uses a learned gate that modulates how much of the
activated signal passes through. This gives the FFN extra expressiveness
at the same parameter count.
SiLU (Sigmoid Linear Unit) = x Β· sigmoid(x)
Smooth, non-monotone; empirically outperforms GELU in this gated form.
PARAMETER COUNT:
Standard 4Γ— FFN: 2 matrices (d_model β†’ d_ff β†’ d_model)
SwiGLU FFN: 3 matrices (w1, w3: d_model β†’ d_ff; w2: d_ff β†’ d_model)
To keep FLOPs equal to a 4Γ— FFN, set d_ff β‰ˆ 8/3 Γ— d_model (see ModelConfig).
BIAS=FALSE:
All projections omit biases following Llama/Mistral practice β€” biases add
negligible capacity at scale and waste memory bandwidth.
"""
import torch
import torch.nn as nn
import torch.nn.functional as F
class SwiGLUFeedForward(nn.Module):
"""
SwiGLU Feed-Forward Network.
Computes: FFN(x) = W2(SiLU(W1 x) βŠ— W3 x)
Args:
d_model: Input/output dimension.
d_ff: Hidden (intermediate) dimension.
Use ModelConfig.d_ff (auto-sized to 8/3 Γ— d_model β†’ nearest 64).
"""
def __init__(self, d_model: int, d_ff: int):
super().__init__()
self.w1 = nn.Linear(d_model, d_ff, bias=False) # gate projection
self.w3 = nn.Linear(d_model, d_ff, bias=False) # value projection
self.w2 = nn.Linear(d_ff, d_model, bias=False) # output projection
def forward(self, x: torch.Tensor) -> torch.Tensor:
# SiLU(gate) βŠ— value, then project back down
return self.w2(F.silu(self.w1(x)) * self.w3(x))