Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
124 changes: 124 additions & 0 deletions tests/quantization/test_quark.py
Original file line number Diff line number Diff line change
Expand Up @@ -82,6 +82,130 @@ def enable_pickle(monkeypatch):
monkeypatch.setenv("VLLM_ALLOW_INSECURE_SERIALIZATION", "1")


def test_quark_nvfp4_moe_forwards_swiglu_params():
from unittest.mock import MagicMock, patch

from vllm.model_executor.layers.fused_moe.activation import MoEActivation
from vllm.model_executor.layers.fused_moe.config import (
FusedMoEConfig,
FusedMoEParallelConfig,
RoutingMethodType,
)
from vllm.model_executor.layers.quantization.quark import quark_moe
from vllm.model_executor.layers.quantization.quark.quark_moe import (
QuarkNvfp4MoEMethod,
)

moe_config = FusedMoEConfig(
num_experts=2,
experts_per_token=1,
hidden_dim=16,
intermediate_size=32,
num_local_experts=2,
num_logical_experts=2,
activation=MoEActivation.SWIGLUOAI,
device="cpu",
routing_method=RoutingMethodType.TopK,
moe_parallel_config=FusedMoEParallelConfig.make_no_parallel(),
in_dtype=torch.bfloat16,
max_num_tokens=8,
intermediate_size_per_partition=32,
)

layer = torch.nn.Module()
layer.w13_weight_scale = torch.ones(2)
layer.w2_weight_scale = torch.ones(2)
layer.w13_weight_scale_2 = torch.ones(2)
layer.w2_weight_scale_2 = torch.ones(2)
layer.w13_input_scale_2 = torch.ones(2)
layer.w2_input_scale_2 = torch.ones(2)
layer.swiglu_limit = 7.0
layer.swiglu_alpha = 1.702
layer.swiglu_beta = 1.0

make_quant_config = MagicMock(return_value=object())
with (
patch.object(
quark_moe,
"select_nvfp4_moe_backend",
return_value=(object(), object()),
),
patch.object(
quark_moe,
"make_nvfp4_moe_quant_config",
make_quant_config,
),
):
method = QuarkNvfp4MoEMethod({}, {}, moe_config, MagicMock())
method.get_fused_moe_quant_config(layer)

_, kwargs = make_quant_config.call_args
assert kwargs["swiglu_limit"] == 7.0
assert kwargs["swiglu_alpha"] == 1.702
assert kwargs["swiglu_beta"] == 1.0


def test_quark_nvfp4_swiglu_limit_allows_emulation_backend(monkeypatch):
from vllm.model_executor.layers.fused_moe.activation import MoEActivation
from vllm.model_executor.layers.fused_moe.config import (
FusedMoEConfig,
FusedMoEParallelConfig,
RoutingMethodType,
)
from vllm.model_executor.layers.fused_moe.oracle import nvfp4
from vllm.model_executor.layers.fused_moe.oracle.nvfp4 import (
NvFp4MoeBackend,
select_nvfp4_moe_backend,
)
from vllm.model_executor.layers.quantization.utils.quant_utils import (
kNvfp4Dynamic,
kNvfp4Static,
)

class UnsupportedExperts:
@staticmethod
def is_supported_config(cls, config, weight_key, activation_key, fmt):
return False, "unsupported"

class SupportedExperts:
@staticmethod
def is_supported_config(cls, config, weight_key, activation_key, fmt):
return True, None

def backend_to_kernel_cls(backend):
if backend == NvFp4MoeBackend.EMULATION:
return [SupportedExperts]
return [UnsupportedExperts]

config = FusedMoEConfig(
num_experts=2,
experts_per_token=1,
hidden_dim=16,
intermediate_size=32,
num_local_experts=2,
num_logical_experts=2,
activation=MoEActivation.SWIGLUOAI,
device="cpu",
routing_method=RoutingMethodType.TopK,
moe_parallel_config=FusedMoEParallelConfig.make_no_parallel(),
in_dtype=torch.bfloat16,
max_num_tokens=8,
swiglu_limit=7.0,
intermediate_size_per_partition=32,
)

monkeypatch.setattr(nvfp4.envs, "is_set", lambda _: False)
monkeypatch.setattr(nvfp4.envs, "VLLM_TEST_FORCE_FP8_MARLIN", False)
monkeypatch.setattr(nvfp4, "backend_to_kernel_cls", backend_to_kernel_cls)

backend, experts_cls = select_nvfp4_moe_backend(
config=config,
weight_key=kNvfp4Static,
activation_key=kNvfp4Dynamic,
)

assert backend == NvFp4MoeBackend.EMULATION
assert experts_cls is SupportedExperts
def test_quark_config_has_no_model_specific_fused_mappings():
config = QuarkConfig({})

Expand Down
4 changes: 4 additions & 0 deletions vllm/model_executor/layers/fused_moe/config.py
Original file line number Diff line number Diff line change
Expand Up @@ -817,6 +817,8 @@
gemm1_alpha: float | None = None,
gemm1_beta: float | None = None,
gemm1_clamp_limit: float | None = None,
gemm1_alpha: float | None = None,

Check failure on line 820 in vllm/model_executor/layers/fused_moe/config.py

View workflow job for this annotation

GitHub Actions / pre-commit

Duplicate parameter "gemm1_alpha" in function definition

Check failure on line 820 in vllm/model_executor/layers/fused_moe/config.py

View workflow job for this annotation

GitHub Actions / pre-commit

Duplicate parameter "gemm1_alpha" in function definition

Check failure on line 820 in vllm/model_executor/layers/fused_moe/config.py

View workflow job for this annotation

GitHub Actions / pre-commit

Duplicate parameter "gemm1_alpha" in function definition

Check failure on line 820 in vllm/model_executor/layers/fused_moe/config.py

View workflow job for this annotation

GitHub Actions / pre-commit

Duplicate parameter "gemm1_alpha" in function definition

Check failure on line 820 in vllm/model_executor/layers/fused_moe/config.py

View workflow job for this annotation

GitHub Actions / pre-commit

Duplicate parameter "gemm1_alpha" in function definition

Check failure on line 820 in vllm/model_executor/layers/fused_moe/config.py

View workflow job for this annotation

GitHub Actions / pre-commit

Duplicate parameter "gemm1_alpha" in function definition
gemm1_beta: float | None = None,
) -> FusedMoEQuantConfig:
"""
Construct a quant config for mxfp4 activations and nvp4 weights.
Expand All @@ -838,6 +840,8 @@
gemm1_alpha=gemm1_alpha,
gemm1_beta=gemm1_beta,
gemm1_clamp_limit=gemm1_clamp_limit,
gemm1_alpha=gemm1_alpha,

Check failure on line 843 in vllm/model_executor/layers/fused_moe/config.py

View workflow job for this annotation

GitHub Actions / pre-commit

Ruff (invalid-syntax)

vllm/model_executor/layers/fused_moe/config.py:843:9: invalid-syntax: Duplicate keyword argument "gemm1_alpha"
gemm1_beta=gemm1_beta,

Check failure on line 844 in vllm/model_executor/layers/fused_moe/config.py

View workflow job for this annotation

GitHub Actions / pre-commit

Ruff (invalid-syntax)

vllm/model_executor/layers/fused_moe/config.py:844:9: invalid-syntax: Duplicate keyword argument "gemm1_beta"
)


Expand Down
2 changes: 2 additions & 0 deletions vllm/model_executor/layers/fused_moe/oracle/nvfp4.py
Original file line number Diff line number Diff line change
Expand Up @@ -559,6 +559,8 @@
gemm1_alpha=swiglu_alpha,
gemm1_beta=swiglu_beta,
gemm1_clamp_limit=swiglu_limit,
gemm1_alpha=swiglu_alpha,

Check failure on line 562 in vllm/model_executor/layers/fused_moe/oracle/nvfp4.py

View workflow job for this annotation

GitHub Actions / pre-commit

Ruff (invalid-syntax)

vllm/model_executor/layers/fused_moe/oracle/nvfp4.py:562:13: invalid-syntax: Duplicate keyword argument "gemm1_alpha"

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

duplicated from L503 above

gemm1_beta=swiglu_beta,

Check failure on line 563 in vllm/model_executor/layers/fused_moe/oracle/nvfp4.py

View workflow job for this annotation

GitHub Actions / pre-commit

Ruff (invalid-syntax)

vllm/model_executor/layers/fused_moe/oracle/nvfp4.py:563:13: invalid-syntax: Duplicate keyword argument "gemm1_beta"
)

if backend == NvFp4MoeBackend.FLASHINFER_CUTEDSL:
Expand Down
Loading