OT-Sufficient Bottleneck Flow Matching / ot_sufficient_flow.py

Mechanism confirmed, baseline not beaten

Raw ⬇ ZIP
  1import json, math, random
  2import numpy as np
  3import torch
  4import torch.nn as nn
  5
  6SEED = 1122
  7np.random.seed(SEED); random.seed(SEED); torch.manual_seed(SEED)
  8torch.set_num_threads(4)
  9device = 'cuda' if torch.cuda.is_available() else 'cpu'
 10
 11def sinkhorn(cost, eps=0.12, iters=120):
 12    # Uniform marginals; log-domain updates are stable for this small toy problem.
 13    n, m = cost.shape
 14    logK = -cost / eps
 15    loga = torch.full((n,), -math.log(n), device=cost.device)
 16    logb = torch.full((m,), -math.log(m), device=cost.device)
 17    u = torch.zeros_like(loga); v = torch.zeros_like(logb)
 18    for _ in range(iters):
 19        u = loga - torch.logsumexp(logK + v[None, :], dim=1)
 20        v = logb - torch.logsumexp(logK + u[:, None], dim=0)
 21    return torch.exp(logK + u[:, None] + v[None, :])
 22
 23def sample_xy(n):
 24    s = torch.rand(n, device=device) * 4 - 2
 25    nuisance = torch.randn(n, device=device)
 26    # Bimodal conditional law, with a mild heteroscedastic component.
 27    sign = torch.where(torch.rand(n, device=device) < .5, -1., 1.)
 28    y = sign * (1.0 + .30*s) + (.12 + .025*s.abs()) * torch.randn(n, device=device)
 29    return torch.stack([s, nuisance], 1), s[:, None], y[:, None]
 30
 31class Flow(nn.Module):
 32    def __init__(self):
 33        super().__init__()
 34        self.net = nn.Sequential(nn.Linear(3,48), nn.Tanh(), nn.Linear(48,48), nn.Tanh(), nn.Linear(48,1))
 35    def forward(self,t,y,z): return self.net(torch.cat([t,y,z],1))
 36
 37class MeanHead(nn.Module):
 38    def __init__(self):
 39        super().__init__(); self.net=nn.Sequential(nn.Linear(1,32),nn.Tanh(),nn.Linear(32,1))
 40    def forward(self,z): return self.net(z)
 41
 42def train_flow(lam, steps=260, n=64):
 43    model=Flow().to(device); opt=torch.optim.Adam(model.parameters(), lr=2e-3)
 44    for step in range(steps):
 45        _, z, y = sample_xy(n); y0=torch.randn_like(y)
 46        # R is normalized representation-locality cost as in the proposal.
 47        R=(z-z.T).pow(2); R=R/(R.mean().detach()+1e-6)
 48        C=(y0-y.T).pow(2)
 49        P=sinkhorn(C + lam*R, eps=.16, iters=55).detach()
 50        t=torch.rand(n,n,device=device)
 51        yt=(1-t)*y0 + t*y.T
 52        u=y.T-y0
 53        pred=model(t.reshape(-1,1),yt.reshape(-1,1),z.T.repeat(n,1).reshape(-1,1))
 54        loss=(P.reshape(-1,1)*(pred-u.reshape(-1,1)).pow(2)).sum()
 55        opt.zero_grad(); loss.backward(); opt.step()
 56    return model
 57
 58def train_mean(steps=260,n=64):
 59    model=MeanHead().to(device); opt=torch.optim.Adam(model.parameters(),lr=3e-3)
 60    for _ in range(steps):
 61        _,z,y=sample_xy(n); loss=(model(z)-y).pow(2).mean()
 62        opt.zero_grad();loss.backward();opt.step()
 63    return model
 64
 65def true_samples(s, n=1600):
 66    s=torch.full((n,),float(s),device=device); sign=torch.where(torch.arange(n,device=device)%2==0,-1.,1.)
 67    return (sign*(1+.30*s)+(.12+.025*abs(s))*torch.randn(n,device=device)).cpu().numpy()
 68
 69def w1(a,b):
 70    a=np.sort(a); b=np.sort(b); return float(np.mean(np.abs(a-b)))
 71
 72def evaluate(flow, mean):
 73    grid=[-1.5,-.5,.5,1.5]; wf=[]; wm=[]; cover=[]
 74    for s in grid:
 75        z=torch.full((512,1),s,device=device); target=true_samples(s,512)
 76        with torch.no_grad():
 77            y=torch.randn(512,1,device=device)
 78            for k in range(30):
 79                t=torch.full_like(y,(k+.5)/30)
 80                y=y+flow(t,y,z)/30
 81            yp=y[:,0].cpu().numpy(); mp=mean(z)[:,0].cpu().numpy()
 82        wf.append(w1(yp,target)); wm.append(w1(mp,target))
 83        # fraction in either true mode neighborhood, a simple mode-coverage proxy
 84        cover.append(float(((yp < -.35) | (yp > .35)).mean()))
 85    return {'flow_w1':float(np.mean(wf)),'mse_w1':float(np.mean(wm)),
 86            'flow_mode_coverage':float(np.mean(cover)), 'per_s_w1':wf}
 87
 88def main():
 89    # Prediction 1: lambda=0 has no representation-locality pressure.
 90    # Prediction 2: increasing lambda lowers paired representation distance.
 91    # Prediction 3: increasing epsilon washes out locality (higher paired distance).
 92    z=torch.linspace(-2,2,40,device=device)[:,None]; y0=torch.randn(40,1,device=device)
 93    rows=[]
 94    for eps in [.06,.16,.40]:
 95        for lam in [0.,.5,2.,8.]:
 96            R=(z-z.T).pow(2); R=R/(R.mean()+1e-6); C=(y0-y0.T).pow(2)
 97            P=sinkhorn(C+lam*R,eps=eps,iters=180)
 98            locality=float((P*R).sum().cpu()); marginal=float(max((P.sum(1)-1/40).abs().max(),(P.sum(0)-1/40).abs().max()).cpu())
 99            rows.append({'epsilon':eps,'lambda':lam,'paired_R':locality,'marginal_err':marginal})
100    # Exact 1D quadratic OT check: sorted coupling is no worse than a random permutation.
101    a=torch.tensor([-.8,.1,1.4,2.0]); b=torch.tensor([-1.1,.4,1.0,2.5])
102    sorted_cost=float(((a-b)**2).mean()); random_cost=float(((a-b[torch.tensor([2,0,3,1])])**2).mean())
103    # Interpolation derivative check.
104    t=.37; y_t=(1-t)*a+t*b; y_t2=(1-(t+1e-4))*a+(t+1e-4)*b
105    deriv_err=float(((y_t2-y_t)/1e-4-(b-a)).abs().max())
106    math_check={'sorted_ot_cost':sorted_cost,'random_coupling_cost':random_cost,'derivative_max_error':deriv_err}
107    mean=train_mean(); flow0=train_flow(0.0); flow8=train_flow(8.0)
108    eval0=evaluate(flow0,mean); eval8=evaluate(flow8,mean)
109    out={'device':device,'math_check':math_check,'locality_sweep':rows,
110         'comparison_lambda0':eval0,'comparison_lambda8':eval8,
111         'predictions':{
112          'lambda_effect_at_eps_0.16': 'paired_R should decrease monotonically with lambda',
113          'epsilon_effect_at_lambda_8': 'paired_R should increase with epsilon',
114          'flow_distributional_effect': 'flow W1 should be lower than deterministic MSE W1 and mode coverage near 1'}}
115    with open('results.json','w') as f: json.dump(out,f,indent=2)
116    print(json.dumps(out,indent=2))
117
118if __name__=='__main__': main()