[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(