Knieper Rollout Stability Metric / experiment.py

Failed on benchmark

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