Skip to content

vllm.model_executor.layers.fused_moe.oracle.fp8

Functions:

_get_priority_backends(moe_config, weight_key, activation_key)

Get available backends in priority order based on platform and config.

This function can be extended to become more complex as needed.

Source code in vllm/model_executor/layers/fused_moe/oracle/fp8.py
def _get_priority_backends(
    moe_config: FusedMoEConfig,
    weight_key: QuantKey | None,
    activation_key: QuantKey | None,
) -> list[Fp8MoeBackend]:
    """Get available backends in priority order based on platform and config.

    This function can be extended to become more complex as needed.
    """
    _AVAILABLE_BACKENDS = [
        Fp8MoeBackend.AITER,
        Fp8MoeBackend.FLASHINFER_TRTLLM,
        Fp8MoeBackend.FLASHINFER_CUTLASS,
        Fp8MoeBackend.DEEPGEMM,
        Fp8MoeBackend.VLLM_CUTLASS,
        Fp8MoeBackend.TRITON,
        Fp8MoeBackend.MARLIN,
        Fp8MoeBackend.HUMMING,
        Fp8MoeBackend.BATCHED_DEEPGEMM,
        Fp8MoeBackend.BATCHED_VLLM_CUTLASS,
        Fp8MoeBackend.BATCHED_TRITON,
        Fp8MoeBackend.XPU,
        Fp8MoeBackend.CPU_W8A8,
        Fp8MoeBackend.CPU,
        Fp8MoeBackend.HPC,
    ]

    def _move_to_front(backends: list[Fp8MoeBackend], backend: Fp8MoeBackend) -> None:
        backends.insert(0, backends.pop(backends.index(backend)))

    # With DeepEP v2 contiguous layout (do_expand=False), tensors are
    # worst-case allocated with padding. TrtLLM's tile-level skipping
    # avoids wasted compute on padding rows; other backends process all rows.
    if (
        current_platform.is_cuda()
        and current_platform.is_device_capability_family(100)
        and moe_config.moe_parallel_config.use_deepep_v2_kernels
        and activation_key == kFp8Dynamic128Sym
        and weight_key == kFp8Static128BlockSym
    ):
        _move_to_front(_AVAILABLE_BACKENDS, Fp8MoeBackend.FLASHINFER_TRTLLM)

    # On Hopper for Block Fp8, prefer Triton for TP and FI CUTLASS for EP.
    if (
        current_platform.is_cuda()
        and current_platform.is_device_capability(90)
        and activation_key == kFp8Dynamic128Sym
        and weight_key == kFp8Static128BlockSym
    ):
        if moe_config.moe_parallel_config.ep_size > 1:
            _move_to_front(_AVAILABLE_BACKENDS, Fp8MoeBackend.FLASHINFER_CUTLASS)
        else:
            _move_to_front(_AVAILABLE_BACKENDS, Fp8MoeBackend.TRITON)

    if current_platform.is_xpu():
        # XPU platform supports TritonExperts and XPUExpertsFp8,
        # move XPU backend to the front.
        _move_to_front(_AVAILABLE_BACKENDS, Fp8MoeBackend.XPU)

    if current_platform.is_cpu():
        # W8A8 first: it falls through to the W8A16 backend whenever the
        # hardware (AMX-FP8) or the config isn't supported.
        _move_to_front(_AVAILABLE_BACKENDS, Fp8MoeBackend.CPU)
        _move_to_front(_AVAILABLE_BACKENDS, Fp8MoeBackend.CPU_W8A8)

    return _AVAILABLE_BACKENDS

_humming_fp8_weight_schema(layer, weight, weight_scale)

Build the humming weight schema from the canonical on-device fp8/mxfp8 tensors (scale dtype/shape, block size), not the producing quant method.

Source code in vllm/model_executor/layers/fused_moe/oracle/fp8.py
def _humming_fp8_weight_schema(
    layer: RoutedExperts, weight: torch.Tensor, weight_scale: torch.Tensor
) -> dict[str, Any]:
    """Build the humming weight schema from the canonical on-device fp8/mxfp8
    tensors (scale dtype/shape, block size), not the producing quant method."""
    # mxfp8: e8m0 group-32 scales (stored as uint8 bytes or e8m0). humming has
    # no compressed-tensors mxfp8 loader; its modelopt schema fits both sources.
    if weight_scale.dtype in (torch.uint8, torch.float8_e8m0fnu):
        return {"quant_method": "modelopt", "quant_algo": "mxfp8"}

    if hasattr(layer, "w13_weight_scale_inv"):
        assert hasattr(layer, "weight_block_size")
        return {"quant_method": "fp8", "weight_block_size": layer.weight_block_size}

    # fp8 (e4m3): recover the strategy from the scale layout (block from
    # weight_block_size, else channel vs tensor by per-expert scale count).
    config: dict[str, Any] = {
        "quant_method": "compressed-tensors",
        "format": "float-quantized",
        "type": "float",
        "num_bits": 8,
        "symmetric": True,
    }
    weight_block_size = getattr(layer, "weight_block_size", None)
    num_experts, num_output = weight.shape[0], weight.shape[-2]
    if weight_block_size is not None:
        config["strategy"] = "block"
        config["block_structure"] = list(weight_block_size)
    elif weight_scale.numel() >= num_experts * num_output:
        config["strategy"] = "channel"
    else:
        config["strategy"] = "tensor"
    return config

make_fp8_moe_quant_config(fp8_backend, w1_scale, w2_scale, a1_scale, a2_scale, w1_bias=None, w2_bias=None, block_shape=None, per_act_token_quant=False, per_out_ch_quant=False, swiglu_limit=None, gemm1_alpha=None, gemm1_beta=None, layer=None)

Create FusedMoEQuantConfig for the specified FP8 Backend. The FusedMoEQuantConfig holds the scales that are used at runtime by the Modular Kernel abstraction.

Note that certain kernels (e.g. Flashinfer CUTLASS) need special Quant configs to handle non-standard inputs to their kernel interfaces.

In a future PR, we will have this function should be a method of the modular kernel itself.

Source code in vllm/model_executor/layers/fused_moe/oracle/fp8.py
def make_fp8_moe_quant_config(
    fp8_backend: Fp8MoeBackend,
    w1_scale: torch.Tensor,
    w2_scale: torch.Tensor,
    a1_scale: torch.Tensor | None,
    a2_scale: torch.Tensor | None,
    w1_bias: torch.Tensor | None = None,
    w2_bias: torch.Tensor | None = None,
    block_shape: list[int] | None = None,
    per_act_token_quant: bool = False,
    per_out_ch_quant: bool = False,
    swiglu_limit: float | None = None,
    gemm1_alpha: float | None = None,
    gemm1_beta: float | None = None,
    layer: torch.nn.Module | None = None,
) -> FusedMoEQuantConfig:
    """Create FusedMoEQuantConfig for the specified FP8 Backend.
    The FusedMoEQuantConfig holds the scales that are used
    at runtime by the Modular Kernel abstraction.

    Note that certain kernels (e.g. Flashinfer CUTLASS) need
    special Quant configs to handle non-standard inputs to
    their kernel interfaces.

    In a future PR, we will have this function should be
    a method of the modular kernel itself.
    """
    if fp8_backend == Fp8MoeBackend.CPU_W8A8:
        return fp8_w8a8_moe_quant_config(
            w1_scale=w1_scale,
            w2_scale=w2_scale,
            a1_scale=a1_scale,
            block_shape=block_shape,
        )

    # MARLIN and CPU (W8A16) are mixed precision W8A16 configs.
    if fp8_backend == Fp8MoeBackend.MARLIN or fp8_backend == Fp8MoeBackend.CPU:
        return fp8_w8a16_moe_quant_config(
            w1_scale=w1_scale,
            w2_scale=w2_scale,
            w1_bias=w1_bias,
            w2_bias=w2_bias,
            block_shape=block_shape,
            gemm1_alpha=gemm1_alpha,
            gemm1_beta=gemm1_beta,
            gemm1_clamp_limit=swiglu_limit,
        )
    elif fp8_backend == Fp8MoeBackend.HUMMING:
        from vllm.model_executor.layers.fused_moe import RoutedExperts
        from vllm.model_executor.layers.quantization.utils.humming import (
            get_humming_moe_quant_config,
        )

        assert isinstance(layer, RoutedExperts)
        return get_humming_moe_quant_config(
            layer,
            gemm1_alpha=gemm1_alpha,
            gemm1_beta=gemm1_beta,
            gemm1_clamp_limit=swiglu_limit,
        )

    # Flashinfer CUTLASS or HPC per-tensor uses single dq scale
    # (alpha = w_scale * a_scale) and inverse a2 scale.
    if (
        fp8_backend in [Fp8MoeBackend.FLASHINFER_CUTLASS, Fp8MoeBackend.HPC]
        and block_shape is None
    ):
        assert a1_scale is not None and a2_scale is not None
        g1_alphas = w1_scale * a1_scale
        g2_alphas = w2_scale * a2_scale
        if layer is not None:
            layer.register_parameter(
                "g1_alphas", torch.nn.Parameter(g1_alphas, requires_grad=False)
            )
            layer.register_parameter(
                "g2_alphas", torch.nn.Parameter(g2_alphas, requires_grad=False)
            )
            g1_alphas = layer.g1_alphas
            g2_alphas = layer.g2_alphas
        return fp8_w8a8_moe_quant_config(
            w1_scale=w1_scale,
            w2_scale=w2_scale,
            w1_bias=w1_bias,
            w2_bias=w2_bias,
            a1_scale=a1_scale,
            a2_scale=a2_scale,
            a1_gscale=(1.0 / a1_scale),
            a2_gscale=(1.0 / a2_scale),
            g1_alphas=g1_alphas,
            g2_alphas=g2_alphas,
            gemm1_clamp_limit=swiglu_limit,
        )
    # MXFP8 (block [1, 32]) dispatches to the mxfp8 activation quant. Scales are
    # the non-swizzled (num_tokens, hidden_dim // 32) uint8 UE8M0 layout for all
    # backends; the DeepGEMM expert permute repacks them for the grouped GEMM.
    if block_shape == [1, 32]:
        return FusedMoEQuantConfig.make(
            "mxfp8",
            w1_scale=w1_scale,
            w2_scale=w2_scale,
            w1_bias=w1_bias,
            w2_bias=w2_bias,
            a1_scale=a1_scale,
            a2_scale=a2_scale,
            block_shape=block_shape,
            is_scale_swizzled=False,
            gemm1_alpha=gemm1_alpha,
            gemm1_beta=gemm1_beta,
            gemm1_clamp_limit=swiglu_limit,
        )

    # All other backends use normal config.
    return fp8_w8a8_moe_quant_config(
        w1_scale=w1_scale,
        w2_scale=w2_scale,
        w1_bias=w1_bias,
        w2_bias=w2_bias,
        a1_scale=a1_scale,
        a2_scale=a2_scale,
        block_shape=block_shape,
        per_act_token_quant=per_act_token_quant,
        per_out_ch_quant=per_out_ch_quant,
        gemm1_alpha=gemm1_alpha,
        gemm1_beta=gemm1_beta,
        gemm1_clamp_limit=swiglu_limit,
    )

map_fp8_backend(runner_backend)

Map user's MoEBackend to Fp8MoeBackend.

Source code in vllm/model_executor/layers/fused_moe/oracle/fp8.py
def map_fp8_backend(runner_backend: MoEBackend) -> Fp8MoeBackend:
    """Map user's MoEBackend to Fp8MoeBackend."""
    mapping = {
        "triton": Fp8MoeBackend.TRITON,
        "deep_gemm": Fp8MoeBackend.DEEPGEMM,
        "cutlass": Fp8MoeBackend.VLLM_CUTLASS,
        "flashinfer_trtllm": Fp8MoeBackend.FLASHINFER_TRTLLM,
        "flashinfer_cutlass": Fp8MoeBackend.FLASHINFER_CUTLASS,
        "marlin": Fp8MoeBackend.MARLIN,
        "humming": Fp8MoeBackend.HUMMING,
        "aiter": Fp8MoeBackend.AITER,
        "hpc": Fp8MoeBackend.HPC,
    }
    if backend := mapping.get(runner_backend):
        return backend
    raise ValueError(
        f"moe_backend='{runner_backend}' is not supported for FP8 MoE. "
        f"Expected one of {list(mapping.keys())}."
    )

pad_tp_shard_to_weight_blocks(config, weight_block_size)

Pad the TP shard to whole checkpoint blocks, keeping scales rank-local.

Source code in vllm/model_executor/layers/fused_moe/oracle/fp8.py
def pad_tp_shard_to_weight_blocks(
    config: FusedMoEConfig,
    weight_block_size: list[int],
) -> bool:
    """Pad the TP shard to whole checkpoint blocks, keeping scales rank-local."""
    block_n, block_k = weight_block_size
    if (
        block_n != block_k
        or config.tp_size == 1
        or config.intermediate_size_per_partition % block_n == 0
    ):
        return False
    if (
        config.intermediate_size % block_n != 0
        or config.hidden_dim % block_n != 0
        or config.ep_size != 1
        or config.is_lora_enabled
        or config.has_bias
    ):
        raise ValueError(
            f"Block-aligned FP8 TP sharding requires {block_n}-aligned "
            "global expert dimensions, pure TP, and no LoRA or "
            "expert bias."
        )
    config.tp_shard_with_padding = True
    config.intermediate_size_per_partition = round_up(
        config.intermediate_size_per_partition, block_n
    )
    return True

refine_fp8_moe_block_shape(config, weight_block_size)

Compute a refined block shape for block-quantized FP8 MoE weights whose checkpoint blocks cannot be sharded exactly across TP ranks.

TP shards the intermediate dim of the expert weights, so a per-shard size that is not a multiple of the checkpoint's block size makes the checkpoint's block scales impossible to shard exactly. When a finer block size (>= 32) divides both the checkpoint blocks and all involved dims, the weight scales can be refined to that granularity at load time (a lossless upsampling, since the refined block divides the checkpoint block). Only Triton-based kernels can consume the refined block shape: they take it as a runtime argument, while the other backends require the native 128x128 blocks. The refined shape is encoded in the QuantKey used for backend selection, so backends that only support 128x128 blocks are rejected by the oracle automatically.

Returns the refined [block_n, block_k] shape, or None if no refinement is needed or possible.

Source code in vllm/model_executor/layers/fused_moe/oracle/fp8.py
def refine_fp8_moe_block_shape(
    config: FusedMoEConfig,
    weight_block_size: list[int],
) -> list[int] | None:
    """Compute a refined block shape for block-quantized FP8 MoE weights whose
    checkpoint blocks cannot be sharded exactly across TP ranks.

    TP shards the intermediate dim of the expert weights, so a per-shard size
    that is not a multiple of the checkpoint's block size makes the
    checkpoint's block scales impossible to shard exactly. When a finer block
    size (>= 32) divides both the checkpoint blocks and all involved dims,
    the weight scales can be refined to that granularity at load time (a
    lossless upsampling, since the refined block divides the checkpoint
    block). Only Triton-based kernels can consume the refined block shape:
    they take it as a runtime argument, while the other backends require the
    native 128x128 blocks. The refined shape is encoded in the QuantKey used
    for backend selection, so backends that only support 128x128 blocks are
    rejected by the oracle automatically.

    Returns the refined [block_n, block_k] shape, or None if no refinement
    is needed or possible.
    """
    block_n, block_k = weight_block_size
    ispp = config.intermediate_size_per_partition
    if ispp % block_n == 0 and (config.tp_size == 1 or ispp % block_k == 0):
        return None
    refine = math.gcd(block_n, block_k, ispp, config.hidden_dim)
    if refine < 32:
        return None
    return [refine, refine]

resolve_fp8_moe_weight_block_shape(config, weight_block_size, activation_key, is_checkpoint_fp8_serialized)

Return the TP-adapted block shape and refine factor: refine if kernels allow, else pad to the TP shard.

Source code in vllm/model_executor/layers/fused_moe/oracle/fp8.py
def resolve_fp8_moe_weight_block_shape(
    config: FusedMoEConfig,
    weight_block_size: list[int],
    activation_key: QuantKey,
    is_checkpoint_fp8_serialized: bool,
) -> tuple[list[int], tuple[int, int] | None]:
    """Return the TP-adapted block shape and refine factor:
    refine if kernels allow, else pad to the TP shard."""
    refined_shape = refine_fp8_moe_block_shape(config, weight_block_size)
    if is_checkpoint_fp8_serialized and config.moe_backend != "auto":
        kernel_classes = backend_to_kernel_cls(map_fp8_backend(config.moe_backend))
        can_refine = refined_shape is not None and any(
            k_cls._supports_quant_scheme(
                create_fp8_quant_key(
                    static=True, group_shape=GroupShape(*refined_shape)
                ),
                activation_key,
            )
            for k_cls in kernel_classes
        )
        if not can_refine and pad_tp_shard_to_weight_blocks(config, weight_block_size):
            logger.info_once(
                "FP8 %s TP loading uses complete checkpoint blocks: "
                "local allocation %d, without weight requantization.",
                config.moe_backend,
                config.intermediate_size_per_partition,
            )
            return weight_block_size, None
    if refined_shape is None:
        return weight_block_size, None
    logger.info_once(
        "FP8 MoE block scales refined from %s to %s to fit "
        "the TP-sharded intermediate size %d.",
        str(weight_block_size),
        str(refined_shape),
        config.intermediate_size_per_partition,
    )
    return refined_shape, (
        weight_block_size[0] // refined_shape[0],
        weight_block_size[1] // refined_shape[1],
    )

select_fp8_moe_backend(config, weight_key, activation_key, allow_vllm_cutlass=False)

Select the primary FP8 MoE backend Note: Shape-specific fallbacks may still occur at runtime.

Source code in vllm/model_executor/layers/fused_moe/oracle/fp8.py
def select_fp8_moe_backend(
    config: FusedMoEConfig,
    weight_key: QuantKey | None,
    activation_key: QuantKey | None,
    allow_vllm_cutlass: bool = False,
) -> tuple[Fp8MoeBackend, type[mk.FusedMoEExperts] | None]:
    """Select the primary FP8 MoE backend
    Note: Shape-specific fallbacks may still occur at runtime.
    """
    # NOTE: the kernels are selected in the following order.
    AVAILABLE_BACKENDS = _get_priority_backends(config, weight_key, activation_key)

    # NOTE(rob): We need to peak into the P/F selection to determine
    # if we are using the batched or standard expert format, which
    # if not ideal. Once we unify TP + DP/EP, we can select P/F first.
    activation_format = (
        mk.FusedMoEActivationFormat.BatchedExperts
        if config.moe_parallel_config.use_batched_activation_format
        else mk.FusedMoEActivationFormat.Standard
    )

    def _make_log_backend(backend: Fp8MoeBackend):
        available_backend_strs = [b.value for b in AVAILABLE_BACKENDS]
        return (
            f"Using {backend.value} Fp8 MoE backend out "
            f"of potential backends: {available_backend_strs}."
        )

    def _make_log_unsupported(backend: Fp8MoeBackend, reason: str | None) -> str:
        if reason:
            return (
                f"FP8 MoE backend {backend.value} does not support the "
                f"deployment configuration since {reason}."
            )
        else:
            return (
                f"FP8 MoE backend '{backend.value}' does not support the "
                "deployment configuration."
            )

    def _return_or_raise(
        backend: Fp8MoeBackend,
        config: FusedMoEConfig,
        weight_key: QuantKey | None,
        activation_key: QuantKey | None,
        activation_format: mk.FusedMoEActivationFormat,
    ) -> tuple[Fp8MoeBackend, type[mk.FusedMoEExperts]]:
        for k_cls in backend_to_kernel_cls(backend):
            supported, reason = k_cls.is_supported_config(
                k_cls, config, weight_key, activation_key, activation_format
            )
            if supported:
                logger.info_once(_make_log_backend(backend))
                return backend, k_cls
        raise ValueError(_make_log_unsupported(backend, reason))

    # Handle explicit moe_backend from user.
    runner_backend = config.moe_backend
    if runner_backend != "auto":
        requested_backend = map_fp8_backend(runner_backend)
        # For batched activation format, use batched variants if available.
        if activation_format == mk.FusedMoEActivationFormat.BatchedExperts:
            if requested_backend == Fp8MoeBackend.DEEPGEMM:
                requested_backend = Fp8MoeBackend.BATCHED_DEEPGEMM
            elif requested_backend == Fp8MoeBackend.TRITON:
                requested_backend = Fp8MoeBackend.BATCHED_TRITON
            elif requested_backend == Fp8MoeBackend.VLLM_CUTLASS:
                requested_backend = Fp8MoeBackend.BATCHED_VLLM_CUTLASS

        if (
            requested_backend
            in [
                Fp8MoeBackend.VLLM_CUTLASS,
                Fp8MoeBackend.BATCHED_VLLM_CUTLASS,
            ]
            and not allow_vllm_cutlass
        ):
            raise ValueError(
                "vLLM CUTLASS FP8 MoE backend is disabled for this configuration."
            )

        return _return_or_raise(
            requested_backend, config, weight_key, activation_key, activation_format
        )

    # Handle explicit DeepGEMM FP8 configuration.
    if envs.is_set("VLLM_USE_DEEP_GEMM") or envs.is_set("VLLM_MOE_USE_DEEP_GEMM"):
        if not envs.VLLM_USE_DEEP_GEMM or not envs.VLLM_MOE_USE_DEEP_GEMM:
            AVAILABLE_BACKENDS.remove(Fp8MoeBackend.DEEPGEMM)
            AVAILABLE_BACKENDS.remove(Fp8MoeBackend.BATCHED_DEEPGEMM)
        else:
            backend = (
                Fp8MoeBackend.DEEPGEMM
                if activation_format == mk.FusedMoEActivationFormat.Standard
                else Fp8MoeBackend.BATCHED_DEEPGEMM
            )
            return _return_or_raise(
                backend, config, weight_key, activation_key, activation_format
            )

    # Handle explicit AITER FP8 configuration.
    if envs.is_set("VLLM_ROCM_USE_AITER") or envs.is_set("VLLM_ROCM_USE_AITER_MOE"):
        skip_aiter_moe = (
            not envs.VLLM_ROCM_USE_AITER
            or not envs.VLLM_ROCM_USE_AITER_MOE
            or rocm_aiter_ops.is_rdna_aiter_enabled()
        )
        if skip_aiter_moe:
            if Fp8MoeBackend.AITER in AVAILABLE_BACKENDS:
                AVAILABLE_BACKENDS.remove(Fp8MoeBackend.AITER)
        else:
            backend = Fp8MoeBackend.AITER
            return _return_or_raise(
                backend, config, weight_key, activation_key, activation_format
            )

    if not allow_vllm_cutlass:
        AVAILABLE_BACKENDS.remove(Fp8MoeBackend.VLLM_CUTLASS)
        AVAILABLE_BACKENDS.remove(Fp8MoeBackend.BATCHED_VLLM_CUTLASS)

    # Select kernels in order of backend.
    for backend in AVAILABLE_BACKENDS:
        for k_cls in backend_to_kernel_cls(backend):
            supported, reason = k_cls.is_supported_config(
                k_cls,
                config,
                weight_key,
                activation_key,
                activation_format,
            )
            if supported:
                logger.info_once(_make_log_backend(backend))
                return backend, k_cls
            else:
                logger.debug_once(_make_log_unsupported(backend, reason))

    # TODO(rob): per discussion with TPU team, we need a way to register
    # MoE backends by OOT plugins, rather than having an explicit list
    # of AVAILABLE_BACKENDS. Enabling returning `Fp8MoeBackend.NONE` is
    # a temporary measure until these register APIs are complete.
    if current_platform.is_cuda() or current_platform.is_rocm():
        raise NotImplementedError(
            "No FP8 MoE backend supports the deployment configuration."
        )

    return Fp8MoeBackend.NONE, None