From 6a9633edf88d298b607350b5bd7e4d858e8f4574 Mon Sep 17 00:00:00 2001 From: xiaohongbo Date: Mon, 14 Sep 2026 05:58:57 -0700 Subject: [PATCH] [python] Optimize batched video frame reads --- .../pypaimon/multimodal/lerobot/dataset.py | 49 ++++++++++++++--- paimon-python/pypaimon/multimodal/video.py | 44 +++++++++++---- .../pypaimon/tests/multimodal_lerobot_test.py | 54 +++++++++++++++++++ .../pypaimon/tests/multimodal_video_test.py | 41 ++++++++++++++ 4 files changed, 171 insertions(+), 17 deletions(-) diff --git a/paimon-python/pypaimon/multimodal/lerobot/dataset.py b/paimon-python/pypaimon/multimodal/lerobot/dataset.py index 6abbd34babac..196445e01b44 100644 --- a/paimon-python/pypaimon/multimodal/lerobot/dataset.py +++ b/paimon-python/pypaimon/multimodal/lerobot/dataset.py @@ -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: diff --git a/paimon-python/pypaimon/multimodal/video.py b/paimon-python/pypaimon/multimodal/video.py index ebaaab74f494..d50ec501df50 100644 --- a/paimon-python/pypaimon/multimodal/video.py +++ b/paimon-python/pypaimon/multimodal/video.py @@ -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. """ @@ -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) @@ -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: @@ -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): @@ -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) diff --git a/paimon-python/pypaimon/tests/multimodal_lerobot_test.py b/paimon-python/pypaimon/tests/multimodal_lerobot_test.py index a9790b4e1607..16a31aa9b66c 100644 --- a/paimon-python/pypaimon/tests/multimodal_lerobot_test.py +++ b/paimon-python/pypaimon/tests/multimodal_lerobot_test.py @@ -47,6 +47,7 @@ from pypaimon.multimodal.lerobot.dataset import ( _PyAVVideoDecoder, _arrow_rows, + _decode_video_rows, _image_tensor, _index_names, _open_video_decoder, @@ -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", diff --git a/paimon-python/pypaimon/tests/multimodal_video_test.py b/paimon-python/pypaimon/tests/multimodal_video_test.py index 080cbc2a9ae0..3913e1900118 100644 --- a/paimon-python/pypaimon/tests/multimodal_video_test.py +++ b/paimon-python/pypaimon/tests/multimodal_video_test.py @@ -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)