[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
|
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(
|
def forward_extend(
|
||||||
self,
|
self,
|
||||||
q: torch.Tensor,
|
q: torch.Tensor,
|
||||||
@@ -2221,6 +2272,9 @@ class AiterAttnBackend(AttentionBackend):
|
|||||||
):
|
):
|
||||||
self.logits_soft_cap = layer.logit_cap
|
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 = (
|
cache_loc = (
|
||||||
forward_batch.out_cache_loc
|
forward_batch.out_cache_loc
|
||||||
if not layer.is_cross_attention
|
if not layer.is_cross_attention
|
||||||
@@ -2332,6 +2386,8 @@ class AiterAttnBackend(AttentionBackend):
|
|||||||
and not forward_batch.forward_mode.is_draft_extend_v2()
|
and not forward_batch.forward_mode.is_draft_extend_v2()
|
||||||
):
|
):
|
||||||
extend_no_prefix = not any(forward_batch.extend_prefix_lens_cpu)
|
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 kv_indices.shape[0] == 0 or extend_no_prefix:
|
||||||
if self.use_fp8_prefill_attn and self.head_pad_mode != "zero":
|
if self.use_fp8_prefill_attn and self.head_pad_mode != "zero":
|
||||||
output = self.mla_fp8_prefill_attn(
|
output = self.mla_fp8_prefill_attn(
|
||||||
|
|||||||
@@ -23,7 +23,7 @@ from sglang.srt.utils import (
|
|||||||
use_intel_amx_backend,
|
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 /
|
# ROCm runs dedicated MHA/MLA implementations (forward_mha_rocm.py /
|
||||||
# forward_mla_rocm.py) so the shared CUDA paths carry no AMD branches. Backend
|
# 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 = {
|
_ROCM_FORWARD_METHODS = {
|
||||||
AttnForwardMethod.MHA: AttnForwardMethod.MHA_ROCM,
|
AttnForwardMethod.MHA: AttnForwardMethod.MHA_ROCM,
|
||||||
AttnForwardMethod.MHA_ONE_SHOT: AttnForwardMethod.MHA_ONE_SHOT_ROCM,
|
AttnForwardMethod.MHA_ONE_SHOT: AttnForwardMethod.MHA_ONE_SHOT_ROCM,
|
||||||
|
AttnForwardMethod.MHA_CHUNKED_KV: AttnForwardMethod.MHA_CHUNKED_KV_ROCM,
|
||||||
AttnForwardMethod.MLA: AttnForwardMethod.MLA_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():
|
if is_in_tc_piecewise_cuda_graph() or is_in_breakable_cuda_graph():
|
||||||
return AttnForwardMethod.MHA
|
return AttnForwardMethod.MHA
|
||||||
if forward_batch.forward_mode.is_extend_without_speculative():
|
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
|
return AttnForwardMethod.MHA
|
||||||
else:
|
else:
|
||||||
return AttnForwardMethod.MLA
|
return AttnForwardMethod.MLA
|
||||||
|
|||||||
@@ -37,5 +37,8 @@ class AttnForwardMethod(IntEnum):
|
|||||||
# Use one-shot multi-head attention for ROCm
|
# Use one-shot multi-head attention for ROCm
|
||||||
MHA_ONE_SHOT_ROCM = auto()
|
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
|
# Use absorbed multi-latent attention for ROCm
|
||||||
MLA_ROCM = auto()
|
MLA_ROCM = auto()
|
||||||
|
|||||||
@@ -612,7 +612,7 @@ class DeepseekMHAForwardMixin:
|
|||||||
dst_dtype: torch.dtype,
|
dst_dtype: torch.dtype,
|
||||||
forward_batch: ForwardBatch,
|
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(
|
kv_a, k_pe = get_token_to_kv_pool().get_mla_kv_buffer(
|
||||||
self.attn_mha, kv_indices, dst_dtype
|
self.attn_mha, kv_indices, dst_dtype
|
||||||
)
|
)
|
||||||
|
|||||||
+12
@@ -268,6 +268,18 @@ class DeepseekMHARocmForwardMixin:
|
|||||||
positions, hidden_states, forward_batch, zero_allocator
|
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(
|
def _concat_and_cast_mha_k_rocm(
|
||||||
self: DeepseekV2AttentionMLA,
|
self: DeepseekV2AttentionMLA,
|
||||||
k_nope: torch.Tensor,
|
k_nope: torch.Tensor,
|
||||||
|
|||||||
@@ -2136,6 +2136,10 @@ class DeepseekV2AttentionMLA(
|
|||||||
inner_state = self.forward_normal_one_shot_rocm_prepare(
|
inner_state = self.forward_normal_one_shot_rocm_prepare(
|
||||||
positions, hidden_states, forward_batch, zero_allocator
|
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:
|
elif attn_forward_method == AttnForwardMethod.MLA_ROCM:
|
||||||
inner_state = self.forward_absorb_rocm_prepare(
|
inner_state = self.forward_absorb_rocm_prepare(
|
||||||
positions,
|
positions,
|
||||||
@@ -2204,6 +2208,8 @@ class DeepseekV2AttentionMLA(
|
|||||||
return self.forward_normal_core(*inner_state)
|
return self.forward_normal_core(*inner_state)
|
||||||
elif attn_forward_method == AttnForwardMethod.MHA_ONE_SHOT_ROCM:
|
elif attn_forward_method == AttnForwardMethod.MHA_ONE_SHOT_ROCM:
|
||||||
return self.forward_normal_one_shot_core(*inner_state)
|
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:
|
elif attn_forward_method == AttnForwardMethod.MLA_ROCM:
|
||||||
return self.forward_absorb_rocm_core(*inner_state)
|
return self.forward_absorb_rocm_core(*inner_state)
|
||||||
elif attn_forward_method == AttnForwardMethod.MLA_FUSED_ROPE_ROCM:
|
elif attn_forward_method == AttnForwardMethod.MLA_FUSED_ROPE_ROCM:
|
||||||
|
|||||||
Reference in New Issue
Block a user