Skip to content
GitHub

Chapter 5 · Learning

Surrogate gradients

LearnChapter 5

Surrogate gradients

A spike is all or nothing, so its derivative is zero almost everywhere. Training spiking networks by gradient descent takes a stand-in for that derivative.

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.

sparx.surrogate

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.

LIFCell + jax.grad

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 jax
import jax.numpy as jnp
import optax
import sparx
from sparx.dynamics import LIFCell, decay
from sparx.surrogate import ATan
T, N = 200, 40
trains = (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 steps
cell = 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 ATan
for _ 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).

All of Learn