diff --git a/python/sglang/srt/hardware_backend/npu/dsv4/dsv4_memory_pool.py b/python/sglang/srt/hardware_backend/npu/dsv4/dsv4_memory_pool.py index ec9fe9e70..884516b49 100644 --- a/python/sglang/srt/hardware_backend/npu/dsv4/dsv4_memory_pool.py +++ b/python/sglang/srt/hardware_backend/npu/dsv4/dsv4_memory_pool.py @@ -321,8 +321,8 @@ class DSV4NPUTokenToKVPool(DeepSeekV4TokenToKVPool): return self.get_indexer_compress_states(layer_id) return self.get_attention_compress_states(layer_id) - def _make_attn_state_pool( - self, ratio: int, enable_memory_saver: bool + def _make_compress_state_pool( + self, ratio: int, *, head_dim: int, enable_memory_saver: bool ) -> NPUCompressStatePool: # ONLINE_C128 (CUDA-only) collapses the c128 ring to size 1; the NPU fused # compressor has no online mode, so assert the config mismatch early. @@ -330,46 +330,26 @@ class DSV4NPUTokenToKVPool(DeepSeekV4TokenToKVPool): "SGLANG_OPT_USE_ONLINE_COMPRESS is incompatible with the " "NPU fused compressor (no online mode in the kernel)." ) + config = self.compressed_pool_configs[ratio] ring_size = self.get_ring_size(ratio) # A5 cache_mode=2 addresses one ring bank per request. The A3 # explicit-location path can share the smaller flat pool, but the A5 # cycle ABI needs enough physical banks for every req_pool_idx. - size = self._state_pool_size(ratio) + size = config.state_size if is_npu_arch35(): size = max(size, self.num_req_slots * ring_size) return NPUCompressStatePool( size=size, ring_size=ring_size, overlap=ratio == 4, - head_dim=self.qk_nope_head_dim + self.qk_rope_head_dim, - dtype=self.c4_state_dtype if ratio == 4 else self.c128_state_dtype, + head_dim=head_dim, + dtype=config.state_dtype, device=self.device, enable_memory_saver=enable_memory_saver, ratio=ratio, swa_page_size=self.swa_page_size, ) - def _make_indexer_state_pool( - self, ratio: int, enable_memory_saver: bool - ) -> NPUCompressStatePool: - # c4 indexer shares the c4 state pool size budget but has its own - # slot_dim (indexer_head_dim vs attention head_dim). - ring_size = self.get_ring_size(ratio) - size = self.c4_state_pool_size - if is_npu_arch35(): - size = max(size, self.num_req_slots * ring_size) - return NPUCompressStatePool( - size=size, - ring_size=ring_size, - overlap=ratio == 4, - head_dim=self.indexer_head_dim, - device=self.device, - dtype=self.c4_state_dtype, - enable_memory_saver=enable_memory_saver, - ratio=ratio, - swa_page_size=self.swa_page_size, - ) - def _make_indexer_pool( self, size: int, @@ -394,10 +374,11 @@ class DSV4NPUTokenToKVPool(DeepSeekV4TokenToKVPool): def get_contiguous_buf_infos(self) -> Tuple[List[int], List[int], List[int]]: """Main PD buffers addressed by the full KV page id.""" + indexer_pool = self._indexer_pool(4) buffers = ( self.c4_kv_pool.kv_buffer - + self.c4_indexer_kv_pool.index_k_buffer - + self.c4_indexer_kv_pool.index_scale_buffer + + indexer_pool.index_k_buffer + + indexer_pool.index_scale_buffer ) return ( [buf.data_ptr() for buf in buffers], @@ -530,16 +511,15 @@ class DSV4NPUTokenToKVPool(DeepSeekV4TokenToKVPool): ``torch.ops.custom.npu_quant_lightning_indexer`` consumes. """ item = self.layer_mapping[layer_id] - if item.compress_ratio == 4: - if from_indexer: - kv = self.c4_indexer_kv_pool.get_index_k(item.compress_layer_id) - else: - kv = self.c4_kv_pool.kv_buffer[item.compress_layer_id] - elif item.compress_ratio == 128: - assert not from_indexer, "c128 has no indexer pool" - kv = self.c128_kv_pool.kv_buffer[item.compress_layer_id] - else: + if item.compress_ratio == 0: return None + if from_indexer: + indexer_pool = self._indexer_pool(item.compress_ratio) + kv = indexer_pool.get_index_k(item.compress_layer_id) + else: + compress_pool = item.compress_kv_pool + assert compress_pool is not None, "Missing compressed KV pool" + kv = compress_pool.kv_buffer[item.compress_layer_id] if loc is not None: kv = kv.flatten(0, 1)[loc] return kv @@ -651,21 +631,17 @@ class DSV4NPUTokenToKVPool(DeepSeekV4TokenToKVPool): ratio, compress_layer_id, _ = self.layer_mapping[layer_id] device_type = kv.device.type if from_indexer: - assert ratio == 4, f"indexer only on c4 layers, got ratio={ratio}" + indexer_pool = self._indexer_pool(ratio) if device_type == "npu": - assert self.c4_indexer_kv_pool.has_npu_storage, ( + assert indexer_pool.has_npu_storage, ( "NPU index buffers not allocated — pool was init'd on CUDA?" ) - self.c4_indexer_kv_pool.set_index_k_scale( - compress_layer_id, loc, kv, kv_scale - ) + indexer_pool.set_index_k_scale(compress_layer_id, loc, kv, kv_scale) return if kv_scale is None: - self.c4_indexer_kv_pool.set_index_fused(compress_layer_id, loc, kv) + indexer_pool.set_index_fused(compress_layer_id, loc, kv) return - self.c4_indexer_kv_pool.set_index_k_scale_buffer( - compress_layer_id, loc, kv, kv_scale - ) + indexer_pool.set_index_k_scale_buffer(compress_layer_id, loc, kv, kv_scale) return compress_pool = self.c4_kv_pool if ratio == 4 else self.c128_kv_pool if device_type == "npu": @@ -690,5 +666,6 @@ class DSV4NPUTokenToKVPool(DeepSeekV4TokenToKVPool): ) -> torch.Tensor: # The indexer scale is fp16 on pre-A5 parts and fp32 on A5. assert from_indexer, "only indexer compress pool has dequant scale" - compress_layer_id = self.layer_mapping[layer_id].compress_layer_id - return self.c4_indexer_kv_pool.get_index_scale(compress_layer_id) + item = self.layer_mapping[layer_id] + indexer_pool = self._indexer_pool(item.compress_ratio) + return indexer_pool.get_index_scale(item.compress_layer_id) diff --git a/python/sglang/srt/layers/attention/deepseek_v4_backend.py b/python/sglang/srt/layers/attention/deepseek_v4_backend.py index 3a3fc1f8e..b76e513e2 100644 --- a/python/sglang/srt/layers/attention/deepseek_v4_backend.py +++ b/python/sglang/srt/layers/attention/deepseek_v4_backend.py @@ -778,6 +778,7 @@ class DeepseekV4AttnBackend( ): return PagedIndexerMetadata( page_size=self.page_size, + compressed_page_size=self.token_to_kv_pool.get_index_k_page_size(), page_table=core_attn_metadata.page_table, compressed_seq_lens=core_attn_metadata.c4_topk_lengths_raw, use_topk_v2=self.dsa_topk_backend.should_use_topk_v2() and not _is_xpu, @@ -1482,8 +1483,6 @@ class DeepseekV4AttnBackend( seq_lens_cpu_list, extend_seq_lens_cpu, strict=True ) ) - # ``swa_window_size`` on the pool is its storage page size, not the - # model's SWA window, so pass both explicitly. return SparsePrefillChunkCache.build( seq_lens=forward_batch.seq_lens.to(torch.int32), extend_seq_lens=forward_batch.extend_seq_lens.to(torch.int32), @@ -1491,7 +1490,7 @@ class DeepseekV4AttnBackend( req_to_token=self.req_to_token, full_to_swa=self.token_to_kv_pool.full_to_swa_index_mapping, swa_window_size=SWA_WINDOW, - swa_page_size=self.token_to_kv_pool.swa_window_size, + swa_page_size=self.token_to_kv_pool.swa_page_size, num_qo_tokens=num_qo_tokens, max_seq_len=max(seq_lens_cpu_list), total_swa=total_swa, @@ -1772,23 +1771,20 @@ class DeepseekV4AttnBackend( 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 + swa_page_size = token_to_kv_pool.swa_page_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 + swa_k_cache = swa_k_cache[:, : swa_page_size * k_cache_total_dim].view( + swa_k_cache.shape[0], swa_page_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_page_size = token_to_kv_pool.get_extra_key_page_size(layer_id) extra_k_cache = extra_k_cache[ - :, : page_sizes[compress_ratio] * k_cache_total_dim + :, : extra_page_size * k_cache_total_dim ].view( extra_k_cache.shape[0], - page_sizes[compress_ratio], + extra_page_size, 1, k_cache_total_dim, ) 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 index f78bd1c73..9ac37dd0a 100644 --- a/python/sglang/srt/layers/attention/deepseek_v4_backend_hip_radix.py +++ b/python/sglang/srt/layers/attention/deepseek_v4_backend_hip_radix.py @@ -507,6 +507,7 @@ class DeepseekV4HipRadixBackend( def init_forward_metadata_indexer(self, core_attn_metadata: DSV4AttnMetadata): return PagedIndexerMetadata( page_size=self.page_size, + compressed_page_size=self.token_to_kv_pool.get_index_k_page_size(), page_table=core_attn_metadata.page_table, compressed_seq_lens=core_attn_metadata.c4_topk_lengths_raw, use_topk_v2=self.dsa_topk_backend.should_use_topk_v2(), @@ -1601,23 +1602,20 @@ class DeepseekV4HipRadixBackend( 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 + swa_page_size = token_to_kv_pool.swa_page_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 + swa_k_cache = swa_k_cache[:, : swa_page_size * k_cache_total_dim].view( + swa_k_cache.shape[0], swa_page_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_page_size = token_to_kv_pool.get_extra_key_page_size(layer_id) extra_k_cache = extra_k_cache[ - :, : page_sizes[compress_ratio] * k_cache_total_dim + :, : extra_page_size * k_cache_total_dim ].view( extra_k_cache.shape[0], - page_sizes[compress_ratio], + extra_page_size, 1, k_cache_total_dim, ) diff --git a/python/sglang/srt/layers/attention/dsv4/compressor_v2.py b/python/sglang/srt/layers/attention/dsv4/compressor_v2.py index 207fdb7e5..b7e2040e9 100644 --- a/python/sglang/srt/layers/attention/dsv4/compressor_v2.py +++ b/python/sglang/srt/layers/attention/dsv4/compressor_v2.py @@ -265,7 +265,7 @@ class CompressorBackendMixin: bf16_store = False kv_scale_cache = None if compressor.is_in_indexer: - page_size = token_to_kv_pool.get_index_k_page_size() + page_size = token_to_kv_pool.get_index_k_page_size(compressor.ratio) if use_hip_fp4: kv_cache = token_to_kv_pool.get_index_k_fp4_payload_buffer(layer_id) kv_scale_cache = token_to_kv_pool.get_index_k_fp4_scale_buffer(layer_id) diff --git a/python/sglang/srt/layers/attention/dsv4/metadata.py b/python/sglang/srt/layers/attention/dsv4/metadata.py index 67e07be27..4c4741bde 100644 --- a/python/sglang/srt/layers/attention/dsv4/metadata.py +++ b/python/sglang/srt/layers/attention/dsv4/metadata.py @@ -2,7 +2,7 @@ from __future__ import annotations import warnings from dataclasses import dataclass, field, fields -from typing import TYPE_CHECKING, Any, List, Optional +from typing import Any, List, Optional import torch @@ -11,10 +11,6 @@ from sglang.srt.utils import is_hip, is_sm120_supported, is_xpu _IS_SM120 = is_sm120_supported() -if TYPE_CHECKING: - pass - - """ Some comments on the common terms used in DeepSeekV4Backend: @@ -45,7 +41,7 @@ positions: Some other notes: c4_ / c128_: means "compressed by 4" / "compressed by 128". - compressed_page_size: page_size // 4 + compressed_page_size: physical indexer pool page size compressed_seq_lens: seq_lens // 4, but bounded by at least 1, due to flash_mla requirement. c4_sparse: means "compressed by 4" but only attend to top-512 tokens. all related length will be clipped to 512. @@ -114,6 +110,7 @@ class NonPagedIndexerPlan: @dataclass class PagedIndexerMetadata: page_size: int + compressed_page_size: int page_table: torch.Tensor compressed_seq_lens: torch.Tensor use_topk_v2: bool @@ -177,10 +174,6 @@ class PagedIndexerMetadata: assert self.page_size == 256, "the system hardcodes page_size=256" - @property - def compressed_page_size(self) -> int: - return self.page_size // 4 - @property def max_seq_len(self) -> int: return self.page_table.shape[1] * self.page_size @@ -202,6 +195,7 @@ class PagedIndexerMetadata: dst=self, check_eq_fields=[ "page_size", + "compressed_page_size", "force_deep_gemm_metadata", "use_prefill_cuda_graph", "use_topk_v2", 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 a6f9b5672..c96085db2 100644 --- a/python/sglang/srt/mem_cache/deepseek_v4_memory_pool.py +++ b/python/sglang/srt/mem_cache/deepseek_v4_memory_pool.py @@ -2,7 +2,7 @@ from __future__ import annotations import logging from contextlib import nullcontext -from typing import List, Literal, NamedTuple, Optional, Sequence, Tuple +from typing import List, NamedTuple, Optional, Sequence, Tuple import torch @@ -498,8 +498,15 @@ class DeepSeekV4IndexerPool(KVCache): ) +class _CompressedPoolConfig(NamedTuple): + kv_size: int + state_size: int + state_dtype: torch.dtype + indexer_size: Optional[int] = None + + class DeepSeekV4LayerItem(NamedTuple): - compress_ratio: Literal[0, 4, 128] + compress_ratio: int compress_layer_id: int compress_kv_pool: Optional[DeepSeekV4SingleKVPool] = None @@ -623,7 +630,6 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool): self.num_req_slots = ( num_req_slots if num_req_slots is not None else max_num_reqs + 1 ) - self.c4_size = c4_size self.c4_logical_size = c4_logical_size self.c128_size = c128_size from sglang.kernels.ops.attention.dsv4.unified_kv_kernels.env_gate import ( @@ -642,7 +648,6 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool): # so the caller-supplied, SWA-scaled size does not apply here. c4_state_pool_size = self.num_req_slots * c4_ring_size # Non-unified (fp8) keeps the caller-supplied, SWA-addressed size. - self.c4_state_pool_size = c4_state_pool_size c128_ring_size = self.get_ring_size(128) if ONLINE_C128: # Request-scoped C128 state must also cover PD preallocation slots. @@ -652,9 +657,19 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool): c128_state_pool_size = max( c128_state_pool_size, self.num_req_slots * c128_ring_size ) - self.c128_state_pool_size = c128_state_pool_size - self.c4_state_dtype = c4_state_dtype - self.c128_state_dtype = c128_state_dtype + self.compressed_pool_configs = { + 4: _CompressedPoolConfig( + kv_size=c4_size, + state_size=c4_state_pool_size, + state_dtype=c4_state_dtype, + indexer_size=c4_logical_size, + ), + 128: _CompressedPoolConfig( + kv_size=c128_size, + state_size=c128_state_pool_size, + state_dtype=c128_state_dtype, + ), + } self.compression_ratios = compression_ratios self.online_mtp_max_draft_tokens = online_mtp_max_draft_tokens self.online_c128_state_num_req_slots = c128_state_pool_size @@ -681,24 +696,17 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool): self.sliding_window = sliding_window self.swa_size = swa_size - self.swa_window_size = swa_page_size self.swa_page_size = swa_page_size - self.scale_pad = 1 self.qk_nope_head_dim = qk_nope_head_dim self.qk_rope_head_dim = qk_rope_head_dim self.indexer_head_dim = indexer_head_dim stage_layer_num = len(stage_ratios) - c4_layer_num = sum(1 for r in stage_ratios if r == 4) - c128_layer_num = sum(1 for r in stage_ratios if r == 128) - c4_page_size = page_size // 4 - c128_page_size = page_size // 128 + kv_pool_cls: type = DeepSeekV4SingleKVPool if self._unified_kv: self.swa_kv_pool = None - self.c4_kv_pool = None - self.c128_kv_pool = None swa_ring_size = get_swa_ring_size( self.sliding_window, get_spec().speculative_algorithm is not None ) @@ -721,7 +729,6 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool): self.swa_req_ring_size = self.unified_swa_ring_size else: self.unified_kv_pool = None - kv_pool_cls: type = DeepSeekV4SingleKVPool if self.uniform_fp8: assert dtype == torch.float8_e4m3fn, ( "--dsv4-attn-backend trtllm requires " @@ -739,43 +746,14 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool): cls=kv_pool_cls, ) - c4_kv_pool_type = kv_pool_cls - if enable_hisparse: - assert not self.uniform_fp8, ( - "enable_hisparse is not supported with --dsv4-attn-backend trtllm." - ) - c4_kv_pool_type = HiSparseC4DevicePool - self.c4_kv_pool = self._make_kv_pool( - size=c4_size, - page_size=c4_page_size, - dtype=dtype, - layer_num=c4_layer_num, - device=device, - enable_memory_saver=enable_memory_saver, - global_page_size=page_size, - cls=c4_kv_pool_type, - ) - - self.c128_kv_pool = self._make_kv_pool( - size=c128_size, - page_size=c128_page_size, - dtype=dtype, - layer_num=c128_layer_num, - device=device, - enable_memory_saver=enable_memory_saver, - global_page_size=page_size, - cls=kv_pool_cls, - ) - - indexer_size = self.c4_logical_size - self.c4_indexer_kv_pool = self._make_indexer_pool( - indexer_size, - c4_page_size, - dtype, - indexer_head_dim, - c4_layer_num, - device, - enable_memory_saver, + self._init_compressed_pools( + stage_ratios=stage_ratios, + page_size=page_size, + dtype=dtype, + device=device, + enable_memory_saver=enable_memory_saver, + enable_hisparse=enable_hisparse, + kv_pool_cls=kv_pool_cls, ) self._init_compressed_layer_mapping() @@ -804,53 +782,39 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool): data_lens: List[int] = [] item_lens: List[int] = [] - if self._unified_kv: - # Unified buffer per layer: [swa_pages + padded_compress_rows, head_dim]. - # Compressed region [swa_pages:] is page-contiguous (row swa_pages + - # loc//ratio), so reuse the page-block PD transfer by offsetting the ptr - # past the SWA ring and setting item_len = one page of rows. The SWA ring - # ships separately as StateType.SWA_RING. Order [c4, c4_indexer, c128] - # mirrors the non-unified kv_data layout (keeps PP ptr-slicing valid). - stage_ratios = self.compression_ratios[self._stage_start : self._stage_end] - swa_pages = self.unified_kv_pool.swa_pages + def append_page_buffer(buf: torch.Tensor) -> None: + assert buf.ndim == 2, f"expected 2D buffer, got {buf.ndim}D" + data_ptrs.append(buf.data_ptr()) + data_lens.append(buf.nbytes) + item_lens.append(buf[0].nbytes) - def _append_compressed_entry(local_layer_id: int, ratio: int) -> None: - buf = self.unified_kv_pool.kv_buffer[local_layer_id] - assert buf.ndim == 2, f"expected 2D buffer, got {buf.ndim}D" - row_bytes = buf[0].nbytes - rows_per_page = self.page_size // ratio - compress_rows = buf.shape[0] - swa_pages - data_ptrs.append(buf.data_ptr() + swa_pages * row_bytes) - data_lens.append(compress_rows * row_bytes) - item_lens.append(rows_per_page * row_bytes) + stage_ratios = self.compression_ratios[self._stage_start : self._stage_end] + # Registration order defines the PD wire layout: C4 KV, C4 indexer, C128 KV. + # Keep each indexer immediately after the KV buffers of the same ratio. + for ratio, kv_pool in self.kv_pools.items(): + if self._unified_kv: + # Unified buffers store token rows after the SWA ring. Transfer + # compressed pages from the offset; SWA ships as StateType.SWA_RING. + swa_pages = self.unified_kv_pool.swa_pages + for local_layer_id, layer_ratio in enumerate(stage_ratios): + if layer_ratio != ratio: + continue + buf = self.unified_kv_pool.kv_buffer[local_layer_id] + assert buf.ndim == 2, f"expected 2D buffer, got {buf.ndim}D" + row_bytes = buf[0].nbytes + rows_per_page = self.page_size // ratio + compress_rows = buf.shape[0] - swa_pages + data_ptrs.append(buf.data_ptr() + swa_pages * row_bytes) + data_lens.append(compress_rows * row_bytes) + item_lens.append(rows_per_page * row_bytes) + else: + for buf in kv_pool.kv_buffer: + append_page_buffer(buf) - c4_locals = [i for i, r in enumerate(stage_ratios) if r == 4] - c128_locals = [i for i, r in enumerate(stage_ratios) if r == 128] - - for i in c4_locals: - _append_compressed_entry(i, 4) - for buf in self.c4_indexer_kv_pool.contiguous_page_row_buffers(): - assert buf.ndim == 2, f"expected 2D buffer, got {buf.ndim}D" - data_ptrs.append(buf.data_ptr()) - data_lens.append(buf.nbytes) - item_lens.append(buf[0].nbytes) - for i in c128_locals: - _append_compressed_entry(i, 128) - - return data_ptrs, data_lens, item_lens - - buf_groups = [ - self.c4_kv_pool.kv_buffer, - self.c4_indexer_kv_pool.contiguous_page_row_buffers(), - self.c128_kv_pool.kv_buffer, - ] - - for bufs in buf_groups: - for buf in bufs: - assert buf.ndim == 2, f"expected 2D buffer, got {buf.ndim}D" - data_ptrs.append(buf.data_ptr()) - data_lens.append(buf.nbytes) - item_lens.append(buf[0].nbytes) + indexer_pool = self.index_pools.get(ratio) + if indexer_pool is not None: + for buf in indexer_pool.contiguous_page_row_buffers(): + append_page_buffer(buf) return data_ptrs, data_lens, item_lens @@ -950,6 +914,62 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool): item_lens.append(t[0].nbytes if ONLINE_C128 else t[0].nbytes * 128) return data_ptrs, data_lens, item_lens + def _init_compressed_pools( + self, + *, + stage_ratios: Sequence[int], + page_size: int, + dtype: torch.dtype, + device: str, + enable_memory_saver: bool, + enable_hisparse: bool, + kv_pool_cls: type, + ) -> None: + configs = self.compressed_pool_configs + layer_counts = {ratio: stage_ratios.count(ratio) for ratio in configs} + # Keep empty pools and allocation order for PP stages without a given ratio. + self.kv_pools: dict[int, Optional[DeepSeekV4SingleKVPool]] = { + ratio: None for ratio in configs + } + + if not self._unified_kv: + for ratio, config in configs.items(): + pool_cls = kv_pool_cls + if ratio == 4 and enable_hisparse: + assert not self.uniform_fp8, ( + "enable_hisparse is not supported with --dsv4-attn-backend trtllm." + ) + pool_cls = HiSparseC4DevicePool + self.kv_pools[ratio] = self._make_kv_pool( + size=config.kv_size, + page_size=page_size // ratio, + dtype=dtype, + layer_num=layer_counts[ratio], + device=device, + enable_memory_saver=enable_memory_saver, + global_page_size=page_size, + cls=pool_cls, + ) + + self.index_pools: dict[int, DeepSeekV4IndexerPool] = { + ratio: self._make_indexer_pool( + config.indexer_size, + page_size // ratio, + dtype, + self.indexer_head_dim, + layer_counts[ratio], + device, + enable_memory_saver, + ) + for ratio, config in configs.items() + if config.indexer_size is not None + } + + # HiCache and hardware backends still access the per-ratio attributes. + self.c4_kv_pool = self.kv_pools[4] + self.c128_kv_pool = self.kv_pools[128] + self.c4_indexer_kv_pool = self.index_pools[4] + def _make_kv_pool( self, *, @@ -1002,21 +1022,17 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool): enable_memory_saver, ) - def _state_pool_size(self, ratio: int) -> int: - return self.c4_state_pool_size if ratio == 4 else self.c128_state_pool_size - - def _make_attn_state_pool( - self, ratio: int, enable_memory_saver: bool + def _make_compress_state_pool( + self, ratio: int, *, head_dim: int, enable_memory_saver: bool ) -> CompressStatePool: - """Build the per-layer attention compress-state pool for ``ratio`` - (4 or 128). Overridden by :class:`DSV4NPUTokenToKVPool` to swap the - ring-buffered pool for the NPU paged one.""" + """Build attention or indexer state; hardware backends override this factory.""" + config = self.compressed_pool_configs[ratio] return CompressStatePool( - size=self._state_pool_size(ratio), + size=config.state_size, ring_size=self.get_ring_size(ratio), overlap=ratio == 4, - head_dim=self.qk_nope_head_dim + self.qk_rope_head_dim, - dtype=self.c4_state_dtype if ratio == 4 else self.c128_state_dtype, + head_dim=head_dim, + dtype=config.state_dtype, device=self.device, enable_memory_saver=enable_memory_saver, ratio=ratio, @@ -1027,25 +1043,7 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool): ), ) - def _make_indexer_state_pool( - self, ratio: int, enable_memory_saver: bool - ) -> CompressStatePool: - """Build the per-layer indexer compress-state pool (c4 only).""" - return CompressStatePool( - size=self._state_pool_size(ratio), - ring_size=self.get_ring_size(ratio), - overlap=ratio == 4, - head_dim=self.indexer_head_dim, - device=self.device, - dtype=self.c4_state_dtype, - enable_memory_saver=enable_memory_saver, - ratio=ratio, - swa_page_size=self.swa_page_size, - ) - def _init_paged_compress_states(self, enable_memory_saver: bool): - c4_state_pool_size = self.c4_state_pool_size - c128_state_pool_size = self.c128_state_pool_size total_L = len(self.compression_ratios) self.compress_state_pools: List[Optional[CompressStatePool]] = [None] * total_L self.indexer_compress_state_pools: List[Optional[CompressStatePool]] = [ @@ -1057,44 +1055,34 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool): if ratio == 0: continue - self.compress_state_pools[idx] = self._make_attn_state_pool( - ratio, enable_memory_saver + self.compress_state_pools[idx] = self._make_compress_state_pool( + ratio, + head_dim=self.qk_nope_head_dim + self.qk_rope_head_dim, + enable_memory_saver=enable_memory_saver, ) - if ratio == 4: - self.indexer_compress_state_pools[idx] = self._make_indexer_state_pool( - ratio, enable_memory_saver + if ratio in self.index_pools: + self.indexer_compress_state_pools[idx] = self._make_compress_state_pool( + ratio, + head_dim=self.indexer_head_dim, + enable_memory_saver=enable_memory_saver, ) def _init_compressed_layer_mapping(self): - c0_cnt = c4_cnt = c128_cnt = 0 + layer_counts = {0: 0, **{ratio: 0 for ratio in self.kv_pools}} total_L = len(self.compression_ratios) self.layer_mapping: List[Optional[DeepSeekV4LayerItem]] = [None] * total_L for idx in range(self._stage_start, self._stage_end): ratio = self.compression_ratios[idx] - if ratio == 0: - self.layer_mapping[idx] = DeepSeekV4LayerItem( - compress_ratio=0, - compress_layer_id=c0_cnt, - ) - c0_cnt += 1 - elif ratio == 4: - self.layer_mapping[idx] = DeepSeekV4LayerItem( - compress_ratio=4, - compress_layer_id=c4_cnt, - compress_kv_pool=self.c4_kv_pool, - ) - c4_cnt += 1 - elif ratio == 128: - self.layer_mapping[idx] = DeepSeekV4LayerItem( - compress_ratio=128, - compress_layer_id=c128_cnt, - compress_kv_pool=self.c128_kv_pool, - ) - c128_cnt += 1 - else: + if ratio not in layer_counts: raise ValueError(f"Unsupported compression ratio: {ratio}") + self.layer_mapping[idx] = DeepSeekV4LayerItem( + compress_ratio=ratio, + compress_layer_id=layer_counts[ratio], + compress_kv_pool=self.kv_pools.get(ratio), + ) + layer_counts[ratio] += 1 def wait_layer_transfer(self, layer_id: int) -> None: if self.layer_transfer_counter is not None: @@ -1213,20 +1201,6 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool): def get_swa_raw_buffer(self, layer_id: int) -> torch.Tensor: return self.swa_kv_pool.kv_buffer[self._swa_local_layer_id(layer_id)] - def get_swa_key_buffer(self, layer_id: int) -> torch.Tensor: - self.wait_layer_transfer(layer_id) - return self.swa_kv_pool.get_key_buffer(self._swa_local_layer_id(layer_id)) - - def set_swa_key_buffer( - self, - layer_id: int, - loc: torch.Tensor, - cache_nope_fp8_rope_bf16_pack: NopeFp8RopeBf16Pack, - ) -> None: - self.swa_kv_pool.set_key_buffer( - self._swa_local_layer_id(layer_id), loc, cache_nope_fp8_rope_bf16_pack - ) - def get_extra_key_page_size(self, layer_id: int) -> int: _, _, compress_kv_pool = self.layer_mapping[layer_id] assert compress_kv_pool is not None @@ -1250,26 +1224,36 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool): compress_layer_id, loc, cache_nope_fp8_rope_bf16_pack ) - def get_index_k_page_size(self) -> int: - return self.c4_indexer_kv_pool.page_size + def _indexer_pool(self, compress_ratio: int) -> DeepSeekV4IndexerPool: + pool = self.index_pools.get(compress_ratio) + assert pool is not None, ( + f"No indexer pool for compression ratio {compress_ratio}" + ) + return pool + + def get_index_k_page_size(self, compress_ratio: int = 4) -> int: + return self._indexer_pool(compress_ratio).page_size def get_index_k_with_scale_buffer(self, layer_id: int) -> torch.Tensor: self.wait_layer_transfer(layer_id) compress_ratio, compress_layer_id, _ = self.layer_mapping[layer_id] - assert compress_ratio == 4, f"only c4 has indexer, got {compress_ratio = }" - return self.c4_indexer_kv_pool.get_index_k_with_scale_buffer(compress_layer_id) + return self._indexer_pool(compress_ratio).get_index_k_with_scale_buffer( + compress_layer_id + ) def get_index_k_fp4_payload_buffer(self, layer_id: int) -> torch.Tensor: self.wait_layer_transfer(layer_id) compress_ratio, compress_layer_id, _ = self.layer_mapping[layer_id] - assert compress_ratio == 4, f"only c4 has indexer, got {compress_ratio = }" - return self.c4_indexer_kv_pool.get_index_k_fp4_payload_buffer(compress_layer_id) + return self._indexer_pool(compress_ratio).get_index_k_fp4_payload_buffer( + compress_layer_id + ) def get_index_k_fp4_scale_buffer(self, layer_id: int) -> torch.Tensor: self.wait_layer_transfer(layer_id) compress_ratio, compress_layer_id, _ = self.layer_mapping[layer_id] - assert compress_ratio == 4, f"only c4 has indexer, got {compress_ratio = }" - return self.c4_indexer_kv_pool.get_index_k_fp4_scale_buffer(compress_layer_id) + return self._indexer_pool(compress_ratio).get_index_k_fp4_scale_buffer( + compress_layer_id + ) def get_index_k_scale_buffer( self, @@ -1281,8 +1265,7 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool): ) -> Tuple[torch.Tensor, torch.Tensor]: self.wait_layer_transfer(layer_id) compress_ratio, compress_layer_id, _ = self.layer_mapping[layer_id] - assert compress_ratio == 4, f"only c4 has indexer, got {compress_ratio = }" - return self.c4_indexer_kv_pool.get_index_k_scale_buffer( + return self._indexer_pool(compress_ratio).get_index_k_scale_buffer( compress_layer_id, seq_len_tensor, page_indices, @@ -1298,8 +1281,7 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool): index_k_scale: torch.Tensor, ) -> None: compress_ratio, compress_layer_id, _ = self.layer_mapping[layer_id] - assert compress_ratio == 4, f"only c4 has indexer, got {compress_ratio = }" - self.c4_indexer_kv_pool.set_index_k_scale_buffer( + self._indexer_pool(compress_ratio).set_index_k_scale_buffer( compress_layer_id, loc, index_k, index_k_scale ) @@ -1425,8 +1407,9 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool): cache_k: torch.Tensor, ) -> None: compress_ratio, compress_layer_id, _ = self.layer_mapping[layer_id] - assert compress_ratio == 4, f"only c4 has indexer, got {compress_ratio = }" - return self.c4_indexer_kv_pool.set_index_fused(compress_layer_id, loc, cache_k) + return self._indexer_pool(compress_ratio).set_index_fused( + compress_layer_id, loc, cache_k + ) def set_index_k_fp4( self, @@ -1435,5 +1418,6 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool): cache_k: torch.Tensor, ) -> None: compress_ratio, compress_layer_id, _ = self.layer_mapping[layer_id] - assert compress_ratio == 4, f"only c4 has indexer, got {compress_ratio = }" - return self.c4_indexer_kv_pool.set_index_fp4(compress_layer_id, loc, cache_k) + return self._indexer_pool(compress_ratio).set_index_fp4( + compress_layer_id, loc, cache_k + ) diff --git a/test/registered/unit/layers/test_dsv4_nonpaged_indexer.py b/test/registered/unit/layers/test_dsv4_nonpaged_indexer.py index 06fe11e66..baa490a6a 100644 --- a/test/registered/unit/layers/test_dsv4_nonpaged_indexer.py +++ b/test/registered/unit/layers/test_dsv4_nonpaged_indexer.py @@ -45,6 +45,7 @@ class TestDSV4PagedIndexerMetadata(CustomTestCase): ): metadata = PagedIndexerMetadata( page_size=256, + compressed_page_size=64, page_table=torch.zeros((1, 1), dtype=torch.int32), compressed_seq_lens=torch.tensor([65], dtype=torch.int32), use_topk_v2=False, @@ -59,22 +60,7 @@ class TestDSV4PagedIndexerMetadata(CustomTestCase): self.assertEqual(args[1:], (64, 1)) jit_metadata.assert_not_called() - def test_sm120_fp8_torch_fallback_keeps_metadata_none(self): - with ( - envs.SGLANG_FP8_PAGED_MQA_LOGITS_TORCH.override(True), - envs.SGLANG_OPT_USE_AITER_INDEXER.override(False), - envs.SGLANG_OPT_USE_TOPK_V2.override(False), - ): - metadata = PagedIndexerMetadata( - page_size=256, - page_table=torch.zeros((1, 1), dtype=torch.int32), - compressed_seq_lens=torch.tensor([65], dtype=torch.int32), - use_topk_v2=False, - ) - - self.assertIsNone(metadata.deep_gemm_metadata) - - def test_topk_v2_ineligible_backend_skips_plan(self): + def test_torch_fallback_skips_deep_gemm_and_ineligible_topk_plan(self): with ( envs.SGLANG_FP8_PAGED_MQA_LOGITS_TORCH.override(True), envs.SGLANG_OPT_USE_AITER_INDEXER.override(False), @@ -83,14 +69,54 @@ class TestDSV4PagedIndexerMetadata(CustomTestCase): ): metadata = PagedIndexerMetadata( page_size=256, + compressed_page_size=64, page_table=torch.zeros((1, 1), dtype=torch.int32), compressed_seq_lens=torch.tensor([65], dtype=torch.int32), use_topk_v2=False, ) + self.assertIsNone(metadata.deep_gemm_metadata) plan_topk_v2.assert_not_called() self.assertEqual(metadata.topk_metadata.numel(), 0) + def test_physical_page_size_controls_metadata_and_replay(self): + planner = MagicMock(return_value=torch.zeros((1, 2), dtype=torch.int32)) + deep_gemm = SimpleNamespace( + get_num_sms=MagicMock(return_value=1), + get_paged_mqa_logits_metadata=planner, + ) + with patch.dict(sys.modules, {"deep_gemm": deep_gemm}): + metadata = [ + PagedIndexerMetadata( + page_size=256, + compressed_page_size=page_size, + page_table=torch.zeros((1, 3), dtype=torch.int32), + compressed_seq_lens=torch.tensor([65], dtype=torch.int32), + use_topk_v2=False, + force_deep_gemm_metadata=True, + ) + for page_size in (64, 32, 32) + ] + + self.assertEqual( + [call.args[1] for call in planner.call_args_list], [64, 32, 32] + ) + self.assertEqual([m.max_compressed_seq_len for m in metadata], [192, 96, 96]) + self.assertEqual([m.max_seq_len for m in metadata], [768, 768, 768]) + with self.assertRaisesRegex(AssertionError, "compressed_page_size"): + metadata[0].copy_(metadata[1]) + + destination, source = metadata[1:] + source.page_table.fill_(7) + source.compressed_seq_lens.fill_(17) + page_table_ptr = destination.page_table.data_ptr() + destination.copy_(source) + self.assertEqual(destination.page_table.data_ptr(), page_table_ptr) + torch.testing.assert_close(destination.page_table, source.page_table) + torch.testing.assert_close( + destination.compressed_seq_lens, source.compressed_seq_lens + ) + class TestDSV4FlashInferTopK(CustomTestCase): def test_compact_page_transform_respects_fuse_topk(self): diff --git a/test/registered/unit/mem_cache/test_dsv4_compressed_pools.py b/test/registered/unit/mem_cache/test_dsv4_compressed_pools.py new file mode 100644 index 000000000..65eaf1c79 --- /dev/null +++ b/test/registered/unit/mem_cache/test_dsv4_compressed_pools.py @@ -0,0 +1,180 @@ +import unittest +from itertools import product +from types import SimpleNamespace +from unittest.mock import MagicMock, patch + +import torch + +from sglang.srt.mem_cache.deepseek_v4_memory_pool import ( + DeepSeekV4SingleKVPool, + DeepSeekV4TokenToKVPool, + _CompressedPoolConfig, +) +from sglang.test.ci.ci_register import register_cpu_ci +from sglang.test.test_utils import CustomTestCase + +register_cpu_ci(est_time=5, suite="base-a-test-cpu") + + +class TestDSV4CompressedPools(CustomTestCase): + def test_pp_mapping_and_pd_buffer_order(self): + for unified, stage_ratios in product( + (False, True), ([4, 0, 128, 4], [128], [0]) + ): + with self.subTest(unified=unified, stage_ratios=stage_ratios): + pool = DeepSeekV4TokenToKVPool.__new__(DeepSeekV4TokenToKVPool) + pool._unified_kv = unified + pool.uniform_fp8 = False + pool.compressed_pool_configs = { + 4: _CompressedPoolConfig( + 256, 64, torch.bfloat16, indexer_size=1024 + ), + 128: _CompressedPoolConfig(512, 8, torch.float32), + } + pool.indexer_head_dim = 128 + pool.page_size = 256 + pool.compression_ratios = [128] + stage_ratios + [4] + pool._stage_start = 1 + pool._stage_end = 1 + len(stage_ratios) + + def kv_factory(**kwargs): + return SimpleNamespace( + kv_buffer=[ + torch.empty((3, kwargs["page_size"]), dtype=torch.uint8) + for _ in range(kwargs["layer_num"]) + ] + ) + + # Separate payload/scale buffers exercise the FP4 transfer contract. + indexer_buffers = [ + torch.empty((3, width), dtype=torch.uint8) + for _ in range(stage_ratios.count(4)) + for width in (32, 2) + ] + indexer = SimpleNamespace( + contiguous_page_row_buffers=lambda: indexer_buffers + ) + with ( + patch.object(pool, "_make_kv_pool", side_effect=kv_factory), + patch.object(pool, "_make_indexer_pool", return_value=indexer), + ): + pool._init_compressed_pools( + stage_ratios=stage_ratios, + page_size=256, + dtype=torch.float8_e4m3fn, + device="cpu", + enable_memory_saver=False, + enable_hisparse=False, + kv_pool_cls=DeepSeekV4SingleKVPool, + ) + pool._init_compressed_layer_mapping() + self.assertIsNone(pool.layer_mapping[0]) + self.assertIsNone(pool.layer_mapping[-1]) + for local_id, ratio in enumerate(stage_ratios): + item = pool.layer_mapping[local_id + 1] + self.assertEqual( + item.compress_layer_id, stage_ratios[:local_id].count(ratio) + ) + self.assertIs(item.compress_kv_pool, pool.kv_pools.get(ratio)) + self.assertIs(pool.c4_kv_pool, pool.kv_pools[4]) + self.assertIs(pool.c128_kv_pool, pool.kv_pools[128]) + self.assertIs(pool.c4_indexer_kv_pool, pool.index_pools[4]) + + if unified: + buffers = [ + torch.empty((9, 8), dtype=torch.uint8) for _ in stage_ratios + ] + pool.unified_kv_pool = SimpleNamespace( + swa_pages=2, kv_buffer=buffers + ) + + def kv_entries(ratio): + return [ + (buf.data_ptr() + 16, 56, 256 // ratio * 8) + for buf, r in zip(buffers, stage_ratios) + if r == ratio + ] + else: + + def kv_entries(ratio): + return [ + (b.data_ptr(), b.nbytes, b[0].nbytes) + for b in pool.kv_pools[ratio].kv_buffer + ] + + indexer_entries = [ + (b.data_ptr(), b.nbytes, b[0].nbytes) for b in indexer_buffers + ] + expected = kv_entries(4) + indexer_entries + kv_entries(128) + actual = list(zip(*pool.get_contiguous_buf_infos())) + self.assertEqual(actual, expected) + + def test_shared_state_factory_preserves_layouts(self): + pool = DeepSeekV4TokenToKVPool.__new__(DeepSeekV4TokenToKVPool) + pool.compressed_pool_configs = { + 4: _CompressedPoolConfig(256, 64, torch.bfloat16, indexer_size=1024), + 128: _CompressedPoolConfig(512, 8, torch.float32), + } + pool.compression_ratios = [0, 4, 128] + pool._stage_start, pool._stage_end = 0, 3 + pool.index_pools = {4: object()} + pool.qk_nope_head_dim, pool.qk_rope_head_dim = 448, 64 + pool.indexer_head_dim = 128 + pool.device = "cpu" + pool.swa_page_size = 128 + pool.online_mtp_max_draft_tokens = 3 + for online in (False, True): + with ( + self.subTest(online=online), + patch( + "sglang.srt.mem_cache.deepseek_v4_memory_pool.ONLINE_C128", online + ), + patch.object( + pool, + "get_ring_size", + side_effect=lambda r: 8 if r == 4 else (1 if online else 128), + ), + ): + pool._init_paged_compress_states(False) + c4 = pool.compress_state_pools[1].kv_score_buffer.kv_score + indexer = pool.indexer_compress_state_pools[1].kv_score_buffer.kv_score + c128 = pool.compress_state_pools[2].kv_score_buffer.kv_score + self.assertEqual(c4.shape, (76, 2048)) + self.assertEqual(indexer.shape, (76, 512)) + self.assertEqual(c128.shape, (40, 1536) if online else (256, 1024)) + self.assertEqual(c4.dtype, torch.bfloat16) + self.assertEqual(indexer.dtype, torch.bfloat16) + self.assertEqual(c128.dtype, torch.float32) + self.assertNotEqual(c4.data_ptr(), indexer.data_ptr()) + self.assertIsNone(pool.compress_state_pools[0]) + self.assertIsNone(pool.indexer_compress_state_pools[2]) + + def test_indexer_access_uses_layer_ratio_and_waits_only_for_reads(self): + pool = DeepSeekV4TokenToKVPool.__new__(DeepSeekV4TokenToKVPool) + pool.kv_pools = {4: None, 128: None} + pool.compression_ratios = [4, 128, 4] + pool._stage_start, pool._stage_end = 0, 3 + pool._init_compressed_layer_mapping() + indexer = MagicMock(page_size=32) + pool.index_pools = {4: indexer} + trace = MagicMock() + trace.attach_mock(indexer, "indexer") + with patch.object(pool, "wait_layer_transfer") as wait: + trace.attach_mock(wait, "wait") + pool.get_index_k_fp4_payload_buffer(2) + pool.set_index_k_fp4(2, "loc", "cache") + self.assertEqual( + trace.mock_calls, + [ + unittest.mock.call.wait(2), + unittest.mock.call.indexer.get_index_k_fp4_payload_buffer(1), + unittest.mock.call.indexer.set_index_fp4(1, "loc", "cache"), + ], + ) + self.assertEqual(pool.get_index_k_page_size(4), 32) + with self.assertRaisesRegex(AssertionError, "No indexer pool"): + pool.get_index_k_page_size(128) + + +if __name__ == "__main__": + unittest.main()