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