[mem] Introduce PoolStats dataclass; unify pool metrics and token_usage (#22554)
This commit is contained in:
@@ -2354,26 +2354,10 @@ class Scheduler(
|
|||||||
def get_new_batch_prefill(self) -> Optional[ScheduleBatch]:
|
def get_new_batch_prefill(self) -> Optional[ScheduleBatch]:
|
||||||
prefill_delayer_single_pass = None
|
prefill_delayer_single_pass = None
|
||||||
if self.prefill_delayer:
|
if self.prefill_delayer:
|
||||||
# Get token usage from several pools
|
# Get max usage across all pools for prefill delay decision
|
||||||
token_usage = None
|
max_pool_usage = self.get_pool_stats().get_max_pool_usage()
|
||||||
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
|
|
||||||
prefill_delayer_single_pass = PrefillDelayerSinglePassExecutor(
|
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(
|
ret = self._get_new_batch_prefill_raw(
|
||||||
|
|||||||
@@ -1,9 +1,10 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import dataclasses
|
||||||
import logging
|
import logging
|
||||||
import time
|
import time
|
||||||
import warnings
|
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.disaggregation.utils import DisaggregationMode
|
||||||
from sglang.srt.environ import envs
|
from sglang.srt.environ import envs
|
||||||
@@ -20,6 +21,90 @@ if TYPE_CHECKING:
|
|||||||
logger = logging.getLogger(__name__)
|
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:
|
class SchedulerRuntimeCheckerMixin:
|
||||||
def _session_held_tokens(self: Scheduler) -> int:
|
def _session_held_tokens(self: Scheduler) -> int:
|
||||||
if isinstance(self.tree_cache, SessionAwareCache):
|
if isinstance(self.tree_cache, SessionAwareCache):
|
||||||
@@ -41,12 +126,36 @@ class SchedulerRuntimeCheckerMixin:
|
|||||||
return self.tree_cache.session_held_req_count()
|
return self.tree_cache.session_held_req_count()
|
||||||
return 0
|
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()
|
available_size = self.token_to_kv_pool_allocator.available_size()
|
||||||
evictable_size = self.tree_cache.evictable_size()
|
evictable_size = self.tree_cache.evictable_size()
|
||||||
num_used = self.max_total_num_tokens - (available_size + evictable_size)
|
num_used = self.max_total_num_tokens - (available_size + evictable_size)
|
||||||
token_usage = num_used / self.max_total_num_tokens
|
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):
|
def _get_mamba_token_info(self: Scheduler):
|
||||||
is_mamba_radix_cache = (
|
is_mamba_radix_cache = (
|
||||||
@@ -68,18 +177,20 @@ class SchedulerRuntimeCheckerMixin:
|
|||||||
)
|
)
|
||||||
full_token_usage = full_num_used / self.token_to_kv_pool_allocator.size
|
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
|
mamba_usage = mamba_num_used / self.req_to_token_pool.mamba_pool.size
|
||||||
return (
|
|
||||||
full_num_used,
|
return PoolStats(
|
||||||
mamba_num_used,
|
is_hybrid_ssm=True,
|
||||||
full_token_usage,
|
full_num_used=full_num_used,
|
||||||
mamba_usage,
|
full_token_usage=full_token_usage,
|
||||||
full_available_size,
|
full_available_size=full_available_size,
|
||||||
full_evictable_size,
|
full_evictable_size=full_evictable_size,
|
||||||
mamba_available_size,
|
mamba_num_used=mamba_num_used,
|
||||||
mamba_evictable_size,
|
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_available_size = self.token_to_kv_pool_allocator.full_available_size()
|
||||||
full_evictable_size = self.tree_cache.full_evictable_size()
|
full_evictable_size = self.tree_cache.full_evictable_size()
|
||||||
swa_available_size = self.token_to_kv_pool_allocator.swa_available_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
|
full_token_usage = full_num_used / self.full_tokens_per_layer
|
||||||
swa_token_usage = swa_num_used / self.swa_tokens_per_layer
|
swa_token_usage = swa_num_used / self.swa_tokens_per_layer
|
||||||
return (
|
|
||||||
full_num_used,
|
return PoolStats(
|
||||||
swa_num_used,
|
is_hybrid_swa=True,
|
||||||
full_token_usage,
|
full_num_used=full_num_used,
|
||||||
swa_token_usage,
|
full_token_usage=full_token_usage,
|
||||||
full_available_size,
|
full_available_size=full_available_size,
|
||||||
full_evictable_size,
|
full_evictable_size=full_evictable_size,
|
||||||
swa_available_size,
|
swa_num_used=swa_num_used,
|
||||||
swa_evictable_size,
|
swa_token_usage=swa_token_usage,
|
||||||
|
swa_available_size=swa_available_size,
|
||||||
|
swa_evictable_size=swa_evictable_size,
|
||||||
)
|
)
|
||||||
|
|
||||||
def _check_hybrid_memory(self: Scheduler):
|
def _check_hybrid_memory(self: Scheduler):
|
||||||
(
|
pool_stats = self._get_swa_token_info()
|
||||||
full_num_used,
|
full_num_used = pool_stats.full_num_used
|
||||||
swa_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
|
||||||
full_available_size,
|
swa_available_size = pool_stats.swa_available_size
|
||||||
full_evictable_size,
|
swa_evictable_size = pool_stats.swa_evictable_size
|
||||||
swa_available_size,
|
|
||||||
swa_evictable_size,
|
|
||||||
) = self._get_swa_token_info()
|
|
||||||
session_held_full = self._session_held_full_tokens()
|
session_held_full = self._session_held_full_tokens()
|
||||||
session_held_swa = self._session_held_swa_tokens()
|
session_held_swa = self._session_held_swa_tokens()
|
||||||
|
|
||||||
@@ -132,16 +242,13 @@ class SchedulerRuntimeCheckerMixin:
|
|||||||
return memory_leak, token_msg
|
return memory_leak, token_msg
|
||||||
|
|
||||||
def _check_mamba_memory(self: Scheduler):
|
def _check_mamba_memory(self: Scheduler):
|
||||||
(
|
pool_stats = self._get_mamba_token_info()
|
||||||
full_num_used,
|
full_num_used = pool_stats.full_num_used
|
||||||
mamba_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
|
||||||
full_available_size,
|
mamba_available_size = pool_stats.mamba_available_size
|
||||||
full_evictable_size,
|
mamba_evictable_size = pool_stats.mamba_evictable_size
|
||||||
mamba_available_size,
|
|
||||||
mamba_evictable_size,
|
|
||||||
) = self._get_mamba_token_info()
|
|
||||||
session_held = self._session_held_tokens()
|
session_held = self._session_held_tokens()
|
||||||
memory_leak = (
|
memory_leak = (
|
||||||
full_num_used != self.tree_cache.full_protected_size() + session_held
|
full_num_used != self.tree_cache.full_protected_size() + session_held
|
||||||
@@ -181,7 +288,9 @@ class SchedulerRuntimeCheckerMixin:
|
|||||||
return memory_leak, token_msg
|
return memory_leak, token_msg
|
||||||
|
|
||||||
def _check_radix_cache_memory(self: Scheduler):
|
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()
|
protected_size = self.tree_cache.protected_size()
|
||||||
session_held = self._session_held_tokens()
|
session_held = self._session_held_tokens()
|
||||||
memory_leak = (available_size + evictable_size) != (
|
memory_leak = (available_size + evictable_size) != (
|
||||||
@@ -219,7 +328,9 @@ class SchedulerRuntimeCheckerMixin:
|
|||||||
)
|
)
|
||||||
return
|
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()
|
protected_size = self.tree_cache.protected_size()
|
||||||
|
|
||||||
uncached_size = self._get_batch_uncached_size(current_batch)
|
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
|
and time.perf_counter() > self.metrics_collector.last_log_time + 30
|
||||||
):
|
):
|
||||||
# During idle time, also collect metrics every 30 seconds.
|
# During idle time, also collect metrics every 30 seconds.
|
||||||
if self.is_hybrid_swa:
|
pool_stats = self.get_pool_stats()
|
||||||
(
|
num_used, _ = pool_stats.get_kv_token_stats()
|
||||||
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()
|
|
||||||
|
|
||||||
priority_enabled = self.enable_priority_scheduling
|
priority_enabled = self.enable_priority_scheduling
|
||||||
self.stats.num_running_reqs = QueueCount.from_reqs(
|
self.stats.num_running_reqs = QueueCount.from_reqs(
|
||||||
self.running_batch.reqs, priority_enabled
|
self.running_batch.reqs, priority_enabled
|
||||||
)
|
)
|
||||||
self.stats.num_used_tokens = num_used
|
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.gen_throughput = 0
|
||||||
self.stats.num_queue_reqs = QueueCount.from_reqs(
|
self.stats.num_queue_reqs = QueueCount.from_reqs(
|
||||||
self.waiting_queue, priority_enabled
|
self.waiting_queue, priority_enabled
|
||||||
|
|||||||
@@ -17,7 +17,7 @@ from __future__ import annotations
|
|||||||
|
|
||||||
import logging
|
import logging
|
||||||
from abc import ABC, abstractmethod
|
from abc import ABC, abstractmethod
|
||||||
from typing import TYPE_CHECKING, List, Optional
|
from typing import TYPE_CHECKING, List, Optional, Tuple
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
@@ -86,7 +86,7 @@ class BaseTpWorker(ABC):
|
|||||||
def get_pad_input_ids_func(self):
|
def get_pad_input_ids_func(self):
|
||||||
return getattr(self.model_runner.model, "pad_input_ids", None)
|
return getattr(self.model_runner.model, "pad_input_ids", None)
|
||||||
|
|
||||||
def get_memory_pool(self):
|
def get_memory_pool(self) -> Tuple[ReqToTokenPool, BaseTokenToKVPoolAllocator]:
|
||||||
return (
|
return (
|
||||||
self.model_runner.req_to_token_pool,
|
self.model_runner.req_to_token_pool,
|
||||||
self.model_runner.token_to_kv_pool_allocator,
|
self.model_runner.token_to_kv_pool_allocator,
|
||||||
|
|||||||
@@ -30,7 +30,6 @@ from sglang.srt.observability.metrics_collector import (
|
|||||||
SchedulerStats,
|
SchedulerStats,
|
||||||
compute_routing_key_stats,
|
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.device_timer import DeviceTimer, GapTimer
|
||||||
from sglang.srt.utils.scheduler_status_logger import SchedulerStatusLogger
|
from sglang.srt.utils.scheduler_status_logger import SchedulerStatusLogger
|
||||||
|
|
||||||
@@ -41,7 +40,7 @@ if TYPE_CHECKING:
|
|||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
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()
|
LOG_FORWARD_ITERS = envs.SGLANG_LOG_FORWARD_ITERS.get()
|
||||||
ENABLE_METRICS_DEVICE_TIMER = envs.SGLANG_ENABLE_METRICS_DEVICE_TIMER.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_input_throughput = self.last_prefill_tokens / gap_latency
|
||||||
self.last_prefill_tokens = prefill_stats.log_input_tokens
|
self.last_prefill_tokens = prefill_stats.log_input_tokens
|
||||||
|
|
||||||
# TODO: generalize this for various memory pools
|
pool_stats = self.get_pool_stats()
|
||||||
msg_parts = []
|
num_used, _ = pool_stats.get_kv_token_stats()
|
||||||
num_used = token_usage = full_token_usage = None
|
max_pool_usage = pool_stats.get_max_pool_usage()
|
||||||
|
full_token_usage = pool_stats.full_token_usage
|
||||||
if self.is_hybrid_swa:
|
token_usage_msg = ", ".join(pool_stats.get_prefill_usage_msg_parts()) + ", "
|
||||||
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) + ", "
|
|
||||||
|
|
||||||
self.stats.new_token_ratio = prefill_stats.new_token_ratio
|
self.stats.new_token_ratio = prefill_stats.new_token_ratio
|
||||||
iter_msg = f" [{self.forward_ct + 1}]" if LOG_FORWARD_ITERS else ""
|
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 = prefill_stats.num_running_reqs
|
||||||
self.stats.num_running_reqs_offline_batch = 0
|
self.stats.num_running_reqs_offline_batch = 0
|
||||||
self.stats.num_used_tokens = num_used
|
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
|
self.stats.full_token_usage = full_token_usage
|
||||||
if self.is_hybrid_swa:
|
if pool_stats.is_hybrid_swa:
|
||||||
self.stats.swa_token_usage = swa_token_usage
|
self.stats.swa_token_usage = pool_stats.swa_token_usage
|
||||||
if self.is_hybrid_ssm:
|
if pool_stats.is_hybrid_ssm:
|
||||||
self.stats.mamba_usage = mamba_usage
|
self.stats.mamba_usage = pool_stats.mamba_usage
|
||||||
|
|
||||||
priority_enabled = self.enable_priority_scheduling
|
priority_enabled = self.enable_priority_scheduling
|
||||||
self.stats.num_queue_reqs = QueueCount.from_reqs(
|
self.stats.num_queue_reqs = QueueCount.from_reqs(
|
||||||
@@ -551,57 +514,11 @@ class SchedulerMetricsMixin:
|
|||||||
num_running_reqs = len(batch.reqs)
|
num_running_reqs = len(batch.reqs)
|
||||||
num_running_reqs_offline_batch = 0
|
num_running_reqs_offline_batch = 0
|
||||||
|
|
||||||
# TODO: generalize this for various memory pools
|
pool_stats = self.get_pool_stats()
|
||||||
msg_parts = []
|
num_used, _ = pool_stats.get_kv_token_stats()
|
||||||
num_used = token_usage = full_token_usage = None
|
max_pool_usage = pool_stats.get_max_pool_usage()
|
||||||
|
full_token_usage = pool_stats.full_token_usage
|
||||||
if self.is_hybrid_swa:
|
token_usage_msg = ", ".join(pool_stats.get_decode_usage_msg_parts()) + ", "
|
||||||
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) + ", "
|
|
||||||
|
|
||||||
if RECORD_STEP_TIME:
|
if RECORD_STEP_TIME:
|
||||||
self.step_time_dict[num_running_reqs].append(
|
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_running_reqs_offline_batch = num_running_reqs_offline_batch
|
||||||
self.stats.num_used_tokens = num_used
|
self.stats.num_used_tokens = num_used
|
||||||
# maximum usage of all pools
|
# maximum usage of all pools
|
||||||
self.stats.token_usage = token_usage
|
self.stats.token_usage = max_pool_usage
|
||||||
# usage of full attention
|
# usage of full attention
|
||||||
self.stats.full_token_usage = full_token_usage
|
self.stats.full_token_usage = full_token_usage
|
||||||
if self.is_hybrid_swa:
|
if pool_stats.is_hybrid_swa:
|
||||||
self.stats.swa_token_usage = swa_token_usage
|
self.stats.swa_token_usage = pool_stats.swa_token_usage
|
||||||
if self.is_hybrid_ssm:
|
if pool_stats.is_hybrid_ssm:
|
||||||
self.stats.mamba_usage = mamba_usage
|
self.stats.mamba_usage = pool_stats.mamba_usage
|
||||||
self.stats.decode_sum_seq_lens = batch.seq_lens_cpu.sum().item()
|
self.stats.decode_sum_seq_lens = batch.seq_lens_cpu.sum().item()
|
||||||
self.stats.gen_throughput = self.last_gen_throughput
|
self.stats.gen_throughput = self.last_gen_throughput
|
||||||
self.stats.num_queue_reqs = QueueCount.from_reqs(
|
self.stats.num_queue_reqs = QueueCount.from_reqs(
|
||||||
@@ -887,14 +804,7 @@ class SchedulerMetricsMixin:
|
|||||||
return num_pending_tokens
|
return num_pending_tokens
|
||||||
|
|
||||||
def get_load(self: Scheduler, _: GetLoadReqInput = None) -> GetLoadReqOutput:
|
def get_load(self: Scheduler, _: GetLoadReqInput = None) -> GetLoadReqOutput:
|
||||||
if self.is_hybrid_swa:
|
num_tokens, _ = self.get_pool_stats().get_kv_token_stats()
|
||||||
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_pending_tokens = self._get_num_pending_tokens()
|
num_pending_tokens = self._get_num_pending_tokens()
|
||||||
|
|
||||||
# Tokens and request count in waiting queue, bootstrap queue, prealloc queue
|
# 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)
|
waiting_queues.append(self.disagg_decode_prealloc_queue.retracted_queue)
|
||||||
|
|
||||||
num_waiting_reqs = sum(len(queue) for queue in waiting_queues)
|
num_waiting_reqs = sum(len(queue) for queue in waiting_queues)
|
||||||
|
num_used_tokens, kv_token_usage = self.get_pool_stats().get_kv_token_stats()
|
||||||
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
|
|
||||||
)
|
|
||||||
|
|
||||||
memory = None
|
memory = None
|
||||||
if include_all or "memory" in include:
|
if include_all or "memory" in include:
|
||||||
@@ -1044,7 +941,7 @@ class SchedulerMetricsMixin:
|
|||||||
num_waiting_reqs=num_waiting_reqs,
|
num_waiting_reqs=num_waiting_reqs,
|
||||||
num_used_tokens=num_used_tokens,
|
num_used_tokens=num_used_tokens,
|
||||||
max_total_num_tokens=self.max_total_num_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),
|
gen_throughput=round(self.stats.gen_throughput, 2),
|
||||||
cache_hit_rate=round(self.stats.cache_hit_rate, 4),
|
cache_hit_rate=round(self.stats.cache_hit_rate, 4),
|
||||||
utilization=round(self.stats.utilization, 4),
|
utilization=round(self.stats.utilization, 4),
|
||||||
|
|||||||
@@ -9,6 +9,7 @@ maybe_stub_sgl_kernel()
|
|||||||
|
|
||||||
from sglang.srt.managers.io_struct import PauseGenerationReqInput
|
from sglang.srt.managers.io_struct import PauseGenerationReqInput
|
||||||
from sglang.srt.managers.scheduler import Scheduler
|
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")
|
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 = MagicMock()
|
||||||
scheduler.token_to_kv_pool_allocator.available_size.return_value = 1000
|
scheduler.token_to_kv_pool_allocator.available_size.return_value = 1000
|
||||||
scheduler.max_total_num_tokens = 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
|
return scheduler
|
||||||
|
|
||||||
def test_inplace_only_sets_flag(self):
|
def test_inplace_only_sets_flag(self):
|
||||||
|
|||||||
Reference in New Issue
Block a user