The problem
A neuron spikes when v - threshold crosses zero: the output is the step function of it. Nudge the weights a little and the membrane moves a little, but the spike stays exactly 0 or exactly 1 until the membrane crosses. The step's derivative is zero on both sides and undefined at the crossing, so backpropagation gets nothing back through a spike.
The fix keeps the forward pass exact and replaces only the derivative. In the backward pass, sparx's spike uses the slope of a smooth function shaped like a bump around the threshold. Neurons close to firing get a large gradient; neurons far below it get almost none.
ATan(alpha=2.0)
sparx.spike is a jax.custom_jvp, so the same rule serves jax.grad, jax.jvp and vmap. Every neuron takes its surrogate as an argument: sparx.nn.LIF(surrogate=FastSigmoid(100.0)).
Teach a neuron when to fire
Here one neuron listens to 40 inputs, each a random spike train, through 40 weights. Mark the steps where it should fire on the strip under its membrane (click to add or remove a mark), then train. Each update takes the gradient of how far its spikes are from the marks, both smoothed by an exponential filter, through the neuron and its spikes, and moves the weights by Adam.
The rows on top are the inputs, darker for a stronger excitatory weight and blue for an inhibitory one. Watch the weights of inputs that fire just before a mark grow, and spikes slide toward the marks. With StraightThrough, whose slope is 1 everywhere, every step counts equally, wherever the membrane is, and training wanders.
The code
import jaximport jax.numpy as jnpimport optax
import sparxfrom sparx.dynamics import LIFCell, decayfrom sparx.surrogate import ATan
T, N = 200, 40trains = (jax.random.uniform(jax.random.key(0), (T, N)) < 0.04).astype(jnp.float32)target = jnp.zeros(T).at[jnp.array([40, 90, 150])].set(1.0) # fire at these stepscell = LIFCell(decay=decay(tau=10.0), threshold=1.0, surrogate=ATan())keep = decay(tau=10.0) # an exponential filter of 10 steps
def smooth(spikes): return jax.lax.scan(lambda f, s: (keep * f + s,) * 2, 0.0, spikes)[1]
def loss(w): out, _ = sparx.run(cell, trains @ w) # spikes are exactly 0 or 1 return jnp.mean((smooth(out.value) - smooth(target)) ** 2)
w = 0.35 * jax.random.normal(jax.random.key(1), (N,))adam = optax.adam(0.04)state = adam.init(w)grad = jax.jit(jax.grad(loss)) # the slope comes from ATanfor _ in range(400): updates, state = adam.update(grad(w), state) w = optax.apply_updates(w, updates)print("fires at", jnp.flatnonzero(sparx.run(cell, trains @ w)[0].value))Choosing one
Each surrogate has a sharpness. A wide one passes gradient from neurons far from threshold, which helps a silent network start learning. A narrow one is closer to the true step. Through long recurrence the choice decides whether training works at all. In sparx's recurrent SHD run, ATan let the gradient norm pass 1e8 within 300 steps, and FastSigmoid(100) kept it below 10.
Checked
The figure computes its gradient by hand in the browser. site/test/learning.test.ts compares it with jax.gradthrough sparx.run on the same inputs, weights and targets, for ATan, FastSigmoid and Triangle. Every entry agrees to within 5e-16.
sparx's LIF spikes and gradients match SpikingJelly and snnTorch to within 3.6e-7, with soft, hard and detached resets (fidelity ledger).