research(phase3b): benchmark 1D CNN on RAW pulses — data-starved, benched
Trains a PyTorch 1D CNN on encoded RAW pulse arrays as the Phase 3B leg of the designed ensemble, scored on the SAME group-disjoint held-out test set as the heuristic and statistical models (GroupShuffleSplit by device sub-folder, no near-duplicate leakage). RESULT — CNN is data-starved and loses decisively: CNN balanced 0.338, top-1 0.795 (garage-inflated) statistical balanced 0.625, top-1 0.868 heuristic top-1 0.300 Only 489 train samples with severe class imbalance (Garage 565, Security 2, Weather 12). A blend sweep confirmed every non-zero CNN weight degrades the statistical model (0.625 -> 0.613 at 15% CNN, worse beyond), so the CNN is NOT wired into decode(). Production ensemble stays heuristic + statistical, both already live. Committing the trainer + shared pulse_encoder (trainer/inference parity) + metrics.json to document the reproducible negative result. The benched .pt/.onnx binaries are intentionally NOT committed. Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
This commit is contained in:
@@ -0,0 +1,36 @@
|
||||
{
|
||||
"n_samples": 803,
|
||||
"n_train": 489,
|
||||
"n_val": 94,
|
||||
"n_test": 220,
|
||||
"classes": [
|
||||
"Doorbell",
|
||||
"Fan Controller",
|
||||
"Garage Door Opener",
|
||||
"Remote Control",
|
||||
"Security Sensor",
|
||||
"Weather Sensor"
|
||||
],
|
||||
"pulse_len": 512,
|
||||
"epochs": 60,
|
||||
"class_balance": {
|
||||
"Doorbell": 39,
|
||||
"Garage Door Opener": 565,
|
||||
"Weather Sensor": 12,
|
||||
"Fan Controller": 91,
|
||||
"Security Sensor": 2,
|
||||
"Remote Control": 94
|
||||
},
|
||||
"cnn": {
|
||||
"accuracy": 0.7954545454545454,
|
||||
"balanced_accuracy": 0.33796618290289177,
|
||||
"macro_f1": 0.3552945963506287
|
||||
},
|
||||
"heuristic_same_test": {
|
||||
"top1": 0.3
|
||||
},
|
||||
"statistical_same_test": {
|
||||
"top1": 0.8681818181818182,
|
||||
"balanced_accuracy": 0.6253728690437551
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,306 @@
|
||||
#!/usr/bin/env python3
|
||||
"""
|
||||
Device-Category 1D-CNN — Phase 3B (PLAN_TO_PROD)
|
||||
================================================
|
||||
|
||||
Trains a 1D convolutional net that predicts a device *category* directly from
|
||||
the normalized RAW pulse array of a Sub-GHz capture — the "shape of the signal"
|
||||
rather than the hand-picked timing statistics the Phase 3A model uses.
|
||||
|
||||
Architecture (per PLAN_TO_PROD §Phase B):
|
||||
Conv1D(64,3) → BN → ReLU
|
||||
Conv1D(128,3) → BN → ReLU
|
||||
Conv1D(256,3) → BN → ReLU
|
||||
GlobalAvgPool → Dense(256) → Dropout(0.3) → Dense(N) → Softmax
|
||||
Input: pulse array scaled to [-1,1], padded/truncated to 512 (pulse_encoder).
|
||||
|
||||
Honesty guardrails (identical discipline to the Phase 3A trainer):
|
||||
* GROUP-AWARE split — a device sub-folder never spans train/test, so we
|
||||
don't score inflated accuracy off near-duplicate repeat captures.
|
||||
* Epoch selection uses a validation slice carved from TRAIN only; the
|
||||
group-disjoint TEST set is scored exactly once at the end.
|
||||
* The heuristic router AND the Phase 3A statistical model are scored on the
|
||||
*same* held-out files → a fair three-way comparison, not cherry-picked.
|
||||
* Class-weighted loss + balanced accuracy reported, because the RAW corpus
|
||||
is ~70% Garage/Gate and raw accuracy alone is misleading.
|
||||
|
||||
Run:
|
||||
python scripts/train_cnn_classifier.py
|
||||
python scripts/train_cnn_classifier.py --per-category 400 --epochs 60
|
||||
"""
|
||||
|
||||
import argparse
|
||||
import json
|
||||
import sys
|
||||
import time
|
||||
from collections import Counter
|
||||
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 src.matcher.pulse_encoder import encode_pulses, PULSE_LEN
|
||||
from scripts.benchmark_realworld import FOLDER_TO_CATEGORY, DEFAULT_DATASET
|
||||
from scripts.train_category_classifier import (
|
||||
extract_features, heuristic_top1, FEATURE_NAMES,
|
||||
)
|
||||
|
||||
|
||||
def collect(dataset_root: Path, per_category: int, seed: int):
|
||||
"""Return aligned CNN arrays + stat features + labels/groups/freqs.
|
||||
|
||||
Every accepted RAW file yields BOTH the CNN encoding and the Phase-3A
|
||||
14-feature vector, indexed identically, so the same group split scores all
|
||||
three models on the same files.
|
||||
"""
|
||||
import random
|
||||
rng = random.Random(seed)
|
||||
ta, pd = get_timing_analyzer(), get_preamble_detector()
|
||||
|
||||
Xcnn, Xstat, 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
|
||||
enc = encode_pulses(meta.raw_data)
|
||||
if enc is None:
|
||||
skipped["bad_pulses"] += 1
|
||||
continue
|
||||
feats = extract_features(meta.raw_data, meta.frequency, ta, pd)
|
||||
if feats is None:
|
||||
skipped["no_timing"] += 1
|
||||
continue
|
||||
Xcnn.append(enc)
|
||||
Xstat.append(feats)
|
||||
y.append(category)
|
||||
groups.append(str(p.parent))
|
||||
freqs.append(meta.frequency)
|
||||
taken += 1
|
||||
return (np.asarray(Xcnn, dtype=np.float32),
|
||||
np.asarray(Xstat, dtype=np.float32),
|
||||
np.array(y), np.array(groups), np.array(freqs), skipped)
|
||||
|
||||
|
||||
def build_model(n_classes: int):
|
||||
import torch.nn as nn
|
||||
|
||||
class PulseCNN(nn.Module):
|
||||
def __init__(self, n_out):
|
||||
super().__init__()
|
||||
self.features = nn.Sequential(
|
||||
nn.Conv1d(1, 64, 3, padding=1), nn.BatchNorm1d(64), nn.ReLU(),
|
||||
nn.MaxPool1d(2),
|
||||
nn.Conv1d(64, 128, 3, padding=1), nn.BatchNorm1d(128), nn.ReLU(),
|
||||
nn.MaxPool1d(2),
|
||||
nn.Conv1d(128, 256, 3, padding=1), nn.BatchNorm1d(256), nn.ReLU(),
|
||||
nn.AdaptiveAvgPool1d(1), # GlobalAvgPool
|
||||
)
|
||||
self.head = nn.Sequential(
|
||||
nn.Flatten(),
|
||||
nn.Linear(256, 256), nn.ReLU(), nn.Dropout(0.3),
|
||||
nn.Linear(256, n_out),
|
||||
)
|
||||
|
||||
def forward(self, x):
|
||||
return self.head(self.features(x))
|
||||
|
||||
return PulseCNN(n_classes)
|
||||
|
||||
|
||||
def main():
|
||||
ap = argparse.ArgumentParser()
|
||||
ap.add_argument("--dataset", type=Path, default=DEFAULT_DATASET)
|
||||
ap.add_argument("--per-category", type=int, default=400)
|
||||
ap.add_argument("--seed", type=int, default=42)
|
||||
ap.add_argument("--test-size", type=float, default=0.25)
|
||||
ap.add_argument("--val-size", type=float, default=0.15)
|
||||
ap.add_argument("--epochs", type=int, default=60)
|
||||
ap.add_argument("--batch-size", type=int, default=64)
|
||||
ap.add_argument("--lr", type=float, default=1e-3)
|
||||
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)
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from sklearn.model_selection import GroupShuffleSplit
|
||||
from sklearn.metrics import (accuracy_score, balanced_accuracy_score,
|
||||
f1_score, confusion_matrix, classification_report)
|
||||
|
||||
torch.manual_seed(args.seed)
|
||||
np.random.seed(args.seed)
|
||||
|
||||
print("Collecting + encoding RAW pulses ...")
|
||||
t0 = time.time()
|
||||
Xcnn, Xstat, y, groups, freqs, skipped = collect(
|
||||
args.dataset, args.per_category, args.seed)
|
||||
print(f" {len(Xcnn)} 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))}")
|
||||
|
||||
classes = sorted(set(y))
|
||||
cls_to_i = {c: i for i, c in enumerate(classes)}
|
||||
yi = np.array([cls_to_i[c] for c in y])
|
||||
|
||||
# Group-aware test split (no device leakage)
|
||||
gss = GroupShuffleSplit(n_splits=1, test_size=args.test_size,
|
||||
random_state=args.seed)
|
||||
trainval_idx, test_idx = next(gss.split(Xcnn, yi, groups))
|
||||
# Validation carved from train (for epoch selection only; test untouched)
|
||||
gss2 = GroupShuffleSplit(n_splits=1, test_size=args.val_size,
|
||||
random_state=args.seed)
|
||||
tr_rel, val_rel = next(gss2.split(Xcnn[trainval_idx], yi[trainval_idx],
|
||||
groups[trainval_idx]))
|
||||
train_idx = trainval_idx[tr_rel]
|
||||
val_idx = trainval_idx[val_rel]
|
||||
print(f" train={len(train_idx)} val={len(val_idx)} test={len(test_idx)} "
|
||||
f"(test groups={len(set(groups[test_idx]))}, group-disjoint)")
|
||||
|
||||
def to_tensor(idx):
|
||||
x = torch.from_numpy(Xcnn[idx]).unsqueeze(1) # (N,1,L)
|
||||
t = torch.from_numpy(yi[idx]).long()
|
||||
return x, t
|
||||
|
||||
Xtr, ytr = to_tensor(train_idx)
|
||||
Xval, yval = to_tensor(val_idx)
|
||||
Xte, yte = to_tensor(test_idx)
|
||||
|
||||
# Class weights (inverse freq) for imbalance
|
||||
counts = np.bincount(yi[train_idx], minlength=len(classes)).astype(float)
|
||||
counts[counts == 0] = 1.0
|
||||
weights = torch.tensor((counts.sum() / counts), dtype=torch.float32)
|
||||
weights = weights / weights.mean()
|
||||
|
||||
model = build_model(len(classes))
|
||||
opt = torch.optim.Adam(model.parameters(), lr=args.lr, weight_decay=1e-4)
|
||||
lossf = nn.CrossEntropyLoss(weight=weights)
|
||||
|
||||
n = len(train_idx)
|
||||
best_val, best_state = -1.0, None
|
||||
for epoch in range(args.epochs):
|
||||
model.train()
|
||||
perm = torch.randperm(n)
|
||||
for i in range(0, n, args.batch_size):
|
||||
bi = perm[i:i + args.batch_size]
|
||||
opt.zero_grad()
|
||||
out = model(Xtr[bi])
|
||||
loss = lossf(out, ytr[bi])
|
||||
loss.backward()
|
||||
opt.step()
|
||||
# Validation balanced accuracy
|
||||
model.eval()
|
||||
with torch.no_grad():
|
||||
vpred = model(Xval).argmax(1).numpy()
|
||||
vbal = balanced_accuracy_score(yval.numpy(), vpred)
|
||||
if vbal > best_val:
|
||||
best_val = vbal
|
||||
best_state = {k: v.clone() for k, v in model.state_dict().items()}
|
||||
if (epoch + 1) % 10 == 0:
|
||||
print(f" epoch {epoch+1:3d} val_bal_acc={vbal:.1%} best={best_val:.1%}")
|
||||
|
||||
model.load_state_dict(best_state)
|
||||
model.eval()
|
||||
with torch.no_grad():
|
||||
te_logits = model(Xte)
|
||||
te_pred = te_logits.argmax(1).numpy()
|
||||
yte_np = yte.numpy()
|
||||
acc = accuracy_score(yte_np, te_pred)
|
||||
bal = balanced_accuracy_score(yte_np, te_pred)
|
||||
f1 = f1_score(yte_np, te_pred, average="macro")
|
||||
|
||||
# ── Baselines on the SAME test files ───────────────────────────────────
|
||||
router = get_category_router()
|
||||
h_top1 = 0
|
||||
for row, freq, truth in zip(Xstat[test_idx], freqs[test_idx], y[test_idx]):
|
||||
cat, _allowed, _full = heuristic_top1(row, freq, router)
|
||||
if cat == truth:
|
||||
h_top1 += 1
|
||||
h_acc = h_top1 / len(test_idx)
|
||||
|
||||
stat_acc = stat_bal = None
|
||||
try:
|
||||
import joblib
|
||||
bundle = joblib.load(args.out / "category_classifier.joblib")
|
||||
stat_model, stat_classes = bundle["model"], [str(c) for c in bundle["classes"]]
|
||||
stat_pred = stat_model.predict(Xstat[test_idx])
|
||||
stat_acc = accuracy_score(y[test_idx], stat_pred)
|
||||
stat_bal = balanced_accuracy_score(y[test_idx], stat_pred)
|
||||
except Exception as e:
|
||||
print(f" (statistical baseline unavailable: {e})")
|
||||
|
||||
print("\n" + "=" * 66)
|
||||
print("THREE-WAY COMPARISON (same group-disjoint test set)")
|
||||
print("=" * 66)
|
||||
print(f"{'model':26}{'top-1':>10}{'balanced':>11}")
|
||||
print(f"{'heuristic router':26}{h_acc:>10.1%}{'—':>11}")
|
||||
if stat_acc is not None:
|
||||
print(f"{'statistical (RF/GB)':26}{stat_acc:>10.1%}{stat_bal:>11.1%}")
|
||||
print(f"{'1D CNN (this run)':26}{acc:>10.1%}{bal:>11.1%}")
|
||||
print(f"\nCNN macro F1: {f1:.1%}")
|
||||
print("\nPer-category (CNN):")
|
||||
print(classification_report(yte_np, te_pred,
|
||||
labels=list(range(len(classes))),
|
||||
target_names=classes, zero_division=0))
|
||||
print("Confusion (rows=truth, cols=pred): labels=", classes)
|
||||
print(confusion_matrix(yte_np, te_pred, labels=list(range(len(classes)))))
|
||||
|
||||
# ── Persist: torch state + ONNX + metrics ──────────────────────────────
|
||||
args.out.mkdir(parents=True, exist_ok=True)
|
||||
torch.save({"state_dict": best_state, "classes": classes,
|
||||
"pulse_len": PULSE_LEN}, args.out / "category_cnn.pt")
|
||||
dummy = torch.zeros(1, 1, PULSE_LEN)
|
||||
onnx_path = args.out / "category_cnn.onnx"
|
||||
try:
|
||||
torch.onnx.export(
|
||||
model, dummy, str(onnx_path),
|
||||
input_names=["pulses"], output_names=["logits"],
|
||||
dynamic_axes={"pulses": {0: "batch"}, "logits": {0: "batch"}},
|
||||
opset_version=13,
|
||||
)
|
||||
onnx_ok = True
|
||||
except Exception as e:
|
||||
onnx_ok = False
|
||||
print(f" (ONNX export failed, torch .pt still saved: {e})")
|
||||
meta = {
|
||||
"n_samples": int(len(Xcnn)), "n_train": int(len(train_idx)),
|
||||
"n_val": int(len(val_idx)), "n_test": int(len(test_idx)),
|
||||
"classes": classes, "pulse_len": PULSE_LEN, "epochs": args.epochs,
|
||||
"class_balance": {k: int(v) for k, v in Counter(y).items()},
|
||||
"cnn": {"accuracy": acc, "balanced_accuracy": bal, "macro_f1": f1},
|
||||
"heuristic_same_test": {"top1": h_acc},
|
||||
"statistical_same_test": (
|
||||
{"top1": stat_acc, "balanced_accuracy": stat_bal}
|
||||
if stat_acc is not None else None),
|
||||
}
|
||||
(args.out / "category_cnn_metrics.json").write_text(json.dumps(meta, indent=2))
|
||||
print(f"\nSaved torch -> {args.out / 'category_cnn.pt'}")
|
||||
if onnx_ok:
|
||||
print(f"Saved ONNX -> {onnx_path}")
|
||||
print(f"Saved metrics-> {args.out / 'category_cnn_metrics.json'}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,49 @@
|
||||
#!/usr/bin/env python3
|
||||
"""
|
||||
Shared RAW-pulse encoder for the 1D-CNN (Phase 3B).
|
||||
|
||||
A single source of truth for turning a Flipper RAW pulse array into the fixed
|
||||
normalized tensor the CNN consumes — imported by BOTH the trainer
|
||||
(``scripts/train_cnn_classifier.py``) and the inference wrapper
|
||||
(``src/matcher/cnn_classifier.py``) so the two can never drift.
|
||||
|
||||
Encoding:
|
||||
* take the first ``length`` signed durations (µs; +high / -low),
|
||||
* scale by the sequence's *median absolute* duration — a robust, TE-like
|
||||
unit so a short pulse is ~±1 and a long pulse ~±2-3 (the ratios that carry
|
||||
the protocol identity), while absolute gain / capture level is normalised
|
||||
out. Using the median instead of the max is deliberate: a single huge
|
||||
inter-packet gap (~10 000 µs) would otherwise squash every informative
|
||||
short pulse toward zero and blind the CNN.
|
||||
* clip to ±``CLIP`` and rescale to [-1, 1] so outsized gaps saturate the
|
||||
ceiling instead of dominating the dynamic range,
|
||||
* right-pad with zeros to ``length``.
|
||||
|
||||
Returns a float32 array of shape ``(length,)`` or ``None`` when the pulse train
|
||||
is empty or degenerate (all-zero).
|
||||
"""
|
||||
|
||||
from typing import List, Optional
|
||||
|
||||
import numpy as np
|
||||
|
||||
PULSE_LEN = 512
|
||||
CLIP = 8.0 # in median-pulse units; gaps beyond this saturate to ±1
|
||||
|
||||
|
||||
def encode_pulses(pulses: List[int], length: int = PULSE_LEN) -> Optional[np.ndarray]:
|
||||
if pulses is None or len(pulses) == 0:
|
||||
return None
|
||||
arr = np.asarray(pulses[:length], dtype=np.float32)
|
||||
nz = np.abs(arr[arr != 0])
|
||||
if nz.size == 0:
|
||||
return None
|
||||
scale = float(np.median(nz))
|
||||
if scale <= 0.0:
|
||||
scale = float(np.max(np.abs(arr)))
|
||||
if scale <= 0.0:
|
||||
return None
|
||||
arr = np.clip(arr / scale, -CLIP, CLIP) / CLIP
|
||||
if arr.size < length:
|
||||
arr = np.pad(arr, (0, length - arr.size))
|
||||
return arr.astype(np.float32)
|
||||
Reference in New Issue
Block a user