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