Skip to content

Commit adfd1d7

Browse files
authored
Use the backend's copy hook in single-backend to_backend (pytorch#23339)
## What is wrong today Before ExecuTorch calls a backend's `preprocess`, it copies the program, so a backend that edits the graph cannot change the caller's program. `BackendDetails.copy_exported_program_for_preprocess` lets a backend decide how that copy is made. The default is `copy.deepcopy`. A backend can override it to share large constant tensors instead of copying them. The CUDA backend does that in its low-memory mode, because copying model weights costs as much memory as the weights themselves. `EdgeProgramManager.to_backend`, which `to_edge_transform_and_lower` also uses, already calls the hook. But two other paths do not: - `to_backend(backend_id, edge_program, compile_specs)` calls `copy.deepcopy` directly. - `to_backend(edge_program, partitioner)` calls that same function for every partition. On those paths a backend's override is ignored, and every constant is copied. For a large delegate payload that can be gigabytes of extra memory during lowering. ## What this change does The single-backend `to_backend` now calls `cls.copy_exported_program_for_preprocess(edge_program, compile_specs)` instead of `copy.deepcopy(edge_program)`. ```python copied_edge_program = cls.copy_exported_program_for_preprocess( edge_program, compile_specs ) ``` The default hook is still `copy.deepcopy`, so nothing changes for a backend that does not override it. Only backends that opted in see a difference, and they now get the same behavior on every lowering path. ## What was tested - New unit test in `test_lowered_backend_module.py`: `to_backend` calls the backend's copy hook once with the program and the compile specs, and passes what the hook returns to `preprocess`. It fails before this change (the hook is never called) and passes after it. - The other tests in that file give the same result before and after. One end-to-end test fails both ways in my environment, because the installed runtime does not include the test-only demo backend.
1 parent 3ebf896 commit adfd1d7

2 files changed

Lines changed: 24 additions & 1 deletion

File tree

‎exir/backend/backend_api.py‎

Lines changed: 3 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -111,7 +111,9 @@ def to_backend(
111111
# All backend implementation are final, so we don't need to consider nested subclasses.
112112
for cls in BackendDetails.__subclasses__():
113113
if backend_id == cls.__name__:
114-
copied_edge_program = copy.deepcopy(edge_program)
114+
copied_edge_program = cls.copy_exported_program_for_preprocess(
115+
edge_program, compile_specs
116+
)
115117
preprocess_result: PreprocessResult = cls.preprocess(
116118
copied_edge_program,
117119
compile_specs,

‎exir/backend/test/test_lowered_backend_module.py‎

Lines changed: 21 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -6,6 +6,7 @@
66

77
import operator
88
import unittest
9+
from unittest.mock import patch
910

1011
import executorch.exir.tests.models as models
1112

@@ -279,3 +280,23 @@ def test_arrange_graph_outputs_reorders_mutations_before_user_outputs(self):
279280
self.assertEqual(gi0.args[1], 1)
280281
self.assertEqual(gi1.args[1], 0)
281282
self.assertEqual(gi2.args[1], 2)
283+
284+
def test_to_backend_copies_the_program_through_the_backend_hook(self):
285+
# A backend overrides copy_exported_program_for_preprocess to avoid copying
286+
# large constants, so to_backend must not deep-copy the program around it.
287+
model = models.MLP()
288+
edge_program = to_edge(
289+
export(model, model.get_random_inputs(), strict=True)
290+
).exported_program()
291+
292+
with patch.object(
293+
DemoBackend,
294+
"copy_exported_program_for_preprocess",
295+
return_value=edge_program,
296+
) as copy_hook, patch.object(
297+
DemoBackend, "preprocess", wraps=DemoBackend.preprocess
298+
) as preprocess:
299+
to_backend(DemoBackend.__name__, edge_program, [])
300+
301+
copy_hook.assert_called_once_with(edge_program, [])
302+
self.assertIs(preprocess.call_args.args[0], edge_program)

0 commit comments

Comments
 (0)