Skip to content
GitHub

sparx.losses

Differentiable losses over a network’s time-major outputs [T, B, ...].

The classification losses take outputs [T, B, C] and integer labels; van_rossum compares spikes with target spike trains, the loss of fitting a network to recordings. The statistics that only measure spike trains after a run, without a gradient, are sparx.spiketrains.

The classification losses return one loss per example, [B], so a caller chooses the reduction (dew’s objectives sum it over a Ratio). For a loss on one readout of the outputs, reduce time first and use optax: the maximum membrane of a leaky integrator readout (jnp.max(v, axis=0), Cramer et al., IEEE TNNLS 2020), its mean, or the spike count (jnp.sum(spikes, axis=0)), each as logits for optax.softmax_cross_entropy_with_integer_labels. softmax_sum_cross_entropy is the readout of Hammouamri et al.’s SNN-delays, which takes the softmax at every step before summing over time.

Readout names those choices, and readout_logits and readout_losses apply one, so a classifier’s objective, its trained task and a recipe’s command line share one list.

Name
READOUTSEvery Readout, for a value read from a record or a command line.
ReadoutHow the outputs [T, B, C] are scored against the labels.
per_step_cross_entropyCross entropy of the outputs at every step against labels, averaged over time.
rate_mseMean squared error between each output neuron’s firing rate and its target rate.
readout_logitsThe class scores [B, C] a readout predicts from outputs [T, B, C], in float32.
readout_lossesEach example’s loss [B] under readout, for outputs [T, B, C] and integer labels.
softmax_sumThe class probabilities of every step summed over time, [T, B, C] -> [B, C], in float32.
softmax_sum_cross_entropyCross entropy of softmax_sum(outputs) taken as logits, SNN-delays’ loss='sum'.
van_rossumVan Rossum’s (Neural Computation 2001) distance between spike trains, squared; differentiable.
READOUTS: tuple[Readout, ...] = ('mean', 'max', 'sum', 'softmax_sum', 'per_step')

sparx.losses on GitHub

Every Readout, for a value read from a record or a command line.

Readout

sparx.losses on GitHub

How the outputs [T, B, C] are scored against the labels.

  • mean: cross entropy of the time-averaged outputs, a firing rate for spikes or the mean membrane of a leaky integrator readout.
  • max: cross entropy of each class’s largest output over time, the readout Cramer et al. (IEEE TNNLS 2020) use on a leaky integrator for SHD.
  • sum: cross entropy of the summed outputs, the spike count.
  • softmax_sum: the softmax of every step summed over time, scored as logits (softmax_sum_cross_entropy), SNN-delays’ loss='sum'.
  • per_step: cross entropy at every step, averaged (per_step_cross_entropy). Predictions use the mean.
def per_step_cross_entropy(outputs: jax.Array, labels: jax.Array) -> jax.Array

sparx.losses on GitHub

Cross entropy of the outputs at every step against labels, averaged over time.

The temporal term of Deng et al., “Temporal Efficient Training of Spiking Neural Network via Gradient Re-weighting” (ICLR 2022), and snnTorch’s ce_rate_loss. Each step is asked to classify alone, which they find generalizes better than the loss of the time-averaged output. Computed in float32.

def rate_mse(spikes: jax.Array, labels: jax.Array, correct: float = 0.8, incorrect: float = 0.2) -> jax.Array

sparx.losses on GitHub

Mean squared error between each output neuron’s firing rate and its target rate.

The labelled class is asked to fire at correct of the steps and every other class at incorrect, so no neuron is pushed to silence or to saturation. It is snnTorch’s mse_count_loss divided by T (which returns the squared error of spike counts over T), so it does not grow with T, where T * correct and T * incorrect are whole numbers (snnTorch rounds its target counts down).

def readout_logits(readout: Readout, outputs: jax.Array) -> jax.Array

sparx.losses on GitHub

The class scores [B, C] a readout predicts from outputs [T, B, C], in float32.

def readout_losses(readout: Readout, outputs: jax.Array, labels: jax.Array) -> jax.Array

sparx.losses on GitHub

Each example’s loss [B] under readout, for outputs [T, B, C] and integer labels.

def softmax_sum(outputs: jax.Array) -> jax.Array

sparx.losses on GitHub

The class probabilities of every step summed over time, [T, B, C] -> [B, C], in float32.

Each step votes with a distribution that sums to 1, so no single step’s large membrane outweighs the rest, as it does in a sum or maximum of the raw outputs. The argmax is SNN-delays’ prediction.

def softmax_sum_cross_entropy(outputs: jax.Array, labels: jax.Array) -> jax.Array

sparx.losses on GitHub

Cross entropy of softmax_sum(outputs) taken as logits, SNN-delays’ loss='sum'.

Their calc_loss passes the summed probabilities to torch’s CrossEntropyLoss, which applies a log-softmax to them again; this keeps that, so the loss is theirs. The summed probabilities lie in [0, T], so the loss cannot fall below log(1 + (C - 1) exp(-T)).

def van_rossum(spikes: jax.Array, target: jax.Array, tau: float, dt: float = 1.0) -> jax.Array

sparx.losses on GitHub

Van Rossum’s (Neural Computation 2001) distance between spike trains, squared; differentiable.

D^2 = (1 / tau) * integral of (f(t) - g(t))^2 dt, f = sum_k exp(-(t - t_k) / tau) H(t - t_k)

for time-major spike counts [T, ...] on a grid of dt, compared elementwise (one train per trailing index) and summed over trains. Between grid points the difference of the filtered trains decays exponentially, so the integral is exact for spikes on the grid, its tail after the last step included: h_n^2 (1 - exp(-2 dt / tau)) / 2 per step and h^2 / 2 for the tail, h the filtered difference. Elephant’s van_rossum_distance is sqrt(2) times this distance’s square root. Spikes through a surrogate train against recorded ones by it.