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/cli.py from Zipeng365/WISP: direct link, hf CLI and curl.
- Browser
- Download file 17 kB
-
https://huggingface.co/Zipeng365/WISP/resolve/main/src/wisp_release/cli.py
- Command line
-
hf download hf://Zipeng365/WISP/src/wisp_release/cli.py
-
curl -L -o cli.py https://huggingface.co/Zipeng365/WISP/resolve/main/src/wisp_release/cli.py
17 kB
| """Command-line release entry points. No network calls or baseline imports.""" | |
| from __future__ import annotations | |
| import argparse | |
| import csv | |
| import json | |
| import pickle | |
| from collections import Counter | |
| from dataclasses import replace | |
| from pathlib import Path | |
| from typing import Any | |
| from .checkpoints import load_checkpoint, predict_checkpoint, sha256_file | |
| from .data import as_har_data, read_npz | |
| from .methods import FIXED_METHODS, SEARCH_METHODS, SELECTION_METHODS, build_family, list_methods | |
| def _json(payload: Any) -> str: | |
| return json.dumps(payload, ensure_ascii=False, indent=2, default=str) | |
| def _write_new(path: Path, payload: str) -> None: | |
| path.parent.mkdir(parents=True, exist_ok=True) | |
| with path.open("x", encoding="utf-8") as stream: | |
| stream.write(payload + "\n") | |
| def _metadata(args: argparse.Namespace, data: dict[str, Any]) -> dict[str, Any]: | |
| path = args.metadata | |
| if path is None: | |
| candidates = [args.checkpoint.parent / "checkpoint.json", | |
| args.checkpoint.with_name(args.checkpoint.stem + ".metadata.json")] | |
| path = next((candidate for candidate in candidates if candidate.is_file()), None) | |
| out = {} if path is None else json.loads(path.read_text(encoding="utf-8")) | |
| if not isinstance(out, dict): | |
| raise ValueError("Checkpoint metadata must be a JSON object") | |
| if "label_names" in data: | |
| if out.get("label_names") is not None and out["label_names"] != data["label_names"]: | |
| raise ValueError("Input NPZ label_names differ from checkpoint metadata") | |
| out.setdefault("label_names", data["label_names"]) | |
| return out | |
| def inventory(root: Path, manifest: Path | None = None) -> dict[str, Any]: | |
| """Inspect model paths/status without loading a pickle or claiming usability.""" | |
| manifest = manifest or root / "model_inventory.csv" | |
| if manifest.is_file(): | |
| with manifest.open(newline="", encoding="utf-8") as stream: | |
| rows = list(csv.DictReader(stream)) | |
| return {"manifest": str(manifest), "expected_cells": len(rows), | |
| "source_status_counts": dict(Counter(row.get("source_status", "unknown") for row in rows)), | |
| "method_counts": dict(Counter(row.get("method", "unknown") for row in rows)), | |
| "note": "Inventory status is not a checkpoint load or prediction test.", | |
| "models": rows} | |
| if not root.is_dir(): | |
| raise FileNotFoundError(root) | |
| files = sorted(root.rglob("*.pkl")) | |
| return {"root": str(root), "files_present": len(files), | |
| "bytes_present": sum(path.stat().st_size for path in files), | |
| "note": "Files present is not the same as verified usable models.", | |
| "models": [{"model_path": str(path.relative_to(root)), "bytes": path.stat().st_size} | |
| for path in files]} | |
| def _predict(args: argparse.Namespace, *, evaluate: bool = False) -> dict[str, Any]: | |
| import numpy as np | |
| if args.output.exists(): | |
| raise FileExistsError(f"Refusing to overwrite {args.output}") | |
| data = read_npz(args.input, require_y=evaluate) | |
| metadata = _metadata(args, data) | |
| model = load_checkpoint(args.checkpoint, trust_pickle=args.trust_pickle, | |
| expected_sha256=args.sha256) | |
| prediction = predict_checkpoint(model, data["X"], subject=data.get("subject"), | |
| time_index=data.get("time_index"), metadata=metadata, | |
| n_jobs=args.n_jobs) | |
| if evaluate: | |
| from wisp.core.metrics import classification_report_dict | |
| from .scoring import score | |
| result = classification_report_dict(data["y"], prediction, | |
| label_names=metadata.get("label_names")) | |
| classes = getattr(model, "classes_", None) | |
| n_classes = (len(metadata["label_names"]) if metadata.get("label_names") is not None | |
| else int(np.max(classes)) + 1 if classes is not None and len(classes) | |
| else int(max(np.max(data["y"]), np.max(prediction))) + 1) | |
| result.update(score(data["y"], prediction, n_classes)) | |
| result.update({"checkpoint": str(args.checkpoint), "input": str(args.input), | |
| "decoder": "sequence" if getattr(model, "smoother_", None) is not None else "direct", | |
| "metric_contract": "paper: macro_F1 on observed truth/prediction classes; worst_F1 on full global vocabulary", | |
| "test_usage": "This command evaluates; it does not select or refit a model."}) | |
| _write_new(args.output, _json(result)) | |
| return result | |
| arrays = {"y_pred": prediction} | |
| for key in ("subject", "time_index"): | |
| if key in data: | |
| arrays[key] = data[key] | |
| if metadata.get("label_names") is not None: | |
| names = np.asarray(metadata["label_names"], dtype=str) | |
| arrays["label_names"] = names | |
| arrays["y_pred_name"] = names[prediction] | |
| args.output.parent.mkdir(parents=True, exist_ok=True) | |
| with args.output.open("xb") as stream: | |
| np.savez_compressed(stream, **arrays) | |
| return {"output": str(args.output), "n_predictions": len(prediction), | |
| "decoder": "sequence" if getattr(model, "smoother_", None) is not None else "direct"} | |
| def _training_plan(args: argparse.Namespace) -> dict[str, Any]: | |
| result = {key: value for key, value in vars(args).items() if key != "execute"} | |
| result["execute"] = bool(args.execute) | |
| result["note"] = ("Training/search is enabled explicitly. Run on allocated compute resources." | |
| if args.execute else "Plan only. Add --execute to fit/search on allocated compute resources.") | |
| return result | |
| def _fit(args: argparse.Namespace) -> dict[str, Any]: | |
| if not args.execute: | |
| return _training_plan(args) | |
| metadata_path = args.output.with_name(args.output.stem + ".metadata.json") | |
| for path in (args.output, metadata_path): | |
| if path.exists(): | |
| raise FileExistsError(f"Refusing to overwrite {path}") | |
| data = read_npz(args.input, require_y=True) | |
| dataset = as_har_data(data, n_classes=args.n_classes, source=str(args.input)) | |
| overrides = {} if args.overrides is None else json.loads(args.overrides) | |
| if not isinstance(overrides, dict): | |
| raise ValueError("--overrides must be a JSON object") | |
| model = build_family(args.method, seed=args.seed, n_jobs=args.n_jobs, | |
| direct=args.direct, overrides=overrides) | |
| if getattr(model, "use_hmm", False) and (dataset.subject is None or dataset.time_index is None): | |
| raise ValueError("Sequence-decoder fitting requires both subject/session IDs and time_index; " | |
| "use --direct only if intentionally fitting a different direct-decoder variant") | |
| model.fit(dataset.X, dataset.y, groups=dataset.subject, time_index=dataset.time_index, | |
| n_classes=dataset.n_classes) | |
| args.output.parent.mkdir(parents=True, exist_ok=True) | |
| with args.output.open("xb") as stream: | |
| pickle.dump(model, stream, protocol=pickle.HIGHEST_PROTOCOL) | |
| metadata = {"method_id": args.method, "input_shape": list(dataset.X.shape[1:]), | |
| "label_names": dataset.label_names, "model_seed": args.seed, | |
| "fit_scope": "explicit input NPZ; no automatic benchmark split or refit", | |
| "use_hmm": getattr(model, "smoother_", None) is not None, | |
| "source_status": "newly_fitted", "sha256": sha256_file(args.output), | |
| "params": model.get_params(deep=False)} | |
| _write_new(metadata_path, _json(metadata)) | |
| return {"checkpoint": str(args.output), "metadata": str(metadata_path), | |
| "sha256": metadata["sha256"], "fit_seconds": getattr(model, "fit_seconds_", None)} | |
| def _select_fit(args: argparse.Namespace) -> dict[str, Any]: | |
| if not args.execute: | |
| return _training_plan(args) | |
| from .selection import fit_selector | |
| metadata_path = args.output.with_name(args.output.stem + ".metadata.json") | |
| selection_path = args.output.with_name(args.output.stem + ".selection.json") | |
| for path in (args.output, metadata_path, selection_path): | |
| if path.exists(): | |
| raise FileExistsError(f"Refusing to overwrite {path}") | |
| train, valid = [as_har_data(read_npz(path, require_y=True), n_classes=args.n_classes, | |
| source=str(path)) for path in (args.train, args.valid)] | |
| model, report = fit_selector(args.method, train, valid, seed=args.seed, n_jobs=args.n_jobs) | |
| args.output.parent.mkdir(parents=True, exist_ok=True) | |
| with args.output.open("xb") as stream: | |
| pickle.dump(model, stream, protocol=pickle.HIGHEST_PROTOCOL) | |
| metadata = {"method_id": args.method, "selected_method_id": report["selected_method_id"], | |
| "input_shape": list(train.X.shape[1:]), "label_names": train.label_names, | |
| "model_seed": args.seed, "fit_scope": "train+validation", | |
| "use_hmm": getattr(model, "smoother_", None) is not None, | |
| "source_status": "newly_fitted", "sha256": sha256_file(args.output), | |
| "selection_uses_test": False} | |
| _write_new(metadata_path, _json(metadata)) | |
| _write_new(selection_path, _json(report)) | |
| return {"checkpoint": str(args.output), "metadata": str(metadata_path), | |
| "selection": str(selection_path), "selected_method_id": report["selected_method_id"], | |
| "sha256": metadata["sha256"]} | |
| def _search(args: argparse.Namespace) -> dict[str, Any]: | |
| if not args.execute: | |
| return _training_plan(args) | |
| from wisp.search.controller import WISPSearchController | |
| from wisp.search.profiles import get_profile | |
| if args.candidate_workers > args.n_jobs: | |
| raise ValueError("candidate_workers cannot exceed the explicit total n_jobs budget") | |
| if args.max_evaluations < args.population: | |
| raise ValueError("max_evaluations must be at least the initial population") | |
| parts = [as_har_data(read_npz(path, require_y=True), n_classes=args.n_classes, source=str(path)) | |
| for path in (args.train, args.valid, args.test)] | |
| for part in parts: | |
| if part.subject is None or part.time_index is None: | |
| raise ValueError("Search datasets require subject/session IDs and time_index") | |
| if part.label_names != parts[0].label_names: | |
| raise ValueError("Train/validation/test must share the same global label_names") | |
| if part.X.shape[1:] != parts[0].X.shape[1:]: | |
| raise ValueError("Train/validation/test must share the same time/channel shape") | |
| groups = [set(part.subject.astype(str).tolist()) for part in parts] | |
| if any(groups[first] & groups[second] for first, second in ((0, 1), (0, 2), (1, 2))): | |
| raise ValueError("Participant/session leakage between train/validation/test") | |
| # Users supply locked partitions; this command never invents a split, and | |
| # selection stays on validation inside the preserved search controller. | |
| profile = replace(get_profile(args.method), max_evaluations=args.max_evaluations) | |
| controller = WISPSearchController( | |
| train=parts[0], valid=parts[1], test=parts[2], output=args.output, | |
| profile=profile, population=args.population, generations=args.generations, | |
| seed=args.seed, candidate_workers=args.candidate_workers, | |
| n_jobs_per_candidate=max(1, args.n_jobs // args.candidate_workers), | |
| ) | |
| result = controller.run() | |
| model_path = args.output / "selected_model.pkl" | |
| if model_path.is_file(): | |
| sidecar = args.output / "selected_model.metadata.json" | |
| metadata = {"method_id": args.method, "input_shape": list(parts[0].X.shape[1:]), | |
| "label_names": parts[0].label_names, "model_seed": args.seed, | |
| "fit_scope": "train+validation; selection uses validation only", | |
| "sha256": sha256_file(model_path), "source_status": "newly_fitted", | |
| "search_profile": profile.to_dict(), "selection_uses_test": False} | |
| if not sidecar.exists(): | |
| _write_new(sidecar, _json(metadata)) | |
| return result | |
| def _positive(value: str) -> int: | |
| result = int(value) | |
| if result < 1: | |
| raise argparse.ArgumentTypeError("must be positive") | |
| return result | |
| def make_parser() -> argparse.ArgumentParser: | |
| parser = argparse.ArgumentParser(prog="wisp", description="WISP method checkpoints and CPU workflows") | |
| parser.add_argument("--version", action="version", version="WISP release 0.1.0") | |
| sub = parser.add_subparsers(dest="command", required=True) | |
| sub.add_parser("list", help="List the full paper method series; no model loads") | |
| inv = sub.add_parser("inventory", help="Read inventory/file presence; no model loads") | |
| inv.add_argument("--root", type=Path, default=Path("models")) | |
| inv.add_argument("--manifest", type=Path) | |
| for name in ("predict", "evaluate"): | |
| command = sub.add_parser(name, help=f"{name.capitalize()} using an explicitly trusted fitted checkpoint") | |
| command.add_argument("--checkpoint", type=Path, required=True) | |
| command.add_argument("--input", type=Path, required=True) | |
| command.add_argument("--output", type=Path, required=True) | |
| command.add_argument("--trust-pickle", action="store_true") | |
| command.add_argument("--sha256", help="Expected digest from an independently trusted source") | |
| command.add_argument("--metadata", type=Path, help="JSON sidecar; checkpoint.json auto-detected when present") | |
| command.add_argument("--n-jobs", type=_positive, default=1) | |
| fit = sub.add_parser("family-fit", help="Plan or explicitly fit one of the seven fixed WISP families") | |
| fit.add_argument("--method", choices=tuple(FIXED_METHODS), required=True) | |
| fit.add_argument("--input", type=Path, required=True) | |
| fit.add_argument("--output", type=Path, required=True) | |
| fit.add_argument("--seed", type=int, required=True) | |
| fit.add_argument("--n-jobs", type=_positive, required=True) | |
| fit.add_argument("--n-classes", type=_positive) | |
| fit.add_argument("--direct", action="store_true", help="Intentionally fit without the sequence decoder") | |
| fit.add_argument("--overrides", help="Explicit JSON parameter overrides (changes the trained variant)") | |
| fit.add_argument("--execute", action="store_true", help="Actually fit; otherwise print a plan only") | |
| select = sub.add_parser("select-fit", help="Plan or explicitly select a fixed family on validation and refit once") | |
| select.add_argument("--method", choices=tuple(SELECTION_METHODS), required=True) | |
| select.add_argument("--train", type=Path, required=True) | |
| select.add_argument("--valid", type=Path, required=True) | |
| select.add_argument("--output", type=Path, required=True) | |
| select.add_argument("--seed", type=int, required=True) | |
| select.add_argument("--n-jobs", type=_positive, required=True) | |
| select.add_argument("--n-classes", type=_positive) | |
| select.add_argument("--execute", action="store_true", help="Actually select/refit; otherwise print a plan only") | |
| search = sub.add_parser("search", help="Plan or explicitly run the preserved CPU search controller") | |
| search.add_argument("--method", choices=tuple(SEARCH_METHODS), required=True) | |
| for part in ("train", "valid", "test"): | |
| search.add_argument(f"--{part}", type=Path, required=True) | |
| search.add_argument("--output", type=Path, required=True) | |
| search.add_argument("--population", type=_positive, required=True) | |
| search.add_argument("--generations", type=_positive, required=True) | |
| search.add_argument("--max-evaluations", type=_positive, required=True) | |
| search.add_argument("--seed", type=int, required=True) | |
| search.add_argument("--n-jobs", type=_positive, required=True, help="Total CPU thread budget") | |
| search.add_argument("--candidate-workers", type=_positive, default=1) | |
| search.add_argument("--n-classes", type=_positive) | |
| search.add_argument("--execute", action="store_true", help="Actually search; otherwise print a plan only") | |
| return parser | |
| def main(argv: list[str] | None = None) -> int: | |
| parser = make_parser() | |
| args = parser.parse_args(argv) | |
| try: | |
| if args.command == "list": | |
| result = list_methods() | |
| elif args.command == "inventory": | |
| result = inventory(args.root, args.manifest) | |
| elif args.command in ("predict", "evaluate"): | |
| result = _predict(args, evaluate=args.command == "evaluate") | |
| elif args.command == "family-fit": | |
| result = _fit(args) | |
| elif args.command == "select-fit": | |
| result = _select_fit(args) | |
| else: | |
| result = _search(args) | |
| except (ValueError, TypeError, FileNotFoundError, FileExistsError, KeyError) as exc: | |
| parser.error(str(exc)) | |
| print(_json(result)) | |
| return 0 | |
| if __name__ == "__main__": | |
| raise SystemExit(main()) | |