Skip to content
Closed
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
50 changes: 50 additions & 0 deletions exir/backend/test/test_partitioner.py
Original file line number Diff line number Diff line change
Expand Up @@ -620,6 +620,56 @@ def partition(
]
self.assertEqual(len(copy_node), 1)

def test_buffer_mutated_indirectly_is_not_taken_for_constant(self):
"""A buffer written a step after it is read still counts as mutated.

`tag_constant_data` used to decide by looking at a buffer's direct users for the
node that performs the mutation. A cache that is concatenated with its new values
and written back further down has something else as its user, so it was taken for
a constant and handed to the delegate as a buffer. A backend that compiles mutable
buffers into state then produced a model asking for state that the partitioner had
been told not to take over.
"""

class IndirectlyMutated(torch.nn.Module):
def __init__(self):
super().__init__()
self.register_buffer("cache", torch.zeros(1, 4))

def forward(self, x):
# The buffer's user is the cat, not the write.
joined = torch.cat([self.cache, x], dim=1)
self.cache.copy_(joined[:, -4:])
return joined.sum(1)

edge = exir.to_edge(
torch.export.export(IndirectlyMutated(), (torch.zeros(1, 2),), strict=True)
)
program = edge.exported_program()
signature = program.graph_signature
self.assertIn("cache", signature.buffers_to_mutate.values())

placeholder = next(
node
for node in program.graph.nodes
if node.op == "placeholder" and signature.inputs_to_buffers.get(node.name)
)
self.assertNotIn(
placeholder.name,
{user.name for user in placeholder.users} & set(signature.buffers_to_mutate),
"this test is only meaningful while the mutation is not a direct user",
)

for node in program.graph.nodes:
if node.op == "call_function":
node.meta["delegation_tag"] = "tag0"
tag_constant_data(program)

self.assertIsNone(
placeholder.meta.get("delegation_tag"),
"a mutated buffer must not be tagged as constant data",
)

def test_buffer_mutation1(self):
class TestModule(torch.nn.Module):
def __init__(self):
Expand Down
41 changes: 34 additions & 7 deletions exir/backend/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -342,20 +342,23 @@ def format_delegated_graph(graph_module: torch.fx.GraphModule) -> str:
return graph_format_str


def tag_constant_data(edge_program: ExportedProgram) -> None:
def _mutated_buffer_placeholders(edge_program: ExportedProgram) -> Set[torch.fx.Node]:
"""
Util function for partitioners. This function tags the const/param/buffers nodes
whose users all belong within the same partition. This should be called after tagging all other nodes.
Any const/param/buffer which is used as input to a subgraph, will be tagged with the same tag as that
subgraph. Throw error when const/param/buffers is used across different partitions. That is the
underlying data will be owned by multiple delegates.
Placeholders whose data the program mutates, which must not be treated as
constant. A buffer is matched against the signature's mutation targets rather
than by looking through its users for the mutating node: a buffer that is not
written in one step — a cache first concatenated with new values and written
back further down the chain — has something other than the mutation as its
user, and would otherwise be taken for a constant and folded into a delegate
instead of remaining a mutable buffer input. Parameters and lifted constants
have no mutation target, so for them the direct-user check remains.
"""
# Cache signature lookups to avoid rebuilding dicts on every access.
sig = edge_program.graph_signature
params_map = sig.inputs_to_parameters
buffers_map = sig.inputs_to_buffers
constants_map = sig.inputs_to_lifted_tensor_constants
buffers_to_mutate = sig.buffers_to_mutate
mutated_targets = set(buffers_to_mutate.values())

mutated_buffer = set()
for node in edge_program.graph.nodes:
Expand All @@ -364,12 +367,36 @@ def tag_constant_data(edge_program: ExportedProgram) -> None:
or node.name in buffers_map
or node.name in constants_map
):
if buffers_map.get(node.name) in mutated_targets:
logging.info(
"The buffer node is a mutated buffer node, which is not constant."
)
mutated_buffer.add(node)
continue
for node_user in node.users:
if node_user.name in buffers_to_mutate:
logging.info(
"The buffer node is a mutated buffer node, which is not constant."
)
mutated_buffer.add(node)
return mutated_buffer


def tag_constant_data(edge_program: ExportedProgram) -> None:
"""
Util function for partitioners. This function tags the const/param/buffers nodes
whose users all belong within the same partition. This should be called after tagging all other nodes.
Any const/param/buffer which is used as input to a subgraph, will be tagged with the same tag as that
subgraph. Throw error when const/param/buffers is used across different partitions. That is the
underlying data will be owned by multiple delegates.
"""
# Cache signature lookups to avoid rebuilding dicts on every access.
sig = edge_program.graph_signature
params_map = sig.inputs_to_parameters
buffers_map = sig.inputs_to_buffers
constants_map = sig.inputs_to_lifted_tensor_constants

mutated_buffer = _mutated_buffer_placeholders(edge_program)

for node in edge_program.graph.nodes:
# go through const/param/buffer nodes, if all users of const/param/buffer nodes are partitioned then partition
Expand Down
Loading