Transfer-Spectrum Pseudo-Transition Scheduler / bench_experiment.py

Mechanism confirmed, baseline not beaten

Raw ⬇ ZIP
  1import sys, json, math, random
  2import numpy as np
  3import torch
  4import torch.nn as nn
  5sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
  6from bench import get_dataset, make_model, sweep_baseline, evaluate, make_report
  7
  8TRACK, MODEL = 'dynamics', 'rnn_small'
  9SEEDS, SWEEP_SEEDS = tuple(range(8)), (0, 1, 2, 3)
 10EPOCHS, BATCH = 10, 128
 11# This is the complete shared union: baseline is evaluated at every idea lr/decay.
 12GRID = [{'lr': lr, 'weight_decay': wd} for lr in (1e-3, 3e-3, 1e-2)
 13        for wd in (0.0, 1e-4)]
 14
 15
 16def seed_all(seed):
 17    random.seed(seed); np.random.seed(seed); torch.manual_seed(seed)
 18    if torch.cuda.is_available(): torch.cuda.manual_seed_all(seed)
 19
 20
 21def transfer_stats(net, x, device):
 22    """Empirical local linear transfer T from hidden trajectories of trained GRU."""
 23    net.eval()
 24    with torch.no_grad():
 25        seq = x.to(device).view(x.shape[0], -1, 3)
 26        h = torch.zeros(1, seq.shape[0], net.rnn.hidden_size, device=device)
 27        hs = []
 28        for t in range(seq.shape[1]):
 29            _, h = net.rnn(seq[:, t:t+1], h); hs.append(h[-1])
 30        A = torch.stack(hs[:-1], 1).reshape(-1, net.rnn.hidden_size)
 31        B = torch.stack(hs[1:], 1).reshape(-1, net.rnn.hidden_size)
 32        eye = torch.eye(A.shape[1], device=device)
 33        T = torch.linalg.solve(A.T @ A + 1e-3 * eye, A.T @ B)
 34        ev, vec = torch.linalg.eig(T)
 35        order = torch.argsort(ev.abs(), descending=True); ev, vec = ev[order], vec[:, order]
 36        mags = ev.abs().real; l0 = max(float(mags[0]), 1e-8)
 37        l1 = float(mags[1]) if len(mags) > 1 else 0.0
 38        gap = (l0-l1)/l0
 39        xi = 1.0/max(math.log(l0/max(l1, 1e-8)), 1e-8)
 40        V = vec[:, :min(3, vec.shape[1])].real
 41        z = A @ V; p = (z*z).mean(0); p = p/(p.sum()+1e-8)
 42        entropy = float(-(p*torch.log(p+1e-8)).sum())
 43        return gap, entropy, xi
 44
 45
 46def run(seed, cfg, scheduler, return_state=False):
 47    seed_all(seed)
 48    ds = get_dataset(TRACK, seed, n_train=4000, n_test=1000)
 49    device = 'cuda' if torch.cuda.is_available() else 'cpu'
 50    try:
 51        net = make_model(MODEL, ds['input_shape'], 1).to(device)
 52        opt = torch.optim.Adam(net.parameters(), lr=cfg['lr'], weight_decay=cfg['weight_decay'])
 53        lossf = nn.MSELoss(); xtr, ytr = ds['xtr'].to(device), ds['ytr'].to(device)
 54        current_lr = float(cfg['lr']); stats=[]
 55        for _ in range(EPOCHS):
 56            net.train(); perm=torch.randperm(len(xtr), device=device)
 57            for i in range(0, len(xtr), BATCH):
 58                idx=perm[i:i+BATCH]; loss=lossf(net(xtr[idx]), ytr[idx])
 59                opt.zero_grad(set_to_none=True); loss.backward()
 60                torch.nn.utils.clip_grad_norm_(net.parameters(), 5.0); opt.step()
 61            if scheduler:
 62                gap, entropy, xi=transfer_stats(net, xtr[:min(512,len(xtr))], device)
 63                stats.append((gap, entropy, xi))
 64                current_lr=float(np.clip(current_lr*np.clip(gap/.25,.5,1.25), 1e-5, cfg['lr']))
 65                if entropy>.70 and gap<.12: current_lr=max(.5*current_lr,1e-5)
 66                for pg in opt.param_groups: pg['lr']=current_lr
 67        net.eval()
 68        with torch.no_grad(): metric=float(((net(ds['xte'].to(device))-ds['yte'].to(device))**2).mean())
 69        state={'gap':stats[-1][0], 'entropy':stats[-1][1], 'xi':stats[-1][2],
 70               'lr_final':current_lr, 'trigger_count':sum(g<.12 and h>.70 for g,h,_ in stats)} if stats else {}
 71        return (metric, state) if return_state else metric
 72    except RuntimeError:
 73        if device == 'cuda':
 74            torch.cuda.empty_cache(); old=torch.cuda.is_available; torch.cuda.is_available=lambda:False
 75            try: return run(seed,cfg,scheduler,return_state)
 76            finally: torch.cuda.is_available=old
 77        raise
 78
 79
 80def main():
 81    base=sweep_baseline(lambda c: lambda s: run(s,c,False), GRID, seeds=SWEEP_SEEDS)
 82    # Three idea settings, all present in baseline sweep; report the best idea.
 83    idea_candidates=[]
 84    for cfg in GRID:
 85        r=evaluate(lambda s, c=cfg: run(s,c,True), seeds=SEEDS)
 86        idea_candidates.append({'cfg':cfg, 'result':r})
 87    best=min(idea_candidates, key=lambda z:z['result']['mean'])
 88    idea=best['result']; sig=[]
 89    for s in SEEDS: _, st=run(s,best['cfg'],True,return_state=True); sig.append(st)
 90    gaps=np.array([x['gap'] for x in sig]); ents=np.array([x['entropy'] for x in sig])
 91    corr=float(np.corrcoef(gaps,ents)[0,1]) if np.std(gaps)>0 and np.std(ents)>0 else 0.0
 92    signature={'prediction':'low normalized transfer gap with high projection entropy marks a pseudo-transition',
 93      'predicted_thresholds':{'gap_lt':0.12,'entropy_gt':0.70},
 94      'observed_mean_gap':float(gaps.mean()),'observed_mean_entropy':float(ents.mean()),
 95      'gap_entropy_correlation':corr,'trigger_count':int(sum(x['trigger_count'] for x in sig)),
 96      'confirmed':bool(np.any(gaps<.12) and np.any(ents>.70))}
 97    base['idea_grid']= [{'cfg':z['cfg'],'mean':z['result']['mean']} for z in idea_candidates]
 98    report=make_report(TRACK,MODEL,base,idea,signature)
 99    report['idea']['best_cfg']=best['cfg']
100    with open('bench_report.json','w') as f: json.dump(report,f,indent=2)
101    print(json.dumps(report,indent=2))
102
103if __name__=='__main__': main()