Scikit-learn
human-activity-recognition
wearable
wrist
time-series
cpu
scikit-learn
WISP / tests /test_release_api.py
Zipeng365's picture
Add files using upload-large-folder tool
10ef792 verified
Raw History Blame Contribute Delete
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
@pytest.mark.parametrize("missing", ["subject", "time_index", "both"])
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
@pytest.mark.parametrize("classes", [np.array([0, 2]), np.array([1, 0]), np.array([0, 0]),
np.array([0.0, 1.0]), np.array([], dtype=np.int64)])
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)
@pytest.mark.parametrize("method,count", [("wisp_select5", 5), ("wisp_select7", 7)])
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))