diff --git a/demos/modular_arithmetic_grokking_demo.py b/demos/modular_arithmetic_grokking_demo.py new file mode 100755 index 0000000000..11e18e0e2f --- /dev/null +++ b/demos/modular_arithmetic_grokking_demo.py @@ -0,0 +1,206 @@ +#!/usr/bin/env python3 +"""End-to-end modular addition grokking demo. + +This script prepares a held-out modular-addition dataset, optionally trains a +small nanoGPT model, and evaluates a saved checkpoint by querying every +``a+b=`` prompt and checking whether the generated answer equals ``(a+b) % p``. + +The defaults follow common grokking replications: prime modulus 113, a limited +30% training split, no dropout, AdamW with substantial weight decay, and a small +1-layer transformer trained for many iterations. +""" + +from __future__ import annotations + +import argparse +import json +import os +import pickle +import random +import subprocess +import sys +from pathlib import Path + +ROOT = Path(__file__).resolve().parents[1] +if str(ROOT) not in sys.path: + sys.path.insert(0, str(ROOT)) + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser(description="Prepare, train, and evaluate modular-addition grokking.") + parser.add_argument("--modulo", "--modulus", dest="modulus", type=int, default=113, help="Modulo p for a+b mod p.") + parser.add_argument("--train-fraction", type=float, default=0.3, help="Fraction of ordered pairs used for training.") + parser.add_argument("--train-repeats", type=int, default=200, help="How often to repeat training equations in train.bin.") + parser.add_argument("--val-repeats", type=int, default=20, help="How often to repeat held-out equations in val.bin.") + parser.add_argument("--seed", type=int, default=42, help="Dataset split seed and evaluation sample seed.") + parser.add_argument("--data-dir", type=Path, default=ROOT / "data" / "modular_arithmetic", help="Dataset output directory.") + parser.add_argument("--out-dir", type=Path, default=None, help="Training/checkpoint output directory.") + parser.add_argument("--ckpt", type=Path, default=None, help="Checkpoint to evaluate. Defaults to OUT_DIR/ckpt.pt.") + parser.add_argument("--skip-train", action="store_true", help="Only prepare data and evaluate an existing checkpoint.") + parser.add_argument("--skip-prepare", action="store_true", help="Reuse existing data files instead of regenerating them.") + parser.add_argument("--device", default=os.environ.get("DEVICE", "cuda:0"), help="Torch device for training/eval.") + parser.add_argument("--dtype", default=os.environ.get("DTYPE", "float16"), choices=["float32", "float16", "bfloat16"], help="Training/eval dtype.") + parser.add_argument("--max-iters", type=int, default=20000, help="Training iterations.") + parser.add_argument("--batch-size", type=int, default=256, help="Training batch size.") + parser.add_argument("--eval-interval", type=int, default=500, help="Training eval interval.") + parser.add_argument("--eval-iters", type=int, default=100, help="Training eval iterations.") + parser.add_argument("--learning-rate", type=float, default=1e-3, help="AdamW learning rate.") + parser.add_argument("--weight-decay", type=float, default=1.0, help="AdamW weight decay; high values encourage grokking.") + parser.add_argument("--n-layer", type=int, default=1, help="Transformer layers.") + parser.add_argument("--n-head", type=int, default=4, help="Attention heads.") + parser.add_argument("--n-embd", type=int, default=128, help="Embedding width.") + parser.add_argument("--block-size", type=int, default=32, help="Context window.") + parser.add_argument("--eval-split", choices=["all", "train", "val"], default="all", help="Which equation split to evaluate.") + parser.add_argument("--max-eval-examples", type=int, default=None, help="Optional cap on evaluated equations.") + parser.add_argument("--show-examples", type=int, default=12, help="Number of predictions to print.") + return parser.parse_args() + + +def run(cmd: list[str]) -> None: + print("+", " ".join(cmd), flush=True) + subprocess.run(cmd, cwd=ROOT, check=True) + + +def prepare_data(args: argparse.Namespace) -> None: + run([ + sys.executable, str(ROOT / "data" / "modular_arithmetic" / "prepare.py"), + "--out-dir", str(args.data_dir), + "--modulus", str(args.modulus), + "--train-fraction", str(args.train_fraction), + "--train-repeats", str(args.train_repeats), + "--val-repeats", str(args.val_repeats), + "--seed", str(args.seed), + ]) + + +def train(args: argparse.Namespace, out_dir: Path) -> None: + run([ + sys.executable, "train.py", + "--dataset", "modular_arithmetic", + "--out_dir", str(out_dir), + "--device", args.device, + "--dtype", args.dtype, + "--block_size", str(args.block_size), + "--batch_size", str(args.batch_size), + "--n_layer", str(args.n_layer), + "--n_head", str(args.n_head), + "--n_embd", str(args.n_embd), + "--dropout", "0.0", + "--bias", + "--max_iters", str(args.max_iters), + "--eval_interval", str(args.eval_interval), + "--eval_iters", str(args.eval_iters), + "--learning_rate", str(args.learning_rate), + "--weight_decay", str(args.weight_decay), + "--warmup_iters", "100", + "--decay_lr", + "--min_lr", "1e-5", + "--always_save_checkpoint", + "--only_save_checkpoint_at_end", + "--no-compile", + ]) + + +def load_model(ckpt_path: Path, device: str): + import torch + from model import GPT + from gpt_conf import GPTConfig + + from inspect import signature + + load_kwargs = {"map_location": device} + if "weights_only" in signature(torch.load).parameters: + load_kwargs["weights_only"] = False + checkpoint = torch.load(ckpt_path, **load_kwargs) + checkpoint["model_args"]["dropout"] = 0.0 + model = GPT(GPTConfig(**checkpoint["model_args"])) + state_dict = checkpoint["model"] + for key in list(state_dict.keys()): + if key.startswith("_orig_mod."): + state_dict[key[len("_orig_mod."):]] = state_dict.pop(key) + model.load_state_dict(state_dict, strict=False) + model.to(device) + model.eval() + return model, checkpoint + + +def split_pairs(args: argparse.Namespace): + rng = random.Random(args.seed) + pairs = [(a, b) for a in range(args.modulus) for b in range(args.modulus)] + rng.shuffle(pairs) + split = int(len(pairs) * args.train_fraction) + return pairs[:split], pairs[split:] + + +def encode(text: str, stoi: dict[str, int]): + import torch + + return torch.tensor([stoi[ch] for ch in text], dtype=torch.long).unsqueeze(0) + + +def predict_answer(model, prompt: str, stoi: dict[str, int], itos: dict[int, str], device: str, max_digits: int) -> str: + import torch + from torch.nn import functional as F + + idx = encode(prompt, stoi).to(device) + out_chars: list[str] = [] + for _ in range(max_digits + 1): + idx_cond = idx if idx.size(1) <= model.config.block_size else idx[:, -model.config.block_size:] + logits, _ = model(idx_cond) + probs = F.softmax(logits[:, -1, :], dim=-1) + next_id = int(torch.argmax(probs, dim=-1).item()) + ch = itos[next_id] + if ch == "\n": + break + out_chars.append(ch) + idx = torch.cat([idx, torch.tensor([[next_id]], device=device)], dim=1) + return "".join(out_chars) + + +def evaluate(args: argparse.Namespace, ckpt_path: Path) -> dict[str, float]: + import torch + + meta_path = args.data_dir / "meta.pkl" + with open(meta_path, "rb") as f: + meta = pickle.load(f) + stoi = meta["stoi"] + itos = {int(k): v for k, v in meta["itos"].items()} + model, _ = load_model(ckpt_path, args.device) + train_pairs, val_pairs = split_pairs(args) + pairs = {"train": train_pairs, "val": val_pairs, "all": train_pairs + val_pairs}[args.eval_split] + if args.max_eval_examples is not None: + pairs = pairs[: args.max_eval_examples] + max_digits = len(str(args.modulus - 1)) + correct = 0 + examples = [] + with torch.no_grad(): + for a, b in pairs: + expected = str((a + b) % args.modulus) + pred = predict_answer(model, f"{a}+{b}=", stoi, itos, args.device, max_digits) + ok = pred == expected + correct += int(ok) + if len(examples) < args.show_examples: + examples.append({"prompt": f"{a}+{b}=", "prediction": pred, "expected": expected, "correct": ok}) + accuracy = correct / max(1, len(pairs)) + result = {"split": args.eval_split, "accuracy": accuracy, "correct": correct, "total": len(pairs), "examples": examples} + print(json.dumps(result, indent=2)) + return result + + +def main() -> None: + args = parse_args() + out_dir = args.out_dir or ROOT / "out" / f"modular_arithmetic_grokking_p{args.modulus}" + ckpt_path = args.ckpt or out_dir / "ckpt.pt" + args.data_dir.mkdir(parents=True, exist_ok=True) + out_dir.mkdir(parents=True, exist_ok=True) + if not args.skip_prepare: + prepare_data(args) + if not args.skip_train: + train(args, out_dir) + if not ckpt_path.exists(): + raise FileNotFoundError(f"Checkpoint not found: {ckpt_path}") + evaluate(args, ckpt_path) + + +if __name__ == "__main__": + main() diff --git a/explorations/modular_arithmetic_grokking.yaml b/explorations/modular_arithmetic_grokking.yaml new file mode 100644 index 0000000000..6cfcca80c4 --- /dev/null +++ b/explorations/modular_arithmetic_grokking.yaml @@ -0,0 +1,50 @@ +# Modular addition grokking sweep. +# Literature-informed defaults: +# - prime moduli p=97 and p=113 are common modular-arithmetic grokking testbeds +# - limited training data (around 30%-50% of ordered pairs) plus substantial +# AdamW weight decay encourages delayed generalization/grokking +# - a small 1-layer transformer with d_model=128 is enough for the task +--- +common_group: + dataset: ["modular_arithmetic"] + block_size: [32] + batch_size: [256] + n_layer: [1] + n_head: [4] + n_embd: [128] + dropout: [0.0] + learning_rate: [0.001] + min_lr: [0.00001] + eval_interval: [500] + eval_iters: [100] + max_iters: [20000] + always_save_checkpoint: [true] + only_save_checkpoint_at_end: [true] + compile: [true] + dtype: ["float16"] + device: ["cuda:0"] + log_lm_head_vector_stats: [true] + log_post_norm_lm_head_vector_stats: [true] + +parameter_groups: + # Main replication setting; prepare with: + # python data/modular_arithmetic/prepare.py --modulus 113 --train-fraction 0.3 --train-repeats 200 --val-repeats 20 + - run_name_override: ["mod_add_p113_train30_wd1"] + tensorboard_run_name: ["mod_add_p113_train30_wd1"] + norm_variant_lm_head: ["cappedhyperspherenorm"] + + - run_name_override: ["mod_add_p113_train30_wd1"] + tensorboard_run_name: ["mod_add_p113_train30_wd1"] + weight_decay: [1.0, 5.0, 10.0, 50.0, 100.0] + norm_variant_lm_head: ["cappedhyperspherenorm"] + + - run_name_override: ["mod_add_p113_train30_wd1"] + tensorboard_run_name: ["mod_add_p113_train30_wd1"] + optimizer: ["muon"] + weight_decay: [0] + + - run_name_override: ["mod_add_p113_train30_wd1"] + tensorboard_run_name: ["mod_add_p113_train30_wd1"] + norm_variant_lm_head: ["cappedhyperspherenorm"] + optimizer: ["muon"] + weight_decay: [0] diff --git a/optimization_and_search/run_experiments.py b/optimization_and_search/run_experiments.py index 379514328e..59e5fbc78a 100644 --- a/optimization_and_search/run_experiments.py +++ b/optimization_and_search/run_experiments.py @@ -44,6 +44,16 @@ "areq", "distillation_val_loss", "ntp_val_loss", + "lm_head_magnitude_max", + "lm_head_magnitude_avg", + "lm_head_magnitude_median", + "lm_head_magnitude_std", + "lm_head_magnitude_min", + "post_norm_lm_head_magnitude_max", + "post_norm_lm_head_magnitude_avg", + "post_norm_lm_head_magnitude_median", + "post_norm_lm_head_magnitude_std", + "post_norm_lm_head_magnitude_min", "zeus_total_energy_j", "zeus_total_time_s", "zeus_avg_power_w", @@ -604,6 +614,26 @@ def read_metrics(out_dir: str) -> dict: "areq", "distillation_val_loss", "ntp_val_loss", + "overall_weight_stdev", + "overall_weight_kurtosis", + "overall_weight_max", + "overall_weight_min", + "overall_weight_abs_max", + "overall_activation_stdev", + "overall_activation_kurtosis", + "overall_activation_max", + "overall_activation_min", + "overall_activation_abs_max", + "lm_head_magnitude_max", + "lm_head_magnitude_avg", + "lm_head_magnitude_median", + "lm_head_magnitude_std", + "lm_head_magnitude_min", + "post_norm_lm_head_magnitude_max", + "post_norm_lm_head_magnitude_avg", + "post_norm_lm_head_magnitude_median", + "post_norm_lm_head_magnitude_std", + "post_norm_lm_head_magnitude_min", ] casts = [ float, @@ -631,19 +661,38 @@ def read_metrics(out_dir: str) -> dict: float, float, float, + float, + float, + float, + float, + float, + float, + float, + float, + float, + float, + float, + float, + float, + float, + float, + float, + float, + float, + float, + float, ] if len(base_metric_keys) != len(casts): raise ValueError( f"Metric schema mismatch: {len(base_metric_keys)} keys vs {len(casts)} casts." ) - if len(parts) == len(base_metric_keys) - 1: + if len(parts) in {24, 34, len(base_metric_keys) - 1}: # Backward compatibility for runs created before best_val_bits_per_byte. parts.insert(1, "") if len(parts) < len(base_metric_keys): - raise ValueError( - f"Expected at least {len(base_metric_keys)} metrics in {path}, got {len(parts)}." - ) + # Backward compatibility for runs created before newer optional metrics. + parts.extend([""] * (len(base_metric_keys) - len(parts))) metrics: dict[str, float] = {} for key, typ, value in zip(base_metric_keys, casts, parts): diff --git a/run_exploration_monitor.py b/run_exploration_monitor.py index d91ff967f7..4d4c12066f 100644 --- a/run_exploration_monitor.py +++ b/run_exploration_monitor.py @@ -206,6 +206,16 @@ def on_mount(self) -> None: "areq", "distillation_val_loss", "ntp_val_loss", + "lm_head_magnitude_max", + "lm_head_magnitude_avg", + "lm_head_magnitude_median", + "lm_head_magnitude_std", + "lm_head_magnitude_min", + "post_norm_lm_head_magnitude_max", + "post_norm_lm_head_magnitude_avg", + "post_norm_lm_head_magnitude_median", + "post_norm_lm_head_magnitude_std", + "post_norm_lm_head_magnitude_min", "zeus_total_energy_j", "zeus_total_time_s", "zeus_avg_power_w", diff --git a/train.py b/train.py index 4923f9c308..4aa02190da 100644 --- a/train.py +++ b/train.py @@ -157,6 +157,8 @@ def __init__(self, args, model_group, training_group, logging_group): self.latest_areq = float('nan') self.latest_bits_per_byte = float('nan') self.bits_per_byte_scales = {} + self.latest_lm_head_magnitude_stats = self._empty_lm_head_magnitude_stats() + self.latest_post_norm_lm_head_magnitude_stats = self._empty_lm_head_magnitude_stats() # store overall statistics for weights and activations self.latest_overall_weight_stats = { @@ -1607,6 +1609,88 @@ def _safe_better_than_chance(self, vocab_size: float, loss_value: float) -> floa return 0.0 return vocab_size / math.exp(loss_value) + @staticmethod + def _empty_lm_head_magnitude_stats() -> dict[str, float]: + return { + "max": float('nan'), + "avg": float('nan'), + "median": float('nan'), + "std": float('nan'), + "min": float('nan'), + } + + def _iter_lm_head_modules(self): + if getattr(self.raw_model.config, "multicontext", False): + idx = 0 + while f'lm_head_{idx}' in self.raw_model.transformer: + yield self.raw_model.transformer[f'lm_head_{idx}'] + idx += 1 + elif hasattr(self.raw_model, 'lm_head'): + yield self.raw_model.lm_head + + def _compute_lm_head_magnitude_stats(self, apply_lm_head_norm: bool = False) -> dict[str, float]: + vectors = [] + with torch.no_grad(): + for lm_head in self._iter_lm_head_modules(): + weight = lm_head.weight.detach() + if apply_lm_head_norm: + weight = self.raw_model.apply_lm_head_norm(weight) + vectors.append(weight.float()) + if not vectors: + return self._empty_lm_head_magnitude_stats() + magnitudes = torch.linalg.vector_norm(torch.cat(vectors, dim=0), ord=2, dim=1) + return { + "max": magnitudes.max().item(), + "avg": magnitudes.mean().item(), + "median": magnitudes.median().item(), + "std": magnitudes.std(unbiased=False).item(), + "min": magnitudes.min().item(), + } + + def _update_lm_head_magnitude_stats(self) -> None: + if self.args.log_lm_head_vector_stats: + self.latest_lm_head_magnitude_stats = self._compute_lm_head_magnitude_stats(False) + if ( + self.args.log_post_norm_lm_head_vector_stats + and getattr(self.raw_model.config, "norm_variant_lm_head", None) is not None + ): + self.latest_post_norm_lm_head_magnitude_stats = self._compute_lm_head_magnitude_stats(True) + + def _log_lm_head_magnitude_stats_tensorboard(self, target_dataset: str, tokens_trained: float) -> None: + if not self.args.tensorboard_log or self.writer is None: + return + stat_groups = [] + if self.args.log_lm_head_vector_stats: + stat_groups.append(("lm_head_magnitude", self.latest_lm_head_magnitude_stats)) + if ( + self.args.log_post_norm_lm_head_vector_stats + and getattr(self.raw_model.config, "norm_variant_lm_head", None) is not None + ): + stat_groups.append(("post_norm_lm_head_magnitude", self.latest_post_norm_lm_head_magnitude_stats)) + for tag, stats in stat_groups: + for stat_name, stat_value in stats.items(): + self.writer.add_scalar(f"{target_dataset}/{tag}_{stat_name}_iters", stat_value, self.iter_num) + self.writer.add_scalar(f"{target_dataset}/{tag}_{stat_name}_tokens", stat_value, tokens_trained) + + def _lm_head_magnitude_csv_values(self) -> list[float]: + values = [] + for enabled, stats in ( + (self.args.log_lm_head_vector_stats, self.latest_lm_head_magnitude_stats), + (self.args.log_post_norm_lm_head_vector_stats, self.latest_post_norm_lm_head_magnitude_stats), + ): + if enabled: + values.extend(stats[name] for name in ("max", "avg", "median", "std", "min")) + return values + + def _format_lm_head_magnitude_best_values(self) -> list[str]: + values = [] + for stats in ( + self.latest_lm_head_magnitude_stats, + self.latest_post_norm_lm_head_magnitude_stats, + ): + values.extend(f"{stats[name]:.6f}" for name in ("max", "avg", "median", "std", "min")) + return values + def log_metrics(self, losses, running_mfu, epoch, tokens_trained, target_dataset, val_better_than_chance): compute_rankme = self.args.log_rankme or self.args.log_areq @@ -1743,6 +1827,7 @@ def log_metrics(self, losses, running_mfu, epoch, tokens_trained, target_dataset self._log_zeus_tensorboard(target_dataset, tokens_trained) self._log_bit_metrics(target_dataset, tokens_trained) + self._log_lm_head_magnitude_stats_tensorboard(target_dataset, tokens_trained) if self.args.csv_log: @@ -1759,6 +1844,7 @@ def log_metrics(self, losses, running_mfu, epoch, tokens_trained, target_dataset losses['val'].item(), rankme_value, areq_value, + *self._lm_head_magnitude_csv_values(), prefix=f"{target_dataset}_", ) @@ -1770,6 +1856,7 @@ def log_metrics(self, losses, running_mfu, epoch, tokens_trained, target_dataset rankme_value, areq_value, running_mfu, + *self._lm_head_magnitude_csv_values(), prefix="bulk_", ) @@ -2141,6 +2228,7 @@ def run_validation_step(self, running_mfu, current_epoch, current_dataset, num_s self.latest_areq = self._to_scalar(losses.get('areq', float('nan'))) self.latest_distillation_val_loss = self._to_scalar(losses.get('distillation_val_loss', float('nan'))) self.latest_ntp_val_loss = self._to_scalar(losses.get('ntp_val_loss', float('nan'))) + self._update_lm_head_magnitude_stats() if self.args.gns_type is not None: self.gns = self.gns_ema.get_gns() @@ -2284,6 +2372,7 @@ def run_validation_step(self, running_mfu, current_epoch, current_dataset, num_s f"{self.latest_overall_activation_stats['max']:.6f}", f"{self.latest_overall_activation_stats['min']:.6f}", f"{self.latest_overall_activation_stats['abs_max']:.6f}", + *self._format_lm_head_magnitude_best_values(), ] best_loss_file.write(", ".join(metrics) + "\n") num_steps_with_worse_loss = 0 diff --git a/train_args.py b/train_args.py index 8e4497f7d9..83edbb5031 100644 --- a/train_args.py +++ b/train_args.py @@ -1522,6 +1522,8 @@ def parse_args(): logging_group.add_argument('--log_all_metrics', default=False, action=argparse.BooleanOptionalAction, help='Enable logging of all metrics including gns') logging_group.add_argument('--log_rankme', default=True, action=argparse.BooleanOptionalAction, help='Log RankMe representation metric during validation') logging_group.add_argument('--log_areq', default=True, action=argparse.BooleanOptionalAction, help='Log aReQ representation metric during validation') + logging_group.add_argument('--log_lm_head_vector_stats', default=False, action=argparse.BooleanOptionalAction, help='Log raw lm_head row-vector magnitude summary statistics to TensorBoard/CSV') + logging_group.add_argument('--log_post_norm_lm_head_vector_stats', default=False, action=argparse.BooleanOptionalAction, help='Log lm_head row-vector magnitude summary statistics after norm_variant_lm_head is applied') # Turn activation/weight statistics off to save CPU RAM and wall time. training_group.add_argument( diff --git a/variations/norm_variations.py b/variations/norm_variations.py index 23eee6f867..d83e54f5a3 100644 --- a/variations/norm_variations.py +++ b/variations/norm_variations.py @@ -208,7 +208,13 @@ class CappedHyperSphereNorm(nn.Module): def __init__(self, config): super().__init__() - self.radius = math.sqrt(config.n_embd) + # Determine radius initialization value + radius_init = None + if config.hsnorm_radius is not None: + self.radius = config.hsnorm_radius + else: + self.radius = math.sqrt(config.n_embd) + print(self.radius) def forward(self, x): norms = x.norm(2, dim=-1, keepdim=True)