Source code for arcade.data.transitions.result

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