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/checkpoints.py from Zipeng365/WISP: direct link, hf CLI and curl.
- Browser
- Download file 7.74 kB
-
https://huggingface.co/Zipeng365/WISP/resolve/main/src/wisp_release/checkpoints.py
- Command line
-
hf download hf://Zipeng365/WISP/src/wisp_release/checkpoints.py
-
curl -L -o checkpoints.py https://huggingface.co/Zipeng365/WISP/resolve/main/src/wisp_release/checkpoints.py
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 | |