[AMD]: Support MLA with nhead<16 and FP8 KV cache for TP=8 (Kimi K2.5… (#21213)
Co-authored-by: RoyWang <RoyWang@amd.com>
This commit is contained in:
@@ -234,13 +234,25 @@ class AiterAttnBackend(AttentionBackend):
|
||||
self.forward_metadata: ForwardMetadata = None
|
||||
|
||||
if self.use_mla:
|
||||
_valid_heads = self.num_head in (4, 8) or (
|
||||
self.num_head % 16 == 0 and 16 <= self.num_head <= 128
|
||||
)
|
||||
assert _valid_heads, (
|
||||
f"Aiter MLA supports num_head of 4, 8, or multiples of 16 "
|
||||
f"in [16, 128].\n"
|
||||
f"Provided {self.num_head} number of heads.\n"
|
||||
"Try adjusting tensor_parallel_size value."
|
||||
)
|
||||
self.num_head_padded = 16 if self.num_head < 16 else self.num_head
|
||||
self.head_repeat_factor = 16 // self.num_head if self.num_head < 16 else 1
|
||||
|
||||
self.enable_dp_attention = is_dp_attention_enabled()
|
||||
self.qo_indptr_ = torch.zeros(
|
||||
(max_bs + 1,), dtype=torch.int32, device=model_runner.device
|
||||
)
|
||||
global _use_mla_ps_kernel, fast_mode, intra_batch_mode
|
||||
|
||||
# current mla_decode_fwd onln support fake-nps in self.num_head == 16
|
||||
# current mla_decode_fwd only support fake-nps in self.num_head == 16
|
||||
# so all num_head size does not use qh16 kernel to simulate
|
||||
# it should not use fake-nps (fast_mode = False, intra_batch_mode = True)
|
||||
# it will cause gpu-fault or accuracy issue
|
||||
@@ -254,7 +266,7 @@ class AiterAttnBackend(AttentionBackend):
|
||||
# for non-fp8 kv_cache on tp8, use non-persist kernel to avoid performance degradation
|
||||
# head_num=16 (tp8 perf issue), head_num=128 (unsupported, like tp1 or --enable-dp-attention with tp8-dp8)
|
||||
if (
|
||||
self.num_head == 16 or self.num_head == 128
|
||||
self.num_head_padded == 16 or self.num_head_padded == 128
|
||||
) and self.kv_cache_dtype is not fp8_dtype:
|
||||
_use_mla_ps_kernel = False
|
||||
fast_mode = False
|
||||
@@ -268,7 +280,7 @@ class AiterAttnBackend(AttentionBackend):
|
||||
self.fix_max_split_per_batch = self.max_split_per_batch
|
||||
|
||||
def make_mla_decode_meta_data_buffer(self, max_seqlen_qo, batch_size):
|
||||
nhead = self.num_head
|
||||
nhead = self.num_head_padded
|
||||
dtype = self.kv_cache_dtype
|
||||
|
||||
if self.enable_dp_attention:
|
||||
@@ -355,7 +367,7 @@ class AiterAttnBackend(AttentionBackend):
|
||||
qo_indptr,
|
||||
kv_indptr,
|
||||
kv_last_page_len,
|
||||
self.num_head // nhead_kv,
|
||||
self.num_head_padded // nhead_kv,
|
||||
nhead_kv,
|
||||
False,
|
||||
work_metadata,
|
||||
@@ -541,6 +553,36 @@ class AiterAttnBackend(AttentionBackend):
|
||||
f"Got topk={self.topk}."
|
||||
)
|
||||
|
||||
def _mla_decode_fwd_with_head_pad(
|
||||
self,
|
||||
q: torch.Tensor,
|
||||
k_buffer_flat: torch.Tensor,
|
||||
layer,
|
||||
**kwargs,
|
||||
):
|
||||
"""Wrap mla_decode_fwd with head-dimension padding for num_head < 16.
|
||||
|
||||
When head_repeat_factor > 1 (i.e. num_head is 4 or 8), q is
|
||||
repeat-interleaved to reach num_head_padded (16) before the kernel
|
||||
call, and the corresponding output columns are sliced back afterward.
|
||||
q / o must already be shaped (..., num_head, head_dim).
|
||||
"""
|
||||
if self.head_repeat_factor > 1:
|
||||
q_in = q.repeat_interleave(self.head_repeat_factor, dim=1)
|
||||
o = q.new_empty(
|
||||
(q.shape[0], self.num_head_padded, layer.v_head_dim),
|
||||
dtype=self.input_dtype,
|
||||
)
|
||||
mla_decode_fwd(q_in, k_buffer_flat, o, **kwargs)
|
||||
return o[:, :: self.head_repeat_factor, :]
|
||||
else:
|
||||
o = q.new_empty(
|
||||
(q.shape[0], layer.tp_q_head_num, layer.v_head_dim),
|
||||
dtype=self.input_dtype,
|
||||
)
|
||||
mla_decode_fwd(q, k_buffer_flat, o, **kwargs)
|
||||
return o
|
||||
|
||||
def mla_fp8_prefill_attn(
|
||||
self,
|
||||
q: torch.Tensor,
|
||||
@@ -2178,11 +2220,6 @@ class AiterAttnBackend(AttentionBackend):
|
||||
K_Buffer = K_Buffer.view(-1, layer.tp_k_head_num, layer.qk_head_dim)
|
||||
return o
|
||||
elif forward_batch.forward_mode.is_target_verify():
|
||||
o = q.new_empty(
|
||||
(q.shape[0], layer.tp_q_head_num, layer.v_head_dim),
|
||||
dtype=self.input_dtype,
|
||||
)
|
||||
|
||||
work_metadata = self.forward_metadata.work_metadata
|
||||
work_indptr = self.forward_metadata.work_indptr
|
||||
work_info_set = self.forward_metadata.work_info_set
|
||||
@@ -2193,15 +2230,15 @@ class AiterAttnBackend(AttentionBackend):
|
||||
|
||||
num_kv_splits = self.forward_metadata.num_kv_splits
|
||||
|
||||
mla_decode_fwd(
|
||||
o = self._mla_decode_fwd_with_head_pad(
|
||||
q,
|
||||
K_Buffer.view(-1, 1, 1, layer.qk_head_dim),
|
||||
o,
|
||||
self.forward_metadata.qo_indptr,
|
||||
self.forward_metadata.kv_indptr,
|
||||
self.forward_metadata.kv_indices,
|
||||
self.forward_metadata.kv_last_page_len,
|
||||
self.forward_metadata.max_q_len,
|
||||
layer,
|
||||
qo_indptr=self.forward_metadata.qo_indptr,
|
||||
kv_indptr=self.forward_metadata.kv_indptr,
|
||||
kv_indices=self.forward_metadata.kv_indices,
|
||||
kv_last_page_lens=self.forward_metadata.kv_last_page_len,
|
||||
max_seqlen_q=self.forward_metadata.max_q_len,
|
||||
sm_scale=layer.scaling,
|
||||
logit_cap=layer.logit_cap,
|
||||
work_meta_data=work_metadata,
|
||||
@@ -2232,30 +2269,21 @@ class AiterAttnBackend(AttentionBackend):
|
||||
num_kv_splits = self.forward_metadata.num_kv_splits
|
||||
|
||||
if self.forward_metadata.run_graph is not True:
|
||||
|
||||
bs, q_pad, q_mask = pad_sequence_with_mask(
|
||||
q.view(q.shape[0], -1),
|
||||
qo_indptr[:-1],
|
||||
forward_batch.extend_seq_lens,
|
||||
self.forward_metadata.max_q_len,
|
||||
)
|
||||
o = q.new_empty(
|
||||
(
|
||||
bs * self.forward_metadata.max_q_len,
|
||||
layer.tp_q_head_num,
|
||||
layer.v_head_dim,
|
||||
),
|
||||
dtype=self.input_dtype,
|
||||
)
|
||||
mla_decode_fwd(
|
||||
o = self._mla_decode_fwd_with_head_pad(
|
||||
q_pad.view(-1, layer.tp_q_head_num, layer.qk_head_dim),
|
||||
K_Buffer.view(-1, 1, 1, layer.qk_head_dim),
|
||||
o,
|
||||
self.forward_metadata.qo_indptr,
|
||||
self.forward_metadata.kv_indptr,
|
||||
self.forward_metadata.kv_indices,
|
||||
self.forward_metadata.kv_last_page_len,
|
||||
self.forward_metadata.max_q_len,
|
||||
layer,
|
||||
qo_indptr=self.forward_metadata.qo_indptr,
|
||||
kv_indptr=self.forward_metadata.kv_indptr,
|
||||
kv_indices=self.forward_metadata.kv_indices,
|
||||
kv_last_page_lens=self.forward_metadata.kv_last_page_len,
|
||||
max_seqlen_q=self.forward_metadata.max_q_len,
|
||||
sm_scale=layer.scaling,
|
||||
logit_cap=layer.logit_cap,
|
||||
work_meta_data=work_metadata,
|
||||
@@ -2273,20 +2301,15 @@ class AiterAttnBackend(AttentionBackend):
|
||||
total_valid_q = int(qo_indptr[-1].item())
|
||||
return o[:total_valid_q]
|
||||
else:
|
||||
o = q.new_empty(
|
||||
(q.shape[0], layer.tp_q_head_num, layer.v_head_dim),
|
||||
dtype=self.input_dtype,
|
||||
)
|
||||
|
||||
mla_decode_fwd(
|
||||
o = self._mla_decode_fwd_with_head_pad(
|
||||
q,
|
||||
K_Buffer.view(-1, 1, 1, layer.qk_head_dim),
|
||||
o,
|
||||
self.forward_metadata.qo_indptr,
|
||||
self.forward_metadata.kv_indptr,
|
||||
self.forward_metadata.kv_indices,
|
||||
self.forward_metadata.kv_last_page_len,
|
||||
self.forward_metadata.max_q_len,
|
||||
layer,
|
||||
qo_indptr=self.forward_metadata.qo_indptr,
|
||||
kv_indptr=self.forward_metadata.kv_indptr,
|
||||
kv_indices=self.forward_metadata.kv_indices,
|
||||
kv_last_page_lens=self.forward_metadata.kv_last_page_len,
|
||||
max_seqlen_q=self.forward_metadata.max_q_len,
|
||||
sm_scale=layer.scaling,
|
||||
logit_cap=layer.logit_cap,
|
||||
work_meta_data=work_metadata,
|
||||
@@ -2395,17 +2418,8 @@ class AiterAttnBackend(AttentionBackend):
|
||||
save_kv_cache=True,
|
||||
sinks=None,
|
||||
):
|
||||
|
||||
q = q.reshape(-1, layer.tp_q_head_num * layer.qk_head_dim)
|
||||
|
||||
if layer.qk_head_dim != layer.v_head_dim:
|
||||
o = q.new_empty(
|
||||
(q.shape[0], layer.tp_q_head_num * layer.v_head_dim),
|
||||
dtype=self.input_dtype,
|
||||
)
|
||||
else:
|
||||
o = torch.empty_like(q, dtype=self.input_dtype)
|
||||
|
||||
k_descale = None
|
||||
v_descale = None
|
||||
if self.kv_cache_dtype == fp8_dtype:
|
||||
@@ -2458,15 +2472,15 @@ class AiterAttnBackend(AttentionBackend):
|
||||
|
||||
num_kv_splits = self.forward_metadata.num_kv_splits
|
||||
|
||||
mla_decode_fwd(
|
||||
o = self._mla_decode_fwd_with_head_pad(
|
||||
q.view(-1, layer.tp_q_head_num, layer.qk_head_dim),
|
||||
k_buffer.view(-1, 1, 1, layer.qk_head_dim),
|
||||
o.view(-1, layer.tp_q_head_num, layer.v_head_dim),
|
||||
self.forward_metadata.qo_indptr,
|
||||
self.forward_metadata.kv_indptr,
|
||||
self.forward_metadata.kv_indices,
|
||||
self.forward_metadata.kv_last_page_len,
|
||||
self.forward_metadata.max_q_len,
|
||||
layer,
|
||||
qo_indptr=self.forward_metadata.qo_indptr,
|
||||
kv_indptr=self.forward_metadata.kv_indptr,
|
||||
kv_indices=self.forward_metadata.kv_indices,
|
||||
kv_last_page_lens=self.forward_metadata.kv_last_page_len,
|
||||
max_seqlen_q=self.forward_metadata.max_q_len,
|
||||
sm_scale=layer.scaling,
|
||||
logit_cap=layer.logit_cap,
|
||||
work_meta_data=work_metadata,
|
||||
@@ -2487,6 +2501,8 @@ class AiterAttnBackend(AttentionBackend):
|
||||
layer.layer_id
|
||||
)
|
||||
|
||||
o = torch.empty_like(q, dtype=self.input_dtype)
|
||||
|
||||
if self.use_triton_unified_attention:
|
||||
|
||||
bs = forward_batch.batch_size
|
||||
@@ -2501,8 +2517,6 @@ class AiterAttnBackend(AttentionBackend):
|
||||
if self.forward_metadata.swa_page_table is not None:
|
||||
page_table = self.forward_metadata.swa_page_table
|
||||
|
||||
o = torch.empty_like(q, dtype=self.input_dtype)
|
||||
|
||||
max_kv_len = page_table.shape[1] * self.page_size
|
||||
|
||||
unified_attention(
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
"""MI35x Kimi-K2.5-MXFP4 aiter MLA backend accuracy tests (4-GPU)
|
||||
"""MI35x Kimi-K2.5-MXFP4 aiter MLA backend accuracy tests (8-GPU)
|
||||
|
||||
Tests Kimi-K2.5-MXFP4 with the aiter unified attention backend on MI35x,
|
||||
covering both default and FP8 KV cache configurations.
|
||||
@@ -7,13 +7,6 @@ The FP8 KV cache variant validates the fix for assertion failure
|
||||
`q_scale.has_value() && kv_scale.has_value()` in aiter ASM MLA decode
|
||||
when layer.k_scale is None (the RadixAttention default).
|
||||
|
||||
NOTE: TP must be <= 4 for Kimi-K2.5 with the aiter MLA kernel.
|
||||
Kimi-K2.5 has num_attention_heads=64; with tp_size=8 that gives
|
||||
64/8 = 8 heads per GPU, but the aiter ASM MLA kernel requires
|
||||
heads_per_gpu % 16 == 0. With tp_size=4: 64/4 = 16 heads, which
|
||||
satisfies the constraint. (DeepSeek-R1/V3 has 128 heads so TP=8
|
||||
yields 128/8 = 16 heads and works fine.)
|
||||
|
||||
Registry: nightly-amd-8-gpu-mi35x-kimi-k25-mxfp4-aiter-mla suite
|
||||
"""
|
||||
|
||||
@@ -61,7 +54,7 @@ class ModelConfig:
|
||||
"""Configuration for a model variant to test."""
|
||||
|
||||
model_path: str
|
||||
tp_size: int = 4
|
||||
tp_size: int = 8
|
||||
accuracy_threshold: float = 0.92
|
||||
other_args: Optional[List[str]] = None
|
||||
env_vars: Optional[dict] = None
|
||||
@@ -85,9 +78,7 @@ def get_kimi_k25_mxfp4_models() -> List[ModelConfig]:
|
||||
model_path = get_model_path()
|
||||
common_kwargs = {
|
||||
"model_path": model_path,
|
||||
# TP=4 required: Kimi-K2.5 has 64 attn heads; aiter ASM MLA needs
|
||||
# heads_per_gpu % 16 == 0 -> 64/4=16 works, 64/8=8 does not.
|
||||
"tp_size": 4,
|
||||
"tp_size": 8,
|
||||
"accuracy_threshold": 0.92,
|
||||
"timeout": 3600,
|
||||
}
|
||||
|
||||
@@ -1,14 +1,8 @@
|
||||
"""Kimi-K2.5-MXFP4 aiter MLA backend test (4-GPU, FP8 KV cache)
|
||||
"""Kimi-K2.5-MXFP4 aiter MLA backend test (8-GPU, FP8 KV cache)
|
||||
|
||||
PR-level test for Kimi-K2.5-MXFP4 with aiter unified attention backend
|
||||
and fp8_e4m3 KV cache on MI35x.
|
||||
|
||||
NOTE: TP must be <= 4 for Kimi-K2.5 with the aiter MLA kernel.
|
||||
Kimi-K2.5 has num_attention_heads=64; with tp_size=8 that gives
|
||||
64/8 = 8 heads per GPU, but the aiter ASM MLA kernel requires
|
||||
heads_per_gpu % 16 == 0. With tp_size=4: 64/4 = 16 heads, which
|
||||
satisfies the constraint. (DeepSeek-R1/V3 has 128 heads so TP=8
|
||||
yields 128/8 = 16 heads and works fine.)
|
||||
"""
|
||||
|
||||
import os
|
||||
@@ -41,10 +35,9 @@ class TestKimiK25MXFP4(CustomTestCase):
|
||||
def setUpClass(cls):
|
||||
cls.model = KIMI_K25_MXFP4_MODEL_PATH
|
||||
cls.base_url = DEFAULT_URL_FOR_TEST
|
||||
# TP=4 required: 64 attn heads / 4 = 16 heads per GPU (aiter MLA needs % 16 == 0)
|
||||
other_args = [
|
||||
"--tp",
|
||||
"4",
|
||||
"8",
|
||||
"--attention-backend",
|
||||
"aiter",
|
||||
"--kv-cache-dtype",
|
||||
|
||||
Reference in New Issue
Block a user