"""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())