Tangential Bellman Tie Resolver / bench_stage2.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
  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()