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.
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 sparxfrom 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 stepspikes, state = sparx.run(cell, drive) # spikes.value: 1 if it firedprint(int(spikes.value.sum()), "spikes in 300 steps")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 nnimport jaximport jax.numpy as jnpimport 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)) # [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) xent = optax.softmax_cross_entropy_with_integer_labels return xent(logits, labels).mean()
# Through the spikes, by their surrogategrads = jax.grad(loss)(params)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 a cost led by 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 nnimport jaximport jax.numpy as jnp
from sparx.nn import LI, LIF
# 7 readings in: the way to the target, velocity, attitude, spin.# 2 membranes out: the rotors' thrust.pilot = nn.Sequential([ nn.Dense(64), LIF(tau=3.0, reset="zero"), nn.Dense(64), LIF(tau=3.0, reset="zero"), nn.Dense(2), LI(tau=5.0),])
def step(params, carried, readings): """10 ms of the network, its membranes carried in `state`.""" 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"]# 256 drones, every neuron at restmembranes, carried = step(params, {}, jnp.zeros((256, 7)))# Training scans step() and the drone's physics over 2 s of flight# and takes jax.grad of a cost led by the distance to the target:# through the spikes by their surrogate, and through the physics.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 the code builds, in a worker in your browser.
Current-based LIF, Izhikevich, STDP and the cortical microcircuit on one drawn network match NEST spike for spike. Conductance and adaptive models agree within stated margins. Chaotic networks, like this one, match in rate, irregularity and synchrony over seeds.
The code
import jax
from sparx.graph import PopulationRate, SpikeRaster, simulatefrom sparx.graph.models import brunel
# A tenth of Brunel's network, with synapses ten times hisnetwork = brunel(250, g=5.0, eta=2.0, j=1.0) # 1,250 LIF neuronsmonitors = {"spikes": SpikeRaster("e"), "rate": PopulationRate("e")}result = simulate(network, network.init(jax.random.key(0)), duration=400.0, key=jax.random.key(1), monitors=monitors)
# [4000, 1000]: one row of booleans per 0.1 ms stepspikes = result.records["spikes"]print(float(result.records["rate"][1000:].mean()), "Hz")· rate … · CV … · Fano …
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.
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 comparisonsLearn by changing things
- Why spikes?Where spiking networks beat conventional ones today, where they lose, and how to choose.
- What a neuron doesSynapses, a membrane that keeps charge, and the all-or-none spike.
- Membranes and time constantsThe RC circuit, the exact solution, and what a step of dt does to it.
- Spikes and thresholdsThreshold, reset, refractoriness, the f-I curve and adaptation.
- CodingRate, timing and change: three ways to put a number into spikes.
- Surrogate gradientsWhy a spike has no gradient, and the stand-in that lets gradient descent through.
- Backpropagation through timeUnrolling a network over time, what it costs, and why gradients explode.
- Local learning rulesSTDP, three-factor rules and e-prop: learning from what each synapse can see.
- DelaysSpikes take time to travel. Learning how long turns sequences into coincidences.
- Networks and dynamicsBalanced excitation and inhibition, irregular firing, and chaos.
- Physical units and biologyMillivolts and nanosiemens: conductances, receptors and real neurons.
- Simulators and fidelityHow a simulator steps time, and how to tell whether two of them agree.
- Hardware and eventsEvent cameras, neuromorphic chips, and NIR, the format that moves a network between them.
- Case study: the droneA spiking network that learned to fly by gradients through its physics.
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