sparx
Sparx: spiking neural networks in JAX and Flax, trained, simulated and served through dew.
Networks are Flax linen modules over time-major spike trains [T, ...].
Deep spiking networks:
sparx.nn: Flax layers over the neuron models, parallel spiking neurons and delayed synapses.sparx.models: architectures built from them (SEWResNet,SpikingMLP).sparx.surrogate: the spike and its surrogate gradients.sparx.encode: the encoders that turn data into spike trains.sparx.lossesandsparx.rates: losses over time, firing-rate readouts and penalties.sparx.learn: rules beyond backpropagation through time (e-prop, OTTT, EventProp, conversion).
Circuits in physical units:
sparx.dynamics: neuron, synapse and plasticity models as pure JAX, the dimensionless family deep networks train with and the physical one, andrun, which scans any of them over time.sparx.graph: populations and projections wired into aNetwork,simulate, and connectomes.sparx.spiketrains: statistics of and distances between recorded spike trains.
Around them:
sparx.objectives: the objectives that train spiking networks under dew’sTrainer, withsparx.metrics(their accuracy),sparx.tasks(the trained classifierdew.pipelineloads) andsparx.config(SNNRunConfig, the run a recipe trains andrun.jsonrecords).sparx.datasets: spiking datasets as dew datasets (SHD).sparx.serve:StreamServer, many streaming sessions in one batch.sparx.nir: exchange through the Neuromorphic Intermediate Representation.
sparx.graph, sparx.learn, sparx.objectives, sparx.metrics,
sparx.tasks, sparx.config, sparx.datasets, sparx.serve and
sparx.nir load on first access (sparx.graph.Network after import sparx). The graph, the objectives and the datasets import dew’s trainer
and data stack, about 0.9 s on a 4-core CPU, which a script that only
trains a network in its own loop does not need.
Modules
Section titled “Modules”| Module | |
|---|---|
sparx.config | The run a spiking classifier trains as, one typed record that run.json holds. |
sparx.datasets | Neuromorphic datasets as dense, binned spike counts. |
sparx.dynamics | Neuron, synapse and plasticity models, and the runner that scans them over time. |
sparx.encode | Turn a batch field into the time-major input [T, B, ...] of a spiking network. |
sparx.graph | Circuits and connectomes: populations of neurons joined by projections, simulated on one clock. |
sparx.learn | Learning rules beyond surrogate-gradient backpropagation through time (design.md section 7). |
sparx.losses | Differentiable losses over a network’s time-major outputs [T, B, ...]. |
sparx.metrics | Metrics over a spiking objective’s evaluation, which dew’s Trainer.fit(metrics=...) takes. |
sparx.models | Spiking network architectures built from sparx.nn layers. |
sparx.nir | Exchanging networks through NIR, the Neuromorphic Intermediate Representation (Pedersen et al. 2024). |
sparx.nn | Flax linen layers for spiking networks, over time-major inputs [T, ...]. |
sparx.objectives | Spiking networks as dew objectives, trained by dew’s Trainer. |
sparx.rates | Read and regularize the firing rates spiking layers sow. |
sparx.serve | Serving streaming spiking models: many sessions, each with its own neuron state, in one batched program. |
sparx.spiketrains | Statistics of and distances between spike trains, for comparing networks that cannot match spike for spike. |
sparx.surrogate | The spike nonlinearity and the gradients it trains with. |
sparx.tasks | Trained spiking networks as inference tasks, which dew.pipeline loads from a run. |
Contents
Section titled “Contents”| Name | |
|---|---|
firing_rates | The mean firing rate of each spiking layer, in spikes per step, keyed by layer path. |
rate_penalty | The squared distance of each neuron’s rate outside [lower, upper], averaged over neurons. |
run | Scan model over time-major inputs; return its output [T, ...] and the final state. |
spike | The Heaviside step of x, 1 where x >= 0, differentiated through surrogate. |
firing_rates
Section titled “firing_rates”def firing_rates(sown: Sown) -> dict[str, jax.Array]The mean firing rate of each spiking layer, in spikes per step, keyed by layer path.
rate_penalty
Section titled “rate_penalty”def rate_penalty(sown: Sown, lower: float = 0.0, upper: float = 1.0, rows: jax.Array | None = None) -> jax.ArrayThe squared distance of each neuron’s rate outside [lower, upper], averaged over neurons.
A neuron’s rate is its time-averaged spikes averaged over the batch (the
leading axis of each sown array), in spikes per step as firing_rates
reports it, and so are lower and upper. Each example is weighed by
rows when given, [B], so a batch’s repeated rows can weigh nothing. Silent
neurons below lower receive gradient to fire and saturated ones above
upper to stop, the role of the activity regularizers of Zenke and
Vogels (Neural Computation 2021). Every neuron of every layer weighs the
same.
def run(model: Model[State, Inputs], inputs: Inputs | jax.Array | np.ndarray, state: State | None = None, dt: float = 1.0, record: Callable[[State], Record] | None = None, unroll: int | bool = 1) -> tuple[Output | tuple[Output, Record], State]Scan model over time-major inputs; return its output [T, ...] and the final state.
inputs is what model.step takes for one step, with every leaf
[T, ...] or a scalar held over time: a SynapticInput for a neuron
model, Arrivals for a PointNeuron. An array [T, ...] stands for
SynapticInput(jump=...), the dimensionless family’s input. The
population’s per-step shape and its dtype are those of the first input
with a time axis. state None starts the population at rest; a run over
the first k steps and one over the rest from its final state equal one
run over all. With record, each step’s output is paired with
record(state) after the step (the membrane voltage, say). unroll is
jax.lax.scan’s: how many steps one loop iteration holds.
def spike(x: jax.Array, surrogate: Surrogate) -> jax.ArrayThe Heaviside step of x, 1 where x >= 0, differentiated through surrogate.