Files
Model-Optimizer/tests/unit/torch/quantization/plugins/test_attention_quant.py
T
sychen52 d0c01a4e96 Add p quantization to our triton fa kernel (#1757)
### What does this PR do?

Type of change: new feature

- New P_QDQ feature in the Triton FA kernel: fake-quant softmax P before
P·V, modes fp8/nvfp4; API attention(..., p_qdq,
  p_qdq_scale); denominator unquantized, backward is STE.
- New quantization/attention/p_qdq.py +
quantization/common/fp8_quant.py; reuses nvfp4_quant FP4 rounding.
- _QuantAttention: softmax_quantizer→p_bmm_quantizer, dispatches
FP8/NVFP4 to the kernel (no kitchen); adds
TensorQuantizer.is_fp8/is_nvfp4_dynamic; envelope guards for unsupported
cases.
- Recipe wildcard + vLLM reload updated for the rename; nvfp4_tensor.py
comment typo fixed; ruff ignores + tests added.

### Usage

### Testing
added unit tests

### Before your PR is "*Ready for review*"

Make sure you read and follow [Contributor
guidelines](https://github.com/NVIDIA/Model-Optimizer/blob/main/CONTRIBUTING.md)
and your commits are signed (`git commit -s -S`).

Make sure you read and follow the [Security Best
Practices](https://github.com/NVIDIA/Model-Optimizer/blob/main/SECURITY.md#security-coding-practices-for-contributors)
(e.g. avoiding hardcoded `trust_remote_code=True`, `torch.load(...,
weights_only=False)`, `pickle`, etc.).

- Is this change backward compatible?: ✅ 
- If you copied code from any other sources or added a new PIP
dependency, did you follow guidance in `CONTRIBUTING.md`: N/A
- Did you write any new necessary tests?: ✅
- Did you update
[Changelog](https://github.com/NVIDIA/Model-Optimizer/blob/main/CHANGELOG.rst)?:
N/A
- Did you get Claude approval on this PR?: ✅ / ❌ / N/A <!--- Run
`/claude review`. NVIDIA org members can self-trigger for complex
changes; orthogonal to CodeRabbit. -->

### Additional Information
<!-- E.g. related issue. -->


<!-- This is an auto-generated comment: release notes by coderabbit.ai
-->
## Summary by CodeRabbit

## Summary

* **New Features**
* Added softmax probability quant-dequant (P_QDQ) support via `p_qdq`
and `p_qdq_scale`.
* When enabled, quantized attention can route through the Triton P_QDQ
path.
* Added stricter Triton attention “envelope” validation with
`validate_triton_attention_envelope`.

* **Bug Fixes**
* Improved handling of quantizer keys to correctly skip softmax-P
(`p_bmm_quantizer`) entries.

* **Tests**
* Added GPU forward/backward coverage for FP8 (E4M3) and NVFP4 (E2M1),
including reference comparisons and invalid-parameter cases.

* **Documentation**
* Updated attention quantization configs to use `p_bmm` quantizers
instead of `softmax_quantizer`.
<!-- end of auto-generated comment: release notes by coderabbit.ai -->

---------

Signed-off-by: Shiyang Chen <shiychen@nvidia.com>
2026-06-24 12:06:03 -07:00

200 lines
7.6 KiB
Python

# SPDX-FileCopyrightText: Copyright (c) 2024 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
import pytest
import torch
import torch.nn as nn
import torch.nn.functional as F
from _test_utils.torch.transformers_models import get_tiny_bert, get_tiny_llama, get_tiny_t5
from transformers import LlamaConfig
from transformers.models.llama.modeling_llama import LlamaAttention
import modelopt.torch.quantization as mtq
from modelopt.torch.quantization.plugins.huggingface import _QuantAttention
transformers = pytest.importorskip("transformers")
class MatmulAttention(nn.Module):
def forward(self, hidden_states, **kwargs):
q, k, v = hidden_states, hidden_states, hidden_states
a = torch.softmax(torch.matmul(q, k.transpose(-2, -1)), dim=-1)
return torch.matmul(a, v), None
class BMMAttention(nn.Module):
def forward(self, hidden_states, **kwargs):
q, k, v = hidden_states, hidden_states, hidden_states
a = torch.softmax(torch.bmm(q, k.transpose(-2, -1)), dim=-1)
return torch.bmm(a, v), None
class BinMatmulAttention(nn.Module):
def forward(self, hidden_states, **kwargs):
q, k, v = hidden_states, hidden_states, hidden_states
return torch.softmax(q @ k.transpose(-2, -1), dim=-1) @ v, None
class SDPAAttention(nn.Module):
def forward(self, hidden_states, **kwargs):
q, k, v = hidden_states, hidden_states, hidden_states
return F.scaled_dot_product_attention(q, k, v), None
kv_cache_config = {
"quant_cfg": [
{"quantizer_name": "*[kv]_bmm_quantizer", "cfg": {"num_bits": 4}, "enable": True},
{"quantizer_name": "*p_bmm_quantizer", "enable": False},
],
"algorithm": "max",
}
@pytest.mark.parametrize(
("model_getter", "attn_cls"),
[
(get_tiny_llama, None),
(get_tiny_llama, MatmulAttention),
(get_tiny_llama, BMMAttention),
(get_tiny_llama, BinMatmulAttention),
(get_tiny_llama, SDPAAttention),
(get_tiny_t5, None),
],
)
def test_kv_quant_hf(model_getter, attn_cls):
model_test = model_getter()
print(model_test)
input_ids = torch.randint(0, model_test.config.vocab_size, (1, 4))
if getattr(model_test.config, "is_encoder_decoder", False):
kwargs = {"decoder_input_ids": input_ids}
attention_module = "SelfAttention"
else:
kwargs = {}
attention_module = "self_attn"
original_is_compatible_attention = None
if attn_cls is not None:
# Test case for transformers < 4.48
# This needs:
# 1) replace the attention class with the test attention class
# 2) set _QuantAttention.is_compatible_attention output to False to fall back to the transformers < 4.48 support
for name, module in model_test.named_modules():
if name.endswith(attention_module):
if original_is_compatible_attention is None:
original_is_compatible_attention = _QuantAttention.is_compatible_attention
_QuantAttention.is_compatible_attention = classmethod(lambda cls, x: False)
parent = model_test.get_submodule(name.split(f".{attention_module}")[0])
setattr(parent, attention_module, attn_cls())
model_test(input_ids, **kwargs)
mtq.quantize(model_test, kv_cache_config, lambda model: model(input_ids, **kwargs))
for name, module in model_test.named_modules():
if name.endswith(attention_module):
assert hasattr(module, "k_bmm_quantizer")
assert hasattr(module, "v_bmm_quantizer")
assert module.k_bmm_quantizer.amax is not None
assert module.v_bmm_quantizer.amax is not None
model_test(input_ids, **kwargs)
if attn_cls is not None:
_QuantAttention.is_compatible_attention = original_is_compatible_attention
mtq.unregister(attn_cls)
def test_kv_quant_bert():
"""Test KV cache quantization on BERT model with decorated attention."""
model_test = get_tiny_bert()
input_ids = torch.randint(0, model_test.config.vocab_size, (1, 8))
attention_mask = torch.ones_like(input_ids)
# Run forward pass before quantization
model_test(input_ids, attention_mask=attention_mask)
# Quantize with KV cache quantization
mtq.quantize(
model_test,
kv_cache_config,
lambda model: model(input_ids, attention_mask=attention_mask),
)
# BERT attention modules are at encoder.layer.X.attention.self
found_quantized_attention = False
for name, module in model_test.named_modules():
if "attention.self" in name or name.endswith(".self"):
if hasattr(module, "k_bmm_quantizer") and hasattr(module, "v_bmm_quantizer"):
found_quantized_attention = True
# Verify quantizers were calibrated
assert module.k_bmm_quantizer.amax is not None, f"k_bmm not calibrated in {name}"
assert module.v_bmm_quantizer.amax is not None, f"v_bmm not calibrated in {name}"
assert module.q_bmm_quantizer.amax is not None, f"q_bmm not calibrated in {name}"
assert found_quantized_attention, "No quantized attention modules found in BERT model"
# Run forward pass after quantization to ensure it works
output = model_test(input_ids, attention_mask=attention_mask)
assert output is not None
assert output.start_logits is not None
assert output.end_logits is not None
def _make_quant_attention(hidden_size=128, num_q_heads=4, num_kv_heads=2):
config = LlamaConfig(
hidden_size=hidden_size,
num_attention_heads=num_q_heads,
num_key_value_heads=num_kv_heads,
)
quant_attention = _QuantAttention.convert(LlamaAttention(config, layer_idx=0))
quant_attention.config._attn_implementation = "sdpa"
return quant_attention
def test_p_qdq_mode_detection():
"""p_bmm_quantizer config maps to the right Triton softmax qdq mode."""
quant_attention = _make_quant_attention()
sq = quant_attention.p_bmm_quantizer
# Default int8 quantizer: not a supported Triton qdq format
assert quant_attention._p_qdq_mode() is None
# Per-tensor FP8 E4M3
sq.num_bits = (4, 3)
assert quant_attention._p_qdq_mode() == "fp8"
# NVFP4: E2M1 with dynamic block-16 E4M3 scales
sq.num_bits = (2, 1)
sq.block_sizes = {-1: 16, "type": "dynamic", "scale_bits": (4, 3)}
assert quant_attention._p_qdq_mode() == "nvfp4"
# Static (calibrated) NVFP4 must not map to the dynamic-scale kernel,
# including configs where a missing "type" key means static.
sq.block_sizes = {-1: 16, "type": "static", "scale_bits": (4, 3)}
assert quant_attention._p_qdq_mode() is None
sq.block_sizes = {-1: 16, "scale_bits": (4, 3)}
assert quant_attention._p_qdq_mode() is None
# MXFP8 stays on the kitchen path
sq.num_bits = (4, 3)
sq.block_sizes = {-1: 32, "type": "dynamic", "scale_bits": (8, 0)}
assert quant_attention._p_qdq_mode() is None
# Disabled quantizer never maps to a mode
sq.num_bits = (4, 3)
sq.block_sizes = None
sq.disable()
assert quant_attention._p_qdq_mode() is None