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()