Skip to content
GitHub

sparx.tasks

Trained spiking networks as inference tasks, which dew.pipeline loads from a run.

SpikingClassification is what a run of SpikingClassifierObjective loads as (its saved_task). It follows dew’s SavedTask protocol, so dew.pipeline(run_dir, trust=("sparx",)) builds it in a fresh process from the run’s record and checkpoint, and objective.pipeline(state) builds it from a state still in memory.

Name
SpikingClassificationA trained spiking classifier: class scores and predictions for a batch field.
bound_callThe model’s __call__ with the keyword arguments kwargs bound, and train when it takes one, as apply’s method.
class SpikingClassification

sparx.tasks on GitHub

A trained spiking classifier: class scores and predictions for a batch field.

call holds the model keyword arguments it runs with (a trained DelayedDense’s sigma, say); 0 for sigma reads the rounded delays, the network as deployed.

FieldTypeDefault
modelnn.Module
variablesVariables
encoderSpikeEncoder
readoutReadout'mean'
callMapping[str, float]field(default_factory=dict)
def logits(x: jax.Array, key: jax.Array | int = 0) -> jax.Array

Class scores [B, classes] for a batch field [B, ...]; key drives a random encoder.

def __call__(x: jax.Array, key: jax.Array | int = 0) -> jax.Array

Predicted classes [B] for a batch field [B, ...].

def from_run(directory: str, ema: bool | None = None, step: int | str | None = None, mesh: MeshSpec | None = None, layout: Layout | None = None, dtype: DTypeLike | None = None, param_dtype: DTypeLike | None = None) -> SpikingClassification

Load the classifier a run of SpikingClassifierObjective in directory saved.

The model is the one its record names, over the selected checkpoint’s weights (the average when the run kept one, unless ema is False), placed on mesh under layout; the record and the weights are read from one pinned step (dew.inference.tasks.run_record). dtype replaces the recorded compute dtype and param_dtype the dtype the weights are read in.

def bound_call(model: nn.Module, train: bool, kwargs: Mapping[str, jax.Array | float]) -> functools.partial[jax.Array]

sparx.tasks on GitHub

The model’s __call__ with the keyword arguments kwargs bound, and train when it takes one, as apply’s method.

A model with BatchNorm or dropout takes train, as dew’s models do, so they know which pass they are in. A stack without them, a flax nn.Sequential of sparx layers say, need not.