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:
Examples
>>> from pyqit.utils.diagnostic import check_barren_plateau >>> result = check_barren_plateau(model, dm, num_samples=100) >>> print(result)