Source code for pyqit.core.callbacks.base

"""Backend-neutral callback protocol."""

from dataclasses import dataclass, field
from typing import Any

from pyqit.base.base_object import _PyQitObject


[docs] @dataclass class LoopState: """Everything a callback may read, and the one flag it may write. Attributes ---------- model : BaseModel The model being trained. datamodule : DataModule The data it is training on, already set up. history : TrainingHistory Metrics recorded so far this run. reporter : Reporter Console output, for callbacks that announce something. max_epochs : int Epoch budget for the run. epoch : int Zero-based index of the epoch just finished; ``-1`` before the first. metrics : dict of {str: float} This epoch's metrics, keyed ``train_loss``, ``val_loss``, ``train_acc``, ``val_acc``, ``epoch_time``. stop : bool Set by a callback to end training after this epoch. Both loops check it; the Lightning loop forwards it to ``trainer.should_stop``. optimizer : object The live optimizer, once the loop has built it: a ``qml`` optimizer or a ``torch.optim`` one. ``None`` during ``on_fit_start``. optimizer_state : object Set during ``on_fit_start`` by a callback restoring a run; the loop loads it into the optimizer it builds. Backend-specific. """ model: Any datamodule: Any history: Any reporter: Any max_epochs: int epoch: int = -1 metrics: dict[str, float] = field(default_factory=dict) stop: bool = False optimizer: Any = None optimizer_state: Any = None
[docs] class BaseCallback(_PyQitObject): """Base class for pyqit callbacks. Override any of the three hooks. Each takes one `LoopState`, and both training loops call them, so a subclass works on either backend. Examples -------- >>> from pyqit.core import BaseCallback >>> class StopWhenConverged(BaseCallback): ... def on_epoch_end(self, state): ... if state.metrics["train_loss"] < 0.01: ... state.stop = True """ _tags = { "object_type": "callback", }
[docs] def on_fit_start(self, state: LoopState) -> None: """Called once, after setup and before the first epoch."""
[docs] def on_epoch_end(self, state: LoopState) -> None: """Called once per epoch, with ``state.metrics`` filled for that epoch."""
[docs] def on_fit_end(self, state: LoopState) -> None: """Called once, after the last epoch, including after an early stop."""
[docs] @classmethod def get_test_params(cls): """List constructor kwargs used to parametrize this class in the test suite.""" return [{}]