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
1 change: 1 addition & 0 deletions qa/L0_pytorch_unittest/test.sh
Original file line number Diff line number Diff line change
Expand Up @@ -60,6 +60,7 @@ NVTE_ALLOW_UNSAFE_PICKLE_EXTRA_STATE=1 python3 -m pytest --tb=auto --junitxml=$X
python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_flex_attention.xml $TE_PATH/tests/pytorch/attention/test_flex_attention.py || test_fail "test_flex_attention.py"
NVTE_ALLOW_UNSAFE_PICKLE_EXTRA_STATE=1 NVTE_ALLOW_NONDETERMINISTIC_ALGO=0 python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_attention_deterministic.xml $TE_PATH/tests/pytorch/attention/test_attention.py || test_fail "NVTE_ALLOW_NONDETERMINISTIC_ALGO=0 test_attention.py"
python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_linear_mxfp8_attention.xml $TE_PATH/tests/pytorch/attention/test_linear_mxfp8_attention.py || test_fail "test_linear_mxfp8_attention.py"
python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_fused_mla_q_uproj.xml $TE_PATH/tests/pytorch/attention/test_fused_mla_q_uproj.py || test_fail "test_fused_mla_q_uproj.py"
python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_kv_cache.xml $TE_PATH/tests/pytorch/attention/test_kv_cache.py || test_fail "test_kv_cache.py"
python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_cu_seqlens_cache.xml $TE_PATH/tests/pytorch/attention/test_cu_seqlens_cache.py || test_fail "test_cu_seqlens_cache.py"
python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_hf_integration.xml $TE_PATH/tests/pytorch/test_hf_integration.py || test_fail "test_hf_integration.py"
Expand Down
137 changes: 137 additions & 0 deletions tests/pytorch/attention/test_fused_mla_q_uproj.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,137 @@
# Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved.
#
# See LICENSE for license information.

"""Unit tests for FusedMLAQUpProjRopeQuant.

Run:
pytest tests/pytorch/attention/test_fused_mla_q_uproj.py -v
"""

import pytest
import torch

import transformer_engine.pytorch # registers transformer_engine_torch
import transformer_engine_torch as tex
from transformer_engine.pytorch.attention import FusedMLAQUpProjRopeQuant
from transformer_engine.pytorch.tensor.mxfp8_tensor import MXFP8Quantizer, MXFP8Tensor

# DSv3 671B MLA dims
NUM_HEADS = 128
HEAD_DIM_NOPE = 128
HEAD_DIM_ROPE = 64
HEAD_DIM = HEAD_DIM_NOPE + HEAD_DIM_ROPE # 192
Q_LORA_RANK = 1536
PROJ_DIM = NUM_HEADS * HEAD_DIM # 24576

SEED = 42

fused_supported, reason_not_supported = (
(True, "")
if FusedMLAQUpProjRopeQuant.is_supported()
else (
False,
(
"FusedMLAQUpProjRopeQuant.is_supported() returned False "
"(SM100+, cudnn-frontend >= 1.27.0, and NVTE_FUSED_MLA_Q_UPROJ=1 required)"
),
)
)


def _dequantize_fused_output(query: MXFP8Tensor, s: int, b: int) -> torch.Tensor:
"""Dequantize the rowwise fused output to bf16 [s, b, nh, head_dim].

TE's C++ dequantize kernel requires 2D layout, so reshape before calling dequantize().
"""
tokens = s * b
q_2d = MXFP8Tensor(
shape=(tokens, PROJ_DIM),
dtype=torch.bfloat16,
rowwise_data=query._rowwise_data.view(tokens, PROJ_DIM),
rowwise_scale_inv=query._rowwise_scale_inv.view(tokens, PROJ_DIM // 32),
columnwise_data=None,
columnwise_scale_inv=None,
quantizer=query._quantizer,
requires_grad=False,
fp8_dtype=query._fp8_dtype,
with_gemm_swizzled_scales=False,
)
return q_2d.dequantize().to(torch.bfloat16).view(s, b, NUM_HEADS, HEAD_DIM)


def _reference_q_uproj(
x: torch.Tensor,
w_mxfp8: MXFP8Tensor,
cos: torch.Tensor,
sin: torch.Tensor,
s: int,
b: int,
) -> torch.Tensor:
"""Unfused bf16 reference: dequantize-then-GEMM + RoPE. Returns [s, b, nh, head_dim] bf16."""
x_dq = (
MXFP8Quantizer(fp8_dtype=tex.DType.kFloat8E4M3, rowwise=True, columnwise=False)(x)
.dequantize()
.to(torch.bfloat16)
)
w_dq = w_mxfp8.dequantize().to(torch.bfloat16)
out = (x_dq @ w_dq.t()).view(s, b, NUM_HEADS, HEAD_DIM)

q_nope = out[..., :HEAD_DIM_NOPE]
q_rope = out[..., HEAD_DIM_NOPE:]
cos_ = cos[:, None, None, :].to(q_rope.dtype)
sin_ = sin[:, None, None, :].to(q_rope.dtype)
half = HEAD_DIM_ROPE // 2
x1, x2 = q_rope[..., 0::2], q_rope[..., 1::2]
q_rope_out = torch.cat(
[
x1 * cos_[..., :half] - x2 * sin_[..., :half],
x2 * cos_[..., half:] + x1 * sin_[..., half:],
],
dim=-1,
)
return torch.cat([q_nope, q_rope_out], dim=-1)


def _build_rope_tables(tokens: int, device: torch.device) -> tuple[torch.Tensor, torch.Tensor]:
inv_freq = 1.0 / (
10000
** (torch.arange(0, HEAD_DIM_ROPE, 2, dtype=torch.float32, device=device) / HEAD_DIM_ROPE)
)
freqs = torch.cat(
[torch.outer(torch.arange(tokens, device=device, dtype=torch.float32), inv_freq)] * 2,
dim=-1,
)
return freqs.cos().to(torch.bfloat16), freqs.sin().to(torch.bfloat16)


@pytest.mark.skipif(not fused_supported, reason=reason_not_supported)
@pytest.mark.parametrize("tokens", [256])
def test_fused_mla_q_uproj(tokens: int) -> None:
"""Forward numerics and x_saved properties for FusedMLAQUpProjRopeQuant.run().

Full forward+backward autograd testing (via _FusedMLAQUpProjFunction) lives in
Megatron-Core.
"""
s, b = tokens, 1
device = torch.device("cuda")
torch.manual_seed(SEED)
torch.cuda.manual_seed(SEED)

x = torch.randn(tokens, Q_LORA_RANK, dtype=torch.bfloat16, device=device)
w = MXFP8Quantizer(fp8_dtype=tex.DType.kFloat8E4M3, rowwise=True, columnwise=False)(
torch.randn(PROJ_DIM, Q_LORA_RANK, dtype=torch.bfloat16, device=device)
)
cos, sin = _build_rope_tables(tokens, device)

query, x_saved = FusedMLAQUpProjRopeQuant.run(x, w, cos, sin, s, b)

# Forward numerics: FP8 GEMM + output quantize introduce ~10% relative error.
fused_dq = _dequantize_fused_output(query, s, b)
ref_dq = _reference_q_uproj(x, w, cos, sin, s, b)
torch.testing.assert_close(fused_dq, ref_dq, atol=0.5, rtol=0.1)

# x_saved: must be MXFP8 with only columnwise data retained for wgrad.
assert isinstance(x_saved, MXFP8Tensor)
assert x_saved._columnwise_data is not None, "x_saved must retain columnwise data for wgrad"
assert x_saved._rowwise_data is None, "x_saved rowwise data should be dropped after forward"
1 change: 1 addition & 0 deletions transformer_engine/pytorch/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -29,6 +29,7 @@
from transformer_engine.pytorch.module import destroy_ub
from transformer_engine.pytorch.module import UserBufferQuantizationMode
from transformer_engine.pytorch.attention import DotProductAttention
from transformer_engine.pytorch.attention import FusedMLAQUpProjRopeQuant
from transformer_engine.pytorch.attention import MultiheadAttention
from transformer_engine.pytorch.attention import InferenceParams
from transformer_engine.pytorch.attention import RotaryPositionEmbedding
Expand Down
2 changes: 2 additions & 0 deletions transformer_engine/pytorch/attention/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,12 +5,14 @@
"""Python interface for attention"""

from .dot_product_attention import DotProductAttention
from .fused_mla_q_uproj import FusedMLAQUpProjRopeQuant
from .multi_head_attention import MultiheadAttention
from .inference import InferenceParams
from .rope import RotaryPositionEmbedding

__all__ = [
"DotProductAttention",
"FusedMLAQUpProjRopeQuant",
"MultiheadAttention",
"InferenceParams",
"RotaryPositionEmbedding",
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -1363,6 +1363,7 @@ def forward(
deterministic,
softmax_offset,
fp8_output,
bf16_backward,
layer_number,
return_max_logit,
packed_qkv=None,
Expand Down Expand Up @@ -1427,6 +1428,9 @@ def forward(
# fp8_dtype = tex.DType.kFloat8E4M3
if is_input_fp8:
q_fp8, k_fp8, v_fp8 = q, k, v

if fp8_recipe.mxfp8():
qkv_scale_inv_format = "bhsd" # Same as what combine_and_quantize would give
else:
q_fp8, k_fp8, v_fp8, qkv_layout, qkv_scale_inv_format = combine_and_quantize(
qkv_layout,
Expand Down Expand Up @@ -1602,6 +1606,8 @@ def forward(

ctx.is_input_fp8 = is_input_fp8
ctx.is_output_fp8 = is_output_fp8
# Return dQ/dK/dV in bf16 even if is_input_fp8
ctx.bf16_backward = bf16_backward

tensors_to_save, tensor_objects = prepare_for_saving(
*fp8_tensors,
Expand Down Expand Up @@ -1860,7 +1866,8 @@ def backward(ctx, d_out, *_args):
# dq, dk, dv: torch.Tensor; dtype = torch.float16 or torch.bfloat16
dq, dk, dv = dq_, dk_, dv_
is_quantized_tensor = isinstance(dq_, QuantizedTensorStorage)
if is_quantized_tensor and not ctx.is_input_fp8:

if is_quantized_tensor and (not ctx.is_input_fp8 or ctx.bf16_backward):
# return in F16
dq, dk, dv = combine_and_dequantize(
ctx.dqkv_layout,
Expand All @@ -1869,7 +1876,7 @@ def backward(ctx, d_out, *_args):
dv_,
src_nominal_dtype=dq_.dtype,
)
if not is_quantized_tensor and ctx.is_input_fp8:
if not is_quantized_tensor and ctx.is_input_fp8 and not ctx.bf16_backward:
# return in FP8
dq, dk, dv, _, _ = combine_and_quantize(
ctx.dqkv_layout, dq_, dk_, dv_, ctx.dQKV_quantizer
Expand Down Expand Up @@ -1968,6 +1975,7 @@ def backward(ctx, d_out, *_args):
None,
None, # packed_qkv
None, # packed_kv
None,
)


Expand Down Expand Up @@ -2064,6 +2072,7 @@ def forward(
score_mod_bprop_tensors: Optional[Dict[str, torch.Tensor]] = None,
packed_qkv: Optional[torch.Tensor] = None,
packed_kv: Optional[torch.Tensor] = None,
bf16_backward: bool = False,
) -> torch.Tensor:
"""fused attention fprop"""
assert (
Expand Down Expand Up @@ -2280,6 +2289,7 @@ def forward(
self.deterministic,
softmax_offset,
fp8_output,
bf16_backward,
self.layer_number,
self.return_max_logit,
packed_qkv,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -32,6 +32,7 @@
Float8BlockScalingRecipeState,
)
from transformer_engine.pytorch.tensor.storage.float8_tensor_storage import Float8TensorStorage
from transformer_engine.pytorch.tensor.storage.mxfp8_tensor_storage import MXFP8TensorStorage
from transformer_engine.pytorch.module.base import TransformerEngineBaseModule
from transformer_engine.pytorch.export import is_in_onnx_export_mode
from transformer_engine.pytorch.constants import AttnMaskTypes, AttnTypes, dist_group_type, DType
Expand Down Expand Up @@ -1347,6 +1348,7 @@ def forward(
inference_params: Optional[InferenceParams] = None,
pad_between_seqs: Optional[bool] = None,
fp8_output: Optional[bool] = False,
bf16_backward: Optional[bool] = False,
num_splits: Optional[int] = 1,
score_mod: Optional[Callable] = None,
score_mod_bprop: Optional[Callable] = None,
Expand Down Expand Up @@ -1817,6 +1819,25 @@ def forward(
qkv_format=qkv_format,
inference_params=inference_params,
)
elif all(
isinstance(x, MXFP8TensorStorage) for x in [query_layer, key_layer, value_layer]
):
# Pre-quantized MXFP8 q/k/v: the wrapper has no real storage, so run
# layout detection on the underlying rowwise data (mirrors the Float8 path).
(
qkv_layout,
query_layer._rowwise_data,
key_layer._rowwise_data,
value_layer._rowwise_data,
q_format,
kv_format,
) = dpa_utils.get_qkv_layout(
query_layer._rowwise_data,
key_layer._rowwise_data,
value_layer._rowwise_data,
qkv_format=qkv_format,
inference_params=inference_params,
)
else:
(
qkv_layout,
Expand Down Expand Up @@ -2190,6 +2211,7 @@ def forward(
fp8_output=fp8_output,
packed_qkv=qkv_layer,
packed_kv=kv_layer,
bf16_backward=bf16_backward,
)
return self.fused_attention(
query_layer,
Expand Down Expand Up @@ -2227,6 +2249,7 @@ def forward(
score_mod_bprop_tensors=score_mod_bprop_tensors,
packed_qkv=qkv_layer,
packed_kv=kv_layer,
bf16_backward=bf16_backward,
)

if use_unfused_attention:
Expand Down
Loading
Loading