[DSv4] Move mHC tilelang kernels & Don't use CustomOP in dsv4/nvidia (#43905)

Signed-off-by: Woosuk Kwon <[email protected]>
This commit is contained in:
Woosuk Kwon
2026-05-29 10:25:02 +08:00
committed by GitHub
parent bf18d7e0b4
commit 7bd45da585
6 changed files with 72 additions and 102 deletions
+3 -8
View File
@@ -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)
+34 -15
View File
@@ -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,
)
+9 -18
View File
@@ -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,
+20 -54
View File
@@ -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,
+6 -7
View File
@@ -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,