Add scripted-runtime KV-pool and lock-ref exhauster primitives (#27412)
This commit is contained in:
@@ -9,6 +9,12 @@ from sglang.test.scripted_runtime.context import (
|
||||
queries,
|
||||
radix,
|
||||
)
|
||||
from sglang.test.scripted_runtime.context.kv_pool_exhauster import (
|
||||
ScriptedKvPoolExhauster,
|
||||
)
|
||||
from sglang.test.scripted_runtime.context.lock_ref_exhauster import (
|
||||
ScriptedLockRefExhauster,
|
||||
)
|
||||
from sglang.test.scripted_runtime.context.req_starter import ScriptedContextReqStarter
|
||||
|
||||
if TYPE_CHECKING:
|
||||
@@ -41,6 +47,8 @@ class ScriptedContext:
|
||||
self._http_poster = http_poster
|
||||
|
||||
self._seen_rids: set[str] = set()
|
||||
self._kv_exhauster = ScriptedKvPoolExhauster(self.scheduler)
|
||||
self._lock_ref_exhauster = ScriptedLockRefExhauster(self.scheduler)
|
||||
self._req_starter = ScriptedContextReqStarter(self)
|
||||
|
||||
def start_req(
|
||||
@@ -93,6 +101,16 @@ class ScriptedContext:
|
||||
), "evict_radix currently supports only full eviction (prefix_tokens=None)"
|
||||
return lifecycle.flush_cache(self)
|
||||
|
||||
def exhaust_kv(self, *, leave_pages: int) -> None:
|
||||
return self._kv_exhauster.exhaust(leave_pages=leave_pages)
|
||||
|
||||
def exhaust_lock_refs(self, *, leave_refs: int) -> None:
|
||||
return self._lock_ref_exhauster.exhaust(leave_refs=leave_refs)
|
||||
|
||||
def _release_exhausted_pools(self) -> None:
|
||||
self._kv_exhauster.release()
|
||||
self._lock_ref_exhauster.release()
|
||||
|
||||
def get_all_node_hit_counts(self) -> Dict[int, int]:
|
||||
return radix.get_all_node_hit_counts(self)
|
||||
|
||||
|
||||
@@ -0,0 +1,34 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import TYPE_CHECKING, List
|
||||
|
||||
if TYPE_CHECKING:
|
||||
import torch
|
||||
|
||||
from sglang.srt.managers.scheduler import Scheduler
|
||||
|
||||
|
||||
class ScriptedKvPoolExhauster:
|
||||
|
||||
def __init__(self, scheduler: "Scheduler") -> None:
|
||||
self.scheduler = scheduler
|
||||
self._held: List["torch.Tensor"] = []
|
||||
|
||||
def exhaust(self, *, leave_pages: int) -> None:
|
||||
allocator = self.scheduler.token_to_kv_pool_allocator
|
||||
|
||||
leave_tokens = leave_pages * self.scheduler.page_size
|
||||
need = allocator.available_size() - leave_tokens
|
||||
if need <= 0:
|
||||
return
|
||||
|
||||
held = allocator.alloc(need)
|
||||
assert (
|
||||
held is not None
|
||||
), f"exhaust_kv: allocator could not grab {need} tokens to create pressure"
|
||||
self._held.append(held)
|
||||
|
||||
def release(self) -> None:
|
||||
for held in self._held:
|
||||
self.scheduler.token_to_kv_pool_allocator.free(held)
|
||||
self._held.clear()
|
||||
@@ -0,0 +1,49 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import TYPE_CHECKING, Any, List
|
||||
|
||||
from sglang.test.scripted_runtime.context.radix import _node_lock_ref
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from sglang.srt.managers.scheduler import Scheduler
|
||||
|
||||
|
||||
class ScriptedLockRefExhauster:
|
||||
|
||||
def __init__(self, scheduler: "Scheduler") -> None:
|
||||
self.scheduler = scheduler
|
||||
self._locked: List[Any] = []
|
||||
|
||||
def exhaust(self, *, leave_refs: int) -> None:
|
||||
tree_cache = self.scheduler.tree_cache
|
||||
if tree_cache.disable:
|
||||
return
|
||||
|
||||
while True:
|
||||
evictable = self._evictable_nodes()
|
||||
if len(evictable) <= leave_refs:
|
||||
return
|
||||
|
||||
target = evictable[0]
|
||||
tree_cache.inc_lock_ref(target)
|
||||
|
||||
newly_locked = [node for node in evictable if _node_lock_ref(node) > 0]
|
||||
if not newly_locked:
|
||||
return
|
||||
self._locked.append(target)
|
||||
|
||||
def release(self) -> None:
|
||||
tree_cache = self.scheduler.tree_cache
|
||||
for node in self._locked:
|
||||
tree_cache.dec_lock_ref(node)
|
||||
self._locked.clear()
|
||||
|
||||
def _evictable_nodes(self) -> List[Any]:
|
||||
evictable: List[Any] = []
|
||||
stack = list(self.scheduler.tree_cache.root_node.children.values())
|
||||
while stack:
|
||||
node = stack.pop()
|
||||
if _node_lock_ref(node) == 0:
|
||||
evictable.append(node)
|
||||
stack.extend(node.children.values())
|
||||
return evictable
|
||||
@@ -50,6 +50,7 @@ class ScriptedBatchRecord:
|
||||
def _reset_engine_state(ctx: ScriptedContext) -> Generator:
|
||||
scheduler = ctx.scheduler
|
||||
|
||||
ctx._release_exhausted_pools()
|
||||
ctx.abort_all()
|
||||
for _ in range(RESET_DRAIN_MAX_STEPS):
|
||||
yield
|
||||
|
||||
Reference in New Issue
Block a user