From 628c43630155b9fa8a8ec8871eb90c06bf8b36e6 Mon Sep 17 00:00:00 2001 From: Hexiang Wang <56632993+whx-sjtu@users.noreply.github.com> Date: Tue, 5 May 2026 23:55:37 +0800 Subject: [PATCH] [New Model][ROCm] Add AMD support for DeepSeek V4 (#40871) Signed-off-by: ganyi Signed-off-by: whx-sjtu Signed-off-by: tjtanaa Signed-off-by: tjtanaavllm Co-authored-by: ganyi Co-authored-by: tjtanaa Co-authored-by: tjtanaavllm --- CMakeLists.txt | 12 +- ...deepseek_v4_qnorm_rope_kv_insert_kernel.cu | 38 +- csrc/moe/topk_softplus_sqrt_kernels.cu | 53 +- csrc/moe/torch_bindings.cpp | 3 +- csrc/torch_bindings.cpp | 2 - requirements/rocm.txt | 3 + tests/kernels/moe/test_topk_softplus_sqrt.py | 6 +- vllm/config/kernel.py | 2 + .../kernels/linear/scaled_mm/aiter.py | 15 + vllm/model_executor/layers/activation.py | 4 +- .../layers/deepseek_compressor.py | 5 +- .../layers/deepseek_v4_attention.py | 92 ++- .../layers/fused_moe/oracle/mxfp4.py | 81 ++- vllm/model_executor/layers/mhc.py | 107 +++- .../layers/quantization/utils/fp8_utils.py | 9 + .../layers/sparse_attn_indexer.py | 30 +- vllm/model_executor/models/deepseek_v4.py | 7 +- vllm/model_executor/models/deepseek_v4_mtp.py | 8 +- vllm/platforms/rocm.py | 1 + vllm/v1/attention/backends/mla/sparse_swa.py | 3 +- .../fused_inv_rope_fp8_quant.py | 4 +- .../v1/attention/ops/rocm_aiter_mla_sparse.py | 588 ++++++++++++++++-- 22 files changed, 939 insertions(+), 134 deletions(-) diff --git a/CMakeLists.txt b/CMakeLists.txt index bf4ac05e4f2..13788fa8743 100644 --- a/CMakeLists.txt +++ b/CMakeLists.txt @@ -307,12 +307,12 @@ set(VLLM_EXT_SRC "csrc/quantization/activation_kernels.cu" "csrc/cuda_utils_kernels.cu" "csrc/custom_all_reduce.cu" - "csrc/torch_bindings.cpp") + "csrc/torch_bindings.cpp" + "csrc/fused_deepseek_v4_qnorm_rope_kv_insert_kernel.cu") if(VLLM_GPU_LANG STREQUAL "CUDA") list(APPEND VLLM_EXT_SRC - "csrc/minimax_reduce_rms_kernel.cu" - "csrc/fused_deepseek_v4_qnorm_rope_kv_insert_kernel.cu") + "csrc/minimax_reduce_rms_kernel.cu") SET(CUTLASS_ENABLE_HEADERS_ONLY ON CACHE BOOL "Enable only the header library") @@ -1047,13 +1047,13 @@ endif() set(VLLM_MOE_EXT_SRC "csrc/moe/torch_bindings.cpp" "csrc/moe/moe_align_sum_kernels.cu" - "csrc/moe/topk_softmax_kernels.cu") + "csrc/moe/topk_softmax_kernels.cu" + "csrc/moe/topk_softplus_sqrt_kernels.cu") if(VLLM_GPU_LANG STREQUAL "CUDA") list(APPEND VLLM_MOE_EXT_SRC "csrc/moe/moe_wna16.cu" - "csrc/moe/grouped_topk_kernels.cu" - "csrc/moe/topk_softplus_sqrt_kernels.cu") + "csrc/moe/grouped_topk_kernels.cu") endif() if(VLLM_GPU_LANG STREQUAL "CUDA") diff --git a/csrc/fused_deepseek_v4_qnorm_rope_kv_insert_kernel.cu b/csrc/fused_deepseek_v4_qnorm_rope_kv_insert_kernel.cu index e96017d86da..2f2e7ecc182 100644 --- a/csrc/fused_deepseek_v4_qnorm_rope_kv_insert_kernel.cu +++ b/csrc/fused_deepseek_v4_qnorm_rope_kv_insert_kernel.cu @@ -29,7 +29,11 @@ */ #include -#include +#ifndef USE_ROCM + #include +#else + #include +#endif #include #include @@ -42,7 +46,23 @@ #include "type_convert.cuh" #ifndef FINAL_MASK - #define FINAL_MASK 0xffffffffu + #ifdef USE_ROCM + #define FINAL_MASK 0xffffffffffffffffULL + #else + #define FINAL_MASK 0xffffffffu + #endif +#endif + +#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 + __hip_fp8_e4m3_fnuz fp8_val(val); + #endif + return reinterpret_cast(fp8_val); +} #endif namespace vllm { @@ -314,9 +334,13 @@ __global__ void fusedDeepseekV4QNormRopeKVRopeQuantInsertKernel( for (int i = 0; i < kElemsPerLane; i++) { float scaled = elements[i] * inv_scale; scaled = fminf(fmaxf(scaled, -kFp8Max), kFp8Max); +#ifndef USE_ROCM __nv_fp8_storage_t s = __nv_cvt_float_to_fp8(scaled, __NV_SATFINITE, __NV_E4M3); out_bytes[i] = static_cast(s); +#else + out_bytes[i] = rocm_cvt_float_to_fp8_e4m3(scaled); +#endif } // One 16-byte STG per lane. *reinterpret_cast(token_fp8_ptr + dim_base) = @@ -384,6 +408,7 @@ void launchFusedDeepseekV4QNormRopeKVRopeQuantInsert( // PDL: enable programmatic stream serialization whenever the hardware // supports it (SM90+). On pre-Hopper GPUs the attribute is unavailable, // so leave numAttrs = 0 and launch as a regular kernel. +#ifndef USE_ROCM static int const sm_version = getSMVersion(); // Host-side guard: the device kernel body is compiled as a no-op for // bf16 on pre-Ampere (sm_70/sm_75) because _typeConvert is @@ -410,6 +435,15 @@ void launchFusedDeepseekV4QNormRopeKVRopeQuantInsert( q_inout, kv_in, k_cache, slot_mapping, position_ids, cos_sin_cache, eps, num_tokens_full, num_tokens_insert, num_heads_q, cache_block_size, kv_block_stride); +#else + // ROCm: use standard kernel launch syntax (no PDL/stream serialization) + // clang-format off + fusedDeepseekV4QNormRopeKVRopeQuantInsertKernel + <<>>( + q_inout, kv_in, k_cache, slot_mapping, position_ids, cos_sin_cache, + eps, num_tokens_full, num_tokens_insert, num_heads_q, + cache_block_size, kv_block_stride); +#endif } } // namespace deepseek_v4_fused_ops diff --git a/csrc/moe/topk_softplus_sqrt_kernels.cu b/csrc/moe/topk_softplus_sqrt_kernels.cu index 50a8540a737..43d461a0179 100644 --- a/csrc/moe/topk_softplus_sqrt_kernels.cu +++ b/csrc/moe/topk_softplus_sqrt_kernels.cu @@ -60,15 +60,6 @@ __device__ __forceinline__ float toFloat(T value) { } } -#define FINAL_MASK 0xffffffff -template -__inline__ __device__ T warpReduceSum(T val) { -#pragma unroll - for (int mask = 16; mask > 0; mask >>= 1) - val += __shfl_xor_sync(FINAL_MASK, val, mask, 32); - return val; -} - // ====================== TopK softplus_sqrt things // =============================== @@ -272,8 +263,14 @@ __launch_bounds__(WARPS_PER_CTA* WARP_SIZE_PARAM) __global__ } } // Compute per-thread scale (using warp reduction when renormalizing). + // THREADS_PER_ROW-parameterized butterfly works for both warp sizes (32 + // on CUDA, 64 on ROCm CDNA) and any THREADS_PER_ROW the dispatch picks. if (renormalize) { - selected_sum = warpReduceSum(selected_sum); +#pragma unroll + for (int mask = THREADS_PER_ROW / 2; mask > 0; mask /= 2) { + selected_sum += + VLLM_SHFL_XOR_SYNC_WIDTH(selected_sum, mask, THREADS_PER_ROW); + } } float scale = static_cast(routed_scaling_factor); if (renormalize) { @@ -544,7 +541,6 @@ void topkGatingSoftplusSqrtKernelLauncher( const IndType* tid2eid, cudaStream_t stream) { static constexpr int WARPS_PER_TB = 4; static constexpr int BYTES_PER_LDG_POWER_OF_2 = 16; -#ifndef USE_ROCM // for bfloat16 dtype, we need 4 bytes loading to make sure num_experts // elements can be loaded by a warp static constexpr int BYTES_PER_LDG_MULTIPLE_64 = @@ -552,6 +548,19 @@ void topkGatingSoftplusSqrtKernelLauncher( std::is_same_v) ? 4 : 8; + // Narrower LDG (ELTS_PER_LDG=1) used by 192/320/448/576 on ROCm WARP_SIZE=64 + // where ELTS_PER_LDG=2 fails the EXPERTS%(ELTS_PER_LDG*WARP_SIZE)==0 check. + // On CUDA WARP_SIZE=32 the wider LDG already aligns, so the alias collapses + // back to BYTES_PER_LDG_MULTIPLE_64 — no behavioral change for CUDA. +#ifdef USE_ROCM + static constexpr int BYTES_PER_LDG_MULTIPLE_64_NARROW = + (std::is_same_v || + std::is_same_v) + ? 2 + : 4; +#else + static constexpr int BYTES_PER_LDG_MULTIPLE_64_NARROW = + BYTES_PER_LDG_MULTIPLE_64; #endif switch (num_experts) { case 1: @@ -584,27 +593,29 @@ void topkGatingSoftplusSqrtKernelLauncher( case 512: LAUNCH_SOFTPLUS_SQRT(512, WARPS_PER_TB, BYTES_PER_LDG_POWER_OF_2); break; - // (CUDA only) support multiples of 64 when num_experts is not power of 2. - // ROCm uses WARP_SIZE 64 so 8 bytes loading won't fit for some of - // num_experts, alternatively we can test 4 bytes loading and enable it in - // future. -#ifndef USE_ROCM + // Multiples of 64 that are not powers of 2. The kernel requires + // EXPERTS % (ELTS_PER_LDG * WARP_SIZE) == 0. With ELTS_PER_LDG=2 + // (BYTES_PER_LDG_MULTIPLE_64), this holds for all five values on CUDA + // WARP_SIZE=32 but only for 384 on ROCm WARP_SIZE=64. The other four + // use BYTES_PER_LDG_MULTIPLE_64_NARROW (ELTS_PER_LDG=1), which + // satisfies the assertion for any multiple of 64 on either backend; + // on CUDA the narrow alias collapses back to the wider load, so CUDA + // behavior is unchanged. case 192: - LAUNCH_SOFTPLUS_SQRT(192, WARPS_PER_TB, BYTES_PER_LDG_MULTIPLE_64); + LAUNCH_SOFTPLUS_SQRT(192, WARPS_PER_TB, BYTES_PER_LDG_MULTIPLE_64_NARROW); break; case 320: - LAUNCH_SOFTPLUS_SQRT(320, WARPS_PER_TB, BYTES_PER_LDG_MULTIPLE_64); + LAUNCH_SOFTPLUS_SQRT(320, WARPS_PER_TB, BYTES_PER_LDG_MULTIPLE_64_NARROW); break; case 384: LAUNCH_SOFTPLUS_SQRT(384, WARPS_PER_TB, BYTES_PER_LDG_MULTIPLE_64); break; case 448: - LAUNCH_SOFTPLUS_SQRT(448, WARPS_PER_TB, BYTES_PER_LDG_MULTIPLE_64); + LAUNCH_SOFTPLUS_SQRT(448, WARPS_PER_TB, BYTES_PER_LDG_MULTIPLE_64_NARROW); break; case 576: - LAUNCH_SOFTPLUS_SQRT(576, WARPS_PER_TB, BYTES_PER_LDG_MULTIPLE_64); + LAUNCH_SOFTPLUS_SQRT(576, WARPS_PER_TB, BYTES_PER_LDG_MULTIPLE_64_NARROW); break; -#endif default: { TORCH_CHECK(false, "Unsupported expert number: ", num_experts); } diff --git a/csrc/moe/torch_bindings.cpp b/csrc/moe/torch_bindings.cpp index b737cb54353..8940e341cd0 100644 --- a/csrc/moe/torch_bindings.cpp +++ b/csrc/moe/torch_bindings.cpp @@ -16,14 +16,13 @@ TORCH_LIBRARY_EXPAND(TORCH_EXTENSION_NAME, m) { "bias) -> ()"); m.impl("topk_sigmoid", torch::kCUDA, &topk_sigmoid); -#ifndef USE_ROCM m.def( "topk_softplus_sqrt(Tensor! topk_weights, Tensor! topk_indices, Tensor! " "token_expert_indices, Tensor gating_output, bool renormalize, float " "routed_scaling_factor, Tensor? " "bias, Tensor? input_ids, Tensor? tid2eid) -> ()"); m.impl("topk_softplus_sqrt", torch::kCUDA, &topk_softplus_sqrt); -#endif + // Calculate the result of moe by summing up the partial results // from all selected experts. m.def("moe_sum(Tensor input, Tensor! output) -> ()"); diff --git a/csrc/torch_bindings.cpp b/csrc/torch_bindings.cpp index 8d8f7bed044..e695497fd88 100644 --- a/csrc/torch_bindings.cpp +++ b/csrc/torch_bindings.cpp @@ -183,7 +183,6 @@ TORCH_LIBRARY_EXPAND(TORCH_EXTENSION_NAME, ops) { "int forced_token_heads_per_warp=-1) -> ()"); ops.impl("fused_qk_norm_rope", torch::kCUDA, &fused_qk_norm_rope); -#ifndef USE_ROCM // Horizontally-fused DeepseekV4-MLA: per-head RMSNorm + GPT-J RoPE for Q, and // GPT-J RoPE + UE8M0 FP8 quant + paged cache insert for KV, all in one // kernel launch. @@ -194,7 +193,6 @@ TORCH_LIBRARY_EXPAND(TORCH_EXTENSION_NAME, ops) { "float eps, int cache_block_size) -> ()"); ops.impl("fused_deepseek_v4_qnorm_rope_kv_rope_quant_insert", torch::kCUDA, &fused_deepseek_v4_qnorm_rope_kv_rope_quant_insert); -#endif // Apply repetition penalties to logits in-place ops.def( diff --git a/requirements/rocm.txt b/requirements/rocm.txt index 0b472b90c02..037b20874b5 100644 --- a/requirements/rocm.txt +++ b/requirements/rocm.txt @@ -21,3 +21,6 @@ timm>=1.0.17 # amd-quark: required for Quark quantization on ROCm # To be consistent with test_quark.py amd-quark>=0.8.99 +# tilelang has to be installed for mhc module to be +# imported correctly. +tilelang==0.1.9 diff --git a/tests/kernels/moe/test_topk_softplus_sqrt.py b/tests/kernels/moe/test_topk_softplus_sqrt.py index 7f5aacb383d..1b68213fafe 100644 --- a/tests/kernels/moe/test_topk_softplus_sqrt.py +++ b/tests/kernels/moe/test_topk_softplus_sqrt.py @@ -70,7 +70,8 @@ def test_sqrtsoftplus_bias_uses_deepseek_v4_routing_method(): @pytest.mark.skipif( - not current_platform.is_cuda(), reason="This test is skipped on non-CUDA platform." + not current_platform.is_cuda_alike(), + reason="This test is skipped on non-CUDA platform.", ) @pytest.mark.parametrize("num_tokens", [1, 33, 128]) @pytest.mark.parametrize("hidden_size", [1024, 2048]) @@ -125,7 +126,8 @@ def test_fused_topk_softplus_sqrt( @pytest.mark.skipif( - not current_platform.is_cuda(), reason="This test is skipped on non-CUDA platform." + not current_platform.is_cuda_alike(), + reason="This test is skipped on non-CUDA platform.", ) @pytest.mark.parametrize("num_tokens", [1, 33, 128]) @pytest.mark.parametrize("hidden_size", [1024, 2048]) diff --git a/vllm/config/kernel.py b/vllm/config/kernel.py index f7d5f19e238..da1b1f9f1b1 100644 --- a/vllm/config/kernel.py +++ b/vllm/config/kernel.py @@ -119,6 +119,7 @@ MoEBackend = Literal[ "flashinfer_cutedsl", "marlin", "humming", + "triton_unfused", "aiter", "emulation", ] @@ -150,6 +151,7 @@ class KernelConfig: - "flashinfer_cutedsl": Use FlashInfer with CuteDSL kernels (FP4 only) - "marlin": Use Marlin kernels (weight-only quantization) - "humming": Use Humming Mixed Precision kernels + - "triton_unfused": Use Triton unfused MoE kernels - "aiter": Use AMD AITer kernels (ROCm only) - "emulation": use BF16/FP16 GEMM, dequantizing weights and running QDQ on activations. diff --git a/vllm/model_executor/kernels/linear/scaled_mm/aiter.py b/vllm/model_executor/kernels/linear/scaled_mm/aiter.py index 8a8650d2213..5ded5ca798a 100644 --- a/vllm/model_executor/kernels/linear/scaled_mm/aiter.py +++ b/vllm/model_executor/kernels/linear/scaled_mm/aiter.py @@ -312,6 +312,21 @@ class AiterFp8BlockScaledMMKernel(Fp8BlockScaledMMLinearKernel): As: torch.Tensor, Bs: torch.Tensor, ) -> torch.Tensor: + if As.dtype != Bs.dtype: + from vllm.model_executor.layers.quantization.utils.fp8_utils import ( + _upcast_e8m0_to_fp32, + ) + + if As.dtype == torch.float8_e8m0fnu: + As = _upcast_e8m0_to_fp32(As).contiguous() + else: + As = As.to(torch.float32) + + if Bs.dtype == torch.float8_e8m0fnu: + Bs = _upcast_e8m0_to_fp32(Bs).contiguous() + else: + Bs = Bs.to(torch.float32) + out_dtype = self.config.out_dtype if self.use_triton: gemm_a8w8_blockscale_op = rocm_aiter_ops.triton_gemm_a8w8_blockscale diff --git a/vllm/model_executor/layers/activation.py b/vllm/model_executor/layers/activation.py index 59cc95f18c5..df9459012ae 100644 --- a/vllm/model_executor/layers/activation.py +++ b/vllm/model_executor/layers/activation.py @@ -169,7 +169,9 @@ class SiluAndMulWithClamp(CustomOp): def __init__(self, swiglu_limit: float, *, compile_native: bool = True): super().__init__(compile_native=compile_native) self.swiglu_limit = float(swiglu_limit) - if current_platform.is_cuda_alike() or current_platform.is_xpu(): + if current_platform.is_rocm(): + self._forward_method = self.forward_native + elif current_platform.is_cuda_alike() or current_platform.is_xpu(): self.op = torch.ops._C.silu_and_mul_with_clamp elif current_platform.is_cpu(): self._forward_method = self.forward_native diff --git a/vllm/model_executor/layers/deepseek_compressor.py b/vllm/model_executor/layers/deepseek_compressor.py index cae80c35316..48628fec46e 100644 --- a/vllm/model_executor/layers/deepseek_compressor.py +++ b/vllm/model_executor/layers/deepseek_compressor.py @@ -300,6 +300,7 @@ class DeepseekCompressor(nn.Module): state_cache = self.state_cache.kv_cache # kv_state stored in first half, score_state stored in second half state_width = state_cache.shape[-1] // 2 + pdl_kwargs = {} if current_platform.is_rocm() else {"launch_pdl": False} # Store the KV and score (with fused APE addition) in the state. # NOTE: PDL is disabled — both this kernel and _fused_kernel below @@ -324,7 +325,7 @@ class DeepseekCompressor(nn.Module): TRITON_BLOCK_SIZE=triton.next_power_of_2(kv.shape[-1]), STATE_WIDTH=state_width, COMPRESS_RATIO=self.compress_ratio, - launch_pdl=False, + **pdl_kwargs, ) # Fused: compress → RMSNorm → RoPE → FP8 quant → KV cache write. @@ -373,7 +374,7 @@ class DeepseekCompressor(nn.Module): SCALE_DIM=self._scale_dim, KV_BLOCK_STRIDE=kv_cache.stride(0), num_warps=self._num_warps, - launch_pdl=False, + **pdl_kwargs, ) diff --git a/vllm/model_executor/layers/deepseek_v4_attention.py b/vllm/model_executor/layers/deepseek_v4_attention.py index 847c3eee55a..494d6133808 100644 --- a/vllm/model_executor/layers/deepseek_v4_attention.py +++ b/vllm/model_executor/layers/deepseek_v4_attention.py @@ -28,6 +28,11 @@ from vllm.v1.attention.ops.deepseek_v4_ops import ( fused_inv_rope_fp8_quant, fused_q_kv_rmsnorm, ) +from vllm.v1.attention.ops.rocm_aiter_mla_sparse import ( + rocm_forward_decode_fallback, + rocm_inv_rope_einsum, + rocm_sparse_attn_prefill, +) if TYPE_CHECKING: from vllm.v1.attention.backends.mla.sparse_swa import ( @@ -53,6 +58,7 @@ from vllm.model_executor.layers.quantization.input_quant_fp8 import ( from vllm.model_executor.layers.quantization.utils.quant_utils import ( GroupShape, ) +from vllm.platforms import current_platform from vllm.utils.multi_stream_utils import ( execute_in_parallel, maybe_execute_in_parallel, @@ -198,8 +204,6 @@ class DeepseekV4MultiHeadLatentAttentionWrapper(PluggableLayer): # Pick fp8_einsum recipe based on GPU arch: # SM90: FP32 block scales stay [g, r/128, d/128] → sfb_gran_mn=128 # SM100: INT32 packed scales become [g, r, ...] → sfb_gran_mn=1 - from vllm.platforms import current_platform - cap = current_platform.get_device_capability() assert cap is not None, "DeepseekV4 attention requires a CUDA device" self._einsum_recipe = (1, 128, 128) if cap.major <= 9 else (1, 1, 128) @@ -222,6 +226,7 @@ class DeepseekV4MultiHeadLatentAttentionWrapper(PluggableLayer): + 1 # 1B pad ) + # Will be None on ROCm for now. self.aux_stream_list = mla_modules.aux_stream_list # [0]: GEMM start / post-GEMM event0. [1..3]: GEMM done events; # [1] doubles as post-GEMM event1. Reuse is safe: GEMM fully joins @@ -303,6 +308,19 @@ class DeepseekV4MultiHeadLatentAttentionWrapper(PluggableLayer): ) o = o_padded[:, : self.n_local_heads, :] + # Keep ROCm on the BF16 reference wo_a path util kernel ready. + if current_platform.is_rocm(): + z = rocm_inv_rope_einsum( + self.rotary_emb, + o, + positions, + self.rope_head_dim, + self.n_local_groups, + self.o_lora_rank, + self.wo_a, + ) + return self.wo_b(z.flatten(1)) + # O projection: inverse RoPE + FP8 quant + einsum + wo_b o_fp8, o_scale = fused_inv_rope_fp8_quant( o, @@ -336,12 +354,15 @@ class DeepseekV4MultiHeadLatentAttentionWrapper(PluggableLayer): return self.wo_b(z.flatten(1)) def attn_gemm_parallel_execute(self, hidden_states) -> tuple[Any, ...]: - assert self.aux_stream_list is not None - assert len(self.aux_stream_list) >= 3 + aux_streams = self.aux_stream_list + if aux_streams is not None: + assert len(aux_streams) >= 3 + aux_streams = aux_streams[:3] # fused_wqa_wkv (heaviest) on default; the three lighter input GEMMs # on aux streams 0..2 when their owning module exists. ln_events[0] # is the fan-out start event; ln_events[1..3] are per-aux done events. + # On ROCm, aux_streams is None and execute_in_parallel runs serially. aux_fns: list[Callable[[], Any] | None] = [None, None, None] if self.compressor is not None: @@ -385,7 +406,7 @@ class DeepseekV4MultiHeadLatentAttentionWrapper(PluggableLayer): aux_fns, self.ln_events[0], self.ln_events[1:4], - self.aux_stream_list[:3], + aux_streams, enable=hidden_states.shape[0] <= envs.VLLM_MULTI_STREAM_GEMM_TOKEN_THRESHOLD, ) @@ -419,8 +440,9 @@ class DeepseekV4MultiHeadLatentAttentionWrapper(PluggableLayer): # downstream reads q on default). Indexer/compressor go on aux for # overlap with default's GEMM + cache write. if self.indexer is not None: - assert self.aux_stream_list is not None - aux_stream = self.aux_stream_list[0] + aux_stream = ( + self.aux_stream_list[0] if self.aux_stream_list is not None else None + ) indexer = self.indexer # Local ref so the closure keeps a non-None type for mypy. assert self.compressor is not None @@ -448,8 +470,9 @@ class DeepseekV4MultiHeadLatentAttentionWrapper(PluggableLayer): ) elif self.compressor is not None: # wq_b + kv_insert on default, compressor on aux. - assert self.aux_stream_list is not None - aux_stream = self.aux_stream_list[0] + aux_stream = ( + self.aux_stream_list[0] if self.aux_stream_list is not None else None + ) compressor = self.compressor def wq_b_kv_insert() -> torch.Tensor: @@ -668,7 +691,7 @@ class DeepseekV4MLAAttention(nn.Module, AttentionLayerBase): vllm_config.scheduler_config.max_num_batched_tokens ) self.max_model_len = vllm_config.model_config.max_model_len - # DeepseekV4 only supports fp8 kv-cache format for now + # DeepseekV4 only supports fp8 kv-cache format for now. kv_cache_dtype = cache_config.cache_dtype if cache_config is not None else "fp8" assert kv_cache_dtype.startswith("fp8"), ( @@ -816,6 +839,25 @@ class DeepseekV4MLAAttention(nn.Module, AttentionLayerBase): swa_indices = swa_metadata.decode_swa_indices swa_lens = swa_metadata.decode_swa_lens + if current_platform.is_rocm(): + rocm_forward_decode_fallback( + q=q, + kv_cache=kv_cache, + swa_k_cache=self.swa_cache_layer.kv_cache, + swa_only=swa_only, + topk_indices=topk_indices, + topk_lens=topk_lens, + swa_indices=swa_indices, + swa_lens=swa_lens, + attn_sink=self.attn_sink, + scale=self.scale, + head_dim=self.head_dim, + nope_head_dim=self.nope_head_dim, + rope_head_dim=self.rope_head_dim, + output=output, + ) + return + # We treat queries in the same seq as different queries # and later we only attend by generated indices. # q arrives pre-padded to self.padded_heads by the outer wrapper. @@ -980,15 +1022,27 @@ class DeepseekV4MLAAttention(nn.Module, AttentionLayerBase): N, ) - output_chunk, _, _ = flash_mla_sparse_fwd( - q=q[query_start:query_end], - kv=kv.view(-1, 1, q.shape[-1]), - indices=combined_indices.unsqueeze(1), - sm_scale=self.scale, - attn_sink=self.attn_sink, - topk_length=combined_lens, - out=output[query_start:query_end], - ) + if current_platform.is_rocm(): + rocm_sparse_attn_prefill( + q=q[query_start:query_end], + kv=kv.view(-1, 1, q.shape[-1]), + indices=combined_indices.unsqueeze(1), + topk_length=combined_lens, + scale=self.scale, + head_dim=self.head_dim, + attn_sink=self.attn_sink, + output=output[query_start:query_end], + ) + else: + output_chunk, _, _ = flash_mla_sparse_fwd( + q=q[query_start:query_end], + kv=kv.view(-1, 1, q.shape[-1]), + indices=combined_indices.unsqueeze(1), + sm_scale=self.scale, + attn_sink=self.attn_sink, + topk_length=combined_lens, + out=output[query_start:query_end], + ) class DeepseekV4IndexerCache(torch.nn.Module, AttentionLayerBase): diff --git a/vllm/model_executor/layers/fused_moe/oracle/mxfp4.py b/vllm/model_executor/layers/fused_moe/oracle/mxfp4.py index 437da8e6438..7c596d52a65 100644 --- a/vllm/model_executor/layers/fused_moe/oracle/mxfp4.py +++ b/vllm/model_executor/layers/fused_moe/oracle/mxfp4.py @@ -18,6 +18,7 @@ from vllm.model_executor.layers.fused_moe.all2all_utils import ( from vllm.model_executor.layers.fused_moe.config import ( FusedMoEQuantConfig, FusedMoEQuantDesc, + RoutingMethodType, mxfp4_mxfp8_moe_quant_config, mxfp4_w4a8_moe_quant_config, mxfp4_w4a16_moe_quant_config, @@ -64,6 +65,8 @@ class Mxfp4MoeBackend(Enum): MARLIN = "MARLIN" # ROCm AITER backends AITER_MXFP4_BF16 = "AITER_MXFP4_BF16" # W4A16: CK kernel + # Keep the legacy name as an alias while the ROCm split backend rename settles. + AITER = "AITER_MXFP4_BF16" AITER_MXFP4_FP8 = "AITER_MXFP4_FP8" # W4A8: triton kernel # Triton TRITON = "TRITON" @@ -253,6 +256,8 @@ def _get_priority_backends() -> list[Mxfp4MoeBackend]: TRTLLM MXFP8; SM90 falls through to Triton_unfused or Marlin (the backend-level ``is_supported_config`` check filters by device capability). """ + if current_platform.is_rocm(): + return [Mxfp4MoeBackend.AITER_MXFP4_BF16] _AVAILABLE_BACKENDS = [ Mxfp4MoeBackend.FLASHINFER_TRTLLM_MXFP4_MXFP8, Mxfp4MoeBackend.DEEPGEMM_MXFP4, @@ -543,8 +548,22 @@ def select_deepseek_v4_mxfp4_moe_backend( activation_format, ) + # DeepSeek-V4 on ROCm is more accurate with the unfused Triton MXFP4 path + # than the default AITER path. Prefer Triton-unfused for this routing mode, + # while keeping AITER as a fallback if Triton-unfused rejects the config. + if ( + current_platform.is_rocm() + and config.routing_method == RoutingMethodType.DeepseekV4 + ): + priority_backends = [ + Mxfp4MoeBackend.TRITON_UNFUSED, + Mxfp4MoeBackend.AITER_MXFP4_BF16, + ] + else: + priority_backends = _get_priority_backends() + # Iterate priority backends: TRTLLM MXFP8, then Triton. - for backend in _get_priority_backends(): + for backend in priority_backends: activation_key = _backend_activation_key(backend) for k_cls in backend_to_kernel_cls(backend): supported, reason = k_cls.is_supported_config( @@ -1252,6 +1271,64 @@ def convert_weight_to_mxfp4_moe_kernel_format( w2_bias, ) + elif mxfp4_backend == Mxfp4MoeBackend.AITER_MXFP4_BF16: + from vllm._aiter_ops import rocm_aiter_ops + + if w13_bias is not None: + w13_bias = w13_bias.data.to(torch.float32) + if w2_bias is not None: + w2_bias = w2_bias.data.to(torch.float32) + + e, n, k = w13_weight.shape + + w13_weight.view(torch.uint8).copy_( + w13_weight.data.view(torch.uint8) + .view(e, n // 2, 2, k) + .permute(0, 2, 1, 3) + .contiguous() + .view(e, n, k) + ) + w13_weight_scale.data = ( + w13_weight_scale.data.view(e, n // 2, 2, -1) + .permute(0, 2, 1, 3) + .contiguous() + .view(e, n, -1) + ) + + w13_weight.data = w13_weight.data.view(torch.float4_e2m1fn_x2) + w2_weight.data = w2_weight.data.view(torch.float4_e2m1fn_x2) + + w13_weight.data = rocm_aiter_ops.shuffle_weight_a16w4(w13_weight, 16, True) + shuffled_w13_scale = rocm_aiter_ops.shuffle_scale_a16w4( + w13_weight_scale.view(-1, w13_weight_scale.shape[-1]), + num_experts, + True, + ) + + w2_weight.data = rocm_aiter_ops.shuffle_weight_a16w4(w2_weight, 16, False) + shuffled_w2_scale = rocm_aiter_ops.shuffle_scale_a16w4( + w2_weight_scale.view(-1, w2_weight_scale.shape[-1]), + num_experts, + False, + ) + + if w13_bias is not None: + w13_bias = ( + w13_bias.data.view(-1, n // 2, 2) + .permute(0, 2, 1) + .contiguous() + .view(-1, n) + ) + + return ( + w13_weight, + w2_weight, + shuffled_w13_scale, + shuffled_w2_scale, + w13_bias, + w2_bias, + ) + elif mxfp4_backend in TRITON_BACKENDS: from triton_kernels.matmul_ogs import FlexCtx, PrecisionConfig @@ -1307,7 +1384,7 @@ def convert_weight_to_mxfp4_moe_kernel_format( else: raise ValueError( f"Unsupported mxfp4_backend for Mxfp4MoEMethod: {mxfp4_backend}. " - f"Expected TRTLLM or Triton backend." + f"Expected TRTLLM, Triton, or AITER backend." ) diff --git a/vllm/model_executor/layers/mhc.py b/vllm/model_executor/layers/mhc.py index f5c1f06844b..cbc5ec2962e 100644 --- a/vllm/model_executor/layers/mhc.py +++ b/vllm/model_executor/layers/mhc.py @@ -234,6 +234,39 @@ def mhc_pre( num_tokens = residual_flat.shape[0] fn_flat = fn + if current_platform.is_rocm(): + x = residual_flat.view(num_tokens, hc_mult * hidden_size).to(torch.float32) + mixes = torch.matmul(x, fn_flat.t()) + sqrsum = x.square().sum(dim=-1, keepdim=True) + mixes = mixes * torch.rsqrt(sqrsum / (hc_mult * hidden_size) + rms_eps) + + pre_logits = mixes[:, :hc_mult] * hc_scale[0] + hc_base[:hc_mult] + pre_mix = torch.sigmoid(pre_logits) + hc_pre_eps + + post_logits = ( + mixes[:, hc_mult : 2 * hc_mult] * hc_scale[1] + + hc_base[hc_mult : 2 * hc_mult] + ) + post_mix = torch.sigmoid(post_logits) * hc_post_mult_value + + comb_logits = mixes[:, 2 * hc_mult :].view( + num_tokens, hc_mult, hc_mult + ) * hc_scale[2] + hc_base[2 * hc_mult :].view(1, hc_mult, hc_mult) + comb_mix = torch.softmax(comb_logits, dim=-1) + hc_sinkhorn_eps + comb_mix = comb_mix / (comb_mix.sum(dim=-2, keepdim=True) + hc_sinkhorn_eps) + for _ in range(sinkhorn_repeat - 1): + comb_mix = comb_mix / (comb_mix.sum(dim=-1, keepdim=True) + hc_sinkhorn_eps) + comb_mix = comb_mix / (comb_mix.sum(dim=-2, keepdim=True) + hc_sinkhorn_eps) + + layer_input = torch.sum( + pre_mix.unsqueeze(-1) * residual_flat.to(torch.float32), dim=1 + ).to(torch.bfloat16) + return ( + post_mix.view(*outer_shape, hc_mult, 1), + comb_mix.view(*outer_shape, hc_mult, hc_mult), + layer_input.view(*outer_shape, hidden_size), + ) + # these number are from deepgemm kernel impl block_k = 64 block_m = 64 @@ -414,6 +447,14 @@ def mhc_post( post_layer_mix: torch.Tensor, comb_res_mix: torch.Tensor, ) -> torch.Tensor: + if current_platform.is_rocm(): + mixed_residual = torch.einsum( + "...ij,...ih->...jh", + comb_res_mix.to(torch.float32), + residual.to(torch.float32), + ) + post_term = post_layer_mix.to(torch.float32) * x.unsqueeze(-2).to(torch.float32) + return (mixed_residual + post_term).to(residual.dtype) out = torch.empty_like(residual) mhc_post_tilelang( comb_res_mix, @@ -551,6 +592,49 @@ def hc_head_fuse_tilelang( T.pdl_trigger() +def _hc_head_fused_reference( + hs_flat: torch.Tensor, + fn: torch.Tensor, + hc_scale: torch.Tensor, + hc_base: torch.Tensor, + out: torch.Tensor, + hidden_size: int, + rms_eps: float, + hc_eps: float, + hc_mult: int, +) -> None: + """Pure-PyTorch reference for `hc_head_fuse_tilelang`. + + Used on platforms where the tilelang HIP/CUDA backend is not available + (e.g. ROCm builds shipping a tilelang wheel without `target.build.tilelang_hip`). + Mirrors the math of the tilelang kernel exactly: + + x = hs_flat.flatten(-2, -1) # (T, hc_mult * H), fp32 + mixes = x @ fn.T # (T, hc_mult) + rsqrt = 1 / sqrt(||x||^2 / (hc_mult * H) + rms_eps) + pre[m] = sigmoid(mixes[m] * rsqrt * hc_scale[0] + hc_base[m]) + hc_eps + out = sum_m pre[m] * hs_flat[:, m, :] # cast back to bf16 + + `out` is mutated in place to keep the same op contract + (`mutates_args=["out"]`). + """ + num_tokens = hs_flat.shape[0] + if num_tokens == 0: + return + x = hs_flat.reshape(num_tokens, hc_mult * hidden_size).to(torch.float32) + # fn: (hc_mult, hc_mult * hidden_size) → mixes: (T, hc_mult) + mixes = torch.matmul(x, fn.t()) + sqrsum = x.square().sum(dim=-1, keepdim=True) + rsqrt = torch.rsqrt(sqrsum / (hc_mult * hidden_size) + rms_eps) + # hc_scale has shape (1,); hc_base has shape (hc_mult,) + pre_mix = torch.sigmoid(mixes * rsqrt * hc_scale[0] + hc_base) + hc_eps + # weighted sum over the hc_mult channel dim + result = torch.sum(pre_mix.unsqueeze(-1) * hs_flat.to(torch.float32), dim=1).to( + out.dtype + ) + out.copy_(result) + + def _hc_head_fused_kernel( hs_flat: torch.Tensor, fn: torch.Tensor, @@ -563,8 +647,15 @@ def _hc_head_fused_kernel( hc_mult: int, ) -> None: """Fill pre-allocated `out` (T, H) in-place with the hc_head result.""" - if hs_flat.shape[0] > 0: - hc_head_fuse_tilelang( + if hs_flat.shape[0] == 0: + return + if current_platform.is_rocm(): + # tilelang ships only the CUDA codegen in upstream wheels, so the HIP + # FFI target (`target.build.tilelang_hip`) is missing and the JIT call + # would raise `ValueError: Cannot find global function ...`. Use a + # numerically equivalent torch fallback instead. `mhc_pre` and + # `mhc_post` already follow this same pattern above. + _hc_head_fused_reference( hs_flat, fn, hc_scale, @@ -575,6 +666,18 @@ def _hc_head_fused_kernel( hc_eps, hc_mult, ) + return + hc_head_fuse_tilelang( + hs_flat, + fn, + hc_scale, + hc_base, + out, + hidden_size, + rms_eps, + hc_eps, + hc_mult, + ) direct_register_custom_op( diff --git a/vllm/model_executor/layers/quantization/utils/fp8_utils.py b/vllm/model_executor/layers/quantization/utils/fp8_utils.py index 9613b11d35e..d9aab35c25f 100644 --- a/vllm/model_executor/layers/quantization/utils/fp8_utils.py +++ b/vllm/model_executor/layers/quantization/utils/fp8_utils.py @@ -843,6 +843,15 @@ def w8a8_triton_block_scaled_mm( assert len(block_size) == 2 block_n, block_k = block_size[0], block_size[1] + # Triton cannot currently bind E8M0 scale tensors directly. On ROCm, + # DeepSeek-V4 checkpoints store block scales in exponent-only E8M0 format, + # so decode them to fp32 before launching the kernel. + if current_platform.is_rocm(): + if As.dtype == torch.float8_e8m0fnu: + As = _upcast_e8m0_to_fp32(As).contiguous() + if Bs.dtype == torch.float8_e8m0fnu: + Bs = _upcast_e8m0_to_fp32(Bs).contiguous() + assert A.shape[-1] == B.shape[-1] assert A.shape[:-1] == As.shape[:-1] and A.is_contiguous() assert triton.cdiv(A.shape[-1], block_k) == As.shape[-1] diff --git a/vllm/model_executor/layers/sparse_attn_indexer.py b/vllm/model_executor/layers/sparse_attn_indexer.py index ca82f2feb7e..4bf52a49c43 100644 --- a/vllm/model_executor/layers/sparse_attn_indexer.py +++ b/vllm/model_executor/layers/sparse_attn_indexer.py @@ -499,13 +499,31 @@ class SparseAttnIndexer(CustomOp): k: torch.Tensor, weights: torch.Tensor, ): - assert not self.skip_k_cache_insert, ( - "AMD platform doesn't support skip cache insert yet" - ) assert not self.use_fp4_cache, "AMD platform doesn't support fp4 cache yet" assert isinstance(q_quant, torch.Tensor), ( "AMD sparse_attn_indexer expects a single FP8 q_quant tensor" ) + if self.skip_k_cache_insert or not rocm_aiter_ops.is_enabled(): + from vllm.v1.attention.ops.rocm_aiter_mla_sparse import ( + rocm_aiter_sparse_attn_indexer_native, + ) + + return rocm_aiter_sparse_attn_indexer_native( + hidden_states, + _encode_layer_name(self.k_cache.prefix), + self.k_cache.kv_cache, + q_quant, + k, + weights, + self.quant_block_size, + self.scale_fmt, + self.topk_tokens, + self.head_dim, + self.max_model_len, + self.max_total_seq_len, + self.topk_indices_buffer, + skip_k_cache_insert=self.skip_k_cache_insert, + ) if rocm_aiter_ops.is_enabled(): return torch.ops.vllm.rocm_aiter_sparse_attn_indexer( hidden_states, @@ -522,8 +540,4 @@ class SparseAttnIndexer(CustomOp): self.max_total_seq_len, self.topk_indices_buffer, ) - else: - raise RuntimeError( - "Sparse attention indexer ROCm custom op requires ROCm " - "Aiter ops to be enabled." - ) + raise RuntimeError("Sparse attention indexer ROCm path could not be selected.") diff --git a/vllm/model_executor/models/deepseek_v4.py b/vllm/model_executor/models/deepseek_v4.py index 01aa922f3c2..cef4038dc2e 100644 --- a/vllm/model_executor/models/deepseek_v4.py +++ b/vllm/model_executor/models/deepseek_v4.py @@ -1245,7 +1245,12 @@ class DeepseekV4Model(nn.Module): # DeepseekV4MultiHeadLatentAttentionWrapper.attn_gemm_parallel_execute # (compressor kv_score, indexer.weights_proj, indexer.compressor # kv_score). fused_wqa_wkv stays on the default stream. - aux_stream_list = [torch.cuda.Stream() for _ in range(3)] + # Disable them on ROCm because of hang issues. + aux_stream_list = ( + None + if current_platform.is_rocm() + else [torch.cuda.Stream() for _ in range(3)] + ) self.device = current_platform.device_type # Reserved topk indices buffer for all Indexer layers to reuse. diff --git a/vllm/model_executor/models/deepseek_v4_mtp.py b/vllm/model_executor/models/deepseek_v4_mtp.py index a3724e5ebe8..195709c9dac 100644 --- a/vllm/model_executor/models/deepseek_v4_mtp.py +++ b/vllm/model_executor/models/deepseek_v4_mtp.py @@ -167,8 +167,12 @@ class DeepSeekV4MultiTokenPredictor(nn.Module): ) # Three aux streams shared across all MTP layers, mirroring - # DeepseekV4Model. - aux_stream_list = [torch.cuda.Stream() for _ in range(3)] + # DeepseekV4Model. ROCm runs the same work serially for now. + aux_stream_list = ( + None + if current_platform.is_rocm() + else [torch.cuda.Stream() for _ in range(3)] + ) # to map the exact layer index from weights self.layers = torch.nn.ModuleDict( diff --git a/vllm/platforms/rocm.py b/vllm/platforms/rocm.py index 7200a7698d6..984b706e704 100644 --- a/vllm/platforms/rocm.py +++ b/vllm/platforms/rocm.py @@ -409,6 +409,7 @@ class RocmPlatform(Platform): "gptq", "gptq_marlin", # will be overwritten with gptq "fp8", + "deepseek_v4_fp8", "compressed-tensors", "fbgemm_fp8", "gguf", diff --git a/vllm/v1/attention/backends/mla/sparse_swa.py b/vllm/v1/attention/backends/mla/sparse_swa.py index b17fd5d3441..28564e6a97d 100644 --- a/vllm/v1/attention/backends/mla/sparse_swa.py +++ b/vllm/v1/attention/backends/mla/sparse_swa.py @@ -7,6 +7,7 @@ import torch from vllm.config import CacheConfig, VllmConfig, get_current_vllm_config from vllm.model_executor.layers.attention_layer_base import AttentionLayerBase +from vllm.platforms import current_platform from vllm.triton_utils import tl, triton from vllm.v1.attention.backend import ( AttentionBackend, @@ -360,7 +361,7 @@ class DeepseekSparseSWAMetadataBuilder(AttentionMetadataBuilder): _LAYER_TYPE_C4A: None, _LAYER_TYPE_C128A: None, } - if num_decode_tokens == 0: + if num_decode_tokens == 0 or current_platform.is_rocm(): return out for layer_type in self._layer_types: # get_mla_metadata() is the official FlashMLA entry point that diff --git a/vllm/v1/attention/ops/deepseek_v4_ops/fused_inv_rope_fp8_quant.py b/vllm/v1/attention/ops/deepseek_v4_ops/fused_inv_rope_fp8_quant.py index 84647d6120d..68d33f1aa10 100644 --- a/vllm/v1/attention/ops/deepseek_v4_ops/fused_inv_rope_fp8_quant.py +++ b/vllm/v1/attention/ops/deepseek_v4_ops/fused_inv_rope_fp8_quant.py @@ -9,6 +9,7 @@ INT32-packed UE8M0 on SM100) so fp8_einsum skips transform_sf_into_required_layo import torch +from vllm.platforms import current_platform from vllm.triton_utils import tl, triton from vllm.utils.torch_utils import direct_register_custom_op @@ -242,6 +243,7 @@ def _fused_inv_rope_fp8_quant_kernel_impl( (scale_inner * tma_aligned_T, 1, tma_aligned_T), ) grid = (tma_aligned_T, n_groups * heads_per_group) + pdl_kwargs = {} if current_platform.is_rocm() else {"launch_pdl": False} _fused_inv_rope_fp8_quant_per_head[grid]( o, positions, @@ -265,7 +267,7 @@ def _fused_inv_rope_fp8_quant_kernel_impl( HALF_ROPE=half_rope, TMA_ALIGNED_SCALES=tma_aligned_scales, num_stages=1, - launch_pdl=False, + **pdl_kwargs, num_warps=1, ) return fp8_buf, scale_buf diff --git a/vllm/v1/attention/ops/rocm_aiter_mla_sparse.py b/vllm/v1/attention/ops/rocm_aiter_mla_sparse.py index 627d870b62f..5d0343ffd60 100644 --- a/vllm/v1/attention/ops/rocm_aiter_mla_sparse.py +++ b/vllm/v1/attention/ops/rocm_aiter_mla_sparse.py @@ -2,9 +2,11 @@ # SPDX-FileCopyrightText: Copyright contributors to the vLLM project import functools import importlib +import math from importlib.util import find_spec import torch +import torch.nn.functional as F from vllm.forward_context import get_forward_context from vllm.platforms import current_platform @@ -13,6 +15,11 @@ from vllm.utils.torch_utils import LayerNameType from vllm.v1.attention.backends.mla.indexer import DeepseekV32IndexerMetadata from vllm.v1.attention.ops.common import pack_seq_triton, unpack_seq_triton +if current_platform.is_rocm(): + from vllm.platforms.rocm import _ON_GFX942 +else: + _ON_GFX942 = False + @triton.jit def _indexer_k_quant_and_cache_kernel( @@ -230,6 +237,43 @@ def fp8_paged_mqa_logits_torch( fp8_dtype = current_platform.fp8_dtype() batch_size, next_n, _, dim = q.size() + if next_n == 1: + block_size = kv_cache.shape[1] + logits = torch.full( + [batch_size, max_model_len], + float("-inf"), + device=q.device, + dtype=torch.float32, + ) + if context_lens.dim() > 1: + context_lens = context_lens.squeeze(-1) + kv_cache_flat = kv_cache.view(-1, block_size * (dim + 4)) + for i in range(batch_size): + q_i = q[i, 0].to(torch.float32) + q_scale = weights[i] + seq_len = int(context_lens[i].item()) + assert seq_len <= max_model_len + num_pages = cdiv(seq_len, block_size) + padded_seq_len = num_pages * block_size + pages = block_tables[i, :num_pages] + cache = kv_cache_flat[pages] + scale_offset = block_size * dim + cache_value = ( + cache[..., :scale_offset].view(dtype=fp8_dtype).to(torch.float32) + ) + cache_scale = ( + cache[..., scale_offset:].view(dtype=torch.float32).contiguous() + ) + cache_value = cache_value.view(padded_seq_len, dim) + cache_scale = cache_scale.view(padded_seq_len) + score = F.linear(cache_value, q_i) + score = F.relu(score) + score *= q_scale[None, :] + score = score.sum(dim=1) + score *= cache_scale + logits[i, :seq_len] = score[:seq_len] + return logits + kv_cache, scale = kv_cache[..., :dim], kv_cache[..., dim:] scale = scale.contiguous().view(torch.float) q = q.float() @@ -241,20 +285,30 @@ def fp8_paged_mqa_logits_torch( device=q.device, dtype=torch.float32, ) - context_lens = context_lens.tolist() for i in range(batch_size): context_len = context_lens[i] - q_offsets = torch.arange(context_len - next_n, context_len, device="cuda") + if context_len.ndim == 0: + context_len_i = int(context_len.item()) + q_offsets = torch.arange( + context_len_i - next_n, context_len_i, device=q.device + ) + context_limit = torch.full( + (next_n,), context_len_i, dtype=torch.int32, device=q.device + ) + else: + context_limit = context_len.to(device=q.device, dtype=torch.int32) + q_offsets = context_limit - 1 weight_slice = ( weights[i * next_n : (i + 1) * next_n, :].transpose(0, 1).contiguous() ) - for block_rk in range(cdiv(context_len, block_size)): + max_context_len = int(context_limit.max().item()) + for block_rk in range(cdiv(max_context_len, block_size)): block_idx = block_tables[i][block_rk] qx, kx = q[i], kv_cache[block_idx] k_offsets = torch.arange( - block_rk * block_size, (block_rk + 1) * block_size, device="cuda" + block_rk * block_size, (block_rk + 1) * block_size, device=q.device ) - mask = (k_offsets[None, :] < context_len) & ( + mask = (k_offsets[None, :] < context_limit[:, None]) & ( k_offsets[None, :] <= q_offsets[:, None] ) s = torch.where( @@ -331,30 +385,52 @@ def rocm_fp8_paged_mqa_logits( aiter_paged_mqa_logits_module = paged_mqa_logits_module() if aiter_paged_mqa_logits_module is not None: - deepgemm_fp8_paged_mqa_logits = ( - aiter_paged_mqa_logits_module.deepgemm_fp8_paged_mqa_logits + if _ON_GFX942: + deepgemm_fp8_paged_mqa_logits = ( + aiter_paged_mqa_logits_module.deepgemm_fp8_paged_mqa_logits + ) + batch_size, next_n, heads, _ = q_fp8.shape + out_logits = torch.full( + [batch_size * next_n, max_model_len], + float("-inf"), + device="cuda", + dtype=torch.float32, + ) + deepgemm_fp8_paged_mqa_logits( + q_fp8, + kv_cache_fp8, + weights, + out_logits, + context_lens, + block_tables, + max_model_len, + ChunkK=256, + Preshuffle=block_size == 64, + KVBlockSize=block_size, + WavePerEU=2, + ) + return out_logits + deepgemm_fp8_paged_mqa_logits_stage1 = ( + aiter_paged_mqa_logits_module.deepgemm_fp8_paged_mqa_logits_stage1 ) batch_size, next_n, heads, _ = q_fp8.shape - out_logits = torch.full( - [batch_size * next_n, max_model_len], + out_qk = torch.full( + (heads, batch_size * next_n, max_model_len), float("-inf"), device="cuda", dtype=torch.float32, ) - deepgemm_fp8_paged_mqa_logits( + deepgemm_fp8_paged_mqa_logits_stage1( q_fp8, kv_cache_fp8, weights, - out_logits, + out_qk, context_lens, block_tables, max_model_len, - ChunkK=256, - Preshuffle=block_size == 64, - KVBlockSize=block_size, - WavePerEU=2, + ChunkQ=heads, ) - return out_logits + return out_qk.sum(dim=0) else: return fp8_paged_mqa_logits_torch( q_fp8, kv_cache_fp8, weights, context_lens, block_tables, max_model_len @@ -464,6 +540,27 @@ def rocm_fp8_mqa_logits( return fp8_mqa_logits_torch(q, kv, weights, cu_seqlen_ks, cu_seqlen_ke) +def _topk_indices_torch(logits: torch.Tensor, topk_tokens: int) -> torch.Tensor: + k = min(topk_tokens, logits.shape[-1]) + values, indices = torch.topk(logits, k=k, dim=-1) + indices = indices.to(torch.int32) + indices = torch.where( + values == float("-inf"), + torch.full_like(indices, -1, dtype=torch.int32), + indices, + ) + if k == topk_tokens: + return indices + padded = torch.full( + (logits.shape[0], topk_tokens), + -1, + dtype=torch.int32, + device=logits.device, + ) + padded[:, :k] = indices + return padded + + def rocm_aiter_sparse_attn_indexer_fake( hidden_states: torch.Tensor, k_cache_prefix: LayerNameType, @@ -482,8 +579,9 @@ def rocm_aiter_sparse_attn_indexer_fake( # profile run # NOTE(Chen): create the max possible flattened_kv. So that # profile_run can get correct memory usage. + device = hidden_states.device if k is None else k.device _flattened_kv = torch.empty( - [total_seq_lens, head_dim + 4], device=k.device, dtype=torch.uint8 + [total_seq_lens, head_dim + 4], device=device, dtype=torch.uint8 ) fp8_dtype = current_platform.fp8_dtype() _k_fp8 = _flattened_kv[..., :head_dim].view(fp8_dtype).contiguous() @@ -491,7 +589,7 @@ def rocm_aiter_sparse_attn_indexer_fake( return topk_indices_buffer -def rocm_aiter_sparse_attn_indexer( +def rocm_aiter_sparse_attn_indexer_native( hidden_states: torch.Tensor, k_cache_prefix: LayerNameType, kv_cache: torch.Tensor, @@ -505,10 +603,12 @@ def rocm_aiter_sparse_attn_indexer( max_model_len: int, total_seq_lens: int, topk_indices_buffer: torch.Tensor | None, + skip_k_cache_insert: bool = False, ) -> torch.Tensor: # careful! this will be None in dummy run attn_metadata = get_forward_context().attn_metadata fp8_dtype = current_platform.fp8_dtype() + from vllm import _custom_ops as ops from vllm.utils.torch_utils import _resolve_layer_name k_cache_prefix = _resolve_layer_name(k_cache_prefix) @@ -537,19 +637,33 @@ def rocm_aiter_sparse_attn_indexer( has_decode = layer_attn_metadata.num_decodes > 0 has_prefill = layer_attn_metadata.num_prefills > 0 num_decode_tokens = layer_attn_metadata.num_decode_tokens + device = hidden_states.device if k is None else k.device # during speculative decoding, k may be padded to the CUDA graph batch # size while slot_mapping only covers actual tokens. num_tokens = slot_mapping.shape[0] - k = k[:num_tokens] + if k is not None: + k = k[:num_tokens] + elif not skip_k_cache_insert: + raise ValueError("k must be provided when skip_k_cache_insert is False") - indexer_k_quant_and_cache_triton( - k, - kv_cache, - slot_mapping, - quant_block_size, - scale_fmt, - ) + if not skip_k_cache_insert: + if _ON_GFX942: + ops.indexer_k_quant_and_cache( + k, + kv_cache, + slot_mapping, + quant_block_size, + scale_fmt, + ) + else: + indexer_k_quant_and_cache_triton( + k, + kv_cache, + slot_mapping, + quant_block_size, + scale_fmt, + ) topk_indices_buffer[: hidden_states.shape[0]] = -1 if has_prefill: @@ -558,22 +672,31 @@ def rocm_aiter_sparse_attn_indexer( for chunk in prefill_metadata.chunks: k_fp8 = torch.empty( [chunk.total_seq_lens, head_dim], - device=k.device, + device=device, dtype=fp8_dtype, ) k_scale = torch.empty( [chunk.total_seq_lens, 4], - device=k.device, + device=device, dtype=torch.uint8, ) - cp_gather_indexer_k_quant_cache_triton( - kv_cache, - k_fp8, - k_scale, - chunk.block_table, - chunk.cu_seq_lens, - chunk.token_to_seq, - ) + if _ON_GFX942: + ops.cp_gather_indexer_k_quant_cache( + kv_cache, + k_fp8, + k_scale, + chunk.block_table, + chunk.cu_seq_lens, + ) + else: + cp_gather_indexer_k_quant_cache_triton( + kv_cache, + k_fp8, + k_scale, + chunk.block_table, + chunk.cu_seq_lens, + token_to_seq=chunk.token_to_seq, + ) logits = rocm_fp8_mqa_logits( q_fp8[chunk.token_start : chunk.token_end], @@ -582,21 +705,10 @@ def rocm_aiter_sparse_attn_indexer( chunk.cu_seqlen_ks, chunk.cu_seqlen_ke, ) - num_rows = logits.shape[0] - assert topk_tokens == 2048, "top_k_per_row assumes size 2048" topk_indices = topk_indices_buffer[ chunk.token_start : chunk.token_end, :topk_tokens ] - torch.ops._C.top_k_per_row_prefill( - logits, - chunk.cu_seqlen_ks, - chunk.cu_seqlen_ke, - topk_indices, - num_rows, - logits.stride(0), - logits.stride(1), - topk_tokens, - ) + topk_indices.copy_(_topk_indices_torch(logits, topk_tokens)) if has_decode: decode_metadata = layer_attn_metadata.decode @@ -633,19 +745,8 @@ def rocm_aiter_sparse_attn_indexer( max_model_len=max_model_len, ) - num_rows = logits.shape[0] - assert topk_tokens == 2048, "top_k_per_row assumes size 2048" topk_indices = topk_indices_buffer[:num_decode_tokens, :topk_tokens] - torch.ops._C.top_k_per_row_decode( - logits, - next_n, - decode_metadata.seq_lens, - topk_indices, - num_rows, - logits.stride(0), - logits.stride(1), - topk_tokens, - ) + topk_indices.copy_(_topk_indices_torch(logits, topk_tokens)[:num_decode_tokens]) if decode_metadata.requires_padding: # if padded, we need to unpack @@ -659,3 +760,370 @@ def rocm_aiter_sparse_attn_indexer( ) return topk_indices_buffer + + +def rocm_aiter_sparse_attn_indexer( + hidden_states: torch.Tensor, + k_cache_prefix: LayerNameType, + kv_cache: torch.Tensor, + q_fp8: torch.Tensor, + k: torch.Tensor, + weights: torch.Tensor, + quant_block_size: int, + scale_fmt: str | None, + topk_tokens: int, + head_dim: int, + max_model_len: int, + total_seq_lens: int, + topk_indices_buffer: torch.Tensor | None, +) -> torch.Tensor: + return rocm_aiter_sparse_attn_indexer_native( + hidden_states, + k_cache_prefix, + kv_cache, + q_fp8, + k, + weights, + quant_block_size, + scale_fmt, + topk_tokens, + head_dim, + max_model_len, + total_seq_lens, + topk_indices_buffer, + skip_k_cache_insert=False, + ) + + +def _decode_e8m0_scales(scale: torch.Tensor) -> torch.Tensor: + if scale.dtype == torch.float8_e8m0fnu: + from vllm.model_executor.layers.quantization.utils.fp8_utils import ( + _upcast_e8m0_to_fp32, + ) + + return _upcast_e8m0_to_fp32(scale).contiguous() + return scale.to(torch.float32) + + +def _expand_2d_block_scales( + scale: torch.Tensor, + rows: int, + cols: int, +) -> torch.Tensor: + scale = _decode_e8m0_scales(scale) + row_blocks, col_blocks = scale.shape[-2:] + row_block = math.ceil(rows / row_blocks) + col_block = math.ceil(cols / col_blocks) + scale = torch.repeat_interleave(scale, row_block, dim=-2)[..., :rows, :] + scale = torch.repeat_interleave(scale, col_block, dim=-1)[..., :, :cols] + return scale + + +def _apply_gptj_inv_rope_ref( + x: torch.Tensor, + positions: torch.Tensor, + cos_sin_cache: torch.Tensor, + rope_dim: int, +) -> torch.Tensor: + if rope_dim == 0 or x.numel() == 0: + return x + half_rot = rope_dim // 2 + nope_dim = x.shape[-1] - rope_dim + dtype = x.dtype + x = x.to(torch.float32) + cache = cos_sin_cache.index_select(0, positions.to(torch.long)) + cos = cache[:, :half_rot].to(torch.float32) + sin = cache[:, half_rot : 2 * half_rot].to(torch.float32) + view_shape = (positions.shape[0],) + (1,) * (x.dim() - 2) + (half_rot,) + cos = cos.view(view_shape) + sin = sin.view(view_shape) + rope = x[..., nope_dim:] + y_even = rope[..., 0::2] + y_odd = rope[..., 1::2] + rope_out = torch.stack( + (y_even * cos + y_odd * sin, y_odd * cos - y_even * sin), + dim=-1, + ).flatten(-2) + x = x.clone() + x[..., nope_dim:] = rope_out + return x.to(dtype) + + +def _apply_inv_rope_ref( + rotary_emb: torch.nn.Module, + x: torch.Tensor, + positions: torch.Tensor, + rope_dim: int, +) -> torch.Tensor: + if hasattr(rotary_emb, "forward_native"): + try: + query, _ = rotary_emb.forward_native( + positions, + x.clone(), + None, + inverse=True, + ) + return query + except TypeError: + pass + return _apply_gptj_inv_rope_ref(x, positions, rotary_emb.cos_sin_cache, rope_dim) + + +def rocm_inv_rope_einsum( + rotary_emb: torch.nn.Module, + o: torch.Tensor, + positions: torch.Tensor, + rope_head_dim: int, + n_local_groups: int, + o_lora_rank: int, + wo_a: torch.nn.Module, +) -> torch.Tensor: + """Reference inverse-RoPE + WO_A einsum path used on ROCm.""" + o_ref = _apply_inv_rope_ref(rotary_emb, o, positions, rope_head_dim).to( + torch.bfloat16 + ) + o_ref = o_ref.view(o.shape[0], n_local_groups, -1) + + hidden_dim = o_ref.shape[-1] + if hasattr(wo_a, "weight_scale_inv"): + wo_a_weight = wo_a.weight.view(n_local_groups, o_lora_rank, hidden_dim).to( + torch.float32 + ) + wo_a_scale = _expand_2d_block_scales( + wo_a.weight_scale_inv.view( + n_local_groups, -1, wo_a.weight_scale_inv.shape[-1] + ), + o_lora_rank, + hidden_dim, + ) + wo_a_weight = (wo_a_weight * wo_a_scale).to(torch.bfloat16) + else: + wo_a_weight = wo_a.weight.view(n_local_groups, o_lora_rank, hidden_dim).to( + torch.bfloat16 + ) + + return torch.einsum("tgd,grd->tgr", o_ref, wo_a_weight) + + +def rocm_ref_sparse_attn_prefill( + q: torch.Tensor, + kv: torch.Tensor, + indices: torch.Tensor, + topk_length: torch.Tensor | None, + scale: float, + head_dim: int, + attn_sink: torch.Tensor | None, +) -> torch.Tensor: + indices = indices.clone().squeeze(1) + s_q, h_q, d_qk = q.shape + topk = indices.shape[-1] + s_kv = kv.shape[0] + if topk_length is not None: + mask = torch.arange(topk, device=indices.device).unsqueeze( + 0 + ) >= topk_length.unsqueeze(1) + indices[mask] = -1 + invalid_mask = (indices < 0) | (indices >= s_kv) + indices[invalid_mask] = 0 + + qf = q.float() + gathered_kv = kv.index_select(0, indices.flatten()).reshape(s_q, topk, d_qk).float() + scores = qf @ gathered_kv.transpose(1, 2) + scores *= scale + scores[invalid_mask.unsqueeze(1).expand_as(scores)] = float("-inf") + + orig_lse = torch.logsumexp(scores, dim=-1) + lse_for_o = orig_lse + if attn_sink is not None: + lse_for_o = torch.logsumexp( + torch.stack( + [orig_lse, attn_sink[:h_q].view(1, h_q).expand_as(orig_lse)], + dim=0, + ), + dim=0, + ) + lse_for_o = lse_for_o.clone() + lse_for_o[lse_for_o == float("-inf")] = float("+inf") + probs = torch.exp(scores - lse_for_o.unsqueeze(-1)) + out = probs @ gathered_kv[..., :head_dim] + lonely_q_mask = orig_lse == float("-inf") + out[lonely_q_mask.unsqueeze(-1).expand_as(out)] = 0.0 + return out.to(torch.bfloat16) + + +def rocm_sparse_attn_prefill( + q: torch.Tensor, + kv: torch.Tensor, + indices: torch.Tensor, + topk_length: torch.Tensor | None, + scale: float, + head_dim: int, + attn_sink: torch.Tensor | None, + output: torch.Tensor, +) -> None: + output_chunk = rocm_ref_sparse_attn_prefill( + q=q, + kv=kv, + indices=indices, + topk_length=topk_length, + scale=scale, + head_dim=head_dim, + attn_sink=attn_sink, + ) + output.copy_(output_chunk.to(output.dtype)) + + +def rocm_dequantize_blocked_k_cache( + quant_k_cache: torch.Tensor, + head_dim: int, + nope_head_dim: int, + rope_head_dim: int, +) -> torch.Tensor: + fp8_dtype = current_platform.fp8_dtype() + tile_size = 64 + num_tiles = nope_head_dim // tile_size + + num_blocks, block_size, _ = quant_k_cache.shape + quant_k_cache = quant_k_cache.view(num_blocks, -1) + input_nope_rope = quant_k_cache[ + :, : block_size * (nope_head_dim + 2 * rope_head_dim) + ].view(num_blocks, block_size, nope_head_dim + 2 * rope_head_dim) + input_nope = input_nope_rope[:, :, :nope_head_dim].view(fp8_dtype) + input_rope = input_nope_rope[:, :, nope_head_dim:].view(torch.bfloat16) + input_scale = ( + quant_k_cache[:, block_size * (nope_head_dim + 2 * rope_head_dim) :] + .view(num_blocks, block_size, 8)[:, :, :num_tiles] + .view(torch.float8_e8m0fnu) + ) + + result = torch.empty( + (num_blocks, block_size, 1, head_dim), + dtype=torch.bfloat16, + device=quant_k_cache.device, + ) + result[..., nope_head_dim:] = input_rope.unsqueeze(2) + for tile_idx in range(num_tiles): + cur_nope = input_nope[ + ..., tile_idx * tile_size : (tile_idx + 1) * tile_size + ].to(torch.bfloat16) + cur_scales = input_scale[:, :, tile_idx].to(torch.bfloat16).unsqueeze(-1) + result[..., tile_idx * tile_size : (tile_idx + 1) * tile_size] = ( + cur_nope * cur_scales + ).unsqueeze(2) + return result + + +def rocm_ref_sparse_attn_decode( + q: torch.Tensor, + blocked_k: torch.Tensor, + indices_in_kvcache: torch.Tensor, + topk_length: torch.Tensor | None, + scale: float, + head_dim: int, + attn_sink: torch.Tensor | None, + extra_blocked_k: torch.Tensor | None = None, + extra_indices_in_kvcache: torch.Tensor | None = None, + extra_topk_length: torch.Tensor | None = None, +) -> torch.Tensor: + b, s_q, h_q, d_qk = q.shape + + def process_scope( + cur_blocked_k: torch.Tensor, + cur_indices: torch.Tensor, + cur_topk_length: torch.Tensor | None, + ) -> tuple[torch.Tensor, torch.Tensor]: + cur_indices = cur_indices.reshape(b, s_q, -1) + topk = cur_indices.size(-1) + fixed_indices = torch.clamp_min(cur_indices, 0) + gathered_kv = ( + cur_blocked_k.view(-1, d_qk) + .index_select(0, fixed_indices.view(-1)) + .view(b, s_q, topk, d_qk) + ) + invalid_mask = cur_indices == -1 + if cur_topk_length is not None: + cur_topk_length = cur_topk_length.reshape(b) + invalid_mask |= torch.arange(0, topk, device=invalid_mask.device).view( + 1, 1, topk + ) >= cur_topk_length.view(b, 1, 1) + return gathered_kv, invalid_mask + + gathered_kv, invalid_mask = process_scope( + blocked_k, indices_in_kvcache, topk_length + ) + if extra_blocked_k is not None: + assert extra_indices_in_kvcache is not None + gathered_kv1, invalid_mask1 = process_scope( + extra_blocked_k, extra_indices_in_kvcache, extra_topk_length + ) + gathered_kv = torch.cat([gathered_kv, gathered_kv1], dim=2) + invalid_mask = torch.cat([invalid_mask, invalid_mask1], dim=2) + + gathered_kv = gathered_kv.view(b * s_q, -1, d_qk).float() + gathered_kv[gathered_kv != gathered_kv] = 0.0 + qf = q.float().view(b * s_q, h_q, d_qk) + attn_weight = qf @ gathered_kv.transpose(-1, -2) + attn_weight *= scale + attn_weight[ + invalid_mask.view(b * s_q, 1, -1).expand(b * s_q, h_q, invalid_mask.size(-1)) + ] = float("-inf") + lse = attn_weight.logsumexp(dim=-1) + attn_weight = torch.exp(attn_weight - lse.unsqueeze(-1)) + output = attn_weight @ gathered_kv[..., :head_dim] + output = output.view(b, s_q, h_q, head_dim) + lse = lse.view(b, s_q, h_q) + + if attn_sink is not None: + output *= (1.0 / (1.0 + torch.exp(attn_sink.view(1, 1, h_q) - lse))).unsqueeze( + -1 + ) + + lonely_q_mask = lse == float("-inf") + output[lonely_q_mask.unsqueeze(-1).expand_as(output)] = 0.0 + return output.squeeze(1).to(torch.bfloat16) + + +def rocm_forward_decode_fallback( + q: torch.Tensor, + kv_cache: torch.Tensor | None, + swa_k_cache: torch.Tensor, + swa_only: bool, + topk_indices: torch.Tensor | None, + topk_lens: torch.Tensor | None, + swa_indices: torch.Tensor, + swa_lens: torch.Tensor, + attn_sink: torch.Tensor | None, + scale: float, + head_dim: int, + nope_head_dim: int, + rope_head_dim: int, + output: torch.Tensor, +) -> None: + blocked_swa = rocm_dequantize_blocked_k_cache( + swa_k_cache, + head_dim=head_dim, + nope_head_dim=nope_head_dim, + rope_head_dim=rope_head_dim, + ) + blocked_extra = None + if not swa_only: + assert kv_cache is not None + blocked_extra = rocm_dequantize_blocked_k_cache( + kv_cache, + head_dim=head_dim, + nope_head_dim=nope_head_dim, + rope_head_dim=rope_head_dim, + ) + attn_out = rocm_ref_sparse_attn_decode( + q=q.unsqueeze(1), + blocked_k=blocked_swa, + indices_in_kvcache=swa_indices.unsqueeze(1), + topk_length=swa_lens, + scale=scale, + head_dim=head_dim, + attn_sink=attn_sink[: q.shape[1]] if attn_sink is not None else None, + extra_blocked_k=blocked_extra, + extra_indices_in_kvcache=topk_indices, + extra_topk_length=topk_lens, + ) + output.copy_(attn_out.to(output.dtype))