import sys, json, random, math from pathlib import Path import numpy as np import torch sys.path.insert(0, '/home/maxwelhelp/all/math2nn') import bench from bench import train_model, sweep_baseline, make_report SEEDS = tuple(range(8)) # Union is shared: every idea learning rate is evaluated by the baseline sweep. LR_GRID = [0.0015, 0.003, 0.006] EPOCHS = 3 BATCH = 64 HIDDEN = 64 M = 32 T_OBS = 2 LAMBDA = 0.001 EPS = 1e-3 def seed_all(seed): random.seed(seed); np.random.seed(seed); torch.manual_seed(seed) if torch.cuda.is_available(): torch.cuda.manual_seed_all(seed) def new_model(): return bench.make_model('rnn_small', (24,), 1) def baseline_fn(cfg): def run(seed): seed_all(seed) ds = bench.get_dataset('dynamics', seed=seed, n_train=200, n_test=200) model = new_model() _, metric, _ = train_model(model, ds, epochs=EPOCHS, lr=cfg['lr'], batch=BATCH, weight_decay=0.0, log=lambda *a: None) return float(metric) return run def gru_observation_jacobian(model, flat_x, m=M, tsteps=T_OBS, create_graph=False): """J of selected hidden coordinates over a real trained GRU trajectory wrt h0.""" seq = flat_x.reshape(1, -1, 3) # This function is intentionally model-dependent, not an analytic toy graph. def obs(h0): h = h0.reshape(1, 1, -1) ys = [] for t in range(min(tsteps, seq.shape[1])): _, h = model.rnn(seq[:, t:t+1, :], h) ys.append(h[0, 0, :m]) return torch.cat(ys) h0 = torch.zeros(HIDDEN, device=flat_x.device, dtype=flat_x.dtype, requires_grad=True) old = torch.backends.cudnn.enabled torch.backends.cudnn.enabled = False try: J = torch.autograd.functional.jacobian(obs, h0, create_graph=create_graph) finally: torch.backends.cudnn.enabled = old return J def obs_penalty(model, flat_x): J = gru_observation_jacobian(model, flat_x, M, T_OBS, True) G = J.T @ J + EPS * torch.eye(J.shape[1], device=J.device, dtype=J.dtype) return -torch.linalg.slogdet(G)[1] def idea_train(model, ds, lr): # Own loop is required because the intervention is a new training loss. ladder = ['cuda', 'cpu'] if torch.cuda.is_available() else ['cpu'] last = None for dev in ladder: try: model = model.to(dev) x, y = ds['xtr'].to(dev), ds['ytr'].to(dev) opt = torch.optim.Adam(model.parameters(), lr=lr) model.train() for ep in range(EPOCHS): perm = torch.randperm(len(x), device=dev) for bi in range(0, len(x), BATCH): ix = perm[bi:bi+BATCH] pred = model(x[ix]) task = (pred - y[ix]).pow(2).mean() # One representative trajectory per minibatch keeps this small. reg = obs_penalty(model, x[ix[0]]) if (ep == EPOCHS-1 and bi == 0) else torch.zeros((), device=dev) loss = task + LAMBDA * reg opt.zero_grad(set_to_none=True); loss.backward() torch.nn.utils.clip_grad_norm_(model.parameters(), 5.0) opt.step() model.eval() with torch.no_grad(): metric = float((model(x if False else ds['xte'].to(dev)) - ds['yte'].to(dev)).pow(2).mean().cpu()) return model, metric except Exception as e: last = e try: model = model.to('cpu') except Exception: pass raise RuntimeError(last) def idea_fn(cfg, keep=False): def run(seed): seed_all(seed) ds = bench.get_dataset('dynamics', seed=seed, n_train=200, n_test=200) model, metric = idea_train(new_model(), ds, cfg['lr']) if keep: kept[(seed, cfg['lr'])] = (model, ds) return float(metric) return run def evaluate8(fn): vals = [float(fn(s)) for s in SEEDS] return {'per_seed': vals, 'mean': float(np.mean(vals)), 'std': float(np.std(vals, ddof=1))} def signature(model, ds): model.eval(); x = ds['xte'][0].to(next(model.parameters()).device) rows=[] with torch.no_grad(): pass for m in [8, 16, 32]: J = gru_observation_jacobian(model, x, m, 2, False).detach() sv = torch.linalg.svdvals(J) rank = int((sv > 1e-5).sum().item()) G = J.T @ J + EPS*torch.eye(HIDDEN, device=J.device) rows.append({'m':m, 'predicted_rank_upper_bound_T1':min(HIDDEN,2*m), 'observed_rank_T1':rank, 'observed_smin':float(sv[-1]), 'observed_logdet':float(torch.linalg.slogdet(G)[1])}) return {'state_dim':HIDDEN, 'horizon_T':1, 'predicted_counting_threshold_m':HIDDEN//2, 'observed':rows, 'confirmed': any(r['m']==HIDDEN//2 and r['observed_rank_T1']>=HIDDEN for r in rows)} kept={} def main(): grid=[{'lr':v} for v in LR_GRID] base=sweep_baseline(baseline_fn, grid, seeds=(0,1,2,3)) idea_results={} for cfg in grid: idea_results[str(cfg['lr'])]=evaluate8(idea_fn(cfg, keep=(cfg['lr']==base['best_cfg']['lr']))) best_key=min(idea_results, key=lambda k: idea_results[k]['mean']) idea=idea_results[best_key] # Recover a trained baseline model at the selected configuration for signature. bmodel=new_model(); seed_all(0) bds=bench.get_dataset('dynamics',0,n_train=400,n_test=400) bmodel,_,_=train_model(bmodel,bds,epochs=EPOCHS,lr=base['best_cfg']['lr'],batch=BATCH,log=lambda *a:None) imodel, ids = kept.get((0, base['best_cfg']['lr']), (None,None)) if imodel is None: imodel, _ = idea_train(new_model(), bds, base['best_cfg']['lr']); ids=bds extra={'prediction':'For T=1, m >= n/2 is the counting threshold for full rank.', 'baseline_trained_model':signature(bmodel,bds), 'idea_trained_model':signature(imodel,ids)} report=make_report('dynamics','rnn_small',base,idea,extra) report['idea_sweep']=idea_results report['protocol_notes']='Baseline sweep uses the same three learning rates as the idea; idea adds only finite-horizon logdet loss.' Path('bench_report.json').write_text(json.dumps(report,indent=2)) print(json.dumps(report,indent=2)) if __name__=='__main__': main()