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.
Contents
Section titled “Contents”| Name | |
|---|---|
firing_rates | The mean firing rate of each spiking layer, in spikes per step, keyed by layer path. |
rate_penalty | The squared distance of each neuron’s rate outside [lower, upper], averaged over neurons. |
firing_rates
Section titled “firing_rates”def firing_rates(sown: Sown) -> dict[str, jax.Array]The mean firing rate of each spiking layer, in spikes per step, keyed by layer path.
rate_penalty
Section titled “rate_penalty”def rate_penalty(sown: Sown, lower: float = 0.0, upper: float = 1.0, rows: jax.Array | None = None) -> jax.ArrayThe 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.