snnlab.sim.execution
Complete declared API of the execution module, with signatures, data fields, validation and source.
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
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:
| Field | Annotation | Default | Meaning |
|---|---|---|---|
input_id | str | required | Exact declared graph input id. |
value | torch.Tensor | required | Stored member of this data contract; see the class docstring and serialization methods. |
source | Mapping[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
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:
| Field | Annotation | Default | Meaning |
|---|---|---|---|
target_id | str | required | Exact named training target id. |
value | torch.Tensor | required | Stored member of this data contract; see the class docstring and serialization methods. |
source | Mapping[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
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:
| Field | Annotation | Default | Meaning |
|---|---|---|---|
input_id | str | required | Exact declared graph input id. |
steps | torch.Tensor | required | Stored member of this data contract; see the class docstring and serialization methods. |
batches | torch.Tensor | required | Stored member of this data contract; see the class docstring and serialization methods. |
channels | torch.Tensor | required | Stored member of this data contract; see the class docstring and serialization methods. |
steps_count | int | required | Number of simulation timesteps in an event or Poisson binding. |
batch_size | int | required | Number of presentations in a binding or mini-batch. |
source | Mapping[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
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:
| Field | Annotation | Default | Meaning |
|---|---|---|---|
input_id | str | required | Exact declared graph input id. |
steps_count | int | required | Number of simulation timesteps in an event or Poisson binding. |
batch_size | int | required | Number of presentations in a binding or mini-batch. |
rates_hz | Sequence[float] | required | Configured input rates in spikes per second. |
seed | int | required | Seed controlling this operation’s random stream. |
categorical | bool | False | Stored 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 = FalseDatasetEncoder
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:
| Field | Annotation | Default | Meaning |
|---|---|---|---|
kind | Literal['rate_poisson', 'prebinned_spikes', 'event_bin'] | required | Stored member of this data contract; see the class docstring and serialization methods. |
duration_ms | float | None | None | Physical presentation duration in milliseconds. |
max_rate_hz | float | None | None | Maximum rate in spikes per second for encoded input. |
seed | int | 0 | Seed 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 = 0DatasetSnapshotBinding
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:
| Field | Annotation | Default | Meaning |
|---|---|---|---|
path | Path | required | Filesystem source or destination path, as described below. |
input_id | str | required | Exact declared graph input id. |
dataset_id | str | required | Stored member of this data contract; see the class docstring and serialization methods. |
split | str | required | Stored member of this data contract; see the class docstring and serialization methods. |
encoder | DatasetEncoder | required | Stored member of this data contract; see the class docstring and serialization methods. |
target_id | str | None | None | Exact named training target id. |
feature_key | str | 'features' | Stored member of this data contract; see the class docstring and serialization methods. |
label_key | str | 'labels' | Stored member of this data contract; see the class docstring and serialization methods. |
sample_cap | int | None | None | Stored member of this data contract; see the class docstring and serialization methods. |
shuffle | bool | False | Stored member of this data contract; see the class docstring and serialization methods. |
order_seed | int | 0 | Seed 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 = 0ResolvedDenseInputs
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:
| Field | Annotation | Default | Meaning |
|---|---|---|---|
tensors | Mapping[str, torch.Tensor] | required | Stored member of this data contract; see the class docstring and serialization methods. |
protocol | Mapping[str, Any] | required | Stored 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
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:
| Field | Annotation | Default | Meaning |
|---|---|---|---|
kind | RequestKind | required | Stored member of this data contract; see the class docstring and serialization methods. |
executor | ExecutorName | 'legacy' | Execution route: legacy or graph; graph must be selected explicitly. |
bundle | Path | None | None | Compiled data bundle or bundle path, as annotated. |
graph | Mapping[str, Any] | None | None | Serialized graph mapping. |
inputs | Mapping[str, torch.Tensor] | field(default_factory=dict) | Input tensors keyed by graph input id. |
input_bindings | Sequence[DenseArrayBinding] | field(default_factory=tuple) | Named dense input bindings, resolved against graph contracts. |
event_bindings | Sequence[EventStreamBinding] | field(default_factory=tuple) | Named sparse event bindings with zero-based step/batch/channel coordinates. |
poisson_bindings | Sequence[PoissonInputBinding] | field(default_factory=tuple) | Declared generated Poisson input bindings. |
dataset_binding | DatasetSnapshotBinding | None | None | Stored member of this data contract; see the class docstring and serialization methods. |
protocol | Mapping[str, Any] | field(default_factory=dict) | Stored member of this data contract; see the class docstring and serialization methods. |
training | Mapping[str, Any] | None | None | Authored training declaration or serialized training recipe. |
targets | Mapping[str, torch.Tensor] | field(default_factory=dict) | Named integer targets or target objects, according to this contract. |
target_bindings | Sequence[TargetArrayBinding] | field(default_factory=tuple) | Stored member of this data contract; see the class docstring and serialization methods. |
seed | int | 0 | Seed controlling this operation’s random stream. |
device | str | 'auto' | Requested or resolved tensor execution device. |
recording | RecordingProfile | 'full' | Retained Recording or recording selection, as annotated. |
recording_fields | Sequence[str] | None | None | Explicit field names to retain. |
checkpoint | Path | None | None | Training checkpoint record or authenticated checkpoint path. |
runtime_state | GraphRuntimeState | None | None | Dynamic graph state for causal continuation; distinct from training checkpoints. |
options | Mapping[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
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:
| Field | Annotation | Default | Meaning |
|---|---|---|---|
executor | ExecutorName | required | Execution route: legacy or graph; graph must be selected explicitly. |
outputs | dict[str, torch.Tensor] | field(default_factory=dict) | Stored member of this data contract; see the class docstring and serialization methods. |
recordings | dict[str, torch.Tensor] | field(default_factory=dict) | Stored member of this data contract; see the class docstring and serialization methods. |
parameters | dict[str, torch.Tensor] | field(default_factory=dict) | Stored member of this data contract; see the class docstring and serialization methods. |
gradients | dict[str, torch.Tensor] | field(default_factory=dict) | Stored member of this data contract; see the class docstring and serialization methods. |
optimizer_state | dict[str, Any] | field(default_factory=dict) | Stored member of this data contract; see the class docstring and serialization methods. |
training_checkpoint | TrainingCheckpoint | None | None | Stored member of this data contract; see the class docstring and serialization methods. |
selected_checkpoint | TrainingCheckpoint | None | None | Stored member of this data contract; see the class docstring and serialization methods. |
final_state | dict[str, torch.Tensor] | field(default_factory=dict) | Stored member of this data contract; see the class docstring and serialization methods. |
runtime_state | GraphRuntimeState | None | None | Dynamic graph state for causal continuation; distinct from training checkpoints. |
metrics | dict[str, Any] | field(default_factory=dict) | Stored member of this data contract; see the class docstring and serialization methods. |
model | nn.Module | None | None | Stored 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 = NoneCapabilityIssue
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:
| Field | Annotation | Default | Meaning |
|---|---|---|---|
element | str | required | Stored member of this data contract; see the class docstring and serialization methods. |
capability | str | required | Stored member of this data contract; see the class docstring and serialization methods. |
message | str | required | Stored member of this data contract; see the class docstring and serialization methods. |
Complete class implementation
class CapabilityIssue:
element: str
capability: str
message: strTrainingCheckpoint
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:
| Field | Annotation | Default | Meaning |
|---|---|---|---|
graph_digest | str | required | Stored member of this data contract; see the class docstring and serialization methods. |
training_digest | str | required | Stored member of this data contract; see the class docstring and serialization methods. |
completed_updates | int | required | Stored member of this data contract; see the class docstring and serialization methods. |
execution_protocol | Mapping[str, Any] | required | Stored member of this data contract; see the class docstring and serialization methods. |
initialization | Mapping[str, Any] | required | Stored member of this data contract; see the class docstring and serialization methods. |
parameters | dict[str, torch.Tensor] | required | Stored member of this data contract; see the class docstring and serialization methods. |
optimizer_state | dict[str, dict[str, Any]] | required | Stored member of this data contract; see the class docstring and serialization methods. |
rng_state | torch.Tensor | required | Stored member of this data contract; see the class docstring and serialization methods. |
rng_backend | str | 'cpu' | Stored member of this data contract; see the class docstring and serialization methods. |
accelerator_rng_states | dict[str, torch.Tensor] | field(default_factory=dict) | Stored member of this data contract; see the class docstring and serialization methods. |
data_state | Mapping[str, Any] | field(default_factory=dict) | Stored member of this data contract; see the class docstring and serialization methods. |
selected_loss | float | None | None | Stored 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 = NoneParameterInterchange
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:
| Field | Annotation | Default | Meaning |
|---|---|---|---|
parameters | dict[str, torch.Tensor] | required | Stored member of this data contract; see the class docstring and serialization methods. |
provenance | Mapping[str, Any] | required | Stored 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
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.| Parameter | Annotation | Default | Meaning |
|---|---|---|---|
path | str | Path | required | Filesystem source or destination path, as described below. |
graph | Mapping[str, Any] | required | Serialized 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
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.| Parameter | Annotation | Default | Meaning |
|---|---|---|---|
path | str | Path | required | Filesystem source or destination path, as described below. |
training | Mapping[str, Any] | required | Authored 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
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.| Parameter | Annotation | Default | Meaning |
|---|---|---|---|
training | Mapping[str, Any] | required | Authored training declaration or serialized training recipe. |
bindings | Sequence[TargetArrayBinding] | () | Defined by the source contract and implementation below. |
targets | Mapping[str, torch.Tensor] | None | None | Named integer targets or target objects, according to this contract. |
sample_count | int | required | Defined by the source contract and implementation below. |
device | str | 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, rowsload_event_stream_bindings
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.| Parameter | Annotation | Default | Meaning |
|---|---|---|---|
path | str | Path | required | Filesystem source or destination path, as described below. |
graph | Mapping[str, Any] | required | Serialized 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
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) -> ResolvedDenseInputsSource docstring:
Validate dense arrays, resolve symbolic axes, and freeze run provenance.| Parameter | Annotation | Default | Meaning |
|---|---|---|---|
graph | Mapping[str, Any] | required | Serialized graph mapping. |
bindings | Sequence[DenseArrayBinding] | () | Defined by the source contract and implementation below. |
inputs | Mapping[str, torch.Tensor] | None | None | Input tensors keyed by graph input id. |
device | str | torch.device | 'cpu' | Requested or resolved tensor execution device. |
seed | int | 0 | Seed controlling this operation’s random stream. |
protocol | Mapping[str, Any] | None | None | Defined 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
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) -> ResolvedDenseInputsSource docstring:
Validate sparse spike coordinates and materialize binary graph inputs.| Parameter | Annotation | Default | Meaning |
|---|---|---|---|
graph | Mapping[str, Any] | required | Serialized graph mapping. |
bindings | Sequence[EventStreamBinding] | required | Defined by the source contract and implementation below. |
device | str | torch.device | 'cpu' | Requested or resolved tensor execution device. |
seed | int | 0 | Seed controlling this operation’s random stream. |
protocol | Mapping[str, Any] | None | None | Defined 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
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) -> ResolvedDenseInputsSource docstring:
Resolve dense, event-stream, generated-Poisson, or mixed graph inputs.| Parameter | Annotation | Default | Meaning |
|---|---|---|---|
graph | Mapping[str, Any] | required | Serialized graph mapping. |
dense_bindings | Sequence[DenseArrayBinding] | () | Defined by the source contract and implementation below. |
event_bindings | Sequence[EventStreamBinding] | () | Named sparse event bindings with zero-based step/batch/channel coordinates. |
poisson_bindings | Sequence[PoissonInputBinding] | () | Declared generated Poisson input bindings. |
inputs | Mapping[str, torch.Tensor] | None | None | Input tensors keyed by graph input id. |
device | str | torch.device | 'cpu' | Requested or resolved tensor execution device. |
seed | int | 0 | Seed controlling this operation’s random stream. |
protocol | Mapping[str, Any] | None | None | Defined 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
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) -> ResolvedDenseInputsSource docstring:
Generate reproducible fixed or per-presentation categorical Poisson spikes.| Parameter | Annotation | Default | Meaning |
|---|---|---|---|
graph | Mapping[str, Any] | required | Serialized graph mapping. |
bindings | Sequence[PoissonInputBinding] | required | Defined by the source contract and implementation below. |
device | str | torch.device | 'cpu' | Requested or resolved tensor execution device. |
seed | int | 0 | Seed controlling this operation’s random stream. |
protocol | Mapping[str, Any] | None | None | Defined 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
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.| Parameter | Annotation | Default | Meaning |
|---|---|---|---|
graph | Mapping[str, Any] | required | Serialized graph mapping. |
binding | DatasetSnapshotBinding | required | Defined by the source contract and implementation below. |
device | str | torch.device | 'cpu' | Requested or resolved tensor execution device. |
execution_seed | int | 0 | Defined by the source contract and implementation below. |
protocol | Mapping[str, Any] | None | None | Defined 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), targetsgraph_capability_issues
def graph_capability_issues(graph: Mapping[str, Any]) -> list[CapabilityIssue]Source docstring:
Return precise graph-executor capability failures.| Parameter | Annotation | Default | Meaning |
|---|---|---|---|
graph | Mapping[str, Any] | required | Serialized graph mapping. |
Return annotation: list[CapabilityIssue].
Return expressions (branch-dependent; names refer to the linked implementation):
issuesImplementation
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 issuesPlannedProjection
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:
| Field | Annotation | Default | Meaning |
|---|---|---|---|
id | str | required | Stable identifier in the relevant graph or data contract. |
source | str | required | Stored member of this data contract; see the class docstring and serialization methods. |
target | str | required | Stored member of this data contract; see the class docstring and serialization methods. |
polarity | str | required | Stored member of this data contract; see the class docstring and serialization methods. |
decay | float | required | Stored member of this data contract; see the class docstring and serialization methods. |
delay_steps | int | required | Stored member of this data contract; see the class docstring and serialization methods. |
parameter | str | required | Stored member of this data contract; see the class docstring and serialization methods. |
enabled | bool | required | Whether 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: boolGraphPlan
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:
| Field | Annotation | Default | Meaning |
|---|---|---|---|
graph | Mapping[str, Any] | required | Serialized graph mapping. |
dt_ms | float | required | Simulation timestep in milliseconds. |
populations | tuple[Mapping[str, Any], ...] | required | Stored member of this data contract; see the class docstring and serialization methods. |
projections | tuple[PlannedProjection, ...] | required | Stored member of this data contract; see the class docstring and serialization methods. |
observables | tuple[Mapping[str, Any], ...] | required | Stored member of this data contract; see the class docstring and serialization methods. |
outputs | tuple[Mapping[str, Any], ...] | required | Stored 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
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:
| Field | Annotation | Default | Meaning |
|---|---|---|---|
signature | str | required | Stored member of this data contract; see the class docstring and serialization methods. |
compatibility | dict[str, Any] | required | Stored member of this data contract; see the class docstring and serialization methods. |
completed_steps | int | required | Stored member of this data contract; see the class docstring and serialization methods. |
voltages | dict[str, torch.Tensor] | required | Stored member of this data contract; see the class docstring and serialization methods. |
refractory | dict[str, torch.Tensor] | required | Stored member of this data contract; see the class docstring and serialization methods. |
conductances | dict[str, torch.Tensor] | required | Stored member of this data contract; see the class docstring and serialization methods. |
population_histories | dict[str, torch.Tensor] | required | Stored member of this data contract; see the class docstring and serialization methods. |
input_histories | dict[str, torch.Tensor] | required | Stored member of this data contract; see the class docstring and serialization methods. |
GraphRuntimeState.detached
def GraphRuntimeState.detached(self, *, device: str | torch.device='cpu') -> GraphRuntimeState| Parameter | Annotation | Default | Meaning |
|---|---|---|---|
device | str | 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
def runtime_state_compatibility(plan: GraphPlan) -> dict[str, Any]Source docstring:
Describe state-layout and dynamical semantics, excluding parameter values.| Parameter | Annotation | Default | Meaning |
|---|---|---|---|
plan | GraphPlan | required | Lowered 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
def runtime_state_signature(plan: GraphPlan) -> strHash the canonical runtime compatibility mapping for a planned graph; use this signature to reject continuation against an incompatible plan.
| Parameter | Annotation | Default | Meaning |
|---|---|---|---|
plan | GraphPlan | required | Lowered 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
def save_runtime_state(path: str | Path, state: GraphRuntimeState) -> PathSource docstring:
Atomically publish a portable JSON/NPZ graph-runtime state directory.| Parameter | Annotation | Default | Meaning |
|---|---|---|---|
path | str | Path | required | Filesystem source or destination path, as described below. |
state | GraphRuntimeState | required | Defined by the source contract and implementation below. |
Return annotation: Path.
Return expressions (branch-dependent; names refer to the linked implementation):
rootImplementation
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 rootload_runtime_state
def load_runtime_state(path: str | Path, *, device: str | torch.device='cpu') -> GraphRuntimeStateSource docstring:
Load and authenticate a portable graph-runtime state artifact.| Parameter | Annotation | Default | Meaning |
|---|---|---|---|
path | str | Path | required | Filesystem source or destination path, as described below. |
device | str | 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
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.| Parameter | Annotation | Default | Meaning |
|---|---|---|---|
path | str | Path | required | Filesystem source or destination path, as described below. |
result | ExecutionResult | required | Defined by the source contract and implementation below. |
graph | Mapping[str, Any] | required | Serialized graph mapping. |
seed | int | required | Seed controlling this operation’s random stream. |
Return annotation: Mapping[str, Any].
Return expressions (branch-dependent; names refer to the linked implementation):
manifestImplementation
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 manifestvalidate_inference_artifacts
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.| Parameter | Annotation | Default | Meaning |
|---|---|---|---|
path | str | Path | required | Filesystem source or destination path, as described below. |
graph | Mapping[str, Any] | None | None | Serialized graph mapping. |
seed | int | None | None | Seed controlling this operation’s random stream. |
Return annotation: Mapping[str, Any].
Return expressions (branch-dependent; names refer to the linked implementation):
manifestExplicit 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 manifestderive_inference_products
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.| Parameter | Annotation | Default | Meaning |
|---|---|---|---|
source | str | Path | required | Defined by the source contract and implementation below. |
destination | str | Path | required | Defined by the source contract and implementation below. |
logits_id | str | required | Defined by the source contract and implementation below. |
labels | np.ndarray | torch.Tensor | required | Defined by the source contract and implementation below. |
spike_recordings | Sequence[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
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.| Parameter | Annotation | Default | Meaning |
|---|---|---|---|
path | str | Path | required | Filesystem source or destination path, as described below. |
source_artifact_digest | str | None | None | Defined by the source contract and implementation below. |
Return annotation: Mapping[str, Any].
Return expressions (branch-dependent; names refer to the linked implementation):
manifestExplicit 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 manifestsave_training_checkpoint
def save_training_checkpoint(path: str | Path, checkpoint: TrainingCheckpoint) -> PathSource docstring:
Atomically write a named, authenticated graph-training checkpoint.| Parameter | Annotation | Default | Meaning |
|---|---|---|---|
path | str | Path | required | Filesystem source or destination path, as described below. |
checkpoint | TrainingCheckpoint | required | Training checkpoint record or authenticated checkpoint path. |
Return annotation: Path.
Return expressions (branch-dependent; names refer to the linked implementation):
rootExplicit 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 rootload_training_checkpoint
def load_training_checkpoint(path: str | Path, *, device: str | torch.device='cpu') -> TrainingCheckpointSource docstring:
Load and authenticate a portable named graph-training checkpoint.| Parameter | Annotation | Default | Meaning |
|---|---|---|---|
path | str | Path | required | Filesystem source or destination path, as described below. |
device | str | 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
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.| Parameter | Annotation | Default | Meaning |
|---|---|---|---|
device | str | torch.device | required | Requested 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
def restore_training_rng_state(checkpoint: TrainingCheckpoint, device: str | torch.device) -> NoneSource docstring:
Restore CPU and exact-matching accelerator streams or fail closed.| Parameter | Annotation | Default | Meaning |
|---|---|---|---|
checkpoint | TrainingCheckpoint | required | Training checkpoint record or authenticated checkpoint path. |
device | str | torch.device | required | Requested 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
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.| Parameter | Annotation | Default | Meaning |
|---|---|---|---|
graph | Mapping[str, Any] | required | Serialized 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
def import_legacy_parameters_v1(graph: Mapping[str, Any], state_dict: Mapping[str, torch.Tensor], *, device: str | torch.device='cpu') -> ParameterInterchangeSource docstring:
Import the exact supported one-layer legacy parameter state by semantic name.| Parameter | Annotation | Default | Meaning |
|---|---|---|---|
graph | Mapping[str, Any] | required | Serialized graph mapping. |
state_dict | Mapping[str, torch.Tensor] | required | Defined by the source contract and implementation below. |
device | str | 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
def export_legacy_parameters_v1(graph: Mapping[str, Any], parameters: Mapping[str, torch.Tensor]) -> ParameterInterchangeSource docstring:
Export a complete supported graph parameter set under legacy state keys.| Parameter | Annotation | Default | Meaning |
|---|---|---|---|
graph | Mapping[str, Any] | required | Serialized graph mapping. |
parameters | Mapping[str, torch.Tensor] | required | Defined 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
Source docstring:
Fixed causal delay used by recurrent and feedback projections.Constructor:
DelayBuffer(self, delay_steps: int, prototype: torch.Tensor)| Parameter | Annotation | Default | Meaning |
|---|---|---|---|
delay_steps | int | required | Defined by the constructor implementation below. |
prototype | torch.Tensor | required | Defined by the constructor implementation below. |
Instance members assigned by the constructor (expressions are evaluated when constructed):
| Member | Annotation | Initial expression |
|---|---|---|
delay_steps | unannotated | delay_steps |
Constructor/initialization exception expressions:
| Explicit exception expression |
|---|
ValueError('causal delay buffer requires at least one step') |
DelayBuffer.read
def DelayBuffer.read(self) -> torch.TensorReturn 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
def DelayBuffer.push(self, value: torch.Tensor) -> None| Parameter | Annotation | Default | Meaning |
|---|---|---|---|
value | torch.Tensor | required | Defined 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
def DelayBuffer.export(self) -> torch.TensorReturn 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
Decorators: classmethod.
def DelayBuffer.restore(cls, values: torch.Tensor) -> DelayBuffer| Parameter | Annotation | Default | Meaning |
|---|---|---|---|
values | torch.Tensor | required | Defined by the source contract and implementation below. |
Return annotation: DelayBuffer.
Return expressions (branch-dependent; names refer to the linked implementation):
resultExplicit 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 resultComplete 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 resultplan_graph
def plan_graph(graph: Mapping[str, Any]) -> GraphPlanCheck capabilities and lower the full topology before simulation, resolving dimensions, receptor polarity, integral delays and deterministic zero-delay feedforward ordering. Return GraphPlan.
| Parameter | Annotation | Default | Meaning |
|---|---|---|---|
graph | Mapping[str, Any] | required | Serialized 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
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)| Parameter | Annotation | Default | Meaning |
|---|---|---|---|
plan | GraphPlan | required | Lowered GraphPlan for the complete graph. |
seed | int | 0 | Seed controlling this operation’s random stream. |
trainable_parameters | Sequence[str] | () | Defined by the constructor implementation below. |
surrogate_slope | float | M.SURROGATE_SLOPE | Defined by the constructor implementation below. |
Instance members assigned by the constructor (expressions are evaluated when constructed):
| Member | Annotation | Initial expression |
|---|---|---|
plan | unannotated | plan |
surrogate_slope | unannotated | float(surrogate_slope) |
weights | unannotated | nn.ParameterDict() |
initialization_metadata | dict[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
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
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| Parameter | Annotation | Default | Meaning |
|---|---|---|---|
inputs | Mapping[str, torch.Tensor] | required | Input tensors keyed by graph input id. |
record | bool | RecordingProfile | True | Defined by the source contract and implementation below. |
recording_fields | Sequence[str] | None | None | Explicit field names to retain. |
runtime_state | GraphRuntimeState | None | None | Dynamic graph state for causal continuation; distinct from training checkpoints. |
interventions | Sequence[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
def build(spec: ExecutionSpec) -> ExecutionResultPlan 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.
| Parameter | Annotation | Default | Meaning |
|---|---|---|---|
spec | ExecutionSpec | required | Defined 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
def simulate(spec: ExecutionSpec, *, runtime_state: GraphRuntimeState | None=None) -> ExecutionResultExecute 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.
| Parameter | Annotation | Default | Meaning |
|---|---|---|---|
spec | ExecutionSpec | required | Defined by the source contract and implementation below. |
runtime_state | GraphRuntimeState | None | None | Dynamic 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'})resultExplicit 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 resulttrain
def train(spec: ExecutionSpec) -> ExecutionResultExecute 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.
| Parameter | Annotation | Default | Meaning |
|---|---|---|---|
spec | ExecutionSpec | required | Defined 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
def infer(spec: ExecutionSpec) -> ExecutionResultExecute graph-native inference through simulate and mark the request as inference in the result metrics. Legacy requests return routing metadata.
| Parameter | Annotation | Default | Meaning |
|---|---|---|---|
spec | ExecutionSpec | required | Defined 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
def execution_spec_from_args(args: Any, *, kind: RequestKind | None=None) -> ExecutionSpecSource docstring:
Compatibility adapter: resolved CLI arguments become one typed request.| Parameter | Annotation | Default | Meaning |
|---|---|---|---|
args | Any | required | Defined by the source contract and implementation below. |
kind | RequestKind | None | None | Defined 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
def resolve_device(requested: str | torch.device='auto') -> strSource docstring:
Resolve an explicit device or select the fastest available accelerator.| Parameter | Annotation | Default | Meaning |
|---|---|---|---|
requested | str | 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'nameExplicit 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 nameexecute_request
def execute_request(spec: ExecutionSpec, *, legacy: Callable[[], ExecutionResult] | None=None) -> ExecutionResultSource docstring:
Dispatch one typed request; the CLI supplies its unchanged legacy body.| Parameter | Annotation | Default | Meaning |
|---|---|---|---|
spec | ExecutionSpec | required | Defined by the source contract and implementation below. |
legacy | Callable[[], ExecutionResult] | None | None | Defined 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.
| Name | Annotation | Initial expression | Source |
|---|---|---|---|
ExecutorName | unannotated | Literal['legacy', 'graph'] | Source |
RequestKind | unannotated | Literal['build', 'simulate', 'train', 'infer'] | Source |
RecordingProfile | unannotated | Literal['full', 'observables', 'none'] | Source |
DENSE_ARRAY_BINDING_SCHEMA | unannotated | 'tools/snnsim.dense-array-binding/v1' | Source |
EVENT_STREAM_BINDING_SCHEMA | unannotated | 'tools/snnsim.event-stream-binding/v1' | Source |
MIXED_INPUT_BINDING_SCHEMA | unannotated | 'tools/snnsim.mixed-input-bindings/v1' | Source |
POISSON_INPUT_BINDING_SCHEMA | unannotated | 'tools/snnsim.poisson-input-binding/v1' | Source |
DATASET_SNAPSHOT_BINDING_SCHEMA | unannotated | 'tools/snnsim.dataset-snapshot-binding/v1' | Source |
EXECUTION_PROTOCOL_SCHEMA | unannotated | 'tools/snnsim.execution-protocol/v1' | Source |
INFERENCE_OVERRIDE_SCHEMA | unannotated | 'tools/snnsim.inference-overrides/v1' | Source |
INFERENCE_INTERVENTION_SCHEMA | unannotated | 'tools/snnsim.inference-interventions/v1' | Source |
INFERENCE_ARTIFACT_SCHEMA | unannotated | 'tools/snnsim.inference-artifacts/v1' | Source |
DERIVED_INFERENCE_SCHEMA | unannotated | 'tools/snnsim.derived-inference/v1' | Source |
TRAINING_CHECKPOINT_SCHEMA | unannotated | 'tools/snnsim.training-checkpoint/v1' | Source |
LEGACY_PARAMETER_INTERCHANGE_SCHEMA | unannotated | 'tools/snnsim.legacy-parameter-interchange/v1' | Source |
GRAPH_CAPABILITIES_V1 | unannotated | {'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_SCHEMA | unannotated | 'tools/snnsim.graph-runtime-state/v1' | Source |