Dimension-Free Brenier Transport Layer / bench_experiment.py

Mechanism confirmed, baseline not beaten

Raw ⬇ ZIP
  1import sys, json, math, random
  2from pathlib import Path
  3import numpy as np
  4import torch
  5import torch.nn as nn
  6
  7sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
  8from bench import get_dataset, make_model, train_model, sweep_baseline, make_report
  9from bench.protocol import DEFAULT_SEEDS
 10
 11SEEDS = tuple(range(8))
 12EPOCHS = 12
 13BATCH = 128
 14# Union of all learning rates is used for both systems (search-space parity).
 15LR_GRID = [1e-3, 3e-3, 1e-2]
 16CAPS = [0.587 * 2.0 * math.sqrt(3.0), 0.587 * 2.0 * math.sqrt(3.0) * 0.75,
 17        0.587 * 2.0 * math.sqrt(3.0) * 1.25]
 18
 19
 20def seed_all(s):
 21    random.seed(s); np.random.seed(s); torch.manual_seed(s)
 22    if torch.cuda.is_available(): torch.cuda.manual_seed_all(s)
 23
 24
 25def jacobian_norm(net, x, max_n=24):
 26    """Observed operator norm of d(output)/d(input), on trained model probes."""
 27    net.eval(); dev = next(net.parameters()).device; x = x[:max_n].to(dev).detach().clone().requires_grad_(True)
 28    vals = []
 29    for i in range(len(x)):
 30        g = torch.autograd.grad(net(x[i:i+1]).sum(), x, retain_graph=True,
 31                                create_graph=False, allow_unused=False)[0][i]
 32        vals.append(float(torch.linalg.vector_norm(g).detach().cpu()))
 33    return float(max(vals)) if vals else float('nan')
 34
 35
 36def baseline_one(seed, lr, keep_model=False):
 37    seed_all(seed)
 38    ds = get_dataset('dynamics', seed, n_train=4000, n_test=1000)
 39    model = make_model('rnn_small', ds['input_shape'], ds['out_dim'])
 40    net, metric, hist = train_model(model, ds, epochs=EPOCHS, lr=lr, batch=BATCH,
 41                                    weight_decay=0.0, log=lambda *_: None)
 42    if net is None: return {'metric': float('nan'), 'jacobian': float('nan')}
 43    j = jacobian_norm(net, ds['xte'])
 44    return {'metric': float(metric), 'jacobian': j, 'final_train_loss': float(hist[-1])}
 45
 46
 47def capped_one(seed, lr, cap, keep_model=False):
 48    seed_all(seed)
 49    ds = get_dataset('dynamics', seed, n_train=4000, n_test=1000)
 50    # Same rnn_small architecture and Adam budget as baseline; only intervention differs.
 51    net = make_model('rnn_small', ds['input_shape'], ds['out_dim'])
 52    device = 'cuda' if torch.cuda.is_available() else 'cpu'
 53    try:
 54        net = net.to(device)
 55        xtr, ytr = ds['xtr'].to(device), ds['ytr'].to(device)
 56        opt = torch.optim.Adam(net.parameters(), lr=lr)
 57        lossf = nn.MSELoss(); violations = 0; peak = 0.0
 58        for _ in range(EPOCHS):
 59            net.train(); perm = torch.randperm(len(xtr), device=device)
 60            for i in range(0, len(xtr), BATCH):
 61                idx = perm[i:i+BATCH]
 62                loss = lossf(net(xtr[idx]), ytr[idx])
 63                opt.zero_grad(); loss.backward(); opt.step()
 64                # Certificate projection: scale all parameters if observed probe Jacobian exceeds L*.
 65                # This is deliberately a training-time update, not an alternate readout.
 66                if i != 0: continue
 67                probe = xtr[idx[:8]].detach().clone().requires_grad_(True)
 68                j = jacobian_norm(net, probe, max_n=8)
 69                peak = max(peak, j)
 70                if j > cap:
 71                    violations += 1
 72                    with torch.no_grad():
 73                        scale = math.sqrt(cap / max(j, 1e-12))
 74                        for p in net.parameters(): p.mul_(scale)
 75        net.eval()
 76        with torch.no_grad(): metric = float(((net(ds['xte'].to(device)) - ds['yte'].to(device))**2).mean())
 77        jtest = jacobian_norm(net, ds['xte'].to(device))
 78        return {'metric': metric, 'jacobian': jtest, 'train_probe_peak': peak,
 79                'cap': cap, 'cap_events': violations}
 80    except Exception:
 81        # explicit CPU fallback, matching the harness requirement
 82        return capped_cpu(seed, lr, cap)
 83
 84
 85def capped_cpu(seed, lr, cap):
 86    seed_all(seed); ds = get_dataset('dynamics', seed, 4000, 1000)
 87    net = make_model('rnn_small', ds['input_shape'], ds['out_dim']).cpu()
 88    opt = torch.optim.Adam(net.parameters(), lr=lr); lossf = nn.MSELoss(); peak=0.; events=0
 89    for _ in range(EPOCHS):
 90        perm=torch.randperm(len(ds['xtr']))
 91        for i in range(0,len(perm),BATCH):
 92            idx=perm[i:i+BATCH]; loss=lossf(net(ds['xtr'][idx]),ds['ytr'][idx])
 93            opt.zero_grad(); loss.backward(); opt.step()
 94            j=jacobian_norm(net,ds['xtr'][idx[:8]]); peak=max(peak,j)
 95            if j>cap:
 96                events+=1
 97                with torch.no_grad():
 98                    for p in net.parameters(): p.mul_(math.sqrt(cap/j))
 99    with torch.no_grad(): m=float(((net(ds['xte'])-ds['yte'])**2).mean())
100    return {'metric':m,'jacobian':jacobian_norm(net,ds['xte']),'train_probe_peak':peak,'cap':cap,'cap_events':events}
101
102
103def main():
104    # Baseline sweep on four seeds; all three rates are explicitly evaluated on baseline.
105    sweep = {}
106    for lr in LR_GRID:
107        vals=[baseline_one(s,lr)['metric'] for s in (0,1,2,3)]
108        sweep[str(lr)]={'lr':lr,'mean_metric':float(np.nanmean(vals)), 'per_seed':vals}
109    best_lr=min(LR_GRID,key=lambda z:sweep[str(z)]['mean_metric'])
110    # Idea sweep over the same three lrs and three a-priori certificate multipliers.
111    idea_cfg=[]
112    for lr in LR_GRID:
113        for cap in CAPS:
114            vals=[capped_one(s,lr,cap)['metric'] for s in (0,1,2,3)]
115            idea_cfg.append({'lr':lr,'cap':cap,'mean_metric':float(np.nanmean(vals)), 'per_seed':vals})
116    best=min(idea_cfg,key=lambda z:z['mean_metric'])
117    base_full=[baseline_one(s,best_lr) for s in SEEDS]
118    idea_full=[capped_one(s,best['lr'],best['cap']) for s in SEEDS]
119    base_block={'best':{'lr':best_lr,'mean_metric':sweep[str(best_lr)]['mean_metric']},
120                'sweep':sweep,'full':{'per_seed':[x['metric'] for x in base_full],
121                'details':base_full,'config':{'lr':best_lr,'epochs':EPOCHS}}}
122    idea_res={'per_seed':[x['metric'] for x in idea_full], 'details':idea_full,
123              'config':best}
124    # Signature is measured from the trained systems: predicted cap behavior vs observed norms.
125    pred=float(best['cap']); obs=float(np.nanmax([x['jacobian'] for x in idea_full]))
126    bobs=float(np.nanmax([x['jacobian'] for x in base_full]))
127    sig={'prediction': 'certificate cap limits Jacobian operator norm',
128         'predicted_max_jacobian':pred,'observed_idea_max_jacobian':obs,
129         'observed_baseline_max_jacobian':bobs,
130         'confirmed': bool(np.isfinite(obs) and obs <= pred*1.10)}
131    report=make_report('dynamics','rnn_small',base_block,idea_res,
132                       {'mechanism_signature':sig,
133                        'budget':{'epochs':EPOCHS,'batch':BATCH,'lr_grid':LR_GRID},
134                        'selection':{'idea_best':best}})
135    Path('bench_report.json').write_text(json.dumps(report,indent=2))
136    print(json.dumps(report,indent=2))
137
138if __name__=='__main__': main()