vllm.model_executor.kernels.mhc.aiter ¶
Functions:
-
mhc_fused_post_pre_aiter–Fused mHC post + next mHC pre on ROCm via AITER.
-
mhc_pre_aiter–Forward pass for mHC pre block.
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=1, tile_n=1, norm_weight=None, norm_eps=0.0) ¶
Fused mHC post + next mHC pre on ROCm via AITER.
Returns residual_cur, post_mix_cur, comb_mix_cur, layer_input_cur.
Source code in vllm/model_executor/kernels/mhc/aiter.py
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=1, norm_weight=None, norm_eps=0.0) ¶
Forward pass for mHC pre block.
Parameters:
-
(residual¶Tensor) –shape (..., hc_mult, hidden_size), dtype torch.bfloat16
-
(fn¶Tensor) –shape (hc_mult3, hc_mult * hidden_size), dtype torch.float32
-
(hc_scale¶Tensor) –shape (3,), dtype torch.float32
-
(hc_base¶Tensor) –shape (hc_mult3,), dtype torch.float32
-
(rms_eps¶float) –RMS normalization epsilon
-
(hc_pre_eps¶float) –pre-mix epsilon
-
(hc_sinkhorn_eps¶float) –sinkhorn epsilon
-
(hc_post_mult_value¶float) –post-mix multiplier value
-
(sinkhorn_repeat¶int) –number of sinkhorn iterations
-
(n_splits¶int, default:1) –split-k factor;
-
(norm_weight¶Tensor | None, default:None) –optional RMSNorm weight fused into the pre kernel
-
(norm_eps¶float, default:0.0) –epsilon for the fused RMSNorm when norm_weight is set
Returns:
-
post_mix(Tensor) –shape (..., hc_mult), dtype torch.float32
-
comb_mix(Tensor) –shape (..., hc_mult, hc_mult), dtype torch.float32
-
layer_input(Tensor) –shape (..., hidden_size), dtype torch.bfloat16