Skip to content
2 changes: 1 addition & 1 deletion python/sglang/srt/disaggregation/prefill.py
Original file line number Diff line number Diff line change
Expand Up @@ -189,7 +189,7 @@ def __init__(
self.scheduler = scheduler
self.scheduler_stage_metrics = scheduler_stage_metrics
self.max_total_num_tokens = (
self.scheduler.tp_worker.model_runner.effective_max_total_num_tokens
self.scheduler.tp_worker.model_runner.effective_logical_max_total_num_tokens
)
self.transfer_backend = transfer_backend
if envs.SGLANG_DISAGG_STAGING_BUFFER.get():
Expand Down
8 changes: 4 additions & 4 deletions python/sglang/srt/managers/scheduler.py
Original file line number Diff line number Diff line change
Expand Up @@ -1250,7 +1250,8 @@ def emit_metrics_constants(self) -> None:
# TODO: max_running_requests_under_SLO has no setter — dead chain.
max_running_requests_under_SLO=None,
page_size=self.page_size,
num_pages=self.max_total_num_tokens // self.page_size,
num_pages=self.max_total_num_tokens
// self.token_to_kv_pool_allocator.page_size,
context_len=self.model_config.context_len,
startup_available_gpu_memory_gb=self.startup_available_gpu_memory_gb,
)
Expand Down Expand Up @@ -2345,8 +2346,7 @@ def init_pool_stats_observer(self) -> None:
enable_hisparse=self.enable_hisparse,
full_tokens_per_layer=self.full_tokens_per_layer,
swa_tokens_per_layer=self.swa_tokens_per_layer,
max_total_num_tokens=self.max_total_num_tokens
* get_parallel().attn_dcp_size,
max_total_num_tokens=self.max_total_num_tokens,
get_last_batch=lambda: self.last_batch,
get_running_batch=lambda: self.running_batch,
)
Expand Down Expand Up @@ -2527,7 +2527,7 @@ def init_req_max_new_tokens(self, req):
max_new_tokens = self.token_to_kv_pool_allocator.max_new_tokens_for_memory(
input_len,
max_new_tokens,
token_capacity=self.max_total_num_tokens * get_parallel().attn_dcp_size,
token_capacity=self.max_total_num_tokens,
sliding_window_size=self.sliding_window_size,
chunk_size=self.chunked_prefill_size,
)
Expand Down
12 changes: 3 additions & 9 deletions python/sglang/srt/managers/tp_worker.py
Original file line number Diff line number Diff line change
Expand Up @@ -409,9 +409,7 @@ def alloc_memory_pool(
assert self.model_runner.max_running_requests > 0, "max_running_request is zero"
max_req_len = min(
self.model_config.context_len - 1,
self.model_runner.effective_max_total_num_tokens
* get_parallel().attn_dcp_size
- 1,
self.model_runner.effective_logical_max_total_num_tokens - 1,
)
assert max_req_len > 0, "Memory pool size is too small"

Expand Down Expand Up @@ -539,17 +537,13 @@ def register_hisparse_coordinator(self, coordinator):
def get_worker_info(self):
max_req_len = min(
self.model_config.context_len - 1,
self.model_runner.effective_max_total_num_tokens
* get_parallel().attn_dcp_size
- 1,
self.model_runner.effective_logical_max_total_num_tokens - 1,
)
max_req_input_len = max_req_len - 5
if self.dllm_algorithm is not None:
max_req_input_len -= self.dllm_algorithm.block_size
return (
self.model_runner.req_to_token_pool.schedulable_token_capacity(
self.model_runner.max_total_num_tokens
),
self.model_runner.logical_max_total_num_tokens,
get_schedule().max_prefill_tokens,
self.model_runner.max_running_requests,
get_schedule().max_queued_requests,
Expand Down
11 changes: 11 additions & 0 deletions python/sglang/srt/mem_cache/kv_cache_configurator.py
Original file line number Diff line number Diff line change
Expand Up @@ -316,6 +316,17 @@ def hybrid_swa_token_capacity(
)
return full_capacity or swa_capacity

def logical_token_capacity(self, *, max_total_num_tokens: int) -> int:
Comment thread
chromecast56 marked this conversation as resolved.
"""Request tokens a pool of `max_total_num_tokens` per-rank rows holds."""
# SWA allocators never widen under DCP.
if self.is_hybrid_swa:
return max_total_num_tokens
# Target rows widen into attn_dcp_size ids; draft sizes already carry
# loc_space_scale.
return (
max_total_num_tokens * get_parallel().attn_dcp_size // self.loc_space_scale
)

def _build_fp4_quant_method(self, *, num_layers: int):
if not is_float4_e2m1fn_x2(self.kv_cache_dtype):
return None
Expand Down
16 changes: 16 additions & 0 deletions python/sglang/srt/model_executor/model_runner.py
Original file line number Diff line number Diff line change
Expand Up @@ -1390,6 +1390,22 @@ def unload_lora_adapter(self, lora_ref: LoRARef):
"""Unload a lora adapter that was previously loaded during initialization or dynamic loading."""
return self.lora_manager.unload_lora_adapter(lora_ref)

@property
def logical_max_total_num_tokens(self):
"""Request-token capacity in logical tokens, not per-rank DCP rows."""
return self.req_to_token_pool.schedulable_token_capacity(
self.kv_cache_configurator.logical_token_capacity(
max_total_num_tokens=self.max_total_num_tokens
)
)

@property
def effective_logical_max_total_num_tokens(self):
"""Logical request limit, preserving hybrid SWA's separate pool bounds."""
if self.is_hybrid_swa:
return self.effective_max_total_num_tokens
return self.logical_max_total_num_tokens

@property
def effective_max_total_num_tokens(self):
"""Return the max token pool size considering hybrid swa settings."""
Expand Down
139 changes: 139 additions & 0 deletions test/registered/unit/managers/test_dcp_logical_capacity.py
Comment thread
chromecast56 marked this conversation as resolved.
Original file line number Diff line number Diff line change
@@ -0,0 +1,139 @@
"""Logical DCP capacities must agree across validation, admission and telemetry."""

import unittest
from types import SimpleNamespace as NS
from unittest.mock import patch

import torch

from sglang.srt.managers.tp_worker import TpModelWorker
from sglang.srt.mem_cache.allocator.unified_mamba import (
UnifiedMambaTokenToKVPoolAllocator,
)
from sglang.srt.mem_cache.kv_cache_configurator import KVCacheConfigurator
from sglang.srt.mem_cache.memory_pool import ReqToTokenPool
from sglang.srt.mem_cache.unified_memory_pool import (
MambaSubPoolSpec,
MLASubPoolSpec,
UnifiedKVPool,
)
from sglang.srt.model_executor.model_runner import ModelRunner
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase, enter_scope, published_topology

register_cpu_ci(est_time=5, suite="base-a-test-cpu")

PHYSICAL = 690240
CONTEXT = 1048576


def make_configurator(*, is_hybrid_swa=False, is_draft_worker=False):
configurator = KVCacheConfigurator.__new__(KVCacheConfigurator)
configurator.is_hybrid_swa = is_hybrid_swa
configurator.is_draft_worker = is_draft_worker
return configurator


def make_unified_mamba_allocator(*, n_full_tokens, n_mamba_slots):
full = MLASubPoolSpec(
name="full",
layer_num=3,
kv_lora_rank=6,
qk_rope_head_dim=2,
store_dtype=torch.bfloat16,
grow_direction="down",
)
mamba = MambaSubPoolSpec(
name="mamba",
layer_num=2,
conv_state_shapes=((4, 3),),
conv_dtype=torch.float32,
temporal_state_shape=(2, 2, 2),
temporal_dtype=torch.float32,
grow_direction="up",
)
pool = UnifiedKVPool(
total_bytes=full.entry_bytes() * n_full_tokens
+ mamba.entry_bytes() * n_mamba_slots,
sub_pool_specs=[full, mamba],
device="cpu",
enable_memory_saver=False,
page_size=1,
)
kvcache = NS(
full_kv_pool=NS(buf=torch.empty(pool.max_slots("full"))),
mamba_pool=NS(buf=torch.empty(pool.max_slots("mamba"))),
)
return UnifiedMambaTokenToKVPoolAllocator(
unified_buffer=pool, kvcache=kvcache, device="cpu", page_size=1
)


def make_worker(dcp_size):
kv = NS(size=PHYSICAL, mem_usage=8.89)
runner = ModelRunner.__new__(ModelRunner)
runner.server_args = NS(dcp_size=dcp_size)
runner.kv_cache_configurator = make_configurator()
runner.is_hybrid_swa = False
runner.max_total_num_tokens = PHYSICAL
runner.max_running_requests = 64
runner.token_to_kv_pool = kv
runner.req_to_token_pool = ReqToTokenPool.__new__(ReqToTokenPool)
runner.req_to_token_pool.size = 96
runner.req_to_token_pool.max_context_len = CONTEXT
runner.req_to_token_pool._aux_cache = None
runner.forward_stream = None
return NS(
model_runner=runner,
model_config=NS(context_len=CONTEXT),
random_seed=0,
device="cpu",
dllm_algorithm=None,
)


class TestDcpLogicalCapacity(CustomTestCase):
def make_worker(self, dcp_size):
enter_scope(self, published_topology(tp_size=dcp_size, dcp_size=dcp_size))
return make_worker(dcp_size)

def setUp(self):
for config in (
patch(
"sglang.srt.managers.tp_worker.get_schedule",
return_value=NS(max_prefill_tokens=16384, max_queued_requests=None),
),
):
config.start()
self.addCleanup(config.stop)

def test_logical_capacity(self):
for dcp_size in (1, 2, 8):
with self.subTest(dcp_size=dcp_size):
info = TpModelWorker.get_worker_info(self.make_worker(dcp_size))
self.assertEqual(info[0], PHYSICAL * dcp_size)
self.assertEqual(info[4], min(CONTEXT, PHYSICAL * dcp_size) - 1)

# Unified Mamba's allocator.size also counts Mamba state bytes, so
# capacity must come from the configured rows.
runner = self.make_worker(8).model_runner
runner.max_total_num_tokens = 64
runner.token_to_kv_pool_allocator = make_unified_mamba_allocator(
n_full_tokens=64, n_mamba_slots=8
)
self.assertGreater(runner.token_to_kv_pool_allocator.size, 64 * 8)
self.assertEqual(runner.logical_max_total_num_tokens, 64 * 8)

# Draft sizes already carry loc_space_scale; SWA never widens.
runner.kv_cache_configurator = make_configurator(is_draft_worker=True)
self.assertEqual(runner.logical_max_total_num_tokens, 64)
runner.is_hybrid_swa = True
runner.kv_cache_configurator = make_configurator(is_hybrid_swa=True)
runner.full_max_total_num_tokens = 32
runner.swa_max_total_num_tokens = 16
self.assertEqual(runner.logical_max_total_num_tokens, 64)
self.assertEqual(runner.effective_logical_max_total_num_tokens, 32)


if __name__ == "__main__":
unittest.main()
Loading