""" 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))