From e9a0b0069b2f37857b4ce8ee92039fe8128bc1cd Mon Sep 17 00:00:00 2001 From: Siddartha Pothapragada Date: Tue, 29 Sep 2026 13:30:06 -0700 Subject: [PATCH] Qualcomm: skip schema-less higher-order ops in LiftConstantScalarOperands LiftConstantScalarOperands reads node.target._schema to decide which scalar operands to lift. A torch.no_grad() region is captured as a higher-order op (wrap_with_set_grad_enabled) whose target carries no schema, so the pass raised AttributeError on an otherwise valid program. This blocks quantization for any model that traces a no_grad block, which includes the transformers decoder-only LLMs going through the ExecuTorch exporter. Skip nodes whose target has no _schema. Testing for the attribute the pass actually reads, rather than for the OpOverload type, keeps EdgeOpOverload targets lifting should the pass ever move into the capture pipeline. Adds a unit test in tests/test_passes.py and the requested subtest in the rework framework; both fail before the change and pass after. --- .../_passes/lift_constant_scalar_operands.py | 5 ++++ backends/qualcomm/tests/rework/src/pattern.py | 14 +++++++++++ backends/qualcomm/tests/test_passes.py | 23 +++++++++++++++++++ 3 files changed, 42 insertions(+) diff --git a/backends/qualcomm/_passes/lift_constant_scalar_operands.py b/backends/qualcomm/_passes/lift_constant_scalar_operands.py index bcd0817e253..5a224195695 100644 --- a/backends/qualcomm/_passes/lift_constant_scalar_operands.py +++ b/backends/qualcomm/_passes/lift_constant_scalar_operands.py @@ -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 diff --git a/backends/qualcomm/tests/rework/src/pattern.py b/backends/qualcomm/tests/rework/src/pattern.py index 9f833ec45a3..75d3034d946 100644 --- a/backends/qualcomm/tests/rework/src/pattern.py +++ b/backends/qualcomm/tests/rework/src/pattern.py @@ -4323,6 +4323,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( @@ -4392,6 +4398,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"} diff --git a/backends/qualcomm/tests/test_passes.py b/backends/qualcomm/tests/test_passes.py index abeb82f0e13..0fe21503e8e 100644 --- a/backends/qualcomm/tests/test_passes.py +++ b/backends/qualcomm/tests/test_passes.py @@ -11,6 +11,7 @@ FoldQDQ, InsertIOQDQ, InsertReshapeForReduceOps, + LiftConstantScalarOperands, RemoveRedundancy, ) from executorch.backends.qualcomm._passes.qnn_pass_manager import ( @@ -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"