diff --git a/tests/models/test_deepseek_v4_mega_moe.py b/tests/models/test_deepseek_v4_mega_moe.py index 58d5239285dd..7d95df8a8d8b 100644 --- a/tests/models/test_deepseek_v4_mega_moe.py +++ b/tests/models/test_deepseek_v4_mega_moe.py @@ -17,7 +17,9 @@ make_deepseek_v4_expert_params_mapping, ) from vllm.models.deepseek_v4.nvidia.mtp import DeepSeekV4MTP -from vllm.models.deepseek_v4.nvidia.ops.prepare_megamoe import prepare_megamoe_inputs +from vllm.models.deepseek_v4.nvidia.ops.prepare_megamoe import ( + prepare_megamoe_inputs, +) from vllm.models.deepseek_v4_1.common.mm_preprocess import ( IMAGE_PAD_ID, IMAGE_SENTINEL_BASE_ID, @@ -55,7 +57,10 @@ def v41_moe_config(dist_init): ), ), quant_config=None, - kernel_config=SimpleNamespace(moe_backend="deep_gemm_mega_moe"), + kernel_config=SimpleNamespace( + moe_backend="deep_gemm_mega_moe", + enable_jit_warmup=True, + ), parallel_config=SimpleNamespace( enable_expert_parallel=True, enable_eplb=False, @@ -129,6 +134,7 @@ def test_deepseek_v41_fused_moe_uses_draft_counts_or_main_defaults( config.dspark_n_routed_experts = draft_experts config.dspark_num_experts_per_tok = draft_top_k v41_moe_config.kernel_config.moe_backend = "auto" + v41_moe_config.kernel_config.enable_jit_warmup = False captured = {} def make_experts(**kwargs): diff --git a/vllm/model_executor/layers/fused_moe/deep_gemm_utils.py b/vllm/model_executor/layers/fused_moe/deep_gemm_utils.py index 70158dafa1c7..a41bc6c5c00c 100644 --- a/vllm/model_executor/layers/fused_moe/deep_gemm_utils.py +++ b/vllm/model_executor/layers/fused_moe/deep_gemm_utils.py @@ -6,13 +6,26 @@ """ import math +from typing import Any import torch import vllm.model_executor.layers.fused_moe.modular_kernel as mk from vllm.model_executor.layers.fused_moe.utils import count_expert_num_tokens +from vllm.model_executor.warmup.jit_warmup import ( + WarmupChoices, + WarmupIntRange, +) +from vllm.model_executor.warmup.jit_warmup_triton_helper import ( + DispatchSpec, + TritonWarmupTensor, + triton_kernel_dispatcher_with_warmup, +) from vllm.triton_utils import tl, triton -from vllm.utils.deep_gemm import get_mk_alignment_for_contiguous_layout +from vllm.utils.deep_gemm import ( + get_mk_alignment_for_contiguous_layout, + get_theoretical_mk_alignment_for_contiguous_layout, +) from vllm.utils.math_utils import round_up @@ -111,6 +124,22 @@ def apply_expert_map(expert_id, expert_map): return expert_id +def _deep_gemm_local_num_experts(vllm_config: Any) -> int: + num_experts = vllm_config.model_config.hf_config.n_routed_experts + parallel_config = vllm_config.parallel_config + eplb_config = getattr(parallel_config, "eplb_config", None) + num_experts += int(getattr(eplb_config, "num_redundant_experts", 0) or 0) + if parallel_config.enable_expert_parallel: + try: + from vllm.distributed.parallel_state import get_ep_group + + world_size = get_ep_group().world_size + except Exception: + world_size = parallel_config.data_parallel_size + num_experts //= max(world_size, 1) + return num_experts + + @triton.jit def _fwd_kernel_ep_scatter_1( num_recv_tokens_per_expert, @@ -157,6 +186,52 @@ def _fwd_kernel_ep_scatter_1( ) +def _deepgemm_ep_scatter_start_kernel_warmup_inputs(vllm_config: Any) -> dict[str, Any]: + num_experts = _deep_gemm_local_num_experts(vllm_config) + top_k = vllm_config.model_config.hf_config.num_experts_per_tok + num_tokens: Any = WarmupIntRange( + 1, vllm_config.scheduler_config.max_num_batched_tokens + 1 + ) + max_align_m = get_mk_alignment_for_contiguous_layout()[0] + align_m = min( + get_theoretical_mk_alignment_for_contiguous_layout( + expected_m=num_tokens * top_k, num_groups=num_experts + ) + or max_align_m, + max_align_m, + ) + return dict( + num_recv_tokens_per_expert=TritonWarmupTensor( + torch.int32, shape=(num_experts,) + ), + expert_start_loc=TritonWarmupTensor(torch.int32), + m_indices=TritonWarmupTensor(torch.int32), + align_m=align_m, + ) + + +@triton_kernel_dispatcher_with_warmup( + kernel=_fwd_kernel_ep_scatter_1, + warmup_inputs=_deepgemm_ep_scatter_start_kernel_warmup_inputs, +) +def _DEEPGEMM_EP_SCATTER_START_KERNEL( + num_recv_tokens_per_expert: torch.Tensor, + expert_start_loc: torch.Tensor, + m_indices: torch.Tensor, + *, + align_m: int, +) -> DispatchSpec: + num_experts = num_recv_tokens_per_expert.shape[0] + # BLOCK_E is the m_indices fill-loop tile (masked), independent of align_m. + return (num_experts,), dict( + num_experts=num_experts, + num_warps=8, + BLOCK_E=128, + BLOCK_EXPERT_NUM=triton.next_power_of_2(num_experts), + ALIGN_M=align_m, + ) + + @triton.jit def _fwd_kernel_ep_scatter_2( total_token_num, @@ -271,76 +346,81 @@ def _fwd_kernel_ep_scatter_2( ) -@torch.no_grad() -def ep_scatter( +def _deepgemm_ep_scatter_copy_kernel_warmup_inputs(vllm_config: Any) -> dict[str, Any]: + hidden_size = vllm_config.model_config.hf_config.hidden_size + topk_num = vllm_config.model_config.hf_config.num_experts_per_tok + total_token_num: Any = WarmupIntRange( + 1, vllm_config.scheduler_config.max_num_batched_tokens + 1 + ) + pack_ue8m0: Any = WarmupChoices(False, True) + has_expert_map: Any = WarmupChoices(False, True) + default_block_size = get_mk_alignment_for_contiguous_layout()[1] + block_size = 32 if pack_ue8m0 else default_block_size + scale_hidden_size = hidden_size // block_size + scale_dtype = torch.int32 if pack_ue8m0 else torch.float32 + output_scale_strides = (1, 16) if pack_ue8m0 else (scale_hidden_size, 1) + return dict( + recv_x=TritonWarmupTensor( + torch.float8_e4m3fn, shape=(total_token_num, hidden_size) + ), + recv_x_scale=TritonWarmupTensor( + scale_dtype, shape=(total_token_num, scale_hidden_size) + ), + recv_topk=TritonWarmupTensor(torch.int32, shape=(total_token_num, topk_num)), + expert_map=TritonWarmupTensor(torch.int32) if has_expert_map else None, + expert_start_loc=TritonWarmupTensor(torch.int32), + output_tensor=TritonWarmupTensor( + torch.float8_e4m3fn, shape=(total_token_num, hidden_size) + ), + output_tensor_scale=TritonWarmupTensor( + scale_dtype, + shape=(total_token_num, scale_hidden_size), + strides=output_scale_strides, + ), + output_index=TritonWarmupTensor(torch.int32, shape=(total_token_num, topk_num)), + block_size=block_size, + pack_ue8m0=pack_ue8m0, + ) + + +@triton_kernel_dispatcher_with_warmup( + kernel=_fwd_kernel_ep_scatter_2, + warmup_inputs=_deepgemm_ep_scatter_copy_kernel_warmup_inputs, +) +def _DEEPGEMM_EP_SCATTER_COPY_KERNEL( recv_x: torch.Tensor, recv_x_scale: torch.Tensor, recv_topk: torch.Tensor, - num_recv_tokens_per_expert: torch.Tensor, expert_map: torch.Tensor | None, expert_start_loc: torch.Tensor, output_tensor: torch.Tensor, output_tensor_scale: torch.Tensor, - m_indices: torch.Tensor, output_index: torch.Tensor, - align_m: int = 128, - block_size: int = 128, - pack_ue8m0: bool = False, -): - # BLOCK_E is the m_indices fill-loop tile (masked), independent of align_m. - BLOCK_E = 128 - BLOCK_D = block_size # block size of activation-scale quantization - num_warps = 8 - num_experts = num_recv_tokens_per_expert.shape[0] + *, + block_size: int, + pack_ue8m0: bool, +) -> DispatchSpec: hidden_size = recv_x.shape[1] - # grid = (triton.cdiv(hidden_size, BLOCK_D), num_experts) - grid = num_experts - - assert m_indices.shape[0] % align_m == 0 - assert expert_start_loc.shape[0] == num_experts - # pack_ue8m0: scatter packs 4 UE8M0 bytes per int32; else copies scales as-is. - scale_hidden_size = hidden_size // BLOCK_D + scale_hidden_size = hidden_size // block_size scale_packed_size = (scale_hidden_size + 3) // 4 if pack_ue8m0 else 1 - - _fwd_kernel_ep_scatter_1[(grid,)]( - num_recv_tokens_per_expert, - expert_start_loc, - m_indices, - num_experts=num_experts, - num_warps=num_warps, - BLOCK_E=BLOCK_E, - BLOCK_EXPERT_NUM=triton.next_power_of_2(num_experts), - ALIGN_M=align_m, - ) - - grid = min(recv_topk.shape[0], 1024 * 8) - - _fwd_kernel_ep_scatter_2[(grid,)]( - recv_topk.shape[0], - expert_start_loc, - recv_x, - recv_x.stride(0), - recv_x.stride(1), - recv_x_scale, - recv_x_scale.stride(0), - recv_x_scale.stride(1), - recv_topk, - recv_topk.stride(0), - recv_topk.stride(1), - output_tensor, - output_tensor.stride(0), - output_tensor.stride(1), - output_tensor_scale, - output_tensor_scale.stride(0), - output_tensor_scale.stride(1), - output_index, - output_index.stride(0), - output_index.stride(1), + return (min(recv_topk.shape[0], 1024 * 8),), dict( + total_token_num=recv_topk.shape[0], + recv_x_stride0=recv_x.stride(0), + recv_x_stride1=recv_x.stride(1), + recv_x_scale_stride0=recv_x_scale.stride(0), + recv_x_scale_stride1=recv_x_scale.stride(1), + recv_topk_stride0=recv_topk.stride(0), + recv_topk_stride1=recv_topk.stride(1), + output_tensor_stride0=output_tensor.stride(0), + output_tensor_stride1=output_tensor.stride(1), + output_tensor_scale_stride0=output_tensor_scale.stride(0), + output_tensor_scale_stride1=output_tensor_scale.stride(1), + output_index_stride0=output_index.stride(0), + output_index_stride1=output_index.stride(1), topk_num=recv_topk.shape[1], - expert_map=expert_map, HAS_EXPERT_MAP=expert_map is not None, - num_warps=num_warps, + num_warps=8, HIDDEN_SIZE=hidden_size, HIDDEN_SIZE_PAD=triton.next_power_of_2(hidden_size), SCALE_HIDDEN_SIZE=scale_hidden_size, @@ -349,7 +429,52 @@ def ep_scatter( SCALE_PACKED_SIZE=scale_packed_size, SCALE_PACKED_SIZE_PAD=triton.next_power_of_2(scale_packed_size), ) - return + + +class DeepGemmEPScatter: + def __init__( + self, + *, + start: Any, + copy: Any, + ) -> None: + self.start = start + self.copy = copy + + def __call__( + self, + recv_x: torch.Tensor, + recv_x_scale: torch.Tensor, + recv_topk: torch.Tensor, + num_recv_tokens_per_expert: torch.Tensor, + expert_map: torch.Tensor | None, + expert_start_loc: torch.Tensor, + output_tensor: torch.Tensor, + output_tensor_scale: torch.Tensor, + m_indices: torch.Tensor, + output_index: torch.Tensor, + align_m: int, + block_size: int, + pack_ue8m0: bool, + ) -> None: + self.start( + num_recv_tokens_per_expert, + expert_start_loc, + m_indices, + align_m=align_m, + ) + self.copy( + recv_x, + recv_x_scale, + recv_topk, + expert_map, + expert_start_loc, + output_tensor, + output_tensor_scale, + output_index, + block_size=block_size, + pack_ue8m0=pack_ue8m0, + ) @triton.jit @@ -414,44 +539,119 @@ def _fwd_kernel_ep_gather( ) -@torch.no_grad() -def ep_gather( +def _deepgemm_ep_gather_kernel_warmup_inputs(vllm_config: Any) -> dict[str, Any]: + hidden_size = vllm_config.model_config.hf_config.hidden_size + topk_num = vllm_config.model_config.hf_config.num_experts_per_tok + dtype = vllm_config.model_config.dtype + total_token_num: Any = WarmupIntRange( + 1, vllm_config.scheduler_config.max_num_batched_tokens + 1 + ) + has_expert_map: Any = WarmupChoices(False, True) + topk_shape = (total_token_num, topk_num) + return dict( + input_tensor=TritonWarmupTensor(dtype, shape=(total_token_num, hidden_size)), + recv_topk_ids=TritonWarmupTensor(torch.int32, shape=topk_shape), + recv_topk_weight=TritonWarmupTensor(torch.float32, shape=topk_shape), + input_index=TritonWarmupTensor(torch.int32, shape=topk_shape), + expert_map=TritonWarmupTensor(torch.int32) if has_expert_map else None, + output_tensor=TritonWarmupTensor(dtype, shape=(total_token_num, hidden_size)), + ) + + +@triton_kernel_dispatcher_with_warmup( + kernel=_fwd_kernel_ep_gather, + warmup_inputs=_deepgemm_ep_gather_kernel_warmup_inputs, +) +def _DEEPGEMM_EP_GATHER_KERNEL( input_tensor: torch.Tensor, recv_topk_ids: torch.Tensor, recv_topk_weight: torch.Tensor, input_index: torch.Tensor, expert_map: torch.Tensor | None, output_tensor: torch.Tensor, -): +) -> DispatchSpec: num_warps = 2 num_tokens = output_tensor.shape[0] hidden_size = input_tensor.shape[1] - BLOCK_D = math.gcd(hidden_size, 1024) - assert hidden_size % BLOCK_D == 0 - grid = (triton.cdiv(hidden_size, BLOCK_D), min(num_tokens, 1024)) + block_d = math.gcd(hidden_size, 1024) + assert hidden_size % block_d == 0 + grid = (triton.cdiv(hidden_size, block_d), min(num_tokens, 1024)) + + return grid, dict( + total_token_num=num_tokens, + input_tensor_stride0=input_tensor.stride(0), + input_tensor_stride1=input_tensor.stride(1), + recv_topk_ids_stride0=recv_topk_ids.stride(0), + recv_topk_ids_stride1=recv_topk_ids.stride(1), + recv_topk_weight_stride0=recv_topk_weight.stride(0), + recv_topk_weight_stride1=recv_topk_weight.stride(1), + input_index_stride0=input_index.stride(0), + input_index_stride1=input_index.stride(1), + output_tensor_stride0=output_tensor.stride(0), + output_tensor_stride1=output_tensor.stride(1), + topk_num=recv_topk_ids.shape[1], + HAS_EXPERT_MAP=expert_map is not None, + num_warps=num_warps, + BLOCK_D=block_d, + ) + + +@torch.no_grad() +def ep_scatter( + recv_x: torch.Tensor, + recv_x_scale: torch.Tensor, + recv_topk: torch.Tensor, + num_recv_tokens_per_expert: torch.Tensor, + expert_map: torch.Tensor | None, + expert_start_loc: torch.Tensor, + output_tensor: torch.Tensor, + output_tensor_scale: torch.Tensor, + m_indices: torch.Tensor, + output_index: torch.Tensor, + align_m: int = 128, + block_size: int = 128, + pack_ue8m0: bool = False, +): + block_d = block_size # block size of activation-scale quantization + num_experts = num_recv_tokens_per_expert.shape[0] - _fwd_kernel_ep_gather[grid]( - num_tokens, + assert m_indices.shape[0] % align_m == 0 + assert expert_start_loc.shape[0] == num_experts + + _DEEPGEMM_EP_SCATTER( + recv_x, + recv_x_scale, + recv_topk, + num_recv_tokens_per_expert, + expert_map, + expert_start_loc, + output_tensor, + output_tensor_scale, + m_indices, + output_index, + align_m, + block_d, + pack_ue8m0, + ) + return + + +@torch.no_grad() +def ep_gather( + input_tensor: torch.Tensor, + recv_topk_ids: torch.Tensor, + recv_topk_weight: torch.Tensor, + input_index: torch.Tensor, + expert_map: torch.Tensor | None, + output_tensor: torch.Tensor, +): + _DEEPGEMM_EP_GATHER_KERNEL( input_tensor, - input_tensor.stride(0), - input_tensor.stride(1), recv_topk_ids, - recv_topk_ids.stride(0), - recv_topk_ids.stride(1), recv_topk_weight, - recv_topk_weight.stride(0), - recv_topk_weight.stride(1), input_index, - input_index.stride(0), - input_index.stride(1), + expert_map, output_tensor, - output_tensor.stride(0), - output_tensor.stride(1), - topk_num=recv_topk_ids.shape[1], - expert_map=expert_map, - HAS_EXPERT_MAP=expert_map is not None, - num_warps=num_warps, - BLOCK_D=BLOCK_D, ) return @@ -567,3 +767,9 @@ def deepgemm_unpermute_and_reduce( expert_map=expert_map, output_tensor=output, ) + + +_DEEPGEMM_EP_SCATTER = DeepGemmEPScatter( + start=_DEEPGEMM_EP_SCATTER_START_KERNEL, + copy=_DEEPGEMM_EP_SCATTER_COPY_KERNEL, +) diff --git a/vllm/model_executor/layers/fused_moe/experts/fused_batched_moe.py b/vllm/model_executor/layers/fused_moe/experts/fused_batched_moe.py index f16b351d353d..0aef0e3f5a1a 100644 --- a/vllm/model_executor/layers/fused_moe/experts/fused_batched_moe.py +++ b/vllm/model_executor/layers/fused_moe/experts/fused_batched_moe.py @@ -2,6 +2,8 @@ # SPDX-FileCopyrightText: Copyright contributors to the vLLM project """Fused batched MoE kernel.""" +from typing import Any + import torch import vllm.model_executor.layers.fused_moe.modular_kernel as mk @@ -13,8 +15,13 @@ FusedMoEConfig, FusedMoEParallelConfig, FusedMoEQuantConfig, + _get_config_dtype_str, +) +from vllm.model_executor.layers.fused_moe.fused_moe import ( + _triton_moe_compute_type, + _triton_moe_config, + try_get_optimal_moe_config, ) -from vllm.model_executor.layers.fused_moe.fused_moe import try_get_optimal_moe_config from vllm.model_executor.layers.fused_moe.topk_weight_and_reduce import ( TopKWeightAndReduceDelegate, ) @@ -33,6 +40,15 @@ kFp8StaticChannelSym, kFp8StaticTensorSym, ) +from vllm.model_executor.warmup.jit_warmup import ( + WarmupChoices, + WarmupIntRange, +) +from vllm.model_executor.warmup.jit_warmup_triton_helper import ( + DispatchSpec, + TritonWarmupTensor, + triton_kernel_dispatcher_with_warmup, +) from vllm.platforms import current_platform from vllm.triton_utils import tl, triton, use_tensor_descriptor from vllm.triton_utils.allocation import set_triton_allocator @@ -424,6 +440,148 @@ def batched_triton_kernel( ) +def _batched_triton_kernel_warmup_inputs(vllm_config: Any) -> dict[str, Any]: + hf = vllm_config.model_config.hf_config + hidden_size = hf.hidden_size + intermediate_size = hf.moe_intermediate_size + num_experts = hf.n_routed_experts + top_k = hf.num_experts_per_tok + max_num_tokens: Any = WarmupIntRange( + 1, vllm_config.scheduler_config.max_num_batched_tokens + 1 + ) + scenario: Any = WarmupChoices(0, 1, 2, 3) + use_td: Any = WarmupChoices(False, True) + second_gemm = scenario % 2 == 1 + use_fp8 = scenario >= 2 + n = hidden_size if second_gemm else 2 * intermediate_size + k = intermediate_size if second_gemm else hidden_size + dtype = torch.float8_e4m3fn if use_fp8 else vllm_config.model_config.dtype + group = 128 if use_fp8 else 0 + config = _triton_moe_config( + num_experts=num_experts, + hidden_size=hidden_size, + intermediate_size=intermediate_size, + top_k=top_k, + config_dtype=_get_config_dtype_str( + dtype=dtype, + use_fp8_w8a8=use_fp8, + use_int8_w8a16=False, + use_int4_w4a16=False, + ), + num_tokens=max_num_tokens, + group_n=group, + group_k=group, + ) + scale_cols = triton.cdiv(k, group) if group > 0 else 1 + return dict( + A=TritonWarmupTensor(dtype, shape=(num_experts, max_num_tokens, k)), + B=TritonWarmupTensor(dtype, shape=(num_experts, n, k)), + C=TritonWarmupTensor( + vllm_config.model_config.dtype, shape=(num_experts, max_num_tokens, n) + ), + expert_num_tokens=TritonWarmupTensor(torch.int32, shape=(num_experts,)), + compute_type=_triton_moe_compute_type(dtype), + max_num_tokens=max_num_tokens, + K=k, + N=n, + A_scale=TritonWarmupTensor( + torch.float32, shape=(num_experts, max_num_tokens, scale_cols) + ) + if use_fp8 + else None, + B_scale=TritonWarmupTensor(torch.float32, shape=(num_experts, n, scale_cols)) + if use_fp8 + else None, + B_zp=None, + stride_ae=max_num_tokens * k, + stride_am=k, + stride_ak=1, + stride_be=n * k, + stride_bk=1, + stride_bn=k, + stride_ce=max_num_tokens * n, + stride_cm=n, + stride_cn=1, + stride_ase=max_num_tokens * scale_cols, + stride_asm=scale_cols, + stride_ask=1, + stride_bse=n * scale_cols, + stride_bsk=1, + stride_bsn=scale_cols, + group_n=group, + group_k=group, + use_fp8_w8a8=use_fp8, + use_int8_w8a16=False, + per_act_token_quant=False, + BLOCK_M=config["BLOCK_SIZE_M"], + BLOCK_N=config["BLOCK_SIZE_N"], + BLOCK_K=config["BLOCK_SIZE_K"], + USE_TD=use_td, + num_warps=config["num_warps"], + num_stages=config["num_stages"], + ) + + +@triton_kernel_dispatcher_with_warmup( + kernel=batched_triton_kernel, + warmup_inputs=_batched_triton_kernel_warmup_inputs, +) +def _BATCHED_TRITON_KERNEL( + A: torch.Tensor, + B: torch.Tensor, + C: torch.Tensor, + expert_num_tokens: torch.Tensor, + compute_type: tl.dtype, + max_num_tokens: int, + K: int, + N: int, + A_scale: torch.Tensor | None, + B_scale: torch.Tensor | None, + B_zp: torch.Tensor | None, + stride_ae: int, + stride_am: int, + stride_ak: int, + stride_be: int, + stride_bk: int, + stride_bn: int, + stride_ce: int, + stride_cm: int, + stride_cn: int, + stride_ase: int, + stride_asm: int, + stride_ask: int, + stride_bse: int, + stride_bsk: int, + stride_bsn: int, + group_n: int, + group_k: int, + use_fp8_w8a8: bool, + use_int8_w8a16: bool, + per_act_token_quant: bool, + *, + BLOCK_M: int, + BLOCK_N: int, + BLOCK_K: int, + USE_TD: bool, + num_warps: int, + num_stages: int, +) -> DispatchSpec: + grid = ( + expert_num_tokens.shape[0], + triton.cdiv(max_num_tokens, BLOCK_M) * triton.cdiv(B.shape[1], BLOCK_N), + ) + return grid, dict( + a_ptr=A, + b_ptr=B, + c_ptr=C, + a_scale_ptr=A_scale, + b_scale_ptr=B_scale, + b_zp_ptr=B_zp, + num_warps=num_warps, + num_stages=num_stages, + ) + + def invoke_moe_batched_triton_kernel( A: torch.Tensor, # [E, max_tokens, K] B: torch.Tensor, # [E, N, K] @@ -433,7 +591,7 @@ def invoke_moe_batched_triton_kernel( # Quantization data A_scale: torch.Tensor | None, B_scale: torch.Tensor | None, - B_zp: torch.Tensor, + B_zp: torch.Tensor | None, # Quantization schemes use_fp8_w8a8: bool, use_int8_w8a16: bool, @@ -451,11 +609,6 @@ def invoke_moe_batched_triton_kernel( BLOCK_N = config["BLOCK_SIZE_N"] BLOCK_K = config["BLOCK_SIZE_K"] - grid = ( - expert_num_tokens.size(0), - triton.cdiv(max_num_tokens, BLOCK_M) * triton.cdiv(B.size(1), BLOCK_N), - ) - A_scale = normalize_batched_scales_shape(A_scale, expert_num_tokens.shape[0]) if B_scale is not None and B_scale.ndim == 1: @@ -505,7 +658,7 @@ def invoke_moe_batched_triton_kernel( if use_td: set_triton_allocator(A.device) - batched_triton_kernel[grid]( + _BATCHED_TRITON_KERNEL( A, B, C, @@ -547,6 +700,8 @@ def invoke_moe_batched_triton_kernel( BLOCK_N=BLOCK_N, BLOCK_K=BLOCK_K, USE_TD=use_td, + num_warps=config.get("num_warps", 4), + num_stages=config.get("num_stages", 3), ) diff --git a/vllm/model_executor/layers/fused_moe/experts/nvfp4_emulation_moe.py b/vllm/model_executor/layers/fused_moe/experts/nvfp4_emulation_moe.py index cd862cac5958..bb5daf084d4f 100644 --- a/vllm/model_executor/layers/fused_moe/experts/nvfp4_emulation_moe.py +++ b/vllm/model_executor/layers/fused_moe/experts/nvfp4_emulation_moe.py @@ -24,6 +24,9 @@ ) from vllm.model_executor.layers.fused_moe.experts.triton_moe import TritonExperts from vllm.model_executor.layers.fused_moe.fused_moe import ( + _triton_moe_compute_type, + _triton_moe_config, + _triton_moe_em, try_get_optimal_moe_config, write_zeros_to_output, ) @@ -42,6 +45,15 @@ kNvfp4Dynamic, kNvfp4Static, ) +from vllm.model_executor.warmup.jit_warmup import ( + WarmupChoices, + WarmupIntRange, +) +from vllm.model_executor.warmup.jit_warmup_triton_helper import ( + DispatchSpec, + TritonWarmupTensor, + triton_kernel_dispatcher_with_warmup, +) from vllm.triton_utils import tl, triton logger = init_logger(__name__) @@ -254,6 +266,122 @@ def fused_moe_nvfp4_emulation_kernel( tl.store(c_ptrs, accumulator, mask=c_mask) +def _fused_moe_nvfp4_emulation_kernel_warmup_inputs(vllm_config: Any) -> dict[str, Any]: + hf = vllm_config.model_config.hf_config + hidden_size = hf.hidden_size + intermediate_size = hf.moe_intermediate_size + num_experts = hf.n_routed_experts + config_top_k = hf.num_experts_per_tok + batch_tokens: Any = WarmupIntRange( + 1, vllm_config.scheduler_config.max_num_batched_tokens + 1 + ) + second_gemm: Any = WarmupChoices(False, True) + routed_multiplier = config_top_k if second_gemm else 1 + n = hidden_size if second_gemm else 2 * intermediate_size + k = intermediate_size if second_gemm else hidden_size + top_k = 1 if second_gemm else config_top_k + mul_routed_weight = second_gemm + config = _triton_moe_config( + num_experts=num_experts, + hidden_size=hidden_size, + intermediate_size=intermediate_size, + top_k=config_top_k, + config_dtype=None, + num_tokens=batch_tokens, + group_n=0, + group_k=0, + ) + block_m = config["BLOCK_SIZE_M"] + block_n = config["BLOCK_SIZE_N"] + block_k = config["BLOCK_SIZE_K"] + em = _triton_moe_em(batch_tokens * routed_multiplier, top_k, block_m, False) + valid = batch_tokens * routed_multiplier * top_k + k_packed = k // 2 + k_scale = max(1, k // 16) + return dict( + A=TritonWarmupTensor(vllm_config.model_config.dtype, shape=(1, k)), + B=TritonWarmupTensor(torch.uint8, shape=(num_experts, n, k_packed)), + C=TritonWarmupTensor( + vllm_config.model_config.dtype, + shape=(batch_tokens * routed_multiplier, top_k, n), + ), + B_scale=TritonWarmupTensor(torch.uint8, shape=(num_experts, n, k_scale)), + w_global_scale=TritonWarmupTensor(torch.float32, shape=(num_experts,)), + topk_weights=TritonWarmupTensor(torch.float32, shape=(valid,)) + if mul_routed_weight + else None, + sorted_token_ids=TritonWarmupTensor(torch.int32, shape=(em,)), + expert_ids=TritonWarmupTensor(torch.int32, shape=(triton.cdiv(em, block_m),)), + num_tokens_post_padded=TritonWarmupTensor(torch.int32), + N=n, + K=k, + EM=em, + num_valid_tokens=valid, + stride_am=k, + stride_ak=1, + stride_be=n * k_packed, + stride_bk=1, + stride_bn=k_packed, + stride_cm=n, + stride_cn=1, + stride_bse=n * k_scale, + stride_bsk=1, + stride_bsn=k_scale, + block_k_diviable=k % block_k == 0, + MUL_ROUTED_WEIGHT=mul_routed_weight, + top_k=top_k, + compute_type=_triton_moe_compute_type(vllm_config.model_config.dtype), + group_size=16, + BLOCK_SIZE_M=block_m, + BLOCK_SIZE_N=block_n, + BLOCK_SIZE_K=block_k, + GROUP_SIZE_M=config["GROUP_SIZE_M"], + ) + + +@triton_kernel_dispatcher_with_warmup( + kernel=fused_moe_nvfp4_emulation_kernel, + warmup_inputs=_fused_moe_nvfp4_emulation_kernel_warmup_inputs, +) +def _fused_moe_nvfp4_emulation( + A: torch.Tensor, + B: torch.Tensor, + C: torch.Tensor, + B_scale: torch.Tensor, + w_global_scale: torch.Tensor, + topk_weights: torch.Tensor | None, + sorted_token_ids: torch.Tensor, + expert_ids: torch.Tensor, + num_tokens_post_padded: torch.Tensor, + N: int, + K: int, + EM: int, + num_valid_tokens: int, + stride_am: int, + stride_ak: int, + stride_be: int, + stride_bk: int, + stride_bn: int, + stride_cm: int, + stride_cn: int, + stride_bse: int, + stride_bsk: int, + stride_bsn: int, + *, + block_k_diviable: bool, + MUL_ROUTED_WEIGHT: bool, + top_k: int, + compute_type: tl.dtype, + group_size: int, + BLOCK_SIZE_M: int, + BLOCK_SIZE_N: int, + BLOCK_SIZE_K: int, + GROUP_SIZE_M: int, +) -> DispatchSpec: + grid = (triton.cdiv(EM, BLOCK_SIZE_M) * triton.cdiv(N, BLOCK_SIZE_N),) + return grid, dict(a_ptr=A, b_ptr=B, c_ptr=C, b_scale_ptr=B_scale) + + def invoke_fused_moe_nvfp4_emulation_kernel( A: torch.Tensor, B: torch.Tensor, @@ -280,7 +408,6 @@ def invoke_fused_moe_nvfp4_emulation_kernel( N = B.size(1) K = A.size(1) - M = A.size(0) num_tokens = M * top_k @@ -291,11 +418,7 @@ def invoke_fused_moe_nvfp4_emulation_kernel( A.size(0) * top_k * config["BLOCK_SIZE_M"], ) - grid = lambda META: ( - triton.cdiv(EM, META["BLOCK_SIZE_M"]) * triton.cdiv(N, META["BLOCK_SIZE_N"]), - ) - - fused_moe_nvfp4_emulation_kernel[grid]( + return _fused_moe_nvfp4_emulation( A, B, C, diff --git a/vllm/model_executor/layers/fused_moe/experts/trtllm_lora_moe.py b/vllm/model_executor/layers/fused_moe/experts/trtllm_lora_moe.py index a3d50011d919..c9adf10e0fc3 100644 --- a/vllm/model_executor/layers/fused_moe/experts/trtllm_lora_moe.py +++ b/vllm/model_executor/layers/fused_moe/experts/trtllm_lora_moe.py @@ -20,6 +20,7 @@ """ from abc import abstractmethod +from typing import Any import torch @@ -38,6 +39,11 @@ from vllm.model_executor.layers.fused_moe.topk_weight_and_reduce import ( TopKWeightAndReduceNoOP, ) +from vllm.model_executor.warmup.jit_warmup_triton_helper import ( + DispatchSpec, + TritonWarmupTensor, + triton_kernel_dispatcher_with_warmup, +) from vllm.platforms import current_platform from vllm.triton_utils import tl, triton from vllm.utils.flashinfer import has_flashinfer_trtllm_fused_moe @@ -67,6 +73,37 @@ def _unpermute_activation_kernel( tl.store(out_ptrs, zeros, mask=col_mask) +def _trtllm_lora_unpermute_activation_warmup_inputs(vllm_config: Any) -> dict[str, Any]: + num_cols = vllm_config.model_config.hf_config.moe_intermediate_size + dtype = vllm_config.model_config.dtype + return dict( + act_permuted=TritonWarmupTensor(dtype, shape=(1, num_cols)), + idx_map=TritonWarmupTensor(torch.int64), + out=TritonWarmupTensor(dtype, shape=(1, num_cols)), + intermediate_size=num_cols, + ) + + +@triton_kernel_dispatcher_with_warmup( + kernel=_unpermute_activation_kernel, + warmup_inputs=_trtllm_lora_unpermute_activation_warmup_inputs, +) +def _TRTLLM_LORA_UNPERMUTE_ACTIVATION_KERNEL( + act_permuted: torch.Tensor, + idx_map: torch.Tensor, + out: torch.Tensor, + intermediate_size: int, +) -> DispatchSpec: + return (out.shape[0], triton.cdiv(intermediate_size, 1024)), dict( + act_ptr=act_permuted, + idx_ptr=idx_map, + num_cols=intermediate_size, + stride_ar=act_permuted.stride(0), + stride_or=out.stride(0), + BLOCK_I=1024, + ) + + @triton.jit def _finalize_lora_kernel( gemm2_ptr, # (num_permuted, K) base FC2 output, permuted, unweighted @@ -108,6 +145,52 @@ def _finalize_lora_kernel( ) +def _trtllm_lora_finalize_warmup_inputs(vllm_config: Any) -> dict[str, Any]: + hidden_size = vllm_config.model_config.hf_config.hidden_size + top_k = vllm_config.model_config.hf_config.num_experts_per_tok + dtype = vllm_config.model_config.dtype + return dict( + gemm2_permuted=TritonWarmupTensor(dtype, shape=(1, hidden_size)), + expert_weights=TritonWarmupTensor(torch.float32, shape=(1, top_k)), + idx_map=TritonWarmupTensor(torch.int64, shape=(top_k,)), + w2_delta=TritonWarmupTensor(dtype, shape=(1, top_k, hidden_size)), + output=TritonWarmupTensor(dtype, shape=(1, hidden_size)), + top_k=top_k, + scale=1.0, + ) + + +@triton_kernel_dispatcher_with_warmup( + kernel=_finalize_lora_kernel, + warmup_inputs=_trtllm_lora_finalize_warmup_inputs, +) +def _TRTLLM_LORA_FINALIZE_KERNEL( + gemm2_permuted: torch.Tensor, + expert_weights: torch.Tensor, + idx_map: torch.Tensor, + w2_delta: torch.Tensor, + output: torch.Tensor, + *, + top_k: int, + scale: float, +) -> DispatchSpec: + hidden_size = gemm2_permuted.shape[1] + return (output.shape[0], triton.cdiv(hidden_size, 512)), dict( + gemm2_ptr=gemm2_permuted, + weight_ptr=expert_weights, + idx_ptr=idx_map, + delta_ptr=w2_delta, + out_ptr=output, + K=hidden_size, + stride_g0=gemm2_permuted.stride(0), + stride_d0=w2_delta.stride(0), + stride_d1=w2_delta.stride(1), + stride_o0=output.stride(0), + TOP_K=top_k, + BLOCK_K=512, + ) + + class _TrtLlmLoRAExpertsBase(LoRAExpertsMixin, mk.FusedMoEExpertsModular): """LoRA-aware trtllm MoE experts""" @@ -434,16 +517,11 @@ def _unpermute_activation( dtype=act_permuted.dtype, device=act_permuted.device, ) - BLOCK_I = 1024 - grid = (num_rows, triton.cdiv(intermediate_size, BLOCK_I)) - _unpermute_activation_kernel[grid]( + _TRTLLM_LORA_UNPERMUTE_ACTIVATION_KERNEL( act_permuted, idx_map, out, intermediate_size, - act_permuted.stride(0), - out.stride(0), - BLOCK_I=BLOCK_I, ) return out @@ -464,23 +542,14 @@ def _finalize_with_w2_lora( (``expert_weights`` in expanded order, ``idx_map < 0`` dropped), scale by ``scale``, and add the already-weighted ``w2_delta`` reduced over top_k. """ - K = gemm2_permuted.size(1) - BLOCK_K = 512 - grid = (num_tokens, triton.cdiv(K, BLOCK_K)) - _finalize_lora_kernel[grid]( + _TRTLLM_LORA_FINALIZE_KERNEL( gemm2_permuted, - expert_weights.reshape(-1), + expert_weights, idx_map, w2_delta, output, - K, - gemm2_permuted.stride(0), - w2_delta.stride(0), - w2_delta.stride(1), - output.stride(0), - scale, - TOP_K=top_k, - BLOCK_K=BLOCK_K, + top_k=top_k, + scale=scale, ) diff --git a/vllm/model_executor/layers/fused_moe/fused_moe.py b/vllm/model_executor/layers/fused_moe/fused_moe.py index dc619b2b12d2..c9b88b450227 100644 --- a/vllm/model_executor/layers/fused_moe/fused_moe.py +++ b/vllm/model_executor/layers/fused_moe/fused_moe.py @@ -31,6 +31,15 @@ resolve_moe_use_td, warn_if_moe_use_td_ineffective, ) +from vllm.model_executor.warmup.jit_warmup import ( + WarmupChoices, + WarmupIntRange, +) +from vllm.model_executor.warmup.jit_warmup_triton_helper import ( + DispatchSpec, + TritonWarmupTensor, + triton_kernel_dispatcher_with_warmup, +) from vllm.platforms import current_platform from vllm.triton_utils import tl, triton from vllm.triton_utils.allocation import set_triton_allocator @@ -295,6 +304,8 @@ def fused_moe_kernel_gptq_awq( tl.store(c_ptrs, accumulator, mask=c_mask) +# NOTE(zyongye): we can remove all the wna16 kernel +# once we drop off sm75 support @triton.jit def fused_moe_kernel( # Pointers to matrices @@ -610,8 +621,187 @@ def fused_moe_kernel( tl.store(c_ptrs, accumulator, mask=c_mask) -# NOTE(zyongye): we can remove all the wna16 kernel -# once we drop off sm75 support +def _fused_moe_triton_kernel_warmup_inputs(vllm_config: Any) -> dict[str, Any]: + hf = vllm_config.model_config.hf_config + hidden_size = hf.hidden_size + intermediate_size = hf.moe_intermediate_size + num_experts = hf.n_routed_experts + config_top_k = hf.num_experts_per_tok + batch_tokens: Any = WarmupIntRange( + 1, vllm_config.scheduler_config.max_num_batched_tokens + 1 + ) + scenario: Any = WarmupChoices(0, 1, 2, 3, 4, 5, 6, 7) + second_gemm = scenario % 2 == 1 + naive = scenario % 4 >= 2 + use_fp8 = scenario >= 4 + routed_multiplier = config_top_k if second_gemm else 1 + n = hidden_size if second_gemm else 2 * intermediate_size + k = intermediate_size if second_gemm else hidden_size + top_k = 1 if second_gemm else config_top_k + mul_weight = second_gemm + dtype = torch.float8_e4m3fn if use_fp8 else vllm_config.model_config.dtype + group = 128 if use_fp8 else 0 + config = _triton_moe_config( + num_experts=num_experts, + hidden_size=hidden_size, + intermediate_size=intermediate_size, + top_k=config_top_k, + config_dtype=_get_config_dtype_str( + use_fp8_w8a8=use_fp8, + use_int8_w8a16=False, + use_int4_w4a16=False, + dtype=dtype, + ), + num_tokens=batch_tokens, + group_n=group, + group_k=group, + ) + block_m = config["BLOCK_SIZE_M"] + em = _triton_moe_em(batch_tokens * routed_multiplier, top_k, block_m, naive) + valid = batch_tokens * routed_multiplier * top_k + scale_cols = triton.cdiv(k, group) if group > 0 else 1 + return dict( + A=TritonWarmupTensor(dtype, shape=(1, k)), + B=TritonWarmupTensor(dtype, shape=(1, n, k)), + C=TritonWarmupTensor(torch.bfloat16, shape=(1, top_k, n)), + B_bias=None, + A_scale=TritonWarmupTensor(torch.float32, shape=(1, scale_cols)) + if use_fp8 + else None, + B_scale=TritonWarmupTensor(torch.float32, shape=(1, n, scale_cols)) + if use_fp8 + else None, + topk_weights=TritonWarmupTensor(torch.float32) if mul_weight else None, + sorted_token_ids=None if naive else TritonWarmupTensor(torch.int32), + expert_ids=TritonWarmupTensor(torch.int32), + num_tokens_post_padded=TritonWarmupTensor(torch.int32), + N=n, + K=k, + EM=em, + num_valid_tokens=valid, + stride_am=k, + stride_ak=1, + stride_be=n * k, + stride_bk=1, + stride_bn=k, + stride_cm=n, + stride_cn=1, + stride_asm=scale_cols if group else 0, + stride_ask=1 if group else 0, + stride_bse=n * scale_cols if group else 0, + stride_bsk=1 if group else 0, + stride_bsn=scale_cols if group else 0, + stride_bbe=0, + stride_bbn=0, + group_n=group, + group_k=group, + dtype=dtype, + A_ROWS=valid, + naive_block_assignment=naive, + BLOCK_SIZE_M=block_m, + BLOCK_SIZE_N=config["BLOCK_SIZE_N"], + BLOCK_SIZE_K=config["BLOCK_SIZE_K"], + GROUP_SIZE_M=config["GROUP_SIZE_M"], + SPLIT_K=config["SPLIT_K"], + MUL_ROUTED_WEIGHT=mul_weight, + top_k=top_k, + compute_type=_triton_moe_compute_type(dtype), + use_fp8_w8a8=use_fp8, + use_int8_w8a8=False, + use_int8_w8a16=False, + per_channel_quant=False, + HAS_BIAS=False, + SWAP_AB=use_fp8 and enable_swap_ab(block_m, config["BLOCK_SIZE_N"]), + USE_TD=resolve_moe_use_td() and not use_fp8 and k % config["BLOCK_SIZE_K"] == 0, + num_warps=config["num_warps"], + num_stages=config["num_stages"], + waves_per_eu=config.get("waves_per_eu"), + ) + + +@triton_kernel_dispatcher_with_warmup( + kernel=fused_moe_kernel, + warmup_inputs=_fused_moe_triton_kernel_warmup_inputs, +) +def _FUSED_MOE_TRITON_KERNEL( + A: torch.Tensor, + B: torch.Tensor, + C: torch.Tensor, + B_bias: torch.Tensor | None, + A_scale: torch.Tensor | None, + B_scale: torch.Tensor | None, + topk_weights: torch.Tensor | None, + sorted_token_ids: torch.Tensor | None, + expert_ids: torch.Tensor, + num_tokens_post_padded: torch.Tensor, + N: int, + K: int, + EM: int, + num_valid_tokens: int, + stride_am: int, + stride_ak: int, + stride_be: int, + stride_bk: int, + stride_bn: int, + stride_cm: int, + stride_cn: int, + stride_asm: int, + stride_ask: int, + stride_bse: int, + stride_bsk: int, + stride_bsn: int, + stride_bbe: int, + stride_bbn: int, + group_n: int, + group_k: int, + *, + dtype: torch.dtype, + A_ROWS: int, + naive_block_assignment: bool, + BLOCK_SIZE_M: int, + BLOCK_SIZE_N: int, + BLOCK_SIZE_K: int, + GROUP_SIZE_M: int, + SPLIT_K: int, + MUL_ROUTED_WEIGHT: bool, + top_k: int, + compute_type: tl.dtype, + use_fp8_w8a8: bool, + use_int8_w8a8: bool, + use_int8_w8a16: bool, + per_channel_quant: bool, + HAS_BIAS: bool, + SWAP_AB: bool, + USE_TD: bool = False, + num_warps: int = 4, + num_stages: int = 3, + waves_per_eu: int | None = None, + matrix_instr_nonkdim: int | None = None, + kpack: int | None = None, +) -> DispatchSpec: + grid: Any = lambda META: ( + triton.cdiv(EM, META["BLOCK_SIZE_M"]) + * triton.cdiv(B.size(1), META["BLOCK_SIZE_N"]), + ) + launch_kwargs = dict( + a_ptr=A, + b_ptr=B, + c_ptr=C, + b_bias_ptr=B_bias, + a_scale_ptr=A_scale, + b_scale_ptr=B_scale, + num_warps=num_warps, + num_stages=num_stages, + ) + if waves_per_eu is not None: + launch_kwargs["waves_per_eu"] = waves_per_eu + if matrix_instr_nonkdim is not None: + launch_kwargs["matrix_instr_nonkdim"] = matrix_instr_nonkdim + if kpack is not None: + launch_kwargs["kpack"] = kpack + return grid, launch_kwargs + + def invoke_fused_moe_wna16_cuda_kernel( A: torch.Tensor, B: torch.Tensor, @@ -833,10 +1023,6 @@ def invoke_fused_moe_triton_kernel( ) else: EM = num_tokens * config["BLOCK_SIZE_M"] - grid = lambda META: ( - triton.cdiv(EM, META["BLOCK_SIZE_M"]) - * triton.cdiv(B.size(1), META["BLOCK_SIZE_N"]), - ) HAS_BIAS = B_bias is not None config = config.copy() @@ -857,13 +1043,12 @@ def invoke_fused_moe_triton_kernel( BLOCK_SIZE_K, ) use_td = False - # Triton treats 0-D tensor arguments as scalar values, but the kernel # loads tensor-wise activation scales through a pointer. if A_scale is not None and A_scale.ndim == 0: A_scale = A_scale.reshape(1) - fused_moe_kernel[grid]( + _FUSED_MOE_TRITON_KERNEL( A, B, C, @@ -894,6 +1079,8 @@ def invoke_fused_moe_triton_kernel( B_bias.stride(1) if B_bias is not None else 0, 0 if block_shape is None else block_shape[0], 0 if block_shape is None else block_shape[1], + dtype=A.dtype, + A_ROWS=A.size(0), MUL_ROUTED_WEIGHT=mul_routed_weight, top_k=top_k, compute_type=compute_type, @@ -1052,6 +1239,41 @@ def compute_identity_kernel( ) +def _compute_identity_warmup_inputs(vllm_config: Any) -> dict[str, Any]: + hidden_dim = vllm_config.model_config.hf_config.hidden_size + top_k = vllm_config.model_config.hf_config.num_experts_per_tok + dtype = vllm_config.model_config.dtype + num_tokens: Any = WarmupIntRange( + 1, min(vllm_config.scheduler_config.max_num_batched_tokens, 16) + 1 + ) + return dict( + top_k=top_k, + hidden_states=TritonWarmupTensor(dtype, shape=(num_tokens, hidden_dim)), + expert_scales=TritonWarmupTensor(torch.float32, shape=(num_tokens, top_k)), + num_tokens=num_tokens, + output=TritonWarmupTensor(dtype, shape=(num_tokens, hidden_dim)), + hidden_dim=hidden_dim, + scales_stride=top_k, + ) + + +@triton_kernel_dispatcher_with_warmup( + kernel=compute_identity_kernel, warmup_inputs=_compute_identity_warmup_inputs +) +def _COMPUTE_IDENTITY_KERNEL( + *, + top_k: int, + hidden_states: torch.Tensor, + expert_scales: torch.Tensor, + num_tokens: int, + output: torch.Tensor, + hidden_dim: int, + scales_stride: int, +) -> DispatchSpec: + block_size = 256 + return (num_tokens * (hidden_dim // block_size),), dict(BLOCK_SIZE=block_size) + + def zero_experts_compute_triton( expert_indices: torch.Tensor, expert_scales: torch.Tensor, @@ -1059,9 +1281,7 @@ def zero_experts_compute_triton( zero_expert_type: str, hidden_states: torch.Tensor, ) -> torch.Tensor: - N = expert_indices.numel() top_k = expert_indices.size(-1) - grid = lambda meta: (triton.cdiv(N, meta["BLOCK_SIZE"]),) if zero_expert_type == "identity": zero_expert_mask = expert_indices < num_experts @@ -1076,16 +1296,14 @@ def zero_experts_compute_triton( hidden_dim = hidden_states.size(-1) num_tokens = hidden_states.size(0) - grid = lambda meta: (num_tokens * (hidden_dim // meta["BLOCK_SIZE"]),) - compute_identity_kernel[grid]( - top_k, - hidden_states, - zero_expert_scales, - num_tokens, - output, - hidden_dim, - zero_expert_scales.stride(0), - BLOCK_SIZE=256, + _COMPUTE_IDENTITY_KERNEL( + top_k=top_k, + hidden_states=hidden_states, + expert_scales=zero_expert_scales, + num_tokens=num_tokens, + output=output, + hidden_dim=hidden_dim, + scales_stride=zero_expert_scales.stride(0), ) return output @@ -1451,6 +1669,59 @@ def try_get_optimal_moe_config( return config +def _triton_moe_config( + *, + num_experts: int, + hidden_size: int, + intermediate_size: int, + top_k: int, + config_dtype: str | None, + num_tokens: int, + group_n: int, + group_k: int, +) -> dict[str, int]: + block_shape = [group_n, group_k] if group_n > 0 and group_k > 0 else None + config = try_get_optimal_moe_config( + (num_experts, 2 * intermediate_size, hidden_size), + (num_experts, hidden_size, intermediate_size), + top_k, + config_dtype, + num_tokens, + block_shape=block_shape, + ) + block_size_k = ( + min(config["BLOCK_SIZE_K"], min(group_n, group_k)) + if block_shape is not None + else config["BLOCK_SIZE_K"] + ) + return { + **config, + "BLOCK_SIZE_K": block_size_k, + "num_warps": config.get("num_warps", 4), + "num_stages": config.get("num_stages", 3), + } + + +def _triton_moe_compute_type(dtype: torch.dtype) -> tl.dtype: + if dtype == torch.float32: + return tl.float32 + if dtype == torch.float16: + return tl.float16 + return tl.bfloat16 + + +def _triton_moe_em( + num_tokens: int, + top_k: int, + block_size_m: int, + naive_block_assignment: bool, +) -> int: + routed_tokens = num_tokens * top_k + if naive_block_assignment: + return routed_tokens * block_size_m + return triton.cdiv(routed_tokens, block_size_m) * block_size_m + + def fused_experts_op( hidden_states: torch.Tensor, w1: torch.Tensor, diff --git a/vllm/model_executor/layers/fused_moe/moe_fused_mul_sum.py b/vllm/model_executor/layers/fused_moe/moe_fused_mul_sum.py index d8ad9b94557d..4616255a0a4c 100644 --- a/vllm/model_executor/layers/fused_moe/moe_fused_mul_sum.py +++ b/vllm/model_executor/layers/fused_moe/moe_fused_mul_sum.py @@ -1,8 +1,16 @@ # SPDX-License-Identifier: Apache-2.0 # SPDX-FileCopyrightText: Copyright contributors to the vLLM project +from typing import Any + import torch from torch._subclasses.fake_tensor import FakeTensor +from vllm.model_executor.warmup.jit_warmup import WarmupChoices, _when +from vllm.model_executor.warmup.jit_warmup_triton_helper import ( + DispatchSpec, + TritonWarmupTensor, + triton_kernel_dispatcher_with_warmup, +) from vllm.triton_utils import tl, triton @@ -72,6 +80,72 @@ def moe_fused_mul_sum_kernel( tl.store(out_row + offs_k, acc.to(outputs_ptr.dtype.element_ty), mask=kmask) +def _moe_fused_mul_sum_warmup_inputs(vllm_config: Any) -> dict[str, Any]: + hf_config = vllm_config.model_config.hf_config + hidden_size = hf_config.hidden_size + top_k = hf_config.num_experts_per_tok + input_dtype: Any = WarmupChoices(vllm_config.model_config.dtype, torch.float32) + output_dtype: Any = WarmupChoices(vllm_config.model_config.dtype, torch.float32) + has_topk_ids: Any = WarmupChoices(False, True) + has_expert_map: Any = WarmupChoices(False, True) + has_num_valid: Any = WarmupChoices(False, True) + _when( + input_dtype == output_dtype + and (has_topk_ids or not has_expert_map) + and (has_topk_ids or not has_num_valid) + ) + return dict( + inputs=TritonWarmupTensor( + input_dtype, + shape=(1, top_k, hidden_size), + ), + topk_weights=TritonWarmupTensor( + torch.float32, + shape=(1, top_k), + ), + outputs=TritonWarmupTensor( + output_dtype, + shape=(1, hidden_size), + ), + topk_ids=( + TritonWarmupTensor(torch.int32, shape=(1, top_k)) if has_topk_ids else None + ), + expert_map=TritonWarmupTensor(torch.int32) if has_expert_map else None, + num_valid_tokens=(TritonWarmupTensor(torch.int32) if has_num_valid else None), + ) + + +@triton_kernel_dispatcher_with_warmup( + kernel=moe_fused_mul_sum_kernel, + warmup_inputs=_moe_fused_mul_sum_warmup_inputs, +) +def _MOE_FUSED_MUL_SUM_KERNEL( + inputs: torch.Tensor, + topk_weights: torch.Tensor, + outputs: torch.Tensor, + topk_ids: torch.Tensor | None, + expert_map: torch.Tensor | None, + num_valid_tokens: torch.Tensor | None, +) -> DispatchSpec: + num_tokens, top_k, hidden_size = inputs.shape + block_k, num_warps, num_stages = _heuristic_config( + hidden_size, + inputs.dtype.itemsize, + ) + return (num_tokens,), dict( + top_ids_ptr=topk_ids, + stride_m=top_k * hidden_size, + has_topk_ids=topk_ids is not None, + has_expert_map=expert_map is not None, + has_num_valid=num_valid_tokens is not None, + top_k=top_k, + hidden_size=hidden_size, + BLOCK_K=block_k, + num_warps=num_warps, + num_stages=num_stages, + ) + + def _heuristic_config( hidden_size: int, element_size: int, @@ -146,27 +220,13 @@ def moe_fused_mul_sum( assert topk_ids.dtype in (torch.int32, torch.int64) if not isinstance(inputs, FakeTensor): - BLOCK_K, num_warps, num_stages = _heuristic_config( - hidden_size, - inputs.element_size(), - ) - grid = (num_tokens,) - moe_fused_mul_sum_kernel[grid]( + _MOE_FUSED_MUL_SUM_KERNEL( inputs, topk_weights, outputs, topk_ids, expert_map, num_valid_tokens, - top_k * hidden_size, - topk_ids is not None, - expert_map is not None, - num_valid_tokens is not None, - top_k, - hidden_size, - BLOCK_K, - num_warps=num_warps, - num_stages=num_stages, ) return outputs diff --git a/vllm/model_executor/layers/fused_moe/prepare_finalize/deepep_v2.py b/vllm/model_executor/layers/fused_moe/prepare_finalize/deepep_v2.py index d834c6eff73c..64a3b9e23c02 100644 --- a/vllm/model_executor/layers/fused_moe/prepare_finalize/deepep_v2.py +++ b/vllm/model_executor/layers/fused_moe/prepare_finalize/deepep_v2.py @@ -1,6 +1,7 @@ # SPDX-License-Identifier: Apache-2.0 # SPDX-FileCopyrightText: Copyright contributors to the vLLM project from collections.abc import Callable +from typing import Any import deep_ep import torch @@ -17,6 +18,15 @@ MXFP8_BLOCK_SIZE, swizzle_mxfp8_scale, ) +from vllm.model_executor.warmup.jit_warmup import ( + WarmupChoices, + WarmupIntRange, +) +from vllm.model_executor.warmup.jit_warmup_triton_helper import ( + DispatchSpec, + TritonWarmupTensor, + triton_kernel_dispatcher_with_warmup, +) from vllm.triton_utils import tl, triton from vllm.utils.math_utils import round_up from vllm.v1.worker.ubatching import ( @@ -540,24 +550,59 @@ def _globalize_recv_topk_idx_kernel( tl.store(topk_idx_ptr + offs, tl.where(valid, g, -1), mask=mask) +def _globalize_recv_topk_idx_warmup_inputs(vllm_config: Any) -> dict[str, Any]: + topk = vllm_config.model_config.hf_config.num_experts_per_tok + num_experts = vllm_config.model_config.hf_config.n_routed_experts + num_tokens: Any = WarmupIntRange( + 1, vllm_config.scheduler_config.max_num_batched_tokens + 1 + ) + rank_expert_offset: Any = WarmupChoices(0, 1, 2) + return dict( + recv_topk_idx=TritonWarmupTensor(torch.int64, shape=(num_tokens, topk)), + psum_recv_per_rank=TritonWarmupTensor( + torch.int32, shape=(vllm_config.parallel_config.data_parallel_size,) + ), + rank_expert_offset=rank_expert_offset, + num_experts=num_experts, + ) + + +@triton_kernel_dispatcher_with_warmup( + kernel=_globalize_recv_topk_idx_kernel, + warmup_inputs=_globalize_recv_topk_idx_warmup_inputs, +) +def _GLOBALIZE_RECV_TOPK_IDX_KERNEL( + recv_topk_idx: torch.Tensor, + psum_recv_per_rank: torch.Tensor, + rank_expert_offset: int, + num_experts: int, + *, + n_elements: int | None = None, + block: int = 1024, +) -> DispatchSpec: + topk = recv_topk_idx.shape[1] + if n_elements is None: + n_elements = recv_topk_idx.shape[0] * topk + return (triton.cdiv(n_elements, block),), dict( + topk_idx_ptr=recv_topk_idx, + psum_ptr=psum_recv_per_rank, + P=psum_recv_per_rank.shape[0], + n_elements=n_elements, + topk=topk, + BLOCK=block, + ) + + def _globalize_recv_topk_idx( recv_topk_idx: torch.Tensor, # [N, topk] local expert IDs, -1 = non-local psum_recv_per_rank: torch.Tensor, rank_expert_offset: int, num_experts: int, ) -> torch.Tensor: - N, topk = recv_topk_idx.shape - n = N * topk - BLOCK = 1024 - grid = (triton.cdiv(n, BLOCK),) - _globalize_recv_topk_idx_kernel[grid]( + _GLOBALIZE_RECV_TOPK_IDX_KERNEL( recv_topk_idx, psum_recv_per_rank, - psum_recv_per_rank.shape[0], rank_expert_offset, num_experts, - n, - topk=topk, - BLOCK=BLOCK, ) return recv_topk_idx diff --git a/vllm/model_executor/layers/fused_moe/router/base_router.py b/vllm/model_executor/layers/fused_moe/router/base_router.py index 01e674b2b131..29134b7ed651 100644 --- a/vllm/model_executor/layers/fused_moe/router/base_router.py +++ b/vllm/model_executor/layers/fused_moe/router/base_router.py @@ -2,6 +2,7 @@ # SPDX-FileCopyrightText: Copyright contributors to the vLLM project from abc import abstractmethod from collections.abc import Callable +from typing import Any import torch @@ -9,11 +10,18 @@ from vllm.model_executor.layers.fused_moe.router.fused_moe_router import ( FusedMoERouter, ) +from vllm.model_executor.warmup.jit_warmup import WarmupChoices +from vllm.model_executor.warmup.jit_warmup_triton_helper import ( + DispatchSpec, + TritonWarmupTensor, + triton_kernel_dispatcher_with_warmup, +) from vllm.platforms import current_platform from vllm.triton_utils import tl, triton from vllm.v1.worker.ubatching import dbo_current_ubatch_id if current_platform.is_cuda_alike(): + _EPLB_MAP_BLOCK_SIZE = 256 @triton.jit def _eplb_map_and_record_i32_kernel( @@ -92,6 +100,58 @@ def _eplb_map_and_record_i32_kernel( safe_physical_id = tl.where(physical_id >= 0, physical_id, 0) tl.atomic_add(out_ptr + safe_physical_id, 1, mask=valid) + def _eplb_map_and_record_warmup_inputs(vllm_config: Any) -> dict[str, Any]: + top_k = vllm_config.model_config.hf_config.num_experts_per_tok + num_logical_experts = vllm_config.model_config.hf_config.n_routed_experts + num_redundant_experts = ( + vllm_config.parallel_config.eplb_config.num_redundant_experts + ) + has_num_unpadded: Any = WarmupChoices(False, True) + output_dtype: Any = WarmupChoices(torch.int32, torch.int64) + return dict( + topk_ids=TritonWarmupTensor(torch.int32), + logical_replica_count=TritonWarmupTensor(torch.int64), + logical_to_physical_map=TritonWarmupTensor(torch.int64), + out=TritonWarmupTensor(output_dtype), + expert_load_view=TritonWarmupTensor(torch.int32), + record_enabled=TritonWarmupTensor(torch.bool), + num_unpadded_tokens=TritonWarmupTensor(torch.int32) + if has_num_unpadded + else None, + num_logical_experts=num_logical_experts, + map_slots=1024, + out_size=num_logical_experts + num_redundant_experts, + numel=1, + num_active_experts=top_k, + ) + + @triton_kernel_dispatcher_with_warmup( + kernel=_eplb_map_and_record_i32_kernel, + warmup_inputs=_eplb_map_and_record_warmup_inputs, + ) + def _EPLB_MAP_AND_RECORD_KERNEL( + topk_ids: torch.Tensor, + logical_replica_count: torch.Tensor, + logical_to_physical_map: torch.Tensor, + out: torch.Tensor, + expert_load_view: torch.Tensor, + record_enabled: torch.Tensor, + num_unpadded_tokens: torch.Tensor | None, + num_logical_experts: int, + map_slots: int, + out_size: int, + numel: int, + num_active_experts: int, + ) -> DispatchSpec: + grid = (triton.cdiv(numel, _EPLB_MAP_BLOCK_SIZE),) + return grid, dict( + logical_to_physical_ptr=logical_to_physical_map, + out_ids_ptr=out, + out_ptr=expert_load_view, + HAS_NUM_UNPADDED=num_unpadded_tokens is not None, + BLOCK_SIZE=_EPLB_MAP_BLOCK_SIZE, + ) + def _eplb_map_and_record_triton( topk_ids: torch.Tensor, logical_to_physical_map: torch.Tensor, @@ -106,9 +166,8 @@ def _eplb_map_and_record_triton( return topk_ids num_active_experts = topk_ids_in.shape[-1] out_flat = torch.empty((numel,), device=topk_ids.device, dtype=topk_ids.dtype) - grid = lambda meta: (triton.cdiv(numel, meta["BLOCK_SIZE"]),) assert expert_load_view.is_contiguous() - _eplb_map_and_record_i32_kernel[grid]( + _EPLB_MAP_AND_RECORD_KERNEL( topk_ids_in, logical_replica_count.contiguous(), logical_to_physical_map.contiguous(), @@ -121,8 +180,6 @@ def _eplb_map_and_record_triton( expert_load_view.shape[0], numel, num_active_experts, - HAS_NUM_UNPADDED=num_unpadded_tokens is not None, - BLOCK_SIZE=256, ) return out_flat.reshape(topk_ids.shape) diff --git a/vllm/model_executor/layers/fused_moe/router/dsv4_topk.py b/vllm/model_executor/layers/fused_moe/router/dsv4_topk.py index dbd29f3d08ab..a57db80a6478 100644 --- a/vllm/model_executor/layers/fused_moe/router/dsv4_topk.py +++ b/vllm/model_executor/layers/fused_moe/router/dsv4_topk.py @@ -1,10 +1,19 @@ # SPDX-License-Identifier: Apache-2.0 # SPDX-FileCopyrightText: Copyright contributors to the vLLM project +from typing import Any + import torch +from vllm.model_executor.warmup.jit_warmup import WarmupChoices +from vllm.model_executor.warmup.jit_warmup_triton_helper import ( + DispatchSpec, + TritonWarmupTensor, + triton_kernel_dispatcher_with_warmup, +) from vllm.platforms import current_platform from vllm.triton_utils import tl, triton +from vllm.utils.math_utils import next_power_of_2 _TOPK = 6 @@ -35,6 +44,14 @@ def can_use_dsv4_topk( ) +def _image_sentinel_base_id() -> int: + from vllm.models.deepseek_v4.common.mm_preprocess import ( + IMAGE_SENTINEL_BASE_ID, + ) + + return IMAGE_SENTINEL_BASE_ID + + if current_platform.is_cuda(): @triton.jit @@ -112,6 +129,63 @@ def _dsv4_topk_kernel( tl.store(topk_weights_ptr + output_offsets, selected_weights, mask=output_mask) tl.store(topk_ids_ptr + output_offsets, selected_ids, mask=output_mask) + def _dsv4_topk_warmup_inputs(vllm_config: Any) -> dict[str, Any]: + hf_config = vllm_config.model_config.hf_config + num_experts = hf_config.n_routed_experts + has_vl: Any = WarmupChoices( + False, + getattr(vllm_config.model_config.hf_config, "vision_n_layers", 0) > 0, + ) + launch_pdl: Any = WarmupChoices(False, True) + return dict( + gating_output=TritonWarmupTensor(torch.float32, shape=(1, num_experts)), + correction_bias=TritonWarmupTensor(torch.float32, shape=(num_experts,)), + topk_weights=TritonWarmupTensor(torch.float32, shape=(1, _TOPK)), + topk_ids=TritonWarmupTensor( + torch.int64 + if vllm_config.kernel_config.moe_backend == "deep_gemm_mega_moe" + else torch.int32, + shape=(1, _TOPK), + ), + routed_scaling_factor=float( + getattr(hf_config, "routed_scaling_factor", 1.0) + ), + input_ids=TritonWarmupTensor(torch.int64) if has_vl else None, + bias_vl=TritonWarmupTensor(torch.float32, shape=(num_experts,)) + if has_vl + else None, + image_sentinel_lo=_image_sentinel_base_id() if has_vl else 0, + launch_pdl=launch_pdl, + ) + + @triton_kernel_dispatcher_with_warmup( + kernel=_dsv4_topk_kernel, + warmup_inputs=_dsv4_topk_warmup_inputs, + ) + def _DSV4_TOPK_KERNEL( + gating_output: torch.Tensor, + correction_bias: torch.Tensor, + topk_weights: torch.Tensor, + topk_ids: torch.Tensor, + routed_scaling_factor: float, + input_ids: torch.Tensor | None = None, + bias_vl: torch.Tensor | None = None, + image_sentinel_lo: int = 0, + launch_pdl: bool | None = None, + ) -> DispatchSpec: + num_tokens, num_experts = gating_output.shape + return (num_tokens,), dict( + NUM_EXPERTS=num_experts, + BLOCK_N=next_power_of_2(num_experts), + HAS_VL=bias_vl is not None and image_sentinel_lo > 0, + num_warps=1, + launch_pdl=( + current_platform.is_arch_support_pdl() + if launch_pdl is None + else launch_pdl + ), + ) + def dsv4_topk( gating_output: torch.Tensor, @@ -123,7 +197,6 @@ def dsv4_topk( image_sentinel_lo: int = 0, ) -> tuple[torch.Tensor, torch.Tensor]: num_tokens, num_experts = gating_output.shape - has_vl = bias_vl is not None and image_sentinel_lo > 0 if bias_vl is not None: assert input_ids is not None, "bias_vl routing requires input_ids" assert bias_vl.dtype == torch.float32 and bias_vl.is_contiguous() @@ -133,7 +206,7 @@ def dsv4_topk( topk_weights = gating_output.new_empty(shape, dtype=torch.float32) topk_ids = gating_output.new_empty(shape, dtype=indices_dtype) if num_tokens > 0: - _dsv4_topk_kernel[(num_tokens,)]( + _DSV4_TOPK_KERNEL( gating_output, correction_bias, topk_weights, @@ -142,10 +215,5 @@ def dsv4_topk( input_ids, bias_vl, image_sentinel_lo, - NUM_EXPERTS=num_experts, - BLOCK_N=triton.next_power_of_2(num_experts), - HAS_VL=has_vl, - num_warps=1, - launch_pdl=current_platform.is_arch_support_pdl(), ) return topk_weights, topk_ids diff --git a/vllm/model_executor/layers/fused_moe/utils.py b/vllm/model_executor/layers/fused_moe/utils.py index 48e217b1a91a..a92d4be5a043 100644 --- a/vllm/model_executor/layers/fused_moe/utils.py +++ b/vllm/model_executor/layers/fused_moe/utils.py @@ -3,7 +3,7 @@ import functools from collections.abc import Iterable from math import prod -from typing import TYPE_CHECKING +from typing import TYPE_CHECKING, Any import torch import torch.nn.functional as F @@ -41,6 +41,15 @@ per_tensor_dequantize, ) from vllm.model_executor.models.utils import PPMissingLayer +from vllm.model_executor.warmup.jit_warmup import ( + WarmupChoices, + WarmupIntRange, +) +from vllm.model_executor.warmup.jit_warmup_triton_helper import ( + DispatchSpec, + TritonWarmupTensor, + triton_kernel_dispatcher_with_warmup, +) from vllm.platforms import current_platform from vllm.triton_utils import tl, triton from vllm.utils.math_utils import cdiv @@ -161,6 +170,49 @@ def _count_expert_num_tokens( tl.store(expert_num_tokens_ptr + curr_expert, tl.sum(acc)) +def _count_expert_num_tokens_warmup_inputs(vllm_config: Any) -> dict[str, Any]: + top_k = vllm_config.model_config.hf_config.num_experts_per_tok + num_tokens: Any = WarmupIntRange( + 1, + min( + vllm_config.scheduler_config.max_num_batched_tokens, + 4096 // vllm_config.model_config.hf_config.num_experts_per_tok, + ) + + 1, + ) + num_experts: Any = WarmupChoices(1, 2, 16) + has_expert_map: Any = WarmupChoices(False, True) + return dict( + topk_ids=TritonWarmupTensor(torch.int32, shape=(num_tokens, top_k)), + expert_num_tokens=TritonWarmupTensor(torch.int32, shape=(num_experts,)), + num_local_experts=num_experts, + expert_map=TritonWarmupTensor(torch.int32) if has_expert_map else None, + ) + + +@triton_kernel_dispatcher_with_warmup( + kernel=_count_expert_num_tokens, + warmup_inputs=_count_expert_num_tokens_warmup_inputs, +) +def _COUNT_EXPERT_NUM_TOKENS_KERNEL( + topk_ids: torch.Tensor, + expert_num_tokens: torch.Tensor, + num_local_experts: int, + expert_map: torch.Tensor | None, + *, + block_size: int | None = None, +) -> DispatchSpec: + topk_numel = prod(topk_ids.shape) + if block_size is None: + block_size = triton.next_power_of_2(min(topk_numel, 1024)) + return (num_local_experts,), dict( + num_experts=num_local_experts, + topk_numel=topk_numel, + HAS_EXPERT_MAP=expert_map is not None, + BLOCK_SIZE=block_size, + ) + + def count_expert_num_tokens( topk_ids: torch.Tensor, num_local_experts: int, expert_map: torch.Tensor | None ) -> torch.Tensor: @@ -184,18 +236,11 @@ def count_expert_num_tokens( (num_local_experts), device=topk_ids.device, dtype=torch.int32 ) - grid = num_local_experts - BLOCK_SIZE = min(topk_ids.numel(), 1024) - BLOCK_SIZE = triton.next_power_of_2(BLOCK_SIZE) - - _count_expert_num_tokens[(grid,)]( + _COUNT_EXPERT_NUM_TOKENS_KERNEL( topk_ids, expert_num_tokens, num_local_experts, - topk_ids.numel(), expert_map, - HAS_EXPERT_MAP=expert_map is not None, - BLOCK_SIZE=BLOCK_SIZE, ) return expert_num_tokens @@ -520,7 +565,13 @@ def _swiglu_limit_torch( output.copy_(F.silu(gate) * up) -@triton.jit +@triton.jit( + do_not_specialize=[ + "hidden_size", + "input_row_stride", + "swiglu_limit", + ] +) def _swiglu_limit_pad_aware_kernel( input_ptr, # [num_tokens, 2 * hidden_size] output_ptr, # [num_tokens, hidden_size] @@ -576,6 +627,48 @@ def _swiglu_limit_pad_aware_kernel( ) +def _swiglu_limit_pad_aware_warmup_inputs( + vllm_config: Any, +) -> dict[str, Any]: + num_tokens: Any = WarmupIntRange( + 1, vllm_config.scheduler_config.max_num_batched_tokens + 1 + ) + has_limit: Any = WarmupChoices(False, True) + has_expert_map: Any = WarmupChoices(False, True) + return dict( + output=TritonWarmupTensor(torch.bfloat16, shape=(num_tokens, 1)), + input=TritonWarmupTensor(torch.bfloat16, shape=(num_tokens, 2), strides=(2, 1)), + topk_ids=TritonWarmupTensor(torch.int32), + swiglu_limit=1.0 if has_limit else 0.0, + expert_map=TritonWarmupTensor(torch.int32) if has_expert_map else None, + ) + + +@triton_kernel_dispatcher_with_warmup( + kernel=_swiglu_limit_pad_aware_kernel, + warmup_inputs=_swiglu_limit_pad_aware_warmup_inputs, +) +def _SWIGLU_LIMIT_PAD_AWARE_KERNEL( + output: torch.Tensor, + input: torch.Tensor, + topk_ids: torch.Tensor, + swiglu_limit: float, + expert_map: torch.Tensor | None, +) -> DispatchSpec: + num_tokens, gate_up_size = input.shape + hidden_size = gate_up_size // 2 + block_size = 1024 + return (min(num_tokens, 256), triton.cdiv(hidden_size, block_size)), dict( + hidden_size=hidden_size, + input_row_stride=gate_up_size, + num_tokens=num_tokens, + HAS_LIMIT=swiglu_limit > 0, + HAS_EXPERT_MAP=expert_map is not None, + BLOCK_SIZE=block_size, + num_warps=4, + ) + + def _swiglu_limit_pad_aware( output: torch.Tensor, input: torch.Tensor, @@ -583,26 +676,16 @@ def _swiglu_limit_pad_aware( swiglu_limit: float, expert_map: torch.Tensor | None = None, ) -> None: - num_tokens, gate_up_size = input.shape - hidden_size = gate_up_size // 2 + num_tokens = input.shape[0] if num_tokens == 0: return - BLOCK_SIZE = 1024 - grid = (min(num_tokens, 256), triton.cdiv(hidden_size, BLOCK_SIZE)) - _swiglu_limit_pad_aware_kernel[grid]( - input, + _SWIGLU_LIMIT_PAD_AWARE_KERNEL( output, + input, topk_ids, - expert_map, - hidden_size, - gate_up_size, - num_tokens, swiglu_limit, - HAS_LIMIT=swiglu_limit > 0, - HAS_EXPERT_MAP=expert_map is not None, - BLOCK_SIZE=BLOCK_SIZE, - num_warps=4, + expert_map, ) diff --git a/vllm/models/deepseek_v4/nvidia/model.py b/vllm/models/deepseek_v4/nvidia/model.py index 021001b9ce01..b9dc8946936a 100644 --- a/vllm/models/deepseek_v4/nvidia/model.py +++ b/vllm/models/deepseek_v4/nvidia/model.py @@ -82,7 +82,9 @@ DeepseekV4FlashInferSM120Attention, ) from vllm.models.deepseek_v4.nvidia.flashmla import DeepseekV4FlashMLAAttention -from vllm.models.deepseek_v4.nvidia.ops.prepare_megamoe import prepare_megamoe_inputs +from vllm.models.deepseek_v4.nvidia.ops.prepare_megamoe import ( + _PREPARE_MEGAMOE_INPUTS_KERNEL, +) from vllm.platforms import current_platform from vllm.sequence import IntermediateTensors from vllm.utils.flashinfer_moe_ep import ( @@ -729,19 +731,19 @@ def forward( self.top_k, "fp8xfp4", ) - - prepare_megamoe_inputs( - hidden_states, - topk_weights, - topk_ids, - symm_buffer.x[:num_tokens], - symm_buffer.x_sf[:num_tokens], - symm_buffer.topk_idx[:num_tokens], - symm_buffer.topk_weights[:num_tokens], - is_padding=is_padding, - shared_x_sf=shared_x_sf, - shared_block_m=shared_block_m, - ) + if num_tokens > 0: + _PREPARE_MEGAMOE_INPUTS_KERNEL( + hidden_states, + topk_weights, + topk_ids, + symm_buffer.x[:num_tokens], + symm_buffer.x_sf[:num_tokens], + symm_buffer.topk_idx[:num_tokens], + symm_buffer.topk_weights[:num_tokens], + is_padding=is_padding, + shared_x_sf=shared_x_sf, + shared_block_m=shared_block_m, + ) assert self._transformed_l1_weights is not None assert self._transformed_l2_weights is not None @@ -898,6 +900,115 @@ def __init__( else: self._init_fused_moe_experts(vllm_config, config, quant_config, prefix) + if vllm_config.kernel_config.enable_jit_warmup: + from vllm.model_executor.layers.fused_moe.router.dsv4_topk import ( + _DSV4_TOPK_KERNEL, + ) + + if vllm_config.parallel_config.enable_eplb: + from vllm.model_executor.layers.fused_moe.router.base_router import ( + _EPLB_MAP_AND_RECORD_KERNEL, + ) + + _EPLB_MAP_AND_RECORD_KERNEL.register_warmup() + if ( + config.n_routed_experts in (256, 384) + and config.num_experts_per_tok == 6 + and config.norm_topk_prob + and config.scoring_func == "sqrtsoftplus" + ): + _DSV4_TOPK_KERNEL.register_warmup() + if self.use_mega_moe: + from vllm.model_executor.layers.fused_moe.deep_gemm_utils import ( + _DEEPGEMM_EP_GATHER_KERNEL, + _DEEPGEMM_EP_SCATTER_COPY_KERNEL, + _DEEPGEMM_EP_SCATTER_START_KERNEL, + ) + + _PREPARE_MEGAMOE_INPUTS_KERNEL.register_warmup() + _DEEPGEMM_EP_SCATTER_START_KERNEL.register_warmup() + _DEEPGEMM_EP_SCATTER_COPY_KERNEL.register_warmup() + _DEEPGEMM_EP_GATHER_KERNEL.register_warmup() + else: + from vllm.model_executor.layers.fused_moe.experts.deep_gemm_moe import ( # noqa: E501 + DeepGemmExperts, + DeepGemmFP4Experts, + ) + from vllm.model_executor.layers.fused_moe.experts.fused_batched_moe import ( # noqa: E501 + _BATCHED_TRITON_KERNEL, + BatchedTritonExperts, + ) + from vllm.model_executor.layers.fused_moe.experts.triton_moe import ( + TritonExperts, + ) + from vllm.model_executor.layers.fused_moe.fused_moe import ( + _COMPUTE_IDENTITY_KERNEL, + _FUSED_MOE_TRITON_KERNEL, + ) + from vllm.model_executor.layers.fused_moe.moe_fused_mul_sum import ( + _MOE_FUSED_MUL_SUM_KERNEL, + ) + from vllm.model_executor.layers.fused_moe.utils import ( + _COUNT_EXPERT_NUM_TOKENS_KERNEL, + _SWIGLU_LIMIT_PAD_AWARE_KERNEL, + ) + + experts_cls = getattr( + self.experts.routed_experts.quant_method, + "experts_cls", + None, + ) + if not isinstance(experts_cls, type) or issubclass( + experts_cls, TritonExperts + ): + _FUSED_MOE_TRITON_KERNEL.register_warmup() + if not isinstance(experts_cls, type) or issubclass( + experts_cls, BatchedTritonExperts + ): + _BATCHED_TRITON_KERNEL.register_warmup() + if isinstance(experts_cls, type) and issubclass( + experts_cls, (DeepGemmExperts, DeepGemmFP4Experts) + ): + from vllm.model_executor.layers.fused_moe.deep_gemm_utils import ( # noqa: E501 + _DEEPGEMM_EP_GATHER_KERNEL, + _DEEPGEMM_EP_SCATTER_COPY_KERNEL, + _DEEPGEMM_EP_SCATTER_START_KERNEL, + ) + + _DEEPGEMM_EP_SCATTER_START_KERNEL.register_warmup() + _DEEPGEMM_EP_SCATTER_COPY_KERNEL.register_warmup() + _DEEPGEMM_EP_GATHER_KERNEL.register_warmup() + _COMPUTE_IDENTITY_KERNEL.register_warmup() + _MOE_FUSED_MUL_SUM_KERNEL.register_warmup() + _COUNT_EXPERT_NUM_TOKENS_KERNEL.register_warmup() + _SWIGLU_LIMIT_PAD_AWARE_KERNEL.register_warmup() + + from vllm.model_executor.layers.fused_moe.experts.nvfp4_emulation_moe import ( # noqa: E501 + _fused_moe_nvfp4_emulation, + Nvfp4QuantizationEmulationTritonExperts, + ) + + if isinstance(experts_cls, type) and issubclass( + experts_cls, Nvfp4QuantizationEmulationTritonExperts + ): + _fused_moe_nvfp4_emulation.register_warmup() + + if vllm_config.lora_config is not None: + from vllm.model_executor.layers.fused_moe.experts.trtllm_lora_moe import ( # noqa: E501 + _TRTLLM_LORA_FINALIZE_KERNEL, + _TRTLLM_LORA_UNPERMUTE_ACTIVATION_KERNEL, + ) + + _TRTLLM_LORA_UNPERMUTE_ACTIVATION_KERNEL.register_warmup() + _TRTLLM_LORA_FINALIZE_KERNEL.register_warmup() + + if vllm_config.parallel_config.all2all_backend == "deepep_v2": + from vllm.model_executor.layers.fused_moe.prepare_finalize.deepep_v2 import ( # noqa: E501 + _GLOBALIZE_RECV_TOPK_IDX_KERNEL, + ) + + _GLOBALIZE_RECV_TOPK_IDX_KERNEL.register_warmup() + def _init_mega_moe_experts( self, vllm_config: VllmConfig, diff --git a/vllm/models/deepseek_v4/nvidia/ops/prepare_megamoe.py b/vllm/models/deepseek_v4/nvidia/ops/prepare_megamoe.py index d254dad9d6a1..222ebe4833eb 100644 --- a/vllm/models/deepseek_v4/nvidia/ops/prepare_megamoe.py +++ b/vllm/models/deepseek_v4/nvidia/ops/prepare_megamoe.py @@ -7,10 +7,21 @@ MegaMoE kernels consume. """ +from typing import Any + import torch +from vllm.model_executor.warmup.jit_warmup import WarmupChoices +from vllm.model_executor.warmup.jit_warmup_triton_helper import ( + DispatchSpec, + TritonWarmupTensor, + triton_kernel_dispatcher_with_warmup, +) from vllm.triton_utils import tl, triton +_PREPARE_MEGAMOE_BLOCK_K = 128 +_PREPARE_MEGAMOE_GROUP_K = 32 + @triton.jit def _prepare_megamoe_inputs_kernel( @@ -145,7 +156,59 @@ def _prepare_megamoe_inputs_kernel( ) -def prepare_megamoe_inputs( +def _prepare_megamoe_inputs_warmup_inputs(vllm_config: Any) -> dict[str, Any]: + hf_config = vllm_config.model_config.hf_text_config + hidden_size = hf_config.hidden_size + top_k = hf_config.num_experts_per_tok + max_tokens = vllm_config.scheduler_config.max_num_batched_tokens + shared_block_m: Any = WarmupChoices( + 1, + *( + (8, 16, 32, 64, 96, 128, 192) + if getattr( + vllm_config.model_config.hf_text_config, + "n_shared_experts", + None, + ) + is not None + else () + ), + ) + has_padding: Any = WarmupChoices(False, True) + padding_aligned: Any = WarmupChoices(False, True) + x_scale_width = hidden_size // _PREPARE_MEGAMOE_BLOCK_K + shared_rows = triton.cdiv(shared_block_m, 128) * 128 + # Mirrors DeepGEMM's get_num_max_shared_sf_tokens buffer layout. + shared_stride_k = triton.cdiv(max_tokens, 384) * 384 * 16 + return dict( + hidden_states=TritonWarmupTensor(torch.bfloat16, shape=(1, hidden_size)), + topk_weights=TritonWarmupTensor(torch.float32, shape=(1, top_k)), + topk_ids=TritonWarmupTensor(torch.int64, shape=(1, top_k)), + x_fp8=TritonWarmupTensor(torch.float8_e4m3fn, shape=(1, hidden_size)), + x_sf=TritonWarmupTensor(torch.int32, shape=(1, x_scale_width)), + topk_idx_out=TritonWarmupTensor(torch.int64, shape=(1, top_k)), + topk_weights_out=TritonWarmupTensor(torch.float32, shape=(1, top_k)), + is_padding=( + TritonWarmupTensor(torch.bool, aligned=padding_aligned) + if has_padding + else None + ), + shared_x_sf=TritonWarmupTensor( + torch.int32, + shape=(shared_rows, x_scale_width), + strides=(1, shared_stride_k), + ) + if shared_block_m != 1 + else None, + shared_block_m=shared_block_m if shared_block_m != 1 else None, + ) + + +@triton_kernel_dispatcher_with_warmup( + kernel=_prepare_megamoe_inputs_kernel, + warmup_inputs=_prepare_megamoe_inputs_warmup_inputs, +) +def _PREPARE_MEGAMOE_INPUTS_KERNEL( hidden_states: torch.Tensor, topk_weights: torch.Tensor, topk_ids: torch.Tensor, @@ -156,10 +219,8 @@ def prepare_megamoe_inputs( is_padding: torch.Tensor | None = None, shared_x_sf: torch.Tensor | None = None, shared_block_m: int | None = None, -) -> None: +) -> DispatchSpec: num_tokens, hidden_size = hidden_states.shape - if num_tokens == 0: - return if hidden_size % 128 != 0: raise ValueError( "DeepSeek V4 MegaMoE input staging requires hidden_size to be " @@ -180,8 +241,8 @@ def prepare_megamoe_inputs( assert shared_block_m is not None if shared_block_m <= 0: raise ValueError("MegaMoE shared_block_m must be positive.") - expected_sf_k = hidden_size // 128 - if shared_x_sf.ndim != 2 or shared_x_sf.shape[1] != expected_sf_k: + expected_sf_k = hidden_size // _PREPARE_MEGAMOE_BLOCK_K + if len(shared_x_sf.shape) != 2 or shared_x_sf.shape[1] != expected_sf_k: raise ValueError( "MegaMoE shared_x_sf must have shape " f"(*, {expected_sf_k}), got {tuple(shared_x_sf.shape)}." @@ -194,41 +255,32 @@ def prepare_megamoe_inputs( f"{required_rows}, got {shared_x_sf.shape[0]}." ) - block_k = 128 - grid = (num_tokens, triton.cdiv(hidden_size, block_k)) + block_k = _PREPARE_MEGAMOE_BLOCK_K block_topk = triton.next_power_of_2(top_k) + grid = (num_tokens, triton.cdiv(hidden_size, block_k)) padding_stride_m = is_padding.stride(0) if is_padding is not None else 0 - _prepare_megamoe_inputs_kernel[grid]( - hidden_states, - x_fp8, - x_sf, - shared_x_sf, - topk_ids, - topk_weights, - is_padding, - topk_idx_out, - topk_weights_out, - hidden_states.stride(0), - hidden_states.stride(1), - x_fp8.stride(0), - x_fp8.stride(1), - x_sf.stride(0), - x_sf.stride(1), - shared_x_sf.stride(0) if shared_x_sf is not None else 0, - shared_x_sf.stride(1) if shared_x_sf is not None else 0, - topk_ids.stride(0), - topk_ids.stride(1), - topk_weights.stride(0), - topk_weights.stride(1), - padding_stride_m, - topk_idx_out.stride(0), - topk_idx_out.stride(1), - topk_weights_out.stride(0), - topk_weights_out.stride(1), - hidden_size, - top_k, + return grid, dict( + hidden_stride_m=hidden_states.stride(0), + hidden_stride_k=hidden_states.stride(1), + x_stride_m=x_fp8.stride(0), + x_stride_k=x_fp8.stride(1), + x_sf_stride_m=x_sf.stride(0), + x_sf_stride_k=x_sf.stride(1), + shared_x_sf_stride_m=shared_x_sf.stride(0) if shared_x_sf is not None else 0, + shared_x_sf_stride_k=shared_x_sf.stride(1) if shared_x_sf is not None else 0, + topk_ids_stride_m=topk_ids.stride(0), + topk_ids_stride_k=topk_ids.stride(1), + topk_weights_stride_m=topk_weights.stride(0), + topk_weights_stride_k=topk_weights.stride(1), + is_padding_stride_m=padding_stride_m, + topk_idx_stride_m=topk_idx_out.stride(0), + topk_idx_stride_k=topk_idx_out.stride(1), + topk_weights_out_stride_m=topk_weights_out.stride(0), + topk_weights_out_stride_k=topk_weights_out.stride(1), + hidden_size=hidden_size, + top_k=top_k, BLOCK_K=block_k, - GROUP_K=32, + GROUP_K=_PREPARE_MEGAMOE_GROUP_K, BLOCK_TOPK=block_topk, SHARED_BLOCK_M=shared_block_m or 1, num_warps=4, diff --git a/vllm/models/kimi_k3/nvidia/model.py b/vllm/models/kimi_k3/nvidia/model.py index c4f0092f8d9b..f1dcd6ede52a 100644 --- a/vllm/models/kimi_k3/nvidia/model.py +++ b/vllm/models/kimi_k3/nvidia/model.py @@ -98,7 +98,9 @@ DeepseekV4MegaMoEExperts, DeepseekV4MLP, ) -from vllm.models.deepseek_v4.nvidia.ops.prepare_megamoe import prepare_megamoe_inputs +from vllm.models.deepseek_v4.nvidia.ops.prepare_megamoe import ( + _PREPARE_MEGAMOE_INPUTS_KERNEL, +) from vllm.models.kimi_k3.nvidia.kda import KimiK3DeltaAttention from vllm.models.kimi_k3.nvidia.latent_moe_runner import ( LatentMoERunner, @@ -512,16 +514,17 @@ def forward( else None, ) - prepare_megamoe_inputs( - hidden_states, - topk_weights, - topk_ids, - symm_buffer.x[:num_tokens], - symm_buffer.x_sf[:num_tokens], - symm_buffer.topk_idx[:num_tokens], - symm_buffer.topk_weights[:num_tokens], - is_padding=is_padding, - ) + if num_tokens > 0: + _PREPARE_MEGAMOE_INPUTS_KERNEL( + hidden_states, + topk_weights, + topk_ids, + symm_buffer.x[:num_tokens], + symm_buffer.x_sf[:num_tokens], + symm_buffer.topk_idx[:num_tokens], + symm_buffer.topk_weights[:num_tokens], + is_padding=is_padding, + ) self.finalize_weights() assert self._transformed_l1_weights is not None assert self._transformed_l2_weights is not None