diff --git a/python/sglang/srt/observability/req_time_stats.py b/python/sglang/srt/observability/req_time_stats.py index 8caf21c32..aa067152f 100644 --- a/python/sglang/srt/observability/req_time_stats.py +++ b/python/sglang/srt/observability/req_time_stats.py @@ -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