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
+129 -148
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.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]