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