[RL] Fix FlashInfer TRTLLM MXFP8 dense weight layout (#28459)

This commit is contained in:
Ziang Li
2026-06-17 10:33:35 +00:00
committed by GitHub
parent 2f1390fcb1
commit 3fb65ebabd
3 changed files with 45 additions and 6 deletions
+22 -4
View File
@@ -593,6 +593,21 @@ class Fp8LinearMethod(LinearMethodBase):
scale_u8 = layer.weight_scale_inv.data
n, k = weight.shape
epilogue_tile_m = 128
sf_cols = k // 32
scale_u8 = scale_u8.contiguous().view(torch.uint8).reshape(n, sf_cols)
padded_n = ((n + epilogue_tile_m - 1) // epilogue_tile_m) * (
epilogue_tile_m
)
pad_rows = padded_n - n
if pad_rows:
scale_u8 = F.pad(
scale_u8,
(0, 0, 0, pad_rows),
mode="constant",
value=0,
)
copy_or_rebind_param(
layer,
@@ -603,9 +618,9 @@ class Fp8LinearMethod(LinearMethodBase):
)
copy_or_rebind_param(
layer,
"weight_scale_inv",
"weight_scale_inv_shuffled",
shuffle_matrix_sf_a(
scale_u8.contiguous().view(torch.uint8).reshape(n, k // 32),
scale_u8,
epilogue_tile_m,
num_elts_per_sf=32,
)
@@ -775,8 +790,11 @@ class Fp8LinearMethod(LinearMethodBase):
)
if self.use_mxfp8:
if get_fp8_gemm_runner_backend().is_flashinfer_cutlass():
weight_scale = layer.weight_scale_inv_swizzled
if (
get_fp8_gemm_runner_backend().is_flashinfer_cutlass()
or get_fp8_gemm_runner_backend().is_flashinfer_trtllm()
):
weight_scale = layer.weight_scale_inv_shuffled
else:
weight_scale = layer.weight_scale_inv
if isinstance(x, tuple):
@@ -171,7 +171,7 @@ class FlashinferTrtllmGenMoeBackendMXFP8MixedBF16Base:
"--kv-cache-dtype",
"bf16",
"--fp8-gemm-backend",
"flashinfer_cutlass",
"flashinfer_trtllm",
"--moe-runner-backend",
cls.backend,
"--tp-size",
@@ -2,6 +2,7 @@ from sglang.test.ci.ci_register import register_cuda_ci
register_cuda_ci(est_time=320, stage="extra-b", runner_config="4-gpu-b200")
import time
import unittest
import requests
@@ -20,6 +21,7 @@ class UpdateWeightsFromDiskBase:
base_url = DEFAULT_URL_FOR_TEST
request_timeout = 120
update_timeout = 240
idle_timeout = 30
launch_env = None
decode_payload = {
"text": "The capital of France is",
@@ -74,6 +76,19 @@ class UpdateWeightsFromDiskBase:
def _run_decode(self):
return self._post_json("/generate", self.decode_payload)["text"]
def _wait_until_idle(self):
deadline = time.monotonic() + self.idle_timeout
last_loads = None
while time.monotonic() < deadline:
last_loads = self._get_json("/v1/loads?include=core")["loads"]
if last_loads and all(
load["num_running_reqs"] == 0 and load["num_waiting_reqs"] == 0
for load in last_loads
):
return
time.sleep(0.1)
self.fail(f"Server did not become idle before weight update: {last_loads=}")
def _assert_non_empty_decode(self):
self.assertTrue(len(self._run_decode()) > 0)
@@ -138,6 +153,7 @@ class UpdateWeightsFromDiskBase:
for update_test_suite in self.update_test_suites:
with self.subTest(case_name=case_name, **update_test_suite):
self._wait_until_idle()
ret = self._run_update_weights(
self.model,
flush_cache=update_test_suite["flush_cache"],
@@ -157,7 +173,8 @@ class UpdateWeightsFromDiskBase:
class TestServerUpdateWeightsFromDiskMXFP8(UpdateWeightsFromDiskBase, CustomTestCase):
model = "zianglih/Qwen3-30B-A3B-Instruct-2507-MXFP8-last-8-BF16"
model = "zianglih/JoyAI-LLM-Flash-MXFP8-last-6-BF16"
decode_payload = {**UpdateWeightsFromDiskBase.decode_payload, "routed_dp_rank": 0}
backend_test_suites = (
{
"name": "flashinfer_trtllm_routed_mxfp8",
@@ -166,10 +183,14 @@ class TestServerUpdateWeightsFromDiskMXFP8(UpdateWeightsFromDiskBase, CustomTest
"0",
"--tp-size",
"4",
"--dp-size",
"4",
"--enable-dp-attention",
"--fp8-gemm-backend",
"flashinfer_trtllm",
"--moe-runner-backend",
"flashinfer_trtllm_routed",
"--trust-remote-code",
),
},
)