Add num_waiting_uncached_tokens load metric (#27174)
Co-authored-by: cctry <cctry@fb.com> Co-authored-by: Lianmin Zheng <lianminzheng@gmail.com>
This commit is contained in:
co-authored by
cctry
Lianmin Zheng
parent
687baf9471
commit
f6cd1a9822
@@ -2071,6 +2071,14 @@ class GetLoadsReqOutput(BaseReq):
|
||||
num_waiting_reqs: int = field(
|
||||
metadata={"metric": ("gauge", "Number of waiting requests")}
|
||||
)
|
||||
num_waiting_uncached_tokens: int = field(
|
||||
metadata={
|
||||
"metric": (
|
||||
"gauge",
|
||||
"Number of uncached input tokens waiting for prefill compute",
|
||||
)
|
||||
}
|
||||
)
|
||||
num_used_tokens: int = field(
|
||||
metadata={"metric": ("gauge", "Number of tokens in use")}
|
||||
)
|
||||
|
||||
@@ -144,6 +144,7 @@ CORE_METRIC_FIELDS = (
|
||||
"dp_rank",
|
||||
"num_running_reqs",
|
||||
"num_waiting_reqs",
|
||||
"num_waiting_uncached_tokens",
|
||||
"num_used_tokens",
|
||||
"num_total_tokens",
|
||||
"max_total_num_tokens",
|
||||
@@ -218,6 +219,7 @@ class LoadSnapshot(msgspec.Struct, omit_defaults=True):
|
||||
dp_rank: int = 0
|
||||
num_running_reqs: int = 0
|
||||
num_waiting_reqs: int = 0
|
||||
num_waiting_uncached_tokens: int = 0
|
||||
num_used_tokens: int = 0
|
||||
num_total_tokens: int = 0
|
||||
max_total_num_tokens: int = 0
|
||||
@@ -292,6 +294,7 @@ class LoadSnapshot(msgspec.Struct, omit_defaults=True):
|
||||
"dp_rank": self.dp_rank,
|
||||
"num_running_reqs": self.num_running_reqs,
|
||||
"num_waiting_reqs": self.num_waiting_reqs,
|
||||
"num_waiting_uncached_tokens": self.num_waiting_uncached_tokens,
|
||||
"num_used_tokens": self.num_used_tokens,
|
||||
"num_total_tokens": self.num_total_tokens,
|
||||
"max_total_num_tokens": self.max_total_num_tokens,
|
||||
|
||||
@@ -816,6 +816,11 @@ class Req(ReqDllmMixin):
|
||||
self.last_host_node: Any = None
|
||||
self.best_match_node: Any = None
|
||||
self.host_hit_length = 0
|
||||
# Total cached prefix length (on-device prefix_indices + host_hit_length),
|
||||
# capped at the max allowed prefix. Set during prefix matching at schedule
|
||||
# time and used to estimate uncached tokens / sort by longest prefix for
|
||||
# load reporting.
|
||||
self.num_matched_prefix_tokens = 0
|
||||
# Tokens loaded from storage backend (L3) during prefetch for this request
|
||||
self.storage_hit_length = 0
|
||||
# The node to lock until for swa radix tree lock ref
|
||||
@@ -1315,6 +1320,7 @@ class Req(ReqDllmMixin):
|
||||
self.indexer_topk = None
|
||||
self.last_node = None
|
||||
self.cache_protected_len = 0
|
||||
self.num_matched_prefix_tokens = 0
|
||||
self.swa_uuid_for_lock = None
|
||||
self.swa_prefix_lock_released = False
|
||||
self.extend_input_len = 0
|
||||
|
||||
@@ -50,7 +50,7 @@ from sglang.srt.mem_cache.hisparse_memory_pool import (
|
||||
)
|
||||
from sglang.srt.mem_cache.radix_cache import RadixCache, RadixKey, TreeNode
|
||||
from sglang.srt.mem_cache.swa_memory_pool import SWATokenToKVPoolAllocator
|
||||
from sglang.srt.server_args import ServerArgs
|
||||
from sglang.srt.server_args import ServerArgs, get_global_server_args
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from sglang.srt.mem_cache.allocator import BaseTokenToKVPoolAllocator
|
||||
@@ -115,6 +115,10 @@ def match_prefix_for_req(
|
||||
match_result.best_match_node,
|
||||
match_result.host_hit_length,
|
||||
)
|
||||
max_len = req._compute_max_prefix_len(len(token_ids))
|
||||
req.num_matched_prefix_tokens = min(
|
||||
len(req.prefix_indices) + req.host_hit_length, max_len
|
||||
)
|
||||
if match_result.mamba_branching_seqlen is not None:
|
||||
req.mamba_branching_seqlen = match_result.mamba_branching_seqlen
|
||||
if match_result.cache_protected_len is not None:
|
||||
@@ -162,6 +166,20 @@ class SchedulePolicy:
|
||||
def calc_priority(
|
||||
self, waiting_queue: List[Req], running_batch: Optional[ScheduleBatch] = None
|
||||
) -> None:
|
||||
policy = self._determine_active_policy(waiting_queue)
|
||||
|
||||
# Populate req.num_matched_prefix_tokens at schedule time. Cache-aware policies
|
||||
# set it in _compute_prefix_matches; do the same full match for
|
||||
# cache-agnostic policies when the radix supports it, so the load
|
||||
# snapshot has it. Skip on decode (never prefills).
|
||||
if (
|
||||
not isinstance(policy, CacheAwarePolicy)
|
||||
and self.tree_cache.supports_fast_match_prefix()
|
||||
and get_global_server_args().disaggregation_mode != "decode"
|
||||
):
|
||||
for r in waiting_queue:
|
||||
match_prefix_for_req(self.tree_cache, r)
|
||||
|
||||
if self.policy == CacheAgnosticPolicy.FCFS:
|
||||
if self.enable_priority_scheduling:
|
||||
SchedulePolicy._sort_by_priority_and_fcfs(
|
||||
@@ -169,8 +187,6 @@ class SchedulePolicy:
|
||||
)
|
||||
return
|
||||
|
||||
policy = self._determine_active_policy(waiting_queue)
|
||||
|
||||
if isinstance(policy, CacheAwarePolicy):
|
||||
temporary_deprioritized = self._compute_prefix_matches(
|
||||
waiting_queue, policy
|
||||
@@ -279,7 +295,7 @@ class SchedulePolicy:
|
||||
"""Sorts the waiting queue based on the longest prefix match."""
|
||||
waiting_queue.sort(
|
||||
key=lambda r: (
|
||||
-len(r.prefix_indices)
|
||||
-r.num_matched_prefix_tokens
|
||||
if r.rid not in temporary_deprioritized
|
||||
else float("inf")
|
||||
)
|
||||
|
||||
@@ -72,6 +72,19 @@ class SchedulerLoadInquirer:
|
||||
num_pending_tokens += req.seqlen - len(req.prefix_indices) - chunk_deduct
|
||||
return num_pending_tokens
|
||||
|
||||
def get_num_waiting_uncached_tokens(self) -> int:
|
||||
"""Get uncached input tokens waiting for prefill compute."""
|
||||
if self.disaggregation_mode == DisaggregationMode.DECODE:
|
||||
return 0
|
||||
num_tokens = 0
|
||||
for req in self.get_waiting_queue():
|
||||
# if match-in-waiting-queue disabled, this metric returns seq_lens
|
||||
num_tokens += max(0, req.seqlen - req.num_matched_prefix_tokens)
|
||||
cr = self.get_chunked_req()
|
||||
if cr is not None:
|
||||
num_tokens += max(0, cr.seqlen - len(cr.prefix_indices))
|
||||
return num_tokens
|
||||
|
||||
def get_loads(self, req: GetLoadsReqInput = None) -> GetLoadsReqOutput:
|
||||
"""
|
||||
Get comprehensive load metrics for /v1/loads endpoint.
|
||||
@@ -101,6 +114,7 @@ class SchedulerLoadInquirer:
|
||||
)
|
||||
|
||||
num_waiting_reqs = sum(len(queue) for queue in waiting_queues)
|
||||
num_waiting_uncached_tokens = self.get_num_waiting_uncached_tokens()
|
||||
num_used_tokens, kv_token_usage = (
|
||||
self.pool_stats_observer.get_pool_stats().get_kv_token_stats()
|
||||
)
|
||||
@@ -193,6 +207,7 @@ class SchedulerLoadInquirer:
|
||||
timestamp=time.time(),
|
||||
num_running_reqs=num_running_reqs,
|
||||
num_waiting_reqs=num_waiting_reqs,
|
||||
num_waiting_uncached_tokens=num_waiting_uncached_tokens,
|
||||
num_used_tokens=num_used_tokens,
|
||||
num_total_tokens=num_total_tokens,
|
||||
max_total_num_tokens=self.max_total_num_tokens,
|
||||
|
||||
@@ -237,6 +237,9 @@ class BasePrefixCache(ABC, PrefixCacheTrait):
|
||||
def match_prefix(self, params: MatchPrefixParams) -> MatchResult:
|
||||
pass
|
||||
|
||||
def supports_fast_match_prefix(self) -> bool:
|
||||
return False
|
||||
|
||||
@abstractmethod
|
||||
def cache_finished_req(self, req: Req, is_insert: bool = True, **kwargs):
|
||||
pass
|
||||
|
||||
@@ -36,9 +36,13 @@ class _StubReq:
|
||||
self.last_host_node = None
|
||||
self.best_match_node = None
|
||||
self.host_hit_length = None
|
||||
self.num_matched_prefix_tokens = 0
|
||||
self.mamba_branching_seqlen = None
|
||||
self.cache_protected_len = None
|
||||
|
||||
def _compute_max_prefix_len(self, input_len):
|
||||
return max(input_len - 1, 0)
|
||||
|
||||
|
||||
class TestZeroMatchResult(unittest.TestCase):
|
||||
def test_zero_replaces_indices_and_nodes(self):
|
||||
|
||||
Reference in New Issue
Block a user