Extensions
Versioned Python registrations for custom dynamics, weights, operations, training and dataset encoders.
snnlab.extensions connects user-defined Python implementations to portable graph specifications. Register a stable name such as my_lab.adaptive/v1, then use it in a lang.CustomNeuron or another custom declaration. Import the registering module before compiling, loading or executing a bundle that requires it.
Bundles contain names and JSON configuration, not Python code or pickled functions. manifest["extensions"] lists required definitions. Loading a bundle never automatically imports modules or downloads implementations. Changing an implementation’s semantics requires a new versioned name; the library cannot authenticate the source code behind a name.
Registration
Use the specialized registration functions below. Names must match a qualified identifier ending in /v and a positive version number. Duplicate category/name pairs raise ValueError; built-ins remain separate. extensions.definitions(category=None) returns registered immutable definition records for inspection.
All registrations accept optional validate(config). It should raise ValueError for invalid configuration. Configuration must serialize as JSON; use numeric values or explicit unit/value mappings. Implement differentiable runtime functions using PyTorch operations, and use the supplied device/dtype when initializing tensors.
Neurons
def register_neuron(name, step, *, initialize, input_unit='nA', state_units=None, validate=None):
...initialize(StateContext) returns a mapping containing voltage and any additional tensor state. refractory defaults to a zero-valued integer tensor. step(NeuronContext) returns (new_state, spikes). Return a new state mapping and avoid mutating previous tensors in place. State keys, tensor shapes and dtypes must remain stable during execution; spikes must be binary and have (batch, cells) shape.
input_unit is "nA" for current or "uS" for conductance. state_units maps additional state names to units; these exposed ports must have (batch, cells) floating tensors. Declare a diagnostic with net.expose(cells.state("adaptation"), name="adaptation"). The standard voltage/spikes/refractory ports cannot be redeclared.
class StateContext:
shape: tuple[int, int]
device: Any
dtype: Any
dt_ms: float
config: Mapping[str, Any]class NeuronContext:
state: Mapping[str, Any]
excitatory: Any
inhibitory: Any
dt_ms: float
config: Mapping[str, Any]
spike: Callable[[Any], Any]NeuronContext.spike(value) emits hard threshold spikes with the configured backward surrogate. Custom neurons choose their integration, reset and refractory behaviour; receiving a refractory state tensor does not implement those dynamics automatically. Extra hidden state can have other fixed shapes and dtypes, but must remain finite and on the execution device. All state is detached and saved at continuation boundaries.
neuron = lang.CustomNeuron("my_lab.adaptive/v1", tau_ms=20.0)See Customisation for a complete adaptive neuron definition and its response plots.
Synapses
def register_synapse(name, step, *, initialize=None, output_unit='nA', state_units=None, validate=None):
...class SynapseContext:
state: Mapping[str, Any]
drive: Any
dt_ms: float
config: Mapping[str, Any]initialize(StateContext) optionally returns synapse state with a required value tensor. Without it, value starts at zero. step(SynapseContext) receives delayed source activity already multiplied by the projection weights as drive. Return the complete updated state mapping, including value, which is added to the target’s excitatory or inhibitory input.
output_unit is "nA" or "uS"; it must match the target neuron. Additional state_units expose diagnostic ports through projection.state(name). The standard value, current and conductance names are reserved. Arbitrary auxiliary tensor state survives continuation; built-in delay routing and graph topology remain in effect.
synapse = lang.CustomSynapse("my_lab.synapse/v1", tau_ms=5.0)Initializers and constraints
def register_initializer(name, function, *, validate=None):
...def register_constraint(name, function, *, validate=None):
...Initializer callback: function(shape, config, *, device, dtype) -> Tensor. Return finite values with the exact requested shape/device/dtype. Use PyTorch’s seeded generator for reproducible draws, rather than hidden random state. Projection runtime shape is (source, target); ordinary custom-operation parameters retain their declared shape.
Constraint callback: function(parameter, config) -> Tensor. Return the constrained tensor with unchanged shape/device/dtype. Constraints apply after initialization and fan-in normalization, and after built-in training optimizer steps. An external PyTorch loop should call module.enforce_constraints() after its optimizer step. A masking constraint can enforce persistent zero connections; execution still uses dense matrices.
initializer = lang.CustomInitializer("my_lab.distribution/v1", scale=1.0)
constraint = lang.CustomConstraint("my_lab.mask/v1", mask=[[1, 0], [0, 1]])Operations and readouts
def register_operation(name, function, *, validate=None):
...def custom(definition: str, sources, *, name: str, shape, unit: str, parameters=(), signal_type='continuous', **config) -> Signal:
...Callback: function(sources, parameters, config) -> Tensor. sources is a tuple of input tensors; parameters maps declared graph parameter IDs to live tensors. Return the declared output shape on the execution device/dtype, retaining autograd. Explicitly list any trainable parameters through parameters=(ref, ...) and select them in the training recipe. Captured Python modules or tensors are not automatically registered model parameters.
Custom readouts can be operations, explicit graph populations, or ordinary external PyTorch layers. Custom architecture builders are normal Python functions that call the Network methods; they expand before serialization. This API does not add arbitrary connection routing or sparse storage backends.
Objectives and regularizers
def register_objective(name, function, *, classification=False, validate=None):
...def register_regularizer(name, function, *, validate=None):
...Objective callback: function(prediction, target, config) -> scalar Tensor. Use training.CustomObjective(definition, prediction=signal, target="name", weight=1.0, **config). Non-classification objectives accept finite real target tensors with a leading sample axis, including multidimensional regression targets. classification=True requires one-dimensional integer labels and enables argmax accuracy metrics.
Regularizer callback: function(signals, duration_s, config) -> scalar Tensor. Use training.CustomRegularizer(definition, signals=(signal, ...), strength=..., **config). Signals are supplied even when diagnostics are disabled. The executor multiplies objective and regularizer values by their declared weight or strength. Returned scalars must be finite, on the execution device and use the prediction/signal dtype.
Optimizers
def register_optimizer(name, factory, *, validate=None):
...Factory: factory(parameter_groups, config) -> torch.optim.Optimizer. Use training.CustomOptimizer(definition, **config). The executor supplies selected parameter groups and their learning rates, then calls standard zero_grad, step and constraint enforcement.
Checkpointing supports per-parameter tensor state and JSON scalar state, including stateless optimizers. Reconstruct the same recipe and registration before resuming. Scheduler state, arbitrary Python objects and custom global optimizer state are not automatically captured. External training loops can manage these separately.
Surrogate gradients
def register_surrogate(name, derivative, *, validate=None):
...Callback: derivative(value, config) -> Tensor, with the same shape/device/dtype as value. Use training.CustomSurrogate(definition, **config). The executor keeps a hard threshold in the forward pass and multiplies the backward gradient by this derivative. This applies to built-in spiking neurons and custom neurons using context.spike; a custom neuron that thresholds independently owns its own backward behaviour.
For direct PyTorch construction, pass the serialized custom specification as GraphExecutor(..., surrogate=recipe["surrogate"]).
Dataset encoders and inputs
def register_encoder(name, function, *, validate=None):
...Callback: function(arrays, selected, *, dt_ms, channels, config, seed) -> Tensor. Return binary (time, selected_samples, channels) spikes. Use DatasetEncoder(kind="custom", definition="my_lab.encoder/v1", config={...}, seed=...) inside DatasetSnapshotBinding. The resolved protocol retains the definition, configuration, seed and snapshot identity. Dataset snapshot labels remain integer labels.
For arbitrary generators, temporal patterns or preprocessing, a normal Python function can instead produce a tensor supplied through DenseArrayBinding. This also supports non-spike continuous inputs. No new generator interface is necessary.
Diagnostics and saved state
Official outputs, exposed diagnostics and .numpy() work for registered signal ports. Custom dynamic state and current traces use runtime-state artifact version 2; existing conductance-only state artifacts remain version 1 and are still readable. Continuation detaches previous state, so it does not perform backpropagation across saved trajectory boundaries.
The graph executor keeps runtime_state.currents separate from .conductances; registered hidden neuron/synapse state lives in .custom_state. Registered callbacks must not keep mutable trajectory state outside their supplied state mapping if they need reproducible continuation. GPU support depends on using device-compatible PyTorch operations in your implementation.