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.
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 jaximport jax.numpy as jnpimport 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 oncelayer, 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 roundedIn 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).