Gramian-Regularized Latent State Models / stage2_gramian_bench.py

Failed on benchmark

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