diff --git a/Makefile b/Makefile deleted file mode 100644 index d6ef19420..000000000 --- a/Makefile +++ /dev/null @@ -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 - @if [ -z "$(filter-out $@,$(MAKECMDGOALS))" ]; then \ - echo "Version required. Usage: make update "; \ - 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" - -%: - @: diff --git a/python/sglang/srt/layers/attention/aiter_backend.py b/python/sglang/srt/layers/attention/aiter_backend.py index dafe5ee19..d28f8e4f9 100644 --- a/python/sglang/srt/layers/attention/aiter_backend.py +++ b/python/sglang/srt/layers/attention/aiter_backend.py @@ -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, diff --git a/python/sglang/srt/layers/quantization/fp8.py b/python/sglang/srt/layers/quantization/fp8.py index ad6d81cb8..a56c9dc06 100644 --- a/python/sglang/srt/layers/quantization/fp8.py +++ b/python/sglang/srt/layers/quantization/fp8.py @@ -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]: diff --git a/python/sglang/srt/layers/quantization/quark/quark.py b/python/sglang/srt/layers/quantization/quark/quark.py index 37500e687..783f8ea4b 100644 --- a/python/sglang/srt/layers/quantization/quark/quark.py +++ b/python/sglang/srt/layers/quantization/quark/quark.py @@ -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): diff --git a/python/sglang/srt/layers/rocm_linear_utils.py b/python/sglang/srt/layers/rocm_linear_utils.py index ee7dd1f59..6c8a6a367 100644 --- a/python/sglang/srt/layers/rocm_linear_utils.py +++ b/python/sglang/srt/layers/rocm_linear_utils.py @@ -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( diff --git a/python/sglang/srt/model_executor/model_runner.py b/python/sglang/srt/model_executor/model_runner.py index 4d58278b7..ac5e363ef 100644 --- a/python/sglang/srt/model_executor/model_runner.py +++ b/python/sglang/srt/model_executor/model_runner.py @@ -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"): diff --git a/python/sglang/srt/models/deepseek_v2.py b/python/sglang/srt/models/deepseek_v2.py index 9fdfbcb6a..ea8ec92a5 100644 --- a/python/sglang/srt/models/deepseek_v2.py +++ b/python/sglang/srt/models/deepseek_v2.py @@ -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)