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 tests/test_release_api.py from Zipeng365/WISP: direct link, hf CLI and curl.
- Browser
- Download file 15.2 kB
-
https://huggingface.co/Zipeng365/WISP/resolve/main/tests/test_release_api.py
- Command line
-
hf download hf://Zipeng365/WISP/tests/test_release_api.py
-
curl -L -o test_release_api.py https://huggingface.co/Zipeng365/WISP/resolve/main/tests/test_release_api.py
15.2 kB
| """Small synthetic tests; run only on allocated compute, never login nodes.""" | |
| from __future__ import annotations | |
| import hashlib | |
| import json | |
| import pickle | |
| import sys | |
| import numpy as np | |
| import pytest | |
| from wisp_release import list_methods, load_checkpoint, predict_checkpoint | |
| from wisp_release.checkpoints import require_sequence_metadata | |
| from wisp_release.cli import main | |
| from wisp_release.data import as_har_data, read_npz | |
| from wisp_release.methods import build_family | |
| from wisp_release.selection import choose_validation_winner, fit_selector | |
| from wisp_release.scoring import score | |
| class RecordingModel: | |
| def __init__(self, sequence: bool = False): | |
| self.smoother_ = object() if sequence else None | |
| self.n_jobs = 99 | |
| self.called = False | |
| def predict(self, X, *, groups=None, time_index=None): | |
| self.called = True | |
| self.last_groups = groups | |
| self.last_time_index = time_index | |
| return np.zeros(len(X), dtype=np.int64) | |
| class SelectorFixture: | |
| fit_records = [] | |
| def __init__(self, family): | |
| self.family = family | |
| self.smoother_ = None | |
| self.n_jobs = 1 | |
| def fit(self, X, y, *, groups, time_index, n_classes): | |
| self.fit_records.append((self.family, len(X), n_classes)) | |
| self.classes_ = np.arange(n_classes) | |
| return self | |
| def predict(self, X, *, groups=None, time_index=None): | |
| return X[:, 0, 0].astype(np.int64) if self.family == "wisp_rc" else np.zeros(len(X), dtype=np.int64) | |
| def test_full_series_has_only_paper_methods(): | |
| rows = list_methods() | |
| assert len(rows) == 11 | |
| assert len({row["method_id"] for row in rows}) == 11 | |
| assert all(row["paper_name"].startswith("WISP-") for row in rows) | |
| assert sum(row["family_fit"] for row in rows) == 7 | |
| assert sum(row["search"] for row in rows) == 2 | |
| assert all("baseline" not in row["method_id"] for row in rows) | |
| def test_pickle_load_requires_explicit_trust_before_open(tmp_path): | |
| with pytest.raises(ValueError, match="arbitrary code"): | |
| load_checkpoint(tmp_path / "does_not_exist.pkl") | |
| def test_sha256_verified_before_unpickling(tmp_path, monkeypatch): | |
| path = tmp_path / "model.pkl" | |
| path.write_bytes(pickle.dumps(RecordingModel())) | |
| with pytest.raises(ValueError, match="SHA256 mismatch"): | |
| load_checkpoint(path, trust_pickle=True, expected_sha256="0" * 64) | |
| with pytest.raises(ValueError, match="hexadecimal"): | |
| load_checkpoint(path, trust_pickle=True, expected_sha256="not-a-digest") | |
| digest = hashlib.sha256(path.read_bytes()).hexdigest() | |
| assert isinstance(load_checkpoint(path, trust_pickle=True, expected_sha256=digest), RecordingModel) | |
| def test_legacy_class_namespace_without_global_alias(tmp_path): | |
| from wisp.cpu.algorithms import HarSculptForestHMM | |
| model = HarSculptForestHMM(seed=123, n_jobs=1) | |
| original = pickle.dumps(model, protocol=2) | |
| legacy = original.replace(b"cwisp.cpu.algorithms\n", b"ctsevolve.cpu.algorithms\n") | |
| assert original != legacy | |
| path = tmp_path / "legacy.pkl" | |
| path.write_bytes(legacy) | |
| before_modules = {key for key in sys.modules if key == "tsevolve" or key.startswith("tsevolve.")} | |
| loaded = load_checkpoint(path, trust_pickle=True) | |
| assert isinstance(loaded, HarSculptForestHMM) | |
| assert loaded.seed == 123 | |
| assert path.read_bytes() == legacy | |
| after_modules = {key for key in sys.modules if key == "tsevolve" or key.startswith("tsevolve.")} | |
| assert after_modules == before_modules | |
| def test_sequence_decoder_cannot_silently_be_disabled(missing): | |
| model = RecordingModel(sequence=True) | |
| subject = None if missing in ("subject", "both") else np.array([1, 1, 2]) | |
| time_index = None if missing in ("time_index", "both") else np.arange(3) | |
| with pytest.raises(ValueError, match="both subject"): | |
| predict_checkpoint(model, np.ones((3, 8, 3)), subject=subject, time_index=time_index) | |
| assert not model.called | |
| def test_sequence_prediction_retains_groups_times_and_limits_jobs(): | |
| model = RecordingModel(sequence=True) | |
| subject = np.array(["recording_a", "recording_a", "recording_b"]) | |
| time = np.array([0, 1, 0]) | |
| pred = predict_checkpoint(model, np.ones((3, 8, 3)), subject=subject, time_index=time, | |
| metadata={"input_shape": [8, 3], "label_names": ["rest", "walk"]}) | |
| np.testing.assert_array_equal(pred, [0, 0, 0]) | |
| np.testing.assert_array_equal(model.last_groups, subject) | |
| np.testing.assert_array_equal(model.last_time_index, time) | |
| assert model.n_jobs == 1 | |
| def test_shape_and_finite_value_guards(): | |
| model = RecordingModel() | |
| with pytest.raises(ValueError, match="time/channel"): | |
| predict_checkpoint(model, np.ones((3, 8, 3)), metadata={"input_shape": [16, 3]}) | |
| with pytest.raises(ValueError, match="NaN"): | |
| predict_checkpoint(model, np.full((3, 8, 3), np.nan)) | |
| with pytest.raises(ValueError, match="shape"): | |
| predict_checkpoint(model, np.ones((3, 8))) | |
| with pytest.raises(ValueError, match="one-dimensional"): | |
| require_sequence_metadata(model, [[1, 1, 1]], None, n=3) | |
| def test_short_vocabulary_rejected_even_when_missing_class_not_predicted(): | |
| model = RecordingModel() | |
| model.classes_ = np.arange(3) | |
| with pytest.raises(ValueError, match="preserve the complete fitted vocabulary"): | |
| predict_checkpoint(model, np.ones((3, 8, 3)), metadata={"label_names": ["rest", "walk"]}) | |
| assert model.called is False | |
| def test_nonstandard_checkpoint_class_axis_rejected_without_remapping(classes): | |
| model = RecordingModel() | |
| model.classes_ = classes | |
| with pytest.raises(ValueError, match="contiguous encoded global axis"): | |
| predict_checkpoint(model, np.ones((3, 8, 3))) | |
| assert model.called is False | |
| np.testing.assert_array_equal(model.classes_, classes) | |
| def test_npz_and_checkpoint_sidecar_vocab_conflict(tmp_path): | |
| model = RecordingModel() | |
| model.classes_ = np.arange(3) | |
| checkpoint = tmp_path / "model.pkl" | |
| checkpoint.write_bytes(pickle.dumps(model)) | |
| (tmp_path / "checkpoint.json").write_text(json.dumps({"input_shape": [8, 3], | |
| "label_names": ["rest", "walk", "run"]})) | |
| data = tmp_path / "test.npz" | |
| np.savez(data, X=np.ones((3, 8, 3)), y=np.zeros(3, dtype=np.int64), | |
| label_names=np.array(["rest", "run", "walk"])) | |
| with pytest.raises(SystemExit) as error: | |
| main(["evaluate", "--checkpoint", str(checkpoint), "--input", str(data), | |
| "--output", str(tmp_path / "metrics.json"), "--trust-pickle"]) | |
| assert error.value.code == 2 | |
| assert not (tmp_path / "metrics.json").exists() | |
| def test_numeric_npz_no_embedded_pickle(tmp_path): | |
| path = tmp_path / "safe.npz" | |
| np.savez(path, X=np.ones((4, 8, 3)), y=np.array([0, 2, 2, 0]), | |
| subject=np.array(["a", "a", "b", "b"]), time_index=np.array([0, 1, 0, 1]), | |
| label_names=np.array(["rest", "absent-in-fold", "walk"])) | |
| data = read_npz(path, require_y=True) | |
| dataset = as_har_data(data) | |
| assert dataset.n_classes == 3 | |
| assert dataset.label_names == ["rest", "absent-in-fold", "walk"] | |
| unsafe = tmp_path / "unsafe.npz" | |
| np.savez(unsafe, X=np.ones((4, 8, 3)), y=np.array([0, 2, 2, 0]), | |
| subject=np.asarray([{}, {}, {}, {}], dtype=object)) | |
| with pytest.raises(ValueError, match="object arrays"): | |
| read_npz(unsafe) | |
| def test_global_class_axis_requires_explicit_vocabulary(): | |
| data = {"X": np.ones((4, 8, 3)), "y": np.array([0, 2, 2, 0])} | |
| with pytest.raises(ValueError, match="global label_names"): | |
| as_har_data(data) | |
| assert as_har_data(data, n_classes=4).n_classes == 4 | |
| def test_cli_training_is_plan_only_by_default(capsys, tmp_path): | |
| assert main(["family-fit", "--method", "wisp_cis", "--input", "not-present.npz", | |
| "--output", str(tmp_path / "model.pkl"), "--seed", "42", "--n-jobs", "1"]) == 0 | |
| result = json.loads(capsys.readouterr().out) | |
| assert result["execute"] is False | |
| assert not (tmp_path / "model.pkl").exists() | |
| assert main(["search", "--method", "wisp_random", "--train", "train.npz", "--valid", "valid.npz", | |
| "--test", "test.npz", "--output", str(tmp_path / "search"), "--population", "2", | |
| "--generations", "1", "--max-evaluations", "4", "--seed", "42", "--n-jobs", "1"]) == 0 | |
| result = json.loads(capsys.readouterr().out) | |
| assert result["execute"] is False | |
| assert not (tmp_path / "search").exists() | |
| assert main(["select-fit", "--method", "wisp_select5", "--train", "train.npz", "--valid", "valid.npz", | |
| "--output", str(tmp_path / "selector.pkl"), "--seed", "42", "--n-jobs", "1"]) == 0 | |
| result = json.loads(capsys.readouterr().out) | |
| assert result["execute"] is False | |
| assert not (tmp_path / "selector.pkl").exists() | |
| def test_cli_evaluate_and_predict(tmp_path, capsys): | |
| model = RecordingModel(sequence=True) | |
| checkpoint = tmp_path / "model.pkl" | |
| checkpoint.write_bytes(pickle.dumps(model)) | |
| (tmp_path / "checkpoint.json").write_text(json.dumps({"input_shape": [8, 3], | |
| "label_names": ["rest", "walk"]})) | |
| data = tmp_path / "test.npz" | |
| np.savez(data, X=np.ones((4, 8, 3)), y=np.zeros(4, dtype=np.int64), | |
| subject=np.array(["a", "a", "b", "b"]), time_index=np.array([0, 1, 0, 1])) | |
| output = tmp_path / "metrics.json" | |
| base = ["--checkpoint", str(checkpoint), "--input", str(data), "--trust-pickle"] | |
| assert main(["evaluate", *base, "--output", str(output)]) == 0 | |
| assert json.loads(output.read_text())["macro_f1"] == 1.0 | |
| capsys.readouterr() | |
| predictions = tmp_path / "predictions.npz" | |
| assert main(["predict", *base, "--output", str(predictions)]) == 0 | |
| with np.load(predictions, allow_pickle=False) as archive: | |
| assert archive["y_pred_name"].tolist() == ["rest"] * 4 | |
| capsys.readouterr() | |
| with pytest.raises(SystemExit): | |
| main(["predict", *base, "--output", str(predictions)]) | |
| def test_inventory_does_not_unpickle(tmp_path, capsys): | |
| (tmp_path / "not_a_valid_model.pkl").write_bytes(b"not a pickle") | |
| assert main(["inventory", "--root", str(tmp_path)]) == 0 | |
| result = json.loads(capsys.readouterr().out) | |
| assert result["files_present"] == 1 | |
| assert "verified usable" in result["note"] | |
| def test_family_constructor_preserves_names_and_forces_positive_jobs(): | |
| model = build_family("wisp_cis", seed=7, n_jobs=1, direct=True, | |
| overrides={"n_kernels": 8, "n_intervals": 4}) | |
| assert model.seed == 7 | |
| assert model.use_hmm is False | |
| assert model.n_kernels == 8 | |
| with pytest.raises(ValueError, match="fixed families only"): | |
| build_family("wisp_random", seed=7) | |
| def test_synthetic_fixed_family_round_trip(tmp_path, capsys): | |
| random = np.random.default_rng(9) | |
| data = tmp_path / "tiny.npz" | |
| X = random.normal(size=(24, 24, 3)) | |
| y = np.tile([0, 2], 12).astype(np.int64) | |
| np.savez(data, X=X, y=y, label_names=np.array(["rest", "absent", "walk"])) | |
| output = tmp_path / "tiny.pkl" | |
| assert main(["family-fit", "--method", "wisp_rc", "--input", str(data), "--output", str(output), | |
| "--seed", "3", "--n-jobs", "1", "--direct", "--overrides", '{"n_kernels": 8}', | |
| "--execute"]) == 0 | |
| capsys.readouterr() | |
| metadata = json.loads((tmp_path / "tiny.metadata.json").read_text()) | |
| assert metadata["label_names"] == ["rest", "absent", "walk"] | |
| fitted = load_checkpoint(output, trust_pickle=True, expected_sha256=metadata["sha256"]) | |
| assert fitted.classes_.tolist() == [0, 1, 2] | |
| pred = predict_checkpoint(fitted, X, metadata=metadata) | |
| assert pred.shape == (24,) | |
| assert set(pred.tolist()) <= {0, 2} | |
| def test_paper_worst_class_includes_unobserved_global_class(): | |
| result = score(np.array([0, 2, 0, 2]), np.array([0, 2, 0, 2]), 3) | |
| assert result["macro_f1"] == 1.0 | |
| assert result["worst_class_f1"] == 0.0 | |
| assert len(result["confusion_matrix"]) == 3 | |
| def test_selector_tie_break_exact_precedence(): | |
| base = {"macro_f1": 0.7, "worst_class_f1": 0.4, "model_size_mb": 2.0} | |
| def pair(**changes): | |
| return [{**base, "method_id": "wisp_ss"}, {**base, "method_id": "wisp_rc", **changes}] | |
| assert choose_validation_winner(pair()) == "wisp_ss" | |
| assert choose_validation_winner(pair(model_size_mb=1.0)) == "wisp_rc" | |
| assert choose_validation_winner(pair(worst_class_f1=0.5, model_size_mb=100.0)) == "wisp_rc" | |
| assert choose_validation_winner(pair(macro_f1=0.8, worst_class_f1=0.0, model_size_mb=100.0)) == "wisp_rc" | |
| with pytest.raises(ValueError, match="finite"): | |
| choose_validation_winner(pair(macro_f1=float("nan"))) | |
| def test_selector_pickle_size_matches_preserved_default_protocol(): | |
| from wisp.cpu.algorithms import HarSculptForestHMM, estimate_pickle_size_mb | |
| from wisp_release.selection import _model_size_mb | |
| model = HarSculptForestHMM(seed=42, n_jobs=1) | |
| assert _model_size_mb(model) == estimate_pickle_size_mb(model) | |
| assert _model_size_mb(model) == len(pickle.dumps(model)) / (1024 * 1024) | |
| def test_selector_fits_train_candidates_then_refits_once_without_test(method, count, monkeypatch): | |
| from wisp.core.data import HARData | |
| import wisp_release.selection as selection | |
| y = np.array([0, 1, 2, 2]) | |
| X = np.ones((4, 8, 1)) | |
| X[:, 0, 0] = y | |
| train = HARData(X=X.copy(), y=y, subject=np.array(["train"] * 4), time_index=np.arange(4), | |
| label_names=["a", "b", "c"]) | |
| valid = HARData(X=X.copy(), y=y, subject=np.array(["valid"] * 4), time_index=np.arange(4), | |
| label_names=["a", "b", "c"]) | |
| SelectorFixture.fit_records = [] | |
| monkeypatch.setattr(selection, "_model_size_mb", lambda model: 1.0) | |
| fitted, report = fit_selector(method, train, valid, seed=42, n_jobs=1, | |
| builder=lambda family, **kwargs: SelectorFixture(family)) | |
| assert report["selected_method_id"] == "wisp_rc" | |
| assert report["selection_uses_test"] is False | |
| assert len(report["validation_candidates"]) == count | |
| assert len(SelectorFixture.fit_records) == count + 1 | |
| assert all(size == 4 for _, size, _ in SelectorFixture.fit_records[:-1]) | |
| assert SelectorFixture.fit_records[-1] == ("wisp_rc", 8, 3) | |
| assert fitted.family == "wisp_rc" | |
| if count == 5: | |
| assert {row["method_id"] for row in report["validation_candidates"]} == { | |
| "wisp_ss", "wisp_rc", "wisp_so", "wisp_gis", "wisp_cis"} | |
| leaking = HARData(X=X, y=y, subject=np.array(["train"] * 4), time_index=np.arange(4), | |
| label_names=["a", "b", "c"]) | |
| with pytest.raises(ValueError, match="leakage"): | |
| fit_selector(method, train, leaking, seed=42, builder=lambda family, **kwargs: SelectorFixture(family)) | |