diff --git a/tests/kernels/test_mhc_kernels.py b/tests/kernels/test_mhc_kernels.py index e7d4cde43f1..0e0e3769f49 100644 --- a/tests/kernels/test_mhc_kernels.py +++ b/tests/kernels/test_mhc_kernels.py @@ -340,22 +340,17 @@ def test_hc_head_tilelang(num_tokens, hidden_size, hc_mult): hc_base = torch.randn((hc_mult,), dtype=torch.float32) * 0.1 rms_eps = hc_eps = 1e-6 - out = torch.empty((num_tokens, hidden_size), dtype=torch.bfloat16) - out.fill_(float("nan")) - - result = torch.ops.vllm.hc_head_fused_kernel_tilelang( + out = torch.ops.vllm.hc_head_fused_kernel_tilelang( residual, fn, hc_scale, hc_base, - out, - hidden_size, rms_eps, hc_eps, - hc_mult, ) - assert result is None + assert out.shape == (num_tokens, hidden_size) + assert out.dtype == torch.bfloat16 assert not torch.isnan(out).any() out_ref = hc_head_ref(residual, fn, hc_scale, hc_base, rms_eps, hc_eps) diff --git a/vllm/model_executor/kernels/mhc/tilelang.py b/vllm/model_executor/kernels/mhc/tilelang.py index d76123bb762..e0007141d53 100644 --- a/vllm/model_executor/kernels/mhc/tilelang.py +++ b/vllm/model_executor/kernels/mhc/tilelang.py @@ -29,7 +29,7 @@ def _tilelang_hc_prenorm_gemm( n_thr: int = 512, n_splits: int = 1, ) -> None: - from vllm._tilelang_ops import ( + from vllm.model_executor.kernels.mhc.tilelang_kernels import ( hc_prenorm_gemm_block_m_tilelang, hc_prenorm_gemm_tilelang, ) @@ -126,7 +126,7 @@ def mhc_pre_tilelang( comb_mix: shape (..., hc_mult, hc_mult), dtype torch.float32 layer_input: shape (..., hidden_size), dtype torch.bfloat16 """ - from vllm._tilelang_ops import ( + from vllm.model_executor.kernels.mhc.tilelang_kernels import ( compute_num_split, mhc_pre_big_fuse_tilelang, mhc_pre_big_fuse_with_norm_tilelang, @@ -306,7 +306,9 @@ def mhc_post_tilelang( post_layer_mix: torch.Tensor, comb_res_mix: torch.Tensor, ) -> torch.Tensor: - from vllm._tilelang_ops import mhc_post_tilelang as _mhc_post_kernel + from vllm.model_executor.kernels.mhc.tilelang_kernels import ( + mhc_post_tilelang as _mhc_post_kernel, + ) out = torch.empty_like(residual) _mhc_post_kernel( @@ -353,7 +355,7 @@ def mhc_fused_post_pre_tilelang( layer_input_cur: shape (..., hidden_size) """ - from vllm._tilelang_ops import ( + from vllm.model_executor.kernels.mhc.tilelang_kernels import ( compute_num_split, mhc_fused_tilelang, mhc_post_tilelang, @@ -608,21 +610,22 @@ def _mhc_post_tilelang_fake( return torch.empty_like(residual) -def _hc_head_fused_kernel_tilelang( +def hc_head_fused_kernel_tilelang( hs_flat: torch.Tensor, fn: torch.Tensor, hc_scale: torch.Tensor, hc_base: torch.Tensor, - out: torch.Tensor, - hidden_size: int, rms_eps: float, hc_eps: float, - hc_mult: int, -) -> None: - """Fill pre-allocated `out` (T, H) in-place with the hc_head result.""" - if hs_flat.shape[0] == 0: - return - from vllm._tilelang_ops import hc_head_fuse_tilelang +) -> torch.Tensor: + """Apply the fused hc_head kernel and return the (T, H) bf16 result.""" + num_tokens, hc_mult, hidden_size = hs_flat.shape + out = torch.empty( + num_tokens, hidden_size, dtype=torch.bfloat16, device=hs_flat.device + ) + if num_tokens == 0: + return out + from vllm.model_executor.kernels.mhc.tilelang_kernels import hc_head_fuse_tilelang hc_head_fuse_tilelang( hs_flat, @@ -635,6 +638,21 @@ def _hc_head_fused_kernel_tilelang( hc_eps, hc_mult, ) + return out + + +def _hc_head_fused_kernel_tilelang_fake( + hs_flat: torch.Tensor, + fn: torch.Tensor, + hc_scale: torch.Tensor, + hc_base: torch.Tensor, + rms_eps: float, + hc_eps: float, +) -> torch.Tensor: + num_tokens, _, hidden_size = hs_flat.shape + return torch.empty( + num_tokens, hidden_size, dtype=torch.bfloat16, device=hs_flat.device + ) direct_register_custom_op( @@ -659,6 +677,7 @@ direct_register_custom_op( direct_register_custom_op( op_name="hc_head_fused_kernel_tilelang", - op_func=_hc_head_fused_kernel_tilelang, - mutates_args=["out"], + op_func=hc_head_fused_kernel_tilelang, + mutates_args=[], + fake_impl=_hc_head_fused_kernel_tilelang_fake, ) diff --git a/vllm/_tilelang_ops.py b/vllm/model_executor/kernels/mhc/tilelang_kernels.py similarity index 100% rename from vllm/_tilelang_ops.py rename to vllm/model_executor/kernels/mhc/tilelang_kernels.py diff --git a/vllm/model_executor/layers/mhc.py b/vllm/model_executor/layers/mhc.py index b720fa1f6fe..5249481293a 100644 --- a/vllm/model_executor/layers/mhc.py +++ b/vllm/model_executor/layers/mhc.py @@ -243,21 +243,13 @@ class HCHeadOp(CustomOp): 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.vllm.hc_head_fused_kernel_tilelang( + out = torch.ops.vllm.hc_head_fused_kernel_tilelang( 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) @@ -273,25 +265,24 @@ class HCHeadOp(CustomOp): 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 - ) if HAS_TILELANG: - torch.ops.vllm.hc_head_fused_kernel_tilelang( + out = torch.ops.vllm.hc_head_fused_kernel_tilelang( hs_flat, hc_fn, hc_scale, hc_base, - out, - hidden_size, rms_norm_eps, hc_eps, - hc_mult, ) 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, diff --git a/vllm/models/deepseek_v4/nvidia/model.py b/vllm/models/deepseek_v4/nvidia/model.py index 6ade4caf9d9..30a7e6e747f 100644 --- a/vllm/models/deepseek_v4/nvidia/model.py +++ b/vllm/models/deepseek_v4/nvidia/model.py @@ -15,6 +15,12 @@ from vllm.distributed import ( get_tensor_model_parallel_rank, get_tensor_model_parallel_world_size, ) +from vllm.model_executor.kernels.mhc.tilelang import ( + hc_head_fused_kernel_tilelang, + mhc_fused_post_pre_tilelang, + mhc_post_tilelang, + mhc_pre_tilelang, +) from vllm.model_executor.layers.activation import SiluAndMul, SiluAndMulWithClamp from vllm.model_executor.layers.fused_moe import FusedMoE from vllm.model_executor.layers.fused_moe.router.fused_topk_bias_router import ( @@ -28,12 +34,6 @@ from vllm.model_executor.layers.linear import ( RowParallelLinear, ) from vllm.model_executor.layers.logits_processor import LogitsProcessor -from vllm.model_executor.layers.mhc import ( - HCHeadOp, - MHCFusedPostPreOp, - MHCPostOp, - MHCPreOp, -) from vllm.model_executor.layers.quantization import QuantizationConfig from vllm.model_executor.layers.rotary_embedding import get_rope from vllm.model_executor.layers.vocab_parallel_embedding import ( @@ -794,10 +794,6 @@ class DeepseekV4DecoderLayer(nn.Module): ): super().__init__() - # Lazy import to avoid top-level tilelang dependency. - # Registers both torch.ops.vllm.mhc_pre and mhc_post - import vllm.model_executor.layers.mhc # noqa: F401 - config = vllm_config.model_config.hf_config self.hidden_size = config.hidden_size @@ -860,42 +856,6 @@ class DeepseekV4DecoderLayer(nn.Module): ), requires_grad=False, ) - self.mhc_pre = MHCPreOp() - self.mhc_post = MHCPostOp() - self.mhc_fused_post_pre = MHCFusedPostPreOp() - - def hc_pre( - self, - x: torch.Tensor, - hc_fn: torch.Tensor, - hc_scale: torch.Tensor, - hc_base: torch.Tensor, - norm_weight: torch.Tensor | None = None, - norm_eps: float = 1e-6, - ): - post_mix, res_mix, layer_input = self.mhc_pre( - residual=x, - fn=hc_fn, - hc_scale=hc_scale, - hc_base=hc_base, - rms_eps=self.rms_norm_eps, - hc_pre_eps=self.hc_eps, - hc_sinkhorn_eps=self.hc_eps, - hc_post_mult_value=self.hc_post_alpha, - sinkhorn_repeat=self.hc_sinkhorn_iters, - norm_weight=norm_weight, - norm_eps=norm_eps, - ) - return layer_input, post_mix, res_mix - - def hc_post( - self, - x: torch.Tensor, - residual: torch.Tensor, - post: torch.Tensor, - comb: torch.Tensor, - ): - return self.mhc_post(x, residual, post, comb) def forward( self, @@ -909,18 +869,23 @@ class DeepseekV4DecoderLayer(nn.Module): attn_norm_weight = self.attn_norm.weight.data attn_norm_eps = self.attn_norm.variance_epsilon if residual is None: - # Run standalone hc_pre on first layer + # Run standalone mhc_pre on first layer residual = x - x, post_mix, res_mix = self.hc_pre( + post_mix, res_mix, x = mhc_pre_tilelang( x, self.hc_attn_fn, self.hc_attn_scale, self.hc_attn_base, + self.rms_norm_eps, + self.hc_eps, + self.hc_eps, + self.hc_post_alpha, + self.hc_sinkhorn_iters, norm_weight=attn_norm_weight, norm_eps=attn_norm_eps, ) else: - residual, post_mix, res_mix, x = self.mhc_fused_post_pre( + residual, post_mix, res_mix, x = mhc_fused_post_pre_tilelang( x, residual, post_mix, @@ -939,12 +904,12 @@ class DeepseekV4DecoderLayer(nn.Module): norm_eps=attn_norm_eps, ) - # attn_norm is fused into hc_pre / mhc_fused_post_pre above. + # attn_norm is fused into mhc_pre_tilelang / mhc_fused_post_pre above. x = self.attn(positions, x, None) ffn_norm_weight = self.ffn_norm.weight.data ffn_norm_eps = self.ffn_norm.variance_epsilon - residual, post_mix, res_mix, x = self.mhc_fused_post_pre( + residual, post_mix, res_mix, x = mhc_fused_post_pre_tilelang( x, residual, post_mix, @@ -1047,7 +1012,6 @@ class DeepseekV4Model(nn.Module): torch.empty(1, dtype=torch.float32), requires_grad=False, ) - self.hc_head_op = HCHeadOp() # Pre-hc_head residual stream buffer for the MTP draft. Stable # address (outside the cudagraph pool) so the copy_ in forward() # refreshes it correctly across captured shapes. @@ -1117,7 +1081,9 @@ class DeepseekV4Model(nn.Module): residual, ) if layer is not None: - hidden_states = layer.hc_post(hidden_states, residual, post_mix, res_mix) + hidden_states = mhc_post_tilelang( + hidden_states, residual, post_mix, res_mix + ) if not get_pp_group().is_last_rank: return IntermediateTensors({"hidden_states": hidden_states}) @@ -1126,7 +1092,7 @@ class DeepseekV4Model(nn.Module): num_tokens = hidden_states.shape[0] self._mtp_hidden_buffer[:num_tokens].copy_(hidden_states.flatten(1)) - hidden_states = self.hc_head_op( + hidden_states = hc_head_fused_kernel_tilelang( hidden_states, self.hc_head_fn, self.hc_head_scale, diff --git a/vllm/models/deepseek_v4/nvidia/mtp.py b/vllm/models/deepseek_v4/nvidia/mtp.py index 3db831f0261..133a96e3acd 100644 --- a/vllm/models/deepseek_v4/nvidia/mtp.py +++ b/vllm/models/deepseek_v4/nvidia/mtp.py @@ -24,11 +24,14 @@ from vllm.distributed import ( get_tensor_model_parallel_world_size, ) from vllm.logger import init_logger +from vllm.model_executor.kernels.mhc.tilelang import ( + hc_head_fused_kernel_tilelang, + mhc_post_tilelang, +) from vllm.model_executor.layers.fused_moe import FusedMoE from vllm.model_executor.layers.layernorm import RMSNorm from vllm.model_executor.layers.linear import ReplicatedLinear from vllm.model_executor.layers.logits_processor import LogitsProcessor -from vllm.model_executor.layers.mhc import HCHeadOp from vllm.model_executor.layers.vocab_parallel_embedding import ( VocabParallelEmbedding, ) @@ -122,8 +125,6 @@ class DeepSeekV4MultiTokenPredictorLayer(nn.Module): aux_stream_list=aux_stream_list, ) - self.hc_head_op = HCHeadOp() - def forward( self, input_ids: torch.Tensor, @@ -155,9 +156,7 @@ class DeepSeekV4MultiTokenPredictorLayer(nn.Module): hidden_states, residual, post_mix, res_mix = self.mtp_block( positions=positions, x=hidden_states, input_ids=None ) - hidden_states = self.mtp_block.hc_post( - hidden_states, residual, post_mix, res_mix - ) + hidden_states = mhc_post_tilelang(hidden_states, residual, post_mix, res_mix) # Return the flat pre-hc_head residual so it can be re-fed as the # next spec step's `previous_hidden_states` when # num_speculative_tokens > 1. hc_head is deferred to compute_logits. @@ -237,7 +236,7 @@ class DeepSeekV4MultiTokenPredictor(nn.Module): hidden_states = hidden_states.view( -1, mtp_layer.hc_mult, mtp_layer.config.hidden_size ) - hidden_states = mtp_layer.hc_head_op( + hidden_states = hc_head_fused_kernel_tilelang( hidden_states, mtp_layer.hc_head_fn, mtp_layer.hc_head_scale,