""" Config for the redesigned run: small custom vocab (not GPT-2's 50k) so embedding overhead doesn't dominate the parameter budget, sized to hit a genuine ~20:1 token:param ratio on a free-tier T4. Run `python model.py` after building this to confirm the exact param count before launching a long run. """ from dataclasses import dataclass @dataclass class ModelConfig: vocab_size: int = 8192 # custom BPE, trained on YOUR corpus (tokenizer_train.py) # -- NOT GPT-2's 50304. At small model sizes, a 50k vocab's # embedding table alone eats 60-75% of total params, leaving # almost nothing for actual transformer capacity. 8192 keeps # embedding overhead to ~17-20% of total. context_len: int = 256 # countdown prompts are short; halving context vs. the previous # 512 also halves the quadratic attention cost, for free d_model: int = 384 n_layer: int = 10 n_head: int = 6 n_kv_head: int = 2 # GQA d_ff: int = 1024 # SwiGLU inner dim rope_theta: float = 10000.0 dropout: float = 0.0 tie_embeddings: bool = True # -> this config lands at ~18.9M params, verified by model.py @dataclass class TrainConfig: data_dir: str = "data" train_bin: str = "data/train.bin" val_bin: str = "data/val.bin" # ---- the ratio that actually matters ---- target_tokens: int = 380_000_000 # ~20:1 tokens:params -- genuinely Chinchilla-optimal, # not a compromise like the 1.3:1 ratio last time # ---- optimization ---- micro_batch_size: int = 32 # smaller model + shorter context = bigger batch fits grad_accum_steps: int = 4 # effective batch = 32*4*256 = 32,768 tokens/step max_lr: float = 8e-4 # slightly higher than the 116M run's 6e-4 -- smaller # models generally tolerate a higher LR min_lr: float = 8e-5 warmup_steps: int = 300 weight_decay: float = 0.1 grad_clip: float = 1.0 beta1: float = 0.9 beta2: float = 0.95 precision: str = "fp16" # T4 = Turing, no bf16 tensor cores -- same reasoning as before ckpt_dir: str = "checkpoints" log_path: str = "logs/train_log.csv" save_every_steps: int = 250 eval_every_steps: int = 250 eval_iters: int = 50 log_every_steps: int = 20 seed: int = 1337