feat: Priority-based scheduling optimization (including default priority, preemption toggle, priority-based metrics, etc.) (#17026)

Co-authored-by: hnyls2002 <lsyincs@gmail.com>
This commit is contained in:
zhuxinjie-nz
2026-03-05 06:52:08 +08:00
committed by GitHub
parent 9457c049e1
commit 28c931e1a5
8 changed files with 354 additions and 75 deletions

View File

@@ -27,7 +27,7 @@ class SchedulerDllmMixin:
def get_new_batch_dllm(self: Scheduler) -> Optional[ScheduleBatch]:
"""Generate a new batch for DLLM (Diffusion LLM) scheduling."""
if self.try_preemption:
if self.enable_priority_preemption:
self.running_batch.batch_is_full = False
# Early exit if batch is full or no requests available
@@ -82,7 +82,7 @@ class SchedulerDllmMixin:
if (
self.get_num_allocatable_reqs(running_bs) <= 0
and self.dllm_manager.is_empty()
and not self.try_preemption
and not self.enable_priority_preemption
):
self.running_batch.batch_is_full = True
return True
@@ -186,12 +186,8 @@ class SchedulerDllmMixin:
# Record prefill stats for logging after forward
from sglang.srt.observability.scheduler_metrics_mixin import PrefillStats
new_batch.prefill_stats = PrefillStats(
log_input_tokens=self.adder.log_input_tokens,
log_hit_tokens=self.adder.log_hit_tokens,
new_token_ratio=self.adder.new_token_ratio,
running_bs=len(self.running_batch.reqs),
num_new_seqs=len(can_run_list),
new_batch.prefill_stats = PrefillStats.from_adder(
self.adder, self.running_batch.reqs, self.enable_priority_scheduling
)
return new_batch
@@ -209,8 +205,9 @@ class SchedulerDllmMixin:
# Try preemption if batch is full
if self.running_batch.batch_is_full:
if not self.try_preemption or not adder.preempt_to_schedule(
req, self.server_args
if (
not self.enable_priority_preemption
or not adder.preempt_to_schedule(req, self.server_args)
):
break

View File

@@ -813,8 +813,13 @@ class Scheduler(
else "cpu"
),
)
# Enable preemption for priority scheduling.
self.try_preemption = self.enable_priority_scheduling
# NOTE: preemption is enabled by default for priority scheduling.
self.enable_priority_preemption = (
self.enable_priority_scheduling
and not self.server_args.disable_priority_preemption
)
self.init_new_token_ratio = min(
envs.SGLANG_INIT_NEW_TOKEN_RATIO.get()
* self.server_args.schedule_conservativeness,
@@ -1994,7 +1999,7 @@ class Scheduler(
for req in ready_grammar_requests:
self._add_request_to_queue(req)
if self.try_preemption:
if self.enable_priority_preemption:
# Reset batch_is_full to try preemption with a prefill adder.
self.running_batch.batch_is_full = False
@@ -2004,6 +2009,7 @@ class Scheduler(
return None
running_bs = len(self.running_batch.reqs)
# Ignore the check if self.chunked_req is not None.
# In the non-PP case, when self.chunked_req is not None, num_allocatable_reqs should always be greater than 0,
# as the space for the chunked requests has just been released.
@@ -2012,7 +2018,7 @@ class Scheduler(
if (
self.get_num_allocatable_reqs(running_bs) <= 0
and self.chunked_req is not None
and not self.try_preemption
and not self.enable_priority_preemption
):
self.running_batch.batch_is_full = True
return None
@@ -2090,8 +2096,9 @@ class Scheduler(
self.running_batch.batch_is_full = True
if self.running_batch.batch_is_full:
if not self.try_preemption or not adder.preempt_to_schedule(
req, self.server_args
if (
not self.enable_priority_preemption
or not adder.preempt_to_schedule(req, self.server_args)
):
break
@@ -2174,12 +2181,8 @@ class Scheduler(
new_batch.prepare_for_extend()
# Record prefill stats for logging after forward
new_batch.prefill_stats = PrefillStats(
log_input_tokens=adder.log_input_tokens,
log_hit_tokens=adder.log_hit_tokens,
new_token_ratio=adder.new_token_ratio,
running_bs=len(self.running_batch.reqs),
num_new_seqs=len(can_run_list),
new_batch.prefill_stats = PrefillStats.from_adder(
adder, self.running_batch.reqs, self.enable_priority_scheduling
)
# Mixed-style chunked prefill

View File

@@ -9,6 +9,7 @@ from sglang.srt.disaggregation.utils import DisaggregationMode
from sglang.srt.environ import envs
from sglang.srt.managers.schedule_batch import ScheduleBatch
from sglang.srt.mem_cache.session_aware_cache import SessionAwareCache
from sglang.srt.observability.metrics_collector import QueueCount
from sglang.srt.utils.common import ceil_align, raise_error_or_warn
from sglang.srt.utils.request_logger import disable_request_logging
from sglang.srt.utils.watchdog import WatchdogRaw
@@ -301,26 +302,31 @@ class SchedulerRuntimeCheckerMixin:
) = self._get_mamba_token_info()
else:
num_used, token_usage, _, _ = self._get_token_info()
num_running_reqs = len(self.running_batch.reqs)
self.stats.num_running_reqs = num_running_reqs
_enable_ps = self.enable_priority_scheduling
self.stats.num_running_reqs = QueueCount.from_reqs(
self.running_batch.reqs, _enable_ps
)
self.stats.num_used_tokens = num_used
self.stats.token_usage = round(token_usage, 2)
self.stats.gen_throughput = 0
self.stats.num_queue_reqs = len(self.waiting_queue)
self.stats.num_queue_reqs = QueueCount.from_reqs(
self.waiting_queue, _enable_ps
)
self.stats.num_grammar_queue_reqs = len(self.grammar_manager)
if self.disaggregation_mode == DisaggregationMode.PREFILL:
self.stats.num_prefill_prealloc_queue_reqs = len(
self.disagg_prefill_bootstrap_queue.queue
self.stats.num_prefill_prealloc_queue_reqs = QueueCount.from_reqs(
self.disagg_prefill_bootstrap_queue.queue, _enable_ps
)
self.stats.num_prefill_inflight_queue_reqs = len(
self.disagg_prefill_inflight_queue
self.stats.num_prefill_inflight_queue_reqs = QueueCount.from_reqs(
self.disagg_prefill_inflight_queue, _enable_ps
)
if self.disaggregation_mode == DisaggregationMode.DECODE:
self.stats.num_decode_prealloc_queue_reqs = len(
self.disagg_decode_prealloc_queue.queue
self.stats.num_decode_prealloc_queue_reqs = QueueCount.from_reqs(
self.disagg_decode_prealloc_queue.queue, _enable_ps
)
self.stats.num_decode_transfer_queue_reqs = len(
self.disagg_decode_transfer_queue.queue
self.stats.num_decode_transfer_queue_reqs = QueueCount.from_reqs(
self.disagg_decode_transfer_queue.queue, _enable_ps
)
self.metrics_collector.log_stats(self.stats)
self._publish_kv_events()

View File

@@ -233,6 +233,8 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi
self.context_len = self.model_config.context_len
self.image_token_id = self.model_config.image_token_id
self.max_req_input_len = None # Will be set later in engine.py
self.enable_priority_scheduling = server_args.enable_priority_scheduling
self.default_priority_value = server_args.default_priority_value
speculative_algorithm = SpeculativeAlgorithm.from_string(
server_args.speculative_algorithm
)
@@ -421,6 +423,8 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi
"model_name": self.server_args.served_model_name,
# TODO: Add lora name/path in the future,
}
if self.enable_priority_scheduling:
labels["priority"] = ""
if self.server_args.tokenizer_metrics_allowed_custom_labels:
for label in self.server_args.tokenizer_metrics_allowed_custom_labels:
labels[label] = ""
@@ -482,6 +486,7 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi
# Normalize the request
obj.normalize_batch_and_arguments()
self._set_default_priority(obj)
self._validate_rid(obj)
if isinstance(obj, GenerateReqInput) and obj.routed_dp_rank is not None:
@@ -1894,11 +1899,13 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi
)
custom_labels = getattr(state.obj, "custom_labels", None)
labels = (
{**self.metrics_collector.labels, **custom_labels}
if custom_labels
else self.metrics_collector.labels
)
labels = dict(self.metrics_collector.labels)
if custom_labels:
labels.update(custom_labels)
if self.enable_priority_scheduling:
priority = getattr(state.obj, "priority", None)
if priority is not None:
labels["priority"] = str(priority)
if (
state.time_stats.first_token_time == 0.0
and self.disaggregation_mode != DisaggregationMode.PREFILL
@@ -2367,6 +2374,15 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi
return span_attrs
def _set_default_priority(self, obj: Union[GenerateReqInput, EmbeddingReqInput]):
"""Set the default priority value."""
if (
self.enable_priority_scheduling
and obj.priority is None
and self.default_priority_value is not None
):
obj.priority = self.default_priority_value
class ServerStatus(Enum):
Up = "Up"

View File

@@ -13,12 +13,15 @@
# ==============================================================================
"""Utilities for Prometheus Metrics Collection."""
from __future__ import annotations
import dataclasses
import logging
import os
import time
from collections import Counter
from dataclasses import dataclass, field
from typing import Any, Dict, List, Optional, Union
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Set, Union
from sglang.srt.environ import envs
from sglang.srt.model_executor.forward_batch_info import ForwardMode
@@ -27,8 +30,12 @@ from sglang.srt.server_args import ServerArgs
from sglang.srt.utils import get_bool_env_var
from sglang.srt.utils.gauge_histogram import GaugeHistogram
SGLANG_TEST_REQUEST_TIME_STATS = get_bool_env_var("SGLANG_TEST_REQUEST_TIME_STATS")
if TYPE_CHECKING:
from prometheus_client import Gauge
from sglang.srt.managers.schedule_batch import Req
SGLANG_TEST_REQUEST_TIME_STATS = get_bool_env_var("SGLANG_TEST_REQUEST_TIME_STATS")
logger = logging.getLogger(__name__)
@@ -47,10 +54,30 @@ def get_histogram_conf_from_env(env_var_name: str) -> Optional[List[float]]:
return [float(x) for x in env_var_value.split(",")]
@dataclass
class QueueCount:
"""Holds both the total count and optional per-priority breakdown for a queue."""
total: int = 0
by_priority: Optional[Dict[int, int]] = None
@classmethod
def from_reqs(cls, reqs: List[Req], enable_priority_scheduling: bool = False):
# NOTE: If requests have priority=None (no --default-priority-value set),
# Counter will produce {None: N}, resulting in priority="None" Prometheus labels.
# Set --default-priority-value when enabling priority scheduling to avoid this.
by_priority = (
dict(Counter(req.priority for req in reqs))
if enable_priority_scheduling
else None
)
return cls(total=len(reqs), by_priority=by_priority)
@dataclass
class SchedulerStats:
# Basics
num_running_reqs: int = 0
num_running_reqs: QueueCount = field(default_factory=QueueCount)
num_used_tokens: int = 0
token_usage: float = 0.0
pending_prealloc_token_usage: float = 0.0
@@ -58,7 +85,7 @@ class SchedulerStats:
mamba_usage: float = 0.0
decode_sum_seq_lens: int = 0
gen_throughput: float = 0.0
num_queue_reqs: int = 0
num_queue_reqs: QueueCount = field(default_factory=QueueCount)
num_grammar_queue_reqs: int = 0
num_running_reqs_offline_batch: int = 0
cache_hit_rate: float = 0.0
@@ -74,10 +101,10 @@ class SchedulerStats:
num_paused_reqs: int = 0
# PD disaggregation
num_prefill_prealloc_queue_reqs: int = 0
num_prefill_inflight_queue_reqs: int = 0
num_decode_prealloc_queue_reqs: int = 0
num_decode_transfer_queue_reqs: int = 0
num_prefill_prealloc_queue_reqs: QueueCount = field(default_factory=QueueCount)
num_prefill_inflight_queue_reqs: QueueCount = field(default_factory=QueueCount)
num_decode_prealloc_queue_reqs: QueueCount = field(default_factory=QueueCount)
num_decode_transfer_queue_reqs: QueueCount = field(default_factory=QueueCount)
kv_transfer_speed_gb_s: float = 0.0
kv_transfer_latency_ms: float = 0.0
@@ -155,6 +182,7 @@ class SchedulerMetricsCollector:
self.enable_lora = enable_lora
self.enable_hierarchical_cache = enable_hierarchical_cache
self.last_log_time = time.perf_counter()
self._known_priorities: Set[int] = set()
self.num_running_reqs = Gauge(
name="sglang:num_running_reqs",
@@ -730,9 +758,22 @@ class SchedulerMetricsCollector:
multiprocess_mode="mostrecent",
)
def _log_gauge(self, gauge, data: Union[int, float]) -> None:
def _log_gauge(self, gauge: Gauge, data: Union[int, float, QueueCount]) -> None:
# Convenience function for logging to gauge.
gauge.labels(**self.labels).set(data)
if isinstance(data, QueueCount):
# NOTE: When priority scheduling is enabled, the total is recorded under
# priority="" (the default label value). Per-priority breakdowns are recorded
# with priority="<int>". Grafana queries should use priority="" for totals.
gauge.labels(**self.labels).set(data.total)
if data.by_priority is not None:
self._known_priorities.update(data.by_priority.keys())
for priority in self._known_priorities:
value = data.by_priority.get(priority, 0)
labels = dict(self.labels)
labels["priority"] = str(priority)
gauge.labels(**labels).set(value)
else:
gauge.labels(**self.labels).set(data)
def _log_histogram(self, histogram, data: Union[int, float]) -> None:
histogram.labels(**self.labels).observe(data)

View File

@@ -5,7 +5,7 @@ import logging
import time
from collections import defaultdict
from contextlib import contextmanager
from typing import TYPE_CHECKING, Optional, Union
from typing import TYPE_CHECKING, List, Optional, Union
from sglang.srt.disaggregation.kv_events import EventPublisherFactory, KVEventBatch
from sglang.srt.disaggregation.utils import DisaggregationMode
@@ -25,6 +25,7 @@ from sglang.srt.managers.scheduler import ScheduleBatch
from sglang.srt.managers.utils import GenerationBatchResult
from sglang.srt.observability.metrics_collector import (
DPCooperationInfo,
QueueCount,
SchedulerMetricsCollector,
SchedulerStats,
compute_routing_key_stats,
@@ -34,6 +35,8 @@ 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, Scheduler
logger = logging.getLogger(__name__)
@@ -50,9 +53,26 @@ class PrefillStats:
log_input_tokens: int
log_hit_tokens: int
new_token_ratio: float
running_bs: int
num_running_reqs: QueueCount
num_new_seqs: int # len(can_run_list)
@classmethod
def from_adder(
cls,
adder: PrefillAdder,
running_reqs: List[Req],
enable_priority_scheduling: bool = False,
):
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),
)
class KvMetrics:
def __init__(self):
@@ -115,6 +135,8 @@ class SchedulerMetricsMixin:
"pp_rank": pp_rank,
"moe_ep_rank": self.moe_ep_rank,
}
if self.enable_priority_scheduling:
labels["priority"] = ""
if dp_rank is not None:
labels["dp_rank"] = dp_rank
if self.server_args.extra_metric_labels:
@@ -214,7 +236,7 @@ class SchedulerMetricsMixin:
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.running_bs}, "
f"#running-req: {prefill_stats.num_running_reqs.total}, "
f"#queue-req: {len(self.waiting_queue)}, "
)
@@ -252,7 +274,7 @@ class SchedulerMetricsMixin:
prefill_stats.log_hit_tokens / total_tokens if total_tokens > 0 else 0.0
)
self.stats.num_running_reqs = prefill_stats.running_bs
self.stats.num_running_reqs = prefill_stats.num_running_reqs
self.stats.num_running_reqs_offline_batch = 0
self.stats.num_used_tokens = num_used
self.stats.token_usage = token_usage
@@ -260,7 +282,11 @@ class SchedulerMetricsMixin:
self.stats.swa_token_usage = swa_token_usage
if self.is_hybrid_ssm:
self.stats.mamba_usage = mamba_usage
self.stats.num_queue_reqs = len(self.waiting_queue)
_enable_ps = self.enable_priority_scheduling
self.stats.num_queue_reqs = QueueCount.from_reqs(
self.waiting_queue, _enable_ps
)
self.stats.num_grammar_queue_reqs = len(self.grammar_manager)
self.stats.cache_hit_rate = cache_hit_rate
@@ -273,20 +299,20 @@ class SchedulerMetricsMixin:
# PD disaggregation
if self.disaggregation_mode == DisaggregationMode.PREFILL:
self.stats.num_prefill_prealloc_queue_reqs = len(
self.disagg_prefill_bootstrap_queue.queue
self.stats.num_prefill_prealloc_queue_reqs = QueueCount.from_reqs(
self.disagg_prefill_bootstrap_queue.queue, _enable_ps
)
self.stats.num_prefill_inflight_queue_reqs = len(
self.disagg_prefill_inflight_queue
self.stats.num_prefill_inflight_queue_reqs = QueueCount.from_reqs(
self.disagg_prefill_inflight_queue, _enable_ps
)
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.disaggregation_mode == DisaggregationMode.DECODE:
self.stats.num_decode_prealloc_queue_reqs = len(
self.disagg_decode_prealloc_queue.queue
self.stats.num_decode_prealloc_queue_reqs = QueueCount.from_reqs(
self.disagg_decode_prealloc_queue.queue, _enable_ps
)
self.stats.num_decode_transfer_queue_reqs = len(
self.disagg_decode_transfer_queue.queue
self.stats.num_decode_transfer_queue_reqs = QueueCount.from_reqs(
self.disagg_decode_transfer_queue.queue, _enable_ps
)
# Others
@@ -410,8 +436,9 @@ class SchedulerMetricsMixin:
logger.info(msg)
if self.enable_metrics:
_enable_ps = self.enable_priority_scheduling
# Basics
self.stats.num_running_reqs = num_running_reqs
self.stats.num_running_reqs = QueueCount.from_reqs(batch.reqs, _enable_ps)
self.stats.num_running_reqs_offline_batch = num_running_reqs_offline_batch
self.stats.num_used_tokens = num_used
self.stats.token_usage = token_usage
@@ -421,7 +448,9 @@ class SchedulerMetricsMixin:
self.stats.mamba_usage = mamba_usage
self.stats.decode_sum_seq_lens = batch.seq_lens_cpu.sum().item()
self.stats.gen_throughput = self.last_gen_throughput
self.stats.num_queue_reqs = len(self.waiting_queue)
self.stats.num_queue_reqs = QueueCount.from_reqs(
self.waiting_queue, _enable_ps
)
self.stats.num_grammar_queue_reqs = len(self.grammar_manager)
self.stats.cache_hit_rate = cache_hit_rate
@@ -438,20 +467,19 @@ class SchedulerMetricsMixin:
# PD disaggregation
if self.disaggregation_mode == DisaggregationMode.PREFILL:
self.stats.num_prefill_prealloc_queue_reqs = len(
self.disagg_prefill_bootstrap_queue.queue
self.stats.num_prefill_prealloc_queue_reqs = QueueCount.from_reqs(
self.disagg_prefill_bootstrap_queue.queue, _enable_ps
)
self.stats.num_prefill_inflight_queue_reqs = len(
self.disagg_prefill_inflight_queue
self.stats.num_prefill_inflight_queue_reqs = QueueCount.from_reqs(
self.disagg_prefill_inflight_queue, _enable_ps
)
elif self.disaggregation_mode == DisaggregationMode.DECODE:
self.stats.num_decode_prealloc_queue_reqs = len(
self.disagg_decode_prealloc_queue.queue
self.stats.num_decode_prealloc_queue_reqs = QueueCount.from_reqs(
self.disagg_decode_prealloc_queue.queue, _enable_ps
)
self.stats.num_decode_transfer_queue_reqs = len(
self.disagg_decode_transfer_queue.queue
self.stats.num_decode_transfer_queue_reqs = QueueCount.from_reqs(
self.disagg_decode_transfer_queue.queue, _enable_ps
)
running_routing_keys = [r.routing_key for r in batch.reqs]
waiting_routing_keys = [r.routing_key for r in self.waiting_queue]
(
@@ -504,13 +532,13 @@ class SchedulerMetricsMixin:
return
kv_metrics = KvMetrics()
kv_metrics.request_active_slots = self.stats.num_running_reqs
kv_metrics.request_active_slots = self.stats.num_running_reqs.total
kv_metrics.request_total_slots = self.max_running_requests
kv_metrics.kv_active_blocks = int(
self.stats.token_usage * self.max_total_num_tokens
)
kv_metrics.kv_total_blocks = self.max_total_num_tokens
kv_metrics.num_requests_waiting = self.stats.num_queue_reqs
kv_metrics.num_requests_waiting = self.stats.num_queue_reqs.total
kv_metrics.gpu_cache_usage_perc = self.stats.token_usage
kv_metrics.gpu_prefix_cache_hit_rate = self.stats.cache_hit_rate
kv_metrics.data_parallel_rank = self.dp_rank if self.dp_rank is not None else 0
@@ -594,7 +622,7 @@ class SchedulerMetricsMixin:
and self.stats.max_running_requests_under_SLO > 0
):
self.stats.utilization = max(
self.stats.num_running_reqs
self.stats.num_running_reqs.total
/ self.stats.max_running_requests_under_SLO,
self.stats.token_usage / 0.9,
)

View File

@@ -333,6 +333,8 @@ class ServerArgs:
prefill_max_requests: Optional[int] = None
schedule_policy: str = "fcfs"
enable_priority_scheduling: bool = False
disable_priority_preemption: bool = False
default_priority_value: Optional[int] = None
abort_on_priority_when_disabled: bool = False
schedule_low_priority_values_first: bool = False
priority_scheduling_preemption_threshold: int = 10
@@ -3386,6 +3388,18 @@ class ServerArgs:
default=ServerArgs.enable_priority_scheduling,
help="Enable priority scheduling. Requests with higher priority integer values will be scheduled first by default.",
)
parser.add_argument(
"--disable-priority-preemption",
action="store_true",
default=ServerArgs.disable_priority_preemption,
help="Disable priority scheduling preemption.",
)
parser.add_argument(
"--default-priority-value",
type=int,
default=ServerArgs.default_priority_value,
help="Default priority for requests without explicit priority.",
)
parser.add_argument(
"--abort-on-priority-when-disabled",
action="store_true",