[AMD] Support aiter fa mha chunked kv for Kimi-K3 (#37691)
Co-authored-by: HAI <hixiao@gmail.com>
This commit is contained in:
@@ -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
|
||||
)
|
||||
|
||||
+12
@@ -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,
|
||||
|
||||
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user