diff --git a/python/sglang/srt/model_executor/model_runner.py b/python/sglang/srt/model_executor/model_runner.py index 856aec750..03bffbf4f 100644 --- a/python/sglang/srt/model_executor/model_runner.py +++ b/python/sglang/srt/model_executor/model_runner.py @@ -201,6 +201,7 @@ from sglang.srt.utils import ( set_cuda_arch, slow_rank_detector, ) +from sglang.srt.utils.device_timer import device_timer_ctx from sglang.srt.utils.nvtx_pytorch_hooks import PytHooks from sglang.srt.utils.nvtx_utils import profile_range from sglang.srt.utils.offloader import ( @@ -1319,12 +1320,7 @@ class ModelRunner: forward_batch.split_index + forward_count, self.model_config.num_hidden_layers, ) - ctx = ( - self.device_timer.wrap(metadata={"category": "split_prefill"}) - if self.device_timer - else contextlib.nullcontext() - ) - with ctx: + with device_timer_ctx(self.device_timer, "split_prefill"): ret = self.model.forward_split_prefill( forward_batch.input_ids, forward_batch.positions, @@ -1547,22 +1543,17 @@ class ModelRunner: and self.prefill_cuda_graph_runner.can_run_graph(forward_batch) and get_cp_strategy() is None ): + # Prefill cuda graph (piecewise). + kwargs = self._extend_forward_kwargs(forward_batch, pp_proxy_tensors) category = ( "target_verify" if forward_batch.forward_mode.is_target_verify() else "extend" ) - # Prefill cuda graph (piecewise). - kwargs = self._extend_forward_kwargs(forward_batch, pp_proxy_tensors) - # TODO: device_timer.wrap is too broad here — it also includes - # load_batch time. Move timing into the prefill cuda graph runner + # TODO: the timing here is too broad -- it also includes + # load_batch time. Move it into the prefill cuda graph runner # to capture only the model.forward part. - ctx = ( - self.device_timer.wrap(metadata={"category": category}) - if self.device_timer - else contextlib.nullcontext() - ) - with ctx: + with device_timer_ctx(self.device_timer, category): ret = self.prefill_cuda_graph_runner.execute( forward_batch, **kwargs ) diff --git a/python/sglang/srt/model_executor/runner/decode_cuda_graph_runner.py b/python/sglang/srt/model_executor/runner/decode_cuda_graph_runner.py index 60f4cab53..9da4ba737 100644 --- a/python/sglang/srt/model_executor/runner/decode_cuda_graph_runner.py +++ b/python/sglang/srt/model_executor/runner/decode_cuda_graph_runner.py @@ -99,6 +99,7 @@ from sglang.srt.utils import ( require_attn_tp_gather, require_mlp_tp_gather, ) +from sglang.srt.utils.device_timer import device_timer_ctx from sglang.srt.utils.profile_utils import export_cuda_graph_capture_trace try: @@ -1212,12 +1213,8 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner): forward_batch: ForwardBatch, pp_proxy_tensors: Optional[PPProxyTensors] = None, ) -> Union[LogitsProcessorOutput, PPProxyTensors]: - timer_ctx = ( - self.model_runner.device_timer.wrap( - metadata={"category": forward_batch.forward_mode.name.lower()} - ) - if self.model_runner.device_timer - else contextlib.nullcontext() + timer_ctx = device_timer_ctx( + self.model_runner.device_timer, forward_batch.forward_mode.name.lower() ) # Publish a read-done event for the WAR barrier: a cuda-graph forward # finishes its shared req_to_token / SWA reads at this pre-replay diff --git a/python/sglang/srt/model_executor/runner/eager_runner.py b/python/sglang/srt/model_executor/runner/eager_runner.py index 55d40dc87..ff1c300e9 100644 --- a/python/sglang/srt/model_executor/runner/eager_runner.py +++ b/python/sglang/srt/model_executor/runner/eager_runner.py @@ -55,6 +55,7 @@ from sglang.srt.utils.common import ( get_eager_max_batch_size, require_mlp_sync, ) +from sglang.srt.utils.device_timer import device_timer_ctx logger = logging.getLogger(__name__) @@ -237,11 +238,7 @@ class EagerRunner(BaseRunner): # FIXME: add pp_proxy_tensors arg to all models kwargs = model_runner._pp_kwargs(pp_proxy_tensors) - ctx = ( - model_runner.device_timer.wrap(metadata={"category": "decode"}) - if model_runner.device_timer - else contextlib.nullcontext() - ) + ctx = device_timer_ctx(model_runner.device_timer, "decode") with ctx, pdmux_ctx: return model_runner.model.forward( @@ -297,12 +294,7 @@ class EagerRunner(BaseRunner): if forward_batch.forward_mode.is_target_verify() else "extend" ) - ctx = ( - model_runner.device_timer.wrap(metadata={"category": category}) - if model_runner.device_timer - else contextlib.nullcontext() - ) - with ctx: + with device_timer_ctx(model_runner.device_timer, category): pcg_runner = model_runner.prefill_cuda_graph_runner if ( _is_hip @@ -401,12 +393,7 @@ class EagerRunner(BaseRunner): model_runner.attn_backend.forward_metadata = None kwargs = model_runner._pp_kwargs(pp_proxy_tensors) - ctx = ( - model_runner.device_timer.wrap(metadata={"category": "idle"}) - if model_runner.device_timer - else contextlib.nullcontext() - ) - with ctx: + with device_timer_ctx(model_runner.device_timer, "idle"): return model_runner.model.forward( forward_batch.input_ids, forward_batch.positions, diff --git a/python/sglang/srt/speculative/eagle_draft_cuda_graph_runner.py b/python/sglang/srt/speculative/eagle_draft_cuda_graph_runner.py index 0e47b2546..8de99d832 100644 --- a/python/sglang/srt/speculative/eagle_draft_cuda_graph_runner.py +++ b/python/sglang/srt/speculative/eagle_draft_cuda_graph_runner.py @@ -1,6 +1,5 @@ from __future__ import annotations -import contextlib from dataclasses import dataclass from typing import TYPE_CHECKING, Callable, Optional @@ -47,6 +46,7 @@ from sglang.srt.utils import ( require_mlp_tp_gather, ) from sglang.srt.utils.async_probe import maybe_detect_nan, maybe_detect_oob +from sglang.srt.utils.device_timer import device_timer_ctx if TYPE_CHECKING: from sglang.srt.speculative.eagle_worker_v2 import EagleDraftWorker @@ -653,12 +653,7 @@ class EAGLEDraftCudaGraphRunner(DecodeCudaGraphRunner): # Replay via backend shape_key = self._make_graph_key(bs) - timer_ctx = ( - self.model_runner.device_timer.wrap(metadata={"category": "eagle_draft"}) - if self.model_runner.device_timer - else contextlib.nullcontext() - ) - with timer_ctx: + with device_timer_ctx(self.model_runner.device_timer, "eagle_draft"): out = self._replay_graph(shape_key, forward_batch) if self.buffers.dsa_seed_topk is not None: forward_batch.spec_info.dsa_topk_indices = None diff --git a/python/sglang/srt/speculative/eagle_draft_extend_cuda_graph_runner.py b/python/sglang/srt/speculative/eagle_draft_extend_cuda_graph_runner.py index 1eea9b969..c4b7f92cc 100644 --- a/python/sglang/srt/speculative/eagle_draft_extend_cuda_graph_runner.py +++ b/python/sglang/srt/speculative/eagle_draft_extend_cuda_graph_runner.py @@ -49,6 +49,7 @@ from sglang.srt.utils import ( require_mlp_sync, require_mlp_tp_gather, ) +from sglang.srt.utils.device_timer import device_timer_ctx _is_hip = is_hip() @@ -615,14 +616,7 @@ class EAGLEDraftExtendCudaGraphRunner(DecodeCudaGraphRunner): self.raw_bs = raw_bs self.bs = bs shape_key = self._make_graph_key(bs) - timer_ctx = ( - self.model_runner.device_timer.wrap( - metadata={"category": "eagle_draft_extend"} - ) - if self.model_runner.device_timer - else contextlib.nullcontext() - ) - with timer_ctx: + with device_timer_ctx(self.model_runner.device_timer, "eagle_draft_extend"): out = self._replay_graph(shape_key, forward_batch) out = LogitsProcessorOutput( diff --git a/python/sglang/srt/speculative/frozen_kv_mtp_cuda_graph_runner.py b/python/sglang/srt/speculative/frozen_kv_mtp_cuda_graph_runner.py index 8eb3d18bc..0cce8c07f 100644 --- a/python/sglang/srt/speculative/frozen_kv_mtp_cuda_graph_runner.py +++ b/python/sglang/srt/speculative/frozen_kv_mtp_cuda_graph_runner.py @@ -41,6 +41,7 @@ from sglang.srt.utils import ( require_mlp_sync, require_mlp_tp_gather, ) +from sglang.srt.utils.device_timer import device_timer_ctx if TYPE_CHECKING: from sglang.srt.speculative.frozen_kv_mtp_worker_v2 import FrozenKVMTPDraftWorker @@ -446,11 +447,12 @@ class FrozenKVMTPCudaGraphRunner(DecodeCudaGraphRunner): shape_key = self._make_graph_key(bs) # NVTX span: the graph bypasses `model_runner.forward`'s record_function. span_name = f"step[DRAFT_LOOP raw_bs={raw_bs} bs={bs} topk={self.topk}]" - if torch.autograd._profiler_enabled(): - with torch.profiler.record_function(span_name): + with device_timer_ctx(self.model_runner.device_timer, "frozen_kv_draft"): + if torch.autograd._profiler_enabled(): + with torch.profiler.record_function(span_name): + out = self._replay_graph(shape_key, forward_batch) + else: out = self._replay_graph(shape_key, forward_batch) - else: - out = self._replay_graph(shape_key, forward_batch) if bs != raw_bs: out = self._postprocess_output_to_raw_bs(out, raw_bs) diff --git a/python/sglang/srt/speculative/multi_layer_eagle_draft_extend_cuda_graph_runner.py b/python/sglang/srt/speculative/multi_layer_eagle_draft_extend_cuda_graph_runner.py index f4bd3fb16..54844777a 100644 --- a/python/sglang/srt/speculative/multi_layer_eagle_draft_extend_cuda_graph_runner.py +++ b/python/sglang/srt/speculative/multi_layer_eagle_draft_extend_cuda_graph_runner.py @@ -82,6 +82,7 @@ from sglang.srt.utils import ( require_mlp_sync, require_mlp_tp_gather, ) +from sglang.srt.utils.device_timer import device_timer_ctx if is_npu(): from sglang.srt.speculative.multi_layer_eagle_utils import ( @@ -498,7 +499,8 @@ class MultiLayerEagleDraftExtendCudaGraphRunner(DecodeCudaGraphRunner): self.bs = bs shape_key = self._make_graph_key(bs) - return self._replay_graph(shape_key, fb_view) + with device_timer_ctx(self.model_runner.device_timer, "eagle_draft_extend"): + return self._replay_graph(shape_key, fb_view) class MultiLayerEagleMultiStepDraftExtendCudaGraphRunner: @@ -973,7 +975,10 @@ class OneGraphMultiLayerEagleMultiStepDraftExtendCudaGraphRunner( if r is not None: r.deepep_adapter.replay() shape_key = first._make_graph_key(self.bs) - outs = first.backend.replay(shape_key, self._replay_spec_info) + with device_timer_ctx( + first.model_runner.device_timer, "eagle_draft_extend" + ): + outs = first.backend.replay(shape_key, self._replay_spec_info) raw_bs = self.raw_bs self._cached = {} non_null = [r for r in self.runners if r is not None] diff --git a/python/sglang/srt/utils/device_timer.py b/python/sglang/srt/utils/device_timer.py index 3562e4df3..5ef5e14ad 100644 --- a/python/sglang/srt/utils/device_timer.py +++ b/python/sglang/srt/utils/device_timer.py @@ -1,26 +1,44 @@ from collections import deque -from contextlib import contextmanager +from contextlib import contextmanager, nullcontext from dataclasses import dataclass from typing import Callable, Deque, Dict, List, Optional import torch +def device_timer_ctx(timer: Optional["DeviceTimer"], category: str): + """Timing context for one forward segment; no-op when the timer is absent. + + A segment that skips this stays out of the fwd_occupancy numerator while + still counting in its wall-clock denominator, i.e. reads as GPU idle. + """ + if timer is None: + return nullcontext() + return timer.wrap(metadata={"category": category}) + + class DeviceTimer: def __init__(self, reporter: Callable): self._intervals: Deque[_TimingInterval] = deque() self._reporters: List[Callable] = [reporter] + self._in_wrap = False def add_reporter(self, reporter: Callable): self._reporters.append(reporter) @contextmanager def wrap(self, metadata: Dict): - self._intervals.append(_TimingInterval.create()) + # Not re-entrant: a nested wrap would end the wrong interval and leave + # an un-ended one at the head of the queue for _report() to trip over. + assert not self._in_wrap, "DeviceTimer.wrap is not re-entrant" + interval = _TimingInterval.create() + self._intervals.append(interval) + self._in_wrap = True try: yield finally: - self._intervals[-1].end(metadata=metadata) + self._in_wrap = False + interval.end(metadata=metadata) self._report() def _report(self):