Solver-Trajectory Flow Matching / run_experiment.py

Mechanism confirmed, baseline not beaten

Raw ⬇ ZIP
  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()