from sglang.test.ci.ci_register import register_cuda_ci register_cuda_ci(est_time=200, suite="stage-c-test-4-gpu-b200") import unittest import requests from sglang.srt.utils import kill_process_tree from sglang.test.test_utils import ( DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, DEFAULT_URL_FOR_TEST, CustomTestCase, popen_launch_server, ) class TestServerUpdateWeightsFromDiskMXFP8(CustomTestCase): model = "zianglih/Qwen3-30B-A3B-Instruct-2507-MXFP8-last-8-BF16" base_url = DEFAULT_URL_FOR_TEST request_timeout = 120 update_timeout = 240 decode_payload = { "text": "The capital of France is", "sampling_params": {"temperature": 0, "max_new_tokens": 16}, } backend_test_suites = ( { "fp8_gemm_backend": "flashinfer_trtllm", "moe_runner_backend": "flashinfer_trtllm_routed", }, ) def _launch_server(self, fp8_gemm_backend, moe_runner_backend): return popen_launch_server( self.model, self.base_url, timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, other_args=[ "--base-gpu-id", "0", "--tp-size", "4", "--fp8-gemm-backend", fp8_gemm_backend, "--moe-runner-backend", moe_runner_backend, ], ) def _get_json(self, endpoint, timeout=None): response = requests.get( f"{self.base_url}{endpoint}", timeout=timeout or self.request_timeout, ) response.raise_for_status() return response.json() def _post_json(self, endpoint, payload, timeout=None): response = requests.post( f"{self.base_url}{endpoint}", json=payload, timeout=timeout or self.request_timeout, ) response.raise_for_status() return response.json() def _run_decode(self): return self._post_json("/generate", self.decode_payload)["text"] def _assert_non_empty_decode(self): self.assertTrue(len(self._run_decode()) > 0) def _get_decode_logprob_signature(self): ret = self._post_json( "/generate", {**self.decode_payload, "return_logprob": True}, ) output_token_logprobs = ret["meta_info"].get("output_token_logprobs") self.assertIsNotNone(output_token_logprobs) self.assertGreater( len(output_token_logprobs), 0, "Expected non-empty output_token_logprobs.", ) return { "text": ret["text"], "token_ids": [int(x[1]) for x in output_token_logprobs], "logprobs": [float(x[0]) for x in output_token_logprobs], } def _assert_decode_logprob_unchanged(self, before, after, atol=1e-4): self.assertEqual(after["text"], before["text"]) self.assertEqual(after["token_ids"], before["token_ids"]) self.assertEqual(len(after["logprobs"]), len(before["logprobs"])) for idx, (a, b) in enumerate(zip(after["logprobs"], before["logprobs"])): self.assertLessEqual( abs(a - b), atol, f"Output token logprob changed at idx={idx}: before={b}, after={a}", ) def _get_model_info(self): return self._get_json("/get_model_info")["model_path"] def _run_update_weights( self, model_path, flush_cache=True, abort_all_requests=False, ): return self._post_json( "/update_weights_from_disk", { "model_path": model_path, "flush_cache": flush_cache, "abort_all_requests": abort_all_requests, }, timeout=self.update_timeout, ) def test_parameterized_update_weights_mxfp8(self): update_test_suites = ( {"flush_cache": True, "abort_all_requests": False}, {"flush_cache": False, "abort_all_requests": False}, ) for backend_test_suite in self.backend_test_suites: with self.subTest(**backend_test_suite): process = self._launch_server( backend_test_suite["fp8_gemm_backend"], backend_test_suite["moe_runner_backend"], ) try: origin_model_path = self._get_model_info() self.assertEqual(origin_model_path, self.model) self._assert_non_empty_decode() baseline_sig = self._get_decode_logprob_signature() for update_test_suite in update_test_suites: with self.subTest( fp8_gemm_backend=backend_test_suite["fp8_gemm_backend"], moe_runner_backend=backend_test_suite["moe_runner_backend"], flush_cache=update_test_suite["flush_cache"], abort_all_requests=update_test_suite["abort_all_requests"], ): ret = self._run_update_weights( self.model, flush_cache=update_test_suite["flush_cache"], abort_all_requests=update_test_suite[ "abort_all_requests" ], ) self.assertTrue(ret.get("success"), f"{ret=}") self.assertEqual(self._get_model_info(), self.model) self._assert_non_empty_decode() updated_sig = self._get_decode_logprob_signature() self._assert_decode_logprob_unchanged( baseline_sig, updated_sig ) finally: kill_process_tree(process.pid) if __name__ == "__main__": unittest.main()