Skip to content
Merged
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
5 changes: 5 additions & 0 deletions backends/qualcomm/_passes/lift_constant_scalar_operands.py
Original file line number Diff line number Diff line change
Expand Up @@ -184,6 +184,11 @@ def _lift(self, gm: torch.fx.GraphModule) -> None:
if (
n.op != "call_function"
or isinstance(n.target, (BuiltinMethodType, BuiltinFunctionType))
# Higher-order ops (e.g. wrap_with_set_grad_enabled, emitted for a
# torch.no_grad block) carry no schema, and lifting reads
# node.target._schema. Test for the schema rather than the
# OpOverload type so EdgeOpOverload targets keep lifting.
or not hasattr(n.target, "_schema")
or n.target in SKIP_LIFT_OPS
):
continue
Expand Down
14 changes: 14 additions & 0 deletions backends/qualcomm/tests/rework/src/pattern.py
Original file line number Diff line number Diff line change
Expand Up @@ -4457,6 +4457,12 @@ class _Where(torch.nn.Module):
def forward(self, x):
return torch.where(x > 0, 1.0, 0.0)

class _HigherOrderOps(torch.nn.Module):
def forward(self, x):
with torch.no_grad():
y = x + 1
return y * 2

@staticmethod
@unpack_pass_fixtures
def test(
Expand Down Expand Up @@ -4526,6 +4532,14 @@ def lower(module):
f"got {n.meta['val'].dtype} (use_self_dtype broken?)"
)

with subtests.test(msg="skip_higher_order_ops"):
# The pass reads node.target._schema; a higher-order op has none, so it
# must be skipped rather than crash the annotation pipeline.
gm = lower(LiftConstantScalarOperands._HigherOrderOps())
assertions.assert_target_count(
gm, torch.ops.higher_order.wrap_with_set_grad_enabled, 1
)


class LpaiPartitionFallbackSupport:
_SKIP_NODE_ID_SET = {"aten_add_tensor", "aten_mean_dim"}
Expand Down
23 changes: 23 additions & 0 deletions backends/qualcomm/tests/test_passes.py
Original file line number Diff line number Diff line change
Expand Up @@ -11,6 +11,7 @@
FoldQDQ,
InsertIOQDQ,
InsertReshapeForReduceOps,
LiftConstantScalarOperands,
RemoveRedundancy,
)
from executorch.backends.qualcomm._passes.qnn_pass_manager import (
Expand Down Expand Up @@ -117,6 +118,28 @@ def test_context_loader_edge_op_is_delegated(self):
QnnOperatorSupport.is_node_supported(support, None, context_loader_nodes[0])
)

def test_lift_constant_scalar_operands_skips_higher_order_ops(self):
# A `with torch.no_grad()` block is captured as a higher-order op
# (wrap_with_set_grad_enabled) whose target carries no schema.
# LiftConstantScalarOperands reads node.target._schema, so it must skip
# such nodes instead of crashing with AttributeError.
class Model(torch.nn.Module):
def forward(self, x):
with torch.no_grad():
y = x + 1
return y * 2

gm = torch.export.export(Model(), (torch.randn(3),), strict=True).graph_module
schemaless = [
node
for node in gm.graph.nodes
if node.op == "call_function" and not hasattr(node.target, "_schema")
]
self.assertTrue(schemaless, "expected a higher-order op in the graph")

# Must not raise.
LiftConstantScalarOperands().call(gm)

def test_build_op_wrappers_returns_context_binary(self):
op_name = "ctx_loader_build"
ctx_bin = b"qnn_context_binary"
Expand Down
Loading