Source code for arcade.stages.classify

"""Classifier training and evaluation stage."""

from __future__ import annotations

from typing import Any

import numpy as np

from .._logging import get_logger
from ..classifiers.base import (
    FILTER_SET_CLASSIFIERS,
    BaseClassifier,
    ensure_classifiers_registered,
    get_classifier,
    iq_trace_classifier_names,
    seq_trace_classifier_names,
)
from .base import BaseClassifyStage

log = get_logger(__name__)


def _ensure_classifiers_registered() -> None:
    ensure_classifiers_registered()


def _format_classifier_metrics(metrics: Any, num_qubits: int) -> str:
    """Format joint, F5Q, and per-qubit test accuracy for logging."""
    pq = getattr(metrics, "per_qubit_accuracy", {}) or {}
    f5q = getattr(metrics, "fidelity_f5q_geometric_mean", None)
    parts = [
        f"Q{q}={pq[q]:.4f}" if q in pq else f"Q{q}=—"
        for q in range(num_qubits)
    ]
    f5q_s = f"F5Q={f5q:.4f}" if f5q is not None else "F5Q=—"
    return f"joint={metrics.accuracy:.4f}  {f5q_s}  " + "  ".join(parts)


def _classifier_config(cfg: Any, clf_name: str, filter_set: Any) -> dict[str, Any]:
    """Build the per-classifier config dict from the global config."""
    from arcade.nn.utils import load_nn_config

    ccfg = cfg.classifiers
    dcfg = cfg.data
    clf_specific = getattr(ccfg, clf_name, None)
    if isinstance(clf_specific, dict):
        out = dict(clf_specific)
    elif clf_specific and hasattr(clf_specific, "model_dump"):
        out = clf_specific.model_dump()
    else:
        out = {}

    if clf_name in {"path_signature", "ngrc", "reservoir", "leakage", "multilevel"}:
        out.setdefault("num_qubits", dcfg.num_qubits)
        out.setdefault("num_levels", dcfg.num_levels)
        out.setdefault("qubit_bit_order", dcfg.qubit_bit_order)

    if clf_name in {"path_signature", "ngrc"}:
        out.setdefault("qubit_bit_order", dcfg.qubit_bit_order)

    if filter_set is not None and clf_name in FILTER_SET_CLASSIFIERS:
        out["filter_set"] = filter_set
        if clf_name == "threshold":
            out.setdefault("strategy", "per_qubit")
            out["threshold_fit_mode"] = ccfg.threshold_mode
            tp = ccfg.threshold_percentile
            if tp is not None:
                out["threshold_percentile"] = tp

    try:
        cls = get_classifier(clf_name)
        needs_nn = getattr(cls, "REQUIRES_NN_CONFIG", False)
    except KeyError:
        needs_nn = False

    if needs_nn or out.get("config_file"):
        nn_cfg_path = out.get("config_file")
        if nn_cfg_path:
            out["nn_config"] = load_nn_config(nn_cfg_path)
    return out


[docs] class DefaultClassifyStage(BaseClassifyStage): """Train and evaluate classifiers from config. Runs the feature stage for matched filters, then trains each classifier listed in ``cfg.classifiers.run`` (plus an optional custom classifier). """
[docs] def run(self, cfg: Any, data_out: dict[str, Any]) -> dict[str, Any]: """Train/evaluate all configured classifiers. Args: cfg: Pipeline config with ``data`` and ``classifiers`` blocks. data_out: Output of the data stage (``splits``, ``transition_result``, optional ``post_measured``). Returns: Dict with ``filter_set`` and ``classifiers`` (per-name results containing ``model``, ``metrics``, ``history``, ``sweep``, …). If ``splits`` is missing, returns empty classifier results. """ from ..features import kinds as feature_mod _ensure_classifiers_registered() log.info("=== Stage 2: Classify ===") dcfg = cfg.data ccfg = cfg.classifiers splits = data_out["splits"] transition_result = data_out["transition_result"] if splits is None: log.warning("No labeled data → skipping classify stage") return {"filter_set": None, "classifiers": {}} from .features import DefaultFeatureStage feature_out = DefaultFeatureStage().run(cfg, data_out) filter_set = feature_out["filter_set"] clf_results: dict[str, dict[str, Any]] = {} for clf_name in ccfg.run: log.info("--- Classifier: %s ---", clf_name) result = self.run_single_classifier( clf_name, cfg=cfg, filter_set=filter_set, splits=splits, transition_result=transition_result, feature_registry=feature_mod, post_measured=data_out.get("post_measured"), ) timing = self._benchmark_latency(result) result["timing"] = timing clf_results[clf_name] = result metrics = result.get("metrics") if metrics is not None: log.info(" %s%s", clf_name, _format_classifier_metrics(metrics, dcfg.num_qubits)) else: log.info(" %s → (no metrics)", clf_name) if ccfg.custom_classifier: log.info("--- Custom classifier: %s ---", ccfg.custom_classifier) result = self.run_single_classifier( ccfg.custom_classifier, cfg=cfg, filter_set=filter_set, splits=splits, transition_result=transition_result, feature_registry=feature_mod, ) result["timing"] = self._benchmark_latency(result) clf_results[ccfg.custom_classifier] = result return { "filter_set": filter_set, "classifiers": clf_results, }
@staticmethod def _benchmark_latency(result: dict[str, Any]) -> dict[str, float]: from arcade.training import benchmark_inference_latency timing = benchmark_inference_latency( result["model"], result.get("benchmark_features"), traces_sample=result.get("benchmark_traces"), ) if timing is not None: return { "inference_ms": timing.mean_ms, "mean_ms": timing.mean_ms, "total_ms": timing.mean_ms, "feature_ms": 0.0, "p95_ms": timing.p95_ms, } return {}
[docs] def run_single_classifier( self, clf_name: str, *, cfg: Any, filter_set: Any, splits: Any, transition_result: Any, feature_registry: Any, post_measured: dict[str, np.ndarray] | None = None, ) -> dict[str, Any]: """Train and evaluate one classifier, optionally with length sweep. Args: clf_name: Registered classifier name (e.g. ``threshold``, ``ngrc``). cfg: Pipeline config (training, sweep, tuning, data knobs). filter_set: Matched-filter bank from the feature stage (may be ``None`` for trace-only classifiers). splits: Train/val/test split object from the data stage. transition_result: Transition detection result used for length sweeps (may be ``None`` when sweep is disabled). feature_registry: Module exposing ``get_feature_fn`` (usually :mod:`arcade.features.kinds`). post_measured: Optional post-measurement labels keyed by split name (``train`` / ``val`` / ``test``), used by REMF. Returns: Dict with ``model``, ``metrics``, ``history``, ``sweep``, ``best_params``, ``benchmark_features``, and ``benchmark_traces``. """ from arcade.nn.builder import validate_nn_config from arcade.optimization import maybe_tune from arcade.training import ( evaluate, sweep_tail_zeroing, sweep_trace_length, train, training_kwargs, ) dcfg = cfg.data ccfg = cfg.classifiers clf_cls = get_classifier(clf_name) clf_config = _classifier_config(cfg, clf_name, filter_set) best_params = None if cfg.tuning.enabled and clf_name in cfg.tuning.search_space: train_features_probe = feature_registry.get_feature_fn( getattr(clf_cls, "FEATURE_KIND", "filter_features"), )(splits.train_traces, filter_set) val_features_probe = feature_registry.get_feature_fn( getattr(clf_cls, "FEATURE_KIND", "filter_features"), )(splits.val_traces, filter_set) best_params = maybe_tune( cfg, clf_name, clf_cls, train_features_probe, splits.train_labels, val_features_probe, splits.val_labels, filter_set, ) if best_params: clf_config = {**clf_config, **best_params} nn_cfg = clf_config.get("nn_config") if nn_cfg is not None: validate_nn_config(nn_cfg) clf = clf_cls.from_config(clf_config) feat_kind = "filter_features" if isinstance(clf, BaseClassifier): feat_kind = getattr(clf, "FEATURE_KIND", "filter_features") extract_fn = feature_registry.get_feature_fn(feat_kind) train_features = extract_fn(splits.train_traces, filter_set) val_features = extract_fn(splits.val_traces, filter_set) test_features = extract_fn(splits.test_traces, filter_set) is_threshold = clf_name == "threshold" history: dict[str, Any] = {} threshold_fit_kw: dict[str, Any] = {} if clf_name == "threshold": threshold_fit_kw["refit_thresholds"] = True threshold_fit_kw["threshold_fit_mode"] = ccfg.threshold_mode threshold_fit_kw["qubit_bit_order"] = dcfg.qubit_bit_order iq_trace_clfs = iq_trace_classifier_names() seq_trace_clfs = seq_trace_classifier_names() trace_fit_kw: dict[str, Any] = {} if clf_name in iq_trace_clfs: trace_fit_kw["train_traces"] = splits.train_traces trace_fit_kw["val_traces"] = splits.val_traces trace_fit_kw["num_levels"] = dcfg.num_levels trace_fit_kw["num_qubits"] = dcfg.num_qubits trace_fit_kw["qubit_bit_order"] = dcfg.qubit_bit_order if clf_name == "remf" and post_measured is not None: trace_fit_kw["train_post_measured"] = post_measured.get("train") trace_fit_kw["val_post_measured"] = post_measured.get("val") if not is_threshold: train_kw = training_kwargs(cfg, clf_name, clf_config) train_kw.update(trace_fit_kw) history = train( clf, train_features, splits.train_labels, val_features, splits.val_labels, **train_kw, ) elif is_threshold and hasattr(clf, "fit"): clf.fit( train_features, splits.train_labels, val_features, splits.val_labels, **threshold_fit_kw, ) eval_kw: dict[str, Any] = { "num_qubits": dcfg.num_qubits, "num_levels": dcfg.num_levels, "qubit_bit_order": dcfg.qubit_bit_order, } if clf_name in iq_trace_clfs: eval_kw["traces"] = splits.test_traces metrics = evaluate( clf, test_features, splits.test_labels, **eval_kw, ) sweep_result = None if cfg.sweep.enabled and transition_result is not None: sweep_kw: dict[str, Any] = { "num_qubits": dcfg.num_qubits, "num_levels": dcfg.num_levels, "filter_types": tuple(ccfg.filter_types), "qubit_bit_order": dcfg.qubit_bit_order, } if ccfg.threshold_percentile is not None: sweep_kw["threshold_percentile"] = ccfg.threshold_percentile if cfg.sweep.lengths: sweep_kw["lengths"] = cfg.sweep.lengths else: sweep_kw["n_points"] = cfg.sweep.n_points st = cfg.sweep.timing sweep_kw["sampling_period_ns"] = st.sampling_period_ns sweep_kw["data_transfer_ns"] = st.data_transfer_ns sweep_kw["inference_ns"] = st.inference_ns seq_kinds = ( "iq_trace", "raw_trace", "raw_trace_seq", "mf_trace_seq", "mf_rmf_trace_seq", ) if (feat_kind in seq_kinds or clf_name in seq_trace_clfs) and not is_threshold: sweep_kw["iq_trace_sweep"] = clf_name in iq_trace_clfs if clf_name in iq_trace_clfs: sweep_kw["train_kwargs"] = { **training_kwargs( cfg, clf_name, _classifier_config(cfg, clf_name, filter_set), ), **trace_fit_kw, "num_qubits": dcfg.num_qubits, } sweep_result = sweep_tail_zeroing( clf_cls, _classifier_config(cfg, clf_name, filter_set), extract_fn, splits, filter_set, qubit_bit_order=dcfg.qubit_bit_order, **sweep_kw, ) else: if not is_threshold: clf_config_eff = _classifier_config(cfg, clf_name, filter_set) sweep_kw["model_factory"] = lambda: clf_cls.from_config(clf_config_eff) sweep_kw["train_kwargs"] = training_kwargs(cfg, clf_name, clf_config_eff) sweep_result = sweep_trace_length( transition_result, splits.train_traces, splits.train_labels, splits.test_traces, splits.test_labels, val_traces=splits.val_traces, val_labels=splits.val_labels, **sweep_kw, ) bench_traces = None if clf_name in iq_trace_clfs or clf_name in seq_trace_clfs: bench_traces = splits.test_traces return { "model": clf, "metrics": metrics, "history": history, "sweep": sweep_result, "best_params": best_params, "benchmark_features": test_features, "benchmark_traces": bench_traces, }
[docs] def run_classify_stage(cfg: Any, data_out: dict) -> dict[str, Any]: """Run the default classify stage. Thin wrapper around :class:`DefaultClassifyStage`. Args: cfg: Pipeline config with ``data`` and ``classifiers`` blocks. data_out: Output of the data stage. Returns: Classify-stage dict (``filter_set``, ``classifiers``). """ return DefaultClassifyStage().run(cfg, data_out)
[docs] def run_single_classifier(*args: Any, **kwargs: Any) -> dict[str, Any]: """Train and evaluate one classifier (module-level entry point). Backward-compatible wrapper around :meth:`DefaultClassifyStage.run_single_classifier`. Args: *args: Positional args forwarded to :meth:`DefaultClassifyStage.run_single_classifier`. **kwargs: Keyword args forwarded to :meth:`DefaultClassifyStage.run_single_classifier`. Returns: Per-classifier result dict (``model``, ``metrics``, …). """ return DefaultClassifyStage().run_single_classifier(*args, **kwargs)