PyTorch integration
GraphExecutor modules, execution plans, parameters, gradients and saved state.
Import GraphExecutor, GraphPlan and plan_graph from snnlab.sim.execution. GraphExecutor is a torch.nn.Module implementing the complete supported compiled graph, including multiple populations, connections and operations.
plan_graph
def plan_graph(graph: Mapping[str, Any]) -> GraphPlan: ...Checks backend capabilities, timebase, projection units and delays, then prepares population ordering and projection decay. It does not initialize model weights. Delays must be integer multiples of dt; recurrent/feedback population projections remain causal. Unsupported graph features raise ValueError.
GraphPlan holds graph, dt_ms, populations, projections, observables and outputs. Treat it as execution planning data rather than a replacement authoring API.
GraphExecutor
GraphExecutor(
plan,
*,
seed=0,
trainable_parameters=(),
surrogate_slope=5.0,
surrogate=None,
)plan is the result of plan_graph. seed seeds PyTorch initialization. trainable_parameters selects graph parameter IDs whose tensors require gradients; by default none are selected. surrogate_slope controls spike backward derivatives. When loading a training bundle, pass its resolved trainable IDs and compiled surrogate configuration explicitly. Construction calls torch.manual_seed, so it affects the global PyTorch RNG.
module = GraphExecutor(
plan_graph(bundle.graph),
seed=17,
trainable_parameters=bundle.training["resolved_parameters"]["trainable"],
surrogate_slope=bundle.training["surrogate"]["slope"],
)forward
result = module(
inputs,
diagnostics=True,
runtime_state=None,
interventions=(),
)inputs maps declared input names to tensors with (time, batch, channels) leading layout. Inputs must share time length, batch size and device; use a compatible module device/dtype. This call returns ExecutionResult, whose official output tensors retain gradients. Diagnostics and continuation state are detached. Runtime state continues a trajectory when explicitly passed; otherwise each call starts with fresh dynamic state.
interventions supplies supported runtime intervention mappings; it is a lower-level equivalent of the execution request’s inference interventions.
Parameter map
module.parameter_map() returns a mapping from graph names to live parameter tensors. It preserves familiar IDs such as w_out or input_to_E.weight, whereas PyTorch registration names use __ in place of periods. Do not assume runtime projection orientation equals graph orientation: graph matrices use (target, source), runtime projection tensors use (source, target) and are fan-in normalized at initialization.
Ordinary PyTorch layers
Register the executor and an ordinary nn.Sequential as child modules of your own nn.Module. Feed result.outputs into the head. Include both parameter sets in the optimizer, use tensor outputs directly for loss calculation, and enforce graph constraints after optimizer steps. Neither .numpy() nor detached diagnostics carries a differentiable feature path.
.parameters(), .to(...) and .state_dict() work through ordinary module registration. Save the complete wrapper’s state dictionary to include the SNN and external head. Reconstruct the same architecture before loading it. A PyTorch state dictionary is not the built-in snnlab training checkpoint format.
See PyTorch Integration for a complete hybrid training loop, epoch curves and reloaded inference.
Registered extensions and constraints
Pass a serialized custom surrogate through surrogate=... when constructing an executor directly. Import registered implementations before planning the graph. module.enforce_constraints() applies built-in and custom constraints after an external optimizer step. Extra state and current-based neurons/synapses participate in the same module and state-dictionary workflow. See Extensions.