OT-Sufficient Bottleneck Flow Matching / ot_sufficient_flow.py
Mechanism confirmed, baseline not beaten
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()