Conversation
🔗 Helpful Links🧪 See artifacts and rendered test results at hud.pytorch.org/pr/pytorch/executorch/22252
Note: Links to docs will display an error until the docs builds have been completed. ⏳ 2 Pending, 1 Unrelated FailureAs of commit 700b54a with merge base 273cb33 ( BROKEN TRUNK - The following job failed but were present on the merge base:👉 Rebase onto the `viable/strict` branch to avoid these failures
This comment was automatically generated by Dr. CI and updates every 15 minutes. |
|
|
This PR needs a
|
shewu-quic
left a comment
There was a problem hiding this comment.
Thanks for looking into this issue. Could you please add a test to our new test framework?
diff --git a/backends/qualcomm/tests/rework/src/pattern.py b/backends/qualcomm/tests/rework/src/pattern.py
index f6ea8c2fae..3f44d8b786 100644
--- a/backends/qualcomm/tests/rework/src/pattern.py
+++ b/backends/qualcomm/tests/rework/src/pattern.py
@@ -4456,6 +4456,12 @@ class LiftConstantScalarOperands:
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
@@ -4526,6 +4532,10 @@ class LiftConstantScalarOperands:
f"got {n.meta['val'].dtype} (use_self_dtype broken?)"
)
+ with subtests.test(msg="skip_higher_order_ops"):
+ 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"}
…ands 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.
563428a to
e9a0b00
Compare
|
@shewu-quic added the subtest to the rework framework as suggested : _HigherOrderOps + skip_higher_order_ops in pattern.py. pytest backends/qualcomm/tests/rework/passes/test.py::test_lift_constant_scalar_operands gives 16/16 with the fix, 12/16 without |
The pass inspects node.target._schema to lift scalar operands, but higher-order ops (e.g. wrap_with_set_grad_enabled from a torch.no_grad block) have no OpOverload schema, so the pass crashed with AttributeError on a valid exported program (surfaced exporting Mamba2 / Mixtral MoE via the QNN backend). Skip non-OpOverload targets.
cc @cbilgin