Files
sglang/test/srt/test_prefill_delayer.py
2026-01-07 09:53:52 +08:00

187 lines
5.7 KiB
Python

import os
import unittest
from types import SimpleNamespace
import requests
from sglang.bench_serving import run_benchmark
from sglang.srt.environ import envs
from sglang.srt.utils import kill_process_tree
from sglang.test.run_eval import run_eval
from sglang.test.test_utils import (
DEFAULT_MLA_MODEL_NAME_FOR_TEST,
DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
DEFAULT_URL_FOR_TEST,
CustomTestCase,
get_benchmark_args,
popen_launch_server,
)
class TestPrefillDelayerThroughputOnlineServing(CustomTestCase):
def test_1_has_prefill_delayer(self):
self._run(prefill_delayer=True)
def test_2_no_prefill_delayer(self):
self._run(prefill_delayer=False)
def _run(self, prefill_delayer: bool):
_run_throughput_test(
debug_name=f"online_serving ({prefill_delayer=})",
prefill_delayer=prefill_delayer,
other_launch_args=[
# Not really needed, only to test support non-FCFS algorithms
"--schedule-policy",
"lpm",
],
other_benchmark_args=dict(
num_prompts=500,
random_input_len=30000,
random_output_len=256,
request_rate=32,
),
)
class TestPrefillDelayerThroughputOfflineGen(CustomTestCase):
def test_1_has_prefill_delayer(self):
self._run(prefill_delayer=True)
def test_2_no_prefill_delayer(self):
self._run(prefill_delayer=False)
def _run(self, prefill_delayer: bool):
_run_throughput_test(
debug_name=f"offline_gen ({prefill_delayer=})",
prefill_delayer=prefill_delayer,
other_benchmark_args=dict(
num_prompts=800,
random_input_len=30000,
random_output_len=500,
),
other_launch_args=[
"--max-total-tokens",
"200000",
],
)
def _run_throughput_test(
debug_name: str,
prefill_delayer: bool,
other_launch_args,
other_benchmark_args,
):
model = "Qwen/Qwen3-0.6B"
base_url = DEFAULT_URL_FOR_TEST
process = _launch_server(
prefill_delayer=prefill_delayer,
model=model,
base_url=base_url,
other_args=other_launch_args,
)
try:
args = get_benchmark_args(
base_url=base_url,
dataset_name="random",
tokenizer=model,
**other_benchmark_args,
)
res = run_benchmark(args)
_print_prefill_delayer_metrics(base_url, expect_metrics=prefill_delayer)
finally:
kill_process_tree(process.pid)
print(f"=== {debug_name} ===")
print(f"Input throughput: {res['input_throughput']:.2f} token/s")
print(f"Output throughput: {res['output_throughput']:.2f} token/s")
class TestPrefillDelayerAccuracy(CustomTestCase):
def test_1_mgsm_en_has_prefill_delayer(self):
self._run_accuracy_test(prefill_delayer=True)
def test_2_mgsm_en_no_prefill_delayer(self):
self._run_accuracy_test(prefill_delayer=False)
def _run_accuracy_test(self, prefill_delayer: bool):
model = DEFAULT_MLA_MODEL_NAME_FOR_TEST
base_url = DEFAULT_URL_FOR_TEST
process = _launch_server(
prefill_delayer=prefill_delayer,
model=model,
base_url=base_url,
other_args=[
# Not really needed, only to test support non-FCFS algorithms
"--schedule-policy",
"lpm",
# Use this to ensure prefill delayer will be run
"--max-total-tokens",
"4096",
],
)
try:
args = SimpleNamespace(
base_url=base_url,
model=model,
eval_name="mgsm_en",
num_examples=None,
num_threads=1024,
)
metrics = run_eval(args)
print(f"=== mgsm_en ({prefill_delayer=}) ===")
print(f"{metrics=}")
self.assertGreater(metrics["score"], 0.87)
finally:
kill_process_tree(process.pid)
def _launch_server(*, model, base_url, prefill_delayer: bool, other_args):
os.environ["SGLANG_PREFILL_DELAYER_DEBUG_LOG"] = "1"
world_size = os.environ.get("SGLANG_TEST_WORLD_SIZE", "8")
with envs.SGLANG_SCHEDULER_DECREASE_PREFILL_IDLE.override(
prefill_delayer
), envs.SGLANG_PREFILL_DELAYER_MAX_DELAY_PASSES.override(100):
return popen_launch_server(
model,
base_url,
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
other_args=[
"--trust-remote-code",
"--tp",
world_size,
"--enable-dp-attention",
"--dp",
world_size,
"--chunked-prefill-size",
"131072",
"--mem-fraction-static",
"0.6",
"--enable-metrics",
*(other_args or []),
],
)
def _print_prefill_delayer_metrics(base_url: str, expect_metrics: bool):
metrics_response = requests.get(f"{base_url}/metrics")
assert metrics_response.status_code == 200
metrics_text = metrics_response.text
prefill_delayer_metrics = [
line for line in metrics_text.split("\n") if "prefill_delayer" in line
]
print("=== PrefillDelayer Metrics ===")
for line in prefill_delayer_metrics:
print(line)
if expect_metrics:
assert "sglang:prefill_delayer_wait_forward_passes" in metrics_text
assert "sglang:prefill_delayer_wait_seconds" in metrics_text
assert "sglang:prefill_delayer_timeouts_total" in metrics_text
if __name__ == "__main__":
unittest.main()