diff --git a/test/registered/disaggregation/test_disaggregation_decode_radix_cache.py b/test/registered/disaggregation/test_disaggregation_decode_radix_cache.py index dd173a7b4..a78a7a933 100644 --- a/test/registered/disaggregation/test_disaggregation_decode_radix_cache.py +++ b/test/registered/disaggregation/test_disaggregation_decode_radix_cache.py @@ -26,15 +26,7 @@ from sglang.test.test_utils import ( try_cached_model, ) -register_cuda_ci(est_time=509, stage="base-c", runner_config="8-gpu-h20") - - -def _has_nixl(): - try: - import nixl._api # noqa: F401 - except ImportError: - return False - return True +register_cuda_ci(est_time=170, stage="base-c", runner_config="8-gpu-h20") 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( is_in_ci() or _has_mooncake(), "Mooncake is required for decode radix cache disaggregation coverage.", diff --git a/test/registered/disaggregation/test_disaggregation_dp_attention.py b/test/registered/disaggregation/test_disaggregation_dp_attention.py index fcbab186f..e22b15e67 100644 --- a/test/registered/disaggregation/test_disaggregation_dp_attention.py +++ b/test/registered/disaggregation/test_disaggregation_dp_attention.py @@ -1,7 +1,6 @@ import unittest from types import SimpleNamespace -from sglang.benchmark.serving import run_benchmark from sglang.srt.environ import envs from sglang.test.ci.ci_register import register_cuda_ci 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 ( DEFAULT_MODEL_NAME_FOR_TEST_MLA, DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, - get_benchmark_args, popen_launch_pd_server, 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): - """PD-disagg + DP-attention e2e on `total_tokens` LB — the most complex - dispatch (token accounting + tie-break + estimated_tokens). Simpler - algorithms are unit-tested in + """The dispatch algorithm itself is covered in test/registered/unit/managers/test_data_parallel_controller.py. """ @@ -113,22 +109,6 @@ class TestDisaggregationDPAttention(PDDisaggregationServerBase): 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__": unittest.main() diff --git a/test/registered/disaggregation/test_disaggregation_nixl.py b/test/registered/disaggregation/test_disaggregation_nixl.py index 1820840a4..8754d5145 100644 --- a/test/registered/disaggregation/test_disaggregation_nixl.py +++ b/test/registered/disaggregation/test_disaggregation_nixl.py @@ -22,7 +22,7 @@ from sglang.test.test_utils import ( 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. NIXL_PREFILL_TP_SIZE = 4 @@ -198,12 +198,8 @@ class NixlPDDisaggregationServerBase(PDDisaggregationServerBase): "NIXL with the configured backend is required for this test.", ) class TestDisaggregationNixlBasic(NixlPDDisaggregationServerBase): - """Small NIXL PD E2E coverage. - - 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. + """Mooncake owns the broad disaggregation functional matrix in + test_disaggregation_basic.py; this is the NIXL-only subset. """ @classmethod @@ -256,21 +252,6 @@ class TestDisaggregationNixlBasic(NixlPDDisaggregationServerBase): self.assertEqual(len(output_logprobs), completion_tokens) 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): args = SimpleNamespace( base_url=f"http://{self.base_host}:{self.lb_port}", @@ -305,7 +286,10 @@ class TestDisaggregationNixlFailure(NixlPDDisaggregationServerBase): def setUpClass(cls): _require_configured_nixl_backend() 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) configure_nixl_pd_backend(cls) cls.launch_all() @@ -321,12 +305,12 @@ class TestDisaggregationNixlFailure(NixlPDDisaggregationServerBase): eval_name="gsm8k", api="completion", max_tokens=512, - num_examples=200, + num_examples=50, num_threads=128, ) - # Match TestDisaggregationMooncakeFailure: inject many transfer failures - # and tolerate eval/request errors as long as workers remain healthy. + # Tolerate eval/request errors; the gate is that workers stay healthy + # after the injected transfer failures, not the score itself. try: metrics = run_eval(args) print(f"Evaluation metrics: {metrics}") diff --git a/test/registered/disaggregation/test_disaggregation_pp.py b/test/registered/disaggregation/test_disaggregation_pp.py index cb40378c8..a08c678ff 100644 --- a/test/registered/disaggregation/test_disaggregation_pp.py +++ b/test/registered/disaggregation/test_disaggregation_pp.py @@ -14,84 +14,7 @@ from sglang.test.test_utils import ( try_cached_model, ) -register_cuda_ci(est_time=357, 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) +register_cuda_ci(est_time=250, stage="base-c", runner_config="8-gpu-h20") class TestDisaggregationPrefillPPDynamicChunkAccuracy(PDDisaggregationServerBase): diff --git a/test/registered/unit/managers/test_data_parallel_controller.py b/test/registered/unit/managers/test_data_parallel_controller.py index 85c373f45..c17ad1f71 100644 --- a/test/registered/unit/managers/test_data_parallel_controller.py +++ b/test/registered/unit/managers/test_data_parallel_controller.py @@ -1,16 +1,17 @@ """DPBudget + DataParallelController dispatch tests. -`total_tokens` (the most complex algorithm) is exercised end-to-end in -test/registered/disaggregation/test_disaggregation_dp_attention.py; its -tie-break on `total_requests` transitively covers that state. +The e2e counterpart, over the real scheduler load-report path, is +test/registered/disaggregation/test_disaggregation_dp_attention.py. Fragility: scheduler tests bypass `DataParallelController.__init__` via `__new__` and inject only the attrs the schedulers read (`workers`, `status`, `_active_workers`, `round_robin_counter`, `dp_budget`). Update `_make_controller` 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 from types import SimpleNamespace from unittest.mock import MagicMock @@ -279,5 +280,77 @@ class TestStatusAwarenessInconsistency(CustomTestCase): 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__": unittest.main()