From afdcbd5d39eaf2b37b616c8ee8aabc51e15e70ef Mon Sep 17 00:00:00 2001 From: Tuukka Sarvi Date: Thu, 18 Jun 2026 15:21:14 +0300 Subject: [PATCH] [ROCm][DSv4] Functional fixes for DeepSeek V4 on MI300X/MI325X (#45681) Signed-off-by: ganyi Signed-off-by: Markus Hartikainen Signed-off-by: Tuukka Sarvi Co-authored-by: ganyi Co-authored-by: Cursor Co-authored-by: Markus Hartikainen Co-authored-by: Jin Tao --- ...deepseek_v4_qnorm_rope_kv_insert_kernel.cu | 15 +- ..._fused_deepseek_v4_qnorm_rope_kv_insert.py | 171 ++++++++++-- .../layers/quantization/utils/fp8_utils.py | 25 +- vllm/models/deepseek_v4/amd/rocm.py | 4 + .../deepseek_v4/common/ops/cache_utils.py | 55 +++- vllm/models/deepseek_v4/nvidia/ops/o_proj.py | 4 +- .../v1/attention/ops/rocm_aiter_mla_sparse.py | 61 +++- .../v1/attention/ops/triton_fp8_mqa_logits.py | 262 ++++++++++++++++++ 8 files changed, 545 insertions(+), 52 deletions(-) create mode 100644 vllm/v1/attention/ops/triton_fp8_mqa_logits.py diff --git a/csrc/libtorch_stable/fused_deepseek_v4_qnorm_rope_kv_insert_kernel.cu b/csrc/libtorch_stable/fused_deepseek_v4_qnorm_rope_kv_insert_kernel.cu index 4d34b4b6b50..7bc435b8e0d 100644 --- a/csrc/libtorch_stable/fused_deepseek_v4_qnorm_rope_kv_insert_kernel.cu +++ b/csrc/libtorch_stable/fused_deepseek_v4_qnorm_rope_kv_insert_kernel.cu @@ -18,7 +18,7 @@ * ROPE_DIM = 64 (RoPE applied to dims [NOPE_DIM, HEAD_DIM)) * NOPE_DIM = 448 * QUANT_BLOCK = 64 (UE8M0 FP8 quant block) - * FP8_MAX = 448.0f + * FP8_MAX = 224.0f on ROCm FNUZ / 448.0f on OCP * is_neox=false (GPT-J interleaved pairs) * cos_sin_cache layout [max_pos, rope_dim] = cos || sin (cos first, sin * second along last dim; each half is rope_dim/2 = 32 values) @@ -61,10 +61,11 @@ #ifdef USE_ROCM // ROCm-compatible FP8 conversion helpers __device__ __forceinline__ uint8_t rocm_cvt_float_to_fp8_e4m3(float val) { - #if defined(HIP_FP8_TYPE_OCP) - __hip_fp8_e4m3 fp8_val(val); - #else + // gfx942 uses FNUZ FP8; other ROCm targets use OCP E4M3. + #if defined(__gfx942__) __hip_fp8_e4m3_fnuz fp8_val(val); + #else + __hip_fp8_e4m3 fp8_val(val); #endif return reinterpret_cast(fp8_val); } @@ -90,7 +91,13 @@ constexpr int kQuantBlock = 64; constexpr int kNumQuantBlocks = kNopeDim / kQuantBlock; // 7 constexpr int kScaleBytesPerToken = kNumQuantBlocks + 1; // 8 (7 real + 1 pad) constexpr int kTokenDataBytes = kNopeDim + kRopeDim * 2; // 448 + 128 = 576 +// FNUZ on gfx942 / OCP elsewhere. FNUZ uses 224.0 (not the dtype's raw +// 240.0) to match the rest of vLLM's FNUZ pipeline. +#if defined(USE_ROCM) && defined(__gfx942__) +constexpr float kFp8Max = 224.0f; +#else constexpr float kFp8Max = 448.0f; +#endif #ifndef USE_ROCM // When num_tokens is less than this threshold, diff --git a/tests/kernels/test_fused_deepseek_v4_qnorm_rope_kv_insert.py b/tests/kernels/test_fused_deepseek_v4_qnorm_rope_kv_insert.py index e568ce57638..d2919185519 100644 --- a/tests/kernels/test_fused_deepseek_v4_qnorm_rope_kv_insert.py +++ b/tests/kernels/test_fused_deepseek_v4_qnorm_rope_kv_insert.py @@ -19,17 +19,28 @@ The kernel is imported via import pytest import torch +from vllm.model_executor.layers.quantization.utils.quant_utils import ( + get_fp8_min_max, +) from vllm.models.deepseek_v4.common.ops import ( dequantize_and_gather_k_cache, quantize_and_insert_k_cache, ) +from vllm.platforms import current_platform # ── Constants matching the kernel ──────────────────────────────────────────── HEAD_DIM = 512 ROPE_DIM = 64 NOPE_DIM = HEAD_DIM - ROPE_DIM # 448 QUANT_BLOCK = 64 -FP8_MAX = 448.0 +# Match the C++ SWA-K encoder: FNUZ on gfx942, OCP elsewhere. +USE_FNUZ = current_platform.is_fp8_fnuz() +_, FP8_MAX = get_fp8_min_max() +# The kernel emits FNUZ-encoded fp8 bytes on gfx942 (rocm_cvt_float_to_fp8_e4m3) +# but stores them into float8_e4m3fn-typed tensors, matching vLLM's ROCm cache +# convention. References must encode under the same scheme and the kernel's +# e4m3fn-typed outputs must be reinterpreted under it before decoding. +FP8_STORE_DTYPE = torch.float8_e4m3fnuz if USE_FNUZ else torch.float8_e4m3fn HEAD_BYTES = NOPE_DIM + ROPE_DIM * 2 + 8 # 448 + 128 + 8 = 584 @@ -81,10 +92,11 @@ def apply_rope_gptj_last_k( cos = cos.unsqueeze(1) sin = sin.unsqueeze(1) - # Use addcmul (compiles to FMA on CUDA) for the 2x2 rotation. nvcc lowers - # the kernel's `e*c - o*s` to fma(e, c, -o*s); matching that here keeps - # near-cancellation pairs on the same bf16 grid as the kernel output and - # avoids spurious 1-ULP boundary flips at high num_tokens. + # Use addcmul (an FMA) for the 2x2 rotation to mirror the kernel's + # `e*c - o*s` fused form. This keeps the reference close to the kernel, but + # the fp32 reference and the fp32 GPU kernel can still round to bf16 on + # opposite sides of a round-to-nearest tie for a tiny number of elements at + # high positions, so callers compare the RoPE region within 1 bf16 ULP. new_even = torch.addcmul(-odd * sin, even, cos) new_odd = torch.addcmul(odd * cos, even, sin) rope_rotated = torch.stack((new_even, new_odd), dim=-1).reshape(shape) @@ -148,6 +160,86 @@ def _call_fused( ) +def _bf16_ulp_distance(a: torch.Tensor, b: torch.Tensor) -> torch.Tensor: + """Representable-step distance between two bf16 tensors. + + Reinterprets the bf16 bit patterns under the IEEE-754 total ordering so + that adjacent representable values differ by exactly 1. + """ + + def key(t: torch.Tensor) -> torch.Tensor: + u = t.contiguous().view(torch.int16).to(torch.int64) & 0xFFFF + return torch.where(u >= 0x8000, 0xFFFF - u, u + 0x8000) + + return (key(a) - key(b)).abs() + + +def _fp8_ulp_distance(a: torch.Tensor, b: torch.Tensor) -> torch.Tensor: + """Representable-step distance between two 8-bit fp8 tensors. + + Reinterprets the fp8 bytes under a sign-magnitude total ordering so that + adjacent representable values differ by exactly 1. Inputs must already share + the same fp8 encoding (e.g. both FP8_STORE_DTYPE). + """ + + def key(t: torch.Tensor) -> torch.Tensor: + u = t.contiguous().view(torch.uint8).to(torch.int64) + return torch.where(u >= 0x80, 0xFF - u, u + 0x80) + + return (key(a) - key(b)).abs() + + +def _as_stored_fp8(t: torch.Tensor) -> torch.Tensor: + """Reinterpret a float8_e4m3fn-typed kernel output under the real (FNUZ on + gfx942) encoding the kernel actually wrote, without touching the bytes.""" + return t.contiguous().view(torch.uint8).view(FP8_STORE_DTYPE) + + +def _dequant_cache(k_cache_2d, num_tokens, num_blocks, block_size): + """Round-trip a [num_blocks, block_size*HEAD_BYTES] K-cache back to bf16.""" + device = k_cache_2d.device + out = torch.zeros(1, num_tokens, HEAD_DIM, dtype=torch.bfloat16, device=device) + seq_lens = torch.tensor([num_tokens], dtype=torch.int32, device=device) + block_table = torch.arange(num_blocks, dtype=torch.int32, device=device).unsqueeze( + 0 + ) + k_cache_3d = k_cache_2d.view(num_blocks, block_size, HEAD_BYTES) + dequantize_and_gather_k_cache( + out, + k_cache_3d, + seq_lens, + None, + block_table, + block_size, + offset=0, + use_fnuz=USE_FNUZ, + ) + return out[0, :num_tokens] + + +def _assert_kv_cache_parity( + k_cache_fused, k_cache_ref, num_tokens, num_blocks, block_size +): + """Assert the fused and reference K-caches agree after decoding. + + The NoPE region is deterministic UE8M0 FP8, so its round-trip must be + bit-identical. The RoPE region is stored as bf16 after an fp32 rotation: + the GPU kernel and the PyTorch reference can fall on opposite sides of a + round-to-nearest tie and differ by at most one bf16 ULP. (Spot checks show + the kernel value is the correctly-rounded one; the fp32 torch reference is + the one that lands on the wrong side near a midpoint.) Allow <=1 ULP there. + """ + rec_fused = _dequant_cache(k_cache_fused, num_tokens, num_blocks, block_size) + rec_ref = _dequant_cache(k_cache_ref, num_tokens, num_blocks, block_size) + torch.testing.assert_close( + rec_fused[:, :NOPE_DIM], rec_ref[:, :NOPE_DIM], rtol=0, atol=0 + ) + max_ulp = int( + _bf16_ulp_distance(rec_fused[:, NOPE_DIM:], rec_ref[:, NOPE_DIM:]).max().item() + ) + assert max_ulp <= 1, f"RoPE bf16 region differs by {max_ulp} ULP (>1)" + + # ── Test 1: Q path numerical parity ────────────────────────────────────────── @@ -241,7 +333,7 @@ def test_kv_path_matches_reference(num_tokens: int, block_size: int): num_blocks, block_size * HEAD_BYTES, dtype=torch.uint8, device=device ) quantize_and_insert_k_cache( - kv_ref, k_cache_ref, slot_mapping, block_size=block_size + kv_ref, k_cache_ref, slot_mapping, block_size=block_size, use_fnuz=USE_FNUZ ) # ── Fused path (dummy q, padded to FlashMLA's min head count 64) ─────── @@ -273,7 +365,14 @@ def test_kv_path_matches_reference(num_tokens: int, block_size: int): # gather_lens arg is None (use seq_lens) k_cache_3d = k_cache_2d.view(num_blocks, block_size, HEAD_BYTES) dequantize_and_gather_k_cache( - out, k_cache_3d, seq_lens, None, block_table, block_size, offset=0 + out, + k_cache_3d, + seq_lens, + None, + block_table, + block_size, + offset=0, + use_fnuz=USE_FNUZ, ) return out[0, :num_tokens] @@ -297,12 +396,10 @@ def test_kv_path_matches_reference(num_tokens: int, block_size: int): f"fused NoPE token {t} diff {diff_fused} > {max_allowed}" ) - # RoPE region: bf16 stored exactly → zero diff. - rope_diff = (recovered_fused[:, NOPE_DIM:] - kv_ref[:, NOPE_DIM:]).abs().max() - assert rope_diff.item() == 0.0, f"RoPE portion not exact: {rope_diff.item()}" - - # Exact byte equality of the two cache buffers — strong parity. - torch.testing.assert_close(k_cache_fused, k_cache_ref, rtol=0, atol=0) + # Strong parity: NoPE FP8 round-trip bit-identical, RoPE bf16 within 1 ULP. + _assert_kv_cache_parity( + k_cache_fused, k_cache_ref, num_tokens, num_blocks, block_size + ) # ── Test 2b: DP padding (slot_mapping shorter than q/kv) ───────────────────── @@ -336,7 +433,7 @@ def test_kv_path_with_dp_padding(num_tokens: int, pad: int, block_size: int): num_blocks, block_size * HEAD_BYTES, dtype=torch.uint8, device=device ) quantize_and_insert_k_cache( - kv_ref, k_cache_ref, slot_mapping, block_size=block_size + kv_ref, k_cache_ref, slot_mapping, block_size=block_size, use_fnuz=USE_FNUZ ) # Fused: pass full-sized q/kv/positions, shorter slot_mapping. @@ -354,7 +451,9 @@ def test_kv_path_with_dp_padding(num_tokens: int, pad: int, block_size: int): block_size, ) - torch.testing.assert_close(k_cache_fused, k_cache_ref, rtol=0, atol=0) + _assert_kv_cache_parity( + k_cache_fused, k_cache_ref, num_tokens, num_blocks, block_size + ) # ── Test 3: combined single-call Q + KV parity ─────────────────────────────── @@ -403,7 +502,7 @@ def test_combined_q_and_kv( num_blocks, block_size * HEAD_BYTES, dtype=torch.uint8, device=device ) quantize_and_insert_k_cache( - kv_ref, k_cache_ref, slot_mapping, block_size=block_size + kv_ref, k_cache_ref, slot_mapping, block_size=block_size, use_fnuz=USE_FNUZ ) # Fused single call. @@ -426,7 +525,9 @@ def test_combined_q_and_kv( assert pad_region.abs().max().item() == 0.0, ( "padded head slots must be exact zero" ) - torch.testing.assert_close(k_cache_fused, k_cache_ref, rtol=0, atol=0) + _assert_kv_cache_parity( + k_cache_fused, k_cache_ref, num_tokens, num_blocks, block_size + ) # ── Full-cache (FlashInfer) path parity ────────────────────────────────────── @@ -499,7 +600,7 @@ def _fp8_full_cache_reference( q_ref = apply_rope_gptj_last_k(q_ref, positions, cos_sin_cache) q_fp8.copy_( torch.clamp(q_ref.float() * q_fp8_scale_inv, -FP8_MAX, FP8_MAX).to( - torch.float8_e4m3fn + FP8_STORE_DTYPE ) ) @@ -510,7 +611,7 @@ def _fp8_full_cache_reference( pos_in_block = slots % block_size k_cache[block_idx, pos_in_block] = torch.clamp( kv_ref[valid].float() / fp8_scale, -FP8_MAX, FP8_MAX - ).to(torch.float8_e4m3fn) + ).to(FP8_STORE_DTYPE) def _bf16_full_cache_reference( @@ -565,12 +666,17 @@ def test_full_cache_per_tensor_fp8_matches_reference( fp8_scale = torch.tensor([1.0], dtype=torch.float32, device=device) q_fp8_scale_inv = torch.tensor([1.0], dtype=torch.float32, device=device) - q_fp8_ref = torch.empty_like(q, dtype=torch.float8_e4m3fn) + # References are encoded under the scheme the kernel actually writes + # (FNUZ on gfx942); the kernel's own outputs must stay float8_e4m3fn-typed + # because the op asserts that dtype. + q_fp8_ref = torch.empty_like(q, dtype=FP8_STORE_DTYPE) q_fp8_fused = torch.empty_like(q, dtype=torch.float8_e4m3fn) k_cache_ref = torch.zeros( + num_blocks, block_size, HEAD_DIM, dtype=FP8_STORE_DTYPE, device=device + ) + k_cache_fused = torch.zeros( num_blocks, block_size, HEAD_DIM, dtype=torch.float8_e4m3fn, device=device ) - k_cache_fused = torch.zeros_like(k_cache_ref) _fp8_full_cache_reference( q, @@ -599,12 +705,29 @@ def test_full_cache_per_tensor_fp8_matches_reference( block_size, ) + # Q is RMSNorm(no-weight)+RoPE in fp32 before fp8 quant; the RMSNorm + # reduction and RoPE rotation can land the kernel and the torch reference on + # opposite sides of an fp8 round-to-nearest tie, so allow <=1 fp8 ULP. + q_fused = _as_stored_fp8(q_fp8_fused) + q_max_ulp = int(_fp8_ulp_distance(q_fused, q_fp8_ref).max().item()) + assert q_max_ulp <= 1, f"Q fp8 differs by {q_max_ulp} ULP (>1)" + + # K-cache NoPE region [0, NOPE_DIM) is a deterministic per-tensor fp8 quant + # of the (un-rotated) KV input, so it must be bit-identical. The RoPE region + # [NOPE_DIM, HEAD_DIM) is rotated in fp32 and may differ by <=1 fp8 ULP. + k_fused = _as_stored_fp8(k_cache_fused) torch.testing.assert_close( - q_fp8_fused.float(), q_fp8_ref.float(), rtol=0, atol=0.25 + k_fused[..., :NOPE_DIM].float(), + k_cache_ref[..., :NOPE_DIM].float(), + rtol=0, + atol=0, ) - torch.testing.assert_close( - k_cache_fused.float(), k_cache_ref.float(), rtol=0, atol=0.25 + k_max_ulp = int( + _fp8_ulp_distance(k_fused[..., NOPE_DIM:], k_cache_ref[..., NOPE_DIM:]) + .max() + .item() ) + assert k_max_ulp <= 1, f"K-cache RoPE fp8 differs by {k_max_ulp} ULP (>1)" @pytest.mark.skipif( diff --git a/vllm/model_executor/layers/quantization/utils/fp8_utils.py b/vllm/model_executor/layers/quantization/utils/fp8_utils.py index 66a9aa86bde..be1167332ed 100644 --- a/vllm/model_executor/layers/quantization/utils/fp8_utils.py +++ b/vllm/model_executor/layers/quantization/utils/fp8_utils.py @@ -1363,9 +1363,28 @@ def process_fp8_weight_block_strategy( ) if current_platform.is_fp8_fnuz() and weight.dtype == torch.float8_e4m3fn: - weight, weight_scale, _ = normalize_e4m3fn_to_e4m3fnuz( - weight=weight, weight_scale=weight_scale - ) + if weight_scale.dtype == torch.float8_e8m0fnu: + # UE8M0 scales: e8m0 stores exponent-only values (2^(exp-127)), + # so doubling the dequant scale == incrementing the exponent byte + # by 1. Convert the OCP E4M3 weight bytes to FNUZ in place by + # reinterpreting and patching the NaN sentinel (-128 in int8), + # then double the UE8M0 exponent so the dequantized magnitudes + # match. + weight_as_int8 = weight.view(torch.int8) + ROCM_FP8_NAN_AS_INT = -128 + weight_as_int8[weight_as_int8 == ROCM_FP8_NAN_AS_INT] = 0 + weight = weight_as_int8.view(torch.float8_e4m3fnuz) + exp_bytes = weight_scale.view(torch.uint8) + weight_scale = ( + (exp_bytes.to(torch.int16) + 1) + .clamp(max=254) + .to(torch.uint8) + .view(torch.float8_e8m0fnu) + ) + else: + weight, weight_scale, _ = normalize_e4m3fn_to_e4m3fnuz( + weight=weight, weight_scale=weight_scale + ) weight = _maybe_pad_fp8_weight(weight) return weight, weight_scale diff --git a/vllm/models/deepseek_v4/amd/rocm.py b/vllm/models/deepseek_v4/amd/rocm.py index 7b300c60ced..641b3da68bd 100644 --- a/vllm/models/deepseek_v4/amd/rocm.py +++ b/vllm/models/deepseek_v4/amd/rocm.py @@ -14,6 +14,7 @@ from vllm.models.deepseek_v4.sparse_mla import ( DeepseekV4FlashMLAMetadata, DeepseekV4FlashMLAMetadataBuilder, ) +from vllm.platforms import current_platform from vllm.triton_utils import tl, triton from vllm.v1.attention.backend import ( CommonAttentionMetadata, @@ -796,6 +797,7 @@ class DeepseekV4ROCMAiterMLAAttention(DeepseekV4Attention): assert attn_metadata is not None assert compressed_k_cache is not None block_table = attn_metadata.block_table[num_decodes:] + # compressed_k_cache is OCP on every platform (Triton encoder). dequantize_and_gather_k_cache( kv[:chunk_size], compressed_k_cache, @@ -804,6 +806,7 @@ class DeepseekV4ROCMAiterMLAAttention(DeepseekV4Attention): block_table=block_table[chunk_start:chunk_end], block_size=attn_metadata.block_size // self.compress_ratio, offset=0, + use_fnuz=False, ) swa_block_table = swa_metadata.block_table[num_decodes:] @@ -815,6 +818,7 @@ class DeepseekV4ROCMAiterMLAAttention(DeepseekV4Attention): block_table=swa_block_table[chunk_start:chunk_end], block_size=swa_metadata.block_size, offset=N, + use_fnuz=current_platform.is_fp8_fnuz(), ) query_start = ( diff --git a/vllm/models/deepseek_v4/common/ops/cache_utils.py b/vllm/models/deepseek_v4/common/ops/cache_utils.py index 8adf219dbbe..ffaec528aa8 100644 --- a/vllm/models/deepseek_v4/common/ops/cache_utils.py +++ b/vllm/models/deepseek_v4/common/ops/cache_utils.py @@ -16,6 +16,10 @@ preparation. import torch +from vllm.model_executor.layers.quantization.utils.quant_utils import ( + get_fp8_min_max, +) +from vllm.platforms import current_platform from vllm.triton_utils import tl, triton from vllm.utils.import_utils import has_cutedsl @@ -39,6 +43,7 @@ def quantize_and_insert_k_kernel( block_stride: tl.constexpr, # total bytes per block (padded) fp8_max: tl.constexpr, n_quant_blocks: tl.constexpr, # 8 (7 real + 1 padding) + use_fnuz: tl.constexpr = False, ): """ Quantize K tensor and insert into paged K cache. @@ -49,6 +54,9 @@ def quantize_and_insert_k_kernel( - [64*576 + 64*8, block_stride): Padding One program per token. + + ``use_fnuz=True`` selects FNUZ (``tl.float8e4b8``); default OCP + (``tl.float8e4nv``) matches every production caller. """ pid = tl.program_id(0) @@ -112,8 +120,11 @@ def quantize_and_insert_k_kernel( x_scaled = x / scale x_clamped = tl.clamp(x_scaled, -fp8_max, fp8_max) - # Convert to fp8, then bitcast to uint8 for storage - x_fp8 = x_clamped.to(tl.float8e4nv) + # Convert to fp8 (FNUZ on gfx942, OCP elsewhere), then bitcast to uint8. + if use_fnuz: + x_fp8 = x_clamped.to(tl.float8e4b8) + else: + x_fp8 = x_clamped.to(tl.float8e4nv) x_uint8 = x_fp8.to(tl.uint8, bitcast=True) # Store as uint8 (1 byte each) @@ -145,6 +156,7 @@ def quantize_and_insert_k_cache( slot_mapping: torch.Tensor, # [num_tokens] int64 block_size: int = 64, is_ue8m0: bool = True, + use_fnuz: bool = False, ): """ Quantize K tensor and insert into paged K cache. @@ -155,6 +167,10 @@ def quantize_and_insert_k_cache( - Next 64 * 8 = 512 bytes: Scales - Each token: 8 bytes (uint8 scales, 7 real + 1 padding) - Padded to multiple of 576 + + ``use_fnuz=True`` selects FNUZ E4M3 cache encoding and is only valid on + platforms whose FP8 format is FNUZ. ``use_fnuz=False`` selects OCP E4M3, + which is used by OCP-encoded caches even on gfx942. """ assert k.dim() == 2 and k.shape[1] == 512, ( f"K must be [num_tokens, 512], got {k.shape}" @@ -171,7 +187,12 @@ def quantize_and_insert_k_cache( TOKEN_BF16_DIM = 64 TOKEN_SCALE_DIM = 8 QUANT_BLOCK_SIZE = 64 - FP8_MAX = 448.0 + if use_fnuz: + if not current_platform.is_fp8_fnuz(): + raise ValueError("use_fnuz=True requires a platform using FNUZ FP8") + _, FP8_MAX = get_fp8_min_max() + else: + FP8_MAX = torch.finfo(torch.float8_e4m3fn).max TOKEN_DATA_SIZE = TOKEN_FP8_DIM + TOKEN_BF16_DIM * 2 grid = (num_tokens,) @@ -191,6 +212,7 @@ def quantize_and_insert_k_cache( block_stride=block_stride, fp8_max=FP8_MAX, n_quant_blocks=8, + use_fnuz=use_fnuz, ) @@ -216,6 +238,7 @@ def _dequantize_and_gather_k_kernel( output_dim: tl.constexpr, # 512 fp8_max: tl.constexpr, n_quant_blocks: tl.constexpr, # 7 real blocks + use_fnuz: tl.constexpr = False, ): batch_idx = tl.program_id(0) worker_id = tl.program_id(1) @@ -273,8 +296,11 @@ def _dequantize_and_gather_k_kernel( # Load quantized fp8 values (stored as uint8) x_uint8 = tl.load(token_fp8_ptr + offsets, mask=mask, other=0) - # Bitcast uint8 back to fp8 - x_fp8 = x_uint8.to(tl.float8e4nv, bitcast=True) + # Bitcast uint8 back to fp8 (FNUZ on gfx942, OCP elsewhere). + if use_fnuz: + x_fp8 = x_uint8.to(tl.float8e4b8, bitcast=True) + else: + x_fp8 = x_uint8.to(tl.float8e4nv, bitcast=True) # Convert fp8 to float32 for computation x_float = x_fp8.to(tl.float32) @@ -317,6 +343,7 @@ def dequantize_and_gather_k_cache_triton( block_table: torch.Tensor, block_size: int, offset: int, + use_fnuz: bool = False, ) -> None: TOKEN_FP8_DIM = 448 TOKEN_BF16_DIM = 64 @@ -347,6 +374,7 @@ def dequantize_and_gather_k_cache_triton( output_dim=512, fp8_max=FP8_MAX, n_quant_blocks=7, + use_fnuz=use_fnuz, ) @@ -363,7 +391,15 @@ def dequantize_and_gather_k_cache( block_table: torch.Tensor, block_size: int, offset: int, + use_fnuz: bool = False, ) -> None: + """Dequantize and gather a paged DSv4 K cache. + + ``use_fnuz`` MUST match the encoder of the specific cache being read: + ``False`` for ``compressed_k_cache`` (Triton encoder is OCP everywhere), + ``current_platform.is_fp8_fnuz()`` for ``swa_k_cache`` (C++ encoder + writes FNUZ on gfx942 and OCP on gfx950). + """ if has_cutedsl(): # lazily import, otherwise some tests fail due to CUDA driver init failure. from vllm.models.deepseek_v4.nvidia.ops.dequant_gather_k_cutedsl import ( @@ -376,7 +412,14 @@ def dequantize_and_gather_k_cache( return dequantize_and_gather_k_cache_triton( - out, k_cache, seq_lens, gather_lens, block_table, block_size, offset + out, + k_cache, + seq_lens, + gather_lens, + block_table, + block_size, + offset, + use_fnuz=use_fnuz, ) diff --git a/vllm/models/deepseek_v4/nvidia/ops/o_proj.py b/vllm/models/deepseek_v4/nvidia/ops/o_proj.py index a0b4e2c678e..18e3b10562b 100644 --- a/vllm/models/deepseek_v4/nvidia/ops/o_proj.py +++ b/vllm/models/deepseek_v4/nvidia/ops/o_proj.py @@ -3,7 +3,9 @@ import torch import torch.nn as nn -from vllm.models.deepseek_v4.common.ops import fused_inv_rope_fp8_quant +from vllm.models.deepseek_v4.common.ops.fused_inv_rope_fp8_quant import ( + fused_inv_rope_fp8_quant, +) from vllm.platforms import current_platform from vllm.utils.deep_gemm import fp8_einsum diff --git a/vllm/v1/attention/ops/rocm_aiter_mla_sparse.py b/vllm/v1/attention/ops/rocm_aiter_mla_sparse.py index 51513a5a9f4..dbd4d8d1d4c 100644 --- a/vllm/v1/attention/ops/rocm_aiter_mla_sparse.py +++ b/vllm/v1/attention/ops/rocm_aiter_mla_sparse.py @@ -504,7 +504,13 @@ def fp8_mqa_logits_torch( ) mask = mask_lo & mask_hi - score = torch.einsum("mhd,nd->hmn", q, k).float() * scale + # ``score`` is [H, M, N]; ``scale`` is the per-KV-token scale, which + # vLLM callers hand us as ``[N, 1]`` (a ``[N, 4]`` uint8 buffer cast + # to fp32). PyTorch right-aligns dimensions for broadcasting, so a + # naked ``score * scale`` would align ``scale``'s leading dim with + # ``score``'s M dim and raise a shape mismatch. Flatten to ``[N]`` so + # broadcasting lines up with the last dim of ``score``. + score = torch.einsum("mhd,nd->hmn", q, k).float() * scale.reshape(-1) logits = (score.relu() * weights.unsqueeze(-1).transpose(0, 1)).sum(dim=0) logits = logits.masked_fill(~mask, float("-inf")) @@ -557,13 +563,26 @@ def rocm_fp8_mqa_logits( # path after aiter merge this kernel into main from vllm._aiter_ops import rocm_aiter_ops + k_fp8, scale = kv + + # Temporarily route gfx942 to the vendored ROCm/aiter#3257 workaround. + # Remove this branch once vLLM bumps AITER to a version that includes + # ROCm/aiter#3257. + if _ON_GFX942 and rocm_aiter_ops.is_enabled(): + from vllm.v1.attention.ops.triton_fp8_mqa_logits import ( + fp8_mqa_logits_gfx942, + ) + + return fp8_mqa_logits_gfx942( + q, k_fp8, scale, weights, cu_seqlen_ks, cu_seqlen_ke + ) + aiter_mqa_logits_module = None if rocm_aiter_ops.is_enabled(): aiter_mqa_logits_module = mqa_logits_module() if aiter_mqa_logits_module is not None: fp8_mqa_logits = aiter_mqa_logits_module.fp8_mqa_logits - k_fp8, scale = kv return fp8_mqa_logits(q, k_fp8, scale, weights, cu_seqlen_ks, cu_seqlen_ke) else: return fp8_mqa_logits_torch(q, kv, weights, cu_seqlen_ks, cu_seqlen_ke) @@ -1249,7 +1268,10 @@ def _sparse_attn_decode_ragged_kernel( NOPE_DIM: tl.constexpr, NOPE_BLOCK: tl.constexpr, ROPE_DIM: tl.constexpr, - IS_FNUZ: tl.constexpr, + # SWA K-cache (main): C++ encoder writes FNUZ on gfx942, OCP on gfx950. + # Compressed K-cache (extra): Triton encoder writes OCP everywhere. + IS_FNUZ_MAIN: tl.constexpr, + IS_FNUZ_EXTRA: tl.constexpr, BLOCK_H: tl.constexpr, BLOCK_K: tl.constexpr, ): @@ -1306,8 +1328,8 @@ def _sparse_attn_decode_ragged_kernel( mask=valid[:, None] & nope_mask[None, :], other=0, ) - if IS_FNUZ: - x_fp8 = x_uint8.to(tl.float8e4b15, bitcast=True) + if IS_FNUZ_MAIN: + x_fp8 = x_uint8.to(tl.float8e4b8, bitcast=True) else: x_fp8 = x_uint8.to(tl.float8e4nv, bitcast=True) encoded_scales = tl.load( @@ -1374,8 +1396,8 @@ def _sparse_attn_decode_ragged_kernel( mask=valid[:, None] & nope_mask[None, :], other=0, ) - if IS_FNUZ: - x_fp8 = x_uint8.to(tl.float8e4b15, bitcast=True) + if IS_FNUZ_EXTRA: + x_fp8 = x_uint8.to(tl.float8e4b8, bitcast=True) else: x_fp8 = x_uint8.to(tl.float8e4nv, bitcast=True) encoded_scales = tl.load( @@ -1485,7 +1507,12 @@ def _sparse_attn_decode_partial_kernel( NOPE_DIM: tl.constexpr, NOPE_BLOCK: tl.constexpr, ROPE_DIM: tl.constexpr, - IS_FNUZ: tl.constexpr, + # `main_cache` is the SWA K-cache (written by the C++ encoder, FNUZ on + # gfx942 / OCP on gfx950). `extra_cache` is the compressed K-cache + # (Triton encoder, OCP on every platform). Reading both with the same + # `IS_FNUZ` would decode one of them with the wrong FNUZ/OCP scale ratio. + IS_FNUZ_MAIN: tl.constexpr, + IS_FNUZ_EXTRA: tl.constexpr, BLOCK_H: tl.constexpr, BLOCK_K: tl.constexpr, NUM_SPLITS: tl.constexpr, @@ -1551,8 +1578,8 @@ def _sparse_attn_decode_partial_kernel( mask=valid[:, None] & nope_mask[None, :], other=0, ) - if IS_FNUZ: - x_fp8 = x_uint8.to(tl.float8e4b15, bitcast=True) + if IS_FNUZ_MAIN: + x_fp8 = x_uint8.to(tl.float8e4b8, bitcast=True) else: x_fp8 = x_uint8.to(tl.float8e4nv, bitcast=True) encoded_scales = tl.load( @@ -1622,8 +1649,8 @@ def _sparse_attn_decode_partial_kernel( mask=valid[:, None] & nope_mask[None, :], other=0, ) - if IS_FNUZ: - x_fp8 = x_uint8.to(tl.float8e4b15, bitcast=True) + if IS_FNUZ_EXTRA: + x_fp8 = x_uint8.to(tl.float8e4b8, bitcast=True) else: x_fp8 = x_uint8.to(tl.float8e4nv, bitcast=True) encoded_scales = tl.load( @@ -2095,7 +2122,8 @@ def _rocm_sparse_attn_decode_ragged_triton( NOPE_DIM=nope_head_dim, NOPE_BLOCK=nope_block, ROPE_DIM=rope_head_dim, - IS_FNUZ=is_fnuz, + IS_FNUZ_MAIN=is_fnuz, + IS_FNUZ_EXTRA=False, BLOCK_H=block_h, BLOCK_K=block_k, num_warps=8, @@ -2153,7 +2181,12 @@ def _rocm_sparse_attn_decode_ragged_triton( NOPE_DIM=nope_head_dim, NOPE_BLOCK=nope_block, ROPE_DIM=rope_head_dim, - IS_FNUZ=is_fnuz, + # main_cache = swa_k_cache (C++ encoder, FNUZ on gfx942 / OCP on gfx950). + # extra_cache = compressed kv_cache (Triton encoder, OCP everywhere). + # Reading both with a single IS_FNUZ would decode one of them with the + # wrong FNUZ/OCP scale ratio (~1.87×). + IS_FNUZ_MAIN=is_fnuz, + IS_FNUZ_EXTRA=False, BLOCK_H=block_h, BLOCK_K=block_k, NUM_SPLITS=num_splits, diff --git a/vllm/v1/attention/ops/triton_fp8_mqa_logits.py b/vllm/v1/attention/ops/triton_fp8_mqa_logits.py new file mode 100644 index 00000000000..619d0ec50a9 --- /dev/null +++ b/vllm/v1/attention/ops/triton_fp8_mqa_logits.py @@ -0,0 +1,262 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""Temporary gfx942 fallback for AITER's fp8_mqa_logits kernel. + +This module vendors AITER's Triton fp8_mqa_logits kernel with the gfx942 +tile-size workaround from ROCm/aiter#3257. It is used only while vLLM's +pinned AITER version lacks that fix. + +TODO: Remove this vendored copy once vLLM pins an AITER version that includes +ROCm/aiter#3257 bugfix for gfx942. +""" + +import torch + +from vllm.triton_utils import tl, triton + +# gfx942 (MI300X) has 64 KiB of LDS per CU. We accept the default +# (BLOCK_KV=128, num_stages=2) tile only when *both* of these hold: +# +# 1. Occupancy gate. With waves_per_eu=2 and num_warps=4 we target two +# workgroups co-resident on a CU -> per-WG LDS budget = 32 KiB. Triton +# keeps Q in registers (loop-invariant) and the fp32 scores accumulator +# in VGPRs (heavy VALU), so only the double-buffered KV tile is +# expected to live in LDS. A 0.9 safety factor leaves headroom for any +# LDS overhead the compiler may add. +# +# 2. Hardware ceiling. Defensive upper bound that also counts Q and +# scores against the 64 KiB CU limit, in case a Triton version (older +# or future) decides to spill them to LDS. False positives here only +# shrink the tile; false negatives are JIT-aborts, so we lean +# conservative. +_GFX942_CU_LDS_BYTES = 64 * 1024 +_GFX942_PER_WG_LDS_BUDGET_BYTES = _GFX942_CU_LDS_BYTES * 9 // 20 # ~28.8 KiB + + +def _gfx942_default_tile_fits_lds(num_heads: int, head_size: int) -> bool: + """Return True iff (BLOCK_KV=128, num_stages=2) fits in MI300X LDS.""" + BLOCK_KV = 128 + NUM_STAGES = 2 + kv_bytes = head_size * BLOCK_KV * NUM_STAGES + scores_bytes = num_heads * BLOCK_KV * 4 + q_bytes = num_heads * head_size + fits_occupancy = kv_bytes < _GFX942_PER_WG_LDS_BUDGET_BYTES + fits_hardware = q_bytes + kv_bytes + scores_bytes <= _GFX942_CU_LDS_BYTES + return fits_occupancy and fits_hardware + + +@triton.jit +def _fp8_mqa_logits_kernel( + Q_ptr, # fp8e4m3 [seq_len, H, D] + KV_ptr, # fp8e4m3 [seq_len_kv, D] + kv_scales_ptr, # fp32 [seq_len_kv] + weights_ptr, # fp32 [seq_len, H] + cu_start_ptr, # int32 [seq_len] + cu_end_ptr, # int32 [seq_len] + logits_ptr, # fp32 [seq_len, seq_len_kv] + seq_len, + seq_len_kv, + NUM_HEADS: tl.constexpr, + HEAD_SIZE: tl.constexpr, + # strides + stride_q_s: tl.int64, + stride_q_h: tl.constexpr, + stride_q_d: tl.constexpr, + stride_kv_s: tl.int64, + stride_kv_d: tl.constexpr, + stride_w_s: tl.int64, + stride_w_h: tl.constexpr, + stride_logits_s: tl.int64, + stride_logits_k: tl.int64, + # block sizes + BLOCK_KV: tl.constexpr, +): + row_id = tl.program_id(0) + # go from larger to smaller in terms of work + # to reduce the tail effect + row_id = tl.num_programs(0) - row_id - 1 + tl.assume(row_id >= 0) + tl.assume(stride_q_s > 0) + tl.assume(stride_q_h > 0) + tl.assume(stride_q_d > 0) + tl.assume(stride_kv_s > 0) + tl.assume(stride_kv_d > 0) + tl.assume(stride_w_s > 0) + tl.assume(stride_w_h > 0) + + logits_row_ptrs = logits_ptr + row_id * stride_logits_s + + h_inds = tl.arange(0, NUM_HEADS)[:, None] + d_inds = tl.arange(0, HEAD_SIZE) + + # load Q[BLOCK_Q, NUM_HEADS, HEAD_SIZE] + q_ptrs = ( + Q_ptr + row_id * stride_q_s + h_inds * stride_q_h + d_inds[None, :] * stride_q_d + ) + + q_block = tl.load(q_ptrs, cache_modifier=".cg") + w_ptrs = weights_ptr + row_id * stride_w_s + h_inds * stride_w_h + w_block = tl.load(w_ptrs, cache_modifier=".cg").to(tl.float32) + + # Load start/end for each row in this block + start_ind = tl.load(cu_start_ptr + row_id) + end_ind = tl.load(cu_end_ptr + row_id) + + start_ind = tl.maximum(start_ind, 0) + end_ind = tl.minimum(end_ind, seq_len_kv) + shifted_end = end_ind - start_ind + shifted_unmasked_end = shifted_end // BLOCK_KV * BLOCK_KV + + kv_col_offsets = tl.arange(0, BLOCK_KV) + start_ind + kv_ptrs = ( + KV_ptr + kv_col_offsets[None, :] * stride_kv_s + d_inds[:, None] * stride_kv_d + ) + + kv_scales_ptrs = kv_scales_ptr + kv_col_offsets + + logits_ptrs = logits_row_ptrs + kv_col_offsets * stride_logits_k + + # Loop over KV tiles + for _ in tl.range(0, shifted_unmasked_end, BLOCK_KV): + kv_block = tl.load(kv_ptrs) + kv_scales = tl.load(kv_scales_ptrs) + + # [NUM_HEADS, BLOCK_KV] = [NUM_HEADS, HEAD_SIZE] x [HEAD_SIZE, BLOCK_KV] + scores = tl.dot(q_block, kv_block, input_precision="ieee") + # Multiply by kv_scales (broadcast along rows) + scores = scores * kv_scales[None, :] + # ReLU + scores = tl.maximum(scores, 0.0) + scores = scores * w_block + # [NUM_HEADS, BLOCK_KV] -> [BLOCK_KV, ] + scores = tl.sum(scores, axis=0) + tl.store(logits_ptrs, scores) + + kv_ptrs += BLOCK_KV * stride_kv_s + kv_scales_ptrs += BLOCK_KV + logits_ptrs += BLOCK_KV * stride_logits_k + kv_col_offsets += BLOCK_KV + + # masked load + kv_col_mask = kv_col_offsets < end_ind + kv_block = tl.load(kv_ptrs, mask=kv_col_mask[None, :], other=0.0) + kv_scales = tl.load(kv_scales_ptrs, mask=kv_col_mask, other=0.0) + + # [NUM_HEADS, BLOCK_KV] = [NUM_HEADS, HEAD_SIZE] x [HEAD_SIZE, BLOCK_KV] + scores = tl.dot(q_block, kv_block, input_precision="ieee") + # Multiply by kv_scales (broadcast along rows) + scores = scores * kv_scales[None, :] + # ReLU + scores = tl.maximum(scores, 0.0) + scores = scores * w_block + # [NUM_HEADS, BLOCK_KV] -> [BLOCK_KV, ] + scores = tl.sum(scores, axis=0) + # masked store + in_window = (kv_col_offsets >= start_ind) & (kv_col_offsets < end_ind) + tl.store(logits_ptrs, scores, mask=in_window) + + +def fp8_mqa_logits_gfx942( + q: torch.Tensor, + k_fp8: torch.Tensor, + kv_scales: torch.Tensor, + weights: torch.Tensor, + cu_starts: torch.Tensor, + cu_ends: torch.Tensor, +) -> torch.Tensor: + """Compute FP8 MQA logits on MI300X (gfx942) using the vendored kernel. + + Drop-in replacement for ``aiter.ops.triton.attention.fp8_mqa_logits. + fp8_mqa_logits`` on MI300X. Selects ``(BLOCK_KV, num_stages)`` based on + whether the default tile fits within the 64 KiB LDS budget of a gfx942 + CU (see module docstring). + + Args: + q: Query tensor of shape ``[M, H, D]``, FP8 dtype. + k_fp8: Key tensor of shape ``[N, D]``, FP8 dtype. + kv_scales: K scales of shape ``[N]`` (or ``[N, 1]`` -- viewed as + ``[N]``), float32. + weights: Per-head weights of shape ``[M, H]``, float32. + cu_starts: Start indices (inclusive) of shape ``[M]``, int32. + cu_ends: End indices (exclusive) of shape ``[M]``, int32. + + Returns: + Logits of shape ``[M, N]``, float32 -- positions outside + ``[cu_starts[i], cu_ends[i])`` for row ``i`` are pre-filled with + ``-inf`` so the caller can run a top-k without masking. + """ + seq_len, num_heads, head_size = q.shape + seq_len_kv = k_fp8.shape[0] + assert num_heads & (num_heads - 1) == 0, ( + f"num_heads must be a power of two (got {num_heads})" + ) + assert head_size & (head_size - 1) == 0, ( + f"head_size must be a power of two (got {head_size})" + ) + + # The kernel walks ``kv_scales`` as a 1-D contiguous array of size N + # (it indexes by ``kv_scales_ptr + kv_col_offsets``). The vLLM caller + # passes a ``[N, 4]`` uint8 view-cast-to-float32 which lands as + # ``[N, 1]`` contiguous -- byte-identical to ``[N]`` -- but flatten + # explicitly to keep the kernel's pointer arithmetic intent clear. + kv_scales_1d = kv_scales.reshape(-1) + + # Initialise with -inf so positions outside [cu_starts, cu_ends) read + # as ``-inf`` after the masked store path -- this matches AITER's + # ``fp8_mqa_logits`` semantics and is what the top-k consumer expects. + logits = torch.full( + (seq_len, seq_len_kv), + fill_value=-float("inf"), + dtype=torch.float32, + device=q.device, + ) + + if _gfx942_default_tile_fits_lds(num_heads, head_size): + block_kv = 128 + num_stages = 2 + else: + # DSv4 sparse indexer (NUM_HEADS=64, HEAD_SIZE=128) lands here: + # default tile spills past gfx942's 64 KiB LDS budget. (64, 1) + # needs ~33 KiB and clears the per-WG budget with margin. + block_kv = 64 + num_stages = 1 + + # heuristic for MFMA instruction shape, identical to AITER's choice + matrix_instr_nonkdim = 32 + if seq_len <= 1024: + matrix_instr_nonkdim = 16 + + stride_q_s, stride_q_h, stride_q_d = q.stride() + stride_kv_s, stride_kv_d = k_fp8.stride() + stride_w_s, stride_w_h = weights.stride() + stride_logits_s, stride_logits_k = logits.stride() + + _fp8_mqa_logits_kernel[(seq_len,)]( + Q_ptr=q, + KV_ptr=k_fp8, + kv_scales_ptr=kv_scales_1d, + weights_ptr=weights, + cu_start_ptr=cu_starts, + cu_end_ptr=cu_ends, + logits_ptr=logits, + seq_len=seq_len, + seq_len_kv=seq_len_kv, + NUM_HEADS=num_heads, + HEAD_SIZE=head_size, + stride_q_s=stride_q_s, + stride_q_h=stride_q_h, + stride_q_d=stride_q_d, + stride_kv_s=stride_kv_s, + stride_kv_d=stride_kv_d, + stride_w_s=stride_w_s, + stride_w_h=stride_w_h, + stride_logits_s=stride_logits_s, + stride_logits_k=stride_logits_k, + BLOCK_KV=block_kv, + num_warps=4, + num_stages=num_stages, + waves_per_eu=2, + matrix_instr_nonkdim=matrix_instr_nonkdim, + ) + + return logits