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.
Contents
Section titled “Contents”| Name | |
|---|---|
ACTIVATIONS | The activations a RateCell takes, by name. |
IZHIKEVICH_2003 | (a, b, c, d) of each class in Izhikevich (2003), Figure 2. |
IZHIKEVICH_2004 | 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 | 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. |
ALIFCell | LIF with an adaptive threshold that rises with each spike and decays back. |
ALIFState | |
AdEx | The adaptive exponential integrate-and-fire neuron (Brette and Gerstner, J. Neurophysiol. 2005). |
AdExState | |
Alpha | w e / tau * s * exp(-s / tau), peaking at the weight w at s = tau: NEST’s iaf_psc_alpha. |
Arrivals | 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). |
BernoulliCell | A leaky integrate-and-fire neuron with escape noise: it fires with a probability of its membrane. |
BernoulliState | |
BiExponential | w f (exp(-s / tau_decay) - exp(-s / tau_rise)), normalized by f to peak at the weight w. |
DecayingHebb | A Hebbian trace that decays toward the latest coactivity at the rate eta. |
Delta | A voltage jump of the arriving weight (mV), with no kinetics. |
Dense | Every unit to every unit through weight, [F, F], weight[i, j] from unit i to unit j. |
DopamineSTDP | Reward-modulated STDP (Izhikevich 2007): NEST’s stdp_dopamine_synapse (Potjans et al. 2010). |
DopamineTraces | |
EligibleHebb | |
Exponential | Jumps by the arriving weight and decays with tau (ms): NEST’s iaf_psc_exp and iaf_cond_exp. |
FastWeights | Fast weights on a wiring: alpha times the Hebbian trace that rule keeps, added to each weight. |
Gap | Gap-junction coupling onto each neuron over a step: the current drive(s) - conductance * v(s) (pA). |
Graded | Transmission proportional to a graded presynaptic value, with first-order kinetics. |
GradedPotential | A neuron that never spikes and releases transmitter as a sigmoidal function of its voltage. |
GradedPotentialState | |
GradedState | |
HebbianRule | How a recurrent layer’s Hebbian trace changes over a step, on any Wiring. |
HodgkinHuxley | The squid giant axon (Hodgkin and Huxley, J. Physiol. 1952), as NEST’s hh_psc_alpha states it. |
HodgkinHuxleyState | |
Izhikevich | Izhikevich’s simple model (IEEE Trans. Neural Netw. 2003), in ms and mV. |
IzhikevichState | |
LICell | A leaky integrator that never fires; its output is its membrane, a graded value. |
LIFCell | Leaky integrate-and-fire. decay=1 integrates without leak (IF). |
Landing | 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"). |
LeakyIntegrateAndFire | Leaky integrate-and-fire with current and conductance input. |
LeakyIntegrateAndFireState | |
MembraneState | |
MgBlock | The fraction of NMDA receptors not blocked by magnesium at voltage v (mV). |
Model | Anything run steps: a neuron model, or a neuron with the synapses onto it (PointNeuron). |
ModulatedHebb | A Hebbian trace whose rate the network’s own activity sets through a neuromodulator, clipped. |
NeuronModel | One population of neurons, advanced one step at a time on the synaptic input it receives. |
OjaHebb | Oja’s rule: a Hebbian trace that each postsynaptic unit’s own activity bounds. |
Output | What a population sends in a step, and when within the step. |
PairSTDP | All-to-all pair STDP with soft or hard bounds: NEST’s stdp_synapse (Guetig et al. 2003). |
Plasticity | A rule that changes the weights of a projection’s edges with the spikes on either side. |
PointNeuron | A neuron model and the synapses onto it, by receptor, stepped together. |
PointNeuronState | A neuron’s state and its synapses’ by receptor name. |
PulseCell | RNeuralNet’s neuron: each step it outputs a function of the sum of what arrived, and empties the sum. |
RateCell | A leaky rate unit; its output is its activity h, a graded value. |
RateState | |
Receptor | A synapse model and how its output reaches the membrane. |
RecurrentCell | Feed a model’s output back to its own input through wiring, fixed or plastic. |
RecurrentState | |
Reset | 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. |
RetroactiveHebb | A Hebbian trace that a neuromodulator writes from an eligibility trace of recent coactivity. |
STDPTraces | |
Serial | Two models in series: each step, the first’s Output.value is the second’s input jump. |
Sparse | Units wired along the edges pre[e] -> post[e] with weight[e], among size units. |
StochasticRelease | Release that succeeds at random: each spike, each synapse releases with probability p. |
SynapseModel | |
SynapticInput | What a population receives in one step. |
Term | A current (amplitude + slope * s) * exp(-s / tau) over a step, s in [0, dt] ms from its start. |
TripletSTDP | All-to-all triplet STDP (Pfister and Gerstner 2006): NEST’s stdp_triplet_synapse. |
TripletTraces | |
TsodyksMarkram | Short-term depression and facilitation (Tsodyks and Markram 1997; Markram et al. 1998). |
TsodyksMarkramState | |
Wiring | Which units of a recurrent layer reach which, and with what weights: Dense or Sparse. |
decay | exp(-dt / tau): what a time constant tau leaves of a value after a step dt in the same unit. |
izhikevich_2003 | The cortical and thalamic classes of Izhikevich (2003), Figure 2, by their names there. |
izhikevich_2004 | 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"). |
run | Scan model over time-major inputs; return its output [T, ...] and the final state. |
ACTIVATIONS
Section titled “ACTIVATIONS”ACTIVATIONS = {'tanh': jnp.tanh, 'relu': jax.nn.relu, 'sigmoid': jax.nn.sigmoid}The activations a RateCell takes, by name.
IZHIKEVICH_2003
Section titled “IZHIKEVICH_2003”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
Section titled “IZHIKEVICH_2004”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
Section titled “RECEPTORS”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.
ALIFCell
Section titled “ALIFCell”class ALIFCellLIF 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.
| Field | Type | Default |
|---|---|---|
decay | jax.Array | float | |
adapt_decay | jax.Array | float | |
beta | jax.Array | float | 1.8 |
threshold | jax.Array | float | 1.0 |
reset | Reset | struct.field(pytree_node=False, default='subtract') |
surrogate | Surrogate | struct.field(pytree_node=False, default=ATan()) |
detach_reset | bool | struct.field(pytree_node=False, default=False) |
refractory | float | struct.field(pytree_node=False, default=0) |
ALIFCell.init_state
Section titled “ALIFCell.init_state”def init_state(shape: tuple[int, ...], dtype: jnp.dtype) -> ALIFStateALIFCell.step
Section titled “ALIFCell.step”def step(state: ALIFState, inputs: SynapticInput, dt: float) -> tuple[ALIFState, Output]ALIFCell.is_refractory
Section titled “ALIFCell.is_refractory”def is_refractory(state: ALIFState, dt: float) -> jax.ArrayALIFCell.after_threshold
Section titled “ALIFCell.after_threshold”def after_threshold(state: ALIFState, jump: jax.Array, fired: jax.Array) -> ALIFStateALIFState
Section titled “ALIFState”class ALIFState(NamedTuple)| Field | Type | Default |
|---|---|---|
v | jax.Array | |
a | jax.Array | |
r | jax.Array |
class AdExsparx.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) - wAt 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.
| Field | Type | Default |
|---|---|---|
c_m | jax.Array | float | 281.0 |
g_l | jax.Array | float | 30.0 |
e_l | jax.Array | float | -70.6 |
v_t | jax.Array | float | -50.4 |
delta_t | jax.Array | float | 2.0 |
v_peak | jax.Array | float | 0.0 |
v_reset | jax.Array | float | -60.0 |
a | jax.Array | float | 4.0 |
b | jax.Array | float | 80.5 |
tau_w | jax.Array | float | 144.0 |
t_ref | float | struct.field(pytree_node=False, default=0.0) |
substep | float | None | struct.field(pytree_node=False, default=0.01) |
reversal | Mapping[str, float] | |
gates | Mapping[str, MgBlock] | |
surrogate | Surrogate | struct.field(pytree_node=False, default=ATan()) |
AdEx.init_state
Section titled “AdEx.init_state”def init_state(shape: tuple[int, ...], dtype: jnp.dtype) -> AdExStateAdEx.step
Section titled “AdEx.step”def step(state: AdExState, inputs: SynapticInput, dt: float) -> tuple[AdExState, Output]AdEx.is_refractory
Section titled “AdEx.is_refractory”def is_refractory(state: AdExState, dt: float) -> jax.ArrayAdEx.after_threshold
Section titled “AdEx.after_threshold”def after_threshold(state: AdExState, jump: jax.Array, fired: jax.Array) -> AdExStateAdExState
Section titled “AdExState”class AdExState(NamedTuple)sparx.dynamics.neurons on GitHub
| Field | Type | Default |
|---|---|---|
v | jax.Array | |
w | jax.Array | |
refractory | jax.Array |
class Alphasparx.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.
| Field | Type | Default |
|---|---|---|
tau | jax.Array | float | 2.0 |
lands | Landing |
Alpha.init_state
Section titled “Alpha.init_state”def init_state(shape: tuple[int, ...], dtype: jnp.dtype) -> AlphaStateAlpha.output
Section titled “Alpha.output”def output(state: AlphaState) -> tuple[Term, ...]Alpha.step
Section titled “Alpha.step”def step(state: AlphaState, arriving: jax.Array | float, dt: float) -> AlphaStateArrivals
Section titled “Arrivals”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).
| Field | Type | Default |
|---|---|---|
current | jax.Array | float | 0.0 |
spikes | Mapping[str, jax.Array] | {} |
noise | jax.Array | None | None |
BernoulliCell
Section titled “BernoulliCell”class BernoulliCellA 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 0v[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.
| Field | Type | Default |
|---|---|---|
decay | jax.Array | float | |
threshold | jax.Array | float | 1.0 |
beta | jax.Array | float | 1.0 |
reset | Reset | struct.field(pytree_node=False, default='subtract') |
BernoulliCell.init_state
Section titled “BernoulliCell.init_state”def init_state(shape: tuple[int, ...], dtype: jnp.dtype) -> BernoulliStateBernoulliCell.step
Section titled “BernoulliCell.step”def step(state: BernoulliState, inputs: SynapticInput, dt: float) -> tuple[BernoulliState, Output]BernoulliCell.is_refractory
Section titled “BernoulliCell.is_refractory”def is_refractory(state: BernoulliState, dt: float) -> jax.ArrayBernoulliCell.after_threshold
Section titled “BernoulliCell.after_threshold”def after_threshold(state: BernoulliState, jump: jax.Array, fired: jax.Array) -> BernoulliStateBernoulliState
Section titled “BernoulliState”class BernoulliState(NamedTuple)| Field | Type | Default |
|---|---|---|
v | jax.Array | |
p | jax.Array |
BiExponential
Section titled “BiExponential”class BiExponentialsparx.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.
| Field | Type | Default |
|---|---|---|
tau_rise | float | struct.field(pytree_node=False, default=0.5) |
tau_decay | float | struct.field(pytree_node=False, default=5.0) |
lands | Landing | |
peak_factor | float |
BiExponential.init_state
Section titled “BiExponential.init_state”def init_state(shape: tuple[int, ...], dtype: jnp.dtype) -> BiExponentialStateBiExponential.output
Section titled “BiExponential.output”def output(state: BiExponentialState) -> tuple[Term, ...]BiExponential.step
Section titled “BiExponential.step”def step(state: BiExponentialState, arriving: jax.Array | float, dt: float) -> BiExponentialStateDecayingHebb
Section titled “DecayingHebb”class DecayingHebbA Hebbian trace that decays toward the latest coactivity at the rate eta.
keep = (1 - eta) ** dthebb[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).
| Field | Type | Default |
|---|---|---|
eta | jax.Array | float |
DecayingHebb.init_trace
Section titled “DecayingHebb.init_trace”def init_trace(wiring: Wiring, shape: tuple[int, ...], dtype: jnp.dtype) -> jax.ArrayDecayingHebb.hebb
Section titled “DecayingHebb.hebb”def hebb(trace: jax.Array) -> jax.ArrayDecayingHebb.update
Section titled “DecayingHebb.update”def update(trace: jax.Array, wiring: Wiring, pre: jax.Array, post: jax.Array, dt: float) -> jax.Arrayclass Deltasparx.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.
| Field | Type | Default |
|---|---|---|
after_threshold | bool | struct.field(pytree_node=False, default=False) |
lands | Landing |
Delta.init_state
Section titled “Delta.init_state”def init_state(shape: tuple[int, ...], dtype: jnp.dtype) -> tuple[()]Delta.output
Section titled “Delta.output”def output(state: tuple[()]) -> tuple[Term, ...]Delta.step
Section titled “Delta.step”def step(state: tuple[()], arriving: jax.Array | float, dt: float) -> tuple[()]class DenseEvery 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].
| Field | Type | Default |
|---|---|---|
weight | jax.Array | |
precision | PrecisionLike | struct.field(pytree_node=False, default=None) |
longest_delay | int |
Dense.connections
Section titled “Dense.connections”def connections(shape: tuple[int, ...]) -> tuple[int, ...]Dense.send
Section titled “Dense.send”def send(output: jax.Array, fast: jax.Array | None) -> jax.ArrayDense.presynaptic
Section titled “Dense.presynaptic”def presynaptic(x: jax.Array) -> jax.ArrayDense.postsynaptic
Section titled “Dense.postsynaptic”def postsynaptic(x: jax.Array) -> jax.ArrayDense.per_example
Section titled “Dense.per_example”def per_example(x: jax.Array) -> jax.ArrayDopamineSTDP
Section titled “DopamineSTDP”class DopamineSTDPsparx.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.
| Field | Type | Default |
|---|---|---|
modulator | str | struct.field(pytree_node=False, default='dopamine') |
tau_plus | jax.Array | float | 20.0 |
tau_minus | jax.Array | float | 20.0 |
tau_c | jax.Array | float | 1000.0 |
tau_n | float | struct.field(pytree_node=False, default=200.0) |
a_plus | jax.Array | float | 1.0 |
a_minus | jax.Array | float | 1.5 |
b | jax.Array | float | 0.0 |
w_min | jax.Array | float | 0.0 |
w_max | jax.Array | float | 200.0 |
DopamineSTDP.init_state
Section titled “DopamineSTDP.init_state”def init_state(pre: int, post: int, edges: int, dtype: jnp.dtype = jnp.float32) -> DopamineTracesDopamineSTDP.modulated_by
Section titled “DopamineSTDP.modulated_by”def modulated_by() -> Mapping[str, float | None]DopamineSTDP.step
Section titled “DopamineSTDP.step”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.
DopamineTraces
Section titled “DopamineTraces”class DopamineTraces(NamedTuple)sparx.dynamics.plasticity on GitHub
| Field | Type | Default |
|---|---|---|
pre | jax.Array | |
post | jax.Array | |
eligibility | jax.Array | |
dopamine | jax.Array |
EligibleHebb
Section titled “EligibleHebb”class EligibleHebb(NamedTuple)| Field | Type | Default |
|---|---|---|
hebb | jax.Array | |
eligibility | jax.Array |
Exponential
Section titled “Exponential”class Exponentialsparx.dynamics.synapses on GitHub
Jumps by the arriving weight and decays with tau (ms): NEST’s iaf_psc_exp and iaf_cond_exp.
| Field | Type | Default |
|---|---|---|
tau | jax.Array | float | 5.0 |
lands | Landing |
Exponential.init_state
Section titled “Exponential.init_state”def init_state(shape: tuple[int, ...], dtype: jnp.dtype) -> jax.ArrayExponential.output
Section titled “Exponential.output”def output(state: jax.Array) -> tuple[Term, ...]Exponential.step
Section titled “Exponential.step”def step(state: jax.Array, arriving: jax.Array | float, dt: float) -> jax.ArrayFastWeights
Section titled “FastWeights”class FastWeightsFast 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.
| Field | Type | Default |
|---|---|---|
alpha | jax.Array | float | |
rule | HebbianRule[Trace] | |
connections | jax.Array | None | None |
class Gap(NamedTuple)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.
| Field | Type | Default |
|---|---|---|
conductance | jax.Array | |
drive | Term |
Graded
Section titled “Graded”class Gradedsparx.dynamics.synapses on GitHub
Transmission proportional to a graded presynaptic value, with first-order kinetics.
tau ds/dt = sum_j w_j r_j - sr_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.
| Field | Type | Default |
|---|---|---|
tau | jax.Array | float | 5.0 |
lands | Landing |
Graded.init_state
Section titled “Graded.init_state”def init_state(shape: tuple[int, ...], dtype: jnp.dtype) -> GradedStateGraded.output
Section titled “Graded.output”def output(state: GradedState) -> tuple[Term, ...]Graded.step
Section titled “Graded.step”def step(state: GradedState, arriving: jax.Array | float, dt: float) -> GradedStateGradedPotential
Section titled “GradedPotential”class GradedPotentialsparx.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_mrelease(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.
| Field | Type | Default |
|---|---|---|
tau_m | jax.Array | float | 20.0 |
c_m | jax.Array | float | 200.0 |
e_l | jax.Array | float | -60.0 |
v_half | jax.Array | float | -35.0 |
slope | jax.Array | float | 5.0 |
i_e | jax.Array | float | 0.0 |
reversal | Mapping[str, float] | |
gates | Mapping[str, MgBlock] |
GradedPotential.release
Section titled “GradedPotential.release”def release(v: jax.Array) -> jax.ArrayThe release at voltage v, as a fraction of the maximum.
GradedPotential.init_state
Section titled “GradedPotential.init_state”def init_state(shape: tuple[int, ...], dtype: jnp.dtype) -> GradedPotentialStateGradedPotential.step
Section titled “GradedPotential.step”def step(state: GradedPotentialState, inputs: SynapticInput, dt: float) -> tuple[GradedPotentialState, Output]GradedPotential.is_refractory
Section titled “GradedPotential.is_refractory”def is_refractory(state: GradedPotentialState, dt: float) -> jax.ArrayGradedPotential.after_threshold
Section titled “GradedPotential.after_threshold”def after_threshold(state: GradedPotentialState, jump: jax.Array, fired: jax.Array) -> GradedPotentialStateGradedPotentialState
Section titled “GradedPotentialState”class GradedPotentialState(NamedTuple)sparx.dynamics.neurons on GitHub
| Field | Type | Default |
|---|---|---|
v | jax.Array |
GradedState
Section titled “GradedState”class GradedState(NamedTuple)sparx.dynamics.synapses on GitHub
| Field | Type | Default |
|---|---|---|
value | jax.Array | |
target | jax.Array |
HebbianRule
Section titled “HebbianRule”class HebbianRule(Protocol)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.
HebbianRule.init_trace
Section titled “HebbianRule.init_trace”def init_trace(wiring: Wiring, shape: tuple[int, ...], dtype: jnp.dtype) -> TraceHebbianRule.hebb
Section titled “HebbianRule.hebb”def hebb(trace: Trace) -> jax.ArrayHebbianRule.update
Section titled “HebbianRule.update”def update(trace: Trace, wiring: Wiring, pre: jax.Array, post: jax.Array, dt: float) -> TraceHodgkinHuxley
Section titled “HodgkinHuxley”class HodgkinHuxleysparx.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) + Idx/dt = alpha_x(v) (1 - x) - beta_x(v) x, x in m, h, nwith 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_mis 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’sexponential_euler(withsubstep=None, Brian2’s step exactly); stable and first order.
| Field | Type | Default |
|---|---|---|
c_m | jax.Array | float | 100.0 |
g_na | jax.Array | float | 12000.0 |
g_k | jax.Array | float | 3600.0 |
g_l | jax.Array | float | 30.0 |
e_na | jax.Array | float | 50.0 |
e_k | jax.Array | float | -77.0 |
e_l | jax.Array | float | -54.402 |
v_init | jax.Array | float | -65.0 |
v_spike | jax.Array | float | 0.0 |
t_ref | float | struct.field(pytree_node=False, default=2.0) |
scheme | Literal['strang', 'rk4', 'exponential_euler'] | struct.field(pytree_node=False, default='strang') |
substep | float | None | struct.field(pytree_node=False, default=0.01) |
reversal | Mapping[str, float] | |
gates | Mapping[str, MgBlock] | |
surrogate | Surrogate | struct.field(pytree_node=False, default=ATan()) |
HodgkinHuxley.rates
Section titled “HodgkinHuxley.rates”def rates(v: jax.Array) -> tuple[tuple[jax.Array, jax.Array], ...](alpha, beta) of m, h and n at v, in 1/ms.
HodgkinHuxley.init_state
Section titled “HodgkinHuxley.init_state”def init_state(shape: tuple[int, ...], dtype: jnp.dtype) -> HodgkinHuxleyStateHodgkinHuxley.step
Section titled “HodgkinHuxley.step”def step(state: HodgkinHuxleyState, inputs: SynapticInput, dt: float) -> tuple[HodgkinHuxleyState, Output]HodgkinHuxley.is_refractory
Section titled “HodgkinHuxley.is_refractory”def is_refractory(state: HodgkinHuxleyState, dt: float) -> jax.ArrayHodgkinHuxley.after_threshold
Section titled “HodgkinHuxley.after_threshold”def after_threshold(state: HodgkinHuxleyState, jump: jax.Array, fired: jax.Array) -> HodgkinHuxleyStateHodgkinHuxleyState
Section titled “HodgkinHuxleyState”class HodgkinHuxleyState(NamedTuple)sparx.dynamics.neurons on GitHub
| Field | Type | Default |
|---|---|---|
v | jax.Array | |
m | jax.Array | |
h | jax.Array | |
n | jax.Array | |
refractory | jax.Array |
Izhikevich
Section titled “Izhikevich”class Izhikevichsparx.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)).
| Field | Type | Default |
|---|---|---|
a | jax.Array | float | 0.02 |
b | jax.Array | float | 0.2 |
c | jax.Array | float | -65.0 |
d | jax.Array | float | 8.0 |
v_th | jax.Array | float | 30.0 |
v_init | jax.Array | float | -65.0 |
v_u | jax.Array | float | 0.0 |
u_decay | jax.Array | float | 1.0 |
quadratic | tuple[float, float, float] | struct.field(pytree_node=False, default=(0.04, 5.0, 140.0)) |
scheme | Literal['published', 'euler', 'semi_implicit'] | struct.field(pytree_node=False, default='published') |
order | Literal['izhikevich', 'nest'] | struct.field(pytree_node=False, default='izhikevich') |
reversal | Mapping[str, float] | |
gates | Mapping[str, MgBlock] | |
surrogate | Surrogate | struct.field(pytree_node=False, default=ATan()) |
Izhikevich.init_state
Section titled “Izhikevich.init_state”def init_state(shape: tuple[int, ...], dtype: jnp.dtype) -> IzhikevichStateIzhikevich.step
Section titled “Izhikevich.step”def step(state: IzhikevichState, inputs: SynapticInput, dt: float) -> tuple[IzhikevichState, Output]Izhikevich.is_refractory
Section titled “Izhikevich.is_refractory”def is_refractory(state: IzhikevichState, dt: float) -> jax.ArrayIzhikevich.after_threshold
Section titled “Izhikevich.after_threshold”def after_threshold(state: IzhikevichState, jump: jax.Array, fired: jax.Array) -> IzhikevichStateIzhikevichState
Section titled “IzhikevichState”class IzhikevichState(NamedTuple)sparx.dynamics.neurons on GitHub
| Field | Type | Default |
|---|---|---|
v | jax.Array | |
u | jax.Array |
LICell
Section titled “LICell”class LICellA 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.
| Field | Type | Default |
|---|---|---|
decay | jax.Array | float |
LICell.init_state
Section titled “LICell.init_state”def init_state(shape: tuple[int, ...], dtype: jnp.dtype) -> MembraneStateLICell.step
Section titled “LICell.step”def step(state: MembraneState, inputs: SynapticInput, dt: float) -> tuple[MembraneState, Output]LICell.is_refractory
Section titled “LICell.is_refractory”def is_refractory(state: MembraneState, dt: float) -> jax.ArrayLICell.after_threshold
Section titled “LICell.after_threshold”def after_threshold(state: MembraneState, jump: jax.Array, fired: jax.Array) -> MembraneStateLIFCell
Section titled “LIFCell”class LIFCellLeaky integrate-and-fire. decay=1 integrates without leak (IF).
| Field | Type | Default |
|---|---|---|
decay | jax.Array | float | |
threshold | jax.Array | float | 1.0 |
reset | Reset | struct.field(pytree_node=False, default='subtract') |
surrogate | Surrogate | struct.field(pytree_node=False, default=ATan()) |
detach_reset | bool | struct.field(pytree_node=False, default=False) |
LIFCell.init_state
Section titled “LIFCell.init_state”def init_state(shape: tuple[int, ...], dtype: jnp.dtype) -> MembraneStateLIFCell.step
Section titled “LIFCell.step”def step(state: MembraneState, inputs: SynapticInput, dt: float) -> tuple[MembraneState, Output]LIFCell.is_refractory
Section titled “LIFCell.is_refractory”def is_refractory(state: MembraneState, dt: float) -> jax.ArrayLIFCell.after_threshold
Section titled “LIFCell.after_threshold”def after_threshold(state: MembraneState, jump: jax.Array, fired: jax.Array) -> MembraneStateLanding
Section titled “Landing”Landingsparx.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").
LeakyIntegrateAndFire
Section titled “LeakyIntegrateAndFire”class LeakyIntegrateAndFiresparx.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_mAt 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.
| Field | Type | Default |
|---|---|---|
tau_m | jax.Array | float | 20.0 |
c_m | jax.Array | float | 200.0 |
e_l | jax.Array | float | -60.0 |
v_th | jax.Array | float | -50.0 |
v_reset | jax.Array | float | -60.0 |
t_ref | jax.Array | float | 5.0 |
i_e | jax.Array | float | 0.0 |
reversal | Mapping[str, float] | |
gates | Mapping[str, MgBlock] | |
surrogate | Surrogate | struct.field(pytree_node=False, default=ATan()) |
detach_reset | bool | struct.field(pytree_node=False, default=False) |
LeakyIntegrateAndFire.init_state
Section titled “LeakyIntegrateAndFire.init_state”def init_state(shape: tuple[int, ...], dtype: jnp.dtype) -> LeakyIntegrateAndFireStateLeakyIntegrateAndFire.step
Section titled “LeakyIntegrateAndFire.step”def step(state: LeakyIntegrateAndFireState, inputs: SynapticInput, dt: float) -> tuple[LeakyIntegrateAndFireState, Output]LeakyIntegrateAndFire.is_refractory
Section titled “LeakyIntegrateAndFire.is_refractory”def is_refractory(state: LeakyIntegrateAndFireState, dt: float) -> jax.ArrayLeakyIntegrateAndFire.after_threshold
Section titled “LeakyIntegrateAndFire.after_threshold”def after_threshold(state: LeakyIntegrateAndFireState, jump: jax.Array, fired: jax.Array) -> LeakyIntegrateAndFireStateLeakyIntegrateAndFireState
Section titled “LeakyIntegrateAndFireState”class LeakyIntegrateAndFireState(NamedTuple)sparx.dynamics.neurons on GitHub
| Field | Type | Default |
|---|---|---|
v | jax.Array | |
refractory | jax.Array |
MembraneState
Section titled “MembraneState”class MembraneState(NamedTuple)| Field | Type | Default |
|---|---|---|
v | jax.Array |
MgBlock
Section titled “MgBlock”class MgBlocksparx.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.
| Field | Type | Default |
|---|---|---|
mg | float | 1.0 |
MgBlock.__call__
Section titled “MgBlock.__call__”def __call__(v: jax.Array) -> jax.Arrayclass Model(Protocol)Anything run steps: a neuron model, or a neuron with the synapses onto it (PointNeuron).
| Field | Type | Default |
|---|---|---|
graded | bool |
Model.init_state
Section titled “Model.init_state”def init_state(shape: tuple[int, ...], dtype: jnp.dtype) -> StateThe population at rest, for inputs of per-step shape shape and dtype dtype.
Model.step
Section titled “Model.step”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.
ModulatedHebb
Section titled “ModulatedHebb”class ModulatedHebbA 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.
| Field | Type | Default |
|---|---|---|
modulator | jax.Array | |
modulator_bias | jax.Array | float | |
fanout | jax.Array | float | |
fanout_bias | jax.Array | float | |
clip | float | struct.field(pytree_node=False, default=2.0) |
ModulatedHebb.init_trace
Section titled “ModulatedHebb.init_trace”def init_trace(wiring: Wiring, shape: tuple[int, ...], dtype: jnp.dtype) -> jax.ArrayModulatedHebb.hebb
Section titled “ModulatedHebb.hebb”def hebb(trace: jax.Array) -> jax.ArrayModulatedHebb.update
Section titled “ModulatedHebb.update”def update(trace: jax.Array, wiring: Wiring, pre: jax.Array, post: jax.Array, dt: float) -> jax.ArrayNeuronModel
Section titled “NeuronModel”class NeuronModel(Model[State, SynapticInput], Protocol)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.
NeuronModel.is_refractory
Section titled “NeuronModel.is_refractory”def is_refractory(state: State, dt: float) -> jax.ArrayWhere the neuron cannot fire in the coming step of dt, as booleans; all False for a model
without refractoriness.
NeuronModel.after_threshold
Section titled “NeuronModel.after_threshold”def after_threshold(state: State, jump: jax.Array, fired: jax.Array) -> Statestate 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.
OjaHebb
Section titled “OjaHebb”class OjaHebbOja’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).
| Field | Type | Default |
|---|---|---|
eta | jax.Array | float |
OjaHebb.init_trace
Section titled “OjaHebb.init_trace”def init_trace(wiring: Wiring, shape: tuple[int, ...], dtype: jnp.dtype) -> jax.ArrayOjaHebb.hebb
Section titled “OjaHebb.hebb”def hebb(trace: jax.Array) -> jax.ArrayOjaHebb.update
Section titled “OjaHebb.update”def update(trace: jax.Array, wiring: Wiring, pre: jax.Array, post: jax.Array, dt: float) -> jax.ArrayOutput
Section titled “Output”class Output(NamedTuple)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.
| Field | Type | Default |
|---|---|---|
value | jax.Array | |
offset | jax.Array |
PairSTDP
Section titled “PairSTDP”class PairSTDPsparx.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).
| Field | Type | Default |
|---|---|---|
tau_plus | jax.Array | float | 20.0 |
tau_minus | jax.Array | float | 20.0 |
lambda_ | jax.Array | float | 0.01 |
alpha | jax.Array | float | 1.0 |
mu_plus | jax.Array | float | 1.0 |
mu_minus | jax.Array | float | 1.0 |
w_max | jax.Array | float | 100.0 |
PairSTDP.init_state
Section titled “PairSTDP.init_state”def init_state(pre: int, post: int, edges: int, dtype: jnp.dtype = jnp.float32) -> STDPTracesPairSTDP.modulated_by
Section titled “PairSTDP.modulated_by”def modulated_by() -> Mapping[str, float | None]PairSTDP.step
Section titled “PairSTDP.step”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.
Plasticity
Section titled “Plasticity”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.
Plasticity.init_state
Section titled “Plasticity.init_state”def init_state(pre: int, post: int, edges: int, dtype: jnp.dtype = jnp.float32) -> TracesThe traces with no spike yet.
Plasticity.modulated_by
Section titled “Plasticity.modulated_by”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.
Plasticity.step
Section titled “Plasticity.step”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).
PointNeuron
Section titled “PointNeuron”class PointNeuronsparx.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.
| Field | Type | Default |
|---|---|---|
neuron | NeuronModel[State] | |
receptors | Mapping[str, Receptor] | |
hold | Literal['mean', 'start'] | struct.field(pytree_node=False, default='mean') |
reset_synapses | bool | struct.field(pytree_node=False, default=False) |
freeze_synapses | bool | struct.field(pytree_node=False, default=False) |
graded | bool |
PointNeuron.init_state
Section titled “PointNeuron.init_state”def init_state(shape: tuple[int, ...], dtype: jnp.dtype) -> PointNeuronState[State]PointNeuron.landing
Section titled “PointNeuron.landing”def landing(where: Landing) -> list[str]The receptors whose arrivals land where.
PointNeuron.advance
Section titled “PointNeuron.advance”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.
PointNeuron.receive
Section titled “PointNeuron.receive”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.
PointNeuron.frozen
Section titled “PointNeuron.frozen”def frozen(state: PointNeuronState[State], dt: float) -> jax.ArrayWhere the neuron is refractory for the coming step, for freeze_synapses.
PointNeuron.delta
Section titled “PointNeuron.delta”def delta(arriving: Mapping[str, jax.Array]) -> jax.ArrayThe voltage jump (mV) of the weights arriving on receptors that land before the threshold.
PointNeuron.step
Section titled “PointNeuron.step”def step(state: PointNeuronState[State], inputs: Arrivals, dt: float) -> tuple[PointNeuronState[State], Output]PointNeuronState
Section titled “PointNeuronState”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.
| Field | Type | Default |
|---|---|---|
neuron | State | |
synapses | Mapping[str, object] |
PulseCell
Section titled “PulseCell”class PulseCellRNeuralNet’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 otherwisev[t] = 0The 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.
| Field | Type | Default |
|---|---|---|
threshold | jax.Array | float | 2.0 |
PulseCell.init_state
Section titled “PulseCell.init_state”def init_state(shape: tuple[int, ...], dtype: jnp.dtype) -> MembraneStatePulseCell.step
Section titled “PulseCell.step”def step(state: MembraneState, inputs: SynapticInput, dt: float) -> tuple[MembraneState, Output]PulseCell.is_refractory
Section titled “PulseCell.is_refractory”def is_refractory(state: MembraneState, dt: float) -> jax.ArrayPulseCell.after_threshold
Section titled “PulseCell.after_threshold”def after_threshold(state: MembraneState, jump: jax.Array, fired: jax.Array) -> MembraneStateRateCell
Section titled “RateCell”class RateCellA leaky rate unit; its output is its activity h, a graded value.
alpha = decay ** dth[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.
| Field | Type | Default |
|---|---|---|
decay | jax.Array | float | |
bias | jax.Array | float | 0.0 |
activation | Literal['tanh', 'relu', 'sigmoid'] | struct.field(pytree_node=False, default='tanh') |
RateCell.init_state
Section titled “RateCell.init_state”def init_state(shape: tuple[int, ...], dtype: jnp.dtype) -> RateStateRateCell.step
Section titled “RateCell.step”def step(state: RateState, inputs: SynapticInput, dt: float) -> tuple[RateState, Output]RateCell.is_refractory
Section titled “RateCell.is_refractory”def is_refractory(state: RateState, dt: float) -> jax.ArrayRateCell.after_threshold
Section titled “RateCell.after_threshold”def after_threshold(state: RateState, jump: jax.Array, fired: jax.Array) -> RateStateRateState
Section titled “RateState”class RateState(NamedTuple)| Field | Type | Default |
|---|---|---|
h | jax.Array |
Receptor
Section titled “Receptor”class Receptorsparx.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.
| Field | Type | Default |
|---|---|---|
synapse | SynapseModel | struct.field(default_factory=Exponential) |
kind | Literal['current', 'conductance'] | struct.field(pytree_node=False, default='current') |
RecurrentCell
Section titled “RecurrentCell”class RecurrentCellFeed 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).
| Field | Type | Default |
|---|---|---|
inner | NeuronModel[State] | |
wiring | Wiring | |
fast_weights | FastWeights[Trace] | None | None |
cut_gradient | bool | struct.field(pytree_node=False, default=False) |
graded | bool |
RecurrentCell.init_state
Section titled “RecurrentCell.init_state”def init_state(shape: tuple[int, ...], dtype: jnp.dtype) -> RecurrentState[State, Trace]RecurrentCell.step
Section titled “RecurrentCell.step”def step(state: RecurrentState[State, Trace], inputs: SynapticInput, dt: float) -> tuple[RecurrentState[State, Trace], Output]RecurrentCell.is_refractory
Section titled “RecurrentCell.is_refractory”def is_refractory(state: RecurrentState[State, Trace], dt: float) -> jax.ArrayRecurrentCell.after_threshold
Section titled “RecurrentCell.after_threshold”def after_threshold(state: RecurrentState[State, Trace], jump: jax.Array, fired: jax.Array) -> RecurrentState[State, Trace]RecurrentState
Section titled “RecurrentState”class RecurrentState(NamedTuple)| Field | Type | Default |
|---|---|---|
inner | State | |
output | jax.Array | |
arriving | jax.Array | |
trace | Trace | None | None |
history | jax.Array | None | None |
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.
RetroactiveHebb
Section titled “RetroactiveHebb”class RetroactiveHebbA 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) ** dteligibility[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.
| Field | Type | Default |
|---|---|---|
modulator | jax.Array | |
modulator_bias | jax.Array | float | |
eta | jax.Array | float | |
clip | float | struct.field(pytree_node=False, default=1.0) |
RetroactiveHebb.init_trace
Section titled “RetroactiveHebb.init_trace”def init_trace(wiring: Wiring, shape: tuple[int, ...], dtype: jnp.dtype) -> EligibleHebbRetroactiveHebb.hebb
Section titled “RetroactiveHebb.hebb”def hebb(trace: EligibleHebb) -> jax.ArrayRetroactiveHebb.update
Section titled “RetroactiveHebb.update”def update(trace: EligibleHebb, wiring: Wiring, pre: jax.Array, post: jax.Array, dt: float) -> EligibleHebbSTDPTraces
Section titled “STDPTraces”class STDPTraces(NamedTuple)sparx.dynamics.plasticity on GitHub
| Field | Type | Default |
|---|---|---|
pre | jax.Array | |
post | jax.Array |
Serial
Section titled “Serial”class SerialTwo 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.
| Field | Type | Default |
|---|---|---|
first | NeuronModel[First] | |
second | NeuronModel[Second] | |
graded | bool |
Serial.init_state
Section titled “Serial.init_state”def init_state(shape: tuple[int, ...], dtype: jnp.dtype) -> tuple[First, Second]Serial.step
Section titled “Serial.step”def step(state: tuple[First, Second], inputs: SynapticInput, dt: float) -> tuple[tuple[First, Second], Output]Serial.is_refractory
Section titled “Serial.is_refractory”def is_refractory(state: tuple[First, Second], dt: float) -> jax.ArraySerial.after_threshold
Section titled “Serial.after_threshold”def after_threshold(state: tuple[First, Second], jump: jax.Array, fired: jax.Array) -> tuple[First, Second]Sparse
Section titled “Sparse”class SparseUnits 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).
| Field | Type | Default |
|---|---|---|
pre | jax.Array | |
post | jax.Array | |
weight | jax.Array | |
size | int | struct.field(pytree_node=False) |
delay | jax.Array | None | None |
longest_delay | int | struct.field(pytree_node=False, default=1) |
Sparse.connections
Section titled “Sparse.connections”def connections(shape: tuple[int, ...]) -> tuple[int, ...]Sparse.send
Section titled “Sparse.send”def send(output: jax.Array, fast: jax.Array | None) -> jax.ArraySparse.presynaptic
Section titled “Sparse.presynaptic”def presynaptic(x: jax.Array) -> jax.ArraySparse.postsynaptic
Section titled “Sparse.postsynaptic”def postsynaptic(x: jax.Array) -> jax.ArraySparse.per_example
Section titled “Sparse.per_example”def per_example(x: jax.Array) -> jax.ArrayStochasticRelease
Section titled “StochasticRelease”class StochasticReleasesparx.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.
| Field | Type | Default |
|---|---|---|
p | jax.Array | float | 0.5 |
quantal | jax.Array | float | 1.0 |
StochasticRelease.transmit
Section titled “StochasticRelease.transmit”def transmit(key: jax.Array, weight: jax.Array, sent: jax.Array) -> jax.ArrayWhat edges of weight transmit for presynaptic values sent (0, 1 or an efficacy), one draw
each.
SynapseModel
Section titled “SynapseModel”class SynapseModel(Protocol)sparx.dynamics.synapses on GitHub
| Field | Type | Default |
|---|---|---|
lands | Landing |
SynapseModel.init_state
Section titled “SynapseModel.init_state”def init_state(shape: tuple[int, ...], dtype: jnp.dtype) -> StateSynapseModel.output
Section titled “SynapseModel.output”def output(state: State) -> tuple[Term, ...]SynapseModel.step
Section titled “SynapseModel.step”def step(state: State, arriving: jax.Array | float, dt: float) -> StateSynapticInput
Section titled “SynapticInput”class SynapticInputWhat 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.
| Field | Type | Default |
|---|---|---|
current | jax.Array | float | 0.0 |
currents | tuple[Term, ...] | () |
conductance | Mapping[str, jax.Array] | struct.field(default_factory=dict) |
jump | jax.Array | float | 0.0 |
gap | Gap | None | None |
noise | jax.Array | None | None |
waveforms | tuple[Term, ...] |
SynapticInput.current_at
Section titled “SynapticInput.current_at”def current_at(s: jax.Array | float) -> jax.ArrayThe total current (pA) s ms into the step, the drive of gap junctions included.
class Term(NamedTuple)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.
| Field | Type | Default |
|---|---|---|
amplitude | jax.Array | |
slope | jax.Array | |
tau | jax.Array | float |
Term.at
Section titled “Term.at”def at(s: jax.Array | float) -> jax.ArrayTerm.mean
Section titled “Term.mean”def mean(dt: float) -> jax.ArrayThe term’s average over [0, dt]: exact.
TripletSTDP
Section titled “TripletSTDP”class TripletSTDPsparx.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.
| Field | Type | Default |
|---|---|---|
tau_plus | jax.Array | float | 16.8 |
tau_x | jax.Array | float | 101.0 |
tau_minus | jax.Array | float | 33.7 |
tau_y | jax.Array | float | 125.0 |
a2_plus | jax.Array | float | 5e-10 |
a3_plus | jax.Array | float | 0.0062 |
a2_minus | jax.Array | float | 0.007 |
a3_minus | jax.Array | float | 0.00023 |
w_max | jax.Array | float | 100.0 |
TripletSTDP.init_state
Section titled “TripletSTDP.init_state”def init_state(pre: int, post: int, edges: int, dtype: jnp.dtype = jnp.float32) -> TripletTracesTripletSTDP.modulated_by
Section titled “TripletSTDP.modulated_by”def modulated_by() -> Mapping[str, float | None]TripletSTDP.step
Section titled “TripletSTDP.step”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.
TripletTraces
Section titled “TripletTraces”class TripletTraces(NamedTuple)sparx.dynamics.plasticity on GitHub
| Field | Type | Default |
|---|---|---|
pre | jax.Array | |
pre_triplet | jax.Array | |
post | jax.Array | |
post_triplet | jax.Array |
TsodyksMarkram
Section titled “TsodyksMarkram”class TsodyksMarkramsparx.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.
| Field | Type | Default |
|---|---|---|
U | jax.Array | float | 0.5 |
tau_rec | jax.Array | float | 800.0 |
tau_fac | jax.Array | float | 0.0 |
TsodyksMarkram.init_state
Section titled “TsodyksMarkram.init_state”def init_state(shape: tuple[int, ...], dtype: jnp.dtype = jnp.float32) -> TsodyksMarkramStateFully recovered and unfacilitated, per presynaptic neuron.
TsodyksMarkram.step
Section titled “TsodyksMarkram.step”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.
TsodyksMarkramState
Section titled “TsodyksMarkramState”class TsodyksMarkramState(NamedTuple)sparx.dynamics.plasticity on GitHub
| Field | Type | Default |
|---|---|---|
recovered | jax.Array | |
facilitation | jax.Array |
Wiring
Section titled “Wiring”class Wiring(Protocol)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.
| Field | Type | Default |
|---|---|---|
longest_delay | int |
Wiring.connections
Section titled “Wiring.connections”def connections(shape: tuple[int, ...]) -> tuple[int, ...]Wiring.send
Section titled “Wiring.send”def send(output: jax.Array, fast: jax.Array | None) -> jax.ArrayWiring.presynaptic
Section titled “Wiring.presynaptic”def presynaptic(x: jax.Array) -> jax.ArrayWiring.postsynaptic
Section titled “Wiring.postsynaptic”def postsynaptic(x: jax.Array) -> jax.ArrayWiring.per_example
Section titled “Wiring.per_example”def per_example(x: jax.Array) -> jax.Arraydef decay(tau: float, dt: float = 1.0) -> floatexp(-dt / tau): what a time constant tau leaves of a value after a step dt in the same unit.
izhikevich_2003
Section titled “izhikevich_2003”def izhikevich_2003(kind: str, **fields={}) -> Izhikevichsparx.dynamics.neurons on GitHub
The cortical and thalamic classes of Izhikevich (2003), Figure 2, by their names there.
izhikevich_2004
Section titled “izhikevich_2004”def izhikevich_2004(pattern: str, **fields={}) -> Izhikevichsparx.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]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.