Download model/block.py from Vivid86/MiniTransformer-91M: direct link, hf CLI and curl.
- Browser
- Download file 4.74 kB
-
https://huggingface.co/Vivid86/MiniTransformer-91M/resolve/main/model/block.py
- Command line
-
hf download hf://Vivid86/MiniTransformer-91M/model/block.py
-
curl -L -o block.py https://huggingface.co/Vivid86/MiniTransformer-91M/resolve/main/model/block.py
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 | |