[FlashInfer v0.6.6][RL] Support fp8-last-n-bf16 RL for flashinfer_trtllm_routed moe backend (#20214)
This commit is contained in:
@@ -693,6 +693,16 @@ class FusedMoE(torch.nn.Module):
|
|||||||
if method.__class__.__name__ == "KTEPWrapperMethod":
|
if method.__class__.__name__ == "KTEPWrapperMethod":
|
||||||
method = method.gpu_method
|
method = method.gpu_method
|
||||||
|
|
||||||
|
# For flashinfer TRT-LLM BF16 path, process_weights_after_loading reshapes
|
||||||
|
# expert weights into block layout. During weight update, we must restore
|
||||||
|
# canonical load-time shapes before copying checkpoint tensors.
|
||||||
|
if isinstance(method, UnquantizedFusedMoEMethod):
|
||||||
|
method.maybe_restore_flashinfer_trtllm_bf16_weight_shape_for_load(
|
||||||
|
layer=self,
|
||||||
|
param=param,
|
||||||
|
weight_name=weight_name,
|
||||||
|
)
|
||||||
|
|
||||||
loaded_weight = (
|
loaded_weight = (
|
||||||
loaded_weight.t().contiguous()
|
loaded_weight.t().contiguous()
|
||||||
if (
|
if (
|
||||||
|
|||||||
@@ -675,10 +675,24 @@ def fused_experts_none_to_flashinfer_trtllm_bf16(
|
|||||||
dispatch_output: StandardDispatchOutput,
|
dispatch_output: StandardDispatchOutput,
|
||||||
quant_info: FlashInferTrtllmBf16MoeQuantInfo,
|
quant_info: FlashInferTrtllmBf16MoeQuantInfo,
|
||||||
runner_config: MoeRunnerConfig,
|
runner_config: MoeRunnerConfig,
|
||||||
|
use_routed_topk: bool = False,
|
||||||
) -> StandardCombineInput:
|
) -> StandardCombineInput:
|
||||||
# lazy import
|
# lazy import
|
||||||
from sglang.srt.layers.moe.token_dispatcher.standard import StandardCombineInput
|
from sglang.srt.layers.moe.token_dispatcher.standard import StandardCombineInput
|
||||||
|
from sglang.srt.layers.moe.topk import TopKOutputChecker
|
||||||
|
from sglang.srt.layers.moe.utils import RoutingMethodType
|
||||||
|
|
||||||
|
trtllm_bf16_routed_moe = None
|
||||||
|
trtllm_bf16_moe = None
|
||||||
|
if use_routed_topk:
|
||||||
|
try:
|
||||||
|
from flashinfer.fused_moe import trtllm_bf16_routed_moe
|
||||||
|
except ImportError as e:
|
||||||
|
raise ImportError(
|
||||||
|
"Can't import trtllm_bf16_routed_moe from flashinfer. "
|
||||||
|
"Please check flashinfer version to use bf16 with flashinfer_trtllm_routed backend."
|
||||||
|
) from e
|
||||||
|
else:
|
||||||
try:
|
try:
|
||||||
from flashinfer.fused_moe import trtllm_bf16_moe
|
from flashinfer.fused_moe import trtllm_bf16_moe
|
||||||
except ImportError as e:
|
except ImportError as e:
|
||||||
@@ -690,6 +704,7 @@ def fused_experts_none_to_flashinfer_trtllm_bf16(
|
|||||||
assert (
|
assert (
|
||||||
runner_config.activation == "silu"
|
runner_config.activation == "silu"
|
||||||
), "Only silu is supported for flashinfer trtllm moe"
|
), "Only silu is supported for flashinfer trtllm moe"
|
||||||
|
if not use_routed_topk:
|
||||||
assert (
|
assert (
|
||||||
dispatch_output.topk_output.topk_config.renormalize
|
dispatch_output.topk_output.topk_config.renormalize
|
||||||
), "Renormalize is required for flashinfer trtllm moe"
|
), "Renormalize is required for flashinfer trtllm moe"
|
||||||
@@ -699,15 +714,49 @@ def fused_experts_none_to_flashinfer_trtllm_bf16(
|
|||||||
assert (
|
assert (
|
||||||
runner_config.is_gated
|
runner_config.is_gated
|
||||||
), "Only gated MoEs are supported for flashinfer trtllm moe"
|
), "Only gated MoEs are supported for flashinfer trtllm moe"
|
||||||
from sglang.srt.layers.moe.topk import TopKOutputChecker
|
|
||||||
|
|
||||||
assert TopKOutputChecker.format_is_bypassed(dispatch_output.topk_output)
|
|
||||||
|
|
||||||
hidden_states = dispatch_output.hidden_states
|
hidden_states = dispatch_output.hidden_states
|
||||||
topk_output = dispatch_output.topk_output
|
topk_output = dispatch_output.topk_output
|
||||||
topk_config = topk_output.topk_config
|
|
||||||
|
|
||||||
with use_symmetric_memory(get_tp_group(), disabled=not is_allocation_symmetric()):
|
with use_symmetric_memory(get_tp_group(), disabled=not is_allocation_symmetric()):
|
||||||
|
if use_routed_topk:
|
||||||
|
assert (
|
||||||
|
runner_config.top_k is not None
|
||||||
|
), "runner_config.top_k is required for flashinfer_trtllm_routed."
|
||||||
|
assert TopKOutputChecker.format_is_standard(topk_output)
|
||||||
|
routing_method_type = runner_config.routing_method_type
|
||||||
|
if routing_method_type is None:
|
||||||
|
routing_method_type = RoutingMethodType.Default
|
||||||
|
elif routing_method_type == RoutingMethodType.DeepSeekV3:
|
||||||
|
routing_method_type = RoutingMethodType.TopK
|
||||||
|
|
||||||
|
packed_topk_ids = _pack_topk_for_flashinfer_routed(
|
||||||
|
topk_ids=topk_output.topk_ids,
|
||||||
|
topk_weights=topk_output.topk_weights,
|
||||||
|
)
|
||||||
|
final_hidden_states = trtllm_bf16_routed_moe(
|
||||||
|
topk_ids=packed_topk_ids,
|
||||||
|
hidden_states=hidden_states,
|
||||||
|
gemm1_weights=quant_info.gemm1_weights,
|
||||||
|
gemm2_weights=quant_info.gemm2_weights,
|
||||||
|
num_experts=quant_info.global_num_experts,
|
||||||
|
top_k=runner_config.top_k,
|
||||||
|
n_group=None,
|
||||||
|
topk_group=None,
|
||||||
|
intermediate_size=runner_config.intermediate_size_per_partition,
|
||||||
|
local_expert_offset=quant_info.local_expert_offset,
|
||||||
|
local_num_experts=runner_config.num_local_experts,
|
||||||
|
routing_method_type=routing_method_type,
|
||||||
|
routed_scaling_factor=(
|
||||||
|
runner_config.routed_scaling_factor
|
||||||
|
if runner_config.routed_scaling_factor is not None
|
||||||
|
else 1.0
|
||||||
|
),
|
||||||
|
tune_max_num_tokens=next_power_of_2(hidden_states.shape[0]),
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
assert TopKOutputChecker.format_is_bypassed(topk_output)
|
||||||
|
topk_config = topk_output.topk_config
|
||||||
|
|
||||||
# Call the fused kernel
|
# Call the fused kernel
|
||||||
final_hidden_states = trtllm_bf16_moe(
|
final_hidden_states = trtllm_bf16_moe(
|
||||||
@@ -768,6 +817,13 @@ def fused_experts_none_to_flashinfer_trtllm_routed(
|
|||||||
runner_config,
|
runner_config,
|
||||||
use_routed_topk=True,
|
use_routed_topk=True,
|
||||||
)
|
)
|
||||||
|
if isinstance(quant_info, FlashInferTrtllmBf16MoeQuantInfo):
|
||||||
|
return fused_experts_none_to_flashinfer_trtllm_bf16(
|
||||||
|
dispatch_output,
|
||||||
|
quant_info,
|
||||||
|
runner_config,
|
||||||
|
use_routed_topk=True,
|
||||||
|
)
|
||||||
raise TypeError(
|
raise TypeError(
|
||||||
f"Unexpected quant_info type for flashinfer_trtllm_routed: {type(quant_info)}"
|
f"Unexpected quant_info type for flashinfer_trtllm_routed: {type(quant_info)}"
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -323,12 +323,58 @@ class UnquantizedFusedMoEMethod(FusedMoEMethodBase, MultiPlatformOp):
|
|||||||
|
|
||||||
return
|
return
|
||||||
|
|
||||||
|
def maybe_restore_flashinfer_trtllm_bf16_weight_shape_for_load(
|
||||||
|
self,
|
||||||
|
layer: torch.nn.Module,
|
||||||
|
param: torch.nn.Parameter,
|
||||||
|
weight_name: str,
|
||||||
|
) -> None:
|
||||||
|
"""Restore canonical BF16 MoE load shapes before hot weight copy.
|
||||||
|
|
||||||
|
The flashinfer TRT-LLM BF16 postprocess reshapes expert weights into
|
||||||
|
block layout. During weight update, checkpoint tensors are in
|
||||||
|
canonical layout and need a temporary shape restore for copy.
|
||||||
|
"""
|
||||||
|
if not get_moe_runner_backend().is_flashinfer_trtllm_routed():
|
||||||
|
return
|
||||||
|
|
||||||
|
expected_shape = None
|
||||||
|
if weight_name.endswith(".experts.w13_weight"):
|
||||||
|
w13_rows = (
|
||||||
|
2 * layer.intermediate_size_per_partition
|
||||||
|
if layer.moe_runner_config.is_gated
|
||||||
|
else layer.intermediate_size_per_partition
|
||||||
|
)
|
||||||
|
expected_shape = (layer.num_local_experts, w13_rows, layer.hidden_size)
|
||||||
|
elif weight_name.endswith(".experts.w2_weight"):
|
||||||
|
expected_shape = (
|
||||||
|
layer.num_local_experts,
|
||||||
|
layer.hidden_size,
|
||||||
|
layer.intermediate_size_per_partition,
|
||||||
|
)
|
||||||
|
|
||||||
|
if expected_shape is None or tuple(param.data.shape) == expected_shape:
|
||||||
|
return
|
||||||
|
|
||||||
|
expected_numel = expected_shape[0] * expected_shape[1] * expected_shape[2]
|
||||||
|
if param.data.numel() != expected_numel:
|
||||||
|
raise RuntimeError(
|
||||||
|
f"Cannot restore flashinfer TRT-LLM BF16 MoE weight shape for {weight_name}: "
|
||||||
|
f"current shape={tuple(param.data.shape)}, expected shape={expected_shape}."
|
||||||
|
)
|
||||||
|
|
||||||
|
param.data = param.data.reshape(expected_shape)
|
||||||
|
|
||||||
def create_moe_runner(
|
def create_moe_runner(
|
||||||
self, layer: torch.nn.Module, moe_runner_config: MoeRunnerConfig
|
self, layer: torch.nn.Module, moe_runner_config: MoeRunnerConfig
|
||||||
):
|
):
|
||||||
self.moe_runner_config = moe_runner_config
|
self.moe_runner_config = moe_runner_config
|
||||||
if self.use_flashinfer_trtllm_moe:
|
if self.use_flashinfer_trtllm_moe:
|
||||||
backend = MoeRunnerBackend.FLASHINFER_TRTLLM
|
backend = (
|
||||||
|
MoeRunnerBackend.FLASHINFER_TRTLLM_ROUTED
|
||||||
|
if get_moe_runner_backend().is_flashinfer_trtllm_routed()
|
||||||
|
else MoeRunnerBackend.FLASHINFER_TRTLLM
|
||||||
|
)
|
||||||
elif self.use_triton_kernels:
|
elif self.use_triton_kernels:
|
||||||
backend = MoeRunnerBackend.TRITON_KERNELS
|
backend = MoeRunnerBackend.TRITON_KERNELS
|
||||||
else:
|
else:
|
||||||
|
|||||||
@@ -2638,7 +2638,8 @@ class ServerArgs:
|
|||||||
assert self.quantization in [
|
assert self.quantization in [
|
||||||
"fp8",
|
"fp8",
|
||||||
"mxfp8",
|
"mxfp8",
|
||||||
], f"Invalid quantization '{self.quantization}'. \nFlashInfer TRTLLM routed MOE supports only: 'fp8' or 'mxfp8'."
|
None,
|
||||||
|
], f"Invalid quantization '{self.quantization}'. \nFlashInfer TRTLLM routed MOE supports only: 'fp8', 'mxfp8', or bfloat16 (None)."
|
||||||
self.disable_shared_experts_fusion = True
|
self.disable_shared_experts_fusion = True
|
||||||
logger.warning(
|
logger.warning(
|
||||||
"FlashInfer TRTLLM routed MoE is enabled. --disable-shared-experts-fusion is automatically set."
|
"FlashInfer TRTLLM routed MoE is enabled. --disable-shared-experts-fusion is automatically set."
|
||||||
|
|||||||
@@ -12,7 +12,7 @@ from sglang.test.test_utils import (
|
|||||||
popen_launch_server,
|
popen_launch_server,
|
||||||
)
|
)
|
||||||
|
|
||||||
register_cuda_ci(est_time=500, suite="nightly-4-gpu-b200", nightly=True)
|
register_cuda_ci(est_time=600, suite="nightly-4-gpu-b200", nightly=True)
|
||||||
|
|
||||||
|
|
||||||
class FlashinferTrtllmGenMoeBackendFP8Base:
|
class FlashinferTrtllmGenMoeBackendFP8Base:
|
||||||
@@ -187,5 +187,11 @@ class TestFlashinferTrtllmGenMoeBackendMXFP8Routed(
|
|||||||
backend = "flashinfer_trtllm_routed"
|
backend = "flashinfer_trtllm_routed"
|
||||||
|
|
||||||
|
|
||||||
|
class TestFlashinferTrtllmGenMoeBackendBF16Routed(
|
||||||
|
FlashinferTrtllmGenMoeBackendBF16Base, CustomTestCase
|
||||||
|
):
|
||||||
|
backend = "flashinfer_trtllm_routed"
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
unittest.main()
|
unittest.main()
|
||||||
|
|||||||
@@ -0,0 +1,165 @@
|
|||||||
|
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()
|
||||||
Reference in New Issue
Block a user