Source code for arcade.data.processors.splitter

"""Stratified train / validation / test splitting."""

from __future__ import annotations

from dataclasses import dataclass

import numpy as np

from ..._logging import get_logger

log = get_logger(__name__)


@dataclass
class SplitData:
    """Train / val / test arrays and reconstruction indices."""

    train_traces: np.ndarray | dict[int, np.ndarray]
    train_labels: np.ndarray
    val_traces: np.ndarray | dict[int, np.ndarray]
    val_labels: np.ndarray
    test_traces: np.ndarray | dict[int, np.ndarray]
    test_labels: np.ndarray
    train_indices: np.ndarray
    val_indices: np.ndarray
    test_indices: np.ndarray


def _subset_traces(
    traces: np.ndarray | dict[int, np.ndarray], idx: np.ndarray
) -> np.ndarray | dict[int, np.ndarray]:
    if isinstance(traces, dict):
        return {q: arr[idx] for q, arr in traces.items()}
    return traces[idx]


def split(
    traces: np.ndarray | dict[int, np.ndarray],
    labels: np.ndarray,
    *,
    mode: str = "ratio",
    train_ratio: float = 0.60,
    val_ratio: float = 0.15,
    test_ratio: float = 0.25,
    train_per_class: int | None = None,
    val_per_class: int | None = None,
    test_per_class: int | None = None,
    stratified: bool = True,
    seed: int = 42,
    shuffle: bool = True,
) -> SplitData:
    """Split traces into train / validation / test."""

    labels = np.asarray(labels)
    n = len(labels)
    if isinstance(traces, dict):
        lengths = {q: len(arr) for q, arr in traces.items()}
        if lengths and set(lengths.values()) != {n}:
            raise ValueError(f"All per-qubit traces must have {n} shots, got {lengths}")
    else:
        traces_arr = np.asarray(traces)
        if len(traces_arr) != n:
            raise ValueError(
                f"traces ({len(traces_arr)}) and labels ({n}) length mismatch",
            )

    rng = np.random.default_rng(seed)

    if mode == "per_class":
        if train_per_class is None:
            raise ValueError("train_per_class is required when mode='per_class'")
        strat_labels = labels if labels.ndim == 1 else labels[:, 0]
        train_idx, val_idx, test_idx = _per_class_split(
            strat_labels,
            train_per_class=train_per_class,
            val_per_class=val_per_class,
            test_per_class=test_per_class,
            rng=rng,
            shuffle=shuffle,
        )
    elif mode == "ratio":
        total = train_ratio + val_ratio + test_ratio
        if abs(total - 1.0) > 1e-6:
            raise ValueError(f"Ratios must sum to 1.0, got {total:.6f}")
        if stratified:
            strat_labels = labels if labels.ndim == 1 else labels[:, 0]
            train_idx, val_idx, test_idx = _stratified_split(
                strat_labels,
                train_ratio,
                val_ratio,
                rng,
                shuffle,
            )
        else:
            indices = np.arange(n)
            if shuffle:
                rng.shuffle(indices)
            n_train = int(n * train_ratio)
            n_val = int(n * val_ratio)
            train_idx = indices[:n_train]
            val_idx = indices[n_train : n_train + n_val]
            test_idx = indices[n_train + n_val :]
    else:
        raise ValueError(f"Unknown split mode: {mode!r}")

    log.info(
        "Split: train=%d, val=%d, test=%d (mode=%s stratified=%s)",
        len(train_idx),
        len(val_idx),
        len(test_idx),
        mode,
        stratified if mode == "ratio" else True,
    )

    return SplitData(
        train_traces=_subset_traces(traces, train_idx),
        train_labels=labels[train_idx],
        val_traces=_subset_traces(traces, val_idx),
        val_labels=labels[val_idx],
        test_traces=_subset_traces(traces, test_idx),
        test_labels=labels[test_idx],
        train_indices=train_idx,
        val_indices=val_idx,
        test_indices=test_idx,
    )


def _stratified_split(
    strat_labels: np.ndarray,
    train_ratio: float,
    val_ratio: float,
    rng: np.random.Generator,
    shuffle: bool,
) -> tuple[np.ndarray, np.ndarray, np.ndarray]:
    unique_classes = np.unique(strat_labels)
    train_parts: list[np.ndarray] = []
    val_parts: list[np.ndarray] = []
    test_parts: list[np.ndarray] = []

    for cls in unique_classes:
        cls_indices = np.where(strat_labels == cls)[0]
        if shuffle:
            rng.shuffle(cls_indices)

        n_cls = len(cls_indices)
        n_train = max(1, int(n_cls * train_ratio))
        n_val = max(1, int(n_cls * val_ratio)) if n_cls > 2 else 0

        train_parts.append(cls_indices[:n_train])
        val_parts.append(cls_indices[n_train : n_train + n_val])
        test_parts.append(cls_indices[n_train + n_val :])

    train_idx = np.concatenate(train_parts)
    val_idx = np.concatenate(val_parts)
    test_idx = np.concatenate(test_parts)

    if shuffle:
        rng.shuffle(train_idx)
        rng.shuffle(val_idx)
        rng.shuffle(test_idx)

    return train_idx, val_idx, test_idx


def _per_class_split(
    strat_labels: np.ndarray,
    *,
    train_per_class: int,
    val_per_class: int | None,
    test_per_class: int | None,
    rng: np.random.Generator,
    shuffle: bool,
) -> tuple[np.ndarray, np.ndarray, np.ndarray]:
    unique = np.unique(strat_labels)
    train_parts: list[np.ndarray] = []
    val_parts: list[np.ndarray] = []
    test_parts: list[np.ndarray] = []

    n_val = 0 if val_per_class is None else val_per_class

    for cls in unique:
        cls_idx = np.where(strat_labels == cls)[0]
        if shuffle:
            rng.shuffle(cls_idx)
        n_cls = len(cls_idx)
        if test_per_class is None:
            need = train_per_class + n_val
            if need > n_cls:
                raise ValueError(
                    f"Class {cls}: need {need} samples but only {n_cls} available",
                )
            train_parts.append(cls_idx[:train_per_class])
            val_parts.append(cls_idx[train_per_class : train_per_class + n_val])
            test_parts.append(cls_idx[train_per_class + n_val :])
        else:
            need = train_per_class + n_val + test_per_class
            if need > n_cls:
                raise ValueError(
                    f"Class {cls}: need {need} samples but only {n_cls} available",
                )
            train_parts.append(cls_idx[:train_per_class])
            val_parts.append(cls_idx[train_per_class : train_per_class + n_val])
            test_parts.append(
                cls_idx[train_per_class + n_val : train_per_class + n_val + test_per_class],
            )

    train_idx = np.concatenate(train_parts)
    val_idx = np.concatenate(val_parts) if val_parts else np.array([], dtype=np.intp)
    test_idx = np.concatenate(test_parts)

    if shuffle:
        rng.shuffle(train_idx)
        rng.shuffle(val_idx)
        rng.shuffle(test_idx)

    return train_idx, val_idx, test_idx


def split_from_indices(
    traces: np.ndarray | dict[int, np.ndarray],
    labels: np.ndarray,
    train_indices: np.ndarray,
    val_indices: np.ndarray,
    test_indices: np.ndarray,
) -> SplitData:
    """Partition *traces* using precomputed index arrays."""
    labels = np.asarray(labels)
    return SplitData(
        train_traces=_subset_traces(traces, train_indices),
        train_labels=labels[train_indices],
        val_traces=_subset_traces(traces, val_indices),
        val_labels=labels[val_indices],
        test_traces=_subset_traces(traces, test_indices),
        test_labels=labels[test_indices],
        train_indices=train_indices,
        val_indices=val_indices,
        test_indices=test_indices,
    )

from typing import Any

from ..base import BaseDataProcessor, register_processor
from ..bundle import DataBundle


[docs] @register_processor("split") class SplitProcessor(BaseDataProcessor): """Stratified train / validation / test splitting (processor wrapper)."""
[docs] def process(self, bundle: DataBundle, **kwargs: Any) -> DataBundle: if bundle.labels is None: raise ValueError("SplitProcessor requires bundle.labels") traces = bundle.demod_traces if bundle.demod_traces is not None else bundle.traces sc = kwargs.get("split", kwargs) split_data = split( traces, bundle.labels, mode=sc.get("mode", "ratio"), train_ratio=sc.get("train_ratio", 0.60), val_ratio=sc.get("val_ratio", 0.15), test_ratio=sc.get("test_ratio", 0.25), train_per_class=sc.get("train_per_class"), val_per_class=sc.get("val_per_class"), test_per_class=sc.get("test_per_class"), stratified=sc.get("stratified", True), seed=sc.get("seed", 42), shuffle=sc.get("shuffle", True), ) bundle.splits = { "train_traces": split_data.train_traces, "train_labels": split_data.train_labels, "val_traces": split_data.val_traces, "val_labels": split_data.val_labels, "test_traces": split_data.test_traces, "test_labels": split_data.test_labels, "train_indices": split_data.train_indices, "val_indices": split_data.val_indices, "test_indices": split_data.test_indices, } return bundle