Enable B12x backend for non-gated MoEs (like Nemotron) (#43328)

Signed-off-by: Andrii Skliar <[email protected]>
Co-authored-by: Andrii Skliar <[email protected]>
This commit is contained in:
Andrii Skliar
2026-07-06 12:40:07 -07:00
committed by GitHub
co-authored by Andrii Skliar
parent b136cc2c2c
commit b1384f5ec6
3 changed files with 225 additions and 40 deletions
+157 -22
View File
@@ -1,6 +1,8 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
from types import SimpleNamespace
import pytest
import torch
@@ -8,8 +10,7 @@ from vllm.platforms import current_platform
if not current_platform.is_device_capability_family(120):
pytest.skip(
reason="FlashInfer CuteDSL SM12x MoE requires SM120 "
"(RTX Pro 6000 / DGX Spark).",
reason="FlashInfer B12x MoE requires SM120 (RTX Pro 6000 / DGX Spark).",
allow_module_level=True,
)
@@ -18,8 +19,8 @@ from vllm.utils.flashinfer import has_flashinfer_b12x_moe
if not has_flashinfer_b12x_moe():
pytest.skip(
reason=(
"FlashInfer cute_dsl_fused_moe_nvfp4 / convert_sf_to_mma_layout "
"not available in installed FlashInfer (needs PRs #3051 and #3066)."
"FlashInfer B12xMoEWrapper not available in installed "
"FlashInfer (needs PR #3080)."
),
allow_module_level=True,
)
@@ -40,7 +41,6 @@ from vllm.model_executor.layers.fused_moe.config import nvfp4_moe_quant_config
from vllm.model_executor.layers.fused_moe.experts.flashinfer_b12x_moe import (
FlashInferB12xExperts,
)
from vllm.utils.flashinfer import flashinfer_convert_sf_to_mma_layout
from vllm.utils.torch_utils import set_random_seed
# Dimensions chosen to satisfy FP4 alignment requirements (k multiple of 256,
@@ -59,7 +59,7 @@ def _reorder_gate_up_to_up_gate(
) -> tuple[torch.Tensor, torch.Tensor]:
"""Swap gate and up-projection halves along dim=1 to [up, gate] order.
The SM12x kernel expects weights in [up (w3), gate (w1)] order while the
The B12x kernel expects weights in [up (w3), gate (w1)] order while the
BF16 reference uses [gate (w1), up (w3)]. This replicates the reordering
done at model-load time by ``prepare_nvfp4_moe_layer_for_fi_or_cutlass``.
"""
@@ -70,6 +70,22 @@ def _reorder_gate_up_to_up_gate(
)
def _process_b12x_weights(
experts: FlashInferB12xExperts,
w1_scale: torch.Tensor,
w2_scale: torch.Tensor,
w1_scale_2: torch.Tensor,
w2_scale_2: torch.Tensor,
) -> None:
layer = SimpleNamespace(
w13_weight_scale=w1_scale,
w13_weight_scale_2=w1_scale_2,
w2_weight_scale=w2_scale,
w2_weight_scale_2=w2_scale_2,
)
experts.process_weights_after_loading(layer)
@pytest.mark.parametrize("m,n,k", MNK_FACTORS)
@pytest.mark.parametrize("e", [8, 16])
@pytest.mark.parametrize("topk", [1, 2, 4])
@@ -174,22 +190,12 @@ def test_flashinfer_b12x_moe(
moe_config=moe_config,
quant_config=quant_config,
)
# In production, process_weights_after_loading computes these after
# normalizing block scales. In the test the scales are already in final
# form (global_scale=1.0), so we compute the MMA layouts directly.
num_experts_w1, m1, k1_sf = w1_blockscale.shape
experts.w1_sf_mma = flashinfer_convert_sf_to_mma_layout(
w1_blockscale.reshape(num_experts_w1 * m1, k1_sf),
m=m1,
k=k1_sf * 16,
num_groups=num_experts_w1,
)
num_experts_w2, m2, k2_sf = w2_blockscale.shape
experts.w2_sf_mma = flashinfer_convert_sf_to_mma_layout(
w2_blockscale.reshape(num_experts_w2 * m2, k2_sf),
m=m2,
k=k2_sf * 16,
num_groups=num_experts_w2,
_process_b12x_weights(
experts,
w1_blockscale,
w2_blockscale,
ones_e,
ones_e,
)
kernel = mk.FusedMoEKernel(
@@ -224,5 +230,134 @@ def test_flashinfer_b12x_moe(
torch.testing.assert_close(sm12x_output, torch_output, atol=2e-1, rtol=2e-1)
@pytest.mark.parametrize("m,n,k", MNK_FACTORS)
@pytest.mark.parametrize("e", [8, 16])
@pytest.mark.parametrize("topk", [1, 2, 4])
@pytest.mark.parametrize("dtype", [torch.bfloat16])
@torch.inference_mode()
def test_flashinfer_b12x_moe_relu2(
m: int,
n: int,
k: int,
e: int,
topk: int,
dtype: torch.dtype,
workspace_init,
):
"""Test FlashInferB12xExperts with ReLU2 (non-gated) activation.
ReLU2 is used by Nemotron-H style models. Unlike the gated SiLU
path, w1 has shape [E, N, K] (not [E, 2N, K]) and the activation
is relu(x)^2 without a gate/up split.
"""
set_random_seed(7)
with set_current_vllm_config(
VllmConfig(parallel_config=ParallelConfig(pipeline_parallel_size=1))
):
a = torch.randn((m, k), device="cuda", dtype=dtype) / 10
# Non-gated: w1 shape is (e, n, k), not (e, 2n, k).
w1_bf16 = torch.randn((e, n, k), device="cuda", dtype=dtype) / 15
w2_bf16 = torch.randn((e, k, n), device="cuda", dtype=dtype) / 15
gs = torch.ones(1, device="cuda", dtype=torch.float32)
sf_vec_size = 16
# W1: no gate/up reordering for non-gated.
w1_flat = w1_bf16.reshape(e * n, k)
w1_q_flat, w1_sf_flat = fp4_quantize(
w1_flat,
global_scale=gs,
sf_vec_size=sf_vec_size,
is_sf_swizzled_layout=True,
)
w1_q = w1_q_flat.view(e, n, k // 2)
w1_blockscale = w1_sf_flat.view(e, n, w1_sf_flat.shape[1])
w2_flat = w2_bf16.reshape(e * k, n)
w2_q_flat, w2_sf_flat = fp4_quantize(
w2_flat,
global_scale=gs,
sf_vec_size=sf_vec_size,
is_sf_swizzled_layout=True,
)
w2_q = w2_q_flat.view(e, k, n // 2)
w2_blockscale = w2_sf_flat.view(e, k, w2_sf_flat.shape[1])
ones_e = torch.ones(e, device="cuda", dtype=torch.float32)
quant_config = nvfp4_moe_quant_config(
g1_alphas=ones_e,
g2_alphas=ones_e,
a1_gscale=ones_e,
a2_gscale=ones_e,
w1_scale=w1_blockscale,
w2_scale=w2_blockscale,
)
moe_config = make_dummy_moe_config(
num_experts=e,
experts_per_token=topk,
hidden_dim=k,
intermediate_size=n,
in_dtype=dtype,
activation=MoEActivation.RELU2_NO_MUL,
)
experts = FlashInferB12xExperts(
moe_config=moe_config,
quant_config=quant_config,
)
_process_b12x_weights(
experts,
w1_blockscale,
w2_blockscale,
ones_e,
ones_e,
)
kernel = mk.FusedMoEKernel(
maybe_make_prepare_finalize(
moe=moe_config,
quant_config=quant_config,
allow_new_interface=True,
use_monolithic=False,
),
experts,
inplace=False,
)
score = torch.randn((m, e), device="cuda", dtype=dtype)
topk_weights, topk_ids, _ = fused_topk(a, score, topk, renormalize=False)
b12x_output = kernel.apply(
hidden_states=a,
w1=w1_q,
w2=w2_q,
topk_weights=topk_weights,
topk_ids=topk_ids,
global_num_experts=e,
activation=MoEActivation.RELU2_NO_MUL,
apply_router_weight_on_input=False,
expert_map=None,
)
torch_output = torch_moe(
a,
w1_bf16,
w2_bf16,
score,
topk,
activation=MoEActivation.RELU2_NO_MUL,
)
torch.testing.assert_close(
b12x_output,
torch_output,
atol=2e-1,
rtol=2e-1,
)
if __name__ == "__main__":
test_flashinfer_b12x_moe(16, 128, 256, 8, 2, torch.bfloat16)
+2 -1
View File
@@ -55,6 +55,7 @@ def make_dummy_moe_config(
intermediate_size: int = 1,
in_dtype: torch.dtype = torch.bfloat16,
max_num_tokens: int = 512,
activation: MoEActivation = MoEActivation.SILU,
) -> FusedMoEConfig:
"""
This is a dummy config for the mk constructor interface
@@ -73,7 +74,7 @@ def make_dummy_moe_config(
else num_experts,
num_logical_experts=num_experts,
moe_parallel_config=FusedMoEParallelConfig.make_no_parallel(),
activation=MoEActivation.SILU,
activation=activation,
in_dtype=in_dtype,
device="cuda",
routing_method=RoutingMethodType.TopK,
@@ -1,6 +1,8 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
from typing import Any
import torch
import vllm.model_executor.layers.fused_moe.modular_kernel as mk
@@ -20,7 +22,6 @@ from vllm.model_executor.layers.quantization.utils.quant_utils import (
)
from vllm.platforms import current_platform
from vllm.utils.flashinfer import (
flashinfer_b12x_fused_moe,
flashinfer_convert_sf_to_mma_layout,
has_flashinfer_b12x_moe,
)
@@ -42,6 +43,11 @@ class FlashInferB12xExperts(mk.FusedMoEExpertsModular):
Only NVFP4 (kNvfp4Static/kNvfp4Dynamic) quantization is supported.
"""
_ACTIVATION_MAP: dict[MoEActivation, str] = {
MoEActivation.SILU: "silu",
MoEActivation.RELU2_NO_MUL: "relu2",
}
def __init__(
self,
moe_config: FusedMoEConfig,
@@ -60,6 +66,30 @@ class FlashInferB12xExperts(mk.FusedMoEExpertsModular):
# one. Holding it on the instance keeps apply() alloc-free.
self._fc2_input_scale: torch.Tensor | None = None
# Shape params for B12xMoEWrapper construction.
self.global_num_experts = moe_config.num_experts
self.topk = moe_config.experts_per_token
self.hidden_dim = moe_config.hidden_dim
self.intermediate_size_per_partition = (
moe_config.intermediate_size_per_partition
)
self.max_num_tokens = moe_config.max_num_tokens
self.local_expert_offset = self.ep_rank * self.num_local_experts
activation = moe_config.activation
if activation not in self._ACTIVATION_MAP:
raise ValueError(
f"FlashInferB12xExperts does not support "
f"activation {activation!r}. "
f"Supported: {list(self._ACTIVATION_MAP.keys())}"
)
self._activation_str = self._ACTIVATION_MAP[activation]
# Lazily created on first apply() call.
self._wrapper: Any | None = None
self.w1_sf_mma: torch.Tensor | None = None
self.w2_sf_mma: torch.Tensor | None = None
def process_weights_after_loading(self, layer: torch.nn.Module) -> None:
# Normalise block scales to absorb the per-expert weight global scale
# (w_gs). vLLM's NVFP4 convention stores:
@@ -141,7 +171,7 @@ class FlashInferB12xExperts(mk.FusedMoEExpertsModular):
@staticmethod
def _supports_no_act_and_mul() -> bool:
return False
return True
@staticmethod
def _supports_quant_scheme(
@@ -158,11 +188,13 @@ class FlashInferB12xExperts(mk.FusedMoEExpertsModular):
@staticmethod
def _supports_activation(activation: MoEActivation) -> bool:
return activation == MoEActivation.SILU
return activation in (MoEActivation.SILU, MoEActivation.RELU2_NO_MUL)
@staticmethod
def _supports_parallel_config(moe_parallel_config: FusedMoEParallelConfig) -> bool:
return True
# B12xMoEWrapper does not yet support expert parallelism: its local
# expert count must equal the global expert count.
return not moe_parallel_config.use_ep
def supports_expert_map(self) -> bool:
return False
@@ -190,13 +222,29 @@ class FlashInferB12xExperts(mk.FusedMoEExpertsModular):
@property
def expects_unquantized_inputs(self) -> bool:
# b12x_fused_moe expects BF16 hidden states and performs its own FP4
# B12xMoEWrapper expects BF16 hidden states and performs its own FP4
# quantization internally. Returning True prevents the modular kernel
# from pre-quantizing activations, which would produce an FP4-packed
# tensor with size(-1)=k//2 and break the scale-factor conversion that
# expects size(-1)=k.
# from pre-quantizing activations.
return True
def _ensure_wrapper(self) -> None:
"""Lazily create B12xMoEWrapper on first use."""
if self._wrapper is not None:
return
from flashinfer.fused_moe import B12xMoEWrapper
self._wrapper = B12xMoEWrapper(
num_experts=self.global_num_experts,
top_k=self.topk,
hidden_size=self.hidden_dim,
intermediate_size=self.intermediate_size_per_partition,
use_cuda_graph=True,
max_num_tokens=self.max_num_tokens,
num_local_experts=self.num_local_experts,
activation=self._activation_str,
)
def apply(
self,
output: torch.Tensor,
@@ -224,13 +272,16 @@ class FlashInferB12xExperts(mk.FusedMoEExpertsModular):
assert self._fc2_input_scale is not None, (
"_fc2_input_scale must be set by process_weights_after_loading"
)
assert self.w1_sf_mma is not None and self.w2_sf_mma is not None, (
"process_weights_after_loading must run before FlashInferB12xExperts.apply"
)
top_k = topk_ids.shape[1]
self._ensure_wrapper()
wrapper = self._wrapper
assert wrapper is not None
flashinfer_b12x_fused_moe(
wrapper_output = wrapper.run(
x=hidden_states,
token_selected_experts=topk_ids.to(torch.int32),
token_final_scales=topk_weights,
w1_weight=w1,
w1_weight_sf=self.w1_sf_mma,
w1_alpha=self.g1_alphas,
@@ -238,9 +289,7 @@ class FlashInferB12xExperts(mk.FusedMoEExpertsModular):
w2_weight=w2,
w2_weight_sf=self.w2_sf_mma,
w2_alpha=self.g2_alphas,
num_experts=global_num_experts,
top_k=top_k,
num_local_experts=self.num_local_experts,
output_dtype=self.out_dtype,
output=output,
token_selected_experts=topk_ids.to(torch.int32),
token_final_scales=topk_weights,
)
output.copy_(wrapper_output)