Correlated stochastic integrate-and-fire recurrent layer / bench_exp.py

Mechanism confirmed, baseline not beaten

Raw ⬇ ZIP
  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()