Chapter 8 · Learning
Delays
A spike takes time to travel from one neuron to the next, and the time differs from synapse to synapse. This chapter shows how a neuron can use those differences to recognize a sequence, and how a network can learn them by gradient descent.
Time spent travelling
A spike leaves the cell body and runs down the axon at a speed set by the axon’s diameter and insulation: thin, bare axons conduct slowly, and thick ones wrapped in myelin conduct fast. Crossing the synapse at the end takes a little more time. So every connection has its own delay, which depends on how far the spike goes and how fast.
The models so far ignored this, or gave every connection the same delay. But differences in delay are information. Chapter 1’s neuron was a coincidence detector: it fired for two inputs that arrived within a few milliseconds of each other and ignored the same two spread out. Put a delay on each input and the neuron fires for the inputs that arrive together, which means inputs that were sent in a particular order, at particular intervals. A coincidence detector behind delay lines is a sequence detector.
A circuit that measures time
In 1948 Jeffress proposed that the brain locates sounds this way. A sound from the left reaches the left ear a fraction of a millisecond before the right. Suppose signals from each ear travel along delay lines in opposite directions, past a row of coincidence detectors. Each detector sits where the two paths’ delays differ by a particular amount, so the one that fires says how much sooner the sound reached one ear, and so where it came from. Carr and Konishi found this circuit in the barn owl’s brainstem: in the nucleus laminaris, the incoming axons act as delay lines, and the laminaris neurons as coincidence detectors.
Making delays learnable
To train delays by gradient descent you need the loss to change smoothly as a delay changes. A delay of a whole number of steps has no gradient: move it by 0.1 step and nothing happens until it rounds to the next step, when everything jumps.
Hammouamri, Khalfaoui-Hassani and Masquelier made it smooth by spreading each synapse’s spike over several steps. A synapse from input to output has a weight and a delay , and delivers through a kernel over the lags : a Gaussian of width centred on the delay and normalized to sum to 1,
Moving slides the bump along the lags, which changes smoothly, so the loss has a gradient in every delay. During training shrinks. At the kernel is a single 1 at the rounded delay and each synapse delivers to exactly one step: that is the network you deploy. sparx’s DelayedDense is this layer.
Line them up
Here is the smallest version of the problem. Inputs A, B and C fire once each, at steps 5, 18 and 30. They reach one leaky integrator, with a decay time of 4 steps, through weights of 0.6, and the integrator is read at step 50. The reading is highest when all three spikes arrive at step 50 exactly. The violet bars are each synapse’s kernel: where its spike lands, and how spread out it is.
Drag the delays yourself, or press “Learn the delays”. Gradient ascent on the reading then moves all three at once while shrinks from 8 to 0.5 over 260 updates, and at the end the delays are rounded. Scramble them and learn again: from most starts they find the same answer.
Why start wide? A spike that arrives far from step 50 contributes almost nothing to the reading, and with a narrow kernel moving it a little changes nothing either. Its gradient is close to zero and it never moves. A wide kernel smears every spike across many steps, so some of it always reaches the reading and the gradient always points the right way. As the delays settle, narrowing the kernel sharpens them.
import jaximport jax.numpy as jnpimport optax
from sparx.nn import LI, DelayedDense
# A, B and C fire once each, at steps 5, 18 and 30x = jnp.zeros((80, 1, 3))x = x.at[jnp.array([5, 18, 30]), 0, jnp.arange(3)].set(1.0)layer = DelayedDense(1, max_delay=45, use_bias=False)readout = LI(tau=4.0)kernel = jnp.full((3, 1), 0.6)delay = jnp.array([[4.0], [16.0], [22.0]])
def loss(delay, sigma): # minus the readout at step 50 params = {"kernel": kernel, "delay": delay} y = layer.apply({"params": params}, x, sigma) return -readout.apply({}, y)[50, 0, 0]
adam = optax.adam(0.6)state = adam.init(delay)grad = jax.jit(jax.grad(loss))# The Gaussians narrow as the delays train.for sigma in jnp.geomspace(8.0, 0.5, 260): updates, state = adam.update(grad(delay, sigma), state) delay = jnp.clip(optax.apply_updates(delay, updates), 0, 45)
# Deployed, each delay is rounded to a whole step.print(delay[:, 0].round(), -loss(delay.round(), 0))It prints the delays 45, 32 and 20, which put all three arrivals on step 50, and a reading of 1.8, three times 0.6.
In a real network
Spoken words are sequences in time, so delays suit them. The Spiking Heidelberg Digits are 10,420 recordings of spoken digits, 0 to 9 in English and German, converted into the spikes of 700 simulated channels of the inner ear. Hammouamri et al. trained a network with two hidden layers of LIF neurons and learned delays, about 0.2 million parameters, and reported 95.07% on its test set, the mean of ten seeds. SHD has no validation split, so they used the test set for validation too.
sparx.models.SpikingMLP(delays=...) is their network. On a 4-core CPU, 20 of their 150 epochs with their recipe reached 91.9%, and their own code on the same machine 93.6%. The full-length comparison over several seeds has not been run yet; the status page tracks it.
Try this
- Without the figure: what delays put all three arrivals exactly on step 50?
- One arrival lands one step early, at step 49. How much does it add to the reading at step 50?
- Two ears 20 cm apart hear a sound from directly to one side. If sound travels at 343 m/s, how much sooner does it reach the near ear? What range of delays must a Jeffress circuit cover?
- After training, a delay settles at 4.3 steps. What does the deployed network use, and why can’t it use 4.3?
Answers
- 45, 32 and 20: step 5 plus 45, 18 plus 32 and 30 plus 20 are all 50.
- Its 0.6 decays for one step at , so it adds about 0.47 instead of 0.6.
- ms. A sound straight ahead reaches both ears at once, so the circuit must cover differences from about −0.58 to +0.58 ms.
- 4 steps. A spike exists only at whole steps, so at the kernel is a one-hot at the rounded delay. Training with a shrinking is what lets the network get used to that rounding before it happens.
Summary
Every connection takes time, and a neuron that detects coincidences turns differences in delay into sensitivity to sequences, as the barn owl does to locate sounds. Spreading each synapse over a Gaussian of lags makes its delay differentiable, and shrinking the Gaussian during training leaves every synapse delivering to one step. The next chapter puts thousands of neurons together and looks at what they do as a network.
References
- L. A. Jeffress, “A place theory of sound localization”, Journal of Comparative and Physiological Psychology 41(1), 1948, doi:10.1037/h0061495.
- C. E. Carr and M. Konishi, “A circuit for detection of interaural time differences in the brain stem of the barn owl”, Journal of Neuroscience 10(10), 1990, doi:10.1523/JNEUROSCI.10-10-03227.1990.
- I. Hammouamri, I. Khalfaoui-Hassani and T. Masquelier, “Learning delays in spiking neural networks using dilated convolutions with learnable spacings”, ICLR 2024, arXiv:2306.17670.
- B. Cramer, Y. Stradmann, J. Schemmel and F. Zenke, “The Heidelberg spiking data sets for the systematic evaluation of spiking neural networks”, IEEE TNNLS 33(7), 2022, doi:10.1109/TNNLS.2020.3044364.