import sys, json, math from pathlib import Path import numpy as np import torch import torch.nn as nn import torch.nn.functional as F sys.path.insert(0, '/home/maxwelhelp/all/math2nn') from bench import get_dataset, make_report, evaluate, sweep_baseline TRACK = 'geometry_support_regression' SEEDS = tuple(range(8)) GRID = [{'lr': 1e-3}, {'lr': 3e-3}, {'lr': 6e-3}] R2 = 7.5 class MatchedMLP(nn.Module): def __init__(self, seed, idea=False): super().__init__() torch.manual_seed(seed) self.idea = idea self.features = nn.Sequential(nn.Linear(2, 64), nn.ReLU(), nn.Linear(64, 5), nn.Tanh()) self.head = nn.Linear(5, 1) def encode(self, x): return self.features(x) def forward(self, x): return self.head(self.encode(x)) def compact_score(z): d = z.shape[1] a = (R2 - d - 2.0) / 2.0 mu = z.detach().mean(0) q = z.detach() - mu cov = q.T @ q / max(1, z.shape[0] - 1) + 1e-3 * torch.eye(d, device=z.device) L = torch.linalg.cholesky(cov) v = torch.linalg.solve_triangular(L, (z - mu).T, upper=False).T r2 = (v * v).sum(1) logc = (torch.lgamma(torch.tensor(d / 2 + 1 + a, device=z.device)) - d / 2 * math.log(math.pi) - d / 2 * math.log(R2) - torch.lgamma(torch.tensor(1 + a, device=z.device))) score = -logc - a * torch.log(torch.clamp(1 - r2 / R2, min=1e-6)) + F.softplus(r2 - R2) return score, r2 def baseline_train(cfg, seed): np.random.seed(seed); torch.manual_seed(seed) ds = get_dataset(TRACK, seed, n_train=400, n_test=400) net = MatchedMLP(seed, idea=False) net, metric, _ = __import__('bench').train_model(net, ds, epochs=25, lr=cfg['lr'], batch=128, log=lambda *_: None) return float(metric) def idea_train(cfg, seed, capture=False): np.random.seed(seed); torch.manual_seed(seed) ds = get_dataset(TRACK, seed, n_train=400, n_test=400) try: device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') except Exception: device = torch.device('cpu') try: net = MatchedMLP(seed, idea=True).to(device) x = ds['xtr'].to(device); y = ds['ytr'].to(device) xt = ds['xte'].to(device); yt = ds['yte'].to(device) opt = torch.optim.Adam(net.parameters(), lr=cfg['lr']) for _ in range(25): net.train(); opt.zero_grad() z = net.encode(x); out = net.head(z) loss = F.mse_loss(out, y) + 0.002 * compact_score(z)[0].mean() loss.backward(); opt.step() net.eval() with torch.no_grad(): zte = net.encode(xt); pred = net.head(zte); metric = F.mse_loss(pred, yt).item() _, r2 = compact_score(zte) sig = {'mean_test_radius2': float(r2.mean().cpu()), 'test_inside_support_fraction': float((r2 < R2).float().mean().cpu())} return (float(metric), sig) if capture else float(metric) except RuntimeError: torch.backends.cudnn.enabled = False net = MatchedMLP(seed, idea=True) x, y, xt, yt = ds['xtr'], ds['ytr'], ds['xte'], ds['yte'] opt = torch.optim.Adam(net.parameters(), lr=cfg['lr']) for _ in range(25): opt.zero_grad(); z = net.encode(x); out = net.head(z) (F.mse_loss(out, y) + 0.002 * compact_score(z)[0].mean()).backward(); opt.step() with torch.no_grad(): zte = net.encode(xt); pred = net.head(zte); _, r2 = compact_score(zte) metric = F.mse_loss(pred, yt).item() sig = {'mean_test_radius2': float(r2.mean()), 'test_inside_support_fraction': float((r2 < R2).float().mean())} return (float(metric), sig) if capture else float(metric) def main(): base = sweep_baseline(lambda cfg: (lambda seed: baseline_train(cfg, seed)), GRID, seeds=(0,1,2,3)) best = base['best_cfg'] idea = evaluate(lambda seed: idea_train(best, seed), seeds=SEEDS) sigrows = [idea_train(best, s, capture=True)[1] for s in SEEDS] observed = float(np.mean([r['mean_test_radius2'] for r in sigrows])) extra = {'prediction': 'calibrated compact density predicts nominal E[r2]=d=5', 'predicted_mean_radius2': 5.0, 'observed_mean_radius2_on_trained_models': observed, 'absolute_error': abs(observed - 5.0), 'mean_inside_support_fraction': float(np.mean([r['test_inside_support_fraction'] for r in sigrows])), 'confirmed': bool(abs(observed - 5.0) < 0.5)} report = make_report(TRACK, 'matched_mlp_embedding_regressor', base, idea, extra) report['idea']['per_seed_signatures'] = sigrows Path('bench_report.json').write_text(json.dumps(report, indent=2)) print(json.dumps(report, indent=2)) if __name__ == '__main__': main()