[Session] Gate streaming sessions with --enable-streaming-session and spec v2 guard (#19531)

This commit is contained in:
Liangsheng Yin
2026-02-27 18:14:55 -08:00
committed by GitHub
parent b5a8e4179e
commit e08ef06758
6 changed files with 62 additions and 17 deletions

View File

@@ -172,6 +172,7 @@ from sglang.srt.managers.utils import GenerationBatchResult, validate_input_leng
from sglang.srt.mem_cache.cache_init_params import CacheInitParams
from sglang.srt.mem_cache.common import release_kv_cache
from sglang.srt.mem_cache.radix_cache import RadixCache
from sglang.srt.mem_cache.session_aware_cache import SessionAwareCache
from sglang.srt.model_executor.forward_batch_info import ForwardMode, PPProxyTensors
from sglang.srt.multiplex.multiplexing_mixin import SchedulerMultiplexMixin
from sglang.srt.observability.req_time_stats import (
@@ -721,9 +722,9 @@ class Scheduler(
else:
self.tree_cache = RadixCache(params)
from sglang.srt.mem_cache.session_aware_cache import SessionAwareCache
if server_args.enable_streaming_session:
self.tree_cache = SessionAwareCache(self.tree_cache)
self.tree_cache = SessionAwareCache(self.tree_cache)
if (
server_args.disaggregation_mode == "decode"
@@ -2937,7 +2938,6 @@ class Scheduler(
return ExpertDistributionReqOutput()
def open_session(self, recv_req: OpenSessionReqInput):
# handle error
session_id = recv_req.session_id
if session_id in self.sessions:
logger.warning(f"session id {session_id} already exist, cannot open.")
@@ -2968,7 +2968,8 @@ class Scheduler(
req = next(iter(session.req_nodes.values())).req
if not req.finished():
req.session = None
self.tree_cache.release_session(session_id)
if isinstance(self.tree_cache, SessionAwareCache):
self.tree_cache.release_session(session_id)
del self.sessions[session_id]
def reap_timed_out_sessions(self):

View File

@@ -8,6 +8,7 @@ from typing import TYPE_CHECKING
from sglang.srt.disaggregation.utils import DisaggregationMode
from sglang.srt.environ import envs
from sglang.srt.managers.schedule_batch import ScheduleBatch
from sglang.srt.mem_cache.session_aware_cache import SessionAwareCache
from sglang.srt.utils.common import ceil_align, raise_error_or_warn
from sglang.srt.utils.request_logger import disable_request_logging
from sglang.srt.utils.watchdog import WatchdogRaw
@@ -19,6 +20,16 @@ logger = logging.getLogger(__name__)
class SchedulerRuntimeCheckerMixin:
def _session_held_tokens(self: Scheduler) -> int:
if isinstance(self.tree_cache, SessionAwareCache):
return self.tree_cache.session_held_tokens()
return 0
def _session_held_req_count(self: Scheduler) -> int:
if isinstance(self.tree_cache, SessionAwareCache):
return self.tree_cache.session_held_req_count()
return 0
def _get_token_info(self: Scheduler):
available_size = self.token_to_kv_pool_allocator.available_size()
evictable_size = self.tree_cache.evictable_size()
@@ -92,7 +103,7 @@ class SchedulerRuntimeCheckerMixin:
swa_available_size,
swa_evictable_size,
) = self._get_swa_token_info()
session_held = self.tree_cache.session_held_tokens()
session_held = self._session_held_tokens()
memory_leak = (full_num_used - session_held) != 0 or swa_num_used != 0
token_msg = (
f"{self.full_tokens_per_layer=}, {full_available_size=}, {full_evictable_size=}, {self.tree_cache.full_protected_size()=}\n"
@@ -111,7 +122,7 @@ class SchedulerRuntimeCheckerMixin:
mamba_available_size,
mamba_evictable_size,
) = self._get_mamba_token_info()
session_held = self.tree_cache.session_held_tokens()
session_held = self._session_held_tokens()
memory_leak = (
full_num_used != self.tree_cache.full_protected_size() + session_held
or mamba_num_used != self.tree_cache.mamba_protected_size()
@@ -152,7 +163,7 @@ class SchedulerRuntimeCheckerMixin:
def _check_radix_cache_memory(self: Scheduler):
_, _, available_size, evictable_size = self._get_token_info()
protected_size = self.tree_cache.protected_size()
session_held = self.tree_cache.session_held_tokens()
session_held = self._session_held_tokens()
memory_leak = (available_size + evictable_size) != (
self.max_total_num_tokens - protected_size - session_held
)
@@ -204,7 +215,7 @@ class SchedulerRuntimeCheckerMixin:
log_msg = f"[Mem Check (BUSY)] {available_size=}, {evictable_size=}, {protected_size=}, {uncached_size=}"
logger.info(log_msg)
session_held = self.tree_cache.session_held_tokens()
session_held = self._session_held_tokens()
total_tokens = (
available_size
+ evictable_size
@@ -224,7 +235,7 @@ class SchedulerRuntimeCheckerMixin:
else:
req_total_size = self.req_to_token_pool.size
session_req_count = self.tree_cache.session_held_req_count()
session_req_count = self._session_held_req_count()
if len(self.req_to_token_pool.free_slots) + session_req_count != req_total_size:
msg = (
"req_to_token_pool memory leak detected!"

View File

@@ -480,7 +480,7 @@ class TokenizerCommunicatorMixin:
return _Communicator.merge_results(results)
async def destroy_weights_update_group(
self,
self: TokenizerManager,
obj: DestroyWeightsUpdateGroupReqInput,
request: Optional[fastapi.Request] = None,
) -> Tuple[bool, str]:
@@ -523,7 +523,7 @@ class TokenizerCommunicatorMixin:
return success, message
async def init_weights_send_group_for_remote_instance(
self,
self: TokenizerManager,
obj: InitWeightsSendGroupForRemoteInstanceReqInput,
request: Optional[fastapi.Request] = None,
) -> Tuple[bool, str]:
@@ -538,7 +538,7 @@ class TokenizerCommunicatorMixin:
return result.success, result.message
async def send_weights_to_remote_instance(
self,
self: TokenizerManager,
obj: SendWeightsToRemoteInstanceReqInput,
request: Optional[fastapi.Request] = None,
) -> Tuple[bool, str]:
@@ -581,7 +581,7 @@ class TokenizerCommunicatorMixin:
return success, message
async def update_weights_from_ipc(
self,
self: TokenizerManager,
obj: UpdateWeightsFromIPCReqInput,
request: Optional[fastapi.Request] = None,
) -> Tuple[bool, str]:
@@ -907,9 +907,26 @@ class TokenizerCommunicatorMixin:
return results
async def open_session(
self, obj: OpenSessionReqInput, request: Optional[fastapi.Request] = None
self: TokenizerManager,
obj: OpenSessionReqInput,
request: Optional[fastapi.Request] = None,
):
self.auto_create_handle_loop()
if obj.streaming:
if not self.server_args.enable_streaming_session:
raise ValueError(
"Streaming sessions are disabled. "
"Please relaunch with --enable-streaming-session."
)
if (
self.server_args.speculative_algorithm is not None
and not self.server_args.disable_overlap_schedule
):
raise ValueError(
"Streaming sessions are incompatible with speculative decoding v2 "
"(overlap + speculative). Use --disable-overlap-schedule or "
"disable speculative decoding."
)
if obj.session_id is None:
obj.session_id = uuid.uuid4().hex
@@ -924,11 +941,15 @@ class TokenizerCommunicatorMixin:
return session_id
async def close_session(
self, obj: CloseSessionReqInput, request: Optional[fastapi.Request] = None
self: TokenizerManager,
obj: CloseSessionReqInput,
request: Optional[fastapi.Request] = None,
):
await self.send_to_scheduler.send_pyobj(obj)
def _update_weight_version_if_provided(self, weight_version: Optional[str]) -> None:
def _update_weight_version_if_provided(
self: TokenizerManager, weight_version: Optional[str]
) -> None:
"""Update weight version if provided."""
if weight_version is not None:
self.server_args.weight_version = weight_version

View File

@@ -354,6 +354,7 @@ class ServerArgs:
pp_async_batch_depth: int = 0
stream_interval: int = 1
stream_output: bool = False
enable_streaming_session: bool = False
random_seed: Optional[int] = None
constrained_json_whitespace_pattern: Optional[str] = None
constrained_json_disable_any_whitespace: bool = False
@@ -3387,6 +3388,12 @@ class ServerArgs:
action="store_true",
help="Whether to output as a sequence of disjoint segments.",
)
parser.add_argument(
"--enable-streaming-session",
action="store_true",
default=ServerArgs.enable_streaming_session,
help="Enable streaming session mode and SessionAwareCache wrapper.",
)
parser.add_argument(
"--random-seed",
type=int,