Spectral Basin Allocation for Multimodal Neural Memories / bench_experiment.py

Mechanism confirmed, baseline not beaten

Raw ⬇ ZIP
  1import sys, json, random
  2from pathlib import Path
  3import numpy as np
  4import torch
  5import torch.nn as nn
  6sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
  7from bench import get_dataset, make_model, train_model, evaluate, sweep_baseline, make_report
  8
  9SEEDS = tuple(range(8))
 10# Same lr union on both sides; baseline also sweeps its central Adam knob.
 11LR_GRID = (1e-3, 3e-3, 6e-3)
 12BASE_GRID = [{'lr': lr, 'weight_decay': wd, 'epochs': 10} for lr in LR_GRID for wd in (0.0, 1e-4)]
 13IDEA_GRID = [{'lr': lr, 'weight_decay': 0.0, 'epochs': 10, 'lam': 0.02} for lr in LR_GRID]
 14
 15
 16def seed_all(seed):
 17    random.seed(seed); np.random.seed(seed); torch.manual_seed(seed)
 18    if torch.cuda.is_available():
 19        torch.cuda.manual_seed_all(seed)
 20
 21
 22def device():
 23    return 'cuda' if torch.cuda.is_available() else 'cpu'
 24
 25
 26def spectral_rate(phi, K=1.0):
 27    """Differentiable smallest non-gauge eigenvalue for an 8-node ring.
 28    phi: [batch, 8]. The symmetric composite Laplacian is PSD when locked.
 29    A unit ring is used, matching one shared phase-delay graph for all samples.
 30    """
 31    n = phi.shape[1]
 32    A = torch.zeros((n, n), device=phi.device, dtype=phi.dtype)
 33    idx = torch.arange(n, device=phi.device)
 34    A[idx, (idx + 1) % n] = 1.0
 35    A[idx, (idx - 1) % n] = 1.0
 36    C = A[None] * torch.cos(phi[:, None, :] - phi[:, :, None])
 37    L = torch.diag_embed(C.sum(-1)) - C
 38    ev = torch.linalg.eigvalsh(L)
 39    return K * ev[:, 1], ev
 40
 41
 42def forward_hidden(net, x):
 43    seq = x.view(x.shape[0], -1, 3)
 44    try:
 45        out, h = net.rnn(seq)
 46    except RuntimeError:
 47        old = torch.backends.cudnn.enabled
 48        torch.backends.cudnn.enabled = False
 49        try: out, h = net.rnn(seq)
 50        finally: torch.backends.cudnn.enabled = old
 51    return net.head(h[-1]), out
 52
 53
 54def train_idea(seed, cfg, return_signature=False):
 55    seed_all(seed)
 56    ds = get_dataset('dynamics', seed, n_train=400, n_test=400)
 57    net = make_model('rnn_small', ds['input_shape'], ds['out_dim'])
 58    dev = device()
 59    x, y = ds['xtr'].to(dev), ds['ytr'].to(dev)
 60    xt, yt = ds['xte'].to(dev), ds['yte'].to(dev)
 61    opt = torch.optim.Adam(net.parameters(), lr=cfg['lr'], weight_decay=cfg['weight_decay'])
 62    mse = nn.MSELoss()
 63    batch = 128
 64    net.to(dev)
 65    rates, gaps = [], []
 66    try:
 67        for _ in range(cfg['epochs']):
 68            perm = torch.randperm(len(x), device=dev)
 69            net.train()
 70            for ix in perm.split(batch):
 71                pred, hseq = forward_hidden(net, x[ix])
 72                # Hidden coordinates are a learned phase plane. The intervention
 73                # penalizes insufficient composite-Laplacian contraction.
 74                phi = torch.atan2(hseq[..., 1], hseq[..., 0])
 75                rq, _ = spectral_rate(phi)
 76                spec_loss = torch.relu(0.20 - rq).pow(2).mean()
 77                loss = mse(pred, y[ix]) + cfg['lam'] * spec_loss
 78                opt.zero_grad(); loss.backward(); torch.nn.utils.clip_grad_norm_(net.parameters(), 5.0); opt.step()
 79        net.eval()
 80        with torch.no_grad():
 81            pred, hs = forward_hidden(net, xt)
 82            metric = float(mse(pred, yt).cpu())
 83            ph = torch.atan2(hs[..., 1], hs[..., 0])
 84            rr, _ = spectral_rate(ph)
 85            # A directly observed NN-scale locking statistic: adjacent phase
 86            # disagreement in the trained recurrent trajectories.
 87            gap = torch.mean(torch.abs(torch.atan2(torch.sin(ph[:, 1:] - ph[:, :-1]), torch.cos(ph[:, 1:] - ph[:, :-1]))))
 88            rates.append(float(rr.mean().cpu())); gaps.append(float(gap.cpu()))
 89    except RuntimeError:
 90        # Robust shared-GPU fallback: rerun this small experiment on CPU.
 91        if dev != 'cpu':
 92            return train_idea_cpu(seed, cfg, return_signature)
 93        raise
 94    result = (metric, {'rate': float(np.mean(rates)), 'phase_gap': float(np.mean(gaps))})
 95    return result if return_signature else metric
 96
 97
 98def train_idea_cpu(seed, cfg, return_signature=False):
 99    old = torch.cuda.is_available
100    # The actual function is device-selected by CUDA availability; make a CPU
101    # equivalent explicit for environments where CUDA allocation fails.
102    seed_all(seed); ds = get_dataset('dynamics', seed, 400, 400)
103    net = make_model('rnn_small', ds['input_shape'], ds['out_dim']).cpu()
104    x,y,xt,yt = ds['xtr'],ds['ytr'],ds['xte'],ds['yte']; opt=torch.optim.Adam(net.parameters(),lr=cfg['lr'],weight_decay=cfg['weight_decay']); mse=nn.MSELoss()
105    for _ in range(cfg['epochs']):
106        for ix in torch.randperm(len(x)).split(128):
107            pred,hs=forward_hidden(net,x[ix]); ph=torch.atan2(hs[...,1],hs[...,0]); rr,_=spectral_rate(ph); loss=mse(pred,y[ix])+cfg['lam']*torch.relu(.20-rr).pow(2).mean(); opt.zero_grad(); loss.backward(); opt.step()
108    with torch.no_grad():
109        pred,hs=forward_hidden(net,xt); ph=torch.atan2(hs[...,1],hs[...,0]); rr,_=spectral_rate(ph); gap=torch.mean(torch.abs(torch.atan2(torch.sin(ph[:,1:]-ph[:,:-1]),torch.cos(ph[:,1:]-ph[:,:-1]))))
110    out=(float(mse(pred,yt)),{'rate':float(rr.mean()),'phase_gap':float(gap)})
111    return out if return_signature else out[0]
112
113
114def make_base(cfg):
115    def run(seed):
116        seed_all(seed); ds=get_dataset('dynamics',seed,400,400); net,metric,_=train_model(make_model('rnn_small',ds['input_shape'],ds['out_dim']),ds,epochs=cfg['epochs'],lr=cfg['lr'],weight_decay=cfg['weight_decay'],batch=128,log=lambda *_:None); return metric
117    return run
118
119
120def main():
121    base=sweep_baseline(make_base, BASE_GRID, seeds=(0,1,2,3))
122    # Evaluate all idea settings on all paired seeds; select by full-seed mean.
123    idea_runs=[]
124    for cfg in IDEA_GRID:
125        vals=[train_idea(s,cfg) for s in SEEDS]
126        idea_runs.append((float(np.mean(vals)),cfg,vals))
127    _,best_cfg,best_vals=min(idea_runs,key=lambda z:z[0])
128    idea={'mean':float(np.mean(best_vals)),'std':float(np.std(best_vals)),'per_seed':[float(v) for v in best_vals],'n':8}
129    sig=[]
130    for s in SEEDS:
131        v= train_idea(s,best_cfg,True); sig.append(v[1])
132    # Baseline trained-model signature uses the same hidden observables.
133    bsig=[]
134    for s in SEEDS:
135        seed_all(s); ds=get_dataset('dynamics',s,400,400); net,_,_=train_model(make_model('rnn_small',ds['input_shape'],ds['out_dim']),ds,epochs=base['best_cfg']['epochs'],lr=base['best_cfg']['lr'],weight_decay=base['best_cfg']['weight_decay'],batch=128,log=lambda *_:None)
136        net.eval(); dev=next(net.parameters()).device
137        with torch.no_grad():
138            _,hs=forward_hidden(net,ds['xte'].to(dev)); ph=torch.atan2(hs[...,1],hs[...,0]); rr,_=spectral_rate(ph); gap=torch.mean(torch.abs(torch.atan2(torch.sin(ph[:,1:]-ph[:,:-1]),torch.cos(ph[:,1:]-ph[:,:-1])))); bsig.append({'rate':float(rr.mean()),'phase_gap':float(gap)})
139    br={'rate':float(np.mean([z['rate'] for z in bsig])),'phase_gap':float(np.mean([z['phase_gap'] for z in bsig]))}
140    ir={'rate':float(np.mean([z['rate'] for z in sig])),'phase_gap':float(np.mean([z['phase_gap'] for z in sig]))}
141    signature={'prediction':'spectral penalty should raise hidden phase rate and reduce adjacent phase gap','baseline_observed':br,'idea_observed':ir,'predicted_rate_change':ir['rate']-br['rate'],'observed_gap_change':ir['phase_gap']-br['phase_gap'],'confirmed':bool(ir['rate']>br['rate'] and ir['phase_gap']<br['phase_gap'])}
142    rep=make_report('dynamics','rnn_small',base,idea,{'mechanism_signature':signature,'idea_sweep':[{'cfg':c,'mean':m} for m,c,_ in idea_runs],'structural_match':'dynamics track: controlled pendulum rollout and recurrent stability'})
143    Path('bench_report.json').write_text(json.dumps(rep,indent=2)); print(json.dumps(rep,indent=2))
144if __name__=='__main__': main()