Skip to content

Commit 338ac76

Browse files
authored
Erase the cat/slice/select nop ops after memory planning (pytorch#22651)
Summary: `_cat_nop`, `_slice_copy_nop` and `_select_copy_nop` do nothing at runtime. A cat is only rewritten to its nop form once memory planning can place every input at a contiguous offset inside the output, and a slice or select likewise once its output can be colocated inside its input, so by the time planning has run they hold by construction. The instructions that remain are pure dispatch overhead. This erases them, pointing their consumers at the output buffer the producers have already written in place. A nop that cannot be erased raises rather than being skipped: these ops exist only to carry a placement constraint from constraint generation to here, so one reaching the emitter is a compiler bug. **Why it has to live inside the memory planning pass.** A standalone pass is not possible: - Before planning the nodes cannot go, because the node is what carries the placement constraint. Remove it early and the aliasing never happens. - After planning nothing may run at all: # WARNING: DO NOT ADD ANY MORE PASSES AFTER MEMORY PLANNING PASS. # THERE ARE A LOT OF ASSUMPTIONS IN THE STACK THAT MEMORY PLANNING IS # THE LAST PASS BEFORE THE EMITTER. (exir/program/_program.py) `CadenceMemoryPlanning.run` already mutates the graph after `mem_planning.run` - `SimplifyIdmaOpsPass` retargets nodes and runs dead code elimination there - so that slot is the one point where placement is decided but the program is not yet emitted. This joins it. **The assumptions that warning refers to are real.** Spec lifetimes are node indices, so erasing nodes leaves them pointing past the end of the graph. That trips `find_peak_memory_usage` and, in executorch/util, the activation memory profiler. `update_all_tensors_lifetime` alone does not fix it, because `update_tensor_lifetime` only ever widens a lifetime: end = node_idx if end is None or end < node_idx else end so a stale larger bound survives. `_refresh_lifetimes` uses the first call to identify exactly which specs the recompute reaches, clears those, and rebuilds them. Clearing a wider set would leave specs the recompute never revisits stuck at None, which reads as "no lifetime" and silently drops them from the memory reports. With that, no change is needed outside the Cadence backend - the shared executorch diagnostics work unmodified. The nop targets are looked up inside `call` rather than at class scope, because their schemas come from `ops_registrations` and `memory_planning` does not import it - resolving at import time breaks any module that imports `CadenceMemoryPlanning` without having registered the ops first. **The kernels stay, deliberately.** Reviewed By: DrJessop Differential Revision: D118922409 Pull Request resolved: pytorch#22651
1 parent 8449c90 commit 338ac76

2 files changed

Lines changed: 273 additions & 23 deletions

File tree

‎backends/cadence/aot/memory_planning.py‎

Lines changed: 94 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -26,7 +26,11 @@
2626
)
2727

2828
from executorch.exir import ExecutorchProgramManager
29-
from executorch.exir.memory_planning import collect_specs_from_nodes, Verifier
29+
from executorch.exir.memory_planning import (
30+
collect_specs_from_nodes,
31+
update_all_tensors_lifetime,
32+
Verifier,
33+
)
3034
from executorch.exir.pass_base import PassBase
3135
from executorch.exir.pass_manager import PassManager
3236
from executorch.exir.passes import MemoryPlanningPass
@@ -388,6 +392,89 @@ def call(self, graph_module: torch.fx.GraphModule) -> Optional[PassResult]:
388392
return PassResult(graph_module, modified)
389393

390394

395+
class RemoveNopOpsPass(PassBase):
396+
"""Erase the cat/slice/select nop nodes once memory planning is done.
397+
398+
These ops never touch data. A cat is only rewritten to its nop form once
399+
memory planning can place every input at a contiguous offset inside the
400+
output, and a slice likewise once its output can be colocated inside its
401+
input, so by the time planning has run both hold by construction and the
402+
instructions are pure dispatch overhead.
403+
404+
This has to happen here rather than as its own pass. Before planning the
405+
nodes cannot go, because they are what carries the placement constraint;
406+
after it nothing may run at all (see the warning in
407+
exir/program/_program.py). Running inside the memory planning pass, in the
408+
slot SimplifyIdmaOpsPass already occupies, is the one point where the
409+
placement is decided but the program is not yet emitted.
410+
"""
411+
412+
def __init__(self, graph_signature: Optional[ExportGraphSignature] = None) -> None:
413+
self.graph_signature = graph_signature
414+
415+
def call(self, graph_module: torch.fx.GraphModule) -> Optional[PassResult]:
416+
# Every nop target memory_constraints.py can produce. Keep in sync
417+
# with compute_cat_contiguity_constraints and
418+
# compute_slice_and_select_loc_constraints. Looked up here rather than
419+
# at class scope because the schemas come from ops_registrations, which
420+
# this module does not import.
421+
targets = (
422+
torch.ops.aten._cat_nop.out,
423+
torch.ops.aten._slice_copy_nop.Tensor_out,
424+
torch.ops.aten._select_copy_nop.int_out,
425+
)
426+
modified = False
427+
for target in targets:
428+
for node in graph_module.graph.find_nodes(
429+
op="call_function", target=target
430+
):
431+
out = node.kwargs.get("out")
432+
if out is None:
433+
# These ops exist only to carry a placement constraint
434+
# between constraint generation and here. One that cannot
435+
# be erased would reach the runtime, where there is no
436+
# kernel to service it, so fail the build instead.
437+
raise RuntimeError(
438+
f"cannot erase {node.name} ({node.target}): no out "
439+
"kwarg to redirect its consumers to. A nop op must "
440+
"never reach the emitter."
441+
)
442+
# Consumers read the output buffer, which the producers have
443+
# already written in place. Point them straight at it.
444+
node.replace_all_uses_with(out)
445+
graph_module.graph.erase_node(node)
446+
modified = True
447+
448+
if not modified:
449+
return PassResult(graph_module, False)
450+
451+
graph_module.recompile()
452+
self._refresh_lifetimes(graph_module)
453+
return PassResult(graph_module, True)
454+
455+
def _refresh_lifetimes(self, graph_module: torch.fx.GraphModule) -> None:
456+
"""Recompute spec lifetimes against the graph we actually emit.
457+
458+
Lifetimes are node indices, so erasing nodes leaves them pointing past
459+
the end of the graph, which trips the peak-memory reporter and the
460+
activation profiler.
461+
462+
update_tensor_lifetime only ever widens a lifetime, so a single
463+
recompute leaves the stale upper bounds in place. The first call
464+
identifies exactly which specs the recompute reaches (it returns
465+
them); those are cleared and the second call then rebuilds them from
466+
scratch. Clearing a wider set than that would leave specs the
467+
recompute never revisits stuck at None, which reads as "no lifetime"
468+
and silently drops them from the memory reports.
469+
470+
Placement is already decided at this point - only lifetimes change.
471+
"""
472+
specs = update_all_tensors_lifetime(graph_module, self.graph_signature)
473+
for spec in specs:
474+
spec.lifetime = [None, None]
475+
update_all_tensors_lifetime(graph_module, self.graph_signature)
476+
477+
391478
ConstraintGenPassType: TypeAlias = Callable[
392479
[MemConstraints],
393480
Callable[[torch.fx.GraphModule], Optional[PassResult]],
@@ -465,8 +552,11 @@ def run(
465552
)
466553
mem_planning.run(graph_module, graph_signature)
467554

468-
graph_module = PassManager(passes=[SimplifyIdmaOpsPass()])(
469-
graph_module
470-
).graph_module
555+
graph_module = PassManager(
556+
passes=[
557+
SimplifyIdmaOpsPass(),
558+
RemoveNopOpsPass(graph_signature),
559+
]
560+
)(graph_module).graph_module
471561

472562
return PassResult(graph_module, True)

0 commit comments

Comments
 (0)