Clean up req_time_stats: reduce overhead and simplify (#22186)

This commit is contained in:
Lianmin Zheng
2026-04-06 14:20:51 -07:00
committed by GitHub
parent 0bc4e0ea75
commit a80961333b
+87 -101
View File
@@ -55,12 +55,8 @@ def calibrate_time_diff():
global_diff_realtime_monotonic = time.time() - time.perf_counter()
def real_time():
return time.time()
def monotonic_time():
return time.perf_counter()
real_time = time.time
monotonic_time = time.perf_counter
def convert_time_to_realtime(time_value: float) -> float:
@@ -82,9 +78,19 @@ def convert_time_cross_thread(
@dataclass
class RequestStageConfig:
"""Configuration for a request pipeline stage.
Attributes:
stage_name: Name used for metrics labels and trace span names.
level: Trace hierarchy depth.
1 = leaf stages (atomic operations, e.g. TOKENIZE, PREFILL_FORWARD),
2 = parent/dispatch stages (e.g. API_SERVER_DISPATCH, REQUEST_PROCESS),
3 = composite/nested stages (e.g. DECODE_LOOP, PREFILL_CHUNKED_FORWARD).
metrics_is_observed: Whether to call metrics_collector.observe_per_stage_req_latency.
"""
stage_name: str
level: int = 0
# whether to call metrics_collector.observe_per_stage_req_latency
metrics_is_observed: bool = False
@@ -95,13 +101,13 @@ class RequestStage:
level=1,
)
API_SERVER_DISPATCH = RequestStageConfig(
"dispatch",
"api_server_dispatch",
level=2,
)
# DP controller
DC_DISPATCH = RequestStageConfig(
"dc_dispatch",
DPC_DISPATCH = RequestStageConfig(
"dpc_dispatch",
level=2,
)
@@ -210,7 +216,8 @@ class ReqTimeStatsBase:
if hasattr(new_obj, key):
setattr(new_obj, key, value)
new_obj.trace_ctx.rebuild_thread_context()
if new_obj.trace_ctx.tracing_enable:
new_obj.trace_ctx.rebuild_thread_context()
return new_obj
@@ -316,64 +323,60 @@ class APIServerReqTimeStats(ReqTimeStatsBase):
return state
def set_created_time(self, ts=None):
if ts is None:
ts = time.perf_counter()
ts = ts or time.perf_counter()
self.created_time = ts
self.trace_ctx.trace_req_start(convert_time_to_realtime_ns(ts))
if self.trace_ctx.tracing_enable:
self.trace_ctx.trace_req_start(convert_time_to_realtime_ns(ts))
def set_finished_time(self, ts=None):
if ts is None:
ts = time.perf_counter()
ts = ts or time.perf_counter()
self.finished_time = ts
self.trace_ctx.trace_req_finish(convert_time_to_realtime_ns(ts))
if self.trace_ctx.tracing_enable:
self.trace_ctx.trace_req_finish(convert_time_to_realtime_ns(ts))
def set_first_token_time(self, ts=None):
if ts is None:
ts = time.perf_counter()
ts = ts or time.perf_counter()
self.first_token_time = ts
self.last_time = ts
def set_last_time(self, ts=None):
if ts is None:
ts = time.perf_counter()
ts = ts or time.perf_counter()
self.last_time = ts
def set_tokenize_finish_time(self, ts=None):
if ts is None:
ts = time.perf_counter()
ts = ts or time.perf_counter()
self.tokenize_finish_time = ts
stage = RequestStage.TOKENIZE
self.trace_slice(stage, self.created_time, ts)
def set_api_server_dispatch_time(self, ts=None):
if ts is None:
ts = time.perf_counter()
ts = ts or time.perf_counter()
self.api_server_dispatch_time = ts
self.trace_ctx.trace_slice_start(
RequestStage.API_SERVER_DISPATCH.stage_name,
RequestStage.API_SERVER_DISPATCH.level,
convert_time_to_realtime_ns(ts),
)
if self.trace_ctx.tracing_enable:
self.trace_ctx.trace_slice_start(
RequestStage.API_SERVER_DISPATCH.stage_name,
RequestStage.API_SERVER_DISPATCH.level,
convert_time_to_realtime_ns(ts),
)
def set_api_server_dispatch_finish_time(self, ts=None):
if ts is None:
ts = time.perf_counter()
ts = ts or time.perf_counter()
self.api_server_dispatch_finish_time = ts
self.trace_ctx.trace_slice_end(
RequestStage.API_SERVER_DISPATCH.stage_name,
RequestStage.API_SERVER_DISPATCH.level,
convert_time_to_realtime_ns(ts),
thread_finish_flag=True,
)
if self.trace_ctx.tracing_enable:
self.trace_ctx.trace_slice_end(
RequestStage.API_SERVER_DISPATCH.stage_name,
RequestStage.API_SERVER_DISPATCH.level,
convert_time_to_realtime_ns(ts),
thread_finish_flag=True,
)
def set_response_sent_to_client_time(self, ts=None):
if ts is None:
ts = time.perf_counter()
ts = ts or time.perf_counter()
self.response_sent_to_client_time = ts
def get_interval(self):
@@ -463,8 +466,8 @@ class DPControllerReqTimeStats(ReqTimeStatsBase):
api_server_dispatch_time: float = 0.0
# new timestamp, get by time.perf_counter()
dc_dispatch_time: float = 0.0
dc_dispatch_finish_time: float = 0.0
dpc_dispatch_time: float = 0.0
dpc_dispatch_finish_time: float = 0.0
def __getstate__(self) -> object:
state = {}
@@ -473,33 +476,33 @@ class DPControllerReqTimeStats(ReqTimeStatsBase):
# state = {
# "created_time": self.created_time,
# "api_server_dispatch_time": self.api_server_dispatch_time,
# "dc_dispatch_time": self.dc_dispatch_time,
# "dpc_dispatch_time": self.dpc_dispatch_time,
# }
state.update(super().__getstate__())
return state
def set_dp_dispatch_time(self, ts=None):
if ts is None:
ts = time.perf_counter()
self.dc_dispatch_time = ts
ts = ts or time.perf_counter()
self.dpc_dispatch_time = ts
self.trace_ctx.trace_slice_start(
RequestStage.DC_DISPATCH.stage_name,
RequestStage.DC_DISPATCH.level,
convert_time_to_realtime_ns(ts),
)
if self.trace_ctx.tracing_enable:
self.trace_ctx.trace_slice_start(
RequestStage.DPC_DISPATCH.stage_name,
RequestStage.DPC_DISPATCH.level,
convert_time_to_realtime_ns(ts),
)
def set_dp_dispatch_finish_time(self, ts=None):
if ts is None:
ts = time.perf_counter()
self.dc_dispatch_finish_time = ts
ts = ts or time.perf_counter()
self.dpc_dispatch_finish_time = ts
self.trace_ctx.trace_slice_end(
RequestStage.DC_DISPATCH.stage_name,
RequestStage.DC_DISPATCH.level,
convert_time_to_realtime_ns(ts),
thread_finish_flag=True,
)
if self.trace_ctx.tracing_enable:
self.trace_ctx.trace_slice_end(
RequestStage.DPC_DISPATCH.stage_name,
RequestStage.DPC_DISPATCH.level,
convert_time_to_realtime_ns(ts),
thread_finish_flag=True,
)
@dataclass
@@ -516,7 +519,7 @@ class SchedulerReqTimeStats(ReqTimeStatsBase):
# propagated from tokenizer/grpc_server or dp controller
created_time: float = 0.0
api_server_dispatch_time: float = 0.0
dc_dispatch_time: float = 0.0
dpc_dispatch_time: float = 0.0
# common, get by time.perf_counter()
wait_queue_entry_time: float = 0.0
@@ -571,13 +574,11 @@ class SchedulerReqTimeStats(ReqTimeStatsBase):
def set_scheduler_recv_time(self, ts=None):
calibrate_time_diff()
if ts is None:
ts = time.perf_counter()
ts = ts or time.perf_counter()
self.scheduler_recv_time = ts
def set_retract_time(self, ts=None):
if ts is None:
ts = time.perf_counter()
ts = ts or time.perf_counter()
# retract
self.last_forward_entry_time = 0.0
self.last_prefill_finished_time = 0.0
@@ -585,11 +586,11 @@ class SchedulerReqTimeStats(ReqTimeStatsBase):
self.last_decode_finish_time = 0.0
self.last_decode_scheduled_time = 0.0
self.trace_ctx.trace_event("retract", 1, convert_time_to_realtime_ns(ts))
if self.trace_ctx.tracing_enable:
self.trace_ctx.trace_event("retract", 1, convert_time_to_realtime_ns(ts))
def set_wait_queue_entry_time(self, ts=None):
if ts is None:
ts = time.perf_counter()
ts = ts or time.perf_counter()
if self.wait_queue_entry_time == 0.0:
if self.enable_metrics or self.trace_ctx.tracing_enable:
if self.disagg_mode == DisaggregationMode.PREFILL:
@@ -610,8 +611,7 @@ class SchedulerReqTimeStats(ReqTimeStatsBase):
self.wait_queue_entry_time = ts
def set_forward_entry_time(self, ts=None):
if ts is None:
ts = time.perf_counter()
ts = ts or time.perf_counter()
if self.forward_entry_time == 0.0:
self.forward_entry_time = ts
self.last_forward_entry_time = ts
@@ -645,18 +645,15 @@ class SchedulerReqTimeStats(ReqTimeStatsBase):
self.last_forward_entry_time = ts
def set_prefill_run_batch_start_time(self, ts=None):
if ts is None:
ts = time.perf_counter()
ts = ts or time.perf_counter()
self.prefill_run_batch_start_time = ts
def set_prefill_run_batch_end_time(self, ts=None):
if ts is None:
ts = time.perf_counter()
ts = ts or time.perf_counter()
self.prefill_run_batch_end_time = ts
def set_last_chunked_prefill_finish_time(self, ts=None):
if ts is None:
ts = time.perf_counter()
ts = ts or time.perf_counter()
last_time = self.last_chunked_prefill_finish_time
self.last_chunked_prefill_finish_time = ts
@@ -668,8 +665,7 @@ class SchedulerReqTimeStats(ReqTimeStatsBase):
self.trace_slice(stage, last_time, ts)
def set_prefill_finished_time(self, ts=None):
if ts is None:
ts = time.perf_counter()
ts = ts or time.perf_counter()
if self.prefill_finished_time == 0.0:
self.prefill_finished_time = ts
self.last_prefill_finished_time = ts
@@ -712,8 +708,7 @@ class SchedulerReqTimeStats(ReqTimeStatsBase):
)
def set_last_decode_finish_time(self, ts=None):
if ts is None:
ts = time.perf_counter()
ts = ts or time.perf_counter()
last_time = self.last_decode_finish_time
self.last_decode_finish_time = ts
@@ -736,8 +731,7 @@ class SchedulerReqTimeStats(ReqTimeStatsBase):
self.decode_ct += 1
def set_last_scheduled_time(self, forward_mode: ForwardMode, ts=None, attrs=None):
if ts is None:
ts = time.perf_counter()
ts = ts or time.perf_counter()
if self.trace_ctx.tracing_enable:
if (
@@ -764,11 +758,11 @@ class SchedulerReqTimeStats(ReqTimeStatsBase):
self.last_decode_scheduled_time = ts
def set_completion_time(self, ts=None):
if ts is None:
ts = time.perf_counter()
ts = ts or time.perf_counter()
self.completion_time = ts
self.trace_ctx.abort()
if self.trace_ctx.tracing_enable:
self.trace_ctx.abort()
def compute_and_observe_kv_transfer_metrics(
self,
@@ -835,14 +829,12 @@ class SchedulerReqTimeStats(ReqTimeStatsBase):
return result if result else None
def set_quick_finish_time(self, ts=None):
if ts is None:
ts = time.perf_counter()
ts = ts or time.perf_counter()
self.set_completion_time(ts)
self.forward_entry_time = ts
def set_prefill_bootstrap_queue_entry_time(self, ts=None):
if ts is None:
ts = time.perf_counter()
ts = ts or time.perf_counter()
self.prefill_bootstrap_queue_entry_time = ts
stage = RequestStage.PREFILL_PREPARE
@@ -850,13 +842,11 @@ class SchedulerReqTimeStats(ReqTimeStatsBase):
self.trace_slice(stage, self.scheduler_recv_time, ts)
def set_prefill_transfer_queue_entry_time(self, ts=None):
if ts is None:
ts = time.perf_counter()
ts = ts or time.perf_counter()
self.prefill_transfer_queue_entry_time = ts
def set_prefill_kv_transfer_finish_time(self, ts=None):
if ts is None:
ts = time.perf_counter()
ts = ts or time.perf_counter()
self.prefill_kv_transfer_finish_time = ts
stage = RequestStage.PREFILL_TRANSFER_KV_CACHE
@@ -866,8 +856,7 @@ class SchedulerReqTimeStats(ReqTimeStatsBase):
self.trace_slice(stage, self.prefill_transfer_queue_entry_time, ts)
def set_decode_prealloc_queue_entry_time(self, ts=None):
if ts is None:
ts = time.perf_counter()
ts = ts or time.perf_counter()
self.decode_prealloc_queue_entry_time = ts
stage = RequestStage.DECODE_PREPARE
@@ -875,8 +864,7 @@ class SchedulerReqTimeStats(ReqTimeStatsBase):
self.trace_slice(stage, self.scheduler_recv_time, ts)
def set_decode_transfer_queue_entry_time(self, ts=None):
if ts is None:
ts = time.perf_counter()
ts = ts or time.perf_counter()
self.decode_transfer_queue_entry_time = ts
stage = RequestStage.DECODE_BOOTSTRAP
@@ -886,14 +874,12 @@ class SchedulerReqTimeStats(ReqTimeStatsBase):
self.trace_slice(stage, self.decode_prealloc_queue_entry_time, ts)
def set_bootstrap_done_time(self, ts=None):
if ts is None:
ts = time.perf_counter()
ts = ts or time.perf_counter()
if self.bootstrap_done_time == 0.0:
self.bootstrap_done_time = ts
def set_decode_prebuilt_finish_time(self, ts=None):
if ts is None:
ts = time.perf_counter()
ts = ts or time.perf_counter()
self.decode_prebuilt_finish_time = ts
stage = RequestStage.DECODE_FAKE_OUTPUT