fix: streaming session race condition + some metrics (#21875)
Co-authored-by: Claude Opus 4.6 (1M context) <noreply@anthropic.com> Co-authored-by: hnyls2002 <lsyincs@gmail.com> Co-authored-by: Liangsheng Yin <hnyls2002@gmail.com>
This commit is contained in:
co-authored by
Claude Opus 4.6
hnyls2002
Liangsheng Yin
parent
37fc47c645
commit
c1ab68b45e
@@ -1848,8 +1848,11 @@ class Scheduler(
|
|||||||
self.stream_output([req], req.return_logprob)
|
self.stream_output([req], req.return_logprob)
|
||||||
return
|
return
|
||||||
|
|
||||||
elif session_id in self.session_controller:
|
elif (
|
||||||
# Session exists: create request from session
|
session_id in self.session_controller
|
||||||
|
and not self.session_controller.get(session_id).close_on_finish
|
||||||
|
):
|
||||||
|
# Session exists and is not closing: create request from session
|
||||||
session = self.session_controller.get(session_id)
|
session = self.session_controller.get(session_id)
|
||||||
req = session.create_req(
|
req = session.create_req(
|
||||||
recv_req,
|
recv_req,
|
||||||
@@ -1866,7 +1869,13 @@ class Scheduler(
|
|||||||
return
|
return
|
||||||
|
|
||||||
else:
|
else:
|
||||||
# Session ID provided but session not found
|
# Session not found, or session is closing
|
||||||
|
if session_id in self.session_controller:
|
||||||
|
error_msg = (
|
||||||
|
f"Invalid request: close was requested for session {session_id}"
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
error_msg = f"Invalid request: session id {session_id} does not exist"
|
||||||
req = Req(
|
req = Req(
|
||||||
recv_req.rid,
|
recv_req.rid,
|
||||||
recv_req.input_text,
|
recv_req.input_text,
|
||||||
@@ -1875,9 +1884,7 @@ class Scheduler(
|
|||||||
vocab_size=self.model_config.vocab_size,
|
vocab_size=self.model_config.vocab_size,
|
||||||
)
|
)
|
||||||
req.tokenizer = self.tokenizer
|
req.tokenizer = self.tokenizer
|
||||||
req.set_finish_with_abort(
|
req.set_finish_with_abort(error_msg)
|
||||||
f"Invalid request: session id {session_id} does not exist"
|
|
||||||
)
|
|
||||||
self.init_req_max_new_tokens(req)
|
self.init_req_max_new_tokens(req)
|
||||||
self._add_request_to_queue(req)
|
self._add_request_to_queue(req)
|
||||||
return
|
return
|
||||||
@@ -3461,7 +3468,10 @@ class Scheduler(
|
|||||||
return ExpertDistributionReqOutput()
|
return ExpertDistributionReqOutput()
|
||||||
|
|
||||||
def open_session(self, recv_req: OpenSessionReqInput):
|
def open_session(self, recv_req: OpenSessionReqInput):
|
||||||
return self.session_controller.open(recv_req)
|
output = self.session_controller.open(recv_req)
|
||||||
|
if self.pp_rank == 0 and self.tp_rank == 0 and self.attn_cp_rank == 0:
|
||||||
|
return output
|
||||||
|
return None
|
||||||
|
|
||||||
def close_session(self, recv_req: CloseSessionReqInput):
|
def close_session(self, recv_req: CloseSessionReqInput):
|
||||||
self.session_controller.close(recv_req)
|
self.session_controller.close(recv_req)
|
||||||
|
|||||||
@@ -117,6 +117,13 @@ class PoolStats:
|
|||||||
|
|
||||||
|
|
||||||
class SchedulerRuntimeCheckerMixin:
|
class SchedulerRuntimeCheckerMixin:
|
||||||
|
def _alive_streaming_session_count(self: Scheduler) -> int:
|
||||||
|
return sum(
|
||||||
|
1
|
||||||
|
for session in self.session_controller.sessions.values()
|
||||||
|
if session.streaming
|
||||||
|
)
|
||||||
|
|
||||||
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):
|
||||||
return self.tree_cache.session_held_tokens()
|
return self.tree_cache.session_held_tokens()
|
||||||
@@ -451,6 +458,8 @@ class SchedulerRuntimeCheckerMixin:
|
|||||||
return
|
return
|
||||||
|
|
||||||
self.get_pool_stats().update_scheduler_stats(self.stats)
|
self.get_pool_stats().update_scheduler_stats(self.stats)
|
||||||
|
self.stats.num_streaming_sessions = self._alive_streaming_session_count()
|
||||||
|
self.stats.streaming_session_held_tokens = self._session_held_tokens()
|
||||||
|
|
||||||
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(
|
||||||
|
|||||||
@@ -25,6 +25,7 @@ from sglang.srt.managers.io_struct import (
|
|||||||
)
|
)
|
||||||
from sglang.srt.managers.schedule_batch import FINISH_ABORT, Req
|
from sglang.srt.managers.schedule_batch import FINISH_ABORT, Req
|
||||||
from sglang.srt.mem_cache.session_aware_cache import SessionAwareCache
|
from sglang.srt.mem_cache.session_aware_cache import SessionAwareCache
|
||||||
|
from sglang.srt.utils.common import log_info_on_rank0
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from sglang.srt.mem_cache.base_prefix_cache import BasePrefixCache
|
from sglang.srt.mem_cache.base_prefix_cache import BasePrefixCache
|
||||||
@@ -92,6 +93,7 @@ class Session:
|
|||||||
self.timeout = timeout
|
self.timeout = timeout
|
||||||
self.last_active_time: float = time.monotonic()
|
self.last_active_time: float = time.monotonic()
|
||||||
self.req_nodes: Dict[str, SessionReqNode] = {}
|
self.req_nodes: Dict[str, SessionReqNode] = {}
|
||||||
|
self.close_on_finish: bool = False
|
||||||
|
|
||||||
def is_timed_out(self) -> bool:
|
def is_timed_out(self) -> bool:
|
||||||
if self.timeout is None:
|
if self.timeout is None:
|
||||||
@@ -275,6 +277,9 @@ class SessionController:
|
|||||||
streaming=bool(recv_req.streaming),
|
streaming=bool(recv_req.streaming),
|
||||||
timeout=recv_req.timeout,
|
timeout=recv_req.timeout,
|
||||||
)
|
)
|
||||||
|
log_info_on_rank0(
|
||||||
|
logger, f"Session opened: {session_id} (active={len(self.sessions)})"
|
||||||
|
)
|
||||||
return OpenSessionReqOutput(session_id, True)
|
return OpenSessionReqOutput(session_id, True)
|
||||||
|
|
||||||
def close(self, recv_req: CloseSessionReqInput):
|
def close(self, recv_req: CloseSessionReqInput):
|
||||||
@@ -286,10 +291,30 @@ class SessionController:
|
|||||||
|
|
||||||
def _close(self, session_id: str):
|
def _close(self, session_id: str):
|
||||||
session = self.sessions[session_id]
|
session = self.sessions[session_id]
|
||||||
|
req = None
|
||||||
|
has_unfinished_request = False
|
||||||
if session.streaming and session.req_nodes:
|
if session.streaming and session.req_nodes:
|
||||||
assert len(session.req_nodes) == 1
|
assert len(session.req_nodes) == 1
|
||||||
req = next(iter(session.req_nodes.values())).req
|
req = next(iter(session.req_nodes.values())).req
|
||||||
if not req.finished():
|
if not req.finished():
|
||||||
|
has_unfinished_request = True
|
||||||
|
|
||||||
|
if has_unfinished_request:
|
||||||
|
# An in-flight request is still decoding on this session's KV
|
||||||
|
# memory. Freeing now would corrupt the scheduler. Mark the
|
||||||
|
# session for deferred cleanup: the request keeps its session
|
||||||
|
# reference so cache_finished_req takes the streaming path,
|
||||||
|
# and we schedule release_session for after it completes.
|
||||||
|
session.close_on_finish = True
|
||||||
|
logger.info(
|
||||||
|
"Deferring session close for %s (unfinished request)",
|
||||||
|
session_id,
|
||||||
|
)
|
||||||
|
return
|
||||||
|
|
||||||
|
# No active request -- safe to release immediately.
|
||||||
|
if session.streaming and session.req_nodes:
|
||||||
|
req = next(iter(session.req_nodes.values())).req
|
||||||
req.session = None
|
req.session = None
|
||||||
|
|
||||||
# Release multimodal features held by session requests.
|
# Release multimodal features held by session requests.
|
||||||
@@ -304,20 +329,46 @@ class SessionController:
|
|||||||
node.req.multimodal_inputs = None
|
node.req.multimodal_inputs = None
|
||||||
|
|
||||||
if isinstance(self.tree_cache, SessionAwareCache):
|
if isinstance(self.tree_cache, SessionAwareCache):
|
||||||
self.tree_cache.release_session(session_id)
|
self.tree_cache.release_session(
|
||||||
|
session_id, req if session.streaming else None
|
||||||
|
)
|
||||||
del self.sessions[session_id]
|
del self.sessions[session_id]
|
||||||
|
log_info_on_rank0(
|
||||||
|
logger, f"Session closed: {session_id} (active={len(self.sessions)})"
|
||||||
|
)
|
||||||
|
|
||||||
def maybe_reap(self, now: float, interval: float = 1.0):
|
def maybe_reap(self, now: float, interval: float = 1.0):
|
||||||
# reap sessions every second
|
# reap sessions every second
|
||||||
if now - self._last_reap_time > interval:
|
if now - self._last_reap_time > interval:
|
||||||
self._last_reap_time = now
|
self._last_reap_time = now
|
||||||
|
|
||||||
|
# Finish deferred closes for sessions whose requests completed.
|
||||||
|
pending = [
|
||||||
|
sid
|
||||||
|
for sid, session in self.sessions.items()
|
||||||
|
if session.close_on_finish and self._all_requests_finished(session)
|
||||||
|
]
|
||||||
|
for sid in pending:
|
||||||
|
log_info_on_rank0(
|
||||||
|
logger, f"Deferred close ready for session {sid}, releasing."
|
||||||
|
)
|
||||||
|
# Reset close_on_finish so _close proceeds with the release.
|
||||||
|
self.sessions[sid].close_on_finish = False
|
||||||
|
self._close(sid)
|
||||||
|
|
||||||
timed_out = [
|
timed_out = [
|
||||||
sid for sid, session in self.sessions.items() if session.is_timed_out()
|
sid for sid, session in self.sessions.items() if session.is_timed_out()
|
||||||
]
|
]
|
||||||
for sid in timed_out:
|
for sid in timed_out:
|
||||||
logger.info(f"Session {sid} timed out, closing.")
|
log_info_on_rank0(logger, f"Session {sid} timed out, closing.")
|
||||||
self._close(sid)
|
self._close(sid)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _all_requests_finished(session: "Session") -> bool:
|
||||||
|
if not session.req_nodes:
|
||||||
|
return True
|
||||||
|
return all(node.req.finished() for node in session.req_nodes.values())
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def adjust_mm_offsets(recv_req: TokenizedGenerateReqInput, req: Req, image_inputs):
|
def adjust_mm_offsets(recv_req: TokenizedGenerateReqInput, req: Req, image_inputs):
|
||||||
# For session requests, adjust mm_inputs offsets by the prefix length.
|
# For session requests, adjust mm_inputs offsets by the prefix length.
|
||||||
|
|||||||
@@ -1090,12 +1090,14 @@ class TokenizerCommunicatorMixin:
|
|||||||
elif obj.session_id in self.session_futures:
|
elif obj.session_id in self.session_futures:
|
||||||
return None
|
return None
|
||||||
|
|
||||||
|
future = asyncio.Future()
|
||||||
|
self.session_futures[obj.session_id] = future
|
||||||
self.send_to_scheduler.send_pyobj(obj)
|
self.send_to_scheduler.send_pyobj(obj)
|
||||||
|
|
||||||
self.session_futures[obj.session_id] = asyncio.Future()
|
try:
|
||||||
session_id = await self.session_futures[obj.session_id]
|
return await future
|
||||||
del self.session_futures[obj.session_id]
|
finally:
|
||||||
return session_id
|
self.session_futures.pop(obj.session_id, None)
|
||||||
|
|
||||||
async def close_session(
|
async def close_session(
|
||||||
self: TokenizerManager,
|
self: TokenizerManager,
|
||||||
|
|||||||
@@ -2365,9 +2365,15 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerScoreMixin):
|
|||||||
self.send_to_scheduler.send_pyobj(ranks)
|
self.send_to_scheduler.send_pyobj(ranks)
|
||||||
|
|
||||||
def _handle_open_session_req_output(self, recv_obj):
|
def _handle_open_session_req_output(self, recv_obj):
|
||||||
self.session_futures[recv_obj.session_id].set_result(
|
future = self.session_futures.get(recv_obj.session_id)
|
||||||
recv_obj.session_id if recv_obj.success else None
|
if future is None:
|
||||||
|
logger.warning(
|
||||||
|
"Open session response arrived after waiter cleanup: %s",
|
||||||
|
recv_obj.session_id,
|
||||||
)
|
)
|
||||||
|
return
|
||||||
|
if not future.done():
|
||||||
|
future.set_result(recv_obj.session_id if recv_obj.success else None)
|
||||||
|
|
||||||
def _handle_update_weights_from_disk_req_output(self, recv_obj):
|
def _handle_update_weights_from_disk_req_output(self, recv_obj):
|
||||||
if self.server_args.dp_size == 1:
|
if self.server_args.dp_size == 1:
|
||||||
|
|||||||
@@ -9,6 +9,7 @@ import triton.language as tl
|
|||||||
|
|
||||||
from sglang.srt.mem_cache.base_prefix_cache import BasePrefixCache, EvictParams
|
from sglang.srt.mem_cache.base_prefix_cache import BasePrefixCache, EvictParams
|
||||||
from sglang.srt.mem_cache.memory_pool import HybridReqToTokenPool, ReqToTokenPool
|
from sglang.srt.mem_cache.memory_pool import HybridReqToTokenPool, ReqToTokenPool
|
||||||
|
from sglang.srt.mem_cache.session_aware_cache import SessionAwareCache
|
||||||
from sglang.srt.mem_cache.swa_memory_pool import SWATokenToKVPoolAllocator
|
from sglang.srt.mem_cache.swa_memory_pool import SWATokenToKVPoolAllocator
|
||||||
from sglang.srt.server_args import get_global_server_args
|
from sglang.srt.server_args import get_global_server_args
|
||||||
from sglang.srt.utils import support_triton
|
from sglang.srt.utils import support_triton
|
||||||
@@ -476,13 +477,52 @@ def release_kv_cache(req: Req, tree_cache: BasePrefixCache, is_insert: bool = Tr
|
|||||||
req.mamba_pool_idx = None
|
req.mamba_pool_idx = None
|
||||||
return
|
return
|
||||||
|
|
||||||
|
# Streaming sessions transfer req_pool ownership into SessionSlot objects.
|
||||||
|
# Trim any speculative tail before that transfer, otherwise later turns
|
||||||
|
# restore only the committed prefix and can strand unreachable KV pages.
|
||||||
|
#
|
||||||
|
# Aborted streaming-session requests (e.g. input too long) skip the
|
||||||
|
# streaming path entirely. match_prefix did not restore the slot's KV
|
||||||
|
# state, so the request has a fresh pool slot that should be freed by
|
||||||
|
# cache_finished_req below (which also sets req_pool_idx = None).
|
||||||
|
from sglang.srt.managers.schedule_batch import FINISH_ABORT
|
||||||
|
|
||||||
|
is_streaming_session = (
|
||||||
|
isinstance(tree_cache, SessionAwareCache)
|
||||||
|
and getattr(req, "session", None) is not None
|
||||||
|
and req.session.streaming
|
||||||
|
)
|
||||||
|
is_aborted_streaming = is_streaming_session and isinstance(
|
||||||
|
getattr(req, "finished_reason", None), FINISH_ABORT
|
||||||
|
)
|
||||||
|
if is_streaming_session and not is_aborted_streaming:
|
||||||
|
start_p, end_p = req.pop_overallocated_kv_cache()
|
||||||
|
page_size = get_global_server_args().page_size
|
||||||
|
if page_size > 1:
|
||||||
|
start_p = ceil_align(start_p, page_size)
|
||||||
|
if start_p < end_p:
|
||||||
|
indices_to_free = tree_cache.req_to_token_pool.req_to_token[
|
||||||
|
req.req_pool_idx
|
||||||
|
][start_p:end_p]
|
||||||
|
tree_cache.token_to_kv_pool_allocator.free(indices_to_free)
|
||||||
|
req.kv_allocated_len = req.kv_committed_len
|
||||||
|
|
||||||
tree_cache.cache_finished_req(req, is_insert=is_insert)
|
tree_cache.cache_finished_req(req, is_insert=is_insert)
|
||||||
|
|
||||||
# FIXME: SessionAwareCache.cache_finished_req sets req_pool_idx = None to
|
# SessionAwareCache.cache_finished_req sets req_pool_idx = None to transfer
|
||||||
# transfer KV ownership to the SessionSlot, so we skip the remaining
|
# KV ownership to the SessionSlot, so the remaining cleanup is skipped.
|
||||||
# cleanup (overalloc free + pool slot free). This means over-allocated
|
# Streaming-session specific overalloc trimming must therefore happen
|
||||||
# tokens from speculative decoding are NOT freed between turns.
|
# before cache_finished_req above.
|
||||||
if req.req_pool_idx is None:
|
if req.req_pool_idx is None:
|
||||||
|
if is_streaming_session:
|
||||||
|
# The request no longer owns any KV once SessionAwareCache either
|
||||||
|
# transfers it into the session slot or frees it on abort. Mark
|
||||||
|
# both bookkeeping flags so busy-time memory checks do not keep
|
||||||
|
# counting this finished request as uncached KV.
|
||||||
|
if not req.kv_committed_freed:
|
||||||
|
req.pop_committed_kv_cache()
|
||||||
|
if not req.kv_overallocated_freed:
|
||||||
|
req.pop_overallocated_kv_cache()
|
||||||
return
|
return
|
||||||
|
|
||||||
start_p, end_p = req.pop_overallocated_kv_cache()
|
start_p, end_p = req.pop_overallocated_kv_cache()
|
||||||
|
|||||||
@@ -1,5 +1,6 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import logging
|
||||||
from dataclasses import dataclass, field
|
from dataclasses import dataclass, field
|
||||||
from typing import TYPE_CHECKING, Any, Dict, Optional
|
from typing import TYPE_CHECKING, Any, Dict, Optional
|
||||||
|
|
||||||
@@ -22,6 +23,9 @@ if TYPE_CHECKING:
|
|||||||
from sglang.srt.managers.schedule_batch import Req
|
from sglang.srt.managers.schedule_batch import Req
|
||||||
|
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
class _VirtualNode:
|
class _VirtualNode:
|
||||||
"""Sentinel node for streaming session requests.
|
"""Sentinel node for streaming session requests.
|
||||||
|
|
||||||
@@ -187,6 +191,16 @@ class SessionAwareCache(BasePrefixCache):
|
|||||||
if slot is None or slot.req_pool_idx is None:
|
if slot is None or slot.req_pool_idx is None:
|
||||||
return self.inner.match_prefix(params)
|
return self.inner.match_prefix(params)
|
||||||
|
|
||||||
|
# If the request is destined for abort (e.g. input too long),
|
||||||
|
# do NOT restore the slot's KV state. set_finish_with_abort
|
||||||
|
# truncates origin_input_ids to [0], so alloc_for_extend would
|
||||||
|
# overwrite the slot's req_to_token row with a 1-token prefix,
|
||||||
|
# destroying the session's accumulated KV mapping. By skipping
|
||||||
|
# restore, the request gets a fresh pool slot from alloc_for_extend
|
||||||
|
# and the session slot remains untouched.
|
||||||
|
if req.to_finish is not None:
|
||||||
|
return self.inner.match_prefix(params)
|
||||||
|
|
||||||
slot.restore_to_req(req)
|
slot.restore_to_req(req)
|
||||||
|
|
||||||
# logprob_start_len is already forced to -1 for streaming sessions
|
# logprob_start_len is already forced to -1 for streaming sessions
|
||||||
@@ -208,13 +222,56 @@ class SessionAwareCache(BasePrefixCache):
|
|||||||
if not _is_streaming(req):
|
if not _is_streaming(req):
|
||||||
return self.inner.cache_finished_req(req, is_insert=is_insert, **kwargs)
|
return self.inner.cache_finished_req(req, is_insert=is_insert, **kwargs)
|
||||||
|
|
||||||
|
from sglang.srt.managers.schedule_batch import FINISH_ABORT
|
||||||
|
|
||||||
session_id = req.session.session_id
|
session_id = req.session.session_id
|
||||||
slot = self.slots.get(session_id)
|
slot = self.slots.get(session_id)
|
||||||
is_first = slot is None
|
is_first = slot is None
|
||||||
|
|
||||||
|
# When an aborted streaming-session request was scheduled (e.g.
|
||||||
|
# input too long), match_prefix skipped restore_to_req so the
|
||||||
|
# request got a fresh pool slot from alloc_for_extend. Don't
|
||||||
|
# overwrite the session slot -- free the transient KV and pool slot.
|
||||||
|
if not is_first and isinstance(req.finished_reason, FINISH_ABORT):
|
||||||
|
if req.req_pool_idx is not None:
|
||||||
|
# Free all KV pages allocated for this aborted request.
|
||||||
|
end = req.kv_allocated_len
|
||||||
|
if end > 0:
|
||||||
|
kv_indices = self.req_to_token_pool.req_to_token[
|
||||||
|
req.req_pool_idx, :end
|
||||||
|
]
|
||||||
|
self.token_to_kv_pool_allocator.free(kv_indices)
|
||||||
|
self.req_to_token_pool.free_slots.append(req.req_pool_idx)
|
||||||
|
req.req_pool_idx = None
|
||||||
|
return
|
||||||
|
|
||||||
if is_first:
|
if is_first:
|
||||||
slot = SessionSlot()
|
slot = SessionSlot()
|
||||||
self.slots[session_id] = slot
|
self.slots[session_id] = slot
|
||||||
|
|
||||||
|
# If the session's KV is shrinking (e.g. client sent a shorter
|
||||||
|
# prompt after an abort), free the orphaned tail pages before
|
||||||
|
# save_from_req overwrites the slot's committed length.
|
||||||
|
# Never free tree-protected tokens — those are managed by the tree.
|
||||||
|
if (
|
||||||
|
not is_first
|
||||||
|
and slot.is_holding_kv
|
||||||
|
and req.kv_committed_len < slot.kv_committed_len
|
||||||
|
):
|
||||||
|
old_end = slot.kv_allocated_len
|
||||||
|
new_end = req.kv_committed_len
|
||||||
|
if self.page_size > 1:
|
||||||
|
new_end = ceil_align(new_end, self.page_size)
|
||||||
|
new_end = max(new_end, slot.cache_protected_len)
|
||||||
|
if new_end < old_end:
|
||||||
|
kv_indices = self.req_to_token_pool.req_to_token[
|
||||||
|
slot.req_pool_idx, new_end:old_end
|
||||||
|
]
|
||||||
|
self.token_to_kv_pool_allocator.free(kv_indices)
|
||||||
|
slot.cache_protected_len = min(
|
||||||
|
slot.cache_protected_len, req.kv_committed_len
|
||||||
|
)
|
||||||
|
|
||||||
slot.save_from_req(req, is_first=is_first)
|
slot.save_from_req(req, is_first=is_first)
|
||||||
|
|
||||||
def cache_unfinished_req(self, req: Req, **kwargs):
|
def cache_unfinished_req(self, req: Req, **kwargs):
|
||||||
@@ -251,23 +308,99 @@ class SessionAwareCache(BasePrefixCache):
|
|||||||
|
|
||||||
# -- Session lifecycle --
|
# -- Session lifecycle --
|
||||||
|
|
||||||
def release_session(self, session_id: str):
|
def _resolve_release_state(
|
||||||
|
self, slot: SessionSlot, req: Optional[Req]
|
||||||
|
) -> tuple[int, Any]:
|
||||||
|
"""Resolve the currently tree-owned prefix for a session slot.
|
||||||
|
|
||||||
|
A long-lived session can outlive radix-tree splits caused by unrelated
|
||||||
|
traffic. In that case, the saved `last_node` may no longer represent the
|
||||||
|
full protected prefix even though the slot's req_to_token row still
|
||||||
|
contains tree-owned indices at the front. Re-match the current request
|
||||||
|
text, then intersect the returned tree indices with the slot's row so
|
||||||
|
release uses the prefix that is still actually backed by the tree.
|
||||||
|
"""
|
||||||
|
protected_len = slot.cache_protected_len
|
||||||
|
lock_node = slot.last_node
|
||||||
|
|
||||||
|
# TODO: re-match logic disabled — match_prefix has side effects
|
||||||
|
# (splits) that disturb tree accounting. Directly using
|
||||||
|
# slot.last_node + cache_protected_len is safe after split analysis.
|
||||||
|
return protected_len, lock_node
|
||||||
|
|
||||||
|
if (
|
||||||
|
req is None
|
||||||
|
or not slot.is_holding_kv
|
||||||
|
or slot.req_pool_idx is None
|
||||||
|
or protected_len <= 0
|
||||||
|
):
|
||||||
|
return protected_len, lock_node
|
||||||
|
|
||||||
|
from sglang.srt.mem_cache.radix_cache import RadixKey
|
||||||
|
|
||||||
|
token_ids = (req.origin_input_ids + req.output_ids)[: slot.kv_committed_len]
|
||||||
|
if not token_ids:
|
||||||
|
return 0, None
|
||||||
|
|
||||||
|
match = self.inner.match_prefix(
|
||||||
|
MatchPrefixParams(
|
||||||
|
key=RadixKey(token_ids=token_ids, extra_key=req.extra_key),
|
||||||
|
req=None,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
if len(match.device_indices) == 0:
|
||||||
|
return 0, None
|
||||||
|
|
||||||
|
max_protected_len = min(len(match.device_indices), protected_len)
|
||||||
|
row_indices = self.req_to_token_pool.req_to_token[
|
||||||
|
slot.req_pool_idx, :max_protected_len
|
||||||
|
].to(dtype=torch.int64)
|
||||||
|
match_indices = match.device_indices[:max_protected_len]
|
||||||
|
mismatches = (match_indices != row_indices).nonzero(as_tuple=False)
|
||||||
|
if mismatches.numel() == 0 and max_protected_len == len(match.device_indices):
|
||||||
|
common_len = max_protected_len
|
||||||
|
return common_len, match.last_device_node
|
||||||
|
|
||||||
|
common_len = (
|
||||||
|
int(mismatches[0].item()) if mismatches.numel() > 0 else max_protected_len
|
||||||
|
)
|
||||||
|
if self.page_size > 1:
|
||||||
|
common_len = (common_len // self.page_size) * self.page_size
|
||||||
|
if common_len <= 0:
|
||||||
|
return 0, None
|
||||||
|
|
||||||
|
rematch = self.inner.match_prefix(
|
||||||
|
MatchPrefixParams(
|
||||||
|
key=RadixKey(token_ids=token_ids[:common_len], extra_key=req.extra_key),
|
||||||
|
req=None,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
return len(rematch.device_indices), rematch.last_device_node
|
||||||
|
|
||||||
|
def release_session(self, session_id: str, req: Optional[Req] = None):
|
||||||
"""Release all KV resources held by a streaming session."""
|
"""Release all KV resources held by a streaming session."""
|
||||||
slot = self.slots.pop(session_id, None)
|
slot = self.slots.pop(session_id, None)
|
||||||
if slot is None:
|
if slot is None:
|
||||||
return
|
return
|
||||||
|
protected_len, lock_node = self._resolve_release_state(slot, req)
|
||||||
|
tokens_freed = (
|
||||||
|
max(0, slot.kv_allocated_len - protected_len) if slot.is_holding_kv else 0
|
||||||
|
)
|
||||||
|
logger.info(
|
||||||
|
"Session KV released: %s (%d tokens freed)", session_id, tokens_freed
|
||||||
|
)
|
||||||
|
|
||||||
if slot.last_node is not None:
|
if lock_node is not None:
|
||||||
if slot.swa_uuid_for_lock is not None:
|
if slot.swa_uuid_for_lock is not None:
|
||||||
self.inner.dec_lock_ref(
|
self.inner.dec_lock_ref(
|
||||||
slot.last_node,
|
lock_node,
|
||||||
DecLockRefParams(swa_uuid_for_lock=slot.swa_uuid_for_lock),
|
DecLockRefParams(swa_uuid_for_lock=slot.swa_uuid_for_lock),
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
self.inner.dec_lock_ref(slot.last_node)
|
self.inner.dec_lock_ref(lock_node)
|
||||||
|
|
||||||
if slot.is_holding_kv:
|
if slot.is_holding_kv:
|
||||||
start = slot.cache_protected_len
|
start = protected_len
|
||||||
end = slot.kv_allocated_len
|
end = slot.kv_allocated_len
|
||||||
if start < end:
|
if start < end:
|
||||||
kv_indices = self.req_to_token_pool.req_to_token[
|
kv_indices = self.req_to_token_pool.req_to_token[
|
||||||
|
|||||||
@@ -132,6 +132,10 @@ class SchedulerStats:
|
|||||||
hicache_host_used_tokens: int = 0
|
hicache_host_used_tokens: int = 0
|
||||||
hicache_host_total_tokens: int = 0
|
hicache_host_total_tokens: int = 0
|
||||||
|
|
||||||
|
# Streaming session metrics
|
||||||
|
num_streaming_sessions: int = 0
|
||||||
|
streaming_session_held_tokens: int = 0
|
||||||
|
|
||||||
# Routing key metrics
|
# Routing key metrics
|
||||||
num_unique_running_routing_keys: int = 0
|
num_unique_running_routing_keys: int = 0
|
||||||
routing_key_running_req_counts: List[int] = field(default_factory=list)
|
routing_key_running_req_counts: List[int] = field(default_factory=list)
|
||||||
@@ -176,6 +180,7 @@ class SchedulerMetricsCollector:
|
|||||||
labels: Dict[str, str],
|
labels: Dict[str, str],
|
||||||
enable_lora: bool = False,
|
enable_lora: bool = False,
|
||||||
enable_hierarchical_cache: bool = False,
|
enable_hierarchical_cache: bool = False,
|
||||||
|
enable_streaming_session: bool = False,
|
||||||
server_args: Optional["ServerArgs"] = None,
|
server_args: Optional["ServerArgs"] = None,
|
||||||
) -> None:
|
) -> None:
|
||||||
# We need to import prometheus_client after setting the env variable `PROMETHEUS_MULTIPROC_DIR`
|
# We need to import prometheus_client after setting the env variable `PROMETHEUS_MULTIPROC_DIR`
|
||||||
@@ -184,6 +189,7 @@ class SchedulerMetricsCollector:
|
|||||||
self.labels = labels
|
self.labels = labels
|
||||||
self.enable_lora = enable_lora
|
self.enable_lora = enable_lora
|
||||||
self.enable_hierarchical_cache = enable_hierarchical_cache
|
self.enable_hierarchical_cache = enable_hierarchical_cache
|
||||||
|
self.enable_streaming_session = enable_streaming_session
|
||||||
self.last_log_time = time.perf_counter()
|
self.last_log_time = time.perf_counter()
|
||||||
self._known_priorities: Set[int] = set()
|
self._known_priorities: Set[int] = set()
|
||||||
|
|
||||||
@@ -654,6 +660,21 @@ class SchedulerMetricsCollector:
|
|||||||
multiprocess_mode="mostrecent",
|
multiprocess_mode="mostrecent",
|
||||||
)
|
)
|
||||||
|
|
||||||
|
# Streaming session metrics (only created when streaming sessions are enabled)
|
||||||
|
if self.enable_streaming_session:
|
||||||
|
self.num_streaming_sessions = Gauge(
|
||||||
|
name="sglang:num_streaming_sessions",
|
||||||
|
documentation="The number of active streaming sessions.",
|
||||||
|
labelnames=labels.keys(),
|
||||||
|
multiprocess_mode="mostrecent",
|
||||||
|
)
|
||||||
|
self.streaming_session_held_tokens = Gauge(
|
||||||
|
name="sglang:streaming_session_held_tokens",
|
||||||
|
documentation="The number of KV tokens currently held by streaming session slots.",
|
||||||
|
labelnames=labels.keys(),
|
||||||
|
multiprocess_mode="mostrecent",
|
||||||
|
)
|
||||||
|
|
||||||
self.num_unique_running_routing_keys = Gauge(
|
self.num_unique_running_routing_keys = Gauge(
|
||||||
name="sglang:num_unique_running_routing_keys",
|
name="sglang:num_unique_running_routing_keys",
|
||||||
documentation="Number of unique routing keys in running batch.",
|
documentation="Number of unique routing keys in running batch.",
|
||||||
@@ -1049,6 +1070,13 @@ class SchedulerMetricsCollector:
|
|||||||
self.hicache_host_total_tokens, stats.hicache_host_total_tokens
|
self.hicache_host_total_tokens, stats.hicache_host_total_tokens
|
||||||
)
|
)
|
||||||
|
|
||||||
|
# Streaming session metrics (only logged if streaming sessions are enabled)
|
||||||
|
if self.enable_streaming_session:
|
||||||
|
self._log_gauge(self.num_streaming_sessions, stats.num_streaming_sessions)
|
||||||
|
self._log_gauge(
|
||||||
|
self.streaming_session_held_tokens, stats.streaming_session_held_tokens
|
||||||
|
)
|
||||||
|
|
||||||
self._log_gauge(
|
self._log_gauge(
|
||||||
self.num_unique_running_routing_keys, stats.num_unique_running_routing_keys
|
self.num_unique_running_routing_keys, stats.num_unique_running_routing_keys
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -148,6 +148,7 @@ class SchedulerMetricsMixin:
|
|||||||
labels=labels,
|
labels=labels,
|
||||||
enable_lora=self.enable_lora,
|
enable_lora=self.enable_lora,
|
||||||
enable_hierarchical_cache=self.enable_hierarchical_cache,
|
enable_hierarchical_cache=self.enable_hierarchical_cache,
|
||||||
|
enable_streaming_session=self.server_args.enable_streaming_session,
|
||||||
server_args=self.server_args,
|
server_args=self.server_args,
|
||||||
)
|
)
|
||||||
self.enable_mfu_metrics = bool(self.server_args.enable_mfu_metrics)
|
self.enable_mfu_metrics = bool(self.server_args.enable_mfu_metrics)
|
||||||
@@ -602,6 +603,8 @@ class SchedulerMetricsMixin:
|
|||||||
self.stats.cache_hit_rate = cache_hit_rate
|
self.stats.cache_hit_rate = cache_hit_rate
|
||||||
|
|
||||||
self.stats.max_total_num_tokens = self.max_total_num_tokens
|
self.stats.max_total_num_tokens = self.max_total_num_tokens
|
||||||
|
self.stats.num_streaming_sessions = self._alive_streaming_session_count()
|
||||||
|
self.stats.streaming_session_held_tokens = self._session_held_tokens()
|
||||||
|
|
||||||
# Speculative decoding
|
# Speculative decoding
|
||||||
self.stats.spec_accept_rate = spec_accept_rate
|
self.stats.spec_accept_rate = spec_accept_rate
|
||||||
|
|||||||
@@ -10,8 +10,12 @@ Usage:
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
import asyncio
|
import asyncio
|
||||||
|
import json
|
||||||
|
import os
|
||||||
|
import tempfile
|
||||||
import time
|
import time
|
||||||
import unittest
|
import unittest
|
||||||
|
from pathlib import Path
|
||||||
from typing import Any, Optional
|
from typing import Any, Optional
|
||||||
|
|
||||||
import aiohttp
|
import aiohttp
|
||||||
@@ -66,6 +70,21 @@ LEAK_FILLER = (
|
|||||||
"We promptly judged antique ivory buckles for the next prize. "
|
"We promptly judged antique ivory buckles for the next prize. "
|
||||||
) * 20
|
) * 20
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Abort-heavy chunked prefill leak repro constants
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
ABORT_REPRO_CONTEXT_LEN = 512
|
||||||
|
ABORT_REPRO_PAGE_SIZE = 16
|
||||||
|
ABORT_REPRO_GEN_LEN = 8
|
||||||
|
ABORT_REPRO_SESSIONS = 4
|
||||||
|
ABORT_REPRO_WARMUP_TURNS = 2
|
||||||
|
ABORT_REPRO_ROUNDS = 8
|
||||||
|
ABORT_REPRO_STREAM_TOKENS = 150
|
||||||
|
ABORT_REPRO_ABORT_TOKENS = 320
|
||||||
|
ABORT_REPRO_NON_STREAMING_TOKENS = 96
|
||||||
|
ABORT_REPRO_CHUNKED_PREFILL_SIZE = 128
|
||||||
|
|
||||||
|
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
# Logprob leak helpers
|
# Logprob leak helpers
|
||||||
@@ -206,6 +225,198 @@ async def _leak_run_all(base_url: str, tokenizer: Any) -> None:
|
|||||||
assert resp.status == 200
|
assert resp.status == 200
|
||||||
|
|
||||||
|
|
||||||
|
def _make_token_sized_ids(
|
||||||
|
tokenizer: Any, prefix: str, min_tokens: int, max_tokens: Optional[int] = None
|
||||||
|
) -> list[int]:
|
||||||
|
text = prefix
|
||||||
|
chunk = " pack quartz wizard sphinx zebra fox " * 16
|
||||||
|
token_ids = tokenizer.encode(text)
|
||||||
|
while len(token_ids) < min_tokens:
|
||||||
|
text += chunk
|
||||||
|
token_ids = tokenizer.encode(text)
|
||||||
|
if max_tokens is not None:
|
||||||
|
token_ids = token_ids[:max_tokens]
|
||||||
|
return token_ids
|
||||||
|
|
||||||
|
|
||||||
|
async def _abort_repro_generate(
|
||||||
|
base_url: str,
|
||||||
|
session: aiohttp.ClientSession,
|
||||||
|
input_ids: list[int],
|
||||||
|
max_new_tokens: int,
|
||||||
|
session_params: Optional[dict[str, Any]] = None,
|
||||||
|
expect_abort: bool = False,
|
||||||
|
) -> Optional[dict[str, Any]]:
|
||||||
|
payload: dict[str, Any] = {
|
||||||
|
"input_ids": input_ids,
|
||||||
|
"sampling_params": {
|
||||||
|
"temperature": 0,
|
||||||
|
"max_new_tokens": max_new_tokens,
|
||||||
|
"no_stop_trim": True,
|
||||||
|
"skip_special_tokens": False,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
if session_params:
|
||||||
|
payload["session_params"] = session_params
|
||||||
|
|
||||||
|
async with session.post(base_url + "/generate", json=payload) as resp:
|
||||||
|
text = await resp.text()
|
||||||
|
if expect_abort:
|
||||||
|
if resp.status == 200:
|
||||||
|
data = json.loads(text)
|
||||||
|
finish_reason = data.get("meta_info", {}).get("finish_reason", {})
|
||||||
|
assert finish_reason.get("type") == "abort", text
|
||||||
|
assert "maximum allowed length" in finish_reason.get(
|
||||||
|
"message", ""
|
||||||
|
), text
|
||||||
|
return data
|
||||||
|
assert resp.status == 400, text
|
||||||
|
assert "maximum allowed length" in text, text
|
||||||
|
return None
|
||||||
|
|
||||||
|
assert resp.status == 200, text
|
||||||
|
data = json.loads(text)
|
||||||
|
finish_reason = data.get("meta_info", {}).get("finish_reason", {})
|
||||||
|
assert finish_reason.get("type") != "abort", text
|
||||||
|
return data
|
||||||
|
|
||||||
|
|
||||||
|
def _read_tail(path: str, num_lines: int = 120) -> str:
|
||||||
|
if not path or not os.path.exists(path):
|
||||||
|
return "<missing log>"
|
||||||
|
lines = Path(path).read_text(errors="replace").splitlines()
|
||||||
|
return "\n".join(lines[-num_lines:])
|
||||||
|
|
||||||
|
|
||||||
|
async def _abort_repro_run_all(base_url: str, tokenizer: Any) -> None:
|
||||||
|
timeout = aiohttp.ClientTimeout(total=300)
|
||||||
|
async with aiohttp.ClientSession(timeout=timeout) as http:
|
||||||
|
session_ids = []
|
||||||
|
for _ in range(ABORT_REPRO_SESSIONS):
|
||||||
|
async with http.post(
|
||||||
|
base_url + "/open_session",
|
||||||
|
json={"capacity_of_str_len": 50000, "streaming": True},
|
||||||
|
) as resp:
|
||||||
|
assert resp.status == 200, await resp.text()
|
||||||
|
session_ids.append(await resp.json())
|
||||||
|
|
||||||
|
try:
|
||||||
|
for warmup_turn in range(ABORT_REPRO_WARMUP_TURNS):
|
||||||
|
warmup_tasks = []
|
||||||
|
for session_idx, session_id in enumerate(session_ids):
|
||||||
|
input_ids = _make_token_sized_ids(
|
||||||
|
tokenizer,
|
||||||
|
prefix=f"[warmup={warmup_turn} session={session_idx}]",
|
||||||
|
min_tokens=ABORT_REPRO_STREAM_TOKENS,
|
||||||
|
max_tokens=ABORT_REPRO_STREAM_TOKENS + 8,
|
||||||
|
)
|
||||||
|
warmup_tasks.append(
|
||||||
|
_abort_repro_generate(
|
||||||
|
base_url,
|
||||||
|
http,
|
||||||
|
input_ids,
|
||||||
|
ABORT_REPRO_GEN_LEN,
|
||||||
|
session_params={"id": session_id, "rid": None},
|
||||||
|
)
|
||||||
|
)
|
||||||
|
await asyncio.gather(*warmup_tasks)
|
||||||
|
|
||||||
|
for round_idx in range(ABORT_REPRO_ROUNDS):
|
||||||
|
mixed_tasks = []
|
||||||
|
for session_idx, session_id in enumerate(session_ids):
|
||||||
|
input_ids = _make_token_sized_ids(
|
||||||
|
tokenizer,
|
||||||
|
prefix=f"[round={round_idx} ok session={session_idx}]",
|
||||||
|
min_tokens=ABORT_REPRO_STREAM_TOKENS,
|
||||||
|
max_tokens=ABORT_REPRO_STREAM_TOKENS + 8,
|
||||||
|
)
|
||||||
|
mixed_tasks.append(
|
||||||
|
_abort_repro_generate(
|
||||||
|
base_url,
|
||||||
|
http,
|
||||||
|
input_ids,
|
||||||
|
ABORT_REPRO_GEN_LEN,
|
||||||
|
session_params={"id": session_id, "rid": None},
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
for ns_idx in range(2):
|
||||||
|
input_ids = _make_token_sized_ids(
|
||||||
|
tokenizer,
|
||||||
|
prefix=f"[round={round_idx} ns={ns_idx}]",
|
||||||
|
min_tokens=ABORT_REPRO_NON_STREAMING_TOKENS,
|
||||||
|
max_tokens=ABORT_REPRO_NON_STREAMING_TOKENS + 8,
|
||||||
|
)
|
||||||
|
mixed_tasks.append(
|
||||||
|
_abort_repro_generate(
|
||||||
|
base_url,
|
||||||
|
http,
|
||||||
|
input_ids,
|
||||||
|
ABORT_REPRO_GEN_LEN,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
await asyncio.gather(*mixed_tasks)
|
||||||
|
|
||||||
|
abort_tasks = []
|
||||||
|
for session_idx, session_id in enumerate(session_ids):
|
||||||
|
input_ids = _make_token_sized_ids(
|
||||||
|
tokenizer,
|
||||||
|
prefix=f"[round={round_idx} abort session={session_idx}]",
|
||||||
|
min_tokens=ABORT_REPRO_ABORT_TOKENS,
|
||||||
|
)
|
||||||
|
abort_tasks.append(
|
||||||
|
_abort_repro_generate(
|
||||||
|
base_url,
|
||||||
|
http,
|
||||||
|
input_ids,
|
||||||
|
ABORT_REPRO_GEN_LEN,
|
||||||
|
session_params={"id": session_id, "rid": None},
|
||||||
|
expect_abort=True,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
await asyncio.gather(*abort_tasks)
|
||||||
|
|
||||||
|
recovery_tasks = []
|
||||||
|
for session_idx, session_id in enumerate(session_ids):
|
||||||
|
input_ids = _make_token_sized_ids(
|
||||||
|
tokenizer,
|
||||||
|
prefix=f"[round={round_idx} recover session={session_idx}]",
|
||||||
|
min_tokens=ABORT_REPRO_NON_STREAMING_TOKENS,
|
||||||
|
max_tokens=ABORT_REPRO_NON_STREAMING_TOKENS + 8,
|
||||||
|
)
|
||||||
|
recovery_tasks.append(
|
||||||
|
_abort_repro_generate(
|
||||||
|
base_url,
|
||||||
|
http,
|
||||||
|
input_ids,
|
||||||
|
ABORT_REPRO_GEN_LEN,
|
||||||
|
session_params={"id": session_id, "rid": None},
|
||||||
|
)
|
||||||
|
)
|
||||||
|
recovery_results = await asyncio.gather(*recovery_tasks)
|
||||||
|
for result in recovery_results:
|
||||||
|
assert result is not None
|
||||||
|
assert result["meta_info"]["cached_tokens"] > 0, result
|
||||||
|
|
||||||
|
health = requests.get(base_url + "/health", timeout=10)
|
||||||
|
if health.status_code != 200:
|
||||||
|
raise RuntimeError(
|
||||||
|
f"server unhealthy after round={round_idx}: "
|
||||||
|
f"{health.status_code} {health.text}"
|
||||||
|
)
|
||||||
|
finally:
|
||||||
|
for session_id in session_ids:
|
||||||
|
async with http.post(
|
||||||
|
base_url + "/close_session", json={"session_id": session_id}
|
||||||
|
) as resp:
|
||||||
|
assert resp.status == 200, await resp.text()
|
||||||
|
|
||||||
|
|
||||||
|
# ===================================================================
|
||||||
|
# Test class
|
||||||
|
# ===================================================================
|
||||||
|
|
||||||
|
|
||||||
class TestStreamingSession(CustomTestCase):
|
class TestStreamingSession(CustomTestCase):
|
||||||
@classmethod
|
@classmethod
|
||||||
def setUpClass(cls):
|
def setUpClass(cls):
|
||||||
@@ -467,5 +678,94 @@ class TestStreamingSessionRetractMixedChunk(TestStreamingSession):
|
|||||||
kill_process_tree(cls.process.pid)
|
kill_process_tree(cls.process.pid)
|
||||||
|
|
||||||
|
|
||||||
|
class TestStreamingSessionAbortLeakRepro(CustomTestCase):
|
||||||
|
@classmethod
|
||||||
|
def setUpClass(cls):
|
||||||
|
cls.model = DEFAULT_SMALL_MODEL_NAME_FOR_TEST
|
||||||
|
cls.base_url = DEFAULT_URL_FOR_TEST
|
||||||
|
cls.stdout = tempfile.NamedTemporaryFile(
|
||||||
|
prefix="streaming-session-abort-repro.",
|
||||||
|
suffix=".stdout.log",
|
||||||
|
delete=False,
|
||||||
|
mode="w+",
|
||||||
|
encoding="utf-8",
|
||||||
|
)
|
||||||
|
cls.stderr = tempfile.NamedTemporaryFile(
|
||||||
|
prefix="streaming-session-abort-repro.",
|
||||||
|
suffix=".stderr.log",
|
||||||
|
delete=False,
|
||||||
|
mode="w+",
|
||||||
|
encoding="utf-8",
|
||||||
|
)
|
||||||
|
cls.process = popen_launch_server(
|
||||||
|
cls.model,
|
||||||
|
cls.base_url,
|
||||||
|
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
|
||||||
|
other_args=[
|
||||||
|
"--enable-streaming-session",
|
||||||
|
"--chunked-prefill-size",
|
||||||
|
str(ABORT_REPRO_CHUNKED_PREFILL_SIZE),
|
||||||
|
"--context-length",
|
||||||
|
str(ABORT_REPRO_CONTEXT_LEN),
|
||||||
|
"--page-size",
|
||||||
|
str(ABORT_REPRO_PAGE_SIZE),
|
||||||
|
"--max-running-requests",
|
||||||
|
"32",
|
||||||
|
"--log-level",
|
||||||
|
"info",
|
||||||
|
],
|
||||||
|
return_stdout_stderr=(cls.stdout, cls.stderr),
|
||||||
|
)
|
||||||
|
cls.tokenizer = get_tokenizer(cls.model)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def tearDownClass(cls):
|
||||||
|
kill_process_tree(cls.process.pid)
|
||||||
|
for handle in (cls.stdout, cls.stderr):
|
||||||
|
path = handle.name
|
||||||
|
handle.close()
|
||||||
|
if os.path.exists(path):
|
||||||
|
os.remove(path)
|
||||||
|
|
||||||
|
def test_abort_heavy_chunked_prefill_does_not_leak(self) -> None:
|
||||||
|
requests.post(self.base_url + "/flush_cache")
|
||||||
|
|
||||||
|
asyncio.run(_abort_repro_run_all(self.base_url, self.tokenizer))
|
||||||
|
|
||||||
|
for i in range(3):
|
||||||
|
ids = self.tokenizer.encode(f"Post-session cleanup request {i}.")
|
||||||
|
response = requests.post(
|
||||||
|
self.base_url + "/generate",
|
||||||
|
json={
|
||||||
|
"input_ids": ids,
|
||||||
|
"sampling_params": {"temperature": 0, "max_new_tokens": 4},
|
||||||
|
},
|
||||||
|
timeout=30,
|
||||||
|
)
|
||||||
|
self.assertEqual(response.status_code, 200, response.text)
|
||||||
|
|
||||||
|
time.sleep(5)
|
||||||
|
self.assertIsNone(
|
||||||
|
self.process.poll(),
|
||||||
|
"Server crashed during abort-heavy streaming session repro.\n"
|
||||||
|
f"---- stderr tail ----\n{_read_tail(self.stderr.name)}",
|
||||||
|
)
|
||||||
|
|
||||||
|
health = requests.get(self.base_url + "/health", timeout=10)
|
||||||
|
self.assertEqual(
|
||||||
|
health.status_code,
|
||||||
|
200,
|
||||||
|
"Server unhealthy after abort-heavy streaming session cleanup.\n"
|
||||||
|
f"---- stderr tail ----\n{_read_tail(self.stderr.name)}",
|
||||||
|
)
|
||||||
|
|
||||||
|
stderr_tail = _read_tail(self.stderr.name)
|
||||||
|
self.assertNotIn(
|
||||||
|
"token_to_kv_pool_allocator memory leak detected",
|
||||||
|
stderr_tail,
|
||||||
|
stderr_tail,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
unittest.main()
|
unittest.main()
|
||||||
|
|||||||
@@ -0,0 +1,358 @@
|
|||||||
|
from types import SimpleNamespace
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from sglang.srt.managers.schedule_batch import FINISH_ABORT
|
||||||
|
from sglang.srt.mem_cache.base_prefix_cache import MatchResult
|
||||||
|
from sglang.srt.mem_cache.common import release_kv_cache
|
||||||
|
from sglang.srt.mem_cache.session_aware_cache import SessionAwareCache, SessionSlot
|
||||||
|
from sglang.test.ci.ci_register import register_cpu_ci
|
||||||
|
|
||||||
|
register_cpu_ci(est_time=8, suite="stage-a-test-cpu")
|
||||||
|
|
||||||
|
|
||||||
|
class _FakeAllocator:
|
||||||
|
def __init__(self):
|
||||||
|
self.freed = []
|
||||||
|
|
||||||
|
def free(self, free_index: torch.Tensor):
|
||||||
|
self.freed.append(free_index.clone())
|
||||||
|
|
||||||
|
|
||||||
|
class _FakeInnerCache:
|
||||||
|
def __init__(self, req_to_token_pool, allocator, page_size, match_results=None):
|
||||||
|
self.req_to_token_pool = req_to_token_pool
|
||||||
|
self.token_to_kv_pool_allocator = allocator
|
||||||
|
self.page_size = page_size
|
||||||
|
self.match_results = list(match_results or [])
|
||||||
|
self.dec_lock_ref_calls = []
|
||||||
|
|
||||||
|
def cache_finished_req(self, *args, **kwargs):
|
||||||
|
raise AssertionError("Streaming requests should not delegate to inner cache")
|
||||||
|
|
||||||
|
def match_prefix(self, *args, **kwargs):
|
||||||
|
if not self.match_results:
|
||||||
|
raise AssertionError("Unexpected match_prefix call")
|
||||||
|
return self.match_results.pop(0)
|
||||||
|
|
||||||
|
def dec_lock_ref(self, node, *args, **kwargs):
|
||||||
|
self.dec_lock_ref_calls.append(node)
|
||||||
|
|
||||||
|
def supports_mamba(self):
|
||||||
|
return False
|
||||||
|
|
||||||
|
def sanity_check(self):
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
class _FakeReq:
|
||||||
|
def __init__(
|
||||||
|
self, session_id: str, req_pool_idx: int, committed: int, allocated: int
|
||||||
|
):
|
||||||
|
self.session = SimpleNamespace(session_id=session_id, streaming=True)
|
||||||
|
self.req_pool_idx = req_pool_idx
|
||||||
|
self.kv_committed_len = committed
|
||||||
|
self.kv_allocated_len = allocated
|
||||||
|
self.kv_committed_freed = False
|
||||||
|
self.kv_overallocated_freed = False
|
||||||
|
self.origin_input_ids = list(range(committed))
|
||||||
|
self.output_ids = []
|
||||||
|
self.extra_key = None
|
||||||
|
self.swa_evicted_seqlen = 0
|
||||||
|
self.last_node = None
|
||||||
|
self.cache_protected_len = 0
|
||||||
|
self.swa_uuid_for_lock = None
|
||||||
|
self.mamba_pool_idx = None
|
||||||
|
self.mamba_ping_pong_track_buffer = None
|
||||||
|
self.mamba_next_track_idx = None
|
||||||
|
self.mamba_last_track_seqlen = None
|
||||||
|
self.mamba_branching_seqlen = None
|
||||||
|
self.pop_overallocated_calls = 0
|
||||||
|
self.to_finish = None
|
||||||
|
self.finished_reason = None
|
||||||
|
|
||||||
|
def pop_committed_kv_cache(self):
|
||||||
|
assert not self.kv_committed_freed
|
||||||
|
self.kv_committed_freed = True
|
||||||
|
return self.kv_committed_len
|
||||||
|
|
||||||
|
def pop_overallocated_kv_cache(self):
|
||||||
|
assert not self.kv_overallocated_freed
|
||||||
|
self.pop_overallocated_calls += 1
|
||||||
|
self.kv_overallocated_freed = True
|
||||||
|
return self.kv_committed_len, self.kv_allocated_len
|
||||||
|
|
||||||
|
|
||||||
|
def test_streaming_release_kv_cache_trims_overallocated_tail(monkeypatch):
|
||||||
|
page_size = 16
|
||||||
|
req_to_token = torch.arange(128, dtype=torch.int32).reshape(1, 128)
|
||||||
|
req_to_token_pool = SimpleNamespace(req_to_token=req_to_token, free_slots=[])
|
||||||
|
allocator = _FakeAllocator()
|
||||||
|
tree_cache = SessionAwareCache(
|
||||||
|
_FakeInnerCache(req_to_token_pool, allocator, page_size)
|
||||||
|
)
|
||||||
|
req = _FakeReq("session-a", req_pool_idx=0, committed=17, allocated=40)
|
||||||
|
|
||||||
|
monkeypatch.setattr(
|
||||||
|
"sglang.srt.mem_cache.common.get_global_server_args",
|
||||||
|
lambda: SimpleNamespace(page_size=page_size, speculative_algorithm="eagle"),
|
||||||
|
)
|
||||||
|
|
||||||
|
release_kv_cache(req, tree_cache)
|
||||||
|
|
||||||
|
slot = tree_cache.slots["session-a"]
|
||||||
|
assert req.pop_overallocated_calls == 1
|
||||||
|
assert req.kv_committed_freed is True
|
||||||
|
assert req.kv_overallocated_freed is True
|
||||||
|
assert req.req_pool_idx is None
|
||||||
|
assert slot.kv_committed_len == 17
|
||||||
|
assert slot.kv_allocated_len == 17
|
||||||
|
assert len(allocator.freed) == 1
|
||||||
|
assert allocator.freed[0].tolist() == list(range(32, 40))
|
||||||
|
|
||||||
|
|
||||||
|
def test_release_session_recomputes_current_tree_owned_prefix():
|
||||||
|
page_size = 16
|
||||||
|
req_to_token = torch.arange(128, dtype=torch.int32).reshape(1, 128)
|
||||||
|
req_to_token_pool = SimpleNamespace(req_to_token=req_to_token, free_slots=[])
|
||||||
|
allocator = _FakeAllocator()
|
||||||
|
|
||||||
|
full_match = MatchResult(
|
||||||
|
device_indices=torch.tensor(list(range(16)) + list(range(64, 96))),
|
||||||
|
last_device_node="stale-expanded",
|
||||||
|
last_host_node="stale-expanded",
|
||||||
|
)
|
||||||
|
protected_match = MatchResult(
|
||||||
|
device_indices=torch.tensor(list(range(16))),
|
||||||
|
last_device_node="current-protected",
|
||||||
|
last_host_node="current-protected",
|
||||||
|
)
|
||||||
|
inner = _FakeInnerCache(
|
||||||
|
req_to_token_pool,
|
||||||
|
allocator,
|
||||||
|
page_size,
|
||||||
|
match_results=[full_match, protected_match],
|
||||||
|
)
|
||||||
|
tree_cache = SessionAwareCache(inner)
|
||||||
|
|
||||||
|
tree_cache.slots["session-a"] = SessionSlot(
|
||||||
|
req_pool_idx=0,
|
||||||
|
kv_committed_len=48,
|
||||||
|
kv_allocated_len=48,
|
||||||
|
last_node="outdated-node",
|
||||||
|
cache_protected_len=32,
|
||||||
|
)
|
||||||
|
req = _FakeReq("session-a", req_pool_idx=0, committed=48, allocated=48)
|
||||||
|
|
||||||
|
tree_cache.release_session("session-a", req)
|
||||||
|
|
||||||
|
assert inner.dec_lock_ref_calls == ["current-protected"]
|
||||||
|
assert req_to_token_pool.free_slots == [0]
|
||||||
|
assert len(allocator.freed) == 1
|
||||||
|
assert allocator.freed[0].tolist() == list(range(16, 48))
|
||||||
|
|
||||||
|
|
||||||
|
def test_release_session_never_grows_tree_owned_prefix():
|
||||||
|
page_size = 16
|
||||||
|
req_to_token = torch.arange(128, dtype=torch.int32).reshape(1, 128)
|
||||||
|
req_to_token_pool = SimpleNamespace(req_to_token=req_to_token, free_slots=[])
|
||||||
|
allocator = _FakeAllocator()
|
||||||
|
|
||||||
|
overmatched = MatchResult(
|
||||||
|
device_indices=torch.tensor(list(range(48))),
|
||||||
|
last_device_node="overmatched-node",
|
||||||
|
last_host_node="overmatched-node",
|
||||||
|
)
|
||||||
|
capped_match = MatchResult(
|
||||||
|
device_indices=torch.tensor(list(range(16))),
|
||||||
|
last_device_node="original-lock-node",
|
||||||
|
last_host_node="original-lock-node",
|
||||||
|
)
|
||||||
|
inner = _FakeInnerCache(
|
||||||
|
req_to_token_pool,
|
||||||
|
allocator,
|
||||||
|
page_size,
|
||||||
|
match_results=[overmatched, capped_match],
|
||||||
|
)
|
||||||
|
tree_cache = SessionAwareCache(inner)
|
||||||
|
|
||||||
|
tree_cache.slots["session-a"] = SessionSlot(
|
||||||
|
req_pool_idx=0,
|
||||||
|
kv_committed_len=48,
|
||||||
|
kv_allocated_len=48,
|
||||||
|
last_node="outdated-node",
|
||||||
|
cache_protected_len=16,
|
||||||
|
)
|
||||||
|
req = _FakeReq("session-a", req_pool_idx=0, committed=48, allocated=48)
|
||||||
|
|
||||||
|
tree_cache.release_session("session-a", req)
|
||||||
|
|
||||||
|
assert inner.dec_lock_ref_calls == ["original-lock-node"]
|
||||||
|
assert req_to_token_pool.free_slots == [0]
|
||||||
|
assert len(allocator.freed) == 1
|
||||||
|
assert allocator.freed[0].tolist() == list(range(16, 48))
|
||||||
|
|
||||||
|
|
||||||
|
def test_match_prefix_abort_does_not_restore_live_session_slot():
|
||||||
|
req_to_token = torch.arange(256, dtype=torch.int32).reshape(2, 128)
|
||||||
|
req_to_token_pool = SimpleNamespace(req_to_token=req_to_token, free_slots=[])
|
||||||
|
allocator = _FakeAllocator()
|
||||||
|
inner = _FakeInnerCache(
|
||||||
|
req_to_token_pool,
|
||||||
|
allocator,
|
||||||
|
page_size=16,
|
||||||
|
match_results=[
|
||||||
|
MatchResult(
|
||||||
|
device_indices=torch.tensor([], dtype=torch.int64),
|
||||||
|
last_device_node=None,
|
||||||
|
last_host_node=None,
|
||||||
|
)
|
||||||
|
],
|
||||||
|
)
|
||||||
|
tree_cache = SessionAwareCache(inner)
|
||||||
|
tree_cache.slots["session-a"] = SessionSlot(
|
||||||
|
req_pool_idx=0,
|
||||||
|
kv_committed_len=48,
|
||||||
|
kv_allocated_len=48,
|
||||||
|
cache_protected_len=16,
|
||||||
|
)
|
||||||
|
|
||||||
|
req = _FakeReq("session-a", req_pool_idx=1, committed=1, allocated=1)
|
||||||
|
req.to_finish = FINISH_ABORT("too long")
|
||||||
|
|
||||||
|
result = tree_cache.match_prefix(
|
||||||
|
SimpleNamespace(
|
||||||
|
req=req,
|
||||||
|
key=SimpleNamespace(token_ids=list(range(64))),
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
slot = tree_cache.slots["session-a"]
|
||||||
|
assert req.req_pool_idx == 1
|
||||||
|
assert req.kv_committed_len == 1
|
||||||
|
assert req.kv_allocated_len == 1
|
||||||
|
assert slot.req_pool_idx == 0
|
||||||
|
assert slot.kv_committed_len == 48
|
||||||
|
assert slot.kv_allocated_len == 48
|
||||||
|
assert len(result.device_indices) == 0
|
||||||
|
|
||||||
|
|
||||||
|
def test_aborted_streaming_turn_preserves_slot_and_accounting(monkeypatch):
|
||||||
|
page_size = 16
|
||||||
|
req_to_token = torch.arange(256, dtype=torch.int32).reshape(2, 128)
|
||||||
|
req_to_token_pool = SimpleNamespace(req_to_token=req_to_token, free_slots=[])
|
||||||
|
allocator = _FakeAllocator()
|
||||||
|
tree_cache = SessionAwareCache(
|
||||||
|
_FakeInnerCache(req_to_token_pool, allocator, page_size)
|
||||||
|
)
|
||||||
|
tree_cache.slots["session-a"] = SessionSlot(
|
||||||
|
req_pool_idx=0,
|
||||||
|
kv_committed_len=48,
|
||||||
|
kv_allocated_len=48,
|
||||||
|
cache_protected_len=16,
|
||||||
|
swa_evicted_seqlen=8,
|
||||||
|
last_node="lock-node",
|
||||||
|
)
|
||||||
|
|
||||||
|
req = _FakeReq("session-a", req_pool_idx=1, committed=5, allocated=23)
|
||||||
|
req.finished_reason = FINISH_ABORT("too long")
|
||||||
|
|
||||||
|
monkeypatch.setattr(
|
||||||
|
"sglang.srt.mem_cache.common.get_global_server_args",
|
||||||
|
lambda: SimpleNamespace(page_size=page_size, speculative_algorithm="eagle"),
|
||||||
|
)
|
||||||
|
|
||||||
|
release_kv_cache(req, tree_cache)
|
||||||
|
|
||||||
|
slot = tree_cache.slots["session-a"]
|
||||||
|
assert slot.req_pool_idx == 0
|
||||||
|
assert slot.kv_committed_len == 48
|
||||||
|
assert slot.kv_allocated_len == 48
|
||||||
|
assert req.kv_committed_freed is True
|
||||||
|
assert req.kv_overallocated_freed is True
|
||||||
|
assert req.req_pool_idx is None
|
||||||
|
assert req.pop_overallocated_calls == 1
|
||||||
|
assert tree_cache.session_held_tokens() == 32
|
||||||
|
assert tree_cache.session_held_full_tokens() == 32
|
||||||
|
assert tree_cache.session_held_swa_tokens() == 32
|
||||||
|
assert tree_cache.session_held_req_count() == 1
|
||||||
|
assert req_to_token_pool.free_slots == [1]
|
||||||
|
assert len(allocator.freed) == 1
|
||||||
|
assert allocator.freed[0].tolist() == list(range(128, 151))
|
||||||
|
|
||||||
|
tree_cache.release_session("session-a")
|
||||||
|
|
||||||
|
assert tree_cache.session_held_tokens() == 0
|
||||||
|
assert tree_cache.session_held_swa_tokens() == 0
|
||||||
|
assert tree_cache.session_held_req_count() == 0
|
||||||
|
assert req_to_token_pool.free_slots == [1, 0]
|
||||||
|
assert len(allocator.freed) == 2
|
||||||
|
assert allocator.freed[1].tolist() == list(range(16, 48))
|
||||||
|
|
||||||
|
|
||||||
|
def test_session_shrink_frees_orphaned_tail():
|
||||||
|
"""When a session's KV shrinks (client retried with shorter prompt),
|
||||||
|
the orphaned tail pages must be freed before save_from_req overwrites
|
||||||
|
the slot."""
|
||||||
|
page_size = 16
|
||||||
|
pool_size = 256
|
||||||
|
req_to_token = torch.arange(pool_size, dtype=torch.int32).reshape(1, pool_size)
|
||||||
|
req_to_token_pool = SimpleNamespace(req_to_token=req_to_token, free_slots=[])
|
||||||
|
allocator = _FakeAllocator()
|
||||||
|
inner = _FakeInnerCache(req_to_token_pool, allocator, page_size)
|
||||||
|
tree_cache = SessionAwareCache(inner)
|
||||||
|
|
||||||
|
# Session slot has 128 tokens committed
|
||||||
|
tree_cache.slots["session-a"] = SessionSlot(
|
||||||
|
req_pool_idx=0,
|
||||||
|
kv_committed_len=128,
|
||||||
|
kv_allocated_len=128,
|
||||||
|
last_node="lock-node",
|
||||||
|
cache_protected_len=16,
|
||||||
|
)
|
||||||
|
|
||||||
|
# New request finished with only 48 tokens (client truncated)
|
||||||
|
req = _FakeReq("session-a", req_pool_idx=0, committed=48, allocated=48)
|
||||||
|
|
||||||
|
tree_cache.cache_finished_req(req)
|
||||||
|
|
||||||
|
slot = tree_cache.slots["session-a"]
|
||||||
|
# Slot should now reflect the shrunk state
|
||||||
|
assert slot.kv_committed_len == 48
|
||||||
|
assert slot.kv_allocated_len == 48
|
||||||
|
# The tail [48:128] should have been freed (page-aligned: [48:128])
|
||||||
|
assert len(allocator.freed) == 1
|
||||||
|
assert allocator.freed[0].tolist() == list(range(48, 128))
|
||||||
|
|
||||||
|
|
||||||
|
def test_session_shrink_page_aligns_free_start():
|
||||||
|
"""The shrink free should page-align the start to avoid freeing
|
||||||
|
tokens that are still part of the new committed prefix."""
|
||||||
|
page_size = 16
|
||||||
|
pool_size = 256
|
||||||
|
req_to_token = torch.arange(pool_size, dtype=torch.int32).reshape(1, pool_size)
|
||||||
|
req_to_token_pool = SimpleNamespace(req_to_token=req_to_token, free_slots=[])
|
||||||
|
allocator = _FakeAllocator()
|
||||||
|
inner = _FakeInnerCache(req_to_token_pool, allocator, page_size)
|
||||||
|
tree_cache = SessionAwareCache(inner)
|
||||||
|
|
||||||
|
# Session slot has 128 tokens
|
||||||
|
tree_cache.slots["session-a"] = SessionSlot(
|
||||||
|
req_pool_idx=0,
|
||||||
|
kv_committed_len=128,
|
||||||
|
kv_allocated_len=128,
|
||||||
|
last_node="lock-node",
|
||||||
|
cache_protected_len=16,
|
||||||
|
)
|
||||||
|
|
||||||
|
# New request committed 50 tokens (not page-aligned)
|
||||||
|
req = _FakeReq("session-a", req_pool_idx=0, committed=50, allocated=50)
|
||||||
|
|
||||||
|
tree_cache.cache_finished_req(req)
|
||||||
|
|
||||||
|
slot = tree_cache.slots["session-a"]
|
||||||
|
assert slot.kv_committed_len == 50
|
||||||
|
# Free start should be ceil_align(50, 16) = 64, not 50
|
||||||
|
# So freed range is [64:128]
|
||||||
|
assert len(allocator.freed) == 1
|
||||||
|
assert allocator.freed[0].tolist() == list(range(64, 128))
|
||||||
Reference in New Issue
Block a user