File size: 7,498 Bytes
35d483e
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
"""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()