""" TransformerBlock ───────────────── A single decoder layer combining: - RMSNorm (normalisation) - Pre-LayerNorm ordering (stability) - CausalSelfAttention (with RoPE + GQA + SDPA) - SwiGLU FeedForward (gated activation) ───────────────────────────────────────── WHY PRE-LAYERNORM (PRE-LN) vs POST-LN: ───────────────────────────────────────── Original "Attention is All You Need" used Post-LN: x = LayerNorm(x + Attention(x)) The norm is applied AFTER the residual addition. This means the residual signal must pass through the normaliser at every layer. With many layers, gradients at the early layers become vanishingly small (the norm acts as a bottleneck), causing instability or requiring extremely small learning rates. Pre-LN applies the norm BEFORE the sublayer: x = x + Attention(LayerNorm(x)) Now the residual path is clean — the gradient flows straight through the addition without being rescaled. This makes training far more stable at any depth and is the standard in every modern LLM (Llama, Mistral, GPT-NeoX, Falcon, Gemma). ───────────────────────────────────────── WHY RMSNORM vs LAYERNORM: ───────────────────────────────────────── LayerNorm computes: mean = mean(x) std = std(x) out = (x - mean) / std * gamma + beta RMSNorm simplifies to: rms = sqrt(mean(x²)) out = x / rms * gamma RMSNorm omits: 1. Mean subtraction (re-centering) — empirically unnecessary for LLMs 2. Beta parameter — saves a small amount of memory Result: ~15% faster than LayerNorm with equivalent training dynamics. Used in: Llama, Mistral, Falcon, Gemma, Qwen. """ import torch import torch.nn as nn from model.attention import CausalSelfAttention from model.feed_forward import SwiGLUFeedForward class RMSNorm(nn.Module): """ Root Mean Square Layer Normalisation. Args: dim: Feature dimension to normalise over. eps: Small constant for numerical stability (prevents division by zero when the RMS is extremely small). """ def __init__(self, dim: int, eps: float = 1e-5): super().__init__() self.eps = eps self.weight = nn.Parameter(torch.ones(dim)) # learnable scale (γ) def forward(self, x: torch.Tensor) -> torch.Tensor: # Compute RMS along the last dimension, normalise, then re-scale rms = x.pow(2).mean(-1, keepdim=True).add(self.eps).rsqrt() return x * rms * self.weight class TransformerBlock(nn.Module): """ One decoder layer (Pre-LN + RMSNorm + CausalSelfAttention + SwiGLU FFN). Args: config: ModelConfig instance. All hyperparameters are read from here. Forward signature: x — input tensor [B, T, d_model] cos, sin — RoPE tables from RotaryEmbedding.forward(T) past_kv — optional KV-cache tuple for inference Returns: (x, present_kv) x — output tensor [B, T, d_model] present_kv — updated (k, v) tuple to pass back during inference """ def __init__(self, config): super().__init__() self.attn_norm = RMSNorm(config.d_model, eps=config.norm_eps) self.ff_norm = RMSNorm(config.d_model, eps=config.norm_eps) self.attn = CausalSelfAttention( d_model = config.d_model, n_heads = config.n_heads, n_kv_heads = config.n_kv_heads, dropout = config.dropout, ) self.ff = SwiGLUFeedForward(config.d_model, config.d_ff) def forward( self, x: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor, past_kv: tuple | None = None, ) -> tuple[torch.Tensor, tuple]: # ── Attention sublayer (Pre-LN) ─────────────────────────────────────── # Normalise BEFORE attention so residual stream stays unscaled. attn_out, present_kv = self.attn(self.attn_norm(x), cos, sin, past_kv) x = x + attn_out # residual connection # ── Feed-Forward sublayer (Pre-LN) ──────────────────────────────────── x = x + self.ff(self.ff_norm(x)) # residual connection return x, present_kv