Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
10 changes: 8 additions & 2 deletions exir/passes/constant_prop_pass.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down
26 changes: 26 additions & 0 deletions exir/tests/test_passes.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down
Loading