Finite-Excitation Latent Replay / bench_felr.py

Failed on benchmark

Raw ⬇ ZIP
  1import sys, json, math, random
  2from pathlib import Path
  3import numpy as np
  4import torch
  5import torch.nn as nn
  6
  7sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
  8from bench import get_dataset, make_model, train_model, evaluate, sweep_baseline, make_report
  9
 10OUT = Path('bench_report.json')
 11SEEDS = tuple(range(8))
 12# Same learning-rate union is used by baseline and idea.
 13LR_GRID = [1e-3, 3e-3, 1e-2]
 14EPOCHS = 14
 15NTRAIN, NTEST = 900, 300
 16BATCH = 128
 17
 18
 19def seed_all(seed):
 20    random.seed(seed); np.random.seed(seed); torch.manual_seed(seed)
 21    if torch.cuda.is_available():
 22        try: torch.cuda.manual_seed_all(seed)
 23        except Exception: pass
 24
 25
 26def baseline_metric(cfg, seed):
 27    seed_all(seed)
 28    d = get_dataset('dynamics', seed, n_train=NTRAIN, n_test=NTEST)
 29    try:
 30        _, metric, _ = train_model(make_model('rnn_small', d['input_shape'], d['out_dim']),
 31                                   d, epochs=EPOCHS, lr=cfg['lr'], batch=BATCH,
 32                                   weight_decay=cfg['weight_decay'], log=lambda *_: None)
 33        return float(metric)
 34    except (RuntimeError, torch.cuda.OutOfMemoryError):
 35        return _baseline_cpu(cfg, seed)
 36
 37
 38def _baseline_cpu(cfg, seed):
 39    # train_model already has fallback; this path is only a defensive retry.
 40    seed_all(seed); d = get_dataset('dynamics', seed, n_train=NTRAIN, n_test=NTEST)
 41    old = torch.cuda.is_available
 42    try:
 43        net = make_model('rnn_small', d['input_shape'], d['out_dim']).cpu()
 44        opt = torch.optim.Adam(net.parameters(), lr=cfg['lr'], weight_decay=cfg['weight_decay'])
 45        for _ in range(EPOCHS):
 46            p = torch.randperm(len(d['xtr']))
 47            for j in range(0, len(p), BATCH):
 48                ix=p[j:j+BATCH]; loss=((net(d['xtr'][ix])-d['ytr'][ix])**2).mean()
 49                opt.zero_grad(); loss.backward(); opt.step()
 50        with torch.no_grad(): return float(((net(d['xte'])-d['yte'])**2).mean())
 51    finally:
 52        pass
 53
 54
 55def integral_regressor(x):
 56    """Omega = dt sum Phi(theta, omega, u), with Phi=[theta,omega,u,sin(theta)]."""
 57    z=x.reshape(x.shape[0], 8, 3)
 58    phi=torch.stack((z[:,:,0], z[:,:,1], z[:,:,2], torch.sin(z[:,:,0])), dim=-1)
 59    return phi.mean(dim=1)
 60
 61
 62def gram_score(omega, eps=0.035):
 63    # Computable conservative certificate q=lambda_min(Ghat)-sum(2||O||e+e^2).
 64    g=omega.T @ omega
 65    lam=torch.linalg.eigvalsh(g)[0]
 66    err=2*torch.linalg.matrix_norm(omega, ord=2)*eps + eps*eps
 67    return lam - err
 68
 69
 70def idea_train(cfg, seed, return_info=False):
 71    seed_all(seed)
 72    d=get_dataset('dynamics', seed, n_train=NTRAIN, n_test=NTEST)
 73    device=torch.device('cuda' if torch.cuda.is_available() else 'cpu')
 74    try:
 75        net=make_model('rnn_small', d['input_shape'], d['out_dim']).to(device)
 76        xtr,ytr=d['xtr'].to(device),d['ytr'].to(device)
 77        xte,yte=d['xte'].to(device),d['yte'].to(device)
 78        opt=torch.optim.Adam(net.parameters(),lr=cfg['lr'],weight_decay=cfg['weight_decay'])
 79        rng=torch.Generator(device=device); rng.manual_seed(seed+10000)
 80        # fixed-size history; greedy replacement maximizes current minimum eigenvalue
 81        history=[]; activated=[]; losses=[]; qvals=[]; activation_epoch=None
 82        for ep in range(EPOCHS):
 83            net.train(); perm=torch.randperm(len(xtr),generator=rng,device=device); ep_loss=0.
 84            for j in range(0,len(perm),BATCH):
 85                ix=perm[j:j+BATCH]; xb,yb=xtr[ix],ytr[ix]
 86                om=integral_regressor(xb); q=gram_score(om, cfg['eps'])
 87                qvals.append(float(q.detach().cpu()))
 88                exciting=bool(q.item()>cfg['gamma'])
 89                if exciting:
 90                    if activation_epoch is None: activation_epoch=ep
 91                    activated.append(1)
 92                    # Replay selected history plus current batch, not arbitrary old data.
 93                    cand=(float(torch.linalg.eigvalsh(om.T@om)[0].detach().cpu()), xb.detach(), yb.detach(), om.detach())
 94                    history.append(cand); history.sort(key=lambda a:a[0],reverse=True); history=history[:cfg['replay']]
 95                    batches=[(xb,yb)] + [(h[1],h[2]) for h in history[:-1]]
 96                    opt.zero_grad(); loss=sum(((net(a)-b)**2).mean() for a,b in batches)/len(batches)
 97                    loss.backward(); opt.step()
 98                else:
 99                    activated.append(0)
100                    # Conservative policy: freeze adapter update until finite excitation.
101                    loss=torch.zeros((),device=device)
102                ep_loss += float(loss.detach().cpu())*len(ix)
103            losses.append(ep_loss/len(xtr))
104        net.eval()
105        with torch.no_grad(): metric=float(((net(xte)-yte)**2).mean().cpu())
106        info={'metric':metric,'activation_epoch':activation_epoch,
107              'activation_rate':float(np.mean(activated)),'mean_q':float(np.mean(qvals)),
108              'positive_q_rate':float(np.mean(np.asarray(qvals)>cfg['gamma'])),
109              'loss_first':losses[0],'loss_last':losses[-1],
110              'post_activation_loss_drop': (float(losses[activation_epoch]-losses[-1]) if activation_epoch is not None else 0.0),
111              'model_params':sum(p.numel() for p in net.parameters())}
112        return info if return_info else metric
113    except (RuntimeError, torch.cuda.OutOfMemoryError):
114        # Shared GPU can fail; repeat entirely on CPU.
115        torch.cuda.empty_cache() if torch.cuda.is_available() else None
116        old=torch.cuda.is_available
117        # identical loop through a temporary CPU-only recursive implementation
118        if device.type=='cuda':
119            torch.cuda.is_available=lambda: False
120            try: return idea_train(cfg,seed,return_info)
121            finally: torch.cuda.is_available=old
122        raise
123
124
125def main():
126    base_grid=[{'lr':lr,'weight_decay':wd} for lr in LR_GRID for wd in [0.0,1e-4]]
127    # Baseline decisive Adam knob (lr and weight decay) is swept. Idea uses same union.
128    base=sweep_baseline(lambda cfg: lambda s: baseline_metric(cfg,s), base_grid, seeds=(0,1,2,3))
129    idea_cfgs=[{'lr':lr,'weight_decay':base['best_cfg']['weight_decay'], 'gamma':g, 'eps':0.035, 'replay':4}
130               for lr in LR_GRID for g in [0.0]]
131    # Evaluate the three idea settings on all paired seeds; choose by the same 4-seed tuning split.
132    idea_trials=[]
133    for cfg in idea_cfgs:
134        r=evaluate(lambda s,cfg=cfg: idea_train(cfg,s), seeds=(0,1,2,3))
135        idea_trials.append({'cfg':cfg,'mean':r['mean']})
136    best=min(idea_trials,key=lambda z:z['mean'])['cfg']
137    idea=evaluate(lambda s: idea_train(best,s), seeds=SEEDS)
138    # Signature is measured from trained models, not the algebraic toy: aggregate per-seed behavior.
139    sig=[]
140    for s in SEEDS:
141        sig.append(idea_train(best,s,True))
142    signature={'prediction':'parameter updates become active after q>gamma and loss then decreases',
143      'predicted_vs_observed':{'predicted_positive_q_activation':True,
144        'observed_positive_q_rate_mean':float(np.mean([z['positive_q_rate'] for z in sig])),
145        'observed_activation_rate_mean':float(np.mean([z['activation_rate'] for z in sig])),
146        'observed_post_activation_loss_drop_mean':float(np.mean([z['post_activation_loss_drop'] for z in sig])),
147        'activation_epoch_values':[z['activation_epoch'] for z in sig]},
148      'confirmed':bool(np.mean([z['positive_q_rate'] for z in sig])>0 and np.mean([z['post_activation_loss_drop'] for z in sig])>0)}
149    # Report idea sweep alongside canonical make_report output.
150    rep=make_report('dynamics','rnn_small',base,idea,signature)
151    rep['idea_sweep']=idea_trials; rep['protocol']={'seeds':list(SEEDS),'n_train':NTRAIN,'n_test':NTEST,'epochs':EPOCHS,
152      'structural_match':'controlled pendulum rollout / latent dynamics', 'custom_track':None}
153    OUT.write_text(json.dumps(rep,indent=2))
154    print(json.dumps(rep,indent=2))
155
156if __name__=='__main__': main()