import json, math, os, random, time import numpy as np import torch from torch import nn SEED = 648 random.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED) torch.set_num_threads(min(12, os.cpu_count() or 1)) def device(): return torch.device('cuda' if torch.cuda.is_available() else 'cpu') class CorrelatedIF(nn.Module): def __init__(self, inp, hidden, rho=0.0, dt=0.1, tau_m=1.0, tau_k=1.0, noise=0.18, refractory=1): super().__init__() self.hidden, self.dt, self.tau_m, self.tau_k = hidden, dt, tau_m, tau_k self.noise, self.refractory = noise, refractory self.rho = float(rho) self.inp = nn.Linear(inp, hidden) self.rec = nn.Linear(hidden, hidden, bias=False) self.feedback = nn.Parameter(torch.tensor(0.15)) self.readout = nn.Linear(hidden, 2) self.rho_logit = nn.Parameter(torch.tensor(math.log(rho/(1-rho))) if 0 < rho < 1 else torch.tensor(-12.0)) def forward(self, seq, return_stats=False): B, T, _ = seq.shape; dev = seq.device x = torch.full((B,self.hidden), -0.65, device=dev) refr = torch.zeros((B,self.hidden), device=dev, dtype=torch.long) prev = torch.zeros_like(x); f = torch.zeros(B, device=dev) spikes=[] alpha = math.exp(-self.dt/self.tau_k) # Fixed rho is used for named ablations; learned rho is bounded but initialized at requested value. rho = torch.sigmoid(self.rho_logit) for t in range(T): f = alpha*f + prev.mean(1)/self.tau_k active = (refr == 0) & (x < 0) drift = -x/self.tau_m + self.inp(seq[:,t]) + self.rec(prev) + self.feedback*f[:,None] z0 = torch.randn(B,1,device=dev); zi = torch.randn(B,self.hidden,device=dev) dz = self.noise * (rho*z0 + torch.sqrt(1-rho*rho)*zi) * math.sqrt(self.dt) xn = torch.where(active, x + drift*self.dt + dz, x) hard = (xn >= 0).float() # straight-through fast sigmoid surrogate around threshold soft = torch.sigmoid(12.0*xn) s = hard + soft - soft.detach() x = torch.where(hard.bool(), torch.full_like(xn, -0.75), xn) refr = torch.where(hard.bool(), torch.full_like(refr, self.refractory), torch.clamp(refr-1,min=0)) prev = s spikes.append(hard) logits = self.readout(x) if return_stats: sp = torch.stack(spikes,1) return logits, {'rate': sp.mean().item(), 'rate_var': sp.mean((2,)).var().item(), 'active': (sp>0).float().mean().item()} return logits class LeakyRNN(nn.Module): def __init__(self, inp, hidden, dt=0.1): super().__init__(); self.hidden=hidden; self.dt=dt self.inp=nn.Linear(inp,hidden); self.rec=nn.Linear(hidden,hidden,bias=False); self.readout=nn.Linear(hidden,2) def forward(self, seq, return_stats=False): B,T,_=seq.shape; x=torch.zeros(B,self.hidden,device=seq.device) for t in range(T): x=x + self.dt*(-x + self.inp(seq[:,t]) + self.rec(x)) return self.readout(x) def make_data(n, T=30, noise=0.0): g=torch.Generator().manual_seed(SEED+n+int(noise*1000)) x=torch.randn(n,T,1,generator=g); w=torch.linspace(-1,1,T) score=(x[:,:,0]*w).sum(1); y=(score>0).long() if noise: x=x+noise*torch.randn(x.shape,generator=g) return x,y def train(kind, train_x, train_y, test_x, test_y, dev, epochs=12): torch.manual_seed(SEED+{'det':1,'ind':2,'corr':3}[kind]) model = LeakyRNN(1,32).to(dev) if kind=='det' else CorrelatedIF(1,32,rho=(0.65 if kind=='corr' else 0.0)).to(dev) opt=torch.optim.Adam(model.parameters(),lr=3e-3); bs=64 losses=[]; grad_samples=[]; t0=time.time() for ep in range(epochs): p=torch.randperm(len(train_x),device=dev) for ix in p.split(bs): opt.zero_grad(set_to_none=True); out=model(train_x[ix]); loss=nn.functional.cross_entropy(out,train_y[ix]); loss.backward() grad_samples.append(float(torch.nn.utils.clip_grad_norm_(model.parameters(),5.0))); opt.step(); losses.append(loss.item()) with torch.no_grad(): out=model(test_x); acc=(out.argmax(1)==test_y).float().mean().item() stats=model(test_x,True)[1] if kind!='det' else {} noisy_x,_=make_data(len(test_x),noise=0.5); noisy_x=noisy_x.to(dev) noisy_acc=(model(noisy_x).argmax(1)==test_y).float().mean().item() return {'accuracy':acc,'noisy_accuracy':noisy_acc,'final_loss':float(np.mean(losses[-20:])), 'grad_norm_var':float(np.var(grad_samples[-100:])), 'seconds':time.time()-t0, **stats} def math_check(): # Conditional covariance of sigma*(rho*z0+sqrt(1-rho²)*zi)*sqrt(dt). torch.manual_seed(SEED); n=200000; rho=.65; sigma=.18; dt=.1 z0=torch.randn(n,1); zi=torch.randn(n,2) d=sigma*(rho*z0+math.sqrt(1-rho*rho)*zi)*math.sqrt(dt) cov=np.cov(d.numpy(),rowvar=False) expected=np.array([[sigma*sigma*dt, sigma*sigma*rho*rho*dt],[sigma*sigma*rho*rho*dt,sigma*sigma*dt]]) cov_err=float(np.max(np.abs(cov-expected))) # Exponential kernel recurrence gives exact sampled impulse response alpha^k/tau. tau=1.7; dt=.1; alpha=math.exp(-dt/tau); f=0.; vals=[] for k in range(8): f=alpha*f+(1.0/tau if k==0 else 0.0); vals.append(f) filt_err=float(max(abs(vals[k]-(alpha**k)/tau) for k in range(8))) return {'empirical_cov':cov.tolist(),'expected_cov':expected.tolist(),'cov_max_abs_error':cov_err,'filter_max_abs_error':filt_err} def main(): dev=device() try: train_x,train_y=make_data(1024); test_x,test_y=make_data(512,noise=0.0) train_x,train_y,test_x,test_y=[z.to(dev) for z in (train_x,train_y,test_x,test_y)] result={'device':str(dev),'math_check':math_check()} for k in ('det','ind','corr'): try: result[k]=train(k,train_x,train_y,test_x,test_y,dev) except Exception as e: if dev.type=='cuda': 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)] result['device']='cpu_fallback'; result[k]=train(k,train_x,train_y,test_x,test_y,dev) else: raise print(json.dumps(result,indent=2)) except Exception: if dev.type=='cuda': torch.cuda.empty_cache(); os.environ['CUDA_VISIBLE_DEVICES']='' main() else: raise if __name__=='__main__': main()