diff --git a/exir/passes/constant_prop_pass.py b/exir/passes/constant_prop_pass.py index 19c25a838af..ea600e1c222 100644 --- a/exir/passes/constant_prop_pass.py +++ b/exir/passes/constant_prop_pass.py @@ -213,11 +213,17 @@ def replace_with_constant_node( # If prop_constant_tensor_fqn already exists in the state dict, we need # to create a new name. Find the largest suffix of "_prop_tensor_constant", # and increment it by 1 to form the new name. - if prop_constant_tensor_fqn in exported_program.constants: + # A program re-entering this pass can already carry a + # `_prop_tensor_constant*` in state_dict from an earlier run. Emission + # resolves placeholders against state_dict before constants, so reusing + # such a name lets the stale entry shadow the tensor written here. Both + # the check and the suffix scan therefore span state_dict too. + taken = exported_program.constants.keys() | exported_program.state_dict.keys() + if prop_constant_tensor_fqn in taken: suffix = 1 + max( ( int(name[len(prefix) :]) - for name in exported_program.constants.keys() + for name in taken if name.startswith(prefix) and name[len(prefix) :].isdigit() ), default=-1, diff --git a/exir/tests/test_passes.py b/exir/tests/test_passes.py index 0f0caa3736b..0ebd71c1690 100644 --- a/exir/tests/test_passes.py +++ b/exir/tests/test_passes.py @@ -1690,6 +1690,32 @@ def forward(self, x: torch.Tensor) -> torch.Tensor: new_ep.graph_module.code ) + def test_constant_prop_pass_avoids_state_dict_name_collision(self) -> None: + class Add(torch.nn.Module): + def forward(self, x: torch.Tensor) -> torch.Tensor: + return x + 3 + + edge = to_edge( + export(Add(), (torch.ones(1),), strict=True), + compile_config=EdgeCompileConfig(_skip_dim_order=False), + ) + edge = edge.transform([ScalarToTensorPass(), RemoveMixedTypeOperators()]) + exported_program = lift_constant_tensor_pass(edge.exported_program()) + + # A program re-entering this pass can already carry a + # `_prop_tensor_constant*` in state_dict from an earlier run. + # Emission resolves placeholders against state_dict before + # constants, so reusing the name lets the stale entry shadow the + # propagated tensor and the two disagree on size. + stale = torch.zeros(12) + exported_program.state_dict["_prop_tensor_constant0"] = stale + + new_ep = constant_prop_pass(exported_program) + + for name in new_ep.constants: + self.assertNotIn(name, new_ep.state_dict) + self.assertIs(new_ep.state_dict["_prop_tensor_constant0"], stale) + def test_pass_no_user_inputs(self) -> None: class NoUserInputs(torch.nn.Module): def __init__(self):