Solver-Trajectory Flow Matching / run_experiment.py
Mechanism confirmed, baseline not beaten
1import os, json, math, random
2import numpy as np
3import torch
4from torch import nn
5
6SEED=926
7random.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED)
8torch.set_num_threads(min(8, os.cpu_count() or 1))
9device=torch.device('cuda' if torch.cuda.is_available() else 'cpu')
10
11# A 2-D conditional refinement problem. c is the target/context and z is source noise.
12# The teacher follows a smooth, non-straight feasible path with zero endpoint bump.
13def teacher(z, c, t):
14 # curved path: target is c; curvature direction depends on c, and source noise is retained
15 tt=t[...,None]
16 base=(1-tt)*z+tt*c
17 bump=0.65*torch.sin(math.pi*t)[...,None]*torch.stack((c[...,1], -c[...,0]), -1)
18 return base+bump
19
20def anchors(z,c,K):
21 ts=torch.linspace(0,1,K+1,device=z.device)
22 return torch.stack([teacher(z,c,ts[k].expand(z.shape[0])) for k in range(K+1)],1),ts
23
24# Analytic math verification, before learning.
25def math_checks():
26 z=torch.tensor([[0.4,-0.7],[1.1,0.2]])
27 c=torch.tensor([[0.8,0.3],[-0.2,1.0]])
28 rows=[]
29 for K in [2,4,8,16,32]:
30 a,ts=anchors(z,c,K)
31 # Random interpolation points, and exact segment velocity identity.
32 errs=[]; teachererrs=[]; velerrs=[]
33 for k in range(K):
34 r=torch.tensor([0.17+0.61*((k+1)%5)/5, 0.31+0.53*((k+2)%7)/7])
35 x=(1-r[:,None])*a[:,k]+r[:,None]*a[:,k+1]
36 t=(1-r)*ts[k]+r*ts[k+1]
37 errs.append((x-((1-r[:,None])*a[:,k]+r[:,None]*a[:,k+1])).abs().max().item())
38 teachererrs.append((x-teacher(z,c,t)).abs().max().item())
39 vel=(a[:,k+1]-a[:,k])/(ts[k+1]-ts[k])
40 # finite difference on the piecewise path is exactly velocity
41 eps=(ts[k+1]-ts[k])*0.23
42 x2=(1-(r+0.23)[:,None])*a[:,k]+(r+0.23)[:,None]*a[:,k+1]
43 velerrs.append(((x2-x)/eps-vel).abs().max().item())
44 rows.append({'K':K,'piecewise_interpolation_identity_max_error':max(errs),'curved_teacher_deviation_max_error':max(teachererrs),'velocity_max_error':max(velerrs)})
45 # Polygonal approximation error at midpoints: expected O(K^-2) for C2 path.
46 dev=[]
47 z1=torch.tensor([[0.7,-0.4]]); c1=torch.tensor([[0.3,1.2]])
48 for K in [2,4,8,16,32,64]:
49 a,ts=anchors(z1,c1,K); mids=(ts[:-1]+ts[1:])/2
50 exact=teacher(z1.repeat(K,1),c1.repeat(K,1),mids)
51 poly=(a[:, :-1]+a[:,1:])[0]/2
52 dev.append((K,float(((exact-poly)**2).sum(1).sqrt().max())))
53 slopes=[]
54 for (k1,e1),(k2,e2) in zip(dev[:-1],dev[1:]): slopes.append(math.log(e1/e2,2))
55 return {'identity_checks':rows,'polygonal_deviation':dev,
56 'predicted_polygonal_log2_slope':2.0,'observed_slopes':slopes,
57 'prediction_tolerance':0.25}
58
59class Field(nn.Module):
60 def __init__(self):
61 super().__init__(); self.net=nn.Sequential(nn.Linear(5,96),nn.Tanh(),nn.Linear(96,96),nn.Tanh(),nn.Linear(96,2))
62 def forward(self,x,t,c):
63 return self.net(torch.cat((x,t[:,None],c),1))
64
65def make_batch(n,K,trajectory,weighted=False):
66 z=torch.randn(n,2,device=device); c=torch.randn(n,2,device=device)
67 if trajectory:
68 a,ts=anchors(z,c,K); k=torch.randint(0,K,(n,),device=device); r=torch.rand(n,device=device)
69 x=(1-r[:,None])*a[torch.arange(n),k]+r[:,None]*a[torch.arange(n),k+1]
70 t=ts[k]+r*(ts[k+1]-ts[k]); u=(a[torch.arange(n),k+1]-a[torch.arange(n),k])/(ts[k+1]-ts[k])[:,None]
71 # q increases toward later, more reliable refinement segments; lambda is tested separately.
72 q=(k.float()+0.5)/K
73 return x,t,c,u,q
74 r=torch.rand(n,device=device); t=r; x=(1-r[:,None])*z+r[:,None]*c; u=c-z
75 return x,t,c,u,torch.ones(n,device=device)
76
77def train(trajectory,K,lam,steps=1800):
78 m=Field().to(device); opt=torch.optim.Adam(m.parameters(),lr=2e-3)
79 m.train()
80 for it in range(steps):
81 x,t,c,u,q=make_batch(128,K,trajectory)
82 pred=m(x,t,c); w=1+lam*q if trajectory else torch.ones_like(q)
83 loss=(w[:,None]*(pred-u)**2).mean(); opt.zero_grad(); loss.backward(); opt.step()
84 return m
85
86@torch.no_grad()
87def evaluate(m,K,N=1500,steps_list=(4,8,16)):
88 if isinstance(N, tuple):
89 steps_list, N = N, 1500
90 m.eval(); z=torch.randn(N,2,device=device); c=torch.randn(N,2,device=device)
91 out={}
92 for S in steps_list:
93 x=z.clone(); path=[]
94 for j in range(S):
95 t=torch.full((N,),j/S,device=device); path.append(x.clone()); x=x+m(x,t,c)/S
96 path.append(x.clone())
97 # Compare to exact teacher at solver times; this is the promised trajectory metric.
98 ts=torch.arange(S+1,device=device).float()/S
99 exact=torch.stack([teacher(z,c,ts[j].expand(N)) for j in range(S+1)],1)
100 pred=torch.stack(path,1)
101 adherence=((pred-exact)**2).sum(2).sqrt().mean().item()
102 final=((x-c)**2).sum(1).sqrt().mean().item()
103 out[str(S)]={'trajectory_distance':adherence,'final_target_distance':final}
104 return out
105
106def segment_errors(m,K,N=3000):
107 m.eval(); vals=[]
108 with torch.no_grad():
109 for k in range(K):
110 x,t,c,u,q=make_batch(N//K,K,True)
111 mask=torch.full((N//K,),k,device=device,dtype=torch.long)
112 # make_batch chooses random k, so directly regenerate fixed segment
113 z=torch.randn(N//K,2,device=device); c=torch.randn(N//K,2,device=device); a,ts=anchors(z,c,K); r=torch.rand(N//K,device=device)
114 x=(1-r[:,None])*a[:,k]+r[:,None]*a[:,k+1]; tt=ts[k]+r*(ts[k+1]-ts[k]); u=(a[:,k+1]-a[:,k])/(ts[k+1]-ts[k])
115 vals.append(float(((m(x,tt,c)-u)**2).mean().item()))
116 return vals
117
118def main():
119 checks=math_checks()
120 # Endpoint baseline and trajectory models, identical architecture/training budget.
121 base=train(False,8,0,1800); traj=train(True,8,0,1800); weighted=train(True,8,1.0,1800)
122 results={'device':str(device),'checks':checks,
123 'baseline_endpoint':evaluate(base,8),'trajectory_lambda0':evaluate(traj,8),'trajectory_lambda1':evaluate(weighted,8),
124 'segment_mse_lambda0':segment_errors(traj,8),'segment_mse_lambda1':segment_errors(weighted,8)}
125 # A small lambda sweep gives the weight prediction: lambda=0 and constant q would share optimum;
126 # with informative q, high lambda preferentially reduces late/high-q segment error.
127 sweep={}
128 for lam in [0.0,0.25,1.0,4.0]:
129 mm=train(True,8,lam,1200); sweep[str(lam)]={'eval':evaluate(mm,8,(4,)),'segment_mse':segment_errors(mm,8,1600)}
130 results['lambda_sweep']=sweep
131 with open('results.json','w') as f: json.dump(results,f,indent=2)
132 print(json.dumps(results,indent=2))
133
134if __name__=='__main__': main()