Skip to content
Closed
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
49 changes: 42 additions & 7 deletions paimon-python/pypaimon/multimodal/lerobot/dataset.py
Original file line number Diff line number Diff line change
Expand Up @@ -1547,18 +1547,53 @@ def _identity(values):

def _decode_video_rows(row_groups, collators):
for collator in collators:
locations = []
for rows in row_groups:
indices = [
index for index, row in rows.items()
locations.extend(
(rows, index) for index, row in rows.items()
if collator.video_column in row
]
if not indices:
continue
decoded = collator([rows[index] for index in indices])
for index, row in zip(indices, decoded):
)
if locations:
input_rows = [rows[index] for rows, index in locations]
decoded = _decode_unique_video_rows(collator, input_rows)
for (rows, index), row in zip(locations, decoded):
rows[index] = row


def _decode_unique_video_rows(collator, rows):
unique = OrderedDict()
keys = []
for row in rows:
raw = row[collator.video_column]
if hasattr(raw, "as_py"):
raw = raw.as_py()
try:
key = None if raw is None else bytes(raw)
except (TypeError, ValueError):
key = object()
keys.append(key)
unique.setdefault(key, row)

decoded = collator(list(unique.values()))
values = {
key: row[collator.output_column]
for key, row in zip(unique, decoded)
}
result = []
emitted = set()
for row, key in zip(rows, keys):
output = dict(row)
value = values[key]
if key in emitted:
clone = getattr(value, "clone", None)
if callable(clone):
value = clone()
emitted.add(key)
output[collator.output_column] = value
result.append(output)
return result


def _normalize_index(index, size):
index = operator.index(index)
if index < 0:
Expand Down
44 changes: 34 additions & 10 deletions paimon-python/pypaimon/multimodal/video.py
Original file line number Diff line number Diff line change
Expand Up @@ -34,6 +34,8 @@ class VideoFrameCollator:
PyAV, TorchCodec, or an application decoder to be plugged in.

The cache is process-local and keyed by physical video payload identity.
Batch inputs are decoded in frame order per payload and returned in their
original order.
``collate_fn`` defaults to PyTorch's ``default_collate`` and may be replaced
for decoders that already return batched objects.
"""
Expand Down Expand Up @@ -82,7 +84,7 @@ def __call__(self, rows):
self._ensure_process_local_cache()
single_row = isinstance(rows, Mapping)
input_rows = [rows] if single_row else list(rows)
decoded_rows = [self._decode_row(row) for row in input_rows]
decoded_rows = self._decode_rows(input_rows)
if single_row:
return decoded_rows[0]
return self._collate(decoded_rows)
Expand All @@ -108,7 +110,35 @@ def __getstate__(self):
state["_owner_pid"] = None
return state

def _decode_row(self, row):
def _decode_rows(self, rows):
decoded = [None] * len(rows)
grouped = OrderedDict()
for position, row in enumerate(rows):
descriptor = self._descriptor(row)
if descriptor is None:
output = dict(row)
output[self.output_column] = None
decoded[position] = output
continue
grouped.setdefault(descriptor.payload_descriptor, []).append(
(descriptor.frame_index, position, row)
)

# Decode groups in last-use order so the LRU retains the videos used
# latest by the caller after the batch completes.
groups = sorted(
grouped.items(), key=lambda item: item[1][-1][1])
for payload, requests in groups:
decoder = self._decoder(payload)
for frame_index, position, row in sorted(requests):
output = dict(row)
output[self.output_column] = self.decode_fn(
decoder, frame_index, output
)
decoded[position] = output
return decoded

def _descriptor(self, row):
if not isinstance(row, Mapping):
raise ValueError("VideoFrameCollator expects row dictionaries.")
if self.video_column not in row:
Expand All @@ -117,10 +147,8 @@ def _decode_row(self, row):
)

raw = row[self.video_column]
output = dict(row)
if raw is None:
output[self.output_column] = None
return output
return None
if hasattr(raw, "as_py"):
raw = raw.as_py()
if not VideoFrameDescriptor.is_video_frame_descriptor(raw):
Expand All @@ -138,11 +166,7 @@ def _decode_row(self, row):
"VideoFrameDescriptor without trailing bytes."
% self.video_column
)
decoder = self._decoder(descriptor.payload_descriptor)
output[self.output_column] = self.decode_fn(
decoder, descriptor.frame_index, output
)
return output
return descriptor

def _decoder(self, descriptor):
resource = self._decoders.pop(descriptor, None)
Expand Down
54 changes: 54 additions & 0 deletions paimon-python/pypaimon/tests/multimodal_lerobot_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -47,6 +47,7 @@
from pypaimon.multimodal.lerobot.dataset import (
_PyAVVideoDecoder,
_arrow_rows,
_decode_video_rows,
_image_tensor,
_index_names,
_open_video_decoder,
Expand Down Expand Up @@ -144,6 +145,59 @@ def _catalog_metadata(connection, name):

class LeRobotValidationTest(unittest.TestCase):

def test_video_rows_are_decoded_in_one_batch(self):
class Collator:

video_column = "camera"
output_column = "camera"

def __init__(self):
self.calls = []

def __call__(self, rows):
self.calls.append(rows)
return [dict(row, camera="decoded") for row in rows]

collator = Collator()
base = {2: {"camera": "base"}, 3: {"action": 3}}
delta = {1: {"camera": "delta"}}

_decode_video_rows([base, delta], [collator])

self.assertEqual(1, len(collator.calls))
self.assertEqual(["base", "delta"], [
row["camera"] for row in collator.calls[0]
])
self.assertEqual("decoded", base[2]["camera"])
self.assertEqual("decoded", delta[1]["camera"])
self.assertEqual(3, base[3]["action"])

def test_duplicate_video_frames_are_decoded_once(self):
class Collator:

video_column = "camera"
output_column = "camera"

def __init__(self):
self.calls = []

def __call__(self, rows):
self.calls.append(rows)
return [dict(row, camera="decoded") for row in rows]

collator = Collator()
base = {2: {"camera": b"same", "action": 2}}
delta = {1: {"camera": b"same", "action": 1}}

_decode_video_rows([base, delta], [collator])

self.assertEqual(1, len(collator.calls))
self.assertEqual(1, len(collator.calls[0]))
self.assertEqual("decoded", base[2]["camera"])
self.assertEqual("decoded", delta[1]["camera"])
self.assertEqual(2, base[2]["action"])
self.assertEqual(1, delta[1]["action"])

@unittest.skipUnless(
av is not None and importlib.util.find_spec("torch") is not None,
"PyAV and Torch are required for video decoding",
Expand Down
41 changes: 41 additions & 0 deletions paimon-python/pypaimon/tests/multimodal_video_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -94,6 +94,47 @@ def factory(stream):
)
self.assertEqual(descriptors[0], result[0]["video"])

def test_decodes_each_video_in_frame_order_and_restores_rows(self):
descriptors = [
self._descriptor("a.mp4", b"video-a", frame)
for frame in (9, 1, 4)
]
descriptor_b = self._descriptor("b.mp4", b"video-b", 5)
calls = []

class Decoder:

def __init__(self, stream):
self.name = stream.read()

def decode(self, frame_index):
calls.append((self.name, frame_index))
return frame_index

collator = VideoFrameCollator(
self.table,
video_column="video",
decoder_factory=Decoder,
decode_fn=lambda decoder, frame, row: decoder.decode(frame),
collate_fn=lambda rows: rows,
)
try:
result = collator([
{"video": descriptors[0]},
{"video": descriptors[1]},
{"video": descriptors[2]},
{"video": descriptor_b},
])
finally:
collator.close()

self.assertEqual(
[(b"video-a", 1), (b"video-a", 4), (b"video-a", 9),
(b"video-b", 5)],
calls,
)
self.assertEqual([9, 1, 4, 5], [row["frame"] for row in result])

def test_evicts_least_recently_used_decoder(self):
descriptors = [
self._descriptor("episode-%d.mp4" % index, bytes([index]), index)
Expand Down
Loading