diff --git a/vllm/envs.py b/vllm/envs.py index cdb6fb163e0..3d0efffaee9 100755 --- a/vllm/envs.py +++ b/vllm/envs.py @@ -217,7 +217,7 @@ if TYPE_CHECKING: VLLM_USE_FLASHINFER_MOE_MXFP4_MXFP8_CUTLASS: bool = False VLLM_ALLREDUCE_USE_SYMM_MEM: bool = True VLLM_ALLREDUCE_USE_FLASHINFER: bool = False - VLLM_USE_SPECIALIZED_MODELS: bool = True + VLLM_USE_SPECIALIZED_MODELS: bool = False VLLM_TUNED_CONFIG_FOLDER: str | None = None VLLM_GPT_OSS_SYSTEM_TOOL_MCP_LABELS: set[str] = set() VLLM_USE_EXPERIMENTAL_PARSER_CONTEXT: bool = False @@ -1530,7 +1530,7 @@ environment_variables: dict[str, Callable[[], Any]] = { ), # Whether to enable specialized model implementations when available. "VLLM_USE_SPECIALIZED_MODELS": lambda: bool( - int(os.getenv("VLLM_USE_SPECIALIZED_MODELS", "1")) + int(os.getenv("VLLM_USE_SPECIALIZED_MODELS", "0")) ), # Experimental: use this to enable MCP tool calling for non harmony models "VLLM_USE_EXPERIMENTAL_PARSER_CONTEXT": lambda: bool( diff --git a/vllm/model_executor/specialized_models/deepseek_v3_2_nvfp4/README.md b/vllm/model_executor/specialized_models/deepseek_v3_2_nvfp4/README.md new file mode 100644 index 00000000000..d3c6d571cfc --- /dev/null +++ b/vllm/model_executor/specialized_models/deepseek_v3_2_nvfp4/README.md @@ -0,0 +1,28 @@ +# nvidia/DeepSeek-V3.2-NVFP4 + +An optimized implementation for `nvidia/DeepSeek-V3.2-NVFP4` with FP8 FlashInfer MLA on Blackwell GPUs. + +The main win comes from aggressively fusing ops in the attention path, across the MLA and sparse-indexer boundary, which is critical for low latency. +On top of manual fusions, the implementation uses `torch.compile` with vLLM's custom fusion passes to fuse remaining miscellaneous ops. +It is compatible with piecewise CUDA graphs for prefill and full CUDA graphs for decode. + +TP and EP are supported; PP is not. +MTP is supported. + +## Usage + +```bash +# With TP +VLLM_USE_SPECIALIZED_MODELS=1 VLLM_USE_V2_MODEL_RUNNER=1 vllm serve nvidia/DeepSeek-V3.2-NVFP4 \ + -tp 8 \ + --compilation-config '{"max_cudagraph_capture_size": 1024}' \ + --speculative-config '{"method": "mtp", "num_speculative_tokens": 1}' \ + --kernel-config.enable_flashinfer_autotune=False + +# With attention DP + MoE EP +VLLM_USE_SPECIALIZED_MODELS=1 VLLM_USE_V2_MODEL_RUNNER=1 vllm serve nvidia/DeepSeek-V3.2-NVFP4 \ + -dp 8 -ep \ + --compilation-config '{"max_cudagraph_capture_size": 1024}' \ + --speculative-config '{"method": "mtp", "num_speculative_tokens": 1}' \ + --kernel-config.enable_flashinfer_autotune=False +``` diff --git a/vllm/model_executor/specialized_models/deepseek_v3_2_nvfp4/decoder_layer.py b/vllm/model_executor/specialized_models/deepseek_v3_2_nvfp4/decoder_layer.py index 9e5f5fb8bb5..de1024a7502 100644 --- a/vllm/model_executor/specialized_models/deepseek_v3_2_nvfp4/decoder_layer.py +++ b/vllm/model_executor/specialized_models/deepseek_v3_2_nvfp4/decoder_layer.py @@ -18,7 +18,7 @@ from vllm.utils.torch_utils import direct_register_custom_op from vllm.v1.attention.backends.mla.indexer import get_max_prefill_buffer_size from .attention import MonolithicMLAAttention -from .ops import fused_norm_rope, fused_q +from .kernels import fused_norm_rope, fused_q from .sparse_indexer import sparse_attn_indexer @@ -384,7 +384,7 @@ class MonolithicDecoderLayer(nn.Module): from vllm.utils.flashinfer import flashinfer_scaled_fp4_mm - from .ops import silu_and_mul_nvfp4_quant + from .kernels import silu_and_mul_nvfp4_quant dp = shared_experts.down_proj diff --git a/vllm/model_executor/specialized_models/deepseek_v3_2_nvfp4/ops.py b/vllm/model_executor/specialized_models/deepseek_v3_2_nvfp4/kernels.py similarity index 75% rename from vllm/model_executor/specialized_models/deepseek_v3_2_nvfp4/ops.py rename to vllm/model_executor/specialized_models/deepseek_v3_2_nvfp4/kernels.py index 00ec7fa9bd7..cde9f6eb71a 100644 --- a/vllm/model_executor/specialized_models/deepseek_v3_2_nvfp4/ops.py +++ b/vllm/model_executor/specialized_models/deepseek_v3_2_nvfp4/kernels.py @@ -14,193 +14,6 @@ def _rms_norm(x, w, eps, HIDDEN_SIZE: tl.constexpr): return (x * rrms) * w -@triton.jit -def _rms_norm_small_dim_kernel( - x_ptr, - x_stride, - w_ptr, - y_ptr, - y_stride, - eps, - HIDDEN_SIZE: tl.constexpr, - BLOCK_SIZE: tl.constexpr, -): - row_idx = tl.program_id(0) - x_row_ptr = x_ptr + row_idx * x_stride - y_row_ptr = y_ptr + row_idx * y_stride - - block = tl.arange(0, BLOCK_SIZE) - mask = block < HIDDEN_SIZE - x = tl.load(x_row_ptr + block, mask=mask, other=0.0) - w = tl.load(w_ptr + block, mask=mask) - y = _rms_norm(x, w, eps, HIDDEN_SIZE) - tl.store(y_row_ptr + block, y, mask=mask) - - -def rms_norm_small( - x: torch.Tensor, - w: torch.Tensor, - eps: float, -) -> torch.Tensor: - assert x.ndim == 2 - assert w.ndim == 1 - num_tokens, hidden_size = x.shape - y = torch.empty_like(x) - - _rms_norm_small_dim_kernel[(num_tokens,)]( - x, - x.stride(0), - w, - y, - y.stride(0), - eps, - hidden_size, - BLOCK_SIZE=triton.next_power_of_2(hidden_size), - ) - return y - - -@triton.jit -def _rms_norm_kernel( - x_ptr, - x_stride, - w_ptr, - y_ptr, - y_stride, - eps, - NUM_COLS: tl.constexpr, - BLOCK_SIZE: tl.constexpr, -): - row_idx = tl.program_id(0) - x_row_ptr = x_ptr + row_idx * x_stride - y_row_ptr = y_ptr + row_idx * y_stride - - sq_sum = tl.zeros([BLOCK_SIZE], dtype=tl.float32) - for i in range(0, NUM_COLS, BLOCK_SIZE): - offset = i + tl.arange(0, BLOCK_SIZE) - mask = offset < NUM_COLS - x = tl.load(x_row_ptr + offset, mask=mask, other=0.0).to(tl.float32) - sq_sum += x * x - - mean_sq = tl.sum(sq_sum, axis=0) / NUM_COLS - rrms = tl.rsqrt(mean_sq + eps) - - for i in range(0, NUM_COLS, BLOCK_SIZE): - offset = i + tl.arange(0, BLOCK_SIZE) - mask = offset < NUM_COLS - x = tl.load(x_row_ptr + offset, mask=mask).to(tl.float32) - w = tl.load(w_ptr + offset, mask=mask).to(tl.float32) - y = (x * rrms) * w - tl.store(y_row_ptr + offset, y, mask=mask) - - -def rms_norm( - x: torch.Tensor, - w: torch.Tensor, - eps: float, -) -> torch.Tensor: - assert x.ndim == 2 - assert w.ndim == 1 - num_tokens, hidden_size = x.shape - - y = torch.empty_like(x) - BLOCK_SIZE = 1024 # TODO: Tune this - _rms_norm_kernel[(num_tokens,)]( - x, - x.stride(0), - w, - y, - y.stride(0), - eps, - hidden_size, - BLOCK_SIZE, - ) - return y - - -@triton.jit -def _fused_add_rms_norm_kernel( - x_ptr, - x_stride, - residual_ptr, - residual_stride, - w_ptr, - y_ptr, - y_stride, - residual_new_ptr, - residual_new_stride, - eps, - W_OFFSET: tl.constexpr, - NUM_COLS: tl.constexpr, - BLOCK_SIZE: tl.constexpr, -): - row_idx = tl.program_id(0) - x_row_ptr = x_ptr + row_idx * x_stride - r_row_ptr = residual_ptr + row_idx * residual_stride - y_row_ptr = y_ptr + row_idx * y_stride - r_new_row_ptr = residual_new_ptr + row_idx * residual_new_stride - - sq_sum = tl.zeros([BLOCK_SIZE], dtype=tl.float32) - for i in range(0, NUM_COLS, BLOCK_SIZE): - offset = i + tl.arange(0, BLOCK_SIZE) - mask = offset < NUM_COLS - x = tl.load(x_row_ptr + offset, mask=mask, other=0.0).to(tl.float32) - r = tl.load(r_row_ptr + offset, mask=mask, other=0.0).to(tl.float32) - x = x + r - sq_sum += x * x - - mean_sq = tl.sum(sq_sum, axis=0) / NUM_COLS - rrms = tl.rsqrt(mean_sq + eps) - - for i in range(0, NUM_COLS, BLOCK_SIZE): - offset = i + tl.arange(0, BLOCK_SIZE) - mask = offset < NUM_COLS - - # Recompute x + r - x = tl.load(x_row_ptr + offset, mask=mask).to(tl.float32) - r = tl.load(r_row_ptr + offset, mask=mask).to(tl.float32) - x = x + r - - w = tl.load(w_ptr + offset, mask=mask).to(tl.float32) - y = (x * rrms) * (w + W_OFFSET) - tl.store(y_row_ptr + offset, y, mask=mask) - - tl.store(r_new_row_ptr + offset, x, mask=mask) - - -def fused_add_rms_norm( - x: torch.Tensor, - residual: torch.Tensor, - weight: torch.Tensor, - eps: float, - w_offset: float = 0.0, -) -> tuple[torch.Tensor, torch.Tensor]: - assert x.ndim == 2 - assert residual.shape == x.shape - assert weight.ndim == 1 - num_tokens, hidden_size = x.shape - - y = torch.empty_like(x) - residual_new = torch.empty_like(x) - BLOCK_SIZE = 1024 # TODO: Tune this - _fused_add_rms_norm_kernel[(num_tokens,)]( - x, - x.stride(0), - residual, - residual.stride(0), - weight, - y, - y.stride(0), - residual_new, - residual_new.stride(0), - eps, - w_offset, - hidden_size, - BLOCK_SIZE, - ) - return y, residual_new - - @triton.jit def _layer_norm(x, w, b, eps, mask, HIDDEN_SIZE: tl.constexpr): x = x.to(tl.float32) @@ -214,61 +27,8 @@ def _layer_norm(x, w, b, eps, mask, HIDDEN_SIZE: tl.constexpr): return (x - mean) * rstd * w + b -# Optimized for small hidden size @triton.jit -def _layer_norm_kernel( - x_ptr, - x_stride, - w_ptr, - bias_ptr, - y_ptr, - y_stride, - eps, - HIDDEN_SIZE: tl.constexpr, - BLOCK_SIZE: tl.constexpr, -): - row_idx = tl.program_id(0) - x_row_ptr = x_ptr + row_idx * x_stride - y_row_ptr = y_ptr + row_idx * y_stride - - block = tl.arange(0, BLOCK_SIZE) - mask = block < HIDDEN_SIZE - - x = tl.load(x_row_ptr + block, mask=mask, other=0.0) - w = tl.load(w_ptr + block, mask=mask) - b = tl.load(bias_ptr + block, mask=mask) - - y = _layer_norm(x, w, b, eps, mask, HIDDEN_SIZE) - tl.store(y_row_ptr + block, y, mask=mask) - - -def layer_norm( - x: torch.Tensor, - w: torch.Tensor, - bias: torch.Tensor, - eps: float, -) -> torch.Tensor: - assert x.ndim == 2 - assert w.ndim == 1 - assert bias.ndim == 1 - num_tokens, hidden_size = x.shape - y = torch.empty_like(x) - _layer_norm_kernel[(num_tokens,)]( - x, - x.stride(0), - w, - bias, - y, - y.stride(0), - eps, - hidden_size, - BLOCK_SIZE=triton.next_power_of_2(hidden_size), - ) - return y - - -@triton.jit -def _rope_kernel( +def _rope( base_ptr, head_stride, cos, @@ -294,7 +54,7 @@ def _rope_kernel( @triton.jit -def _cos_sin_cache_kernel( +def _get_cos_sin( cos_sin_cache_ptr, cos_sin_cache_stride, pos, @@ -308,89 +68,6 @@ def _cos_sin_cache_kernel( return cos, sin -@triton.jit -def _qk_rope_kernel( - q_ptr, - q_stride0, - q_stride1, - NUM_Q_HEADS: tl.constexpr, - Q_START_OFFSET: tl.constexpr, - k_ptr, - k_stride0, - k_stride1, - NUM_K_HEADS: tl.constexpr, - pos_ptr, - cos_sin_ptr, - cos_sin_stride, - HALF_ROT_DIM: tl.constexpr, - INTERLEAVED: tl.constexpr, -): - tok_idx = tl.program_id(1) - pos = tl.load(pos_ptr + tok_idx) - - cos, sin = _cos_sin_cache_kernel( - cos_sin_ptr, - cos_sin_stride, - pos, - HALF_ROT_DIM, - ) - - if tl.program_id(0) == 0: - # Handle Q [NUM_Q_HEADS, ROT_DIM] - q_base_ptr = q_ptr + tok_idx * q_stride0 - _rope_kernel( - q_base_ptr, - q_stride1, - cos, - sin, - NUM_Q_HEADS, - HALF_ROT_DIM, - Q_START_OFFSET, - INTERLEAVED, - ) - elif tl.program_id(0) == 1: - # Handle K [NUM_K_HEADS, ROT_DIM] - k_base_ptr = k_ptr + tok_idx * k_stride0 - _rope_kernel( - k_base_ptr, k_stride1, cos, sin, NUM_K_HEADS, HALF_ROT_DIM, 0, INTERLEAVED - ) - - -def qk_rope( - positions: torch.Tensor, - q: torch.Tensor, - k: torch.Tensor, - cos_sin_cache: torch.Tensor, - q_start_offset: int, - interleaved: bool, -) -> None: - assert q.ndim == 3 - assert k.ndim == 3 - assert q.shape[0] == k.shape[0] - assert cos_sin_cache.ndim == 2 - assert positions.ndim == 1 - assert q_start_offset < q.shape[-1] - num_tokens, num_q_heads, _ = q.shape - num_tokens, num_k_heads, _ = k.shape - rot_dim = cos_sin_cache.shape[-1] - _qk_rope_kernel[(2, num_tokens)]( - q, - q.stride(0), - q.stride(1), - num_q_heads, - q_start_offset, - k, - k.stride(0), - k.stride(1), - num_k_heads, - positions, - cos_sin_cache, - cos_sin_cache.stride(0), - rot_dim // 2, - interleaved, - ) - - @triton.jit def _fp8_ue8m0_quantize(vals): """Quantize float32 values to FP8 E4M3 with a ue8m0 (power-of-2) scale. @@ -432,88 +109,6 @@ def _fp8_quant_and_cache_write( tl.store(kv_cache_scale_ptr + scale_byte_off // 4, scale) -@triton.jit -def _concat_quant_fp8_kernel( - ql_nope_ptr, # [B, N, L] - q_pe_ptr, # [B, N, P] - out_ptr, # [B, N, L+P], fp8 - scale_ptr, # [1], float32 (static per-tensor scale) - nope_stride_b, - nope_stride_n, - pe_stride_b, - pe_stride_n, - out_stride_b, - out_stride_n, - L: tl.constexpr, # kv_lora_rank (512) - P: tl.constexpr, # qk_rope_head_dim (64) - L_BLOCK: tl.constexpr, - P_BLOCK: tl.constexpr, -): - """Concatenate ql_nope and q_pe, then quantize to FP8 with a static scale. - - Grid: (N, B) where N = num_heads, B = num_tokens. - """ - head_idx = tl.program_id(0) - tok_idx = tl.program_id(1) - scale = tl.load(scale_ptr) - - # Load ql_nope [L] and quantize - l_off = tl.arange(0, L_BLOCK) - l_mask = l_off < L - nope = tl.load( - ql_nope_ptr + tok_idx * nope_stride_b + head_idx * nope_stride_n + l_off, - mask=l_mask, - ).to(tl.float32) - nope_fp8 = (nope / scale).to(tl.float8e4nv) - tl.store( - out_ptr + tok_idx * out_stride_b + head_idx * out_stride_n + l_off, - nope_fp8, - mask=l_mask, - ) - - # Load q_pe [P] and quantize - p_off = tl.arange(0, P_BLOCK) - p_mask = p_off < P - pe = tl.load( - q_pe_ptr + tok_idx * pe_stride_b + head_idx * pe_stride_n + p_off, - mask=p_mask, - ).to(tl.float32) - pe_fp8 = (pe / scale).to(tl.float8e4nv) - tl.store( - out_ptr + tok_idx * out_stride_b + head_idx * out_stride_n + L + p_off, - pe_fp8, - mask=p_mask, - ) - - -def concat_quant_fp8( - ql_nope: torch.Tensor, # [B, N, L] - q_pe: torch.Tensor, # [B, N, P] - scale: torch.Tensor, # [1] -) -> torch.Tensor: - """Fused concat + per-tensor FP8 quantization for MLA decode query.""" - B, N, L = ql_nope.shape - P = q_pe.shape[2] - out = torch.empty(B, N, L + P, dtype=torch.float8_e4m3fn, device=ql_nope.device) - _concat_quant_fp8_kernel[(N, B)]( - ql_nope, - q_pe, - out, - scale, - ql_nope.stride(0), - ql_nope.stride(1), - q_pe.stride(0), - q_pe.stride(1), - out.stride(0), - out.stride(1), - L=L, - P=P, - L_BLOCK=triton.next_power_of_2(L), - P_BLOCK=triton.next_power_of_2(P), - ) - return out - - @triton.jit def _fused_norm_rope_kernel( pos_ptr, @@ -616,7 +211,7 @@ def _fused_norm_rope_kernel( # k_pe is not needed after the cache write (MLA decode reads # from kv_cache), so we skip writing back to kpe_ptr. pos = tl.load(pos_ptr + tok_idx) - cos, sin = _cos_sin_cache_kernel( + cos, sin = _get_cos_sin( kpe_rope_cos_sin_cache_ptr, kpe_rope_cos_sin_cache_stride, pos, @@ -945,7 +540,7 @@ def _fused_q_kernel( return pos = tl.load(pos_ptr + tok_idx) - cos, sin = _cos_sin_cache_kernel( + cos, sin = _get_cos_sin( q_pe_cos_sin_ptr, q_pe_cos_sin_stride, pos, @@ -996,13 +591,13 @@ def _fused_q_kernel( return pos = tl.load(pos_ptr + tok_idx) - cos, sin = _cos_sin_cache_kernel( + cos, sin = _get_cos_sin( index_q_cos_sin_ptr, index_q_cos_sin_stride, pos, INDEX_Q_HALF_ROT_DIM, ) - _rope_kernel( + _rope( index_q_ptr + tok_idx * index_q_stride0 + head_idx * index_q_stride1, 0, cos,