Skip to content

vllm.models.deepseek_v4.nvidia.ops.prepare_megamoe

Triton input-staging kernel for DeepSeek V4 MegaMoE.

Quantizes hidden states to fp8 with E8M0 group scales and repacks the routing top-k tensors into the int64/float32 layout that the DeepGEMM MegaMoE kernels consume.

Functions:

_resolve_hidden_quant(hidden_quant, block_k)

Map a hidden-state QuantKey to the kernel's (GROUP_K, USE_UE8M0).

Source code in vllm/models/deepseek_v4/nvidia/ops/prepare_megamoe.py
def _resolve_hidden_quant(hidden_quant: QuantKey, block_k: int) -> tuple[int, bool]:
    """Map a hidden-state ``QuantKey`` to the kernel's ``(GROUP_K, USE_UE8M0)``."""
    scale = hidden_quant.scale
    group_shape = scale.group_shape
    if (
        hidden_quant.dtype != FP8_DTYPE
        or not hidden_quant.symmetric
        or hidden_quant.scale2 is not None
        or scale.static
        or not group_shape.is_per_group()
        or block_k % group_shape.col != 0
        or scale.dtype not in (*_E8M0_SCALE_DTYPES, torch.float32)
    ):
        raise ValueError(
            "DeepSeek V4 MegaMoE input staging requires symmetric dynamic fp8 "
            f"per-group quantization with a group size dividing {block_k} and "
            f"E8M0 or fp32 scales, got {hidden_quant}."
        )
    return group_shape.col, scale.dtype in _E8M0_SCALE_DTYPES

prepare_megamoe_inputs(hidden_states, topk_weights, topk_ids, x_fp8, x_sf, topk_idx_out, topk_weights_out, is_padding=None, shared_x_sf=None, shared_block_m=None, hidden_quant=kMxfp8Dynamic)

Quantize hidden states and repack top-k routing for DeepGEMM MegaMoE.

Parameters:

  • hidden_states

    (Tensor) –

    Input activations of shape [num_tokens, hidden_dim].

  • topk_weights

    (Tensor) –

    Router top-k weights of shape [num_tokens, top_k].

  • topk_ids

    (Tensor) –

    Router top-k expert ids of shape [num_tokens, top_k].

  • x_fp8

    (Tensor) –

    Output buffer for the fp8-quantized hidden states.

  • x_sf

    (Tensor) –

    Output buffer for the hidden-state scale factors.

  • topk_idx_out

    (Tensor) –

    Output buffer for the repacked top-k expert ids.

  • topk_weights_out

    (Tensor) –

    Output buffer for the repacked top-k weights.

  • is_padding

    (Tensor | None, default: None ) –

    Optional per-token mask; padded tokens are not routed.

  • shared_x_sf

    (Tensor | None, default: None ) –

    Optional output buffer for shared-expert scale factors. Must be given together with shared_block_m.

  • shared_block_m

    (int | None, default: None ) –

    Block M used to lay out shared_x_sf.

  • hidden_quant

    (QuantKey, default: kMxfp8Dynamic ) –

    Hidden-state quantization scheme. E8M0 scales are packed four per int32 into x_sf; fp32 scales are stored directly.

Source code in vllm/models/deepseek_v4/nvidia/ops/prepare_megamoe.py
def prepare_megamoe_inputs(
    hidden_states: torch.Tensor,
    topk_weights: torch.Tensor,
    topk_ids: torch.Tensor,
    x_fp8: torch.Tensor,
    x_sf: torch.Tensor,
    topk_idx_out: torch.Tensor,
    topk_weights_out: torch.Tensor,
    is_padding: torch.Tensor | None = None,
    shared_x_sf: torch.Tensor | None = None,
    shared_block_m: int | None = None,
    hidden_quant: QuantKey = kMxfp8Dynamic,
) -> None:
    """Quantize hidden states and repack top-k routing for DeepGEMM MegaMoE.

    Args:
        hidden_states: Input activations of shape ``[num_tokens, hidden_dim]``.
        topk_weights: Router top-k weights of shape ``[num_tokens, top_k]``.
        topk_ids: Router top-k expert ids of shape ``[num_tokens, top_k]``.
        x_fp8: Output buffer for the fp8-quantized hidden states.
        x_sf: Output buffer for the hidden-state scale factors.
        topk_idx_out: Output buffer for the repacked top-k expert ids.
        topk_weights_out: Output buffer for the repacked top-k weights.
        is_padding: Optional per-token mask; padded tokens are not routed.
        shared_x_sf: Optional output buffer for shared-expert scale factors.
            Must be given together with ``shared_block_m``.
        shared_block_m: Block M used to lay out ``shared_x_sf``.
        hidden_quant: Hidden-state quantization scheme. E8M0 scales are packed
            four per int32 into ``x_sf``; fp32 scales are stored directly.

    """
    num_tokens, hidden_size = hidden_states.shape
    if num_tokens == 0:
        return
    if hidden_size % 128 != 0:
        raise ValueError(
            "DeepSeek V4 MegaMoE input staging requires hidden_size to be "
            "a multiple of 128."
        )
    block_k = 128
    group_k, use_ue8m0 = _resolve_hidden_quant(hidden_quant, block_k)
    top_k = topk_ids.shape[1]
    if topk_weights.shape != topk_ids.shape:
        raise ValueError(
            "DeepSeek V4 MegaMoE input staging requires topk_weights and "
            "topk_ids to have the same shape."
        )
    if (shared_x_sf is None) != (shared_block_m is None):
        raise ValueError(
            "DeepSeek V4 MegaMoE shared input staging requires both "
            "shared_x_sf and shared_block_m."
        )
    if shared_x_sf is not None and not use_ue8m0:
        raise ValueError(
            "DeepSeek V4 MegaMoE shared input staging currently requires "
            "UE8M0-packed hidden scales."
        )
    if shared_x_sf is not None:
        assert shared_block_m is not None
        if shared_block_m <= 0:
            raise ValueError("MegaMoE shared_block_m must be positive.")
        expected_sf_k = hidden_size // 128
        if shared_x_sf.ndim != 2 or shared_x_sf.shape[1] != expected_sf_k:
            raise ValueError(
                "MegaMoE shared_x_sf must have shape "
                f"(*, {expected_sf_k}), got {tuple(shared_x_sf.shape)}."
            )
        aligned_block_m = triton.cdiv(shared_block_m, 128) * 128
        required_rows = triton.cdiv(num_tokens, shared_block_m) * aligned_block_m
        if shared_x_sf.shape[0] < required_rows:
            raise ValueError(
                "MegaMoE shared_x_sf has insufficient rows: requires "
                f"{required_rows}, got {shared_x_sf.shape[0]}."
            )

    # On GB200, eight-row tiles win from 64 tokens; keep smaller batches untiled.
    block_m = 8 if num_tokens >= 64 else 1
    grid = (triton.cdiv(num_tokens, block_m), triton.cdiv(hidden_size, block_k))
    block_topk = triton.next_power_of_2(top_k)
    padding_stride_m = is_padding.stride(0) if is_padding is not None else 0
    _prepare_megamoe_inputs_kernel[grid](
        hidden_states,
        x_fp8,
        x_sf,
        shared_x_sf,
        topk_ids,
        topk_weights,
        is_padding,
        topk_idx_out,
        topk_weights_out,
        hidden_states.stride(0),
        hidden_states.stride(1),
        x_fp8.stride(0),
        x_fp8.stride(1),
        x_sf.stride(0),
        x_sf.stride(1),
        shared_x_sf.stride(0) if shared_x_sf is not None else 0,
        shared_x_sf.stride(1) if shared_x_sf is not None else 0,
        topk_ids.stride(0),
        topk_ids.stride(1),
        topk_weights.stride(0),
        topk_weights.stride(1),
        padding_stride_m,
        topk_idx_out.stride(0),
        topk_idx_out.stride(1),
        topk_weights_out.stride(0),
        topk_weights_out.stride(1),
        num_tokens,
        hidden_size,
        top_k,
        BLOCK_M=block_m,
        BLOCK_K=block_k,
        GROUP_K=group_k,
        BLOCK_TOPK=block_topk,
        SHARED_BLOCK_M=shared_block_m or 1,
        USE_UE8M0=use_ue8m0,
        num_warps=4,
    )