Skip to content

vllm.models.kimi_k3.common.mtp

Fused Kimi-K3 MTP input preparation.

Functions:

  • fused_mtp_input

    Mask and normalize both MTP inputs into the projection layout.

fused_mtp_input(positions, inputs_embeds, previous_hidden_states, enorm_weight, hnorm_weight, eps)

Mask and normalize both MTP inputs into the projection layout.

Source code in vllm/models/kimi_k3/common/mtp.py
def fused_mtp_input(
    positions: torch.Tensor,
    inputs_embeds: torch.Tensor,
    previous_hidden_states: torch.Tensor,
    enorm_weight: torch.Tensor,
    hnorm_weight: torch.Tensor,
    eps: float,
) -> torch.Tensor:
    """Mask and normalize both MTP inputs into the projection layout."""
    num_tokens, hidden_size = inputs_embeds.shape
    output = torch.empty(
        num_tokens,
        2 * hidden_size,
        dtype=inputs_embeds.dtype,
        device=inputs_embeds.device,
    )
    if num_tokens == 0:
        return output

    _fused_mtp_input_kernel[(num_tokens, 2)](
        positions,
        inputs_embeds,
        previous_hidden_states,
        enorm_weight,
        hnorm_weight,
        output,
        eps,
        inputs_embeds.stride(0),
        previous_hidden_states.stride(0),
        output.stride(0),
        hidden_size,
        triton.next_power_of_2(hidden_size),
    )
    return output