diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index 6a69c2b02..5493c0aa4 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -2354,26 +2354,10 @@ class Scheduler( def get_new_batch_prefill(self) -> Optional[ScheduleBatch]: prefill_delayer_single_pass = None if self.prefill_delayer: - # Get token usage from several pools - token_usage = None - if self.is_hybrid_swa: - _, _, full_token_usage, swa_token_usage, *_ = self._get_swa_token_info() - token_usage = max(full_token_usage, swa_token_usage) - if self.is_hybrid_ssm: - _, _, full_token_usage, mamba_token_usage, *_ = ( - self._get_mamba_token_info() - ) - token_usage = ( - max(token_usage, mamba_token_usage) - if token_usage is not None - else max(full_token_usage, mamba_token_usage) - ) - if token_usage is None: - _, token_usage, _, _ = self._get_token_info() - - assert token_usage is not None + # Get max usage across all pools for prefill delay decision + max_pool_usage = self.get_pool_stats().get_max_pool_usage() prefill_delayer_single_pass = PrefillDelayerSinglePassExecutor( - self.prefill_delayer, token_usage=token_usage + self.prefill_delayer, token_usage=max_pool_usage ) ret = self._get_new_batch_prefill_raw( diff --git a/python/sglang/srt/managers/scheduler_runtime_checker_mixin.py b/python/sglang/srt/managers/scheduler_runtime_checker_mixin.py index d9672ffa0..c1c1ab48b 100644 --- a/python/sglang/srt/managers/scheduler_runtime_checker_mixin.py +++ b/python/sglang/srt/managers/scheduler_runtime_checker_mixin.py @@ -1,9 +1,10 @@ from __future__ import annotations +import dataclasses import logging import time import warnings -from typing import TYPE_CHECKING +from typing import TYPE_CHECKING, List, Optional, Tuple from sglang.srt.disaggregation.utils import DisaggregationMode from sglang.srt.environ import envs @@ -20,6 +21,90 @@ if TYPE_CHECKING: logger = logging.getLogger(__name__) +@dataclasses.dataclass +class PoolStats: + # For full pools (required) + full_num_used: int + full_token_usage: float + full_available_size: int + full_evictable_size: int + + is_hybrid_swa: bool = False + is_hybrid_ssm: bool = False + + # For hybrid-swa pools + swa_num_used: Optional[int] = None + swa_token_usage: Optional[float] = None + swa_available_size: Optional[int] = None + swa_evictable_size: Optional[int] = None + + # For mamba pools + mamba_num_used: Optional[int] = None + mamba_usage: Optional[float] = None + mamba_available_size: Optional[int] = None + mamba_evictable_size: Optional[int] = None + + def get_kv_token_stats(self) -> Tuple[int, float]: + # NOTE: mamba pool is not included in the "token usage" calculation. + if self.is_hybrid_swa: + num_used = max(self.full_num_used, self.swa_num_used) + token_usage = max(self.full_token_usage, self.swa_token_usage) + else: + num_used = self.full_num_used + token_usage = self.full_token_usage + + return num_used, token_usage + + def get_max_pool_usage(self) -> float: + usage = self.full_token_usage + if self.is_hybrid_swa: + usage = max(usage, self.swa_token_usage) + if self.is_hybrid_ssm: + usage = max(usage, self.mamba_usage) + assert usage is not None and usage >= 0, f"{usage=} is not valid" + return usage + + def get_prefill_usage_msg_parts(self) -> List[str]: + parts = [] + if self.is_hybrid_swa: + parts += [ + f"full token usage: {self.full_token_usage:.2f}", + f"swa token usage: {self.swa_token_usage:.2f}", + ] + if self.is_hybrid_ssm: + if not self.is_hybrid_swa: + parts.append(f"full token usage: {self.full_token_usage:.2f}") + parts.append(f"mamba usage: {self.mamba_usage:.2f}") + if not parts: + parts.append(f"token usage: {self.full_token_usage:.2f}") + return parts + + def get_decode_usage_msg_parts(self) -> List[str]: + parts = [] + if self.is_hybrid_swa: + parts += [ + f"#full token: {self.full_num_used}", + f"full token usage: {self.full_token_usage:.2f}", + f"#swa token: {self.swa_num_used}", + f"swa token usage: {self.swa_token_usage:.2f}", + ] + if self.is_hybrid_ssm: + if not self.is_hybrid_swa: + parts += [ + f"#full token: {self.full_num_used}", + f"full token usage: {self.full_token_usage:.2f}", + ] + parts += [ + f"mamba num: {self.mamba_num_used}", + f"mamba usage: {self.mamba_usage:.2f}", + ] + if not parts: + parts.append( + f"#token: {self.full_num_used}, token usage: {self.full_token_usage:.2f}" + ) + return parts + + class SchedulerRuntimeCheckerMixin: def _session_held_tokens(self: Scheduler) -> int: if isinstance(self.tree_cache, SessionAwareCache): @@ -41,12 +126,36 @@ class SchedulerRuntimeCheckerMixin: return self.tree_cache.session_held_req_count() return 0 - def _get_token_info(self: Scheduler): + def get_pool_stats(self: Scheduler) -> PoolStats: + if self.is_hybrid_swa: + pool_stats = self._get_swa_token_info() + elif self.is_hybrid_ssm: + return self._get_mamba_token_info() + else: + return self._get_token_info() + + # swa + ssm can coexist: overlay mamba fields onto swa stats + if self.is_hybrid_ssm: + mamba_stats = self._get_mamba_token_info() + pool_stats.is_hybrid_ssm = True + pool_stats.mamba_num_used = mamba_stats.mamba_num_used + pool_stats.mamba_usage = mamba_stats.mamba_usage + pool_stats.mamba_available_size = mamba_stats.mamba_available_size + pool_stats.mamba_evictable_size = mamba_stats.mamba_evictable_size + + return pool_stats + + def _get_token_info(self: Scheduler) -> PoolStats: available_size = self.token_to_kv_pool_allocator.available_size() evictable_size = self.tree_cache.evictable_size() num_used = self.max_total_num_tokens - (available_size + evictable_size) token_usage = num_used / self.max_total_num_tokens - return num_used, token_usage, available_size, evictable_size + return PoolStats( + full_num_used=num_used, + full_token_usage=token_usage, + full_available_size=available_size, + full_evictable_size=evictable_size, + ) def _get_mamba_token_info(self: Scheduler): is_mamba_radix_cache = ( @@ -68,18 +177,20 @@ class SchedulerRuntimeCheckerMixin: ) full_token_usage = full_num_used / self.token_to_kv_pool_allocator.size mamba_usage = mamba_num_used / self.req_to_token_pool.mamba_pool.size - return ( - full_num_used, - mamba_num_used, - full_token_usage, - mamba_usage, - full_available_size, - full_evictable_size, - mamba_available_size, - mamba_evictable_size, + + return PoolStats( + is_hybrid_ssm=True, + full_num_used=full_num_used, + full_token_usage=full_token_usage, + full_available_size=full_available_size, + full_evictable_size=full_evictable_size, + mamba_num_used=mamba_num_used, + mamba_usage=mamba_usage, + mamba_available_size=mamba_available_size, + mamba_evictable_size=mamba_evictable_size, ) - def _get_swa_token_info(self: Scheduler): + def _get_swa_token_info(self: Scheduler) -> PoolStats: full_available_size = self.token_to_kv_pool_allocator.full_available_size() full_evictable_size = self.tree_cache.full_evictable_size() swa_available_size = self.token_to_kv_pool_allocator.swa_available_size() @@ -92,28 +203,27 @@ class SchedulerRuntimeCheckerMixin: ) full_token_usage = full_num_used / self.full_tokens_per_layer swa_token_usage = swa_num_used / self.swa_tokens_per_layer - return ( - full_num_used, - swa_num_used, - full_token_usage, - swa_token_usage, - full_available_size, - full_evictable_size, - swa_available_size, - swa_evictable_size, + + return PoolStats( + is_hybrid_swa=True, + full_num_used=full_num_used, + full_token_usage=full_token_usage, + full_available_size=full_available_size, + full_evictable_size=full_evictable_size, + swa_num_used=swa_num_used, + swa_token_usage=swa_token_usage, + swa_available_size=swa_available_size, + swa_evictable_size=swa_evictable_size, ) def _check_hybrid_memory(self: Scheduler): - ( - full_num_used, - swa_num_used, - _, - _, - full_available_size, - full_evictable_size, - swa_available_size, - swa_evictable_size, - ) = self._get_swa_token_info() + pool_stats = self._get_swa_token_info() + full_num_used = pool_stats.full_num_used + swa_num_used = pool_stats.swa_num_used + full_available_size = pool_stats.full_available_size + full_evictable_size = pool_stats.full_evictable_size + swa_available_size = pool_stats.swa_available_size + swa_evictable_size = pool_stats.swa_evictable_size session_held_full = self._session_held_full_tokens() session_held_swa = self._session_held_swa_tokens() @@ -132,16 +242,13 @@ class SchedulerRuntimeCheckerMixin: return memory_leak, token_msg def _check_mamba_memory(self: Scheduler): - ( - full_num_used, - mamba_num_used, - _, - _, - full_available_size, - full_evictable_size, - mamba_available_size, - mamba_evictable_size, - ) = self._get_mamba_token_info() + pool_stats = self._get_mamba_token_info() + full_num_used = pool_stats.full_num_used + mamba_num_used = pool_stats.mamba_num_used + full_available_size = pool_stats.full_available_size + full_evictable_size = pool_stats.full_evictable_size + mamba_available_size = pool_stats.mamba_available_size + mamba_evictable_size = pool_stats.mamba_evictable_size session_held = self._session_held_tokens() memory_leak = ( full_num_used != self.tree_cache.full_protected_size() + session_held @@ -181,7 +288,9 @@ class SchedulerRuntimeCheckerMixin: return memory_leak, token_msg def _check_radix_cache_memory(self: Scheduler): - _, _, available_size, evictable_size = self._get_token_info() + pool_stats = self._get_token_info() + available_size = pool_stats.full_available_size + evictable_size = pool_stats.full_evictable_size protected_size = self.tree_cache.protected_size() session_held = self._session_held_tokens() memory_leak = (available_size + evictable_size) != ( @@ -219,7 +328,9 @@ class SchedulerRuntimeCheckerMixin: ) return - _, _, available_size, evictable_size = self._get_token_info() + pool_stats = self._get_token_info() + available_size = pool_stats.full_available_size + evictable_size = pool_stats.full_evictable_size protected_size = self.tree_cache.protected_size() uncached_size = self._get_batch_uncached_size(current_batch) @@ -294,40 +405,15 @@ class SchedulerRuntimeCheckerMixin: and time.perf_counter() > self.metrics_collector.last_log_time + 30 ): # During idle time, also collect metrics every 30 seconds. - if self.is_hybrid_swa: - ( - full_num_used, - swa_num_used, - full_token_usage, - swa_token_usage, - _, - _, - _, - _, - ) = self._get_swa_token_info() - num_used = max(full_num_used, swa_num_used) - token_usage = max(full_token_usage, swa_token_usage) - elif self.is_hybrid_ssm: - ( - num_used, - _, - full_token_usage, - mamba_usage, - _, - _, - _, - _, - ) = self._get_mamba_token_info() - token_usage = max(full_token_usage, mamba_usage) - else: - num_used, token_usage, _, _ = self._get_token_info() + pool_stats = self.get_pool_stats() + num_used, _ = pool_stats.get_kv_token_stats() priority_enabled = self.enable_priority_scheduling self.stats.num_running_reqs = QueueCount.from_reqs( self.running_batch.reqs, priority_enabled ) self.stats.num_used_tokens = num_used - self.stats.token_usage = round(token_usage, 2) + self.stats.token_usage = round(pool_stats.get_max_pool_usage(), 2) self.stats.gen_throughput = 0 self.stats.num_queue_reqs = QueueCount.from_reqs( self.waiting_queue, priority_enabled diff --git a/python/sglang/srt/managers/tp_worker.py b/python/sglang/srt/managers/tp_worker.py index 347805f63..226139dbb 100644 --- a/python/sglang/srt/managers/tp_worker.py +++ b/python/sglang/srt/managers/tp_worker.py @@ -17,7 +17,7 @@ from __future__ import annotations import logging from abc import ABC, abstractmethod -from typing import TYPE_CHECKING, List, Optional +from typing import TYPE_CHECKING, List, Optional, Tuple import torch @@ -86,7 +86,7 @@ class BaseTpWorker(ABC): def get_pad_input_ids_func(self): return getattr(self.model_runner.model, "pad_input_ids", None) - def get_memory_pool(self): + def get_memory_pool(self) -> Tuple[ReqToTokenPool, BaseTokenToKVPoolAllocator]: return ( self.model_runner.req_to_token_pool, self.model_runner.token_to_kv_pool_allocator, diff --git a/python/sglang/srt/observability/scheduler_metrics_mixin.py b/python/sglang/srt/observability/scheduler_metrics_mixin.py index f28c74d81..dc28c57ea 100644 --- a/python/sglang/srt/observability/scheduler_metrics_mixin.py +++ b/python/sglang/srt/observability/scheduler_metrics_mixin.py @@ -30,7 +30,6 @@ from sglang.srt.observability.metrics_collector import ( SchedulerStats, compute_routing_key_stats, ) -from sglang.srt.utils import get_bool_env_var from sglang.srt.utils.device_timer import DeviceTimer, GapTimer from sglang.srt.utils.scheduler_status_logger import SchedulerStatusLogger @@ -41,7 +40,7 @@ if TYPE_CHECKING: logger = logging.getLogger(__name__) -RECORD_STEP_TIME = get_bool_env_var("SGLANG_RECORD_STEP_TIME") +RECORD_STEP_TIME = envs.SGLANG_RECORD_STEP_TIME.get() LOG_FORWARD_ITERS = envs.SGLANG_LOG_FORWARD_ITERS.get() ENABLE_METRICS_DEVICE_TIMER = envs.SGLANG_ENABLE_METRICS_DEVICE_TIMER.get() @@ -343,47 +342,11 @@ class SchedulerMetricsMixin: self.last_input_throughput = self.last_prefill_tokens / gap_latency self.last_prefill_tokens = prefill_stats.log_input_tokens - # TODO: generalize this for various memory pools - msg_parts = [] - num_used = token_usage = full_token_usage = None - - if self.is_hybrid_swa: - full_num_used, swa_num_used, full_tok, swa_token_usage, *_ = ( - self._get_swa_token_info() - ) - num_used = max(full_num_used, swa_num_used) - token_usage = max(full_tok, swa_token_usage) - full_token_usage = full_tok - msg_parts += [ - f"full token usage: {full_tok:.2f}", - f"swa token usage: {swa_token_usage:.2f}", - ] - - if self.is_hybrid_ssm: - num_used_m, _, full_tok_m, mamba_usage, *_ = self._get_mamba_token_info() - num_used = max(num_used, num_used_m) if num_used is not None else num_used_m - token_usage = ( - max(token_usage, mamba_usage) - if token_usage is not None - else max(full_tok_m, mamba_usage) - ) - if full_token_usage is None: - full_token_usage = full_tok_m - msg_parts.append(f"full token usage: {full_tok_m:.2f}") - msg_parts.append(f"mamba usage: {mamba_usage:.2f}") - - if full_token_usage is None: - num_used, tok, _, _ = self._get_token_info() - full_token_usage = tok - token_usage = tok - msg_parts.append(f"token usage: {tok:.2f}") - - assert ( - num_used is not None - and token_usage is not None - and full_token_usage is not None - ) - token_usage_msg = ", ".join(msg_parts) + ", " + pool_stats = self.get_pool_stats() + num_used, _ = pool_stats.get_kv_token_stats() + max_pool_usage = pool_stats.get_max_pool_usage() + full_token_usage = pool_stats.full_token_usage + token_usage_msg = ", ".join(pool_stats.get_prefill_usage_msg_parts()) + ", " self.stats.new_token_ratio = prefill_stats.new_token_ratio iter_msg = f" [{self.forward_ct + 1}]" if LOG_FORWARD_ITERS else "" @@ -455,12 +418,12 @@ class SchedulerMetricsMixin: self.stats.num_running_reqs = prefill_stats.num_running_reqs self.stats.num_running_reqs_offline_batch = 0 self.stats.num_used_tokens = num_used - self.stats.token_usage = token_usage + self.stats.token_usage = max_pool_usage self.stats.full_token_usage = full_token_usage - if self.is_hybrid_swa: - self.stats.swa_token_usage = swa_token_usage - if self.is_hybrid_ssm: - self.stats.mamba_usage = mamba_usage + if pool_stats.is_hybrid_swa: + self.stats.swa_token_usage = pool_stats.swa_token_usage + if pool_stats.is_hybrid_ssm: + self.stats.mamba_usage = pool_stats.mamba_usage priority_enabled = self.enable_priority_scheduling self.stats.num_queue_reqs = QueueCount.from_reqs( @@ -551,57 +514,11 @@ class SchedulerMetricsMixin: num_running_reqs = len(batch.reqs) num_running_reqs_offline_batch = 0 - # TODO: generalize this for various memory pools - msg_parts = [] - num_used = token_usage = full_token_usage = None - - if self.is_hybrid_swa: - full_num_used, swa_num_used, full_tok, swa_token_usage, *_ = ( - self._get_swa_token_info() - ) - num_used = max(full_num_used, swa_num_used) - token_usage = max(full_tok, swa_token_usage) - full_token_usage = full_tok - msg_parts += [ - f"#full token: {full_num_used}", - f"full token usage: {full_tok:.2f}", - f"#swa token: {swa_num_used}", - f"swa token usage: {swa_token_usage:.2f}", - ] - - if self.is_hybrid_ssm: - num_used_m, mamba_num, full_tok_m, mamba_usage, *_ = ( - self._get_mamba_token_info() - ) - num_used = max(num_used, num_used_m) if num_used is not None else num_used_m - token_usage = ( - max(token_usage, mamba_usage) - if token_usage is not None - else max(full_tok_m, mamba_usage) - ) - if full_token_usage is None: - full_token_usage = full_tok_m - msg_parts += [ - f"#full token: {num_used_m}", - f"full token usage: {full_tok_m:.2f}", - ] - msg_parts += [ - f"mamba num: {mamba_num}", - f"mamba usage: {mamba_usage:.2f}", - ] - - if full_token_usage is None: - num_used, tok, _, _ = self._get_token_info() - full_token_usage = tok - token_usage = tok - msg_parts.append(f"#token: {num_used}, token usage: {tok:.2f}") - - assert ( - num_used is not None - and token_usage is not None - and full_token_usage is not None - ) - token_usage_msg = ", ".join(msg_parts) + ", " + pool_stats = self.get_pool_stats() + num_used, _ = pool_stats.get_kv_token_stats() + max_pool_usage = pool_stats.get_max_pool_usage() + full_token_usage = pool_stats.full_token_usage + token_usage_msg = ", ".join(pool_stats.get_decode_usage_msg_parts()) + ", " if RECORD_STEP_TIME: self.step_time_dict[num_running_reqs].append( @@ -688,13 +605,13 @@ class SchedulerMetricsMixin: self.stats.num_running_reqs_offline_batch = num_running_reqs_offline_batch self.stats.num_used_tokens = num_used # maximum usage of all pools - self.stats.token_usage = token_usage + self.stats.token_usage = max_pool_usage # usage of full attention self.stats.full_token_usage = full_token_usage - if self.is_hybrid_swa: - self.stats.swa_token_usage = swa_token_usage - if self.is_hybrid_ssm: - self.stats.mamba_usage = mamba_usage + if pool_stats.is_hybrid_swa: + self.stats.swa_token_usage = pool_stats.swa_token_usage + if pool_stats.is_hybrid_ssm: + self.stats.mamba_usage = pool_stats.mamba_usage self.stats.decode_sum_seq_lens = batch.seq_lens_cpu.sum().item() self.stats.gen_throughput = self.last_gen_throughput self.stats.num_queue_reqs = QueueCount.from_reqs( @@ -887,14 +804,7 @@ class SchedulerMetricsMixin: return num_pending_tokens def get_load(self: Scheduler, _: GetLoadReqInput = None) -> GetLoadReqOutput: - if self.is_hybrid_swa: - full_num_used, swa_num_used, *_ = self._get_swa_token_info() - num_tokens = max(full_num_used, swa_num_used) - elif self.is_hybrid_ssm: - num_tokens = self._get_mamba_token_info()[0] - else: - num_tokens = self._get_token_info()[0] - + num_tokens, _ = self.get_pool_stats().get_kv_token_stats() num_pending_tokens = self._get_num_pending_tokens() # Tokens and request count in waiting queue, bootstrap queue, prealloc queue @@ -945,20 +855,7 @@ class SchedulerMetricsMixin: waiting_queues.append(self.disagg_decode_prealloc_queue.retracted_queue) num_waiting_reqs = sum(len(queue) for queue in waiting_queues) - - if self.is_hybrid_swa: - full_num_used, swa_num_used, *_ = self._get_swa_token_info() - num_used_tokens = max(full_num_used, swa_num_used) - elif self.is_hybrid_ssm: - num_used_tokens = self._get_mamba_token_info()[0] - else: - num_used_tokens = self._get_token_info()[0] - - token_usage = ( - num_used_tokens / self.max_total_num_tokens - if self.max_total_num_tokens > 0 - else 0.0 - ) + num_used_tokens, kv_token_usage = self.get_pool_stats().get_kv_token_stats() memory = None if include_all or "memory" in include: @@ -1044,7 +941,7 @@ class SchedulerMetricsMixin: num_waiting_reqs=num_waiting_reqs, num_used_tokens=num_used_tokens, max_total_num_tokens=self.max_total_num_tokens, - token_usage=round(token_usage, 4), + token_usage=round(kv_token_usage, 4), gen_throughput=round(self.stats.gen_throughput, 2), cache_hit_rate=round(self.stats.cache_hit_rate, 4), utilization=round(self.stats.utilization, 4), diff --git a/test/registered/unit/managers/test_scheduler_pause_generation.py b/test/registered/unit/managers/test_scheduler_pause_generation.py index 4e94626d4..46bc664eb 100644 --- a/test/registered/unit/managers/test_scheduler_pause_generation.py +++ b/test/registered/unit/managers/test_scheduler_pause_generation.py @@ -9,6 +9,7 @@ maybe_stub_sgl_kernel() from sglang.srt.managers.io_struct import PauseGenerationReqInput from sglang.srt.managers.scheduler import Scheduler +from sglang.srt.managers.scheduler_runtime_checker_mixin import PoolStats register_cpu_ci(est_time=10, suite="stage-a-test-cpu") @@ -33,7 +34,14 @@ class TestSchedulerPauseGeneration(unittest.TestCase): scheduler.token_to_kv_pool_allocator = MagicMock() scheduler.token_to_kv_pool_allocator.available_size.return_value = 1000 scheduler.max_total_num_tokens = 1000 - scheduler._get_token_info = MagicMock(return_value=(0, 0, 1000, 0)) + scheduler._get_token_info = MagicMock( + return_value=PoolStats( + full_num_used=0, + full_token_usage=0, + full_available_size=1000, + full_evictable_size=0, + ) + ) return scheduler def test_inplace_only_sets_flag(self):