"""Top-level data loading dispatch."""
from __future__ import annotations
from pathlib import Path
from typing import Any
from .base import get_loader
from .bundle import DataBundle
from . import loaders # noqa: F401 — register built-in loaders
_FORMAT_MAP = {
".h5": "hdf5",
".hdf5": "hdf5",
".npy": "npz",
".npz": "npz",
".pkl": "pickle",
".pickle": "pickle",
}
def load(path: str | Path, fmt: str | None = None, **kwargs: Any) -> DataBundle:
"""Load data from disk, auto-detecting format from extension."""
path = Path(path)
if not path.exists():
raise FileNotFoundError(f"Data file not found: {path}")
if fmt is None:
suffix = path.suffix.lower()
fmt = _FORMAT_MAP.get(suffix)
if fmt is None:
raise ValueError(
f"Cannot infer format from extension '{suffix}'. Pass fmt= explicitly."
)
loader_cls = get_loader(fmt)
return loader_cls().load(path, **kwargs)
def _resolve_format(path: str | Path, config: dict[str, Any]) -> str:
fmt = config.get("format", "auto")
if fmt not in (None, "auto"):
return str(fmt)
suffix = Path(path).suffix.lower()
resolved = _FORMAT_MAP.get(suffix)
if resolved is None:
raise ValueError(
f"Cannot infer format from extension '{suffix}'. Set data.format in config."
)
return resolved
[docs]
def load_data(path: str | Path, config: dict[str, Any]) -> DataBundle:
"""Dispatch to the configured loader (``default``, registered name, or dotted path)."""
loader_name = config.get("loader", "default")
path = Path(path)
if loader_name == "default":
fmt = _resolve_format(path, config)
loader_cls = get_loader(fmt)
if fmt == "hdf5":
return loader_cls.from_config(config).load(str(path))
return loader_cls().load(str(path), config=config)
loader_cls = get_loader(loader_name)
if loader_name in {"hdf5", "path_signature"} and hasattr(loader_cls, "from_config"):
return loader_cls.from_config(config).load(str(path))
return loader_cls().load(str(path), config=config)