Clean up TokenizerManager and req_time_stats: reduce overhead and simplify (#21646)
This commit is contained in:
@@ -84,7 +84,6 @@ from sglang.srt.observability.cpu_monitor import start_cpu_monitor_thread
|
||||
from sglang.srt.observability.metrics_collector import TokenizerMetricsCollector
|
||||
from sglang.srt.observability.req_time_stats import (
|
||||
APIServerReqTimeStats,
|
||||
calibrate_time_diff,
|
||||
convert_time_to_realtime,
|
||||
real_time,
|
||||
set_time_batch,
|
||||
@@ -205,22 +204,6 @@ def _slice_streaming_output_meta_info(
|
||||
meta_info[key] = meta_info[key][last_output_offset:]
|
||||
|
||||
|
||||
def _merge_incremental_stream_meta_info(
|
||||
out_list: list[dict[str, Any]],
|
||||
) -> dict[str, Any]:
|
||||
"""Preserve delta-style output metadata when queued chunks are coalesced."""
|
||||
meta_info_list = [chunk["meta_info"] for chunk in out_list]
|
||||
meta_info = dict(meta_info_list[-1])
|
||||
for key in _INCREMENTAL_STREAMING_META_INFO_KEYS:
|
||||
if any(key in chunk_meta_info for chunk_meta_info in meta_info_list):
|
||||
meta_info[key] = [
|
||||
item
|
||||
for chunk_meta_info in meta_info_list
|
||||
for item in chunk_meta_info.get(key, [])
|
||||
]
|
||||
return meta_info
|
||||
|
||||
|
||||
class InputFormat(Enum):
|
||||
"""Input format types for tokenization handling."""
|
||||
|
||||
@@ -268,9 +251,6 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerScoreMixin):
|
||||
# Init PD disaggregation and encoder disaggregation
|
||||
self.init_disaggregation()
|
||||
|
||||
# Subprocess liveness watchdog — set by Engine or http_server after construction
|
||||
self._subprocess_watchdog = None
|
||||
|
||||
# Init metric collector and watchdog
|
||||
self.init_metric_collector_watchdog()
|
||||
|
||||
@@ -395,6 +375,9 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerScoreMixin):
|
||||
# Session
|
||||
self.session_futures = {} # session_id -> asyncio event
|
||||
|
||||
# Subprocess liveness watchdog — set by Engine or http_server after construction
|
||||
self._subprocess_watchdog = None
|
||||
|
||||
def init_request_logging_and_dumping(self):
|
||||
# TODO: Refactor and organize the log export code.
|
||||
# Request logging
|
||||
@@ -539,7 +522,6 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerScoreMixin):
|
||||
# Normalize the request
|
||||
obj.normalize_batch_and_arguments()
|
||||
self._set_default_priority(obj)
|
||||
self._validate_rid_not_in_flight(obj)
|
||||
|
||||
if isinstance(obj, GenerateReqInput) and obj.routed_dp_rank is not None:
|
||||
dp_size = self.server_args.dp_size
|
||||
@@ -552,7 +534,7 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerScoreMixin):
|
||||
f"routed_dp_rank={obj.routed_dp_rank} out of range [0, {dp_size})"
|
||||
)
|
||||
|
||||
self._req_stats_init(obj, request)
|
||||
self._init_req_state(obj, request)
|
||||
if self.server_args.language_only:
|
||||
self._handle_epd_disaggregation_encode_request(obj)
|
||||
if self.server_args.tokenizer_worker_num > 1:
|
||||
@@ -570,9 +552,8 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerScoreMixin):
|
||||
# Tokenize the request and send it to the scheduler
|
||||
if obj.is_single:
|
||||
tokenized_obj = await self._tokenize_one_request(obj)
|
||||
state = self.rid_to_state[obj.rid]
|
||||
self._send_one_request(tokenized_obj)
|
||||
async for response in self._wait_one_response(obj, state, request):
|
||||
async for response in self._wait_one_response(obj, request):
|
||||
yield response
|
||||
else:
|
||||
async for response in self._handle_batch_request(obj, request):
|
||||
@@ -825,17 +806,6 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerScoreMixin):
|
||||
obj, input_text, input_ids, input_embeds, mm_inputs, token_type_ids
|
||||
)
|
||||
|
||||
def _validate_rid_not_in_flight(
|
||||
self, obj: Union[GenerateReqInput, EmbeddingReqInput]
|
||||
) -> None:
|
||||
"""Validate that request IDs are not already in flight."""
|
||||
if obj.rid is None:
|
||||
return
|
||||
rids = obj.rid if isinstance(obj.rid, list) else [obj.rid]
|
||||
conflicts = set(rids) & self.rid_to_state.keys()
|
||||
if conflicts:
|
||||
raise ValueError(f"Duplicate request IDs detected: {list(conflicts)}")
|
||||
|
||||
def _validate_one_request(
|
||||
self, obj: Union[GenerateReqInput, EmbeddingReqInput], input_ids: List[int]
|
||||
) -> None:
|
||||
@@ -1204,13 +1174,90 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerScoreMixin):
|
||||
self.send_to_scheduler.send_pyobj(batch_req)
|
||||
set_time_batch(tokenized_objs, "set_api_server_dispatch_finish_time")
|
||||
|
||||
def _coalesce_streaming_chunks(
|
||||
self,
|
||||
out_list: list,
|
||||
rid: str,
|
||||
) -> dict:
|
||||
"""Coalesce multiple incremental streaming chunks into one.
|
||||
|
||||
Both text and output_ids are incremental deltas, so we concatenate them;
|
||||
all other fields (meta_info, etc.) are taken from the last chunk.
|
||||
"""
|
||||
if len(out_list) >= 20:
|
||||
logger.warning(
|
||||
"Streaming backlog: rid=%s, coalescing %d queued chunks into one. "
|
||||
"This may inflate P99 ITL for affected requests.",
|
||||
rid,
|
||||
len(out_list),
|
||||
)
|
||||
out = dict(out_list[-1])
|
||||
if "output_ids" in out:
|
||||
out["output_ids"] = [id for chunk in out_list for id in chunk["output_ids"]]
|
||||
if "text" in out:
|
||||
out["text"] = "".join(chunk["text"] for chunk in out_list)
|
||||
if "meta_info" in out:
|
||||
meta_info_list = [chunk["meta_info"] for chunk in out_list]
|
||||
meta_info = dict(meta_info_list[-1])
|
||||
for key in _INCREMENTAL_STREAMING_META_INFO_KEYS:
|
||||
if any(key in m for m in meta_info_list):
|
||||
meta_info[key] = [
|
||||
item for m in meta_info_list for item in m.get(key, [])
|
||||
]
|
||||
out["meta_info"] = meta_info
|
||||
return out
|
||||
|
||||
async def _handle_abort_finish_reason(
|
||||
self,
|
||||
out: dict,
|
||||
state: ReqState,
|
||||
is_stream: bool,
|
||||
) -> Optional[dict]:
|
||||
"""Handle abort/error finish reasons from the scheduler.
|
||||
|
||||
Returns the output dict if it should be yielded (stream abort), or None
|
||||
for normal flow. Raises ValueError or HTTPException for non-stream aborts.
|
||||
"""
|
||||
finish_reason = out["meta_info"]["finish_reason"]
|
||||
|
||||
if (
|
||||
finish_reason.get("type") == "abort"
|
||||
and finish_reason.get("status_code") == HTTPStatus.BAD_REQUEST
|
||||
):
|
||||
if not is_stream:
|
||||
raise ValueError(finish_reason["message"])
|
||||
return out
|
||||
|
||||
if finish_reason.get("type") == "abort" and finish_reason.get(
|
||||
"status_code"
|
||||
) in (
|
||||
HTTPStatus.SERVICE_UNAVAILABLE,
|
||||
HTTPStatus.INTERNAL_SERVER_ERROR,
|
||||
):
|
||||
# Delete the key to prevent resending abort request to the scheduler and
|
||||
# to ensure aborted request state is cleaned up.
|
||||
if state.obj.rid in self.rid_to_state:
|
||||
del self.rid_to_state[state.obj.rid]
|
||||
|
||||
# Mark ongoing LoRA request as finished.
|
||||
if self.server_args.enable_lora and state.obj.lora_path:
|
||||
await self.lora_registry.release(state.obj.lora_id)
|
||||
if not is_stream:
|
||||
raise fastapi.HTTPException(
|
||||
status_code=finish_reason["status_code"],
|
||||
detail=finish_reason["message"],
|
||||
)
|
||||
return out
|
||||
|
||||
return None
|
||||
|
||||
async def _wait_one_response(
|
||||
self,
|
||||
obj: Union[GenerateReqInput, EmbeddingReqInput],
|
||||
state: ReqState,
|
||||
request: Optional[fastapi.Request] = None,
|
||||
):
|
||||
"""Wait for the response of one request."""
|
||||
state = self.rid_to_state[obj.rid]
|
||||
# Not all request types have `stream` (e.g., EmbeddingReqInput). Default to non-streaming.
|
||||
is_stream = getattr(obj, "stream", False)
|
||||
while True:
|
||||
@@ -1233,38 +1280,18 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerScoreMixin):
|
||||
continue
|
||||
|
||||
# Drain all pending outputs atomically.
|
||||
# With incremental streaming output, each chunk carries only a
|
||||
# delta, so every queued chunk must be yielded to avoid dropping
|
||||
# token ids. Without it, outputs are cumulative and only the
|
||||
# latest chunk contains the full result, so we can safely skip
|
||||
# intermediate ones.
|
||||
incremental_stream = (
|
||||
is_stream and self.server_args.incremental_streaming_output
|
||||
)
|
||||
out_list = state.out_list
|
||||
state.out_list = []
|
||||
finished = state.finished
|
||||
state.event.clear()
|
||||
|
||||
# With incremental streaming, each chunk is a delta — coalesce
|
||||
# multiple queued chunks to avoid dropping token ids.
|
||||
incremental_stream = (
|
||||
is_stream and self.server_args.incremental_streaming_output
|
||||
)
|
||||
if incremental_stream and len(out_list) > 1:
|
||||
if len(out_list) >= 20:
|
||||
logger.warning(
|
||||
"Streaming backlog: rid=%s, coalescing %d queued chunks into one. "
|
||||
"This may inflate P99 ITL for affected requests.",
|
||||
obj.rid,
|
||||
len(out_list),
|
||||
)
|
||||
# Coalesce all deltas into a single chunk. Text, output_ids,
|
||||
# and output-side incremental metadata all need to be merged.
|
||||
out = dict(out_list[-1])
|
||||
if "output_ids" in out:
|
||||
out["output_ids"] = [
|
||||
id for chunk in out_list for id in chunk["output_ids"]
|
||||
]
|
||||
if "text" in out:
|
||||
out["text"] = "".join(chunk["text"] for chunk in out_list)
|
||||
if "meta_info" in out:
|
||||
out["meta_info"] = _merge_incremental_stream_meta_info(out_list)
|
||||
out = self._coalesce_streaming_chunks(out_list, obj.rid)
|
||||
else:
|
||||
out = out_list[-1]
|
||||
|
||||
@@ -1280,7 +1307,6 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerScoreMixin):
|
||||
out["text"] = state.get_text()
|
||||
|
||||
if finished:
|
||||
# For non-streaming cases, response has not been sent yet (`response_sent_to_client_time` has not been set yet).
|
||||
# Record response sent time right before we log finished results and metrics.
|
||||
if not state.time_stats.response_sent_to_client_time:
|
||||
state.time_stats.set_response_sent_to_client_time()
|
||||
@@ -1294,47 +1320,19 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerScoreMixin):
|
||||
)
|
||||
|
||||
if self.request_metrics_exporter_manager.exporter_enabled():
|
||||
# Asynchronously write metrics for this request using the exporter manager.
|
||||
asyncio.create_task(
|
||||
self.request_metrics_exporter_manager.write_record(obj, out)
|
||||
)
|
||||
|
||||
# Check if this was an abort/error created by scheduler
|
||||
if isinstance(out["meta_info"].get("finish_reason"), dict):
|
||||
finish_reason = out["meta_info"]["finish_reason"]
|
||||
if (
|
||||
finish_reason.get("type") == "abort"
|
||||
and finish_reason.get("status_code") == HTTPStatus.BAD_REQUEST
|
||||
):
|
||||
if not is_stream:
|
||||
raise ValueError(finish_reason["message"])
|
||||
else:
|
||||
yield out
|
||||
break
|
||||
abort_out = await self._handle_abort_finish_reason(
|
||||
out, state, is_stream
|
||||
)
|
||||
if abort_out is not None:
|
||||
yield abort_out
|
||||
break
|
||||
|
||||
if finish_reason.get("type") == "abort" and finish_reason.get(
|
||||
"status_code"
|
||||
) in (
|
||||
HTTPStatus.SERVICE_UNAVAILABLE,
|
||||
HTTPStatus.INTERNAL_SERVER_ERROR,
|
||||
):
|
||||
# This is an abort request initiated by scheduler.
|
||||
# Delete the key to prevent resending abort request to the scheduler and
|
||||
# to ensure aborted request state is cleaned up.
|
||||
if state.obj.rid in self.rid_to_state:
|
||||
del self.rid_to_state[state.obj.rid]
|
||||
|
||||
# Mark ongoing LoRA request as finished.
|
||||
if self.server_args.enable_lora and state.obj.lora_path:
|
||||
await self.lora_registry.release(state.obj.lora_id)
|
||||
if not is_stream:
|
||||
raise fastapi.HTTPException(
|
||||
status_code=finish_reason["status_code"],
|
||||
detail=finish_reason["message"],
|
||||
)
|
||||
else:
|
||||
yield out
|
||||
break
|
||||
yield out
|
||||
break
|
||||
|
||||
@@ -1346,8 +1344,7 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerScoreMixin):
|
||||
"response_sent_to_client_ts"
|
||||
] = state.time_stats.get_response_sent_to_client_realtime()
|
||||
yield out
|
||||
|
||||
if not is_stream:
|
||||
else:
|
||||
if (
|
||||
request is not None
|
||||
and not obj.background
|
||||
@@ -1377,9 +1374,7 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerScoreMixin):
|
||||
# Set up generators for each request in the batch
|
||||
for i in range(batch_size):
|
||||
tmp_obj = obj[i]
|
||||
state = self.rid_to_state[tmp_obj.rid]
|
||||
state.obj = tmp_obj
|
||||
generators.append(self._wait_one_response(tmp_obj, state, request))
|
||||
generators.append(self._wait_one_response(tmp_obj, request))
|
||||
rids.append(tmp_obj.rid)
|
||||
else:
|
||||
# Sequential tokenization and processing
|
||||
@@ -1391,12 +1386,8 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerScoreMixin):
|
||||
for i in range(batch_size):
|
||||
tmp_obj = obj[i]
|
||||
tokenized_obj = await self._tokenize_one_request(tmp_obj)
|
||||
state = self.rid_to_state[tmp_obj.rid]
|
||||
state.obj = tmp_obj
|
||||
self._send_one_request(tokenized_obj)
|
||||
generators.append(
|
||||
self._wait_one_response(tmp_obj, state, request)
|
||||
)
|
||||
generators.append(self._wait_one_response(tmp_obj, request))
|
||||
rids.append(tmp_obj.rid)
|
||||
else:
|
||||
# FIXME: When using batch and parallel_sample_num together, the perf is not optimal.
|
||||
@@ -1421,11 +1412,9 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerScoreMixin):
|
||||
tokenized_obj.sampling_params = copy.copy(tokenized_obj.sampling_params)
|
||||
tokenized_obj.sampling_params.max_new_tokens = 0
|
||||
tokenized_obj.stream = False
|
||||
self._req_stats_init(tmp_obj)
|
||||
state = self.rid_to_state[tmp_obj.rid]
|
||||
tokenized_obj.time_stats = state.time_stats
|
||||
self._init_req_state(tmp_obj)
|
||||
self._send_one_request(tokenized_obj)
|
||||
await self._wait_one_response(tmp_obj, state, request).__anext__()
|
||||
await self._wait_one_response(tmp_obj, request).__anext__()
|
||||
|
||||
# Expand requests, assign new rids for them, and send them
|
||||
for i in range(batch_size):
|
||||
@@ -1433,11 +1422,10 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerScoreMixin):
|
||||
tmp_obj = copy.copy(objs[i])
|
||||
tokenized_obj = copy.copy(tokenized_objs[i])
|
||||
tokenized_obj.rid = tmp_obj.regenerate_rid()
|
||||
self._req_stats_init(tmp_obj)
|
||||
state = self.rid_to_state[tmp_obj.rid]
|
||||
tokenized_obj.time_stats = state.time_stats
|
||||
self._init_req_state(tmp_obj)
|
||||
tokenized_obj.time_stats = self.rid_to_state[tmp_obj.rid].time_stats
|
||||
self._send_one_request(tokenized_obj)
|
||||
generators.append(self._wait_one_response(tmp_obj, state, request))
|
||||
generators.append(self._wait_one_response(tmp_obj, request))
|
||||
rids.append(tmp_obj.rid)
|
||||
|
||||
self.rid_to_state[objs[i].rid].time_stats.set_finished_time()
|
||||
@@ -1795,9 +1783,10 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerScoreMixin):
|
||||
state.time_stats.set_first_token_time()
|
||||
|
||||
if state.finished:
|
||||
state.time_stats.trace_ctx.trace_set_root_attrs(
|
||||
self.convert_to_span_attrs(state, recv_obj, i)
|
||||
)
|
||||
if state.time_stats.trace_ctx.tracing_enable:
|
||||
state.time_stats.trace_ctx.trace_set_root_attrs(
|
||||
self.convert_to_span_attrs(state, recv_obj, i)
|
||||
)
|
||||
state.time_stats.set_finished_time()
|
||||
meta_info["e2e_latency"] = state.time_stats.get_e2e_latency()
|
||||
|
||||
@@ -2459,12 +2448,11 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerScoreMixin):
|
||||
obj.lora_id[i] if isinstance(obj.lora_id, list) else obj.lora_id
|
||||
)
|
||||
|
||||
def _req_stats_init(
|
||||
def _init_req_state(
|
||||
self,
|
||||
obj: Union[GenerateReqInput, EmbeddingReqInput],
|
||||
request: Optional[fastapi.Request] = None,
|
||||
):
|
||||
calibrate_time_diff()
|
||||
created_time = obj.received_time
|
||||
|
||||
external_trace_header = None
|
||||
@@ -2478,39 +2466,32 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerScoreMixin):
|
||||
external_trace_header = extract_trace_headers(request.headers)
|
||||
obj.external_trace_header = external_trace_header
|
||||
|
||||
# Normalize single/batch into a uniform list of (rid, sub_obj, bootstrap_room)
|
||||
if not hasattr(obj, "is_single") or obj.is_single:
|
||||
time_stats = APIServerReqTimeStats(disagg_mode=self.disaggregation_mode)
|
||||
state = ReqState([], False, asyncio.Event(), obj, time_stats)
|
||||
self.rid_to_state[obj.rid] = state
|
||||
|
||||
if self.server_args.enable_trace:
|
||||
bootstrap_room = (
|
||||
obj.bootstrap_room if hasattr(obj, "bootstrap_room") else None
|
||||
)
|
||||
time_stats.init_trace_ctx(
|
||||
obj.rid,
|
||||
bootstrap_room,
|
||||
external_trace_header,
|
||||
)
|
||||
time_stats.set_created_time(created_time)
|
||||
items = [(obj.rid, obj, getattr(obj, "bootstrap_room", None))]
|
||||
else:
|
||||
for i in range(len(obj.rid)):
|
||||
time_stats = APIServerReqTimeStats(disagg_mode=self.disaggregation_mode)
|
||||
state = ReqState([], False, asyncio.Event(), obj[i], time_stats)
|
||||
self.rid_to_state[obj.rid[i]] = state
|
||||
|
||||
if self.server_args.enable_trace:
|
||||
bootstrap_room = (
|
||||
items = [
|
||||
(
|
||||
obj.rid[i],
|
||||
obj[i],
|
||||
(
|
||||
obj.bootstrap_room[i]
|
||||
if hasattr(obj, "bootstrap_room") and obj.bootstrap_room
|
||||
else None
|
||||
)
|
||||
time_stats.init_trace_ctx(
|
||||
obj.rid[i],
|
||||
bootstrap_room,
|
||||
external_trace_header,
|
||||
)
|
||||
time_stats.set_created_time(created_time)
|
||||
),
|
||||
)
|
||||
for i in range(len(obj.rid))
|
||||
]
|
||||
|
||||
for rid, sub_obj, bootstrap_room in items:
|
||||
if rid in self.rid_to_state:
|
||||
raise ValueError(f"Duplicate request ID detected: {rid}")
|
||||
time_stats = APIServerReqTimeStats(disagg_mode=self.disaggregation_mode)
|
||||
state = ReqState([], False, asyncio.Event(), sub_obj, time_stats)
|
||||
self.rid_to_state[rid] = state
|
||||
if self.server_args.enable_trace:
|
||||
time_stats.init_trace_ctx(rid, bootstrap_room, external_trace_header)
|
||||
time_stats.set_created_time(created_time)
|
||||
|
||||
def _should_dispatch_to_encoder(
|
||||
self, obj: Union[GenerateReqInput, EmbeddingReqInput]
|
||||
|
||||
Reference in New Issue
Block a user