Scikit-learn
human-activity-recognition
wearable
wrist
time-series
cpu
scikit-learn
WISP / src /wisp_release /selection.py
Zipeng365's picture
Add files using upload-large-folder tool
80b01cc verified
Raw History Blame Contribute Delete
4.82 kB
"""Validation-only WISP-Select5/Select7 and one final train+validation refit."""
from __future__ import annotations
import math
from typing import Any, Callable
from .checkpoints import predict_checkpoint
from .methods import FIXED_METHODS, build_family
from .scoring import score
SELECTOR_ORDER = tuple(FIXED_METHODS)
SELECT5_FAMILIES = ("wisp_ss", "wisp_rc", "wisp_so", "wisp_gis", "wisp_cis")
def selection_key(record: dict[str, Any]) -> tuple[float, float, float, int]:
"""Historical deterministic tie break: macro, worst F1, size, paper order."""
macro = float(record["macro_f1"])
worst = float(record["worst_class_f1"])
size = float(record["model_size_mb"])
if not all(math.isfinite(value) for value in (macro, worst, size)) or size < 0:
raise ValueError("Validation selection requires finite scores and nonnegative model sizes")
return (-macro, -worst, size, SELECTOR_ORDER.index(record["method_id"]))
def choose_validation_winner(records: list[dict[str, Any]]) -> str:
if not records:
raise ValueError("No validation candidates were evaluated")
return min(records, key=selection_key)["method_id"]
def _model_size_mb(model: Any) -> float:
# The original fixed-method runner called this preserved helper, which uses
# pickle's DEFAULT protocol (not HIGHEST_PROTOCOL). Renaming a Python module
# can change pickle byte counts: refits are not claimed binary-identical.
from wisp.cpu.algorithms import estimate_pickle_size_mb
return estimate_pickle_size_mb(model)
def fit_selector(method_id: str, train: Any, valid: Any, *, seed: int, n_jobs: int = 1,
builder: Callable[..., Any] | None = None) -> tuple[Any, dict[str, Any]]:
"""Fit/evaluate fixed families on validation, then refit the winner once.
No test input is accepted. ``train`` and ``valid`` are HARData objects with
the same global vocabulary, shape and disjoint participant/session IDs.
``builder`` is a testing/customization seam; default canonical settings are
the preserved paper registry. Changing a builder changes the experiment.
"""
import numpy as np
from wisp.core.data import concatenate_har
if method_id not in ("wisp_select5", "wisp_select7"):
raise ValueError("Expected wisp_select5 or wisp_select7")
if train.label_names != valid.label_names or train.n_classes != valid.n_classes:
raise ValueError("Train/validation must share the same global label vocabulary")
if train.X.shape[1:] != valid.X.shape[1:]:
raise ValueError("Train/validation must share the same time/channel shape")
for data in (train, valid):
if data.subject is None or data.time_index is None:
raise ValueError("Selector requires subject/session IDs and time_index")
train_groups = set(np.asarray(train.subject).astype(str).tolist())
valid_groups = set(np.asarray(valid.subject).astype(str).tolist())
if train_groups & valid_groups:
raise ValueError("Participant/session leakage between train and validation")
eligible = SELECT5_FAMILIES if method_id == "wisp_select5" else SELECTOR_ORDER
builder = builder or build_family
candidates = []
for family in eligible:
model = builder(family, seed=seed, n_jobs=n_jobs)
model.fit(train.X, train.y, groups=train.subject, time_index=train.time_index,
n_classes=train.n_classes)
prediction = predict_checkpoint(model, valid.X, subject=valid.subject,
time_index=valid.time_index, n_jobs=n_jobs)
candidates.append({"method_id": family,
**score(valid.y, prediction, valid.n_classes),
"model_size_mb": _model_size_mb(model)})
winner = choose_validation_winner(candidates)
combined = concatenate_har(train, valid, source="train+validation")
model = builder(winner, seed=seed, n_jobs=n_jobs)
model.fit(combined.X, combined.y, groups=combined.subject, time_index=combined.time_index,
n_classes=combined.n_classes)
report = {"method_id": method_id, "selected_method_id": winner,
"validation_candidates": candidates, "fit_scope": "train+validation",
"selection_uses_test": False, "model_seed": seed,
"tie_break": ["higher validation macro_f1", "higher full-vocabulary worst_class_f1",
"smaller fitted model default-protocol pickle size", "canonical family order"],
"serialized_size_caveat": "Module namespace renaming may alter serialized byte counts; "
"refitted models are not claimed binary-identical to historical files.",
"canonical_family_order": list(SELECTOR_ORDER)}
return model, report