[Qwen3.5][MTP] Support FlashInfer CuTe DSL for online NVFP4 draft MoE (#31382)
This commit is contained in:
@@ -1436,22 +1436,33 @@ class ModelConfig:
|
|||||||
"quant_method", "" if not self.quantization else self.quantization
|
"quant_method", "" if not self.quantization else self.quantization
|
||||||
).lower()
|
).lower()
|
||||||
|
|
||||||
|
# ModelOpt FP4 checkpoints quantize only the target model; an
|
||||||
|
# embedded MTP draft may stay unquantized, so an explicit
|
||||||
|
# nvfp4_online opt-in for the draft wins over checkpoint detection.
|
||||||
|
# The online loader rejects already-packed weights at load time.
|
||||||
|
preserve_online_draft_quantization = (
|
||||||
|
self.is_draft_model
|
||||||
|
and self.quantization == "nvfp4_online"
|
||||||
|
and quant_method == "modelopt_fp4"
|
||||||
|
)
|
||||||
|
|
||||||
# Detect which checkpoint is it
|
# Detect which checkpoint is it
|
||||||
for _, method in QUANTIZATION_METHODS.items():
|
if not preserve_online_draft_quantization:
|
||||||
quantization_override = method.override_quantization_method(
|
for _, method in QUANTIZATION_METHODS.items():
|
||||||
quant_cfg, self.quantization
|
quantization_override = method.override_quantization_method(
|
||||||
)
|
quant_cfg, self.quantization
|
||||||
if quantization_override:
|
)
|
||||||
quant_method = quantization_override
|
if quantization_override:
|
||||||
self.quantization = quantization_override
|
quant_method = quantization_override
|
||||||
break
|
self.quantization = quantization_override
|
||||||
|
break
|
||||||
|
|
||||||
# Verify quantization configurations.
|
# Verify quantization configurations.
|
||||||
if self.quantization is None:
|
if self.quantization is None:
|
||||||
self.quantization = quant_method
|
self.quantization = quant_method
|
||||||
elif self.quantization != quant_method:
|
elif self.quantization != quant_method:
|
||||||
# Check if the CLI-specified quantization is compatible with HF config's quant_method
|
# Check if the CLI-specified quantization is compatible with HF config's quant_method
|
||||||
is_compatible = (
|
is_compatible = preserve_online_draft_quantization or (
|
||||||
self.quantization in compatible_quantization_methods
|
self.quantization in compatible_quantization_methods
|
||||||
and quant_method
|
and quant_method
|
||||||
in compatible_quantization_methods[self.quantization]
|
in compatible_quantization_methods[self.quantization]
|
||||||
|
|||||||
@@ -107,7 +107,8 @@ class DeepEPMoE(FusedMoE):
|
|||||||
elif (
|
elif (
|
||||||
get_moe_runner_backend().is_flashinfer_cutedsl()
|
get_moe_runner_backend().is_flashinfer_cutedsl()
|
||||||
and quant_config is not None
|
and quant_config is not None
|
||||||
and quant_config.get_name() in ("modelopt_fp4", "modelopt_mixed")
|
and quant_config.get_name()
|
||||||
|
in ("modelopt_fp4", "modelopt_mixed", "nvfp4_online")
|
||||||
):
|
):
|
||||||
self.deprecate_flag = True
|
self.deprecate_flag = True
|
||||||
elif (
|
elif (
|
||||||
|
|||||||
@@ -16,6 +16,8 @@ from sglang.srt.layers.moe import (
|
|||||||
MoeRunner,
|
MoeRunner,
|
||||||
MoeRunnerBackend,
|
MoeRunnerBackend,
|
||||||
MoeRunnerConfig,
|
MoeRunnerConfig,
|
||||||
|
get_deepep_mode,
|
||||||
|
get_moe_a2a_backend,
|
||||||
get_moe_runner_backend,
|
get_moe_runner_backend,
|
||||||
)
|
)
|
||||||
from sglang.srt.layers.moe.moe_runner.triton import TritonMoeQuantInfo
|
from sglang.srt.layers.moe.moe_runner.triton import TritonMoeQuantInfo
|
||||||
@@ -1988,6 +1990,17 @@ class ModelOptNvFp4FusedMoEMethod(FusedMoEMethodBase):
|
|||||||
|
|
||||||
return get_moe_runner_backend().is_flashinfer_cutedsl()
|
return get_moe_runner_backend().is_flashinfer_cutedsl()
|
||||||
|
|
||||||
|
@property
|
||||||
|
def supports_nvfp4_online_moe(self) -> bool:
|
||||||
|
a2a_backend = get_moe_a2a_backend()
|
||||||
|
return self.enable_flashinfer_trtllm_moe or (
|
||||||
|
self.enable_flashinfer_cutedsl_moe
|
||||||
|
and (
|
||||||
|
a2a_backend.is_flashinfer()
|
||||||
|
or (a2a_backend.is_deepep() and get_deepep_mode().is_low_latency())
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
# ----- CuteDSL v1 vs v2 path helpers -----
|
# ----- CuteDSL v1 vs v2 path helpers -----
|
||||||
#
|
#
|
||||||
# "v1": cutedsl + deepep low-latency.
|
# "v1": cutedsl + deepep low-latency.
|
||||||
@@ -2026,16 +2039,14 @@ class ModelOptNvFp4FusedMoEMethod(FusedMoEMethodBase):
|
|||||||
"NVFP4 quantization was selected, "
|
"NVFP4 quantization was selected, "
|
||||||
" dynamic quantization is not supported."
|
" dynamic quantization is not supported."
|
||||||
)
|
)
|
||||||
# `nvfp4_online` is not a serialized checkpoint format, but after the
|
# Online conversion changes only weight loading; downstream tensors use
|
||||||
# online loader converts each expert it uses the same packed NVFP4
|
# the same packed weight and scale layout as serialized ModelOpt NVFP4.
|
||||||
# weights, block scales, and per-tensor scales as serialized ModelOpt
|
|
||||||
# NVFP4 checkpoints. Reuse this layout and swap only the weight loader.
|
|
||||||
if is_nvfp4_online:
|
if is_nvfp4_online:
|
||||||
if not self.enable_flashinfer_trtllm_moe:
|
if not self.supports_nvfp4_online_moe:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
"--quantization nvfp4_online supports only "
|
"--quantization nvfp4_online supports flashinfer_trtllm, "
|
||||||
"--moe-runner-backend flashinfer_trtllm or "
|
"flashinfer_trtllm_routed, or flashinfer_cutedsl with "
|
||||||
"flashinfer_trtllm_routed."
|
"FlashInfer A2A or DeepEP low_latency."
|
||||||
)
|
)
|
||||||
|
|
||||||
# TODO(ch-wan): check if this is needed
|
# TODO(ch-wan): check if this is needed
|
||||||
@@ -2167,17 +2178,19 @@ class ModelOptNvFp4FusedMoEMethod(FusedMoEMethodBase):
|
|||||||
{"quant_method": FusedMoeWeightScaleSupported.TENSOR.value}
|
{"quant_method": FusedMoeWeightScaleSupported.TENSOR.value}
|
||||||
)
|
)
|
||||||
|
|
||||||
w13_input_scale_shape = (layer.num_experts, num_shards)
|
input_scale_fill = 1.0 if is_nvfp4_online else None
|
||||||
w13_input_scale = PerTensorScaleParameter(
|
w13_input_scale = _make_per_tensor_scale_parameter(
|
||||||
data=torch.empty(w13_input_scale_shape, dtype=torch.float32),
|
(layer.num_experts, num_shards),
|
||||||
weight_loader=weight_loader,
|
weight_loader=weight_loader,
|
||||||
|
fill_value=input_scale_fill,
|
||||||
)
|
)
|
||||||
w13_input_scale._sglang_require_global_experts = True
|
w13_input_scale._sglang_require_global_experts = True
|
||||||
layer.register_parameter("w13_input_scale", w13_input_scale)
|
layer.register_parameter("w13_input_scale", w13_input_scale)
|
||||||
|
|
||||||
w2_input_scale = PerTensorScaleParameter(
|
w2_input_scale = _make_per_tensor_scale_parameter(
|
||||||
data=torch.empty(layer.num_experts, dtype=torch.float32),
|
(layer.num_experts,),
|
||||||
weight_loader=weight_loader,
|
weight_loader=weight_loader,
|
||||||
|
fill_value=input_scale_fill,
|
||||||
)
|
)
|
||||||
w2_input_scale._sglang_require_global_experts = True
|
w2_input_scale._sglang_require_global_experts = True
|
||||||
layer.register_parameter("w2_input_scale", w2_input_scale)
|
layer.register_parameter("w2_input_scale", w2_input_scale)
|
||||||
|
|||||||
@@ -165,11 +165,11 @@ class ModelOptNvFp4OnlineFusedMoEMethod(ModelOptNvFp4FusedMoEMethod):
|
|||||||
if layer_match is not None
|
if layer_match is not None
|
||||||
else layer_prefix
|
else layer_prefix
|
||||||
)
|
)
|
||||||
if not self.enable_flashinfer_trtllm_moe:
|
if not self.supports_nvfp4_online_moe:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
"--quantization nvfp4_online supports only "
|
"--quantization nvfp4_online supports flashinfer_trtllm, "
|
||||||
"--moe-runner-backend flashinfer_trtllm or "
|
"flashinfer_trtllm_routed, or flashinfer_cutedsl with "
|
||||||
"flashinfer_trtllm_routed."
|
"FlashInfer A2A or DeepEP low_latency."
|
||||||
)
|
)
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
@@ -190,6 +190,12 @@ class ModelOptNvFp4OnlineFusedMoEMethod(ModelOptNvFp4FusedMoEMethod):
|
|||||||
"--quantization nvfp4_online expects 2D expert weights, "
|
"--quantization nvfp4_online expects 2D expert weights, "
|
||||||
f"got shape {tuple(weight.shape)}."
|
f"got shape {tuple(weight.shape)}."
|
||||||
)
|
)
|
||||||
|
if not weight.is_floating_point():
|
||||||
|
raise ValueError(
|
||||||
|
"--quantization nvfp4_online expects floating-point source "
|
||||||
|
f"expert weights, got dtype {weight.dtype}. Serialized packed "
|
||||||
|
"FP4 weights must use --quantization modelopt_fp4."
|
||||||
|
)
|
||||||
if weight.shape[-1] % 16 != 0:
|
if weight.shape[-1] % 16 != 0:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
"--quantization nvfp4_online requires expert weight K to be "
|
"--quantization nvfp4_online requires expert weight K to be "
|
||||||
|
|||||||
@@ -50,18 +50,23 @@ def should_run_flashinfer_autotune(
|
|||||||
if mr.server_args.disable_flashinfer_autotune:
|
if mr.server_args.disable_flashinfer_autotune:
|
||||||
return False
|
return False
|
||||||
|
|
||||||
# CuteDSL v1 (cutedsl runner + deepep a2a) bypasses MoeRunner and must not
|
server_args = mr.server_args
|
||||||
# be autotuned -- its _dummy_run would dispatch more tokens per rank than
|
if for_speculative_draft:
|
||||||
# SGLANG_DEEPEP_NUM_MAX_DISPATCH_TOKENS_PER_RANK, tripping a DeepEP assert.
|
backend_str = (
|
||||||
# Read server_args directly to avoid depending on initialize_moe_config()
|
server_args.speculative_moe_runner_backend or server_args.moe_runner_backend
|
||||||
# having already populated the MoE backend globals.
|
)
|
||||||
if (
|
a2a_backend_str = (
|
||||||
mr.server_args.moe_runner_backend == "flashinfer_cutedsl"
|
server_args.speculative_moe_a2a_backend or server_args.moe_a2a_backend
|
||||||
and mr.server_args.moe_a2a_backend == "deepep"
|
)
|
||||||
):
|
else:
|
||||||
return False
|
backend_str = server_args.moe_runner_backend
|
||||||
|
a2a_backend_str = server_args.moe_a2a_backend
|
||||||
|
|
||||||
backend_str = mr.server_args.moe_runner_backend
|
# Autotune can run before the MoE backend globals are initialized, so read
|
||||||
|
# the target or draft backend from server_args. CuteDSL v1 bypasses
|
||||||
|
# MoeRunner, and its dummy dispatch can exceed DeepEP low-latency's token limit.
|
||||||
|
if backend_str == "flashinfer_cutedsl" and a2a_backend_str == "deepep":
|
||||||
|
return False
|
||||||
|
|
||||||
# TODO smor- support other cases for flashinfer autotune, such as, mamba backend
|
# TODO smor- support other cases for flashinfer autotune, such as, mamba backend
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user