snnlab
Examples

Train a two-class readout

Fit a tiny synthetic fixture, save its checkpoint and run inference.

Train a weighted spike-count readout to distinguish activity on input channel zero from activity on channel one. The four presentations are generated in memory, so no dataset download is needed. The script checks that cross-entropy decreases, saves the trained parameters and runs inference through the saved checkpoint.

Download the standalone script · View source

Run it

After installation, run from the repository root:

uv run python examples/train_readout.py

For the download, save the Python file and run it in an environment with snnlab installed. Every example is self-contained, uses CPU, and runs without a dataset download, Graphviz or FFmpeg.

Files, when produced, are written below artifacts/examples/train_readout/ relative to your working directory. A repeated run replaces that example’s outputs.

Complete code

"""Train a weighted spike-count readout on a tiny two-class fixture."""

from pathlib import Path

import torch

from snnlab import lang as snn
from snnlab.lang import training
from snnlab.sim.execution import ExecutionSpec, infer, train


def main(out=Path("artifacts/examples/train_readout")):
    net = snn.Network("two_class_readout", dt=1 * snn.ms)
    events = net.input(
        "events", shape=("time", "batch", 2), signal_type="spikes", unit="spike"
    )
    scores = snn.readouts.SpikeCount(source=events, classes=2, name="classifier")
    net.output("scores", scores)
    recipe = snn.TrainSpec(
        objectives=[training.CrossEntropy(prediction=scores, target="label")],
        parameter_groups=[
            training.ParameterGroup(
                [row["id"] for row in net.parameters],
                name="readout",
                lr=0.1,
            )
        ],
        optimizer=training.AdamW(weight_decay=0.0),
        presentation_duration=3 * snn.ms,
    )
    bundle = snn.compile(net, training=recipe, target="tools/snnsim")
    labels = torch.tensor([0, 1, 0, 1])
    inputs = torch.zeros(3, 4, 2)
    inputs[:, torch.arange(4), labels] = 1
    checkpoint = Path(out) / "checkpoint"
    result = train(
        ExecutionSpec(
            kind="train",
            executor="graph",
            graph=bundle.graph,
            training=bundle.training,
            inputs={"events": inputs},
            targets={"label": labels},
            seed=17,
            device="cpu",
            options={"updates": 10, "save_final_checkpoint": checkpoint},
        )
    )
    losses = [row["loss"] for row in result.metrics["updates"]]
    assert losses[-1] < losses[0]
    prediction = (
        infer(
            ExecutionSpec(
                kind="infer",
                executor="graph",
                graph=bundle.graph,
                inputs={"events": inputs},
                checkpoint=checkpoint,
                seed=17,
                device="cpu",
            )
        )
        .outputs["scores"]
        .argmax(dim=-1)
    )
    assert torch.equal(prediction, labels)
    print(f"Cross-entropy: {losses[0]:.4f} -> {losses[-1]:.4f}")
    print("Predictions:", prediction.tolist())
    print("Labels:     ", labels.tolist())


if __name__ == "__main__":
    main()

What to notice

  1. TrainSpec declares the objective, complete parameter scope, optimizer and physical presentation duration. Targets and sample tensors belong to the execution request.
  2. The linear readout starts at zero. AdamW learns signed projection weights; this is an operation rather than an AMPA conductance projection.
  3. updates=10 means ten optimizer updates. This fixture uses the complete four-presentation batch.
  4. Inference uses the checkpoint's parameters without resuming the optimizer. argmax converts the two scores into predicted class indices.

This intentionally simple task exercises the training interface. It has no hidden spiking population and therefore does not demonstrate surrogate-gradient learning through a spike threshold. Predictions are checked on the same fixture used for training; they are not a held-out accuracy measurement. For experiment design, add a separate evaluation protocol and dataset split.

Expected result

Cross-entropy decreases from approximately 0.6931.
Predictions: [0, 1, 0, 1]
Labels:      [0, 1, 0, 1]

SpikeCount, training recipes, train, infer.

On this page