diff --git a/level_0_baseline/README.md b/level_0_baseline/README.md index cacde29..2795c09 100644 --- a/level_0_baseline/README.md +++ b/level_0_baseline/README.md @@ -1,76 +1,121 @@ -# Level 0 Baseline +# 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 an independent, auditable nanoGPT baseline. It does not import the repository's WW-PGD experiment framework or apply WW-PGD. Its purpose is to establish a credible AdamW language-model baseline and measure WeightWatcher layer spectra before adding an optimizer extension. -## Scope +## Corrected baseline -- 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 original isolated baseline was an 82K-parameter, one-layer byte model. This version uses a materially more realistic MacBook-scale configuration: -## Install +- pinned FineWeb-Edu `sample-10BT` source; +- GPT-2 BPE tokenization (`tiktoken`, vocabulary 50,257); +- 16M training, 1M validation, and 1M test tokens stored as `uint16`; +- 4 transformer blocks, 4 attention heads, width 128, context 256; +- 7,253,248 parameters with tied token embedding/output weights; +- AdamW with decoupled weight decay, gradient clipping, 100-step warmup, and cosine decay; +- microbatch 4 with 8 accumulation steps: 8,192 tokens per optimizer update and 16.384M tokens over 2,000 updates; +- fixed, independent train/validation/test probes that do not advance the training RNG; +- test evaluation only at the final and validation-selected checkpoints; +- non-randomized WeightWatcher analysis of the 24 transformer block matrices only. The large embedding/output matrix is excluded from periodic spectral analysis. -```bash -cd level_0_baseline -python -m venv .venv -source .venv/bin/activate -pip install -e '.[data,analysis,test]' -``` - -## Paths +The default layer learning-rate multiplier is flat (`layer_lr_decay: 1.0`). Layerwise decay is available as an explicit ablation, not silently enabled in the baseline. -Defaults are under `/tmp/nanogpt-level0`. Override them without editing code: +## Install in the existing Conda environment ```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 +conda activate ww_prod310 +cd ~/Desktop/work/nanoGPT/nanogpt-experiments/level_0_baseline +python -m pip install -e '.[data,analysis,test]' ``` -## Prepare the real corpus +## Paths -This prepares fixed 50 MB training, 2 MB validation, and 2 MB test byte-token splits from streamed FineWeb-Edu: +The corrected format uses a new root so the earlier raw-byte files cannot be mistaken for GPT-2-tokenized data: ```bash -level0-prepare-data --dataset fineweb-edu +export NANOGPT_LEVEL0_ROOT=/tmp/nanogpt-level0-gpt2 +export NANOGPT_LEVEL0_DATA_ROOT=$NANOGPT_LEVEL0_ROOT/data +export NANOGPT_LEVEL0_RESULTS_ROOT=$NANOGPT_LEVEL0_ROOT/results ``` -To monitor the streamed download and preparation, enable heartbeat logging: +## Prepare the pinned FineWeb-Edu corpus ```bash level0-prepare-data \ + --config configs/level0.yaml \ --dataset fineweb-edu \ --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 preparer reports documents, GPT-2 tokens, elapsed time, throughput, ETA, and time since the stream last produced tokens. It validates compatible existing data and reuses it; pass `--force` to rebuild it. + +A successful preparation produces: -## Run one seed +```text +/tmp/nanogpt-level0-gpt2/data/train.bin +/tmp/nanogpt-level0-gpt2/data/val.bin +/tmp/nanogpt-level0-gpt2/data/test.bin +/tmp/nanogpt-level0-gpt2/data/meta.json +``` + +## Run one AdamW seed ```bash -./scripts/run_one.sh adamw 1337 -./scripts/run_one.sh muon 1337 +./scripts/run_one.sh adamw 1337 \ + 2>&1 | tee /tmp/level0-gpt2-adamw-seed1337.log ``` -## Run multiple seeds +The command refuses to overwrite an existing nonempty run. To intentionally replace it: ```bash -NANOGPT_LEVEL0_SEEDS=1337,2027,4099 ./scripts/run_multiseed.sh +NANOGPT_LEVEL0_OVERWRITE=1 ./scripts/run_one.sh adamw 1337 ``` -The notebooks read `NANOGPT_LEVEL0_RESULTS_ROOT`. Select the single-seed run with `NANOGPT_LEVEL0_NOTEBOOK_OPTIMIZER` and `NANOGPT_LEVEL0_NOTEBOOK_SEED`. +The run writes periodic metrics and checkpoints, validation-selected and final test metrics, WeightWatcher CSV files, and `run_complete.json` under: + +```text +/tmp/nanogpt-level0-gpt2/results/adamw_seed_1337 +``` -For a bounded infrastructure smoke test, override the run length and batch size: +## Run several AdamW 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 \ + ./scripts/run_multiseed.sh ``` -Next-token error is `1 - next-token accuracy`; the notebooks derive and plot it explicitly. +Muon remains available for a later baseline comparison: + +```bash +NANOGPT_LEVEL0_OPTIMIZERS=adamw,muon ./scripts/run_multiseed.sh +``` + +## Analyze one seed + +```bash +export NANOGPT_LEVEL0_RESULTS_ROOT=/tmp/nanogpt-level0-gpt2/results +export NANOGPT_LEVEL0_NOTEBOOK_OPTIMIZER=adamw +export NANOGPT_LEVEL0_NOTEBOOK_SEED=1337 +jupyter lab notebooks/01_single_seed.ipynb +``` + +The notebook plots train/validation loss, perplexity, exact next-GPT-2-token accuracy, generalization gap, learning rate, final versus validation-selected test metrics, and WeightWatcher alpha by transformer matrix. + +## Bounded smoke run + +The smoke run checks the software path only; it is not a scientific result: + +```bash +NANOGPT_LEVEL0_MAX_STEPS=2 \ +NANOGPT_LEVEL0_BATCH_SIZE=1 \ +NANOGPT_LEVEL0_GRAD_ACCUM_STEPS=1 \ +NANOGPT_LEVEL0_EVAL_INTERVAL=1 \ +NANOGPT_LEVEL0_WEIGHTWATCHER=0 \ +NANOGPT_LEVEL0_OVERWRITE=1 \ + ./scripts/run_one.sh adamw 1337 +``` + +## Primary outcomes + +Use held-out cross-entropy and perplexity as the principal language-model outcomes. Exact token accuracy is reported as a secondary diagnostic and should not be compared numerically with the old 256-byte-vocabulary accuracy. diff --git a/level_0_baseline/configs/level0.yaml b/level_0_baseline/configs/level0.yaml index 5d385c9..19459e4 100644 --- a/level_0_baseline/configs/level0.yaml +++ b/level_0_baseline/configs/level0.yaml @@ -1,17 +1,30 @@ +data: + dataset_name: HuggingFaceFW/fineweb-edu + dataset_config: sample-10BT + dataset_split: train + dataset_revision: 593b3a867298afb8ce42625a270ef20ddcad28f9 + tokenizer: gpt2 + dtype: uint16 + train_tokens: 16000000 + val_tokens: 1000000 + test_tokens: 1000000 model: - vocab_size: 256 + vocab_size: 50257 block_size: 256 - n_layer: 1 - n_head: 1 - n_embd: 64 + n_layer: 4 + n_head: 4 + n_embd: 128 dropout: 0.0 bias: false + tie_weights: true training: - batch_size: 16 - grad_accum_steps: 1 + batch_size: 4 + grad_accum_steps: 8 max_steps: 2000 eval_interval: 50 eval_batches: 20 + eval_batch_size: 4 + test_eval_batches: 40 checkpoint_interval: 250 learning_rate: 0.0006 muon_learning_rate: 0.02 @@ -21,7 +34,9 @@ training: weight_decay: 0.1 beta1: 0.9 beta2: 0.95 + epsilon: 1.0e-8 grad_clip: 1.0 + layer_lr_decay: 1.0 optimizer: adamw muon_momentum: 0.95 muon_nesterov: true @@ -30,4 +45,4 @@ training: analysis: weightwatcher: true weightwatcher_interval: 250 - randomize: true + randomize: false diff --git a/level_0_baseline/notebooks/01_single_seed.ipynb b/level_0_baseline/notebooks/01_single_seed.ipynb index d887603..ad5ca4a 100644 --- a/level_0_baseline/notebooks/01_single_seed.ipynb +++ b/level_0_baseline/notebooks/01_single_seed.ipynb @@ -1,10 +1,188 @@ { "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", + "id": "a2e4ab7d", + "metadata": {}, + "source": [ + "# Isolated Level 0 — single-seed diagnostics\n", + "\n", + "GPT-2 BPE AdamW/Muon baseline on pinned FineWeb-Edu." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "e145cca3", + "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-gpt2/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", + "\n", + "required = ['manifest.json', 'metrics.csv', 'run_complete.json', 'final_metrics.json', 'selected_checkpoint_metrics.json']\n", + "missing = [name for name in required if not (RUN / name).is_file()]\n", + "if missing:\n", + " raise FileNotFoundError(f'Missing run artifacts in {RUN}: {missing}')\n", + "\n", + "manifest = json.loads((RUN / 'manifest.json').read_text())\n", + "completion = json.loads((RUN / 'run_complete.json').read_text())\n", + "final_metrics = json.loads((RUN / 'final_metrics.json').read_text())\n", + "selected_metrics = json.loads((RUN / 'selected_checkpoint_metrics.json').read_text())\n", + "metrics = pd.read_csv(RUN / 'metrics.csv')\n", + "\n", + "pd.DataFrame([{\n", + " 'run': str(RUN),\n", + " 'parameters': manifest['parameter_count'],\n", + " 'optimizer_steps': completion['optimizer_steps'],\n", + " 'tokens_seen': completion['tokens_seen'],\n", + " 'best_validation_step': completion['best_validation_step'],\n", + " 'best_validation_loss': completion['best_validation_loss'],\n", + " 'selected_test_loss': completion['selected_test_loss'],\n", + " 'final_test_loss': completion['final_test_loss'],\n", + " 'weightwatcher_failures': completion['weightwatcher_failures'],\n", + "}])" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "afdada3c", + "metadata": {}, + "outputs": [], + "source": [ + "for metric, ylabel in [\n", + " ('loss', 'cross-entropy loss'),\n", + " ('perplexity', 'perplexity'),\n", + " ('accuracy', 'exact next-token accuracy (%)'),\n", + "]:\n", + " plt.figure(figsize=(10, 5))\n", + " for split in ('train', 'val'):\n", + " values = metrics[f'{split}_{metric}']\n", + " if metric == 'accuracy':\n", + " values = 100 * values\n", + " plt.plot(metrics['tokens_seen'], values, label=split)\n", + " plt.xlabel('training tokens')\n", + " plt.ylabel(ylabel)\n", + " plt.title(f'{OPTIMIZER} seed {SEED}: {metric}')\n", + " plt.grid(alpha=0.25)\n", + " plt.legend()\n", + " plt.show()" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "2a698b15", + "metadata": {}, + "outputs": [], + "source": [ + "plt.figure(figsize=(10, 5))\n", + "plt.plot(metrics['tokens_seen'], metrics['val_generalization_gap'])\n", + "plt.axhline(0.0, linewidth=1, linestyle='--')\n", + "plt.xlabel('training tokens')\n", + "plt.ylabel('validation loss − training loss')\n", + "plt.title(f'{OPTIMIZER} seed {SEED}: generalization gap')\n", + "plt.grid(alpha=0.25)\n", + "plt.show()\n", + "\n", + "plt.figure(figsize=(10, 5))\n", + "plt.plot(metrics['step'], metrics['learning_rate'], label='base learning rate')\n", + "plt.plot(metrics['step'], metrics['min_group_learning_rate'], linestyle='--', label='minimum parameter-group LR')\n", + "plt.plot(metrics['step'], metrics['max_group_learning_rate'], linestyle=':', label='maximum parameter-group LR')\n", + "plt.xlabel('optimizer step')\n", + "plt.ylabel('learning rate')\n", + "plt.title(f'{OPTIMIZER} seed {SEED}: learning-rate schedule')\n", + "plt.grid(alpha=0.25)\n", + "plt.legend()\n", + "plt.show()" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "11dbb708", + "metadata": {}, + "outputs": [], + "source": [ + "checkpoint_comparison = pd.DataFrame([\n", + " {\n", + " 'checkpoint': 'validation-selected',\n", + " 'step': selected_metrics['selected_step'],\n", + " 'validation_loss': selected_metrics['validation_loss'],\n", + " 'test_loss': selected_metrics['test_loss'],\n", + " 'test_perplexity': selected_metrics['test_perplexity'],\n", + " 'test_accuracy_percent': 100 * selected_metrics['test_accuracy'],\n", + " },\n", + " {\n", + " 'checkpoint': 'final',\n", + " 'step': final_metrics['step'],\n", + " 'validation_loss': final_metrics['validation_loss'],\n", + " 'test_loss': final_metrics['test_loss'],\n", + " 'test_perplexity': final_metrics['test_perplexity'],\n", + " 'test_accuracy_percent': 100 * final_metrics['test_accuracy'],\n", + " },\n", + "])\n", + "checkpoint_comparison" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "2ddfc553", + "metadata": {}, + "outputs": [], + "source": [ + "ww_files = sorted(RUN.glob('weightwatcher_step_*.csv'))\n", + "if not ww_files:\n", + " print('No successful WeightWatcher measurements found.')\n", + "else:\n", + " ww = pd.concat([pd.read_csv(path) for path in ww_files], ignore_index=True)\n", + " ww['alpha'] = pd.to_numeric(ww.get('alpha'), errors='coerce')\n", + " ww['step'] = pd.to_numeric(ww['step'], errors='coerce')\n", + " layer_column = next((name for name in ('source_layer', 'longname', 'name', 'layer_id') if name in ww), None)\n", + " if layer_column is None:\n", + " raise KeyError('WeightWatcher output has no layer identifier column')\n", + " valid = ww[np.isfinite(ww['alpha'])].copy()\n", + " plt.figure(figsize=(12, 7))\n", + " for layer, frame in valid.groupby(layer_column):\n", + " frame = frame.sort_values('step')\n", + " plt.plot(frame['step'], frame['alpha'], marker='o', linewidth=1, markersize=3, label=str(layer))\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}: transformer-matrix alpha trajectories')\n", + " plt.grid(alpha=0.25)\n", + " plt.legend(fontsize=7, ncol=2, bbox_to_anchor=(1.02, 1), loc='upper left')\n", + " plt.tight_layout()\n", + " plt.show()\n", + "\n", + " display_columns = [column for column in ('step', layer_column, 'alpha', 'D', 'xmin', 'num_evals') if column in valid]\n", + " display(valid[display_columns].sort_values(['step', layer_column]).tail(30))" + ] + } ], - "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 } diff --git a/level_0_baseline/notebooks/02_multiseed.ipynb b/level_0_baseline/notebooks/02_multiseed.ipynb index d2ac6d9..3490f18 100644 --- a/level_0_baseline/notebooks/02_multiseed.ipynb +++ b/level_0_baseline/notebooks/02_multiseed.ipynb @@ -1,10 +1,152 @@ { "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", + "id": "c9a2de26", + "metadata": {}, + "source": [ + "# Isolated Level 0 — multi-seed diagnostics" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "85d1bd38", + "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-gpt2/results'))\n", + "rows = []\n", + "summary_rows = []\n", + "for path in sorted(ROOT.glob('*_seed_*/metrics.csv')):\n", + " run = path.parent\n", + " completion_path = run / 'run_complete.json'\n", + " selected_path = run / 'selected_checkpoint_metrics.json'\n", + " if not completion_path.is_file() or not selected_path.is_file():\n", + " continue\n", + " optimizer, seed_text = run.name.rsplit('_seed_', 1)\n", + " frame = pd.read_csv(path)\n", + " frame['optimizer'] = optimizer\n", + " frame['seed'] = int(seed_text)\n", + " rows.append(frame)\n", + " selected = json.loads(selected_path.read_text())\n", + " completion = json.loads(completion_path.read_text())\n", + " summary_rows.append({\n", + " 'optimizer': optimizer,\n", + " 'seed': int(seed_text),\n", + " 'selected_step': selected['selected_step'],\n", + " 'selected_validation_loss': selected['validation_loss'],\n", + " 'selected_test_loss': selected['test_loss'],\n", + " 'selected_test_perplexity': selected['test_perplexity'],\n", + " 'selected_test_accuracy_percent': 100 * selected['test_accuracy'],\n", + " 'final_test_loss': completion['final_test_loss'],\n", + " })\n", + "if not rows:\n", + " raise FileNotFoundError(f'No completed runs under {ROOT}')\n", + "all_metrics = pd.concat(rows, ignore_index=True)\n", + "summary = pd.DataFrame(summary_rows).sort_values(['optimizer', 'seed'])\n", + "summary" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "0aecaa20", + "metadata": {}, + "outputs": [], + "source": [ + "def band_plot(column, ylabel):\n", + " plt.figure(figsize=(10, 6))\n", + " for optimizer, frame in all_metrics.groupby('optimizer'):\n", + " aggregated = frame.groupby('tokens_seen')[column].agg(['mean', 'std']).reset_index()\n", + " spread = aggregated['std'].fillna(0.0)\n", + " line, = plt.plot(aggregated['tokens_seen'], aggregated['mean'], label=optimizer)\n", + " plt.fill_between(\n", + " aggregated['tokens_seen'],\n", + " aggregated['mean'] - spread,\n", + " aggregated['mean'] + spread,\n", + " alpha=0.2,\n", + " color=line.get_color(),\n", + " )\n", + " plt.xlabel('training tokens')\n", + " plt.ylabel(ylabel)\n", + " plt.title(f'{column}: mean ± one standard deviation')\n", + " plt.grid(alpha=0.25)\n", + " plt.legend()\n", + " plt.show()\n", + "\n", + "band_plot('val_loss', 'validation cross-entropy loss')\n", + "band_plot('val_perplexity', 'validation perplexity')\n", + "band_plot('val_accuracy', 'exact next-token accuracy (fraction)')\n", + "band_plot('val_generalization_gap', 'validation loss − training loss')" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "b870db14", + "metadata": {}, + "outputs": [], + "source": [ + "ww_rows = []\n", + "for run in sorted(ROOT.glob('*_seed_*')):\n", + " if '_seed_' not in run.name:\n", + " continue\n", + " optimizer, seed_text = run.name.rsplit('_seed_', 1)\n", + " for path in sorted(run.glob('weightwatcher_step_*.csv')):\n", + " frame = pd.read_csv(path)\n", + " frame['optimizer'] = optimizer\n", + " frame['seed'] = int(seed_text)\n", + " ww_rows.append(frame)\n", + "if not ww_rows:\n", + " print('No successful WeightWatcher measurements found.')\n", + "else:\n", + " ww = pd.concat(ww_rows, ignore_index=True)\n", + " ww['alpha'] = pd.to_numeric(ww.get('alpha'), errors='coerce')\n", + " layer_column = next((name for name in ('source_layer', 'longname', 'name', 'layer_id') if name in ww), None)\n", + " valid = ww[np.isfinite(ww['alpha'])].copy()\n", + " for layer, layer_frame in valid.groupby(layer_column):\n", + " plt.figure(figsize=(10, 5))\n", + " for optimizer, frame in layer_frame.groupby('optimizer'):\n", + " aggregated = frame.groupby('step')['alpha'].agg(['mean', 'std']).reset_index()\n", + " spread = aggregated['std'].fillna(0.0)\n", + " line, = plt.plot(aggregated['step'], aggregated['mean'], label=optimizer)\n", + " plt.fill_between(\n", + " aggregated['step'],\n", + " aggregated['mean'] - spread,\n", + " aggregated['mean'] + spread,\n", + " alpha=0.2,\n", + " color=line.get_color(),\n", + " )\n", + " plt.axhline(2.0, linestyle='--', linewidth=1)\n", + " plt.xlabel('optimizer step')\n", + " plt.ylabel('WeightWatcher alpha')\n", + " plt.title(str(layer))\n", + " plt.grid(alpha=0.25)\n", + " plt.legend()\n", + " plt.show()" + ] + } ], - "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 } diff --git a/level_0_baseline/pyproject.toml b/level_0_baseline/pyproject.toml index 5840ab6..23e3440 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 = "Self-contained GPT-2-BPE Level 0 nanoGPT baseline for AdamW and Muon" 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] 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..2593861 --- a/level_0_baseline/scripts/run_multiseed.sh +++ b/level_0_baseline/scripts/run_multiseed.sh @@ -1,9 +1,17 @@ #!/usr/bin/env bash set -euo pipefail + +SCRIPT_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)" +REPO_ROOT="$(cd "$SCRIPT_DIR/.." && pwd)" +cd "$REPO_ROOT" + SEEDS="${NANOGPT_LEVEL0_SEEDS:-1337,2027,4099}" -for optimizer in adamw muon; do - IFS=',' read -ra xs <<< "$SEEDS" - for seed in "${xs[@]}"; do +OPTIMIZERS="${NANOGPT_LEVEL0_OPTIMIZERS:-adamw}" +IFS=',' read -r -a SEED_ARRAY <<< "$SEEDS" +IFS=',' read -r -a OPTIMIZER_ARRAY <<< "$OPTIMIZERS" + +for optimizer in "${OPTIMIZER_ARRAY[@]}"; do + for seed in "${SEED_ARRAY[@]}"; do ./scripts/run_one.sh "$optimizer" "$seed" 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..624dfa8 --- a/level_0_baseline/scripts/run_one.sh +++ b/level_0_baseline/scripts/run_one.sh @@ -1,6 +1,28 @@ #!/usr/bin/env bash set -euo pipefail -ROOT="${NANOGPT_LEVEL0_ROOT:-/tmp/nanogpt-level0}" + +SCRIPT_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)" +REPO_ROOT="$(cd "$SCRIPT_DIR/.." && pwd)" +cd "$REPO_ROOT" + +ROOT="${NANOGPT_LEVEL0_ROOT:-/tmp/nanogpt-level0-gpt2}" 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="${NANOGPT_LEVEL0_DEVICE:-auto}" +DATA_ROOT="${NANOGPT_LEVEL0_DATA_ROOT:-$ROOT/data}" +RESULTS_ROOT="${NANOGPT_LEVEL0_RESULTS_ROOT:-$ROOT/results}" + +ARGS=( + --config configs/level0.yaml + --optimizer "$OPTIMIZER" + --seed "$SEED" + --device "$DEVICE" + --data-root "$DATA_ROOT" + --results-root "$RESULTS_ROOT" +) +if [[ "${NANOGPT_LEVEL0_OVERWRITE:-0}" == "1" ]]; then + ARGS+=(--overwrite) +fi + +export PYTHONUNBUFFERED=1 +python -u -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..2c6e23e 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" +"""Self-contained 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..f91122c 100644 --- a/level_0_baseline/src/level0_baseline/config.py +++ b/level_0_baseline/src/level0_baseline/config.py @@ -1,31 +1,172 @@ from __future__ import annotations + +import copy import os from pathlib import Path +from typing import Any, Callable + import yaml -DEFAULT_ROOT = Path("/tmp/nanogpt-level0") +DEFAULT_ROOT = Path("/tmp/nanogpt-level0-gpt2") + + +def _parse_bool(value: str) -> bool: + normalized = value.strip().lower() + if normalized in {"1", "true", "yes", "on"}: + return True + if normalized in {"0", "false", "no", "off"}: + return False + raise ValueError(f"invalid boolean value: {value!r}") + def roots() -> dict[str, Path]: root = Path(os.getenv("NANOGPT_LEVEL0_ROOT", DEFAULT_ROOT)) return { "root": root, "data": Path(os.getenv("NANOGPT_LEVEL0_DATA_ROOT", root / "data")), - "results": Path(os.getenv("NANOGPT_LEVEL0_RESULTS_ROOT", root / "results")), + "results": Path( + os.getenv("NANOGPT_LEVEL0_RESULTS_ROOT", root / "results") + ), "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 = { + +def _set_nested(cfg: dict[str, Any], section: str, key: str, value: Any) -> None: + if section not in cfg or not isinstance(cfg[section], dict): + raise ValueError(f"configuration is missing section {section!r}") + cfg[section][key] = value + + +def load_config(path: str | Path) -> dict[str, Any]: + with open(path, "r", encoding="utf-8") as handle: + loaded = yaml.safe_load(handle) + if not isinstance(loaded, dict): + raise ValueError("configuration root must be a mapping") + cfg: dict[str, Any] = copy.deepcopy(loaded) + + env_map: dict[str, tuple[str, str, Callable[[str], Any]]] = { "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_MIN_LR": ("training", "min_lr", float), + "NANOGPT_LEVEL0_WARMUP_STEPS": ("training", "warmup_steps", int), + "NANOGPT_LEVEL0_WEIGHT_DECAY": ("training", "weight_decay", float), "NANOGPT_LEVEL0_EVAL_INTERVAL": ("training", "eval_interval", int), + "NANOGPT_LEVEL0_EVAL_BATCHES": ("training", "eval_batches", int), + "NANOGPT_LEVEL0_EVAL_BATCH_SIZE": ( + "training", + "eval_batch_size", + int, + ), + "NANOGPT_LEVEL0_TEST_EVAL_BATCHES": ( + "training", + "test_eval_batches", + int, + ), + "NANOGPT_LEVEL0_CHECKPOINT_INTERVAL": ( + "training", + "checkpoint_interval", + int, + ), + "NANOGPT_LEVEL0_LAYER_LR_DECAY": ( + "training", + "layer_lr_decay", + float, + ), + "NANOGPT_LEVEL0_N_LAYER": ("model", "n_layer", int), + "NANOGPT_LEVEL0_N_HEAD": ("model", "n_head", int), + "NANOGPT_LEVEL0_N_EMBD": ("model", "n_embd", int), + "NANOGPT_LEVEL0_BLOCK_SIZE": ("model", "block_size", int), + "NANOGPT_LEVEL0_WEIGHTWATCHER": ( + "analysis", + "weightwatcher", + _parse_bool, + ), + "NANOGPT_LEVEL0_WW_INTERVAL": ( + "analysis", + "weightwatcher_interval", + int, + ), } for name, (section, key, cast) in env_map.items(): if name in os.environ: - cfg[section][key] = cast(os.environ[name]) + _set_nested(cfg, section, key, cast(os.environ[name])) + + validate_config(cfg) return cfg + + +def validate_config(cfg: dict[str, Any]) -> None: + for section in ("data", "model", "training", "analysis"): + if section not in cfg or not isinstance(cfg[section], dict): + raise ValueError(f"configuration is missing section {section!r}") + + data = cfg["data"] + model = cfg["model"] + training = cfg["training"] + analysis = cfg["analysis"] + + if data.get("tokenizer") != "gpt2": + raise ValueError("the corrected Level 0 baseline requires tokenizer: gpt2") + if data.get("dtype") != "uint16": + raise ValueError("the corrected Level 0 baseline requires dtype: uint16") + for key in ("train_tokens", "val_tokens", "test_tokens"): + if int(data.get(key, 0)) <= 0: + raise ValueError(f"data.{key} must be positive") + + for key in ("vocab_size", "block_size", "n_layer", "n_head", "n_embd"): + if int(model.get(key, 0)) <= 0: + raise ValueError(f"model.{key} must be positive") + if int(model["n_embd"]) % int(model["n_head"]) != 0: + raise ValueError("model.n_embd must be divisible by model.n_head") + if int(model["vocab_size"]) < 50257: + raise ValueError("model.vocab_size must cover the GPT-2 tokenizer") + dropout = float(model.get("dropout", 0.0)) + if not 0.0 <= dropout < 1.0: + raise ValueError("model.dropout must satisfy 0 <= dropout < 1") + + positive_ints = ( + "batch_size", + "grad_accum_steps", + "max_steps", + "eval_interval", + "eval_batches", + "eval_batch_size", + "test_eval_batches", + "checkpoint_interval", + ) + for key in positive_ints: + if int(training.get(key, 0)) <= 0: + raise ValueError(f"training.{key} must be positive") + if int(training.get("warmup_steps", 0)) < 0: + raise ValueError("training.warmup_steps must be nonnegative") + if int(training["warmup_steps"]) >= int(training["max_steps"]): + raise ValueError("training.warmup_steps must be smaller than max_steps") + if float(training.get("learning_rate", 0.0)) <= 0: + raise ValueError("training.learning_rate must be positive") + if not 0 < float(training.get("min_lr", 0.0)) <= float( + training["learning_rate"] + ): + raise ValueError("training.min_lr must be in (0, learning_rate]") + if float(training.get("weight_decay", -1.0)) < 0: + raise ValueError("training.weight_decay must be nonnegative") + if float(training.get("epsilon", 0.0)) <= 0: + raise ValueError("training.epsilon must be positive") + if float(training.get("grad_clip", -1.0)) < 0: + raise ValueError("training.grad_clip must be nonnegative") + layer_lr_decay = float(training.get("layer_lr_decay", 1.0)) + if not 0 < layer_lr_decay <= 1: + raise ValueError("training.layer_lr_decay must be in (0, 1]") + if str(training.get("optimizer", "")).lower() not in {"adamw", "muon"}: + raise ValueError("training.optimizer must be adamw or muon") + + interval = int(analysis.get("weightwatcher_interval", 0)) + if bool(analysis.get("weightwatcher", False)) and interval <= 0: + raise ValueError("analysis.weightwatcher_interval must be positive") diff --git a/level_0_baseline/src/level0_baseline/data.py b/level_0_baseline/src/level0_baseline/data.py index ed96451..4fe63b5 100644 --- a/level_0_baseline/src/level0_baseline/data.py +++ b/level_0_baseline/src/level0_baseline/data.py @@ -1,22 +1,35 @@ from __future__ import annotations import argparse +import gc +import hashlib import json +import os import sys import threading import time from pathlib import Path +from typing import Any, Iterable, Iterator, Protocol import numpy as np -from .config import roots +from .config import load_config, roots +DATA_SCHEMA_VERSION = 2 +UINT16_MAX = np.iinfo(np.uint16).max -def encode(text: str) -> np.ndarray: - return np.frombuffer(text.encode("utf-8", errors="replace"), dtype=np.uint8) +class Tokenizer(Protocol): + name: str + n_vocab: int + eot_token: int -def _format_duration(seconds: float) -> str: + def encode_ordinary(self, text: str) -> list[int]: ... + + +def _format_duration(seconds: float | None) -> str: + if seconds is None: + return "unknown" seconds = max(0, int(seconds)) hours, remainder = divmod(seconds, 3600) minutes, secs = divmod(remainder, 60) @@ -29,41 +42,38 @@ def _format_duration(seconds: float) -> str: def _progress_message( *, - collected_bytes: int, - required_bytes: int, + collected_tokens: int, + required_tokens: int, documents: int, elapsed_seconds: float, stalled_seconds: float, ) -> str: - elapsed_seconds = max(elapsed_seconds, 1e-9) - rate = collected_bytes / elapsed_seconds - remaining = max(required_bytes - collected_bytes, 0) + elapsed_seconds = max(float(elapsed_seconds), 1e-9) + rate = collected_tokens / elapsed_seconds + remaining = max(required_tokens - collected_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 * collected_tokens / max(required_tokens, 1) return ( "[level0-prepare-data] progress " f"documents={documents:,} " - f"bytes={collected_bytes:,}/{required_bytes:,} " + f"tokens={collected_tokens:,}/{required_tokens:,} " 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"speed={rate:,.0f} tokens/s " + f"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.""" + """Emit heartbeats even while dataset metadata or streaming is blocked.""" - 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.collected_tokens = 0 self.documents = 0 self._lock = threading.Lock() self._stop = threading.Event() @@ -76,37 +86,37 @@ def __init__(self, required_bytes: int, interval_seconds: float): 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 update(self, documents: int, collected_tokens: int) -> None: now = time.monotonic() with self._lock: - if collected_bytes > self.collected_bytes: + if collected_tokens > self.collected_tokens: self.last_progress_at = now self.documents = int(documents) - self.collected_bytes = int(collected_bytes) + self.collected_tokens = int(collected_tokens) def _snapshot(self) -> tuple[int, int, float, float]: now = time.monotonic() with self._lock: return ( self.documents, - self.collected_bytes, + self.collected_tokens, now - self.started_at, now - self.last_progress_at, ) def _run(self) -> None: while not self._stop.wait(self.interval_seconds): - documents, collected_bytes, elapsed, stalled = self._snapshot() + documents, tokens, elapsed, stalled = self._snapshot() print( _progress_message( - collected_bytes=collected_bytes, - required_bytes=self.required_bytes, + collected_tokens=tokens, + required_tokens=self.required_tokens, documents=documents, elapsed_seconds=elapsed, stalled_seconds=stalled, @@ -121,161 +131,397 @@ def stop(self) -> tuple[int, int, float, float]: return self._snapshot() -def _fineweb_texts(load_dataset, verbose: bool): +def _sha256_file(path: Path) -> str: + digest = hashlib.sha256() + with path.open("rb") as handle: + for block in iter(lambda: handle.read(1024 * 1024), b""): + digest.update(block) + return digest.hexdigest() + + +def _load_tokenizer(name: str) -> Tokenizer: + try: + import tiktoken + except ImportError as exc: # pragma: no cover - exercised by CLI users. + raise SystemExit( + "Install data support from level_0_baseline: " + "python -m pip install -e '.[data]'" + ) from exc + tokenizer = tiktoken.get_encoding(name) + if tokenizer.n_vocab - 1 > UINT16_MAX: + raise ValueError( + f"tokenizer {name!r} has IDs that do not fit in uint16" + ) + return tokenizer + + +def _fineweb_texts( + load_dataset, + *, + dataset_name: str, + dataset_config: str, + dataset_split: str, + dataset_revision: str, + verbose: bool, +) -> Iterator[str]: if verbose: print( "[level0-prepare-data] resolving streamed dataset " - "HuggingFaceFW/fineweb-edu sample-10BT train", + f"{dataset_name} {dataset_config} {dataset_split} " + f"revision={dataset_revision}", file=sys.stderr, flush=True, ) dataset = load_dataset( - "HuggingFaceFW/fineweb-edu", - name="sample-10BT", - split="train", + dataset_name, + name=dataset_config, + split=dataset_split, + revision=dataset_revision, streaming=True, ) if verbose: print( - "[level0-prepare-data] dataset stream ready; collecting documents", + "[level0-prepare-data] dataset stream ready; tokenizing documents", file=sys.stderr, flush=True, ) - for row in dataset: - yield row["text"] + try: + for row in dataset: + text = row.get("text") + if isinstance(text, str) and text: + yield text + finally: + del dataset + gc.collect() + + +def _compatible_metadata( + metadata: dict[str, Any], + expected: dict[str, Any], +) -> bool: + keys = ( + "data_schema_version", + "dataset_name", + "dataset_config", + "dataset_split", + "dataset_revision", + "tokenizer", + "vocab_size", + "dtype", + ) + if any(metadata.get(key) != expected.get(key) for key in keys): + return False + return metadata.get("split_tokens") == expected.get("split_tokens") + + +def validate_prepared_data( + output_dir: str | Path, + *, + expected: dict[str, Any] | None = None, + verify_hashes: bool = False, +) -> dict[str, Any]: + output = Path(output_dir) + metadata_path = output / "meta.json" + if not metadata_path.is_file(): + raise FileNotFoundError(f"missing prepared-data metadata: {metadata_path}") + metadata = json.loads(metadata_path.read_text(encoding="utf-8")) + if metadata.get("data_schema_version") != DATA_SCHEMA_VERSION: + raise ValueError( + "prepared data is not the corrected GPT-2-BPE Level 0 format; " + "prepare a fresh dataset in a new directory" + ) + if metadata.get("dtype") != "uint16": + raise ValueError("prepared data must use uint16 GPT-2 token IDs") + if expected is not None and not _compatible_metadata(metadata, expected): + raise ValueError("prepared data does not match the requested identity") + + split_tokens = metadata.get("split_tokens") or {} + hashes = metadata.get("sha256") or {} + for split in ("train", "val", "test"): + path = output / f"{split}.bin" + if not path.is_file(): + raise FileNotFoundError(f"missing prepared split: {path}") + expected_bytes = int(split_tokens.get(split, 0)) * np.dtype(np.uint16).itemsize + if expected_bytes <= 0 or path.stat().st_size != expected_bytes: + raise ValueError( + f"prepared split {split} has {path.stat().st_size:,} bytes; " + f"expected {expected_bytes:,}" + ) + if verify_hashes and _sha256_file(path) != hashes.get(split): + raise ValueError(f"prepared split hash mismatch: {split}") + return metadata -def write_splits( - texts, - out: Path, - train_bytes: int, - val_bytes: int, - test_bytes: int, +def write_token_splits( + texts: Iterable[str], + output_dir: str | Path, *, + tokenizer: Tokenizer, + train_tokens: int, + val_tokens: int, + test_tokens: int, + dataset_metadata: dict[str, str], 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 +) -> dict[str, Any]: + output = Path(output_dir) + output.mkdir(parents=True, exist_ok=True) + split_sizes = { + "train": int(train_tokens), + "val": int(val_tokens), + "test": int(test_tokens), + } + required_tokens = sum(split_sizes.values()) + if required_tokens <= 0: + raise ValueError("the requested token count must be positive") + reporter = ( - _ProgressReporter(need, log_interval_seconds) if verbose else None + _ProgressReporter(required_tokens, log_interval_seconds) + if verbose + else None ) if reporter is not None: - reporter.start(out) - + reporter.start(output) + + split_order = ("train", "val", "test") + split_chunks: dict[str, list[np.ndarray]] = {name: [] for name in split_order} + split_collected = {name: 0 for name in split_order} + split_documents = {name: 0 for name in split_order} + split_index = 0 + used_tokens = 0 + encoded_tokens = 0 + discarded_boundary_tokens = 0 + documents = 0 + snapshot: tuple[int, int, float, float] | None = None try: for text in texts: + ids = tokenizer.encode_ordinary(text) + ids.append(int(tokenizer.eot_token)) + if ids and max(ids) > UINT16_MAX: + raise ValueError("token ID exceeds uint16 capacity") + chunk = np.asarray(ids, dtype=np.uint16) + encoded_tokens += int(chunk.size) documents += 1 - x = encode(text + "\n") - chunks.append(x) - total += len(x) + + if split_index >= len(split_order): + break + split = split_order[split_index] + remaining = split_sizes[split] - split_collected[split] + take = min(remaining, int(chunk.size)) + if take > 0: + split_chunks[split].append(chunk[:take]) + split_collected[split] += take + split_documents[split] += 1 + used_tokens += take + if take < int(chunk.size): + # Never let one source document cross a scientific split. The + # unused suffix is deliberately discarded at the boundary. + discarded_boundary_tokens += int(chunk.size) - take + if split_collected[split] == split_sizes[split]: + split_index += 1 + if reporter is not None: - reporter.update(documents, total) - if total >= need: + reporter.update(documents, used_tokens) + if split_index >= len(split_order): break finally: - snapshot = reporter.stop() if reporter is not None else None - - if total < need: - raise RuntimeError(f"corpus supplied {total:,} bytes; need {need:,}") - + close = getattr(texts, "close", None) + if callable(close): + close() + if reporter is not None: + snapshot = reporter.stop() + + if split_collected != split_sizes: + raise RuntimeError( + "corpus did not supply all requested split tokens: " + f"collected={split_collected}, requested={split_sizes}" + ) 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", + "[level0-prepare-data] tokenization complete " + f"documents={documents:,} used_tokens={used_tokens:,} " + f"elapsed={_format_duration(snapshot[2])}; writing 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), + temporary_paths: dict[str, Path] = {} + hashes: dict[str, str] = {} + for split in split_order: + temporary_path = output / f".{split}.bin.tmp" + np.concatenate(split_chunks[split]).tofile(temporary_path) + temporary_paths[split] = temporary_path + hashes[split] = _sha256_file(temporary_path) + for split in split_order: + os.replace(temporary_paths[split], output / f"{split}.bin") + + metadata: dict[str, Any] = { + "data_schema_version": DATA_SCHEMA_VERSION, + **dataset_metadata, + "tokenizer": tokenizer.name, + "vocab_size": int(tokenizer.n_vocab), + "eot_token": int(tokenizer.eot_token), + "dtype": "uint16", + "split_tokens": split_sizes, + "split_documents": split_documents, + "document_disjoint_splits": True, + "documents_consumed": documents, + "tokens_encoded": encoded_tokens, + "tokens_used": used_tokens, + "boundary_tokens_discarded": discarded_boundary_tokens, + "sha256": hashes, + "created_unix_time": time.time(), } - 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, - ) + temporary_metadata = output / ".meta.json.tmp" + temporary_metadata.write_text( + json.dumps(metadata, indent=2, sort_keys=True), + encoding="utf-8", ) + os.replace(temporary_metadata, output / "meta.json") if verbose: - sizes = ", ".join( - f"{name}={end - start:,}" - for name, (start, end) in boundaries.items() - ) print( - f"[level0-prepare-data] complete output={out} {sizes}", + "[level0-prepare-data] complete " + f"output={output} " + + " ".join( + f"{split}={count:,}" for split, count in split_sizes.items() + ), file=sys.stderr, flush=True, ) + return metadata + + +def expected_metadata(config: dict[str, Any], tokenizer: Tokenizer) -> dict[str, Any]: + data = config["data"] + return { + "data_schema_version": DATA_SCHEMA_VERSION, + "dataset_name": data["dataset_name"], + "dataset_config": data["dataset_config"], + "dataset_split": data["dataset_split"], + "dataset_revision": data["dataset_revision"], + "tokenizer": tokenizer.name, + "vocab_size": int(tokenizer.n_vocab), + "dtype": data["dtype"], + "split_tokens": { + "train": int(data["train_tokens"]), + "val": int(data["val_tokens"]), + "test": int(data["test_tokens"]), + }, + } -def main(): +def main() -> None: parser = argparse.ArgumentParser() - parser.add_argument("--dataset", default="fineweb-edu") + parser.add_argument("--config", default="configs/level0.yaml") 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("--dataset", default=None) 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("--train-tokens", type=int) + parser.add_argument("--val-tokens", type=int) + parser.add_argument("--test-tokens", type=int) + parser.add_argument("--force", action="store_true") + 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"] + config = load_config(args.config) + data = config["data"] + if args.dataset not in (None, "fineweb-edu"): + parser.error("the isolated baseline supports only --dataset fineweb-edu") + for argument, key in ( + (args.train_tokens, "train_tokens"), + (args.val_tokens, "val_tokens"), + (args.test_tokens, "test_tokens"), + ): + if argument is not None: + data[key] = int(argument) + + resolved_roots = roots() + output = Path(args.output_dir) if args.output_dir else resolved_roots["data"] + cache_root = resolved_roots["cache"] + cache_root.mkdir(parents=True, exist_ok=True) + huggingface_root = cache_root / "huggingface" + os.environ.setdefault("HF_HOME", str(huggingface_root)) + os.environ.setdefault("HF_DATASETS_CACHE", str(huggingface_root / "datasets")) + os.environ.setdefault("HF_HUB_CACHE", str(huggingface_root / "hub")) + os.environ.setdefault( + "HUGGINGFACE_HUB_CACHE", + str(huggingface_root / "hub"), + ) + os.environ.setdefault("TIKTOKEN_CACHE_DIR", str(cache_root / "tiktoken")) + Path(os.environ["TIKTOKEN_CACHE_DIR"]).mkdir(parents=True, exist_ok=True) + tokenizer = _load_tokenizer(data["tokenizer"]) + expected = expected_metadata(config, tokenizer) + if not args.force and (output / "meta.json").is_file(): + try: + validate_prepared_data( + output, + expected=expected, + verify_hashes=True, + ) + except (FileNotFoundError, ValueError): + pass + else: + print( + f"[level0-prepare-data] compatible data already exists: {output}", + file=sys.stderr, + flush=True, + ) + print(output) + return + if args.local_text: text = Path(args.local_text).read_text(encoding="utf-8") - def repeat(): + def repeat_text() -> Iterator[str]: while True: yield text - texts = repeat() + texts: Iterable[str] = repeat_text() + dataset_metadata = { + "dataset_name": "local-text", + "dataset_config": "local-text", + "dataset_split": "local-text", + "dataset_revision": _sha256_file(Path(args.local_text)), + } 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) + except ImportError as exc: # pragma: no cover - exercised by CLI users. + raise SystemExit( + "Install data support from level_0_baseline: " + "python -m pip install -e '.[data]'" + ) from exc + dataset_metadata = { + "dataset_name": data["dataset_name"], + "dataset_config": data["dataset_config"], + "dataset_split": data["dataset_split"], + "dataset_revision": data["dataset_revision"], + } + texts = _fineweb_texts( + load_dataset, + dataset_name=data["dataset_name"], + dataset_config=data["dataset_config"], + dataset_split=data["dataset_split"], + dataset_revision=data["dataset_revision"], + verbose=args.verbose, + ) - write_splits( + write_token_splits( texts, - out, - args.train_bytes, - args.val_bytes, - args.test_bytes, + output, + tokenizer=tokenizer, + train_tokens=int(data["train_tokens"]), + val_tokens=int(data["val_tokens"]), + test_tokens=int(data["test_tokens"]), + dataset_metadata=dataset_metadata, verbose=args.verbose, log_interval_seconds=args.log_interval_seconds, ) - print(out) + print(output) 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..5032e08 100644 --- a/level_0_baseline/src/level0_baseline/model.py +++ b/level_0_baseline/src/level0_baseline/model.py @@ -1,45 +1,88 @@ from __future__ import annotations -from dataclasses import dataclass + import math +from dataclasses import dataclass + import torch import torch.nn as nn import torch.nn.functional as F + @dataclass class GPTConfig: - vocab_size: int = 256 + vocab_size: int = 50257 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 = 128 dropout: float = 0.0 bias: bool = False + tie_weights: bool = True + 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.register_buffer( + "mask", + torch.tril(torch.ones(cfg.block_size, cfg.block_size, dtype=torch.bool)).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 hasattr(F, "scaled_dot_product_attention"): + 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: # pragma: no cover - modern supported PyTorch uses the branch above. + attention = (q @ k.transpose(-2, -1)) / math.sqrt(head_size) + attention = attention.masked_fill( + ~self.mask[:, :, :sequence_length, :sequence_length], + 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 +91,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 +115,65 @@ 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) + residual_std = 0.02 / math.sqrt(2 * cfg.n_layer) + for block in self.blocks: + nn.init.normal_(block.attn.out_proj.weight, mean=0.0, std=residual_std) + nn.init.normal_(block.mlp.proj.weight, mean=0.0, std=residual_std) + if cfg.tie_weights: + self.lm_head.weight = self.token_embedding.weight - def _init(self, module: nn.Module) -> None: + @staticmethod + def _init(module: nn.Module) -> None: if isinstance(module, (nn.Linear, 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: + 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("sequence exceeds block_size") - pos = torch.arange(t, device=idx.device) - x = self.drop(self.token_embedding(idx) + self.position_embedding(pos)) + 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 num_parameters(self, *, exclude_position_embedding: bool = False) -> int: + count = sum(parameter.numel() for parameter in self.parameters()) + if exclude_position_embedding: + count -= self.position_embedding.weight.numel() + return count + + def spectral_layers(self) -> list[tuple[str, nn.Linear]]: + layers: list[tuple[str, nn.Linear]] = [] + for index, block in enumerate(self.blocks): + prefix = f"block_{index:02d}" + layers.extend( + [ + (f"{prefix}_W_Q", block.attn.q_proj), + (f"{prefix}_W_K", block.attn.k_proj), + (f"{prefix}_W_V", block.attn.v_proj), + (f"{prefix}_W_O", block.attn.out_proj), + (f"{prefix}_W_MLP_IN", block.mlp.fc), + (f"{prefix}_W_MLP_OUT", block.mlp.proj), + ] + ) + return layers diff --git a/level_0_baseline/src/level0_baseline/optim.py b/level_0_baseline/src/level0_baseline/optim.py index fbc19ae..e676826 100644 --- a/level_0_baseline/src/level0_baseline/optim.py +++ b/level_0_baseline/src/level0_baseline/optim.py @@ -1,59 +1,214 @@ from __future__ import annotations + +import re +from collections import defaultdict +from typing import Any, Iterable + import torch +_BLOCK_PATTERN = re.compile(r"^blocks\.(\d+)\.") + + @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("Muon 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] | list[dict[str, Any]], + *, + lr: float = 0.02, + momentum: float = 0.95, + nesterov: bool = True, + weight_decay: float = 0.0, + ) -> None: + 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) + momentum_buffer = self.state[parameter].setdefault( + "momentum_buffer", + torch.zeros_like(parameter), + ) + momentum_buffer.mul_(group["momentum"]).add_(parameter.grad) + update_input = ( + parameter.grad.add( + momentum_buffer, + alpha=group["momentum"], + ) + if group["nesterov"] + else momentum_buffer + ) + update = zeropower_via_newtonschulz5(update_input) + 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: + +def _layer_multiplier(name: str, n_layer: int, decay: float) -> float: + if decay == 1.0: + return 1.0 + match = _BLOCK_PATTERN.match(name) + if match: + block_index = int(match.group(1)) + return decay ** max(n_layer - 1 - block_index, 0) + if name.startswith(("token_embedding", "position_embedding")): + return decay**n_layer + return 1.0 + + +def _parameter_records(model, training: dict[str, Any]): + n_layer = int(model.cfg.n_layer) + layer_decay = float(training.get("layer_lr_decay", 1.0)) + for name, parameter in model.named_parameters(): + if not parameter.requires_grad: continue - if p.ndim == 2 and "embedding" not in n and "lm_head" not in n: - muon.append(p) + yield { + "name": name, + "parameter": parameter, + "decay": parameter.ndim >= 2, + "lr_multiplier": _layer_multiplier(name, n_layer, layer_decay), + } + + +def _adamw_groups(model, training: dict[str, Any]) -> list[dict[str, Any]]: + grouped: dict[tuple[bool, float], list[torch.nn.Parameter]] = defaultdict(list) + names: dict[tuple[bool, float], list[str]] = defaultdict(list) + for record in _parameter_records(model, training): + key = (bool(record["decay"]), float(record["lr_multiplier"])) + grouped[key].append(record["parameter"]) + names[key].append(str(record["name"])) + + base_lr = float(training["learning_rate"]) + groups: list[dict[str, Any]] = [] + for (use_decay, multiplier), parameters in sorted( + grouped.items(), + key=lambda item: (item[0][1], item[0][0]), + ): + groups.append( + { + "params": parameters, + "lr": base_lr * multiplier, + "initial_lr": base_lr * multiplier, + "lr_multiplier": multiplier, + "weight_decay": ( + float(training["weight_decay"]) if use_decay else 0.0 + ), + "group_name": ( + f"{'decay' if use_decay else 'no_decay'}_lr_{multiplier:.6f}" + ), + "parameter_names": names[(use_decay, multiplier)], + } + ) + return groups + + +def make_optimizers(model, config: dict[str, Any]) -> list[torch.optim.Optimizer]: + training = config["training"] + optimizer_name = str(training["optimizer"]).lower() + if optimizer_name == "adamw": + return [ + torch.optim.AdamW( + _adamw_groups(model, training), + lr=float(training["learning_rate"]), + betas=(float(training["beta1"]), float(training["beta2"])), + eps=float(training["epsilon"]), + ) + ] + if optimizer_name != "muon": + raise ValueError(f"unsupported optimizer: {optimizer_name}") + + muon_grouped: dict[float, list[torch.nn.Parameter]] = defaultdict(list) + auxiliary_grouped: dict[tuple[bool, float], list[torch.nn.Parameter]] = defaultdict(list) + for record in _parameter_records(model, training): + name = str(record["name"]) + parameter = record["parameter"] + multiplier = float(record["lr_multiplier"]) + use_muon = ( + parameter.ndim == 2 + and not name.startswith(("token_embedding", "position_embedding", "lm_head")) + ) + if use_muon: + muon_grouped[multiplier].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_grouped[(bool(record["decay"]), multiplier)].append(parameter) + + muon_base_lr = float(training["muon_learning_rate"]) + muon_groups = [ + { + "params": parameters, + "lr": muon_base_lr * multiplier, + "initial_lr": muon_base_lr * multiplier, + "lr_multiplier": multiplier, + "weight_decay": float(training["weight_decay"]), + "group_name": f"muon_lr_{multiplier:.6f}", + } + for multiplier, parameters in sorted(muon_grouped.items()) + ] + auxiliary_base_lr = float(training["muon_aux_adamw_learning_rate"]) + auxiliary_groups = [ + { + "params": parameters, + "lr": auxiliary_base_lr * multiplier, + "initial_lr": auxiliary_base_lr * multiplier, + "lr_multiplier": multiplier, + "weight_decay": ( + float(training["weight_decay"]) if use_decay else 0.0 + ), + "group_name": ( + f"aux_{'decay' if use_decay else 'no_decay'}_lr_{multiplier:.6f}" + ), + } + for (use_decay, multiplier), parameters in sorted( + auxiliary_grouped.items(), + key=lambda item: (item[0][1], item[0][0]), + ) + ] + return [ + Muon( + muon_groups, + lr=muon_base_lr, + momentum=float(training["muon_momentum"]), + nesterov=bool(training["muon_nesterov"]), + weight_decay=float(training["weight_decay"]), + ), + torch.optim.AdamW( + auxiliary_groups, + lr=auxiliary_base_lr, + betas=(float(training["beta1"]), float(training["beta2"])), + eps=float(training["epsilon"]), + ), + ] diff --git a/level_0_baseline/src/level0_baseline/train.py b/level_0_baseline/src/level0_baseline/train.py index 2123405..a3b1aac 100644 --- a/level_0_baseline/src/level0_baseline/train.py +++ b/level_0_baseline/src/level0_baseline/train.py @@ -1,126 +1,756 @@ from __future__ import annotations -import argparse, csv, json, math, platform, random, time + +import argparse +import csv +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 +import torch.nn as nn + from .config import load_config, roots +from .data import validate_prepared_data from .model import GPT, GPTConfig from .optim import make_optimizers -def device_auto(): +RUN_SCHEMA_VERSION = 2 + + +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 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 learning_rate_at(step_index: int, training: dict[str, Any]) -> float: + warmup_steps = int(training["warmup_steps"]) + max_steps = int(training["max_steps"]) + learning_rate = float(training["learning_rate"]) + minimum_learning_rate = float(training["min_lr"]) + if step_index < warmup_steps: + return learning_rate * (step_index + 1) / max(1, warmup_steps) + if step_index >= max_steps: + return minimum_learning_rate + ratio = (step_index - warmup_steps) / max(1, max_steps - warmup_steps) + cosine = 0.5 * (1.0 + math.cos(math.pi * ratio)) + return minimum_learning_rate + cosine * ( + learning_rate - minimum_learning_rate + ) + + +def _assert_data_matches_config( + metadata: dict[str, Any], + config: dict[str, Any], +) -> None: + data = config["data"] + for key in ( + "dataset_name", + "dataset_config", + "dataset_split", + "dataset_revision", + "tokenizer", + "dtype", + ): + if metadata.get(key) != data.get(key): + raise ValueError( + f"prepared data field {key!r} does not match the configured identity" + ) + expected_splits = { + "train": int(data["train_tokens"]), + "val": int(data["val_tokens"]), + "test": int(data["test_tokens"]), + } + if metadata.get("split_tokens") != expected_splits: + raise ValueError( + "prepared data split sizes do not match the configured token counts" + ) + + +def _sample_cpu_batch( + data: np.ndarray, + *, + batch_size: int, + block_size: int, + generator: torch.Generator, +) -> tuple[torch.Tensor, torch.Tensor]: + upper = len(data) - block_size - 1 + if upper <= 0: + raise ValueError("prepared split is shorter than the configured block size") + indices = torch.randint(upper, (batch_size,), generator=generator) + x = torch.stack( + [ + torch.from_numpy( + np.asarray(data[int(index) : int(index) + block_size], dtype=np.int64) + ) + for index in indices + ] + ) + y = torch.stack( + [ + torch.from_numpy( + np.asarray( + data[int(index) + 1 : int(index) + 1 + block_size], + dtype=np.int64, + ) + ) + for index in indices + ] + ) + return x, y + + +def _to_device( + batch: tuple[torch.Tensor, torch.Tensor], + device: torch.device, +) -> tuple[torch.Tensor, torch.Tensor]: + x, y = batch + return x.to(device, non_blocking=True), y.to(device, non_blocking=True) + + +def make_fixed_probe( + data: np.ndarray, + *, + batch_size: int, + block_size: int, + batches: int, + seed: int, +) -> list[tuple[torch.Tensor, torch.Tensor]]: + generator = torch.Generator(device="cpu").manual_seed(seed) + return [ + _sample_cpu_batch( + data, + batch_size=batch_size, + block_size=block_size, + generator=generator, + ) + for _ in range(batches) + ] + @torch.no_grad() -def evaluate(model, data, batch_size, block_size, n_batches, device, generator): +def evaluate_probe( + model: nn.Module, + probe: Iterable[tuple[torch.Tensor, torch.Tensor]], + 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) - logits, loss = model(x, y) - losses.append(loss.item()) - correct += (logits.argmax(-1) == y).sum().item() - total += y.numel() - model.train() - loss = float(np.mean(losses)) - return loss, math.exp(min(20, loss)), correct / total - -def weightwatch(model, out, step, randomize): + total_nll = 0.0 + correct = 0 + token_count = 0 + try: + for cpu_batch in probe: + x, y = _to_device(cpu_batch, device) + logits, loss = model(x, y) + if loss is None: + raise RuntimeError("evaluation did not return a loss") + count = y.numel() + total_nll += float(loss.detach().cpu()) * count + correct += int((logits.argmax(dim=-1) == y).sum().detach().cpu()) + token_count += count + finally: + model.train(was_training) + loss_value = total_nll / max(token_count, 1) + return { + "loss": loss_value, + "perplexity": math.exp(min(20.0, loss_value)), + "accuracy": correct / max(token_count, 1), + "bits_per_token": loss_value / math.log(2.0), + "tokens": float(token_count), + } + + +def _gradient_norm(parameters: Iterable[torch.nn.Parameter]) -> torch.Tensor: + norms = [ + parameter.grad.detach().float().norm(2) + for parameter in parameters + if parameter.grad is not None + ] + if not norms: + return torch.tensor(0.0) + return torch.linalg.vector_norm(torch.stack(norms), ord=2) + + +def _weight_norm(parameters: Iterable[torch.nn.Parameter]) -> float: + squares = sum( + float((parameter.detach().float() ** 2).sum().cpu()) + for parameter in parameters + ) + return math.sqrt(squares) + + +def _atomic_json(path: Path, payload: dict[str, Any]) -> None: + temporary = path.with_name(f".{path.name}.tmp") + temporary.write_text( + json.dumps(payload, indent=2, sort_keys=True, allow_nan=False), + encoding="utf-8", + ) + os.replace(temporary, path) + + +class _SpectralSnapshot(nn.Module): + def __init__(self, model: GPT): + super().__init__() + layers: dict[str, nn.Linear] = {} + cpu_rng = torch.get_rng_state() + try: + for name, source in model.spectral_layers(): + copied = nn.Linear( + source.in_features, + source.out_features, + bias=False, + device="cpu", + ) + copied.weight.data.copy_( + source.weight.detach().to(device="cpu", dtype=torch.float32) + ) + copied.weight.requires_grad_(False) + layers[name] = copied + finally: + torch.set_rng_state(cpu_rng) + self.layers = nn.ModuleDict(layers) + + +def _source_layer_from_row(row: dict[str, Any], names: list[str]) -> str: + text = " ".join( + str(row.get(key, "")) for key in ("longname", "name", "layer_id") + ) + matches = [name for name in names if name in text] + return matches[0] if len(matches) == 1 else "" + + +def weightwatch( + model: GPT, + output_dir: Path, + *, + step: int, + randomize: bool, +) -> dict[str, Any]: 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']}" + except ImportError as exc: + raise RuntimeError( + "WeightWatcher is enabled but not installed; install the analysis extra" + ) from exc + + started = time.perf_counter() + snapshot = _SpectralSnapshot(model) + source_names = [name for name, _ in model.spectral_layers()] + try: + details = ww.WeightWatcher(model=snapshot).analyze( + randomize=bool(randomize) + ) + details.insert(0, "step", int(step)) + rows = details.to_dict(orient="records") + details.insert( + 1, + "source_layer", + [_source_layer_from_row(row, source_names) for row in rows], + ) + path = output_dir / f"weightwatcher_step_{step:07d}.csv" + details.to_csv(path, index=False) + return { + "step": step, + "success": True, + "rows": int(len(details)), + "path": str(path), + "elapsed_sec": time.perf_counter() - started, + } + except Exception as exc: # WeightWatcher should not erase the training run. + payload = { + "step": step, + "success": False, + "exception_type": type(exc).__name__, + "exception_message": str(exc), + "elapsed_sec": time.perf_counter() - started, + } + _atomic_json( + output_dir / f"weightwatcher_error_step_{step:07d}.json", + payload, + ) + print( + "[level0-train] WeightWatcher failed " + f"step={step}: {type(exc).__name__}: {exc}", + file=sys.stderr, + flush=True, + ) + return payload + finally: + del snapshot + + +def _optimizer_manifest(optimizers: list[torch.optim.Optimizer]) -> list[dict[str, Any]]: + rows: list[dict[str, Any]] = [] + for optimizer_index, optimizer in enumerate(optimizers): + for group_index, group in enumerate(optimizer.param_groups): + rows.append( + { + "optimizer_index": optimizer_index, + "optimizer_class": type(optimizer).__name__, + "group_index": group_index, + "group_name": group.get("group_name", f"group_{group_index}"), + "parameter_count": sum( + parameter.numel() for parameter in group["params"] + ), + "initial_lr": float(group.get("initial_lr", group["lr"])), + "lr_multiplier": float(group.get("lr_multiplier", 1.0)), + "weight_decay": float(group.get("weight_decay", 0.0)), + } + ) + return rows + + +def _save_checkpoint( + path: Path, + *, + model: GPT, + optimizers: list[torch.optim.Optimizer], + step: int, + config: dict[str, Any], + train_generator: torch.Generator, + best_validation_loss: float, + best_validation_step: int, +) -> None: + torch.save( + { + "run_schema_version": RUN_SCHEMA_VERSION, + "model": model.state_dict(), + "optimizers": [optimizer.state_dict() for optimizer in optimizers], + "step": int(step), + "config": config, + "train_generator_state": train_generator.get_state(), + "cpu_rng_state": torch.get_rng_state(), + "best_validation_loss": float(best_validation_loss), + "best_validation_step": int(best_validation_step), + }, + path, + ) + + +def _prepare_run_directory(run: Path, *, overwrite: bool) -> None: + if run.exists() and any(run.iterdir()): + if not overwrite: + raise FileExistsError( + f"run directory already contains files: {run}; " + "use --overwrite or choose a fresh results root" + ) + shutil.rmtree(run) 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) + + +def run_training( + config: dict[str, Any], + *, + data_root: str | Path, + results_root: str | Path, + device: str = "auto", + overwrite: bool = False, +) -> Path: + training = config["training"] + model_config = GPTConfig(**config["model"]) + seed = int(training["seed"]) + optimizer_name = str(training["optimizer"]).lower() + selected_device = device_auto() if device == "auto" else torch.device(device) + torch.set_float32_matmul_precision("high") + seed_all(seed) + + data_path = Path(data_root) + data_metadata = validate_prepared_data(data_path, verify_hashes=True) + _assert_data_matches_config(data_metadata, config) + if int(data_metadata["vocab_size"]) > model_config.vocab_size: + raise ValueError("model vocabulary does not cover the prepared tokenizer") + arrays = { + split: np.memmap( + data_path / f"{split}.bin", + dtype=np.uint16, + mode="r", + ) + for split in ("train", "val", "test") + } + + run = Path(results_root) / f"{optimizer_name}_seed_{seed}" + _prepare_run_directory(run, overwrite=overwrite) + raw_model = GPT(model_config).to(selected_device) + optimizers = make_optimizers(raw_model, config) + model: nn.Module = raw_model + if bool(training.get("compile", False)): + if not hasattr(torch, "compile"): + raise RuntimeError("torch.compile was requested but is unavailable") + model = torch.compile(raw_model) + + train_generator = torch.Generator(device="cpu").manual_seed(seed + 101) + probe_batch_size = int(training.get("eval_batch_size", training["batch_size"])) + train_probe = make_fixed_probe( + arrays["train"], + batch_size=probe_batch_size, + block_size=model_config.block_size, + batches=int(training["eval_batches"]), + seed=1001, + ) + validation_probe = make_fixed_probe( + arrays["val"], + batch_size=probe_batch_size, + block_size=model_config.block_size, + batches=int(training["eval_batches"]), + seed=2001, + ) + test_probe = make_fixed_probe( + arrays["test"], + batch_size=probe_batch_size, + block_size=model_config.block_size, + batches=int(training["test_eval_batches"]), + seed=3001, + ) + + tokens_per_step = ( + int(training["batch_size"]) + * model_config.block_size + * int(training["grad_accum_steps"]) + ) + max_steps = int(training["max_steps"]) + manifest = { + "run_schema_version": RUN_SCHEMA_VERSION, + "profile": "level0_gpt2_bpe_realistic_v2", + "config": config, + "device": str(selected_device), + "torch_version": torch.__version__, + "platform": platform.platform(), + "parameter_count": raw_model.num_parameters(), + "non_embedding_parameter_count": raw_model.num_parameters( + exclude_position_embedding=True + ), + "data_root": str(data_path.resolve()), + "data_metadata": data_metadata, + "tokens_per_optimizer_step": tokens_per_step, + "planned_optimizer_steps": max_steps, + "planned_training_tokens": tokens_per_step * max_steps, + "evaluation_protocol": { + "train_and_validation": "fixed_probe_periodic", + "test": "final_and_validation_selected_only", + "training_rng_isolated_from_evaluation": True, + "fixed_probe_seeds": {"train": 1001, "val": 2001, "test": 3001}, + }, + "optimizer_groups": _optimizer_manifest(optimizers), + "accuracy_definition": "exact_top1_next_gpt2_bpe_token", + } + _atomic_json(run / "manifest.json", manifest) + + metric_fields = [ + "step", + "tokens_seen", + "elapsed_sec", + "tokens_per_sec", + "learning_rate", + "min_group_learning_rate", + "max_group_learning_rate", + "train_loss", + "train_perplexity", + "train_accuracy", + "train_bits_per_token", + "val_loss", + "val_perplexity", + "val_accuracy", + "val_bits_per_token", + "val_generalization_gap", + "grad_norm", + "weight_norm", + ] + best_validation_loss = float("inf") + best_validation_step = -1 + last_gradient_norm = 0.0 + weightwatcher_records: list[dict[str, Any]] = [] + start_time = time.perf_counter() + latest_metrics: dict[str, Any] = {} + + initial_scale = learning_rate_at(0, training) / float(training["learning_rate"]) + for optimizer in optimizers: + for group in optimizer.param_groups: + group["lr"] = float(group["initial_lr"]) * initial_scale + + print( + "[level0-train] starting " + f"optimizer={optimizer_name} seed={seed} device={selected_device} " + f"parameters={manifest['parameter_count']:,} steps={max_steps:,} " + f"tokens_per_step={tokens_per_step:,} " + f"planned_tokens={tokens_per_step * max_steps:,}", + file=sys.stderr, + flush=True, + ) + + with (run / "metrics.csv").open("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"]: + + for step in range(max_steps + 1): + if step == 0 or step % int(training["eval_interval"]) == 0 or step == max_steps: + train_metrics = evaluate_probe(model, train_probe, selected_device) + validation_metrics = evaluate_probe( + model, + validation_probe, + selected_device, + ) + elapsed = time.perf_counter() - start_time + group_lrs = [ + float(group["lr"]) + for optimizer in optimizers + for group in optimizer.param_groups + ] + current_lr = learning_rate_at(max(step - 1, 0), training) + latest_metrics = { + "step": step, + "tokens_seen": step * tokens_per_step, + "elapsed_sec": elapsed, + "tokens_per_sec": ( + step * tokens_per_step / max(elapsed, 1e-9) + ), + "learning_rate": current_lr, + "min_group_learning_rate": min(group_lrs), + "max_group_learning_rate": max(group_lrs), + "train_loss": train_metrics["loss"], + "train_perplexity": train_metrics["perplexity"], + "train_accuracy": train_metrics["accuracy"], + "train_bits_per_token": train_metrics["bits_per_token"], + "val_loss": validation_metrics["loss"], + "val_perplexity": validation_metrics["perplexity"], + "val_accuracy": validation_metrics["accuracy"], + "val_bits_per_token": validation_metrics["bits_per_token"], + "val_generalization_gap": ( + validation_metrics["loss"] - train_metrics["loss"] + ), + "grad_norm": last_gradient_norm, + "weight_norm": _weight_norm(raw_model.parameters()), + } + writer.writerow(latest_metrics) + handle.flush() + + if validation_metrics["loss"] < best_validation_loss: + best_validation_loss = validation_metrics["loss"] + best_validation_step = step + torch.save( + { + "model": raw_model.state_dict(), + "step": step, + "validation_metrics": validation_metrics, + "config": config, + }, + run / "checkpoint_best.pt", + ) + + remaining_steps = max_steps - step + steps_per_second = step / max(elapsed, 1e-9) + eta = ( + remaining_steps / steps_per_second + if steps_per_second > 0 + else None + ) + eta_text = f"{eta:.1f}s" if eta is not None else "unknown" + print( + "[level0-train] progress " + f"step={step:,}/{max_steps:,} " + f"tokens={step * tokens_per_step:,} " + f"train_loss={train_metrics['loss']:.4f} " + f"val_loss={validation_metrics['loss']:.4f} " + f"val_ppl={validation_metrics['perplexity']:.2f} " + f"val_acc={100 * validation_metrics['accuracy']:.2f}% " + f"elapsed={elapsed:.1f}s eta={eta_text}", + file=sys.stderr, + flush=True, + ) + + weightwatcher_due = ( + bool(config["analysis"]["weightwatcher"]) + and ( + step % int(config["analysis"]["weightwatcher_interval"]) == 0 + or step == max_steps + ) + ) + if weightwatcher_due: + weightwatcher_records.append( + weightwatch( + raw_model, + run, + step=step, + randomize=bool(config["analysis"]["randomize"]), + ) + ) + + if step == 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) + + scheduled_base_lr = learning_rate_at(step, training) + scale = scheduled_base_lr / float(training["learning_rate"]) + for optimizer in optimizers: + for group in optimizer.param_groups: + group["lr"] = float(group["initial_lr"]) * scale + optimizer.zero_grad(set_to_none=True) + + model.train() + for _ in range(int(training["grad_accum_steps"])): + cpu_batch = _sample_cpu_batch( + arrays["train"], + batch_size=int(training["batch_size"]), + block_size=model_config.block_size, + generator=train_generator, + ) + x, y = _to_device(cpu_batch, selected_device) _, 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") + if loss is None: + raise RuntimeError("training did not return a loss") + (loss / int(training["grad_accum_steps"])).backward() + + if float(training["grad_clip"]) > 0: + gradient = torch.nn.utils.clip_grad_norm_( + raw_model.parameters(), + float(training["grad_clip"]), + ) + else: + gradient = _gradient_norm(raw_model.parameters()) + last_gradient_norm = float(gradient.detach().cpu()) + for optimizer in optimizers: + optimizer.step() + + completed_step = step + 1 + if completed_step % int(training["checkpoint_interval"]) == 0: + _save_checkpoint( + run / f"checkpoint_{completed_step:07d}.pt", + model=raw_model, + optimizers=optimizers, + step=completed_step, + config=config, + train_generator=train_generator, + best_validation_loss=best_validation_loss, + best_validation_step=best_validation_step, + ) + + _save_checkpoint( + run / "checkpoint_final.pt", + model=raw_model, + optimizers=optimizers, + step=max_steps, + config=config, + train_generator=train_generator, + best_validation_loss=best_validation_loss, + best_validation_step=best_validation_step, + ) + + final_test_metrics = evaluate_probe(model, test_probe, selected_device) + final_metrics = { + "step": max_steps, + "validation_loss": float(latest_metrics["val_loss"]), + "validation_perplexity": float(latest_metrics["val_perplexity"]), + "validation_accuracy": float(latest_metrics["val_accuracy"]), + "test_loss": final_test_metrics["loss"], + "test_perplexity": final_test_metrics["perplexity"], + "test_accuracy": final_test_metrics["accuracy"], + "test_bits_per_token": final_test_metrics["bits_per_token"], + } + _atomic_json(run / "final_metrics.json", final_metrics) + + best_checkpoint = torch.load( + run / "checkpoint_best.pt", + map_location="cpu", + weights_only=False, + ) + raw_model.load_state_dict(best_checkpoint["model"]) + selected_test_metrics = evaluate_probe(model, test_probe, selected_device) + selected_metrics = { + "selected_step": int(best_checkpoint["step"]), + "validation_loss": float( + best_checkpoint["validation_metrics"]["loss"] + ), + "validation_perplexity": float( + best_checkpoint["validation_metrics"]["perplexity"] + ), + "validation_accuracy": float( + best_checkpoint["validation_metrics"]["accuracy"] + ), + "test_loss": selected_test_metrics["loss"], + "test_perplexity": selected_test_metrics["perplexity"], + "test_accuracy": selected_test_metrics["accuracy"], + "test_bits_per_token": selected_test_metrics["bits_per_token"], + } + _atomic_json(run / "selected_checkpoint_metrics.json", selected_metrics) + + completion = { + "completed": True, + "run_schema_version": RUN_SCHEMA_VERSION, + "optimizer": optimizer_name, + "seed": seed, + "optimizer_steps": max_steps, + "tokens_seen": max_steps * tokens_per_step, + "best_validation_step": selected_metrics["selected_step"], + "best_validation_loss": selected_metrics["validation_loss"], + "final_validation_loss": final_metrics["validation_loss"], + "selected_test_loss": selected_metrics["test_loss"], + "final_test_loss": final_metrics["test_loss"], + "validation_collapse_final_minus_selected": ( + final_metrics["validation_loss"] - selected_metrics["validation_loss"] + ), + "test_collapse_final_minus_selected": ( + final_metrics["test_loss"] - selected_metrics["test_loss"] + ), + "weightwatcher_measurements": len(weightwatcher_records), + "weightwatcher_failures": sum( + not bool(record.get("success")) for record in weightwatcher_records + ), + "elapsed_sec": time.perf_counter() - start_time, + } + _atomic_json(run / "run_complete.json", completion) + print( + "[level0-train] complete " + f"run={run} selected_step={selected_metrics['selected_step']} " + f"selected_test_loss={selected_metrics['test_loss']:.4f} " + f"final_test_loss={final_metrics['test_loss']:.4f}", + file=sys.stderr, + flush=True, + ) + return run + + +def main() -> None: + parser = argparse.ArgumentParser() + 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("--overwrite", action="store_true") + args = parser.parse_args() + + config = load_config(args.config) + if args.optimizer: + config["training"]["optimizer"] = args.optimizer + if args.seed is not None: + config["training"]["seed"] = args.seed + resolved = roots() + run = run_training( + config, + data_root=args.data_root or resolved["data"], + results_root=args.results_root or resolved["results"], + device=args.device, + overwrite=args.overwrite, + ) + print(run) + if __name__ == "__main__": main() diff --git a/level_0_baseline/tests/test_baseline.py b/level_0_baseline/tests/test_baseline.py index 76d8995..abc2c19 100644 --- a/level_0_baseline/tests/test_baseline.py +++ b/level_0_baseline/tests/test_baseline.py @@ -1,97 +1,285 @@ +from __future__ import annotations + import json import sys from pathlib import Path +import numpy as np +import pandas as pd +import pytest import torch -sys.path.insert(0, str(Path(__file__).resolve().parents[1] / "src")) +PACKAGE_ROOT = Path(__file__).resolve().parents[1] +sys.path.insert(0, str(PACKAGE_ROOT / "src")) -from level0_baseline.data import _progress_message, write_splits +from level0_baseline.config import load_config +from level0_baseline.data import ( + _progress_message, + validate_prepared_data, + write_token_splits, +) from level0_baseline.model import GPT, GPTConfig from level0_baseline.optim import make_optimizers +from level0_baseline.train import ( + _SpectralSnapshot, + make_fixed_probe, + run_training, +) + + +class FakeTokenizer: + name = "gpt2" + n_vocab = 32 + eot_token = 0 + @staticmethod + def encode_ordinary(text: str) -> list[int]: + return [1 + (ord(character) % 31) for character in text] -def cfg(opt): + +def tiny_training_config() -> dict: return { + "data": { + "dataset_name": "unit", + "dataset_config": "unit", + "dataset_split": "train", + "dataset_revision": "unit", + "tokenizer": "gpt2", + "dtype": "uint16", + "train_tokens": 256, + "val_tokens": 128, + "test_tokens": 128, + }, + "model": { + "vocab_size": 32, + "block_size": 8, + "n_layer": 1, + "n_head": 1, + "n_embd": 16, + "dropout": 0.0, + "bias": False, + "tie_weights": True, + }, "training": { - "optimizer": opt, + "batch_size": 2, + "grad_accum_steps": 2, + "max_steps": 2, + "eval_interval": 1, + "eval_batches": 2, + "eval_batch_size": 2, + "test_eval_batches": 2, + "checkpoint_interval": 1, "learning_rate": 0.001, - "weight_decay": 0.1, + "muon_learning_rate": 0.02, + "muon_aux_adamw_learning_rate": 0.001, + "min_lr": 0.0001, + "warmup_steps": 1, + "weight_decay": 0.01, "beta1": 0.9, "beta2": 0.95, + "epsilon": 1e-8, + "grad_clip": 1.0, + "layer_lr_decay": 1.0, + "optimizer": "adamw", "muon_momentum": 0.95, "muon_nesterov": True, - "muon_learning_rate": 0.02, - "muon_aux_adamw_learning_rate": 0.001, - } + "seed": 7, + "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 prepare_tiny_data(path: Path) -> None: + write_token_splits( + iter(["abcdefghijklmnopqrstuvwxyz"] * 64), + path, + tokenizer=FakeTokenizer(), + train_tokens=256, + val_tokens=128, + test_tokens=128, + dataset_metadata={ + "dataset_name": "unit", + "dataset_config": "unit", + "dataset_split": "train", + "dataset_revision": "unit", + }, + ) -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 test_muon_partition_and_step(): - model = GPT(GPTConfig(block_size=8, n_embd=16)) - optimizers = make_optimizers(model, cfg("muon")) - assert len(optimizers) == 2 - _, loss = model( - torch.randint(0, 256, (2, 8)), - torch.randint(0, 256, (2, 8)), - ) - loss.backward() - for optimizer in optimizers: - optimizer.step() +def test_default_config_is_realistic_gpt2_bpe_baseline(): + config = load_config(PACKAGE_ROOT / "configs" / "level0.yaml") + assert config["data"]["tokenizer"] == "gpt2" + assert config["data"]["dtype"] == "uint16" + assert config["model"] == { + "vocab_size": 50257, + "block_size": 256, + "n_layer": 4, + "n_head": 4, + "n_embd": 128, + "dropout": 0.0, + "bias": False, + "tie_weights": True, + } + assert config["training"]["batch_size"] == 4 + assert config["training"]["grad_accum_steps"] == 8 + assert config["training"]["layer_lr_decay"] == 1.0 + assert config["analysis"]["randomize"] is False + + +def test_model_size_and_weightwatcher_scope_are_not_toy_byte_baseline(): + model = GPT(GPTConfig()) + assert model.num_parameters() == 7_253_248 + assert len(model.blocks) == 4 + assert len(model.spectral_layers()) == 24 + names = [name for name, _ in model.spectral_layers()] + assert "block_00_W_Q" in names + assert "block_03_W_MLP_OUT" in names + snapshot = _SpectralSnapshot(model) + assert len(snapshot.layers) == 24 + assert "token_embedding" not in snapshot.layers + + +def test_adamw_uses_decoupled_decay_and_explicit_lr_groups(): + config = tiny_training_config() + model = GPT(GPTConfig(**config["model"])) + optimizer = make_optimizers(model, config)[0] + assert isinstance(optimizer, torch.optim.AdamW) + assert {float(group["weight_decay"]) for group in optimizer.param_groups} == { + 0.0, + 0.01, + } + assert all(float(group["lr_multiplier"]) == 1.0 for group in optimizer.param_groups) -def test_progress_message_reports_elapsed_eta_and_stall(): +def test_progress_message_reports_tokens_elapsed_eta_and_stall(): message = _progress_message( - collected_bytes=50, - required_bytes=100, + collected_tokens=50, + required_tokens=100, documents=4, elapsed_seconds=10, stalled_seconds=3, ) assert "documents=4" in message + assert "tokens=50/100" 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_gpt2_token_split_preparation_is_atomic_and_valid(tmp_path, capsys): + metadata = write_token_splits( + iter(["abcdef"] * 10), tmp_path, - train_bytes=2, - val_bytes=2, - test_bytes=2, + tokenizer=FakeTokenizer(), + train_tokens=8, + val_tokens=4, + test_tokens=4, + dataset_metadata={ + "dataset_name": "unit", + "dataset_config": "unit", + "dataset_split": "train", + "dataset_revision": "unit", + }, verbose=True, log_interval_seconds=60, ) - captured = capsys.readouterr() assert "[level0-prepare-data] starting" in captured.err - assert "[level0-prepare-data] collection complete" in captured.err - assert "[level0-prepare-data] complete" in captured.err - assert (tmp_path / "train.bin").stat().st_size == 2 - assert (tmp_path / "val.bin").stat().st_size == 2 - assert (tmp_path / "test.bin").stat().st_size == 2 - metadata = json.loads((tmp_path / "meta.json").read_text()) - assert metadata["sizes"] == {"train": 2, "val": 2, "test": 2} + assert "tokenization complete" in captured.err + assert metadata["tokenizer"] == "gpt2" + assert metadata["dtype"] == "uint16" + assert metadata["document_disjoint_splits"] is True + assert metadata["boundary_tokens_discarded"] > 0 + assert (tmp_path / "train.bin").stat().st_size == 16 + assert (tmp_path / "val.bin").stat().st_size == 8 + assert (tmp_path / "test.bin").stat().st_size == 8 + assert not list(tmp_path.glob("*.tmp")) + validated = validate_prepared_data(tmp_path, verify_hashes=True) + assert validated["split_tokens"] == {"train": 8, "val": 4, "test": 4} + + +def test_old_byte_data_is_rejected(tmp_path): + (tmp_path / "meta.json").write_text( + json.dumps({"tokenizer": "utf8-byte", "dtype": "uint8"}) + ) + with pytest.raises(ValueError, match="corrected GPT-2-BPE"): + validate_prepared_data(tmp_path) + + +def test_fixed_probes_do_not_advance_training_rng(): + data = np.arange(256, dtype=np.uint16) + training_generator = torch.Generator().manual_seed(123) + state_before = training_generator.get_state().clone() + first = make_fixed_probe( + data, + batch_size=2, + block_size=8, + batches=2, + seed=456, + ) + second = make_fixed_probe( + data, + batch_size=2, + block_size=8, + batches=2, + seed=456, + ) + assert torch.equal(training_generator.get_state(), state_before) + assert all( + torch.equal(a_tensor, b_tensor) + for a_batch, b_batch in zip(first, second) + for a_tensor, b_tensor in zip(a_batch, b_batch) + ) + + +def test_two_step_training_writes_complete_scientific_artifacts(tmp_path): + data_root = tmp_path / "data" + results_root = tmp_path / "results" + prepare_tiny_data(data_root) + run = run_training( + tiny_training_config(), + data_root=data_root, + results_root=results_root, + device="cpu", + ) + + completion = json.loads((run / "run_complete.json").read_text()) + selected = json.loads((run / "selected_checkpoint_metrics.json").read_text()) + final = json.loads((run / "final_metrics.json").read_text()) + metrics = pd.read_csv(run / "metrics.csv") + manifest = json.loads((run / "manifest.json").read_text()) + + assert completion["completed"] is True + assert completion["optimizer_steps"] == 2 + assert completion["tokens_seen"] == 64 + assert completion["weightwatcher_measurements"] == 0 + assert len(metrics) == 3 + assert list(metrics["step"]) == [0, 1, 2] + assert not any(column.startswith("test_") for column in metrics.columns) + assert selected["selected_step"] in {0, 1, 2} + assert np.isfinite(selected["test_loss"]) + assert np.isfinite(final["test_loss"]) + assert manifest["evaluation_protocol"]["test"] == ( + "final_and_validation_selected_only" + ) + assert manifest["evaluation_protocol"][ + "training_rng_isolated_from_evaluation" + ] is True + assert (run / "checkpoint_best.pt").is_file() + assert (run / "checkpoint_final.pt").is_file() + + +def test_nonempty_run_requires_explicit_overwrite(tmp_path): + data_root = tmp_path / "data" + results_root = tmp_path / "results" + prepare_tiny_data(data_root) + config = tiny_training_config() + run_training(config, data_root=data_root, results_root=results_root, device="cpu") + with pytest.raises(FileExistsError, match="--overwrite"): + run_training(config, data_root=data_root, results_root=results_root, device="cpu")