finished integration

Signed-off-by: Duncan Moss <djm.moss@gmail.com>
This commit is contained in:
Duncan Moss
2025-11-04 19:52:25 +00:00
committed by Alexander Matveev
parent 1fb4217a05
commit ee45599ec8
4 changed files with 52 additions and 24 deletions
@@ -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,
)
@@ -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)
@@ -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
@@ -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
),
)