diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index 714d38b48..6a13e3fec 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -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): diff --git a/python/sglang/srt/managers/scheduler_runtime_checker_mixin.py b/python/sglang/srt/managers/scheduler_runtime_checker_mixin.py index 0ac5563c9..8ec82c49a 100644 --- a/python/sglang/srt/managers/scheduler_runtime_checker_mixin.py +++ b/python/sglang/srt/managers/scheduler_runtime_checker_mixin.py @@ -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!" diff --git a/python/sglang/srt/managers/tokenizer_communicator_mixin.py b/python/sglang/srt/managers/tokenizer_communicator_mixin.py index f2ffa9909..d6892431e 100644 --- a/python/sglang/srt/managers/tokenizer_communicator_mixin.py +++ b/python/sglang/srt/managers/tokenizer_communicator_mixin.py @@ -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 diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index 8b9c0a9ff..38f1ab56f 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -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, diff --git a/test/registered/sessions/test_session_control.py b/test/registered/sessions/test_session_control.py index d74a39f67..af166c552 100644 --- a/test/registered/sessions/test_session_control.py +++ b/test/registered/sessions/test_session_control.py @@ -45,6 +45,7 @@ class TestSessionControl(unittest.TestCase): other_args=[ "--attention-backend", "flashinfer", + "--enable-streaming-session", ], ) diff --git a/test/registered/sessions/test_session_latency.py b/test/registered/sessions/test_session_latency.py index 7298641db..63e1e7746 100644 --- a/test/registered/sessions/test_session_latency.py +++ b/test/registered/sessions/test_session_latency.py @@ -309,7 +309,11 @@ class BenchSessionLatency(CustomTestCase): cls.model, cls.base_url, timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, - other_args=["--attention-backend", "flashinfer"], + other_args=[ + "--attention-backend", + "flashinfer", + "--enable-streaming-session", + ], ) cls.tokenizer = get_tokenizer(cls.model)