class FlashInferCuteDSLExperts(mk.FusedMoEExpertsModular):
"""CuteDSL NvFP4 MoE experts using the FlashInfer functional API.
Uses Standard activation format (non-batched). The kernel handles
routing, expert computation, and reduction internally.
Supports expert parallelism natively.
"""
def __init__(
self,
moe_config: FusedMoEConfig,
quant_config: FusedMoEQuantConfig,
per_token_activation: bool = False,
):
super().__init__(
moe_config=moe_config,
quant_config=quant_config,
)
assert quant_config.quant_dtype == "nvfp4", (
"Only nvfp4 quantization is currently supported."
)
self.out_dtype = moe_config.in_dtype
self.hidden_dim = moe_config.hidden_dim
self.intermediate_size_per_partition = (
moe_config.intermediate_size_per_partition
)
self.topk = moe_config.experts_per_token
self.local_num_experts = moe_config.num_local_experts
self.global_num_experts = moe_config.num_experts
self.ep_rank = moe_config.moe_parallel_config.ep_rank
self.local_expert_offset = self.ep_rank * self.local_num_experts
self.gemm1_alpha = quant_config.gemm1_alpha
self.gemm1_beta = quant_config.gemm1_beta
self.gemm1_clamp_limit = quant_config.gemm1_clamp_limit
self.situ_beta = moe_config.activation_situ_beta
self.situ_linear_beta = moe_config.activation_situ_linear_beta
self.per_token_activation = per_token_activation
self.per_token_global_scale = None
if per_token_activation:
assert quant_config.a2_gscale is not None
self.per_token_global_scale = quant_config.a2_gscale.new_full(
(1,),
NVFP4_PER_TOKEN_BASE_GLOBAL_SCALE,
)
def process_weights_after_loading(self, layer: torch.nn.Module) -> None:
layer.w13_weight_scale_2.data.mul_(layer.w13_input_scale)
layer.w2_weight_scale_2.data.mul_(layer.w2_input_scale)
@staticmethod
def activation_format() -> mk.FusedMoEActivationFormat:
return mk.FusedMoEActivationFormat.Standard
@staticmethod
def _supports_current_device() -> bool:
p = current_platform
return (
p.is_cuda()
and p.is_device_capability_family(100)
and has_flashinfer_cutedsl_moe_nvfp4()
)
@staticmethod
def _supports_no_act_and_mul() -> bool:
return True
@staticmethod
def _supports_quant_scheme(
weight_key: QuantKey | None,
activation_key: QuantKey | None,
) -> bool:
SUPPORTED_W_A = [
(kNvfp4Static, kNvfp4Dynamic),
(kNvfp4Static, kNvfp4DynamicToken),
]
return (weight_key, activation_key) in SUPPORTED_W_A
@staticmethod
def _supports_activation(activation: MoEActivation) -> bool:
return activation in (
MoEActivation.SILU,
MoEActivation.SWIGLUOAI,
MoEActivation.SWIGLUOAI_UNINTERLEAVE,
MoEActivation.RELU2_NO_MUL,
MoEActivation.SITU,
)
@staticmethod
def _supports_parallel_config(
moe_parallel_config: FusedMoEParallelConfig,
) -> bool:
return True
def finalize_weight_and_reduce_impl(self) -> mk.TopKWeightAndReduce:
return TopKWeightAndReduceNoOP()
@property
def expects_unquantized_inputs(self) -> bool:
return self.per_token_activation
def workspace_shapes(
self,
M: int,
N: int,
K: int,
topk: int,
global_num_experts: int,
local_num_experts: int,
expert_tokens_meta: mk.ExpertTokensMetadata | None,
activation: MoEActivation,
) -> tuple[tuple[int, ...], tuple[int, ...], tuple[int, ...]]:
workspace1 = (0,)
workspace2 = (0,)
expected_hidden_dim = K if self.expects_unquantized_inputs else K * 2
assert self.hidden_dim == expected_hidden_dim
output = (M, self.hidden_dim)
return (workspace1, workspace2, output)
def apply(
self,
output: torch.Tensor,
hidden_states: torch.Tensor,
w1: torch.Tensor,
w2: torch.Tensor,
topk_weights: torch.Tensor,
topk_ids: torch.Tensor,
activation: MoEActivation,
global_num_experts: int,
expert_map: torch.Tensor | None,
a1q_scale: torch.Tensor | None,
a2_scale: torch.Tensor | None,
workspace13: torch.Tensor | None,
workspace2: torch.Tensor | None,
expert_tokens_meta: mk.ExpertTokensMetadata | None,
apply_router_weight_on_input: bool | None,
):
assert self.quant_dtype == "nvfp4"
assert self.w1_scale is not None
assert self.w2_scale is not None
if self.expects_unquantized_inputs:
hidden_states, block_scale, per_token_scale = (
quantize_nvfp4_per_token_input(hidden_states)
)
fc2_input_scale = self.per_token_global_scale
else:
assert a1q_scale is not None
block_scale = a1q_scale
per_token_scale = None
fc2_input_scale = self.a2_gscale
assert block_scale is not None
assert fc2_input_scale is not None
# a1q_scale is (M, K//16) float8_e4m3fn from fp4_quantize.
# The functional API expects x_sf with trailing dim: (M, K//16, 1).
x_sf = block_scale.unsqueeze(-1)
# The kernel defaults swiglu_{alpha,beta,limit} to the plain-SwiGLU
# values, so only forward the ones the model actually sets.
swiglu_params: dict[str, float | None] = {}
if activation == MoEActivation.SILU:
swiglu_params = {"swiglu_limit": self.gemm1_clamp_limit}
elif activation in (
MoEActivation.SWIGLUOAI,
MoEActivation.SWIGLUOAI_UNINTERLEAVE,
):
swiglu_params = {
"swiglu_alpha": self.gemm1_alpha,
"swiglu_beta": self.gemm1_beta,
"swiglu_limit": self.gemm1_clamp_limit,
}
elif activation == MoEActivation.SITU:
# The cute_dsl kernel keys SiTU on situ_beta and requires
# activation_type to stay a base type (ActivationType.Situ is
# rejected by normalize_cute_dsl_moe_activation_type), so the
# Swiglu base type is passed below and SiTU rides the betas.
if self.situ_beta is None:
raise ValueError(
"SITU activation requires moe_config.activation_situ_beta"
)
swiglu_params = {
"situ_beta": self.situ_beta,
"situ_linear_beta": self.situ_linear_beta,
}
swiglu_kwargs = {k: v for k, v in swiglu_params.items() if v is not None}
per_token_kwargs = (
{"per_token_scale": per_token_scale} if self.per_token_activation else {}
)
flashinfer_cute_dsl_fused_moe_nvfp4(
x=hidden_states,
x_sf=x_sf,
token_selected_experts=topk_ids.to(torch.int32),
token_final_scales=topk_weights.float(),
w1_weight=w1,
w1_weight_sf=self.w1_scale,
w1_alpha=self.g1_alphas,
fc2_input_scale=fc2_input_scale,
w2_weight=w2,
w2_weight_sf=self.w2_scale,
w2_alpha=self.g2_alphas,
num_experts=self.global_num_experts,
top_k=self.topk,
num_local_experts=self.local_num_experts,
local_expert_offset=self.local_expert_offset,
moe_output=output,
activation_type=activation_to_flashinfer_int(
MoEActivation.SILU if activation == MoEActivation.SITU else activation
),
**swiglu_kwargs,
**per_token_kwargs,
)