Skip to content
Open
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
3 changes: 3 additions & 0 deletions backends/qualcomm/quantizer/quantizer.py
Original file line number Diff line number Diff line change
Expand Up @@ -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 (
Expand Down Expand Up @@ -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)
Expand Down
32 changes: 32 additions & 0 deletions backends/qualcomm/quantizer/rules.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down
48 changes: 48 additions & 0 deletions backends/qualcomm/tests/rework/src/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -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"
4 changes: 4 additions & 0 deletions backends/qualcomm/tests/rework/utils/test.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
67 changes: 67 additions & 0 deletions backends/qualcomm/tests/test_passes.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down Expand Up @@ -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.
Expand Down
Loading