#!/usr/bin/env python3 """ Device-Category Classifier — Phase 3A (statistical ML) ====================================================== PLAN_TO_PROD Phase 3A: "Statistical Features (Lowest effort, highest ROI)". Trains a supervised classifier that predicts a device *category* from the timing/statistical features of a RAW Sub-GHz capture — the gap the heuristic category router struggles with, because shared line-encoders (Princeton, EV1527, Holtek, ...) reuse identical timing across device types. Why RAW-only: KEY .sub files already carry a decoded ``Protocol:`` name and are handled well by ``CategoryRouter.route_by_protocol``. ML's job is the RAW signals with no protocol match, so we train and evaluate on RAW captures only. Honesty guardrails: * GROUP-AWARE split — the same device sub-folder (e.g. one physical remote captured many times) never appears in both train and test, so we don't score inflated accuracy off near-duplicate captures. * The heuristic ``CategoryRouter`` is scored on the *exact same* held-out files, so the ML number is directly comparable, not cherry-picked. Run: python scripts/train_category_classifier.py python scripts/train_category_classifier.py --per-category 400 --out models/ """ import argparse import json import sys import time from collections import Counter, defaultdict from pathlib import Path import numpy as np sys.path.insert(0, str(Path(__file__).parent.parent)) from src.parser.sub_parser import parse_sub_file from src.matcher.timing_analyzer import get_timing_analyzer from src.matcher.preamble_detector import get_preamble_detector from src.matcher.category_router import get_category_router from scripts.benchmark_realworld import FOLDER_TO_CATEGORY, DEFAULT_DATASET # Feature order is frozen — inference must build vectors the same way. FEATURE_NAMES = [ "freq_mhz", "short_pulse_us", "long_pulse_us", "pulse_ratio", "short_gap_us", "long_gap_us", "gap_ratio", "duty_cycle", "pulse_count", "pulse_mean_abs", "pulse_std_abs", "pulse_min_abs", "pulse_max_abs", "preamble_type", ] _PREAMBLE_ID = {"none": 0, "long_burst": 1, "alternating": 2, "sync_word": 3, "custom": 4} def extract_features(pulses, frequency, ta, pd): """Timing + statistical feature vector for one RAW capture (or None).""" if not pulses: return None timing = ta.extract_timing(pulses) short, long = timing.short_pulse_us, timing.long_pulse_us if short == 0 or long == 0: return None detected = pd.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) return [ frequency / 1_000_000, short, long, timing.pulse_ratio, timing.short_gap_us, timing.long_gap_us, gap_ratio, timing.duty_cycle, len(pulses), float(abs_p.mean()), float(abs_p.std()), float(abs_p.min()), float(abs_p.max()), _PREAMBLE_ID.get(ptype, 0), ] def collect(dataset_root: Path, per_category: int, seed: int): """Return (X, y, groups, freqs) over RAW files in the mapped folders. groups = device sub-folder path (keeps repeat captures of one device together across the train/test split). """ import random rng = random.Random(seed) ta, pd = get_timing_analyzer(), get_preamble_detector() X, y, groups, freqs = [], [], [], [] skipped = Counter() for folder, category in FOLDER_TO_CATEGORY.items(): folder_path = dataset_root / folder if not folder_path.is_dir(): continue subs = sorted(folder_path.rglob("*.sub")) rng.shuffle(subs) taken = 0 for p in subs: if taken >= per_category: break try: meta = parse_sub_file(str(p)) except Exception: skipped["parse_error"] += 1 continue if not getattr(meta, "has_raw_data", False): skipped["no_raw"] += 1 continue feats = extract_features(meta.raw_data, meta.frequency, ta, pd) if feats is None: skipped["no_timing"] += 1 continue X.append(feats) y.append(category) groups.append(str(p.parent)) freqs.append(meta.frequency) taken += 1 return np.array(X), np.array(y), np.array(groups), np.array(freqs), skipped def heuristic_top1(X_row, freq, router): """Heuristic router's top-1 category for one feature row (same features).""" # X columns: 0 freq_mhz,1 short,2 long,...,8 pulse_count,...,13 preamble inv = {v: k for k, v in _PREAMBLE_ID.items()} pred = router.predict( frequency=int(freq), short_pulse_us=int(X_row[1]), long_pulse_us=int(X_row[2]), pulse_count=int(X_row[8]), preamble_type=inv.get(int(X_row[13]), "none"), ) return pred.primary_category, list(pred.allowed_categories or []), pred.use_full_db def main(): ap = argparse.ArgumentParser() ap.add_argument("--dataset", type=Path, default=DEFAULT_DATASET) ap.add_argument("--per-category", type=int, default=400, help="max RAW files sampled per folder") ap.add_argument("--seed", type=int, default=42) ap.add_argument("--test-size", type=float, default=0.25) ap.add_argument("--out", type=Path, default=Path("models")) args = ap.parse_args() if not args.dataset.is_dir(): print(f"Dataset not found: {args.dataset}", file=sys.stderr) sys.exit(1) from sklearn.ensemble import RandomForestClassifier, GradientBoostingClassifier from sklearn.model_selection import GroupShuffleSplit from sklearn.metrics import (accuracy_score, balanced_accuracy_score, f1_score, confusion_matrix, classification_report) import joblib print("Collecting + extracting features ...") t0 = time.time() X, y, groups, freqs, skipped = collect(args.dataset, args.per_category, args.seed) print(f" {len(X)} RAW samples in {time.time()-t0:.1f}s (skipped: {dict(skipped)})") print(f" class balance: {dict(Counter(y))}") print(f" distinct device groups: {len(set(groups))}") # Group-aware split (no device leakage between train/test) gss = GroupShuffleSplit(n_splits=1, test_size=args.test_size, random_state=args.seed) train_idx, test_idx = next(gss.split(X, y, groups)) Xtr, Xte = X[train_idx], X[test_idx] ytr, yte = y[train_idx], y[test_idx] fte = freqs[test_idx] print(f" train={len(Xtr)} test={len(Xte)} " f"(group-disjoint, {len(set(groups[test_idx]))} test groups)") models = { "random_forest": RandomForestClassifier( n_estimators=300, class_weight="balanced", random_state=args.seed, n_jobs=-1), "gradient_boost": GradientBoostingClassifier(random_state=args.seed), } results = {} best_name, best_bal = None, -1.0 for name, clf in models.items(): clf.fit(Xtr, ytr) pred = clf.predict(Xte) acc = accuracy_score(yte, pred) bal = balanced_accuracy_score(yte, pred) f1 = f1_score(yte, pred, average="macro") results[name] = (acc, bal, f1) print(f"\n── {name} ──") print(f" accuracy : {acc:.1%}") print(f" balanced accuracy : {bal:.1%}") print(f" macro F1 : {f1:.1%}") if bal > best_bal: best_name, best_bal, best_clf, best_pred = name, bal, clf, pred # ── Heuristic baseline on the SAME test files ────────────────────────── h_top1, h_routed = 0, 0 labels = sorted(set(y)) for row, freq, truth in zip(Xte, fte, yte): cat, allowed, full = heuristic_top1(row, freq, router=get_category_router()) if cat == truth: h_top1 += 1 if truth in allowed or full: h_routed += 1 n = len(Xte) print("\n" + "=" * 66) print(f"BEST MODEL: {best_name} (balanced acc {best_bal:.1%})") print("=" * 66) print(f"{'':26}{'top-1':>8}{'routed':>9}") print(f"{'heuristic router':26}{h_top1/n:>8.1%}{h_routed/n:>9.1%}") print(f"{'ML classifier (top-1)':26}" f"{accuracy_score(yte, best_pred):>8.1%}{'—':>9}") print("\nPer-category (best model):") print(classification_report(yte, best_pred, zero_division=0)) print("Confusion (rows=truth, cols=pred): labels=", labels) print(confusion_matrix(yte, best_pred, labels=labels)) print("\nFeature importances (best model, if available):") if hasattr(best_clf, "feature_importances_"): for nm, imp in sorted(zip(FEATURE_NAMES, best_clf.feature_importances_), key=lambda t: -t[1]): print(f" {nm:16} {imp:.3f}") # ── Persist ──────────────────────────────────────────────────────────── args.out.mkdir(parents=True, exist_ok=True) model_path = args.out / "category_classifier.joblib" joblib.dump({"model": best_clf, "features": FEATURE_NAMES, "classes": list(best_clf.classes_)}, model_path) meta = { "best_model": best_name, "n_samples": int(len(X)), "n_train": int(len(Xtr)), "n_test": int(len(Xte)), "features": FEATURE_NAMES, "class_balance": {k: int(v) for k, v in Counter(y).items()}, "metrics": {k: {"accuracy": v[0], "balanced_accuracy": v[1], "macro_f1": v[2]} for k, v in results.items()}, "heuristic_baseline_same_test": { "top1": h_top1 / n, "routed": h_routed / n, "n_test": n}, } (args.out / "category_classifier_metrics.json").write_text( json.dumps(meta, indent=2)) print(f"\nSaved model -> {model_path}") print(f"Saved metrics-> {args.out / 'category_classifier_metrics.json'}") if __name__ == "__main__": main()