[AMD] Support aiter fa mha chunked kv for Kimi-K3 (#37691)

Co-authored-by: HAI <hixiao@gmail.com>
This commit is contained in:
billishyahao
2026-09-07 00:31:24 -07:00
committed by GitHub
co-authored by HAI
parent 4d23a4fa6d
commit 644841c50c
6 changed files with 82 additions and 2 deletions
@@ -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(
@@ -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
@@ -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()
@@ -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
)
@@ -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,
+6
View File
@@ -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: