From 866793c502b76699109df739345454ecd79a0d45 Mon Sep 17 00:00:00 2001 From: kk <43161300+kkHuang-amd@users.noreply.github.com> Date: Tue, 19 May 2026 00:15:07 +0800 Subject: [PATCH] Amd/deepseek v4 rebase main 0509 (#24933) Co-authored-by: root Co-authored-by: wunhuang Co-authored-by: Thomas Wang <1am9trash@gmail.com> Co-authored-by: Xinyi Song <86638975+RolaoDenthu@users.noreply.github.com> Co-authored-by: HaiShaw Co-authored-by: amd-danli103 Co-authored-by: Lin, Soga Co-authored-by: Raiden-Makoto Co-authored-by: Hubert Lu <55214931+hubertlu-tw@users.noreply.github.com> Co-authored-by: yichiche@amd.com Co-authored-by: yctseng0211 Co-authored-by: Bingxu Chen --- python/sglang/jit_kernel/deepseek_v4.py | 26 + python/sglang/srt/environ.py | 7 + .../layers/attention/attention_registry.py | 21 +- .../deepseek_v4_backend_hip_radix.py | 1265 +++++++++++++++++ .../srt/layers/attention/dsv4/compress_hip.py | 455 ++++++ .../srt/layers/attention/dsv4/compressor.py | 18 +- .../srt/layers/attention/dsv4/indexer.py | 19 +- .../srt/layers/attention/hip_flash_mla.py | 197 +++ .../attention/nsa/index_buf_accessor.py | 4 - .../layers/attention/nsa/tilelang_kernel.py | 1216 +++++++++++++++- python/sglang/srt/layers/deepseek_v4_rope.py | 168 +++ python/sglang/srt/layers/moe/topk.py | 4 +- python/sglang/srt/layers/quantization/fp8.py | 159 ++- .../mem_cache/deepseek_v4_compress_state.py | 107 +- .../srt/mem_cache/deepseek_v4_memory_pool.py | 17 +- python/sglang/srt/models/deepseek_v2.py | 5 + python/sglang/srt/models/deepseek_v4.py | 58 +- 17 files changed, 3677 insertions(+), 69 deletions(-) create mode 100644 python/sglang/srt/layers/attention/deepseek_v4_backend_hip_radix.py create mode 100644 python/sglang/srt/layers/attention/dsv4/compress_hip.py create mode 100644 python/sglang/srt/layers/attention/hip_flash_mla.py diff --git a/python/sglang/jit_kernel/deepseek_v4.py b/python/sglang/jit_kernel/deepseek_v4.py index 6f9d772ad..5ff07e88a 100644 --- a/python/sglang/jit_kernel/deepseek_v4.py +++ b/python/sglang/jit_kernel/deepseek_v4.py @@ -13,6 +13,13 @@ from sglang.jit_kernel.utils import ( make_cpp_args, ) from sglang.srt.environ import envs +from sglang.srt.utils import get_bool_env_var, is_hip + +_is_hip = is_hip() +_use_aiter = get_bool_env_var("SGLANG_USE_AITER") and _is_hip + +if _use_aiter: + from aiter.tuned_gemm import tgemm if TYPE_CHECKING: from tvm_ffi.module import Module @@ -644,6 +651,23 @@ def fused_rope( positions: torch.Tensor, inverse: bool = False, ) -> None: + """Apply rotary embeddings to both Q and K in a single fused CUDA kernel. + + Args: + q: [batch_size, num_q_heads, rope_dim] bfloat16 + k: [batch_size, num_k_heads, rope_dim] bfloat16 or None + freqs_cis: [max_seq_len, rope_dim // 2] complex64 (full table) + positions: [batch_size] int32 or int64, indices into freqs_cis + inverse: if True, apply inverse rotation (conjugate freqs) + """ + if _is_hip: + from sglang.srt.layers.deepseek_v4_rope import apply_rotary_emb_triton + + apply_rotary_emb_triton(q, freqs_cis, positions=positions, inverse=inverse) + if k is not None: + apply_rotary_emb_triton(k, freqs_cis, positions=positions, inverse=inverse) + return + freqs_real = torch.view_as_real(freqs_cis).flatten(-2).contiguous() module = _jit_fused_rope_module() module.forward(q, k, freqs_real, positions, inverse) @@ -1029,5 +1053,7 @@ def _dispatch_bf16_fp32_backend( z = x.new_empty(x.size(0), y.size(0), dtype=torch.float32) deep_gemm.bf16_gemm_nt(x, y, z) return z + elif _use_aiter: + return tgemm.mm(x, y, otype=torch.float32) else: return torch.nn.functional.linear(x.float(), y.float()) diff --git a/python/sglang/srt/environ.py b/python/sglang/srt/environ.py index eb2292bd5..be06b3d02 100644 --- a/python/sglang/srt/environ.py +++ b/python/sglang/srt/environ.py @@ -571,6 +571,13 @@ class Envs: # ==================================================================== # DeepSeek V4 + SGLANG_OPT_DPSK_V4_RADIX = EnvBool(True) + SGLANG_OPT_USE_OLD_COMPRESSOR = EnvBool(False) + SGLANG_OPT_USE_TRITON_SWA_PREPARE = EnvBool(True) + SGLANG_OPT_USE_AITER_MHC_PRE = EnvBool(True) + SGLANG_OPT_USE_AITER_MHC_POST = EnvBool(True) + SGLANG_OPT_USE_FUSED_COMPRESS = EnvBool(False) + SGLANG_FIX_MTP_HC_HIDDEN = EnvBool(False) # ==================================================================== # Set False when using FP4-to-FP8 converted DeepSeek V4 checkpoint. diff --git a/python/sglang/srt/layers/attention/attention_registry.py b/python/sglang/srt/layers/attention/attention_registry.py index 47350e403..6e0f48bc6 100644 --- a/python/sglang/srt/layers/attention/attention_registry.py +++ b/python/sglang/srt/layers/attention/attention_registry.py @@ -105,11 +105,24 @@ def create_nsa_backend(runner): @register_attention_backend("dsv4") def create_dsv4_backend(runner): - from sglang.srt.layers.attention.deepseek_v4_backend import ( - DeepseekV4AttnBackend, - ) + from sglang.srt.utils import is_hip - return DeepseekV4AttnBackend(runner) + if is_hip(): + from sglang.srt.layers.attention.deepseek_v4_backend_hip_radix import ( + DeepseekV4HipRadixBackend, + ) + + logger.info( + "Using DeepseekV4HipRadixBackend for compressed attention backend (HIP)." + ) + return DeepseekV4HipRadixBackend(runner) + else: + from sglang.srt.layers.attention.deepseek_v4_backend import ( + DeepseekV4AttnBackend, + ) + + logger.info("Using DeepseekV4AttnBackend for dsv4 attention backend (CUDA).") + return DeepseekV4AttnBackend(runner) @register_attention_backend("triton") diff --git a/python/sglang/srt/layers/attention/deepseek_v4_backend_hip_radix.py b/python/sglang/srt/layers/attention/deepseek_v4_backend_hip_radix.py new file mode 100644 index 000000000..9a9a7225a --- /dev/null +++ b/python/sglang/srt/layers/attention/deepseek_v4_backend_hip_radix.py @@ -0,0 +1,1265 @@ +from __future__ import annotations + +import enum +import functools +import logging +from dataclasses import dataclass, field +from typing import ( + TYPE_CHECKING, + Dict, + List, + Literal, + Optional, + Tuple, + TypeVar, + Union, +) + +import torch +import torch.nn.functional as F + +from sglang.srt.environ import envs +from sglang.srt.layers.attention.base_attn_backend import AttentionBackend +from sglang.srt.layers.attention.dsv4.compressor import ( + CompressorBackendMixin, + FusedCompressMetadata, + create_paged_compressor_data, +) +from sglang.srt.layers.attention.dsv4.indexer import C4IndexerBackendMixin +from sglang.srt.layers.attention.dsv4.metadata import ( + PagedIndexerMetadata, + copy_metadata, + maybe_copy_inplace, +) +from sglang.srt.layers.attention.dsv4.metadata_kernel import ( + init_compression_metadata as _init_compression_metadata_triton, +) +from sglang.srt.layers.attention.dsv4.quant_k_cache import ( + quant_to_nope_fp8_rope_bf16_pack_triton, +) +from sglang.srt.layers.dp_attention import ( + get_attention_cp_rank, + get_attention_cp_size, +) +from sglang.srt.mem_cache.deepseek_v4_memory_pool import DeepSeekV4TokenToKVPool +from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode +from sglang.srt.speculative.spec_info import SpecInput +from sglang.srt.utils import ceil_align + +if TYPE_CHECKING: + from flash_mla.flash_mla_interface import FlashMLASchedMeta + + from sglang.srt.layers.radix_attention import RadixAttention + from sglang.srt.model_executor.model_runner import ModelRunner + +logger = logging.getLogger(__name__) + +SWA_WINDOW = 128 +C4_TOPK = 512 +PAGE_INDEX_ALIGNED_SIZE = 64 + + +T = TypeVar("T", bound=Optional[torch.Tensor]) + + +def _pad_last_dim(x: T, multiples_of: int = PAGE_INDEX_ALIGNED_SIZE) -> T: + if x is None: + return None + curr_size = x.shape[-1] + target_size = ceil_align(curr_size, multiples_of) + return F.pad(x, pad=(0, target_size - curr_size), mode="constant", value=-1) + + +def _create_flashmla_metadata(): + from sglang.srt.utils import is_hip + + if is_hip(): + return None + import flash_mla + + return flash_mla.get_mla_metadata()[0] + + +def _create_dummy_paged_compress_data(compress_ratio: int): + return None + + +@dataclass +class DSV4AttnMetadata: + page_size: int + page_table: torch.Tensor + raw_out_loc: torch.Tensor + cuda_int32_kwargs: dict + + seq_lens_casual: torch.Tensor + positions_casual: torch.Tensor + + swa_page_indices: torch.Tensor + swa_topk_lengths: torch.Tensor + + c4_sparse_topk: int + c4_out_loc: Optional[torch.Tensor] = None + c4_topk_lengths_raw: Optional[torch.Tensor] = None + c4_topk_lengths_clamp1: Optional[torch.Tensor] = None + c4_sparse_topk_lengths: torch.Tensor = field(init=False) + c4_sparse_page_indices: torch.Tensor = field(init=False) + + c128_out_loc: Optional[torch.Tensor] = None + c128_page_indices: Optional[torch.Tensor] = None + c128_topk_lengths_clamp1: Optional[torch.Tensor] = None + + c1_flashmla_metadata: FlashMLASchedMeta = field(init=False, repr=False) + c4_flashmla_metadata: FlashMLASchedMeta = field(init=False, repr=False) + c128_flashmla_metadata: FlashMLASchedMeta = field(init=False, repr=False) + + @property + def positions(self) -> torch.Tensor: + return self.positions_casual + + def get_flashmla_metadata(self, compress_ratio: Literal[0, 4, 128]): + if compress_ratio == 0: + return self.c1_flashmla_metadata + elif compress_ratio == 4: + return self.c4_flashmla_metadata + elif compress_ratio == 128: + return self.c128_flashmla_metadata + else: + raise ValueError(f"invalid {compress_ratio=}") + + def copy_(self, other: DSV4AttnMetadata) -> None: + copy_metadata( + src=other, + dst=self, + check_eq_fields=[ + "c4_sparse_topk", + "page_size", + "cuda_int32_kwargs", + ], + copy_fields=[ + "raw_out_loc", + "seq_lens_casual", + "positions_casual", + "c4_out_loc", + "c128_out_loc", + "page_table", + "swa_page_indices", + "swa_topk_lengths", + "c128_page_indices", + "c128_topk_lengths_clamp1", + "c4_topk_lengths_raw", + "c4_topk_lengths_clamp1", + "c4_sparse_topk_lengths", + "c4_sparse_page_indices", + ], + assign_fields=[ + "c1_flashmla_metadata", + "c4_flashmla_metadata", + "c128_flashmla_metadata", + ], + ) + + def init_compression_metadata(self): + assert self.page_table.dim() == 2 + assert ( + self.raw_out_loc.shape == self.seq_lens_casual.shape + ), f"{self.raw_out_loc.shape=}, {self.seq_lens_casual.shape=}" + + ( + self.c4_out_loc, + _, + self.c4_topk_lengths_raw, + self.c4_topk_lengths_clamp1, + self.c128_out_loc, + _, + self.c128_topk_lengths_clamp1, + self.c128_page_indices, + ) = _init_compression_metadata_triton( + self.seq_lens_casual, + self.positions_casual, + self.raw_out_loc, + self.page_table, + self.page_size, + compute_page_indices=True, + ) + + self.c128_page_indices = _pad_last_dim(self.c128_page_indices) + self.swa_page_indices = _pad_last_dim(self.swa_page_indices) + + _CP_REINDEX_FIELDS = [ + "seq_lens_casual", + "positions_casual", + "swa_page_indices", + "swa_topk_lengths", + "page_table", + "c4_topk_lengths_raw", + "c4_topk_lengths_clamp1", + "c128_page_indices", + "c128_topk_lengths_clamp1", + ] + _CP_GLOBAL_FIELDS = [ + "raw_out_loc", + "c4_out_loc", + "c128_out_loc", + ] + + def apply_cp_reindex(self) -> None: + cp_rank = get_attention_cp_rank() + cp_size = get_attention_cp_size() + idx = slice(cp_rank, None, cp_size) + pre_global_len = self.seq_lens_casual.shape[0] + assert pre_global_len % cp_size == 0, ( + f"apply_cp_reindex: global token count {pre_global_len} is not divisible by cp_size={cp_size}. " + "CP round-robin requires padding to ensure divisibility." + ) + expected_local_len = pre_global_len // cp_size + for field_name in self._CP_REINDEX_FIELDS: + val = getattr(self, field_name, None) + assert isinstance( + val, torch.Tensor + ), f"CP reindex: {field_name} is {type(val)}, expected Tensor" + setattr(self, field_name, val[idx].contiguous()) + + for field_name in self._CP_REINDEX_FIELDS: + val = getattr(self, field_name) + assert val.shape[0] == expected_local_len, ( + f"apply_cp_reindex post-condition: {field_name}.shape[0]={val.shape[0]} " + f"!= expected_local_len={expected_local_len} (cp_size={cp_size})" + ) + for field_name in self._CP_GLOBAL_FIELDS: + val = getattr(self, field_name, None) + if val is None: + continue + assert val.shape[0] == pre_global_len, ( + f"apply_cp_reindex post-condition: global field {field_name}.shape[0]={val.shape[0]} " + f"!= pre_global_len={pre_global_len} (must remain global for compressor write path)" + ) + + def init_flashmla_related(self): + # c4_sparse_topk is set from model_config.index_topk per-model + # (small model: 512, large model: 1024). + assert self.c4_sparse_topk in (512, 1024), ( + f"unexpected c4_sparse_topk={self.c4_sparse_topk}; " + "supported: 512 (small) or 1024 (large)" + ) + assert self.c4_topk_lengths_clamp1 is not None + self.c4_sparse_topk_lengths = torch.clamp( + self.c4_topk_lengths_clamp1, max=self.c4_sparse_topk + ) + self.c4_sparse_page_indices = torch.full( + (self.c4_topk_lengths_clamp1.size(0), self.c4_sparse_topk), + -1, + dtype=torch.int32, + device=self.c4_topk_lengths_clamp1.device, + ) + self.c4_sparse_page_indices = _pad_last_dim(self.c4_sparse_page_indices) + self.c1_flashmla_metadata = _create_flashmla_metadata() + self.c4_flashmla_metadata = _create_flashmla_metadata() + self.c128_flashmla_metadata = _create_flashmla_metadata() + + +@dataclass +class DSV4Metadata: + core_attn_metadata: DSV4AttnMetadata + indexer_metadata: Optional[PagedIndexerMetadata] + + c4_compress_metadata: Optional[FusedCompressMetadata] = None + c128_compress_metadata: Optional[FusedCompressMetadata] = None + + @property + def core_metadata(self) -> DSV4AttnMetadata: + return self.core_attn_metadata + + def copy_(self, other: DSV4Metadata): + self.core_attn_metadata.copy_(other.core_attn_metadata) + maybe_copy_inplace(self.indexer_metadata, src=other.indexer_metadata) + maybe_copy_inplace(self.c4_compress_metadata, src=other.c4_compress_metadata) + maybe_copy_inplace( + self.c128_compress_metadata, src=other.c128_compress_metadata + ) + + +@dataclass +class DSV4RawVerifyMetadata: + req_pool_indices: torch.Tensor + seq_lens: torch.Tensor + out_cache_loc: torch.Tensor + + extend_seq_lens: Optional[torch.Tensor] = None + + def copy_(self, other: DSV4RawVerifyMetadata): + self.req_pool_indices.copy_(other.req_pool_indices) + self.seq_lens.copy_(other.seq_lens) + self.out_cache_loc.copy_(other.out_cache_loc) + + self.extend_seq_lens = other.extend_seq_lens + + +@dataclass +class DSV4RawDecodeMetadata: + req_pool_indices: torch.Tensor + seq_lens: torch.Tensor + out_cache_loc: torch.Tensor + + def copy_(self, other: DSV4RawDecodeMetadata): + self.req_pool_indices.copy_(other.req_pool_indices) + self.seq_lens.copy_(other.seq_lens) + self.out_cache_loc.copy_(other.out_cache_loc) + + +class _GraphBucket(enum.Enum): + DECODE_OR_IDLE = "decode_or_idle" + TARGET_VERIFY = "target_verify" + DRAFT_EXTEND = "draft_extend" + + @classmethod + def of(cls, forward_mode: ForwardMode) -> _GraphBucket: + if forward_mode.is_decode_or_idle(): + return cls.DECODE_OR_IDLE + if forward_mode.is_target_verify(): + return cls.TARGET_VERIFY + if forward_mode.is_draft_extend(include_v2=True): + return cls.DRAFT_EXTEND + raise NotImplementedError(f"unsupported {forward_mode=}") + + +class DeepseekV4HipRadixBackend( + AttentionBackend, C4IndexerBackendMixin, CompressorBackendMixin +): + def __init__( + self, + model_runner: ModelRunner, + skip_prefill: bool = False, + speculative_step_id=0, + topk=0, + speculative_num_steps=0, + ): + super().__init__() + self.device = torch.device(model_runner.device) + head_dim = model_runner.model_config.head_dim + assert ( + head_dim == 512 + ), "DSV4 MQA head_dim = qk_nope_head_dim(448) + qk_rope_head_dim(64) = 512" + self.softmax_scale: float = head_dim**-0.5 + self.head_dim_v: int = model_runner.model_config.v_head_dim + self.cuda_int32_kwargs = {"device": self.device, "dtype": torch.int32} + self.swa_page_size = 128 + assert model_runner.page_size is not None + assert model_runner.req_to_token_pool is not None + self.page_size = model_runner.page_size + assert self.page_size == 256, "the system hardcodes page_size=256" + + self.req_to_token = model_runner.req_to_token_pool.req_to_token + self.token_to_kv_pool: DeepSeekV4TokenToKVPool = model_runner.token_to_kv_pool + self.MAX_SEQ_LEN_FOR_CAPTURE = self.req_to_token.shape[1] + + assert isinstance(self.token_to_kv_pool, DeepSeekV4TokenToKVPool) + self.c4_topk = getattr( + model_runner.model_config.hf_text_config, "index_topk", C4_TOPK + ) + + self.topk = model_runner.server_args.speculative_eagle_topk or 0 + assert self.topk in [0, 1], "MTP Topk > 1 not supported for DeepSeek V4" + self.mtp_enabled = self.topk > 0 + self.speculative_num_steps = speculative_num_steps + self.speculative_num_draft_tokens: int = ( + model_runner.server_args.speculative_num_draft_tokens + ) + self.speculative_step_id = speculative_step_id + self.forward_metadata: Union[ + DSV4Metadata, + DSV4RawVerifyMetadata, + DSV4RawDecodeMetadata, + ] = None + self._replay_forward_batch: Optional[ForwardBatch] = None # FIXME: out-of-band + + def _move_to_device(self, x: List[int]) -> torch.Tensor: + pin_tensor = torch.tensor(x, dtype=torch.int32, pin_memory=True) + return pin_tensor.to(self.device, non_blocking=True) + + def init_forward_metadata_indexer(self, core_attn_metadata: DSV4AttnMetadata): + return PagedIndexerMetadata( + page_size=self.page_size, + page_table=core_attn_metadata.page_table, + c4_seq_lens=core_attn_metadata.c4_topk_lengths_raw, + ) + + def init_forward_metadata_decode( + self, + max_seq_len: int, + req_pool_indices: torch.Tensor, + seq_lens: torch.Tensor, + out_cache_loc: torch.Tensor, + ) -> Union[DSV4Metadata, DSV4RawDecodeMetadata]: + assert ( + req_pool_indices.shape[0] == seq_lens.shape[0] == out_cache_loc.shape[0] + ), f"{req_pool_indices.shape=} {seq_lens.shape=} {out_cache_loc.shape=}" + + if envs.SGLANG_PREP_IN_CUDA_GRAPH.get(): + return DSV4RawDecodeMetadata( + req_pool_indices=req_pool_indices, + seq_lens=seq_lens, + out_cache_loc=out_cache_loc, + ) + + core_attn_metadata = self.make_core_attn_metadata( + req_to_token=self.req_to_token, + req_pool_indices_repeated=req_pool_indices, + seq_lens_casual=seq_lens, + max_seq_len=max_seq_len, + out_loc=out_cache_loc, + need_compress=True, + ) + + indexer_metadata = self.init_forward_metadata_indexer(core_attn_metadata) + + create = functools.partial( + create_paged_compressor_data, + is_prefill=False, + token_to_kv_pool=self.token_to_kv_pool, + req_to_token=self.req_to_token, + req_pool_indices=req_pool_indices, + seq_lens=seq_lens, + ) + + return DSV4Metadata( + core_attn_metadata, + indexer_metadata, + c4_compress_metadata=create(compress_ratio=4), + c128_compress_metadata=create(compress_ratio=128), + ) + + def init_forward_metadata_prefill( + self, + max_seq_len: int, + req_pool_indices: torch.Tensor, + seq_lens: torch.Tensor, + seq_lens_cpu: List[int], + out_cache_loc: torch.Tensor, + num_tokens: int, + extend_seq_lens: torch.Tensor, + extend_seq_lens_cpu: List[int], + need_compress: bool = True, + use_prefill_cuda_graph: bool = False, + ) -> DSV4Metadata: + seq_lens_casual, req_pool_indices_repeated = self.expand_prefill_casually( + num_tokens=num_tokens, + seq_lens=seq_lens_cpu, + extend_seq_lens=extend_seq_lens_cpu, + req_pool_indices=req_pool_indices, + padded_num_tokens=out_cache_loc.shape[0], + ) + core_attn_metadata = self.make_core_attn_metadata( + req_to_token=self.req_to_token, + req_pool_indices_repeated=req_pool_indices_repeated, + seq_lens_casual=seq_lens_casual, + max_seq_len=max_seq_len, + out_loc=out_cache_loc, + need_compress=need_compress, + is_prefill=True, + ) + indexer_metadata = ( + self.init_forward_metadata_indexer(core_attn_metadata) + if need_compress + else None + ) + if not need_compress: + create = _create_dummy_paged_compress_data + else: + create = functools.partial( + create_paged_compressor_data, + is_prefill=True, + token_to_kv_pool=self.token_to_kv_pool, + req_to_token=self.req_to_token, + req_pool_indices=req_pool_indices, + seq_lens=seq_lens, + seq_lens_cpu=seq_lens_cpu, + extend_lens=extend_seq_lens, + extend_lens_cpu=extend_seq_lens_cpu, + use_prefill_cuda_graph=use_prefill_cuda_graph, + ) + return DSV4Metadata( + core_attn_metadata, + indexer_metadata, + c4_compress_metadata=create(compress_ratio=4), + c128_compress_metadata=create(compress_ratio=128), + ) + + def init_forward_metadata_target_verify( + self, + max_seq_len: int, + req_pool_indices: torch.Tensor, + seq_lens: torch.Tensor, + out_cache_loc: Optional[torch.Tensor] = None, + use_prefill_cuda_graph: bool = False, + ) -> Union[DSV4Metadata, DSV4RawVerifyMetadata]: + if envs.SGLANG_PREP_IN_CUDA_GRAPH.get(): + assert out_cache_loc is not None + if not hasattr(self, "extend_seq_lens_buffer"): + self.extend_seq_lens_buffer = torch.tensor( + [self.speculative_num_draft_tokens] * 1025, device=self.device + ) + extend_seq_lens = self.extend_seq_lens_buffer[: len(seq_lens)] + + return DSV4RawVerifyMetadata( + req_pool_indices=req_pool_indices, + seq_lens=seq_lens, + out_cache_loc=out_cache_loc, + extend_seq_lens=extend_seq_lens, + ) + else: + seq_lens_cpu = seq_lens.tolist() + return self.init_forward_metadata_target_verify_old( + max_seq_len=max_seq_len, + req_pool_indices=req_pool_indices, + seq_lens=seq_lens, + seq_lens_cpu=seq_lens_cpu, + out_cache_loc=out_cache_loc, + use_prefill_cuda_graph=use_prefill_cuda_graph, + ) + + def init_forward_metadata_target_verify_old( + self, + max_seq_len: int, + req_pool_indices: torch.Tensor, + seq_lens: torch.Tensor, + seq_lens_cpu: Optional[List[int]] = None, + out_cache_loc: Optional[torch.Tensor] = None, + use_prefill_cuda_graph: bool = False, + ) -> DSV4Metadata: + batch_size = len(seq_lens) + seq_lens = seq_lens + self.speculative_num_draft_tokens + seq_lens_cpu = [x + self.speculative_num_draft_tokens for x in seq_lens_cpu] + extend_seq_lens_cpu = [self.speculative_num_draft_tokens] * batch_size + extend_seq_lens = self._move_to_device(extend_seq_lens_cpu) + num_tokens = self.speculative_num_draft_tokens * batch_size + if out_cache_loc is None: + out_cache_loc = seq_lens.new_zeros(num_tokens) + return self.init_forward_metadata_prefill( + max_seq_len=max_seq_len, + req_pool_indices=req_pool_indices, + seq_lens=seq_lens, + seq_lens_cpu=seq_lens_cpu, + out_cache_loc=out_cache_loc, + num_tokens=num_tokens, + extend_seq_lens=extend_seq_lens, + extend_seq_lens_cpu=extend_seq_lens_cpu, + need_compress=True, + use_prefill_cuda_graph=use_prefill_cuda_graph, + ) + + def make_forward_metadata_from_raw_verify( + self, raw_metadata: DSV4RawVerifyMetadata + ) -> DSV4Metadata: + req_pool_indices = raw_metadata.req_pool_indices + seq_lens = raw_metadata.seq_lens + out_cache_loc = raw_metadata.out_cache_loc + + bs, num_draft_tokens = len(seq_lens), self.speculative_num_draft_tokens + seq_lens = seq_lens + self.speculative_num_draft_tokens + extend_seq_lens = raw_metadata.extend_seq_lens + + seq_lens_casual, req_pool_indices_repeated = ( + self.expand_extend_with_same_length( + bs, num_draft_tokens, seq_lens, req_pool_indices + ) + ) + core_attn_metadata = self.make_core_attn_metadata( + req_to_token=self.req_to_token, + req_pool_indices_repeated=req_pool_indices_repeated, + seq_lens_casual=seq_lens_casual, + max_seq_len=self.MAX_SEQ_LEN_FOR_CAPTURE, + out_loc=out_cache_loc, + need_compress=True, + ) + indexer_metadata = self.init_forward_metadata_indexer(core_attn_metadata) + create = functools.partial( + create_paged_compressor_data, + is_prefill=True, + token_to_kv_pool=self.token_to_kv_pool, + req_to_token=self.req_to_token, + req_pool_indices=req_pool_indices, + seq_lens=seq_lens, + extend_lens=extend_seq_lens, + seq_lens_cpu=None, + extend_lens_cpu=None, + use_prefill_cuda_graph=True, + num_q_tokens=num_draft_tokens * bs, + ) + return DSV4Metadata( + core_attn_metadata, + indexer_metadata, + c4_compress_metadata=create(compress_ratio=4), + c128_compress_metadata=create(compress_ratio=128), + ) + + def make_forward_metadata_from_raw_decode( + self, raw_metadata: DSV4RawDecodeMetadata + ) -> DSV4Metadata: + req_pool_indices = raw_metadata.req_pool_indices + seq_lens = raw_metadata.seq_lens + out_cache_loc = raw_metadata.out_cache_loc + + core_attn_metadata = self.make_core_attn_metadata( + req_to_token=self.req_to_token, + req_pool_indices_repeated=req_pool_indices, + seq_lens_casual=seq_lens, + max_seq_len=self.MAX_SEQ_LEN_FOR_CAPTURE, + out_loc=out_cache_loc, + need_compress=True, + ) + indexer_metadata = self.init_forward_metadata_indexer(core_attn_metadata) + + create = functools.partial( + create_paged_compressor_data, + is_prefill=False, + token_to_kv_pool=self.token_to_kv_pool, + req_to_token=self.req_to_token, + req_pool_indices=req_pool_indices, + seq_lens=seq_lens, + ) + + return DSV4Metadata( + core_attn_metadata, + indexer_metadata, + c4_compress_metadata=create(compress_ratio=4), + c128_compress_metadata=create(compress_ratio=128), + ) + + def init_forward_metadata_draft_extend( + self, + max_seq_len: int, + req_pool_indices: torch.Tensor, + seq_lens: torch.Tensor, + seq_lens_cpu: List[int], + num_tokens_per_bs: int, + out_cache_loc: Optional[torch.Tensor] = None, + use_prefill_cuda_graph: bool = False, + ) -> DSV4Metadata: + batch_size = len(seq_lens) + extend_seq_lens_cpu = [num_tokens_per_bs] * batch_size + extend_seq_lens = self._move_to_device(extend_seq_lens_cpu) + num_tokens = num_tokens_per_bs * batch_size + if out_cache_loc is None: + out_cache_loc = seq_lens.new_zeros(num_tokens) + return self.init_forward_metadata_prefill( + seq_lens=seq_lens, + max_seq_len=max_seq_len, + req_pool_indices=req_pool_indices, + seq_lens_cpu=seq_lens_cpu, + out_cache_loc=out_cache_loc, + num_tokens=num_tokens, + extend_seq_lens=extend_seq_lens, + extend_seq_lens_cpu=extend_seq_lens_cpu, + need_compress=False, + use_prefill_cuda_graph=use_prefill_cuda_graph, + ) + + def init_forward_metadata(self, forward_batch: ForwardBatch) -> None: + if self.mtp_enabled and forward_batch.forward_mode.is_idle(): + return + + req_pool_indices = forward_batch.req_pool_indices + seq_lens = forward_batch.seq_lens.to(torch.int32) + seq_lens_cpu = forward_batch.seq_lens_cpu + assert forward_batch.req_to_token_pool.req_to_token is self.req_to_token + + assert self.swa_page_size % SWA_WINDOW == 0 and self.page_size % 128 == 0 + assert seq_lens_cpu is not None + max_seq_len = int(seq_lens_cpu.max().item()) + + if forward_batch.forward_mode.is_decode_or_idle(): + metadata = self.init_forward_metadata_decode( + max_seq_len=max_seq_len, + req_pool_indices=req_pool_indices, + seq_lens=seq_lens, + out_cache_loc=forward_batch.out_cache_loc, + ) + elif forward_batch.forward_mode.is_target_verify(): + metadata = self.init_forward_metadata_target_verify( + max_seq_len=max_seq_len, + req_pool_indices=req_pool_indices, + seq_lens=seq_lens, + out_cache_loc=forward_batch.out_cache_loc, + ) + elif forward_batch.forward_mode.is_prefill(include_draft_extend_v2=True): + extend_seq_lens_cpu = forward_batch.extend_seq_lens_cpu + extend_seq_lens = forward_batch.extend_seq_lens + assert ( + seq_lens is not None + and seq_lens_cpu is not None + and extend_seq_lens is not None + and extend_seq_lens_cpu is not None + ) + is_draft = forward_batch.forward_mode.is_draft_extend(include_v2=True) + metadata = self.init_forward_metadata_prefill( + max_seq_len=max_seq_len, + req_pool_indices=req_pool_indices, + seq_lens=seq_lens, + seq_lens_cpu=seq_lens_cpu.tolist(), + out_cache_loc=forward_batch.out_cache_loc, + num_tokens=sum(extend_seq_lens_cpu), + extend_seq_lens=extend_seq_lens, + extend_seq_lens_cpu=extend_seq_lens_cpu, + need_compress=not is_draft, + ) + else: + raise NotImplementedError(f"unsupported mode {forward_batch.forward_mode=}") + + self.forward_metadata = metadata + + def init_cuda_graph_state(self, max_bs: int, max_num_tokens: int) -> None: + self.cuda_graph_metadata_of_bucket_and_bs: Dict[ + _GraphBucket, + Dict[ + int, + Union[DSV4Metadata, DSV4RawDecodeMetadata, DSV4RawVerifyMetadata], + ], + ] = {bucket: {} for bucket in _GraphBucket} + self.draft_extend_num_tokens_per_bs = ( + max_num_tokens // max_bs if max_bs > 0 else 1 + ) + + def init_forward_metadata_capture_cuda_graph( + self, + bs: int, + num_tokens: int, + req_pool_indices: torch.Tensor, + seq_lens: torch.Tensor, + encoder_lens: Optional[torch.Tensor], + forward_mode: ForwardMode, + spec_info: Optional[SpecInput], + ) -> None: + assert req_pool_indices.size(0) == bs + assert seq_lens.size(0) == bs + + bucket = _GraphBucket.of(forward_mode) + raw_type: Optional[type] = None + if bucket == _GraphBucket.DECODE_OR_IDLE: + metadata = self.init_forward_metadata_decode( + max_seq_len=self.MAX_SEQ_LEN_FOR_CAPTURE, + req_pool_indices=req_pool_indices, + seq_lens=seq_lens, + out_cache_loc=torch.zeros_like(seq_lens), + ) + raw_type = DSV4RawDecodeMetadata + elif bucket == _GraphBucket.TARGET_VERIFY: + out_cache_loc = torch.zeros(num_tokens, **self.cuda_int32_kwargs) + metadata = self.init_forward_metadata_target_verify( + max_seq_len=self.MAX_SEQ_LEN_FOR_CAPTURE, + req_pool_indices=req_pool_indices, + seq_lens=seq_lens, + out_cache_loc=out_cache_loc, + use_prefill_cuda_graph=True, + ) + raw_type = DSV4RawVerifyMetadata + elif bucket == _GraphBucket.DRAFT_EXTEND: + num_tokens_per_bs = num_tokens // bs + metadata = self.init_forward_metadata_draft_extend( + max_seq_len=self.MAX_SEQ_LEN_FOR_CAPTURE, + req_pool_indices=req_pool_indices, + seq_lens=seq_lens, + seq_lens_cpu=seq_lens.tolist(), + num_tokens_per_bs=num_tokens_per_bs, + use_prefill_cuda_graph=True, + ) + else: + raise NotImplementedError(f"{forward_mode=} not supported yet") + + self.cuda_graph_metadata_of_bucket_and_bs[bucket][bs] = metadata + self.forward_metadata = metadata + if raw_type is not None: + self._current_capture_raw = ( + metadata if isinstance(metadata, raw_type) else None + ) + + def init_forward_metadata_replay_cuda_graph( + self, + bs: int, + req_pool_indices: torch.Tensor, + seq_lens: torch.Tensor, + seq_lens_sum: int, + encoder_lens: Optional[torch.Tensor], + forward_mode: ForwardMode, + spec_info: Optional[SpecInput], + seq_lens_cpu: Optional[torch.Tensor], + ) -> None: + bucket = _GraphBucket.of(forward_mode) + + # FIXME: see cuda_graph_runner — this attribute is set out-of-band. + fb = self._replay_forward_batch + out_cache_loc = fb.out_cache_loc + actual_forward_mode = fb.forward_mode + + if actual_forward_mode == ForwardMode.IDLE: + logger.debug( + f"[IDLE replay] bs={bs}, " + f"local_seq_lens_len={len(seq_lens)}, " + f"has_graph={bs in self.cuda_graph_metadata_of_bucket_and_bs[_GraphBucket.DECODE_OR_IDLE]}" + ) + device = seq_lens.device + seq_lens = torch.ones(bs, dtype=seq_lens.dtype, device=device) + seq_lens_cpu = torch.ones(bs, dtype=torch.int64) + seq_lens_sum = bs + req_pool_indices = torch.zeros( + bs, dtype=req_pool_indices.dtype, device=device + ) + out_cache_loc = torch.zeros(bs, dtype=torch.int64, device=device) + + assert seq_lens_cpu is not None + seq_lens = seq_lens[:bs] + seq_lens_cpu = seq_lens_cpu[:bs] + req_pool_indices = req_pool_indices[:bs] + + actual_max_seq_len = seq_lens_cpu.max().item() + chosen_max_seq_len = self.MAX_SEQ_LEN_FOR_CAPTURE + assert actual_max_seq_len <= chosen_max_seq_len + + if bucket == _GraphBucket.DECODE_OR_IDLE: + assert out_cache_loc is not None + assert len(out_cache_loc.shape) == 1, f"{out_cache_loc.shape=}" + out_cache_loc_padded = torch.nn.functional.pad( + out_cache_loc, + pad=(0, bs - len(out_cache_loc)), + mode="constant", + value=0, + ) + temp_metadata = self.init_forward_metadata_decode( + max_seq_len=chosen_max_seq_len, + req_pool_indices=req_pool_indices, + seq_lens=seq_lens, + out_cache_loc=out_cache_loc_padded, + ) + elif bucket == _GraphBucket.TARGET_VERIFY: + assert out_cache_loc is not None + num_tokens = self.speculative_num_draft_tokens * bs + out_cache_loc_padded = torch.nn.functional.pad( + out_cache_loc, + pad=(0, num_tokens - len(out_cache_loc)), + mode="constant", + value=0, + ) + temp_metadata = self.init_forward_metadata_target_verify( + max_seq_len=chosen_max_seq_len, + req_pool_indices=req_pool_indices, + seq_lens=seq_lens, + out_cache_loc=out_cache_loc_padded, + use_prefill_cuda_graph=True, + ) + elif bucket == _GraphBucket.DRAFT_EXTEND: + num_tokens_per_bs = self.draft_extend_num_tokens_per_bs + temp_metadata = self.init_forward_metadata_draft_extend( + max_seq_len=chosen_max_seq_len, + req_pool_indices=req_pool_indices, + seq_lens=seq_lens, + seq_lens_cpu=seq_lens_cpu.tolist(), + num_tokens_per_bs=num_tokens_per_bs, + use_prefill_cuda_graph=True, + ) + else: + raise NotImplementedError + + self.replay_cuda_graph_metadata_from( + bs=bs, temp_metadata=temp_metadata, bucket=bucket + ) + + def replay_cuda_graph_metadata_from( + self, + bs: int, + temp_metadata: Union[ + DSV4Metadata, + DSV4RawVerifyMetadata, + DSV4RawDecodeMetadata, + ], + bucket: _GraphBucket, + ) -> None: + chosen_metadata = self.cuda_graph_metadata_of_bucket_and_bs[bucket][bs] + chosen_metadata.copy_(temp_metadata) + self.forward_metadata = chosen_metadata + + def get_cuda_graph_seq_len_fill_value(self): + return 1 + + def on_after_cuda_graph_warmup(self): + metadata = self.forward_metadata + if isinstance(metadata, DSV4Metadata) and isinstance( + metadata.core_attn_metadata, DSV4AttnMetadata + ): + core = metadata.core_attn_metadata + core.c1_flashmla_metadata = _create_flashmla_metadata() + core.c4_flashmla_metadata = _create_flashmla_metadata() + core.c128_flashmla_metadata = _create_flashmla_metadata() + + # PREP_IN_CUDA_GRAPH=True: warmup upgraded raw->full on the host; + # restore raw so capture re-runs the upgrade inside the graph. + current_raw = getattr(self, "_current_capture_raw", None) + if current_raw is not None: + self.forward_metadata = current_raw + + def store_cache( + self, layer_id: int, swa_k: torch.Tensor, forward_batch: ForwardBatch + ) -> None: + raw_loc = forward_batch.out_cache_loc + if envs.SGLANG_OPT_USE_FUSED_STORE_CACHE.get(): + self.token_to_kv_pool.set_swa_key_buffer_radix_fused( + layer_id=layer_id, + raw_loc=raw_loc, + cache_k=swa_k, + ) + else: + swa_k_pack = quant_to_nope_fp8_rope_bf16_pack_triton(swa_k) + self.token_to_kv_pool.set_swa_key_buffer_radix( + layer_id=layer_id, + raw_loc=raw_loc, + cache_nope_fp8_rope_bf16_pack=swa_k_pack, + ) + + def _maybe_upgrade_forward_metadata(self) -> None: + # With SGLANG_PREP_IN_CUDA_GRAPH=1, init_forward_metadata_* + # returns a Raw metadata that only carries a few tensors. The + # full DSV4Metadata (including c4/c128 compress + core_attn + + # indexer metadata) must be materialized before any caller that + # touches those fields. For 1.6T the first two layers have + # compress_ratio=128, so forward_core_compressor / forward_c4_indexer + # can fire before attn_backend.forward(), and must trigger the + # upgrade themselves. + if isinstance(self.forward_metadata, DSV4RawVerifyMetadata): + self.forward_metadata = self.make_forward_metadata_from_raw_verify( + raw_metadata=self.forward_metadata, + ) + elif isinstance(self.forward_metadata, DSV4RawDecodeMetadata): + self.forward_metadata = self.make_forward_metadata_from_raw_decode( + raw_metadata=self.forward_metadata, + ) + + def forward( + self, + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + layer: RadixAttention, + forward_batch: ForwardBatch, + compress_ratio: Literal[0, 4, 128], + save_kv_cache: bool = True, + attn_sink: Optional[torch.Tensor] = None, + **_, + ) -> torch.Tensor: + self._maybe_upgrade_forward_metadata() + + if self.mtp_enabled and forward_batch.forward_mode.is_idle(): + return q.new_empty(q.shape[0], q.shape[1], layer.v_head_dim) + + assert k is v, "DeepseekV4 shares k and v" + swa_k = k + + layer_id = layer.layer_id + metadata = self.forward_metadata + core_attn_metadata = metadata.core_attn_metadata + token_to_kv_pool = forward_batch.token_to_kv_pool + assert isinstance(token_to_kv_pool, DeepSeekV4TokenToKVPool) + + if isinstance(core_attn_metadata, DSV4AttnMetadata): + if save_kv_cache: + self.store_cache(layer_id, swa_k, forward_batch) + swa_k_cache = token_to_kv_pool.get_swa_key_buffer_radix(layer_id) + + extra_k_cache, extra_indices, extra_topk_lengths = None, None, None + if compress_ratio == 4: + extra_k_cache = token_to_kv_pool.get_extra_key_buffer(layer_id) + extra_indices = core_attn_metadata.c4_sparse_page_indices + extra_topk_lengths = core_attn_metadata.c4_sparse_topk_lengths + elif compress_ratio == 128: + extra_k_cache = token_to_kv_pool.get_extra_key_buffer(layer_id) + extra_indices = core_attn_metadata.c128_page_indices + extra_topk_lengths = core_attn_metadata.c128_topk_lengths_clamp1 + + swa_window_size = token_to_kv_pool.swa_window_size + assert swa_k_cache.ndim == 2 + k_cache_total_dim = token_to_kv_pool.swa_kv_pool.kv_cache_total_dim + swa_k_cache = swa_k_cache[:, : swa_window_size * k_cache_total_dim].view( + swa_k_cache.shape[0], swa_window_size, 1, k_cache_total_dim + ) + + if extra_k_cache is not None: + page_sizes = { + 4: token_to_kv_pool.page_size // 4, + 128: token_to_kv_pool.page_size // 128, + } + extra_k_cache = extra_k_cache[ + :, : page_sizes[compress_ratio] * k_cache_total_dim + ].view( + extra_k_cache.shape[0], + page_sizes[compress_ratio], + 1, + k_cache_total_dim, + ) + swa_page_indices = core_attn_metadata.swa_page_indices + swa_topk_lengths = core_attn_metadata.swa_topk_lengths + + if self.mtp_enabled: + if swa_page_indices.shape[0] != q.shape[0]: + swa_page_indices = _pad_tensor_to_size( + swa_page_indices, q.shape[0], value=0 + ) + + if swa_topk_lengths.shape[0] != q.shape[0]: + swa_topk_lengths = _pad_tensor_to_size( + swa_topk_lengths, q.shape[0], value=1 + ) + + if q.ndim == 3: + q = q.unsqueeze(1) + if swa_page_indices.ndim == 2: + swa_page_indices = swa_page_indices.unsqueeze(1) + if extra_indices is not None and extra_indices.ndim == 2: + extra_indices = extra_indices.unsqueeze(1) + + assert attn_sink is not None + + flashmla_metadata = core_attn_metadata.get_flashmla_metadata(compress_ratio) + + assert ( + swa_page_indices.shape[-1] % 64 == 0 + ), f"{swa_page_indices.shape=}'s last dimension is not aligned to 64" + if extra_indices is not None: + assert ( + extra_indices.shape[-1] % 64 == 0 + ), f"{extra_indices.shape=}'s last dimension is not aligned to 64" + + import os + + from sglang.srt.layers.attention.hip_flash_mla import ( + flash_mla_with_kvcache_entrypoint, + ) + + backend = os.environ.get("SGLANG_HACK_FLASHMLA_BACKEND", "kernel") + input_dict = dict( + q=q, + k_cache=swa_k_cache, + head_dim_v=self.head_dim_v, + block_table=None, + cache_seqlens=None, + tile_scheduler_metadata=flashmla_metadata, + softmax_scale=self.softmax_scale, + is_fp8_kvcache=True, + indices=swa_page_indices, + topk_length=swa_topk_lengths, + attn_sink=attn_sink, + extra_k_cache=extra_k_cache, + extra_indices_in_kvcache=extra_indices, + extra_topk_length=extra_topk_lengths, + ) + o = flash_mla_with_kvcache_entrypoint(**input_dict, backend=backend)[0] + + o = o.squeeze(1) + return o + + raise NotImplementedError("ragged attention") + + def expand_prefill_casually( + self, + num_tokens: int, + seq_lens: List[int], + extend_seq_lens: List[int], + req_pool_indices: torch.Tensor, + padded_num_tokens: Optional[int], + ) -> Tuple[torch.Tensor, torch.Tensor]: + seq_lens_casual = torch.empty(num_tokens, **self.cuda_int32_kwargs) + idx_to_req_repeated = torch.empty(num_tokens, **self.cuda_int32_kwargs) + offset = 0 + for i, (kv_len, qo_len) in enumerate(zip(seq_lens, extend_seq_lens)): + out = seq_lens_casual[offset : offset + qo_len] + offset += qo_len + torch.arange(kv_len - qo_len + 1, kv_len + 1, out=out) + idx_to_req_repeated[offset - qo_len : offset].fill_(i) + + assert offset == num_tokens + req_pool_indices_repeated = req_pool_indices[idx_to_req_repeated] + + if padded_num_tokens is not None and padded_num_tokens > num_tokens: + pad_size = padded_num_tokens - num_tokens + seq_lens_casual = torch.nn.functional.pad( + seq_lens_casual, + (0, pad_size), + value=1, + ) + req_pool_indices_repeated = torch.nn.functional.pad( + req_pool_indices_repeated, + (0, pad_size), + value=req_pool_indices_repeated[-1].item(), + ) + + return seq_lens_casual, req_pool_indices_repeated + + def expand_extend_with_same_length( + self, + bs: int, + qo_len: int, + seq_lens: torch.Tensor, + req_pool_indices: torch.Tensor, + ): + seq_lens_casual = seq_lens[:, None] + torch.arange( + -qo_len + 1, 1, **self.cuda_int32_kwargs + ) + seq_lens_casual = seq_lens_casual.flatten() + idx_to_req_repeated = torch.arange( + bs, **self.cuda_int32_kwargs + ).repeat_interleave(qo_len) + req_pool_indices_repeated = req_pool_indices[idx_to_req_repeated] + return seq_lens_casual, req_pool_indices_repeated + + def make_core_attn_metadata( + self, + req_to_token: torch.Tensor, + req_pool_indices_repeated: torch.Tensor, + seq_lens_casual: torch.Tensor, + max_seq_len: int, + out_loc: torch.Tensor, + need_compress: bool = True, + is_prefill: bool = False, + ) -> DSV4AttnMetadata: + assert self.swa_page_size == SWA_WINDOW + + swa_page_indices = self.get_swa_page_indices( + seq_lens_casual=seq_lens_casual, + req_pool_indices_repeated=req_pool_indices_repeated, + ) + + swa_page_indices = _pad_last_dim( + swa_page_indices, multiples_of=PAGE_INDEX_ALIGNED_SIZE + ) + + raw_positions = seq_lens_casual - 1 + swa_topk_lengths = torch.clamp(seq_lens_casual, max=SWA_WINDOW) + + page_table = req_to_token[ + req_pool_indices_repeated, : max_seq_len : self.page_size + ] + page_table = (page_table // self.page_size).to(torch.int32) + + core_attn_metadata = DSV4AttnMetadata( + page_size=self.page_size, + raw_out_loc=out_loc, + seq_lens_casual=seq_lens_casual, + cuda_int32_kwargs=self.cuda_int32_kwargs, + positions_casual=raw_positions, + page_table=page_table, + swa_page_indices=swa_page_indices, + swa_topk_lengths=swa_topk_lengths, + c4_sparse_topk=self.c4_topk, + ) + + if need_compress: + core_attn_metadata.init_compression_metadata() + core_attn_metadata.init_flashmla_related() + else: + core_attn_metadata.c4_sparse_topk_lengths = None + core_attn_metadata.c4_sparse_page_indices = None + core_attn_metadata.c1_flashmla_metadata = _create_flashmla_metadata() + core_attn_metadata.c4_flashmla_metadata = None + core_attn_metadata.c128_flashmla_metadata = None + return core_attn_metadata + + def get_swa_page_indices( + self, + seq_lens_casual: torch.Tensor, + req_pool_indices_repeated: torch.Tensor, + ) -> torch.Tensor: + pos_causal = seq_lens_casual - 1 + num_qo_tokens = seq_lens_casual.size(0) + offsets = pos_causal.unsqueeze(1) - torch.arange( + SWA_WINDOW, **self.cuda_int32_kwargs + ).unsqueeze(0) + invalid_offset_mask = offsets < 0 + offsets.masked_fill_(invalid_offset_mask, 0) + raw_indices = self.req_to_token[req_pool_indices_repeated[:, None], offsets] + assert raw_indices.shape == (num_qo_tokens, SWA_WINDOW) + raw_indices.masked_fill_(invalid_offset_mask, -1) + swa_indices = self.token_to_kv_pool.translate_loc_from_full_to_swa(raw_indices) + return swa_indices + + +class DeepseekV4MultiStepBackend(DeepseekV4HipRadixBackend): + def __init__( + self, model_runner: ModelRunner, topk: int, speculative_num_steps: int + ): + super().__init__(model_runner) + self.model_runner = model_runner + self.topk = topk + self.speculative_num_steps = speculative_num_steps + self.attn_backends: List[DeepseekV4HipRadixBackend] = [] + for i in range(self.speculative_num_steps): + self.attn_backends.append( + DeepseekV4HipRadixBackend( + model_runner, + speculative_step_id=i, + topk=self.topk, + speculative_num_steps=self.speculative_num_steps, + ) + ) + + def init_forward_metadata(self, forward_batch: ForwardBatch): + for i in range(self.speculative_num_steps - 1): + self.attn_backends[i].init_forward_metadata(forward_batch) + + def init_cuda_graph_state(self, max_bs: int, max_num_tokens: int): + for i in range(self.speculative_num_steps): + self.attn_backends[i].init_cuda_graph_state(max_bs, max_num_tokens) + + def init_forward_metadata_capture_cuda_graph(self, forward_batch: ForwardBatch): + for i in range(self.speculative_num_steps): + self.attn_backends[i].init_forward_metadata_capture_cuda_graph( + forward_batch.batch_size, + forward_batch.batch_size * self.topk, + forward_batch.req_pool_indices, + forward_batch.seq_lens, + encoder_lens=None, + forward_mode=ForwardMode.DECODE, + spec_info=forward_batch.spec_info, + ) + + def on_after_cuda_graph_warmup(self): + for backend in self.attn_backends: + backend.on_after_cuda_graph_warmup() + + def init_forward_metadata_replay_cuda_graph( + self, forward_batch: ForwardBatch, bs: int + ): + if self.speculative_num_steps == 1: + return + + self.attn_backends[0]._replay_forward_batch = forward_batch + self.attn_backends[0].init_forward_metadata_replay_cuda_graph( + bs=bs, + req_pool_indices=forward_batch.req_pool_indices, + seq_lens=forward_batch.seq_lens, + seq_lens_sum=forward_batch.seq_lens_sum, + encoder_lens=None, + forward_mode=ForwardMode.DECODE, + spec_info=forward_batch.spec_info, + seq_lens_cpu=forward_batch.seq_lens_cpu, + ) + self.attn_backends[0]._replay_forward_batch = None + temp_metadata = self.attn_backends[0].forward_metadata + + for i in range(1, self.speculative_num_steps - 1): + self.attn_backends[i].replay_cuda_graph_metadata_from( + bs=bs, + temp_metadata=temp_metadata, + bucket=_GraphBucket.DECODE_OR_IDLE, + ) + + +def _pad_tensor_to_size(tensor: torch.Tensor, size: int, *, value: int = 0): + if value == 0: + return torch.cat( + [tensor, tensor.new_zeros(size - tensor.shape[0], *tensor.shape[1:])], + dim=0, + ) + else: + return torch.cat( + [ + tensor, + tensor.new_full((size - tensor.shape[0], *tensor.shape[1:]), value), + ], + dim=0, + ) diff --git a/python/sglang/srt/layers/attention/dsv4/compress_hip.py b/python/sglang/srt/layers/attention/dsv4/compress_hip.py new file mode 100644 index 000000000..b796de3a5 --- /dev/null +++ b/python/sglang/srt/layers/attention/dsv4/compress_hip.py @@ -0,0 +1,455 @@ +from __future__ import annotations + +import os +from functools import cached_property +from typing import TYPE_CHECKING, Any + +import torch +import torch.nn as nn +import triton +import triton.language as tl + +from sglang.srt.environ import envs +from sglang.srt.layers.attention.dsv4.compressor import Compressor as _CompressorBase +from sglang.srt.layers.attention.nsa.nsa_indexer import rotate_activation +from sglang.srt.layers.deepseek_v4_rope import ( + apply_rotary_emb_triton, + fused_norm_rope_inplace_triton, +) +from sglang.srt.mem_cache.deepseek_v4_compress_state import ( + CompressStatePool, + KVAndScore, +) +from sglang.srt.mem_cache.deepseek_v4_memory_pool import DeepSeekV4TokenToKVPool + +if TYPE_CHECKING: + from sglang.srt.layers.attention.deepseek_v4_backend_hip_radix import ( + DeepseekV4HipRadixBackend, + ) + from sglang.srt.model_executor.forward_batch_info import ForwardBatch + + +@triton.jit +def _rms_normalize_kernel( + x_ptr, + weight_ptr, + eps, + stride_row, + dim, + BLOCK_SIZE: tl.constexpr, + HAS_WEIGHT: tl.constexpr, +): + pid = tl.program_id(0) + offs = tl.arange(0, BLOCK_SIZE) + mask = offs < dim + base = pid * stride_row + x = tl.load(x_ptr + base + offs, mask=mask, other=0.0).to(tl.float32) + mean_sq = tl.sum(x * x, axis=0) / dim + rms_inv = tl.rsqrt(mean_sq + eps) + out = x * rms_inv + if HAS_WEIGHT: + weight = tl.load(weight_ptr + offs, mask=mask, other=0.0) + out = out * weight + tl.store(x_ptr + base + offs, out, mask=mask) + + +def rms_normalize_triton( + x: torch.Tensor, eps: float, weight: torch.Tensor = None +) -> torch.Tensor: + dim = x.shape[-1] + x_flat = x.view(-1, dim) + num_rows = x_flat.shape[0] + BLOCK_SIZE = triton.next_power_of_2(dim) + grid = (num_rows,) + _rms_normalize_kernel[grid]( + x_flat, + weight, + eps, + x_flat.stride(0), + dim, + BLOCK_SIZE=BLOCK_SIZE, + HAS_WEIGHT=(weight is not None), + ) + return x + + +class DeepseekRefRMSNorm(nn.Module): + def __init__(self, dim: int, eps: float = 1e-6): + super().__init__() + self.dim = dim + self.eps = eps + self.weight = nn.Parameter(torch.ones(dim, dtype=torch.float32)) + + def forward(self, x: torch.Tensor): + return rms_normalize_triton(x, self.eps, self.weight) + + +class CompressorHip(_CompressorBase): + """HIP (ROCm) specific Compressor implementation.""" + + def __init__(self, *args, **kwargs) -> None: + super().__init__(*args, **kwargs) + self.norm = DeepseekRefRMSNorm(self.head_dim, eps=self.norm.variance_epsilon) + + @cached_property + def use_fused_compress(self) -> bool: + return False + + @cached_property + def use_hip_fused_compress(self) -> bool: + return envs.SGLANG_OPT_USE_FUSED_COMPRESS.get() + + def _get_states(self, forward_batch: ForwardBatch) -> KVAndScore: + token_to_kv_pool = forward_batch.token_to_kv_pool + assert isinstance(token_to_kv_pool, DeepSeekV4TokenToKVPool) + if self.is_in_indexer: + return token_to_kv_pool.get_indexer_compress_states(self.layer_id) + else: + return token_to_kv_pool.get_attention_compress_states(self.layer_id) + + def _get_state_pool(self, forward_batch: ForwardBatch) -> CompressStatePool: + token_to_kv_pool = forward_batch.token_to_kv_pool + assert isinstance(token_to_kv_pool, DeepSeekV4TokenToKVPool) + if self.is_in_indexer: + ret = token_to_kv_pool.get_indexer_compress_states(self.layer_id) + else: + ret = token_to_kv_pool.get_attention_compress_states(self.layer_id) + + assert isinstance(ret, CompressStatePool) + + return ret + + def overlap_transform(self, tensor: torch.Tensor, fill_value: Any) -> torch.Tensor: + assert tensor.dim() == 3 + assert tensor.shape[1:] == (self.ratio, 2 * self.head_dim) + + s, r, d = tensor.size(0), self.ratio, self.head_dim + new_tensor = tensor.new_full((s, 2 * r, d), fill_value) + new_tensor[:, r:] = tensor[:, :, d:] + new_tensor[1:, :r] = tensor[:-1, :, :d] + return new_tensor + + def overlap_transform_decode(self, tensor: torch.Tensor) -> torch.Tensor: + assert tensor.dim() == 3 + assert tensor.shape[1:] == (2 * self.ratio, 2 * self.head_dim) + r, d = self.ratio, self.head_dim + ret = torch.cat((tensor[:, :r, :d], tensor[:, r:, d:]), dim=1) + return ret + + @staticmethod + def compute_state_len(seq_len: int, ratio: int): + return seq_len % ratio + (ratio == 4) * ratio + + @staticmethod + def compute_state_len_indices(seq_len: int, ratio: int): + state_len = seq_len % ratio + (ratio == 4) * ratio + return torch.arange(seq_len - state_len, seq_len).clamp(min=-1) + + def print_tensor(self, y: torch.Tensor, name: str): + enable = int(os.environ.get("SGLANG_ENABLE_PRINT_TENSOR", 0)) + if enable: + print(f"[sgl] {name}: shape={y.shape}, dtype={y.dtype}, device={y.device}") + print(f"{y.flatten()[:10]}...{y.flatten()[-10:]}") + + def compress_extend_paged( + self, + kv_and_scores: KVAndScore, + forward_batch: ForwardBatch, + ): + backend = forward_batch.attn_backend + if TYPE_CHECKING: + assert isinstance(backend, DeepseekV4HipRadixBackend) + token_to_kv_pool = forward_batch.token_to_kv_pool + assert isinstance(token_to_kv_pool, DeepSeekV4TokenToKVPool) + + state_pool = self._get_state_pool(forward_batch) + prefix_lens = forward_batch.extend_prefix_lens_cpu + extend_lens = forward_batch.extend_seq_lens_cpu + req_pool_indices = forward_batch.req_pool_indices + req_to_token = forward_batch.req_to_token_pool.req_to_token + assert not self.forward_mode.is_target_verify() + + assert extend_lens is not None and prefix_lens is not None + device = kv_and_scores.kv.device + + assert kv_and_scores.kv.shape[-1] == self.head_dim * self.coff + compressed_kv_output = torch.full( + (kv_and_scores.kv.size(0), self.head_dim), + fill_value=10000.0, + dtype=kv_and_scores.kv.dtype, + device=device, + ) + + bs = forward_batch.batch_size + pt = 0 + for i in range(bs): + kv_and_score = kv_and_scores[pt : pt + extend_lens[i]] + pre_state_indices = self.compute_state_len_indices( + seq_len=prefix_lens[i], ratio=self.ratio + ).to(device) + raw_loc = torch.where( + pre_state_indices < 0, + -1, + req_to_token[req_pool_indices[i], pre_state_indices], + ) + swa_loc = token_to_kv_pool.translate_loc_from_full_to_swa(raw_loc) + state_loc = state_pool.translate_from_swa_loc_to_state_loc(swa_loc) + pre_kv_state = state_pool.get_state_by_state_loc(state_loc) + kv_and_score_buffer = KVAndScore.cat([pre_kv_state, kv_and_score], dim=0) + valid_kv_len = kv_and_score_buffer.kv.size(0) + + post_state_indices = self.compute_state_len_indices( + seq_len=prefix_lens[i] + extend_lens[i], ratio=self.ratio + ).to(device) + post_state_len = post_state_indices.size(0) + + assert post_state_len <= valid_kv_len + post_raw_loc = torch.where( + post_state_indices < 0, + -1, + req_to_token[req_pool_indices[i], post_state_indices], + ) + post_swa_loc = token_to_kv_pool.translate_loc_from_full_to_swa(post_raw_loc) + post_state_loc = state_pool.translate_from_swa_loc_to_state_loc( + post_swa_loc + ) + post_state_to_set = kv_and_score_buffer[valid_kv_len - post_state_len :] + state_pool.set_state_by_state_loc(post_state_loc, post_state_to_set) + + compress_len = valid_kv_len // self.ratio * self.ratio + if compress_len == 0: + pt += extend_lens[i] + continue + + kv_and_score_to_compress = kv_and_score_buffer[:compress_len].view( + compress_len // self.ratio, self.ratio, -1 + ) + kv_and_score_to_compress.score.add_(self.ape.unsqueeze(0)) + + if self.overlap: + new_kv = self.overlap_transform( + kv_and_score_to_compress.kv, fill_value=0 + ) + new_score = self.overlap_transform( + kv_and_score_to_compress.score, fill_value=float("-inf") + ) + kv_and_score_to_compress = KVAndScore.from_kv_score( + kv=new_kv, score=new_score + ) + del new_kv, new_score + kv_and_score_to_compress = kv_and_score_to_compress[1:] + + if kv_and_score_to_compress.kv.size(0) == 0: + pt += extend_lens[i] + continue + + kv_compressed = ( + kv_and_score_to_compress.kv + * kv_and_score_to_compress.score.softmax(dim=1) + ).sum(dim=1) + + assert kv_compressed.dtype == torch.float32 + + beg_idx = prefix_lens[i] // self.ratio * self.ratio + end_idx = (prefix_lens[i] + extend_lens[i]) // self.ratio * self.ratio + freqs_cis = self.freqs_cis[beg_idx : end_idx : self.ratio] + assert freqs_cis.size(0) == kv_compressed.size( + 0 + ), f"{freqs_cis.shape=} {kv_compressed.shape=}" + if self.use_hip_fused_compress: + fused_norm_rope_inplace_triton( + kv_compressed, self.norm.weight, self.norm.eps, freqs_cis + ) + else: + kv_compressed = self.norm(kv_compressed) + apply_rotary_emb_triton( + kv_compressed[..., -self.rope_head_dim :], freqs_cis + ) + del beg_idx, end_idx + + if self.rotate: + kv_compressed = rotate_activation(kv_compressed) + + start = prefix_lens[i] + start = start + self.ratio - 1 - start % self.ratio + indices_in_seq = torch.arange( + start, + prefix_lens[i] + extend_lens[i], + self.ratio, + device=kv_and_scores.kv.device, + ) + assert indices_in_seq.size(0) == kv_compressed.size(0) + compressed_kv_output[indices_in_seq - prefix_lens[i] + pt] = kv_compressed + + pt += extend_lens[i] + + return compressed_kv_output + + def compress_decode_paged( + self, + kv_and_scores: KVAndScore, + forward_batch: ForwardBatch, + ): + """Paged and cudagraph compatible version of compress_decode""" + assert self.ape_converted + state_pool = self._get_state_pool(forward_batch) + token_to_kv_pool = forward_batch.token_to_kv_pool + assert isinstance(token_to_kv_pool, DeepSeekV4TokenToKVPool) + req_pool_indices = forward_batch.req_pool_indices + req_to_token = forward_batch.req_to_token_pool.req_to_token + seq_lens = forward_batch.seq_lens + + if forward_batch.forward_mode.is_target_verify(): + draft_tokens = forward_batch.attn_backend.speculative_num_draft_tokens + offsets = torch.arange(1, draft_tokens + 1, device=seq_lens.device) + seq_lens_2d = seq_lens[:, None] + offsets[None, :] + seq_lens = seq_lens_2d.view(-1) + req_pool_indices = req_pool_indices.repeat_interleave(draft_tokens) + + raw_locs = req_to_token[req_pool_indices, seq_lens - 1] + + swa_locs = token_to_kv_pool.translate_loc_from_full_to_swa(raw_locs) + state_locs = state_pool.translate_from_swa_loc_to_state_loc(swa_locs) + state_pool.set_state_by_state_loc(state_locs, kv_and_scores) + + compress_bulk_len = self.ratio * self.coff + compress_indices = seq_lens[:, None] + torch.arange( + -compress_bulk_len, 0, device=seq_lens.device + ) + compress_indices.clamp_(min=-1) + compress_indices_raw = torch.where( + compress_indices < 0, + -1, + req_to_token[req_pool_indices[:, None], compress_indices], + ) + compress_indices_swa = token_to_kv_pool.translate_loc_from_full_to_swa( + compress_indices_raw + ) + compress_indices_state = state_pool.translate_from_swa_loc_to_state_loc( + compress_indices_swa + ) + kv_and_score_to_compress = state_pool.get_state_by_state_loc( + compress_indices_state.view(-1) + ).view(-1, self.ratio, self.coff * self.head_dim) + kv_and_score_to_compress.score.add_(self.ape.unsqueeze(0)) + + bs = seq_lens.size(0) + if self.overlap: + kv_and_score_to_compress = kv_and_score_to_compress.view( + bs, self.coff * self.ratio, self.coff * self.head_dim + ) + kv_and_score_to_compress = KVAndScore.from_kv_score( + kv=self.overlap_transform_decode(kv_and_score_to_compress.kv), + score=self.overlap_transform_decode(kv_and_score_to_compress.score), + ) + + self.print_tensor(kv_and_score_to_compress.kv, "kv_to_compress") + self.print_tensor(kv_and_score_to_compress.score, "score_to_compress") + + kv_and_score_to_compress = kv_and_score_to_compress.view( + bs, self.ratio * self.coff, self.head_dim + ) + + kv_compressed = ( + kv_and_score_to_compress.kv * kv_and_score_to_compress.score.softmax(dim=1) + ).sum(dim=1) + self.print_tensor(kv_compressed, "kv_before_norm") + if self.use_hip_fused_compress: + freqs_cis = self._init_freqs_cis_per_decode_step(forward_batch, seq_lens) + fused_norm_rope_inplace_triton( + kv_compressed, self.norm.weight, self.norm.eps, freqs_cis + ) + else: + kv_compressed = self.norm(kv_compressed) + self.print_tensor(kv_compressed, "kv_after_norm") + freqs_cis = self.freqs_cis[(seq_lens - 1) // self.ratio * self.ratio] + self.print_tensor(freqs_cis, "freqs_cis") + apply_rotary_emb_triton( + kv_compressed[..., -self.rope_head_dim :], freqs_cis + ) + self.print_tensor(kv_compressed, "kv_after_rope") + if self.rotate: + kv_compressed = rotate_activation(kv_compressed) + + self.print_tensor(kv_compressed, "compressed_kv_output") + return kv_compressed + + def compress_fused( + self, + kv_score: torch.Tensor, + forward_batch: ForwardBatch, + ) -> torch.Tensor: + backend = forward_batch.attn_backend + if TYPE_CHECKING: + assert isinstance(backend, DeepseekV4HipRadixBackend) + kv_score_buffer = self._get_state_pool(forward_batch) + kv_score_buffer = kv_score_buffer.kv_score_buffer.kv_score + + return backend.forward_compress( + kv_score_buffer=kv_score_buffer, + kv_score_input=kv_score, + ape=self.ape.view(-1, self.head_dim), + head_dim=self.head_dim, + norm=self.norm, + freqs_cis_cache=self.freqs_cis, + rotate=self.rotate, + compress_ratio=self.ratio, + forward_batch=forward_batch, + is_paged=True, + ) + + def compress_dispatch( + self, + kv_score: torch.Tensor, + forward_batch: ForwardBatch, + ) -> torch.Tensor: + if self.use_fused_compress: + return self.compress_fused(kv_score, forward_batch) + + self.compress_decode = self.compress_decode_paged + self.compress_extend = self.compress_extend_paged + kv_and_scores = KVAndScore(kv_score) + + if TYPE_CHECKING: + assert isinstance(kv_and_scores, KVAndScore) + + if ( + forward_batch.forward_mode.is_decode() + or forward_batch.forward_mode.is_target_verify() + ): + result = self.compress_decode( + kv_and_scores=kv_and_scores, + forward_batch=forward_batch, + ) + elif forward_batch.forward_mode.is_extend(): + result = self.compress_extend( + kv_and_scores=kv_and_scores, + forward_batch=forward_batch, + ) + else: + msg = f"Forward mode {forward_batch.forward_mode} not supported in Compressor." + raise NotImplementedError(msg) + + return result + + def _init_freqs_cis_per_decode_step( + self, + forward_batch: ForwardBatch, + seq_lens: torch.Tensor, + ) -> torch.Tensor: + attr = f"freqs_cis_c{self.ratio}" + cached = getattr(forward_batch, attr, None) + if cached is not None: + return cached + decoded = self.freqs_cis[(seq_lens - 1) // self.ratio * self.ratio] + setattr(forward_batch, attr, decoded) + return decoded + + def forward(self, x: torch.Tensor, forward_batch: ForwardBatch) -> torch.Tensor: + if forward_batch.forward_mode.is_idle(): + assert x.shape[0] == 0 + return x.new_empty(0, self.head_dim) + + kv_score = self.compute_kv_score(x, forward_batch) + self.forward_mode = forward_batch.forward_mode + return self.compress_dispatch(kv_score, forward_batch) diff --git a/python/sglang/srt/layers/attention/dsv4/compressor.py b/python/sglang/srt/layers/attention/dsv4/compressor.py index 332f52977..56174f0e5 100644 --- a/python/sglang/srt/layers/attention/dsv4/compressor.py +++ b/python/sglang/srt/layers/attention/dsv4/compressor.py @@ -24,12 +24,16 @@ from sglang.srt.layers.dp_attention import get_attention_cp_size from sglang.srt.layers.layernorm import RMSNorm from sglang.srt.layers.linear import ReplicatedLinear from sglang.srt.layers.utils.cp_utils import cp_all_gather_rerange_output -from sglang.srt.mem_cache.deepseek_v4_compress_state import CompressStatePool +from sglang.srt.mem_cache.deepseek_v4_compress_state import ( + CompressStatePool, +) from sglang.srt.mem_cache.deepseek_v4_memory_pool import DeepSeekV4TokenToKVPool +from sglang.srt.models.deepseek_v2 import _is_hip from sglang.srt.utils import add_prefix if TYPE_CHECKING: from sglang.srt.layers.attention.deepseek_v4_backend import DeepseekV4AttnBackend + from sglang.srt.layers.rotary_embedding import RotaryEmbedding from sglang.srt.model_executor.forward_batch_info import ForwardBatch @@ -293,6 +297,7 @@ class Compressor(nn.Module): head_dim: int, rotate: bool = False, prefix: str = "", + rotary_emb: Optional[RotaryEmbedding] = None, ) -> None: super().__init__() self.layer_id = layer_id @@ -304,7 +309,7 @@ class Compressor(nn.Module): self.ratio = compress_ratio self.overlap = self.ratio == 4 self.rotate = rotate - coff = 1 + self.overlap + self.coff = coff = 1 + self.overlap self.ape = nn.Parameter( torch.empty(self.ratio, coff * self.head_dim, dtype=torch.float32) @@ -321,6 +326,7 @@ class Compressor(nn.Module): self.norm = RMSNorm( self.head_dim, eps=config.rms_norm_eps, weight_dtype=torch.float32 ) + self.rotary_emb = rotary_emb self.freqs_cis = freqs_cis self.ape_converted = False @@ -350,6 +356,8 @@ class Compressor(nn.Module): # NOTE: used by v2 compressor backend def compute_kv_score(self, x: torch.Tensor, forward_batch: ForwardBatch): kv_score = linear_bf16_fp32(x, self.wkv_gate.weight) + + # CUDA path: delegate to backend if nsa_use_prefill_cp(forward_batch): kv_score = cp_all_gather_rerange_output( kv_score, @@ -383,3 +391,9 @@ class Compressor(nn.Module): forward_batch=forward_batch, is_paged=True, ) + + +if _is_hip: + from sglang.srt.layers.attention.dsv4.compress_hip import ( # noqa: F811 + CompressorHip as Compressor, + ) diff --git a/python/sglang/srt/layers/attention/dsv4/indexer.py b/python/sglang/srt/layers/attention/dsv4/indexer.py index 5c538631e..a6a5b9af8 100644 --- a/python/sglang/srt/layers/attention/dsv4/indexer.py +++ b/python/sglang/srt/layers/attention/dsv4/indexer.py @@ -22,7 +22,6 @@ from sglang.srt.state_capturer.indexer_topk import get_global_indexer_capturer from sglang.srt.utils import add_prefix, is_hip if TYPE_CHECKING: - from sglang.srt.layers.attention.deepseek_v4_backend import DeepseekV4AttnBackend from sglang.srt.layers.attention.dsv4.compressor import ( CompressorBackendMixin, ) @@ -100,7 +99,7 @@ def topk_transform_512_pytorch_vectorized( out_raw_indices: Optional[torch.Tensor] = None, ) -> None: - TOPK = 512 + TOPK = out_page_indices.shape[1] batch_size = scores.shape[0] max_seq_len = scores.shape[1] device = scores.device @@ -332,11 +331,6 @@ class C4IndexerBackendMixin: indexer_metadata = metadata.indexer_metadata core_metadata = metadata.core_metadata - from sglang.srt.layers.attention.deepseek_v4_backend import ( - DSV4AttnMetadata, - ) - - assert isinstance(core_metadata, DSV4AttnMetadata) assert isinstance(indexer_metadata, PagedIndexerMetadata) if enable_multi_stream: @@ -374,7 +368,7 @@ class C4IndexerBackendMixin: assert len(weights.shape) == 3 weights = weights.squeeze(2) if envs.SGLANG_OPT_USE_TILELANG_INDEXER.get(): - from sglang.srt.layers.attention.dsv4.tilelang_kernel import ( + from sglang.srt.layers.attention.nsa.tilelang_kernel import ( tilelang_fp8_paged_mqa_logits as fn, ) elif envs.SGLANG_FP8_PAGED_MQA_LOGITS_TORCH.get(): @@ -383,7 +377,8 @@ class C4IndexerBackendMixin: from deep_gemm import fp8_paged_mqa_logits as fn _c4sl = indexer_metadata.c4_seq_lens - if _c4sl.dim() == 1: + _use_tilelang = envs.SGLANG_OPT_USE_TILELANG_INDEXER.get() + if _c4sl.dim() == 1 and not _use_tilelang: _c4sl = _c4sl.unsqueeze(-1) logits = fn( q_fp8, @@ -479,6 +474,7 @@ class C4Indexer(nn.Module): quant_config: Optional[QuantizationConfig] = None, prefix: str = "", alt_streams: Optional[List[torch.cuda.Stream]] = None, + rotary_emb=None, ): super().__init__() self.layer_id = layer_id @@ -486,6 +482,7 @@ class C4Indexer(nn.Module): self.n_heads = config.index_n_heads self.head_dim = config.index_head_dim self.rope_head_dim = config.qk_rope_head_dim + self.index_topk = config.index_topk self.q_lora_rank = config.q_lora_rank self.softmax_scale = self.head_dim**-0.5 self.n_local_heads = self.n_heads @@ -514,7 +511,9 @@ class C4Indexer(nn.Module): head_dim=self.head_dim, rotate=True, prefix=add_prefix("compressor", prefix), + rotary_emb=rotary_emb, ) + self.rotary_emb = rotary_emb self.freqs_cis = freqs_cis self.weight_scale: float = self.softmax_scale * self.n_heads**-0.5 self.alt_streams = alt_streams @@ -545,8 +544,6 @@ class C4Indexer(nn.Module): enable_multi_stream: bool = False, q_lora_ready: Optional[torch.cuda.Event] = None, ) -> None: - if TYPE_CHECKING: - assert isinstance(forward_batch.attn_backend, DeepseekV4AttnBackend) return forward_batch.attn_backend.forward_c4_indexer( x=x, q_lora=q_lora, diff --git a/python/sglang/srt/layers/attention/hip_flash_mla.py b/python/sglang/srt/layers/attention/hip_flash_mla.py new file mode 100644 index 000000000..8f1fbd117 --- /dev/null +++ b/python/sglang/srt/layers/attention/hip_flash_mla.py @@ -0,0 +1,197 @@ +from typing import Any, Optional + +import torch + +from sglang.srt.layers.quantization.fp8_kernel import is_fp8_fnuz +from sglang.srt.utils import is_hip + +FP8_DTYPE = torch.float8_e4m3fnuz if is_fp8_fnuz() else torch.float8_e4m3fn + + +def flash_mla_with_kvcache_entrypoint(backend: str, **kwargs): + if is_hip(): + import os + + from sglang.srt.layers.attention.nsa.tilelang_kernel import ( + dpsk_v4_fp8_attention_fwd, + ) + + backend = os.environ.get("SGLANG_HACK_FLASHMLA_BACKEND", "tilelang") + else: + import flash_mla + + if backend == "comparison": + pack_ref, pack_fast_via_tester = flash_mla_with_kvcache_entrypoint( + backend="torch", **kwargs + ) + pack_fast_via_api = flash_mla_with_kvcache_entrypoint( + backend="kernel", **kwargs + ) + _assert_close(pack_ref=pack_fast_via_tester, pack_fast=pack_fast_via_api) + _assert_close(pack_ref=pack_ref, pack_fast=pack_fast_via_tester) + _assert_close(pack_ref=pack_ref, pack_fast=pack_fast_via_api) + return pack_ref + + if backend == "torch": + return flash_mla_with_kvcache_torch(**kwargs) + + if backend == "tilelang": + return dpsk_v4_fp8_attention_fwd(**kwargs) + + if backend == "kernel": + return flash_mla.flash_mla_with_kvcache(**kwargs) + + raise NotImplementedError(f"unknown backend: {backend!r}") + + +def flash_mla_with_kvcache_torch( + q: torch.Tensor, + k_cache: torch.Tensor, + block_table: Optional[torch.Tensor], + cache_seqlens: Optional[torch.Tensor], + head_dim_v: int, + tile_scheduler_metadata: Any, + num_splits: None = None, + softmax_scale: Optional[float] = None, + causal: bool = False, + is_fp8_kvcache: bool = False, + indices: Optional[torch.Tensor] = None, + attn_sink: Optional[torch.Tensor] = None, + extra_k_cache: Optional[torch.Tensor] = None, + extra_indices_in_kvcache: Optional[torch.Tensor] = None, + topk_length: Optional[torch.Tensor] = None, + extra_topk_length: Optional[torch.Tensor] = None, +): + + from sglang.srt.flashmla_tests import quant as flashmla_quant + from sglang.srt.flashmla_tests.lib import ( + ExtraTestParamForDecode, + KVScope, + TestcaseForDecode, + TestParam, + ) + from sglang.srt.flashmla_tests.ref import ref_sparse_attn_decode + + assert block_table is None + assert cache_seqlens is None + assert is_fp8_kvcache + + b, s_q, h_q, d_qk = q.shape + d_v = head_dim_v + + fp8_layout = flashmla_quant.FP8KVCacheLayout.MODEL1_FP8Sparse + + p = TestParam( + s_q=s_q, + s_kv="unused", + topk="unused", + h_q=h_q, + h_kv=1, + d_qk=d_qk, + d_v=d_v, + decode=ExtraTestParamForDecode( + b=b, + is_varlen="unused", + have_zero_seqlen_k="unused", + extra_s_k="unused", + extra_topk="unused", + extra_block_size="unused", + have_extra_topk_length="unused", + ), + # unused? + seed=-1, + check_correctness=True, + is_all_indices_invalid=False, + num_runs=10, + have_attn_sink=True, + have_topk_length=True, + ) + + blocked_k_quantized = k_cache + blocked_k = flashmla_quant.dequantize_k_cache( + blocked_k_quantized.view(FP8_DTYPE), fp8_layout + ) + # blocked_k_requantized = flashmla_quant.quantize_k_cache(blocked_k, fp8_layout) + # assert torch.testing.assert_allclose(blocked_k_requantized.byte(), blocked_k_quantized.byte()) + kv_scope = KVScope( + t="unused", + cache_seqlens="unused", + block_table="unused", + blocked_k=blocked_k, + blocked_k_quantized=blocked_k_quantized, + abs_indices="unused", + indices_in_kvcache=indices, + topk_length=topk_length, + ) + + extra_kv_scope = None + if extra_k_cache is not None: + extra_blocked_k_quantized = extra_k_cache + extra_blocked_k = flashmla_quant.dequantize_k_cache( + extra_blocked_k_quantized.view(FP8_DTYPE), fp8_layout + ) + # extra_blocked_k_requantized = flashmla_quant.quantize_k_cache(extra_blocked_k, fp8_layout) + # assert torch.testing.assert_allclose(extra_blocked_k_requantized.byte(), extra_blocked_k_quantized.byte()) + extra_kv_scope = KVScope( + t="unused", + cache_seqlens="unused", + block_table="unused", + blocked_k=extra_blocked_k, + blocked_k_quantized=extra_blocked_k_quantized, + abs_indices="unused", + indices_in_kvcache=extra_indices_in_kvcache, + topk_length=extra_topk_length, + ) + + t = TestcaseForDecode( + p="unused", + q=q, + attn_sink=attn_sink, + sm_scale=softmax_scale, + kv_scope=kv_scope, + extra_kv_scope=extra_kv_scope, + ) + # print(f"hi {p=} {t=}") + # print( + # f"hi info " + # f"{get_tensor_info(t.kv_scope.blocked_k)=} " + # f"{get_tensor_info(t.kv_scope.blocked_k_quantized)=} " + # f"{get_tensor_info(t.extra_kv_scope.blocked_k) if t.extra_kv_scope is not None else None=} " + # f"{get_tensor_info(t.extra_kv_scope.blocked_k_quantized) if t.extra_kv_scope is not None else None=} " + # ) + + pack_ref = ref_sparse_attn_decode(p, t) + + # tile_scheduler_metadata, _ = flash_mla.get_mla_metadata() + # pack_fast_via_tester = flashmla_lib.run_flash_mla_decode( + # p, t, tile_scheduler_metadata, num_splits=None + # ) + + # return pack_ref, pack_fast_via_tester + return pack_ref + + +def _assert_close(pack_ref, pack_fast): + import sglang.srt.flashmla_tests.kernelkit as kk + + out_ref, lse_ref = pack_ref + out_fast, lse_fast = pack_fast + + # the copied threshold is too strict, not checked why + # copied from: test_flash_mla_sparse_decoding.py + # is_out_correct = kk.check_is_allclose( + # "out", out_fast, out_ref, abs_tol=1e-3, rel_tol=2.01 / 128, cos_diff_tol=5e-6 + # ) + # is_lse_correct = kk.check_is_allclose( + # "lse", lse_fast, lse_ref, abs_tol=1e-6, rel_tol=8.01 / 65536 + # ) + + # loosen thresh + is_out_correct = kk.check_is_allclose( + "out", out_fast, out_ref, abs_tol=1e-2, rel_tol=10.0, cos_diff_tol=5e-6 + ) + is_lse_correct = kk.check_is_allclose( + "lse", lse_fast, lse_ref, abs_tol=1e-6, rel_tol=8.01 / 65536 + ) + + assert is_out_correct and is_lse_correct, f"{is_out_correct=} {is_lse_correct=}" diff --git a/python/sglang/srt/layers/attention/nsa/index_buf_accessor.py b/python/sglang/srt/layers/attention/nsa/index_buf_accessor.py index b5b97147e..db1e80ee3 100644 --- a/python/sglang/srt/layers/attention/nsa/index_buf_accessor.py +++ b/python/sglang/srt/layers/attention/nsa/index_buf_accessor.py @@ -432,10 +432,6 @@ def _set_k_and_s_triton( assert ( page_size % 16 == 0 ), f"HIP preshuffle requires page_size to be a multiple of 16, got {page_size}" - else: - assert ( - page_size == 1 - ), f"HIP legacy NSA path requires page_size == 1, got {page_size}" else: assert page_size == 64 diff --git a/python/sglang/srt/layers/attention/nsa/tilelang_kernel.py b/python/sglang/srt/layers/attention/nsa/tilelang_kernel.py index bfc62d7f0..62509c308 100644 --- a/python/sglang/srt/layers/attention/nsa/tilelang_kernel.py +++ b/python/sglang/srt/layers/attention/nsa/tilelang_kernel.py @@ -1,5 +1,6 @@ +import functools from functools import lru_cache -from typing import Optional, Tuple +from typing import Any, Optional, Tuple import tilelang import tilelang.language as T @@ -10,6 +11,28 @@ from sglang.srt.utils import is_gfx95_supported, is_hip tilelang.set_log_level("WARNING") +# Workaround a tilelang bug: BaseKernelAdapter._legalize_result_idx mutates the +# `out_idx` list in place when normalising negative indices to positive ones. +# That breaks any @tilelang.jit factory that compiles two prim_funcs with +# different param counts (e.g. our unified single/dual partial kernel) — the +# second compile sees indices already-converted for the first's len(params) +# and silently builds the wrong adapter, leading to IndexError at call time. +# Patch once on import to copy the list before mutation. +from tilelang.jit.adapter.base import ( # noqa: E402 + BaseKernelAdapter as _BaseKernelAdapter, +) + +if not getattr(_BaseKernelAdapter, "_legalize_result_idx_patched", False): + _orig_legalize = _BaseKernelAdapter._legalize_result_idx + + def _legalize_result_idx_safe(self, result_idx): + if isinstance(result_idx, list): + result_idx = list(result_idx) + return _orig_legalize(self, result_idx) + + _BaseKernelAdapter._legalize_result_idx = _legalize_result_idx_safe + _BaseKernelAdapter._legalize_result_idx_patched = True + pass_configs = { tilelang.PassConfigKey.TL_DISABLE_WARP_SPECIALIZED: True, tilelang.PassConfigKey.TL_DISABLE_TMA_LOWER: True, @@ -25,8 +48,11 @@ _is_gfx95_supported = is_gfx95_supported() _is_fp8_fnuz = is_fp8_fnuz() BF16 = "bfloat16" -FP8 = "float8_e4m3fnuz" if _is_fp8_fnuz else "float8_e4m3" +FP8 = "float8_e4m3fnuz" if _is_fp8_fnuz else "float8_e4m3fn" +FP8_DTYPE = torch.float8_e4m3fnuz if _is_fp8_fnuz else torch.float8_e4m3fn FP32 = "float32" +INT32 = "int32" +UINT8 = "uint8" def fast_log2_ceil(x): @@ -1375,3 +1401,1189 @@ def tilelang_sparse_fwd( ) out = kernel(q.unsqueeze(0), kv.unsqueeze(0), indices.unsqueeze(0)) # type: ignore return out + + +@functools.cache +def fp8_paged_mqa_logits_kernel( + head_dim: int = 128, + num_heads: int = 64, + block_size: int = 64, + clear_accum: bool = True, + split_kv: int = 1, +) -> Any: + N = T.symbolic("batch_size") + L = T.symbolic("max_table_length") + S = T.symbolic("max_seq_len") + C = T.symbolic("num_blocks") + B = block_size + D = head_dim + H = num_heads + SK = int(split_kv) + BLOCK_BYTES = B * (D + 4) + SCALE_OFFSET = B * D + + assert D % 4 == 0 + assert H % 4 == 0 + assert D == 128 + assert SK >= 1 + + @tilelang.jit( + pass_configs={ + **pass_configs, + tilelang.PassConfigKey.TL_DISABLE_SAFE_MEMORY_ACCESS: True, + } + ) + def fp8_paged_mqa_logits( + q: T.Tensor[(N, H, D), FP8], + kvcache_u8: T.Tensor[(C, BLOCK_BYTES), UINT8], + weight: T.Tensor[(N, H), FP32], + seq_lens: T.Tensor[(N,), INT32], + page_table: T.Tensor[(N, L), INT32], + o: T.Tensor[(N, S), FP32], + ) -> None: + _ = N, L, S, C, D, H, B + with T.Kernel(N * SK) as bxs: + bx = bxs % N + pid_split = bxs // N + seq_len = seq_lens[bx] + np_total = T.ceildiv(seq_len, B) + stride = T.ceildiv(np_total, SK) + i_start = pid_split * stride + n_iters = T.max(0, T.min(stride, np_total - i_start)) + + q_smem = T.alloc_shared((H, D), FP8) + q_s_frag = T.alloc_fragment((H,), FP32) + T.copy(q[bx, 0, 0], q_smem) + T.copy(weight[bx, 0], q_s_frag) + + for j in T.Pipelined(n_iters, num_stages=2): + i = i_start + j + page = page_table[bx, i] + k_smem_u8 = T.alloc_shared((B * D,), UINT8) + T.copy(kvcache_u8[page, 0:SCALE_OFFSET], k_smem_u8) + k_smem = T.view(k_smem_u8, (B, D), FP8) + k_s_smem_u8 = T.alloc_shared((B * 4,), UINT8) + T.copy(kvcache_u8[page, SCALE_OFFSET:BLOCK_BYTES], k_s_smem_u8) + k_s_smem = T.view(k_s_smem_u8, (B,), FP32) + k_s_frag = T.alloc_fragment((B,), FP32) + T.copy(k_s_smem, k_s_frag) + + logits = T.alloc_fragment((B, H), FP32) + if not clear_accum: + T.fill(logits, 0.0) + T.gemm( + k_smem, + q_smem, + logits, + transpose_A=False, + transpose_B=True, + clear_accum=clear_accum, + ) + + # post processing + for h, j2 in T.Parallel(H, B): + logits[j2, h] = T.max(logits[j2, h], 0.0) * q_s_frag[h] + logits_sum = T.alloc_fragment((B,), FP32) + T.reduce_sum(logits, logits_sum, dim=1) + for j2 in T.Parallel(B): + logits_sum[j2] *= k_s_frag[j2] + T.copy(logits_sum, o[bx, i * B]) + + return fp8_paged_mqa_logits + + +def tilelang_fp8_paged_mqa_logits( + q_fp8: torch.Tensor, + kvcache_fp8: torch.Tensor, + weight: torch.Tensor, + seq_lens: torch.Tensor, + page_table: torch.Tensor, + deep_gemm_metadata: Any, + max_seq_len: int, + clean_logits: bool = True, +) -> torch.Tensor: + _ = deep_gemm_metadata + batch_size, _, num_heads, head_dim = q_fp8.shape + block_size = kvcache_fp8.shape[1] + assert head_dim == 128, "TODO" + assert block_size == 64, "TODO" + assert q_fp8.shape == (batch_size, 1, num_heads, head_dim) + assert kvcache_fp8.shape[1:] == (block_size, 1, head_dim + 4) + assert weight.shape == (batch_size, num_heads) + assert seq_lens.shape == (batch_size,) + assert page_table.shape[0] == batch_size + assert clean_logits == False + + logits = page_table.new_empty((batch_size, max_seq_len), dtype=torch.float32) + + NUM_CU = 256 + split_kv = split_kv = max(1, min(max_seq_len // block_size, NUM_CU // batch_size)) + kernel = fp8_paged_mqa_logits_kernel( + head_dim=head_dim, + num_heads=num_heads, + block_size=block_size, + clear_accum=clean_logits, + split_kv=split_kv, + ) + q_fp8 = q_fp8.view(batch_size, num_heads, head_dim) + kvcache_u8 = kvcache_fp8.view(-1, block_size * (head_dim + 4)) + kernel(q_fp8, kvcache_u8, weight, seq_lens, page_table, logits) + return logits + + +def _build_fp8_combined_view(k_cache: torch.Tensor) -> Tuple[torch.Tensor, int, int]: + """ + Reinterpret a MODEL1_FP8Sparse KV cache as a contiguous uint32 view. + Input: k_cache (num_blocks, block_size, 1, d_qk) fp8/uint8 + — per-block storage also holds scales + padding past d_qk. + Output: (num_blocks, block_pad_u32) uint32 covering the full block + stride. Same storage ashe input, no copy. + """ + k_u8 = k_cache.view(torch.uint8) if k_cache.dtype != torch.uint8 else k_cache + num_blocks = k_u8.shape[0] + block_size = k_u8.shape[1] + block_pad_u32 = k_u8.stride(0) // 4 + storage = k_u8.untyped_storage() + flat_u32 = torch.empty(0, dtype=torch.uint32, device=k_u8.device).set_( + storage, 0, (storage.nbytes() // 4,), (1,) + ) + k_combined = torch.as_strided( + flat_u32, + size=(num_blocks, block_pad_u32), + stride=(block_pad_u32, 1), + storage_offset=k_u8.storage_offset() // 4, + ) + return k_combined, num_blocks, block_size + + +_TOPK_LEN_SENTINEL_CACHE: dict = {} +_INT32_MAX = 2**30 + + +def _topk_length_sentinel(device: torch.device, batch: int) -> torch.Tensor: + """Cached `(batch,) int32 INT_MAX` tensor used when `topk_length` is None.""" + cur = _TOPK_LEN_SENTINEL_CACHE.get(device) + if cur is None or cur.numel() < batch: + cur = torch.full( + (max(batch, 256),), _INT32_MAX, dtype=torch.int32, device=device + ) + _TOPK_LEN_SENTINEL_CACHE[device] = cur + return cur[:batch] + + +@tilelang.jit( + out_idx=[-2, -1], + pass_configs={ + tilelang.PassConfigKey.TL_DISABLE_TMA_LOWER: True, + tilelang.PassConfigKey.TL_DISABLE_WARP_SPECIALIZED: True, + tilelang.PassConfigKey.TL_ENABLE_FAST_MATH: True, + }, +) +def dpsk_v4_fp8_partial_kernel( + num_heads: int, + topk_1: int, + block_size_kv_1: int, + topk_2: int = 0, + block_size_kv_2: int = 0, + *, + dim: int = 448, + tail_dim: int = 64, + sm_scale: float = 0.0, + block_I: int = 64, + inner_iter_1: int = 1, + inner_iter_2: int = 0, + num_stages: int = 0, + threads: int = 512, +) -> Any: + """ + Read FP8 K cache directly, dequantise to BF16 in-kernel, do flash-attn + online softmax with split-K. Supports a second cache (`topk_2>0`) and + `attn_sink` is folded later by the combine kernel. + """ + log2e: float = 1.44269504 + if sm_scale <= 0.0: + sm_scale = (1.0 / (dim + tail_dim)) ** 0.5 * log2e + else: + sm_scale = sm_scale * log2e + assert dim == 448 and tail_dim == 64 + assert topk_1 % block_I == 0 + assert ( + topk_1 // block_I + ) % inner_iter_1 == 0, ( + f"NI_1={topk_1 // block_I} must be divisible by inner_iter_1={inner_iter_1}" + ) + assert block_size_kv_1 > 0 and (block_size_kv_1 & (block_size_kv_1 - 1)) == 0 + + is_dual = topk_2 > 0 + if is_dual: + assert inner_iter_2 > 0, "dual-cache call requires inner_iter_2 > 0" + assert topk_2 % block_I == 0 + assert ( + topk_2 // block_I + ) % inner_iter_2 == 0, ( + f"NI_2={topk_2 // block_I} must be divisible by inner_iter_2={inner_iter_2}" + ) + assert block_size_kv_2 > 0 and (block_size_kv_2 & (block_size_kv_2 - 1)) == 0 + + PACKED_W = dim + 2 * tail_dim + NOPE_TILE = 64 + NUM_TILES = dim // NOPE_TILE + SCALE_W = 8 + PACKED_W4 = PACKED_W // 4 + SCALE_W4 = SCALE_W // 4 + + kv_group = 1 + batch = T.symbolic("batch") + seq_len = T.symbolic("seq_len") + num_blocks_kv_1 = T.symbolic("num_blocks_kv_1") + block_pad_u32_1 = T.symbolic("block_pad_u32_1") + if is_dual: + num_blocks_kv_2 = T.symbolic("num_blocks_kv_2") + block_pad_u32_2 = T.symbolic("block_pad_u32_2") + + head_kv = num_heads // kv_group + D = dim + D_tail = tail_dim + BI = block_I + padded_H = max(tilelang.math.next_power_of_2(head_kv), 16) + if head_kv > 64: + assert head_kv % 64 == 0 + REPLICATE_H = (head_kv + 63) // 64 if head_kv > 64 else 1 + H_per_block = 64 if REPLICATE_H > 1 else padded_H + + NI_1 = topk_1 // BI + n_groups_1 = NI_1 // inner_iter_1 + NI_2 = (topk_2 // BI) if is_dual else 0 + n_groups_2 = (NI_2 // inner_iter_2) if is_dual else 0 + n_groups = n_groups_1 + n_groups_2 + + BS_KV_1 = block_size_kv_1 + NOPE_ROPE_U32_PER_BLOCK_1 = BS_KV_1 * PACKED_W4 + if is_dual: + BS_KV_2 = block_size_kv_2 + NOPE_ROPE_U32_PER_BLOCK_2 = BS_KV_2 * PACKED_W4 + + q_shape = [batch, seq_len, num_heads, D + D_tail] + k1_shape = [num_blocks_kv_1, block_pad_u32_1] + indices1_shape = [batch, seq_len, topk_1] + topk_length_shape = [batch] + partial_o_shape = [batch, seq_len, n_groups, num_heads, D + D_tail] + partial_lse_shape = [batch, seq_len, n_groups, num_heads] + if is_dual: + k2_shape = [num_blocks_kv_2, block_pad_u32_2] + indices2_shape = [batch, seq_len, topk_2] + + accum_dtype = "float" + indices_dtype = INT32 + + if is_dual: + + @T.prim_func + def main( + Q: T.Tensor(q_shape, BF16), # type: ignore + K_combined_1: T.Tensor(k1_shape, "uint32"), # type: ignore + Indices_1: T.Tensor(indices1_shape, indices_dtype), # type: ignore + Topk_length_1: T.Tensor(topk_length_shape, indices_dtype), # type: ignore + K_combined_2: T.Tensor(k2_shape, "uint32"), # type: ignore + Indices_2: T.Tensor(indices2_shape, indices_dtype), # type: ignore + Topk_length_2: T.Tensor(topk_length_shape, indices_dtype), # type: ignore + Partial_O: T.Tensor(partial_o_shape, BF16), # type: ignore + Partial_LSE: T.Tensor(partial_lse_shape, accum_dtype), # type: ignore + ) -> None: + """ + grid: (seq_len * REPLICATE_H * n_groups, batch, 1) + Each block processes `inner_iter_1` (or `inner_iter_2`) consecutive + KV tiles of one phase and writes one (partial_o, partial_lse) entry. + """ + with T.Kernel( + seq_len * REPLICATE_H * n_groups, batch, kv_group, threads=threads + ) as (bx, by, bz): + Q_shared = T.alloc_fragment([H_per_block, D], BF16) + Q_tail_shared = T.alloc_fragment([H_per_block, D_tail], BF16) + K_packed_shared = T.alloc_shared([BI, PACKED_W4], "uint32") + K_scale_shared = T.alloc_shared([BI, SCALE_W4], "uint32") + KV_shared = T.alloc_shared([BI, D], BF16) + K_tail_shared = T.alloc_shared([BI, D_tail], BF16) + S_shared = T.alloc_shared([H_per_block, BI], BF16) + page_idx_shared = T.alloc_shared([BI], INT32) + + mask = T.alloc_fragment([BI], "bool") + scale_byte_local = T.alloc_fragment([BI, NUM_TILES], "uint32") + + acc_o = T.alloc_fragment([H_per_block, D], accum_dtype) + acc_o_tail = T.alloc_fragment([H_per_block, D_tail], accum_dtype) + acc_s = T.alloc_fragment([H_per_block, BI], accum_dtype) + sumexp = T.alloc_fragment([H_per_block], accum_dtype) + sumexp_i = T.alloc_fragment([H_per_block], accum_dtype) + alpha = T.alloc_fragment([H_per_block], accum_dtype) + m_i = T.alloc_fragment([H_per_block], accum_dtype) + m_i_prev = T.alloc_fragment([H_per_block], accum_dtype) + + T.fill(acc_o, 0) + T.fill(acc_o_tail, 0) + T.fill(sumexp, 0) + T.fill(m_i, -(2**30)) + + b_i, g_i = by, bz + # bx encodes (s_i, h_replicate, group_i). + spans_per_seq = REPLICATE_H * n_groups + s_i = bx // spans_per_seq + rest = bx % spans_per_seq + group_i = rest // REPLICATE_H + h_rep = rest % REPLICATE_H + H0 = g_i * padded_H + (0 if REPLICATE_H == 1 else h_rep * 64) + H1 = H0 + H_per_block + + tk_len_1 = Topk_length_1[b_i] + tk_len_2 = Topk_length_2[b_i] + actual_n_groups_1 = T.ceildiv(tk_len_1, BI * inner_iter_1) + actual_n_groups_2 = T.ceildiv(tk_len_2, BI * inner_iter_2) + + if (group_i < n_groups_1) & (group_i < actual_n_groups_1): + # Phase 1 active: SWA cache work + Partial_O write. + T.copy(Q[b_i, s_i, H0:H1, :D], Q_shared) + T.copy(Q[b_i, s_i, H0:H1, D : D + D_tail], Q_tail_shared) + for k_i in T.Pipelined(inner_iter_1, num_stages=num_stages): + iter_i = group_i * inner_iter_1 + k_i + for bi_i in T.Parallel(BI): + pos = iter_i * BI + bi_i + idx = Indices_1[b_i, s_i, pos] + valid = (idx >= 0) & (pos < tk_len_1) + page_idx_shared[bi_i] = T.if_then_else(valid, idx, 0) + mask[bi_i] = valid + + for bi_i, w_i in T.Parallel(BI, PACKED_W4): + page = page_idx_shared[bi_i] + block_id = page // BS_KV_1 + t_in_block = page % BS_KV_1 + K_packed_shared[bi_i, w_i] = K_combined_1[ + block_id, t_in_block * PACKED_W4 + w_i + ] + + for bi_i, w_i in T.Parallel(BI, SCALE_W4): + page = page_idx_shared[bi_i] + block_id = page // BS_KV_1 + t_in_block = page % BS_KV_1 + K_scale_shared[bi_i, w_i] = K_combined_1[ + block_id, + NOPE_ROPE_U32_PER_BLOCK_1 + t_in_block * SCALE_W4 + w_i, + ] + + for bi_i, ti in T.Parallel(BI, NUM_TILES): + word_idx = ti // 4 + byte_in_word = ti % 4 + word = K_scale_shared[bi_i, word_idx] + scale_byte_local[bi_i, ti] = ( + word >> T.Cast("uint32", byte_in_word * 8) + ) & T.uint32(0xFF) + + for bi_i, d_i in T.Parallel(BI, D): + word_idx = d_i // 4 + byte_in_word = d_i % 4 + word = K_packed_shared[bi_i, word_idx] + b_u32 = ( + word >> T.Cast("uint32", byte_in_word * 8) + ) & T.uint32(0xFF) + sign_bf = (b_u32 & T.uint32(0x80)) * T.uint32(0x100) + exp_e4 = (b_u32 & T.uint32(0x78)) >> T.uint32(3) + mant_bf = (b_u32 & T.uint32(0x7)) * T.uint32(0x10) + scale_byte = scale_byte_local[bi_i, d_i // NOPE_TILE] + exp_combined = exp_e4 + scale_byte - T.uint32(7) + bf16_bits = ( + sign_bf | (exp_combined << T.uint32(7)) | mant_bf + ) + KV_shared[bi_i, d_i] = T.reinterpret( + BF16, T.Cast("uint16", bf16_bits) + ) + + for bi_i, j in T.Parallel(BI, D_tail): + abs_off = D + 2 * j + word_idx = abs_off // 4 + word_off = abs_off % 4 + word = K_packed_shared[bi_i, word_idx] + half_u32 = T.if_then_else( + word_off == 0, + word & T.uint32(0xFFFF), + (word >> T.uint32(16)) & T.uint32(0xFFFF), + ) + K_tail_shared[bi_i, j] = T.reinterpret( + BF16, T.Cast("uint16", half_u32) + ) + + for h_i, bi_i in T.Parallel(H_per_block, BI): + acc_s[h_i, bi_i] = T.if_then_else( + mask[bi_i], 0, -T.infinity(acc_s.dtype) + ) + T.gemm( + Q_shared, + KV_shared, + acc_s, + transpose_B=True, + policy=T.GemmWarpPolicy.FullRow, + ) + T.gemm( + Q_tail_shared, + K_tail_shared, + acc_s, + transpose_B=True, + policy=T.GemmWarpPolicy.FullRow, + ) + T.copy(m_i, m_i_prev) + T.reduce_max(acc_s, m_i, dim=1, clear=False) + for h_i in T.Parallel(H_per_block): + m_i[h_i] = T.max(m_i[h_i], m_i_prev[h_i]) + for h_i in T.Parallel(H_per_block): + alpha[h_i] = T.exp2((m_i_prev[h_i] - m_i[h_i]) * sm_scale) + for h_i, bi_i in T.Parallel(H_per_block, BI): + acc_s[h_i, bi_i] = T.exp2( + acc_s[h_i, bi_i] * sm_scale - m_i[h_i] * sm_scale + ) + T.reduce_sum(acc_s, sumexp_i, dim=1) + for h_i in T.Parallel(H_per_block): + sumexp[h_i] = sumexp[h_i] * alpha[h_i] + sumexp_i[h_i] + for h_i, d_i in T.Parallel(H_per_block, D): + acc_o[h_i, d_i] *= alpha[h_i] + for h_i, d_i in T.Parallel(H_per_block, D_tail): + acc_o_tail[h_i, d_i] *= alpha[h_i] + T.copy(acc_s, S_shared) + T.gemm( + S_shared, + KV_shared, + acc_o, + policy=T.GemmWarpPolicy.FullRow, + ) + T.gemm( + S_shared, + K_tail_shared, + acc_o_tail, + policy=T.GemmWarpPolicy.FullRow, + ) + # ---- finalize phase 1 (active) ---- + for h_i, d_i in T.Parallel(H_per_block, D): + acc_o[h_i, d_i] = acc_o[h_i, d_i] / T.if_then_else( + sumexp[h_i] == 0.0, 1.0, sumexp[h_i] + ) + for h_i, d_i in T.Parallel(H_per_block, D_tail): + acc_o_tail[h_i, d_i] = acc_o_tail[h_i, d_i] / T.if_then_else( + sumexp[h_i] == 0.0, 1.0, sumexp[h_i] + ) + for h_i in T.Parallel(H_per_block): + m_i[h_i] = T.if_then_else( + sumexp[h_i] == 0.0, + -(2.0**30), + T.log2(sumexp[h_i]) + m_i[h_i] * sm_scale, + ) + T.copy(acc_o, Partial_O[b_i, s_i, group_i, H0:H1, :D]) + T.copy( + acc_o_tail, + Partial_O[b_i, s_i, group_i, H0:H1, D : D + D_tail], + ) + T.copy(m_i, Partial_LSE[b_i, s_i, group_i, H0:H1]) + elif group_i < n_groups_1: + # Phase 1 skipped: m_i is still the -2^30 + T.copy(m_i, Partial_LSE[b_i, s_i, group_i, H0:H1]) + elif (group_i - n_groups_1) < actual_n_groups_2: + # Phase 2 active: c128 cache work + Partial_O write. + T.copy(Q[b_i, s_i, H0:H1, :D], Q_shared) + T.copy(Q[b_i, s_i, H0:H1, D : D + D_tail], Q_tail_shared) + for k_i in T.Pipelined(inner_iter_2, num_stages=num_stages): + iter_i = (group_i - n_groups_1) * inner_iter_2 + k_i + for bi_i in T.Parallel(BI): + pos = iter_i * BI + bi_i + idx = Indices_2[b_i, s_i, pos] + valid = (idx >= 0) & (pos < tk_len_2) + page_idx_shared[bi_i] = T.if_then_else(valid, idx, 0) + mask[bi_i] = valid + + for bi_i, w_i in T.Parallel(BI, PACKED_W4): + page = page_idx_shared[bi_i] + block_id = page // BS_KV_2 + t_in_block = page % BS_KV_2 + K_packed_shared[bi_i, w_i] = K_combined_2[ + block_id, t_in_block * PACKED_W4 + w_i + ] + + for bi_i, w_i in T.Parallel(BI, SCALE_W4): + page = page_idx_shared[bi_i] + block_id = page // BS_KV_2 + t_in_block = page % BS_KV_2 + K_scale_shared[bi_i, w_i] = K_combined_2[ + block_id, + NOPE_ROPE_U32_PER_BLOCK_2 + t_in_block * SCALE_W4 + w_i, + ] + + for bi_i, ti in T.Parallel(BI, NUM_TILES): + word_idx = ti // 4 + byte_in_word = ti % 4 + word = K_scale_shared[bi_i, word_idx] + scale_byte_local[bi_i, ti] = ( + word >> T.Cast("uint32", byte_in_word * 8) + ) & T.uint32(0xFF) + + for bi_i, d_i in T.Parallel(BI, D): + word_idx = d_i // 4 + byte_in_word = d_i % 4 + word = K_packed_shared[bi_i, word_idx] + b_u32 = ( + word >> T.Cast("uint32", byte_in_word * 8) + ) & T.uint32(0xFF) + sign_bf = (b_u32 & T.uint32(0x80)) * T.uint32(0x100) + exp_e4 = (b_u32 & T.uint32(0x78)) >> T.uint32(3) + mant_bf = (b_u32 & T.uint32(0x7)) * T.uint32(0x10) + scale_byte = scale_byte_local[bi_i, d_i // NOPE_TILE] + exp_combined = exp_e4 + scale_byte - T.uint32(7) + bf16_bits = ( + sign_bf | (exp_combined << T.uint32(7)) | mant_bf + ) + KV_shared[bi_i, d_i] = T.reinterpret( + BF16, T.Cast("uint16", bf16_bits) + ) + + for bi_i, j in T.Parallel(BI, D_tail): + abs_off = D + 2 * j + word_idx = abs_off // 4 + word_off = abs_off % 4 + word = K_packed_shared[bi_i, word_idx] + half_u32 = T.if_then_else( + word_off == 0, + word & T.uint32(0xFFFF), + (word >> T.uint32(16)) & T.uint32(0xFFFF), + ) + K_tail_shared[bi_i, j] = T.reinterpret( + BF16, T.Cast("uint16", half_u32) + ) + + for h_i, bi_i in T.Parallel(H_per_block, BI): + acc_s[h_i, bi_i] = T.if_then_else( + mask[bi_i], 0, -T.infinity(acc_s.dtype) + ) + T.gemm( + Q_shared, + KV_shared, + acc_s, + transpose_B=True, + policy=T.GemmWarpPolicy.FullRow, + ) + T.gemm( + Q_tail_shared, + K_tail_shared, + acc_s, + transpose_B=True, + policy=T.GemmWarpPolicy.FullRow, + ) + T.copy(m_i, m_i_prev) + T.reduce_max(acc_s, m_i, dim=1, clear=False) + for h_i in T.Parallel(H_per_block): + m_i[h_i] = T.max(m_i[h_i], m_i_prev[h_i]) + for h_i in T.Parallel(H_per_block): + alpha[h_i] = T.exp2((m_i_prev[h_i] - m_i[h_i]) * sm_scale) + for h_i, bi_i in T.Parallel(H_per_block, BI): + acc_s[h_i, bi_i] = T.exp2( + acc_s[h_i, bi_i] * sm_scale - m_i[h_i] * sm_scale + ) + T.reduce_sum(acc_s, sumexp_i, dim=1) + for h_i in T.Parallel(H_per_block): + sumexp[h_i] = sumexp[h_i] * alpha[h_i] + sumexp_i[h_i] + for h_i, d_i in T.Parallel(H_per_block, D): + acc_o[h_i, d_i] *= alpha[h_i] + for h_i, d_i in T.Parallel(H_per_block, D_tail): + acc_o_tail[h_i, d_i] *= alpha[h_i] + T.copy(acc_s, S_shared) + T.gemm( + S_shared, + KV_shared, + acc_o, + policy=T.GemmWarpPolicy.FullRow, + ) + T.gemm( + S_shared, + K_tail_shared, + acc_o_tail, + policy=T.GemmWarpPolicy.FullRow, + ) + # ---- finalize phase 2 (active) ---- + for h_i, d_i in T.Parallel(H_per_block, D): + acc_o[h_i, d_i] = acc_o[h_i, d_i] / T.if_then_else( + sumexp[h_i] == 0.0, 1.0, sumexp[h_i] + ) + for h_i, d_i in T.Parallel(H_per_block, D_tail): + acc_o_tail[h_i, d_i] = acc_o_tail[h_i, d_i] / T.if_then_else( + sumexp[h_i] == 0.0, 1.0, sumexp[h_i] + ) + for h_i in T.Parallel(H_per_block): + m_i[h_i] = T.if_then_else( + sumexp[h_i] == 0.0, + -(2.0**30), + T.log2(sumexp[h_i]) + m_i[h_i] * sm_scale, + ) + T.copy(acc_o, Partial_O[b_i, s_i, group_i, H0:H1, :D]) + T.copy( + acc_o_tail, + Partial_O[b_i, s_i, group_i, H0:H1, D : D + D_tail], + ) + T.copy(m_i, Partial_LSE[b_i, s_i, group_i, H0:H1]) + else: + # Phase 2 skipped: m_i is still the -2^30 + T.copy(m_i, Partial_LSE[b_i, s_i, group_i, H0:H1]) + + return main + + @T.prim_func + def main( + Q: T.Tensor(q_shape, BF16), # type: ignore + K_combined_1: T.Tensor(k1_shape, "uint32"), # type: ignore + Indices_1: T.Tensor(indices1_shape, indices_dtype), # type: ignore + Topk_length_1: T.Tensor(topk_length_shape, indices_dtype), # type: ignore + Partial_O: T.Tensor(partial_o_shape, BF16), # type: ignore + Partial_LSE: T.Tensor(partial_lse_shape, accum_dtype), # type: ignore + ) -> None: + """ + grid: (seq_len * REPLICATE_H * n_groups, batch, 1) + Each block processes `inner_iter_1` consecutive KV tiles and writes + one (partial_o, partial_lse) entry. + """ + with T.Kernel( + seq_len * REPLICATE_H * n_groups, batch, kv_group, threads=threads + ) as (bx, by, bz): + Q_shared = T.alloc_fragment([H_per_block, D], BF16) + Q_tail_shared = T.alloc_fragment([H_per_block, D_tail], BF16) + K_packed_shared = T.alloc_shared([BI, PACKED_W4], "uint32") + K_scale_shared = T.alloc_shared([BI, SCALE_W4], "uint32") + KV_shared = T.alloc_shared([BI, D], BF16) + K_tail_shared = T.alloc_shared([BI, D_tail], BF16) + S_shared = T.alloc_shared([H_per_block, BI], BF16) + page_idx_shared = T.alloc_shared([BI], INT32) + + mask = T.alloc_fragment([BI], "bool") + scale_byte_local = T.alloc_fragment([BI, NUM_TILES], "uint32") + + acc_o = T.alloc_fragment([H_per_block, D], accum_dtype) + acc_o_tail = T.alloc_fragment([H_per_block, D_tail], accum_dtype) + acc_s = T.alloc_fragment([H_per_block, BI], accum_dtype) + sumexp = T.alloc_fragment([H_per_block], accum_dtype) + sumexp_i = T.alloc_fragment([H_per_block], accum_dtype) + alpha = T.alloc_fragment([H_per_block], accum_dtype) + m_i = T.alloc_fragment([H_per_block], accum_dtype) + m_i_prev = T.alloc_fragment([H_per_block], accum_dtype) + + T.fill(acc_o, 0) + T.fill(acc_o_tail, 0) + T.fill(sumexp, 0) + T.fill(m_i, -(2**30)) + + b_i, g_i = by, bz + spans_per_seq = REPLICATE_H * n_groups + s_i = bx // spans_per_seq + rest = bx % spans_per_seq + group_i = rest // REPLICATE_H + h_rep = rest % REPLICATE_H + H0 = g_i * padded_H + (0 if REPLICATE_H == 1 else h_rep * 64) + H1 = H0 + H_per_block + + T.copy(Q[b_i, s_i, H0:H1, :D], Q_shared) + T.copy(Q[b_i, s_i, H0:H1, D : D + D_tail], Q_tail_shared) + + tk_len_1 = Topk_length_1[b_i] + + for k_i in T.Pipelined(inner_iter_1, num_stages=num_stages): + iter_i = group_i * inner_iter_1 + k_i + for bi_i in T.Parallel(BI): + pos = iter_i * BI + bi_i + idx = Indices_1[b_i, s_i, pos] + valid = (idx >= 0) & (pos < tk_len_1) + page_idx_shared[bi_i] = T.if_then_else(valid, idx, 0) + mask[bi_i] = valid + + for bi_i, w_i in T.Parallel(BI, PACKED_W4): + page = page_idx_shared[bi_i] + block_id = page // BS_KV_1 + t_in_block = page % BS_KV_1 + K_packed_shared[bi_i, w_i] = K_combined_1[ + block_id, t_in_block * PACKED_W4 + w_i + ] + + for bi_i, w_i in T.Parallel(BI, SCALE_W4): + page = page_idx_shared[bi_i] + block_id = page // BS_KV_1 + t_in_block = page % BS_KV_1 + K_scale_shared[bi_i, w_i] = K_combined_1[ + block_id, + NOPE_ROPE_U32_PER_BLOCK_1 + t_in_block * SCALE_W4 + w_i, + ] + + for bi_i, ti in T.Parallel(BI, NUM_TILES): + word_idx = ti // 4 + byte_in_word = ti % 4 + word = K_scale_shared[bi_i, word_idx] + scale_byte_local[bi_i, ti] = ( + word >> T.Cast("uint32", byte_in_word * 8) + ) & T.uint32(0xFF) + + for bi_i, d_i in T.Parallel(BI, D): + word_idx = d_i // 4 + byte_in_word = d_i % 4 + word = K_packed_shared[bi_i, word_idx] + b_u32 = (word >> T.Cast("uint32", byte_in_word * 8)) & T.uint32( + 0xFF + ) + sign_bf = (b_u32 & T.uint32(0x80)) * T.uint32(0x100) + exp_e4 = (b_u32 & T.uint32(0x78)) >> T.uint32(3) + mant_bf = (b_u32 & T.uint32(0x7)) * T.uint32(0x10) + scale_byte = scale_byte_local[bi_i, d_i // NOPE_TILE] + exp_combined = exp_e4 + scale_byte - T.uint32(7) + bf16_bits = sign_bf | (exp_combined << T.uint32(7)) | mant_bf + KV_shared[bi_i, d_i] = T.reinterpret( + BF16, T.Cast("uint16", bf16_bits) + ) + + for bi_i, j in T.Parallel(BI, D_tail): + abs_off = D + 2 * j + word_idx = abs_off // 4 + word_off = abs_off % 4 + word = K_packed_shared[bi_i, word_idx] + half_u32 = T.if_then_else( + word_off == 0, + word & T.uint32(0xFFFF), + (word >> T.uint32(16)) & T.uint32(0xFFFF), + ) + K_tail_shared[bi_i, j] = T.reinterpret( + BF16, T.Cast("uint16", half_u32) + ) + + for h_i, bi_i in T.Parallel(H_per_block, BI): + acc_s[h_i, bi_i] = T.if_then_else( + mask[bi_i], 0, -T.infinity(acc_s.dtype) + ) + T.gemm( + Q_shared, + KV_shared, + acc_s, + transpose_B=True, + policy=T.GemmWarpPolicy.FullRow, + ) + T.gemm( + Q_tail_shared, + K_tail_shared, + acc_s, + transpose_B=True, + policy=T.GemmWarpPolicy.FullRow, + ) + T.copy(m_i, m_i_prev) + T.reduce_max(acc_s, m_i, dim=1, clear=False) + for h_i in T.Parallel(H_per_block): + m_i[h_i] = T.max(m_i[h_i], m_i_prev[h_i]) + for h_i in T.Parallel(H_per_block): + alpha[h_i] = T.exp2((m_i_prev[h_i] - m_i[h_i]) * sm_scale) + for h_i, bi_i in T.Parallel(H_per_block, BI): + acc_s[h_i, bi_i] = T.exp2( + acc_s[h_i, bi_i] * sm_scale - m_i[h_i] * sm_scale + ) + T.reduce_sum(acc_s, sumexp_i, dim=1) + for h_i in T.Parallel(H_per_block): + sumexp[h_i] = sumexp[h_i] * alpha[h_i] + sumexp_i[h_i] + for h_i, d_i in T.Parallel(H_per_block, D): + acc_o[h_i, d_i] *= alpha[h_i] + for h_i, d_i in T.Parallel(H_per_block, D_tail): + acc_o_tail[h_i, d_i] *= alpha[h_i] + T.copy(acc_s, S_shared) + T.gemm(S_shared, KV_shared, acc_o, policy=T.GemmWarpPolicy.FullRow) + T.gemm( + S_shared, K_tail_shared, acc_o_tail, policy=T.GemmWarpPolicy.FullRow + ) + + for h_i, d_i in T.Parallel(H_per_block, D): + acc_o[h_i, d_i] = acc_o[h_i, d_i] / T.if_then_else( + sumexp[h_i] == 0.0, 1.0, sumexp[h_i] + ) + for h_i, d_i in T.Parallel(H_per_block, D_tail): + acc_o_tail[h_i, d_i] = acc_o_tail[h_i, d_i] / T.if_then_else( + sumexp[h_i] == 0.0, 1.0, sumexp[h_i] + ) + for h_i in T.Parallel(H_per_block): + m_i[h_i] = T.if_then_else( + sumexp[h_i] == 0.0, + -(2.0**30), + T.log2(sumexp[h_i]) + m_i[h_i] * sm_scale, + ) + T.copy(acc_o, Partial_O[b_i, s_i, group_i, H0:H1, :D]) + T.copy(acc_o_tail, Partial_O[b_i, s_i, group_i, H0:H1, D : D + D_tail]) + T.copy(m_i, Partial_LSE[b_i, s_i, group_i, H0:H1]) + + return main + + +@tilelang.jit( + out_idx=[-2, -1], + pass_configs={ + tilelang.PassConfigKey.TL_DISABLE_TMA_LOWER: True, + tilelang.PassConfigKey.TL_DISABLE_WARP_SPECIALIZED: True, + tilelang.PassConfigKey.TL_ENABLE_FAST_MATH: True, + }, +) +def dpsk_v4_combine_kernel( + num_heads: int, + n_groups_1: int, + n_groups_2: int = 0, + *, + block_I: int = 64, + inner_iter_1: int = 1, + inner_iter_2: int = 1, + dim: int = 448, + tail_dim: int = 64, + head_per_block: int = 16, + threads: int = 256, + use_attn_sink: bool = False, +) -> Any: + """ + Combine `n_groups` flash-attention partials into the final output. + + Inputs: + Partial_O : (batch, seq_len, n_groups, num_heads, dim+tail_dim) bf16 + Partial_LSE : (batch, seq_len, n_groups, num_heads) fp32, log2 form + Topk_length_1: (batch,) int32, actual phase-1 length + Topk_length_2: (batch,) int32, actual phase-2 length (dual only) + Attn_sink : (num_heads,) fp32 + Outputs: + Output : (batch, seq_len, num_heads, dim+tail_dim) bf16 + LSE : (batch, seq_len, num_heads) fp32, natural log + + Each grid block handles `head_per_block` heads of one (batch, seq) row. + """ + log2e: float = 1.44269504 + ln2: float = 0.69314718 + assert num_heads % head_per_block == 0 + + is_dual = n_groups_2 > 0 + n_groups = n_groups_1 + n_groups_2 + + H_per_block = head_per_block + HEAD_BLOCKS = num_heads // H_per_block + DT = dim + tail_dim + + batch = T.symbolic("batch") + seq_len = T.symbolic("seq_len") + + accum_dtype = "float" + + if is_dual: + + @T.prim_func + def main( + Partial_O: T.Tensor( + [batch, seq_len, n_groups, num_heads, DT], BF16 + ), # type: ignore + Partial_LSE: T.Tensor( + [batch, seq_len, n_groups, num_heads], accum_dtype + ), # type: ignore + Topk_length_1: T.Tensor([batch], INT32), # type: ignore + Topk_length_2: T.Tensor([batch], INT32), # type: ignore + Attn_sink: T.Tensor([num_heads], FP32), # type: ignore + Output: T.Tensor([batch, seq_len, num_heads, DT], BF16), # type: ignore + LSE: T.Tensor([batch, seq_len, num_heads], accum_dtype), # type: ignore + ) -> None: + with T.Kernel(seq_len * HEAD_BLOCKS, batch, threads=threads) as ( + bx, + by, + ): + shared_lse = T.alloc_shared([n_groups, H_per_block], accum_dtype) + lse_max = T.alloc_fragment([H_per_block], accum_dtype) + lse_sum = T.alloc_fragment([H_per_block], accum_dtype) + scale = T.alloc_fragment([H_per_block, n_groups], accum_dtype) + acc_o = T.alloc_fragment([H_per_block, DT], accum_dtype) + attn_sink_frag = T.alloc_fragment([H_per_block], accum_dtype) + o_scale_frag = T.alloc_fragment([H_per_block], accum_dtype) + final_lse = T.alloc_fragment([H_per_block], accum_dtype) + + b_i = by + s_i = bx // HEAD_BLOCKS + head_block = bx % HEAD_BLOCKS + H0 = head_block * H_per_block + H1 = H0 + H_per_block + + # Clamp to the captured-shape upper bounds so callers passing + # the INT32_MAX sentinel (= "all valid") still iterate exactly + # n_groups groups, not 33M. + actual_n_groups_1 = T.min( + T.ceildiv(Topk_length_1[b_i], block_I * inner_iter_1), + n_groups_1, + ) + actual_n_groups_2 = T.min( + T.ceildiv(Topk_length_2[b_i], block_I * inner_iter_2), + n_groups - n_groups_1, + ) + actual_n_groups = actual_n_groups_1 + actual_n_groups_2 + + # Pass 1: load only active groups' LSE into compact slots. + for k_c in T.serial(actual_n_groups): + k = T.if_then_else( + k_c < actual_n_groups_1, + k_c, + n_groups_1 + (k_c - actual_n_groups_1), + ) + T.copy(Partial_LSE[b_i, s_i, k, H0:H1], shared_lse[k_c, :]) + + T.fill(lse_max, -(2**30)) + for k_c in T.serial(actual_n_groups): + for h_i in T.Parallel(H_per_block): + lse_max[h_i] = T.max(lse_max[h_i], shared_lse[k_c, h_i]) + T.fill(lse_sum, 0) + for k_c in T.serial(actual_n_groups): + for h_i in T.Parallel(H_per_block): + lse_sum[h_i] = lse_sum[h_i] + T.exp2( + shared_lse[k_c, h_i] - lse_max[h_i] + ) + for k_c in T.serial(actual_n_groups): + for h_i in T.Parallel(H_per_block): + scale[h_i, k_c] = T.exp2( + shared_lse[k_c, h_i] - lse_max[h_i] - T.log2(lse_sum[h_i]) + ) + + T.fill(acc_o, 0) + for k_c in T.serial(actual_n_groups): + k = T.if_then_else( + k_c < actual_n_groups_1, + k_c, + n_groups_1 + (k_c - actual_n_groups_1), + ) + for h_i, d_i in T.Parallel(H_per_block, DT): + acc_o[h_i, d_i] = acc_o[h_i, d_i] + scale[h_i, k_c] * Partial_O[ + b_i, s_i, k, H0 + h_i, d_i + ].astype(accum_dtype) + + for h_i in T.Parallel(H_per_block): + empty = lse_max[h_i] <= -(2**29) + final_lse[h_i] = T.if_then_else( + empty, + T.infinity(accum_dtype), + (lse_max[h_i] + T.log2(lse_sum[h_i])) * ln2, + ) + + if use_attn_sink: + for h_i in T.Parallel(H_per_block): + attn_sink_frag[h_i] = Attn_sink[H0 + h_i] + for h_i in T.Parallel(H_per_block): + empty = lse_max[h_i] <= -(2**29) + o_scale_frag[h_i] = T.if_then_else( + empty, + 0.0, + 1.0 + / ( + 1.0 + + T.exp2((attn_sink_frag[h_i] - final_lse[h_i]) * log2e) + ), + ) + for h_i, d_i in T.Parallel(H_per_block, DT): + acc_o[h_i, d_i] = acc_o[h_i, d_i] * o_scale_frag[h_i] + + T.copy(acc_o, Output[b_i, s_i, H0:H1, :]) + T.copy(final_lse, LSE[b_i, s_i, H0:H1]) + + return main + + @T.prim_func + def main( + Partial_O: T.Tensor( + [batch, seq_len, n_groups, num_heads, DT], BF16 + ), # type: ignore + Partial_LSE: T.Tensor( + [batch, seq_len, n_groups, num_heads], accum_dtype + ), # type: ignore + Attn_sink: T.Tensor([num_heads], FP32), # type: ignore + Output: T.Tensor([batch, seq_len, num_heads, DT], BF16), # type: ignore + LSE: T.Tensor([batch, seq_len, num_heads], accum_dtype), # type: ignore + ) -> None: + with T.Kernel(seq_len * HEAD_BLOCKS, batch, threads=threads) as (bx, by): + shared_lse = T.alloc_shared([n_groups, H_per_block], accum_dtype) + + lse_max = T.alloc_fragment([H_per_block], accum_dtype) + lse_sum = T.alloc_fragment([H_per_block], accum_dtype) + scale = T.alloc_fragment([H_per_block, n_groups], accum_dtype) + acc_o = T.alloc_fragment([H_per_block, DT], accum_dtype) + attn_sink_frag = T.alloc_fragment([H_per_block], accum_dtype) + o_scale_frag = T.alloc_fragment([H_per_block], accum_dtype) + final_lse = T.alloc_fragment([H_per_block], accum_dtype) + + b_i = by + s_i = bx // HEAD_BLOCKS + head_block = bx % HEAD_BLOCKS + H0 = head_block * H_per_block + H1 = H0 + H_per_block + + for k in T.serial(n_groups): + T.copy(Partial_LSE[b_i, s_i, k, H0:H1], shared_lse[k, :]) + + T.fill(lse_max, -(2**30)) + for k in T.serial(n_groups): + for h_i in T.Parallel(H_per_block): + lse_max[h_i] = T.max(lse_max[h_i], shared_lse[k, h_i]) + T.fill(lse_sum, 0) + for k in T.serial(n_groups): + for h_i in T.Parallel(H_per_block): + lse_sum[h_i] = lse_sum[h_i] + T.exp2( + shared_lse[k, h_i] - lse_max[h_i] + ) + for k in T.serial(n_groups): + for h_i in T.Parallel(H_per_block): + scale[h_i, k] = T.exp2( + shared_lse[k, h_i] - lse_max[h_i] - T.log2(lse_sum[h_i]) + ) + + T.fill(acc_o, 0) + for k in T.serial(n_groups): + for h_i, d_i in T.Parallel(H_per_block, DT): + acc_o[h_i, d_i] = acc_o[h_i, d_i] + scale[h_i, k] * Partial_O[ + b_i, s_i, k, H0 + h_i, d_i + ].astype(accum_dtype) + + for h_i in T.Parallel(H_per_block): + empty = lse_max[h_i] <= -(2**29) + final_lse[h_i] = T.if_then_else( + empty, + T.infinity(accum_dtype), + (lse_max[h_i] + T.log2(lse_sum[h_i])) * ln2, + ) + + if use_attn_sink: + for h_i in T.Parallel(H_per_block): + attn_sink_frag[h_i] = Attn_sink[H0 + h_i] + for h_i in T.Parallel(H_per_block): + empty = lse_max[h_i] <= -(2**29) + o_scale_frag[h_i] = T.if_then_else( + empty, + 0.0, + 1.0 + / ( + 1.0 + T.exp2((attn_sink_frag[h_i] - final_lse[h_i]) * log2e) + ), + ) + for h_i, d_i in T.Parallel(H_per_block, DT): + acc_o[h_i, d_i] = acc_o[h_i, d_i] * o_scale_frag[h_i] + + T.copy(acc_o, Output[b_i, s_i, H0:H1, :]) + T.copy(final_lse, LSE[b_i, s_i, H0:H1]) + + return main + + +""" +2-stage attention kernel (partial + combine) over an FP8 KV cache, +with optional second cache (`extra_k_cache`). +""" + + +def dpsk_v4_fp8_attention_fwd( + q: torch.Tensor, + k_cache: torch.Tensor, + block_table: Optional[torch.Tensor], + cache_seqlens: Optional[torch.Tensor], + head_dim_v: int, + tile_scheduler_metadata: Any, + num_splits: None = None, + softmax_scale: Optional[float] = None, + causal: bool = False, + is_fp8_kvcache: bool = False, + indices: Optional[torch.Tensor] = None, + attn_sink: Optional[torch.Tensor] = None, + extra_k_cache: Optional[torch.Tensor] = None, + extra_indices_in_kvcache: Optional[torch.Tensor] = None, + topk_length: Optional[torch.Tensor] = None, + extra_topk_length: Optional[torch.Tensor] = None, +) -> Tuple[torch.Tensor, torch.Tensor]: + """ + Follows the original `flash_mla.flash_mla_with_kvcache` signature. + """ + if _is_gfx95_supported: + block_I, threads, num_stages, block_per_cu, cu = 64, 512, 0, 2, 256 + else: + block_I, threads, num_stages, block_per_cu, cu = 32, 128, 1, 1, 304 + + batch, seq_len, num_heads, _ = q.shape + # Partial grid is (seq_len * REPLICATE_H * n_groups, batch, kv_group); the + # heuristic in _pick_inner_iter assumes `total_blocks = seq * ni / inner_iter`, + # so `seq` must include REPLICATE_H or n_groups doubles for medium batches. + replicate_h = max((num_heads + 63) // 64, 1) + seq = batch * seq_len * replicate_h + + k1, _, bs_kv_1 = _build_fp8_combined_view(k_cache) + topk_1 = indices.shape[-1] + ni_1 = topk_1 // block_I + tk_len_1 = ( + topk_length + if topk_length is not None + else _topk_length_sentinel(q.device, batch) + ) + if attn_sink is None: + attn_sink = torch.full( + (num_heads,), float("-inf"), dtype=torch.float32, device=q.device + ) + + has_extra = extra_k_cache is not None + if not has_extra: + inner_iter_1 = _pick_inner_iter(seq, ni_1, cu, block_per_cu) + inner_iter_2 = 1 + n_groups_1 = ni_1 // inner_iter_1 + n_groups_2 = 0 + partial = dpsk_v4_fp8_partial_kernel( + num_heads, + topk_1, + bs_kv_1, + sm_scale=softmax_scale, + block_I=block_I, + inner_iter_1=inner_iter_1, + num_stages=num_stages, + threads=threads, + ) + partial_o, partial_lse = partial(q, k1, indices, tk_len_1) + else: + k2, _, bs_kv_2 = _build_fp8_combined_view(extra_k_cache) + topk_2 = extra_indices_in_kvcache.shape[-1] + ni_2 = topk_2 // block_I + # Each phase picks its own optimal split-K independently — kernel + # body uses two T.Pipelined loops with separate compile-time iter + # counts, no shared-divisor constraint. + inner_iter_1 = _pick_inner_iter(seq, ni_1, cu, block_per_cu) + inner_iter_2 = _pick_inner_iter(seq, ni_2, cu, block_per_cu) + n_groups_1 = ni_1 // inner_iter_1 + n_groups_2 = ni_2 // inner_iter_2 + tk_len_2 = ( + extra_topk_length + if extra_topk_length is not None + else _topk_length_sentinel(q.device, batch) + ) + partial = dpsk_v4_fp8_partial_kernel( + num_heads, + topk_1, + bs_kv_1, + topk_2, + bs_kv_2, + sm_scale=softmax_scale, + block_I=block_I, + inner_iter_1=inner_iter_1, + inner_iter_2=inner_iter_2, + num_stages=num_stages, + threads=threads, + ) + partial_o, partial_lse = partial( + q, + k1, + indices, + tk_len_1, + k2, + extra_indices_in_kvcache, + tk_len_2, + ) + + combine = dpsk_v4_combine_kernel( + num_heads, + n_groups_1, + n_groups_2, + block_I=block_I, + inner_iter_1=inner_iter_1, + inner_iter_2=inner_iter_2, + head_per_block=4, + threads=256, + use_attn_sink=True, + ) + if has_extra: + return combine(partial_o, partial_lse, tk_len_1, tk_len_2, attn_sink) + return combine(partial_o, partial_lse, attn_sink) diff --git a/python/sglang/srt/layers/deepseek_v4_rope.py b/python/sglang/srt/layers/deepseek_v4_rope.py index c717850c6..c8d391426 100644 --- a/python/sglang/srt/layers/deepseek_v4_rope.py +++ b/python/sglang/srt/layers/deepseek_v4_rope.py @@ -177,3 +177,171 @@ def apply_rotary_emb_triton( ) return x + + +@triton.jit +def _fused_norm_rope_kernel( + x_ptr, + weight_ptr, + freqs_real_ptr, + positions_ptr, + eps, + stride_x_row, + stride_freq_row, + HEAD_DIM: tl.constexpr, + ROPE_DIM: tl.constexpr, + HEAD_BLOCK: tl.constexpr, + ROPE_PAIR_BLOCK: tl.constexpr, + HAS_WEIGHT: tl.constexpr, + USE_POS: tl.constexpr, +): + # NOTE: avoids store-then-reload on the same kernel: rope-segment values + # are loaded a 2nd time as (real, imag) pairs straight from the input, + # rms_inv/weight applied in register, and all stores happen at the end. + pid = tl.program_id(0) + base = pid.to(tl.int64) * stride_x_row + + offs = tl.arange(0, HEAD_BLOCK) + mask = offs < HEAD_DIM + x = tl.load(x_ptr + base + offs, mask=mask, other=0.0).to(tl.float32) + + sum_sq = tl.sum(x * x, axis=0) + rms_inv = tl.rsqrt(sum_sq / HEAD_DIM + eps) + + if HAS_WEIGHT: + w = tl.load(weight_ptr + offs, mask=mask, other=0.0).to(tl.float32) + x_normed = x * rms_inv * w + else: + x_normed = x * rms_inv + + rope_start = HEAD_DIM - ROPE_DIM + + pair_offs = tl.arange(0, ROPE_PAIR_BLOCK) + pair_mask = pair_offs < (ROPE_DIM // 2) + + x_real = tl.load( + x_ptr + base + rope_start + 2 * pair_offs, + mask=pair_mask, + other=0.0, + ).to(tl.float32) + x_imag = tl.load( + x_ptr + base + rope_start + 2 * pair_offs + 1, + mask=pair_mask, + other=0.0, + ).to(tl.float32) + + if HAS_WEIGHT: + w_real = tl.load( + weight_ptr + rope_start + 2 * pair_offs, + mask=pair_mask, + other=1.0, + ).to(tl.float32) + w_imag = tl.load( + weight_ptr + rope_start + 2 * pair_offs + 1, + mask=pair_mask, + other=1.0, + ).to(tl.float32) + x_real = x_real * rms_inv * w_real + x_imag = x_imag * rms_inv * w_imag + else: + x_real = x_real * rms_inv + x_imag = x_imag * rms_inv + + if USE_POS: + position = tl.load(positions_ptr + pid).to(tl.int64) + else: + position = pid.to(tl.int64) + + freq_base = position * stride_freq_row + f_real = tl.load( + freqs_real_ptr + freq_base + 2 * pair_offs, + mask=pair_mask, + other=0.0, + ).to(tl.float32) + f_imag = tl.load( + freqs_real_ptr + freq_base + 2 * pair_offs + 1, + mask=pair_mask, + other=0.0, + ).to(tl.float32) + + out_real = x_real * f_real - x_imag * f_imag + out_imag = x_real * f_imag + x_imag * f_real + + is_non_rope = offs < rope_start + tl.store( + x_ptr + base + offs, + x_normed.to(x_ptr.dtype.element_ty), + mask=mask & is_non_rope, + ) + tl.store( + x_ptr + base + rope_start + 2 * pair_offs, + out_real.to(x_ptr.dtype.element_ty), + mask=pair_mask, + ) + tl.store( + x_ptr + base + rope_start + 2 * pair_offs + 1, + out_imag.to(x_ptr.dtype.element_ty), + mask=pair_mask, + ) + + +def fused_norm_rope_inplace_triton( + kv: torch.Tensor, + weight: Optional[torch.Tensor], + eps: float, + freqs_cis: torch.Tensor, + positions: Optional[torch.Tensor] = None, +) -> None: + """Fused RMSNorm (over head_dim) + RoPE (on last rope_dim of head_dim), in-place. + + Equivalent to:: + + kv = rms_normalize(kv, eps, weight) + apply_rotary_emb_triton(kv[..., -rope_dim:], freqs_cis, positions=positions) + + Args: + kv: [M, head_dim], any float dtype, contiguous along last dim. Modified in-place. + weight: [head_dim] or None. + eps: RMSNorm epsilon. + freqs_cis: complex tensor. + - If ``positions`` is None: shape [M, rope_dim // 2], one freq per token. + - Else: shape [max_seq, rope_dim // 2], full table; indexed by ``positions``. + positions: optional [M] int tensor, absolute positions to index into ``freqs_cis``. + """ + assert kv.dim() == 2 and kv.stride(-1) == 1 + M, head_dim = kv.shape + + freqs_real = torch.view_as_real(freqs_cis).flatten(-2) + rope_dim = freqs_real.shape[-1] + assert head_dim >= rope_dim and rope_dim % 2 == 0 + if weight is not None: + assert weight.shape == (head_dim,) + if positions is None: + assert ( + freqs_real.shape[0] == M + ), f"freqs_cis row count {freqs_real.shape[0]} != M={M}" + else: + assert positions.shape == (M,) and positions.dim() == 1 + + if M == 0: + return + + HEAD_BLOCK = triton.next_power_of_2(head_dim) + ROPE_PAIR_BLOCK = max(triton.next_power_of_2(rope_dim // 2), 1) + + grid = (M,) + _fused_norm_rope_kernel[grid]( + kv, + weight, + freqs_real, + positions, + eps, + kv.stride(0), + freqs_real.stride(0), + HEAD_DIM=head_dim, + ROPE_DIM=rope_dim, + HEAD_BLOCK=HEAD_BLOCK, + ROPE_PAIR_BLOCK=ROPE_PAIR_BLOCK, + HAS_WEIGHT=(weight is not None), + USE_POS=(positions is not None), + ) diff --git a/python/sglang/srt/layers/moe/topk.py b/python/sglang/srt/layers/moe/topk.py index f47758996..62a22e572 100644 --- a/python/sglang/srt/layers/moe/topk.py +++ b/python/sglang/srt/layers/moe/topk.py @@ -1407,7 +1407,9 @@ def select_experts( scoring_func=scoring_func, ) elif custom_routing_function is None: - assert not apply_routed_scaling_factor_on_output, "Not implemented" + if scoring_func != "sqrtsoftplus": + assert not apply_routed_scaling_factor_on_output, "Not implemented" + if scoring_func == "sqrtsoftplus": _biased_topk = ( biased_topk_jit_kernel_impl diff --git a/python/sglang/srt/layers/quantization/fp8.py b/python/sglang/srt/layers/quantization/fp8.py index 0eb85e734..8d7dfa2d3 100644 --- a/python/sglang/srt/layers/quantization/fp8.py +++ b/python/sglang/srt/layers/quantization/fp8.py @@ -84,6 +84,7 @@ from sglang.srt.utils import ( get_bool_env_var, is_cpu, is_cuda, + is_gfx95_supported, is_hip, is_musa, is_npu, @@ -111,9 +112,21 @@ _is_cpu = is_cpu() _is_fp8_fnuz = is_fp8_fnuz() _use_hip_int4 = get_bool_env_var("SGLANG_INT4_WEIGHT") and _is_hip _use_aiter = envs.SGLANG_USE_AITER.get() and _is_hip +_is_shuffle_moe_mxfp4 = is_gfx95_supported() + + +def _require_fp4_dtype(): + fp4_dtype = getattr(torch, "float4_e2m1fn_x2", None) + if fp4_dtype is None: + raise RuntimeError( + "DeepSeek-V4 FP4 experts require torch.float4_e2m1fn_x2 support." + ) + return fp4_dtype + if _use_aiter or _use_hip_int4: from aiter.ops.shuffle import shuffle_weight + from aiter.utility.fp4_utils import e8m0_shuffle if _use_aiter: from sglang.srt.layers.quantization.fp8_utils import ( @@ -998,12 +1011,13 @@ class Fp8MoEMethod(FusedMoEMethodBase): # WEIGHT_SCALES if self.is_fp4_expert: fp4_block_k = 32 + fp4_scale_dtype = torch.float8_e8m0fnu if _use_aiter else torch.float32 w13_weight_scale = torch.nn.Parameter( torch.ones( num_experts, 2 * intermediate_size_per_partition, hidden_size // fp4_block_k, - dtype=torch.float32, + dtype=fp4_scale_dtype, ), requires_grad=False, ) @@ -1012,7 +1026,7 @@ class Fp8MoEMethod(FusedMoEMethodBase): num_experts, hidden_size, intermediate_size_per_partition // fp4_block_k, - dtype=torch.float32, + dtype=fp4_scale_dtype, ), requires_grad=False, ) @@ -1123,6 +1137,105 @@ class Fp8MoEMethod(FusedMoEMethodBase): layer.w2_input_scale = None def process_weights_after_loading_block_quant(self, layer: Module) -> None: + # AMD FP4 experts: use aiter's native MXFP4 MoE path + if _use_aiter and self.is_fp4_expert: + fp4_weight_dtype = _require_fp4_dtype() + + # CK FP4 MoE kernel requires K_packed divisible by 128 + # (i.e., K_logical divisible by 256). + # Pad intermediate_size_per_partition if needed. + fp4_k_align = 256 + E, w13_N, w13_K_packed = layer.w13_weight.shape + _, w2_N, w2_K_packed = layer.w2_weight.shape + inter_per_part = w13_N // 2 + padded_inter = ( + (inter_per_part + fp4_k_align - 1) // fp4_k_align * fp4_k_align + ) + if padded_inter != inter_per_part: + pad_amount = padded_inter - inter_per_part + fp4_block_k = 32 + + # Pad w13_weight: (E, 2*inter, K_packed) → (E, 2*padded, K_packed) + old_w13 = layer.w13_weight.data + new_w13 = torch.zeros( + E, + 2 * padded_inter, + w13_K_packed, + dtype=old_w13.dtype, + device=old_w13.device, + ) + new_w13[:, :inter_per_part, :] = old_w13[:, :inter_per_part, :] + new_w13[:, padded_inter : padded_inter + inter_per_part, :] = old_w13[ + :, inter_per_part:, : + ] + layer.w13_weight = torch.nn.Parameter(new_w13, requires_grad=False) + + # Pad w2_weight: (E, N, inter_packed) → (E, N, padded_packed) + old_w2 = layer.w2_weight.data + new_w2 = torch.zeros( + E, + w2_N, + padded_inter // 2, + dtype=old_w2.dtype, + device=old_w2.device, + ) + new_w2[:, :, :w2_K_packed] = old_w2 + layer.w2_weight = torch.nn.Parameter(new_w2, requires_grad=False) + + # Pad w13 scale: (E, 2*inter, K/block_k) → (E, 2*padded, K/block_k) + old_s13 = layer.w13_weight_scale_inv.data + _, _, s13_K = old_s13.shape + new_s13 = torch.zeros( + E, + 2 * padded_inter, + s13_K, + dtype=old_s13.dtype, + device=old_s13.device, + ) + new_s13[:, :inter_per_part, :] = old_s13[:, :inter_per_part, :] + new_s13[:, padded_inter : padded_inter + inter_per_part, :] = old_s13[ + :, inter_per_part:, : + ] + layer.w13_weight_scale_inv = torch.nn.Parameter( + new_s13, requires_grad=False + ) + + # Pad w2 scale: (E, N, inter/block_k) → (E, N, padded/block_k) + old_s2 = layer.w2_weight_scale_inv.data + new_s2 = torch.zeros( + E, + w2_N, + padded_inter // fp4_block_k, + dtype=old_s2.dtype, + device=old_s2.device, + ) + new_s2[:, :, : old_s2.shape[2]] = old_s2 + layer.w2_weight_scale_inv = torch.nn.Parameter( + new_s2, requires_grad=False + ) + + for scale_name in ("w13_weight_scale_inv", "w2_weight_scale_inv"): + scale = getattr(layer, scale_name) + num_experts, num_rows, _ = scale.shape + scale.data = e8m0_shuffle(scale.view(num_experts * num_rows, -1)).view( + num_experts, num_rows, -1 + ) + + layer.w13_weight.data = layer.w13_weight.data.view(fp4_weight_dtype) + layer.w2_weight.data = layer.w2_weight.data.view(fp4_weight_dtype) + + is_shuffled = _is_shuffle_moe_mxfp4 + if is_shuffled: + layer.w13_weight.data = shuffle_weight( + layer.w13_weight.contiguous(), (16, 16) + ) + layer.w2_weight.data = shuffle_weight( + layer.w2_weight.contiguous(), (16, 16) + ) + layer.w13_weight.is_shuffled = is_shuffled + layer.w2_weight.is_shuffled = is_shuffled + return + # If ROCm, normalize the weights and scales to e4m3fnuz if _is_fp8_fnuz: # activation_scheme: dynamic @@ -1148,8 +1261,6 @@ class Fp8MoEMethod(FusedMoEMethodBase): ) layer.w2_input_scale = None if _use_aiter: - # add this section for MI300 - # Pre-shuffle weights layer.w13_weight.data = shuffle_weight( layer.w13_weight.contiguous(), (16, 16) ) @@ -1158,12 +1269,12 @@ class Fp8MoEMethod(FusedMoEMethodBase): ) elif _use_aiter: # Pre-shuffle weights - layer.w13_weight.data = shuffle_weight( - layer.w13_weight.contiguous(), (16, 16) - ) - layer.w2_weight.data = shuffle_weight( - layer.w2_weight.contiguous(), (16, 16) - ) + t = shuffle_weight(layer.w13_weight, (16, 16)) + layer.w13_weight.copy_(t) + del t + t = shuffle_weight(layer.w2_weight, (16, 16)) + layer.w2_weight.copy_(t) + del t elif _is_cpu: assert ( _is_cpu_amx_available @@ -1190,8 +1301,9 @@ class Fp8MoEMethod(FusedMoEMethodBase): layer.w2_weight.data = layer.w2_weight.data.view(torch.int8) return - layer.w13_weight.data = layer.w13_weight.data.view(torch.int8) - layer.w2_weight.data = layer.w2_weight.data.view(torch.int8) + fp4_weight_dtype = _require_fp4_dtype() if _use_aiter else torch.int8 + layer.w13_weight.data = layer.w13_weight.data.view(fp4_weight_dtype) + layer.w2_weight.data = layer.w2_weight.data.view(fp4_weight_dtype) if get_moe_a2a_backend().is_megamoe(): from sglang.srt.layers.moe.mega_moe import ( @@ -1930,8 +2042,23 @@ class Fp8MoEMethod(FusedMoEMethodBase): AiterQuantType, ) - if _use_aiter and self.block_quant: - quant_type = AiterQuantType.PER_128X128 + w13_weight = layer.w13_weight + w2_weight = layer.w2_weight + + if self.block_quant: + quant_type = ( + AiterQuantType.PER_1X32 + if self.is_fp4_expert + else AiterQuantType.PER_128X128 + ) + + if self.is_fp4_expert: + fp4_weight_dtype = _require_fp4_dtype() + w13_weight = w13_weight.view(fp4_weight_dtype) + w2_weight = w2_weight.view(fp4_weight_dtype) + if getattr(layer.w13_weight, "is_shuffled", False): + w13_weight.is_shuffled = True + w2_weight.is_shuffled = True w13_scale = layer.w13_weight_scale_inv w2_scale = layer.w2_weight_scale_inv else: @@ -1939,8 +2066,8 @@ class Fp8MoEMethod(FusedMoEMethodBase): w13_scale = layer.w13_weight_scale1 w2_scale = layer.w2_weight_scale1 return AiterMoeQuantInfo( - w13_weight=layer.w13_weight, - w2_weight=layer.w2_weight, + w13_weight=w13_weight, + w2_weight=w2_weight, quant_type=quant_type, w13_scale=w13_scale, w2_scale=w2_scale, diff --git a/python/sglang/srt/mem_cache/deepseek_v4_compress_state.py b/python/sglang/srt/mem_cache/deepseek_v4_compress_state.py index 3c865fe84..4543c48fd 100644 --- a/python/sglang/srt/mem_cache/deepseek_v4_compress_state.py +++ b/python/sglang/srt/mem_cache/deepseek_v4_compress_state.py @@ -7,8 +7,11 @@ import torch from sglang.srt.constants import GPU_MEMORY_TYPE_KV_CACHE from sglang.srt.mem_cache.utils import maybe_init_custom_mem_pool +from sglang.srt.utils import is_hip from sglang.srt.utils.torch_memory_saver_adapter import TorchMemorySaverAdapter +_is_hip = is_hip() + @dataclasses.dataclass class KVAndScore: @@ -22,16 +25,55 @@ class KVAndScore: def score(self) -> torch.Tensor: return self.kv_score[..., self._item_size :] + @property + def shape(self): + return self.kv_score.shape + def __post_init__(self): self._item_size = self.kv_score.shape[-1] // 2 + @staticmethod + def from_kv_score(*, kv: torch.Tensor, score: torch.Tensor) -> KVAndScore: + assert kv.shape == score.shape + return KVAndScore(torch.cat([kv, score], dim=-1)) + + def new_empty(self, new_shape) -> KVAndScore: + assert new_shape[-1] == self._item_size + new_shape = list(new_shape) + new_shape[-1] = 2 * self._item_size + return KVAndScore(self.kv_score.new_empty(new_shape, requires_grad=False)) + def __getitem__(self, index) -> KVAndScore: return KVAndScore(self.kv_score[index]) + def __setitem__(self, index, value: KVAndScore): + self.kv_score[index] = value.kv_score + def clear(self): self.kv.zero_() self.score.fill_(float("-inf")) + def view(self, *args): + args = list(args) + if isinstance(args[-1], int) and args[-1] != -1: + args[-1] = 2 * self._item_size + return KVAndScore(self.kv_score.view(*args)) + + def clone(self) -> KVAndScore: + return KVAndScore(self.kv_score.clone()) + + @staticmethod + def cat(tensors: list[KVAndScore], dim: int) -> KVAndScore: + assert dim != -1, "Concatenation along last dim is not supported." + assert len(tensors) > 0, "At least one tensor is required for concatenation." + item_size = tensors[0]._item_size + for v in tensors: + assert ( + v._item_size == item_size + ), "All tensors must have the same item size." + + return KVAndScore(torch.cat([v.kv_score for v in tensors], dim=dim)) + class CompressStatePool: def __init__( @@ -45,8 +87,11 @@ class CompressStatePool: enable_memory_saver: bool, ratio: int, online: bool = False, + swa_page_size: int = 0, ): self.ring_size = ring_size + self.swa_page_size = swa_page_size + self.enable_memory_saver = enable_memory_saver if online: assert ring_size == 1, "online compress requires ring_size=1" @@ -57,25 +102,47 @@ class CompressStatePool: self._size = (self._size + ratio - 1) // ratio * ratio last_dim = 2 * (1 + overlap) * head_dim - self.memory_saver_adapter = TorchMemorySaverAdapter.create( - enable=enable_memory_saver - ) - self.enable_custom_mem_pool, self.custom_mem_pool, _ = ( - maybe_init_custom_mem_pool(device=device) - ) + if _is_hip: + self.kv_score_buffer = KVAndScore( + torch.empty((self._size, last_dim), dtype=dtype, device=device) + ) + if not online: + self.kv_score_buffer[-1].clear() + else: + self.memory_saver_adapter = TorchMemorySaverAdapter.create( + enable=enable_memory_saver + ) + self.enable_custom_mem_pool, self.custom_mem_pool, _ = ( + maybe_init_custom_mem_pool(device=device) + ) - with self.memory_saver_adapter.region(GPU_MEMORY_TYPE_KV_CACHE): - with ( - torch.cuda.use_mem_pool(self.custom_mem_pool) - if self.custom_mem_pool - else nullcontext() - ): - self.kv_score_buffer = KVAndScore( - torch.empty( - (self._size, last_dim), - dtype=dtype, - device=device, + with self.memory_saver_adapter.region(GPU_MEMORY_TYPE_KV_CACHE): + with ( + torch.cuda.use_mem_pool(self.custom_mem_pool) + if self.custom_mem_pool + else nullcontext() + ): + self.kv_score_buffer = KVAndScore( + torch.empty( + (self._size, last_dim), + dtype=dtype, + device=device, + ) ) - ) - if not online: - self.kv_score_buffer[-1].clear() + if not online: + self.kv_score_buffer[-1].clear() + + def translate_from_swa_loc_to_state_loc( + self, swa_loc: torch.Tensor + ) -> torch.Tensor: + swa_pages = swa_loc // self.swa_page_size + state_loc = swa_pages * self.ring_size + (swa_loc % self.ring_size) + state_loc = torch.where(swa_loc < 0, -1, state_loc) + return state_loc + + def get_state_by_state_loc(self, state_loc: torch.Tensor) -> KVAndScore: + return self.kv_score_buffer[state_loc] + + def set_state_by_state_loc(self, state_loc: torch.Tensor, value: KVAndScore): + self.kv_score_buffer[state_loc] = value + self.kv_score_buffer[-1].clear() diff --git a/python/sglang/srt/mem_cache/deepseek_v4_memory_pool.py b/python/sglang/srt/mem_cache/deepseek_v4_memory_pool.py index a185b02c5..a6b826d7d 100644 --- a/python/sglang/srt/mem_cache/deepseek_v4_memory_pool.py +++ b/python/sglang/srt/mem_cache/deepseek_v4_memory_pool.py @@ -18,11 +18,13 @@ from sglang.srt.mem_cache.base_swa_memory_pool import BaseSWAKVPool from sglang.srt.mem_cache.deepseek_v4_compress_state import CompressStatePool from sglang.srt.mem_cache.memory_pool import KVCache from sglang.srt.server_args import get_global_server_args -from sglang.srt.utils import ceil_div +from sglang.srt.utils import ceil_div, is_hip logger = logging.getLogger(__name__) -ONLINE_C128 = envs.SGLANG_OPT_USE_ONLINE_COMPRESS.get() +_is_hip = is_hip() + +ONLINE_C128 = not _is_hip and envs.SGLANG_OPT_USE_ONLINE_COMPRESS.get() def get_compress_state_ring_size( @@ -144,6 +146,9 @@ class DeepSeekV4SingleKVPool(KVCache): ) def get_key_buffer(self, layer_id: int): + if self.store_dtype != self.dtype: + return self.kv_buffer[layer_id - self.start_layer].view(self.dtype) + return self.kv_buffer[layer_id] def set_kv_buffer(self, *args, **kwargs) -> None: @@ -466,7 +471,7 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool): ) self.c4_indexer_kv_pool = DeepSeekV4IndexerPool( - self.c4_logical_size, + self.c4_logical_size if not _is_hip else c4_size, c4_page_size, dtype, indexer_head_dim, @@ -477,7 +482,10 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool): self._init_compressed_layer_mapping() - self._init_paged_compress_states(enable_memory_saver) + if _is_hip: + self._init_paged_compress_states(False) + else: + self._init_paged_compress_states(enable_memory_saver) self._should_cache_swa = envs.SGLANG_OPT_CACHE_SWA_TRANSLATION.get() self.cached_loc = None @@ -585,6 +593,7 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool): dtype=self.state_dtype, enable_memory_saver=enable_memory_saver, ratio=ratio, + swa_page_size=self.swa_page_size, ) def _init_compressed_layer_mapping(self): diff --git a/python/sglang/srt/models/deepseek_v2.py b/python/sglang/srt/models/deepseek_v2.py index de7b88afe..da96f9ba3 100644 --- a/python/sglang/srt/models/deepseek_v2.py +++ b/python/sglang/srt/models/deepseek_v2.py @@ -575,6 +575,11 @@ class DeepseekV2MoE(nn.Module): use_grouped_topk=False, scoring_func=config.scoring_func, is_fp4_experts=getattr(quant_config, "is_fp4_experts", False), + apply_routed_scaling_factor_on_output=( + True + if _use_aiter + else self.experts.should_fuse_routed_scaling_factor_in_topk + ), ) self.topk = TopK(**topk_kwargs) diff --git a/python/sglang/srt/models/deepseek_v4.py b/python/sglang/srt/models/deepseek_v4.py index b8df0c313..1535a81db 100644 --- a/python/sglang/srt/models/deepseek_v4.py +++ b/python/sglang/srt/models/deepseek_v4.py @@ -58,6 +58,7 @@ from sglang.srt.layers.logits_processor import LogitsProcessor from sglang.srt.layers.moe import get_moe_a2a_backend from sglang.srt.layers.moe.fused_moe_triton import FusedMoE from sglang.srt.layers.quantization.fp8_kernel import sglang_per_token_group_quant_fp8 +from sglang.srt.layers.rotary_embedding import get_rope_wrapper from sglang.srt.layers.utils import PPMissingLayer, get_layer_id from sglang.srt.layers.utils.cp_utils import ( cp_all_gather_rerange_output, @@ -76,6 +77,12 @@ from sglang.srt.model_loader.utils import maybe_executor_submit, should_async_lo from sglang.srt.model_loader.weight_utils import default_weight_loader from sglang.srt.models.dbrx import ReplicatedLinear from sglang.srt.models.deepseek_v2 import ParallelLMHead, _is_cuda, _is_hip, _is_npu + +if not _is_hip: + from sglang.srt.layers.utils.cp_utils import ( + prepare_context_parallel_metadata, + ) + from sglang.srt.server_args import get_global_server_args from sglang.srt.utils import ( LazyValue, @@ -94,6 +101,9 @@ if TYPE_CHECKING: from sglang.srt.layers.attention.deepseek_v4_backend import ( DeepseekV4AttnBackend, ) + from sglang.srt.layers.attention.deepseek_v4_backend_hip_radix import ( + DeepseekV4HipRadixBackend, + ) from sglang.srt.layers.quantization import QuantizationConfig from sglang.srt.mem_cache.deepseek_v4_memory_pool import DeepSeekV4TokenToKVPool from sglang.srt.model_executor.forward_batch_info import ForwardBatch @@ -200,6 +210,16 @@ class MQALayer(nn.Module): rope_base = config.compress_rope_theta if self.compress_ratio else rope_theta + self.rotary_emb = get_rope_wrapper( + head_size=self.rope_head_dim, + rotary_dim=self.rope_head_dim, + max_position=config.max_position_embeddings, + base=rope_base, + rope_scaling=rope_scaling, + is_neox_style=False, + device=get_global_server_args().device, + ) + from sglang.srt.layers.deepseek_v4_rope import precompute_freqs_cis assert self.compress_ratio in {0, 4, 128} @@ -243,6 +263,7 @@ class MQALayer(nn.Module): head_dim=self.head_dim, rotate=False, prefix=add_prefix("compressor", prefix), + rotary_emb=getattr(self, "rotary_emb", None), ) if self.compress_ratio == 4: self.indexer = C4Indexer( @@ -252,10 +273,11 @@ class MQALayer(nn.Module): quant_config=quant_config, prefix=add_prefix("indexer", prefix), alt_streams=self.alt_streams_indexer, + rotary_emb=getattr(self, "rotary_emb", None), ) self.attn_sink = nn.Parameter(torch.empty(self.n_heads, dtype=torch.float32)) - self.fuse_wqa_wkv = envs.SGLANG_OPT_FUSE_WQA_WKV.get() + self.fuse_wqa_wkv = not _is_hip and envs.SGLANG_OPT_FUSE_WQA_WKV.get() if self.fuse_wqa_wkv: self.wqkv_a = ReplicatedLinear( self.hidden_size, @@ -409,7 +431,7 @@ class MQALayer(nn.Module): x: torch.Tensor, positions: torch.Tensor, forward_batch: ForwardBatch, - attn_backend: DeepseekV4AttnBackend, + attn_backend, q_out: Optional[torch.Tensor] = None, ) -> torch.Tensor: assert self.alt_streams is not None @@ -469,7 +491,7 @@ class MQALayer(nn.Module): x: torch.Tensor, positions: torch.Tensor, forward_batch: ForwardBatch, - attn_backend: DeepseekV4AttnBackend, + attn_backend, q_out: Optional[torch.Tensor] = None, ) -> Tuple[torch.Tensor, Optional[torch.Tensor]]: if self.fuse_wqa_wkv: @@ -530,7 +552,10 @@ class MQALayer(nn.Module): attn_backend = forward_batch.attn_backend if TYPE_CHECKING: - assert isinstance(attn_backend, DeepseekV4AttnBackend) + assert isinstance( + attn_backend, + (DeepseekV4AttnBackend, DeepseekV4HipRadixBackend), + ) enable_multi_stream = ( envs.SGLANG_OPT_USE_MULTI_STREAM_OVERLAP.get() @@ -717,6 +742,22 @@ class DeepseekV4DecoderLayer(nn.Module): ) return y, post.squeeze(-1), comb, norm is not None + if _is_hip and envs.SGLANG_OPT_USE_AITER_MHC_PRE.get(): + from aiter.ops.mhc import mhc_pre + + post, comb, y = mhc_pre( + residual=x, + fn=hc_fn, + hc_scale=hc_scale, + hc_base=hc_base, + rms_eps=self.rms_norm_eps, + hc_pre_eps=self.hc_eps, + hc_sinkhorn_eps=self.hc_eps, + hc_post_mult_value=2.0, + sinkhorn_repeat=self.hc_sinkhorn_iters, + ) + return y, post.squeeze(-1), comb, False + if envs.SGLANG_OPT_DEEPGEMM_HC_PRENORM.get(): import deep_gemm @@ -765,6 +806,13 @@ class DeepseekV4DecoderLayer(nn.Module): return mhc_post(x, residual, post, comb) + elif _is_hip and envs.SGLANG_OPT_USE_AITER_MHC_POST.get(): + from aiter.ops.mhc import mhc_post + + result = torch.empty_like(residual) + mhc_post(result, x, residual, post, comb) + return result + assert residual.shape == (x.shape[0], self.hc_mult, x.shape[-1]) assert post.shape == (x.shape[0], self.hc_mult) assert comb.shape == (x.shape[0], self.hc_mult, self.hc_mult) @@ -1284,7 +1332,7 @@ class DeepseekV4ForCausalLM(nn.Module): cache_compressor_weight = {} COMPRESSOR_PART = ".compressor.w" - fuse_wqa_wkv = envs.SGLANG_OPT_FUSE_WQA_WKV.get() + fuse_wqa_wkv = not _is_hip and envs.SGLANG_OPT_FUSE_WQA_WKV.get() cache_wqkv_a_weight: dict[str, dict[str, torch.Tensor]] = {} def auto_weight_loader(module):