OT-Sufficient Bottleneck Flow Matching / ot_registered_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
 7from bench import sweep_baseline, make_report, custom_tracks
 8
 9def getds(seed, ntr=400, nte=400):
10    d=custom_tracks()['conditional_multitoken_diffusion'].get_dataset(seed,ntr,nte)
11    return {k: torch.as_tensor(d[k],dtype=torch.float32) for k in ('xtr','ytr','xte','yte')} | {'seed':seed}
12
13class Base(nn.Module):
14    def __init__(self):
15        super().__init__(); self.enc=nn.Sequential(nn.Linear(10,32),nn.Tanh(),nn.Linear(32,1)); self.head=nn.Sequential(nn.Linear(1,32),nn.Tanh(),nn.Linear(32,8))
16    def forward(self,x): return self.head(self.enc(x))
17class Flow(nn.Module):
18    def __init__(self):
19        super().__init__(); self.enc=nn.Sequential(nn.Linear(10,32),nn.Tanh(),nn.Linear(32,1)); self.v=nn.Sequential(nn.Linear(10,32),nn.Tanh(),nn.Linear(32,8))
20    def forward(self,t,y,z): return self.v(torch.cat([t,y,z],-1))
21def seed(s):
22    random.seed(s+9000); np.random.seed(s+9000); torch.manual_seed(s+9000)
23    if torch.cuda.is_available(): torch.cuda.manual_seed_all(s+9000)
24def sh(cost,eps=.16,iters=30):
25    n,m=cost.shape; lk=-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)
26    for _ in range(iters): u=la-torch.logsumexp(lk+v[None],1); v=lb-torch.logsumexp(lk+u[:,None],0)
27    return torch.exp(lk+u[:,None]+v[None])
28def run_dev(fn):
29    try: return fn(torch.device('cuda' if torch.cuda.is_available() else 'cpu'))
30    except RuntimeError: return fn(torch.device('cpu'))
31def base(ds,lr,epochs=12,ret=False):
32    seed(ds['seed'])
33    def go(dev):
34        m=Base().to(dev); x=ds['xtr'].to(dev); y=ds['ytr'].to(dev); o=torch.optim.Adam(m.parameters(),lr=lr)
35        for _ in range(epochs):
36            for ix in torch.randperm(len(x),device=dev).split(64):
37                l=((m(x[ix])-y[ix])**2).mean();o.zero_grad();l.backward();o.step()
38        with torch.no_grad(): z=m(ds['xte'].to(dev)); metric=float(((z-ds['yte'].to(dev))**2).mean())
39        return (metric,m) if ret else metric
40    return run_dev(go)
41def idea(ds,lr,lam,epochs=12,ret=False):
42    seed(ds['seed'])
43    def go(dev):
44        m=Flow().to(dev); x=ds['xtr'].to(dev); y=ds['ytr'].to(dev); o=torch.optim.Adam(m.parameters(),lr=lr)
45        for _ in range(epochs):
46            for ix in torch.randperm(len(x),device=dev).split(64):
47                xb,yb=x[ix],y[ix]; z=m.enc(xb); y0=torch.randn_like(yb); C=((y0[:,None,:]-yb[None,:,:])**2).mean(-1); R=(z-z.T)**2; R=R/(R.mean().detach()+1e-6); P=sh(C+lam*R).detach(); n=len(xb)
48                t=torch.rand(n,n,1,device=dev); yt=(1-t)*y0[:,None,:]+t*yb[None,:,:]; u=yb[None,:,:]-y0[:,None,:]; zz=z[None,:,:].expand(n,n,1); pred=m.v(torch.cat([t,yt,zz],-1)); l=(P[:,:,None]*(pred-u)**2).sum();o.zero_grad();l.backward();o.step()
49        with torch.no_grad():
50            xe=ds['xte'].to(dev); z=m.enc(xe); a=torch.randn(len(xe),8,device=dev)
51            for k in range(25): a=a+m(torch.full((len(xe),1), (k+.5)/25,device=dev),a,z)/25
52            metric=float(((a-ds['yte'].to(dev))**2).mean())
53        return (metric,m) if ret else metric
54    return run_dev(go)
55def ev(fn,seeds=tuple(range(8))):
56    a=[float(fn(s)) for s in seeds]; return {'per_seed':a,'mean':float(np.mean(a)),'std':float(np.std(a,ddof=1))}
57def main():
58    seeds=tuple(range(8)); lrs=[.001,.003,.01]; grid=[{'lr':q,'epochs':12} for q in lrs]
59    def mk(c): return lambda s: base(dict(getds(s,400,400),seed=s),c['lr'],c['epochs'])
60    b=sweep_baseline(mk,grid,seeds=tuple(range(4))); lr=b['best_cfg']['lr']; rs=[]
61    for lam in [0.,2.,8.]: rs.append((ev(lambda s,lam=lam: idea(dict(getds(s,400,400),seed=s),lr,lam)),lam))
62    ir,lam=min(rs,key=lambda q:q[0]['mean'])
63    _,m=idea(dict(getds(0,400,400),seed=0),lr,lam,ret=True); m.eval(); d=getds(77,400,400); dev=next(m.parameters()).device
64    with torch.no_grad():
65        z=m.enc(d['xte'].to(dev)); samples=[]
66        for _ in range(8):
67            a=torch.randn(400,8,device=dev)
68            for k in range(25): a=a+m(torch.full((400,1),(k+.5)/25,device=dev),a,z)/25
69            samples.append(a.cpu().numpy())
70    pred=np.stack(samples); obs=d['yte'].numpy(); sig={'observed_token_std':float(obs.std()),'predicted_token_std':float(pred.std()),'observed_sample_mean_std':float(obs.mean(1).std()),'predicted_sample_mean_std':float(pred.mean(2).std()),'confirmed':bool(abs(pred.std()-obs.std())<.25)}
71    extra={'mechanism_signature':sig,'idea_lambda':lam,'idea_runs':[{'lambda':q,'result':r} for r,q in rs]}
72    rep=make_report('conditional_multitoken_diffusion','shared_encoder_mlp',b,ir,extra)
73    json.dump(rep,open('bench_report.json','w'),indent=2); print(json.dumps(rep,indent=2))
74if __name__=='__main__': main()