Aiter fp8 kv cache (#13147)

This commit is contained in:
kk
2025-12-08 16:39:53 -08:00
committed by GitHub
parent 119fd956fb
commit c106b54b57
7 changed files with 594 additions and 96 deletions
-49
View File
@@ -1,49 +0,0 @@
.PHONY: check-deps install-deps format update help
# Show help for each target
help:
@echo "Available targets:"
@grep -E '^[a-zA-Z0-9_-]+:.*?## .*$$' $(MAKEFILE_LIST) | sort | awk 'BEGIN {FS = ":.*?## "}; {printf "\033[36m%-20s\033[0m %s\n", $$1, $$2}'
check-deps: ## Check and install required Python formatting dependencies
@command -v isort >/dev/null 2>&1 || (echo "Installing isort..." && pip install isort)
@command -v black >/dev/null 2>&1 || (echo "Installing black..." && pip install black)
install-deps: ## Install Python formatting tools (isort and black)
pip install isort black
format: check-deps ## Format modified Python files using isort and black
@echo "Formatting modified Python files..."
git diff --name-only --diff-filter=M | grep '\.py$$' | xargs -I {} sh -c 'isort {} && black {}'
FILES_TO_UPDATE = docker/rocm.Dockerfile \
python/pyproject.toml \
python/pyproject_other.toml \
python/sglang/version.py \
docs/developer_guide/setup_github_runner.md \
docs/get_started/install.md \
docs/platforms/amd_gpu.md \
docs/platforms/ascend_npu.md \
docs/platforms/cpu_server.md \
docs/platforms/xpu.md \
benchmark/deepseek_v3/README.md
update: ## Update version numbers across project files. Usage: make update <new_version>
@if [ -z "$(filter-out $@,$(MAKECMDGOALS))" ]; then \
echo "Version required. Usage: make update <new_version>"; \
exit 1; \
fi
@OLD_VERSION=$$(grep "version" python/sglang/version.py | cut -d '"' -f2); \
NEW_VERSION=$(filter-out $@,$(MAKECMDGOALS)); \
echo "Updating version from $$OLD_VERSION to $$NEW_VERSION"; \
for file in $(FILES_TO_UPDATE); do \
if [ "$(shell uname)" = "Darwin" ]; then \
sed -i '' -e "s/$$OLD_VERSION/$$NEW_VERSION/g" $$file; \
else \
sed -i -e "s/$$OLD_VERSION/$$NEW_VERSION/g" $$file; \
fi \
done; \
echo "Version update complete"
%:
@:
@@ -4,6 +4,7 @@ from __future__ import annotations
end to end attention solution with aiter kernels end to end attention solution with aiter kernels
""" """
import logging
from dataclasses import dataclass from dataclasses import dataclass
from enum import Enum, auto from enum import Enum, auto
from typing import TYPE_CHECKING, Optional from typing import TYPE_CHECKING, Optional
@@ -27,6 +28,8 @@ if TYPE_CHECKING:
try: try:
from aiter import ( from aiter import (
flash_attn_varlen_func, flash_attn_varlen_func,
get_mla_metadata_info_v1,
get_mla_metadata_v1,
mha_batch_prefill_func, mha_batch_prefill_func,
paged_attention_ragged, paged_attention_ragged,
) )
@@ -37,6 +40,21 @@ except ImportError:
) )
from sglang.srt.configs.model_config import AttentionArch 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): class WrapperDispatch(Enum):
@@ -52,6 +70,14 @@ class ForwardMetadata:
kv_last_page_len: torch.Tensor kv_last_page_len: torch.Tensor
max_q_len: int max_q_len: int
max_kv_len: Optional[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 global_workspace_buffer = None
@@ -72,6 +98,10 @@ class AiterAttnBackend(AttentionBackend):
extend_attention_fwd, 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.extend_attention_fwd = torch.compiler.disable(extend_attention_fwd)
self.device = model_runner.device self.device = model_runner.device
@@ -154,6 +184,118 @@ class AiterAttnBackend(AttentionBackend):
self.enable_dp_attention = is_dp_attention_enabled() 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): def init_forward_metadata(self, forward_batch: ForwardBatch):
"""Init auxiliary variables for triton attention backend.""" """Init auxiliary variables for triton attention backend."""
@@ -164,6 +306,16 @@ class AiterAttnBackend(AttentionBackend):
kv_last_page_len = None kv_last_page_len = None
max_q_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 forward_batch.forward_mode.is_decode_or_idle():
if spec_info is None: if spec_info is None:
kv_indptr[1 : bs + 1] = torch.cumsum(forward_batch.seq_lens, dim=0) 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] kv_last_page_len = self.kv_last_page_len[:bs]
max_q_len = 1 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( self.forward_metadata = ForwardMetadata(
kv_indptr, kv_indptr,
kv_indices, kv_indices,
@@ -197,6 +376,13 @@ class AiterAttnBackend(AttentionBackend):
kv_last_page_len, kv_last_page_len,
max_q_len, max_q_len,
None, 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(): elif forward_batch.forward_mode.is_draft_extend():
@@ -209,6 +395,35 @@ class AiterAttnBackend(AttentionBackend):
self.req_to_token, 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( self.forward_metadata = ForwardMetadata(
kv_indptr, kv_indptr,
kv_indices, kv_indices,
@@ -217,6 +432,14 @@ class AiterAttnBackend(AttentionBackend):
self.kv_last_page_len[:bs], self.kv_last_page_len[:bs],
max(forward_batch.extend_seq_lens_cpu), max(forward_batch.extend_seq_lens_cpu),
forward_batch.seq_lens_cpu.max().item(), 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: else:
self.indices_updater_prefill.update( self.indices_updater_prefill.update(
@@ -266,6 +489,36 @@ class AiterAttnBackend(AttentionBackend):
kv_indices, kv_indices,
self.req_to_token.stride(0), 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( self.forward_metadata = ForwardMetadata(
kv_indptr, kv_indptr,
kv_indices, kv_indices,
@@ -274,6 +527,14 @@ class AiterAttnBackend(AttentionBackend):
self.kv_last_page_len[:bs], self.kv_last_page_len[:bs],
draft_num, draft_num,
None, 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: else:
self.indices_updater_prefill.update( self.indices_updater_prefill.update(
@@ -361,6 +622,31 @@ class AiterAttnBackend(AttentionBackend):
device=self.device, 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( def init_forward_metadata_capture_cuda_graph(
self, self,
bs: int, bs: int,
@@ -371,6 +657,18 @@ class AiterAttnBackend(AttentionBackend):
forward_mode: ForwardMode, forward_mode: ForwardMode,
spec_info: Optional[SpecInput], 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(): if forward_mode.is_decode_or_idle():
qo_indptr = None qo_indptr = None
kv_last_page_len = 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] kv_last_page_len = self.cuda_graph_kv_last_page_len[:bs]
max_q_len = 1 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( self.forward_metadata = ForwardMetadata(
kv_indptr, kv_indptr,
kv_indices, kv_indices,
qo_indptr, qo_indptr,
kv_last_page_len, kv_last_page_len,
max_q_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(): 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] kv_last_page_len = self.cuda_graph_kv_last_page_len[:bs]
max_q_len = self.num_draft_tokens 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( self.forward_metadata = ForwardMetadata(
kv_indptr, kv_indptr,
kv_indices, kv_indices,
qo_indptr, qo_indptr,
kv_last_page_len, kv_last_page_len,
max_q_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: else:
seq_lens_sum = seq_lens.sum().item() 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] kv_last_page_len = self.cuda_graph_kv_last_page_len[:bs]
max_q_len = num_tokens_per_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( self.forward_metadata = ForwardMetadata(
kv_indptr, kv_indptr,
kv_indices, kv_indices,
qo_indptr, qo_indptr,
kv_last_page_len, kv_last_page_len,
max_q_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: else:
raise ValueError(f"Invalid mode: {forward_mode=}") raise ValueError(f"Invalid mode: {forward_mode=}")
@@ -507,6 +911,7 @@ class AiterAttnBackend(AttentionBackend):
spec_info: Optional[SpecInput], spec_info: Optional[SpecInput],
seq_lens_cpu: Optional[torch.Tensor], seq_lens_cpu: Optional[torch.Tensor],
): ):
if forward_mode.is_decode_or_idle(): if forward_mode.is_decode_or_idle():
kv_indptr = self.kv_indptr kv_indptr = self.kv_indptr
kv_indices = self.cuda_graph_kv_indices kv_indices = self.cuda_graph_kv_indices
@@ -549,6 +954,7 @@ class AiterAttnBackend(AttentionBackend):
kv_indices, kv_indices,
self.req_to_token.stride(0), self.req_to_token.stride(0),
) )
elif forward_mode.is_draft_extend(): elif forward_mode.is_draft_extend():
seq_lens = seq_lens[:bs] seq_lens = seq_lens[:bs]
accept_lens = spec_info.accept_length[:bs] accept_lens = spec_info.accept_length[:bs]
@@ -566,6 +972,7 @@ class AiterAttnBackend(AttentionBackend):
kv_indices, kv_indices,
self.req_to_token.stride(0), self.req_to_token.stride(0),
) )
else: else:
raise ValueError("Invalid forward mode") 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_target_verify()
and not forward_batch.forward_mode.is_draft_extend() 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( o = flash_attn_varlen_func(
q, q,
k, k,
@@ -637,6 +1045,13 @@ class AiterAttnBackend(AttentionBackend):
kvc, k_pe = torch.split( kvc, k_pe = torch.split(
K_Buffer, [kv_lora_rank, qk_rope_head_dim], dim=-1 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 = layer.kv_b_proj(kvc.contiguous())[0]
kvprefix = kvprefix.view( 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) K_Buffer = K_Buffer.view(-1, layer.tp_k_head_num, layer.qk_head_dim)
return o return o
elif forward_batch.forward_mode.is_target_verify(): 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( mla_decode_fwd(
q, q,
K_Buffer.view(-1, 1, 1, layer.qk_head_dim), K_Buffer.view(-1, 1, 1, layer.qk_head_dim),
@@ -711,16 +1156,51 @@ class AiterAttnBackend(AttentionBackend):
self.forward_metadata.max_q_len, self.forward_metadata.max_q_len,
layer.scaling, layer.scaling,
layer.logit_cap, 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 return o
elif forward_batch.forward_mode.is_draft_extend(): elif forward_batch.forward_mode.is_draft_extend():
o = q.new_empty((q.shape[0], layer.tp_q_head_num, layer.v_head_dim)) o = q.new_empty(
causal = True (q.shape[0], layer.tp_q_head_num, layer.v_head_dim),
sliding_window_size = -1 dtype=self.input_dtype,
kv_indptr = self.forward_metadata.kv_indptr )
kv_indices = self.forward_metadata.kv_indices
mla_prefill_fwd( 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, q,
K_Buffer.view(-1, 1, 1, layer.qk_head_dim), K_Buffer.view(-1, 1, 1, layer.qk_head_dim),
o, o,
@@ -731,28 +1211,18 @@ class AiterAttnBackend(AttentionBackend):
self.forward_metadata.max_q_len, self.forward_metadata.max_q_len,
layer.scaling, layer.scaling,
layer.logit_cap, 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 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: else:
raise ValueError( raise ValueError(
f"Invalid forward mode for MLA prefill: {forward_batch.forward_mode=}" f"Invalid forward mode for MLA prefill: {forward_batch.forward_mode=}"
@@ -764,6 +1234,12 @@ class AiterAttnBackend(AttentionBackend):
bs0 = forward_batch.batch_size + 1 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( o = mha_batch_prefill_func(
q.contiguous().view(-1, layer.tp_q_head_num, layer.head_dim), q.contiguous().view(-1, layer.tp_q_head_num, layer.head_dim),
k_cache, k_cache,
@@ -795,9 +1271,12 @@ class AiterAttnBackend(AttentionBackend):
q = q.reshape(-1, layer.tp_q_head_num * layer.qk_head_dim) q = q.reshape(-1, layer.tp_q_head_num * layer.qk_head_dim)
if layer.qk_head_dim != layer.v_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: else:
o = torch.empty_like(q) o = torch.empty_like(q, dtype=self.input_dtype)
if save_kv_cache: if save_kv_cache:
forward_batch.token_to_kv_pool.set_kv_buffer( forward_batch.token_to_kv_pool.set_kv_buffer(
@@ -806,6 +1285,33 @@ class AiterAttnBackend(AttentionBackend):
if self.use_mla: if self.use_mla:
k_buffer = forward_batch.token_to_kv_pool.get_key_buffer(layer.layer_id) 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( mla_decode_fwd(
q.view(-1, layer.tp_q_head_num, layer.qk_head_dim), q.view(-1, layer.tp_q_head_num, layer.qk_head_dim),
k_buffer.view(-1, 1, 1, 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, self.forward_metadata.max_q_len,
layer.scaling, layer.scaling,
layer.logit_cap, 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: else:
self.logits_soft_cap = layer.logit_cap 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( 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,
q.view(-1, layer.tp_q_head_num, layer.qk_head_dim), 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( k_cache.view(-1, 1, layer.tp_k_head_num, layer.qk_head_dim),
-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),
),
forward_batch.token_to_kv_pool.get_value_buffer(layer.layer_id).view(
-1, 1, layer.tp_v_head_num, layer.v_head_dim
),
self.scale, self.scale,
self.forward_metadata.kv_indptr, self.forward_metadata.kv_indptr,
self.forward_metadata.kv_indices, self.forward_metadata.kv_indices,
@@ -175,6 +175,7 @@ class Fp8Config(QuantizationConfig):
) -> Optional[QuantizeMethodBase]: ) -> Optional[QuantizeMethodBase]:
from sglang.srt.layers.linear import LinearBase from sglang.srt.layers.linear import LinearBase
from sglang.srt.layers.moe.fused_moe_triton import FusedMoE from sglang.srt.layers.moe.fused_moe_triton import FusedMoE
from sglang.srt.layers.radix_attention import RadixAttention
if isinstance(layer, LinearBase): if isinstance(layer, LinearBase):
if is_layer_skipped(prefix, self.ignored_layers): if is_layer_skipped(prefix, self.ignored_layers):
@@ -182,6 +183,8 @@ class Fp8Config(QuantizationConfig):
return Fp8LinearMethod(self) return Fp8LinearMethod(self)
elif isinstance(layer, FusedMoE): elif isinstance(layer, FusedMoE):
return Fp8MoEMethod(self) return Fp8MoEMethod(self)
elif isinstance(layer, RadixAttention):
return Fp8KVCacheMethod(self)
return None return None
def get_scaled_act_names(self) -> List[str]: def get_scaled_act_names(self) -> List[str]:
@@ -71,6 +71,8 @@ class QuarkConfig(QuantizationConfig):
): ):
if isinstance(layer, LinearBase): if isinstance(layer, LinearBase):
return UnquantizedLinearMethod() return UnquantizedLinearMethod()
elif isinstance(layer, RadixAttention):
return QuarkKVCacheMethod(self)
return None return None
if isinstance(layer, LinearBase): if isinstance(layer, LinearBase):
@@ -1,11 +1,12 @@
import torch 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.fused_qk_concat import fused_qk_rope_cat
from aiter.ops.triton.gemm_a16w16 import gemm_a16w16 from aiter.ops.triton.gemm_a16w16 import gemm_a16w16
from aiter.ops.triton.gemm_a16w16_atomic import gemm_a16w16_atomic from aiter.ops.triton.gemm_a16w16_atomic import gemm_a16w16_atomic
from sglang.srt.utils import BumpAllocator 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( 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.logits_processor import LogitsProcessorOutput
from sglang.srt.layers.pooler import EmbeddingPoolerOutput 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.sampler import Sampler
from sglang.srt.layers.torchao_utils import apply_torchao_config_to_model from sglang.srt.layers.torchao_utils import apply_torchao_config_to_model
from sglang.srt.lora.lora_manager import LoRAManager from sglang.srt.lora.lora_manager import LoRAManager
@@ -1629,19 +1630,19 @@ class ModelRunner:
and kv_cache_quant_algo.upper() == "FP8" and kv_cache_quant_algo.upper() == "FP8"
): ):
if _is_hip: if _is_hip:
self.kv_cache_dtype = torch.float8_e4m3fnuz self.kv_cache_dtype = fp8_dtype
else: else:
self.kv_cache_dtype = torch.float8_e4m3fn self.kv_cache_dtype = torch.float8_e4m3fn
else: else:
self.kv_cache_dtype = self.dtype self.kv_cache_dtype = self.dtype
elif self.server_args.kv_cache_dtype == "fp8_e5m2": elif self.server_args.kv_cache_dtype == "fp8_e5m2":
if _is_hip: # Using natively supported format if _is_hip: # Using natively supported format
self.kv_cache_dtype = torch.float8_e5m2fnuz self.kv_cache_dtype = fp8_dtype
else: else:
self.kv_cache_dtype = torch.float8_e5m2 self.kv_cache_dtype = torch.float8_e5m2
elif self.server_args.kv_cache_dtype == "fp8_e4m3": elif self.server_args.kv_cache_dtype == "fp8_e4m3":
if _is_hip: # Using natively supported format if _is_hip: # Using natively supported format
self.kv_cache_dtype = torch.float8_e4m3fnuz self.kv_cache_dtype = fp8_dtype
else: else:
self.kv_cache_dtype = torch.float8_e4m3fn self.kv_cache_dtype = torch.float8_e4m3fn
elif self.server_args.kv_cache_dtype in ("bf16", "bfloat16"): elif self.server_args.kv_cache_dtype in ("bf16", "bfloat16"):
+19 -2
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.base_config import QuantizationConfig
from sglang.srt.layers.quantization.fp8 import Fp8Config from sglang.srt.layers.quantization.fp8 import Fp8Config
from sglang.srt.layers.quantization.fp8_kernel import ( from sglang.srt.layers.quantization.fp8_kernel import (
fp8_dtype,
is_fp8_fnuz, is_fp8_fnuz,
per_tensor_quant_mla_fp8, per_tensor_quant_mla_fp8,
per_token_group_quant_mla_deep_gemm_masked_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 ( from sglang.srt.layers.rocm_linear_utils import (
aiter_dsv3_router_gemm, aiter_dsv3_router_gemm,
fused_qk_rope_cat, fused_qk_rope_cat_and_cache_mla,
get_dsv3_gemm_output_zero_allocator_size, get_dsv3_gemm_output_zero_allocator_size,
) )
@@ -2006,6 +2007,8 @@ class DeepseekV2AttentionMLA(nn.Module):
topk_indices, topk_indices,
llama_4_scaling, llama_4_scaling,
): ):
save_kv_cache = True
if self.current_attention_backend in FORWARD_ABSORB_CORE_ATTENTION_BACKENDS: if self.current_attention_backend in FORWARD_ABSORB_CORE_ATTENTION_BACKENDS:
extra_args = {} extra_args = {}
if self._fuse_rope_for_trtllm_mla(forward_batch): if self._fuse_rope_for_trtllm_mla(forward_batch):
@@ -2029,16 +2032,29 @@ class DeepseekV2AttentionMLA(nn.Module):
if _use_aiter_gfx95: if _use_aiter_gfx95:
cos = self.rotary_emb.cos_cache cos = self.rotary_emb.cos_cache
sin = self.rotary_emb.sin_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_nope_out,
q_pe, q_pe,
k_nope, k_nope,
k_pe, k_pe,
forward_batch.token_to_kv_pool.get_key_buffer(
self.attn_mqa.layer_id
),
forward_batch.out_cache_loc,
positions, positions,
cos, cos,
sin, sin,
self.attn_mqa.k_scale,
self.rotary_emb.is_neox_style, self.rotary_emb.is_neox_style,
q_out_dtype=kv_cache_dtype,
) )
save_kv_cache = False
else: else:
q = torch.cat([q_nope_out, q_pe], dim=-1) q = torch.cat([q_nope_out, q_pe], dim=-1)
k = torch.cat([k_nope, k_pe], dim=-1) k = torch.cat([k_nope, k_pe], dim=-1)
@@ -2052,6 +2068,7 @@ class DeepseekV2AttentionMLA(nn.Module):
k, k,
k_nope, k_nope,
forward_batch, forward_batch,
save_kv_cache=save_kv_cache,
**(dict(topk_indices=topk_indices) if topk_indices is not None else {}), **(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) attn_output = attn_output.view(-1, self.num_local_heads, self.kv_lora_rank)