Skip to content
GitHub

sparx.learn

Learning rules beyond surrogate-gradient backpropagation through time (design.md section 7).

Module
sparx.learn.convertConverting a trained ReLU network to a spiking one (Rueckauer et al., Frontiers in Neuroscience 2017).
sparx.learn.reinforceREINFORCE for a recurrent layer of stochastic spiking neurons, its gradient carried as eligibilities.
Name
DiffusionWhat a reward spread to: each connection’s weight change, each unit’s local reward and each connection’s share of its target’s.
EPropParamsA 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.
EventLIFThe membrane and synaptic time constants tau_m and tau_syn (ms), and the threshold v_th (mV above rest).
OTTTLayer
PredictiveCodingPredictive coding’s inference and weight update; alpha > 0 makes it PC-ALM.
RNeuralNetRNeuralNet-Research’s network, clocked: neurons internal neurons, inputs input neurons and the reward feeder, in that order, as one recurrent layer.
ReinforceParamsA 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.
ResidualBlockA layer of Seely and Gould’s residual MLP: scale * f(below) @ kernel, plus below when skip.
Settled
SpikingMaxPoolMax pooling of spikes [T, ..., H, W, C] over unpadded window patches, gated by spike counts.
accumulateSum each step’s loss and its gradient over a sequence, holding the carry fixed in each step’s gradient.
bptt_lossThe summed loss of eprop_forward’s readout; its gradient is BPTT’s, or e-prop’s with cut_recurrence.
eligibility_tracesEvery 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.
eprope-prop’s gradients for eprop_forward’s network, computed online.
eprop_forwardThe 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_entropyCross 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_normThe stack with each nn.BatchNorm folded into the dense or convolutional layer before it.
normalizeScale weights and biases so each layer’s activations on inputs stay at or below the threshold 1.
otttOTTT’s gradients for a feedforward stack, computed online.
ottt_densespikes @ weight whose weight gradient pairs the error with the presynaptic trace, OTTT’s op.
policy_gradientREINFORCE’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_mlpTheir residual MLP of depth layers at their mean-field scales (pcalm.model.model_scales).
reward_diffusionSpread reward from the unit root backward along wiring, and each connection’s weight change.
reward_sharesEach connection’s share of its target’s reward, [E]: a softmax of |activity| at the sources of the target’s incoming connections.
run_convertedThe converted network’s output firing rates [B, ...] over steps steps of constant inputs.
sequential_blocksThe blocks and parameters of a Flax nn.Sequential, one layer of the stack per element.
spike_timesThe output spike times of a layer of neurons driven by input spike times.
squared_errorHalf the squared error per example, the loss of their experiments.
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.

FieldTypeDefault
changejax.Array
creditjax.Array
sharejax.Array
class EPropParams(NamedTuple)

sparx.learn.online on GitHub

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.

FieldTypeDefault
w_injax.Array
w_recjax.Array
w_outjax.Array
b_outjax.Array
class EventLIF

sparx.learn.events on GitHub

The 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.

FieldTypeDefault
tau_mfloat20.0
tau_synfloat5.0
v_thfloat1.0
iterationsint60
class OTTTLayer(NamedTuple)

sparx.learn.online on GitHub

FieldTypeDefault
weightjax.Array
biasjax.Array
class PredictiveCoding

sparx.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.

FieldTypeDefault
budgetintstruct.field(pytree_node=False)
state_lrfloat | jax.Array
rhofloat | jax.Array1.0
alphafloat | jax.Array0.0
inner_stepsintstruct.field(pytree_node=False, default=1)
creditLiteral['before', 'after']struct.field(pytree_node=False, default='before')
def energy(blocks: Sequence[Block], params: Sequence[Params], x: jax.Array, y: jax.Array, settled: Settled, loss: Loss = squared_error) -> jax.Array

Each example’s energy [B] at settled: the output’s loss and the shifted prediction errors.

def settle(blocks: Sequence[Block], params: Sequence[Params], x: jax.Array, y: jax.Array, loss: Loss = squared_error) -> Settled

The hidden activity after inference from the forward pass, and the multipliers the update reads.

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.

class RNeuralNet

sparx.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.

FieldTypeDefault
cellRecurrentCell[MembraneState, None]
outputsjax.Array
neuronsintstruct.field(pytree_node=False)
inputsintstruct.field(pytree_node=False)
wiringSparse
feederint
def wire(threshold: np.ndarray, pre: np.ndarray, post: np.ndarray, weight: np.ndarray, myelin: np.ndarray, inputs: int, outputs: np.ndarray) -> RNeuralNet

The 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.

def random(seed: int, neurons: int, inputs: int, outputs: int, fan: int = 16, input_fan: int = 5, myelin: int | None = None) -> RNeuralNet

A 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.

def drive(x: jax.Array) -> jax.Array

The layer’s input [..., size] for the input neurons’ values x [..., inputs].

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.

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.

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.

FieldTypeDefault
w_injax.Array
w_recjax.Array
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.

FieldTypeDefault
featuresint
scalefloat
activationLiteral['linear', 'tanh', 'relu'] | NoneNone
skipboolFalse
def __call__(below: jax.Array) -> jax.Array
class Settled(NamedTuple)

sparx.learn.predictive on GitHub

FieldTypeDefault
activitytuple[jax.Array, ...]
multiplierstuple[jax.Array, ...]
class SpikingMaxPool(Neuron)

sparx.learn.convert on GitHub

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.

FieldTypeDefault
windowPair(2, 2)
stridesPair(2, 2)
def build(x: jax.Array) -> _Gate
def accumulate(step: Callable[[P, C, X], tuple[jax.Array, C]], params: P, carry: C, xs: X) -> tuple[jax.Array, P, C]

sparx.learn.online on GitHub

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.

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.Array

sparx.learn.online on GitHub

The summed loss of eprop_forward’s readout; its gradient is BPTT’s, or e-prop’s with cut_recurrence.

def eligibility_traces(cell: NeuronModel, params: EPropParams, inputs: jax.Array, dt: float = 1.0, precision: PrecisionLike = None) -> jax.Array

sparx.learn.online on GitHub

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.

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]

sparx.learn.online on GitHub

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.

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]

sparx.learn.online on GitHub

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’.

def first_spike_cross_entropy(times: jax.Array, labels: jax.Array, tau: float = 5.0, silent: float = 10000.0) -> jax.Array

sparx.learn.events on GitHub

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.

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.

def fold_batch_norm(model: nn.Sequential, variables: StackVariables) -> tuple[nn.Sequential, dict[str, dict[str, Params]]]

sparx.learn.convert on GitHub

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.

def normalize(model: nn.Sequential, variables: StackVariables, inputs: jax.Array, percentile: float = 99.9) -> dict[str, dict[str, Params]]

sparx.learn.convert on GitHub

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]]

sparx.learn.online on GitHub

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.

def ottt_dense(spikes: jax.Array, trace: jax.Array, weight: jax.Array) -> jax.Array

sparx.learn.online on GitHub

spikes @ weight whose weight gradient pairs the error with the presynaptic trace, OTTT’s op.

def policy_gradient(eligibility: ReinforceParams, rewards: jax.Array, baseline: jax.Array | float = 0.0) -> ReinforceParams

sparx.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.

def residual_mlp(width: int, depth: int, inputs: int, outputs: int, activation: Literal['linear', 'tanh', 'relu'] = 'relu') -> nn.Sequential

sparx.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).

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) -> Diffusion

sparx.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.

def reward_shares(wiring: Sparse, activity: jax.Array) -> jax.Array

sparx.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.

def run_converted(model: nn.Sequential, variables: StackVariables, inputs: jax.Array, steps: int, chunk: int = 50) -> jax.Array

sparx.learn.convert on GitHub

The 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.

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.

def spike_times(inputs: jax.Array, weights: jax.Array, neuron: EventLIF, horizon: float, capacity: int) -> tuple[jax.Array, jax.Array]

sparx.learn.events on GitHub

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.

def squared_error(prediction: jax.Array, target: jax.Array) -> jax.Array

sparx.learn.predictive on GitHub

Half the squared error per example, the loss of their experiments.