Skip to content

vllm.model_executor.layers.mhc

Classes:

Functions:

  • hc_contract –

    [s, n * hidden_size] -> [s, hidden_size] by averaging.

  • hc_expand –

    [s, hidden_size] -> [s, n * hidden_size] by replication.

HCHeadOp

Bases: CustomOp

HC head reduction for DeepSeek V4.

Computes gates from the RMS-normalized flattened HC residual and returns out = sum_i gate_i * residual_i, collapsing hc_mult streams to one.

Source code in vllm/model_executor/layers/mhc.py
@CustomOp.register("hc_head")
class HCHeadOp(CustomOp):
    """HC head reduction for DeepSeek V4.

    Computes gates from the RMS-normalized flattened HC residual and
    returns out = sum_i gate_i * residual_i, collapsing hc_mult streams
    to one.
    """

    # --8<-- [end:hc_head]
    @classmethod
    def enabled(cls) -> bool:
        return True

    def forward_cuda(
        self,
        hidden_states: torch.Tensor,
        hc_fn: torch.Tensor,
        hc_scale: torch.Tensor,
        hc_base: torch.Tensor,
        rms_norm_eps: float,
        hc_eps: float,
    ) -> torch.Tensor:
        hc_mult, hidden_size = hidden_states.shape[-2:]
        outer_shape = hidden_states.shape[:-2]
        hs_flat = hidden_states.view(-1, hc_mult, hidden_size)
        out = torch.ops.vllm.hc_head_fused_kernel_tilelang(
            hs_flat,
            hc_fn,
            hc_scale,
            hc_base,
            rms_norm_eps,
            hc_eps,
        )
        return out.view(*outer_shape, hidden_size)

    def forward_hip(
        self,
        hidden_states: torch.Tensor,
        hc_fn: torch.Tensor,
        hc_scale: torch.Tensor,
        hc_base: torch.Tensor,
        rms_norm_eps: float,
        hc_eps: float,
    ) -> torch.Tensor:
        hc_mult, hidden_size = hidden_states.shape[-2:]
        outer_shape = hidden_states.shape[:-2]
        hs_flat = hidden_states.view(-1, hc_mult, hidden_size)

        if HAS_TILELANG_MHC:
            out = torch.ops.vllm.hc_head_fused_kernel_tilelang(
                hs_flat,
                hc_fn,
                hc_scale,
                hc_base,
                rms_norm_eps,
                hc_eps,
            )
        else:
            num_tokens = hs_flat.shape[0]
            out = torch.empty(
                num_tokens,
                hidden_size,
                dtype=torch.bfloat16,
                device=hidden_states.device,
            )
            torch.ops.vllm.hc_head_triton(
                hs_flat,
                hc_fn,
                hc_scale,
                hc_base,
                out,
                hidden_size,
                rms_norm_eps,
                hc_eps,
                hc_mult,
            )

        return out.view(*outer_shape, hidden_size)

    def forward_native(self, *args, **kwargs):
        raise NotImplementedError("Native implementation of hc_head is not available")

    def forward_xpu(
        self,
        hidden_states: torch.Tensor,
        hc_fn: torch.Tensor,
        hc_scale: torch.Tensor,
        hc_base: torch.Tensor,
        rms_norm_eps: float,
        hc_eps: float,
    ) -> torch.Tensor:
        hc_mult, hidden_size = hidden_states.shape[-2:]
        outer_shape = hidden_states.shape[:-2]
        hs_flat = hidden_states.view(-1, hc_mult, hidden_size)
        num_tokens = hs_flat.shape[0]

        out = torch.empty(
            num_tokens, hidden_size, dtype=torch.bfloat16, device=hidden_states.device
        )
        torch.ops._xpu_C.hc_head_fused(
            hs_flat, hc_fn, hc_scale, hc_base, out, rms_norm_eps, hc_eps
        )
        return out.view(*outer_shape, hidden_size)

    def forward_cpu(
        self,
        hidden_states: torch.Tensor,
        hc_fn: torch.Tensor,
        hc_scale: torch.Tensor,
        hc_base: torch.Tensor,
        rms_norm_eps: float,
        hc_eps: float,
    ) -> torch.Tensor:
        return mhc_kernels.hc_head_fused_cpu(
            hidden_states, hc_fn, hc_scale, hc_base, rms_norm_eps, hc_eps
        )

MHCFusedPostPreOp

Bases: CustomOp

Fused MHC post block followed by the next MHC pre block.

Equivalent to applying MHCPostOp and then MHCPreOp to the updated residual streams, returning residual_cur, post_mix_cur, comb_mix_cur, and layer_input_cur.

Source code in vllm/model_executor/layers/mhc.py
@CustomOp.register("mhc_fused_post_pre")
class MHCFusedPostPreOp(CustomOp):
    """Fused MHC post block followed by the next MHC pre block.

    Equivalent to applying MHCPostOp and then MHCPreOp to the updated
    residual streams, returning residual_cur, post_mix_cur, comb_mix_cur,
    and layer_input_cur.
    """

    # --8<-- [end:mhc_fused_post_pre]
    @classmethod
    def enabled(cls) -> bool:
        return True

    def forward_cuda(
        self,
        x: torch.Tensor,
        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,
        n_splits: int = 1,
        tile_n: int = 1,
        norm_weight: torch.Tensor | None = None,
        norm_eps: float = 0.0,
    ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
        return torch.ops.vllm.mhc_fused_post_pre_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,
            n_splits,
            tile_n,
            norm_weight,
            norm_eps,
        )

    def forward_hip(
        self,
        x: torch.Tensor,
        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,
        n_splits: int = 1,
        tile_n: int = 1,
        norm_weight: torch.Tensor | None = None,
        norm_eps: float = 0.0,
    ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
        if HAS_AITER_MHC_FUSED and _aiter_mhc_supported(
            residual,
            norm_weight,
            supports_norm=HAS_AITER_MHC_FUSED_NORM,
        ):
            return torch.ops.vllm.mhc_fused_post_pre_aiter(
                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,
                n_splits,
                tile_n,
                norm_weight,
                norm_eps,
            )
        if HAS_TILELANG_MHC:
            return torch.ops.vllm.mhc_fused_post_pre_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,
                n_splits,
                tile_n,
                norm_weight,
                norm_eps,
            )
        residual_cur, post_mix_cur, comb_mix_cur, layer_input_cur = self.forward_native(
            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,
            n_splits,
            tile_n,
            norm_weight,
            norm_eps,
        )
        return (
            residual_cur,
            post_mix_cur,
            comb_mix_cur,
            _apply_mhc_norm(layer_input_cur, norm_weight, norm_eps),
        )

    def forward_native(
        self,
        x: torch.Tensor,
        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,
        n_splits: int = 1,
        tile_n: int = 1,
        norm_weight: torch.Tensor | None = None,
        norm_eps: float = 0.0,
    ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
        # Decompose into post + pre (no fused kernel available).
        residual_cur = mhc_kernels.mhc_post_torch(
            x, residual, post_layer_mix, comb_res_mix
        )
        post_mix_cur, comb_mix_cur, layer_input_cur = mhc_kernels.mhc_pre_torch(
            residual_cur,
            fn,
            hc_scale,
            hc_base,
            rms_eps,
            hc_pre_eps,
            hc_sinkhorn_eps,
            hc_post_mult_value,
            sinkhorn_repeat,
        )
        return residual_cur, post_mix_cur, comb_mix_cur, layer_input_cur

    def forward_xpu(
        self,
        x: torch.Tensor,
        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,
        n_splits: int = 1,
        tile_n: int = 1,
        norm_weight: torch.Tensor | None = None,
        norm_eps: float = 0.0,
    ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
        return torch.ops._xpu_C.mhc_fused_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,
        )

    def forward_cpu(
        self,
        x: torch.Tensor,
        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,
        n_splits: int = 1,
        tile_n: int = 1,
        norm_weight: torch.Tensor | None = None,
        norm_eps: float = 0.0,
    ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
        # Decompose into post + pre (no fused kernel available).
        residual_cur = mhc_kernels.mhc_post_cpu(
            x, residual, post_layer_mix, comb_res_mix
        )
        post_mix_cur, comb_mix_cur, layer_input_cur = mhc_kernels.mhc_pre_cpu(
            residual_cur,
            fn,
            hc_scale,
            hc_base,
            rms_eps,
            hc_pre_eps,
            hc_sinkhorn_eps,
            hc_post_mult_value,
            sinkhorn_repeat,
        )
        return residual_cur, post_mix_cur, comb_mix_cur, layer_input_cur

MHCPostOp

Bases: CustomOp

MHC post block.

Combines the layer output with the HC residual streams: out_j = post_layer_mix_j * x + sum_i comb_res_mix_ij * residual_i.

Source code in vllm/model_executor/layers/mhc.py
@CustomOp.register("mhc_post")
class MHCPostOp(CustomOp):
    """MHC post block.

    Combines the layer output with the HC residual streams:
    out_j = post_layer_mix_j * x + sum_i comb_res_mix_ij * residual_i.
    """

    # --8<-- [end:mhc_post]

    @classmethod
    def enabled(cls) -> bool:
        return True

    def forward_cuda(
        self,
        x: torch.Tensor,
        residual: torch.Tensor,
        post_layer_mix: torch.Tensor,
        comb_res_mix: torch.Tensor,
    ) -> torch.Tensor:
        return torch.ops.vllm.mhc_post_tilelang(
            x, residual, post_layer_mix, comb_res_mix
        )

    def forward_hip(
        self,
        x: torch.Tensor,
        residual: torch.Tensor,
        post_layer_mix: torch.Tensor,
        comb_res_mix: torch.Tensor,
    ) -> torch.Tensor:
        if _aiter_mhc_supported(residual, None, supports_norm=True):
            return torch.ops.vllm.mhc_post_aiter(
                x,
                residual,
                post_layer_mix,
                comb_res_mix,
            )
        if HAS_TILELANG_MHC:
            return torch.ops.vllm.mhc_post_tilelang(
                x, residual, post_layer_mix, comb_res_mix
            )
        else:
            return self.forward_native(x, residual, post_layer_mix, comb_res_mix)

    def forward_native(
        self,
        x: torch.Tensor,
        residual: torch.Tensor,
        post_layer_mix: torch.Tensor,
        comb_res_mix: torch.Tensor,
    ) -> torch.Tensor:
        return mhc_kernels.mhc_post_torch(
            x,
            residual,
            post_layer_mix,
            comb_res_mix,
        )

    def forward_xpu(
        self,
        x: torch.Tensor,
        residual: torch.Tensor,
        post_layer_mix: torch.Tensor,
        comb_res_mix: torch.Tensor,
    ) -> torch.Tensor:
        return torch.ops._xpu_C.mhc_post(
            x,
            residual,
            post_layer_mix,
            comb_res_mix,
        )

    def forward_cpu(
        self,
        x: torch.Tensor,
        residual: torch.Tensor,
        post_layer_mix: torch.Tensor,
        comb_res_mix: torch.Tensor,
    ) -> torch.Tensor:
        return mhc_kernels.mhc_post_cpu(x, residual, post_layer_mix, comb_res_mix)

MHCPreDelayedOp

Bases: CustomOp

MHC pre block using the pre-mix carried from the previous sublayer.

Same gates as :class:MHCPreOp, but the stream collapse applies the caller's pre_mix and this sublayer's pre-mix is returned for the next sublayer seam.

Passing sublayer_out / post_layer_mix / comb_res_mix also applies the preceding post block, which lets AITER fold it into the pre projection and saves a launch. Returns residual, post_mix, comb_mix, layer_input, next_pre_mix; residual is the caller's own tensor when no post is requested.

Source code in vllm/model_executor/layers/mhc.py
@CustomOp.register("mhc_pre_delayed")
class MHCPreDelayedOp(CustomOp):
    """MHC pre block using the pre-mix carried from the previous sublayer.

    Same gates as :class:`MHCPreOp`, but the stream collapse applies the
    caller's ``pre_mix`` and this sublayer's pre-mix is returned for the next
    sublayer seam.

    Passing ``sublayer_out`` / ``post_layer_mix`` / ``comb_res_mix`` also
    applies the preceding post block, which lets AITER fold it into the pre
    projection and saves a launch. Returns residual, post_mix, comb_mix,
    layer_input, next_pre_mix; ``residual`` is the caller's own tensor when
    no post is requested.
    """

    # --8<-- [end:mhc_pre_delayed]
    @classmethod
    def enabled(cls) -> bool:
        return True

    def __init__(self) -> None:
        super().__init__()
        # Built here, not lazily: CustomOp resolves dispatch against the
        # current vLLM config, which is only set during model construction.
        self._post = MHCPostOp()

    def _maybe_post(
        self,
        residual: torch.Tensor,
        sublayer_out: torch.Tensor | None,
        post_layer_mix: torch.Tensor | None,
        comb_res_mix: torch.Tensor | None,
    ) -> torch.Tensor:
        if sublayer_out is None:
            return residual
        assert post_layer_mix is not None and comb_res_mix is not None
        return self._post(sublayer_out, residual, post_layer_mix, comb_res_mix)

    def forward_cuda(
        self,
        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,
        sublayer_out: torch.Tensor | None = None,
        post_layer_mix: torch.Tensor | None = None,
        comb_res_mix: torch.Tensor | None = None,
    ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
        residual = self._maybe_post(
            residual, sublayer_out, post_layer_mix, comb_res_mix
        )
        return residual, *torch.ops.vllm.mhc_pre_delayed_tilelang(
            residual,
            fn,
            hc_scale,
            hc_base,
            rms_eps,
            hc_pre_eps,
            hc_sinkhorn_eps,
            hc_post_mult_value,
            sinkhorn_repeat,
            pre_mix,
            x,
            norm_weight,
            norm_eps,
        )

    def forward_hip(
        self,
        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,
        sublayer_out: torch.Tensor | None = None,
        post_layer_mix: torch.Tensor | None = None,
        comb_res_mix: torch.Tensor | None = None,
    ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
        # aiter's fused delayed seam: post-mix, gate projection, and the
        # collapse with the carried pre-mix and its RMSNorm. It always applies
        # the norm and projects `residual` itself.
        if (
            HAS_AITER_MHC_FUSED_POST_PRE_DELAYED_RMS_NORM
            and x is None
            and norm_weight is not None
            and _aiter_mhc_supported(residual, norm_weight, supports_norm=True)
        ):
            residual_out = (
                torch.empty_like(residual) if sublayer_out is not None else None
            )
            rest = torch.ops.vllm.mhc_fused_post_pre_delayed_rms_norm_aiter(
                residual,
                fn,
                hc_scale,
                hc_base,
                rms_eps,
                hc_pre_eps,
                hc_sinkhorn_eps,
                hc_post_mult_value,
                sinkhorn_repeat,
                pre_mix,
                sublayer_out,
                post_layer_mix,
                comb_res_mix,
                norm_weight,
                norm_eps,
                residual_out,
            )
            return (residual if residual_out is None else residual_out), *rest
        # The unfused AITER delayed path drives mhc_pre_gemm_sqrsum against
        # `residual` and folds no RMSNorm, so it cannot serve the model-entry
        # broadcast (which projects a narrower `x`) or a fused norm the branch
        # above did not take. Both are handled by TileLang, or by the reference
        # when TileLang is unavailable.
        if (
            x is None
            and norm_weight is None
            and _aiter_mhc_supported(residual, None, supports_norm=False)
        ):
            from vllm._aiter_ops import rocm_aiter_ops

            num_tokens = residual.numel() // (residual.shape[-1] * residual.shape[-2])
            # Folding the post in only pays while the residual still fits in
            # cache; past that AITER's heuristic wants the separate kernels,
            # which can ask for a non-temporal store and a large-m split-k.
            if sublayer_out is not None and (
                rocm_aiter_ops.mhc_fused_post_pre_delayed_prefers_unfused(num_tokens)
            ):
                residual = self._maybe_post(
                    residual, sublayer_out, post_layer_mix, comb_res_mix
                )
                sublayer_out = None
            # The op writes the folded post's residual into this buffer rather
            # than returning it, since on the unfused path it has none of its
            # own and an op must not hand back one of its inputs.
            residual_out = (
                torch.empty_like(residual) if sublayer_out is not None else None
            )
            rest = torch.ops.vllm.mhc_pre_delayed_aiter(
                residual,
                fn,
                hc_scale,
                hc_base,
                rms_eps,
                hc_pre_eps,
                hc_sinkhorn_eps,
                hc_post_mult_value,
                sinkhorn_repeat,
                pre_mix,
                sublayer_out,
                post_layer_mix,
                comb_res_mix,
                residual_out,
            )
            return (residual if residual_out is None else residual_out), *rest
        residual = self._maybe_post(
            residual, sublayer_out, post_layer_mix, comb_res_mix
        )
        if HAS_TILELANG_MHC:
            return residual, *torch.ops.vllm.mhc_pre_delayed_tilelang(
                residual,
                fn,
                hc_scale,
                hc_base,
                rms_eps,
                hc_pre_eps,
                hc_sinkhorn_eps,
                hc_post_mult_value,
                sinkhorn_repeat,
                pre_mix,
                x,
                norm_weight,
                norm_eps,
            )
        # The post is already applied, so it must not be requested again.
        return self.forward_native(
            residual,
            fn,
            hc_scale,
            hc_base,
            rms_eps,
            hc_pre_eps,
            hc_sinkhorn_eps,
            hc_post_mult_value,
            sinkhorn_repeat,
            pre_mix,
            x,
            norm_weight,
            norm_eps,
        )

    def forward_native(
        self,
        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,
        sublayer_out: torch.Tensor | None = None,
        post_layer_mix: torch.Tensor | None = None,
        comb_res_mix: torch.Tensor | None = None,
    ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
        residual = self._maybe_post(
            residual, sublayer_out, post_layer_mix, comb_res_mix
        )
        post_mix, comb_mix, layer_input, next_pre_mix = (
            mhc_kernels.mhc_pre_delayed_torch(
                residual,
                fn,
                hc_scale,
                hc_base,
                rms_eps,
                hc_pre_eps,
                hc_sinkhorn_eps,
                hc_post_mult_value,
                sinkhorn_repeat,
                pre_mix=pre_mix,
                x=x,
            )
        )
        return (
            residual,
            post_mix,
            comb_mix,
            _apply_mhc_norm(layer_input, norm_weight, norm_eps),
            next_pre_mix,
        )

MHCPreOp

Bases: CustomOp

MHC pre block.

Computes mix logits from RMS-normalized HC residual streams, then returns post_mix, comb_mix, and layer_input = sum_i pre_mix_i * residual_i.

Source code in vllm/model_executor/layers/mhc.py
@CustomOp.register("mhc_pre")
class MHCPreOp(CustomOp):
    """MHC pre block.

    Computes mix logits from RMS-normalized HC residual streams, then
    returns post_mix, comb_mix, and
    layer_input = sum_i pre_mix_i * residual_i.
    """

    # --8<-- [end:mhc_pre]
    @classmethod
    def enabled(cls) -> bool:
        return True

    def forward_cuda(
        self,
        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,
        n_splits: int = 1,
        norm_weight: torch.Tensor | None = None,
        norm_eps: float = 0.0,
    ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
        return torch.ops.vllm.mhc_pre_tilelang(
            residual,
            fn,
            hc_scale,
            hc_base,
            rms_eps,
            hc_pre_eps,
            hc_sinkhorn_eps,
            hc_post_mult_value,
            sinkhorn_repeat,
            n_splits,
            norm_weight,
            norm_eps,
        )

    def forward_hip(
        self,
        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,
        n_splits: int = 1,
        norm_weight: torch.Tensor | None = None,
        norm_eps: float = 0.0,
    ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
        if _aiter_mhc_supported(
            residual,
            norm_weight,
            supports_norm=HAS_AITER_MHC_PRE_NORM,
        ):
            return torch.ops.vllm.mhc_pre_aiter(
                residual,
                fn,
                hc_scale,
                hc_base,
                rms_eps,
                hc_pre_eps,
                hc_sinkhorn_eps,
                hc_post_mult_value,
                sinkhorn_repeat,
                n_splits,
                norm_weight,
                norm_eps,
            )
        elif HAS_TILELANG_MHC:
            return torch.ops.vllm.mhc_pre_tilelang(
                residual,
                fn,
                hc_scale,
                hc_base,
                rms_eps,
                hc_pre_eps,
                hc_sinkhorn_eps,
                hc_post_mult_value,
                sinkhorn_repeat,
                n_splits,
                norm_weight,
                norm_eps,
            )
        else:
            post_mix, comb_mix, layer_input = self.forward_native(
                residual,
                fn,
                hc_scale,
                hc_base,
                rms_eps,
                hc_pre_eps,
                hc_sinkhorn_eps,
                hc_post_mult_value,
                sinkhorn_repeat,
                n_splits,
                norm_weight,
                norm_eps,
            )
            return (
                post_mix,
                comb_mix,
                _apply_mhc_norm(layer_input, norm_weight, norm_eps),
            )

    def forward_native(
        self,
        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,
        n_splits: int = 1,
        norm_weight: torch.Tensor | None = None,
        norm_eps: float = 0.0,
    ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
        return mhc_kernels.mhc_pre_torch(
            residual,
            fn,
            hc_scale,
            hc_base,
            rms_eps,
            hc_pre_eps,
            hc_sinkhorn_eps,
            hc_post_mult_value,
            sinkhorn_repeat,
        )

    def forward_xpu(
        self,
        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,
        n_splits: int = 1,
        norm_weight: torch.Tensor | None = None,
        norm_eps: float = 0.0,
    ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
        return torch.ops._xpu_C.mhc_pre(
            residual,
            fn,
            hc_scale,
            hc_base,
            rms_eps,
            hc_pre_eps,
            hc_sinkhorn_eps,
            hc_post_mult_value,
            sinkhorn_repeat,
        )

    def forward_cpu(
        self,
        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,
        n_splits: int = 1,
        norm_weight: torch.Tensor | None = None,
        norm_eps: float = 0.0,
    ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
        return mhc_kernels.mhc_pre_cpu(
            residual,
            fn,
            hc_scale,
            hc_base,
            rms_eps,
            hc_pre_eps,
            hc_sinkhorn_eps,
            hc_post_mult_value,
            sinkhorn_repeat,
            n_splits,
            norm_weight,
            norm_eps,
        )

hc_contract(x, n)

[s, n * hidden_size] -> [s, hidden_size] by averaging.

Source code in vllm/model_executor/layers/mhc.py
def hc_contract(x: torch.Tensor, n: int) -> torch.Tensor:
    """[s, n * hidden_size] -> [s, hidden_size] by averaging."""
    return x.mean(dim=1)

hc_expand(x, n)

[s, hidden_size] -> [s, n * hidden_size] by replication.

Source code in vllm/model_executor/layers/mhc.py
def hc_expand(x: torch.Tensor, n: int) -> torch.Tensor:
    """[s, hidden_size] -> [s, n * hidden_size] by replication."""
    return x.unsqueeze(1).expand(-1, n, -1).contiguous()