From 9d34c2809f58f3d84ef5dd343733e3f5e86395d5 Mon Sep 17 00:00:00 2001 From: Ziang Li Date: Thu, 13 Aug 2026 17:33:46 -0700 Subject: [PATCH] [FlashInfer v0.6.16] Support FlashInfer CuTe DSL NVFP4 MoE quantization (#28354) --- .../docs/references/environment_variables.mdx | 7 +- python/sglang/srt/arg_groups/overrides.py | 3 +- python/sglang/srt/configs/model_config.py | 2 +- python/sglang/srt/environ.py | 5 +- .../moe/moe_runner/flashinfer_cutedsl.py | 140 ++++++++++++++++-- .../moe/moe_runner/flashinfer_cutlass.py | 3 + .../srt/layers/quantization/base_config.py | 3 + .../srt/layers/quantization/kv_cache.py | 4 +- .../srt/layers/quantization/modelopt_quant.py | 38 +++-- .../srt/layers/quantization/nvfp4_online.py | 29 ++-- python/sglang/srt/server_args.py | 14 +- ...est_flashinfer_nvfp4_online_moe_backend.py | 136 +++++++++++++++++ .../test_flashinfer_trtllm_gen_moe_backend.py | 82 +--------- ...test_update_weights_from_disk_blackwell.py | 72 +++++++-- 14 files changed, 405 insertions(+), 133 deletions(-) create mode 100644 test/registered/backends/test_flashinfer_nvfp4_online_moe_backend.py diff --git a/docs/docs/references/environment_variables.mdx b/docs/docs/references/environment_variables.mdx index 574d3541d..7bca0227e 100644 --- a/docs/docs/references/environment_variables.mdx +++ b/docs/docs/references/environment_variables.mdx @@ -761,9 +761,14 @@ SGLang supports various environment variables that can be used to configure its SGLANG_FLASHINFER_NVFP4_PER_TOKEN_ACTIVATION - Enable FlashInfer TRTLLM NVFP4 per-token activation scaling for serialized modelopt_fp4 checkpoints; checkpoint FP32 activation scales are treated as 1 + Enable FlashInfer TRT-LLM or CuTe DSL v2 (no A2A or FlashInfer A2A) per-token FP32 activation scaling for serialized modelopt_fp4 checkpoints; checkpoint activation scales are treated as 1 false + + SGLANG_FLASHINFER_MOE_FUSED_FINALIZE + Use FlashInfer's fused atomic CUTLASS and CuTe DSL MoE finalize for best performance. Deterministic inference overrides this to false. + true + FLASHINFER_NVFP4_4OVER6 Enable FlashInfer NVFP4 4over6 scaling for NVFP4 per-token activation and online NVFP4 MoE weight quantization paths diff --git a/python/sglang/srt/arg_groups/overrides.py b/python/sglang/srt/arg_groups/overrides.py index e665e8a25..07329b6cd 100644 --- a/python/sglang/srt/arg_groups/overrides.py +++ b/python/sglang/srt/arg_groups/overrides.py @@ -2396,11 +2396,12 @@ def _moe_runner_backend_quant_constraints(view: Any) -> dict: elif moe_runner_backend not in [ "flashinfer_trtllm", "flashinfer_trtllm_routed", + "flashinfer_cutedsl", ]: raise ValueError( "--quantization nvfp4_online supports only " "--moe-runner-backend flashinfer_trtllm or " - "flashinfer_trtllm_routed." + "flashinfer_trtllm_routed, or flashinfer_cutedsl." ) # Ascend runs MXFP8 MoE on the Ascend runner; every backend selected below is # CUDA/ROCm-only. Forcing one here would not merely pick the wrong runner: diff --git a/python/sglang/srt/configs/model_config.py b/python/sglang/srt/configs/model_config.py index 44c4356f5..7e3f53c7b 100644 --- a/python/sglang/srt/configs/model_config.py +++ b/python/sglang/srt/configs/model_config.py @@ -1466,7 +1466,7 @@ class ModelConfig: # so eligible MoE experts are requantized online. "modelopt_fp4": ["modelopt", "fp8"], "modelopt_mixed": ["modelopt"], - "nvfp4_online": ["fp8"], + "nvfp4_online": ["fp8", "modelopt_fp8"], "petit_nvfp4": ["modelopt"], "w8a8_int8": ["compressed-tensors", "compressed_tensors"], "w8a8_fp8": ["compressed-tensors", "compressed_tensors"], diff --git a/python/sglang/srt/environ.py b/python/sglang/srt/environ.py index 150687781..09c406584 100644 --- a/python/sglang/srt/environ.py +++ b/python/sglang/srt/environ.py @@ -866,11 +866,14 @@ class Envs: # Per-rank dispatch capacity of the FlashInfer MoE A2A dispatcher. Unset # means each call site keeps its own default. SGLANG_FLASHINFER_NUM_MAX_DISPATCH_TOKENS_PER_RANK = EnvInt(None) - # Enable NVFP4 per-token activation scaling path for FlashInfer TRT-LLM MoE. + # Enable per-token FP32 activation scaling for serialized ModelOpt FP4 with + # FlashInfer TRT-LLM or CuTe DSL v2 MoE. SGLANG_FLASHINFER_NVFP4_PER_TOKEN_ACTIVATION = EnvBool(False) # Launch the TRT-LLM MoE grouped GEMMs with PDL only at or below this # token count. SGLANG_TRTLLM_MOE_PDL_MAX_TOKENS = EnvInt(8192) + # Use FlashInfer's fused atomic CUTLASS/CuTe DSL MoE finalize. + SGLANG_FLASHINFER_MOE_FUSED_FINALIZE = EnvBool(True) # Master switch for the experimental TRT-LLM LoRA fast path; when OFF (default) every # fine-grained opt switch reads False, keeping non-experimental paths byte-identical. SGLANG_EXPERIMENTAL_LORA_OPTI = EnvBool(False) diff --git a/python/sglang/srt/layers/moe/moe_runner/flashinfer_cutedsl.py b/python/sglang/srt/layers/moe/moe_runner/flashinfer_cutedsl.py index 851cf57b5..82b635db6 100644 --- a/python/sglang/srt/layers/moe/moe_runner/flashinfer_cutedsl.py +++ b/python/sglang/srt/layers/moe/moe_runner/flashinfer_cutedsl.py @@ -6,6 +6,7 @@ from typing import TYPE_CHECKING, Any, Optional import torch +from sglang.srt.environ import envs from sglang.srt.layers.moe.moe_runner.base import ( MoeQuantInfo, MoeRunnerConfig, @@ -237,6 +238,64 @@ def _cutedsl_wrapper_activation_type(activation: str, activation_type_cls: Any) ) +def _make_per_token_global_scale(input_tensor: torch.Tensor) -> torch.Tensor: + from flashinfer.quantization.nvfp4_quantization_utils import ( + current_nvfp4_4over6_config, + make_nvfp4_global_scale, + ) + + return make_nvfp4_global_scale( + input_tensor, + per_token_activation=True, + nvfp4_4over6_config=current_nvfp4_4over6_config(), + ) + + +def refresh_cutedsl_standard_scales_for_weight_update( + layer: torch.nn.Module, +) -> None: + w1_alpha, fc2_input_scale, w2_alpha, used_input_scale = ( + resolve_cutedsl_standard_scales(layer) + ) + if layer.quant_config.use_per_token_activation: + used_input_scale = _make_per_token_global_scale(used_input_scale) + + new_scales = (w1_alpha, fc2_input_scale, w2_alpha) + current_scales = layer._cutedsl_scales + current_input_scale = layer._cutedsl_input_scale + + # Decode CUDA graphs capture these tensor addresses, so reloads must update + # their values without replacing the tensors. + if ( + not isinstance(current_scales, tuple) + or len(current_scales) != len(new_scales) + or not isinstance(current_input_scale, torch.Tensor) + ): + raise RuntimeError( + "CuTe DSL scale metadata changed during weight reload; " + "CUDA graph recapture is required." + ) + scale_pairs = ( + *zip(current_scales, new_scales), + (current_input_scale, used_input_scale), + ) + for current, new in scale_pairs: + if ( + not isinstance(current, torch.Tensor) + or current.shape != new.shape + or current.dtype != new.dtype + or current.device != new.device + ): + raise RuntimeError( + "CuTe DSL scale metadata changed during weight reload; " + "CUDA graph recapture is required." + ) + + with torch.no_grad(): + for current, new in scale_pairs: + current.copy_(new) + + def ensure_cutedsl_wrapper(layer: torch.nn.Module) -> None: """Lazily create CuteDslMoEWrapper and resolve scales on first forward. @@ -246,9 +305,10 @@ def ensure_cutedsl_wrapper(layer: torch.nn.Module) -> None: typically runs during the autotune dummy forward under inference_mode(). We wrap the creation in inference_mode(False) so that those pre-allocated buffers are normal tensors -- inference tensors cannot be inplace-updated - during later CUDA graph capture, which runs outside inference_mode. + during later CUDA graph capture, which runs outside inference_mode. The + resolved scale tensors share this scope because reload updates them in place. """ - if getattr(layer, "_cutedsl_wrapper", None) is not None: + if layer._cutedsl_wrapper is not None: return try: @@ -296,14 +356,17 @@ def ensure_cutedsl_wrapper(layer: torch.nn.Module) -> None: local_expert_offset=layer.moe_ep_rank * layer.num_local_experts, output_dtype=layer.moe_runner_config.params_dtype, device=str(layer.w13_weight.device), + use_fused_finalize=envs.SGLANG_FLASHINFER_MOE_FUSED_FINALIZE.get(), activation_type=_cutedsl_wrapper_activation_type( layer.moe_runner_config.activation, ActivationType ), ) - w1_alpha, fc2_input_scale, w2_alpha, used_input_scale = ( - resolve_cutedsl_standard_scales(layer) - ) + w1_alpha, fc2_input_scale, w2_alpha, used_input_scale = ( + resolve_cutedsl_standard_scales(layer) + ) + if layer.quant_config.use_per_token_activation: + used_input_scale = _make_per_token_global_scale(used_input_scale) layer._cutedsl_scales = (w1_alpha, fc2_input_scale, w2_alpha) layer._cutedsl_input_scale = used_input_scale @@ -356,6 +419,9 @@ class CuteDslFp4MoeQuantInfo(MoeQuantInfo): # v1 only: True when DeepEP pre-quantizes activations to NVFP4. use_nvfp4_dispatch: bool = False + # v2 only: quantize hidden states with per-token dynamic activation scales. + use_per_token_activation: bool = False + # v1 only: SBO down-GEMM overlap args. down_gemm_overlap_args: Optional[DownGemmOverlapArgs] = None @@ -385,11 +451,29 @@ def fused_experts_none_to_flashinfer_cutedsl_fp4( if topk_ids.dtype != torch.int32: topk_ids = topk_ids.to(torch.int32) - x_fp4, x_sf = fp4_quantize( - hidden_states, - quant_info.a1_scale, - sf_vec_size=_FP4_SF_VEC_SIZE, - is_sf_swizzled_layout=False, + if quant_info.use_per_token_activation: + from flashinfer import SfLayout, nvfp4_quantize + + x_fp4, x_sf, per_token_scale = nvfp4_quantize( + hidden_states, + quant_info.a1_scale, + sfLayout=SfLayout.layout_linear, + per_token_activation=True, + backend="cute-dsl", + ) + else: + x_fp4, x_sf = fp4_quantize( + hidden_states, + quant_info.a1_scale, + sf_vec_size=_FP4_SF_VEC_SIZE, + is_sf_swizzled_layout=False, + ) + per_token_scale = None + + seq_len, hidden_size = hidden_states.shape + x_fp4 = x_fp4.reshape(seq_len, hidden_size // 2) + x_sf = x_sf.view(torch.float8_e4m3fn).reshape( + seq_len, hidden_size // _FP4_SF_VEC_SIZE ) output = quant_info.wrapper.run( @@ -404,6 +488,7 @@ def fused_experts_none_to_flashinfer_cutedsl_fp4( w2_weight=quant_info.w2_weight, w2_weight_sf=quant_info.w2_weight_sf, w2_alpha=quant_info.w2_alpha, + per_token_scale=per_token_scale, ) return StandardCombineInput(hidden_states=output) @@ -444,14 +529,38 @@ def fused_experts_flashinfer_to_flashinfer_cutedsl_fp4( topk_ids = topk_ids.to(torch.int32) if x_sf is not None: + if quant_info.use_per_token_activation: + raise ValueError( + "flashinfer_cutedsl per-token activation requires BF16 dispatch " + "so the runner can forward per_token_scale to FlashInfer." + ) # NVFP4 dispatch, inputs are already quantized. x_fp4 = hidden_states + per_token_scale = None else: - x_fp4, x_sf = fp4_quantize( - hidden_states, - quant_info.a1_scale, - sf_vec_size=_FP4_SF_VEC_SIZE, - is_sf_swizzled_layout=False, + if quant_info.use_per_token_activation: + from flashinfer import SfLayout, nvfp4_quantize + + x_fp4, x_sf, per_token_scale = nvfp4_quantize( + hidden_states, + quant_info.a1_scale, + sfLayout=SfLayout.layout_linear, + per_token_activation=True, + backend="cute-dsl", + ) + else: + x_fp4, x_sf = fp4_quantize( + hidden_states, + quant_info.a1_scale, + sf_vec_size=_FP4_SF_VEC_SIZE, + is_sf_swizzled_layout=False, + ) + per_token_scale = None + + seq_len, hidden_size = hidden_states.shape + x_fp4 = x_fp4.reshape(seq_len, hidden_size // 2) + x_sf = x_sf.view(torch.float8_e4m3fn).reshape( + seq_len, hidden_size // _FP4_SF_VEC_SIZE ) output = quant_info.wrapper.run( @@ -466,6 +575,7 @@ def fused_experts_flashinfer_to_flashinfer_cutedsl_fp4( w2_weight=quant_info.w2_weight, w2_weight_sf=quant_info.w2_weight_sf, w2_alpha=quant_info.w2_alpha, + per_token_scale=per_token_scale, ) # Note: output contains routed expert results; shared_expert is handled separately diff --git a/python/sglang/srt/layers/moe/moe_runner/flashinfer_cutlass.py b/python/sglang/srt/layers/moe/moe_runner/flashinfer_cutlass.py index cca9331c3..6a592fb5e 100644 --- a/python/sglang/srt/layers/moe/moe_runner/flashinfer_cutlass.py +++ b/python/sglang/srt/layers/moe/moe_runner/flashinfer_cutlass.py @@ -18,6 +18,7 @@ from sglang.srt.distributed import get_tp_group from sglang.srt.distributed.device_communicators.pynccl_allocator import ( use_symmetric_memory, ) +from sglang.srt.environ import envs from sglang.srt.layers.dp_attention import is_allocation_symmetric from sglang.srt.layers.moe.moe_runner.base import ( MoeQuantInfo, @@ -235,6 +236,7 @@ def _run_flashinfer_cutlass( tune_max_num_tokens=next_power_of_2(x.shape[0]), activation_type=_activation_type(runner_config), enable_alltoall=enable_alltoall, + use_fused_finalize=envs.SGLANG_FLASHINFER_MOE_FUSED_FINALIZE.get(), )[0] if quant_info.quant_type in ("bf16", "fp8"): @@ -387,6 +389,7 @@ def fused_experts_none_to_flashinfer_mxfp4( activation_type=ActivationType.Swiglu, tune_max_num_tokens=next_power_of_2(x.shape[0]), output=out, + use_fused_finalize=envs.SGLANG_FLASHINFER_MOE_FUSED_FINALIZE.get(), ) if do_pad: diff --git a/python/sglang/srt/layers/quantization/base_config.py b/python/sglang/srt/layers/quantization/base_config.py index 56e95c50b..6a81f42b5 100644 --- a/python/sglang/srt/layers/quantization/base_config.py +++ b/python/sglang/srt/layers/quantization/base_config.py @@ -185,6 +185,9 @@ class QuantizationConfig(ABC): if hf_quant_config is None: return None + if user_quant == "nvfp4_online": + return None + # Check if this is a ModelOpt config quant_algo = hf_quant_config.get("quant_algo", "").upper() diff --git a/python/sglang/srt/layers/quantization/kv_cache.py b/python/sglang/srt/layers/quantization/kv_cache.py index 105cb9839..853f0764e 100644 --- a/python/sglang/srt/layers/quantization/kv_cache.py +++ b/python/sglang/srt/layers/quantization/kv_cache.py @@ -56,8 +56,8 @@ class BaseKVCacheMethod(QuantizeMethodBase): if is_fp8_fnuz(): k_scale *= 2 v_scale *= 2 - elif layer.k_scale < 0.0 and layer.v_scale < 0.0: - # If no scales were loaded (both scales are invalid negative + elif layer.k_scale <= 0.0 and layer.v_scale <= 0.0: + # If no scales were loaded (both scales are invalid non-positive # values), use the default value of 1.0 k_scale = 1.0 v_scale = 1.0 diff --git a/python/sglang/srt/layers/quantization/modelopt_quant.py b/python/sglang/srt/layers/quantization/modelopt_quant.py index 10cb2b894..b510394a6 100755 --- a/python/sglang/srt/layers/quantization/modelopt_quant.py +++ b/python/sglang/srt/layers/quantization/modelopt_quant.py @@ -1330,7 +1330,8 @@ class ModelOptFp4Config(ModelOptQuantConfig): and checkpoint-provided scales. - Serialized + per-token FP32 activation scales: set `SGLANG_FLASHINFER_NVFP4_PER_TOKEN_ACTIVATION=1`; use - `flashinfer_trtllm` or `flashinfer_trtllm_routed`. + `flashinfer_trtllm`, `flashinfer_trtllm_routed`, or `flashinfer_cutedsl` + v2 with no A2A or FlashInfer A2A. - BF16/FP16/FP8 MoE + per-tensor FP32 activation scales: quantize expert weights on load, keep dense weights in source precision or FP8, and use 1.0 when the checkpoint has no NVFP4 activation scale. @@ -2498,13 +2499,14 @@ class ModelOptNvFp4FusedMoEMethod(FusedMoEMethodBase): ) # TODO: for flashinfer always do MOE_NVFP4_DISPATCH + use_dispatch_fp4 = not self.quant_config.use_per_token_activation and ( + MOE_NVFP4_DISPATCH or should_use_flashinfer_cutlass_moe_fp4_allgather() + ) + layer.dispatcher.set_quant_config( { "input_global_scale": ( - layer.w13_input_scale_quant - if MOE_NVFP4_DISPATCH - or should_use_flashinfer_cutlass_moe_fp4_allgather() - else None + layer.w13_input_scale_quant if use_dispatch_fp4 else None ) } ) @@ -2565,17 +2567,19 @@ class ModelOptNvFp4FusedMoEMethod(FusedMoEMethodBase): interleave_w13_halves, ) - layer.w13_weight = Parameter( + copy_or_rebind_param( + layer, + "w13_weight", interleave_w13_halves( layer.w13_weight.view(torch.uint8), group_size=64, dim=1 ).contiguous(), - requires_grad=False, ) - layer.w13_weight_scale = Parameter( + copy_or_rebind_param( + layer, + "w13_weight_scale", interleave_w13_halves( layer.w13_weight_scale, group_size=64, dim=1 ).contiguous(), - requires_grad=False, ) # Process w13 weights @@ -2636,6 +2640,7 @@ class ModelOptNvFp4FusedMoEMethod(FusedMoEMethodBase): from sglang.srt.layers.moe.moe_runner.flashinfer_cutedsl import ( _FP4_SF_VEC_SIZE, + refresh_cutedsl_standard_scales_for_weight_update, ) sf_vec_size = _FP4_SF_VEC_SIZE @@ -2644,7 +2649,9 @@ class ModelOptNvFp4FusedMoEMethod(FusedMoEMethodBase): w13_k = layer.w13_weight.shape[2] * 2 w2_m = layer.w2_weight.shape[1] w2_k = layer.w2_weight.shape[2] * 2 - layer.w13_blockscale_mma = Parameter( + copy_or_rebind_param( + layer, + "w13_blockscale_mma", convert_sf_to_mma_layout( layer.w13_blockscale_swizzled.contiguous() .view(torch.uint8) @@ -2654,9 +2661,10 @@ class ModelOptNvFp4FusedMoEMethod(FusedMoEMethodBase): num_groups=num_local_experts, sf_vec_size=sf_vec_size, ), - requires_grad=False, ) - layer.w2_blockscale_mma = Parameter( + copy_or_rebind_param( + layer, + "w2_blockscale_mma", convert_sf_to_mma_layout( layer.w2_blockscale_swizzled.contiguous() .view(torch.uint8) @@ -2666,8 +2674,9 @@ class ModelOptNvFp4FusedMoEMethod(FusedMoEMethodBase): num_groups=num_local_experts, sf_vec_size=sf_vec_size, ), - requires_grad=False, ) + if layer._cutedsl_wrapper is not None: + refresh_cutedsl_standard_scales_for_weight_update(layer) @property def load_up_proj_weight_first(self) -> bool: @@ -2696,6 +2705,8 @@ class ModelOptNvFp4FusedMoEMethod(FusedMoEMethodBase): if moe_runner_backend.is_flashinfer_cutedsl(): import sglang.srt.layers.moe.moe_runner.flashinfer_cutedsl # noqa: F401 – triggers @register_fused_func + layer._cutedsl_wrapper = None + if moe_runner_backend.is_flashinfer_cutlass(): import sglang.srt.layers.moe.moe_runner.flashinfer_cutlass # noqa: F401 @@ -2840,6 +2851,7 @@ class ModelOptNvFp4FusedMoEMethod(FusedMoEMethodBase): a1_scale=layer._cutedsl_input_scale, a2_scale=fc2_input_scale, wrapper=layer._cutedsl_wrapper, + use_per_token_activation=self.quant_config.use_per_token_activation, ) return self.runner.run(dispatch_output, quant_info) diff --git a/python/sglang/srt/layers/quantization/nvfp4_online.py b/python/sglang/srt/layers/quantization/nvfp4_online.py index 84f9d11c2..3ad711f96 100644 --- a/python/sglang/srt/layers/quantization/nvfp4_online.py +++ b/python/sglang/srt/layers/quantization/nvfp4_online.py @@ -69,7 +69,9 @@ class NvFp4OnlineConfig(ModelOptQuantConfig): use_mxfp8: bool = False, ) -> None: source_ignored_layers = self._normalize_ignored_layers(exclude_modules) - fp4_ignored_layers = list(source_ignored_layers) + fp4_ignored_layers = ( + [] if self._use_per_token_activation else list(source_ignored_layers) + ) if ignored_layers_str := envs.SGLANG_FP4_IGNORED_LAYERS.get(): fp4_ignored_layers.extend( layer.strip() @@ -114,8 +116,10 @@ class NvFp4OnlineConfig(ModelOptQuantConfig): quant_method = str(config.get("quant_method", "")).lower() use_mxfp8 = "mxfp8" in quant_method is_checkpoint_fp8_serialized = "fp8" in quant_method or use_mxfp8 - ignored_layers = config.get("ignored_layers") or config.get( - "modules_to_not_convert" + ignored_layers = ( + config.get("ignore") + or config.get("ignored_layers") + or config.get("modules_to_not_convert") ) if isinstance(ignored_layers, str): ignored_layers = [ignored_layers] @@ -142,14 +146,15 @@ class NvFp4OnlineConfig(ModelOptQuantConfig): return Fp8LinearMethod(self) return UnquantizedLinearMethod() if isinstance(layer, FusedMoE): - if is_layer_skipped( + source_layer_ignored = is_layer_skipped( prefix, self.exclude_modules, self.packed_modules_mapping - ) or self.is_layer_excluded(prefix): + ) or self.is_layer_excluded(prefix) + if source_layer_ignored and not self.use_per_token_activation: return None if is_layer_skipped( prefix, self.fp4_ignored_layers, self.packed_modules_mapping ): - if self.is_checkpoint_fp8_serialized: + if self.is_checkpoint_fp8_serialized and not source_layer_ignored: return Fp8MoEMethod(self) return None return ModelOptNvFp4OnlineFusedMoEMethod(self, prefix) @@ -201,15 +206,15 @@ class ModelOptNvFp4OnlineFusedMoEMethod(ModelOptNvFp4FusedMoEMethod): if layer_match is not None else layer_prefix ) - if ( - quant_config.use_per_token_activation - and not self.enable_flashinfer_trtllm_moe + if quant_config.use_per_token_activation and not ( + self.enable_flashinfer_trtllm_moe or self._is_cutedsl_v2_standard ): raise ValueError( "--quantization nvfp4_online requires online per-token FP32 " - "activation scales and supports only flashinfer_trtllm or " - "flashinfer_trtllm_routed. Use --quantization modelopt_fp4 " - "for per-tensor FP32 activation scales." + "activation scales and supports flashinfer_trtllm, " + "flashinfer_trtllm_routed, or flashinfer_cutedsl with no A2A " + "or FlashInfer A2A. Use --quantization modelopt_fp4 for " + "per-tensor FP32 activation scales." ) def prepare_weight_loader(self, layer, weight_loader): diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index 20909a953..afad543b1 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -6670,9 +6670,9 @@ class ServerArgs: if view.moe_runner_backend == "flashinfer_cutedsl": # modelopt_mixed with non-NVFP4 MoE layers is rejected at load time. assert ( - view.quantization in ["modelopt_fp4", "modelopt_mixed"] + view.quantization in ["modelopt_fp4", "modelopt_mixed", "nvfp4_online"] or self.get_model_config().nvfp4_moe_meta is not None - ), f"Invalid quantization '{view.quantization}'. \nFlashInfer CuteDSL MOE currently supports only: 'modelopt_fp4', 'modelopt_mixed' (with NVFP4 MoE layers), or hybrid NVFP4 models." + ), f"Invalid quantization '{view.quantization}'. \nFlashInfer CuteDSL MOE currently supports only: 'modelopt_fp4', 'modelopt_mixed' (with NVFP4 MoE layers), 'nvfp4_online', or hybrid NVFP4 models." assert view.ep_size in [ 1, self.tp_size, @@ -6685,6 +6685,14 @@ class ServerArgs: f"flashinfer_cutedsl supports moe_a2a_backend='none', 'deepep', or 'flashinfer', " f"got '{view.moe_a2a_backend}'." ) + if view.moe_a2a_backend == "deepep" and ( + view.quantization == "nvfp4_online" + or envs.SGLANG_FLASHINFER_NVFP4_PER_TOKEN_ACTIVATION.get() + ): + raise ValueError( + "flashinfer_cutedsl per-token NVFP4 activation requires " + "moe_a2a_backend='none' or 'flashinfer'." + ) if view.moe_runner_backend in ["flashinfer_trtllm", "experimental_sgl_trtllm"]: assert view.quantization in [ @@ -7983,6 +7991,8 @@ class ServerArgs: "1" if self.enable_deterministic_inference else "0" ) self._handle_custom_all_reduce_v2_multinode() + if self.enable_deterministic_inference: + envs.SGLANG_FLASHINFER_MOE_FUSED_FINALIZE.set("0") if self.debug_cuda_graph: if not (is_cuda() or is_hip()): logger.warning( diff --git a/test/registered/backends/test_flashinfer_nvfp4_online_moe_backend.py b/test/registered/backends/test_flashinfer_nvfp4_online_moe_backend.py new file mode 100644 index 000000000..229238f49 --- /dev/null +++ b/test/registered/backends/test_flashinfer_nvfp4_online_moe_backend.py @@ -0,0 +1,136 @@ +import os +import unittest +from types import SimpleNamespace + +import requests + +from sglang.srt.utils import kill_process_tree +from sglang.test.ci.ci_register import register_cuda_ci +from sglang.test.run_eval import run_eval +from sglang.test.test_utils import ( + DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, + DEFAULT_URL_FOR_TEST, + CustomTestCase, + popen_launch_server, +) + +register_cuda_ci(est_time=800, stage="nightly", runner_config="4-gpu-b200") + + +class FlashinferNvFp4OnlineMoeBackendBase: + backend = None + model = None + extra_args = [] + extra_env = {} + eval_args = {} + spec_accept_length_threshold = None + + @classmethod + def setUpClass(cls): + cls.base_url = DEFAULT_URL_FOR_TEST + cls.process = popen_launch_server( + cls.model, + cls.base_url, + timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, + env={**os.environ, **cls.extra_env, "SGLANG_ENABLE_JIT_DEEPGEMM": "False"}, + other_args=[ + *cls.extra_args, + "--moe-runner-backend", + cls.backend, + "--cuda-graph-max-bs-decode", + "128", + "--tp-size", + "4", + "--ep-size", + "4", + "--quantization", + "nvfp4_online", + "--mem-fraction-static", + "0.7", + ], + ) + + @classmethod + def tearDownClass(cls): + kill_process_tree(cls.process.pid) + + def test_gsm8k(self): + args = SimpleNamespace( + base_url=self.base_url, + model=self.model, + eval_name="gsm8k", + num_examples=200, + num_threads=128, + **self.eval_args, + ) + metrics = run_eval(args) + print(f"{metrics=}") + self.assertGreater(metrics["score"], 0.90) + if self.spec_accept_length_threshold is not None: + server_info = requests.get(self.base_url + "/server_info").json() + avg_spec_accept_length = server_info["internal_states"][0][ + "avg_spec_accept_length" + ] + print(f"{avg_spec_accept_length=}") + self.assertGreater( + avg_spec_accept_length, self.spec_accept_length_threshold + ) + + +# Only this class is affected, but the file runs with failfast, so leaving it +# enabled also cuts off the class that sorts after it. +@unittest.skip( + "flashinfer-ai/flashinfer#4486: on SM100/SM103 the TRTLLM_GEN tile-192 BMM " + "path returns non-finite MoE output from FlashInfer 0.6.16.post4 on, so the " + "first real prefill trips the sampler NaN assert and gsm8k scores 0.0. " + "See #34629 for the package bisect." +) +class TestFlashinferTrtllmGenMoeBackendNvFp4Online( + FlashinferNvFp4OnlineMoeBackendBase, CustomTestCase +): + backend = "flashinfer_trtllm" + model = "Qwen/Qwen3-30B-A3B-Instruct-2507-FP8" + eval_args = {"api": "completion", "max_tokens": 512} + extra_env = { + "FLASHINFER_NVFP4_4OVER6": "1", + "FLASHINFER_NVFP4_4OVER6_ERR_MODE": "MSE", + "FLASHINFER_NVFP4_4OVER6_ERR_USE_FAST_MATH": "1", + "FLASHINFER_NVFP4_4OVER6_E4M3_USE_256": "1", + "SGLANG_FP4_IGNORED_LAYERS": ",".join( + ["shared_expert"] + + [f"model.layers.{layer_id}" for layer_id in range(40, 48)] + ), + } + + +class TestFlashinferCuteDSLMoeBackendNvFp4Online( + FlashinferNvFp4OnlineMoeBackendBase, CustomTestCase +): + backend = "flashinfer_cutedsl" + model = "nvidia/NVIDIA-Nemotron-3-Super-120B-A12B-FP8" + extra_args = [ + "--reasoning-parser", + "nemotron_3", + "--tool-call-parser", + "qwen3_coder", + "--speculative-algorithm", + "EAGLE", + "--speculative-num-steps", + "3", + "--speculative-eagle-topk", + "1", + "--speculative-num-draft-tokens", + "4", + ] + eval_args = {"max_tokens": 16000, "temperature": 1.0, "top_p": 0.95} + spec_accept_length_threshold = 2.5 + extra_env = { + "FLASHINFER_NVFP4_4OVER6": "1", + "FLASHINFER_NVFP4_4OVER6_ERR_MODE": "MSE", + "FLASHINFER_NVFP4_4OVER6_ERR_USE_FAST_MATH": "1", + "FLASHINFER_NVFP4_4OVER6_E4M3_USE_256": "1", + } + + +if __name__ == "__main__": + unittest.main() diff --git a/test/registered/backends/test_flashinfer_trtllm_gen_moe_backend.py b/test/registered/backends/test_flashinfer_trtllm_gen_moe_backend.py index 0105ba6bc..2c6cb0bcd 100644 --- a/test/registered/backends/test_flashinfer_trtllm_gen_moe_backend.py +++ b/test/registered/backends/test_flashinfer_trtllm_gen_moe_backend.py @@ -243,58 +243,6 @@ class FlashinferTrtllmGenMoeBackendNVFP4Base: self.assertGreater(metrics["score"], 0.89) -class FlashinferTrtllmGenMoeBackendNvFp4OnlineBase: - backend = None - extra_env = {} - - @classmethod - def setUpClass(cls): - cls.model = "Qwen/Qwen3-Next-80B-A3B-Instruct-FP8" - cls.base_url = DEFAULT_URL_FOR_TEST - cls.process = popen_launch_server( - cls.model, - cls.base_url, - timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, - env={**os.environ, **cls.extra_env, "SGLANG_ENABLE_JIT_DEEPGEMM": "False"}, - other_args=[ - "--attention-backend", - "triton", - "--moe-runner-backend", - cls.backend, - "--cuda-graph-max-bs-decode", - "128", - "--tp-size", - "4", - "--ep-size", - "2", - "--quantization", - "nvfp4_online", - "--mem-fraction-static", - "0.7", - "--mamba-ssm-dtype", - "bfloat16", - ], - ) - - @classmethod - def tearDownClass(cls): - kill_process_tree(cls.process.pid) - - def test_gsm8k(self): - args = SimpleNamespace( - base_url=self.base_url, - model=self.model, - eval_name="gsm8k", - api="completion", - max_tokens=512, - num_examples=200, - num_threads=128, - ) - metrics = run_eval(args) - print(f"{metrics=}") - self.assertGreater(metrics["score"], 0.90) - - class TestFlashinferTrtllmGenMoeBackendFP8( FlashinferTrtllmGenMoeBackendFP8Base, CustomTestCase ): @@ -325,6 +273,12 @@ class TestFlashinferTrtllmGenMoeBackendBF16Routed( backend = "flashinfer_trtllm_routed" +@unittest.skip( + "flashinfer-ai/flashinfer#4486: on SM100/SM103 the TRTLLM_GEN tile-192 BMM " + "path returns non-finite MoE output from FlashInfer 0.6.16.post4 on, so the " + "first real prefill trips the sampler NaN assert and gsm8k scores 0.0. " + "See #34629 for the package bisect." +) class TestFlashinferTrtllmGenMoeBackendNvFp4PerTokenActivationRouted( FlashinferTrtllmGenMoeBackendNVFP4Base, CustomTestCase ): @@ -332,29 +286,5 @@ class TestFlashinferTrtllmGenMoeBackendNvFp4PerTokenActivationRouted( backend = "flashinfer_trtllm_routed" -# Only this class is affected, but the file runs with failfast, so leaving it -# enabled also cuts off the two classes that sort after it. -@unittest.skip( - "flashinfer-ai/flashinfer#4486: on SM100/SM103 the TRTLLM_GEN tile-192 BMM " - "path returns non-finite MoE output from FlashInfer 0.6.16.post4 on, so the " - "first real prefill trips the sampler NaN assert and gsm8k scores 0.0. " - "See #34629 for the package bisect." -) -class TestFlashinferTrtllmGenMoeBackendNvFp4Online( - FlashinferTrtllmGenMoeBackendNvFp4OnlineBase, CustomTestCase -): - extra_env = { - "FLASHINFER_NVFP4_4OVER6": "1", - "FLASHINFER_NVFP4_4OVER6_ERR_MODE": "MSE", - "FLASHINFER_NVFP4_4OVER6_ERR_USE_FAST_MATH": "1", - "FLASHINFER_NVFP4_4OVER6_E4M3_USE_256": "1", - "SGLANG_FP4_IGNORED_LAYERS": ",".join( - ["shared_expert"] - + [f"model.layers.{layer_id}" for layer_id in range(40, 48)] - ), - } - backend = "flashinfer_trtllm" - - if __name__ == "__main__": unittest.main() diff --git a/test/registered/rl/test_update_weights_from_disk_blackwell.py b/test/registered/rl/test_update_weights_from_disk_blackwell.py index d12676e45..3c7eec8f7 100644 --- a/test/registered/rl/test_update_weights_from_disk_blackwell.py +++ b/test/registered/rl/test_update_weights_from_disk_blackwell.py @@ -7,6 +7,11 @@ import unittest import requests +from sglang.srt.constants import ( + GPU_MEMORY_TYPE_CUDA_GRAPH, + GPU_MEMORY_TYPE_KV_CACHE, + GPU_MEMORY_TYPE_WEIGHTS, +) from sglang.srt.utils import kill_process_tree from sglang.test.test_utils import ( DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, @@ -44,10 +49,17 @@ class UpdateWeightsFromDiskBase: ) def _launch_server(self, backend_test_suite): - launch_kwargs = {} - if self.launch_env is not None: - launch_kwargs["env"] = self.launch_env - other_args = backend_test_suite.get("other_args") + launch_kwargs = { + "env": { + "SGLANG_MEMORY_SAVER_CUDA_GRAPH": "1", + **(self.launch_env or {}), + } + } + other_args = ( + *backend_test_suite.get("other_args", ()), + "--enable-memory-saver", + "--cuda-graph-backend-prefill=disabled", + ) return popen_launch_server( self.model, self.base_url, @@ -140,6 +152,18 @@ class UpdateWeightsFromDiskBase: timeout=self.update_timeout, ) + def _offload_engine_and_resume_weights(self): + self._post_json("/release_memory_occupation", {}) + self._post_json( + "/resume_memory_occupation", {"tags": [GPU_MEMORY_TYPE_WEIGHTS]} + ) + + def _resume_kv_cache_and_cuda_graph(self): + self._post_json( + "/resume_memory_occupation", + {"tags": [GPU_MEMORY_TYPE_KV_CACHE, GPU_MEMORY_TYPE_CUDA_GRAPH]}, + ) + def test_parameterized_update_weights_from_disk(self): for backend_test_suite in self.backend_test_suites: case_name = backend_test_suite.get("name", "default") @@ -154,6 +178,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() + self._offload_engine_and_resume_weights() ret = self._run_update_weights( self.model, flush_cache=update_test_suite["flush_cache"], @@ -161,6 +186,7 @@ class UpdateWeightsFromDiskBase: "abort_all_requests" ], ) + self._resume_kv_cache_and_cuda_graph() self.assertTrue(ret.get("success"), f"{ret=}") self.assertEqual(self._get_model_info(), self.model) self._assert_non_empty_decode() @@ -169,7 +195,7 @@ class UpdateWeightsFromDiskBase: baseline_sig, updated_sig ) finally: - kill_process_tree(process.pid) + kill_process_tree(process.pid, wait_timeout=60) class TestServerUpdateWeightsFromDiskMXFP8(UpdateWeightsFromDiskBase, CustomTestCase): @@ -179,8 +205,6 @@ class TestServerUpdateWeightsFromDiskMXFP8(UpdateWeightsFromDiskBase, CustomTest { "name": "flashinfer_trtllm_routed_mxfp8", "other_args": ( - "--base-gpu-id", - "0", "--tp-size", "4", "--dp-size", @@ -202,8 +226,6 @@ class TestServerUpdateWeightsFromDiskNVFP4(UpdateWeightsFromDiskBase, CustomTest { "name": "flashinfer_trtllm_nvfp4", "other_args": ( - "--base-gpu-id", - "0", "--tp-size", "4", "--fp4-gemm-backend", @@ -215,5 +237,37 @@ class TestServerUpdateWeightsFromDiskNVFP4(UpdateWeightsFromDiskBase, CustomTest ) +class TestServerUpdateWeightsFromDiskNVFP4CuteDSL( + UpdateWeightsFromDiskBase, CustomTestCase +): + model = "nvidia/Qwen3-30B-A3B-NVFP4" + launch_env = { + "SGLANG_FLASHINFER_NVFP4_PER_TOKEN_ACTIVATION": "1", + "SGLANG_FLASHINFER_MOE_FUSED_FINALIZE": "1", + "FLASHINFER_NVFP4_4OVER6": "1", + "FLASHINFER_NVFP4_4OVER6_ERR_MODE": "MSE", + "FLASHINFER_NVFP4_4OVER6_ERR_USE_FAST_MATH": "1", + "FLASHINFER_NVFP4_4OVER6_E4M3_USE_256": "1", + } + backend_test_suites = ( + { + "name": "flashinfer_cutedsl_nvfp4", + "other_args": ( + "--tp-size", + "4", + "--ep-size", + "4", + "--fp4-gemm-backend", + "flashinfer_cutedsl", + "--moe-runner-backend", + "flashinfer_cutedsl", + "--moe-a2a-backend", + "none", + "--enable-deterministic-inference", + ), + }, + ) + + if __name__ == "__main__": unittest.main()