Transport-PDE Predictor for Delayed Neural State Updates / stage2_bench.py

Mechanism confirmed, baseline not beaten

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