Rename & fix

Signed-off-by: Woosuk Kwon <woosuk@inferact.ai>
This commit is contained in:
Woosuk Kwon
2026-04-08 21:58:38 +00:00
parent 88fa073594
commit 4f1d426261
4 changed files with 38 additions and 415 deletions
+2 -2
View File
@@ -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(
@@ -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
```
@@ -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
@@ -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,