From e42b37f817ef7dd7f3a1999fde57859937ddd834 Mon Sep 17 00:00:00 2001 From: Siddartha Pothapragada Date: Wed, 30 Sep 2026 08:39:54 -0700 Subject: [PATCH] Qualcomm AI Engine Direct - do not annotate non-float activations Observers only work with float tensors, so a quantization spec on an integer or boolean activation makes convert_pt2e emit a quantize_per_tensor whose dtype assert fires on the following export. Annotators guard the operands they look at - Cat, Stack, Embedding, Where and IndexPut each carry their own check - but annotate_single_in, Chunk, IndexCopy and SliceScatter annotate whatever they are given, so an int64 activation still reaches an observer. slice_scatter over an int64 cache index, the shape HuggingFace's cache_position update takes, is the case that surfaced this. Enforce the invariant once after annotation instead of adding a fifth per-op guard. Note _is_float_tensor cannot simply be negated for this: it also returns False for the list and tuple outputs of native_layer_norm and max.dim, which are legitimately annotated, so the sweep tests the dtype of single tensor values only and leaves every float annotation untouched. Covered by a unittest and by a test in the rework framework; both fail before the change with an output_qspec on the int64 slice_scatter. --- backends/qualcomm/quantizer/quantizer.py | 3 + backends/qualcomm/quantizer/rules.py | 32 ++++++++++ backends/qualcomm/tests/rework/src/utils.py | 48 ++++++++++++++ backends/qualcomm/tests/rework/utils/test.py | 4 ++ backends/qualcomm/tests/test_passes.py | 67 ++++++++++++++++++++ 5 files changed, 154 insertions(+) diff --git a/backends/qualcomm/quantizer/quantizer.py b/backends/qualcomm/quantizer/quantizer.py index 217d039ff60..1a6a1791498 100644 --- a/backends/qualcomm/quantizer/quantizer.py +++ b/backends/qualcomm/quantizer/quantizer.py @@ -32,6 +32,7 @@ from executorch.backends.qualcomm.quantizer.registry_loader import ( load_backend_rules_and_constraints, ) +from executorch.backends.qualcomm.quantizer.rules import drop_non_float_annotations from executorch.backends.qualcomm.quantizer.validators import NormalizedConstraints from executorch.backends.qualcomm.serialization.qc_schema import ( @@ -532,6 +533,8 @@ def annotate(self, model: GraphModule) -> GraphModule: self._annotate(model) self._annotate_custom_annotation(model) + drop_non_float_annotations(model) + # This is the only place we have sufficient information for min-max ranges. # This has to be done before calibration since this affects scale/offset. self._replace_inf(model) diff --git a/backends/qualcomm/quantizer/rules.py b/backends/qualcomm/quantizer/rules.py index f3c33d544f3..c37b40917c4 100644 --- a/backends/qualcomm/quantizer/rules.py +++ b/backends/qualcomm/quantizer/rules.py @@ -61,6 +61,38 @@ def _is_float_tensor(node: Node): return node.meta["val"].dtype in (torch.bfloat16, torch.float32) +def _is_non_float_tensor(node: Node): + """Check if the node's tensor is definitely not a float tensor. + + This is not the negation of _is_float_tensor: that one also returns False for a + node whose value is a list or tuple (aten.native_layer_norm.default, aten.max.dim), + and those outputs are legitimately annotated. + """ + if not isinstance(node, Node) or not isinstance(node.meta.get("val"), FakeTensor): + return False + return not node.meta["val"].dtype.is_floating_point + + +def drop_non_float_annotations(graph_module: torch.fx.GraphModule) -> None: + """Remove every quantization spec that landed on an integer or boolean tensor. + + Observers only work with float tensors, so such a spec makes convert_pt2e emit a + quantize_per_tensor whose input assert fires at export time. Annotators guard the + operands they know about, but an annotator cannot see which of its operands a + neighbouring op will later claim, so the invariant is enforced once here. + """ + for node in graph_module.graph.nodes: + annotation = node.meta.get(Q_ANNOTATION_KEY) + if annotation is None: + continue + if _is_non_float_tensor(node): + annotation.output_qspec = None + for arg in [ + arg for arg in annotation.input_qspec_map if _is_non_float_tensor(arg) + ]: + del annotation.input_qspec_map[arg] + + def annotate_in_out_obs_sharing_op( node: Node, quantization_config: QuantizationConfig ) -> None: diff --git a/backends/qualcomm/tests/rework/src/utils.py b/backends/qualcomm/tests/rework/src/utils.py index 148f78003d5..50f279d2edc 100644 --- a/backends/qualcomm/tests/rework/src/utils.py +++ b/backends/qualcomm/tests/rework/src/utils.py @@ -646,3 +646,51 @@ def test(compile_spec): with mock.patch("torch.export.export", wraps=torch.export.export) as export_spy: to_edge_transform_and_lower_to_qnn(exported, None, compile_spec) export_spy.assert_not_called() + + +class NonFloatActivationAnnotation: + @staticmethod + def test(quantizer): + # Observers only handle float tensors, so a spec on an int64 activation makes + # convert_pt2e emit a quantize_per_tensor whose dtype assert fires at export. + # slice_scatter over an int64 index - the shape HuggingFace's cache_position + # update takes - is annotated whole, with no per-operand guard to catch it. + class _SliceScatterInt64(torch.nn.Module): + def forward(self, x, ids): + updated = torch.slice_scatter(ids, ids[:2] + 1, dim=0, start=0, end=2) + return x * updated.to(torch.float32) + + inputs = (torch.randn(4), torch.arange(4, dtype=torch.int64)) + module = _SliceScatterInt64().eval() + gm = torch.export.export(module, inputs, strict=True).module() + gm = quantizer.transform_for_annotation(gm) + quantizer.annotate(gm) + + def _is_non_float(node): + val = node.meta.get("val") if isinstance(node, torch.fx.Node) else None + return isinstance(val, torch.Tensor) and not val.dtype.is_floating_point + + # Guard against a vacuous test. + assert any( + n.target == torch.ops.aten.slice_scatter.default and _is_non_float(n) + for n in gm.graph.nodes + ), "expected an int64 aten.slice_scatter.default in the graph" + + annotated_float_nodes = 0 + for node in gm.graph.nodes: + annotation = node.meta.get(Q_ANNOTATION_KEY) + if annotation is None: + continue + if _is_non_float(node): + assert ( + annotation.output_qspec is None + ), f"non-float node {node.name} carries an output_qspec" + else: + annotated_float_nodes += 1 + for arg in annotation.input_qspec_map: + assert not _is_non_float( + arg + ), f"{node.name} quantizes non-float input {arg.name}" + + # The float half of the graph must still be quantized. + assert annotated_float_nodes > 0, "no float node was annotated" diff --git a/backends/qualcomm/tests/rework/utils/test.py b/backends/qualcomm/tests/rework/utils/test.py index 093c69ab210..bf881997632 100644 --- a/backends/qualcomm/tests/rework/utils/test.py +++ b/backends/qualcomm/tests/rework/utils/test.py @@ -43,3 +43,7 @@ def test_qat(subtests): def test_lowering_with_exported_program(compile_spec): LoweringWithExportedProgram.test(compile_spec) # noqa: F405 + + +def test_non_float_activation_annotation(quantizer): + NonFloatActivationAnnotation.test(quantizer) # noqa: F405 diff --git a/backends/qualcomm/tests/test_passes.py b/backends/qualcomm/tests/test_passes.py index 606fed01d72..31278f383e9 100644 --- a/backends/qualcomm/tests/test_passes.py +++ b/backends/qualcomm/tests/test_passes.py @@ -60,6 +60,7 @@ from torch.export.exported_program import OutputKind from torch.library import Library from torchao.quantization.pt2e.quantize_pt2e import convert_pt2e, prepare_pt2e +from torchao.quantization.pt2e.quantizer.quantizer import Q_ANNOTATION_KEY class TestPasses(unittest.TestCase): @@ -898,6 +899,72 @@ def run_backend(backend): self.skipTest(f"LPAI quantizer unavailable: {e}") raise + def test_annotate_skips_non_float_activations(self): + """No integer or boolean activation may carry a quantization spec. + + Observers only work with float tensors, so a spec on an int64 activation makes + convert_pt2e emit a quantize_per_tensor whose dtype assert fires at export. + The per-op guards cover the operands an annotator inspects; slice_scatter on an + int64 cache index (the HuggingFace cache_position update) is annotated whole. + """ + + class SliceScatterInt64(torch.nn.Module): + def forward(self, x, ids): + updated = torch.slice_scatter(ids, ids[:2] + 1, dim=0, start=0, end=2) + return x * updated.to(torch.float32) + + module = SliceScatterInt64().eval() + sample_input = (torch.randn(4), torch.arange(4, dtype=torch.int64)) + gm = torch.export.export(module, sample_input, strict=True).module() + + quantizer = QnnQuantizer() + quantizer.set_default_quant_config(quant_dtype=QuantDtype.use_8a8w) + gm = quantizer.transform_for_annotation(gm) + quantizer.annotate(gm) + + def is_non_float(node): + val = node.meta.get("val") + return isinstance(val, torch.Tensor) and not val.dtype.is_floating_point + + # Guard against a vacuous test: the int64 slice_scatter must be in the graph. + self.assertTrue( + any( + n.target == torch.ops.aten.slice_scatter.default and is_non_float(n) + for n in gm.graph.nodes + ), + "expected an int64 aten.slice_scatter.default in the graph", + ) + + annotated_float_nodes = 0 + for node in gm.graph.nodes: + annotation = node.meta.get(Q_ANNOTATION_KEY) + if annotation is None: + continue + if is_non_float(node): + self.assertIsNone( + annotation.output_qspec, + f"non-float node {node.name} carries an output_qspec", + ) + else: + annotated_float_nodes += 1 + for arg in annotation.input_qspec_map: + self.assertFalse( + is_non_float(arg), + f"{node.name} quantizes non-float input {arg.name}", + ) + + # The float half of the graph must still be quantized. + self.assertGreater(annotated_float_nodes, 0) + + # The full pipeline: re-exporting runs the quantize_per_tensor meta kernel, + # which is where the annotation on the int64 tensor surfaces as + # "Expecting input to have dtype torch.float32, but got dtype: torch.int64". + prepared = prepare_pt2e( + torch.export.export(module, sample_input, strict=True).module(), quantizer + ) + prepared(*sample_input) + torch.export.export(convert_pt2e(prepared), sample_input, strict=True) + def test_expand_broadcast_preserves_rank0_input_mutation(self): """A rank-0 user input that is BOTH broadcast (needs rank promotion) AND mutated in place must keep its USER_INPUT_MUTATION write-back on the rank-0 value.