This is an automated email from the ASF dual-hosted git repository.

JingsongLi pushed a commit to branch master
in repository https://gitbox.apache.org/repos/asf/paimon.git


The following commit(s) were added to refs/heads/master by this push:
     new 37232889aa [python] Support custom dataset readers in 
PaimonLeRobotDataset (#9785)
37232889aa is described below

commit 37232889aa5e8c267acff28173a19b9e9639c7d8
Author: XiaoHongbo <[email protected]>
AuthorDate: Mon Sep 14 10:44:16 2026 +0800

    [python] Support custom dataset readers in PaimonLeRobotDataset (#9785)
---
 docs/docs/pypaimon/lerobot.md                      |  27 +
 paimon-python/pypaimon/multimodal/__init__.py      |   6 +-
 .../pypaimon/multimodal/lerobot/__init__.py        |   6 +-
 .../pypaimon/multimodal/lerobot/dataset.py         | 616 +++++++++++++--------
 .../pypaimon/multimodal/lerobot/reader.py          | 187 +++++++
 .../pypaimon/tests/multimodal_lerobot_test.py      | 161 +++++-
 6 files changed, 762 insertions(+), 241 deletions(-)

diff --git a/docs/docs/pypaimon/lerobot.md b/docs/docs/pypaimon/lerobot.md
index b1e487fff1..99110cdb5d 100644
--- a/docs/docs/pypaimon/lerobot.md
+++ b/docs/docs/pypaimon/lerobot.md
@@ -226,3 +226,30 @@ Without `tag_name`, the latest snapshots are used. Frame 
lookups use the BTree
 on `index`; payloads remain lazy. Video decoding prefers TorchCodec, falls back
 to PyAV, and reuses a bounded decoder cache. Set `video_backend` to force
 either decoder.
+
+Subclass `PaimonDatasetReader` for a custom logical frame layout:
+
+```python
+from pypaimon.multimodal import PaimonDatasetReader, PaimonLeRobotDataset
+
+class CustomDatasetReader(PaimonDatasetReader):
+    def __init__(self, metadata, source, **kwargs):
+        self._source = source
+        super().__init__(
+            metadata,
+            file_io=getattr(source, "file_io", None),
+            **kwargs,
+        )
+
+    def read_indices(self, indices, columns):
+        return self._source.read_indices(indices, columns)
+
+reader = CustomDatasetReader(
+    metadata, source, delta_timestamps=delta_timestamps
+)
+dataset = PaimonLeRobotDataset(reader)
+```
+
+`source` exposes `read_indices` and optional `file_io`. Pass
+`schema=source.schema` for eager schema validation. `PaimonDatasetReader`
+reuses the standard Episode, delta-window, media, and Torch handling.
diff --git a/paimon-python/pypaimon/multimodal/__init__.py 
b/paimon-python/pypaimon/multimodal/__init__.py
index 00d59d5c71..25cfc784c5 100644
--- a/paimon-python/pypaimon/multimodal/__init__.py
+++ b/paimon-python/pypaimon/multimodal/__init__.py
@@ -29,7 +29,10 @@ from pypaimon.multimodal.hdf5 import (
     Hdf5File,
     Hdf5LoadResult,
 )
-from pypaimon.multimodal.lerobot.dataset import PaimonLeRobotDataset
+from pypaimon.multimodal.lerobot.dataset import (
+    PaimonDatasetReader,
+    PaimonLeRobotDataset,
+)
 from pypaimon.multimodal.rosbag import (
     RosbagLoadResult,
     RosbagSource,
@@ -67,6 +70,7 @@ __all__ = [
     "MultimodalTable",
     "NoSuchKey",
     "ObjectInfo",
+    "PaimonDatasetReader",
     "PaimonLeRobotDataset",
     "PutObjectResult",
     "RosbagLoadResult",
diff --git a/paimon-python/pypaimon/multimodal/lerobot/__init__.py 
b/paimon-python/pypaimon/multimodal/lerobot/__init__.py
index 25e196dbb8..482184a52b 100644
--- a/paimon-python/pypaimon/multimodal/lerobot/__init__.py
+++ b/paimon-python/pypaimon/multimodal/lerobot/__init__.py
@@ -17,11 +17,15 @@
 """LeRobot Dataset v3 integration for multimodal Paimon tables."""
 
 from pypaimon.multimodal.lerobot.api import load_from_lerobot
-from pypaimon.multimodal.lerobot.dataset import PaimonLeRobotDataset
+from pypaimon.multimodal.lerobot.dataset import (
+    PaimonDatasetReader,
+    PaimonLeRobotDataset,
+)
 from pypaimon.multimodal.lerobot.writer import PaimonLeRobotWriter
 
 
 __all__ = [
+    "PaimonDatasetReader",
     "PaimonLeRobotDataset",
     "PaimonLeRobotWriter",
     "load_from_lerobot",
diff --git a/paimon-python/pypaimon/multimodal/lerobot/dataset.py 
b/paimon-python/pypaimon/multimodal/lerobot/dataset.py
index 7a13731e5f..6abbd34bab 100644
--- a/paimon-python/pypaimon/multimodal/lerobot/dataset.py
+++ b/paimon-python/pypaimon/multimodal/lerobot/dataset.py
@@ -22,14 +22,14 @@ import io
 import json
 import math
 import operator
-import os
 import sys
+from abc import ABC, abstractmethod
 from collections import OrderedDict
+from collections.abc import Mapping
 from functools import partial
 
 import pyarrow as pa
 
-from pypaimon.common.options.core_options import CoreOptions
 from pypaimon.multimodal.lerobot.metadata import (
     _companion_table_identifiers,
     _restore_pandas_metadata,
@@ -37,6 +37,7 @@ from pypaimon.multimodal.lerobot.metadata import (
     _validate_tag_name,
 )
 from pypaimon.multimodal.lerobot.loader import _DECLARED_NUMERIC_RANGES
+from pypaimon.multimodal.lerobot.reader import _PaimonTableFrameReader
 from pypaimon.multimodal.lerobot.schema import (
     _feature_shape,
     _require_v3,
@@ -45,7 +46,6 @@ from pypaimon.multimodal.lerobot.schema import (
 )
 from pypaimon.multimodal.table import _target_schema, _time_travel_table
 from pypaimon.multimodal.video import VideoFrameCollator
-from pypaimon.read.query_auth_split import QueryAuthSplit
 
 
 _TORCH_DTYPE_NAMES = {
@@ -75,11 +75,12 @@ _CONTROL_FEATURES = frozenset({
 })
 
 
-class PaimonLeRobotDataset:
-    """Map-style LeRobot reader backed by indexed Paimon reads.
+class PaimonDatasetReader(ABC):
+    """Read-side implementation for Paimon-backed LeRobot datasets.
 
-    LeRobot metadata is resolved from the Paimon table group and remains
-    available through :attr:`meta`.
+    Subclasses provide batched ``read_indices``. Resolved LeRobot metadata
+    remains available through :attr:`meta`. Readers must be picklable for
+    DataLoader workers.
 
     Set ``return_uint8=True`` to keep 8-bit visual frames in their decoded
     ``torch.uint8`` representation instead of normalizing them to float32.
@@ -88,23 +89,68 @@ class PaimonLeRobotDataset:
 
     def __init__(
             self,
-            table,
+            meta,
             *,
-            tag_name=None,
+            schema=None,
+            file_io=None,
             episodes=None,
             image_transforms=None,
             delta_timestamps=None,
             tolerance_s=1e-4,
             blob_parallelism=16,
             video_backend=None,
-            return_uint8=False):
-        if sys.version_info < (3, 10):
-            raise RuntimeError(
-                "PaimonLeRobotDataset requires Python 3.10 or newer; "
-                "install and run 'pypaimon[lerobot]' on a supported Python "
-                "version.")
-        raw_table, self.meta = _load_dataset(table, tag_name)
-        self.tag_name = tag_name
+            return_uint8=False,
+            _resolved_meta=False):
+        _require_dataset_python()
+        metadata = meta if _resolved_meta else _reader_metadata(meta)
+        if schema is not None and not isinstance(schema, pa.Schema):
+            raise TypeError(
+                "PaimonDatasetReader schema must be a pyarrow.Schema.")
+        self.file_io = file_io
+        info = self._init_dataset(
+            metadata,
+            episodes,
+            image_transforms,
+            delta_timestamps,
+            tolerance_s,
+            blob_parallelism,
+            video_backend,
+            return_uint8,
+        )
+        schema = schema if schema is not None else _schema_from_info(info)
+        self.schema = schema
+        projection, validation_context, subtasks = \
+            self._init_frame_contract(
+                schema, info, self._validate_physical_metadata())
+        rows = self._open_frame_rows(projection)
+        self._set_frame_rows(
+            rows, projection, validation_context, subtasks)
+
+    @abstractmethod
+    def read_indices(self, indices, columns):
+        """Return one row per requested absolute index as a PyArrow Table.
+
+        ``indices`` are unique; result order is unrestricted. Result columns
+        must match ``schema`` and contain every requested index exactly once.
+        """
+
+    def _validate_physical_metadata(self):
+        return False
+
+    def _open_frame_rows(self, projection):
+        return None
+
+    def _init_dataset(
+            self,
+            metadata,
+            episodes,
+            image_transforms,
+            delta_timestamps,
+            tolerance_s,
+            blob_parallelism,
+            video_backend,
+            return_uint8):
+        self.meta = metadata
         self.repo_id = self.meta.repo_id
         self.image_transforms = image_transforms
         self.delta_timestamps = delta_timestamps
@@ -125,7 +171,7 @@ class PaimonLeRobotDataset:
 
         info = self._init_metadata()
         self._init_episodes(episodes)
-        self._init_reader(raw_table, info)
+        return info
 
     def _init_metadata(self):
         info = dict(_metadata_member(self.meta, "info", {}))
@@ -169,6 +215,11 @@ class PaimonLeRobotDataset:
     def _init_episodes(self, episodes):
         self._episode_ranges = _episode_ranges(
             self.meta, self._total_frames, self._total_episodes)
+        if (self._episode_ranges is None
+                and (self._total_frames or self._total_episodes)):
+            raise ValueError(
+                "LeRobot metadata must define episodes for a non-empty "
+                "dataset.")
         self._episode_ends = [end for _, end in self._episode_ranges] \
             if self._episode_ranges is not None else None
         self.episodes = _selected_episodes(episodes, self._total_episodes)
@@ -197,15 +248,45 @@ class PaimonLeRobotDataset:
         if self._delta_indices and self._episode_ranges is None:
             raise ValueError("delta_timestamps requires episode metadata.")
 
-    def _init_reader(self, raw_table, info):
-        target_schema = _target_schema(raw_table)
-        table_fields = set(target_schema.names)
+    def _set_frame_rows(
+            self, rows, projection, validation_context, subtasks):
+        self._frame_rows = rows
+        access = rows if rows is not None else self
+        self._snapshot_id = getattr(access, "snapshot_id", None)
+        self._read_table = getattr(access, "_table", None)
+        self._frame_locator = getattr(access, "_locator", None)
+        self._projection = projection
+        self._validation_context = validation_context
+        self._file_io = getattr(access, "file_io", None)
+        if self._video_keys and self._file_io is None:
+            raise ValueError(
+                "A video-backed PaimonDatasetReader must expose file_io.")
+        self._video_collators = [
+            VideoFrameCollator(
+                access,
+                video_column=key,
+                decoder_factory=partial(
+                    _open_video_decoder, backend=self.video_backend),
+                decode_fn=_decode_video_frame,
+                output_column=key,
+                collate_fn=_identity,
+            )
+            for key in self._video_keys
+        ]
+        self._init_delta_projection(validation_context, subtasks)
+
+    def _init_frame_contract(self, target_schema, info, validate_metadata):
         tasks = _metadata_member(self.meta, "tasks")
         subtasks = _metadata_member(self.meta, "subtasks")
         _validate_component_metadata(
             self._features, self._total_tasks, tasks, subtasks)
         source_schema = _schema_from_info(info)
-        _validate_lerobot_schema(source_schema, target_schema, self.repo_id)
+        if validate_metadata:
+            _validate_lerobot_schema(
+                source_schema, target_schema, self.repo_id)
+        else:
+            _validate_reader_schema(
+                source_schema, target_schema, self.repo_id)
         validation_context = _build_frame_validation_context(
             self.meta,
             self._episode_ranges,
@@ -215,37 +296,14 @@ class PaimonLeRobotDataset:
             source_schema.field("timestamp").type,
         )
         projection = list(self._features)
-        missing = set(projection) - table_fields
+        missing = set(projection) - set(target_schema.names)
         if missing:
             raise ValueError(
-                "Paimon table is missing LeRobot fields: %s"
+                "LeRobot frame schema is missing fields: %s"
                 % sorted(missing))
+        return projection, validation_context, subtasks
 
-        self._read_table, self._snapshot_id, splits = _indexed_read_table(
-            raw_table, projection)
-        snapshot = self._read_table.snapshot_manager().get_snapshot_by_id(
-            self._snapshot_id)
-        if snapshot.next_row_id != self._total_frames:
-            raise ValueError(
-                "Paimon table has %d rows but metadata declares %d frames."
-                % (snapshot.next_row_id, self._total_frames))
-        self._projection = projection
-        self._frame_locator = _FrameLocator(
-            self._read_table, snapshot, splits)
-        self._validation_context = validation_context
-        self._file_io = self._read_table.file_io
-        self._video_collators = [
-            VideoFrameCollator(
-                self._read_table,
-                video_column=key,
-                decoder_factory=partial(
-                    _open_video_decoder, backend=self.video_backend),
-                decode_fn=_decode_video_frame,
-                output_column=key,
-                collate_fn=_identity,
-            )
-            for key in self._video_keys
-        ]
+    def _init_delta_projection(self, validation_context, subtasks):
         self._task_names = validation_context["task_names"]
         self._subtask_names = validation_context["subtask_names"]
         self._delta_projection = None
@@ -281,12 +339,24 @@ class PaimonLeRobotDataset:
     def __len__(self):
         return self.num_frames
 
-    def __getitem__(self, index):
-        if isinstance(index, slice):
-            return self.__getitems__(range(*index.indices(len(self))))
-        return self.__getitems__([index])[0]
+    @property
+    def absolute_to_relative_idx(self):
+        if self._selected_ranges is None:
+            return None
+        result = {}
+        relative = 0
+        for begin, end in self._selected_ranges:
+            for absolute in range(begin, end):
+                result[absolute] = relative
+                relative += 1
+        return result
 
-    def __getitems__(self, indices):
+    def get_item(self, index):
+        """Return one fully assembled frame."""
+        return self.get_items([index])[0]
+
+    def get_items(self, indices):
+        """Return fully assembled frames for one batch."""
         dataset_indices = [
             _normalize_index(index, len(self)) for index in indices
         ]
@@ -307,9 +377,7 @@ class PaimonLeRobotDataset:
             if position not in unique_frame_index_set
         })
         lookup_indices = sorted(unique_frame_index_set.union(delta_indices))
-        splits, needs_filter = self._frame_locator.locate(lookup_indices)
-        rows = self._read_rows(
-            lookup_indices, self._projection, splits, needs_filter)
+        rows = self._read_rows(lookup_indices, self._projection)
         base_rows = {
             index: rows[index] for index in unique_frame_indices
         }
@@ -323,28 +391,36 @@ class PaimonLeRobotDataset:
         _attach_task_labels(
             base_rows, self._task_names, self._subtask_names)
         row_groups = [base_rows, delta_rows]
-        image_sources = _image_blob_sources(
-            row_groups, self._image_keys)
-        for attempt in range(_IMAGE_READ_ATTEMPTS):
-            if attempt:
-                _restore_image_blob_sources(image_sources)
-            try:
-                _resolve_image_blobs(
-                    self._file_io,
-                    row_groups,
-                    self._image_keys,
-                    self.blob_parallelism,
-                )
-                _decode_image_rows(
-                    row_groups,
-                    self._image_keys,
-                    self._features,
-                    self.return_uint8,
-                )
-                break
-            except OSError:
-                if attempt + 1 == _IMAGE_READ_ATTEMPTS:
-                    raise
+        if self._file_io is not None:
+            image_sources = _image_blob_sources(
+                row_groups, self._image_keys)
+            for attempt in range(_IMAGE_READ_ATTEMPTS):
+                if attempt:
+                    _restore_image_blob_sources(image_sources)
+                try:
+                    _resolve_image_blobs(
+                        self._file_io,
+                        row_groups,
+                        self._image_keys,
+                        self.blob_parallelism,
+                    )
+                    _decode_image_rows(
+                        row_groups,
+                        self._image_keys,
+                        self._features,
+                        self.return_uint8,
+                    )
+                    break
+                except OSError:
+                    if attempt + 1 == _IMAGE_READ_ATTEMPTS:
+                        raise
+        else:
+            _decode_image_rows(
+                row_groups,
+                self._image_keys,
+                self._features,
+                self.return_uint8,
+            )
 
         _decode_video_rows(
             row_groups, getattr(self, "_video_collators", ()))
@@ -380,20 +456,30 @@ class PaimonLeRobotDataset:
             result.append(item)
         return result
 
+    def __getitem__(self, index):
+        if isinstance(index, slice):
+            return self.get_items(range(*index.indices(len(self))))
+        return self.get_item(index)
+
+    def __getitems__(self, indices):
+        return self.get_items(indices)
+
     def close(self):
         first_error = None
-        locator = getattr(self, "_frame_locator", None)
-        if locator is not None:
-            try:
-                locator.close()
-            except Exception as error:
-                first_error = error
         for collator in getattr(self, "_video_collators", ()):
             try:
                 collator.close()
             except Exception as error:
                 if first_error is None:
                     first_error = error
+        rows = getattr(self, "_frame_rows", None)
+        self._frame_rows = None
+        if rows is not None:
+            try:
+                rows.close()
+            except Exception as error:
+                if first_error is None:
+                    first_error = error
         if first_error is not None:
             raise first_error
 
@@ -403,19 +489,18 @@ class PaimonLeRobotDataset:
         except Exception:
             pass
 
-    def _read_rows(
-            self, indices, projection, splits=None, needs_filter=True):
+    def _read_rows(self, indices, projection):
         if not indices:
             return {}
-        return _read_rows_by_index(
-            self._read_table,
+        rows = self._frame_rows if self._frame_rows is not None else self
+        return _read_reader_rows(
+            rows,
             projection,
             indices,
+            self.schema,
             self._validation_context,
             self.tolerance_s,
             self._features,
-            splits,
-            needs_filter,
         )
 
     def set_image_transforms(self, image_transforms):
@@ -458,113 +543,141 @@ class PaimonLeRobotDataset:
                self.num_frames, list(self.features)))
 
 
-class _FrameLocator:
-    """Locate LeRobot frame rows in one fixed Paimon snapshot."""
-
-    def __init__(self, table, snapshot, splits):
-        self._table = table
-        self._snapshot = snapshot
-        self._scanner = None
-        self._scanner_initialized = False
-        self._process_id = os.getpid()
-        self._set_splits(splits)
+class _PaimonTableDatasetReader(PaimonDatasetReader):
 
-    def _set_splits(self, splits):
-        from pypaimon.read.datasource.torch_dataset import (
-            SplitRangeIndex,
-            row_ranges_for_split,
+    def __init__(
+            self,
+            table,
+            *,
+            tag_name=None,
+            episodes=None,
+            image_transforms=None,
+            delta_timestamps=None,
+            tolerance_s=1e-4,
+            blob_parallelism=16,
+            video_backend=None,
+            return_uint8=False):
+        self._frames_table, meta = _load_dataset(table, tag_name)
+        self.tag_name = tag_name
+        super().__init__(
+            meta,
+            schema=_target_schema(self._frames_table),
+            episodes=episodes,
+            image_transforms=image_transforms,
+            delta_timestamps=delta_timestamps,
+            tolerance_s=tolerance_s,
+            blob_parallelism=blob_parallelism,
+            video_backend=video_backend,
+            return_uint8=return_uint8,
+            _resolved_meta=True,
         )
 
-        self._splits = splits
-        self._split_ranges = [
-            row_ranges_for_split(split) for split in splits
-        ]
-        self._split_range_index = SplitRangeIndex(self._split_ranges)
+    def read_indices(self, indices, columns):
+        return self._frame_rows.read_indices(indices, columns)
 
-    def locate(self, indices):
-        """Return narrowed splits and whether rows still need filtering."""
-        self._ensure_process()
-        predicate = _index_predicate(self._table, indices)
-        try:
-            scanner = self._index_scanner(predicate)
-        except Exception as error:
-            raise RuntimeError(
-                "Failed to open the Paimon global index for LeRobot frame "
-                "lookups.") from error
-        if scanner is None:
-            raise RuntimeError(
-                "PaimonLeRobotDataset requires a readable global index on "
-                "the frame 'index' column.")
-        try:
-            evaluation = scanner.scan_with_coverage(predicate)
-            if evaluation is None:
-                raise RuntimeError(
-                    "The Paimon global index could not evaluate the LeRobot "
-                    "frame index predicate.")
-            unindexed = scanner.unindexed_ranges(
-                predicate,
-                search_mode=self._table.options.scalar_index_search_mode(),
-                contributing_field_ids=evaluation.contributing_field_ids,
-            )
-            ranges = evaluation.result.results().to_range_list() + unindexed
-            from pypaimon.read.datasource.torch_dataset import (
-                select_indexed_splits,
-            )
-            from pypaimon.utils.range import Range
-            return select_indexed_splits(
-                self._splits,
-                self._split_ranges,
-                self._split_range_index,
-                Range.sort_and_merge_overlap(ranges, True),
-            ), bool(unindexed)
-        except RuntimeError:
-            raise
-        except Exception as error:
-            raise RuntimeError(
-                "Failed to query the Paimon global index for LeRobot "
-                "frames.") from error
-
-    def _ensure_process(self):
-        process_id = os.getpid()
-        if process_id == self._process_id:
-            return
-        self._scanner = None
-        self._scanner_initialized = False
-        self._set_splits(self._splits)
-        self._process_id = process_id
-
-    def _index_scanner(self, predicate):
-        if not self._scanner_initialized:
-            from pypaimon.globalindex import DataEvolutionGlobalIndexScanner
-            self._scanner = DataEvolutionGlobalIndexScanner.create(
-                self._table,
-                predicate=predicate,
-                snapshot=self._snapshot,
+    def _validate_physical_metadata(self):
+        return True
+
+    def _open_frame_rows(self, projection):
+        rows = _PaimonTableFrameReader(
+            self._frames_table, columns=projection)
+        if rows.num_rows != self._total_frames:
+            raise ValueError(
+                "Paimon table has %d rows but metadata declares %d frames."
+                % (rows.num_rows, self._total_frames))
+        return rows
+
+
+class PaimonLeRobotDataset:
+    """Map-style Dataset facade backed by :class:`PaimonDatasetReader`."""
+
+    def __init__(
+            self,
+            table,
+            *,
+            tag_name=None,
+            episodes=None,
+            image_transforms=None,
+            delta_timestamps=None,
+            tolerance_s=1e-4,
+            blob_parallelism=16,
+            video_backend=None,
+            return_uint8=False):
+        _require_dataset_python()
+        if isinstance(table, PaimonDatasetReader):
+            if (
+                tag_name is not None
+                or episodes is not None
+                or image_transforms is not None
+                or delta_timestamps is not None
+                or tolerance_s != 1e-4
+                or blob_parallelism != 16
+                or video_backend is not None
+                or return_uint8
+            ):
+                raise ValueError(
+                    "Configure Dataset options on PaimonDatasetReader.")
+            self.reader = table
+        else:
+            self.reader = _PaimonTableDatasetReader(
+                table,
+                tag_name=tag_name,
+                episodes=episodes,
+                image_transforms=image_transforms,
+                delta_timestamps=delta_timestamps,
+                tolerance_s=tolerance_s,
+                blob_parallelism=blob_parallelism,
+                video_backend=video_backend,
+                return_uint8=return_uint8,
             )
-            self._scanner_initialized = True
-        return self._scanner
+
+    def __len__(self):
+        return len(self.reader)
+
+    def __getitem__(self, index):
+        return self.reader[index]
+
+    def __getitems__(self, indices):
+        return self.reader.get_items(indices)
+
+    @property
+    def return_uint8(self):
+        return self.reader.return_uint8
+
+    @return_uint8.setter
+    def return_uint8(self, value):
+        if not isinstance(value, bool):
+            raise TypeError("return_uint8 must be a boolean.")
+        self.reader.return_uint8 = value
+
+    @property
+    def image_transforms(self):
+        return self.reader.image_transforms
+
+    @image_transforms.setter
+    def image_transforms(self, value):
+        self.reader.set_image_transforms(value)
+
+    def set_image_transforms(self, image_transforms):
+        self.reader.set_image_transforms(image_transforms)
+
+    def clear_image_transforms(self):
+        self.reader.clear_image_transforms()
 
     def close(self):
-        scanner = self._scanner
-        self._scanner = None
-        self._scanner_initialized = False
-        if scanner is not None and self._process_id == os.getpid():
-            scanner.close()
-
-    def __getstate__(self):
-        state = self.__dict__.copy()
-        state["_scanner"] = None
-        state["_scanner_initialized"] = False
-        state["_process_id"] = None
-        state["_split_ranges"] = None
-        state["_split_range_index"] = None
-        return state
+        self.reader.close()
 
-    def __del__(self):
-        try:
-            self.close()
-        except Exception:
-            pass
+    def __getattr__(self, name):
+        if name.startswith("__") and name.endswith("__"):
+            raise AttributeError(name)
+        reader = self.__dict__.get("reader")
+        if reader is None:
+            raise AttributeError(name)
+        return getattr(reader, name)
+
+    def __repr__(self):
+        return repr(self.reader).replace(
+            self.reader.__class__.__name__, self.__class__.__name__, 1)
 
 
 class _PaimonLeRobotMetadata:
@@ -623,9 +736,40 @@ class _PaimonLeRobotMetadata:
         }
 
     def get_task_index(self, task):
-        if task not in self.tasks.index:
+        if hasattr(self.tasks, "loc"):
+            if task not in self.tasks.index:
+                return None
+            return int(self.tasks.loc[task].task_index)
+        try:
+            return list(self.tasks).index(task)
+        except ValueError:
             return None
-        return int(self.tasks.loc[task].task_index)
+
+
+def _reader_metadata(metadata):
+    info = _metadata_member(metadata, "info")
+    if not isinstance(info, Mapping):
+        raise TypeError(
+            "PaimonDatasetReader metadata must contain an info map.")
+    info = dict(info)
+    features = info.get("features")
+    if isinstance(features, Mapping):
+        info["features"] = {
+            name: dict(feature) for name, feature in features.items()
+        }
+        for feature in info["features"].values():
+            if "shape" in feature:
+                feature["shape"] = tuple(feature["shape"])
+    stats = _metadata_member(metadata, "stats")
+    return _PaimonLeRobotMetadata(
+        str(_metadata_member(metadata, "repo_id", "custom-reader")),
+        _metadata_member(metadata, "revision"),
+        info,
+        _numpy_stats(stats) if stats is not None else None,
+        _metadata_member(metadata, "episodes"),
+        _metadata_member(metadata, "tasks"),
+        _metadata_member(metadata, "subtasks"),
+    )
 
 
 def _load_dataset(table, tag_name):
@@ -729,7 +873,8 @@ def _numpy_stats(value):
 
 
 def _metadata_member(metadata, name, default=None):
-    value = getattr(metadata, name, None)
+    value = metadata.get(name) if isinstance(metadata, Mapping) \
+        else getattr(metadata, name, None)
     return default if value is None else value
 
 
@@ -905,59 +1050,54 @@ def _delta_indices(delta_timestamps, fps, tolerance_s, 
features):
     return result
 
 
-def _indexed_read_table(raw_table, projection):
-    read_table = raw_table.copy({
-        CoreOptions.BLOB_AS_DESCRIPTOR.key(): "true"
-    })
-    plan = read_table.new_read_builder().with_projection(
-        projection).new_scan().plan()
-    splits = plan.splits()
-    if any(
-            isinstance(split, QueryAuthSplit)
-            and (
-                getattr(split.auth_result, "filter", None)
-                or getattr(split.auth_result, "column_masking", None)
-            )
-            for split in splits):
+def _validate_reader_schema(expected_schema, actual_schema, source):
+    for expected_field in expected_schema:
+        target_index = actual_schema.get_field_index(expected_field.name)
+        if target_index < 0:
+            continue
+        target_type = actual_schema.field(target_index).type
+        if expected_field.type != target_type:
+            raise ValueError(
+                "LeRobot feature %s from %s expects %s, found %s."
+                % (expected_field.name, source, expected_field.type,
+                   target_type))
+
+
+def _read_reader_rows(
+        reader, projection, indices, expected_schema, validation_context,
+        tolerance_s, features):
+    values = reader.read_indices(tuple(indices), tuple(projection))
+    if not isinstance(values, pa.Table):
+        raise TypeError(
+            "PaimonDatasetReader.read_indices() must return a pyarrow.Table.")
+    missing = set(projection) - set(values.column_names)
+    if missing:
         raise ValueError(
-            "PaimonLeRobotDataset does not support query authorization "
-            "filters or column masking.")
-    if plan.snapshot_id is None:
-        raise ValueError("Paimon LeRobot frames table has no snapshot.")
-    if read_table.options.scan_tag_name() is None:
-        read_table = _time_travel_table(
-            read_table, snapshot_id=plan.snapshot_id)
-    return read_table, plan.snapshot_id, splits
-
-
-def _index_predicate(table, indices):
-    return table.new_read_builder().new_predicate_builder().is_in(
-        "index", indices)
-
-
-def _read_rows_by_index(
-        table, projection, indices, validation_context, tolerance_s, features,
-        splits=None, needs_filter=True):
-    builder = table.new_read_builder().with_projection(projection)
-    if needs_filter:
-        builder = builder.with_filter(_index_predicate(table, indices))
-    if splits is None:
-        splits = builder.new_scan().plan().splits()
-    rows = _arrow_rows(builder.new_read().to_arrow(splits), features)
+            "PaimonDatasetReader result is missing fields: %s"
+            % sorted(missing))
+    for name in projection:
+        expected_type = expected_schema.field(name).type
+        actual_type = values.schema.field(name).type
+        if actual_type != expected_type:
+            raise ValueError(
+                "PaimonDatasetReader field %s expects %s, found %s."
+                % (name, expected_type, actual_type))
+    rows = _arrow_rows(values.select(projection), features)
     expected = set(indices)
     result = {}
     for row in rows:
         index = _control_index(row, "index", -1)
         if index not in expected or index in result:
             raise ValueError(
-                "Paimon BTree returned an unexpected or duplicate LeRobot "
+                "PaimonDatasetReader returned an unexpected or duplicate "
                 "index: %d." % index)
-        _validate_control_row(index, row, validation_context, tolerance_s)
+        _validate_control_row(
+            index, row, validation_context, tolerance_s)
         result[index] = row
     missing = expected - set(result)
     if missing:
         raise RuntimeError(
-            "Paimon index lookup did not return LeRobot indices %s."
+            "PaimonDatasetReader did not return indices %s."
             % sorted(missing))
     return result
 
@@ -1428,6 +1568,14 @@ def _normalize_index(index, size):
     return index
 
 
+def _require_dataset_python():
+    if sys.version_info < (3, 10):
+        raise RuntimeError(
+            "PaimonLeRobotDataset requires Python 3.10 or newer; "
+            "install and run 'pypaimon[lerobot]' on a supported Python "
+            "version.")
+
+
 def _positive_int(value, name):
     try:
         value = operator.index(value)
diff --git a/paimon-python/pypaimon/multimodal/lerobot/reader.py 
b/paimon-python/pypaimon/multimodal/lerobot/reader.py
new file mode 100644
index 0000000000..a63945fd86
--- /dev/null
+++ b/paimon-python/pypaimon/multimodal/lerobot/reader.py
@@ -0,0 +1,187 @@
+# Licensed to the Apache Software Foundation (ASF) under one
+# or more contributor license agreements.  See the NOTICE file
+# distributed with this work for additional information
+# regarding copyright ownership.  The ASF licenses this file
+# to you under the Apache License, Version 2.0 (the
+# "License"); you may not use this file except in compliance
+# with the License.  You may obtain a copy of the License at
+#
+#   http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing,
+# software distributed under the License is distributed on an
+# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
+# KIND, either express or implied.  See the License for the
+# specific language governing permissions and limitations
+# under the License.
+
+"""Indexed frame-row reader used by :class:`PaimonDatasetReader`."""
+
+import os
+
+from pypaimon.common.options.core_options import CoreOptions
+from pypaimon.multimodal.table import _time_travel_table
+from pypaimon.read.query_auth_split import QueryAuthSplit
+
+
+class _PaimonTableFrameReader:
+    """Read logical frame rows from one indexed Paimon table."""
+
+    def __init__(self, frames_table, *, columns):
+        self._table, self.snapshot_id, splits = _indexed_read_table(
+            frames_table, columns)
+        snapshot = self._table.snapshot_manager().get_snapshot_by_id(
+            self.snapshot_id)
+        self.num_rows = snapshot.next_row_id
+        self.file_io = self._table.file_io
+        self._locator = _FrameLocator(self._table, snapshot, splits)
+
+    def read_indices(self, indices, columns):
+        splits, needs_filter = self._locator.locate(indices)
+        builder = self._table.new_read_builder().with_projection(columns)
+        if needs_filter:
+            builder = builder.with_filter(
+                _index_predicate(self._table, indices))
+        return builder.new_read().to_arrow(splits)
+
+    def close(self):
+        self._locator.close()
+
+
+class _FrameLocator:
+    """Locate LeRobot frame rows in one fixed Paimon snapshot."""
+
+    def __init__(self, table, snapshot, splits):
+        self._table = table
+        self._snapshot = snapshot
+        self._scanner = None
+        self._scanner_initialized = False
+        self._process_id = os.getpid()
+        self._set_splits(splits)
+
+    def _set_splits(self, splits):
+        from pypaimon.read.datasource.torch_dataset import (
+            SplitRangeIndex,
+            row_ranges_for_split,
+        )
+
+        self._splits = splits
+        self._split_ranges = [
+            row_ranges_for_split(split) for split in splits
+        ]
+        self._split_range_index = SplitRangeIndex(self._split_ranges)
+
+    def locate(self, indices):
+        """Return narrowed splits and whether rows still need filtering."""
+        self._ensure_process()
+        predicate = _index_predicate(self._table, indices)
+        try:
+            scanner = self._index_scanner(predicate)
+        except Exception as error:
+            raise RuntimeError(
+                "Failed to open the Paimon global index for LeRobot frame "
+                "lookups.") from error
+        if scanner is None:
+            raise RuntimeError(
+                "PaimonLeRobotDataset requires a readable global index on "
+                "the frame 'index' column.")
+        try:
+            evaluation = scanner.scan_with_coverage(predicate)
+            if evaluation is None:
+                raise RuntimeError(
+                    "The Paimon global index could not evaluate the LeRobot "
+                    "frame index predicate.")
+            unindexed = scanner.unindexed_ranges(
+                predicate,
+                search_mode=self._table.options.scalar_index_search_mode(),
+                contributing_field_ids=evaluation.contributing_field_ids,
+            )
+            ranges = evaluation.result.results().to_range_list() + unindexed
+            from pypaimon.read.datasource.torch_dataset import (
+                select_indexed_splits,
+            )
+            from pypaimon.utils.range import Range
+            return select_indexed_splits(
+                self._splits,
+                self._split_ranges,
+                self._split_range_index,
+                Range.sort_and_merge_overlap(ranges, True),
+            ), bool(unindexed)
+        except RuntimeError:
+            raise
+        except Exception as error:
+            raise RuntimeError(
+                "Failed to query the Paimon global index for LeRobot "
+                "frames.") from error
+
+    def _ensure_process(self):
+        process_id = os.getpid()
+        if process_id == self._process_id:
+            return
+        self._scanner = None
+        self._scanner_initialized = False
+        self._set_splits(self._splits)
+        self._process_id = process_id
+
+    def _index_scanner(self, predicate):
+        if not self._scanner_initialized:
+            from pypaimon.globalindex import DataEvolutionGlobalIndexScanner
+            self._scanner = DataEvolutionGlobalIndexScanner.create(
+                self._table,
+                predicate=predicate,
+                snapshot=self._snapshot,
+            )
+            self._scanner_initialized = True
+        return self._scanner
+
+    def close(self):
+        scanner = self._scanner
+        self._scanner = None
+        self._scanner_initialized = False
+        if scanner is not None and self._process_id == os.getpid():
+            scanner.close()
+
+    def __getstate__(self):
+        state = self.__dict__.copy()
+        state["_scanner"] = None
+        state["_scanner_initialized"] = False
+        state["_process_id"] = None
+        state["_split_ranges"] = None
+        state["_split_range_index"] = None
+        return state
+
+    def __del__(self):
+        try:
+            self.close()
+        except Exception:
+            pass
+
+
+def _indexed_read_table(raw_table, projection):
+    read_table = raw_table.copy({
+        CoreOptions.BLOB_AS_DESCRIPTOR.key(): "true"
+    })
+    plan = read_table.new_read_builder().with_projection(
+        projection).new_scan().plan()
+    splits = plan.splits()
+    if any(
+            isinstance(split, QueryAuthSplit)
+            and (
+                getattr(split.auth_result, "filter", None)
+                or getattr(split.auth_result, "column_masking", None)
+            )
+            for split in splits):
+        raise ValueError(
+            "PaimonLeRobotDataset does not support query authorization "
+            "filters or column masking.")
+    if plan.snapshot_id is None:
+        raise ValueError("Paimon LeRobot frames table has no snapshot.")
+    if read_table.options.scan_tag_name() is None:
+        read_table = _time_travel_table(
+            read_table, snapshot_id=plan.snapshot_id)
+    return read_table, plan.snapshot_id, splits
+
+
+def _index_predicate(table, indices):
+    return table.new_read_builder().new_predicate_builder().is_in(
+        "index", indices)
diff --git a/paimon-python/pypaimon/tests/multimodal_lerobot_test.py 
b/paimon-python/pypaimon/tests/multimodal_lerobot_test.py
index 97b180d39e..a9790b4e16 100644
--- a/paimon-python/pypaimon/tests/multimodal_lerobot_test.py
+++ b/paimon-python/pypaimon/tests/multimodal_lerobot_test.py
@@ -97,6 +97,21 @@ except ImportError:
     av = None
 
 
+class _ManualDatasetReader(pmm.PaimonDatasetReader):
+
+    def read_indices(self, indices, columns):
+        raise NotImplementedError
+
+
+class _PickleDatasetReader(_ManualDatasetReader):
+
+    def __init__(self, value):
+        self.value = value
+
+    def __getstate__(self):
+        return {"value": self.value}
+
+
 def _replaced_contract(field, old, new):
     description = field.metadata[b"description"].decode("utf-8")
     if old not in description:
@@ -325,6 +340,141 @@ class LeRobotValidationTest(unittest.TestCase):
                 pmm.PaimonLeRobotDataset(Mock())
             load.assert_not_called()
 
+    def test_dataset_reads_one_batch_from_custom_reader(self):
+        try:
+            import torch
+        except ImportError as error:
+            self.skipTest(str(error))
+
+        info = {
+            "codebase_version": "v3.0",
+            "total_frames": 3,
+            "total_episodes": 1,
+            "total_tasks": 1,
+            "fps": 10,
+            "features": {
+                "index": {"dtype": "int64", "shape": [1]},
+                "episode_index": {"dtype": "int64", "shape": [1]},
+                "frame_index": {"dtype": "int64", "shape": [1]},
+                "timestamp": {"dtype": "float32", "shape": [1]},
+                "task_index": {"dtype": "int64", "shape": [1]},
+                "observation.state": {"dtype": "float32", "shape": [2]},
+                "action": {"dtype": "float32", "shape": [1]},
+            },
+        }
+        rows = {
+            index: {
+                "index": index,
+                "episode_index": 0,
+                "frame_index": index,
+                "timestamp": index / 10,
+                "task_index": 0,
+                "observation.state": [index, index + 1],
+                "action": float(index),
+            }
+            for index in range(3)
+        }
+
+        class Reader(pmm.PaimonDatasetReader):
+
+            def __init__(self, metadata, **kwargs):
+                self.calls = []
+                self.closed = False
+                super().__init__(metadata, **kwargs)
+
+            def read_indices(self, indices, columns):
+                self.calls.append((indices, columns))
+                return pa.Table.from_pylist([
+                    {name: rows[index][name] for name in columns}
+                    for index in indices
+                ], schema=self.schema)
+
+            def close(self):
+                super().close()
+                self.closed = True
+
+        metadata = {
+            "repo_id": "logical/multi-table",
+            "revision": "dataset-version-12",
+            "info": info,
+            "episodes": [{
+                "episode_index": 0,
+                "dataset_from_index": 0,
+                "dataset_to_index": 3,
+                "length": 3,
+                "tasks": ["pick"],
+            }],
+            "tasks": ["pick"],
+            "stats": {"action": {"mean": [1.0]}},
+        }
+        missing_episodes = dict(metadata)
+        missing_episodes.pop("episodes")
+        with self.assertRaisesRegex(ValueError, "must define episodes"):
+            Reader(missing_episodes)
+        reader = Reader(
+            metadata,
+            delta_timestamps={"action": [-0.1, 0.0, 0.1]},
+        )
+        with self.assertRaisesRegex(TypeError, "tag_name"):
+            Reader(metadata, tag_name="snapshot-b")
+        dataset = pmm.PaimonLeRobotDataset(reader)
+
+        self.assertIsInstance(dataset.reader, pmm.PaimonDatasetReader)
+        self.assertEqual(_schema_from_info(info), dataset.reader.schema)
+        self.assertIsNone(dataset.reader.absolute_to_relative_idx)
+        self.assertTrue(repr(dataset).startswith("PaimonLeRobotDataset("))
+        sample, _ = dataset.__getitems__([1, 2])
+
+        self.assertEqual([((0, 1, 2), tuple(info["features"]))],
+                         reader.calls)
+        self.assertFalse(hasattr(dataset, "tag_name"))
+        self.assertEqual("dataset-version-12", dataset.meta.revision)
+        self.assertEqual((2,), dataset.features["observation.state"]["shape"])
+        self.assertEqual([1.0], dataset.meta.stats["action"]["mean"].tolist())
+        self.assertEqual(0, dataset.meta.get_task_index("pick"))
+        self.assertEqual("pick", sample["task"])
+        torch.testing.assert_close(
+            sample["observation.state"], torch.tensor([1.0, 2.0]))
+        torch.testing.assert_close(
+            sample["action"], torch.tensor([0.0, 1.0, 2.0]))
+        self.assertEqual([False, False, False],
+                         sample["action_is_pad"].tolist())
+        dataset.return_uint8 = True
+        self.assertTrue(reader.return_uint8)
+        with self.assertRaisesRegex(TypeError, "return_uint8"):
+            dataset.return_uint8 = 1
+        image_transforms = Mock()
+        dataset.image_transforms = image_transforms
+        self.assertIs(image_transforms, reader.image_transforms)
+        dataset.image_transforms = None
+        self.assertIsNone(reader.image_transforms)
+        with self.assertRaisesRegex(TypeError, "image_transforms"):
+            dataset.image_transforms = 1
+
+        wrong_schema = reader.schema.set(
+            reader.schema.get_field_index("action"),
+            pa.field("action", pa.int64()),
+        )
+        with patch.object(reader, "read_indices") as read:
+            read.return_value = pa.Table.from_pylist(
+                [{name: rows[0][name] for name in info["features"]}],
+                schema=wrong_schema,
+            )
+            with self.assertRaisesRegex(
+                    ValueError, "field action expects float, found int64"):
+                dataset[0]
+        dataset.close()
+        self.assertTrue(reader.closed)
+
+    def test_dataset_does_not_proxy_pickle_protocol(self):
+        dataset = pmm.PaimonLeRobotDataset(_PickleDatasetReader(7))
+        with self.assertRaises(AttributeError):
+            dataset.__getattr__("__getstate__")
+        restored = pickle.loads(pickle.dumps(dataset))
+
+        self.assertIsInstance(restored.reader, _PickleDatasetReader)
+        self.assertEqual(7, restored.reader.value)
+
     def test_metadata_json_preserves_nested_values(self):
         values = {
             "name": "机器人",
@@ -632,7 +782,7 @@ class LeRobotValidationTest(unittest.TestCase):
                     "L", np.full((4, 5), 80 + index, np.uint8)),
             })
 
-        dataset = object.__new__(pmm.PaimonLeRobotDataset)
+        dataset = object.__new__(_ManualDatasetReader)
         dataset._total_frames = 2
         dataset.episodes = None
         dataset._selected_ranges = None
@@ -692,7 +842,7 @@ class LeRobotValidationTest(unittest.TestCase):
             "observation.image": descriptor,
         }]
 
-        dataset = object.__new__(pmm.PaimonLeRobotDataset)
+        dataset = object.__new__(_ManualDatasetReader)
         dataset._total_frames = 1
         dataset.episodes = None
         dataset._selected_ranges = None
@@ -752,6 +902,8 @@ class LeRobotValidationTest(unittest.TestCase):
                 "pypaimon.multimodal.lerobot.dataset."
                 "_load_dataset",
                 return_value=loaded), patch(
+                "pypaimon.multimodal.lerobot.dataset._target_schema",
+                return_value=pa.schema([])), patch(
                 "pypaimon.multimodal.lerobot.dataset.sys.version_info",
                 (3, 10)):
             for invalid in (0, 1, None, "true"):
@@ -2936,8 +3088,8 @@ class LeRobotImportTest(unittest.TestCase):
             return original_plan(scan)
 
         with patch.object(TableScan, "plan", new=counted_plan), patch.object(
-                dataset, "_read_rows",
-                wraps=dataset._read_rows) as read, patch(
+                dataset.reader, "_read_rows",
+                wraps=dataset.reader._read_rows) as read, patch(
                 "pypaimon.multimodal.blob_read.fetch_blob_bodies",
                 wraps=fetch_blob_bodies) as fetch:
             last, first = dataset.__getitems__([4, 0])
@@ -2947,7 +3099,6 @@ class LeRobotImportTest(unittest.TestCase):
         self.assertIs(scanner, dataset._frame_locator._scanner)
         self.assertEqual(1, read.call_count)
         self.assertEqual([0, 1, 3, 4], read.call_args.args[0])
-        self.assertFalse(read.call_args.args[3])
         self.assertEqual(1, fetch.call_count)
         self.assertEqual(3, fetch.call_args.args[3])
         self.assertEqual("place", last["task"])

Reply via email to