Tangential Bellman Tie Resolver / bench_stage2.py
Failed on benchmark
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
8
9SEED=2137
10EPOCHS=10
11NTR=1200
12NTE=400
13BATCH=128
14BETA=.8
15LAMBDA=.08
16
17def H(z, delta, P, beta):
18 # z,delta [B,S], P [B,S,S]
19 m=np.min(z+delta,axis=0)
20 return delta+beta*np.einsum('bxy,y->bx',P,m)
21
22def math_check():
23 rng=np.random.default_rng(SEED); B,S=4,7; beta=.73
24 P=rng.random((B,S,S)); P/=P.sum(2,keepdims=True); delta=rng.random((B,S))
25 z=rng.normal(size=(B,S)); w=rng.normal(size=(B,S))
26 ratio=np.max(abs(H(z,delta,P,beta)-H(w,delta,P,beta)))/np.max(abs(z-w))
27 star=np.zeros_like(delta)
28 for _ in range(1000):
29 new=H(star,delta,P,beta)
30 if np.max(abs(new-star))<1e-13: break
31 star=new
32 cur=np.zeros_like(delta); errs=[]
33 for _ in range(9): errs.append(float(np.max(abs(cur-star)))); cur=H(cur,delta,P,beta)
34 ratios=[errs[i+1]/errs[i] for i in range(8) if errs[i]>1e-14]
35 return {'operator_lipschitz_ratio':float(ratio),'beta':beta,'iteration_ratios':ratios,
36 'max_ratio_over_beta':float(max(ratios)/beta)}
37
38def seed_all(seed):
39 random.seed(seed); np.random.seed(seed); torch.manual_seed(seed)
40 if torch.cuda.is_available(): torch.cuda.manual_seed_all(seed)
41
42def branch_loss(net, x, y, beta=BETA):
43 # Three nearby action branches: perturb only the most recent control u.
44 xb=x.view(-1,8,3); offsets=torch.tensor([-0.35,0.,0.35],device=x.device)
45 outs=[]
46 for off in offsets:
47 q=xb.clone(); q[:,-1,2]+=off; outs.append(net(q.reshape(len(x),-1)))
48 pred=torch.cat(outs,1) # [n,3]
49 # Local deficit marks relative to best candidate, with target supervision.
50 delta=(pred-y[:,None]).pow(2); delta=delta-delta.min(1,keepdim=True).values
51 # Learned transition consequence proxy: branch predictions are continuation outcomes.
52 # Iterate finite Bellman map independently per sample (one-state branch pool).
53 z=torch.zeros_like(delta)
54 for _ in range(8):
55 m=(z+delta).min(1,keepdim=True).values
56 z=delta+beta*m
57 scores=z+delta
58 weights=torch.softmax(-scores/.12,dim=1)
59 resolved=(weights*pred).sum(1)
60 # Central action remains the task action; resolver consistency is auxiliary.
61 return (pred[:,1]-y).pow(2).mean()+LAMBDA*(resolved-y).pow(2).mean(), weights.detach(), pred.detach()
62
63def train_idea(seed, lr, wd, collect=False):
64 seed_all(seed); ds=get_dataset('dynamics',seed,n_train=NTR,n_test=NTE)
65 net=make_model('rnn_small',tuple(ds['xtr'].shape[1:]),1)
66 device='cuda' if torch.cuda.is_available() else 'cpu'
67 try:
68 net.to(device); x,y=ds['xtr'].to(device),ds['ytr'].to(device)
69 opt=torch.optim.Adam(net.parameters(),lr=lr,weight_decay=wd)
70 for _ in range(EPOCHS):
71 net.train(); p=torch.randperm(len(x),device=device)
72 for i in range(0,len(x),BATCH):
73 idx=p[i:i+BATCH]; loss,_,_=branch_loss(net,x[idx],y[idx])
74 opt.zero_grad(); loss.backward(); opt.step()
75 net.eval()
76 with torch.no_grad():
77 xt,yt=ds['xte'].to(device),ds['yte'].to(device)
78 loss,w,pred=branch_loss(net,xt,yt)
79 metric=float((pred[:,1]-yt).pow(2).mean())
80 # behavior signature: smoothness of resolver weights under action perturbation
81 smooth=float(torch.mean(torch.abs(w[:,2]-w[:,0])))
82 hard=float((pred.argmin(1)==0).float().mean())
83 if collect: return metric, {'resolver_weight_span':smooth,'hard_best_frequency':hard}
84 return metric
85 except RuntimeError:
86 # CPU retry is explicit for shared-GPU failures.
87 seed_all(seed); ds=get_dataset('dynamics',seed,n_train=NTR,n_test=NTE)
88 net=make_model('rnn_small',tuple(ds['xtr'].shape[1:]),1).cpu(); x,y=ds['xtr'],ds['ytr']
89 opt=torch.optim.Adam(net.parameters(),lr=lr,weight_decay=wd)
90 for _ in range(EPOCHS):
91 p=torch.randperm(len(x))
92 for i in range(0,len(x),BATCH):
93 idx=p[i:i+BATCH]; loss,_,_=branch_loss(net,x[idx],y[idx]); opt.zero_grad(); loss.backward(); opt.step()
94 with torch.no_grad(): pred=torch.cat([net(x.view(-1,8,3).clone().reshape(len(x),-1)) for _ in [0]],1)
95 return float(((net(ds['xte'])[:,0]-ds['yte'])**2).mean())
96
97def baseline_fn(cfg):
98 def run(seed):
99 seed_all(seed); ds=get_dataset('dynamics',seed,n_train=NTR,n_test=NTE)
100 net=make_model('rnn_small',tuple(ds['xtr'].shape[1:]),1)
101 _,metric,_=train_model(net,ds,epochs=EPOCHS,lr=cfg['lr'],batch=BATCH,weight_decay=cfg['wd'],log=lambda *_:None)
102 return metric
103 return run
104
105def main():
106 grid=[{'lr':lr,'wd':wd} for lr in [1e-3,3e-3,5e-3] for wd in [0.,1e-4]]
107 base=sweep_baseline(baseline_fn,grid)
108 # Same lr union is evaluated by baseline; idea has 3 nearby settings at best wd.
109 wd=base['best_cfg']['wd']; idea_cfgs=[{'lr':lr,'wd':wd} for lr in [1e-3,3e-3,5e-3]]
110 idea_cfg_results=[]
111 for cfg in idea_cfgs:
112 r=evaluate(lambda s,cfg=cfg:train_idea(s,cfg['lr'],cfg['wd']),seeds=range(8))
113 idea_cfg_results.append((cfg,r))
114 idea_cfg,best=min(idea_cfg_results,key=lambda q:q[1]['mean'])
115 sig=[]
116 for s in range(8): sig.append(train_idea(s,idea_cfg['lr'],idea_cfg['wd'],True)[1])
117 sigkeys=sig[0].keys(); signature={k:{'observed_mean':float(np.mean([a[k] for a in sig])),'per_seed': [float(a[k]) for a in sig]} for k in sigkeys}
118 signature.update({'prediction':'resolver branch weights vary continuously with counterfactual action consequences; measured on trained models','confirmed':False})
119 rep=make_report('dynamics','rnn_small',base,best,{'math_check':math_check(),'trained_behavior':signature,'idea_sweep':[{'cfg':c,'mean':r['mean']} for c,r in idea_cfg_results]})
120 rep['protocol_notes']={'n_train':NTR,'n_test':NTE,'epochs':EPOCHS,'structural_match':'controlled pendulum dynamics; action branches are final-step control perturbations','baseline_knobs_swept':['lr','weight_decay']}
121 Path('bench_report.json').write_text(json.dumps(rep,indent=2))
122 print(json.dumps(rep,indent=2))
123if __name__=='__main__': main()