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.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()
|
||||||
|
|
||||||
|
# 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 incremental_stream and len(out_list) > 1:
|
||||||
if len(out_list) >= 20:
|
out = self._coalesce_streaming_chunks(out_list, obj.rid)
|
||||||
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)
|
|
||||||
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 abort_out is not None:
|
||||||
):
|
yield abort_out
|
||||||
if not is_stream:
|
break
|
||||||
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:
|
|
||||||
yield out
|
|
||||||
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,9 +1783,10 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerScoreMixin):
|
|||||||
state.time_stats.set_first_token_time()
|
state.time_stats.set_first_token_time()
|
||||||
|
|
||||||
if state.finished:
|
if state.finished:
|
||||||
state.time_stats.trace_ctx.trace_set_root_attrs(
|
if state.time_stats.trace_ctx.tracing_enable:
|
||||||
self.convert_to_span_attrs(state, recv_obj, i)
|
state.time_stats.trace_ctx.trace_set_root_attrs(
|
||||||
)
|
self.convert_to_span_attrs(state, recv_obj, i)
|
||||||
|
)
|
||||||
state.time_stats.set_finished_time()
|
state.time_stats.set_finished_time()
|
||||||
meta_info["e2e_latency"] = state.time_stats.get_e2e_latency()
|
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
|
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,39 +2466,32 @@ 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(
|
)
|
||||||
obj.rid[i],
|
for i in range(len(obj.rid))
|
||||||
bootstrap_room,
|
]
|
||||||
external_trace_header,
|
|
||||||
)
|
for rid, sub_obj, bootstrap_room in items:
|
||||||
time_stats.set_created_time(created_time)
|
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(
|
def _should_dispatch_to_encoder(
|
||||||
self, obj: Union[GenerateReqInput, EmbeddingReqInput]
|
self, obj: Union[GenerateReqInput, EmbeddingReqInput]
|
||||||
|
|||||||
Reference in New Issue
Block a user