forked from Karylab-cklius/vllm
committed by
Alexander Matveev
parent
1fb4217a05
commit
ee45599ec8
@@ -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
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user