Source code for arcade.data.loaders.pickle

"""Pickle data loader."""

from __future__ import annotations

import pickle
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("pickle") class PickleLoader(BaseDataLoader): """Load readout traces from pickle files.""" supported_formats = (".pkl", ".pickle")
[docs] def load(self, path: str | Path, **kwargs: Any) -> DataBundle: path = Path(path) meta: dict[str, Any] = {"source": str(path), "format": "pickle"} with open(path, "rb") as f: data = pickle.load(f) # noqa: S301 if isinstance(data, dict): return DataBundle( traces=data.get("traces", np.empty(0)), labels=data.get("labels"), time=data.get("time"), metadata={**meta, **data.get("metadata", {})}, ) if isinstance(data, np.ndarray): return DataBundle(traces=data, metadata=meta) return DataBundle(traces=np.empty(0), metadata={**meta, "raw": data})