Skip to content
GitHub

sparx.dynamics

Neuron, synapse and plasticity models, and the runner that scans them over time.

Every neuron model meets one contract (sparx.dynamics.core): init_state(shape, dtype) and step(state, SynapticInput, dt) -> (state, Output), and run(model, inputs) scans one over time. A model’s output is its spikes, or a graded value each step when the model is graded. Two families meet it. sparx.dynamics.ml holds the dimensionless, one-step-is-one-unit models deep spiking networks train with: soft resets, detached resets, learnable decays, input as a jump of the membrane. sparx.dynamics.neurons holds models of neurons as biology measures them: membrane equations in mV and ms with conductances, refractoriness and reversal potentials. Synapses with receptor kinetics and plasticity rules complete them; sparx.nn builds layers from these models and sparx.graph builds circuits and connectomes (design.md sections 4 and 5).

This package exports the models and the contract. The arithmetic the models share (fire, exact_linear, rk4, substeps and the rest), which a new model is written with, stays in sparx.dynamics.core.

Name
ACTIVATIONSThe activations a RateCell takes, by name.
IZHIKEVICH_2003(a, b, c, d) of each class in Izhikevich (2003), Figure 2.
IZHIKEVICH_2004The fields of each of Izhikevich’s (2004) twenty firing patterns, Figure 1 (A) to (T), as his code sets them. Each pattern also has its own input protocol and step, which tests/test_simulators.py reads from his code’s run.
RECEPTORSReversal potentials (mV) of the common receptors, the defaults models read conductances against: AMPA and NMDA 0 mV, GABA-A -80 mV (Cl-), GABA-B -95 mV (K+), as in Brette et al.’s simulator benchmarks (2007) and the cortical models built on them.
ALIFCellLIF with an adaptive threshold that rises with each spike and decays back.
ALIFState
AdExThe adaptive exponential integrate-and-fire neuron (Brette and Gerstner, J. Neurophysiol. 2005).
AdExState
Alphaw e / tau * s * exp(-s / tau), peaking at the weight w at s = tau: NEST’s iaf_psc_alpha.
ArrivalsOne step’s input to a PointNeuron: a held current (pA), the weighted spikes arriving at the end of the step on each receptor, by name, and the noise a stochastic model fires by (SynapticInput.noise).
BernoulliCellA leaky integrate-and-fire neuron with escape noise: it fires with a probability of its membrane.
BernoulliState
BiExponentialw f (exp(-s / tau_decay) - exp(-s / tau_rise)), normalized by f to peak at the weight w.
DecayingHebbA Hebbian trace that decays toward the latest coactivity at the rate eta.
DeltaA voltage jump of the arriving weight (mV), with no kinetics.
DenseEvery unit to every unit through weight, [F, F], weight[i, j] from unit i to unit j.
DopamineSTDPReward-modulated STDP (Izhikevich 2007): NEST’s stdp_dopamine_synapse (Potjans et al. 2010).
DopamineTraces
EligibleHebb
ExponentialJumps by the arriving weight and decays with tau (ms): NEST’s iaf_psc_exp and iaf_cond_exp.
FastWeightsFast weights on a wiring: alpha times the Hebbian trace that rule keeps, added to each weight.
GapGap-junction coupling onto each neuron over a step: the current drive(s) - conductance * v(s) (pA).
GradedTransmission proportional to a graded presynaptic value, with first-order kinetics.
GradedPotentialA neuron that never spikes and releases transmitter as a sigmoidal function of its voltage.
GradedPotentialState
GradedState
HebbianRuleHow a recurrent layer’s Hebbian trace changes over a step, on any Wiring.
HodgkinHuxleyThe squid giant axon (Hodgkin and Huxley, J. Physiol. 1952), as NEST’s hh_psc_alpha states it.
HodgkinHuxleyState
IzhikevichIzhikevich’s simple model (IEEE Trans. Neural Netw. 2003), in ms and mV.
IzhikevichState
LICellA leaky integrator that never fires; its output is its membrane, a graded value.
LIFCellLeaky integrate-and-fire. decay=1 integrates without leak (IF).
LandingWhere a synapse’s arrivals reach the neuron. synapse: into the synapse’s state at the end of the step, whose output shapes the membrane from the next step on. before_threshold: as a voltage jump at the end of the step, before the threshold test (NEST’s delta synapse). after_threshold: as a voltage jump after the test and before the reset (Brian2’s on_pre="v += w").
LeakyIntegrateAndFireLeaky integrate-and-fire with current and conductance input.
LeakyIntegrateAndFireState
MembraneState
MgBlockThe fraction of NMDA receptors not blocked by magnesium at voltage v (mV).
ModelAnything run steps: a neuron model, or a neuron with the synapses onto it (PointNeuron).
ModulatedHebbA Hebbian trace whose rate the network’s own activity sets through a neuromodulator, clipped.
NeuronModelOne population of neurons, advanced one step at a time on the synaptic input it receives.
OjaHebbOja’s rule: a Hebbian trace that each postsynaptic unit’s own activity bounds.
OutputWhat a population sends in a step, and when within the step.
PairSTDPAll-to-all pair STDP with soft or hard bounds: NEST’s stdp_synapse (Guetig et al. 2003).
PlasticityA rule that changes the weights of a projection’s edges with the spikes on either side.
PointNeuronA neuron model and the synapses onto it, by receptor, stepped together.
PointNeuronStateA neuron’s state and its synapses’ by receptor name.
PulseCellRNeuralNet’s neuron: each step it outputs a function of the sum of what arrived, and empties the sum.
RateCellA leaky rate unit; its output is its activity h, a graded value.
RateState
ReceptorA synapse model and how its output reaches the membrane.
RecurrentCellFeed a model’s output back to its own input through wiring, fixed or plastic.
RecurrentState
ResetWhat a spike does to a dimensionless membrane: subtract the threshold (soft reset, which keeps the overshoot), set it to zero (hard reset), or leave it, none.
RetroactiveHebbA Hebbian trace that a neuromodulator writes from an eligibility trace of recent coactivity.
STDPTraces
SerialTwo models in series: each step, the first’s Output.value is the second’s input jump.
SparseUnits wired along the edges pre[e] -> post[e] with weight[e], among size units.
StochasticReleaseRelease that succeeds at random: each spike, each synapse releases with probability p.
SynapseModel
SynapticInputWhat a population receives in one step.
TermA current (amplitude + slope * s) * exp(-s / tau) over a step, s in [0, dt] ms from its start.
TripletSTDPAll-to-all triplet STDP (Pfister and Gerstner 2006): NEST’s stdp_triplet_synapse.
TripletTraces
TsodyksMarkramShort-term depression and facilitation (Tsodyks and Markram 1997; Markram et al. 1998).
TsodyksMarkramState
WiringWhich units of a recurrent layer reach which, and with what weights: Dense or Sparse.
decayexp(-dt / tau): what a time constant tau leaves of a value after a step dt in the same unit.
izhikevich_2003The cortical and thalamic classes of Izhikevich (2003), Figure 2, by their names there.
izhikevich_2004The neuron of one of the twenty firing patterns of Izhikevich (2004), Figure 1, by its name (IZHIKEVICH_2004), integrated as the paper’s code does (scheme="semi_implicit").
runScan model over time-major inputs; return its output [T, ...] and the final state.
ACTIVATIONS = {'tanh': jnp.tanh, 'relu': jax.nn.relu, 'sigmoid': jax.nn.sigmoid}

sparx.dynamics.ml on GitHub

The activations a RateCell takes, by name.

IZHIKEVICH_2003: Mapping[str, tuple[float, float, float, float]]

sparx.dynamics.neurons on GitHub

(a, b, c, d) of each class in Izhikevich (2003), Figure 2.

IZHIKEVICH_2004: Mapping[str, Mapping[str, float | tuple[float, float, float]]]

sparx.dynamics.neurons on GitHub

The fields of each of Izhikevich’s (2004) twenty firing patterns, Figure 1 (A) to (T), as his code sets them. Each pattern also has its own input protocol and step, which tests/test_simulators.py reads from his code’s run.

RECEPTORS: Mapping[str, float] = {'ampa': 0.0, 'nmda': 0.0, 'gaba_a': -80.0, 'gaba_b': -95.0}

sparx.dynamics.neurons on GitHub

Reversal potentials (mV) of the common receptors, the defaults models read conductances against: AMPA and NMDA 0 mV, GABA-A -80 mV (Cl-), GABA-B -95 mV (K+), as in Brette et al.’s simulator benchmarks (2007) and the cortical models built on them.

class ALIFCell

sparx.dynamics.ml on GitHub

LIF with an adaptive threshold that rises with each spike and decays back.

theta[t] = threshold + beta * a[t-1]
s[t] = H(v[t] - theta[t])
v[t] <- v[t] - s[t] * threshold (a soft reset by the baseline)
a[t] = adapt_decay * a[t-1] + s[t]

The adaptive neurons of Bellec et al., “A solution to the learning dilemma for recurrent networks of spiking neurons” (Nature Communications 2020), whose adaptation time constants of hundreds of steps give a recurrent network memory beyond its membranes’. As in their equations and code, a spike subtracts the baseline threshold, not the adaptive one. Their reset lands one step later, undecayed; the reset here lands at the spike, like every model of this family, which makes this their model with a reset of decay * threshold (docs/fidelity.md). Gradients flow through a.

refractory is their n_refractory as a duration in the unit of dt, round(refractory / dt) steps: a spike and the silence after it span that many steps, during which the membrane integrates but cannot fire, and no gradient passes the spike (they use 2 to 5 at 1 ms a step). 0 and 1 step leave the neuron free to fire on the next step.

FieldTypeDefault
decayjax.Array | float
adapt_decayjax.Array | float
betajax.Array | float1.8
thresholdjax.Array | float1.0
resetResetstruct.field(pytree_node=False, default='subtract')
surrogateSurrogatestruct.field(pytree_node=False, default=ATan())
detach_resetboolstruct.field(pytree_node=False, default=False)
refractoryfloatstruct.field(pytree_node=False, default=0)
def init_state(shape: tuple[int, ...], dtype: jnp.dtype) -> ALIFState
def step(state: ALIFState, inputs: SynapticInput, dt: float) -> tuple[ALIFState, Output]
def is_refractory(state: ALIFState, dt: float) -> jax.Array
def after_threshold(state: ALIFState, jump: jax.Array, fired: jax.Array) -> ALIFState
class ALIFState(NamedTuple)

sparx.dynamics.ml on GitHub

FieldTypeDefault
vjax.Array
ajax.Array
rjax.Array
class AdEx

sparx.dynamics.neurons on GitHub

The adaptive exponential integrate-and-fire neuron (Brette and Gerstner, J. Neurophysiol. 2005).

C dv/dt = -g_L (v - E_L) + g_L D_T exp((v - v_T) / D_T) - w + I + sum_k g_k (E_k - v)
tau_w dw/dt = a (v - E_L) - w

At v >= v_peak it fires, v is set to v_reset and w grows by b; v then holds for t_ref (0 by default). The right-hand side reads the voltage capped at v_peak, and reads v_reset while refractory, as NEST’s aeif_* models do, which keeps the exponential finite through the upswing. Integrated by RK4 in substeps of at most substep ms, with synaptic currents evaluated at each stage and conductances held; a spike is detected and reset at the substep that crosses, and the rest of the step continues from the reset, as NEST’s adaptive integration does.

The upswing makes the equation stiff. Against NEST’s adaptive RK45 over 500 ms of Naud et al.’s patterns, spike times drift by up to 5.3 ms with substeps of 0.1 ms, 0.4 ms with 0.01 ms (the default) and 0.1 ms with 0.001 ms (tests/test_simulators.py); an exponential Rosenbrock step was less accurate than RK4 at every length tried. Lengthen substep (None: one per step) to trade that accuracy for speed in large networks. The defaults are NEST’s, the parameters of Brette and Gerstner’s Figure 2.

FieldTypeDefault
c_mjax.Array | float281.0
g_ljax.Array | float30.0
e_ljax.Array | float-70.6
v_tjax.Array | float-50.4
delta_tjax.Array | float2.0
v_peakjax.Array | float0.0
v_resetjax.Array | float-60.0
ajax.Array | float4.0
bjax.Array | float80.5
tau_wjax.Array | float144.0
t_reffloatstruct.field(pytree_node=False, default=0.0)
substepfloat | Nonestruct.field(pytree_node=False, default=0.01)
reversalMapping[str, float]
gatesMapping[str, MgBlock]
surrogateSurrogatestruct.field(pytree_node=False, default=ATan())
def init_state(shape: tuple[int, ...], dtype: jnp.dtype) -> AdExState
def step(state: AdExState, inputs: SynapticInput, dt: float) -> tuple[AdExState, Output]
def is_refractory(state: AdExState, dt: float) -> jax.Array
def after_threshold(state: AdExState, jump: jax.Array, fired: jax.Array) -> AdExState
class AdExState(NamedTuple)

sparx.dynamics.neurons on GitHub

FieldTypeDefault
vjax.Array
wjax.Array
refractoryjax.Array
class Alpha

sparx.dynamics.synapses on GitHub

w e / tau * s * exp(-s / tau), peaking at the weight w at s = tau: NEST’s iaf_psc_alpha.

The state is the waveform (value + slope * s) exp(-s / tau) from the start of the step, which a step carries forward exactly.

FieldTypeDefault
taujax.Array | float2.0
landsLanding
def init_state(shape: tuple[int, ...], dtype: jnp.dtype) -> AlphaState
def output(state: AlphaState) -> tuple[Term, ...]
def step(state: AlphaState, arriving: jax.Array | float, dt: float) -> AlphaState
class Arrivals(NamedTuple)

sparx.dynamics.synapses on GitHub

One step’s input to a PointNeuron: a held current (pA), the weighted spikes arriving at the end of the step on each receptor, by name, and the noise a stochastic model fires by (SynapticInput.noise).

FieldTypeDefault
currentjax.Array | float0.0
spikesMapping[str, jax.Array]{}
noisejax.Array | NoneNone
class BernoulliCell

sparx.dynamics.ml on GitHub

A leaky integrate-and-fire neuron with escape noise: it fires with a probability of its membrane.

v[t] = decay ** dt * v[t-1] + x[t]
p[t] = sigmoid(beta * (v[t] - threshold))
s[t] = 1 where noise[t] < p[t], else 0
v[t] <- reset(v[t], s[t])

noise[t] is uniform on [0, 1), one draw per neuron per step, which the caller passes as SynapticInput.noise (jax.random.uniform of a key), so the cell is a pure function of its inputs and a given noise replays one trajectory; noise 1 - s forces the spikes s, since p lies strictly between 0 and 1. p is a probability per step, the discrete-time escape noise of Pfister et al. (Neural Computation 2006), whatever dt is; beta is the inverse of the noise’s temperature, and as it grows the cell approaches LIFCell. The spike is a sample and passes no gradient: sparx.learn.reinforce trains a layer of these cells from the probability of the spikes they drew.

FieldTypeDefault
decayjax.Array | float
thresholdjax.Array | float1.0
betajax.Array | float1.0
resetResetstruct.field(pytree_node=False, default='subtract')
def init_state(shape: tuple[int, ...], dtype: jnp.dtype) -> BernoulliState
def step(state: BernoulliState, inputs: SynapticInput, dt: float) -> tuple[BernoulliState, Output]
def is_refractory(state: BernoulliState, dt: float) -> jax.Array
def after_threshold(state: BernoulliState, jump: jax.Array, fired: jax.Array) -> BernoulliState
class BernoulliState(NamedTuple)

sparx.dynamics.ml on GitHub

FieldTypeDefault
vjax.Array
pjax.Array
class BiExponential

sparx.dynamics.synapses on GitHub

w f (exp(-s / tau_decay) - exp(-s / tau_rise)), normalized by f to peak at the weight w.

NEST’s iaf_cond_beta and the common dual-exponential AMPA and NMDA kinetics. Needs tau_rise < tau_decay; equal time constants are Alpha.

FieldTypeDefault
tau_risefloatstruct.field(pytree_node=False, default=0.5)
tau_decayfloatstruct.field(pytree_node=False, default=5.0)
landsLanding
peak_factorfloat
def init_state(shape: tuple[int, ...], dtype: jnp.dtype) -> BiExponentialState
def output(state: BiExponentialState) -> tuple[Term, ...]
def step(state: BiExponentialState, arriving: jax.Array | float, dt: float) -> BiExponentialState
class DecayingHebb

sparx.dynamics.ml on GitHub

A Hebbian trace that decays toward the latest coactivity at the rate eta.

keep = (1 - eta) ** dt
hebb[i, j] <- keep hebb[i, j] + (1 - keep) pre[i] post[j]

Differentiable plasticity’s trace (Miconi et al. 2018, eq. 2), at dt = 1 the hebb = (1 - eta) * hebb + eta * outer(yin, yout) of their simple/simple.py, with one eta for every connection, learned with the weights. A step of dt keeps what dt unit steps would of the old trace, so eta is a rate per unit of time and lies in [0, 1).

FieldTypeDefault
etajax.Array | float
def init_trace(wiring: Wiring, shape: tuple[int, ...], dtype: jnp.dtype) -> jax.Array
def hebb(trace: jax.Array) -> jax.Array
def update(trace: jax.Array, wiring: Wiring, pre: jax.Array, post: jax.Array, dt: float) -> jax.Array
class Delta

sparx.dynamics.synapses on GitHub

A voltage jump of the arriving weight (mV), with no kinetics.

By default the jump lands before the step’s threshold test, as NEST’s iaf_psc_delta input does, so a jump due at the end of a step can fire the neuron in that step. after_threshold=True lands it after the test and before the reset, as Brian2’s on_pre="v += w" does: it takes effect from the next step, decaying over it first, and is lost if the neuron fired.

FieldTypeDefault
after_thresholdboolstruct.field(pytree_node=False, default=False)
landsLanding
def init_state(shape: tuple[int, ...], dtype: jnp.dtype) -> tuple[()]
def output(state: tuple[()]) -> tuple[Term, ...]
def step(state: tuple[()], arriving: jax.Array | float, dt: float) -> tuple[()]
class Dense

sparx.dynamics.ml on GitHub

Every unit to every unit through weight, [F, F], weight[i, j] from unit i to unit j.

The product multiplies from the right at precision, the matrix product precision of jax.lax.dot. A value per connection is [..., F, F].

FieldTypeDefault
weightjax.Array
precisionPrecisionLikestruct.field(pytree_node=False, default=None)
longest_delayint
def connections(shape: tuple[int, ...]) -> tuple[int, ...]
def send(output: jax.Array, fast: jax.Array | None) -> jax.Array
def presynaptic(x: jax.Array) -> jax.Array
def postsynaptic(x: jax.Array) -> jax.Array
def per_example(x: jax.Array) -> jax.Array
class DopamineSTDP

sparx.dynamics.plasticity on GitHub

Reward-modulated STDP (Izhikevich 2007): NEST’s stdp_dopamine_synapse (Potjans et al. 2010).

Pair STDP writes each synapse’s eligibility c instead of its weight, and the weight follows the eligibility gated by the concentration n of the network’s modulator modulator, above a baseline b:

dc/dt = -c / tau_c + a_plus K_pre delta(t - t_post) - a_minus K_post delta(t - t_pre)
dw/dt = c (n - b)

K_pre sums exp(-s / tau_plus) over earlier presynaptic spikes, each s ms before, and K_post sums exp(-s / tau_minus) over earlier postsynaptic arrivals, each read just before the spike that uses it, as in PairSTDP. Between steps c decays with tau_c and n with tau_n, so the weight integrates their product exactly, as NEST does between events; tau_n must be the modulator’s tau, which the network checks, so it is a Python number and not traced. A modulator whose release is 1 / tau_n adds NEST’s increment per dopamine spike.

The weight is clipped to [w_min, w_max] after every step, where NEST clips at events (spikes, dopamine arrivals and the volume transmitter’s deliveries). The two agree while n - b keeps its sign between dopamine releases, which it always does for b = 0, or while the weight stays inside its bounds. The defaults are NEST’s, Izhikevich’s (2007) values.

FieldTypeDefault
modulatorstrstruct.field(pytree_node=False, default='dopamine')
tau_plusjax.Array | float20.0
tau_minusjax.Array | float20.0
tau_cjax.Array | float1000.0
tau_nfloatstruct.field(pytree_node=False, default=200.0)
a_plusjax.Array | float1.0
a_minusjax.Array | float1.5
bjax.Array | float0.0
w_minjax.Array | float0.0
w_maxjax.Array | float200.0
def init_state(pre: int, post: int, edges: int, dtype: jnp.dtype = jnp.float32) -> DopamineTraces
def modulated_by() -> Mapping[str, float | None]
def step(traces: DopamineTraces, weights: jax.Array, pre_spikes: jax.Array, post_arrivals: jax.Array, pre: jax.Array, post: jax.Array, dt: float, modulators: Mapping[str, jax.Array]) -> tuple[DopamineTraces, jax.Array]

As Plasticity.step: the weight over the step, then this step’s spikes on the eligibility.

class DopamineTraces(NamedTuple)

sparx.dynamics.plasticity on GitHub

FieldTypeDefault
prejax.Array
postjax.Array
eligibilityjax.Array
dopaminejax.Array
class EligibleHebb(NamedTuple)

sparx.dynamics.ml on GitHub

FieldTypeDefault
hebbjax.Array
eligibilityjax.Array
class Exponential

sparx.dynamics.synapses on GitHub

Jumps by the arriving weight and decays with tau (ms): NEST’s iaf_psc_exp and iaf_cond_exp.

FieldTypeDefault
taujax.Array | float5.0
landsLanding
def init_state(shape: tuple[int, ...], dtype: jnp.dtype) -> jax.Array
def output(state: jax.Array) -> tuple[Term, ...]
def step(state: jax.Array, arriving: jax.Array | float, dt: float) -> jax.Array
class FastWeights

sparx.dynamics.ml on GitHub

Fast weights on a wiring: alpha times the Hebbian trace that rule keeps, added to each weight.

Differentiable plasticity (Miconi et al. 2018): alpha is how much of its trace each connection adds, one per connection ([F, F] on a Dense wiring, [E] on a Sparse one) or anything that broadcasts to it ([F], one per postsynaptic unit on a dense wiring, as in Backpropamine’s language models), learned by backpropagating through the traces. Each example’s trace starts at zero, so what the network stores in it is what that sequence taught it. The rules are DecayingHebb and OjaHebb (2018), ModulatedHebb and RetroactiveHebb (Backpropamine, 2019), or any HebbianRule.

connections makes only some connections of a Sparse wiring plastic, by their indices in its edge list: the rule keeps a trace for those alone, alpha is one per plastic connection, and the others keep their weights. A network whose few connections adapt fast, as most of a brain’s synapses do not, then holds and backpropagates through traces for those few.

FieldTypeDefault
alphajax.Array | float
ruleHebbianRule[Trace]
connectionsjax.Array | NoneNone
class Gap(NamedTuple)

sparx.dynamics.core on GitHub

Gap-junction coupling onto each neuron over a step: the current drive(s) - conductance * v(s) (pA).

conductance (nS) is the sum of the neuron’s junction conductances, sum_j g_ij, and drive is sum_j g_ij v_j(s) over the step, its partners’ voltages as a waveform (Term, linear in s). The current is linear in the neuron’s own voltage, so a model integrates the conductance with its synaptic conductances and the drive with its synaptic currents; a linear membrane solves both exactly, and stays stable however strong the coupling.

FieldTypeDefault
conductancejax.Array
driveTerm
class Graded

sparx.dynamics.synapses on GitHub

Transmission proportional to a graded presynaptic value, with first-order kinetics.

tau ds/dt = sum_j w_j r_j - s

r_j is the output of presynaptic neuron j (a GradedPotential’s release, a rate unit’s activity), arriving at the end of each step and held over the next, so s relaxes toward the weighted release with time constant tau (ms), and at steady state transmits w at full release. The waveform over a step, target + (value - target) exp(-s / tau), is two Terms, which a linear membrane integrates exactly. Prinz, Bucher and Marder’s (2004) graded synapse has the same form with a time constant that depends on the presynaptic voltage; here it is fixed.

FieldTypeDefault
taujax.Array | float5.0
landsLanding
def init_state(shape: tuple[int, ...], dtype: jnp.dtype) -> GradedState
def output(state: GradedState) -> tuple[Term, ...]
def step(state: GradedState, arriving: jax.Array | float, dt: float) -> GradedState
class GradedPotential

sparx.dynamics.neurons on GitHub

A neuron that never spikes and releases transmitter as a sigmoidal function of its voltage.

C dv/dt = -g_L (v - E_L) + sum_k g_k (E_k - v) + I, g_L = C / tau_m
release(v) = 1 / (1 + exp((v_half - v) / slope))

The membrane is LeakyIntegrateAndFire’s without a threshold, solved exactly over each step for held inputs, with conductances, gates, synaptic current waveforms and gap junctions as LeakyIntegrateAndFire reads them. The output is the release at the end of the step, as a fraction of the maximal rate, in [0, 1]. A Graded synapse holds it over the next step and filters it with its own kinetics, so the weight of a projection from a graded population is the current (pA) or conductance (nS) of its synapses at full release.

The sigmoid is the graded transmission of the stomatogastric network models of Prinz, Bucher and Marder (Nature Neuroscience 2004), s(V_pre) = 1 / (1 + exp((V_th - V_pre) / delta)), whose V_th of -35 mV and delta of 5 mV are the defaults here. Much of the fly’s visual system is non-spiking in this sense, and Lappalainen et al.’s connectome-constrained model of it (Nature 2024) also uses passive point neurons that transmit a function of their voltage, there a rectified linear one. The membrane’s defaults are LeakyIntegrateAndFire’s.

FieldTypeDefault
tau_mjax.Array | float20.0
c_mjax.Array | float200.0
e_ljax.Array | float-60.0
v_halfjax.Array | float-35.0
slopejax.Array | float5.0
i_ejax.Array | float0.0
reversalMapping[str, float]
gatesMapping[str, MgBlock]
def release(v: jax.Array) -> jax.Array

The release at voltage v, as a fraction of the maximum.

def init_state(shape: tuple[int, ...], dtype: jnp.dtype) -> GradedPotentialState
def step(state: GradedPotentialState, inputs: SynapticInput, dt: float) -> tuple[GradedPotentialState, Output]
def is_refractory(state: GradedPotentialState, dt: float) -> jax.Array
def after_threshold(state: GradedPotentialState, jump: jax.Array, fired: jax.Array) -> GradedPotentialState
class GradedPotentialState(NamedTuple)

sparx.dynamics.neurons on GitHub

FieldTypeDefault
vjax.Array
class GradedState(NamedTuple)

sparx.dynamics.synapses on GitHub

FieldTypeDefault
valuejax.Array
targetjax.Array
class HebbianRule(Protocol)

sparx.dynamics.ml on GitHub

How a recurrent layer’s Hebbian trace changes over a step, on any Wiring.

A rule keeps a Trace of what it has seen, init_trace(wiring, shape, dtype) at the start of a sequence for outputs of shape [..., F], and hebb(trace) is the value per connection that FastWeights.alpha scales. update(trace, wiring, pre, post, dt) takes the output fed back this step (pre, [..., F]) and the step’s new output (post), and reads them through the wiring’s presynaptic, postsynaptic and per_example.

def init_trace(wiring: Wiring, shape: tuple[int, ...], dtype: jnp.dtype) -> Trace
def hebb(trace: Trace) -> jax.Array
def update(trace: Trace, wiring: Wiring, pre: jax.Array, post: jax.Array, dt: float) -> Trace
class HodgkinHuxley

sparx.dynamics.neurons on GitHub

The squid giant axon (Hodgkin and Huxley, J. Physiol. 1952), as NEST’s hh_psc_alpha states it.

C dv/dt = -g_Na m^3 h (v - E_Na) - g_K n^4 (v - E_K) - g_L (v - E_L) + I
dx/dt = alpha_x(v) (1 - x) - beta_x(v) x, x in m, h, n

with the rates (1/ms) in the modern convention, rest near -65 mV: alpha_n = 0.01 (v + 55) / (1 - exp(-(v + 55) / 10)), beta_n = 0.125 exp(-(v + 65) / 80), alpha_m = 0.1 (v + 40) / (1 - exp(-(v + 40) / 10)), beta_m = 4 exp(-(v + 65) / 18), alpha_h = 0.07 exp(-(v + 65) / 20), beta_h = 1 / (1 + exp(-(v + 35) / 10)). Conductances are NEST’s, for a 100 pF membrane (1 uF/cm^2 over 1e-4 cm^2).

The membrane has no reset; a spike is its peak. As in NEST, a spike is reported at the step where the voltage is at or above v_spike (0 mV) and has begun to fall, and none is reported within t_ref of one.

scheme names the integration, in substeps of at most substep ms:

  • "strang" (the default): half a substep of each gate with the voltage held (exact, Rush and Larsen 1978), a substep of the voltage with the gates held (exact), and half a substep of the gates again. Second order, and stable at any substep, since each part is solved exactly. At 0.01 ms (the default) it fires with NEST’s adaptive solution within a step.
  • "rk4": classical RK4 of the whole system. At 0.025 ms it is within 0.01 mV of NEST, but the gates grow stiff under hyperpolarization (beta_m is 130/ms at -128 mV) and RK4 diverges there; shorter substeps only move the limit.
  • "exponential_euler": gates and voltage each advanced exactly from the start of the substep, Brian2’s exponential_euler (with substep=None, Brian2’s step exactly); stable and first order.
FieldTypeDefault
c_mjax.Array | float100.0
g_najax.Array | float12000.0
g_kjax.Array | float3600.0
g_ljax.Array | float30.0
e_najax.Array | float50.0
e_kjax.Array | float-77.0
e_ljax.Array | float-54.402
v_initjax.Array | float-65.0
v_spikejax.Array | float0.0
t_reffloatstruct.field(pytree_node=False, default=2.0)
schemeLiteral['strang', 'rk4', 'exponential_euler']struct.field(pytree_node=False, default='strang')
substepfloat | Nonestruct.field(pytree_node=False, default=0.01)
reversalMapping[str, float]
gatesMapping[str, MgBlock]
surrogateSurrogatestruct.field(pytree_node=False, default=ATan())
def rates(v: jax.Array) -> tuple[tuple[jax.Array, jax.Array], ...]

(alpha, beta) of m, h and n at v, in 1/ms.

def init_state(shape: tuple[int, ...], dtype: jnp.dtype) -> HodgkinHuxleyState
def step(state: HodgkinHuxleyState, inputs: SynapticInput, dt: float) -> tuple[HodgkinHuxleyState, Output]
def is_refractory(state: HodgkinHuxleyState, dt: float) -> jax.Array
def after_threshold(state: HodgkinHuxleyState, jump: jax.Array, fired: jax.Array) -> HodgkinHuxleyState
class HodgkinHuxleyState(NamedTuple)

sparx.dynamics.neurons on GitHub

FieldTypeDefault
vjax.Array
mjax.Array
hjax.Array
njax.Array
refractoryjax.Array
class Izhikevich

sparx.dynamics.neurons on GitHub

Izhikevich’s simple model (IEEE Trans. Neural Netw. 2003), in ms and mV.

dv/dt = 0.04 v^2 + 5 v + 140 - u + I, du/dt = a (b v - u)

At v >= v_th (30 mV) it fires, v is set to c and u grows by d. I is in the model’s own units, as in the paper. scheme names the integration: "published" takes two half-steps of v and then one step of u from the new v, the 2003 paper’s code (at dt = 1 with order="izhikevich" it is that code to the last bit, and its firing patterns are this scheme’s); "euler" is the forward Euler step of both from the old values, NEST’s consistent_integration; "semi_implicit" is one Euler step of v and then one of u from the new v, the code of the 2004 paper’s twenty firing patterns (izhikevich_2004). Synaptic current waveforms are read at the start of the step and conductances at the voltage of each update. The defaults are the regular spiking cell; the 2003 paper’s classes are izhikevich_2003.

order names the arithmetic of the quadratic term, which the membrane amplifies from the last bit to whole spikes over a long run: "izhikevich" squares first, 0.04 * v**2, as his MATLAB code does; "nest" multiplies (0.04 * v) * v, as NEST’s izhikevich does.

quadratic holds the coefficients of v^2, v and 1, and du/dt = a (b (v - v_u) - u_decay u): the defaults are the equations above, and two of the 2004 patterns change them (class 1 excitability and the integrator take 4.1 v + 108, accommodation du/dt = a b (v + 65)).

FieldTypeDefault
ajax.Array | float0.02
bjax.Array | float0.2
cjax.Array | float-65.0
djax.Array | float8.0
v_thjax.Array | float30.0
v_initjax.Array | float-65.0
v_ujax.Array | float0.0
u_decayjax.Array | float1.0
quadratictuple[float, float, float]struct.field(pytree_node=False, default=(0.04, 5.0, 140.0))
schemeLiteral['published', 'euler', 'semi_implicit']struct.field(pytree_node=False, default='published')
orderLiteral['izhikevich', 'nest']struct.field(pytree_node=False, default='izhikevich')
reversalMapping[str, float]
gatesMapping[str, MgBlock]
surrogateSurrogatestruct.field(pytree_node=False, default=ATan())
def init_state(shape: tuple[int, ...], dtype: jnp.dtype) -> IzhikevichState
def step(state: IzhikevichState, inputs: SynapticInput, dt: float) -> tuple[IzhikevichState, Output]
def is_refractory(state: IzhikevichState, dt: float) -> jax.Array
def after_threshold(state: IzhikevichState, jump: jax.Array, fired: jax.Array) -> IzhikevichState
class IzhikevichState(NamedTuple)

sparx.dynamics.neurons on GitHub

FieldTypeDefault
vjax.Array
ujax.Array
class LICell

sparx.dynamics.ml on GitHub

A leaky integrator that never fires; its output is its membrane, a graded value.

The usual readout of a spiking classifier: the last layer integrates the spikes it receives and the loss reads its membrane. As the first model of a Serial it is a synapse’s decaying current.

FieldTypeDefault
decayjax.Array | float
def init_state(shape: tuple[int, ...], dtype: jnp.dtype) -> MembraneState
def step(state: MembraneState, inputs: SynapticInput, dt: float) -> tuple[MembraneState, Output]
def is_refractory(state: MembraneState, dt: float) -> jax.Array
def after_threshold(state: MembraneState, jump: jax.Array, fired: jax.Array) -> MembraneState
class LIFCell

sparx.dynamics.ml on GitHub

Leaky integrate-and-fire. decay=1 integrates without leak (IF).

FieldTypeDefault
decayjax.Array | float
thresholdjax.Array | float1.0
resetResetstruct.field(pytree_node=False, default='subtract')
surrogateSurrogatestruct.field(pytree_node=False, default=ATan())
detach_resetboolstruct.field(pytree_node=False, default=False)
def init_state(shape: tuple[int, ...], dtype: jnp.dtype) -> MembraneState
def step(state: MembraneState, inputs: SynapticInput, dt: float) -> tuple[MembraneState, Output]
def is_refractory(state: MembraneState, dt: float) -> jax.Array
def after_threshold(state: MembraneState, jump: jax.Array, fired: jax.Array) -> MembraneState
Landing

sparx.dynamics.synapses on GitHub

Where a synapse’s arrivals reach the neuron. synapse: into the synapse’s state at the end of the step, whose output shapes the membrane from the next step on. before_threshold: as a voltage jump at the end of the step, before the threshold test (NEST’s delta synapse). after_threshold: as a voltage jump after the test and before the reset (Brian2’s on_pre="v += w").

class LeakyIntegrateAndFire

sparx.dynamics.neurons on GitHub

Leaky integrate-and-fire with current and conductance input.

C dv/dt = -g_L (v - E_L) + sum_k g_k (E_k - v) + I, g_L = C / tau_m

At v >= v_th the neuron fires, v is set to v_reset and held there for t_ref. With the inputs constant over a step the equation is linear in v, and the update is its exact solution: v relaxes toward (g_L E_L + sum g_k E_k + I) / (g_L + sum g_k) with time constant C / (g_L + sum g_k), and synaptic current waveforms are integrated exactly against that time constant, as NEST’s iaf_psc_exp and iaf_psc_alpha do. Conductances are read against reversal, by receptor name, held over the step as Brian2’s exponential_euler holds them; a receptor in gates is scaled by its gate at the voltage at the start of the step (NMDA’s magnesium block by default). A voltage jump from delta synapses lands after the step’s integration, before the threshold test, and is lost during refractoriness, as in NEST’s iaf_psc_delta. The defaults are the cortical cell of Brette et al.’s benchmarks (2007): 20 ms, 200 pF, rest -60 mV, threshold -50 mV, reset -60 mV, 5 ms refractory.

FieldTypeDefault
tau_mjax.Array | float20.0
c_mjax.Array | float200.0
e_ljax.Array | float-60.0
v_thjax.Array | float-50.0
v_resetjax.Array | float-60.0
t_refjax.Array | float5.0
i_ejax.Array | float0.0
reversalMapping[str, float]
gatesMapping[str, MgBlock]
surrogateSurrogatestruct.field(pytree_node=False, default=ATan())
detach_resetboolstruct.field(pytree_node=False, default=False)
def init_state(shape: tuple[int, ...], dtype: jnp.dtype) -> LeakyIntegrateAndFireState
def step(state: LeakyIntegrateAndFireState, inputs: SynapticInput, dt: float) -> tuple[LeakyIntegrateAndFireState, Output]
def is_refractory(state: LeakyIntegrateAndFireState, dt: float) -> jax.Array
def after_threshold(state: LeakyIntegrateAndFireState, jump: jax.Array, fired: jax.Array) -> LeakyIntegrateAndFireState
class LeakyIntegrateAndFireState(NamedTuple)

sparx.dynamics.neurons on GitHub

FieldTypeDefault
vjax.Array
refractoryjax.Array
class MembraneState(NamedTuple)

sparx.dynamics.ml on GitHub

FieldTypeDefault
vjax.Array
class MgBlock

sparx.dynamics.neurons on GitHub

The fraction of NMDA receptors not blocked by magnesium at voltage v (mV).

B(v) = 1 / (1 + [Mg] / 3.57 exp(-0.062 v))

Jahr and Stevens (J. Neurosci. 1990), mg the extracellular concentration in mM. A conductance-based model scales its nmda conductance by B at the voltage at the start of each step.

FieldTypeDefault
mgfloat1.0
def __call__(v: jax.Array) -> jax.Array
class Model(Protocol)

sparx.dynamics.core on GitHub

Anything run steps: a neuron model, or a neuron with the synapses onto it (PointNeuron).

FieldTypeDefault
gradedbool
def init_state(shape: tuple[int, ...], dtype: jnp.dtype) -> State

The population at rest, for inputs of per-step shape shape and dtype dtype.

def step(state: State, inputs: Inputs, dt: float) -> tuple[State, Output]

Advance one step of dt on inputs; return the new state and the step’s output.

class ModulatedHebb

sparx.dynamics.ml on GitHub

A Hebbian trace whose rate the network’s own activity sets through a neuromodulator, clipped.

m = tanh(sum_j post[j] modulator[j] + modulator_bias)
eta[j] = m fanout[j] + fanout_bias[j]
hebb[i, j] <- clip(hebb[i, j] + dt eta[j] pre[i] post[j], -clip, clip)

Backpropamine’s simple neuromodulation (Miconi et al. 2019, eq. 3), with the fan-out of their appendix and simplemaze/maze.py: one modulator level per example, read off the new activity (their h2mod, [F] and a scalar), fanned out to a rate per postsynaptic unit (their modfanout, [F] and [F]). The rate can be negative, so the trace can grow, shrink or flip sign; the clip, 2 in their code, bounds it.

FieldTypeDefault
modulatorjax.Array
modulator_biasjax.Array | float
fanoutjax.Array | float
fanout_biasjax.Array | float
clipfloatstruct.field(pytree_node=False, default=2.0)
def init_trace(wiring: Wiring, shape: tuple[int, ...], dtype: jnp.dtype) -> jax.Array
def hebb(trace: jax.Array) -> jax.Array
def update(trace: jax.Array, wiring: Wiring, pre: jax.Array, post: jax.Array, dt: float) -> jax.Array
class NeuronModel(Model[State, SynapticInput], Protocol)

sparx.dynamics.core on GitHub

One population of neurons, advanced one step at a time on the synaptic input it receives.

Besides stepping, a model answers the two questions a PointNeuron asks of it between steps, so the synapses can follow the neuron without reading its state’s fields: where it is refractory, and how a voltage jump that lands after the threshold test changes it.

def is_refractory(state: State, dt: float) -> jax.Array

Where the neuron cannot fire in the coming step of dt, as booleans; all False for a model without refractoriness.

def after_threshold(state: State, jump: jax.Array, fired: jax.Array) -> State

state after a voltage jump that lands past the step’s threshold test, before its reset.

fired is the step’s Output.value. A model whose spike resets the membrane loses the jump where it fired, since the reset that follows overwrites it; one that does not reset keeps it.

class OjaHebb

sparx.dynamics.ml on GitHub

Oja’s rule: a Hebbian trace that each postsynaptic unit’s own activity bounds.

hebb[i, j] <- hebb[i, j] + dt eta post[j] (pre[i] - post[j] hebb[i, j])

Differentiable plasticity’s alternative to the decaying trace (Miconi et al. 2018, eq. 3; their maze/maze.py with rule="oja"), which keeps a memory without input instead of letting it decay to zero (Oja 1982).

FieldTypeDefault
etajax.Array | float
def init_trace(wiring: Wiring, shape: tuple[int, ...], dtype: jnp.dtype) -> jax.Array
def hebb(trace: jax.Array) -> jax.Array
def update(trace: jax.Array, wiring: Wiring, pre: jax.Array, post: jax.Array, dt: float) -> jax.Array
class Output(NamedTuple)

sparx.dynamics.core on GitHub

What a population sends in a step, and when within the step.

For a spiking model value is 0 or 1, and offset is the fraction of the step, in [0, 1], at which the membrane crossed threshold, from linear interpolation of the voltage across the step (Hansel et al., Neural Computation 1998); 1 where nothing fired. A spike’s time is (step + offset) * dt from the start of the run. For a graded model value is real (a release rate, an activity or a membrane), its value at the end of the step, and offset is 1.

FieldTypeDefault
valuejax.Array
offsetjax.Array
class PairSTDP

sparx.dynamics.plasticity on GitHub

All-to-all pair STDP with soft or hard bounds: NEST’s stdp_synapse (Guetig et al. 2003).

With weights normalized by w_max, a postsynaptic spike potentiates and a presynaptic spike depresses:

w <- min(w + lambda (1 - w)^mu_plus K_pre, 1)
w <- max(w - alpha lambda w^mu_minus K_post, 0)

where K_pre sums exp(-s / tau_plus) over earlier presynaptic spikes, each s ms before, and K_post sums exp(-s / tau_minus) over earlier postsynaptic arrivals. mu = 0 is additive STDP (Song et al. 2000), mu = 1 multiplicative (van Rossum et al. 2000).

FieldTypeDefault
tau_plusjax.Array | float20.0
tau_minusjax.Array | float20.0
lambda_jax.Array | float0.01
alphajax.Array | float1.0
mu_plusjax.Array | float1.0
mu_minusjax.Array | float1.0
w_maxjax.Array | float100.0
def init_state(pre: int, post: int, edges: int, dtype: jnp.dtype = jnp.float32) -> STDPTraces
def modulated_by() -> Mapping[str, float | None]
def step(traces: STDPTraces, weights: jax.Array, pre_spikes: jax.Array, post_arrivals: jax.Array, pre: jax.Array, post: jax.Array, dt: float, modulators: Mapping[str, jax.Array]) -> tuple[STDPTraces, jax.Array]

As Plasticity.step.

class Plasticity(Protocol)

sparx.dynamics.plasticity on GitHub

A rule that changes the weights of a projection’s edges with the spikes on either side.

Its traces live on neurons and edges, init_state(pre, post, edges, dtype) for pre presynaptic and post postsynaptic neurons joined by edges synapses, and step advances them and the weights [E] of the edges pre[E] -> post[E] by one step.

def init_state(pre: int, post: int, edges: int, dtype: jnp.dtype = jnp.float32) -> Traces

The traces with no spike yet.

def modulated_by() -> Mapping[str, float | None]

The modulators the rule reads, by name, each with the time constant (ms) the rule assumes its concentration decays with, or None where the rule reads the concentration alone.

A network refuses a rule that reads a modulator it lacks, or one whose time constant differs.

def step(traces: Traces, weights: jax.Array, pre_spikes: jax.Array, post_arrivals: jax.Array, pre: jax.Array, post: jax.Array, dt: float, modulators: Mapping[str, jax.Array]) -> tuple[Traces, jax.Array]

One step: traces decay over dt, then this step’s spikes update weights and traces.

pre_spikes[N_pre] and post_arrivals[N_post] are 0 or 1; weights[E] belong to the edges pre[E] -> post[E]. modulators holds each neuromodulator’s concentration after this step’s release, by name, a scalar each (with a leading trial axis under vmap).

class PointNeuron

sparx.dynamics.synapses on GitHub

A neuron model and the synapses onto it, by receptor, stepped together.

PointNeuron(LIF(...), {"ex": Receptor(Exponential(2.0))}) is NEST’s iaf_psc_exp for one excitatory receptor; conductance receptors are named for their reversal potentials ("ampa", "gaba_a", …).

A conductance changes within a step while the neuron holds it, and hold picks the value held. "mean", the default, is its exact average over the step, which makes the leak’s decay over the step exact and the scheme second order. "start" is its value at the start of the step, Brian2’s exponential_euler, first order.

Two options reproduce models that tie their synapses to the neuron’s spike, as Shiu et al.’s (2024) whole-brain model does in Brian2: reset_synapses clears a neuron’s synaptic state when it fires (after the step’s arrivals), and freeze_synapses holds it while the neuron is refractory: it does not decay, and arrivals are discarded, which is what Brian2’s (unless refractory) flag on a synaptic variable does (a conditional write that stops synaptic updates too). Neither is physiology, where a synaptic current outlives the spike and input during refractoriness is not lost; both are off by default.

FieldTypeDefault
neuronNeuronModel[State]
receptorsMapping[str, Receptor]
holdLiteral['mean', 'start']struct.field(pytree_node=False, default='mean')
reset_synapsesboolstruct.field(pytree_node=False, default=False)
freeze_synapsesboolstruct.field(pytree_node=False, default=False)
gradedbool
def init_state(shape: tuple[int, ...], dtype: jnp.dtype) -> PointNeuronState[State]
def landing(where: Landing) -> list[str]

The receptors whose arrivals land where.

def advance(state: PointNeuronState[State], current: jax.Array | float, jump: jax.Array | float, dt: float, gap: Gap | None = None, noise: jax.Array | None = None) -> tuple[PointNeuronState[State], Output]

Move the membrane over a step on the synapses’ output and gap, with jump (mV) landing at its end and noise for a stochastic model.

def receive(state: PointNeuronState[State], arriving: Mapping[str, jax.Array], dt: float, fired: jax.Array, frozen: jax.Array) -> PointNeuronState[State]

Decay the synapses over the step and add the weights due at its end; jumps after the threshold land on the neuron.

fired is the step’s output, and frozen where the neuron was refractory during the step, as frozen read it before the step; reset_synapses and freeze_synapses act on them.

def frozen(state: PointNeuronState[State], dt: float) -> jax.Array

Where the neuron is refractory for the coming step, for freeze_synapses.

def delta(arriving: Mapping[str, jax.Array]) -> jax.Array

The voltage jump (mV) of the weights arriving on receptors that land before the threshold.

def step(state: PointNeuronState[State], inputs: Arrivals, dt: float) -> tuple[PointNeuronState[State], Output]
class PointNeuronState(NamedTuple)

sparx.dynamics.synapses on GitHub

A neuron’s state and its synapses’ by receptor name.

Each receptor’s synapse model has its own state type, and Python has no way to type a mapping whose values differ by key, so the synapses’ states are left untyped.

FieldTypeDefault
neuronState
synapsesMapping[str, object]
class PulseCell

sparx.dynamics.ml on GitHub

RNeuralNet’s neuron: each step it outputs a function of the sum of what arrived, and empties the sum.

s[t] = v[t-1] + x[t]
output[t] = s[t] if s[t] >= threshold
exp(s[t] - threshold) - 1 otherwise
v[t] = 0

The activation of Soma_t::ActivationFunction in Ashish Kumar Singh’s RNeuralNet-Research (commit d4b7803, with its A_CONST of 1): an ELU whose exponential branch is shifted by the threshold while its linear branch is not, so the output jumps from 0 to threshold there, and an empty sum gives exp(-threshold) - 1, not 0. The output is a graded message, sent each step, after which the sum empties, as the original neuron resets once it has sent on every outgoing connection. A jump that lands after the step (after_threshold) waits in v for the next. The neuron has no time constant, so dt changes nothing; a threshold of minus infinity passes the sum through unchanged, the original’s input neurons. sparx.learn.RNeuralNet wires these neurons as the original does.

FieldTypeDefault
thresholdjax.Array | float2.0
def init_state(shape: tuple[int, ...], dtype: jnp.dtype) -> MembraneState
def step(state: MembraneState, inputs: SynapticInput, dt: float) -> tuple[MembraneState, Output]
def is_refractory(state: MembraneState, dt: float) -> jax.Array
def after_threshold(state: MembraneState, jump: jax.Array, fired: jax.Array) -> MembraneState
class RateCell

sparx.dynamics.ml on GitHub

A leaky rate unit; its output is its activity h, a graded value.

alpha = decay ** dt
h[t] = alpha h[t-1] + (1 - alpha) f(x[t] + bias)

FLYNN’s leaky integrator (Wang and Chen, arXiv 2607.00025, eq. 1), h_{t+1} = alpha h_t + (1 - alpha) tanh(W h_t + x_t + b), trained on the whole fly connectome. The recurrent product W h_t arrives with the external input in the jump x: through a RecurrentCell (sparx.nn.Recurrent(sparx.nn.Rate())), or through a projection onto a Delta receptor of a sparx.graph.Network with a delay of one step, which delivers the activity of the step before. decay is FLYNN’s leak rate alpha per unit of time, exp(-1 / tau) for a time constant of tau (sparx.dynamics.decay); they train it directly, one per cell class. activation is FLYNN’s tanh, or ReLU or the logistic sigmoid (ACTIVATIONS). The activity is a convex combination of its last value and the activation, so it stays within the activation’s range.

FieldTypeDefault
decayjax.Array | float
biasjax.Array | float0.0
activationLiteral['tanh', 'relu', 'sigmoid']struct.field(pytree_node=False, default='tanh')
def init_state(shape: tuple[int, ...], dtype: jnp.dtype) -> RateState
def step(state: RateState, inputs: SynapticInput, dt: float) -> tuple[RateState, Output]
def is_refractory(state: RateState, dt: float) -> jax.Array
def after_threshold(state: RateState, jump: jax.Array, fired: jax.Array) -> RateState
class RateState(NamedTuple)

sparx.dynamics.ml on GitHub

FieldTypeDefault
hjax.Array
class Receptor

sparx.dynamics.synapses on GitHub

A synapse model and how its output reaches the membrane.

kind="current" adds the waveform as a current (pA); kind="conductance" holds its value over the step (as PointNeuron.hold picks it) as a conductance (nS) against the neuron’s reversal potential for this receptor’s name. A synapse that lands as a voltage jump (Delta) is one either way.

FieldTypeDefault
synapseSynapseModelstruct.field(default_factory=Exponential)
kindLiteral['current', 'conductance']struct.field(pytree_node=False, default='current')
class RecurrentCell

sparx.dynamics.ml on GitHub

Feed a model’s output back to its own input through wiring, fixed or plastic.

output[t] = inner.step(x[t] + send(output[t-1], alpha * hebb[t-1]))
hebb[t] = rule.update(hebb[t-1], output[t-1], output[t], dt)

The wiring is Dense, every unit to every unit, or Sparse, along a list of edges such as a connectome’s, where each edge may take its own number of steps (Sparse.delay): the output then arrives that many steps after it was sent, weighted as the wiring was when it left, and the state holds what is on its way. With fast_weights (FastWeights), a Hebbian trace each sequence writes adds fast weights to the wiring’s; without, the trace stays None and the weights fixed. On a wiring whose delays differ, a connection’s trace pairs the new output with what the connection delivers this step, its source’s output delay steps before, which the state keeps (RecurrentState.history), and a message carries the fast weights of the step that sent it. Any model runs inside: ALIFCell gives the recurrent adaptive network (LSNN) of Bellec et al. (2020), RateCell FLYNN’s recurrence (on a Sparse wiring, its connectome), RateCell(0.0) with plasticity Miconi et al.’s networks, and a spiking model fast weights between spikes. The trace is part of the state, so a stream fed in chunks carries it, and backpropagating through a plastic cell holds one per example per step. cut_gradient stops the gradient at the fed-back output (the weights still receive theirs), the gradient e-prop computes online (their stop_z_gradients).

FieldTypeDefault
innerNeuronModel[State]
wiringWiring
fast_weightsFastWeights[Trace] | NoneNone
cut_gradientboolstruct.field(pytree_node=False, default=False)
gradedbool
def init_state(shape: tuple[int, ...], dtype: jnp.dtype) -> RecurrentState[State, Trace]
def step(state: RecurrentState[State, Trace], inputs: SynapticInput, dt: float) -> tuple[RecurrentState[State, Trace], Output]
def is_refractory(state: RecurrentState[State, Trace], dt: float) -> jax.Array
def after_threshold(state: RecurrentState[State, Trace], jump: jax.Array, fired: jax.Array) -> RecurrentState[State, Trace]
class RecurrentState(NamedTuple)

sparx.dynamics.ml on GitHub

FieldTypeDefault
innerState
outputjax.Array
arrivingjax.Array
traceTrace | NoneNone
historyjax.Array | NoneNone
Reset

sparx.dynamics.core on GitHub

What a spike does to a dimensionless membrane: subtract the threshold (soft reset, which keeps the overshoot), set it to zero (hard reset), or leave it, none.

class RetroactiveHebb

sparx.dynamics.ml on GitHub

A Hebbian trace that a neuromodulator writes from an eligibility trace of recent coactivity.

m = tanh(sum_j post[j] modulator[j] + modulator_bias)
hebb[i, j] <- clip(hebb[i, j] + dt m eligibility[i, j], -clip, clip)
keep = (1 - eta) ** dt
eligibility[i, j] <- keep eligibility[i, j] + (1 - keep) pre[i] post[j]

Backpropamine’s retroactive neuromodulation (Miconi et al. 2019, eqs. 4 and 5; their maze/batch.py with type="modul" and the hard clip at 1, addpw=3). The coactivity changes no weight until the modulator arrives, so a signal that comes later (a reward) can still credit the connections that were active before it, as dopamine gates the plasticity recent activity left behind.

FieldTypeDefault
modulatorjax.Array
modulator_biasjax.Array | float
etajax.Array | float
clipfloatstruct.field(pytree_node=False, default=1.0)
def init_trace(wiring: Wiring, shape: tuple[int, ...], dtype: jnp.dtype) -> EligibleHebb
def hebb(trace: EligibleHebb) -> jax.Array
def update(trace: EligibleHebb, wiring: Wiring, pre: jax.Array, post: jax.Array, dt: float) -> EligibleHebb
class STDPTraces(NamedTuple)

sparx.dynamics.plasticity on GitHub

FieldTypeDefault
prejax.Array
postjax.Array
class Serial

sparx.dynamics.ml on GitHub

Two models in series: each step, the first’s Output.value is the second’s input jump.

Serial(LICell(synapse_decay), LIFCell(decay)) is the current-based LIF, whose input charges a decaying synaptic current that the membrane integrates,

i[t] = synapse_decay * i[t-1] + x[t]
v[t] = decay * v[t-1] + i[t]

snnTorch’s Synaptic and the CUBA neurons of Zenke and Vogels (2021). Longer chains nest. The output is the second model’s, in the dtype it gives it, and the pair is graded when the second model is. The step’s noise is the second model’s, which fires: Serial(LICell(...), BernoulliCell(...)) is a current-based neuron with escape noise. In physical units the counterpart is a PointNeuron with an Exponential synapse, where an arrival shapes the membrane from the next step on.

FieldTypeDefault
firstNeuronModel[First]
secondNeuronModel[Second]
gradedbool
def init_state(shape: tuple[int, ...], dtype: jnp.dtype) -> tuple[First, Second]
def step(state: tuple[First, Second], inputs: SynapticInput, dt: float) -> tuple[tuple[First, Second], Output]
def is_refractory(state: tuple[First, Second], dt: float) -> jax.Array
def after_threshold(state: tuple[First, Second], jump: jax.Array, fired: jax.Array) -> tuple[First, Second]
class Sparse

sparx.dynamics.ml on GitHub

Units wired along the edges pre[e] -> post[e] with weight[e], among size units.

For a wiring diagram whose edges are a small part of all pairs, a connectome: FLYNN (Wang and Chen, arXiv 2607.00025) trains the whole fly brain’s 5.3 million edges among 139 thousand neurons. Sending gathers the presynaptic outputs and sums them at their postsynaptic units, in time and memory proportional to the edges. A value per connection is one per edge, [..., E].

delay[e], whole steps from 1 to longest_delay, is how long an output takes to cross edge e; None is one step for every edge. A message is weighted when it is sent, so a weight that changes while it travels leaves it as it was. RNeuralNet’s connections queue their messages so (sparx.learn.RNeuralNet).

FieldTypeDefault
prejax.Array
postjax.Array
weightjax.Array
sizeintstruct.field(pytree_node=False)
delayjax.Array | NoneNone
longest_delayintstruct.field(pytree_node=False, default=1)
def connections(shape: tuple[int, ...]) -> tuple[int, ...]
def send(output: jax.Array, fast: jax.Array | None) -> jax.Array
def presynaptic(x: jax.Array) -> jax.Array
def postsynaptic(x: jax.Array) -> jax.Array
def per_example(x: jax.Array) -> jax.Array
class StochasticRelease

sparx.dynamics.synapses on GitHub

Release that succeeds at random: each spike, each synapse releases with probability p.

A presynaptic spike reaches each of its synapses, and each releases independently, a Bernoulli draw per edge per spike from the network’s noise key. A release transmits quantal times the edge’s weight, so the current one spike sends over n edges of weight w is binomial, with mean n p q w and variance n p (1 - p) (q w)^2; quantal = 1 / p keeps the mean of the deterministic projection. On a projection with short-term plasticity (TsodyksMarkram) the probability is p times the spike’s efficacy u x, so depression and facilitation change how often a synapse releases, and a release always transmits one quantum (p = 1 makes the efficacy the probability). The resources x still deplete by their mean, as in the deterministic model; depletion by the releases that happened, as in Fuhrmann et al.’s (2002) stochastic synapse, is not modeled.

FieldTypeDefault
pjax.Array | float0.5
quantaljax.Array | float1.0
def transmit(key: jax.Array, weight: jax.Array, sent: jax.Array) -> jax.Array

What edges of weight transmit for presynaptic values sent (0, 1 or an efficacy), one draw each.

class SynapseModel(Protocol)

sparx.dynamics.synapses on GitHub

FieldTypeDefault
landsLanding
def init_state(shape: tuple[int, ...], dtype: jnp.dtype) -> State
def output(state: State) -> tuple[Term, ...]
def step(state: State, arriving: jax.Array | float, dt: float) -> State
class SynapticInput

sparx.dynamics.core on GitHub

What a population receives in one step.

current (pA) adds to the membrane equation as it is, held over the step. currents are synaptic current waveforms over the step (Term). conductance (nS) holds one total conductance per receptor, keyed by the receptor’s name, which the neuron model pairs with its reversal potential, so an excitatory and an inhibitory conductance pull the voltage toward different targets. jump (mV) is added to the voltage at the end of the step, before the threshold test. gap is the step’s gap-junction coupling, None for a neuron without junctions. noise is a uniform draw in [0, 1) per neuron, which a stochastic model fires by (sparx.dynamics.BernoulliCell); None for every other model.

FieldTypeDefault
currentjax.Array | float0.0
currentstuple[Term, ...]()
conductanceMapping[str, jax.Array]struct.field(default_factory=dict)
jumpjax.Array | float0.0
gapGap | NoneNone
noisejax.Array | NoneNone
waveformstuple[Term, ...]
def current_at(s: jax.Array | float) -> jax.Array

The total current (pA) s ms into the step, the drive of gap junctions included.

class Term(NamedTuple)

sparx.dynamics.core on GitHub

A current (amplitude + slope * s) * exp(-s / tau) over a step, s in [0, dt] ms from its start.

amplitude in pA, slope in pA/ms, tau in ms.

FieldTypeDefault
amplitudejax.Array
slopejax.Array
taujax.Array | float
def at(s: jax.Array | float) -> jax.Array
def mean(dt: float) -> jax.Array

The term’s average over [0, dt]: exact.

class TripletSTDP

sparx.dynamics.plasticity on GitHub

All-to-all triplet STDP (Pfister and Gerstner 2006): NEST’s stdp_triplet_synapse.

A postsynaptic spike potentiates by K_pre (A2+ + A3+ y) and a presynaptic spike depresses by K_post (A2- + A3- r), weights clipped to [0, w_max], where K_pre, r are presynaptic traces with tau_plus, tau_x and K_post, y postsynaptic traces with tau_minus, tau_y, each read just before the spike that uses it. The defaults are Pfister and Gerstner’s all-to-all fit to visual cortex, whose A2+ (5e-10) is effectively zero. NEST’s synapse has the same defaults, but keeps tau_minus and tau_y on the postsynaptic neuron, where they default to 20 and 110 ms.

FieldTypeDefault
tau_plusjax.Array | float16.8
tau_xjax.Array | float101.0
tau_minusjax.Array | float33.7
tau_yjax.Array | float125.0
a2_plusjax.Array | float5e-10
a3_plusjax.Array | float0.0062
a2_minusjax.Array | float0.007
a3_minusjax.Array | float0.00023
w_maxjax.Array | float100.0
def init_state(pre: int, post: int, edges: int, dtype: jnp.dtype = jnp.float32) -> TripletTraces
def modulated_by() -> Mapping[str, float | None]
def step(traces: TripletTraces, weights: jax.Array, pre_spikes: jax.Array, post_arrivals: jax.Array, pre: jax.Array, post: jax.Array, dt: float, modulators: Mapping[str, jax.Array]) -> tuple[TripletTraces, jax.Array]

As PairSTDP.step.

class TripletTraces(NamedTuple)

sparx.dynamics.plasticity on GitHub

FieldTypeDefault
prejax.Array
pre_tripletjax.Array
postjax.Array
post_tripletjax.Array
class TsodyksMarkram

sparx.dynamics.plasticity on GitHub

Short-term depression and facilitation (Tsodyks and Markram 1997; Markram et al. 1998).

A spike releases a fraction u of the available resources x and transmits w u x. Between spikes x recovers toward 1 with tau_rec and u relaxes toward U with tau_fac (no facilitation when tau_fac is 0). At a spike, with h since the last one:

u = U + u_last (1 - U) exp(-h / tau_fac)
x = 1 + (x_last (1 - u_last) - 1) exp(-h / tau_rec)

the update of NEST’s tsodyks2_synapse, kept per presynaptic neuron and integrated exactly between spikes. A synapse starts at rest, fully recovered and unfacilitated, so its first spike transmits w U. NEST transmits its initial w u x at the first spike, and its initial u is 0.5 whatever U is; it agrees with sparx when u is set to U.

FieldTypeDefault
Ujax.Array | float0.5
tau_recjax.Array | float800.0
tau_facjax.Array | float0.0
def init_state(shape: tuple[int, ...], dtype: jnp.dtype = jnp.float32) -> TsodyksMarkramState

Fully recovered and unfacilitated, per presynaptic neuron.

def step(state: TsodyksMarkramState, spikes: jax.Array, dt: float) -> tuple[TsodyksMarkramState, jax.Array]

Advance dt and release for this step’s spikes.

Returns the efficacy u x per presynaptic neuron, 0 where silent.

class TsodyksMarkramState(NamedTuple)

sparx.dynamics.plasticity on GitHub

FieldTypeDefault
recoveredjax.Array
facilitationjax.Array
class Wiring(Protocol)

sparx.dynamics.ml on GitHub

Which units of a recurrent layer reach which, and with what weights: Dense or Sparse.

send(output, fast) carries a layer’s outputs [..., F] to every unit’s input through the wiring’s weights, plus fast, extra weights per example and connection (a plastic trace’s), or None. An output takes one step to arrive, or as many as its connection’s delay, up to longest_delay, so send returns [longest_delay, ..., F], what arrives 1, 2, … steps after the step that sent it; each message is weighted when it is sent. A plasticity rule reads values at each connection’s presynaptic unit (presynaptic(x)), at its postsynaptic unit (postsynaptic(x)), or one per example at every connection (per_example(x)), so one rule serves every wiring. connections(shape) is the shape of one value per connection for outputs of shape, and refuses outputs the wiring does not fit.

FieldTypeDefault
longest_delayint
def connections(shape: tuple[int, ...]) -> tuple[int, ...]
def send(output: jax.Array, fast: jax.Array | None) -> jax.Array
def presynaptic(x: jax.Array) -> jax.Array
def postsynaptic(x: jax.Array) -> jax.Array
def per_example(x: jax.Array) -> jax.Array
def decay(tau: float, dt: float = 1.0) -> float

sparx.dynamics.core on GitHub

exp(-dt / tau): what a time constant tau leaves of a value after a step dt in the same unit.

def izhikevich_2003(kind: str, **fields={}) -> Izhikevich

sparx.dynamics.neurons on GitHub

The cortical and thalamic classes of Izhikevich (2003), Figure 2, by their names there.

def izhikevich_2004(pattern: str, **fields={}) -> Izhikevich

sparx.dynamics.neurons on GitHub

The neuron of one of the twenty firing patterns of Izhikevich (2004), Figure 1, by its name (IZHIKEVICH_2004), integrated as the paper’s code does (scheme="semi_implicit").

def run(model: Model[State, Inputs], inputs: Inputs | jax.Array | np.ndarray, state: State | None = None, dt: float = 1.0, record: Callable[[State], Record] | None = None, unroll: int | bool = 1) -> tuple[Output | tuple[Output, Record], State]

sparx.dynamics.core on GitHub

Scan model over time-major inputs; return its output [T, ...] and the final state.

inputs is what model.step takes for one step, with every leaf [T, ...] or a scalar held over time: a SynapticInput for a neuron model, Arrivals for a PointNeuron. An array [T, ...] stands for SynapticInput(jump=...), the dimensionless family’s input. The population’s per-step shape and its dtype are those of the first input with a time axis. state None starts the population at rest; a run over the first k steps and one over the rest from its final state equal one run over all. With record, each step’s output is paired with record(state) after the step (the membrane voltage, say). unroll is jax.lax.scan’s: how many steps one loop iteration holds.