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"])