A from-scratch, granular Mixture-of-Experts GPT trainer built around the
Muon optimizer. The live source is the thin root train.py + the moegpt/
package (models, layers, MoE, optim, data, utils — timm-style).
Almost every architectural choice is a field on GPTConfig
(moegpt/models/config.py); the builders branch on the config rather than a
factory. The settled geometry this repo is built around:
- d_model 768, 12 layers (layer 0 dense + 11 MoE), 12 heads (head_dim 64), context window 512.
- 192 routed granular experts (SwiGLU hidden 256), top-4, plus 1 coarse shared expert (hidden 512) always active.
- Gated + RoPE attention, RMSNorm, tied embeddings.
- 1.34B total / 117M active params (~91% sparsity).
MoE routing uses capacity-factor dispatch with static shapes — every expert
is invoked every step, which keeps torch.compile(fullgraph=True) from
recompiling and gives Muon a constant set of same-shape matrices.
Requires Python 3 with torch, numpy, matplotlib, plus tiktoken,
pyarrow, and wandb.
Training reads tokenized memmap .bin streams (train.bin / val.bin,
uint16) from the directory in PERFGPT_DATA_DIR; train.py raises if they're
missing. On the B200 pods the data lives at /path/to/data/fineweb_edu_20b.
export PERFGPT_DATA_DIR=/path/to/data/fineweb_edu_20bpython train.pytorchrun --standalone --nproc_per_node=8 train.py--nproc_per_node should match the GPU count on the node. The global batch is
split across DDP ranks. Muon is DDP-only (each rank holds full unsharded
matrices — it will not work under FSDP sharding).
CLI flags (a launcher can fan out many jobs against one shared train.py
without editing the file):
| Flag | Meaning |
|---|---|
--lr |
Learning rate. With PERFGPT_OPTIMIZER=muon this is the Muon matrix LR (~50× Adam's; sweeps center on ~0.02). |
--global-batch-size |
Override the global batch size (split across ranks). |
--tpp |
Tokens-per-parameter budget: total tokens = active_params × tpp. Hard-stops training at the budget; tpp unset falls back to epoch-based training. |
Example — Muon run, 8 GPUs, fixed token budget:
PERFGPT_OPTIMIZER=muon \
torchrun --standalone --nproc_per_node=8 train.py --lr 0.02 --global-batch-size 128 --tpp 20| Variable | Default | Meaning |
|---|---|---|
PERFGPT_DATA_DIR |
— | Directory holding train.bin / val.bin (required). |
PERFGPT_OPTIMIZER |
adamw |
adamw or muon. |
PERFGPT_GRAD_ACCUM |
1 |
Split each per-rank batch into N micro-batches to cut activation memory. Global batch and TPP budget are unchanged. Use for large-batch cells that OOM on one node. |
PERFGPT_METRICS_DIR |
— | Where per-run JSONL metrics are written. |
PERFORMANT_GPT_PEAK_FLOPS |
31.2e12 |
Hardware peak FLOPS for MFU. Set to 2.25e15 on a B200 node. |
PERFGPT_SMOKE / PERFGPT_SMOKE_STEPS |
— / 5 |
Run a short smoke test of N steps. |
| Path | What |
|---|---|
train.py |
Training loop, DDP/torchrun entry, optimizer + data wiring, CLI/env parsing. |
moegpt/utils/perf.py |
PerfTracker — MoE-aware MFU + token-budget accounting (total_params vs active_params). |
moegpt/models/config.py |
GPTConfig — the dispatch table. |
moegpt/models/{gpt,hyper_gpt}.py |
GPTModel / HyperGPTModel (+ registry in registry.py). |
moegpt/moe/block.py |
MoEBlock: top-k router + capacity-factor static-shape dispatch. |
moegpt/layers/attention/{softmax,gated,rope,mla}.py |
Attention variants + build_attention. |
moegpt/optim/{muon,factory}.py |
Fused MuonAdamW + build_muon_optimizer. |
moegpt/moe/param_count.py |
Analytic param calculator (torch-free; matches torch exactly); use it to size new geometries. |