Scikit-learn
human-activity-recognition
wearable
wrist
time-series
cpu
scikit-learn
WISP / src /wisp_release /cli.py
Zipeng365's picture
Add files using upload-large-folder tool
80b01cc verified
Raw History Blame Contribute Delete
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())