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 [{}]