diff --git a/modelopt/torch/quantization/plugins/huggingface.py b/modelopt/torch/quantization/plugins/huggingface.py index 2515910ec6e..12e7c49b687 100644 --- a/modelopt/torch/quantization/plugins/huggingface.py +++ b/modelopt/torch/quantization/plugins/huggingface.py @@ -247,8 +247,17 @@ 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: + if args: + kwargs.pop("attention_mask", None) + 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..e7b8fb69d61 100644 --- a/tests/unit/torch/quantization/plugins/test_attention_quant.py +++ b/tests/unit/torch/quantization/plugins/test_attention_quant.py @@ -185,3 +185,129 @@ def test_p_qdq_mode_detection(): sq.block_sizes = None sq.disable() assert quant_attention._p_qdq_mode() is None + + +@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() + + pq = quant_attention.p_bmm_quantizer + pq.num_bits = (4, 3) + pq.block_sizes = None + 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", + triton_attention, + ) + monkeypatch.setattr( + quant_attention, + "_init_kitchen_attn_fn", + lambda: pytest.fail("P quantizer dispatch reached Kitchen initialization"), + ) + + 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) + output = quant_attention._quantized_attention( + original_attention, + quant_attention, + states, + states, + states, + expected_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"])