diff --git a/models/category_classifier.joblib b/models/category_classifier.joblib new file mode 100644 index 0000000..75ded70 Binary files /dev/null and b/models/category_classifier.joblib differ diff --git a/models/category_classifier_metrics.json b/models/category_classifier_metrics.json new file mode 100644 index 0000000..9400f43 --- /dev/null +++ b/models/category_classifier_metrics.json @@ -0,0 +1,47 @@ +{ + "best_model": "gradient_boost", + "n_samples": 803, + "n_train": 583, + "n_test": 220, + "features": [ + "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" + ], + "class_balance": { + "Doorbell": 39, + "Garage Door Opener": 565, + "Weather Sensor": 12, + "Fan Controller": 91, + "Security Sensor": 2, + "Remote Control": 94 + }, + "metrics": { + "random_forest": { + "accuracy": 0.8272727272727273, + "balanced_accuracy": 0.5789748226457088, + "macro_f1": 0.5360433604336042 + }, + "gradient_boost": { + "accuracy": 0.8681818181818182, + "balanced_accuracy": 0.6253728690437551, + "macro_f1": 0.5340975664713835 + } + }, + "heuristic_baseline_same_test": { + "top1": 0.3, + "routed": 0.5590909090909091, + "n_test": 220 + } +} \ No newline at end of file diff --git a/scripts/train_category_classifier.py b/scripts/train_category_classifier.py new file mode 100644 index 0000000..ffaad43 --- /dev/null +++ b/scripts/train_category_classifier.py @@ -0,0 +1,260 @@ +#!/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()