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