Correlated stochastic integrate-and-fire recurrent layer / bench_exp.py
Mechanism confirmed, baseline not beaten
1import sys, json, math, random
2import numpy as np
3import torch
4from torch import nn
5sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
6from bench import get_dataset, train_model, evaluate, sweep_baseline, make_report
7
8H = 64
9EPOCHS = 8
10BATCH = 128
11LRS = [1e-3, 3e-3, 1e-2]
12RHO = 0.65
13NOISE = 0.18
14DT = 0.1
15
16class DeterministicLeaky(nn.Module):
17 def __init__(self, input_dim=3, hidden=H, out_dim=1):
18 super().__init__()
19 self.inp = nn.Linear(input_dim, hidden)
20 self.rec = nn.Linear(hidden, hidden, bias=False)
21 self.head = nn.Linear(hidden, out_dim)
22 def forward(self, x):
23 x = x.view(x.shape[0], -1, 3)
24 h = torch.zeros(x.shape[0], H, device=x.device)
25 for t in range(x.shape[1]):
26 h = h + DT * (-h + self.inp(x[:, t]) + self.rec(h))
27 return self.head(h)
28
29class CorrelatedIF(nn.Module):
30 def __init__(self, input_dim=3, hidden=H, out_dim=1, rho=RHO):
31 super().__init__()
32 self.inp = nn.Linear(input_dim, hidden)
33 self.rec = nn.Linear(hidden, hidden, bias=False)
34 self.head = nn.Linear(hidden, out_dim)
35 self.feedback = nn.Parameter(torch.tensor(0.15))
36 self.rho = float(rho)
37 self.noise = NOISE
38 self._last_x = None
39 self._last_spikes = None
40 self._last_noise = None
41 def forward(self, x):
42 x = x.view(x.shape[0], -1, 3)
43 b = x.shape[0]
44 dev = x.device
45 h = torch.full((b, H), -0.65, device=dev)
46 refractory = torch.zeros((b, H), device=dev, dtype=torch.long)
47 prev = torch.zeros_like(h)
48 filt = torch.zeros(b, device=dev)
49 all_spikes, all_noise = [], []
50 alpha = math.exp(-DT)
51 for t in range(x.shape[1]):
52 filt = alpha * filt + prev.mean(1)
53 active = (refractory == 0) & (h < 0)
54 drift = -h + self.inp(x[:, t]) + self.rec(prev) + self.feedback * filt[:, None]
55 z0 = torch.randn(b, 1, device=dev)
56 zi = torch.randn(b, H, device=dev)
57 noise = self.noise * (self.rho * z0 + math.sqrt(1-self.rho*self.rho) * zi) * math.sqrt(DT)
58 hn = torch.where(active, h + DT * drift + noise, h)
59 hard = (hn >= 0).float()
60 surrogate = torch.sigmoid(12 * hn)
61 spike = hard + surrogate - surrogate.detach()
62 h = torch.where(hard.bool(), torch.full_like(hn, -0.75), hn)
63 refractory = torch.where(hard.bool(), torch.ones_like(refractory), torch.clamp(refractory-1, min=0))
64 prev = spike
65 all_spikes.append(hard)
66 all_noise.append(noise)
67 self._last_x = h.detach()
68 self._last_spikes = torch.stack(all_spikes, 1).detach()
69 self._last_noise = torch.stack(all_noise, 1).detach()
70 return self.head(h)
71
72def seed_all(seed):
73 random.seed(seed); np.random.seed(seed); torch.manual_seed(seed)
74 if torch.cuda.is_available(): torch.cuda.manual_seed_all(seed)
75
76def make_base(cfg, seed):
77 seed_all(seed + 10000)
78 return DeterministicLeaky()
79
80def make_idea(cfg, seed):
81 seed_all(seed + 20000)
82 return CorrelatedIF()
83
84def run_one(kind, cfg, seed, n_train=400, n_test=200):
85 seed_all(seed)
86 ds = get_dataset('dynamics', seed, n_train=n_train, n_test=n_test)
87 model = make_base(cfg, seed) if kind == 'base' else make_idea(cfg, seed)
88 net, metric, hist = train_model(model, ds, epochs=EPOCHS, lr=cfg['lr'], batch=BATCH)
89 if net is None:
90 raise RuntimeError('training failed')
91 sig = {}
92 if kind == 'idea':
93 dev = next(net.parameters()).device
94 with torch.no_grad():
95 _ = net(ds['xte'].to(dev))
96 q = net._last_noise[:, :, :2].reshape(-1, 2).cpu().numpy()
97 observed = float(np.corrcoef(q[:, 0], q[:, 1])[0, 1])
98 predicted = RHO * RHO
99 spikes = float(net._last_spikes.mean().cpu())
100 sig = {'observed_noise_corr': observed, 'predicted_noise_corr': predicted, 'spike_rate': spikes}
101 return float(metric), sig
102
103def eval_factory(kind, cfg, seeds):
104 vals=[]
105 for s in seeds:
106 vals.append(run_one(kind, cfg, int(s))[0])
107 return {'mean': float(np.mean(vals)), 'per_seed': vals}
108
109def main():
110 # Baseline sweep uses the same union of learning rates tried for the idea.
111 grid = [{'lr': lr} for lr in LRS]
112 base_block = sweep_baseline(lambda cfg: (lambda seed: run_one('base', cfg, int(seed))[0]), grid, seeds=(0,1,2,3))
113 best_lr = base_block['best_cfg']['lr']
114 idea_grid = [{'lr': lr} for lr in LRS]
115 idea_sweep = [{'cfg': c, 'mean': eval_factory('idea', c, (0,1,2,3))['mean']} for c in idea_grid]
116 idea_lr = min(idea_sweep, key=lambda z:z['mean'])['cfg']['lr']
117 base_full = eval_factory('base', {'lr': best_lr}, tuple(range(8)))
118 idea_full = eval_factory('idea', {'lr': idea_lr}, tuple(range(8)))
119 # Signature is measured from trained benchmark models, averaged over final paired runs.
120 obs=[]; rates=[]
121 for s in range(8):
122 _, sg = run_one('idea', {'lr': idea_lr}, s)
123 obs.append(sg['observed_noise_corr']); rates.append(sg['spike_rate'])
124 signature = {'predicted_shared_covariance_factor': RHO*RHO,
125 'observed_shared_noise_correlation_mean': float(np.mean(obs)),
126 'observed_shared_noise_correlation_per_seed': obs,
127 'spike_rate_mean': float(np.mean(rates)),
128 'confirmed': bool(abs(float(np.mean(obs))-RHO*RHO) < 0.08)}
129 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'})
130 with open('bench_report.json','w') as f: json.dump(report, f, indent=2)
131 print(json.dumps(report, indent=2))
132
133if __name__ == '__main__': main()