Monotone CDT autoencoder bottleneck / stage2_transport_bench.py

Mechanism confirmed, baseline not beaten

Raw ⬇ ZIP
  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()