feat: first-class session identity in SGLang (#29436)
This commit is contained in:
@@ -463,18 +463,21 @@ class Runtime:
|
||||
self,
|
||||
prompt: str,
|
||||
sampling_params: Optional[Dict] = None,
|
||||
session_id: Optional[str] = None,
|
||||
):
|
||||
if self.server_args.skip_tokenizer_init:
|
||||
json_data = {
|
||||
"input_ids": prompt,
|
||||
"sampling_params": sampling_params,
|
||||
"stream": True,
|
||||
"session_id": session_id,
|
||||
}
|
||||
else:
|
||||
json_data = {
|
||||
"text": prompt,
|
||||
"sampling_params": sampling_params,
|
||||
"stream": True,
|
||||
"session_id": session_id,
|
||||
}
|
||||
pos = 0
|
||||
|
||||
@@ -505,6 +508,7 @@ class Runtime:
|
||||
logprob_start_len: Optional[Union[List[int], int]] = None,
|
||||
top_logprobs_num: Optional[Union[List[int], int]] = None,
|
||||
lora_path: Optional[List[Optional[str]]] = None,
|
||||
session_id: Optional[str] = None,
|
||||
):
|
||||
json_data = {
|
||||
"text": prompt,
|
||||
@@ -513,6 +517,7 @@ class Runtime:
|
||||
"logprob_start_len": logprob_start_len,
|
||||
"top_logprobs_num": top_logprobs_num,
|
||||
"lora_path": lora_path,
|
||||
"session_id": session_id,
|
||||
}
|
||||
assert not isinstance(lora_path, list) or len(lora_path) == len(prompt)
|
||||
response = requests.post(
|
||||
|
||||
@@ -33,6 +33,7 @@ class EngineBase(ABC):
|
||||
data_parallel_rank: Optional[int] = None,
|
||||
rid: Optional[Union[List[str], str]] = None,
|
||||
priority: Optional[int] = None,
|
||||
session_id: Optional[str] = None,
|
||||
) -> Union[Dict, Iterator[Dict]]:
|
||||
"""Generate outputs based on given inputs."""
|
||||
pass
|
||||
|
||||
@@ -357,6 +357,7 @@ class Engine(EngineScoreMixin, EngineBase):
|
||||
rid: Optional[Union[List[str], str]] = None,
|
||||
session_params: Optional[Dict] = None,
|
||||
priority: Optional[int] = None,
|
||||
session_id: Optional[str] = None,
|
||||
) -> Union[Dict, Iterator[Dict]]:
|
||||
"""
|
||||
The arguments of this function is the same as `sglang/srt/managers/io_struct.py::GenerateReqInput`.
|
||||
@@ -392,6 +393,7 @@ class Engine(EngineScoreMixin, EngineBase):
|
||||
disagg_prefill_dp_rank=disagg_prefill_dp_rank,
|
||||
external_trace_header=external_trace_header,
|
||||
rid=rid,
|
||||
session_id=session_id,
|
||||
session_params=session_params,
|
||||
priority=priority,
|
||||
)
|
||||
@@ -459,6 +461,7 @@ class Engine(EngineScoreMixin, EngineBase):
|
||||
rid: Optional[Union[List[str], str]] = None,
|
||||
session_params: Optional[Dict] = None,
|
||||
priority: Optional[int] = None,
|
||||
session_id: Optional[str] = None,
|
||||
) -> Union[Dict, AsyncIterator[Dict]]:
|
||||
"""
|
||||
The arguments of this function is the same as `sglang/srt/managers/io_struct.py::GenerateReqInput`.
|
||||
@@ -494,6 +497,7 @@ class Engine(EngineScoreMixin, EngineBase):
|
||||
disagg_prefill_dp_rank=disagg_prefill_dp_rank,
|
||||
external_trace_header=external_trace_header,
|
||||
rid=rid,
|
||||
session_id=session_id,
|
||||
session_params=session_params,
|
||||
priority=priority,
|
||||
)
|
||||
|
||||
@@ -115,6 +115,7 @@ class HttpServerEngineAdapter(EngineBase):
|
||||
lora_path=None,
|
||||
custom_logit_processor=None,
|
||||
priority=None,
|
||||
session_id=None,
|
||||
):
|
||||
payload = {
|
||||
"text": prompt,
|
||||
@@ -128,6 +129,7 @@ class HttpServerEngineAdapter(EngineBase):
|
||||
"lora_path": lora_path,
|
||||
"custom_logit_processor": custom_logit_processor,
|
||||
"priority": priority,
|
||||
"session_id": session_id,
|
||||
}
|
||||
# Filter out None values
|
||||
payload = {k: v for k, v in payload.items() if v is not None}
|
||||
|
||||
@@ -356,6 +356,7 @@ class CompletionRequest(BaseModel):
|
||||
ignore_eos: bool = False
|
||||
skip_special_tokens: bool = True
|
||||
lora_path: Optional[Union[List[Optional[str]], Optional[str]]] = None
|
||||
session_id: Optional[str] = None
|
||||
session_params: Optional[Dict] = None
|
||||
response_format: Optional[Union[ResponseFormat, StructuralTagResponseFormat]] = None
|
||||
custom_params: Optional[Dict] = None
|
||||
@@ -728,6 +729,7 @@ class ChatCompletionRequest(BaseModel):
|
||||
continue_final_message: bool = False
|
||||
skip_special_tokens: bool = True
|
||||
lora_path: Optional[Union[List[Optional[str]], Optional[str]]] = None
|
||||
session_id: Optional[str] = None
|
||||
session_params: Optional[Dict] = None
|
||||
separate_reasoning: bool = True
|
||||
stream_reasoning: bool = True
|
||||
@@ -1395,6 +1397,7 @@ class ResponsesRequest(BaseModel):
|
||||
default_factory=lambda: f"resp_{uuid.uuid4().hex}",
|
||||
description="The request_id related to this request. If the caller does not set it, a random uuid will be generated.",
|
||||
)
|
||||
session_id: Optional[str] = None
|
||||
priority: int = Field(default=0, description="Request priority")
|
||||
extra_key: Optional[str] = Field(
|
||||
default=None,
|
||||
|
||||
@@ -606,6 +606,7 @@ class OpenAIServingChat(OpenAIServingBase):
|
||||
return_routed_experts=request.return_routed_experts,
|
||||
routed_experts_start_len=request.routed_experts_start_len,
|
||||
rid=request.rid,
|
||||
session_id=request.session_id,
|
||||
extra_key=self._compute_extra_key(request),
|
||||
require_reasoning=require_reasoning,
|
||||
priority=request.priority,
|
||||
|
||||
@@ -125,6 +125,7 @@ class OpenAIServingCompletion(OpenAIServingBase):
|
||||
return_routed_experts=request.return_routed_experts,
|
||||
routed_experts_start_len=request.routed_experts_start_len,
|
||||
rid=request.rid,
|
||||
session_id=request.session_id,
|
||||
extra_key=self._compute_extra_key(request),
|
||||
priority=request.priority,
|
||||
routing_key=self.extract_routing_key(raw_request),
|
||||
|
||||
@@ -360,6 +360,7 @@ class OpenAIServingResponses(OpenAIServingChat):
|
||||
sampling_params=sampling_params,
|
||||
stream=request.stream,
|
||||
rid=request.request_id,
|
||||
session_id=request.session_id,
|
||||
extra_key=self._compute_extra_key(request),
|
||||
background=request.background,
|
||||
)
|
||||
@@ -2362,6 +2363,7 @@ class OpenAIServingResponses(OpenAIServingChat):
|
||||
sampling_params=sampling_params,
|
||||
stream=adapted_request.stream,
|
||||
rid=request_id,
|
||||
session_id=adapted_request.session_id,
|
||||
extra_key=adapted_request.extra_key,
|
||||
return_logprob=adapted_request.return_logprob,
|
||||
logprob_start_len=adapted_request.logprob_start_len,
|
||||
|
||||
@@ -153,6 +153,9 @@ class GenerateReqInput:
|
||||
# Request ID(s). If omitted, generated during normalization. For batch
|
||||
# requests, a string is expanded to per-item IDs using it as a prefix.
|
||||
rid: Optional[Union[str, List[str]]] = field(default=None, kw_only=True)
|
||||
# Stable identity shared by requests in the same session. Unlike
|
||||
# session_params, this does not alter or reconstruct the prompt.
|
||||
session_id: Optional[str] = field(default=None, kw_only=True)
|
||||
# The input prompt. It can be a single prompt or a batch of prompts.
|
||||
text: Optional[Union[List[str], str]] = None
|
||||
# The token ids for text.
|
||||
@@ -338,6 +341,8 @@ class GenerateReqInput:
|
||||
"""
|
||||
self._validate_inputs()
|
||||
self._determine_batch_size()
|
||||
if self.session_id is not None and self.session_params is not None:
|
||||
raise ValueError("session_id and session_params cannot both be set.")
|
||||
self._handle_parallel_sampling()
|
||||
|
||||
if self.is_single:
|
||||
@@ -693,6 +698,7 @@ class GenerateReqInput:
|
||||
return cache[i]
|
||||
sub = GenerateReqInput(
|
||||
rid=self.rid[i],
|
||||
session_id=self.session_id,
|
||||
text=self.text[i] if self.text is not None else None,
|
||||
input_ids=self.input_ids[i] if self.input_ids is not None else None,
|
||||
input_embeds=(
|
||||
@@ -800,6 +806,7 @@ class TokenizedGenerateReqInput(BaseReq, kw_only=True):
|
||||
return_indexer_topk: bool = False
|
||||
|
||||
# Session info for continual prompting
|
||||
session_id: Optional[str] = field(default=None, kw_only=True)
|
||||
session_params: Optional[SessionParams] = None
|
||||
|
||||
# LoRA related
|
||||
|
||||
@@ -708,6 +708,7 @@ class Req(ReqDllmMixin):
|
||||
] = None,
|
||||
return_pooled_hidden_states: bool = False,
|
||||
multi_item_delimiter_indices: Optional[List[int]] = None,
|
||||
session_id: Optional[str] = None,
|
||||
):
|
||||
# Input and output info
|
||||
self.rid = rid
|
||||
@@ -730,6 +731,7 @@ class Req(ReqDllmMixin):
|
||||
self.dllm_initialized: bool = False
|
||||
|
||||
self.session = session
|
||||
self.session_id = session_id
|
||||
self.input_embeds = input_embeds
|
||||
self.positional_embed_overrides = positional_embed_overrides
|
||||
self.multi_item_delimiter_indices = multi_item_delimiter_indices
|
||||
|
||||
@@ -123,7 +123,6 @@ from sglang.srt.managers.io_struct import (
|
||||
LoadLoRAAdapterReqInput,
|
||||
LoadLoRAAdapterReqOutput,
|
||||
OpenSessionReqInput,
|
||||
OpenSessionReqOutput,
|
||||
PauseGenerationReqInput,
|
||||
ProfileReq,
|
||||
ReleaseMemoryOccupationReqInput,
|
||||
@@ -2011,36 +2010,11 @@ class Scheduler(
|
||||
session_id = (
|
||||
recv_req.session_params.id if recv_req.session_params is not None else None
|
||||
)
|
||||
# Radix-native session: session_id is just a tag; KV bulk-freed on close.
|
||||
# Radix-native sessions use only the top-level session_id.
|
||||
radix_native_session = (
|
||||
session_id is not None and self.server_args.enable_session_radix_cache
|
||||
recv_req.session_id is not None
|
||||
and self.server_args.enable_session_radix_cache
|
||||
)
|
||||
if radix_native_session:
|
||||
sp = recv_req.session_params
|
||||
if (
|
||||
sp.rid is not None
|
||||
or sp.offset is not None
|
||||
or sp.replace is not None
|
||||
or sp.drop_previous_output is not None
|
||||
):
|
||||
error_msg = (
|
||||
"Invalid request: radix-native sessions do not support "
|
||||
"session_params rid/offset/replace/drop_previous_output; "
|
||||
"send full context each turn."
|
||||
)
|
||||
req = Req(
|
||||
recv_req.rid,
|
||||
recv_req.input_text,
|
||||
recv_req.input_ids,
|
||||
recv_req.sampling_params,
|
||||
vocab_size=self.model_config.vocab_size,
|
||||
http_worker_ipc=recv_req.http_worker_ipc,
|
||||
)
|
||||
req.tokenizer = self.tokenizer
|
||||
req.set_finish_with_abort(error_msg)
|
||||
self.init_req_max_new_tokens(req)
|
||||
self._add_request_to_queue(req)
|
||||
return
|
||||
|
||||
if session_id is None or radix_native_session:
|
||||
# Normal non-session request, or a radix-native session request
|
||||
@@ -2063,6 +2037,7 @@ class Scheduler(
|
||||
token_ids_logprob=recv_req.token_ids_logprob,
|
||||
stream=recv_req.stream,
|
||||
lora_id=recv_req.lora_id,
|
||||
session_id=recv_req.session_id,
|
||||
input_embeds=recv_req.input_embeds,
|
||||
positional_embed_overrides=recv_req.positional_embed_overrides,
|
||||
token_type_ids=recv_req.token_type_ids,
|
||||
@@ -2094,8 +2069,6 @@ class Scheduler(
|
||||
multi_item_delimiter_indices=recv_req.multi_item_delimiter_indices,
|
||||
)
|
||||
req.tokenizer = self.tokenizer
|
||||
if radix_native_session:
|
||||
req.session_id = session_id
|
||||
|
||||
if self.disaggregation_mode != DisaggregationMode.NULL:
|
||||
# Invalid request for disaggregated mode
|
||||
@@ -4082,24 +4055,17 @@ class Scheduler(
|
||||
return ExpertDistributionReqOutput()
|
||||
|
||||
def open_session(self, recv_req: OpenSessionReqInput):
|
||||
if self.server_args.enable_session_radix_cache:
|
||||
# Radix-native: open is implicit; explicit open only permits id reuse.
|
||||
session_id = recv_req.session_id
|
||||
self.tree_cache.register_session(session_id)
|
||||
output = OpenSessionReqOutput(
|
||||
session_id=session_id, success=session_id is not None
|
||||
)
|
||||
else:
|
||||
output = self.session_controller.open(recv_req)
|
||||
output = self.session_controller.open(recv_req)
|
||||
if self.ps.pp_rank == 0 and self.ps.tp_rank == 0 and self.ps.attn_cp_rank == 0:
|
||||
return output
|
||||
return None
|
||||
|
||||
def close_session(self, recv_req: CloseSessionReqInput):
|
||||
if self.server_args.enable_session_radix_cache:
|
||||
# "Close" just triggers eviction of the session's tagged KV.
|
||||
self.tree_cache.release_session(recv_req.session_id)
|
||||
else:
|
||||
self.tree_cache.release_radix_session(recv_req.session_id)
|
||||
if recv_req.session_id in self.session_controller or not (
|
||||
self.server_args.enable_session_radix_cache
|
||||
):
|
||||
self.session_controller.close(recv_req)
|
||||
|
||||
def maybe_sleep_on_idle(self):
|
||||
|
||||
@@ -1165,6 +1165,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
||||
lora_id=obj.lora_id,
|
||||
input_embeds=input_embeds,
|
||||
positional_embed_overrides=obj.positional_embed_overrides,
|
||||
session_id=obj.session_id,
|
||||
session_params=session_params,
|
||||
custom_logit_processor=obj.custom_logit_processor,
|
||||
require_reasoning=obj.require_reasoning,
|
||||
|
||||
@@ -337,6 +337,9 @@ class BasePrefixCache(ABC, PrefixCacheTrait):
|
||||
def release_session(self, session_id: str) -> None:
|
||||
pass
|
||||
|
||||
def release_radix_session(self, session_id: str) -> None:
|
||||
pass
|
||||
|
||||
def session_held_tokens(self, active_pool_idxs: Optional[set] = None) -> int:
|
||||
return 0
|
||||
|
||||
|
||||
@@ -31,6 +31,7 @@ class CacheInitParams:
|
||||
|
||||
enable_metrics: bool = False
|
||||
enable_kv_cache_events: bool = False
|
||||
enable_session_radix_cache: bool = False
|
||||
|
||||
enable_mamba_extra_buffer: bool = False
|
||||
enable_mamba_extra_buffer_lazy: bool = False
|
||||
|
||||
@@ -219,6 +219,7 @@ def build_kv_cache(
|
||||
eviction_policy=server_args.radix_eviction_policy,
|
||||
enable_metrics=enable_metrics,
|
||||
enable_kv_cache_events=enable_kv_cache_events,
|
||||
enable_session_radix_cache=server_args.enable_session_radix_cache,
|
||||
enable_mamba_extra_buffer=server_args.enable_mamba_extra_buffer(),
|
||||
enable_mamba_extra_buffer_lazy=server_args.enable_mamba_extra_buffer_lazy(),
|
||||
pp_rank=ps.pp_rank,
|
||||
|
||||
@@ -290,6 +290,7 @@ class RadixCache(SessionRadixCacheMixin, KVCacheEventMixin, BasePrefixCache):
|
||||
self.token_to_kv_pool_allocator = params.token_to_kv_pool_allocator
|
||||
self.page_size = params.page_size
|
||||
self.enable_kv_cache_events = params.enable_kv_cache_events
|
||||
self.enable_session_radix_cache = params.enable_session_radix_cache
|
||||
self.is_eagle = params.is_eagle
|
||||
self.disable_finished_insert = params.disable_finished_insert
|
||||
self.eviction_policy = params.eviction_policy.lower()
|
||||
|
||||
@@ -34,13 +34,6 @@ class SessionRadixCacheMixin:
|
||||
if not hasattr(self, "_session_leaves"):
|
||||
self._reset_session_radix_state()
|
||||
|
||||
def register_session(self, session_id: str) -> None:
|
||||
self._ensure_session_radix_state()
|
||||
if session_id is None:
|
||||
return
|
||||
self._closed_session_ids.pop(session_id, None)
|
||||
self._session_leaves.setdefault(session_id, set())
|
||||
|
||||
def _remember_closed_session(self, session_id: str) -> None:
|
||||
self._closed_session_ids[session_id] = None
|
||||
self._closed_session_ids.move_to_end(session_id)
|
||||
@@ -62,6 +55,8 @@ class SessionRadixCacheMixin:
|
||||
|
||||
def _tag_session_leaf(self, req: Req, radix_key, node=None) -> None:
|
||||
"""Add this request's session id to its leaf's holder set; no-op for non-session reqs."""
|
||||
if not self.enable_session_radix_cache:
|
||||
return
|
||||
self._ensure_session_radix_state()
|
||||
sid = getattr(req, "session_id", None)
|
||||
if sid is None or sid in self._closed_session_ids:
|
||||
@@ -86,7 +81,7 @@ class SessionRadixCacheMixin:
|
||||
len(self._session_leaves[sid]),
|
||||
)
|
||||
|
||||
def release_session(self, session_id: str) -> int:
|
||||
def release_radix_session(self, session_id: str) -> int:
|
||||
"""Close: drop this session from each of its tagged leaves, freeing a node
|
||||
only once no other session still holds it (last holder). Shared
|
||||
prefixes/leaves kept."""
|
||||
|
||||
@@ -438,6 +438,9 @@ class StreamingSession(BasePrefixCache):
|
||||
|
||||
self._free_slot_mamba(slot)
|
||||
|
||||
def release_radix_session(self, session_id: str) -> None:
|
||||
self.inner.release_radix_session(session_id)
|
||||
|
||||
def session_held_tokens(self, active_pool_idxs: Optional[set] = None) -> int:
|
||||
"""Total KV tokens held by session slots, not tracked by the tree.
|
||||
|
||||
|
||||
Reference in New Issue
Block a user