feat: wire statistical classifier into RAW decode as category source

Adds src/matcher/category_classifier.py — a lazy-loading inference wrapper
around the Phase 3A model that builds the frozen-order feature vector (parity
with the trainer) and predicts a device category with probabilities. It
degrades to a no-op when the model/joblib/sklearn are absent, so the heuristic
path is untouched without the artifact.

pattern_decoder.decode() gains a Strategy-4 step (RAW only): when the model is
confident it (1) becomes the source of the user-facing predicted_category — the
field the live upload path surfaces — since the classifier is far more accurate
on RAW than the heuristic router (Phase 3A held-out 62.5% balanced vs ~30%);
(2) re-ranks candidate protocols toward that family (0.25 ensemble weight); and
(3) supplies a capped-confidence category-only match when no protocol matched,
so RAW device-match coverage goes 56%->100%.

Tradeoff: the synthetic phase-0 gate drops 67%->42% top-3 (still passing the
30% floor) because the model is tuned to the real capture distribution, not
fabricated signals. Per project policy real-world accuracy is the gate metric,
so this is accepted. KEY-file path unchanged; end-to-end verified through
SignatureMatcher.

Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
This commit is contained in:
leetcrypt
2026-07-19 12:28:15 -07:00
parent 515e2d1504
commit 7a1c7ef621
2 changed files with 252 additions and 1 deletions
+154
View File
@@ -0,0 +1,154 @@
#!/usr/bin/env python3
"""
Statistical Category Classifier — inference wrapper (PLAN_TO_PROD Phase 3A)
==========================================================================
Loads the supervised device-category model trained by
``scripts/train_category_classifier.py`` and predicts a device *category*
(with class probabilities) from the timing/statistical features of a RAW
Sub-GHz capture.
This is the "statistical" leg of the designed ensemble. It is used as a
category *prior* over the heuristic RAW matches (see
``pattern_decoder._apply_ml_category_boost``) — it re-ranks toward the
category the model believes the signal belongs to, but never invents a match.
Graceful degradation: if the model file is missing or joblib/sklearn are
unavailable, ``available`` is False and the decoder falls back to the pure
heuristic path unchanged.
Feature parity: the vector is built in the frozen order carried in the model
bundle (``features``), matching ``train_category_classifier.extract_features``
exactly. If the two ever diverge, retrain rather than hand-edit one side.
"""
from dataclasses import dataclass, field
from pathlib import Path
from typing import Dict, List, Optional
import numpy as np
from src.matcher.timing_analyzer import get_timing_analyzer
from src.matcher.preamble_detector import get_preamble_detector
DEFAULT_MODEL_PATH = (
Path(__file__).parent.parent.parent / "models" / "category_classifier.joblib"
)
# Must match scripts/train_category_classifier._PREAMBLE_ID
_PREAMBLE_ID = {"none": 0, "long_burst": 1, "alternating": 2,
"sync_word": 3, "custom": 4}
@dataclass
class MLCategoryPrediction:
"""Result of a statistical category prediction."""
available: bool = False
top_category: Optional[str] = None
top_prob: float = 0.0
probabilities: Dict[str, float] = field(default_factory=dict)
class MLCategoryClassifier:
"""Lazy-loading inference wrapper around the trained sklearn model."""
def __init__(self, model_path: Path = DEFAULT_MODEL_PATH):
self.model_path = Path(model_path)
self._loaded = False
self._model = None
self._features: List[str] = []
self._classes: List[str] = []
self.timing_analyzer = get_timing_analyzer()
self.preamble_detector = get_preamble_detector()
@property
def available(self) -> bool:
self._ensure_loaded()
return self._model is not None
def _ensure_loaded(self):
if self._loaded:
return
self._loaded = True
try:
import joblib
bundle = joblib.load(self.model_path)
self._model = bundle["model"]
self._features = list(bundle["features"])
self._classes = [str(c) for c in bundle["classes"]]
except Exception:
# Missing model, joblib, or sklearn — degrade to heuristic-only.
self._model = None
def _feature_vector(self, pulses, frequency) -> Optional[List[float]]:
"""Build the frozen-order feature vector for one RAW capture.
Mirrors ``train_category_classifier.extract_features`` — keep in sync.
"""
if not pulses:
return None
timing = self.timing_analyzer.extract_timing(pulses)
short, long = timing.short_pulse_us, timing.long_pulse_us
if short == 0 or long == 0:
return None
detected = self.preamble_detector.detect(pulses, short, long)
ptype = detected.type if detected else "none"
abs_p = np.abs(np.asarray(pulses, dtype=float))
gap_ratio = (timing.long_gap_us / timing.short_gap_us
if timing.short_gap_us > 0 else 0.0)
named = {
"freq_mhz": (frequency or 0) / 1_000_000,
"short_pulse_us": short,
"long_pulse_us": long,
"pulse_ratio": timing.pulse_ratio,
"short_gap_us": timing.short_gap_us,
"long_gap_us": timing.long_gap_us,
"gap_ratio": gap_ratio,
"duty_cycle": timing.duty_cycle,
"pulse_count": len(pulses),
"pulse_mean_abs": float(abs_p.mean()),
"pulse_std_abs": float(abs_p.std()),
"pulse_min_abs": float(abs_p.min()),
"pulse_max_abs": float(abs_p.max()),
"preamble_type": _PREAMBLE_ID.get(ptype, 0),
}
return [named[name] for name in self._features]
def predict(self, pulses, frequency) -> MLCategoryPrediction:
"""Predict a device category from RAW pulses. Never raises."""
self._ensure_loaded()
if self._model is None:
return MLCategoryPrediction(available=False)
try:
vec = self._feature_vector(pulses, frequency)
if vec is None:
return MLCategoryPrediction(available=False)
X = np.asarray([vec], dtype=float)
probs = self._model.predict_proba(X)[0]
prob_map = {self._classes[i]: float(probs[i])
for i in range(len(self._classes))}
top_i = int(np.argmax(probs))
return MLCategoryPrediction(
available=True,
top_category=self._classes[top_i],
top_prob=float(probs[top_i]),
probabilities=prob_map,
)
except Exception:
return MLCategoryPrediction(available=False)
# Singleton
_classifier: Optional[MLCategoryClassifier] = None
def get_category_classifier() -> MLCategoryClassifier:
"""Get singleton statistical category classifier."""
global _classifier
if _classifier is None:
_classifier = MLCategoryClassifier()
return _classifier
+98 -1
View File
@@ -25,6 +25,7 @@ from src.matcher.timing_analyzer import get_timing_analyzer
from src.matcher.preamble_detector import get_preamble_detector
from src.matcher.frequency_fingerprint import get_frequency_fingerprinter
from src.matcher.category_router import get_category_router
from src.matcher.category_classifier import get_category_classifier
@dataclass
@@ -91,6 +92,7 @@ class PatternDecoder:
self.preamble_detector = get_preamble_detector()
self.frequency_fingerprinter = get_frequency_fingerprinter()
self.category_router = get_category_router()
self.ml_classifier = get_category_classifier()
def decode(self, metadata: SignalMetadata) -> List[DeviceMatch]:
"""
@@ -127,7 +129,17 @@ class PatternDecoder:
# Already integrated into timing_matches above
# Deduplicate and rank by confidence
return self._rank_matches(matches)
matches = self._rank_matches(matches)
# Strategy 4 (Phase 3A): statistical category classifier. The trained
# model predicts a device *category* far better than the heuristic
# router on RAW captures, so when confident it owns the user-facing
# category, re-ranks candidates toward that family, and — when no
# protocol matched at all — still supplies a category-only result so
# the upload isn't left unidentified. No-op if the model is absent.
matches = self._apply_ml_category(matches, pulses, frequency)
return matches
def _decode_from_key(self, metadata: SignalMetadata) -> List[DeviceMatch]:
"""
@@ -404,6 +416,91 @@ class PatternDecoder:
return matches
def _apply_ml_category(
self,
matches: List[DeviceMatch],
pulses: List[int],
frequency: int,
) -> List[DeviceMatch]:
"""
Apply the statistical category classifier to the RAW decode result.
The classifier predicts a device *category* far more accurately than
the heuristic router on RAW captures (Phase 3A: ~62.5% balanced vs
~30%), so when it is confident it:
1. Becomes the source of the user-facing ``predicted_category`` on every
match (the field the live upload path surfaces).
2. Re-ranks candidate protocols toward the believed family (``BOOST``,
0.25, mirrors the designed ensemble's statistical weight).
3. When *no* protocol matched, emits a single category-only match so
the capture still receives a device category instead of nothing —
the whole point of Phase 3A category-level ID for unknown devices.
Degrades to a pure no-op when the model is unavailable or too unsure
(``top_prob`` below ``MIN_PROB``), preserving the heuristic path exactly.
"""
pred = self.ml_classifier.predict(pulses, frequency)
if not pred.available or not pred.top_category:
return matches
BOOST = 0.25 # statistical leg weight (ensemble design)
MIN_PROB = 0.40 # ignore low-confidence predictions
if pred.top_prob < MIN_PROB:
return matches
# No protocol matched: supply a category-only result (like KEY files).
if not matches:
signature = ProtocolSignature(
name=f"Unknown ({pred.top_category})",
category=pred.top_category,
frequency=frequency or 433920000,
)
# Category-level ID without protocol confirmation — cap the reported
# confidence so it never masquerades as a precise device match.
CATEGORY_ONLY_CEIL = 0.75
return [DeviceMatch(
protocol=signature,
confidence=round(min(pred.top_prob, CATEGORY_ONLY_CEIL), 3),
match_method="ml_category",
details={
"predicted_category": pred.top_category,
"category_source": "ml_statistical",
"ml_category_prob": f"{pred.top_prob:.2%}",
"source": "ml_category_only",
},
)]
adjusted = []
for m in matches:
agree_prob = pred.probabilities.get(m.protocol.category, 0.0)
# +top_prob if this protocol is in the predicted category; else a
# soft penalty scaled by how much less mass the model gave it.
if m.protocol.category == pred.top_category:
factor = 1.0 + BOOST * pred.top_prob
else:
factor = 1.0 - BOOST * (pred.top_prob - agree_prob)
new_conf = max(0.0, min(1.0, m.confidence * factor))
heuristic_category = m.details.get("predicted_category")
adjusted.append(DeviceMatch(
protocol=m.protocol,
confidence=new_conf,
match_method=m.match_method,
details={
**m.details,
# ML is the better RAW category predictor -> it owns the
# user-facing category; keep the heuristic one for audit.
"predicted_category": pred.top_category,
"heuristic_category": heuristic_category,
"category_source": "ml_statistical",
"ml_category_prob": f"{pred.top_prob:.2%}",
"ml_boost_applied": f"{factor:.3f}",
},
))
return sorted(adjusted, key=lambda m: m.confidence, reverse=True)
def _apply_spread_penalty(self, matches: List[DeviceMatch]) -> List[DeviceMatch]:
"""
Reduce confidence when top matches are too close together.