mirror of
https://github.com/NVIDIA/Model-Optimizer.git
synced 2026-10-02 03:14:52 +08:00
### 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>
200 lines
7.6 KiB
Python
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
|