sparx.learn
Learning rules beyond surrogate-gradient backpropagation through time (design.md section 7).
Modules
Section titled “Modules”| Module | |
|---|---|
sparx.learn.convert | Converting a trained ReLU network to a spiking one (Rueckauer et al., Frontiers in Neuroscience 2017). |
sparx.learn.reinforce | REINFORCE for a recurrent layer of stochastic spiking neurons, its gradient carried as eligibilities. |
Contents
Section titled “Contents”| Name | |
|---|---|
Diffusion | What a reward spread to: each connection’s weight change, each unit’s local reward and each connection’s share of its target’s. |
EPropParams | A recurrent layer, x_t = u_t @ w_in + z_{t-1} @ w_rec, and its leaky readout, y_t = decay(tau, dt) y_{t-1} + z_t @ w_out + b_out. |
EventLIF | The membrane and synaptic time constants tau_m and tau_syn (ms), and the threshold v_th (mV above rest). |
OTTTLayer | |
PredictiveCoding | Predictive coding’s inference and weight update; alpha > 0 makes it PC-ALM. |
RNeuralNet | RNeuralNet-Research’s network, clocked: neurons internal neurons, inputs input neurons and the reward feeder, in that order, as one recurrent layer. |
ReinforceParams | A recurrent layer, x_t = u_t @ w_in + s_{t-1} @ w_rec, [in, N] and [N, N]; or each example’s eligibility for them, with a leading batch axis. |
ResidualBlock | A layer of Seely and Gould’s residual MLP: scale * f(below) @ kernel, plus below when skip. |
Settled | |
SpikingMaxPool | Max pooling of spikes [T, ..., H, W, C] over unpadded window patches, gated by spike counts. |
accumulate | Sum each step’s loss and its gradient over a sequence, holding the carry fixed in each step’s gradient. |
bptt_loss | The summed loss of eprop_forward’s readout; its gradient is BPTT’s, or e-prop’s with cut_recurrence. |
eligibility_traces | Every step’s eligibility traces dz_t/dW through each neuron’s own state, [T, B, N, in + N], for analysis; eprop uses them as they are made instead of storing them. |
eprop | e-prop’s gradients for eprop_forward’s network, computed online. |
eprop_forward | The network eprop trains, over time-major inputs [T, B, in]: a RecurrentCell of cell and an LICell readout of time constant tau, both stepped at dt. |
first_spike_cross_entropy | Cross entropy of softmax(-t_first / tau) over output neurons, the time-to-first-spike loss of Göltz et al. (2021) and Wunderlich and Pehle’s latency tasks. |
fold_batch_norm | The stack with each nn.BatchNorm folded into the dense or convolutional layer before it. |
normalize | Scale weights and biases so each layer’s activations on inputs stay at or below the threshold 1. |
ottt | OTTT’s gradients for a feedforward stack, computed online. |
ottt_dense | spikes @ weight whose weight gradient pairs the error with the presynaptic trace, OTTT’s op. |
policy_gradient | REINFORCE’s estimate of the gradient of the expected reward: the mean over the batch of (rewards - baseline) * eligibility, for each example’s reward rewards [B]. |
residual_mlp | Their residual MLP of depth layers at their mean-field scales (pcalm.model.model_scales). |
reward_diffusion | Spread reward from the unit root backward along wiring, and each connection’s weight change. |
reward_shares | Each connection’s share of its target’s reward, [E]: a softmax of |activity| at the sources of the target’s incoming connections. |
run_converted | The converted network’s output firing rates [B, ...] over steps steps of constant inputs. |
sequential_blocks | The blocks and parameters of a Flax nn.Sequential, one layer of the stack per element. |
spike_times | The output spike times of a layer of neurons driven by input spike times. |
squared_error | Half the squared error per example, the loss of their experiments. |
Diffusion
Section titled “Diffusion”class Diffusion(NamedTuple)sparx.learn.diffusion on GitHub
What a reward spread to: each connection’s weight change, each unit’s local reward and each connection’s share of its target’s.
| Field | Type | Default |
|---|---|---|
change | jax.Array | |
credit | jax.Array | |
share | jax.Array |
EPropParams
Section titled “EPropParams”class EPropParams(NamedTuple)A recurrent layer, x_t = u_t @ w_in + z_{t-1} @ w_rec, and its leaky readout,
y_t = decay(tau, dt) y_{t-1} + z_t @ w_out + b_out.
| Field | Type | Default |
|---|---|---|
w_in | jax.Array | |
w_rec | jax.Array | |
w_out | jax.Array | |
b_out | jax.Array |
EventLIF
Section titled “EventLIF”class EventLIFThe membrane and synaptic time constants tau_m and tau_syn (ms), and the threshold v_th (mV
above rest).
iterations is the bisection steps that find each crossing. The closed
form of the membrane divides by tau_syn - tau_m, so the two must
differ.
| Field | Type | Default |
|---|---|---|
tau_m | float | 20.0 |
tau_syn | float | 5.0 |
v_th | float | 1.0 |
iterations | int | 60 |
OTTTLayer
Section titled “OTTTLayer”class OTTTLayer(NamedTuple)| Field | Type | Default |
|---|---|---|
weight | jax.Array | |
bias | jax.Array |
PredictiveCoding
Section titled “PredictiveCoding”class PredictiveCodingsparx.learn.predictive on GitHub
Predictive coding’s inference and weight update; alpha > 0 makes it PC-ALM.
The activity relaxes in budget cycles of inner_steps steps of
state_lr down each example’s own energy, every cycle but the last
followed by a multiplier step lambda_l += alpha r_l. alpha = 0 with
one inner step is PC with budget activity steps. credit says which
multipliers the weight update reads: those its last activity steps ran
with ("before", their pre_dual_energy and the paper’s Algorithm 1),
or those one more multiplier step gives ("after", their
post_dual_energy). rho is the penalty on the prediction errors, 1
in their experiments, where state_lr is near 1 / lambda_max of the
energy’s Hessian in the activity (their eta_best_by_cell.csv) and
budget is 2 L.
| Field | Type | Default |
|---|---|---|
budget | int | struct.field(pytree_node=False) |
state_lr | float | jax.Array | |
rho | float | jax.Array | 1.0 |
alpha | float | jax.Array | 0.0 |
inner_steps | int | struct.field(pytree_node=False, default=1) |
credit | Literal['before', 'after'] | struct.field(pytree_node=False, default='before') |
PredictiveCoding.energy
Section titled “PredictiveCoding.energy”def energy(blocks: Sequence[Block], params: Sequence[Params], x: jax.Array, y: jax.Array, settled: Settled, loss: Loss = squared_error) -> jax.ArrayEach example’s energy [B] at settled: the output’s loss and the shifted prediction errors.
PredictiveCoding.settle
Section titled “PredictiveCoding.settle”def settle(blocks: Sequence[Block], params: Sequence[Params], x: jax.Array, y: jax.Array, loss: Loss = squared_error) -> SettledThe hidden activity after inference from the forward pass, and the multipliers the update reads.
PredictiveCoding.gradient
Section titled “PredictiveCoding.gradient”def gradient(blocks: Sequence[Block], params: Sequence[Params], x: jax.Array, y: jax.Array, loss: Loss = squared_error, rows: jax.Array | None = None) -> list[Params]The weight update: the gradient of the energy summed over the batch’s rows at the settled state.
rows weighs each example (1 for a real row, 0 for a repeat); their
reference takes the mean over the batch, the gradient here divided by
the batch size.
RNeuralNet
Section titled “RNeuralNet”class RNeuralNetsparx.learn.diffusion on GitHub
RNeuralNet-Research’s network, clocked: neurons internal neurons, inputs input neurons and the
reward feeder, in that order, as one recurrent layer.
cell is a PulseCell per unit fed back through a Sparse wiring
with each connection’s delay; the input neurons pass their values on
unchanged (a threshold of minus infinity) and the feeder sends nothing.
outputs are the internal neurons wired to the feeder, whose last
outputs share the reward first. Build one with wire, from connections
in the order the original makes them, or random, as its
NeuralNet_init draws them; run it over inputs [T, ..., inputs]
and learn from a reward.
| Field | Type | Default |
|---|---|---|
cell | RecurrentCell[MembraneState, None] | |
outputs | jax.Array | |
neurons | int | struct.field(pytree_node=False) |
inputs | int | struct.field(pytree_node=False) |
wiring | Sparse | |
feeder | int |
RNeuralNet.wire
Section titled “RNeuralNet.wire”def wire(threshold: np.ndarray, pre: np.ndarray, post: np.ndarray, weight: np.ndarray, myelin: np.ndarray, inputs: int, outputs: np.ndarray) -> RNeuralNetThe network of the connections pre[e] -> post[e], in the order they were made.
Units count the len(threshold) internal neurons, then inputs
input neurons, then the feeder. Each connection has its weight and
the original’s Myelin; its delay follows from its place in the
pass (see the module). A neuron sends once a tick, so its outgoing
connections must be made one after another: NeuralNet_init makes
an output neuron’s connection to the feeder after every other, which
in the pass splits its sending in two, and here it is made with the
neuron’s others. No connection may enter an input neuron, leave the
feeder, or return to its source.
RNeuralNet.random
Section titled “RNeuralNet.random”def random(seed: int, neurons: int, inputs: int, outputs: int, fan: int = 16, input_fan: int = 5, myelin: int | None = None) -> RNeuralNetA network drawn as NeuralNet_init draws it, from a numpy generator seeded with seed.
Thresholds are normal around 2 with deviation 0.5. Each neuron
connects to fan others (MAX_AXIONS), none receiving more than
fan (MAX_DENDRITES), with a normal weight of deviation
1 / sqrt(fan) and a Myelin from 0 to 19; one that finds no free
target makes fewer. Each input neuron connects to input_fan
neurons with weight 1 and Myelin 1, and outputs neurons connect
to the feeder with weight 1 and Myelin 4, no neuron taking two of
these. The original wires all but its last neuron, which an
off-by-one leaves unconnected; here every neuron is wired. myelin
gives every connection between neurons that Myelin in place of
the draw, one delay throughout but for the pass order’s tick.
RNeuralNet.drive
Section titled “RNeuralNet.drive”def drive(x: jax.Array) -> jax.ArrayThe layer’s input [..., size] for the input neurons’ values x [..., inputs].
RNeuralNet.run
Section titled “RNeuralNet.run”def run(x: jax.Array, state: RecurrentState[MembraneState, None] | None = None) -> tuple[jax.Array, RecurrentState[MembraneState, None]]Each unit’s output each tick [T, ..., size] over input values x [T, ..., inputs], and the
final state; state None starts with nothing on its way.
RNeuralNet.learn
Section titled “RNeuralNet.learn”def learn(activity: jax.Array, reward: jax.Array | float, paths: Paths = 'first', discount: float = 1.0, eta: float = 0.01) -> tuple[RNeuralNet, Diffusion]The network after reward spreads from the feeder over the outputs activity [size] of the
tick it came in (reward_diffusion), and what it spread to.
ReinforceParams
Section titled “ReinforceParams”class ReinforceParams(NamedTuple)sparx.learn.reinforce on GitHub
A recurrent layer, x_t = u_t @ w_in + s_{t-1} @ w_rec, [in, N] and [N, N]; or each
example’s eligibility for them, with a leading batch axis.
| Field | Type | Default |
|---|---|---|
w_in | jax.Array | |
w_rec | jax.Array |
ResidualBlock
Section titled “ResidualBlock”class ResidualBlock(nn.Module)sparx.learn.predictive on GitHub
A layer of Seely and Gould’s residual MLP: scale * f(below) @ kernel, plus below when skip.
The input layer reads the input as it is (activation=None); every
other layer applies the activation to the activity below first. The
kernel starts at N(0, 1) entrywise and scale sets its effective size.
| Field | Type | Default |
|---|---|---|
features | int | |
scale | float | |
activation | Literal['linear', 'tanh', 'relu'] | None | None |
skip | bool | False |
ResidualBlock.__call__
Section titled “ResidualBlock.__call__”def __call__(below: jax.Array) -> jax.ArraySettled
Section titled “Settled”class Settled(NamedTuple)sparx.learn.predictive on GitHub
| Field | Type | Default |
|---|---|---|
activity | tuple[jax.Array, ...] | |
multipliers | tuple[jax.Array, ...] |
SpikingMaxPool
Section titled “SpikingMaxPool”class SpikingMaxPool(Neuron)Max pooling of spikes [T, ..., H, W, C] over unpadded window patches, gated by spike counts.
Rueckauer et al. (section 2.2.6) gate a spiking max-pool unit so it passes only the spikes of its most active input, judged by an estimate of the input rates, for example their online average. Here the estimate is the spike count so far, the online average times the step count, which ranks inputs the same way. Counting before this step keeps an input from winning a tie by the spike it is firing now, which would let the unit fire faster than any of its inputs. The counts are the layer’s state, carried across calls like a neuron’s.
| Field | Type | Default |
|---|---|---|
window | Pair | (2, 2) |
strides | Pair | (2, 2) |
SpikingMaxPool.build
Section titled “SpikingMaxPool.build”def build(x: jax.Array) -> _Gateaccumulate
Section titled “accumulate”def accumulate(step: Callable[[P, C, X], tuple[jax.Array, C]], params: P, carry: C, xs: X) -> tuple[jax.Array, P, C]Sum each step’s loss and its gradient over a sequence, holding the carry fixed in each step’s gradient.
step(params, carry, x_t) -> (loss_t, carry). Exact for a model whose
carry the gradient does not cross (a detached membrane, as OTTT has),
and then memory stays that of one step. Returns the total loss, the
summed gradient, and the final carry.
bptt_loss
Section titled “bptt_loss”def bptt_loss(cell: NeuronModel, params: EPropParams, inputs: jax.Array, targets: jax.Array, loss: Loss, tau: float, dt: float = 1.0, cut_recurrence: bool = False, precision: PrecisionLike = None) -> jax.ArrayThe summed loss of eprop_forward’s readout; its gradient is BPTT’s, or e-prop’s with
cut_recurrence.
eligibility_traces
Section titled “eligibility_traces”def eligibility_traces(cell: NeuronModel, params: EPropParams, inputs: jax.Array, dt: float = 1.0, precision: PrecisionLike = None) -> jax.ArrayEvery step’s eligibility traces dz_t/dW through each neuron’s own state, [T, B, N, in + N],
for analysis; eprop uses them as they are made instead of storing them.
def eprop(cell: NeuronModel, params: EPropParams, inputs: jax.Array, targets: jax.Array, loss: Loss, tau: float, dt: float = 1.0, feedback: jax.Array | None = None, precision: PrecisionLike = None) -> tuple[jax.Array, EPropParams]e-prop’s gradients for eprop_forward’s network, computed online.
Each synapse i -> j keeps an eligibility vector, the derivative of
neuron j’s state by the weight through j’s own dynamics, advanced
every step by the neuron’s state Jacobian; its eligibility trace is the
spike’s derivative through it. The readout’s leak filters the traces,
so the learning signal dloss_t/dy_t @ feedback.T of each step weights
them directly. feedback is w_out (symmetric e-prop) unless given
(random e-prop). The readout’s own gradients are exact.
The vectors are derived for any elementwise model from the forward
derivative of its step. A state variable whose input enters with
constant coefficients shared by all neurons (the membrane of a model
with a detached reset) needs only the filtered presynaptic activity,
and one the gradient never reaches (a refractory count) needs nothing.
Memory is B x N x P for the filtered traces, plus as much for each
remaining state variable, P = in + N, whatever the sequence length.
It runs in the widest dtype of the inputs, the parameters and float32,
and precision is its matrix products’. Returns the summed loss and the
gradients, each in its parameter’s dtype.
eprop_forward
Section titled “eprop_forward”def eprop_forward(cell: NeuronModel, params: EPropParams, inputs: jax.Array, tau: float, dt: float = 1.0, cut_recurrence: bool = False, precision: PrecisionLike = None) -> tuple[jax.Array, jax.Array]The network eprop trains, over time-major inputs [T, B, in]: a RecurrentCell of cell
and an LICell readout of time constant tau, both stepped at dt.
Returns the readout [T, B, out] and the recurrent layer’s spikes
[T, B, N], in the widest dtype of the inputs, the parameters and
float32. With cut_recurrence the gradient stops at the fed-back
spikes, which makes BPTT’s gradient e-prop’s. precision is the
matrix products’.
first_spike_cross_entropy
Section titled “first_spike_cross_entropy”def first_spike_cross_entropy(times: jax.Array, labels: jax.Array, tau: float = 5.0, silent: float = 10000.0) -> jax.ArrayCross entropy of softmax(-t_first / tau) over output neurons, the time-to-first-spike loss of
Göltz et al. (2021) and Wunderlich and Pehle’s latency tasks.
times [B, N, K], labels [B]; returns the batch mean. A silent
neuron counts as firing at silent ms. A silent neuron carries no
gradient whatever this is, so a large value makes silence a cliff in
the loss; the horizon of the simulation keeps the loss bounded.
fold_batch_norm
Section titled “fold_batch_norm”def fold_batch_norm(model: nn.Sequential, variables: StackVariables) -> tuple[nn.Sequential, dict[str, dict[str, Params]]]The stack with each nn.BatchNorm folded into the dense or convolutional layer before it.
Inference-mode batch norm computes scale (z - mean) / sigma + offset
per output channel with sigma = sqrt(var + epsilon); applied to
z = x W + b it is again affine in x, with W scale / sigma and
scale (b - mean) / sigma + offset (their section 2.2.3). The running
statistics come from variables["batch_stats"] and epsilon from the
layer. Returns the stack without its batch norms, and its variables.
normalize
Section titled “normalize”def normalize(model: nn.Sequential, variables: StackVariables, inputs: jax.Array, percentile: float = 99.9) -> dict[str, dict[str, Params]]Scale weights and biases so each layer’s activations on inputs stay at or below the threshold 1.
Dense or convolutional layer l with scale lambda_l, the
percentile-th percentile of its positive activations, becomes
W lambda_{l-1} / lambda_l and b / lambda_l (lambda_0 = 1 for
inputs in [0, 1]). Pooling and flattening commute with a positive
scale, so they keep the scale of the layer before them. The last layer
is scaled the same way, which keeps the argmax. Fold batch norm first.
def ottt(cells: Sequence[NeuronModel], layers: Sequence[OTTTLayer], inputs: jax.Array, targets: jax.Array, loss: Loss, tau: float, dt: float = 1.0) -> tuple[jax.Array, list[OTTTLayer]]OTTT’s gradients for a feedforward stack, computed online.
layers[k] feeds cells[k], and one more layer reads the last cell’s
spikes out, so len(layers) == len(cells) + 1. The first layer sees the
input as it is; every later one sees spikes, and learns from their
trace a_t = decay(tau, dt) a_{t-1} + s_t (their rate_tracking).
Their trace decays by 1 - dt / tau, the forward Euler step of the
same leak; tau = -dt / log(1 - dt / tau_theirs) gives their decay.
Membranes carry no gradient between steps, so each step’s loss
differentiates through that step alone; give the cells
detach_reset=True, as their OnlineLIFNode has.
Returns the summed loss and the gradients.
ottt_dense
Section titled “ottt_dense”def ottt_dense(spikes: jax.Array, trace: jax.Array, weight: jax.Array) -> jax.Arrayspikes @ weight whose weight gradient pairs the error with the presynaptic trace, OTTT’s op.
policy_gradient
Section titled “policy_gradient”def policy_gradient(eligibility: ReinforceParams, rewards: jax.Array, baseline: jax.Array | float = 0.0) -> ReinforceParamssparx.learn.reinforce on GitHub
REINFORCE’s estimate of the gradient of the expected reward: the mean over the batch of
(rewards - baseline) * eligibility, for each example’s reward rewards [B].
Every eligibility has mean zero over the spikes it was drawn with, so a
baseline that does not depend on those spikes (a running mean of the
reward, say) leaves the estimate unbiased and can shrink its variance.
Ascend it to raise the reward.
residual_mlp
Section titled “residual_mlp”def residual_mlp(width: int, depth: int, inputs: int, outputs: int, activation: Literal['linear', 'tanh', 'relu'] = 'relu') -> nn.Sequentialsparx.learn.predictive on GitHub
Their residual MLP of depth layers at their mean-field scales (pcalm.model.model_scales).
The input layer scales by 1 / sqrt(inputs), the depth - 2 interior
layers, each with an identity skip, by 1 / sqrt(width * depth), and
the readout by 1 / width (Innocenti et al.’s parameterization, their
Appendix B).
reward_diffusion
Section titled “reward_diffusion”def reward_diffusion(wiring: Sparse, activity: jax.Array, reward: jax.Array | float, root: int | jax.Array, paths: Paths = 'first', discount: float = 1.0, eta: float = 0.01) -> Diffusionsparx.learn.diffusion on GitHub
Spread reward from the unit root backward along wiring, and each connection’s weight change.
For one example: activity is each unit’s last output, [size],
reward a scalar and root a unit, which may differ by example;
jax.vmap it over a batch. paths="first" is the
original’s depth-first spread, where each unit passes on the reward it
holds when the spread first reaches it, through its incoming
connections in their order in the wiring; a unit that holds no reward
then, or has no inputs, passes nothing and may be reached again.
paths="all" passes on every path’s reward, which needs a discount
below 1 to converge on a cycle. Either way each connection passes
discount times its share (1 in the original), and a connection whose
target holds no reward keeps its weight.
reward_shares
Section titled “reward_shares”def reward_shares(wiring: Sparse, activity: jax.Array) -> jax.Arraysparx.learn.diffusion on GitHub
Each connection’s share of its target’s reward, [E]: a softmax of |activity| at the sources
of the target’s incoming connections.
activity is each unit’s last output, [size]. The exponentials are
taken less each target’s largest, so they stay finite where the
original’s float exponentials overflow (an activity past 88); elsewhere
the shares are the original’s up to rounding.
run_converted
Section titled “run_converted”def run_converted(model: nn.Sequential, variables: StackVariables, inputs: jax.Array, steps: int, chunk: int = 50) -> jax.ArrayThe converted network’s output firing rates [B, ...] over steps steps of constant inputs.
The input is analog, the same every step, as Rueckauer et al. drive
the first layer. The steps run chunk at a time, the neurons’ state
carried between calls in the "state" collection, so memory holds
chunk steps of every layer’s activity whatever steps is; the
result equals one call over all steps.
sequential_blocks
Section titled “sequential_blocks”def sequential_blocks(model: nn.Sequential, params: Mapping[str, Params]) -> tuple[list[Block], list[Params]]sparx.learn.predictive on GitHub
The blocks and parameters of a Flax nn.Sequential, one layer of the stack per element.
An element without parameters, such as nn.relu, is a layer of its own
with empty parameters, so its output is held free like any other.
spike_times
Section titled “spike_times”def spike_times(inputs: jax.Array, weights: jax.Array, neuron: EventLIF, horizon: float, capacity: int) -> tuple[jax.Array, jax.Array]The output spike times of a layer of neurons driven by input spike times.
inputs [B, M, K] holds each input’s spike times (ms, inf where
absent), weights [M, N] (pA). Returns the spike times [B, N, capacity]
up to horizon ms, inf where absent, and the number of neurons that
fire more than capacity times [B], whose first capacity spikes are
still exact. Differentiable in the weights and the input times.
Each neuron walks its own way through the inputs: a move takes it to its
next crossing, or, when the membrane stays below threshold until the next
input, to that input. A neuron makes one move per input before the
horizon, one per spike and one to the horizon, so M K + capacity + 1
moves finish every neuron within capacity and reach the spike past it of
every other.
squared_error
Section titled “squared_error”def squared_error(prediction: jax.Array, target: jax.Array) -> jax.Arraysparx.learn.predictive on GitHub
Half the squared error per example, the loss of their experiments.