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