Skip to content
Open
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
50 changes: 50 additions & 0 deletions tests/quantization/test_online.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,9 @@
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
"""Tests online quantization."""

from types import SimpleNamespace
from unittest.mock import Mock

import pytest
import torch

Expand All @@ -23,6 +26,53 @@
from vllm.utils.flashinfer import has_flashinfer_trtllm_fused_moe


def test_online_nvfp4_reuses_kernel_when_weights_are_reprocessed(
monkeypatch,
) -> None:
method = object.__new__(Nvfp4OnlineMoEMethod)
method.moe = SimpleNamespace(is_act_and_mul=True)
method.nvfp4_backend = object()
method.experts_cls = object
method.moe_quant_config = None
method.moe_kernel = None

layer = Mock()
converted_weights = tuple(object() for _ in range(8))
convert_weights = Mock(return_value=converted_weights)
process_weights = Mock()
kernel = SimpleNamespace(
fused_experts=SimpleNamespace(
process_weights_after_loading=process_weights,
)
)
make_kernel = Mock(return_value=kernel)
get_quant_config = Mock(return_value=object())
method.get_fused_moe_quant_config = get_quant_config

monkeypatch.setattr(
"vllm.model_executor.layers.quantization.online.nvfp4."
"convert_to_nvfp4_moe_kernel_format",
convert_weights,
)
monkeypatch.setattr(
"vllm.model_executor.layers.quantization.online.nvfp4.replace_parameter",
Mock(),
)
monkeypatch.setattr(
"vllm.model_executor.layers.quantization.online.nvfp4.make_nvfp4_moe_kernel",
make_kernel,
)

method._setup_kernel(layer)
method._setup_kernel(layer)

assert method.moe_kernel is kernel
assert convert_weights.call_count == 2
make_kernel.assert_called_once()
get_quant_config.assert_called_once()
assert process_weights.call_count == 2


@pytest.mark.skipif(
not is_quant_method_supported("fp8"),
reason="FP8 is not supported on this GPU type.",
Expand Down
24 changes: 13 additions & 11 deletions vllm/model_executor/layers/quantization/online/nvfp4.py
Original file line number Diff line number Diff line change
Expand Up @@ -144,17 +144,19 @@ def _setup_kernel(self, layer: RoutedExperts) -> None:
replace_parameter(layer, "w2_weight_scale_2", w2_scale_2)
replace_parameter(layer, "w2_input_scale", a2_scale)

self.moe_quant_config = self.get_fused_moe_quant_config(layer)
assert self.experts_cls is not None
self.moe_kernel = make_nvfp4_moe_kernel(
moe_quant_config=self.moe_quant_config,
moe_config=self.moe,
experts_cls=self.experts_cls,
backend=self.nvfp4_backend,
routing_tables=layer._expert_routing_tables(),
layer=layer,
per_token_activation=True,
)
if self.moe_kernel is None:
self.moe_quant_config = self.get_fused_moe_quant_config(layer)
assert self.experts_cls is not None
self.moe_kernel = make_nvfp4_moe_kernel(
moe_quant_config=self.moe_quant_config,
moe_config=self.moe,
experts_cls=self.experts_cls,
backend=self.nvfp4_backend,
routing_tables=layer._expert_routing_tables(),
layer=layer,
per_token_activation=True,
)

self.moe_kernel.fused_experts.process_weights_after_loading(layer)

def get_fused_moe_quant_config(self, layer: torch.nn.Module) -> FusedMoEQuantConfig:
Expand Down
Loading