"""Transition detection processor (relaxation, excitation, leakage)."""
from __future__ import annotations
from dataclasses import dataclass, field
from typing import Any, Union
import numpy as np
from ..base import BaseDataProcessor, register_processor
from ..bundle import DataBundle
TraceLike = Union[np.ndarray, dict[int, np.ndarray]]
@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)
def clean_traces(self, qubit: int, state: int) -> np.ndarray:
return self.per_qubit[qubit][f"traces_{state}"]
def centroids(self, qubit: int) -> dict[int, tuple[float, float]]:
return self.per_qubit[qubit]["centroids"]
def _euclidean(i: np.ndarray, q: np.ndarray, cx: float, cy: float) -> np.ndarray:
return np.sqrt((i - cx) ** 2 + (q - cy) ** 2)
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,
) -> dict[int, dict[int, tuple[float, float]]]:
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 _transitions_centroid(
traces: TraceLike,
labels: np.ndarray,
num_qubits: int,
num_levels: int,
decode_levels: int,
radius_scale: float,
) -> dict[int, dict[str, Any]]:
"""Generalized centroid-radius separation for n-level systems."""
centroids = _find_centroids(traces, labels, num_qubits, decode_levels)
result: dict[int, dict[str, Any]] = {}
for q in range(num_qubits):
qs = _qubit_state(labels, q, decode_levels)
tq = np.asarray(traces[q], dtype=np.float64) if isinstance(traces, dict) else traces
cx = {s: centroids[q][s][0] for s in range(num_levels)}
cy = {s: centroids[q][s][1] for s in range(num_levels)}
pairwise_dist: dict[tuple[int, int], float] = {}
for s1 in range(num_levels):
for s2 in range(s1 + 1, num_levels):
d = np.sqrt((cx[s1] - cx[s2]) ** 2 + (cy[s1] - cy[s2]) ** 2)
pairwise_dist[(s1, s2)] = d / 2.0
pairwise_dist[(s2, s1)] = d / 2.0
clean_radius: dict[int, float] = {}
for s in range(num_levels):
neighbors = [pairwise_dist[(s, s2)] for s2 in range(num_levels) if s2 != s]
clean_radius[s] = radius_scale * min(neighbors) if neighbors else 0.0
state_traces: dict[int, np.ndarray] = {}
mean_iq: dict[int, np.ndarray] = {}
for s in range(num_levels):
state_traces[s] = tq[qs == s]
mean_iq[s] = state_traces[s].mean(axis=1) if len(state_traces[s]) > 0 else np.empty((0, 2))
info: dict[str, Any] = {}
incorrect: dict[int, np.ndarray] = {}
inc_iq: dict[int, np.ndarray] = {}
for s in range(num_levels):
if len(mean_iq[s]) == 0:
info[f"traces_{s}"] = np.empty((0, *tq.shape[1:]))
incorrect[s] = np.empty((0, *tq.shape[1:]))
inc_iq[s] = np.empty((0, 2))
continue
dist = _euclidean(mean_iq[s][:, 0], mean_iq[s][:, 1], cx[s], cy[s])
clean_mask = dist < clean_radius[s]
info[f"traces_{s}"] = state_traces[s][clean_mask]
incorrect[s] = state_traces[s][~clean_mask]
inc_iq[s] = incorrect[s].mean(axis=1) if len(incorrect[s]) > 0 else np.empty((0, 2))
for s in range(num_levels):
if len(incorrect[s]) == 0:
continue
for t in range(num_levels):
if t == s:
continue
pair_r = pairwise_dist.get((s, t), clean_radius[s])
d = _euclidean(inc_iq[s][:, 0], inc_iq[s][:, 1], cx[t], cy[t])
trans = incorrect[s][d < radius_scale * pair_r]
key = f"relax_{s}_{t}" if s > t else f"excite_{s}_{t}"
info[key] = trans
info["centroids"] = {s: (cx[s], cy[s]) for s in range(num_levels)}
result[q] = info
return result
def _detect_spectral(
traces: TraceLike,
labels: np.ndarray,
num_qubits: int,
num_levels: int,
decode_levels: int,
n_clusters: int,
radius_scale: float,
) -> dict[int, dict[str, Any]]:
"""Spectral-clustering based transition identification."""
from sklearn.cluster import SpectralClustering
centroids_initial = _find_centroids(traces, labels, num_qubits, decode_levels)
result: dict[int, dict[str, Any]] = {}
for q in range(num_qubits):
qs = _qubit_state(labels, q, decode_levels)
tq = np.asarray(traces[q], dtype=np.float64) if isinstance(traces, dict) else traces
mtv = tq.mean(axis=1)
state1_mask = qs == 1
state1_mean_iq = mtv[state1_mask]
if len(state1_mean_iq) < n_clusters:
result[q] = _transitions_centroid(
traces, labels, num_qubits, num_levels, decode_levels, radius_scale
)[q]
continue
sc = SpectralClustering(n_clusters=n_clusters, random_state=42)
cluster_labels = sc.fit_predict(state1_mean_iq)
c0_i, c0_q = centroids_initial[q][0]
cluster_sizes: dict[int, int] = {}
cluster_dist_to_0: dict[int, float] = {}
for c in range(n_clusters):
members = state1_mean_iq[cluster_labels == c]
cluster_sizes[c] = len(members)
if len(members) > 0:
cm = members.mean(axis=0)
cluster_dist_to_0[c] = float(np.sqrt((cm[0] - c0_i) ** 2 + (cm[1] - c0_q) ** 2))
else:
cluster_dist_to_0[c] = float("inf")
true_1_cluster = max(cluster_sizes, key=cluster_sizes.get) # type: ignore[arg-type]
remaining = [c for c in range(n_clusters) if c != true_1_cluster]
relax_cluster = min(remaining, key=lambda c: cluster_dist_to_0[c])
leak_clusters = [c for c in remaining if c != relax_cluster]
cx: dict[int, float] = {}
cy: dict[int, float] = {}
cx[0], cy[0] = centroids_initial[q][0]
t1_members = state1_mean_iq[cluster_labels == true_1_cluster]
cx[1] = float(t1_members[:, 0].mean())
cy[1] = float(t1_members[:, 1].mean())
if leak_clusters:
all_leak = state1_mean_iq[np.isin(cluster_labels, leak_clusters)]
if len(all_leak) > 0:
cx[2] = float(all_leak[:, 0].mean())
cy[2] = float(all_leak[:, 1].mean())
else:
cx[2], cy[2] = 0.0, 0.0
else:
cx[2], cy[2] = 0.0, 0.0
effective_levels = max(num_levels, 3)
new_qs = qs.copy()
state1_indices = np.where(state1_mask)[0]
for idx, cl in zip(state1_indices, cluster_labels):
if cl in leak_clusters:
new_qs[idx] = 2
pairwise_dist: dict[tuple[int, int], float] = {}
for s1 in range(effective_levels):
for s2 in range(s1 + 1, effective_levels):
d = np.sqrt((cx[s1] - cx[s2]) ** 2 + (cy[s1] - cy[s2]) ** 2) / 2.0
pairwise_dist[(s1, s2)] = d
pairwise_dist[(s2, s1)] = d
clean_radius: dict[int, float] = {}
for s in range(effective_levels):
neighbors = [
pairwise_dist[(s, s2)]
for s2 in range(effective_levels)
if s2 != s and pairwise_dist.get((s, s2), 0) > 0
]
clean_radius[s] = radius_scale * min(neighbors) if neighbors else 0.0
state_traces_map = {s: tq[new_qs == s] for s in range(effective_levels)}
info: dict[str, Any] = {}
for s in range(effective_levels):
if len(state_traces_map[s]) == 0:
info[f"traces_{s}"] = np.empty((0, *tq.shape[1:]))
continue
m = state_traces_map[s].mean(axis=1)
dist = _euclidean(m[:, 0], m[:, 1], cx[s], cy[s])
info[f"traces_{s}"] = state_traces_map[s][dist < clean_radius[s]]
for s in range(effective_levels):
if len(state_traces_map[s]) == 0:
continue
m = state_traces_map[s].mean(axis=1)
dist = _euclidean(m[:, 0], m[:, 1], cx[s], cy[s])
incorrect_arr = state_traces_map[s][dist >= clean_radius[s]]
if len(incorrect_arr) == 0:
continue
inc_iq = incorrect_arr.mean(axis=1)
for t in range(effective_levels):
if t == s:
continue
pair_r = pairwise_dist.get((s, t), clean_radius.get(s, 0.0))
d = _euclidean(inc_iq[:, 0], inc_iq[:, 1], cx[t], cy[t])
trans = incorrect_arr[d < radius_scale * pair_r]
key = f"relax_{s}_{t}" if s > t else f"excite_{s}_{t}"
info[key] = trans
info["centroids"] = {s: (cx.get(s, 0.0), cy.get(s, 0.0)) for s in range(effective_levels)}
result[q] = info
return result
def detect_transitions(
traces: TraceLike,
labels: np.ndarray,
num_qubits: int,
num_levels: int = 2,
*,
method: str = "centroid",
include_leakage: bool = False,
radius_scale: float = 1.0,
n_clusters: int = 3,
label_num_levels: int | None = None,
) -> TransitionResult:
"""Unified transition detection: relaxation + optional leakage.
Args:
method: ``"centroid"`` or ``"spectral"``.
include_leakage: When True and method is ``"spectral"``, runs
spectral clustering on the |1> population to find leakage states.
"""
labels = np.asarray(labels)
decode_levels = label_num_levels if label_num_levels is not None else num_levels
if include_leakage and method == "spectral":
effective_levels = max(num_levels, 3)
per_qubit = _detect_spectral(
traces, labels, num_qubits, effective_levels, decode_levels,
n_clusters=n_clusters, radius_scale=radius_scale,
)
else:
per_qubit = _transitions_centroid(
traces, labels, num_qubits, num_levels, decode_levels,
radius_scale=radius_scale,
)
return TransitionResult(
num_qubits=num_qubits,
num_levels=num_levels,
per_qubit=per_qubit,
)
[docs]
@register_processor("transitions")
class TransitionProcessor(BaseDataProcessor):
"""Detect transition trajectories (relaxation, excitation, leakage)."""
def __init__(
self,
method: str = "centroid",
spectral: bool = False,
include_leakage: bool = False,
) -> None:
self.method = method
self.spectral = spectral
self.include_leakage = include_leakage
[docs]
def process(self, bundle: DataBundle, **kwargs: Any) -> DataBundle:
if bundle.labels is None:
raise ValueError("TransitionProcessor requires bundle.labels")
traces = bundle.demod_traces if bundle.demod_traces is not None else bundle.traces
num_qubits = kwargs.get("num_qubits", bundle.num_qubits or 1)
num_levels = kwargs.get("num_levels", bundle.metadata.get("num_levels", 2))
method = "spectral" if self.spectral else self.method
bundle.transitions = detect_transitions(
traces,
bundle.labels,
num_qubits,
num_levels,
method=method,
include_leakage=kwargs.get("include_leakage", self.include_leakage),
radius_scale=kwargs.get("radius_scale", 1.0),
n_clusters=kwargs.get("n_clusters", 3),
)
return bundle