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