import os, json, math, random import numpy as np import torch from torch import nn SEED=926 random.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED) torch.set_num_threads(min(8, os.cpu_count() or 1)) device=torch.device('cuda' if torch.cuda.is_available() else 'cpu') # A 2-D conditional refinement problem. c is the target/context and z is source noise. # The teacher follows a smooth, non-straight feasible path with zero endpoint bump. def teacher(z, c, t): # curved path: target is c; curvature direction depends on c, and source noise is retained tt=t[...,None] base=(1-tt)*z+tt*c bump=0.65*torch.sin(math.pi*t)[...,None]*torch.stack((c[...,1], -c[...,0]), -1) return base+bump def anchors(z,c,K): ts=torch.linspace(0,1,K+1,device=z.device) return torch.stack([teacher(z,c,ts[k].expand(z.shape[0])) for k in range(K+1)],1),ts # Analytic math verification, before learning. def math_checks(): z=torch.tensor([[0.4,-0.7],[1.1,0.2]]) c=torch.tensor([[0.8,0.3],[-0.2,1.0]]) rows=[] for K in [2,4,8,16,32]: a,ts=anchors(z,c,K) # Random interpolation points, and exact segment velocity identity. errs=[]; teachererrs=[]; velerrs=[] for k in range(K): r=torch.tensor([0.17+0.61*((k+1)%5)/5, 0.31+0.53*((k+2)%7)/7]) x=(1-r[:,None])*a[:,k]+r[:,None]*a[:,k+1] t=(1-r)*ts[k]+r*ts[k+1] errs.append((x-((1-r[:,None])*a[:,k]+r[:,None]*a[:,k+1])).abs().max().item()) teachererrs.append((x-teacher(z,c,t)).abs().max().item()) vel=(a[:,k+1]-a[:,k])/(ts[k+1]-ts[k]) # finite difference on the piecewise path is exactly velocity eps=(ts[k+1]-ts[k])*0.23 x2=(1-(r+0.23)[:,None])*a[:,k]+(r+0.23)[:,None]*a[:,k+1] velerrs.append(((x2-x)/eps-vel).abs().max().item()) rows.append({'K':K,'piecewise_interpolation_identity_max_error':max(errs),'curved_teacher_deviation_max_error':max(teachererrs),'velocity_max_error':max(velerrs)}) # Polygonal approximation error at midpoints: expected O(K^-2) for C2 path. dev=[] z1=torch.tensor([[0.7,-0.4]]); c1=torch.tensor([[0.3,1.2]]) for K in [2,4,8,16,32,64]: a,ts=anchors(z1,c1,K); mids=(ts[:-1]+ts[1:])/2 exact=teacher(z1.repeat(K,1),c1.repeat(K,1),mids) poly=(a[:, :-1]+a[:,1:])[0]/2 dev.append((K,float(((exact-poly)**2).sum(1).sqrt().max()))) slopes=[] for (k1,e1),(k2,e2) in zip(dev[:-1],dev[1:]): slopes.append(math.log(e1/e2,2)) return {'identity_checks':rows,'polygonal_deviation':dev, 'predicted_polygonal_log2_slope':2.0,'observed_slopes':slopes, 'prediction_tolerance':0.25} class Field(nn.Module): def __init__(self): super().__init__(); self.net=nn.Sequential(nn.Linear(5,96),nn.Tanh(),nn.Linear(96,96),nn.Tanh(),nn.Linear(96,2)) def forward(self,x,t,c): return self.net(torch.cat((x,t[:,None],c),1)) def make_batch(n,K,trajectory,weighted=False): z=torch.randn(n,2,device=device); c=torch.randn(n,2,device=device) if trajectory: a,ts=anchors(z,c,K); k=torch.randint(0,K,(n,),device=device); r=torch.rand(n,device=device) x=(1-r[:,None])*a[torch.arange(n),k]+r[:,None]*a[torch.arange(n),k+1] 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] # q increases toward later, more reliable refinement segments; lambda is tested separately. q=(k.float()+0.5)/K return x,t,c,u,q r=torch.rand(n,device=device); t=r; x=(1-r[:,None])*z+r[:,None]*c; u=c-z return x,t,c,u,torch.ones(n,device=device) def train(trajectory,K,lam,steps=1800): m=Field().to(device); opt=torch.optim.Adam(m.parameters(),lr=2e-3) m.train() for it in range(steps): x,t,c,u,q=make_batch(128,K,trajectory) pred=m(x,t,c); w=1+lam*q if trajectory else torch.ones_like(q) loss=(w[:,None]*(pred-u)**2).mean(); opt.zero_grad(); loss.backward(); opt.step() return m @torch.no_grad() def evaluate(m,K,N=1500,steps_list=(4,8,16)): if isinstance(N, tuple): steps_list, N = N, 1500 m.eval(); z=torch.randn(N,2,device=device); c=torch.randn(N,2,device=device) out={} for S in steps_list: x=z.clone(); path=[] for j in range(S): t=torch.full((N,),j/S,device=device); path.append(x.clone()); x=x+m(x,t,c)/S path.append(x.clone()) # Compare to exact teacher at solver times; this is the promised trajectory metric. ts=torch.arange(S+1,device=device).float()/S exact=torch.stack([teacher(z,c,ts[j].expand(N)) for j in range(S+1)],1) pred=torch.stack(path,1) adherence=((pred-exact)**2).sum(2).sqrt().mean().item() final=((x-c)**2).sum(1).sqrt().mean().item() out[str(S)]={'trajectory_distance':adherence,'final_target_distance':final} return out def segment_errors(m,K,N=3000): m.eval(); vals=[] with torch.no_grad(): for k in range(K): x,t,c,u,q=make_batch(N//K,K,True) mask=torch.full((N//K,),k,device=device,dtype=torch.long) # make_batch chooses random k, so directly regenerate fixed segment 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) 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]) vals.append(float(((m(x,tt,c)-u)**2).mean().item())) return vals def main(): checks=math_checks() # Endpoint baseline and trajectory models, identical architecture/training budget. base=train(False,8,0,1800); traj=train(True,8,0,1800); weighted=train(True,8,1.0,1800) results={'device':str(device),'checks':checks, 'baseline_endpoint':evaluate(base,8),'trajectory_lambda0':evaluate(traj,8),'trajectory_lambda1':evaluate(weighted,8), 'segment_mse_lambda0':segment_errors(traj,8),'segment_mse_lambda1':segment_errors(weighted,8)} # A small lambda sweep gives the weight prediction: lambda=0 and constant q would share optimum; # with informative q, high lambda preferentially reduces late/high-q segment error. sweep={} for lam in [0.0,0.25,1.0,4.0]: mm=train(True,8,lam,1200); sweep[str(lam)]={'eval':evaluate(mm,8,(4,)),'segment_mse':segment_errors(mm,8,1600)} results['lambda_sweep']=sweep with open('results.json','w') as f: json.dump(results,f,indent=2) print(json.dumps(results,indent=2)) if __name__=='__main__': main()