Skip to content
Merged
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
13 changes: 11 additions & 2 deletions modelopt/torch/quantization/plugins/huggingface.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Comment thread
yingguo-trt marked this conversation as resolved.
)

if kitchen is not None and self.kitchen_attn_fn is None:
Expand Down
126 changes: 126 additions & 0 deletions tests/unit/torch/quantization/plugins/test_attention_quant.py
Original file line number Diff line number Diff line change
Expand Up @@ -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"])
Loading