mirror of
https://github.com/vllm-project/vllm.git
synced 2026-08-13 09:18:12 +00:00
[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:
@@ -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)
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user