Dimension-Free Brenier Transport Layer / bench_experiment.py
Mechanism confirmed, baseline not beaten
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()