Scikit-learn
human-activity-recognition
wearable
wrist
time-series
cpu
scikit-learn
WISP / src /wisp_release /checkpoints.py
Zipeng365's picture
Add files using upload-large-folder tool
80b01cc verified
Raw History Blame Contribute Delete
7.74 kB
"""Explicitly trusted legacy checkpoint loading and guarded inference.
Pickle executes code. Namespace compatibility and SHA256 comparison do not make
an untrusted pickle safe; users must opt in before any unpickling takes place.
"""
from __future__ import annotations
import hashlib
import hmac
import pickle
import re
from pathlib import Path
from typing import Any, Mapping
class WISPUnpickler(pickle.Unpickler):
"""Map historical method modules, without installing process-wide aliases.
This is a compatibility loader, NOT a restricted or sandboxed unpickler.
"""
def find_class(self, module: str, name: str) -> Any:
if module == "tsevolve.search.ts_evolve_v3":
module = "wisp.search.controller"
elif module == "tsevolve" or module.startswith("tsevolve."):
module = "wisp" + module[len("tsevolve"):]
return super().find_class(module, name)
def sha256_file(path: str | Path) -> str:
digest = hashlib.sha256()
with Path(path).open("rb") as stream:
for block in iter(lambda: stream.read(1024 * 1024), b""):
digest.update(block)
return digest.hexdigest()
def load_checkpoint(path: str | Path, *, trust_pickle: bool = False,
expected_sha256: str | None = None) -> Any:
"""Load a trusted fitted WISP estimator, optionally checking its SHA256.
Supply ``expected_sha256`` from an independently trusted source. A digest
supplied alongside a checkpoint proves integrity, not publisher identity.
Historical checkpoint files are left unchanged on disk.
"""
if not trust_pickle:
raise ValueError("Pickle can execute arbitrary code. Load only trusted WISP "
"checkpoints and explicitly set trust_pickle=True.")
path = Path(path)
if expected_sha256 is not None:
if not re.fullmatch(r"[0-9a-fA-F]{64}", expected_sha256):
raise ValueError("expected_sha256 must be a 64-character hexadecimal digest")
if not hmac.compare_digest(sha256_file(path), expected_sha256.lower()):
raise ValueError(f"SHA256 mismatch for {path.name}; refusing to unpickle")
with path.open("rb") as stream:
model = WISPUnpickler(stream).load()
if not callable(getattr(model, "predict", None)):
raise TypeError("Checkpoint does not contain a fitted estimator with predict()")
return model
def _array_vector(value: Any, *, name: str, n: int) -> Any:
import numpy as np
array = np.asarray(value)
if array.ndim != 1 or len(array) != n:
raise ValueError(f"{name} must be a one-dimensional vector with {n} entries")
return array
def _check_X(X: Any, metadata: Mapping[str, Any] | None = None) -> Any:
import numpy as np
array = np.asarray(X, dtype=np.float64)
if array.ndim != 3 or not all(array.shape):
raise ValueError("X must have nonempty shape (n_windows, n_timepoints, n_channels)")
if not np.isfinite(array).all():
raise ValueError("X contains NaN or infinite values; apply the dataset preprocessing first")
if metadata and metadata.get("input_shape") is not None:
expected = tuple(int(v) for v in metadata["input_shape"])
# A complete source array shape may include a variable number of windows.
if len(expected) == 3:
expected = expected[1:]
if len(expected) != 2 or array.shape[1:] != expected:
raise ValueError(f"Checkpoint expects time/channel shape {expected}; got {array.shape[1:]}")
return array
def _validate_class_axis(model: Any, metadata: Mapping[str, Any] | None) -> None:
"""Never redefine a fitted global class axis using a shorter input vocabulary."""
import numpy as np
classes = getattr(model, "classes_", None)
if classes is not None:
classes = np.asarray(classes)
if (classes.ndim != 1 or not len(classes)
or not np.issubdtype(classes.dtype, np.integer)
or not np.array_equal(classes, np.arange(len(classes)))):
raise ValueError("Checkpoint classes_ must be the contiguous encoded global "
"axis 0..K-1; refusing to silently remap a nonstandard class axis")
names = None if metadata is None else metadata.get("label_names")
if names is not None:
if isinstance(names, (str, bytes)) or np.asarray(names).ndim != 1:
raise ValueError("label_names must be a one-dimensional global class vocabulary")
if classes is not None and len(names) != len(classes):
raise ValueError(f"label_names has {len(names)} entries but checkpoint global "
f"classes_ has {len(classes)}; preserve the complete fitted vocabulary")
def require_sequence_metadata(model: Any, subject: Any, time_index: Any, *, n: int) -> tuple[Any, Any]:
"""Do not silently bypass the sequence decoder stored in a checkpoint."""
needs_sequence = getattr(model, "smoother_", None) is not None
if needs_sequence and (subject is None or time_index is None):
raise ValueError("This checkpoint uses the sequence decoder: both subject "
"(sequence/session IDs) and time_index are required. "
"Do not merge unrelated recording sessions into one sequence.")
if subject is not None:
subject = _array_vector(subject, name="subject", n=n)
if time_index is not None:
time_index = _array_vector(time_index, name="time_index", n=n)
return subject, time_index
def set_inference_jobs(model: Any, n_jobs: int) -> None:
"""Limit the known fitted estimator/ensemble members without refitting."""
if n_jobs < 1:
raise ValueError("n_jobs must be positive")
seen: set[int] = set()
def visit(estimator: Any) -> None:
if estimator is None or id(estimator) in seen:
return
seen.add(id(estimator))
if hasattr(estimator, "n_jobs"):
estimator.n_jobs = n_jobs
for attr in ("model_", "members_", "estimators_", "steps"):
child = getattr(estimator, attr, None)
if isinstance(child, (tuple, list)):
for item in child:
visit(item[-1] if isinstance(item, tuple) else item)
elif child is not None:
visit(child)
visit(model)
def predict_checkpoint(model: Any, X: Any, *, subject: Any = None,
time_index: Any = None, metadata: Mapping[str, Any] | None = None,
n_jobs: int = 1) -> Any:
"""Predict encoded class IDs using the fitted checkpoint's actual decoder.
``metadata`` may contain ``input_shape=[T,C]`` and ``label_names``; it never
authorizes pickle loading, downloads data, or silently changes a decoder.
"""
import numpy as np
_validate_class_axis(model, metadata)
X = _check_X(X, metadata)
subject, time_index = require_sequence_metadata(model, subject, time_index, n=len(X))
set_inference_jobs(model, n_jobs)
prediction = np.asarray(model.predict(X, groups=subject, time_index=time_index))
if prediction.ndim != 1 or len(prediction) != len(X):
raise ValueError("Model returned an invalid prediction shape")
if not np.issubdtype(prediction.dtype, np.integer):
raise ValueError("WISP checkpoints must return encoded integer class IDs")
prediction = prediction.astype(np.int64, copy=False)
names = None if metadata is None else metadata.get("label_names")
if names is not None and len(prediction) and (prediction.min() < 0 or prediction.max() >= len(names)):
raise ValueError("Predictions are incompatible with the supplied label_names")
return prediction