[srt] Add sglang:weight_load_duration_seconds gauge with source label (#27363)
This commit is contained in:
@@ -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:
|
||||
|
||||
@@ -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,6 +105,7 @@ class SchedulerWeightUpdaterManager:
|
||||
|
||||
def update_weights_from_disk(self, recv_req: UpdateWeightFromDiskReqInput):
|
||||
"""In-place update of the weights from disk."""
|
||||
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:
|
||||
@@ -115,6 +134,7 @@ class SchedulerWeightUpdaterManager:
|
||||
recv_req: UpdateWeightsFromDistributedReqInput,
|
||||
) -> Tuple[bool, str]:
|
||||
"""Update the online model parameter."""
|
||||
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)
|
||||
@@ -124,6 +144,7 @@ class SchedulerWeightUpdaterManager:
|
||||
|
||||
def update_weights_from_tensor(self, recv_req: UpdateWeightsFromTensorReqInput):
|
||||
"""Update the online model parameter from tensors."""
|
||||
with self._observe_weight_load("tensor"):
|
||||
if recv_req.disable_draft_model:
|
||||
worker = self.tp_worker
|
||||
else:
|
||||
@@ -138,6 +159,7 @@ class SchedulerWeightUpdaterManager:
|
||||
|
||||
def update_weights_from_ipc(self, recv_req: UpdateWeightsFromIPCReqInput):
|
||||
"""Update the online model parameter from IPC for checkpoint-engine integration."""
|
||||
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:
|
||||
|
||||
@@ -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_<source> call on "
|
||||
"this scheduler rank (seconds). `source` label is one of: disk, "
|
||||
"distributed, tensor, ipc. Event-detection via "
|
||||
"changes(...[<range>]) > 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,
|
||||
|
||||
Reference in New Issue
Block a user