[misc] Trim redundant variants from the 8-gpu-h20 disaggregation test suite (#39544)
This commit is contained in:
@@ -26,15 +26,7 @@ from sglang.test.test_utils import (
|
|||||||
try_cached_model,
|
try_cached_model,
|
||||||
)
|
)
|
||||||
|
|
||||||
register_cuda_ci(est_time=509, stage="base-c", runner_config="8-gpu-h20")
|
register_cuda_ci(est_time=170, stage="base-c", runner_config="8-gpu-h20")
|
||||||
|
|
||||||
|
|
||||||
def _has_nixl():
|
|
||||||
try:
|
|
||||||
import nixl._api # noqa: F401
|
|
||||||
except ImportError:
|
|
||||||
return False
|
|
||||||
return True
|
|
||||||
|
|
||||||
|
|
||||||
def _has_mooncake():
|
def _has_mooncake():
|
||||||
@@ -134,16 +126,6 @@ class DisaggregationDecodeRadixCacheTestMixin:
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
@unittest.skipUnless(
|
|
||||||
is_in_ci() or _has_nixl(),
|
|
||||||
"NIXL is required for decode radix cache disaggregation coverage.",
|
|
||||||
)
|
|
||||||
class TestDisaggregationDecodeRadixCacheNixl(
|
|
||||||
DisaggregationDecodeRadixCacheTestMixin, PDDisaggregationServerBase
|
|
||||||
):
|
|
||||||
transfer_backend_name = "nixl"
|
|
||||||
|
|
||||||
|
|
||||||
@unittest.skipUnless(
|
@unittest.skipUnless(
|
||||||
is_in_ci() or _has_mooncake(),
|
is_in_ci() or _has_mooncake(),
|
||||||
"Mooncake is required for decode radix cache disaggregation coverage.",
|
"Mooncake is required for decode radix cache disaggregation coverage.",
|
||||||
|
|||||||
@@ -1,7 +1,6 @@
|
|||||||
import unittest
|
import unittest
|
||||||
from types import SimpleNamespace
|
from types import SimpleNamespace
|
||||||
|
|
||||||
from sglang.benchmark.serving import run_benchmark
|
|
||||||
from sglang.srt.environ import envs
|
from sglang.srt.environ import envs
|
||||||
from sglang.test.ci.ci_register import register_cuda_ci
|
from sglang.test.ci.ci_register import register_cuda_ci
|
||||||
from sglang.test.run_eval import run_eval
|
from sglang.test.run_eval import run_eval
|
||||||
@@ -11,18 +10,15 @@ from sglang.test.server_fixtures.disaggregation_fixture import (
|
|||||||
from sglang.test.test_utils import (
|
from sglang.test.test_utils import (
|
||||||
DEFAULT_MODEL_NAME_FOR_TEST_MLA,
|
DEFAULT_MODEL_NAME_FOR_TEST_MLA,
|
||||||
DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
|
DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
|
||||||
get_benchmark_args,
|
|
||||||
popen_launch_pd_server,
|
popen_launch_pd_server,
|
||||||
try_cached_model,
|
try_cached_model,
|
||||||
)
|
)
|
||||||
|
|
||||||
register_cuda_ci(est_time=191, stage="base-c", runner_config="8-gpu-h20")
|
register_cuda_ci(est_time=140, stage="base-c", runner_config="8-gpu-h20")
|
||||||
|
|
||||||
|
|
||||||
class TestDisaggregationDPAttention(PDDisaggregationServerBase):
|
class TestDisaggregationDPAttention(PDDisaggregationServerBase):
|
||||||
"""PD-disagg + DP-attention e2e on `total_tokens` LB — the most complex
|
"""The dispatch algorithm itself is covered in
|
||||||
dispatch (token accounting + tie-break + estimated_tokens). Simpler
|
|
||||||
algorithms are unit-tested in
|
|
||||||
test/registered/unit/managers/test_data_parallel_controller.py.
|
test/registered/unit/managers/test_data_parallel_controller.py.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
@@ -113,22 +109,6 @@ class TestDisaggregationDPAttention(PDDisaggregationServerBase):
|
|||||||
|
|
||||||
self.assertGreater(metrics["score"], 0.60)
|
self.assertGreater(metrics["score"], 0.60)
|
||||||
|
|
||||||
def test_bench_serving(self):
|
|
||||||
args = get_benchmark_args(
|
|
||||||
base_url=f"http://{self.base_host}:{self.lb_port}",
|
|
||||||
dataset_name="random",
|
|
||||||
tokenizer=self.model,
|
|
||||||
num_prompts=1000,
|
|
||||||
random_input_len=4096,
|
|
||||||
random_output_len=1024,
|
|
||||||
request_rate=float("inf"),
|
|
||||||
max_concurrency=256,
|
|
||||||
)
|
|
||||||
result = run_benchmark(args)
|
|
||||||
|
|
||||||
self.assertLess(result["mean_tpot_ms"], 20)
|
|
||||||
self.assertEqual(result["completed"], 1000)
|
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
unittest.main()
|
unittest.main()
|
||||||
|
|||||||
@@ -22,7 +22,7 @@ from sglang.test.test_utils import (
|
|||||||
try_cached_model,
|
try_cached_model,
|
||||||
)
|
)
|
||||||
|
|
||||||
register_cuda_ci(est_time=506, stage="base-c", runner_config="8-gpu-h20")
|
register_cuda_ci(est_time=350, stage="base-c", runner_config="8-gpu-h20")
|
||||||
# base-c 8-GPU runner is required for TP4 prefill + TP4 decode.
|
# base-c 8-GPU runner is required for TP4 prefill + TP4 decode.
|
||||||
|
|
||||||
NIXL_PREFILL_TP_SIZE = 4
|
NIXL_PREFILL_TP_SIZE = 4
|
||||||
@@ -198,12 +198,8 @@ class NixlPDDisaggregationServerBase(PDDisaggregationServerBase):
|
|||||||
"NIXL with the configured backend is required for this test.",
|
"NIXL with the configured backend is required for this test.",
|
||||||
)
|
)
|
||||||
class TestDisaggregationNixlBasic(NixlPDDisaggregationServerBase):
|
class TestDisaggregationNixlBasic(NixlPDDisaggregationServerBase):
|
||||||
"""Small NIXL PD E2E coverage.
|
"""Mooncake owns the broad disaggregation functional matrix in
|
||||||
|
test_disaggregation_basic.py; this is the NIXL-only subset.
|
||||||
Mooncake already owns the broad disaggregation functional matrix in
|
|
||||||
test_disaggregation_basic.py. This class intentionally mirrors only the
|
|
||||||
subset that proves NIXL can launch, transfer KV, serve a request, return
|
|
||||||
logprobs, and keep all workers alive.
|
|
||||||
"""
|
"""
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
@@ -256,21 +252,6 @@ class TestDisaggregationNixlBasic(NixlPDDisaggregationServerBase):
|
|||||||
self.assertEqual(len(output_logprobs), completion_tokens)
|
self.assertEqual(len(output_logprobs), completion_tokens)
|
||||||
self.assertGreater(len(input_logprobs), 0)
|
self.assertGreater(len(input_logprobs), 0)
|
||||||
|
|
||||||
|
|
||||||
@unittest.skipUnless(
|
|
||||||
_HAS_CONFIGURED_NIXL_BACKEND,
|
|
||||||
"NIXL with the configured backend is required for this test.",
|
|
||||||
)
|
|
||||||
class TestDisaggregationNixlAccuracy(NixlPDDisaggregationServerBase):
|
|
||||||
@classmethod
|
|
||||||
def setUpClass(cls):
|
|
||||||
_require_configured_nixl_backend()
|
|
||||||
_clear_disagg_failure_env()
|
|
||||||
super().setUpClass()
|
|
||||||
cls.model = try_cached_model(DEFAULT_MODEL_NAME_FOR_TEST)
|
|
||||||
configure_nixl_pd_backend(cls)
|
|
||||||
cls.launch_all()
|
|
||||||
|
|
||||||
def test_gsm8k_accuracy(self):
|
def test_gsm8k_accuracy(self):
|
||||||
args = SimpleNamespace(
|
args = SimpleNamespace(
|
||||||
base_url=f"http://{self.base_host}:{self.lb_port}",
|
base_url=f"http://{self.base_host}:{self.lb_port}",
|
||||||
@@ -305,7 +286,10 @@ class TestDisaggregationNixlFailure(NixlPDDisaggregationServerBase):
|
|||||||
def setUpClass(cls):
|
def setUpClass(cls):
|
||||||
_require_configured_nixl_backend()
|
_require_configured_nixl_backend()
|
||||||
super().setUpClass()
|
super().setUpClass()
|
||||||
os.environ["SGLANG_TEST_DISAGG_FAILURE_PROB"] = "0.05"
|
# 0.2, not the mooncake twin's 0.05: this test never inspects the eval
|
||||||
|
# result, so the rate alone decides how much of the failure path runs;
|
||||||
|
# over 50 requests 0.05 would leave an 8% chance of exercising none.
|
||||||
|
os.environ["SGLANG_TEST_DISAGG_FAILURE_PROB"] = "0.2"
|
||||||
cls.model = try_cached_model(DEFAULT_MODEL_NAME_FOR_TEST)
|
cls.model = try_cached_model(DEFAULT_MODEL_NAME_FOR_TEST)
|
||||||
configure_nixl_pd_backend(cls)
|
configure_nixl_pd_backend(cls)
|
||||||
cls.launch_all()
|
cls.launch_all()
|
||||||
@@ -321,12 +305,12 @@ class TestDisaggregationNixlFailure(NixlPDDisaggregationServerBase):
|
|||||||
eval_name="gsm8k",
|
eval_name="gsm8k",
|
||||||
api="completion",
|
api="completion",
|
||||||
max_tokens=512,
|
max_tokens=512,
|
||||||
num_examples=200,
|
num_examples=50,
|
||||||
num_threads=128,
|
num_threads=128,
|
||||||
)
|
)
|
||||||
|
|
||||||
# Match TestDisaggregationMooncakeFailure: inject many transfer failures
|
# Tolerate eval/request errors; the gate is that workers stay healthy
|
||||||
# and tolerate eval/request errors as long as workers remain healthy.
|
# after the injected transfer failures, not the score itself.
|
||||||
try:
|
try:
|
||||||
metrics = run_eval(args)
|
metrics = run_eval(args)
|
||||||
print(f"Evaluation metrics: {metrics}")
|
print(f"Evaluation metrics: {metrics}")
|
||||||
|
|||||||
@@ -14,84 +14,7 @@ from sglang.test.test_utils import (
|
|||||||
try_cached_model,
|
try_cached_model,
|
||||||
)
|
)
|
||||||
|
|
||||||
register_cuda_ci(est_time=357, stage="base-c", runner_config="8-gpu-h20")
|
register_cuda_ci(est_time=250, stage="base-c", runner_config="8-gpu-h20")
|
||||||
|
|
||||||
|
|
||||||
class TestDisaggregationPrefillPPAccuracy(PDDisaggregationServerBase):
|
|
||||||
@classmethod
|
|
||||||
def setUpClass(cls):
|
|
||||||
super().setUpClass()
|
|
||||||
cls.model = try_cached_model(DEFAULT_MODEL_NAME_FOR_TEST)
|
|
||||||
|
|
||||||
# Non blocking start servers
|
|
||||||
cls.start_prefill()
|
|
||||||
cls.start_decode()
|
|
||||||
|
|
||||||
# Block until both
|
|
||||||
cls.wait_server_ready(cls.prefill_url + "/health", process=cls.process_prefill)
|
|
||||||
cls.wait_server_ready(cls.decode_url + "/health", process=cls.process_decode)
|
|
||||||
|
|
||||||
cls.launch_lb()
|
|
||||||
|
|
||||||
@classmethod
|
|
||||||
def start_prefill(cls):
|
|
||||||
prefill_args = [
|
|
||||||
"--trust-remote-code",
|
|
||||||
"--disaggregation-mode",
|
|
||||||
"prefill",
|
|
||||||
"--disaggregation-bootstrap-port",
|
|
||||||
cls.bootstrap_port,
|
|
||||||
"--tp-size",
|
|
||||||
"2",
|
|
||||||
"--pp-size",
|
|
||||||
"2",
|
|
||||||
"--disable-overlap-schedule",
|
|
||||||
]
|
|
||||||
prefill_args += cls.transfer_backend + cls.rdma_devices
|
|
||||||
cls.process_prefill = popen_launch_pd_server(
|
|
||||||
cls.model,
|
|
||||||
cls.prefill_url,
|
|
||||||
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
|
|
||||||
other_args=prefill_args,
|
|
||||||
)
|
|
||||||
|
|
||||||
@classmethod
|
|
||||||
def start_decode(cls):
|
|
||||||
decode_args = [
|
|
||||||
"--trust-remote-code",
|
|
||||||
"--disaggregation-mode",
|
|
||||||
"decode",
|
|
||||||
"--disaggregation-bootstrap-port",
|
|
||||||
cls.bootstrap_port,
|
|
||||||
"--tp-size",
|
|
||||||
"2",
|
|
||||||
"--base-gpu-id",
|
|
||||||
"4",
|
|
||||||
]
|
|
||||||
decode_args += cls.transfer_backend + cls.rdma_devices
|
|
||||||
cls.process_decode = popen_launch_pd_server(
|
|
||||||
cls.model,
|
|
||||||
cls.decode_url,
|
|
||||||
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
|
|
||||||
other_args=decode_args,
|
|
||||||
)
|
|
||||||
|
|
||||||
def test_gsm8k(self):
|
|
||||||
args = SimpleNamespace(
|
|
||||||
base_url=self.base_url,
|
|
||||||
model=self.model,
|
|
||||||
eval_name="gsm8k",
|
|
||||||
api="completion",
|
|
||||||
max_tokens=512,
|
|
||||||
num_examples=200,
|
|
||||||
num_threads=128,
|
|
||||||
)
|
|
||||||
metrics = run_eval(args)
|
|
||||||
print(f"{metrics=}")
|
|
||||||
|
|
||||||
self.assertGreater(metrics["score"], 0.24)
|
|
||||||
# Wait a little bit so that the memory check happens.
|
|
||||||
time.sleep(5)
|
|
||||||
|
|
||||||
|
|
||||||
class TestDisaggregationPrefillPPDynamicChunkAccuracy(PDDisaggregationServerBase):
|
class TestDisaggregationPrefillPPDynamicChunkAccuracy(PDDisaggregationServerBase):
|
||||||
|
|||||||
@@ -1,16 +1,17 @@
|
|||||||
"""DPBudget + DataParallelController dispatch tests.
|
"""DPBudget + DataParallelController dispatch tests.
|
||||||
|
|
||||||
`total_tokens` (the most complex algorithm) is exercised end-to-end in
|
The e2e counterpart, over the real scheduler load-report path, is
|
||||||
test/registered/disaggregation/test_disaggregation_dp_attention.py; its
|
test/registered/disaggregation/test_disaggregation_dp_attention.py.
|
||||||
tie-break on `total_requests` transitively covers that state.
|
|
||||||
|
|
||||||
Fragility: scheduler tests bypass `DataParallelController.__init__` via
|
Fragility: scheduler tests bypass `DataParallelController.__init__` via
|
||||||
`__new__` and inject only the attrs the schedulers read (`workers`, `status`,
|
`__new__` and inject only the attrs the schedulers read (`workers`, `status`,
|
||||||
`_active_workers`, `round_robin_counter`, `dp_budget`). Update `_make_controller`
|
`_active_workers`, `round_robin_counter`, `dp_budget`). Update `_make_controller`
|
||||||
if a scheduler starts reading another attr. `maybe_external_dp_rank_routing`
|
if a scheduler starts reading another attr. `maybe_external_dp_rank_routing`
|
||||||
is exercised as the real method, no mock.
|
is exercised as the real method, no mock. The refresh-throttle tests inject
|
||||||
|
`load_snapshot_reader` and `_last_refresh_time` on top of those.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
import time
|
||||||
import unittest
|
import unittest
|
||||||
from types import SimpleNamespace
|
from types import SimpleNamespace
|
||||||
from unittest.mock import MagicMock
|
from unittest.mock import MagicMock
|
||||||
@@ -279,5 +280,77 @@ class TestStatusAwarenessInconsistency(CustomTestCase):
|
|||||||
ctl.workers[2].send_pyobj.assert_called_once()
|
ctl.workers[2].send_pyobj.assert_called_once()
|
||||||
|
|
||||||
|
|
||||||
|
class TestRefreshLoadBudgetThrottle(CustomTestCase):
|
||||||
|
@staticmethod
|
||||||
|
def _controller_with_reader(dp_size, snapshots):
|
||||||
|
ctl = _make_controller(dp_size)
|
||||||
|
ctl.load_snapshot_reader = MagicMock()
|
||||||
|
ctl.load_snapshot_reader.read_all.return_value = snapshots
|
||||||
|
return ctl
|
||||||
|
|
||||||
|
def test_throttled_refresh_spreads_a_burst_across_ranks(self):
|
||||||
|
idle = [_load(dp_rank=i, timestamp=1.0, num_total_tokens=0) for i in range(4)]
|
||||||
|
ctl = self._controller_with_reader(dp_size=4, snapshots=idle)
|
||||||
|
# A refresh stamp in the future keeps every call inside the window, so
|
||||||
|
# the burst runs entirely on speculative counters.
|
||||||
|
ctl._last_refresh_time = time.perf_counter() + 3600.0
|
||||||
|
|
||||||
|
for _ in range(8):
|
||||||
|
ctl.refresh_load_budget()
|
||||||
|
ctl.total_tokens_scheduler(_req(input_ids=[0] * 100))
|
||||||
|
|
||||||
|
ctl.load_snapshot_reader.read_all.assert_not_called()
|
||||||
|
self.assertEqual(
|
||||||
|
ctl.dp_budget.total_tokens,
|
||||||
|
[200, 200, 200, 200],
|
||||||
|
"speculative increments should spread the burst evenly",
|
||||||
|
)
|
||||||
|
for i, worker in enumerate(ctl.workers):
|
||||||
|
self.assertEqual(
|
||||||
|
worker.send_pyobj.call_count, 2, f"worker {i} should get 2 of 8 reqs"
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_refresh_outside_window_overwrites_speculative_increments(self):
|
||||||
|
reported = [
|
||||||
|
_load(dp_rank=0, timestamp=2.0, num_total_tokens=10),
|
||||||
|
_load(dp_rank=1, timestamp=2.0, num_total_tokens=20),
|
||||||
|
]
|
||||||
|
ctl = self._controller_with_reader(dp_size=2, snapshots=reported)
|
||||||
|
ctl._last_refresh_time = 0.0 # window has long passed
|
||||||
|
ctl.dp_budget.total_tokens = [999, 999]
|
||||||
|
|
||||||
|
ctl.refresh_load_budget()
|
||||||
|
|
||||||
|
ctl.load_snapshot_reader.read_all.assert_called_once()
|
||||||
|
self.assertEqual(
|
||||||
|
ctl.dp_budget.total_tokens,
|
||||||
|
[10, 20],
|
||||||
|
"a fresh snapshot must replace the speculative state",
|
||||||
|
)
|
||||||
|
self.assertGreater(ctl._last_refresh_time, 0.0)
|
||||||
|
|
||||||
|
def test_unchanged_snapshot_does_not_reset_the_burst(self):
|
||||||
|
frozen = [_load(dp_rank=i, timestamp=1.0, num_total_tokens=0) for i in range(2)]
|
||||||
|
ctl = self._controller_with_reader(dp_size=2, snapshots=frozen)
|
||||||
|
ctl._last_refresh_time = 0.0
|
||||||
|
ctl.refresh_load_budget() # adopts timestamp 1.0
|
||||||
|
|
||||||
|
for _ in range(4):
|
||||||
|
ctl.total_tokens_scheduler(_req(input_ids=[0] * 50))
|
||||||
|
after_burst = list(ctl.dp_budget.total_tokens)
|
||||||
|
|
||||||
|
ctl._last_refresh_time = 0.0 # let the next refresh through the throttle
|
||||||
|
ctl.refresh_load_budget() # same timestamp -> update_budget skips it
|
||||||
|
|
||||||
|
self.assertEqual(
|
||||||
|
after_burst, [100, 100], "burst should have spread over both ranks"
|
||||||
|
)
|
||||||
|
self.assertEqual(
|
||||||
|
ctl.dp_budget.total_tokens,
|
||||||
|
after_burst,
|
||||||
|
"a stale-timestamp snapshot must not wipe the speculative state",
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
unittest.main()
|
unittest.main()
|
||||||
|
|||||||
Reference in New Issue
Block a user