OT-Sufficient Bottleneck Flow Matching / ot_sufficient_bench.py

Mechanism confirmed, baseline not beaten

Raw ⬇ ZIP
  1import sys
  2sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
  3import json, math, random
  4import numpy as np
  5import torch
  6import torch.nn as nn
  7
  8META = {"name": "conditional_bimodal_regression", "domain": "probabilistic_regression", "description": "A scalar covariate with nuisance features and a bimodal conditional target; tests preservation of conditional laws."}
  9
 10def get_dataset(seed, n_train=400, n_test=400):
 11    rng = np.random.RandomState(seed)
 12    def make(n):
 13        s = rng.uniform(-2, 2, n).astype(np.float32)
 14        nuisance = rng.normal(0, 1, n).astype(np.float32)
 15        p = (0.5 + 0.22*np.sin(1.5*s)).astype(np.float32)
 16        branch = (rng.rand(n) < p).astype(np.float32)*2-1
 17        scale = 0.10 + 0.025*np.abs(s)
 18        y = branch*(0.9 + .28*s) + rng.normal(0, scale).astype(np.float32)
 19        return np.stack([s, nuisance], 1), y[:, None].astype(np.float32)
 20    xtr,ytr=make(n_train); xte,yte=make(n_test)
 21    return {"xtr":xtr,"ytr":ytr,"xte":xte,"yte":yte,"task":"regression","metric":"mse","out_dim":1}
 22
 23def seed_all(seed):
 24    random.seed(seed); np.random.seed(seed); torch.manual_seed(seed)
 25    if torch.cuda.is_available(): torch.cuda.manual_seed_all(seed)
 26
 27class Baseline(nn.Module):
 28    def __init__(self):
 29        super().__init__(); self.enc=nn.Sequential(nn.Linear(2,32),nn.Tanh(),nn.Linear(32,1)); self.head=nn.Sequential(nn.Linear(1,32),nn.Tanh(),nn.Linear(32,1))
 30    def forward(self,x): return self.head(self.enc(x))
 31
 32class OTFlow(nn.Module):
 33    def __init__(self):
 34        super().__init__(); self.enc=nn.Sequential(nn.Linear(2,32),nn.Tanh(),nn.Linear(32,1)); self.v=nn.Sequential(nn.Linear(3,32),nn.Tanh(),nn.Linear(32,1))
 35    def forward(self,t,y,z): return self.v(torch.cat([t,y,z],1))
 36
 37def sinkhorn(cost, eps=.16, iters=35):
 38    n,m=cost.shape; logk=-cost/eps; la=torch.full((n,),-math.log(n),device=cost.device); lb=torch.full((m,),-math.log(m),device=cost.device); u=torch.zeros_like(la); v=torch.zeros_like(lb)
 39    for _ in range(iters):
 40        u=la-torch.logsumexp(logk+v[None,:],1); v=lb-torch.logsumexp(logk+u[:,None],0)
 41    return torch.exp(logk+u[:,None]+v[None,:])
 42
 43def device_run(fn):
 44    try: return fn(torch.device("cuda" if torch.cuda.is_available() else "cpu"))
 45    except RuntimeError: return fn(torch.device("cpu"))
 46
 47def train_base(ds, lr, epochs=12, return_model=False):
 48    seed_all(ds.get("seed",0)+9000)
 49    def run(dev):
 50        m=Baseline().to(dev); opt=torch.optim.Adam(m.parameters(),lr=lr); x=torch.as_tensor(ds['xtr'],device=dev); y=torch.as_tensor(ds['ytr'],device=dev)
 51        for _ in range(epochs):
 52            for ix in torch.randperm(len(x),device=dev).split(64):
 53                loss=((m(x[ix])-y[ix])**2).mean(); opt.zero_grad(); loss.backward(); opt.step()
 54        with torch.no_grad(): metric=float(((m(torch.as_tensor(ds['xte'],device=dev))-torch.as_tensor(ds['yte'],device=dev))**2).mean())
 55        return (metric,m) if return_model else metric
 56    return device_run(run)
 57
 58def train_idea(ds, lr, lam, eps=.16, epochs=12, return_model=False):
 59    seed_all(ds.get("seed",0)+9000)
 60    def run(dev):
 61        m=OTFlow().to(dev); opt=torch.optim.Adam(m.parameters(),lr=lr); x=torch.as_tensor(ds['xtr'],device=dev); y=torch.as_tensor(ds['ytr'],device=dev)
 62        for _ in range(epochs):
 63            for ix in torch.randperm(len(x),device=dev).split(64):
 64                xb,yb=x[ix],y[ix]; z=m.enc(xb); y0=torch.randn_like(yb); R=(z-z.T).pow(2); R=R/(R.mean().detach()+1e-6); C=(y0-yb.T).pow(2); P=sinkhorn(C+lam*R,eps).detach()
 65                n=len(xb); t=torch.rand(n,n,1,device=dev); yt=(1-t)*y0[:,None,:]+t*yb[None,:,:]; u=yb[None,:,:]-y0[:,None,:]; pred=m.v(torch.cat([t.expand(n,n,1),yt,z[None,:,:].expand(n,n,1)],-1)); loss=(P[:,:,None]*(pred-u)**2).sum()
 66                opt.zero_grad(); loss.backward(); opt.step()
 67        with torch.no_grad():
 68            xe=torch.as_tensor(ds['xte'],device=dev); z=m.enc(xe); sample=torch.randn_like(z)
 69            for k in range(25): sample=sample+m(torch.full_like(sample,(k+.5)/25),sample,z)/25
 70            metric=float(((sample-torch.as_tensor(ds['yte'],device=dev))**2).mean())
 71        return (metric,m) if return_model else metric
 72    return device_run(run)
 73
 74def evaluate(fn,seeds=tuple(range(8))):
 75    vals=[float(fn(s)) for s in seeds]; return {"per_seed":vals,"mean":float(np.mean(vals)),"std":float(np.std(vals,ddof=1))}
 76
 77def main():
 78    from bench import sweep_baseline, make_report
 79    seeds=tuple(range(8)); lrs=[1e-3,3e-3,1e-2]
 80    def mk(cfg): return lambda s: train_base(dict(get_dataset(s),seed=s),cfg['lr'],cfg['epochs'])
 81    base=sweep_baseline(mk,[{'lr':lr,'epochs':12} for lr in lrs],seeds=tuple(range(4)))
 82    bestlr=base['best_cfg']['lr']; runs=[]
 83    for lam in [0.,2.,8.]:
 84        r=evaluate(lambda s,lam=lam: train_idea(dict(get_dataset(s),seed=s),bestlr,lam,epochs=12),seeds); runs.append((r,lam))
 85    idea,lam=min(runs,key=lambda q:q[0]['mean'])
 86    metric,m=train_idea(dict(get_dataset(0),seed=0),bestlr,lam,epochs=12,return_model=True)
 87    m.eval(); d=get_dataset(77,400,400); xe=torch.tensor(d['xte']); yt=d['yte'][:,0]; dev=next(m.parameters()).device
 88    with torch.no_grad():
 89        z=m.enc(xe.to(dev)); ys=[]
 90        for q in range(8):
 91            a=torch.randn_like(z)
 92            for k in range(25): a=a+m(torch.full_like(a,(k+.5)/25),a,z)/25
 93            ys.append(a[:,0].cpu().numpy())
 94        yp=np.stack(ys); s=xe[:,0].numpy(); bins=np.digitize(s,[-1,0,1]); obs=[]; pred=[]
 95        for b in range(1,4):
 96            mask=bins==b; obs.append(float(np.std(yt[mask]))); pred.append(float(np.std(yp[:,mask])))
 97    sig={"observed_conditional_std_by_bin":obs,"predicted_conditional_std_by_bin":pred,"predicted_mean_std":float(np.mean(pred)),"observed_mean_std":float(np.mean(obs)),"confirmed":bool(abs(np.mean(pred)-np.mean(obs))<0.35)}
 98    extra={"custom_track":{"name":META['name'],"file":"ot_sufficient_bench.py","domain":META['domain']},"mechanism_signature":sig,"idea_lambda":lam,"idea_runs":[{"lambda":q,"result":r} for r,q in runs]}
 99    rep=make_report('conditional_bimodal_regression','shared_encoder_mlp',base,idea,extra)
100    with open('bench_report.json','w') as f: json.dump(rep,f,indent=2)
101    print(json.dumps(rep,indent=2))
102if __name__=='__main__': main()