class Mxfp4MoEMethod(FusedMoEMethodBase):
"""MXFP4 MoE quantization method."""
def __init__(self, moe: FusedMoEConfig):
super().__init__(moe)
self.weight_dtype = "mxfp4"
self.mxfp4_backend, self.experts_cls = select_deepseek_v4_mxfp4_moe_backend(moe)
self.max_capture_size = moe.max_capture_size
self._cache_permute_indices: dict[torch.Size, torch.Tensor] = {}
self.moe_kernel: mk.FusedMoEKernel | None = None
# Used for triton kernel precision configs
self.w13_precision_config = None
self.w2_precision_config = None
@property
def supports_eplb(self) -> bool:
return True
@property
def skip_forward_padding(self) -> bool:
# SM100_FI_MXFP4_MXFP8_TRTLLM supports padding with mxfp8 quant
# so can skip the padding in the forward before applying the moe method
return self.mxfp4_backend == Mxfp4MoeBackend.FLASHINFER_TRTLLM_MXFP4_MXFP8
# TODO(bnell): move to MK/expert_class?
@property
def has_unpadded_output(self) -> bool:
return self.mxfp4_backend in [
Mxfp4MoeBackend.FLASHINFER_TRTLLM_MXFP4_MXFP8,
Mxfp4MoeBackend.FLASHINFER_TRTLLM_MXFP4_BF16,
]
def maybe_roundup_sizes(
self,
hidden_size: int,
intermediate_size_per_partition: int,
act_dtype: torch.dtype,
moe_parallel_config: FusedMoEParallelConfig,
) -> tuple[int, int]:
hidden_size, intermediate_size_per_partition = super().maybe_roundup_sizes(
hidden_size=hidden_size,
intermediate_size_per_partition=intermediate_size_per_partition,
act_dtype=act_dtype,
moe_parallel_config=moe_parallel_config,
)
return mxfp4_round_up_hidden_size_and_intermediate_size(
self.mxfp4_backend,
hidden_size,
intermediate_size_per_partition,
activation=self.moe.activation,
)
@staticmethod
def _encode_mxfp4_weight_scale(loaded_weight: torch.Tensor) -> torch.Tensor:
if loaded_weight.dtype == torch.uint8:
return loaded_weight
if loaded_weight.dtype == torch.float8_e8m0fnu:
return loaded_weight.view(torch.uint8)
if loaded_weight.is_floating_point():
return loaded_weight.to(torch.float8_e8m0fnu).view(torch.uint8)
return loaded_weight
@staticmethod
def get_scale_weight_loader(weight_loader):
def mxfp4_weight_loader(
param: torch.nn.Parameter,
loaded_weight: torch.Tensor,
weight_name: str,
shard_id: str,
expert_id: int,
return_success: bool = False,
) -> bool | None:
loaded_weight = Mxfp4MoEMethod._encode_mxfp4_weight_scale(loaded_weight)
return weight_loader(
param,
loaded_weight,
weight_name,
shard_id,
expert_id,
return_success=return_success,
)
return mxfp4_weight_loader
def create_weights(
self,
layer: RoutedExperts,
num_experts: int,
hidden_size: int,
intermediate_size_per_partition: int,
params_dtype: torch.dtype,
**extra_weight_attrs,
):
self.num_experts = num_experts
weight_dtype = torch.uint8
scale_dtype = torch.uint8
mxfp4_block = 32
layer.params_dtype = params_dtype
layer.num_experts = num_experts
self.intermediate_size = intermediate_size_per_partition
self.hidden_size = hidden_size
weight_loader = extra_weight_attrs.pop("weight_loader")
scale_weight_loader = Mxfp4MoEMethod.get_scale_weight_loader(weight_loader)
# Fused gate_up_proj (column parallel)
w13_weight = torch.nn.Parameter(
torch.zeros(
num_experts,
self.moe.w13_num_shards * intermediate_size_per_partition,
hidden_size // 2,
dtype=weight_dtype,
),
requires_grad=False,
)
layer.register_parameter("w13_weight", w13_weight)
set_weight_attrs(w13_weight, extra_weight_attrs)
set_weight_attrs(w13_weight, {"weight_loader": weight_loader})
w13_weight_scale = torch.nn.Parameter(
torch.zeros(
num_experts,
self.moe.w13_num_shards * intermediate_size_per_partition,
hidden_size // mxfp4_block,
dtype=scale_dtype,
),
requires_grad=False,
)
layer.register_parameter("w13_weight_scale", w13_weight_scale)
set_weight_attrs(w13_weight_scale, extra_weight_attrs)
set_weight_attrs(w13_weight_scale, {"weight_loader": scale_weight_loader})
w13_weight_scale.quant_method = "block"
# down_proj (row parallel)
w2_weight = torch.nn.Parameter(
torch.zeros(
num_experts,
hidden_size,
intermediate_size_per_partition // 2,
dtype=weight_dtype,
),
requires_grad=False,
)
layer.register_parameter("w2_weight", w2_weight)
set_weight_attrs(w2_weight, extra_weight_attrs)
set_weight_attrs(w2_weight, {"weight_loader": weight_loader})
w2_weight_scale = torch.nn.Parameter(
torch.zeros(
num_experts,
hidden_size,
intermediate_size_per_partition // mxfp4_block,
dtype=scale_dtype,
),
requires_grad=False,
)
layer.register_parameter("w2_weight_scale", w2_weight_scale)
set_weight_attrs(w2_weight_scale, extra_weight_attrs)
set_weight_attrs(w2_weight_scale, {"weight_loader": scale_weight_loader})
w2_weight_scale.quant_method = "block"
if self.moe.has_bias:
w13_bias = torch.nn.Parameter(
torch.zeros(
num_experts,
self.moe.w13_num_shards * intermediate_size_per_partition,
dtype=torch.bfloat16,
),
requires_grad=False,
)
layer.register_parameter("w13_bias", w13_bias)
set_weight_attrs(w13_bias, extra_weight_attrs)
set_weight_attrs(w13_bias, {"weight_loader": weight_loader})
w2_bias = torch.nn.Parameter(
torch.zeros(
num_experts,
hidden_size,
dtype=torch.bfloat16,
),
requires_grad=False,
)
layer.register_parameter("w2_bias", w2_bias)
set_weight_attrs(w2_bias, extra_weight_attrs)
set_weight_attrs(w2_bias, {"weight_loader": weight_loader})
def _setup_kernel(
self,
layer: RoutedExperts,
w13: torch.Tensor,
w2: torch.Tensor,
w13_scale: torch.Tensor,
w2_scale: torch.Tensor,
w13_bias: torch.Tensor | None = None,
w2_bias: torch.Tensor | None = None,
) -> None:
num_experts = self.num_experts
intermediate_size = self.intermediate_size
hidden_size = self.hidden_size
sf_block_size = 32
# Shape assertions — skipped for SITU since its kernel handles native
# (non-256-aligned) intermediate sizes without prior round-up.
from vllm.model_executor.layers.fused_moe.activation import MoEActivation
if self.moe.activation != MoEActivation.SITU:
assert (
w13.dim() == 3
and w13.shape[0] == num_experts
and w13.shape[1] == intermediate_size * self.moe.w13_num_shards
and w13.shape[2] == hidden_size // 2
)
assert (
w13_scale.dim() == 3
and w13_scale.shape[0] == num_experts
and w13_scale.shape[1] == intermediate_size * self.moe.w13_num_shards
and w13_scale.shape[2] == hidden_size // sf_block_size
)
assert (
w2.dim() == 3
and w2.shape[0] == num_experts
and w2.shape[1] == hidden_size
and w2.shape[2] == intermediate_size // 2
)
assert (
w2_scale.dim() == 3
and w2_scale.shape[1] == hidden_size
and w2_scale.shape[2] == intermediate_size // sf_block_size
)
if w13_bias is not None:
assert (
w13_bias.dim() == 2
and w13_bias.shape[0] == num_experts
and w13_bias.shape[1] == intermediate_size * self.moe.w13_num_shards
)
if w2_bias is not None:
assert (
w2_bias.dim() == 2
and w2_bias.shape[0] == num_experts
and w2_bias.shape[1] == hidden_size
)
# Convert weights to kernel format
w13, w2, w13_scale, w2_scale, w13_bias, w2_bias = (
convert_weight_to_mxfp4_moe_kernel_format(
mxfp4_backend=self.mxfp4_backend,
layer=layer,
w13_weight=w13,
w2_weight=w2,
w13_weight_scale=w13_scale,
w2_weight_scale=w2_scale,
w13_bias=w13_bias,
w2_bias=w2_bias,
_cache_permute_indices=self._cache_permute_indices,
activation=self.moe.activation,
)
)
# For TRITON backends, weights are wrapped tensors from triton_kernels
# that don't support .detach(). Manually assign parameters.
is_gfx1250 = False
if current_platform.is_rocm():
from vllm.platforms.rocm import on_gfx1250
is_gfx1250 = on_gfx1250()
uses_triton_weight_format = self.mxfp4_backend in TRITON_BACKENDS or (
self.mxfp4_backend == Mxfp4MoeBackend.AITER_MXFP4_BF16 and is_gfx1250
)
if not uses_triton_weight_format:
replace_parameter(layer, "w13_weight", w13)
replace_parameter(layer, "w2_weight", w2)
replace_parameter(layer, "w13_weight_scale", w13_scale)
replace_parameter(layer, "w2_weight_scale", w2_scale)
else:
layer.w13_weight = w13
layer.w2_weight = w2
self.w13_precision_config = w13_scale
self.w2_precision_config = w2_scale
if w13_bias is not None and w2_bias is not None:
replace_parameter(layer, "w13_bias", w13_bias)
replace_parameter(layer, "w2_bias", w2_bias)
# Build quant config
self.moe_quant_config = self.get_fused_moe_quant_config(layer)
# Build kernel (modular or monolithic)
if self.moe_quant_config is not None and self.experts_cls is not None:
self.moe_kernel = make_mxfp4_moe_kernel(
moe_quant_config=self.moe_quant_config,
moe_config=self.moe,
mxfp4_backend=self.mxfp4_backend,
experts_cls=self.experts_cls,
routing_tables=layer._expert_routing_tables(),
)
self.moe_kernel.fused_experts.process_weights_after_loading(layer)
def process_weights_after_loading(self, layer):
w13 = layer.w13_weight
w2 = layer.w2_weight
w13_scale = layer.w13_weight_scale
w2_scale = layer.w2_weight_scale
w13_bias = getattr(layer, "w13_bias", None)
w2_bias = getattr(layer, "w2_bias", None)
if self.mxfp4_backend == Mxfp4MoeBackend.NONE:
return
self._setup_kernel(layer, w13, w2, w13_scale, w2_scale, w13_bias, w2_bias)
def get_fused_moe_quant_config(
self,
layer: RoutedExperts,
) -> FusedMoEQuantConfig | None:
w1_bias = getattr(layer, "w13_bias", None)
w2_bias = getattr(layer, "w2_bias", None)
swiglu_limit = getattr(layer, "swiglu_limit", None)
is_gfx1250 = False
if current_platform.is_rocm():
from vllm.platforms.rocm import on_gfx1250
is_gfx1250 = on_gfx1250()
if self.mxfp4_backend in TRITON_BACKENDS or (
self.mxfp4_backend == Mxfp4MoeBackend.AITER_MXFP4_BF16 and is_gfx1250
):
# TRITON backends free w13/w2_weight_scale after swizzling; the
# swizzled scales live inside the precision configs instead.
assert self.w13_precision_config is not None
assert self.w2_precision_config is not None
w1_scale = self.w13_precision_config
w2_scale = self.w2_precision_config
else:
w1_scale = layer.w13_weight_scale
w2_scale = layer.w2_weight_scale
if self.mxfp4_backend == Mxfp4MoeBackend.EMULATION:
# Canonical ``mxfp4`` checkpoints are weight-only W4A16. The
# generic EMULATION config is W4A4, so preserve BF16 activations
# while the fallback dequantizes only the weights.
return mxfp4_w4a16_moe_quant_config(
w1_scale=w1_scale,
w2_scale=w2_scale,
w1_bias=w1_bias,
w2_bias=w2_bias,
gemm1_clamp_limit=swiglu_limit,
)
return make_mxfp4_moe_quant_config(
mxfp4_backend=self.mxfp4_backend,
w1_scale=w1_scale,
w2_scale=w2_scale,
w1_bias=w1_bias,
w2_bias=w2_bias,
swiglu_limit=swiglu_limit,
layer=layer,
)
def apply(
self,
layer: RoutedExperts,
x: torch.Tensor,
topk_weights: torch.Tensor,
topk_ids: torch.Tensor,
shared_experts: SharedExperts | None,
shared_experts_input: torch.Tensor | None,
) -> torch.Tensor:
assert not self.is_monolithic
assert self.moe_kernel is not None
return self.moe_kernel.apply(
hidden_states=x,
w1=layer.w13_weight,
w2=layer.w2_weight,
topk_weights=topk_weights,
topk_ids=topk_ids,
activation=layer.activation,
global_num_experts=layer.global_num_experts,
apply_router_weight_on_input=layer.apply_router_weight_on_input,
expert_map=layer.expert_map,
shared_experts=shared_experts,
shared_experts_input=shared_experts_input,
)
def apply_monolithic(
self,
layer: RoutedExperts,
x: torch.Tensor,
router_logits: torch.Tensor,
input_ids: torch.Tensor | None = None,
) -> torch.Tensor | UnfinalizedMoEOutput:
assert self.is_monolithic
assert self.moe_kernel is not None
return self.moe_kernel.apply_monolithic(
hidden_states=x,
w1=layer.w13_weight,
w2=layer.w2_weight,
router_logits=router_logits,
activation=layer.activation,
global_num_experts=layer.global_num_experts,
expert_map=layer.expert_map,
apply_router_weight_on_input=layer.apply_router_weight_on_input,
num_expert_group=layer.num_expert_group,
topk_group=layer.topk_group,
e_score_correction_bias=layer.e_score_correction_bias,
routed_scaling_factor=layer.routed_scaling_factor,
)