snnlab
API referencesnnlab.lang

snnlab.lang.training

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

Back to lang reference

Training declarations are serialized by compile(training=...). They specify objectives, parameter scope, optimizer and gradient rules; concrete datasets, targets and execution ordering remain request inputs.

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.

SignalLike

View source

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

Declared fields, including fields inherited from local data classes:

FieldAnnotationDefaultMeaning
idstrrequiredStable identifier in the relevant graph or data contract.
Complete class implementation
class SignalLike(Protocol):
    id: str

Objective

View source

Class decorators: dataclass(frozen=True).

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

Objective(kind: str, prediction: str, target: str, weight: float = 1.0)

Declared fields, including fields inherited from local data classes:

FieldAnnotationDefaultMeaning
kindstrrequiredStored member of this data contract; see the class docstring and serialization methods.
predictionstrrequiredStored member of this data contract; see the class docstring and serialization methods.
targetstrrequiredStored member of this data contract; see the class docstring and serialization methods.
weightfloat1.0Stored member of this data contract; see the class docstring and serialization methods.
Complete class implementation
class Objective:
    kind: str
    prediction: str
    target: str
    weight: float = 1.0

CrossEntropy

View source

def CrossEntropy(*, prediction: SignalLike | str, target: str, weight: float=1.0) -> Objective

Create a cross_entropy objective binding prediction.id (or a supplied signal id string) to a named integer target, with the supplied scalar objective weight.

ParameterAnnotationDefaultMeaning
predictionSignalLike | strrequiredDefined by the source contract and implementation below.
targetstrrequiredDefined by the source contract and implementation below.
weightfloat1.0Defined by the source contract and implementation below.

Return annotation: Objective.

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

Objective('cross_entropy', value, target, weight)
Implementation
def CrossEntropy(
    *, prediction: SignalLike | str, target: str, weight: float = 1.0
) -> Objective:
    value = prediction if isinstance(prediction, str) else prediction.id
    return Objective("cross_entropy", value, target, weight)

ParameterGroup

View source

Named trainable or frozen parameter scope. Groups must be exhaustive and non-overlapping. Frozen groups use zero learning rate; trainable groups require a positive finite rate.

Class decorators: dataclass(frozen=True).

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

ParameterGroup(parameters: Sequence[ParameterRef | str], name: str, lr: float, frozen: bool = False)

Declared fields, including fields inherited from local data classes:

FieldAnnotationDefaultMeaning
parametersSequence[ParameterRef | str]requiredStored member of this data contract; see the class docstring and serialization methods.
namestrrequiredName used to identify the authored or rendered object.
lrfloatrequiredLearning rate; frozen groups require zero and trainable groups require a positive value.
frozenboolFalseWhether this parameter scope is excluded from optimizer updates.

ParameterGroup.ids

View source

def ParameterGroup.ids(self) -> list[str]

Return annotation: list[str].

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

[p.id if isinstance(p, ParameterRef) else p for p in self.parameters]
Implementation
def ids(self) -> list[str]:
        return [p.id if isinstance(p, ParameterRef) else p for p in self.parameters]
Complete class implementation
class ParameterGroup:
    parameters: Sequence[ParameterRef | str]
    name: str
    lr: float
    frozen: bool = False

    def ids(self) -> list[str]:
        return [p.id if isinstance(p, ParameterRef) else p for p in self.parameters]

Regularizer

View source

Class decorators: dataclass(frozen=True).

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

Regularizer(kind: str, signals: tuple[str, ...], strength: float, config: dict[str, Any] = field(default_factory=dict))

Declared fields, including fields inherited from local data classes:

FieldAnnotationDefaultMeaning
kindstrrequiredStored member of this data contract; see the class docstring and serialization methods.
signalstuple[str, ...]requiredStored member of this data contract; see the class docstring and serialization methods.
strengthfloatrequiredRegularizer or adaptation scale, according to the containing contract.
configdict[str, Any]field(default_factory=dict)Stored member of this data contract; see the class docstring and serialization methods.
Complete class implementation
class Regularizer:
    kind: str
    signals: tuple[str, ...]
    strength: float
    config: dict[str, Any] = field(default_factory=dict)

UpperRatePenalty

View source

def UpperRatePenalty(*, signal: Signal, threshold: float, strength: float) -> Regularizer

Compatibility convenience for SpikeBudgetPenalty on a single signal, using threshold as the rate ceiling in Hz.

ParameterAnnotationDefaultMeaning
signalSignalrequiredDefined by the source contract and implementation below.
thresholdfloatrequiredDefined by the source contract and implementation below.
strengthfloatrequiredRegularizer or adaptation scale, according to the containing contract.

Return annotation: Regularizer.

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

SpikeBudgetPenalty(signals=(signal,), ceiling_hz=threshold, strength=strength)
Implementation
def UpperRatePenalty(
    *, signal: Signal, threshold: float, strength: float
) -> Regularizer:
    return SpikeBudgetPenalty(
        signals=(signal,), ceiling_hz=threshold, strength=strength
    )

SpikeBudgetPenalty

View source

def SpikeBudgetPenalty(*, signals: Sequence[Signal | str], ceiling_hz: float, strength: float) -> Regularizer

Declare a squared-hinge penalty above ceiling_hz, aggregated as the mean over presentations and layers of each population's mean-rate overshoot squared. strength scales the penalty.

ParameterAnnotationDefaultMeaning
signalsSequence[Signal | str]requiredDefined by the source contract and implementation below.
ceiling_hzfloatrequiredMean-rate ceiling in spikes per second.
strengthfloatrequiredRegularizer or adaptation scale, according to the containing contract.

Return annotation: Regularizer.

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

Regularizer('spike_budget', ids, strength, {'ceiling': {'value': ceiling_hz, 'unit': 'Hz'}, 'penalty': 'squared_hinge', 'aggregation': 'mean_presentations_then_layers_of_population_mean_rate'})
Implementation
def SpikeBudgetPenalty(
    *, signals: Sequence[Signal | str], ceiling_hz: float, strength: float
) -> Regularizer:
    ids = tuple(signal if isinstance(signal, str) else signal.id for signal in signals)
    return Regularizer(
        "spike_budget",
        ids,
        strength,
        {
            "ceiling": {"value": ceiling_hz, "unit": "Hz"},
            "penalty": "squared_hinge",
            "aggregation": "mean_presentations_then_layers_of_population_mean_rate",
        },
    )

Optimizer

View source

Class decorators: dataclass(frozen=True).

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

Optimizer(kind: str, config: dict[str, Any] = field(default_factory=dict))

Declared fields, including fields inherited from local data classes:

FieldAnnotationDefaultMeaning
kindstrrequiredStored member of this data contract; see the class docstring and serialization methods.
configdict[str, Any]field(default_factory=dict)Stored member of this data contract; see the class docstring and serialization methods.
Complete class implementation
class Optimizer:
    kind: str
    config: dict[str, Any] = field(default_factory=dict)

AdamW

View source

def AdamW(**config: Any) -> Optimizer

Create an adamw optimizer declaration. Keyword configuration is stored for recipe validation and execution rather than creating a torch optimizer immediately.

ParameterAnnotationDefaultMeaning
**configAnyvariadicDefined by the source contract and implementation below.

Return annotation: Optimizer.

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

Optimizer('adamw', config)
Implementation
def AdamW(**config: Any) -> Optimizer:
    return Optimizer("adamw", config)

FastSigmoid

View source

def FastSigmoid(*, slope: float=1.0) -> Spec

Source docstring:

Fast-sigmoid surrogate used by the collection's spike backward pass.
ParameterAnnotationDefaultMeaning
slopefloat1.0Defined by the source contract and implementation below.

Return annotation: Spec.

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

Spec('fast_sigmoid', {'slope': slope})
Implementation
def FastSigmoid(*, slope: float = 1.0) -> Spec:
    """Fast-sigmoid surrogate used by the collection's spike backward pass."""
    return Spec("fast_sigmoid", {"slope": slope})

StopGradient

View source

Class decorators: dataclass(frozen=True).

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

StopGradient(signal: str)

Declared fields, including fields inherited from local data classes:

FieldAnnotationDefaultMeaning
signalstrrequiredStored member of this data contract; see the class docstring and serialization methods.

StopGradient.at

View source

Decorators: classmethod.

def StopGradient.at(cls, signal: Signal) -> 'StopGradient'

Declare a stop-gradient boundary at a supplied Signal's id.

ParameterAnnotationDefaultMeaning
signalSignalrequiredDefined by the source contract and implementation below.

Return annotation: 'StopGradient'.

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

cls(signal.id)
Implementation
def at(cls, signal: Signal) -> "StopGradient":
        return cls(signal.id)
Complete class implementation
class StopGradient:
    signal: str

    @classmethod
    def at(cls, signal: Signal) -> "StopGradient":
        return cls(signal.id)

TrainSpec

View source

Full training recipe. Compile it with the graph to resolve parameter groups and check objective/regularizer reachability. presentation_duration is physical time; epochs and gradient clipping are recipe settings.

Class decorators: dataclass.

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

TrainSpec(objectives: Sequence[Objective], parameter_groups: Sequence[ParameterGroup], optimizer: Optimizer, regularizers: Sequence[Regularizer] = (), stop_gradients: Sequence[StopGradient] = (), epochs: int = 1, gradient_clip: float | None = None, surrogate: Spec | None = None, presentation_duration: Quantity | None = None)

Declared fields, including fields inherited from local data classes:

FieldAnnotationDefaultMeaning
objectivesSequence[Objective]requiredStored member of this data contract; see the class docstring and serialization methods.
parameter_groupsSequence[ParameterGroup]requiredStored member of this data contract; see the class docstring and serialization methods.
optimizerOptimizerrequiredStored member of this data contract; see the class docstring and serialization methods.
regularizersSequence[Regularizer]()Stored member of this data contract; see the class docstring and serialization methods.
stop_gradientsSequence[StopGradient]()Stored member of this data contract; see the class docstring and serialization methods.
epochsint1Number of training passes over the selected presentations.
gradient_clipfloat | NoneNoneStored member of this data contract; see the class docstring and serialization methods.
surrogateSpec | NoneNoneStored member of this data contract; see the class docstring and serialization methods.
presentation_durationQuantity | NoneNoneStored member of this data contract; see the class docstring and serialization methods.
Complete class implementation
class TrainSpec:
    objectives: Sequence[Objective]
    parameter_groups: Sequence[ParameterGroup]
    optimizer: Optimizer
    regularizers: Sequence[Regularizer] = ()
    stop_gradients: Sequence[StopGradient] = ()
    epochs: int = 1
    gradient_clip: float | None = None
    surrogate: Spec | None = None
    presentation_duration: Quantity | None = None

On this page