[Session] Extract SessionController and clean up session logic in Scheduler (#19547)
This commit is contained in:
@@ -22,7 +22,7 @@ import time
|
|||||||
from collections import deque
|
from collections import deque
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
from http import HTTPStatus
|
from http import HTTPStatus
|
||||||
from typing import Any, Deque, Dict, List, Optional, Tuple, Union
|
from typing import Any, Deque, List, Optional, Tuple, Union
|
||||||
|
|
||||||
import psutil
|
import psutil
|
||||||
import setproctitle
|
import setproctitle
|
||||||
@@ -112,7 +112,6 @@ from sglang.srt.managers.io_struct import (
|
|||||||
LoadLoRAAdapterReqInput,
|
LoadLoRAAdapterReqInput,
|
||||||
LoadLoRAAdapterReqOutput,
|
LoadLoRAAdapterReqOutput,
|
||||||
OpenSessionReqInput,
|
OpenSessionReqInput,
|
||||||
OpenSessionReqOutput,
|
|
||||||
PauseGenerationReqInput,
|
PauseGenerationReqInput,
|
||||||
ProfileReq,
|
ProfileReq,
|
||||||
ReleaseMemoryOccupationReqInput,
|
ReleaseMemoryOccupationReqInput,
|
||||||
@@ -167,7 +166,7 @@ from sglang.srt.managers.scheduler_runtime_checker_mixin import (
|
|||||||
from sglang.srt.managers.scheduler_update_weights_mixin import (
|
from sglang.srt.managers.scheduler_update_weights_mixin import (
|
||||||
SchedulerUpdateWeightsMixin,
|
SchedulerUpdateWeightsMixin,
|
||||||
)
|
)
|
||||||
from sglang.srt.managers.session_controller import Session
|
from sglang.srt.managers.session_controller import SessionController
|
||||||
from sglang.srt.managers.utils import GenerationBatchResult, validate_input_length
|
from sglang.srt.managers.utils import GenerationBatchResult, validate_input_length
|
||||||
from sglang.srt.mem_cache.cache_init_params import CacheInitParams
|
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.common import release_kv_cache
|
||||||
@@ -755,8 +754,7 @@ class Scheduler(
|
|||||||
self.return_health_check_ct = 0
|
self.return_health_check_ct = 0
|
||||||
self.num_retracted_reqs: int = 0
|
self.num_retracted_reqs: int = 0
|
||||||
self.num_paused_reqs: int = 0
|
self.num_paused_reqs: int = 0
|
||||||
self.sessions: Dict[str, Session] = {}
|
self.session_controller = SessionController(self.tree_cache)
|
||||||
self._last_reap_sessions: float = 0.0
|
|
||||||
self.forward_sleep_time = None
|
self.forward_sleep_time = None
|
||||||
self._engine_paused = False
|
self._engine_paused = False
|
||||||
|
|
||||||
@@ -1131,7 +1129,6 @@ class Scheduler(
|
|||||||
self.process_batch_result(batch, result)
|
self.process_batch_result(batch, result)
|
||||||
else:
|
else:
|
||||||
# When the server is idle, do self-check and re-init some states.
|
# When the server is idle, do self-check and re-init some states.
|
||||||
# Skip if there are any streaming sessions (latency sensitive).
|
|
||||||
self.self_check_during_idle()
|
self.self_check_during_idle()
|
||||||
|
|
||||||
# Update last_batch
|
# Update last_batch
|
||||||
@@ -1363,9 +1360,7 @@ class Scheduler(
|
|||||||
|
|
||||||
def process_input_requests(self, recv_reqs: List):
|
def process_input_requests(self, recv_reqs: List):
|
||||||
now = time.monotonic()
|
now = time.monotonic()
|
||||||
if now - self._last_reap_sessions > 1.0: # reap sessions every second
|
self.session_controller.maybe_reap(now)
|
||||||
self._last_reap_sessions = now
|
|
||||||
self.reap_timed_out_sessions()
|
|
||||||
for recv_req in recv_reqs:
|
for recv_req in recv_reqs:
|
||||||
# If it is a health check generation request and there are running requests, ignore it.
|
# If it is a health check generation request and there are running requests, ignore it.
|
||||||
if is_health_check_generate_req(recv_req) and (
|
if is_health_check_generate_req(recv_req) and (
|
||||||
@@ -1485,12 +1480,13 @@ class Scheduler(
|
|||||||
self,
|
self,
|
||||||
recv_req: TokenizedGenerateReqInput,
|
recv_req: TokenizedGenerateReqInput,
|
||||||
):
|
):
|
||||||
# Create a new request
|
# Route: normal request / session request / session-not-found
|
||||||
if (
|
session_id = (
|
||||||
recv_req.session_params is None
|
recv_req.session_params.id if recv_req.session_params is not None else None
|
||||||
or recv_req.session_params.id is None
|
)
|
||||||
or recv_req.session_params.id not in self.sessions
|
|
||||||
):
|
if session_id is None:
|
||||||
|
# Normal non-session request
|
||||||
if recv_req.input_embeds is not None:
|
if recv_req.input_embeds is not None:
|
||||||
# Generate fake input_ids based on the length of input_embeds
|
# Generate fake input_ids based on the length of input_embeds
|
||||||
seq_length = len(recv_req.input_embeds)
|
seq_length = len(recv_req.input_embeds)
|
||||||
@@ -1553,21 +1549,14 @@ class Scheduler(
|
|||||||
self.stream_output([req], req.return_logprob)
|
self.stream_output([req], req.return_logprob)
|
||||||
return
|
return
|
||||||
|
|
||||||
if (
|
elif session_id in self.session_controller:
|
||||||
recv_req.session_params is not None
|
# Session exists: create request from session
|
||||||
and recv_req.session_params.id is not None
|
session = self.session_controller.get(session_id)
|
||||||
):
|
|
||||||
req.set_finish_with_abort(
|
|
||||||
f"Invalid request: session id {recv_req.session_params.id} does not exist"
|
|
||||||
)
|
|
||||||
self.init_req_max_new_tokens(req)
|
|
||||||
self._add_request_to_queue(req)
|
|
||||||
return
|
|
||||||
else:
|
|
||||||
# Create a new request from a previous session
|
|
||||||
session = self.sessions[recv_req.session_params.id]
|
|
||||||
req = session.create_req(
|
req = session.create_req(
|
||||||
recv_req, self.tokenizer, self.model_config.vocab_size
|
recv_req,
|
||||||
|
self.tokenizer,
|
||||||
|
self.model_config.vocab_size,
|
||||||
|
eos_token_ids=self.model_config.hf_eos_token_id,
|
||||||
)
|
)
|
||||||
# TODO: set trace context
|
# TODO: set trace context
|
||||||
if self.enable_metrics:
|
if self.enable_metrics:
|
||||||
@@ -1577,22 +1566,28 @@ class Scheduler(
|
|||||||
self._add_request_to_queue(req)
|
self._add_request_to_queue(req)
|
||||||
return
|
return
|
||||||
|
|
||||||
|
else:
|
||||||
|
# Session ID provided but session not found
|
||||||
|
req = Req(
|
||||||
|
recv_req.rid,
|
||||||
|
recv_req.input_text,
|
||||||
|
recv_req.input_ids,
|
||||||
|
recv_req.sampling_params,
|
||||||
|
vocab_size=self.model_config.vocab_size,
|
||||||
|
)
|
||||||
|
req.tokenizer = self.tokenizer
|
||||||
|
req.set_finish_with_abort(
|
||||||
|
f"Invalid request: session id {session_id} does not exist"
|
||||||
|
)
|
||||||
|
self.init_req_max_new_tokens(req)
|
||||||
|
self._add_request_to_queue(req)
|
||||||
|
return
|
||||||
|
|
||||||
# Handle multimodal inputs
|
# Handle multimodal inputs
|
||||||
if recv_req.mm_inputs is not None:
|
if recv_req.mm_inputs is not None:
|
||||||
image_inputs = self._get_multimodal_inputs(recv_req.mm_inputs)
|
image_inputs = self._get_multimodal_inputs(recv_req.mm_inputs)
|
||||||
|
|
||||||
# For session requests, adjust mm_inputs offsets by the prefix length.
|
SessionController.adjust_mm_offsets(recv_req, req, image_inputs)
|
||||||
# Session.create_req prepends previous context to origin_input_ids,
|
|
||||||
# so offsets from the new prompt need to be shifted.
|
|
||||||
if len(recv_req.input_ids) < len(req.origin_input_ids):
|
|
||||||
assert recv_req.session_params.id in self.sessions
|
|
||||||
prefix_len = len(req.origin_input_ids) - len(recv_req.input_ids)
|
|
||||||
for mm_item in image_inputs.mm_items:
|
|
||||||
if mm_item.offsets:
|
|
||||||
mm_item.offsets = [
|
|
||||||
(start + prefix_len, end + prefix_len)
|
|
||||||
for start, end in mm_item.offsets
|
|
||||||
]
|
|
||||||
|
|
||||||
# The following steps are already fast, execute locally on each rank.
|
# The following steps are already fast, execute locally on each rank.
|
||||||
# Expand a single image token into multiple dummy tokens for receiving image embeddings
|
# Expand a single image token into multiple dummy tokens for receiving image embeddings
|
||||||
@@ -2938,47 +2933,10 @@ class Scheduler(
|
|||||||
return ExpertDistributionReqOutput()
|
return ExpertDistributionReqOutput()
|
||||||
|
|
||||||
def open_session(self, recv_req: OpenSessionReqInput):
|
def open_session(self, recv_req: OpenSessionReqInput):
|
||||||
session_id = recv_req.session_id
|
return self.session_controller.open(recv_req)
|
||||||
if session_id in self.sessions:
|
|
||||||
logger.warning(f"session id {session_id} already exist, cannot open.")
|
|
||||||
return OpenSessionReqOutput(session_id, False)
|
|
||||||
elif session_id is None:
|
|
||||||
logger.warning("session id is None, cannot open.")
|
|
||||||
return OpenSessionReqOutput(session_id, False)
|
|
||||||
else:
|
|
||||||
self.sessions[session_id] = Session(
|
|
||||||
recv_req.capacity_of_str_len,
|
|
||||||
session_id,
|
|
||||||
streaming=bool(recv_req.streaming),
|
|
||||||
timeout=recv_req.timeout,
|
|
||||||
)
|
|
||||||
return OpenSessionReqOutput(session_id, True)
|
|
||||||
|
|
||||||
def close_session(self, recv_req: CloseSessionReqInput):
|
def close_session(self, recv_req: CloseSessionReqInput):
|
||||||
session_id = recv_req.session_id
|
self.session_controller.close(recv_req)
|
||||||
if session_id not in self.sessions:
|
|
||||||
logger.warning(f"session id {session_id} does not exist, cannot delete.")
|
|
||||||
else:
|
|
||||||
self._close_session(session_id)
|
|
||||||
|
|
||||||
def _close_session(self, session_id: str):
|
|
||||||
session = self.sessions[session_id]
|
|
||||||
if session.streaming and session.req_nodes:
|
|
||||||
assert len(session.req_nodes) == 1
|
|
||||||
req = next(iter(session.req_nodes.values())).req
|
|
||||||
if not req.finished():
|
|
||||||
req.session = None
|
|
||||||
if isinstance(self.tree_cache, SessionAwareCache):
|
|
||||||
self.tree_cache.release_session(session_id)
|
|
||||||
del self.sessions[session_id]
|
|
||||||
|
|
||||||
def reap_timed_out_sessions(self):
|
|
||||||
timed_out = [
|
|
||||||
sid for sid, session in self.sessions.items() if session.is_timed_out()
|
|
||||||
]
|
|
||||||
for sid in timed_out:
|
|
||||||
logger.info(f"Session {sid} timed out, closing.")
|
|
||||||
self._close_session(sid)
|
|
||||||
|
|
||||||
def maybe_sleep_on_idle(self):
|
def maybe_sleep_on_idle(self):
|
||||||
if self.idle_sleeper is not None:
|
if self.idle_sleeper is not None:
|
||||||
|
|||||||
@@ -10,13 +10,26 @@
|
|||||||
# limitations under the License.
|
# limitations under the License.
|
||||||
# ==============================================================================
|
# ==============================================================================
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
import logging
|
import logging
|
||||||
import time
|
import time
|
||||||
import uuid
|
import uuid
|
||||||
from typing import Dict, Optional
|
from typing import TYPE_CHECKING, Dict, Optional
|
||||||
|
|
||||||
from sglang.srt.managers.io_struct import TokenizedGenerateReqInput
|
from sglang.srt.managers.io_struct import (
|
||||||
|
CloseSessionReqInput,
|
||||||
|
OpenSessionReqInput,
|
||||||
|
OpenSessionReqOutput,
|
||||||
|
TokenizedGenerateReqInput,
|
||||||
|
)
|
||||||
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
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from sglang.srt.mem_cache.base_prefix_cache import BasePrefixCache
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
class SessionReqNode:
|
class SessionReqNode:
|
||||||
@@ -85,7 +98,13 @@ class Session:
|
|||||||
return False
|
return False
|
||||||
return time.monotonic() - self.last_active_time > self.timeout
|
return time.monotonic() - self.last_active_time > self.timeout
|
||||||
|
|
||||||
def create_req(self, req: TokenizedGenerateReqInput, tokenizer, vocab_size: int):
|
def create_req(
|
||||||
|
self,
|
||||||
|
req: TokenizedGenerateReqInput,
|
||||||
|
tokenizer,
|
||||||
|
vocab_size: int,
|
||||||
|
eos_token_ids=None,
|
||||||
|
):
|
||||||
assert req.session_params is not None
|
assert req.session_params is not None
|
||||||
self.last_active_time = time.monotonic()
|
self.last_active_time = time.monotonic()
|
||||||
session_params = req.session_params
|
session_params = req.session_params
|
||||||
@@ -190,6 +209,14 @@ class Session:
|
|||||||
top_logprobs_num=req.top_logprobs_num,
|
top_logprobs_num=req.top_logprobs_num,
|
||||||
token_ids_logprob=req.token_ids_logprob,
|
token_ids_logprob=req.token_ids_logprob,
|
||||||
vocab_size=vocab_size,
|
vocab_size=vocab_size,
|
||||||
|
eos_token_ids=eos_token_ids,
|
||||||
|
require_reasoning=req.require_reasoning,
|
||||||
|
return_hidden_states=req.return_hidden_states,
|
||||||
|
return_routed_experts=req.return_routed_experts,
|
||||||
|
priority=req.priority,
|
||||||
|
routing_key=req.routing_key,
|
||||||
|
http_worker_ipc=req.http_worker_ipc,
|
||||||
|
time_stats=req.time_stats,
|
||||||
)
|
)
|
||||||
if last_req is not None:
|
if last_req is not None:
|
||||||
new_req.multimodal_inputs = last_req.multimodal_inputs
|
new_req.multimodal_inputs = last_req.multimodal_inputs
|
||||||
@@ -206,3 +233,77 @@ class Session:
|
|||||||
self.req_nodes[req.rid] = new_req_node
|
self.req_nodes[req.rid] = new_req_node
|
||||||
|
|
||||||
return new_req
|
return new_req
|
||||||
|
|
||||||
|
|
||||||
|
class SessionController:
|
||||||
|
def __init__(self, tree_cache: BasePrefixCache):
|
||||||
|
self.sessions: Dict[str, Session] = {}
|
||||||
|
self._last_reap_time: float = 0.0
|
||||||
|
self.tree_cache = tree_cache
|
||||||
|
|
||||||
|
def __contains__(self, session_id: str) -> bool:
|
||||||
|
return session_id in self.sessions
|
||||||
|
|
||||||
|
def get(self, session_id: str) -> Optional[Session]:
|
||||||
|
return self.sessions.get(session_id)
|
||||||
|
|
||||||
|
def open(self, recv_req: OpenSessionReqInput) -> OpenSessionReqOutput:
|
||||||
|
session_id = recv_req.session_id
|
||||||
|
if session_id in self.sessions:
|
||||||
|
logger.warning(f"session id {session_id} already exist, cannot open.")
|
||||||
|
return OpenSessionReqOutput(session_id, False)
|
||||||
|
elif session_id is None:
|
||||||
|
logger.warning("session id is None, cannot open.")
|
||||||
|
return OpenSessionReqOutput(session_id, False)
|
||||||
|
else:
|
||||||
|
self.sessions[session_id] = Session(
|
||||||
|
recv_req.capacity_of_str_len,
|
||||||
|
session_id,
|
||||||
|
streaming=bool(recv_req.streaming),
|
||||||
|
timeout=recv_req.timeout,
|
||||||
|
)
|
||||||
|
return OpenSessionReqOutput(session_id, True)
|
||||||
|
|
||||||
|
def close(self, recv_req: CloseSessionReqInput):
|
||||||
|
session_id = recv_req.session_id
|
||||||
|
if session_id not in self.sessions:
|
||||||
|
logger.warning(f"session id {session_id} does not exist, cannot delete.")
|
||||||
|
else:
|
||||||
|
self._close(session_id)
|
||||||
|
|
||||||
|
def _close(self, session_id: str):
|
||||||
|
session = self.sessions[session_id]
|
||||||
|
if session.streaming and session.req_nodes:
|
||||||
|
assert len(session.req_nodes) == 1
|
||||||
|
req = next(iter(session.req_nodes.values())).req
|
||||||
|
if not req.finished():
|
||||||
|
req.session = None
|
||||||
|
if isinstance(self.tree_cache, SessionAwareCache):
|
||||||
|
self.tree_cache.release_session(session_id)
|
||||||
|
del self.sessions[session_id]
|
||||||
|
|
||||||
|
def maybe_reap(self, now: float, interval: float = 1.0):
|
||||||
|
# reap sessions every second
|
||||||
|
if now - self._last_reap_time > interval:
|
||||||
|
self._last_reap_time = now
|
||||||
|
timed_out = [
|
||||||
|
sid for sid, session in self.sessions.items() if session.is_timed_out()
|
||||||
|
]
|
||||||
|
for sid in timed_out:
|
||||||
|
logger.info(f"Session {sid} timed out, closing.")
|
||||||
|
self._close(sid)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def adjust_mm_offsets(recv_req: TokenizedGenerateReqInput, req: Req, image_inputs):
|
||||||
|
# For session requests, adjust mm_inputs offsets by the prefix length.
|
||||||
|
# Session.create_req prepends previous context to origin_input_ids,
|
||||||
|
# so offsets from the new prompt need to be shifted.
|
||||||
|
if len(recv_req.input_ids) >= len(req.origin_input_ids):
|
||||||
|
return
|
||||||
|
prefix_len = len(req.origin_input_ids) - len(recv_req.input_ids)
|
||||||
|
for mm_item in image_inputs.mm_items:
|
||||||
|
if mm_item.offsets:
|
||||||
|
mm_item.offsets = [
|
||||||
|
(start + prefix_len, end + prefix_len)
|
||||||
|
for start, end in mm_item.offsets
|
||||||
|
]
|
||||||
|
|||||||
Reference in New Issue
Block a user