feat: Phase 3A statistical device-category classifier (86.8% top-1 on RAW)
Trains sklearn RandomForest + GradientBoosting on 14 timing/statistical features from RAW-only UberGuidoZ captures, with a group-aware split (GroupShuffleSplit by device sub-folder) so near-duplicate captures never leak across train/test. The heuristic CategoryRouter is scored on the exact same held-out files for a fair comparison. Best model (gradient_boost): 86.8% top-1 accuracy, 62.5% balanced accuracy vs heuristic 30.0% top-1 / 55.9% routed. Raw accuracy is inflated by the 70%-Garage RAW class balance; balanced accuracy (~2x the heuristic) is the honest number and meets the Phase 3A 60-75% target. Not yet wired into the decoder ensemble. Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
This commit is contained in:
Binary file not shown.
@@ -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
|
||||
}
|
||||
}
|
||||
@@ -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()
|
||||
Reference in New Issue
Block a user