Skip to content
GitHub

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.losses and sparx.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, and run, which scans any of them over time.
  • sparx.graph: populations and projections wired into a Network, 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’s Trainer, with sparx.metrics (their accuracy), sparx.tasks (the trained classifier dew.pipeline loads) and sparx.config (SNNRunConfig, the run a recipe trains and run.json records).
  • 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.

Module
sparx.configThe run a spiking classifier trains as, one typed record that run.json holds.
sparx.datasetsNeuromorphic datasets as dense, binned spike counts.
sparx.dynamicsNeuron, synapse and plasticity models, and the runner that scans them over time.
sparx.encodeTurn a batch field into the time-major input [T, B, ...] of a spiking network.
sparx.graphCircuits and connectomes: populations of neurons joined by projections, simulated on one clock.
sparx.learnLearning rules beyond surrogate-gradient backpropagation through time (design.md section 7).
sparx.lossesDifferentiable losses over a network’s time-major outputs [T, B, ...].
sparx.metricsMetrics over a spiking objective’s evaluation, which dew’s Trainer.fit(metrics=...) takes.
sparx.modelsSpiking network architectures built from sparx.nn layers.
sparx.nirExchanging networks through NIR, the Neuromorphic Intermediate Representation (Pedersen et al. 2024).
sparx.nnFlax linen layers for spiking networks, over time-major inputs [T, ...].
sparx.objectivesSpiking networks as dew objectives, trained by dew’s Trainer.
sparx.ratesRead and regularize the firing rates spiking layers sow.
sparx.serveServing streaming spiking models: many sessions, each with its own neuron state, in one batched program.
sparx.spiketrainsStatistics of and distances between spike trains, for comparing networks that cannot match spike for spike.
sparx.surrogateThe spike nonlinearity and the gradients it trains with.
sparx.tasksTrained spiking networks as inference tasks, which dew.pipeline loads from a run.
Name
firing_ratesThe mean firing rate of each spiking layer, in spikes per step, keyed by layer path.
rate_penaltyThe squared distance of each neuron’s rate outside [lower, upper], averaged over neurons.
runScan model over time-major inputs; return its output [T, ...] and the final state.
spikeThe Heaviside step of x, 1 where x >= 0, differentiated through surrogate.
def firing_rates(sown: Sown) -> dict[str, jax.Array]

sparx.rates on GitHub

The mean firing rate of each spiking layer, in spikes per step, keyed by layer path.

def rate_penalty(sown: Sown, lower: float = 0.0, upper: float = 1.0, rows: jax.Array | None = None) -> jax.Array

sparx.rates on GitHub

The 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]

sparx.dynamics.core on GitHub

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.Array

sparx.surrogate on GitHub

The Heaviside step of x, 1 where x >= 0, differentiated through surrogate.