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

Mechanism confirmed, baseline not beaten

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