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.
Contents
Section titled “Contents”| Name | |
|---|---|
SpikingClassification | A trained spiking classifier: class scores and predictions for a batch field. |
bound_call | The model’s __call__ with the keyword arguments kwargs bound, and train when it takes one, as apply’s method. |
SpikingClassification
Section titled “SpikingClassification”class SpikingClassificationA 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.
| Field | Type | Default |
|---|---|---|
model | nn.Module | |
variables | Variables | |
encoder | SpikeEncoder | |
readout | Readout | 'mean' |
call | Mapping[str, float] | field(default_factory=dict) |
SpikingClassification.logits
Section titled “SpikingClassification.logits”def logits(x: jax.Array, key: jax.Array | int = 0) -> jax.ArrayClass scores [B, classes] for a batch field [B, ...]; key drives a random encoder.
SpikingClassification.__call__
Section titled “SpikingClassification.__call__”def __call__(x: jax.Array, key: jax.Array | int = 0) -> jax.ArrayPredicted classes [B] for a batch field [B, ...].
SpikingClassification.from_run
Section titled “SpikingClassification.from_run”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) -> SpikingClassificationLoad 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.
bound_call
Section titled “bound_call”def bound_call(model: nn.Module, train: bool, kwargs: Mapping[str, jax.Array | float]) -> functools.partial[jax.Array]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.