Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
34 changes: 29 additions & 5 deletions python/sglang/srt/model_executor/pool_configurator.py
Original file line number Diff line number Diff line change
Expand Up @@ -1212,6 +1212,25 @@ def _resolve_swa_cap_tokens(self) -> Optional[int]:
headroom = self.swa_prefix_tails * (self.sliding_window_size + self.page_size)
return ceil_align(cap + headroom, self.page_size)

def _get_paged_kv_bytes_per_token(self, compress_ratio: int = 0) -> float:
# Unified rings, the NPU pool and the trtllm uniform-FP8 pool do not go
# through DeepSeekV4SingleKVPool.create_buffer, so they carry no page pad.
if self._unified or _is_npu or get_exec().kernel.dsv4_attn_backend == "trtllm":
return self.kv_bytes
from sglang.srt.mem_cache.deepseek_v4_memory_pool import (
resolve_compressed_kv_layout,
select_dsv4_kv_layout,
)

layout, compressed_option = select_dsv4_kv_layout()
page_size = self.page_size
if compress_ratio:
layout = resolve_compressed_kv_layout(
layout, compress_ratio, compressed_option
)
page_size = self.page_size // compress_ratio
return layout.page_bytes(page_size) / page_size

def _get_bytes_per_swa_token(self) -> float:
"""Bytes one SWA slot costs across the stage. c4_state_pool_size = swa_tokens
// page_size * ring, so c4 compress state is priced per SWA slot too."""
Expand All @@ -1224,9 +1243,11 @@ def _get_bytes_per_swa_token(self) -> float:
c4_indexer_state_bytes = 2 * 2 * self.indexer_head_dim * c4_state_dtype_size

c4_state_ratio = self.ring_sizes.get(4, 0) / self.page_size
return self.kv_bytes * self.num_layers_total + c4_state_ratio * (
c4_state_bytes + c4_indexer_state_bytes
) * self.num_layers(4)
return self._get_paged_kv_bytes_per_token() * self.num_layers_total + (
c4_state_ratio
* (c4_state_bytes + c4_indexer_state_bytes)
* self.num_layers(4)
)

def _compressed_bytes_per_full_token(self, ratio: int) -> float:
"""Compressed KV (+ indexer) bytes one full token adds per layer of `ratio`;
Expand All @@ -1235,9 +1256,12 @@ def _compressed_bytes_per_full_token(self, ratio: int) -> float:
return (self.kv_bytes + self.low_ratio_index_bytes) / ratio
if ratio == 4:
c4_frac = 1 / (4 * self.c4_shrink_factor)
return c4_frac * self.kv_bytes + 1 / 4 * self.indexer_bytes_per_token
return (
c4_frac * self._get_paged_kv_bytes_per_token(4)
+ 1 / 4 * self.indexer_bytes_per_token
)
assert ratio == 128, f"unsupported compression ratio: {ratio}"
return 1 / 128 * self.kv_bytes
return 1 / 128 * self._get_paged_kv_bytes_per_token(128)

def _get_bytes_per_full_token(self) -> float:
# Cap mode and ring mode both move the SWA pool and the c4 state that
Expand Down
24 changes: 24 additions & 0 deletions test/registered/unit/model_executor/test_pool_configurator.py
Original file line number Diff line number Diff line change
Expand Up @@ -1268,6 +1268,30 @@ def test_dsv4_owner_layers_follow_the_pool_rule(self):
owners = collect_sources_by_ratio(ratios, [2, 4, 9], range(8, 10))
self.assertEqual(owners, {4: [8], 2: [9]})

def test_dsv4_paged_kv_budget(self):
from sglang.srt.environ import envs

cfg = self._dsv4_configurator_for_budget()
cfg._unified = False
cfg.encoder_replay = False
cfg.page_size = 256
cfg.kv_bytes = 584
cfg.attn_head_dim = 512
cfg.num_layers_total = 3
cfg.stage_owner_layers = {4: [1], 128: [2]}
cfg.indexer_bytes_per_token = 132
_publish_config(self, dsv4_attn_backend="flashmla")
with (
envs.SGLANG_DSV4_KV_LAYOUT.override("v4"),
envs.SGLANG_DSV4_COMPRESS_STATE_DTYPE.override("float32"),
):
cfg.bytes_per_swa_token = cfg._get_bytes_per_swa_token()
self.assertEqual(cfg.bytes_per_swa_token, 3 * 585 + 320)
self.assertEqual(
cfg._get_bytes_per_full_token(),
0.1 * cfg.bytes_per_swa_token + 585 / 4 + 864 / 128 + 132 / 4,
)

def test_dsv4_paged_dspark_budget_reserves_window_and_draft_layers(self):
from sglang.srt.model_executor.pool_configurator import DSV4PoolConfigurator

Expand Down
Loading