Skip to content

vllm.model_executor.layers.fused_moe.moe_fused_mul_sum

Functions:

  • moe_fused_mul_sum –

    Fused kernel for MoE (Mixture of Experts) to perform weighted summation

moe_fused_mul_sum(inputs, topk_weights, outputs=None, topk_ids=None, expert_map=None, num_valid_tokens=None)

Fused kernel for MoE (Mixture of Experts) to perform weighted summation of expert outputs.

Parameters:

  • inputs

    (Tensor) –

    The output from experts. Shape: (num_tokens, top_k, hidden_size).

  • topk_weights

    (Tensor) –

    The weights assigned to each expert for each token. Shape: (num_tokens, top_k).

  • outputs

    (Tensor | None, default: None ) –

    Optional pre-allocated output tensor. Shape: (num_tokens, hidden_size).

  • topk_ids

    (Tensor | None, default: None ) –

    Optional indices of the top-k experts. Shape: (num_tokens, top_k). A value of -1 marks a slot the expert GEMM skipped; those slots are excluded from the sum. When provided, rows with all top ids < 0 (worst-case padding) are skipped and their output rows left untouched. Required when expert_map is provided.

  • expert_map

    (Tensor | None, default: None ) –

    Optional mapping for Expert Parallelism. A value < 0 indicates an invalid token/expert pair that will be skipped. Only needed when topk_ids may contain non-local expert ids; if every non-(-1) id is already a local expert, leave it None to skip the redundant per-slot lookup.

  • num_valid_tokens

    (Tensor | None, default: None ) –

    Optional device scalar (1-element tensor) holding the number of real token rows (num_recv for a decode dispatch). When provided, rows past it are left untouched, so the static cudagraph grid never sums stale padding rows. Pass the token count, not token*top_k.

Returns:

  • Tensor –

    The fused weighted sum of expert outputs.

  • Shape ( Tensor ) –

    (num_tokens, hidden_size).

Source code in vllm/model_executor/layers/fused_moe/moe_fused_mul_sum.py
def moe_fused_mul_sum(
    inputs: torch.Tensor,
    topk_weights: torch.Tensor,
    outputs: torch.Tensor | None = None,
    topk_ids: torch.Tensor | None = None,
    expert_map: torch.Tensor | None = None,
    num_valid_tokens: torch.Tensor | None = None,
) -> torch.Tensor:
    """Fused kernel for MoE (Mixture of Experts) to perform weighted summation
    of expert outputs.

    Args:
        inputs: The output from experts.
            Shape: (num_tokens, top_k, hidden_size).
        topk_weights: The weights assigned to each expert for each token.
            Shape: (num_tokens, top_k).
        outputs: Optional pre-allocated output tensor.
            Shape: (num_tokens, hidden_size).
        topk_ids: Optional indices of the top-k experts. Shape:
            (num_tokens, top_k). A value of -1 marks a slot the expert GEMM
            skipped; those slots are excluded from the sum. When provided, rows
            with all top ids < 0 (worst-case padding) are skipped and their
            output rows left untouched. Required when `expert_map` is provided.
        expert_map: Optional mapping for Expert Parallelism. A value < 0
            indicates an invalid token/expert pair that will be skipped. Only
            needed when `topk_ids` may contain non-local expert ids; if every
            non-(-1) id is already a local expert, leave it None to skip the
            redundant per-slot lookup.
        num_valid_tokens: Optional device scalar (1-element tensor) holding the
            number of real token rows (num_recv for a decode dispatch). When
            provided, rows past it are left untouched, so the static cudagraph
            grid never sums stale padding rows. Pass the token count, not
            token*top_k.

    Returns:
        The fused weighted sum of expert outputs.
        Shape: (num_tokens, hidden_size).

    """
    assert inputs.ndim == 3
    assert topk_weights.ndim == 2
    assert inputs.is_contiguous()
    assert topk_weights.is_contiguous()
    assert inputs.dtype in (torch.float32, torch.float16, torch.bfloat16)
    assert topk_weights.dtype in (torch.float32, torch.float16, torch.bfloat16)

    num_tokens, top_k, hidden_size = inputs.shape
    output_shape = (num_tokens, hidden_size)
    if outputs is None:
        outputs = torch.empty(output_shape, dtype=inputs.dtype, device=inputs.device)

    assert outputs.shape == output_shape
    assert topk_weights.shape == (num_tokens, top_k)
    assert expert_map is None or topk_ids is not None, (
        "topk_ids is required to interpret expert_map"
    )
    if topk_ids is not None:
        assert topk_ids.shape == (num_tokens, top_k)
        assert topk_ids.is_contiguous()
        assert topk_ids.dtype in (torch.int32, torch.int64)

    if not isinstance(inputs, FakeTensor):
        BLOCK_K, num_warps, num_stages = _heuristic_config(
            hidden_size,
            inputs.element_size(),
            inputs.device.index or 0,
        )
        grid = (num_tokens,)
        moe_fused_mul_sum_kernel[grid](
            inputs,
            topk_weights,
            outputs,
            topk_ids,
            expert_map,
            num_valid_tokens,
            top_k * hidden_size,
            topk_ids is not None,
            expert_map is not None,
            num_valid_tokens is not None,
            top_k,
            hidden_size,
            BLOCK_K,
            num_warps=num_warps,
            num_stages=num_stages,
        )

    return outputs