Finite-Horizon Hidden-State Observability Regularizer / observability_bench.py
Failed on benchmark
1import sys, json, random, math
2from pathlib import Path
3import numpy as np
4import torch
5sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
6import bench
7from bench import train_model, sweep_baseline, make_report
8
9SEEDS = tuple(range(8))
10# Union is shared: every idea learning rate is evaluated by the baseline sweep.
11LR_GRID = [0.0015, 0.003, 0.006]
12EPOCHS = 3
13BATCH = 64
14HIDDEN = 64
15M = 32
16T_OBS = 2
17LAMBDA = 0.001
18EPS = 1e-3
19
20def seed_all(seed):
21 random.seed(seed); np.random.seed(seed); torch.manual_seed(seed)
22 if torch.cuda.is_available(): torch.cuda.manual_seed_all(seed)
23
24def new_model():
25 return bench.make_model('rnn_small', (24,), 1)
26
27def baseline_fn(cfg):
28 def run(seed):
29 seed_all(seed)
30 ds = bench.get_dataset('dynamics', seed=seed, n_train=200, n_test=200)
31 model = new_model()
32 _, metric, _ = train_model(model, ds, epochs=EPOCHS, lr=cfg['lr'],
33 batch=BATCH, weight_decay=0.0, log=lambda *a: None)
34 return float(metric)
35 return run
36
37def gru_observation_jacobian(model, flat_x, m=M, tsteps=T_OBS, create_graph=False):
38 """J of selected hidden coordinates over a real trained GRU trajectory wrt h0."""
39 seq = flat_x.reshape(1, -1, 3)
40 # This function is intentionally model-dependent, not an analytic toy graph.
41 def obs(h0):
42 h = h0.reshape(1, 1, -1)
43 ys = []
44 for t in range(min(tsteps, seq.shape[1])):
45 _, h = model.rnn(seq[:, t:t+1, :], h)
46 ys.append(h[0, 0, :m])
47 return torch.cat(ys)
48 h0 = torch.zeros(HIDDEN, device=flat_x.device, dtype=flat_x.dtype, requires_grad=True)
49 old = torch.backends.cudnn.enabled
50 torch.backends.cudnn.enabled = False
51 try:
52 J = torch.autograd.functional.jacobian(obs, h0, create_graph=create_graph)
53 finally:
54 torch.backends.cudnn.enabled = old
55 return J
56
57def obs_penalty(model, flat_x):
58 J = gru_observation_jacobian(model, flat_x, M, T_OBS, True)
59 G = J.T @ J + EPS * torch.eye(J.shape[1], device=J.device, dtype=J.dtype)
60 return -torch.linalg.slogdet(G)[1]
61
62def idea_train(model, ds, lr):
63 # Own loop is required because the intervention is a new training loss.
64 ladder = ['cuda', 'cpu'] if torch.cuda.is_available() else ['cpu']
65 last = None
66 for dev in ladder:
67 try:
68 model = model.to(dev)
69 x, y = ds['xtr'].to(dev), ds['ytr'].to(dev)
70 opt = torch.optim.Adam(model.parameters(), lr=lr)
71 model.train()
72 for ep in range(EPOCHS):
73 perm = torch.randperm(len(x), device=dev)
74 for bi in range(0, len(x), BATCH):
75 ix = perm[bi:bi+BATCH]
76 pred = model(x[ix])
77 task = (pred - y[ix]).pow(2).mean()
78 # One representative trajectory per minibatch keeps this small.
79 reg = obs_penalty(model, x[ix[0]]) if (ep == EPOCHS-1 and bi == 0) else torch.zeros((), device=dev)
80 loss = task + LAMBDA * reg
81 opt.zero_grad(set_to_none=True); loss.backward()
82 torch.nn.utils.clip_grad_norm_(model.parameters(), 5.0)
83 opt.step()
84 model.eval()
85 with torch.no_grad(): metric = float((model(x if False else ds['xte'].to(dev)) - ds['yte'].to(dev)).pow(2).mean().cpu())
86 return model, metric
87 except Exception as e:
88 last = e
89 try: model = model.to('cpu')
90 except Exception: pass
91 raise RuntimeError(last)
92
93def idea_fn(cfg, keep=False):
94 def run(seed):
95 seed_all(seed)
96 ds = bench.get_dataset('dynamics', seed=seed, n_train=200, n_test=200)
97 model, metric = idea_train(new_model(), ds, cfg['lr'])
98 if keep: kept[(seed, cfg['lr'])] = (model, ds)
99 return float(metric)
100 return run
101
102def evaluate8(fn):
103 vals = [float(fn(s)) for s in SEEDS]
104 return {'per_seed': vals, 'mean': float(np.mean(vals)), 'std': float(np.std(vals, ddof=1))}
105
106def signature(model, ds):
107 model.eval(); x = ds['xte'][0].to(next(model.parameters()).device)
108 rows=[]
109 with torch.no_grad():
110 pass
111 for m in [8, 16, 32]:
112 J = gru_observation_jacobian(model, x, m, 2, False).detach()
113 sv = torch.linalg.svdvals(J)
114 rank = int((sv > 1e-5).sum().item())
115 G = J.T @ J + EPS*torch.eye(HIDDEN, device=J.device)
116 rows.append({'m':m, 'predicted_rank_upper_bound_T1':min(HIDDEN,2*m),
117 'observed_rank_T1':rank, 'observed_smin':float(sv[-1]),
118 'observed_logdet':float(torch.linalg.slogdet(G)[1])})
119 return {'state_dim':HIDDEN, 'horizon_T':1,
120 'predicted_counting_threshold_m':HIDDEN//2,
121 'observed':rows,
122 'confirmed': any(r['m']==HIDDEN//2 and r['observed_rank_T1']>=HIDDEN for r in rows)}
123
124kept={}
125def main():
126 grid=[{'lr':v} for v in LR_GRID]
127 base=sweep_baseline(baseline_fn, grid, seeds=(0,1,2,3))
128 idea_results={}
129 for cfg in grid:
130 idea_results[str(cfg['lr'])]=evaluate8(idea_fn(cfg, keep=(cfg['lr']==base['best_cfg']['lr'])))
131 best_key=min(idea_results, key=lambda k: idea_results[k]['mean'])
132 idea=idea_results[best_key]
133 # Recover a trained baseline model at the selected configuration for signature.
134 bmodel=new_model(); seed_all(0)
135 bds=bench.get_dataset('dynamics',0,n_train=400,n_test=400)
136 bmodel,_,_=train_model(bmodel,bds,epochs=EPOCHS,lr=base['best_cfg']['lr'],batch=BATCH,log=lambda *a:None)
137 imodel, ids = kept.get((0, base['best_cfg']['lr']), (None,None))
138 if imodel is None:
139 imodel, _ = idea_train(new_model(), bds, base['best_cfg']['lr']); ids=bds
140 extra={'prediction':'For T=1, m >= n/2 is the counting threshold for full rank.',
141 'baseline_trained_model':signature(bmodel,bds),
142 'idea_trained_model':signature(imodel,ids)}
143 report=make_report('dynamics','rnn_small',base,idea,extra)
144 report['idea_sweep']=idea_results
145 report['protocol_notes']='Baseline sweep uses the same three learning rates as the idea; idea adds only finite-horizon logdet loss.'
146 Path('bench_report.json').write_text(json.dumps(report,indent=2))
147 print(json.dumps(report,indent=2))
148if __name__=='__main__': main()