mirror of
https://github.com/vllm-project/vllm.git
synced 2026-08-17 11:10:16 +00:00
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:
co-authored by
Andrii Skliar
parent
b136cc2c2c
commit
b1384f5ec6
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user