"""Read numeric NPZ input without allowing embedded Python pickles.""" from __future__ import annotations from pathlib import Path from typing import Any def read_npz(path: str | Path, *, require_y: bool = False) -> dict[str, Any]: import numpy as np path = Path(path) with np.load(path, allow_pickle=False) as archive: if "X" not in archive: raise ValueError("NPZ must contain X with shape (n_windows,n_timepoints,n_channels)") if require_y and "y" not in archive: raise ValueError("NPZ must contain encoded integer y for this command") try: out = {key: np.asarray(archive[key]) for key in ("X", "y", "subject", "time_index", "label_names", "sampling_rate_hz") if key in archive} except ValueError as exc: raise ValueError("NPZ object arrays are not supported; export numeric arrays and " "Unicode strings, not pickled metadata") from exc X = out["X"] if X.ndim != 3 or not all(X.shape): raise ValueError("X must have nonempty shape (n_windows,n_timepoints,n_channels)") if "y" in out: y = out["y"] if y.ndim != 1 or len(y) != len(X) or not np.issubdtype(y.dtype, np.integer): raise ValueError("y must be a one-dimensional integer array aligned with X") if y.size and y.min() < 0: raise ValueError("y must use nonnegative global class IDs") for key in ("subject", "time_index"): if key in out and (out[key].ndim != 1 or len(out[key]) != len(X)): raise ValueError(f"{key} must be one-dimensional and aligned with X") if "label_names" in out: if out["label_names"].ndim != 1: raise ValueError("label_names must be a one-dimensional global class vocabulary") out["label_names"] = [str(value) for value in out["label_names"].tolist()] if "y" in out and out["y"].size and out["y"].max() >= len(out["label_names"]): raise ValueError("y is incompatible with label_names") return out def as_har_data(data: dict[str, Any], *, n_classes: int | None = None, source: str = "NPZ") -> Any: import numpy as np from wisp.core.data import HARData if "y" not in data: raise ValueError("y is required to construct a training/evaluation dataset") names = data.get("label_names") if names is None: if n_classes is None or n_classes < 1: raise ValueError("Supply the global label_names in NPZ or explicit --n-classes; " "a fold's observed labels must not redefine the class axis") names = [str(index) for index in range(n_classes)] elif n_classes is not None and n_classes != len(names): raise ValueError("n_classes differs from the supplied label_names vocabulary") if data["y"].size and data["y"].max() >= len(names): raise ValueError("y contains a class ID outside the global vocabulary") rate = data.get("sampling_rate_hz") return HARData(X=data["X"], y=data["y"], subject=data.get("subject"), time_index=data.get("time_index"), label_names=names, source=source, sampling_rate_hz=None if rate is None else float(np.asarray(rate).reshape(-1)[0]))