From 213c547ec3d740cc2b91dda29200aa9513c190ed Mon Sep 17 00:00:00 2001 From: Jacob Szwejbka Date: Sat, 22 Aug 2026 15:14:37 -0700 Subject: [PATCH] Ignore symbolic metadata for lifted constants (#22056) Summary: D116957343 began lifting `get_attr` tensor constants after edge transforms. Constant-backed `FakeTensor` values can retain symbolic-looking shape metadata even though pass replay uses their attached concrete constants, causing a false symbolic-input validation failure. Exclude constant-backed `FakeTensor` values from symbolic runtime-input snapshots while retaining validation for actual symbolic inputs. Reviewed By: shoumikhin Differential Revision: D117081177 --- exir/pass_base.py | 2 ++ exir/tests/test_pass_infra.py | 22 ++++++++++++++++++++++ 2 files changed, 24 insertions(+) diff --git a/exir/pass_base.py b/exir/pass_base.py index 7d026acf878..6071aae2be8 100644 --- a/exir/pass_base.py +++ b/exir/pass_base.py @@ -121,6 +121,8 @@ def _leaf_symbolic_snapshot(value: Argument) -> Any: return scalar_snapshot if isinstance(value, FakeTensor): + if value.constant is not None: + return None dims = [] has_symbolic_dim = False for dim in value.shape: diff --git a/exir/tests/test_pass_infra.py b/exir/tests/test_pass_infra.py index 59406b13f8f..16ed5af4180 100644 --- a/exir/tests/test_pass_infra.py +++ b/exir/tests/test_pass_infra.py @@ -24,6 +24,7 @@ from executorch.exir.passes import ScalarToTensorPass from executorch.exir.passes.pass_registry import PassRegistry from executorch.exir.program import to_edge +from torch._subclasses.fake_tensor import FakeTensor from torch.export import Dim, export, ExportedProgram from torch.export.graph_signature import InputKind, InputSpec, TensorArgument from torch.fx.passes.infra.pass_base import PassBase, PassResult @@ -513,6 +514,27 @@ def placeholder( any(dim is not None for dim in self._symbolic_input_shape(new_input)) ) + def test_export_pass_ignores_symbolic_metadata_for_constant_input(self) -> None: + graph_module = self._export_dynamic_graph_module() + original_input = self._find_input_node(graph_module) + original_value = original_input.meta["val"] + self.assertIsInstance(original_value, FakeTensor) + assert isinstance(original_value, FakeTensor) + constant = torch.randn(2, 3) + original_value.constant = constant + + new_graph_module = ExportPass()(graph_module).graph_module + new_input = self._find_input_node(new_graph_module) + new_value = new_input.meta["val"] + + self.assertIsInstance(new_value, FakeTensor) + assert isinstance(new_value, FakeTensor) + self.assertIs(new_value.constant, constant) + self.assertEqual( + self._symbolic_input_shape(new_input), + self._symbolic_input_shape(original_input), + ) + def test_export_pass_rejects_collapsed_symbolic_input_metadata(self) -> None: class CollapseSymbolicInputPass(ExportPass): def placeholder(