From 86e26220978545921fd16e70cc057de66b276b01 Mon Sep 17 00:00:00 2001 From: kk <43161300+kkHuang-amd@users.noreply.github.com> Date: Wed, 25 Mar 2026 13:38:02 +0800 Subject: [PATCH] [AMD] Add mha fp8-kv support (#21253) Co-authored-by: wunhuang --- .../srt/layers/attention/aiter_backend.py | 85 ++++++++++--------- 1 file changed, 45 insertions(+), 40 deletions(-) diff --git a/python/sglang/srt/layers/attention/aiter_backend.py b/python/sglang/srt/layers/attention/aiter_backend.py index 7c9afd310..23875a653 100755 --- a/python/sglang/srt/layers/attention/aiter_backend.py +++ b/python/sglang/srt/layers/attention/aiter_backend.py @@ -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)