Skip to content
GitHub

sparx.rates

Read and regularize the firing rates spiking layers sow.

Applying a network with the "spike_rates" collection mutable collects every spiking layer’s rate per example and neuron, averaged over time:

outputs, sown = net.apply(params, spikes, mutable=["spike_rates"])
sown["spike_rates"] # {"LIF_0": {"rate": (rates [B, 128],)}, ...}

firing_rates summarizes them per layer for logging; rate_penalty keeps each neuron’s rate in a band as a differentiable loss term.

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