diff --git a/python/sglang/srt/layers/attention/nsa_backend.py b/python/sglang/srt/layers/attention/nsa_backend.py index dd3754f1b4bb..aa93f20ca31e 100644 --- a/python/sglang/srt/layers/attention/nsa_backend.py +++ b/python/sglang/srt/layers/attention/nsa_backend.py @@ -35,6 +35,7 @@ seqlens_expand_triton, ) from sglang.srt.layers.dp_attention import get_attention_tp_size +from sglang.srt.mem_cache.memory_pool import MLAKVCacheLayout from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode from sglang.srt.utils import is_cuda, is_hip @@ -302,15 +303,13 @@ def __init__( ) self.use_nsa = is_deepseek_nsa(model_runner.model_config.hf_config) assert self.use_nsa, "NSA backend only supports DeepSeek NSA" - self.nsa_kv_cache_store_fp8 = ( - model_runner.token_to_kv_pool.nsa_kv_cache_store_fp8 - ) + self.kv_cache_layout = model_runner.token_to_kv_pool.kv_cache_layout + self.kv_cache_size = model_runner.token_to_kv_pool.kv_cache_size self.nsa_index_topk = get_nsa_index_topk(model_runner.model_config.hf_config) self.max_context_len = model_runner.model_config.context_len self.num_q_heads = ( model_runner.model_config.num_attention_heads // get_attention_tp_size() ) - self.kv_cache_dim = model_runner.token_to_kv_pool.kv_cache_dim self.qk_nope_head_dim = model_runner.model_config.qk_nope_head_dim self.kv_lora_rank = model_runner.model_config.kv_lora_rank self.qk_rope_head_dim = model_runner.model_config.qk_rope_head_dim @@ -319,11 +318,12 @@ def __init__( self.req_to_token = model_runner.req_to_token_pool.req_to_token self.use_mha: bool = False + nsa_prefill_backend = model_runner.server_args.nsa_prefill_backend + self.prefill_is_flashmla_auto = nsa_prefill_backend == "flashmla_auto" self.nsa_prefill_impl: _NSA_IMPL_T = ( - model_runner.server_args.nsa_prefill_backend + None if self.prefill_is_flashmla_auto else nsa_prefill_backend ) self.nsa_decode_impl: _NSA_IMPL_T = model_runner.server_args.nsa_decode_backend - self.enable_auto_select_prefill_impl = self.nsa_prefill_impl == "flashmla_auto" self._arange_buf = torch.arange(16384, device=self.device, dtype=torch.int32) @@ -383,8 +383,22 @@ def _transform_table_1_to_real(self, page_table: torch.Tensor) -> torch.Tensor: ) return page_table[:, strided_indices] // page_size + @staticmethod + def should_use_decode_backend(forward_mode: ForwardMode) -> bool: + """Whether the given forward mode dispatches to nsa_decode_backend. + + This is answerable from static config alone (does not depend on + set_nsa_impl), so it is safe to call before init_forward_metadata. + """ + return ( + forward_mode.is_decode_or_idle() + or forward_mode.is_target_verify() + or forward_mode.is_draft_extend(include_v2=True) + ) + def init_forward_metadata(self, forward_batch: ForwardBatch): - """Init the metadata for a forward pass.""" + """Sets the metadata as fields of NativeSparseAttnBackend for a single forward pass. + Using the forward metadata for more than one forward pass is a bug.""" batch_size = forward_batch.batch_size device = forward_batch.seq_lens.device @@ -403,10 +417,16 @@ def init_forward_metadata(self, forward_batch: ForwardBatch): ] page_table_1_flattened = None + # Topk indices needs to transform from indices in each request to indices + # in the flattened array of all requests by + # topk_indices = topk_indices + topk_indices_offset + # Offset is repeated once for each token in the request in a flattened array. + # e.g, for three requests with lengths [2, 3, 4], + # topk_indices_offset = [0, 0, 2, 2, 2, 5, 5, 5, 5] + # Only used for prefill with TopkTransformMethod.RAGGED. topk_indices_offset = None - # Centralized dispatch: decide all strategies for this batch - self.set_nsa_prefill_impl(forward_batch) + self.set_nsa_impl(forward_batch) topk_transform_method = self.get_topk_transform_method( forward_batch.forward_mode ) @@ -643,6 +663,7 @@ def init_forward_metadata(self, forward_batch: ForwardBatch): seq_len_q=1, ) if self.nsa_decode_impl == "flashmla_kv" + or self.nsa_prefill_impl == "flashmla_kv" else None ), paged_mqa_schedule_metadata=paged_mqa_schedule_metadata, @@ -787,9 +808,8 @@ def init_forward_metadata_capture_cuda_graph( forward_mode: ForwardMode, spec_info: Optional[SpecInput], ): - self.set_nsa_prefill_impl(forward_batch=None) + self.set_nsa_impl(forward_batch=None) - """Initialize forward metadata for capturing CUDA graph.""" if forward_mode.is_decode_or_idle(): # Normal Decode # Get sequence information @@ -944,10 +964,9 @@ def init_forward_metadata_replay_cuda_graph( seq_lens_cpu: Optional[torch.Tensor], out_cache_loc: Optional[torch.Tensor] = None, ): - """Initialize forward metadata for replaying CUDA graph.""" assert seq_lens_cpu is not None - self.set_nsa_prefill_impl(forward_batch=None) + self.set_nsa_impl(forward_batch=None) seq_lens = seq_lens[:bs] seq_lens_cpu = seq_lens_cpu[:bs] @@ -1104,16 +1123,15 @@ def init_forward_metadata_replay_cuda_graph_from_precomputed( precomputed: PrecomputedMetadata, forward_mode: ForwardMode, ): - """Fast path: copy precomputed metadata to this backend's metadata. - - This function only performs copy operations, no computation. + """Compared to init_forward_metadata_replay_cuda_graph, + this function copies precomputed metadata instead of computing them. Args: bs: Batch size precomputed: Precomputed metadata to copy from forward_mode: Forward mode """ - self.set_nsa_prefill_impl(forward_batch=None) + self.set_nsa_impl(forward_batch=None) metadata = self.decode_cuda_graph_metadata[bs] @@ -1275,10 +1293,7 @@ def forward_extend( nsa_impl = ( self.nsa_decode_impl - if ( - forward_batch.forward_mode.is_target_verify() - or forward_batch.forward_mode.is_draft_extend(include_v2=True) - ) + if self.should_use_decode_backend(forward_batch.forward_mode) else self.nsa_prefill_impl ) @@ -1399,18 +1414,18 @@ def forward_extend( if q_rope is not None: q_all = concat_mla_absorb_q_general(q_nope, q_rope) - if topk_transform_method == TopkTransformMethod.RAGGED: - if any(forward_batch.extend_prefix_lens_cpu): - page_table_1_flattened = ( - self.forward_metadata.page_table_1_flattened - ) - assert page_table_1_flattened is not None - kv_cache = dequantize_k_cache_paged( - kv_cache, page_table_1_flattened - ) - else: - kv_cache = _cat([k, k_rope], dim=-1) - page_table_1 = topk_indices + if topk_transform_method != TopkTransformMethod.RAGGED: + raise ValueError( + "Internal error: Unexpected topk transform method for NSA backend flashmla_sparse." + ) + + if any(forward_batch.extend_prefix_lens_cpu): + page_table_1_flattened = self.forward_metadata.page_table_1_flattened + assert page_table_1_flattened is not None + kv_cache = dequantize_k_cache_paged(kv_cache, page_table_1_flattened) + else: + kv_cache = _cat([k, k_rope], dim=-1) + page_table_1 = topk_indices return self._forward_flashmla_sparse( q_all=q_all, @@ -1712,10 +1727,14 @@ def _forward_flashmla_kv( # TODO the 2nd dim is seq_len_q, need to be >1 when MTP q_all = q_all.view(-1, 1, layer.tp_q_head_num, layer.head_dim) - kv_cache = kv_cache.view(-1, self.real_page_size, 1, self.kv_cache_dim) + kv_cache = kv_cache.view(-1, self.real_page_size, 1, self.kv_cache_size) assert self.real_page_size == 64, "only page size 64 is supported" - if not self.nsa_kv_cache_store_fp8: + if self.kv_cache_layout != MLAKVCacheLayout.FP8_NOPE_WITH_BLOCK_SCALE_BF16_ROPE: + assert ( + self.kv_cache_layout != MLAKVCacheLayout.FP8_NOPE_FP8_ROPE + ), "Internal error: NSA backend flashmla_kv does not support FP8_NOPE_FP8_ROPE" + # inefficiently quantize the whole cache kv_cache = quantize_k_cache(kv_cache) @@ -1937,6 +1956,9 @@ def _forward_trtllm( merge_query = q_rope is not None if self.kv_cache_dtype == torch.float8_e4m3fn: + assert ( + self.kv_cache_layout == MLAKVCacheLayout.FP8_NOPE_FP8_ROPE + ), "Internal error: trtllm mla backend only supports FP8_NOPE_FP8_ROPE" # For FP8 path, we quantize the query and rope parts and merge them into a single tensor # Note: rope application in deepseek_v2.py:forward_absorb_prepare is skipped for FP8 decode path of this trtllm_mla backend assert q_rope is not None, "For FP8 path q_rope should not be None." @@ -1973,7 +1995,9 @@ def _forward_trtllm( ) k_cache = forward_batch.token_to_kv_pool.get_key_buffer(layer.layer_id) - kv_cache = k_cache.view(-1, self.real_page_size, self.kv_cache_dim).unsqueeze(1) + kv_cache = k_cache.view(-1, self.real_page_size, self.kv_cache_size).unsqueeze( + 1 + ) if merge_query: q_nope = q.view(-1, layer.tp_q_head_num, layer.v_head_dim) @@ -2009,7 +2033,7 @@ def _forward_trtllm( _, num_heads, head_dim = q_all.shape q = q_all.view(batch_size, 1, num_heads, head_dim) - kv = kv_cache.view(-1, 1, self.real_page_size, self.kv_cache_dim) + kv = kv_cache.view(-1, 1, self.real_page_size, self.kv_cache_size) block_tables = page_table_1.unsqueeze(1) seq_lens = metadata.cache_seqlens_int32 if seq_lens is None else seq_lens @@ -2056,9 +2080,10 @@ def get_cuda_graph_seq_len_fill_value(self): """Get the fill value for sequence length in CUDA graph.""" return 1 - def set_nsa_prefill_impl(self, forward_batch: Optional[ForwardBatch] = None): + def set_nsa_impl(self, forward_batch: Optional[ForwardBatch] = None): """ - Decide all attention prefill dispatch strategies for this batch. + Decide all attention dispatch strategies for this batch. + Sets nsa_prefill_impl, nsa_decode_impl and use_mha depending on forward mode. """ from sglang.srt.utils import get_device_sm, is_blackwell @@ -2087,9 +2112,17 @@ def set_nsa_prefill_impl(self, forward_batch: Optional[ForwardBatch] = None): else: self.use_mha = False # Decode/verify always use MLA + # forward_batch is None only for cudagraph + if forward_batch is None or forward_batch.forward_mode.is_decode_or_idle(): + assert self.nsa_decode_impl != "flashmla_auto" + return + # Set MLA implementation only if not using MHA - if not self.use_mha and self.enable_auto_select_prefill_impl: - if self.nsa_kv_cache_store_fp8: + if not self.use_mha and self.prefill_is_flashmla_auto: + if ( + self.kv_cache_layout + == MLAKVCacheLayout.FP8_NOPE_WITH_BLOCK_SCALE_BF16_ROPE + ): if ( is_blackwell() and forward_batch is not None @@ -2109,19 +2142,13 @@ def set_nsa_prefill_impl(self, forward_batch: Optional[ForwardBatch] = None): def get_topk_transform_method( self, forward_mode: Optional[ForwardMode] = None ) -> TopkTransformMethod: - """ - SGLANG_NSA_FUSE_TOPK controls whether to fuse the topk transform into the topk kernel. - This method is used to select the topk transform method which can be fused or unfused. - """ - if ( - # disable for MTP - self.nsa_kv_cache_store_fp8 + if forward_mode is None or forward_mode.is_decode_or_idle(): + return TopkTransformMethod.PAGED + elif ( + self.kv_cache_layout == MLAKVCacheLayout.FP8_NOPE_WITH_BLOCK_SCALE_BF16_ROPE and self.nsa_prefill_impl == "flashmla_sparse" ): topk_transform_method = TopkTransformMethod.RAGGED - - if forward_mode is not None and (forward_mode.is_decode_or_idle()): - topk_transform_method = TopkTransformMethod.PAGED else: topk_transform_method = TopkTransformMethod.PAGED return topk_transform_method @@ -2133,11 +2160,12 @@ def get_indexer_metadata( forward_batch.hisparse_coordinator is not None and forward_batch.forward_mode.is_decode_or_idle() ) + topk_transform_method = self.get_topk_transform_method( + forward_batch.forward_mode + ) return NSAIndexerMetadata( attn_metadata=self.forward_metadata, - topk_transform_method=self.get_topk_transform_method( - forward_batch.forward_mode - ), + topk_transform_method=topk_transform_method, paged_mqa_schedule_metadata=self.forward_metadata.paged_mqa_schedule_metadata, force_unfused_topk=force_unfused, ) @@ -2229,7 +2257,7 @@ def init_forward_metadata_replay_cuda_graph( # Set nsa_prefill_impl for first 3 backends (required by the method) for i in range(3): - self.attn_backends[i].set_nsa_prefill_impl(forward_batch=None) + self.attn_backends[i].set_nsa_impl(forward_batch=None) # Prepare FlashMLA tensors if needed flashmla_num_splits_src = None diff --git a/python/sglang/srt/managers/hisparse_coordinator.py b/python/sglang/srt/managers/hisparse_coordinator.py index 92ef22f404cf..cc8c41eef2b2 100644 --- a/python/sglang/srt/managers/hisparse_coordinator.py +++ b/python/sglang/srt/managers/hisparse_coordinator.py @@ -53,7 +53,7 @@ def __init__( host_size=0, page_size=1, # for simplicity, we set page size to 1 to enable backup one token at a time layout="layer_first", - override_kv_cache_dim=self.mem_pool_device.kv_cache_dim, + override_kv_cache_dim=self.mem_pool_device.kv_cache_size, ) max_num_reqs = req_to_token_pool.size diff --git a/python/sglang/srt/mem_cache/hisparse_memory_pool.py b/python/sglang/srt/mem_cache/hisparse_memory_pool.py index 5af8d257ad6b..542044fb5afd 100644 --- a/python/sglang/srt/mem_cache/hisparse_memory_pool.py +++ b/python/sglang/srt/mem_cache/hisparse_memory_pool.py @@ -10,7 +10,7 @@ BaseTokenToKVPoolAllocator, PagedTokenToKVPoolAllocator, ) -from sglang.srt.mem_cache.memory_pool import NSATokenToKVPool +from sglang.srt.mem_cache.memory_pool import MLAKVCacheLayout, NSATokenToKVPool from sglang.srt.utils import is_cuda, is_hip # sgl_kernel.kvcacheio is only available in CUDA/ROCm sgl-kernel builds (not XPU/MPS/NPU/CPU). @@ -28,9 +28,10 @@ def transfer_kv_all_layer_mla(*args, **kwargs): class HiSparseNSATokenToKVPool(NSATokenToKVPool): + def __init__( self, - size: int, + max_total_num_tokens: int, page_size: int, kv_lora_rank: int, dtype: torch.dtype, @@ -39,13 +40,14 @@ def __init__( device: str, index_head_dim: int, enable_memory_saver: bool, - kv_cache_dim: int, + kv_cache_layout: MLAKVCacheLayout, + kv_cache_size: int, start_layer: Optional[int] = None, end_layer: Optional[int] = None, host_to_device_ratio: int = 2, ): super().__init__( - size=size, + max_total_num_tokens=max_total_num_tokens, page_size=page_size, kv_lora_rank=kv_lora_rank, dtype=dtype, @@ -54,12 +56,13 @@ def __init__( device=device, index_head_dim=index_head_dim, enable_memory_saver=enable_memory_saver, - kv_cache_dim=kv_cache_dim, + kv_cache_layout=kv_cache_layout, + kv_cache_size=kv_cache_size, start_layer=start_layer, end_layer=end_layer, - index_buf_size=size * host_to_device_ratio, + index_buf_size=max_total_num_tokens * host_to_device_ratio, ) - self.bytes_per_token = self.kv_cache_dim * self.dtype.itemsize + self.bytes_per_token = self.kv_cache_size * self.dtype.itemsize def register_mapping(self, full_to_hisparse_device_index_mapping: torch.Tensor): self.full_to_hisparse_device_index_mapping = ( diff --git a/python/sglang/srt/mem_cache/memory_pool.py b/python/sglang/srt/mem_cache/memory_pool.py index df5d54223c22..80998b6207d9 100644 --- a/python/sglang/srt/mem_cache/memory_pool.py +++ b/python/sglang/srt/mem_cache/memory_pool.py @@ -29,6 +29,7 @@ import logging from contextlib import contextmanager, nullcontext from dataclasses import dataclass, fields +from enum import Enum, auto from typing import TYPE_CHECKING, Any, List, Optional, Tuple, Union import numpy as np @@ -42,7 +43,6 @@ from sglang.srt.environ import envs from sglang.srt.layers.attention.nsa import index_buf_accessor from sglang.srt.layers.attention.nsa.quant_k_cache import ( - quantize_k_cache, quantize_k_cache_separate, ) from sglang.srt.layers.radix_attention import RadixAttention @@ -60,8 +60,32 @@ is_npu, next_power_of_2, ) +from sglang.srt.utils.common import is_float4_e2m1fn_x2 +from sglang.srt.utils.custom_op import register_custom_op from sglang.srt.utils.torch_memory_saver_adapter import TorchMemorySaverAdapter +store_cache = register_custom_op(store_cache, mutates_args=["k_cache", "v_cache"]) + + +class MLAKVCacheLayout(Enum): + """Layout of MLA kv cache.""" + + FP4 = auto() # fp4 k_nope + fp4 k_rope + BF16 = auto() # bf16 k_nope + bf16 k_rope + FP8_NOPE_FP8_ROPE = auto() # fp8 k_nope + fp8 k_rope + FP8_NOPE_WITH_BLOCK_SCALE_BF16_ROPE = ( + auto() + ) # fp8 k_nope + fp32 block-scale + bf16 k_rope + + def is_homogeneous(self: MLAKVCacheLayout) -> bool: + """Check if the layout is homogeneous, i.e. all tensors have the same dtype.""" + return self in [ + MLAKVCacheLayout.FP8_NOPE_FP8_ROPE, + MLAKVCacheLayout.BF16, + MLAKVCacheLayout.FP4, + ] + + if TYPE_CHECKING: from sglang.srt.managers.cache_controller import LayerDoneCounter from sglang.srt.managers.schedule_batch import Req @@ -1432,9 +1456,10 @@ def get_mla_kv_buffer( class MLATokenToKVPool(KVCache): + def __init__( self, - size: int, + max_total_num_tokens: int, page_size: int, dtype: torch.dtype, kv_lora_rank: int, @@ -1442,13 +1467,14 @@ def __init__( layer_num: int, device: str, enable_memory_saver: bool, + kv_cache_layout: MLAKVCacheLayout, + kv_cache_size: int, start_layer: Optional[int] = None, end_layer: Optional[int] = None, use_nsa: bool = False, - override_kv_cache_dim: Optional[int] = None, ): super().__init__( - size, + max_total_num_tokens, page_size, dtype, layer_num, @@ -1461,18 +1487,8 @@ def __init__( self.kv_lora_rank = kv_lora_rank self.qk_rope_head_dim = qk_rope_head_dim self.use_nsa = use_nsa - self.nsa_kv_cache_store_fp8 = ( - use_nsa - and dtype == torch.float8_e4m3fn - and override_kv_cache_dim is not None - ) - # When override_kv_cache_dim is provided with nsa model, we assume the - # override kv cache dim is correct and use it directly. - self.kv_cache_dim = ( - override_kv_cache_dim - if self.nsa_kv_cache_store_fp8 - else (kv_lora_rank + qk_rope_head_dim) - ) + self.kv_cache_layout = kv_cache_layout + self.kv_cache_size = kv_cache_size self._create_buffers() @@ -1483,7 +1499,7 @@ def __init__( ) if not use_nsa: # NSA will allocate indexer KV cache later and then log the total size - self._finalize_allocation_log(size) + self._finalize_allocation_log(max_total_num_tokens) def _create_buffers(self): with self.memory_saver_adapter.region(GPU_MEMORY_TYPE_KV_CACHE): @@ -1495,7 +1511,7 @@ def _create_buffers(self): # The padded slot 0 is used for writing dummy outputs from padded tokens. self.kv_buffer = [ torch.zeros( - (self.size + self.page_size, 1, self.kv_cache_dim), + (self.size + self.page_size, 1, self.kv_cache_size), dtype=self.store_dtype, device=self.device, ) @@ -1552,7 +1568,13 @@ def set_kv_buffer( cache_v: torch.Tensor, ): layer_id = layer.layer_id - assert not self.nsa_kv_cache_store_fp8 + assert ( + self.kv_cache_layout != MLAKVCacheLayout.FP8_NOPE_WITH_BLOCK_SCALE_BF16_ROPE + ), "Internal error: FP8_NOPE_WITH_BLOCK_SCALE_BF16_ROPE must use set_mla_kv_buffer" + assert ( + self.kv_cache_layout.is_homogeneous() + ), "Internal error: MLAKVCacheLayout must be homogeneous" + if cache_k.dtype != self.dtype: cache_k = cache_k.to(self.dtype) @@ -1572,7 +1594,7 @@ def set_mla_kv_buffer( ): layer_id = layer.layer_id - if self.nsa_kv_cache_store_fp8: + if self.kv_cache_layout == MLAKVCacheLayout.FP8_NOPE_WITH_BLOCK_SCALE_BF16_ROPE: # OPTIMIZATION: Quantize k_nope and k_rope separately to avoid concat overhead # This also enables reuse of set_mla_kv_buffer_triton two-tensor write path # quantize_k_cache_separate returns (nope_part, rope_part) as uint8 bytes @@ -1590,6 +1612,9 @@ def set_mla_kv_buffer( cache_k_rope_fp8, ) else: + assert ( + self.kv_cache_layout.is_homogeneous() + ), "Internal error: MLAKVCacheLayout must be homogeneous" if cache_k_nope.dtype != self.dtype: cache_k_nope = cache_k_nope.to(self.dtype) cache_k_rope = cache_k_rope.to(self.dtype) @@ -1657,6 +1682,43 @@ def load_cpu_copy(self, kv_cache_cpu, indices): class MLATokenToKVPoolFP4(MLATokenToKVPool): + def __init__( + self, + max_total_num_tokens: int, + page_size: int, + dtype: torch.dtype, + kv_lora_rank: int, + qk_rope_head_dim: int, + layer_num: int, + device: str, + enable_memory_saver: bool, + kv_cache_layout: MLAKVCacheLayout, + kv_cache_size: int, + start_layer: Optional[int] = None, + end_layer: Optional[int] = None, + ): + assert is_float4_e2m1fn_x2( + dtype + ), "Internal error: dtype must be float4_e2m1fn_x2" + assert ( + kv_cache_layout == MLAKVCacheLayout.FP4 + ), "Internal error: kv_cache_layout must be FP4" + super().__init__( + max_total_num_tokens, + page_size, + dtype, + kv_lora_rank, + qk_rope_head_dim, + layer_num, + device, + enable_memory_saver, + kv_cache_layout, + kv_cache_size, + start_layer, + end_layer, + use_nsa=False, + ) + def _create_buffers(self): with self.memory_saver_adapter.region(GPU_MEMORY_TYPE_KV_CACHE): with ( @@ -1667,7 +1729,7 @@ def _create_buffers(self): # The padded slot 0 is used for writing dummy outputs from padded tokens. m = self.size + self.page_size n = 1 # head_num - k = self.kv_cache_dim # head_dim + k = self.kv_cache_size # head_dim scale_block_size = 16 self.store_dtype = torch.uint8 @@ -1721,7 +1783,6 @@ def set_kv_buffer( cache_v: torch.Tensor, ): layer_id = layer.layer_id - assert not self.nsa_kv_cache_store_fp8 if cache_k.dtype != self.dtype: from sglang.srt.layers.quantization.kvfp4_tensor import KVFP4QuantizeUtil @@ -1746,43 +1807,35 @@ def set_mla_kv_buffer( ): layer_id = layer.layer_id - if self.nsa_kv_cache_store_fp8: - # original cache_k: (num_tokens, num_heads 1, hidden 576); we unsqueeze the page_size=1 dim here - # TODO no need to cat - cache_k = torch.cat([cache_k_nope, cache_k_rope], dim=-1) - cache_k = quantize_k_cache(cache_k.unsqueeze(1)).squeeze(1) - cache_k = cache_k.view(self.store_dtype) - self.kv_buffer[layer_id - self.start_layer][loc] = cache_k - else: - if cache_k_nope.dtype != self.dtype: - from sglang.srt.layers.quantization.kvfp4_tensor import ( - KVFP4QuantizeUtil, - ) - - cache_k_nope_fp4, cache_k_nope_fp4_sf = ( - KVFP4QuantizeUtil.batched_quantize(cache_k_nope) - ) - cache_k_rope_fp4, cache_k_rope_fp4_sf = ( - KVFP4QuantizeUtil.batched_quantize(cache_k_rope) - ) - - if self.store_dtype != self.dtype: - cache_k_nope = cache_k_nope.view(self.store_dtype) - cache_k_rope = cache_k_rope.view(self.store_dtype) + if cache_k_nope.dtype != self.dtype: + from sglang.srt.layers.quantization.kvfp4_tensor import ( + KVFP4QuantizeUtil, + ) - set_mla_kv_buffer_triton( - self.kv_buffer[layer_id - self.start_layer], - loc, - cache_k_nope_fp4, - cache_k_rope_fp4, + cache_k_nope_fp4, cache_k_nope_fp4_sf = KVFP4QuantizeUtil.batched_quantize( + cache_k_nope ) - set_mla_kv_scale_buffer_triton( - self.kv_scale_buffer[layer_id - self.start_layer], - loc, - cache_k_nope_fp4_sf, - cache_k_rope_fp4_sf, + cache_k_rope_fp4, cache_k_rope_fp4_sf = KVFP4QuantizeUtil.batched_quantize( + cache_k_rope ) + if self.store_dtype != self.dtype: + cache_k_nope = cache_k_nope.view(self.store_dtype) + cache_k_rope = cache_k_rope.view(self.store_dtype) + + set_mla_kv_buffer_triton( + self.kv_buffer[layer_id - self.start_layer], + loc, + cache_k_nope_fp4, + cache_k_rope_fp4, + ) + set_mla_kv_scale_buffer_triton( + self.kv_scale_buffer[layer_id - self.start_layer], + loc, + cache_k_nope_fp4_sf, + cache_k_rope_fp4_sf, + ) + class NSATokenToKVPool(MLATokenToKVPool): quant_block_size = 128 @@ -1791,7 +1844,7 @@ class NSATokenToKVPool(MLATokenToKVPool): def __init__( self, - size: int, + max_total_num_tokens: int, page_size: int, kv_lora_rank: int, dtype: torch.dtype, @@ -1800,18 +1853,14 @@ def __init__( device: str, index_head_dim: int, enable_memory_saver: bool, - kv_cache_dim: int, + kv_cache_layout: MLAKVCacheLayout, + kv_cache_size: int, start_layer: Optional[int] = None, end_layer: Optional[int] = None, index_buf_size: Optional[int] = None, ): - - override_dim = ( - kv_cache_dim if kv_cache_dim != kv_lora_rank + qk_rope_head_dim else None - ) - super().__init__( - size, + max_total_num_tokens, page_size, dtype, kv_lora_rank, @@ -1819,16 +1868,17 @@ def __init__( layer_num, device, enable_memory_saver, - start_layer, - end_layer, + kv_cache_layout=kv_cache_layout, + kv_cache_size=kv_cache_size, + start_layer=start_layer, + end_layer=end_layer, use_nsa=True, - override_kv_cache_dim=override_dim, ) # self.index_k_dtype = torch.float8_e4m3fn # self.index_k_scale_dtype = torch.float32 self.index_head_dim = index_head_dim if index_buf_size is None: - index_buf_size = size + index_buf_size = max_total_num_tokens # num head == 1 and head dim == 128 for index_k in NSA assert index_head_dim == 128 @@ -1861,7 +1911,7 @@ def __init__( ) for _ in range(layer_num) ] - self._finalize_allocation_log(size) + self._finalize_allocation_log(max_total_num_tokens) def get_index_k_with_scale_buffer(self, layer_id: int) -> torch.Tensor: if self.layer_transfer_counter is not None: diff --git a/python/sglang/srt/mem_cache/memory_pool_host.py b/python/sglang/srt/mem_cache/memory_pool_host.py index 9666080d3f72..c65a6916168e 100644 --- a/python/sglang/srt/mem_cache/memory_pool_host.py +++ b/python/sglang/srt/mem_cache/memory_pool_host.py @@ -1726,7 +1726,7 @@ def __init__( pin_memory, device, allocator_type, - override_kv_cache_dim=device_pool.kv_cache_dim, + override_kv_cache_dim=device_pool.kv_cache_size, ) self.indexer_page_stride_size = ( self.indexer_size_per_token * self.page_size * self.indexer_dtype.itemsize diff --git a/python/sglang/srt/model_executor/model_runner_kv_cache_mixin.py b/python/sglang/srt/model_executor/model_runner_kv_cache_mixin.py index 291d9515df17..583b2b476ec8 100644 --- a/python/sglang/srt/model_executor/model_runner_kv_cache_mixin.py +++ b/python/sglang/srt/model_executor/model_runner_kv_cache_mixin.py @@ -24,6 +24,7 @@ HybridReqToTokenPool, MHATokenToKVPool, MHATokenToKVPoolFP4, + MLAKVCacheLayout, MLATokenToKVPool, MLATokenToKVPoolFP4, NSATokenToKVPool, @@ -235,44 +236,54 @@ def handle_max_mamba_cache(self: ModelRunner, total_rest_memory): ) return total_rest_memory - mamba_state_memory - def calculate_mla_kv_cache_dim(self: ModelRunner) -> int: + def get_mla_kv_cache_layout(self: ModelRunner) -> tuple[MLAKVCacheLayout, int]: + """ + Returns: + - MLAKVCacheLayout: layout of MLA kv cache + - int: dimension of MLA kv cache if MLAKVCacheLayout is BF16/FP4 or number of bytes otherwise. + """ + is_nsa_model = is_deepseek_nsa(self.model_config.hf_config) kv_cache_dtype = self.kv_cache_dtype kv_lora_rank = self.model_config.kv_lora_rank qk_rope_head_dim = self.model_config.qk_rope_head_dim kv_cache_dim = kv_lora_rank + qk_rope_head_dim # default mla kv cache dim - # For non-NSA models, MLA kv cache dim is simply kv_lora_rank + qk_rope_head_dim - if not is_nsa_model: - return kv_cache_dim - - # TRTLLM backend does not override kv_cache_dim for MLA kv cache - # Assuming nsa prefill and decode backends are the same when using trtllm MLA backend, - # since it is not compatible for trtllm and other mla attn backend due to the different - # kv cache layout. - if ( - self.server_args.nsa_prefill_backend == "trtllm" - or self.server_args.nsa_decode_backend == "trtllm" - ): - return kv_cache_dim + if kv_cache_dtype == torch.bfloat16: + return MLAKVCacheLayout.BF16, kv_cache_dim + + if is_float4_e2m1fn_x2(kv_cache_dtype): + return MLAKVCacheLayout.FP4, kv_cache_dim + + if is_nsa_model: + prefill_is_trtllm = self.server_args.nsa_prefill_backend == "trtllm" + decode_is_trtllm = self.server_args.nsa_decode_backend == "trtllm" + assert ( + prefill_is_trtllm == decode_is_trtllm + ), "NSA backend trtllm cannot be mixed with other NSA backends." + + use_block_scale = not prefill_is_trtllm + else: + use_block_scale = False + + if not use_block_scale: + return MLAKVCacheLayout.FP8_NOPE_FP8_ROPE, kv_cache_dim quant_block_size = NSATokenToKVPool.quant_block_size rope_storage_dtype = NSATokenToKVPool.rope_storage_dtype - # Calculate override_kv_cache_dim for FP8 storage for non-trtllm attention backends: - # kv_lora_rank + scale storage (kv_lora_rank // quant_block_size * 4 bytes) + rope dimension storage - # Note: rope dimension is stored in original dtype (bf16), not quantized to fp8 - if kv_cache_dtype == torch.float8_e4m3fn: - assert ( - kv_lora_rank % quant_block_size == 0 - ), f"kv_lora_rank {kv_lora_rank} must be multiple of quant_block_size {quant_block_size}" - return ( - kv_lora_rank - + kv_lora_rank // quant_block_size * 4 - + qk_rope_head_dim * rope_storage_dtype.itemsize - ) + assert ( + kv_lora_rank % quant_block_size == 0 + ), f"kv_lora_rank {kv_lora_rank} must be multiple of quant_block_size {quant_block_size}" + assert ( + rope_storage_dtype == torch.bfloat16 + ), "Internal error: unexpected rope storage dtype" - return kv_cache_dim + return MLAKVCacheLayout.FP8_NOPE_WITH_BLOCK_SCALE_BF16_ROPE, ( + kv_lora_rank + + kv_lora_rank // quant_block_size * 4 + + qk_rope_head_dim * rope_storage_dtype.itemsize + ) def _resolve_hybrid_swa_tokens( self: ModelRunner, token_capacity: int @@ -539,15 +550,17 @@ def _init_pools(self: ModelRunner): end_layer=self.end_layer, ) elif self.use_mla_backend and is_nsa_model: + kv_cache_layout, kv_cache_size = self.get_mla_kv_cache_layout() nsa_pool_kwargs = dict( - size=self.max_total_num_tokens, + max_total_num_tokens=self.max_total_num_tokens, page_size=self.page_size, dtype=self.kv_cache_dtype, kv_lora_rank=self.model_config.kv_lora_rank, qk_rope_head_dim=self.model_config.qk_rope_head_dim, layer_num=self.num_effective_layers, device=self.device, - kv_cache_dim=self.calculate_mla_kv_cache_dim(), + kv_cache_layout=kv_cache_layout, + kv_cache_size=kv_cache_size, enable_memory_saver=self.server_args.enable_memory_saver, start_layer=self.start_layer, end_layer=self.end_layer, @@ -565,6 +578,7 @@ def _init_pools(self: ModelRunner): self.token_to_kv_pool = NSATokenToKVPool(**nsa_pool_kwargs) elif self.use_mla_backend and not self.mambaish_config: assert not is_nsa_model + kv_cache_layout, kv_cache_size = self.get_mla_kv_cache_layout() if is_float4_e2m1fn_x2(self.kv_cache_dtype): self.token_to_kv_pool = MLATokenToKVPoolFP4( self.max_total_num_tokens, @@ -575,6 +589,8 @@ def _init_pools(self: ModelRunner): layer_num=self.num_effective_layers, device=self.device, enable_memory_saver=self.server_args.enable_memory_saver, + kv_cache_layout=kv_cache_layout, + kv_cache_size=kv_cache_size, start_layer=self.start_layer, end_layer=self.end_layer, ) @@ -588,6 +604,8 @@ def _init_pools(self: ModelRunner): layer_num=self.num_effective_layers, device=self.device, enable_memory_saver=self.server_args.enable_memory_saver, + kv_cache_layout=kv_cache_layout, + kv_cache_size=kv_cache_size, start_layer=self.start_layer, end_layer=self.end_layer, ) diff --git a/python/sglang/srt/models/deepseek_common/attention_backend_handler.py b/python/sglang/srt/models/deepseek_common/attention_backend_handler.py index dfc1d4c97241..3e3702f9267d 100644 --- a/python/sglang/srt/models/deepseek_common/attention_backend_handler.py +++ b/python/sglang/srt/models/deepseek_common/attention_backend_handler.py @@ -143,7 +143,7 @@ def handle_attention_aiter(attn, forward_batch): def handle_attention_nsa(attn, forward_batch): """ - Dispatch logic is centralized in NativeSparseAttnBackend.set_nsa_prefill_impl and executed + Dispatch logic is centralized in NativeSparseAttnBackend.set_nsa_impl and executed in init_forward_metadata. Read the decision from backend.use_mha. """ diff --git a/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mla.py b/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mla.py index fe1b3c966a71..daaab123b3c5 100644 --- a/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mla.py +++ b/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mla.py @@ -519,10 +519,23 @@ def _fuse_rope_for_trtllm_mla( Check if we should skip rope and do fused rope+quantize for TRTLLM MLA decode in fp8_e4m3 path. """ if self.current_attention_backend == "nsa": - return ( - get_global_server_args().nsa_decode_backend == "trtllm" - or get_global_server_args().nsa_prefill_backend == "trtllm" - ) and forward_batch.attn_backend.kv_cache_dtype == torch.float8_e4m3fn + is_fp8_kv_cache = ( + forward_batch.attn_backend.kv_cache_dtype == torch.float8_e4m3fn + ) + if not is_fp8_kv_cache: + return False + + # Local import to avoid circular dependency (nsa_backend imports from model layer) + from sglang.srt.layers.attention.nsa_backend import ( + NativeSparseAttnBackend, + ) + + if NativeSparseAttnBackend.should_use_decode_backend( + forward_batch.forward_mode + ): + return get_global_server_args().nsa_decode_backend == "trtllm" + else: + return get_global_server_args().nsa_prefill_backend == "trtllm" return ( self.current_attention_backend == "trtllm_mla"