Koopman Deadline Controller / koopman_bench.py

Failed on benchmark

Raw ⬇ ZIP
 1import sys, math, json
 2from pathlib import Path
 3import numpy as np
 4import torch
 5sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
 6from bench import get_dataset, make_model, train_model, sweep_baseline, make_report
 7
 8TRACK='dynamics'; MODEL='rnn_small'; EPOCHS=8; BATCH=128
 9SEEDS=tuple(range(8)); SWEEP_SEEDS=tuple(range(4))
10# The union of all idea and baseline settings is evaluated by the baseline sweep.
11GRID=[{'lr':1e-3,'rounds':4},{'lr':3e-3,'rounds':4},{'lr':1e-2,'rounds':4}]
12
13# Iterative refinement map: each round applies the same trained rnn to a shifted
14# state window. The controller uses only the model's observed output trajectory.
15def rollout(net, x, rounds=4, controller=False, eps=0.01, delta=0.02):
16    device=next(net.parameters()).device
17    q=x.to(device).clone(); preds=[]; changes=[]; transitions=[]; stable=0
18    with torch.no_grad():
19        for t in range(rounds):
20            pred=net(q); val=pred[:,0]
21            if preds:
22                change=float(torch.mean(torch.abs(pred-preds[-1])).item())
23                changes.append(change); transitions.append((changes[-2] if len(changes)>1 else change, change))
24                stable=stable+1 if change < delta else 0
25            preds.append(pred.detach())
26            q=q.clone(); q[:,:-3]=q[:,3:]; q[:,-3]=val
27            if controller and len(changes)>=2 and stable>=2:
28                # Fit the scalar Koopman map for the observed disagreement D_t.
29                d0=np.asarray([a for a,b in transitions],dtype=float)
30                d1=np.asarray([b for a,b in transitions],dtype=float)
31                X=np.c_[d0,np.ones_like(d0)]
32                coef=np.linalg.lstsq(X,d1,rcond=None)[0]
33                lam=abs(float(coef[0])); D=changes[-1]
34                remaining=max(0,rounds-t-1)
35                future=D*(lam**remaining) if 0 < lam < 1 else D
36                # Conservative gate: only stop when predicted disagreement and
37                # recent observed changes are both below tolerance.
38                if future <= eps and change < delta:
39                    break
40    arr=torch.stack(preds,dim=1)[:, :, 0]
41    return arr, len(preds), changes
42
43def train_score(seed, lr, controller, rounds=4):
44    torch.manual_seed(seed); np.random.seed(seed)
45    ds=get_dataset(TRACK, seed, n_train=400, n_test=400)
46    net=make_model(MODEL, ds['input_shape'], ds['out_dim'])
47    net, _, _=train_model(net, ds, epochs=EPOCHS, lr=lr, batch=BATCH, log=lambda *_: None)
48    if net is None: return float('nan'), {'rounds':float('nan')}, None
49    net.eval(); device=next(net.parameters()).device
50    pred,n,ch=rollout(net,ds['xte'],rounds,controller=controller)
51    y=ds['yte'][:,0].to(device)
52    mse=float(torch.mean((pred[:,-1]-y)**2).item())
53    return mse, {'rounds':n,'changes':ch}, net
54
55def one(seed,cfg,controller): return train_score(seed,cfg['lr'],controller,cfg['rounds'])
56def baseline_sweep():
57    def factory(cfg):
58        return lambda seed: one(int(seed),cfg,False)[0]
59    return sweep_baseline(factory,GRID,seeds=SWEEP_SEEDS)
60
61def idea_eval(cfg,seeds=SEEDS):
62    vals=[]; aux=[]
63    for s in seeds:
64        m,a,_=one(s,cfg,True); vals.append(m); aux.append(a)
65    return {'mean':float(np.mean(vals)),'std':float(np.std(vals)),'per_seed':vals,'n':len(vals)},aux
66
67def signature(cfg):
68    fitted=[]; observed=[]
69    for s in SEEDS:
70        _,_,net=train_score(s,cfg['lr'],False,cfg['rounds'])
71        ds=get_dataset(TRACK,s,n_train=400,n_test=80); net.eval()
72        _,_,ch=rollout(net,ds['xte'],cfg['rounds'],False)
73        if len(ch)>=2:
74            x=np.asarray(ch[:-1]); y=np.asarray(ch[1:])
75            lam=float(np.linalg.lstsq(np.c_[x,np.ones_like(x)],y,rcond=None)[0][0])
76            fitted.append(abs(lam)); observed.append(float(np.mean(y/np.maximum(x,1e-8))))
77    pf=float(np.mean(fitted)); po=float(np.mean(observed))
78    return {'quantity':'NN observed disagreement contraction','predicted_mean_ratio':pf,'observed_mean_ratio':po,'absolute_error':abs(pf-po),'n_trajectories':len(fitted),'confirmed':bool(abs(pf-po)<0.10)}
79
80def main():
81    sweep=baseline_sweep(); base_cfg=sweep['best_cfg']
82    idea_grid=[]
83    for cfg in GRID:
84        r,_=idea_eval(cfg); idea_grid.append({'cfg':cfg,'mean':r['mean']})
85    best_idea_cfg=min(idea_grid,key=lambda z:z['mean'])['cfg']
86    idea_res,aux=idea_eval(best_idea_cfg,SEEDS)
87    base_block={'sweep':sweep,'best_cfg':base_cfg,'full':sweep['full']}
88    rep=make_report(TRACK,MODEL,base_block,idea_res,{'mechanism_signature':signature(best_idea_cfg),'protocol_note':'Matched dynamics task and freshly trained identical rnn_small architecture; only the inference stopping rule differs.'})
89    rep['idea_sweep']=idea_grid; rep['idea_best_cfg']=best_idea_cfg; rep['idea_auxiliary_rounds']=aux
90    Path('bench_report.json').write_text(json.dumps(rep,indent=2)); print(json.dumps(rep,indent=2))
91if __name__=='__main__': main()