"""User-facing dataset descriptors and detectors."""
from __future__ import annotations
import urllib.parse
from collections.abc import Mapping
from dataclasses import dataclass, field
from pathlib import Path
from typing import Any
import numpy as np
from zephon._internal.io.catalog import (
DatasetHeader as _DatasetHeader,
)
from zephon._internal.io.catalog import (
ShardCatalogHandle as _ShardCatalogHandle,
)
from zephon._internal.io.formats import (
ensure_builtin_formats as _ensure_builtin_formats,
)
from zephon._internal.io.formats.base import (
FormatHandler as _FormatHandler,
)
from zephon._internal.io.formats.base import (
get_format as _get_format,
)
from zephon._internal.io.index import find_and_load_index as _find_and_load_index
from zephon._internal.io.index.index_types import IndexData as _IndexData
from zephon._internal.io.protocols import RandomAccessShard as _RandomAccessShard
from zephon._internal.io.storage import (
RouterStorageBackend as _RouterStorageBackend,
)
from zephon._internal.io.storage import (
StorageBackend as _StorageBackend,
)
from zephon._internal.io.suffixes import detect_format as _detect_format
from zephon.io.memory import InMemoryShard
[docs]
@dataclass(frozen=True)
class Dataset:
"""Descriptor for a dataset with index and backend info.
Instances are created via ``from_path`` (file-backed) or ``from_dict``
(in-memory/testing). They always provide:
- ``name``: a user-facing identifier used in mixtures
- ``backend``: opaque metadata that lets FetchOp build a reader later
- ``path``: original filesystem path if file-backed, otherwise ``None``
File-backed datasets also carry a few-KB handle to the node-local shard
catalog (set by ``from_path``); this is what travels in ctx, not
``shard_meta``.
Shard counts are not stored as a mapping; read them via :meth:`ids` /
:meth:`counts` / :meth:`total` / :meth:`max_count` / :meth:`shard_count`.
Backend kinds used by the internal store builder:
- ``litdata``/``mds``/``jsonl``/``parquet``/``vortex``: ``kind`` and
``path`` (no per-shard ``shards`` graph — that lives in the catalog now)
- ``inmem``: ``kind`` and ``shards`` (``dict[int, InMemoryShard]``)
Note: this class does not expose any method to fetch rows; IO is delegated
to an internal shard store owned by the FetchOp.
"""
name: str
backend: Mapping[str, object]
path: str | None = None
# Node-local shard-catalog handle (file-backed datasets only); set by
# from_path(), not a constructor argument.
_catalog_handle: _ShardCatalogHandle | None = field(
default=None, compare=False, init=False
)
# _ids/_counts are constructor-seeded discovery output, not a cache: at
# planning time the catalog doesn't exist yet (preflight builds it once per
# node, later), so they are the only copy the driver can read. Pickling
# drops them: by spawn time the catalog exists, and attach() maps pages
# shared node-wide where pickled arrays would be a private copy per worker
# (in-memory datasets re-derive from the shipped shards). Reads stay uncached.
_ids: np.ndarray | None = field(default=None, compare=False, repr=False)
_counts: np.ndarray | None = field(default=None, compare=False, repr=False)
def _inmem_shards(self) -> Mapping[int, InMemoryShard]:
shards = self.backend.get("shards")
assert self.backend.get("kind") == "inmem" and isinstance(shards, Mapping)
return shards
[docs]
def ids(self) -> np.ndarray:
"""Return the sorted shard ids as an ``int64`` array."""
if self._ids is not None:
return self._ids
if self._catalog_handle is not None:
return self._catalog_handle.attach().ids()
return _inmem_ids_counts(self._inmem_shards())[0]
[docs]
def counts(self) -> np.ndarray:
"""Return per-shard sample counts (``int64``), aligned with :meth:`ids`."""
if self._counts is not None:
return self._counts
if self._catalog_handle is not None:
return self._catalog_handle.attach().num_rows()
return _inmem_ids_counts(self._inmem_shards())[1]
[docs]
def raw_bytes(self) -> np.ndarray:
"""Return per-shard byte sizes (``int64``), aligned with :meth:`ids`.
File-backed datasets read the catalog's on-disk shard sizes; in-memory
shards size their resident payloads lazily (:attr:`InMemoryShard.raw_bytes`).
"""
if self._catalog_handle is not None:
return self._catalog_handle.ensure_attached().raw_bytes()
shards = self._inmem_shards()
return np.array([shards[i].raw_bytes for i in sorted(shards)], dtype=np.int64)
[docs]
def total(self) -> int:
"""Total sample count across all shards."""
return int(self.counts().sum())
[docs]
def max_count(self) -> int:
"""Largest single-shard sample count (0 when empty)."""
return int(self.counts().max(initial=0))
[docs]
def shard_count(self) -> int:
"""Number of shards."""
return self.ids().size
def __len__(self) -> int:
return self.total()
def __getstate__(self) -> dict[str, object]:
# Counts never travel — see the _ids/_counts field comment.
return {
"name": self.name,
"backend": self.backend,
"path": self.path,
"_catalog_handle": self._catalog_handle,
}
def __deepcopy__(self, memo: dict[int, Any]) -> "Dataset":
# Share, don't copy: a base-WorkSource deepcopy-clone must not duplicate
# the (immutable) handle and count arrays.
memo[id(self)] = self
return self
[docs]
@classmethod
def from_path(cls, name: str, path: str, *, fmt: str | None = None) -> "Dataset":
"""Construct a file-backed dataset descriptor.
Performs a *count-only* discovery: it obtains ``shard_id`` + ``num_rows``
as small numpy arrays (the work source's input) without materializing the
per-shard ``shard_meta`` graph. The full columnar catalog is built once
per node later, by the Engine's ``finalize()`` (or lazily by the store
builder for Engine-less use).
Args:
name: Logical dataset name used in mixtures and debugging.
path: Filesystem directory containing the dataset.
fmt: Optional explicit format. When ``None``, auto-detects.
Supported formats:
- ``litdata`` directories containing ``index.json`` structured with
``config`` and ``chunks``
- ``mds`` directories containing ``index.json`` structured with
``shards``
- ``jsonl`` directories where ``*.jsonl`` files act as shards
Special URI schemes:
- ``hf://org/name[@rev]/[config/]split`` is served through the
HuggingFace backend: the split's uploaded files when their format is
readable, else HuggingFace's Parquet conversion when it is complete
and built from the requested commit. ``fmt`` picks the source in that
format; appending ``~original`` or ``~parquet`` to the revision
(``hf://org/name@~parquet/split``) forces one. A partial conversion is
used only with ``ZEPHON_HF_ALLOW_PARTIAL=1``. Uploaded files are read
as stored: the dataset card's reader options and ``features`` casting
are not applied. The dataset's ``path`` becomes a URI pinned to the
resolved commit, source and config.
Returns:
Dataset: a descriptor populated with shard counts and a catalog
handle for later IO.
Raises:
FileNotFoundError: if ``path`` does not exist.
ValueError: if ``path`` is not a directory or format unsupported.
"""
url = urllib.parse.urlparse(path)
is_remote = bool(url.scheme)
root_path: Path | None = None
if not is_remote:
root_path = Path(path)
if not root_path.exists():
raise FileNotFoundError(f"Dataset path does not exist: {root_path}")
if not root_path.is_dir():
raise ValueError(f"Dataset path must be a directory: {root_path}")
root_path = root_path.resolve()
root_str = str(root_path)
else:
root_str = path.rstrip("/") or path
storage = _RouterStorageBackend()
root_str = storage.canonical_root(root_str, fmt=fmt)
kind = fmt
if kind is None:
kind = _auto_detect_format(storage, root_str, root_path)
if kind is None:
raise ValueError(f"Unsupported dataset format at path: {root_str}")
_ensure_builtin_formats(required={kind})
handler = _get_format(kind)
ids, counts = _discover_counts(handler, root_str, storage)
# ids()/counts() hand these arrays out shared (the work-source cursor
# gathers from them in place); freeze them so an in-place edit fails loud.
ids.setflags(write=False)
counts.setflags(write=False)
backend = {"kind": kind, "path": root_str}
header = _DatasetHeader(name=name, root=root_str, format=kind, path=root_str)
handle = _ShardCatalogHandle(dataset=header)
dataset = cls(
name=name,
backend=backend,
path=root_str,
_ids=ids,
_counts=counts,
)
# Non-init private field: set on the frozen instance post-construction.
object.__setattr__(dataset, "_catalog_handle", handle)
return dataset
[docs]
@classmethod
def from_dict(cls, name: str, shards: Mapping[int, InMemoryShard]) -> "Dataset":
"""Construct an in-memory dataset descriptor."""
norm = {int(sid): shard for sid, shard in shards.items()}
backend: Mapping[str, object] = {"kind": "inmem", "shards": norm}
ids, counts = _inmem_ids_counts(norm)
return cls(name=name, backend=backend, path=None, _ids=ids, _counts=counts)
def _inmem_ids_counts(
shards: Mapping[int, _RandomAccessShard],
) -> tuple[np.ndarray, np.ndarray]:
"""Derive aligned ``(ids, counts)`` from resident in-memory shards."""
ids = np.array(sorted(int(k) for k in shards), dtype=np.int64)
counts = np.array([len(shards[int(i)]) for i in ids], dtype=np.int64)
# Handed out shared (same contract as the from_path arrays): freeze them.
ids.setflags(write=False)
counts.setflags(write=False)
return ids, counts
def _discover_counts(
handler: _FormatHandler, path: str, storage: _StorageBackend
) -> tuple[np.ndarray, np.ndarray]:
"""Obtain ``(ids, counts)`` without materializing per-shard ``shard_meta``.
``FormatHandler.discover_counts`` gives index formats a metadata-only fast
path; the protocol default runs the full ``discover()`` and keeps only the
counts, so the heavy graph is at most transiently built, never retained.
Results are normalized to ``int64`` arrays.
"""
ids, counts = handler.discover_counts(path, storage)
return np.asarray(ids, dtype=np.int64), np.asarray(counts, dtype=np.int64)
__all__ = ["Dataset"]
def _auto_detect_format(
storage: _RouterStorageBackend, root_str: str, root_path: Path | None
) -> str | None:
result = _find_and_load_index(root_str, storage)
if result is not None:
return _classify_index_payload(result)
if root_path is not None:
entries = [p.name for p in root_path.iterdir() if p.is_file()]
else:
try:
entries = storage.listdir(root_str)
except Exception:
entries = []
return _detect_format(entries)
def _classify_index_payload(data: _IndexData) -> str:
if isinstance(data, dict):
if "chunks" in data and "config" in data:
return "litdata"
if "shards" in data:
# Could be MDS, Parquet, or Vortex index
# Check if shards have Parquet-specific fields in extra
shards = data.get("shards", [])
if shards:
first_shard = shards[0] if isinstance(shards, list) else shards.get(0)
if isinstance(first_shard, dict):
extra = first_shard.get("extra", {})
if isinstance(extra, dict) and "row_groups" in extra:
return "parquet"
return "mds"
return "mds"