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.
Contents
Section titled “Contents”| Name | |
|---|---|
SEWBlock | A basic SEW block: two 3x3 conv-BN-neuron layers, joined to the shortcut by connect. |
SEWResNet | A SEW ResNet over [T, B, H, W, C] inputs, returning per-step logits [T, B, classes]. |
SpikingMLP | Dense spiking layers over [T, B, ...] with a leaky integrator readout [T, B, classes]. |
sew_resnet18 | SEW-ResNet-18: stages of (2, 2, 2, 2) basic blocks. |
sew_resnet34 | SEW-ResNet-34: stages of (3, 4, 6, 3) basic blocks. |
SEWBlock
Section titled “SEWBlock”class SEWBlock(nn.Module)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.
| Field | Type | Default |
|---|---|---|
features | int | |
strides | int | 1 |
connect | Connect | 'add' |
neuron | Neuron | LIF() |
dtype | Dtype | None | None |
param_dtype | Dtype | jnp.float32 |
precision | PrecisionLike | None |
SEWBlock.__call__
Section titled “SEWBlock.__call__”def __call__(x: jax.Array, train: bool) -> jax.ArraySEWResNet
Section titled “SEWResNet”class SEWResNet(nn.Module)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.
| Field | Type | Default |
|---|---|---|
stages | Sequence[int] | |
classes | int | |
width | int | 64 |
connect | Connect | 'add' |
stem | Literal['imagenet', 'small'] | 'imagenet' |
neuron | Neuron | LIF() |
dtype | Dtype | None | None |
param_dtype | Dtype | jnp.float32 |
precision | PrecisionLike | None |
SEWResNet.__call__
Section titled “SEWResNet.__call__”def __call__(x: jax.Array, train: bool) -> jax.ArraySpikingMLP
Section titled “SpikingMLP”class SpikingMLP(nn.Module)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.
| Field | Type | Default |
|---|---|---|
hidden | Sequence[int] | |
classes | int | |
neuron | Neuron | LIF() |
recurrent | bool | False |
delays | int | Sequence[int] | 0 |
extend | bool | False |
batch_norm | bool | False |
use_bias | bool | True |
weight_init | Literal['lecun_normal', 'kaiming_uniform'] | 'lecun_normal' |
dropout | float | 0.0 |
dropout_mask | Literal['step', 'sequence'] | 'step' |
readout_tau | float | 2.0 |
learn_readout_tau | bool | False |
dtype | Dtype | None | None |
param_dtype | Dtype | jnp.float32 |
precision | PrecisionLike | None |
SpikingMLP.max_delays
Section titled “SpikingMLP.max_delays”def max_delays() -> tuple[int, ...]Each synapse’s largest delay, hidden layers then the readout; 0 for a dense synapse.
SpikingMLP.__call__
Section titled “SpikingMLP.__call__”def __call__(x: jax.Array, train: bool = False, sigma: float | jax.Array = 0) -> jax.Arraysew_resnet18
Section titled “sew_resnet18”def sew_resnet18(classes: int, **kwargs={}) -> SEWResNetSEW-ResNet-18: stages of (2, 2, 2, 2) basic blocks.
sew_resnet34
Section titled “sew_resnet34”def sew_resnet34(classes: int, **kwargs={}) -> SEWResNetSEW-ResNet-34: stages of (3, 4, 6, 3) basic blocks.