Skip to content
amanmprojectsPublic

About

32.5M story model trained from scratch in seven hours on an RTX 4060

Resources

Stars

0 stars

Watchers

0 watching

Forks

Repository files navigation

Sprout 🌱

A 32,514,560-parameter story-completion transformer trained from random initialization on a downloaded prefix of TinyStories. No pretrained model, hosted API, or paid service is used. This is a small creative-writing experiment, not an instruction-tuned chatbot.

Published results and weights

Completed seven hours of training: 1.659 billion tokens processed, best monitored validation loss 1.4100. See FINDINGS.md for evaluation limitations. Public run artifacts are in results/.

Try the released checkpoint

uv venv --python 3.12 .venv312
uv pip install --python .venv312/bin/python torch==2.8.0 tokenizers==0.23.2 numpy huggingface_hub
.venv312/bin/python -c "from huggingface_hub import snapshot_download; snapshot_download('amanm10000/sprout', local_dir='release')"
.venv312/bin/python sample.py --checkpoint release/best.pt --tokenizer release/tokenizer.json \
  --prompt "Once upon a time, there was a little fox who lived in the forest."

The inference checkpoint omits optimizer state and is not resumable. It uses the custom model.py implementation, not Transformers AutoModel. The local full training checkpoints remain unchanged.

Model and data

  • 8 blocks, width 512, 8 attention heads; RoPE, RMSNorm, SwiGLU, untied embedding/output weights.
  • 8,192-token byte-level BPE trained on 50,000 training stories; context length 512.
  • bf16 autocast, PyTorch scaled-dot-product/flash attention, fused AdamW, gradient clipping, warmup and time-based cosine LR decay.
  • Batch size selected by measured throughput while retaining at least 700 MiB of free GPU memory; gradient accumulation targets approximately 16k tokens/update.
  • Source: https://huggingface.co/datasets/roneneldan/TinyStories (CDLA-Sharing-1.0 dataset; review upstream terms before redistributing).
  • Network throughput was slow, so training uses only complete stories from the downloaded prefix, not the entire dataset. data/manifest.json records the actual story/token counts. Validation comes from the independent official validation file.
  • Training text is synthetic, may contain biases/errors, and does not guarantee safe or factually correct outputs. Validation is a small fixed sample used for monitoring, not a comprehensive quality evaluation.

Environment

Use .venv312/bin/python (Python 3.12, PyTorch 2.8.0 CUDA 12.8). The earlier .venv setup was incomplete and is not used. requirements.lock.txt records installed package versions. Nothing in the system Python environment was modified.

cd /home/aman/Work/tries/2026-09-17-model/sprout
.venv312/bin/python -m unittest test_model.py -v
.venv312/bin/python train.py --smoke --out runs/new-smoke-test

The CPU smoke test overfits one sentence, checks checkpoint reload, and generates the sentence. It is a correctness test, not evidence of generalization.

Train

.venv312/bin/python prepare_data.py
.venv312/bin/python train.py --hours 7 --out runs/sprout

The seven-hour wall-clock budget starts after benchmarking and includes training, periodic validation, checkpointing, and sample generation. Final validation/checkpointing can add a short tail. Machine suspend pauses computation; keep the laptop plugged in, awake, and well ventilated. Other GPU processes are not terminated or reconfigured.

Checkpoints are written atomically roughly every 15 minutes:

  • runs/sprout/latest.pt: latest resumable state
  • runs/sprout/best.pt: best measured validation checkpoint
  • runs/sprout/diary.md: fixed-prompt stories showing learning over time
  • runs/sprout/metrics.jsonl: training loss, throughput, LR, token counts
  • runs/sprout/status.json: running/completed/stopped/failed state
  • runs/sprout/config.json: model, timing, and hardware configuration
  • runs/sprout/benchmark.json: tested batch sizes and throughput

Monitor / stop

When launched as a user service named sprout-training:

systemctl --user status sprout-training
journalctl --user -u sprout-training -f
systemctl --user stop sprout-training

SIGTERM requests a graceful checkpoint after the current optimizer step. A crash/OOM preserves the previous complete checkpoint rather than saving a partially updated model. The service is not automatically restarted and does not survive a reboot; closing this chat alone does not stop it. System logout may stop user services depending on the machine's user-manager configuration.

Resume within the original total seven-hour compute budget (not seven additional hours):

.venv312/bin/python train.py --hours 7 --out runs/sprout --resume runs/sprout/latest.pt

Generate

CPU sampling is the default so you can sample while training without competing for GPU memory:

.venv312/bin/python sample.py --checkpoint runs/sprout/latest.pt \
  --prompt "Once upon a time, a tiny robot found a seed."

After training, use --checkpoint runs/sprout/best.pt --device cuda for faster sampling. Other controls: --tokens 220 --temperature 0.8 --top-k 40 --seed 42. Only load checkpoints you trust.

About

32.5M story model trained from scratch in seven hours on an RTX 4060

Resources

Stars

0 stars

Watchers

0 watching

Forks

Releases

Packages

Contributors

Languages