Skip to content
GitHub

sparx.objectives

Spiking networks as dew objectives, trained by dew’s Trainer.

Dew splits a run into a model (a Flax module), an objective (parameters, loss, evaluation) and a trainer (mesh, compiled step, EMA, checkpoints, logging). A sparx network is a Flax module, so what sparx adds is the objectives:

  • SpikingClassifierObjective encodes a batch into spikes, runs the network over time, and reads its outputs as class scores.
  • ActivityFitObjective fits a network’s spikes to recorded ones.
  • EPropObjective trains a recurrent spiking layer with e-prop’s online gradients (sparx.learn.eprop) in place of backpropagation.
  • PredictiveCodingObjective trains a stack of layers with predictive coding’s or PC-ALM’s local weight updates (sparx.learn.PredictiveCoding).
  • RNeuralNetObjective rewards an RNeuralNet’s choice and learns from the reward by reward diffusion or by REINFORCE (sparx.learn.RNeuralNet).
import optax
from dew import Field, Trainer
from dew.data import Dataset
from sparx.encode import RateEncoder
from sparx.metrics import Accuracy
from sparx.objectives import SpikingClassifierObjective
objective = SpikingClassifierObjective(net, Field("image", (28, 28, 1)), RateEncoder(steps=16))
trainer = Trainer(objective, optax.adam(1e-3), key=0)
state = trainer.fit(Dataset.from_records({"image": x, "label": y}, batch=128),
steps=2000, metrics=[Accuracy()])

A run’s record names each by import path, sparx.objectives:EPropObjective, as dew records any objective. Each counts a batch’s rows through Objective.row_mean, so the repeats that fill a validation split’s last batch count for nothing in the loss.

Name
ActivityFitObjectiveFit a network’s spiking to recorded spike trains: model fitting to recordings.
EPropObjectiveClassify recordings with a recurrent spiking network whose gradients are e-prop’s.
PredictiveCodingObjectiveClassify with a stack of layers whose weight updates are predictive coding’s or PC-ALM’s.
RNeuralNetObjectiveReward an RNeuralNet’s choice among its output neurons after a sequence, and learn from the reward.
RateBandKeep each spiking neuron’s rate, in spikes per step, within [lower, upper], adding weight * sparx.rate_penalty.
SpikingClassifierObjectiveClassify a batch field with a spiking network.
class ActivityFitObjective(Objective[Ratio])

sparx.objectives on GitHub

Fit a network’s spiking to recorded spike trains: model fitting to recordings.

Each example holds a stimulus, [T, in] under stimulus.key, and the spikes recorded in response, [T, N] under recording; model maps the time-major stimulus [T, B, in] to spikes [T, B, N] of the recorded neurons (through surrogate gradients, so it trains).

loss="van_rossum" sums van Rossum’s (2001) distance over neurons and examples (sparx.losses.van_rossum, time constant tau ms, steps of dt ms). It compares spike timing at the scale tau, and its gradient moves spikes toward their recorded times. loss="psth" takes the squared difference of the trial-averaged rates over the batch, smoothed over window steps, for recordings repeated over trials where only the rate is reproducible, over the batch’s real rows. The loss is reported per example; metrics give the model’s and the recording’s mean rates (spikes per step). Evaluation returns one TokenScores row per example, its loss.

FieldTypeDefault
shownMapping[str, Shown]{'distance': Shown(better='lower')}
def fresh_variables(key: jax.Array, held: Variables | None) -> Variables
def loss(variables: Variables, batch: Batch, step: Step) -> tuple[Ratio, Aux]
def evaluate(params: Variables, batch: Batch, step: Step) -> TokenScores
class EPropObjective(Objective[Ratio])

sparx.objectives on GitHub

Classify recordings with a recurrent spiking network whose gradients are e-prop’s.

model is Bellec et al.’s (2020) network as a sparx.models.SpikingMLP with one recurrent hidden layer (recurrent=True) and dense synapses: an input synapse, a recurrent layer of model.neuron, and a synapse into a leaky readout, over the [T, channels] field sample. Each step’s cross entropy, divided by the steps, is summed over time, so the loss is the per-step mean, averaged over the batch’s rows. The class is the argmax of the readout averaged over time.

The loss runs sparx.learn.eprop on the model’s weights, which computes the loss and its gradients online in memory that does not grow with the recording, and hands the gradients to dew’s trainer as the loss’s own (Objective.with_gradients), so the trainer’s one gradient path applies them with its accumulation, sharding and logging. A bias is a synapse from an input that is always 1, and e-prop trains it as one. e-prop computes no gradient for a time constant, so a model that learns one is refused, as are delays, batch norm, dropout and plasticity rules. It computes as the model does: in the model’s dtype when it has one, float32 or wider, else in the widest of float32 and the parameters’ dtypes, at the model’s precision. rule="bptt" differentiates sparx.learn.bptt_loss instead. With rule="random" the feedback weights are drawn once by init and kept in the feedback collection, which no update touches.

The recurrent layer has no self-connections, as in Bellec et al.: init zeroes the recurrent matrix’s diagonal, which is masked in the forward pass and gets no gradient. The trained network is the SpikingMLP itself, so pipeline(state) and dew.pipeline(run_dir, trust=("sparx",)) load it as a sparx.tasks.SpikingClassification, which streams, serves and exports as any other. Evaluation runs the model and returns TokenScores for sparx.metrics.Accuracy.

FieldTypeDefault
shownMapping[str, Shown]{'accuracy': Shown(better='higher', percent=True)}
ruleEPropRulerule
cellNeuronModelmodel.neuron.bind({}).model(jnp.zeros((1, self.hidden)))
def fresh_variables(key: jax.Array, held: Variables | None) -> Variables
def loss(variables: Variables, batch: Batch, step: Step) -> tuple[Ratio, Aux]
def evaluate(params: Variables, batch: Batch, step: Step) -> TokenScores
def task_record() -> Mapping[str, JSON]

The recordings as they are, scored by the readout averaged over time.

def build_task(variables: Variables, processor: Processor | None | Omitted = OMITTED) -> SpikingClassification

The trained network as a classifier over variables; it reads no text, so takes no processor.

class PredictiveCodingObjective(Objective[Ratio])

sparx.objectives on GitHub

Classify with a stack of layers whose weight updates are predictive coding’s or PC-ALM’s.

model is a Flax nn.Sequential over the field sample, flattened per example, each element one layer of the stack (sparx.learn.residual_mlp builds Seely and Gould’s residual MLP). Its output scores classes classes against the one-hot labels under labels by half the squared error, the loss of their experiments, averaged over the batch’s rows. rule (sparx.learn.PredictiveCoding) relaxes the hidden activity from the forward pass, and its weight update, the energy’s gradient at the relaxed activity, goes to dew’s trainer as the loss’s own (Objective.with_gradients), so the trainer’s optimizer, accumulation and logging apply it. rule=None trains the same stack by backpropagation, their baseline. The reported loss is the forward pass’s whatever the rule, beside the batch’s accuracy. Evaluation returns TokenScores for sparx.metrics.Accuracy.

FieldTypeDefault
shownMapping[str, Shown]{'accuracy': Shown(better='higher', percent=True)}
def fresh_variables(key: jax.Array, held: Variables | None) -> Variables
def loss(variables: Variables, batch: Batch, step: Step) -> tuple[Ratio, Aux]
def evaluate(params: Variables, batch: Batch, step: Step) -> TokenScores
class RNeuralNetObjective(Objective[Ratio])

sparx.objectives on GitHub

Reward an RNeuralNet’s choice among its output neurons after a sequence, and learn from the reward.

The network (sparx.learn.RNeuralNet) runs over the field sample, one sequence of input values [T, inputs] per example, from nothing on its way. At the last tick it chooses one of its output neurons: in training it draws the choice from a softmax of beta times their outputs, in evaluation it takes the largest. A choice that matches the label under labels earns a reward of 1 and any other -1, the one scalar per example that every rule receives. The weights of all its connections are the parameters.

rule="first" or "all" is reward diffusion (sparx.learn.reward_diffusion, with discount): each example’s reward spreads from the feeder over its last tick’s outputs, and the mean of the examples’ weight changes goes to dew’s trainer as the loss’s gradient with its sign flipped (Objective.with_gradients), so optax.sgd(0.01 * batch) applies the original’s W_CONST of 0.01 for every reward.

Two rules add the changes attention-gated reinforcement learning (AGREL, Roelfsema and van Ooyen 2005) makes to a spread of reward. Both spread the reward prediction error R - b, with b the mean reward of the batch’s other examples, where AGREL takes an expansive function of its error. rule="gated" spreads it from the chosen output neuron instead of the feeder, along every path (paths="all", with discount), still shared by the softmax of absolute activity. rule="agrel" also sends it back through the connections’ weights and the neurons’ slopes along every delayed path, so each weight changes by R - b times the derivative of the chosen output’s last activity. AGREL’s feedback computes that update layer by layer in a layered network (Pozzi, Bohte and Roelfsema 2020); here it is differentiated through the network over time. rule="reinforce" is REINFORCE (Williams 1992) on the choice, -(R - b) log p(choice) per example, differentiated the same way. Whatever the rule, the reported loss is the negative mean reward of the choices drawn, and evaluation returns TokenScores for sparx.metrics.Accuracy.

FieldTypeDefault
shownMapping[str, Shown]{'accuracy': Shown(better='higher', percent=True)}
ruleRewardRulerule
def fresh_variables(key: jax.Array, held: Variables | None) -> Variables
def loss(variables: Variables, batch: Batch, step: Step) -> tuple[Ratio, Aux]
def evaluate(params: Variables, batch: Batch, step: Step) -> TokenScores
class RateBand

sparx.objectives on GitHub

Keep each spiking neuron’s rate, in spikes per step, within [lower, upper], adding weight * sparx.rate_penalty.

FieldTypeDefault
lowerfloat0.0
upperfloat1.0
weightfloat1.0
class SpikingClassifierObjective(Objective[Ratio])

sparx.objectives on GitHub

Classify a batch field with a spiking network.

model maps the encoder’s time-major input [T, B, ...] to outputs [T, B, classes]: spikes, or the membrane of a sparx.nn.LI readout. sample names the field and its per-example shape; labels names the integer class field. encoder is one of sparx.encode’s, which read a uint8 field of intensities as x / 255.

A model that takes train, as dew’s models do, gets True in the loss, with rngs={"dropout": ...}, and False in evaluation; one without BatchNorm or dropout need not take it, so an nn.Sequential stack of sparx layers trains as it is, and exports to NIR (sparx.nir.to_nir). A model that keeps batch_stats (BatchNorm) has them updated by the loss. ema_decay keeps an exponential moving average of the parameters, which evaluate scores when the trainer passes it.

The loss is the mean cross entropy of the readout over the batch’s rows, plus the rates penalty when given. Metrics report the batch accuracy, each spiking layer’s mean firing rate (rate/<layer>) and the penalty. Evaluation returns TokenScores with one row per example (loss, weight, whether the argmax is the label), which sparx.metrics.Accuracy reads.

schedules names model keyword arguments that follow one of dew’s schedules over schedule_steps steps, as a learning rate does: {"sigma": Linear(peak=7.5, end=0.5)} anneals a sparx.nn.DelayedDense, {"masking": ...} a MaskedPSN. The value is read at Step.step, the count of accepted microbatches, in the loss and in evaluation alike, and a schedule’s every holds each value for that many steps, as a torch scheduler stepped once an epoch does. deployed holds model keyword arguments that evaluation and the trained classifier run with in place of the schedules’ values, such as {"sigma": 0} to score every delay rounded to a whole step, the network as deployed and as SNN-delays evaluates it.

Parameter groups with optimizers of their own, as SNN-delays trains its delay positions, are the trainer’s: OptimConfig(param_groups=...), each a dew ParamGroup with its own schedule, momentum, weight decay and bounds.

Every checkpoint records the model, encoder, readout, schedules and deployed arguments (task_record), and the run loads back as a sparx.tasks.SpikingClassification (dew.pipeline(run_dir, trust=("sparx",)), or pipeline(state) after training).

FieldTypeDefault
shownMapping[str, Shown]{'accuracy': Shown(better='higher', percent=True)}
readoutReadoutreadout

SpikingClassifierObjective.fresh_variables

Section titled “SpikingClassifierObjective.fresh_variables”
def fresh_variables(key: jax.Array, held: Variables | None) -> Variables
def loss(variables: Variables, batch: Batch, step: Step) -> tuple[Ratio, Aux]
def evaluate(params: Variables, batch: Batch, step: Step) -> TokenScores
def task_record() -> Mapping[str, JSON]

The encoder, readout, schedules and deployed arguments of the classifier.

def build_task(variables: Variables, processor: Processor | None | Omitted = OMITTED) -> SpikingClassification

The trained classifier over variables.

The schedules stand at their final values, with the deployed arguments over them. A spiking classifier reads no text, so it takes no processor.