snnlab
API referencesnnlab.sim

snnlab.sim.execution

Complete declared API of the execution module, with signatures, data fields, validation and source.

Back to sim reference

Graph-native entry points and typed data contracts. Explicitly select executor=graph for typed graph execution; legacy requests route through the established CLI handlers. Inputs use time, batch, feature axes; events use zero-based integer steps. Checkpoints and runtime states authenticate different contracts. All filesystem loaders treat bundles as data rather than importing authoring code.

The signatures, defaults, fields, docstrings and implementation excerpts below are generated from the Python source. Annotations are shown as declared; unannotated means the source supplies no type annotation. These pages document callable surfaces, including legacy support utilities, without promising backend support for every declaration.

DenseArrayBinding

View source

Source docstring:

One concrete dense tensor resolved against a named graph input.

Class decorators: dataclass(frozen=True).

Dataclass constructor parameters. Factory defaults are shown as field declarations; omit these arguments to create fresh values per instance:

DenseArrayBinding(input_id: str, value: torch.Tensor, source: Mapping[str, Any] = field(default_factory=dict))

Declared fields, including fields inherited from local data classes:

FieldAnnotationDefaultMeaning
input_idstrrequiredExact declared graph input id.
valuetorch.TensorrequiredStored member of this data contract; see the class docstring and serialization methods.
sourceMapping[str, Any]field(default_factory=dict)Stored member of this data contract; see the class docstring and serialization methods.
Complete class implementation
class DenseArrayBinding:
    """One concrete dense tensor resolved against a named graph input."""

    input_id: str
    value: torch.Tensor
    source: Mapping[str, Any] = field(default_factory=dict)

TargetArrayBinding

View source

Source docstring:

One concrete target vector resolved against a named training target.

Class decorators: dataclass(frozen=True).

Dataclass constructor parameters. Factory defaults are shown as field declarations; omit these arguments to create fresh values per instance:

TargetArrayBinding(target_id: str, value: torch.Tensor, source: Mapping[str, Any] = field(default_factory=dict))

Declared fields, including fields inherited from local data classes:

FieldAnnotationDefaultMeaning
target_idstrrequiredExact named training target id.
valuetorch.TensorrequiredStored member of this data contract; see the class docstring and serialization methods.
sourceMapping[str, Any]field(default_factory=dict)Stored member of this data contract; see the class docstring and serialization methods.
Complete class implementation
class TargetArrayBinding:
    """One concrete target vector resolved against a named training target."""

    target_id: str
    value: torch.Tensor
    source: Mapping[str, Any] = field(default_factory=dict)

EventStreamBinding

View source

Source docstring:

Sparse binary spike events resolved against one named graph input.

Class decorators: dataclass(frozen=True).

Dataclass constructor parameters. Factory defaults are shown as field declarations; omit these arguments to create fresh values per instance:

EventStreamBinding(input_id: str, steps: torch.Tensor, batches: torch.Tensor, channels: torch.Tensor, steps_count: int, batch_size: int, source: Mapping[str, Any] = field(default_factory=dict))

Declared fields, including fields inherited from local data classes:

FieldAnnotationDefaultMeaning
input_idstrrequiredExact declared graph input id.
stepstorch.TensorrequiredStored member of this data contract; see the class docstring and serialization methods.
batchestorch.TensorrequiredStored member of this data contract; see the class docstring and serialization methods.
channelstorch.TensorrequiredStored member of this data contract; see the class docstring and serialization methods.
steps_countintrequiredNumber of simulation timesteps in an event or Poisson binding.
batch_sizeintrequiredNumber of presentations in a binding or mini-batch.
sourceMapping[str, Any]field(default_factory=dict)Stored member of this data contract; see the class docstring and serialization methods.
Complete class implementation
class EventStreamBinding:
    """Sparse binary spike events resolved against one named graph input."""

    input_id: str
    steps: torch.Tensor
    batches: torch.Tensor
    channels: torch.Tensor
    steps_count: int
    batch_size: int
    source: Mapping[str, Any] = field(default_factory=dict)

PoissonInputBinding

View source

Source docstring:

Generated Bernoulli-discretised Poisson spikes for one graph input.

Class decorators: dataclass(frozen=True).

Dataclass constructor parameters. Factory defaults are shown as field declarations; omit these arguments to create fresh values per instance:

PoissonInputBinding(input_id: str, steps_count: int, batch_size: int, rates_hz: Sequence[float], seed: int, categorical: bool = False)

Declared fields, including fields inherited from local data classes:

FieldAnnotationDefaultMeaning
input_idstrrequiredExact declared graph input id.
steps_countintrequiredNumber of simulation timesteps in an event or Poisson binding.
batch_sizeintrequiredNumber of presentations in a binding or mini-batch.
rates_hzSequence[float]requiredConfigured input rates in spikes per second.
seedintrequiredSeed controlling this operation’s random stream.
categoricalboolFalseStored member of this data contract; see the class docstring and serialization methods.
Complete class implementation
class PoissonInputBinding:
    """Generated Bernoulli-discretised Poisson spikes for one graph input."""

    input_id: str
    steps_count: int
    batch_size: int
    rates_hz: Sequence[float]
    seed: int
    categorical: bool = False

DatasetEncoder

View source

Source docstring:

Portable standard encoding recipe for an immutable dataset snapshot.

Class decorators: dataclass(frozen=True).

Dataclass constructor parameters. Factory defaults are shown as field declarations; omit these arguments to create fresh values per instance:

DatasetEncoder(kind: Literal['rate_poisson', 'prebinned_spikes', 'event_bin'], duration_ms: float | None = None, max_rate_hz: float | None = None, seed: int = 0)

Declared fields, including fields inherited from local data classes:

FieldAnnotationDefaultMeaning
kindLiteral['rate_poisson', 'prebinned_spikes', 'event_bin']requiredStored member of this data contract; see the class docstring and serialization methods.
duration_msfloat | NoneNonePhysical presentation duration in milliseconds.
max_rate_hzfloat | NoneNoneMaximum rate in spikes per second for encoded input.
seedint0Seed controlling this operation’s random stream.
Complete class implementation
class DatasetEncoder:
    """Portable standard encoding recipe for an immutable dataset snapshot."""

    kind: Literal["rate_poisson", "prebinned_spikes", "event_bin"]
    duration_ms: float | None = None
    max_rate_hz: float | None = None
    seed: int = 0

DatasetSnapshotBinding

View source

Source docstring:

Bind one digest-identified NPZ snapshot to a graph input and target.

Class decorators: dataclass(frozen=True).

Dataclass constructor parameters. Factory defaults are shown as field declarations; omit these arguments to create fresh values per instance:

DatasetSnapshotBinding(path: Path, input_id: str, dataset_id: str, split: str, encoder: DatasetEncoder, target_id: str | None = None, feature_key: str = 'features', label_key: str = 'labels', sample_cap: int | None = None, shuffle: bool = False, order_seed: int = 0)

Declared fields, including fields inherited from local data classes:

FieldAnnotationDefaultMeaning
pathPathrequiredFilesystem source or destination path, as described below.
input_idstrrequiredExact declared graph input id.
dataset_idstrrequiredStored member of this data contract; see the class docstring and serialization methods.
splitstrrequiredStored member of this data contract; see the class docstring and serialization methods.
encoderDatasetEncoderrequiredStored member of this data contract; see the class docstring and serialization methods.
target_idstr | NoneNoneExact named training target id.
feature_keystr'features'Stored member of this data contract; see the class docstring and serialization methods.
label_keystr'labels'Stored member of this data contract; see the class docstring and serialization methods.
sample_capint | NoneNoneStored member of this data contract; see the class docstring and serialization methods.
shuffleboolFalseStored member of this data contract; see the class docstring and serialization methods.
order_seedint0Seed for deterministic sample ordering.
Complete class implementation
class DatasetSnapshotBinding:
    """Bind one digest-identified NPZ snapshot to a graph input and target."""

    path: Path
    input_id: str
    dataset_id: str
    split: str
    encoder: DatasetEncoder
    target_id: str | None = None
    feature_key: str = "features"
    label_key: str = "labels"
    sample_cap: int | None = None
    shuffle: bool = False
    order_seed: int = 0

ResolvedDenseInputs

View source

Validated input tensors together with the execution protocol that authenticates their source, shape, dtype and timing.

Class decorators: dataclass(frozen=True).

Dataclass constructor parameters. Factory defaults are shown as field declarations; omit these arguments to create fresh values per instance:

ResolvedDenseInputs(tensors: Mapping[str, torch.Tensor], protocol: Mapping[str, Any])

Declared fields, including fields inherited from local data classes:

FieldAnnotationDefaultMeaning
tensorsMapping[str, torch.Tensor]requiredStored member of this data contract; see the class docstring and serialization methods.
protocolMapping[str, Any]requiredStored member of this data contract; see the class docstring and serialization methods.
Complete class implementation
class ResolvedDenseInputs:
    tensors: Mapping[str, torch.Tensor]
    protocol: Mapping[str, Any]

ExecutionSpec

View source

Complete request record. Supply either a bundle path or graph mapping, named input bindings, external targets, seed/device, recording selection and optional checkpoint/runtime state. executor defaults to legacy: select graph explicitly for graph-native methods. options carries validated execution overrides rather than arbitrary graph mutation.

Class decorators: dataclass(frozen=True).

Dataclass constructor parameters. Factory defaults are shown as field declarations; omit these arguments to create fresh values per instance:

ExecutionSpec(kind: RequestKind, executor: ExecutorName = 'legacy', bundle: Path | None = None, graph: Mapping[str, Any] | None = None, inputs: Mapping[str, torch.Tensor] = field(default_factory=dict), input_bindings: Sequence[DenseArrayBinding] = field(default_factory=tuple), event_bindings: Sequence[EventStreamBinding] = field(default_factory=tuple), poisson_bindings: Sequence[PoissonInputBinding] = field(default_factory=tuple), dataset_binding: DatasetSnapshotBinding | None = None, protocol: Mapping[str, Any] = field(default_factory=dict), training: Mapping[str, Any] | None = None, targets: Mapping[str, torch.Tensor] = field(default_factory=dict), target_bindings: Sequence[TargetArrayBinding] = field(default_factory=tuple), seed: int = 0, device: str = 'auto', recording: RecordingProfile = 'full', recording_fields: Sequence[str] | None = None, checkpoint: Path | None = None, runtime_state: GraphRuntimeState | None = None, options: Mapping[str, Any] = field(default_factory=dict))

Declared fields, including fields inherited from local data classes:

FieldAnnotationDefaultMeaning
kindRequestKindrequiredStored member of this data contract; see the class docstring and serialization methods.
executorExecutorName'legacy'Execution route: legacy or graph; graph must be selected explicitly.
bundlePath | NoneNoneCompiled data bundle or bundle path, as annotated.
graphMapping[str, Any] | NoneNoneSerialized graph mapping.
inputsMapping[str, torch.Tensor]field(default_factory=dict)Input tensors keyed by graph input id.
input_bindingsSequence[DenseArrayBinding]field(default_factory=tuple)Named dense input bindings, resolved against graph contracts.
event_bindingsSequence[EventStreamBinding]field(default_factory=tuple)Named sparse event bindings with zero-based step/batch/channel coordinates.
poisson_bindingsSequence[PoissonInputBinding]field(default_factory=tuple)Declared generated Poisson input bindings.
dataset_bindingDatasetSnapshotBinding | NoneNoneStored member of this data contract; see the class docstring and serialization methods.
protocolMapping[str, Any]field(default_factory=dict)Stored member of this data contract; see the class docstring and serialization methods.
trainingMapping[str, Any] | NoneNoneAuthored training declaration or serialized training recipe.
targetsMapping[str, torch.Tensor]field(default_factory=dict)Named integer targets or target objects, according to this contract.
target_bindingsSequence[TargetArrayBinding]field(default_factory=tuple)Stored member of this data contract; see the class docstring and serialization methods.
seedint0Seed controlling this operation’s random stream.
devicestr'auto'Requested or resolved tensor execution device.
recordingRecordingProfile'full'Retained Recording or recording selection, as annotated.
recording_fieldsSequence[str] | NoneNoneExplicit field names to retain.
checkpointPath | NoneNoneTraining checkpoint record or authenticated checkpoint path.
runtime_stateGraphRuntimeState | NoneNoneDynamic graph state for causal continuation; distinct from training checkpoints.
optionsMapping[str, Any]field(default_factory=dict)Stored member of this data contract; see the class docstring and serialization methods.
Complete class implementation
class ExecutionSpec:
    kind: RequestKind
    executor: ExecutorName = "legacy"
    bundle: Path | None = None
    graph: Mapping[str, Any] | None = None
    inputs: Mapping[str, torch.Tensor] = field(default_factory=dict)
    input_bindings: Sequence[DenseArrayBinding] = field(default_factory=tuple)
    event_bindings: Sequence[EventStreamBinding] = field(default_factory=tuple)
    poisson_bindings: Sequence[PoissonInputBinding] = field(default_factory=tuple)
    dataset_binding: DatasetSnapshotBinding | None = None
    protocol: Mapping[str, Any] = field(default_factory=dict)
    training: Mapping[str, Any] | None = None
    targets: Mapping[str, torch.Tensor] = field(default_factory=dict)
    target_bindings: Sequence[TargetArrayBinding] = field(default_factory=tuple)
    seed: int = 0
    device: str = "auto"
    recording: RecordingProfile = "full"
    recording_fields: Sequence[str] | None = None
    checkpoint: Path | None = None
    runtime_state: GraphRuntimeState | None = None
    options: Mapping[str, Any] = field(default_factory=dict)

ExecutionResult

View source

Named outputs, recordings, parameters, gradients, optimizer state and metrics, plus optional selected/training checkpoints and runtime state. Tensor dictionaries are keyed by stable graph ids; do not assume legacy tensor names.

Class decorators: dataclass.

Dataclass constructor parameters. Factory defaults are shown as field declarations; omit these arguments to create fresh values per instance:

ExecutionResult(executor: ExecutorName, outputs: dict[str, torch.Tensor] = field(default_factory=dict), recordings: dict[str, torch.Tensor] = field(default_factory=dict), parameters: dict[str, torch.Tensor] = field(default_factory=dict), gradients: dict[str, torch.Tensor] = field(default_factory=dict), optimizer_state: dict[str, Any] = field(default_factory=dict), training_checkpoint: TrainingCheckpoint | None = None, selected_checkpoint: TrainingCheckpoint | None = None, final_state: dict[str, torch.Tensor] = field(default_factory=dict), runtime_state: GraphRuntimeState | None = None, metrics: dict[str, Any] = field(default_factory=dict), model: nn.Module | None = None)

Declared fields, including fields inherited from local data classes:

FieldAnnotationDefaultMeaning
executorExecutorNamerequiredExecution route: legacy or graph; graph must be selected explicitly.
outputsdict[str, torch.Tensor]field(default_factory=dict)Stored member of this data contract; see the class docstring and serialization methods.
recordingsdict[str, torch.Tensor]field(default_factory=dict)Stored member of this data contract; see the class docstring and serialization methods.
parametersdict[str, torch.Tensor]field(default_factory=dict)Stored member of this data contract; see the class docstring and serialization methods.
gradientsdict[str, torch.Tensor]field(default_factory=dict)Stored member of this data contract; see the class docstring and serialization methods.
optimizer_statedict[str, Any]field(default_factory=dict)Stored member of this data contract; see the class docstring and serialization methods.
training_checkpointTrainingCheckpoint | NoneNoneStored member of this data contract; see the class docstring and serialization methods.
selected_checkpointTrainingCheckpoint | NoneNoneStored member of this data contract; see the class docstring and serialization methods.
final_statedict[str, torch.Tensor]field(default_factory=dict)Stored member of this data contract; see the class docstring and serialization methods.
runtime_stateGraphRuntimeState | NoneNoneDynamic graph state for causal continuation; distinct from training checkpoints.
metricsdict[str, Any]field(default_factory=dict)Stored member of this data contract; see the class docstring and serialization methods.
modelnn.Module | NoneNoneStored member of this data contract; see the class docstring and serialization methods.
Complete class implementation
class ExecutionResult:
    executor: ExecutorName
    outputs: dict[str, torch.Tensor] = field(default_factory=dict)
    recordings: dict[str, torch.Tensor] = field(default_factory=dict)
    parameters: dict[str, torch.Tensor] = field(default_factory=dict)
    gradients: dict[str, torch.Tensor] = field(default_factory=dict)
    optimizer_state: dict[str, Any] = field(default_factory=dict)
    training_checkpoint: TrainingCheckpoint | None = None
    selected_checkpoint: TrainingCheckpoint | None = None
    final_state: dict[str, torch.Tensor] = field(default_factory=dict)
    runtime_state: GraphRuntimeState | None = None
    metrics: dict[str, Any] = field(default_factory=dict)
    model: nn.Module | None = None

CapabilityIssue

View source

Element-level diagnostic identifying a required capability and why it is unsupported.

Class decorators: dataclass(frozen=True).

Dataclass constructor parameters. Factory defaults are shown as field declarations; omit these arguments to create fresh values per instance:

CapabilityIssue(element: str, capability: str, message: str)

Declared fields, including fields inherited from local data classes:

FieldAnnotationDefaultMeaning
elementstrrequiredStored member of this data contract; see the class docstring and serialization methods.
capabilitystrrequiredStored member of this data contract; see the class docstring and serialization methods.
messagestrrequiredStored member of this data contract; see the class docstring and serialization methods.
Complete class implementation
class CapabilityIssue:
    element: str
    capability: str
    message: str

TrainingCheckpoint

View source

Portable graph training state with authenticated graph/recipe identity, update coordinate, input protocol, initialized parameter metadata, parameter/optimizer tensors and backend-specific RNG state. It is distinct from dynamic GraphRuntimeState.

Class decorators: dataclass.

Dataclass constructor parameters. Factory defaults are shown as field declarations; omit these arguments to create fresh values per instance:

TrainingCheckpoint(graph_digest: str, training_digest: str, completed_updates: int, execution_protocol: Mapping[str, Any], initialization: Mapping[str, Any], parameters: dict[str, torch.Tensor], optimizer_state: dict[str, dict[str, Any]], rng_state: torch.Tensor, rng_backend: str = 'cpu', accelerator_rng_states: dict[str, torch.Tensor] = field(default_factory=dict), data_state: Mapping[str, Any] = field(default_factory=dict), selected_loss: float | None = None)

Declared fields, including fields inherited from local data classes:

FieldAnnotationDefaultMeaning
graph_digeststrrequiredStored member of this data contract; see the class docstring and serialization methods.
training_digeststrrequiredStored member of this data contract; see the class docstring and serialization methods.
completed_updatesintrequiredStored member of this data contract; see the class docstring and serialization methods.
execution_protocolMapping[str, Any]requiredStored member of this data contract; see the class docstring and serialization methods.
initializationMapping[str, Any]requiredStored member of this data contract; see the class docstring and serialization methods.
parametersdict[str, torch.Tensor]requiredStored member of this data contract; see the class docstring and serialization methods.
optimizer_statedict[str, dict[str, Any]]requiredStored member of this data contract; see the class docstring and serialization methods.
rng_statetorch.TensorrequiredStored member of this data contract; see the class docstring and serialization methods.
rng_backendstr'cpu'Stored member of this data contract; see the class docstring and serialization methods.
accelerator_rng_statesdict[str, torch.Tensor]field(default_factory=dict)Stored member of this data contract; see the class docstring and serialization methods.
data_stateMapping[str, Any]field(default_factory=dict)Stored member of this data contract; see the class docstring and serialization methods.
selected_lossfloat | NoneNoneStored member of this data contract; see the class docstring and serialization methods.
Complete class implementation
class TrainingCheckpoint:
    graph_digest: str
    training_digest: str
    completed_updates: int
    execution_protocol: Mapping[str, Any]
    initialization: Mapping[str, Any]
    parameters: dict[str, torch.Tensor]
    optimizer_state: dict[str, dict[str, Any]]
    rng_state: torch.Tensor
    rng_backend: str = "cpu"
    accelerator_rng_states: dict[str, torch.Tensor] = field(default_factory=dict)
    data_state: Mapping[str, Any] = field(default_factory=dict)
    selected_loss: float | None = None

ParameterInterchange

View source

Named graph parameter tensors plus versioned provenance from import/export of the supported legacy parameter map.

Class decorators: dataclass(frozen=True).

Dataclass constructor parameters. Factory defaults are shown as field declarations; omit these arguments to create fresh values per instance:

ParameterInterchange(parameters: dict[str, torch.Tensor], provenance: Mapping[str, Any])

Declared fields, including fields inherited from local data classes:

FieldAnnotationDefaultMeaning
parametersdict[str, torch.Tensor]requiredStored member of this data contract; see the class docstring and serialization methods.
provenanceMapping[str, Any]requiredStored member of this data contract; see the class docstring and serialization methods.
Complete class implementation
class ParameterInterchange:
    parameters: dict[str, torch.Tensor]
    provenance: Mapping[str, Any]

load_dense_array_bindings

View source

def load_dense_array_bindings(path: str | Path, graph: Mapping[str, Any]) -> tuple[DenseArrayBinding, ...]

Source docstring:

Load a replayable NPY/NPZ file without weakening named-input semantics.
ParameterAnnotationDefaultMeaning
pathstr | PathrequiredFilesystem source or destination path, as described below.
graphMapping[str, Any]requiredSerialized graph mapping.

Return annotation: tuple[DenseArrayBinding, ...].

Return expressions (branch-dependent; names refer to the linked implementation):

(DenseArrayBinding(input_ids[0], torch.as_tensor(loaded), {**source_base, 'array': None}),)
tuple((DenseArrayBinding(input_id, torch.as_tensor(value), {**source_base, 'array': keys[input_id]}) for input_id, value in arrays.items()))

Explicit exceptions in this implementation; called helpers may raise additional errors:

Explicit exception expression
ValueError(f'a dense NPY can bind only a graph with exactly one input; graph inputs are {input_ids}')
Implementation
def load_dense_array_bindings(
    path: str | Path, graph: Mapping[str, Any]
) -> tuple[DenseArrayBinding, ...]:
    """Load a replayable NPY/NPZ file without weakening named-input semantics."""
    source_path = Path(path)
    digest = "sha256:" + hashlib.sha256(source_path.read_bytes()).hexdigest()
    loaded = np.load(source_path, allow_pickle=False)
    input_ids = [row["id"] for row in graph.get("inputs", [])]
    source_base = {
        "kind": "file",
        "path": str(source_path),
        "digest": digest,
    }
    if isinstance(loaded, np.ndarray):
        if len(input_ids) != 1:
            raise ValueError(
                "a dense NPY can bind only a graph with exactly one input; "
                f"graph inputs are {input_ids}"
            )
        return (
            DenseArrayBinding(
                input_ids[0],
                torch.as_tensor(loaded),
                {**source_base, "array": None},
            ),
        )
    try:
        arrays = {key: loaded[key] for key in loaded.files}
    finally:
        loaded.close()
    if set(arrays) == {"input_spikes"} and len(input_ids) == 1:
        arrays = {input_ids[0]: arrays["input_spikes"]}
        keys = {input_ids[0]: "input_spikes"}
    else:
        keys = {key: key for key in arrays}
    return tuple(
        DenseArrayBinding(
            input_id,
            torch.as_tensor(value),
            {**source_base, "array": keys[input_id]},
        )
        for input_id, value in arrays.items()
    )

load_target_array_bindings

View source

def load_target_array_bindings(path: str | Path, training: Mapping[str, Any]) -> tuple[TargetArrayBinding, ...]

Source docstring:

Load NPY/NPZ class targets against the recipe's named objectives.
ParameterAnnotationDefaultMeaning
pathstr | PathrequiredFilesystem source or destination path, as described below.
trainingMapping[str, Any]requiredAuthored training declaration or serialized training recipe.

Return annotation: tuple[TargetArrayBinding, ...].

Return expressions (branch-dependent; names refer to the linked implementation):

tuple((TargetArrayBinding(target_id, torch.as_tensor(arrays[target_id]), {**source, 'array': keys[target_id]}) for target_id in target_ids))

Explicit exceptions in this implementation; called helpers may raise additional errors:

Explicit exception expression
ValueError(f'target array ids do not match recipe targets; expected={target_ids}, got={sorted(arrays)}')
ValueError(f'a target NPY requires exactly one recipe target; targets are {target_ids}')
Implementation
def load_target_array_bindings(
    path: str | Path, training: Mapping[str, Any]
) -> tuple[TargetArrayBinding, ...]:
    """Load NPY/NPZ class targets against the recipe's named objectives."""
    source_path = Path(path)
    digest = "sha256:" + hashlib.sha256(source_path.read_bytes()).hexdigest()
    target_ids = sorted(
        {row["target"] for row in training.get("objectives", []) if "target" in row}
    )
    loaded = np.load(source_path, allow_pickle=False)
    source = {"kind": "file", "path": str(source_path), "digest": digest}
    if isinstance(loaded, np.ndarray):
        if len(target_ids) != 1:
            raise ValueError(
                f"a target NPY requires exactly one recipe target; targets are {target_ids}"
            )
        arrays = {target_ids[0]: loaded}
        keys = {target_ids[0]: None}
    else:
        try:
            arrays = {key: loaded[key] for key in loaded.files}
        finally:
            loaded.close()
        keys = {key: key for key in arrays}
    if set(arrays) != set(target_ids):
        raise ValueError(
            f"target array ids do not match recipe targets; expected={target_ids}, got={sorted(arrays)}"
        )
    return tuple(
        TargetArrayBinding(
            target_id,
            torch.as_tensor(arrays[target_id]),
            {**source, "array": keys[target_id]},
        )
        for target_id in target_ids
    )

resolve_target_array_bindings

View source

def resolve_target_array_bindings(training: Mapping[str, Any], *, bindings: Sequence[TargetArrayBinding]=(), targets: Mapping[str, torch.Tensor] | None=None, sample_count: int, device: str | torch.device='cpu') -> tuple[dict[str, torch.Tensor], list[dict[str, Any]]]

Source docstring:

Validate named one-dimensional integer targets and retain their provenance.
ParameterAnnotationDefaultMeaning
trainingMapping[str, Any]requiredAuthored training declaration or serialized training recipe.
bindingsSequence[TargetArrayBinding]()Defined by the source contract and implementation below.
targetsMapping[str, torch.Tensor] | NoneNoneNamed integer targets or target objects, according to this contract.
sample_countintrequiredDefined by the source contract and implementation below.
devicestr | torch.device'cpu'Requested or resolved tensor execution device.

Return annotation: tuple[dict[str, torch.Tensor], list[dict[str, Any]]].

Return expressions (branch-dependent; names refer to the linked implementation):

(resolved, rows)

Explicit exceptions in this implementation; called helpers may raise additional errors:

Explicit exception expression
ValueError('provide target bindings or target tensors, not both')
ValueError('duplicate target array binding')
ValueError(f'target ids do not match recipe; missing={sorted(expected - set(by_name))}, unexpected={sorted(set(by_name) - expected)}')
ValueError(f'training target {target_id} expected shape [{sample_count}], got {list(value.shape)}')
ValueError(f'training target {target_id} must use an integer dtype')
Implementation
def resolve_target_array_bindings(
    training: Mapping[str, Any],
    *,
    bindings: Sequence[TargetArrayBinding] = (),
    targets: Mapping[str, torch.Tensor] | None = None,
    sample_count: int,
    device: str | torch.device = "cpu",
) -> tuple[dict[str, torch.Tensor], list[dict[str, Any]]]:
    """Validate named one-dimensional integer targets and retain their provenance."""
    if bindings and targets:
        raise ValueError("provide target bindings or target tensors, not both")
    if not bindings:
        bindings = tuple(
            TargetArrayBinding(name, value, {"kind": "memory"})
            for name, value in (targets or {}).items()
        )
    expected = {
        row["target"] for row in training.get("objectives", []) if "target" in row
    }
    by_name = {binding.target_id: binding for binding in bindings}
    if len(by_name) != len(bindings):
        raise ValueError("duplicate target array binding")
    if set(by_name) != expected:
        raise ValueError(
            f"target ids do not match recipe; missing={sorted(expected - set(by_name))}, unexpected={sorted(set(by_name) - expected)}"
        )
    resolved = {}
    rows = []
    integer_dtypes = {torch.int8, torch.int16, torch.int32, torch.int64, torch.uint8}
    for target_id in sorted(by_name):
        binding = by_name[target_id]
        value = binding.value
        if value.ndim != 1 or value.shape[0] != sample_count:
            raise ValueError(
                f"training target {target_id} expected shape [{sample_count}], got {list(value.shape)}"
            )
        if value.dtype not in integer_dtypes:
            raise ValueError(f"training target {target_id} must use an integer dtype")
        resolved[target_id] = value.to(device=device, dtype=torch.long)
        rows.append(
            {
                "target_id": target_id,
                "shape": list(value.shape),
                "dtype": str(value.dtype).removeprefix("torch."),
                "source": dict(binding.source),
            }
        )
    return resolved, rows

load_event_stream_bindings

View source

def load_event_stream_bindings(path: str | Path, graph: Mapping[str, Any]) -> tuple[EventStreamBinding, ...]

Source docstring:

Load named sparse spike coordinates from a replayable NPZ file.
ParameterAnnotationDefaultMeaning
pathstr | PathrequiredFilesystem source or destination path, as described below.
graphMapping[str, Any]requiredSerialized graph mapping.

Return annotation: tuple[EventStreamBinding, ...].

Return expressions (branch-dependent; names refer to the linked implementation):

tuple(bindings)

Explicit exceptions in this implementation; called helpers may raise additional errors:

Explicit exception expression
ValueError('event-stream replay requires an NPZ file')
ValueError(f'event-stream NPZ keys do not match graph inputs; expected={sorted(expected)}, got={sorted(arrays)}')
ValueError(f'event input {input_id} steps_count and batch_size must be scalars')
ValueError(f'event input {input_id} steps_count and batch_size must use integer dtypes')
Implementation
def load_event_stream_bindings(
    path: str | Path, graph: Mapping[str, Any]
) -> tuple[EventStreamBinding, ...]:
    """Load named sparse spike coordinates from a replayable NPZ file."""
    source_path = Path(path)
    if source_path.suffix.lower() != ".npz":
        raise ValueError("event-stream replay requires an NPZ file")
    digest = "sha256:" + hashlib.sha256(source_path.read_bytes()).hexdigest()
    loaded = np.load(source_path, allow_pickle=False)
    try:
        arrays = {key: loaded[key] for key in loaded.files}
    finally:
        loaded.close()
    input_ids = [row["id"] for row in graph.get("inputs", [])]
    fields = ("steps", "batches", "channels", "steps_count", "batch_size")
    plain = set(fields)
    use_plain = len(input_ids) == 1 and set(arrays) == plain
    expected = (
        plain
        if use_plain
        else {f"{input_id}.{field}" for input_id in input_ids for field in fields}
    )
    if set(arrays) != expected:
        raise ValueError(
            "event-stream NPZ keys do not match graph inputs; "
            f"expected={sorted(expected)}, got={sorted(arrays)}"
        )
    source_base = {"kind": "file", "path": str(source_path), "digest": digest}
    bindings = []
    for input_id in input_ids:

        def key(field: str) -> str:
            return field if use_plain else f"{input_id}.{field}"

        steps_count_value = np.asarray(arrays[key("steps_count")])
        batch_size_value = np.asarray(arrays[key("batch_size")])
        if steps_count_value.size != 1 or batch_size_value.size != 1:
            raise ValueError(
                f"event input {input_id} steps_count and batch_size must be scalars"
            )
        if not np.issubdtype(steps_count_value.dtype, np.integer) or not np.issubdtype(
            batch_size_value.dtype, np.integer
        ):
            raise ValueError(
                f"event input {input_id} steps_count and batch_size must use integer dtypes"
            )
        bindings.append(
            EventStreamBinding(
                input_id=input_id,
                steps=torch.as_tensor(arrays[key("steps")]),
                batches=torch.as_tensor(arrays[key("batches")]),
                channels=torch.as_tensor(arrays[key("channels")]),
                steps_count=int(steps_count_value.item()),
                batch_size=int(batch_size_value.item()),
                source={
                    **source_base,
                    "arrays": {field: key(field) for field in fields},
                },
            )
        )
    return tuple(bindings)

resolve_dense_array_bindings

View source

def resolve_dense_array_bindings(graph: Mapping[str, Any], *, bindings: Sequence[DenseArrayBinding]=(), inputs: Mapping[str, torch.Tensor] | None=None, device: str | torch.device='cpu', seed: int=0, protocol: Mapping[str, Any] | None=None) -> ResolvedDenseInputs

Source docstring:

Validate dense arrays, resolve symbolic axes, and freeze run provenance.
ParameterAnnotationDefaultMeaning
graphMapping[str, Any]requiredSerialized graph mapping.
bindingsSequence[DenseArrayBinding]()Defined by the source contract and implementation below.
inputsMapping[str, torch.Tensor] | NoneNoneInput tensors keyed by graph input id.
devicestr | torch.device'cpu'Requested or resolved tensor execution device.
seedint0Seed controlling this operation’s random stream.
protocolMapping[str, Any] | NoneNoneDefined by the source contract and implementation below.

Return annotation: ResolvedDenseInputs.

Return expressions (branch-dependent; names refer to the linked implementation):

ResolvedDenseInputs(tensors=resolved, protocol=execution_protocol)

Explicit exceptions in this implementation; called helpers may raise additional errors:

Explicit exception expression
ValueError('provide dense input bindings or input tensors, not both')
ValueError(f'dense input ids do not match graph inputs; missing={missing}, unexpected={unexpected}')
ValueError(f'execution protocol cannot override reserved fields {sorted(reserved & supplied.keys())}')
ValueError(f'duplicate dense binding for input {binding.input_id}')
ValueError(f"input {input_id} dense binding requires declared shape beginning with ['time', 'batch']")
ValueError(f'input {input_id} rank expected {len(declared)}, got {value.ndim}')
ValueError(f'input {input_id} trailing shape expected {expected_tail}, got {tuple(value.shape[2:])}')
ValueError(f'input {input_id} contains non-finite values')
ValueError(f'input {input_id} spike values must be boolean or zero/one')
ValueError(f'execution protocol must be JSON-serializable: {exc}')
ValueError(f'input {input_id} leading shape expected {leading_shape}, got {current_leading}')
ValueError(f'input {input_id} mask values must be boolean or zero/one')
Implementation
def resolve_dense_array_bindings(
    graph: Mapping[str, Any],
    *,
    bindings: Sequence[DenseArrayBinding] = (),
    inputs: Mapping[str, torch.Tensor] | None = None,
    device: str | torch.device = "cpu",
    seed: int = 0,
    protocol: Mapping[str, Any] | None = None,
) -> ResolvedDenseInputs:
    """Validate dense arrays, resolve symbolic axes, and freeze run provenance."""
    if bindings and inputs:
        raise ValueError("provide dense input bindings or input tensors, not both")
    if not bindings:
        bindings = tuple(
            DenseArrayBinding(name, value, {"kind": "memory"})
            for name, value in (inputs or {}).items()
        )
    specs = {row["id"]: row for row in graph.get("inputs", [])}
    by_name: dict[str, DenseArrayBinding] = {}
    for binding in bindings:
        if binding.input_id in by_name:
            raise ValueError(f"duplicate dense binding for input {binding.input_id}")
        by_name[binding.input_id] = binding
    if set(by_name) != set(specs):
        missing = sorted(set(specs) - set(by_name))
        unexpected = sorted(set(by_name) - set(specs))
        raise ValueError(
            f"dense input ids do not match graph inputs; missing={missing}, unexpected={unexpected}"
        )

    resolved: dict[str, torch.Tensor] = {}
    rows: list[dict[str, Any]] = []
    leading_shape: tuple[int, int] | None = None
    masks: list[str] = []
    for input_id in sorted(specs):
        spec = specs[input_id]
        binding = by_name[input_id]
        value = binding.value
        declared = spec.get("shape", [])
        if len(declared) < 2 or declared[:2] != ["time", "batch"]:
            raise ValueError(
                f"input {input_id} dense binding requires declared shape beginning with ['time', 'batch']"
            )
        if value.ndim == len(declared) - 1 and declared[1] == "batch":
            value = value.unsqueeze(1)
        if value.ndim != len(declared):
            raise ValueError(
                f"input {input_id} rank expected {len(declared)}, got {value.ndim}"
            )
        expected_tail = tuple(int(axis) for axis in declared[2:])
        if tuple(value.shape[2:]) != expected_tail:
            raise ValueError(
                f"input {input_id} trailing shape expected {expected_tail}, got {tuple(value.shape[2:])}"
            )
        current_leading = (int(value.shape[0]), int(value.shape[1]))
        if leading_shape is None:
            leading_shape = current_leading
        elif current_leading != leading_shape:
            raise ValueError(
                f"input {input_id} leading shape expected {leading_shape}, got {current_leading}"
            )
        signal_type = spec.get("signal_type")
        if signal_type == "mask":
            if value.dtype != torch.bool:
                if not torch.all((value == 0) | (value == 1)):
                    raise ValueError(
                        f"input {input_id} mask values must be boolean or zero/one"
                    )
                value = value.bool()
            masks.append(input_id)
        elif not (value.is_floating_point() or value.dtype == torch.bool):
            value = value.float()
        if value.is_floating_point() and not torch.isfinite(value).all():
            raise ValueError(f"input {input_id} contains non-finite values")
        if signal_type == "spikes" and not torch.all((value == 0) | (value == 1)):
            raise ValueError(
                f"input {input_id} spike values must be boolean or zero/one"
            )
        if signal_type != "mask":
            value = value.float()
        value = value.to(device)
        resolved[input_id] = value
        rows.append(
            {
                "input_id": input_id,
                "representation": "dense_array",
                "shape": list(value.shape),
                "dtype": str(value.dtype).removeprefix("torch."),
                "signal_type": signal_type,
                "unit": spec.get("unit"),
                "source": dict(binding.source),
            }
        )
    assert leading_shape is not None
    dt_ms = float(graph["timebase"]["dt"]["value"])
    supplied = dict(protocol or {})
    dataset = dict(supplied.pop("dataset", {}))
    reserved = {
        "schema",
        "binding_schema",
        "representation",
        "inputs",
        "timing",
        "masks",
        "seeds",
    }
    if reserved & supplied.keys():
        raise ValueError(
            f"execution protocol cannot override reserved fields {sorted(reserved & supplied.keys())}"
        )
    dataset.setdefault("identity", None)
    dataset.setdefault("split", None)
    dataset.setdefault("sample_cap", leading_shape[1])
    dataset.setdefault("batch_size", leading_shape[1])
    dataset.setdefault("shuffle", None)
    execution_protocol = {
        "schema": EXECUTION_PROTOCOL_SCHEMA,
        "binding_schema": DENSE_ARRAY_BINDING_SCHEMA,
        "representation": "dense_array",
        "inputs": rows,
        "dataset": dataset,
        "timing": {
            "dt_ms": dt_ms,
            "steps": leading_shape[0],
            "duration_ms": leading_shape[0] * dt_ms,
        },
        "masks": masks,
        "seeds": {"execution": int(seed)},
        **supplied,
    }
    try:
        json.dumps(execution_protocol, sort_keys=True)
    except TypeError as exc:
        raise ValueError(
            f"execution protocol must be JSON-serializable: {exc}"
        ) from exc
    return ResolvedDenseInputs(tensors=resolved, protocol=execution_protocol)

resolve_event_stream_bindings

View source

def resolve_event_stream_bindings(graph: Mapping[str, Any], *, bindings: Sequence[EventStreamBinding], device: str | torch.device='cpu', seed: int=0, protocol: Mapping[str, Any] | None=None) -> ResolvedDenseInputs

Source docstring:

Validate sparse spike coordinates and materialize binary graph inputs.
ParameterAnnotationDefaultMeaning
graphMapping[str, Any]requiredSerialized graph mapping.
bindingsSequence[EventStreamBinding]requiredDefined by the source contract and implementation below.
devicestr | torch.device'cpu'Requested or resolved tensor execution device.
seedint0Seed controlling this operation’s random stream.
protocolMapping[str, Any] | NoneNoneDefined by the source contract and implementation below.

Return annotation: ResolvedDenseInputs.

Return expressions (branch-dependent; names refer to the linked implementation):

ResolvedDenseInputs(tensors=resolved, protocol=execution_protocol)

Explicit exceptions in this implementation; called helpers may raise additional errors:

Explicit exception expression
ValueError(f'event input ids do not match graph inputs; missing={missing}, unexpected={unexpected}')
ValueError('graph execution requires at least one event input binding')
ValueError(f'execution protocol cannot override reserved fields {sorted(reserved & supplied.keys())}')
ValueError(f'duplicate event binding for input {binding.input_id}')
ValueError(f"input {input_id} event binding requires shape ['time', 'batch', channels]")
ValueError(f'input {input_id} event binding requires signal_type spikes')
ValueError(f'event input {input_id} steps_count and batch_size must be integers')
ValueError(f'event input {input_id} steps_count and batch_size must be positive')
ValueError(f'event input {input_id} coordinates must be one-dimensional')
ValueError(f'event input {input_id} coordinate lengths must match')
ValueError(f'event input {input_id} coordinates must use integer dtypes')
ValueError(f'execution protocol must be JSON-serializable: {exc}')
ValueError(f'event input {input_id} {label} coordinates must be in [0, {upper})')
ValueError(f'event input {input_id} coordinates must be ordered by step, batch, channel')
ValueError(f'event input {input_id} contains duplicate coordinates')
ValueError(f'event input {input_id} leading shape expected {leading_shape}, got {current_leading}')
Implementation
def resolve_event_stream_bindings(
    graph: Mapping[str, Any],
    *,
    bindings: Sequence[EventStreamBinding],
    device: str | torch.device = "cpu",
    seed: int = 0,
    protocol: Mapping[str, Any] | None = None,
) -> ResolvedDenseInputs:
    """Validate sparse spike coordinates and materialize binary graph inputs."""
    specs = {row["id"]: row for row in graph.get("inputs", [])}
    by_name: dict[str, EventStreamBinding] = {}
    for binding in bindings:
        if binding.input_id in by_name:
            raise ValueError(f"duplicate event binding for input {binding.input_id}")
        by_name[binding.input_id] = binding
    if set(by_name) != set(specs):
        missing = sorted(set(specs) - set(by_name))
        unexpected = sorted(set(by_name) - set(specs))
        raise ValueError(
            f"event input ids do not match graph inputs; missing={missing}, unexpected={unexpected}"
        )

    resolved: dict[str, torch.Tensor] = {}
    rows: list[dict[str, Any]] = []
    leading_shape: tuple[int, int] | None = None
    integer_dtypes = {
        torch.int8,
        torch.int16,
        torch.int32,
        torch.int64,
        torch.uint8,
    }
    for input_id in sorted(specs):
        spec = specs[input_id]
        binding = by_name[input_id]
        declared = spec.get("shape", [])
        if declared[:2] != ["time", "batch"] or len(declared) != 3:
            raise ValueError(
                f"input {input_id} event binding requires shape ['time', 'batch', channels]"
            )
        if spec.get("signal_type") != "spikes":
            raise ValueError(
                f"input {input_id} event binding requires signal_type spikes"
            )
        channels_count = int(declared[2])
        if (
            not isinstance(binding.steps_count, Integral)
            or isinstance(binding.steps_count, bool)
            or not isinstance(binding.batch_size, Integral)
            or isinstance(binding.batch_size, bool)
        ):
            raise ValueError(
                f"event input {input_id} steps_count and batch_size must be integers"
            )
        steps_count = int(binding.steps_count)
        batch_size = int(binding.batch_size)
        if steps_count <= 0 or batch_size <= 0:
            raise ValueError(
                f"event input {input_id} steps_count and batch_size must be positive"
            )
        coordinates = (binding.steps, binding.batches, binding.channels)
        if any(value.ndim != 1 for value in coordinates):
            raise ValueError(
                f"event input {input_id} coordinates must be one-dimensional"
            )
        lengths = {int(value.numel()) for value in coordinates}
        if len(lengths) != 1:
            raise ValueError(f"event input {input_id} coordinate lengths must match")
        if any(value.dtype not in integer_dtypes for value in coordinates):
            raise ValueError(
                f"event input {input_id} coordinates must use integer dtypes"
            )
        steps = binding.steps.to(dtype=torch.int64, device="cpu")
        batches = binding.batches.to(dtype=torch.int64, device="cpu")
        channels = binding.channels.to(dtype=torch.int64, device="cpu")
        bounds = (
            ("step", steps, steps_count),
            ("batch", batches, batch_size),
            ("channel", channels, channels_count),
        )
        for label, values, upper in bounds:
            if torch.any(values < 0) or torch.any(values >= upper):
                raise ValueError(
                    f"event input {input_id} {label} coordinates must be in [0, {upper})"
                )
        flat = (steps * batch_size + batches) * channels_count + channels
        if flat.numel() > 1:
            differences = flat[1:] - flat[:-1]
            if torch.any(differences < 0):
                raise ValueError(
                    f"event input {input_id} coordinates must be ordered by step, batch, channel"
                )
            if torch.any(differences == 0):
                raise ValueError(
                    f"event input {input_id} contains duplicate coordinates"
                )
        current_leading = (steps_count, batch_size)
        if leading_shape is None:
            leading_shape = current_leading
        elif current_leading != leading_shape:
            raise ValueError(
                f"event input {input_id} leading shape expected {leading_shape}, got {current_leading}"
            )
        value = torch.zeros(
            (steps_count, batch_size, channels_count),
            dtype=torch.float32,
            device=device,
        )
        if flat.numel():
            value[
                steps.to(device=device),
                batches.to(device=device),
                channels.to(device=device),
            ] = 1.0
        resolved[input_id] = value
        rows.append(
            {
                "input_id": input_id,
                "representation": "event_stream",
                "shape": list(value.shape),
                "dtype": "float32",
                "signal_type": "spikes",
                "unit": spec.get("unit"),
                "event_count": int(flat.numel()),
                "source": dict(binding.source),
            }
        )
    if leading_shape is None:
        raise ValueError("graph execution requires at least one event input binding")
    dt_ms = float(graph["timebase"]["dt"]["value"])
    supplied = dict(protocol or {})
    dataset = dict(supplied.pop("dataset", {}))
    reserved = {
        "schema",
        "binding_schema",
        "representation",
        "inputs",
        "dataset",
        "timing",
        "masks",
        "seeds",
        "resolution",
    }
    if reserved & supplied.keys():
        raise ValueError(
            f"execution protocol cannot override reserved fields {sorted(reserved & supplied.keys())}"
        )
    dataset.setdefault("identity", None)
    dataset.setdefault("split", None)
    dataset.setdefault("sample_cap", leading_shape[1])
    dataset.setdefault("batch_size", leading_shape[1])
    dataset.setdefault("shuffle", None)
    execution_protocol = {
        "schema": EXECUTION_PROTOCOL_SCHEMA,
        "binding_schema": EVENT_STREAM_BINDING_SCHEMA,
        "representation": "event_stream",
        "inputs": rows,
        "dataset": dataset,
        "timing": {
            "dt_ms": dt_ms,
            "steps": leading_shape[0],
            "duration_ms": leading_shape[0] * dt_ms,
        },
        "masks": [],
        "seeds": {"execution": int(seed)},
        "resolution": {
            "coordinates": "zero_based_integer_steps",
            "ordering": "step,batch,channel",
            "duplicates": "reject",
            "materialization": "binary_dense",
        },
        **supplied,
    }
    try:
        json.dumps(execution_protocol, sort_keys=True)
    except TypeError as exc:
        raise ValueError(
            f"execution protocol must be JSON-serializable: {exc}"
        ) from exc
    return ResolvedDenseInputs(tensors=resolved, protocol=execution_protocol)

resolve_input_bindings

View source

def resolve_input_bindings(graph: Mapping[str, Any], *, dense_bindings: Sequence[DenseArrayBinding]=(), event_bindings: Sequence[EventStreamBinding]=(), poisson_bindings: Sequence[PoissonInputBinding]=(), inputs: Mapping[str, torch.Tensor] | None=None, device: str | torch.device='cpu', seed: int=0, protocol: Mapping[str, Any] | None=None) -> ResolvedDenseInputs

Source docstring:

Resolve dense, event-stream, generated-Poisson, or mixed graph inputs.
ParameterAnnotationDefaultMeaning
graphMapping[str, Any]requiredSerialized graph mapping.
dense_bindingsSequence[DenseArrayBinding]()Defined by the source contract and implementation below.
event_bindingsSequence[EventStreamBinding]()Named sparse event bindings with zero-based step/batch/channel coordinates.
poisson_bindingsSequence[PoissonInputBinding]()Declared generated Poisson input bindings.
inputsMapping[str, torch.Tensor] | NoneNoneInput tensors keyed by graph input id.
devicestr | torch.device'cpu'Requested or resolved tensor execution device.
seedint0Seed controlling this operation’s random stream.
protocolMapping[str, Any] | NoneNoneDefined by the source contract and implementation below.

Return annotation: ResolvedDenseInputs.

Return expressions (branch-dependent; names refer to the linked implementation):

resolve_poisson_input_bindings(graph, bindings=poisson_bindings, device=device, seed=seed, protocol=protocol)
resolve_dense_array_bindings(graph, bindings=dense_bindings, device=device, seed=seed, protocol=protocol)
resolve_event_stream_bindings(graph, bindings=event_bindings, device=device, seed=seed, protocol=protocol)
ResolvedDenseInputs(tensors={**dense.tensors, **events.tensors}, protocol=execution_protocol)

Explicit exceptions in this implementation; called helpers may raise additional errors:

Explicit exception expression
ValueError('provide dense input bindings or input tensors, not both')
ValueError(f'graph inputs cannot have dense and event bindings: {overlap}')
ValueError(f'input ids do not match graph inputs; missing={missing}, unexpected={unexpected}')
ValueError('dense and event input bindings must resolve to the same timestep, duration, and batch shape')
ValueError('Poisson bindings cannot yet be mixed with replay bindings')
Implementation
def resolve_input_bindings(
    graph: Mapping[str, Any],
    *,
    dense_bindings: Sequence[DenseArrayBinding] = (),
    event_bindings: Sequence[EventStreamBinding] = (),
    poisson_bindings: Sequence[PoissonInputBinding] = (),
    inputs: Mapping[str, torch.Tensor] | None = None,
    device: str | torch.device = "cpu",
    seed: int = 0,
    protocol: Mapping[str, Any] | None = None,
) -> ResolvedDenseInputs:
    """Resolve dense, event-stream, generated-Poisson, or mixed graph inputs."""
    if dense_bindings and inputs:
        raise ValueError("provide dense input bindings or input tensors, not both")
    if inputs:
        dense_bindings = tuple(
            DenseArrayBinding(name, value, {"kind": "memory"})
            for name, value in inputs.items()
        )
    dense_ids = {binding.input_id for binding in dense_bindings}
    event_ids = {binding.input_id for binding in event_bindings}
    poisson_ids = {binding.input_id for binding in poisson_bindings}
    overlap = sorted(
        (dense_ids & event_ids) | (dense_ids & poisson_ids) | (event_ids & poisson_ids)
    )
    if overlap:
        raise ValueError(
            f"graph inputs cannot have dense and event bindings: {overlap}"
        )
    graph_ids = {row["id"] for row in graph.get("inputs", [])}
    if dense_ids | event_ids | poisson_ids != graph_ids:
        missing = sorted(graph_ids - dense_ids - event_ids - poisson_ids)
        unexpected = sorted((dense_ids | event_ids | poisson_ids) - graph_ids)
        raise ValueError(
            f"input ids do not match graph inputs; missing={missing}, unexpected={unexpected}"
        )
    if poisson_bindings:
        if dense_bindings or event_bindings:
            raise ValueError(
                "Poisson bindings cannot yet be mixed with replay bindings"
            )
        return resolve_poisson_input_bindings(
            graph,
            bindings=poisson_bindings,
            device=device,
            seed=seed,
            protocol=protocol,
        )
    if not event_bindings:
        return resolve_dense_array_bindings(
            graph,
            bindings=dense_bindings,
            device=device,
            seed=seed,
            protocol=protocol,
        )
    if not dense_bindings:
        return resolve_event_stream_bindings(
            graph,
            bindings=event_bindings,
            device=device,
            seed=seed,
            protocol=protocol,
        )

    def graph_with_inputs(input_ids: set[str]) -> dict[str, Any]:
        return {
            **graph,
            "inputs": [
                row for row in graph.get("inputs", []) if row["id"] in input_ids
            ],
        }

    dense = resolve_dense_array_bindings(
        graph_with_inputs(dense_ids),
        bindings=dense_bindings,
        device=device,
        seed=seed,
        protocol=protocol,
    )
    events = resolve_event_stream_bindings(
        graph_with_inputs(event_ids),
        bindings=event_bindings,
        device=device,
        seed=seed,
        protocol=protocol,
    )
    if dense.protocol["timing"] != events.protocol["timing"]:
        raise ValueError(
            "dense and event input bindings must resolve to the same timestep, duration, and batch shape"
        )
    execution_protocol = {
        "schema": EXECUTION_PROTOCOL_SCHEMA,
        "binding_schema": MIXED_INPUT_BINDING_SCHEMA,
        "representation": "mixed",
        "inputs": sorted(
            [*dense.protocol["inputs"], *events.protocol["inputs"]],
            key=lambda row: row["input_id"],
        ),
        "dataset": dense.protocol["dataset"],
        "timing": dense.protocol["timing"],
        "masks": dense.protocol["masks"],
        "seeds": dense.protocol["seeds"],
        "resolution": {"event_stream": events.protocol["resolution"]},
        **{
            key: value
            for key, value in dense.protocol.items()
            if key
            not in {
                "schema",
                "binding_schema",
                "representation",
                "inputs",
                "dataset",
                "timing",
                "masks",
                "seeds",
                "resolution",
            }
        },
    }
    return ResolvedDenseInputs(
        tensors={**dense.tensors, **events.tensors}, protocol=execution_protocol
    )

resolve_poisson_input_bindings

View source

def resolve_poisson_input_bindings(graph: Mapping[str, Any], *, bindings: Sequence[PoissonInputBinding], device: str | torch.device='cpu', seed: int=0, protocol: Mapping[str, Any] | None=None) -> ResolvedDenseInputs

Source docstring:

Generate reproducible fixed or per-presentation categorical Poisson spikes.
ParameterAnnotationDefaultMeaning
graphMapping[str, Any]requiredSerialized graph mapping.
bindingsSequence[PoissonInputBinding]requiredDefined by the source contract and implementation below.
devicestr | torch.device'cpu'Requested or resolved tensor execution device.
seedint0Seed controlling this operation’s random stream.
protocolMapping[str, Any] | NoneNoneDefined by the source contract and implementation below.

Return annotation: ResolvedDenseInputs.

Return expressions (branch-dependent; names refer to the linked implementation):

ResolvedDenseInputs(tensors=tensors, protocol=execution_protocol)

Explicit exceptions in this implementation; called helpers may raise additional errors:

Explicit exception expression
ValueError('duplicate Poisson binding for a graph input')
ValueError(f'Poisson input ids do not match graph inputs; missing={missing}, unexpected={unexpected}')
ValueError('execution protocol cannot override reserved Poisson fields')
ValueError(f"input {input_id} Poisson binding requires shape ['time', 'batch', channels]")
ValueError(f'input {input_id} Poisson binding requires signal_type spikes')
ValueError(f'Poisson input {input_id} steps_count and batch_size must be positive')
ValueError(f'Poisson input {input_id} rates must be finite and non-negative')
ValueError(f'fixed-rate Poisson input {input_id} requires exactly one rate')
ValueError(f'Poisson input {input_id} rate times dt exceeds probability one')
ValueError(f'Poisson input {input_id} leading shape expected {leading_shape}, got {current_leading}')
Implementation
def resolve_poisson_input_bindings(
    graph: Mapping[str, Any],
    *,
    bindings: Sequence[PoissonInputBinding],
    device: str | torch.device = "cpu",
    seed: int = 0,
    protocol: Mapping[str, Any] | None = None,
) -> ResolvedDenseInputs:
    """Generate reproducible fixed or per-presentation categorical Poisson spikes."""
    specs = {row["id"]: row for row in graph.get("inputs", [])}
    by_name = {binding.input_id: binding for binding in bindings}
    if len(by_name) != len(bindings):
        raise ValueError("duplicate Poisson binding for a graph input")
    if set(by_name) != set(specs):
        missing = sorted(set(specs) - set(by_name))
        unexpected = sorted(set(by_name) - set(specs))
        raise ValueError(
            f"Poisson input ids do not match graph inputs; missing={missing}, unexpected={unexpected}"
        )
    dt_ms = float(graph["timebase"]["dt"]["value"])
    tensors: dict[str, torch.Tensor] = {}
    rows: list[dict[str, Any]] = []
    leading_shape: tuple[int, int] | None = None
    for input_id in sorted(specs):
        spec = specs[input_id]
        binding = by_name[input_id]
        declared = spec.get("shape", [])
        if declared[:2] != ["time", "batch"] or len(declared) != 3:
            raise ValueError(
                f"input {input_id} Poisson binding requires shape ['time', 'batch', channels]"
            )
        if spec.get("signal_type") != "spikes":
            raise ValueError(
                f"input {input_id} Poisson binding requires signal_type spikes"
            )
        if binding.steps_count <= 0 or binding.batch_size <= 0:
            raise ValueError(
                f"Poisson input {input_id} steps_count and batch_size must be positive"
            )
        rates = tuple(float(rate) for rate in binding.rates_hz)
        if not rates or any(not math.isfinite(rate) or rate < 0 for rate in rates):
            raise ValueError(
                f"Poisson input {input_id} rates must be finite and non-negative"
            )
        if not binding.categorical and len(rates) != 1:
            raise ValueError(
                f"fixed-rate Poisson input {input_id} requires exactly one rate"
            )
        if max(rates) * dt_ms / 1000.0 > 1.0:
            raise ValueError(
                f"Poisson input {input_id} rate times dt exceeds probability one"
            )
        current_leading = (int(binding.steps_count), int(binding.batch_size))
        if leading_shape is None:
            leading_shape = current_leading
        elif current_leading != leading_shape:
            raise ValueError(
                f"Poisson input {input_id} leading shape expected {leading_shape}, got {current_leading}"
            )
        generator = torch.Generator(device="cpu").manual_seed(int(binding.seed))
        if binding.categorical:
            indices = torch.randint(
                len(rates), (binding.batch_size,), generator=generator
            )
            realized = torch.tensor(rates, dtype=torch.float32)[indices]
        else:
            realized = torch.full((binding.batch_size,), rates[0], dtype=torch.float32)
        probability = realized.reshape(1, -1, 1) * dt_ms / 1000.0
        value = (
            (
                torch.rand(
                    binding.steps_count,
                    binding.batch_size,
                    int(declared[2]),
                    generator=generator,
                )
                < probability
            )
            .float()
            .to(device)
        )
        tensors[input_id] = value
        rows.append(
            {
                "input_id": input_id,
                "representation": "poisson",
                "shape": list(value.shape),
                "dtype": "float32",
                "signal_type": "spikes",
                "unit": spec.get("unit"),
                "protocol": "categorical_rate" if binding.categorical else "fixed_rate",
                "rates_hz": list(rates),
                "realized_rates_hz": realized.tolist(),
                "seed": int(binding.seed),
                "selection": "uniform_independent_per_presentation"
                if binding.categorical
                else "constant",
            }
        )
    assert leading_shape is not None
    supplied = dict(protocol or {})
    dataset = dict(supplied.pop("dataset", {}))
    if {
        "schema",
        "binding_schema",
        "representation",
        "inputs",
        "timing",
        "seeds",
    } & supplied.keys():
        raise ValueError("execution protocol cannot override reserved Poisson fields")
    dataset.setdefault("identity", None)
    dataset.setdefault("split", None)
    dataset.setdefault("sample_cap", leading_shape[1])
    dataset.setdefault("batch_size", leading_shape[1])
    dataset.setdefault("shuffle", None)
    execution_protocol = {
        "schema": EXECUTION_PROTOCOL_SCHEMA,
        "binding_schema": POISSON_INPUT_BINDING_SCHEMA,
        "representation": "poisson",
        "inputs": rows,
        "dataset": dataset,
        "timing": {
            "dt_ms": dt_ms,
            "steps": leading_shape[0],
            "duration_ms": leading_shape[0] * dt_ms,
        },
        "masks": [],
        "seeds": {
            "execution": int(seed),
            "poisson": {row["input_id"]: row["seed"] for row in rows},
        },
        "resolution": {
            "distribution": "Bernoulli discretization of homogeneous Poisson",
            "rate_selection": "per_presentation",
        },
        **supplied,
    }
    json.dumps(execution_protocol, sort_keys=True)
    return ResolvedDenseInputs(tensors=tensors, protocol=execution_protocol)

resolve_dataset_snapshot_binding

View source

def resolve_dataset_snapshot_binding(graph: Mapping[str, Any], binding: DatasetSnapshotBinding, *, device: str | torch.device='cpu', execution_seed: int=0, protocol: Mapping[str, Any] | None=None) -> tuple[ResolvedDenseInputs, tuple[TargetArrayBinding, ...]]

Source docstring:

Load, select, and encode one immutable MNIST/SHD-style NPZ snapshot.
ParameterAnnotationDefaultMeaning
graphMapping[str, Any]requiredSerialized graph mapping.
bindingDatasetSnapshotBindingrequiredDefined by the source contract and implementation below.
devicestr | torch.device'cpu'Requested or resolved tensor execution device.
execution_seedint0Defined by the source contract and implementation below.
protocolMapping[str, Any] | NoneNoneDefined by the source contract and implementation below.

Return annotation: tuple[ResolvedDenseInputs, tuple[TargetArrayBinding, ...]].

Return expressions (branch-dependent; names refer to the linked implementation):

(ResolvedDenseInputs(resolved.tensors, protocol), targets)

Explicit exceptions in this implementation; called helpers may raise additional errors:

Explicit exception expression
ValueError('dataset snapshot binding requires exactly its one named graph input')
ValueError("dataset snapshot binding requires graph input shape ['time', 'batch', channels]")
ValueError('dataset snapshot binding requires a spike graph input')
ValueError('dataset snapshot binding requires an NPZ file')
ValueError('dataset snapshot identity and split must be non-empty')
ValueError(f'dataset snapshot is missing labels key {binding.label_key!r}')
ValueError('dataset snapshot labels must be a one-dimensional integer array')
ValueError('dataset snapshot must contain at least one sample')
ValueError(f'dataset snapshot sample cap must be in [1, {sample_count}]')
ValueError('rate-Poisson encoder requires explicit duration_ms and max_rate_hz')
ValueError('prebinned encoder does not accept duration, max rate, or a stochastic seed')
ValueError('event-bin encoder requires duration_ms and does not accept max rate or a stochastic seed')
ValueError(f'dataset execution protocol metadata does not match binding; unknown={unknown_dataset}, conflicts={conflicts}')
ValueError(f'dataset execution protocol cannot override reserved fields {sorted(reserved & supplied.keys())}')
ValueError(f'dataset snapshot is missing feature key {binding.feature_key!r}')
ValueError(f'rate-Poisson dataset features expected shape {(sample_count, channels)}, got {features.shape}')
ValueError('rate-Poisson dataset features must use a floating dtype')
ValueError('rate-Poisson dataset features must be finite in [0, 1]')
ValueError('rate-Poisson encoder requires timestep-aligned positive duration and a supported finite non-negative max rate')
ValueError(f'prebinned dataset features expected [time, {sample_count}, {channels}]')
ValueError('prebinned dataset spikes must be binary')
ValueError(f'unsupported dataset encoder {encoder.kind!r}')
ValueError(f'event dataset snapshot is missing keys {missing}')
ValueError('event dataset coordinates must be equal-length sample/channel integer and time floating arrays')
ValueError('event-bin encoder duration must be a positive integer number of timesteps')
ValueError('event dataset coordinates are out of snapshot bounds')
Implementation
def resolve_dataset_snapshot_binding(
    graph: Mapping[str, Any],
    binding: DatasetSnapshotBinding,
    *,
    device: str | torch.device = "cpu",
    execution_seed: int = 0,
    protocol: Mapping[str, Any] | None = None,
) -> tuple[ResolvedDenseInputs, tuple[TargetArrayBinding, ...]]:
    """Load, select, and encode one immutable MNIST/SHD-style NPZ snapshot."""
    specs = {row["id"]: row for row in graph.get("inputs", [])}
    if set(specs) != {binding.input_id}:
        raise ValueError(
            "dataset snapshot binding requires exactly its one named graph input"
        )
    spec = specs[binding.input_id]
    declared = spec.get("shape", [])
    if declared[:2] != ["time", "batch"] or len(declared) != 3:
        raise ValueError(
            "dataset snapshot binding requires graph input shape ['time', 'batch', channels]"
        )
    if spec.get("signal_type") != "spikes":
        raise ValueError("dataset snapshot binding requires a spike graph input")
    source_path = Path(binding.path)
    if source_path.suffix.lower() != ".npz":
        raise ValueError("dataset snapshot binding requires an NPZ file")
    source_digest = _file_digest(source_path)
    loaded = np.load(source_path, allow_pickle=False)
    try:
        arrays = {key: loaded[key] for key in loaded.files}
    finally:
        loaded.close()
    if not binding.dataset_id or not binding.split:
        raise ValueError("dataset snapshot identity and split must be non-empty")
    if binding.label_key not in arrays:
        raise ValueError(
            f"dataset snapshot is missing labels key {binding.label_key!r}"
        )
    labels = np.asarray(arrays[binding.label_key])
    if labels.ndim != 1 or not np.issubdtype(labels.dtype, np.integer):
        raise ValueError(
            "dataset snapshot labels must be a one-dimensional integer array"
        )
    sample_count = int(labels.shape[0])
    if sample_count <= 0:
        raise ValueError("dataset snapshot must contain at least one sample")
    cap = sample_count if binding.sample_cap is None else int(binding.sample_cap)
    if cap <= 0 or cap > sample_count:
        raise ValueError(f"dataset snapshot sample cap must be in [1, {sample_count}]")
    order = torch.arange(sample_count)
    if binding.shuffle:
        generator = torch.Generator(device="cpu").manual_seed(int(binding.order_seed))
        order = torch.randperm(sample_count, generator=generator)
    selected = order[:cap].numpy()
    selected_labels = labels[selected].astype(np.int64, copy=False)
    encoder = binding.encoder
    dt_ms = float(graph["timebase"]["dt"]["value"])
    channels = int(declared[2])
    encoder_row: dict[str, Any]
    if encoder.kind == "rate_poisson" and (
        encoder.duration_ms is None or encoder.max_rate_hz is None
    ):
        raise ValueError(
            "rate-Poisson encoder requires explicit duration_ms and max_rate_hz"
        )
    if encoder.kind == "prebinned_spikes" and (
        encoder.duration_ms is not None
        or encoder.max_rate_hz is not None
        or encoder.seed != 0
    ):
        raise ValueError(
            "prebinned encoder does not accept duration, max rate, or a stochastic seed"
        )
    if encoder.kind == "event_bin" and (
        encoder.duration_ms is None
        or encoder.max_rate_hz is not None
        or encoder.seed != 0
    ):
        raise ValueError(
            "event-bin encoder requires duration_ms and does not accept max rate or a stochastic seed"
        )
    if encoder.kind in {"rate_poisson", "prebinned_spikes"}:
        if binding.feature_key not in arrays:
            raise ValueError(
                f"dataset snapshot is missing feature key {binding.feature_key!r}"
            )
        features = np.asarray(arrays[binding.feature_key])
    if encoder.kind == "rate_poisson":
        if features.shape != (sample_count, channels):
            raise ValueError(
                f"rate-Poisson dataset features expected shape {(sample_count, channels)}, got {features.shape}"
            )
        if not np.issubdtype(features.dtype, np.floating):
            raise ValueError("rate-Poisson dataset features must use a floating dtype")
        if (
            not np.isfinite(features).all()
            or np.any(features < 0)
            or np.any(features > 1)
        ):
            raise ValueError("rate-Poisson dataset features must be finite in [0, 1]")
        duration_ms = float(encoder.duration_ms or 0)
        max_rate_hz = float(encoder.max_rate_hz or 0)
        raw_steps = duration_ms / dt_ms
        if (
            duration_ms <= 0
            or not math.isclose(raw_steps, round(raw_steps), abs_tol=1e-9)
            or not math.isfinite(max_rate_hz)
            or max_rate_hz < 0
            or max_rate_hz * dt_ms / 1000.0 > 1
        ):
            raise ValueError(
                "rate-Poisson encoder requires timestep-aligned positive duration and a supported finite non-negative max rate"
            )
        rates = torch.as_tensor(features[selected], dtype=torch.float32)
        generator = torch.Generator(device="cpu").manual_seed(int(encoder.seed))
        probability = rates.unsqueeze(0) * max_rate_hz * dt_ms / 1000.0
        spikes = (
            torch.rand(int(round(raw_steps)), cap, channels, generator=generator)
            < probability
        ).float()
        encoder_row = {
            "kind": encoder.kind,
            "duration_ms": duration_ms,
            "max_rate_hz": max_rate_hz,
            "seed": int(encoder.seed),
            "distribution": "Bernoulli discretization of feature-scaled homogeneous Poisson",
        }
    elif encoder.kind == "prebinned_spikes":
        if features.ndim != 3 or features.shape[1:] != (sample_count, channels):
            raise ValueError(
                f"prebinned dataset features expected [time, {sample_count}, {channels}]"
            )
        if not np.all((features == 0) | (features == 1)):
            raise ValueError("prebinned dataset spikes must be binary")
        spikes = torch.as_tensor(features[:, selected, :], dtype=torch.float32)
        encoder_row = {
            "kind": encoder.kind,
            "duration_ms": int(features.shape[0]) * dt_ms,
            "seed": None,
        }
    elif encoder.kind == "event_bin":
        event_keys = ("event_sample", "event_time_ms", "event_channel")
        missing = [key for key in event_keys if key not in arrays]
        if missing:
            raise ValueError(f"event dataset snapshot is missing keys {missing}")
        samples = np.asarray(arrays["event_sample"])
        times = np.asarray(arrays["event_time_ms"])
        event_channels = np.asarray(arrays["event_channel"])
        if (
            samples.ndim != 1
            or times.ndim != 1
            or event_channels.ndim != 1
            or not (len(samples) == len(times) == len(event_channels))
            or not np.issubdtype(samples.dtype, np.integer)
            or not np.issubdtype(event_channels.dtype, np.integer)
            or not np.issubdtype(times.dtype, np.floating)
        ):
            raise ValueError(
                "event dataset coordinates must be equal-length sample/channel integer and time floating arrays"
            )
        duration_ms = float(encoder.duration_ms or 0)
        raw_steps = duration_ms / dt_ms
        if duration_ms <= 0 or not math.isclose(
            raw_steps, round(raw_steps), abs_tol=1e-9
        ):
            raise ValueError(
                "event-bin encoder duration must be a positive integer number of timesteps"
            )
        if (
            not np.isfinite(times).all()
            or np.any(samples < 0)
            or np.any(samples >= sample_count)
            or np.any(event_channels < 0)
            or np.any(event_channels >= channels)
            or np.any(times < 0)
            or np.any(times >= duration_ms)
        ):
            raise ValueError("event dataset coordinates are out of snapshot bounds")
        selected_position = {
            int(sample): index for index, sample in enumerate(selected)
        }
        spikes = torch.zeros(int(round(raw_steps)), cap, channels)
        retained = 0
        collisions = 0
        for sample, time_ms, channel in zip(
            samples, times, event_channels, strict=True
        ):
            batch = selected_position.get(int(sample))
            if batch is None:
                continue
            step = min(int(float(time_ms) / dt_ms), spikes.shape[0] - 1)
            collisions += int(spikes[step, batch, int(channel)] != 0)
            spikes[step, batch, int(channel)] = 1
            retained += 1
        encoder_row = {
            "kind": encoder.kind,
            "duration_ms": duration_ms,
            "seed": None,
            "timestamp_unit": "ms",
            "binning": "floor_left_closed_right_open",
            "retained_events": retained,
            "binary_collisions": collisions,
        }
    else:
        raise ValueError(f"unsupported dataset encoder {encoder.kind!r}")
    source = {
        "kind": "dataset_snapshot",
        "path": str(source_path),
        "digest": source_digest,
        "arrays": {
            "labels": binding.label_key,
            **(
                {"features": binding.feature_key}
                if encoder.kind != "event_bin"
                else {
                    "sample": "event_sample",
                    "time_ms": "event_time_ms",
                    "channel": "event_channel",
                }
            ),
        },
    }
    resolved = resolve_dense_array_bindings(
        graph,
        bindings=(DenseArrayBinding(binding.input_id, spikes, source),),
        device=device,
        seed=execution_seed,
        protocol={
            "dataset": {
                "identity": binding.dataset_id,
                "split": binding.split,
                "sample_cap": cap,
                "batch_size": cap,
                "shuffle": bool(binding.shuffle),
            }
        },
    )
    supplied = dict(protocol or {})
    supplied_dataset = dict(supplied.pop("dataset", {}))
    expected_dataset = dict(resolved.protocol["dataset"])
    unknown_dataset = sorted(set(supplied_dataset) - set(expected_dataset))
    conflicts = sorted(
        key
        for key, value in supplied_dataset.items()
        if value is not None and value != expected_dataset[key]
    )
    if unknown_dataset or conflicts:
        raise ValueError(
            "dataset execution protocol metadata does not match binding; "
            f"unknown={unknown_dataset}, conflicts={conflicts}"
        )
    reserved = {
        "schema",
        "binding_schema",
        "representation",
        "inputs",
        "timing",
        "masks",
        "seeds",
        "dataset_binding",
    }
    if reserved & supplied.keys():
        raise ValueError(
            f"dataset execution protocol cannot override reserved fields {sorted(reserved & supplied.keys())}"
        )
    protocol = {
        **resolved.protocol,
        "binding_schema": DATASET_SNAPSHOT_BINDING_SCHEMA,
        "representation": "dataset_snapshot",
        "dataset_binding": {
            "source": source,
            "sample_count": sample_count,
            "selected_indices": selected.tolist(),
            "order_seed": int(binding.order_seed),
            "encoder": encoder_row,
            "target_id": binding.target_id,
        },
        "seeds": {
            **resolved.protocol["seeds"],
            "dataset_order": int(binding.order_seed),
            "encoder": int(encoder.seed) if encoder.kind == "rate_poisson" else None,
        },
        **supplied,
    }
    targets = (
        (
            TargetArrayBinding(
                binding.target_id,
                torch.as_tensor(selected_labels),
                source={**source, "array": binding.label_key},
            ),
        )
        if binding.target_id
        else ()
    )
    return ResolvedDenseInputs(resolved.tensors, protocol), targets

graph_capability_issues

View source

def graph_capability_issues(graph: Mapping[str, Any]) -> list[CapabilityIssue]

Source docstring:

Return precise graph-executor capability failures.
ParameterAnnotationDefaultMeaning
graphMapping[str, Any]requiredSerialized graph mapping.

Return annotation: list[CapabilityIssue].

Return expressions (branch-dependent; names refer to the linked implementation):

issues
Implementation
def graph_capability_issues(graph: Mapping[str, Any]) -> list[CapabilityIssue]:
    """Return precise graph-executor capability failures."""
    issues: list[CapabilityIssue] = []
    neuron_capabilities: set[str] = {"coba_lif", "leaky_integrator"}
    synapse_capabilities: set[str] = {"ampa", "gaba", "leaky_integrator"}
    operation_capabilities: set[str] = {
        "linear",
        "reduce_mean",
        "reduce_sum",
        "select_final",
        "duration_normalise",
        "cumulative_sum",
    }
    connection_capabilities: set[str] = {"feedforward", "recurrent", "feedback"}
    for pop in graph.get("populations", []):
        kind = pop.get("neuron", {}).get("kind")
        if kind not in neuron_capabilities:
            issues.append(
                CapabilityIssue(pop["id"], f"neuron:{kind}", "unsupported neuron kind")
            )
    for projection in graph.get("projections", []):
        synapse = projection.get("synapse", {}).get("kind")
        if synapse not in synapse_capabilities:
            issues.append(
                CapabilityIssue(
                    projection["id"], f"synapse:{synapse}", "unsupported synapse kind"
                )
            )
        connection = projection.get("connection")
        if connection not in connection_capabilities:
            issues.append(
                CapabilityIssue(
                    projection["id"],
                    f"connection:{connection}",
                    "unsupported connection kind",
                )
            )
    for operation in graph.get("operations", []):
        kind = operation.get("kind")
        if kind not in operation_capabilities:
            issues.append(
                CapabilityIssue(
                    operation["id"], f"operation:{kind}", "unsupported operation kind"
                )
            )
    return issues

PlannedProjection

View source

One lowered projection, including source and target ids, receptor, parameter key, delay, decay and enabled state.

Class decorators: dataclass(frozen=True).

Dataclass constructor parameters. Factory defaults are shown as field declarations; omit these arguments to create fresh values per instance:

PlannedProjection(id: str, source: str, target: str, polarity: str, decay: float, delay_steps: int, parameter: str, enabled: bool)

Declared fields, including fields inherited from local data classes:

FieldAnnotationDefaultMeaning
idstrrequiredStable identifier in the relevant graph or data contract.
sourcestrrequiredStored member of this data contract; see the class docstring and serialization methods.
targetstrrequiredStored member of this data contract; see the class docstring and serialization methods.
polaritystrrequiredStored member of this data contract; see the class docstring and serialization methods.
decayfloatrequiredStored member of this data contract; see the class docstring and serialization methods.
delay_stepsintrequiredStored member of this data contract; see the class docstring and serialization methods.
parameterstrrequiredStored member of this data contract; see the class docstring and serialization methods.
enabledboolrequiredWhether an authored projection contributes conductance during execution.
Complete class implementation
class PlannedProjection:
    id: str
    source: str
    target: str
    polarity: str
    decay: float
    delay_steps: int
    parameter: str
    enabled: bool

GraphPlan

View source

Frozen lowered execution plan containing graph data, timestep, ordered populations, projection plans, observable bindings and output bindings.

Class decorators: dataclass(frozen=True).

Dataclass constructor parameters. Factory defaults are shown as field declarations; omit these arguments to create fresh values per instance:

GraphPlan(graph: Mapping[str, Any], dt_ms: float, populations: tuple[Mapping[str, Any], ...], projections: tuple[PlannedProjection, ...], observables: tuple[Mapping[str, Any], ...], outputs: tuple[Mapping[str, Any], ...])

Declared fields, including fields inherited from local data classes:

FieldAnnotationDefaultMeaning
graphMapping[str, Any]requiredSerialized graph mapping.
dt_msfloatrequiredSimulation timestep in milliseconds.
populationstuple[Mapping[str, Any], ...]requiredStored member of this data contract; see the class docstring and serialization methods.
projectionstuple[PlannedProjection, ...]requiredStored member of this data contract; see the class docstring and serialization methods.
observablestuple[Mapping[str, Any], ...]requiredStored member of this data contract; see the class docstring and serialization methods.
outputstuple[Mapping[str, Any], ...]requiredStored member of this data contract; see the class docstring and serialization methods.
Complete class implementation
class GraphPlan:
    graph: Mapping[str, Any]
    dt_ms: float
    populations: tuple[Mapping[str, Any], ...]
    projections: tuple[PlannedProjection, ...]
    observables: tuple[Mapping[str, Any], ...]
    outputs: tuple[Mapping[str, Any], ...]

GraphRuntimeState

View source

Source docstring:

Complete dynamic state required to continue one graph trajectory.

Static parameters are deliberately excluded: weight checkpoints and runtime
state are orthogonal, allowing a mature state to branch across compatible
parameterisations of the same graph structure.

Class decorators: dataclass.

Dataclass constructor parameters. Factory defaults are shown as field declarations; omit these arguments to create fresh values per instance:

GraphRuntimeState(signature: str, compatibility: dict[str, Any], completed_steps: int, voltages: dict[str, torch.Tensor], refractory: dict[str, torch.Tensor], conductances: dict[str, torch.Tensor], population_histories: dict[str, torch.Tensor], input_histories: dict[str, torch.Tensor])

Declared fields, including fields inherited from local data classes:

FieldAnnotationDefaultMeaning
signaturestrrequiredStored member of this data contract; see the class docstring and serialization methods.
compatibilitydict[str, Any]requiredStored member of this data contract; see the class docstring and serialization methods.
completed_stepsintrequiredStored member of this data contract; see the class docstring and serialization methods.
voltagesdict[str, torch.Tensor]requiredStored member of this data contract; see the class docstring and serialization methods.
refractorydict[str, torch.Tensor]requiredStored member of this data contract; see the class docstring and serialization methods.
conductancesdict[str, torch.Tensor]requiredStored member of this data contract; see the class docstring and serialization methods.
population_historiesdict[str, torch.Tensor]requiredStored member of this data contract; see the class docstring and serialization methods.
input_historiesdict[str, torch.Tensor]requiredStored member of this data contract; see the class docstring and serialization methods.

GraphRuntimeState.detached

View source

def GraphRuntimeState.detached(self, *, device: str | torch.device='cpu') -> GraphRuntimeState
ParameterAnnotationDefaultMeaning
devicestr | torch.device'cpu'Requested or resolved tensor execution device.

Return annotation: GraphRuntimeState.

Return expressions (branch-dependent; names refer to the linked implementation):

GraphRuntimeState(signature=self.signature, compatibility=self.compatibility, completed_steps=self.completed_steps, voltages=moved(self.voltages), refractory=moved(self.refractory), conductances=moved(self.conductances), population_histories=moved(self.population_histories), input_histories=moved(self.input_histories))
Implementation
def detached(self, *, device: str | torch.device = "cpu") -> GraphRuntimeState:
        def moved(values: Mapping[str, torch.Tensor]) -> dict[str, torch.Tensor]:
            return {
                name: value.detach().to(device).clone()
                for name, value in values.items()
            }

        return GraphRuntimeState(
            signature=self.signature,
            compatibility=self.compatibility,
            completed_steps=self.completed_steps,
            voltages=moved(self.voltages),
            refractory=moved(self.refractory),
            conductances=moved(self.conductances),
            population_histories=moved(self.population_histories),
            input_histories=moved(self.input_histories),
        )
Complete class implementation
class GraphRuntimeState:
    """Complete dynamic state required to continue one graph trajectory.

    Static parameters are deliberately excluded: weight checkpoints and runtime
    state are orthogonal, allowing a mature state to branch across compatible
    parameterisations of the same graph structure.
    """

    signature: str
    compatibility: dict[str, Any]
    completed_steps: int
    voltages: dict[str, torch.Tensor]
    refractory: dict[str, torch.Tensor]
    conductances: dict[str, torch.Tensor]
    population_histories: dict[str, torch.Tensor]
    input_histories: dict[str, torch.Tensor]

    def detached(self, *, device: str | torch.device = "cpu") -> GraphRuntimeState:
        def moved(values: Mapping[str, torch.Tensor]) -> dict[str, torch.Tensor]:
            return {
                name: value.detach().to(device).clone()
                for name, value in values.items()
            }

        return GraphRuntimeState(
            signature=self.signature,
            compatibility=self.compatibility,
            completed_steps=self.completed_steps,
            voltages=moved(self.voltages),
            refractory=moved(self.refractory),
            conductances=moved(self.conductances),
            population_histories=moved(self.population_histories),
            input_histories=moved(self.input_histories),
        )

runtime_state_compatibility

View source

def runtime_state_compatibility(plan: GraphPlan) -> dict[str, Any]

Source docstring:

Describe state-layout and dynamical semantics, excluding parameter values.
ParameterAnnotationDefaultMeaning
planGraphPlanrequiredLowered GraphPlan for the complete graph.

Return annotation: dict[str, Any].

Return expressions (branch-dependent; names refer to the linked implementation):

{'schema': RUNTIME_STATE_SCHEMA, 'dt_ms': plan.dt_ms, 'populations': [{'id': row['id'], 'size': row['size'], 'neuron': row['neuron']} for row in plan.populations], 'inputs': [{'id': row['id'], 'shape': row['shape'], 'signal_type': row.get('signal_type')} for row in plan.graph.get('inputs', [])], 'projections': [{'id': row.id, 'source': row.source, 'target': row.target, 'polarity': row.polarity, 'synapse': next((item['synapse'] for item in plan.graph.get('projections', []) if item['id'] == row.id)), 'delay_steps': row.delay_steps, 'parameter': row.parameter, 'parameter_shape': parameters[row.parameter]['shape'], 'enabled': row.enabled} for row in plan.projections]}
Implementation
def runtime_state_compatibility(plan: GraphPlan) -> dict[str, Any]:
    """Describe state-layout and dynamical semantics, excluding parameter values."""
    parameters = {row["id"]: row for row in plan.graph.get("parameters", [])}
    return {
        "schema": RUNTIME_STATE_SCHEMA,
        "dt_ms": plan.dt_ms,
        "populations": [
            {"id": row["id"], "size": row["size"], "neuron": row["neuron"]}
            for row in plan.populations
        ],
        "inputs": [
            {
                "id": row["id"],
                "shape": row["shape"],
                "signal_type": row.get("signal_type"),
            }
            for row in plan.graph.get("inputs", [])
        ],
        "projections": [
            {
                "id": row.id,
                "source": row.source,
                "target": row.target,
                "polarity": row.polarity,
                "synapse": next(
                    item["synapse"]
                    for item in plan.graph.get("projections", [])
                    if item["id"] == row.id
                ),
                "delay_steps": row.delay_steps,
                "parameter": row.parameter,
                "parameter_shape": parameters[row.parameter]["shape"],
                "enabled": row.enabled,
            }
            for row in plan.projections
        ],
    }

runtime_state_signature

View source

def runtime_state_signature(plan: GraphPlan) -> str

Hash the canonical runtime compatibility mapping for a planned graph; use this signature to reject continuation against an incompatible plan.

ParameterAnnotationDefaultMeaning
planGraphPlanrequiredLowered GraphPlan for the complete graph.

Return annotation: str.

Return expressions (branch-dependent; names refer to the linked implementation):

'sha256:' + hashlib.sha256(encoded).hexdigest()
Implementation
def runtime_state_signature(plan: GraphPlan) -> str:
    payload = runtime_state_compatibility(plan)
    encoded = json.dumps(payload, sort_keys=True, separators=(",", ":")).encode()
    return "sha256:" + hashlib.sha256(encoded).hexdigest()

save_runtime_state

View source

def save_runtime_state(path: str | Path, state: GraphRuntimeState) -> Path

Source docstring:

Atomically publish a portable JSON/NPZ graph-runtime state directory.
ParameterAnnotationDefaultMeaning
pathstr | PathrequiredFilesystem source or destination path, as described below.
stateGraphRuntimeStaterequiredDefined by the source contract and implementation below.

Return annotation: Path.

Return expressions (branch-dependent; names refer to the linked implementation):

root
Implementation
def save_runtime_state(path: str | Path, state: GraphRuntimeState) -> Path:
    """Atomically publish a portable JSON/NPZ graph-runtime state directory."""
    root = Path(path)
    root.mkdir(parents=True, exist_ok=True)
    groups = {
        "voltages": state.voltages,
        "refractory": state.refractory,
        "conductances": state.conductances,
        "population_histories": state.population_histories,
        "input_histories": state.input_histories,
    }
    arrays: dict[str, np.ndarray] = {}
    tensors: list[dict[str, Any]] = []
    for group_name, values in groups.items():
        for name in sorted(values):
            key = f"tensor_{len(arrays):04d}"
            value = values[name].detach().cpu().contiguous()
            arrays[key] = value.numpy()
            tensors.append(
                {
                    "group": group_name,
                    "name": name,
                    "key": key,
                    "shape": list(value.shape),
                    "dtype": str(value.dtype).removeprefix("torch."),
                }
            )
    fd, temporary_name = tempfile.mkstemp(prefix=".tensors-", suffix=".npz", dir=root)
    os.close(fd)
    temporary_tensors = Path(temporary_name)
    try:
        np.savez_compressed(temporary_tensors, **arrays)
        tensors_digest = _file_digest(temporary_tensors)
        os.replace(temporary_tensors, root / "tensors.npz")
    finally:
        temporary_tensors.unlink(missing_ok=True)
    manifest = {
        "schema": RUNTIME_STATE_SCHEMA,
        "schema_version": 1,
        "signature": state.signature,
        "compatibility": state.compatibility,
        "completed_steps": state.completed_steps,
        "tensors_file": "tensors.npz",
        "tensors_digest": tensors_digest,
        "tensors": tensors,
    }
    fd, temporary_name = tempfile.mkstemp(prefix=".manifest-", suffix=".json", dir=root)
    temporary_manifest = Path(temporary_name)
    try:
        with os.fdopen(fd, "w") as handle:
            json.dump(manifest, handle, sort_keys=True, separators=(",", ":"))
            handle.write("\n")
        os.replace(temporary_manifest, root / "manifest.json")
    finally:
        temporary_manifest.unlink(missing_ok=True)
    return root

load_runtime_state

View source

def load_runtime_state(path: str | Path, *, device: str | torch.device='cpu') -> GraphRuntimeState

Source docstring:

Load and authenticate a portable graph-runtime state artifact.
ParameterAnnotationDefaultMeaning
pathstr | PathrequiredFilesystem source or destination path, as described below.
devicestr | torch.device'cpu'Requested or resolved tensor execution device.

Return annotation: GraphRuntimeState.

Return expressions (branch-dependent; names refer to the linked implementation):

GraphRuntimeState(signature=manifest['signature'], compatibility=manifest['compatibility'], completed_steps=int(manifest['completed_steps']), voltages=groups['voltages'], refractory=groups['refractory'], conductances=groups['conductances'], population_histories=groups['population_histories'], input_histories=groups['input_histories'])

Explicit exceptions in this implementation; called helpers may raise additional errors:

Explicit exception expression
ValueError(f"unsupported runtime-state schema: {manifest.get('schema')}")
ValueError(f"runtime-state tensors digest expected {manifest.get('tensors_digest')}, got {actual_digest}")
ValueError(f'runtime-state tensor keys expected {sorted(expected_keys)}, got {sorted(archive.files)}')
ValueError(f"runtime-state tensor {row['group']}.{row['name']} metadata does not match tensors.npz")
Implementation
def load_runtime_state(
    path: str | Path,
    *,
    device: str | torch.device = "cpu",
) -> GraphRuntimeState:
    """Load and authenticate a portable graph-runtime state artifact."""
    root = Path(path)
    manifest = json.loads((root / "manifest.json").read_text())
    if (
        manifest.get("schema") != RUNTIME_STATE_SCHEMA
        or manifest.get("schema_version") != 1
    ):
        raise ValueError(f"unsupported runtime-state schema: {manifest.get('schema')}")
    tensors_path = root / manifest.get("tensors_file", "tensors.npz")
    actual_digest = _file_digest(tensors_path)
    if actual_digest != manifest.get("tensors_digest"):
        raise ValueError(
            f"runtime-state tensors digest expected {manifest.get('tensors_digest')}, got {actual_digest}"
        )
    groups: dict[str, dict[str, torch.Tensor]] = {
        "voltages": {},
        "refractory": {},
        "conductances": {},
        "population_histories": {},
        "input_histories": {},
    }
    with np.load(tensors_path, allow_pickle=False) as archive:
        expected_keys = {row["key"] for row in manifest["tensors"]}
        if set(archive.files) != expected_keys:
            raise ValueError(
                f"runtime-state tensor keys expected {sorted(expected_keys)}, got {sorted(archive.files)}"
            )
        for row in manifest["tensors"]:
            array = archive[row["key"]]
            if list(array.shape) != row["shape"] or str(array.dtype) != row["dtype"]:
                raise ValueError(
                    f"runtime-state tensor {row['group']}.{row['name']} metadata does not match tensors.npz"
                )
            groups[row["group"]][row["name"]] = torch.from_numpy(array.copy()).to(
                device
            )
    return GraphRuntimeState(
        signature=manifest["signature"],
        compatibility=manifest["compatibility"],
        completed_steps=int(manifest["completed_steps"]),
        voltages=groups["voltages"],
        refractory=groups["refractory"],
        conductances=groups["conductances"],
        population_histories=groups["population_histories"],
        input_histories=groups["input_histories"],
    )

write_inference_artifacts

View source

def write_inference_artifacts(path: str | Path, result: ExecutionResult, *, graph: Mapping[str, Any], seed: int) -> Mapping[str, Any]

Source docstring:

Persist graph inference tensors and a digest-bound cache manifest.
ParameterAnnotationDefaultMeaning
pathstr | PathrequiredFilesystem source or destination path, as described below.
resultExecutionResultrequiredDefined by the source contract and implementation below.
graphMapping[str, Any]requiredSerialized graph mapping.
seedintrequiredSeed controlling this operation’s random stream.

Return annotation: Mapping[str, Any].

Return expressions (branch-dependent; names refer to the linked implementation):

manifest
Implementation
def write_inference_artifacts(
    path: str | Path,
    result: ExecutionResult,
    *,
    graph: Mapping[str, Any],
    seed: int,
) -> Mapping[str, Any]:
    """Persist graph inference tensors and a digest-bound cache manifest."""
    root = Path(path)
    root.mkdir(parents=True, exist_ok=True)
    payloads = {
        "recording.npz": result.recordings,
        "outputs.npz": result.outputs,
        "parameters.npz": result.parameters,
    }
    files = []
    for filename, tensors in payloads.items():
        arrays = {
            name: value.detach().cpu().numpy()
            for name, value in sorted(tensors.items())
        }
        destination = root / filename
        np.savez_compressed(destination, **arrays)
        files.append(
            {
                "path": filename,
                "digest": _file_digest(destination),
                "arrays": [
                    {
                        "name": name,
                        "shape": list(value.shape),
                        "dtype": str(value.dtype),
                    }
                    for name, value in arrays.items()
                ],
            }
        )
    metrics_path = root / "metrics.json"
    metrics_path.write_text(json.dumps(result.metrics, indent=2, sort_keys=True) + "\n")
    files.append({"path": "metrics.json", "digest": _file_digest(metrics_path)})
    request = {
        "seed": int(seed),
        "execution_protocol": result.metrics.get("execution_protocol"),
        "checkpoint": result.metrics.get("checkpoint"),
        "inference_overrides": result.metrics.get("inference_overrides"),
        "inference_interventions": result.metrics.get("inference_interventions"),
        "recording": result.metrics.get("recording"),
        "device": result.metrics.get("device"),
        "source_graph_digest": result.metrics.get("source_graph_digest"),
        "effective_graph_digest": result.metrics.get("effective_graph_digest"),
    }
    manifest = {
        "schema": INFERENCE_ARTIFACT_SCHEMA,
        "schema_version": 1,
        "graph_digest": _json_digest(graph),
        "request_seed": int(seed),
        "request_digest": _json_digest(request),
        "files": files,
    }
    manifest["artifact_digest"] = _json_digest(manifest)
    (root / "inference-manifest.json").write_text(
        json.dumps(manifest, indent=2, sort_keys=True) + "\n"
    )
    return manifest

validate_inference_artifacts

View source

def validate_inference_artifacts(path: str | Path, *, graph: Mapping[str, Any] | None=None, seed: int | None=None) -> Mapping[str, Any]

Source docstring:

Authenticate a persisted graph inference artifact set before reuse.
ParameterAnnotationDefaultMeaning
pathstr | PathrequiredFilesystem source or destination path, as described below.
graphMapping[str, Any] | NoneNoneSerialized graph mapping.
seedint | NoneNoneSeed controlling this operation’s random stream.

Return annotation: Mapping[str, Any].

Return expressions (branch-dependent; names refer to the linked implementation):

manifest

Explicit exceptions in this implementation; called helpers may raise additional errors:

Explicit exception expression
ValueError(f"unsupported inference-artifact schema: {manifest.get('schema')}")
ValueError(f'inference artifact manifest digest expected {claimed_artifact_digest}, got {actual_artifact_digest}')
ValueError(f"inference artifact request seed expected {int(seed)}, got {manifest.get('request_seed')}")
ValueError('inference artifact manifest file set is incomplete')
ValueError(f"inference artifact request digest expected {manifest.get('request_digest')}, got {actual_request_digest}")
ValueError(f"inference artifact graph digest expected {expected_graph_digest}, got {manifest.get('graph_digest')}")
ValueError(f'inference artifact path must be a filename: {filename}')
ValueError(f"inference artifact {filename} digest expected {row.get('digest')}, got {actual_digest}")
ValueError(f'inference artifact {filename} array inventory does not match manifest')
Implementation
def validate_inference_artifacts(
    path: str | Path,
    *,
    graph: Mapping[str, Any] | None = None,
    seed: int | None = None,
) -> Mapping[str, Any]:
    """Authenticate a persisted graph inference artifact set before reuse."""
    root = Path(path)
    manifest_path = root / "inference-manifest.json"
    manifest = json.loads(manifest_path.read_text())
    if (
        manifest.get("schema") != INFERENCE_ARTIFACT_SCHEMA
        or manifest.get("schema_version") != 1
    ):
        raise ValueError(
            f"unsupported inference-artifact schema: {manifest.get('schema')}"
        )
    claimed_artifact_digest = manifest.get("artifact_digest")
    unsigned = dict(manifest)
    unsigned.pop("artifact_digest", None)
    actual_artifact_digest = _json_digest(unsigned)
    if claimed_artifact_digest != actual_artifact_digest:
        raise ValueError(
            f"inference artifact manifest digest expected {claimed_artifact_digest}, got {actual_artifact_digest}"
        )
    if graph is not None:
        expected_graph_digest = _json_digest(graph)
        if manifest.get("graph_digest") != expected_graph_digest:
            raise ValueError(
                f"inference artifact graph digest expected {expected_graph_digest}, got {manifest.get('graph_digest')}"
            )
    if seed is not None and manifest.get("request_seed") != int(seed):
        raise ValueError(
            f"inference artifact request seed expected {int(seed)}, got {manifest.get('request_seed')}"
        )
    expected_files = {
        "recording.npz",
        "outputs.npz",
        "parameters.npz",
        "metrics.json",
    }
    rows = manifest.get("files", [])
    if {row.get("path") for row in rows} != expected_files:
        raise ValueError("inference artifact manifest file set is incomplete")
    for row in rows:
        filename = row["path"]
        if Path(filename).name != filename:
            raise ValueError(f"inference artifact path must be a filename: {filename}")
        payload_path = root / filename
        actual_digest = _file_digest(payload_path)
        if actual_digest != row.get("digest"):
            raise ValueError(
                f"inference artifact {filename} digest expected {row.get('digest')}, got {actual_digest}"
            )
        if filename.endswith(".npz"):
            loaded = np.load(payload_path, allow_pickle=False)
            try:
                actual_arrays = {
                    name: {
                        "name": name,
                        "shape": list(loaded[name].shape),
                        "dtype": str(loaded[name].dtype),
                    }
                    for name in sorted(loaded.files)
                }
            finally:
                loaded.close()
            expected_arrays = {row["name"]: row for row in row.get("arrays", [])}
            if actual_arrays != expected_arrays:
                raise ValueError(
                    f"inference artifact {filename} array inventory does not match manifest"
                )
    metrics = json.loads((root / "metrics.json").read_text())
    request = {
        "seed": int(manifest["request_seed"]),
        "execution_protocol": metrics.get("execution_protocol"),
        "checkpoint": metrics.get("checkpoint"),
        "inference_overrides": metrics.get("inference_overrides"),
        "inference_interventions": metrics.get("inference_interventions"),
        "recording": metrics.get("recording"),
        "device": metrics.get("device"),
        "source_graph_digest": metrics.get("source_graph_digest"),
        "effective_graph_digest": metrics.get("effective_graph_digest"),
    }
    actual_request_digest = _json_digest(request)
    if actual_request_digest != manifest.get("request_digest"):
        raise ValueError(
            f"inference artifact request digest expected {manifest.get('request_digest')}, got {actual_request_digest}"
        )
    return manifest

derive_inference_products

View source

def derive_inference_products(source: str | Path, destination: str | Path, *, logits_id: str, labels: np.ndarray | torch.Tensor, spike_recordings: Sequence[str]=()) -> Mapping[str, Any]

Source docstring:

Derive named accuracy, per-cell rates, and sparse rasters from a cache.
ParameterAnnotationDefaultMeaning
sourcestr | PathrequiredDefined by the source contract and implementation below.
destinationstr | PathrequiredDefined by the source contract and implementation below.
logits_idstrrequiredDefined by the source contract and implementation below.
labelsnp.ndarray | torch.TensorrequiredDefined by the source contract and implementation below.
spike_recordingsSequence[str]()Defined by the source contract and implementation below.

Return annotation: Mapping[str, Any].

Return expressions (branch-dependent; names refer to the linked implementation):

{**summary, 'artifact_digest': manifest['artifact_digest']}

Explicit exceptions in this implementation; called helpers may raise additional errors:

Explicit exception expression
ValueError(f'derived inference destination already exists: {root}')
ValueError(f'inference outputs do not contain logits {logits_id!r}')
ValueError(f'inference logits {logits_id} must have floating shape [batch, classes]')
ValueError(f'inference labels shape expected [{logits.shape[0]}], got {list(label_values.shape)}')
ValueError('inference labels must use an integer dtype')
ValueError(f'inference labels must be in [0, {logits.shape[1]})')
ValueError('inference artifact requires a positive execution duration')
ValueError(f'inference recordings do not contain spikes {recording_id!r}')
ValueError(f'inference spike recording {recording_id} must be binary [time, batch, cells]')
Implementation
def derive_inference_products(
    source: str | Path,
    destination: str | Path,
    *,
    logits_id: str,
    labels: np.ndarray | torch.Tensor,
    spike_recordings: Sequence[str] = (),
) -> Mapping[str, Any]:
    """Derive named accuracy, per-cell rates, and sparse rasters from a cache."""
    source_root = Path(source)
    source_manifest = validate_inference_artifacts(source_root)
    output_file = np.load(source_root / "outputs.npz", allow_pickle=False)
    recording_file = np.load(source_root / "recording.npz", allow_pickle=False)
    try:
        if logits_id not in output_file.files:
            raise ValueError(f"inference outputs do not contain logits {logits_id!r}")
        logits = np.asarray(output_file[logits_id])
        if logits.ndim != 2 or not np.issubdtype(logits.dtype, np.floating):
            raise ValueError(
                f"inference logits {logits_id} must have floating shape [batch, classes]"
            )
        label_values = (
            labels.detach().cpu().numpy()
            if isinstance(labels, torch.Tensor)
            else np.asarray(labels)
        )
        if label_values.ndim != 1 or label_values.shape[0] != logits.shape[0]:
            raise ValueError(
                f"inference labels shape expected [{logits.shape[0]}], got {list(label_values.shape)}"
            )
        if not np.issubdtype(label_values.dtype, np.integer):
            raise ValueError("inference labels must use an integer dtype")
        if np.any(label_values < 0) or np.any(label_values >= logits.shape[1]):
            raise ValueError(f"inference labels must be in [0, {logits.shape[1]})")
        metrics = json.loads((source_root / "metrics.json").read_text())
        timing = metrics.get("execution_protocol", {}).get("timing", {})
        duration_s = float(timing.get("duration_ms", 0.0)) / 1000.0
        if not math.isfinite(duration_s) or duration_s <= 0:
            raise ValueError(
                "inference artifact requires a positive execution duration"
            )
        rates: dict[str, np.ndarray] = {}
        rasters: dict[str, np.ndarray] = {}
        recording_rows = []
        for recording_id in spike_recordings:
            if recording_id not in recording_file.files:
                raise ValueError(
                    f"inference recordings do not contain spikes {recording_id!r}"
                )
            spikes = np.asarray(recording_file[recording_id])
            if (
                spikes.ndim != 3
                or spikes.shape[1] != logits.shape[0]
                or not np.all((spikes == 0) | (spikes == 1))
            ):
                raise ValueError(
                    f"inference spike recording {recording_id} must be binary [time, batch, cells]"
                )
            rates[recording_id] = spikes.sum(axis=0, dtype=np.float64) / duration_s
            coordinates = np.argwhere(spikes != 0)
            rasters[f"{recording_id}.steps"] = coordinates[:, 0].astype(np.int64)
            rasters[f"{recording_id}.batches"] = coordinates[:, 1].astype(np.int64)
            rasters[f"{recording_id}.cells"] = coordinates[:, 2].astype(np.int64)
            rasters[f"{recording_id}.shape"] = np.asarray(spikes.shape, dtype=np.int64)
            recording_rows.append(
                {
                    "id": recording_id,
                    "spike_shape": list(spikes.shape),
                    "rate_shape": list(rates[recording_id].shape),
                    "spike_count": int(coordinates.shape[0]),
                }
            )
    finally:
        output_file.close()
        recording_file.close()

    root = Path(destination)
    if root.exists():
        raise ValueError(f"derived inference destination already exists: {root}")
    root.mkdir(parents=True)
    predictions = logits.argmax(axis=1).astype(np.int64)
    np.save(root / "labels.npy", label_values.astype(np.int64, copy=False))
    np.save(root / "predictions.npy", predictions)
    np.savez_compressed(root / "rates.npz", **rates)
    np.savez_compressed(root / "rasters.npz", **rasters)
    summary = {
        "schema": DERIVED_INFERENCE_SCHEMA,
        "source_artifact_digest": source_manifest["artifact_digest"],
        "logits_id": logits_id,
        "logits_shape": list(logits.shape),
        "labels_shape": list(label_values.shape),
        "accuracy": float(np.mean(predictions == label_values)),
        "duration_s": duration_s,
        "spike_recordings": recording_rows,
    }
    (root / "summary.json").write_text(
        json.dumps(summary, indent=2, sort_keys=True) + "\n"
    )
    files = []
    for filename in (
        "labels.npy",
        "predictions.npy",
        "rates.npz",
        "rasters.npz",
        "summary.json",
    ):
        path = root / filename
        files.append({"path": filename, "digest": _file_digest(path)})
    manifest = {
        "schema": DERIVED_INFERENCE_SCHEMA,
        "schema_version": 1,
        "source_artifact_digest": source_manifest["artifact_digest"],
        "files": files,
    }
    manifest["artifact_digest"] = _json_digest(manifest)
    (root / "derived-manifest.json").write_text(
        json.dumps(manifest, indent=2, sort_keys=True) + "\n"
    )
    return {**summary, "artifact_digest": manifest["artifact_digest"]}

validate_derived_inference_products

View source

def validate_derived_inference_products(path: str | Path, *, source_artifact_digest: str | None=None) -> Mapping[str, Any]

Source docstring:

Authenticate one derived-inference directory before downstream reuse.
ParameterAnnotationDefaultMeaning
pathstr | PathrequiredFilesystem source or destination path, as described below.
source_artifact_digeststr | NoneNoneDefined by the source contract and implementation below.

Return annotation: Mapping[str, Any].

Return expressions (branch-dependent; names refer to the linked implementation):

manifest

Explicit exceptions in this implementation; called helpers may raise additional errors:

Explicit exception expression
ValueError(f"unsupported derived-inference schema: {manifest.get('schema')}")
ValueError(f'derived inference manifest digest expected {claimed}, got {actual}')
ValueError('derived inference source artifact digest does not match expected cache')
ValueError('derived inference manifest file set is incomplete')
ValueError('derived inference summary identity does not match manifest')
ValueError(f"derived inference {row['path']} digest expected {row.get('digest')}, got {actual_digest}")
Implementation
def validate_derived_inference_products(
    path: str | Path, *, source_artifact_digest: str | None = None
) -> Mapping[str, Any]:
    """Authenticate one derived-inference directory before downstream reuse."""
    root = Path(path)
    manifest = json.loads((root / "derived-manifest.json").read_text())
    if (
        manifest.get("schema") != DERIVED_INFERENCE_SCHEMA
        or manifest.get("schema_version") != 1
    ):
        raise ValueError(
            f"unsupported derived-inference schema: {manifest.get('schema')}"
        )
    claimed = manifest.get("artifact_digest")
    unsigned = dict(manifest)
    unsigned.pop("artifact_digest", None)
    actual = _json_digest(unsigned)
    if claimed != actual:
        raise ValueError(
            f"derived inference manifest digest expected {claimed}, got {actual}"
        )
    if (
        source_artifact_digest is not None
        and manifest.get("source_artifact_digest") != source_artifact_digest
    ):
        raise ValueError(
            "derived inference source artifact digest does not match expected cache"
        )
    expected_files = {
        "labels.npy",
        "predictions.npy",
        "rates.npz",
        "rasters.npz",
        "summary.json",
    }
    rows = manifest.get("files", [])
    if {row.get("path") for row in rows} != expected_files:
        raise ValueError("derived inference manifest file set is incomplete")
    for row in rows:
        path = root / row["path"]
        actual_digest = _file_digest(path)
        if actual_digest != row.get("digest"):
            raise ValueError(
                f"derived inference {row['path']} digest expected {row.get('digest')}, got {actual_digest}"
            )
    summary = json.loads((root / "summary.json").read_text())
    if summary.get("schema") != DERIVED_INFERENCE_SCHEMA or summary.get(
        "source_artifact_digest"
    ) != manifest.get("source_artifact_digest"):
        raise ValueError("derived inference summary identity does not match manifest")
    return manifest

save_training_checkpoint

View source

def save_training_checkpoint(path: str | Path, checkpoint: TrainingCheckpoint) -> Path

Source docstring:

Atomically write a named, authenticated graph-training checkpoint.
ParameterAnnotationDefaultMeaning
pathstr | PathrequiredFilesystem source or destination path, as described below.
checkpointTrainingCheckpointrequiredTraining checkpoint record or authenticated checkpoint path.

Return annotation: Path.

Return expressions (branch-dependent; names refer to the linked implementation):

root

Explicit exceptions in this implementation; called helpers may raise additional errors:

Explicit exception expression
ValueError('training checkpoint CPU RNG state must be one-dimensional uint8')
ValueError(f'training checkpoint has unsupported RNG backend {checkpoint.rng_backend!r}')
ValueError('CPU training checkpoint cannot contain accelerator RNG states')
ValueError('CUDA training checkpoint RNG devices must be contiguous from cuda:0')
ValueError('MPS training checkpoint requires exactly the mps RNG state')
ValueError(f'training checkpoint {name} RNG state must be one-dimensional uint8')
Implementation
def save_training_checkpoint(path: str | Path, checkpoint: TrainingCheckpoint) -> Path:
    """Atomically write a named, authenticated graph-training checkpoint."""
    if checkpoint.rng_state.dtype != torch.uint8 or checkpoint.rng_state.ndim != 1:
        raise ValueError(
            "training checkpoint CPU RNG state must be one-dimensional uint8"
        )
    if checkpoint.rng_backend not in {"cpu", "cuda", "mps"}:
        raise ValueError(
            f"training checkpoint has unsupported RNG backend {checkpoint.rng_backend!r}"
        )
    devices = sorted(checkpoint.accelerator_rng_states)
    if checkpoint.rng_backend == "cpu" and devices:
        raise ValueError(
            "CPU training checkpoint cannot contain accelerator RNG states"
        )
    if checkpoint.rng_backend == "cuda" and (
        not devices or devices != [f"cuda:{index}" for index in range(len(devices))]
    ):
        raise ValueError(
            "CUDA training checkpoint RNG devices must be contiguous from cuda:0"
        )
    if checkpoint.rng_backend == "mps" and devices != ["mps"]:
        raise ValueError("MPS training checkpoint requires exactly the mps RNG state")
    for name, state in checkpoint.accelerator_rng_states.items():
        if state.dtype != torch.uint8 or state.ndim != 1:
            raise ValueError(
                f"training checkpoint {name} RNG state must be one-dimensional uint8"
            )
    root = Path(path)
    root.mkdir(parents=True, exist_ok=True)
    arrays: dict[str, np.ndarray] = {}
    tensors: list[dict[str, Any]] = []

    def append(group: str, name: str, value: torch.Tensor, state: str | None = None):
        key = f"tensor_{len(arrays):04d}"
        tensor = value.detach().cpu().contiguous()
        arrays[key] = tensor.numpy()
        row = {
            "group": group,
            "name": name,
            "key": key,
            "shape": list(tensor.shape),
            "dtype": str(tensor.dtype).removeprefix("torch."),
        }
        if state is not None:
            row["state"] = state
        tensors.append(row)

    for name in sorted(checkpoint.parameters):
        append("parameters", name, checkpoint.parameters[name])
    optimizer_scalars: dict[str, dict[str, Any]] = {}
    for name in sorted(checkpoint.optimizer_state):
        optimizer_scalars[name] = {}
        for state, value in sorted(checkpoint.optimizer_state[name].items()):
            if isinstance(value, torch.Tensor):
                append("optimizer", name, value, state)
            else:
                optimizer_scalars[name][state] = value
    append("rng", "cpu", checkpoint.rng_state)
    for name in sorted(checkpoint.accelerator_rng_states):
        append("rng", name, checkpoint.accelerator_rng_states[name])
    fd, temporary_name = tempfile.mkstemp(prefix=".tensors-", suffix=".npz", dir=root)
    os.close(fd)
    temporary_tensors = Path(temporary_name)
    try:
        np.savez_compressed(temporary_tensors, **arrays)
        tensors_digest = _file_digest(temporary_tensors)
        os.replace(temporary_tensors, root / "tensors.npz")
    finally:
        temporary_tensors.unlink(missing_ok=True)
    manifest = {
        "schema": TRAINING_CHECKPOINT_SCHEMA,
        "schema_version": 2,
        "backend": "tools/snnsim.graph-training/v1",
        "rng_backend": checkpoint.rng_backend,
        "accelerator_rng_devices": sorted(checkpoint.accelerator_rng_states),
        "graph_digest": checkpoint.graph_digest,
        "training_digest": checkpoint.training_digest,
        "completed_updates": checkpoint.completed_updates,
        "selected_loss": checkpoint.selected_loss,
        "execution_protocol": checkpoint.execution_protocol,
        "initialization": checkpoint.initialization,
        "data_state": checkpoint.data_state,
        "optimizer_scalars": optimizer_scalars,
        "tensors_file": "tensors.npz",
        "tensors_digest": tensors_digest,
        "tensors": tensors,
    }
    fd, temporary_name = tempfile.mkstemp(prefix=".manifest-", suffix=".json", dir=root)
    temporary_manifest = Path(temporary_name)
    try:
        with os.fdopen(fd, "w") as handle:
            json.dump(manifest, handle, sort_keys=True, separators=(",", ":"))
            handle.write("\n")
        os.replace(temporary_manifest, root / "manifest.json")
    finally:
        temporary_manifest.unlink(missing_ok=True)
    return root

load_training_checkpoint

View source

def load_training_checkpoint(path: str | Path, *, device: str | torch.device='cpu') -> TrainingCheckpoint

Source docstring:

Load and authenticate a portable named graph-training checkpoint.
ParameterAnnotationDefaultMeaning
pathstr | PathrequiredFilesystem source or destination path, as described below.
devicestr | torch.device'cpu'Requested or resolved tensor execution device.

Return annotation: TrainingCheckpoint.

Return expressions (branch-dependent; names refer to the linked implementation):

TrainingCheckpoint(graph_digest=manifest['graph_digest'], training_digest=manifest['training_digest'], completed_updates=int(manifest['completed_updates']), selected_loss=manifest.get('selected_loss'), execution_protocol=manifest['execution_protocol'], initialization=manifest['initialization'], parameters=parameters, optimizer_state=optimizer_state, rng_state=rng_state, rng_backend=rng_backend, accelerator_rng_states=accelerator_rng_states, data_state=manifest.get('data_state', {}))

Explicit exceptions in this implementation; called helpers may raise additional errors:

Explicit exception expression
ValueError(f"unsupported training-checkpoint schema: {manifest.get('schema')}")
ValueError(f"training-checkpoint tensors digest expected {manifest.get('tensors_digest')}, got {actual_digest}")
ValueError('training checkpoint is missing CPU RNG state')
ValueError(f'training-checkpoint tensor keys expected {sorted(expected_keys)}, got {sorted(archive.files)}')
ValueError(f'training checkpoint has unsupported RNG backend {rng_backend!r}')
ValueError('training checkpoint accelerator RNG device inventory does not match tensors')
ValueError('CPU training checkpoint cannot contain accelerator RNG states')
ValueError('CUDA training checkpoint is missing accelerator RNG states')
ValueError('MPS training checkpoint requires exactly the mps RNG state')
ValueError(f"training-checkpoint tensor {row['group']}.{row['name']} metadata does not match tensors.npz")
ValueError('CUDA training checkpoint RNG devices must be contiguous from cuda:0')
ValueError('training checkpoint CPU RNG state must be one-dimensional uint8')
ValueError(f"unsupported training-checkpoint tensor group {row['group']}")
ValueError(f"training checkpoint {row['name']} RNG state must be one-dimensional uint8")
Implementation
def load_training_checkpoint(
    path: str | Path, *, device: str | torch.device = "cpu"
) -> TrainingCheckpoint:
    """Load and authenticate a portable named graph-training checkpoint."""
    root = Path(path)
    manifest = json.loads((root / "manifest.json").read_text())
    if manifest.get("schema") != TRAINING_CHECKPOINT_SCHEMA or manifest.get(
        "schema_version"
    ) not in {1, 2}:
        raise ValueError(
            f"unsupported training-checkpoint schema: {manifest.get('schema')}"
        )
    tensors_path = root / manifest.get("tensors_file", "tensors.npz")
    actual_digest = _file_digest(tensors_path)
    if actual_digest != manifest.get("tensors_digest"):
        raise ValueError(
            f"training-checkpoint tensors digest expected {manifest.get('tensors_digest')}, got {actual_digest}"
        )
    parameters: dict[str, torch.Tensor] = {}
    optimizer_state: dict[str, dict[str, Any]] = {
        name: dict(values)
        for name, values in manifest.get("optimizer_scalars", {}).items()
    }
    rng_state = None
    accelerator_rng_states: dict[str, torch.Tensor] = {}
    with np.load(tensors_path, allow_pickle=False) as archive:
        expected_keys = {row["key"] for row in manifest["tensors"]}
        if set(archive.files) != expected_keys:
            raise ValueError(
                f"training-checkpoint tensor keys expected {sorted(expected_keys)}, got {sorted(archive.files)}"
            )
        for row in manifest["tensors"]:
            array = archive[row["key"]]
            if list(array.shape) != row["shape"] or str(array.dtype) != row["dtype"]:
                raise ValueError(
                    f"training-checkpoint tensor {row['group']}.{row['name']} metadata does not match tensors.npz"
                )
            value = torch.from_numpy(array.copy())
            if row["group"] == "parameters":
                parameters[row["name"]] = value.to(device)
            elif row["group"] == "optimizer":
                optimizer_state.setdefault(row["name"], {})[row["state"]] = value.to(
                    device
                )
            elif row["group"] == "rng" and row["name"] == "cpu":
                if value.dtype != torch.uint8 or value.ndim != 1:
                    raise ValueError(
                        "training checkpoint CPU RNG state must be one-dimensional uint8"
                    )
                rng_state = value.cpu()
            elif row["group"] == "rng":
                if value.dtype != torch.uint8 or value.ndim != 1:
                    raise ValueError(
                        f"training checkpoint {row['name']} RNG state must be one-dimensional uint8"
                    )
                accelerator_rng_states[row["name"]] = value.cpu()
            else:
                raise ValueError(
                    f"unsupported training-checkpoint tensor group {row['group']}"
                )
    if rng_state is None:
        raise ValueError("training checkpoint is missing CPU RNG state")
    schema_version = int(manifest["schema_version"])
    rng_backend = manifest.get("rng_backend", "cpu")
    declared_devices = set(manifest.get("accelerator_rng_devices", []))
    if schema_version == 2:
        if rng_backend not in {"cpu", "cuda", "mps"}:
            raise ValueError(
                f"training checkpoint has unsupported RNG backend {rng_backend!r}"
            )
        if declared_devices != set(accelerator_rng_states):
            raise ValueError(
                "training checkpoint accelerator RNG device inventory does not match tensors"
            )
        if rng_backend == "cpu" and accelerator_rng_states:
            raise ValueError(
                "CPU training checkpoint cannot contain accelerator RNG states"
            )
        if rng_backend == "cuda" and not accelerator_rng_states:
            raise ValueError(
                "CUDA training checkpoint is missing accelerator RNG states"
            )
        if rng_backend == "cuda":
            expected_cuda = [
                f"cuda:{index}" for index in range(len(accelerator_rng_states))
            ]
            if sorted(accelerator_rng_states) != expected_cuda:
                raise ValueError(
                    "CUDA training checkpoint RNG devices must be contiguous from cuda:0"
                )
        if rng_backend == "mps" and set(accelerator_rng_states) != {"mps"}:
            raise ValueError(
                "MPS training checkpoint requires exactly the mps RNG state"
            )
    return TrainingCheckpoint(
        graph_digest=manifest["graph_digest"],
        training_digest=manifest["training_digest"],
        completed_updates=int(manifest["completed_updates"]),
        selected_loss=manifest.get("selected_loss"),
        execution_protocol=manifest["execution_protocol"],
        initialization=manifest["initialization"],
        parameters=parameters,
        optimizer_state=optimizer_state,
        rng_state=rng_state,
        rng_backend=rng_backend,
        accelerator_rng_states=accelerator_rng_states,
        data_state=manifest.get("data_state", {}),
    )

capture_training_rng_state

View source

def capture_training_rng_state(device: str | torch.device) -> tuple[str, dict[str, torch.Tensor]]

Source docstring:

Capture every stochastic stream required for exact resume on one backend.
ParameterAnnotationDefaultMeaning
devicestr | torch.devicerequiredRequested or resolved tensor execution device.

Return annotation: tuple[str, dict[str, torch.Tensor]].

Return expressions (branch-dependent; names refer to the linked implementation):

('cuda', {f'cuda:{index}': state.detach().cpu().clone() for index, state in enumerate(states)})
('mps', {'mps': torch.mps.get_rng_state().detach().cpu().clone()})
('cpu', {})

Explicit exceptions in this implementation; called helpers may raise additional errors:

Explicit exception expression
ValueError(f'CUDA RNG capture expected {expected} device states, got {len(states)}')
ValueError('MPS RNG capture requires an available MPS backend')
Implementation
def capture_training_rng_state(
    device: str | torch.device,
) -> tuple[str, dict[str, torch.Tensor]]:
    """Capture every stochastic stream required for exact resume on one backend."""
    name = str(device).lower()
    if name.startswith("cuda"):
        states = torch.cuda.get_rng_state_all()
        expected = torch.cuda.device_count()
        if len(states) != expected or expected <= 0:
            raise ValueError(
                f"CUDA RNG capture expected {expected} device states, got {len(states)}"
            )
        return "cuda", {
            f"cuda:{index}": state.detach().cpu().clone()
            for index, state in enumerate(states)
        }
    if name == "mps":
        if not torch.backends.mps.is_available():
            raise ValueError("MPS RNG capture requires an available MPS backend")
        return "mps", {"mps": torch.mps.get_rng_state().detach().cpu().clone()}
    return "cpu", {}

restore_training_rng_state

View source

def restore_training_rng_state(checkpoint: TrainingCheckpoint, device: str | torch.device) -> None

Source docstring:

Restore CPU and exact-matching accelerator streams or fail closed.
ParameterAnnotationDefaultMeaning
checkpointTrainingCheckpointrequiredTraining checkpoint record or authenticated checkpoint path.
devicestr | torch.devicerequiredRequested or resolved tensor execution device.

Return annotation: None.

Explicit exceptions in this implementation; called helpers may raise additional errors:

Explicit exception expression
ValueError(f'training checkpoint RNG backend {checkpoint.rng_backend} cannot resume on {requested_backend}')
ValueError('training checkpoint CUDA RNG topology does not match available devices')
ValueError('MPS RNG restore requires an available MPS backend')
ValueError('training checkpoint MPS RNG state is missing')
ValueError('CPU resume cannot restore accelerator RNG states')
Implementation
def restore_training_rng_state(
    checkpoint: TrainingCheckpoint, device: str | torch.device
) -> None:
    """Restore CPU and exact-matching accelerator streams or fail closed."""
    name = str(device).lower()
    requested_backend = (
        "cuda" if name.startswith("cuda") else "mps" if name == "mps" else "cpu"
    )
    if checkpoint.rng_backend != requested_backend:
        raise ValueError(
            f"training checkpoint RNG backend {checkpoint.rng_backend} cannot resume on {requested_backend}"
        )
    if requested_backend == "cuda":
        expected = [f"cuda:{index}" for index in range(torch.cuda.device_count())]
        if sorted(checkpoint.accelerator_rng_states) != expected:
            raise ValueError(
                "training checkpoint CUDA RNG topology does not match available devices"
            )
        torch.set_rng_state(checkpoint.rng_state.cpu())
        torch.cuda.set_rng_state_all(
            [checkpoint.accelerator_rng_states[key].cpu() for key in expected]
        )
    elif requested_backend == "mps":
        if not torch.backends.mps.is_available():
            raise ValueError("MPS RNG restore requires an available MPS backend")
        if set(checkpoint.accelerator_rng_states) != {"mps"}:
            raise ValueError("training checkpoint MPS RNG state is missing")
        torch.set_rng_state(checkpoint.rng_state.cpu())
        torch.mps.set_rng_state(checkpoint.accelerator_rng_states["mps"].cpu())
    else:
        if checkpoint.accelerator_rng_states:
            raise ValueError("CPU resume cannot restore accelerator RNG states")
        torch.set_rng_state(checkpoint.rng_state.cpu())

legacy_parameter_map_v1

View source

def legacy_parameter_map_v1(graph: Mapping[str, Any]) -> dict[str, str]

Source docstring:

Map the supported one-layer legacy COBANet by semantic graph roles.
ParameterAnnotationDefaultMeaning
graphMapping[str, Any]requiredSerialized graph mapping.

Return annotation: dict[str, str].

Return expressions (branch-dependent; names refer to the linked implementation):

dict(sorted(mapping.items()))

Explicit exceptions in this implementation; called helpers may raise additional errors:

Explicit exception expression
ValueError(f'legacy parameter mapping must be complete; missing={missing}, extra={extra}')
ValueError(f"legacy parameter mapping requires one parameter on projection {projection.get('id')}")
ValueError(f'legacy parameter mapping duplicates role {legacy}')
ValueError(f"legacy parameter mapping cannot classify projection {projection['id']}")
ValueError(f"legacy parameter mapping cannot classify recurrent projection {projection['id']}")
Implementation
def legacy_parameter_map_v1(graph: Mapping[str, Any]) -> dict[str, str]:
    """Map the supported one-layer legacy COBANet by semantic graph roles."""
    populations = {row["id"]: row for row in graph.get("populations", [])}
    input_ids = {row["id"] for row in graph.get("inputs", [])}
    mapping: dict[str, str] = {}
    recurrent_index = 1
    for projection in graph.get("projections", []):
        parameters = projection.get("parameters", [])
        if len(parameters) != 1:
            raise ValueError(
                f"legacy parameter mapping requires one parameter on projection {projection.get('id')}"
            )
        parameter = parameters[0]
        source = projection["source"].partition(".")[0]
        target = projection["target"].partition(".")[0]
        if source in input_ids:
            legacy = "W_ff.0"
        elif populations[target]["neuron"]["kind"] == "leaky_integrator":
            legacy = "W_ff.1"
        elif projection.get("connection") == "recurrent":
            source_size = int(populations[source]["size"])
            target_size = int(populations[target]["size"])
            if source == target and projection["polarity"] == "excitatory":
                legacy = f"W_ee.{recurrent_index}"
            elif source == target and projection["polarity"] == "inhibitory":
                legacy = f"W_ii.{recurrent_index}"
            elif source_size >= target_size and projection["polarity"] == "excitatory":
                legacy = f"W_ei.{recurrent_index}"
            elif source_size <= target_size and projection["polarity"] == "inhibitory":
                legacy = f"W_ie.{recurrent_index}"
            else:
                raise ValueError(
                    f"legacy parameter mapping cannot classify recurrent projection {projection['id']}"
                )
        else:
            raise ValueError(
                f"legacy parameter mapping cannot classify projection {projection['id']}"
            )
        if legacy in mapping.values():
            raise ValueError(f"legacy parameter mapping duplicates role {legacy}")
        mapping[parameter] = legacy
    for operation in graph.get("operations", []):
        if operation.get("kind") != "linear":
            continue
        for parameter in operation.get("parameters", []):
            if parameter in mapping:
                continue
            legacy = "W_ff.1"
            if legacy in mapping.values():
                raise ValueError(f"legacy parameter mapping duplicates role {legacy}")
            mapping[parameter] = legacy
    graph_parameters = {row["id"] for row in graph.get("parameters", [])}
    if set(mapping) != graph_parameters:
        missing = sorted(graph_parameters - set(mapping))
        extra = sorted(set(mapping) - graph_parameters)
        raise ValueError(
            f"legacy parameter mapping must be complete; missing={missing}, extra={extra}"
        )
    return dict(sorted(mapping.items()))

import_legacy_parameters_v1

View source

def import_legacy_parameters_v1(graph: Mapping[str, Any], state_dict: Mapping[str, torch.Tensor], *, device: str | torch.device='cpu') -> ParameterInterchange

Source docstring:

Import the exact supported one-layer legacy parameter state by semantic name.
ParameterAnnotationDefaultMeaning
graphMapping[str, Any]requiredSerialized graph mapping.
state_dictMapping[str, torch.Tensor]requiredDefined by the source contract and implementation below.
devicestr | torch.device'cpu'Requested or resolved tensor execution device.

Return annotation: ParameterInterchange.

Return expressions (branch-dependent; names refer to the linked implementation):

ParameterInterchange(parameters=dict(sorted(parameters.items())), provenance={'schema': LEGACY_PARAMETER_INTERCHANGE_SCHEMA, 'mapping_version': 1, 'direction': 'legacy_to_graph', 'mapping': mapping})

Explicit exceptions in this implementation; called helpers may raise additional errors:

Explicit exception expression
ValueError(f'legacy parameter interchange requires exact keys; missing={sorted(set(reverse) - set(state_dict))}, extra={sorted(set(state_dict) - set(reverse))}')
Implementation
def import_legacy_parameters_v1(
    graph: Mapping[str, Any],
    state_dict: Mapping[str, torch.Tensor],
    *,
    device: str | torch.device = "cpu",
) -> ParameterInterchange:
    """Import the exact supported one-layer legacy parameter state by semantic name."""
    mapping = legacy_parameter_map_v1(graph)
    reverse = {legacy: graph_name for graph_name, legacy in mapping.items()}
    if set(state_dict) != set(reverse):
        raise ValueError(
            f"legacy parameter interchange requires exact keys; missing={sorted(set(reverse) - set(state_dict))}, extra={sorted(set(state_dict) - set(reverse))}"
        )
    parameters = {
        reverse[name]: value.detach().clone().to(device)
        for name, value in state_dict.items()
    }
    _validate_interchange_parameters(graph, parameters)
    return ParameterInterchange(
        parameters=dict(sorted(parameters.items())),
        provenance={
            "schema": LEGACY_PARAMETER_INTERCHANGE_SCHEMA,
            "mapping_version": 1,
            "direction": "legacy_to_graph",
            "mapping": mapping,
        },
    )

export_legacy_parameters_v1

View source

def export_legacy_parameters_v1(graph: Mapping[str, Any], parameters: Mapping[str, torch.Tensor]) -> ParameterInterchange

Source docstring:

Export a complete supported graph parameter set under legacy state keys.
ParameterAnnotationDefaultMeaning
graphMapping[str, Any]requiredSerialized graph mapping.
parametersMapping[str, torch.Tensor]requiredDefined by the source contract and implementation below.

Return annotation: ParameterInterchange.

Return expressions (branch-dependent; names refer to the linked implementation):

ParameterInterchange(parameters=dict(sorted(exported.items())), provenance={'schema': LEGACY_PARAMETER_INTERCHANGE_SCHEMA, 'mapping_version': 1, 'direction': 'graph_to_legacy', 'mapping': mapping})
Implementation
def export_legacy_parameters_v1(
    graph: Mapping[str, Any], parameters: Mapping[str, torch.Tensor]
) -> ParameterInterchange:
    """Export a complete supported graph parameter set under legacy state keys."""
    _validate_interchange_parameters(graph, parameters)
    mapping = legacy_parameter_map_v1(graph)
    exported = {
        mapping[name]: value.detach().clone() for name, value in parameters.items()
    }
    return ParameterInterchange(
        parameters=dict(sorted(exported.items())),
        provenance={
            "schema": LEGACY_PARAMETER_INTERCHANGE_SCHEMA,
            "mapping_version": 1,
            "direction": "graph_to_legacy",
            "mapping": mapping,
        },
    )

DelayBuffer

View source

Source docstring:

Fixed causal delay used by recurrent and feedback projections.

Constructor:

DelayBuffer(self, delay_steps: int, prototype: torch.Tensor)
ParameterAnnotationDefaultMeaning
delay_stepsintrequiredDefined by the constructor implementation below.
prototypetorch.TensorrequiredDefined by the constructor implementation below.

Instance members assigned by the constructor (expressions are evaluated when constructed):

MemberAnnotationInitial expression
delay_stepsunannotateddelay_steps

Constructor/initialization exception expressions:

Explicit exception expression
ValueError('causal delay buffer requires at least one step')

DelayBuffer.read

View source

def DelayBuffer.read(self) -> torch.Tensor

Return annotation: torch.Tensor.

Return expressions (branch-dependent; names refer to the linked implementation):

self._values[0]
Implementation
def read(self) -> torch.Tensor:
        return self._values[0]

DelayBuffer.push

View source

def DelayBuffer.push(self, value: torch.Tensor) -> None
ParameterAnnotationDefaultMeaning
valuetorch.TensorrequiredDefined by the source contract and implementation below.

Return annotation: None.

Implementation
def push(self, value: torch.Tensor) -> None:
        self._values.append(value)
        self._values.pop(0)

DelayBuffer.export

View source

def DelayBuffer.export(self) -> torch.Tensor

Return annotation: torch.Tensor.

Return expressions (branch-dependent; names refer to the linked implementation):

torch.stack([value.detach().clone() for value in self._values])
Implementation
def export(self) -> torch.Tensor:
        return torch.stack([value.detach().clone() for value in self._values])

DelayBuffer.restore

View source

Decorators: classmethod.

def DelayBuffer.restore(cls, values: torch.Tensor) -> DelayBuffer
ParameterAnnotationDefaultMeaning
valuestorch.TensorrequiredDefined by the source contract and implementation below.

Return annotation: DelayBuffer.

Return expressions (branch-dependent; names refer to the linked implementation):

result

Explicit exceptions in this implementation; called helpers may raise additional errors:

Explicit exception expression
ValueError('delay history must have shape [delay, batch, ...]')
Implementation
def restore(cls, values: torch.Tensor) -> DelayBuffer:
        if values.ndim < 2 or values.shape[0] < 1:
            raise ValueError("delay history must have shape [delay, batch, ...]")
        result = cls(int(values.shape[0]), values[0])
        result._values = [value.detach().clone() for value in values.unbind(0)]
        return result
Complete class implementation
class DelayBuffer:
    """Fixed causal delay used by recurrent and feedback projections."""

    def __init__(self, delay_steps: int, prototype: torch.Tensor):
        if delay_steps < 1:
            raise ValueError("causal delay buffer requires at least one step")
        self.delay_steps = delay_steps
        self._values = [torch.zeros_like(prototype) for _ in range(delay_steps)]

    def read(self) -> torch.Tensor:
        return self._values[0]

    def push(self, value: torch.Tensor) -> None:
        self._values.append(value)
        self._values.pop(0)

    def export(self) -> torch.Tensor:
        return torch.stack([value.detach().clone() for value in self._values])

    @classmethod
    def restore(cls, values: torch.Tensor) -> DelayBuffer:
        if values.ndim < 2 or values.shape[0] < 1:
            raise ValueError("delay history must have shape [delay, batch, ...]")
        result = cls(int(values.shape[0]), values[0])
        result._values = [value.detach().clone() for value in values.unbind(0)]
        return result

plan_graph

View source

def plan_graph(graph: Mapping[str, Any]) -> GraphPlan

Check capabilities and lower the full topology before simulation, resolving dimensions, receptor polarity, integral delays and deterministic zero-delay feedforward ordering. Return GraphPlan.

ParameterAnnotationDefaultMeaning
graphMapping[str, Any]requiredSerialized graph mapping.

Return annotation: GraphPlan.

Return expressions (branch-dependent; names refer to the linked implementation):

GraphPlan(graph=graph, dt_ms=dt, populations=tuple(ordered), projections=tuple(planned), observables=tuple(graph.get('observables', [])), outputs=tuple(graph.get('outputs', [])))

Explicit exceptions in this implementation; called helpers may raise additional errors:

Explicit exception expression
ValueError(f'graph executor capability failure: {detail}')
ValueError('graph timebase dt must be positive')
ValueError(f"{row['id']}: delay {delay_ms} ms is not an integer number of dt={dt} ms steps")
ValueError('zero-delay population projections form an algebraic cycle')
ValueError(f"{row['id']}: projection parameter {parameter_id} requires unit uS, got {unit}")
Implementation
def plan_graph(graph: Mapping[str, Any]) -> GraphPlan:
    issues = graph_capability_issues(graph)
    if issues:
        detail = "; ".join(
            f"{x.element} requires {x.capability}: {x.message}" for x in issues
        )
        raise ValueError(f"graph executor capability failure: {detail}")
    dt = float(graph["timebase"]["dt"]["value"])
    if dt <= 0:
        raise ValueError("graph timebase dt must be positive")
    parameter_rows = {row["id"]: row for row in graph.get("parameters", [])}
    planned = []
    for row in graph.get("projections", []):
        for parameter_id in row.get("parameters", []):
            unit = parameter_rows.get(parameter_id, {}).get("unit")
            if unit != "uS":
                raise ValueError(
                    f"{row['id']}: projection parameter {parameter_id} requires unit uS, got {unit}"
                )
        delay = row.get("delay")
        delay_ms = 0.0 if delay is None else float(delay["value"])
        raw_steps = delay_ms / dt
        steps = int(round(raw_steps))
        if not math.isclose(raw_steps, steps, abs_tol=1e-9):
            raise ValueError(
                f"{row['id']}: delay {delay_ms} ms is not an integer number of dt={dt} ms steps"
            )
        source_owner = row["source"].partition(".")[0]
        # Recurrent/feedback edges are causal even when the author declares
        # zero additional delay. Feedforward population edges may consume the
        # current-step source after topological scheduling.
        if (
            source_owner in {p["id"] for p in graph.get("populations", [])}
            and row.get("connection") != "feedforward"
        ):
            steps = max(1, steps)
        tau = float(row["synapse"]["tau"]["value"])
        decay = (
            0.0 if row["synapse"]["kind"] == "leaky_integrator" else math.exp(-dt / tau)
        )
        planned.append(
            PlannedProjection(
                id=row["id"],
                source=row["source"],
                target=row["target"],
                polarity=row["polarity"],
                decay=decay,
                delay_steps=steps,
                parameter=row["parameters"][0],
                enabled=row.get("enabled", True),
            )
        )
    populations = list(graph.get("populations", []))
    population_ids = {p["id"] for p in populations}
    zero_edges = [
        (p.source.partition(".")[0], p.target.partition(".")[0])
        for p in planned
        if p.enabled
        and p.delay_steps == 0
        and p.source.partition(".")[0] in population_ids
    ]
    ordered: list[Mapping[str, Any]] = []
    remaining = {p["id"]: p for p in populations}
    while remaining:
        ready = sorted(
            name
            for name in remaining
            if not any(dst == name and src in remaining for src, dst in zero_edges)
        )
        if not ready:
            raise ValueError(
                "zero-delay population projections form an algebraic cycle"
            )
        for name in ready:
            ordered.append(remaining.pop(name))
    return GraphPlan(
        graph=graph,
        dt_ms=dt,
        populations=tuple(ordered),
        projections=tuple(planned),
        observables=tuple(graph.get("observables", [])),
        outputs=tuple(graph.get("outputs", [])),
    )

GraphExecutor

View source

Source docstring:

Dense graph executor whose graph topology is lowered before simulation.

Bases: nn.Module. Inherited third-party framework APIs follow their owning library.

Constructor:

GraphExecutor(self, plan: GraphPlan, *, seed: int=0, trainable_parameters: Sequence[str]=(), surrogate_slope: float=M.SURROGATE_SLOPE)
ParameterAnnotationDefaultMeaning
planGraphPlanrequiredLowered GraphPlan for the complete graph.
seedint0Seed controlling this operation’s random stream.
trainable_parametersSequence[str]()Defined by the constructor implementation below.
surrogate_slopefloatM.SURROGATE_SLOPEDefined by the constructor implementation below.

Instance members assigned by the constructor (expressions are evaluated when constructed):

MemberAnnotationInitial expression
planunannotatedplan
surrogate_slopeunannotatedfloat(surrogate_slope)
weightsunannotatednn.ParameterDict()
initialization_metadatadict[str, dict[str, Any]]{}

Constructor/initialization exception expressions:

Explicit exception expression
ValueError(f"{row['id']}: initial_zero_fraction must satisfy 0 <= fraction < 1")
ValueError(f"{row['id']}: unsupported constraint {constraint.get('kind')}")
ValueError(f"{row['id']}: exact_k zeroing requires a matrix")
ValueError(f"{row['id']}: unsupported initializer zeroing {zeroing}")
ValueError(f"{row['id']}: unsupported initializer {init['kind']}")

GraphExecutor.parameter_map

View source

def GraphExecutor.parameter_map(self) -> dict[str, torch.Tensor]

Return annotation: dict[str, torch.Tensor].

Return expressions (branch-dependent; names refer to the linked implementation):

{name.replace('__', '.'): value for name, value in self.weights.items()}
Implementation
def parameter_map(self) -> dict[str, torch.Tensor]:
        return {name.replace("__", "."): value for name, value in self.weights.items()}

GraphExecutor.forward

View source

def GraphExecutor.forward(self, inputs: Mapping[str, torch.Tensor], *, record: bool | RecordingProfile=True, recording_fields: Sequence[str] | None=None, runtime_state: GraphRuntimeState | None=None, interventions: Sequence[Mapping[str, Any]]=()) -> ExecutionResult
ParameterAnnotationDefaultMeaning
inputsMapping[str, torch.Tensor]requiredInput tensors keyed by graph input id.
recordbool | RecordingProfileTrueDefined by the source contract and implementation below.
recording_fieldsSequence[str] | NoneNoneExplicit field names to retain.
runtime_stateGraphRuntimeState | NoneNoneDynamic graph state for causal continuation; distinct from training checkpoints.
interventionsSequence[Mapping[str, Any]]()Defined by the source contract and implementation below.

Return annotation: ExecutionResult.

Return expressions (branch-dependent; names refer to the linked implementation):

ExecutionResult(executor='graph', outputs=outputs, recordings=packed, parameters={k: v.detach().clone() for k, v in self.parameter_map().items()}, final_state={f'{k}.voltage': v.detach().clone() for k, v in voltage.items()}, runtime_state=next_runtime_state, metrics={'resolved_interventions': resolved_interventions}, model=self)

Explicit exceptions in this implementation; called helpers may raise additional errors:

Explicit exception expression
ValueError(f'recording profile expected full, observables, or none; got {recording!r}')
ValueError('graph execution requires at least one input tensor')
ValueError(f'input {name} leading shape expected {(steps, batch)}, got {tuple(value.shape[:2])}')
ValueError(f'input {name} device expected {device}, got {value.device}')
ValueError(f'input {name} dtype expected {parameter_dtype}, got {value.dtype}')
ValueError(f'inference intervention {index} has unsupported kind {kind!r}')
ValueError(f'inference intervention {index} targets unknown population {population_id!r}')
ValueError(f'inference intervention {index} population {population_id!r} does not emit spikes')
ValueError(f'inference intervention repeats {kind} for population {population_id}')
ValueError(f'inference intervention {index} has unsupported fields: {unknown}')
ValueError('runtime state is incompatible with graph plan: ' + (detail or f'signature expected {expected_signature}, got {runtime_state.signature}'))
ValueError(f'unavailable recording fields: {sorted(unknown)}')
ValueError(f'{op_id}: valid-time mask leading shape expected {tuple(target.shape[:2])}, got {tuple(mask.shape[:2])}')
ValueError(f'{op_id}: valid-time mask must have shape [time, batch]')
ValueError(f'{op_id}: valid-time mask contains an empty reduction window')
ValueError(f'operation dependencies are unresolved: {unresolved}')
ValueError(f'inference intervention {index} drop probability must be finite and between zero and one')
ValueError(f'inference intervention {index} Poisson rate must be finite, non-negative, and satisfy rate times dt <= 1')
ValueError(f'runtime state {label} keys expected {sorted(shapes)}, got {sorted(values)}')
ValueError(f'runtime state voltages.{name} dtype expected {parameter_dtype}, got {value.dtype}')
ValueError(f'runtime state refractory.{name} dtype expected torch.int64, got {value.dtype}')
ValueError(f'runtime state conductances.{name} dtype expected {parameter_dtype}, got {value.dtype}')
ValueError(f'runtime state population_histories.{name} dtype expected {parameter_dtype}, got {value.dtype}')
ValueError(f'runtime state input_histories.{name} dtype expected {inputs[name].dtype}, got {value.dtype}')
ValueError(f'runtime state {label}.{name} shape expected {expected_shape}, got {tuple(value.shape)}')
ValueError(f"{op['id']}: unsupported operation {kind}")
ValueError(f"{op['id']}: valid-time mask must have shape [time, batch]")
ValueError(f"{op['id']}: valid-time mask contains zero valid duration")
ValueError(f"{op['id']}: spike-rate duration must be positive seconds")
Implementation
def forward(
        self,
        inputs: Mapping[str, torch.Tensor],
        *,
        record: bool | RecordingProfile = True,
        recording_fields: Sequence[str] | None = None,
        runtime_state: GraphRuntimeState | None = None,
        interventions: Sequence[Mapping[str, Any]] = (),
    ) -> ExecutionResult:
        recording: RecordingProfile = (
            "full" if record is True else "none" if record is False else record
        )
        if recording not in {"full", "observables", "none"}:
            raise ValueError(
                f"recording profile expected full, observables, or none; got {recording!r}"
            )
        if not inputs:
            raise ValueError("graph execution requires at least one input tensor")
        first = next(iter(inputs.values()))
        steps, batch = first.shape[:2]
        device = first.device
        parameter_dtype = (
            next(iter(self.weights.values())).dtype if self.weights else first.dtype
        )
        input_specs = {row["id"]: row for row in self.plan.graph.get("inputs", [])}
        for name, value in inputs.items():
            if value.shape[:2] != (steps, batch):
                raise ValueError(
                    f"input {name} leading shape expected {(steps, batch)}, got {tuple(value.shape[:2])}"
                )
            if value.device != device:
                raise ValueError(
                    f"input {name} device expected {device}, got {value.device}"
                )
            is_mask = input_specs.get(name, {}).get("signal_type") == "mask"
            if value.dtype != parameter_dtype and not (
                is_mask and value.dtype == torch.bool
            ):
                raise ValueError(
                    f"input {name} dtype expected {parameter_dtype}, got {value.dtype}"
                )
        populations = {p["id"]: p for p in self.plan.populations}
        resolved_interventions: list[dict[str, Any]] = []
        intervention_keys: set[tuple[str, str]] = set()
        for index, raw in enumerate(interventions):
            row = dict(raw)
            kind = str(row.get("kind", ""))
            population_id = str(row.get("population_id", ""))
            if kind not in {"drop_spikes", "add_poisson_spikes"}:
                raise ValueError(
                    f"inference intervention {index} has unsupported kind {kind!r}"
                )
            if population_id not in populations:
                raise ValueError(
                    f"inference intervention {index} targets unknown population {population_id!r}"
                )
            if populations[population_id]["neuron"]["kind"] != "coba_lif":
                raise ValueError(
                    f"inference intervention {index} population {population_id!r} does not emit spikes"
                )
            key = (kind, population_id)
            if key in intervention_keys:
                raise ValueError(
                    f"inference intervention repeats {kind} for population {population_id}"
                )
            intervention_keys.add(key)
            allowed = (
                {"kind", "population_id", "probability", "seed"}
                if kind == "drop_spikes"
                else {"kind", "population_id", "rate_hz", "seed"}
            )
            unknown = sorted(set(row) - allowed)
            if unknown:
                raise ValueError(
                    f"inference intervention {index} has unsupported fields: {unknown}"
                )
            seed = int(row.get("seed", 0))
            if kind == "drop_spikes":
                value = float(row.get("probability", float("nan")))
                if not math.isfinite(value) or not 0 <= value <= 1:
                    raise ValueError(
                        f"inference intervention {index} drop probability must be finite and between zero and one"
                    )
                resolved_interventions.append(
                    {
                        "kind": kind,
                        "population_id": population_id,
                        "probability": value,
                        "seed": seed,
                    }
                )
            else:
                value = float(row.get("rate_hz", float("nan")))
                probability = value * self.plan.dt_ms / 1000.0
                if not math.isfinite(value) or value < 0 or probability > 1:
                    raise ValueError(
                        f"inference intervention {index} Poisson rate must be finite, non-negative, and satisfy rate times dt <= 1"
                    )
                resolved_interventions.append(
                    {
                        "kind": kind,
                        "population_id": population_id,
                        "rate_hz": value,
                        "probability_per_step": probability,
                        "seed": seed,
                    }
                )
        population_history_lengths = {
            name: max(
                (
                    p.delay_steps
                    for p in self.plan.projections
                    if p.source.startswith(name + ".")
                ),
                default=1,
            )
            for name in populations
        }
        input_history_lengths = {
            row["id"]: max(
                (
                    p.delay_steps
                    for p in self.plan.projections
                    if p.source.partition(".")[0] == row["id"]
                ),
                default=0,
            )
            for row in self.plan.graph.get("inputs", [])
        }
        expected_compatibility = runtime_state_compatibility(self.plan)
        expected_signature = runtime_state_signature(self.plan)
        if runtime_state is None:
            voltage = {
                name: (
                    torch.zeros((batch, p["size"]), device=device)
                    if p["neuron"]["kind"] == "leaky_integrator"
                    else torch.full((batch, p["size"]), M.E_L, device=device)
                )
                for name, p in populations.items()
            }
            refractory = {
                name: torch.zeros((batch, p["size"]), dtype=torch.long, device=device)
                for name, p in populations.items()
            }
            spikes = {
                name: torch.zeros((batch, p["size"]), device=device)
                for name, p in populations.items()
            }
            conductance = {
                (p.id, p.polarity): torch.zeros(
                    (batch, populations[p.target.partition(".")[0]]["size"]),
                    device=device,
                )
                for p in self.plan.projections
            }
            histories = {
                name: DelayBuffer(population_history_lengths[name], value)
                for name, value in spikes.items()
            }
            input_histories = {
                name: torch.zeros(
                    (length, *inputs[name].shape[1:]),
                    device=device,
                    dtype=inputs[name].dtype,
                )
                for name, length in input_history_lengths.items()
                if length > 0
            }
            completed_steps = 0
        else:
            if runtime_state.signature != expected_signature:
                detail = _compatibility_mismatch(
                    expected_compatibility, runtime_state.compatibility
                )
                raise ValueError(
                    "runtime state is incompatible with graph plan: "
                    + (
                        detail
                        or f"signature expected {expected_signature}, got {runtime_state.signature}"
                    )
                )

            def restore_group(
                label: str,
                values: Mapping[str, torch.Tensor],
                shapes: Mapping[str, tuple[int, ...]],
            ) -> dict[str, torch.Tensor]:
                if set(values) != set(shapes):
                    raise ValueError(
                        f"runtime state {label} keys expected {sorted(shapes)}, got {sorted(values)}"
                    )
                restored = {}
                for name, expected_shape in shapes.items():
                    value = values[name]
                    if tuple(value.shape) != expected_shape:
                        raise ValueError(
                            f"runtime state {label}.{name} shape expected {expected_shape}, got {tuple(value.shape)}"
                        )
                    restored[name] = value.detach().to(device).clone()
                return restored

            pop_shapes = {
                name: (batch, int(row["size"])) for name, row in populations.items()
            }
            voltage = restore_group("voltages", runtime_state.voltages, pop_shapes)
            for name, value in voltage.items():
                if value.dtype != parameter_dtype:
                    raise ValueError(
                        f"runtime state voltages.{name} dtype expected {parameter_dtype}, got {value.dtype}"
                    )
            refractory = restore_group(
                "refractory", runtime_state.refractory, pop_shapes
            )
            for name, value in refractory.items():
                if value.dtype != torch.long:
                    raise ValueError(
                        f"runtime state refractory.{name} dtype expected torch.int64, got {value.dtype}"
                    )
            conductance_by_id = restore_group(
                "conductances",
                runtime_state.conductances,
                {
                    p.id: (batch, int(populations[p.target.partition(".")[0]]["size"]))
                    for p in self.plan.projections
                },
            )
            conductance = {
                (p.id, p.polarity): conductance_by_id[p.id]
                for p in self.plan.projections
            }
            for name, value in conductance_by_id.items():
                if value.dtype != parameter_dtype:
                    raise ValueError(
                        f"runtime state conductances.{name} dtype expected {parameter_dtype}, got {value.dtype}"
                    )
            population_history_values = restore_group(
                "population_histories",
                runtime_state.population_histories,
                {
                    name: (population_history_lengths[name], batch, int(row["size"]))
                    for name, row in populations.items()
                },
            )
            histories = {
                name: DelayBuffer.restore(value)
                for name, value in population_history_values.items()
            }
            for name, value in population_history_values.items():
                if value.dtype != parameter_dtype:
                    raise ValueError(
                        f"runtime state population_histories.{name} dtype expected {parameter_dtype}, got {value.dtype}"
                    )
            spikes = {name: histories[name]._values[-1].clone() for name in populations}
            input_histories = restore_group(
                "input_histories",
                runtime_state.input_histories,
                {
                    name: (length, *inputs[name].shape[1:])
                    for name, length in input_history_lengths.items()
                    if length > 0
                },
            )
            for name, value in input_histories.items():
                if value.dtype != inputs[name].dtype:
                    raise ValueError(
                        f"runtime state input_histories.{name} dtype expected {inputs[name].dtype}, got {value.dtype}"
                    )
            completed_steps = int(runtime_state.completed_steps)
        recordings: dict[str, list[torch.Tensor]] = {
            o["id"]: [] for o in self.plan.observables
        }
        state_recordings: dict[str, list[torch.Tensor]] = {
            f"{name}.voltage": [] for name in populations
        }
        projection_recordings: dict[str, list[torch.Tensor]] = {
            f"{p.id}.conductance": [] for p in self.plan.projections
        }
        selected_fields = None if recording_fields is None else set(recording_fields)
        if selected_fields is not None:
            available = set(recordings) if recording != "none" else set()
            if recording == "full":
                available |= set(state_recordings) | set(projection_recordings)
                available |= {f"{name}.spikes" for name in populations}
            unknown = selected_fields - available
            if unknown:
                raise ValueError(f"unavailable recording fields: {sorted(unknown)}")
            recordings = {k: v for k, v in recordings.items() if k in selected_fields}
            state_recordings = {
                k: v for k, v in state_recordings.items() if k in selected_fields
            }
            projection_recordings = {
                k: v for k, v in projection_recordings.items() if k in selected_fields
            }
        integrator_sum: dict[str, torch.Tensor] = {}
        spike_traces: dict[str, list[torch.Tensor]] = {name: [] for name in populations}
        voltage_traces: dict[str, list[torch.Tensor]] = {
            name: [] for name in populations
        }
        if selected_fields is not None:
            required_signals = {row["signal"] for row in self.plan.outputs}
            for operation in self.plan.graph.get("operations", []):
                required_signals.update(operation["sources"])
            named_spikes = selected_fields if recording == "full" else set()
            spike_traces = {
                name: []
                for name in populations
                if f"{name}.spikes" in required_signals | named_spikes
            }
            voltage_traces = {
                name: []
                for name in populations
                if f"{name}.voltage" in required_signals
            }

        for t in range(steps):
            new_spikes: dict[str, torch.Tensor] = {}
            for pop in self.plan.populations:
                name = pop["id"]
                incoming = {
                    "excitatory": torch.zeros_like(voltage[name]),
                    "inhibitory": torch.zeros_like(voltage[name]),
                }
                for projection in self.plan.projections:
                    if projection.target.partition(".")[0] != name:
                        continue
                    key = (projection.id, projection.polarity)
                    if not projection.enabled:
                        conductance[key].zero_()
                        continue
                    source_owner = projection.source.partition(".")[0]
                    if source_owner in populations:
                        if projection.delay_steps == 0:
                            source = new_spikes[source_owner]
                        else:
                            history = histories[source_owner]._values
                            source = history[-projection.delay_steps]
                    else:
                        source_t = t - projection.delay_steps
                        source = (
                            inputs[source_owner][source_t]
                            if source_t >= 0
                            else input_histories[source_owner][source_t]
                        )
                    drive = (
                        source @ self.weights[projection.parameter.replace(".", "__")]
                    )
                    conductance[key] = conductance[key] * projection.decay + drive
                    incoming[projection.polarity] += conductance[key]
                neuron = pop["neuron"]
                if neuron["kind"] == "leaky_integrator":
                    beta = math.exp(-self.plan.dt_ms / float(neuron["tau"]["value"]))
                    voltage[name] = (
                        beta * voltage[name]
                        + (1.0 - beta) / self.plan.dt_ms * incoming["excitatory"]
                    )
                    new_spikes[name] = torch.zeros_like(spikes[name])
                    integrator_sum[name] = (
                        integrator_sum.get(name, torch.zeros_like(voltage[name]))
                        + voltage[name]
                    )
                    threshold = neuron.get("soft_reset_threshold")
                    if threshold is not None:
                        reset = M.fast_sigmoid_spike(
                            voltage[name] - float(threshold),
                            float(neuron.get("surrogate_slope", M.SURROGATE_SLOPE)),
                        )
                        if pop.get("spiking"):
                            new_spikes[name] = reset
                        voltage[name] = voltage[name] - reset * float(threshold)
                    continue
                tau_mem = float(neuron["tau_mem"]["value"])
                c_m = float(neuron.get("capacitance_nf", 1.0 if tau_mem >= 15 else 0.5))
                g_l = float(neuron.get("leak_us", c_m / tau_mem))
                ref_steps = int(
                    neuron.get(
                        "refractory_steps",
                        max(
                            1,
                            round(
                                (M.ref_ms_E if tau_mem >= 15 else M.ref_ms_I)
                                / self.plan.dt_ms
                            ),
                        ),
                    )
                )
                dampen = float(neuron.get("voltage_grad_dampen", M.V_GRAD_DAMPEN))
                threshold = float(neuron.get("threshold_mv", M.V_th))
                voltage[name], new_spikes[name], refractory[name] = M.lif_step_expeuler(
                    voltage[name],
                    refractory[name],
                    incoming["excitatory"],
                    incoming["inhibitory"],
                    c_m,
                    g_l,
                    ref_steps,
                    lambda value, threshold_offset=0.0, threshold=threshold: (
                        M.fast_sigmoid_spike(
                            value - threshold - threshold_offset, self.surrogate_slope
                        )
                    ),
                    dt_override=self.plan.dt_ms,
                    v_grad_dampen=dampen,
                )
                for intervention_index, intervention in enumerate(
                    resolved_interventions
                ):
                    if intervention["population_id"] != name:
                        continue
                    absolute_step = completed_steps + t
                    seed_material = (
                        f"{intervention['seed']}:{intervention_index}:"
                        f"{intervention['kind']}:{name}:{absolute_step}"
                    ).encode()
                    step_seed = int.from_bytes(
                        hashlib.sha256(seed_material).digest()[:8], "big"
                    ) % (2**63 - 1)
                    generator = torch.Generator(device=device).manual_seed(step_seed)
                    sample = torch.rand(
                        new_spikes[name].shape,
                        device=device,
                        generator=generator,
                    )
                    if intervention["kind"] == "drop_spikes":
                        new_spikes[name] = new_spikes[name] * (
                            sample >= intervention["probability"]
                        )
                    else:
                        added = (sample < intervention["probability_per_step"]).to(
                            new_spikes[name].dtype
                        )
                        new_spikes[name] = torch.maximum(new_spikes[name], added)
            spikes = new_spikes
            for name in spike_traces:
                spike_traces[name].append(spikes[name])
            for name in voltage_traces:
                voltage_traces[name].append(voltage[name])
            for name in populations:
                histories[name].push(spikes[name])
            if recording != "none":
                for observable in self.plan.observables:
                    if observable["id"] not in recordings:
                        continue
                    owner, _, port = observable["signal"].partition(".")
                    recordings[observable["id"]].append(
                        (spikes if port == "spikes" else voltage)[owner]
                        .detach()
                        .clone()
                    )
            if recording == "full":
                for name in populations:
                    if f"{name}.voltage" not in state_recordings:
                        continue
                    state_recordings[f"{name}.voltage"].append(
                        voltage[name].detach().clone()
                    )
                for projection in self.plan.projections:
                    if f"{projection.id}.conductance" not in projection_recordings:
                        continue
                    projection_recordings[f"{projection.id}.conductance"].append(
                        conductance[(projection.id, projection.polarity)]
                        .detach()
                        .clone()
                    )

        outputs: dict[str, torch.Tensor] = {}
        signal_values: dict[str, torch.Tensor] = {
            f"{name}.value": value for name, value in inputs.items()
        }
        for name, values in spike_traces.items():
            signal_values[f"{name}.spikes"] = torch.stack(values)
        for name, values in voltage_traces.items():
            signal_values[f"{name}.voltage"] = torch.stack(values)

        def time_mask(
            mask: torch.Tensor, *, target: torch.Tensor, op_id: str
        ) -> torch.Tensor:
            if mask.shape[:2] != target.shape[:2]:
                raise ValueError(
                    f"{op_id}: valid-time mask leading shape expected {tuple(target.shape[:2])}, got {tuple(mask.shape[:2])}"
                )
            if mask.ndim != 2:
                raise ValueError(
                    f"{op_id}: valid-time mask must have shape [time, batch]"
                )
            mask_value = mask.to(device=target.device, dtype=target.dtype)
            return mask_value.reshape(
                mask_value.shape[0], mask_value.shape[1], *([1] * (target.ndim - 2))
            )

        def reduce_time(
            source: torch.Tensor, *, kind: str, mask: torch.Tensor | None, op_id: str
        ) -> torch.Tensor:
            if mask is None:
                return source.sum(dim=0) if kind == "reduce_sum" else source.mean(dim=0)
            weights = time_mask(mask, target=source, op_id=op_id)
            numerator = (source * weights).sum(dim=0)
            if kind == "reduce_sum":
                return numerator
            counts = weights.sum(dim=0)
            if torch.any(counts <= 0):
                raise ValueError(
                    f"{op_id}: valid-time mask contains an empty reduction window"
                )
            return numerator / counts

        remaining_ops = list(self.plan.graph.get("operations", []))
        while remaining_ops:
            ready_index = next(
                (
                    index
                    for index, op in enumerate(remaining_ops)
                    if all(source in signal_values for source in op["sources"])
                ),
                None,
            )
            if ready_index is None:
                unresolved = {
                    op["id"]: [
                        source
                        for source in op["sources"]
                        if source not in signal_values
                    ]
                    for op in remaining_ops
                }
                raise ValueError(f"operation dependencies are unresolved: {unresolved}")
            op = remaining_ops.pop(ready_index)
            sources = [signal_values[source] for source in op["sources"]]
            kind = op["kind"]
            if kind == "linear":
                parameter = op["parameters"][0].replace(".", "__")
                signal_values[f"{op['id']}.value"] = (
                    sources[0] @ self.weights[parameter]
                )
            elif kind in {"reduce_mean", "reduce_sum"}:
                mask_name = op.get("config", {}).get("mask")
                mask = signal_values.get(mask_name) if mask_name else None
                source_id = op["sources"][0]
                owner, _, port = source_id.partition(".")
                if (
                    kind == "reduce_mean"
                    and mask is None
                    and port == "voltage"
                    and owner in integrator_sum
                ):
                    signal_values[f"{op['id']}.value"] = integrator_sum[owner] / steps
                else:
                    signal_values[f"{op['id']}.value"] = reduce_time(
                        sources[0], kind=kind, mask=mask, op_id=op["id"]
                    )
            elif kind == "select_final":
                signal_values[f"{op['id']}.value"] = sources[0][-1]
            elif kind == "duration_normalise":
                config = op.get("config", {})
                mask_name = config.get("mask")
                if mask_name:
                    mask = signal_values[mask_name]
                    if mask.ndim != 2:
                        raise ValueError(
                            f"{op['id']}: valid-time mask must have shape [time, batch]"
                        )
                    mask_seconds = mask.to(
                        device=sources[0].device, dtype=sources[0].dtype
                    ).sum(dim=0) * (self.plan.dt_ms / 1000.0)
                    mask_seconds = mask_seconds.reshape(
                        mask_seconds.shape[0], *([1] * (sources[0].ndim - 1))
                    )
                    if torch.any(mask_seconds <= 0):
                        raise ValueError(
                            f"{op['id']}: valid-time mask contains zero valid duration"
                        )
                    signal_values[f"{op['id']}.value"] = sources[0] / mask_seconds
                else:
                    duration_s = float(config["duration"])
                    if duration_s <= 0:
                        raise ValueError(
                            f"{op['id']}: spike-rate duration must be positive seconds"
                        )
                    signal_values[f"{op['id']}.value"] = sources[0] / duration_s
            elif kind == "cumulative_sum":
                signal_values[f"{op['id']}.value"] = sources[0].cumsum(dim=0)
            else:
                raise ValueError(f"{op['id']}: unsupported operation {kind}")
        for output in self.plan.outputs:
            outputs[output["id"]] = signal_values[output["signal"]]
        packed = {k: torch.stack(v) for k, v in recordings.items() if v}
        packed.update({k: torch.stack(v) for k, v in state_recordings.items() if v})
        packed.update(
            {k: torch.stack(v) for k, v in projection_recordings.items() if v}
        )
        if recording == "full":
            packed.update(
                {
                    f"{name}.spikes": torch.stack(values)
                    for name, values in spike_traces.items()
                    if selected_fields is None or f"{name}.spikes" in selected_fields
                }
            )
        next_input_histories = {
            name: torch.cat((history, inputs[name]), dim=0)[-history.shape[0] :]
            .detach()
            .clone()
            for name, history in input_histories.items()
        }
        next_runtime_state = GraphRuntimeState(
            signature=expected_signature,
            compatibility=expected_compatibility,
            completed_steps=completed_steps + steps,
            voltages={name: value.detach().clone() for name, value in voltage.items()},
            refractory={
                name: value.detach().clone() for name, value in refractory.items()
            },
            conductances={
                p.id: conductance[(p.id, p.polarity)].detach().clone()
                for p in self.plan.projections
            },
            population_histories={
                name: history.export() for name, history in histories.items()
            },
            input_histories=next_input_histories,
        )
        return ExecutionResult(
            executor="graph",
            outputs=outputs,
            recordings=packed,
            parameters={k: v.detach().clone() for k, v in self.parameter_map().items()},
            final_state={
                f"{k}.voltage": v.detach().clone() for k, v in voltage.items()
            },
            runtime_state=next_runtime_state,
            metrics={"resolved_interventions": resolved_interventions},
            model=self,
        )
Complete class implementation
class GraphExecutor(nn.Module):
    """Dense graph executor whose graph topology is lowered before simulation."""

    def __init__(
        self,
        plan: GraphPlan,
        *,
        seed: int = 0,
        trainable_parameters: Sequence[str] = (),
        surrogate_slope: float = M.SURROGATE_SLOPE,
    ):
        super().__init__()
        self.plan = plan
        self.surrogate_slope = float(surrogate_slope)
        torch.manual_seed(seed)
        rows = {row["id"]: row for row in plan.graph.get("parameters", [])}
        self.weights = nn.ParameterDict()
        self.initialization_metadata: dict[str, dict[str, Any]] = {}
        pop_ids = {p["id"] for p in plan.populations}
        trainable = set(trainable_parameters)

        def initialise(
            row: Mapping[str, Any],
            *,
            runtime_shape: tuple[int, ...],
            scale_by_fanin: bool,
        ) -> torch.Tensor:
            init = row["initializer"]
            kind = init["kind"]
            if kind in {"normal", "lower_clamped_normal"}:
                value = (
                    torch.randn(*runtime_shape)
                    .mul_(float(init["std"]))
                    .add_(float(init["mean"]))
                    .clamp_(min=0)
                )
            elif kind == "signed_normal":
                value = (
                    torch.randn(*runtime_shape)
                    .mul_(float(init["std"]))
                    .add_(float(init["mean"]))
                )
            elif kind == "uniform":
                value = (
                    torch.rand(*runtime_shape)
                    .mul_(float(init["high"]) - float(init["low"]))
                    .add_(float(init["low"]))
                )
            elif init["kind"] == "constant":
                value = torch.full(runtime_shape, float(init["value"]))
            elif kind == "zeros":
                value = torch.zeros(runtime_shape)
            else:
                raise ValueError(f"{row['id']}: unsupported initializer {init['kind']}")
            zero_fraction = float(init.get("initial_zero_fraction", 0.0))
            if not 0 <= zero_fraction < 1:
                raise ValueError(
                    f"{row['id']}: initial_zero_fraction must satisfy 0 <= fraction < 1"
                )
            zeroing = init.get("zeroing", "bernoulli")
            if zero_fraction:
                if zeroing == "exact_k":
                    if len(runtime_shape) != 2:
                        raise ValueError(
                            f"{row['id']}: exact_k zeroing requires a matrix"
                        )
                    fan_in, fan_out = runtime_shape
                    kept = max(1, int(round((1.0 - zero_fraction) * fan_in)))
                    mask = torch.zeros(runtime_shape)
                    for column in range(fan_out):
                        mask[torch.randperm(fan_in)[:kept], column] = 1
                    value = value * mask * (fan_in / kept)
                elif zeroing == "bernoulli":
                    value = (
                        value
                        * (torch.rand(*runtime_shape) > zero_fraction)
                        / (1 - zero_fraction)
                    )
                else:
                    raise ValueError(
                        f"{row['id']}: unsupported initializer zeroing {zeroing}"
                    )
            constraint = row.get("constraint")
            if constraint is not None:
                if constraint.get("kind") != "non_negative":
                    raise ValueError(
                        f"{row['id']}: unsupported constraint {constraint.get('kind')}"
                    )
                value = value.clamp(min=0)
            if scale_by_fanin:
                value = value / runtime_shape[0]
            flat = value.reshape(-1)
            self.initialization_metadata[row["id"]] = {
                "initializer": dict(init),
                "constraint": dict(constraint) if constraint else None,
                "unit": row.get("unit"),
                "runtime_shape": list(runtime_shape),
                "scaling": "fan_in_normalized" if scale_by_fanin else "direct",
                "statistics": {
                    "count": flat.numel(),
                    "zero_fraction": float((flat == 0).float().mean()),
                    "mean": float(flat.mean()),
                    "std": float(flat.std(unbiased=False)),
                    "min": float(flat.min()),
                    "max": float(flat.max()),
                },
            }
            return value

        def init_priority(projection: PlannedProjection) -> tuple[int, str]:
            source = projection.source.partition(".")[0]
            target = projection.target.partition(".")[0]
            target_kind = next(
                p["neuron"]["kind"] for p in plan.populations if p["id"] == target
            )
            if source not in pop_ids:
                return (0, projection.id)
            if target_kind == "leaky_integrator":
                return (1, projection.id)
            if projection.polarity == "excitatory":
                return (2, projection.id)
            return (3, projection.id)

        realised: dict[str, torch.Tensor] = {}
        for projection in sorted(plan.projections, key=init_priority):
            row = rows[projection.parameter]
            shape = tuple(reversed(row["shape"]))  # runtime is [source, target]
            realised[projection.parameter] = initialise(
                row, runtime_shape=shape, scale_by_fanin=True
            )
        for operation in plan.graph.get("operations", []):
            if operation.get("kind") != "linear":
                continue
            for parameter in operation.get("parameters", []):
                if parameter in realised:
                    continue
                row = rows[parameter]
                shape = tuple(reversed(row["shape"]))  # runtime is [source, target]
                realised[parameter] = initialise(
                    row, runtime_shape=shape, scale_by_fanin=False
                )
        for projection in plan.projections:
            parameter_id = projection.parameter
            self.weights[projection.parameter.replace(".", "__")] = nn.Parameter(
                realised[projection.parameter], requires_grad=parameter_id in trainable
            )
        for operation in plan.graph.get("operations", []):
            for parameter in operation.get("parameters", []):
                self.weights[parameter.replace(".", "__")] = nn.Parameter(
                    realised[parameter], requires_grad=parameter in trainable
                )

    def parameter_map(self) -> dict[str, torch.Tensor]:
        return {name.replace("__", "."): value for name, value in self.weights.items()}

    def forward(
        self,
        inputs: Mapping[str, torch.Tensor],
        *,
        record: bool | RecordingProfile = True,
        recording_fields: Sequence[str] | None = None,
        runtime_state: GraphRuntimeState | None = None,
        interventions: Sequence[Mapping[str, Any]] = (),
    ) -> ExecutionResult:
        recording: RecordingProfile = (
            "full" if record is True else "none" if record is False else record
        )
        if recording not in {"full", "observables", "none"}:
            raise ValueError(
                f"recording profile expected full, observables, or none; got {recording!r}"
            )
        if not inputs:
            raise ValueError("graph execution requires at least one input tensor")
        first = next(iter(inputs.values()))
        steps, batch = first.shape[:2]
        device = first.device
        parameter_dtype = (
            next(iter(self.weights.values())).dtype if self.weights else first.dtype
        )
        input_specs = {row["id"]: row for row in self.plan.graph.get("inputs", [])}
        for name, value in inputs.items():
            if value.shape[:2] != (steps, batch):
                raise ValueError(
                    f"input {name} leading shape expected {(steps, batch)}, got {tuple(value.shape[:2])}"
                )
            if value.device != device:
                raise ValueError(
                    f"input {name} device expected {device}, got {value.device}"
                )
            is_mask = input_specs.get(name, {}).get("signal_type") == "mask"
            if value.dtype != parameter_dtype and not (
                is_mask and value.dtype == torch.bool
            ):
                raise ValueError(
                    f"input {name} dtype expected {parameter_dtype}, got {value.dtype}"
                )
        populations = {p["id"]: p for p in self.plan.populations}
        resolved_interventions: list[dict[str, Any]] = []
        intervention_keys: set[tuple[str, str]] = set()
        for index, raw in enumerate(interventions):
            row = dict(raw)
            kind = str(row.get("kind", ""))
            population_id = str(row.get("population_id", ""))
            if kind not in {"drop_spikes", "add_poisson_spikes"}:
                raise ValueError(
                    f"inference intervention {index} has unsupported kind {kind!r}"
                )
            if population_id not in populations:
                raise ValueError(
                    f"inference intervention {index} targets unknown population {population_id!r}"
                )
            if populations[population_id]["neuron"]["kind"] != "coba_lif":
                raise ValueError(
                    f"inference intervention {index} population {population_id!r} does not emit spikes"
                )
            key = (kind, population_id)
            if key in intervention_keys:
                raise ValueError(
                    f"inference intervention repeats {kind} for population {population_id}"
                )
            intervention_keys.add(key)
            allowed = (
                {"kind", "population_id", "probability", "seed"}
                if kind == "drop_spikes"
                else {"kind", "population_id", "rate_hz", "seed"}
            )
            unknown = sorted(set(row) - allowed)
            if unknown:
                raise ValueError(
                    f"inference intervention {index} has unsupported fields: {unknown}"
                )
            seed = int(row.get("seed", 0))
            if kind == "drop_spikes":
                value = float(row.get("probability", float("nan")))
                if not math.isfinite(value) or not 0 <= value <= 1:
                    raise ValueError(
                        f"inference intervention {index} drop probability must be finite and between zero and one"
                    )
                resolved_interventions.append(
                    {
                        "kind": kind,
                        "population_id": population_id,
                        "probability": value,
                        "seed": seed,
                    }
                )
            else:
                value = float(row.get("rate_hz", float("nan")))
                probability = value * self.plan.dt_ms / 1000.0
                if not math.isfinite(value) or value < 0 or probability > 1:
                    raise ValueError(
                        f"inference intervention {index} Poisson rate must be finite, non-negative, and satisfy rate times dt <= 1"
                    )
                resolved_interventions.append(
                    {
                        "kind": kind,
                        "population_id": population_id,
                        "rate_hz": value,
                        "probability_per_step": probability,
                        "seed": seed,
                    }
                )
        population_history_lengths = {
            name: max(
                (
                    p.delay_steps
                    for p in self.plan.projections
                    if p.source.startswith(name + ".")
                ),
                default=1,
            )
            for name in populations
        }
        input_history_lengths = {
            row["id"]: max(
                (
                    p.delay_steps
                    for p in self.plan.projections
                    if p.source.partition(".")[0] == row["id"]
                ),
                default=0,
            )
            for row in self.plan.graph.get("inputs", [])
        }
        expected_compatibility = runtime_state_compatibility(self.plan)
        expected_signature = runtime_state_signature(self.plan)
        if runtime_state is None:
            voltage = {
                name: (
                    torch.zeros((batch, p["size"]), device=device)
                    if p["neuron"]["kind"] == "leaky_integrator"
                    else torch.full((batch, p["size"]), M.E_L, device=device)
                )
                for name, p in populations.items()
            }
            refractory = {
                name: torch.zeros((batch, p["size"]), dtype=torch.long, device=device)
                for name, p in populations.items()
            }
            spikes = {
                name: torch.zeros((batch, p["size"]), device=device)
                for name, p in populations.items()
            }
            conductance = {
                (p.id, p.polarity): torch.zeros(
                    (batch, populations[p.target.partition(".")[0]]["size"]),
                    device=device,
                )
                for p in self.plan.projections
            }
            histories = {
                name: DelayBuffer(population_history_lengths[name], value)
                for name, value in spikes.items()
            }
            input_histories = {
                name: torch.zeros(
                    (length, *inputs[name].shape[1:]),
                    device=device,
                    dtype=inputs[name].dtype,
                )
                for name, length in input_history_lengths.items()
                if length > 0
            }
            completed_steps = 0
        else:
            if runtime_state.signature != expected_signature:
                detail = _compatibility_mismatch(
                    expected_compatibility, runtime_state.compatibility
                )
                raise ValueError(
                    "runtime state is incompatible with graph plan: "
                    + (
                        detail
                        or f"signature expected {expected_signature}, got {runtime_state.signature}"
                    )
                )

            def restore_group(
                label: str,
                values: Mapping[str, torch.Tensor],
                shapes: Mapping[str, tuple[int, ...]],
            ) -> dict[str, torch.Tensor]:
                if set(values) != set(shapes):
                    raise ValueError(
                        f"runtime state {label} keys expected {sorted(shapes)}, got {sorted(values)}"
                    )
                restored = {}
                for name, expected_shape in shapes.items():
                    value = values[name]
                    if tuple(value.shape) != expected_shape:
                        raise ValueError(
                            f"runtime state {label}.{name} shape expected {expected_shape}, got {tuple(value.shape)}"
                        )
                    restored[name] = value.detach().to(device).clone()
                return restored

            pop_shapes = {
                name: (batch, int(row["size"])) for name, row in populations.items()
            }
            voltage = restore_group("voltages", runtime_state.voltages, pop_shapes)
            for name, value in voltage.items():
                if value.dtype != parameter_dtype:
                    raise ValueError(
                        f"runtime state voltages.{name} dtype expected {parameter_dtype}, got {value.dtype}"
                    )
            refractory = restore_group(
                "refractory", runtime_state.refractory, pop_shapes
            )
            for name, value in refractory.items():
                if value.dtype != torch.long:
                    raise ValueError(
                        f"runtime state refractory.{name} dtype expected torch.int64, got {value.dtype}"
                    )
            conductance_by_id = restore_group(
                "conductances",
                runtime_state.conductances,
                {
                    p.id: (batch, int(populations[p.target.partition(".")[0]]["size"]))
                    for p in self.plan.projections
                },
            )
            conductance = {
                (p.id, p.polarity): conductance_by_id[p.id]
                for p in self.plan.projections
            }
            for name, value in conductance_by_id.items():
                if value.dtype != parameter_dtype:
                    raise ValueError(
                        f"runtime state conductances.{name} dtype expected {parameter_dtype}, got {value.dtype}"
                    )
            population_history_values = restore_group(
                "population_histories",
                runtime_state.population_histories,
                {
                    name: (population_history_lengths[name], batch, int(row["size"]))
                    for name, row in populations.items()
                },
            )
            histories = {
                name: DelayBuffer.restore(value)
                for name, value in population_history_values.items()
            }
            for name, value in population_history_values.items():
                if value.dtype != parameter_dtype:
                    raise ValueError(
                        f"runtime state population_histories.{name} dtype expected {parameter_dtype}, got {value.dtype}"
                    )
            spikes = {name: histories[name]._values[-1].clone() for name in populations}
            input_histories = restore_group(
                "input_histories",
                runtime_state.input_histories,
                {
                    name: (length, *inputs[name].shape[1:])
                    for name, length in input_history_lengths.items()
                    if length > 0
                },
            )
            for name, value in input_histories.items():
                if value.dtype != inputs[name].dtype:
                    raise ValueError(
                        f"runtime state input_histories.{name} dtype expected {inputs[name].dtype}, got {value.dtype}"
                    )
            completed_steps = int(runtime_state.completed_steps)
        recordings: dict[str, list[torch.Tensor]] = {
            o["id"]: [] for o in self.plan.observables
        }
        state_recordings: dict[str, list[torch.Tensor]] = {
            f"{name}.voltage": [] for name in populations
        }
        projection_recordings: dict[str, list[torch.Tensor]] = {
            f"{p.id}.conductance": [] for p in self.plan.projections
        }
        selected_fields = None if recording_fields is None else set(recording_fields)
        if selected_fields is not None:
            available = set(recordings) if recording != "none" else set()
            if recording == "full":
                available |= set(state_recordings) | set(projection_recordings)
                available |= {f"{name}.spikes" for name in populations}
            unknown = selected_fields - available
            if unknown:
                raise ValueError(f"unavailable recording fields: {sorted(unknown)}")
            recordings = {k: v for k, v in recordings.items() if k in selected_fields}
            state_recordings = {
                k: v for k, v in state_recordings.items() if k in selected_fields
            }
            projection_recordings = {
                k: v for k, v in projection_recordings.items() if k in selected_fields
            }
        integrator_sum: dict[str, torch.Tensor] = {}
        spike_traces: dict[str, list[torch.Tensor]] = {name: [] for name in populations}
        voltage_traces: dict[str, list[torch.Tensor]] = {
            name: [] for name in populations
        }
        if selected_fields is not None:
            required_signals = {row["signal"] for row in self.plan.outputs}
            for operation in self.plan.graph.get("operations", []):
                required_signals.update(operation["sources"])
            named_spikes = selected_fields if recording == "full" else set()
            spike_traces = {
                name: []
                for name in populations
                if f"{name}.spikes" in required_signals | named_spikes
            }
            voltage_traces = {
                name: []
                for name in populations
                if f"{name}.voltage" in required_signals
            }

        for t in range(steps):
            new_spikes: dict[str, torch.Tensor] = {}
            for pop in self.plan.populations:
                name = pop["id"]
                incoming = {
                    "excitatory": torch.zeros_like(voltage[name]),
                    "inhibitory": torch.zeros_like(voltage[name]),
                }
                for projection in self.plan.projections:
                    if projection.target.partition(".")[0] != name:
                        continue
                    key = (projection.id, projection.polarity)
                    if not projection.enabled:
                        conductance[key].zero_()
                        continue
                    source_owner = projection.source.partition(".")[0]
                    if source_owner in populations:
                        if projection.delay_steps == 0:
                            source = new_spikes[source_owner]
                        else:
                            history = histories[source_owner]._values
                            source = history[-projection.delay_steps]
                    else:
                        source_t = t - projection.delay_steps
                        source = (
                            inputs[source_owner][source_t]
                            if source_t >= 0
                            else input_histories[source_owner][source_t]
                        )
                    drive = (
                        source @ self.weights[projection.parameter.replace(".", "__")]
                    )
                    conductance[key] = conductance[key] * projection.decay + drive
                    incoming[projection.polarity] += conductance[key]
                neuron = pop["neuron"]
                if neuron["kind"] == "leaky_integrator":
                    beta = math.exp(-self.plan.dt_ms / float(neuron["tau"]["value"]))
                    voltage[name] = (
                        beta * voltage[name]
                        + (1.0 - beta) / self.plan.dt_ms * incoming["excitatory"]
                    )
                    new_spikes[name] = torch.zeros_like(spikes[name])
                    integrator_sum[name] = (
                        integrator_sum.get(name, torch.zeros_like(voltage[name]))
                        + voltage[name]
                    )
                    threshold = neuron.get("soft_reset_threshold")
                    if threshold is not None:
                        reset = M.fast_sigmoid_spike(
                            voltage[name] - float(threshold),
                            float(neuron.get("surrogate_slope", M.SURROGATE_SLOPE)),
                        )
                        if pop.get("spiking"):
                            new_spikes[name] = reset
                        voltage[name] = voltage[name] - reset * float(threshold)
                    continue
                tau_mem = float(neuron["tau_mem"]["value"])
                c_m = float(neuron.get("capacitance_nf", 1.0 if tau_mem >= 15 else 0.5))
                g_l = float(neuron.get("leak_us", c_m / tau_mem))
                ref_steps = int(
                    neuron.get(
                        "refractory_steps",
                        max(
                            1,
                            round(
                                (M.ref_ms_E if tau_mem >= 15 else M.ref_ms_I)
                                / self.plan.dt_ms
                            ),
                        ),
                    )
                )
                dampen = float(neuron.get("voltage_grad_dampen", M.V_GRAD_DAMPEN))
                threshold = float(neuron.get("threshold_mv", M.V_th))
                voltage[name], new_spikes[name], refractory[name] = M.lif_step_expeuler(
                    voltage[name],
                    refractory[name],
                    incoming["excitatory"],
                    incoming["inhibitory"],
                    c_m,
                    g_l,
                    ref_steps,
                    lambda value, threshold_offset=0.0, threshold=threshold: (
                        M.fast_sigmoid_spike(
                            value - threshold - threshold_offset, self.surrogate_slope
                        )
                    ),
                    dt_override=self.plan.dt_ms,
                    v_grad_dampen=dampen,
                )
                for intervention_index, intervention in enumerate(
                    resolved_interventions
                ):
                    if intervention["population_id"] != name:
                        continue
                    absolute_step = completed_steps + t
                    seed_material = (
                        f"{intervention['seed']}:{intervention_index}:"
                        f"{intervention['kind']}:{name}:{absolute_step}"
                    ).encode()
                    step_seed = int.from_bytes(
                        hashlib.sha256(seed_material).digest()[:8], "big"
                    ) % (2**63 - 1)
                    generator = torch.Generator(device=device).manual_seed(step_seed)
                    sample = torch.rand(
                        new_spikes[name].shape,
                        device=device,
                        generator=generator,
                    )
                    if intervention["kind"] == "drop_spikes":
                        new_spikes[name] = new_spikes[name] * (
                            sample >= intervention["probability"]
                        )
                    else:
                        added = (sample < intervention["probability_per_step"]).to(
                            new_spikes[name].dtype
                        )
                        new_spikes[name] = torch.maximum(new_spikes[name], added)
            spikes = new_spikes
            for name in spike_traces:
                spike_traces[name].append(spikes[name])
            for name in voltage_traces:
                voltage_traces[name].append(voltage[name])
            for name in populations:
                histories[name].push(spikes[name])
            if recording != "none":
                for observable in self.plan.observables:
                    if observable["id"] not in recordings:
                        continue
                    owner, _, port = observable["signal"].partition(".")
                    recordings[observable["id"]].append(
                        (spikes if port == "spikes" else voltage)[owner]
                        .detach()
                        .clone()
                    )
            if recording == "full":
                for name in populations:
                    if f"{name}.voltage" not in state_recordings:
                        continue
                    state_recordings[f"{name}.voltage"].append(
                        voltage[name].detach().clone()
                    )
                for projection in self.plan.projections:
                    if f"{projection.id}.conductance" not in projection_recordings:
                        continue
                    projection_recordings[f"{projection.id}.conductance"].append(
                        conductance[(projection.id, projection.polarity)]
                        .detach()
                        .clone()
                    )

        outputs: dict[str, torch.Tensor] = {}
        signal_values: dict[str, torch.Tensor] = {
            f"{name}.value": value for name, value in inputs.items()
        }
        for name, values in spike_traces.items():
            signal_values[f"{name}.spikes"] = torch.stack(values)
        for name, values in voltage_traces.items():
            signal_values[f"{name}.voltage"] = torch.stack(values)

        def time_mask(
            mask: torch.Tensor, *, target: torch.Tensor, op_id: str
        ) -> torch.Tensor:
            if mask.shape[:2] != target.shape[:2]:
                raise ValueError(
                    f"{op_id}: valid-time mask leading shape expected {tuple(target.shape[:2])}, got {tuple(mask.shape[:2])}"
                )
            if mask.ndim != 2:
                raise ValueError(
                    f"{op_id}: valid-time mask must have shape [time, batch]"
                )
            mask_value = mask.to(device=target.device, dtype=target.dtype)
            return mask_value.reshape(
                mask_value.shape[0], mask_value.shape[1], *([1] * (target.ndim - 2))
            )

        def reduce_time(
            source: torch.Tensor, *, kind: str, mask: torch.Tensor | None, op_id: str
        ) -> torch.Tensor:
            if mask is None:
                return source.sum(dim=0) if kind == "reduce_sum" else source.mean(dim=0)
            weights = time_mask(mask, target=source, op_id=op_id)
            numerator = (source * weights).sum(dim=0)
            if kind == "reduce_sum":
                return numerator
            counts = weights.sum(dim=0)
            if torch.any(counts <= 0):
                raise ValueError(
                    f"{op_id}: valid-time mask contains an empty reduction window"
                )
            return numerator / counts

        remaining_ops = list(self.plan.graph.get("operations", []))
        while remaining_ops:
            ready_index = next(
                (
                    index
                    for index, op in enumerate(remaining_ops)
                    if all(source in signal_values for source in op["sources"])
                ),
                None,
            )
            if ready_index is None:
                unresolved = {
                    op["id"]: [
                        source
                        for source in op["sources"]
                        if source not in signal_values
                    ]
                    for op in remaining_ops
                }
                raise ValueError(f"operation dependencies are unresolved: {unresolved}")
            op = remaining_ops.pop(ready_index)
            sources = [signal_values[source] for source in op["sources"]]
            kind = op["kind"]
            if kind == "linear":
                parameter = op["parameters"][0].replace(".", "__")
                signal_values[f"{op['id']}.value"] = (
                    sources[0] @ self.weights[parameter]
                )
            elif kind in {"reduce_mean", "reduce_sum"}:
                mask_name = op.get("config", {}).get("mask")
                mask = signal_values.get(mask_name) if mask_name else None
                source_id = op["sources"][0]
                owner, _, port = source_id.partition(".")
                if (
                    kind == "reduce_mean"
                    and mask is None
                    and port == "voltage"
                    and owner in integrator_sum
                ):
                    signal_values[f"{op['id']}.value"] = integrator_sum[owner] / steps
                else:
                    signal_values[f"{op['id']}.value"] = reduce_time(
                        sources[0], kind=kind, mask=mask, op_id=op["id"]
                    )
            elif kind == "select_final":
                signal_values[f"{op['id']}.value"] = sources[0][-1]
            elif kind == "duration_normalise":
                config = op.get("config", {})
                mask_name = config.get("mask")
                if mask_name:
                    mask = signal_values[mask_name]
                    if mask.ndim != 2:
                        raise ValueError(
                            f"{op['id']}: valid-time mask must have shape [time, batch]"
                        )
                    mask_seconds = mask.to(
                        device=sources[0].device, dtype=sources[0].dtype
                    ).sum(dim=0) * (self.plan.dt_ms / 1000.0)
                    mask_seconds = mask_seconds.reshape(
                        mask_seconds.shape[0], *([1] * (sources[0].ndim - 1))
                    )
                    if torch.any(mask_seconds <= 0):
                        raise ValueError(
                            f"{op['id']}: valid-time mask contains zero valid duration"
                        )
                    signal_values[f"{op['id']}.value"] = sources[0] / mask_seconds
                else:
                    duration_s = float(config["duration"])
                    if duration_s <= 0:
                        raise ValueError(
                            f"{op['id']}: spike-rate duration must be positive seconds"
                        )
                    signal_values[f"{op['id']}.value"] = sources[0] / duration_s
            elif kind == "cumulative_sum":
                signal_values[f"{op['id']}.value"] = sources[0].cumsum(dim=0)
            else:
                raise ValueError(f"{op['id']}: unsupported operation {kind}")
        for output in self.plan.outputs:
            outputs[output["id"]] = signal_values[output["signal"]]
        packed = {k: torch.stack(v) for k, v in recordings.items() if v}
        packed.update({k: torch.stack(v) for k, v in state_recordings.items() if v})
        packed.update(
            {k: torch.stack(v) for k, v in projection_recordings.items() if v}
        )
        if recording == "full":
            packed.update(
                {
                    f"{name}.spikes": torch.stack(values)
                    for name, values in spike_traces.items()
                    if selected_fields is None or f"{name}.spikes" in selected_fields
                }
            )
        next_input_histories = {
            name: torch.cat((history, inputs[name]), dim=0)[-history.shape[0] :]
            .detach()
            .clone()
            for name, history in input_histories.items()
        }
        next_runtime_state = GraphRuntimeState(
            signature=expected_signature,
            compatibility=expected_compatibility,
            completed_steps=completed_steps + steps,
            voltages={name: value.detach().clone() for name, value in voltage.items()},
            refractory={
                name: value.detach().clone() for name, value in refractory.items()
            },
            conductances={
                p.id: conductance[(p.id, p.polarity)].detach().clone()
                for p in self.plan.projections
            },
            population_histories={
                name: history.export() for name, history in histories.items()
            },
            input_histories=next_input_histories,
        )
        return ExecutionResult(
            executor="graph",
            outputs=outputs,
            recordings=packed,
            parameters={k: v.detach().clone() for k, v in self.parameter_map().items()},
            final_state={
                f"{k}.voltage": v.detach().clone() for k, v in voltage.items()
            },
            runtime_state=next_runtime_state,
            metrics={"resolved_interventions": resolved_interventions},
            model=self,
        )

build

View source

def build(spec: ExecutionSpec) -> ExecutionResult

Plan and initialize a graph request, returning ExecutionResult with model, parameters and build metrics. A legacy request returns routing metadata instead of executing the legacy CLI.

ParameterAnnotationDefaultMeaning
specExecutionSpecrequiredDefined by the source contract and implementation below.

Return annotation: ExecutionResult.

Return expressions (branch-dependent; names refer to the linked implementation):

ExecutionResult(executor='legacy', metrics={'request': 'build', 'routing': 'legacy'})
ExecutionResult(executor='graph', model=model, parameters=model.parameter_map(), metrics={'build_s': time.perf_counter() - started, 'initialization': model.initialization_metadata, 'training_schema': training.get('schema') if training else None})

Explicit exceptions in this implementation; called helpers may raise additional errors:

Explicit exception expression
ValueError('graph execution requires graph data or a bundle')
Implementation
def build(spec: ExecutionSpec) -> ExecutionResult:
    if spec.executor == "legacy":
        return ExecutionResult(
            executor="legacy", metrics={"request": "build", "routing": "legacy"}
        )
    graph = spec.graph
    training = spec.training
    if graph is None and spec.bundle is not None:
        manifest, graph = load_graph_bundle(spec.bundle)
        if spec.kind == "train" and training is None:
            training = load_training_recipe(spec.bundle, manifest, graph)
    if graph is None:
        raise ValueError("graph execution requires graph data or a bundle")
    device = resolve_device(spec.device)
    started = time.perf_counter()
    trainable = (
        training.get("resolved_parameters", {}).get("trainable", []) if training else []
    )
    surrogate = (training or {}).get("surrogate") or {}
    surrogate_slope = float(surrogate.get("slope", M.SURROGATE_SLOPE))
    model = GraphExecutor(
        plan_graph(graph),
        seed=spec.seed,
        trainable_parameters=trainable,
        surrogate_slope=surrogate_slope,
    ).to(device)
    return ExecutionResult(
        executor="graph",
        model=model,
        parameters=model.parameter_map(),
        metrics={
            "build_s": time.perf_counter() - started,
            "initialization": model.initialization_metadata,
            "training_schema": training.get("schema") if training else None,
        },
    )

simulate

View source

def simulate(spec: ExecutionSpec, *, runtime_state: GraphRuntimeState | None=None) -> ExecutionResult

Execute graph-native forward dynamics with validated inputs, optional checkpoints, inference overrides/interventions and optional runtime continuation. Return outputs, selected recordings, final/runtime state and provenance metrics. A typed legacy request returns routing metadata.

ParameterAnnotationDefaultMeaning
specExecutionSpecrequiredDefined by the source contract and implementation below.
runtime_stateGraphRuntimeState | NoneNoneDynamic graph state for causal continuation; distinct from training checkpoints.

Return annotation: ExecutionResult.

Return expressions (branch-dependent; names refer to the linked implementation):

ExecutionResult(executor='legacy', metrics={'request': 'simulate', 'routing': 'legacy'})
result

Explicit exceptions in this implementation; called helpers may raise additional errors:

Explicit exception expression
ValueError(f'unsupported inference overrides: {unknown_overrides}')
ValueError('inference timestep must be finite and positive')
ValueError('inference timestep recompilation cannot convert runtime state')
ValueError('inference timestep recompilation requires resampleable Poisson input bindings')
ValueError('duration and input-rate inference overrides require Poisson input bindings')
ValueError(f'inference duration {duration_ms} ms must be a positive integer multiple of dt={dt_ms} ms')
ValueError('inference input rate must be finite and non-negative')
ValueError(f'inference projection scales target unknown projections: {unknown}')
ValueError('dataset snapshot binding cannot be combined with other input bindings')
ValueError('graph execution requires graph data or a bundle')
ValueError(f'inference checkpoint graph digest expected {graph_digest}, got {checkpoint.graph_digest}')
ValueError('inference checkpoint parameter names do not match graph')
ValueError(f'inference projection scale {projection_id} must be finite and non-negative')
ValueError(f'inference checkpoint parameter {name} expected shape={list(parameter.shape)} dtype={parameter.dtype}, got shape={list(restored.shape)} dtype={restored.dtype}')
Implementation
def simulate(
    spec: ExecutionSpec, *, runtime_state: GraphRuntimeState | None = None
) -> ExecutionResult:
    if spec.executor != "graph":
        return ExecutionResult(
            executor="legacy", metrics={"request": "simulate", "routing": "legacy"}
        )
    overrides = dict(spec.options.get("inference_overrides", {}))
    requested_interventions = tuple(spec.options.get("inference_interventions", ()))
    interventions = tuple(
        {**dict(row), "seed": dict(row).get("seed", spec.seed)}
        for row in requested_interventions
    )
    allowed_overrides = {
        "duration_ms",
        "input_rate_hz",
        "projection_scales",
        "timestep_ms",
    }
    unknown_overrides = sorted(set(overrides) - allowed_overrides)
    if unknown_overrides:
        raise ValueError(f"unsupported inference overrides: {unknown_overrides}")
    source_graph: Mapping[str, Any] | None = None
    source_graph_digest: str | None = None
    source_dt_ms: float | None = None
    build_spec = spec
    if "timestep_ms" in overrides:
        timestep_ms = float(overrides["timestep_ms"])
        if not math.isfinite(timestep_ms) or timestep_ms <= 0:
            raise ValueError("inference timestep must be finite and positive")
        if runtime_state is not None or spec.runtime_state is not None:
            raise ValueError(
                "inference timestep recompilation cannot convert runtime state"
            )
        if (
            not spec.poisson_bindings
            or spec.input_bindings
            or spec.event_bindings
            or spec.inputs
        ):
            raise ValueError(
                "inference timestep recompilation requires resampleable Poisson input bindings"
            )
        if spec.graph is not None:
            source_graph = spec.graph
        elif spec.bundle is not None:
            _, source_graph = load_graph_bundle(spec.bundle)
        else:
            raise ValueError("graph execution requires graph data or a bundle")
        source_graph_digest = _json_digest(source_graph)
        source_dt_ms = float(source_graph["timebase"]["dt"]["value"])
        recompiled_graph = copy.deepcopy(source_graph)
        recompiled_graph["timebase"]["dt"] = {
            "value": timestep_ms,
            "unit": "ms",
        }
        build_spec = replace(spec, graph=recompiled_graph, bundle=None)
    built = build(build_spec)
    assert isinstance(built.model, GraphExecutor)
    device = resolve_device(spec.device)
    poisson_bindings = spec.poisson_bindings
    if (
        "duration_ms" in overrides
        or "input_rate_hz" in overrides
        or "timestep_ms" in overrides
    ):
        if (
            not poisson_bindings
            or spec.input_bindings
            or spec.event_bindings
            or spec.inputs
        ):
            raise ValueError(
                "duration and input-rate inference overrides require Poisson input bindings"
            )
        dt_ms = built.model.plan.dt_ms
        duration_ms = float(
            overrides.get(
                "duration_ms",
                poisson_bindings[0].steps_count * (source_dt_ms or dt_ms),
            )
        )
        raw_steps = duration_ms / dt_ms
        if duration_ms <= 0 or not math.isclose(
            raw_steps, round(raw_steps), abs_tol=1e-9
        ):
            raise ValueError(
                f"inference duration {duration_ms} ms must be a positive integer multiple of dt={dt_ms} ms"
            )
        rate = overrides.get("input_rate_hz")
        if rate is not None and (not math.isfinite(float(rate)) or float(rate) < 0):
            raise ValueError("inference input rate must be finite and non-negative")
        poisson_bindings = tuple(
            replace(
                binding,
                steps_count=int(round(raw_steps)),
                rates_hz=(float(rate),) if rate is not None else binding.rates_hz,
                categorical=False if rate is not None else binding.categorical,
            )
            for binding in poisson_bindings
        )
    checkpoint_provenance = None
    if spec.checkpoint:
        checkpoint_path = Path(spec.checkpoint)
        if checkpoint_path.is_dir():
            checkpoint = load_training_checkpoint(checkpoint_path, device=device)
            graph_digest = source_graph_digest or _json_digest(built.model.plan.graph)
            if checkpoint.graph_digest != graph_digest:
                raise ValueError(
                    f"inference checkpoint graph digest expected {graph_digest}, got {checkpoint.graph_digest}"
                )
            parameter_map = built.model.parameter_map()
            if set(checkpoint.parameters) != set(parameter_map):
                raise ValueError(
                    "inference checkpoint parameter names do not match graph"
                )
            with torch.no_grad():
                for name, parameter in parameter_map.items():
                    restored = checkpoint.parameters[name]
                    if (
                        restored.shape != parameter.shape
                        or restored.dtype != parameter.dtype
                    ):
                        raise ValueError(
                            f"inference checkpoint parameter {name} expected shape={list(parameter.shape)} dtype={parameter.dtype}, "
                            f"got shape={list(restored.shape)} dtype={restored.dtype}"
                        )
                    parameter.copy_(restored)
            checkpoint_provenance = {
                "format": TRAINING_CHECKPOINT_SCHEMA,
                "path": str(checkpoint_path),
                "graph_digest": checkpoint.graph_digest,
                "training_digest": checkpoint.training_digest,
                "completed_updates": checkpoint.completed_updates,
                "selected_loss": checkpoint.selected_loss,
            }
        else:
            state_dict = torch.load(
                checkpoint_path, map_location=device, weights_only=True
            )
            checkpoint_provenance = {
                "format": "graph_torch_state_dict",
                "path": str(checkpoint_path),
            }
            if "W_ff.0" in state_dict:
                imported = import_legacy_parameters_v1(
                    built.model.plan.graph, state_dict, device=device
                )
                with torch.no_grad():
                    for name, parameter in built.model.parameter_map().items():
                        parameter.copy_(imported.parameters[name])
                checkpoint_provenance.update(
                    format="legacy_torch_state_dict",
                    interchange=imported.provenance,
                )
            else:
                built.model.load_state_dict(state_dict)
    scales = dict(overrides.get("projection_scales", {}))
    if scales:
        projection_parameters = {
            row["id"]: row["parameters"][0]
            for row in built.model.plan.graph.get("projections", [])
        }
        unknown = sorted(set(scales) - set(projection_parameters))
        if unknown:
            raise ValueError(
                f"inference projection scales target unknown projections: {unknown}"
            )
        with torch.no_grad():
            parameters = built.model.parameter_map()
            for projection_id, factor in sorted(scales.items()):
                factor = float(factor)
                if not math.isfinite(factor) or factor < 0:
                    raise ValueError(
                        f"inference projection scale {projection_id} must be finite and non-negative"
                    )
                parameters[projection_parameters[projection_id]].mul_(factor)
    tracemalloc.start()
    started = time.perf_counter()
    if spec.dataset_binding is not None:
        if (
            spec.input_bindings
            or spec.event_bindings
            or poisson_bindings
            or spec.inputs
        ):
            raise ValueError(
                "dataset snapshot binding cannot be combined with other input bindings"
            )
        resolved_inputs, _ = resolve_dataset_snapshot_binding(
            built.model.plan.graph,
            spec.dataset_binding,
            device=device,
            execution_seed=spec.seed,
            protocol=spec.protocol,
        )
    else:
        resolved_inputs = resolve_input_bindings(
            built.model.plan.graph,
            dense_bindings=spec.input_bindings,
            event_bindings=spec.event_bindings,
            poisson_bindings=poisson_bindings,
            inputs=spec.inputs,
            device=device,
            seed=spec.seed,
            protocol=spec.protocol,
        )
    result = built.model(
        resolved_inputs.tensors,
        record=spec.recording,
        **(
            {"recording_fields": spec.recording_fields}
            if spec.recording_fields is not None
            else {}
        ),
        runtime_state=runtime_state
        if runtime_state is not None
        else spec.runtime_state,
        interventions=interventions,
    )
    elapsed = time.perf_counter() - started
    _, peak = tracemalloc.get_traced_memory()
    tracemalloc.stop()
    resolved_interventions = result.metrics.pop("resolved_interventions", [])
    result.metrics.update(
        {
            "simulate_s": elapsed,
            "peak_python_bytes": peak,
            "device": device,
            "recording": spec.recording,
            "execution_protocol": resolved_inputs.protocol,
            "checkpoint": checkpoint_provenance,
            "inference_overrides": {
                "schema": INFERENCE_OVERRIDE_SCHEMA,
                "requested": overrides,
                "resolved": {
                    "duration_ms": resolved_inputs.protocol["timing"]["duration_ms"],
                    "timestep_ms": built.model.plan.dt_ms,
                    "projection_scales": scales,
                    **(
                        {"input_rate_hz": float(overrides["input_rate_hz"])}
                        if "input_rate_hz" in overrides
                        else {}
                    ),
                },
            }
            if overrides
            else None,
            "inference_interventions": {
                "schema": INFERENCE_INTERVENTION_SCHEMA,
                "requested": [dict(row) for row in requested_interventions],
                "resolved": resolved_interventions,
            }
            if interventions
            else None,
            "source_graph_digest": source_graph_digest,
            "effective_graph_digest": _json_digest(built.model.plan.graph),
            **built.metrics,
        }
    )
    if result.runtime_state is not None:
        result.metrics.update(
            {
                "runtime_state_schema": RUNTIME_STATE_SCHEMA,
                "runtime_state_signature": result.runtime_state.signature,
                "completed_steps": result.runtime_state.completed_steps,
            }
        )
    return result

train

View source

def train(spec: ExecutionSpec) -> ExecutionResult

Execute validated graph-native training with deterministic mini-batches, named integer targets, constrained parameter updates and portable checkpoint state. Return training metrics, tensors, gradients and selected/final checkpoints. A typed legacy request returns routing metadata.

ParameterAnnotationDefaultMeaning
specExecutionSpecrequiredDefined by the source contract and implementation below.

Return annotation: ExecutionResult.

Return expressions (branch-dependent; names refer to the linked implementation):

ExecutionResult(executor='legacy', metrics={'request': 'train', 'routing': 'legacy'})
ExecutionResult(executor='graph', outputs=final_forward.outputs, recordings=final_forward.recordings, parameters={name: value.detach().clone() for name, value in parameter_map.items()}, gradients=last_gradients, optimizer_state=optimizer_state_by_name(), training_checkpoint=final_checkpoint, selected_checkpoint=selected_checkpoint, model=model, metrics={**built.metrics, 'updates': history, 'execution_protocol': resolved_inputs.protocol, 'trainable_parameters': sorted(last_gradients), 'optimizer': optimizer_spec, 'training_checkpoint_schema': TRAINING_CHECKPOINT_SCHEMA, 'training_checkpoint_rng': {'backend': final_checkpoint.rng_backend, 'devices': sorted(final_checkpoint.accelerator_rng_states)}, 'resumed_from_update': completed_updates})

Explicit exceptions in this implementation; called helpers may raise additional errors:

Explicit exception expression
ValueError('graph training requires a training recipe or training bundle')
ValueError('graph training requires external target tensors')
ValueError('graph training inputs must share one dataset sample axis')
ValueError('graph training batch size must be positive')
ValueError('graph training requires at least one trainable parameter group')
ValueError(f"graph training unsupported optimizer {optimizer_spec.get('kind')}")
ValueError('graph training updates must be positive')
ValueError('dataset snapshot binding cannot be combined with other input or target bindings')
ValueError(f'training checkpoint graph digest expected {graph_digest}, got {resumed.graph_digest}')
ValueError(f'training checkpoint recipe digest expected {training_digest}, got {resumed.training_digest}')
ValueError('training checkpoint execution protocol does not match request')
ValueError('training checkpoint initializer metadata does not match graph build')
ValueError(f'training checkpoint parameter names mismatch; missing={missing}, extra={extra}')
ValueError(f'training checkpoint optimizer parameter names mismatch; missing={missing}, extra={extra}')
ValueError('dataset training checkpoint is already at the requested end epoch')
ValueError('dataset training checkpoint requires epoch and batch data-order state')
ValueError(f'dataset training checkpoint has invalid data-order state {data_state}')
ValueError('single-batch training cannot resume a dataset-iteration checkpoint')
ValueError(f"objective[{index}] unsupported kind {objective.get('kind')}")
ValueError(f"objective[{index}] prediction {objective['prediction']} is not a graph output")
ValueError(f'objective[{index}] missing target tensor {target_name}')
ValueError(f"regularizer[{index}] unsupported kind {regularizer.get('kind')}")
ValueError(f'training checkpoint parameter {name} expected shape={list(parameter.shape)} dtype={parameter.dtype}, got shape={list(restored.shape)} dtype={restored.dtype}')
ValueError(f'regularizer[{index}] spike signal {signal} is not recorded')
Implementation
def train(spec: ExecutionSpec) -> ExecutionResult:
    if spec.executor != "graph":
        return ExecutionResult(
            executor="legacy", metrics={"request": "train", "routing": "legacy"}
        )
    built = build(spec)
    assert isinstance(built.model, GraphExecutor)
    model = built.model
    graph = model.plan.graph
    training = spec.training
    if training is None and spec.bundle is not None:
        manifest, _ = load_graph_bundle(spec.bundle)
        training = load_training_recipe(spec.bundle, manifest, graph)
    if training is None:
        raise ValueError("graph training requires a training recipe or training bundle")
    if (
        not spec.targets
        and not spec.target_bindings
        and not (spec.dataset_binding and spec.dataset_binding.target_id)
    ):
        raise ValueError("graph training requires external target tensors")
    device = resolve_device(spec.device)
    dataset_targets: tuple[TargetArrayBinding, ...] = ()
    if spec.dataset_binding is not None:
        if (
            spec.input_bindings
            or spec.event_bindings
            or spec.poisson_bindings
            or spec.inputs
            or spec.target_bindings
            or spec.targets
        ):
            raise ValueError(
                "dataset snapshot binding cannot be combined with other input or target bindings"
            )
        resolved_inputs, dataset_targets = resolve_dataset_snapshot_binding(
            graph,
            spec.dataset_binding,
            device=device,
            execution_seed=spec.seed,
            protocol=spec.protocol,
        )
    else:
        resolved_inputs = resolve_input_bindings(
            graph,
            dense_bindings=spec.input_bindings,
            event_bindings=spec.event_bindings,
            poisson_bindings=spec.poisson_bindings,
            inputs=spec.inputs,
            device=device,
            seed=spec.seed,
            protocol=spec.protocol,
        )
    dataset_epochs = int(spec.options.get("epochs", 0) or 0)
    dataset_mode = dataset_epochs > 0
    dataset_size = next(iter(resolved_inputs.tensors.values())).shape[1]
    if any(
        value.shape[1] != dataset_size for value in resolved_inputs.tensors.values()
    ):
        raise ValueError("graph training inputs must share one dataset sample axis")
    resolved_targets, target_rows = resolve_target_array_bindings(
        training,
        bindings=dataset_targets or spec.target_bindings,
        targets=spec.targets,
        sample_count=dataset_size,
        device=device,
    )
    batch_size = int(spec.options.get("batch_size") or dataset_size)
    if batch_size <= 0:
        raise ValueError("graph training batch size must be positive")
    batches_per_epoch = math.ceil(dataset_size / batch_size)
    shuffle = bool(spec.options.get("shuffle", False))
    protocol = {**resolved_inputs.protocol, "targets": target_rows}
    if dataset_mode:
        protocol = {
            **protocol,
            "dataset": {
                **protocol["dataset"],
                "sample_cap": dataset_size,
                "batch_size": batch_size,
                "shuffle": shuffle,
            },
            "training_iteration": {
                "schema": "tools/snnsim.dataset-iteration/v1",
                "epochs": dataset_epochs,
                "drop_last": False,
                "order_seed": int(spec.seed),
            },
        }
    resolved_inputs = ResolvedDenseInputs(resolved_inputs.tensors, protocol)
    parameter_map = model.parameter_map()
    graph_digest = training["graph_digest"]
    training_digest = _json_digest(training)
    resumed: TrainingCheckpoint | None = None
    completed_updates = 0
    data_state: dict[str, Any] = {"epoch": 0, "batch": 0} if dataset_mode else {}
    if spec.checkpoint is not None:
        resumed = load_training_checkpoint(spec.checkpoint, device=device)
        if resumed.graph_digest != graph_digest:
            raise ValueError(
                f"training checkpoint graph digest expected {graph_digest}, got {resumed.graph_digest}"
            )
        if resumed.training_digest != training_digest:
            raise ValueError(
                f"training checkpoint recipe digest expected {training_digest}, got {resumed.training_digest}"
            )
        if resumed.execution_protocol != resolved_inputs.protocol:
            raise ValueError(
                "training checkpoint execution protocol does not match request"
            )
        if resumed.initialization != built.metrics["initialization"]:
            raise ValueError(
                "training checkpoint initializer metadata does not match graph build"
            )
        if set(resumed.parameters) != set(parameter_map):
            missing = sorted(set(parameter_map) - set(resumed.parameters))
            extra = sorted(set(resumed.parameters) - set(parameter_map))
            raise ValueError(
                f"training checkpoint parameter names mismatch; missing={missing}, extra={extra}"
            )
        with torch.no_grad():
            for name, parameter in parameter_map.items():
                restored = resumed.parameters[name]
                if (
                    restored.shape != parameter.shape
                    or restored.dtype != parameter.dtype
                ):
                    raise ValueError(
                        f"training checkpoint parameter {name} expected shape={list(parameter.shape)} dtype={parameter.dtype}, "
                        f"got shape={list(restored.shape)} dtype={restored.dtype}"
                    )
                parameter.copy_(restored)
        completed_updates = resumed.completed_updates
        if dataset_mode:
            if set(resumed.data_state) != {"epoch", "batch"}:
                raise ValueError(
                    "dataset training checkpoint requires epoch and batch data-order state"
                )
            data_state = {
                "epoch": int(resumed.data_state["epoch"]),
                "batch": int(resumed.data_state["batch"]),
            }
            epoch = data_state["epoch"]
            batch = data_state["batch"]
            if (
                epoch < 0
                or epoch > dataset_epochs
                or batch < 0
                or batch >= batches_per_epoch
                or (epoch == dataset_epochs and batch != 0)
            ):
                raise ValueError(
                    f"dataset training checkpoint has invalid data-order state {data_state}"
                )
        elif resumed.data_state:
            raise ValueError(
                "single-batch training cannot resume a dataset-iteration checkpoint"
            )
    groups = []
    for group in sorted(
        training.get("parameter_groups", []), key=lambda row: row["id"]
    ):
        if group.get("frozen"):
            continue
        groups.append(
            {
                "params": [parameter_map[name] for name in sorted(group["parameters"])],
                "lr": float(group["lr"]),
                "name": group["id"],
            }
        )
    if not groups:
        raise ValueError(
            "graph training requires at least one trainable parameter group"
        )
    optimizer_spec = training.get("optimizer", {})
    if optimizer_spec.get("kind") != "adamw":
        raise ValueError(
            f"graph training unsupported optimizer {optimizer_spec.get('kind')}"
        )
    optimizer = torch.optim.AdamW(groups, **dict(optimizer_spec.get("config", {})))
    if resumed is not None:
        trainable_names = {
            name for name, parameter in parameter_map.items() if parameter.requires_grad
        }
        if set(resumed.optimizer_state) != trainable_names:
            missing = sorted(trainable_names - set(resumed.optimizer_state))
            extra = sorted(set(resumed.optimizer_state) - trainable_names)
            raise ValueError(
                f"training checkpoint optimizer parameter names mismatch; missing={missing}, extra={extra}"
            )
        for name in sorted(trainable_names):
            optimizer.state[parameter_map[name]] = {
                key: value.to(device) if isinstance(value, torch.Tensor) else value
                for key, value in resumed.optimizer_state[name].items()
            }
        restore_training_rng_state(resumed, device)
    output_ids = {row["signal"]: row["id"] for row in graph.get("outputs", [])}
    updates_option = spec.options.get("updates")
    updates = int(
        updates_option
        if updates_option is not None
        else (
            dataset_epochs * math.ceil(dataset_size / batch_size) if dataset_mode else 1
        )
    )
    if updates <= 0:
        raise ValueError("graph training updates must be positive")
    history = []
    last_gradients: dict[str, torch.Tensor] = {}
    final_forward: ExecutionResult | None = None
    selected_checkpoint: TrainingCheckpoint | None = None

    def optimizer_state_by_name() -> dict[str, dict[str, Any]]:
        packed = {}
        for name, parameter in parameter_map.items():
            if parameter not in optimizer.state:
                continue
            packed[name] = {
                key: value.detach().clone()
                if isinstance(value, torch.Tensor)
                else value
                for key, value in optimizer.state[parameter].items()
            }
        return packed

    def checkpoint_at(
        update_count: int, loss_value: float, next_data_state: Mapping[str, Any]
    ) -> TrainingCheckpoint:
        rng_backend, accelerator_rng_states = capture_training_rng_state(device)
        return TrainingCheckpoint(
            graph_digest=graph_digest,
            training_digest=training_digest,
            completed_updates=update_count,
            selected_loss=loss_value,
            execution_protocol=resolved_inputs.protocol,
            initialization=built.metrics["initialization"],
            parameters={
                name: value.detach().clone() for name, value in parameter_map.items()
            },
            optimizer_state=optimizer_state_by_name(),
            rng_state=torch.get_rng_state().clone(),
            rng_backend=rng_backend,
            accelerator_rng_states=accelerator_rng_states,
            data_state=dict(next_data_state),
        )

    scheduled: list[tuple[int, int, torch.Tensor]] = []
    if dataset_mode:
        for epoch in range(data_state["epoch"], dataset_epochs):
            generator = torch.Generator(device="cpu").manual_seed(spec.seed + epoch)
            order = (
                torch.randperm(dataset_size, generator=generator)
                if shuffle
                else torch.arange(dataset_size)
            )
            first_batch = data_state["batch"] if epoch == data_state["epoch"] else 0
            for batch in range(first_batch, batches_per_epoch):
                scheduled.append(
                    (epoch, batch, order[batch * batch_size : (batch + 1) * batch_size])
                )
        scheduled = scheduled[:updates]
        if not scheduled:
            raise ValueError(
                "dataset training checkpoint is already at the requested end epoch"
            )
    else:
        scheduled = [
            (-1, update, torch.arange(dataset_size)) for update in range(updates)
        ]

    for update, (epoch, batch, sample_indices) in enumerate(scheduled):
        optimizer.zero_grad(set_to_none=True)
        batch_inputs = {
            name: value.index_select(1, sample_indices.to(value.device))
            for name, value in resolved_inputs.tensors.items()
        }
        batch_targets = {
            name: value.index_select(0, sample_indices.to(value.device))
            for name, value in resolved_targets.items()
        }
        forward = model(batch_inputs, record="full")
        components: dict[str, torch.Tensor] = {}
        loss = torch.zeros((), device=device)
        for index, objective in enumerate(training.get("objectives", [])):
            if objective.get("kind") != "cross_entropy":
                raise ValueError(
                    f"objective[{index}] unsupported kind {objective.get('kind')}"
                )
            output_id = output_ids.get(objective["prediction"])
            if output_id is None:
                raise ValueError(
                    f"objective[{index}] prediction {objective['prediction']} is not a graph output"
                )
            target_name = objective["target"]
            if target_name not in resolved_targets:
                raise ValueError(
                    f"objective[{index}] missing target tensor {target_name}"
                )
            target = batch_targets[target_name].to(device=device, dtype=torch.long)
            value = torch.nn.functional.cross_entropy(
                forward.outputs[output_id], target
            ) * float(objective.get("weight", 1.0))
            components[f"objective[{index}]"] = value
            loss = loss + value
        duration = training.get("presentation_duration")
        duration_s = (
            float(duration["value"]) / 1000.0
            if duration
            else resolved_inputs.protocol["timing"]["duration_ms"] / 1000.0
        )
        for index, regularizer in enumerate(training.get("regularizers", [])):
            if regularizer.get("kind") != "spike_budget":
                raise ValueError(
                    f"regularizer[{index}] unsupported kind {regularizer.get('kind')}"
                )
            ceiling = float(regularizer["config"]["ceiling"]["value"])
            penalties = []
            for signal in regularizer["signals"]:
                spikes = forward.recordings.get(signal)
                if spikes is None:
                    raise ValueError(
                        f"regularizer[{index}] spike signal {signal} is not recorded"
                    )
                sample_rates = spikes.sum(dim=0).mean(dim=1) / duration_s
                penalties.append(torch.relu(sample_rates - ceiling).square())
            value = float(regularizer["strength"]) * torch.stack(penalties).mean()
            components[f"regularizer[{index}]"] = value
            loss = loss + value
        loss.backward()
        last_gradients = {
            name: parameter.grad.detach().clone()
            for name, parameter in parameter_map.items()
            if parameter.grad is not None
        }
        clip = training.get("gradient_clip")
        if clip is not None:
            torch.nn.utils.clip_grad_norm_(
                [parameter for group in groups for parameter in group["params"]],
                float(clip),
            )
        optimizer.step()
        rows = {row["id"]: row for row in graph.get("parameters", [])}
        with torch.no_grad():
            for name, parameter in parameter_map.items():
                if (rows[name].get("constraint") or {}).get("kind") == "non_negative":
                    parameter.clamp_(min=0)
        absolute_update = completed_updates + update + 1
        next_data_state: dict[str, Any] = {}
        if dataset_mode:
            next_data_state = {"epoch": epoch, "batch": batch + 1}
            if next_data_state["batch"] == batches_per_epoch:
                next_data_state = {"epoch": epoch + 1, "batch": 0}
        loss_value = float(loss.detach())
        history.append(
            {
                "update": absolute_update,
                **({"epoch": epoch + 1, "batch": batch + 1} if dataset_mode else {}),
                "loss": loss_value,
                "components": {
                    name: float(value.detach()) for name, value in components.items()
                },
            }
        )
        candidate = checkpoint_at(absolute_update, loss_value, next_data_state)
        if selected_checkpoint is None or loss_value < float(
            selected_checkpoint.selected_loss
        ):
            selected_checkpoint = candidate
        final_forward = forward
    assert final_forward is not None
    completed_this_call = len(scheduled)
    final_checkpoint = checkpoint_at(
        completed_updates + completed_this_call,
        history[-1]["loss"],
        next_data_state,
    )
    if save_final := spec.options.get("save_final_checkpoint"):
        save_training_checkpoint(save_final, final_checkpoint)
    if save_selected := spec.options.get("save_selected_checkpoint"):
        assert selected_checkpoint is not None
        save_training_checkpoint(save_selected, selected_checkpoint)

    return ExecutionResult(
        executor="graph",
        outputs=final_forward.outputs,
        recordings=final_forward.recordings,
        parameters={
            name: value.detach().clone() for name, value in parameter_map.items()
        },
        gradients=last_gradients,
        optimizer_state=optimizer_state_by_name(),
        training_checkpoint=final_checkpoint,
        selected_checkpoint=selected_checkpoint,
        model=model,
        metrics={
            **built.metrics,
            "updates": history,
            "execution_protocol": resolved_inputs.protocol,
            "trainable_parameters": sorted(last_gradients),
            "optimizer": optimizer_spec,
            "training_checkpoint_schema": TRAINING_CHECKPOINT_SCHEMA,
            "training_checkpoint_rng": {
                "backend": final_checkpoint.rng_backend,
                "devices": sorted(final_checkpoint.accelerator_rng_states),
            },
            "resumed_from_update": completed_updates,
        },
    )

infer

View source

def infer(spec: ExecutionSpec) -> ExecutionResult

Execute graph-native inference through simulate and mark the request as inference in the result metrics. Legacy requests return routing metadata.

ParameterAnnotationDefaultMeaning
specExecutionSpecrequiredDefined by the source contract and implementation below.

Return annotation: ExecutionResult.

Return expressions (branch-dependent; names refer to the linked implementation):

simulate(spec) if spec.executor == 'graph' else ExecutionResult(executor='legacy', metrics={'request': 'infer', 'routing': 'legacy'})
Implementation
def infer(spec: ExecutionSpec) -> ExecutionResult:
    return (
        simulate(spec)
        if spec.executor == "graph"
        else ExecutionResult(
            executor="legacy", metrics={"request": "infer", "routing": "legacy"}
        )
    )

execution_spec_from_args

View source

def execution_spec_from_args(args: Any, *, kind: RequestKind | None=None) -> ExecutionSpec

Source docstring:

Compatibility adapter: resolved CLI arguments become one typed request.
ParameterAnnotationDefaultMeaning
argsAnyrequiredDefined by the source contract and implementation below.
kindRequestKind | NoneNoneDefined by the source contract and implementation below.

Return annotation: ExecutionSpec.

Return expressions (branch-dependent; names refer to the linked implementation):

ExecutionSpec(kind=resolved_kind, executor=getattr(args, 'executor', 'legacy'), bundle=Path(args.bundle) if getattr(args, 'bundle', None) else None, seed=int(getattr(args, 'seed', 0) or 0), device=resolve_device(getattr(args, 'device', 'auto')), recording=getattr(args, 'recording', 'full'), checkpoint=Path(args.load_weights) if getattr(args, 'load_weights', None) else None, options={key: value for key, value in vars(args).items() if key not in {'bundle', 'executor'}})
Implementation
def execution_spec_from_args(
    args: Any, *, kind: RequestKind | None = None
) -> ExecutionSpec:
    """Compatibility adapter: resolved CLI arguments become one typed request."""
    resolved_kind = kind or ("infer" if getattr(args, "infer", False) else args.mode)
    if resolved_kind == "sim":
        resolved_kind = "simulate"
    return ExecutionSpec(
        kind=resolved_kind,
        executor=getattr(args, "executor", "legacy"),
        bundle=Path(args.bundle) if getattr(args, "bundle", None) else None,
        seed=int(getattr(args, "seed", 0) or 0),
        device=resolve_device(getattr(args, "device", "auto")),
        recording=getattr(args, "recording", "full"),
        checkpoint=(
            Path(args.load_weights) if getattr(args, "load_weights", None) else None
        ),
        options={
            key: value
            for key, value in vars(args).items()
            if key not in {"bundle", "executor"}
        },
    )

resolve_device

View source

def resolve_device(requested: str | torch.device='auto') -> str

Source docstring:

Resolve an explicit device or select the fastest available accelerator.
ParameterAnnotationDefaultMeaning
requestedstr | torch.device'auto'Defined by the source contract and implementation below.

Return annotation: str.

Return expressions (branch-dependent; names refer to the linked implementation):

resolve_device(forced)
'cuda'
'cpu'
name

Explicit exceptions in this implementation; called helpers may raise additional errors:

Explicit exception expression
ValueError('CUDA was requested but torch.cuda.is_available() is false')
ValueError('MPS was requested but torch.backends.mps.is_available() is false')
ValueError(f'device expected auto, cpu, cuda, cuda:N, or mps; got {requested!r}')
ValueError(f'{name} was requested but torch.cuda.is_available() is false')
Implementation
def resolve_device(requested: str | torch.device = "auto") -> str:
    """Resolve an explicit device or select the fastest available accelerator."""
    name = str(requested).lower()
    if name == "auto":
        forced = os.environ.get("PINGLAB_DEVICE")
        if forced:
            return resolve_device(forced)
        if torch.cuda.is_available():
            return "cuda"
        # Graph execution launches several small kernels from Python per timestep.
        # On the representative 800E/200I graph MPS is slower than CPU, so keep it
        # available explicitly without selecting it automatically.
        return "cpu"
    if name == "cuda" and not torch.cuda.is_available():
        raise ValueError("CUDA was requested but torch.cuda.is_available() is false")
    if name == "mps" and not torch.backends.mps.is_available():
        raise ValueError(
            "MPS was requested but torch.backends.mps.is_available() is false"
        )
    if (
        name != "cpu"
        and name != "cuda"
        and name != "mps"
        and not name.startswith("cuda:")
    ):
        raise ValueError(
            f"device expected auto, cpu, cuda, cuda:N, or mps; got {requested!r}"
        )
    if name.startswith("cuda:") and not torch.cuda.is_available():
        raise ValueError(f"{name} was requested but torch.cuda.is_available() is false")
    return name

execute_request

View source

def execute_request(spec: ExecutionSpec, *, legacy: Callable[[], ExecutionResult] | None=None) -> ExecutionResult

Source docstring:

Dispatch one typed request; the CLI supplies its unchanged legacy body.
ParameterAnnotationDefaultMeaning
specExecutionSpecrequiredDefined by the source contract and implementation below.
legacyCallable[[], ExecutionResult] | NoneNoneDefined by the source contract and implementation below.

Return annotation: ExecutionResult.

Return expressions (branch-dependent; names refer to the linked implementation):

legacy()
handlers[spec.kind](spec)

Explicit exceptions in this implementation; called helpers may raise additional errors:

Explicit exception expression
ValueError('legacy execution requires the registered legacy request body')
Implementation
def execute_request(
    spec: ExecutionSpec,
    *,
    legacy: Callable[[], ExecutionResult] | None = None,
) -> ExecutionResult:
    """Dispatch one typed request; the CLI supplies its unchanged legacy body."""
    if spec.executor == "legacy":
        if legacy is None:
            raise ValueError(
                "legacy execution requires the registered legacy request body"
            )
        return legacy()
    handlers = {"build": build, "simulate": simulate, "train": train, "infer": infer}
    return handlers[spec.kind](spec)

Constants and type aliases

Initial source expressions are shown, not evaluated runtime values. Legacy configuration may mutate module defaults.

NameAnnotationInitial expressionSource
ExecutorNameunannotatedLiteral['legacy', 'graph']Source
RequestKindunannotatedLiteral['build', 'simulate', 'train', 'infer']Source
RecordingProfileunannotatedLiteral['full', 'observables', 'none']Source
DENSE_ARRAY_BINDING_SCHEMAunannotated'tools/snnsim.dense-array-binding/v1'Source
EVENT_STREAM_BINDING_SCHEMAunannotated'tools/snnsim.event-stream-binding/v1'Source
MIXED_INPUT_BINDING_SCHEMAunannotated'tools/snnsim.mixed-input-bindings/v1'Source
POISSON_INPUT_BINDING_SCHEMAunannotated'tools/snnsim.poisson-input-binding/v1'Source
DATASET_SNAPSHOT_BINDING_SCHEMAunannotated'tools/snnsim.dataset-snapshot-binding/v1'Source
EXECUTION_PROTOCOL_SCHEMAunannotated'tools/snnsim.execution-protocol/v1'Source
INFERENCE_OVERRIDE_SCHEMAunannotated'tools/snnsim.inference-overrides/v1'Source
INFERENCE_INTERVENTION_SCHEMAunannotated'tools/snnsim.inference-interventions/v1'Source
INFERENCE_ARTIFACT_SCHEMAunannotated'tools/snnsim.inference-artifacts/v1'Source
DERIVED_INFERENCE_SCHEMAunannotated'tools/snnsim.derived-inference/v1'Source
TRAINING_CHECKPOINT_SCHEMAunannotated'tools/snnsim.training-checkpoint/v1'Source
LEGACY_PARAMETER_INTERCHANGE_SCHEMAunannotated'tools/snnsim.legacy-parameter-interchange/v1'Source
GRAPH_CAPABILITIES_V1unannotated{'schema': 'tools/snnsim.capabilities/v1', 'neurons': {'coba_lif', 'leaky_integrator'}, 'synapses': {'ampa', 'gaba', 'leaky_integrator'}, 'operations': {'linear', 'reduce_mean', 'reduce_sum', 'select_final', 'duration_normalise', 'cumulative_sum'}, 'connections': {'feedforward', 'recurrent', 'feedback'}, 'recordings': {'spikes', 'voltage'}, 'delays': 'integer_steps', 'training': {'objectives': {'cross_entropy'}, 'regularizers': {'spike_budget'}, 'optimizers': {'adamw'}, 'parameter_groups': 'named_trainable_and_frozen', 'updates': 'deterministic_epochs_and_minibatches', 'targets': 'named_integer_arrays'}}Source
RUNTIME_STATE_SCHEMAunannotated'tools/snnsim.graph-runtime-state/v1'Source

On this page