Knieper Rollout Stability Metric / experiment.py
Failed on benchmark
1import json, random
2import numpy as np
3
4
5def toy_checks():
6 out = {}; H = 8
7 av = np.array([.80,.95,.99,1.,1.01,1.05,1.20]); obs=[]
8 for a in av:
9 z=g=1.
10 for _ in range(H): z*=abs(a); g=max(g,z)
11 obs.append(g)
12 out['scalar_boundary']={'H':H,'predicted_boundary_abs_a':1.,'sweep_a':av.tolist(),
13 'observed_gain':obs,'predicted_gain':[max(1.,abs(float(a))**H) for a in av]}
14 Ks=np.array([.25,.75,1.,2.,4.]); trans=[]; final=[]
15 for K in Ks:
16 A=np.array([[0.,K],[0.,0.]]); d=np.array([0.,1.]); tr=[d.copy()]
17 for _ in range(2): d=A@d; tr.append(d.copy())
18 trans.append(max(np.linalg.norm(q) for q in tr)); final.append(np.linalg.norm(tr[-1]))
19 out['transient_blind_spot']={'K':Ks.tolist(),'predicted_DH_over_eps':[max(1.,float(k)) for k in Ks],
20 'observed_DH_over_eps':trans,'final_step_gain':final}
21 rho,H2=.8,5; K2=np.array([0.,.5,1.,2.,4.]); gains=[]; pred=[]
22 for K in K2:
23 A=np.array([[rho,K],[0.,rho]]); d=np.array([0.,1.]); peak=1.
24 for k in range(1,H2+1):
25 d=A@d; peak=max(peak,np.linalg.norm(d))
26 gains.append(peak); pred.append(max([1.]+[float(np.linalg.norm([k*K*rho**(k-1),rho**k])) for k in range(1,H2+1)]))
27 out['nonnormal_scaling']={'rho':rho,'H':H2,'K':K2.tolist(),'observed_gain':gains,
28 'predicted_gain':pred,'slope_observed_large_K':float((gains[-1]-gains[1])/(K2[-1]-K2[1]))}
29 return out
30
31
32def run_gru(device):
33 import torch
34 import torch.nn as nn
35 torch.manual_seed(7); np.random.seed(7); random.seed(7)
36 T,N,B,inp,hid=18,512,64,3,16
37 x=torch.randn(N,T,inp,device=device); y=torch.roll(x,-1,1); y[:,-1]=0
38 class Model(nn.Module):
39 def __init__(self):
40 super().__init__(); self.rnn=nn.GRU(inp,hid,batch_first=True); self.head=nn.Linear(hid,inp)
41 def forward(self,z,h=None):
42 q,h=self.rnn(z,h); return self.head(q),h,q
43 def train(lam):
44 torch.manual_seed(11); m=Model().to(device); opt=torch.optim.Adam(m.parameters(),lr=.01)
45 for step in range(100):
46 ix=torch.randint(0,N,(B,),device=device); z=x[ix]; target=y[ix]; p,_,_=m(z)
47 task=(p-target).pow(2).mean(); eps=torch.randn(B,hid,device=device)*.02
48 h0=torch.zeros(1,B,hid,device=device); h1=h0.clone(); h1[0]+=eps; u=z[:,:6]
49 _,_,a=m(u,h0); _,_,b=m(u,h1)
50 roll=torch.sqrt((a-b).pow(2).sum(-1)+1e-12).amax(1)/(eps.norm(dim=1)+1e-8)
51 loss=task+lam*roll.mean(); opt.zero_grad(); loss.backward(); opt.step()
52 with torch.no_grad():
53 ix=torch.arange(128,device=device); p,_,_=m(x[ix]); mse=(p-y[ix]).pow(2).mean().item()
54 eps=torch.randn(128,hid,device=device)*.02; h0=torch.zeros(1,128,hid,device=device); h1=h0.clone(); h1[0]+=eps
55 _,_,a=m(x[ix,:8],h0); _,_,b=m(x[ix,:8],h1)
56 gain=(torch.sqrt((a-b).pow(2).sum(-1)+1e-12).amax(1)/eps.norm(dim=1)).median().item()
57 return {'loss':mse,'median_G8':gain}
58 return {'device':device,'baseline':train(0.),'rollout_lambda_0.03':train(.03)}
59
60
61def gru_experiment():
62 try:
63 import torch
64 requested='cuda' if torch.cuda.is_available() else 'cpu'
65 try: return run_gru(requested)
66 except Exception as first:
67 if requested=='cuda':
68 return {'device':'cpu_fallback','cuda_error':repr(first),'result':run_gru('cpu')}
69 return {'error':repr(first)}
70 except Exception as e: return {'error':repr(e)}
71
72if __name__=='__main__':
73 result={'toy':toy_checks(),'gru':gru_experiment()}
74 with open('results.json','w') as f: json.dump(result,f,indent=2)
75 print(json.dumps(result,indent=2))