Source code for arcade.data.loaders.npz

"""NumPy (.npy / .npz) data loader."""

from __future__ import annotations

from pathlib import Path
from typing import Any

import numpy as np

from ..base import BaseDataLoader, register_loader
from ..bundle import DataBundle


[docs] @register_loader("npz") class NPZLoader(BaseDataLoader): """Load readout traces from .npy or .npz files.""" supported_formats = (".npy", ".npz")
[docs] def load(self, path: str | Path, **kwargs: Any) -> DataBundle: path = Path(path) if path.suffix.lower() == ".npy": return self._load_npy(path) return self._load_npz(path)
def _load_npy(self, path: Path) -> DataBundle: data = np.load(path, allow_pickle=False) return DataBundle( traces=data, metadata={"source": str(path), "format": "npy"}, ) def _load_npz(self, path: Path) -> DataBundle: data = np.load(path, allow_pickle=False) meta: dict[str, Any] = {"source": str(path), "format": "npz"} if "traces" in data: traces = data["traces"] elif len(data.files) == 1: traces = data[data.files[0]] else: traces = data[data.files[0]] meta["all_keys"] = data.files labels = data["labels"] if "labels" in data else None time = data["time"] if "time" in data else None return DataBundle(traces=traces, labels=labels, time=time, metadata=meta)