Skip to content

vllm.models.deepseek_v41.nvidia.ops.mhc

Dispatch DSV4.1 mHC operations and overlap coefficient generation.

Functions:

init_mhc_all_reduce(vllm_config)

Build the fused all-reduce + mHC kernel the overlap path reduces with.

Collective over the TP group, so every rank must call it, and only when supports_mhc_all_reduce holds.

Source code in vllm/models/deepseek_v41/nvidia/ops/mhc.py
def init_mhc_all_reduce(vllm_config: VllmConfig) -> None:
    """Build the fused all-reduce + mHC kernel the overlap path reduces with.

    Collective over the TP group, so every rank must call it, and only when
    ``supports_mhc_all_reduce`` holds.
    """
    global _all_reduce_mhc
    from .cute_dsl import AllReduceMHC

    config = vllm_config.model_config.hf_config
    _all_reduce_mhc = AllReduceMHC(
        hidden_size=config.hidden_size,
        hc_mult=config.hc_mult,
        max_num_tokens=MHC_OVERLAP_MAX_TOKENS,
        top_k=config.num_experts_per_tok,
        device=current_platform.current_device(),
    )
    logger.info_once(
        "DSV4.1 mHC: CuTe DSL Lamport all-reduce fused with mHC for up to %d tokens.",
        MHC_OVERLAP_MAX_TOKENS,
    )

mhc_pre_delayed_overlap(residual, fn, hc_scale, hc_base, rms_eps, hc_pre_eps, hc_sinkhorn_eps, hc_post_mult_value, sinkhorn_repeat, pre_mix=None, x=None, norm_weight=None, norm_eps=1e-06, *, stream, layer_input=None)

Prepare the input on the caller stream and coefficients on another stream.

Returns post mix, residual mix, normalized input, and next pre-mix. Only the input is ready on the caller stream; join stream before using coefficients. The caller must retain the weights and coefficient outputs until the join.

Source code in vllm/models/deepseek_v41/nvidia/ops/mhc.py
def mhc_pre_delayed_overlap(
    residual: torch.Tensor,
    fn: torch.Tensor,
    hc_scale: torch.Tensor,
    hc_base: torch.Tensor,
    rms_eps: float,
    hc_pre_eps: float,
    hc_sinkhorn_eps: float,
    hc_post_mult_value: float,
    sinkhorn_repeat: int,
    pre_mix: torch.Tensor | None = None,
    x: torch.Tensor | None = None,
    norm_weight: torch.Tensor | None = None,
    norm_eps: float = 1e-6,
    *,
    stream: torch.cuda.Stream,
    layer_input: torch.Tensor | None = None,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
    """Prepare the input on the caller stream and coefficients on another stream.

    Returns post mix, residual mix, normalized input, and next pre-mix. Only the
    input is ready on the caller stream; join stream before using coefficients.
    The caller must retain the weights and coefficient outputs until the join.
    """
    from vllm.model_executor.kernels.mhc.warmup import (
        MHC_PRE_NORM_KERNEL,
        compute_mhc_pre_num_splits,
    )
    from vllm.utils.deep_gemm import tf32_hc_prenorm_gemm

    n, hc, hidden = residual.shape
    assert residual.is_contiguous() and residual.dtype == torch.bfloat16
    assert norm_weight is not None
    if x is None:
        x = residual.view(n, hc * hidden)
    assert x.is_contiguous()
    post = torch.empty((n, hc), device=residual.device, dtype=torch.float32)
    comb = torch.empty((n, hc * hc), device=residual.device, dtype=torch.float32)
    next_pre = torch.empty_like(post)
    compute_input = layer_input is None
    if layer_input is None:
        layer_input = torch.empty(
            (n, hidden), device=residual.device, dtype=torch.bfloat16
        )
    outputs = post.unsqueeze(-1), comb.view(n, hc, hc), layer_input, next_pre
    if n == 0:
        return outputs
    splits = compute_mhc_pre_num_splits(x.shape[1], n)
    mix = torch.empty(
        (splits, n, hc * (hc + 2)), device=residual.device, dtype=torch.float32
    )
    sqr = torch.empty((splits, n), device=residual.device, dtype=torch.float32)
    epilogue = partial(
        MHC_PRE_NORM_KERNEL,
        mix,
        sqr,
        hc_scale,
        hc_base,
        residual,
        post,
        comb,
        layer_input,
        norm_weight,
        pre_mix if pre_mix is not None else post,
        next_pre,
        layer_input,  # Unused aux output in split modes.
        hidden_size=hidden,
        rms_eps=rms_eps,
        hc_pre_eps=hc_pre_eps,
        hc_sinkhorn_eps=hc_sinkhorn_eps,
        hc_post_mult_value=hc_post_mult_value,
        sinkhorn_repeat=sinkhorn_repeat,
        norm_eps=norm_eps,
        hc_mult=hc,
        use_pre_mix_in=pre_mix is not None,
        save_pre_mix=True,
        rms_numel=x.shape[1],
    )
    main = torch.cuda.current_stream()
    # Prioritize input readiness for small batches before releasing statistics.
    if compute_input and n <= 8:
        epilogue(split_mode="input")
    stream.wait_stream(main)
    with torch.cuda.stream(stream):
        tf32_hc_prenorm_gemm(x, fn, mix, sqr, splits)
        epilogue(split_mode="stats")
    for tensor in (x, mix, sqr):
        tensor.record_stream(stream)
    if compute_input and n > 8:
        epilogue(split_mode="input")
    return outputs

mhc_shifted_post_pre(x, residual, post_layer_mix, comb_res_mix, fn, hc_scale, hc_base, rms_eps, hc_pre_eps, hc_sinkhorn_eps, hc_post_mult_value, sinkhorn_repeat, pre_mix=None, norm_weight=None, norm_eps=1e-06, capture_aux=False, *, stream=None, reduce_results=False)

Dispatch shifted post/pre to overlap, Mega-mHC, or fused TileLang.

A MoE output left unfinalized is finalized inside the fused all-reduce. When stream is supplied, join it before consuming the returned coefficients.

Source code in vllm/models/deepseek_v41/nvidia/ops/mhc.py
def mhc_shifted_post_pre(
    x: torch.Tensor | MoEOutput,
    residual: torch.Tensor,
    post_layer_mix: torch.Tensor,
    comb_res_mix: torch.Tensor,
    fn: torch.Tensor,
    hc_scale: torch.Tensor,
    hc_base: torch.Tensor,
    rms_eps: float,
    hc_pre_eps: float,
    hc_sinkhorn_eps: float,
    hc_post_mult_value: float,
    sinkhorn_repeat: int,
    pre_mix: torch.Tensor | None = None,
    norm_weight: torch.Tensor | None = None,
    norm_eps: float = 1e-6,
    capture_aux: bool = False,
    *,
    stream: torch.cuda.Stream | None = None,
    reduce_results: bool = False,
) -> tuple[
    torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor
]:
    """Dispatch shifted post/pre to overlap, Mega-mHC, or fused TileLang.

    A MoE output left unfinalized is finalized inside the fused all-reduce.
    When stream is supplied, join it before consuming the returned coefficients.
    """
    layer_input = None
    if isinstance(x, MoEOutput):
        assert _all_reduce_mhc is not None, "init_mhc_all_reduce was not called"
        assert pre_mix is not None and norm_weight is not None
        routed = x.routed
        assert isinstance(routed, UnfinalizedMoEOutput)
        assert x.shared_output is not None
        residual, layer_input = _all_reduce_mhc.finalize(
            routed.gemm2_permuted,
            routed.expert_weights,
            routed.expanded_idx_to_permuted_idx,
            x.shared_output,
            residual,
            post_layer_mix,
            comb_res_mix,
            pre_mix,
            norm_weight,
            norm_eps,
        )
        if stream is None:
            # Outside FULL graphs, compute the next mixes inline.
            stream = torch.cuda.current_stream()
    elif reduce_results:
        if stream is not None and 0 < x.shape[0] <= MHC_OVERLAP_MAX_TOKENS:
            assert _all_reduce_mhc is not None, "init_mhc_all_reduce was not called"
            assert pre_mix is not None and norm_weight is not None
            residual, layer_input = _all_reduce_mhc(
                x,
                residual,
                post_layer_mix,
                comb_res_mix,
                pre_mix,
                norm_weight,
                norm_eps,
            )
        else:
            x = get_tp_group().all_reduce(x)
    if stream is not None:
        if layer_input is None:
            assert isinstance(x, torch.Tensor)
            residual = mhc_post_tilelang(x, residual, post_layer_mix, comb_res_mix)
        aux = (
            residual.mean(dim=1)
            if capture_aux
            else residual.new_empty(0, residual.shape[-1])
        )
        pre_outputs = mhc_pre_delayed_overlap(
            residual,
            fn,
            hc_scale,
            hc_base,
            rms_eps,
            hc_pre_eps,
            hc_sinkhorn_eps,
            hc_post_mult_value,
            sinkhorn_repeat,
            pre_mix=pre_mix,
            norm_weight=norm_weight,
            norm_eps=norm_eps,
            stream=stream,
            layer_input=layer_input,
        )
        return residual, *pre_outputs, aux

    assert isinstance(x, torch.Tensor)
    if can_use_mega_mhc(x, residual, pre_mix, norm_weight, capture_aux):
        assert pre_mix is not None and norm_weight is not None
        outputs = mhc_shifted_post_pre_deep_gemm(
            x,
            residual,
            pre_mix,
            post_layer_mix,
            comb_res_mix,
            fn,
            hc_scale,
            hc_base,
            rms_eps,
            hc_pre_eps,
            hc_post_mult_value,
            hc_sinkhorn_eps,
            sinkhorn_repeat,
            norm_weight,
            norm_eps,
        )
        return *outputs, x.new_empty(0, x.shape[1])

    return mhc_fused_post_pre_delayed_tilelang(
        x,
        residual,
        post_layer_mix,
        comb_res_mix,
        fn,
        hc_scale,
        hc_base,
        rms_eps,
        hc_pre_eps,
        hc_sinkhorn_eps,
        hc_post_mult_value,
        sinkhorn_repeat,
        pre_mix=pre_mix,
        norm_weight=norm_weight,
        norm_eps=norm_eps,
        capture_aux=capture_aux,
    )

supports_mhc_overlap(vllm_config)

Check kernel requirements and safety of sharing the coefficient stream.

Source code in vllm/models/deepseek_v41/nvidia/ops/mhc.py
def supports_mhc_overlap(vllm_config: VllmConfig) -> bool:
    """Check kernel requirements and safety of sharing the coefficient stream."""
    config = vllm_config.model_config.hf_config
    # DeepGEMM's prenorm kernel requires K % 64 == 0, N % 8 == 0, and N <= 32.
    mix_size = config.hc_mult * (config.hc_mult + 2)
    return (
        current_platform.is_device_capability_family(100)
        and is_deep_gemm_supported()
        and config.hidden_size % 64 == 0
        and 0 < mix_size <= 32
        and mix_size % 8 == 0
        and not vllm_config.parallel_config.use_ubatching
    )