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