From 48dbc24cbff1caaed646d577f2c7e6fd394b9116 Mon Sep 17 00:00:00 2001 From: YAMY <74099316+YAMY1234@users.noreply.github.com> Date: Thu, 30 Jul 2026 17:19:48 -0700 Subject: [PATCH] [Qwen3.5][MTP] Support FlashInfer CuTe DSL for online NVFP4 draft MoE (#31382) --- python/sglang/srt/configs/model_config.py | 29 +++++++++----- python/sglang/srt/layers/moe/ep_moe/layer.py | 3 +- .../srt/layers/quantization/modelopt_quant.py | 39 ++++++++++++------- .../srt/layers/quantization/nvfp4_online.py | 14 +++++-- .../runner/flashinfer_autotune.py | 27 +++++++------ 5 files changed, 74 insertions(+), 38 deletions(-) diff --git a/python/sglang/srt/configs/model_config.py b/python/sglang/srt/configs/model_config.py index 37566da40..c0fd8f6c5 100644 --- a/python/sglang/srt/configs/model_config.py +++ b/python/sglang/srt/configs/model_config.py @@ -1436,22 +1436,33 @@ class ModelConfig: "quant_method", "" if not self.quantization else self.quantization ).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 - for _, method in QUANTIZATION_METHODS.items(): - quantization_override = method.override_quantization_method( - quant_cfg, self.quantization - ) - if quantization_override: - quant_method = quantization_override - self.quantization = quantization_override - break + if not preserve_online_draft_quantization: + for _, method in QUANTIZATION_METHODS.items(): + quantization_override = method.override_quantization_method( + quant_cfg, self.quantization + ) + if quantization_override: + quant_method = quantization_override + self.quantization = quantization_override + break # Verify quantization configurations. if self.quantization is None: self.quantization = quant_method elif self.quantization != 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 and quant_method in compatible_quantization_methods[self.quantization] diff --git a/python/sglang/srt/layers/moe/ep_moe/layer.py b/python/sglang/srt/layers/moe/ep_moe/layer.py index 13dba0388..f656fe9ec 100644 --- a/python/sglang/srt/layers/moe/ep_moe/layer.py +++ b/python/sglang/srt/layers/moe/ep_moe/layer.py @@ -107,7 +107,8 @@ class DeepEPMoE(FusedMoE): elif ( get_moe_runner_backend().is_flashinfer_cutedsl() 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 elif ( diff --git a/python/sglang/srt/layers/quantization/modelopt_quant.py b/python/sglang/srt/layers/quantization/modelopt_quant.py index 6cefdc8ee..45c0eef71 100755 --- a/python/sglang/srt/layers/quantization/modelopt_quant.py +++ b/python/sglang/srt/layers/quantization/modelopt_quant.py @@ -16,6 +16,8 @@ from sglang.srt.layers.moe import ( MoeRunner, MoeRunnerBackend, MoeRunnerConfig, + get_deepep_mode, + get_moe_a2a_backend, get_moe_runner_backend, ) 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() + @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 ----- # # "v1": cutedsl + deepep low-latency. @@ -2026,16 +2039,14 @@ class ModelOptNvFp4FusedMoEMethod(FusedMoEMethodBase): "NVFP4 quantization was selected, " " dynamic quantization is not supported." ) - # `nvfp4_online` is not a serialized checkpoint format, but after the - # online loader converts each expert it uses the same packed NVFP4 - # weights, block scales, and per-tensor scales as serialized ModelOpt - # NVFP4 checkpoints. Reuse this layout and swap only the weight loader. + # Online conversion changes only weight loading; downstream tensors use + # the same packed weight and scale layout as serialized ModelOpt NVFP4. if is_nvfp4_online: - if not self.enable_flashinfer_trtllm_moe: + if not self.supports_nvfp4_online_moe: raise ValueError( - "--quantization nvfp4_online supports only " - "--moe-runner-backend flashinfer_trtllm or " - "flashinfer_trtllm_routed." + "--quantization nvfp4_online supports flashinfer_trtllm, " + "flashinfer_trtllm_routed, or flashinfer_cutedsl with " + "FlashInfer A2A or DeepEP low_latency." ) # TODO(ch-wan): check if this is needed @@ -2167,17 +2178,19 @@ class ModelOptNvFp4FusedMoEMethod(FusedMoEMethodBase): {"quant_method": FusedMoeWeightScaleSupported.TENSOR.value} ) - w13_input_scale_shape = (layer.num_experts, num_shards) - w13_input_scale = PerTensorScaleParameter( - data=torch.empty(w13_input_scale_shape, dtype=torch.float32), + input_scale_fill = 1.0 if is_nvfp4_online else None + w13_input_scale = _make_per_tensor_scale_parameter( + (layer.num_experts, num_shards), weight_loader=weight_loader, + fill_value=input_scale_fill, ) w13_input_scale._sglang_require_global_experts = True layer.register_parameter("w13_input_scale", w13_input_scale) - w2_input_scale = PerTensorScaleParameter( - data=torch.empty(layer.num_experts, dtype=torch.float32), + w2_input_scale = _make_per_tensor_scale_parameter( + (layer.num_experts,), weight_loader=weight_loader, + fill_value=input_scale_fill, ) w2_input_scale._sglang_require_global_experts = True layer.register_parameter("w2_input_scale", w2_input_scale) diff --git a/python/sglang/srt/layers/quantization/nvfp4_online.py b/python/sglang/srt/layers/quantization/nvfp4_online.py index 3d118f53d..ce607386f 100644 --- a/python/sglang/srt/layers/quantization/nvfp4_online.py +++ b/python/sglang/srt/layers/quantization/nvfp4_online.py @@ -165,11 +165,11 @@ class ModelOptNvFp4OnlineFusedMoEMethod(ModelOptNvFp4FusedMoEMethod): if layer_match is not None else layer_prefix ) - if not self.enable_flashinfer_trtllm_moe: + if not self.supports_nvfp4_online_moe: raise ValueError( - "--quantization nvfp4_online supports only " - "--moe-runner-backend flashinfer_trtllm or " - "flashinfer_trtllm_routed." + "--quantization nvfp4_online supports flashinfer_trtllm, " + "flashinfer_trtllm_routed, or flashinfer_cutedsl with " + "FlashInfer A2A or DeepEP low_latency." ) @staticmethod @@ -190,6 +190,12 @@ class ModelOptNvFp4OnlineFusedMoEMethod(ModelOptNvFp4FusedMoEMethod): "--quantization nvfp4_online expects 2D expert weights, " 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: raise ValueError( "--quantization nvfp4_online requires expert weight K to be " diff --git a/python/sglang/srt/model_executor/runner/flashinfer_autotune.py b/python/sglang/srt/model_executor/runner/flashinfer_autotune.py index 69cf50758..09668a2e1 100644 --- a/python/sglang/srt/model_executor/runner/flashinfer_autotune.py +++ b/python/sglang/srt/model_executor/runner/flashinfer_autotune.py @@ -50,18 +50,23 @@ def should_run_flashinfer_autotune( if mr.server_args.disable_flashinfer_autotune: return False - # CuteDSL v1 (cutedsl runner + deepep a2a) bypasses MoeRunner and must not - # be autotuned -- its _dummy_run would dispatch more tokens per rank than - # SGLANG_DEEPEP_NUM_MAX_DISPATCH_TOKENS_PER_RANK, tripping a DeepEP assert. - # Read server_args directly to avoid depending on initialize_moe_config() - # having already populated the MoE backend globals. - if ( - mr.server_args.moe_runner_backend == "flashinfer_cutedsl" - and mr.server_args.moe_a2a_backend == "deepep" - ): - return False + server_args = mr.server_args + if for_speculative_draft: + backend_str = ( + server_args.speculative_moe_runner_backend or server_args.moe_runner_backend + ) + a2a_backend_str = ( + server_args.speculative_moe_a2a_backend or server_args.moe_a2a_backend + ) + else: + 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