Vivid86's picture
Upload folder using huggingface_hub
a2ec932 verified
Raw History Blame Contribute Delete
4.74 kB
"""
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