import json import os import sys from pathlib import Path import numpy as np import torch import torch.nn as nn # Import the read-only shared benchmark package. sys.path.insert(0, "/home/maxwelhelp/all/math2nn") from bench import get_dataset, train_model, evaluate, sweep_baseline, make_report SEEDS = [0, 1, 2, 3, 4, 5, 6, 7] # The union of baseline and idea learning rates is identical. GRID = [{"lr": 1e-3}, {"lr": 3e-3}, {"lr": 6e-3}] EPOCHS = 8 BATCH = 128 NTRAIN = 400 NTEST = 400 def make_phi(device): z = torch.linspace(-1.0, 1.0, 8, device=device) z1, z2 = torch.meshgrid(z, z, indexing="ij") return torch.stack((torch.ones_like(z1), z1, z2, 0.5 * (z1 * z1 + z2 * z2))) def conserve_hidden(h, rank=2): """Rank compress each 8x8 hidden state and restore four weighted moments.""" b = h.shape[0] x = h.reshape(b, 8, 8) # Batched truncated SVD is the low-rank compression intervention. u, s, vh = torch.linalg.svd(x, full_matrices=False) xt = (u[:, :, :rank] * s[:, None, :rank]) @ vh[:, :rank, :] phi = make_phi(h.device) flat_phi = phi.reshape(4, -1) flat_x = x.reshape(b, -1) flat_t = xt.reshape(b, -1) gram = flat_phi @ flat_phi.T delta = (flat_phi @ (flat_x - flat_t).T).T coeff = torch.linalg.solve(gram + 1e-8 * torch.eye(4, device=h.device), delta.T).T return (flat_t + coeff @ flat_phi).reshape(b, 64) class IdeaRNN(nn.Module): """Same rnn_small architecture, with conservative hidden compression.""" def __init__(self, hidden=64, rank=2): super().__init__() self.rnn = nn.GRU(3, hidden, batch_first=True) self.head = nn.Linear(hidden, 1) self.rank = rank self.last_signature = {} def forward(self, x): seq = x.view(x.shape[0], -1, 3) try: _, h = self.rnn(seq) except RuntimeError: old = torch.backends.cudnn.enabled torch.backends.cudnn.enabled = False try: _, h = self.rnn(seq) finally: torch.backends.cudnn.enabled = old raw = h[-1] out = conserve_hidden(raw, self.rank) if not self.training: with torch.no_grad(): phi = make_phi(raw.device).reshape(4, -1) self.last_signature = { "raw_norm": float(raw.norm().cpu()), "compressed_relative_error": float((out - raw).norm().cpu() / (raw.norm().cpu() + 1e-12)), "moment_residual": float(((out - raw) @ phi.T).abs().max().cpu()), } return self.head(out) def make_baseline(cfg): def fn(seed): torch.manual_seed(seed) np.random.seed(seed) from bench import make_model ds = get_dataset("dynamics", seed, NTRAIN, NTEST) model = make_model("rnn_small", ds["input_shape"], ds["out_dim"]) _, metric, _ = train_model(model, ds, epochs=EPOCHS, lr=cfg["lr"], batch=BATCH, log=lambda *_: None) return metric return fn def make_idea(cfg): def fn(seed): torch.manual_seed(seed) np.random.seed(seed) ds = get_dataset("dynamics", seed, NTRAIN, NTEST) model = IdeaRNN(rank=2) _, metric, _ = train_model(model, ds, epochs=EPOCHS, lr=cfg["lr"], batch=BATCH, log=lambda *_: None) return metric return fn def inspect_signature(seed, cfg): """Measure the proposed preservation on a trained benchmark model.""" torch.manual_seed(seed) ds = get_dataset("dynamics", seed, NTRAIN, NTEST) model = IdeaRNN(rank=2) trained, _, _ = train_model(model, ds, epochs=EPOCHS, lr=cfg["lr"], batch=BATCH, log=lambda *_: None) trained.eval() device = next(trained.parameters()).device with torch.no_grad(): _ = trained(ds["xte"].to(device)) sig = dict(trained.last_signature) sig.update({"target": "hidden moments", "predicted_residual": 0.0, "observed_residual": sig.get("moment_residual", float("nan"))}) sig["confirmed"] = bool(sig.get("observed_residual", 1.0) < 1e-5) return sig def main(): # Cheap numerical verification before any benchmark training. torch.manual_seed(329) h = torch.randn(5, 64) before = h.reshape(5, 8, 8) phi = make_phi(h.device).reshape(4, -1) after = conserve_hidden(h).reshape(5, -1) math_resid = float(((after - before.reshape(5, -1)) @ phi.T).abs().max()) assert math_resid < 1e-5 base = sweep_baseline(make_baseline, GRID, seeds=SEEDS) idea_candidates = [] for cfg in GRID: r = evaluate(make_idea(cfg), seeds=SEEDS) idea_candidates.append((r["mean"], cfg, r)) _, best_cfg, idea = min(idea_candidates, key=lambda q: q[0]) trained_sig = inspect_signature(SEEDS[0], best_cfg) report = make_report( "dynamics", "rnn_small", base, idea, {"mechanism_signature": { "math_check_max_moment_residual": math_resid, "trained_model": trained_sig, "predicted": "four hidden-state moments preserved after rank-2 compression", "observed": "test-time hidden-state moment residual measured from trained model", "confirmed": bool(math_resid < 1e-5 and trained_sig.get("confirmed", False)), }, "protocol": {"epochs": EPOCHS, "n_train": NTRAIN, "n_test": NTEST, "grid": GRID, "seeds": SEEDS}}) report["math_verification"] = {"max_moment_residual": math_resid, "passed": math_resid < 1e-5} Path("bench_report.json").write_text(json.dumps(report, indent=2)) print(json.dumps(report, indent=2)) if __name__ == "__main__": main()