pyqit.utils.diagnostic.check_barren_plateau#

pyqit.utils.diagnostic.check_barren_plateau(model, datamodule_or_X, y=None, num_samples: int = 200, loss_name: str = 'mse', plot: bool = True) → BPResult[source]#

Monte-Carlo sample gradients at random weights and compare to baseline.

Runs at construction-time weights, so it never touches or requires fitting. Trainer(check_bp=True) runs this as a pre-flight check.

Parameters:
  • model (BaseModel)

  • datamodule_or_X (DataModule or array-like) – A set-up DataModule, or raw X, in which case y is required.

  • y (array-like, optional) – Required when datamodule_or_X is raw X.

  • num_samples (int, default 200) – Random weight draws to average the gradient variance over.

  • loss_name (str, default "mse") – A name from loss_registry().

  • plot (bool, default True) – Draw a gradient-variance histogram if matplotlib is installed.

Return type:

BPResult

Examples

>>> from pyqit.utils.diagnostic import check_barren_plateau
>>> result = check_barren_plateau(model, dm, num_samples=100)
>>> print(result)