# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
"""Test model set-up and weight loading for quark-quantized models.

Run `pytest tests/quantization/test_quark.py`.

See also `tests/kernels/moe/test_ocp_mx_moe.py`.
"""

import importlib.metadata
from dataclasses import dataclass
from importlib.util import find_spec
from types import SimpleNamespace
from unittest.mock import MagicMock, Mock, patch

import huggingface_hub
import lm_eval
import pytest
import torch
from packaging import version

from tests.quantization.utils import load_model_without_vllm_runner
from vllm._aiter_ops import is_aiter_found_and_supported, rocm_aiter_ops
from vllm.config import VllmConfig, set_current_vllm_config
from vllm.config.cache import CacheConfig, CacheDType
from vllm.forward_context import set_forward_context
from vllm.model_executor import parameter
from vllm.model_executor.kernels.linear.scaled_mm.aiter import (
    AiterHipbMMPerTokenFp8ScaledMMLinearKernel,
    AiterPerTokenFp8ScaledMMLinearKernel,
    AiterPreshuffledPerTokenFp8ScaledMMLinearKernel,
)
from vllm.model_executor.kernels.linear.scaled_mm.ScaledMMLinearKernel import (
    FP8ScaledMMLinearLayerConfig,
)
from vllm.model_executor.layers.attention import Attention
from vllm.model_executor.layers.fused_moe import (
    FusedMoeWeightScaleSupported,
    RoutedExperts,
    UnquantizedFusedMoEMethod,
)
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.linear import LinearBase, UnquantizedLinearMethod
from vllm.model_executor.layers.quantization.quark.quark import (  # noqa: E501
    QuarkConfig,
    QuarkLinearMethod,
    QuarkNVFP4,
    QuarkOCP_MX,
    QuarkW8A8Fp8,
    QuarkW8A8Fp8PerBlock,
    QuarkW8A8Int8,
)
from vllm.model_executor.layers.quantization.quark.quark_moe import (  # noqa: E501
    QuarkMoEMethod,
    QuarkW4A8Fp8MoEMethod,
    QuarkW4A16Int4MoEMethod,
    QuarkW8A8Fp8MoEMethod,
    QuarkW8A8Int8MoEMethod,
)
from vllm.model_executor.layers.quantization.quark.schemes import (
    QuarkScheme,
    QuarkW4A16Int4,
)
from vllm.model_executor.layers.quantization.quark.utils import (
    QuarkQTensorHint,
    canonicalize_quark_packed_int4,
    should_ignore_layer,
)
from vllm.model_executor.layers.quantization.utils.mxfp4_utils import (
    quant_dequant_mxfp4,
)
from vllm.model_executor.layers.quantization.utils.quant_utils import (
    QuantKey,
    is_layer_skipped,
    kFp8Dynamic128Sym,
    kFp8DynamicTensorSym,
    kFp8DynamicTokenSym,
    kFp8Static128BlockE8M0Sym,
    kFp8Static128BlockSym,
    kFp8StaticChannelSym,
    kFp8StaticTensorSym,
    kInt4W4A8StaticChannelSym,
    kInt8DynamicTensorAsym,
    kInt8DynamicTensorSym,
    kInt8DynamicTokenAsym,
    kInt8DynamicTokenSym,
    kInt8StaticChannelSym,
    kInt8StaticTensorAsym,
    kInt8StaticTensorSym,
    kMxfp4Dynamic,
    kMxfp4Static,
    kMxfp6E2M3Dynamic,
    kMxfp6E2M3Static,
    kMxfp6E3M2Dynamic,
    kMxfp6E3M2Static,
    kNvfp4Dynamic,
    kNvfp4Static,
)
from vllm.model_executor.models.llama import LlamaForCausalLM
from vllm.model_executor.models.utils import WeightsMapper
from vllm.platforms import current_platform
from vllm.transformers_utils.repo_utils import hf_api

if current_platform.is_rocm():
    from vllm.platforms.rocm import on_gfx942, on_gfx950
else:

    def on_gfx942() -> bool:
        return False

    def on_gfx950() -> bool:
        return False


from .reference_mxfp4 import dq_mxfp4_torch, qdq_mxfp4_torch

# Minimum amd-quark version for MXFP4/OCP_MX tests (single source of truth).
QUARK_MXFP4_MIN_VERSION = "0.12"

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()

AITER_PTPC_KERNELS = (
    AiterHipbMMPerTokenFp8ScaledMMLinearKernel,
    AiterPreshuffledPerTokenFp8ScaledMMLinearKernel,
    AiterPerTokenFp8ScaledMMLinearKernel,
)

DEVICE_TYPE = current_platform.device_type


@dataclass(frozen=True)
class QTensorConfig:
    name: str
    weight: QuarkQTensorHint
    input_tensors: QuarkQTensorHint
    weight_quant_key: QuantKey | None = None
    act_quant_key: QuantKey | None = None
    dispatch_cls: type[QuarkScheme] | type[QuarkMoEMethod] | None = None
    expected_error: tuple[type[Exception], str] | None = None


QTENSOR_CONFIGS = [
    QTensorConfig(
        name="fp8_w8a8_static_tensor",
        weight={"dtype": "fp8_e4m3", "qscheme": "per_tensor", "is_dynamic": False},
        input_tensors={
            "dtype": "fp8_e4m3",
            "qscheme": "per_tensor",
            "is_dynamic": False,
        },
        weight_quant_key=kFp8StaticTensorSym,
        act_quant_key=kFp8StaticTensorSym,
        dispatch_cls=QuarkW8A8Fp8,
    ),
    QTensorConfig(
        name="fp8_w8a8_static_tensor_single_entry_lists",
        weight=[
            {
                "dtype": "fp8_e4m3",
                "qscheme": "per_tensor",
                "is_dynamic": False,
            }
        ],
        input_tensors=[
            {
                "dtype": "fp8_e4m3",
                "qscheme": "per_tensor",
                "is_dynamic": False,
            }
        ],
        weight_quant_key=kFp8StaticTensorSym,
        act_quant_key=kFp8StaticTensorSym,
        dispatch_cls=QuarkW8A8Fp8,
    ),
    QTensorConfig(
        name="fp8_w8a8_dynamic_tensor",
        weight={"dtype": "fp8_e4m3", "qscheme": "per_tensor", "is_dynamic": False},
        input_tensors={
            "dtype": "fp8_e4m3",
            "qscheme": "per_tensor",
            "is_dynamic": True,
        },
        weight_quant_key=kFp8StaticTensorSym,
        act_quant_key=kFp8DynamicTensorSym,
        dispatch_cls=QuarkW8A8Fp8,
    ),
    QTensorConfig(
        name="fp8_w8a8_dynamic_token",
        weight={"dtype": "fp8_e4m3", "qscheme": "per_channel", "is_dynamic": False},
        input_tensors={
            "dtype": "fp8_e4m3",
            "qscheme": "per_channel",
            "is_dynamic": True,
        },
        weight_quant_key=kFp8StaticChannelSym,
        act_quant_key=kFp8DynamicTokenSym,
        dispatch_cls=QuarkW8A8Fp8,
    ),
    QTensorConfig(
        name="fp8_w8a8_channel_static_tensor",
        weight={"dtype": "fp8_e4m3", "qscheme": "per_channel", "is_dynamic": False},
        input_tensors={
            "dtype": "fp8_e4m3",
            "qscheme": "per_tensor",
            "is_dynamic": False,
        },
        weight_quant_key=kFp8StaticChannelSym,
        act_quant_key=kFp8StaticTensorSym,
        dispatch_cls=QuarkW8A8Fp8,
    ),
    QTensorConfig(
        name="fp8_w8a8_tensor_dynamic_token",
        weight={"dtype": "fp8_e4m3", "qscheme": "per_tensor", "is_dynamic": False},
        input_tensors={
            "dtype": "fp8_e4m3",
            "qscheme": "per_channel",
            "is_dynamic": True,
        },
        weight_quant_key=kFp8StaticTensorSym,
        act_quant_key=kFp8DynamicTokenSym,
        dispatch_cls=QuarkW8A8Fp8,
    ),
    QTensorConfig(
        name="fp8_w8a8_dynamic_block_fp32",
        weight={
            "dtype": "fp8_e4m3",
            "qscheme": "per_block",
            "is_dynamic": False,
            "block_size": [128, 128],
            "symmetric": True,
        },
        input_tensors={
            "dtype": "fp8_e4m3",
            "qscheme": "per_group",
            "is_dynamic": True,
            "group_size": 128,
            "symmetric": True,
        },
        weight_quant_key=kFp8Static128BlockSym,
        act_quant_key=kFp8Dynamic128Sym,
        dispatch_cls=QuarkW8A8Fp8PerBlock,
    ),
    QTensorConfig(
        name="fp8_w8a8_dynamic_block_e8m0",
        weight={
            "dtype": "fp8_e4m3",
            "qscheme": "per_block",
            "is_dynamic": False,
            "block_size": [128, 128],
            "symmetric": True,
            "scale_type": "float8_e8m0fnu",
        },
        input_tensors={
            "dtype": "fp8_e4m3",
            "qscheme": "per_group",
            "is_dynamic": True,
            "group_size": 128,
            "symmetric": True,
        },
        weight_quant_key=kFp8Static128BlockE8M0Sym,
        act_quant_key=kFp8Dynamic128Sym,
        dispatch_cls=QuarkW8A8Fp8PerBlock,
    ),
    QTensorConfig(
        name="fp8_w8a8_dynamic_block_fp32_moe",
        weight={
            "dtype": "fp8_e4m3",
            "qscheme": "per_block",
            "is_dynamic": False,
            "block_size": [128, 128],
            "symmetric": True,
        },
        input_tensors={
            "dtype": "fp8_e4m3",
            "qscheme": "per_group",
            "is_dynamic": True,
            "group_size": 128,
            "symmetric": True,
        },
        weight_quant_key=kFp8Static128BlockSym,
        act_quant_key=kFp8Dynamic128Sym,
        dispatch_cls=QuarkW8A8Fp8MoEMethod,
    ),
    QTensorConfig(
        name="fp8_w8a8_block_static_input",
        weight={
            "dtype": "fp8_e4m3",
            "qscheme": "per_block",
            "is_dynamic": False,
            "block_size": [128, 128],
            "symmetric": True,
        },
        input_tensors={
            "dtype": "fp8_e4m3",
            "qscheme": "per_group",
            "is_dynamic": False,
            "group_size": 128,
            "symmetric": True,
        },
        expected_error=(NotImplementedError, "No quark compatible scheme"),
    ),
    QTensorConfig(
        name="fp8_w8a8_block_group_size_mismatch",
        weight={
            "dtype": "fp8_e4m3",
            "qscheme": "per_block",
            "is_dynamic": False,
            "block_size": [128, 128],
            "symmetric": True,
        },
        input_tensors={
            "dtype": "fp8_e4m3",
            "qscheme": "per_group",
            "is_dynamic": True,
            "group_size": 64,
            "symmetric": True,
        },
        expected_error=(NotImplementedError, "No quark compatible scheme"),
    ),
    QTensorConfig(
        name="fp8_w8a8_block_missing_block_size",
        weight={
            "dtype": "fp8_e4m3",
            "qscheme": "per_block",
            "is_dynamic": False,
            "symmetric": True,
        },
        input_tensors={
            "dtype": "fp8_e4m3",
            "qscheme": "per_group",
            "is_dynamic": True,
            "group_size": 128,
            "symmetric": True,
        },
        expected_error=(ValueError, "requires `block_size`"),
    ),
    QTensorConfig(
        name="int8_w8a8_static_symmetric",
        weight={
            "dtype": "int8",
            "qscheme": "per_tensor",
            "is_dynamic": False,
            "symmetric": True,
        },
        input_tensors={
            "dtype": "int8",
            "qscheme": "per_tensor",
            "is_dynamic": False,
            "symmetric": True,
        },
        weight_quant_key=kInt8StaticTensorSym,
        act_quant_key=kInt8StaticTensorSym,
        dispatch_cls=QuarkW8A8Int8,
    ),
    QTensorConfig(
        name="int8_w8a8_static_asymmetric",
        weight={
            "dtype": "int8",
            "qscheme": "per_tensor",
            "is_dynamic": False,
            "symmetric": True,
        },
        input_tensors={
            "dtype": "int8",
            "qscheme": "per_tensor",
            "is_dynamic": False,
            "symmetric": False,
        },
        weight_quant_key=kInt8StaticTensorSym,
        act_quant_key=kInt8StaticTensorAsym,
        dispatch_cls=QuarkW8A8Int8,
    ),
    QTensorConfig(
        name="int8_w8a8_channel_static_symmetric",
        weight={
            "dtype": "int8",
            "qscheme": "per_channel",
            "is_dynamic": False,
            "symmetric": True,
        },
        input_tensors={
            "dtype": "int8",
            "qscheme": "per_tensor",
            "is_dynamic": False,
            "symmetric": True,
        },
        weight_quant_key=kInt8StaticChannelSym,
        act_quant_key=kInt8StaticTensorSym,
        dispatch_cls=QuarkW8A8Int8,
    ),
    QTensorConfig(
        name="int8_w8a8_channel_static_asymmetric",
        weight={
            "dtype": "int8",
            "qscheme": "per_channel",
            "is_dynamic": False,
            "symmetric": True,
        },
        input_tensors={
            "dtype": "int8",
            "qscheme": "per_tensor",
            "is_dynamic": False,
            "symmetric": False,
        },
        weight_quant_key=kInt8StaticChannelSym,
        act_quant_key=kInt8StaticTensorAsym,
        dispatch_cls=QuarkW8A8Int8,
    ),
    QTensorConfig(
        name="int8_w8a8_dynamic_tensor_symmetric",
        weight={
            "dtype": "int8",
            "qscheme": "per_tensor",
            "is_dynamic": False,
            "symmetric": True,
        },
        input_tensors={
            "dtype": "int8",
            "qscheme": "per_channel",
            "is_dynamic": True,
            "symmetric": True,
        },
        weight_quant_key=kInt8StaticTensorSym,
        act_quant_key=kInt8DynamicTensorSym,
        dispatch_cls=QuarkW8A8Int8,
    ),
    QTensorConfig(
        name="int8_w8a8_dynamic_tensor_asymmetric",
        weight={
            "dtype": "int8",
            "qscheme": "per_tensor",
            "is_dynamic": False,
            "symmetric": True,
        },
        input_tensors={
            "dtype": "int8",
            "qscheme": "per_channel",
            "is_dynamic": True,
            "symmetric": False,
        },
        weight_quant_key=kInt8StaticTensorSym,
        act_quant_key=kInt8DynamicTensorAsym,
        dispatch_cls=QuarkW8A8Int8,
    ),
    QTensorConfig(
        name="int8_w8a8_dynamic_token",
        weight={
            "dtype": "int8",
            "qscheme": "per_channel",
            "is_dynamic": False,
            "symmetric": True,
        },
        input_tensors={
            "dtype": "int8",
            "qscheme": "per_channel",
            "is_dynamic": True,
            "symmetric": True,
        },
        weight_quant_key=kInt8StaticChannelSym,
        act_quant_key=kInt8DynamicTokenSym,
        dispatch_cls=QuarkW8A8Int8,
    ),
    QTensorConfig(
        name="int8_w8a8_dynamic_token_asymmetric",
        weight={
            "dtype": "int8",
            "qscheme": "per_channel",
            "is_dynamic": False,
            "symmetric": True,
        },
        input_tensors={
            "dtype": "int8",
            "qscheme": "per_channel",
            "is_dynamic": True,
            "symmetric": False,
        },
        weight_quant_key=kInt8StaticChannelSym,
        act_quant_key=kInt8DynamicTokenAsym,
        dispatch_cls=QuarkW8A8Int8,
    ),
    QTensorConfig(
        name="ocp_mx_mxfp4_weight_only",
        weight={
            "dtype": "fp4",
            "qscheme": "per_group",
            "group_size": 32,
            "scale_format": "e8m0",
            "is_dynamic": False,
        },
        input_tensors=None,
        weight_quant_key=kMxfp4Static,
        act_quant_key=None,
        dispatch_cls=QuarkOCP_MX,
    ),
    QTensorConfig(
        name="ocp_mx_mxfp4_activation",
        weight={
            "dtype": "fp4",
            "qscheme": "per_group",
            "group_size": 32,
            "scale_format": "e8m0",
            "is_dynamic": False,
        },
        input_tensors={
            "dtype": "fp4",
            "qscheme": "per_group",
            "group_size": 32,
            "scale_format": "e8m0",
            "is_dynamic": True,
        },
        weight_quant_key=kMxfp4Static,
        act_quant_key=kMxfp4Dynamic,
        dispatch_cls=QuarkOCP_MX,
    ),
    QTensorConfig(
        name="ocp_mx_mxfp6_e3m2",
        weight={
            "dtype": "fp6_e3m2",
            "qscheme": "per_group",
            "group_size": 32,
            "scale_format": "e8m0",
            "is_dynamic": False,
        },
        input_tensors={
            "dtype": "fp6_e3m2",
            "qscheme": "per_group",
            "group_size": 32,
            "scale_format": "e8m0",
            "is_dynamic": True,
        },
        weight_quant_key=kMxfp6E3M2Static,
        act_quant_key=kMxfp6E3M2Dynamic,
        dispatch_cls=QuarkOCP_MX,
    ),
    QTensorConfig(
        name="ocp_mx_mxfp4_mxfp6_e3m2_activation",
        weight={
            "dtype": "fp4",
            "qscheme": "per_group",
            "group_size": 32,
            "scale_format": "e8m0",
            "is_dynamic": False,
        },
        input_tensors={
            "dtype": "fp6_e3m2",
            "qscheme": "per_group",
            "group_size": 32,
            "scale_format": "e8m0",
            "is_dynamic": True,
        },
        weight_quant_key=kMxfp4Static,
        act_quant_key=kMxfp6E3M2Dynamic,
        dispatch_cls=QuarkOCP_MX,
    ),
    QTensorConfig(
        name="ocp_mx_mxfp4_mxfp6_e2m3_activation",
        weight={
            "dtype": "fp4",
            "qscheme": "per_group",
            "group_size": 32,
            "scale_format": "e8m0",
            "is_dynamic": False,
        },
        input_tensors={
            "dtype": "fp6_e2m3",
            "qscheme": "per_group",
            "group_size": 32,
            "scale_format": "e8m0",
            "is_dynamic": True,
        },
        weight_quant_key=kMxfp4Static,
        act_quant_key=kMxfp6E2M3Dynamic,
        dispatch_cls=QuarkOCP_MX,
    ),
    QTensorConfig(
        name="ocp_mx_mxfp6_e2m3",
        weight={
            "dtype": "fp6_e2m3",
            "qscheme": "per_group",
            "group_size": 32,
            "scale_format": "e8m0",
            "is_dynamic": False,
        },
        input_tensors={
            "dtype": "fp6_e2m3",
            "qscheme": "per_group",
            "group_size": 32,
            "scale_format": "e8m0",
            "is_dynamic": True,
        },
        weight_quant_key=kMxfp6E2M3Static,
        act_quant_key=kMxfp6E2M3Dynamic,
        dispatch_cls=QuarkOCP_MX,
    ),
    QTensorConfig(
        name="nvfp4",
        weight=[
            {
                "dtype": "fp4",
                "qscheme": "per_group",
                "group_size": 16,
                "is_dynamic": False,
            },
            {"dtype": "fp8_e4m3", "qscheme": "per_tensor", "is_dynamic": False},
        ],
        input_tensors=[
            {
                "dtype": "fp4",
                "qscheme": "per_group",
                "group_size": 16,
                "is_dynamic": True,
            },
            {"dtype": "fp8_e4m3", "qscheme": "per_tensor", "is_dynamic": False},
        ],
        weight_quant_key=kNvfp4Static,
        act_quant_key=kNvfp4Dynamic,
        dispatch_cls=QuarkNVFP4,
    ),
    QTensorConfig(
        name="w4a8_fp8_static",
        weight=[
            {"dtype": "fp8_e4m3", "qscheme": "per_tensor", "is_dynamic": False},
            {
                "dtype": "int4",
                "qscheme": "per_channel",
                "is_dynamic": False,
                "symmetric": True,
                "ch_axis": 0,
            },
        ],
        input_tensors={
            "dtype": "fp8_e4m3",
            "qscheme": "per_tensor",
            "is_dynamic": False,
        },
        weight_quant_key=kInt4W4A8StaticChannelSym,
        act_quant_key=kFp8StaticTensorSym,
        dispatch_cls=QuarkW4A8Fp8MoEMethod,
    ),
    QTensorConfig(
        name="w4a8_fp8_dynamic",
        weight=[
            {"dtype": "fp8_e4m3", "qscheme": "per_tensor", "is_dynamic": False},
            {
                "dtype": "int4",
                "qscheme": "per_channel",
                "is_dynamic": False,
                "symmetric": True,
                "ch_axis": 0,
            },
        ],
        input_tensors={
            "dtype": "fp8_e4m3",
            "qscheme": "per_channel",
            "is_dynamic": True,
        },
        weight_quant_key=kInt4W4A8StaticChannelSym,
        act_quant_key=kFp8DynamicTokenSym,
        dispatch_cls=QuarkW4A8Fp8MoEMethod,
    ),
    QTensorConfig(
        name="w4a8_fp8_static_single_entry_input",
        weight=[
            {"dtype": "fp8_e4m3", "qscheme": "per_tensor", "is_dynamic": False},
            {
                "dtype": "int4",
                "qscheme": "per_channel",
                "is_dynamic": False,
                "symmetric": True,
                "ch_axis": 0,
            },
        ],
        input_tensors=[
            {
                "dtype": "fp8_e4m3",
                "qscheme": "per_tensor",
                "is_dynamic": False,
            }
        ],
        weight_quant_key=kInt4W4A8StaticChannelSym,
        act_quant_key=kFp8StaticTensorSym,
        dispatch_cls=QuarkW4A8Fp8MoEMethod,
    ),
]


def _make_qtensor_config(
    weight: QuarkQTensorHint,
    input_tensors: QuarkQTensorHint,
    exclude: list[str] | None = None,
) -> QuarkConfig:
    return QuarkConfig(
        {
            "global_quant_config": {
                "weight": weight,
                "input_tensors": input_tensors,
            },
            "layer_type_quant_config": {},
            "exclude": exclude or [],
        }
    )


def _make_test_moe_config() -> FusedMoEConfig:
    return FusedMoEConfig(
        num_experts=8,
        experts_per_token=2,
        hidden_dim=256,
        intermediate_size=256,
        num_local_experts=8,
        num_logical_experts=8,
        activation=MoEActivation.SILU,
        device=current_platform.device_type,
        routing_method=RoutingMethodType.Renormalize,
        moe_parallel_config=FusedMoEParallelConfig.make_no_parallel(),
        in_dtype=torch.bfloat16,
    )


if QUARK_MXFP4_AVAILABLE:
    from quark.torch.export.nn.modules.realquantizer import StaticScaledRealQuantizer
    from quark.torch.kernel import mx as mx_kernel
    from quark.torch.quantization.config.config import FP4PerGroupSpec

try:
    hf_api().list_repo_refs(
        "amd/Llama-3.3-70B-Instruct-WMXFP4-AMXFP4-KVFP8-Scale-UINT8-SQ"
    )
    HF_HUB_AMD_ORG_ACCESS = True
except huggingface_hub.errors.RepositoryNotFoundError:
    HF_HUB_AMD_ORG_ACCESS = False


@pytest.fixture(scope="function", autouse=True)
def enable_pickle(monkeypatch):
    """`LLM.apply_model` requires pickling a function."""
    monkeypatch.setenv("VLLM_ALLOW_INSECURE_SERIALIZATION", "1")


def test_quark_w8a8_fp8_per_block_registers_weight_scale(monkeypatch):
    from vllm.model_executor.layers.quantization.utils.fp8_utils import (
        get_fp8_block_weight_scale,
    )

    monkeypatch.setattr(
        "vllm.model_executor.layers.quantization.quark.schemes."
        "quark_w8a8_fp8.get_current_vllm_config",
        lambda: SimpleNamespace(model_config=SimpleNamespace(dtype=torch.bfloat16)),
    )
    scheme = QuarkW8A8Fp8PerBlock(kFp8Static128BlockSym, kFp8Dynamic128Sym)

    layer = torch.nn.Module()
    layer.weight_scale = torch.tensor([2.0])
    assert get_fp8_block_weight_scale(layer) is None
    layer.scheme = scheme
    assert get_fp8_block_weight_scale(layer) is layer.weight_scale
    layer.weight_scale_inv = torch.tensor([3.0])
    assert get_fp8_block_weight_scale(layer) is layer.weight_scale
    layer.scheme = None
    assert get_fp8_block_weight_scale(layer) is layer.weight_scale_inv

    loaded = torch.nn.Module()

    def weight_loader(param, loaded_weight):
        return None

    dummy_param = torch.nn.Parameter(torch.empty(1), requires_grad=False)
    with (
        patch(
            "vllm.model_executor.layers.quantization.quark.schemes.quark_w8a8_fp8."
            "validate_fp8_block_shape"
        ),
        patch(
            "vllm.model_executor.layers.quantization.quark.schemes.quark_w8a8_fp8."
            "create_fp8_weight_parameter",
            return_value=dummy_param,
        ),
        patch(
            "vllm.model_executor.layers.quantization.quark.schemes.quark_w8a8_fp8."
            "create_fp8_scale_parameter",
            return_value=dummy_param,
        ),
        patch(
            "vllm.model_executor.layers.quantization.quark.schemes.quark_w8a8_fp8."
            "init_fp8_linear_kernel",
            return_value=MagicMock(),
        ),
    ):
        scheme.create_weights(
            loaded,
            output_partition_sizes=[256],
            input_size_per_partition=256,
            params_dtype=torch.bfloat16,
            weight_loader=weight_loader,
            input_size=256,
            output_size=256,
        )
    assert hasattr(loaded, "weight_scale")
    assert not hasattr(loaded, "weight_scale_inv")


def test_quark_config_has_no_model_specific_fused_mappings():
    config = QuarkConfig({})

    assert "gate_up_proj" not in config.packed_modules_mapping
    assert "fused_wqa_wkv" not in config.packed_modules_mapping


def test_quark_config_preserves_existing_packed_modules_mapping():
    class CustomQuarkConfig(QuarkConfig):
        packed_modules_mapping = {"custom_proj": ["a", "b"]}

    config = CustomQuarkConfig({})

    assert config.packed_modules_mapping["custom_proj"] == ["a", "b"]


def test_quant_method_dispatch_ignored(default_vllm_config):
    config = _make_qtensor_config(None, None, exclude=["linear", "experts"])

    class TestLinear(LinearBase):
        def __init__(self):
            torch.nn.Module.__init__(self)

    class TestRoutedExperts(RoutedExperts):
        def __init__(self):
            torch.nn.Module.__init__(self)
            self.moe_config = _make_test_moe_config()

    assert config.get_quant_method_target("linear", LinearBase) == (
        None,
        None,
        UnquantizedLinearMethod,
    )
    assert isinstance(
        config.get_quant_method(TestLinear(), "linear"), UnquantizedLinearMethod
    )

    assert config.get_quant_method_target("experts", RoutedExperts) == (
        None,
        None,
        UnquantizedFusedMoEMethod,
    )
    assert isinstance(
        config.get_quant_method(TestRoutedExperts(), "experts"),
        UnquantizedFusedMoEMethod,
    )

    mxfp4_config = _make_qtensor_config(
        {
            "dtype": "fp4",
            "qscheme": "per_group",
            "group_size": 32,
            "scale_format": "e8m0",
            "is_dynamic": False,
        },
        None,
        exclude=["self_attn.q_proj", "mlp.down_proj"],
    )
    assert mxfp4_config.get_quant_method_target("self_attn.q_proj", LinearBase) == (
        None,
        None,
        UnquantizedLinearMethod,
    )
    assert isinstance(
        mxfp4_config.get_quant_method(TestLinear(), "self_attn.q_proj"),
        UnquantizedLinearMethod,
    )

    assert mxfp4_config.get_quant_method_target("mlp.down_proj", LinearBase) == (
        None,
        None,
        UnquantizedLinearMethod,
    )
    assert isinstance(
        mxfp4_config.get_quant_method(TestLinear(), "mlp.down_proj"),
        UnquantizedLinearMethod,
    )


@pytest.mark.parametrize("case", QTENSOR_CONFIGS, ids=lambda case: case.name)
def test_quant_method_dispatch_target(case):
    config = _make_qtensor_config(case.weight, case.input_tensors)
    if case.expected_error is not None:
        error_type, error_message = case.expected_error
        with pytest.raises(error_type, match=error_message):
            config.get_quant_method_target("linear", LinearBase)
        return

    assert case.dispatch_cls is not None
    is_linear = issubclass(case.dispatch_cls, QuarkScheme)

    weight_quant_key, act_quant_key, method_cls = config.get_quant_method_target(
        "linear" if is_linear else "experts",
        LinearBase if is_linear else RoutedExperts,
    )

    assert weight_quant_key == case.weight_quant_key
    assert act_quant_key == case.act_quant_key
    assert method_cls is (QuarkLinearMethod if is_linear else case.dispatch_cls)


def test_quant_method_dispatch_mxfp8_2d_block(default_vllm_config):
    """Shape B (DeepSeek-V4.1): 32x32 per-block e8m0 routes to QuarkOCP_MX.

    Verifies that get_quant_method and get_quant_method_target agree (the
    short-circuit that previously caused them to diverge has been removed).
    """
    from vllm.model_executor.layers.quantization.quark.schemes import QuarkOCP_MX
    from vllm.model_executor.layers.quantization.utils.quant_utils import kMxfp8Static

    default_vllm_config.model_config = SimpleNamespace(dtype=torch.bfloat16)
    mxfp8_spec = {
        "weight": {
            "dtype": "fp8_e4m3",
            "qscheme": "per_block",
            "block_size": [32, 32],
            "symmetric": True,
            "is_dynamic": False,
            "scale_type": "float8_e8m0fnu",
        },
        "input_tensors": {
            "dtype": "fp8_e4m3",
            "qscheme": "per_group",
            "group_size": 32,
            "symmetric": True,
            "is_dynamic": True,
        },
    }
    config = QuarkConfig(
        {
            "global_quant_config": {
                "weight": {
                    "dtype": "fp4",
                    "qscheme": "per_group",
                    "group_size": 32,
                    "scale_format": "e8m0",
                    "is_dynamic": False,
                },
                "input_tensors": {
                    "dtype": "fp4",
                    "qscheme": "per_group",
                    "group_size": 32,
                    "scale_format": "e8m0",
                    "is_dynamic": True,
                },
            },
            "layer_type_quant_config": {},
            "layer_quant_config": {"layers.0.attn.wkv": mxfp8_spec},
            "exclude": [],
        }
    )

    class TestLinear(LinearBase):
        def __init__(self):
            torch.nn.Module.__init__(self)

    # get_quant_method_target now routes MXFP8 through the matcher chain.
    wk, ak, mcls = config.get_quant_method_target("layers.0.attn.wkv", LinearBase)
    assert wk == kMxfp8Static
    assert mcls is QuarkLinearMethod

    linear = TestLinear()
    method = config.get_quant_method(linear, "layers.0.attn.wkv")
    assert isinstance(method, QuarkLinearMethod)
    assert isinstance(linear.scheme, QuarkOCP_MX)
    assert linear.scheme.weight_quant_key == kMxfp8Static
    assert linear.scheme.scale_block_rows == 32

    # Experts still fall through to the global MXFP4 spec.
    assert (
        config.get_quant_method_target("layers.0.ffn.experts", RoutedExperts)[2]
        is not QuarkLinearMethod
    )

    # 128x128 per-block FP8 (DeepSeek V4) keeps its existing Quark scheme.
    v4_spec = {
        "weight": {**mxfp8_spec["weight"], "block_size": [128, 128]},
        "input_tensors": {**mxfp8_spec["input_tensors"], "group_size": 128},
    }
    v4_config = _make_qtensor_config(v4_spec["weight"], v4_spec["input_tensors"])
    assert v4_config.get_quant_method_target("linear", LinearBase)[2] is (
        QuarkLinearMethod
    )
    linear = TestLinear()
    assert isinstance(v4_config.get_quant_method(linear, "linear"), QuarkLinearMethod)
    assert isinstance(linear.scheme, QuarkW8A8Fp8PerBlock)


def test_quant_method_dispatch_mxfp8_canonical(default_vllm_config):
    """Shape A (canonical MX): per_group/group_size=32/e8m0 FP8 routes
    to QuarkOCP_MX with scale_block_rows == 1.

    This covers the canonical MXFP8 spelling (no 2-D block). Validated by
    synthetic dispatch test only — no canonical-MXFP8 checkpoint on hand.
    """
    from vllm.model_executor.layers.quantization.quark.schemes import QuarkOCP_MX
    from vllm.model_executor.layers.quantization.utils.quant_utils import kMxfp8Static

    default_vllm_config.model_config = SimpleNamespace(dtype=torch.bfloat16)
    mxfp8_canonical_spec = {
        "weight": {
            "dtype": "fp8_e4m3",
            "qscheme": "per_group",
            "group_size": 32,
            "scale_format": "e8m0",
            "is_dynamic": False,
        },
        "input_tensors": {
            "dtype": "fp8_e4m3",
            "qscheme": "per_group",
            "group_size": 32,
            "scale_format": "e8m0",
            "is_dynamic": True,
        },
    }
    config = QuarkConfig(
        {
            "global_quant_config": mxfp8_canonical_spec,
            "layer_type_quant_config": {},
            "layer_quant_config": {},
            "exclude": [],
        }
    )

    wk, ak, mcls = config.get_quant_method_target("linear", LinearBase)
    assert wk == kMxfp8Static
    assert mcls is QuarkLinearMethod

    class TestLinear(LinearBase):
        def __init__(self):
            torch.nn.Module.__init__(self)

    linear = TestLinear()
    method = config.get_quant_method(linear, "linear")
    assert isinstance(method, QuarkLinearMethod)
    assert isinstance(linear.scheme, QuarkOCP_MX)
    assert linear.scheme.weight_quant_key == kMxfp8Static
    assert linear.scheme.scale_block_rows == 1


def test_quant_method_dispatch_mxfp8_moe_raises(default_vllm_config):
    """MXFP8 in a MoE config raises ValueError — experts are unsupported."""
    from vllm.model_executor.layers.quantization.quark.quark_moe import (
        QuarkOCP_MX_MoEMethod,
    )

    default_vllm_config.model_config = SimpleNamespace(dtype=torch.bfloat16)
    mxfp8_spec = {
        "weight": {
            "dtype": "fp8_e4m3",
            "qscheme": "per_group",
            "group_size": 32,
            "scale_format": "e8m0",
            "is_dynamic": False,
        },
        "input_tensors": {
            "dtype": "fp8_e4m3",
            "qscheme": "per_group",
            "group_size": 32,
            "scale_format": "e8m0",
            "is_dynamic": True,
        },
    }
    config = QuarkConfig(
        {
            "global_quant_config": mxfp8_spec,
            "layer_type_quant_config": {},
            "layer_quant_config": {},
            "exclude": [],
        }
    )
    wk, ak, mcls = config.get_quant_method_target("experts", RoutedExperts)
    assert mcls is QuarkOCP_MX_MoEMethod
    assert wk is not None
    # The OCP MX MoE constructor should fail loudly for MXFP8.
    fake_moe_config = MagicMock()
    with pytest.raises(ValueError, match="MXFP8 experts are not supported"):
        QuarkOCP_MX_MoEMethod(fake_moe_config, wk, ak)


@pytest.mark.parametrize(
    ("weight", "input_tensors"),
    [
        pytest.param(
            {
                "dtype": "int8",
                "qscheme": "per_group",
                "is_dynamic": False,
                "symmetric": True,
            },
            {
                "dtype": "int8",
                "qscheme": "per_tensor",
                "is_dynamic": False,
                "symmetric": True,
            },
            id="single_entry",
        ),
        pytest.param(
            [
                {"dtype": "int8", "qscheme": "per_tensor"},
                {"dtype": "int8", "qscheme": "per_tensor"},
            ],
            [
                {"dtype": "int8", "qscheme": "per_tensor"},
                {"dtype": "int8", "qscheme": "per_tensor"},
            ],
            id="multi_entry",
        ),
    ],
)
def test_quant_method_dispatch_unsupported(weight, input_tensors):
    config = _make_qtensor_config(weight, input_tensors)

    class TestRoutedExperts(RoutedExperts):
        def __init__(self):
            torch.nn.Module.__init__(self)

    with pytest.raises(RuntimeError, match="^Unsupported FusedMoe scheme$"):
        config.get_quant_method_target("experts", RoutedExperts)

    with pytest.raises(RuntimeError, match="^Unsupported FusedMoe scheme$"):
        config.get_quant_method(TestRoutedExperts(), "experts")


@pytest.mark.parametrize(
    "case",
    [case for case in QTENSOR_CONFIGS if case.expected_error is None],
    ids=lambda case: case.name,
)
def test_quant_method_dispatch_instantiation(case, monkeypatch, default_vllm_config):
    config = _make_qtensor_config(case.weight, case.input_tensors)
    assert case.dispatch_cls is not None
    if issubclass(case.dispatch_cls, QuarkScheme):

        class TestLinear(LinearBase):
            def __init__(self):
                torch.nn.Module.__init__(self)

        monkeypatch.setattr(
            "vllm.model_executor.layers.quantization.quark.schemes."
            "quark_w8a8_fp8.get_current_vllm_config",
            lambda: SimpleNamespace(model_config=SimpleNamespace(dtype=torch.bfloat16)),
        )
        layer = TestLinear()
        method = config.get_quant_method(layer, "linear")

        assert isinstance(method, QuarkLinearMethod)
        assert isinstance(layer.scheme, case.dispatch_cls)
        if case.weight_quant_key == kFp8Static128BlockE8M0Sym:
            # TODO: Remove once E8M0 quant key is properly handled in oracle
            assert layer.scheme.weight_quant_key == kFp8Static128BlockSym
        else:
            assert layer.scheme.weight_quant_key == case.weight_quant_key
        assert layer.scheme.activation_quant_key == case.act_quant_key
    else:

        class TestRoutedExperts(RoutedExperts):
            def __init__(self):
                torch.nn.Module.__init__(self)
                self.moe_config = _make_test_moe_config()

        for target in (
            "select_fp8_moe_backend",
            "select_int8_moe_backend",
            "select_mxfp4_moe_backend",
            "backend_to_kernel_cls",
            "select_nvfp4_moe_backend",
        ):
            monkeypatch.setattr(
                f"vllm.model_executor.layers.quantization.quark.quark_moe.{target}",
                lambda *args, **kwargs: (object(), object()),
            )

        # AssertionError: W4A8 FP8 MoE requires ROCm AITER fused MoE support
        monkeypatch.setattr(
            "vllm.model_executor.layers.quantization.quark.quark_moe."
            "rocm_aiter_ops.is_fused_moe_enabled",
            lambda: True,
        )

        # default_vllm_config carries no model, but the FP8 and OCP MX methods
        # read the model type off the HF config.
        monkeypatch.setattr(
            "vllm.model_executor.layers.quantization.quark.quark_moe."
            "get_current_vllm_config",
            lambda: SimpleNamespace(
                model_config=SimpleNamespace(hf_config=SimpleNamespace())
            ),
        )

        layer = TestRoutedExperts()
        method = config.get_quant_method(layer, "experts")

        assert isinstance(method, case.dispatch_cls)


QUARK_MOE_MODULE = "vllm.model_executor.layers.quantization.quark.quark_moe"
# validate_fp8_block_shape_moe imports this from vllm.distributed when called,
# so the patch has to target the source module rather than a local binding.
TP_WORLD_SIZE = "vllm.distributed.get_tensor_model_parallel_world_size"


def _make_per_block_fp8_moe_method(
    activation_quant_key: QuantKey = kFp8Dynamic128Sym,
) -> QuarkW8A8Fp8MoEMethod:
    return QuarkW8A8Fp8MoEMethod(
        _make_test_moe_config(),
        kFp8Static128BlockSym,
        activation_quant_key,
    )


def test_quark_w8a8_fp8_moe_per_block_requires_dynamic_group_input():
    with (
        patch(
            f"{QUARK_MOE_MODULE}.select_fp8_moe_backend",
            return_value=(Mock(), Mock()),
        ),
        patch(f"{QUARK_MOE_MODULE}.get_current_vllm_config"),
        pytest.raises(ValueError, match="per-block scales"),
    ):
        _make_per_block_fp8_moe_method(kFp8StaticTensorSym)


def test_quark_w8a8_fp8_moe_per_block_weight_shapes():
    with (
        patch(
            f"{QUARK_MOE_MODULE}.select_fp8_moe_backend",
            return_value=(Mock(), Mock()),
        ),
        patch(f"{QUARK_MOE_MODULE}.get_current_vllm_config"),
        patch(TP_WORLD_SIZE, return_value=1),
    ):
        method = _make_per_block_fp8_moe_method()
        assert method.block_quant
        assert method.weight_block_size == [128, 128]

        layer = torch.nn.Module()
        method.create_weights(
            layer,
            num_experts=4,
            hidden_size=512,
            intermediate_size_per_partition=256,
            params_dtype=torch.bfloat16,
        )

    w13_num_shards = method.moe.w13_num_shards
    assert layer.weight_block_size == [128, 128]
    assert layer.w13_weight.shape == (4, w13_num_shards * 256, 512)
    assert layer.w2_weight.shape == (4, 512, 256)
    # Quark exports block scales as `weight_scale`, like the per-tensor and
    # per-channel schemes, so all schemes register the same parameter name.
    assert not hasattr(layer, "w13_weight_scale_inv")
    assert not hasattr(layer, "w2_weight_scale_inv")
    # One scale per 128x128 tile of each expert's weight.
    assert layer.w13_weight_scale.shape == (
        4,
        w13_num_shards * (256 // 128),
        512 // 128,
    )
    assert layer.w2_weight_scale.shape == (4, 512 // 128, 256 // 128)
    # The loader shards block scales on the block grid, not per row.
    for scale in (layer.w13_weight_scale, layer.w2_weight_scale):
        assert scale.quant_method == FusedMoeWeightScaleSupported.BLOCK.value


def test_quark_w8a8_fp8_moe_per_block_rejects_misaligned_partition():
    with (
        patch(
            f"{QUARK_MOE_MODULE}.select_fp8_moe_backend",
            return_value=(Mock(), Mock()),
        ),
        patch(f"{QUARK_MOE_MODULE}.get_current_vllm_config"),
        patch(TP_WORLD_SIZE, return_value=1),
        pytest.raises(ValueError, match="not divisible by"),
    ):
        _make_per_block_fp8_moe_method().create_weights(
            torch.nn.Module(),
            num_experts=4,
            hidden_size=512,
            intermediate_size_per_partition=192,
            params_dtype=torch.bfloat16,
        )


def test_quark_fp8_ptpc_exposes_kernel_input_quant_key(monkeypatch):
    """QuarkW8A8Fp8 must advertise the key its kernel consumes pre-quantized.

    Only kernel selection runs, no GEMM. The shape matters: AITER PTPC
    kernels decline untuned shapes, leaving a torch kernel that has no key.
    """
    # Llama-3.1-70B qkv_proj at TP1 (aiter-tuned N, K on gfx950).
    N, K = 10240, 8192
    dtype = torch.bfloat16
    kernel_config = FP8ScaledMMLinearLayerConfig(
        weight_quant_key=kFp8StaticChannelSym,
        activation_quant_key=kFp8DynamicTokenSym,
        weight_shape=(N, K),
        input_dtype=dtype,
        out_dtype=dtype,
    )
    if not any(
        cls.is_supported()[0] and cls.can_implement(kernel_config)[0]
        for cls in AITER_PTPC_KERNELS
    ):
        pytest.skip("no AITER PTPC kernel is usable for this shape and environment")

    monkeypatch.setattr(parameter, "get_tensor_model_parallel_rank", lambda: 0)
    monkeypatch.setattr(parameter, "get_tensor_model_parallel_world_size", lambda: 1)

    scheme = QuarkW8A8Fp8.__new__(QuarkW8A8Fp8)
    scheme.weight_qscheme = "per_channel"
    scheme.is_static_input_scheme = False
    scheme.activation_quant_key = kFp8DynamicTokenSym
    scheme.weight_quant_key = kFp8StaticChannelSym
    scheme.out_dtype = dtype
    scheme.input_dtype = dtype

    layer = torch.nn.Module()
    with set_current_vllm_config(VllmConfig()):
        scheme.create_weights(
            layer,
            output_partition_sizes=[N],
            input_size_per_partition=K,
            params_dtype=dtype,
            weight_loader=lambda *args, **kwargs: None,
        )

    assert isinstance(scheme.fp8_linear, AITER_PTPC_KERNELS)
    assert layer.input_quant_key == kFp8DynamicTokenSym


@pytest.mark.parametrize("kv_cache_dtype", ["auto", "fp8"])
def test_quark_fp8_w_per_tensor_a_per_tensor(
    kv_cache_dtype: CacheDType, monkeypatch, dist_init, workspace_init
):
    model_path = "amd/Llama-3.1-8B-Instruct-FP8-KV-Quark-test"
    checkpoint_scales = {}
    scale_names = {
        "model.layers.0.self_attn.k_proj.output_scale",
        "model.layers.0.self_attn.v_proj.output_scale",
    }
    original_load_weights = LlamaForCausalLM.load_weights

    def load_weights(self, weights):
        def capture_scales():
            for name, weight in weights:
                if name in scale_names:
                    checkpoint_scales[name] = weight.detach().cpu()
                yield name, weight

        return original_load_weights(self, capture_scales())

    monkeypatch.setattr(LlamaForCausalLM, "load_weights", load_weights)
    model, vllm_config = load_model_without_vllm_runner(
        model_path,
        model_config_kwargs={"hf_overrides": {"num_hidden_layers": 3}},
        vllm_config_kwargs={"cache_config": CacheConfig(cache_dtype=kv_cache_dtype)},
    )

    qkv_proj = model.model.layers[0].self_attn.qkv_proj
    assert isinstance(qkv_proj.quant_method, QuarkLinearMethod)
    assert isinstance(qkv_proj.scheme, QuarkW8A8Fp8)
    assert len(qkv_proj.input_scale.shape) == 0
    assert qkv_proj.weight.dtype is current_platform.fp8_dtype()
    assert len(qkv_proj.weight_scale.shape) == 0

    attn = model.model.layers[0].self_attn.attn
    if kv_cache_dtype == "fp8":
        assert checkpoint_scales.keys() == scale_names
        scale_multiplier = 2 if current_platform.is_fp8_fnuz() else 1
        assert attn._k_scale_float == (
            checkpoint_scales["model.layers.0.self_attn.k_proj.output_scale"].item()
            * scale_multiplier
        )
        assert attn._v_scale_float == (
            checkpoint_scales["model.layers.0.self_attn.v_proj.output_scale"].item()
            * scale_multiplier
        )
    else:
        assert attn._k_scale_float == 1.0
        assert attn._v_scale_float == 1.0

    monkeypatch.setattr(Attention, "forward", lambda _, q, k, v: q.contiguous())
    input_ids = torch.tensor([1, 2, 3, 4], device=DEVICE_TYPE)
    positions = torch.arange(input_ids.numel(), device=DEVICE_TYPE)
    with (
        set_current_vllm_config(vllm_config),
        set_forward_context(None, vllm_config, num_tokens=input_ids.numel()),
    ):
        hidden_states = model(input_ids, positions, None)
        logits = model.compute_logits(hidden_states)
    assert torch.isfinite(logits).all()


def test_quark_fp8_w_per_channel_a_per_token(monkeypatch, dist_init, workspace_init):
    model_path = "amd/Qwen2.5-1.5B-Instruct-ptpc-Quark-ts"
    model, vllm_config = load_model_without_vllm_runner(
        model_path,
        model_config_kwargs={"hf_overrides": {"num_hidden_layers": 3}},
    )

    qkv_proj = model.model.layers[0].self_attn.qkv_proj
    assert isinstance(qkv_proj.quant_method, QuarkLinearMethod)
    assert isinstance(qkv_proj.scheme, QuarkW8A8Fp8)
    assert qkv_proj.weight.dtype is current_platform.fp8_dtype()
    assert qkv_proj.weight_scale.shape[0] == qkv_proj.weight.shape[1]
    assert qkv_proj.weight_scale.shape[1] == 1

    monkeypatch.setattr(Attention, "forward", lambda _, q, k, v: q.contiguous())
    input_ids = torch.tensor([1, 2, 3, 4], device=DEVICE_TYPE)
    positions = torch.arange(input_ids.numel(), device=DEVICE_TYPE)
    with (
        set_current_vllm_config(vllm_config),
        set_forward_context(None, vllm_config, num_tokens=input_ids.numel()),
    ):
        hidden_states = model(input_ids, positions, None)
        logits = model.compute_logits(hidden_states)
    assert torch.isfinite(logits).all()


def test_quark_int8_w_per_tensor_a_per_tensor(monkeypatch, dist_init, workspace_init):
    model_path = "amd/Llama-3.1-8B-Instruct-w-int8-a-int8-sym-test"
    model, vllm_config = load_model_without_vllm_runner(
        model_path,
        model_config_kwargs={"hf_overrides": {"num_hidden_layers": 3}},
    )
    with set_current_vllm_config(vllm_config):
        qkv_proj = model.model.layers[0].self_attn.qkv_proj
        assert isinstance(qkv_proj.quant_method, QuarkLinearMethod)
        assert isinstance(qkv_proj.scheme, QuarkW8A8Int8)

        monkeypatch.setattr(Attention, "forward", lambda _, q, k, v: q.contiguous())
        input_ids = torch.tensor([1, 2, 3, 4], device=DEVICE_TYPE)
        positions = torch.arange(input_ids.numel(), device=DEVICE_TYPE)
        with set_forward_context(None, vllm_config, num_tokens=input_ids.numel()):
            hidden_states = model(input_ids, positions, None)
            logits = model.compute_logits(hidden_states)
        assert torch.isfinite(logits).all()


@pytest.mark.parametrize("tp", [1])
def test_quark_int8_w8a8_moe(vllm_runner, tp):
    """Test W8A8 INT8 MoE quantization with a tiny Qwen3 MoE model."""
    model_path = "amd/tiny-qwen3-moe-w8a8-int8"
    with vllm_runner(
        model_path,
        enforce_eager=True,
        tensor_parallel_size=tp,
        gpu_memory_utilization=0.1,
    ) as llm:

        def check_model(model):
            layer = model.model.layers[0]
            # MoE experts should use QuarkW8A8Int8MoEMethod
            moe = layer.mlp.experts
            assert isinstance(moe._quant_method, QuarkW8A8Int8MoEMethod), (
                f"Expected QuarkW8A8Int8MoEMethod, got {type(moe._quant_method)}"
            )
            # Non-MoE linear layers should use QuarkW8A8Int8
            qkv_proj = layer.self_attn.qkv_proj
            assert isinstance(qkv_proj.scheme, QuarkW8A8Int8)

        llm.apply_model(check_model)

        output = llm.generate_greedy("Hello", max_tokens=4)
        assert output


@pytest.mark.parametrize("tp", [1])
def test_quark_fp8_w8a8_per_block_moe(vllm_runner, tp):
    """Test per-block (128x128) FP8 MoE quantization with a tiny Qwen3 MoE model."""
    model_path = "Adamji/tiny-qwen3-moe-fp8-per-block"
    with vllm_runner(
        model_path,
        enforce_eager=True,
        tensor_parallel_size=tp,
        gpu_memory_utilization=0.1,
    ) as llm:

        def check_model(model):
            experts = model.model.layers[0].mlp.experts
            method = experts._quant_method
            assert isinstance(method, QuarkW8A8Fp8MoEMethod), (
                f"Expected QuarkW8A8Fp8MoEMethod, got {type(method)}"
            )
            assert method.weight_qscheme == "per_block"
            assert method.weight_block_size == [128, 128]

            # hidden_size=128, moe_intermediate_size=256 and 4 experts, so one
            # scale per 128x128 tile, with w13 stacking gate on top of up.
            routed_experts = experts.routed_experts
            assert routed_experts.w13_weight_scale.shape == (4, 4, 1)
            assert routed_experts.w2_weight_scale.shape == (4, 1, 2)
            # Quark exports the block scales under the same name as the
            # per-tensor and per-channel schemes.
            assert not hasattr(routed_experts, "w13_weight_scale_inv")
            assert not hasattr(routed_experts, "w2_weight_scale_inv")

        llm.apply_model(check_model)

        output = llm.generate_greedy("Hello", max_tokens=4)
        assert output


@pytest.mark.skipif(
    not (on_gfx950() or on_gfx942()),
    reason="Quark W4A8 (INT4-FP8) MoE requires the AITER kernel on gfx942/gfx950",
)
def test_quark_w4a8_fp8_moe(monkeypatch, dist_init, workspace_init):
    """Test W4A8 (INT4 weight + FP8 activation) MoE with a tiny Qwen3 MoE model.

    W4A8 dispatches through the AITER fused MoE kernel, so AITER must be on.
    """
    monkeypatch.setenv("VLLM_ROCM_USE_AITER", "1")
    monkeypatch.setenv("VLLM_ROCM_USE_AITER_MOE", "1")
    rocm_aiter_ops.refresh_env_variables()

    model_path = "amd/tiny-qwen3-moe-w4a8"
    model, vllm_config = load_model_without_vllm_runner(
        model_path,
    )
    with set_current_vllm_config(vllm_config):
        moe = model.model.layers[0].mlp.experts
        assert isinstance(moe._quant_method, QuarkW4A8Fp8MoEMethod), (
            f"Expected QuarkW4A8Fp8MoEMethod, got {type(moe._quant_method)}"
        )

        monkeypatch.setattr(Attention, "forward", lambda _, q, k, v: q.contiguous())
        input_ids = torch.tensor([1, 2, 3, 4], device=DEVICE_TYPE)
        positions = torch.arange(input_ids.numel(), device=DEVICE_TYPE)
        with set_forward_context(None, vllm_config, num_tokens=input_ids.numel()):
            hidden_states = model(input_ids, positions, None)
            logits = model.compute_logits(hidden_states)
        assert torch.isfinite(logits).all()


def test_quark_fp8_parity(dist_init, workspace_init):
    quark_model_id = "amd-quark/llama-tiny-fp8-quark-quant-method"
    fp8_model_id = "amd-quark/llama-tiny-fp8-quant-method"

    def load_state_dict(model_id: str) -> dict[str, torch.Tensor]:
        model, _ = load_model_without_vllm_runner(model_id)
        return {k: v.cpu() for k, v in model.state_dict().items()}

    quark_state_dict = load_state_dict(quark_model_id)
    fp8_state_dict = load_state_dict(fp8_model_id)

    assert fp8_state_dict.keys() == quark_state_dict.keys()

    for key in fp8_state_dict:
        assert torch.equal(fp8_state_dict[key], quark_state_dict[key])


@dataclass
class AccuracyTestConfig:
    model_name: str
    excepted_value: float

    def get_model_args(
        self,
        tp_size: int,
        model_max_len: int | None = None,
        kwargs: dict | None = None,
    ) -> dict:
        if kwargs is None:
            kwargs = {}

        model_args = {
            "pretrained": self.model_name,
            "dtype": "auto",
            "add_bos_token": True,
            "tensor_parallel_size": tp_size,
            "gpu_memory_utilization": 0.7,
            **kwargs,
        }
        if model_max_len is not None:
            model_args["max_model_len"] = model_max_len

        return model_args


WIKITEXT_ACCURACY_CONFIGS = [
    AccuracyTestConfig(
        model_name="fxmarty/qwen1.5_moe_a2.7b_chat_w_fp4_a_fp6_e2m3",
        excepted_value=11.3,
    ),
    AccuracyTestConfig(
        model_name="fxmarty/qwen1.5_moe_a2.7b_chat_w_fp6_e3m2_a_fp6_e3m2",
        excepted_value=10.6,
    ),
]


@pytest.mark.skipif(
    not QUARK_MXFP4_AVAILABLE,
    reason=f"amd-quark>={QUARK_MXFP4_MIN_VERSION} is not available",
)
@pytest.mark.parametrize(
    "config", WIKITEXT_ACCURACY_CONFIGS, ids=lambda config: config.model_name
)
@pytest.mark.parametrize("tp_size", [1, 2])
def test_ocp_mx_wikitext_correctness(config: AccuracyTestConfig, tp_size: int):
    device_count = torch.accelerator.device_count()
    if device_count < tp_size:
        pytest.skip(f"This test requires >={tp_size} gpus, got only {device_count}")

    results = lm_eval.simple_evaluate(
        model="vllm",
        model_args=config.get_model_args(
            tp_size=tp_size, kwargs={"cudagraph_capture_sizes": [16]}
        ),
        tasks="wikitext",
        batch_size=64,
    )

    measured_value = results["results"]["wikitext"]["word_perplexity,none"]
    assert measured_value == pytest.approx(config.excepted_value, abs=0.1)


GSM8K_ACCURACY_CONFIGS = [
    # Private model.
    AccuracyTestConfig(
        model_name="amd/DeepSeek-R1-WMXFP4-AMXFP4-Scale-UINT8-MoE-Quant",
        excepted_value=0.96,
    ),
]


@pytest.mark.parametrize("config", GSM8K_ACCURACY_CONFIGS)
@pytest.mark.skipif(
    not QUARK_MXFP4_AVAILABLE,
    reason=f"amd-quark>={QUARK_MXFP4_MIN_VERSION} is not available",
)
@pytest.mark.skipif(
    not HF_HUB_AMD_ORG_ACCESS,
    reason="Read access to huggingface.co/amd is required for this test.",
)
def test_mxfp4_gsm8k_correctness(config: AccuracyTestConfig):
    device_count = torch.accelerator.device_count()
    if device_count < 8:
        pytest.skip(f"This test requires >=8 gpus, got only {device_count}")

    task = "gsm8k"
    rtol = 0.03

    results = lm_eval.simple_evaluate(
        model="vllm",
        model_args=config.get_model_args(tp_size=8, model_max_len=38768),
        tasks=task,
        batch_size=64,
        num_fewshot=8,
    )

    EXPECTED_VALUE = config.excepted_value
    measured_value = results["results"][task]["exact_match,strict-match"]
    assert (
        measured_value - rtol < EXPECTED_VALUE
        and measured_value + rtol > EXPECTED_VALUE
    ), f"Expected: {EXPECTED_VALUE} |  Measured: {measured_value}"


@pytest.mark.skipif(
    not QUARK_MXFP4_AVAILABLE,
    reason=f"amd-quark>={QUARK_MXFP4_MIN_VERSION} is not available",
)
@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_fused_qdq_match_quark(float_dtype: torch.dtype, scalings: list[int]):
    torch.manual_seed(0)

    hidden_size = 64 * 32
    inp = (torch.rand(1, 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)]
        )

    inp_kernel = inp.clone()
    inp_kernel_clone = inp_kernel.clone()

    res_hip = mx_kernel.qdq_mxfp4_hip(inp_kernel_clone, "even")
    res_torch = qdq_mxfp4_torch(inp_kernel, "even")

    for i in range(hidden_size // 32):
        assert torch.all(torch.isfinite(res_hip[:, i * 32 : (i + 1) * 32]))
        assert torch.all(torch.isfinite(res_torch[:, i * 32 : (i + 1) * 32]))

        torch.testing.assert_close(
            res_hip[:, i * 32 : (i + 1) * 32], res_torch[:, i * 32 : (i + 1) * 32]
        )


@pytest.mark.skipif(
    not QUARK_MXFP4_AVAILABLE,
    reason=f"amd-quark>={QUARK_MXFP4_MIN_VERSION} is not available",
)
@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_dequant_kernel_match_quark(
    float_dtype: torch.dtype, scalings: list[int]
):
    qspec = FP4PerGroupSpec(
        ch_axis=-1,
        group_size=32,
        scale_format="e8m0",
        scale_calculation_mode="even",
        is_dynamic=False,
    ).to_quantization_spec()

    weight_quantizer = StaticScaledRealQuantizer(
        qspec=qspec,
        quantizer=None,
        reorder=False,
        real_quantized=True,
        float_dtype=float_dtype,
        device=DEVICE_TYPE,
    )

    observer = qspec.observer_cls(qspec, device=DEVICE_TYPE)

    hidden_size = 512
    shape = (11008, hidden_size)

    w = (torch.rand(shape, device=DEVICE_TYPE, dtype=float_dtype) - 0.5) * 2

    # Make it so that different groups have different scales.
    for i in range(hidden_size // 32):
        w[:, i * 32 : (i + 1) * 32] = (
            w[:, i * 32 : (i + 1) * 32] * scalings[i % len(scalings)]
        )

    observer(w)
    scale, _ = observer._calculate_qparams()
    weight_quantizer.scale = scale

    w_mxfp4 = weight_quantizer.to_real_quantize_params(w).to(DEVICE_TYPE)
    weight_quantizer.maybe_convert_and_transpose_scale()

    scale = weight_quantizer.scale

    out_hip = mx_kernel.dq_mxfp4_hip(w_mxfp4, scale, float_dtype)

    out_torch = dq_mxfp4_torch(w_mxfp4, scale, float_dtype)

    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 = {
    "qkv_proj": ["q_proj", "k_proj", "v_proj"],
    "gate_up_proj": ["gate_proj", "up_proj"],
}


def test_quark_should_ignore_layer_checks_children():
    assert should_ignore_layer(
        "model.layers.78.mlp.experts",
        ["model.layers.78.mlp.experts.0.down_proj"],
        check_children=True,
    )


def test_quark_should_ignore_layer_rejects_partial_fused_matches():
    with pytest.raises(ValueError, match="different quantization schemes"):
        should_ignore_layer(
            "model.layers.0.self_attn.qkv_proj",
            ["model.layers.0.self_attn.q_proj"],
            FUSED_MAPPING,
        )


def test_fused_name_listed_directly_is_skipped():
    # Regression for Step-3.5-Flash-FP8: the checkpoint lists the fused
    # name (``qkv_proj``) directly in ``modules_to_not_convert``. When a
    # ``packed_modules_mapping`` is registered on the model, the fused
    # match must still win over per-shard expansion.
    ignored = ["model.layers.0.self_attn.qkv_proj"]
    assert is_layer_skipped(
        prefix="model.layers.0.self_attn.qkv_proj",
        ignored_layers=ignored,
        fused_mapping=FUSED_MAPPING,
    )
    assert is_layer_skipped(
        prefix="model.layers.0.mlp.gate_up_proj",
        ignored_layers=["model.layers.0.mlp.gate_up_proj"],
        fused_mapping=FUSED_MAPPING,
    )


def test_unfused_shards_listed_is_skipped():
    # Quark INT8 style: per-shard names listed; all shards present means
    # the fused layer is skipped via expansion.
    ignored = [
        "model.layers.0.self_attn.q_proj",
        "model.layers.0.self_attn.k_proj",
        "model.layers.0.self_attn.v_proj",
    ]
    assert is_layer_skipped(
        prefix="model.layers.0.self_attn.qkv_proj",
        ignored_layers=ignored,
        fused_mapping=FUSED_MAPPING,
    )


def test_partial_shards_raises():
    # Only some shards listed -> ambiguous, must raise. Fused name is
    # not in ignored_layers, so we fall through to per-shard expansion.
    ignored = ["model.layers.0.self_attn.q_proj"]
    with pytest.raises(ValueError):
        is_layer_skipped(
            prefix="model.layers.0.self_attn.qkv_proj",
            ignored_layers=ignored,
            fused_mapping=FUSED_MAPPING,
        )


def test_not_skipped_when_nothing_listed():
    assert not is_layer_skipped(
        prefix="model.layers.0.self_attn.qkv_proj",
        ignored_layers=["model.layers.0.mlp.gate_up_proj"],
        fused_mapping=FUSED_MAPPING,
    )


def test_non_fused_layer_unaffected():
    assert is_layer_skipped(
        prefix="model.layers.0.self_attn.o_proj",
        ignored_layers=["model.layers.0.self_attn.o_proj"],
        fused_mapping=FUSED_MAPPING,
    )
    assert not is_layer_skipped(
        prefix="model.layers.0.self_attn.o_proj",
        ignored_layers=["model.layers.1.self_attn.o_proj"],
        fused_mapping=FUSED_MAPPING,
    )


def test_substr_match_on_fused_name():
    # Substring matching: a fused-name match should also
    # short-circuit before shard expansion.
    assert is_layer_skipped(
        prefix="model.layers.0.self_attn.qkv_proj",
        ignored_layers=["self_attn.qkv_proj"],
        fused_mapping=FUSED_MAPPING,
        match_mode="substring",
    )


@pytest.mark.parametrize(
    ("prefix", "ignored_layer", "expected"),
    [
        ("model.layers.0.self_attn.b_proj", "b_proj", True),
        ("model.layers.0.self_attn.q_b_proj", "b_proj", False),
        ("model.layers.0.self_attn.kv_b_proj", "b_proj", False),
        ("model.layers.5.self_attn.g_proj", "5.self_attn.g_proj", True),
        ("model.layers.6.self_attn.g_proj", "5.self_attn.g_proj", False),
    ],
)
def test_suffix_match_at_module_boundary(prefix, ignored_layer, expected):
    assert (
        is_layer_skipped(
            prefix=prefix,
            ignored_layers=[ignored_layer],
            match_mode="suffix",
        )
        is expected
    )


_GLM5_MXFP4_WEIGHT = {
    "dtype": "fp4",
    "qscheme": "per_group",
    "group_size": 32,
    "scale_format": "e8m0",
    "is_dynamic": False,
}
_GLM5_BLOCK_FP8_WEIGHT = {
    "dtype": "fp8_e4m3",
    "qscheme": "per_block",
    "is_dynamic": False,
    "block_size": [128, 128],
    "symmetric": True,
}
_GLM5_BLOCK_FP8_INPUT = {
    "dtype": "fp8_e4m3",
    "qscheme": "per_group",
    "is_dynamic": True,
    "group_size": 128,
    "symmetric": True,
}
_GLM5_GATE_UP = "language_model.model.layers.0.mlp.gate_up_proj"


def _glm5_mixed_precision_config() -> QuarkConfig:
    """Global MXFP4 with the dense-MLP gate/up shards marked block-FP8."""
    fp8 = {"weight": _GLM5_BLOCK_FP8_WEIGHT, "input_tensors": _GLM5_BLOCK_FP8_INPUT}
    return QuarkConfig(
        {
            "global_quant_config": {
                "weight": _GLM5_MXFP4_WEIGHT,
                "input_tensors": None,
            },
            "layer_type_quant_config": {},
            "layer_quant_config": {"*mlp.gate_proj": fp8, "*mlp.up_proj": fp8},
            "exclude": [],
        }
    )


class _GLM5RecordingParam:
    """Stand-in for a BF16 target param that records what gets loaded."""

    def __init__(self):
        self.loaded_weight: torch.Tensor | None = None
        self.loaded_shard: object = "unset"

    def weight_loader(self, param, loaded_weight, shard_id=None):
        self.loaded_weight = loaded_weight
        self.loaded_shard = shard_id


def _glm5_block_fp8(out_dim: int, in_dim: int, block: int = 128):
    """Return a (fp8 weight, f32 per-block scale) pair."""
    weight = (torch.randn(out_dim, in_dim) * 0.1).to(torch.float8_e4m3fn)
    scale = torch.rand(out_dim // block, in_dim // block, dtype=torch.float32) + 0.5
    return weight, scale


def test_glm5next_gate_up_proj_mapping_is_genuine_fusion():
    from vllm.models.glm5next.common.model import Glm5NextForConditionalGeneration

    mapping = Glm5NextForConditionalGeneration.packed_modules_mapping
    assert mapping["gate_up_proj"] == ["gate_proj", "up_proj"]


def test_glm5next_genuine_fusion_resolves_gate_up_to_block_fp8():
    # The fused module must expand to its real shard names so each resolves to
    # the per-layer block-FP8 entry (rather than the global MXFP4 scheme).
    config = _glm5_mixed_precision_config()
    config.packed_modules_mapping = {"gate_up_proj": ["gate_proj", "up_proj"]}
    _, _, scheme_cls = config.get_scheme_cls(LinearBase, _GLM5_GATE_UP)
    assert scheme_cls is QuarkW8A8Fp8PerBlock


def test_glm5next_attn_loader_accepts_quark_weight_scale():
    # Regression for KeyError on '...kv_a_proj_with_mqa.weight_scale': the fused
    # q_a/kv_a projection is kept BF16 and dequantized on load; the Quark scale
    # name must be recognized and routed to the correct fused shard.
    from vllm.models.glm5next.common.model import (
        _dequant_fp8_block,
        _try_load_fp8_attn_proj,
    )

    prefix = "layers.3.self_attn"
    target = _GLM5RecordingParam()
    params_dict = {f"{prefix}.fused_qkv_a_proj.weight": target}
    buf: dict = {}
    loaded: set = set()

    weight, scale = _glm5_block_fp8(256, 256)
    # fp8 weight arrives first -> buffered, nothing loaded yet.
    assert (
        _try_load_fp8_attn_proj(
            f"{prefix}.kv_a_proj_with_mqa.weight", weight, buf, params_dict, loaded, 0
        )
        is True
    )
    assert target.loaded_weight is None
    # Quark-named scale completes the pair -> dequantize + load.
    assert (
        _try_load_fp8_attn_proj(
            f"{prefix}.kv_a_proj_with_mqa.weight_scale",
            scale,
            buf,
            params_dict,
            loaded,
            0,
        )
        is True
    )
    assert target.loaded_weight is not None
    assert target.loaded_weight.dtype == torch.bfloat16
    assert target.loaded_shard == 1  # kv_a is shard 1 of fused_qkv_a_proj
    assert torch.equal(target.loaded_weight, _dequant_fp8_block(weight, scale, 128))
    assert f"{prefix}.fused_qkv_a_proj.weight" in loaded


def test_glm5next_attn_loader_accepts_deepseek_weight_scale_inv():
    # DeepSeek scale name keeps working unchanged.
    from vllm.models.glm5next.common.model import (
        _dequant_fp8_block,
        _try_load_fp8_attn_proj,
    )

    prefix = "layers.7.self_attn"
    target = _GLM5RecordingParam()
    params_dict = {f"{prefix}.o_proj.weight": target}
    buf: dict = {}
    loaded: set = set()

    weight, scale = _glm5_block_fp8(128, 256)
    _try_load_fp8_attn_proj(
        f"{prefix}.o_proj.weight", weight, buf, params_dict, loaded, 0
    )
    assert (
        _try_load_fp8_attn_proj(
            f"{prefix}.o_proj.weight_scale_inv", scale, buf, params_dict, loaded, 0
        )
        is True
    )
    assert target.loaded_shard is None  # o_proj is a direct (non-fused) proj
    assert torch.equal(target.loaded_weight, _dequant_fp8_block(weight, scale, 128))


_REVERSE_AWQ_PACK_ORDER = [0, 4, 1, 5, 2, 6, 3, 7]


def _quark_int4_config(
    *,
    pack_method: str = "reorder",
    symmetric: bool = True,
    exclude: list[str] | None = None,
) -> dict:
    return {
        "quant_method": "quark",
        "export": {"pack_method": pack_method, "kv_cache_group": []},
        "global_quant_config": {
            "weight": {
                "dtype": "int4",
                "group_size": 128,
                "symmetric": symmetric,
            }
        },
        "exclude": exclude or [],
    }


def _sign_extend_int4_nibbles(t: torch.Tensor) -> torch.Tensor:
    mask = (t & 0x8).bool()
    t = t.clone()
    t[mask] = t[mask] | 0xF0
    return t


def _dequantize_quark_signed_awq_torch(
    qweight: torch.Tensor,
    scales: torch.Tensor,
    qzeros: torch.Tensor,
    group_size: int,
    *,
    pack_reorder: bool = True,
) -> torch.Tensor:
    bits = 4
    shifts = torch.arange(0, 32, bits, device=qweight.device)
    iweights = ((qweight[:, :, None] >> shifts[None, None, :]) & 0xF).to(torch.int8)
    iweights = iweights.view(qweight.shape[0], -1)
    zeros = ((qzeros[:, :, None] >> shifts[None, None, :]) & 0xF).to(torch.int8)
    zeros = zeros.view(qzeros.shape[0], -1)

    if pack_reorder:
        order = torch.tensor(_REVERSE_AWQ_PACK_ORDER, device=qweight.device)
    else:
        order = torch.arange(8, device=qweight.device)
    iweights = iweights.view(qweight.shape[0], -1, 8)[:, :, order].reshape(
        qweight.shape[0], -1
    )
    zeros = zeros.view(qzeros.shape[0], -1, 8)[:, :, order].reshape(qzeros.shape[0], -1)
    iweights = _sign_extend_int4_nibbles(iweights & 0xF)
    zeros = _sign_extend_int4_nibbles(zeros & 0xF)

    scales = scales.repeat_interleave(group_size, dim=0)
    zeros = zeros.repeat_interleave(group_size, dim=0)
    return (iweights - zeros) * scales


def _dequantize_awq_unsigned_torch(
    qweight: torch.Tensor,
    scales: torch.Tensor,
    qzeros: torch.Tensor,
    group_size: int,
    *,
    pack_reorder: bool = True,
) -> torch.Tensor:
    bits = 4
    shifts = torch.arange(0, 32, bits, device=qweight.device)
    iweights = ((qweight[:, :, None] >> shifts[None, None, :]) & 0xF).to(torch.int8)
    iweights = iweights.view(qweight.shape[0], -1)
    zeros = ((qzeros[:, :, None] >> shifts[None, None, :]) & 0xF).to(torch.int8)
    zeros = zeros.view(qzeros.shape[0], -1)

    if pack_reorder:
        order = torch.tensor(_REVERSE_AWQ_PACK_ORDER, device=qweight.device)
    else:
        order = torch.arange(8, device=qweight.device)
    iweights = iweights.view(qweight.shape[0], -1, 8)[:, :, order].reshape(
        qweight.shape[0], -1
    )
    zeros = zeros.view(qzeros.shape[0], -1, 8)[:, :, order].reshape(qzeros.shape[0], -1)

    scales = scales.repeat_interleave(group_size, dim=0)
    zeros = zeros.repeat_interleave(group_size, dim=0)
    return (iweights - zeros) * scales


def _pack_int4_nibbles(nibbles: torch.Tensor, *, pack_reorder: bool) -> torch.Tensor:
    pack_order = (
        torch.tensor(_REVERSE_AWQ_PACK_ORDER, device=nibbles.device)
        if pack_reorder
        else torch.arange(8, device=nibbles.device)
    )
    shifts = pack_order * 4
    return ((nibbles.to(torch.int64) & 0xF) << shifts).sum(dim=-1).to(torch.int32)


class TestQuarkInt4Format:
    """Tests for Quark INT4 export format compatibility."""

    def test_quark_order_pack_method_uses_native_int4_scheme(self):
        quant_config = QuarkConfig.from_config(_quark_int4_config(pack_method="order"))
        weight_config = quant_config.quant_config["global_quant_config"]
        weight_key, act_key, scheme_cls = quant_config._get_scheme_cls_from_config(
            weight_config
        )
        assert scheme_cls is QuarkW4A16Int4

        scheme = quant_config.init_scheme(
            scheme_cls,
            weight_quant_key=weight_key,
            activation_quant_key=act_key,
            weight_config=weight_config["weight"],
        )
        assert isinstance(scheme, QuarkW4A16Int4)
        assert not scheme.pack_reorder

    def test_quark_int4_scheme_supports_asymmetric_weights(self):
        quant_config = QuarkConfig.from_config(_quark_int4_config(symmetric=False))
        weight_config = quant_config.quant_config["global_quant_config"]
        weight_key, act_key, scheme_cls = quant_config._get_scheme_cls_from_config(
            weight_config
        )
        assert scheme_cls is QuarkW4A16Int4

        scheme = quant_config.init_scheme(
            scheme_cls,
            weight_quant_key=weight_key,
            activation_quant_key=act_key,
            weight_config=weight_config["weight"],
        )
        assert isinstance(scheme, QuarkW4A16Int4)
        assert not scheme.is_symmetric

    @pytest.mark.parametrize("missing_field", ["group_size", "symmetric"])
    def test_quark_int4_scheme_requires_weight_config_fields(self, missing_field):
        from vllm.model_executor.layers.quantization.utils.quant_utils import (
            kInt4Static,
        )

        weight_config = {
            "dtype": "int4",
            "group_size": 128,
            "symmetric": True,
        }
        weight_config.pop(missing_field)
        quant_config = QuarkConfig.from_config(
            {
                "quant_method": "quark",
                "export": {"pack_method": "reorder", "kv_cache_group": []},
                "global_quant_config": {"weight": weight_config},
                "exclude": [],
            }
        )

        with pytest.raises(ValueError, match=missing_field):
            quant_config.init_scheme(
                QuarkW4A16Int4,
                weight_quant_key=kInt4Static,
                activation_quant_key=None,
                weight_config=weight_config,
            )

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

        from vllm.model_executor.layers.quantization.utils.quant_utils import (
            kInt4StaticAsym,
        )

        quant_config = QuarkConfig.from_config(_quark_int4_config(symmetric=False))
        moe_config = type("MoeConfig", (), {})()
        moe_config.has_bias = False

        mock_backend = MagicMock()
        mock_experts_cls = MagicMock()
        with patch(
            "vllm.model_executor.layers.quantization.quark.quark_moe"
            ".select_wna16_moe_backend",
            return_value=(mock_backend, mock_experts_cls),
        ):
            method = QuarkW4A16Int4MoEMethod(
                kInt4StaticAsym,
                None,
                quant_config.quant_config["global_quant_config"]["weight"],
                quant_config.pack_method,
                moe_config,
            )

        assert method.group_size == 128
        assert method.pack_reorder
        assert method.use_wna16_backend

    def test_triton_wna16_experts_supports_asymmetric_int4(self):
        from vllm.model_executor.layers.fused_moe.experts.triton_moe import (
            TritonWNA16Experts,
        )
        from vllm.model_executor.layers.quantization.utils.quant_utils import (
            kInt4Static,
            kInt4Static32,
            kInt4Static32Asym,
            kInt4StaticAsym,
        )

        for key in [kInt4Static, kInt4Static32, kInt4StaticAsym, kInt4Static32Asym]:
            assert TritonWNA16Experts._supports_quant_scheme(key, None), (
                f"{key} should be supported by TritonWNA16Experts"
            )

    def test_quark_moe_loader_shards_w2_zero_point_across_tp_ranks(self):
        """w2 zero-points are sharded along the intermediate dim per TP rank.

        Loading the same Quark checkpoint tensor at tensor_parallel_size=2 for
        both ranks must jointly reconstruct the tensor_parallel_size=1 (full)
        result, with no overlap or gap. This guards the custom loader's TP
        split for the packed zero-point layout.
        """

        class _FakeMoEConfig:
            has_bias = False

            def __init__(self, tp_size, tp_rank):
                self.tp_size = tp_size
                self.tp_rank = tp_rank

        class _FakeLayer:
            def __init__(self, moe_config):
                self.moe_config = moe_config
                self.intermediate_size_per_partition = 64
                self.group_size_div_factor = 1

        def load_w2_zero_point(raw_zp, tp_size, tp_rank):
            from unittest.mock import MagicMock, patch

            from vllm.model_executor.layers.quantization.utils.quant_utils import (
                kInt4StaticAsym,
            )

            quant_config = QuarkConfig.from_config(_quark_int4_config(symmetric=False))
            moe_config = _FakeMoEConfig(tp_size, tp_rank)
            with patch(
                "vllm.model_executor.layers.quantization.quark.quark_moe"
                ".select_wna16_moe_backend",
                return_value=(MagicMock(), MagicMock()),
            ):
                method = QuarkW4A16Int4MoEMethod(
                    kInt4StaticAsym,
                    None,
                    quant_config.quant_config["global_quant_config"]["weight"],
                    quant_config.pack_method,
                    moe_config,
                )

            layer = _FakeLayer(moe_config)
            # The loader converts the packed zero-point, then writes the
            # per-rank slice into param.data[expert_id]. A single expert dim is
            # enough; the last dim shrinks by tp_size after sharding.
            param = torch.nn.Parameter(
                torch.zeros(1, 16, 8 // tp_size, dtype=torch.uint8),
                requires_grad=False,
            )
            loader = method.get_weight_loader(
                layer,
                weight_loader=None,
            )
            loader(
                param,
                raw_zp.clone(),
                weight_name="w2_weight_zero_point",
                shard_id="w2",
                expert_id=0,
            )
            return param.data[0]

        # Quark-packed zero-point for one expert: rows = hidden // pack_factor,
        # cols carry the (packed) intermediate dim, split evenly across ranks.
        raw_zp = torch.randint(0, 256, (8, 16), dtype=torch.uint8)

        full = load_w2_zero_point(raw_zp, tp_size=1, tp_rank=0)
        shard0 = load_w2_zero_point(raw_zp, tp_size=2, tp_rank=0)
        shard1 = load_w2_zero_point(raw_zp, tp_size=2, tp_rank=1)

        expected = full.view(full.size(0), 2, -1)
        torch.testing.assert_close(shard0, expected[:, 0])
        torch.testing.assert_close(shard1, expected[:, 1])

    def test_quark_shared_expert_gate_keeps_quantized_tensors(self):
        quant_config = QuarkConfig.from_config(_quark_int4_config())
        mapper = quant_config.get_cache_scale_mapper()

        output_names = {
            name
            for name, _ in mapper.apply(
                [
                    (
                        "model.language_model.layers.0.mlp.shared_expert_gate.weight",
                        torch.zeros(1),
                    ),
                    (
                        "model.language_model.layers.0.mlp.shared_expert_gate"
                        ".weight_scale",
                        torch.zeros(1),
                    ),
                    (
                        "model.language_model.layers.0.mlp.shared_expert_gate"
                        ".weight_zero_point",
                        torch.zeros(1),
                    ),
                ]
            )
        }

        assert (
            "model.language_model.layers.0.mlp.shared_expert_gate.weight"
            in output_names
        )
        assert (
            "model.language_model.layers.0.mlp.shared_expert_gate.weight_scale"
            in output_names
        )
        assert (
            "model.language_model.layers.0.mlp.shared_expert_gate.weight_zero_point"
            in output_names
        )

    def test_quark_apply_mapper_updates_exclude_and_layer_quant_config(self):
        quant_config = QuarkConfig.from_config(
            {
                **_quark_int4_config(),
                "exclude": ["lm_head"],
                "layer_quant_config": {
                    "model.language_model.layers.0.mlp.gate": {
                        "weight": {"dtype": "float16"},
                    },
                },
            }
        )
        quant_config.apply_vllm_mapper(
            WeightsMapper(
                orig_to_new_prefix={
                    "lm_head": "language_model.lm_head",
                    "model.language_model.": "language_model.model.",
                },
            )
        )

        layer_quant_config = quant_config.quant_config["layer_quant_config"]
        assert "language_model.model.layers.0.mlp.gate" in layer_quant_config
        assert "model.language_model.layers.0.mlp.gate" not in layer_quant_config
        assert quant_config.quant_config["exclude"] == [
            "language_model.lm_head",
        ]

    def test_quark_apply_mapper_ignores_non_string_list_entries(self):
        quant_config = QuarkConfig.from_config(
            {
                **_quark_int4_config(),
                "algo_config": [
                    {"name": "qronos", "inside_layer_modules": ["self_attn.q_proj"]}
                ],
            }
        )
        quant_config.apply_vllm_mapper(WeightsMapper())

        assert quant_config.quant_config["algo_config"] == [
            {"name": "qronos", "inside_layer_modules": ["self_attn.q_proj"]}
        ]

    def test_quark_exclude_matches_via_mapper_prefix(self):
        quant_config = QuarkConfig.from_config(_quark_int4_config(exclude=["lm_head"]))
        quant_config.apply_vllm_mapper(
            WeightsMapper(orig_to_new_prefix={"lm_head": "language_model.lm_head"})
        )

        exclude_layers = quant_config.quant_config["exclude"]
        assert should_ignore_layer("language_model.lm_head", ignore=exclude_layers)
        assert not should_ignore_layer(
            "language_model.model.layers.0.mlp.gate", ignore=exclude_layers
        )


@pytest.mark.parametrize("symmetric", [False, True])
@pytest.mark.parametrize("pack_method", ["order", "reorder"])
def test_quark_int4_canonicalizes_pack_for_kernel_layout(pack_method, symmetric):
    """Quark pack order is normalized to the layout expected by awq_* ops.

    This reuses the existing AWQ dequant/gemm kernels for compute only; loading
    still goes through the native Quark quantization path, not AutoAWQ.
    """
    pack_reorder = pack_method == "reorder"
    group_size = 2
    packed_values = torch.tensor(
        [
            [0, 1, 7, 8, 9, 15, 2, 14],
            [15, 8, 0, 3, 12, 7, 1, 9],
        ],
        dtype=torch.int32,
    )
    packed_zeros = torch.zeros((1, 8), dtype=torch.int32)
    qweight = _pack_int4_nibbles(packed_values, pack_reorder=pack_reorder).view(2, 1)
    qzeros = _pack_int4_nibbles(packed_zeros, pack_reorder=pack_reorder).view(1, 1)
    scales = torch.ones((1, 8), dtype=torch.float16)

    dequantize = (
        _dequantize_quark_signed_awq_torch
        if symmetric
        else _dequantize_awq_unsigned_torch
    )
    expected = dequantize(
        qweight,
        scales,
        qzeros,
        group_size,
        pack_reorder=pack_reorder,
    )
    canonical_weight = canonicalize_quark_packed_int4(
        qweight,
        pack_reorder=pack_reorder,
        is_symmetric=symmetric,
    )
    canonical_zero = canonicalize_quark_packed_int4(
        qzeros,
        pack_reorder=pack_reorder,
        is_symmetric=symmetric,
    )
    actual = _dequantize_awq_unsigned_torch(
        canonical_weight, scales, canonical_zero, group_size
    )

    assert torch.equal(actual, expected)


# ---------------------------------------------------------------------------
# override_quantization_method: Quark always delegates to QuarkConfig
# ---------------------------------------------------------------------------

_QUARK_MXFP4_CFG = {
    "quant_method": "quark",
    "global_quant_config": {
        "weight": {"dtype": "fp4", "qscheme": "per_group", "group_size": 32},
    },
}

_QUARK_MXFP8_CFG = {
    "quant_method": "quark",
    "global_quant_config": {
        "weight": {
            "dtype": "fp8_e4m3",
            "qscheme": "per_block",
            "block_size": [32, 32],
            "scale_type": "float8_e8m0fnu",
            "symmetric": True,
        },
    },
}

_QUARK_MIXED_CFG = {
    "quant_method": "quark",
    "global_quant_config": {
        "weight": {"dtype": "fp4", "qscheme": "per_group", "group_size": 32},
    },
    "layer_quant_config": {"layers.0.attn": {"weight": {"dtype": "fp8_e4m3"}}},
}

_FP8_CFG = {"quant_method": "fp8"}

_DSV4_FP8_CFG = {"quant_method": "deepseek_v4_fp8"}


@pytest.mark.parametrize(
    "quant_config_cls",
    [
        "deepseek_v41",
        "deepseek_v4",
    ],
)
@pytest.mark.parametrize(
    "hf_quant_cfg, expect_none",
    [
        (_QUARK_MXFP4_CFG, True),
        (_QUARK_MXFP8_CFG, True),
        (_QUARK_MIXED_CFG, True),
        (_FP8_CFG, False),
        (_DSV4_FP8_CFG, False),
    ],
    ids=[
        "quark-mxfp4",
        "quark-mxfp8",
        "quark-mixed",
        "native-fp8",
        "deepseek-v4-fp8",
    ],
)
def test_quark_override_delegates_to_quark_config(
    quant_config_cls, hf_quant_cfg, expect_none
):
    """Any quant_method=='quark' must return None (delegate to QuarkConfig).

    Native FP8 configs should still be claimed by the model-specific config.
    """
    import importlib

    if quant_config_cls == "deepseek_v41":
        # Module was renamed from deepseek_v4_1 to deepseek_v41 upstream.
        for mod_name in (
            "vllm.models.deepseek_v41.quant_config",
            "vllm.models.deepseek_v4_1.quant_config",
        ):
            try:
                mod = importlib.import_module(mod_name)
                break
            except ModuleNotFoundError:
                continue
        else:
            pytest.skip("deepseek_v41 quant_config not found")
        model_type = "deepseek_v41"
    else:
        mod = importlib.import_module("vllm.models.deepseek_v4.quant_config")
        model_type = "deepseek_v4"

    config_cls = mod.DeepseekV4FP8Config
    hf_config = SimpleNamespace(model_type=model_type)
    result = config_cls.override_quantization_method(
        hf_quant_cfg, None, hf_config=hf_config
    )

    if expect_none:
        assert result is None, (
            f"Quark config should delegate to QuarkConfig (return None), got {result!r}"
        )
    else:
        assert result == "deepseek_v4_fp8", (
            f"Native FP8 should be claimed by DeepseekV4FP8Config, got {result!r}"
        )
