"""Feature / filter-bank stage."""
from __future__ import annotations
from typing import Any
from .._logging import get_logger
from .base import BaseFeatureStage
log = get_logger(__name__)
[docs]
class DefaultFeatureStage(BaseFeatureStage):
"""Compute MF / RMF / EMF matched filters from transition results."""
[docs]
def run(self, cfg: Any, data_out: dict[str, Any]) -> dict[str, Any]:
"""Build the matched-filter bank used by filter-set classifiers.
Args:
cfg: Pipeline config (``data.transitions``, ``classifiers.filter_types``).
data_out: Data-stage dict; uses ``transition_result`` and optional
``spectral_labels`` for effective level count.
Returns:
Dict with ``filter_set``, ``filter_types``, and
``transition_result``. ``filter_set`` is ``None`` when no
transition result is available.
"""
from arcade.features.matched_filter import compute_filters
dcfg = cfg.data
ccfg = cfg.classifiers
transition_result = data_out.get("transition_result")
filter_set = None
if transition_result is None:
log.warning("No transition result → skipping feature stage")
return {"filter_set": None, "filter_types": tuple(ccfg.filter_types)}
filter_types = tuple(ccfg.filter_types)
thr_pct = ccfg.threshold_percentile
tcfg = dcfg.transitions
num_levels_eff = (
tcfg.effective_levels
if (
tcfg.enabled
and tcfg.effective_levels is not None
and data_out.get("spectral_labels") is not None
)
else dcfg.num_levels
)
filt_kw: dict[str, Any] = {"filter_types": filter_types}
if thr_pct is not None:
filt_kw["threshold_percentile"] = thr_pct
filter_set = compute_filters(
transition_result,
dcfg.num_qubits,
num_levels_eff,
**filt_kw,
)
log.info("Filters computed: %d total features", filter_set.total_features)
return {
"filter_set": filter_set,
"filter_types": filter_types,
"transition_result": transition_result,
}