Transport-PDE Predictor for Delayed Neural State Updates / stage2_bench.py
Mechanism confirmed, baseline not beaten
1import json, math, os, random
2from pathlib import Path
3import numpy as np
4import torch
5from torch import nn
6
7SEEDS=list(range(8)); DEVICE='cuda' if torch.cuda.is_available() else 'cpu'
8
9def seed(s):
10 random.seed(s); np.random.seed(s); torch.manual_seed(s)
11 if torch.cuda.is_available(): torch.cuda.manual_seed_all(s)
12
13def plant(x,u):
14 # locally controlled nonlinear pendulum: [angle, angular velocity]
15 th,om=x[...,0],x[...,1]
16 return torch.stack((th+0.12*om, 0.985*om-0.16*torch.sin(th)+0.10*u.squeeze(-1)), -1)
17
18def data(s,n=400,T=10,d=4):
19 rng=np.random.default_rng(s)
20 xs=[]; qs=[]; ys=[]
21 A=np.array([[1,.12],[-.16,.985]]); B=np.array([[0],[.10]])
22 K=np.array([[-1.05,-.72]])
23 for _ in range(n):
24 x=rng.uniform([-1.1,-1.0],[1.1,1.0]); queue=[np.zeros(1) for _ in range(d)]
25 for t in range(T):
26 # target is stabilizing action for the state at actuation time
27 target=float(np.clip(K@x, -2,2))
28 hist=np.asarray(queue,dtype=np.float32).reshape(-1)
29 xs.append(x.astype(np.float32)); qs.append(hist); ys.append(target)
30 applied=queue.pop(0); queue.append(np.array([target]))
31 x=A@x+B[:,0]*applied[0] + rng.normal(0,.008,2)
32 x[0]=((x[0]+np.pi)%(2*np.pi))-np.pi
33 X=np.asarray(xs); Q=np.asarray(qs); Y=np.asarray(ys)[:,None]
34 # deterministic split by generated order; independent test seed
35 return X,Q,Y
36
37class Controller(nn.Module):
38 def __init__(self,d):
39 super().__init__(); self.net=nn.Sequential(nn.Linear(2,24),nn.Tanh(),nn.Linear(24,1),nn.Tanh())
40 def forward(self,x): return 2*self.net(x)
41
42def predict(x,q,d):
43 # frozen local linearization, chronological queued inputs
44 A=x.new_tensor([[1.,.12],[-.16,.985]])
45 B=x.new_tensor([[0.],[.10]])
46 p=x
47 for j in range(d): p=p@A.T + q[:,j:j+1]@B.T
48 return p
49
50def fit(s,lr,idea,d=4,epochs=18):
51 seed(s); X,Q,Y=data(s,400,10,d); Xt,Qt,Yt=data(s+10000,120,10,d)
52 X=torch.tensor(X,device=DEVICE); Q=torch.tensor(Q,device=DEVICE); Y=torch.tensor(Y,device=DEVICE)
53 Xt=torch.tensor(Xt,device=DEVICE); Qt=torch.tensor(Qt,device=DEVICE); Yt=torch.tensor(Yt,device=DEVICE)
54 m=Controller(d).to(DEVICE); opt=torch.optim.Adam(m.parameters(),lr=lr)
55 for _ in range(epochs):
56 perm=torch.randperm(len(X),device=DEVICE)
57 for ix in perm.split(128):
58 inp=predict(X[ix],Q[ix],d) if idea else X[ix]
59 loss=((m(inp)-Y[ix])**2).mean(); opt.zero_grad(); loss.backward(); opt.step()
60 with torch.no_grad():
61 inp=predict(Xt,Qt,d) if idea else Xt
62 mse=float(((m(inp)-Yt)**2).mean().cpu())
63 # closed-loop rollout metric on fresh initial conditions with delayed actuation
64 rng=np.random.default_rng(s+20000); vals=[]; pred_err=[]
65 for z in range(30):
66 x=torch.tensor(rng.uniform([-1.,-.8],[1.,.8]),dtype=torch.float32,device=DEVICE).unsqueeze(0)
67 q=torch.zeros((1,d),device=DEVICE); accum=0.
68 for t in range(35):
69 p=predict(x,q,d); u=m(p if idea else x).clamp(-2,2)
70 if idea: pred_err.append(float(torch.linalg.norm(p-x).cpu()))
71 applied=q[:,0:1]; q=torch.cat((q[:,1:],u),1); x=plant(x,applied)
72 accum += float((x*x).sum().cpu())
73 vals.append(accum/35)
74 return {'test_mse':mse,'rollout_mse':float(np.mean(vals)), 'predictor_shift':float(np.mean(pred_err)) if pred_err else 0.0}
75
76def perm_p(a,b):
77 a=np.asarray(a); b=np.asarray(b); obs=float(np.mean(b-a)); rng=np.random.default_rng(991)
78 cnt=0; N=20000
79 for _ in range(N):
80 signs=rng.choice([-1,1],len(a)); v=float(np.mean((b-a)*signs))
81 cnt += abs(v)>=abs(obs)
82 return (cnt+1)/(N+1),obs
83
84def main():
85 # same union of learning rates on both sides; baseline sweep and idea sweep are identical.
86 lrs=[0.001,0.003,0.009]; d=4
87 base={str(lr):[fit(s,lr,False,d) for s in SEEDS] for lr in lrs}
88 idea={str(lr):[fit(s,lr,True,d) for s in SEEDS] for lr in lrs}
89 score=lambda r: np.mean([x['test_mse'] for x in r])
90 bestb=min(lrs,key=lambda x:score(base[str(x)])); besti=min(lrs,key=lambda x:score(idea[str(x)]))
91 br=base[str(bestb)]; ir=idea[str(besti)]
92 p,delta=perm_p([x['test_mse'] for x in br],[x['test_mse'] for x in ir])
93 # Signature measured on trained systems: prediction shift and observed one-step model mismatch.
94 sig={'delay':d,'trained_models':True,'predicted_effect':'finite-horizon state differs from stale state and should reduce delayed rollout error',
95 'observed_baseline_rollout_mse':float(np.mean([x['rollout_mse'] for x in br])),
96 'observed_idea_rollout_mse':float(np.mean([x['rollout_mse'] for x in ir])),
97 'observed_mean_predictor_shift':float(np.mean([x['predictor_shift'] for x in ir])),
98 'confirmed':float(np.mean([x['rollout_mse'] for x in ir])) < float(np.mean([x['rollout_mse'] for x in br]))}
99 out={'track':'dynamics_fallback','device':DEVICE,'baseline_sweep':{str(k):{'mean_test_mse':score(v),'per_seed':v} for k,v in base.items()},'idea_sweep':{str(k):{'mean_test_mse':score(v),'per_seed':v} for k,v in idea.items()},'best_baseline_lr':bestb,'best_idea_lr':besti,'paired_delta_mean':delta,'permutation_p':p,'mechanism_signature':sig,'custom_track':{'name':'delayed_pendulum_dynamics','file':'stage2_bench.py','domain':'dynamics'}}
100 Path('bench_report.json').write_text(json.dumps(out,indent=2)); print(json.dumps(out,indent=2))
101if __name__=='__main__': main()