Tiny add command line args for prefill delayer and unify names (#16830)
This commit is contained in:
@@ -489,6 +489,18 @@ _warn_deprecated_env_to_cli_flag(
|
|||||||
"SGLANG_SUPPORT_CUTLASS_BLOCK_FP8",
|
"SGLANG_SUPPORT_CUTLASS_BLOCK_FP8",
|
||||||
"It will be completely removed in 0.5.7. Please use '--fp8-gemm-backend=cutlass' instead.",
|
"It will be completely removed in 0.5.7. Please use '--fp8-gemm-backend=cutlass' instead.",
|
||||||
)
|
)
|
||||||
|
_warn_deprecated_env_to_cli_flag(
|
||||||
|
"SGLANG_SCHEDULER_DECREASE_PREFILL_IDLE",
|
||||||
|
"Please use '--enable-prefill-delayer' instead.",
|
||||||
|
)
|
||||||
|
_warn_deprecated_env_to_cli_flag(
|
||||||
|
"SGLANG_PREFILL_DELAYER_MAX_DELAY_PASSES",
|
||||||
|
"Please use '--prefill-delayer-max-delay-passes' instead.",
|
||||||
|
)
|
||||||
|
_warn_deprecated_env_to_cli_flag(
|
||||||
|
"SGLANG_PREFILL_DELAYER_TOKEN_USAGE_LOW_WATERMARK",
|
||||||
|
"Please use '--prefill-delayer-token-usage-low-watermark' instead.",
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def example_with_exit_stack():
|
def example_with_exit_stack():
|
||||||
|
|||||||
@@ -765,7 +765,7 @@ class Scheduler(
|
|||||||
self.schedule_low_priority_values_first,
|
self.schedule_low_priority_values_first,
|
||||||
)
|
)
|
||||||
self.prefill_delayer: Optional[PrefillDelayer] = None
|
self.prefill_delayer: Optional[PrefillDelayer] = None
|
||||||
if envs.SGLANG_SCHEDULER_DECREASE_PREFILL_IDLE.get():
|
if self.server_args.enable_prefill_delayer:
|
||||||
self.prefill_delayer = PrefillDelayer(
|
self.prefill_delayer = PrefillDelayer(
|
||||||
dp_size=self.dp_size,
|
dp_size=self.dp_size,
|
||||||
attn_tp_size=self.attn_tp_size,
|
attn_tp_size=self.attn_tp_size,
|
||||||
@@ -774,10 +774,8 @@ class Scheduler(
|
|||||||
metrics_collector=(
|
metrics_collector=(
|
||||||
self.metrics_collector if self.enable_metrics else None
|
self.metrics_collector if self.enable_metrics else None
|
||||||
),
|
),
|
||||||
max_delay_passes=envs.SGLANG_PREFILL_DELAYER_MAX_DELAY_PASSES.get(),
|
max_delay_passes=self.server_args.prefill_delayer_max_delay_passes,
|
||||||
token_usage_low_watermark=(
|
token_usage_low_watermark=self.server_args.prefill_delayer_token_usage_low_watermark,
|
||||||
envs.SGLANG_PREFILL_DELAYER_TOKEN_USAGE_LOW_WATERMARK.get()
|
|
||||||
),
|
|
||||||
)
|
)
|
||||||
# Enable preemption for priority scheduling.
|
# Enable preemption for priority scheduling.
|
||||||
self.try_preemption = self.enable_priority_scheduling
|
self.try_preemption = self.enable_priority_scheduling
|
||||||
|
|||||||
@@ -95,7 +95,9 @@ class SchedulerMetricsMixin:
|
|||||||
if dp_rank is not None:
|
if dp_rank is not None:
|
||||||
labels["dp_rank"] = dp_rank
|
labels["dp_rank"] = dp_rank
|
||||||
self.metrics_collector = SchedulerMetricsCollector(
|
self.metrics_collector = SchedulerMetricsCollector(
|
||||||
labels=labels, enable_lora=self.enable_lora
|
labels=labels,
|
||||||
|
enable_lora=self.enable_lora,
|
||||||
|
prefill_delayer_max_delay_passes=self.server_args.prefill_delayer_max_delay_passes,
|
||||||
)
|
)
|
||||||
|
|
||||||
if ENABLE_METRICS_DEVICE_TIMER:
|
if ENABLE_METRICS_DEVICE_TIMER:
|
||||||
|
|||||||
@@ -271,6 +271,7 @@ class SchedulerMetricsCollector:
|
|||||||
self,
|
self,
|
||||||
labels: Dict[str, str],
|
labels: Dict[str, str],
|
||||||
enable_lora: bool = False,
|
enable_lora: bool = False,
|
||||||
|
prefill_delayer_max_delay_passes: int = 30,
|
||||||
) -> None:
|
) -> None:
|
||||||
# We need to import prometheus_client after setting the env variable `PROMETHEUS_MULTIPROC_DIR`
|
# We need to import prometheus_client after setting the env variable `PROMETHEUS_MULTIPROC_DIR`
|
||||||
from prometheus_client import Counter, Gauge, Histogram, Summary
|
from prometheus_client import Counter, Gauge, Histogram, Summary
|
||||||
@@ -761,13 +762,12 @@ class SchedulerMetricsCollector:
|
|||||||
labelnames=list(labels.keys()) + ["category", "num_prefill_ranks"],
|
labelnames=list(labels.keys()) + ["category", "num_prefill_ranks"],
|
||||||
)
|
)
|
||||||
|
|
||||||
max_delay_passes = envs.SGLANG_PREFILL_DELAYER_MAX_DELAY_PASSES.get()
|
|
||||||
self.prefill_delayer_wait_forward_passes = Histogram(
|
self.prefill_delayer_wait_forward_passes = Histogram(
|
||||||
name="sglang:prefill_delayer_wait_forward_passes",
|
name="sglang:prefill_delayer_wait_forward_passes",
|
||||||
documentation="Histogram of forward passes waited by prefill delayer.",
|
documentation="Histogram of forward passes waited by prefill delayer.",
|
||||||
labelnames=labels.keys(),
|
labelnames=labels.keys(),
|
||||||
# Need bucket "<=0" for zero-delay cases
|
# Need bucket "<=0" for zero-delay cases
|
||||||
buckets=[0, 5, 20, max_delay_passes - 1],
|
buckets=[0, 5, 20, prefill_delayer_max_delay_passes - 1],
|
||||||
)
|
)
|
||||||
self.prefill_delayer_wait_seconds = Histogram(
|
self.prefill_delayer_wait_seconds = Histogram(
|
||||||
name="sglang:prefill_delayer_wait_seconds",
|
name="sglang:prefill_delayer_wait_seconds",
|
||||||
|
|||||||
@@ -311,6 +311,9 @@ class ServerArgs:
|
|||||||
swa_full_tokens_ratio: float = 0.8
|
swa_full_tokens_ratio: float = 0.8
|
||||||
disable_hybrid_swa_memory: bool = False
|
disable_hybrid_swa_memory: bool = False
|
||||||
radix_eviction_policy: str = "lru"
|
radix_eviction_policy: str = "lru"
|
||||||
|
enable_prefill_delayer: bool = False
|
||||||
|
prefill_delayer_max_delay_passes: int = 30
|
||||||
|
prefill_delayer_token_usage_low_watermark: Optional[float] = None
|
||||||
|
|
||||||
# Runtime options
|
# Runtime options
|
||||||
device: Optional[str] = None
|
device: Optional[str] = None
|
||||||
@@ -665,6 +668,9 @@ class ServerArgs:
|
|||||||
# Handle deprecated arguments.
|
# Handle deprecated arguments.
|
||||||
self._handle_deprecated_args()
|
self._handle_deprecated_args()
|
||||||
|
|
||||||
|
# Handle deprecated environment variables for prefill delayer.
|
||||||
|
self._handle_prefill_delayer_env_compat()
|
||||||
|
|
||||||
# Set missing default values.
|
# Set missing default values.
|
||||||
self._handle_missing_default_values()
|
self._handle_missing_default_values()
|
||||||
|
|
||||||
@@ -778,6 +784,14 @@ class ServerArgs:
|
|||||||
)
|
)
|
||||||
self.tool_call_parser = deprecated_tool_call_parsers[self.tool_call_parser]
|
self.tool_call_parser = deprecated_tool_call_parsers[self.tool_call_parser]
|
||||||
|
|
||||||
|
def _handle_prefill_delayer_env_compat(self):
|
||||||
|
if envs.SGLANG_SCHEDULER_DECREASE_PREFILL_IDLE.get():
|
||||||
|
self.enable_prefill_delayer = True
|
||||||
|
if x := envs.SGLANG_PREFILL_DELAYER_MAX_DELAY_PASSES.get():
|
||||||
|
self.prefill_delayer_max_delay_passes = x
|
||||||
|
if x := envs.SGLANG_PREFILL_DELAYER_TOKEN_USAGE_LOW_WATERMARK.get():
|
||||||
|
self.prefill_delayer_token_usage_low_watermark = x
|
||||||
|
|
||||||
def _handle_missing_default_values(self):
|
def _handle_missing_default_values(self):
|
||||||
if self.tokenizer_path is None:
|
if self.tokenizer_path is None:
|
||||||
self.tokenizer_path = self.model_path
|
self.tokenizer_path = self.model_path
|
||||||
@@ -2884,6 +2898,23 @@ class ServerArgs:
|
|||||||
default=ServerArgs.radix_eviction_policy,
|
default=ServerArgs.radix_eviction_policy,
|
||||||
help="The eviction policy of radix trees. 'lru' stands for Least Recently Used, 'lfu' stands for Least Frequently Used.",
|
help="The eviction policy of radix trees. 'lru' stands for Least Recently Used, 'lfu' stands for Least Frequently Used.",
|
||||||
)
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--enable-prefill-delayer",
|
||||||
|
action="store_true",
|
||||||
|
help="Enable prefill delayer for DP attention to reduce idle time.",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--prefill-delayer-max-delay-passes",
|
||||||
|
type=int,
|
||||||
|
default=ServerArgs.prefill_delayer_max_delay_passes,
|
||||||
|
help="Maximum forward passes to delay prefill.",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--prefill-delayer-token-usage-low-watermark",
|
||||||
|
type=float,
|
||||||
|
default=None,
|
||||||
|
help="Token usage low watermark for prefill delayer.",
|
||||||
|
)
|
||||||
|
|
||||||
# Runtime options
|
# Runtime options
|
||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
|
|||||||
@@ -14,7 +14,6 @@ import torch
|
|||||||
import torch.multiprocessing as mp
|
import torch.multiprocessing as mp
|
||||||
|
|
||||||
from sglang.bench_serving import run_benchmark
|
from sglang.bench_serving import run_benchmark
|
||||||
from sglang.srt.environ import envs
|
|
||||||
from sglang.srt.managers.prefill_delayer import PrefillDelayer
|
from sglang.srt.managers.prefill_delayer import PrefillDelayer
|
||||||
from sglang.srt.utils import kill_process_tree
|
from sglang.srt.utils import kill_process_tree
|
||||||
from sglang.test.run_eval import run_eval
|
from sglang.test.run_eval import run_eval
|
||||||
@@ -502,32 +501,36 @@ def _launch_server(
|
|||||||
):
|
):
|
||||||
os.environ["SGLANG_PREFILL_DELAYER_DEBUG_LOG"] = "1"
|
os.environ["SGLANG_PREFILL_DELAYER_DEBUG_LOG"] = "1"
|
||||||
|
|
||||||
with envs.SGLANG_SCHEDULER_DECREASE_PREFILL_IDLE.override(
|
return popen_launch_server(
|
||||||
prefill_delayer
|
model,
|
||||||
), envs.SGLANG_PREFILL_DELAYER_MAX_DELAY_PASSES.override(
|
base_url,
|
||||||
max_delay_passes
|
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
|
||||||
), envs.SGLANG_PREFILL_DELAYER_TOKEN_USAGE_LOW_WATERMARK.override(
|
other_args=[
|
||||||
token_usage_low_watermark
|
"--trust-remote-code",
|
||||||
):
|
"--tp",
|
||||||
return popen_launch_server(
|
WORLD_SIZE,
|
||||||
model,
|
"--enable-dp-attention",
|
||||||
base_url,
|
"--dp",
|
||||||
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
|
WORLD_SIZE,
|
||||||
other_args=[
|
"--chunked-prefill-size",
|
||||||
"--trust-remote-code",
|
"131072",
|
||||||
"--tp",
|
"--mem-fraction-static",
|
||||||
WORLD_SIZE,
|
"0.6",
|
||||||
"--enable-dp-attention",
|
"--enable-metrics",
|
||||||
"--dp",
|
*(["--enable-prefill-delayer"] if prefill_delayer else []),
|
||||||
WORLD_SIZE,
|
"--prefill-delayer-max-delay-passes",
|
||||||
"--chunked-prefill-size",
|
str(max_delay_passes),
|
||||||
"131072",
|
*(
|
||||||
"--mem-fraction-static",
|
[
|
||||||
"0.6",
|
"--prefill-delayer-token-usage-low-watermark",
|
||||||
"--enable-metrics",
|
str(token_usage_low_watermark),
|
||||||
*(other_args or []),
|
]
|
||||||
],
|
if token_usage_low_watermark is not None
|
||||||
)
|
else []
|
||||||
|
),
|
||||||
|
*(other_args or []),
|
||||||
|
],
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def _print_prefill_delayer_metrics(base_url: str, expect_metrics: bool) -> str:
|
def _print_prefill_delayer_metrics(base_url: str, expect_metrics: bool) -> str:
|
||||||
|
|||||||
Reference in New Issue
Block a user