""" ModelConfig — central configuration dataclass for MiniTransformer. All model hyperparameters live here. Passing this single object to MiniTransformer guarantees that every sub-module is consistent. Use __post_init__ to: - Validate divisibility constraints - Auto-compute d_ff when left at 0 (SwiGLU sizing rule) """ import math from dataclasses import dataclass, field @dataclass class ModelConfig: # ── Vocabulary ────────────────────────────────────────────────────────── vocab_size: int = 8192 # Must match your trained tokenizer # ── Model Dimensions ──────────────────────────────────────────────────── d_model: int = 512 # Embedding / hidden dimension n_layers: int = 6 # Number of transformer blocks n_heads: int = 8 # Number of query heads n_kv_heads: int = 8 # Number of key/value heads # n_kv_heads == n_heads → standard MHA # n_kv_heads < n_heads → GQA (e.g. Llama 3) # n_kv_heads == 1 → MQA d_ff: int = 0 # FFN hidden dim. 0 = auto (SwiGLU rule: ⌈8/3·d_model⌉ → 64) # ── Sequence Length ────────────────────────────────────────────────────── context_len: int = 2048 # Maximum sequence length during training and inference # ── Regularisation ─────────────────────────────────────────────────────── dropout: float = 0.0 # Attention + FFN dropout. 0.0 is standard at scale. # ── Positional Encoding ─────────────────────────────────────────────────── rope_base: int = 10000 # RoPE theta. Increase for longer contexts (e.g. 500000 for Llama 3.1) # ── Normalisation ───────────────────────────────────────────────────────── norm_eps: float = 1e-5 # RMSNorm epsilon (prevents division by zero) # ── Output Head ─────────────────────────────────────────────────────────── tie_weights: bool = True # Share token embedding ↔ lm_head weights (GPT-2, Llama) def __post_init__(self): # ── Divisibility constraints ────────────────────────────────────────── assert self.d_model % self.n_heads == 0, ( f"d_model ({self.d_model}) must be divisible by n_heads ({self.n_heads}). " f"head_dim would be {self.d_model / self.n_heads:.1f}." ) assert self.n_heads % self.n_kv_heads == 0, ( f"n_heads ({self.n_heads}) must be divisible by n_kv_heads ({self.n_kv_heads}) " f"for Grouped-Query Attention." ) # ── Auto-compute d_ff (SwiGLU sizing) ──────────────────────────────── # SwiGLU needs ~2/3 the neurons of a standard 4× FFN to match FLOPs, # because it uses two weight matrices (gate + value) instead of one. # Formula: round(8/3 * d_model) to the nearest multiple of 64. if self.d_ff == 0: raw = (8 / 3) * self.d_model self.d_ff = int(math.ceil(raw / 64) * 64) @property def head_dim(self) -> int: """Dimension of each attention head.""" return self.d_model // self.n_heads @property def num_params_approx(self) -> int: """ Quick parameter estimate (embedding + n_layers × block). Useful for sanity-checking before instantiating the model. """ emb = self.vocab_size * self.d_model attn = self.d_model * (self.n_heads + 2 * self.n_kv_heads) * self.head_dim ff = 3 * self.d_model * self.d_ff # w1, w2, w3 for SwiGLU block = attn + ff total = emb + self.n_layers * block if self.tie_weights: total -= emb # head shares embedding weights return total def __repr__(self) -> str: return ( f"ModelConfig(" f"vocab={self.vocab_size}, " f"d_model={self.d_model}, " f"layers={self.n_layers}, " f"heads={self.n_heads}/{self.n_kv_heads} (Q/KV), " f"d_ff={self.d_ff}, " f"ctx={self.context_len}, " f"~{self.num_params_approx/1e6:.1f}M params" f")" )