[Refactor] Generalize DeepSeek V4 compressed pool management (#38954)

This commit is contained in:
Liangsheng Yin
2026-09-10 17:29:03 -07:00
committed by GitHub
parent d006f40e24
commit 41da06adca
8 changed files with 435 additions and 280 deletions
@@ -321,8 +321,8 @@ class DSV4NPUTokenToKVPool(DeepSeekV4TokenToKVPool):
return self.get_indexer_compress_states(layer_id) return self.get_indexer_compress_states(layer_id)
return self.get_attention_compress_states(layer_id) return self.get_attention_compress_states(layer_id)
def _make_attn_state_pool( def _make_compress_state_pool(
self, ratio: int, enable_memory_saver: bool self, ratio: int, *, head_dim: int, enable_memory_saver: bool
) -> NPUCompressStatePool: ) -> NPUCompressStatePool:
# ONLINE_C128 (CUDA-only) collapses the c128 ring to size 1; the NPU fused # 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. # 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 " "SGLANG_OPT_USE_ONLINE_COMPRESS is incompatible with the "
"NPU fused compressor (no online mode in the kernel)." "NPU fused compressor (no online mode in the kernel)."
) )
config = self.compressed_pool_configs[ratio]
ring_size = self.get_ring_size(ratio) ring_size = self.get_ring_size(ratio)
# A5 cache_mode=2 addresses one ring bank per request. The A3 # A5 cache_mode=2 addresses one ring bank per request. The A3
# explicit-location path can share the smaller flat pool, but the A5 # explicit-location path can share the smaller flat pool, but the A5
# cycle ABI needs enough physical banks for every req_pool_idx. # 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(): if is_npu_arch35():
size = max(size, self.num_req_slots * ring_size) size = max(size, self.num_req_slots * ring_size)
return NPUCompressStatePool( return NPUCompressStatePool(
size=size, size=size,
ring_size=ring_size, ring_size=ring_size,
overlap=ratio == 4, overlap=ratio == 4,
head_dim=self.qk_nope_head_dim + self.qk_rope_head_dim, head_dim=head_dim,
dtype=self.c4_state_dtype if ratio == 4 else self.c128_state_dtype, dtype=config.state_dtype,
device=self.device, device=self.device,
enable_memory_saver=enable_memory_saver, enable_memory_saver=enable_memory_saver,
ratio=ratio, ratio=ratio,
swa_page_size=self.swa_page_size, 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( def _make_indexer_pool(
self, self,
size: int, size: int,
@@ -394,10 +374,11 @@ class DSV4NPUTokenToKVPool(DeepSeekV4TokenToKVPool):
def get_contiguous_buf_infos(self) -> Tuple[List[int], List[int], List[int]]: def get_contiguous_buf_infos(self) -> Tuple[List[int], List[int], List[int]]:
"""Main PD buffers addressed by the full KV page id.""" """Main PD buffers addressed by the full KV page id."""
indexer_pool = self._indexer_pool(4)
buffers = ( buffers = (
self.c4_kv_pool.kv_buffer self.c4_kv_pool.kv_buffer
+ self.c4_indexer_kv_pool.index_k_buffer + indexer_pool.index_k_buffer
+ self.c4_indexer_kv_pool.index_scale_buffer + indexer_pool.index_scale_buffer
) )
return ( return (
[buf.data_ptr() for buf in buffers], [buf.data_ptr() for buf in buffers],
@@ -530,16 +511,15 @@ class DSV4NPUTokenToKVPool(DeepSeekV4TokenToKVPool):
``torch.ops.custom.npu_quant_lightning_indexer`` consumes. ``torch.ops.custom.npu_quant_lightning_indexer`` consumes.
""" """
item = self.layer_mapping[layer_id] item = self.layer_mapping[layer_id]
if item.compress_ratio == 4: if item.compress_ratio == 0:
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:
return None 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: if loc is not None:
kv = kv.flatten(0, 1)[loc] kv = kv.flatten(0, 1)[loc]
return kv return kv
@@ -651,21 +631,17 @@ class DSV4NPUTokenToKVPool(DeepSeekV4TokenToKVPool):
ratio, compress_layer_id, _ = self.layer_mapping[layer_id] ratio, compress_layer_id, _ = self.layer_mapping[layer_id]
device_type = kv.device.type device_type = kv.device.type
if from_indexer: 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": 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?" "NPU index buffers not allocated — pool was init'd on CUDA?"
) )
self.c4_indexer_kv_pool.set_index_k_scale( indexer_pool.set_index_k_scale(compress_layer_id, loc, kv, kv_scale)
compress_layer_id, loc, kv, kv_scale
)
return return
if kv_scale is None: 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 return
self.c4_indexer_kv_pool.set_index_k_scale_buffer( indexer_pool.set_index_k_scale_buffer(compress_layer_id, loc, kv, kv_scale)
compress_layer_id, loc, kv, kv_scale
)
return return
compress_pool = self.c4_kv_pool if ratio == 4 else self.c128_kv_pool compress_pool = self.c4_kv_pool if ratio == 4 else self.c128_kv_pool
if device_type == "npu": if device_type == "npu":
@@ -690,5 +666,6 @@ class DSV4NPUTokenToKVPool(DeepSeekV4TokenToKVPool):
) -> torch.Tensor: ) -> torch.Tensor:
# The indexer scale is fp16 on pre-A5 parts and fp32 on A5. # The indexer scale is fp16 on pre-A5 parts and fp32 on A5.
assert from_indexer, "only indexer compress pool has dequant scale" assert from_indexer, "only indexer compress pool has dequant scale"
compress_layer_id = self.layer_mapping[layer_id].compress_layer_id item = self.layer_mapping[layer_id]
return self.c4_indexer_kv_pool.get_index_scale(compress_layer_id) indexer_pool = self._indexer_pool(item.compress_ratio)
return indexer_pool.get_index_scale(item.compress_layer_id)
@@ -778,6 +778,7 @@ class DeepseekV4AttnBackend(
): ):
return PagedIndexerMetadata( return PagedIndexerMetadata(
page_size=self.page_size, 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, page_table=core_attn_metadata.page_table,
compressed_seq_lens=core_attn_metadata.c4_topk_lengths_raw, 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, 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 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( return SparsePrefillChunkCache.build(
seq_lens=forward_batch.seq_lens.to(torch.int32), seq_lens=forward_batch.seq_lens.to(torch.int32),
extend_seq_lens=forward_batch.extend_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, req_to_token=self.req_to_token,
full_to_swa=self.token_to_kv_pool.full_to_swa_index_mapping, full_to_swa=self.token_to_kv_pool.full_to_swa_index_mapping,
swa_window_size=SWA_WINDOW, 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, num_qo_tokens=num_qo_tokens,
max_seq_len=max(seq_lens_cpu_list), max_seq_len=max(seq_lens_cpu_list),
total_swa=total_swa, total_swa=total_swa,
@@ -1772,23 +1771,20 @@ class DeepseekV4AttnBackend(
extra_indices = core_attn_metadata.c128_page_indices extra_indices = core_attn_metadata.c128_page_indices
extra_topk_lengths = core_attn_metadata.c128_topk_lengths_clamp1 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 assert swa_k_cache.ndim == 2
k_cache_total_dim = token_to_kv_pool.swa_kv_pool.kv_cache_total_dim 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 = swa_k_cache[:, : swa_page_size * k_cache_total_dim].view(
swa_k_cache.shape[0], swa_window_size, 1, k_cache_total_dim swa_k_cache.shape[0], swa_page_size, 1, k_cache_total_dim
) )
if extra_k_cache is not None: if extra_k_cache is not None:
page_sizes = { extra_page_size = token_to_kv_pool.get_extra_key_page_size(layer_id)
4: token_to_kv_pool.page_size // 4,
128: token_to_kv_pool.page_size // 128,
}
extra_k_cache = extra_k_cache[ extra_k_cache = extra_k_cache[
:, : page_sizes[compress_ratio] * k_cache_total_dim :, : extra_page_size * k_cache_total_dim
].view( ].view(
extra_k_cache.shape[0], extra_k_cache.shape[0],
page_sizes[compress_ratio], extra_page_size,
1, 1,
k_cache_total_dim, k_cache_total_dim,
) )
@@ -507,6 +507,7 @@ class DeepseekV4HipRadixBackend(
def init_forward_metadata_indexer(self, core_attn_metadata: DSV4AttnMetadata): def init_forward_metadata_indexer(self, core_attn_metadata: DSV4AttnMetadata):
return PagedIndexerMetadata( return PagedIndexerMetadata(
page_size=self.page_size, 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, page_table=core_attn_metadata.page_table,
compressed_seq_lens=core_attn_metadata.c4_topk_lengths_raw, compressed_seq_lens=core_attn_metadata.c4_topk_lengths_raw,
use_topk_v2=self.dsa_topk_backend.should_use_topk_v2(), 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_indices = core_attn_metadata.c128_page_indices
extra_topk_lengths = core_attn_metadata.c128_topk_lengths_clamp1 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 assert swa_k_cache.ndim == 2
k_cache_total_dim = token_to_kv_pool.swa_kv_pool.kv_cache_total_dim 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 = swa_k_cache[:, : swa_page_size * k_cache_total_dim].view(
swa_k_cache.shape[0], swa_window_size, 1, k_cache_total_dim swa_k_cache.shape[0], swa_page_size, 1, k_cache_total_dim
) )
if extra_k_cache is not None: if extra_k_cache is not None:
page_sizes = { extra_page_size = token_to_kv_pool.get_extra_key_page_size(layer_id)
4: token_to_kv_pool.page_size // 4,
128: token_to_kv_pool.page_size // 128,
}
extra_k_cache = extra_k_cache[ extra_k_cache = extra_k_cache[
:, : page_sizes[compress_ratio] * k_cache_total_dim :, : extra_page_size * k_cache_total_dim
].view( ].view(
extra_k_cache.shape[0], extra_k_cache.shape[0],
page_sizes[compress_ratio], extra_page_size,
1, 1,
k_cache_total_dim, k_cache_total_dim,
) )
@@ -265,7 +265,7 @@ class CompressorBackendMixin:
bf16_store = False bf16_store = False
kv_scale_cache = None kv_scale_cache = None
if compressor.is_in_indexer: 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: if use_hip_fp4:
kv_cache = token_to_kv_pool.get_index_k_fp4_payload_buffer(layer_id) 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) kv_scale_cache = token_to_kv_pool.get_index_k_fp4_scale_buffer(layer_id)
@@ -2,7 +2,7 @@ from __future__ import annotations
import warnings import warnings
from dataclasses import dataclass, field, fields from dataclasses import dataclass, field, fields
from typing import TYPE_CHECKING, Any, List, Optional from typing import Any, List, Optional
import torch import torch
@@ -11,10 +11,6 @@ from sglang.srt.utils import is_hip, is_sm120_supported, is_xpu
_IS_SM120 = is_sm120_supported() _IS_SM120 = is_sm120_supported()
if TYPE_CHECKING:
pass
""" """
Some comments on the common terms used in DeepSeekV4Backend: Some comments on the common terms used in DeepSeekV4Backend:
@@ -45,7 +41,7 @@ positions:
Some other notes: Some other notes:
c4_ / c128_: means "compressed by 4" / "compressed by 128". 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. 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. c4_sparse: means "compressed by 4" but only attend to top-512 tokens.
all related length will be clipped to 512. all related length will be clipped to 512.
@@ -114,6 +110,7 @@ class NonPagedIndexerPlan:
@dataclass @dataclass
class PagedIndexerMetadata: class PagedIndexerMetadata:
page_size: int page_size: int
compressed_page_size: int
page_table: torch.Tensor page_table: torch.Tensor
compressed_seq_lens: torch.Tensor compressed_seq_lens: torch.Tensor
use_topk_v2: bool use_topk_v2: bool
@@ -177,10 +174,6 @@ class PagedIndexerMetadata:
assert self.page_size == 256, "the system hardcodes page_size=256" assert self.page_size == 256, "the system hardcodes page_size=256"
@property
def compressed_page_size(self) -> int:
return self.page_size // 4
@property @property
def max_seq_len(self) -> int: def max_seq_len(self) -> int:
return self.page_table.shape[1] * self.page_size return self.page_table.shape[1] * self.page_size
@@ -202,6 +195,7 @@ class PagedIndexerMetadata:
dst=self, dst=self,
check_eq_fields=[ check_eq_fields=[
"page_size", "page_size",
"compressed_page_size",
"force_deep_gemm_metadata", "force_deep_gemm_metadata",
"use_prefill_cuda_graph", "use_prefill_cuda_graph",
"use_topk_v2", "use_topk_v2",
@@ -2,7 +2,7 @@ from __future__ import annotations
import logging import logging
from contextlib import nullcontext from contextlib import nullcontext
from typing import List, Literal, NamedTuple, Optional, Sequence, Tuple from typing import List, NamedTuple, Optional, Sequence, Tuple
import torch 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): class DeepSeekV4LayerItem(NamedTuple):
compress_ratio: Literal[0, 4, 128] compress_ratio: int
compress_layer_id: int compress_layer_id: int
compress_kv_pool: Optional[DeepSeekV4SingleKVPool] = None compress_kv_pool: Optional[DeepSeekV4SingleKVPool] = None
@@ -623,7 +630,6 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool):
self.num_req_slots = ( self.num_req_slots = (
num_req_slots if num_req_slots is not None else max_num_reqs + 1 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.c4_logical_size = c4_logical_size
self.c128_size = c128_size self.c128_size = c128_size
from sglang.kernels.ops.attention.dsv4.unified_kv_kernels.env_gate import ( 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. # so the caller-supplied, SWA-scaled size does not apply here.
c4_state_pool_size = self.num_req_slots * c4_ring_size c4_state_pool_size = self.num_req_slots * c4_ring_size
# Non-unified (fp8) keeps the caller-supplied, SWA-addressed 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) c128_ring_size = self.get_ring_size(128)
if ONLINE_C128: if ONLINE_C128:
# Request-scoped C128 state must also cover PD preallocation slots. # 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 = max(
c128_state_pool_size, self.num_req_slots * c128_ring_size c128_state_pool_size, self.num_req_slots * c128_ring_size
) )
self.c128_state_pool_size = c128_state_pool_size self.compressed_pool_configs = {
self.c4_state_dtype = c4_state_dtype 4: _CompressedPoolConfig(
self.c128_state_dtype = c128_state_dtype 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.compression_ratios = compression_ratios
self.online_mtp_max_draft_tokens = online_mtp_max_draft_tokens self.online_mtp_max_draft_tokens = online_mtp_max_draft_tokens
self.online_c128_state_num_req_slots = c128_state_pool_size self.online_c128_state_num_req_slots = c128_state_pool_size
@@ -681,24 +696,17 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool):
self.sliding_window = sliding_window self.sliding_window = sliding_window
self.swa_size = swa_size self.swa_size = swa_size
self.swa_window_size = swa_page_size
self.swa_page_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_nope_head_dim = qk_nope_head_dim
self.qk_rope_head_dim = qk_rope_head_dim self.qk_rope_head_dim = qk_rope_head_dim
self.indexer_head_dim = indexer_head_dim self.indexer_head_dim = indexer_head_dim
stage_layer_num = len(stage_ratios) stage_layer_num = len(stage_ratios)
c4_layer_num = sum(1 for r in stage_ratios if r == 4) kv_pool_cls: type = DeepSeekV4SingleKVPool
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
if self._unified_kv: if self._unified_kv:
self.swa_kv_pool = None self.swa_kv_pool = None
self.c4_kv_pool = None
self.c128_kv_pool = None
swa_ring_size = get_swa_ring_size( swa_ring_size = get_swa_ring_size(
self.sliding_window, get_spec().speculative_algorithm is not None 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 self.swa_req_ring_size = self.unified_swa_ring_size
else: else:
self.unified_kv_pool = None self.unified_kv_pool = None
kv_pool_cls: type = DeepSeekV4SingleKVPool
if self.uniform_fp8: if self.uniform_fp8:
assert dtype == torch.float8_e4m3fn, ( assert dtype == torch.float8_e4m3fn, (
"--dsv4-attn-backend trtllm requires " "--dsv4-attn-backend trtllm requires "
@@ -739,43 +746,14 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool):
cls=kv_pool_cls, cls=kv_pool_cls,
) )
c4_kv_pool_type = kv_pool_cls self._init_compressed_pools(
if enable_hisparse: stage_ratios=stage_ratios,
assert not self.uniform_fp8, ( page_size=page_size,
"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, dtype=dtype,
layer_num=c4_layer_num,
device=device, device=device,
enable_memory_saver=enable_memory_saver, enable_memory_saver=enable_memory_saver,
global_page_size=page_size, enable_hisparse=enable_hisparse,
cls=c4_kv_pool_type, kv_pool_cls=kv_pool_cls,
)
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_layer_mapping() self._init_compressed_layer_mapping()
@@ -804,17 +782,23 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool):
data_lens: List[int] = [] data_lens: List[int] = []
item_lens: List[int] = [] item_lens: List[int] = []
if self._unified_kv: def append_page_buffer(buf: torch.Tensor) -> None:
# Unified buffer per layer: [swa_pages + padded_compress_rows, head_dim]. assert buf.ndim == 2, f"expected 2D buffer, got {buf.ndim}D"
# Compressed region [swa_pages:] is page-contiguous (row swa_pages + data_ptrs.append(buf.data_ptr())
# loc//ratio), so reuse the page-block PD transfer by offsetting the ptr data_lens.append(buf.nbytes)
# past the SWA ring and setting item_len = one page of rows. The SWA ring item_lens.append(buf[0].nbytes)
# 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_compressed_entry(local_layer_id: int, ratio: int) -> None: 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] buf = self.unified_kv_pool.kv_buffer[local_layer_id]
assert buf.ndim == 2, f"expected 2D buffer, got {buf.ndim}D" assert buf.ndim == 2, f"expected 2D buffer, got {buf.ndim}D"
row_bytes = buf[0].nbytes row_bytes = buf[0].nbytes
@@ -823,34 +807,14 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool):
data_ptrs.append(buf.data_ptr() + swa_pages * row_bytes) data_ptrs.append(buf.data_ptr() + swa_pages * row_bytes)
data_lens.append(compress_rows * row_bytes) data_lens.append(compress_rows * row_bytes)
item_lens.append(rows_per_page * 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] indexer_pool = self.index_pools.get(ratio)
c128_locals = [i for i, r in enumerate(stage_ratios) if r == 128] if indexer_pool is not None:
for buf in indexer_pool.contiguous_page_row_buffers():
for i in c4_locals: append_page_buffer(buf)
_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)
return data_ptrs, data_lens, item_lens 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) item_lens.append(t[0].nbytes if ONLINE_C128 else t[0].nbytes * 128)
return data_ptrs, data_lens, item_lens 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( def _make_kv_pool(
self, self,
*, *,
@@ -1002,21 +1022,17 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool):
enable_memory_saver, enable_memory_saver,
) )
def _state_pool_size(self, ratio: int) -> int: def _make_compress_state_pool(
return self.c4_state_pool_size if ratio == 4 else self.c128_state_pool_size self, ratio: int, *, head_dim: int, enable_memory_saver: bool
def _make_attn_state_pool(
self, ratio: int, enable_memory_saver: bool
) -> CompressStatePool: ) -> CompressStatePool:
"""Build the per-layer attention compress-state pool for ``ratio`` """Build attention or indexer state; hardware backends override this factory."""
(4 or 128). Overridden by :class:`DSV4NPUTokenToKVPool` to swap the config = self.compressed_pool_configs[ratio]
ring-buffered pool for the NPU paged one."""
return CompressStatePool( return CompressStatePool(
size=self._state_pool_size(ratio), size=config.state_size,
ring_size=self.get_ring_size(ratio), ring_size=self.get_ring_size(ratio),
overlap=ratio == 4, overlap=ratio == 4,
head_dim=self.qk_nope_head_dim + self.qk_rope_head_dim, head_dim=head_dim,
dtype=self.c4_state_dtype if ratio == 4 else self.c128_state_dtype, dtype=config.state_dtype,
device=self.device, device=self.device,
enable_memory_saver=enable_memory_saver, enable_memory_saver=enable_memory_saver,
ratio=ratio, 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): 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) total_L = len(self.compression_ratios)
self.compress_state_pools: List[Optional[CompressStatePool]] = [None] * total_L self.compress_state_pools: List[Optional[CompressStatePool]] = [None] * total_L
self.indexer_compress_state_pools: List[Optional[CompressStatePool]] = [ self.indexer_compress_state_pools: List[Optional[CompressStatePool]] = [
@@ -1057,44 +1055,34 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool):
if ratio == 0: if ratio == 0:
continue continue
self.compress_state_pools[idx] = self._make_attn_state_pool( self.compress_state_pools[idx] = self._make_compress_state_pool(
ratio, enable_memory_saver ratio,
head_dim=self.qk_nope_head_dim + self.qk_rope_head_dim,
enable_memory_saver=enable_memory_saver,
) )
if ratio == 4: if ratio in self.index_pools:
self.indexer_compress_state_pools[idx] = self._make_indexer_state_pool( self.indexer_compress_state_pools[idx] = self._make_compress_state_pool(
ratio, enable_memory_saver ratio,
head_dim=self.indexer_head_dim,
enable_memory_saver=enable_memory_saver,
) )
def _init_compressed_layer_mapping(self): 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) total_L = len(self.compression_ratios)
self.layer_mapping: List[Optional[DeepSeekV4LayerItem]] = [None] * total_L self.layer_mapping: List[Optional[DeepSeekV4LayerItem]] = [None] * total_L
for idx in range(self._stage_start, self._stage_end): for idx in range(self._stage_start, self._stage_end):
ratio = self.compression_ratios[idx] ratio = self.compression_ratios[idx]
if ratio == 0: if ratio not in layer_counts:
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:
raise ValueError(f"Unsupported compression ratio: {ratio}") 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: def wait_layer_transfer(self, layer_id: int) -> None:
if self.layer_transfer_counter is not 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: 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)] 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: def get_extra_key_page_size(self, layer_id: int) -> int:
_, _, compress_kv_pool = self.layer_mapping[layer_id] _, _, compress_kv_pool = self.layer_mapping[layer_id]
assert compress_kv_pool is not None assert compress_kv_pool is not None
@@ -1250,26 +1224,36 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool):
compress_layer_id, loc, cache_nope_fp8_rope_bf16_pack compress_layer_id, loc, cache_nope_fp8_rope_bf16_pack
) )
def get_index_k_page_size(self) -> int: def _indexer_pool(self, compress_ratio: int) -> DeepSeekV4IndexerPool:
return self.c4_indexer_kv_pool.page_size 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: def get_index_k_with_scale_buffer(self, layer_id: int) -> torch.Tensor:
self.wait_layer_transfer(layer_id) self.wait_layer_transfer(layer_id)
compress_ratio, compress_layer_id, _ = self.layer_mapping[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._indexer_pool(compress_ratio).get_index_k_with_scale_buffer(
return self.c4_indexer_kv_pool.get_index_k_with_scale_buffer(compress_layer_id) compress_layer_id
)
def get_index_k_fp4_payload_buffer(self, layer_id: int) -> torch.Tensor: def get_index_k_fp4_payload_buffer(self, layer_id: int) -> torch.Tensor:
self.wait_layer_transfer(layer_id) self.wait_layer_transfer(layer_id)
compress_ratio, compress_layer_id, _ = self.layer_mapping[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._indexer_pool(compress_ratio).get_index_k_fp4_payload_buffer(
return self.c4_indexer_kv_pool.get_index_k_fp4_payload_buffer(compress_layer_id) compress_layer_id
)
def get_index_k_fp4_scale_buffer(self, layer_id: int) -> torch.Tensor: def get_index_k_fp4_scale_buffer(self, layer_id: int) -> torch.Tensor:
self.wait_layer_transfer(layer_id) self.wait_layer_transfer(layer_id)
compress_ratio, compress_layer_id, _ = self.layer_mapping[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._indexer_pool(compress_ratio).get_index_k_fp4_scale_buffer(
return self.c4_indexer_kv_pool.get_index_k_fp4_scale_buffer(compress_layer_id) compress_layer_id
)
def get_index_k_scale_buffer( def get_index_k_scale_buffer(
self, self,
@@ -1281,8 +1265,7 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool):
) -> Tuple[torch.Tensor, torch.Tensor]: ) -> Tuple[torch.Tensor, torch.Tensor]:
self.wait_layer_transfer(layer_id) self.wait_layer_transfer(layer_id)
compress_ratio, compress_layer_id, _ = self.layer_mapping[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._indexer_pool(compress_ratio).get_index_k_scale_buffer(
return self.c4_indexer_kv_pool.get_index_k_scale_buffer(
compress_layer_id, compress_layer_id,
seq_len_tensor, seq_len_tensor,
page_indices, page_indices,
@@ -1298,8 +1281,7 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool):
index_k_scale: torch.Tensor, index_k_scale: torch.Tensor,
) -> None: ) -> None:
compress_ratio, compress_layer_id, _ = self.layer_mapping[layer_id] compress_ratio, compress_layer_id, _ = self.layer_mapping[layer_id]
assert compress_ratio == 4, f"only c4 has indexer, got {compress_ratio = }" self._indexer_pool(compress_ratio).set_index_k_scale_buffer(
self.c4_indexer_kv_pool.set_index_k_scale_buffer(
compress_layer_id, loc, index_k, index_k_scale compress_layer_id, loc, index_k, index_k_scale
) )
@@ -1425,8 +1407,9 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool):
cache_k: torch.Tensor, cache_k: torch.Tensor,
) -> None: ) -> None:
compress_ratio, compress_layer_id, _ = self.layer_mapping[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._indexer_pool(compress_ratio).set_index_fused(
return self.c4_indexer_kv_pool.set_index_fused(compress_layer_id, loc, cache_k) compress_layer_id, loc, cache_k
)
def set_index_k_fp4( def set_index_k_fp4(
self, self,
@@ -1435,5 +1418,6 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool):
cache_k: torch.Tensor, cache_k: torch.Tensor,
) -> None: ) -> None:
compress_ratio, compress_layer_id, _ = self.layer_mapping[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._indexer_pool(compress_ratio).set_index_fp4(
return self.c4_indexer_kv_pool.set_index_fp4(compress_layer_id, loc, cache_k) compress_layer_id, loc, cache_k
)
@@ -45,6 +45,7 @@ class TestDSV4PagedIndexerMetadata(CustomTestCase):
): ):
metadata = PagedIndexerMetadata( metadata = PagedIndexerMetadata(
page_size=256, page_size=256,
compressed_page_size=64,
page_table=torch.zeros((1, 1), dtype=torch.int32), page_table=torch.zeros((1, 1), dtype=torch.int32),
compressed_seq_lens=torch.tensor([65], dtype=torch.int32), compressed_seq_lens=torch.tensor([65], dtype=torch.int32),
use_topk_v2=False, use_topk_v2=False,
@@ -59,22 +60,7 @@ class TestDSV4PagedIndexerMetadata(CustomTestCase):
self.assertEqual(args[1:], (64, 1)) self.assertEqual(args[1:], (64, 1))
jit_metadata.assert_not_called() jit_metadata.assert_not_called()
def test_sm120_fp8_torch_fallback_keeps_metadata_none(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),
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):
with ( with (
envs.SGLANG_FP8_PAGED_MQA_LOGITS_TORCH.override(True), envs.SGLANG_FP8_PAGED_MQA_LOGITS_TORCH.override(True),
envs.SGLANG_OPT_USE_AITER_INDEXER.override(False), envs.SGLANG_OPT_USE_AITER_INDEXER.override(False),
@@ -83,14 +69,54 @@ class TestDSV4PagedIndexerMetadata(CustomTestCase):
): ):
metadata = PagedIndexerMetadata( metadata = PagedIndexerMetadata(
page_size=256, page_size=256,
compressed_page_size=64,
page_table=torch.zeros((1, 1), dtype=torch.int32), page_table=torch.zeros((1, 1), dtype=torch.int32),
compressed_seq_lens=torch.tensor([65], dtype=torch.int32), compressed_seq_lens=torch.tensor([65], dtype=torch.int32),
use_topk_v2=False, use_topk_v2=False,
) )
self.assertIsNone(metadata.deep_gemm_metadata)
plan_topk_v2.assert_not_called() plan_topk_v2.assert_not_called()
self.assertEqual(metadata.topk_metadata.numel(), 0) 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): class TestDSV4FlashInferTopK(CustomTestCase):
def test_compact_page_transform_respects_fuse_topk(self): def test_compact_page_transform_respects_fuse_topk(self):
@@ -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()