Skip to content
GitHub

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.

Name
ATanThe arctangent step’s derivative, alpha / 2 / (1 + (pi / 2 * alpha * x)^2).
FastSigmoidSuperSpike’s derivative, 1 / (slope * |x| + 1)^2.
GaussianThe normal density with standard deviation sigma, which integrates to 1.
RectangleA box of height 1 / width over |x| < width / 2, which integrates to 1.
SigmoidThe logistic step’s derivative, alpha * sigmoid(alpha x) * (1 - sigmoid(alpha x)).
StraightThroughThe identity’s derivative, 1 everywhere: the straight-through estimator.
SurrogateThe derivative that stands in for the Heaviside step’s in a backward pass.
TriangleA piecewise linear bump, scale * max(0, 1 - |x| / width).
spikeThe Heaviside step of x, 1 where x >= 0, differentiated through surrogate.
class ATan(Surrogate)

sparx.surrogate on GitHub

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.

FieldTypeDefault
alphafloat2.0
def derivative(x: jax.Array) -> jax.Array
class FastSigmoid(Surrogate)

sparx.surrogate on GitHub

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.

FieldTypeDefault
slopefloat25.0
def derivative(x: jax.Array) -> jax.Array
class Gaussian(Surrogate)

sparx.surrogate on GitHub

The normal density with standard deviation sigma, which integrates to 1.

The Gaussian window of Wu et al. (Frontiers 2018).

FieldTypeDefault
sigmafloat0.5
def derivative(x: jax.Array) -> jax.Array
class Rectangle(Surrogate)

sparx.surrogate on GitHub

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

FieldTypeDefault
widthfloat1.0
def derivative(x: jax.Array) -> jax.Array
class Sigmoid(Surrogate)

sparx.surrogate on GitHub

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.

FieldTypeDefault
alphafloat4.0
def derivative(x: jax.Array) -> jax.Array
class StraightThrough(Surrogate)

sparx.surrogate on GitHub

The identity’s derivative, 1 everywhere: the straight-through estimator.

def derivative(x: jax.Array) -> jax.Array
class Surrogate(ABC)

sparx.surrogate on GitHub

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

def derivative(x: jax.Array) -> jax.Array

The surrogate’s d spike / d x at x, in x’s dtype.

def __call__(x: jax.Array) -> jax.Array
class Triangle(Surrogate)

sparx.surrogate on GitHub

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.

FieldTypeDefault
widthfloat1.0
scalefloat1.0
def derivative(x: jax.Array) -> jax.Array
def spike(x: jax.Array, surrogate: Surrogate) -> jax.Array

sparx.surrogate on GitHub

The Heaviside step of x, 1 where x >= 0, differentiated through surrogate.