Source code for arcade.stages.data

"""Data preparation stage for Readout 2019-style IQ datasets."""

from __future__ import annotations

from pathlib import Path
from typing import Any

import numpy as np

from .._logging import get_logger
from .base import BaseDataStage
from .helpers import (
    _compute_split_indices,
    resolve_cache_paths,
    trace_shape_descr,
    transition_plan,
)

log = get_logger(__name__)


[docs] class Readout2019DataStage(BaseDataStage): """Load, split (indices first), demodulate, preprocess, detect transitions."""
[docs] def run( self, cfg: Any, *, cache_dir: str | Path | None = None, demod_cache_path: str | Path | None = None, classifiers_run: list[str] | None = None, ) -> dict[str, Any]: """Load and prepare IQ traces for downstream classify / hardware stages. Args: cfg: Pipeline config (``data`` block drives loaders, demod, preprocessing, transitions, and split). cache_dir: Optional root for demod / spectral / transition caches. demod_cache_path: Optional explicit demod-cache file path. classifiers_run: Classifier names used when ``transitions.auto`` plans leakage vs filter-only detection. Returns: Data-stage dict with ``bundle``, ``traces``, ``labels``, ``label_set``, ``relaxation_result``, ``transition_result``, ``spectral_labels``, ``splits``, and flat train/val/test keys when splits are available. May include ``post_measured``. """ import arcade.data # noqa: F401 — register bundled processors import arcade.data.loaders # noqa: F401 — register loaders from arcade.data.base import get_processor from arcade.data.bundle import DataBundle from arcade.data.cache.demod import load_demod_cache, save_demod_cache from arcade.data.cache.spectral import ( load_spectral_cache, save_spectral_cache, spectral_cache_complete, ) from arcade.data.cache.transition import load_transition_cache, save_transition_cache from arcade.data.load import load_data from arcade.data.processors.demodulator import load_or_demodulate from arcade.data.processors.preprocessor import preprocess from arcade.data.processors.splitter import split_from_indices from arcade.data.transitions.centroid import detect_relaxations, detect_transitions from arcade.data.transitions.spectral import detect_leakage from arcade.utils.labels import LabelSet, decode_per_qubit, prepare_labels log.info("=== Stage 1: Data ===") dcfg = cfg.data dm = dcfg.demodulation tcfg = dcfg.transitions cache_paths = resolve_cache_paths( dcfg, cache_dir=cache_dir, demod_cache_path=demod_cache_path, ) demod_cache_file = cache_paths["demod"] spectral_cache_path = cache_paths["spectral"] transition_cache_path = cache_paths["transition"] bundle: DataBundle | None = None label_set: LabelSet | None = None labels: np.ndarray | None = None split_indices: tuple[np.ndarray, np.ndarray, np.ndarray] | None = None used_demod_cache = False traces: Any = None if demod_cache_file and not dm.cache_force_rebuild and demod_cache_file.is_file(): cached = load_demod_cache(demod_cache_file) traces = cached["demod_traces"] labels = cached["labels"] pq = decode_per_qubit( labels, dcfg.num_qubits, dcfg.num_levels, qubit_bit_order=dcfg.qubit_bit_order, ) label_set = LabelSet( labels, pq, dcfg.num_qubits, dcfg.num_levels, dcfg.qubit_bit_order, ) log.info( "Loaded demod cache from %s (%d shots, %d qubits)", demod_cache_file, len(labels), len(traces), ) used_demod_cache = True if not used_demod_cache: log.info("Loading data from %s (loader=%s)", dcfg.path, dcfg.loader) bundle = load_data(dcfg.path, dcfg.model_dump()) raw_labels = bundle.labels log.info( "Loaded: traces %s, labels %s", getattr(bundle.traces, "shape", type(bundle.traces)), raw_labels.shape if raw_labels is not None else None, ) traces, label_set = prepare_labels( bundle.traces, raw_labels, num_qubits=dcfg.num_qubits, num_levels=dcfg.num_levels, label_format=dcfg.label_format, qubit_bit_order=dcfg.qubit_bit_order, ) labels = label_set.joint log.info( "After label preparation (format=%s): unique labels=%d", dcfg.label_format, len(np.unique(labels)), ) if labels is not None: split_indices = _compute_split_indices(labels, dcfg.split) log.info( "Split indices computed on raw labels: train=%d, val=%d, test=%d", len(split_indices[0]), len(split_indices[1]), len(split_indices[2]), ) if dm.enabled: log.info("Demodulating raw ADC traces") traces = load_or_demodulate( traces, dm.frequencies, dm.sampling_rate, pre_demodulated=dm.pre_demodulated, num_qubits=dcfg.num_qubits, time_bin_width=dm.time_bin_width, skip_samples=dm.skip_samples, correct_iq_offset=dm.correct_iq_offset, correct_iq_amplitude=dm.correct_iq_amplitude, ) if demod_cache_file and (dm.cache_auto_save or dm.cache_force_rebuild): save_demod_cache( demod_cache_file, traces, labels, source=str(dcfg.path), frequencies=list(dm.frequencies), sampling_rate=float(dm.sampling_rate), ) elif labels is not None: split_indices = _compute_split_indices(labels, dcfg.split) log.info( "Split indices computed on cached labels: train=%d, val=%d, test=%d", len(split_indices[0]), len(split_indices[1]), len(split_indices[2]), ) pp = dcfg.preprocessing traces = preprocess( traces, boxcar_window=pp.boxcar_window, truncate_at=pp.truncate_at, normalize=pp.normalize, ) log.info("Preprocessed traces: %s", trace_shape_descr(traces)) if dcfg.custom_processor: log.info("Running custom processor: %s", dcfg.custom_processor) proc_cls = get_processor(dcfg.custom_processor) processor = proc_cls() bundle = DataBundle( traces=traces, labels=labels, time=bundle.time if bundle is not None else None, metadata=bundle.metadata if bundle is not None else {}, ) bundle = processor.process(bundle, dcfg.model_dump()) traces = bundle.traces labels = bundle.labels splits = None if labels is not None and split_indices is not None: train_idx, val_idx, test_idx = split_indices splits = split_from_indices(traces, labels, train_idx, val_idx, test_idx) log.info( "Split applied: train=%d, val=%d, test=%d", len(splits.train_labels), len(splits.val_labels), len(splits.test_labels), ) run_relax, run_leak, leak_method = transition_plan(tcfg, classifiers_run) relaxation_result = None transition_result = None spectral_labels_by_q: dict[int, np.ndarray] | None = None if splits is not None and (run_relax or run_leak): train_tr = splits.train_traces train_lbl = splits.train_labels levels_for_detection = ( tcfg.effective_levels if tcfg.effective_levels is not None else dcfg.num_levels ) base_kwargs: dict[str, Any] = { "num_qubits": dcfg.num_qubits, "num_levels": dcfg.num_levels, "method": tcfg.method, "radius_scale": tcfg.radius_scale, } if run_leak and leak_method == "spectral" and spectral_cache_path: if ( not tcfg.spectral_cache_force_rebuild and spectral_cache_complete(spectral_cache_path, dcfg.num_qubits) ): cached = load_spectral_cache(spectral_cache_path, dcfg.num_qubits) transition_result = cached["transition_result"] spectral_labels_by_q = cached["spectral_labels"] log.info("Loaded per-qubit spectral cache from %s", spectral_cache_path) if ( transition_result is None and run_leak and transition_cache_path and not tcfg.cache_force_rebuild ): transition_result = load_transition_cache( str(transition_cache_path), train_tr, train_lbl, fingerprint_extra={"n_clusters": tcfg.n_clusters, "method": leak_method}, **base_kwargs, ) if run_relax and relaxation_result is None: log.info("Detecting relaxations (centroid_radius) on training data") relaxation_result = detect_relaxations( train_tr, train_lbl, dcfg.num_qubits, dcfg.num_levels, label_num_levels=dcfg.num_levels, radius_scale=tcfg.radius_scale, ) if run_leak and transition_result is None: log.info("Detecting leakage (method=%s) on training data", leak_method) transition_result, spectral_labels_by_q = detect_leakage( train_tr, train_lbl, dcfg.num_qubits, levels_for_detection, label_num_levels=dcfg.num_levels, method=leak_method, n_clusters=tcfg.n_clusters, radius_scale=tcfg.radius_scale, ) if ( leak_method == "spectral" and spectral_cache_path and isinstance(train_tr, dict) ): save_spectral_cache( spectral_cache_path, train_tr, train_lbl, num_qubits=dcfg.num_qubits, num_levels=levels_for_detection, decode_levels=dcfg.num_levels, n_clusters=tcfg.n_clusters, radius_scale=tcfg.radius_scale, metadata={"split_seed": dcfg.split.seed}, ) if transition_cache_path: save_transition_cache( str(transition_cache_path), transition_result, train_tr, train_lbl, fingerprint_extra={"n_clusters": tcfg.n_clusters, "method": leak_method}, **base_kwargs, ) elif not run_leak and run_relax: transition_result = relaxation_result elif not tcfg.auto and tcfg.enabled and transition_result is None: log.info( "Detecting transitions (method=%s) on training data only", tcfg.method, ) det_kw: dict[str, Any] = { "method": tcfg.method, "radius_scale": tcfg.radius_scale, } if tcfg.method in ("spectral", "gmm"): det_kw["n_clusters"] = tcfg.n_clusters transition_result = detect_transitions( train_tr, train_lbl, dcfg.num_qubits, levels_for_detection, label_num_levels=dcfg.num_levels, **det_kw, ) relaxation_result = transition_result if transition_cache_path: save_transition_cache( str(transition_cache_path), transition_result, train_tr, train_lbl, fingerprint_extra={"n_clusters": tcfg.n_clusters}, **base_kwargs, ) out: dict[str, Any] = { "bundle": bundle, "traces": traces, "labels": labels, "label_set": label_set, "relaxation_result": relaxation_result, "transition_result": transition_result, "spectral_labels": spectral_labels_by_q, "splits": splits, } if splits is not None: out.update( { "train_traces": splits.train_traces, "train_labels": splits.train_labels, "val_traces": splits.val_traces, "val_labels": splits.val_labels, "test_traces": splits.test_traces, "test_labels": splits.test_labels, "train_indices": splits.train_indices, "val_indices": splits.val_indices, "test_indices": splits.test_indices, } ) post_full = None if bundle is not None: post_full = (bundle.metadata or {}).get("post_measured_labels") if post_full is not None and splits is not None: post_arr = np.asarray(post_full, dtype=np.float64) out["post_measured"] = { "train": post_arr[splits.train_indices], "val": post_arr[splits.val_indices], "test": post_arr[splits.test_indices], } return out
DefaultDataStage = Readout2019DataStage