[AMD] Add mha fp8-kv support (#21253)
Co-authored-by: wunhuang <wunhuang@amd.com>
This commit is contained in:
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user