From 599e75f432e5fd7c77e65dc95587f3441201bdbc Mon Sep 17 00:00:00 2001 From: TJian Date: Mon, 18 May 2026 00:18:50 +0800 Subject: [PATCH] [ROCm] [Bugfix] Fix DeepSeek V4 Functionality and Accuracy (#42810) Signed-off-by: tjtanaa --- vllm/model_executor/layers/mhc.py | 88 ++++++----- .../layers/sparse_attn_indexer.py | 27 +--- vllm/model_executor/models/deepseek_v4.py | 3 +- .../v1/attention/ops/rocm_aiter_mla_sparse.py | 147 ++++-------------- 4 files changed, 88 insertions(+), 177 deletions(-) diff --git a/vllm/model_executor/layers/mhc.py b/vllm/model_executor/layers/mhc.py index 36325d35202..a9f84ffaf04 100644 --- a/vllm/model_executor/layers/mhc.py +++ b/vllm/model_executor/layers/mhc.py @@ -61,31 +61,35 @@ class MHCPreOp(CustomOp): sinkhorn_repeat: int, n_splits: int = 1, ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: - hidden_size = residual.shape[-1] - if hidden_size % 256 == 0: - return torch.ops.vllm.mhc_pre_aiter( - residual, - fn, - hc_scale, - hc_base, - rms_eps, - hc_pre_eps, - hc_sinkhorn_eps, - hc_post_mult_value, - sinkhorn_repeat, - ) - else: - return mhc_kernels.mhc_pre_torch( - residual, - fn, - hc_scale, - hc_base, - rms_eps, - hc_pre_eps, - hc_sinkhorn_eps, - hc_post_mult_value, - sinkhorn_repeat, - ) + # TODO: Reenable aiter after we are at the aiter + # version that has this bugfix + # https://github.com/ROCm/aiter/commit/b639cb63bcac4672dce33a731fad042a65cb3649 + # It has accuracy problem at large number of tokens. + # hidden_size = residual.shape[-1] + # if hidden_size % 256 == 0: + # return torch.ops.vllm.mhc_pre_aiter( + # residual, + # fn, + # hc_scale, + # hc_base, + # rms_eps, + # hc_pre_eps, + # hc_sinkhorn_eps, + # hc_post_mult_value, + # sinkhorn_repeat, + # ) + # else: + return mhc_kernels.mhc_pre_torch( + residual, + fn, + hc_scale, + hc_base, + rms_eps, + hc_pre_eps, + hc_sinkhorn_eps, + hc_post_mult_value, + sinkhorn_repeat, + ) def forward_native(self, *args, **kwargs): raise NotImplementedError("Native implementation of mhc_pre is not available") @@ -124,21 +128,25 @@ class MHCPostOp(CustomOp): post_layer_mix: torch.Tensor, comb_res_mix: torch.Tensor, ) -> torch.Tensor: - hidden_size = residual.shape[-1] - if hidden_size % 256 == 0: - return torch.ops.vllm.mhc_post_aiter( - x, - residual, - post_layer_mix, - comb_res_mix, - ) - else: - return mhc_kernels.mhc_post_torch( - x, - residual, - post_layer_mix, - comb_res_mix, - ) + # TODO: Reenable aiter after we are at the aiter + # version that has this bugfix + # https://github.com/ROCm/aiter/commit/b639cb63bcac4672dce33a731fad042a65cb3649 + # It has accuracy problem at large number of tokens. + # hidden_size = residual.shape[-1] + # if hidden_size % 256 == 0: + # return torch.ops.vllm.mhc_post_aiter( + # x, + # residual, + # post_layer_mix, + # comb_res_mix, + # ) + # else: + return mhc_kernels.mhc_post_torch( + x, + residual, + post_layer_mix, + comb_res_mix, + ) def forward_native(self, *args, **kwargs): raise NotImplementedError("Native implementation of mhc_post is not available") diff --git a/vllm/model_executor/layers/sparse_attn_indexer.py b/vllm/model_executor/layers/sparse_attn_indexer.py index 9f202435a73..381e877fcb9 100644 --- a/vllm/model_executor/layers/sparse_attn_indexer.py +++ b/vllm/model_executor/layers/sparse_attn_indexer.py @@ -505,27 +505,6 @@ class SparseAttnIndexer(CustomOp): 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, @@ -541,5 +520,9 @@ class SparseAttnIndexer(CustomOp): self.max_model_len, self.max_total_seq_len, self.topk_indices_buffer, + skip_k_cache_insert=self.skip_k_cache_insert, ) - raise RuntimeError("Sparse attention indexer ROCm path could not be selected.") + raise RuntimeError( + "Sparse attention indexer ROCm path is only supported on AITER. " + "Please enable aiter with VLLM_ROCM_USE_AITER=1" + ) diff --git a/vllm/model_executor/models/deepseek_v4.py b/vllm/model_executor/models/deepseek_v4.py index 3f3a2fb1702..65606796c9e 100644 --- a/vllm/model_executor/models/deepseek_v4.py +++ b/vllm/model_executor/models/deepseek_v4.py @@ -1277,7 +1277,8 @@ class DeepseekV4DecoderLayer(nn.Module): x, post, comb = self.hc_pre( x, self.hc_ffn_fn, self.hc_ffn_scale, self.hc_ffn_base ) - x = self.ffn_norm(x) + # ffn_norm is now folded into self.ffn.norm_gate; ffn() takes + # the pre-norm activation directly. x = self.ffn(x, input_ids) x = self.hc_post(x, residual, post, comb) return x, None, None, None diff --git a/vllm/v1/attention/ops/rocm_aiter_mla_sparse.py b/vllm/v1/attention/ops/rocm_aiter_mla_sparse.py index c3f1b29465e..8bda5bb7a86 100644 --- a/vllm/v1/attention/ops/rocm_aiter_mla_sparse.py +++ b/vllm/v1/attention/ops/rocm_aiter_mla_sparse.py @@ -542,7 +542,11 @@ 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: +def _topk_indices_torch( + logits: torch.Tensor, + topk_tokens: int, + row_starts: torch.Tensor | None = None, +) -> torch.Tensor: k = min(topk_tokens, logits.shape[-1]) values, indices = torch.topk(logits, k=k, dim=-1) indices = indices.to(torch.int32) @@ -551,6 +555,12 @@ def _topk_indices_torch(logits: torch.Tensor, topk_tokens: int) -> torch.Tensor: torch.full_like(indices, -1, dtype=torch.int32), indices, ) + if row_starts is not None: + # Match the CUDA top_k_per_row_prefill contract: indices are local to + # each row's valid [row_start, row_end) range, not columns in the + # concatenated chunk logits matrix. + starts = row_starts.to(dtype=torch.int32).view(-1, 1) + indices = torch.where(indices < 0, indices, indices - starts) if k == topk_tokens: return indices padded = torch.full( @@ -563,64 +573,6 @@ def _topk_indices_torch(logits: torch.Tensor, topk_tokens: int) -> torch.Tensor: return padded -# topk_tokens values with dedicated fused C++ kernel support. -_TOPK_FAST_PATH_VALUES = frozenset({2048}) - - -def _topk_indices_prefill( - logits: torch.Tensor, - topk_tokens: int, - topk_out: torch.Tensor, - cu_seqlen_ks: torch.Tensor, - cu_seqlen_ke: torch.Tensor, -) -> None: - """Top-k indices for the prefill path. - - Writes ``logits.shape[0]`` rows into ``topk_out``; caller must size the - view accordingly. - """ - if topk_tokens in _TOPK_FAST_PATH_VALUES: - torch.ops._C.top_k_per_row_prefill( - logits, - cu_seqlen_ks, - cu_seqlen_ke, - topk_out, - logits.shape[0], - logits.stride(0), - logits.stride(1), - topk_tokens, - ) - else: - topk_out.copy_(_topk_indices_torch(logits, topk_tokens)) - - -def _topk_indices_decode( - logits: torch.Tensor, - topk_tokens: int, - topk_out: torch.Tensor, - seq_lens: torch.Tensor, - next_n: int, -) -> None: - """Top-k indices for the decode path. - - Writes ``logits.shape[0] == batch_size * next_n`` rows into ``topk_out``; - caller must size the view to ``num_padded_tokens``. - """ - if topk_tokens in _TOPK_FAST_PATH_VALUES: - torch.ops._C.top_k_per_row_decode( - logits, - next_n, - seq_lens, - topk_out, - logits.shape[0], - logits.stride(0), - logits.stride(1), - topk_tokens, - ) - else: - topk_out.copy_(_topk_indices_torch(logits, topk_tokens)) - - def rocm_aiter_sparse_attn_indexer_fake( hidden_states: torch.Tensor, k_cache_prefix: LayerNameType, @@ -635,21 +587,13 @@ def rocm_aiter_sparse_attn_indexer_fake( max_model_len: int, total_seq_lens: int, topk_indices_buffer: torch.Tensor | None, + skip_k_cache_insert: bool = False, ) -> torch.Tensor: - # 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=device, dtype=torch.uint8 - ) - fp8_dtype = current_platform.fp8_dtype() - _k_fp8 = _flattened_kv[..., :head_dim].view(fp8_dtype).contiguous() - _k_scale = _flattened_kv[..., head_dim:].view(torch.float32).contiguous() return topk_indices_buffer -def rocm_aiter_sparse_attn_indexer_native( +@eager_break_during_capture +def rocm_aiter_sparse_attn_indexer( hidden_states: torch.Tensor, k_cache_prefix: LayerNameType, kv_cache: torch.Tensor, @@ -688,6 +632,7 @@ def rocm_aiter_sparse_attn_indexer_native( max_model_len, total_seq_lens, topk_indices_buffer, + skip_k_cache_insert, ) layer_attn_metadata = attn_metadata[k_cache_prefix] assert isinstance(layer_attn_metadata, DeepseekV32IndexerMetadata) @@ -768,12 +713,18 @@ def rocm_aiter_sparse_attn_indexer_native( topk_indices = topk_indices_buffer[ chunk.token_start : chunk.token_end, :topk_tokens ] - _topk_indices_prefill( + + num_rows = logits.shape[0] + + torch.ops._C.top_k_per_row_prefill( logits, - topk_tokens, - topk_indices, chunk.cu_seqlen_ks, chunk.cu_seqlen_ke, + topk_indices, + num_rows, + logits.stride(0), + logits.stride(1), + topk_tokens, ) if has_decode: @@ -811,16 +762,18 @@ def rocm_aiter_sparse_attn_indexer_native( max_model_len=max_model_len, ) - # Size the view to num_padded_tokens: top_k_per_row_decode writes - # logits.shape[0] == num_padded_tokens rows, and the unpack below - # reshapes to (batch_size, next_n, ...). topk_indices = topk_indices_buffer[:num_padded_tokens, :topk_tokens] - _topk_indices_decode( + num_rows = logits.shape[0] + + torch.ops._C.top_k_per_row_decode( logits, - topk_tokens, - topk_indices, - decode_metadata.seq_lens, next_n, + decode_metadata.seq_lens, + topk_indices, + num_rows, + logits.stride(0), + logits.stride(1), + topk_tokens, ) if decode_metadata.requires_padding: @@ -837,40 +790,6 @@ def rocm_aiter_sparse_attn_indexer_native( return topk_indices_buffer -@eager_break_during_capture -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 (