From 5417e14afebaf0d90b6b505aeea81e54486c5f3c Mon Sep 17 00:00:00 2001 From: xiaohongbo Date: Wed, 16 Sep 2026 06:41:42 -0700 Subject: [PATCH 1/4] [python] Reduce LeRobot DataLoader worker startup time --- .../pypaimon/multimodal/lerobot/dataset.py | 39 ++++++++++- .../pypaimon/tests/multimodal_lerobot_test.py | 69 +++++++++++++++++++ 2 files changed, 106 insertions(+), 2 deletions(-) diff --git a/paimon-python/pypaimon/multimodal/lerobot/dataset.py b/paimon-python/pypaimon/multimodal/lerobot/dataset.py index eca82cbfd6e9..a19da2059511 100644 --- a/paimon-python/pypaimon/multimodal/lerobot/dataset.py +++ b/paimon-python/pypaimon/multimodal/lerobot/dataset.py @@ -702,7 +702,7 @@ class _PaimonLeRobotMetadata: def __init__( self, repo_id, tag_name, info, stats, episodes, tasks, - subtasks): + subtasks, *, compact_episodes=False): self.repo_id = repo_id self.revision = tag_name self.info = info @@ -710,6 +710,41 @@ def __init__( self.episodes = episodes self.tasks = tasks self.subtasks = subtasks + self._compact_episodes = compact_episodes + + def __getstate__(self): + state = self.__dict__.copy() + if not self._compact_episodes: + return state + try: + from datasets import Dataset + except ImportError: + return state + episodes = self.episodes + if type(episodes) is not Dataset: + return state + default_format = { + "type": None, "format_kwargs": {}, + "columns": episodes.column_names, "output_all_columns": False, + } + if (not episodes.cache_files + and episodes._indices is None + and not episodes._indexes + and episodes.format == default_format): + # Rebuild Dataset's derived batch index in the worker. + state["episodes"] = ( + episodes.data.table, episodes.info, episodes.split, + episodes._fingerprint) + state["_episodes_as_arrow"] = True + return state + + def __setstate__(self, state): + if state.pop("_episodes_as_arrow", False): + from datasets import Dataset + table, info, split, fingerprint = state["episodes"] + state["episodes"] = Dataset( + table, info=info, split=split, fingerprint=fingerprint) + self.__dict__.update(state) def __getattr__(self, name): info = self.__dict__.get("info", {}) @@ -822,7 +857,7 @@ def _load_dataset(table, tag_name): "stats")) metadata = _PaimonLeRobotMetadata( str(table.identifier), tag_name, info, stats, episodes, tasks, - subtasks) + subtasks, compact_episodes=True) return frames, metadata diff --git a/paimon-python/pypaimon/tests/multimodal_lerobot_test.py b/paimon-python/pypaimon/tests/multimodal_lerobot_test.py index 7c07c0718a3b..15ee3cd9ba62 100644 --- a/paimon-python/pypaimon/tests/multimodal_lerobot_test.py +++ b/paimon-python/pypaimon/tests/multimodal_lerobot_test.py @@ -46,6 +46,7 @@ from pypaimon.multimodal.connection import MultimodalConnection from pypaimon.multimodal.lerobot import load_from_lerobot from pypaimon.multimodal.lerobot.dataset import ( + _PaimonLeRobotMetadata, _PyAVVideoDecoder, _arrow_rows, _decode_video_rows, @@ -148,6 +149,51 @@ def _catalog_metadata(connection, name): class LeRobotValidationTest(unittest.TestCase): + def test_episode_metadata_pickle_stays_small_and_usable(self): + try: + from datasets import Dataset + except ImportError: + self.skipTest("datasets is not installed") + + rows = [{ + "episode_index": index, + "dataset_from_index": index * 400, + "dataset_to_index": (index + 1) * 400, + "length": 400, + "tasks": ["pick", "place"], + } for index in range(50)] + episodes = Dataset(pa.Table.from_pylist(rows)) + metadata = _PaimonLeRobotMetadata( + "robot", "tag", {"fps": 50}, None, episodes, ["pick", "place"], + None, compact_episodes=True) + + payload = pickle.dumps(metadata) + self.assertLess(len(payload), len(pickle.dumps(episodes)) * 3 // 4) + restored = pickle.loads(payload) + self.assertIsInstance(restored.episodes, Dataset) + self.assertEqual(episodes[:], restored.episodes[:]) + self.assertEqual(episodes.features, restored.episodes.features) + self.assertEqual(episodes._fingerprint, restored.episodes._fingerprint) + self.assertEqual("tag", restored.revision) + self.assertEqual(50, restored.fps) + + metadata.episodes = episodes.with_format("numpy") + restored = pickle.loads(pickle.dumps(metadata)) + self.assertEqual("numpy", restored.episodes.format["type"]) + + with tempfile.TemporaryDirectory() as directory: + path = str(Path(directory) / "episodes.arrow") + with pa.OSFile(path, "wb") as output: + with pa.ipc.new_stream( + output, episodes.data.table.schema) as writer: + writer.write_table(episodes.data.table) + metadata.episodes = Dataset.from_file(path) + restored = pickle.loads(pickle.dumps(metadata)) + self.assertEqual( + metadata.episodes.cache_files, + restored.episodes.cache_files, + ) + def test_video_columns_decode_in_parallel(self): barrier = threading.Barrier(2) @@ -2905,6 +2951,29 @@ def _create_image_dataset(root): dataset.save_episode() dataset.finalize() + def test_table_dataset_pickle_preserves_episode_metadata_and_reads(self): + import torch + + self.connection.load_from_lerobot("worker_pickle", self.image_source) + table = self.connection.get_table("worker_pickle") + dataset = pmm.PaimonLeRobotDataset(table, return_uint8=True) + restored = pickle.loads(pickle.dumps(dataset)) + + self.assertEqual(dataset.meta.episodes[:], restored.meta.episodes[:]) + self.assertEqual( + dataset.meta.episodes._fingerprint, + restored.meta.episodes._fingerprint, + ) + for index in (0, 2, 4): + original = dataset[index] + reread = restored[index] + self.assertEqual(original.keys(), reread.keys()) + for key in original: + if torch.is_tensor(original[key]): + self.assertTrue(torch.equal(original[key], reread[key])) + else: + self.assertEqual(original[key], reread[key]) + def test_import_infers_schema_and_preserves_episodes(self): import pandas as pd From 663ee5c99a582caea07727a017c12b18ba152f25 Mon Sep 17 00:00:00 2001 From: xiaohongbo Date: Wed, 16 Sep 2026 19:52:19 -0700 Subject: [PATCH 2/4] [python] Derive LeRobot episode serialization from Dataset state --- paimon-python/pypaimon/multimodal/lerobot/dataset.py | 7 ++----- paimon-python/pypaimon/tests/multimodal_lerobot_test.py | 2 +- 2 files changed, 3 insertions(+), 6 deletions(-) diff --git a/paimon-python/pypaimon/multimodal/lerobot/dataset.py b/paimon-python/pypaimon/multimodal/lerobot/dataset.py index a19da2059511..6f984cc17268 100644 --- a/paimon-python/pypaimon/multimodal/lerobot/dataset.py +++ b/paimon-python/pypaimon/multimodal/lerobot/dataset.py @@ -702,7 +702,7 @@ class _PaimonLeRobotMetadata: def __init__( self, repo_id, tag_name, info, stats, episodes, tasks, - subtasks, *, compact_episodes=False): + subtasks): self.repo_id = repo_id self.revision = tag_name self.info = info @@ -710,12 +710,9 @@ def __init__( self.episodes = episodes self.tasks = tasks self.subtasks = subtasks - self._compact_episodes = compact_episodes def __getstate__(self): state = self.__dict__.copy() - if not self._compact_episodes: - return state try: from datasets import Dataset except ImportError: @@ -857,7 +854,7 @@ def _load_dataset(table, tag_name): "stats")) metadata = _PaimonLeRobotMetadata( str(table.identifier), tag_name, info, stats, episodes, tasks, - subtasks, compact_episodes=True) + subtasks) return frames, metadata diff --git a/paimon-python/pypaimon/tests/multimodal_lerobot_test.py b/paimon-python/pypaimon/tests/multimodal_lerobot_test.py index 15ee3cd9ba62..2e68ae445b0f 100644 --- a/paimon-python/pypaimon/tests/multimodal_lerobot_test.py +++ b/paimon-python/pypaimon/tests/multimodal_lerobot_test.py @@ -165,7 +165,7 @@ def test_episode_metadata_pickle_stays_small_and_usable(self): episodes = Dataset(pa.Table.from_pylist(rows)) metadata = _PaimonLeRobotMetadata( "robot", "tag", {"fps": 50}, None, episodes, ["pick", "place"], - None, compact_episodes=True) + None) payload = pickle.dumps(metadata) self.assertLess(len(payload), len(pickle.dumps(episodes)) * 3 // 4) From e819e38e65bf58f83cabf845703a7390d01dd4d2 Mon Sep 17 00:00:00 2001 From: xiaohongbo Date: Wed, 16 Sep 2026 20:32:31 -0700 Subject: [PATCH 3/4] [python] Avoid Dataset internals in LeRobot worker serialization --- .../pypaimon/multimodal/lerobot/dataset.py | 36 +++++++++++-------- .../pypaimon/tests/multimodal_lerobot_test.py | 15 +++++++- 2 files changed, 35 insertions(+), 16 deletions(-) diff --git a/paimon-python/pypaimon/multimodal/lerobot/dataset.py b/paimon-python/pypaimon/multimodal/lerobot/dataset.py index 6f984cc17268..66bfa24e95e0 100644 --- a/paimon-python/pypaimon/multimodal/lerobot/dataset.py +++ b/paimon-python/pypaimon/multimodal/lerobot/dataset.py @@ -18,6 +18,7 @@ """LeRobot-compatible map-style reads from a multimodal Paimon table.""" import bisect +import hashlib import io import json import math @@ -710,28 +711,24 @@ def __init__( self.episodes = episodes self.tasks = tasks self.subtasks = subtasks + self._episodes_source = None def __getstate__(self): state = self.__dict__.copy() - try: - from datasets import Dataset - except ImportError: + source = state.pop("_episodes_source", None) + if source is None: return state - episodes = self.episodes - if type(episodes) is not Dataset: + source_episodes, table, fingerprint = source + if self.episodes is not source_episodes: return state + episodes = self.episodes default_format = { "type": None, "format_kwargs": {}, "columns": episodes.column_names, "output_all_columns": False, } - if (not episodes.cache_files - and episodes._indices is None - and not episodes._indexes - and episodes.format == default_format): - # Rebuild Dataset's derived batch index in the worker. + if episodes.format == default_format and not episodes.list_indexes(): state["episodes"] = ( - episodes.data.table, episodes.info, episodes.split, - episodes._fingerprint) + table, episodes.info, episodes.split, fingerprint) state["_episodes_as_arrow"] = True return state @@ -739,8 +736,10 @@ def __setstate__(self, state): if state.pop("_episodes_as_arrow", False): from datasets import Dataset table, info, split, fingerprint = state["episodes"] - state["episodes"] = Dataset( + episodes = Dataset( table, info=info, split=split, fingerprint=fingerprint) + state["episodes"] = episodes + state["_episodes_source"] = (episodes, table, fingerprint) self.__dict__.update(state) def __getattr__(self, name): @@ -833,7 +832,7 @@ def _load_dataset(table, tag_name): frames = _component_table(catalog, raw_table, tag_name) episodes_table = _component_table( catalog, catalog.get_table(identifiers["episodes"]), tag_name) - episodes = _episode_dataset(episodes_table) + episodes, episodes_arrow, fingerprint = _episode_dataset(episodes_table) tasks_table = _component_table( catalog, catalog.get_table(identifiers["tasks"]), tag_name) tasks = _component_dataframe(tasks_table, "task_index") @@ -855,6 +854,7 @@ def _load_dataset(table, tag_name): metadata = _PaimonLeRobotMetadata( str(table.identifier), tag_name, info, stats, episodes, tasks, subtasks) + metadata._episodes_source = (episodes, episodes_arrow, fingerprint) return frames, metadata @@ -890,7 +890,13 @@ def _episode_dataset(table): if not name.startswith("stats/") ] data = _read_arrow(table, projection).sort_by("episode_index") - return Dataset(data) + # Keep the fingerprint stable when workers rebuild the Dataset. + with pa.BufferOutputStream() as output: + with pa.ipc.new_stream(output, data.schema) as writer: + writer.write_table(data) + fingerprint = hashlib.blake2b( + output.getvalue(), digest_size=8).hexdigest() + return Dataset(data, fingerprint=fingerprint), data, fingerprint def _component_dataframe(table, index_field): diff --git a/paimon-python/pypaimon/tests/multimodal_lerobot_test.py b/paimon-python/pypaimon/tests/multimodal_lerobot_test.py index 2e68ae445b0f..2b99775daaf5 100644 --- a/paimon-python/pypaimon/tests/multimodal_lerobot_test.py +++ b/paimon-python/pypaimon/tests/multimodal_lerobot_test.py @@ -162,10 +162,13 @@ def test_episode_metadata_pickle_stays_small_and_usable(self): "length": 400, "tasks": ["pick", "place"], } for index in range(50)] - episodes = Dataset(pa.Table.from_pylist(rows)) + episodes_arrow = pa.Table.from_pylist(rows) + fingerprint = "0123456789abcdef" + episodes = Dataset(episodes_arrow, fingerprint=fingerprint) metadata = _PaimonLeRobotMetadata( "robot", "tag", {"fps": 50}, None, episodes, ["pick", "place"], None) + metadata._episodes_source = (episodes, episodes_arrow, fingerprint) payload = pickle.dumps(metadata) self.assertLess(len(payload), len(pickle.dumps(episodes)) * 3 // 4) @@ -177,6 +180,11 @@ def test_episode_metadata_pickle_stays_small_and_usable(self): self.assertEqual("tag", restored.revision) self.assertEqual(50, restored.fps) + episodes.set_format("numpy") + restored = pickle.loads(pickle.dumps(metadata)) + self.assertEqual("numpy", restored.episodes.format["type"]) + episodes.reset_format() + metadata.episodes = episodes.with_format("numpy") restored = pickle.loads(pickle.dumps(metadata)) self.assertEqual("numpy", restored.episodes.format["type"]) @@ -2958,12 +2966,17 @@ def test_table_dataset_pickle_preserves_episode_metadata_and_reads(self): table = self.connection.get_table("worker_pickle") dataset = pmm.PaimonLeRobotDataset(table, return_uint8=True) restored = pickle.loads(pickle.dumps(dataset)) + reopened = pmm.PaimonLeRobotDataset(table, return_uint8=True) self.assertEqual(dataset.meta.episodes[:], restored.meta.episodes[:]) self.assertEqual( dataset.meta.episodes._fingerprint, restored.meta.episodes._fingerprint, ) + self.assertEqual( + dataset.meta.episodes._fingerprint, + reopened.meta.episodes._fingerprint, + ) for index in (0, 2, 4): original = dataset[index] reread = restored[index] From dcb798c2923fe22003ed9ead214626d53363c2bf Mon Sep 17 00:00:00 2001 From: xiaohongbo Date: Wed, 16 Sep 2026 23:36:37 -0700 Subject: [PATCH 4/4] [python] Simplify LeRobot worker startup serialization --- .../pypaimon/multimodal/lerobot/dataset.py | 47 ++++++------------- .../pypaimon/tests/multimodal_lerobot_test.py | 7 +-- 2 files changed, 15 insertions(+), 39 deletions(-) diff --git a/paimon-python/pypaimon/multimodal/lerobot/dataset.py b/paimon-python/pypaimon/multimodal/lerobot/dataset.py index 66bfa24e95e0..4f1b78f21d75 100644 --- a/paimon-python/pypaimon/multimodal/lerobot/dataset.py +++ b/paimon-python/pypaimon/multimodal/lerobot/dataset.py @@ -18,12 +18,13 @@ """LeRobot-compatible map-style reads from a multimodal Paimon table.""" import bisect -import hashlib import io import json import math import operator +import pickle import sys +import zlib from abc import ABC, abstractmethod from collections import OrderedDict from collections.abc import Mapping @@ -711,35 +712,21 @@ def __init__( self.episodes = episodes self.tasks = tasks self.subtasks = subtasks - self._episodes_source = None + self._compress_episodes = False def __getstate__(self): state = self.__dict__.copy() - source = state.pop("_episodes_source", None) - if source is None: - return state - source_episodes, table, fingerprint = source - if self.episodes is not source_episodes: - return state - episodes = self.episodes - default_format = { - "type": None, "format_kwargs": {}, - "columns": episodes.column_names, "output_all_columns": False, - } - if episodes.format == default_format and not episodes.list_indexes(): - state["episodes"] = ( - table, episodes.info, episodes.split, fingerprint) - state["_episodes_as_arrow"] = True + if state.get("_compress_episodes", False): + # Keep worker-startup payloads small without changing Dataset state. + state["episodes"] = zlib.compress( + pickle.dumps(self.episodes, protocol=pickle.HIGHEST_PROTOCOL), + level=1) + state["_episodes_zlib"] = True return state def __setstate__(self, state): - if state.pop("_episodes_as_arrow", False): - from datasets import Dataset - table, info, split, fingerprint = state["episodes"] - episodes = Dataset( - table, info=info, split=split, fingerprint=fingerprint) - state["episodes"] = episodes - state["_episodes_source"] = (episodes, table, fingerprint) + if state.pop("_episodes_zlib", False): + state["episodes"] = pickle.loads(zlib.decompress(state["episodes"])) self.__dict__.update(state) def __getattr__(self, name): @@ -832,7 +819,7 @@ def _load_dataset(table, tag_name): frames = _component_table(catalog, raw_table, tag_name) episodes_table = _component_table( catalog, catalog.get_table(identifiers["episodes"]), tag_name) - episodes, episodes_arrow, fingerprint = _episode_dataset(episodes_table) + episodes = _episode_dataset(episodes_table) tasks_table = _component_table( catalog, catalog.get_table(identifiers["tasks"]), tag_name) tasks = _component_dataframe(tasks_table, "task_index") @@ -854,7 +841,7 @@ def _load_dataset(table, tag_name): metadata = _PaimonLeRobotMetadata( str(table.identifier), tag_name, info, stats, episodes, tasks, subtasks) - metadata._episodes_source = (episodes, episodes_arrow, fingerprint) + metadata._compress_episodes = True return frames, metadata @@ -890,13 +877,7 @@ def _episode_dataset(table): if not name.startswith("stats/") ] data = _read_arrow(table, projection).sort_by("episode_index") - # Keep the fingerprint stable when workers rebuild the Dataset. - with pa.BufferOutputStream() as output: - with pa.ipc.new_stream(output, data.schema) as writer: - writer.write_table(data) - fingerprint = hashlib.blake2b( - output.getvalue(), digest_size=8).hexdigest() - return Dataset(data, fingerprint=fingerprint), data, fingerprint + return Dataset(data) def _component_dataframe(table, index_field): diff --git a/paimon-python/pypaimon/tests/multimodal_lerobot_test.py b/paimon-python/pypaimon/tests/multimodal_lerobot_test.py index 2b99775daaf5..a0fed26ddd7e 100644 --- a/paimon-python/pypaimon/tests/multimodal_lerobot_test.py +++ b/paimon-python/pypaimon/tests/multimodal_lerobot_test.py @@ -168,7 +168,7 @@ def test_episode_metadata_pickle_stays_small_and_usable(self): metadata = _PaimonLeRobotMetadata( "robot", "tag", {"fps": 50}, None, episodes, ["pick", "place"], None) - metadata._episodes_source = (episodes, episodes_arrow, fingerprint) + metadata._compress_episodes = True payload = pickle.dumps(metadata) self.assertLess(len(payload), len(pickle.dumps(episodes)) * 3 // 4) @@ -2966,17 +2966,12 @@ def test_table_dataset_pickle_preserves_episode_metadata_and_reads(self): table = self.connection.get_table("worker_pickle") dataset = pmm.PaimonLeRobotDataset(table, return_uint8=True) restored = pickle.loads(pickle.dumps(dataset)) - reopened = pmm.PaimonLeRobotDataset(table, return_uint8=True) self.assertEqual(dataset.meta.episodes[:], restored.meta.episodes[:]) self.assertEqual( dataset.meta.episodes._fingerprint, restored.meta.episodes._fingerprint, ) - self.assertEqual( - dataset.meta.episodes._fingerprint, - reopened.meta.episodes._fingerprint, - ) for index in (0, 2, 4): original = dataset[index] reread = restored[index]