"""Transition detection results and validation."""
from __future__ import annotations
from dataclasses import dataclass, field
from typing import Any
import numpy as np
from ..types import LevelMismatchError, TraceLike
[docs]
@dataclass
class TransitionResult:
"""Per-qubit separation of clean, relaxation, excitation, and leakage traces."""
num_qubits: int
num_levels: int
per_qubit: dict[int, dict[str, Any]] = field(default_factory=dict)
[docs]
def clean_traces(self, qubit: int, state: int) -> np.ndarray:
return self.per_qubit[qubit][f"traces_{state}"]
[docs]
def centroids(self, qubit: int) -> dict[int, tuple[float, float]]:
return self.per_qubit[qubit]["centroids"]
def _qubit_state(labels: np.ndarray, qubit: int, num_levels: int) -> np.ndarray:
return (labels // (num_levels**qubit)) % num_levels
def find_centroids(
traces: TraceLike,
labels: np.ndarray,
num_qubits: int,
num_levels: int = 2,
) -> dict[int, dict[int, tuple[float, float]]]:
"""Compute IQ centroids for each qubit state."""
labels = np.asarray(labels)
result: dict[int, dict[int, tuple[float, float]]] = {}
for q in range(num_qubits):
result[q] = {}
qs = _qubit_state(labels, q, num_levels)
tq = traces[q] if isinstance(traces, dict) else np.asarray(traces, dtype=np.float64)
mean_iq = tq.mean(axis=1)
for state in range(num_levels):
mask = qs == state
if mask.any():
centroid = mean_iq[mask].mean(axis=0)
result[q][state] = (float(centroid[0]), float(centroid[1]))
else:
result[q][state] = (0.0, 0.0)
return result
def average_traces(
traces: TraceLike,
labels: np.ndarray,
num_qubits: int,
num_levels: int = 2,
) -> dict[int, dict[int, np.ndarray]]:
"""Compute per-state averaged traces for each qubit."""
labels = np.asarray(labels)
result: dict[int, dict[int, np.ndarray]] = {}
for q in range(num_qubits):
result[q] = {}
qs = _qubit_state(labels, q, num_levels)
tq = traces[q] if isinstance(traces, dict) else np.asarray(traces, dtype=np.float64)
for state in range(num_levels):
mask = qs == state
if mask.any():
result[q][state] = tq[mask].mean(axis=0)
else:
result[q][state] = np.zeros(tq.shape[1:])
return result
def validate_transition_result(
result: TransitionResult,
config: dict[str, Any] | Any,
) -> None:
"""Raise :class:`LevelMismatchError` when centroid count mismatches config."""
if hasattr(config, "model_dump"):
cfg = config.model_dump()
elif hasattr(config, "num_levels"):
cfg = {
"num_levels": getattr(config, "num_levels", result.num_levels),
"leakage": getattr(config, "leakage", False),
}
else:
cfg = dict(config)
expected_levels = int(cfg.get("num_levels", result.num_levels))
if cfg.get("leakage", False):
expected_levels = max(expected_levels, 3)
for q in range(result.num_qubits):
info = result.per_qubit.get(q, {})
centroids = info.get("centroids", {})
found = len(centroids)
if found != expected_levels:
raise LevelMismatchError(
f"Qubit {q}: found {found} clusters, expected {expected_levels}. "
"Check num_levels, leakage flag, or provide a custom detector."
)