[RL] Refactor NVFP4 shuffling/swizzling to in-place replacement (#22204)

This commit is contained in:
Ziang Li
2026-04-12 19:08:45 -07:00
committed by GitHub
parent 934e19a610
commit 5593539942
6 changed files with 157 additions and 103 deletions
@@ -157,6 +157,49 @@ class FlashinferTrtllmGenMoeBackendMXFP8Base:
self.assertGreater(metrics["score"], 0.93)
class FlashinferTrtllmGenMoeBackendNVFP4Base:
backend = None
@classmethod
def setUpClass(cls):
cls.model = "nvidia/Qwen3-30B-A3B-NVFP4"
cls.base_url = DEFAULT_URL_FOR_TEST
cls.process = popen_launch_server(
cls.model,
cls.base_url,
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
env={**os.environ, "SGLANG_ENABLE_JIT_DEEPGEMM": "False"},
other_args=[
"--moe-runner-backend",
cls.backend,
"--tp-size",
"4",
"--ep-size",
"4",
"--mem-fraction-static",
"0.7",
],
)
@classmethod
def tearDownClass(cls):
kill_process_tree(cls.process.pid)
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.89)
class TestFlashinferTrtllmGenMoeBackendFP8(
FlashinferTrtllmGenMoeBackendFP8Base, CustomTestCase
):
@@ -175,6 +218,12 @@ class TestFlashinferTrtllmGenMoeBackendBF16(
backend = "flashinfer_trtllm"
class TestFlashinferTrtllmGenMoeBackendNVFP4(
FlashinferTrtllmGenMoeBackendNVFP4Base, CustomTestCase
):
backend = "flashinfer_trtllm"
class TestFlashinferTrtllmGenMoeBackendFP8Routed(
FlashinferTrtllmGenMoeBackendFP8Base, CustomTestCase
):
@@ -193,5 +242,11 @@ class TestFlashinferTrtllmGenMoeBackendBF16Routed(
backend = "flashinfer_trtllm_routed"
class TestFlashinferTrtllmGenMoeBackendNVFP4Routed(
FlashinferTrtllmGenMoeBackendNVFP4Base, CustomTestCase
):
backend = "flashinfer_trtllm_routed"
if __name__ == "__main__":
unittest.main()
@@ -15,37 +15,43 @@ from sglang.test.test_utils import (
)
class TestServerUpdateWeightsFromDiskMXFP8(CustomTestCase):
model = "zianglih/Qwen3-30B-A3B-Instruct-2507-MXFP8-last-8-BF16"
class UpdateWeightsFromDiskBase:
model = None
base_url = DEFAULT_URL_FOR_TEST
request_timeout = 120
update_timeout = 240
launch_env = None
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",
},
backend_test_suites = ()
update_test_suites = (
{"flush_cache": True, "abort_all_requests": False},
{"flush_cache": False, "abort_all_requests": False},
)
def _launch_server(self, fp8_gemm_backend, moe_runner_backend):
@classmethod
def setUpClass(cls):
super().setUpClass()
if cls.model is None:
raise NotImplementedError("Subclass must set 'model' attribute")
if not cls.backend_test_suites:
raise NotImplementedError(
"Subclass must set non-empty 'backend_test_suites'"
)
def _launch_server(self, backend_test_suite):
launch_kwargs = {}
if self.launch_env is not None:
launch_kwargs["env"] = self.launch_env
other_args = backend_test_suite.get("other_args")
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,
],
other_args=other_args,
**launch_kwargs,
)
def _get_json(self, endpoint, timeout=None):
@@ -119,30 +125,19 @@ class TestServerUpdateWeightsFromDiskMXFP8(CustomTestCase):
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},
)
def test_parameterized_update_weights_from_disk(self):
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"],
)
case_name = backend_test_suite.get("name", "default")
with self.subTest(model=self.model, case_name=case_name):
process = self._launch_server(backend_test_suite)
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"],
):
for update_test_suite in self.update_test_suites:
with self.subTest(case_name=case_name, **update_test_suite):
ret = self._run_update_weights(
self.model,
flush_cache=update_test_suite["flush_cache"],
@@ -161,5 +156,43 @@ class TestServerUpdateWeightsFromDiskMXFP8(CustomTestCase):
kill_process_tree(process.pid)
class TestServerUpdateWeightsFromDiskMXFP8(UpdateWeightsFromDiskBase, CustomTestCase):
model = "zianglih/Qwen3-30B-A3B-Instruct-2507-MXFP8-last-8-BF16"
backend_test_suites = (
{
"name": "flashinfer_trtllm_routed_mxfp8",
"other_args": (
"--base-gpu-id",
"0",
"--tp-size",
"4",
"--fp8-gemm-backend",
"flashinfer_trtllm",
"--moe-runner-backend",
"flashinfer_trtllm_routed",
),
},
)
class TestServerUpdateWeightsFromDiskNVFP4(UpdateWeightsFromDiskBase, CustomTestCase):
model = "nvidia/Qwen3-30B-A3B-NVFP4"
backend_test_suites = (
{
"name": "flashinfer_trtllm_nvfp4",
"other_args": (
"--base-gpu-id",
"0",
"--tp-size",
"4",
"--fp4-gemm-backend",
"flashinfer_trtllm",
"--moe-runner-backend",
"flashinfer_trtllm_routed",
),
},
)
if __name__ == "__main__":
unittest.main()