From ee45599ec814a5bc267eea70d31bfc2dd1205689 Mon Sep 17 00:00:00 2001 From: Duncan Moss Date: Mon, 13 Oct 2025 16:49:43 -0700 Subject: [PATCH] finished integration Signed-off-by: Duncan Moss --- .../fused_moe/flashinfer_cutlass_moe.py | 15 +++++- .../flashinfer_cutlass_prepare_finalize.py | 49 ++++++++++++------- .../model_executor/layers/quantization/fp8.py | 1 + .../quantization/utils/flashinfer_utils.py | 11 +++-- 4 files changed, 52 insertions(+), 24 deletions(-) diff --git a/vllm/model_executor/layers/fused_moe/flashinfer_cutlass_moe.py b/vllm/model_executor/layers/fused_moe/flashinfer_cutlass_moe.py index 85ce77fb1f7..85f24ebe6d8 100644 --- a/vllm/model_executor/layers/fused_moe/flashinfer_cutlass_moe.py +++ b/vllm/model_executor/layers/fused_moe/flashinfer_cutlass_moe.py @@ -57,6 +57,7 @@ class FlashInferExperts(mk.FusedMoEPermuteExpertsUnpermute): tp_rank: int = 0, tp_size: int = 1, use_dp: bool = False, + use_deepseek_fp8_block_scale: bool = False, ): super().__init__(quant_config) assert quant_config.quant_dtype in ("nvfp4", torch.float8_e4m3fn, None), ( @@ -69,7 +70,8 @@ class FlashInferExperts(mk.FusedMoEPermuteExpertsUnpermute): self.tp_size = tp_size self.out_dtype = out_dtype self.use_dp = use_dp - + self.use_deepseek_fp8_block_scale = use_deepseek_fp8_block_scale + @property def activation_formats( self, @@ -147,7 +149,7 @@ class FlashInferExperts(mk.FusedMoEPermuteExpertsUnpermute): "Only activation silu is supported in FlashInferExperts" ) - if self.quant_dtype == torch.float8_e4m3fn: + if self.quant_dtype == torch.float8_e4m3fn and not self.use_deepseek_fp8_block_scale: quant_scales = [ self.g1_alphas, self.a2_gscale, @@ -176,6 +178,14 @@ class FlashInferExperts(mk.FusedMoEPermuteExpertsUnpermute): # FlashInfer API requires weight to be long for nvfp4 fc1_expert_weights = w1.view(torch.long) fc2_expert_weights = w2.view(torch.long) + elif self.use_deepseek_fp8_block_scale: + quant_scales = [ + self.w1_scale, + self.w2_scale, + ] + a1q_scale = None + fc1_expert_weights = w1 + fc2_expert_weights = w2 else: quant_scales = None a1q_scale = None @@ -196,6 +206,7 @@ class FlashInferExperts(mk.FusedMoEPermuteExpertsUnpermute): ep_size=self.ep_size, ep_rank=self.ep_rank, output=output, + use_deepseek_fp8_block_scale=self.use_deepseek_fp8_block_scale, ) diff --git a/vllm/model_executor/layers/fused_moe/flashinfer_cutlass_prepare_finalize.py b/vllm/model_executor/layers/fused_moe/flashinfer_cutlass_prepare_finalize.py index 051abbcb794..6293c23acc1 100644 --- a/vllm/model_executor/layers/fused_moe/flashinfer_cutlass_prepare_finalize.py +++ b/vllm/model_executor/layers/fused_moe/flashinfer_cutlass_prepare_finalize.py @@ -28,11 +28,13 @@ class FlashInferCutlassMoEPrepareAndFinalize(mk.FusedMoEPrepareAndFinalize): self, use_dp: bool, num_dispatchers: int = 1, + use_deepseek_fp8_block_scale: bool = False, ): super().__init__() self.num_dispatchers_ = num_dispatchers self.use_dp = use_dp self.local_tokens = None + self.use_deepseek_fp8_block_scale = use_deepseek_fp8_block_scale @property def activation_format(self) -> mk.FusedMoEActivationFormat: @@ -73,8 +75,9 @@ class FlashInferAllToAllMoEPrepareAndFinalize(FlashInferCutlassMoEPrepareAndFina self, use_dp: bool, num_dispatchers: int = 1, + use_deepseek_fp8_block_scale: bool = False, ): - super().__init__(use_dp, num_dispatchers) + super().__init__(use_dp, num_dispatchers, use_deepseek_fp8_block_scale) self.alltoall_info = None # Initialize all2all_manager only for DP case @@ -98,14 +101,15 @@ class FlashInferAllToAllMoEPrepareAndFinalize(FlashInferCutlassMoEPrepareAndFina if not self.use_dp: # Non-DP case: standard quantization - a1q, a1q_scale = moe_kernel_quantize_input( - a1, - quant_config.a1_gscale, - quant_config.quant_dtype, - quant_config.per_act_token_quant, - quant_config.block_shape, - is_fp4_scale_swizzled=not self.use_dp, - ) + if not self.use_deepseek_fp8_block_scale: + a1q, a1q_scale = moe_kernel_quantize_input( + a1, + quant_config.a1_gscale, + quant_config.quant_dtype, + quant_config.per_act_token_quant, + quant_config.block_shape, + is_fp4_scale_swizzled=not self.use_dp, + ) else: # DP case: use FlashInfer AllToAll global_num_tokens_cpu = get_local_sizes() @@ -154,8 +158,9 @@ class FlashInferAllGatherMoEPrepareAndFinalize(FlashInferCutlassMoEPrepareAndFin self, use_dp: bool, num_dispatchers: int = 1, + use_deepseek_fp8_block_scale: bool = False, ): - super().__init__(use_dp, num_dispatchers) + super().__init__(use_dp, num_dispatchers, use_deepseek_fp8_block_scale) def prepare( self, @@ -173,14 +178,19 @@ class FlashInferAllGatherMoEPrepareAndFinalize(FlashInferCutlassMoEPrepareAndFin if not self.use_dp: return a1, None, None, topk_ids, topk_weights - a1q, a1q_scale = moe_kernel_quantize_input( - a1, - quant_config.a1_gscale, - quant_config.quant_dtype, - quant_config.per_act_token_quant, - quant_config.block_shape, - is_fp4_scale_swizzled=not self.use_dp, - ) + if not self.use_deepseek_fp8_block_scale: + a1q, a1q_scale = moe_kernel_quantize_input( + a1, + quant_config.a1_gscale, + quant_config.quant_dtype, + quant_config.per_act_token_quant, + quant_config.block_shape, + is_fp4_scale_swizzled=not self.use_dp, + ) + else: + a1q = a1 + a1q_scale = None + topk_weights, topk_ids, a1q, a1q_scale = get_dp_group().all_gatherv( [topk_weights, topk_ids, a1q, a1q_scale], dim=0, @@ -300,6 +310,7 @@ def create_flashinfer_prepare_finalize( use_dp: bool, use_nvfp4: bool = False, enable_alltoallv: bool = False, + use_deepseek_fp8_block_scale: bool = False, ) -> FlashInferCutlassMoEPrepareAndFinalize: """Factory function to create the appropriate FlashInfer implementation.""" if use_nvfp4: @@ -308,4 +319,4 @@ def create_flashinfer_prepare_finalize( else: return FlashInferAllGatherMoEPrepareAndFinalize(use_dp) # Fp8 only supports AllGather - return FlashInferAllGatherMoEPrepareAndFinalize(use_dp) + return FlashInferAllGatherMoEPrepareAndFinalize(use_dp=use_dp, use_deepseek_fp8_block_scale=use_deepseek_fp8_block_scale) diff --git a/vllm/model_executor/layers/quantization/fp8.py b/vllm/model_executor/layers/quantization/fp8.py index 03eca199d53..dc5844d165d 100644 --- a/vllm/model_executor/layers/quantization/fp8.py +++ b/vllm/model_executor/layers/quantization/fp8.py @@ -1349,6 +1349,7 @@ class Fp8MoEMethod(FusedMoEMethodBase): global_num_experts=global_num_experts, expert_map=expert_map, apply_router_weight_on_input=apply_router_weight_on_input, + use_deepseek_fp8_block_scale=self.block_quant is not None, ) else: from vllm.model_executor.layers.fused_moe import fused_experts diff --git a/vllm/model_executor/layers/quantization/utils/flashinfer_utils.py b/vllm/model_executor/layers/quantization/utils/flashinfer_utils.py index 50ea049c3d5..17e220b4bff 100644 --- a/vllm/model_executor/layers/quantization/utils/flashinfer_utils.py +++ b/vllm/model_executor/layers/quantization/utils/flashinfer_utils.py @@ -186,16 +186,18 @@ def register_moe_scaling_factors(layer: torch.nn.Module) -> None: def build_flashinfer_fp8_cutlass_moe_prepare_finalize( moe: FusedMoEConfig | None, + use_deepseek_fp8_block_scale: bool = False, ) -> mk.FusedMoEPrepareAndFinalize: """Create a FlashInfer CUTLASS fused-MoE prepare finalize kernel""" use_dp = moe.moe_parallel_config.dp_size > 1 if moe is not None else False - return create_flashinfer_prepare_finalize(use_dp) + return create_flashinfer_prepare_finalize(use_dp, use_deepseek_fp8_block_scale=use_deepseek_fp8_block_scale) def select_cutlass_fp8_gemm_impl( moe: FusedMoEConfig | None, quant_config: FusedMoEQuantConfig, out_dtype: torch.dtype | None = None, + use_deepseek_fp8_block_scale: bool = False, ) -> mk.FusedMoEPermuteExpertsUnpermute: """Return a GEMM *experts* implementation for fused-MoE layers""" @@ -207,12 +209,14 @@ def select_cutlass_fp8_gemm_impl( ep_size=moe.moe_parallel_config.ep_size, tp_rank=moe.moe_parallel_config.tp_rank, tp_size=moe.moe_parallel_config.tp_size, + use_deepseek_fp8_block_scale=use_deepseek_fp8_block_scale, ) assert out_dtype is not None, "If moe config is None, out_dtype must be passed" return FlashInferExperts( out_dtype=out_dtype, quant_config=quant_config, + use_deepseek_fp8_block_scale=use_deepseek_fp8_block_scale, ) @@ -226,14 +230,15 @@ def flashinfer_cutlass_moe_fp8( global_num_experts: int = -1, expert_map: torch.Tensor | None = None, apply_router_weight_on_input: bool = False, + use_deepseek_fp8_block_scale: bool = False, ) -> torch.Tensor: quant_config = layer.quant_method.get_fused_moe_quant_config(layer) assert quant_config is not None fused_experts = mk.FusedMoEModularKernel( - build_flashinfer_fp8_cutlass_moe_prepare_finalize(moe=None), + build_flashinfer_fp8_cutlass_moe_prepare_finalize(moe=None, use_deepseek_fp8_block_scale=use_deepseek_fp8_block_scale), select_cutlass_fp8_gemm_impl( - moe=None, quant_config=quant_config, out_dtype=hidden_states.dtype + moe=None, quant_config=quant_config, out_dtype=hidden_states.dtype, use_deepseek_fp8_block_scale=use_deepseek_fp8_block_scale ), )