Wide-tree invariant alignment layer / stage2_bench.py

Failed on benchmark

Raw ⬇ ZIP
  1import sys, json
  2from pathlib import Path
  3import numpy as np
  4import torch
  5import torch.nn as nn
  6sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
  7from bench import get_dataset, make_model, evaluate, sweep_baseline, make_report
  8
  9SEEDS = tuple(range(8))
 10SWEEP_SEEDS = (0, 1, 2, 3)
 11TRACK, MODEL = 'unordered_pointset_denoising', 'mlp_med'
 12EPOCHS, BATCH = 18, 128
 13# Union of step sizes is shared by both sides; baseline also sweeps its weight decay.
 14BASE_GRID = [{'lr': lr, 'weight_decay': wd}
 15             for lr in (1e-3, 3e-3, 6e-3) for wd in (0.0, 1e-4)]
 16IDEA_GRID = [
 17    {'lr': 1e-3, 'weight_decay': 0.0, 'tree_lambda': 0.02},
 18    {'lr': 3e-3, 'weight_decay': 0.0, 'tree_lambda': 0.05},
 19    {'lr': 6e-3, 'weight_decay': 1e-4, 'tree_lambda': 0.10},
 20]
 21
 22def tree_features(z):
 23    """Finite wide rooted contractions over the six unordered set elements."""
 24    c = z - z.mean(dim=1, keepdim=True)
 25    p2, p3, p4 = (c**2).mean(1), (c**3).mean(1), (c**4).mean(1)
 26    return torch.stack((z.mean(1), p2, p3, p4, p2*p2, p2*p3), dim=1)
 27
 28def tree_loss(pred, target):
 29    a, b = tree_features(pred), tree_features(target)
 30    scale = b.detach().std(0, unbiased=False).clamp_min(1e-3)
 31    return (((a-b)/scale)**2).mean()
 32
 33def fit(seed, cfg, tree_lambda=0.0, capture=False):
 34    torch.manual_seed(seed); np.random.seed(seed)
 35    d = get_dataset(TRACK, seed, n_train=400, n_test=160)
 36    # This custom track is a six-coordinate set; bench's generic regression
 37    # adapter flattens y, so restore one target set per input example.
 38    d['ytr'] = d['ytr'].reshape(d['xtr'].shape[0], -1)
 39    d['yte'] = d['yte'].reshape(d['xte'].shape[0], -1)
 40    def run(device):
 41        net = make_model(MODEL, d['input_shape'], d['out_dim']).to(device)
 42        x, y = d['xtr'].to(device), d['ytr'].to(device)
 43        opt = torch.optim.Adam(net.parameters(), lr=cfg['lr'], weight_decay=cfg['weight_decay'])
 44        for _ in range(EPOCHS):
 45            net.train(); perm = torch.randperm(len(x), device=device)
 46            for i in range(0, len(x), BATCH):
 47                ix = perm[i:i+BATCH]; pred = net(x[ix])
 48                loss = ((pred-y[ix])**2).mean()
 49                if tree_lambda: loss = loss + tree_lambda * tree_loss(pred, y[ix])
 50                opt.zero_grad(); loss.backward(); opt.step()
 51        net.eval()
 52        with torch.no_grad():
 53            pred = net(d['xte'].to(device)); metric = float(((pred-d['yte'].to(device))**2).mean())
 54        return metric, pred.detach().cpu()
 55    try:
 56        metric, pred = run('cuda' if torch.cuda.is_available() else 'cpu')
 57    except RuntimeError:
 58        metric, pred = run('cpu')
 59    return (metric, d, pred) if capture else metric
 60
 61def main():
 62    # Numerical check of the Gram-based rooted contractions used by the prior stage.
 63    rng = np.random.default_rng(123); x = rng.normal(size=(32, 16)).astype('float32')
 64    q, _ = np.linalg.qr(rng.normal(size=(16, 16)))
 65    def gram_f(a):
 66        A = torch.tensor(a) @ torch.tensor(a).T / a.shape[1]
 67        one = torch.ones((len(a), 1)); root = A @ one
 68        return torch.cat((A @ one, root**2, A @ (A @ one)), 1)
 69    gx, gq = gram_f(x), gram_f(x @ q.astype('float32'))
 70    inv_err = float((gx-gq).abs().max() / gx.abs().max().clamp_min(1e-8))
 71
 72    base = sweep_baseline(lambda cfg: lambda seed: fit(seed, cfg, 0.0), BASE_GRID,
 73                          seeds=SWEEP_SEEDS)
 74    trials = []
 75    for cfg in IDEA_GRID:
 76        r = evaluate(lambda seed, c=cfg: fit(seed, c, c['tree_lambda']), seeds=SWEEP_SEEDS)
 77        trials.append({'cfg': cfg, 'mean': r['mean']})
 78    best = min(trials, key=lambda r: r['mean'])['cfg']
 79    idea = evaluate(lambda seed: fit(seed, best, best['tree_lambda']), seeds=SEEDS)
 80
 81    sig = []
 82    for s in SEEDS:
 83        bm, bd, bp = fit(s, base['best_cfg'], 0.0, True)
 84        im, _, ip = fit(s, best, best['tree_lambda'], True)
 85        target = tree_features(bd['yte'])
 86        def corr(a, b):
 87            aa, bb = a.numpy().ravel(), b.numpy().ravel()
 88            return float(np.corrcoef(aa, bb)[0, 1])
 89        sig.append({'seed': s, 'baseline_tree_corr': corr(tree_features(bp), target),
 90                    'idea_tree_corr': corr(tree_features(ip), target),
 91                    'baseline_mse': bm, 'idea_mse': im})
 92    cb = float(np.mean([r['baseline_tree_corr'] for r in sig]))
 93    ci = float(np.mean([r['idea_tree_corr'] for r in sig]))
 94    report = make_report(TRACK, MODEL, base, idea, extra={
 95        'custom_track': {'name': TRACK, 'file': 'bench/custom_tracks/unordered_pointset_denoising.py',
 96                         'domain': 'point-set-diffusion'},
 97        'baseline_grid': BASE_GRID, 'idea_grid': IDEA_GRID, 'idea_trials': trials,
 98        'math_check': {'gram_tree_orthogonal_invariance_relative_error': inv_err},
 99        'mechanism_signature': {
100            'claim': 'tree aggregation preserves target set correspondence',
101            'predicted': 'tree-feature correlation is higher for the idea system',
102            'observed_baseline_mean_tree_corr': cb, 'observed_idea_mean_tree_corr': ci,
103            'per_seed': sig, 'confirmed': bool(ci > cb)
104        }})
105    Path('bench_report.json').write_text(json.dumps(report, indent=2))
106    print(json.dumps(report, indent=2))
107
108if __name__ == '__main__': main()