sparx.nn
Flax linen layers for spiking networks, over time-major inputs [T, ...].
Contents
Section titled “Contents”| Name | |
|---|---|
ALIF | Adaptive-threshold LIF (sparx.dynamics.ALIFCell). |
IF | Integrate-and-fire: LIF without leak. |
LI | A leaky integrator readout (sparx.dynamics.LICell); returns its membrane, [T, ...]. |
LIF | Leaky integrate-and-fire (sparx.dynamics.LIFCell) with time constant tau. |
PSN | The parallel spiking neuron over all T x T step pairs, for a fixed T. |
RATES | The collection spiking layers sow their per-example, per-neuron firing rates into. |
STATE | The collection a layer carries its neurons’ state in across apply calls. |
BatchMajor | Run a time-major layer (a sparx layer or a stack of them) on batch-major input [B, T, ...]. |
DecayingTrace | sparx.dynamics.DecayingHebb with a learned rate starting at eta, Miconi et al.’s 0.01. |
DelayedDense | A dense layer whose every synapse has a learnable delay of 0 to max_delay steps. |
Dynamics | Any neuron model of sparx.dynamics as a layer, its fields fixed. |
Flatten | Flatten each example’s trailing ndim axes into one feature axis, channels first. |
Flattens | 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. |
HebbianTrace | The Hebbian trace of a Recurrent layer, which builds its rule with the rule’s learned parameters. |
MaskedPSN | The PSN restricted to the k most recent steps, for a fixed T. |
Modelled | A layer that runs a sparx.dynamics model over time, stepped at dt. |
ModulatedTrace | sparx.dynamics.ModulatedHebb: a learned neuromodulator sets each unit’s rate, clipped at clip. |
Neuron | A population of neurons run over the leading (time) axis of its input. |
OjaTrace | sparx.dynamics.OjaHebb, Oja’s rule, with a learned rate starting at eta. |
Rate | A leaky rate unit (sparx.dynamics.RateCell), FLYNN’s neuron; returns its activity, [T, ...]. |
Recurrent | Feed neuron’s output back into its input through a learned [F, F] matrix, fixed or plastic. |
RetroactiveTrace | sparx.dynamics.RetroactiveHebb: a learned neuromodulator writes recent coactivity into the weights. |
SlidingPSN | k weights slid over time: H[t] = sum_i weight[i] * X[t - k + 1 + i] + bias. |
Synaptic | Current-based LIF, Serial(LICell, LIFCell): synaptic time constant tau_synapse, membrane time constant tau. |
adopt | child (a neuron, a Hebbian trace) as owner’s child called name, so its parameters sit under owner wherever it was built. |
band_mask | M[i, j] = 1 where j <= i <= j + k - 1: each step sees itself and the k - 1 before it. |
delay_kernel | Each synapse’s weights over its max_delay + 1 lags, [max_delay + 1, *delay.shape]. |
history_window | x [T, ...] with the held steps before it prepended, [held + T, ...], in x’s dtype. |
record_rates | Sow spikes’ time-averaged rate into "spike_rates" when that collection is mutable. |
class ALIF(Neuron)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.
| Field | Type | Default |
|---|---|---|
tau | float | 20.0 |
tau_adapt | float | 200.0 |
beta | float | 1.8 |
threshold | float | 1.0 |
reset | Reset | 'subtract' |
surrogate | Surrogate | ATan() |
detach_reset | bool | False |
learn_tau | bool | False |
refractory | float | 0 |
ALIF.build
Section titled “ALIF.build”def build(x: jax.Array) -> ALIFCellclass IF(Neuron)Integrate-and-fire: LIF without leak.
| Field | Type | Default |
|---|---|---|
threshold | float | 1.0 |
reset | Reset | 'subtract' |
surrogate | Surrogate | ATan() |
detach_reset | bool | False |
IF.build
Section titled “IF.build”def build(x: jax.Array) -> LIFCellclass LI(Neuron)A leaky integrator readout (sparx.dynamics.LICell); returns its membrane, [T, ...].
| Field | Type | Default |
|---|---|---|
tau | float | 2.0 |
learn_tau | bool | False |
LI.build
Section titled “LI.build”def build(x: jax.Array) -> LICellclass LIF(Neuron)Leaky integrate-and-fire (sparx.dynamics.LIFCell) with time constant tau.
learn_tau learns one decay per feature (last axis), starting at tau.
| Field | Type | Default |
|---|---|---|
tau | float | 2.0 |
threshold | float | 1.0 |
reset | Reset | 'subtract' |
surrogate | Surrogate | ATan() |
detach_reset | bool | False |
learn_tau | bool | False |
LIF.build
Section titled “LIF.build”def build(x: jax.Array) -> LIFCellclass PSN(nn.Module)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.
| Field | Type | Default |
|---|---|---|
surrogate | Surrogate | ATan() |
precision | jax.lax.Precision | None | None |
PSN.__call__
Section titled “PSN.__call__”def __call__(x: jax.Array) -> jax.ArrayRATES = 'spike_rates'The collection spiking layers sow their per-example, per-neuron firing rates into.
STATE = 'state'The collection a layer carries its neurons’ state in across apply calls.
BatchMajor
Section titled “BatchMajor”class BatchMajor(nn.Module)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.
| Field | Type | Default |
|---|---|---|
layer | nn.Module |
BatchMajor.__call__
Section titled “BatchMajor.__call__”def __call__(x: jax.Array) -> jax.ArrayDecayingTrace
Section titled “DecayingTrace”class DecayingTrace(HebbianTrace)sparx.dynamics.DecayingHebb with a learned rate starting at eta, Miconi et al.’s 0.01.
| Field | Type | Default |
|---|---|---|
eta | float | 0.01 |
DecayingTrace.build
Section titled “DecayingTrace.build”def build(features: int) -> DecayingHebbDelayedDense
Section titled “DelayedDense”class DelayedDense(nn.Module)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.
| Field | Type | Default |
|---|---|---|
features | int | |
max_delay | int | |
use_bias | bool | True |
kernel_init | nn.initializers.Initializer | nn.initializers.lecun_normal() |
dtype | Dtype | None | None |
param_dtype | Dtype | jnp.float32 |
precision | PrecisionLike | None |
DelayedDense.__call__
Section titled “DelayedDense.__call__”def __call__(x: jax.Array, sigma: float | jax.Array) -> jax.ArrayDynamics
Section titled “Dynamics”class Dynamics(Neuron)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).
| Field | Type | Default |
|---|---|---|
neuron | NeuronModel | dynamics.LeakyIntegrateAndFire() |
drive | Literal['current', 'jump'] | 'current' |
Dynamics.inputs
Section titled “Dynamics.inputs”def inputs(x: jax.Array) -> SynapticInputDynamics.build
Section titled “Dynamics.build”def build(x: jax.Array) -> NeuronModelFlatten
Section titled “Flatten”class Flatten(nn.Module)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.
| Field | Type | Default |
|---|---|---|
ndim | int | 3 |
Flatten.flattened_axes
Section titled “Flatten.flattened_axes”def flattened_axes() -> intndim, the trailing axes this layer flattens (Flattens).
Flatten.__call__
Section titled “Flatten.__call__”def __call__(x: jax.Array) -> jax.ArrayFlattens
Section titled “Flattens”class Flattens(Protocol)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.
Flattens.flattened_axes
Section titled “Flattens.flattened_axes”def flattened_axes() -> intFlattens.__call__
Section titled “Flattens.__call__”def __call__(x: jax.Array) -> jax.ArrayHebbianTrace
Section titled “HebbianTrace”class HebbianTrace(nn.Module)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.
HebbianTrace.build
Section titled “HebbianTrace.build”def build(features: int) -> HebbianRuleHebbianTrace.__call__
Section titled “HebbianTrace.__call__”def __call__(features: int) -> HebbianRuleMaskedPSN
Section titled “MaskedPSN”class MaskedPSN(nn.Module)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.
| Field | Type | Default |
|---|---|---|
k | int | |
surrogate | Surrogate | ATan() |
precision | jax.lax.Precision | None | None |
MaskedPSN.__call__
Section titled “MaskedPSN.__call__”def __call__(x: jax.Array, masking: float | jax.Array = 1.0) -> jax.ArrayModelled
Section titled “Modelled”class Modelled(Protocol)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.
| Field | Type | Default |
|---|---|---|
dt | float |
Modelled.model
Section titled “Modelled.model”def model(x: jax.Array) -> NeuronModelModulatedTrace
Section titled “ModulatedTrace”class ModulatedTrace(HebbianTrace)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.
| Field | Type | Default |
|---|---|---|
clip | float | 2.0 |
ModulatedTrace.build
Section titled “ModulatedTrace.build”def build(features: int) -> ModulatedHebbNeuron
Section titled “Neuron”class Neuron(nn.Module)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).
| Field | Type | Default |
|---|---|---|
dt | float | dataclasses.field(default=1.0, kw_only=True) |
unroll | int | dataclasses.field(default=1, kw_only=True) |
Neuron.build
Section titled “Neuron.build”def build(x: jax.Array) -> NeuronModelThe model for inputs like x, [T, ..., features].
Neuron.inputs
Section titled “Neuron.inputs”def inputs(x: jax.Array) -> SynapticInputWhat the model receives from the layer’s input x: a jump of its membrane.
Neuron.model
Section titled “Neuron.model”def model(x: jax.Array) -> NeuronModelThis layer’s model, with its parameters, for inputs like x.
Neuron.__call__
Section titled “Neuron.__call__”def __call__(x: jax.Array) -> jax.ArrayRun over x, [T, ...], and return the outputs [T, ...].
Neuron.run
Section titled “Neuron.run”def run(model: NeuronModel, x: jax.Array) -> jax.ArrayOjaTrace
Section titled “OjaTrace”class OjaTrace(HebbianTrace)sparx.dynamics.OjaHebb, Oja’s rule, with a learned rate starting at eta.
| Field | Type | Default |
|---|---|---|
eta | float | 0.01 |
OjaTrace.build
Section titled “OjaTrace.build”def build(features: int) -> OjaHebbclass Rate(Neuron)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=...)).
| Field | Type | Default |
|---|---|---|
tau | float | 2.0 |
activation | Literal['tanh', 'relu', 'sigmoid'] | 'tanh' |
learn_tau | bool | False |
Rate.build
Section titled “Rate.build”def build(x: jax.Array) -> RateCellRecurrent
Section titled “Recurrent”class Recurrent(Neuron)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.
| Field | Type | Default |
|---|---|---|
neuron | Neuron | LIF() |
kernel_init | nn.initializers.Initializer | nn.initializers.orthogonal() |
precision | PrecisionLike | None |
rule | HebbianTrace | None | None |
alpha_init | nn.initializers.Initializer | nn.initializers.normal(0.01) |
Recurrent.inputs
Section titled “Recurrent.inputs”def inputs(x: jax.Array) -> SynapticInputRecurrent.build
Section titled “Recurrent.build”def build(x: jax.Array) -> RecurrentCellRetroactiveTrace
Section titled “RetroactiveTrace”class RetroactiveTrace(HebbianTrace)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.
| Field | Type | Default |
|---|---|---|
eta | float | 0.01 |
clip | float | 1.0 |
RetroactiveTrace.build
Section titled “RetroactiveTrace.build”def build(features: int) -> RetroactiveHebbSlidingPSN
Section titled “SlidingPSN”class SlidingPSN(nn.Module)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.
| Field | Type | Default |
|---|---|---|
k | int | |
exponential_init | bool | True |
surrogate | Surrogate | ATan() |
SlidingPSN.__call__
Section titled “SlidingPSN.__call__”def __call__(x: jax.Array) -> jax.ArraySynaptic
Section titled “Synaptic”class Synaptic(Neuron)Current-based LIF, Serial(LICell, LIFCell): synaptic time constant tau_synapse,
membrane time constant tau.
| Field | Type | Default |
|---|---|---|
tau | float | 10.0 |
tau_synapse | float | 5.0 |
threshold | float | 1.0 |
reset | Reset | 'subtract' |
surrogate | Surrogate | ATan() |
detach_reset | bool | False |
learn_tau | bool | False |
Synaptic.build
Section titled “Synaptic.build”def build(x: jax.Array) -> Serialdef adopt(child: Child, owner: nn.Module, name: str) -> Childchild (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.
band_mask
Section titled “band_mask”def band_mask(steps: int, k: int) -> jax.ArrayM[i, j] = 1 where j <= i <= j + k - 1: each step sees itself and the k - 1 before it.
delay_kernel
Section titled “delay_kernel”def delay_kernel(delay: jax.Array, max_delay: int, sigma: float | jax.Array) -> jax.ArrayEach 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.
history_window
Section titled “history_window”def history_window(module: nn.Module, x: jax.Array, held: int) -> jax.Arrayx [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.
record_rates
Section titled “record_rates”def record_rates(module: nn.Module, spikes: jax.Array) -> NoneSow spikes’ time-averaged rate into "spike_rates" when that collection is mutable.