Aiter fp8 kv cache (#13147)

This commit is contained in:
kk
2025-12-09 08:39:53 +08:00
committed by GitHub
parent 119fd956fb
commit c106b54b57
7 changed files with 594 additions and 96 deletions

View File

@@ -4,6 +4,7 @@ from __future__ import annotations
end to end attention solution with aiter kernels
"""
import logging
from dataclasses import dataclass
from enum import Enum, auto
from typing import TYPE_CHECKING, Optional
@@ -27,6 +28,8 @@ if TYPE_CHECKING:
try:
from aiter import (
flash_attn_varlen_func,
get_mla_metadata_info_v1,
get_mla_metadata_v1,
mha_batch_prefill_func,
paged_attention_ragged,
)
@@ -37,6 +40,21 @@ except ImportError:
)
from sglang.srt.configs.model_config import AttentionArch
from sglang.srt.layers.quantization.fp8_kernel import fp8_dtype
from sglang.srt.utils import get_bool_env_var
logger = logging.getLogger(__name__)
# Use aiter mla persist design for fp8-kv cache
_use_mla_ps_kernel = get_bool_env_var("SGLANG_AITER_MLA_PERSIST", "True")
# Persist
# fast_mode=True if _use_mla_ps_kernel else False
# intra_batch_mode=False if _use_mla_ps_kernel else True
# fake non-ps, intra_batch_mode needs to be True for non-ps-mode
fast_mode = False
intra_batch_mode = True if _use_mla_ps_kernel else False
class WrapperDispatch(Enum):
@@ -52,6 +70,14 @@ class ForwardMetadata:
kv_last_page_len: torch.Tensor
max_q_len: int
max_kv_len: Optional[int]
work_metadata: Optional[torch.Tensor] = None
work_info_set: Optional[torch.Tensor] = None
work_indptr: Optional[torch.Tensor] = None
reduce_indptr: Optional[torch.Tensor] = None
reduce_final_map: Optional[torch.Tensor] = None
reduce_partial_map: Optional[torch.Tensor] = None
num_kv_splits: Optional[int] = None
# num_kv_splits_indptr: Optional[torch.Tensor] = None
global_workspace_buffer = None
@@ -72,6 +98,10 @@ class AiterAttnBackend(AttentionBackend):
extend_attention_fwd,
)
self.input_dtype = model_runner.model_config.dtype
self.page_size = model_runner.server_args.page_size
self.extend_attention_fwd = torch.compiler.disable(extend_attention_fwd)
self.device = model_runner.device
@@ -154,6 +184,118 @@ class AiterAttnBackend(AttentionBackend):
self.enable_dp_attention = is_dp_attention_enabled()
self.max_split_per_batch = 32 if _use_mla_ps_kernel else None
if self.num_draft_tokens is None and _use_mla_ps_kernel:
self.max_split_per_batch = 64
self.fix_max_split_per_batch = self.max_split_per_batch
def make_mla_decode_meta_data_buffer(self, max_seqlen_qo, batch_size):
nhead = self.num_head
dtype = self.kv_cache_dtype
if self.enable_dp_attention:
gpu = torch.cuda.current_device()
device_properties = torch.cuda.get_device_properties(gpu)
cu_num = device_properties.multi_processor_count
self.max_split_per_batch = min(
(cu_num + batch_size - 1) // batch_size, self.fix_max_split_per_batch
)
(
(work_meta_data_size, work_meta_data_type),
(work_indptr_size, work_indptr_type),
(work_info_set_size, work_info_set_type),
(reduce_indptr_size, reduce_indptr_type),
(reduce_final_map_size, reduce_final_map_type),
(reduce_partial_map_size, reduce_partial_map_type),
) = get_mla_metadata_info_v1(
batch_size,
max_seqlen_qo,
nhead,
dtype,
dtype,
is_sparse=False,
fast_mode=fast_mode,
num_kv_splits=self.max_split_per_batch,
intra_batch_mode=intra_batch_mode,
)
# aiter implementation
# the tensor's meaning please refer aiter/ops/attention.py
work_metadata = torch.empty(
work_meta_data_size, dtype=work_meta_data_type, device="cuda"
)
work_indptr = torch.empty(
work_indptr_size, dtype=work_indptr_type, device="cuda"
)
work_info_set = torch.empty(
work_info_set_size,
dtype=work_info_set_type,
device="cuda",
)
reduce_indptr = torch.empty(
reduce_indptr_size, dtype=reduce_indptr_type, device="cuda"
)
reduce_final_map = torch.empty(
reduce_final_map_size, dtype=reduce_final_map_type, device="cuda"
)
reduce_partial_map = torch.empty(
reduce_partial_map_size, dtype=reduce_partial_map_type, device="cuda"
)
return (
work_metadata,
work_indptr,
work_info_set,
reduce_indptr,
reduce_final_map,
reduce_partial_map,
)
def make_mla_meta_data(
self,
qo_indptr,
kv_indptr,
work_metadata,
work_info_set,
work_indptr,
reduce_indptr,
reduce_final_map,
reduce_partial_map,
max_q_len,
fast_mode,
max_split_per_batch,
intra_batch_mode,
):
nhead_kv = 1
page_size = 1
dtype = self.kv_cache_dtype
meta = get_mla_metadata_v1(
qo_indptr,
kv_indptr,
self.num_head // nhead_kv,
nhead_kv,
True,
work_metadata,
work_info_set,
work_indptr,
reduce_indptr,
reduce_final_map,
reduce_partial_map,
kv_granularity=max(page_size, 16),
max_seqlen_qo=max_q_len,
uni_seqlen_qo=max_q_len,
fast_mode=fast_mode,
max_split_per_batch=max_split_per_batch,
intra_batch_mode=intra_batch_mode,
dtype_q=dtype,
dtype_kv=dtype,
)
def init_forward_metadata(self, forward_batch: ForwardBatch):
"""Init auxiliary variables for triton attention backend."""
@@ -164,6 +306,16 @@ class AiterAttnBackend(AttentionBackend):
kv_last_page_len = None
max_q_len = None
work_metadata = None
work_indptr = None
work_info_set = None
reduce_indptr = None
reduce_final_map = None
reduce_partial_map = None
num_kv_splits = None
# num_kv_splits_indptr = None
if forward_batch.forward_mode.is_decode_or_idle():
if spec_info is None:
kv_indptr[1 : bs + 1] = torch.cumsum(forward_batch.seq_lens, dim=0)
@@ -190,6 +342,33 @@ class AiterAttnBackend(AttentionBackend):
kv_last_page_len = self.kv_last_page_len[:bs]
max_q_len = 1
if _use_mla_ps_kernel:
(
work_metadata,
work_indptr,
work_info_set,
reduce_indptr,
reduce_final_map,
reduce_partial_map,
) = self.make_mla_decode_meta_data_buffer(max_q_len, bs)
num_kv_splits = self.max_split_per_batch
self.make_mla_meta_data(
qo_indptr,
kv_indptr,
work_metadata,
work_info_set,
work_indptr,
reduce_indptr,
reduce_final_map,
reduce_partial_map,
max_q_len,
fast_mode=fast_mode,
max_split_per_batch=num_kv_splits,
intra_batch_mode=intra_batch_mode,
)
self.forward_metadata = ForwardMetadata(
kv_indptr,
kv_indices,
@@ -197,6 +376,13 @@ class AiterAttnBackend(AttentionBackend):
kv_last_page_len,
max_q_len,
None,
work_metadata=work_metadata,
work_info_set=work_info_set,
work_indptr=work_indptr,
reduce_indptr=reduce_indptr,
reduce_final_map=reduce_final_map,
reduce_partial_map=reduce_partial_map,
num_kv_splits=num_kv_splits,
)
elif forward_batch.forward_mode.is_draft_extend():
@@ -209,6 +395,35 @@ class AiterAttnBackend(AttentionBackend):
self.req_to_token,
)
)
if _use_mla_ps_kernel:
max_seqlen_qo = max(forward_batch.extend_seq_lens_cpu)
(
work_metadata,
work_indptr,
work_info_set,
reduce_indptr,
reduce_final_map,
reduce_partial_map,
) = self.make_mla_decode_meta_data_buffer(max_seqlen_qo, bs)
num_kv_splits = self.max_split_per_batch
self.make_mla_meta_data(
qo_indptr,
kv_indptr,
work_metadata,
work_info_set,
work_indptr,
reduce_indptr,
reduce_final_map,
reduce_partial_map,
max_seqlen_qo,
fast_mode=fast_mode,
max_split_per_batch=num_kv_splits,
intra_batch_mode=intra_batch_mode,
)
self.forward_metadata = ForwardMetadata(
kv_indptr,
kv_indices,
@@ -217,6 +432,14 @@ class AiterAttnBackend(AttentionBackend):
self.kv_last_page_len[:bs],
max(forward_batch.extend_seq_lens_cpu),
forward_batch.seq_lens_cpu.max().item(),
work_metadata=work_metadata,
work_info_set=work_info_set,
work_indptr=work_indptr,
reduce_indptr=reduce_indptr,
reduce_final_map=reduce_final_map,
reduce_partial_map=reduce_partial_map,
num_kv_splits=num_kv_splits,
# num_kv_splits_indptr=num_kv_splits_indptr,
)
else:
self.indices_updater_prefill.update(
@@ -266,6 +489,36 @@ class AiterAttnBackend(AttentionBackend):
kv_indices,
self.req_to_token.stride(0),
)
# if self.kv_cache_dtype == fp8_dtype:
if _use_mla_ps_kernel:
max_seqlen_qo = draft_num
(
work_metadata,
work_indptr,
work_info_set,
reduce_indptr,
reduce_final_map,
reduce_partial_map,
) = self.make_mla_decode_meta_data_buffer(max_seqlen_qo, bs)
num_kv_splits = self.max_split_per_batch
self.make_mla_meta_data(
qo_indptr,
kv_indptr,
work_metadata,
work_info_set,
work_indptr,
reduce_indptr,
reduce_final_map,
reduce_partial_map,
max_seqlen_qo,
fast_mode=fast_mode,
max_split_per_batch=num_kv_splits,
intra_batch_mode=intra_batch_mode,
)
self.forward_metadata = ForwardMetadata(
kv_indptr,
kv_indices,
@@ -274,6 +527,14 @@ class AiterAttnBackend(AttentionBackend):
self.kv_last_page_len[:bs],
draft_num,
None,
work_metadata=work_metadata,
work_info_set=work_info_set,
work_indptr=work_indptr,
reduce_indptr=reduce_indptr,
reduce_final_map=reduce_final_map,
reduce_partial_map=reduce_partial_map,
num_kv_splits=num_kv_splits,
# num_kv_splits_indptr=num_kv_splits_indptr,
)
else:
self.indices_updater_prefill.update(
@@ -361,6 +622,31 @@ class AiterAttnBackend(AttentionBackend):
device=self.device,
)
# if self.use_mla and (_use_mla_ps_kernel or self.kv_cache_dtype == fp8_dtype):
if self.use_mla and _use_mla_ps_kernel:
# for persistent mla_decode_fwd
max_seqlen_qo = (
1 if self.num_draft_tokens is None else self.num_draft_tokens
)
(
self.work_metadata,
self.work_indptr,
self.work_info_set,
self.reduce_indptr,
self.reduce_final_map,
self.reduce_partial_map,
) = self.make_mla_decode_meta_data_buffer(max_seqlen_qo, max_bs)
else:
self.work_metadata = None
self.work_indptr = None
self.work_info_set = None
self.reduce_indptr = None
self.reduce_final_map = None
self.reduce_partial_map = None
def init_forward_metadata_capture_cuda_graph(
self,
bs: int,
@@ -371,6 +657,18 @@ class AiterAttnBackend(AttentionBackend):
forward_mode: ForwardMode,
spec_info: Optional[SpecInput],
):
num_kv_splits = None
# num_kv_splits_indptr = None
work_metadata = None
work_info_set = None
work_indptr = None
reduce_indptr = None
reduce_final_map = None
reduce_partial_map = None
if forward_mode.is_decode_or_idle():
qo_indptr = None
kv_last_page_len = None
@@ -401,13 +699,47 @@ class AiterAttnBackend(AttentionBackend):
kv_last_page_len = self.cuda_graph_kv_last_page_len[:bs]
max_q_len = 1
if _use_mla_ps_kernel:
num_kv_splits = self.max_split_per_batch
self.make_mla_meta_data(
qo_indptr,
kv_indptr,
self.work_metadata,
self.work_info_set,
self.work_indptr,
self.reduce_indptr,
self.reduce_final_map,
self.reduce_partial_map,
max_q_len,
fast_mode=fast_mode,
max_split_per_batch=num_kv_splits,
intra_batch_mode=intra_batch_mode,
)
work_metadata = self.work_metadata
work_info_set = self.work_info_set
work_indptr = self.work_indptr
reduce_indptr = self.reduce_indptr
reduce_final_map = self.reduce_final_map
reduce_partial_map = self.reduce_partial_map
self.forward_metadata = ForwardMetadata(
kv_indptr,
kv_indices,
qo_indptr,
kv_last_page_len,
max_q_len,
None,
kv_indptr[-1].item(),
work_metadata=work_metadata,
work_info_set=work_info_set,
work_indptr=work_indptr,
reduce_indptr=reduce_indptr,
reduce_final_map=reduce_final_map,
reduce_partial_map=reduce_partial_map,
num_kv_splits=num_kv_splits,
# num_kv_splits_indptr=num_kv_splits_indptr,
)
elif forward_mode.is_target_verify():
@@ -435,13 +767,49 @@ class AiterAttnBackend(AttentionBackend):
kv_last_page_len = self.cuda_graph_kv_last_page_len[:bs]
max_q_len = self.num_draft_tokens
# if self.kv_cache_dtype == fp8_dtype:
if _use_mla_ps_kernel:
num_kv_splits = self.max_split_per_batch
self.make_mla_meta_data(
qo_indptr,
kv_indptr,
self.work_metadata,
self.work_info_set,
self.work_indptr,
self.reduce_indptr,
self.reduce_final_map,
self.reduce_partial_map,
max_q_len,
fast_mode=fast_mode,
max_split_per_batch=num_kv_splits,
intra_batch_mode=intra_batch_mode,
)
work_metadata = self.work_metadata
work_info_set = self.work_info_set
work_indptr = self.work_indptr
reduce_indptr = self.reduce_indptr
reduce_final_map = self.reduce_final_map
reduce_partial_map = self.reduce_partial_map
self.forward_metadata = ForwardMetadata(
kv_indptr,
kv_indices,
qo_indptr,
kv_last_page_len,
max_q_len,
None,
kv_indptr[-1].item(),
work_metadata=work_metadata,
work_info_set=work_info_set,
work_indptr=work_indptr,
reduce_indptr=reduce_indptr,
reduce_final_map=reduce_final_map,
reduce_partial_map=reduce_partial_map,
num_kv_splits=num_kv_splits,
# num_kv_splits_indptr=num_kv_splits_indptr,
)
else:
seq_lens_sum = seq_lens.sum().item()
@@ -485,13 +853,49 @@ class AiterAttnBackend(AttentionBackend):
)
kv_last_page_len = self.cuda_graph_kv_last_page_len[:bs]
max_q_len = num_tokens_per_bs
if _use_mla_ps_kernel:
num_kv_splits = self.max_split_per_batch
self.make_mla_meta_data(
qo_indptr,
kv_indptr,
self.work_metadata,
self.work_info_set,
self.work_indptr,
self.reduce_indptr,
self.reduce_final_map,
self.reduce_partial_map,
max_q_len,
fast_mode=fast_mode,
max_split_per_batch=num_kv_splits,
intra_batch_mode=intra_batch_mode,
)
work_metadata = self.work_metadata
work_info_set = self.work_info_set
work_indptr = self.work_indptr
reduce_indptr = self.reduce_indptr
reduce_final_map = self.reduce_final_map
reduce_partial_map = self.reduce_partial_map
self.forward_metadata = ForwardMetadata(
kv_indptr,
kv_indices,
qo_indptr,
kv_last_page_len,
max_q_len,
None,
kv_indptr[-1].item(),
work_metadata=work_metadata,
work_info_set=work_info_set,
work_indptr=work_indptr,
reduce_indptr=reduce_indptr,
reduce_final_map=reduce_final_map,
reduce_partial_map=reduce_partial_map,
num_kv_splits=num_kv_splits,
# num_kv_splits_indptr=num_kv_splits_indptr,
)
else:
raise ValueError(f"Invalid mode: {forward_mode=}")
@@ -507,6 +911,7 @@ class AiterAttnBackend(AttentionBackend):
spec_info: Optional[SpecInput],
seq_lens_cpu: Optional[torch.Tensor],
):
if forward_mode.is_decode_or_idle():
kv_indptr = self.kv_indptr
kv_indices = self.cuda_graph_kv_indices
@@ -549,6 +954,7 @@ class AiterAttnBackend(AttentionBackend):
kv_indices,
self.req_to_token.stride(0),
)
elif forward_mode.is_draft_extend():
seq_lens = seq_lens[:bs]
accept_lens = spec_info.accept_length[:bs]
@@ -566,6 +972,7 @@ class AiterAttnBackend(AttentionBackend):
kv_indices,
self.req_to_token.stride(0),
)
else:
raise ValueError("Invalid forward mode")
@@ -619,7 +1026,8 @@ class AiterAttnBackend(AttentionBackend):
and not forward_batch.forward_mode.is_target_verify()
and not forward_batch.forward_mode.is_draft_extend()
):
if kv_indices.shape[0] == 0:
extend_no_prefix = not any(forward_batch.extend_prefix_lens_cpu)
if kv_indices.shape[0] == 0 or extend_no_prefix:
o = flash_attn_varlen_func(
q,
k,
@@ -637,6 +1045,13 @@ class AiterAttnBackend(AttentionBackend):
kvc, k_pe = torch.split(
K_Buffer, [kv_lora_rank, qk_rope_head_dim], dim=-1
)
if self.kv_cache_dtype == fp8_dtype:
dtype = q.dtype
kvc = kvc.to(dtype)
k_pe = k_pe.to(dtype)
kvprefix = layer.kv_b_proj(kvc.contiguous())[0]
kvprefix = kvprefix.view(
@@ -699,7 +1114,37 @@ class AiterAttnBackend(AttentionBackend):
K_Buffer = K_Buffer.view(-1, layer.tp_k_head_num, layer.qk_head_dim)
return o
elif forward_batch.forward_mode.is_target_verify():
o = q.new_empty((q.shape[0], layer.tp_q_head_num, layer.v_head_dim))
o = q.new_empty(
(q.shape[0], layer.tp_q_head_num, layer.v_head_dim),
dtype=self.input_dtype,
)
work_metadata = self.forward_metadata.work_metadata
work_indptr = self.forward_metadata.work_indptr
work_info_set = self.forward_metadata.work_info_set
reduce_indptr = self.forward_metadata.reduce_indptr
reduce_final_map = self.forward_metadata.reduce_final_map
reduce_partial_map = self.forward_metadata.reduce_partial_map
num_kv_splits = self.forward_metadata.num_kv_splits
if layer.layer_id == 0 and _use_mla_ps_kernel:
self.make_mla_meta_data(
self.forward_metadata.qo_indptr,
self.forward_metadata.kv_indptr,
work_metadata,
work_info_set,
work_indptr,
reduce_indptr,
reduce_final_map,
reduce_partial_map,
self.forward_metadata.max_q_len,
fast_mode=fast_mode,
max_split_per_batch=num_kv_splits,
intra_batch_mode=intra_batch_mode,
)
mla_decode_fwd(
q,
K_Buffer.view(-1, 1, 1, layer.qk_head_dim),
@@ -711,16 +1156,51 @@ class AiterAttnBackend(AttentionBackend):
self.forward_metadata.max_q_len,
layer.scaling,
layer.logit_cap,
work_meta_data=work_metadata,
work_indptr=work_indptr,
work_info_set=work_info_set,
reduce_indptr=reduce_indptr,
reduce_final_map=reduce_final_map,
reduce_partial_map=reduce_partial_map,
q_scale=layer.k_scale,
kv_scale=layer.k_scale,
intra_batch_mode=intra_batch_mode,
num_kv_splits=num_kv_splits,
)
K_Buffer = K_Buffer.view(-1, 1, layer.qk_head_dim)
return o
elif forward_batch.forward_mode.is_draft_extend():
o = q.new_empty((q.shape[0], layer.tp_q_head_num, layer.v_head_dim))
causal = True
sliding_window_size = -1
kv_indptr = self.forward_metadata.kv_indptr
kv_indices = self.forward_metadata.kv_indices
mla_prefill_fwd(
o = q.new_empty(
(q.shape[0], layer.tp_q_head_num, layer.v_head_dim),
dtype=self.input_dtype,
)
work_metadata = self.forward_metadata.work_metadata
work_indptr = self.forward_metadata.work_indptr
work_info_set = self.forward_metadata.work_info_set
reduce_indptr = self.forward_metadata.reduce_indptr
reduce_final_map = self.forward_metadata.reduce_final_map
reduce_partial_map = self.forward_metadata.reduce_partial_map
num_kv_splits = self.forward_metadata.num_kv_splits
if layer.layer_id == 0 and _use_mla_ps_kernel:
self.make_mla_meta_data(
self.forward_metadata.qo_indptr,
self.forward_metadata.kv_indptr,
work_metadata,
work_info_set,
work_indptr,
reduce_indptr,
reduce_final_map,
reduce_partial_map,
self.forward_metadata.max_q_len,
fast_mode=fast_mode,
max_split_per_batch=num_kv_splits,
intra_batch_mode=intra_batch_mode,
)
mla_decode_fwd(
q,
K_Buffer.view(-1, 1, 1, layer.qk_head_dim),
o,
@@ -731,28 +1211,18 @@ class AiterAttnBackend(AttentionBackend):
self.forward_metadata.max_q_len,
layer.scaling,
layer.logit_cap,
work_meta_data=work_metadata,
work_indptr=work_indptr,
work_info_set=work_info_set,
reduce_indptr=reduce_indptr,
reduce_final_map=reduce_final_map,
reduce_partial_map=reduce_partial_map,
q_scale=layer.k_scale,
kv_scale=layer.k_scale,
intra_batch_mode=intra_batch_mode,
num_kv_splits=num_kv_splits,
)
K_Buffer = K_Buffer.view(-1, 1, layer.qk_head_dim)
return o
# self.extend_attention_fwd(
# q.view(-1, layer.tp_q_head_num, layer.qk_head_dim),
# k.contiguous(),
# v.contiguous(),
# o.view(-1, layer.tp_q_head_num, layer.v_head_dim),
# forward_batch.token_to_kv_pool.get_key_buffer(layer.layer_id),
# forward_batch.token_to_kv_pool.get_value_buffer(layer.layer_id),
# self.forward_metadata.qo_indptr,
# kv_indptr,
# kv_indices,
# None,
# causal,
# None,
# self.forward_metadata.max_q_len,
# layer.scaling,
# layer.logit_cap,
# sliding_window_size,
# )
# return o
else:
raise ValueError(
f"Invalid forward mode for MLA prefill: {forward_batch.forward_mode=}"
@@ -764,6 +1234,12 @@ class AiterAttnBackend(AttentionBackend):
bs0 = forward_batch.batch_size + 1
# 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)
o = mha_batch_prefill_func(
q.contiguous().view(-1, layer.tp_q_head_num, layer.head_dim),
k_cache,
@@ -795,9 +1271,12 @@ class AiterAttnBackend(AttentionBackend):
q = q.reshape(-1, layer.tp_q_head_num * layer.qk_head_dim)
if layer.qk_head_dim != layer.v_head_dim:
o = q.new_empty((q.shape[0], layer.tp_q_head_num * layer.v_head_dim))
o = q.new_empty(
(q.shape[0], layer.tp_q_head_num * layer.v_head_dim),
dtype=self.input_dtype,
)
else:
o = torch.empty_like(q)
o = torch.empty_like(q, dtype=self.input_dtype)
if save_kv_cache:
forward_batch.token_to_kv_pool.set_kv_buffer(
@@ -806,6 +1285,33 @@ class AiterAttnBackend(AttentionBackend):
if self.use_mla:
k_buffer = forward_batch.token_to_kv_pool.get_key_buffer(layer.layer_id)
work_metadata = self.forward_metadata.work_metadata
work_indptr = self.forward_metadata.work_indptr
work_info_set = self.forward_metadata.work_info_set
reduce_indptr = self.forward_metadata.reduce_indptr
reduce_final_map = self.forward_metadata.reduce_final_map
reduce_partial_map = self.forward_metadata.reduce_partial_map
num_kv_splits = self.forward_metadata.num_kv_splits
if layer.layer_id == 0 and _use_mla_ps_kernel:
self.make_mla_meta_data(
self.forward_metadata.qo_indptr,
self.forward_metadata.kv_indptr,
work_metadata,
work_info_set,
work_indptr,
reduce_indptr,
reduce_final_map,
reduce_partial_map,
self.forward_metadata.max_q_len,
fast_mode=fast_mode,
max_split_per_batch=num_kv_splits,
intra_batch_mode=intra_batch_mode,
)
mla_decode_fwd(
q.view(-1, layer.tp_q_head_num, layer.qk_head_dim),
k_buffer.view(-1, 1, 1, layer.qk_head_dim),
@@ -817,20 +1323,37 @@ class AiterAttnBackend(AttentionBackend):
self.forward_metadata.max_q_len,
layer.scaling,
layer.logit_cap,
work_meta_data=work_metadata,
work_indptr=work_indptr,
work_info_set=work_info_set,
reduce_indptr=reduce_indptr,
reduce_final_map=reduce_final_map,
reduce_partial_map=reduce_partial_map,
q_scale=layer.k_scale,
kv_scale=layer.k_scale,
intra_batch_mode=intra_batch_mode,
num_kv_splits=num_kv_splits,
)
k_buffer = k_buffer.view(-1, 1, layer.qk_head_dim)
else:
self.logits_soft_cap = layer.logit_cap
k_cache, v_cache = forward_batch.token_to_kv_pool.get_kv_buffer(
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)
paged_attention_ragged(
o.view(-1, layer.tp_q_head_num, layer.qk_head_dim),
self.workspace_buffer,
q.view(-1, layer.tp_q_head_num, layer.qk_head_dim),
forward_batch.token_to_kv_pool.get_key_buffer(layer.layer_id).view(
-1, 1, layer.tp_k_head_num, layer.qk_head_dim
),
forward_batch.token_to_kv_pool.get_value_buffer(layer.layer_id).view(
-1, 1, layer.tp_v_head_num, layer.v_head_dim
),
k_cache.view(-1, 1, layer.tp_k_head_num, layer.qk_head_dim),
v_cache.view(-1, 1, layer.tp_v_head_num, layer.v_head_dim),
self.scale,
self.forward_metadata.kv_indptr,
self.forward_metadata.kv_indices,

View File

@@ -175,6 +175,7 @@ class Fp8Config(QuantizationConfig):
) -> Optional[QuantizeMethodBase]:
from sglang.srt.layers.linear import LinearBase
from sglang.srt.layers.moe.fused_moe_triton import FusedMoE
from sglang.srt.layers.radix_attention import RadixAttention
if isinstance(layer, LinearBase):
if is_layer_skipped(prefix, self.ignored_layers):
@@ -182,6 +183,8 @@ class Fp8Config(QuantizationConfig):
return Fp8LinearMethod(self)
elif isinstance(layer, FusedMoE):
return Fp8MoEMethod(self)
elif isinstance(layer, RadixAttention):
return Fp8KVCacheMethod(self)
return None
def get_scaled_act_names(self) -> List[str]:

View File

@@ -71,6 +71,8 @@ class QuarkConfig(QuantizationConfig):
):
if isinstance(layer, LinearBase):
return UnquantizedLinearMethod()
elif isinstance(layer, RadixAttention):
return QuarkKVCacheMethod(self)
return None
if isinstance(layer, LinearBase):

View File

@@ -1,11 +1,12 @@
import torch
from aiter.ops.triton.fused_kv_cache import fused_qk_rope_cat_and_cache_mla
from aiter.ops.triton.fused_qk_concat import fused_qk_rope_cat
from aiter.ops.triton.gemm_a16w16 import gemm_a16w16
from aiter.ops.triton.gemm_a16w16_atomic import gemm_a16w16_atomic
from sglang.srt.utils import BumpAllocator
__all__ = ["fused_qk_rope_cat"]
__all__ = ["fused_qk_rope_cat", "fused_qk_rope_cat_and_cache_mla"]
def aiter_dsv3_router_gemm(

View File

@@ -95,6 +95,7 @@ from sglang.srt.layers.dp_attention import (
)
from sglang.srt.layers.logits_processor import LogitsProcessorOutput
from sglang.srt.layers.pooler import EmbeddingPoolerOutput
from sglang.srt.layers.quantization.fp8_kernel import fp8_dtype
from sglang.srt.layers.sampler import Sampler
from sglang.srt.layers.torchao_utils import apply_torchao_config_to_model
from sglang.srt.lora.lora_manager import LoRAManager
@@ -1629,19 +1630,19 @@ class ModelRunner:
and kv_cache_quant_algo.upper() == "FP8"
):
if _is_hip:
self.kv_cache_dtype = torch.float8_e4m3fnuz
self.kv_cache_dtype = fp8_dtype
else:
self.kv_cache_dtype = torch.float8_e4m3fn
else:
self.kv_cache_dtype = self.dtype
elif self.server_args.kv_cache_dtype == "fp8_e5m2":
if _is_hip: # Using natively supported format
self.kv_cache_dtype = torch.float8_e5m2fnuz
self.kv_cache_dtype = fp8_dtype
else:
self.kv_cache_dtype = torch.float8_e5m2
elif self.server_args.kv_cache_dtype == "fp8_e4m3":
if _is_hip: # Using natively supported format
self.kv_cache_dtype = torch.float8_e4m3fnuz
self.kv_cache_dtype = fp8_dtype
else:
self.kv_cache_dtype = torch.float8_e4m3fn
elif self.server_args.kv_cache_dtype in ("bf16", "bfloat16"):

View File

@@ -105,6 +105,7 @@ from sglang.srt.layers.moe.utils import RoutingMethodType
from sglang.srt.layers.quantization.base_config import QuantizationConfig
from sglang.srt.layers.quantization.fp8 import Fp8Config
from sglang.srt.layers.quantization.fp8_kernel import (
fp8_dtype,
is_fp8_fnuz,
per_tensor_quant_mla_fp8,
per_token_group_quant_mla_deep_gemm_masked_fp8,
@@ -188,7 +189,7 @@ if _use_aiter_gfx95:
)
from sglang.srt.layers.rocm_linear_utils import (
aiter_dsv3_router_gemm,
fused_qk_rope_cat,
fused_qk_rope_cat_and_cache_mla,
get_dsv3_gemm_output_zero_allocator_size,
)
@@ -2006,6 +2007,8 @@ class DeepseekV2AttentionMLA(nn.Module):
topk_indices,
llama_4_scaling,
):
save_kv_cache = True
if self.current_attention_backend in FORWARD_ABSORB_CORE_ATTENTION_BACKENDS:
extra_args = {}
if self._fuse_rope_for_trtllm_mla(forward_batch):
@@ -2029,16 +2032,29 @@ class DeepseekV2AttentionMLA(nn.Module):
if _use_aiter_gfx95:
cos = self.rotary_emb.cos_cache
sin = self.rotary_emb.sin_cache
q, k = fused_qk_rope_cat(
kv_cache_dtype = (
fp8_dtype if self.kv_cache_dtype == "fp8_e4m3" else q_nope_out.dtype
)
q, _, _, k = fused_qk_rope_cat_and_cache_mla(
q_nope_out,
q_pe,
k_nope,
k_pe,
forward_batch.token_to_kv_pool.get_key_buffer(
self.attn_mqa.layer_id
),
forward_batch.out_cache_loc,
positions,
cos,
sin,
self.attn_mqa.k_scale,
self.rotary_emb.is_neox_style,
q_out_dtype=kv_cache_dtype,
)
save_kv_cache = False
else:
q = torch.cat([q_nope_out, q_pe], dim=-1)
k = torch.cat([k_nope, k_pe], dim=-1)
@@ -2052,6 +2068,7 @@ class DeepseekV2AttentionMLA(nn.Module):
k,
k_nope,
forward_batch,
save_kv_cache=save_kv_cache,
**(dict(topk_indices=topk_indices) if topk_indices is not None else {}),
)
attn_output = attn_output.view(-1, self.num_local_heads, self.kv_lora_rank)