[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:
RoyWang
2026-04-04 22:13:29 -07:00
committed by GitHub
co-authored by RoyWang
parent 8cbeacd783
commit dd49127fe6
3 changed files with 81 additions and 83 deletions
@@ -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,
}
+2 -9
View File
@@ -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",