Save and resume training
Compare four uninterrupted updates with a two-plus-two checkpoint resume.
Run the same small readout training job in two ways: four uninterrupted updates, and two updates followed by a checkpoint reload and two more updates. The script compares final parameters, optimizer state and the resumed loss records exactly on CPU.
Download the standalone script · View source
Run it
After installation, run from the repository root:
uv run python examples/resume_training.pyFor 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/resume_training/ relative to your working directory. A repeated run replaces that example’s outputs.
Complete code
"""Compare uninterrupted CPU training with a saved and resumed run."""
from pathlib import Path
import torch
from snnlab import lang as snn
from snnlab.lang import training
from snnlab.sim.execution import ExecutionSpec, load_training_checkpoint, train
def main(out=Path("artifacts/examples/resume_training")):
net = snn.Network("resume_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="all",
lr=0.1,
)
],
optimizer=training.AdamW(weight_decay=0.0),
presentation_duration=3 * snn.ms,
)
bundle = snn.compile(net, training=recipe, target="tools/snnsim")
inputs = torch.zeros(3, 2, 2)
inputs[:, 0, 0] = 1
inputs[:, 1, 1] = 1
common = dict(
kind="train",
executor="graph",
graph=bundle.graph,
training=bundle.training,
inputs={"events": inputs},
targets={"label": torch.tensor([0, 1])},
device="cpu",
seed=17,
)
whole = train(ExecutionSpec(**common, options={"updates": 4}))
checkpoint = Path(out) / "checkpoint"
train(
ExecutionSpec(
**common,
options={
"updates": 2,
"save_final_checkpoint": checkpoint,
},
)
)
loaded = load_training_checkpoint(checkpoint)
assert loaded.completed_updates == 2
resumed = train(
ExecutionSpec(
**common,
checkpoint=checkpoint,
options={"updates": 2},
)
)
for name in whole.parameters:
torch.testing.assert_close(
resumed.parameters[name], whole.parameters[name], rtol=0, atol=0
)
for name, state in whole.optimizer_state.items():
for key, value in state.items():
torch.testing.assert_close(
resumed.optimizer_state[name][key], value, rtol=0, atol=0
)
assert resumed.metrics["updates"] == whole.metrics["updates"][2:]
print("Resumed from update:", resumed.metrics["resumed_from_update"])
print("Final parameters, optimizer state and resumed losses match exactly")
if __name__ == "__main__":
main()What to notice
- A checkpoint records completed optimizer updates as well as named parameters, optimizer state and authenticated execution metadata.
- When resuming,
updates=2requests two additional updates. The resumed update numbers are three and four. - The graph, training recipe, inputs, target binding and seed remain consistent across the two routes. Authentication rejects incompatible continuation requests.
- Training checkpoints and dynamic simulation runtime states solve different problems. A runtime state retains voltages, conductances, histories and refractory counters; this example resumes optimizer training.
The assertions establish equality for this deterministic CPU fixture on the current implementation. They do not establish equality between CPU and accelerator trajectories or between different numerical environments. The script is independent of the preceding training example and builds its own graph and fixture.
Expected result
Resumed from update: 2
Final parameters, optimizer state and resumed losses match exactly