Monotone CDT autoencoder bottleneck / stage2_transport_bench.py
Mechanism confirmed, baseline not beaten
1import json, sys
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, train_model, evaluate, sweep_baseline, make_report
8
9TRACK = 'translated_density_transport'
10MODEL = 'shared_mlp'
11SEEDS = tuple(range(8))
12NTR, NTE, EPOCHS = 400, 100, 25
13GRID = [{'lr': 1e-3}, {'lr': 3e-3}, {'lr': 9e-3}]
14N = 64
15
16class SharedMLP(nn.Module):
17 def __init__(self, inp, out):
18 super().__init__()
19 self.net = nn.Sequential(nn.Linear(inp, 64), nn.Tanh(),
20 nn.Linear(64, 64), nn.Tanh(),
21 nn.Linear(64, out))
22 def forward(self, x): return self.net(x)
23
24def prepare(seed):
25 d = get_dataset(TRACK, seed, NTR, NTE)
26 # The bench adapter flattens custom field targets to [N*N,1]. Restore fields.
27 d['ytr'] = d['ytr'].reshape(NTR, N)
28 d['yte'] = d['yte'].reshape(NTE, N)
29 d['xtr'] = d['xtr'].reshape(NTR, N)
30 d['xte'] = d['xte'].reshape(NTE, N)
31 return d
32
33def quantiles(u):
34 # Inverse CDF at midpoint probabilities; all input fields have unit mass.
35 x = np.linspace(0., 1., N, dtype=np.float32)
36 p = (np.arange(N, dtype=np.float32) + .5) / N
37 out = np.empty_like(u)
38 dx = x[1] - x[0]
39 for i, a in enumerate(u):
40 c = np.cumsum(np.maximum(a, 0.)) * dx
41 c /= max(float(c[-1]), 1e-8)
42 out[i] = np.interp(p, np.r_[0., c], np.r_[x[0]-dx, x])
43 return out
44
45def decode_monotone(h, bw=.012):
46 # Positive increments give a valid monotone quantile map in [0,1].
47 delta = torch.nn.functional.softplus(h) + 1e-5
48 q = torch.cumsum(delta, dim=1) / delta.sum(dim=1, keepdim=True)
49 # Smooth push-forward of uniform reference mass through q.
50 x = torch.linspace(0., 1., N, device=h.device)[None, :, None]
51 z = (x - q[:, None, :]) / bw
52 kern = torch.exp(-.5*z*z) / (bw*np.sqrt(2*np.pi))
53 dens = kern.mean(dim=2)
54 mass = dens.mean(dim=1, keepdim=True)
55 dens = dens / torch.clamp(mass, min=1e-8)
56 return q, dens
57
58def baseline(seed, cfg, return_model=False):
59 torch.manual_seed(seed); np.random.seed(seed)
60 d = prepare(seed)
61 td = {k: torch.as_tensor(d[k], dtype=torch.float32) for k in ('xtr','ytr','xte','yte')}
62 td['task'] = 'regression'
63 net, metric, hist = train_model(SharedMLP(N, N), td, epochs=EPOCHS,
64 lr=cfg['lr'], batch=128, log=lambda *_: None)
65 if return_model: return float(metric), net, d
66 return float(metric)
67
68def idea(seed, cfg, return_model=False):
69 torch.manual_seed(seed); np.random.seed(seed)
70 d = prepare(seed)
71 qtr, qte = quantiles(d['xtr']), quantiles(d['xte'])
72 ytr, yte = d['ytr'], d['yte']
73 td = {'xtr':torch.tensor(qtr), 'ytr':torch.tensor(qtr),
74 'xte':torch.tensor(qte), 'yte':torch.tensor(qte), 'task':'regression'}
75 net, _, hist = train_model(SharedMLP(N, N), td, epochs=EPOCHS,
76 lr=cfg['lr'], batch=128, log=lambda *_: None)
77 dev = next(net.parameters()).device
78 with torch.no_grad():
79 qin = torch.tensor(qte, device=dev)
80 predq, pred = decode_monotone(net(qin))
81 pred, predq = pred.cpu(), predq.cpu()
82 metric = float(((pred - torch.tensor(yte))**2).mean())
83 if return_model: return metric, net, d, qte, predq, pred
84 return metric
85
86def main():
87 # Union parity: every idea lr is also evaluated by baseline sweep.
88 base = sweep_baseline(lambda cfg: lambda s: baseline(s, cfg), GRID, seeds=tuple(range(4)))
89 idea_sweep=[]
90 for cfg in GRID:
91 r=evaluate(lambda s, c=cfg: idea(s,c), seeds=SEEDS)
92 idea_sweep.append({'cfg':cfg, **r})
93 best=min(idea_sweep, key=lambda z:z['mean'])
94 cfg=best['cfg']
95 # Full paired rerun at the selected common setting and trained-model signature.
96 bfull=evaluate(lambda s: baseline(s,cfg), seeds=SEEDS)
97 ifull=evaluate(lambda s: idea(s,cfg), seeds=SEEDS)
98 paired=[]; mono=[]; neg=[]; masserr=[]; shift=[]; bneg=[]
99 for s in SEEDS:
100 bm,bnet,bd=baseline(s,cfg,True)
101 im,inet,idata,qte,pq,pred=idea(s,cfg,True)
102 with torch.no_grad():
103 bp=bnet(torch.tensor(idata['xte'], device=next(bnet.parameters()).device)).cpu()
104 paired.append({'seed':s,'baseline':bm,'idea':im})
105 bneg.append(float((bp<0).float().mean()))
106 mono.append(float(torch.all(torch.diff(pq,dim=1)>=0,dim=1).float().mean()))
107 neg.append(float((pred<0).float().mean()))
108 masserr.append(float(torch.abs(pred.mean(1)-1.).mean()))
109 shift.append(float((pq-torch.tensor(qte)).mean()))
110 extra={'trained_models':True,
111 'prediction':'translated states should produce monotone quantiles with approximately constant +0.08 displacement, nonnegative decoded densities, and unit mass',
112 'predicted_vs_observed':{'expected_quantile_shift':0.08,
113 'observed_mean_quantile_shift':float(np.mean(shift)),
114 'observed_monotone_fraction':float(np.mean(mono)),
115 'observed_idea_negative_fraction':float(np.mean(neg)),
116 'observed_idea_mean_mass_proxy_error':float(np.mean(masserr)),
117 'observed_baseline_negative_fraction':float(np.mean(bneg))},
118 'confirmed':bool(np.mean(mono)>0.99 and np.mean(neg)<1e-8 and np.mean(masserr)<1e-5 and abs(np.mean(shift)-.08)<.03),
119 'paired_values':paired,
120 'idea_sweep':idea_sweep}
121 # make_report computes paired delta and permutation p-value.
122 rep=make_report(TRACK,MODEL,{'best_cfg':base['best_cfg'],'sweep':base['sweep'],'full':bfull},ifull,extra)
123 rep['idea']['best_cfg']=cfg; rep['idea']['sweep']=idea_sweep
124 Path('bench_report.json').write_text(json.dumps(rep,indent=2))
125 print(json.dumps(rep,indent=2))
126
127if __name__=='__main__': main()