Clean up TokenizerManager and req_time_stats: reduce overhead and simplify (#21646)

This commit is contained in:
Lianmin Zheng
2026-04-13 16:47:32 -07:00
committed by GitHub
parent a2b5111962
commit 9fb00ede15
+123 -142
View File
@@ -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.metrics_collector import TokenizerMetricsCollector
from sglang.srt.observability.req_time_stats import ( from sglang.srt.observability.req_time_stats import (
APIServerReqTimeStats, APIServerReqTimeStats,
calibrate_time_diff,
convert_time_to_realtime, convert_time_to_realtime,
real_time, real_time,
set_time_batch, set_time_batch,
@@ -205,22 +204,6 @@ def _slice_streaming_output_meta_info(
meta_info[key] = meta_info[key][last_output_offset:] 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): class InputFormat(Enum):
"""Input format types for tokenization handling.""" """Input format types for tokenization handling."""
@@ -268,9 +251,6 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerScoreMixin):
# Init PD disaggregation and encoder disaggregation # Init PD disaggregation and encoder disaggregation
self.init_disaggregation() self.init_disaggregation()
# Subprocess liveness watchdog — set by Engine or http_server after construction
self._subprocess_watchdog = None
# Init metric collector and watchdog # Init metric collector and watchdog
self.init_metric_collector_watchdog() self.init_metric_collector_watchdog()
@@ -395,6 +375,9 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerScoreMixin):
# Session # Session
self.session_futures = {} # session_id -> asyncio event 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): def init_request_logging_and_dumping(self):
# TODO: Refactor and organize the log export code. # TODO: Refactor and organize the log export code.
# Request logging # Request logging
@@ -539,7 +522,6 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerScoreMixin):
# Normalize the request # Normalize the request
obj.normalize_batch_and_arguments() obj.normalize_batch_and_arguments()
self._set_default_priority(obj) self._set_default_priority(obj)
self._validate_rid_not_in_flight(obj)
if isinstance(obj, GenerateReqInput) and obj.routed_dp_rank is not None: if isinstance(obj, GenerateReqInput) and obj.routed_dp_rank is not None:
dp_size = self.server_args.dp_size 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})" 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: if self.server_args.language_only:
self._handle_epd_disaggregation_encode_request(obj) self._handle_epd_disaggregation_encode_request(obj)
if self.server_args.tokenizer_worker_num > 1: 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 # Tokenize the request and send it to the scheduler
if obj.is_single: if obj.is_single:
tokenized_obj = await self._tokenize_one_request(obj) tokenized_obj = await self._tokenize_one_request(obj)
state = self.rid_to_state[obj.rid]
self._send_one_request(tokenized_obj) 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 yield response
else: else:
async for response in self._handle_batch_request(obj, request): 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 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( def _validate_one_request(
self, obj: Union[GenerateReqInput, EmbeddingReqInput], input_ids: List[int] self, obj: Union[GenerateReqInput, EmbeddingReqInput], input_ids: List[int]
) -> None: ) -> None:
@@ -1204,13 +1174,90 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerScoreMixin):
self.send_to_scheduler.send_pyobj(batch_req) self.send_to_scheduler.send_pyobj(batch_req)
set_time_batch(tokenized_objs, "set_api_server_dispatch_finish_time") 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( async def _wait_one_response(
self, self,
obj: Union[GenerateReqInput, EmbeddingReqInput], obj: Union[GenerateReqInput, EmbeddingReqInput],
state: ReqState,
request: Optional[fastapi.Request] = None, request: Optional[fastapi.Request] = None,
): ):
"""Wait for the response of one request.""" """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. # Not all request types have `stream` (e.g., EmbeddingReqInput). Default to non-streaming.
is_stream = getattr(obj, "stream", False) is_stream = getattr(obj, "stream", False)
while True: while True:
@@ -1233,38 +1280,18 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerScoreMixin):
continue continue
# Drain all pending outputs atomically. # 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 out_list = state.out_list
state.out_list = [] state.out_list = []
finished = state.finished finished = state.finished
state.event.clear() state.event.clear()
if incremental_stream and len(out_list) > 1: # With incremental streaming, each chunk is a delta — coalesce
if len(out_list) >= 20: # multiple queued chunks to avoid dropping token ids.
logger.warning( incremental_stream = (
"Streaming backlog: rid=%s, coalescing %d queued chunks into one. " is_stream and self.server_args.incremental_streaming_output
"This may inflate P99 ITL for affected requests.",
obj.rid,
len(out_list),
) )
# Coalesce all deltas into a single chunk. Text, output_ids, if incremental_stream and len(out_list) > 1:
# and output-side incremental metadata all need to be merged. out = self._coalesce_streaming_chunks(out_list, obj.rid)
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)
else: else:
out = out_list[-1] out = out_list[-1]
@@ -1280,7 +1307,6 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerScoreMixin):
out["text"] = state.get_text() out["text"] = state.get_text()
if finished: 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. # Record response sent time right before we log finished results and metrics.
if not state.time_stats.response_sent_to_client_time: if not state.time_stats.response_sent_to_client_time:
state.time_stats.set_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(): if self.request_metrics_exporter_manager.exporter_enabled():
# Asynchronously write metrics for this request using the exporter manager.
asyncio.create_task( asyncio.create_task(
self.request_metrics_exporter_manager.write_record(obj, out) self.request_metrics_exporter_manager.write_record(obj, out)
) )
# Check if this was an abort/error created by scheduler # Check if this was an abort/error created by scheduler
if isinstance(out["meta_info"].get("finish_reason"), dict): if isinstance(out["meta_info"].get("finish_reason"), dict):
finish_reason = out["meta_info"]["finish_reason"] abort_out = await self._handle_abort_finish_reason(
if ( out, state, is_stream
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
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: if abort_out is not None:
yield out yield abort_out
break break
yield out yield out
break break
@@ -1346,8 +1344,7 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerScoreMixin):
"response_sent_to_client_ts" "response_sent_to_client_ts"
] = state.time_stats.get_response_sent_to_client_realtime() ] = state.time_stats.get_response_sent_to_client_realtime()
yield out yield out
else:
if not is_stream:
if ( if (
request is not None request is not None
and not obj.background and not obj.background
@@ -1377,9 +1374,7 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerScoreMixin):
# Set up generators for each request in the batch # Set up generators for each request in the batch
for i in range(batch_size): for i in range(batch_size):
tmp_obj = obj[i] tmp_obj = obj[i]
state = self.rid_to_state[tmp_obj.rid] generators.append(self._wait_one_response(tmp_obj, request))
state.obj = tmp_obj
generators.append(self._wait_one_response(tmp_obj, state, request))
rids.append(tmp_obj.rid) rids.append(tmp_obj.rid)
else: else:
# Sequential tokenization and processing # Sequential tokenization and processing
@@ -1391,12 +1386,8 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerScoreMixin):
for i in range(batch_size): for i in range(batch_size):
tmp_obj = obj[i] tmp_obj = obj[i]
tokenized_obj = await self._tokenize_one_request(tmp_obj) 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) self._send_one_request(tokenized_obj)
generators.append( generators.append(self._wait_one_response(tmp_obj, request))
self._wait_one_response(tmp_obj, state, request)
)
rids.append(tmp_obj.rid) rids.append(tmp_obj.rid)
else: else:
# FIXME: When using batch and parallel_sample_num together, the perf is not optimal. # 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 = copy.copy(tokenized_obj.sampling_params)
tokenized_obj.sampling_params.max_new_tokens = 0 tokenized_obj.sampling_params.max_new_tokens = 0
tokenized_obj.stream = False tokenized_obj.stream = False
self._req_stats_init(tmp_obj) self._init_req_state(tmp_obj)
state = self.rid_to_state[tmp_obj.rid]
tokenized_obj.time_stats = state.time_stats
self._send_one_request(tokenized_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 # Expand requests, assign new rids for them, and send them
for i in range(batch_size): for i in range(batch_size):
@@ -1433,11 +1422,10 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerScoreMixin):
tmp_obj = copy.copy(objs[i]) tmp_obj = copy.copy(objs[i])
tokenized_obj = copy.copy(tokenized_objs[i]) tokenized_obj = copy.copy(tokenized_objs[i])
tokenized_obj.rid = tmp_obj.regenerate_rid() tokenized_obj.rid = tmp_obj.regenerate_rid()
self._req_stats_init(tmp_obj) self._init_req_state(tmp_obj)
state = self.rid_to_state[tmp_obj.rid] tokenized_obj.time_stats = self.rid_to_state[tmp_obj.rid].time_stats
tokenized_obj.time_stats = state.time_stats
self._send_one_request(tokenized_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) rids.append(tmp_obj.rid)
self.rid_to_state[objs[i].rid].time_stats.set_finished_time() self.rid_to_state[objs[i].rid].time_stats.set_finished_time()
@@ -1795,6 +1783,7 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerScoreMixin):
state.time_stats.set_first_token_time() state.time_stats.set_first_token_time()
if state.finished: if state.finished:
if state.time_stats.trace_ctx.tracing_enable:
state.time_stats.trace_ctx.trace_set_root_attrs( state.time_stats.trace_ctx.trace_set_root_attrs(
self.convert_to_span_attrs(state, recv_obj, i) self.convert_to_span_attrs(state, recv_obj, i)
) )
@@ -2459,12 +2448,11 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerScoreMixin):
obj.lora_id[i] if isinstance(obj.lora_id, list) else obj.lora_id obj.lora_id[i] if isinstance(obj.lora_id, list) else obj.lora_id
) )
def _req_stats_init( def _init_req_state(
self, self,
obj: Union[GenerateReqInput, EmbeddingReqInput], obj: Union[GenerateReqInput, EmbeddingReqInput],
request: Optional[fastapi.Request] = None, request: Optional[fastapi.Request] = None,
): ):
calibrate_time_diff()
created_time = obj.received_time created_time = obj.received_time
external_trace_header = None external_trace_header = None
@@ -2478,38 +2466,31 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerScoreMixin):
external_trace_header = extract_trace_headers(request.headers) external_trace_header = extract_trace_headers(request.headers)
obj.external_trace_header = external_trace_header 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: if not hasattr(obj, "is_single") or obj.is_single:
time_stats = APIServerReqTimeStats(disagg_mode=self.disaggregation_mode) items = [(obj.rid, obj, getattr(obj, "bootstrap_room", None))]
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)
else: else:
for i in range(len(obj.rid)): items = [
time_stats = APIServerReqTimeStats(disagg_mode=self.disaggregation_mode) (
state = ReqState([], False, asyncio.Event(), obj[i], time_stats) obj.rid[i],
self.rid_to_state[obj.rid[i]] = state obj[i],
(
if self.server_args.enable_trace:
bootstrap_room = (
obj.bootstrap_room[i] obj.bootstrap_room[i]
if hasattr(obj, "bootstrap_room") and obj.bootstrap_room if hasattr(obj, "bootstrap_room") and obj.bootstrap_room
else None else None
),
) )
time_stats.init_trace_ctx( for i in range(len(obj.rid))
obj.rid[i], ]
bootstrap_room,
external_trace_header, 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) time_stats.set_created_time(created_time)
def _should_dispatch_to_encoder( def _should_dispatch_to_encoder(