Skip to content
GitHub

sparx.models

Spiking network architectures built from sparx.nn layers.

SEWResNet is the spike-element-wise residual network of Fang et al., “Deep Residual Learning in Spiking Neural Networks” (NeurIPS 2021). A plain spiking ResNet passes the sum of a block’s output and its shortcut through another neuron, which loses the identity map; a SEW block instead combines two spike trains elementwise, so a block whose residual branch is silent passes its input through unchanged. The layout follows SpikingJelly’s model.sew_resnet (commit c6cb8e46): a 7x7 stem with max pooling, four stages of basic blocks at 64, 128, 256 and 512 channels, global average pooling and a linear head, all over time-major [T, B, H, W, C] inputs.

SpikingMLP is the dense network for event data such as SHD: stacked dense or delayed synapses, optionally with batch norm, optionally recurrent spiking layers, and a leaky integrator readout. With every synapse delayed it is the network of Hammouamri et al.’s SNN-delays (ICLR 2024).

neuron is the template every neuron layer of a model copies, such as sparx.nn.LIF(tau=2.0, detach_reset=True). A run’s record holds it as dew records any class, {"class": "sparx.nn.neurons:LIF", "fields": {...}}, and rebuilds the model from it. Each copy (sparx.nn.adopt) belongs to the block that uses it, so its parameters (a learned time constant) are that block’s own.

Both models take train as dew’s objectives pass it, and a run names them by import path (--model sparx.models:SpikingMLP).

Their synapses follow dew’s precision fields, as dew’s models do: dtype is the dtype the convolutions, dense and delayed synapses and batch norms compute in (None infers it from the input and the parameters), param_dtype the dtype their parameters are stored in, and precision their matmuls’ precision. So a run’s --model.dtype bfloat16 reaches them. Neuron membranes and their learned time constants stay float32 whatever these are (dew.nn.precision.at_least_fp32), and spikes come out in the synapses’ dtype, which holds 0 and 1 exactly.

They declare their parameters’ logical axes to dew’s layout (dew.nn.sharding.logical_axes) in dew’s names: a synapse’s presynaptic side is embed and its neurons mlp, a convolution names its output channels embed, and the neurons’ own parameters (time constants, batch norms) name their width as their synapse does. So dew’s default rules split a synapse’s neurons over tensor, or over fsdp on a mesh without tensor parallelism, and its inputs over fsdp when the neurons took tensor. The recurrent matrix names only its postsynaptic side. Below Layout.min_shard elements a parameter stays whole whatever it declares. Each declaration names a module and its parameter, ("readout", "kernel"), since dew matches a declaration in every model in the process.

Name
SEWBlockA basic SEW block: two 3x3 conv-BN-neuron layers, joined to the shortcut by connect.
SEWResNetA SEW ResNet over [T, B, H, W, C] inputs, returning per-step logits [T, B, classes].
SpikingMLPDense spiking layers over [T, B, ...] with a leaky integrator readout [T, B, classes].
sew_resnet18SEW-ResNet-18: stages of (2, 2, 2, 2) basic blocks.
sew_resnet34SEW-ResNet-34: stages of (3, 4, 6, 3) basic blocks.
class SEWBlock(nn.Module)

sparx.models on GitHub

A basic SEW block: two 3x3 conv-BN-neuron layers, joined to the shortcut by connect.

With strides above 1 or a change of width, the shortcut is a 1x1 conv-BN-neuron that downsamples to match, as in SpikingJelly.

FieldTypeDefault
featuresint
stridesint1
connectConnect'add'
neuronNeuronLIF()
dtypeDtype | NoneNone
param_dtypeDtypejnp.float32
precisionPrecisionLikeNone
def __call__(x: jax.Array, train: bool) -> jax.Array
class SEWResNet(nn.Module)

sparx.models on GitHub

A SEW ResNet over [T, B, H, W, C] inputs, returning per-step logits [T, B, classes].

stages gives the number of blocks in each of the four stages, (2, 2, 2, 2) for SEW-ResNet-18. width scales every stage’s channels (64 in the paper), which small inputs and quick experiments can shrink. stem is the paper’s 7x7 stride-2 convolution and 3x3 max pooling for ImageNet; "small" is a single 3x3 stride-1 convolution, the usual stem for 32x32 inputs such as CIFAR-10. Average the logits over time for a rate readout, or pass them to sparx.losses.per_step_cross_entropy.

FieldTypeDefault
stagesSequence[int]
classesint
widthint64
connectConnect'add'
stemLiteral['imagenet', 'small']'imagenet'
neuronNeuronLIF()
dtypeDtype | NoneNone
param_dtypeDtypejnp.float32
precisionPrecisionLikeNone
def __call__(x: jax.Array, train: bool) -> jax.Array
class SpikingMLP(nn.Module)

sparx.models on GitHub

Dense spiking layers over [T, B, ...] with a leaky integrator readout [T, B, classes].

Each width in hidden is a synapse followed by a copy of neuron, and the readout is a synapse followed by an LI integrator. recurrent feeds each hidden layer’s spikes back to itself (sparx.nn.Recurrent). Every layer steps at the neuron’s dt, so readout_tau is in the unit of the neuron’s time constants. Trailing input axes are flattened.

delays makes synapses sparx.nn.DelayedDense, whose Gaussian width is the call’s sigma (0, the rounded delays, by default). An integer above 0 delays the first synapse by up to that many steps. A sequence gives every synapse’s largest delay, hidden layers then the readout, with 0 for a plain dense synapse: Hammouamri et al.’s SNN-delays delays them all. extend appends max_delay // 2 steps of zeros to each delayed synapse’s input, as SNN-delays pads it on the right, so spikes delayed past the input’s end still arrive and each delayed synapse lengthens the sequence by that much; it changes the output’s length, so a stream fed in chunks needs it off.

batch_norm normalizes each hidden synapse’s output over time and batch (nn.BatchNorm, torch’s momentum of 0.1) before the neuron, as SNN-delays does; the readout is not normalized. use_bias gives synapses a bias. weight_init is the weights’ initializer: flax’s lecun_normal, or torch’s kaiming_uniform_(nonlinearity="relu"), uniform within sqrt(6 / fan_in), SNN-delays’ choice.

dropout acts on hidden spikes in training. dropout_mask="step" draws a mask per step; "sequence" draws one per sequence and holds it over time, as SpikingJelly’s multi-step Dropout does, so a dropped neuron is silent for the whole recording.

FieldTypeDefault
hiddenSequence[int]
classesint
neuronNeuronLIF()
recurrentboolFalse
delaysint | Sequence[int]0
extendboolFalse
batch_normboolFalse
use_biasboolTrue
weight_initLiteral['lecun_normal', 'kaiming_uniform']'lecun_normal'
dropoutfloat0.0
dropout_maskLiteral['step', 'sequence']'step'
readout_taufloat2.0
learn_readout_tauboolFalse
dtypeDtype | NoneNone
param_dtypeDtypejnp.float32
precisionPrecisionLikeNone
def max_delays() -> tuple[int, ...]

Each synapse’s largest delay, hidden layers then the readout; 0 for a dense synapse.

def __call__(x: jax.Array, train: bool = False, sigma: float | jax.Array = 0) -> jax.Array
def sew_resnet18(classes: int, **kwargs={}) -> SEWResNet

sparx.models on GitHub

SEW-ResNet-18: stages of (2, 2, 2, 2) basic blocks.

def sew_resnet34(classes: int, **kwargs={}) -> SEWResNet

sparx.models on GitHub

SEW-ResNet-34: stages of (3, 4, 6, 3) basic blocks.