From 744f69c053a0a10cc72b930ddb6d5e8c0d892c13 Mon Sep 17 00:00:00 2001 From: yingguo-trt <244492186+yingguo-trt@users.noreply.github.com> Date: Tue, 15 Sep 2026 15:07:54 +0800 Subject: [PATCH 1/3] Fix disabled P quantization during calibration Signed-off-by: yingguo-trt <244492186+yingguo-trt@users.noreply.github.com> --- .../torch/quantization/plugins/huggingface.py | 11 +++++- .../plugins/test_attention_quant.py | 39 +++++++++++++++++++ 2 files changed, 48 insertions(+), 2 deletions(-) diff --git a/modelopt/torch/quantization/plugins/huggingface.py b/modelopt/torch/quantization/plugins/huggingface.py index 2515910ec6e..b4098c721f1 100644 --- a/modelopt/torch/quantization/plugins/huggingface.py +++ b/modelopt/torch/quantization/plugins/huggingface.py @@ -247,8 +247,15 @@ def _quantized_attention( return self._eager_p_qdq_attention( original_attention_interface, query_states, key_states, value_states, **kwargs ) - return self._triton_qdq_attention( - p_qdq, query_states, key_states, value_states, **kwargs + if self.p_bmm_quantizer._if_quant: + return self._triton_qdq_attention( + p_qdq, query_states, key_states, value_states, **kwargs + ) + + # Fused P-QDQ paths bypass TensorQuantizer.forward(). + if not self.p_bmm_quantizer.is_enabled or not self.p_bmm_quantizer._if_quant: + return original_attention_interface( + self, query_states, key_states, value_states, *args, **kwargs ) if kitchen is not None and self.kitchen_attn_fn is None: diff --git a/tests/unit/torch/quantization/plugins/test_attention_quant.py b/tests/unit/torch/quantization/plugins/test_attention_quant.py index 702cf3ad1db..53e7e2d7b8d 100644 --- a/tests/unit/torch/quantization/plugins/test_attention_quant.py +++ b/tests/unit/torch/quantization/plugins/test_attention_quant.py @@ -185,3 +185,42 @@ def test_p_qdq_mode_detection(): sq.block_sizes = None sq.disable() assert quant_attention._p_qdq_mode() is None + + +def test_causal_p_qdq_respects_disable_quant(monkeypatch): + """Causal attention must bypass fused P QDQ when quantization is inactive.""" + quant_attention = make_quant_attention() + for name in ("q_bmm_quantizer", "k_bmm_quantizer", "v_bmm_quantizer"): + getattr(quant_attention, name).disable() + + pq = quant_attention.p_bmm_quantizer + pq.num_bits = (4, 3) + pq.block_sizes = None + pq.disable_quant() + + monkeypatch.setattr( + quant_attention, + "_triton_qdq_attention", + lambda *args, **kwargs: pytest.fail("inactive P quantizer reached Triton QDQ"), + ) + monkeypatch.setattr( + quant_attention, + "_init_kitchen_attn_fn", + lambda: pytest.fail("inactive P quantizer reached Kitchen initialization"), + ) + expected = object() + + def original_attention(*args, **kwargs): + return expected + + states = torch.zeros(1, 4, 2, 32) + output = quant_attention._quantized_attention( + original_attention, + quant_attention, + states, + states, + states, + attention_mask=None, + ) + + assert output is expected From a93f98d26b359fcdc8e1d28d0a63321488ebfd49 Mon Sep 17 00:00:00 2001 From: yingguo-trt <244492186+yingguo-trt@users.noreply.github.com> Date: Tue, 15 Sep 2026 16:16:47 +0800 Subject: [PATCH 2/3] Preserve positional attention mask on fallback Signed-off-by: yingguo-trt <244492186+yingguo-trt@users.noreply.github.com> --- .../torch/quantization/plugins/huggingface.py | 2 ++ .../plugins/test_attention_quant.py | 31 ++++++++++++++----- 2 files changed, 25 insertions(+), 8 deletions(-) diff --git a/modelopt/torch/quantization/plugins/huggingface.py b/modelopt/torch/quantization/plugins/huggingface.py index b4098c721f1..12e7c49b687 100644 --- a/modelopt/torch/quantization/plugins/huggingface.py +++ b/modelopt/torch/quantization/plugins/huggingface.py @@ -254,6 +254,8 @@ def _quantized_attention( # Fused P-QDQ paths bypass TensorQuantizer.forward(). if not self.p_bmm_quantizer.is_enabled or not self.p_bmm_quantizer._if_quant: + if args: + kwargs.pop("attention_mask", None) return original_attention_interface( self, query_states, key_states, value_states, *args, **kwargs ) diff --git a/tests/unit/torch/quantization/plugins/test_attention_quant.py b/tests/unit/torch/quantization/plugins/test_attention_quant.py index 53e7e2d7b8d..9da7af8b487 100644 --- a/tests/unit/torch/quantization/plugins/test_attention_quant.py +++ b/tests/unit/torch/quantization/plugins/test_attention_quant.py @@ -187,8 +187,9 @@ def test_p_qdq_mode_detection(): assert quant_attention._p_qdq_mode() is None -def test_causal_p_qdq_respects_disable_quant(monkeypatch): - """Causal attention must bypass fused P QDQ when quantization is inactive.""" +@pytest.mark.parametrize("quantization_active", [True, False]) +def test_causal_p_qdq_dispatch_respects_quant_state(monkeypatch, quantization_active): + """Causal attention must dispatch according to the P quantizer runtime state.""" quant_attention = make_quant_attention() for name in ("q_bmm_quantizer", "k_bmm_quantizer", "v_bmm_quantizer"): getattr(quant_attention, name).disable() @@ -196,21 +197,34 @@ def test_causal_p_qdq_respects_disable_quant(monkeypatch): pq = quant_attention.p_bmm_quantizer pq.num_bits = (4, 3) pq.block_sizes = None - pq.disable_quant() + if quantization_active: + pq.enable_quant() + else: + pq.disable_quant() + + calls = [] + expected = object() + expected_attention_mask = object() + + def triton_attention(*args, **kwargs): + calls.append("triton") + assert kwargs["attention_mask"] is expected_attention_mask + return expected monkeypatch.setattr( quant_attention, "_triton_qdq_attention", - lambda *args, **kwargs: pytest.fail("inactive P quantizer reached Triton QDQ"), + triton_attention, ) monkeypatch.setattr( quant_attention, "_init_kitchen_attn_fn", - lambda: pytest.fail("inactive P quantizer reached Kitchen initialization"), + lambda: pytest.fail("P quantizer dispatch reached Kitchen initialization"), ) - expected = object() - def original_attention(*args, **kwargs): + def original_attention(_self, _query, _key, _value, attention_mask): + calls.append("original") + assert attention_mask is expected_attention_mask return expected states = torch.zeros(1, 4, 2, 32) @@ -220,7 +234,8 @@ def original_attention(*args, **kwargs): states, states, states, - attention_mask=None, + expected_attention_mask, ) assert output is expected + assert calls == ["triton" if quantization_active else "original"] From 264a57143a24c9a0fc6cb4a0f9eb63872f043b0b Mon Sep 17 00:00:00 2001 From: yingguo-trt <244492186+yingguo-trt@users.noreply.github.com> Date: Tue, 22 Sep 2026 10:15:04 +0800 Subject: [PATCH 3/3] Test Kitchen dispatch across quantizer runtime states Signed-off-by: yingguo-trt <244492186+yingguo-trt@users.noreply.github.com> --- .../plugins/test_attention_quant.py | 72 +++++++++++++++++++ 1 file changed, 72 insertions(+) diff --git a/tests/unit/torch/quantization/plugins/test_attention_quant.py b/tests/unit/torch/quantization/plugins/test_attention_quant.py index 9da7af8b487..e7b8fb69d61 100644 --- a/tests/unit/torch/quantization/plugins/test_attention_quant.py +++ b/tests/unit/torch/quantization/plugins/test_attention_quant.py @@ -239,3 +239,75 @@ def original_attention(_self, _query, _key, _value, attention_mask): assert output is expected assert calls == ["triton" if quantization_active else "original"] + + +@pytest.mark.parametrize( + ("disable_method", "enable_method"), + [("disable", "enable"), ("disable_quant", "enable_quant")], +) +@pytest.mark.parametrize("preinitialized", [False, True]) +def test_kitchen_dispatch_respects_quantizer_runtime_state( + monkeypatch, disable_method, enable_method, preinitialized +): + """Kitchen dispatch must stop while P quantization is inactive and resume afterward.""" + quant_attention = make_quant_attention() + for name in ("q_bmm_quantizer", "k_bmm_quantizer", "v_bmm_quantizer"): + getattr(quant_attention, name).disable() + + pq = quant_attention.p_bmm_quantizer + pq.num_bits = (4, 3) + pq.block_sizes = {-1: 32, "type": "dynamic", "scale_bits": (8, 0)} + + calls = [] + + def kitchen_attention(query, _key, _value): + calls.append("kitchen") + return query.flatten(2) + + def init_kitchen(): + calls.append("init") + quant_attention.use_kitchen = True + quant_attention.kitchen_attn_fn = kitchen_attention + + monkeypatch.setattr( + "modelopt.torch.quantization.plugins.huggingface.kitchen", + object(), + ) + monkeypatch.setattr(quant_attention, "_init_kitchen_attn_fn", init_kitchen) + + if preinitialized: + quant_attention.use_kitchen = True + quant_attention.kitchen_attn_fn = kitchen_attention + + expected = object() + + def original_attention(_self, _query, _key, _value): + calls.append("original") + return expected + + states = torch.zeros(1, 4, 2, 32) + getattr(pq, disable_method)() + output = quant_attention._quantized_attention( + original_attention, + quant_attention, + states, + states, + states, + ) + + assert output is expected + assert calls == ["original"] + + calls.clear() + getattr(pq, enable_method)() + output = quant_attention._quantized_attention( + original_attention, + quant_attention, + states, + states, + states, + ) + + assert output[0].shape == (1, 2, 4, 32) + assert output[1] is None + assert calls == (["kitchen"] if preinitialized else ["init", "kitchen"])