Skip to content

Commit 89809c4

Browse files
authored
Restore program segments from a view of the input (pytorch#23340)
## What is wrong today `deserialize_pte_binary` reads a `.pte` file back into a Python `Program`. Large data, such as delegate payloads and constant tensors, is stored in segments at the end of the file. To restore them, it first takes everything after the program data as a new `bytes` object, then cuts each segment out of that. Slicing `bytes` makes a copy, so while it runs the process holds: 1. the input data, 2. a copy of all the segments together, 3. a second copy of each segment as it is restored. For a program with several large segments, for example a delegate plus its constants, that is gigabytes of extra memory just to read the file. ## What this change does `deserialize_pte_binary` passes a `memoryview` of the input to `_restore_segments` instead of a sliced copy. A `memoryview` slice points at the same memory and copies nothing. ```python return _restore_segments( program=program, segment_data=memoryview(program_data)[segment_base_offset:], ) ``` `_restore_segments` still copies each segment into its own `bytes`, so: - the restored program carries the same types as before (`bytes` in delegate data, constant buffers and named data), - nothing in the result keeps the whole input alive. Each segment is now copied once instead of twice. The output does not change. When a file has only one segment and it runs to the end of the file, Python does not copy a slice that covers the whole object, so main made only one copy there too, and this change saves nothing. It helps every program with more than one segment, for example a delegate plus constants, or several delegates. ## What was tested - New unit test in `exir/_serialize/test/test_program.py`: `_restore_segments` receives a `memoryview` whose underlying object is the input itself, and the restored delegate data is `bytes` equal to the original blob. It fails before this change (it receives a `bytes` copy) and passes after it. - The rest of `test_program.py` and `backends/cuda/tests/test_merge_ptes.py` (which reads merged programs back with this function) pass before and after. - Read back a 2 GiB program with two 1 GiB delegate segments on Linux aarch64. Peak memory added by `deserialize_pte_binary` fell from 4.0 GiB to 2.0 GiB, and the restored delegates were identical.
1 parent adfd1d7 commit 89809c4

2 files changed

Lines changed: 36 additions & 4 deletions

File tree

‎exir/_serialize/_program.py‎

Lines changed: 8 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -670,7 +670,7 @@ def _restore_named_data(
670670
return named_data_store.get_named_data_store_output()
671671

672672

673-
def _restore_segments(program: Program, segment_data: bytes) -> PTEFile:
673+
def _restore_segments(program: Program, segment_data: memoryview) -> PTEFile:
674674
"""Moves segments from `segment_data` into `program`.
675675
676676
This should recreate the original Program that the segments were extracted
@@ -693,7 +693,9 @@ def _restore_segments(program: Program, segment_data: bytes) -> PTEFile:
693693
raise ValueError(
694694
f"Segment {i} {segment} overflows data length {len(segment_data)}"
695695
)
696-
segments.append(segment_data[segment.offset : segment.offset + segment.size])
696+
segments.append(
697+
bytes(segment_data[segment.offset : segment.offset + segment.size])
698+
)
697699

698700
# Restore delegate segments that weren't inlined previously.
699701
program = _restore_delegates(program, segments)
@@ -754,9 +756,11 @@ def deserialize_pte_binary(program_data: bytes) -> PTEFile:
754756
program: Program = _flatbuffer_to_program(program_data[:program_size])
755757

756758
if segment_base_offset != 0:
757-
# Move segment data back into the Program.
759+
# A view, so the segment data is copied only once, when each segment is
760+
# restored, rather than also as a whole before being split.
758761
return _restore_segments(
759-
program=program, segment_data=program_data[segment_base_offset:]
762+
program=program,
763+
segment_data=memoryview(program_data)[segment_base_offset:],
760764
)
761765

762766
return PTEFile(program=program, mutable_data=None, named_data=None)

‎exir/_serialize/test/test_program.py‎

Lines changed: 28 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -16,13 +16,15 @@
1616
import unittest
1717

1818
from typing import Dict, List, Sequence
19+
from unittest.mock import patch
1920

2021
from executorch.exir._serialize._flatbuffer_program import _flatbuffer_to_program
2122
from executorch.exir._serialize._named_data_store import NamedDataStoreOutput
2223
from executorch.exir._serialize._program import (
2324
_ExtendedHeader,
2425
_get_extended_header,
2526
_program_to_json,
27+
_restore_segments,
2628
deserialize_pte_binary,
2729
PTEFile,
2830
serialize_pte_binary,
@@ -749,6 +751,32 @@ def test_round_trip_with_segments(self) -> None:
749751
self.assertEqual(deserialized.mutable_data, None)
750752
self.assertEqual(deserialized.named_data, None)
751753

754+
def test_deserialize_restores_segments_from_a_view_of_the_input(self) -> None:
755+
program = get_test_program()
756+
blob = self.gen_blob_data(SEGMENT_ALIGNMENT * 4, b"\x10\x11\x01")
757+
add_delegate_data(program, program.execution_plan[0], [blob])
758+
pte_data = bytes(
759+
serialize_pte_binary(
760+
PTEFile(program=program),
761+
extract_delegate_segments=True,
762+
segment_alignment=SEGMENT_ALIGNMENT,
763+
)
764+
)
765+
766+
with patch(
767+
"executorch.exir._serialize._program._restore_segments",
768+
wraps=_restore_segments,
769+
) as restore_segments:
770+
deserialized = deserialize_pte_binary(pte_data)
771+
772+
# A copy of the segment data is as large as all the delegates together.
773+
segment_data = restore_segments.call_args.kwargs["segment_data"]
774+
self.assertIsInstance(segment_data, memoryview)
775+
self.assertIs(segment_data.obj, pte_data)
776+
restored = deserialized.program.backend_delegate_data[-1].data
777+
self.assertIsInstance(restored, bytes)
778+
self.assertEqual(restored, blob)
779+
752780
def test_no_constants(self) -> None:
753781
program = get_test_program()
754782
# Insert placeholder for non-const tensors.

0 commit comments

Comments
 (0)