diff --git a/python/sglang/test/server_fixtures/disaggregation_fixture.py b/python/sglang/test/server_fixtures/disaggregation_fixture.py index d9313ad38..091c26471 100644 --- a/python/sglang/test/server_fixtures/disaggregation_fixture.py +++ b/python/sglang/test/server_fixtures/disaggregation_fixture.py @@ -7,6 +7,8 @@ import warnings from typing import ClassVar, Optional from urllib.parse import urlparse +import requests + from sglang.srt.environ import envs from sglang.srt.utils import kill_process_tree from sglang.test.test_utils import ( @@ -23,6 +25,27 @@ from sglang.utils import wait_for_http_ready logger = logging.getLogger(__name__) +def configure_nixl_pd_backend(test_cls): + test_cls.transfer_backend = ["--disaggregation-transfer-backend", "nixl"] + # NIXL backend/network selection is driven by NIXL environment variables + # such as SGLANG_DISAGGREGATION_NIXL_BACKEND and backend params, not by the + # Mooncake-specific --disaggregation-ib-device argument. + test_cls.rdma_devices = [] + + +def assert_process_healthy(test_case, name, process, url, health_path="/health"): + test_case.assertIsNotNone(process, f"{name} process was not started") + test_case.assertIsNone( + process.poll(), + f"{name} exited unexpectedly with code {process.returncode}", + ) + try: + response = requests.get(f"{url}{health_path}", timeout=10) + except requests.RequestException as e: + test_case.fail(f"Failed to connect to {name} health endpoint: {e}") + test_case.assertEqual(response.status_code, 200, response.text) + + class PDDisaggregationServerBase(CustomTestCase): capture_per_side_logs: ClassVar[bool] = False extra_prefill_env: ClassVar[dict[str, str]] = {} diff --git a/test/registered/disaggregation/test_disaggregation_nixl.py b/test/registered/disaggregation/test_disaggregation_nixl.py new file mode 100644 index 000000000..ca5bf09f9 --- /dev/null +++ b/test/registered/disaggregation/test_disaggregation_nixl.py @@ -0,0 +1,347 @@ +import json +import os +import unittest +import uuid +from types import SimpleNamespace + +import requests + +from sglang.srt.environ import envs +from sglang.test.ci.ci_register import register_cuda_ci +from sglang.test.run_eval import run_eval +from sglang.test.server_fixtures.disaggregation_fixture import ( + PDDisaggregationServerBase, + assert_process_healthy, + configure_nixl_pd_backend, +) +from sglang.test.test_utils import ( + DEFAULT_MODEL_NAME_FOR_TEST, + DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, + is_in_ci, + popen_launch_pd_server, + try_cached_model, +) + +register_cuda_ci(est_time=700, 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 +NIXL_DECODE_TP_SIZE = 4 +NIXL_DECODE_BASE_GPU_ID = 4 +NIXL_PREFILL_UCX_NUM_THREADS = 2 + +# This is a PD transfer functional gate, not the standalone model-quality gate. +# The standalone Llama-3.1-8B GSM8K threshold is 0.80, while existing PD +# disaggregation coverage gates at 0.62. Keep NIXL above catastrophic transfer +# failure while leaving margin for PD/router/evaluator variance on 200 examples. +NIXL_GSM8K_SCORE_THRESHOLD = 0.62 + + +def _nixl_backend_config( + backend, + backend_params_json, + ucx_num_threads=NIXL_PREFILL_UCX_NUM_THREADS, +): + backend_params = json.loads(backend_params_json) + if not isinstance(backend_params, dict) or not all( + isinstance(key, str) and isinstance(value, str) + for key, value in backend_params.items() + ): + raise ValueError( + "SGLANG_DISAGGREGATION_NIXL_BACKEND_PARAMS must be a JSON object " + "with string keys and string values" + ) + + if backend == "UCX": + backend_params["num_threads"] = str(ucx_num_threads) + elif backend == "OBJ": + backend_params.setdefault("num_threads", "8") + elif backend == "GDS_MT": + backend_params.setdefault("thread_count", "8") + elif backend == "UCCL": + backend_params.setdefault("num_cpus", "8") + + return backend, backend_params + + +def _nixl_prefill_ucx_backend_env(): + backend = envs.SGLANG_DISAGGREGATION_NIXL_BACKEND.get() + if backend != "UCX": + return {} + + _, backend_params = _nixl_backend_config( + backend, + envs.SGLANG_DISAGGREGATION_NIXL_BACKEND_PARAMS.get(), + ucx_num_threads=NIXL_PREFILL_UCX_NUM_THREADS, + ) + return {"SGLANG_DISAGGREGATION_NIXL_BACKEND_PARAMS": json.dumps(backend_params)} + + +def _get_configured_nixl_backend_probe_error(): + backend = envs.SGLANG_DISAGGREGATION_NIXL_BACKEND.get() + backend_params_json = envs.SGLANG_DISAGGREGATION_NIXL_BACKEND_PARAMS.get() + + try: + from nixl._api import nixl_agent, nixl_agent_config, nixl_thread_sync_t + except ImportError as e: + return f"NIXL import failed: {e}" + + try: + backend, backend_params = _nixl_backend_config(backend, backend_params_json) + except (json.JSONDecodeError, ValueError) as e: + return str(e) + + try: + probe_num_threads = NIXL_PREFILL_UCX_NUM_THREADS if backend == "UCX" else 8 + agent_config = nixl_agent_config( + backends=[], + num_threads=probe_num_threads, + sync_mode=nixl_thread_sync_t.NIXL_THREAD_SYNC_STRICT, + ) + agent = nixl_agent(f"sglang_nixl_probe_{uuid.uuid4()}", agent_config) + available_plugins = agent.get_plugin_list() + if backend not in available_plugins: + return ( + f"NIXL backend {backend!r} not found. " + f"Available plugins: {available_plugins}." + ) + agent.create_backend(backend, backend_params) + except Exception as e: + return f"NIXL backend probe failed: {e}" + + return None + + +def _has_configured_nixl_backend(): + return _get_configured_nixl_backend_probe_error() is None + + +def _require_configured_nixl_backend(): + error = _get_configured_nixl_backend_probe_error() + if error is not None: + raise RuntimeError(error) + + +def _clear_disagg_failure_env(): + # Failure injection is scoped to the failure test; clear inherited env so + # Basic/Accuracy cannot run with injected transfer failures. + os.environ.pop("SGLANG_TEST_DISAGG_FAILURE_PROB", None) + os.environ.pop("DISAGGREGATION_TEST_FAILURE_PROB", None) + + +_HAS_CONFIGURED_NIXL_BACKEND = is_in_ci() or _has_configured_nixl_backend() + + +class NixlPDDisaggregationServerBase(PDDisaggregationServerBase): + prefill_tp_size = NIXL_PREFILL_TP_SIZE + decode_tp_size = NIXL_DECODE_TP_SIZE + decode_base_gpu_id = NIXL_DECODE_BASE_GPU_ID + + @classmethod + def start_prefill(cls): + prefill_env = { + **cls.extra_prefill_env, + **_nixl_prefill_ucx_backend_env(), + } + prefill_args = [ + "--trust-remote-code", + "--disaggregation-mode", + "prefill", + "--disaggregation-bootstrap-port", + cls.bootstrap_port, + "--tp", + str(cls.prefill_tp_size), + ] + list(cls.extra_prefill_args) + 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, + env=prefill_env, + return_stdout_stderr=( + (cls._prefill_stdout_buf, cls._prefill_stderr_buf) + if cls.capture_per_side_logs + else None + ), + ) + + @classmethod + def start_decode(cls): + decode_args = [ + "--trust-remote-code", + "--disaggregation-mode", + "decode", + "--disaggregation-bootstrap-port", + cls.bootstrap_port, + "--tp", + str(cls.decode_tp_size), + "--base-gpu-id", + str(cls.decode_base_gpu_id), + ] + list(cls.extra_decode_args) + 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, + env=cls.extra_decode_env, + return_stdout_stderr=( + (cls._decode_stdout_buf, cls._decode_stderr_buf) + if cls.capture_per_side_logs + else None + ), + ) + + +@unittest.skipUnless( + _HAS_CONFIGURED_NIXL_BACKEND, + "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. + """ + + @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_completion_returns_text_and_workers_stay_alive(self): + response = requests.post( + self.lb_url + "/generate", + json={ + "text": "The capital of France is", + "sampling_params": {"temperature": 0, "max_new_tokens": 16}, + }, + timeout=60, + ) + self.assertEqual(response.status_code, 200, response.text) + + data = response.json() + self.assertIn("text", data, f"Unexpected response shape: {data}") + self.assertGreater(len(data["text"]), 0, "Generated text should not be empty") + + assert_process_healthy(self, "load balancer", self.process_lb, self.lb_url) + assert_process_healthy(self, "prefill", self.process_prefill, self.prefill_url) + assert_process_healthy(self, "decode", self.process_decode, self.decode_url) + + def test_logprob(self): + response = requests.post( + self.lb_url + "/generate", + json={ + "text": "The capital of france is ", + "sampling_params": {"temperature": 0}, + "return_logprob": True, + "return_input_logprob": True, + "logprob_start_len": 0, + }, + timeout=60, + ) + self.assertEqual(response.status_code, 200, response.text) + + meta_info = response.json()["meta_info"] + completion_tokens = meta_info["completion_tokens"] + input_logprobs = meta_info["input_token_logprobs"] + output_logprobs = meta_info["output_token_logprobs"] + + 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}", + eval_name="gsm8k", + api="completion", + max_tokens=512, + num_examples=200, + num_shots=5, + num_threads=128, + temperature=0.0, + ) + + metrics = run_eval(args) + print(f"Evaluation metrics: {metrics}") + self.assertGreater( + metrics["score"], + NIXL_GSM8K_SCORE_THRESHOLD, + f"Expected NIXL PD transfer to preserve GSM8K accuracy, got {metrics}", + ) + + assert_process_healthy(self, "load balancer", self.process_lb, self.lb_url) + assert_process_healthy(self, "prefill", self.process_prefill, self.prefill_url) + assert_process_healthy(self, "decode", self.process_decode, self.decode_url) + + +@unittest.skipUnless( + _HAS_CONFIGURED_NIXL_BACKEND, + "NIXL with the configured backend is required for this test.", +) +class TestDisaggregationNixlFailure(NixlPDDisaggregationServerBase): + @classmethod + def setUpClass(cls): + _require_configured_nixl_backend() + super().setUpClass() + os.environ["SGLANG_TEST_DISAGG_FAILURE_PROB"] = "0.05" + cls.model = try_cached_model(DEFAULT_MODEL_NAME_FOR_TEST) + configure_nixl_pd_backend(cls) + cls.launch_all() + + @classmethod + def tearDownClass(cls): + os.environ.pop("SGLANG_TEST_DISAGG_FAILURE_PROB", None) + super().tearDownClass() + + def test_gsm8k(self): + args = SimpleNamespace( + base_url=f"http://{self.base_host}:{self.lb_port}", + eval_name="gsm8k", + api="completion", + max_tokens=512, + num_examples=200, + num_threads=128, + ) + + # Match TestDisaggregationMooncakeFailure: inject many transfer failures + # and tolerate eval/request errors as long as workers remain healthy. + try: + metrics = run_eval(args) + print(f"Evaluation metrics: {metrics}") + except Exception as e: + print(f"Test encountered expected errors: {e}") + + assert_process_healthy(self, "load balancer", self.process_lb, self.lb_url) + assert_process_healthy( + self, "prefill", self.process_prefill, self.prefill_url, "/health_generate" + ) + assert_process_healthy( + self, "decode", self.process_decode, self.decode_url, "/health_generate" + ) + + +if __name__ == "__main__": + unittest.main()