diff --git a/level_0_baseline/README.md b/level_0_baseline/README.md index cacde29..1b9403d 100644 --- a/level_0_baseline/README.md +++ b/level_0_baseline/README.md @@ -1,76 +1,132 @@ -# Level 0 Baseline +# Realistic isolated Level 0 nanoGPT baseline -A deliberately self-contained nanoGPT baseline. It does not import the repository's existing experiment framework or WW-PGD code. +This subtree is intentionally independent of the repository's WW-PGD experiment framework. It provides a clean AdamW baseline on natural-language data while retaining deterministic WeightWatcher alpha measurements. -## Scope +## What Level 0 means here -- one transformer block, one ordinary Q/K/V attention head, width 64, context 256 -- byte-level next-token language modeling on a fixed FineWeb-Edu subset -- AdamW or Muon with a global warmup/cosine schedule -- Muon applies only to hidden 2-D matrices; AdamW handles embeddings, tied LM head, LayerNorm parameters, and other non-matrix parameters -- deterministic seeds for initialization and sampled training windows -- immutable train, validation, and test splits -- CSV logging of loss, next-token accuracy, perplexity, validation/test generalization gaps, gradient norm, weight norm, tokens, and elapsed time -- optional checkpoint-time WeightWatcher layer analysis -- single-seed and multi-seed notebooks; multi-seed plots use mean ± one standard deviation shaded bands +The default MacBook preset is designed for an Apple M2 Pro with 16 GB unified memory: -## Install +- GPT-2 BPE tokenization (`tiktoken:gpt2`), padded model vocabulary 50,304 +- four transformer blocks, four attention heads, width 256 +- context length 256 BPE tokens +- approximately 16.1 million trainable parameters +- AdamW with standard matrix/no-matrix weight-decay groups +- 200-step linear warmup followed by cosine decay +- peak learning rate `6e-4`, minimum learning rate `6e-5` +- weight decay `0.1`, betas `(0.9, 0.95)`, gradient clipping `1.0` +- batch size 8 with four gradient-accumulation steps +- 5,000 optimizer steps, or 40.96 million processed training tokens +- fixed train and validation probes with independent RNG streams +- test evaluation only at the final and validation-selected checkpoints +- deterministic non-randomized WeightWatcher alpha measurements every 500 steps + +This replaces the obsolete one-block, width-64, raw-byte experiment. Old `/tmp/nanogpt-level0/data` byte files are rejected by the new trainer. + +## Conda installation ```bash -cd level_0_baseline -python -m venv .venv -source .venv/bin/activate -pip install -e '.[data,analysis,test]' +conda activate ww_prod310 +cd ~/Desktop/work/nanoGPT/nanogpt-experiments/level_0_baseline +python -m pip install -e '.[data,analysis,test]' +python -m pip check ``` ## Paths -Defaults are under `/tmp/nanogpt-level0`. Override them without editing code: +All large artifacts remain under `/tmp` by default: ```bash -export NANOGPT_LEVEL0_DATA_ROOT=/tmp/my-level0/data -export NANOGPT_LEVEL0_RESULTS_ROOT=/tmp/my-level0/results -export NANOGPT_LEVEL0_CACHE_ROOT=/tmp/my-level0/cache +export NANOGPT_LEVEL0_ROOT=/tmp/nanogpt-level0-bpe +export NANOGPT_LEVEL0_DATA_ROOT=$NANOGPT_LEVEL0_ROOT/data +export NANOGPT_LEVEL0_RESULTS_ROOT=$NANOGPT_LEVEL0_ROOT/results +export NANOGPT_LEVEL0_CACHE_ROOT=$NANOGPT_LEVEL0_ROOT/cache ``` -## Prepare the real corpus +## Prepare the pinned FineWeb-Edu corpus -This prepares fixed 50 MB training, 2 MB validation, and 2 MB test byte-token splits from streamed FineWeb-Edu: +The default preparation writes 20 million training tokens and one million tokens each for validation and test. Splits are fixed, atomic, and document-disjoint at boundaries. ```bash -level0-prepare-data --dataset fineweb-edu +./scripts/prepare_data.sh 2>&1 | tee /tmp/level0-bpe-prepare.log ``` -To monitor the streamed download and preparation, enable heartbeat logging: +Equivalent direct command: ```bash level0-prepare-data \ --dataset fineweb-edu \ + --output-dir /tmp/nanogpt-level0-bpe/data \ + --train-tokens 20000000 \ + --val-tokens 1000000 \ + --test-tokens 1000000 \ + --tokenizer gpt2 \ + --model-vocab-size 50304 \ --verbose \ --log-interval-seconds 10 ``` -Verbose output reports documents processed, bytes collected, completion percentage, elapsed time, average throughput, estimated time remaining, and how long the stream has produced no new bytes. The heartbeat continues while the streaming iterator is blocked, making a network or dataset stall visible. +The progress heartbeat reports dataset resolution, current split, documents, tokens, throughput, ETA, and time since the last new tokens arrived. + +## Run AdamW seed 1337 + +```bash +./scripts/run_one.sh adamw 1337 mps \ + 2>&1 | tee /tmp/level0-bpe-adamw-seed1337.log +``` + +The script resumes automatically when `checkpoint_latest.pt` exists. To deliberately discard a prior run: -## Run one seed +```bash +NANOGPT_LEVEL0_OVERWRITE=1 ./scripts/run_one.sh adamw 1337 mps +``` + +For a bounded CPU smoke test: ```bash -./scripts/run_one.sh adamw 1337 -./scripts/run_one.sh muon 1337 +NANOGPT_LEVEL0_MAX_STEPS=2 \ +NANOGPT_LEVEL0_BATCH_SIZE=2 \ +NANOGPT_LEVEL0_GRAD_ACCUM_STEPS=1 \ +NANOGPT_LEVEL0_EVAL_INTERVAL=1 \ +NANOGPT_LEVEL0_DISABLE_WEIGHTWATCHER=1 \ +NANOGPT_LEVEL0_OVERWRITE=1 \ +./scripts/run_one.sh adamw 1337 cpu ``` -## Run multiple seeds +## Output contract + +Each run writes: + +- `manifest.json`: exact model, optimizer, data identity, fixed-probe hashes, and protocol +- `metrics.csv`: train/validation loss, perplexity, bits per token, accuracy, gap, LR, gradient norm, weight norm, and throughput +- `checkpoint_latest.pt`: resumable training state +- `checkpoint_best.pt`: validation-selected model +- `checkpoint_final.pt`: final model +- `final_metrics.json`: final-checkpoint test metrics +- `selected_checkpoint_metrics.json`: validation-selected checkpoint test metrics +- `weightwatcher_step_*.csv`: per-matrix alpha, D, xmin, and tail metadata when available +- `run_complete.json`: transactional completion marker +- `train.log`: persistent progress log + +Test data is not evaluated during training. It is touched only after optimization completes, once for the final checkpoint and once for the validation-selected checkpoint. + +## Plot one run ```bash -NANOGPT_LEVEL0_SEEDS=1337,2027,4099 ./scripts/run_multiseed.sh +export NANOGPT_LEVEL0_RESULTS_ROOT=/tmp/nanogpt-level0-bpe/results +export NANOGPT_LEVEL0_NOTEBOOK_OPTIMIZER=adamw +export NANOGPT_LEVEL0_NOTEBOOK_SEED=1337 +jupyter lab notebooks/01_single_seed.ipynb ``` -The notebooks read `NANOGPT_LEVEL0_RESULTS_ROOT`. Select the single-seed run with `NANOGPT_LEVEL0_NOTEBOOK_OPTIMIZER` and `NANOGPT_LEVEL0_NOTEBOOK_SEED`. +The notebook plots loss, perplexity, exact next-BPE-token accuracy, bits per token, optimization diagnostics, and WeightWatcher alpha trajectories. It also saves PNG files under the run's `plots/` directory. -For a bounded infrastructure smoke test, override the run length and batch size: +## Multiple seeds ```bash -NANOGPT_LEVEL0_MAX_STEPS=2 NANOGPT_LEVEL0_BATCH_SIZE=2 NANOGPT_LEVEL0_EVAL_INTERVAL=1 ./scripts/run_one.sh adamw 1337 +NANOGPT_LEVEL0_SEEDS=1337,2027,4099 \ +NANOGPT_LEVEL0_OPTIMIZERS=adamw \ +NANOGPT_LEVEL0_DEVICE=mps \ +./scripts/run_multiseed.sh ``` -Next-token error is `1 - next-token accuracy`; the notebooks derive and plot it explicitly. +Then run `notebooks/02_multiseed.ipynb` for mean and standard-deviation bands. diff --git a/level_0_baseline/configs/level0.yaml b/level_0_baseline/configs/level0.yaml index 5d385c9..0b2e226 100644 --- a/level_0_baseline/configs/level0.yaml +++ b/level_0_baseline/configs/level0.yaml @@ -1,33 +1,40 @@ model: - vocab_size: 256 + # MacBook-scale GPT baseline: large enough for meaningful BPE language + # modeling and WeightWatcher spectra, while remaining practical on an + # M2 Pro with 16 GB unified memory. + vocab_size: 50304 block_size: 256 - n_layer: 1 - n_head: 1 - n_embd: 64 + n_layer: 4 + n_head: 4 + n_embd: 256 dropout: 0.0 bias: false + training: - batch_size: 16 - grad_accum_steps: 1 - max_steps: 2000 - eval_interval: 50 + batch_size: 8 + grad_accum_steps: 4 + max_steps: 5000 + eval_interval: 100 eval_batches: 20 - checkpoint_interval: 250 + log_interval: 10 + checkpoint_interval: 500 learning_rate: 0.0006 - muon_learning_rate: 0.02 - muon_aux_adamw_learning_rate: 0.0006 min_lr: 0.00006 - warmup_steps: 100 + warmup_steps: 200 weight_decay: 0.1 beta1: 0.9 beta2: 0.95 + epsilon: 1.0e-8 grad_clip: 1.0 optimizer: adamw + muon_learning_rate: 0.02 + muon_aux_adamw_learning_rate: 0.0006 muon_momentum: 0.95 muon_nesterov: true seed: 1337 compile: false + analysis: weightwatcher: true - weightwatcher_interval: 250 - randomize: true + weightwatcher_interval: 500 + randomize: false diff --git a/level_0_baseline/notebooks/01_single_seed.ipynb b/level_0_baseline/notebooks/01_single_seed.ipynb index d887603..108564d 100644 --- a/level_0_baseline/notebooks/01_single_seed.ipynb +++ b/level_0_baseline/notebooks/01_single_seed.ipynb @@ -1,10 +1,235 @@ { "cells": [ - {"cell_type":"markdown","metadata":{},"source":["# Level 0 single-seed diagnostics\n"]}, - {"cell_type":"code","execution_count":null,"metadata":{},"outputs":[],"source":["import os\n","from pathlib import Path\n","import pandas as pd\n","import matplotlib.pyplot as plt\n","ROOT=Path(os.getenv('NANOGPT_LEVEL0_RESULTS_ROOT','/tmp/nanogpt-level0/results'))\n","OPT=os.getenv('NANOGPT_LEVEL0_NOTEBOOK_OPTIMIZER','adamw')\n","SEED=int(os.getenv('NANOGPT_LEVEL0_NOTEBOOK_SEED','1337'))\n","RUN=ROOT/f'{OPT}_seed_{SEED}'\n","df=pd.read_csv(RUN/'metrics.csv')\n","for split in ['train','val','test']: df[f'{split}_error']=1-df[f'{split}_accuracy']\n","df.head()\n"]}, - {"cell_type":"code","execution_count":null,"metadata":{},"outputs":[],"source":["for metric in ['loss','accuracy','error','perplexity']:\n"," plt.figure(figsize=(9,5))\n"," for split in ['train','val','test']: plt.plot(df.step,df[f'{split}_{metric}'],label=split)\n"," plt.xlabel('step'); plt.ylabel(metric); plt.title(f'{OPT} seed {SEED}: {metric}'); plt.legend(); plt.grid(alpha=.25); plt.show()\n"]}, - {"cell_type":"code","execution_count":null,"metadata":{},"outputs":[],"source":["files=sorted(RUN.glob('weightwatcher_step_*.csv'))\n","if files:\n"," ww=pd.concat([pd.read_csv(f) for f in files],ignore_index=True)\n"," layer_col='layer_id' if 'layer_id' in ww else 'layer'\n"," for layer,g in ww.groupby(layer_col): plt.plot(g.step,g.alpha,label=str(layer))\n"," plt.xlabel('step'); plt.ylabel('alpha'); plt.legend(bbox_to_anchor=(1.02,1)); plt.show()\n","else: print('No WeightWatcher files found.')\n"]} + { + "cell_type": "markdown", + "metadata": {}, + "source": [ + "# Realistic Level 0 single-seed diagnostics\n", + "\n", + "This notebook analyzes the isolated GPT-2-BPE AdamW baseline. Test metrics are shown only for the final and validation-selected checkpoints; the training curve itself uses fixed train and validation probes.\n" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [], + "source": [ + "import json\n", + "import os\n", + "import re\n", + "from pathlib import Path\n", + "\n", + "import matplotlib.pyplot as plt\n", + "import numpy as np\n", + "import pandas as pd\n", + "\n", + "ROOT = Path(os.getenv(\"NANOGPT_LEVEL0_RESULTS_ROOT\", \"/tmp/nanogpt-level0-bpe/results\"))\n", + "OPTIMIZER = os.getenv(\"NANOGPT_LEVEL0_NOTEBOOK_OPTIMIZER\", \"adamw\")\n", + "SEED = int(os.getenv(\"NANOGPT_LEVEL0_NOTEBOOK_SEED\", \"1337\"))\n", + "RUN = ROOT / f\"{OPTIMIZER}_seed_{SEED}\"\n", + "REPORT = RUN / \"plots\"\n", + "REPORT.mkdir(parents=True, exist_ok=True)\n", + "\n", + "metrics = pd.read_csv(RUN / \"metrics.csv\")\n", + "manifest = json.loads((RUN / \"manifest.json\").read_text())\n", + "complete = json.loads((RUN / \"run_complete.json\").read_text())\n", + "final = json.loads((RUN / \"final_metrics.json\").read_text())\n", + "selected = json.loads((RUN / \"selected_checkpoint_metrics.json\").read_text())\n", + "\n", + "print(f\"Run: {RUN}\")\n", + "print(f\"Parameters: {manifest['parameter_count']:,}\")\n", + "print(f\"Model: {manifest['model_config']}\")\n", + "print(f\"Final step: {complete['final_step']:,}\")\n", + "print(f\"Final validation loss: {complete['final_validation_loss']:.4f}\")\n", + "print(f\"Final test loss: {final['test_loss']:.4f}\")\n", + "print(f\"Selected step: {selected['selected_step']:,}\")\n", + "print(f\"Selected test loss: {selected['test_loss']:.4f}\")\n", + "metrics.head()\n" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [], + "source": [ + "def save_current(name):\n", + " plt.tight_layout()\n", + " plt.savefig(REPORT / name, dpi=180, bbox_inches=\"tight\")\n", + " plt.show()\n", + "\n", + "plt.figure(figsize=(10, 6))\n", + "plt.plot(metrics.step, metrics.train_loss, label=\"train probe\")\n", + "plt.plot(metrics.step, metrics.val_loss, label=\"validation probe\")\n", + "plt.axvline(selected[\"selected_step\"], linestyle=\"--\", label=\"selected checkpoint\")\n", + "plt.xlabel(\"optimizer step\")\n", + "plt.ylabel(\"cross-entropy loss\")\n", + "plt.title(f\"{OPTIMIZER} seed {SEED}: loss\")\n", + "plt.grid(alpha=0.25)\n", + "plt.legend()\n", + "save_current(\"loss.png\")\n" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [], + "source": [ + "plt.figure(figsize=(10, 6))\n", + "plt.plot(metrics.step, metrics.train_perplexity, label=\"train probe\")\n", + "plt.plot(metrics.step, metrics.val_perplexity, label=\"validation probe\")\n", + "plt.scatter([selected[\"selected_step\"]], [selected[\"test_perplexity\"]], marker=\"x\", s=80, label=\"selected test\")\n", + "plt.scatter([final[\"step\"]], [final[\"test_perplexity\"]], marker=\"x\", s=80, label=\"final test\")\n", + "plt.xlabel(\"optimizer step\")\n", + "plt.ylabel(\"perplexity\")\n", + "plt.title(f\"{OPTIMIZER} seed {SEED}: perplexity\")\n", + "plt.grid(alpha=0.25)\n", + "plt.legend()\n", + "save_current(\"perplexity.png\")\n" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [], + "source": [ + "plt.figure(figsize=(10, 6))\n", + "plt.plot(metrics.step, 100 * metrics.train_accuracy, label=\"train top-1 token accuracy\")\n", + "plt.plot(metrics.step, 100 * metrics.val_accuracy, label=\"validation top-1 token accuracy\")\n", + "plt.scatter([selected[\"selected_step\"]], [100 * selected[\"test_accuracy\"]], marker=\"x\", s=80, label=\"selected test\")\n", + "plt.scatter([final[\"step\"]], [100 * final[\"test_accuracy\"]], marker=\"x\", s=80, label=\"final test\")\n", + "plt.xlabel(\"optimizer step\")\n", + "plt.ylabel(\"exact next-BPE-token accuracy (%)\")\n", + "plt.title(f\"{OPTIMIZER} seed {SEED}: token accuracy\")\n", + "plt.grid(alpha=0.25)\n", + "plt.legend()\n", + "save_current(\"token_accuracy.png\")\n" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [], + "source": [ + "plt.figure(figsize=(10, 6))\n", + "plt.plot(metrics.step, metrics.train_bits_per_token, label=\"train\")\n", + "plt.plot(metrics.step, metrics.val_bits_per_token, label=\"validation\")\n", + "plt.scatter([selected[\"selected_step\"]], [selected[\"test_bits_per_token\"]], marker=\"x\", s=80, label=\"selected test\")\n", + "plt.scatter([final[\"step\"]], [final[\"test_bits_per_token\"]], marker=\"x\", s=80, label=\"final test\")\n", + "plt.xlabel(\"optimizer step\")\n", + "plt.ylabel(\"bits per BPE token\")\n", + "plt.title(f\"{OPTIMIZER} seed {SEED}: bits per token\")\n", + "plt.grid(alpha=0.25)\n", + "plt.legend()\n", + "save_current(\"bits_per_token.png\")\n" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [], + "source": [ + "figure, axes = plt.subplots(2, 2, figsize=(12, 9))\n", + "axes[0, 0].plot(metrics.step, metrics.learning_rate)\n", + "axes[0, 0].set_title(\"warmup + cosine learning rate\")\n", + "axes[0, 1].plot(metrics.step, metrics.grad_norm)\n", + "axes[0, 1].set_title(\"gradient norm\")\n", + "axes[1, 0].plot(metrics.step, metrics.weight_norm)\n", + "axes[1, 0].set_title(\"weight norm\")\n", + "axes[1, 1].plot(metrics.step, metrics.val_generalization_gap)\n", + "axes[1, 1].axhline(0, linewidth=1)\n", + "axes[1, 1].set_title(\"validation loss \u2212 train loss\")\n", + "for axis in axes.ravel():\n", + " axis.set_xlabel(\"optimizer step\")\n", + " axis.grid(alpha=0.25)\n", + "plt.tight_layout()\n", + "plt.savefig(REPORT / \"optimization_diagnostics.png\", dpi=180, bbox_inches=\"tight\")\n", + "plt.show()\n" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [], + "source": [ + "weightwatcher_files = sorted(RUN.glob(\"weightwatcher_step_*.csv\"))\n", + "print(f\"WeightWatcher checkpoints: {len(weightwatcher_files)}\")\n", + "if weightwatcher_files:\n", + " ww = pd.concat([pd.read_csv(path) for path in weightwatcher_files], ignore_index=True)\n", + " ww[\"alpha\"] = pd.to_numeric(ww.get(\"alpha\"), errors=\"coerce\")\n", + " if \"matrix_name\" not in ww:\n", + " source = \"longname\" if \"longname\" in ww else \"name\"\n", + " ww[\"matrix_name\"] = ww[source].astype(str)\n", + " ww[\"matrix_type\"] = ww.matrix_name.str.replace(r\"^L\\d+_\", \"\", regex=True)\n", + " ww[\"block\"] = ww.matrix_name.str.extract(r\"^(L\\d+)\", expand=False)\n", + " ww = ww[np.isfinite(ww.alpha)].copy()\n", + " display_columns = [column for column in [\"step\", \"matrix_name\", \"alpha\", \"D\", \"xmin\", \"num_evals\"] if column in ww]\n", + " display(ww[display_columns].head())\n", + "else:\n", + " ww = pd.DataFrame()\n" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [], + "source": [ + "if not ww.empty:\n", + " plt.figure(figsize=(10, 6))\n", + " aggregate = ww.groupby(\"step\").alpha.agg([\"median\", \"min\", \"max\"]).reset_index()\n", + " plt.plot(aggregate.step, aggregate[\"median\"], label=\"median layer alpha\")\n", + " plt.fill_between(aggregate.step, aggregate[\"min\"], aggregate[\"max\"], alpha=0.15, label=\"layer range\")\n", + " plt.axhline(2.0, linestyle=\"--\", label=\"alpha = 2\")\n", + " plt.xlabel(\"optimizer step\")\n", + " plt.ylabel(\"WeightWatcher alpha\")\n", + " plt.title(f\"{OPTIMIZER} seed {SEED}: aggregate alpha trajectory\")\n", + " plt.grid(alpha=0.25)\n", + " plt.legend()\n", + " save_current(\"alpha_aggregate.png\")\n" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [], + "source": [ + "if not ww.empty:\n", + " for matrix_type, group in ww.groupby(\"matrix_type\"):\n", + " plt.figure(figsize=(10, 6))\n", + " for matrix_name, trajectory in group.groupby(\"matrix_name\"):\n", + " trajectory = trajectory.sort_values(\"step\")\n", + " plt.plot(trajectory.step, trajectory.alpha, marker=\"o\", markersize=3, label=matrix_name)\n", + " plt.axhline(2.0, linestyle=\"--\", linewidth=1, label=\"alpha = 2\")\n", + " plt.xlabel(\"optimizer step\")\n", + " plt.ylabel(\"WeightWatcher alpha\")\n", + " plt.title(f\"{OPTIMIZER} seed {SEED}: {matrix_type}\")\n", + " plt.grid(alpha=0.25)\n", + " plt.legend(ncol=2, fontsize=8)\n", + " save_current(f\"alpha_{matrix_type.lower()}.png\")\n", + "\n", + "print(f\"Saved plots to {REPORT}\")\n" + ] + } ], - "metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":"python","version":"3.10"}}, - "nbformat":4,"nbformat_minor":5 -} + "metadata": { + "kernelspec": { + "display_name": "Python 3", + "language": "python", + "name": "python3" + }, + "language_info": { + "name": "python", + "version": "3.10" + } + }, + "nbformat": 4, + "nbformat_minor": 5 +} \ No newline at end of file diff --git a/level_0_baseline/notebooks/02_multiseed.ipynb b/level_0_baseline/notebooks/02_multiseed.ipynb index d2ac6d9..ea0d4b7 100644 --- a/level_0_baseline/notebooks/02_multiseed.ipynb +++ b/level_0_baseline/notebooks/02_multiseed.ipynb @@ -1,10 +1,102 @@ { "cells": [ - {"cell_type":"markdown","metadata":{},"source":["# Level 0 multi-seed uncertainty bands\n"]}, - {"cell_type":"code","execution_count":null,"metadata":{},"outputs":[],"source":["import os\n","from pathlib import Path\n","import pandas as pd\n","import matplotlib.pyplot as plt\n","ROOT=Path(os.getenv('NANOGPT_LEVEL0_RESULTS_ROOT','/tmp/nanogpt-level0/results'))\n","rows=[]\n","for path in ROOT.glob('*_seed_*/metrics.csv'):\n"," optimizer,seed=path.parent.name.split('_seed_')\n"," d=pd.read_csv(path); d['optimizer']=optimizer; d['seed']=int(seed)\n"," for split in ['train','val','test']: d[f'{split}_error']=1-d[f'{split}_accuracy']\n"," rows.append(d)\n","all_df=pd.concat(rows,ignore_index=True)\n","all_df.groupby('optimizer').seed.nunique()\n"]}, - {"cell_type":"code","execution_count":null,"metadata":{},"outputs":[],"source":["def band_plot(metric):\n"," plt.figure(figsize=(10,6))\n"," for optimizer,d in all_df.groupby('optimizer'):\n"," a=d.groupby('step')[metric].agg(['mean','std']).reset_index(); s=a['std'].fillna(0)\n"," line,=plt.plot(a.step,a['mean'],label=optimizer)\n"," plt.fill_between(a.step,a['mean']-s,a['mean']+s,alpha=.2,color=line.get_color())\n"," plt.xlabel('step'); plt.ylabel(metric); plt.title(f'{metric}: mean ± 1 standard deviation'); plt.legend(); plt.grid(alpha=.25); plt.show()\n","for metric in ['test_loss','test_accuracy','test_error','test_perplexity','test_generalization_gap']: band_plot(metric)\n"]}, - {"cell_type":"code","execution_count":null,"metadata":{},"outputs":[],"source":["ww=[]\n","for run in ROOT.glob('*_seed_*'):\n"," optimizer,seed=run.name.split('_seed_')\n"," for f in run.glob('weightwatcher_step_*.csv'):\n"," d=pd.read_csv(f); d['optimizer']=optimizer; d['seed']=int(seed); ww.append(d)\n","if ww:\n"," ww=pd.concat(ww,ignore_index=True); layer_col='layer_id' if 'layer_id' in ww else 'layer'\n"," for layer,ld in ww.groupby(layer_col):\n"," plt.figure(figsize=(10,5))\n"," for optimizer,d in ld.groupby('optimizer'):\n"," a=d.groupby('step').alpha.agg(['mean','std']).reset_index(); s=a['std'].fillna(0)\n"," line,=plt.plot(a.step,a['mean'],label=optimizer); plt.fill_between(a.step,a['mean']-s,a['mean']+s,alpha=.2,color=line.get_color())\n"," plt.title(f'Layer {layer} alpha'); plt.legend(); plt.show()\n"]} + { + "cell_type": "markdown", + "metadata": {}, + "source": [ + "# Realistic Level 0 multi-seed diagnostics\n", + "\n", + "Aggregate completed AdamW or Muon runs using mean and one-standard-deviation bands. Test results are read from validation-selected checkpoints.\n" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [], + "source": [ + "import json\n", + "import os\n", + "from pathlib import Path\n", + "\n", + "import matplotlib.pyplot as plt\n", + "import numpy as np\n", + "import pandas as pd\n", + "\n", + "ROOT = Path(os.getenv(\"NANOGPT_LEVEL0_RESULTS_ROOT\", \"/tmp/nanogpt-level0-bpe/results\"))\n", + "REPORT = ROOT / \"multiseed-plots\"\n", + "REPORT.mkdir(parents=True, exist_ok=True)\n", + "rows = []\n", + "selected_rows = []\n", + "for path in sorted(ROOT.glob(\"*_seed_*/metrics.csv\")):\n", + " optimizer, seed_text = path.parent.name.split(\"_seed_\")\n", + " frame = pd.read_csv(path)\n", + " frame[\"optimizer\"] = optimizer\n", + " frame[\"seed\"] = int(seed_text)\n", + " rows.append(frame)\n", + " selected_path = path.parent / \"selected_checkpoint_metrics.json\"\n", + " if selected_path.exists():\n", + " selected = json.loads(selected_path.read_text())\n", + " selected.update({\"optimizer\": optimizer, \"seed\": int(seed_text)})\n", + " selected_rows.append(selected)\n", + "if not rows:\n", + " raise RuntimeError(f\"No completed metrics found under {ROOT}\")\n", + "all_metrics = pd.concat(rows, ignore_index=True)\n", + "selected_metrics = pd.DataFrame(selected_rows)\n", + "all_metrics.groupby(\"optimizer\").seed.nunique()\n" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [], + "source": [ + "def band_plot(metric, ylabel=None):\n", + " plt.figure(figsize=(10, 6))\n", + " for optimizer, data in all_metrics.groupby(\"optimizer\"):\n", + " summary = data.groupby(\"step\")[metric].agg([\"mean\", \"std\"]).reset_index()\n", + " spread = summary[\"std\"].fillna(0)\n", + " line, = plt.plot(summary.step, summary[\"mean\"], label=optimizer)\n", + " plt.fill_between(summary.step, summary[\"mean\"] - spread, summary[\"mean\"] + spread, alpha=0.2, color=line.get_color())\n", + " plt.xlabel(\"optimizer step\")\n", + " plt.ylabel(ylabel or metric)\n", + " plt.title(f\"{metric}: mean \u00b1 one standard deviation\")\n", + " plt.grid(alpha=0.25)\n", + " plt.legend()\n", + " plt.tight_layout()\n", + " plt.savefig(REPORT / f\"{metric}.png\", dpi=180, bbox_inches=\"tight\")\n", + " plt.show()\n", + "\n", + "for metric in [\"val_loss\", \"val_perplexity\", \"val_accuracy\", \"val_generalization_gap\"]:\n", + " band_plot(metric)\n" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [], + "source": [ + "if not selected_metrics.empty:\n", + " display(selected_metrics[[\"optimizer\", \"seed\", \"selected_step\", \"validation_loss\", \"test_loss\", \"test_perplexity\", \"test_accuracy\"]].sort_values([\"optimizer\", \"seed\"]))\n", + " summary = selected_metrics.groupby(\"optimizer\")[[\"test_loss\", \"test_perplexity\", \"test_accuracy\"]].agg([\"mean\", \"std\"])\n", + " display(summary)\n", + "print(f\"Saved plots to {REPORT}\")\n" + ] + } ], - "metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":"python","version":"3.10"}}, - "nbformat":4,"nbformat_minor":5 -} + "metadata": { + "kernelspec": { + "display_name": "Python 3", + "language": "python", + "name": "python3" + }, + "language_info": { + "name": "python", + "version": "3.10" + } + }, + "nbformat": 4, + "nbformat_minor": 5 +} \ No newline at end of file diff --git a/level_0_baseline/pyproject.toml b/level_0_baseline/pyproject.toml index 5840ab6..f1dd707 100644 --- a/level_0_baseline/pyproject.toml +++ b/level_0_baseline/pyproject.toml @@ -4,14 +4,20 @@ build-backend = "setuptools.build_meta" [project] name = "nanogpt-level0-baseline" -version = "0.1.0" -description = "Self-contained Level 0 nanoGPT baseline for AdamW and Muon" +version = "0.2.0" +description = "Isolated MacBook-scale BPE nanoGPT baseline with WeightWatcher alpha tracking" requires-python = ">=3.10" -dependencies = ["torch>=2.2", "numpy>=1.24", "pandas>=2.0", "matplotlib>=3.7", "pyyaml>=6.0"] +dependencies = [ + "torch>=2.2", + "numpy>=1.24", + "pandas>=2.0", + "matplotlib>=3.7", + "pyyaml>=6.0", +] [project.optional-dependencies] -data = ["datasets>=2.19"] -analysis = ["weightwatcher>=0.7.5", "jupyter>=1.0"] +data = ["datasets>=2.19", "tiktoken>=0.7"] +analysis = ["weightwatcher>=0.7.5", "jupyter>=1.0", "nbconvert>=7.0"] test = ["pytest>=8.0"] [project.scripts] @@ -20,3 +26,6 @@ level0-train = "level0_baseline.train:main" [tool.setuptools.packages.find] where = ["src"] + +[tool.pytest.ini_options] +testpaths = ["tests"] diff --git a/level_0_baseline/scripts/prepare_data.sh b/level_0_baseline/scripts/prepare_data.sh new file mode 100755 index 0000000..e990d4f --- /dev/null +++ b/level_0_baseline/scripts/prepare_data.sh @@ -0,0 +1,24 @@ +#!/usr/bin/env bash +set -euo pipefail + +ROOT="${NANOGPT_LEVEL0_ROOT:-/tmp/nanogpt-level0-bpe}" +DATA_ROOT="${NANOGPT_LEVEL0_DATA_ROOT:-$ROOT/data}" +CACHE_ROOT="${NANOGPT_LEVEL0_CACHE_ROOT:-$ROOT/cache}" + +export HF_HOME="${HF_HOME:-$CACHE_ROOT/huggingface}" +export HF_DATASETS_CACHE="${HF_DATASETS_CACHE:-$HF_HOME/datasets}" +export HF_HUB_CACHE="${HF_HUB_CACHE:-$HF_HOME/hub}" +export TOKENIZERS_PARALLELISM=false + +mkdir -p "$DATA_ROOT" "$HF_DATASETS_CACHE" "$HF_HUB_CACHE" + +level0-prepare-data \ + --dataset fineweb-edu \ + --output-dir "$DATA_ROOT" \ + --train-tokens "${NANOGPT_LEVEL0_TRAIN_TOKENS:-20000000}" \ + --val-tokens "${NANOGPT_LEVEL0_VAL_TOKENS:-1000000}" \ + --test-tokens "${NANOGPT_LEVEL0_TEST_TOKENS:-1000000}" \ + --tokenizer gpt2 \ + --model-vocab-size 50304 \ + --verbose \ + --log-interval-seconds "${NANOGPT_LEVEL0_DATA_LOG_INTERVAL:-10}" diff --git a/level_0_baseline/scripts/run_multiseed.sh b/level_0_baseline/scripts/run_multiseed.sh old mode 100644 new mode 100755 index 27aefe0..c416ae0 --- a/level_0_baseline/scripts/run_multiseed.sh +++ b/level_0_baseline/scripts/run_multiseed.sh @@ -1,9 +1,14 @@ #!/usr/bin/env bash set -euo pipefail + SEEDS="${NANOGPT_LEVEL0_SEEDS:-1337,2027,4099}" -for optimizer in adamw muon; do - IFS=',' read -ra xs <<< "$SEEDS" - for seed in "${xs[@]}"; do - ./scripts/run_one.sh "$optimizer" "$seed" +OPTIMIZERS="${NANOGPT_LEVEL0_OPTIMIZERS:-adamw}" +DEVICE="${NANOGPT_LEVEL0_DEVICE:-auto}" + +IFS=',' read -r -a optimizer_array <<< "$OPTIMIZERS" +IFS=',' read -r -a seed_array <<< "$SEEDS" +for optimizer in "${optimizer_array[@]}"; do + for seed in "${seed_array[@]}"; do + ./scripts/run_one.sh "$optimizer" "$seed" "$DEVICE" done done diff --git a/level_0_baseline/scripts/run_one.sh b/level_0_baseline/scripts/run_one.sh old mode 100644 new mode 100755 index f9dbed7..9192089 --- a/level_0_baseline/scripts/run_one.sh +++ b/level_0_baseline/scripts/run_one.sh @@ -1,6 +1,31 @@ #!/usr/bin/env bash set -euo pipefail -ROOT="${NANOGPT_LEVEL0_ROOT:-/tmp/nanogpt-level0}" + +ROOT="${NANOGPT_LEVEL0_ROOT:-/tmp/nanogpt-level0-bpe}" OPTIMIZER="${1:-adamw}" SEED="${2:-1337}" -python -m level0_baseline.train --config configs/level0.yaml --optimizer "$OPTIMIZER" --seed "$SEED" --data-root "${NANOGPT_LEVEL0_DATA_ROOT:-$ROOT/data}" --results-root "${NANOGPT_LEVEL0_RESULTS_ROOT:-$ROOT/results}" +DEVICE="${3:-auto}" +DATA_ROOT="${NANOGPT_LEVEL0_DATA_ROOT:-$ROOT/data}" +RESULTS_ROOT="${NANOGPT_LEVEL0_RESULTS_ROOT:-$ROOT/results}" +CONFIG="${NANOGPT_LEVEL0_CONFIG:-configs/level0.yaml}" + +ARGS=( + --config "$CONFIG" + --optimizer "$OPTIMIZER" + --seed "$SEED" + --device "$DEVICE" + --data-root "$DATA_ROOT" + --results-root "$RESULTS_ROOT" +) + +RUN_DIR="$RESULTS_ROOT/${OPTIMIZER}_seed_${SEED}" +if [[ "${NANOGPT_LEVEL0_OVERWRITE:-0}" == "1" ]]; then + ARGS+=(--overwrite) +elif [[ -f "$RUN_DIR/checkpoint_latest.pt" || -f "$RUN_DIR/run_complete.json" ]]; then + ARGS+=(--resume) +fi +if [[ "${NANOGPT_LEVEL0_DISABLE_WEIGHTWATCHER:-0}" == "1" ]]; then + ARGS+=(--no-weightwatcher) +fi + +python -m level0_baseline.train "${ARGS[@]}" diff --git a/level_0_baseline/src/level0_baseline/__init__.py b/level_0_baseline/src/level0_baseline/__init__.py index 8fe33d4..f9a9adb 100644 --- a/level_0_baseline/src/level0_baseline/__init__.py +++ b/level_0_baseline/src/level0_baseline/__init__.py @@ -1,2 +1,3 @@ -"""Self-contained Level 0 nanoGPT baseline.""" -__version__ = "0.1.0" +"""Isolated realistic Level 0 nanoGPT baseline.""" + +__version__ = "0.2.0" diff --git a/level_0_baseline/src/level0_baseline/config.py b/level_0_baseline/src/level0_baseline/config.py index 3d19743..67566b1 100644 --- a/level_0_baseline/src/level0_baseline/config.py +++ b/level_0_baseline/src/level0_baseline/config.py @@ -1,9 +1,14 @@ from __future__ import annotations + +import copy import os from pathlib import Path +from typing import Any + import yaml -DEFAULT_ROOT = Path("/tmp/nanogpt-level0") +DEFAULT_ROOT = Path("/tmp/nanogpt-level0-bpe") + def roots() -> dict[str, Path]: root = Path(os.getenv("NANOGPT_LEVEL0_ROOT", DEFAULT_ROOT)) @@ -14,18 +19,112 @@ def roots() -> dict[str, Path]: "cache": Path(os.getenv("NANOGPT_LEVEL0_CACHE_ROOT", root / "cache")), } -def load_config(path: str | Path) -> dict: - with open(path, "r", encoding="utf-8") as f: - cfg = yaml.safe_load(f) - env_map = { - "NANOGPT_LEVEL0_SEED": ("training", "seed", int), - "NANOGPT_LEVEL0_OPTIMIZER": ("training", "optimizer", str), - "NANOGPT_LEVEL0_MAX_STEPS": ("training", "max_steps", int), - "NANOGPT_LEVEL0_BATCH_SIZE": ("training", "batch_size", int), - "NANOGPT_LEVEL0_LR": ("training", "learning_rate", float), - "NANOGPT_LEVEL0_EVAL_INTERVAL": ("training", "eval_interval", int), - } - for name, (section, key, cast) in env_map.items(): + +_ENV_OVERRIDES: dict[str, tuple[str, str, type]] = { + "NANOGPT_LEVEL0_SEED": ("training", "seed", int), + "NANOGPT_LEVEL0_OPTIMIZER": ("training", "optimizer", str), + "NANOGPT_LEVEL0_MAX_STEPS": ("training", "max_steps", int), + "NANOGPT_LEVEL0_BATCH_SIZE": ("training", "batch_size", int), + "NANOGPT_LEVEL0_GRAD_ACCUM_STEPS": ( + "training", + "grad_accum_steps", + int, + ), + "NANOGPT_LEVEL0_LR": ("training", "learning_rate", float), + "NANOGPT_LEVEL0_EVAL_INTERVAL": ("training", "eval_interval", int), + "NANOGPT_LEVEL0_LOG_INTERVAL": ("training", "log_interval", int), + "NANOGPT_LEVEL0_WEIGHTWATCHER_INTERVAL": ( + "analysis", + "weightwatcher_interval", + int, + ), +} + + +def _positive_int(value: Any, name: str) -> int: + parsed = int(value) + if parsed <= 0: + raise ValueError(f"{name} must be positive") + return parsed + + +def validate_config(cfg: dict[str, Any]) -> None: + if not isinstance(cfg, dict): + raise ValueError("configuration must be a mapping") + for section in ("model", "training", "analysis"): + if not isinstance(cfg.get(section), dict): + raise ValueError(f"missing configuration section: {section}") + + model = cfg["model"] + training = cfg["training"] + analysis = cfg["analysis"] + + for key in ("vocab_size", "block_size", "n_layer", "n_head", "n_embd"): + model[key] = _positive_int(model[key], f"model.{key}") + if model["vocab_size"] > 65_535: + raise ValueError("model.vocab_size must fit in uint16 token files") + if model["n_embd"] % model["n_head"] != 0: + raise ValueError("model.n_embd must be divisible by model.n_head") + dropout = float(model.get("dropout", 0.0)) + if not 0.0 <= dropout < 1.0: + raise ValueError("model.dropout must satisfy 0 <= dropout < 1") + model["dropout"] = dropout + model["bias"] = bool(model.get("bias", False)) + + for key in ( + "batch_size", + "grad_accum_steps", + "max_steps", + "eval_interval", + "eval_batches", + "log_interval", + "checkpoint_interval", + "warmup_steps", + ): + training[key] = _positive_int(training[key], f"training.{key}") + if training["warmup_steps"] >= training["max_steps"]: + raise ValueError("training.warmup_steps must be smaller than max_steps") + + for key in ("learning_rate", "min_lr", "weight_decay", "grad_clip"): + training[key] = float(training[key]) + if training["learning_rate"] <= 0: + raise ValueError("training.learning_rate must be positive") + if not 0 <= training["min_lr"] <= training["learning_rate"]: + raise ValueError("training.min_lr must be between 0 and learning_rate") + if training["weight_decay"] < 0: + raise ValueError("training.weight_decay must be nonnegative") + if training["grad_clip"] < 0: + raise ValueError("training.grad_clip must be nonnegative") + + training["beta1"] = float(training["beta1"]) + training["beta2"] = float(training["beta2"]) + training["epsilon"] = float(training.get("epsilon", 1e-8)) + if training["epsilon"] <= 0: + raise ValueError("training.epsilon must be positive") + if not 0 <= training["beta1"] < 1 or not 0 <= training["beta2"] < 1: + raise ValueError("training betas must be in [0, 1)") + training["seed"] = int(training["seed"]) + training["optimizer"] = str(training["optimizer"]).lower() + if training["optimizer"] not in {"adamw", "muon"}: + raise ValueError("training.optimizer must be adamw or muon") + training["compile"] = bool(training.get("compile", False)) + + analysis["weightwatcher"] = bool(analysis.get("weightwatcher", True)) + analysis["weightwatcher_interval"] = _positive_int( + analysis["weightwatcher_interval"], + "analysis.weightwatcher_interval", + ) + # Alpha tracking must be deterministic. Randomized ESDs are a separate trap + # diagnostic and are deliberately not used in this isolated baseline. + analysis["randomize"] = False + + +def load_config(path: str | Path) -> dict[str, Any]: + with open(path, "r", encoding="utf-8") as handle: + loaded = yaml.safe_load(handle) + cfg = copy.deepcopy(loaded) + for name, (section, key, cast) in _ENV_OVERRIDES.items(): if name in os.environ: cfg[section][key] = cast(os.environ[name]) + validate_config(cfg) return cfg diff --git a/level_0_baseline/src/level0_baseline/data.py b/level_0_baseline/src/level0_baseline/data.py index ed96451..5dbffb3 100644 --- a/level_0_baseline/src/level0_baseline/data.py +++ b/level_0_baseline/src/level0_baseline/data.py @@ -1,24 +1,39 @@ from __future__ import annotations import argparse +import gc +import hashlib import json import sys import threading import time +from collections import OrderedDict from pathlib import Path +from typing import Iterable, Protocol import numpy as np from .config import roots +FINEWEB_DATASET = "HuggingFaceFW/fineweb-edu" +FINEWEB_CONFIG = "sample-10BT" +FINEWEB_REVISION = "593b3a867298afb8ce42625a270ef20ddcad28f9" +TOKEN_DTYPE = np.dtype(" np.ndarray: - return np.frombuffer(text.encode("utf-8", errors="replace"), dtype=np.uint8) +class Encoder(Protocol): + name: str + eot_token: int + n_vocab: int -def _format_duration(seconds: float) -> str: - seconds = max(0, int(seconds)) - hours, remainder = divmod(seconds, 3600) + def encode_ordinary(self, text: str) -> list[int]: ... + + +def _format_duration(seconds: float | None) -> str: + if seconds is None or not np.isfinite(seconds): + return "unknown" + whole = max(0, int(seconds)) + hours, remainder = divmod(whole, 3600) minutes, secs = divmod(remainder, 60) if hours: return f"{hours:d}h{minutes:02d}m{secs:02d}s" @@ -29,253 +44,390 @@ def _format_duration(seconds: float) -> str: def _progress_message( *, - collected_bytes: int, - required_bytes: int, + written_tokens: int, + required_tokens: int, documents: int, elapsed_seconds: float, stalled_seconds: float, + phase: str, + split: str, ) -> str: elapsed_seconds = max(elapsed_seconds, 1e-9) - rate = collected_bytes / elapsed_seconds - remaining = max(required_bytes - collected_bytes, 0) + rate = written_tokens / elapsed_seconds + remaining = max(required_tokens - written_tokens, 0) eta = remaining / rate if rate > 0 else None - percent = 100.0 * collected_bytes / max(required_bytes, 1) - speed_mib = rate / (1024 * 1024) - eta_text = _format_duration(eta) if eta is not None else "unknown" - stall_text = _format_duration(stalled_seconds) + percent = 100.0 * written_tokens / max(required_tokens, 1) return ( "[level0-prepare-data] progress " - f"documents={documents:,} " - f"bytes={collected_bytes:,}/{required_bytes:,} " - f"percent={percent:5.1f}% " - f"elapsed={_format_duration(elapsed_seconds)} " - f"speed={speed_mib:.2f} MiB/s " - f"eta={eta_text} " - f"no_new_bytes_for={stall_text}" + f"phase={phase} split={split} documents={documents:,} " + f"tokens={written_tokens:,}/{required_tokens:,} " + f"percent={percent:5.1f}% elapsed={_format_duration(elapsed_seconds)} " + f"speed={rate:,.0f} tok/s eta={_format_duration(eta)} " + f"no_new_tokens_for={_format_duration(stalled_seconds)}" ) -class _ProgressReporter: - """Emit heartbeat logs even while the streaming iterator is blocked.""" +class ProgressReporter: + """Emit heartbeats even when dataset resolution or streaming blocks.""" - def __init__(self, required_bytes: int, interval_seconds: float): - self.required_bytes = int(required_bytes) + def __init__(self, required_tokens: int, interval_seconds: float): + self.required_tokens = int(required_tokens) self.interval_seconds = float(interval_seconds) self.started_at = time.monotonic() self.last_progress_at = self.started_at - self.collected_bytes = 0 + self.written_tokens = 0 self.documents = 0 + self.phase = "starting" + self.split = "none" self._lock = threading.Lock() self._stop = threading.Event() self._thread = threading.Thread( target=self._run, - name="level0-data-progress", + name="level0-bpe-data-progress", daemon=True, ) def start(self, output_dir: Path) -> None: print( "[level0-prepare-data] starting " - f"required_bytes={self.required_bytes:,} output={output_dir}", + f"required_tokens={self.required_tokens:,} output={output_dir}", file=sys.stderr, flush=True, ) self._thread.start() - def update(self, documents: int, collected_bytes: int) -> None: + def set_phase(self, phase: str, split: str | None = None) -> None: + with self._lock: + self.phase = str(phase) + if split is not None: + self.split = str(split) + + def update(self, *, documents: int, written_tokens: int, split: str) -> None: now = time.monotonic() with self._lock: - if collected_bytes > self.collected_bytes: + if written_tokens > self.written_tokens: self.last_progress_at = now self.documents = int(documents) - self.collected_bytes = int(collected_bytes) + self.written_tokens = int(written_tokens) + self.phase = "tokenizing" + self.split = str(split) - def _snapshot(self) -> tuple[int, int, float, float]: + def snapshot(self) -> tuple[int, int, float, float, str, str]: now = time.monotonic() with self._lock: return ( self.documents, - self.collected_bytes, + self.written_tokens, now - self.started_at, now - self.last_progress_at, + self.phase, + self.split, ) def _run(self) -> None: while not self._stop.wait(self.interval_seconds): - documents, collected_bytes, elapsed, stalled = self._snapshot() + documents, written, elapsed, stalled, phase, split = self.snapshot() print( _progress_message( - collected_bytes=collected_bytes, - required_bytes=self.required_bytes, + written_tokens=written, + required_tokens=self.required_tokens, documents=documents, elapsed_seconds=elapsed, stalled_seconds=stalled, + phase=phase, + split=split, ), file=sys.stderr, flush=True, ) - def stop(self) -> tuple[int, int, float, float]: + def stop(self) -> tuple[int, int, float, float, str, str]: self._stop.set() self._thread.join(timeout=max(1.0, self.interval_seconds + 1.0)) - return self._snapshot() - - -def _fineweb_texts(load_dataset, verbose: bool): - if verbose: - print( - "[level0-prepare-data] resolving streamed dataset " - "HuggingFaceFW/fineweb-edu sample-10BT train", - file=sys.stderr, - flush=True, - ) - dataset = load_dataset( - "HuggingFaceFW/fineweb-edu", - name="sample-10BT", - split="train", - streaming=True, - ) - if verbose: - print( - "[level0-prepare-data] dataset stream ready; collecting documents", - file=sys.stderr, - flush=True, - ) - for row in dataset: - yield row["text"] + return self.snapshot() -def write_splits( - texts, - out: Path, - train_bytes: int, - val_bytes: int, - test_bytes: int, +def load_tokenizer(name: str) -> Encoder: + try: + import tiktoken + except ImportError as exc: + raise SystemExit( + "Install BPE data support with: pip install -e '.[data]'" + ) from exc + encoding = tiktoken.get_encoding(name) + return encoding + + +def _validate_targets(split_targets: dict[str, int]) -> OrderedDict[str, int]: + expected = ("train", "val", "test") + if tuple(split_targets) != expected: + raise ValueError(f"split targets must be ordered as {expected}") + out: OrderedDict[str, int] = OrderedDict() + for name, value in split_targets.items(): + parsed = int(value) + if parsed <= 0: + raise ValueError(f"{name} token target must be positive") + out[name] = parsed + return out + + +def prepare_token_splits( + texts: Iterable[str], + output_dir: Path, + split_targets: dict[str, int], + tokenizer: Encoder, *, - verbose: bool = False, - log_interval_seconds: float = 10.0, -): - out.mkdir(parents=True, exist_ok=True) - need = train_bytes + val_bytes + test_bytes - chunks = [] - total = 0 - documents = 0 - reporter = ( - _ProgressReporter(need, log_interval_seconds) if verbose else None - ) - if reporter is not None: - reporter.start(out) + model_vocab_size: int, + reporter: ProgressReporter | None = None, + dataset_metadata: dict[str, object] | None = None, +) -> dict[str, object]: + """Tokenize streamed documents into non-overlapping fixed BPE splits. + + A document is never shared across splits. If a document crosses a split + boundary, only the prefix needed to finish the current split is retained and + the remainder is discarded; the next split starts from the next document. + """ + + targets = _validate_targets(split_targets) + if int(model_vocab_size) < int(tokenizer.n_vocab): + raise ValueError("model_vocab_size must be at least tokenizer.n_vocab") + if int(model_vocab_size) > 65_535: + raise ValueError("model_vocab_size must fit in uint16 token files") + + output_dir.mkdir(parents=True, exist_ok=True) + temporary_paths = { + name: output_dir / f".{name}.bin.tmp" for name in targets + } + final_paths = {name: output_dir / f"{name}.bin" for name in targets} + for path in temporary_paths.values(): + path.unlink(missing_ok=True) + + handles = { + name: open(path, "wb") for name, path in temporary_paths.items() + } + hashers = {name: hashlib.sha256() for name in targets} + written = {name: 0 for name in targets} + documents_by_split = {name: 0 for name in targets} + discarded_boundary_tokens = 0 + documents_seen = 0 + split_names = list(targets) + split_index = 0 try: for text in texts: - documents += 1 - x = encode(text + "\n") - chunks.append(x) - total += len(x) - if reporter is not None: - reporter.update(documents, total) - if total >= need: + if split_index >= len(split_names): break + documents_seen += 1 + token_ids = list(tokenizer.encode_ordinary(str(text))) + token_ids.append(int(tokenizer.eot_token)) + if not token_ids: + continue + minimum = min(token_ids) + maximum = max(token_ids) + if minimum < 0 or maximum >= int(model_vocab_size): + raise ValueError( + "tokenizer emitted an id outside the configured model vocabulary" + ) + + split = split_names[split_index] + remaining = targets[split] - written[split] + take = min(remaining, len(token_ids)) + if take: + values = np.asarray(token_ids[:take], dtype=TOKEN_DTYPE) + payload = values.tobytes(order="C") + handles[split].write(payload) + hashers[split].update(payload) + written[split] += take + documents_by_split[split] += 1 + if take < len(token_ids): + discarded_boundary_tokens += len(token_ids) - take + + total_written = sum(written.values()) + if reporter is not None: + reporter.update( + documents=documents_seen, + written_tokens=total_written, + split=split, + ) + + if written[split] == targets[split]: + split_index += 1 + if split_index < len(split_names) and reporter is not None: + reporter.set_phase("tokenizing", split_names[split_index]) finally: - snapshot = reporter.stop() if reporter is not None else None - - if total < need: - raise RuntimeError(f"corpus supplied {total:,} bytes; need {need:,}") - - if verbose and snapshot is not None: - _, _, elapsed, _ = snapshot - print( - "[level0-prepare-data] collection complete " - f"documents={documents:,} collected_bytes={total:,} " - f"elapsed={_format_duration(elapsed)}; writing fixed splits", - file=sys.stderr, - flush=True, - ) - - all_tokens = np.concatenate(chunks)[:need] - boundaries = { - "train": (0, train_bytes), - "val": (train_bytes, train_bytes + val_bytes), - "test": (train_bytes + val_bytes, need), + for handle in handles.values(): + handle.flush() + handle.close() + + incomplete = { + name: targets[name] - written[name] + for name in targets + if written[name] != targets[name] } - for name, (start, end) in boundaries.items(): - all_tokens[start:end].tofile(out / f"{name}.bin") - (out / "meta.json").write_text( - json.dumps( - { - "tokenizer": "utf8-byte", - "vocab_size": 256, - "sizes": { - name: end - start - for name, (start, end) in boundaries.items() - }, - }, - indent=2, - ) + if incomplete: + for path in temporary_paths.values(): + path.unlink(missing_ok=True) + raise RuntimeError(f"stream ended before fixed splits were filled: {incomplete}") + + for name in targets: + final_paths[name].unlink(missing_ok=True) + temporary_paths[name].replace(final_paths[name]) + + metadata: dict[str, object] = { + "format_version": 2, + "tokenizer": f"tiktoken:{tokenizer.name}", + "tokenizer_vocab_size": int(tokenizer.n_vocab), + "model_vocab_size": int(model_vocab_size), + "eot_token": int(tokenizer.eot_token), + "dtype": TOKEN_DTYPE.str, + "split_tokens": written, + "split_documents": documents_by_split, + "documents_seen": documents_seen, + "discarded_boundary_tokens": discarded_boundary_tokens, + "sha256": {name: hashers[name].hexdigest() for name in targets}, + } + if dataset_metadata: + metadata["dataset"] = dict(dataset_metadata) + (output_dir / "meta.json").write_text( + json.dumps(metadata, indent=2, sort_keys=True), + encoding="utf-8", ) + return metadata - if verbose: - sizes = ", ".join( - f"{name}={end - start:,}" - for name, (start, end) in boundaries.items() - ) - print( - f"[level0-prepare-data] complete output={out} {sizes}", - file=sys.stderr, - flush=True, - ) + +def _local_text_stream(path: Path) -> Iterable[str]: + text = path.read_text(encoding="utf-8") + while True: + yield text -def main(): - parser = argparse.ArgumentParser() +def main() -> None: + parser = argparse.ArgumentParser( + description="Prepare fixed GPT-2-BPE FineWeb-Edu splits for Level 0" + ) parser.add_argument("--dataset", default="fineweb-edu") parser.add_argument("--output-dir") - parser.add_argument("--train-bytes", type=int, default=50_000_000) - parser.add_argument("--val-bytes", type=int, default=2_000_000) - parser.add_argument("--test-bytes", type=int, default=2_000_000) + parser.add_argument("--train-tokens", type=int, default=20_000_000) + parser.add_argument("--val-tokens", type=int, default=1_000_000) + parser.add_argument("--test-tokens", type=int, default=1_000_000) + parser.add_argument("--tokenizer", default="gpt2") + parser.add_argument("--model-vocab-size", type=int, default=50_304) parser.add_argument("--local-text") - parser.add_argument( - "--verbose", - action="store_true", - help="print streaming progress, elapsed time, throughput, ETA, and stall heartbeats", - ) - parser.add_argument( - "--log-interval-seconds", - type=float, - default=10.0, - help="heartbeat interval used with --verbose (default: 10 seconds)", - ) + parser.add_argument("--verbose", action="store_true") + parser.add_argument("--log-interval-seconds", type=float, default=10.0) args = parser.parse_args() if args.log_interval_seconds <= 0: parser.error("--log-interval-seconds must be greater than zero") - - out = Path(args.output_dir) if args.output_dir else roots()["data"] - if args.local_text: - text = Path(args.local_text).read_text(encoding="utf-8") - - def repeat(): - while True: - yield text - - texts = repeat() - else: - try: - from datasets import load_dataset - except ImportError as exc: - raise SystemExit("Install data support: pip install -e '.[data]'") from exc - texts = _fineweb_texts(load_dataset, args.verbose) - - write_splits( - texts, - out, - args.train_bytes, - args.val_bytes, - args.test_bytes, - verbose=args.verbose, - log_interval_seconds=args.log_interval_seconds, + if args.dataset != "fineweb-edu" and not args.local_text: + parser.error("the isolated baseline currently supports --dataset fineweb-edu") + + output_dir = Path(args.output_dir) if args.output_dir else roots()["data"] + targets = OrderedDict( + train=args.train_tokens, + val=args.val_tokens, + test=args.test_tokens, ) - print(out) + required_tokens = sum(targets.values()) + reporter = ( + ProgressReporter(required_tokens, args.log_interval_seconds) + if args.verbose + else None + ) + if reporter is not None: + reporter.start(output_dir) + + dataset = None + try: + if reporter is not None: + reporter.set_phase("loading_tokenizer", "none") + tokenizer = load_tokenizer(args.tokenizer) + if args.verbose: + print( + "[level0-prepare-data] tokenizer ready " + f"name={tokenizer.name} vocab={tokenizer.n_vocab:,} " + f"eot={tokenizer.eot_token}", + file=sys.stderr, + flush=True, + ) + + if args.local_text: + texts = _local_text_stream(Path(args.local_text)) + dataset_metadata = { + "name": "local-text", + "path": str(Path(args.local_text).resolve()), + } + else: + try: + from datasets import load_dataset + except ImportError as exc: + raise SystemExit( + "Install BPE data support with: pip install -e '.[data]'" + ) from exc + if reporter is not None: + reporter.set_phase("resolving_dataset", "none") + if args.verbose: + print( + "[level0-prepare-data] resolving streamed dataset " + f"{FINEWEB_DATASET} {FINEWEB_CONFIG} train " + f"revision={FINEWEB_REVISION}", + file=sys.stderr, + flush=True, + ) + dataset = load_dataset( + FINEWEB_DATASET, + name=FINEWEB_CONFIG, + split="train", + revision=FINEWEB_REVISION, + streaming=True, + ) + texts = (row["text"] for row in dataset) + dataset_metadata = { + "name": FINEWEB_DATASET, + "config": FINEWEB_CONFIG, + "split": "train", + "revision": FINEWEB_REVISION, + "streaming": True, + } + if args.verbose: + print( + "[level0-prepare-data] dataset stream ready; tokenizing documents", + file=sys.stderr, + flush=True, + ) + + if reporter is not None: + reporter.set_phase("tokenizing", "train") + metadata = prepare_token_splits( + texts, + output_dir, + targets, + tokenizer, + model_vocab_size=args.model_vocab_size, + reporter=reporter, + dataset_metadata=dataset_metadata, + ) + if args.verbose: + print( + "[level0-prepare-data] complete " + f"output={output_dir} " + + " ".join( + f"{name}={count:,}" + for name, count in metadata["split_tokens"].items() + ), + file=sys.stderr, + flush=True, + ) + finally: + if reporter is not None: + reporter.stop() + close = getattr(dataset, "close", None) + if callable(close): + close() + dataset = None + gc.collect() + + print(output_dir, flush=True) if __name__ == "__main__": diff --git a/level_0_baseline/src/level0_baseline/model.py b/level_0_baseline/src/level0_baseline/model.py index 907ae02..b1834fb 100644 --- a/level_0_baseline/src/level0_baseline/model.py +++ b/level_0_baseline/src/level0_baseline/model.py @@ -1,45 +1,89 @@ from __future__ import annotations -from dataclasses import dataclass + import math +from dataclasses import dataclass +from typing import Iterator + import torch import torch.nn as nn import torch.nn.functional as F -@dataclass + +@dataclass(frozen=True) class GPTConfig: - vocab_size: int = 256 + vocab_size: int = 50_304 block_size: int = 256 - n_layer: int = 1 - n_head: int = 1 - n_embd: int = 64 + n_layer: int = 4 + n_head: int = 4 + n_embd: int = 256 dropout: float = 0.0 bias: bool = False + class CausalSelfAttention(nn.Module): def __init__(self, cfg: GPTConfig): super().__init__() - assert cfg.n_embd % cfg.n_head == 0 - self.n_head, self.n_embd, self.dropout = cfg.n_head, cfg.n_embd, cfg.dropout + if cfg.n_embd % cfg.n_head != 0: + raise ValueError("n_embd must be divisible by n_head") + self.n_head = cfg.n_head + self.n_embd = cfg.n_embd + self.dropout = cfg.dropout self.q_proj = nn.Linear(cfg.n_embd, cfg.n_embd, bias=cfg.bias) self.k_proj = nn.Linear(cfg.n_embd, cfg.n_embd, bias=cfg.bias) self.v_proj = nn.Linear(cfg.n_embd, cfg.n_embd, bias=cfg.bias) self.out_proj = nn.Linear(cfg.n_embd, cfg.n_embd, bias=cfg.bias) self.resid_dropout = nn.Dropout(cfg.dropout) - self.register_buffer("mask", torch.tril(torch.ones(cfg.block_size, cfg.block_size)).view(1, 1, cfg.block_size, cfg.block_size)) + self.has_sdpa = hasattr(F, "scaled_dot_product_attention") + if not self.has_sdpa: + mask = torch.tril(torch.ones(cfg.block_size, cfg.block_size)) + self.register_buffer( + "causal_mask", + mask.view(1, 1, cfg.block_size, cfg.block_size), + persistent=False, + ) def forward(self, x: torch.Tensor) -> torch.Tensor: - b, t, c = x.shape - hs = c // self.n_head - q = self.q_proj(x).view(b, t, self.n_head, hs).transpose(1, 2) - k = self.k_proj(x).view(b, t, self.n_head, hs).transpose(1, 2) - v = self.v_proj(x).view(b, t, self.n_head, hs).transpose(1, 2) - att = (q @ k.transpose(-2, -1)) / math.sqrt(hs) - att = att.masked_fill(self.mask[:, :, :t, :t] == 0, float("-inf")) - att = F.softmax(att, dim=-1) - att = F.dropout(att, p=self.dropout, training=self.training) - y = (att @ v).transpose(1, 2).contiguous().view(b, t, c) + batch_size, sequence_length, channels = x.shape + head_size = channels // self.n_head + q = self.q_proj(x).view( + batch_size, sequence_length, self.n_head, head_size + ).transpose(1, 2) + k = self.k_proj(x).view( + batch_size, sequence_length, self.n_head, head_size + ).transpose(1, 2) + v = self.v_proj(x).view( + batch_size, sequence_length, self.n_head, head_size + ).transpose(1, 2) + + if self.has_sdpa: + y = F.scaled_dot_product_attention( + q, + k, + v, + attn_mask=None, + dropout_p=self.dropout if self.training else 0.0, + is_causal=True, + ) + else: + attention = (q @ k.transpose(-2, -1)) / math.sqrt(head_size) + attention = attention.masked_fill( + self.causal_mask[:, :, :sequence_length, :sequence_length] == 0, + float("-inf"), + ) + attention = F.softmax(attention, dim=-1) + attention = F.dropout( + attention, + p=self.dropout, + training=self.training, + ) + y = attention @ v + + y = y.transpose(1, 2).contiguous().view( + batch_size, sequence_length, channels + ) return self.resid_dropout(self.out_proj(y)) + class MLP(nn.Module): def __init__(self, cfg: GPTConfig): super().__init__() @@ -48,20 +92,22 @@ def __init__(self, cfg: GPTConfig): self.dropout = nn.Dropout(cfg.dropout) def forward(self, x: torch.Tensor) -> torch.Tensor: - return self.dropout(self.proj(F.gelu(self.fc(x)))) + return self.dropout(self.proj(F.gelu(self.fc(x), approximate="tanh"))) + class Block(nn.Module): def __init__(self, cfg: GPTConfig): super().__init__() - self.ln1 = nn.LayerNorm(cfg.n_embd) + self.ln1 = nn.LayerNorm(cfg.n_embd, bias=cfg.bias) self.attn = CausalSelfAttention(cfg) - self.ln2 = nn.LayerNorm(cfg.n_embd) + self.ln2 = nn.LayerNorm(cfg.n_embd, bias=cfg.bias) self.mlp = MLP(cfg) def forward(self, x: torch.Tensor) -> torch.Tensor: x = x + self.attn(self.ln1(x)) return x + self.mlp(self.ln2(x)) + class GPT(nn.Module): def __init__(self, cfg: GPTConfig): super().__init__() @@ -70,27 +116,67 @@ def __init__(self, cfg: GPTConfig): self.position_embedding = nn.Embedding(cfg.block_size, cfg.n_embd) self.drop = nn.Dropout(cfg.dropout) self.blocks = nn.ModuleList([Block(cfg) for _ in range(cfg.n_layer)]) - self.ln_f = nn.LayerNorm(cfg.n_embd) + self.ln_f = nn.LayerNorm(cfg.n_embd, bias=cfg.bias) self.lm_head = nn.Linear(cfg.n_embd, cfg.vocab_size, bias=False) self.lm_head.weight = self.token_embedding.weight - self.apply(self._init) - def _init(self, module: nn.Module) -> None: - if isinstance(module, (nn.Linear, nn.Embedding)): + self.apply(self._init_weights) + residual_std = 0.02 / math.sqrt(2 * cfg.n_layer) + for name, parameter in self.named_parameters(): + if name.endswith("attn.out_proj.weight") or name.endswith( + "mlp.proj.weight" + ): + nn.init.normal_(parameter, mean=0.0, std=residual_std) + + @staticmethod + def _init_weights(module: nn.Module) -> None: + if isinstance(module, nn.Linear): + nn.init.normal_(module.weight, mean=0.0, std=0.02) + if module.bias is not None: + nn.init.zeros_(module.bias) + elif isinstance(module, nn.Embedding): nn.init.normal_(module.weight, mean=0.0, std=0.02) - if isinstance(module, nn.Linear) and module.bias is not None: - nn.init.zeros_(module.bias) - - def forward(self, idx: torch.Tensor, targets: torch.Tensor | None = None): - _, t = idx.shape - if t > self.cfg.block_size: - raise ValueError("sequence exceeds block_size") - pos = torch.arange(t, device=idx.device) - x = self.drop(self.token_embedding(idx) + self.position_embedding(pos)) + + def num_parameters(self, *, exclude_position_embeddings: bool = False) -> int: + count = sum(parameter.numel() for parameter in self.parameters()) + if exclude_position_embeddings: + count -= self.position_embedding.weight.numel() + return count + + def forward( + self, + idx: torch.Tensor, + targets: torch.Tensor | None = None, + ) -> tuple[torch.Tensor, torch.Tensor | None]: + _, sequence_length = idx.shape + if sequence_length > self.cfg.block_size: + raise ValueError( + f"sequence length {sequence_length} exceeds block_size " + f"{self.cfg.block_size}" + ) + positions = torch.arange(sequence_length, device=idx.device) + x = self.drop( + self.token_embedding(idx) + self.position_embedding(positions) + ) for block in self.blocks: x = block(x) logits = self.lm_head(self.ln_f(x)) loss = None if targets is not None: - loss = F.cross_entropy(logits.view(-1, logits.size(-1)), targets.view(-1)) + loss = F.cross_entropy( + logits.reshape(-1, logits.size(-1)), + targets.reshape(-1), + ) return logits, loss + + def spectral_matrices(self) -> Iterator[tuple[str, torch.Tensor]]: + """Yield only transformer block matrices used for alpha tracking.""" + + for index, block in enumerate(self.blocks): + prefix = f"L{index:02d}" + yield f"{prefix}_W_Q", block.attn.q_proj.weight + yield f"{prefix}_W_K", block.attn.k_proj.weight + yield f"{prefix}_W_V", block.attn.v_proj.weight + yield f"{prefix}_W_O", block.attn.out_proj.weight + yield f"{prefix}_W_MLP_IN", block.mlp.fc.weight + yield f"{prefix}_W_MLP_OUT", block.mlp.proj.weight diff --git a/level_0_baseline/src/level0_baseline/optim.py b/level_0_baseline/src/level0_baseline/optim.py index fbc19ae..bea409d 100644 --- a/level_0_baseline/src/level0_baseline/optim.py +++ b/level_0_baseline/src/level0_baseline/optim.py @@ -1,59 +1,159 @@ from __future__ import annotations + +import inspect +from typing import Iterable + import torch + @torch.no_grad() -def zeropower_via_newtonschulz5(g: torch.Tensor, steps: int = 5) -> torch.Tensor: - assert g.ndim == 2 - x = g.float() - if x.shape[0] > x.shape[1]: +def zeropower_via_newtonschulz5( + gradient: torch.Tensor, + steps: int = 5, +) -> torch.Tensor: + if gradient.ndim != 2: + raise ValueError("Newton-Schulz zero-power update requires a matrix") + x = gradient.float() + transposed = x.shape[0] > x.shape[1] + if transposed: x = x.T x = x / (x.norm() + 1e-7) a, b, c = 3.4445, -4.7750, 2.0315 for _ in range(steps): - A = x @ x.T - x = a * x + (b * A + c * (A @ A)) @ x - if g.shape[0] > g.shape[1]: + gram = x @ x.T + x = a * x + (b * gram + c * (gram @ gram)) @ x + if transposed: x = x.T - return x.to(g.dtype) + return x.to(gradient.dtype) + class Muon(torch.optim.Optimizer): - def __init__(self, params, lr=0.02, momentum=0.95, nesterov=True, weight_decay=0.0): - super().__init__(params, dict(lr=lr, momentum=momentum, nesterov=nesterov, weight_decay=weight_decay)) + def __init__( + self, + params: Iterable[torch.nn.Parameter], + *, + lr: float = 0.02, + momentum: float = 0.95, + nesterov: bool = True, + weight_decay: float = 0.0, + ): + defaults = { + "lr": lr, + "momentum": momentum, + "nesterov": nesterov, + "weight_decay": weight_decay, + } + super().__init__(params, defaults) @torch.no_grad() def step(self, closure=None): loss = closure() if closure is not None else None for group in self.param_groups: - for p in group["params"]: - if p.grad is None: + for parameter in group["params"]: + if parameter.grad is None: continue - if p.ndim != 2: + if parameter.ndim != 2: raise ValueError("Muon received a non-matrix parameter") - buf = self.state[p].setdefault("momentum_buffer", torch.zeros_like(p)) - buf.mul_(group["momentum"]).add_(p.grad) - g = p.grad.add(buf, alpha=group["momentum"]) if group["nesterov"] else buf - update = zeropower_via_newtonschulz5(g) - update.mul_(max(1, p.shape[0] / p.shape[1]) ** 0.5) + buffer = self.state[parameter].setdefault( + "momentum_buffer", + torch.zeros_like(parameter), + ) + buffer.mul_(group["momentum"]).add_(parameter.grad) + update_source = ( + parameter.grad.add(buffer, alpha=group["momentum"]) + if group["nesterov"] + else buffer + ) + update = zeropower_via_newtonschulz5(update_source) + update.mul_(max(1.0, parameter.shape[0] / parameter.shape[1]) ** 0.5) if group["weight_decay"]: - p.mul_(1 - group["lr"] * group["weight_decay"]) - p.add_(update, alpha=-group["lr"]) + parameter.mul_(1 - group["lr"] * group["weight_decay"]) + parameter.add_(update, alpha=-group["lr"]) return loss -def make_optimizers(model, cfg): - t = cfg["training"] - name = t["optimizer"].lower() - decay = [p for _, p in model.named_parameters() if p.requires_grad and p.ndim >= 2] - nodecay = [p for _, p in model.named_parameters() if p.requires_grad and p.ndim < 2] - if name == "adamw": - return [torch.optim.AdamW([{"params": decay, "weight_decay": t["weight_decay"]}, {"params": nodecay, "weight_decay": 0.0}], lr=t["learning_rate"], betas=(t["beta1"], t["beta2"]))] - if name != "muon": - raise ValueError(f"unsupported optimizer: {name}") - muon, adam = [], [] - for n, p in model.named_parameters(): - if not p.requires_grad: - continue - if p.ndim == 2 and "embedding" not in n and "lm_head" not in n: - muon.append(p) + +def _adamw( + groups: list[dict], + *, + learning_rate: float, + betas: tuple[float, float], + epsilon: float, + device_type: str, +) -> torch.optim.AdamW: + fused_available = "fused" in inspect.signature(torch.optim.AdamW).parameters + use_fused = fused_available and device_type == "cuda" + extra = {"fused": True} if use_fused else {} + return torch.optim.AdamW( + groups, + lr=learning_rate, + betas=betas, + eps=epsilon, + **extra, + ) + + +def make_optimizers( + model: torch.nn.Module, + cfg: dict, + *, + device_type: str = "cpu", +) -> list[torch.optim.Optimizer]: + training = cfg["training"] + optimizer_name = training["optimizer"].lower() + parameters = { + name: parameter + for name, parameter in model.named_parameters() + if parameter.requires_grad + } + decay = [parameter for parameter in parameters.values() if parameter.ndim >= 2] + no_decay = [parameter for parameter in parameters.values() if parameter.ndim < 2] + betas = (training["beta1"], training["beta2"]) + epsilon = float(training.get("epsilon", 1e-8)) + + if optimizer_name == "adamw": + groups = [ + {"params": decay, "weight_decay": training["weight_decay"]}, + {"params": no_decay, "weight_decay": 0.0}, + ] + return [ + _adamw( + groups, + learning_rate=training["learning_rate"], + betas=betas, + epsilon=epsilon, + device_type=device_type, + ) + ] + + if optimizer_name != "muon": + raise ValueError(f"unsupported optimizer: {optimizer_name}") + + muon_parameters: list[torch.nn.Parameter] = [] + auxiliary_parameters: list[torch.nn.Parameter] = [] + for name, parameter in parameters.items(): + if ( + parameter.ndim == 2 + and "token_embedding" not in name + and "position_embedding" not in name + and "lm_head" not in name + ): + muon_parameters.append(parameter) else: - adam.append(p) - return [Muon(muon, lr=t["muon_learning_rate"], momentum=t["muon_momentum"], nesterov=t["muon_nesterov"], weight_decay=t["weight_decay"]), torch.optim.AdamW(adam, lr=t["muon_aux_adamw_learning_rate"], betas=(t["beta1"], t["beta2"]), weight_decay=0.0)] + auxiliary_parameters.append(parameter) + + return [ + Muon( + muon_parameters, + lr=training["muon_learning_rate"], + momentum=training["muon_momentum"], + nesterov=training["muon_nesterov"], + weight_decay=training["weight_decay"], + ), + _adamw( + [{"params": auxiliary_parameters, "weight_decay": 0.0}], + learning_rate=training["muon_aux_adamw_learning_rate"], + betas=betas, + epsilon=epsilon, + device_type=device_type, + ), + ] diff --git a/level_0_baseline/src/level0_baseline/train.py b/level_0_baseline/src/level0_baseline/train.py index 2123405..732bc90 100644 --- a/level_0_baseline/src/level0_baseline/train.py +++ b/level_0_baseline/src/level0_baseline/train.py @@ -1,126 +1,973 @@ from __future__ import annotations -import argparse, csv, json, math, platform, random, time + +import argparse +import csv +import hashlib +import json +import math +import os +import platform +import random +import shutil +import sys +import time +from dataclasses import asdict from pathlib import Path +from typing import Any, Iterable + import numpy as np import torch -from .config import load_config, roots +import torch.nn as nn + +from .config import load_config, roots, validate_config from .model import GPT, GPTConfig from .optim import make_optimizers -def device_auto(): +PROTOCOL_VERSION = "isolated_level0_bpe_v2" +METRIC_FIELDS = [ + "step", + "tokens_seen", + "elapsed_sec", + "learning_rate", + "train_loss", + "train_perplexity", + "train_bits_per_token", + "train_accuracy", + "val_loss", + "val_perplexity", + "val_bits_per_token", + "val_accuracy", + "val_generalization_gap", + "grad_norm", + "weight_norm", + "tokens_per_second", +] + + +def device_auto() -> torch.device: if torch.cuda.is_available(): return torch.device("cuda") if torch.backends.mps.is_available(): return torch.device("mps") return torch.device("cpu") -def seed_all(seed): + +def select_device(requested: str) -> torch.device: + if requested == "auto": + return device_auto() + device = torch.device(requested) + if device.type == "cuda" and not torch.cuda.is_available(): + raise RuntimeError("CUDA was requested but is unavailable") + if device.type == "mps" and not torch.backends.mps.is_available(): + raise RuntimeError("MPS was requested but is unavailable") + return device + + +def synchronize(device: torch.device) -> None: + if device.type == "cuda": + torch.cuda.synchronize(device) + elif device.type == "mps" and hasattr(torch, "mps"): + torch.mps.synchronize() + + +def seed_all(seed: int) -> None: random.seed(seed) np.random.seed(seed) torch.manual_seed(seed) -def lr_at(step, t): - if step < t["warmup_steps"]: - return t["learning_rate"] * (step + 1) / max(1, t["warmup_steps"]) - if step >= t["max_steps"]: - return t["min_lr"] - ratio = (step - t["warmup_steps"]) / max(1, t["max_steps"] - t["warmup_steps"]) - return t["min_lr"] + 0.5 * (1 + math.cos(math.pi * ratio)) * (t["learning_rate"] - t["min_lr"]) -def batch(data, batch_size, block_size, device, generator): - ix = torch.randint(len(data) - block_size - 1, (batch_size,), generator=generator) - x = torch.stack([torch.from_numpy(np.array(data[i:i+block_size], dtype=np.int64)) for i in ix]) - y = torch.stack([torch.from_numpy(np.array(data[i+1:i+1+block_size], dtype=np.int64)) for i in ix]) - return x.to(device), y.to(device) +def stable_seed(seed: int, label: str) -> int: + digest = hashlib.sha256(f"{seed}:{label}".encode()).digest() + return int.from_bytes(digest[:8], "little") & ((1 << 63) - 1) + + +def config_sha256(cfg: dict[str, Any]) -> str: + payload = json.dumps(cfg, sort_keys=True, separators=(",", ":")).encode() + return hashlib.sha256(payload).hexdigest() + + +def lr_at(update_index: int, training: dict[str, Any]) -> float: + """Warmup followed by cosine decay, indexed by zero-based update number.""" + + if update_index < training["warmup_steps"]: + return training["learning_rate"] * (update_index + 1) / max( + 1, training["warmup_steps"] + ) + if update_index >= training["max_steps"]: + return training["min_lr"] + ratio = (update_index - training["warmup_steps"]) / max( + 1, + training["max_steps"] - training["warmup_steps"], + ) + coefficient = 0.5 * (1.0 + math.cos(math.pi * ratio)) + return training["min_lr"] + coefficient * ( + training["learning_rate"] - training["min_lr"] + ) + + +def _batch_from_starts( + data: np.memmap, + starts: Iterable[int], + block_size: int, + device: torch.device, +) -> tuple[torch.Tensor, torch.Tensor]: + starts_list = [int(value) for value in starts] + x = np.stack( + [np.asarray(data[start : start + block_size], dtype=np.int64) for start in starts_list] + ) + y = np.stack( + [ + np.asarray( + data[start + 1 : start + 1 + block_size], + dtype=np.int64, + ) + for start in starts_list + ] + ) + return torch.from_numpy(x).to(device), torch.from_numpy(y).to(device) + + +def random_batch( + data: np.memmap, + batch_size: int, + block_size: int, + device: torch.device, + generator: torch.Generator, +) -> tuple[torch.Tensor, torch.Tensor]: + upper = len(data) - block_size - 1 + if upper <= 0: + raise ValueError("token split is shorter than block_size + 1") + starts = torch.randint(upper, (batch_size,), generator=generator).tolist() + return _batch_from_starts(data, starts, block_size, device) + + +def fixed_eval_starts( + data_length: int, + *, + batch_size: int, + block_size: int, + eval_batches: int, + seed: int, +) -> np.ndarray: + upper = data_length - block_size - 1 + if upper <= 0: + raise ValueError("token split is shorter than block_size + 1") + generator = np.random.default_rng(seed) + return generator.integers( + 0, + upper, + size=(eval_batches, batch_size), + endpoint=False, + dtype=np.int64, + ) + @torch.no_grad() -def evaluate(model, data, batch_size, block_size, n_batches, device, generator): +def evaluate( + model: GPT, + data: np.memmap, + starts: np.ndarray, + block_size: int, + device: torch.device, +) -> dict[str, float]: + was_training = model.training model.eval() - losses, correct, total = [], 0, 0 - for _ in range(n_batches): - x, y = batch(data, batch_size, block_size, device, generator) + losses: list[float] = [] + correct = 0 + total = 0 + for batch_starts in starts: + x, y = _batch_from_starts(data, batch_starts, block_size, device) logits, loss = model(x, y) - losses.append(loss.item()) - correct += (logits.argmax(-1) == y).sum().item() + if loss is None: + raise RuntimeError("evaluation did not produce a loss") + losses.append(float(loss.detach().cpu())) + correct += int((logits.argmax(dim=-1) == y).sum().item()) total += y.numel() - model.train() - loss = float(np.mean(losses)) - return loss, math.exp(min(20, loss)), correct / total + model.train(was_training) + mean_loss = float(np.mean(losses)) + return { + "loss": mean_loss, + "perplexity": math.exp(min(20.0, mean_loss)), + "bits_per_token": mean_loss / math.log(2.0), + "accuracy": correct / max(1, total), + } -def weightwatch(model, out, step, randomize): + +def weight_norm(model: nn.Module) -> float: + return math.sqrt( + sum( + float((parameter.detach().float() ** 2).sum().cpu()) + for parameter in model.parameters() + ) + ) + + +class _MatrixProbe(nn.Module): + def __init__(self, matrices: list[tuple[str, torch.Tensor]]): + super().__init__() + for name, matrix in matrices: + layer = nn.Linear(matrix.shape[1], matrix.shape[0], bias=False) + layer.weight = nn.Parameter( + matrix.detach().float().cpu().clone(), + requires_grad=False, + ) + self.add_module(name, layer) + + +def run_weightwatcher( + model: GPT, + run_dir: Path, + *, + step: int, + tokens_seen: int, +) -> tuple[bool, str]: try: import weightwatcher as ww except ImportError: - return - df = ww.WeightWatcher(model=model).analyze(randomize=randomize) - df.insert(0, "step", step) - df.to_csv(out / f"weightwatcher_step_{step:07d}.csv", index=False) - -def main(): - p = argparse.ArgumentParser() - p.add_argument("--config", default="configs/level0.yaml") - p.add_argument("--data-root") - p.add_argument("--results-root") - p.add_argument("--optimizer", choices=["adamw", "muon"]) - p.add_argument("--seed", type=int) - p.add_argument("--device", default="auto") - a = p.parse_args() - cfg = load_config(a.config) - t = cfg["training"] - if a.optimizer: - t["optimizer"] = a.optimizer - if a.seed is not None: - t["seed"] = a.seed - resolved = roots() - data_root = Path(a.data_root or resolved["data"]) - base = Path(a.results_root or resolved["results"]) - run = base / f"{t['optimizer']}_seed_{t['seed']}" - run.mkdir(parents=True, exist_ok=True) - device = device_auto() if a.device == "auto" else torch.device(a.device) - seed_all(t["seed"]) - generator = torch.Generator().manual_seed(t["seed"]) - arrays = {s: np.memmap(data_root / f"{s}.bin", dtype=np.uint8, mode="r") for s in ("train", "val", "test")} - model = GPT(GPTConfig(**cfg["model"])).to(device) - optimizers = make_optimizers(model, cfg) - base_lrs = [group["lr"] for opt in optimizers for group in opt.param_groups] - manifest = {"config": cfg, "device": str(device), "torch": torch.__version__, "platform": platform.platform(), "parameter_count": sum(p.numel() for p in model.parameters()), "data_root": str(data_root.resolve())} - (run / "manifest.json").write_text(json.dumps(manifest, indent=2)) - fields = ["step", "tokens_seen", "elapsed_sec", "learning_rate", "train_loss", "train_perplexity", "train_accuracy", "val_loss", "val_perplexity", "val_accuracy", "test_loss", "test_perplexity", "test_accuracy", "val_generalization_gap", "test_generalization_gap", "grad_norm", "weight_norm"] - with open(run / "metrics.csv", "w", newline="") as f: - writer = csv.DictWriter(f, fieldnames=fields) + return False, "WeightWatcher is not installed" + + matrices = list(model.spectral_matrices()) + probe = _MatrixProbe(matrices) + try: + details = ww.WeightWatcher(model=probe).analyze(randomize=False) + except Exception as exc: # diagnostic failure must not destroy training + message = f"{type(exc).__name__}: {exc}" + (run_dir / f"weightwatcher_error_step_{step:07d}.json").write_text( + json.dumps({"step": step, "error": message}, indent=2), + encoding="utf-8", + ) + return False, message + + details.insert(0, "tokens_seen", tokens_seen) + details.insert(0, "step", step) + source_column = next( + (name for name in ("longname", "name") if name in details.columns), + None, + ) + matrix_names = [name for name, _ in matrices] + if source_column is not None: + details["matrix_name"] = details[source_column].astype(str).map( + lambda value: next( + (name for name in matrix_names if name in value), + value, + ) + ) + else: + details["matrix_name"] = "unknown" + + output = run_dir / f"weightwatcher_step_{step:07d}.csv" + details.to_csv(output, index=False) + return True, "" + + +def _write_metrics(path: Path, rows: list[dict[str, Any]]) -> None: + temporary = path.with_suffix(".csv.tmp") + with open(temporary, "w", newline="", encoding="utf-8") as handle: + writer = csv.DictWriter(handle, fieldnames=METRIC_FIELDS) writer.writeheader() - start = time.time() - last_grad_norm = float("nan") - for step in range(t["max_steps"] + 1): - lr = lr_at(step, t) - scale = lr / t["learning_rate"] - j = 0 - for opt in optimizers: - for group in opt.param_groups: - group["lr"] = base_lrs[j] * scale - j += 1 - if step % t["eval_interval"] == 0 or step == t["max_steps"]: - values = {s: evaluate(model, arrays[s], t["batch_size"], cfg["model"]["block_size"], t["eval_batches"], device, generator) for s in ("train", "val", "test")} - weight_norm = math.sqrt(sum(float((p.detach().float() ** 2).sum()) for p in model.parameters())) - writer.writerow({"step": step, "tokens_seen": step * t["batch_size"] * cfg["model"]["block_size"] * t["grad_accum_steps"], "elapsed_sec": time.time() - start, "learning_rate": lr, "train_loss": values["train"][0], "train_perplexity": values["train"][1], "train_accuracy": values["train"][2], "val_loss": values["val"][0], "val_perplexity": values["val"][1], "val_accuracy": values["val"][2], "test_loss": values["test"][0], "test_perplexity": values["test"][1], "test_accuracy": values["test"][2], "val_generalization_gap": values["val"][0] - values["train"][0], "test_generalization_gap": values["test"][0] - values["train"][0], "grad_norm": last_grad_norm, "weight_norm": weight_norm}) - f.flush() - print(step, {k: round(v[0], 4) for k, v in values.items()}) - if cfg["analysis"]["weightwatcher"] and step % cfg["analysis"]["weightwatcher_interval"] == 0: - weightwatch(model, run, step, cfg["analysis"]["randomize"]) - if step == t["max_steps"]: - break - for opt in optimizers: - opt.zero_grad(set_to_none=True) - for _ in range(t["grad_accum_steps"]): - x, y = batch(arrays["train"], t["batch_size"], cfg["model"]["block_size"], device, generator) - _, loss = model(x, y) - (loss / t["grad_accum_steps"]).backward() - last_grad_norm = float(torch.nn.utils.clip_grad_norm_(model.parameters(), t["grad_clip"])) - for opt in optimizers: - opt.step() - if (step + 1) % t["checkpoint_interval"] == 0: - torch.save({"model": model.state_dict(), "step": step + 1, "config": cfg}, run / f"checkpoint_{step+1:07d}.pt") - torch.save({"model": model.state_dict(), "step": t["max_steps"], "config": cfg}, run / "checkpoint_final.pt") + writer.writerows(rows) + handle.flush() + os.fsync(handle.fileno()) + temporary.replace(path) + + + +def _torch_load(path: Path, *, map_location): + try: + return torch.load(path, map_location=map_location, weights_only=False) + except TypeError: + return torch.load(path, map_location=map_location) + + +def _atomic_torch_save(payload: dict[str, Any], path: Path) -> None: + temporary = path.with_suffix(path.suffix + ".tmp") + torch.save(payload, temporary) + temporary.replace(path) + + +def _optimizer_state_dicts( + optimizers: list[torch.optim.Optimizer], +) -> list[dict[str, Any]]: + return [optimizer.state_dict() for optimizer in optimizers] + + +def _load_optimizer_state_dicts( + optimizers: list[torch.optim.Optimizer], + states: list[dict[str, Any]], +) -> None: + if len(optimizers) != len(states): + raise ValueError("checkpoint optimizer count does not match configuration") + for optimizer, state in zip(optimizers, states, strict=True): + optimizer.load_state_dict(state) + + +def _checkpoint_payload( + *, + model: GPT, + optimizers: list[torch.optim.Optimizer], + step: int, + cfg: dict[str, Any], + config_hash: str, + train_generator: torch.Generator, + metrics_rows: list[dict[str, Any]], + best_validation_loss: float, + best_validation_step: int, + best_validation_row: dict[str, Any], + elapsed_sec: float, + base_lrs: list[float], + weightwatcher_successes: int, + weightwatcher_failures: int, +) -> dict[str, Any]: + payload: dict[str, Any] = { + "protocol_version": PROTOCOL_VERSION, + "config": cfg, + "config_sha256": config_hash, + "model": model.state_dict(), + "optimizers": _optimizer_state_dicts(optimizers), + "step": int(step), + "train_generator_state": train_generator.get_state(), + "torch_rng_state": torch.get_rng_state(), + "metrics_rows": metrics_rows, + "best_validation_loss": float(best_validation_loss), + "best_validation_step": int(best_validation_step), + "best_validation_row": best_validation_row, + "elapsed_sec": float(elapsed_sec), + "base_lrs": base_lrs, + "weightwatcher_successes": int(weightwatcher_successes), + "weightwatcher_failures": int(weightwatcher_failures), + } + if torch.cuda.is_available(): + payload["cuda_rng_state_all"] = torch.cuda.get_rng_state_all() + if ( + torch.backends.mps.is_available() + and hasattr(torch, "mps") + and hasattr(torch.mps, "get_rng_state") + ): + payload["mps_rng_state"] = torch.mps.get_rng_state() + return payload + + +def _restore_rng_state(payload: dict[str, Any]) -> None: + if "torch_rng_state" in payload: + torch.set_rng_state(payload["torch_rng_state"]) + if "cuda_rng_state_all" in payload and torch.cuda.is_available(): + torch.cuda.set_rng_state_all(payload["cuda_rng_state_all"]) + if ( + "mps_rng_state" in payload + and torch.backends.mps.is_available() + and hasattr(torch, "mps") + and hasattr(torch.mps, "set_rng_state") + ): + torch.mps.set_rng_state(payload["mps_rng_state"]) + + +def _load_data( + data_root: Path, + cfg: dict[str, Any], +) -> tuple[dict[str, np.memmap], dict[str, Any]]: + metadata_path = data_root / "meta.json" + if not metadata_path.is_file(): + raise FileNotFoundError( + f"missing {metadata_path}; run level0-prepare-data first" + ) + metadata = json.loads(metadata_path.read_text(encoding="utf-8")) + if int(metadata.get("format_version", 0)) < 2: + raise RuntimeError( + "prepared data is the obsolete byte-token baseline; prepare the new " + "GPT-2-BPE dataset under /tmp/nanogpt-level0-bpe/data" + ) + if not str(metadata.get("tokenizer", "")).startswith("tiktoken:"): + raise RuntimeError("Level 0 requires tiktoken BPE data") + expected_vocab = int(cfg["model"]["vocab_size"]) + if int(metadata.get("model_vocab_size", -1)) != expected_vocab: + raise RuntimeError( + "prepared data model_vocab_size does not match the model configuration" + ) + dtype = np.dtype(metadata["dtype"]) + arrays: dict[str, np.memmap] = {} + for split in ("train", "val", "test"): + path = data_root / f"{split}.bin" + if not path.is_file(): + raise FileNotFoundError(f"missing prepared split: {path}") + arrays[split] = np.memmap(path, dtype=dtype, mode="r") + expected_tokens = int(metadata["split_tokens"][split]) + if len(arrays[split]) != expected_tokens: + raise RuntimeError( + f"{split}.bin has {len(arrays[split]):,} tokens; " + f"metadata requires {expected_tokens:,}" + ) + return arrays, metadata + + +def _log(run_dir: Path, message: str) -> None: + line = f"[level0-train] {message}" + print(line, file=sys.stderr, flush=True) + with open(run_dir / "train.log", "a", encoding="utf-8") as handle: + handle.write(line + "\n") + handle.flush() + + +def _record_evaluation( + *, + model: GPT, + arrays: dict[str, np.memmap], + eval_starts: dict[str, np.ndarray], + cfg: dict[str, Any], + device: torch.device, + step: int, + tokens_seen: int, + elapsed_sec: float, + learning_rate: float, + grad_norm: float, +) -> dict[str, Any]: + block_size = cfg["model"]["block_size"] + train_metrics = evaluate( + model, + arrays["train"], + eval_starts["train"], + block_size, + device, + ) + validation_metrics = evaluate( + model, + arrays["val"], + eval_starts["val"], + block_size, + device, + ) + return { + "step": step, + "tokens_seen": tokens_seen, + "elapsed_sec": elapsed_sec, + "learning_rate": learning_rate, + "train_loss": train_metrics["loss"], + "train_perplexity": train_metrics["perplexity"], + "train_bits_per_token": train_metrics["bits_per_token"], + "train_accuracy": train_metrics["accuracy"], + "val_loss": validation_metrics["loss"], + "val_perplexity": validation_metrics["perplexity"], + "val_bits_per_token": validation_metrics["bits_per_token"], + "val_accuracy": validation_metrics["accuracy"], + "val_generalization_gap": validation_metrics["loss"] + - train_metrics["loss"], + "grad_norm": grad_norm, + "weight_norm": weight_norm(model), + "tokens_per_second": tokens_seen / max(elapsed_sec, 1e-9), + } + + +def _evaluate_test( + model: GPT, + arrays: dict[str, np.memmap], + eval_starts: dict[str, np.ndarray], + cfg: dict[str, Any], + device: torch.device, +) -> dict[str, float]: + return evaluate( + model, + arrays["test"], + eval_starts["test"], + cfg["model"]["block_size"], + device, + ) + + +def run_experiment( + cfg: dict[str, Any], + *, + data_root: Path, + results_root: Path, + device: torch.device, + resume: bool = False, + overwrite: bool = False, + disable_weightwatcher: bool = False, +) -> Path: + validate_config(cfg) + if resume and overwrite: + raise ValueError("resume and overwrite are mutually exclusive") + + training = cfg["training"] + analysis = cfg["analysis"] + run_dir = results_root / f"{training['optimizer']}_seed_{training['seed']}" + if overwrite and run_dir.exists(): + shutil.rmtree(run_dir) + run_dir.mkdir(parents=True, exist_ok=True) + + complete_path = run_dir / "run_complete.json" + if complete_path.exists(): + if resume: + _log(run_dir, "run is already complete; nothing to resume") + return run_dir + raise FileExistsError( + f"completed run already exists: {run_dir}; use --overwrite explicitly" + ) + + existing_files = [path for path in run_dir.iterdir() if path.name != "train.log"] + if existing_files and not resume: + raise FileExistsError( + f"nonempty run directory exists: {run_dir}; use --resume or --overwrite" + ) + + arrays, data_metadata = _load_data(data_root, cfg) + seed = int(training["seed"]) + seed_all(seed) + train_generator = torch.Generator(device="cpu") + train_generator.manual_seed(stable_seed(seed, "training_windows")) + + model_config = GPTConfig(**cfg["model"]) + model = GPT(model_config).to(device) + if training.get("compile", False): + if device.type != "cuda": + raise RuntimeError("torch.compile is only enabled for the CUDA preset") + model = torch.compile(model) # type: ignore[assignment] + + optimizers = make_optimizers(model, cfg, device_type=device.type) + base_lrs = [ + float(group["lr"]) + for optimizer in optimizers + for group in optimizer.param_groups + ] + config_hash = config_sha256(cfg) + parameter_count = sum(parameter.numel() for parameter in model.parameters()) + tokens_per_step = ( + training["batch_size"] + * cfg["model"]["block_size"] + * training["grad_accum_steps"] + ) + + eval_starts = { + split: fixed_eval_starts( + len(arrays[split]), + batch_size=training["batch_size"], + block_size=cfg["model"]["block_size"], + eval_batches=training["eval_batches"], + seed=stable_seed(seed, f"fixed_{split}_probe"), + ) + for split in ("train", "val", "test") + } + eval_probe_hashes = { + split: hashlib.sha256(starts.tobytes()).hexdigest() + for split, starts in eval_starts.items() + } + + manifest = { + "protocol_version": PROTOCOL_VERSION, + "config": cfg, + "config_sha256": config_hash, + "device": str(device), + "torch_version": torch.__version__, + "platform": platform.platform(), + "parameter_count": parameter_count, + "model_config": asdict(model_config), + "tokens_per_optimizer_step": tokens_per_step, + "planned_train_tokens": tokens_per_step * training["max_steps"], + "data_root": str(data_root.resolve()), + "data_metadata": data_metadata, + "fixed_probe_hashes": eval_probe_hashes, + "test_policy": "final_and_validation_selected_checkpoint_only", + "alpha_policy": "deterministic_nonrandomized_transformer_matrices", + } + manifest_path = run_dir / "manifest.json" + if manifest_path.exists() and resume: + existing_manifest = json.loads(manifest_path.read_text(encoding="utf-8")) + if existing_manifest.get("config_sha256") != config_hash: + raise RuntimeError("resume configuration does not match the existing run") + else: + manifest_path.write_text( + json.dumps(manifest, indent=2, sort_keys=True), + encoding="utf-8", + ) + + metrics_rows: list[dict[str, Any]] = [] + best_validation_loss = float("inf") + best_validation_step = 0 + best_validation_row: dict[str, Any] = {} + weightwatcher_successes = 0 + weightwatcher_failures = 0 + elapsed_prior = 0.0 + start_step = 1 + last_grad_norm = 0.0 + latest_path = run_dir / "checkpoint_latest.pt" + + if resume: + if not latest_path.is_file(): + raise FileNotFoundError( + f"resume requested but checkpoint is missing: {latest_path}" + ) + checkpoint = _torch_load(latest_path, map_location=device) + if checkpoint.get("config_sha256") != config_hash: + raise RuntimeError("checkpoint configuration does not match") + model.load_state_dict(checkpoint["model"]) + _load_optimizer_state_dicts(optimizers, checkpoint["optimizers"]) + train_generator.set_state(checkpoint["train_generator_state"]) + _restore_rng_state(checkpoint) + metrics_rows = list(checkpoint.get("metrics_rows", [])) + best_validation_loss = float(checkpoint["best_validation_loss"]) + best_validation_step = int(checkpoint["best_validation_step"]) + best_validation_row = dict(checkpoint["best_validation_row"]) + elapsed_prior = float(checkpoint.get("elapsed_sec", 0.0)) + base_lrs = [float(value) for value in checkpoint.get("base_lrs", base_lrs)] + weightwatcher_successes = int( + checkpoint.get("weightwatcher_successes", 0) + ) + weightwatcher_failures = int( + checkpoint.get("weightwatcher_failures", 0) + ) + start_step = int(checkpoint["step"]) + 1 + _log(run_dir, f"resuming from optimizer step {start_step - 1}") + else: + start_time = time.monotonic() + initial_row = _record_evaluation( + model=model, + arrays=arrays, + eval_starts=eval_starts, + cfg=cfg, + device=device, + step=0, + tokens_seen=0, + elapsed_sec=0.0, + learning_rate=0.0, + grad_norm=0.0, + ) + metrics_rows.append(initial_row) + best_validation_loss = float(initial_row["val_loss"]) + best_validation_step = 0 + best_validation_row = dict(initial_row) + _write_metrics(run_dir / "metrics.csv", metrics_rows) + _atomic_torch_save( + { + "protocol_version": PROTOCOL_VERSION, + "config_sha256": config_hash, + "model": model.state_dict(), + "step": 0, + "validation_row": initial_row, + }, + run_dir / "checkpoint_best.pt", + ) + if analysis["weightwatcher"] and not disable_weightwatcher: + success, error = run_weightwatcher(model, run_dir, step=0, tokens_seen=0) + if success: + weightwatcher_successes += 1 + else: + weightwatcher_failures += 1 + _log(run_dir, f"WeightWatcher step 0 failed: {error}") + elapsed_prior = time.monotonic() - start_time + _atomic_torch_save( + _checkpoint_payload( + model=model, + optimizers=optimizers, + step=0, + cfg=cfg, + config_hash=config_hash, + train_generator=train_generator, + metrics_rows=metrics_rows, + best_validation_loss=best_validation_loss, + best_validation_step=best_validation_step, + best_validation_row=best_validation_row, + elapsed_sec=elapsed_prior, + base_lrs=base_lrs, + weightwatcher_successes=weightwatcher_successes, + weightwatcher_failures=weightwatcher_failures, + ), + latest_path, + ) + + _log( + run_dir, + "starting " + f"optimizer={training['optimizer']} seed={seed} device={device} " + f"parameters={parameter_count:,} steps={training['max_steps']:,} " + f"tokens_per_step={tokens_per_step:,} " + f"planned_tokens={tokens_per_step * training['max_steps']:,}", + ) + + started_at = time.monotonic() + for step in range(start_step, training["max_steps"] + 1): + update_index = step - 1 + learning_rate = lr_at(update_index, training) + scale = learning_rate / training["learning_rate"] + group_index = 0 + for optimizer in optimizers: + for group in optimizer.param_groups: + group["lr"] = base_lrs[group_index] * scale + group_index += 1 + optimizer.zero_grad(set_to_none=True) + + minibatch_loss = 0.0 + for _ in range(training["grad_accum_steps"]): + x, y = random_batch( + arrays["train"], + training["batch_size"], + cfg["model"]["block_size"], + device, + train_generator, + ) + _, loss = model(x, y) + if loss is None: + raise RuntimeError("training forward pass did not produce a loss") + (loss / training["grad_accum_steps"]).backward() + minibatch_loss += float(loss.detach().cpu()) / training[ + "grad_accum_steps" + ] + + if training["grad_clip"] > 0: + gradient_norm = torch.nn.utils.clip_grad_norm_( + model.parameters(), + training["grad_clip"], + ) + last_grad_norm = float(gradient_norm.detach().cpu()) + else: + gradients = [ + parameter.grad.detach().float().norm() + for parameter in model.parameters() + if parameter.grad is not None + ] + last_grad_norm = ( + float(torch.linalg.vector_norm(torch.stack(gradients)).cpu()) + if gradients + else 0.0 + ) + + for optimizer in optimizers: + optimizer.step() + synchronize(device) + + elapsed = elapsed_prior + time.monotonic() - started_at + tokens_seen = step * tokens_per_step + if step % training["log_interval"] == 0 or step == 1: + rate = tokens_seen / max(elapsed, 1e-9) + remaining_steps = training["max_steps"] - step + eta = remaining_steps * tokens_per_step / max(rate, 1e-9) + _log( + run_dir, + f"step={step:,}/{training['max_steps']:,} " + f"loss={minibatch_loss:.4f} lr={learning_rate:.3e} " + f"grad={last_grad_norm:.3f} tok/s={rate:,.0f} " + f"eta_min={eta / 60:.1f}", + ) + + evaluation_due = ( + step % training["eval_interval"] == 0 + or step == training["max_steps"] + ) + if evaluation_due: + row = _record_evaluation( + model=model, + arrays=arrays, + eval_starts=eval_starts, + cfg=cfg, + device=device, + step=step, + tokens_seen=tokens_seen, + elapsed_sec=elapsed, + learning_rate=learning_rate, + grad_norm=last_grad_norm, + ) + metrics_rows.append(row) + _write_metrics(run_dir / "metrics.csv", metrics_rows) + _log( + run_dir, + f"eval step={step:,} train_loss={row['train_loss']:.4f} " + f"val_loss={row['val_loss']:.4f} " + f"val_ppl={row['val_perplexity']:.2f} " + f"val_acc={100 * row['val_accuracy']:.2f}%", + ) + if float(row["val_loss"]) < best_validation_loss: + best_validation_loss = float(row["val_loss"]) + best_validation_step = step + best_validation_row = dict(row) + _atomic_torch_save( + { + "protocol_version": PROTOCOL_VERSION, + "config_sha256": config_hash, + "model": model.state_dict(), + "step": step, + "validation_row": row, + }, + run_dir / "checkpoint_best.pt", + ) + + weightwatcher_due = ( + analysis["weightwatcher"] + and not disable_weightwatcher + and ( + step % analysis["weightwatcher_interval"] == 0 + or step == training["max_steps"] + ) + ) + if weightwatcher_due: + success, error = run_weightwatcher( + model, + run_dir, + step=step, + tokens_seen=tokens_seen, + ) + if success: + weightwatcher_successes += 1 + else: + weightwatcher_failures += 1 + _log(run_dir, f"WeightWatcher step {step} failed: {error}") + + checkpoint_due = ( + step % training["checkpoint_interval"] == 0 + or step == training["max_steps"] + ) + if checkpoint_due: + payload = _checkpoint_payload( + model=model, + optimizers=optimizers, + step=step, + cfg=cfg, + config_hash=config_hash, + train_generator=train_generator, + metrics_rows=metrics_rows, + best_validation_loss=best_validation_loss, + best_validation_step=best_validation_step, + best_validation_row=best_validation_row, + elapsed_sec=elapsed, + base_lrs=base_lrs, + weightwatcher_successes=weightwatcher_successes, + weightwatcher_failures=weightwatcher_failures, + ) + _atomic_torch_save(payload, latest_path) + _atomic_torch_save( + payload, + run_dir / f"checkpoint_{step:07d}.pt", + ) + + final_step = training["max_steps"] + final_tokens = final_step * tokens_per_step + final_test = _evaluate_test(model, arrays, eval_starts, cfg, device) + final_validation_row = metrics_rows[-1] + final_metrics = { + "checkpoint": "final", + "step": final_step, + "tokens_seen": final_tokens, + "train_loss": final_validation_row["train_loss"], + "validation_loss": final_validation_row["val_loss"], + "test_loss": final_test["loss"], + "test_perplexity": final_test["perplexity"], + "test_bits_per_token": final_test["bits_per_token"], + "test_accuracy": final_test["accuracy"], + "test_generalization_gap": final_test["loss"] + - final_validation_row["train_loss"], + } + (run_dir / "final_metrics.json").write_text( + json.dumps(final_metrics, indent=2, sort_keys=True), + encoding="utf-8", + ) + + best_checkpoint = _torch_load( + run_dir / "checkpoint_best.pt", map_location=device + ) + if best_validation_step == final_step: + selected_test = final_test + else: + selected_model = GPT(model_config).to(device) + selected_model.load_state_dict(best_checkpoint["model"]) + selected_test = _evaluate_test( + selected_model, + arrays, + eval_starts, + cfg, + device, + ) + del selected_model + selected_metrics = { + "checkpoint": "validation_selected", + "selected_step": best_validation_step, + "selected_tokens_seen": best_validation_step * tokens_per_step, + "train_loss": best_validation_row["train_loss"], + "validation_loss": best_validation_row["val_loss"], + "test_loss": selected_test["loss"], + "test_perplexity": selected_test["perplexity"], + "test_bits_per_token": selected_test["bits_per_token"], + "test_accuracy": selected_test["accuracy"], + "test_generalization_gap": selected_test["loss"] + - best_validation_row["train_loss"], + } + (run_dir / "selected_checkpoint_metrics.json").write_text( + json.dumps(selected_metrics, indent=2, sort_keys=True), + encoding="utf-8", + ) + + final_checkpoint = { + "protocol_version": PROTOCOL_VERSION, + "config_sha256": config_hash, + "model": model.state_dict(), + "step": final_step, + "final_metrics": final_metrics, + } + _atomic_torch_save(final_checkpoint, run_dir / "checkpoint_final.pt") + + completion = { + "status": "complete", + "protocol_version": PROTOCOL_VERSION, + "optimizer": training["optimizer"], + "seed": seed, + "final_step": final_step, + "tokens_seen": final_tokens, + "best_validation_step": best_validation_step, + "best_validation_loss": best_validation_loss, + "final_validation_loss": final_validation_row["val_loss"], + "final_test_loss": final_test["loss"], + "selected_checkpoint_test_loss": selected_test["loss"], + "weightwatcher_successes": weightwatcher_successes, + "weightwatcher_failures": weightwatcher_failures, + "test_evaluation_policy": "final_and_validation_selected_checkpoint_only", + } + complete_path.write_text( + json.dumps(completion, indent=2, sort_keys=True), + encoding="utf-8", + ) + _log( + run_dir, + f"complete final_val={final_validation_row['val_loss']:.4f} " + f"final_test={final_test['loss']:.4f} " + f"selected_step={best_validation_step:,} " + f"selected_test={selected_test['loss']:.4f}", + ) + return run_dir + + +def main() -> None: + parser = argparse.ArgumentParser( + description="Train the isolated realistic Level 0 nanoGPT baseline" + ) + parser.add_argument("--config", default="configs/level0.yaml") + parser.add_argument("--data-root") + parser.add_argument("--results-root") + parser.add_argument("--optimizer", choices=["adamw", "muon"]) + parser.add_argument("--seed", type=int) + parser.add_argument("--device", default="auto") + parser.add_argument("--resume", action="store_true") + parser.add_argument("--overwrite", action="store_true") + parser.add_argument("--no-weightwatcher", action="store_true") + args = parser.parse_args() + + cfg = load_config(args.config) + if args.optimizer: + cfg["training"]["optimizer"] = args.optimizer + if args.seed is not None: + cfg["training"]["seed"] = args.seed + validate_config(cfg) + + resolved = roots() + data_root = Path(args.data_root or resolved["data"]) + results_root = Path(args.results_root or resolved["results"]) + results_root.mkdir(parents=True, exist_ok=True) + device = select_device(args.device) + run_dir = run_experiment( + cfg, + data_root=data_root, + results_root=results_root, + device=device, + resume=args.resume, + overwrite=args.overwrite, + disable_weightwatcher=args.no_weightwatcher, + ) + print(run_dir) + if __name__ == "__main__": main() diff --git a/level_0_baseline/tests/conftest.py b/level_0_baseline/tests/conftest.py new file mode 100644 index 0000000..2edad41 --- /dev/null +++ b/level_0_baseline/tests/conftest.py @@ -0,0 +1,8 @@ +from __future__ import annotations + +import sys +from pathlib import Path + +SOURCE_ROOT = Path(__file__).resolve().parents[1] / "src" +if str(SOURCE_ROOT) not in sys.path: + sys.path.insert(0, str(SOURCE_ROOT)) diff --git a/level_0_baseline/tests/test_baseline.py b/level_0_baseline/tests/test_baseline.py index 76d8995..03ff7ab 100644 --- a/level_0_baseline/tests/test_baseline.py +++ b/level_0_baseline/tests/test_baseline.py @@ -1,97 +1,207 @@ +from __future__ import annotations + import json -import sys +from collections import OrderedDict from pathlib import Path +import numpy as np +import pandas as pd import torch -sys.path.insert(0, str(Path(__file__).resolve().parents[1] / "src")) - -from level0_baseline.data import _progress_message, write_splits +from level0_baseline.config import validate_config +from level0_baseline.data import ( + _progress_message, + prepare_token_splits, +) from level0_baseline.model import GPT, GPTConfig from level0_baseline.optim import make_optimizers +from level0_baseline.train import ( + evaluate, + fixed_eval_starts, + random_batch, + run_experiment, +) + + +class FakeTokenizer: + name = "fake" + eot_token = 31 + n_vocab = 32 + def encode_ordinary(self, text: str) -> list[int]: + return [ord(character) % 31 for character in text] -def cfg(opt): + +def tiny_config() -> dict: return { + "model": { + "vocab_size": 64, + "block_size": 8, + "n_layer": 2, + "n_head": 4, + "n_embd": 32, + "dropout": 0.0, + "bias": False, + }, "training": { - "optimizer": opt, + "batch_size": 2, + "grad_accum_steps": 1, + "max_steps": 2, + "eval_interval": 1, + "eval_batches": 2, + "log_interval": 1, + "checkpoint_interval": 1, "learning_rate": 0.001, + "min_lr": 0.0001, + "warmup_steps": 1, "weight_decay": 0.1, "beta1": 0.9, "beta2": 0.95, - "muon_momentum": 0.95, - "muon_nesterov": True, + "epsilon": 1e-8, + "grad_clip": 1.0, + "optimizer": "adamw", "muon_learning_rate": 0.02, "muon_aux_adamw_learning_rate": 0.001, - } + "muon_momentum": 0.95, + "muon_nesterov": True, + "seed": 1337, + "compile": False, + }, + "analysis": { + "weightwatcher": False, + "weightwatcher_interval": 1, + "randomize": False, + }, } -def test_forward_and_accuracy_shape(): - model = GPT(GPTConfig(block_size=8, n_embd=16, n_head=1, n_layer=1)) - x = torch.randint(0, 256, (2, 8)) - logits, loss = model(x, x) - assert logits.shape == (2, 8, 256) - assert torch.isfinite(loss) - - -def test_adamw_step(): - model = GPT(GPTConfig(block_size=8, n_embd=16)) - optimizers = make_optimizers(model, cfg("adamw")) - _, loss = model( - torch.randint(0, 256, (2, 8)), - torch.randint(0, 256, (2, 8)), - ) - loss.backward() - for optimizer in optimizers: - optimizer.step() +def write_tiny_data(root: Path, vocab_size: int = 64) -> None: + root.mkdir(parents=True, exist_ok=True) + split_tokens = {"train": 512, "val": 128, "test": 128} + for offset, (split, count) in enumerate(split_tokens.items()): + values = (np.arange(count, dtype=np.uint16) + offset) % vocab_size + values.astype(" 10_000_000 + assert model.num_parameters() < 25_000_000 + x = torch.randint(0, model.cfg.vocab_size, (2, 16)) + logits, loss = model(x, x) + assert logits.shape == (2, 16, model.cfg.vocab_size) + assert torch.isfinite(loss) + assert len(list(model.spectral_matrices())) == 24 + + +def test_adamw_step_uses_standard_decay_groups(): + cfg = tiny_config() + model = GPT(GPTConfig(**cfg["model"])) + optimizers = make_optimizers(model, cfg, device_type="cpu") + assert len(optimizers) == 1 + assert len(optimizers[0].param_groups) == 2 + x = torch.randint(0, 64, (2, 8)) + _, loss = model(x, x) + assert loss is not None loss.backward() - for optimizer in optimizers: - optimizer.step() + optimizers[0].step() -def test_progress_message_reports_elapsed_eta_and_stall(): +def test_progress_message_reports_phase_eta_and_stall(): message = _progress_message( - collected_bytes=50, - required_bytes=100, + written_tokens=50, + required_tokens=100, documents=4, elapsed_seconds=10, stalled_seconds=3, + phase="tokenizing", + split="train", ) + assert "phase=tokenizing" in message + assert "split=train" in message assert "documents=4" in message assert "percent= 50.0%" in message - assert "elapsed=10s" in message assert "eta=10s" in message - assert "no_new_bytes_for=3s" in message + assert "no_new_tokens_for=3s" in message -def test_verbose_split_preparation_logs_and_writes_files(tmp_path, capsys): - write_splits( - iter(["abcdef"]), +def test_bpe_split_preparation_is_fixed_and_nonoverlapping(tmp_path): + texts = iter(["abcdefghij", "klmnopqrst", "uvwxyz"]) + metadata = prepare_token_splits( + texts, tmp_path, - train_bytes=2, - val_bytes=2, - test_bytes=2, - verbose=True, - log_interval_seconds=60, + OrderedDict(train=8, val=5, test=4), + FakeTokenizer(), + model_vocab_size=64, + dataset_metadata={"name": "unit"}, + ) + assert metadata["split_tokens"] == {"train": 8, "val": 5, "test": 4} + assert metadata["discarded_boundary_tokens"] > 0 + assert (tmp_path / "train.bin").stat().st_size == 16 + assert (tmp_path / "val.bin").stat().st_size == 10 + assert (tmp_path / "test.bin").stat().st_size == 8 + on_disk = json.loads((tmp_path / "meta.json").read_text()) + assert on_disk["tokenizer"] == "tiktoken:fake" + assert on_disk["model_vocab_size"] == 64 + + +def test_evaluation_does_not_advance_training_rng(tmp_path): + data_path = tmp_path / "train.bin" + (np.arange(512, dtype=np.uint16) % 64).astype("