Fisher-Observable Latent State Training / stage2_bench.py

Failed on benchmark

Raw ⬇ ZIP
 1import sys, json, random
 2from pathlib import Path
 3import numpy as np
 4import torch
 5import torch.nn as nn
 6sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
 7from bench import get_dataset, make_model, train_model, evaluate, sweep_baseline, make_report
 8TRACK='dynamics'; MODEL='rnn_small'; EPOCHS=4; NTR=240; NTE=240; BATCH=64
 9LRS=[1e-3,3e-3,1e-2]; ALPHAS=[1e-4,3e-4,1e-3]
10
11def seedall(s):
12    random.seed(s); np.random.seed(s); torch.manual_seed(s)
13
14def fisher_loss(net,x,nsamp=2,eps=1e-3):
15    # Fisher proxy for the learned recurrent rollout: output hidden trajectory
16    # sensitivity to the initial observed state (theta, omega).
17    xx=x[:nsamp].detach().clone().requires_grad_(True)
18    seq=xx.view(xx.shape[0],-1,3); out,_=net.rnn(seq)
19    vals=[]
20    for i in range(nsamp):
21        rows=[]
22        # one Jacobian row per observed rollout time, preserving the chain rule
23        for t in range(out.shape[1]):
24            g=torch.autograd.grad(out[i,t,0],xx,retain_graph=True,create_graph=True)[0][i,:2]
25            rows.append(g)
26        J=torch.stack(rows); I=J.T@J+eps*torch.eye(2,device=x.device)
27        ev=torch.linalg.eigvalsh(I)
28        vals.append(-torch.logdet(I)+1e-3*ev[-1]/(ev[0]+eps))
29    return torch.stack(vals).mean()
30
31def idea_train(seed,lr,alpha,return_net=False):
32    seedall(seed); ds=get_dataset(TRACK,seed,n_train=NTR,n_test=NTE)
33    net=make_model(MODEL,ds['input_shape'],ds['out_dim'])
34    try:
35        device='cuda' if torch.cuda.is_available() else 'cpu'; net.to(device)
36        x,y=ds['xtr'].to(device),ds['ytr'].to(device); opt=torch.optim.Adam(net.parameters(),lr=lr); mse=nn.MSELoss(); gen=torch.Generator().manual_seed(seed)
37        for _ in range(EPOCHS):
38            net.train()
39            for ix in torch.randperm(len(x),generator=gen).split(BATCH):
40                xb,yb=x[ix],y[ix]; opt.zero_grad(set_to_none=True)
41                pred=net(xb); loss=mse(pred,yb)+alpha*fisher_loss(net,xb); loss.backward(); torch.nn.utils.clip_grad_norm_(net.parameters(),5); opt.step()
42        net.eval()
43        with torch.no_grad(): metric=float(mse(net(ds['xte'].to(device)),ds['yte'].to(device)).cpu())
44        return (metric,net,ds,device) if return_net else metric
45    except Exception:
46        # CPU fallback for shared/limited CUDA environments.
47        seedall(seed); ds=get_dataset(TRACK,seed,n_train=NTR,n_test=NTE); net=make_model(MODEL,ds['input_shape'],ds['out_dim']); opt=torch.optim.Adam(net.parameters(),lr=lr); mse=nn.MSELoss(); gen=torch.Generator().manual_seed(seed)
48        for _ in range(EPOCHS):
49            for ix in torch.randperm(len(ds['xtr']),generator=gen).split(BATCH):
50                xb,yb=ds['xtr'][ix],ds['ytr'][ix]; opt.zero_grad(); loss=mse(net(xb),yb)+alpha*fisher_loss(net,xb); loss.backward(); torch.nn.utils.clip_grad_norm_(net.parameters(),5); opt.step()
51        with torch.no_grad(): metric=float(mse(net(ds['xte']),ds['yte']))
52        return (metric,net,ds,'cpu') if return_net else metric
53
54def signature():
55    ans={}
56    for name,a in [('baseline',0.0),('idea',3e-4)]:
57        metric,net,ds,device=idea_train(0,3e-3,a,True); x=ds['xte'][:4].to(device).requires_grad_(True); out,_=net.rnn(x.view(4,-1,3)); evs=[]
58        for i in range(4):
59            J=torch.stack([torch.autograd.grad(out[i,t,0],x,retain_graph=True)[0][i,:2] for t in range(out.shape[1])]); evs.append(torch.linalg.eigvalsh(J.T@J+1e-3*torch.eye(2,device=device)).detach().cpu().numpy())
60        evs=np.asarray(evs); ans[name]={'test_mse_observed':metric,'fisher_lambda_min_predicted':float(evs[:,0].mean()),'fisher_lambda_max_predicted':float(evs[:,1].mean()),'condition_predicted':float(np.mean(evs[:,1]/evs[:,0]))}
61    ans['confirmed']=ans['idea']['fisher_lambda_min_predicted']>ans['baseline']['fisher_lambda_min_predicted'] and ans['idea']['condition_predicted']<ans['baseline']['condition_predicted']; return ans
62
63def main():
64    # Baseline gets the full union of all LRs used by the idea side.
65    base=sweep_baseline(lambda c: lambda s: float(train_model(make_model(MODEL,get_dataset(TRACK,s,NTR,NTE)['input_shape'],1),get_dataset(TRACK,s,NTR,NTE),epochs=EPOCHS,lr=c['lr'],batch=BATCH,log=lambda *_:None)[1]),[{'lr':v} for v in LRS])
66    idea_runs=[]
67    for a in ALPHAS:
68        r=evaluate(lambda s,a=a: idea_train(s,base['best_cfg']['lr'],a)); idea_runs.append((r,a))
69    idea,a=min(idea_runs,key=lambda z:z[0]['mean'])
70    rep=make_report(TRACK,MODEL,base,idea,{'type':'trained_model_fisher_signature','values':signature(),'confirmed_definition':'true only when trained idea increases mean predicted lambda_min and lowers condition number'})
71    rep['idea_sweep']=[{'alpha':a,'full':r} for r,a in idea_runs]; rep['protocol_notes']='Dynamics is the structural match: controlled pendulum multi-step rollout. Same rnn_small, data, epochs, batch, and baseline-selected LR; only Fisher loss differs.'
72    Path('bench_report.json').write_text(json.dumps(rep,indent=2)); print(json.dumps(rep,indent=2))
73if __name__=='__main__': main()