Source code for pyqit.utils.diagnostic

from contextlib import ExitStack, contextmanager
from dataclasses import dataclass
import logging

import numpy as np
from skbase.utils.dependencies import _check_soft_dependencies

from pyqit.core.config import get_backend
from pyqit.core.losses import get_loss_fn
from pyqit.utils.utils import _qnode_of

logger = logging.getLogger("pyqit.diagnostics")

_DEAD_GRADIENT = 1e-12


[docs] @dataclass class BPResult: """Result of `check_barren_plateau`. `repr(result)` or `print(result)` renders a table (rich if installed, ASCII otherwise). Attributes ---------- n_qubits : int Width of the circuit that sets the baseline. For a pipeline, that is its one trainable quantum stage. n_samples : int Random weight draws the gradients were sampled at. layer_variances : dict of {str: float} Gradient variance per weight tensor, keyed by its name in `model.weights`, such as `"main_circuit.weights"`. Circuit tensors count live parameters only. Classical tensors are included. layer_ratios : dict of {str: float} Each entry of `layer_variances` divided by `expected_variance`. The table marks a ratio below 1 as a plateau. overall_variance : float The same value as `quantum_variance`. expected_variance : float Theoretical floor from McClean et al., scaled by `bp_scale_factor`. For a local cost it is about twice the gradient variance of a random circuit, so a flagged model is within a factor of two of random. is_barren : bool True when `quantum_variance` is below `expected_variance`. quantum_variance : float Mean of the circuit tensors' variances. This is the number the verdict is made on. NaN when the model has no circuit weights. classical_variance : float, optional Mean of the classical tensors' variances. `None` when the model has no classical weights. n_parameters : int, optional Circuit parameters sampled, dead ones included. n_dead_parameters : int, optional Circuit parameters whose gradient was zero in every sample because they cannot reach the measured wires. They are left out of the variance. n_executions : int, optional Circuit executions the sampling cost, as counted by the device: one per sample under backprop, one plus two per parameter under parameter-shift. `None` when the model runs no QNode. """ n_qubits: int n_samples: int layer_variances: dict[str, float] layer_ratios: dict[str, float] overall_variance: float expected_variance: float is_barren: bool quantum_variance: float classical_variance: float | None = None n_executions: int | None = None n_parameters: int | None = None n_dead_parameters: int | None = None def __repr__(self) -> str: has_rich = _check_soft_dependencies("rich", severity="none") if has_rich: from rich.console import Console console = Console(force_terminal=False) with console.capture() as capture: console.print(self._build_rich_table()) return "\n" + capture.get().rstrip() else: return self._build_ascii_table() def __rich__(self): return self._build_rich_table() def _dead_text(self) -> str: return f"{self.n_dead_parameters} of {self.n_parameters}" def _build_rich_table(self): from rich.table import Table status_color = "red" if self.is_barren else "green" status_text = "BARREN PLATEAU" if self.is_barren else "HEALTHY" table = Table( title=f"BP Diagnostic Result : [{status_color} bold]{status_text}[/]", show_header=True, header_style="bold cyan", ) table.add_column("Metric / Layer", style="bold") table.add_column("Value", justify="right") table.add_column("Status", justify="right") table.add_row("Qubits", str(self.n_qubits), "") table.add_row("Samples", str(self.n_samples), "") if self.n_executions is not None: table.add_row("Circuit Executions", str(self.n_executions), "") if self.n_dead_parameters: table.add_row("Dead Parameters", self._dead_text(), "[dim]Excluded[/]") table.add_row( "Expected Variance", f"{self.expected_variance:.2e}", "[dim]Baseline[/]" ) q_color = f"[{status_color}]" table.add_row( "Quantum Variance", f"{q_color}{self.quantum_variance:.2e}[/]", f"[{status_color} bold]{status_text}[/]", ) if self.classical_variance is not None: table.add_row( "Classical Variance", f"{self.classical_variance:.2e}", "[dim]N/A[/]" ) table.add_section() for k, r in self.layer_ratios.items(): is_plateau = r < 1.0 r_color = "[red]" if is_plateau else "[green]" tag = "[red bold]← plateau[/]" if is_plateau else "[green]Healthy[/]" table.add_row(f"Layer: {k}", f"{r_color}{r:.3f}x[/]", tag) return table def _build_ascii_table(self) -> str: status = "BARREN PLATEAU" if self.is_barren else "HEALTHY" lines = [ "", "=" * 55, f" BP Diagnostic Result : {status}".center(55), "=" * 55, f"{'Qubits':<25} : {self.n_qubits}", f"{'Samples':<25} : {self.n_samples}", f"{'Expected Variance':<25} : {self.expected_variance:.2e} (Baseline)", f"{'Quantum Variance':<25} : {self.quantum_variance:.2e}", ] if self.n_executions is not None: lines.insert(6, f"{'Circuit Executions':<25} : {self.n_executions}") if self.n_dead_parameters: lines.append(f"{'Dead Parameters':<25} : {self._dead_text()} (Excluded)") if self.classical_variance is not None: lines.append(f"{'Classical Variance':<25} : {self.classical_variance:.2e}") lines.append("-" * 55) lines.append(f"{'Layer':<30} | {'Ratio':<8} | {'Status'}") lines.append("-" * 55) for k, r in self.layer_ratios.items(): tag = "← plateau" if r < 1.0 else "Healthy" layer_name = (k[:27] + "...") if len(k) > 30 else k lines.append(f"{layer_name:<30} | {r:<7.3f}x | {tag}") lines.append("=" * 55) return "\n".join(lines)
[docs] def check_barren_plateau( model, datamodule_or_X, y=None, num_samples: int = 200, loss_name: str = "mse", plot: bool = True, ) -> BPResult: """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. Returns ------- BPResult Examples -------- >>> from pyqit.utils.diagnostic import check_barren_plateau >>> result = check_barren_plateau(model, dm, num_samples=100) # doctest: +SKIP >>> print(result) # doctest: +SKIP """ backend = getattr(model, "backend", get_backend()) X, y_target = _resolve_input(datamodule_or_X, y, model) X, y_target = X[:1], y_target[:1] quantum_keys, classical_keys = _split_weight_keys(model) all_keys = quantum_keys + classical_keys if not all_keys: raise ValueError("Model has no tracked weights to calculate gradients for.") circuit = _circuit_of(model) n_qubits = getattr(circuit, "n_qubits", 1) or 1 measured_wires = getattr(circuit, "_measure_wires", range(n_qubits)) is_local_cost = len(measured_wires) < n_qubits if is_local_cost: raw_baseline = 1.0 / (2**n_qubits) else: raw_baseline = 1.0 / (3.0 * (4 ** (n_qubits - 1))) scale_factor = circuit.get_tag("bp_scale_factor", 1.0, raise_error=False) expected_variance = raw_baseline * scale_factor logger.info( f"Running Barren Plateau Diagnostic | Model: {type(model).__name__} | " f"Backend: {backend} | Samples: {num_samples}" ) with _track_executions(model) as executions: if backend == "torch": layer_gradients = _sample_gradients_torch( model, X, y_target, all_keys, quantum_keys, num_samples, loss_name ) else: layer_gradients = _sample_gradients_pennylane( model, X, y_target, all_keys, quantum_keys, num_samples, loss_name ) layer_variances, n_parameters, n_dead = {}, 0, 0 for k, grads in layer_gradients.items(): per_param = np.asarray(grads).reshape(num_samples, -1) if k in quantum_keys: live = np.abs(per_param).max(axis=0) > _DEAD_GRADIENT n_parameters += live.size n_dead += int(live.size - live.sum()) per_param = per_param[:, live] layer_variances[k] = float(np.var(per_param)) if per_param.size else 0.0 layer_ratios = {k: v / expected_variance for k, v in layer_variances.items()} quantum_vars = [layer_variances[k] for k in quantum_keys if k in layer_variances] classical_vars = [ layer_variances[k] for k in classical_keys if k in layer_variances ] quantum_variance = float(np.mean(quantum_vars)) if quantum_vars else float("nan") classical_variance = float(np.mean(classical_vars)) if classical_vars else None is_barren = quantum_variance < expected_variance if is_barren: logger.warning( f"Severe Barren Plateau detected!" f"Quantum gradient variance ({quantum_variance:.2e}) " f"is below the theoretical random-circuit baseline" f" ({expected_variance:.2e})." ) else: ratio = quantum_variance / expected_variance logger.info( f"Model looks healthy. Variance is {ratio:.1f}x above the random baseline." ) result = BPResult( n_qubits=n_qubits, n_samples=num_samples, layer_variances=layer_variances, layer_ratios=layer_ratios, overall_variance=quantum_variance, expected_variance=expected_variance, is_barren=is_barren, quantum_variance=quantum_variance, classical_variance=classical_variance, n_executions=executions.get("executions"), n_parameters=n_parameters, n_dead_parameters=n_dead, ) if plot: _plot(layer_gradients, quantum_keys, result) return result
@contextmanager def _track_executions(model): """Count circuit executions on every device the model's QNodes run on.""" import pennylane as qml nodes = getattr(model, "_qnodes", {}).values() qnodes = (_qnode_of(node) for node in nodes) devices = {qnode.device for qnode in qnodes if qnode is not None} totals = {} with ExitStack() as stack: trackers = [stack.enter_context(qml.Tracker(dev)) for dev in devices] yield totals if trackers: totals["executions"] = sum(t.totals.get("executions", 0) for t in trackers) def _sample_gradients_torch( model, X, y, weight_keys, random_keys, num_samples, loss_name ): """Mutate the ``nn.Parameter`` objects in ``random_keys`` in place, then restore. Keys in ``weight_keys`` but not ``random_keys`` (classical layers) keep their current values; their gradients are still collected. """ import torch loss_fn = get_loss_fn(loss_name, backend="torch") X_t = torch.as_tensor(np.asarray(X), dtype=torch.float64) y_t = torch.as_tensor(np.asarray(y), dtype=torch.float64) original_state = { k: v.detach().clone() for k, v in model.weights.items() if k in random_keys } layer_gradients = {k: [] for k in weight_keys} try: for _ in range(num_samples): with torch.no_grad(): for k, param in model.weights.items(): if k in random_keys: param.copy_(torch.empty_like(param).uniform_(0, 2 * np.pi)) for k in weight_keys: model.weights[k].grad = None preds = model.forward(X_t) if preds.ndim == 0: preds = preds.unsqueeze(0) loss = loss_fn(preds, y_t) loss.backward() for k in weight_keys: grad = model.weights[k].grad if grad is not None: layer_gradients[k].extend( grad.detach().cpu().numpy().flatten().tolist() ) finally: with torch.no_grad(): for k, param in model.weights.items(): if k in original_state: param.copy_(original_state[k]) return layer_gradients def _sample_gradients_pennylane( model, X, y, weight_keys, random_keys, num_samples, loss_name ): import pennylane as qml import pennylane.numpy as pnp loss_fn = get_loss_fn(loss_name, backend="pennylane") X_p = pnp.array(X, requires_grad=False) y_p = pnp.array(y, requires_grad=False) layer_gradients = {k: [] for k in weight_keys} # Use **flat_kwargs routing to bypass state mutation def cost(*weight_tensors): flat_kwargs = dict(zip(weight_keys, weight_tensors)) preds = model.forward(X_p, **flat_kwargs) if preds.ndim == 0: preds = pnp.expand_dims(preds, axis=0) return loss_fn(preds, y_p) grad_fn = qml.grad(cost) for _ in range(num_samples): weights = [ pnp.random.uniform( 0, 2 * np.pi, size=model.weights[k].shape, requires_grad=True ) if k in random_keys else pnp.array(model.weights[k], requires_grad=True) for k in weight_keys ] grads = grad_fn(*weights) if not isinstance(grads, tuple): grads = (grads,) for k, g in zip(weight_keys, grads): layer_gradients[k].extend(np.asarray(g).flatten().tolist()) return layer_gradients def _resolve_input(datamodule_or_X, y, model): from pyqit.data.datamodule import DataModule if isinstance(datamodule_or_X, DataModule): if not datamodule_or_X._is_setup: raise ValueError( "DataModule is not set up. You must call `dm.setup(stage='fit')` " "with the correct encoder before passing it to the diagnostic tool." ) X_data = ( datamodule_or_X.X_val if datamodule_or_X.X_val is not None else datamodule_or_X.X_train ) y_data = ( datamodule_or_X.y_val if datamodule_or_X.y_val is not None else datamodule_or_X.y_train ) return np.asarray(X_data[:32]), np.asarray(y_data[:32]) if y is None: raise ValueError("y must be provided when datamodule_or_X is a raw array.") return np.asarray(datamodule_or_X)[:32], np.asarray(y)[:32] def _split_weight_keys(model): """Flat weight keys split into those of QNodes and those of classical layers.""" groups = model.weight_groups() return groups.get("quantum", []), groups.get("classical", []) def _circuit_of(model): """The model whose circuit sets the baseline. ``model`` itself, or for a pipeline its one trainable quantum stage. Two such stages would need two floors, so that is rejected. """ steps = getattr(model, "steps", None) if steps is None: return model quantum = [ (name, s.model) for name, s in steps if s.trainable and _split_weight_keys(s.model)[0] ] if len(quantum) > 1: names = [name for name, _ in quantum] raise ValueError( "check_barren_plateau compares gradients against one circuit's " f"baseline, but this pipeline trains {len(quantum)} quantum stages " f"{names}. Run it on each stage's model instead." ) return _circuit_of(quantum[0][1]) if quantum else model def _plot(layer_gradients, quantum_keys, result: BPResult): if not _check_soft_dependencies("matplotlib", severity="none"): logger.warning( "Matplotlib is not installed. Skipping barren plateau histogram plot." ) return else: import matplotlib.pyplot as plt n_layers = len(layer_gradients) fig, axes = plt.subplots(n_layers, 1, figsize=(9, 3 * n_layers), squeeze=False) fig.suptitle( f"Gradient Landscape — {result.n_qubits} Qubits\n" f"Expected variance floor: {result.expected_variance:.2e}", fontsize=13, y=1.02, ) for ax, (key, grads) in zip(axes[:, 0], layer_gradients.items()): ratio = result.layer_ratios[key] is_bp = (key in quantum_keys) and (ratio < 1.0) color = "#E74C3C" if is_bp else "#00CBA9" label = "quantum" if key in quantum_keys else "classical" ax.hist( grads, bins=min(50, max(10, len(grads) // 10)), color=color, alpha=0.75, edgecolor="black", ) ax.axvline(0, color="black", linestyle="--", alpha=0.6) status = ( "PLATEAU" if is_bp else ("Healthy" if key in quantum_keys else "Classical") ) ax.set_title( f"{key} [{label}] — var={result.layer_variances[key]:.2e} \ ({ratio:.2f}x) [{status}]", fontsize=10, ) ax.grid(True, alpha=0.25) plt.tight_layout() plt.show()