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