import json, math, random, sys import numpy as np import torch sys.path.insert(0, '/home/maxwelhelp/all/math2nn') from bench import get_dataset, make_model, train_model, evaluate, sweep_baseline, make_report from moment_sharp_bench import moment_sharp_rescale, mechanism_signature TRACK, MODEL = 'tabular', 'mlp_tiny' EPOCHS, BATCH, NTR, NTE = 12, 128, 1000, 400 # Union of baseline and idea settings: all idea lrs are baseline-evaluated. LRS = [0.0015, 0.003, 0.006] WDS = [0.0, 1e-4, 1e-3] BASE_GRID = [{'lr': lr, 'weight_decay': wd} for lr in LRS for wd in WDS] IDEA_GRID = [{'lr': 0.003, 'weight_decay': 1e-4, 'target_sigma': 2.0}, {'lr': 0.0015, 'weight_decay': 1e-4, 'target_sigma': 2.0}, {'lr': 0.006, 'weight_decay': 1e-4, 'target_sigma': 2.0}] SWEEP_SEEDS = (0, 1, 2, 3) FULL_SEEDS = tuple(range(8)) def seed_all(seed): random.seed(seed); np.random.seed(seed); torch.manual_seed(seed) if torch.cuda.is_available(): try: torch.cuda.manual_seed_all(seed) except Exception: pass def baseline_fn(cfg): def run(seed): seed_all(seed) ds = get_dataset(TRACK, seed, n_train=NTR, n_test=NTE) net = make_model(MODEL, ds['input_shape'], ds['out_dim']) _, metric, _ = train_model(net, ds, epochs=EPOCHS, lr=cfg['lr'], batch=BATCH, weight_decay=cfg['weight_decay'], log=lambda *_: None) return float(metric) return run def idea_run(cfg, seed, want_signature=False): seed_all(seed) ds = get_dataset(TRACK, seed, n_train=NTR, n_test=NTE) net = make_model(MODEL, ds['input_shape'], ds['out_dim']) # Same AdamW and minibatch schedule as bench, with control after each update. device = 'cuda' if torch.cuda.is_available() else 'cpu' try: net = net.to(device) x, y = ds['xtr'].to(device), ds['ytr'].to(device) opt = torch.optim.AdamW(net.parameters(), lr=cfg['lr'], weight_decay=cfg['weight_decay']) gen = torch.Generator(device=device); gen.manual_seed(seed + 10000) ema = {} net.train() for _ in range(EPOCHS): order = torch.randperm(len(x), generator=gen, device=device) for start in range(0, len(x), BATCH): ix = order[start:start+BATCH] opt.zero_grad(set_to_none=True) pred = net(x[ix]); loss = torch.nn.functional.mse_loss(pred, y[ix]) loss.backward(); opt.step() moment_sharp_rescale(net, target_sigma=cfg['target_sigma'], probes=8, ema=ema, generator=gen) net.eval() with torch.no_grad(): metric = float(torch.nn.functional.mse_loss(net(ds['xte'].to(device)), ds['yte'].to(device)).cpu()) sig = mechanism_signature(net, target_sigma=cfg['target_sigma']) if want_signature else None return metric, sig except Exception: # Explicit CUDA -> CPU fallback, preserving deterministic config/seed. seed_all(seed) device = 'cpu'; net = make_model(MODEL, ds['input_shape'], ds['out_dim']) x, y = ds['xtr'], ds['ytr'] opt = torch.optim.AdamW(net.parameters(), lr=cfg['lr'], weight_decay=cfg['weight_decay']) gen = torch.Generator().manual_seed(seed + 10000); ema = {} net.train() for _ in range(EPOCHS): order = torch.randperm(len(x), generator=gen) for start in range(0, len(x), BATCH): ix = order[start:start+BATCH]; opt.zero_grad(set_to_none=True) torch.nn.functional.mse_loss(net(x[ix]), y[ix]).backward(); opt.step() moment_sharp_rescale(net, target_sigma=cfg['target_sigma'], probes=8, ema=ema, generator=gen) net.eval() with torch.no_grad(): metric = float(torch.nn.functional.mse_loss(net(ds['xte']), ds['yte'])) sig = mechanism_signature(net, target_sigma=cfg['target_sigma']) if want_signature else None return metric, sig def idea_fn(cfg): return lambda seed: idea_run(cfg, seed)[0] def main(): base = sweep_baseline(baseline_fn, BASE_GRID, seeds=SWEEP_SEEDS) # Evaluate each idea setting on the full paired seeds; select by mean. idea_trials = [] for cfg in IDEA_GRID: r = evaluate(idea_fn(cfg), seeds=FULL_SEEDS) idea_trials.append({'cfg': cfg, 'result': r}) best_trial = min(idea_trials, key=lambda z: z['result']['mean']) best_cfg = best_trial['cfg'] idea_res = best_trial['result'] signatures = [idea_run(best_cfg, s, want_signature=True)[1] for s in FULL_SEEDS] sig = {'per_seed': signatures, 'predicted_bound_sigma_mean': float(np.mean([np.mean(x['predicted_bound_sigma']) for x in signatures])), 'observed_true_sigma_mean': float(np.mean([np.mean(x['observed_true_sigma']) for x in signatures])), 'max_bound_minus_observed': float(max(x['max_bound_minus_observed'] for x in signatures)), 'confirmed': bool(all(x['confirmed'] for x in signatures))} extra = {'mechanism_signature': sig, 'idea_trials': idea_trials, 'protocol_note': 'tabular is structurally matched: regularization/stability of linear layers.'} report = make_report(TRACK, MODEL, base, idea_res, extra) report['math_check'] = json.load(open('stage2_report.json'))['math_check'] if __import__('os').path.exists('stage2_report.json') else None with open('bench_report.json', 'w') as f: json.dump(report, f, indent=2) print(json.dumps(report, indent=2)) if __name__ == '__main__': main()