Spectral Basin Allocation for Multimodal Neural Memories / bench_experiment.py
Mechanism confirmed, baseline not beaten
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()