Correlated stochastic integrate-and-fire recurrent layer / experiment.py
Mechanism confirmed, baseline not beaten
1import json, math, os, random, time
2import numpy as np
3import torch
4from torch import nn
5
6SEED = 648
7random.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED)
8torch.set_num_threads(min(12, os.cpu_count() or 1))
9
10def device():
11 return torch.device('cuda' if torch.cuda.is_available() else 'cpu')
12
13class CorrelatedIF(nn.Module):
14 def __init__(self, inp, hidden, rho=0.0, dt=0.1, tau_m=1.0, tau_k=1.0,
15 noise=0.18, refractory=1):
16 super().__init__()
17 self.hidden, self.dt, self.tau_m, self.tau_k = hidden, dt, tau_m, tau_k
18 self.noise, self.refractory = noise, refractory
19 self.rho = float(rho)
20 self.inp = nn.Linear(inp, hidden)
21 self.rec = nn.Linear(hidden, hidden, bias=False)
22 self.feedback = nn.Parameter(torch.tensor(0.15))
23 self.readout = nn.Linear(hidden, 2)
24 self.rho_logit = nn.Parameter(torch.tensor(math.log(rho/(1-rho))) if 0 < rho < 1 else torch.tensor(-12.0))
25 def forward(self, seq, return_stats=False):
26 B, T, _ = seq.shape; dev = seq.device
27 x = torch.full((B,self.hidden), -0.65, device=dev)
28 refr = torch.zeros((B,self.hidden), device=dev, dtype=torch.long)
29 prev = torch.zeros_like(x); f = torch.zeros(B, device=dev)
30 spikes=[]
31 alpha = math.exp(-self.dt/self.tau_k)
32 # Fixed rho is used for named ablations; learned rho is bounded but initialized at requested value.
33 rho = torch.sigmoid(self.rho_logit)
34 for t in range(T):
35 f = alpha*f + prev.mean(1)/self.tau_k
36 active = (refr == 0) & (x < 0)
37 drift = -x/self.tau_m + self.inp(seq[:,t]) + self.rec(prev) + self.feedback*f[:,None]
38 z0 = torch.randn(B,1,device=dev); zi = torch.randn(B,self.hidden,device=dev)
39 dz = self.noise * (rho*z0 + torch.sqrt(1-rho*rho)*zi) * math.sqrt(self.dt)
40 xn = torch.where(active, x + drift*self.dt + dz, x)
41 hard = (xn >= 0).float()
42 # straight-through fast sigmoid surrogate around threshold
43 soft = torch.sigmoid(12.0*xn)
44 s = hard + soft - soft.detach()
45 x = torch.where(hard.bool(), torch.full_like(xn, -0.75), xn)
46 refr = torch.where(hard.bool(), torch.full_like(refr, self.refractory), torch.clamp(refr-1,min=0))
47 prev = s
48 spikes.append(hard)
49 logits = self.readout(x)
50 if return_stats:
51 sp = torch.stack(spikes,1)
52 return logits, {'rate': sp.mean().item(), 'rate_var': sp.mean((2,)).var().item(), 'active': (sp>0).float().mean().item()}
53 return logits
54
55class LeakyRNN(nn.Module):
56 def __init__(self, inp, hidden, dt=0.1):
57 super().__init__(); self.hidden=hidden; self.dt=dt
58 self.inp=nn.Linear(inp,hidden); self.rec=nn.Linear(hidden,hidden,bias=False); self.readout=nn.Linear(hidden,2)
59 def forward(self, seq, return_stats=False):
60 B,T,_=seq.shape; x=torch.zeros(B,self.hidden,device=seq.device)
61 for t in range(T): x=x + self.dt*(-x + self.inp(seq[:,t]) + self.rec(x))
62 return self.readout(x)
63
64def make_data(n, T=30, noise=0.0):
65 g=torch.Generator().manual_seed(SEED+n+int(noise*1000))
66 x=torch.randn(n,T,1,generator=g); w=torch.linspace(-1,1,T)
67 score=(x[:,:,0]*w).sum(1); y=(score>0).long()
68 if noise: x=x+noise*torch.randn(x.shape,generator=g)
69 return x,y
70
71def train(kind, train_x, train_y, test_x, test_y, dev, epochs=12):
72 torch.manual_seed(SEED+{'det':1,'ind':2,'corr':3}[kind])
73 model = LeakyRNN(1,32).to(dev) if kind=='det' else CorrelatedIF(1,32,rho=(0.65 if kind=='corr' else 0.0)).to(dev)
74 opt=torch.optim.Adam(model.parameters(),lr=3e-3); bs=64
75 losses=[]; grad_samples=[]; t0=time.time()
76 for ep in range(epochs):
77 p=torch.randperm(len(train_x),device=dev)
78 for ix in p.split(bs):
79 opt.zero_grad(set_to_none=True); out=model(train_x[ix]); loss=nn.functional.cross_entropy(out,train_y[ix]); loss.backward()
80 grad_samples.append(float(torch.nn.utils.clip_grad_norm_(model.parameters(),5.0))); opt.step(); losses.append(loss.item())
81 with torch.no_grad():
82 out=model(test_x); acc=(out.argmax(1)==test_y).float().mean().item()
83 stats=model(test_x,True)[1] if kind!='det' else {}
84 noisy_x,_=make_data(len(test_x),noise=0.5); noisy_x=noisy_x.to(dev)
85 noisy_acc=(model(noisy_x).argmax(1)==test_y).float().mean().item()
86 return {'accuracy':acc,'noisy_accuracy':noisy_acc,'final_loss':float(np.mean(losses[-20:])),
87 'grad_norm_var':float(np.var(grad_samples[-100:])), 'seconds':time.time()-t0, **stats}
88
89def math_check():
90 # Conditional covariance of sigma*(rho*z0+sqrt(1-rho²)*zi)*sqrt(dt).
91 torch.manual_seed(SEED); n=200000; rho=.65; sigma=.18; dt=.1
92 z0=torch.randn(n,1); zi=torch.randn(n,2)
93 d=sigma*(rho*z0+math.sqrt(1-rho*rho)*zi)*math.sqrt(dt)
94 cov=np.cov(d.numpy(),rowvar=False)
95 expected=np.array([[sigma*sigma*dt, sigma*sigma*rho*rho*dt],[sigma*sigma*rho*rho*dt,sigma*sigma*dt]])
96 cov_err=float(np.max(np.abs(cov-expected)))
97 # Exponential kernel recurrence gives exact sampled impulse response alpha^k/tau.
98 tau=1.7; dt=.1; alpha=math.exp(-dt/tau); f=0.; vals=[]
99 for k in range(8): f=alpha*f+(1.0/tau if k==0 else 0.0); vals.append(f)
100 filt_err=float(max(abs(vals[k]-(alpha**k)/tau) for k in range(8)))
101 return {'empirical_cov':cov.tolist(),'expected_cov':expected.tolist(),'cov_max_abs_error':cov_err,'filter_max_abs_error':filt_err}
102
103def main():
104 dev=device()
105 try:
106 train_x,train_y=make_data(1024); test_x,test_y=make_data(512,noise=0.0)
107 train_x,train_y,test_x,test_y=[z.to(dev) for z in (train_x,train_y,test_x,test_y)]
108 result={'device':str(dev),'math_check':math_check()}
109 for k in ('det','ind','corr'):
110 try: result[k]=train(k,train_x,train_y,test_x,test_y,dev)
111 except Exception as e:
112 if dev.type=='cuda':
113 dev=torch.device('cpu'); train_x,train_y,test_x,test_y=[z.cpu() for z in (train_x,train_y,test_x,test_y)]
114 result['device']='cpu_fallback'; result[k]=train(k,train_x,train_y,test_x,test_y,dev)
115 else: raise
116 print(json.dumps(result,indent=2))
117 except Exception:
118 if dev.type=='cuda':
119 torch.cuda.empty_cache(); os.environ['CUDA_VISIBLE_DEVICES']=''
120 main()
121 else: raise
122if __name__=='__main__': main()