Skip to content
GitHub

Chapter 5 · Learning

Surrogate gradients

Gradient descent needs a derivative, and a spike's is zero almost everywhere. This chapter shows why that stops learning cold, how a surrogate derivative gets around it while every spike stays exact, and how to choose one.

The problem

A neuron spikes when its membrane reaches the threshold. Write x=v−θx = v - \theta for how far above threshold the membrane is; the spike is the Heaviside step of xx:

s=H(x)={1x≥00x<0s = H(x) = \begin{cases} 1 & x \geq 0 \\ 0 & x < 0 \end{cases}

Training by gradient descent asks how the loss changes when a weight changes a little. A weight changes the input, the input changes vv, and vv changes ss. The chain rule multiplies the derivatives along that path, and one of them is ds/dx\mathrm{d}s/\mathrm{d}x. It is zero for every xx except 0, where it is undefined. So every gradient that passes through a spike is zero.

This is not a technicality. Nudge a weight by a small amount and every membrane moves by a small amount, but unless one of them happens to cross the threshold, no spike changes, and neither does the loss. Seen as a function of the weights, the loss is flat almost everywhere and jumps where some spike appears or disappears. Gradient descent on that surface has nowhere to go.

The surrogate

The fix is to keep the forward pass exactly as it is, binary spikes and all, and change only the derivative used in the backward pass. In place of H′(x)H'(x), use the derivative σ′(x)\sigma'(x) of a smooth step σ\sigma, a bump centred on the threshold:

∂s∂x  ⟶  σ′(x)\frac{\partial s}{\partial x} \;\longrightarrow\; \sigma'(x)

This is a surrogate gradient. It is a deliberate approximation: the gradient computed this way is not the gradient of anything the network computes. But it points somewhere sensible. A neuron whose membrane is just below threshold could be made to fire by a small push, and the surrogate gives it a large gradient. A neuron far below threshold would need a large push, and the surrogate gives it little. The surrogate answers the question gradient descent needs answered, which neurons are close to changing, while the spikes themselves stay exact.

sparx.surrogate

ATan(alpha=2.0)

The orange line is the spike as the forward pass computes it; the violet curve is the slope the backward pass uses in its place. Each surrogate in sparx is a different bump:

Surrogateσ′(x)\sigma'(x)Note
ATan(alpha=2)α/21+(π2αx)2\dfrac{\alpha/2}{1 + (\tfrac{\pi}{2}\alpha x)^2}Integrates to 1; its tails fall as 1/x21/x^2, so distant neurons still get some gradient
Sigmoid(alpha=4)α sig(αx) (1−sig(αx))\alpha\,\mathrm{sig}(\alpha x)\,(1 - \mathrm{sig}(\alpha x))The logistic’s slope
FastSigmoid(slope=25)1(slope ∣x∣+1)2\dfrac{1}{(\text{slope}\,\lvert x\rvert + 1)^2}SuperSpike’s; peak 1 at any slope
Triangle(width=1)max⁡(0, 1−∣x∣/width)\max(0,\ 1 - \lvert x\rvert/\text{width})Zero beyond width
Rectangle(width=1)1/width1/\text{width} for ∣x∣<width/2\lvert x\rvert < \text{width}/2A box
Gaussian(sigma=0.5)the normal density
StraightThrough()11Passes the gradient as if the spike were the identity

How sparx implements it

sparx.spike is a jax.custom_jvp: its forward pass is the step, and its derivative rule is the surrogate’s. JAX derives reverse mode from that rule by transposing it, so the one definition serves jax.grad, jax.jvp and vmap alike. Every neuron model takes its surrogate as a field, LIFCell(..., surrogate=FastSigmoid(100.0)), and so does every layer, sparx.nn.LIF(surrogate=...).

import jax
import jax.numpy as jnp
from sparx.surrogate import ATan, FastSigmoid, spike
v = jnp.array([-1.0, -0.1, 0.0, 0.1, 1.0]) # membrane - threshold
print(spike(v, ATan())) # [0. 0. 1. 1. 1.]
# The forward pass is the step; the gradient is the surrogate's.
for surrogate in (ATan(), FastSigmoid(25.0)):
slope = jax.vmap(jax.grad(lambda x, s=surrogate: spike(x, s)))(v)
print(type(surrogate).__name__, slope)

The forward pass gives exactly 0 or 1. The gradients are the two bumps sampled at five points: ATan passes 0.09 of the gradient a whole threshold away, FastSigmoid(25) only 0.0015.

Teach a neuron when to fire

Now use it. One neuron listens to 40 inputs, each a random spike train, through 40 weights. The task is to make it fire at three chosen steps and nowhere else. Mark the steps on the strip under its membrane, by clicking to add or remove a mark, and press Train.

The loss compares the neuron’s spikes with the marks after smoothing both with an exponential filter of 10 steps, then takes the mean squared difference. Without the smoothing, a spike one step early would count as fully wrong, with no hint of which way to move it. This filtered distance is van Rossum’s spike distance.

LIFCell + jax.grad

The rows on top are the inputs, darker for a stronger excitatory weight and blue for an inhibitory one. Each update is a full pass of backpropagation through all 200 steps with the chosen surrogate, followed by an Adam step on the weights. Inputs that fire just before a mark grow stronger; inputs that fire before an unwanted spike weaken. With ATan, within 100 updates the neuron fires at the three marks, give or take a few steps. Try the other surrogates from the same starting weights: most get there, at different speeds.

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
inputs = jax.random.uniform(jax.random.key(0), (T, N)) < 0.04
trains = inputs.astype(jnp.float32)
# Fire at these steps
target = jnp.zeros(T).at[jnp.array([40, 90, 150])].set(1.0)
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):
def step(f, s):
f = keep * f + s
return f, f
return jax.lax.scan(step, 0.0, spikes)[1]
def loss(w):
out, _ = sparx.run(cell, trains @ w) # spikes 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)
fired = sparx.run(cell, trains @ w)[0].value
print("fires at", jnp.flatnonzero(fired))

Shape matters less than scale

With so many surrogates to choose from, does the choice matter? Zenke and Vogels tested this systematically and found that learning is robust to the surrogate’s shape, but sensitive to its scale: how wide the bump is and how tall. Too narrow and few neurons receive any gradient; too wide or too tall and gradients through many steps compound.

The compounding is worst in recurrent networks, where the gradient passes through the same neurons at every step. sparx measured it on the Spiking Heidelberg Digits: training a recurrent network with ATan grew the gradient norm past 10810^8 within 300 steps, while the much narrower FastSigmoid(100) kept it below 10. Chapter 6 looks at why gradients through time explode.

Try this

  1. In the teaching figure, reset the weights and train with Triangle. It fires at two of the marks and never at the third. Why?
  2. With ATan(alpha=2), how much gradient passes at x=0.5x = 0.5, half a threshold above?
  3. Why might FastSigmoid(100) keep a recurrent network’s gradients small, when ATan does not?
  4. Why not make training easy by replacing the spike with a sigmoid in the forward pass too?
Answers
  1. Triangle passes no gradient at all where the membrane is more than one threshold away from it. Around the third mark the membrane never comes that close, so no weight receives any signal to make the neuron fire there, and training stops improving. ATan’s long tails do reach it, and so, crudely, does StraightThrough, whose slope is 1 everywhere. On this small problem even that is enough; chapter 6 shows where it is not.
  2. 11+(π/2×2×0.5)2=11+(π/2)2≈0.29\dfrac{1}{1 + (\pi/2 \times 2 \times 0.5)^2} = \dfrac{1}{1 + (\pi/2)^2} \approx 0.29.
  3. At 0.1 from threshold FastSigmoid(100) passes 1/1211/121 of the gradient, and ATan about 0.9. Only neurons within a few hundredths of threshold pass gradient, so far fewer paths through time carry it, and the product of factors along each path stays small.
  4. The network would then compute with real numbers, not spikes. It would lose the sparsity and the additions-only arithmetic of chapter 0, and a network trained that way behaves differently when its spikes are made binary again.

Summary

A spike’s derivative is zero almost everywhere, so ordinary gradients cannot pass it. A surrogate keeps every spike exact in the forward pass and uses a smooth bump’s slope in the backward pass, giving credit to neurons in proportion to how close they are to firing. The shape of the bump matters little, its scale a great deal. The next chapter follows the gradient back through time, step by step, and asks what that costs.

References

  • E. O. Neftci, H. Mostafa and F. Zenke, “Surrogate gradient learning in spiking neural networks”, IEEE Signal Processing Magazine 36(6), 2019, doi:10.1109/MSP.2019.2931595.
  • F. Zenke and T. P. Vogels, “The remarkable robustness of surrogate gradient learning for instilling complex function in spiking neural networks”, Neural Computation 33(4), 2021, doi:10.1162/neco_a_01367.
  • F. Zenke and S. Ganguli, “SuperSpike: supervised learning in multilayer spiking neural networks”, Neural Computation 30(6), 2018, doi:10.1162/neco_a_01086. The fast sigmoid.
  • M. C. W. van Rossum, “A novel spike distance”, Neural Computation 13(4), 2001, doi:10.1162/089976601300014321.
  • sparx’s measurement of gradients through recurrence: design notes.