From 9fb00ede15d5297f08aaa61485b08ae4ddf959f3 Mon Sep 17 00:00:00 2001 From: Lianmin Zheng Date: Mon, 13 Apr 2026 16:47:32 -0700 Subject: [PATCH] Clean up TokenizerManager and req_time_stats: reduce overhead and simplify (#21646) --- .../sglang/srt/managers/tokenizer_manager.py | 277 ++++++++---------- 1 file changed, 129 insertions(+), 148 deletions(-) diff --git a/python/sglang/srt/managers/tokenizer_manager.py b/python/sglang/srt/managers/tokenizer_manager.py index bfc5bef63..09a9df885 100644 --- a/python/sglang/srt/managers/tokenizer_manager.py +++ b/python/sglang/srt/managers/tokenizer_manager.py @@ -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]