Skip to content
GitHub

Chapter 6 · Learning

Backpropagation through time

A spiking network carries its state from one step to the next, so a gradient has to travel back through every step it ran. This chapter follows that journey through one neuron, shows when the gradient fades and when it explodes, and what each costs.

Unrolling

A spiking neuron’s membrane at step tt depends on its membrane at step t−1t-1. A network of them is a recurrent computation, even when no neuron connects back to another: each neuron is recurrent with itself through its own state.

To differentiate a loss computed after TT steps, write the computation out as TT copies of the same step, one after another, all sharing the same weights, and backpropagate through the whole chain. This is backpropagation through time (BPTT). The gradient for a weight is the sum of its gradients at every step.

Two costs follow. The backward pass needs the state of every step, so memory grows in proportion to TT times the number of neurons. And the gradient must pass through every step between where the loss is measured and where the input arrived, multiplied by one factor per step. Those factors decide whether learning over long spans works.

The chain through one neuron

Take chapter 1’s neuron and let its own spike feed back to it with a weight ww, as if its axon looped back onto its own dendrite (neurons with such a synapse exist and are called autapses). With sparx’s order of operations, each step is

ut=β vt−1+xt+w st−1,st=H(ut−θ),vt=ut−θ stu_t = \beta\,v_{t-1} + x_t + w\,s_{t-1}, \qquad s_t = H(u_t - \theta), \qquad v_t = u_t - \theta\,s_t

where utu_t is the membrane before the reset and vtv_t after it. In the backward pass the spike’s derivative is the surrogate’s, σt′=σ′(ut−θ)\sigma'_t = \sigma'(u_t - \theta). The derivative of one step’s membrane with respect to the last is then

∂ut+1∂ut=β (1−θ σt′)⏟leak and reset+w σt′⏟recurrence\frac{\partial u_{t+1}}{\partial u_t} = \underbrace{\beta\,(1 - \theta\,\sigma'_t)}_{\text{leak and reset}} + \underbrace{w\,\sigma'_t}_{\text{recurrence}}

and the gradient that reaches step tt from the end is the product of these factors over every step in between. A product of many numbers either shrinks toward 0 or grows without bound unless the numbers sit very close to 1. There are three terms to watch.

The leak. Without the other two terms, each step multiplies by β=e−1/τ\beta = e^{-1/\tau}, and the gradient from kk steps back is βk=e−k/τ\beta^k = e^{-k/\tau}. The gradient fades over the same time constant as the membrane itself: a neuron cannot learn from an input much older than τ\tau steps through its membrane alone.

The reset. The term 1−θσt′1 - \theta\sigma'_t says that, near threshold, raising the membrane also raises the reset that follows it. With a surrogate whose slope near threshold is 1/θ1/\theta, the two cancel exactly and no gradient passes through the membrane at all. detach_reset=True stops the gradient through the reset, leaving β\beta; SpyTorch’s tutorials and SpikingJelly do the same.

Recurrence. The term wσt′w\sigma'_t is gradient that flows through the neuron’s own spike back into itself. When it pushes ∣∂ut+1/∂ut∣\lvert \partial u_{t+1}/\partial u_t \rvert above 1 for many steps in a row, the gradient explodes. In a network, ww is a matrix of recurrent weights, and the danger is its largest eigenvalue times the surrogate’s slope.

LIFCell + jax.grad

The top shows the neuron’s membrane before each reset, and its spikes. The bottom shows, on a log scale, how strongly the last membrane depends on each earlier step’s input, ∂vT/∂xt\partial v_T / \partial x_t. It starts at the right, where it is about 1, and fades leftward, at a rate of e−1/10e^{-1/10} per step with detach_reset. Raise the self-connection to 0.5 and the gradient at step 0 is about 101410^{14} times larger than at the end with ATan, and larger still with StraightThrough. With FastSigmoid it stays below 1. Turn the gradient through the reset back on, and with ATan the gradient all but vanishes: its slope at threshold is exactly 1.

In a real network

The same arithmetic plays out in matrices. sparx measured it training a recurrent network on the Spiking Heidelberg Digits: once training had grown the recurrent weight matrix’s spectral radius from 1 to 5, backpropagation with ATan drove the gradient norm past 10810^8 within 300 steps. FastSigmoid(100), which passes gradient only within a few hundredths of the threshold, kept it below 10.

The usual defences, roughly in order of how often they are needed:

  • A narrower surrogate passes gradient through fewer neurons at each step, so fewer factors near or above 1.
  • Clipping caps the gradient’s norm before each update. The drone on the front page trained with its gradient clipped to a norm of 1.
  • detach_reset removes the reset’s term, trading a little information for a cleaner product.
  • Shorter sequences mean fewer factors. sparx’s layers stream: with the state collection mutable, a long sequence can be fed in chunks that carry the state forward, and a gradient can be cut between chunks.

Memory has its own remedy. jax.checkpoint stores only some steps’ states and recomputes the rest during the backward pass, trading compute for memory. The racer of chapter 14 trains with every step checkpointed.

import jax
import jax.numpy as jnp
from sparx.dynamics import LIFCell, MembraneState, SynapticInput, decay
from sparx.surrogate import ATan
x = 0.12 + 0.2 * jax.random.uniform(jax.random.key(0), (120,))
def last_membrane(x, w, cell):
"""Run a neuron whose own spike feeds back with weight w."""
def step(carry, xt):
state, spike = carry
drive = SynapticInput(jump=xt + w * spike)
state, out = cell.step(state, drive, 1.0)
return (state, out.value), None
start = (MembraneState(jnp.zeros(())), jnp.zeros(()))
(state, _), _ = jax.lax.scan(step, start, x)
return state.v
cell = LIFCell(decay=decay(tau=10.0), surrogate=ATan(),
detach_reset=True)
for w in (0.0, 0.5):
grad = jax.grad(last_membrane)(x, w, cell) # d v_T / d x_t
print(f"w={w}: {float(abs(grad[0])):.1e} at step 0, "
f"{float(abs(grad[-1])):.1e} at the last step")

Without the self-connection, the gradient at step 0 is e−119/10≈7×10−6e^{-119/10} \approx 7 \times 10^{-6} of its value at the end, the leak alone. With w=0.5w = 0.5 it is many orders of magnitude larger instead. jax.grad computes exactly the product this chapter wrote down.

Avoiding BPTT

BPTT is exact for the surrogate model, but its memory grows with TT and its backward pass must wait for the forward one to finish. Two families of methods avoid it.

Exact spike-time gradients. A LIF neuron in continuous time spikes at a definite moment that moves smoothly as the weights change, so the spike times have true gradients. EventProp (Wunderlich and Pehle) computes them with a backward pass that only visits spike events. sparx’s EventLIF with spike_times gives these gradients for networks of LIF neurons with current synapses.

Online rules. e-prop (Bellec et al.) factors each weight’s gradient into an eligibility trace, computed forward in time at the synapse, times a learning signal from the loss. Its memory does not grow with TT, and it runs as the network runs. Chapter 7 builds it.

Try this

  1. With detach_reset and no self-connection, the gradient at step 0 is about 7×10−67 \times 10^{-6}. What would it be with a membrane time constant of 30 steps?
  2. Find the smallest self-connection at which the gradient at step 0 exceeds 1 with ATan. Why does the neuron’s firing matter?
  3. Set the self-connection to 1.5. The neuron fires on nearly every step, and the gradient no longer explodes. Why?
  4. With no self-connection, StraightThrough and the gradient through the reset, the gradient is exactly zero at every step. Why?
Answers
  1. e−119/30≈0.019e^{-119/30} \approx 0.019. A longer time constant lets the gradient, like the membrane, remember for longer.
  2. Between 0.15 and 0.2 with these inputs. The recurrence term wσt′w\sigma'_t only counts at steps where the membrane is near threshold, where σt′\sigma'_t is large, so the answer depends on how often the neuron hovers there. A neuron that never comes near threshold passes no gradient through its spikes, whatever ww is.
  3. Firing on every step, the neuron’s membrane sits well above threshold after each input: w=1.5w = 1.5 adds more than the threshold takes away. Far from threshold the surrogate’s slope is small, so wσt′w\sigma'_t is small too.
  4. Its slope is 1 everywhere, so ∂vt/∂ut=1−θσt′=0\partial v_t/\partial u_t = 1 - \theta\sigma'_t = 0 at every step, the last one included: whatever raises the membrane raises the reset by exactly as much. Every path from an input to vTv_T passes through uTu_T, so every gradient is zero.

Summary

Backpropagation through time unrolls a network over its steps and multiplies one factor per step. Through a neuron’s membrane the factor is the leak, so gradients fade over the membrane’s time constant; the reset can cancel them near threshold; recurrence through spikes can make them explode, more so with wide surrogates. Narrow surrogates, clipping, detached resets and checkpointing keep it in hand, and spike-time gradients and online rules avoid it. The next chapter turns to rules that need no backward pass at all.

References

  • P. J. Werbos, “Backpropagation through time: what it does and how to do it”, Proceedings of the IEEE 78(10), 1990, doi:10.1109/5.58337.
  • T. C. Wunderlich and C. Pehle, “Event-based backpropagation can compute exact gradients for spiking neural networks”, Scientific Reports 11, 12829, 2021, doi:10.1038/s41598-021-91786-z.
  • G. Bellec et al., “A solution to the learning dilemma for recurrent networks of spiking neurons”, Nature Communications 11, 3625, 2020, doi:10.1038/s41467-020-17236-y.
  • sparx’s measurement of gradients through recurrence: the guide, on the recurrent SHD run.