Skip to content
GitHub

Chapter 7 · Learning

Local learning rules

Backpropagation needs every synapse to know about errors far away and long ago. Brains seem to manage with rules that use only what a synapse can see. This chapter builds three of them, from pure timing to a rule that matches backpropagation's gradients online.

What a synapse can know

Backpropagation through time gives each weight its exact share of the blame for the loss, but it gets there by storing every step and then sending errors backwards through the whole network and all of time. A synapse in a brain has nothing like that. It sees the spikes of the neuron before it, the state of the neuron after it, and the chemicals that wash over both. A learning rule built only from those is called local.

Local rules matter for engineering too. They can run as the network runs, in memory that does not grow with the sequence, which is what on-chip learning on neuromorphic hardware needs.

Spike-timing-dependent plasticity

Bi and Poo stimulated pairs of connected neurons grown in a dish and changed the timing between their spikes. When the receiving neuron fired within about 20 ms after the sending one, the synapse between them grew stronger. When it fired within about 20 ms before, the synapse grew weaker. The closer the two spikes, the larger the change.

That is spike-timing-dependent plasticity (STDP), and it reads as a causal rule: an input that fires just before its neuron fires probably helped, so strengthen it; one that fires just after did not, so weaken it.

To implement it, each neuron keeps a trace of its own recent spikes that jumps at each spike and decays. When the receiving neuron fires, every synapse grows by the sending neuron’s trace, which is large if it fired recently. When the sending neuron fires, its synapse shrinks by the receiving neuron’s trace. With weights scaled so the largest is 1, sparx’s PairSTDP, which is NEST’s stdp_synapse, does

w←min⁡ ⁣(w+λ (1−w)μ Kpre, 1)   at a postsynaptic spike,w←max⁡ ⁣(w−αλ wμ Kpost, 0)   at a presynaptic spikew \leftarrow \min\!\big(w + \lambda\,(1 - w)^{\mu}\,K_\text{pre},\ 1\big) \;\text{ at a postsynaptic spike}, \qquad w \leftarrow \max\!\big(w - \alpha\lambda\,w^{\mu}\,K_\text{post},\ 0\big) \;\text{ at a presynaptic spike}

where KpreK_\text{pre} and KpostK_\text{post} are the two traces, with time constants τ+\tau_+ and τ−\tau_-, λ\lambda sets the step size, and α\alpha how much depression weighs against potentiation.

PairSTDP

Each point runs the rule on one pair of spikes at that lag, from half the largest weight. With additive bounds, μ=0\mu = 0, every change has the same size wherever the weight is, which drives weights to 0 or 1. With multiplicative bounds, μ=1\mu = 1, a strong synapse grows less and a weak one shrinks less, which keeps them in between.

STDP finds a pattern

Masquelier, Guyonneau and Thorpe showed how much this rule can do on its own. In their simulation, 2,000 afferents fired at random. Half of them, now and then, replayed one frozen 50 ms spike pattern, at the same rate as their random firing, so that no single afferent’s rate gave the pattern away. A single neuron listened to all of them through STDP, with nothing telling it the pattern existed. After about 13.5 seconds it was firing selectively to the pattern, and it went on firing earlier and earlier within it.

Here is a smaller version, with 400 afferents and sparx’s LeakyIntegrateAndFire and PairSTDP in place of their neuron and rule. The shaded windows are the pattern’s.

LeakyIntegrateAndFire + PairSTDP

At first the neuron fires anywhere. Inputs that happened to fire before its spikes grow. The pattern repeats the same spikes in the same order every time, so its afferents get that boost repeatedly while random afferents get it only by chance. Within about a minute of simulated time the neuron stops firing outside the pattern and fires in most of its windows. On the right, the weights of the pattern’s afferents (violet) split into strong and weak, and the others fall. At Masquelier et al.’s depression of α=0.85\alpha = 0.85, this smaller network fell silent before finding the pattern, so this one uses 0.55.

import jax
from sparx.dynamics import (
Delta,
LeakyIntegrateAndFire,
PairSTDP,
Receptor,
)
from sparx.graph import (
FixedProbability,
Network,
PoissonInput,
Population,
Projection,
simulate,
)
# NEST's stdp_synapse with additive bounds
stdp = PairSTDP(tau_plus=16.8, tau_minus=33.7, lambda_=0.005,
alpha=0.55, mu_plus=0.0, mu_minus=0.0, w_max=1.0)
delta = {"ampa": Receptor(Delta())}
listener = LeakyIntegrateAndFire(tau_m=10.0, e_l=0.0, v_th=20.0,
v_reset=0.0, t_ref=1.0)
network = Network(
populations=(Population("in", 400, LeakyIntegrateAndFire(), delta),
Population("out", 1, listener, delta)),
projections=(Projection("in", "out", FixedProbability(1.0),
weight=0.475, delay=1.0, receptor="ampa",
plasticity=stdp),),
inputs=(PoissonInput("in", rate=40.0, weight=30.0,
receptor="ampa"),),
dt=1.0,
)
result = simulate(network, network.init(jax.random.key(0)),
duration=2000.0, key=jax.random.key(1))
connections = network.connections(result.variables)
weight = connections["in->out:ampa"].weight
print(weight.min(), weight.max()) # as STDP left them

Three factors

STDP alone learns correlations, whatever they are. To learn a task, the rule needs to know whether things went well. In the brain, neuromodulators such as dopamine broadcast that kind of signal widely, and often seconds after the spikes that earned it.

A three-factor rule bridges the gap. STDP no longer changes the weight; it writes an eligibility trace cc at each synapse that decays slowly. The weight changes only when the third factor, a modulator concentration nn, departs from its baseline bb:

dcdt=−cτc+STDP(t),dwdt=c (n−b)\frac{\mathrm{d}c}{\mathrm{d}t} = -\frac{c}{\tau_c} + \text{STDP}(t), \qquad \frac{\mathrm{d}w}{\mathrm{d}t} = c\,(n - b)

A reward that arrives a second after a lucky pair of spikes can still reinforce them, if τc\tau_c is long enough to remember. Izhikevich showed this solves the distal reward problem in spiking networks. sparx’s DopamineSTDP is NEST’s stdp_dopamine_synapse and reads its modulator from a population’s spikes.

e-prop: gradients without going back

Three-factor rules learn from a scalar reward. Can a local rule learn as precisely as backpropagation? Bellec et al. showed that for a recurrent spiking network the exact gradient factors as

dEdWji=∑tLjt ejit\frac{\mathrm{d}E}{\mathrm{d}W_{ji}} = \sum_t L_j^t\, e_{ji}^t

where ejite_{ji}^t is an eligibility trace, the derivative of neuron jj‘s spike with respect to the weight from ii through jj‘s own state, which the synapse can compute forward in time; and LjtL_j^t is a learning signal, how much the loss depends on neuron jj‘s spike at step tt. The trace is local. The learning signal is not: exactly, it includes how neuron jj‘s spike affects every other neuron later on. e-prop approximates it by its direct effect on the output alone, the error at the readout sent back through the readout’s weights.

For a LIF neuron with a detached reset, the trace is simple: the surrogate slope of neuron jj times the sending neuron’s activity filtered by jj‘s membrane decay. sparx’s eprop keeps it per synapse, filtered once more by the readout’s decay, and adds up the products as the sequence runs. Its memory is two numbers per synapse, whatever the sequence’s length.

Below, two copies of one recurrent network of 50 LIF neurons start from the same weights and learn to draw the grey curve from the same 20 input spike trains, one by e-prop and one by full BPTT.

sparx.learn.eprop · BPTT

The two learn at the same speed: after 100 updates both have cut the loss from 72 to under 0.6. Their gradients point almost the same way, with a cosine similarity of 0.98 at the starting weights and still at the end. What e-prop drops is the gradient through the recurrent spikes, which in this network carries little. In fact, sparx’s eprop is exactly BPTT with the gradient through the recurrent connections cut, which the code checks:

import jax
import jax.numpy as jnp
from sparx.dynamics import LIFCell, decay
from sparx.learn import EPropParams, bptt_loss, eprop
from sparx.surrogate import Triangle
keys = jax.random.split(jax.random.key(0), 4)
T, inputs, size = 200, 20, 50
u = (jax.random.uniform(keys[0], (T, 1, inputs)) < 0.05) * 1.0
t = jnp.arange(T)
target = jnp.sin(2 * jnp.pi * t / 100)[:, None, None]
params = EPropParams(jax.random.normal(keys[1], (inputs, size)),
0.15 * jax.random.normal(keys[2], (size, size)),
0.05 * jax.random.normal(keys[3], (size, 1)),
jnp.zeros(1))
cell = LIFCell(decay(tau=20.0), surrogate=Triangle(scale=0.3),
detach_reset=True)
def loss(y, target):
return 0.5 * jnp.sum((y - target) ** 2)
_, online = eprop(cell, params, u, target, loss, tau=10.0)
def through_time(p, cut):
return bptt_loss(cell, p, u, target, loss, tau=10.0,
cut_recurrence=cut)
full = jax.grad(through_time)(params, False)
cut = jax.grad(through_time)(params, True)
def cosine(a, b):
a = jnp.concatenate([x.ravel() for x in a])
b = jnp.concatenate([x.ravel() for x in b])
return float(a @ b / jnp.linalg.norm(a) / jnp.linalg.norm(b))
print("e-prop vs BPTT:", cosine(online, full))
print("e-prop vs BPTT with the recurrence cut:", cosine(online, cut))

Try this

  1. In the STDP figure, set α\alpha to 1 with additive bounds. Is a random pair of spikes, as likely to come in either order, more likely to strengthen or weaken a synapse? Why?
  2. In the pattern figure, freeze learning once the neuron has found the pattern. Does it still respond to it?
  3. A three-factor rule’s eligibility trace decays with τc=1\tau_c = 1 s. A reward arrives 2 s after the spikes that earned it. What fraction of their eligibility is left?
  4. Why might e-prop struggle on a task where the answer depends on something 2 seconds back, in a network whose membranes decay in 20 ms?
Answers
  1. Weaken. Potentiation’s window has time constant τ+=16.8\tau_+ = 16.8 ms and depression’s τ−=33.7\tau_- = 33.7 ms, so with equal heights the depression window holds twice the area. Random timing then depresses on average, which keeps random inputs from growing.
  2. Yes: freezing stops the weights changing, not the neuron firing. It now detects the pattern with fixed weights, as a trained network would.
  3. e−2≈0.14e^{-2} \approx 0.14, about a seventh.
  4. Its eligibility traces forget at the membrane’s rate, about 20 ms, and it ignores how a spike affects other neurons later, the one path that could carry the information longer. Bellec et al. gave some neurons adaptive thresholds with much longer time constants, chapter 3’s adaptive neuron, which gives the traces a slow component to remember with.

Summary

Local rules learn from what each synapse can see. STDP strengthens inputs that fire just before their neuron and finds repeating patterns without any teacher; three-factor rules add a delayed reward through an eligibility trace; e-prop computes a gradient close to backpropagation’s online, in memory independent of the sequence’s length, by dropping only what flows through the recurrent spikes. The next chapter adds one more thing a synapse can learn: not how strong to be, but how late.

References

  • G. Bi and M. Poo, “Synaptic modifications in cultured hippocampal neurons: dependence on spike timing, synaptic strength, and postsynaptic cell type”, Journal of Neuroscience 18(24), 1998, doi:10.1523/JNEUROSCI.18-24-10464.1998.
  • S. Song, K. D. Miller and L. F. Abbott, “Competitive Hebbian learning through spike-timing-dependent synaptic plasticity”, Nature Neuroscience 3, 2000, doi:10.1038/78829.
  • T. Masquelier, R. Guyonneau and S. J. Thorpe, “Spike timing dependent plasticity finds the start of repeating patterns in continuous spike trains”, PLoS ONE 3(1), e1377, 2008, doi:10.1371/journal.pone.0001377.
  • E. M. Izhikevich, “Solving the distal reward problem through linkage of STDP and dopamine signaling”, Cerebral Cortex 17(10), 2007, doi:10.1093/cercor/bhl152.
  • N. Frémaux and W. Gerstner, “Neuromodulated spike-timing-dependent plasticity, and theory of three-factor learning rules”, Frontiers in Neural Circuits 9, 85, 2016, doi:10.3389/fncir.2015.00085.
  • 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. Equations 1 to 3.