sparx.surrogate
The spike nonlinearity and the gradients it trains with.
A spike is the Heaviside step of x = v - threshold: 1 where the membrane
reaches the threshold, 0 below it. Its true derivative is zero almost
everywhere, so training replaces it with the derivative of a smooth step, the
surrogate, while the forward pass stays exactly binary.
spike is a jax.custom_jvp, so the surrogate serves forward mode
(jax.jvp, forward gradients) and reverse mode (jax.grad, which JAX gets by
transposing the linear tangent rule) from one definition, and it batches under
vmap. A surrogate is a frozen dataclass: hashable, so it is a static
argument of the rule and a field of a Flax module.
Shapes follow the input; the spike has the input’s dtype, which holds 0 and 1 exactly in every float format.
Contents
Section titled “Contents”| Name | |
|---|---|
ATan | The arctangent step’s derivative, alpha / 2 / (1 + (pi / 2 * alpha * x)^2). |
FastSigmoid | SuperSpike’s derivative, 1 / (slope * |x| + 1)^2. |
Gaussian | The normal density with standard deviation sigma, which integrates to 1. |
Rectangle | A box of height 1 / width over |x| < width / 2, which integrates to 1. |
Sigmoid | The logistic step’s derivative, alpha * sigmoid(alpha x) * (1 - sigmoid(alpha x)). |
StraightThrough | The identity’s derivative, 1 everywhere: the straight-through estimator. |
Surrogate | The derivative that stands in for the Heaviside step’s in a backward pass. |
Triangle | A piecewise linear bump, scale * max(0, 1 - |x| / width). |
spike | The Heaviside step of x, 1 where x >= 0, differentiated through surrogate. |
class ATan(Surrogate)The arctangent step’s derivative, alpha / 2 / (1 + (pi / 2 * alpha * x)^2).
Fang et al., “Incorporating Learnable Membrane Time Constant to Enhance
Learning of Spiking Neural Networks” (ICCV 2021), and SpikingJelly’s
surrogate.ATan and snnTorch’s atan with the same alpha. It
integrates to 1 and its tails fall as 1 / x^2, so a neuron far from
threshold still receives gradient.
| Field | Type | Default |
|---|---|---|
alpha | float | 2.0 |
ATan.derivative
Section titled “ATan.derivative”def derivative(x: jax.Array) -> jax.ArrayFastSigmoid
Section titled “FastSigmoid”class FastSigmoid(Surrogate)SuperSpike’s derivative, 1 / (slope * |x| + 1)^2.
Zenke and Ganguli, “SuperSpike: Supervised Learning in Multilayer Spiking
Neural Networks” (Neural Computation 2018), and snnTorch’s fast_sigmoid
(default slope=25); Zenke’s SHD tutorials use slope=100. Its peak is
1 at threshold whatever the slope, so it integrates to 2 / slope, not 1.
| Field | Type | Default |
|---|---|---|
slope | float | 25.0 |
FastSigmoid.derivative
Section titled “FastSigmoid.derivative”def derivative(x: jax.Array) -> jax.ArrayGaussian
Section titled “Gaussian”class Gaussian(Surrogate)The normal density with standard deviation sigma, which integrates to 1.
The Gaussian window of Wu et al. (Frontiers 2018).
| Field | Type | Default |
|---|---|---|
sigma | float | 0.5 |
Gaussian.derivative
Section titled “Gaussian.derivative”def derivative(x: jax.Array) -> jax.ArrayRectangle
Section titled “Rectangle”class Rectangle(Surrogate)A box of height 1 / width over |x| < width / 2, which integrates to 1.
The rectangular window of Wu et al., “Spatio-Temporal Backpropagation for Training High-performance Spiking Neural Networks” (Frontiers 2018).
| Field | Type | Default |
|---|---|---|
width | float | 1.0 |
Rectangle.derivative
Section titled “Rectangle.derivative”def derivative(x: jax.Array) -> jax.ArraySigmoid
Section titled “Sigmoid”class Sigmoid(Surrogate)The logistic step’s derivative, alpha * sigmoid(alpha x) * (1 - sigmoid(alpha x)).
SpikingJelly’s surrogate.Sigmoid (default alpha=4); snnTorch’s
sigmoid names alpha its slope (default 25). It integrates to 1.
| Field | Type | Default |
|---|---|---|
alpha | float | 4.0 |
Sigmoid.derivative
Section titled “Sigmoid.derivative”def derivative(x: jax.Array) -> jax.ArrayStraightThrough
Section titled “StraightThrough”class StraightThrough(Surrogate)The identity’s derivative, 1 everywhere: the straight-through estimator.
StraightThrough.derivative
Section titled “StraightThrough.derivative”def derivative(x: jax.Array) -> jax.ArraySurrogate
Section titled “Surrogate”class Surrogate(ABC)The derivative that stands in for the Heaviside step’s in a backward pass.
derivative(x) is evaluated at x = v - threshold. Calling a surrogate
spikes: ATan()(x) is spike(x, ATan()).
Surrogate.derivative
Section titled “Surrogate.derivative”def derivative(x: jax.Array) -> jax.ArrayThe surrogate’s d spike / d x at x, in x’s dtype.
Surrogate.__call__
Section titled “Surrogate.__call__”def __call__(x: jax.Array) -> jax.ArrayTriangle
Section titled “Triangle”class Triangle(Surrogate)A piecewise linear bump, scale * max(0, 1 - |x| / width).
Bellec et al., “Long short-term memory and learning-to-learn in networks
of spiking neurons” (NeurIPS 2018), use scale=0.3 on a membrane already
divided by its threshold. Zero beyond width, so neurons far from
threshold pass no gradient.
| Field | Type | Default |
|---|---|---|
width | float | 1.0 |
scale | float | 1.0 |
Triangle.derivative
Section titled “Triangle.derivative”def derivative(x: jax.Array) -> jax.Arraydef spike(x: jax.Array, surrogate: Surrogate) -> jax.ArrayThe Heaviside step of x, 1 where x >= 0, differentiated through surrogate.