"""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