Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
19 changes: 19 additions & 0 deletions paimon-python/pypaimon/multimodal/lerobot/dataset.py
Original file line number Diff line number Diff line change
Expand Up @@ -22,7 +22,9 @@
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
Expand Down Expand Up @@ -710,6 +712,22 @@ def __init__(
self.episodes = episodes
self.tasks = tasks
self.subtasks = subtasks
self._compress_episodes = False

def __getstate__(self):
state = self.__dict__.copy()
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_zlib", False):
state["episodes"] = pickle.loads(zlib.decompress(state["episodes"]))
self.__dict__.update(state)

def __getattr__(self, name):
info = self.__dict__.get("info", {})
Expand Down Expand Up @@ -823,6 +841,7 @@ def _load_dataset(table, tag_name):
metadata = _PaimonLeRobotMetadata(
str(table.identifier), tag_name, info, stats, episodes, tasks,
subtasks)
metadata._compress_episodes = True
return frames, metadata


Expand Down
77 changes: 77 additions & 0 deletions paimon-python/pypaimon/tests/multimodal_lerobot_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -148,6 +149,59 @@ 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_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._compress_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)

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

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)

Expand Down Expand Up @@ -2905,6 +2959,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

Expand Down
Loading