Scikit-learn
human-activity-recognition
wearable
wrist
time-series
cpu
scikit-learn
File size: 16,965 Bytes
80b01cc
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
"""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())