Finite-Horizon Hidden-State Observability Regularizer / observability_bench.py

Failed on benchmark

Raw ⬇ ZIP
  1import sys, json, random, math
  2from pathlib import Path
  3import numpy as np
  4import torch
  5sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
  6import bench
  7from bench import train_model, sweep_baseline, make_report
  8
  9SEEDS = tuple(range(8))
 10# Union is shared: every idea learning rate is evaluated by the baseline sweep.
 11LR_GRID = [0.0015, 0.003, 0.006]
 12EPOCHS = 3
 13BATCH = 64
 14HIDDEN = 64
 15M = 32
 16T_OBS = 2
 17LAMBDA = 0.001
 18EPS = 1e-3
 19
 20def seed_all(seed):
 21    random.seed(seed); np.random.seed(seed); torch.manual_seed(seed)
 22    if torch.cuda.is_available(): torch.cuda.manual_seed_all(seed)
 23
 24def new_model():
 25    return bench.make_model('rnn_small', (24,), 1)
 26
 27def baseline_fn(cfg):
 28    def run(seed):
 29        seed_all(seed)
 30        ds = bench.get_dataset('dynamics', seed=seed, n_train=200, n_test=200)
 31        model = new_model()
 32        _, metric, _ = train_model(model, ds, epochs=EPOCHS, lr=cfg['lr'],
 33                                   batch=BATCH, weight_decay=0.0, log=lambda *a: None)
 34        return float(metric)
 35    return run
 36
 37def gru_observation_jacobian(model, flat_x, m=M, tsteps=T_OBS, create_graph=False):
 38    """J of selected hidden coordinates over a real trained GRU trajectory wrt h0."""
 39    seq = flat_x.reshape(1, -1, 3)
 40    # This function is intentionally model-dependent, not an analytic toy graph.
 41    def obs(h0):
 42        h = h0.reshape(1, 1, -1)
 43        ys = []
 44        for t in range(min(tsteps, seq.shape[1])):
 45            _, h = model.rnn(seq[:, t:t+1, :], h)
 46            ys.append(h[0, 0, :m])
 47        return torch.cat(ys)
 48    h0 = torch.zeros(HIDDEN, device=flat_x.device, dtype=flat_x.dtype, requires_grad=True)
 49    old = torch.backends.cudnn.enabled
 50    torch.backends.cudnn.enabled = False
 51    try:
 52        J = torch.autograd.functional.jacobian(obs, h0, create_graph=create_graph)
 53    finally:
 54        torch.backends.cudnn.enabled = old
 55    return J
 56
 57def obs_penalty(model, flat_x):
 58    J = gru_observation_jacobian(model, flat_x, M, T_OBS, True)
 59    G = J.T @ J + EPS * torch.eye(J.shape[1], device=J.device, dtype=J.dtype)
 60    return -torch.linalg.slogdet(G)[1]
 61
 62def idea_train(model, ds, lr):
 63    # Own loop is required because the intervention is a new training loss.
 64    ladder = ['cuda', 'cpu'] if torch.cuda.is_available() else ['cpu']
 65    last = None
 66    for dev in ladder:
 67        try:
 68            model = model.to(dev)
 69            x, y = ds['xtr'].to(dev), ds['ytr'].to(dev)
 70            opt = torch.optim.Adam(model.parameters(), lr=lr)
 71            model.train()
 72            for ep in range(EPOCHS):
 73                perm = torch.randperm(len(x), device=dev)
 74                for bi in range(0, len(x), BATCH):
 75                    ix = perm[bi:bi+BATCH]
 76                    pred = model(x[ix])
 77                    task = (pred - y[ix]).pow(2).mean()
 78                    # One representative trajectory per minibatch keeps this small.
 79                    reg = obs_penalty(model, x[ix[0]]) if (ep == EPOCHS-1 and bi == 0) else torch.zeros((), device=dev)
 80                    loss = task + LAMBDA * reg
 81                    opt.zero_grad(set_to_none=True); loss.backward()
 82                    torch.nn.utils.clip_grad_norm_(model.parameters(), 5.0)
 83                    opt.step()
 84            model.eval()
 85            with torch.no_grad(): metric = float((model(x if False else ds['xte'].to(dev)) - ds['yte'].to(dev)).pow(2).mean().cpu())
 86            return model, metric
 87        except Exception as e:
 88            last = e
 89            try: model = model.to('cpu')
 90            except Exception: pass
 91    raise RuntimeError(last)
 92
 93def idea_fn(cfg, keep=False):
 94    def run(seed):
 95        seed_all(seed)
 96        ds = bench.get_dataset('dynamics', seed=seed, n_train=200, n_test=200)
 97        model, metric = idea_train(new_model(), ds, cfg['lr'])
 98        if keep: kept[(seed, cfg['lr'])] = (model, ds)
 99        return float(metric)
100    return run
101
102def evaluate8(fn):
103    vals = [float(fn(s)) for s in SEEDS]
104    return {'per_seed': vals, 'mean': float(np.mean(vals)), 'std': float(np.std(vals, ddof=1))}
105
106def signature(model, ds):
107    model.eval(); x = ds['xte'][0].to(next(model.parameters()).device)
108    rows=[]
109    with torch.no_grad():
110        pass
111    for m in [8, 16, 32]:
112        J = gru_observation_jacobian(model, x, m, 2, False).detach()
113        sv = torch.linalg.svdvals(J)
114        rank = int((sv > 1e-5).sum().item())
115        G = J.T @ J + EPS*torch.eye(HIDDEN, device=J.device)
116        rows.append({'m':m, 'predicted_rank_upper_bound_T1':min(HIDDEN,2*m),
117                     'observed_rank_T1':rank, 'observed_smin':float(sv[-1]),
118                     'observed_logdet':float(torch.linalg.slogdet(G)[1])})
119    return {'state_dim':HIDDEN, 'horizon_T':1,
120            'predicted_counting_threshold_m':HIDDEN//2,
121            'observed':rows,
122            'confirmed': any(r['m']==HIDDEN//2 and r['observed_rank_T1']>=HIDDEN for r in rows)}
123
124kept={}
125def main():
126    grid=[{'lr':v} for v in LR_GRID]
127    base=sweep_baseline(baseline_fn, grid, seeds=(0,1,2,3))
128    idea_results={}
129    for cfg in grid:
130        idea_results[str(cfg['lr'])]=evaluate8(idea_fn(cfg, keep=(cfg['lr']==base['best_cfg']['lr'])))
131    best_key=min(idea_results, key=lambda k: idea_results[k]['mean'])
132    idea=idea_results[best_key]
133    # Recover a trained baseline model at the selected configuration for signature.
134    bmodel=new_model(); seed_all(0)
135    bds=bench.get_dataset('dynamics',0,n_train=400,n_test=400)
136    bmodel,_,_=train_model(bmodel,bds,epochs=EPOCHS,lr=base['best_cfg']['lr'],batch=BATCH,log=lambda *a:None)
137    imodel, ids = kept.get((0, base['best_cfg']['lr']), (None,None))
138    if imodel is None:
139        imodel, _ = idea_train(new_model(), bds, base['best_cfg']['lr']); ids=bds
140    extra={'prediction':'For T=1, m >= n/2 is the counting threshold for full rank.',
141           'baseline_trained_model':signature(bmodel,bds),
142           'idea_trained_model':signature(imodel,ids)}
143    report=make_report('dynamics','rnn_small',base,idea,extra)
144    report['idea_sweep']=idea_results
145    report['protocol_notes']='Baseline sweep uses the same three learning rates as the idea; idea adds only finite-horizon logdet loss.'
146    Path('bench_report.json').write_text(json.dumps(report,indent=2))
147    print(json.dumps(report,indent=2))
148if __name__=='__main__': main()