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: _PyQitObject

Orchestrates 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 groups model.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) – 0 silent, 1 progress only, 2 progress plus the model table.

  • seed (int, optional) – Seeded before anything stochastic runs. Weights are drawn at model construction, so reproducing them needs pyqit.set_seed before 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 through callbacks.

  • 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_params under 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.Trainer on the torch backend. Rejected on pennylane, which has no underlying trainer.

fit(model: BaseModel | BaseMetaObject, datamodule: DataModule) → TrainingHistory | dict[source]#

Train model on datamodule.

Parameters:
  • model (BaseModel or BaseMetaObject) – Trained in place. A QuantumPipeline trains 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 TrainingHistory per 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 model over the most specific split datamodule holds.

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 a pennylane.numpy.tensor rather than a bare ndarray, 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 model on the test split of datamodule.

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_loss and test_acc under loss_fn, as floats.

Return type:

dict

validate(model: BaseModel | BaseMetaObject, datamodule: DataModule) → dict[source]#

Evaluate model on the validation split of datamodule.

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_loss and val_acc under loss_fn, as floats.

Return type:

dict