Trainer#
- class pyqit.core.Trainer(max_epochs: int = 30, learning_rate: float | dict = 0.01, batch_size: int = 32, optimizer: str = 'adam', loss_fn: str | Callable = 'mse', callbacks: list | None = None, verbose: int = 2, seed: int | None = 42, enable_checkpointing: bool = False, checkpoint_dir: str | None = None, logger: bool | object = False, check_bp: bool = False, bp_samples: int = 200, backend_kwargs: dict | None = None)[source]#
Bases:
_PyQitObjectOrchestrates a training run and dispatches it to a backend loop.
The Trainer trains nothing itself. It seeds, sets the DataModule up, prints the run summary, optionally runs the barren-plateau pre-flight, assembles the callback list, and hands off to the loop registered for the active backend.
Its parameters are the ones that mean the same thing on every backend; anything backend-specific goes through
backend_kwargs.- Parameters:
max_epochs (int, default 30)
learning_rate (float or dict, default 0.01) – One rate for every weight, or one per weight group as
{"quantum": ..., "classical": ...}, naming exactly the groupsmodel.weight_groups()has. A pure circuit model has"quantum"only; a hybrid has both.batch_size (int, default 32) – Applied to the DataModule at
setup, overriding its own setting.optimizer ({"adam", "sgd"}, default "adam")
loss_fn (str or callable, default "mse") – A name from
loss_registry(), or a callable taking(preds, y).callbacks (list of BaseCallback, optional) – Run on both backends. Lightning callbacks are not accepted.
verbose ({0, 1, 2}, default 2) –
0silent,1progress only,2progress plus the model table.seed (int, optional) – Seeded before anything stochastic runs. Weights are drawn at model construction, so reproducing them needs
pyqit.set_seedbefore the model is built; this covers training and diagnostics only.enable_checkpointing (bool, default False) – Installs a default
ModelCheckpoint. For any other policy, pass your own throughcallbacks.checkpoint_dir (str, optional) – Directory for the default checkpoint. Defaults to
"checkpoints".logger (bool or object, default False) – Forwarded to Lightning on the torch backend; ignored on pennylane.
check_bp (bool, default False) – Run the barren-plateau gradient-variance check before training.
bp_samples (int, default 200) – Gradients sampled by the check. Each costs one circuit execution under backprop and
1 + 2 * n_paramsunder parameter-shift, which is what shot-based devices and hardware use. The result reports the count. Gradient samples drawn by that check.backend_kwargs (dict, optional) – Forwarded verbatim to
lightning.pytorch.Traineron the torch backend. Rejected on pennylane, which has no underlying trainer.
- fit(model: BaseModel | BaseMetaObject, datamodule: DataModule) TrainingHistory | dict[source]#
Train
modelondatamodule.- Parameters:
model (BaseModel or BaseMetaObject) – Trained in place. A
QuantumPipelinetrains every trainable stage in turn with this Trainer.datamodule (DataModule) – Set up here if it is not already.
- Returns:
Per-epoch losses, accuracies and timings; for a pipeline, one
TrainingHistoryper trained stage keyed by stage name.- Return type:
TrainingHistory or dict
- classmethod get_test_params()[source]#
List constructor kwargs used to parametrize this class in the test suite.
- predict(model: BaseModel | BaseMetaObject, datamodule: DataModule, return_format: str = 'auto') np.ndarray[source]#
Run
modelover the most specific splitdatamoduleholds.- Parameters:
model (BaseModel or BaseMetaObject) – Used for inference only; weights are not touched.
datamodule (DataModule) – Test split if present, else validation, else train.
return_format ({"auto", "numpy", "torch", "pennylane"}, default "auto") –
"auto"follows whatever the model emits."pennylane"returns apennylane.numpy.tensorrather than a barendarray, for callers composing the result into further pnp-based code (a custom cost function,utils.diagnostic).
- Returns:
Predictions for every sample in the chosen split.
- Return type:
numpy.ndarray or torch.Tensor
- test(model: BaseModel | BaseMetaObject, datamodule: DataModule) dict[source]#
Evaluate
modelon the test split ofdatamodule.- Parameters:
model (BaseModel or BaseMetaObject) – Evaluated with its current weights, which are not touched.
datamodule (DataModule) – Set up here if it is not already. Must hold a test split.
- Returns:
test_lossandtest_accunderloss_fn, as floats.- Return type:
dict
- validate(model: BaseModel | BaseMetaObject, datamodule: DataModule) dict[source]#
Evaluate
modelon the validation split ofdatamodule.- Parameters:
model (BaseModel or BaseMetaObject) – Evaluated with its current weights, which are not touched.
datamodule (DataModule) – Set up here if it is not already. Must hold a validation split.
- Returns:
val_lossandval_accunderloss_fn, as floats.- Return type:
dict