Aiter fp8 kv cache (#13147)
This commit is contained in:
@@ -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,
|
||||
|
||||
@@ -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]:
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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"):
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user