Skip to content
GitHub

Chapter 8 · Learning

Learned delays

LearnChapter 8

Learned delays

A spike takes time to travel. When each synapse learns how long, a neuron can hear a sequence of spikes as one coincidence.

Time as a parameter

In a brain, a spike reaches each target after its own delay, set by the length and speed of the axon. A neuron that only fires when its inputs arrive together can then detect a sequence: the input that fires first takes the longest route. Hammouamri et al. (ICLR 2024) learned these delays by gradient descent beside the weights, and reported a new state of the art on the Spiking Heidelberg Digits with a small network.

A delay is a whole number of steps, which has no gradient. sparx.nn.DelayedDense spreads each synapse over a Gaussian centered at its delay, so moving the delay moves where the spike arrives and the loss can feel it. As training goes on the Gaussians narrow, and at sigma=0 each synapse delivers to exactly one step: that is the network to deploy.

Line them up

Inputs A, B and C fire once each, at steps 5, 18 and 30, into one leaky integrator through a weight of 0.6, and the integrator is read at step 50. The violet bars are each synapse's kernel: where its spike arrives, and how spread out. The reading is highest when all three arrive just at step 50. Drag the delays yourself, or let gradient ascent on the reading learn them while sigma shrinks from 8 to 0.5; at the end it rounds them.

DelayedDense + LI

With a wide Gaussian the arrivals overlap even when the delays are off, so the gradient can see the way to go. With a narrow one it sees only the steps next to each delay. That is why training starts wide.

The code
import jax
import jax.numpy as jnp
import optax
from sparx.nn import LI, DelayedDense
x = jnp.zeros((80, 1, 3)).at[jnp.array([5, 18, 30]), 0, jnp.arange(3)].set(1.0) # A, B, C fire once
layer, readout = DelayedDense(1, max_delay=45, use_bias=False), 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
y = layer.apply({"params": {"kernel": kernel, "delay": delay}}, x, sigma)
return -readout.apply({}, y)[50, 0, 0]
adam = optax.adam(0.6)
state = adam.init(delay)
grad = jax.jit(jax.grad(loss))
for sigma in jnp.geomspace(8.0, 0.5, 260): # the Gaussians narrow as they train
updates, state = adam.update(grad(delay, sigma), state)
delay = jnp.clip(optax.apply_updates(delay, updates), 0, 45)
print(delay[:, 0].round(), -loss(delay.round(), 0)) # deployed: each delay rounded

In a real network

sparx.models.SpikingMLP(delays=...) delays its synapses this way, and with Hammouamri et al.'s settings it is their SNN-delays network, checked against their code. On a 4-core CPU, 20 of their 150 epochs reached 91.9% on SHD, and their code on the same machine 93.6%. The full-length, multi-seed comparison has not been run yet (status).

Checked

site/test/learning.test.ts runs this toy through sparx's DelayedDense and LI at three widths. The browser's kernels match delay_kernel to 1e-14, and its gradient in each delay matches jax.grad to 1e-16.

DelayedDense matches DCLS, the layer SNN-delays uses: outputs within 2.4e-7, and gradients in the inputs, weights and delays within 1.4e-6 (fidelity ledger).

All of Learn