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:
SpikingClassifierObjectiveencodes a batch into spikes, runs the network over time, and reads its outputs as class scores.ActivityFitObjectivefits a network’s spikes to recorded ones.EPropObjectivetrains a recurrent spiking layer with e-prop’s online gradients (sparx.learn.eprop) in place of backpropagation.PredictiveCodingObjectivetrains a stack of layers with predictive coding’s or PC-ALM’s local weight updates (sparx.learn.PredictiveCoding).RNeuralNetObjectiverewards anRNeuralNet’s choice and learns from the reward by reward diffusion or by REINFORCE (sparx.learn.RNeuralNet).
import optaxfrom dew import Field, Trainerfrom dew.data import Datasetfrom sparx.encode import RateEncoderfrom sparx.metrics import Accuracyfrom 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.
Contents
Section titled “Contents”| Name | |
|---|---|
ActivityFitObjective | Fit a network’s spiking to recorded spike trains: model fitting to recordings. |
EPropObjective | Classify recordings with a recurrent spiking network whose gradients are e-prop’s. |
PredictiveCodingObjective | Classify with a stack of layers whose weight updates are predictive coding’s or PC-ALM’s. |
RNeuralNetObjective | Reward an RNeuralNet’s choice among its output neurons after a sequence, and learn from the reward. |
RateBand | Keep each spiking neuron’s rate, in spikes per step, within [lower, upper], adding weight * sparx.rate_penalty. |
SpikingClassifierObjective | Classify a batch field with a spiking network. |
ActivityFitObjective
Section titled “ActivityFitObjective”class ActivityFitObjective(Objective[Ratio])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.
| Field | Type | Default |
|---|---|---|
shown | Mapping[str, Shown] | {'distance': Shown(better='lower')} |
ActivityFitObjective.fresh_variables
Section titled “ActivityFitObjective.fresh_variables”def fresh_variables(key: jax.Array, held: Variables | None) -> VariablesActivityFitObjective.loss
Section titled “ActivityFitObjective.loss”def loss(variables: Variables, batch: Batch, step: Step) -> tuple[Ratio, Aux]ActivityFitObjective.evaluate
Section titled “ActivityFitObjective.evaluate”def evaluate(params: Variables, batch: Batch, step: Step) -> TokenScoresEPropObjective
Section titled “EPropObjective”class EPropObjective(Objective[Ratio])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.
| Field | Type | Default |
|---|---|---|
shown | Mapping[str, Shown] | {'accuracy': Shown(better='higher', percent=True)} |
rule | EPropRule | rule |
cell | NeuronModel | model.neuron.bind({}).model(jnp.zeros((1, self.hidden))) |
EPropObjective.fresh_variables
Section titled “EPropObjective.fresh_variables”def fresh_variables(key: jax.Array, held: Variables | None) -> VariablesEPropObjective.loss
Section titled “EPropObjective.loss”def loss(variables: Variables, batch: Batch, step: Step) -> tuple[Ratio, Aux]EPropObjective.evaluate
Section titled “EPropObjective.evaluate”def evaluate(params: Variables, batch: Batch, step: Step) -> TokenScoresEPropObjective.task_record
Section titled “EPropObjective.task_record”def task_record() -> Mapping[str, JSON]The recordings as they are, scored by the readout averaged over time.
EPropObjective.build_task
Section titled “EPropObjective.build_task”def build_task(variables: Variables, processor: Processor | None | Omitted = OMITTED) -> SpikingClassificationThe trained network as a classifier over variables; it reads no text, so takes no
processor.
PredictiveCodingObjective
Section titled “PredictiveCodingObjective”class PredictiveCodingObjective(Objective[Ratio])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.
| Field | Type | Default |
|---|---|---|
shown | Mapping[str, Shown] | {'accuracy': Shown(better='higher', percent=True)} |
PredictiveCodingObjective.fresh_variables
Section titled “PredictiveCodingObjective.fresh_variables”def fresh_variables(key: jax.Array, held: Variables | None) -> VariablesPredictiveCodingObjective.loss
Section titled “PredictiveCodingObjective.loss”def loss(variables: Variables, batch: Batch, step: Step) -> tuple[Ratio, Aux]PredictiveCodingObjective.evaluate
Section titled “PredictiveCodingObjective.evaluate”def evaluate(params: Variables, batch: Batch, step: Step) -> TokenScoresRNeuralNetObjective
Section titled “RNeuralNetObjective”class RNeuralNetObjective(Objective[Ratio])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.
| Field | Type | Default |
|---|---|---|
shown | Mapping[str, Shown] | {'accuracy': Shown(better='higher', percent=True)} |
rule | RewardRule | rule |
RNeuralNetObjective.fresh_variables
Section titled “RNeuralNetObjective.fresh_variables”def fresh_variables(key: jax.Array, held: Variables | None) -> VariablesRNeuralNetObjective.loss
Section titled “RNeuralNetObjective.loss”def loss(variables: Variables, batch: Batch, step: Step) -> tuple[Ratio, Aux]RNeuralNetObjective.evaluate
Section titled “RNeuralNetObjective.evaluate”def evaluate(params: Variables, batch: Batch, step: Step) -> TokenScoresRateBand
Section titled “RateBand”class RateBandKeep each spiking neuron’s rate, in spikes per step, within [lower, upper], adding
weight * sparx.rate_penalty.
| Field | Type | Default |
|---|---|---|
lower | float | 0.0 |
upper | float | 1.0 |
weight | float | 1.0 |
SpikingClassifierObjective
Section titled “SpikingClassifierObjective”class SpikingClassifierObjective(Objective[Ratio])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).
| Field | Type | Default |
|---|---|---|
shown | Mapping[str, Shown] | {'accuracy': Shown(better='higher', percent=True)} |
readout | Readout | readout |
SpikingClassifierObjective.fresh_variables
Section titled “SpikingClassifierObjective.fresh_variables”def fresh_variables(key: jax.Array, held: Variables | None) -> VariablesSpikingClassifierObjective.loss
Section titled “SpikingClassifierObjective.loss”def loss(variables: Variables, batch: Batch, step: Step) -> tuple[Ratio, Aux]SpikingClassifierObjective.evaluate
Section titled “SpikingClassifierObjective.evaluate”def evaluate(params: Variables, batch: Batch, step: Step) -> TokenScoresSpikingClassifierObjective.task_record
Section titled “SpikingClassifierObjective.task_record”def task_record() -> Mapping[str, JSON]The encoder, readout, schedules and deployed arguments of the classifier.
SpikingClassifierObjective.build_task
Section titled “SpikingClassifierObjective.build_task”def build_task(variables: Variables, processor: Processor | None | Omitted = OMITTED) -> SpikingClassificationThe 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.