simak31's picture
Upload config.py with huggingface_hub
03ea399 verified
Raw History Blame Contribute Delete
2.59 kB
"""
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