import sys, json from pathlib import Path import numpy as np import torch import torch.nn as nn sys.path.insert(0, '/home/maxwelhelp/all/math2nn') from bench import get_dataset, make_model, evaluate, sweep_baseline, make_report SEEDS = tuple(range(8)) SWEEP_SEEDS = (0, 1, 2, 3) TRACK, MODEL = 'unordered_pointset_denoising', 'mlp_med' EPOCHS, BATCH = 18, 128 # Union of step sizes is shared by both sides; baseline also sweeps its weight decay. BASE_GRID = [{'lr': lr, 'weight_decay': wd} for lr in (1e-3, 3e-3, 6e-3) for wd in (0.0, 1e-4)] IDEA_GRID = [ {'lr': 1e-3, 'weight_decay': 0.0, 'tree_lambda': 0.02}, {'lr': 3e-3, 'weight_decay': 0.0, 'tree_lambda': 0.05}, {'lr': 6e-3, 'weight_decay': 1e-4, 'tree_lambda': 0.10}, ] def tree_features(z): """Finite wide rooted contractions over the six unordered set elements.""" c = z - z.mean(dim=1, keepdim=True) p2, p3, p4 = (c**2).mean(1), (c**3).mean(1), (c**4).mean(1) return torch.stack((z.mean(1), p2, p3, p4, p2*p2, p2*p3), dim=1) def tree_loss(pred, target): a, b = tree_features(pred), tree_features(target) scale = b.detach().std(0, unbiased=False).clamp_min(1e-3) return (((a-b)/scale)**2).mean() def fit(seed, cfg, tree_lambda=0.0, capture=False): torch.manual_seed(seed); np.random.seed(seed) d = get_dataset(TRACK, seed, n_train=400, n_test=160) # This custom track is a six-coordinate set; bench's generic regression # adapter flattens y, so restore one target set per input example. d['ytr'] = d['ytr'].reshape(d['xtr'].shape[0], -1) d['yte'] = d['yte'].reshape(d['xte'].shape[0], -1) def run(device): net = make_model(MODEL, d['input_shape'], d['out_dim']).to(device) x, y = d['xtr'].to(device), d['ytr'].to(device) opt = torch.optim.Adam(net.parameters(), lr=cfg['lr'], weight_decay=cfg['weight_decay']) for _ in range(EPOCHS): net.train(); perm = torch.randperm(len(x), device=device) for i in range(0, len(x), BATCH): ix = perm[i:i+BATCH]; pred = net(x[ix]) loss = ((pred-y[ix])**2).mean() if tree_lambda: loss = loss + tree_lambda * tree_loss(pred, y[ix]) opt.zero_grad(); loss.backward(); opt.step() net.eval() with torch.no_grad(): pred = net(d['xte'].to(device)); metric = float(((pred-d['yte'].to(device))**2).mean()) return metric, pred.detach().cpu() try: metric, pred = run('cuda' if torch.cuda.is_available() else 'cpu') except RuntimeError: metric, pred = run('cpu') return (metric, d, pred) if capture else metric def main(): # Numerical check of the Gram-based rooted contractions used by the prior stage. rng = np.random.default_rng(123); x = rng.normal(size=(32, 16)).astype('float32') q, _ = np.linalg.qr(rng.normal(size=(16, 16))) def gram_f(a): A = torch.tensor(a) @ torch.tensor(a).T / a.shape[1] one = torch.ones((len(a), 1)); root = A @ one return torch.cat((A @ one, root**2, A @ (A @ one)), 1) gx, gq = gram_f(x), gram_f(x @ q.astype('float32')) inv_err = float((gx-gq).abs().max() / gx.abs().max().clamp_min(1e-8)) base = sweep_baseline(lambda cfg: lambda seed: fit(seed, cfg, 0.0), BASE_GRID, seeds=SWEEP_SEEDS) trials = [] for cfg in IDEA_GRID: r = evaluate(lambda seed, c=cfg: fit(seed, c, c['tree_lambda']), seeds=SWEEP_SEEDS) trials.append({'cfg': cfg, 'mean': r['mean']}) best = min(trials, key=lambda r: r['mean'])['cfg'] idea = evaluate(lambda seed: fit(seed, best, best['tree_lambda']), seeds=SEEDS) sig = [] for s in SEEDS: bm, bd, bp = fit(s, base['best_cfg'], 0.0, True) im, _, ip = fit(s, best, best['tree_lambda'], True) target = tree_features(bd['yte']) def corr(a, b): aa, bb = a.numpy().ravel(), b.numpy().ravel() return float(np.corrcoef(aa, bb)[0, 1]) sig.append({'seed': s, 'baseline_tree_corr': corr(tree_features(bp), target), 'idea_tree_corr': corr(tree_features(ip), target), 'baseline_mse': bm, 'idea_mse': im}) cb = float(np.mean([r['baseline_tree_corr'] for r in sig])) ci = float(np.mean([r['idea_tree_corr'] for r in sig])) report = make_report(TRACK, MODEL, base, idea, extra={ 'custom_track': {'name': TRACK, 'file': 'bench/custom_tracks/unordered_pointset_denoising.py', 'domain': 'point-set-diffusion'}, 'baseline_grid': BASE_GRID, 'idea_grid': IDEA_GRID, 'idea_trials': trials, 'math_check': {'gram_tree_orthogonal_invariance_relative_error': inv_err}, 'mechanism_signature': { 'claim': 'tree aggregation preserves target set correspondence', 'predicted': 'tree-feature correlation is higher for the idea system', 'observed_baseline_mean_tree_corr': cb, 'observed_idea_mean_tree_corr': ci, 'per_seed': sig, 'confirmed': bool(ci > cb) }}) Path('bench_report.json').write_text(json.dumps(report, indent=2)) print(json.dumps(report, indent=2)) if __name__ == '__main__': main()