From eb646c7b784217d93c75bd4ab1becf86ddf85dd2 Mon Sep 17 00:00:00 2001 From: Leon Gao Date: Mon, 8 Jun 2026 08:41:54 -0700 Subject: [PATCH] [srt] Add sglang:weight_load_duration_seconds gauge with source label (#27363) --- python/sglang/srt/managers/scheduler.py | 1 + .../scheduler_components/weight_updater.py | 96 ++++++++++++------- .../srt/observability/metrics_collector.py | 23 +++++ 3 files changed, 83 insertions(+), 37 deletions(-) diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index 701e0abc6..ea14ee492 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -1578,6 +1578,7 @@ class Scheduler( memory_saver_adapter=self.memory_saver_adapter, flush_cache=self.flush_cache, is_fully_idle=self.is_fully_idle, + metrics_collector=self.metrics_collector, ) def init_lora_drainer(self) -> None: diff --git a/python/sglang/srt/managers/scheduler_components/weight_updater.py b/python/sglang/srt/managers/scheduler_components/weight_updater.py index 59f6597b3..c546a5650 100644 --- a/python/sglang/srt/managers/scheduler_components/weight_updater.py +++ b/python/sglang/srt/managers/scheduler_components/weight_updater.py @@ -2,9 +2,11 @@ from __future__ import annotations import hashlib import logging +import time import traceback +from contextlib import contextmanager from dataclasses import dataclass, field -from typing import Any, Callable, Dict, Tuple +from typing import Any, Callable, Dict, Iterator, Optional, Tuple import torch @@ -75,9 +77,25 @@ class SchedulerWeightUpdaterManager: memory_saver_adapter: Any flush_cache: Callable[..., bool] is_fully_idle: Callable[..., bool] + metrics_collector: Optional[Any] = None offload_tags: set = field(default_factory=set) stashed_model_static_state: Any = None + @contextmanager + def _observe_weight_load(self, source: str) -> Iterator[None]: + # Edge-trigger weight_load_duration_seconds at the end of each + # update_weights_from_* call. Engine is paused during the update so + # the periodic log_stats path can't carry this. + # `source` distinguishes disk vs distributed vs tensor vs ipc. + t0 = time.perf_counter() + try: + yield + finally: + if self.metrics_collector is not None: + self.metrics_collector.observe_weight_load( + time.perf_counter() - t0, source + ) + def flush_cache_after_weight_update(self, recv_req) -> None: if recv_req.flush_cache: flush_cache_success = self.flush_cache( @@ -87,15 +105,16 @@ class SchedulerWeightUpdaterManager: def update_weights_from_disk(self, recv_req: UpdateWeightFromDiskReqInput): """In-place update of the weights from disk.""" - success, message = self.tp_worker.update_weights_from_disk(recv_req) - tp_success = success - if success and self.draft_worker is not None: - success, message = self.draft_worker.update_weights_from_disk(recv_req) - if tp_success: - self.flush_cache_after_weight_update(recv_req) - if not success: - logger.error(message) - return UpdateWeightFromDiskReqOutput(success, message, 0) + with self._observe_weight_load("disk"): + success, message = self.tp_worker.update_weights_from_disk(recv_req) + tp_success = success + if success and self.draft_worker is not None: + success, message = self.draft_worker.update_weights_from_disk(recv_req) + if tp_success: + self.flush_cache_after_weight_update(recv_req) + if not success: + logger.error(message) + return UpdateWeightFromDiskReqOutput(success, message, 0) def init_weights_update_group(self, recv_req: InitWeightsUpdateGroupReqInput): """Initialize the online model parameter update group.""" @@ -115,39 +134,42 @@ class SchedulerWeightUpdaterManager: recv_req: UpdateWeightsFromDistributedReqInput, ) -> Tuple[bool, str]: """Update the online model parameter.""" - success, message = self.tp_worker.update_weights_from_distributed(recv_req) - if success: - self.flush_cache_after_weight_update(recv_req) - else: - logger.error(message) - return UpdateWeightsFromDistributedReqOutput(success, message) + with self._observe_weight_load("distributed"): + success, message = self.tp_worker.update_weights_from_distributed(recv_req) + if success: + self.flush_cache_after_weight_update(recv_req) + else: + logger.error(message) + return UpdateWeightsFromDistributedReqOutput(success, message) def update_weights_from_tensor(self, recv_req: UpdateWeightsFromTensorReqInput): """Update the online model parameter from tensors.""" - if recv_req.disable_draft_model: - worker = self.tp_worker - else: - worker = self.draft_worker or self.tp_worker - success, message = worker.update_weights_from_tensor(recv_req) - if success: - self.flush_cache_after_weight_update(recv_req) - else: - logger.error(message) - torch.distributed.barrier(group=self.tp_cpu_group) - return UpdateWeightsFromTensorReqOutput(success, message) + with self._observe_weight_load("tensor"): + if recv_req.disable_draft_model: + worker = self.tp_worker + else: + worker = self.draft_worker or self.tp_worker + success, message = worker.update_weights_from_tensor(recv_req) + if success: + self.flush_cache_after_weight_update(recv_req) + else: + logger.error(message) + torch.distributed.barrier(group=self.tp_cpu_group) + return UpdateWeightsFromTensorReqOutput(success, message) def update_weights_from_ipc(self, recv_req: UpdateWeightsFromIPCReqInput): """Update the online model parameter from IPC for checkpoint-engine integration.""" - success, message = self.tp_worker.update_weights_from_ipc(recv_req) - tp_success = success - if success and self.draft_worker is not None: - success, message = self.draft_worker.update_weights_from_ipc(recv_req) - if tp_success: - self.flush_cache_after_weight_update(recv_req) - if not success: - logger.error(message) - torch.distributed.barrier(group=self.tp_cpu_group) - return UpdateWeightsFromIPCReqOutput(success, message) + with self._observe_weight_load("ipc"): + success, message = self.tp_worker.update_weights_from_ipc(recv_req) + tp_success = success + if success and self.draft_worker is not None: + success, message = self.draft_worker.update_weights_from_ipc(recv_req) + if tp_success: + self.flush_cache_after_weight_update(recv_req) + if not success: + logger.error(message) + torch.distributed.barrier(group=self.tp_cpu_group) + return UpdateWeightsFromIPCReqOutput(success, message) def get_weights_by_name(self, recv_req: GetWeightsByNameReqInput): parameter = self.tp_worker.get_weights_by_name(recv_req) diff --git a/python/sglang/srt/observability/metrics_collector.py b/python/sglang/srt/observability/metrics_collector.py index 942ad8f5e..58d7b1a88 100644 --- a/python/sglang/srt/observability/metrics_collector.py +++ b/python/sglang/srt/observability/metrics_collector.py @@ -393,6 +393,21 @@ class SchedulerMetricsCollector(_StatLoggerDIMixin): multiprocess_mode="mostrecent", ) + # ================================================================= + # Weight update + # ================================================================= + self.weight_load_duration_seconds = Gauge( + name="sglang:weight_load_duration_seconds", + documentation=( + "Wall time of the most recent update_weights_from_ call on " + "this scheduler rank (seconds). `source` label is one of: disk, " + "distributed, tensor, ipc. Event-detection via " + "changes(...[]) > 0 — no separate counter needed." + ), + labelnames=[*labels.keys(), "source"], + multiprocess_mode="mostrecent", + ) + # ================================================================= # Speculative decoding # ================================================================= @@ -1124,6 +1139,14 @@ class SchedulerMetricsCollector(_StatLoggerDIMixin): def observe_queue_time(self, latency: float) -> None: self._log_histogram(self.queue_time, latency) + def observe_weight_load(self, duration_seconds: float, source: str) -> None: + # Edge-triggered: engine is paused during the update, so log_stats + # won't fire — write the gauge inline at end of update_weights_from_*. + # `source` is "disk" | "distributed" | "tensor" | "ipc". + self.weight_load_duration_seconds.labels(**self.labels, source=source).set( + duration_seconds + ) + def observe_prefill_delayer_outcome( self, forward_passes: int,