Skip to content
GitHub

sparx.nn

Flax linen layers for spiking networks, over time-major inputs [T, ...].

Name
ALIFAdaptive-threshold LIF (sparx.dynamics.ALIFCell).
IFIntegrate-and-fire: LIF without leak.
LIA leaky integrator readout (sparx.dynamics.LICell); returns its membrane, [T, ...].
LIFLeaky integrate-and-fire (sparx.dynamics.LIFCell) with time constant tau.
PSNThe parallel spiking neuron over all T x T step pairs, for a fixed T.
RATESThe collection spiking layers sow their per-example, per-neuron firing rates into.
STATEThe collection a layer carries its neurons’ state in across apply calls.
BatchMajorRun a time-major layer (a sparx layer or a stack of them) on batch-major input [B, T, ...].
DecayingTracesparx.dynamics.DecayingHebb with a learned rate starting at eta, Miconi et al.’s 0.01.
DelayedDenseA dense layer whose every synapse has a learnable delay of 0 to max_delay steps.
DynamicsAny neuron model of sparx.dynamics as a layer, its fields fixed.
FlattenFlatten each example’s trailing ndim axes into one feature axis, channels first.
FlattensA layer that flattens each example’s trailing axes into one feature axis in PyTorch’s order, channels first, and changes nothing else: flattened_axes() is how many trailing axes.
HebbianTraceThe Hebbian trace of a Recurrent layer, which builds its rule with the rule’s learned parameters.
MaskedPSNThe PSN restricted to the k most recent steps, for a fixed T.
ModelledA layer that runs a sparx.dynamics model over time, stepped at dt.
ModulatedTracesparx.dynamics.ModulatedHebb: a learned neuromodulator sets each unit’s rate, clipped at clip.
NeuronA population of neurons run over the leading (time) axis of its input.
OjaTracesparx.dynamics.OjaHebb, Oja’s rule, with a learned rate starting at eta.
RateA leaky rate unit (sparx.dynamics.RateCell), FLYNN’s neuron; returns its activity, [T, ...].
RecurrentFeed neuron’s output back into its input through a learned [F, F] matrix, fixed or plastic.
RetroactiveTracesparx.dynamics.RetroactiveHebb: a learned neuromodulator writes recent coactivity into the weights.
SlidingPSNk weights slid over time: H[t] = sum_i weight[i] * X[t - k + 1 + i] + bias.
SynapticCurrent-based LIF, Serial(LICell, LIFCell): synaptic time constant tau_synapse, membrane time constant tau.
adoptchild (a neuron, a Hebbian trace) as owner’s child called name, so its parameters sit under owner wherever it was built.
band_maskM[i, j] = 1 where j <= i <= j + k - 1: each step sees itself and the k - 1 before it.
delay_kernelEach synapse’s weights over its max_delay + 1 lags, [max_delay + 1, *delay.shape].
history_windowx [T, ...] with the held steps before it prepended, [held + T, ...], in x’s dtype.
record_ratesSow spikes’ time-averaged rate into "spike_rates" when that collection is mutable.
class ALIF(Neuron)

sparx.nn.neurons on GitHub

Adaptive-threshold LIF (sparx.dynamics.ALIFCell).

Each spike raises the threshold by beta, and the rise decays with time constant tau_adapt. The defaults are Bellec et al.’s (2020) in steps of 1 ms. learn_tau learns both decays per feature. refractory is a duration in the unit of dt.

FieldTypeDefault
taufloat20.0
tau_adaptfloat200.0
betafloat1.8
thresholdfloat1.0
resetReset'subtract'
surrogateSurrogateATan()
detach_resetboolFalse
learn_tauboolFalse
refractoryfloat0
def build(x: jax.Array) -> ALIFCell
class IF(Neuron)

sparx.nn.neurons on GitHub

Integrate-and-fire: LIF without leak.

FieldTypeDefault
thresholdfloat1.0
resetReset'subtract'
surrogateSurrogateATan()
detach_resetboolFalse
def build(x: jax.Array) -> LIFCell
class LI(Neuron)

sparx.nn.neurons on GitHub

A leaky integrator readout (sparx.dynamics.LICell); returns its membrane, [T, ...].

FieldTypeDefault
taufloat2.0
learn_tauboolFalse
def build(x: jax.Array) -> LICell
class LIF(Neuron)

sparx.nn.neurons on GitHub

Leaky integrate-and-fire (sparx.dynamics.LIFCell) with time constant tau.

learn_tau learns one decay per feature (last axis), starting at tau.

FieldTypeDefault
taufloat2.0
thresholdfloat1.0
resetReset'subtract'
surrogateSurrogateATan()
detach_resetboolFalse
learn_tauboolFalse
def build(x: jax.Array) -> LIFCell
class PSN(nn.Module)

sparx.nn.parallel on GitHub

The parallel spiking neuron over all T x T step pairs, for a fixed T.

T is the input’s leading axis at init. Each output step reads every input step, earlier and later, so the layer suits classification of a whole sequence, not causal or streaming use.

FieldTypeDefault
surrogateSurrogateATan()
precisionjax.lax.Precision | NoneNone
def __call__(x: jax.Array) -> jax.Array
RATES = 'spike_rates'

sparx.nn.neurons on GitHub

The collection spiking layers sow their per-example, per-neuron firing rates into.

STATE = 'state'

sparx.nn.neurons on GitHub

The collection a layer carries its neurons’ state in across apply calls.

class BatchMajor(nn.Module)

sparx.nn.reshape on GitHub

Run a time-major layer (a sparx layer or a stack of them) on batch-major input [B, T, ...].

dew batches records along the first axis, so a model that dew’s Supervised objective trains reads [B, T, ...] and its loss reads one row per example. This wrapper hands layer the input time-major, [T, B, ...], and returns its output batch-major again. The layer’s parameters sit under layer.

FieldTypeDefault
layernn.Module
def __call__(x: jax.Array) -> jax.Array
class DecayingTrace(HebbianTrace)

sparx.nn.hebbian on GitHub

sparx.dynamics.DecayingHebb with a learned rate starting at eta, Miconi et al.’s 0.01.

FieldTypeDefault
etafloat0.01
def build(features: int) -> DecayingHebb
class DelayedDense(nn.Module)

sparx.nn.delays on GitHub

A dense layer whose every synapse has a learnable delay of 0 to max_delay steps.

Delays start uniform over [0, max_delay]. sigma, given at each call, is the Gaussian’s width in steps: Hammouamri et al. start it near max_delay / 2 and decrease it to 0 over training. It may be traced, as a schedule’s value, as long as it stays positive; pass the Python number 0 for the deployed network. The cost is that of a dense layer applied max_delay + 1 times.

dtype is the dtype the lags’ products run in; None keeps float32, or the input’s dtype when it is wider. param_dtype stores the weights and the bias. The delays stay float32 whatever it is, because a bfloat16 delay moves in steps of 1/16 near 24 and the Gaussian kernel is built from it in float32. precision is the matmuls’ precision.

FieldTypeDefault
featuresint
max_delayint
use_biasboolTrue
kernel_initnn.initializers.Initializernn.initializers.lecun_normal()
dtypeDtype | NoneNone
param_dtypeDtypejnp.float32
precisionPrecisionLikeNone
def __call__(x: jax.Array, sigma: float | jax.Array) -> jax.Array
class Dynamics(Neuron)

sparx.nn.neurons on GitHub

Any neuron model of sparx.dynamics as a layer, its fields fixed.

Dynamics(AdEx(), dt=0.1) runs AdEx in steps of 0.1 ms on input currents [T, ...] in pA, and Dynamics(Izhikevich()) Izhikevich’s neuron in steps of 1 ms on currents in its own units; the model defaults to LeakyIntegrateAndFire with its defaults. drive names what the input is to the model: a "current" held over each step, the physical models’ input, or a "jump" of the membrane, the dimensionless family’s. A layer that learns a model’s constants builds the model from its parameters, as LIF does.

Training through a physical model needs a steeper surrogate than the default. The surrogate is the model’s own (AdEx(surrogate=...)) and reads v - threshold in mV, so its slope is per mV. Backpropagation also runs through the membrane equation, and AdEx’s exponential upswing multiplies the gradient at every step a neuron spends near its peak. A heavy-tailed surrogate passes gradient from all of those steps. For AdEx at 40 Hz over 2000 steps of 0.1 ms, behind a dense layer, the gradient norm reaching that layer was 7.9e8 with the default ATan(), 3.4 with FastSigmoid(25) and 0.42 with FastSigmoid(100); the LeakyIntegrateAndFire, which has no upswing, gave 0.28 with FastSigmoid(25). A wider surrogate (one normalized by a voltage scale, ATan(0.1)) made AdEx’s gradient larger, so a steep FastSigmoid is the recommended start. tests/test_nn.py reruns the AdEx case with ATan() and FastSigmoid(100).

FieldTypeDefault
neuronNeuronModeldynamics.LeakyIntegrateAndFire()
driveLiteral['current', 'jump']'current'
def inputs(x: jax.Array) -> SynapticInput
def build(x: jax.Array) -> NeuronModel
class Flatten(nn.Module)

sparx.nn.reshape on GitHub

Flatten each example’s trailing ndim axes into one feature axis, channels first.

Flax lays images out channels last, [..., H, W, C], and a plain reshape would order their features H, W, C. PyTorch and NIR lay images out channels first and flatten them C, H, W. This layer moves the channel axis (the last) in front of the other flattened axes before the reshape, so feature c * H * W + h * W + w holds channel c at (h, w), the order of torch.nn.Flatten. A dense layer after it then takes PyTorch’s or NIR’s weights as they are, at the cost of one transpose. The axes before the last ndim (time and batch) are kept.

FieldTypeDefault
ndimint3
def flattened_axes() -> int

ndim, the trailing axes this layer flattens (Flattens).

def __call__(x: jax.Array) -> jax.Array
class Flattens(Protocol)

sparx.nn.reshape on GitHub

A layer that flattens each example’s trailing axes into one feature axis in PyTorch’s order, channels first, and changes nothing else: flattened_axes() is how many trailing axes.

A consumer that maps layers to another library’s (NIR’s Flatten) or passes them through a conversion asks a layer for this, not for its class.

def flattened_axes() -> int
def __call__(x: jax.Array) -> jax.Array
class HebbianTrace(nn.Module)

sparx.nn.hebbian on GitHub

The Hebbian trace of a Recurrent layer, which builds its rule with the rule’s learned parameters.

A subclass declares the parameters in build(features), for features units, and returns the rule, a sparx.dynamics.HebbianRule. A layer calls it as its child rule, so they sit under rule.

def build(features: int) -> HebbianRule
def __call__(features: int) -> HebbianRule
class MaskedPSN(nn.Module)

sparx.nn.parallel on GitHub

The PSN restricted to the k most recent steps, for a fixed T.

masking blends the band mask with all ones, masking * M + (1 - masking): 0 is the unmasked PSN, 1 the causal order-k neuron. Fang et al. raise it from 0 to 1 during training (SpikingJelly’s lambda_); pass the schedule’s value at each call. It may be traced.

FieldTypeDefault
kint
surrogateSurrogateATan()
precisionjax.lax.Precision | NoneNone
def __call__(x: jax.Array, masking: float | jax.Array = 1.0) -> jax.Array
class Modelled(Protocol)

sparx.nn.neurons on GitHub

A layer that runs a sparx.dynamics model over time, stepped at dt.

model(x) builds the model with the layer’s parameters for inputs like x, called through apply(variables, x, method="model"). A consumer that maps layers to another library’s (NIR’s LIF) reads the model it gets, a LIFCell or a RecurrentCell, and does not ask for the layer’s class. Every Neuron is one.

FieldTypeDefault
dtfloat
def model(x: jax.Array) -> NeuronModel
class ModulatedTrace(HebbianTrace)

sparx.nn.hebbian on GitHub

sparx.dynamics.ModulatedHebb: a learned neuromodulator sets each unit’s rate, clipped at clip.

The modulator reads the units ([F] and a bias) and fans out to them ([F] and [F]), Backpropamine’s h2mod and modfanout, which start as torch.nn.Linear does; clip is their code’s 2.

FieldTypeDefault
clipfloat2.0
def build(features: int) -> ModulatedHebb
class Neuron(nn.Module)

sparx.nn.neurons on GitHub

A population of neurons run over the leading (time) axis of its input.

A subclass says how to build its model (sparx.dynamics.NeuronModel) from the input, declaring any parameters there; this base runs it with sparx.dynamics.run, carries the "state" collection and sows "spike_rates" for a spiking model. The input reaches the model as inputs(x) says, a jump of the membrane unless a subclass says otherwise. dt is the step in the unit of the model’s time constants. unroll is the number of time steps one iteration of the compiled loop holds (jax.lax.scan).

FieldTypeDefault
dtfloatdataclasses.field(default=1.0, kw_only=True)
unrollintdataclasses.field(default=1, kw_only=True)
def build(x: jax.Array) -> NeuronModel

The model for inputs like x, [T, ..., features].

def inputs(x: jax.Array) -> SynapticInput

What the model receives from the layer’s input x: a jump of its membrane.

def model(x: jax.Array) -> NeuronModel

This layer’s model, with its parameters, for inputs like x.

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

Run over x, [T, ...], and return the outputs [T, ...].

def run(model: NeuronModel, x: jax.Array) -> jax.Array
class OjaTrace(HebbianTrace)

sparx.nn.hebbian on GitHub

sparx.dynamics.OjaHebb, Oja’s rule, with a learned rate starting at eta.

FieldTypeDefault
etafloat0.01
def build(features: int) -> OjaHebb
class Rate(Neuron)

sparx.nn.neurons on GitHub

A leaky rate unit (sparx.dynamics.RateCell), FLYNN’s neuron; returns its activity, [T, ...].

h <- alpha h + (1 - alpha) f(x + bias) with alpha = exp(-dt / tau). bias is learned per feature, starting at 0, as FLYNN learns theirs; learn_tau learns the leak per feature, starting at tau. Recurrent(Rate()) is FLYNN’s recurrence with a dense matrix. tau=0 keeps no memory, h = f(x + bias), the unit of an Elman network and of Miconi et al.’s plastic networks (Recurrent(Rate(tau=0), rule=...)).

FieldTypeDefault
taufloat2.0
activationLiteral['tanh', 'relu', 'sigmoid']'tanh'
learn_tauboolFalse
def build(x: jax.Array) -> RateCell
class Recurrent(Neuron)

sparx.nn.neurons on GitHub

Feed neuron’s output back into its input through a learned [F, F] matrix, fixed or plastic.

Recurrent(ALIF()) is a recurrent adaptive layer (LSNN). The input projection stays outside, as an nn.Dense before this layer, so it runs over all time steps at once; only the feedback product runs inside the loop. The wrapped neuron’s parameters live under neuron, and the layer’s dt must be the neuron’s.

With a rule (sparx.nn.DecayingTrace, OjaTrace, ModulatedTrace, RetroactiveTrace or a HebbianTrace of your own), the feedback runs through recurrent + alpha * hebb: fast weights a Hebbian trace writes from zero in every sequence, differentiable plasticity (Miconi et al. 2018) and Backpropamine (Miconi et al. 2019). alpha, [F, F], starts at alpha_init, small and random as theirs does (.01 * randn in their simple/simple.py; theirs also start recurrent at kernel_init=nn.initializers.normal(0.01)), and the rule’s parameters sit under rule. Recurrent(Rate(tau=0), rule=...) is their tanh network; a spiking neuron learns fast weights between its spikes. Backpropagating through a plastic layer holds a [B, F, F] trace per step.

Backpropagation through the feedback multiplies by the recurrent matrix at every step, and a heavy-tailed surrogate passes gradient through every neuron, even those far from threshold. Training the recurrent network of examples/train_shd.py with ATan grew the matrix’s spectral radius from 1 to 5 and the gradient norm past 1e8 within 300 steps; with FastSigmoid(100) the gradient norm stayed below 10.

FieldTypeDefault
neuronNeuronLIF()
kernel_initnn.initializers.Initializernn.initializers.orthogonal()
precisionPrecisionLikeNone
ruleHebbianTrace | NoneNone
alpha_initnn.initializers.Initializernn.initializers.normal(0.01)
def inputs(x: jax.Array) -> SynapticInput
def build(x: jax.Array) -> RecurrentCell
class RetroactiveTrace(HebbianTrace)

sparx.nn.hebbian on GitHub

sparx.dynamics.RetroactiveHebb: a learned neuromodulator writes recent coactivity into the weights.

The eligibility decays at a learned rate starting at eta, their 0.01. The modulator reads the units ([F] and a bias), Backpropamine’s h2DA, which starts as torch.nn.Linear does; clip is their code’s 1.

FieldTypeDefault
etafloat0.01
clipfloat1.0
def build(features: int) -> RetroactiveHebb
class SlidingPSN(nn.Module)

sparx.nn.parallel on GitHub

k weights slid over time: H[t] = sum_i weight[i] * X[t - k + 1 + i] + bias.

Works for any T, and is causal, so the "state" collection carries the last k - 1 inputs across calls and a stream can be fed in chunks. weight[k - 1] multiplies the current step; the first steps of a fresh sequence see zeros before it. exponential_init starts the weights at (..., 1/4, 1/2, 1), otherwise kaiming uniform (SpikingJelly’s exp_init). The membrane is k weighted slices of the window, O(T k), in membrane precision; SpikingJelly builds the [T, T] Toeplitz matrix of the weights instead.

FieldTypeDefault
kint
exponential_initboolTrue
surrogateSurrogateATan()
def __call__(x: jax.Array) -> jax.Array
class Synaptic(Neuron)

sparx.nn.neurons on GitHub

Current-based LIF, Serial(LICell, LIFCell): synaptic time constant tau_synapse, membrane time constant tau.

FieldTypeDefault
taufloat10.0
tau_synapsefloat5.0
thresholdfloat1.0
resetReset'subtract'
surrogateSurrogateATan()
detach_resetboolFalse
learn_tauboolFalse
def build(x: jax.Array) -> Serial
def adopt(child: Child, owner: nn.Module, name: str) -> Child

sparx.nn.neurons on GitHub

child (a neuron, a Hebbian trace) as owner’s child called name, so its parameters sit under owner wherever it was built.

A module built inside a parent’s compact method belongs to that parent, and one handed down from a parent’s field belongs to the parent; the clone makes it owner’s own. One built outside any module and given to owner as its field name is already that child.

def band_mask(steps: int, k: int) -> jax.Array

sparx.nn.parallel on GitHub

M[i, j] = 1 where j <= i <= j + k - 1: each step sees itself and the k - 1 before it.

def delay_kernel(delay: jax.Array, max_delay: int, sigma: float | jax.Array) -> jax.Array

sparx.nn.delays on GitHub

Each synapse’s weights over its max_delay + 1 lags, [max_delay + 1, *delay.shape].

A normalized Gaussian of width sigma centered at the delay, clipped to [0, max_delay]; at sigma 0 (a Python number), the one-hot of the rounded delay, with no gradient to the delay.

The clip passes the gradient through unchanged. DCLS has no clip in its kernel and clamps its positions after each update, so a delay at an end of the range gets the kernel’s full gradient there; jnp.clip would halve it at the end and zero it past it.

The normalization is a softmax of the log-density: a width far below a step underflows every lag’s density to 0 (at a delay of 0.5 and a width of 0.01 the largest log-density is -1250), where the softmax keeps the nearest lags.

def history_window(module: nn.Module, x: jax.Array, held: int) -> jax.Array

sparx.nn.neurons on GitHub

x [T, ...] with the held steps before it prepended, [held + T, ...], in x’s dtype.

A causal layer that reads held steps back keeps those steps in the "state" collection when it is mutable, so a stream fed in chunks sees the steps the previous chunk ended on; a fresh stream, or a call that does not carry state, sees zeros before its first step. The window’s last held steps are stored for the next call.

def record_rates(module: nn.Module, spikes: jax.Array) -> None

sparx.nn.neurons on GitHub

Sow spikes’ time-averaged rate into "spike_rates" when that collection is mutable.