diff --git a/tests/evals/gsm8k/configs/Qwen3-1.7B-MXFP4.yaml b/tests/evals/gsm8k/configs/Qwen3-1.7B-MXFP4.yaml new file mode 100644 index 00000000000..b0249f41b60 --- /dev/null +++ b/tests/evals/gsm8k/configs/Qwen3-1.7B-MXFP4.yaml @@ -0,0 +1,5 @@ +model_name: "amd-quark/Qwen3-1.7B-MXFP4" +accuracy_threshold: 0.27 +num_questions: 1319 +num_fewshot: 5 +server_args: "--enforce-eager --max-model-len 4096" diff --git a/tests/kernels/quantization/test_mxfp4_kernel_selection.py b/tests/kernels/quantization/test_mxfp4_kernel_selection.py new file mode 100644 index 00000000000..20c171a87ce --- /dev/null +++ b/tests/kernels/quantization/test_mxfp4_kernel_selection.py @@ -0,0 +1,130 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""Tests for MXFP4 linear kernel selection logic (CPU-only) + +Run `pytest tests/kernels/quantization/test_mxfp4_kernel_selection.py`. +""" + +from unittest.mock import patch + +import pytest +import torch + +from vllm.model_executor.kernels.linear import ( + AiterMxfp4LinearKernel, + MxFp4LinearKernel, + MxFp4LinearLayerConfig, + init_mxfp4_linear_kernel, + register_linear_kernel, +) +from vllm.platforms import PlatformEnum + +pytestmark = pytest.mark.cpu_test + + +def test_can_implement_is_abstract(): + """Test that can_implement()/is_supported() are properly defined.""" + assert hasattr(MxFp4LinearKernel, "can_implement") + assert hasattr(MxFp4LinearKernel, "is_supported") + + +def test_aiter_kernel_is_supported_requires_native_mx_support(): + """AiterMxfp4LinearKernel must not be selected on platforms without + native MX compute, even if AITER itself is importable.""" + with patch( + "vllm.model_executor.kernels.linear.mxfp4.aiter.current_platform.supports_mx", + return_value=False, + ): + is_supported, reason = AiterMxfp4LinearKernel.is_supported() + assert not is_supported + assert reason + + +class OOTMxFp4LinearKernel(MxFp4LinearKernel): + @classmethod + def is_supported( + cls, compute_capability: int | None = None + ) -> tuple[bool, str | None]: + return True, None + + @classmethod + def can_implement(cls, config: MxFp4LinearLayerConfig) -> tuple[bool, str | None]: + return True, None + + def process_weights_after_loading(self, layer: torch.nn.Module) -> None: + pass + + def apply_weights( + self, + layer: torch.nn.Module, + x: torch.Tensor, + bias: torch.Tensor | None = None, + ) -> torch.Tensor: + pass + + +@patch("vllm.model_executor.kernels.linear.current_platform") +def test_init_mxfp4_linear_kernel_dispatches_to_registered_kernel(platform_mock): + """init_mxfp4_linear_kernel should select a registered kernel that + reports itself as supported, and construct it with a fresh config.""" + platform_mock._enum = PlatformEnum.OOT + register_linear_kernel(OOTMxFp4LinearKernel, PlatformEnum.OOT, "mxfp4") + + kernel = init_mxfp4_linear_kernel() + + assert isinstance(kernel, OOTMxFp4LinearKernel) + assert kernel.config == MxFp4LinearLayerConfig() + + +class UnsupportedMxFp4LinearKernel(MxFp4LinearKernel): + @classmethod + def is_supported( + cls, compute_capability: int | None = None + ) -> tuple[bool, str | None]: + return False, "never supported" + + @classmethod + def can_implement(cls, config: MxFp4LinearLayerConfig) -> tuple[bool, str | None]: + return True, None + + def process_weights_after_loading(self, layer: torch.nn.Module) -> None: + pass + + def apply_weights( + self, + layer: torch.nn.Module, + x: torch.Tensor, + bias: torch.Tensor | None = None, + ) -> torch.Tensor: + pass + + +@patch("vllm.model_executor.kernels.linear.current_platform") +def test_init_mxfp4_linear_kernel_raises_when_no_kernel_matches(platform_mock): + platform_mock._enum = PlatformEnum.UNSPECIFIED + register_linear_kernel( + UnsupportedMxFp4LinearKernel, PlatformEnum.UNSPECIFIED, "mxfp4" + ) + + with pytest.raises(ValueError, match="Failed to find a kernel"): + init_mxfp4_linear_kernel() + + +@patch("vllm.model_executor.kernels.linear.mxfp4.aiter.is_aiter_found_and_supported") +@patch("vllm.model_executor.kernels.linear.mxfp4.aiter.current_platform") +@patch("vllm.model_executor.kernels.linear.current_platform") +def test_init_mxfp4_linear_kernel_raises_on_rocm_without_aiter( + linear_platform_mock, aiter_platform_mock, is_aiter_found_and_supported_mock +): + """On ROCm, the only registered MXFP4 linear kernel is AITER-based. + If AITER is not found/supported, no kernel should be selected.""" + linear_platform_mock._enum = PlatformEnum.ROCM + aiter_platform_mock.supports_mx.return_value = True + is_aiter_found_and_supported_mock.return_value = False + + with pytest.raises( + ValueError, + match="(?s)Failed to find a kernel.*" + "AITER not found or not supported on the current platform", + ): + init_mxfp4_linear_kernel() diff --git a/tests/quantization/test_quark.py b/tests/quantization/test_quark.py index d2f9dcb8af1..8440db10719 100644 --- a/tests/quantization/test_quark.py +++ b/tests/quantization/test_quark.py @@ -17,6 +17,7 @@ import pytest import torch from packaging import version +from vllm._aiter_ops import is_aiter_found_and_supported from vllm.model_executor.layers.quantization.quark.quark import ( # noqa: E501 QuarkLinearMethod, QuarkW8A8Fp8, @@ -26,6 +27,9 @@ from vllm.model_executor.layers.quantization.quark.quark_moe import ( # noqa: E QuarkW4A8Fp8MoEMethod, QuarkW8A8Int8MoEMethod, ) +from vllm.model_executor.layers.quantization.utils.mxfp4_utils import ( + quant_dequant_mxfp4, +) from vllm.model_executor.layers.quantization.utils.quant_utils import ( is_layer_skipped, ) @@ -52,6 +56,8 @@ QUARK_MXFP4_AVAILABLE = find_spec("quark") is not None and version.parse( importlib.metadata.version("amd-quark") ) >= version.parse(QUARK_MXFP4_MIN_VERSION) +AITER_AVAILABLE = is_aiter_found_and_supported() + DEVICE_TYPE = current_platform.device_type if QUARK_MXFP4_AVAILABLE: @@ -487,6 +493,42 @@ def test_mxfp4_dequant_kernel_match_quark( assert torch.equal(out_hip, out_torch) +@pytest.mark.skipif( + not QUARK_MXFP4_AVAILABLE, + reason=f"amd-quark>={QUARK_MXFP4_MIN_VERSION} is not available", +) +@pytest.mark.skipif( + not AITER_AVAILABLE, + reason="AITER is not found or not supported on the current platform", +) +@pytest.mark.parametrize("float_dtype", [torch.bfloat16, torch.float16]) +@pytest.mark.parametrize("scalings", [[2.3, 0.03, 7.3, 0.1, 0.004, 17.3, 1e4, 1e-4]]) +def test_mxfp4_dynamic_quant_match_quark( + float_dtype: torch.dtype, scalings: list[float] +): + """`AiterMxfp4LinearKernel` quantizes weights dynamically through AITER's + `dynamic_mxfp4_quant`, while the emulation path quantizes/dequantizes + through Quark's `qdq_mxfp4`. Check that both agree on the same input. + """ + from aiter.ops.triton.quant import dynamic_mxfp4_quant + + torch.manual_seed(0) + + hidden_size = 32 * 64 + inp = (torch.rand(48, hidden_size, dtype=float_dtype, device=DEVICE_TYPE) - 0.5) * 2 + for i in range(hidden_size // 32): + inp[:, i * 32 : (i + 1) * 32] = ( + inp[:, i * 32 : (i + 1) * 32] * scalings[i % len(scalings)] + ) + + x_q, x_s = dynamic_mxfp4_quant(inp) + out_dynamic_quant = dq_mxfp4_torch(x_q, x_s, float_dtype) + + out_quark_qdq = quant_dequant_mxfp4(inp) + + assert torch.equal(out_dynamic_quant, out_quark_qdq) + + # Unit tests for ``is_layer_skipped`` fused-name handling. FUSED_MAPPING = { diff --git a/vllm/model_executor/kernels/linear/__init__.py b/vllm/model_executor/kernels/linear/__init__.py index fcc50ffcb0e..39ed4309cad 100644 --- a/vllm/model_executor/kernels/linear/__init__.py +++ b/vllm/model_executor/kernels/linear/__init__.py @@ -74,6 +74,9 @@ from vllm.model_executor.kernels.linear.mxfp4 import ( MxFp4LinearKernel, MxFp4LinearLayerConfig, ) +from vllm.model_executor.kernels.linear.mxfp4.aiter import ( + AiterMxfp4LinearKernel, +) from vllm.model_executor.kernels.linear.mxfp4.flashinfer import ( FlashInferMxFp4LinearKernel, ) @@ -274,6 +277,7 @@ _LINEAR_BACKEND_KERNEL_MAP: dict[str, set[type]] = { AiterFp8BlockScaledMMKernel, AiterPerTokenFp8ScaledMMLinearKernel, AiterPreshuffledPerTokenFp8ScaledMMLinearKernel, + AiterMxfp4LinearKernel, }, "machete": { MacheteLinearKernel, @@ -469,6 +473,9 @@ _POSSIBLE_MXFP4_KERNELS: dict[PlatformEnum, list[type[MxFp4LinearKernel]]] = { MarlinMxFp4LinearKernel, HummingMxFp4LinearKernel, ], + PlatformEnum.ROCM: [ + AiterMxfp4LinearKernel, + ], PlatformEnum.XPU: [ XPUMxFp4LinearKernel, ], @@ -1079,6 +1086,7 @@ __all__ = [ "init_mxfp4_linear_kernel", "MxFp4LinearKernel", "MxFp4LinearLayerConfig", + "AiterMxfp4LinearKernel", "FlashInferMxFp4LinearKernel", "MarlinMxFp4LinearKernel", "FlashInferCutedslMxfp8LinearKernel", diff --git a/vllm/model_executor/kernels/linear/mxfp4/aiter.py b/vllm/model_executor/kernels/linear/mxfp4/aiter.py new file mode 100644 index 00000000000..54f8098b9a5 --- /dev/null +++ b/vllm/model_executor/kernels/linear/mxfp4/aiter.py @@ -0,0 +1,177 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project + +import torch +from torch.nn.parameter import Parameter + +from vllm._aiter_ops import is_aiter_found_and_supported, rocm_aiter_ops +from vllm.platforms import current_platform + +from .base import MxFp4LinearKernel, MxFp4LinearLayerConfig + +# NOTE: Do not import aiter at module scope. Importing aiter eagerly initializes HIP +# which can force the engine core to spawn instead of fork. +# is_aiter_found_and_supported() checks platform + arch + library availability via +# find_spec/amdsmi, so it stays HIP-free. +# Actual aiter imports are deferred to the functions/methods that need them, +# where HIP initialization is expected. +if is_aiter_found_and_supported(): + from vllm.utils.torch_utils import direct_register_custom_op + + def gemm_with_dynamic_quant( + x: torch.Tensor, + weight: torch.Tensor, + weight_scale: torch.Tensor, + rocm_use_aiter_fp4_asm_gemm: bool = False, + out_dtype: torch.dtype | None = torch.bfloat16, + x_scales: torch.Tensor | None = None, + ) -> torch.Tensor: + from aiter.ops.triton.gemm_afp4wfp4 import ( + gemm_afp4wfp4, + gemm_afp4wfp4_preshuffled_weight_scales, + ) + from aiter.ops.triton.quant import dynamic_mxfp4_quant + + if rocm_use_aiter_fp4_asm_gemm: + from aiter import gemm_a4w4, per_1x32_f4_quant_hip + + M = x.shape[0] + N = weight.shape[0] + K = weight.shape[1] + if rocm_use_aiter_fp4_asm_gemm: + if M <= 64 and rocm_aiter_ops.is_triton_gemm_afp4wfp4_presh_ws_tuned(N, K): + if x_scales is None: + # use hip quant kernel for performance + if M >= 32: + x_q, x_s = per_1x32_f4_quant_hip(x, shuffle=True) + else: + x_q, x_s = per_1x32_f4_quant_hip(x, shuffle=False) + else: + x_q = x + x_s = x_scales + + if M >= 32: + x_s = x_s.view(torch.uint8).view(x_s.shape[0] // 32, -1) + else: + x_s = x_s[:M, ...].view(torch.uint8) + + y = torch.empty(M, N, device=x_q.device, dtype=out_dtype) + gemm_afp4wfp4_preshuffled_weight_scales( + x_q.view(torch.uint8), + weight.view(torch.uint8).view(weight.shape[0] // 16, -1), + x_s, + weight_scale.view(torch.uint8).view( + weight_scale.shape[0] // 32, -1 + ), + out_dtype, + y, + ) + else: + if x_scales is None: + # use hip quant kernel for performance + x_q, x_s = per_1x32_f4_quant_hip(x, shuffle=True) + else: + x_q = x + x_s = x_scales + + y = gemm_a4w4( + x_q, + weight.view(x_q.dtype), + x_s, + weight_scale.view(x_s.dtype), + dtype=out_dtype, + bpreshuffle=True, + ) + return y[:M] + else: + if x_scales is None: + x_q, x_s = dynamic_mxfp4_quant(x) + else: + x_q = x + x_s = x_scales + y = torch.empty( + x_q.shape[0], weight.shape[0], device=x_q.device, dtype=out_dtype + ) + + gemm_afp4wfp4(x_q, weight, x_s, weight_scale.T, out_dtype, y) + return y + + def gemm_with_dynamic_quant_fake( + x: torch.Tensor, + weight: torch.Tensor, + weight_scale: torch.Tensor, + x_scales: torch.Tensor = None, + rocm_use_aiter_fp4_asm_gemm: bool = False, + out_dtype: torch.dtype | None = torch.bfloat16, + ) -> torch.Tensor: + return torch.empty( + (*x.shape[:-1], weight.shape[0]), dtype=out_dtype, device=x.device + ) + + direct_register_custom_op( + op_name="gemm_with_dynamic_quant", + op_func=gemm_with_dynamic_quant, + mutates_args=[], + fake_impl=gemm_with_dynamic_quant_fake, + dispatch_key=current_platform.dispatch_key, + ) + + +class AiterMxfp4LinearKernel(MxFp4LinearKernel): + """AITER-based native MXFP4 GEMM kernel for ROCm.""" + + def __init__(self, config: MxFp4LinearLayerConfig) -> None: + super().__init__(config) + self.use_asm_gemm = rocm_aiter_ops.is_asm_fp4_gemm_dynamic_quant_enabled() + self.out_dtype = torch.get_default_dtype() + + @classmethod + def is_supported( + cls, compute_capability: int | None = None + ) -> tuple[bool, str | None]: + if not current_platform.supports_mx(): + return False, "current platform does not support native MXFP4 computation" + if is_aiter_found_and_supported(): + return True, None + return False, "AITER not found or not supported on the current platform" + + @classmethod + def can_implement(cls, c: MxFp4LinearLayerConfig) -> tuple[bool, str | None]: + return True, None + + def process_weights_after_loading(self, layer: torch.nn.Module) -> None: + if self.use_asm_gemm: + from aiter.ops.shuffle import shuffle_weight + + weight_scale = layer.weight_scale.data + sm, sn = weight_scale.shape + weight_scale = weight_scale.view(sm // 32, 2, 16, sn // 8, 2, 4, 1) + weight_scale = weight_scale.permute(0, 3, 5, 2, 4, 1, 6).contiguous() + weight_scale = weight_scale.view(sm, sn) + layer.weight_scale = Parameter(weight_scale, requires_grad=False) + + layer.weight = Parameter( + shuffle_weight(layer.weight.data, layout=(16, 16)), + requires_grad=False, + ) + else: + layer.weight_scale = Parameter( + layer.weight_scale.data.T.contiguous(), requires_grad=False + ) + + def apply_weights( + self, + layer: torch.nn.Module, + x: torch.Tensor, + bias: torch.Tensor | None = None, + ) -> torch.Tensor: + y = torch.ops.vllm.gemm_with_dynamic_quant( + x, + layer.weight, + layer.weight_scale, + self.use_asm_gemm, + self.out_dtype, + ) + if bias is not None: + y = y + bias + return y diff --git a/vllm/model_executor/layers/quantization/quark/schemes/quark_ocp_mx.py b/vllm/model_executor/layers/quantization/quark/schemes/quark_ocp_mx.py index ea63d1bcf12..2d42babeb38 100644 --- a/vllm/model_executor/layers/quantization/quark/schemes/quark_ocp_mx.py +++ b/vllm/model_executor/layers/quantization/quark/schemes/quark_ocp_mx.py @@ -9,8 +9,8 @@ from typing import Any import torch import torch.nn.functional as F -from vllm._aiter_ops import is_aiter_found_and_supported, rocm_aiter_ops from vllm.logger import init_logger +from vllm.model_executor.kernels.linear import init_mxfp4_linear_kernel from vllm.model_executor.layers.quantization.utils.mxfp4_utils import ( dequant_mxfp4, quant_dequant_mxfp4, @@ -36,120 +36,6 @@ from .quark_scheme import QuarkScheme logger = init_logger(__name__) -# NOTE: Do not import aiter at module scope. Importing aiter eagerly initializes HIP -# which can force the engine core to spawn instead of fork. -# is_aiter_found_and_supported() checks platform + arch + library availability via -# find_spec/amdsmi, so it stays HIP-free. -# Actual aiter imports are deferred to the functions/methods that need them, -# where HIP initialization is expected. -if is_aiter_found_and_supported(): - from vllm.utils.torch_utils import direct_register_custom_op - - def gemm_with_dynamic_quant( - x: torch.Tensor, - weight: torch.Tensor, - weight_scale: torch.Tensor, - rocm_use_aiter_fp4_asm_gemm: bool = False, - out_dtype: torch.dtype | None = torch.bfloat16, - x_scales: torch.Tensor | None = None, - ) -> torch.Tensor: - from aiter.ops.triton.gemm_afp4wfp4 import ( - gemm_afp4wfp4, - gemm_afp4wfp4_preshuffled_weight_scales, - ) - from aiter.ops.triton.quant import dynamic_mxfp4_quant - - if rocm_use_aiter_fp4_asm_gemm: - from aiter import gemm_a4w4, per_1x32_f4_quant_hip - - M = x.shape[0] - N = weight.shape[0] - K = weight.shape[1] - if rocm_use_aiter_fp4_asm_gemm: - if M <= 64 and rocm_aiter_ops.is_triton_gemm_afp4wfp4_presh_ws_tuned(N, K): - if x_scales is None: - # use hip quant kernel for performance - if M >= 32: - x_q, x_s = per_1x32_f4_quant_hip(x, shuffle=True) - else: - x_q, x_s = per_1x32_f4_quant_hip(x, shuffle=False) - else: - x_q = x - x_s = x_scales - - if M >= 32: - x_s = x_s.view(torch.uint8).view(x_s.shape[0] // 32, -1) - else: - x_s = x_s[:M, ...].view(torch.uint8) - - y = torch.empty(M, N, device=x_q.device, dtype=out_dtype) - gemm_afp4wfp4_preshuffled_weight_scales( - x_q.view(torch.uint8), - weight.view(torch.uint8).view(weight.shape[0] // 16, -1), - x_s, - weight_scale.view(torch.uint8).view( - weight_scale.shape[0] // 32, -1 - ), - out_dtype, - y, - ) - else: - if x_scales is None: - # use hip quant kernel for performance - x_q, x_s = per_1x32_f4_quant_hip(x, shuffle=True) - else: - x_q = x - x_s = x_scales - - y = gemm_a4w4( - x_q, - weight.view(x_q.dtype), - x_s, - weight_scale.view(x_s.dtype), - dtype=out_dtype, - bpreshuffle=True, - ) - return y[:M] - else: - if x_scales is None: - x_q, x_s = dynamic_mxfp4_quant(x) - else: - x_q = x - x_s = x_scales - y = torch.empty( - x_q.shape[0], weight.shape[0], device=x_q.device, dtype=out_dtype - ) - - gemm_afp4wfp4(x_q, weight, x_s, weight_scale.T, out_dtype, y) - return y - - def gemm_with_dynamic_quant_fake( - x: torch.Tensor, - weight: torch.Tensor, - weight_scale: torch.Tensor, - x_scales: torch.Tensor = None, - rocm_use_aiter_fp4_asm_gemm: bool = False, - out_dtype: torch.dtype | None = torch.bfloat16, - ) -> torch.Tensor: - return torch.empty( - (*x.shape[:-1], weight.shape[0]), dtype=out_dtype, device=x.device - ) - - direct_register_custom_op( - op_name="gemm_with_dynamic_quant", - op_func=gemm_with_dynamic_quant, - mutates_args=[], - fake_impl=gemm_with_dynamic_quant_fake, - dispatch_key=current_platform.dispatch_key, - ) -elif current_platform.is_rocm(): - logger.warning( - "AITER is not found or not supported on the current platform, " - "QuarkOCP_MX will fall back to emulation." - "Native MXFP4/MXFP6 acceleration will not be available." - ) - - class QuarkOCP_MX(QuarkScheme): def __init__( self, @@ -157,8 +43,6 @@ class QuarkOCP_MX(QuarkScheme): input_quant_spec: dict[str, Any] | None, dynamic_mxfp4_quant: bool = False, ): - self.out_dtype = torch.get_default_dtype() - self.qscheme = "per_group" self.weight_quant_spec = weight_quant_spec self.input_quant_spec = input_quant_spec self.dynamic_mxfp4_quant = dynamic_mxfp4_quant @@ -211,17 +95,10 @@ class QuarkOCP_MX(QuarkScheme): self.input_dtype != "mxfp4" or self.weight_dtype != "mxfp4" ) - self.rocm_use_aiter_fp4_asm_gemm = ( - rocm_aiter_ops.is_asm_fp4_gemm_dynamic_quant_enabled() - ) - - if not self.emulate and not is_aiter_found_and_supported(): - # Currently need AITER kernels if not emulating - raise NotImplementedError( - f"{self.__class__.__name__} requires AITER to be installed " - "for non-emulation mode! Please refer to " - "https://github.com/ROCm/aiter for installation details." - ) + # TODO: Move emulation code path as a kernel, and always + # use init_mxfp4_linear_kernel. + if not self.emulate: + self.ocp_mx_linear = init_mxfp4_linear_kernel() if not current_platform.supports_mx(): logger.warning_once( @@ -268,7 +145,7 @@ class QuarkOCP_MX(QuarkScheme): from aiter.ops.triton.quant import dynamic_mxfp4_quant w_q, w_s = dynamic_mxfp4_quant(layer.weight) - layer.weight_scale = torch.nn.Parameter(w_s.T.contiguous(), requires_grad=False) + layer.weight_scale = torch.nn.Parameter(w_s, requires_grad=False) layer.weight = torch.nn.Parameter(w_q, requires_grad=False) def process_weights_after_loading(self, layer: torch.nn.Module) -> None: @@ -284,31 +161,7 @@ class QuarkOCP_MX(QuarkScheme): else: if self.dynamic_mxfp4_quant: self.process_dynamic_mxfp4_weights_after_loading(layer) - elif self.rocm_use_aiter_fp4_asm_gemm: - from aiter.ops.shuffle import shuffle_weight - - # shuffle weight scale - weight_scale_shuffle = layer.weight_scale.data - sm, sn = weight_scale_shuffle.shape - weight_scale_shuffle = weight_scale_shuffle.view( - sm // 32, 2, 16, sn // 8, 2, 4, 1 - ) - weight_scale_shuffle = weight_scale_shuffle.permute( - 0, 3, 5, 2, 4, 1, 6 - ).contiguous() - weight_scale_shuffle = weight_scale_shuffle.view(sm, sn) - layer.weight_scale = torch.nn.Parameter( - weight_scale_shuffle, requires_grad=False - ) - - # shuffle weight - weight_shuffle = layer.weight.data - weight_shuffle = shuffle_weight(weight_shuffle, layout=(16, 16)) - layer.weight = torch.nn.Parameter(weight_shuffle, requires_grad=False) - else: - layer.weight_scale = torch.nn.Parameter( - layer.weight_scale.data.T.contiguous(), requires_grad=False - ) + self.ocp_mx_linear.process_weights_after_loading(layer) def create_weights( self, @@ -375,15 +228,4 @@ class QuarkOCP_MX(QuarkScheme): dq_w = self.dequant_func(layer.weight, layer.weight_scale, x.dtype) qdq_x = self.quant_dequant_func(x) return F.linear(qdq_x, dq_w, bias) - y = torch.ops.vllm.gemm_with_dynamic_quant( - x, - layer.weight, - layer.weight_scale, - self.rocm_use_aiter_fp4_asm_gemm, - self.out_dtype, - ) - # gemm_with_dynamic_quant has no bias argument; add it here so the - # native path matches F.linear (e.g. qkv_proj with qkv_bias=True). - if bias is not None: - y = y + bias - return y + return self.ocp_mx_linear.apply_weights(layer, x, bias)