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>
This commit is contained in:
sychen52
2026-06-24 12:06:03 -07:00
committed by GitHub
parent aa2a6a1b5d
commit d0c01a4e96
20 changed files with 1100 additions and 112 deletions
+2 -2
View File
@@ -89,8 +89,8 @@ def _convert_key_for_vllm(key: str, value: Any) -> tuple[str, str | None, Any]:
if "quantizer" not in key:
return ("copy", key, value)
# Skip softmax_quantizer and lm_head quantizers (not needed in vLLM).
if "softmax_quantizer" in key or (key.startswith("lm_head.") and "quantizer" in key):
# Skip p_bmm_quantizer (softmax-P) and lm_head quantizers (not needed in vLLM).
if "p_bmm_quantizer" in key or (key.startswith("lm_head.") and "quantizer" in key):
return ("skip", None, None)
# Check if this is a q/k/v projection that needs merging
@@ -15,14 +15,17 @@
"""Shared Triton kernels for modelopt (attention, quantization, etc.)."""
from collections.abc import Callable
import torch
from modelopt.torch.utils import import_plugin
IS_AVAILABLE = False
attention = None
attention_calibrate = None
register_triton_attention = None
attention: Callable | None = None
register_triton_attention: Callable | None = None
triton_attention_forward: Callable | None = None
validate_triton_attention_envelope: Callable | None = None
if torch.cuda.is_available():
with import_plugin(
@@ -32,26 +35,19 @@ if torch.cuda.is_available():
"kernel. Try to install triton with `pip install triton`."
),
):
from .triton_fa import attention as _attention
attention = _attention
IS_AVAILABLE = True
from .hf_triton_attention import register_triton_attention as _register_triton_attention
register_triton_attention = _register_triton_attention
# Calibration lives in the sparsity subpackage (skip-softmax specific).
# Imported here so ``from modelopt.torch.kernels.common.attention import
# attention_calibrate`` keeps working.
from modelopt.torch.kernels.sparsity.attention.calibrate import (
attention_calibrate as _attention_calibrate,
from .hf_triton_attention import (
register_triton_attention,
triton_attention_forward,
validate_triton_attention_envelope,
)
from .triton_fa import attention
attention_calibrate = _attention_calibrate
IS_AVAILABLE = True
__all__ = [
"IS_AVAILABLE",
"attention",
"attention_calibrate",
"register_triton_attention",
"triton_attention_forward",
"validate_triton_attention_envelope",
]
@@ -50,6 +50,107 @@ def _seq_lens_from_mask(
return None, False
def _check_mask_supported(attention_mask: torch.Tensor | None, seq_q: int) -> None:
"""Reject attention masks this wrapper would silently misread.
The wrapper only derives right-padded per-sequence lengths from 2D
``[batch, q_len]`` masks; anything else either loses padding info (4D
masks) or corrupts the varlen metadata (FA2-style ``[batch, kv_len]``
masks during cached decode).
"""
def _unsupported(reason):
return NotImplementedError(
f"The ModelOpt Triton attention kernel does not support {reason}. "
"Use unpadded (or uniform-length) right-padded inputs."
)
if attention_mask is None:
return
if attention_mask.dim() == 2:
if attention_mask.shape[1] != seq_q:
# FA2-style [batch, kv_len] mask during cached decode: the wrapper
# would misread KV lengths as query lengths (out-of-bounds access).
raise _unsupported("padded batches during cached decode")
mask_bool = attention_mask.to(torch.bool)
if not mask_bool[:, 0].all():
raise _unsupported("left-padded inputs")
# ``_seq_lens_from_mask`` derives lengths via ``sum(dim=1)``, which is only
# correct when each row is a contiguous run of valid tokens followed by
# padding. A hole (e.g. ``[1, 0, 1]``) would sum to the right count but
# place the valid tokens at the wrong positions, so reject non-right-padded
# masks (any valid token after a pad == row not monotonically non-increasing).
if not (mask_bool[:, :-1].int() >= mask_bool[:, 1:].int()).all():
raise _unsupported("non-contiguously padded inputs")
return
# 4D [batch, 1, q, kv] masks are ignored by the wrapper, which is safe only
# when they encode pure causal structure (the kernel masks causally itself).
# In a causal mask the newest query row sees every position; any masked
# entry there means padding, windowing, or a non-causal/bias pattern.
last_row = attention_mask[..., -1, :]
hidden = ~last_row if attention_mask.dtype == torch.bool else last_row != 0
if hidden.any():
raise _unsupported("masks carrying padding or non-causal structure")
def validate_triton_attention_envelope(
module: nn.Module,
query: torch.Tensor,
key: torch.Tensor,
attention_mask: torch.Tensor | None,
**kwargs,
) -> None:
"""Raise ``NotImplementedError`` for inputs outside this wrapper/kernel envelope.
These limits do not come from the quantization or sparsity features layered
on top — they document what the ``triton_fa`` kernel (causal or single-token
decode only; no sliding window, attention sinks, logit softcapping, or
dropout; head_dim >= 16) and this wrapper's varlen-metadata derivation
(right-padded 2D masks only; no multi-token forwards over a longer KV cache)
support. Callers that route arbitrary HF models onto the kernel dynamically
(e.g. the quantization plugin's p_bmm_quantizer dispatch) should call this
before dispatching, so unsupported models fail loudly instead of silently
computing wrong attention. The sparse-attention path predates these checks
and does not yet enforce them.
"""
# Mistral-style models pass sliding_window as an interface kwarg instead of
# setting it on the attention module, so check both.
if getattr(module, "sliding_window", None) or kwargs.get("sliding_window"):
raise NotImplementedError(
"The ModelOpt Triton attention kernel does not support sliding-window attention layers."
)
# Semantic attention arguments the kernel does not implement: dropping them
# would change the attention math.
for name, reason in (("s_aux", "attention sinks"), ("softcap", "logit softcapping")):
if kwargs.get(name) is not None:
raise NotImplementedError(
f"The ModelOpt Triton attention kernel does not support {reason} ('{name}')."
)
if kwargs.get("is_causal") is False or getattr(module, "is_causal", True) is False:
raise NotImplementedError(
"The ModelOpt Triton attention kernel does not support non-causal attention."
)
if kwargs.get("dropout"):
raise NotImplementedError(
"The ModelOpt Triton attention kernel does not support attention dropout; "
"set attention_dropout=0 for training."
)
if query.shape[-1] < 16:
raise NotImplementedError(
f"The ModelOpt Triton attention kernel requires head_dim >= 16, got {query.shape[-1]}."
)
seq_q, seq_k = query.shape[2], key.shape[2]
if seq_q > 1 and seq_k != seq_q:
# The wrapper only passes K-side varlen metadata for single-token decode;
# multi-token forwards over a longer KV cache would mis-index K/V.
raise NotImplementedError(
"The ModelOpt Triton attention kernel does not support multi-token "
"forwards over a longer KV cache (chunked prefill or "
"assisted/speculative decoding)."
)
_check_mask_supported(attention_mask, seq_q)
def triton_attention_forward(
module: nn.Module,
query: torch.Tensor,
@@ -58,6 +159,8 @@ def triton_attention_forward(
attention_mask: torch.Tensor | None,
scaling: float,
dropout: float = 0.0,
p_qdq: str | None = None,
p_qdq_amax: float | None = None,
**kwargs,
) -> tuple[torch.Tensor, None]:
"""Attention forward compatible with HF AttentionInterface.
@@ -75,6 +178,12 @@ def triton_attention_forward(
Other formats (e.g. 4D causal masks) are ignored.
scaling: Softmax scale (e.g. 1/sqrt(head_dim)).
dropout: Ignored (kernel has no dropout); use 0 for eval.
p_qdq: Optional softmax fake quant-dequant mode ("fp8" or
"nvfp4") forwarded to the kernel. Not passed by HF dispatch;
used by direct callers such as the quantization plugin.
p_qdq_amax: Optional per-tensor amax for the softmax-P qdq; None uses
the kernel default of 1.0 (the theoretical upper bound of the
unnormalized P's amax).
**kwargs: Reserved for future extensions.
Returns:
@@ -121,7 +230,7 @@ def triton_attention_forward(
trials = getattr(method, "_threshold_trials", None)
# Deferred: the package __init__ imports this module, so importing
# attention_calibrate at module top would be circular.
from modelopt.torch.kernels.common.attention import attention_calibrate
from modelopt.torch.kernels.sparsity.attention.calibrate import attention_calibrate
if trials and attention_calibrate is not None:
o, counters = attention_calibrate(q, k, v, **kw, threshold_trials=trials)
@@ -153,6 +262,11 @@ def triton_attention_forward(
if threshold:
kw["skip_softmax_threshold"] = threshold
if p_qdq is not None:
kw["p_qdq"] = p_qdq
if p_qdq_amax is not None:
kw["p_qdq_amax"] = p_qdq_amax
o = attention(q, k, v, **kw)
attn_output = o.view(batch, seq_len, num_heads, head_dim)
@@ -188,4 +302,5 @@ def register_triton_attention() -> bool:
__all__ = [
"register_triton_attention",
"triton_attention_forward",
"validate_triton_attention_envelope",
]
@@ -42,6 +42,8 @@ import triton.language as tl
_apply_sparse_nm_to_qk_tile: Any = None
_is_dense_region: Any = None
_skip_softmax_decision: Any = None
_p_qdq_fp8: Any = None
_p_qdq_nvfp4: Any = None
def _load_sparsity_helpers() -> None:
@@ -62,6 +64,20 @@ def _load_sparsity_helpers() -> None:
_skip_softmax_decision = _skip
def _load_p_qdq_helpers() -> None:
global _p_qdq_fp8, _p_qdq_nvfp4
if _p_qdq_fp8 is None:
from modelopt.torch.kernels.quantization.attention.p_qdq import _p_qdq_nvfp4 as _nvfp4
from modelopt.torch.kernels.quantization.common.fp8_quant import fp8_scalar_qdq as _fp8
_p_qdq_fp8 = _fp8
_p_qdq_nvfp4 = _nvfp4
# Maps the public p_qdq option to the kernel's P_QDQ constexpr.
_P_QDQ_MODES = {None: 0, "fp8": 1, "nvfp4": 2}
LOG2E: float = 1.44269504088896
# ---------------------------------------------------------------------------
@@ -246,6 +262,8 @@ def _attn_fwd(
DENSE_RECENT_TOKENS: tl.constexpr = 64, # Recent KV tokens kept dense (BLOCK_N-independent)
APPLY_SKIP_SOFTMAX: tl.constexpr = False, # Skip KV tiles with negligible scores
SKIP_THRESHOLD_LOG2: tl.constexpr = 0.0, # log2(lambda) in the kernel's scaled log2 score space
P_QDQ: tl.constexpr = 0, # Fake quant-dequant of softmax P: 0=off, 1=FP8 E4M3, 2=NVFP4
p_qdq_scale=1.0, # Per-tensor scale for softmax qdq (runtime scalar; amax/448 or amax/(6*448))
Sparsity_total=None, # Optional int64 scalar for counting total tiles (atomic)
Sparsity_skipped=None, # Optional int64 scalar for counting skipped tiles (atomic)
MEASURE_SPARSITY: tl.constexpr = False, # When True, count total/skipped tiles via atomic adds
@@ -383,6 +401,14 @@ def _attn_fwd(
row_sum = row_sum * correction + l_new
acc = acc * correction[:, None]
# --- Optional softmax quant-dequant (emulates quantized P @ V) ---
# row_sum keeps the unquantized p: the softmax denominator stays in
# fp32 and only the quantized P is fed to BMM2.
if P_QDQ == 1:
p = _p_qdq_fp8(p, p_qdq_scale)
elif P_QDQ == 2:
p = _p_qdq_nvfp4(p, p_qdq_scale, BLOCK_M, BLOCK_N)
# Load V and accumulate
if IS_PAGED:
v = _load_paged_v_tile(
@@ -806,6 +832,8 @@ class _Attention(torch.autograd.Function):
dense_recent_tokens,
skip_softmax_threshold,
measure_sparsity,
p_qdq_mode,
p_qdq_scale,
k_cache,
v_cache,
block_table,
@@ -903,6 +931,8 @@ class _Attention(torch.autograd.Function):
"DENSE_RECENT_TOKENS": dense_recent_tokens,
"APPLY_SKIP_SOFTMAX": apply_skip,
"SKIP_THRESHOLD_LOG2": skip_threshold_log2,
"P_QDQ": p_qdq_mode,
"p_qdq_scale": p_qdq_scale,
"Sparsity_total": sparsity_total,
"Sparsity_skipped": sparsity_skipped,
"MEASURE_SPARSITY": do_measure,
@@ -1106,6 +1136,8 @@ class _Attention(torch.autograd.Function):
None, # dense_recent_tokens
None, # skip_softmax_threshold
None, # measure_sparsity
None, # p_qdq_mode
None, # p_qdq_scale
None, # k_cache
None, # v_cache
None, # block_table
@@ -1132,6 +1164,8 @@ def attention(
dense_recent_tokens: int = 64,
skip_softmax_threshold: float | None = None,
measure_sparsity: bool = False,
p_qdq: str | None = None,
p_qdq_amax: float = 1.0,
k_cache: torch.Tensor | None = None,
v_cache: torch.Tensor | None = None,
block_table: torch.Tensor | None = None,
@@ -1169,6 +1203,25 @@ def attention(
and skipped tiles via atomic counters. The counts are stored as
``_sparsity_total`` and ``_sparsity_skipped`` attributes on the
returned output tensor.
p_qdq: Fake quant-dequant of the softmax probabilities ``P``
before the ``P @ V`` matmul (BMM2), emulating quantized attention.
``"fp8"`` round-trips P through FP8 E4M3 with a static per-tensor
scale (see ``p_qdq_amax``). ``"nvfp4"`` applies the two-level NVFP4
recipe: E2M1 elements with one FP8 E4M3 scale per 16 elements along
the key dimension (the BMM2 contraction axis; every autotuned
BLOCK_N is a multiple of 16). The softmax denominator stays
unquantized. The backward pass uses the straight-through estimator:
gradients are computed from the unquantized P, matching QAT
references that keep the backward dots in high precision.
Set to ``None`` to disable.
p_qdq_amax: Per-tensor amax for the softmax-P quant-dequant. The
kernel's unnormalized P lies in [0, 1] (the max-subtraction caps
every entry at ``exp2(0) = 1``), so 1 is the theoretical upper
bound of its amax — hence the default of 1.0. It is converted to
the standard per-tensor scale internally: ``amax / 448`` for FP8,
and the global scale ``amax / (6 * 448)`` for NVFP4. A runtime
scalar — user-set or calibrated values do not recompile the
kernel. Values above amax saturate.
k_cache: Paged K cache [num_blocks, page_size, num_kv_heads, head_dim].
When provided, K/V are read from paged cache via block_table
instead of from contiguous k/v tensors.
@@ -1186,7 +1239,27 @@ def attention(
require grad, because the saved ``k``/``v`` are dummy tensors in paged
mode and dK/dV would be silently incorrect.
"""
# Both loaders must run unconditionally: Triton computes a kernel's
# dependency hash once, on the first call, walking the full AST. If the
# qdq helpers were still None at that point, their source would be
# permanently excluded from the cache key and later edits to them would
# silently reuse stale compiled kernels from the on-disk cache.
_load_sparsity_helpers()
_load_p_qdq_helpers()
if p_qdq not in _P_QDQ_MODES:
raise ValueError(
f"p_qdq must be one of {sorted(k for k in _P_QDQ_MODES if k)} or None, got {p_qdq!r}"
)
p_qdq_mode = _P_QDQ_MODES[p_qdq]
# Convert the per-tensor amax to the kernel's scale convention
# (``q = cast(p / scale) * scale``): FP8 uses ``amax / 448``; NVFP4 uses the
# global scale ``amax / (6 * 448)``. amax=1 (the default, the theoretical
# upper bound of P's amax) therefore maps to the standard full-range scale.
p_qdq_scale = 1.0
if p_qdq_mode:
if not (math.isfinite(p_qdq_amax) and p_qdq_amax > 0):
raise ValueError(f"p_qdq_amax must be a finite positive value, got {p_qdq_amax}")
p_qdq_scale = p_qdq_amax / 448.0 if p_qdq == "fp8" else p_qdq_amax / (6.0 * 448.0)
sm_scale = 1.0 / (q.shape[2] ** 0.5) if softmax_scale is None else softmax_scale
return _Attention.apply(
q,
@@ -1206,6 +1279,8 @@ def attention(
dense_recent_tokens,
skip_softmax_threshold,
measure_sparsity,
p_qdq_mode,
p_qdq_scale,
k_cache,
v_cache,
block_table,
@@ -13,4 +13,12 @@
# See the License for the specific language governing permissions and
# limitations under the License.
"""Quantization-specific attention kernel pieces (placeholder for combined sparse+quant path)."""
"""Quantization-specific attention kernel pieces.
``p_qdq.py`` holds the softmax-P (``p_bmm_quantizer``) quant-dequant
``@triton.jit`` helpers invoked by the unified flash-attention kernel in
``common/attention/triton_fa.py`` under its ``P_QDQ`` constexpr guard.
Only NVFP4 needs a P-specific helper (tiling and block-amax policy on top of
``quantization/gemm/nvfp4_quant.py``); the FP8 mode uses
``quantization/common/fp8_quant.fp8_scalar_qdq`` directly.
"""
@@ -0,0 +1,65 @@
# SPDX-FileCopyrightText: Copyright (c) 2026 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.
"""Softmax-P quant-dequant helpers for the unified flash attention kernel.
These ``@triton.jit`` helpers fake-quantize the softmax probabilities ``P``
before the ``P @ V`` matmul (BMM2) — the in-kernel counterpart of the
``p_bmm_quantizer`` config. They are called conditionally from the baseline
flash-attention kernel in ``common/attention/triton_fa.py`` under the
``P_QDQ`` constexpr guard, following the same composition pattern as
the sparsity helpers in ``sparsity/attention/skip_softmax_helpers.py``.
Only NVFP4 needs a P-specific helper (tiling policy and block amaxes); the
per-tensor FP8 mode uses ``quantization/common/fp8_quant.fp8_scalar_qdq``
directly. What is P-specific here: the kernel's online-softmax ``p`` is
unnormalized and bounded (``0 <= p <= 1``, since the max-subtraction caps
every entry at ``exp2(0) = 1``), so 1 is the theoretical upper bound of its
amax; block amaxes need no ``abs``; and the NVFP4 scale blocks of 16 run
along the key dimension — the contraction axis of ``P @ V``. The caller
(``attention()``) converts the amax to the ``global_scale`` below.
"""
import triton
import triton.language as tl
from modelopt.torch.kernels.quantization.gemm.nvfp4_quant import nvfp4_scalar_qdq
@triton.jit
def _p_qdq_nvfp4(
p,
global_scale,
BLOCK_M: tl.constexpr,
BLOCK_N: tl.constexpr,
):
"""NVFP4 fake quant-dequant of softmax probabilities.
Two-level scaling per the NVFP4 recipe: E2M1 elements with one FP8 E4M3
scale per 16 contiguous elements along the key dimension (the contraction
axis of ``P @ V``), and a per-tensor ``global_scale`` (runtime scalar,
``amax / (6 * 448)``; ``attention()`` derives it from ``p_qdq_amax``,
which defaults to 1, the theoretical upper bound of P's amax).
``p >= 0``, so the block amax is a plain max (no ``abs``), and
``nvfp4_scalar_qdq`` guards the degenerate all-zero blocks of fully
masked or padded positions.
"""
tl.static_assert(BLOCK_N % 16 == 0, "BLOCK_N must be divisible by 16 for NVFP4")
grouped = tl.reshape(p, (BLOCK_M, BLOCK_N // 16, 16))
block_amax = tl.expand_dims(tl.max(grouped, axis=2), 2) # p >= 0, so max == amax
q = nvfp4_scalar_qdq(grouped, block_amax, global_scale, 16)
return tl.reshape(q, (BLOCK_M, BLOCK_N))
@@ -0,0 +1,20 @@
# SPDX-FileCopyrightText: Copyright (c) 2026 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.
"""Shared composable Triton JIT fake-quantization functions.
Format-level building blocks (FP8 E4M3, NVFP4/E2M1) reused across the gemm
and attention kernel packages.
"""
@@ -0,0 +1,43 @@
# SPDX-FileCopyrightText: Copyright (c) 2026 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.
"""Composable Triton JIT functions for FP8 (E4M3) fake quantization.
Counterpart of ``gemm/nvfp4_quant.py`` for per-tensor FP8. Used by the unified
flash-attention kernel's softmax-P qdq (``common/attention/triton_fa.py``).
"""
import triton
import triton.language as tl
@triton.jit
def fp8_scalar_qdq(x, scale):
"""Per-tensor FP8 E4M3 fake quant-dequant: ``cast(x / scale) * scale``.
Standard quantizer convention with ``scale = amax / 448``. Works with any
tensor shape and sign (all ops are element-wise); out-of-range values
saturate to +-448 like a real quantizer.
Args:
x: Tensor of values to fake-quantize.
scale: Per-tensor scale (runtime scalar or broadcastable tensor).
Returns:
Fake-quantized tensor of the same shape as ``x``, in float32.
"""
FP8_E4M3_MAX: tl.constexpr = 448.0
x_scaled = tl.clamp(x / scale, -FP8_E4M3_MAX, FP8_E4M3_MAX)
return x_scaled.to(tl.float8e4nv).to(tl.float32) * scale
@@ -16,9 +16,10 @@
"""Composable Triton JIT functions for NVFP4 (E2M1) fake quantization.
Single source of truth for FP4 decision-boundary rounding. Used by:
- ``fp4_kernel.py`` (standalone blockwise fake quant)
- ``fp4_kernel_hopper.py`` (Hopper block-pointer variant)
- ``gptq_fused_kernel.py`` (fused GPTQ scalar path)
- ``fp4_kernel.py`` (standalone blockwise fake quant)
- ``fp4_kernel_hopper.py`` (Hopper block-pointer variant)
- ``gptq_fused_kernel.py`` (fused GPTQ scalar path)
- ``../attention/p_qdq.py`` (softmax-P qdq in the flash-attention kernel)
FP4 (E2M1) representable magnitudes: {0.0, 0.5, 1.0, 1.5, 2.0, 3.0, 4.0, 6.0}
"""
@@ -73,9 +74,13 @@ def nvfp4_scalar_quant(
Quantizes each element independently: divide by scale, round to nearest
FP4 (E2M1) value via ``fp4_round_magnitude``, multiply by scale.
All ops are element-wise, so any shape works with a broadcastable
``scale`` (e.g. a [M, B, 16] tile with [M, B, 1] per-block scales, as
used by the attention softmax-P qdq).
Args:
x: [N] float32 tensor of values to quantize (already in registers).
scale: float32 scalar block scale.
scale: float32 scalar block scale (or broadcastable tensor of scales).
N: Compile-time number of elements.
Returns:
@@ -131,9 +136,13 @@ def nvfp4_scalar_qdq(
:func:`fp8_quantize_scale`, then quantizes each element to the nearest
FP4 (E2M1) value.
All ops are element-wise, so any shape works with a broadcastable
``block_amax``.
Args:
x: [N] float32 tensor of values to quantize.
block_amax: Per-block amax (absolute maximum of the block).
block_amax: Per-block amax (absolute maximum of the block);
scalar or broadcastable tensor.
global_scale: Pre-computed ``global_amax / (6.0 * 448.0)``.
N: Compile-time number of elements.
@@ -164,7 +164,7 @@ def _diffusers_triton_attention(
calib_mode = getattr(_thread_local, "calibration_mode", False)
if calib_mode:
trials = getattr(_thread_local, "threshold_trials", None)
from modelopt.torch.kernels.common.attention import attention_calibrate
from .calibrate import attention_calibrate
if trials and attention_calibrate is not None:
o, counters = attention_calibrate(q, k, v, **kw, threshold_trials=trials)
@@ -126,7 +126,7 @@ def _ltx_triton_attention(
calib_mode = getattr(_thread_local, "calibration_mode", False)
if calib_mode:
trials = getattr(_thread_local, "threshold_trials", None)
from modelopt.torch.kernels.common.attention import attention_calibrate
from .calibrate import attention_calibrate
if trials and attention_calibrate is not None:
o, counters = attention_calibrate(q_flat, k_flat, v_flat, **kw, threshold_trials=trials)
@@ -530,6 +530,26 @@ class TensorQuantizer(nn.Module):
and self.block_sizes.get("scale_bits", None) == (8, 0)
)
@property
def is_fp8(self):
"""Check if is per-tensor FP8 E4M3 (no block scales, no per-channel axis)."""
return self._num_bits == (4, 3) and self._block_sizes is None and self._axis is None
@property
def is_nvfp4_dynamic(self):
"""Check if is dynamic NVFP4: E2M1 with E4M3 per-block scales computed dynamically.
Mirror of ``is_nvfp4_static`` for the dynamic-scale layout; like it, this does
not constrain the block size. Consumers that require a specific block size
(e.g. the block-16 Triton kernels) check ``block_sizes[-1]`` downstream.
"""
return (
self._block_sizes is not None
and self._block_sizes.get("type", None) == "dynamic"
and self._num_bits == (2, 1)
and self._block_sizes.get("scale_bits", None) == (4, 3)
)
@property
def is_nvfp4_static(self):
"""True for E2M1 weights + E4M3 per-block scales in static layout (format-only check)."""
@@ -30,6 +30,11 @@ from torch import Tensor
from torch.nn.functional import linear
from transformers.models.t5.modeling_t5 import T5Attention
from modelopt.torch.kernels.common.attention import IS_AVAILABLE as TRITON_FA_AVAILABLE
from modelopt.torch.kernels.common.attention import (
triton_attention_forward,
validate_triton_attention_envelope,
)
from modelopt.torch.kernels.quantization.gemm import IS_AVAILABLE as IS_TRITON_AVAILABLE
from modelopt.torch.opt.dynamic import DynamicModule
from modelopt.torch.utils.distributed import ParallelState
@@ -77,16 +82,16 @@ class _QuantAttention(QuantModule):
self.q_bmm_quantizer = TensorQuantizer()
self.k_bmm_quantizer = TensorQuantizer()
self.v_bmm_quantizer = TensorQuantizer()
self.softmax_quantizer = TensorQuantizer()
self.p_bmm_quantizer = TensorQuantizer()
self.kitchen_attn_fn = None
self.use_kitchen = False
def _init_kitchen_attn_fn(self):
if not self.softmax_quantizer.is_enabled:
if not self.p_bmm_quantizer.is_enabled:
self.kitchen_attn_fn = "disabled"
return
self.use_kitchen = True
if self.softmax_quantizer.is_mxfp(8):
if self.p_bmm_quantizer.is_mxfp(8):
qfa_params = triton_fa_params.QTritonFAParams(
backend="triton",
qk_dot_precisions="bf16@bf16",
@@ -99,7 +104,7 @@ class _QuantAttention(QuantModule):
use_natural_transcendental_func=False, # Different from default
)
else:
raise NotImplementedError(f"softmax_quantizer not supported: {self.softmax_quantizer}")
raise NotImplementedError(f"p_bmm_quantizer not supported: {self.p_bmm_quantizer}")
self.kitchen_attn_fn = KitchenFlashAttentionModule(
num_attention_heads=self.config.num_attention_heads,
@@ -117,6 +122,78 @@ class _QuantAttention(QuantModule):
qfa_params=qfa_params,
)
def _p_qdq_mode(self) -> str | None:
"""Map the p_bmm_quantizer config to a Triton P quant-dequant mode.
Returns "fp8" for per-tensor E4M3, "nvfp4" for dynamic E2M1 with
block-16 E4M3 scales, or None when the p_bmm_quantizer is disabled
or its format is not supported by the built-in Triton kernel
(e.g. MXFP8, which goes through kitchen).
"""
pq = self.p_bmm_quantizer
if not pq.is_enabled:
return None
if pq.is_fp8:
return "fp8"
# Only dynamic NVFP4 maps to the kernel, and only at block size 16 (the kernel
# hardcodes it). Static (calibrated) NVFP4 is excluded because is_nvfp4_dynamic
# requires dynamically-computed block scales.
if pq.is_nvfp4_dynamic and (pq.block_sizes or {}).get(-1, None) == 16:
return "nvfp4"
return None
def _triton_qdq_attention(self, p_qdq, query_states, key_states, value_states, **kwargs):
"""Quantized attention via the built-in Triton kernel (no kitchen required).
Fake quant-dequant of the softmax probabilities (P) is fused into the
flash-attention kernel; see ``p_qdq`` in
:func:`modelopt.torch.kernels.common.attention.triton_fa.attention`.
Inputs outside the kernel/wrapper envelope (sliding window, sinks,
softcapping, non-causal masks, ...) raise ``NotImplementedError``
instead of silently computing wrong attention; see
:func:`validate_triton_attention_envelope
<modelopt.torch.kernels.common.attention.hf_triton_attention.validate_triton_attention_envelope>`.
"""
if not TRITON_FA_AVAILABLE:
raise RuntimeError(
f"p_bmm_quantizer ({p_qdq}) requires the Triton attention kernel. "
"Install triton with `pip install triton` and run on a CUDA device."
)
assert (
triton_attention_forward is not None and validate_triton_attention_envelope is not None
)
attention_mask = kwargs.pop("attention_mask", None)
validate_triton_attention_envelope(self, query_states, key_states, attention_mask, **kwargs)
# Forward a user-set or calibrated per-tensor amax to the kernel, which
# converts it to the FP8 / NVFP4 scale. Without one, the kernel default
# amax of 1.0 applies -- the theoretical upper bound of the unnormalized
# P's amax (P lies in [0, 1]).
p_qdq_amax = None
pq_amax = getattr(self.p_bmm_quantizer, "_amax", None)
if pq_amax is not None:
if pq_amax.numel() != 1:
raise NotImplementedError(
"p_bmm_quantizer via the Triton attention kernel only supports a "
f"per-tensor (scalar) amax, got shape {tuple(pq_amax.shape)}."
)
p_qdq_amax = float(pq_amax)
scaling = kwargs.get("scaling")
if scaling is None:
scaling = query_states.shape[-1] ** -0.5
return triton_attention_forward(
self,
query_states,
key_states,
value_states,
attention_mask,
scaling,
p_qdq=p_qdq,
p_qdq_amax=p_qdq_amax,
)
@staticmethod
def _quantized_attention(
original_attention_interface,
@@ -127,12 +204,24 @@ class _QuantAttention(QuantModule):
*args,
**kwargs,
):
if kitchen is not None and self.kitchen_attn_fn is None:
self._init_kitchen_attn_fn()
query_states = self.q_bmm_quantizer(query_states)
key_states = self.k_bmm_quantizer(key_states)
value_states = self.v_bmm_quantizer(value_states)
# FP8 / NVFP4 P quant-dequant runs on the built-in Triton kernel
# and takes priority over kitchen (which handles MXFP8).
p_qdq = self._p_qdq_mode()
if p_qdq is not None:
# The attention interface passes attention_mask as the only
# positional argument after q/k/v; everything else is a kwarg.
if args:
kwargs["attention_mask"] = args[0]
return self._triton_qdq_attention(
p_qdq, query_states, key_states, value_states, **kwargs
)
if kitchen is not None and self.kitchen_attn_fn is None:
self._init_kitchen_attn_fn()
if not self.use_kitchen:
return original_attention_interface(
self, query_states, key_states, value_states, *args, **kwargs
@@ -239,9 +239,9 @@ class NVFP4QTensor(BaseQuantizedTensor):
e2m1_bounds = cls.get_e2m1_bounds(device)
ord = torch.searchsorted(e2m1_bounds, weight_abs, out_int32=True).to(torch.uint8)
# Efficiently check for rounding at odd-indexed bounds [0.75, 1.75, 2.5]
# Efficiently check for rounding at odd-indexed bounds [0.75, 1.75, 3.5]
# Only need to check bounds at indices 1, 3, 5
odd_bounds = e2m1_bounds[[1, 3, 5]] # [0.75, 1.75, 2.5]
odd_bounds = e2m1_bounds[[1, 3, 5]] # [0.75, 1.75, 3.5]
equals_odd_bounds = torch.any(weight_abs.unsqueeze(-1) == odd_bounds, dim=-1).to(
torch.uint8
)
@@ -14,7 +14,10 @@
# limitations under the License.
# QuantizerCfgList snippet that enables per-tensor FP8 E4M3 on attention q/k/v
# bmm and softmax quantizers. Pair with a model preset to add bmm2-output entries.
# bmm and the softmax-P quantizer. The softmax-P quantizer has a different attribute
# name per backend (diffusers attention: softmax_quantizer; HF attention:
# p_bmm_quantizer), so both are listed -- each is a no-op on the backend that lacks it.
# Pair with a model preset to add bmm2-output entries.
# modelopt-schema: modelopt.torch.quantization.config.QuantizerCfgListConfig
imports:
@@ -23,6 +26,11 @@ imports:
- quantizer_name: '*[qkv]_bmm_quantizer'
cfg:
$import: fp8
# softmax-P quantizer, diffusion attention (no-op for HF).
- quantizer_name: '*softmax_quantizer'
cfg:
$import: fp8
# softmax-P quantizer, HF (LLM) attention (no-op for diffusion).
- quantizer_name: '*p_bmm_quantizer'
cfg:
$import: fp8
+8
View File
@@ -236,6 +236,14 @@ extend-ignore = [
"N803",
"N806",
] # triton kernel style
"modelopt/torch/kernels/quantization/attention/*" = [
"N803",
"N806",
] # triton kernel style
"modelopt/torch/kernels/quantization/common/*" = [
"N803",
"N806",
] # triton kernel style
[tool.ruff.lint.pycodestyle]
max-line-length = 120 # Line length limit for comments and docstrings
@@ -0,0 +1,324 @@
# SPDX-FileCopyrightText: Copyright (c) 2026 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.
"""GPU tests for the softmax quant-dequant (P_QDQ) feature of the Triton FA kernel."""
import pytest
import torch
from conftest import make_qkv, make_varlen_meta, sdpa_reference
from modelopt.torch.kernels.common.attention import IS_AVAILABLE as TRITON_KERNEL_AVAILABLE
from modelopt.torch.quantization.qtensor.nvfp4_tensor import NVFP4QTensor, e2m1_values
from modelopt.torch.quantization.tensor_quant import fp8_eager
if TRITON_KERNEL_AVAILABLE:
from modelopt.torch.kernels.common.attention import attention
from modelopt.torch.kernels.common.attention.triton_fa import LOG2E
# The kernel runs with a single pinned config under pytest (see _FWD_CONFIGS):
# BLOCK_M=128, BLOCK_N=64. The tile-looped reference below relies on it.
BLOCK_N = 64
FP8_E4M3_MAX = 448.0
def _qdq_fp8(p, scale=1.0):
"""Per-tensor FP8 E4M3 qdq, mirroring the kernel's fp8_scalar_qdq.
Reuses modelopt's ``fp8_eager`` (native E4M3 cast). The kernel uses the
quantizer convention ``q = cast(p / scale) * scale``; fp8_eager parametrizes
by amax with ``scale = amax / 448``, so the equivalent amax is ``scale * 448``.
"""
return fp8_eager(p, torch.tensor(scale * FP8_E4M3_MAX, device=p.device))
def _fp4_round(x):
"""Round to the nearest E2M1 value, reusing modelopt's ``NVFP4QTensor._cast_fp4``.
``_cast_fp4`` implements the same round-half-to-even on the FP4 grid as the
kernel's Triton ``fp4_round_magnitude`` (verified bit-for-bit on the grid
boundaries and a dense random sweep), so this also guards that the two stay in
sync. ``_cast_fp4`` does an in-place ``abs_`` and returns uint4 indices, so pass
a clone and map the indices back through ``e2m1_values``.
"""
idx = NVFP4QTensor._cast_fp4(x.clone()).long()
return e2m1_values.to(device=x.device, dtype=x.dtype)[idx]
def _qdq_nvfp4(p, global_scale=1.0):
"""NVFP4 qdq with per-16 E4M3 block scales, mirroring _p_qdq_nvfp4.
p: [..., n] with n % 16 == 0 and p >= 0. The per-block E4M3 scale is the
kernel's ``fp8_quantize_scale(block_amax, global_scale)``, which equals
``fp8_eager(block_amax / 6, amax=448 * global_scale)`` — so the FP8 scale
quantization is reused from modelopt rather than re-spelled here.
"""
shape = p.shape
g = p.reshape(*shape[:-1], shape[-1] // 16, 16)
block_amax = g.amax(dim=-1, keepdim=True) # p >= 0, so max == amax
scale = fp8_eager(block_amax / 6.0, torch.tensor(FP8_E4M3_MAX * global_scale, device=p.device))
scale = torch.where(scale == 0.0, torch.ones_like(scale), scale)
q = _fp4_round(g / scale) * scale
return q.reshape(shape)
def _apply_qdq(p, mode, qdq_scale=1.0):
if mode == "fp8":
return _qdq_fp8(p, qdq_scale)
assert mode == "nvfp4"
return _qdq_nvfp4(p, qdq_scale)
def qdq_attention_reference(q, k, v, scale, mode, is_causal=True, amax=1.0):
"""Tile-looped online-softmax reference replicating kernel P_QDQ semantics.
Single sequence: q [s, h, d], k/v [s_kv, h_kv, d] (fp16). Walks KV tiles
of BLOCK_N exactly like the kernel, keeps the softmax denominator
unquantized, applies qdq to the unnormalized p of each tile, and mirrors
the kernel's ``p.to(v.dtype)`` cast before the P @ V dot.
Returns [s, h, d] float32.
``amax`` mirrors the kernel's ``p_qdq_amax`` and is converted to the same
per-mode scale the wrapper uses: ``amax / 448`` (FP8) or ``amax / (6 * 448)``
(NVFP4 global scale).
"""
qdq_scale = amax / 448.0 if mode == "fp8" else amax / (6.0 * 448.0)
s, h, d = q.shape
s_kv = k.shape[0]
r = h // k.shape[1]
kk = k.repeat_interleave(r, dim=1) if r > 1 else k
vv = v.repeat_interleave(r, dim=1) if r > 1 else v
# Scores in the kernel's scaled log2 space: Q K^T * scale * log2(e)
t = torch.einsum("qhd,khd->hqk", q.float(), kk.float()) * (scale * LOG2E)
if is_causal:
offset = s_kv - s
causal = (
torch.arange(s, device=q.device)[:, None] + offset
>= torch.arange(s_kv, device=q.device)[None, :]
)
t = t.masked_fill(~causal[None], float("-inf"))
row_max = torch.full((h, s), float("-inf"), device=q.device)
row_sum = torch.zeros(h, s, device=q.device)
acc = torch.zeros(h, s, d, device=q.device)
for start in range(0, s_kv, BLOCK_N):
tile = t[:, :, start : start + BLOCK_N]
m_new = torch.maximum(row_max, tile.amax(dim=-1))
p = torch.exp2(tile - m_new[..., None])
l_new = p.sum(dim=-1)
corr = torch.exp2(row_max - m_new)
row_sum = row_sum * corr + l_new
acc = acc * corr[..., None]
p = _apply_qdq(p, mode, qdq_scale)
# Kernel casts p to v.dtype for the BMM2 dot
p = p.to(v.dtype).float()
acc = acc + torch.einsum("hqk,khd->hqd", p, vv[start : start + BLOCK_N].float())
row_max = m_new
out = acc / row_sum[..., None]
return out.permute(1, 0, 2)
@pytest.mark.skipif(not TRITON_KERNEL_AVAILABLE, reason="Need CUDA + triton")
class TestSoftmaxQdqForward:
"""Forward correctness of FP8/NVFP4 softmax quant-dequant."""
@pytest.mark.parametrize("mode", ["fp8", "nvfp4"])
def test_prefill_matches_tile_reference(self, mode):
"""Kernel qdq output matches the tile-looped torch reference."""
seq_len, num_heads, num_kv_heads, head_dim = 128, 4, 2, 64
scale = 1.0 / (head_dim**0.5)
torch.manual_seed(7)
q, k, v = make_qkv(seq_len, num_heads, num_kv_heads, head_dim, dtype=torch.float16)
locs, lens = make_varlen_meta([seq_len])
o = attention(q, k, v, locs, lens, seq_len, softmax_scale=scale, p_qdq=mode)
ref = qdq_attention_reference(q, k, v, scale, mode)
torch.testing.assert_close(o.float(), ref, rtol=5e-3, atol=5e-3)
# The feature must actually change the output vs dense attention.
o_dense = attention(q, k, v, locs, lens, seq_len, softmax_scale=scale)
assert not torch.equal(o, o_dense)
@pytest.mark.parametrize("mode", ["fp8", "nvfp4"])
def test_varlen_partial_tiles(self, mode):
"""Variable-length batch with partial KV tiles (seq % BLOCK_N != 0)."""
seq_lens = [96, 80]
total = sum(seq_lens)
num_heads, num_kv_heads, head_dim = 4, 2, 64
scale = 1.0 / (head_dim**0.5)
torch.manual_seed(11)
q, k, v = make_qkv(total, num_heads, num_kv_heads, head_dim, dtype=torch.float16)
locs, lens = make_varlen_meta(seq_lens)
o = attention(q, k, v, locs, lens, max(seq_lens), softmax_scale=scale, p_qdq=mode)
for b, n in enumerate(seq_lens):
s = int(locs[b].item())
ref = qdq_attention_reference(q[s : s + n], k[s : s + n], v[s : s + n], scale, mode)
torch.testing.assert_close(o[s : s + n].float(), ref, rtol=5e-3, atol=5e-3)
@pytest.mark.parametrize(("mode", "tol"), [("fp8", 5e-2), ("nvfp4", 0.25)])
def test_qdq_close_to_dense(self, mode, tol):
"""Quantization is an approximation: output stays near dense attention."""
seq_len, num_heads, num_kv_heads, head_dim = 128, 4, 2, 64
scale = 1.0 / (head_dim**0.5)
torch.manual_seed(13)
q, k, v = make_qkv(seq_len, num_heads, num_kv_heads, head_dim, dtype=torch.float16)
locs, lens = make_varlen_meta([seq_len])
o = attention(q, k, v, locs, lens, seq_len, softmax_scale=scale, p_qdq=mode)
ref = sdpa_reference(q, k, v, locs, lens)
torch.testing.assert_close(o, ref, rtol=tol, atol=tol)
@pytest.mark.parametrize("mode", ["fp8", "nvfp4"])
def test_decode(self, mode):
"""Decode (seq_q=1 vs KV cache) matches the non-causal tile reference."""
batch, num_heads, num_kv_heads, head_dim = 2, 4, 2, 64
seq_lens_k = [80, 64]
scale = 1.0 / (head_dim**0.5)
torch.manual_seed(17)
q = torch.randn(batch, num_heads, head_dim, device="cuda", dtype=torch.float16)
total_kv = sum(seq_lens_k)
k = torch.randn(total_kv, num_kv_heads, head_dim, device="cuda", dtype=torch.float16)
v = torch.randn(total_kv, num_kv_heads, head_dim, device="cuda", dtype=torch.float16)
locs_k, lens_k = make_varlen_meta(seq_lens_k)
out = attention(
q,
k,
v,
torch.arange(batch, device="cuda", dtype=torch.int32),
torch.ones(batch, device="cuda", dtype=torch.int32),
1,
is_causal=False,
softmax_scale=scale,
b_start_loc_k=locs_k,
b_seq_len_k=lens_k,
max_input_len_k=max(seq_lens_k),
p_qdq=mode,
)
for b, n in enumerate(seq_lens_k):
s = int(locs_k[b].item())
ref = qdq_attention_reference(
q[b : b + 1], k[s : s + n], v[s : s + n], scale, mode, is_causal=False
)
torch.testing.assert_close(out[b : b + 1].float(), ref, rtol=5e-3, atol=5e-3)
@pytest.mark.parametrize(
("mode", "amax"),
[
# Non-power-of-2 amax vs the default of 1.0: a pure power-of-2 change
# is a fixed point of FP quantization (exponent shift only) and would
# not change the output, so use non-power-of-2 values. amax < 1
# (0.3) also exercises FP8 saturation (P entries above amax clamp).
("fp8", 0.3),
("fp8", 3.0),
("nvfp4", 0.7),
],
)
def test_custom_amax_matches_tile_reference(self, mode, amax):
"""User-supplied p_qdq_amax changes the grid and matches the reference."""
seq_len, num_heads, num_kv_heads, head_dim = 128, 4, 2, 64
scale = 1.0 / (head_dim**0.5)
torch.manual_seed(31)
q, k, v = make_qkv(seq_len, num_heads, num_kv_heads, head_dim, dtype=torch.float16)
locs, lens = make_varlen_meta([seq_len])
o = attention(
q,
k,
v,
locs,
lens,
seq_len,
softmax_scale=scale,
p_qdq=mode,
p_qdq_amax=amax,
)
ref = qdq_attention_reference(q, k, v, scale, mode, amax=amax)
torch.testing.assert_close(o.float(), ref, rtol=5e-3, atol=5e-3)
# The amax knob must actually change the quantization grid vs the default (amax=1).
o_default = attention(q, k, v, locs, lens, seq_len, softmax_scale=scale, p_qdq=mode)
assert not torch.equal(o, o_default)
def test_invalid_amax_raises(self):
q, k, v = make_qkv(8, 2, 2, 32, dtype=torch.float16)
locs, lens = make_varlen_meta([8])
with pytest.raises(ValueError, match="p_qdq_amax"):
attention(q, k, v, locs, lens, 8, p_qdq="fp8", p_qdq_amax=0.0)
def test_composes_with_skip_softmax(self):
"""p_qdq composes with the skip-softmax feature in one launch."""
seq_len, num_heads, num_kv_heads, head_dim = 256, 4, 2, 64
scale = 1.0 / (head_dim**0.5)
torch.manual_seed(19)
q, k, v = make_qkv(seq_len, num_heads, num_kv_heads, head_dim, dtype=torch.float16)
locs, lens = make_varlen_meta([seq_len])
o = attention(
q,
k,
v,
locs,
lens,
seq_len,
softmax_scale=scale,
p_qdq="fp8",
skip_softmax_threshold=1e-3,
)
ref = sdpa_reference(q, k, v, locs, lens)
torch.testing.assert_close(o, ref, rtol=5e-2, atol=5e-2)
def test_invalid_mode_raises(self):
q, k, v = make_qkv(8, 2, 2, 32, dtype=torch.float16)
locs, lens = make_varlen_meta([8])
with pytest.raises(ValueError, match="p_qdq"):
attention(q, k, v, locs, lens, 8, p_qdq="int8")
@pytest.mark.skipif(not TRITON_KERNEL_AVAILABLE, reason="Need CUDA + triton")
class TestSoftmaxQdqBackward:
"""Backward uses the straight-through estimator (no qdq re-applied)."""
@pytest.mark.parametrize("mode", ["fp8", "nvfp4"])
def test_ste_gradients_close_to_dense(self, mode):
seq_len, num_heads, num_kv_heads, head_dim = 128, 4, 2, 32
scale = 1.0 / (head_dim**0.5)
torch.manual_seed(23)
q, k, v = make_qkv(seq_len, num_heads, num_kv_heads, head_dim, dtype=torch.float32)
locs, lens = make_varlen_meta([seq_len])
q1, k1, v1 = (t.clone().requires_grad_(True) for t in (q, k, v))
attention(q1, k1, v1, locs, lens, seq_len, softmax_scale=scale, p_qdq=mode).sum().backward()
q2, k2, v2 = (t.clone().requires_grad_(True) for t in (q, k, v))
attention(q2, k2, v2, locs, lens, seq_len, softmax_scale=scale).sum().backward()
# The backward recomputes the unquantized P (straight-through estimator),
# so gradients match dense attention up to the quantization perturbation
# that enters through the saved output (delta = rowsum(O * dO)). A few
# individual elements can shift; the overall norm error stays small.
for g_qdq, g_dense in ((q1.grad, q2.grad), (k1.grad, k2.grad), (v1.grad, v2.grad)):
assert torch.isfinite(g_qdq).all()
rel_err = (g_qdq - g_dense).norm() / g_dense.norm()
assert rel_err < 5e-2, f"relative gradient error too large: {rel_err:.4f}"
@@ -30,7 +30,8 @@ from conftest import make_qkv, make_varlen_meta
from modelopt.torch.kernels.common.attention import IS_AVAILABLE as TRITON_KERNEL_AVAILABLE
if TRITON_KERNEL_AVAILABLE:
from modelopt.torch.kernels.common.attention import attention, attention_calibrate
from modelopt.torch.kernels.common.attention import attention
from modelopt.torch.kernels.sparsity.attention.calibrate import attention_calibrate
@pytest.mark.skipif(not TRITON_KERNEL_AVAILABLE, reason="Need CUDA + triton")
@@ -0,0 +1,241 @@
# 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.
"""GPU tests for HF attention quantization (softmax-P qdq via the Triton kernel).
The CPU-only config-detection test lives in the mirror unit test
(tests/unit/torch/quantization/plugins/test_attention_quant.py::test_p_qdq_mode_detection).
"""
import inspect
import pytest
import torch
from transformers import LlamaConfig
from transformers.models.llama.modeling_llama import LlamaAttention
try:
import kitchen
except ImportError:
kitchen = None
from modelopt.torch.kernels.common.attention import IS_AVAILABLE as TRITON_FA_AVAILABLE
from modelopt.torch.quantization.plugins.huggingface import _QuantAttention
pytest.importorskip("transformers")
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
@pytest.mark.skipif(not TRITON_FA_AVAILABLE, reason="Triton attention kernel unavailable")
def test_p_qdq_fa():
"""FP8/NVFP4 p_bmm_quantizer runs on the built-in Triton kernel (no kitchen)."""
batch_size, num_q_heads, num_kv_heads, seqlen, head_dim = 2, 4, 2, 32, 64
quant_attention = _make_quant_attention(num_q_heads=num_q_heads, num_kv_heads=num_kv_heads)
for name in ("q_bmm_quantizer", "k_bmm_quantizer", "v_bmm_quantizer"):
getattr(quant_attention, name).disable()
torch.manual_seed(29)
q_states = torch.randn(
batch_size, num_q_heads, seqlen, head_dim, dtype=torch.bfloat16, device="cuda"
)
k_states = torch.randn(
batch_size, num_kv_heads, seqlen, head_dim, dtype=torch.bfloat16, device="cuda"
)
v_states = torch.randn(
batch_size, num_kv_heads, seqlen, head_dim, dtype=torch.bfloat16, device="cuda"
)
module = inspect.getmodule(quant_attention.get_attn_type(quant_attention))
orig_attn_fn = module.ALL_ATTENTION_FUNCTIONS["sdpa"]
def run():
return quant_attention._quantized_attention(
orig_attn_fn,
quant_attention,
q_states,
k_states,
v_states,
attention_mask=None,
)[0]
quant_attention.p_bmm_quantizer.disable()
expected = run()
quant_attention.p_bmm_quantizer.enable()
for num_bits, block_sizes, tol in [
((4, 3), None, 0.1), # FP8
((2, 1), {-1: 16, "type": "dynamic", "scale_bits": (4, 3)}, 0.4), # NVFP4
]:
quant_attention.p_bmm_quantizer.num_bits = num_bits
quant_attention.p_bmm_quantizer.block_sizes = block_sizes
output = run()
assert output.shape == expected.shape
assert not torch.equal(output, expected), "softmax qdq should perturb the output"
torch.testing.assert_close(output, expected, atol=tol, rtol=tol)
# A user-set (or calibrated) per-tensor amax on the quantizer overrides the
# kernel's default of 1.0 and changes the quantization grid. Use a non-power-of-2
# amax (3.0): a power-of-2 change is a fixed point of FP quant and wouldn't differ.
quant_attention.p_bmm_quantizer.num_bits = (4, 3)
quant_attention.p_bmm_quantizer.block_sizes = None
out_default = run()
quant_attention.p_bmm_quantizer.amax = torch.tensor(3.0)
out_amax = run()
assert not torch.equal(out_amax, out_default), "user-set amax should change the output"
torch.testing.assert_close(out_amax, expected, atol=0.1, rtol=0.1)
@pytest.mark.skipif(not TRITON_FA_AVAILABLE, reason="Triton attention kernel unavailable")
def test_p_qdq_unsupported_cases_raise():
"""The Triton qdq dispatch rejects attention semantics the kernel cannot honor."""
batch_size, num_q_heads, num_kv_heads, seqlen, head_dim = 2, 4, 2, 32, 64
quant_attention = _make_quant_attention(num_q_heads=num_q_heads, num_kv_heads=num_kv_heads)
for name in ("q_bmm_quantizer", "k_bmm_quantizer", "v_bmm_quantizer"):
getattr(quant_attention, name).disable()
quant_attention.p_bmm_quantizer.num_bits = (4, 3) # FP8 mode
def make_qkv(seq_q=seqlen, seq_k=seqlen):
q = torch.randn(batch_size, num_q_heads, seq_q, head_dim, device="cuda")
k = torch.randn(batch_size, num_kv_heads, seq_k, head_dim, device="cuda")
v = torch.randn(batch_size, num_kv_heads, seq_k, head_dim, device="cuda")
return q, k, v
def run(seq_q=seqlen, seq_k=seqlen, **kwargs):
q, k, v = make_qkv(seq_q, seq_k)
return quant_attention._quantized_attention(None, quant_attention, q, k, v, **kwargs)
with pytest.raises(NotImplementedError, match="sliding-window"):
run(sliding_window=128)
with pytest.raises(NotImplementedError, match="attention sinks"):
run(s_aux=torch.zeros(num_q_heads, device="cuda"))
with pytest.raises(NotImplementedError, match="softcapping"):
run(softcap=50.0)
with pytest.raises(NotImplementedError, match="non-causal"):
run(is_causal=False)
with pytest.raises(NotImplementedError, match="dropout"):
run(dropout=0.1)
with pytest.raises(NotImplementedError, match="KV cache"):
run(seq_q=4, seq_k=20) # chunked prefill / assisted decoding
with pytest.raises(NotImplementedError, match="cached decode"):
# FA2-style [batch, kv_len] padding mask during single-token decode
run(seq_q=1, seq_k=20, attention_mask=torch.ones(batch_size, 20, device="cuda"))
with pytest.raises(NotImplementedError, match="left-padded"):
left_pad = torch.ones(batch_size, seqlen, device="cuda")
left_pad[0, :5] = 0
run(attention_mask=left_pad)
with pytest.raises(NotImplementedError, match="non-contiguously padded"):
# A hole in the mask would sum to the wrong per-sequence length.
holey = torch.ones(batch_size, seqlen, device="cuda")
holey[0, 3] = 0
run(attention_mask=holey)
with pytest.raises(NotImplementedError, match="padding or non-causal"):
padded_4d = torch.zeros(batch_size, 1, seqlen, seqlen, device="cuda")
padded_4d[..., -4:] = torch.finfo(torch.float32).min # last 4 kv positions padded
run(attention_mask=padded_4d)
# A purely causal 4D mask is safe to ignore: the kernel masks causally itself.
causal_4d = torch.zeros(batch_size, 1, seqlen, seqlen, device="cuda")
causal_4d.masked_fill_(
torch.triu(torch.ones(seqlen, seqlen, dtype=torch.bool, device="cuda"), diagonal=1),
torch.finfo(torch.float32).min,
)
output = run(attention_mask=causal_4d)[0]
assert torch.isfinite(output).all()
@pytest.mark.skipif(kitchen is None, reason="kitchen is not installed.")
def test_kitchen_fa():
batch_size = 2
num_q_heads = 4
num_kv_heads = 2
seqlen = 8
hidden_size = 128
config = LlamaConfig(
hidden_size=hidden_size,
num_attention_heads=num_q_heads,
num_key_value_heads=num_kv_heads,
)
original_attention = LlamaAttention(config, layer_idx=0)
q_states = torch.randn(
batch_size, num_q_heads, seqlen, hidden_size, dtype=torch.bfloat16, device="cuda"
)
k_states = torch.randn(
batch_size, num_kv_heads, seqlen, hidden_size, dtype=torch.bfloat16, device="cuda"
)
v_states = torch.randn(
batch_size, num_kv_heads, seqlen, hidden_size, dtype=torch.bfloat16, device="cuda"
)
# Convert it to _QuantAttention using the convert() class method
quant_attention = _QuantAttention.convert(original_attention)
quant_attention.config._attn_implementation = "sdpa"
assert hasattr(quant_attention, "q_bmm_quantizer")
assert hasattr(quant_attention, "k_bmm_quantizer")
assert hasattr(quant_attention, "v_bmm_quantizer")
assert hasattr(quant_attention, "p_bmm_quantizer")
quant_attention.p_bmm_quantizer.disable()
module = inspect.getmodule(quant_attention.get_attn_type(quant_attention))
orig_attn_fn = module.ALL_ATTENTION_FUNCTIONS["sdpa"]
output = quant_attention._quantized_attention(
orig_attn_fn,
quant_attention,
q_states,
k_states,
v_states,
attention_mask=None,
)
expected = output[0]
config = LlamaConfig(
hidden_size=hidden_size,
num_attention_heads=num_q_heads,
num_key_value_heads=num_kv_heads,
)
original_attention = LlamaAttention(config, layer_idx=0)
quant_attention = _QuantAttention.convert(original_attention)
quant_attention.config._attn_implementation = "sdpa"
quant_attention.p_bmm_quantizer.num_bits = (4, 3)
quant_attention.p_bmm_quantizer.block_sizes = {
-1: 32,
"type": "dynamic",
"scale_bits": (8, 0),
}
output = quant_attention._quantized_attention(
None,
quant_attention,
q_states,
k_states,
v_states,
attention_mask=None,
)
diff = (expected - output[0]).abs()
assert torch.allclose(expected, output[0], atol=0.75, rtol=0.75), (
f"{diff.max().item(), diff.mean().item(), diff.std().item()}"
)
@@ -13,8 +13,6 @@
# See the License for the specific language governing permissions and
# limitations under the License.
import inspect
import pytest
import torch
import torch.nn as nn
@@ -23,11 +21,6 @@ from _test_utils.torch.transformers_models import get_tiny_bert, get_tiny_llama,
from transformers import LlamaConfig
from transformers.models.llama.modeling_llama import LlamaAttention
try:
import kitchen
except ImportError:
kitchen = None
import modelopt.torch.quantization as mtq
from modelopt.torch.quantization.plugins.huggingface import _QuantAttention
@@ -63,7 +56,7 @@ class SDPAAttention(nn.Module):
kv_cache_config = {
"quant_cfg": [
{"quantizer_name": "*[kv]_bmm_quantizer", "cfg": {"num_bits": 4}, "enable": True},
{"quantizer_name": "*softmax_quantizer", "enable": False},
{"quantizer_name": "*p_bmm_quantizer", "enable": False},
],
"algorithm": "max",
}
@@ -159,75 +152,48 @@ def test_kv_quant_bert():
assert output.end_logits is not None
@pytest.mark.skipif(kitchen is None, reason="kitchen is not installed.")
def test_kitchen_fa():
batch_size = 2
num_q_heads = 4
num_kv_heads = 2
seqlen = 8
hidden_size = 128
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,
)
original_attention = LlamaAttention(config, layer_idx=0)
q_states = torch.randn(
batch_size, num_q_heads, seqlen, hidden_size, dtype=torch.bfloat16, device="cuda"
)
k_states = torch.randn(
batch_size, num_kv_heads, seqlen, hidden_size, dtype=torch.bfloat16, device="cuda"
)
v_states = torch.randn(
batch_size, num_kv_heads, seqlen, hidden_size, dtype=torch.bfloat16, device="cuda"
)
# Convert it to _QuantAttention using the convert() class method
quant_attention = _QuantAttention.convert(original_attention)
quant_attention = _QuantAttention.convert(LlamaAttention(config, layer_idx=0))
quant_attention.config._attn_implementation = "sdpa"
assert hasattr(quant_attention, "q_bmm_quantizer")
assert hasattr(quant_attention, "k_bmm_quantizer")
assert hasattr(quant_attention, "v_bmm_quantizer")
assert hasattr(quant_attention, "softmax_quantizer")
quant_attention.softmax_quantizer.disable()
module = inspect.getmodule(quant_attention.get_attn_type(quant_attention))
orig_attn_fn = module.ALL_ATTENTION_FUNCTIONS["sdpa"]
return quant_attention
output = quant_attention._quantized_attention(
orig_attn_fn,
quant_attention,
q_states,
k_states,
v_states,
attention_mask=None,
)
expected = output[0]
config = LlamaConfig(
hidden_size=hidden_size,
num_attention_heads=num_q_heads,
num_key_value_heads=num_kv_heads,
)
original_attention = LlamaAttention(config, layer_idx=0)
quant_attention = _QuantAttention.convert(original_attention)
quant_attention.config._attn_implementation = "sdpa"
quant_attention.softmax_quantizer.num_bits = (4, 3)
quant_attention.softmax_quantizer.block_sizes = {
-1: 32,
"type": "dynamic",
"scale_bits": (8, 0),
}
output = quant_attention._quantized_attention(
None,
quant_attention,
q_states,
k_states,
v_states,
attention_mask=None,
)
diff = (expected - output[0]).abs()
assert torch.allclose(expected, output[0], atol=0.75, rtol=0.75), (
f"{diff.max().item(), diff.mean().item(), diff.std().item()}"
)
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