Upload README.md with huggingface_hub
Browse files
README.md
ADDED
|
@@ -0,0 +1,88 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
---
|
| 2 |
+
license: mit
|
| 3 |
+
tags:
|
| 4 |
+
- text-generation
|
| 5 |
+
- from-scratch
|
| 6 |
+
- grpo
|
| 7 |
+
- reinforcement-learning
|
| 8 |
+
- arithmetic-reasoning
|
| 9 |
+
---
|
| 10 |
+
|
| 11 |
+
# tinyzero-countdown-19m
|
| 12 |
+
|
| 13 |
+
A ~18.9M-parameter decoder-only transformer, pretrained from scratch and
|
| 14 |
+
post-trained with GRPO (Group Relative Policy Optimization) to solve
|
| 15 |
+
Countdown-style arithmetic puzzles: given a set of numbers and a target,
|
| 16 |
+
find an equation using each number exactly once that reaches the target.
|
| 17 |
+
|
| 18 |
+
## Architecture
|
| 19 |
+
RoPE positional embeddings, RMSNorm, grouped-query attention (via
|
| 20 |
+
`F.scaled_dot_product_attention`), SwiGLU MLP, tied embeddings. Custom
|
| 21 |
+
8192-token BPE vocabulary trained on the pretraining corpus (not GPT-2's
|
| 22 |
+
tokenizer -- see rationale below).
|
| 23 |
+
|
| 24 |
+
- Parameters: ~18.88M (verified exactly, not estimated)
|
| 25 |
+
- Context length: 256
|
| 26 |
+
- Vocab size: 8192 (custom-trained BPE)
|
| 27 |
+
- d_model: 384, layers: 10, heads: 6 (2 KV heads, GQA)
|
| 28 |
+
|
| 29 |
+
## Training pipeline and what I learned building it
|
| 30 |
+
|
| 31 |
+
**Pretraining**: ~380M tokens on FineWeb-Edu + synthetic arithmetic
|
| 32 |
+
text, at a ~20:1 token:parameter ratio (Chinchilla-optimal). An earlier
|
| 33 |
+
attempt at 116M params / 150M tokens (1.3:1 ratio) showed the failure mode
|
| 34 |
+
directly: healthy train/val loss gap but weak generalization. This version
|
| 35 |
+
also fixes a subtler issue -- at small model scale, a standard 50k-token
|
| 36 |
+
vocabulary's embedding table dominates the parameter budget (60-75% of
|
| 37 |
+
total params); training a small custom vocab instead keeps embedding
|
| 38 |
+
overhead to ~17%, leaving actual capacity for reasoning.
|
| 39 |
+
|
| 40 |
+
**SFT**: an instruction-format fine-tune initially looked successful by
|
| 41 |
+
loss (train 1.02->0.34) but generation accuracy was 0% -- a real
|
| 42 |
+
loss/accuracy divergence caused by a response template that was mostly
|
| 43 |
+
easy-to-predict boilerplate, diluting the loss signal on the tokens that
|
| 44 |
+
actually mattered (the numbers/operators). Root-caused to a large,
|
| 45 |
+
un-bridged distribution shift between the pretraining corpus's format and
|
| 46 |
+
the instruction phrasing; fixed by skipping the instruction wrapper and
|
| 47 |
+
running GRPO directly on the pretrained checkpoint's native prompt format
|
| 48 |
+
instead.
|
| 49 |
+
|
| 50 |
+
**GRPO**: trained directly on the pretrained checkpoint, using the
|
| 51 |
+
verifier (exact equation checker) as a binary+partial-credit reward, group-
|
| 52 |
+
relative advantage normalization, PPO-style clipping, and a KL penalty
|
| 53 |
+
against a frozen reference to prevent collapse. Result: 31.6% ->
|
| 54 |
+
34.4% accuracy on a held-out 250-problem set (+2.8pp), with stable
|
| 55 |
+
KL throughout (no collapse). This is a modest, honestly-reported effect --
|
| 56 |
+
run at only ~500 steps on an 18.9M model, not a large or highly significant
|
| 57 |
+
result, and reported with that caveat intentionally.
|
| 58 |
+
|
| 59 |
+
## Usage
|
| 60 |
+
|
| 61 |
+
```python
|
| 62 |
+
import torch
|
| 63 |
+
from tokenizers import ByteLevelBPETokenizer
|
| 64 |
+
# adapt these imports to wherever you place model.py / config.py from this repo
|
| 65 |
+
from model import TinyTransformer
|
| 66 |
+
from config import ModelConfig
|
| 67 |
+
|
| 68 |
+
mcfg = ModelConfig()
|
| 69 |
+
model = TinyTransformer(mcfg)
|
| 70 |
+
ckpt = torch.load("pytorch_model.pt", map_location="cpu")
|
| 71 |
+
model.load_state_dict(ckpt["model_state_dict"])
|
| 72 |
+
model.eval()
|
| 73 |
+
|
| 74 |
+
tokenizer = ByteLevelBPETokenizer("vocab.json", "merges.txt")
|
| 75 |
+
prompt = "Numbers: [12, 45, 7, 3], Target: 88, Equation:"
|
| 76 |
+
ids = tokenizer.encode(prompt).ids
|
| 77 |
+
x = torch.tensor([ids])
|
| 78 |
+
out = model.generate(x, max_new_tokens=40, temperature=1.0, top_k=1)
|
| 79 |
+
print(tokenizer.decode(out[0, len(ids):].tolist()))
|
| 80 |
+
```
|
| 81 |
+
|
| 82 |
+
## Limitations
|
| 83 |
+
- Small model (~19M params) -- general text fluency is weak; this is
|
| 84 |
+
specialized for the countdown arithmetic task, not general-purpose use.
|
| 85 |
+
- Only handles the raw prompt format shown above; natural-language
|
| 86 |
+
instruction phrasing was found to significantly degrade output quality
|
| 87 |
+
(see training notes above) and was not used for the released checkpoint.
|
| 88 |
+
- Evaluated on synthetically generated countdown problems only.
|