[AMD] Add mha fp8-kv support (#21253)

Co-authored-by: wunhuang <wunhuang@amd.com>
This commit is contained in:
kk
2026-03-24 22:38:02 -07:00
committed by GitHub
co-authored by wunhuang
parent 2b75fed0dd
commit 86e2622097
@@ -753,9 +753,7 @@ class AiterAttnBackend(AttentionBackend):
self._ensure_spec_v2_topk_supported()
if self.use_mla:
device = forward_batch.seq_lens.device
num_draft_tokens = self._resolve_v2_num_draft_tokens(
extend_seq_lens=forward_batch.extend_seq_lens
)
num_draft_tokens = self._resolve_v2_num_draft_tokens()
qo_indptr = self._set_uniform_qo_indptr(bs, num_draft_tokens, device)
kv_indptr = self.kv_indptr[: bs + 1]
@@ -1153,7 +1151,9 @@ class AiterAttnBackend(AttentionBackend):
max_num_tokens: int,
kv_indices_buf: Optional[torch.Tensor] = None,
):
self.cuda_graph_kv_last_page_len = torch.ones(max_bs, dtype=torch.int)
self.cuda_graph_kv_last_page_len = torch.ones(
max_bs, dtype=torch.int, device=self.device
)
if kv_indices_buf is None:
max_num_blocks_per_seq = (
self.max_context_len + self.page_size - 1
@@ -1994,6 +1994,12 @@ class AiterAttnBackend(AttentionBackend):
else forward_batch.encoder_out_cache_loc
)
k_descale = None
v_descale = None
if self.kv_cache_dtype == fp8_dtype:
k_descale = layer.k_scale if layer.k_scale is not None else self.k_scale
v_descale = layer.v_scale if layer.v_scale is not None else self.k_scale
if k is not None:
assert v is not None
if save_kv_cache:
@@ -2027,13 +2033,15 @@ class AiterAttnBackend(AttentionBackend):
if layer.sliding_window_size > 0
else None
),
k_scale=k_descale,
v_scale=v_descale,
)
elif self.use_mla:
forward_batch.token_to_kv_pool.set_kv_buffer(layer, cache_loc, k, v)
else:
forward_batch.token_to_kv_pool.set_kv_buffer(
layer, cache_loc, k, v, layer.k_scale, layer.v_scale
layer, cache_loc, k, v, k_descale, v_descale
)
if self.use_mla:
@@ -2202,12 +2210,8 @@ class AiterAttnBackend(AttentionBackend):
reduce_indptr=reduce_indptr,
reduce_final_map=reduce_final_map,
reduce_partial_map=reduce_partial_map,
q_scale=(
layer.k_scale if layer.k_scale is not None else self.k_scale
),
kv_scale=(
layer.k_scale if layer.k_scale is not None else self.k_scale
),
q_scale=k_descale,
kv_scale=k_descale,
intra_batch_mode=intra_batch_mode,
num_kv_splits=num_kv_splits,
)
@@ -2260,12 +2264,8 @@ class AiterAttnBackend(AttentionBackend):
reduce_indptr=reduce_indptr,
reduce_final_map=reduce_final_map,
reduce_partial_map=reduce_partial_map,
q_scale=(
layer.k_scale if layer.k_scale is not None else self.k_scale
),
kv_scale=(
layer.k_scale if layer.k_scale is not None else self.k_scale
),
q_scale=k_descale,
kv_scale=k_descale,
intra_batch_mode=intra_batch_mode,
num_kv_splits=num_kv_splits,
)
@@ -2295,12 +2295,8 @@ class AiterAttnBackend(AttentionBackend):
reduce_indptr=reduce_indptr,
reduce_final_map=reduce_final_map,
reduce_partial_map=reduce_partial_map,
q_scale=(
layer.k_scale if layer.k_scale is not None else self.k_scale
),
kv_scale=(
layer.k_scale if layer.k_scale is not None else self.k_scale
),
q_scale=k_descale,
kv_scale=k_descale,
intra_batch_mode=intra_batch_mode,
num_kv_splits=num_kv_splits,
)
@@ -2349,11 +2345,14 @@ class AiterAttnBackend(AttentionBackend):
bs0 = forward_batch.batch_size + 1
# To keep the mha_batch_prefill_func function parameters
# declare the necessary parameter and assign None as default value
q_descale = None
# TODO kkhuang-amd need to remove it when mha_batch_prefill_func support fp8-kv
if self.kv_cache_dtype == fp8_dtype:
dtype = q.dtype
k_cache = k_cache.to(dtype)
v_cache = v_cache.to(dtype)
q = q.to(fp8_dtype)
q_descale = layer.k_scale if layer.k_scale is not None else self.k_scale
window_size = (-1, -1)
page_table = self.forward_metadata.kv_indices
@@ -2379,6 +2378,9 @@ class AiterAttnBackend(AttentionBackend):
return_attn_probs=False,
window_size=window_size,
sink_ptr=sinks,
q_descale=q_descale,
k_descale=k_descale,
v_descale=v_descale,
)
return o.view(-1, layer.tp_q_head_num * layer.head_dim)
@@ -2404,6 +2406,12 @@ class AiterAttnBackend(AttentionBackend):
else:
o = torch.empty_like(q, dtype=self.input_dtype)
k_descale = None
v_descale = None
if self.kv_cache_dtype == fp8_dtype:
k_descale = layer.k_scale if layer.k_scale is not None else self.k_scale
v_descale = layer.v_scale if layer.v_scale is not None else self.k_scale
if save_kv_cache:
# Only use SWA-specific kv cache write (reshape_and_cache_flash) when
# both unified attention and sliding window kv pool are active.
@@ -2428,6 +2436,8 @@ class AiterAttnBackend(AttentionBackend):
),
forward_batch.out_cache_loc,
slot_mapping_swa.long() if layer.sliding_window_size > 0 else None,
k_scale=k_descale,
v_scale=v_descale,
)
else:
forward_batch.token_to_kv_pool.set_kv_buffer(
@@ -2464,8 +2474,8 @@ class AiterAttnBackend(AttentionBackend):
reduce_indptr=reduce_indptr,
reduce_final_map=reduce_final_map,
reduce_partial_map=reduce_partial_map,
q_scale=layer.k_scale if layer.k_scale is not None else self.k_scale,
kv_scale=layer.k_scale if layer.k_scale is not None else self.k_scale,
q_scale=k_descale,
kv_scale=k_descale,
intra_batch_mode=intra_batch_mode,
num_kv_splits=num_kv_splits,
)
@@ -2476,13 +2486,6 @@ class AiterAttnBackend(AttentionBackend):
layer.layer_id
)
# TODO kkhuang-amd need to remove it when paged_attention_ragged support fp8-kv
if self.kv_cache_dtype == fp8_dtype:
dtype = q.dtype
k_cache = k_cache.to(dtype)
v_cache = v_cache.to(dtype)
if self.use_triton_unified_attention:
bs = forward_batch.batch_size
@@ -2497,7 +2500,7 @@ 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)
o = torch.empty_like(q, dtype=self.input_dtype)
max_kv_len = page_table.shape[1]
@@ -2520,11 +2523,15 @@ class AiterAttnBackend(AttentionBackend):
block_table=page_table,
softcap=0,
q_descale=None,
k_descale=None,
v_descale=None,
k_descale=k_descale,
v_descale=v_descale,
sinks=sinks,
)
else:
if self.kv_cache_dtype == fp8_dtype:
k_cache = k_cache.to(self.input_dtype)
v_cache = v_cache.to(self.input_dtype)
paged_attention_ragged(
o.view(-1, layer.tp_q_head_num, layer.qk_head_dim),
self.workspace_buffer,
@@ -2634,8 +2641,6 @@ class AiterIndicesUpdaterPrefill:
token_num = kv_indptr[-1]
kv_indices[token_num:] = kv_indices[0]
# self.max_kv_len = torch.max(paged_kernel_lens).item()
extend_lens = seq_lens - prefix_lens
qo_indptr[1 : bs + 1] = torch.cumsum(extend_lens, dim=0)