[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,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)
|
||||
|
||||
@@ -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