[Session] Gate streaming sessions with --enable-streaming-session and spec v2 guard (#19531)
This commit is contained in:
@@ -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):
|
||||
|
||||
@@ -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!"
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user