suvradeepp's picture
Publish Tiny Hinglish Turn Detector development preview
35d483e verified
Raw
History Blame Contribute Delete
7.5 kB
"""Lazy resolution of manifest provenance back to raw records.
Parquet has row-group rather than true single-row random access. The resolver
therefore keeps a small LRU of row groups. Iterating a manifest in its original
audit order is efficient; arbitrary random access is intended for evaluation
and inspection rather than full-data training shuffles.
"""
from __future__ import annotations
from bisect import bisect_right
from collections import OrderedDict
from pathlib import Path
from typing import Any, Iterable, Iterator, Mapping, Sequence
from .ingest import DatasetReadError, _import_pyarrow_parquet, iter_local_records
class ManifestResolutionError(DatasetReadError):
"""A manifest row cannot be resolved to its source record."""
class ManifestRecordResolver:
"""Resolve ``source_file``/``source_row`` provenance with bounded caching."""
def __init__(
self,
*,
source_root: str | Path | None = None,
max_cached_row_groups: int = 1,
) -> None:
if max_cached_row_groups < 0:
raise ValueError("max_cached_row_groups cannot be negative")
self.source_root = Path(source_root) if source_root is not None else None
self.max_cached_row_groups = max_cached_row_groups
self._parquet_files: dict[Path, Any] = {}
self._row_group_starts: dict[Path, list[int]] = {}
self._row_group_cache: OrderedDict[tuple[Path, int, tuple[str, ...] | None], Any] = OrderedDict()
self._basename_cache: dict[str, Path] = {}
def clear(self) -> None:
"""Drop cached Parquet handles and row groups."""
self._parquet_files.clear()
self._row_group_starts.clear()
self._row_group_cache.clear()
def _resolve_source_path(self, source_file: Any) -> Path:
if not source_file:
raise ManifestResolutionError("manifest row has no source_file")
source_text = str(source_file)
if source_text.startswith("hf://"):
raise ManifestResolutionError(
"remote Hugging Face provenance is not random-accessible; audit a downloaded snapshot instead"
)
original = Path(source_text).expanduser()
candidates = [original]
if self.source_root is not None:
candidates.extend((self.source_root / original, self.source_root / original.name))
for candidate in candidates:
if candidate.is_file():
return candidate.resolve()
if self.source_root is not None:
cached = self._basename_cache.get(original.name)
if cached is not None:
return cached
matches = sorted(self.source_root.rglob(original.name))
if len(matches) == 1:
resolved = matches[0].resolve()
self._basename_cache[original.name] = resolved
return resolved
if len(matches) > 1:
raise ManifestResolutionError(
f"source basename {original.name!r} is ambiguous below {self.source_root}"
)
raise ManifestResolutionError(f"source file does not exist: {source_text}")
def _parquet_handle(self, path: Path):
handle = self._parquet_files.get(path)
if handle is None:
parquet = _import_pyarrow_parquet()
try:
handle = parquet.ParquetFile(str(path))
except Exception as exc:
raise ManifestResolutionError(f"cannot open Parquet source {path}: {exc}") from exc
self._parquet_files[path] = handle
starts = [0]
running = 0
for index in range(handle.metadata.num_row_groups):
running += handle.metadata.row_group(index).num_rows
starts.append(running)
self._row_group_starts[path] = starts
return handle
def _read_parquet_row(
self,
path: Path,
row_index: int,
columns: Sequence[str] | None,
) -> dict[str, Any]:
handle = self._parquet_handle(path)
starts = self._row_group_starts[path]
if row_index < 0 or row_index >= starts[-1]:
raise ManifestResolutionError(
f"source_row {row_index} is outside [0, {starts[-1]}) for {path}"
)
row_group = bisect_right(starts, row_index) - 1
column_key = tuple(columns) if columns is not None else None
cache_key = (path, row_group, column_key)
table = self._row_group_cache.get(cache_key)
if table is None:
try:
table = handle.read_row_group(row_group, columns=columns)
except Exception as exc:
raise ManifestResolutionError(
f"cannot read row group {row_group} from {path}: {exc}"
) from exc
if self.max_cached_row_groups:
self._row_group_cache[cache_key] = table
self._row_group_cache.move_to_end(cache_key)
while len(self._row_group_cache) > self.max_cached_row_groups:
self._row_group_cache.popitem(last=False)
else:
self._row_group_cache.move_to_end(cache_key)
offset = row_index - starts[row_group]
rows = table.slice(offset, 1).to_pylist()
if not rows:
raise ManifestResolutionError(f"failed to resolve row {row_index} from {path}")
return dict(rows[0])
def resolve(
self,
manifest_row: Mapping[str, Any],
*,
columns: Sequence[str] | None = None,
) -> dict[str, Any]:
"""Load one raw source record referenced by a manifest row."""
path = self._resolve_source_path(manifest_row.get("source_file"))
try:
row_index = int(manifest_row.get("source_row"))
except (TypeError, ValueError) as exc:
raise ManifestResolutionError("manifest source_row is missing or invalid") from exc
if path.suffix.lower() == ".parquet":
record = self._read_parquet_row(path, row_index, columns)
else:
record = {}
for candidate in iter_local_records(path, columns=columns):
if int(candidate.get("__source_row", -1)) == row_index:
record = candidate
break
if not record:
raise ManifestResolutionError(f"failed to resolve row {row_index} from {path}")
record["__source_file"] = str(path)
record["__source_row"] = row_index
return record
def resolve_audio(self, manifest_row: Mapping[str, Any]) -> Any:
"""Resolve only the raw ``audio`` value for a manifest row."""
record = self.resolve(manifest_row, columns=("audio",))
if "audio" not in record:
raise ManifestResolutionError("source record has no audio column")
return record["audio"]
def iter_manifest_records(
rows: Iterable[Mapping[str, Any]],
*,
source_root: str | Path | None = None,
columns: Sequence[str] | None = None,
max_cached_row_groups: int = 1,
) -> Iterator[dict[str, Any]]:
"""Resolve manifest rows lazily; source-order manifests reuse row groups."""
resolver = ManifestRecordResolver(
source_root=source_root,
max_cached_row_groups=max_cached_row_groups,
)
try:
for row in rows:
yield resolver.resolve(row, columns=columns)
finally:
resolver.clear()