[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(
|
||||
|
||||
Reference in New Issue
Block a user