import sys, json, math, random import numpy as np import torch from torch import nn sys.path.insert(0, '/home/maxwelhelp/all/math2nn') from bench import get_dataset, train_model, evaluate, sweep_baseline, make_report H = 64 EPOCHS = 8 BATCH = 128 LRS = [1e-3, 3e-3, 1e-2] RHO = 0.65 NOISE = 0.18 DT = 0.1 class DeterministicLeaky(nn.Module): def __init__(self, input_dim=3, hidden=H, out_dim=1): super().__init__() self.inp = nn.Linear(input_dim, hidden) self.rec = nn.Linear(hidden, hidden, bias=False) self.head = nn.Linear(hidden, out_dim) def forward(self, x): x = x.view(x.shape[0], -1, 3) h = torch.zeros(x.shape[0], H, device=x.device) for t in range(x.shape[1]): h = h + DT * (-h + self.inp(x[:, t]) + self.rec(h)) return self.head(h) class CorrelatedIF(nn.Module): def __init__(self, input_dim=3, hidden=H, out_dim=1, rho=RHO): super().__init__() self.inp = nn.Linear(input_dim, hidden) self.rec = nn.Linear(hidden, hidden, bias=False) self.head = nn.Linear(hidden, out_dim) self.feedback = nn.Parameter(torch.tensor(0.15)) self.rho = float(rho) self.noise = NOISE self._last_x = None self._last_spikes = None self._last_noise = None def forward(self, x): x = x.view(x.shape[0], -1, 3) b = x.shape[0] dev = x.device h = torch.full((b, H), -0.65, device=dev) refractory = torch.zeros((b, H), device=dev, dtype=torch.long) prev = torch.zeros_like(h) filt = torch.zeros(b, device=dev) all_spikes, all_noise = [], [] alpha = math.exp(-DT) for t in range(x.shape[1]): filt = alpha * filt + prev.mean(1) active = (refractory == 0) & (h < 0) drift = -h + self.inp(x[:, t]) + self.rec(prev) + self.feedback * filt[:, None] z0 = torch.randn(b, 1, device=dev) zi = torch.randn(b, H, device=dev) noise = self.noise * (self.rho * z0 + math.sqrt(1-self.rho*self.rho) * zi) * math.sqrt(DT) hn = torch.where(active, h + DT * drift + noise, h) hard = (hn >= 0).float() surrogate = torch.sigmoid(12 * hn) spike = hard + surrogate - surrogate.detach() h = torch.where(hard.bool(), torch.full_like(hn, -0.75), hn) refractory = torch.where(hard.bool(), torch.ones_like(refractory), torch.clamp(refractory-1, min=0)) prev = spike all_spikes.append(hard) all_noise.append(noise) self._last_x = h.detach() self._last_spikes = torch.stack(all_spikes, 1).detach() self._last_noise = torch.stack(all_noise, 1).detach() return self.head(h) 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 make_base(cfg, seed): seed_all(seed + 10000) return DeterministicLeaky() def make_idea(cfg, seed): seed_all(seed + 20000) return CorrelatedIF() def run_one(kind, cfg, seed, n_train=400, n_test=200): seed_all(seed) ds = get_dataset('dynamics', seed, n_train=n_train, n_test=n_test) model = make_base(cfg, seed) if kind == 'base' else make_idea(cfg, seed) net, metric, hist = train_model(model, ds, epochs=EPOCHS, lr=cfg['lr'], batch=BATCH) if net is None: raise RuntimeError('training failed') sig = {} if kind == 'idea': dev = next(net.parameters()).device with torch.no_grad(): _ = net(ds['xte'].to(dev)) q = net._last_noise[:, :, :2].reshape(-1, 2).cpu().numpy() observed = float(np.corrcoef(q[:, 0], q[:, 1])[0, 1]) predicted = RHO * RHO spikes = float(net._last_spikes.mean().cpu()) sig = {'observed_noise_corr': observed, 'predicted_noise_corr': predicted, 'spike_rate': spikes} return float(metric), sig def eval_factory(kind, cfg, seeds): vals=[] for s in seeds: vals.append(run_one(kind, cfg, int(s))[0]) return {'mean': float(np.mean(vals)), 'per_seed': vals} def main(): # Baseline sweep uses the same union of learning rates tried for the idea. grid = [{'lr': lr} for lr in LRS] base_block = sweep_baseline(lambda cfg: (lambda seed: run_one('base', cfg, int(seed))[0]), grid, seeds=(0,1,2,3)) best_lr = base_block['best_cfg']['lr'] idea_grid = [{'lr': lr} for lr in LRS] idea_sweep = [{'cfg': c, 'mean': eval_factory('idea', c, (0,1,2,3))['mean']} for c in idea_grid] idea_lr = min(idea_sweep, key=lambda z:z['mean'])['cfg']['lr'] base_full = eval_factory('base', {'lr': best_lr}, tuple(range(8))) idea_full = eval_factory('idea', {'lr': idea_lr}, tuple(range(8))) # Signature is measured from trained benchmark models, averaged over final paired runs. obs=[]; rates=[] for s in range(8): _, sg = run_one('idea', {'lr': idea_lr}, s) obs.append(sg['observed_noise_corr']); rates.append(sg['spike_rate']) signature = {'predicted_shared_covariance_factor': RHO*RHO, 'observed_shared_noise_correlation_mean': float(np.mean(obs)), 'observed_shared_noise_correlation_per_seed': obs, 'spike_rate_mean': float(np.mean(rates)), 'confirmed': bool(abs(float(np.mean(obs))-RHO*RHO) < 0.08)} report = make_report('dynamics', 'rnn_small', {'best_cfg': base_block['best_cfg'], 'sweep': base_block['sweep'], 'full': base_full}, idea_full, {'mechanism_signature': signature, 'idea_sweep': idea_sweep, 'architecture': 'matched 64-unit recurrent state and scalar head; only deterministic leaky update vs correlated IF differs'}) with open('bench_report.json','w') as f: json.dump(report, f, indent=2) print(json.dumps(report, indent=2)) if __name__ == '__main__': main()