[RL] Fix FlashInfer TRTLLM MXFP8 dense weight layout (#28459)
This commit is contained in:
@@ -593,6 +593,21 @@ class Fp8LinearMethod(LinearMethodBase):
|
|||||||
scale_u8 = layer.weight_scale_inv.data
|
scale_u8 = layer.weight_scale_inv.data
|
||||||
n, k = weight.shape
|
n, k = weight.shape
|
||||||
epilogue_tile_m = 128
|
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(
|
copy_or_rebind_param(
|
||||||
layer,
|
layer,
|
||||||
@@ -603,9 +618,9 @@ class Fp8LinearMethod(LinearMethodBase):
|
|||||||
)
|
)
|
||||||
copy_or_rebind_param(
|
copy_or_rebind_param(
|
||||||
layer,
|
layer,
|
||||||
"weight_scale_inv",
|
"weight_scale_inv_shuffled",
|
||||||
shuffle_matrix_sf_a(
|
shuffle_matrix_sf_a(
|
||||||
scale_u8.contiguous().view(torch.uint8).reshape(n, k // 32),
|
scale_u8,
|
||||||
epilogue_tile_m,
|
epilogue_tile_m,
|
||||||
num_elts_per_sf=32,
|
num_elts_per_sf=32,
|
||||||
)
|
)
|
||||||
@@ -775,8 +790,11 @@ class Fp8LinearMethod(LinearMethodBase):
|
|||||||
)
|
)
|
||||||
|
|
||||||
if self.use_mxfp8:
|
if self.use_mxfp8:
|
||||||
if get_fp8_gemm_runner_backend().is_flashinfer_cutlass():
|
if (
|
||||||
weight_scale = layer.weight_scale_inv_swizzled
|
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:
|
else:
|
||||||
weight_scale = layer.weight_scale_inv
|
weight_scale = layer.weight_scale_inv
|
||||||
if isinstance(x, tuple):
|
if isinstance(x, tuple):
|
||||||
|
|||||||
@@ -171,7 +171,7 @@ class FlashinferTrtllmGenMoeBackendMXFP8MixedBF16Base:
|
|||||||
"--kv-cache-dtype",
|
"--kv-cache-dtype",
|
||||||
"bf16",
|
"bf16",
|
||||||
"--fp8-gemm-backend",
|
"--fp8-gemm-backend",
|
||||||
"flashinfer_cutlass",
|
"flashinfer_trtllm",
|
||||||
"--moe-runner-backend",
|
"--moe-runner-backend",
|
||||||
cls.backend,
|
cls.backend,
|
||||||
"--tp-size",
|
"--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")
|
register_cuda_ci(est_time=320, stage="extra-b", runner_config="4-gpu-b200")
|
||||||
|
|
||||||
|
import time
|
||||||
import unittest
|
import unittest
|
||||||
|
|
||||||
import requests
|
import requests
|
||||||
@@ -20,6 +21,7 @@ class UpdateWeightsFromDiskBase:
|
|||||||
base_url = DEFAULT_URL_FOR_TEST
|
base_url = DEFAULT_URL_FOR_TEST
|
||||||
request_timeout = 120
|
request_timeout = 120
|
||||||
update_timeout = 240
|
update_timeout = 240
|
||||||
|
idle_timeout = 30
|
||||||
launch_env = None
|
launch_env = None
|
||||||
decode_payload = {
|
decode_payload = {
|
||||||
"text": "The capital of France is",
|
"text": "The capital of France is",
|
||||||
@@ -74,6 +76,19 @@ class UpdateWeightsFromDiskBase:
|
|||||||
def _run_decode(self):
|
def _run_decode(self):
|
||||||
return self._post_json("/generate", self.decode_payload)["text"]
|
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):
|
def _assert_non_empty_decode(self):
|
||||||
self.assertTrue(len(self._run_decode()) > 0)
|
self.assertTrue(len(self._run_decode()) > 0)
|
||||||
|
|
||||||
@@ -138,6 +153,7 @@ class UpdateWeightsFromDiskBase:
|
|||||||
|
|
||||||
for update_test_suite in self.update_test_suites:
|
for update_test_suite in self.update_test_suites:
|
||||||
with self.subTest(case_name=case_name, **update_test_suite):
|
with self.subTest(case_name=case_name, **update_test_suite):
|
||||||
|
self._wait_until_idle()
|
||||||
ret = self._run_update_weights(
|
ret = self._run_update_weights(
|
||||||
self.model,
|
self.model,
|
||||||
flush_cache=update_test_suite["flush_cache"],
|
flush_cache=update_test_suite["flush_cache"],
|
||||||
@@ -157,7 +173,8 @@ class UpdateWeightsFromDiskBase:
|
|||||||
|
|
||||||
|
|
||||||
class TestServerUpdateWeightsFromDiskMXFP8(UpdateWeightsFromDiskBase, CustomTestCase):
|
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 = (
|
backend_test_suites = (
|
||||||
{
|
{
|
||||||
"name": "flashinfer_trtllm_routed_mxfp8",
|
"name": "flashinfer_trtllm_routed_mxfp8",
|
||||||
@@ -166,10 +183,14 @@ class TestServerUpdateWeightsFromDiskMXFP8(UpdateWeightsFromDiskBase, CustomTest
|
|||||||
"0",
|
"0",
|
||||||
"--tp-size",
|
"--tp-size",
|
||||||
"4",
|
"4",
|
||||||
|
"--dp-size",
|
||||||
|
"4",
|
||||||
|
"--enable-dp-attention",
|
||||||
"--fp8-gemm-backend",
|
"--fp8-gemm-backend",
|
||||||
"flashinfer_trtllm",
|
"flashinfer_trtllm",
|
||||||
"--moe-runner-backend",
|
"--moe-runner-backend",
|
||||||
"flashinfer_trtllm_routed",
|
"flashinfer_trtllm_routed",
|
||||||
|
"--trust-remote-code",
|
||||||
),
|
),
|
||||||
},
|
},
|
||||||
)
|
)
|
||||||
|
|||||||
Reference in New Issue
Block a user