From 2ba1b36540a87893b3dbd03f21e8c3e833e12636 Mon Sep 17 00:00:00 2001 From: Charles Martin Date: Tue, 11 Aug 2026 13:06:21 -0700 Subject: [PATCH] Add nanoGPT Muon-HyperBall baseline --- .github/workflows/nanogpt-muon-hyperball.yml | 44 ++ baseline/nanogpt_muon_hyperball/.gitignore | 6 + baseline/nanogpt_muon_hyperball/README.md | 208 ++++++++ .../configs/reference.yaml | 105 +++++ .../notebooks/01_run_baseline.ipynb | 68 +++ .../notebooks/02_compare_muon_hyperball.ipynb | 132 ++++++ .../nanogpt_muon_hyperball/pyproject.toml | 32 ++ .../nanogpt_muon_hyperball/requirements.txt | 11 + .../src/rg_nanogpt_muon_hyperball/__init__.py | 13 + .../src/rg_nanogpt_muon_hyperball/analysis.py | 405 ++++++++++++++++ .../base_optimizers.py | 289 ++++++++++++ .../rg_nanogpt_muon_hyperball/checkpoints.py | 145 ++++++ .../rg_nanogpt_muon_hyperball/completion.py | 323 +++++++++++++ .../src/rg_nanogpt_muon_hyperball/config.py | 278 +++++++++++ .../src/rg_nanogpt_muon_hyperball/data.py | 328 +++++++++++++ .../src/rg_nanogpt_muon_hyperball/engine.py | 325 +++++++++++++ .../rg_nanogpt_muon_hyperball/evaluation.py | 175 +++++++ .../src/rg_nanogpt_muon_hyperball/model.py | 183 +++++++ .../rg_nanogpt_muon_hyperball/optimizers.py | 289 ++++++++++++ .../rg_nanogpt_muon_hyperball/run_utils.py | 179 +++++++ .../src/rg_nanogpt_muon_hyperball/runtime.py | 101 ++++ .../src/rg_nanogpt_muon_hyperball/spectral.py | 317 +++++++++++++ .../rg_nanogpt_muon_hyperball/train_loop.py | 446 ++++++++++++++++++ .../src/rg_nanogpt_muon_hyperball/training.py | 131 +++++ .../tests/test_config.py | 45 ++ .../tests/test_hyperball.py | 111 +++++ .../tests/test_metrics.py | 17 + .../tests/test_partition.py | 47 ++ 28 files changed, 4753 insertions(+) create mode 100644 .github/workflows/nanogpt-muon-hyperball.yml create mode 100644 baseline/nanogpt_muon_hyperball/.gitignore create mode 100644 baseline/nanogpt_muon_hyperball/README.md create mode 100644 baseline/nanogpt_muon_hyperball/configs/reference.yaml create mode 100644 baseline/nanogpt_muon_hyperball/notebooks/01_run_baseline.ipynb create mode 100644 baseline/nanogpt_muon_hyperball/notebooks/02_compare_muon_hyperball.ipynb create mode 100644 baseline/nanogpt_muon_hyperball/pyproject.toml create mode 100644 baseline/nanogpt_muon_hyperball/requirements.txt create mode 100644 baseline/nanogpt_muon_hyperball/src/rg_nanogpt_muon_hyperball/__init__.py create mode 100644 baseline/nanogpt_muon_hyperball/src/rg_nanogpt_muon_hyperball/analysis.py create mode 100644 baseline/nanogpt_muon_hyperball/src/rg_nanogpt_muon_hyperball/base_optimizers.py create mode 100644 baseline/nanogpt_muon_hyperball/src/rg_nanogpt_muon_hyperball/checkpoints.py create mode 100644 baseline/nanogpt_muon_hyperball/src/rg_nanogpt_muon_hyperball/completion.py create mode 100644 baseline/nanogpt_muon_hyperball/src/rg_nanogpt_muon_hyperball/config.py create mode 100644 baseline/nanogpt_muon_hyperball/src/rg_nanogpt_muon_hyperball/data.py create mode 100644 baseline/nanogpt_muon_hyperball/src/rg_nanogpt_muon_hyperball/engine.py create mode 100644 baseline/nanogpt_muon_hyperball/src/rg_nanogpt_muon_hyperball/evaluation.py create mode 100644 baseline/nanogpt_muon_hyperball/src/rg_nanogpt_muon_hyperball/model.py create mode 100644 baseline/nanogpt_muon_hyperball/src/rg_nanogpt_muon_hyperball/optimizers.py create mode 100644 baseline/nanogpt_muon_hyperball/src/rg_nanogpt_muon_hyperball/run_utils.py create mode 100644 baseline/nanogpt_muon_hyperball/src/rg_nanogpt_muon_hyperball/runtime.py create mode 100644 baseline/nanogpt_muon_hyperball/src/rg_nanogpt_muon_hyperball/spectral.py create mode 100644 baseline/nanogpt_muon_hyperball/src/rg_nanogpt_muon_hyperball/train_loop.py create mode 100644 baseline/nanogpt_muon_hyperball/src/rg_nanogpt_muon_hyperball/training.py create mode 100644 baseline/nanogpt_muon_hyperball/tests/test_config.py create mode 100644 baseline/nanogpt_muon_hyperball/tests/test_hyperball.py create mode 100644 baseline/nanogpt_muon_hyperball/tests/test_metrics.py create mode 100644 baseline/nanogpt_muon_hyperball/tests/test_partition.py diff --git a/.github/workflows/nanogpt-muon-hyperball.yml b/.github/workflows/nanogpt-muon-hyperball.yml new file mode 100644 index 0000000..d61e1e3 --- /dev/null +++ b/.github/workflows/nanogpt-muon-hyperball.yml @@ -0,0 +1,44 @@ +name: nanoGPT Muon-HyperBall baseline + +on: + pull_request: + paths: + - 'baseline/nanogpt_muon_hyperball/**' + - '.github/workflows/nanogpt-muon-hyperball.yml' + push: + branches: [main] + paths: + - 'baseline/nanogpt_muon_hyperball/**' + - '.github/workflows/nanogpt-muon-hyperball.yml' + +jobs: + test: + runs-on: ubuntu-latest + steps: + - uses: actions/checkout@v4 + - uses: actions/setup-python@v5 + with: + python-version: '3.11' + cache: pip + cache-dependency-path: baseline/nanogpt_muon_hyperball/pyproject.toml + - name: Install baseline + working-directory: baseline/nanogpt_muon_hyperball + run: python -m pip install -e '.[dev]' + - name: Run tests + working-directory: baseline/nanogpt_muon_hyperball + run: python -m pytest -q tests + - name: Validate notebooks + working-directory: baseline/nanogpt_muon_hyperball + run: | + python - <<'PY' + import ast + import json + from pathlib import Path + + for path in Path('notebooks').glob('*.ipynb'): + notebook = json.loads(path.read_text(encoding='utf-8')) + for index, cell in enumerate(notebook['cells']): + if cell.get('cell_type') == 'code': + ast.parse(''.join(cell.get('source', [])), filename=f'{path}:cell{index}') + print('validated', path) + PY diff --git a/baseline/nanogpt_muon_hyperball/.gitignore b/baseline/nanogpt_muon_hyperball/.gitignore new file mode 100644 index 0000000..ef8b36a --- /dev/null +++ b/baseline/nanogpt_muon_hyperball/.gitignore @@ -0,0 +1,6 @@ +__pycache__/ +*.py[cod] +.ipynb_checkpoints/ +results/ +plots/ +*.out.ipynb diff --git a/baseline/nanogpt_muon_hyperball/README.md b/baseline/nanogpt_muon_hyperball/README.md new file mode 100644 index 0000000..e7672d8 --- /dev/null +++ b/baseline/nanogpt_muon_hyperball/README.md @@ -0,0 +1,208 @@ +# nanoGPT Muon-HyperBall baseline + +This is a **separate** one-block, one-head nanoGPT experiment family for a +matched comparison of: + +1. ordinary Muon + auxiliary AdamW; and +2. Muon + a relative Frobenius HyperBall projection + auxiliary AdamW. + +It does not modify the existing `baseline/nanogpt_one_head` controls or their +results. + +## Why this experiment exists + +The one-epoch Muon runs produced highly variable layer-wise power-law exponents. +A naïvely stretched ten-epoch learning-rate schedule was also unstable: because +the original five-percent warmup was interpreted as five percent of the full +ten-epoch horizon, Muon's matrix LR kept increasing until epoch 0.5 and the run +became non-finite before the first half epoch completed. + +This baseline separates the **training horizon** from the **LR-schedule +horizon**: + +- train for ten corpus-equivalent epochs; +- reproduce the original 488-step Muon warmup; +- finish the original cosine decay at epoch 1; +- hold the configured LR floors through epochs 1–10. + +That preserves the known one-epoch trajectory and then gives the weights a long +low-LR relaxation period in which alpha can stabilize. + +## HyperBall update + +For each hidden transformer matrix, ordinary Muon first proposes the complete +next value, including multiplicative matrix weight decay: + +\[ +W_t^\star = \operatorname{MuonStep}(W_t), +\qquad +\Delta W_t = W_t^\star-W_t . +\] + +HyperBall then projects that displacement into a relative Frobenius ball: + +\[ +\Delta W_t^{\mathrm{HB}} += +\Delta W_t +\min\left( +1, +\frac{\rho\lVert W_t\rVert_F} +{\lVert\Delta W_t\rVert_F+\epsilon} +\right). +\] + +The reference radius is `rho = 0.01`, so each hidden matrix moves by at most one +percent of its current Frobenius norm per optimizer step. The projection is +radial: it does not rotate Muon's update and does not use alpha, ERG gap, +trace-log, validation loss, or any test measurement. + +Muon acts on: + +```text +W_Q, W_K, W_V, W_O, W_MLP_IN, W_MLP_OUT +``` + +The tied embedding/head, normalization gains, and all remaining parameters use +the same auxiliary AdamW path in both arms. + +## Reference protocol + +- pinned FineWeb-Edu 80M / 1M / 1M token train/validation/test split; +- one transformer block, one attention head, width 128, context 256; +- seed `1337`; +- ten corpus-equivalent training epochs; +- WeightWatcher every quarter epoch; +- matrix LR `0.02 -> 0.002` over the first epoch, then floor; +- auxiliary LR `3e-4 -> 3e-5` over the first epoch, then floor; +- validation loss selects `checkpoint_best.pt`; +- fixed test probes are monitoring-only. + +The default output root is: + +```text +/tmp/rg-nanogpt-muon-hyperball +``` + +## Install + +From the repository root: + +```bash +cd baseline/nanogpt_muon_hyperball +python -m pip install -e . +export PYTORCH_ENABLE_MPS_FALLBACK=1 +``` + +Reuse the already verified corpus: + +```text +/tmp/rg-nanogpt-one-head/data +``` + +## Run the ordinary long-Muon control + +```bash +python -u -m rg_nanogpt_muon_hyperball.training \ + --config configs/reference.yaml \ + --optimizer muon \ + --seeds 1337 \ + --data-root /tmp/rg-nanogpt-one-head/data \ + --results-root /tmp/rg-nanogpt-muon-hyperball/results \ + --device auto +``` + +## Run Muon-HyperBall + +```bash +python -u -m rg_nanogpt_muon_hyperball.training \ + --config configs/reference.yaml \ + --optimizer muon_hyperball \ + --seeds 1337 \ + --data-root /tmp/rg-nanogpt-one-head/data \ + --results-root /tmp/rg-nanogpt-muon-hyperball/results \ + --device auto +``` + +For an unattended MacBook run: + +```bash +mkdir -p /tmp/rg-nanogpt-muon-hyperball/logs + +nohup caffeinate -i \ +python -u -m rg_nanogpt_muon_hyperball.training \ + --config configs/reference.yaml \ + --optimizer muon_hyperball \ + --seeds 1337 \ + --data-root /tmp/rg-nanogpt-one-head/data \ + --results-root /tmp/rg-nanogpt-muon-hyperball/results \ + --device auto \ + > /tmp/rg-nanogpt-muon-hyperball/logs/muon_hyperball_seed_1337.log 2>&1 & + +echo $! +``` + +Watch it with: + +```bash +tail -f /tmp/rg-nanogpt-muon-hyperball/logs/muon_hyperball_seed_1337.log +``` + +Do not resume the failed stretched-schedule run. It already contains non-finite +weights. This experiment uses a new result namespace and a different protocol +fingerprint. + +## Recorded HyperBall diagnostics + +Each evaluation row records: + +```text +hyperball_relative_radius +hyperball_matrix_updates_since_eval +hyperball_active_fraction +hyperball_mean_scale +hyperball_min_scale +hyperball_mean_radius +hyperball_max_proposed_update_to_weight_ratio +hyperball_max_applied_update_to_weight_ratio +hyperball_max_proposed_update_norm +hyperball_max_applied_update_norm +``` + +The training loop aborts on non-finite train or validation metrics **before** +calling WeightWatcher, so a numerical optimizer failure is reported directly +rather than surfacing later as `numpy.linalg.LinAlgError: SVD did not converge`. + +## Analysis + +Run: + +```bash +jupyter lab notebooks +``` + +The comparison notebook plots task metrics, per-layer alpha and `D`, ERG gap, +traps, and HyperBall activation statistics. + +## Radius ablation + +Use separate config files and separate result roots for: + +```text +rho = 0.005 +rho = 0.010 # reference +rho = 0.020 +``` + +Select only on validation loss. Test metrics remain protected monitoring +measurements. + +## Tests + +```bash +python -m pytest -q tests +``` + +The tests cover the Frobenius cap, direction preservation, equivalence to +ordinary Muon at effectively infinite radius, optimizer partitioning, schedule +horizon, configuration validation, and notebook syntax. diff --git a/baseline/nanogpt_muon_hyperball/configs/reference.yaml b/baseline/nanogpt_muon_hyperball/configs/reference.yaml new file mode 100644 index 0000000..b2e8506 --- /dev/null +++ b/baseline/nanogpt_muon_hyperball/configs/reference.yaml @@ -0,0 +1,105 @@ +protocol: + name: rg_nanogpt_one_head_muon_hyperball + version: 1 + description: >- + Matched ten-corpus-equivalent one-block, one-head nanoGPT controls for + ordinary Muon and Muon plus a relative Frobenius HyperBall projection. + +dataset: + name: HuggingFaceFW/fineweb-edu + config: sample-10BT + split: train + revision: 593b3a867298afb8ce42625a270ef20ddcad28f9 + tokenizer: gpt2 + train_tokens: 80000000 + val_tokens: 1000000 + test_tokens: 1000000 + +model: + vocab_size: 50257 + block_size: 256 + n_layer: 1 + n_head: 1 + n_embd: 128 + dropout: 0.0 + bias: false + tie_weights: true + +training: + seeds: [1337] + batch_size: 4 + grad_accum_steps: 8 + target_epochs: 10.0 + epoch_interval: 0.25 + eval_interval_steps: 250 + eval_batches: 64 + checkpoint_interval_steps: 250 + grad_clip: 1.0 + +optimizer_profiles: + muon: + display_name: Long Muon + auxiliary AdamW control + family: muon + matrix_learning_rate: 0.02 + matrix_min_learning_rate: 0.002 + aux_learning_rate: 0.0003 + aux_min_learning_rate: 0.00003 + warmup_fraction: 0.05 + warmup_steps: 488 + lr_schedule_epochs: 1.0 + schedule: warmup_cosine_then_floor + momentum: 0.95 + nesterov: true + newton_schulz_steps: 5 + muon_epsilon: 1.0e-7 + matrix_weight_decay: 0.01 + beta1: 0.90 + beta2: 0.95 + epsilon: 1.0e-8 + aux_weight_decay: 0.01 + + muon_hyperball: + display_name: Muon + relative Frobenius HyperBall + auxiliary AdamW + family: muon_hyperball + matrix_learning_rate: 0.02 + matrix_min_learning_rate: 0.002 + aux_learning_rate: 0.0003 + aux_min_learning_rate: 0.00003 + warmup_fraction: 0.05 + warmup_steps: 488 + lr_schedule_epochs: 1.0 + schedule: warmup_cosine_then_floor + momentum: 0.95 + nesterov: true + newton_schulz_steps: 5 + muon_epsilon: 1.0e-7 + matrix_weight_decay: 0.01 + beta1: 0.90 + beta2: 0.95 + epsilon: 1.0e-8 + aux_weight_decay: 0.01 + hyperball_relative_radius: 0.01 + hyperball_epsilon: 1.0e-12 + +evaluation: + train_probe_seed: 21001 + validation_probe_seed: 22001 + test_probe_seed: 23001 + bleu_probe_seed: 24001 + bleu_examples: 64 + bleu_prompt_tokens: 64 + bleu_continuation_tokens: 32 + bleu_batch_size: 4 + +weightwatcher: + enabled: true + ERG: true + randomize: true + strict: true + min_evals: 20 + +runtime: + matmul_precision: high + mps_fallback: true + deterministic_algorithms: false + empty_mps_cache_after_weightwatcher: true diff --git a/baseline/nanogpt_muon_hyperball/notebooks/01_run_baseline.ipynb b/baseline/nanogpt_muon_hyperball/notebooks/01_run_baseline.ipynb new file mode 100644 index 0000000..ce9c810 --- /dev/null +++ b/baseline/nanogpt_muon_hyperball/notebooks/01_run_baseline.ipynb @@ -0,0 +1,68 @@ +{ + "cells": [ + { + "cell_type": "markdown", + "metadata": {}, + "source": [ + "# Run the nanoGPT Muon-HyperBall baseline\n\n", + "This notebook launches one selected arm with the committed reference configuration. ", + "The default is `muon_hyperball`, seed 1337.\n" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [], + "source": [ + "from pathlib import Path\n", + "from rg_nanogpt_muon_hyperball.config import load_config\n", + "from rg_nanogpt_muon_hyperball.training import run_optimizer_replicates\n", + "\n", + "CONFIG = Path('../configs/reference.yaml').resolve()\n", + "DATA_ROOT = Path('/tmp/rg-nanogpt-one-head/data')\n", + "RESULTS_ROOT = Path('/tmp/rg-nanogpt-muon-hyperball/results')\n", + "OPTIMIZER = 'muon_hyperball' # or 'muon'\n", + "SEEDS = (1337,)\n", + "DEVICE = 'auto'\n", + "\n", + "cfg = load_config(CONFIG)\n", + "cfg['training']['target_epochs'], cfg['optimizer_profiles'][OPTIMIZER]\n" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [], + "source": [ + "run_optimizer_replicates(\n", + " cfg=cfg,\n", + " config_path=CONFIG,\n", + " optimizer_name=OPTIMIZER,\n", + " seeds=SEEDS,\n", + " data_root=DATA_ROOT,\n", + " results_root=RESULTS_ROOT,\n", + " device=DEVICE,\n", + " resume=True,\n", + " overwrite=False,\n", + " prepare_data=False,\n", + " progress=True,\n", + ")\n" + ] + } + ], + "metadata": { + "kernelspec": { + "display_name": "Python 3", + "language": "python", + "name": "python3" + }, + "language_info": { + "name": "python", + "version": "3.10" + } + }, + "nbformat": 4, + "nbformat_minor": 5 +} \ No newline at end of file diff --git a/baseline/nanogpt_muon_hyperball/notebooks/02_compare_muon_hyperball.ipynb b/baseline/nanogpt_muon_hyperball/notebooks/02_compare_muon_hyperball.ipynb new file mode 100644 index 0000000..bceea71 --- /dev/null +++ b/baseline/nanogpt_muon_hyperball/notebooks/02_compare_muon_hyperball.ipynb @@ -0,0 +1,132 @@ +{ + "cells": [ + { + "cell_type": "markdown", + "metadata": {}, + "source": [ + "# Compare long Muon with Muon-HyperBall\n\n", + "Plots task metrics, per-layer alpha and fit distance, ERG gap, traps, and ", + "HyperBall activation diagnostics. Incomplete runs can be inspected by setting ", + "`REQUIRE_COMPLETE = False`.\n" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [], + "source": [ + "from pathlib import Path\n", + "import matplotlib.pyplot as plt\n", + "import pandas as pd\n", + "\n", + "from rg_nanogpt_muon_hyperball.analysis import (\n", + " load_epoch_metrics,\n", + " load_layer_metrics,\n", + " plot_epoch_metric,\n", + " plot_layer_metric,\n", + " run_status_table,\n", + ")\n", + "\n", + "RESULTS_ROOT = Path('/tmp/rg-nanogpt-muon-hyperball/results')\n", + "OPTIMIZERS = ('muon', 'muon_hyperball')\n", + "SEEDS = (1337,)\n", + "REQUIRE_COMPLETE = False\n", + "\n", + "run_status_table(RESULTS_ROOT, optimizers=OPTIMIZERS, seeds=SEEDS)\n" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [], + "source": [ + "epochs = load_epoch_metrics(\n", + " RESULTS_ROOT,\n", + " optimizers=OPTIMIZERS,\n", + " seeds=SEEDS,\n", + " require_complete=REQUIRE_COMPLETE,\n", + ")\n", + "layers = load_layer_metrics(\n", + " RESULTS_ROOT,\n", + " optimizers=OPTIMIZERS,\n", + " seeds=SEEDS,\n", + " require_complete=REQUIRE_COMPLETE,\n", + ")\n", + "epochs.tail(), layers.tail()\n" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [], + "source": [ + "for metric in ('val_loss', 'val_accuracy', 'test_accuracy', 'primary_lr'):\n", + " if not epochs.empty and metric in epochs.columns:\n", + " plot_epoch_metric(\n", + " epochs,\n", + " metric=metric,\n", + " optimizers=OPTIMIZERS,\n", + " title=f'{metric}: long Muon versus Muon-HyperBall',\n", + " )\n", + " plt.show()\n" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [], + "source": [ + "for optimizer in OPTIMIZERS:\n", + " for metric in ('alpha', 'D', 'ERG_gap', 'num_traps'):\n", + " if not layers.empty and metric in layers.columns:\n", + " plot_layer_metric(\n", + " layers,\n", + " optimizer=optimizer,\n", + " metric=metric,\n", + " )\n", + " plt.show()\n" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [], + "source": [ + "hyperball_metrics = (\n", + " 'hyperball_active_fraction',\n", + " 'hyperball_mean_scale',\n", + " 'hyperball_min_scale',\n", + " 'hyperball_max_proposed_update_to_weight_ratio',\n", + " 'hyperball_max_applied_update_to_weight_ratio',\n", + ")\n", + "for metric in hyperball_metrics:\n", + " if not epochs.empty and metric in epochs.columns:\n", + " plot_epoch_metric(\n", + " epochs,\n", + " metric=metric,\n", + " optimizers=('muon_hyperball',),\n", + " title=metric,\n", + " )\n", + " plt.show()\n" + ] + } + ], + "metadata": { + "kernelspec": { + "display_name": "Python 3", + "language": "python", + "name": "python3" + }, + "language_info": { + "name": "python", + "version": "3.10" + } + }, + "nbformat": 4, + "nbformat_minor": 5 +} \ No newline at end of file diff --git a/baseline/nanogpt_muon_hyperball/pyproject.toml b/baseline/nanogpt_muon_hyperball/pyproject.toml new file mode 100644 index 0000000..9563431 --- /dev/null +++ b/baseline/nanogpt_muon_hyperball/pyproject.toml @@ -0,0 +1,32 @@ +[build-system] +requires = ["setuptools>=68", "wheel"] +build-backend = "setuptools.build_meta" + +[project] +name = "rg-nanogpt-muon-hyperball" +version = "0.1.0" +description = "Matched long-Muon and Muon-HyperBall one-head nanoGPT baselines" +requires-python = ">=3.10" +dependencies = [ + "torch>=2.2", + "numpy>=1.24", + "pandas>=2.0", + "matplotlib>=3.7", + "pyyaml>=6.0", + "datasets>=2.19", + "tiktoken>=0.7", + "sacrebleu>=2.4", + "weightwatcher==0.7.7", + "jupyter>=1.0", + "ipykernel>=6.29", +] + +[project.optional-dependencies] +dev = ["pytest>=8.0", "nbformat>=5.10"] + +[project.scripts] +rg-hyperball-prepare = "rg_nanogpt_muon_hyperball.data:main" +rg-hyperball-train = "rg_nanogpt_muon_hyperball.training:main" + +[tool.setuptools.packages.find] +where = ["src"] diff --git a/baseline/nanogpt_muon_hyperball/requirements.txt b/baseline/nanogpt_muon_hyperball/requirements.txt new file mode 100644 index 0000000..fd50a21 --- /dev/null +++ b/baseline/nanogpt_muon_hyperball/requirements.txt @@ -0,0 +1,11 @@ +torch>=2.2 +numpy>=1.24 +pandas>=2.0 +matplotlib>=3.7 +pyyaml>=6.0 +datasets>=2.19 +tiktoken>=0.7 +sacrebleu>=2.4 +weightwatcher==0.7.7 +jupyter>=1.0 +ipykernel>=6.29 diff --git a/baseline/nanogpt_muon_hyperball/src/rg_nanogpt_muon_hyperball/__init__.py b/baseline/nanogpt_muon_hyperball/src/rg_nanogpt_muon_hyperball/__init__.py new file mode 100644 index 0000000..2b5d9e3 --- /dev/null +++ b/baseline/nanogpt_muon_hyperball/src/rg_nanogpt_muon_hyperball/__init__.py @@ -0,0 +1,13 @@ +"""One-head nanoGPT Muon and Muon-HyperBall comparison baseline.""" + +from .config import SUPPORTED_OPTIMIZERS, canonical_seeds, load_config, roots +from .optimizers import Muon, MuonHyperBall + +__all__ = [ + "Muon", + "MuonHyperBall", + "SUPPORTED_OPTIMIZERS", + "canonical_seeds", + "load_config", + "roots", +] diff --git a/baseline/nanogpt_muon_hyperball/src/rg_nanogpt_muon_hyperball/analysis.py b/baseline/nanogpt_muon_hyperball/src/rg_nanogpt_muon_hyperball/analysis.py new file mode 100644 index 0000000..b19b018 --- /dev/null +++ b/baseline/nanogpt_muon_hyperball/src/rg_nanogpt_muon_hyperball/analysis.py @@ -0,0 +1,405 @@ +from __future__ import annotations + +import json +import math +from pathlib import Path +from typing import Iterable, Sequence + +import matplotlib.pyplot as plt +import numpy as np +import pandas as pd + +from .config import SUPPORTED_OPTIMIZERS +from .run_utils import run_directory, run_is_complete + +OPTIMIZER_LABELS = { + "muon": "Long Muon + auxiliary AdamW", + "muon_hyperball": "Muon + HyperBall + auxiliary AdamW", +} +OPTIMIZER_COLORS = { + "muon": "#009E73", + "muon_hyperball": "#0072B2", +} +MATRIX_COLORS = { + "W_Q": "#0072B2", + "W_K": "#E69F00", + "W_V": "#009E73", + "W_O": "#D55E00", + "W_MLP_IN": "#CC79A7", + "W_MLP_OUT": "#56B4E9", +} + +_T_975 = { + 1: 12.7062047364, + 2: 4.3026527297, + 3: 3.1824463053, + 4: 2.7764451052, + 5: 2.5705818356, + 6: 2.4469118511, + 7: 2.3646242510, + 8: 2.3060041352, + 9: 2.2621571629, + 10: 2.2281388520, +} + + +def mean_ci95(values: Iterable[float]) -> dict[str, float]: + array = np.asarray(list(values), dtype=float) + array = array[np.isfinite(array)] + n = int(array.size) + if n == 0: + return { + "n": 0, + "mean": np.nan, + "sd": np.nan, + "sem": np.nan, + "ci95_half_width": np.nan, + "ci95_lower": np.nan, + "ci95_upper": np.nan, + } + mean = float(array.mean()) + if n == 1: + return { + "n": 1, + "mean": mean, + "sd": np.nan, + "sem": np.nan, + "ci95_half_width": np.nan, + "ci95_lower": np.nan, + "ci95_upper": np.nan, + } + sd = float(array.std(ddof=1)) + sem = sd / math.sqrt(n) + half = _T_975.get(n - 1, 1.9599639845) * sem + return { + "n": n, + "mean": mean, + "sd": sd, + "sem": sem, + "ci95_half_width": half, + "ci95_lower": mean - half, + "ci95_upper": mean + half, + } + + +def _require_complete(results_root: str | Path, optimizer: str, seed: int) -> None: + if not run_is_complete(results_root, optimizer, seed): + raise FileNotFoundError( + f"missing completed run for optimizer={optimizer} seed={seed}: " + f"{run_directory(results_root, optimizer, seed)}" + ) + + +def _load_csvs( + results_root: str | Path, + relative_path: str, + *, + optimizers: Sequence[str], + seeds: Sequence[int], + require_complete: bool, +) -> pd.DataFrame: + frames = [] + for optimizer in optimizers: + for seed in seeds: + if require_complete: + _require_complete(results_root, optimizer, seed) + path = run_directory(results_root, optimizer, seed) / relative_path + if not path.is_file(): + if require_complete: + raise FileNotFoundError(path) + continue + frame = pd.read_csv(path) + frame["optimizer"] = optimizer + frame["optimizer_label"] = OPTIMIZER_LABELS.get(optimizer, optimizer) + frame["seed"] = int(seed) + frames.append(frame) + return pd.concat(frames, ignore_index=True, sort=False) if frames else pd.DataFrame() + + +def run_status_table( + results_root: str | Path, + *, + optimizers: Sequence[str] = SUPPORTED_OPTIMIZERS, + seeds: Sequence[int] = (1337,), +) -> pd.DataFrame: + rows = [] + for optimizer in optimizers: + for seed in seeds: + run_dir = run_directory(results_root, optimizer, seed) + path = run_dir / "run_complete.json" + payload = json.loads(path.read_text()) if path.is_file() else {} + rows.append( + { + "optimizer": optimizer, + "optimizer_label": OPTIMIZER_LABELS.get(optimizer, optimizer), + "seed": int(seed), + "complete": bool(payload.get("completed", False)), + "steps": payload.get("optimizer_steps", np.nan), + "final_test_loss": payload.get("final_test_loss", np.nan), + "final_test_accuracy": payload.get("final_test_accuracy", np.nan), + "run_dir": str(run_dir), + } + ) + return pd.DataFrame(rows) + + +def load_metrics( + results_root: str | Path, + *, + optimizers: Sequence[str] = SUPPORTED_OPTIMIZERS, + seeds: Sequence[int] = (1337,), + require_complete: bool = True, +) -> pd.DataFrame: + frame = _load_csvs( + results_root, + "metrics.csv", + optimizers=optimizers, + seeds=seeds, + require_complete=require_complete, + ) + if frame.empty: + return frame + return frame.sort_values(["optimizer", "seed", "step"]).drop_duplicates( + ["optimizer", "seed", "step"], keep="last" + ) + + +def load_epoch_metrics( + results_root: str | Path, + *, + optimizers: Sequence[str] = SUPPORTED_OPTIMIZERS, + seeds: Sequence[int] = (1337,), + require_complete: bool = True, +) -> pd.DataFrame: + frame = _load_csvs( + results_root, + "epoch_metrics.csv", + optimizers=optimizers, + seeds=seeds, + require_complete=require_complete, + ) + if frame.empty: + return frame + return frame.sort_values( + ["optimizer", "seed", "nominal_epoch"] + ).drop_duplicates(["optimizer", "seed", "nominal_epoch"], keep="last") + + +def load_layer_metrics( + results_root: str | Path, + *, + optimizers: Sequence[str] = SUPPORTED_OPTIMIZERS, + seeds: Sequence[int] = (1337,), + require_complete: bool = True, +) -> pd.DataFrame: + frame = _load_csvs( + results_root, + "spectral/layers.csv", + optimizers=optimizers, + seeds=seeds, + require_complete=require_complete, + ) + if frame.empty: + return frame + return frame.sort_values( + ["optimizer", "seed", "epoch", "matrix_type"] + ).drop_duplicates( + ["optimizer", "seed", "step", "matrix_name"], keep="last" + ) + + +def load_spectral_summary( + results_root: str | Path, + *, + optimizers: Sequence[str] = SUPPORTED_OPTIMIZERS, + seeds: Sequence[int] = (1337,), + require_complete: bool = True, +) -> pd.DataFrame: + frame = _load_csvs( + results_root, + "spectral/summary.csv", + optimizers=optimizers, + seeds=seeds, + require_complete=require_complete, + ) + if frame.empty: + return frame + return frame.sort_values(["optimizer", "seed", "epoch"]).drop_duplicates( + ["optimizer", "seed", "step"], keep="last" + ) + + +def load_test_results( + results_root: str | Path, + *, + optimizers: Sequence[str] = SUPPORTED_OPTIMIZERS, + seeds: Sequence[int] = (1337,), +) -> pd.DataFrame: + rows = [] + for optimizer in optimizers: + for seed in seeds: + _require_complete(results_root, optimizer, seed) + path = run_directory(results_root, optimizer, seed) / "test_results.json" + payload = json.loads(path.read_text(encoding="utf-8")) + for checkpoint in ("final", "validation_selected"): + values = payload[checkpoint] + rows.append( + { + "optimizer": optimizer, + "optimizer_label": OPTIMIZER_LABELS[optimizer], + "seed": int(seed), + "checkpoint": checkpoint, + "step": int(values["step"]), + "test_loss": float(values["loss"]), + "test_perplexity": float(values["perplexity"]), + "test_accuracy": float(values["accuracy"]), + "test_bleu": float(values["bleu"]), + } + ) + return pd.DataFrame(rows) + + +def summarize_by_epoch( + frame: pd.DataFrame, + metric: str, + *, + x: str = "nominal_epoch", + group: Sequence[str] = ("optimizer",), +) -> pd.DataFrame: + rows = [] + keys = [*group, x] + subset = frame[[*keys, "seed", metric]].copy() + subset[metric] = pd.to_numeric(subset[metric], errors="coerce") + for values, group_frame in subset.groupby(keys, sort=True): + values_tuple = values if isinstance(values, tuple) else (values,) + row = dict(zip(keys, values_tuple, strict=True)) + row.update(mean_ci95(group_frame[metric])) + rows.append(row) + return pd.DataFrame(rows) + + +def plot_epoch_metric( + frame: pd.DataFrame, + *, + metric: str, + x: str = "nominal_epoch", + optimizers: Sequence[str] = SUPPORTED_OPTIMIZERS, + title: str | None = None, + output: str | Path | None = None, +): + figure, axis = plt.subplots(figsize=(9, 5)) + for optimizer in optimizers: + subset = frame[frame["optimizer"] == optimizer] + if subset.empty: + continue + for _, seed_frame in subset.groupby("seed"): + axis.plot( + seed_frame[x], + seed_frame[metric], + color=OPTIMIZER_COLORS[optimizer], + alpha=0.30, + linewidth=1.0, + ) + summary = summarize_by_epoch(subset, metric, x=x) + axis.plot( + summary[x], + summary["mean"], + color=OPTIMIZER_COLORS[optimizer], + linewidth=2.0, + label=OPTIMIZER_LABELS[optimizer], + ) + if summary["ci95_lower"].notna().any(): + axis.fill_between( + summary[x], + summary["ci95_lower"], + summary["ci95_upper"], + color=OPTIMIZER_COLORS[optimizer], + alpha=0.16, + ) + axis.set_xlabel(x.replace("_", " ").title()) + axis.set_ylabel(metric.replace("_", " ").title()) + axis.set_title(title or metric.replace("_", " ").title()) + axis.grid(alpha=0.25) + axis.legend(frameon=False) + figure.tight_layout() + if output is not None: + output = Path(output) + output.parent.mkdir(parents=True, exist_ok=True) + figure.savefig(output, dpi=170, bbox_inches="tight") + return figure, axis + + +def plot_layer_metric( + frame: pd.DataFrame, + *, + optimizer: str, + metric: str, + title: str | None = None, + output: str | Path | None = None, +): + subset = frame[frame["optimizer"] == optimizer].copy() + if subset.empty: + raise ValueError(f"no layer data for optimizer={optimizer}") + figure, axis = plt.subplots(figsize=(10, 5.5)) + for matrix_type, color in MATRIX_COLORS.items(): + matrix = subset[subset["matrix_type"] == matrix_type] + if matrix.empty: + continue + summary = summarize_by_epoch( + matrix, metric, x="epoch", group=("matrix_type",) + ) + axis.plot( + summary["epoch"], + summary["mean"], + color=color, + linewidth=2.0, + marker="o", + markersize=3, + label=matrix_type, + ) + if summary["ci95_lower"].notna().any(): + axis.fill_between( + summary["epoch"], + summary["ci95_lower"], + summary["ci95_upper"], + color=color, + alpha=0.13, + ) + if metric == "alpha": + axis.axhline(2.0, linestyle="--", linewidth=1.0, label="alpha = 2") + if metric == "ERG_gap": + axis.axhline(0.0, linestyle="--", linewidth=1.0) + axis.set_xlabel("Epoch") + axis.set_ylabel(metric) + axis.set_title(title or f"{OPTIMIZER_LABELS[optimizer]} layer {metric}") + axis.grid(alpha=0.25) + axis.legend(frameon=False, ncol=2) + figure.tight_layout() + if output is not None: + output = Path(output) + output.parent.mkdir(parents=True, exist_ok=True) + figure.savefig(output, dpi=170, bbox_inches="tight") + return figure, axis + + +def final_test_summary(test_results: pd.DataFrame) -> pd.DataFrame: + rows = [] + for (optimizer, checkpoint), group in test_results.groupby( + ["optimizer", "checkpoint"] + ): + for metric in ( + "test_loss", + "test_perplexity", + "test_accuracy", + "test_bleu", + ): + rows.append( + { + "optimizer": optimizer, + "optimizer_label": OPTIMIZER_LABELS[optimizer], + "checkpoint": checkpoint, + "metric": metric, + **mean_ci95(group[metric]), + } + ) + return pd.DataFrame(rows) diff --git a/baseline/nanogpt_muon_hyperball/src/rg_nanogpt_muon_hyperball/base_optimizers.py b/baseline/nanogpt_muon_hyperball/src/rg_nanogpt_muon_hyperball/base_optimizers.py new file mode 100644 index 0000000..5b0b4f3 --- /dev/null +++ b/baseline/nanogpt_muon_hyperball/src/rg_nanogpt_muon_hyperball/base_optimizers.py @@ -0,0 +1,289 @@ +from __future__ import annotations + +from dataclasses import dataclass +import math +from typing import Iterable + +import torch + + +@dataclass +class OptimizerHandle: + role: str + optimizer: torch.optim.Optimizer + peak_lr: float + min_lr: float + + def set_lr(self, value: float) -> None: + for group in self.optimizer.param_groups: + group["lr"] = float(value) + + @property + def lr(self) -> float: + return float(self.optimizer.param_groups[0]["lr"]) + + +@torch.no_grad() +def zeropower_via_newton_schulz_5( + update: torch.Tensor, + *, + steps: int = 5, + eps: float = 1e-7, +) -> torch.Tensor: + """Approximate the polar factor of a matrix update using Muon's quintic map. + + The calculation is deliberately performed in float32 for Apple-MPS + reliability, then converted back to the parameter dtype. + """ + if update.ndim != 2: + raise ValueError(f"Muon requires a matrix update, got shape={tuple(update.shape)}") + if steps < 1 or eps <= 0: + raise ValueError("steps and eps must be positive") + original_dtype = update.dtype + transposed = update.shape[0] > update.shape[1] + x = update.T if transposed else update + x = x.float() + x = x / torch.linalg.vector_norm(x).clamp_min(float(eps)) + a, b, c = 3.4445, -4.7750, 2.0315 + for _ in range(int(steps)): + gram = x @ x.T + x = a * x + (b * gram + c * (gram @ gram)) @ x + if transposed: + x = x.T + return x.to(original_dtype) + + +class Muon(torch.optim.Optimizer): + """Single-device Muon for hidden 2-D transformer matrices.""" + + def __init__( + self, + params: Iterable[torch.nn.Parameter], + *, + lr: float, + momentum: float = 0.95, + nesterov: bool = True, + weight_decay: float = 0.01, + newton_schulz_steps: int = 5, + eps: float = 1e-7, + ) -> None: + params = list(params) + if not params: + raise ValueError("Muon requires at least one parameter") + if any(parameter.ndim != 2 for parameter in params): + raise ValueError("Muon accepts only 2-D parameters") + defaults = { + "lr": float(lr), + "momentum": float(momentum), + "nesterov": bool(nesterov), + "weight_decay": float(weight_decay), + "newton_schulz_steps": int(newton_schulz_steps), + "eps": float(eps), + } + super().__init__(params, defaults) + + @torch.no_grad() + def step(self, closure=None): + loss = None + if closure is not None: + with torch.enable_grad(): + loss = closure() + for group in self.param_groups: + lr = float(group["lr"]) + momentum_value = float(group["momentum"]) + for parameter in group["params"]: + if parameter.grad is None: + continue + gradient = parameter.grad.detach() + if gradient.is_sparse: + raise RuntimeError("Muon does not support sparse gradients") + state = self.state[parameter] + buffer = state.get("momentum_buffer") + if buffer is None: + buffer = torch.zeros_like(gradient) + state["momentum_buffer"] = buffer + buffer.lerp_(gradient, 1.0 - momentum_value) + update_source = ( + gradient.lerp(buffer, momentum_value) + if bool(group["nesterov"]) + else buffer + ) + update = zeropower_via_newton_schulz_5( + update_source, + steps=int(group["newton_schulz_steps"]), + eps=float(group["eps"]), + ) + update.mul_(math.sqrt(max(1.0, parameter.shape[0] / parameter.shape[1]))) + decay = float(group["weight_decay"]) + if decay: + parameter.mul_(max(0.0, 1.0 - lr * decay)) + parameter.add_(update, alpha=-lr) + return loss + + +def cosine_learning_rate( + update_index: int, + *, + total_steps: int, + warmup_steps: int, + peak_lr: float, + min_lr: float, +) -> float: + """Linear warm-up followed by cosine decay to a nonzero floor.""" + if total_steps < 1: + raise ValueError("total_steps must be positive") + if not 0 <= warmup_steps < total_steps: + raise ValueError("warmup_steps must be in [0, total_steps)") + if update_index < 0: + raise ValueError("update_index must be nonnegative") + if warmup_steps and update_index < warmup_steps: + return float(peak_lr) * (update_index + 1) / warmup_steps + progress = (update_index - warmup_steps) / max(1, total_steps - warmup_steps - 1) + progress = min(1.0, max(0.0, progress)) + cosine = 0.5 * (1.0 + math.cos(math.pi * progress)) + return float(min_lr) + cosine * (float(peak_lr) - float(min_lr)) + + +def _named_parameters(model) -> list[tuple[str, torch.nn.Parameter]]: + return [ + (name, parameter) + for name, parameter in model.named_parameters() + if parameter.requires_grad + ] + + +def _decay_groups( + named_parameters: list[tuple[str, torch.nn.Parameter]], + weight_decay: float, +) -> list[dict]: + decay = [parameter for _, parameter in named_parameters if parameter.ndim >= 2] + no_decay = [parameter for _, parameter in named_parameters if parameter.ndim < 2] + return [ + {"params": decay, "weight_decay": float(weight_decay)}, + {"params": no_decay, "weight_decay": 0.0}, + ] + + +def make_optimizer_handles(model, profile: dict) -> list[OptimizerHandle]: + named = _named_parameters(model) + family = str(profile["family"]) + + if family == "sgd": + optimizer = torch.optim.SGD( + _decay_groups(named, float(profile["weight_decay"])), + lr=float(profile["learning_rate"]), + momentum=float(profile["momentum"]), + dampening=float(profile.get("dampening", 0.0)), + nesterov=bool(profile.get("nesterov", True)), + ) + return [ + OptimizerHandle( + role="primary", + optimizer=optimizer, + peak_lr=float(profile["learning_rate"]), + min_lr=float(profile["min_learning_rate"]), + ) + ] + + if family == "adamw": + optimizer = torch.optim.AdamW( + _decay_groups(named, float(profile["weight_decay"])), + lr=float(profile["learning_rate"]), + betas=(float(profile["beta1"]), float(profile["beta2"])), + eps=float(profile["epsilon"]), + ) + return [ + OptimizerHandle( + role="primary", + optimizer=optimizer, + peak_lr=float(profile["learning_rate"]), + min_lr=float(profile["min_learning_rate"]), + ) + ] + + if family != "muon": + raise ValueError(f"unsupported optimizer family: {family}") + + hidden = [ + parameter + for name, parameter in named + if name.startswith("blocks.") and parameter.ndim == 2 + ] + hidden_ids = {id(parameter) for parameter in hidden} + auxiliary_named = [ + (name, parameter) for name, parameter in named if id(parameter) not in hidden_ids + ] + if not hidden or not auxiliary_named: + raise ValueError("Muon partition must contain both hidden matrices and auxiliary parameters") + + muon = Muon( + hidden, + lr=float(profile["matrix_learning_rate"]), + momentum=float(profile["momentum"]), + nesterov=bool(profile["nesterov"]), + weight_decay=float(profile["matrix_weight_decay"]), + newton_schulz_steps=int(profile["newton_schulz_steps"]), + eps=float(profile.get("muon_epsilon", 1e-7)), + ) + auxiliary = torch.optim.AdamW( + _decay_groups(auxiliary_named, float(profile["aux_weight_decay"])), + lr=float(profile["aux_learning_rate"]), + betas=(float(profile["beta1"]), float(profile["beta2"])), + eps=float(profile["epsilon"]), + ) + return [ + OptimizerHandle( + role="primary", + optimizer=muon, + peak_lr=float(profile["matrix_learning_rate"]), + min_lr=float(profile["matrix_min_learning_rate"]), + ), + OptimizerHandle( + role="auxiliary", + optimizer=auxiliary, + peak_lr=float(profile["aux_learning_rate"]), + min_lr=float(profile["aux_min_learning_rate"]), + ), + ] + + +def set_learning_rates( + handles: list[OptimizerHandle], + *, + update_index: int, + total_steps: int, + warmup_steps: int, +) -> dict[str, float]: + values: dict[str, float] = {} + for handle in handles: + value = cosine_learning_rate( + update_index, + total_steps=total_steps, + warmup_steps=warmup_steps, + peak_lr=handle.peak_lr, + min_lr=handle.min_lr, + ) + handle.set_lr(value) + values[handle.role] = value + return values + + +def zero_grad(handles: list[OptimizerHandle]) -> None: + for handle in handles: + handle.optimizer.zero_grad(set_to_none=True) + + +def optimizer_step(handles: list[OptimizerHandle]) -> None: + for handle in handles: + handle.optimizer.step() + + +def optimizer_state_dict(handles: list[OptimizerHandle]) -> list[dict]: + return [handle.optimizer.state_dict() for handle in handles] + + +def load_optimizer_state_dict(handles: list[OptimizerHandle], states: list[dict]) -> None: + if len(handles) != len(states): + raise RuntimeError("optimizer-handle count changed across resume") + for handle, state in zip(handles, states, strict=True): + handle.optimizer.load_state_dict(state) diff --git a/baseline/nanogpt_muon_hyperball/src/rg_nanogpt_muon_hyperball/checkpoints.py b/baseline/nanogpt_muon_hyperball/src/rg_nanogpt_muon_hyperball/checkpoints.py new file mode 100644 index 0000000..29b6cb1 --- /dev/null +++ b/baseline/nanogpt_muon_hyperball/src/rg_nanogpt_muon_hyperball/checkpoints.py @@ -0,0 +1,145 @@ +from __future__ import annotations + +from pathlib import Path +import random +from typing import Any + +import numpy as np +import torch + +from .optimizers import ( + OptimizerHandle, + load_optimizer_state_dict, + optimizer_state_dict, +) + + +def _atomic_torch_save(payload: dict[str, Any], path: Path) -> Path: + path.parent.mkdir(parents=True, exist_ok=True) + temporary = path.with_suffix(path.suffix + ".tmp") + torch.save(payload, temporary) + temporary.replace(path) + return path + + +def _capture_accelerator_rng_state() -> dict[str, Any]: + state: dict[str, Any] = {} + if torch.cuda.is_available(): + state["cuda_random_state_all"] = torch.cuda.get_rng_state_all() + if ( + hasattr(torch, "mps") + and hasattr(torch.mps, "get_rng_state") + and torch.backends.mps.is_available() + ): + state["mps_random_state"] = torch.mps.get_rng_state() + return state + + +def _restore_accelerator_rng_state(payload: dict[str, Any]) -> None: + if torch.cuda.is_available() and "cuda_random_state_all" in payload: + torch.cuda.set_rng_state_all(payload["cuda_random_state_all"]) + if ( + "mps_random_state" in payload + and hasattr(torch, "mps") + and hasattr(torch.mps, "set_rng_state") + and torch.backends.mps.is_available() + ): + torch.mps.set_rng_state(payload["mps_random_state"]) + + +def save_training_checkpoint( + path: str | Path, + *, + model, + handles: list[OptimizerHandle], + step: int, + best_validation_loss: float, + best_validation_step: int, + elapsed_seconds: float, + fingerprint: str, + cfg: dict, + optimizer_name: str, + seed: int, + train_generator: torch.Generator, +) -> Path: + payload: dict[str, Any] = { + "schema_version": 2, + "model": model.state_dict(), + "optimizers": optimizer_state_dict(handles), + "step": int(step), + "best_validation_loss": float(best_validation_loss), + "best_validation_step": int(best_validation_step), + "elapsed_seconds": float(elapsed_seconds), + "fingerprint": str(fingerprint), + "config": cfg, + "optimizer_name": str(optimizer_name), + "seed": int(seed), + "python_random_state": random.getstate(), + "numpy_random_state": np.random.get_state(), + "torch_random_state": torch.random.get_rng_state(), + "train_generator_state": train_generator.get_state(), + **_capture_accelerator_rng_state(), + } + return _atomic_torch_save(payload, Path(path)) + + +def load_training_checkpoint( + path: str | Path, + *, + model, + handles: list[OptimizerHandle], + expected_fingerprint: str, + train_generator: torch.Generator, +) -> tuple[int, float, int, float]: + path = Path(path) + payload = torch.load(path, map_location="cpu", weights_only=False) + if str(payload.get("fingerprint")) != str(expected_fingerprint): + raise RuntimeError( + "checkpoint protocol fingerprint does not match the requested run" + ) + model.load_state_dict(payload["model"]) + load_optimizer_state_dict(handles, payload["optimizers"]) + random.setstate(payload["python_random_state"]) + np.random.set_state(payload["numpy_random_state"]) + torch.random.set_rng_state(payload["torch_random_state"]) + train_generator.set_state(payload["train_generator_state"]) + _restore_accelerator_rng_state(payload) + return ( + int(payload["step"]), + float(payload["best_validation_loss"]), + int(payload["best_validation_step"]), + float(payload["elapsed_seconds"]), + ) + + +def save_epoch_model_checkpoint( + run_dir: str | Path, + *, + model, + step: int, + nominal_epoch: float, + actual_epoch: float, + fingerprint: str, + cfg: dict, + optimizer_name: str, + seed: int, +) -> Path: + epoch_text = f"{float(nominal_epoch):06.3f}".replace(".", "p") + path = ( + Path(run_dir) + / "epoch_checkpoints" + / f"model_epoch_{epoch_text}_step_{int(step):07d}.pt" + ) + payload = { + "schema_version": 1, + "model": model.state_dict(), + "step": int(step), + "nominal_epoch": float(nominal_epoch), + "actual_epoch": float(actual_epoch), + "fingerprint": str(fingerprint), + "config": cfg, + "optimizer_name": str(optimizer_name), + "seed": int(seed), + "purpose": "per_epoch_model_only_analysis_checkpoint", + } + return _atomic_torch_save(payload, path) diff --git a/baseline/nanogpt_muon_hyperball/src/rg_nanogpt_muon_hyperball/completion.py b/baseline/nanogpt_muon_hyperball/src/rg_nanogpt_muon_hyperball/completion.py new file mode 100644 index 0000000..fa62bc0 --- /dev/null +++ b/baseline/nanogpt_muon_hyperball/src/rg_nanogpt_muon_hyperball/completion.py @@ -0,0 +1,323 @@ +"""Validation for completed one-head nanoGPT experiment directories.""" + +from __future__ import annotations + +import json +import math +from pathlib import Path +from typing import Any, NoReturn + +import numpy as np +import pandas as pd +import torch + +_REQUIRED_FILES = ( + "run_complete.json", + "manifest.json", + "metrics.csv", + "epoch_metrics.csv", + "checkpoint_latest.pt", + "checkpoint_best.pt", + "checkpoint_final.pt", + "test_results.json", + "spectral/layers.csv", + "spectral/summary.csv", +) + + +class CompletedRunValidationError(RuntimeError): + """A nominally completed run is missing, stale, or inconsistent.""" + + +def _fail(message: str) -> NoReturn: + raise CompletedRunValidationError( + "completed one-head nanoGPT run is stale or inconsistent: " + + message + + ". Use a new results directory or rerun with explicit overwrite." + ) + + +def _read_json(path: Path) -> dict[str, Any]: + try: + payload = json.loads(path.read_text(encoding="utf-8")) + except (OSError, json.JSONDecodeError) as exc: + _fail(f"could not read valid JSON from {path}: {exc}") + if not isinstance(payload, dict): + _fail(f"{path} does not contain a JSON object") + return payload + + +def _read_csv(path: Path) -> pd.DataFrame: + try: + frame = pd.read_csv(path) + except Exception as exc: + _fail(f"could not read {path}: {exc}") + if frame.empty: + _fail(f" {path} is empty") + return frame + + +def _as_int(value: Any, label: str) -> int: + try: + return int(value) + except (TypeError, ValueError): + _fail(f"{label} is not an integer: {value!r}") + + +def _expect(observed: Any, expected: Any, label: str) -> None: + if observed != expected: + _fail( + f"{label} mismatch; observed {observed!r}, expected {expected!r}" + ) + + +def _step_tuple(frame: pd.DataFrame, label: str) -> tuple[int, ...]: + if "step" not in frame.columns: + _fail(f"{label} has no step column") + values = pd.to_numeric(frame["step"], errors="coerce").to_numpy(dtype=float) + if not np.isfinite(values).all(): + _fail(f"{label} contains non-finite step values") + rounded = np.rint(values) + if not np.allclose(values, rounded): + _fail(f"{label} contains non-integer step values") + steps = tuple(int(value) for value in rounded) + if len(steps) != len(set(steps)): + _fail(f"{label} contains duplicate step rows") + return steps + + +def _load_checkpoint(path: Path) -> dict[str, Any]: + try: + payload = torch.load(path, map_location="cpu", weights_only=False) + except Exception as exc: + _fail(f"could not load checkpoint {path}: {exc}") + if not isinstance(payload, dict): + _fail(f"checkpoint {path} is not a mapping") + return payload + + +def validate_completed_run( + run_dir: str | Path, + *, + expected_fingerprint: str | None = None, + expected_optimizer: str | None = None, + expected_seed: int | None = None, + expected_total_steps: int | None = None, + verify_checkpoints: bool = True, +) -> dict[str, Any]: + """Validate a completed run before it is skipped or analyzed. + + Expected values supplied by the current configuration make this a stale-run + guard. Without them, the function still verifies internal consistency. + """ + + root = Path(run_dir) + missing = [ + str(root / relative) + for relative in _REQUIRED_FILES + if not (root / relative).is_file() + or (root / relative).stat().st_size == 0 + ] + if missing: + _fail("missing required artifacts: " + ", ".join(missing)) + + completion = _read_json(root / "run_complete.json") + manifest = _read_json(root / "manifest.json") + test_results = _read_json(root / "test_results.json") + if completion.get("completed") is not True: + _fail("run_complete.json does not declare completed=true") + + recorded_fingerprint = str(completion.get("fingerprint", "")) + fingerprint = ( + str(expected_fingerprint) + if expected_fingerprint is not None + else recorded_fingerprint + ) + if not fingerprint or not recorded_fingerprint: + _fail("the completion record has no protocol fingerprint") + + recorded_optimizer = str(completion.get("optimizer", "")) + optimizer = ( + str(expected_optimizer) + if expected_optimizer is not None + else recorded_optimizer + ) + if not optimizer or not recorded_optimizer: + _fail("the completion record has no optimizer") + + recorded_seed = _as_int(completion.get("seed"), "completion seed") + seed = int(expected_seed) if expected_seed is not None else recorded_seed + recorded_steps = _as_int( + completion.get("optimizer_steps"), "completion optimizer_steps" + ) + total_steps = ( + int(expected_total_steps) + if expected_total_steps is not None + else recorded_steps + ) + best_step = _as_int( + completion.get("best_validation_step"), + "completion best_validation_step", + ) + + _expect(recorded_fingerprint, fingerprint, "completion fingerprint") + _expect(recorded_optimizer, optimizer, "completion optimizer") + _expect(recorded_seed, seed, "completion seed") + _expect(recorded_steps, total_steps, "completion optimizer_steps") + _expect( + str(manifest.get("protocol_fingerprint", "")), + fingerprint, + "manifest fingerprint", + ) + _expect(str(manifest.get("optimizer", "")), optimizer, "manifest optimizer") + _expect(_as_int(manifest.get("seed"), "manifest seed"), seed, "manifest seed") + _expect( + _as_int(manifest.get("max_steps"), "manifest max_steps"), + total_steps, + "manifest max_steps", + ) + + final_test = test_results.get("final") + selected_test = test_results.get("validation_selected") + if not isinstance(final_test, dict) or not isinstance(selected_test, dict): + _fail("test_results.json lacks final or validation_selected results") + _expect( + _as_int(final_test.get("step"), "final test step"), + total_steps, + "final test step", + ) + _expect( + _as_int(selected_test.get("step"), "selected test step"), + best_step, + "selected test step", + ) + + metrics = _read_csv(root / "metrics.csv") + epoch_metrics = _read_csv(root / "epoch_metrics.csv") + layers = _read_csv(root / "spectral" / "layers.csv") + summary = _read_csv(root / "spectral" / "summary.csv") + metric_steps = _step_tuple(metrics, "metrics.csv") + epoch_steps = _step_tuple(epoch_metrics, "epoch_metrics.csv") + summary_steps = _step_tuple(summary, "spectral/summary.csv") + + for label, steps in ( + ("metrics.csv", metric_steps), + ("epoch_metrics.csv", epoch_steps), + ): + if 0 not in steps or total_steps not in steps or max(steps) != total_steps: + _fail(f"{label} does not span step zero through {total_steps}") + + if "test_monitoring_only" not in epoch_metrics.columns: + _fail("epoch_metrics.csv has no test_monitoring_only column") + policy = pd.to_numeric(epoch_metrics["test_monitoring_only"], errors="coerce") + if policy.isna().any() or not policy.astype(int).eq(1).all(): + _fail("epoch_metrics.csv violates the monitoring-only test policy") + + required_layer_columns = { + "step", + "matrix_name", + "alpha", + "ERG_gap", + "num_traps", + } + missing_columns = required_layer_columns.difference(layers.columns) + if missing_columns: + _fail( + "spectral/layers.csv is missing columns " + + ", ".join(sorted(missing_columns)) + ) + if layers.duplicated(["step", "matrix_name"]).any(): + _fail("spectral/layers.csv has duplicate step/matrix rows") + layer_steps = _step_tuple( + layers[["step"]].drop_duplicates().sort_values("step"), + "spectral/layers.csv", + ) + if set(summary_steps) != set(epoch_steps) or set(layer_steps) != set(epoch_steps): + _fail("spectral steps do not match epoch_metrics.csv") + if not layers.groupby("step")["matrix_name"].nunique().eq(6).all(): + _fail("spectral/layers.csv does not contain six matrices per epoch") + if "n_matrices" not in summary.columns: + _fail("spectral/summary.csv has no n_matrices column") + matrix_counts = pd.to_numeric(summary["n_matrices"], errors="coerce") + if matrix_counts.isna().any() or not matrix_counts.astype(int).eq(6).all(): + _fail("spectral/summary.csv does not report six matrices per epoch") + + if "checkpoint_path" not in epoch_metrics.columns: + _fail("epoch_metrics.csv has no checkpoint_path column") + recorded_checkpoint_paths = [ + Path(str(value)) + for value in epoch_metrics["checkpoint_path"] + ] + if len(recorded_checkpoint_paths) != len(set(epoch_steps)): + _fail( + "epoch checkpoint inventory does not match " + "epoch_metrics.csv" + ) + resolved_checkpoint_paths: list[Path] = [] + for recorded in recorded_checkpoint_paths: + candidate = recorded + if not candidate.is_file(): + candidate = root / "epoch_checkpoints" / recorded.name + if not candidate.is_file() or candidate.stat().st_size == 0: + _fail( + f"missing epoch checkpoint recorded by " + f"epoch_metrics.csv: {recorded}" + ) + resolved_checkpoint_paths.append(candidate.resolve()) + if len(resolved_checkpoint_paths) != len( + set(resolved_checkpoint_paths) + ): + _fail( + "epoch_metrics.csv references duplicate epoch " + "checkpoints" + ) + + if verify_checkpoints: + try: + best_loss = float(completion.get("best_validation_loss")) + except (TypeError, ValueError): + _fail("run_complete.json has invalid best_validation_loss") + for filename, expected_step in ( + ("checkpoint_latest.pt", total_steps), + ("checkpoint_final.pt", total_steps), + ("checkpoint_best.pt", best_step), + ): + payload = _load_checkpoint(root / filename) + _expect( + str(payload.get("fingerprint", "")), + fingerprint, + f"{filename} fingerprint", + ) + _expect( + str(payload.get("optimizer_name", "")), + optimizer, + f"{filename} optimizer", + ) + _expect( + _as_int(payload.get("seed"), f"{filename} seed"), + seed, + f"{filename} seed", + ) + _expect( + _as_int(payload.get("step"), f"{filename} step"), + expected_step, + f"{filename} step", + ) + _expect( + _as_int( + payload.get("best_validation_step"), + f"{filename} best_validation_step", + ), + best_step, + f"{filename} best_validation_step", + ) + try: + stored_loss = float(payload.get("best_validation_loss")) + except (TypeError, ValueError): + _fail(f"{filename} has invalid best_validation_loss") + if not math.isclose( + stored_loss, best_loss, rel_tol=1e-12, abs_tol=1e-12 + ): + _fail(f"{filename} best_validation_loss does not match completion") + + return completion diff --git a/baseline/nanogpt_muon_hyperball/src/rg_nanogpt_muon_hyperball/config.py b/baseline/nanogpt_muon_hyperball/src/rg_nanogpt_muon_hyperball/config.py new file mode 100644 index 0000000..8e37d76 --- /dev/null +++ b/baseline/nanogpt_muon_hyperball/src/rg_nanogpt_muon_hyperball/config.py @@ -0,0 +1,278 @@ +from __future__ import annotations + +from copy import deepcopy +import hashlib +import json +import math +import os +from pathlib import Path +from typing import Any + +import yaml + +SUPPORTED_OPTIMIZERS = ("muon", "muon_hyperball") +DEFAULT_ROOT = Path("/tmp/rg-nanogpt-muon-hyperball") + + +def roots() -> dict[str, Path]: + root = Path(os.environ.get("RG_NANOGPT_MUON_HYPERBALL_ROOT", DEFAULT_ROOT)) + return { + "root": root, + "data": Path( + os.environ.get("RG_NANOGPT_MUON_HYPERBALL_DATA_ROOT", root / "data") + ), + "results": Path( + os.environ.get( + "RG_NANOGPT_MUON_HYPERBALL_RESULTS_ROOT", root / "results" + ) + ), + "plots": Path( + os.environ.get("RG_NANOGPT_MUON_HYPERBALL_PLOTS_ROOT", root / "plots") + ), + } + + +def load_config(path: str | Path) -> dict[str, Any]: + with Path(path).open("r", encoding="utf-8") as handle: + cfg = yaml.safe_load(handle) + if not isinstance(cfg, dict): + raise ValueError("configuration root must be a mapping") + validate_config(cfg) + return cfg + + +def validate_config(cfg: dict[str, Any]) -> None: + for section in ( + "protocol", + "dataset", + "model", + "training", + "optimizer_profiles", + "evaluation", + "weightwatcher", + "runtime", + ): + if section not in cfg: + raise ValueError(f"missing configuration section: {section}") + + if int(cfg["protocol"].get("version", 0)) < 1: + raise ValueError("protocol.version must be positive") + + model = cfg["model"] + for key in ("vocab_size", "block_size", "n_layer", "n_head", "n_embd"): + if int(model[key]) < 1: + raise ValueError(f"model.{key} must be positive") + if int(model["n_head"]) != 1 or int(model["n_layer"]) != 1: + raise ValueError("this experiment is fixed to one block and one head") + if int(model["n_embd"]) % int(model["n_head"]) != 0: + raise ValueError("model.n_embd must be divisible by model.n_head") + if not 0.0 <= float(model.get("dropout", 0.0)) < 1.0: + raise ValueError("model.dropout must be in [0, 1)") + + dataset = cfg["dataset"] + for key in ("train_tokens", "val_tokens", "test_tokens"): + if int(dataset[key]) <= int(model["block_size"]) + 1: + raise ValueError(f"dataset.{key} is too small") + + training = cfg["training"] + for key in ( + "batch_size", + "grad_accum_steps", + "target_epochs", + "eval_interval_steps", + "eval_batches", + "checkpoint_interval_steps", + "epoch_interval", + ): + if float(training[key]) <= 0: + raise ValueError(f"training.{key} must be positive") + seeds = [int(seed) for seed in training["seeds"]] + if not seeds or len(set(seeds)) != len(seeds): + raise ValueError("training.seeds must contain unique values") + if float(training["grad_clip"]) < 0: + raise ValueError("training.grad_clip must be nonnegative") + + profiles = cfg["optimizer_profiles"] + for name in SUPPORTED_OPTIMIZERS: + if name not in profiles: + raise ValueError(f"missing optimizer profile: {name}") + validate_optimizer_profile({**profiles[name], "name": name}) + + evaluation = cfg["evaluation"] + for key in ( + "bleu_examples", + "bleu_prompt_tokens", + "bleu_continuation_tokens", + "bleu_batch_size", + ): + if int(evaluation[key]) < 1: + raise ValueError(f"evaluation.{key} must be positive") + if ( + int(evaluation["bleu_prompt_tokens"]) + + int(evaluation["bleu_continuation_tokens"]) + > int(model["block_size"]) + ): + raise ValueError("BLEU prompt plus continuation exceeds context length") + + probe_keys = ( + "train_probe_seed", + "validation_probe_seed", + "test_probe_seed", + "bleu_probe_seed", + ) + probe_seeds = [int(evaluation[key]) for key in probe_keys] + if any(seed < 0 for seed in probe_seeds) or len(set(probe_seeds)) != 4: + raise ValueError("evaluation probe seeds must be distinct and nonnegative") + + ww = cfg["weightwatcher"] + if not bool(ww.get("ERG", False)) or not bool(ww.get("randomize", False)): + raise ValueError("WeightWatcher ERG and randomize must both be enabled") + if int(ww["min_evals"]) < 5: + raise ValueError("weightwatcher.min_evals must be at least 5") + + +def validate_optimizer_profile(profile: dict[str, Any]) -> None: + family = str(profile.get("family", "")) + if family not in {"muon", "muon_hyperball"}: + raise ValueError(f"unsupported optimizer family: {family}") + if str(profile.get("schedule")) != "warmup_cosine_then_floor": + raise ValueError("schedule must be warmup_cosine_then_floor") + + warmup_fraction = float(profile.get("warmup_fraction", -1.0)) + if not 0.0 <= warmup_fraction < 1.0: + raise ValueError("warmup_fraction must be in [0, 1)") + if "warmup_steps" in profile and int(profile["warmup_steps"]) < 0: + raise ValueError("warmup_steps must be nonnegative") + if float(profile.get("lr_schedule_epochs", 0.0)) <= 0: + raise ValueError("lr_schedule_epochs must be positive") + + for peak_key, floor_key in ( + ("matrix_learning_rate", "matrix_min_learning_rate"), + ("aux_learning_rate", "aux_min_learning_rate"), + ): + peak = float(profile[peak_key]) + floor = float(profile[floor_key]) + if peak <= 0 or floor < 0 or floor > peak: + raise ValueError(f"{peak_key}/{floor_key} values are inconsistent") + + if int(profile["newton_schulz_steps"]) < 1: + raise ValueError("newton_schulz_steps must be positive") + if float(profile.get("muon_epsilon", 0.0)) <= 0: + raise ValueError("muon_epsilon must be positive") + + if family == "muon_hyperball": + if float(profile.get("hyperball_relative_radius", 0.0)) <= 0: + raise ValueError("hyperball_relative_radius must be positive") + if float(profile.get("hyperball_epsilon", 0.0)) <= 0: + raise ValueError("hyperball_epsilon must be positive") + + +def optimizer_profile(cfg: dict[str, Any], optimizer: str) -> dict[str, Any]: + optimizer = str(optimizer).lower() + if optimizer not in SUPPORTED_OPTIMIZERS: + raise ValueError( + f"unsupported optimizer {optimizer!r}; choose from {SUPPORTED_OPTIMIZERS}" + ) + profile = deepcopy(cfg["optimizer_profiles"][optimizer]) + profile["name"] = optimizer + validate_optimizer_profile(profile) + return profile + + +def canonical_seeds(cfg: dict[str, Any]) -> tuple[int, ...]: + return tuple(int(seed) for seed in cfg["training"]["seeds"]) + + +def tokens_per_step(cfg: dict[str, Any]) -> int: + return ( + int(cfg["training"]["batch_size"]) + * int(cfg["training"]["grad_accum_steps"]) + * int(cfg["model"]["block_size"]) + ) + + +def _steps_for_epochs(cfg: dict[str, Any], epochs: float, train_tokens: int) -> int: + target_tokens = float(epochs) * int(train_tokens) + return max(1, int(math.ceil(target_tokens / tokens_per_step(cfg)))) + + +def max_steps(cfg: dict[str, Any], train_tokens: int | None = None) -> int: + train_tokens = int(train_tokens or cfg["dataset"]["train_tokens"]) + return _steps_for_epochs( + cfg, float(cfg["training"]["target_epochs"]), train_tokens + ) + + +def lr_schedule_steps( + cfg: dict[str, Any], + profile: dict[str, Any], + train_tokens: int | None = None, +) -> int: + train_tokens = int(train_tokens or cfg["dataset"]["train_tokens"]) + steps = _steps_for_epochs( + cfg, float(profile["lr_schedule_epochs"]), train_tokens + ) + return min(max_steps(cfg, train_tokens), steps) + + +def warmup_steps(profile: dict[str, Any], total_steps: int) -> int: + if total_steps < 2: + return 0 + if "warmup_steps" in profile: + requested = int(profile["warmup_steps"]) + if requested < 0: + raise ValueError("warmup_steps must be nonnegative") + return min(total_steps - 1, requested) + return min( + total_steps - 1, + max(1, int(round(total_steps * float(profile["warmup_fraction"])))), + ) + + +def epoch_step_map( + cfg: dict[str, Any], train_tokens: int | None = None +) -> dict[int, float]: + train_tokens = int(train_tokens or cfg["dataset"]["train_tokens"]) + total_steps = max_steps(cfg, train_tokens) + step_tokens = tokens_per_step(cfg) + target_epochs = float(cfg["training"]["target_epochs"]) + interval = float(cfg["training"]["epoch_interval"]) + + points = [0.0] + current = interval + while current < target_epochs - 1e-12: + points.append(round(current, 12)) + current += interval + points.append(target_epochs) + + result: dict[int, float] = {} + for epoch in points: + step = 0 if epoch == 0 else int(round(epoch * train_tokens / step_tokens)) + result[min(total_steps, max(0, step))] = float(epoch) + result[total_steps] = target_epochs + return dict(sorted(result.items())) + + +def protocol_fingerprint( + cfg: dict[str, Any], + *, + optimizer: str, + seed: int, + data_metadata: dict[str, Any], +) -> str: + payload = { + "protocol": cfg["protocol"], + "dataset": cfg["dataset"], + "model": cfg["model"], + "training": cfg["training"], + "optimizer_profile": optimizer_profile(cfg, optimizer), + "evaluation": cfg["evaluation"], + "weightwatcher": cfg["weightwatcher"], + "optimizer": str(optimizer), + "seed": int(seed), + "data_metadata": data_metadata, + } + canonical = json.dumps( + payload, sort_keys=True, separators=(",", ":"), default=str + ) + return hashlib.sha256(canonical.encode("utf-8")).hexdigest() diff --git a/baseline/nanogpt_muon_hyperball/src/rg_nanogpt_muon_hyperball/data.py b/baseline/nanogpt_muon_hyperball/src/rg_nanogpt_muon_hyperball/data.py new file mode 100644 index 0000000..63015c4 --- /dev/null +++ b/baseline/nanogpt_muon_hyperball/src/rg_nanogpt_muon_hyperball/data.py @@ -0,0 +1,328 @@ +from __future__ import annotations + +import argparse +import hashlib +import json +import os +from pathlib import Path +import sys +import time +from typing import Iterable, Protocol + +import numpy as np + +from .config import load_config, roots + +TOKEN_DTYPE = np.dtype(np.uint16) +SPLIT_NAMES = ("train", "val", "test") + + +class Encoder(Protocol): + n_vocab: int + eot_token: int + + def encode_ordinary(self, text: str) -> list[int]: ... + + +def _sha256(path: Path, chunk_size: int = 4 * 1024 * 1024) -> str: + digest = hashlib.sha256() + with path.open("rb") as handle: + while chunk := handle.read(chunk_size): + digest.update(chunk) + return digest.hexdigest() + + +def _encode_document(text: str, encoder: Encoder) -> np.ndarray: + tokens = encoder.encode_ordinary(text) + tokens.append(int(encoder.eot_token)) + if tokens and max(tokens) > np.iinfo(TOKEN_DTYPE).max: + raise ValueError("token id exceeds uint16 storage capacity") + return np.asarray(tokens, dtype=TOKEN_DTYPE) + + +def write_token_splits( + texts: Iterable[str], + encoder: Encoder, + output_dir: str | Path, + *, + train_tokens: int, + val_tokens: int, + test_tokens: int, + dataset_metadata: dict[str, object] | None = None, + progress_every_documents: int = 2_000, +) -> dict[str, object]: + """Write exact, document-disjoint splits without loading the corpus in RAM.""" + + output_dir = Path(output_dir) + output_dir.mkdir(parents=True, exist_ok=True) + targets = { + "train": int(train_tokens), + "val": int(val_tokens), + "test": int(test_tokens), + } + if any(value <= 0 for value in targets.values()): + raise ValueError("all split sizes must be positive") + + partial = {name: output_dir / f"{name}.bin.partial" for name in targets} + final = {name: output_dir / f"{name}.bin" for name in targets} + for path in partial.values(): + path.unlink(missing_ok=True) + + handles = {name: path.open("wb") for name, path in partial.items()} + written = {name: 0 for name in targets} + document_counts = {name: 0 for name in targets} + split_names = list(targets) + split_index = 0 + documents = 0 + started = time.monotonic() + + try: + for text in texts: + if split_index >= len(split_names): + break + documents += 1 + encoded = _encode_document(str(text), encoder) + split = split_names[split_index] + remaining = targets[split] - written[split] + take = min(remaining, len(encoded)) + if take: + encoded[:take].tofile(handles[split]) + written[split] += int(take) + document_counts[split] += 1 + # Never carry a document remainder into the next split. This is the + # invariant that makes train/validation/test document-disjoint. + if written[split] == targets[split]: + split_index += 1 + if progress_every_documents and documents % int(progress_every_documents) == 0: + total = sum(written.values()) + required = sum(targets.values()) + elapsed = time.monotonic() - started + rate = total / max(elapsed, 1e-9) + print( + f"[one-head-data] documents={documents:,} " + f"tokens={total:,}/{required:,} " + f"({100 * total / max(required, 1):.1f}%) " + f"rate={rate:,.0f} tok/s", + file=sys.stderr, + flush=True, + ) + finally: + for handle in handles.values(): + handle.flush() + os.fsync(handle.fileno()) + handle.close() + + if written != targets: + for path in partial.values(): + path.unlink(missing_ok=True) + raise RuntimeError( + f"stream ended before exact splits were filled: {written} != {targets}" + ) + + for name in split_names: + os.replace(partial[name], final[name]) + + metadata: dict[str, object] = { + "schema_version": 2, + "tokenizer": "gpt2", + "vocab_size": int(encoder.n_vocab), + "eot_token": int(encoder.eot_token), + "dtype": TOKEN_DTYPE.name, + "splits": written, + "split_document_counts": document_counts, + "document_disjoint_splits": True, + "documents_consumed": int(documents), + "files": { + name: { + "path": final[name].name, + "sha256": _sha256(final[name]), + "bytes": int(final[name].stat().st_size), + } + for name in split_names + }, + } + if dataset_metadata: + metadata.update(dataset_metadata) + temporary = output_dir / "meta.json.tmp" + temporary.write_text( + json.dumps(metadata, indent=2, sort_keys=True), encoding="utf-8" + ) + temporary.replace(output_dir / "meta.json") + return metadata + + +def _validate_file_identity( + output_dir: Path, + metadata: dict[str, object], + expected_splits: dict[str, int], +) -> None: + file_metadata = metadata.get("files") + if not isinstance(file_metadata, dict): + raise RuntimeError( + "prepared corpus metadata does not contain file hashes; " + "re-run data preparation with --force" + ) + + for split in SPLIT_NAMES: + record = file_metadata.get(split) + if not isinstance(record, dict): + raise RuntimeError(f"prepared corpus metadata is missing files.{split}") + path = output_dir / str(record.get("path", f"{split}.bin")) + if path.name != f"{split}.bin" or not path.is_file(): + raise RuntimeError(f"prepared {split} file path is invalid: {path}") + + expected_bytes = expected_splits[split] * TOKEN_DTYPE.itemsize + recorded_bytes = int(record.get("bytes", -1)) + actual_bytes = int(path.stat().st_size) + if recorded_bytes != expected_bytes or actual_bytes != expected_bytes: + raise RuntimeError( + f"prepared {split} byte size mismatch: " + f"metadata={recorded_bytes}, actual={actual_bytes}, " + f"expected={expected_bytes}" + ) + + recorded_hash = str(record.get("sha256", "")) + actual_hash = _sha256(path) + if not recorded_hash or actual_hash != recorded_hash: + raise RuntimeError( + f"prepared {split} SHA-256 mismatch; the cached corpus is corrupt " + "or was modified. Re-run data preparation with --force." + ) + + +def validate_prepared_data( + output_dir: str | Path, + cfg: dict, +) -> dict[str, object]: + """Validate dataset identity, exact sizes, and every persisted file hash.""" + + output_dir = Path(output_dir) + metadata_path = output_dir / "meta.json" + required = [ + metadata_path, + *(output_dir / f"{split}.bin" for split in SPLIT_NAMES), + ] + missing = [str(path) for path in required if not path.is_file()] + if missing: + raise FileNotFoundError( + "prepared data are incomplete: " + ", ".join(missing) + ) + + metadata = json.loads(metadata_path.read_text(encoding="utf-8")) + expected_splits = { + "train": int(cfg["dataset"]["train_tokens"]), + "val": int(cfg["dataset"]["val_tokens"]), + "test": int(cfg["dataset"]["test_tokens"]), + } + if metadata.get("splits") != expected_splits: + raise RuntimeError( + "prepared split sizes do not match config: " + f"{metadata.get('splits')} != {expected_splits}" + ) + if metadata.get("dataset_name") != cfg["dataset"]["name"]: + raise RuntimeError("prepared dataset identity does not match config") + if metadata.get("dataset_config") != cfg["dataset"]["config"]: + raise RuntimeError("prepared dataset configuration does not match config") + if metadata.get("dataset_revision") != cfg["dataset"]["revision"]: + raise RuntimeError("prepared dataset revision does not match config") + if metadata.get("tokenizer") != "gpt2": + raise RuntimeError("prepared tokenizer must be GPT-2 BPE") + if metadata.get("dtype") != TOKEN_DTYPE.name: + raise RuntimeError("prepared token dtype must be uint16") + if metadata.get("document_disjoint_splits") is not True: + raise RuntimeError("prepared splits are not marked document-disjoint") + + _validate_file_identity(output_dir, metadata, expected_splits) + return metadata + + +def prepare_fineweb_edu( + cfg: dict, + output_dir: str | Path, + *, + force: bool = False, +) -> dict[str, object]: + output_dir = Path(output_dir) + if not force: + try: + metadata = validate_prepared_data(output_dir, cfg) + print(f"[one-head-data] reusing verified data at {output_dir}") + return metadata + except FileNotFoundError: + pass + + try: + import tiktoken + from datasets import load_dataset + except ImportError as exc: + raise RuntimeError( + "data preparation requires datasets and tiktoken; " + "install dependencies into the active conda environment with " + "`python -m pip install -e .`" + ) from exc + + dataset_cfg = cfg["dataset"] + print( + "[one-head-data] streaming pinned FineWeb-Edu; " + "the first run requires internet access", + flush=True, + ) + stream = load_dataset( + str(dataset_cfg["name"]), + name=str(dataset_cfg["config"]), + split=str(dataset_cfg.get("split", "train")), + revision=str(dataset_cfg["revision"]), + streaming=True, + ) + encoder = tiktoken.get_encoding("gpt2") + metadata = write_token_splits( + (row["text"] for row in stream), + encoder, + output_dir, + train_tokens=int(dataset_cfg["train_tokens"]), + val_tokens=int(dataset_cfg["val_tokens"]), + test_tokens=int(dataset_cfg["test_tokens"]), + dataset_metadata={ + "dataset_name": str(dataset_cfg["name"]), + "dataset_config": str(dataset_cfg["config"]), + "dataset_split": str(dataset_cfg.get("split", "train")), + "dataset_revision": str(dataset_cfg["revision"]), + }, + ) + validate_prepared_data(output_dir, cfg) + print(f"[one-head-data] complete and verified: {output_dir}") + return metadata + + +def load_memmaps( + output_dir: str | Path, + cfg: dict, +) -> tuple[dict[str, object], dict[str, np.memmap]]: + output_dir = Path(output_dir) + metadata = validate_prepared_data(output_dir, cfg) + arrays = { + split: np.memmap( + output_dir / f"{split}.bin", + dtype=TOKEN_DTYPE, + mode="r", + ) + for split in SPLIT_NAMES + } + return metadata, arrays + + +def main() -> None: + parser = argparse.ArgumentParser( + description="Prepare the pinned FineWeb-Edu one-head baseline corpus" + ) + parser.add_argument("--config", required=True) + parser.add_argument("--output-dir") + parser.add_argument("--force", action="store_true") + args = parser.parse_args() + cfg = load_config(args.config) + output = Path(args.output_dir) if args.output_dir else roots()["data"] + prepare_fineweb_edu(cfg, output, force=args.force) + + +if __name__ == "__main__": + main() diff --git a/baseline/nanogpt_muon_hyperball/src/rg_nanogpt_muon_hyperball/engine.py b/baseline/nanogpt_muon_hyperball/src/rg_nanogpt_muon_hyperball/engine.py new file mode 100644 index 0000000..d9d5297 --- /dev/null +++ b/baseline/nanogpt_muon_hyperball/src/rg_nanogpt_muon_hyperball/engine.py @@ -0,0 +1,325 @@ +from __future__ import annotations + +import csv +import json +from pathlib import Path +import shutil + +import torch + +from .completion import validate_completed_run +from .checkpoints import load_training_checkpoint, save_training_checkpoint +from .config import ( + SUPPORTED_OPTIMIZERS, + epoch_step_map, + max_steps, + optimizer_profile, + protocol_fingerprint, + tokens_per_step, + warmup_steps, +) +from .data import load_memmaps +from .evaluation import fixed_bleu_probe, fixed_probe +from .model import GPT, GPTConfig +from .optimizers import make_optimizer_handles +from .run_utils import ( + EPOCH_FIELDS, + METRIC_FIELDS, + checkpoint_eval, + prepare_csv, + run_directory, + truncate_spectral_after, + write_manifest, +) +from .runtime import choose_device, configure_runtime, seed_everything +from .train_loop import execute_training_loop + + +def run_one( + *, + cfg: dict, + data_root: str | Path, + results_root: str | Path, + optimizer_name: str, + seed: int, + device: str = "auto", + resume: bool = True, + overwrite: bool = False, + progress: bool = True, +) -> Path: + optimizer_name = str(optimizer_name).lower() + if optimizer_name not in SUPPORTED_OPTIMIZERS: + raise ValueError(f"unsupported optimizer: {optimizer_name}") + if resume and overwrite: + raise ValueError("resume and overwrite are mutually exclusive") + + data_root = Path(data_root) + results_root = Path(results_root) + run_dir = run_directory(results_root, optimizer_name, int(seed)) + completion_path = run_dir / "run_complete.json" + if run_dir.exists() and overwrite: + shutil.rmtree(run_dir) + if run_dir.exists() and not resume: + raise FileExistsError( + f"incomplete run exists: {run_dir}; enable resume or overwrite" + ) + run_dir.mkdir(parents=True, exist_ok=True) + + data_metadata, arrays = load_memmaps(data_root, cfg) + train_tokens = int(data_metadata["splits"]["train"]) + total_steps = max_steps(cfg, train_tokens) + profile = optimizer_profile(cfg, optimizer_name) + warmup = warmup_steps(profile, total_steps) + fingerprint = protocol_fingerprint( + cfg, + optimizer=optimizer_name, + seed=int(seed), + data_metadata=data_metadata, + ) + if completion_path.is_file(): + validate_completed_run( + run_dir, + expected_fingerprint=fingerprint, + expected_optimizer=optimizer_name, + expected_seed=int(seed), + expected_total_steps=total_steps, + verify_checkpoints=True, + ) + if progress: + print( + "[one-head-train] reuse verified completed " + f"{optimizer_name} seed={seed}" + ) + return run_dir + + resolved_device = choose_device(device) + configure_runtime(resolved_device, cfg) + seed_everything(int(seed)) + model = GPT(GPTConfig(**cfg["model"])).to(resolved_device) + handles = make_optimizer_handles(model, profile) + train_generator = torch.Generator(device="cpu").manual_seed(int(seed) + 11) + + batch_size = int(cfg["training"]["batch_size"]) + eval_batches = int(cfg["training"]["eval_batches"]) + block_size = int(cfg["model"]["block_size"]) + epoch_steps = epoch_step_map(cfg, train_tokens) + eval_cfg = cfg["evaluation"] + + # Evaluation examples are deliberately independent of the training seed. + # This makes paired optimizer comparisons and across-seed uncertainty use + # the same train/validation/test probe windows in every complete run. + train_probe = fixed_probe( + arrays["train"], + batch_size=batch_size, + block_size=block_size, + n_batches=eval_batches, + seed=int(eval_cfg["train_probe_seed"]), + ) + val_probe = fixed_probe( + arrays["val"], + batch_size=batch_size, + block_size=block_size, + n_batches=eval_batches, + seed=int(eval_cfg["validation_probe_seed"]), + ) + test_probe = fixed_probe( + arrays["test"], + batch_size=batch_size, + block_size=block_size, + n_batches=eval_batches, + seed=int(eval_cfg["test_probe_seed"]), + ) + bleu_probe = fixed_bleu_probe( + arrays["test"], + examples=int(eval_cfg["bleu_examples"]), + prompt_tokens=int(eval_cfg["bleu_prompt_tokens"]), + continuation_tokens=int(eval_cfg["bleu_continuation_tokens"]), + seed=int(eval_cfg["bleu_probe_seed"]), + ) + + start_step = 0 + best_validation_loss = float("inf") + best_validation_step = 0 + elapsed_offset = 0.0 + latest_checkpoint = run_dir / "checkpoint_latest.pt" + best_checkpoint = run_dir / "checkpoint_best.pt" + final_checkpoint = run_dir / "checkpoint_final.pt" + if resume and latest_checkpoint.is_file(): + ( + start_step, + best_validation_loss, + best_validation_step, + elapsed_offset, + ) = load_training_checkpoint( + latest_checkpoint, + model=model, + handles=handles, + expected_fingerprint=fingerprint, + train_generator=train_generator, + ) + model.to(resolved_device) + truncate_spectral_after(run_dir, start_step) + if progress: + print( + f"[one-head-train] resume {optimizer_name} " + f"seed={seed} step={start_step}" + ) + elif run_dir.exists() and any(run_dir.iterdir()) and resume: + nontrivial = [ + path for path in run_dir.iterdir() if path.name != "manifest.json" + ] + if nontrivial and not latest_checkpoint.is_file(): + raise FileNotFoundError( + f"cannot resume {run_dir}: checkpoint_latest.pt is missing" + ) + + write_manifest( + run_dir, + cfg=cfg, + data_metadata=data_metadata, + optimizer_name=optimizer_name, + profile=profile, + seed=int(seed), + device=resolved_device, + total_steps=total_steps, + warmup=warmup, + fingerprint=fingerprint, + model=model, + ) + + metrics_path = run_dir / "metrics.csv" + epoch_metrics_path = run_dir / "epoch_metrics.csv" + prepare_csv( + metrics_path, + METRIC_FIELDS, + start_step if start_step else None, + ) + prepare_csv( + epoch_metrics_path, + EPOCH_FIELDS, + start_step if start_step else None, + ) + with ( + metrics_path.open("a", newline="", encoding="utf-8") as metrics_handle, + epoch_metrics_path.open( + "a", newline="", encoding="utf-8" + ) as epoch_handle, + ): + best_validation_loss, best_validation_step, elapsed_total = ( + execute_training_loop( + cfg=cfg, + model=model, + handles=handles, + arrays=arrays, + train_probe=train_probe, + val_probe=val_probe, + test_probe=test_probe, + bleu_probe=bleu_probe, + device=resolved_device, + optimizer_name=optimizer_name, + seed=int(seed), + train_tokens=train_tokens, + total_steps=total_steps, + warmup=warmup, + start_step=start_step, + best_validation_loss=best_validation_loss, + best_validation_step=best_validation_step, + elapsed_offset=elapsed_offset, + fingerprint=fingerprint, + train_generator=train_generator, + epoch_steps=epoch_steps, + metrics_writer=csv.DictWriter( + metrics_handle, fieldnames=METRIC_FIELDS + ), + metrics_handle=metrics_handle, + epoch_writer=csv.DictWriter( + epoch_handle, fieldnames=EPOCH_FIELDS + ), + epoch_handle=epoch_handle, + run_dir=run_dir, + latest_checkpoint=latest_checkpoint, + best_checkpoint=best_checkpoint, + progress=progress, + ) + ) + + for checkpoint in (final_checkpoint, latest_checkpoint): + save_training_checkpoint( + checkpoint, + model=model, + handles=handles, + step=total_steps, + best_validation_loss=best_validation_loss, + best_validation_step=best_validation_step, + elapsed_seconds=elapsed_total, + fingerprint=fingerprint, + cfg=cfg, + optimizer_name=optimizer_name, + seed=int(seed), + train_generator=train_generator, + ) + + final_state = torch.load( + final_checkpoint, + map_location="cpu", + weights_only=False, + )["model"] + final_test = checkpoint_eval( + final_checkpoint, + model=model, + test_probe=test_probe, + bleu_probe=bleu_probe, + device=resolved_device, + bleu_batch_size=int(eval_cfg["bleu_batch_size"]), + ) + best_test = checkpoint_eval( + best_checkpoint, + model=model, + test_probe=test_probe, + bleu_probe=bleu_probe, + device=resolved_device, + bleu_batch_size=int(eval_cfg["bleu_batch_size"]), + ) + model.load_state_dict(final_state) + model.to(resolved_device) + + test_results = { + "policy": ( + "test is monitoring-only; validation loss selects " + "checkpoint_best.pt" + ), + "final": final_test, + "validation_selected": best_test, + } + (run_dir / "test_results.json").write_text( + json.dumps(test_results, indent=2, sort_keys=True), + encoding="utf-8", + ) + completion = { + "completed": True, + "optimizer": optimizer_name, + "seed": int(seed), + "optimizer_steps": int(total_steps), + "train_epochs": float( + total_steps * tokens_per_step(cfg) / train_tokens + ), + "elapsed_seconds": float(elapsed_total), + "best_validation_step": int(best_validation_step), + "best_validation_loss": float(best_validation_loss), + "final_test_loss": float(final_test["loss"]), + "final_test_perplexity": float(final_test["perplexity"]), + "final_test_accuracy": float(final_test["accuracy"]), + "final_test_bleu": float(final_test["bleu"]), + "fingerprint": fingerprint, + } + temporary = run_dir / "run_complete.json.tmp" + temporary.write_text( + json.dumps(completion, indent=2, sort_keys=True), + encoding="utf-8", + ) + temporary.replace(completion_path) + if progress: + print( + f"[one-head-train] complete {optimizer_name} seed={seed}: {run_dir}" + ) + return run_dir diff --git a/baseline/nanogpt_muon_hyperball/src/rg_nanogpt_muon_hyperball/evaluation.py b/baseline/nanogpt_muon_hyperball/src/rg_nanogpt_muon_hyperball/evaluation.py new file mode 100644 index 0000000..c1f1cfd --- /dev/null +++ b/baseline/nanogpt_muon_hyperball/src/rg_nanogpt_muon_hyperball/evaluation.py @@ -0,0 +1,175 @@ +from __future__ import annotations + +from dataclasses import dataclass +import math +from typing import Iterable + +import numpy as np +import torch +import torch.nn as nn + + +@dataclass(frozen=True) +class BleuProbe: + prompts: torch.Tensor + references: torch.Tensor + prompt_tokens: int + continuation_tokens: int + + +def random_batch( + data: np.memmap, + *, + batch_size: int, + block_size: int, + generator: torch.Generator, +) -> tuple[torch.Tensor, torch.Tensor]: + if len(data) <= block_size + 1: + raise ValueError("data split is too short for the configured block size") + starts = torch.randint( + len(data) - block_size - 1, + (int(batch_size),), + generator=generator, + ).tolist() + x = torch.stack( + [torch.from_numpy(np.asarray(data[start : start + block_size], dtype=np.int64)) for start in starts] + ) + y = torch.stack( + [ + torch.from_numpy( + np.asarray(data[start + 1 : start + 1 + block_size], dtype=np.int64) + ) + for start in starts + ] + ) + return x, y + + +def fixed_probe( + data: np.memmap, + *, + batch_size: int, + block_size: int, + n_batches: int, + seed: int, +) -> list[tuple[torch.Tensor, torch.Tensor]]: + generator = torch.Generator(device="cpu").manual_seed(int(seed)) + return [ + random_batch( + data, + batch_size=int(batch_size), + block_size=int(block_size), + generator=generator, + ) + for _ in range(int(n_batches)) + ] + + +def fixed_bleu_probe( + data: np.memmap, + *, + examples: int, + prompt_tokens: int, + continuation_tokens: int, + seed: int, +) -> BleuProbe: + total = int(prompt_tokens) + int(continuation_tokens) + if len(data) <= total + 1: + raise ValueError("test split is too short for BLEU continuation probes") + generator = torch.Generator(device="cpu").manual_seed(int(seed)) + starts = torch.randint(len(data) - total - 1, (int(examples),), generator=generator).tolist() + prompts = torch.stack( + [torch.from_numpy(np.asarray(data[start : start + prompt_tokens], dtype=np.int64)) for start in starts] + ) + references = torch.stack( + [ + torch.from_numpy( + np.asarray( + data[start + prompt_tokens : start + prompt_tokens + continuation_tokens], + dtype=np.int64, + ) + ) + for start in starts + ] + ) + return BleuProbe( + prompts=prompts, + references=references, + prompt_tokens=int(prompt_tokens), + continuation_tokens=int(continuation_tokens), + ) + + +@torch.inference_mode() +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: list[float] = [] + correct = 0 + total = 0 + for x_cpu, y_cpu in probe: + x = x_cpu.to(device) + y = y_cpu.to(device) + logits, loss = model(x, y) + if loss is None: + raise RuntimeError("evaluation forward pass did not return loss") + losses.append(float(loss.detach().cpu())) + correct += int((logits.argmax(dim=-1) == y).sum().detach().cpu()) + total += int(y.numel()) + model.train(was_training) + mean_loss = float(np.mean(losses)) + return { + "loss": mean_loss, + "perplexity": float(math.exp(min(20.0, mean_loss))), + "accuracy": correct / max(1, total), + } + + +@torch.inference_mode() +def evaluate_bleu( + model, + probe: BleuProbe, + *, + device: torch.device, + batch_size: int, +) -> dict[str, float]: + """Greedy fixed-continuation BLEU diagnostic on preregistered test segments. + + This is not a translation benchmark. It measures exact lexical overlap + between deterministic model continuations and the held-out continuation. + """ + try: + import tiktoken + from sacrebleu.metrics import BLEU + except ImportError as exc: + raise RuntimeError( + "BLEU evaluation requires tiktoken and sacrebleu; run scripts/setup_mac.sh" + ) from exc + + was_training = model.training + model.eval() + encoder = tiktoken.get_encoding("gpt2") + hypotheses: list[str] = [] + references: list[str] = [] + for start in range(0, len(probe.prompts), int(batch_size)): + prompts = probe.prompts[start : start + int(batch_size)].to(device) + generated = model.generate_greedy(prompts, probe.continuation_tokens) + continuation = generated[:, -probe.continuation_tokens :].detach().cpu() + reference_batch = probe.references[start : start + int(batch_size)] + for predicted_tokens, reference_tokens in zip(continuation, reference_batch, strict=True): + hypotheses.append(encoder.decode(predicted_tokens.tolist())) + references.append(encoder.decode(reference_tokens.tolist())) + model.train(was_training) + + bleu = BLEU(tokenize="13a", effective_order=True) + score = bleu.corpus_score(hypotheses, [references]) + return { + "bleu": float(score.score), + "bleu_examples": float(len(hypotheses)), + "bleu_sys_len": float(score.sys_len), + "bleu_ref_len": float(score.ref_len), + } diff --git a/baseline/nanogpt_muon_hyperball/src/rg_nanogpt_muon_hyperball/model.py b/baseline/nanogpt_muon_hyperball/src/rg_nanogpt_muon_hyperball/model.py new file mode 100644 index 0000000..eb696bd --- /dev/null +++ b/baseline/nanogpt_muon_hyperball/src/rg_nanogpt_muon_hyperball/model.py @@ -0,0 +1,183 @@ +from __future__ import annotations + +from dataclasses import dataclass +import math + +import torch +import torch.nn as nn +import torch.nn.functional as F + + +@dataclass(frozen=True) +class GPTConfig: + vocab_size: int = 50_257 + block_size: int = 256 + n_layer: int = 1 + n_head: int = 1 + n_embd: int = 128 + dropout: float = 0.0 + bias: bool = False + tie_weights: bool = True + + def __post_init__(self) -> None: + if self.n_layer != 1 or self.n_head != 1: + raise ValueError("the reference architecture is fixed to one block and one attention head") + if self.n_embd % self.n_head != 0: + raise ValueError("n_embd must be divisible by n_head") + if self.block_size < 2 or self.vocab_size < 2 or self.n_embd < 1: + raise ValueError("invalid GPT configuration") + if not 0.0 <= self.dropout < 1.0: + raise ValueError("dropout must be in [0, 1)") + + +class LayerNorm(nn.Module): + def __init__(self, width: int, bias: bool) -> None: + super().__init__() + self.weight = nn.Parameter(torch.ones(width)) + self.bias = nn.Parameter(torch.zeros(width)) if bias else None + + def forward(self, x: torch.Tensor) -> torch.Tensor: + return F.layer_norm(x, self.weight.shape, self.weight, self.bias, 1e-5) + + +class CausalSelfAttention(nn.Module): + def __init__(self, cfg: GPTConfig) -> None: + super().__init__() + self.n_head = cfg.n_head + self.n_embd = cfg.n_embd + self.dropout = float(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) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + batch, sequence, channels = x.shape + head_width = channels // self.n_head + q = self.q_proj(x).view(batch, sequence, self.n_head, head_width).transpose(1, 2) + k = self.k_proj(x).view(batch, sequence, self.n_head, head_width).transpose(1, 2) + v = self.v_proj(x).view(batch, sequence, self.n_head, head_width).transpose(1, 2) + 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, + ) + y = y.transpose(1, 2).contiguous().view(batch, sequence, channels) + return self.resid_dropout(self.out_proj(y)) + + +class MLP(nn.Module): + def __init__(self, cfg: GPTConfig) -> None: + super().__init__() + self.fc = nn.Linear(cfg.n_embd, 4 * cfg.n_embd, bias=cfg.bias) + self.proj = nn.Linear(4 * cfg.n_embd, cfg.n_embd, bias=cfg.bias) + self.dropout = nn.Dropout(cfg.dropout) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + return self.dropout(self.proj(F.gelu(self.fc(x), approximate="tanh"))) + + +class Block(nn.Module): + def __init__(self, cfg: GPTConfig) -> None: + super().__init__() + self.ln1 = LayerNorm(cfg.n_embd, cfg.bias) + self.attn = CausalSelfAttention(cfg) + self.ln2 = LayerNorm(cfg.n_embd, 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) -> None: + super().__init__() + self.cfg = cfg + self.token_embedding = nn.Embedding(cfg.vocab_size, cfg.n_embd) + 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 = LayerNorm(cfg.n_embd, cfg.bias) + self.lm_head = nn.Linear(cfg.n_embd, cfg.vocab_size, bias=False) + if cfg.tie_weights: + self.lm_head.weight = self.token_embedding.weight + + self.apply(self._init_module) + 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) + + @staticmethod + def _init_module(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 hidden_states(self, idx: torch.Tensor) -> torch.Tensor: + _, sequence = idx.shape + if sequence > self.cfg.block_size: + raise ValueError("input sequence exceeds model.block_size") + positions = torch.arange(sequence, device=idx.device) + x = self.drop(self.token_embedding(idx) + self.position_embedding(positions)) + for block in self.blocks: + x = block(x) + return self.ln_f(x) + + def forward( + self, + idx: torch.Tensor, + targets: torch.Tensor | None = None, + ) -> tuple[torch.Tensor, torch.Tensor | None]: + logits = self.lm_head(self.hidden_states(idx)) + loss = None + if targets is not None: + loss = F.cross_entropy(logits.reshape(-1, logits.size(-1)), targets.reshape(-1)) + return logits, loss + + def next_token_logits(self, idx: torch.Tensor) -> torch.Tensor: + # Apply the expensive vocabulary projection only to the final position. + hidden = self.hidden_states(idx)[:, -1:, :] + return self.lm_head(hidden) + + @torch.inference_mode() + def generate_greedy(self, prompts: torch.Tensor, max_new_tokens: int) -> torch.Tensor: + if prompts.ndim != 2: + raise ValueError("prompts must be [batch, sequence]") + if max_new_tokens < 0: + raise ValueError("max_new_tokens must be nonnegative") + idx = prompts + for _ in range(int(max_new_tokens)): + idx_cond = idx[:, -self.cfg.block_size :] + logits = self.next_token_logits(idx_cond) + next_token = logits[:, -1, :].argmax(dim=-1, keepdim=True) + idx = torch.cat((idx, next_token), dim=1) + return idx + + def parameter_count(self) -> int: + return sum(parameter.numel() for parameter in self.parameters()) + + +def transformer_matrix_items( + model: GPT, +) -> list[tuple[str, str, int, torch.Tensor]]: + """Return the six transformer matrices used by WeightWatcher and Muon.""" + items: list[tuple[str, str, int, torch.Tensor]] = [] + for block_index, block in enumerate(model.blocks): + matrices = ( + ("W_Q", block.attn.q_proj.weight), + ("W_K", block.attn.k_proj.weight), + ("W_V", block.attn.v_proj.weight), + ("W_O", block.attn.out_proj.weight), + ("W_MLP_IN", block.mlp.fc.weight), + ("W_MLP_OUT", block.mlp.proj.weight), + ) + for matrix_type, weight in matrices: + items.append((f"L{block_index:02d}_{matrix_type}", matrix_type, block_index, weight)) + return items diff --git a/baseline/nanogpt_muon_hyperball/src/rg_nanogpt_muon_hyperball/optimizers.py b/baseline/nanogpt_muon_hyperball/src/rg_nanogpt_muon_hyperball/optimizers.py new file mode 100644 index 0000000..20b3fd3 --- /dev/null +++ b/baseline/nanogpt_muon_hyperball/src/rg_nanogpt_muon_hyperball/optimizers.py @@ -0,0 +1,289 @@ +from __future__ import annotations + +import math +from typing import Iterable + +import torch + +from .base_optimizers import ( + Muon, + OptimizerHandle, + cosine_learning_rate, + load_optimizer_state_dict, + optimizer_state_dict, + set_learning_rates, + zero_grad, + zeropower_via_newton_schulz_5, +) + + +class MuonHyperBall(torch.optim.Optimizer): + """Muon followed by a relative Frobenius trust-region projection. + + Ordinary Muon, including its multiplicative matrix weight decay, first + proposes a complete displacement ``delta``. HyperBall applies + + delta * min(1, rho * ||W||_F / (||delta||_F + eps)). + + The map is radial, so it changes only displacement magnitude. + """ + + def __init__( + self, + params: Iterable[torch.nn.Parameter], + *, + lr: float, + momentum: float = 0.95, + nesterov: bool = True, + weight_decay: float = 0.01, + newton_schulz_steps: int = 5, + eps: float = 1e-7, + relative_radius: float = 0.01, + hyperball_eps: float = 1e-12, + ) -> None: + params = list(params) + if not params: + raise ValueError("Muon-HyperBall requires at least one parameter") + if any(parameter.ndim != 2 for parameter in params): + raise ValueError("Muon-HyperBall accepts only 2-D parameters") + if relative_radius <= 0: + raise ValueError("relative_radius must be positive") + if hyperball_eps <= 0: + raise ValueError("hyperball_eps must be positive") + defaults = { + "lr": float(lr), + "momentum": float(momentum), + "nesterov": bool(nesterov), + "weight_decay": float(weight_decay), + "newton_schulz_steps": int(newton_schulz_steps), + "eps": float(eps), + "relative_radius": float(relative_radius), + "hyperball_eps": float(hyperball_eps), + } + super().__init__(params, defaults) + self._last_hyperball_summary: dict[str, torch.Tensor] | None = None + + @torch.no_grad() + def step(self, closure=None): + loss = None + if closure is not None: + with torch.enable_grad(): + loss = closure() + + scales: list[torch.Tensor] = [] + active: list[torch.Tensor] = [] + radii: list[torch.Tensor] = [] + proposed_uwr: list[torch.Tensor] = [] + applied_uwr: list[torch.Tensor] = [] + proposed_norms: list[torch.Tensor] = [] + applied_norms: list[torch.Tensor] = [] + + for group in self.param_groups: + lr = float(group["lr"]) + momentum_value = float(group["momentum"]) + relative_radius = float(group["relative_radius"]) + hyperball_eps = float(group["hyperball_eps"]) + + for parameter in group["params"]: + if parameter.grad is None: + continue + gradient = parameter.grad.detach() + if gradient.is_sparse: + raise RuntimeError( + "Muon-HyperBall does not support sparse gradients" + ) + + state = self.state[parameter] + buffer = state.get("momentum_buffer") + if buffer is None: + buffer = torch.zeros_like(gradient) + state["momentum_buffer"] = buffer + buffer.lerp_(gradient, 1.0 - momentum_value) + update_source = ( + gradient.lerp(buffer, momentum_value) + if bool(group["nesterov"]) + else buffer + ) + update = zeropower_via_newton_schulz_5( + update_source, + steps=int(group["newton_schulz_steps"]), + eps=float(group["eps"]), + ) + update.mul_( + math.sqrt(max(1.0, parameter.shape[0] / parameter.shape[1])) + ) + + # Complete reference-Muon displacement, including matrix decay. + delta = update.mul(-lr) + decay = float(group["weight_decay"]) + decay_factor = max(0.0, 1.0 - lr * decay) if decay else 1.0 + if decay_factor != 1.0: + delta.add_(parameter, alpha=decay_factor - 1.0) + + weight_norm = torch.linalg.vector_norm(parameter) + proposed_norm = torch.linalg.vector_norm(delta) + eps_tensor = proposed_norm.new_tensor(hyperball_eps) + radius = weight_norm * relative_radius + raw_scale = radius / (proposed_norm + eps_tensor) + scale = torch.where( + proposed_norm == 0, + proposed_norm.new_ones(()), + raw_scale.clamp(max=1.0), + ) + + applied_norm = proposed_norm * scale + denominator = weight_norm.clamp_min(eps_tensor) + proposed_ratio = proposed_norm / denominator + applied_ratio = applied_norm / denominator + + delta.mul_(scale) + parameter.add_(delta) + + scales.append(scale.detach()) + active.append((scale < (1.0 - 1e-7)).to(scale.dtype).detach()) + radii.append(radius.detach()) + proposed_uwr.append(proposed_ratio.detach()) + applied_uwr.append(applied_ratio.detach()) + proposed_norms.append(proposed_norm.detach()) + applied_norms.append(applied_norm.detach()) + + if scales: + scale_tensor = torch.stack(scales) + active_tensor = torch.stack(active) + radius_tensor = torch.stack(radii) + proposed_uwr_tensor = torch.stack(proposed_uwr) + applied_uwr_tensor = torch.stack(applied_uwr) + proposed_norm_tensor = torch.stack(proposed_norms) + applied_norm_tensor = torch.stack(applied_norms) + self._last_hyperball_summary = { + "matrix_updates": scale_tensor.new_tensor( + float(scale_tensor.numel()) + ), + "active_updates": active_tensor.sum(), + "scale_sum": scale_tensor.sum(), + "scale_min": scale_tensor.min(), + "radius_sum": radius_tensor.sum(), + "proposed_uwr_max": proposed_uwr_tensor.max(), + "applied_uwr_max": applied_uwr_tensor.max(), + "proposed_update_norm_max": proposed_norm_tensor.max(), + "applied_update_norm_max": applied_norm_tensor.max(), + } + else: + self._last_hyperball_summary = None + return loss + + def last_hyperball_summary(self) -> dict[str, torch.Tensor] | None: + if self._last_hyperball_summary is None: + return None + return { + key: value.detach().clone() + for key, value in self._last_hyperball_summary.items() + } + + +def _named_parameters(model) -> list[tuple[str, torch.nn.Parameter]]: + return [ + (name, parameter) + for name, parameter in model.named_parameters() + if parameter.requires_grad + ] + + +def _decay_groups( + named_parameters: list[tuple[str, torch.nn.Parameter]], + weight_decay: float, +) -> list[dict]: + decay = [parameter for _, parameter in named_parameters if parameter.ndim >= 2] + no_decay = [parameter for _, parameter in named_parameters if parameter.ndim < 2] + return [ + {"params": decay, "weight_decay": float(weight_decay)}, + {"params": no_decay, "weight_decay": 0.0}, + ] + + +def make_optimizer_handles(model, profile: dict) -> list[OptimizerHandle]: + named = _named_parameters(model) + family = str(profile["family"]) + if family not in {"muon", "muon_hyperball"}: + raise ValueError(f"unsupported optimizer family: {family}") + + hidden = [ + parameter + for name, parameter in named + if name.startswith("blocks.") and parameter.ndim == 2 + ] + hidden_ids = {id(parameter) for parameter in hidden} + auxiliary_named = [ + (name, parameter) + for name, parameter in named + if id(parameter) not in hidden_ids + ] + if not hidden or not auxiliary_named: + raise ValueError( + "Muon partition must contain hidden matrices and auxiliary parameters" + ) + + common = dict( + lr=float(profile["matrix_learning_rate"]), + momentum=float(profile["momentum"]), + nesterov=bool(profile["nesterov"]), + weight_decay=float(profile["matrix_weight_decay"]), + newton_schulz_steps=int(profile["newton_schulz_steps"]), + eps=float(profile.get("muon_epsilon", 1e-7)), + ) + if family == "muon": + primary = Muon(hidden, **common) + else: + primary = MuonHyperBall( + hidden, + **common, + relative_radius=float(profile["hyperball_relative_radius"]), + hyperball_eps=float(profile["hyperball_epsilon"]), + ) + + auxiliary = torch.optim.AdamW( + _decay_groups(auxiliary_named, float(profile["aux_weight_decay"])), + lr=float(profile["aux_learning_rate"]), + betas=(float(profile["beta1"]), float(profile["beta2"])), + eps=float(profile["epsilon"]), + ) + return [ + OptimizerHandle( + role="primary", + optimizer=primary, + peak_lr=float(profile["matrix_learning_rate"]), + min_lr=float(profile["matrix_min_learning_rate"]), + ), + OptimizerHandle( + role="auxiliary", + optimizer=auxiliary, + peak_lr=float(profile["aux_learning_rate"]), + min_lr=float(profile["aux_min_learning_rate"]), + ), + ] + + +def optimizer_step( + handles: list[OptimizerHandle], +) -> dict[str, torch.Tensor] | None: + summary: dict[str, torch.Tensor] | None = None + for handle in handles: + handle.optimizer.step() + if isinstance(handle.optimizer, MuonHyperBall): + summary = handle.optimizer.last_hyperball_summary() + return summary + + +__all__ = [ + "Muon", + "MuonHyperBall", + "OptimizerHandle", + "cosine_learning_rate", + "load_optimizer_state_dict", + "make_optimizer_handles", + "optimizer_state_dict", + "optimizer_step", + "set_learning_rates", + "zero_grad", + "zeropower_via_newton_schulz_5", +] diff --git a/baseline/nanogpt_muon_hyperball/src/rg_nanogpt_muon_hyperball/run_utils.py b/baseline/nanogpt_muon_hyperball/src/rg_nanogpt_muon_hyperball/run_utils.py new file mode 100644 index 0000000..69198ed --- /dev/null +++ b/baseline/nanogpt_muon_hyperball/src/rg_nanogpt_muon_hyperball/run_utils.py @@ -0,0 +1,179 @@ +from __future__ import annotations + +import csv +import json +from pathlib import Path + +import pandas as pd +import torch + +from .completion import CompletedRunValidationError, validate_completed_run +from .config import lr_schedule_steps, optimizer_profile, tokens_per_step +from .evaluation import evaluate_bleu, evaluate_probe +from .model import GPT + +METRIC_FIELDS = [ + "step", "tokens_seen", "epoch", "elapsed_sec", "tokens_per_sec", + "primary_lr", "auxiliary_lr", "train_loss", "train_perplexity", + "train_accuracy", "val_loss", "val_perplexity", "val_accuracy", + "test_loss", "test_perplexity", "test_accuracy", "test_bleu", + "val_generalization_gap", "test_generalization_gap", + "grad_norm_pre_clip", "grad_norm_post_clip", "gradient_clipped", + "weight_norm", "update_norm_since_eval", "update_to_weight_ratio", + "hyperball_relative_radius", + "hyperball_matrix_updates_since_eval", + "hyperball_active_fraction", "hyperball_mean_scale", + "hyperball_min_scale", "hyperball_mean_radius", + "hyperball_max_proposed_update_to_weight_ratio", + "hyperball_max_applied_update_to_weight_ratio", + "hyperball_max_proposed_update_norm", + "hyperball_max_applied_update_norm", + "mps_current_allocated_mb", "mps_driver_allocated_mb", +] +EPOCH_FIELDS = [ + *METRIC_FIELDS, "nominal_epoch", "checkpoint_path", "test_monitoring_only" +] + + +def run_directory(results_root: str | Path, optimizer: str, seed: int) -> Path: + return Path(results_root) / str(optimizer) / f"seed_{int(seed)}" + + +def run_is_complete(results_root: str | Path, optimizer: str, seed: int) -> bool: + run_dir = run_directory(results_root, optimizer, seed) + try: + validate_completed_run( + run_dir, + expected_optimizer=str(optimizer), + expected_seed=int(seed), + verify_checkpoints=False, + ) + except (CompletedRunValidationError, OSError): + return False + return True + + +def prepare_csv(path: Path, fields: list[str], resume_step: int | None) -> None: + path.parent.mkdir(parents=True, exist_ok=True) + if path.is_file() and resume_step is not None: + frame = pd.read_csv(path) + if "step" in frame.columns: + frame = frame[ + pd.to_numeric(frame["step"], errors="coerce") < int(resume_step) + ] + temporary = path.with_suffix(path.suffix + ".tmp") + frame.to_csv(temporary, index=False) + temporary.replace(path) + if not path.is_file() or path.stat().st_size == 0: + with path.open("w", newline="", encoding="utf-8") as handle: + csv.DictWriter(handle, fieldnames=fields).writeheader() + + +def truncate_spectral_after(run_dir: Path, resume_step: int) -> None: + spectral_root = run_dir / "spectral" + for filename in ("layers.csv", "summary.csv"): + path = spectral_root / filename + if path.is_file(): + frame = pd.read_csv(path) + if "step" in frame.columns: + frame = frame[ + pd.to_numeric(frame["step"], errors="coerce") + < int(resume_step) + ] + temporary = path.with_suffix(path.suffix + ".tmp") + frame.to_csv(temporary, index=False) + temporary.replace(path) + + raw_root = spectral_root / "raw" + if raw_root.is_dir(): + for path in raw_root.glob("weightwatcher_step_*.csv"): + try: + step = int(path.stem.rsplit("_", 1)[-1]) + except ValueError: + continue + if step >= int(resume_step): + path.unlink(missing_ok=True) + for path in spectral_root.glob("status_step_*.json"): + try: + step = int(path.stem.rsplit("_", 1)[-1]) + except ValueError: + continue + if step >= int(resume_step): + path.unlink(missing_ok=True) + + +def write_manifest( + run_dir: Path, + *, + cfg: dict, + data_metadata: dict, + optimizer_name: str, + profile: dict, + seed: int, + device: torch.device, + total_steps: int, + warmup: int, + fingerprint: str, + model: GPT, +) -> None: + payload = { + "schema_version": 1, + "protocol": cfg["protocol"], + "optimizer": optimizer_name, + "optimizer_profile": profile, + "seed": int(seed), + "device": str(device), + "torch_version": torch.__version__, + "model": cfg["model"], + "parameter_count": model.parameter_count(), + "data_metadata": data_metadata, + "training": cfg["training"], + "evaluation": cfg["evaluation"], + "weightwatcher": cfg["weightwatcher"], + "tokens_per_step": tokens_per_step(cfg), + "max_steps": int(total_steps), + "lr_schedule_steps": int( + lr_schedule_steps(cfg, optimizer_profile(cfg, optimizer_name)) + ), + "warmup_steps": int(warmup), + "planned_training_tokens": int(total_steps * tokens_per_step(cfg)), + "protocol_fingerprint": fingerprint, + "test_policy": ( + "fixed test probes are monitoring-only and never select " + "checkpoints or tune schedules" + ), + "hyperball_policy": ( + "Muon proposes the complete hidden-matrix displacement; a relative " + "Frobenius ball changes only displacement magnitude" + ), + } + temporary = run_dir / "manifest.json.tmp" + temporary.write_text( + json.dumps(payload, indent=2, sort_keys=True, default=str), + encoding="utf-8", + ) + temporary.replace(run_dir / "manifest.json") + + +def checkpoint_eval( + checkpoint: Path, + *, + model: GPT, + test_probe, + bleu_probe, + device: torch.device, + bleu_batch_size: int, +) -> dict[str, float]: + payload = torch.load(checkpoint, map_location="cpu", weights_only=False) + model.load_state_dict(payload["model"]) + metrics = evaluate_probe(model, test_probe, device) + bleu = evaluate_bleu( + model, bleu_probe, device=device, batch_size=bleu_batch_size + ) + return { + "step": int(payload["step"]), + "loss": float(metrics["loss"]), + "perplexity": float(metrics["perplexity"]), + "accuracy": float(metrics["accuracy"]), + "bleu": float(bleu["bleu"]), + } diff --git a/baseline/nanogpt_muon_hyperball/src/rg_nanogpt_muon_hyperball/runtime.py b/baseline/nanogpt_muon_hyperball/src/rg_nanogpt_muon_hyperball/runtime.py new file mode 100644 index 0000000..0ad56b1 --- /dev/null +++ b/baseline/nanogpt_muon_hyperball/src/rg_nanogpt_muon_hyperball/runtime.py @@ -0,0 +1,101 @@ +from __future__ import annotations + +import math +import os +import random +from typing import Iterable + +import numpy as np +import torch +import torch.nn as nn + + +def choose_device(requested: str = "auto") -> torch.device: + requested = str(requested).lower() + if requested != "auto": + device = torch.device(requested) + if device.type == "mps" and not torch.backends.mps.is_available(): + raise RuntimeError("MPS was requested but is not available in this PyTorch build/runtime") + return 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 configure_runtime(device: torch.device, cfg: dict) -> None: + torch.set_float32_matmul_precision(str(cfg["runtime"].get("matmul_precision", "high"))) + if device.type == "mps" and bool(cfg["runtime"].get("mps_fallback", True)): + os.environ.setdefault("PYTORCH_ENABLE_MPS_FALLBACK", "1") + if bool(cfg["runtime"].get("deterministic_algorithms", False)): + torch.use_deterministic_algorithms(True, warn_only=True) + + +def synchronize(device: torch.device) -> None: + if device.type == "cuda": + torch.cuda.synchronize(device) + elif device.type == "mps" and hasattr(torch, "mps"): + torch.mps.synchronize() + + +def seed_everything(seed: int) -> None: + random.seed(int(seed)) + np.random.seed(int(seed) % (2**32 - 1)) + torch.manual_seed(int(seed)) + if torch.cuda.is_available(): + torch.cuda.manual_seed_all(int(seed)) + if torch.backends.mps.is_available() and hasattr(torch, "mps") and hasattr(torch.mps, "manual_seed"): + torch.mps.manual_seed(int(seed)) + + +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 + ] + return torch.linalg.vector_norm(torch.stack(norms), ord=2) if norms else torch.tensor(0.0) + + +def model_weight_norm(model: nn.Module) -> float: + squared = 0.0 + for parameter in model.parameters(): + squared += float((parameter.detach().float() ** 2).sum().cpu()) + return math.sqrt(squared) + + +def parameter_snapshot(model: nn.Module) -> list[torch.Tensor]: + return [parameter.detach().float().cpu().clone() for parameter in model.parameters()] + + +def update_norm(previous: list[torch.Tensor] | None, current: list[torch.Tensor]) -> float: + if previous is None: + return 0.0 + if len(previous) != len(current): + raise RuntimeError("parameter inventory changed during training") + squared = 0.0 + for old, new in zip(previous, current, strict=True): + squared += float(((new - old) ** 2).sum()) + return math.sqrt(squared) + + +def mps_memory_megabytes(device: torch.device) -> tuple[float, float]: + if device.type != "mps" or not hasattr(torch, "mps"): + return float("nan"), float("nan") + current = ( + float(torch.mps.current_allocated_memory()) / (1024**2) + if hasattr(torch.mps, "current_allocated_memory") + else float("nan") + ) + driver = ( + float(torch.mps.driver_allocated_memory()) / (1024**2) + if hasattr(torch.mps, "driver_allocated_memory") + else float("nan") + ) + return current, driver + + +def empty_mps_cache(device: torch.device) -> None: + if device.type == "mps" and hasattr(torch, "mps") and hasattr(torch.mps, "empty_cache"): + torch.mps.empty_cache() diff --git a/baseline/nanogpt_muon_hyperball/src/rg_nanogpt_muon_hyperball/spectral.py b/baseline/nanogpt_muon_hyperball/src/rg_nanogpt_muon_hyperball/spectral.py new file mode 100644 index 0000000..7417468 --- /dev/null +++ b/baseline/nanogpt_muon_hyperball/src/rg_nanogpt_muon_hyperball/spectral.py @@ -0,0 +1,317 @@ +from __future__ import annotations + +import json +from pathlib import Path +import random +from typing import Any + +import numpy as np +import pandas as pd +import torch +import torch.nn as nn + +from .model import GPT, transformer_matrix_items + +SPECTRAL_METRICS = ( + "alpha", + "alpha_weighted", + "ERG_gap", + "num_traps", + "detX_num", + "num_pl_spikes", + "num_ERG_spikes", + "D", + "stable_rank", + "mp_softrank", + "log_norm", + "log_spectral_norm", + "entropy", + "Lambda", + "rank_loss", +) + + +class WeightMatrixHolder(nn.Module): + """CPU-only Linear view of the six one-block transformer matrices.""" + + def __init__(self, model: GPT) -> None: + super().__init__() + self.matrix_metadata: list[dict[str, object]] = [] + for name, matrix_type, block, weight in transformer_matrix_items(model): + layer = nn.Linear(weight.shape[1], weight.shape[0], bias=False) + layer.weight = nn.Parameter( + weight.detach().float().cpu().clone(), + requires_grad=False, + ) + self.add_module(name, layer) + self.matrix_metadata.append( + { + "matrix_name": name, + "matrix_type": matrix_type, + "block": int(block), + } + ) + + +def _attach_matrix_metadata( + frame: pd.DataFrame, + metadata: list[dict[str, object]], +) -> pd.DataFrame: + result = frame.copy().reset_index(drop=True) + names = [str(item["matrix_name"]) for item in metadata] + resolved: list[str | None] = [None] * len(result) + for row_index, row in result.iterrows(): + text = " ".join( + str(row.get(column, "")) for column in ("longname", "name") + ) + for name in names: + if name in text: + resolved[row_index] = name + break + if any(value is None for value in resolved) and len(result) == len(metadata): + order = list(range(len(result))) + if "layer_id" in result.columns: + numeric = pd.to_numeric(result["layer_id"], errors="coerce") + if numeric.notna().all(): + order = list(numeric.sort_values().index) + for metadata_index, row_index in enumerate(order): + resolved[row_index] = names[metadata_index] + if any(value is None for value in resolved): + raise RuntimeError( + "WeightWatcher rows could not be matched to all transformer matrices" + ) + by_name = {str(item["matrix_name"]): item for item in metadata} + result.insert(0, "matrix_name", resolved) + result.insert( + 1, + "matrix_type", + [by_name[str(name)]["matrix_type"] for name in resolved], + ) + result.insert( + 2, + "block", + [by_name[str(name)]["block"] for name in resolved], + ) + return result + + +def _atomic_csv(path: Path, frame: pd.DataFrame) -> None: + path.parent.mkdir(parents=True, exist_ok=True) + temporary = path.with_suffix(path.suffix + ".tmp") + frame.to_csv(temporary, index=False) + temporary.replace(path) + + +def _append_deduplicated( + path: Path, + frame: pd.DataFrame, + keys: list[str], +) -> None: + if path.is_file(): + existing = pd.read_csv(path) + combined = pd.concat([existing, frame], ignore_index=True, sort=False) + else: + combined = frame.copy() + combined = combined.drop_duplicates(keys, keep="last").sort_values(keys) + _atomic_csv(path, combined) + + +def _capture_accelerator_rng() -> dict[str, Any]: + state: dict[str, Any] = {} + if torch.cuda.is_available(): + state["cuda"] = torch.cuda.get_rng_state_all() + if ( + hasattr(torch, "mps") + and hasattr(torch.mps, "get_rng_state") + and torch.backends.mps.is_available() + ): + state["mps"] = torch.mps.get_rng_state() + return state + + +def _restore_accelerator_rng(state: dict[str, Any]) -> None: + if torch.cuda.is_available() and "cuda" in state: + torch.cuda.set_rng_state_all(state["cuda"]) + if ( + "mps" in state + and hasattr(torch, "mps") + and hasattr(torch.mps, "set_rng_state") + and torch.backends.mps.is_available() + ): + torch.mps.set_rng_state(state["mps"]) + + +def summarize_spectral_frame( + frame: pd.DataFrame, + *, + step: int, + tokens_seen: int, + epoch: float, +) -> dict[str, Any]: + summary: dict[str, Any] = { + "step": int(step), + "tokens_seen": int(tokens_seen), + "epoch": float(epoch), + "n_matrices": int(len(frame)), + } + for metric in SPECTRAL_METRICS: + values = ( + pd.to_numeric(frame[metric], errors="coerce") + if metric in frame.columns + else pd.Series(dtype=float) + ) + array = values.to_numpy(dtype=float, na_value=np.nan) + finite = array[np.isfinite(array)] + summary[f"{metric}_n"] = int(finite.size) + for statistic in ("mean", "median", "std", "min", "max"): + summary[f"{metric}_{statistic}"] = float("nan") + if finite.size: + summary[f"{metric}_mean"] = float(np.mean(finite)) + summary[f"{metric}_median"] = float(np.median(finite)) + summary[f"{metric}_std"] = ( + float(np.std(finite, ddof=1)) if finite.size > 1 else 0.0 + ) + summary[f"{metric}_min"] = float(np.min(finite)) + summary[f"{metric}_max"] = float(np.max(finite)) + return summary + + +def run_weightwatcher( + model: GPT, + run_dir: str | Path, + *, + step: int, + tokens_seen: int, + train_tokens: int, + config: dict[str, Any], + seed: int, +) -> dict[str, Any]: + """Run WeightWatcher exactly with ERG=True and randomize=True. + + `alpha`, `ERG_gap`, and `num_traps` are retained directly from WeightWatcher. + No fallback alpha, proxy trap count, or synthesized ERG gap is permitted. + Every CPU and accelerator RNG stream is restored after the randomized + diagnostic so measurement cannot change the subsequent training path. + """ + + try: + import weightwatcher as ww + except ImportError as exc: + raise RuntimeError( + "WeightWatcher is required; run scripts/setup_mac.sh" + ) from exc + + run_dir = Path(run_dir) + spectral_root = run_dir / "spectral" + raw_root = spectral_root / "raw" + raw_root.mkdir(parents=True, exist_ok=True) + raw_path = raw_root / f"weightwatcher_step_{int(step):07d}.csv" + if raw_path.is_file(): + frame = pd.read_csv(raw_path) + return summarize_spectral_frame( + frame, + step=step, + tokens_seen=tokens_seen, + epoch=tokens_seen / max(1, int(train_tokens)), + ) + + py_state = random.getstate() + np_state = np.random.get_state() + torch_state = torch.random.get_rng_state() + accelerator_state = _capture_accelerator_rng() + diagnostic_seed = int(seed) + 1_000_003 + int(step) + random.seed(diagnostic_seed) + np.random.seed(diagnostic_seed % (2**32 - 1)) + torch.manual_seed(diagnostic_seed) + + try: + holder = WeightMatrixHolder(model) + watcher = ww.WeightWatcher(model=holder) + details = watcher.analyze( + ERG=True, + randomize=True, + plot=False, + min_evals=int(config.get("min_evals", 20)), + ) + if details is None or len(details) == 0: + raise RuntimeError( + "WeightWatcher returned no transformer-matrix rows" + ) + frame = _attach_matrix_metadata( + pd.DataFrame(details), holder.matrix_metadata + ) + required_columns = ("alpha", "ERG_gap", "num_traps") + missing = [ + column for column in required_columns if column not in frame.columns + ] + if missing: + raise RuntimeError( + "WeightWatcher did not return required ERG/randomization columns: " + + ", ".join(missing) + ) + if frame[list(required_columns)].isna().any().any(): + raise RuntimeError( + "WeightWatcher required alpha/ERG_gap/num_traps values contain NaN" + ) + epoch = tokens_seen / max(1, int(train_tokens)) + frame.insert(0, "step", int(step)) + frame.insert(1, "tokens_seen", int(tokens_seen)) + frame.insert(2, "epoch", float(epoch)) + frame.insert(3, "diagnostic_seed", int(diagnostic_seed)) + _atomic_csv(raw_path, frame) + _append_deduplicated( + spectral_root / "layers.csv", + frame, + keys=["step", "matrix_name"], + ) + summary = summarize_spectral_frame( + frame, + step=step, + tokens_seen=tokens_seen, + epoch=epoch, + ) + _append_deduplicated( + spectral_root / "summary.csv", + pd.DataFrame([summary]), + keys=["step"], + ) + status = { + "step": int(step), + "tokens_seen": int(tokens_seen), + "epoch": float(epoch), + "completed": True, + "raw_path": str(raw_path), + "alpha_valid_matrices": int(summary["alpha_n"]), + "ERG_gap_valid_matrices": int(summary["ERG_gap_n"]), + "num_traps_valid_matrices": int(summary["num_traps_n"]), + } + (spectral_root / f"status_step_{int(step):07d}.json").write_text( + json.dumps(status, indent=2, sort_keys=True), + encoding="utf-8", + ) + return summary + except Exception as exc: + status = { + "step": int(step), + "tokens_seen": int(tokens_seen), + "completed": False, + "error_type": type(exc).__name__, + "error": str(exc), + } + (spectral_root / f"status_step_{int(step):07d}.json").write_text( + json.dumps(status, indent=2, sort_keys=True), + encoding="utf-8", + ) + if bool(config.get("strict", True)): + raise + print( + f"[one-head-ww] WARNING step={step}: " + f"{type(exc).__name__}: {exc}", + flush=True, + ) + return status + finally: + random.setstate(py_state) + np.random.set_state(np_state) + torch.random.set_rng_state(torch_state) + _restore_accelerator_rng(accelerator_state) diff --git a/baseline/nanogpt_muon_hyperball/src/rg_nanogpt_muon_hyperball/train_loop.py b/baseline/nanogpt_muon_hyperball/src/rg_nanogpt_muon_hyperball/train_loop.py new file mode 100644 index 0000000..3298288 --- /dev/null +++ b/baseline/nanogpt_muon_hyperball/src/rg_nanogpt_muon_hyperball/train_loop.py @@ -0,0 +1,446 @@ +from __future__ import annotations + +import csv +import math +from pathlib import Path +import time + +import torch + +from .checkpoints import save_epoch_model_checkpoint, save_training_checkpoint +from .config import lr_schedule_steps, optimizer_profile +from .evaluation import evaluate_bleu, evaluate_probe, random_batch +from .optimizers import optimizer_step, set_learning_rates, zero_grad +from .runtime import ( + empty_mps_cache, + gradient_norm, + model_weight_norm, + mps_memory_megabytes, + parameter_snapshot, + synchronize, + update_norm, +) +from .spectral import run_weightwatcher + + +def _accumulate_hyperball( + accumulator: dict[str, torch.Tensor] | None, + step_summary: dict[str, torch.Tensor] | None, +) -> dict[str, torch.Tensor] | None: + if step_summary is None: + return accumulator + if accumulator is None: + return {key: value.detach().clone() for key, value in step_summary.items()} + + for key in ("matrix_updates", "active_updates", "scale_sum", "radius_sum"): + accumulator[key] = accumulator[key] + step_summary[key] + accumulator["scale_min"] = torch.minimum( + accumulator["scale_min"], step_summary["scale_min"] + ) + for key in ( + "proposed_uwr_max", + "applied_uwr_max", + "proposed_update_norm_max", + "applied_update_norm_max", + ): + accumulator[key] = torch.maximum(accumulator[key], step_summary[key]) + return accumulator + + +def _scalar(value: torch.Tensor) -> float: + return float(value.detach().cpu()) + + +def _hyperball_metrics( + accumulator: dict[str, torch.Tensor] | None, + handles, +) -> dict[str, float]: + radius = float("nan") + for handle in handles: + if handle.role == "primary" and handle.optimizer.param_groups: + radius = float( + handle.optimizer.param_groups[0].get( + "relative_radius", float("nan") + ) + ) + break + + if accumulator is None: + return { + "hyperball_relative_radius": radius, + "hyperball_matrix_updates_since_eval": 0.0, + "hyperball_active_fraction": float("nan"), + "hyperball_mean_scale": float("nan"), + "hyperball_min_scale": float("nan"), + "hyperball_mean_radius": float("nan"), + "hyperball_max_proposed_update_to_weight_ratio": float("nan"), + "hyperball_max_applied_update_to_weight_ratio": float("nan"), + "hyperball_max_proposed_update_norm": float("nan"), + "hyperball_max_applied_update_norm": float("nan"), + } + + count = max(_scalar(accumulator["matrix_updates"]), 1.0) + return { + "hyperball_relative_radius": radius, + "hyperball_matrix_updates_since_eval": count, + "hyperball_active_fraction": _scalar(accumulator["active_updates"]) / count, + "hyperball_mean_scale": _scalar(accumulator["scale_sum"]) / count, + "hyperball_min_scale": _scalar(accumulator["scale_min"]), + "hyperball_mean_radius": _scalar(accumulator["radius_sum"]) / count, + "hyperball_max_proposed_update_to_weight_ratio": _scalar( + accumulator["proposed_uwr_max"] + ), + "hyperball_max_applied_update_to_weight_ratio": _scalar( + accumulator["applied_uwr_max"] + ), + "hyperball_max_proposed_update_norm": _scalar( + accumulator["proposed_update_norm_max"] + ), + "hyperball_max_applied_update_norm": _scalar( + accumulator["applied_update_norm_max"] + ), + } + + +def _require_finite_metrics( + *, + completed_steps: int, + train_metrics: dict, + val_metrics: dict, +) -> None: + values = { + "train_loss": float(train_metrics["loss"]), + "train_perplexity": float(train_metrics["perplexity"]), + "train_accuracy": float(train_metrics["accuracy"]), + "val_loss": float(val_metrics["loss"]), + "val_perplexity": float(val_metrics["perplexity"]), + "val_accuracy": float(val_metrics["accuracy"]), + } + bad = [name for name, value in values.items() if not math.isfinite(value)] + if bad: + raise FloatingPointError( + "non-finite training state at " + f"step={completed_steps}: {', '.join(bad)}. " + "Aborting before checkpoint selection or WeightWatcher." + ) + + +def execute_training_loop( + *, + cfg: dict, + model, + handles, + arrays: dict, + train_probe, + val_probe, + test_probe, + bleu_probe, + device: torch.device, + optimizer_name: str, + seed: int, + train_tokens: int, + total_steps: int, + warmup: int, + start_step: int, + best_validation_loss: float, + best_validation_step: int, + elapsed_offset: float, + fingerprint: str, + train_generator: torch.Generator, + epoch_steps: dict[int, float], + metrics_writer: csv.DictWriter, + metrics_handle, + epoch_writer: csv.DictWriter, + epoch_handle, + run_dir: Path, + latest_checkpoint: Path, + best_checkpoint: Path, + progress: bool, +) -> tuple[float, int, float]: + batch_size = int(cfg["training"]["batch_size"]) + grad_accum = int(cfg["training"]["grad_accum_steps"]) + block_size = int(cfg["model"]["block_size"]) + step_tokens = batch_size * grad_accum * block_size + eval_cfg = cfg["evaluation"] + + profile = optimizer_profile(cfg, optimizer_name) + schedule_total_steps = lr_schedule_steps(cfg, profile, train_tokens) + if not 0 <= warmup < schedule_total_steps: + raise ValueError( + f"warmup={warmup} must be smaller than " + f"lr_schedule_steps={schedule_total_steps}" + ) + + previous_snapshot = parameter_snapshot(model) + last_grad_pre = float("nan") + last_grad_post = float("nan") + last_clipped = False + hyperball_interval: dict[str, torch.Tensor] | None = None + started = time.time() + + last_update_lrs = { + "primary": 0.0, + "auxiliary": ( + 0.0 + if any(handle.role == "auxiliary" for handle in handles) + else float("nan") + ), + } + if start_step > 0: + for handle in handles: + last_update_lrs[handle.role] = float(handle.lr) + + for completed_steps in range(start_step, total_steps + 1): + schedule_index = min(completed_steps, schedule_total_steps - 1) + next_update_lrs = set_learning_rates( + handles, + update_index=schedule_index, + total_steps=schedule_total_steps, + warmup_steps=warmup, + ) + epoch_due = completed_steps in epoch_steps + evaluation_due = ( + completed_steps + % int(cfg["training"]["eval_interval_steps"]) + == 0 + or epoch_due + or completed_steps == total_steps + ) + + if evaluation_due: + synchronize(device) + train_metrics = evaluate_probe(model, train_probe, device) + val_metrics = evaluate_probe(model, val_probe, device) + _require_finite_metrics( + completed_steps=completed_steps, + train_metrics=train_metrics, + val_metrics=val_metrics, + ) + elapsed = elapsed_offset + time.time() - started + + if val_metrics["loss"] < best_validation_loss: + best_validation_loss = float(val_metrics["loss"]) + best_validation_step = int(completed_steps) + save_training_checkpoint( + best_checkpoint, + model=model, + handles=handles, + step=completed_steps, + best_validation_loss=best_validation_loss, + best_validation_step=best_validation_step, + elapsed_seconds=elapsed, + fingerprint=fingerprint, + cfg=cfg, + optimizer_name=optimizer_name, + seed=int(seed), + train_generator=train_generator, + ) + + test_metrics = { + "loss": float("nan"), + "perplexity": float("nan"), + "accuracy": float("nan"), + } + bleu_metrics = {"bleu": float("nan")} + if epoch_due or completed_steps == total_steps: + test_metrics = evaluate_probe(model, test_probe, device) + bleu_metrics = evaluate_bleu( + model, + bleu_probe, + device=device, + batch_size=int(eval_cfg["bleu_batch_size"]), + ) + + tokens_seen = int(completed_steps * step_tokens) + actual_epoch = tokens_seen / max(1, train_tokens) + current_snapshot = parameter_snapshot(model) + delta_norm = update_norm(previous_snapshot, current_snapshot) + previous_snapshot = current_snapshot + weight_norm = model_weight_norm(model) + current_mps, driver_mps = mps_memory_megabytes(device) + row = { + "step": int(completed_steps), + "tokens_seen": tokens_seen, + "epoch": float(actual_epoch), + "elapsed_sec": float(elapsed), + "tokens_per_sec": tokens_seen / max(elapsed, 1e-9), + "primary_lr": float( + last_update_lrs.get("primary", float("nan")) + ), + "auxiliary_lr": float( + last_update_lrs.get("auxiliary", float("nan")) + ), + "train_loss": float(train_metrics["loss"]), + "train_perplexity": float(train_metrics["perplexity"]), + "train_accuracy": float(train_metrics["accuracy"]), + "val_loss": float(val_metrics["loss"]), + "val_perplexity": float(val_metrics["perplexity"]), + "val_accuracy": float(val_metrics["accuracy"]), + "test_loss": float(test_metrics["loss"]), + "test_perplexity": float(test_metrics["perplexity"]), + "test_accuracy": float(test_metrics["accuracy"]), + "test_bleu": float(bleu_metrics["bleu"]), + "val_generalization_gap": float( + val_metrics["loss"] - train_metrics["loss"] + ), + "test_generalization_gap": float( + test_metrics["loss"] - train_metrics["loss"] + ), + "grad_norm_pre_clip": float(last_grad_pre), + "grad_norm_post_clip": float(last_grad_post), + "gradient_clipped": int(last_clipped), + "weight_norm": float(weight_norm), + "update_norm_since_eval": float(delta_norm), + "update_to_weight_ratio": float( + delta_norm / max(weight_norm, 1e-30) + ), + **_hyperball_metrics(hyperball_interval, handles), + "mps_current_allocated_mb": float(current_mps), + "mps_driver_allocated_mb": float(driver_mps), + } + metrics_writer.writerow(row) + metrics_handle.flush() + hyperball_interval = None + + if epoch_due: + nominal_epoch = float(epoch_steps[completed_steps]) + checkpoint_path = save_epoch_model_checkpoint( + run_dir, + model=model, + step=completed_steps, + nominal_epoch=nominal_epoch, + actual_epoch=actual_epoch, + fingerprint=fingerprint, + cfg=cfg, + optimizer_name=optimizer_name, + seed=int(seed), + ) + epoch_writer.writerow( + { + **row, + "nominal_epoch": nominal_epoch, + "checkpoint_path": str(checkpoint_path), + "test_monitoring_only": 1, + } + ) + epoch_handle.flush() + + ww_summary = run_weightwatcher( + model, + run_dir, + step=completed_steps, + tokens_seen=tokens_seen, + train_tokens=train_tokens, + config=cfg["weightwatcher"], + seed=int(seed), + ) + if progress: + print( + "[muon-hyperball-ww] " + f"optimizer={optimizer_name} seed={seed} " + f"epoch={nominal_epoch:.2f} " + f"alpha={ww_summary.get('alpha_median', float('nan')):.3f} " + f"ERG_gap={ww_summary.get('ERG_gap_median', float('nan')):.3f} " + f"num_traps={ww_summary.get('num_traps_mean', float('nan')):.2f}", + flush=True, + ) + if bool( + cfg["runtime"].get( + "empty_mps_cache_after_weightwatcher", True + ) + ): + empty_mps_cache(device) + + if progress: + remaining = total_steps - completed_steps + rate = completed_steps / max(elapsed, 1e-9) + eta = remaining / rate if rate > 0 else float("nan") + eta_text = ( + "unknown" + if not math.isfinite(eta) + else f"{eta / 60:.1f}m" + ) + active = row["hyperball_active_fraction"] + active_text = ( + "n/a" + if not math.isfinite(active) + else f"{100 * active:.1f}%" + ) + print( + "[muon-hyperball-train] " + f"optimizer={optimizer_name} seed={seed} " + f"step={completed_steps}/{total_steps} " + f"epoch={actual_epoch:.3f} " + f"last_lr={last_update_lrs.get('primary', float('nan')):.3e} " + f"next_lr={next_update_lrs.get('primary', float('nan')):.3e} " + f"train_loss={train_metrics['loss']:.4f} " + f"val_loss={val_metrics['loss']:.4f} " + f"val_ppl={val_metrics['perplexity']:.2f} " + f"val_acc={100 * val_metrics['accuracy']:.2f}% " + f"ball_active={active_text} " + f"eta={eta_text}", + flush=True, + ) + + if completed_steps == total_steps: + break + + zero_grad(handles) + for _ in range(grad_accum): + x_cpu, y_cpu = random_batch( + arrays["train"], + batch_size=batch_size, + block_size=block_size, + generator=train_generator, + ) + x = x_cpu.to(device) + y = y_cpu.to(device) + _, loss = model(x, y) + if loss is None: + raise RuntimeError("training forward pass did not return loss") + (loss / grad_accum).backward() + + grad_pre_tensor = gradient_norm(model.parameters()) + last_grad_pre = float(grad_pre_tensor.detach().cpu()) + clip = float(cfg["training"]["grad_clip"]) + if clip > 0: + torch.nn.utils.clip_grad_norm_(model.parameters(), clip) + grad_post_tensor = gradient_norm(model.parameters()) + last_grad_post = float(grad_post_tensor.detach().cpu()) + last_clipped = bool(last_grad_pre > clip) if clip > 0 else False + + step_hyperball = optimizer_step(handles) + hyperball_interval = _accumulate_hyperball( + hyperball_interval, step_hyperball + ) + last_update_lrs = dict(next_update_lrs) + + new_step = completed_steps + 1 + checkpoint_due = ( + new_step + % int(cfg["training"]["checkpoint_interval_steps"]) + == 0 + or new_step in epoch_steps + or new_step == total_steps + ) + if checkpoint_due: + save_training_checkpoint( + latest_checkpoint, + model=model, + handles=handles, + step=new_step, + best_validation_loss=best_validation_loss, + best_validation_step=best_validation_step, + elapsed_seconds=elapsed_offset + time.time() - started, + fingerprint=fingerprint, + cfg=cfg, + optimizer_name=optimizer_name, + seed=int(seed), + train_generator=train_generator, + ) + + return ( + float(best_validation_loss), + int(best_validation_step), + float(elapsed_offset + time.time() - started), + ) diff --git a/baseline/nanogpt_muon_hyperball/src/rg_nanogpt_muon_hyperball/training.py b/baseline/nanogpt_muon_hyperball/src/rg_nanogpt_muon_hyperball/training.py new file mode 100644 index 0000000..fd3cf8c --- /dev/null +++ b/baseline/nanogpt_muon_hyperball/src/rg_nanogpt_muon_hyperball/training.py @@ -0,0 +1,131 @@ +from __future__ import annotations + +import argparse +from copy import deepcopy +from pathlib import Path +from typing import Sequence + +from .config import SUPPORTED_OPTIMIZERS, canonical_seeds, load_config, roots +from .data import prepare_fineweb_edu +from .engine import run_one +from .run_utils import run_directory, run_is_complete + + +def run_optimizer_replicates( + *, + cfg: dict, + config_path: str | Path, + optimizer_name: str, + seeds: Sequence[int] | None = None, + data_root: str | Path | None = None, + results_root: str | Path | None = None, + device: str = "auto", + resume: bool = True, + overwrite: bool = False, + prepare_data: bool = True, + progress: bool = True, +) -> list[Path]: + resolved = roots() + data_root = Path(data_root or resolved["data"]) + results_root = Path(results_root or resolved["results"]) + selected_seeds = tuple(int(seed) for seed in (seeds or canonical_seeds(cfg))) + if prepare_data: + prepare_fineweb_edu(cfg, data_root) + run_dirs = [] + for seed in selected_seeds: + run_dirs.append( + run_one( + cfg=deepcopy(cfg), + data_root=data_root, + results_root=results_root, + optimizer_name=optimizer_name, + seed=seed, + device=device, + resume=resume, + overwrite=overwrite, + progress=progress, + ) + ) + return run_dirs + + +def run_all_replicates( + *, + cfg: dict, + config_path: str | Path, + seeds: Sequence[int] | None = None, + data_root: str | Path | None = None, + results_root: str | Path | None = None, + device: str = "auto", + resume: bool = True, + overwrite: bool = False, + progress: bool = True, +) -> list[Path]: + resolved = roots() + data_root = Path(data_root or resolved["data"]) + results_root = Path(results_root or resolved["results"]) + prepare_fineweb_edu(cfg, data_root) + outputs: list[Path] = [] + for optimizer_name in SUPPORTED_OPTIMIZERS: + outputs.extend( + run_optimizer_replicates( + cfg=cfg, + config_path=config_path, + optimizer_name=optimizer_name, + seeds=seeds, + data_root=data_root, + results_root=results_root, + device=device, + resume=resume, + overwrite=overwrite, + prepare_data=False, + progress=progress, + ) + ) + return outputs + + +def main() -> None: + parser = argparse.ArgumentParser(description="Run the one-head FineWeb-Edu optimizer baselines") + parser.add_argument("--config", required=True) + parser.add_argument("--optimizer", choices=[*SUPPORTED_OPTIMIZERS, "all"], default="all") + parser.add_argument("--seeds", help="comma-separated seeds; default comes from the config") + parser.add_argument("--data-root") + parser.add_argument("--results-root") + parser.add_argument("--device", default="auto") + parser.add_argument("--overwrite", action="store_true") + parser.add_argument("--no-resume", action="store_true") + args = parser.parse_args() + cfg = load_config(args.config) + seeds = ( + tuple(int(value.strip()) for value in args.seeds.split(",") if value.strip()) + if args.seeds + else canonical_seeds(cfg) + ) + if args.optimizer == "all": + run_all_replicates( + cfg=cfg, + config_path=args.config, + seeds=seeds, + data_root=args.data_root, + results_root=args.results_root, + device=args.device, + resume=not args.no_resume, + overwrite=args.overwrite, + ) + else: + run_optimizer_replicates( + cfg=cfg, + config_path=args.config, + optimizer_name=args.optimizer, + seeds=seeds, + data_root=args.data_root, + results_root=args.results_root, + device=args.device, + resume=not args.no_resume, + overwrite=args.overwrite, + ) + + +if __name__ == "__main__": + main() diff --git a/baseline/nanogpt_muon_hyperball/tests/test_config.py b/baseline/nanogpt_muon_hyperball/tests/test_config.py new file mode 100644 index 0000000..bd485be --- /dev/null +++ b/baseline/nanogpt_muon_hyperball/tests/test_config.py @@ -0,0 +1,45 @@ +from copy import deepcopy +from pathlib import Path + +import pytest + +from rg_nanogpt_muon_hyperball.config import ( + load_config, + lr_schedule_steps, + max_steps, + optimizer_profile, + validate_config, + warmup_steps, +) + +CONFIG = Path(__file__).resolve().parents[1] / "configs" / "reference.yaml" + + +def test_reference_config_has_matched_long_muon_arms() -> None: + cfg = load_config(CONFIG) + assert cfg["training"]["seeds"] == [1337] + assert cfg["training"]["target_epochs"] == 10.0 + assert set(cfg["optimizer_profiles"]) == {"muon", "muon_hyperball"} + + total = max_steps(cfg) + assert total == 97657 + + for name in ("muon", "muon_hyperball"): + profile = optimizer_profile(cfg, name) + assert lr_schedule_steps(cfg, profile) == 9766 + assert warmup_steps(profile, total) == 488 + assert profile["matrix_learning_rate"] == 0.02 + assert profile["matrix_min_learning_rate"] == 0.002 + + hb = optimizer_profile(cfg, "muon_hyperball") + assert hb["hyperball_relative_radius"] == 0.01 + + +def test_nonpositive_hyperball_radius_is_rejected() -> None: + cfg = load_config(CONFIG) + broken = deepcopy(cfg) + broken["optimizer_profiles"]["muon_hyperball"][ + "hyperball_relative_radius" + ] = 0.0 + with pytest.raises(ValueError, match="hyperball_relative_radius"): + validate_config(broken) diff --git a/baseline/nanogpt_muon_hyperball/tests/test_hyperball.py b/baseline/nanogpt_muon_hyperball/tests/test_hyperball.py new file mode 100644 index 0000000..30f985e --- /dev/null +++ b/baseline/nanogpt_muon_hyperball/tests/test_hyperball.py @@ -0,0 +1,111 @@ +import torch + +from rg_nanogpt_muon_hyperball.base_optimizers import Muon +from rg_nanogpt_muon_hyperball.optimizers import MuonHyperBall + + +def _clone_parameter(value: torch.Tensor) -> torch.nn.Parameter: + return torch.nn.Parameter(value.detach().clone()) + + +def test_hyperball_caps_relative_frobenius_displacement() -> None: + torch.manual_seed(7) + initial = torch.randn(16, 8) + parameter = _clone_parameter(initial) + parameter.grad = torch.randn_like(parameter) + + rho = 1.0e-3 + optimizer = MuonHyperBall( + [parameter], + lr=0.02, + momentum=0.95, + nesterov=True, + weight_decay=0.01, + relative_radius=rho, + ) + optimizer.step() + + displacement = torch.linalg.vector_norm(parameter.detach() - initial) + radius = rho * torch.linalg.vector_norm(initial) + assert displacement <= radius * (1.0 + 1e-5) + + summary = optimizer.last_hyperball_summary() + assert summary is not None + assert float(summary["active_updates"]) == 1.0 + assert float(summary["applied_uwr_max"]) <= rho * (1.0 + 1e-5) + + +def test_radial_projection_preserves_muon_direction() -> None: + torch.manual_seed(11) + initial = torch.randn(16, 8) + gradient = torch.randn_like(initial) + + plain_parameter = _clone_parameter(initial) + ball_parameter = _clone_parameter(initial) + plain_parameter.grad = gradient.clone() + ball_parameter.grad = gradient.clone() + + plain = Muon( + [plain_parameter], + lr=0.02, + momentum=0.95, + nesterov=True, + weight_decay=0.01, + ) + ball = MuonHyperBall( + [ball_parameter], + lr=0.02, + momentum=0.95, + nesterov=True, + weight_decay=0.01, + relative_radius=1.0e-4, + ) + plain.step() + ball.step() + + proposed = (plain_parameter.detach() - initial).flatten() + applied = (ball_parameter.detach() - initial).flatten() + cosine = torch.dot(proposed, applied) / ( + torch.linalg.vector_norm(proposed) + * torch.linalg.vector_norm(applied) + ) + assert torch.allclose(cosine, torch.tensor(1.0), atol=2e-5, rtol=2e-5) + + +def test_effectively_infinite_radius_matches_plain_muon() -> None: + torch.manual_seed(19) + initial = torch.randn(12, 12) + gradient = torch.randn_like(initial) + + plain_parameter = _clone_parameter(initial) + ball_parameter = _clone_parameter(initial) + plain_parameter.grad = gradient.clone() + ball_parameter.grad = gradient.clone() + + plain = Muon( + [plain_parameter], + lr=0.01, + momentum=0.95, + nesterov=True, + weight_decay=0.02, + ) + ball = MuonHyperBall( + [ball_parameter], + lr=0.01, + momentum=0.95, + nesterov=True, + weight_decay=0.02, + relative_radius=1.0e6, + ) + plain.step() + ball.step() + + assert torch.allclose( + ball_parameter.detach(), + plain_parameter.detach(), + atol=2e-6, + rtol=2e-6, + ) + summary = ball.last_hyperball_summary() + assert summary is not None + assert float(summary["active_updates"]) == 0.0 diff --git a/baseline/nanogpt_muon_hyperball/tests/test_metrics.py b/baseline/nanogpt_muon_hyperball/tests/test_metrics.py new file mode 100644 index 0000000..74ed420 --- /dev/null +++ b/baseline/nanogpt_muon_hyperball/tests/test_metrics.py @@ -0,0 +1,17 @@ +from rg_nanogpt_muon_hyperball.run_utils import METRIC_FIELDS + + +def test_hyperball_diagnostics_are_persisted() -> None: + required = { + "hyperball_relative_radius", + "hyperball_matrix_updates_since_eval", + "hyperball_active_fraction", + "hyperball_mean_scale", + "hyperball_min_scale", + "hyperball_mean_radius", + "hyperball_max_proposed_update_to_weight_ratio", + "hyperball_max_applied_update_to_weight_ratio", + "hyperball_max_proposed_update_norm", + "hyperball_max_applied_update_norm", + } + assert required.issubset(METRIC_FIELDS) diff --git a/baseline/nanogpt_muon_hyperball/tests/test_partition.py b/baseline/nanogpt_muon_hyperball/tests/test_partition.py new file mode 100644 index 0000000..7fb38c3 --- /dev/null +++ b/baseline/nanogpt_muon_hyperball/tests/test_partition.py @@ -0,0 +1,47 @@ +import torch + +from rg_nanogpt_muon_hyperball.config import load_config, optimizer_profile +from rg_nanogpt_muon_hyperball.optimizers import ( + Muon, + MuonHyperBall, + make_optimizer_handles, +) + + +class TinyBlock(torch.nn.Module): + def __init__(self) -> None: + super().__init__() + self.q = torch.nn.Linear(8, 8, bias=False) + self.k = torch.nn.Linear(8, 8, bias=False) + self.v = torch.nn.Linear(8, 8, bias=False) + self.o = torch.nn.Linear(8, 8, bias=False) + self.mlp_in = torch.nn.Linear(8, 32, bias=False) + self.mlp_out = torch.nn.Linear(32, 8, bias=False) + + +class TinyModel(torch.nn.Module): + def __init__(self) -> None: + super().__init__() + self.blocks = torch.nn.ModuleList([TinyBlock()]) + self.embedding = torch.nn.Embedding(32, 8) + self.norm = torch.nn.LayerNorm(8) + + +def test_matched_partition_uses_six_hidden_matrices() -> None: + from pathlib import Path + + cfg = load_config( + Path(__file__).resolve().parents[1] / "configs" / "reference.yaml" + ) + model = TinyModel() + + plain = make_optimizer_handles(model, optimizer_profile(cfg, "muon")) + ball = make_optimizer_handles( + model, optimizer_profile(cfg, "muon_hyperball") + ) + + assert isinstance(plain[0].optimizer, Muon) + assert isinstance(ball[0].optimizer, MuonHyperBall) + assert len(plain[0].optimizer.param_groups[0]["params"]) == 6 + assert len(ball[0].optimizer.param_groups[0]["params"]) == 6 + assert len(plain) == len(ball) == 2