diff --git a/python/sglang/srt/layers/attention/aiter_backend.py b/python/sglang/srt/layers/attention/aiter_backend.py index 0d7c62d46..3fd9f76a7 100755 --- a/python/sglang/srt/layers/attention/aiter_backend.py +++ b/python/sglang/srt/layers/attention/aiter_backend.py @@ -2209,6 +2209,57 @@ class AiterAttnBackend(AttentionBackend): and layer.qk_head_dim == layer.v_head_dim ) + def init_mha_chunk_metadata( + self, forward_batch: ForwardBatch, disable_flashinfer_ragged: bool = False + ) -> None: + pass + + def _forward_extend_prefix_chunk( + self, + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + layer: RadixAttention, + forward_batch: ForwardBatch, + ): + idx = forward_batch.prefix_chunk_idx + output, lse = flash_attn_varlen_func( + q, + k, + v, + self.forward_metadata.qo_indptr, + forward_batch.prefix_chunk_cu_seq_lens[idx], + self.forward_metadata.max_q_len, + forward_batch.prefix_chunk_max_seq_lens[idx], + softmax_scale=layer.scaling, + causal=False, + return_lse=True, + )[:2] + return output, lse.transpose(0, 1).contiguous() + + def _forward_extend_skip_prefix( + self, + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + layer: RadixAttention, + ): + qo_indptr = self.forward_metadata.qo_indptr + max_q_len = self.forward_metadata.max_q_len + output, lse = flash_attn_varlen_func( + q, + k, + v, + qo_indptr, + qo_indptr, + max_q_len, + max_q_len, + softmax_scale=layer.scaling, + causal=True, + return_lse=True, + )[:2] + return output, lse.transpose(0, 1).contiguous() + def forward_extend( self, q: torch.Tensor, @@ -2221,6 +2272,9 @@ class AiterAttnBackend(AttentionBackend): ): self.logits_soft_cap = layer.logit_cap + if forward_batch.attn_attend_prefix_cache: + return self._forward_extend_prefix_chunk(q, k, v, layer, forward_batch) + cache_loc = ( forward_batch.out_cache_loc if not layer.is_cross_attention @@ -2332,6 +2386,8 @@ class AiterAttnBackend(AttentionBackend): and not forward_batch.forward_mode.is_draft_extend_v2() ): extend_no_prefix = not any(forward_batch.extend_prefix_lens_cpu) + if forward_batch.mha_return_lse: + return self._forward_extend_skip_prefix(q, k, v, layer) if kv_indices.shape[0] == 0 or extend_no_prefix: if self.use_fp8_prefill_attn and self.head_pad_mode != "zero": output = self.mla_fp8_prefill_attn( diff --git a/python/sglang/srt/models/deepseek_common/attention_backend_handler.py b/python/sglang/srt/models/deepseek_common/attention_backend_handler.py index f0d7f32c0..bda5c274e 100644 --- a/python/sglang/srt/models/deepseek_common/attention_backend_handler.py +++ b/python/sglang/srt/models/deepseek_common/attention_backend_handler.py @@ -23,7 +23,7 @@ from sglang.srt.utils import ( use_intel_amx_backend, ) -MHA_ONE_SHOT_SUPPORTED_BACKENDS = ["fa3", "flashinfer", "flashmla"] +MHA_ONE_SHOT_SUPPORTED_BACKENDS = ["fa3", "flashinfer", "flashmla", "aiter"] # ROCm runs dedicated MHA/MLA implementations (forward_mha_rocm.py / # forward_mla_rocm.py) so the shared CUDA paths carry no AMD branches. Backend @@ -33,6 +33,7 @@ MHA_ONE_SHOT_SUPPORTED_BACKENDS = ["fa3", "flashinfer", "flashmla"] _ROCM_FORWARD_METHODS = { AttnForwardMethod.MHA: AttnForwardMethod.MHA_ROCM, AttnForwardMethod.MHA_ONE_SHOT: AttnForwardMethod.MHA_ONE_SHOT_ROCM, + AttnForwardMethod.MHA_CHUNKED_KV: AttnForwardMethod.MHA_CHUNKED_KV_ROCM, AttnForwardMethod.MLA: AttnForwardMethod.MLA_ROCM, } @@ -201,6 +202,8 @@ def handle_attention_aiter(attn, forward_batch): if is_in_tc_piecewise_cuda_graph() or is_in_breakable_cuda_graph(): return AttnForwardMethod.MHA if forward_batch.forward_mode.is_extend_without_speculative(): + if not _support_mha_one_shot(attn, forward_batch, "aiter"): + return AttnForwardMethod.MHA_CHUNKED_KV return AttnForwardMethod.MHA else: return AttnForwardMethod.MLA diff --git a/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_methods.py b/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_methods.py index 09e6e2a5a..942fc37d7 100644 --- a/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_methods.py +++ b/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_methods.py @@ -37,5 +37,8 @@ class AttnForwardMethod(IntEnum): # Use one-shot multi-head attention for ROCm MHA_ONE_SHOT_ROCM = auto() + # Use multi-head attention with chunked kv cache for ROCm + MHA_CHUNKED_KV_ROCM = auto() + # Use absorbed multi-latent attention for ROCm MLA_ROCM = auto() diff --git a/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mha.py b/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mha.py index 13b18488b..5c33248f7 100644 --- a/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mha.py +++ b/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mha.py @@ -612,7 +612,7 @@ class DeepseekMHAForwardMixin: dst_dtype: torch.dtype, forward_batch: ForwardBatch, ): - if _is_cuda: + if _is_cuda or _use_aiter_gfx95: kv_a, k_pe = get_token_to_kv_pool().get_mla_kv_buffer( self.attn_mha, kv_indices, dst_dtype ) diff --git a/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mha_rocm.py b/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mha_rocm.py index 6213cfb86..6b455d1d6 100644 --- a/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mha_rocm.py +++ b/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mha_rocm.py @@ -268,6 +268,18 @@ class DeepseekMHARocmForwardMixin: positions, hidden_states, forward_batch, zero_allocator ) + def forward_normal_chunked_kv_rocm_prepare( + self: DeepseekV2AttentionMLA, + positions: torch.Tensor, + hidden_states: torch.Tensor, + forward_batch: ForwardBatch, + zero_allocator: BumpAllocator, + ): + # First do normal mha forward to get output for extended part + return self.forward_normal_rocm_prepare( + positions, hidden_states, forward_batch, zero_allocator + ) + def _concat_and_cast_mha_k_rocm( self: DeepseekV2AttentionMLA, k_nope: torch.Tensor, diff --git a/python/sglang/srt/models/deepseek_v2.py b/python/sglang/srt/models/deepseek_v2.py index e4f54b59b..27b3da146 100644 --- a/python/sglang/srt/models/deepseek_v2.py +++ b/python/sglang/srt/models/deepseek_v2.py @@ -2136,6 +2136,10 @@ class DeepseekV2AttentionMLA( inner_state = self.forward_normal_one_shot_rocm_prepare( positions, hidden_states, forward_batch, zero_allocator ) + elif attn_forward_method == AttnForwardMethod.MHA_CHUNKED_KV_ROCM: + inner_state = self.forward_normal_chunked_kv_rocm_prepare( + positions, hidden_states, forward_batch, zero_allocator + ) elif attn_forward_method == AttnForwardMethod.MLA_ROCM: inner_state = self.forward_absorb_rocm_prepare( positions, @@ -2204,6 +2208,8 @@ class DeepseekV2AttentionMLA( return self.forward_normal_core(*inner_state) elif attn_forward_method == AttnForwardMethod.MHA_ONE_SHOT_ROCM: return self.forward_normal_one_shot_core(*inner_state) + elif attn_forward_method == AttnForwardMethod.MHA_CHUNKED_KV_ROCM: + return self.forward_normal_chunked_kv_core(*inner_state) elif attn_forward_method == AttnForwardMethod.MLA_ROCM: return self.forward_absorb_rocm_core(*inner_state) elif attn_forward_method == AttnForwardMethod.MLA_FUSED_ROPE_ROCM: