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.
Contents
Section titled “Contents”| Name | |
|---|---|
DeltaEncoder | Spike where a signal rises by at least threshold from the step before. |
DirectEncoder | The values themselves as the input at each of steps steps, as a broadcast. |
EventsEncoder | Data that already holds spikes or currents over time, on axis time_axis of each record. |
LatencyEncoder | One spike per value over steps steps, earlier for larger values: time-to-first-spike coding. |
RateEncoder | Bernoulli spikes that fire with probability x at each of steps steps, independently. |
SpikeEncoder | Turns one batch field [B, ...] into the network’s time-major input [T, B, ...]. |
DeltaEncoder
Section titled “DeltaEncoder”class DeltaEncoder(SpikeEncoder)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.
| Field | Type | Default |
|---|---|---|
threshold | float | |
off_spikes | bool | False |
time_axis | int | 0 |
DeltaEncoder.encode
Section titled “DeltaEncoder.encode”def encode(key: jax.Array, x: ArrayLike) -> jax.ArrayDirectEncoder
Section titled “DirectEncoder”class DirectEncoder(SpikeEncoder)The values themselves as the input at each of steps steps, as a broadcast.
| Field | Type | Default |
|---|---|---|
steps | int |
DirectEncoder.encode
Section titled “DirectEncoder.encode”def encode(key: jax.Array, x: ArrayLike) -> jax.ArrayEventsEncoder
Section titled “EventsEncoder”class EventsEncoder(SpikeEncoder)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.
| Field | Type | Default |
|---|---|---|
time_axis | int | 0 |
EventsEncoder.encode
Section titled “EventsEncoder.encode”def encode(key: jax.Array, x: ArrayLike) -> jax.ArrayLatencyEncoder
Section titled “LatencyEncoder”class LatencyEncoder(SpikeEncoder)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.
| Field | Type | Default |
|---|---|---|
steps | int | |
threshold | float | 0.01 |
LatencyEncoder.encode
Section titled “LatencyEncoder.encode”def encode(key: jax.Array, x: ArrayLike) -> jax.ArrayRateEncoder
Section titled “RateEncoder”class RateEncoder(SpikeEncoder)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.
| Field | Type | Default |
|---|---|---|
steps | int |
RateEncoder.encode
Section titled “RateEncoder.encode”def encode(key: jax.Array, x: ArrayLike) -> jax.ArraySpikeEncoder
Section titled “SpikeEncoder”class SpikeEncoder(ABC)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.
SpikeEncoder.__call__
Section titled “SpikeEncoder.__call__”def __call__(key: jax.Array, x: ArrayLike) -> jax.ArraySpikeEncoder.encode
Section titled “SpikeEncoder.encode”def encode(key: jax.Array, x: ArrayLike) -> jax.Array