Move metrics reporting to SchedulerMetricsReporter and retire metrics mixin (#25630)
This commit is contained in:
@@ -587,8 +587,7 @@ class SchedulerDisaggregationPrefillMixin:
|
||||
req.time_stats.set_last_chunked_prefill_finish_time()
|
||||
|
||||
can_run_cuda_graph = getattr(result, "can_run_cuda_graph", False)
|
||||
self.report_prefill_stats(
|
||||
self.metrics_reporter,
|
||||
self.metrics_reporter.report_prefill_stats(
|
||||
batch=batch,
|
||||
prefill_stats=batch.prefill_stats,
|
||||
can_run_cuda_graph=can_run_cuda_graph,
|
||||
|
||||
@@ -93,8 +93,7 @@ class SchedulerDllmMixin:
|
||||
self.token_to_kv_pool_allocator.free_group_end()
|
||||
|
||||
can_run_cuda_graph = getattr(result, "can_run_cuda_graph", False)
|
||||
self.report_prefill_stats(
|
||||
self.metrics_reporter,
|
||||
self.metrics_reporter.report_prefill_stats(
|
||||
batch=batch,
|
||||
prefill_stats=batch.prefill_stats,
|
||||
can_run_cuda_graph=can_run_cuda_graph,
|
||||
@@ -226,7 +225,9 @@ class SchedulerDllmMixin:
|
||||
new_batch.decoding_reqs = None
|
||||
|
||||
# Record prefill stats for logging after forward
|
||||
from sglang.srt.observability.scheduler_metrics_mixin import PrefillStats
|
||||
from sglang.srt.managers.scheduler_components.metrics_reporter import (
|
||||
PrefillStats,
|
||||
)
|
||||
|
||||
new_batch.prefill_stats = PrefillStats.from_adder(
|
||||
self.adder, self.running_batch.reqs, self.enable_priority_scheduling
|
||||
|
||||
@@ -96,7 +96,7 @@ if TYPE_CHECKING:
|
||||
|
||||
from sglang.srt.configs.model_config import ModelConfig
|
||||
from sglang.srt.managers.hisparse_coordinator import HiSparseCoordinator
|
||||
from sglang.srt.observability.scheduler_metrics_mixin import PrefillStats
|
||||
from sglang.srt.managers.scheduler_components.metrics_reporter import PrefillStats
|
||||
from sglang.srt.session.session_controller import Session
|
||||
from sglang.srt.speculative.eagle_info import EagleDraftInput
|
||||
from sglang.srt.speculative.spec_info import SpecInput, SpeculativeAlgorithm
|
||||
|
||||
@@ -177,6 +177,8 @@ from sglang.srt.managers.scheduler_components.load_inquirer import (
|
||||
SchedulerLoadInquirer,
|
||||
)
|
||||
from sglang.srt.managers.scheduler_components.metrics_reporter import (
|
||||
RECORD_STEP_TIME,
|
||||
PrefillStats,
|
||||
SchedulerMetricsReporter,
|
||||
)
|
||||
from sglang.srt.managers.scheduler_components.pool_stats_observer import (
|
||||
@@ -212,11 +214,6 @@ from sglang.srt.observability.req_time_stats import (
|
||||
set_schedule_time_batch,
|
||||
set_time_batch,
|
||||
)
|
||||
from sglang.srt.observability.scheduler_metrics_mixin import (
|
||||
RECORD_STEP_TIME,
|
||||
PrefillStats,
|
||||
SchedulerMetricsMixin,
|
||||
)
|
||||
from sglang.srt.observability.trace import process_tracing_init, trace_set_thread_info
|
||||
from sglang.srt.parser.reasoning_parser import ReasoningParser
|
||||
from sglang.srt.platforms import current_platform
|
||||
@@ -361,7 +358,6 @@ def create_scheduler_watchdog(
|
||||
|
||||
class Scheduler(
|
||||
SchedulerOutputProcessorMixin,
|
||||
SchedulerMetricsMixin,
|
||||
SchedulerDisaggregationDecodeMixin,
|
||||
SchedulerDisaggregationPrefillMixin,
|
||||
SchedulerMultiplexMixin,
|
||||
@@ -3067,15 +3063,15 @@ class Scheduler(
|
||||
elif batch.forward_mode.is_idle():
|
||||
self.process_batch_result_idle(batch, result)
|
||||
|
||||
self.log_batch_result_stats(self.metrics_reporter, batch, result)
|
||||
self.metrics_reporter.log_batch_result_stats(batch, result)
|
||||
|
||||
# Emit forward pass metrics (every iteration when enabled)
|
||||
if self.enable_fpm:
|
||||
self._emit_forward_pass_metrics(self.metrics_reporter, batch, result)
|
||||
self.metrics_reporter._emit_forward_pass_metrics(batch, result)
|
||||
|
||||
self._maybe_clear_mm_inputs(batch)
|
||||
self.maybe_send_health_check_signal()
|
||||
self.update_device_timer(self.metrics_reporter)
|
||||
self.metrics_reporter.update_device_timer()
|
||||
|
||||
def maybe_send_health_check_signal(self):
|
||||
if self.return_health_check_ipcs:
|
||||
@@ -3201,7 +3197,7 @@ class Scheduler(
|
||||
self.new_token_ratio = self.init_new_token_ratio
|
||||
|
||||
# reset device timer window so idle time isn't counted
|
||||
self.reset_device_timer_window(self.metrics_reporter)
|
||||
self.metrics_reporter.reset_device_timer_window()
|
||||
|
||||
# sleep until next event
|
||||
self.maybe_sleep_on_idle()
|
||||
@@ -3362,7 +3358,7 @@ class Scheduler(
|
||||
self.req_to_token_pool.clear()
|
||||
self.token_to_kv_pool_allocator.clear()
|
||||
self.grammar_manager.clear()
|
||||
self.reset_metrics(self.metrics_reporter)
|
||||
self.metrics_reporter.reset_metrics()
|
||||
|
||||
if self.draft_worker:
|
||||
self.draft_worker.clear_cache_pool()
|
||||
@@ -3981,4 +3977,4 @@ def run_scheduler_process(
|
||||
if scheduler is not None:
|
||||
# FPM has a background ZMQ publisher thread that needs explicit
|
||||
# teardown to flush queued metrics and close the socket cleanly.
|
||||
scheduler._shutdown_fpm(scheduler.metrics_reporter)
|
||||
scheduler.metrics_reporter._shutdown_fpm()
|
||||
|
||||
@@ -1,24 +1,82 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import dataclasses
|
||||
import logging
|
||||
import tempfile
|
||||
import time
|
||||
from collections import defaultdict
|
||||
from dataclasses import dataclass
|
||||
from typing import TYPE_CHECKING, Optional
|
||||
from typing import (
|
||||
TYPE_CHECKING,
|
||||
List,
|
||||
Optional,
|
||||
Tuple,
|
||||
Union,
|
||||
)
|
||||
|
||||
from sglang.srt.disaggregation.utils import DisaggregationMode
|
||||
from sglang.srt.environ import envs
|
||||
from sglang.srt.managers.schedule_batch import ScheduleBatch
|
||||
from sglang.srt.managers.utils import GenerationBatchResult
|
||||
from sglang.srt.observability.metrics_collector import (
|
||||
DPCooperationInfo,
|
||||
QueueCount,
|
||||
SchedulerMetricsCollector,
|
||||
SchedulerMetricsCollectorContext,
|
||||
SchedulerStats,
|
||||
compute_routing_key_stats,
|
||||
)
|
||||
from sglang.srt.observability.scheduler_metrics_mixin import (
|
||||
SchedulerMetricsMixin,
|
||||
)
|
||||
from sglang.srt.utils.device_timer import DeviceTimer
|
||||
from sglang.srt.utils.scheduler_status_logger import SchedulerStatusLogger
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from sglang.srt.managers.scheduler import Scheduler
|
||||
from sglang.srt.managers.schedule_batch import Req
|
||||
from sglang.srt.managers.schedule_policy import PrefillAdder
|
||||
from sglang.srt.managers.scheduler import (
|
||||
EmbeddingBatchResult,
|
||||
Scheduler,
|
||||
)
|
||||
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
RECORD_STEP_TIME = envs.SGLANG_RECORD_STEP_TIME.get()
|
||||
LOG_FORWARD_ITERS = envs.SGLANG_LOG_FORWARD_ITERS.get()
|
||||
ENABLE_METRICS_DEVICE_TIMER = envs.SGLANG_ENABLE_METRICS_DEVICE_TIMER.get()
|
||||
|
||||
|
||||
@dataclasses.dataclass
|
||||
class PrefillStats:
|
||||
"""Stats for logging prefill batch metrics."""
|
||||
|
||||
log_input_tokens: int
|
||||
log_hit_tokens: int
|
||||
new_token_ratio: float
|
||||
num_running_reqs: QueueCount
|
||||
num_new_seqs: int # len(can_run_list)
|
||||
num_pending_tokens: int = 0
|
||||
|
||||
@classmethod
|
||||
def from_adder(
|
||||
cls,
|
||||
adder: PrefillAdder,
|
||||
running_reqs: List[Req],
|
||||
enable_priority_scheduling: bool = False,
|
||||
num_pending_tokens: int = 0,
|
||||
):
|
||||
return cls(
|
||||
log_input_tokens=adder.log_input_tokens,
|
||||
log_hit_tokens=adder.log_hit_tokens,
|
||||
new_token_ratio=adder.new_token_ratio,
|
||||
num_running_reqs=QueueCount.from_reqs(
|
||||
running_reqs, enable_priority_scheduling
|
||||
),
|
||||
num_new_seqs=len(adder.can_run_list),
|
||||
num_pending_tokens=num_pending_tokens,
|
||||
)
|
||||
|
||||
|
||||
@dataclass(kw_only=True)
|
||||
class SchedulerMetricsReporter:
|
||||
scheduler: "Scheduler"
|
||||
@@ -41,7 +99,865 @@ class SchedulerMetricsReporter:
|
||||
self.enable_kv_cache_events = (
|
||||
self.metrics_collector_context.enable_kv_cache_events
|
||||
)
|
||||
SchedulerMetricsMixin._init_metrics(
|
||||
self, self.tp_rank, self.pp_rank, self.dp_rank
|
||||
self._init_metrics(self.tp_rank, self.pp_rank, self.dp_rank)
|
||||
self._install_device_timer_on_runners()
|
||||
|
||||
def _init_metrics(
|
||||
self,
|
||||
tp_rank: int,
|
||||
pp_rank: int,
|
||||
dp_rank: Optional[int],
|
||||
):
|
||||
# Basic stats
|
||||
self.forward_ct_decode = 0
|
||||
self.num_generated_tokens = 0
|
||||
self.last_decode_stats_tic = time.perf_counter()
|
||||
self.last_prefill_stats_tic = time.perf_counter()
|
||||
self.last_gen_throughput: float = 0.0
|
||||
self.last_input_throughput: float = 0.0
|
||||
self.step_time_dict = defaultdict(list) # Dict[batch size -> step time]
|
||||
self.stats = SchedulerStats()
|
||||
self._graph_backend_label = {
|
||||
"cpu": "cpu graph",
|
||||
"npu": "npu graph",
|
||||
"musa": "musa graph",
|
||||
}.get(getattr(self.scheduler, "device", ""), "cuda graph")
|
||||
|
||||
# Cumulative spec-decoding counters (reset every decode_log_interval).
|
||||
# Each update adds (num_correct_drafts + bs, bs).
|
||||
# `*_accept_tokens` = drafts + bonus; `*_correct_drafts` = drafts-only.
|
||||
self.spec_num_accept_tokens = 0 # per-log-interval
|
||||
self.spec_num_forward_ct = 0
|
||||
self.spec_total_num_accept_tokens = 0 # lifetime
|
||||
self.spec_total_num_forward_ct = 0
|
||||
|
||||
# For PD disaggregation
|
||||
self.kv_transfer_speed_gb_s: float = 0.0
|
||||
self.kv_transfer_latency_ms: float = 0.0
|
||||
|
||||
self.enable_mfu_metrics = False
|
||||
|
||||
if self.enable_metrics:
|
||||
self.enable_mfu_metrics = self.scheduler.server_args.enable_mfu_metrics
|
||||
if self.enable_mfu_metrics:
|
||||
self._init_estimated_perf_constants()
|
||||
self._mfu_log_flops = 0.0
|
||||
self._mfu_log_read_bytes = 0.0
|
||||
self._mfu_log_write_bytes = 0.0
|
||||
|
||||
self.fwd_occupancy = float("nan")
|
||||
|
||||
self.forward_pass_device_timer: Optional[DeviceTimer] = None
|
||||
|
||||
if ENABLE_METRICS_DEVICE_TIMER:
|
||||
self._device_timer_window_batch_count = 0
|
||||
self._device_timer_window_gpu_time = 0.0
|
||||
self._device_timer_window_start = None
|
||||
|
||||
def _wrap_execution_reporter(**kwargs):
|
||||
self._device_timer_window_gpu_time += kwargs["t"]
|
||||
if self.enable_metrics:
|
||||
self.metrics_collector.increment_forward_execution_seconds(**kwargs)
|
||||
|
||||
self.forward_pass_device_timer = DeviceTimer(
|
||||
reporter=_wrap_execution_reporter,
|
||||
)
|
||||
|
||||
self._init_fpm()
|
||||
|
||||
self.scheduler_status_logger = SchedulerStatusLogger.maybe_create(
|
||||
enable_metrics=self.enable_metrics
|
||||
)
|
||||
SchedulerMetricsMixin._install_device_timer_on_runners(self)
|
||||
|
||||
def _install_device_timer_on_runners(self):
|
||||
if self.forward_pass_device_timer is None:
|
||||
return
|
||||
timer = self.forward_pass_device_timer
|
||||
self.scheduler.tp_worker.model_runner.device_timer = timer
|
||||
if self.scheduler.draft_worker is not None:
|
||||
dw = getattr(self.scheduler.draft_worker, "draft_worker", None)
|
||||
if dw is not None:
|
||||
if hasattr(dw, "draft_runner"):
|
||||
dw.draft_runner.device_timer = timer
|
||||
for r in getattr(dw, "draft_runner_list", []):
|
||||
r.device_timer = timer
|
||||
|
||||
def _init_fpm(self):
|
||||
"""Initialize Forward Pass Metrics (FPM) publisher if configured."""
|
||||
self.scheduler.enable_fpm = False
|
||||
if (
|
||||
self.scheduler.server_args.enable_forward_pass_metrics
|
||||
and self.scheduler.ps.attn_tp_rank == 0
|
||||
and self.scheduler.ps.pp_rank == self.scheduler.ps.pp_size - 1
|
||||
):
|
||||
from sglang.srt.observability.forward_pass_metrics import (
|
||||
_FpmPublisherThread,
|
||||
)
|
||||
|
||||
self.scheduler._fpm_dp_rank = (
|
||||
self.scheduler.ps.dp_rank
|
||||
if self.scheduler.ps.dp_rank is not None
|
||||
else 0
|
||||
)
|
||||
self.scheduler._fpm_worker_id = (
|
||||
self.scheduler.server_args.forward_pass_metrics_worker_id
|
||||
)
|
||||
base_endpoint = self.scheduler.server_args.forward_pass_metrics_ipc_name
|
||||
if base_endpoint is None:
|
||||
ipc_path = tempfile.NamedTemporaryFile(delete=False).name
|
||||
base_endpoint = f"ipc://{ipc_path}"
|
||||
self.scheduler.server_args.forward_pass_metrics_ipc_name = base_endpoint
|
||||
endpoint = f"{base_endpoint}.{self.scheduler._fpm_dp_rank}"
|
||||
self.scheduler._fpm_publisher = _FpmPublisherThread(
|
||||
endpoint,
|
||||
worker_id=self.scheduler._fpm_worker_id,
|
||||
dp_rank=self.scheduler._fpm_dp_rank,
|
||||
)
|
||||
self.scheduler._fpm_gpu_time_acc = 0.0
|
||||
|
||||
def _fpm_device_timer_reporter(t, **_kwargs):
|
||||
self.scheduler._fpm_gpu_time_acc += t
|
||||
|
||||
if self.forward_pass_device_timer is not None:
|
||||
self.forward_pass_device_timer.add_reporter(_fpm_device_timer_reporter)
|
||||
else:
|
||||
self.forward_pass_device_timer = DeviceTimer(
|
||||
reporter=_fpm_device_timer_reporter,
|
||||
)
|
||||
self.scheduler._fpm_uses_device_timer = True
|
||||
self.scheduler.enable_fpm = True
|
||||
logger.info(
|
||||
"FPM: ZMQ PUB bound on %s (dp_rank=%d, device_timer=%s)",
|
||||
endpoint,
|
||||
self.scheduler._fpm_dp_rank,
|
||||
self.scheduler._fpm_uses_device_timer,
|
||||
)
|
||||
|
||||
def _build_scheduled_request_metrics(self, batch: ScheduleBatch):
|
||||
from sglang.srt.observability.forward_pass_metrics import (
|
||||
ScheduledRequestMetrics,
|
||||
WelfordAccumulator,
|
||||
)
|
||||
|
||||
num_prefill_requests = 0
|
||||
sum_prefill_tokens = 0
|
||||
sum_prefill_kv_tokens = 0
|
||||
prefill_lengths = WelfordAccumulator()
|
||||
|
||||
if batch.forward_mode.is_mixed():
|
||||
decode_req_ids = {id(req) for req in batch.decoding_reqs or []}
|
||||
prefill_reqs = [req for req in batch.reqs if id(req) not in decode_req_ids]
|
||||
elif batch.forward_mode.is_extend():
|
||||
prefill_reqs = batch.reqs
|
||||
else:
|
||||
prefill_reqs = []
|
||||
|
||||
if prefill_reqs:
|
||||
stats = batch.prefill_stats
|
||||
for req in prefill_reqs:
|
||||
prefill_lengths.add(len(req.origin_input_ids))
|
||||
num_prefill_requests = stats.num_new_seqs if stats else len(prefill_reqs)
|
||||
sum_prefill_tokens = stats.log_input_tokens if stats else 0
|
||||
sum_prefill_kv_tokens = sum(len(req.prefix_indices) for req in prefill_reqs)
|
||||
|
||||
decode_kv = WelfordAccumulator()
|
||||
if batch.forward_mode.is_mixed():
|
||||
for req in batch.decoding_reqs or []:
|
||||
decode_kv.add(req.seqlen)
|
||||
elif batch.forward_mode.is_decode():
|
||||
for sl in batch.seq_lens_cpu:
|
||||
decode_kv.add(int(sl))
|
||||
|
||||
return ScheduledRequestMetrics(
|
||||
num_prefill_requests=num_prefill_requests,
|
||||
sum_prefill_tokens=sum_prefill_tokens,
|
||||
var_prefill_length=prefill_lengths.variance(),
|
||||
sum_prefill_kv_tokens=sum_prefill_kv_tokens,
|
||||
num_decode_requests=decode_kv.count,
|
||||
sum_decode_kv_tokens=decode_kv.total,
|
||||
var_decode_kv_tokens=decode_kv.variance(),
|
||||
)
|
||||
|
||||
def _build_queued_request_metrics(self):
|
||||
from sglang.srt.observability.forward_pass_metrics import (
|
||||
QueuedRequestMetrics,
|
||||
WelfordAccumulator,
|
||||
)
|
||||
|
||||
prefill_q = WelfordAccumulator()
|
||||
decode_q = WelfordAccumulator()
|
||||
if self.scheduler.disaggregation_mode == DisaggregationMode.PREFILL:
|
||||
for req in self.scheduler.disagg_prefill_bootstrap_queue.queue:
|
||||
prefill_q.add(len(req.origin_input_ids))
|
||||
elif self.scheduler.disaggregation_mode == DisaggregationMode.DECODE:
|
||||
for req in self.scheduler.disagg_decode_prealloc_queue.queue:
|
||||
decode_q.add(req.seqlen)
|
||||
for req in self.scheduler.disagg_decode_transfer_queue.queue:
|
||||
decode_q.add(req.seqlen)
|
||||
else:
|
||||
for req in self.scheduler.waiting_queue:
|
||||
if len(req.output_ids) > 0:
|
||||
decode_q.add(req.seqlen)
|
||||
else:
|
||||
prefill_q.add(len(req.origin_input_ids))
|
||||
|
||||
return QueuedRequestMetrics(
|
||||
num_prefill_requests=prefill_q.count,
|
||||
sum_prefill_tokens=prefill_q.total,
|
||||
var_prefill_length=prefill_q.variance(),
|
||||
num_decode_requests=decode_q.count,
|
||||
sum_decode_kv_tokens=decode_q.total,
|
||||
var_decode_kv_tokens=decode_q.variance(),
|
||||
)
|
||||
|
||||
def update_spec_metrics(self, bs: int, num_correct_drafts: int):
|
||||
self.spec_num_accept_tokens += num_correct_drafts + bs
|
||||
self.spec_num_forward_ct += bs
|
||||
|
||||
# Bonus tokens updated elsewhere
|
||||
self.num_generated_tokens += num_correct_drafts
|
||||
|
||||
def _init_estimated_perf_constants(self) -> None:
|
||||
model_config = self.scheduler.model_config
|
||||
hf_text_config = model_config.hf_text_config
|
||||
|
||||
hidden_size = float(model_config.hidden_size)
|
||||
num_layers = float(getattr(model_config, "num_attention_layers", 0))
|
||||
head_dim = float(getattr(model_config, "head_dim", 0))
|
||||
num_attn_heads = float(
|
||||
model_config.get_num_attention_heads(self.scheduler.ps.tp_size)
|
||||
)
|
||||
num_kv_heads = float(model_config.get_num_kv_heads(self.scheduler.ps.tp_size))
|
||||
intermediate_size = getattr(hf_text_config, "intermediate_size", None)
|
||||
if intermediate_size is None:
|
||||
intermediate_size = getattr(hf_text_config, "ffn_hidden_size", 0)
|
||||
intermediate_size = float(intermediate_size)
|
||||
|
||||
dtype_num_bytes = getattr(model_config.dtype, "itemsize", None)
|
||||
if dtype_num_bytes is None:
|
||||
dtype_num_bytes = 2
|
||||
# Keep this estimator lightweight and consistent with current server dtype.
|
||||
# KV cache quantization-aware bytes can be added in a follow-up.
|
||||
act_bytes = float(dtype_num_bytes)
|
||||
w_bytes = float(dtype_num_bytes)
|
||||
cache_bytes = float(dtype_num_bytes)
|
||||
|
||||
# Linear-layer FLOPs per token on one GPU.
|
||||
attn_linear_flops = (
|
||||
2.0 * hidden_size * head_dim * (num_attn_heads + 2.0 * num_kv_heads)
|
||||
+ 2.0 * hidden_size * head_dim * num_attn_heads
|
||||
)
|
||||
mlp_flops = (
|
||||
6.0 * hidden_size * intermediate_size if intermediate_size > 0 else 0.0
|
||||
)
|
||||
self._linear_flops_per_token = max(
|
||||
0.0, (attn_linear_flops + mlp_flops) * num_layers
|
||||
)
|
||||
|
||||
# Attention dot-product FLOPs coefficient to multiply token-context product.
|
||||
# attn_qk + attn_av = 4 * q * TC * d * L
|
||||
self._attn_dot_flops_coeff = 4.0 * num_attn_heads * head_dim * num_layers
|
||||
|
||||
# KV cache bytes (write one K and one V vector per generated token).
|
||||
self._kv_cache_bytes_per_token = (
|
||||
2.0 * num_layers * num_kv_heads * head_dim * cache_bytes
|
||||
)
|
||||
|
||||
# Weight read bytes per token.
|
||||
self._weight_read_bytes_per_token = (
|
||||
hidden_size
|
||||
* head_dim
|
||||
* (num_attn_heads + 2.0 * num_kv_heads)
|
||||
* w_bytes
|
||||
* num_layers
|
||||
+ hidden_size * head_dim * num_attn_heads * w_bytes * num_layers
|
||||
+ (
|
||||
3.0 * hidden_size * intermediate_size * w_bytes * num_layers
|
||||
if intermediate_size > 0
|
||||
else 0.0
|
||||
)
|
||||
)
|
||||
|
||||
# Activation movement bytes per token (coarse approximation).
|
||||
self._qkv_act_bytes_per_token = (
|
||||
hidden_size * act_bytes * num_layers
|
||||
+ (num_attn_heads + 2.0 * num_kv_heads) * head_dim * act_bytes * num_layers
|
||||
+ head_dim * num_attn_heads * act_bytes * num_layers
|
||||
+ hidden_size * act_bytes * num_layers
|
||||
)
|
||||
self._ffn_act_bytes_per_token = (
|
||||
3.0 * intermediate_size * act_bytes * num_layers
|
||||
if intermediate_size > 0
|
||||
else 0.0
|
||||
)
|
||||
|
||||
# Prefill reads Q/K/V activations from on-device memory.
|
||||
self._prefill_attn_act_read_per_token = (
|
||||
(num_attn_heads + 2.0 * num_kv_heads) * head_dim * act_bytes * num_layers
|
||||
)
|
||||
|
||||
# Decode reads Q from activation memory; K/V reads are from KV cache.
|
||||
self._decode_q_read_bytes_per_token = (
|
||||
num_attn_heads * head_dim * act_bytes * num_layers
|
||||
)
|
||||
|
||||
def _estimate_prefill_perf(self, num_tokens: int) -> Tuple[float, float, float]:
|
||||
tokens = max(0, int(num_tokens))
|
||||
if tokens == 0:
|
||||
return 0.0, 0.0, 0.0
|
||||
|
||||
# Causal prefill token-context product.
|
||||
context_product = tokens * (tokens + 1) / 2.0
|
||||
flops = (
|
||||
tokens * self._linear_flops_per_token
|
||||
+ self._attn_dot_flops_coeff * context_product
|
||||
)
|
||||
|
||||
read_bytes = (
|
||||
tokens * self._weight_read_bytes_per_token
|
||||
+ tokens * self._qkv_act_bytes_per_token
|
||||
+ tokens * self._prefill_attn_act_read_per_token
|
||||
)
|
||||
write_bytes = (
|
||||
tokens * self._kv_cache_bytes_per_token
|
||||
+ tokens * self._qkv_act_bytes_per_token
|
||||
+ tokens * self._ffn_act_bytes_per_token
|
||||
)
|
||||
return flops, read_bytes, write_bytes
|
||||
|
||||
def _estimate_decode_perf(
|
||||
self, batch: ScheduleBatch, num_tokens: int
|
||||
) -> Tuple[float, float, float]:
|
||||
tokens = max(0, int(num_tokens))
|
||||
if tokens == 0:
|
||||
return 0.0, 0.0, 0.0
|
||||
|
||||
total_context = float(batch.seq_lens_cpu.sum().item())
|
||||
flops = (
|
||||
tokens * self._linear_flops_per_token
|
||||
+ self._attn_dot_flops_coeff * total_context
|
||||
)
|
||||
read_bytes = (
|
||||
tokens * self._weight_read_bytes_per_token
|
||||
+ tokens * self._qkv_act_bytes_per_token
|
||||
+ tokens * self._decode_q_read_bytes_per_token
|
||||
+ total_context * self._kv_cache_bytes_per_token
|
||||
)
|
||||
write_bytes = (
|
||||
tokens * self._kv_cache_bytes_per_token
|
||||
+ tokens * self._qkv_act_bytes_per_token
|
||||
+ tokens * self._ffn_act_bytes_per_token
|
||||
)
|
||||
return flops, read_bytes, write_bytes
|
||||
|
||||
def reset_metrics(self):
|
||||
self.forward_ct_decode = 0
|
||||
self.num_generated_tokens = 0
|
||||
self.spec_num_accept_tokens = 0
|
||||
self.spec_num_forward_ct = 0
|
||||
self.spec_total_num_accept_tokens = 0
|
||||
self.spec_total_num_forward_ct = 0
|
||||
|
||||
def report_prefill_stats(
|
||||
self,
|
||||
batch: Optional[ScheduleBatch],
|
||||
prefill_stats: PrefillStats,
|
||||
can_run_cuda_graph: bool,
|
||||
dp_cooperation_info: Optional[DPCooperationInfo] = None,
|
||||
):
|
||||
if (
|
||||
not self.is_stats_logging_rank
|
||||
and not self.current_scheduler_metrics_enabled
|
||||
):
|
||||
return
|
||||
|
||||
now = time.perf_counter()
|
||||
gap_latency = now - self.last_prefill_stats_tic
|
||||
self.last_prefill_stats_tic = now
|
||||
self.last_input_throughput = (
|
||||
prefill_stats.log_input_tokens / gap_latency if gap_latency > 0 else 0.0
|
||||
)
|
||||
|
||||
pool_stats = self.scheduler.pool_stats_observer.get_pool_stats()
|
||||
token_usage_msg = ", ".join(pool_stats.get_prefill_usage_msg_parts()) + ", "
|
||||
|
||||
self.stats.new_token_ratio = prefill_stats.new_token_ratio
|
||||
batch_iter = (
|
||||
batch.forward_iter
|
||||
if batch is not None and batch.forward_iter is not None
|
||||
else self.scheduler.forward_ct
|
||||
)
|
||||
iter_msg = f" [{batch_iter}]" if LOG_FORWARD_ITERS else ""
|
||||
|
||||
msg = (
|
||||
f"Prefill batch{iter_msg}, "
|
||||
f"#new-seq: {prefill_stats.num_new_seqs}, "
|
||||
f"#new-token: {prefill_stats.log_input_tokens}, "
|
||||
f"#cached-token: {prefill_stats.log_hit_tokens}, "
|
||||
f"{token_usage_msg}"
|
||||
f"#running-req: {prefill_stats.num_running_reqs.total}, "
|
||||
f"#queue-req: {len(self.scheduler.waiting_queue)}, "
|
||||
f"#pending-token: {prefill_stats.num_pending_tokens}, "
|
||||
)
|
||||
|
||||
if self.scheduler.disaggregation_mode == DisaggregationMode.PREFILL:
|
||||
msg += f"#bootstrap-req: {len(self.scheduler.disagg_prefill_bootstrap_queue.queue)}, "
|
||||
msg += (
|
||||
f"#inflight-req: {len(self.scheduler.disagg_prefill_inflight_queue)}, "
|
||||
)
|
||||
|
||||
if (
|
||||
self.scheduler.server_args.language_only
|
||||
and self.scheduler.server_args.encoder_transfer_backend
|
||||
== "zmq_to_scheduler"
|
||||
):
|
||||
msg += (
|
||||
f"waiting-image-req: {len(self.scheduler.mm_receiver.waiting_list)}, "
|
||||
)
|
||||
|
||||
msg += f"{self._graph_backend_label}: {can_run_cuda_graph}, "
|
||||
msg += f"input throughput (token/s): {self.last_input_throughput:.2f}"
|
||||
|
||||
if self.enable_mfu_metrics and gap_latency > 0:
|
||||
flops, _, _ = self._estimate_prefill_perf(prefill_stats.log_input_tokens)
|
||||
tflops_per_s = flops / gap_latency / 1e12
|
||||
msg += f", est. prefill TFLOPS/s (per GPU): {tflops_per_s:.2f}"
|
||||
|
||||
if ENABLE_METRICS_DEVICE_TIMER:
|
||||
msg += f", fwd occupancy: {self.fwd_occupancy:.2f}%"
|
||||
|
||||
if self.is_stats_logging_rank:
|
||||
logger.info(msg)
|
||||
if self.current_scheduler_metrics_enabled:
|
||||
self.metrics_collector.increment_prefill_cuda_graph_pass(
|
||||
value=can_run_cuda_graph
|
||||
)
|
||||
self.metrics_collector.increment_realtime_tokens(
|
||||
prefill_compute_tokens=prefill_stats.log_input_tokens,
|
||||
prefill_cache_tokens=prefill_stats.log_hit_tokens,
|
||||
dp_cooperation_info=dp_cooperation_info,
|
||||
)
|
||||
if self.enable_mfu_metrics:
|
||||
flops, read_bytes, write_bytes = self._estimate_prefill_perf(
|
||||
prefill_stats.log_input_tokens
|
||||
)
|
||||
self.metrics_collector.increment_estimated_perf(
|
||||
num_flops_per_gpu=flops,
|
||||
num_read_bytes_per_gpu=read_bytes,
|
||||
num_write_bytes_per_gpu=write_bytes,
|
||||
)
|
||||
|
||||
priority_enabled = self.scheduler.enable_priority_scheduling
|
||||
total_tokens = prefill_stats.log_input_tokens + prefill_stats.log_hit_tokens
|
||||
cache_hit_rate = (
|
||||
prefill_stats.log_hit_tokens / total_tokens if total_tokens > 0 else 0.0
|
||||
)
|
||||
|
||||
# Basics
|
||||
self.stats.num_running_reqs = prefill_stats.num_running_reqs
|
||||
self.stats.num_queue_reqs = QueueCount.from_reqs(
|
||||
self.scheduler.waiting_queue, priority_enabled
|
||||
)
|
||||
self.stats.num_grammar_queue_reqs = len(self.scheduler.grammar_manager)
|
||||
self.stats.cache_hit_rate = cache_hit_rate
|
||||
|
||||
# Memory pool usage ratios / Absolute token counts
|
||||
pool_stats.update_scheduler_stats(self.stats)
|
||||
|
||||
# Retract
|
||||
self.stats.num_retracted_reqs = self.num_retracted_reqs
|
||||
self.stats.num_paused_reqs = self.num_paused_reqs
|
||||
self.num_retracted_reqs = self.num_paused_reqs = 0
|
||||
|
||||
# PD disaggregation
|
||||
if self.scheduler.disaggregation_mode == DisaggregationMode.PREFILL:
|
||||
self.stats.num_prefill_bootstrap_queue_reqs = QueueCount.from_reqs(
|
||||
self.scheduler.disagg_prefill_bootstrap_queue.queue,
|
||||
priority_enabled,
|
||||
)
|
||||
self.stats.num_prefill_inflight_queue_reqs = QueueCount.from_reqs(
|
||||
self.scheduler.disagg_prefill_inflight_queue, priority_enabled
|
||||
)
|
||||
self.stats.kv_transfer_speed_gb_s = self.kv_transfer_speed_gb_s
|
||||
self.stats.kv_transfer_latency_ms = self.kv_transfer_latency_ms
|
||||
elif self.scheduler.disaggregation_mode == DisaggregationMode.DECODE:
|
||||
self.stats.num_decode_prealloc_queue_reqs = QueueCount.from_reqs(
|
||||
self.scheduler.disagg_decode_prealloc_queue.queue, priority_enabled
|
||||
)
|
||||
self.stats.num_decode_transfer_queue_reqs = QueueCount.from_reqs(
|
||||
self.scheduler.disagg_decode_transfer_queue.queue, priority_enabled
|
||||
)
|
||||
|
||||
# Utilization / LoRA / HiCache
|
||||
self._calculate_utilization()
|
||||
self.stats.fwd_occupancy = self.fwd_occupancy
|
||||
self._update_lora_metrics()
|
||||
self._log_hicache_stats()
|
||||
self.metrics_collector.log_stats(self.stats)
|
||||
self.scheduler.kv_events_publisher.emit_kv_metrics()
|
||||
self.scheduler.kv_events_publisher.publish_kv_events()
|
||||
|
||||
def report_decode_stats(
|
||||
self,
|
||||
can_run_cuda_graph: bool,
|
||||
running_batch: ScheduleBatch = None,
|
||||
num_correct_drafts: int = 0,
|
||||
):
|
||||
batch = running_batch or self.scheduler.running_batch
|
||||
|
||||
# Every-iteration work: realtime token counting + status logger
|
||||
if self.current_scheduler_metrics_enabled:
|
||||
decode_tokens = batch.batch_size() + num_correct_drafts
|
||||
self.metrics_collector.increment_realtime_tokens(
|
||||
# TODO unify this w/ the bumping logic in `Scheduler.num_generated_tokens` accumulator
|
||||
decode_tokens=decode_tokens,
|
||||
dp_cooperation_info=batch.dp_cooperation_info,
|
||||
)
|
||||
if self.enable_mfu_metrics:
|
||||
flops, read_bytes, write_bytes = self._estimate_decode_perf(
|
||||
batch, decode_tokens
|
||||
)
|
||||
self.metrics_collector.increment_estimated_perf(
|
||||
num_flops_per_gpu=flops,
|
||||
num_read_bytes_per_gpu=read_bytes,
|
||||
num_write_bytes_per_gpu=write_bytes,
|
||||
)
|
||||
self._mfu_log_flops += flops
|
||||
self._mfu_log_read_bytes += read_bytes
|
||||
self._mfu_log_write_bytes += write_bytes
|
||||
|
||||
if x := self.scheduler_status_logger:
|
||||
x.maybe_dump(batch, self.scheduler.waiting_queue)
|
||||
|
||||
# Periodic work: log + heavy metrics at decode_log_interval
|
||||
if self.forward_ct_decode % self.scheduler.server_args.decode_log_interval != 0:
|
||||
return
|
||||
if (
|
||||
not self.is_stats_logging_rank
|
||||
and not self.current_scheduler_metrics_enabled
|
||||
):
|
||||
return
|
||||
|
||||
gap_latency = time.perf_counter() - self.last_decode_stats_tic
|
||||
self.last_decode_stats_tic = time.perf_counter()
|
||||
self.last_gen_throughput = self.num_generated_tokens / gap_latency
|
||||
|
||||
self.num_generated_tokens = 0
|
||||
num_running_reqs = len(batch.reqs)
|
||||
|
||||
pool_stats = self.scheduler.pool_stats_observer.get_pool_stats()
|
||||
token_usage_msg = ", ".join(pool_stats.get_decode_usage_msg_parts()) + ", "
|
||||
|
||||
if RECORD_STEP_TIME:
|
||||
self.step_time_dict[num_running_reqs].append(
|
||||
gap_latency / self.scheduler.server_args.decode_log_interval
|
||||
)
|
||||
|
||||
batch_iter = (
|
||||
batch.forward_iter
|
||||
if batch is not None and batch.forward_iter is not None
|
||||
else self.scheduler.forward_ct
|
||||
)
|
||||
iter_msg = f" [{batch_iter}]" if LOG_FORWARD_ITERS else ""
|
||||
msg = f"Decode batch{iter_msg}, #running-req: {num_running_reqs}, {token_usage_msg}"
|
||||
|
||||
if self.scheduler.spec_algorithm.is_none():
|
||||
spec_accept_length = 0
|
||||
spec_accept_rate = 0
|
||||
else:
|
||||
spec_accept_length = self.spec_num_accept_tokens / self.spec_num_forward_ct
|
||||
num_correct_drafts = self.spec_num_accept_tokens - self.spec_num_forward_ct
|
||||
if self.scheduler.server_args.speculative_num_draft_tokens:
|
||||
draft_per_round = (
|
||||
self.scheduler.server_args.speculative_num_draft_tokens - 1
|
||||
)
|
||||
else:
|
||||
draft_per_round = self.scheduler.server_args.speculative_num_steps or 0
|
||||
total_draft_tokens = self.spec_num_forward_ct * draft_per_round
|
||||
spec_accept_rate = (
|
||||
num_correct_drafts / total_draft_tokens if total_draft_tokens > 0 else 0
|
||||
)
|
||||
self.spec_total_num_accept_tokens += self.spec_num_accept_tokens
|
||||
self.spec_total_num_forward_ct += self.spec_num_forward_ct
|
||||
self.spec_num_accept_tokens = self.spec_num_forward_ct = 0
|
||||
msg += f"accept len: {spec_accept_length:.2f}, accept rate: {spec_accept_rate:.2f}, "
|
||||
cache_hit_rate = 0.0
|
||||
|
||||
if self.scheduler.disaggregation_mode == DisaggregationMode.DECODE:
|
||||
msg += f"pre-allocated usage: {self.scheduler.disagg_decode_prealloc_queue.num_tokens_pre_allocated / self.scheduler.max_total_num_tokens:.2f}, "
|
||||
msg += f"#prealloc-req: {len(self.scheduler.disagg_decode_prealloc_queue.queue)}, "
|
||||
msg += f"#transfer-req: {len(self.scheduler.disagg_decode_transfer_queue.queue)}, "
|
||||
msg += f"#retracted-req: {len(self.scheduler.disagg_decode_prealloc_queue.retracted_queue)}, "
|
||||
|
||||
if (
|
||||
self.scheduler.server_args.language_only
|
||||
and self.scheduler.server_args.encoder_transfer_backend
|
||||
== "zmq_to_scheduler"
|
||||
):
|
||||
msg += (
|
||||
f"waiting-image-req: {len(self.scheduler.mm_receiver.waiting_list)}, "
|
||||
)
|
||||
|
||||
msg += (
|
||||
f"{self._graph_backend_label}: {can_run_cuda_graph}, "
|
||||
f"gen throughput (token/s): {self.last_gen_throughput:.2f}, "
|
||||
f"#queue-req: {len(self.scheduler.waiting_queue)}"
|
||||
)
|
||||
|
||||
if self.enable_mfu_metrics and gap_latency > 0:
|
||||
flops_per_s = self._mfu_log_flops / gap_latency
|
||||
read_bytes_per_s = self._mfu_log_read_bytes / gap_latency
|
||||
write_bytes_per_s = self._mfu_log_write_bytes / gap_latency
|
||||
tflops_per_s = flops_per_s / 1e12
|
||||
read_gb_per_s = read_bytes_per_s / 1e9
|
||||
write_gb_per_s = write_bytes_per_s / 1e9
|
||||
msg += (
|
||||
f", est. decode TFLOPS/s (per GPU): {tflops_per_s:.2f}, "
|
||||
f"est. read BW (GB/s per GPU): {read_gb_per_s:.2f}, "
|
||||
f"est. write BW (GB/s per GPU): {write_gb_per_s:.2f}"
|
||||
)
|
||||
self._mfu_log_flops = 0.0
|
||||
self._mfu_log_read_bytes = 0.0
|
||||
self._mfu_log_write_bytes = 0.0
|
||||
|
||||
if ENABLE_METRICS_DEVICE_TIMER:
|
||||
msg += f", fwd occupancy: {self.fwd_occupancy:.2f}%"
|
||||
|
||||
if self.is_stats_logging_rank:
|
||||
logger.info(msg)
|
||||
if self.current_scheduler_metrics_enabled:
|
||||
priority_enabled = self.scheduler.enable_priority_scheduling
|
||||
|
||||
# Basics
|
||||
self.stats.num_running_reqs = QueueCount.from_reqs(
|
||||
batch.reqs, priority_enabled
|
||||
)
|
||||
self.stats.num_queue_reqs = QueueCount.from_reqs(
|
||||
self.scheduler.waiting_queue, priority_enabled
|
||||
)
|
||||
self.stats.num_grammar_queue_reqs = len(self.scheduler.grammar_manager)
|
||||
self.stats.gen_throughput = self.last_gen_throughput
|
||||
self.stats.cache_hit_rate = cache_hit_rate
|
||||
self.stats.decode_sum_seq_lens = batch.seq_lens_cpu.sum().item()
|
||||
|
||||
# Memory pool usage ratios / Absolute token counts
|
||||
pool_stats.update_scheduler_stats(self.stats)
|
||||
|
||||
# Speculative decoding
|
||||
self.stats.spec_accept_length = spec_accept_length
|
||||
self.stats.spec_accept_rate = spec_accept_rate
|
||||
|
||||
# Retract
|
||||
self.stats.num_retracted_reqs = self.num_retracted_reqs
|
||||
self.stats.num_paused_reqs = self.num_paused_reqs
|
||||
self.num_retracted_reqs = self.num_paused_reqs = 0
|
||||
|
||||
# PD disaggregation
|
||||
if self.scheduler.disaggregation_mode == DisaggregationMode.PREFILL:
|
||||
self.stats.num_prefill_bootstrap_queue_reqs = QueueCount.from_reqs(
|
||||
self.scheduler.disagg_prefill_bootstrap_queue.queue,
|
||||
priority_enabled,
|
||||
)
|
||||
self.stats.num_prefill_inflight_queue_reqs = QueueCount.from_reqs(
|
||||
self.scheduler.disagg_prefill_inflight_queue, priority_enabled
|
||||
)
|
||||
elif self.scheduler.disaggregation_mode == DisaggregationMode.DECODE:
|
||||
self.stats.num_decode_prealloc_queue_reqs = QueueCount.from_reqs(
|
||||
self.scheduler.disagg_decode_prealloc_queue.queue, priority_enabled
|
||||
)
|
||||
self.stats.num_decode_transfer_queue_reqs = QueueCount.from_reqs(
|
||||
self.scheduler.disagg_decode_transfer_queue.queue, priority_enabled
|
||||
)
|
||||
|
||||
# Streaming session metrics
|
||||
self.stats.num_streaming_sessions = (
|
||||
self.scheduler.pool_stats_observer.streaming_session_count()
|
||||
)
|
||||
self.stats.streaming_session_held_tokens = (
|
||||
self.scheduler.pool_stats_observer.session_held_tokens()
|
||||
)
|
||||
|
||||
# Routing key metrics
|
||||
# (to reduce the overhead, we only compute this when all requests have routing_key)
|
||||
if all(r.routing_key is not None for r in batch.reqs):
|
||||
running_routing_keys = [r.routing_key for r in batch.reqs]
|
||||
waiting_routing_keys = [
|
||||
r.routing_key for r in self.scheduler.waiting_queue
|
||||
]
|
||||
(
|
||||
self.stats.num_unique_running_routing_keys,
|
||||
self.stats.routing_key_running_req_counts,
|
||||
) = compute_routing_key_stats(running_routing_keys)
|
||||
_, self.stats.routing_key_all_req_counts = compute_routing_key_stats(
|
||||
running_routing_keys + waiting_routing_keys
|
||||
)
|
||||
|
||||
# Utilization / LoRA / HiCache
|
||||
self._calculate_utilization()
|
||||
self.stats.fwd_occupancy = self.fwd_occupancy
|
||||
self._update_lora_metrics()
|
||||
self._log_hicache_stats()
|
||||
self.metrics_collector.log_stats(self.stats)
|
||||
self.scheduler.kv_events_publisher.emit_kv_metrics()
|
||||
self.scheduler.kv_events_publisher.publish_kv_events()
|
||||
|
||||
def log_batch_result_stats(
|
||||
self,
|
||||
batch: ScheduleBatch,
|
||||
result: Union[GenerationBatchResult, EmbeddingBatchResult],
|
||||
):
|
||||
if not self.enable_metrics:
|
||||
return
|
||||
if not isinstance(result, GenerationBatchResult):
|
||||
return
|
||||
|
||||
if (m := result.expert_distribution_metrics) is not None:
|
||||
self.metrics_collector.increment_eplb_balancedness(
|
||||
forward_mode=batch.forward_mode.name.lower(),
|
||||
balancedness=m.eplb_balancedness.item(),
|
||||
)
|
||||
|
||||
def _emit_forward_pass_metrics(
|
||||
self,
|
||||
batch: ScheduleBatch,
|
||||
result=None,
|
||||
):
|
||||
"""Emit per-iteration ForwardPassMetrics over ZMQ PUB.
|
||||
|
||||
Prefers GPU-accurate timing from DeviceTimer (which wraps
|
||||
model_runner.forward / cuda_graph.replay via PR #24197).
|
||||
Falls back to monotonic clock when DeviceTimer is not enabled.
|
||||
"""
|
||||
if not self.scheduler.enable_fpm:
|
||||
return
|
||||
|
||||
from sglang.srt.observability.forward_pass_metrics import (
|
||||
ForwardPassMetrics,
|
||||
)
|
||||
|
||||
if self.scheduler._fpm_uses_device_timer:
|
||||
self.forward_pass_device_timer._report()
|
||||
wall_time = self.scheduler._fpm_gpu_time_acc
|
||||
self.scheduler._fpm_gpu_time_acc = 0.0
|
||||
if wall_time == 0.0:
|
||||
return
|
||||
else:
|
||||
wall_time = max(0.0, time.monotonic() - batch.fpm_start_time)
|
||||
|
||||
fpm = ForwardPassMetrics(
|
||||
worker_id=self.scheduler._fpm_worker_id,
|
||||
dp_rank=self.scheduler._fpm_dp_rank,
|
||||
wall_time=wall_time,
|
||||
scheduled_requests=self._build_scheduled_request_metrics(batch),
|
||||
queued_requests=self._build_queued_request_metrics(),
|
||||
)
|
||||
self.scheduler._fpm_publisher.publish(fpm)
|
||||
|
||||
def _shutdown_fpm(self):
|
||||
"""Shut down the FPM publisher thread."""
|
||||
if self.scheduler.enable_fpm:
|
||||
self.scheduler._fpm_publisher.shutdown()
|
||||
|
||||
def _log_hicache_stats(self):
|
||||
"""Populate HiCache host-tier stats on self.stats.
|
||||
|
||||
These are pushed to Prometheus by SchedulerMetricsCollector.log_stats().
|
||||
"""
|
||||
if not self.scheduler.enable_hierarchical_cache:
|
||||
return
|
||||
|
||||
host_pool = getattr(
|
||||
self.scheduler.tree_cache, "token_to_kv_pool_host", None
|
||||
) or getattr(self.scheduler.tree_cache, "full_kv_pool_host", None)
|
||||
assert host_pool is not None, "Host pool not found"
|
||||
self.stats.hicache_host_used_tokens = (
|
||||
host_pool.size - host_pool.available_size()
|
||||
)
|
||||
self.stats.hicache_host_total_tokens = host_pool.size
|
||||
|
||||
def _update_lora_metrics(self):
|
||||
"""Update LoRA pool metrics for monitoring and autoscaling."""
|
||||
if not self.scheduler.enable_lora:
|
||||
return
|
||||
|
||||
try:
|
||||
# Get LoRA memory pool stats
|
||||
lora_manager = self.scheduler.tp_worker.model_runner.lora_manager
|
||||
if lora_manager is None or lora_manager.memory_pool is None:
|
||||
return
|
||||
|
||||
mem_pool = lora_manager.memory_pool
|
||||
slots_total = mem_pool.max_loras_per_batch
|
||||
|
||||
# Calculate active adapters from running batch
|
||||
# This gives a true measure of current load for autoscaling purposes
|
||||
active_lora_ids = set()
|
||||
|
||||
# For PP mode, check all running micro batches
|
||||
if self.scheduler.server_args.pp_size > 1:
|
||||
for batch in self.scheduler.running_mbs:
|
||||
if batch and hasattr(batch, "reqs"):
|
||||
for req in batch.reqs:
|
||||
if hasattr(req, "lora_id") and req.lora_id is not None:
|
||||
active_lora_ids.add(req.lora_id)
|
||||
# For normal mode, check running_batch
|
||||
elif self.scheduler.running_batch:
|
||||
if hasattr(self.scheduler.running_batch, "reqs"):
|
||||
for req in self.scheduler.running_batch.reqs:
|
||||
if hasattr(req, "lora_id") and req.lora_id is not None:
|
||||
active_lora_ids.add(req.lora_id)
|
||||
|
||||
# Count active adapters (excluding None for base model)
|
||||
slots_used = len(active_lora_ids)
|
||||
utilization = slots_used / slots_total if slots_total > 0 else 0.0
|
||||
|
||||
# Update stats
|
||||
self.stats.lora_pool_slots_used = slots_used
|
||||
self.stats.lora_pool_slots_total = slots_total
|
||||
self.stats.lora_pool_utilization = utilization
|
||||
|
||||
except Exception as e:
|
||||
logger.warning(f"Failed to update LoRA metrics: {e}")
|
||||
|
||||
def _calculate_utilization(self):
|
||||
if self.scheduler.disaggregation_mode == DisaggregationMode.PREFILL:
|
||||
self.stats.utilization = -1
|
||||
else:
|
||||
# TODO: max_running_requests_under_SLO has no setter — sglang:utilization stuck at 0 (regressed #22713).
|
||||
max_under_slo = getattr(
|
||||
self.scheduler, "max_running_requests_under_SLO", None
|
||||
)
|
||||
if max_under_slo is not None and max_under_slo > 0:
|
||||
self.stats.utilization = max(
|
||||
self.stats.num_running_reqs.total / max_under_slo,
|
||||
self.stats.token_usage / 0.9,
|
||||
)
|
||||
|
||||
def update_device_timer(self):
|
||||
if not ENABLE_METRICS_DEVICE_TIMER:
|
||||
return
|
||||
self.forward_pass_device_timer._report()
|
||||
now = time.perf_counter()
|
||||
if self._device_timer_window_batch_count == 0:
|
||||
self._device_timer_window_start = now
|
||||
self._device_timer_window_gpu_time = 0.0
|
||||
cpu_time = 0
|
||||
self.fwd_occupancy = float("nan")
|
||||
else:
|
||||
cpu_time = now - self._device_timer_window_start
|
||||
self.fwd_occupancy = min(
|
||||
self._device_timer_window_gpu_time / cpu_time * 100, 100
|
||||
)
|
||||
# ratio = self._device_timer_window_gpu_time / cpu_time if cpu_time > 0 else float("nan")
|
||||
# print(f"{self._device_timer_window_batch_count=} {self.fwd_occupancy=}, {self._device_timer_window_gpu_time=}, {cpu_time=}, {ratio=}")
|
||||
self._device_timer_window_batch_count += 1
|
||||
if (
|
||||
self._device_timer_window_batch_count
|
||||
>= self.scheduler.server_args.decode_log_interval
|
||||
):
|
||||
self._device_timer_window_batch_count = 0
|
||||
|
||||
def reset_device_timer_window(self):
|
||||
if ENABLE_METRICS_DEVICE_TIMER:
|
||||
self._device_timer_window_batch_count = 0
|
||||
self.fwd_occupancy = float("nan")
|
||||
|
||||
@@ -399,8 +399,7 @@ class SchedulerOutputProcessorMixin:
|
||||
self.stream_output(batch.reqs, batch.return_logprob, skip_stream_req)
|
||||
|
||||
can_run_cuda_graph = getattr(result, "can_run_cuda_graph", False)
|
||||
self.report_prefill_stats(
|
||||
self.metrics_reporter,
|
||||
self.metrics_reporter.report_prefill_stats(
|
||||
batch=batch,
|
||||
prefill_stats=batch.prefill_stats,
|
||||
can_run_cuda_graph=can_run_cuda_graph,
|
||||
@@ -514,8 +513,8 @@ class SchedulerOutputProcessorMixin:
|
||||
|
||||
self.metrics_reporter.num_generated_tokens += len(batch.reqs)
|
||||
if not batch.spec_algorithm.is_none():
|
||||
self.update_spec_metrics(
|
||||
self.metrics_reporter, batch.batch_size(), result.num_correct_drafts
|
||||
self.metrics_reporter.update_spec_metrics(
|
||||
batch.batch_size(), result.num_correct_drafts
|
||||
)
|
||||
if self.metrics_reporter.enable_metrics:
|
||||
self.metrics_collector.increment_decode_cuda_graph_pass(
|
||||
@@ -630,8 +629,7 @@ class SchedulerOutputProcessorMixin:
|
||||
self.metrics_reporter.forward_ct_decode = (
|
||||
self.metrics_reporter.forward_ct_decode + 1
|
||||
) % (1 << 30)
|
||||
self.report_decode_stats(
|
||||
self.metrics_reporter,
|
||||
self.metrics_reporter.report_decode_stats(
|
||||
can_run_cuda_graph,
|
||||
running_batch=batch,
|
||||
num_correct_drafts=result.num_correct_drafts,
|
||||
|
||||
@@ -1,961 +0,0 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import dataclasses
|
||||
import logging
|
||||
import tempfile
|
||||
import time
|
||||
from collections import defaultdict
|
||||
from typing import TYPE_CHECKING, List, Optional, Tuple, Union
|
||||
|
||||
from sglang.srt.disaggregation.utils import DisaggregationMode
|
||||
from sglang.srt.environ import envs
|
||||
from sglang.srt.managers.schedule_batch import ScheduleBatch
|
||||
from sglang.srt.managers.utils import GenerationBatchResult
|
||||
from sglang.srt.observability.metrics_collector import (
|
||||
DPCooperationInfo,
|
||||
QueueCount,
|
||||
SchedulerStats,
|
||||
compute_routing_key_stats,
|
||||
)
|
||||
from sglang.srt.utils.device_timer import DeviceTimer
|
||||
from sglang.srt.utils.scheduler_status_logger import SchedulerStatusLogger
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from sglang.srt.managers.schedule_batch import Req
|
||||
from sglang.srt.managers.schedule_policy import PrefillAdder
|
||||
from sglang.srt.managers.scheduler import EmbeddingBatchResult
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
RECORD_STEP_TIME = envs.SGLANG_RECORD_STEP_TIME.get()
|
||||
LOG_FORWARD_ITERS = envs.SGLANG_LOG_FORWARD_ITERS.get()
|
||||
ENABLE_METRICS_DEVICE_TIMER = envs.SGLANG_ENABLE_METRICS_DEVICE_TIMER.get()
|
||||
|
||||
|
||||
@dataclasses.dataclass
|
||||
class PrefillStats:
|
||||
"""Stats for logging prefill batch metrics."""
|
||||
|
||||
log_input_tokens: int
|
||||
log_hit_tokens: int
|
||||
new_token_ratio: float
|
||||
num_running_reqs: QueueCount
|
||||
num_new_seqs: int # len(can_run_list)
|
||||
num_pending_tokens: int = 0
|
||||
|
||||
@classmethod
|
||||
def from_adder(
|
||||
cls,
|
||||
adder: PrefillAdder,
|
||||
running_reqs: List[Req],
|
||||
enable_priority_scheduling: bool = False,
|
||||
num_pending_tokens: int = 0,
|
||||
):
|
||||
return cls(
|
||||
log_input_tokens=adder.log_input_tokens,
|
||||
log_hit_tokens=adder.log_hit_tokens,
|
||||
new_token_ratio=adder.new_token_ratio,
|
||||
num_running_reqs=QueueCount.from_reqs(
|
||||
running_reqs, enable_priority_scheduling
|
||||
),
|
||||
num_new_seqs=len(adder.can_run_list),
|
||||
num_pending_tokens=num_pending_tokens,
|
||||
)
|
||||
|
||||
|
||||
class SchedulerMetricsMixin:
|
||||
enable_fpm: bool = False
|
||||
|
||||
@staticmethod
|
||||
def _init_metrics(
|
||||
self: "SchedulerMetricsReporter",
|
||||
tp_rank: int,
|
||||
pp_rank: int,
|
||||
dp_rank: Optional[int],
|
||||
):
|
||||
# Basic stats
|
||||
self.forward_ct_decode = 0
|
||||
self.num_generated_tokens = 0
|
||||
self.last_decode_stats_tic = time.perf_counter()
|
||||
self.last_prefill_stats_tic = time.perf_counter()
|
||||
self.last_gen_throughput: float = 0.0
|
||||
self.last_input_throughput: float = 0.0
|
||||
self.step_time_dict = defaultdict(list) # Dict[batch size -> step time]
|
||||
self.stats = SchedulerStats()
|
||||
self._graph_backend_label = {
|
||||
"cpu": "cpu graph",
|
||||
"npu": "npu graph",
|
||||
"musa": "musa graph",
|
||||
}.get(getattr(self.scheduler, "device", ""), "cuda graph")
|
||||
|
||||
# Cumulative spec-decoding counters (reset every decode_log_interval).
|
||||
# Each update adds (num_correct_drafts + bs, bs).
|
||||
# `*_accept_tokens` = drafts + bonus; `*_correct_drafts` = drafts-only.
|
||||
self.spec_num_accept_tokens = 0 # per-log-interval
|
||||
self.spec_num_forward_ct = 0
|
||||
self.spec_total_num_accept_tokens = 0 # lifetime
|
||||
self.spec_total_num_forward_ct = 0
|
||||
|
||||
# For PD disaggregation
|
||||
self.kv_transfer_speed_gb_s: float = 0.0
|
||||
self.kv_transfer_latency_ms: float = 0.0
|
||||
|
||||
self.enable_mfu_metrics = False
|
||||
|
||||
if self.enable_metrics:
|
||||
self.enable_mfu_metrics = self.scheduler.server_args.enable_mfu_metrics
|
||||
if self.enable_mfu_metrics:
|
||||
SchedulerMetricsMixin._init_estimated_perf_constants(self)
|
||||
self._mfu_log_flops = 0.0
|
||||
self._mfu_log_read_bytes = 0.0
|
||||
self._mfu_log_write_bytes = 0.0
|
||||
|
||||
self.fwd_occupancy = float("nan")
|
||||
|
||||
self.forward_pass_device_timer: Optional[DeviceTimer] = None
|
||||
|
||||
if ENABLE_METRICS_DEVICE_TIMER:
|
||||
self._device_timer_window_batch_count = 0
|
||||
self._device_timer_window_gpu_time = 0.0
|
||||
self._device_timer_window_start = None
|
||||
|
||||
def _wrap_execution_reporter(**kwargs):
|
||||
self._device_timer_window_gpu_time += kwargs["t"]
|
||||
if self.enable_metrics:
|
||||
self.metrics_collector.increment_forward_execution_seconds(**kwargs)
|
||||
|
||||
self.forward_pass_device_timer = DeviceTimer(
|
||||
reporter=_wrap_execution_reporter,
|
||||
)
|
||||
|
||||
SchedulerMetricsMixin._init_fpm(self)
|
||||
|
||||
self.scheduler_status_logger = SchedulerStatusLogger.maybe_create(
|
||||
enable_metrics=self.enable_metrics
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _install_device_timer_on_runners(self: "SchedulerMetricsReporter"):
|
||||
if self.forward_pass_device_timer is None:
|
||||
return
|
||||
timer = self.forward_pass_device_timer
|
||||
self.scheduler.tp_worker.model_runner.device_timer = timer
|
||||
if self.scheduler.draft_worker is not None:
|
||||
dw = getattr(self.scheduler.draft_worker, "draft_worker", None)
|
||||
if dw is not None:
|
||||
if hasattr(dw, "draft_runner"):
|
||||
dw.draft_runner.device_timer = timer
|
||||
for r in getattr(dw, "draft_runner_list", []):
|
||||
r.device_timer = timer
|
||||
|
||||
@staticmethod
|
||||
def _init_fpm(self: "SchedulerMetricsReporter"):
|
||||
"""Initialize Forward Pass Metrics (FPM) publisher if configured."""
|
||||
self.scheduler.enable_fpm = False
|
||||
if (
|
||||
self.scheduler.server_args.enable_forward_pass_metrics
|
||||
and self.scheduler.ps.attn_tp_rank == 0
|
||||
and self.scheduler.ps.pp_rank == self.scheduler.ps.pp_size - 1
|
||||
):
|
||||
from sglang.srt.observability.forward_pass_metrics import (
|
||||
_FpmPublisherThread,
|
||||
)
|
||||
|
||||
self.scheduler._fpm_dp_rank = (
|
||||
self.scheduler.ps.dp_rank
|
||||
if self.scheduler.ps.dp_rank is not None
|
||||
else 0
|
||||
)
|
||||
self.scheduler._fpm_worker_id = (
|
||||
self.scheduler.server_args.forward_pass_metrics_worker_id
|
||||
)
|
||||
base_endpoint = self.scheduler.server_args.forward_pass_metrics_ipc_name
|
||||
if base_endpoint is None:
|
||||
ipc_path = tempfile.NamedTemporaryFile(delete=False).name
|
||||
base_endpoint = f"ipc://{ipc_path}"
|
||||
self.scheduler.server_args.forward_pass_metrics_ipc_name = base_endpoint
|
||||
endpoint = f"{base_endpoint}.{self.scheduler._fpm_dp_rank}"
|
||||
self.scheduler._fpm_publisher = _FpmPublisherThread(
|
||||
endpoint,
|
||||
worker_id=self.scheduler._fpm_worker_id,
|
||||
dp_rank=self.scheduler._fpm_dp_rank,
|
||||
)
|
||||
self.scheduler._fpm_gpu_time_acc = 0.0
|
||||
|
||||
def _fpm_device_timer_reporter(t, **_kwargs):
|
||||
self.scheduler._fpm_gpu_time_acc += t
|
||||
|
||||
if self.forward_pass_device_timer is not None:
|
||||
self.forward_pass_device_timer.add_reporter(_fpm_device_timer_reporter)
|
||||
else:
|
||||
self.forward_pass_device_timer = DeviceTimer(
|
||||
reporter=_fpm_device_timer_reporter,
|
||||
)
|
||||
self.scheduler._fpm_uses_device_timer = True
|
||||
self.scheduler.enable_fpm = True
|
||||
logger.info(
|
||||
"FPM: ZMQ PUB bound on %s (dp_rank=%d, device_timer=%s)",
|
||||
endpoint,
|
||||
self.scheduler._fpm_dp_rank,
|
||||
self.scheduler._fpm_uses_device_timer,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _build_scheduled_request_metrics(
|
||||
self: "SchedulerMetricsReporter", batch: ScheduleBatch
|
||||
):
|
||||
from sglang.srt.observability.forward_pass_metrics import (
|
||||
ScheduledRequestMetrics,
|
||||
WelfordAccumulator,
|
||||
)
|
||||
|
||||
num_prefill_requests = 0
|
||||
sum_prefill_tokens = 0
|
||||
sum_prefill_kv_tokens = 0
|
||||
prefill_lengths = WelfordAccumulator()
|
||||
|
||||
if batch.forward_mode.is_mixed():
|
||||
decode_req_ids = {id(req) for req in batch.decoding_reqs or []}
|
||||
prefill_reqs = [req for req in batch.reqs if id(req) not in decode_req_ids]
|
||||
elif batch.forward_mode.is_extend():
|
||||
prefill_reqs = batch.reqs
|
||||
else:
|
||||
prefill_reqs = []
|
||||
|
||||
if prefill_reqs:
|
||||
stats = batch.prefill_stats
|
||||
for req in prefill_reqs:
|
||||
prefill_lengths.add(len(req.origin_input_ids))
|
||||
num_prefill_requests = stats.num_new_seqs if stats else len(prefill_reqs)
|
||||
sum_prefill_tokens = stats.log_input_tokens if stats else 0
|
||||
sum_prefill_kv_tokens = sum(len(req.prefix_indices) for req in prefill_reqs)
|
||||
|
||||
decode_kv = WelfordAccumulator()
|
||||
if batch.forward_mode.is_mixed():
|
||||
for req in batch.decoding_reqs or []:
|
||||
decode_kv.add(req.seqlen)
|
||||
elif batch.forward_mode.is_decode():
|
||||
for sl in batch.seq_lens_cpu:
|
||||
decode_kv.add(int(sl))
|
||||
|
||||
return ScheduledRequestMetrics(
|
||||
num_prefill_requests=num_prefill_requests,
|
||||
sum_prefill_tokens=sum_prefill_tokens,
|
||||
var_prefill_length=prefill_lengths.variance(),
|
||||
sum_prefill_kv_tokens=sum_prefill_kv_tokens,
|
||||
num_decode_requests=decode_kv.count,
|
||||
sum_decode_kv_tokens=decode_kv.total,
|
||||
var_decode_kv_tokens=decode_kv.variance(),
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _build_queued_request_metrics(self: "SchedulerMetricsReporter"):
|
||||
from sglang.srt.observability.forward_pass_metrics import (
|
||||
QueuedRequestMetrics,
|
||||
WelfordAccumulator,
|
||||
)
|
||||
|
||||
prefill_q = WelfordAccumulator()
|
||||
decode_q = WelfordAccumulator()
|
||||
if self.scheduler.disaggregation_mode == DisaggregationMode.PREFILL:
|
||||
for req in self.scheduler.disagg_prefill_bootstrap_queue.queue:
|
||||
prefill_q.add(len(req.origin_input_ids))
|
||||
elif self.scheduler.disaggregation_mode == DisaggregationMode.DECODE:
|
||||
for req in self.scheduler.disagg_decode_prealloc_queue.queue:
|
||||
decode_q.add(req.seqlen)
|
||||
for req in self.scheduler.disagg_decode_transfer_queue.queue:
|
||||
decode_q.add(req.seqlen)
|
||||
else:
|
||||
for req in self.scheduler.waiting_queue:
|
||||
if len(req.output_ids) > 0:
|
||||
decode_q.add(req.seqlen)
|
||||
else:
|
||||
prefill_q.add(len(req.origin_input_ids))
|
||||
|
||||
return QueuedRequestMetrics(
|
||||
num_prefill_requests=prefill_q.count,
|
||||
sum_prefill_tokens=prefill_q.total,
|
||||
var_prefill_length=prefill_q.variance(),
|
||||
num_decode_requests=decode_q.count,
|
||||
sum_decode_kv_tokens=decode_q.total,
|
||||
var_decode_kv_tokens=decode_q.variance(),
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def update_spec_metrics(
|
||||
self: "SchedulerMetricsReporter", bs: int, num_correct_drafts: int
|
||||
):
|
||||
self.spec_num_accept_tokens += num_correct_drafts + bs
|
||||
self.spec_num_forward_ct += bs
|
||||
|
||||
# Bonus tokens updated elsewhere
|
||||
self.num_generated_tokens += num_correct_drafts
|
||||
|
||||
@staticmethod
|
||||
def _init_estimated_perf_constants(self: "SchedulerMetricsReporter") -> None:
|
||||
model_config = self.scheduler.model_config
|
||||
hf_text_config = model_config.hf_text_config
|
||||
|
||||
hidden_size = float(model_config.hidden_size)
|
||||
num_layers = float(getattr(model_config, "num_attention_layers", 0))
|
||||
head_dim = float(getattr(model_config, "head_dim", 0))
|
||||
num_attn_heads = float(
|
||||
model_config.get_num_attention_heads(self.scheduler.ps.tp_size)
|
||||
)
|
||||
num_kv_heads = float(model_config.get_num_kv_heads(self.scheduler.ps.tp_size))
|
||||
intermediate_size = getattr(hf_text_config, "intermediate_size", None)
|
||||
if intermediate_size is None:
|
||||
intermediate_size = getattr(hf_text_config, "ffn_hidden_size", 0)
|
||||
intermediate_size = float(intermediate_size)
|
||||
|
||||
dtype_num_bytes = getattr(model_config.dtype, "itemsize", None)
|
||||
if dtype_num_bytes is None:
|
||||
dtype_num_bytes = 2
|
||||
# Keep this estimator lightweight and consistent with current server dtype.
|
||||
# KV cache quantization-aware bytes can be added in a follow-up.
|
||||
act_bytes = float(dtype_num_bytes)
|
||||
w_bytes = float(dtype_num_bytes)
|
||||
cache_bytes = float(dtype_num_bytes)
|
||||
|
||||
# Linear-layer FLOPs per token on one GPU.
|
||||
attn_linear_flops = (
|
||||
2.0 * hidden_size * head_dim * (num_attn_heads + 2.0 * num_kv_heads)
|
||||
+ 2.0 * hidden_size * head_dim * num_attn_heads
|
||||
)
|
||||
mlp_flops = (
|
||||
6.0 * hidden_size * intermediate_size if intermediate_size > 0 else 0.0
|
||||
)
|
||||
self._linear_flops_per_token = max(
|
||||
0.0, (attn_linear_flops + mlp_flops) * num_layers
|
||||
)
|
||||
|
||||
# Attention dot-product FLOPs coefficient to multiply token-context product.
|
||||
# attn_qk + attn_av = 4 * q * TC * d * L
|
||||
self._attn_dot_flops_coeff = 4.0 * num_attn_heads * head_dim * num_layers
|
||||
|
||||
# KV cache bytes (write one K and one V vector per generated token).
|
||||
self._kv_cache_bytes_per_token = (
|
||||
2.0 * num_layers * num_kv_heads * head_dim * cache_bytes
|
||||
)
|
||||
|
||||
# Weight read bytes per token.
|
||||
self._weight_read_bytes_per_token = (
|
||||
hidden_size
|
||||
* head_dim
|
||||
* (num_attn_heads + 2.0 * num_kv_heads)
|
||||
* w_bytes
|
||||
* num_layers
|
||||
+ hidden_size * head_dim * num_attn_heads * w_bytes * num_layers
|
||||
+ (
|
||||
3.0 * hidden_size * intermediate_size * w_bytes * num_layers
|
||||
if intermediate_size > 0
|
||||
else 0.0
|
||||
)
|
||||
)
|
||||
|
||||
# Activation movement bytes per token (coarse approximation).
|
||||
self._qkv_act_bytes_per_token = (
|
||||
hidden_size * act_bytes * num_layers
|
||||
+ (num_attn_heads + 2.0 * num_kv_heads) * head_dim * act_bytes * num_layers
|
||||
+ head_dim * num_attn_heads * act_bytes * num_layers
|
||||
+ hidden_size * act_bytes * num_layers
|
||||
)
|
||||
self._ffn_act_bytes_per_token = (
|
||||
3.0 * intermediate_size * act_bytes * num_layers
|
||||
if intermediate_size > 0
|
||||
else 0.0
|
||||
)
|
||||
|
||||
# Prefill reads Q/K/V activations from on-device memory.
|
||||
self._prefill_attn_act_read_per_token = (
|
||||
(num_attn_heads + 2.0 * num_kv_heads) * head_dim * act_bytes * num_layers
|
||||
)
|
||||
|
||||
# Decode reads Q from activation memory; K/V reads are from KV cache.
|
||||
self._decode_q_read_bytes_per_token = (
|
||||
num_attn_heads * head_dim * act_bytes * num_layers
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _estimate_prefill_perf(
|
||||
self: "SchedulerMetricsReporter", num_tokens: int
|
||||
) -> Tuple[float, float, float]:
|
||||
tokens = max(0, int(num_tokens))
|
||||
if tokens == 0:
|
||||
return 0.0, 0.0, 0.0
|
||||
|
||||
# Causal prefill token-context product.
|
||||
context_product = tokens * (tokens + 1) / 2.0
|
||||
flops = (
|
||||
tokens * self._linear_flops_per_token
|
||||
+ self._attn_dot_flops_coeff * context_product
|
||||
)
|
||||
|
||||
read_bytes = (
|
||||
tokens * self._weight_read_bytes_per_token
|
||||
+ tokens * self._qkv_act_bytes_per_token
|
||||
+ tokens * self._prefill_attn_act_read_per_token
|
||||
)
|
||||
write_bytes = (
|
||||
tokens * self._kv_cache_bytes_per_token
|
||||
+ tokens * self._qkv_act_bytes_per_token
|
||||
+ tokens * self._ffn_act_bytes_per_token
|
||||
)
|
||||
return flops, read_bytes, write_bytes
|
||||
|
||||
@staticmethod
|
||||
def _estimate_decode_perf(
|
||||
self: "SchedulerMetricsReporter", batch: ScheduleBatch, num_tokens: int
|
||||
) -> Tuple[float, float, float]:
|
||||
tokens = max(0, int(num_tokens))
|
||||
if tokens == 0:
|
||||
return 0.0, 0.0, 0.0
|
||||
|
||||
total_context = float(batch.seq_lens_cpu.sum().item())
|
||||
flops = (
|
||||
tokens * self._linear_flops_per_token
|
||||
+ self._attn_dot_flops_coeff * total_context
|
||||
)
|
||||
read_bytes = (
|
||||
tokens * self._weight_read_bytes_per_token
|
||||
+ tokens * self._qkv_act_bytes_per_token
|
||||
+ tokens * self._decode_q_read_bytes_per_token
|
||||
+ total_context * self._kv_cache_bytes_per_token
|
||||
)
|
||||
write_bytes = (
|
||||
tokens * self._kv_cache_bytes_per_token
|
||||
+ tokens * self._qkv_act_bytes_per_token
|
||||
+ tokens * self._ffn_act_bytes_per_token
|
||||
)
|
||||
return flops, read_bytes, write_bytes
|
||||
|
||||
@staticmethod
|
||||
def reset_metrics(self: "SchedulerMetricsReporter"):
|
||||
self.forward_ct_decode = 0
|
||||
self.num_generated_tokens = 0
|
||||
self.spec_num_accept_tokens = 0
|
||||
self.spec_num_forward_ct = 0
|
||||
self.spec_total_num_accept_tokens = 0
|
||||
self.spec_total_num_forward_ct = 0
|
||||
|
||||
@staticmethod
|
||||
def report_prefill_stats(
|
||||
self: "SchedulerMetricsReporter",
|
||||
batch: Optional[ScheduleBatch],
|
||||
prefill_stats: PrefillStats,
|
||||
can_run_cuda_graph: bool,
|
||||
dp_cooperation_info: Optional[DPCooperationInfo] = None,
|
||||
):
|
||||
if (
|
||||
not self.is_stats_logging_rank
|
||||
and not self.current_scheduler_metrics_enabled
|
||||
):
|
||||
return
|
||||
|
||||
now = time.perf_counter()
|
||||
gap_latency = now - self.last_prefill_stats_tic
|
||||
self.last_prefill_stats_tic = now
|
||||
self.last_input_throughput = (
|
||||
prefill_stats.log_input_tokens / gap_latency if gap_latency > 0 else 0.0
|
||||
)
|
||||
|
||||
pool_stats = self.scheduler.pool_stats_observer.get_pool_stats()
|
||||
token_usage_msg = ", ".join(pool_stats.get_prefill_usage_msg_parts()) + ", "
|
||||
|
||||
self.stats.new_token_ratio = prefill_stats.new_token_ratio
|
||||
batch_iter = (
|
||||
batch.forward_iter
|
||||
if batch is not None and batch.forward_iter is not None
|
||||
else self.scheduler.forward_ct
|
||||
)
|
||||
iter_msg = f" [{batch_iter}]" if LOG_FORWARD_ITERS else ""
|
||||
|
||||
msg = (
|
||||
f"Prefill batch{iter_msg}, "
|
||||
f"#new-seq: {prefill_stats.num_new_seqs}, "
|
||||
f"#new-token: {prefill_stats.log_input_tokens}, "
|
||||
f"#cached-token: {prefill_stats.log_hit_tokens}, "
|
||||
f"{token_usage_msg}"
|
||||
f"#running-req: {prefill_stats.num_running_reqs.total}, "
|
||||
f"#queue-req: {len(self.scheduler.waiting_queue)}, "
|
||||
f"#pending-token: {prefill_stats.num_pending_tokens}, "
|
||||
)
|
||||
|
||||
if self.scheduler.disaggregation_mode == DisaggregationMode.PREFILL:
|
||||
msg += f"#bootstrap-req: {len(self.scheduler.disagg_prefill_bootstrap_queue.queue)}, "
|
||||
msg += (
|
||||
f"#inflight-req: {len(self.scheduler.disagg_prefill_inflight_queue)}, "
|
||||
)
|
||||
|
||||
if (
|
||||
self.scheduler.server_args.language_only
|
||||
and self.scheduler.server_args.encoder_transfer_backend
|
||||
== "zmq_to_scheduler"
|
||||
):
|
||||
msg += (
|
||||
f"waiting-image-req: {len(self.scheduler.mm_receiver.waiting_list)}, "
|
||||
)
|
||||
|
||||
msg += f"{self._graph_backend_label}: {can_run_cuda_graph}, "
|
||||
msg += f"input throughput (token/s): {self.last_input_throughput:.2f}"
|
||||
|
||||
if self.enable_mfu_metrics and gap_latency > 0:
|
||||
flops, _, _ = SchedulerMetricsMixin._estimate_prefill_perf(
|
||||
self, prefill_stats.log_input_tokens
|
||||
)
|
||||
tflops_per_s = flops / gap_latency / 1e12
|
||||
msg += f", est. prefill TFLOPS/s (per GPU): {tflops_per_s:.2f}"
|
||||
|
||||
if ENABLE_METRICS_DEVICE_TIMER:
|
||||
msg += f", fwd occupancy: {self.fwd_occupancy:.2f}%"
|
||||
|
||||
if self.is_stats_logging_rank:
|
||||
logger.info(msg)
|
||||
if self.current_scheduler_metrics_enabled:
|
||||
self.metrics_collector.increment_prefill_cuda_graph_pass(
|
||||
value=can_run_cuda_graph
|
||||
)
|
||||
self.metrics_collector.increment_realtime_tokens(
|
||||
prefill_compute_tokens=prefill_stats.log_input_tokens,
|
||||
prefill_cache_tokens=prefill_stats.log_hit_tokens,
|
||||
dp_cooperation_info=dp_cooperation_info,
|
||||
)
|
||||
if self.enable_mfu_metrics:
|
||||
flops, read_bytes, write_bytes = (
|
||||
SchedulerMetricsMixin._estimate_prefill_perf(
|
||||
self, prefill_stats.log_input_tokens
|
||||
)
|
||||
)
|
||||
self.metrics_collector.increment_estimated_perf(
|
||||
num_flops_per_gpu=flops,
|
||||
num_read_bytes_per_gpu=read_bytes,
|
||||
num_write_bytes_per_gpu=write_bytes,
|
||||
)
|
||||
|
||||
priority_enabled = self.scheduler.enable_priority_scheduling
|
||||
total_tokens = prefill_stats.log_input_tokens + prefill_stats.log_hit_tokens
|
||||
cache_hit_rate = (
|
||||
prefill_stats.log_hit_tokens / total_tokens if total_tokens > 0 else 0.0
|
||||
)
|
||||
|
||||
# Basics
|
||||
self.stats.num_running_reqs = prefill_stats.num_running_reqs
|
||||
self.stats.num_queue_reqs = QueueCount.from_reqs(
|
||||
self.scheduler.waiting_queue, priority_enabled
|
||||
)
|
||||
self.stats.num_grammar_queue_reqs = len(self.scheduler.grammar_manager)
|
||||
self.stats.cache_hit_rate = cache_hit_rate
|
||||
|
||||
# Memory pool usage ratios / Absolute token counts
|
||||
pool_stats.update_scheduler_stats(self.stats)
|
||||
|
||||
# Retract
|
||||
self.stats.num_retracted_reqs = self.num_retracted_reqs
|
||||
self.stats.num_paused_reqs = self.num_paused_reqs
|
||||
self.num_retracted_reqs = self.num_paused_reqs = 0
|
||||
|
||||
# PD disaggregation
|
||||
if self.scheduler.disaggregation_mode == DisaggregationMode.PREFILL:
|
||||
self.stats.num_prefill_bootstrap_queue_reqs = QueueCount.from_reqs(
|
||||
self.scheduler.disagg_prefill_bootstrap_queue.queue,
|
||||
priority_enabled,
|
||||
)
|
||||
self.stats.num_prefill_inflight_queue_reqs = QueueCount.from_reqs(
|
||||
self.scheduler.disagg_prefill_inflight_queue, priority_enabled
|
||||
)
|
||||
self.stats.kv_transfer_speed_gb_s = self.kv_transfer_speed_gb_s
|
||||
self.stats.kv_transfer_latency_ms = self.kv_transfer_latency_ms
|
||||
elif self.scheduler.disaggregation_mode == DisaggregationMode.DECODE:
|
||||
self.stats.num_decode_prealloc_queue_reqs = QueueCount.from_reqs(
|
||||
self.scheduler.disagg_decode_prealloc_queue.queue, priority_enabled
|
||||
)
|
||||
self.stats.num_decode_transfer_queue_reqs = QueueCount.from_reqs(
|
||||
self.scheduler.disagg_decode_transfer_queue.queue, priority_enabled
|
||||
)
|
||||
|
||||
# Utilization / LoRA / HiCache
|
||||
SchedulerMetricsMixin._calculate_utilization(self)
|
||||
self.stats.fwd_occupancy = self.fwd_occupancy
|
||||
SchedulerMetricsMixin._update_lora_metrics(self)
|
||||
SchedulerMetricsMixin._log_hicache_stats(self)
|
||||
self.metrics_collector.log_stats(self.stats)
|
||||
self.scheduler.kv_events_publisher.emit_kv_metrics()
|
||||
self.scheduler.kv_events_publisher.publish_kv_events()
|
||||
|
||||
@staticmethod
|
||||
def report_decode_stats(
|
||||
self: "SchedulerMetricsReporter",
|
||||
can_run_cuda_graph: bool,
|
||||
running_batch: ScheduleBatch = None,
|
||||
num_correct_drafts: int = 0,
|
||||
):
|
||||
batch = running_batch or self.scheduler.running_batch
|
||||
|
||||
# Every-iteration work: realtime token counting + status logger
|
||||
if self.current_scheduler_metrics_enabled:
|
||||
decode_tokens = batch.batch_size() + num_correct_drafts
|
||||
self.metrics_collector.increment_realtime_tokens(
|
||||
# TODO unify this w/ the bumping logic in `Scheduler.num_generated_tokens` accumulator
|
||||
decode_tokens=decode_tokens,
|
||||
dp_cooperation_info=batch.dp_cooperation_info,
|
||||
)
|
||||
if self.enable_mfu_metrics:
|
||||
flops, read_bytes, write_bytes = (
|
||||
SchedulerMetricsMixin._estimate_decode_perf(
|
||||
self, batch, decode_tokens
|
||||
)
|
||||
)
|
||||
self.metrics_collector.increment_estimated_perf(
|
||||
num_flops_per_gpu=flops,
|
||||
num_read_bytes_per_gpu=read_bytes,
|
||||
num_write_bytes_per_gpu=write_bytes,
|
||||
)
|
||||
self._mfu_log_flops += flops
|
||||
self._mfu_log_read_bytes += read_bytes
|
||||
self._mfu_log_write_bytes += write_bytes
|
||||
|
||||
if x := self.scheduler_status_logger:
|
||||
x.maybe_dump(batch, self.scheduler.waiting_queue)
|
||||
|
||||
# Periodic work: log + heavy metrics at decode_log_interval
|
||||
if self.forward_ct_decode % self.scheduler.server_args.decode_log_interval != 0:
|
||||
return
|
||||
if (
|
||||
not self.is_stats_logging_rank
|
||||
and not self.current_scheduler_metrics_enabled
|
||||
):
|
||||
return
|
||||
|
||||
gap_latency = time.perf_counter() - self.last_decode_stats_tic
|
||||
self.last_decode_stats_tic = time.perf_counter()
|
||||
self.last_gen_throughput = self.num_generated_tokens / gap_latency
|
||||
|
||||
self.num_generated_tokens = 0
|
||||
num_running_reqs = len(batch.reqs)
|
||||
|
||||
pool_stats = self.scheduler.pool_stats_observer.get_pool_stats()
|
||||
token_usage_msg = ", ".join(pool_stats.get_decode_usage_msg_parts()) + ", "
|
||||
|
||||
if RECORD_STEP_TIME:
|
||||
self.step_time_dict[num_running_reqs].append(
|
||||
gap_latency / self.scheduler.server_args.decode_log_interval
|
||||
)
|
||||
|
||||
batch_iter = (
|
||||
batch.forward_iter
|
||||
if batch is not None and batch.forward_iter is not None
|
||||
else self.scheduler.forward_ct
|
||||
)
|
||||
iter_msg = f" [{batch_iter}]" if LOG_FORWARD_ITERS else ""
|
||||
msg = f"Decode batch{iter_msg}, #running-req: {num_running_reqs}, {token_usage_msg}"
|
||||
|
||||
if self.scheduler.spec_algorithm.is_none():
|
||||
spec_accept_length = 0
|
||||
spec_accept_rate = 0
|
||||
else:
|
||||
spec_accept_length = self.spec_num_accept_tokens / self.spec_num_forward_ct
|
||||
num_correct_drafts = self.spec_num_accept_tokens - self.spec_num_forward_ct
|
||||
if self.scheduler.server_args.speculative_num_draft_tokens:
|
||||
draft_per_round = (
|
||||
self.scheduler.server_args.speculative_num_draft_tokens - 1
|
||||
)
|
||||
else:
|
||||
draft_per_round = self.scheduler.server_args.speculative_num_steps or 0
|
||||
total_draft_tokens = self.spec_num_forward_ct * draft_per_round
|
||||
spec_accept_rate = (
|
||||
num_correct_drafts / total_draft_tokens if total_draft_tokens > 0 else 0
|
||||
)
|
||||
self.spec_total_num_accept_tokens += self.spec_num_accept_tokens
|
||||
self.spec_total_num_forward_ct += self.spec_num_forward_ct
|
||||
self.spec_num_accept_tokens = self.spec_num_forward_ct = 0
|
||||
msg += f"accept len: {spec_accept_length:.2f}, accept rate: {spec_accept_rate:.2f}, "
|
||||
cache_hit_rate = 0.0
|
||||
|
||||
if self.scheduler.disaggregation_mode == DisaggregationMode.DECODE:
|
||||
msg += f"pre-allocated usage: {self.scheduler.disagg_decode_prealloc_queue.num_tokens_pre_allocated / self.scheduler.max_total_num_tokens:.2f}, "
|
||||
msg += f"#prealloc-req: {len(self.scheduler.disagg_decode_prealloc_queue.queue)}, "
|
||||
msg += f"#transfer-req: {len(self.scheduler.disagg_decode_transfer_queue.queue)}, "
|
||||
msg += f"#retracted-req: {len(self.scheduler.disagg_decode_prealloc_queue.retracted_queue)}, "
|
||||
|
||||
if (
|
||||
self.scheduler.server_args.language_only
|
||||
and self.scheduler.server_args.encoder_transfer_backend
|
||||
== "zmq_to_scheduler"
|
||||
):
|
||||
msg += (
|
||||
f"waiting-image-req: {len(self.scheduler.mm_receiver.waiting_list)}, "
|
||||
)
|
||||
|
||||
msg += (
|
||||
f"{self._graph_backend_label}: {can_run_cuda_graph}, "
|
||||
f"gen throughput (token/s): {self.last_gen_throughput:.2f}, "
|
||||
f"#queue-req: {len(self.scheduler.waiting_queue)}"
|
||||
)
|
||||
|
||||
if self.enable_mfu_metrics and gap_latency > 0:
|
||||
flops_per_s = self._mfu_log_flops / gap_latency
|
||||
read_bytes_per_s = self._mfu_log_read_bytes / gap_latency
|
||||
write_bytes_per_s = self._mfu_log_write_bytes / gap_latency
|
||||
tflops_per_s = flops_per_s / 1e12
|
||||
read_gb_per_s = read_bytes_per_s / 1e9
|
||||
write_gb_per_s = write_bytes_per_s / 1e9
|
||||
msg += (
|
||||
f", est. decode TFLOPS/s (per GPU): {tflops_per_s:.2f}, "
|
||||
f"est. read BW (GB/s per GPU): {read_gb_per_s:.2f}, "
|
||||
f"est. write BW (GB/s per GPU): {write_gb_per_s:.2f}"
|
||||
)
|
||||
self._mfu_log_flops = 0.0
|
||||
self._mfu_log_read_bytes = 0.0
|
||||
self._mfu_log_write_bytes = 0.0
|
||||
|
||||
if ENABLE_METRICS_DEVICE_TIMER:
|
||||
msg += f", fwd occupancy: {self.fwd_occupancy:.2f}%"
|
||||
|
||||
if self.is_stats_logging_rank:
|
||||
logger.info(msg)
|
||||
if self.current_scheduler_metrics_enabled:
|
||||
priority_enabled = self.scheduler.enable_priority_scheduling
|
||||
|
||||
# Basics
|
||||
self.stats.num_running_reqs = QueueCount.from_reqs(
|
||||
batch.reqs, priority_enabled
|
||||
)
|
||||
self.stats.num_queue_reqs = QueueCount.from_reqs(
|
||||
self.scheduler.waiting_queue, priority_enabled
|
||||
)
|
||||
self.stats.num_grammar_queue_reqs = len(self.scheduler.grammar_manager)
|
||||
self.stats.gen_throughput = self.last_gen_throughput
|
||||
self.stats.cache_hit_rate = cache_hit_rate
|
||||
self.stats.decode_sum_seq_lens = batch.seq_lens_cpu.sum().item()
|
||||
|
||||
# Memory pool usage ratios / Absolute token counts
|
||||
pool_stats.update_scheduler_stats(self.stats)
|
||||
|
||||
# Speculative decoding
|
||||
self.stats.spec_accept_length = spec_accept_length
|
||||
self.stats.spec_accept_rate = spec_accept_rate
|
||||
|
||||
# Retract
|
||||
self.stats.num_retracted_reqs = self.num_retracted_reqs
|
||||
self.stats.num_paused_reqs = self.num_paused_reqs
|
||||
self.num_retracted_reqs = self.num_paused_reqs = 0
|
||||
|
||||
# PD disaggregation
|
||||
if self.scheduler.disaggregation_mode == DisaggregationMode.PREFILL:
|
||||
self.stats.num_prefill_bootstrap_queue_reqs = QueueCount.from_reqs(
|
||||
self.scheduler.disagg_prefill_bootstrap_queue.queue,
|
||||
priority_enabled,
|
||||
)
|
||||
self.stats.num_prefill_inflight_queue_reqs = QueueCount.from_reqs(
|
||||
self.scheduler.disagg_prefill_inflight_queue, priority_enabled
|
||||
)
|
||||
elif self.scheduler.disaggregation_mode == DisaggregationMode.DECODE:
|
||||
self.stats.num_decode_prealloc_queue_reqs = QueueCount.from_reqs(
|
||||
self.scheduler.disagg_decode_prealloc_queue.queue, priority_enabled
|
||||
)
|
||||
self.stats.num_decode_transfer_queue_reqs = QueueCount.from_reqs(
|
||||
self.scheduler.disagg_decode_transfer_queue.queue, priority_enabled
|
||||
)
|
||||
|
||||
# Streaming session metrics
|
||||
self.stats.num_streaming_sessions = (
|
||||
self.scheduler.pool_stats_observer.streaming_session_count()
|
||||
)
|
||||
self.stats.streaming_session_held_tokens = (
|
||||
self.scheduler.pool_stats_observer.session_held_tokens()
|
||||
)
|
||||
|
||||
# Routing key metrics
|
||||
# (to reduce the overhead, we only compute this when all requests have routing_key)
|
||||
if all(r.routing_key is not None for r in batch.reqs):
|
||||
running_routing_keys = [r.routing_key for r in batch.reqs]
|
||||
waiting_routing_keys = [
|
||||
r.routing_key for r in self.scheduler.waiting_queue
|
||||
]
|
||||
(
|
||||
self.stats.num_unique_running_routing_keys,
|
||||
self.stats.routing_key_running_req_counts,
|
||||
) = compute_routing_key_stats(running_routing_keys)
|
||||
_, self.stats.routing_key_all_req_counts = compute_routing_key_stats(
|
||||
running_routing_keys + waiting_routing_keys
|
||||
)
|
||||
|
||||
# Utilization / LoRA / HiCache
|
||||
SchedulerMetricsMixin._calculate_utilization(self)
|
||||
self.stats.fwd_occupancy = self.fwd_occupancy
|
||||
SchedulerMetricsMixin._update_lora_metrics(self)
|
||||
SchedulerMetricsMixin._log_hicache_stats(self)
|
||||
self.metrics_collector.log_stats(self.stats)
|
||||
self.scheduler.kv_events_publisher.emit_kv_metrics()
|
||||
self.scheduler.kv_events_publisher.publish_kv_events()
|
||||
|
||||
@staticmethod
|
||||
def log_batch_result_stats(
|
||||
self: "SchedulerMetricsReporter",
|
||||
batch: ScheduleBatch,
|
||||
result: Union[GenerationBatchResult, EmbeddingBatchResult],
|
||||
):
|
||||
if not self.enable_metrics:
|
||||
return
|
||||
if not isinstance(result, GenerationBatchResult):
|
||||
return
|
||||
|
||||
if (m := result.expert_distribution_metrics) is not None:
|
||||
self.metrics_collector.increment_eplb_balancedness(
|
||||
forward_mode=batch.forward_mode.name.lower(),
|
||||
balancedness=m.eplb_balancedness.item(),
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _emit_forward_pass_metrics(
|
||||
self: "SchedulerMetricsReporter",
|
||||
batch: ScheduleBatch,
|
||||
result=None,
|
||||
):
|
||||
"""Emit per-iteration ForwardPassMetrics over ZMQ PUB.
|
||||
|
||||
Prefers GPU-accurate timing from DeviceTimer (which wraps
|
||||
model_runner.forward / cuda_graph.replay via PR #24197).
|
||||
Falls back to monotonic clock when DeviceTimer is not enabled.
|
||||
"""
|
||||
if not self.scheduler.enable_fpm:
|
||||
return
|
||||
|
||||
from sglang.srt.observability.forward_pass_metrics import (
|
||||
ForwardPassMetrics,
|
||||
)
|
||||
|
||||
if self.scheduler._fpm_uses_device_timer:
|
||||
self.forward_pass_device_timer._report()
|
||||
wall_time = self.scheduler._fpm_gpu_time_acc
|
||||
self.scheduler._fpm_gpu_time_acc = 0.0
|
||||
if wall_time == 0.0:
|
||||
return
|
||||
else:
|
||||
wall_time = max(0.0, time.monotonic() - batch.fpm_start_time)
|
||||
|
||||
fpm = ForwardPassMetrics(
|
||||
worker_id=self.scheduler._fpm_worker_id,
|
||||
dp_rank=self.scheduler._fpm_dp_rank,
|
||||
wall_time=wall_time,
|
||||
scheduled_requests=SchedulerMetricsMixin._build_scheduled_request_metrics(
|
||||
self, batch
|
||||
),
|
||||
queued_requests=SchedulerMetricsMixin._build_queued_request_metrics(self),
|
||||
)
|
||||
self.scheduler._fpm_publisher.publish(fpm)
|
||||
|
||||
@staticmethod
|
||||
def _shutdown_fpm(self: "SchedulerMetricsReporter"):
|
||||
"""Shut down the FPM publisher thread."""
|
||||
if self.scheduler.enable_fpm:
|
||||
self.scheduler._fpm_publisher.shutdown()
|
||||
|
||||
@staticmethod
|
||||
def _log_hicache_stats(self: "SchedulerMetricsReporter"):
|
||||
"""Populate HiCache host-tier stats on self.stats.
|
||||
|
||||
These are pushed to Prometheus by SchedulerMetricsCollector.log_stats().
|
||||
"""
|
||||
if not self.scheduler.enable_hierarchical_cache:
|
||||
return
|
||||
|
||||
host_pool = getattr(
|
||||
self.scheduler.tree_cache, "token_to_kv_pool_host", None
|
||||
) or getattr(self.scheduler.tree_cache, "full_kv_pool_host", None)
|
||||
assert host_pool is not None, "Host pool not found"
|
||||
self.stats.hicache_host_used_tokens = (
|
||||
host_pool.size - host_pool.available_size()
|
||||
)
|
||||
self.stats.hicache_host_total_tokens = host_pool.size
|
||||
|
||||
@staticmethod
|
||||
def _update_lora_metrics(self: "SchedulerMetricsReporter"):
|
||||
"""Update LoRA pool metrics for monitoring and autoscaling."""
|
||||
if not self.scheduler.enable_lora:
|
||||
return
|
||||
|
||||
try:
|
||||
# Get LoRA memory pool stats
|
||||
lora_manager = self.scheduler.tp_worker.model_runner.lora_manager
|
||||
if lora_manager is None or lora_manager.memory_pool is None:
|
||||
return
|
||||
|
||||
mem_pool = lora_manager.memory_pool
|
||||
slots_total = mem_pool.max_loras_per_batch
|
||||
|
||||
# Calculate active adapters from running batch
|
||||
# This gives a true measure of current load for autoscaling purposes
|
||||
active_lora_ids = set()
|
||||
|
||||
# For PP mode, check all running micro batches
|
||||
if self.scheduler.server_args.pp_size > 1:
|
||||
for batch in self.scheduler.running_mbs:
|
||||
if batch and hasattr(batch, "reqs"):
|
||||
for req in batch.reqs:
|
||||
if hasattr(req, "lora_id") and req.lora_id is not None:
|
||||
active_lora_ids.add(req.lora_id)
|
||||
# For normal mode, check running_batch
|
||||
elif self.scheduler.running_batch:
|
||||
if hasattr(self.scheduler.running_batch, "reqs"):
|
||||
for req in self.scheduler.running_batch.reqs:
|
||||
if hasattr(req, "lora_id") and req.lora_id is not None:
|
||||
active_lora_ids.add(req.lora_id)
|
||||
|
||||
# Count active adapters (excluding None for base model)
|
||||
slots_used = len(active_lora_ids)
|
||||
utilization = slots_used / slots_total if slots_total > 0 else 0.0
|
||||
|
||||
# Update stats
|
||||
self.stats.lora_pool_slots_used = slots_used
|
||||
self.stats.lora_pool_slots_total = slots_total
|
||||
self.stats.lora_pool_utilization = utilization
|
||||
|
||||
except Exception as e:
|
||||
logger.warning(f"Failed to update LoRA metrics: {e}")
|
||||
|
||||
@staticmethod
|
||||
def _calculate_utilization(self: "SchedulerMetricsReporter"):
|
||||
if self.scheduler.disaggregation_mode == DisaggregationMode.PREFILL:
|
||||
self.stats.utilization = -1
|
||||
else:
|
||||
# TODO: max_running_requests_under_SLO has no setter — sglang:utilization stuck at 0 (regressed #22713).
|
||||
max_under_slo = getattr(
|
||||
self.scheduler, "max_running_requests_under_SLO", None
|
||||
)
|
||||
if max_under_slo is not None and max_under_slo > 0:
|
||||
self.stats.utilization = max(
|
||||
self.stats.num_running_reqs.total / max_under_slo,
|
||||
self.stats.token_usage / 0.9,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def update_device_timer(self: "SchedulerMetricsReporter"):
|
||||
if not ENABLE_METRICS_DEVICE_TIMER:
|
||||
return
|
||||
self.forward_pass_device_timer._report()
|
||||
now = time.perf_counter()
|
||||
if self._device_timer_window_batch_count == 0:
|
||||
self._device_timer_window_start = now
|
||||
self._device_timer_window_gpu_time = 0.0
|
||||
cpu_time = 0
|
||||
self.fwd_occupancy = float("nan")
|
||||
else:
|
||||
cpu_time = now - self._device_timer_window_start
|
||||
self.fwd_occupancy = min(
|
||||
self._device_timer_window_gpu_time / cpu_time * 100, 100
|
||||
)
|
||||
# ratio = self._device_timer_window_gpu_time / cpu_time if cpu_time > 0 else float("nan")
|
||||
# print(f"{self._device_timer_window_batch_count=} {self.fwd_occupancy=}, {self._device_timer_window_gpu_time=}, {cpu_time=}, {ratio=}")
|
||||
self._device_timer_window_batch_count += 1
|
||||
if (
|
||||
self._device_timer_window_batch_count
|
||||
>= self.scheduler.server_args.decode_log_interval
|
||||
):
|
||||
self._device_timer_window_batch_count = 0
|
||||
|
||||
@staticmethod
|
||||
def reset_device_timer_window(self: "SchedulerMetricsReporter"):
|
||||
if ENABLE_METRICS_DEVICE_TIMER:
|
||||
self._device_timer_window_batch_count = 0
|
||||
self.fwd_occupancy = float("nan")
|
||||
@@ -8,9 +8,9 @@ from unittest.mock import patch
|
||||
|
||||
from sglang.srt.disaggregation.utils import DisaggregationMode
|
||||
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
|
||||
from sglang.srt.observability.scheduler_metrics_mixin import (
|
||||
from sglang.srt.managers.scheduler_components.metrics_reporter import (
|
||||
PrefillStats,
|
||||
SchedulerMetricsMixin,
|
||||
SchedulerMetricsReporter,
|
||||
)
|
||||
|
||||
|
||||
@@ -85,14 +85,49 @@ class _DummyPublisherThread:
|
||||
pass
|
||||
|
||||
|
||||
class _DummyScheduler(SchedulerMetricsMixin):
|
||||
pass
|
||||
def _make_reporter(scheduler) -> SchedulerMetricsReporter:
|
||||
if not hasattr(scheduler, "server_args"):
|
||||
scheduler.server_args = types.SimpleNamespace(
|
||||
enable_metrics=False,
|
||||
enable_metrics_for_all_schedulers=False,
|
||||
kv_events_config=None,
|
||||
enable_mfu_metrics=False,
|
||||
enable_forward_pass_metrics=False,
|
||||
)
|
||||
if not hasattr(scheduler, "ps"):
|
||||
scheduler.ps = types.SimpleNamespace(attn_tp_rank=0, attn_cp_rank=0)
|
||||
if not hasattr(scheduler, "kv_events_publisher"):
|
||||
scheduler.kv_events_publisher = types.SimpleNamespace(
|
||||
init_kv_events=lambda *a, **kw: None,
|
||||
)
|
||||
if not hasattr(scheduler, "tp_workers"):
|
||||
scheduler.tp_workers = []
|
||||
if not hasattr(scheduler, "tp_worker"):
|
||||
scheduler.tp_worker = types.SimpleNamespace(
|
||||
model_runner=types.SimpleNamespace(),
|
||||
)
|
||||
if not hasattr(scheduler, "draft_worker"):
|
||||
scheduler.draft_worker = None
|
||||
context = types.SimpleNamespace(
|
||||
enable_metrics=False,
|
||||
is_stats_logging_rank=True,
|
||||
current_scheduler_metrics_enabled=False,
|
||||
enable_kv_cache_events=False,
|
||||
collector=None,
|
||||
)
|
||||
return SchedulerMetricsReporter(
|
||||
scheduler=scheduler,
|
||||
tp_rank=0,
|
||||
pp_rank=0,
|
||||
dp_rank=0,
|
||||
metrics_collector_context=context,
|
||||
metrics_collector=None,
|
||||
)
|
||||
|
||||
|
||||
class TestForwardPassMetrics(unittest.TestCase):
|
||||
def setUp(self):
|
||||
self.scheduler = _DummyScheduler()
|
||||
self.scheduler.enable_fpm = True
|
||||
self.scheduler = types.SimpleNamespace()
|
||||
self.scheduler._fpm_worker_id = "worker-7"
|
||||
self.scheduler._fpm_dp_rank = 0
|
||||
self.scheduler._fpm_publisher = _CollectingPublisher()
|
||||
@@ -100,6 +135,8 @@ class TestForwardPassMetrics(unittest.TestCase):
|
||||
self.scheduler._fpm_gpu_time_acc = 0.0
|
||||
self.scheduler.waiting_queue = []
|
||||
self.scheduler.disaggregation_mode = DisaggregationMode.NULL
|
||||
self.reporter = _make_reporter(self.scheduler)
|
||||
self.scheduler.enable_fpm = True
|
||||
|
||||
def _make_batch(self, **overrides):
|
||||
defaults = dict(
|
||||
@@ -135,10 +172,10 @@ class TestForwardPassMetrics(unittest.TestCase):
|
||||
)
|
||||
|
||||
with patch(
|
||||
"sglang.srt.observability.scheduler_metrics_mixin.time.monotonic",
|
||||
"sglang.srt.managers.scheduler_components.metrics_reporter.time.monotonic",
|
||||
return_value=104.5,
|
||||
):
|
||||
self.scheduler._emit_forward_pass_metrics(batch)
|
||||
self.reporter._emit_forward_pass_metrics(batch)
|
||||
|
||||
self.assertEqual(len(self.scheduler._fpm_publisher.metrics), 1)
|
||||
metrics = self.scheduler._fpm_publisher.metrics[0]
|
||||
@@ -158,12 +195,12 @@ class TestForwardPassMetrics(unittest.TestCase):
|
||||
def test_emit_uses_device_timer_gpu_time(self):
|
||||
self.scheduler._fpm_uses_device_timer = True
|
||||
self.scheduler._fpm_gpu_time_acc = 0.042
|
||||
self.scheduler.forward_pass_device_timer = types.SimpleNamespace(
|
||||
self.reporter.forward_pass_device_timer = types.SimpleNamespace(
|
||||
_report=lambda: None,
|
||||
)
|
||||
batch = self._make_batch()
|
||||
|
||||
self.scheduler._emit_forward_pass_metrics(batch)
|
||||
self.reporter._emit_forward_pass_metrics(batch)
|
||||
|
||||
self.assertEqual(len(self.scheduler._fpm_publisher.metrics), 1)
|
||||
self.assertAlmostEqual(
|
||||
@@ -174,12 +211,12 @@ class TestForwardPassMetrics(unittest.TestCase):
|
||||
def test_emit_skips_when_device_timer_zero(self):
|
||||
self.scheduler._fpm_uses_device_timer = True
|
||||
self.scheduler._fpm_gpu_time_acc = 0.0
|
||||
self.scheduler.forward_pass_device_timer = types.SimpleNamespace(
|
||||
self.reporter.forward_pass_device_timer = types.SimpleNamespace(
|
||||
_report=lambda: None,
|
||||
)
|
||||
batch = self._make_batch()
|
||||
|
||||
self.scheduler._emit_forward_pass_metrics(batch)
|
||||
self.reporter._emit_forward_pass_metrics(batch)
|
||||
|
||||
self.assertEqual(len(self.scheduler._fpm_publisher.metrics), 0)
|
||||
|
||||
@@ -187,10 +224,10 @@ class TestForwardPassMetrics(unittest.TestCase):
|
||||
batch = self._make_batch()
|
||||
|
||||
with patch(
|
||||
"sglang.srt.observability.scheduler_metrics_mixin.time.monotonic",
|
||||
"sglang.srt.managers.scheduler_components.metrics_reporter.time.monotonic",
|
||||
return_value=100.035,
|
||||
):
|
||||
self.scheduler._emit_forward_pass_metrics(batch, result=None)
|
||||
self.reporter._emit_forward_pass_metrics(batch, result=None)
|
||||
|
||||
self.assertEqual(len(self.scheduler._fpm_publisher.metrics), 1)
|
||||
self.assertAlmostEqual(
|
||||
@@ -205,10 +242,10 @@ class TestForwardPassMetrics(unittest.TestCase):
|
||||
batch = self._make_batch()
|
||||
|
||||
with patch(
|
||||
"sglang.srt.observability.scheduler_metrics_mixin.time.monotonic",
|
||||
"sglang.srt.managers.scheduler_components.metrics_reporter.time.monotonic",
|
||||
return_value=101.0,
|
||||
):
|
||||
self.scheduler._emit_forward_pass_metrics(batch)
|
||||
self.reporter._emit_forward_pass_metrics(batch)
|
||||
|
||||
metrics = self.scheduler._fpm_publisher.metrics[0]
|
||||
self.assertEqual(metrics.queued_requests.num_prefill_requests, 3)
|
||||
@@ -226,10 +263,10 @@ class TestForwardPassMetrics(unittest.TestCase):
|
||||
batch = self._make_batch()
|
||||
|
||||
with patch(
|
||||
"sglang.srt.observability.scheduler_metrics_mixin.time.monotonic",
|
||||
"sglang.srt.managers.scheduler_components.metrics_reporter.time.monotonic",
|
||||
return_value=101.0,
|
||||
):
|
||||
self.scheduler._emit_forward_pass_metrics(batch)
|
||||
self.reporter._emit_forward_pass_metrics(batch)
|
||||
|
||||
metrics = self.scheduler._fpm_publisher.metrics[0]
|
||||
self.assertEqual(metrics.queued_requests.num_prefill_requests, 0)
|
||||
@@ -237,7 +274,7 @@ class TestForwardPassMetrics(unittest.TestCase):
|
||||
self.assertEqual(metrics.queued_requests.sum_decode_kv_tokens, 15 + 30 + 45)
|
||||
|
||||
def test_init_metrics_uses_server_worker_id(self):
|
||||
scheduler = _DummyScheduler()
|
||||
scheduler = types.SimpleNamespace()
|
||||
scheduler.server_args = types.SimpleNamespace(
|
||||
enable_metrics=False,
|
||||
enable_metrics_for_all_schedulers=False,
|
||||
@@ -254,7 +291,7 @@ class TestForwardPassMetrics(unittest.TestCase):
|
||||
"sglang.srt.observability.forward_pass_metrics._FpmPublisherThread",
|
||||
_DummyPublisherThread,
|
||||
):
|
||||
scheduler.init_metrics(tp_rank=0, pp_rank=0, dp_rank=2)
|
||||
reporter = _make_reporter(scheduler)
|
||||
|
||||
self.assertTrue(scheduler.enable_fpm)
|
||||
self.assertEqual(scheduler._fpm_worker_id, "endpoint-42")
|
||||
@@ -265,7 +302,7 @@ class TestForwardPassMetrics(unittest.TestCase):
|
||||
self.assertIsNotNone(scheduler.server_args.forward_pass_metrics_ipc_name)
|
||||
|
||||
def test_init_fpm_disabled_on_non_last_pp_rank(self):
|
||||
scheduler = _DummyScheduler()
|
||||
scheduler = types.SimpleNamespace()
|
||||
scheduler.server_args = types.SimpleNamespace(
|
||||
enable_metrics=False,
|
||||
enable_metrics_for_all_schedulers=False,
|
||||
@@ -282,7 +319,7 @@ class TestForwardPassMetrics(unittest.TestCase):
|
||||
"sglang.srt.observability.forward_pass_metrics._FpmPublisherThread",
|
||||
_DummyPublisherThread,
|
||||
):
|
||||
scheduler.init_metrics(tp_rank=0, pp_rank=0, dp_rank=0)
|
||||
reporter = _make_reporter(scheduler)
|
||||
|
||||
self.assertFalse(scheduler.enable_fpm)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user