Skip to content
GitHub

sparx.encode

Turn a batch field into the time-major input [T, B, ...] of a spiking network.

An encoder is a frozen dataclass called as encoder(key, x) on a batch field [B, ...]. The same object encodes in a plain JAX loop and inside sparx.objectives.SpikingClassifierObjective, and a run’s record holds it as dew records any class, {"class": "sparx.encode:RateEncoder", "fields": {"steps": 8}}, which rebuilds it in another process. Each has a short name in sparx.registry.spike_encoders (rate), for the recipe’s command line.

Encoders of static data (DirectEncoder, RateEncoder, LatencyEncoder) add a leading time axis of steps: an image batch [B, H, W, C] becomes [steps, B, H, W, C]. Encoders of data that already runs over time (DeltaEncoder, EventsEncoder) move each record’s time axis to the front: [B, T, F] becomes [T, B, F].

The encoders that read values as intensities (DirectEncoder, RateEncoder, LatencyEncoder, DeltaEncoder) read a uint8 field as x / 255 and expect anything else in [0, 1], so raw image bytes and normalized arrays encode alike. EventsEncoder reads spike counts or currents, which it passes on unscaled. Every encoder returns float32.

Direct encoding feeds the analog values as the input at every step and lets the first layer do the encoding, as DIET-SNN (Rathi and Roy, IEEE TNNLS 2021) and SpikingJelly’s static-image examples do.

Name
DeltaEncoderSpike where a signal rises by at least threshold from the step before.
DirectEncoderThe values themselves as the input at each of steps steps, as a broadcast.
EventsEncoderData that already holds spikes or currents over time, on axis time_axis of each record.
LatencyEncoderOne spike per value over steps steps, earlier for larger values: time-to-first-spike coding.
RateEncoderBernoulli spikes that fire with probability x at each of steps steps, independently.
SpikeEncoderTurns one batch field [B, ...] into the network’s time-major input [T, B, ...].
class DeltaEncoder(SpikeEncoder)

sparx.encode on GitHub

Spike where a signal rises by at least threshold from the step before.

Each record holds the signal over time on axis time_axis. The step before the first is zero, so a signal that starts at or above threshold fires at step 0. With off_spikes, a fall of at least threshold emits -1. This is snnTorch’s spikegen.delta without padding.

FieldTypeDefault
thresholdfloat
off_spikesboolFalse
time_axisint0
def encode(key: jax.Array, x: ArrayLike) -> jax.Array
class DirectEncoder(SpikeEncoder)

sparx.encode on GitHub

The values themselves as the input at each of steps steps, as a broadcast.

FieldTypeDefault
stepsint
def encode(key: jax.Array, x: ArrayLike) -> jax.Array
class EventsEncoder(SpikeEncoder)

sparx.encode on GitHub

Data that already holds spikes or currents over time, on axis time_axis of each record.

A record [T, F] arrives batched as [B, T, F]; the default moves its time axis to the front. The values pass unscaled, as float32: a uint8 field here counts spikes.

FieldTypeDefault
time_axisint0
def encode(key: jax.Array, x: ArrayLike) -> jax.Array
class LatencyEncoder(SpikeEncoder)

sparx.encode on GitHub

One spike per value over steps steps, earlier for larger values: time-to-first-spike coding.

A value x in [0, 1] fires once, at step round((1 - x) * (steps - 1)), so 1 fires at the first step and threshold near the last; values below threshold never fire. This is snnTorch’s spikegen.latency with linear=True, normalize=True, clip=True.

FieldTypeDefault
stepsint
thresholdfloat0.01
def encode(key: jax.Array, x: ArrayLike) -> jax.Array
class RateEncoder(SpikeEncoder)

sparx.encode on GitHub

Bernoulli spikes that fire with probability x at each of steps steps, independently.

x is clipped to [0, 1]. The spike count over steps is binomial with mean steps * x. Each key gives a fresh draw, so a training loop that passes its step’s key encodes every step anew. Sampling has no gradient with respect to x.

FieldTypeDefault
stepsint
def encode(key: jax.Array, x: ArrayLike) -> jax.Array
class SpikeEncoder(ABC)

sparx.encode on GitHub

Turns one batch field [B, ...] into the network’s time-major input [T, B, ...].

key drives the random encoders (RateEncoder); the others ignore it, so every encoder is called the same way. It is a JAX PRNG key, as jax.random’s functions take, never an int seed: an encoder runs inside jitted training steps, where the caller splits one key per step. The entry points called from outside JAX (dew’s Trainer, SpikingClassification.logits) take an int seed, as dew’s do, and make the key. A subclass implements encode.

def __call__(key: jax.Array, x: ArrayLike) -> jax.Array
def encode(key: jax.Array, x: ArrayLike) -> jax.Array