Monotone CDT autoencoder bottleneck / experiment.py

Mechanism confirmed, baseline not beaten

Raw ⬇ ZIP
 1import json
 2import numpy as np
 3
 4SEED=1051
 5rng=np.random.default_rng(SEED)
 6
 7
 8def monotone_quantile(h,qmin=0.,qmax=1.,eps=1e-8):
 9    h=np.asarray(h)
10    delta=np.logaddexp(0.,h)+eps
11    return qmin+(qmax-qmin)*np.cumsum(delta,axis=-1)/delta.sum(axis=-1,keepdims=True)
12
13
14def pushforward_density(q,p,edges):
15    # q and p include both endpoints; each p interval carries exactly its mass.
16    mass=np.zeros(len(edges)-1)
17    for a,b,m in zip(q[:-1],q[1:],np.diff(p)):
18        if b<=a: continue
19        lo,hi=max(a,edges[0]),min(b,edges[-1])
20        if hi<=lo: continue
21        j0=max(0,np.searchsorted(edges,lo,side='right')-1)
22        j1=min(len(mass)-1,np.searchsorted(edges,hi,side='left'))
23        for j in range(j0,j1+1):
24            overlap=max(0.,min(hi,edges[j+1])-max(lo,edges[j]))
25            mass[j]+=m*overlap/(b-a)
26    return mass/np.diff(edges)
27
28
29def invariant_checks():
30    K=64; H=rng.normal(size=(2000,K)); out={}
31    # Prediction 1: positivity of increments implies monotonicity for every logit scale.
32    temps=[.05,.2,1.,5.,20.]; rows=[]
33    for t in temps:
34        q=monotone_quantile(t*H)
35        rows.append({'temperature':t,'monotone_fraction':float(np.mean(np.all(np.diff(q,axis=1)>=0,axis=1))), 'min_increment':float(np.min(np.diff(q,axis=1)))})
36    out['temperature_sweep']=rows
37    # Prediction 2: normalization fixes the final span independent of logits.
38    q=monotone_quantile(H,qmin=-2,qmax=3)
39    out['span_prediction']={'predicted_final_endpoint':3.,'observed_final_endpoint':float(np.mean(q[:,-1])),'endpoint_error':float(np.max(np.abs(q[:,-1]-3)))}
40    # Prediction 3: every increment is >= L*eps/sum(delta); increasing eps raises this bound.
41    epses=[1e-10,1e-7,1e-4,1e-2]; pred=[]; obs=[]
42    for e in epses:
43        delta=np.logaddexp(0.,H)+e; q=monotone_quantile(H,eps=e)
44        pred.append(float(np.min(e/delta.sum(axis=1))))
45        obs.append(float(np.min(np.diff(q,axis=1))))
46    out['epsilon_sweep']={'eps':epses,'predicted_bound':pred,'observed_min_increment':obs,'bound_holds':bool(all(o+1e-15>=p for o,p in zip(obs,pred)))}
47    p=np.linspace(0,1,K); edges=np.linspace(0,1,257); u=pushforward_density(monotone_quantile(H[0]),p,edges)
48    out['pushforward']={'negative_cells':int((u<0).sum()),'mass':float(np.sum(u*np.diff(edges))),'mass_error':float(abs(np.sum(u*np.diff(edges))-1.))}
49    return out
50
51
52def fields(n=800,nx=128):
53    x=np.linspace(0,1,nx); ans=[]
54    for _ in range(n):
55        c=rng.uniform(.15,.85); w=rng.uniform(.025,.10)
56        y=np.exp(-.5*((x-c)/w)**2)+.35*np.exp(-.5*((x-(c-.14))/(.6*w))**2)
57        y=np.maximum(y,0); y/=np.trapz(y,x); ans.append(y)
58    return np.asarray(ans)
59
60
61def quantiles(U,K):
62    x=np.linspace(0,1,U.shape[1]); dx=x[1]-x[0]; p=np.linspace(0,1,K); out=[]
63    for u in U:
64        c=np.cumsum(u)*dx; c=(c-c[0])/(c[-1]-c[0])
65        out.append(np.interp(p,c,x))
66    return np.asarray(out)
67
68
69def experiment():
70    U=fields(); split=600; trainU,testU=U[:split],U[split:]; Q=quantiles(U,64); trainQ,testQ=Q[:split],Q[split:]
71    meanQ=trainQ.mean(0); _,_,VQ=np.linalg.svd(trainQ-meanQ,full_matrices=False)
72    meanU=trainU.mean(0); _,_,VU=np.linalg.svd(trainU-meanU,full_matrices=False)
73    edges=np.linspace(0,1,U.shape[1]+1); p=np.linspace(0,1,64); rows=[]
74    for d in [2,4,8]:
75        # Standard reduced decoder: unconstrained physical POD field.
76        base=meanU+(testU-meanU)@VU[:d].T@VU[:d]
77        # Proposed CDT/POD decoder, with cumulative maximum as a robust monotone projection.
78        qraw=meanQ+(testQ-meanQ)@VQ[:d].T@VQ[:d]
79        qidea=np.clip(np.maximum.accumulate(qraw,axis=1),0,1)
80        rec=np.asarray([pushforward_density(q,p,edges) for q in qidea])
81        rows.append({'d':d,'baseline_physical_L1':float(np.mean(np.abs(base-testU))), 'idea_physical_L1':float(np.mean(np.abs(rec-testU))), 'baseline_wasserstein_proxy':float(np.mean(np.abs(quantiles(np.maximum(base,0),64)-testQ))), 'idea_wasserstein':float(np.mean(np.abs(qidea-testQ))), 'baseline_negative_fraction':float(np.mean(base<0)), 'idea_negative_fraction':float(np.mean(rec<0)), 'baseline_mass_abs_error':float(np.mean(np.abs(np.trapz(base,dx=1/127,axis=1)-1))), 'idea_mass_abs_error':float(np.mean(np.abs(np.sum(rec,axis=1)*(1/128)-1))), 'idea_monotone_fraction':float(np.mean(np.all(np.diff(qidea,axis=1)>=-1e-12,axis=1)))})
82    return rows
83
84if __name__=='__main__':
85    report={'seed':SEED,'checks':invariant_checks(),'experiment':experiment()}
86    with open('results.json','w') as f: json.dump(report,f,indent=2)
87    print(json.dumps(report,indent=2))