Source code for arcade.data.processors.transition

"""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