Instructions to use Zipeng365/WISP with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Scikit-learn
How to use Zipeng365/WISP with Scikit-learn:
from huggingface_hub import hf_hub_download import joblib model = joblib.load( hf_hub_download("Zipeng365/WISP", "sklearn_model.joblib") ) # only load pickle files from sources you trust # read more about it here https://skops.readthedocs.io/en/stable/persistence.html - Notebooks
- Google Colab
- Kaggle
Download src/wisp_release/selection.py from Zipeng365/WISP: direct link, hf CLI and curl.
- Browser
- Download file 4.82 kB
-
https://huggingface.co/Zipeng365/WISP/resolve/main/src/wisp_release/selection.py
- Command line
-
hf download hf://Zipeng365/WISP/src/wisp_release/selection.py
-
curl -L -o selection.py https://huggingface.co/Zipeng365/WISP/resolve/main/src/wisp_release/selection.py
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 | |