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 where 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)) # 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 surrogateATan(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 nnimport jaximport 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.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, simulatefrom sparx.graph.models import brunel
network = brunel(250, g=5.0, eta=2.0) # 1,250 LIF neuronsresult = 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 stepprint(float(result.records["rate"][1000:].mean()), "Hz")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.
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
- NeuronsDrive a LIF neuron and an Izhikevich cell with current, and watch the membrane, the threshold and the reset.
- Surrogate gradientsPick the times a neuron should fire, and watch gradient descent through its spikes teach it.
- Learned delaysEach synapse learns how late to deliver its spikes, so a neuron can hear a pattern as a coincidence.
- PlasticitySTDP strengthens the inputs that fire just before a neuron does, until one neuron finds a pattern hidden in noise.
- NetworksBrunel's balanced network in each of its regimes, simulated as sparx simulates it.
- sparx, NEST and Brian2The same neurons in three simulators, overlaid, with the differences measured and named.
- The pilotHow the drone on the front page learned to fly, what silencing its neurons does, and how the browser matches sparx.
- NIR exportThe pilot as a NIR graph, the format other simulators and neuromorphic hardware toolchains read.
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