Support token low usage watermark in prefill delayer (#16814)
This commit is contained in:
@@ -224,6 +224,7 @@ class Envs:
|
||||
SGLANG_SCHEDULER_SKIP_ALL_GATHER = EnvBool(False)
|
||||
SGLANG_SCHEDULER_DECREASE_PREFILL_IDLE = EnvBool(False)
|
||||
SGLANG_PREFILL_DELAYER_MAX_DELAY_PASSES = EnvInt(30)
|
||||
SGLANG_PREFILL_DELAYER_TOKEN_USAGE_LOW_WATERMARK = EnvFloat(None)
|
||||
SGLANG_DATA_PARALLEL_BUDGET_INTERVAL = EnvInt(1)
|
||||
|
||||
# Test: pd-disaggregation
|
||||
|
||||
@@ -31,6 +31,7 @@ class _NegotiateOutput(NamedTuple):
|
||||
output_allow: bool
|
||||
output_reason: str
|
||||
num_prefillable: int
|
||||
num_token_watermark_force_allow: int
|
||||
|
||||
|
||||
class PrefillDelayer:
|
||||
@@ -41,15 +42,19 @@ class PrefillDelayer:
|
||||
cpu_group,
|
||||
server_args,
|
||||
max_delay_passes: int,
|
||||
token_usage_low_watermark: Optional[float],
|
||||
metrics_collector: Optional["SchedulerMetricsCollector"] = None,
|
||||
):
|
||||
self._max_delay_passes = max_delay_passes
|
||||
self._token_usage_low_watermark = token_usage_low_watermark
|
||||
logger.info(
|
||||
f"PrefillDelayer initialized with max_delay_passes={self._max_delay_passes}"
|
||||
f"PrefillDelayer initialized with "
|
||||
f"max_delay_passes={self._max_delay_passes} "
|
||||
f"token_usage_low_watermark={self._token_usage_low_watermark}"
|
||||
)
|
||||
|
||||
self._global_info_buffer = torch.empty(
|
||||
(dp_size, attn_tp_size, 1),
|
||||
(dp_size, attn_tp_size, 2),
|
||||
dtype=torch.int64,
|
||||
device="cpu",
|
||||
)
|
||||
@@ -70,33 +75,53 @@ class PrefillDelayer:
|
||||
), "To use PrefillDelayer, disable_overlap_schedule must be False."
|
||||
|
||||
def _negotiate_should_allow_prefill(
|
||||
self, local_prefillable: bool
|
||||
self, local_prefillable: bool, token_usage: float
|
||||
) -> _NegotiateOutput:
|
||||
out = self._negotiate_should_allow_prefill_pure(
|
||||
prev_state=self._curr_state,
|
||||
local_prefillable=local_prefillable,
|
||||
token_usage=token_usage,
|
||||
)
|
||||
self._curr_state = out.next_state
|
||||
return out
|
||||
|
||||
# (Almost) pure function, do not modify self state
|
||||
def _negotiate_should_allow_prefill_pure(
|
||||
self,
|
||||
prev_state: Optional[_State],
|
||||
local_prefillable: bool,
|
||||
token_usage: float,
|
||||
) -> _NegotiateOutput:
|
||||
global_prefillable = self._gather_info(local_prefillable=local_prefillable)
|
||||
# Compute local states
|
||||
local_token_watermark_force_allow = (
|
||||
local_prefillable
|
||||
and ((x := self._token_usage_low_watermark) is not None)
|
||||
and (token_usage < x)
|
||||
)
|
||||
|
||||
# Gather global states
|
||||
global_prefillable, global_token_watermark_force_allow = self._gather_info(
|
||||
local_prefillable=local_prefillable,
|
||||
local_token_watermark_force_allow=local_token_watermark_force_allow,
|
||||
)
|
||||
|
||||
# Compute derived global states
|
||||
if global_prefillable.min().item() > 0:
|
||||
prefillable_status = "all"
|
||||
elif global_prefillable.max().item() == 0:
|
||||
prefillable_status = "none"
|
||||
else:
|
||||
prefillable_status = "mixed"
|
||||
global_exists_token_watermark_force_allow = (
|
||||
global_token_watermark_force_allow.max().item() > 0
|
||||
)
|
||||
debug_info = dict(
|
||||
input_estimation=prefillable_status,
|
||||
num_prefillable=global_prefillable.sum().item(),
|
||||
num_token_watermark_force_allow=global_token_watermark_force_allow.sum().item(),
|
||||
)
|
||||
|
||||
# Compute outputs
|
||||
if prefillable_status == "all":
|
||||
exist_previous_wait = prev_state is not None
|
||||
return _NegotiateOutput(
|
||||
@@ -108,11 +133,20 @@ class PrefillDelayer:
|
||||
elif prefillable_status == "none":
|
||||
return _NegotiateOutput(
|
||||
next_state=None,
|
||||
# It does not matter whether we allow or not, thus we allow for simplicity
|
||||
output_allow=True,
|
||||
output_reason="",
|
||||
**debug_info,
|
||||
)
|
||||
elif prefillable_status == "mixed":
|
||||
if global_exists_token_watermark_force_allow:
|
||||
return _NegotiateOutput(
|
||||
next_state=None,
|
||||
output_allow=True,
|
||||
output_reason="token_watermark",
|
||||
**debug_info,
|
||||
)
|
||||
|
||||
prev_delayed_count = prev_state.delayed_count if prev_state else 0
|
||||
if prev_delayed_count < self._max_delay_passes - 1:
|
||||
next_state = prev_state or _State()
|
||||
@@ -133,9 +167,11 @@ class PrefillDelayer:
|
||||
else:
|
||||
raise NotImplementedError
|
||||
|
||||
def _gather_info(self, local_prefillable: bool):
|
||||
def _gather_info(
|
||||
self, local_prefillable: bool, local_token_watermark_force_allow: bool
|
||||
):
|
||||
local_info = torch.tensor(
|
||||
[int(local_prefillable)],
|
||||
[int(local_prefillable), int(local_token_watermark_force_allow)],
|
||||
device="cpu",
|
||||
dtype=torch.int64,
|
||||
)
|
||||
@@ -145,12 +181,13 @@ class PrefillDelayer:
|
||||
group=self._cpu_group,
|
||||
)
|
||||
tp0_info = self._global_info_buffer[:, 0, :]
|
||||
return tp0_info[:, 0]
|
||||
return tp0_info[:, 0], tp0_info[:, 1]
|
||||
|
||||
|
||||
class PrefillDelayerSinglePassExecutor:
|
||||
def __init__(self, prefill_delayer: PrefillDelayer):
|
||||
def __init__(self, prefill_delayer: PrefillDelayer, token_usage: float):
|
||||
self._prefill_delayer = prefill_delayer
|
||||
self._token_usage = token_usage
|
||||
self._result: Optional[_NegotiateOutput] = None
|
||||
|
||||
@property
|
||||
@@ -171,6 +208,7 @@ class PrefillDelayerSinglePassExecutor:
|
||||
if not self._called:
|
||||
self._result = self._prefill_delayer._negotiate_should_allow_prefill(
|
||||
local_prefillable=local_prefillable,
|
||||
token_usage=self._token_usage,
|
||||
)
|
||||
return self._result.output_allow
|
||||
|
||||
@@ -187,6 +225,13 @@ def _record_single_pass_result(
|
||||
f"(num_prefillable={output.num_prefillable}, "
|
||||
f"actual_execution={actual_execution})"
|
||||
)
|
||||
elif output.output_allow and (output.output_reason == "token_watermark"):
|
||||
logger.info(
|
||||
f"PrefillDelayer force allow prefill due to low watermark. "
|
||||
f"(num_prefillable={output.num_prefillable}, "
|
||||
f"num_token_watermark_force_allow={output.num_token_watermark_force_allow}, "
|
||||
f"actual_execution={actual_execution})"
|
||||
)
|
||||
else:
|
||||
assert output.output_reason in {
|
||||
"",
|
||||
|
||||
@@ -775,6 +775,9 @@ class Scheduler(
|
||||
self.metrics_collector if self.enable_metrics else None
|
||||
),
|
||||
max_delay_passes=envs.SGLANG_PREFILL_DELAYER_MAX_DELAY_PASSES.get(),
|
||||
token_usage_low_watermark=(
|
||||
envs.SGLANG_PREFILL_DELAYER_TOKEN_USAGE_LOW_WATERMARK.get()
|
||||
),
|
||||
)
|
||||
# Enable preemption for priority scheduling.
|
||||
self.try_preemption = self.enable_priority_scheduling
|
||||
@@ -1846,8 +1849,9 @@ class Scheduler(
|
||||
def get_new_batch_prefill(self) -> Optional[ScheduleBatch]:
|
||||
prefill_delayer_single_pass = None
|
||||
if self.prefill_delayer:
|
||||
_, token_usage, _, _ = self._get_token_info()
|
||||
prefill_delayer_single_pass = PrefillDelayerSinglePassExecutor(
|
||||
self.prefill_delayer
|
||||
self.prefill_delayer, token_usage=token_usage
|
||||
)
|
||||
|
||||
ret = self._get_new_batch_prefill_raw(
|
||||
|
||||
@@ -826,9 +826,9 @@ class SchedulerMetricsCollector:
|
||||
self.prefill_delayer_outcomes_total.labels(
|
||||
**self.labels,
|
||||
input_estimation=input_estimation,
|
||||
output_allow=str(int(output_allow)),
|
||||
output_allow=str(output_allow).lower(),
|
||||
output_reason=output_reason,
|
||||
actual_execution=str(int(actual_execution)),
|
||||
actual_execution=str(actual_execution).lower(),
|
||||
).inc(1)
|
||||
|
||||
def increment_retracted_reqs(
|
||||
|
||||
@@ -1,11 +1,14 @@
|
||||
import asyncio
|
||||
import os
|
||||
import re
|
||||
import time
|
||||
import unittest
|
||||
from collections import defaultdict
|
||||
from dataclasses import dataclass
|
||||
from types import SimpleNamespace
|
||||
from typing import List
|
||||
from typing import List, Optional
|
||||
|
||||
import openai
|
||||
import requests
|
||||
import torch
|
||||
import torch.multiprocessing as mp
|
||||
@@ -32,12 +35,14 @@ WORLD_SIZE = os.environ.get("SGLANG_TEST_WORLD_SIZE", "8")
|
||||
@dataclass
|
||||
class NegotiateCall:
|
||||
prefillable: List[bool]
|
||||
token_usage: List[float]
|
||||
|
||||
|
||||
@dataclass
|
||||
class NegotiateTestCase:
|
||||
name: str
|
||||
max_delay_passes: int
|
||||
token_usage_low_watermark: Optional[float]
|
||||
calls: List[NegotiateCall]
|
||||
expected_allow: bool
|
||||
expected_reason: str
|
||||
@@ -63,11 +68,13 @@ def _run_negotiate_test(rank, world_size, test_cases, results_queue, port):
|
||||
disable_overlap_schedule=False,
|
||||
),
|
||||
max_delay_passes=case.max_delay_passes,
|
||||
token_usage_low_watermark=case.token_usage_low_watermark,
|
||||
)
|
||||
|
||||
for call in case.calls:
|
||||
result = delayer._negotiate_should_allow_prefill(
|
||||
local_prefillable=call.prefillable[rank],
|
||||
token_usage=call.token_usage[rank],
|
||||
)
|
||||
|
||||
results_queue.put((rank, case.name, result.output_allow, result.output_reason))
|
||||
@@ -79,8 +86,12 @@ _NEGOTIATE_TEST_CASES = [
|
||||
NegotiateTestCase(
|
||||
name="all_prefillable",
|
||||
max_delay_passes=100,
|
||||
token_usage_low_watermark=0.8,
|
||||
calls=[
|
||||
NegotiateCall(prefillable=[True, True, True, True]),
|
||||
NegotiateCall(
|
||||
prefillable=[True, True, True, True],
|
||||
token_usage=[0.9, 0.9, 0.9, 0.9],
|
||||
)
|
||||
],
|
||||
expected_allow=True,
|
||||
expected_reason="no_wait",
|
||||
@@ -88,9 +99,16 @@ _NEGOTIATE_TEST_CASES = [
|
||||
NegotiateTestCase(
|
||||
name="all_prefillable_with_previous_wait",
|
||||
max_delay_passes=100,
|
||||
token_usage_low_watermark=0.8,
|
||||
calls=[
|
||||
NegotiateCall(prefillable=[True, False, True, False]),
|
||||
NegotiateCall(prefillable=[True, True, True, True]),
|
||||
NegotiateCall(
|
||||
prefillable=[True, False, True, False],
|
||||
token_usage=[0.9, 0.9, 0.9, 0.9],
|
||||
),
|
||||
NegotiateCall(
|
||||
prefillable=[True, True, True, True],
|
||||
token_usage=[0.9, 0.9, 0.9, 0.9],
|
||||
),
|
||||
],
|
||||
expected_allow=True,
|
||||
expected_reason="wait_success",
|
||||
@@ -98,8 +116,12 @@ _NEGOTIATE_TEST_CASES = [
|
||||
NegotiateTestCase(
|
||||
name="none_prefillable",
|
||||
max_delay_passes=100,
|
||||
token_usage_low_watermark=0.8,
|
||||
calls=[
|
||||
NegotiateCall(prefillable=[False, False, False, False]),
|
||||
NegotiateCall(
|
||||
prefillable=[False, False, False, False],
|
||||
token_usage=[0.9, 0.9, 0.9, 0.9],
|
||||
)
|
||||
],
|
||||
expected_allow=True,
|
||||
expected_reason="",
|
||||
@@ -107,8 +129,51 @@ _NEGOTIATE_TEST_CASES = [
|
||||
NegotiateTestCase(
|
||||
name="mixed_delay",
|
||||
max_delay_passes=100,
|
||||
token_usage_low_watermark=0.8,
|
||||
calls=[
|
||||
NegotiateCall(prefillable=[True, False, True, False]),
|
||||
NegotiateCall(
|
||||
prefillable=[True, False, True, False],
|
||||
token_usage=[0.9, 0.9, 0.9, 0.9],
|
||||
)
|
||||
],
|
||||
expected_allow=False,
|
||||
expected_reason="delay",
|
||||
),
|
||||
NegotiateTestCase(
|
||||
name="mixed_watermark_force_allow",
|
||||
max_delay_passes=100,
|
||||
token_usage_low_watermark=0.8,
|
||||
calls=[
|
||||
NegotiateCall(
|
||||
prefillable=[True, False, True, False],
|
||||
token_usage=[0.5, 0.9, 0.9, 0.9],
|
||||
)
|
||||
],
|
||||
expected_allow=True,
|
||||
expected_reason="token_watermark",
|
||||
),
|
||||
NegotiateTestCase(
|
||||
name="mixed_watermark_disabled",
|
||||
max_delay_passes=100,
|
||||
token_usage_low_watermark=None,
|
||||
calls=[
|
||||
NegotiateCall(
|
||||
prefillable=[True, False, True, False],
|
||||
token_usage=[0.5, 0.9, 0.9, 0.9],
|
||||
)
|
||||
],
|
||||
expected_allow=False,
|
||||
expected_reason="delay",
|
||||
),
|
||||
NegotiateTestCase(
|
||||
name="mixed_watermark_not_prefillable",
|
||||
max_delay_passes=100,
|
||||
token_usage_low_watermark=0.8,
|
||||
calls=[
|
||||
NegotiateCall(
|
||||
prefillable=[False, False, True, False],
|
||||
token_usage=[0.5, 0.9, 0.9, 0.9],
|
||||
)
|
||||
],
|
||||
expected_allow=False,
|
||||
expected_reason="delay",
|
||||
@@ -116,10 +181,20 @@ _NEGOTIATE_TEST_CASES = [
|
||||
NegotiateTestCase(
|
||||
name="mixed_timeout",
|
||||
max_delay_passes=3,
|
||||
token_usage_low_watermark=0.8,
|
||||
calls=[
|
||||
NegotiateCall(prefillable=[True, False, True, False]),
|
||||
NegotiateCall(prefillable=[True, False, True, False]),
|
||||
NegotiateCall(prefillable=[True, False, True, False]),
|
||||
NegotiateCall(
|
||||
prefillable=[True, False, True, False],
|
||||
token_usage=[0.9, 0.9, 0.9, 0.9],
|
||||
),
|
||||
NegotiateCall(
|
||||
prefillable=[True, False, True, False],
|
||||
token_usage=[0.9, 0.9, 0.9, 0.9],
|
||||
),
|
||||
NegotiateCall(
|
||||
prefillable=[True, False, True, False],
|
||||
token_usage=[0.9, 0.9, 0.9, 0.9],
|
||||
),
|
||||
],
|
||||
expected_allow=True,
|
||||
expected_reason="wait_timeout",
|
||||
@@ -172,6 +247,7 @@ class TestPrefillDelayerThroughputOnlineServing(CustomTestCase):
|
||||
self,
|
||||
test_name="online_serving",
|
||||
other_launch_args=[
|
||||
# Not really needed, only to test support non-FCFS algorithms
|
||||
"--schedule-policy",
|
||||
"lpm",
|
||||
],
|
||||
@@ -196,6 +272,7 @@ class TestPrefillDelayerThroughputOfflineGen(CustomTestCase):
|
||||
random_input_len=30000,
|
||||
random_output_len=500,
|
||||
),
|
||||
token_usage_low_watermark=0.8,
|
||||
min_improvement_pct=20,
|
||||
)
|
||||
|
||||
@@ -206,11 +283,13 @@ def _run_throughput_comparison(
|
||||
other_launch_args,
|
||||
other_benchmark_args,
|
||||
min_improvement_pct: float,
|
||||
token_usage_low_watermark: float = None,
|
||||
):
|
||||
common_kwargs = dict(
|
||||
debug_name=test_name,
|
||||
other_launch_args=other_launch_args,
|
||||
other_benchmark_args=other_benchmark_args,
|
||||
token_usage_low_watermark=token_usage_low_watermark,
|
||||
)
|
||||
res_enabled = _run_throughput_test(prefill_delayer=True, **common_kwargs)
|
||||
res_disabled = _run_throughput_test(prefill_delayer=False, **common_kwargs)
|
||||
@@ -229,15 +308,17 @@ def _run_throughput_test(
|
||||
prefill_delayer: bool,
|
||||
other_launch_args,
|
||||
other_benchmark_args,
|
||||
token_usage_low_watermark: float = None,
|
||||
):
|
||||
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,
|
||||
prefill_delayer=prefill_delayer,
|
||||
token_usage_low_watermark=token_usage_low_watermark,
|
||||
)
|
||||
|
||||
try:
|
||||
@@ -290,6 +371,87 @@ def _assert_throughput_improvement(
|
||||
)
|
||||
|
||||
|
||||
class TestPrefillDelayerTokenUsageLowWatermark(CustomTestCase):
|
||||
def test_1_with_low_watermark(self):
|
||||
# The kv cache size here is deliberately small, thus we use smaller token usage
|
||||
self._run(token_usage_low_watermark=0.5)
|
||||
|
||||
def test_2_without_low_watermark(self):
|
||||
self._run(token_usage_low_watermark=None)
|
||||
|
||||
def _run(self, token_usage_low_watermark):
|
||||
model = "Qwen/Qwen3-0.6B"
|
||||
base_url = DEFAULT_URL_FOR_TEST
|
||||
world_size = int(WORLD_SIZE)
|
||||
|
||||
process = _launch_server(
|
||||
model=model,
|
||||
base_url=base_url,
|
||||
prefill_delayer=True,
|
||||
other_args=["--max-total-tokens", "50000"],
|
||||
# e.g. gen throughput is 370 tok/s on H200.
|
||||
# Will need a different threshold on B200
|
||||
max_delay_passes=3000,
|
||||
token_usage_low_watermark=token_usage_low_watermark,
|
||||
)
|
||||
|
||||
async def run_test():
|
||||
client = openai.AsyncClient(base_url=f"{base_url}/v1", api_key="EMPTY")
|
||||
long_prompt = "Hello " * 5000
|
||||
|
||||
async def send_blocking_request():
|
||||
return await client.chat.completions.create(
|
||||
model=model,
|
||||
messages=[{"role": "user", "content": long_prompt}],
|
||||
max_tokens=10000,
|
||||
extra_body={"data_parallel_rank": 0},
|
||||
)
|
||||
|
||||
async def send_normal_request(dp_rank, req_idx):
|
||||
start = time.time()
|
||||
await client.chat.completions.create(
|
||||
model=model,
|
||||
messages=[{"role": "user", "content": "Say hi"}],
|
||||
max_tokens=10,
|
||||
extra_body={"data_parallel_rank": dp_rank},
|
||||
)
|
||||
elapsed = time.time() - start
|
||||
return dp_rank, req_idx, elapsed
|
||||
|
||||
asyncio.create_task(send_blocking_request())
|
||||
await asyncio.sleep(3)
|
||||
|
||||
num_reqs_per_rank = 10
|
||||
results = await asyncio.gather(
|
||||
*[
|
||||
send_normal_request(dp_rank, req_idx)
|
||||
for dp_rank in range(1, world_size)
|
||||
for req_idx in range(num_reqs_per_rank)
|
||||
]
|
||||
)
|
||||
|
||||
enabled = token_usage_low_watermark is not None
|
||||
thresh = 5
|
||||
for dp_rank, req_idx, elapsed in results:
|
||||
print(f"DP rank {dp_rank} req {req_idx} completed in {elapsed:.2f}s")
|
||||
self.assertTrue(
|
||||
(elapsed < thresh) if enabled else (elapsed > thresh),
|
||||
f"DP rank {dp_rank} req {req_idx}: elapsed={elapsed:.2f}s, thresh={thresh}, enabled={enabled}. "
|
||||
f"Maybe you need a different `max_delay_passes` when using hardware other than H200.",
|
||||
)
|
||||
|
||||
try:
|
||||
asyncio.run(run_test())
|
||||
|
||||
metrics_text = _print_prefill_delayer_metrics(base_url, expect_metrics=True)
|
||||
if token_usage_low_watermark is not None:
|
||||
total = _sum_prometheus_metric_values(metrics_text, "token_watermark")
|
||||
self.assertGreater(total, 0, "Expected token_watermark > 0")
|
||||
print(f"total token_watermark: {total}")
|
||||
finally:
|
||||
kill_process_tree(process.pid)
|
||||
|
||||
|
||||
class TestPrefillDelayerAccuracy(CustomTestCase):
|
||||
def test_1_mgsm_en_has_prefill_delayer(self):
|
||||
self._run_accuracy_test(prefill_delayer=True)
|
||||
@@ -305,8 +467,10 @@ class TestPrefillDelayerAccuracy(CustomTestCase):
|
||||
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",
|
||||
],
|
||||
@@ -334,12 +498,17 @@ def _launch_server(
|
||||
prefill_delayer: bool,
|
||||
other_args,
|
||||
max_delay_passes: int = 100,
|
||||
token_usage_low_watermark: float = None,
|
||||
):
|
||||
os.environ["SGLANG_PREFILL_DELAYER_DEBUG_LOG"] = "1"
|
||||
|
||||
with envs.SGLANG_SCHEDULER_DECREASE_PREFILL_IDLE.override(
|
||||
prefill_delayer
|
||||
), envs.SGLANG_PREFILL_DELAYER_MAX_DELAY_PASSES.override(max_delay_passes):
|
||||
), envs.SGLANG_PREFILL_DELAYER_MAX_DELAY_PASSES.override(
|
||||
max_delay_passes
|
||||
), envs.SGLANG_PREFILL_DELAYER_TOKEN_USAGE_LOW_WATERMARK.override(
|
||||
token_usage_low_watermark
|
||||
):
|
||||
return popen_launch_server(
|
||||
model,
|
||||
base_url,
|
||||
|
||||
Reference in New Issue
Block a user