Skip to content
GitHub

sparx: spiking neural networks in JAX

Spiking neural networks in JAX

Train them with gradients or local rules. Simulate circuits in millivolts and milliseconds. Built on Flax and dew, with every model checked against a reference.

How it worksDocs
pip install git+https://github.com/AshishKumar4/sparx

This drone is flown by 128 spiking neurons, live. Move your pointer and it follows. Click to push it, drag it to throw it, click a neuron to silence it.Tap anywhere to send it there. Tap a neuron to silence it.

firing 0 Hztrained with sparx · reaches its target on 100% of 1,000 starts · how it learned

Neurons

A neuron keeps a charge, and spikes when it is full

A spiking neuron adds its input to a membrane potential that leaks away over time. When the membrane reaches a threshold, the neuron sends a spike, a single 1, and resets. The rest of the time it sends nothing.

Every sparx neuron is a small JAX dataclass with a step, and sparx.run scans one over time. The figure is that step, in your browser.

The code
import jax.numpy as jnp
import sparx
from sparx.dynamics import LIFCell, decay
cell = LIFCell(decay=decay(tau=12.0), threshold=1.0, reset="subtract")
drive = jnp.full((300,), 0.12) # the input of each step
spikes, state = sparx.run(cell, drive) # spikes.value: 1 where it fired
print(int(spikes.value.sum()), "spikes in 300 steps")
Neurons, from LIF to Izhikevich
LIFCell

0 spikes per 100 steps

Training

Layers that train like any Flax module

A spike is a step function, so its derivative is zero almost everywhere and gradients cannot pass it. sparx keeps the spike exact on the way forward and uses a smooth surrogate's slope on the way back. Its layers are Flax modules over time-major arrays, [T, B, ...], so jit, grad, vmap, optax and sharding work as they do for any network.

The code
import flax.linen as nn
import jax
import jax.numpy as jnp
import optax
import sparx
class Net(nn.Module):
@nn.compact
def __call__(self, spikes): # [T, B, 784]
x = sparx.nn.LIF(tau=2.0)(nn.Dense(256)(spikes))
return sparx.nn.LI(tau=2.0)(nn.Dense(10)(x)) # membrane [T, B, 10]
net = Net()
images = jax.random.uniform(jax.random.key(0), (32, 784))
labels = jnp.zeros(32, jnp.int32)
spikes = sparx.encode.RateEncoder(steps=8)(jax.random.key(1), images)
params = net.init(jax.random.key(2), spikes)
def loss(params):
logits = jnp.mean(net.apply(params, spikes), axis=0)
return optax.softmax_cross_entropy_with_integer_labels(logits, labels).mean()
grads = jax.grad(loss)(params) # through the spikes, by their surrogate
Teach a neuron when to fire
sparx.surrogate

ATan(alpha=2.0)

The drone

How the pilot learned to fly

The network flying the drone at the top of this page is three sparx layers: 7 readings in, 64 and 64 LIF neurons, and two leaky integrators whose membranes set the rotors' thrust. Nothing showed it how to fly. Each training step flew 256 drones for 2 s through a model of their physics in JAX, and took the gradient of their distance to their targets, through the spikes and through the physics.

After 4,000 steps it brings the drone within 15 cm of a fixed target, and keeps it there, on 100% of 1,000 random starts, and 100% of those that began upside down, in a median of 0.93 s.

The code
import flax.linen as nn
import jax
import jax.numpy as jnp
from sparx.nn import LI, LIF
pilot = nn.Sequential([
nn.Dense(64), LIF(tau=3.0, reset="zero"), # 7 readings in: the way to the target,
nn.Dense(64), LIF(tau=3.0, reset="zero"), # velocity, attitude and spin
nn.Dense(2), LI(tau=5.0), # 2 membranes out: the rotors' thrust
])
def step(params, carried, readings):
"""10 ms of the network, its membranes carried in the `state` collection."""
out, mutated = pilot.apply({"params": params, "state": carried}, readings[None], mutable=["state"])
return out[0], mutated["state"]
params = pilot.init(jax.random.key(0), jnp.zeros((1, 1, 7)))["params"]
membranes, carried = step(params, {}, jnp.zeros((256, 7))) # 256 drones, every neuron at rest
# Training scans step() and the drone's physics over 2 s of flight and takes jax.grad of the
# distance to the target: through the spikes by their surrogate, and through the physics.
Training, lesions and the check against sparx
100%of 1,000 flights reach the target and stay
100%of 1,000 thrown at 4 m/s and spun at 15 rad/s
0.93 smedian time to arrive
1.6 cmmedian distance after 6 s

Circuits

Circuits in millivolts and milliseconds

The same neuron protocol runs biological models in physical units, wired into populations and projections with delays, and simulated on one clock in NEST's order. Brunel's balanced network shifts between its regimes as inhibition and external drive change. The figure steps the network brunel(250) builds, in a worker in your browser.

Where the dynamics are deterministic, sparx matches NEST and Brian2 spike for spike. Where they are chaotic, as here, it matches them in rate, irregularity and synchrony over seeds.

The code
import jax
from sparx.graph import PopulationRate, SpikeRaster, simulate
from sparx.graph.models import brunel
network = brunel(250, g=5.0, eta=2.0) # 1,250 LIF neurons
result = simulate(network, network.init(jax.random.key(0)), duration=400.0, key=jax.random.key(1),
monitors={"spikes": SpikeRaster("e"), "rate": PopulationRate("e")})
spikes = result.records["spikes"] # [4000, 1000]: a row of booleans per 0.1 ms step
print(float(result.records["rate"][1000:].mean()), "Hz")
Brunel's regimes, and the whole fly brain
brunel(250)

1,250 LIF neurons, 125 inputs each, 0.1 ms steps · excitatory rate …

Fidelity

Every model is checked against a reference

Each model is checked against what defines it: a float64 loop of its equations, the authors' code, or NEST and Brian2. The fidelity ledger lists every check, its observed error and every known difference, and the status page what is not established.

ModelReferenceResult
LIF, current LIF, PSNSpikingJelly and snnTorchSame numbersSpikes exact; gradients within 3.6e-7
Biophysical LIFNEST iaf_psc_exp, _alpha, _deltaSame numbersSpike for spike; voltages within 1e-11 mV over 300 ms
Izhikevich (2004), twenty patternsHis figure1.m in GNU OctaveSame numbersEvery spike on his step; voltages to 1e-9 relative
Pair, triplet and dopamine STDPNEST stdp_*_synapseSame numbersEvery transmitted weight within 1e-10 relative
Cortical microcircuit, a fifthNEST on the network sparx drawsSame numbers12,689 spikes alike over 300 ms, float64
DelayedDense (learned delays)DCLS, as SNN-delays uses itWithin a toleranceOutputs within 2.4e-7; gradients within 1.4e-6
Conductance LIFNEST iaf_cond_* (RK45)Within a toleranceWithin 2e-3 mV; error falls 4x when dt halves
AdEx, Naud et al.’s eight patternsNEST aeif_* (RK45)Within a toleranceSame spike counts; every spike within 0.8 ms
Brunel (2000), four regimesNEST's brunel_delta_nest.pySame statisticsRate, CV and Fano factor within NEST’s spread over 8 seeds
FlyWire whole brain (Shiu et al.)Their published Brian2 runsSame statisticsRates correlate at 0.9989; MN9 at 67.1 Hz vs 67.0 ± 6.6

Results

Measured, with their conditions

These are short, untuned runs on a 4-core CPU, each beside its reference's number. No GPU or TPU numbers exist yet, and the full-length, multi-seed SHD comparison is still open.

Commands, times and comparisons
MNIST, rate-coded, 8 steps784-512-512 LIF97.5% after 2 epochs
SHD, Hammouamri et al.'s recipe, 20 of 150 epochs140-256-256 LIF with learned delays91.9% (their code on the same machine: 93.6%)
SHD, 140 channels140-128 ALIF, with and without learned delays74.6% and 64.5%
Fashion-MNIST, Seely and Gould's headline cellReLU residual MLP, depth 32PC-ALM 75.1%, PC 62.2%, backpropagation 77.8%
Pattern completion, Miconi et al.'s taskplastic recurrent network0.3% of bits wrong; 50.1% without fast weights

Learn by changing things

Install it

Python 3.12 or later, with JAX, Flax and dew. For a GPU or TPU, install the matching JAX first.

pip install git+https://github.com/AshishKumar4/sparx