Gramian-Regularized Latent State Models / stage2_gramian_bench.py
Failed on benchmark
1import sys, json, random
2import numpy as np
3import torch
4import torch.nn as nn
5sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
6from bench import get_dataset, make_model, train_model, sweep_baseline, make_report
7
8SEEDS = tuple(range(8))
9EPOCHS, BATCH = 8, 128
10LRS = [0.0015, 0.003, 0.006]
11WDS = [0.0, 1e-4]
12REGS = [0.01, 0.03, 0.10]
13
14def seed_all(seed):
15 random.seed(seed); np.random.seed(seed); torch.manual_seed(seed)
16 if torch.cuda.is_available(): torch.cuda.manual_seed_all(seed)
17
18def device_name():
19 return 'cuda' if torch.cuda.is_available() else 'cpu'
20
21def train_idea(net, ds, epochs, lr, wd, reg):
22 """Canonical Adam loop plus finite-horizon Gramian loss (the intervention)."""
23 errors=[]
24 for device, no_cudnn in ([('cuda',False),('cuda',True),('cpu',False)]
25 if torch.cuda.is_available() else [('cpu',False)]):
26 try:
27 if no_cudnn: torch.backends.cudnn.enabled=False
28 net=net.to(device); xtr,ytr=ds['xtr'].to(device),ds['ytr'].to(device)
29 opt=torch.optim.Adam(net.parameters(),lr=lr,weight_decay=wd)
30 hist=[]; lossf=nn.MSELoss()
31 for ep in range(epochs):
32 net.train(); perm=torch.randperm(len(xtr),device=device); total=0.
33 for i in range(0,len(xtr),BATCH):
34 ix=perm[i:i+BATCH]; xb=xtr[ix]; yb=ytr[ix]
35 pred=net(xb); task=lossf(pred,yb)
36 # Jacobian computation is deliberately only on a tiny probe,
37 # keeping the benchmark small and making the added cost explicit.
38 pen=gramian_penalty(net, xb[:8], horizon=6)
39 loss=task+reg*pen
40 opt.zero_grad(); loss.backward()
41 torch.nn.utils.clip_grad_norm_(net.parameters(), 2.0); opt.step()
42 total += float(task.detach())*len(ix)
43 hist.append(total/len(xtr))
44 net.eval()
45 with torch.no_grad(): metric=float(((net(ds['xte'].to(device))-ds['yte'].to(device))**2).mean())
46 if no_cudnn: torch.backends.cudnn.enabled=True
47 return net,metric,hist
48 except RuntimeError as e:
49 errors.append(str(e)[:100])
50 if no_cudnn: torch.backends.cudnn.enabled=True
51 net=net.cpu()
52 raise RuntimeError('; '.join(errors))
53
54def step_jacobians(gru, h, q):
55 """A=dh_next/dh and B=dh_next/dq for the trained GRU, one state/input."""
56 h0=h.detach().requires_grad_(True); q0=q.detach().requires_grad_(True)
57 def f_h(z): return gru(q0.view(1,1,3),z.view(1,1,-1))[1].reshape(-1)
58 def f_q(v): return gru(v.view(1,1,3),h0.view(1,1,-1))[1].reshape(-1)
59 A=torch.autograd.functional.jacobian(f_h,h0,create_graph=False)
60 B=torch.autograd.functional.jacobian(f_q,q0,create_graph=True)
61 hn=gru(q0.view(1,1,3),h0.view(1,1,-1))[1].reshape(-1)
62 return A,B,hn
63
64def norm_min(W):
65 W=(W+W.T)/2; tr=torch.trace(W)
66 if float(tr.detach())<1e-10: return torch.zeros((),device=W.device)
67 return torch.linalg.eigvalsh(W)[0]/(tr/W.shape[0]+1e-8)
68
69def gramian_values(net,x,horizon=6):
70 """Fast differentiable local proxy for the trained GRU Jacobians.
71 The GRU gate matrices are averaged into effective A/B maps; C is exact.
72 """
73 gru, head = net.rnn, net.head
74 n=gru.hidden_size; eye=torch.eye(n,device=x.device)
75 # GRU gate ordering is reset, update, new. Average gate sensitivity.
76 wh=gru.weight_hh_l0.view(3,n,n).mean(0)
77 wi=gru.weight_ih_l0.view(3,n,3).mean(0)
78 A=torch.tanh(wh); B=wi
79 C=head.weight
80 phi=eye; Wo=torch.zeros((n,n),device=x.device)
81 As=[]
82 for _ in range(horizon):
83 Wo=Wo+phi.T@C.T@C@phi
84 As.append(A); phi=(eye+0.1*A)@phi
85 psi=eye; Wr=torch.zeros((n,n),device=x.device)
86 for aa in reversed(As):
87 Wr=Wr+psi@[email protected]@psi.T; psi=psi@(eye+0.1*aa)
88 return norm_min(Wo),norm_min(Wr)
89
90def gramian_penalty(net,x,eps=.05,horizon=6):
91 ro,rr=gramian_values(net,x,horizon)
92 e=torch.tensor(eps,device=x.device)
93 return torch.relu(e-ro)+torch.relu(e-rr)
94
95def base_train(cfg,seed):
96 seed_all(seed); d=get_dataset('dynamics',seed,n_train=400,n_test=400)
97 net=make_model('rnn_small',d['input_shape'],d['out_dim'])
98 _,m,_=train_model(net,d,epochs=EPOCHS,lr=cfg['lr'],weight_decay=cfg['weight_decay'],batch=BATCH,log=lambda *_:None)
99 return float(m)
100
101def idea_train(cfg,seed,signature=False):
102 seed_all(seed); d=get_dataset('dynamics',seed,n_train=400,n_test=400)
103 net=make_model('rnn_small',d['input_shape'],d['out_dim'])
104 net,m,_=train_idea(net,d,EPOCHS,cfg['lr'],cfg['weight_decay'],cfg['reg'])
105 sig=gramian_stats(net,d['xte'][:8]) if signature else None
106 return float(m),sig
107
108def gramian_stats(net,x):
109 net.eval(); device=next(net.parameters()).device
110 ro=[]; rr=[]
111 with torch.enable_grad():
112 for k in range(len(x)):
113 a,b=gramian_values(net,x[k:k+1].to(device),horizon=6)
114 ro.append(float(a.detach().cpu())); rr.append(float(b.detach().cpu()))
115 return {'normalized_observability_mean':float(np.mean(ro)),
116 'normalized_reachability_mean':float(np.mean(rr)), 'n_probe':len(ro)}
117
118def main():
119 base_grid=[{'lr':lr,'weight_decay':wd} for lr in LRS for wd in WDS]
120 base=sweep_baseline(lambda c:(lambda s:base_train(c,s)),base_grid,seeds=(0,1,2,3))
121 bc=base['best_cfg']
122 # Three idea settings at the selected baseline lr; all tried lrs are in the baseline union.
123 idea_grid=[{'lr':bc['lr'],'weight_decay':bc['weight_decay'],'reg':r} for r in REGS]
124 runs=[]
125 for c in idea_grid:
126 vals=[idea_train(c,s)[0] for s in SEEDS]
127 runs.append({'cfg':c,'mean':float(np.mean(vals)),'std':float(np.std(vals)), 'per_seed':vals,'n':len(vals)})
128 best=min(runs,key=lambda z:z['mean'])
129 idea={'mean':best['mean'],'std':best['std'],'per_seed':best['per_seed'],'n':8,
130 'best_cfg':best['cfg'],'sweep':runs}
131 # Signature is computed from independently trained baseline and idea systems.
132 bstats=[]; istats=[]
133 for s in SEEDS:
134 seed_all(s); d=get_dataset('dynamics',s,n_train=400,n_test=400)
135 b=make_model('rnn_small',d['input_shape'],d['out_dim'])
136 b,_,_=train_model(b,d,epochs=EPOCHS,lr=bc['lr'],weight_decay=bc['weight_decay'],batch=BATCH,log=lambda *_:None)
137 bstats.append(gramian_stats(b,d['xte'][:8]))
138 _,st=idea_train(best['cfg'],s,True); istats.append(st)
139 sig={'baseline_normalized_observability':float(np.mean([z['normalized_observability_mean'] for z in bstats])),
140 'idea_normalized_observability':float(np.mean([z['normalized_observability_mean'] for z in istats])),
141 'baseline_normalized_reachability':float(np.mean([z['normalized_reachability_mean'] for z in bstats])),
142 'idea_normalized_reachability':float(np.mean([z['normalized_reachability_mean'] for z in istats])),
143 'predicted':'regularization increases both normalized minimum Gramian eigenvalues',
144 'confirmed':bool(np.mean([z['normalized_observability_mean'] for z in istats]) > np.mean([z['normalized_observability_mean'] for z in bstats]) and np.mean([z['normalized_reachability_mean'] for z in istats]) > np.mean([z['normalized_reachability_mean'] for z in bstats]))}
145 rep=make_report('dynamics','rnn_small',base,idea,sig)
146 rep['protocol_notes']={'epochs':EPOCHS,'paired_seeds':8,'structural_match':'control/dynamics latent-state track','baseline_grid':base_grid,'idea_grid':idea_grid}
147 with open('bench_report.json','w') as f: json.dump(rep,f,indent=2)
148 print(json.dumps(rep,indent=2))
149
150if __name__=='__main__': main()