Compare commits

..
Author SHA1 Message Date
NickLucche e0dd86f6c3 deprecate timeout
Signed-off-by: NickLucche <nlucches@redhat.com>
2026-04-28 09:03:27 +02:00
NickLucche 01f19dee7f deprecate timeout
Signed-off-by: NickLucche <nlucches@redhat.com>
2026-04-28 09:01:01 +02:00
Moritz SanftandGitHub 2c06cf3486 [Bugfix] use served_model_name for multimodal error message (#41003)
Signed-off-by: Moritz Sanft <58110325+msanft@users.noreply.github.com>
2026-04-27 08:22:35 -07:00
Harry MellorGitHubmergify[bot] <37929162+mergify[bot]@users.noreply.github.com>
e6f710a87f Deprecate support for Transformers v4 (#40389)
Signed-off-by: Harry Mellor <19981378+hmellor@users.noreply.github.com>
Co-authored-by: mergify[bot] <37929162+mergify[bot]@users.noreply.github.com>
2026-04-27 08:19:57 -07:00
c245d35ff4 [Model] Add MiMo-V2.5 support (#40967)
Signed-off-by: Jee Jee Li <pandaleefree@gmail.com>
Signed-off-by: Isotr0py <mozf@mail2.sysu.edu.cn>
Signed-off-by: Isotr0py <Isotr0py@outlook.com>
Signed-off-by: zjy0516 <riverclouds.zhu@qq.com>
Co-authored-by: Jee Jee Li <pandaleefree@gmail.com>
Co-authored-by: zjy0516 <riverclouds.zhu@qq.com>
Co-authored-by: zjy0516 <zhujiangyun@inferact.ai>
Co-authored-by: yasong <yasong.wang@inferact.ai>
Co-authored-by: Claude Opus 4.6 <noreply@anthropic.com>
Co-authored-by: Copilot <copilot@github.com>
2026-04-27 13:26:51 +00:00
Xiaoshuang WangGitHubmergify[bot] <37929162+mergify[bot]@users.noreply.github.com>
f8ac0c7cf0 [Bugfix] Fix k_norm weight sharding in MiniMaxM2Attention when total_num_kv_heads < tp_size (#38191)
Signed-off-by: wxsIcey <1790571317@qq.com>
Signed-off-by: Icey <1790571317@qq.com>
Co-authored-by: mergify[bot] <37929162+mergify[bot]@users.noreply.github.com>
2026-04-27 05:57:13 -07:00
ebf862c351 Add system_fingerprint field to OpenAI-compatible API responses (#40537)
Co-authored-by: Claude <noreply@anthropic.com>
2026-04-27 16:17:52 +08:00
wang.yuqiandGitHub 8d8062d0a7 [Examples] Resettle generate examples. (#36464)
Signed-off-by: wang.yuqi <yuqi.wang@daocloud.io>
2026-04-27 07:48:37 +00:00
985961345a [Bugfix] Install libcublas-dev in Dockerfile for FlashInfer CuTe DSL JIT (#39855)
Signed-off-by: esmeetu <jasonailu87@gmail.com>
Co-authored-by: Roger Wang <hey@rogerw.io>
2026-04-27 15:47:39 +08:00
Yongye ZhuandGitHub 706a04d34b [DSV4] Add silu clamp limit to shared expert (#40950)
Signed-off-by: Yongye Zhu <zyy1102000@gmail.com>
2026-04-27 00:37:43 -07:00
Isotr0pyandGitHub 22631f80a0 [Bugfix] Remove invalid deepstack boundary check for Qwen3-VL (#40932)
Signed-off-by: Isotr0py <mozf@mail2.sysu.edu.cn>
2026-04-27 07:27:06 +00:00
BhoomitandGitHub 2cc008e7b4 [Attention][TurboQuant] Share dequant buffers, eliminate float16_copy (#40941)
Signed-off-by: Bhoomit Vasani <bhoomit.2010@gmail.com>
Signed-off-by: Vasani Bhoomit <bhoomit.2010@gmail.com>
2026-04-27 13:48:36 +08:00
5d5c776444 [Perf] FP8 FlashInfer Attn for ViT (#38065)
Signed-off-by: Zhanda Zhu <zhandazhu@gmail.com>
Co-authored-by: Yubo Gao <ybgao-nvidia@users.noreply.github.com>
2026-04-27 13:44:15 +08:00
ojhaanshikaandGitHub 592ae6805c Cutlass W4A16 (Machete) Tests (#35450)
Signed-off-by: Anshika Ojha <anshikao@nvidia.com>
2026-04-27 05:15:29 +00:00
7b1bc0a3eb [Bugfix] Cap SWA/chunked-local runtime admission to startup pool-sizing bound (#40946)
Signed-off-by: Dao Le <Dao007forever@gmail.com>
Signed-off-by: Nick Hill <nickhill123@gmail.com>
Co-authored-by: Claude <noreply@anthropic.com>
Co-authored-by: Nick Hill <nickhill123@gmail.com>
2026-04-27 04:33:13 +00:00
Silu PandaandGitHub c0879d9483 [Tests] Gate Isaac under Transformers v5 (#40907)
Signed-off-by: Silu Panda <31051721+SiluPanda@users.noreply.github.com>
2026-04-26 19:26:51 -07:00
Giancarlo DelfinandGitHub f5f9878514 [Model Runner V2] Fix rejection sampling acceptance rate gap vs MRV1 (#40651)
Signed-off-by: Giancarlo Delfin <gdelfin@inferact.ai>
2026-04-26 19:12:08 -07:00
youkaichaoGitHubClaudegemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com>
2ce95a761b Auto-disable expandable_segments around cumem memory pool (#40812)
Signed-off-by: youkaichao <youkaichao@gmail.com>
Co-authored-by: Claude <noreply@anthropic.com>
Co-authored-by: gemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com>
2026-04-27 09:37:22 +08:00
+8 4d51588e23 [Feat] DeepSeek V4 Rebased (#40860)
Signed-off-by: Yifan Qiao <yifanqiao@inferact.ai>
Signed-off-by: Woosuk Kwon <woosuk@inferact.ai>
Signed-off-by: qizixi <zixi@inferact.ai>
Signed-off-by: Jee Jee Li <pandaleefree@gmail.com>
Signed-off-by: Yongye Zhu <zyy1102000@gmail.com>
Co-authored-by: Yongye Zhu <zyy1102000@gmail.com>
Co-authored-by: Yongye Zhu <yongye@inferact.ai>
Co-authored-by: Simon Mo <simon@inferact.ai>
Co-authored-by: Bugen Zhao <i@bugenzhao.com>
Co-authored-by: Giancarlo Delfin <gdelfin@inferact.ai>
Co-authored-by: Jee Jee Li <pandaleefree@gmail.com>
Co-authored-by: Nick Hill <nickhill123@gmail.com>
Co-authored-by: Roger Wang <hey@rogerw.io>
Co-authored-by: Roy Wang <yasong.wang@inferact.ai>
Co-authored-by: Woosuk Kwon <woosuk@inferact.ai>
Co-authored-by: youkaichao <youkaichao@gmail.com>
Co-authored-by: Zhewen Li <jerven.vllm@gmail.com>
Co-authored-by: Zijing Liu <liuzijing2014@gmail.com>
Co-authored-by: khluu <khluu000@gmail.com>
Co-authored-by: qizixi <zixi@inferact.ai>
Co-authored-by: Zhewen Li <zhewenli@inferact.ai>
2026-04-26 18:31:08 -07:00
Xinan MiaoGitHubSouthWest7gemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com>OpenAI CodexWang Xingranmergify[bot] <37929162+mergify[bot]@users.noreply.github.com>
32e45636e3 [torch.compile]: Disable Sequence Parallelism (SP) for piecewise compilation (#38373)
Signed-off-by: SouthWest7 <am1ao@qq.com>
Signed-off-by: Xinan Miao <1403572259@qq.com>
Co-authored-by: SouthWest7 <am1ao@qq.com>
Co-authored-by: gemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com>
Co-authored-by: OpenAI Codex <codex@openai.com>
Co-authored-by: Wang Xingran <72983099+wangxingran222@users.noreply.github.com>
Co-authored-by: mergify[bot] <37929162+mergify[bot]@users.noreply.github.com>
2026-04-26 17:44:42 +00:00
b39c266dae [KV Offload] Offload all KV blocks when doing prefill in P/D (#40346)
Signed-off-by: omerpaz95 <omerpaz95@gmail.com>
Signed-off-by: omerpaz95 <73347585+omerpaz95@users.noreply.github.com>
Co-authored-by: Or Ozeri <or@ozery.com>
2026-04-26 15:06:01 +03:00
Dao007foreverandGitHub 9558f43903 [Bugfix] Size FlashInfer NVLink MNNVL workspace to EP group (#40893)
Signed-off-by: Dao Le <Dao007forever@gmail.com>
2026-04-26 01:26:34 -07:00
146 changed files with 8195 additions and 3606 deletions
+14 -14
View File
@@ -388,10 +388,10 @@ steps:
- python3 basic/offline_inference/embed.py
- python3 basic/offline_inference/score.py
# Multi-modal models
- python3 offline_inference/audio_language.py --seed 0
- python3 offline_inference/vision_language.py --seed 0
- python3 offline_inference/vision_language_multi_image.py --seed 0
- python3 offline_inference/encoder_decoder_multimodal.py --model-type whisper --seed 0
- python3 generate/multimodal/audio_language_offline.py --seed 0
- python3 generate/multimodal/vision_language_offline.py --seed 0
- python3 generate/multimodal/vision_language_multi_image_offline.py --seed 0
- python3 generate/multimodal/encoder_decoder_multimodal_offline.py --model-type whisper --seed 0
# Pooling models
- python3 pooling/embed/vision_embedding_offline.py --seed 0
# Features demo
@@ -1647,10 +1647,10 @@ steps:
- python3 basic/offline_inference/embed.py
- python3 basic/offline_inference/score.py
# Multi-modal models
- python3 offline_inference/audio_language.py --seed 0
- python3 offline_inference/vision_language.py --seed 0
- python3 offline_inference/vision_language_multi_image.py --seed 0
- python3 offline_inference/encoder_decoder_multimodal.py --model-type whisper --seed 0
- python3 generate/multimodal/audio_language_offline.py --seed 0
- python3 generate/multimodal/vision_language_offline.py --seed 0
- python3 generate/multimodal/vision_language_multi_image_offline.py --seed 0
- python3 generate/multimodal/encoder_decoder_multimodal_offline.py --model-type whisper --seed 0
# Pooling models
- python3 pooling/embed/vision_embedding_offline.py --seed 0
# Features demo
@@ -1951,8 +1951,8 @@ steps:
- pytest -v -s tests/models/multimodal/processing/
- pytest -v -s tests/models/multimodal/test_mapping.py
- python3 examples/basic/offline_inference/chat.py
- python3 examples/offline_inference/vision_language.py --model-type qwen2_5_vl
- VLLM_WORKER_MULTIPROC_METHOD=spawn python3 examples/offline_inference/audio_language.py --model-type whisper
- python3 examples/generate/multimodal/vision_language_offline.py --model-type qwen2_5_vl
- VLLM_WORKER_MULTIPROC_METHOD=spawn python3 examples/generate/multimodal/audio_language_offline.py --model-type whisper
#------------------------------------------------------- mi300 · quantization --------------------------------------------------------#
@@ -2930,10 +2930,10 @@ steps:
- python3 basic/offline_inference/embed.py
- python3 basic/offline_inference/score.py
# Multi-modal models
- python3 offline_inference/audio_language.py --seed 0
- python3 offline_inference/vision_language.py --seed 0
- python3 offline_inference/vision_language_multi_image.py --seed 0
- python3 offline_inference/encoder_decoder_multimodal.py --model-type whisper --seed 0
- python3 generate/multimodal/audio_language_offline.py --seed 0
- python3 generate/multimodal/vision_language_offline.py --seed 0
- python3 generate/multimodal/vision_language_multi_image_offline.py --seed 0
- python3 generate/multimodal/encoder_decoder_multimodal_offline.py --model-type whisper --seed 0
# Pooling models
- python3 pooling/embed/vision_embedding_offline.py --seed 0
# Features demo
+2
View File
@@ -95,11 +95,13 @@ steps:
- tests/kernels/moe/test_deepgemm.py
- tests/kernels/moe/test_batched_deepgemm.py
- tests/kernels/attention/test_deepgemm_attention.py
- tests/quantization/test_cutlass_w4a16.py
commands:
- pytest -v -s kernels/quantization/test_block_fp8.py
- pytest -v -s kernels/moe/test_deepgemm.py
- pytest -v -s kernels/moe/test_batched_deepgemm.py
- pytest -v -s kernels/attention/test_deepgemm_attention.py
- pytest -v -s quantization/test_cutlass_w4a16.py
- label: Kernels (B200)
timeout_in_minutes: 30
+4 -4
View File
@@ -113,10 +113,10 @@ steps:
- python3 basic/offline_inference/embed.py
- python3 basic/offline_inference/score.py
# for multi-modal models
- python3 offline_inference/audio_language.py --seed 0
- python3 offline_inference/vision_language.py --seed 0
- python3 offline_inference/vision_language_multi_image.py --seed 0
- python3 offline_inference/encoder_decoder_multimodal.py --model-type whisper --seed 0
- python3 generate/multimodal/audio_language_offline.py --seed 0
- python3 generate/multimodal/vision_language_offline.py --seed 0
- python3 generate/multimodal/vision_language_multi_image_offline.py --seed 0
- python3 generate/multimodal/encoder_decoder_multimodal_offline.py --model-type whisper --seed 0
# for pooling models
- python3 pooling/embed/vision_embedding_offline.py --seed 0
# for features demo
+4 -4
View File
@@ -44,10 +44,10 @@ steps:
#- python3 basic/offline_inference/generate.py --model meta-llama/Llama-2-13b-chat-hf --cpu-offload-gb 10 # TODO
#- python3 basic/offline_inference/embed.py # TODO
# for multi-modal models
- python3 offline_inference/audio_language.py --seed 0
- python3 offline_inference/vision_language.py --seed 0
- python3 offline_inference/vision_language_multi_image.py --seed 0
- python3 offline_inference/encoder_decoder_multimodal.py --model-type whisper --seed 0
- python3 generate/multimodal/audio_language_offline.py --seed 0
- python3 generate/multimodal/vision_language_offline.py --seed 0
- python3 generate/multimodal/vision_language_multi_image_offline.py --seed 0
- python3 generate/multimodal/encoder_decoder_multimodal_offline.py --model-type whisper --seed 0
# for pooling models
- python3 pooling/embed/vision_embedding_offline.py --seed 0
# for features demo
+5 -5
View File
@@ -69,9 +69,9 @@ steps:
- pytest -v -s tests/models/multimodal/processing/
- pytest -v -s tests/models/multimodal/test_mapping.py
- python3 examples/basic/offline_inference/chat.py
- python3 examples/offline_inference/vision_language.py --model-type qwen2_5_vl
- python3 examples/generate/multimodal/vision_language_offline.py --model-type qwen2_5_vl
# Whisper needs spawn method to avoid deadlock
- VLLM_WORKER_MULTIPROC_METHOD=spawn python3 examples/offline_inference/audio_language.py --model-type whisper
- VLLM_WORKER_MULTIPROC_METHOD=spawn python3 examples/generate/multimodal/audio_language_offline.py --model-type whisper
- label: Transformers Backward Compatibility Models Test
working_dir: "/vllm-workspace/"
@@ -83,7 +83,7 @@ steps:
- pytest -v -s tests/models/test_transformers.py
- pytest -v -s tests/models/multimodal/processing/
- pytest -v -s tests/models/multimodal/test_mapping.py
- python3 examples/offline_inference/basic/chat.py
- python3 examples/offline_inference/vision_language.py --model-type qwen2_5_vl
- python3 examples/basic/offline_inference/chat.py
- python3 examples/generate/multimodal/vision_language_offline.py --model-type qwen2_5_vl
# Whisper needs spawn method to avoid deadlock
- VLLM_WORKER_MULTIPROC_METHOD=spawn python3 examples/offline_inference/audio_language.py --model-type whisper
- VLLM_WORKER_MULTIPROC_METHOD=spawn python3 examples/generate/multimodal/audio_language_offline.py --model-type whisper
+1 -5
View File
@@ -389,11 +389,7 @@ pull_request_rules:
- files~=^tests/entrypoints/anthropic/.*tool.*
- files~=^vllm/tool_parsers/
- files=docs/features/tool_calling.md
- files~=^examples/tool_chat_*
- files=examples/offline_inference/chat_with_tools.py
- files=examples/online_serving/openai_chat_completion_client_with_tools_required.py
- files=examples/online_serving/openai_chat_completion_tool_calls_with_reasoning.py
- files=examples/online_serving/openai_chat_completion_client_with_tools.py
- files~=^examples/tool_calling/
actions:
label:
add:
-21
View File
@@ -564,27 +564,6 @@ if(VLLM_GPU_LANG STREQUAL "CUDA")
"in CUDA target architectures.")
endif()
# DeepSeek V4 indexer top-k. Needs thread-block clusters + TMA + PDL, so
# builds for Hopper (sm_90a) and Blackwell datacenter (sm_100/sm_103). Not
# supported on sm_120 (consumer Blackwell, no clusters). Requires CUDA >=
# 12.4 for the cuda::ptx mbarrier wrappers. Ported from sglang.
if(${CMAKE_CUDA_COMPILER_VERSION} VERSION_GREATER_EQUAL 13.0)
cuda_archs_loose_intersection(DSV4_TOPK_ARCHS "9.0a;10.0f;11.0f" "${CUDA_ARCHS}")
else()
cuda_archs_loose_intersection(DSV4_TOPK_ARCHS "9.0a;10.0a;10.1a;10.3a" "${CUDA_ARCHS}")
endif()
if(${CMAKE_CUDA_COMPILER_VERSION} VERSION_GREATER_EQUAL 12.4 AND DSV4_TOPK_ARCHS)
set(DSV4_TOPK_SRC "csrc/deepseek_v4/fast_topk_v2.cu")
set_gencode_flags_for_srcs(
SRCS "${DSV4_TOPK_SRC}"
CUDA_ARCHS "${DSV4_TOPK_ARCHS}")
list(APPEND VLLM_EXT_SRC ${DSV4_TOPK_SRC})
message(STATUS "Building deepseek_v4 fast_topk_v2 for archs: ${DSV4_TOPK_ARCHS}")
else()
message(STATUS "Not building deepseek_v4 fast_topk_v2 (needs CUDA >= 12.4 "
"and a compatible Hopper+ arch).")
endif()
#
# Machete kernels
@@ -1,183 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
"""Microbench: fast_topk_v2 vs persistent_topk for k in {512, 1024}.
Both ops select the top-k entries per row of a `[B, L]` float32 score
tensor. vLLM's `persistent_topk` is the existing path used by the indexer;
`fast_topk_v2` is the sm_90+ port from sglang that adds Hopper thread-block
clusters.
V4-Flash uses `index_topk = 512`; V4-Pro uses `index_topk = 1024`. We bench
both Ks at the realistic shape regimes (small-B, L up to 256K compressed).
Timing uses **CUDA graph replay** to amortize launch overhead (~3-5 µs on
Blackwell). We capture N invocations of the same kernel, replay the graph
many times, divide.
Run::
.venv/bin/python benchmarks/kernels/benchmark_fast_topk_v2.py
"""
from __future__ import annotations
import argparse
import statistics
import sys
import torch
import vllm._C # noqa: F401 ensures schemas are registered
from vllm.v1.attention.ops.deepseek_v4_ops.fast_topk import (
fast_topk_v2_raw,
plan_topk_v2,
workspace_ints_per_batch,
)
RADIX_TOPK_WORKSPACE_SIZE = 1024 * 1024 # bytes; matches sparse_attn_indexer.py
def _capture_graph(callable_fn, *, calls_per_graph: int) -> torch.cuda.CUDAGraph:
for _ in range(3):
callable_fn()
torch.cuda.synchronize()
g = torch.cuda.CUDAGraph()
s = torch.cuda.Stream()
s.wait_stream(torch.cuda.current_stream())
with torch.cuda.stream(s):
with torch.cuda.graph(g, stream=s):
for _ in range(calls_per_graph):
callable_fn()
torch.cuda.current_stream().wait_stream(s)
return g
def time_graph_us(graph: torch.cuda.CUDAGraph, *, calls_per_graph: int,
warmup: int = 5, replays: int = 30) -> float:
for _ in range(warmup):
graph.replay()
torch.cuda.synchronize()
samples = []
start = torch.cuda.Event(enable_timing=True)
end = torch.cuda.Event(enable_timing=True)
for _ in range(replays):
start.record()
graph.replay()
end.record()
end.synchronize()
samples.append(start.elapsed_time(end) * 1000.0 / calls_per_graph)
return statistics.median(samples)
def make_inputs(batch_size: int, seq_len: int, *, seed: int = 0):
device = torch.device("cuda")
g = torch.Generator(device=device).manual_seed(seed)
L = (seq_len + 3) & ~3
scores = torch.randn(batch_size, L, generator=g, dtype=torch.float32,
device=device)
seq_lens = torch.full((batch_size,), seq_len, dtype=torch.int32,
device=device)
return scores, seq_lens, L
def bench_persistent_topk(scores, seq_lens, k, *, calls_per_graph: int) -> float:
B = scores.shape[0]
output = scores.new_empty((B, k), dtype=torch.int32)
workspace = scores.new_empty((RADIX_TOPK_WORKSPACE_SIZE,), dtype=torch.uint8)
max_seq_len = scores.shape[1]
def run():
torch.ops._C.persistent_topk(
scores, seq_lens, output, workspace, k, max_seq_len)
graph = _capture_graph(run, calls_per_graph=calls_per_graph)
return time_graph_us(graph, calls_per_graph=calls_per_graph)
def bench_fast_topk_v2(scores, seq_lens, k, *,
calls_per_graph: int) -> float:
B = scores.shape[0]
metadata = plan_topk_v2(seq_lens)
workspace = scores.new_empty((B, workspace_ints_per_batch()),
dtype=torch.int32)
topk_indices = scores.new_empty((B, k), dtype=torch.int32)
def run():
fast_topk_v2_raw(scores, seq_lens, topk=k,
metadata=metadata, workspace=workspace,
topk_indices=topk_indices)
graph = _capture_graph(run, calls_per_graph=calls_per_graph)
return time_graph_us(graph, calls_per_graph=calls_per_graph)
def fmt(us: float) -> str:
return f"{us:8.2f}"
def main():
parser = argparse.ArgumentParser()
parser.add_argument("--batch-sizes", type=int, nargs="+",
default=[1, 4, 16, 32, 64, 128, 256])
parser.add_argument("--seq-lens", type=int, nargs="+",
default=[1024, 4096, 16384, 32768, 65536, 131072])
parser.add_argument("--ks", type=int, nargs="+",
default=[512, 1024])
parser.add_argument("--calls-per-graph", type=int, default=64)
parser.add_argument("--replays", type=int, default=30)
args = parser.parse_args()
if not torch.cuda.is_available():
print("CUDA is required for this benchmark.", file=sys.stderr)
sys.exit(1)
print(f"GPU: {torch.cuda.get_device_name(0)} "
f"(SM {torch.cuda.get_device_capability(0)})")
print(f"calls_per_graph={args.calls_per_graph}, replays={args.replays}")
print("Per-call medians via CUDA graph replay (host launch overhead "
"amortized).\n")
for k in args.ks:
print(f"=== k = {k} ===")
print(f"{'B':>4} {'L':>7} | {'persistent_topk':>17} | "
f"{'fast_topk_v2':>14} | {'speedup':>8} | {'path':<14}")
print("-" * 80)
for B in args.batch_sizes:
for L in args.seq_lens:
# Skip seq_lens beyond persistent_topk's k-dependent useful
# range. Both kernels handle up to 256K with k=1024.
try:
scores, seq_lens, _ = make_inputs(B, L, seed=B * L * k)
p_us = bench_persistent_topk(
scores, seq_lens, k,
calls_per_graph=args.calls_per_graph)
f_us = bench_fast_topk_v2(
scores, seq_lens, k,
calls_per_graph=args.calls_per_graph)
speedup = p_us / f_us if f_us > 0 else float("inf")
if L <= k:
path = "trivial"
elif L <= 4 * 4 * 1024:
path = "register-1p"
elif L <= 32768:
path = "register-2p"
elif B <= 15:
path = "cluster-fused"
else:
path = "cluster-2stg"
print(
f"{B:>4} {L:>7} | "
f"{fmt(p_us):>14} us | "
f"{fmt(f_us):>11} us | "
f"{speedup:>5.2f}x | {path}"
)
except RuntimeError as e:
print(f"{B:>4} {L:>7} | ERROR: {e}")
print()
if __name__ == "__main__":
main()
@@ -0,0 +1,324 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
# Benchmarks FP8 vs BF16 ViT attention via FlashInfer cuDNN backend.
#
# == Usage Examples ==
#
# Benchmark mode (default, FlashInfer CUDAGraph Bench)
# python3 benchmark_vit_fp8_attn.py
#
# Profile mode (PyTorch profiler, saves TensorBoard traces):
# python3 benchmark_vit_fp8_attn.py --profile
# python3 benchmark_vit_fp8_attn.py --profile --profile-output-dir ./profile_traces
#
# Custom seq_lens:
# python3 benchmark_vit_fp8_attn.py --seq-lens 4096 8192 16384
from functools import partial
import numpy as np
import torch
from torch.profiler import ProfilerActivity, profile, record_function
from vllm.utils.argparse_utils import FlexibleArgumentParser
# Qwen3-VL defaults
NUM_HEADS = 16
HEAD_DIM = 72
DEFAULT_SEQ_LENS = [2304, 4096, 8192, 16384]
def _setup_fp8_attention(num_heads: int, head_dim: int) -> tuple:
"""Create FP8 and BF16 attention modules + workspace."""
from types import SimpleNamespace
from unittest.mock import patch
from vllm.config import VllmConfig, set_current_vllm_config
from vllm.config.multimodal import MultiModalConfig
from vllm.model_executor.layers.attention.mm_encoder_attention import (
MMEncoderAttention,
_get_flashinfer_workspace_buffer,
)
from vllm.v1.attention.backends.registry import AttentionBackendEnum
old_dtype = torch.get_default_dtype()
torch.set_default_dtype(torch.bfloat16)
backend_patch = patch(
"vllm.model_executor.layers.attention.mm_encoder_attention"
".get_vit_attn_backend",
return_value=AttentionBackendEnum.FLASHINFER,
)
# FP8 attention
mm_config_fp8 = MultiModalConfig(mm_encoder_attn_dtype="fp8")
vllm_config_fp8 = VllmConfig()
vllm_config_fp8.model_config = SimpleNamespace(multimodal_config=mm_config_fp8)
with set_current_vllm_config(vllm_config_fp8), backend_patch:
attn_fp8 = MMEncoderAttention(
num_heads=num_heads,
head_size=head_dim,
prefix="visual.blocks.0.attn",
).to("cuda")
# BF16 attention (no FP8)
with set_current_vllm_config(VllmConfig()), backend_patch:
attn_bf16 = MMEncoderAttention(
num_heads=num_heads,
head_size=head_dim,
prefix="visual.blocks.0.attn",
).to("cuda")
torch.set_default_dtype(old_dtype)
workspace = _get_flashinfer_workspace_buffer()
return attn_fp8, attn_bf16, workspace
def _build_meta(
seq_len: int,
num_heads: int,
head_dim: int,
fp8: bool,
):
"""Build cu_seqlens, max_seqlen, sequence_lengths."""
from vllm.model_executor.layers.attention.mm_encoder_attention import (
MMEncoderAttention,
)
from vllm.utils.math_utils import round_up
from vllm.v1.attention.backends.registry import AttentionBackendEnum
cu_np = np.array([0, seq_len], dtype=np.int32)
fp8_padded = num_heads * round_up(head_dim, 16) if fp8 else None
seq_lengths = MMEncoderAttention.maybe_compute_seq_lens(
AttentionBackendEnum.FLASHINFER, cu_np, torch.device("cuda")
)
max_seqlen = torch.tensor(
MMEncoderAttention.compute_max_seqlen(AttentionBackendEnum.FLASHINFER, cu_np),
dtype=torch.int32,
)
cu_seqlens = MMEncoderAttention.maybe_recompute_cu_seqlens(
AttentionBackendEnum.FLASHINFER,
cu_np,
num_heads * head_dim,
1,
torch.device("cuda"),
fp8_padded_hidden_size=fp8_padded,
)
return cu_seqlens, max_seqlen, seq_lengths
def run_benchmark(
seq_lens: list[int],
num_heads: int,
head_dim: int,
method: str,
):
"""Benchmark FP8 vs BF16 attention across seq_lens.
Uses FlashInfer GPU-level timing to measure pure kernel time,
excluding CPU launch overhead.
"""
if method == "cupti":
from flashinfer.testing import bench_gpu_time_with_cupti as bench_fn
bench_fn = partial(bench_fn, use_cuda_graph=True, cold_l2_cache=False)
elif method == "cudagraph":
from flashinfer.testing import (
bench_gpu_time_with_cudagraph as bench_fn,
)
bench_fn = partial(bench_fn, cold_l2_cache=False)
else:
raise ValueError(f"Invalid method: {method}")
attn_fp8, attn_bf16, workspace = _setup_fp8_attention(num_heads, head_dim)
print(f"Timing method: {method}")
print(f"{'seq_len':>8} {'BF16 (us)':>12} {'FP8 (us)':>12} {'Speedup':>10}")
print("-" * 46)
for seq_len in seq_lens:
torch.manual_seed(42)
q = torch.randn(
seq_len,
num_heads,
head_dim,
device="cuda",
dtype=torch.bfloat16,
)
k = torch.randn_like(q)
v = torch.randn_like(q)
cu_fp8, max_s, seq_l = _build_meta(seq_len, num_heads, head_dim, fp8=True)
# we can reuse cu_fp8 for cu_bf16 since q, k, and v are contiguous
cu_bf16 = cu_fp8.clone()
def bf16_fn(q=q, k=k, v=v, cu=cu_bf16, ms=max_s, sl=seq_l):
attn_bf16._forward_flashinfer(q, k, v, cu, ms, sl)
def fp8_fn(q=q, k=k, v=v, cu=cu_fp8, ms=max_s, sl=seq_l):
attn_fp8._forward_flashinfer(q, k, v, cu, ms, sl)
# bench_fn returns List[float] of per-iteration times in ms
bf16_times = bench_fn(bf16_fn)
fp8_times = bench_fn(fp8_fn)
bf16_us = np.median(bf16_times) * 1e3 # ms -> us
fp8_us = np.median(fp8_times) * 1e3
speedup = bf16_us / fp8_us if fp8_us > 0 else float("inf")
print(f"{seq_len:>8} {bf16_us:>12.1f} {fp8_us:>12.1f} {speedup:>9.2f}x")
def _make_trace_handler(output_dir: str, worker_name: str, label: str):
"""Create a trace handler that saves to TensorBoard and prints summary."""
def handler(prof):
torch.profiler.tensorboard_trace_handler(output_dir, worker_name)(prof)
print(f"\n{'=' * 80}")
print(label)
print(f"{'=' * 80}")
print(prof.key_averages().table(sort_by="cuda_time_total", row_limit=20))
return handler
def run_profile(
seq_len: int,
num_heads: int,
head_dim: int,
warmup: int,
output_dir: str,
):
"""Profile FP8 vs BF16 attention with PyTorch profiler."""
attn_fp8, attn_bf16, workspace = _setup_fp8_attention(num_heads, head_dim)
torch.manual_seed(42)
q = torch.randn(
seq_len,
num_heads,
head_dim,
device="cuda",
dtype=torch.bfloat16,
)
k = torch.randn_like(q)
v = torch.randn_like(q)
cu_fp8, max_s, seq_l = _build_meta(seq_len, num_heads, head_dim, fp8=True)
# we can reuse cu_fp8 for cu_bf16 since q, k, and v are contiguous
cu_bf16 = cu_fp8.clone()
sched = torch.profiler.schedule(wait=0, warmup=warmup, active=1)
# Profile BF16 (warmup handled by profiler schedule)
with profile(
activities=[ProfilerActivity.CPU, ProfilerActivity.CUDA],
schedule=sched,
on_trace_ready=_make_trace_handler(
output_dir,
f"bf16_h{head_dim}_s{seq_len}",
f"BF16 Attention (seq_len={seq_len}, heads={num_heads}, "
f"head_dim={head_dim})",
),
) as prof_bf16:
for _ in range(warmup + 1):
with record_function("bf16_attention"):
attn_bf16._forward_flashinfer(
q.clone(), k.clone(), v.clone(), cu_bf16, max_s, seq_l
)
torch.accelerator.synchronize()
prof_bf16.step()
# Profile FP8 (warmup handled by profiler schedule)
with profile(
activities=[ProfilerActivity.CPU, ProfilerActivity.CUDA],
schedule=sched,
on_trace_ready=_make_trace_handler(
output_dir,
f"fp8_h{head_dim}_s{seq_len}",
f"FP8 Attention (seq_len={seq_len}, heads={num_heads}, "
f"head_dim={head_dim})",
),
) as prof_fp8:
for _ in range(warmup + 1):
with record_function("fp8_attention"):
attn_fp8._forward_flashinfer(
q.clone(), k.clone(), v.clone(), cu_fp8, max_s, seq_l
)
torch.accelerator.synchronize()
prof_fp8.step()
print(f"\nTensorBoard traces saved to: {output_dir}")
print(f"View with: tensorboard --logdir={output_dir}")
if __name__ == "__main__":
parser = FlexibleArgumentParser(description="Benchmark FP8 vs BF16 ViT attention.")
parser.add_argument(
"--seq-lens",
type=int,
nargs="+",
default=DEFAULT_SEQ_LENS,
help="Sequence lengths to benchmark",
)
parser.add_argument(
"--num-heads",
type=int,
default=NUM_HEADS,
)
parser.add_argument(
"--head-dim",
type=int,
default=HEAD_DIM,
)
parser.add_argument(
"--method",
choices=["cupti", "cudagraph"],
default="cudagraph",
help="GPU timing method: cupti (CUPTI kernel timing) or "
"cudagraph (CUDA graph capture/replay). Default: cudagraph",
)
parser.add_argument(
"--warmup",
type=int,
default=10,
help="Warmup iterations (profile mode only)",
)
parser.add_argument(
"--profile",
action="store_true",
help="Run PyTorch profiler instead of benchmark",
)
parser.add_argument(
"--profile-seq-len",
type=int,
default=8192,
help="Sequence length for profiling (default: 8192)",
)
parser.add_argument(
"--profile-output-dir",
type=str,
default="./profile_traces",
help="Output directory for TensorBoard traces (default: ./profile_traces)",
)
args = parser.parse_args()
if args.profile:
run_profile(
args.profile_seq_len,
args.num_heads,
args.head_dim,
args.warmup,
args.profile_output_dir,
)
else:
run_benchmark(
args.seq_lens,
args.num_heads,
args.head_dim,
args.method,
)
+82 -25
View File
@@ -11,29 +11,74 @@
namespace vllm {
template <typename scalar_t, scalar_t (*ACT_FN)(const scalar_t&),
bool act_first>
bool act_first, bool HAS_CLAMP>
__device__ __forceinline__ scalar_t compute(const scalar_t& x,
const scalar_t& y) {
return act_first ? ACT_FN(x) * y : x * ACT_FN(y);
const scalar_t& y,
const float limit) {
if constexpr (act_first) {
scalar_t gate = x;
scalar_t up = y;
if constexpr (HAS_CLAMP) {
gate = (scalar_t)fminf((float)gate, limit);
up = (scalar_t)fmaxf(fminf((float)up, limit), -limit);
}
return ACT_FN(gate) * up;
} else {
scalar_t gate = x;
scalar_t up = y;
if constexpr (HAS_CLAMP) {
gate = (scalar_t)fmaxf(fminf((float)gate, limit), -limit);
up = (scalar_t)fminf((float)up, limit);
}
return gate * ACT_FN(up);
}
}
template <typename packed_t, packed_t (*PACKED_ACT_FN)(const packed_t&),
bool act_first>
bool act_first, bool HAS_CLAMP>
__device__ __forceinline__ packed_t packed_compute(const packed_t& x,
const packed_t& y) {
return act_first ? packed_mul(PACKED_ACT_FN(x), y)
: packed_mul(x, PACKED_ACT_FN(y));
const packed_t& y,
const float limit) {
if constexpr (act_first) {
packed_t gate = x;
packed_t up = y;
if constexpr (HAS_CLAMP) {
float2 g = cast_to_float2(gate);
float2 u = cast_to_float2(up);
g.x = fminf(g.x, limit);
g.y = fminf(g.y, limit);
u.x = fmaxf(fminf(u.x, limit), -limit);
u.y = fmaxf(fminf(u.y, limit), -limit);
gate = cast_to_packed<packed_t>(g);
up = cast_to_packed<packed_t>(u);
}
return packed_mul(PACKED_ACT_FN(gate), up);
} else {
packed_t gate = x;
packed_t up = y;
if constexpr (HAS_CLAMP) {
float2 g = cast_to_float2(gate);
float2 u = cast_to_float2(up);
g.x = fmaxf(fminf(g.x, limit), -limit);
g.y = fmaxf(fminf(g.y, limit), -limit);
u.x = fminf(u.x, limit);
u.y = fminf(u.y, limit);
gate = cast_to_packed<packed_t>(g);
up = cast_to_packed<packed_t>(u);
}
return packed_mul(gate, PACKED_ACT_FN(up));
}
}
// Activation and gating kernel template.
template <typename scalar_t, typename packed_t,
scalar_t (*ACT_FN)(const scalar_t&),
packed_t (*PACKED_ACT_FN)(const packed_t&), bool act_first,
bool use_vec, bool use_256b = false>
bool use_vec, bool HAS_CLAMP, bool use_256b = false>
__global__ void act_and_mul_kernel(
scalar_t* __restrict__ out, // [..., d]
const scalar_t* __restrict__ input, // [..., 2, d]
const int d) {
const int d, const float limit) {
const scalar_t* x_ptr = input + blockIdx.x * 2 * d;
const scalar_t* y_ptr = x_ptr + d;
scalar_t* out_ptr = out + blockIdx.x * d;
@@ -58,8 +103,9 @@ __global__ void act_and_mul_kernel(
}
#pragma unroll
for (int j = 0; j < pvec_t::NUM_ELTS; j++) {
x.elts[j] = packed_compute<packed_t, PACKED_ACT_FN, act_first>(
x.elts[j], y.elts[j]);
x.elts[j] =
packed_compute<packed_t, PACKED_ACT_FN, act_first, HAS_CLAMP>(
x.elts[j], y.elts[j], limit);
}
if constexpr (use_256b) {
st256(x, &out_vec[i]);
@@ -72,7 +118,8 @@ __global__ void act_and_mul_kernel(
for (int64_t idx = threadIdx.x; idx < d; idx += blockDim.x) {
const scalar_t x = VLLM_LDG(&x_ptr[idx]);
const scalar_t y = VLLM_LDG(&y_ptr[idx]);
out_ptr[idx] = compute<scalar_t, ACT_FN, act_first>(x, y);
out_ptr[idx] =
compute<scalar_t, ACT_FN, act_first, HAS_CLAMP>(x, y, limit);
}
}
}
@@ -151,8 +198,11 @@ packed_gelu_tanh_kernel(const packed_t& val) {
// Launch activation and gating kernel.
// Use ACT_FIRST (bool) indicating whether to apply the activation function
// first.
#define LAUNCH_ACTIVATION_GATE_KERNEL(KERNEL, PACKED_KERNEL, ACT_FIRST) \
// first. HAS_CLAMP (bool) enables pre-activation clamping: gate input is
// clamped (max only) and up input is clamped (both sides) before the
// activation function is applied.
#define LAUNCH_ACTIVATION_GATE_KERNEL(KERNEL, PACKED_KERNEL, ACT_FIRST, \
HAS_CLAMP, LIMIT) \
auto dtype = input.scalar_type(); \
int d = input.size(-1) / 2; \
int64_t num_tokens = input.numel() / input.size(-1); \
@@ -177,8 +227,8 @@ packed_gelu_tanh_kernel(const packed_t& val) {
scalar_t, typename vllm::PackedTypeConverter<scalar_t>::Type, \
KERNEL<scalar_t>, \
PACKED_KERNEL<typename vllm::PackedTypeConverter<scalar_t>::Type>, \
ACT_FIRST, true, true><<<grid, block, 0, stream>>>( \
out.data_ptr<scalar_t>(), input.data_ptr<scalar_t>(), d); \
ACT_FIRST, true, HAS_CLAMP, true><<<grid, block, 0, stream>>>( \
out.data_ptr<scalar_t>(), input.data_ptr<scalar_t>(), d, LIMIT); \
}); \
} else { \
VLLM_DISPATCH_FLOATING_TYPES(dtype, "act_and_mul_kernel", [&] { \
@@ -186,8 +236,8 @@ packed_gelu_tanh_kernel(const packed_t& val) {
scalar_t, typename vllm::PackedTypeConverter<scalar_t>::Type, \
KERNEL<scalar_t>, \
PACKED_KERNEL<typename vllm::PackedTypeConverter<scalar_t>::Type>, \
ACT_FIRST, true, false><<<grid, block, 0, stream>>>( \
out.data_ptr<scalar_t>(), input.data_ptr<scalar_t>(), d); \
ACT_FIRST, true, HAS_CLAMP, false><<<grid, block, 0, stream>>>( \
out.data_ptr<scalar_t>(), input.data_ptr<scalar_t>(), d, LIMIT); \
}); \
} \
} else { \
@@ -197,8 +247,8 @@ packed_gelu_tanh_kernel(const packed_t& val) {
scalar_t, typename vllm::PackedTypeConverter<scalar_t>::Type, \
KERNEL<scalar_t>, \
PACKED_KERNEL<typename vllm::PackedTypeConverter<scalar_t>::Type>, \
ACT_FIRST, false><<<grid, block, 0, stream>>>( \
out.data_ptr<scalar_t>(), input.data_ptr<scalar_t>(), d); \
ACT_FIRST, false, HAS_CLAMP><<<grid, block, 0, stream>>>( \
out.data_ptr<scalar_t>(), input.data_ptr<scalar_t>(), d, LIMIT); \
}); \
}
@@ -206,7 +256,14 @@ void silu_and_mul(torch::Tensor& out, // [..., d]
torch::Tensor& input) // [..., 2 * d]
{
LAUNCH_ACTIVATION_GATE_KERNEL(vllm::silu_kernel, vllm::packed_silu_kernel,
true);
true, false, 0.0f);
}
void silu_and_mul_clamp(torch::Tensor& out, // [..., d]
torch::Tensor& input, // [..., 2 * d]
double limit) {
LAUNCH_ACTIVATION_GATE_KERNEL(vllm::silu_kernel, vllm::packed_silu_kernel,
true, true, (float)limit);
}
void mul_and_silu(torch::Tensor& out, // [..., d]
@@ -215,21 +272,21 @@ void mul_and_silu(torch::Tensor& out, // [..., d]
// The difference between mul_and_silu and silu_and_mul is that mul_and_silu
// applies the silu to the latter half of the input.
LAUNCH_ACTIVATION_GATE_KERNEL(vllm::silu_kernel, vllm::packed_silu_kernel,
false);
false, false, 0.0f);
}
void gelu_and_mul(torch::Tensor& out, // [..., d]
torch::Tensor& input) // [..., 2 * d]
{
LAUNCH_ACTIVATION_GATE_KERNEL(vllm::gelu_kernel, vllm::packed_gelu_kernel,
true);
true, false, 0.0f);
}
void gelu_tanh_and_mul(torch::Tensor& out, // [..., d]
torch::Tensor& input) // [..., 2 * d]
{
LAUNCH_ACTIVATION_GATE_KERNEL(vllm::gelu_tanh_kernel,
vllm::packed_gelu_tanh_kernel, true);
LAUNCH_ACTIVATION_GATE_KERNEL(
vllm::gelu_tanh_kernel, vllm::packed_gelu_tanh_kernel, true, false, 0.0f);
}
namespace vllm {
-693
View File
@@ -1,693 +0,0 @@
// SPDX-License-Identifier: Apache-2.0
// SPDX-FileCopyrightText: Copyright contributors to the vLLM project
//
// DeepSeek V4 indexer top-k (k = 512 for Flash, k = 1024 for Pro). Ported
// from sglang's jit_kernel/csrc/deepseek_v4/topk_v2.cuh.
//
// Combines three strategies (Register / Streaming / Cluster) dispatched per
// row by a separate plan kernel that decides a `cluster_threshold` from the
// observed seq_lens distribution. The host side picks one of three launch
// shapes:
// 1. all rows fit in the small (register) path -> single short kernel
// 2. small batch (<= kNumClusters) with some long rows -> fused cluster
// kernel (stage 1 + tie-break in one launch)
// 3. larger batch -> persistent cluster stage 1 + non-cluster stage 2
//
// Architecture support: Hopper (sm_90a) and Blackwell datacenter (sm_100/
// sm_103). Requires thread-block clusters, TMA bulk async copy, mbarrier,
// and Programmatic Dependent Launch — sm_120 (consumer Blackwell) lacks
// clusters and is not supported. The heuristic constants in `topk_plan`
// were tuned on B200 (sglang upstream); they are functionally correct on
// H100/H200 too but may be suboptimal until retuned.
#include "topk/cluster.cuh"
#include "topk/common.cuh"
#include "topk/register.cuh"
#include "topk/streaming.cuh"
#include "topk/utils.cuh"
#include "core/registration.h"
#include <ATen/cuda/CUDAContext.h>
#include <c10/util/Exception.h>
#include <cooperative_groups.h>
#include <cuda_runtime.h>
#include <torch/all.h>
#include <torch/library.h>
#include <algorithm>
#include <cstdint>
namespace vllm::dsv4_topk {
// All K-dependent type and constant lookups go through these aliases / vars
// so the kernels can be templated on K. Kernel and Smem sizes happen to be
// K-independent (e.g., kMaxTies, kMax2PassLength, kHistBins are all set in
// terms of kBlockSize/kHistBits, not K), so we don't pay extra smem for the
// 1024 instantiation.
template <uint32_t K> using Large = ClusterTopK<K>;
template <uint32_t K> using Medium = StreamingTopK<K>;
template <uint32_t K> using Small = RegisterTopK<K>;
// Metadata struct layout is K-independent — pick any K to grab the type.
using Metadata = Large<512>::Metadata;
constexpr uint32_t kNumClusters = 15; // hardware-capped persistent count
constexpr uint32_t kClusterSize = Large<512>::kClusterSize;
constexpr uint32_t kMax2PassLength = Small<512>::kMax2PassLength;
constexpr uint32_t kMaxSupportedLength = Large<512>::kMaxLength;
// Row 0 of the metadata tensor stores GlobalMetadata; rows [1..N+1) hold the
// per-item Metadata entries that the persistent stage-1 consumes.
struct alignas(16) GlobalMetadata {
uint32_t cluster_threshold;
uint32_t num_cluster_items;
uint32_t reserved[2];
};
static_assert(sizeof(GlobalMetadata) == sizeof(Metadata),
"metadata row 0 layout must match Metadata stride");
#define VLLM_SMALL_TOPK_KERNEL __global__ __launch_bounds__(kBlockSize, 2)
#define VLLM_LARGE_CLUSTER __cluster_dims__(1, kClusterSize, 1)
// Stage 1 is persistent + cluster -> high smem -> occupancy 1.
#define VLLM_LARGE_TOPK_STAGE_1 \
__global__ __launch_bounds__(kBlockSize, 1) VLLM_LARGE_CLUSTER
// Stage 2 is non-cluster + small smem -> occupancy 2.
#define VLLM_LARGE_TOPK_STAGE_2 __global__ __launch_bounds__(kBlockSize, 2)
#define VLLM_FUSED_COMBINE_KERNEL \
__global__ __launch_bounds__(kBlockSize, 1) VLLM_LARGE_CLUSTER
#define VLLM_PLAN_KERNEL __global__ __launch_bounds__(kBlockSize, 1)
struct TopKParams {
const uint32_t* __restrict__ seq_lens;
const float* __restrict__ scores;
const int32_t* __restrict__ page_table;
int32_t* __restrict__ page_indices;
int64_t score_stride;
int64_t page_table_stride;
uint8_t* __restrict__ workspace;
const Metadata* __restrict__ metadata = nullptr;
int64_t workspace_stride; // bytes per batch
uint32_t batch_size;
uint32_t page_bits;
VLLM_DSV4_DEVICE const float* get_scores(uint32_t batch_id) const {
return scores + batch_id * score_stride;
}
template <uint32_t K, bool kRawOutput>
VLLM_DSV4_DEVICE TransformParamsT<kRawOutput> get_transform(
uint32_t batch_id, int32_t* indices) const {
return {
.page_table = page_table + batch_id * page_table_stride,
.indices_in = indices,
.indices_out = page_indices + batch_id * K,
.page_bits = page_bits,
};
}
VLLM_DSV4_DEVICE const GlobalMetadata& get_global_metadata() const {
return *reinterpret_cast<const GlobalMetadata*>(metadata);
}
VLLM_DSV4_DEVICE const Metadata& get_item_metadata(uint32_t work_id) const {
return metadata[1 + work_id]; // skip the GlobalMetadata row
}
};
VLLM_DSV4_DEVICE uint2 partition_work(uint32_t length, uint32_t rank) {
constexpr uint32_t kTMAAlign = 4;
const auto total_units = (length + kTMAAlign - 1) / kTMAAlign;
const auto base = total_units / kClusterSize;
const auto extra = total_units % kClusterSize;
const auto local_units = base + (rank < extra ? 1u : 0u);
const auto offset_units = rank * base + min(rank, extra);
const auto offset = offset_units * kTMAAlign;
const auto finish = min(offset + local_units * kTMAAlign, length);
return {offset, finish - offset};
}
// --------------------------------------------------------------------------
// Plan kernel: decides cluster_threshold from the observed seq_lens
// distribution and compacts items with seq_len > threshold into metadata[1..].
// --------------------------------------------------------------------------
VLLM_PLAN_KERNEL void topk_plan(const uint32_t* __restrict__ seq_lens,
Metadata* __restrict__ metadata,
uint32_t batch_size,
uint32_t static_cluster_threshold) {
// (threshold, max_batch_size_for_that_threshold). Tuned on B200 by sglang.
struct Pair {
uint32_t threshold;
uint32_t max_batch_size;
};
constexpr Pair kCandidates[] = {
{32768, 30}, {40960, 45}, {49152, 45}, {65536, 60},
{98304, 60}, {131072, 75}, {196608, 90}, {262144, 105},
};
constexpr uint32_t kNumCandidates =
sizeof(kCandidates) / sizeof(kCandidates[0]);
constexpr uint32_t kMinBatchSize = kCandidates[0].max_batch_size;
static_assert(kCandidates[0].threshold == kMax2PassLength);
static_assert(kCandidates[kNumCandidates - 1].threshold ==
kMaxSupportedLength);
__shared__ uint32_t s_count;
__shared__ uint32_t s_counts[kNumCandidates];
__shared__ uint32_t s_threshold;
const auto tx = threadIdx.x;
if (tx == 0) s_count = 0;
if (tx < kNumCandidates) s_counts[tx] = 0;
__syncthreads();
if (static_cluster_threshold > 0) {
if (tx == 0) s_threshold = static_cluster_threshold;
} else if (batch_size <= kMinBatchSize) {
if (tx == 0) s_threshold = kMax2PassLength;
} else {
for (uint32_t i = tx; i < batch_size; i += kBlockSize) {
const uint32_t sl = seq_lens[i];
assert(sl <= kMaxSupportedLength);
uint32_t count = 0;
#pragma unroll
for (uint32_t j = 0; j < kNumCandidates; ++j) {
count += (sl > kCandidates[j].threshold ? 1 : 0);
}
if (count > 0) {
atomicAdd(&s_counts[count - 1], 1);
}
}
__syncthreads();
if (tx == 0) {
uint32_t accum = 0;
uint32_t chosen = kMaxSupportedLength;
#pragma unroll
for (uint32_t i = 0; i < kNumCandidates; ++i) {
const auto j = kNumCandidates - 1 - i;
accum += s_counts[j];
if (accum > kCandidates[j].max_batch_size) break;
chosen = kCandidates[j].threshold;
}
s_threshold = chosen;
}
}
__syncthreads();
const auto cluster_threshold = max(s_threshold, kMax2PassLength);
// Compact items with seq_len > cluster_threshold into metadata[1..N+1).
for (uint32_t i = tx; i < batch_size; i += kBlockSize) {
const uint32_t sl = seq_lens[i];
if (sl > cluster_threshold) {
const auto pos = atomicAdd(&s_count, 1);
metadata[1 + pos] = {i, sl, false};
}
}
__syncthreads();
const auto N = s_count;
// has_next chain for the persistent consumer + sentinel slots.
for (uint32_t i = tx; i < N; i += kBlockSize) {
if (i + kNumClusters < N) metadata[1 + i].has_next = true;
}
if (tx < kNumClusters && tx >= N) metadata[1 + tx] = {0, 0, false};
if (tx == 0) {
auto* g = reinterpret_cast<GlobalMetadata*>(metadata);
*g = {
.cluster_threshold = cluster_threshold,
.num_cluster_items = N,
.reserved = {0, 0},
};
}
}
// --------------------------------------------------------------------------
// Short kernel: all rows fit in the register path (max_seq_len <=
// Small::kMax1PassLength).
// --------------------------------------------------------------------------
template <uint32_t K, bool kRawOutput>
VLLM_SMALL_TOPK_KERNEL void topk_short_transform(
const __grid_constant__ TopKParams params) {
alignas(128) extern __shared__ uint8_t smem[];
__shared__ int32_t s_topk_indices[K];
const auto batch_id = blockIdx.x;
const auto seq_len = params.seq_lens[batch_id];
const auto transform =
params.template get_transform<K, kRawOutput>(batch_id, s_topk_indices);
if (seq_len <= K) {
trivial_transform(transform, seq_len, K);
} else {
Small<K>::run(params.get_scores(batch_id), s_topk_indices, seq_len, smem,
/*use_pdl=*/true);
pdl_trigger_secondary<true>();
Small<K>::transform(transform);
}
}
// --------------------------------------------------------------------------
// Persistent stage 1 (cluster). One CTA per cluster; the persistent block
// walks `metadata[1..N]` round-robin and runs Large::stage1 per item.
// --------------------------------------------------------------------------
template <uint32_t K, bool kRawOutput>
VLLM_LARGE_TOPK_STAGE_1 void topk_combine_preprocess(
const __grid_constant__ TopKParams params) {
alignas(128) extern __shared__ uint8_t smem[];
__shared__ int32_t s_topk_indices[K];
uint32_t work_id = blockIdx.x;
uint32_t batch_id = 0, seq_len = 0, length = 0, offset = 0;
bool has_next = false;
const auto cluster_rank = blockIdx.y;
const auto prefetch_metadata = [&] {
const auto m = params.get_item_metadata(work_id);
batch_id = m.batch_id;
seq_len = m.seq_len;
has_next = m.has_next;
work_id += kNumClusters;
};
const auto launch_prologue = [&] {
const auto partition = partition_work(seq_len, cluster_rank);
offset = partition.x;
length = partition.y;
Large<K>::stage1_prologue(params.get_scores(batch_id) + offset, length,
smem);
};
pdl_wait_primary<true>();
pdl_trigger_secondary<true>();
prefetch_metadata();
if (seq_len == 0) return;
Large<K>::stage1_init(smem);
launch_prologue();
while (true) {
const auto this_length = length;
const auto this_offset = offset;
const auto need_prefetch = has_next;
const auto transform =
params.template get_transform<K, kRawOutput>(batch_id, s_topk_indices);
const auto ws = params.workspace + batch_id * params.workspace_stride;
if (need_prefetch) prefetch_metadata();
Large<K>::stage1(s_topk_indices, this_length, smem, /*reuse=*/true);
if (need_prefetch) launch_prologue();
Large<K>::stage1_epilogue(transform, this_offset, ws, smem);
if (!need_prefetch) break;
}
}
// --------------------------------------------------------------------------
// Stage 2 (non-cluster). Per-row dispatch: trivial / Small / Medium / Large.
// --------------------------------------------------------------------------
template <uint32_t K, bool kRawOutput>
VLLM_LARGE_TOPK_STAGE_2 void topk_combine_transform(
const __grid_constant__ TopKParams params) {
alignas(128) extern __shared__ uint8_t smem[];
__shared__ int32_t s_topk_indices[K];
const auto batch_id = blockIdx.x;
const auto seq_len = params.seq_lens[batch_id];
const auto cluster_threshold = params.get_global_metadata().cluster_threshold;
const auto transform =
params.template get_transform<K, kRawOutput>(batch_id, s_topk_indices);
if (seq_len <= K) {
trivial_transform(transform, seq_len, K);
} else if (seq_len <= kMax2PassLength) {
if (seq_len <= Small<K>::kMax1PassLength) {
Small<K>::run(params.get_scores(batch_id), s_topk_indices, seq_len,
smem);
} else {
__syncwarp();
Small<K>::template run<true>(params.get_scores(batch_id),
s_topk_indices, seq_len, smem);
}
Small<K>::transform(transform);
} else if (seq_len <= cluster_threshold) {
Medium<K>::run(params.get_scores(batch_id), seq_len, s_topk_indices, smem);
Medium<K>::transform(transform, smem);
} else {
const auto ws = params.workspace + batch_id * params.workspace_stride;
pdl_wait_primary<true>();
Large<K>::transform(transform, ws, smem);
}
}
// --------------------------------------------------------------------------
// Fused kernel for small batches. Both stage 1 and the tie-break run inside
// the same launch; cluster rank 0 finishes the row.
// --------------------------------------------------------------------------
template <uint32_t K, bool kRawOutput>
VLLM_FUSED_COMBINE_KERNEL void topk_fused_transform(
const __grid_constant__ TopKParams params) {
alignas(128) extern __shared__ uint8_t smem[];
__shared__ int32_t s_topk_indices[K];
const auto batch_id = blockIdx.x;
const auto cluster_rank = blockIdx.y;
const auto seq_len = params.seq_lens[batch_id];
const auto transform =
params.template get_transform<K, kRawOutput>(batch_id, s_topk_indices);
if (seq_len <= K) {
if (cluster_rank != 0) return;
trivial_transform(transform, seq_len, K);
} else if (seq_len <= Small<K>::kMax1PassLength) {
if (cluster_rank != 0) return;
Small<K>::run(params.get_scores(batch_id), s_topk_indices, seq_len, smem,
/*use_pdl=*/true);
Small<K>::transform(transform);
} else {
const auto partition = partition_work(seq_len, cluster_rank);
const auto offset = partition.x;
const auto length = partition.y;
const auto ws = params.workspace + batch_id * params.workspace_stride;
Large<K>::stage1_init(smem);
pdl_wait_primary<true>();
Large<K>::stage1_prologue(params.get_scores(batch_id) + offset, length,
smem);
Large<K>::stage1(s_topk_indices, length, smem);
Large<K>::stage1_epilogue(transform, offset, ws, smem);
cooperative_groups::this_cluster().sync();
if (cluster_rank != 0) return;
Large<K>::transform(transform, ws, smem);
}
}
template <uint32_t K> constexpr size_t kStage1SMEM = sizeof(typename Large<K>::Smem) + 128;
template <uint32_t K> constexpr size_t kStage2SMEM =
(sizeof(typename Small<K>::Smem) > sizeof(typename Medium<K>::Smem)
? sizeof(typename Small<K>::Smem)
: sizeof(typename Medium<K>::Smem)) +
128;
// Per-(kernel, smem) memoization: each instantiation has its own static. This
// matters because cudaFuncSetAttribute is per-function and we want it to fire
// exactly once per kernel symbol.
template <auto* f, size_t kSmem>
void setup_kernel_smem_once() {
[[maybe_unused]] static const auto result = [] {
return cudaFuncSetAttribute(reinterpret_cast<const void*>(f),
cudaFuncAttributeMaxDynamicSharedMemorySize,
static_cast<int>(kSmem));
}();
TORCH_CHECK(result == cudaSuccess,
"fast_topk_v2: cudaFuncSetAttribute failed: ",
cudaGetErrorString(result));
}
// --------------------------------------------------------------------------
// Host-side launchers
// --------------------------------------------------------------------------
#define CHECK_CUDA(x) TORCH_CHECK(x.is_cuda(), #x " must be a CUDA tensor")
#define CHECK_DTYPE(x, t) \
TORCH_CHECK(x.scalar_type() == (t), #x " must be ", #t)
#define CHECK_CONTIG(x) TORCH_CHECK(x.is_contiguous(), #x " must be contiguous")
} // namespace vllm::dsv4_topk
void fast_topk_v2_plan(const torch::Tensor& seq_lens, torch::Tensor& metadata,
int64_t static_cluster_threshold) {
using namespace vllm::dsv4_topk;
CHECK_CUDA(seq_lens);
CHECK_CUDA(metadata);
CHECK_DTYPE(seq_lens, torch::kInt32);
CHECK_DTYPE(metadata, torch::kInt32);
TORCH_CHECK(seq_lens.dim() == 1);
TORCH_CHECK(metadata.dim() == 2 && metadata.size(1) == 4);
TORCH_CHECK(metadata.size(0) == seq_lens.size(0) + 1,
"metadata must be (batch_size + 1, 4)");
CHECK_CONTIG(seq_lens);
CHECK_CONTIG(metadata);
const auto batch_size = static_cast<uint32_t>(seq_lens.size(0));
if (batch_size <= kNumClusters) return; // metadata unused in fused path
const auto stream = at::cuda::getCurrentCUDAStream().stream();
cudaLaunchConfig_t cfg{};
cfg.gridDim = dim3(1);
cfg.blockDim = dim3(kBlockSize);
cfg.dynamicSmemBytes = 0;
cfg.stream = stream;
cfg.numAttrs = 0;
TORCH_CHECK(cudaLaunchKernelEx(
&cfg, &topk_plan,
reinterpret_cast<const uint32_t*>(seq_lens.data_ptr<int32_t>()),
reinterpret_cast<Metadata*>(metadata.data_ptr<int32_t>()),
batch_size,
static_cast<uint32_t>(static_cluster_threshold)) == cudaSuccess,
"fast_topk_v2_plan launch failed: ",
cudaGetErrorString(cudaGetLastError()));
}
namespace vllm::dsv4_topk {
// Shared dispatch path for fast_topk_v2 and fast_topk_v2_raw. Templated on
// (K, kRawOutput). K is the top-k value (512 for V4-Flash, 1024 for V4-Pro);
// kRawOutput=false folds the page-table gather, kRawOutput=true emits raw
// row-local indices. The set of input tensors is the same modulo
// (page_table, page_size), which the caller has already validated.
template <uint32_t K, bool kRawOutput>
static void launch_dispatch(const TopKParams& params, uint32_t batch_size,
uint32_t max_seq_len, cudaStream_t stream) {
// Helper: build a cudaLaunchConfig with optional PDL + cluster attributes.
// The attribute storage must outlive cudaLaunchKernelEx (cfg.attrs points
// into it), so it lives in each call site below as a stack local.
auto make_cfg = [&](dim3 grid, dim3 block, size_t smem,
cudaLaunchAttribute* attrs, bool enable_cluster,
bool enable_pdl) {
cudaLaunchConfig_t cfg{};
cfg.gridDim = grid;
cfg.blockDim = block;
cfg.dynamicSmemBytes = static_cast<unsigned>(smem);
cfg.stream = stream;
int n = 0;
if (enable_pdl) {
attrs[n].id = cudaLaunchAttributeProgrammaticStreamSerialization;
attrs[n].val.programmaticStreamSerializationAllowed = 1;
++n;
}
if (enable_cluster) {
attrs[n].id = cudaLaunchAttributeClusterDimension;
attrs[n].val.clusterDim = {1, kClusterSize, 1};
++n;
}
cfg.numAttrs = n;
cfg.attrs = n ? attrs : nullptr;
return cfg;
};
auto check_launch = [](cudaError_t err) {
TORCH_CHECK(err == cudaSuccess,
"fast_topk_v2 launch failed: ", cudaGetErrorString(err));
};
constexpr size_t kS1 = kStage1SMEM<K>;
constexpr size_t kS2 = kStage2SMEM<K>;
if (max_seq_len <= Small<K>::kMax1PassLength) {
setup_kernel_smem_once<&topk_short_transform<K, kRawOutput>, kS2>();
cudaLaunchAttribute attrs[2];
auto cfg = make_cfg(dim3(batch_size), dim3(kBlockSize), kS2, attrs,
/*cluster=*/false, /*pdl=*/true);
check_launch(cudaLaunchKernelEx(
&cfg, topk_short_transform<K, kRawOutput>, params));
} else if (batch_size <= kNumClusters) {
constexpr size_t kFusedSMEM = kS1 > kS2 ? kS1 : kS2;
setup_kernel_smem_once<&topk_fused_transform<K, kRawOutput>, kFusedSMEM>();
cudaLaunchAttribute attrs[2];
auto cfg = make_cfg(dim3(batch_size, kClusterSize), dim3(kBlockSize),
kFusedSMEM, attrs, /*cluster=*/true, /*pdl=*/true);
check_launch(cudaLaunchKernelEx(
&cfg, topk_fused_transform<K, kRawOutput>, params));
} else {
const auto num_clusters = std::min<uint32_t>(batch_size, kNumClusters);
setup_kernel_smem_once<&topk_combine_preprocess<K, kRawOutput>, kS1>();
cudaLaunchAttribute attrs1[2];
auto cfg1 = make_cfg(dim3(num_clusters, kClusterSize), dim3(kBlockSize),
kS1, attrs1, /*cluster=*/true, /*pdl=*/true);
check_launch(cudaLaunchKernelEx(
&cfg1, topk_combine_preprocess<K, kRawOutput>, params));
setup_kernel_smem_once<&topk_combine_transform<K, kRawOutput>, kS2>();
cudaLaunchAttribute attrs2[2];
auto cfg2 = make_cfg(dim3(batch_size), dim3(kBlockSize), kS2, attrs2,
/*cluster=*/false, /*pdl=*/true);
check_launch(cudaLaunchKernelEx(
&cfg2, topk_combine_transform<K, kRawOutput>, params));
}
}
// Top-level K dispatcher: validate the runtime topk argument and route to
// the right template instantiation.
template <bool kRawOutput>
static void launch_dispatch_k(int64_t topk, const TopKParams& params,
uint32_t batch_size, uint32_t max_seq_len,
cudaStream_t stream) {
if (topk == 512) {
launch_dispatch<512, kRawOutput>(params, batch_size, max_seq_len, stream);
} else if (topk == 1024) {
launch_dispatch<1024, kRawOutput>(params, batch_size, max_seq_len, stream);
} else {
TORCH_CHECK(false,
"fast_topk_v2 supports topk in {512, 1024}, got ", topk);
}
}
} // namespace vllm::dsv4_topk
void fast_topk_v2(const torch::Tensor& scores, const torch::Tensor& seq_lens,
const torch::Tensor& page_table, torch::Tensor& page_indices,
int64_t page_size, const torch::Tensor& workspace,
const torch::Tensor& metadata, int64_t topk) {
using namespace vllm::dsv4_topk;
CHECK_CUDA(scores);
CHECK_CUDA(seq_lens);
CHECK_CUDA(page_table);
CHECK_CUDA(page_indices);
CHECK_CUDA(workspace);
CHECK_CUDA(metadata);
CHECK_DTYPE(scores, torch::kFloat32);
CHECK_DTYPE(seq_lens, torch::kInt32);
CHECK_DTYPE(page_table, torch::kInt32);
CHECK_DTYPE(page_indices, torch::kInt32);
CHECK_DTYPE(workspace, torch::kInt32);
CHECK_DTYPE(metadata, torch::kInt32);
TORCH_CHECK(scores.dim() == 2 && scores.stride(1) == 1,
"scores must be 2D with last stride 1");
TORCH_CHECK(seq_lens.dim() == 1 && seq_lens.is_contiguous());
TORCH_CHECK(page_table.dim() == 2 && page_table.stride(1) == 1,
"page_table must be 2D with last stride 1");
TORCH_CHECK(page_indices.dim() == 2 && page_indices.is_contiguous() &&
page_indices.size(1) == topk,
"page_indices must be (B, topk) contiguous");
// workspace size is K-independent (it stages cluster-path ties whose
// count is bounded by kMaxTies, not K), so this check uses any K.
TORCH_CHECK(workspace.dim() == 2 && workspace.stride(1) == 1 &&
workspace.size(1) == Large<512>::kWorkspaceInts,
"workspace must be (B, kWorkspaceInts) with last stride 1");
TORCH_CHECK(metadata.dim() == 2 && metadata.size(1) == 4 &&
metadata.is_contiguous(),
"metadata must be (B + 1, 4) contiguous");
const auto batch_size = static_cast<uint32_t>(scores.size(0));
TORCH_CHECK(seq_lens.size(0) == batch_size);
TORCH_CHECK(page_table.size(0) == batch_size);
TORCH_CHECK(page_indices.size(0) == batch_size);
TORCH_CHECK(workspace.size(0) == batch_size);
TORCH_CHECK(metadata.size(0) == batch_size + 1);
const auto max_seq_len = static_cast<uint32_t>(scores.size(1));
TORCH_CHECK(page_size > 0 && (page_size & (page_size - 1)) == 0,
"page_size must be a positive power of 2");
TORCH_CHECK(scores.stride(0) % 4 == 0,
"score stride must be a multiple of 4 (TMA 16-byte alignment)");
// page_bits = log2(page_size). __builtin_ctzll is a host-side compiler
// builtin available under C++17 (vLLM compiles host code with C++17).
const auto page_bits = static_cast<uint32_t>(
__builtin_ctzll(static_cast<unsigned long long>(page_size)));
TopKParams params{
.seq_lens =
reinterpret_cast<const uint32_t*>(seq_lens.data_ptr<int32_t>()),
.scores = scores.data_ptr<float>(),
.page_table = page_table.data_ptr<int32_t>(),
.page_indices = page_indices.data_ptr<int32_t>(),
.score_stride = scores.stride(0),
.page_table_stride = page_table.stride(0),
.workspace = reinterpret_cast<uint8_t*>(workspace.data_ptr<int32_t>()),
.metadata =
reinterpret_cast<const Metadata*>(metadata.data_ptr<int32_t>()),
.workspace_stride =
workspace.stride(0) * static_cast<int64_t>(sizeof(int32_t)),
.batch_size = batch_size,
.page_bits = page_bits,
};
launch_dispatch_k<false>(topk, params, batch_size, max_seq_len,
at::cuda::getCurrentCUDAStream().stream());
}
// Top-k only: skip the page-table gather and emit raw row-local indices.
// Same selection algorithm as fast_topk_v2; just doesn't touch a page
// table. Output semantics match torch.ops._C.persistent_topk and the V4
// indexer's existing topk_indices_buffer contract.
void fast_topk_v2_raw(const torch::Tensor& scores,
const torch::Tensor& seq_lens,
torch::Tensor& topk_indices,
const torch::Tensor& workspace,
const torch::Tensor& metadata,
int64_t topk) {
using namespace vllm::dsv4_topk;
CHECK_CUDA(scores);
CHECK_CUDA(seq_lens);
CHECK_CUDA(topk_indices);
CHECK_CUDA(workspace);
CHECK_CUDA(metadata);
CHECK_DTYPE(scores, torch::kFloat32);
CHECK_DTYPE(seq_lens, torch::kInt32);
CHECK_DTYPE(topk_indices, torch::kInt32);
CHECK_DTYPE(workspace, torch::kInt32);
CHECK_DTYPE(metadata, torch::kInt32);
TORCH_CHECK(scores.dim() == 2 && scores.stride(1) == 1,
"scores must be 2D with last stride 1");
TORCH_CHECK(seq_lens.dim() == 1 && seq_lens.is_contiguous());
TORCH_CHECK(topk_indices.dim() == 2 && topk_indices.is_contiguous() &&
topk_indices.size(1) == topk,
"topk_indices must be (B, topk) contiguous");
TORCH_CHECK(workspace.dim() == 2 && workspace.stride(1) == 1 &&
workspace.size(1) == Large<512>::kWorkspaceInts,
"workspace must be (B, kWorkspaceInts) with last stride 1");
TORCH_CHECK(metadata.dim() == 2 && metadata.size(1) == 4 &&
metadata.is_contiguous(),
"metadata must be (B + 1, 4) contiguous");
const auto batch_size = static_cast<uint32_t>(scores.size(0));
TORCH_CHECK(seq_lens.size(0) == batch_size);
TORCH_CHECK(topk_indices.size(0) == batch_size);
TORCH_CHECK(workspace.size(0) == batch_size);
TORCH_CHECK(metadata.size(0) == batch_size + 1);
const auto max_seq_len = static_cast<uint32_t>(scores.size(1));
TORCH_CHECK(scores.stride(0) % 4 == 0,
"score stride must be a multiple of 4 (TMA 16-byte alignment)");
// page_table / page_bits are unused on the raw path; passing nullptr/0 is
// safe because every kernel call site is gated by `if constexpr
// (kRawOutput)` so the page-table loads are eliminated at compile time.
TopKParams params{
.seq_lens =
reinterpret_cast<const uint32_t*>(seq_lens.data_ptr<int32_t>()),
.scores = scores.data_ptr<float>(),
.page_table = nullptr,
.page_indices = topk_indices.data_ptr<int32_t>(),
.score_stride = scores.stride(0),
.page_table_stride = 0,
.workspace = reinterpret_cast<uint8_t*>(workspace.data_ptr<int32_t>()),
.metadata =
reinterpret_cast<const Metadata*>(metadata.data_ptr<int32_t>()),
.workspace_stride =
workspace.stride(0) * static_cast<int64_t>(sizeof(int32_t)),
.batch_size = batch_size,
.page_bits = 0,
};
launch_dispatch_k<true>(topk, params, batch_size, max_seq_len,
at::cuda::getCurrentCUDAStream().stream());
}
int64_t fast_topk_v2_workspace_ints() {
// Workspace size is K-independent (kMaxTies, not K, drives it).
return static_cast<int64_t>(vllm::dsv4_topk::Large<512>::kWorkspaceInts);
}
// Register impls here (instead of in torch_bindings.cpp) so they only exist
// when CMake compiles this source — i.e., when the target build has a
// compatible Hopper / Blackwell-datacenter arch. On other configs the schema
// remains defined but a call surfaces a clear "no impl" runtime error.
TORCH_LIBRARY_IMPL_EXPAND(TORCH_EXTENSION_NAME, CUDA, m) {
m.impl("fast_topk_v2_plan", &fast_topk_v2_plan);
m.impl("fast_topk_v2", &fast_topk_v2);
m.impl("fast_topk_v2_raw", &fast_topk_v2_raw);
}
TORCH_LIBRARY_IMPL_EXPAND(TORCH_EXTENSION_NAME, CompositeExplicitAutograd, m) {
m.impl("fast_topk_v2_workspace_ints", &fast_topk_v2_workspace_ints);
}
-266
View File
@@ -1,266 +0,0 @@
// SPDX-License-Identifier: Apache-2.0
// SPDX-FileCopyrightText: Copyright contributors to the vLLM project
//
// Cluster top-k strategy for very large N. Uses Hopper thread-block clusters
// (cooperative_groups::this_cluster) to parallelize histogram + scatter across
// up to ``kClusterSize`` blocks per row. Each row is processed in two stages:
// stage 1: per-block histogram, all-reduce across the cluster, threshold
// scatter, and an epilogue that page-translates strictly-above
// entries to global memory and stages ties into a per-row workspace.
// stage 2: tie-break across the cluster's combined ties (run by cluster
// rank 0 in the fused kernel, or as a separate launch otherwise).
// Ported from
// jit_kernel/include/sgl_kernel/deepseek_v4/topk/cluster.cuh.
#pragma once
#include "common.cuh"
#include "ptx.cuh"
#include "utils.cuh"
#include <cooperative_groups.h>
#include <cstdint>
namespace vllm::dsv4_topk {
template <uint32_t K>
struct ClusterTopK {
static constexpr uint32_t kClusterSize = 8;
static constexpr uint32_t kHistBits = 10;
static constexpr uint32_t kHistBins = 1 << kHistBits;
static constexpr uint32_t kElemPerStage = 8;
static constexpr uint32_t kSizePerStage = kElemPerStage * kBlockSize;
static constexpr uint32_t kNumStages = 4;
static constexpr uint32_t kMaxLength = kClusterSize * kNumStages * kSizePerStage;
static constexpr uint32_t kAboveBits = 11;
struct Smem {
uint64_t barrier[kNumStages];
uint32_t local_above_equal[kClusterSize];
uint32_t prefix_above_equal;
alignas(128) uint32_t counter_gt;
alignas(128) uint32_t counter_eq;
alignas(128) MatchBin match;
alignas(128) uint32_t warp_sum[kNumWarps];
uint32_t histogram[kHistBins];
alignas(128) float score_buffer[kNumStages][kSizePerStage];
Tie tie_buffer[kMaxTies];
};
// Per-row metadata produced by the plan kernel and consumed by the fused /
// stage-1 kernels. {batch_id, seq_len, has_next} arranged in an int4-sized
// 16-byte struct so the planner can do contiguous int32x4 stores.
struct alignas(16) Metadata {
uint32_t batch_id;
uint32_t seq_len;
bool has_next;
};
// Per-row workspace storing {(num_above, num_ties)} + the gathered ties.
struct WorkSpace {
uint2 metadata;
Tie ties[kMaxTies];
};
static constexpr uint32_t kWorkspaceInts = sizeof(WorkSpace) / sizeof(uint32_t);
VLLM_DSV4_DEVICE static void stage1_init(void* _smem) {
const auto tx = threadIdx.x;
__builtin_assume(tx < kBlockSize);
const auto smem = static_cast<Smem*>(_smem);
if (tx < kHistBins) smem->histogram[tx] = 0;
if (tx < kNumStages) ptx::mbarrier_init(&smem->barrier[tx], 1);
__syncthreads();
}
VLLM_DSV4_DEVICE static void stage1_prologue(const float* scores,
uint32_t length, void* _smem) {
if (threadIdx.x == 0) {
const auto smem = static_cast<Smem*>(_smem);
const auto num_stages = (length + kSizePerStage - 1) / kSizePerStage;
const auto length_aligned = (length + 3u) & ~3u;
#pragma unroll
for (uint32_t stage = 0; stage < kNumStages; stage++) {
if (stage >= num_stages) break;
const auto offset = stage * kSizePerStage;
const auto size = min(kSizePerStage, length_aligned - offset);
const auto size_bytes = size * sizeof(float);
const auto bar = &smem->barrier[stage];
ptx::tma_load(smem->score_buffer[stage], scores + offset, size_bytes,
bar);
ptx::mbarrier_arrive_expect_tx(bar, size_bytes);
}
}
}
VLLM_DSV4_DEVICE static void stage1(int32_t* indices, uint32_t length,
void* _smem, bool reuse = false) {
const auto smem = static_cast<Smem*>(_smem);
const auto tx = threadIdx.x;
__builtin_assume(tx < kBlockSize);
const auto lane_id = tx % kWarpThreads;
const auto warp_id = tx / kWarpThreads;
// Local histogram.
#pragma unroll
for (uint32_t stage = 0; stage < kNumStages; stage++) {
const auto offset = stage * kSizePerStage;
if (offset >= length) break;
const auto size = min(kSizePerStage, length - offset);
if (lane_id == 0) ptx::mbarrier_wait(&smem->barrier[stage], 0);
__syncwarp();
#pragma unroll
for (uint32_t i = 0; i < kElemPerStage; ++i) {
const auto idx = tx + i * kBlockSize;
if (idx >= size) break;
const auto score = smem->score_buffer[stage][idx];
const auto bin = extract_coarse_bin<kHistBits>(score);
atomicAdd(&smem->histogram[bin], 1);
}
}
static_assert(kHistBins <= kBlockSize);
// Two-shot all-reduce across the cluster.
{
auto cluster = cooperative_groups::this_cluster();
cluster.sync();
const auto cluster_rank = blockIdx.y;
const auto kLocalSize = kHistBins / kClusterSize;
const auto offset = kLocalSize * cluster_rank;
const auto src_tx = tx / kClusterSize;
const auto src_rank = tx % kClusterSize;
if (tx < kHistBins) {
const auto addr = &smem->histogram[offset + src_tx];
const auto src_addr = cluster.map_shared_rank(addr, src_rank);
*src_addr = warp_reduce_sum<kClusterSize>(*src_addr);
}
cluster.sync();
}
// Each block now holds the full cluster histogram. Find the threshold.
{
const auto value = tx < kHistBins ? smem->histogram[tx] : 0;
const auto warp_inc = warp_inclusive_sum(lane_id, value);
if (lane_id == kWarpThreads - 1) {
smem->warp_sum[warp_id] = warp_inc;
}
__syncthreads();
const auto tmp = smem->warp_sum[lane_id];
const auto total_length = warp_reduce_sum(tmp);
uint32_t prefix_sum = warp_reduce_sum(lane_id < warp_id ? tmp : 0);
prefix_sum += warp_inc;
const auto above = total_length - prefix_sum;
if (tx < kHistBins && above < K && above + value >= K) {
smem->counter_gt = smem->counter_eq = 0;
smem->match = {
.bin = tx,
.above_count = above,
.equal_count = value,
};
}
__syncthreads();
}
const auto thr_bin = smem->match.bin;
// Scatter strictly-above entries to `indices`, stash ties in tie_buffer.
#pragma unroll
for (uint32_t stage = 0; stage < kNumStages; stage++) {
const auto offset = stage * kSizePerStage;
if (offset >= length) break;
#pragma unroll
for (uint32_t i = 0; i < kElemPerStage; ++i) {
const auto buf_idx = tx + i * kBlockSize;
const auto global_idx = offset + buf_idx;
if (global_idx >= length) break;
const auto score = smem->score_buffer[stage][buf_idx];
const auto bin = extract_coarse_bin<kHistBits>(score);
if (bin > thr_bin) {
indices[atomicAdd(&smem->counter_gt, 1)] = global_idx;
} else if (bin == thr_bin) {
const auto pos = atomicAdd(&smem->counter_eq, 1);
if (pos < kMaxTies) smem->tie_buffer[pos] = {global_idx, score};
}
}
}
if (reuse) {
const auto num_stages = (length + kSizePerStage - 1) / kSizePerStage;
if (tx < kHistBins) smem->histogram[tx] = 0;
if (tx < num_stages) ptx::mbarrier_arrive(&smem->barrier[tx]);
}
__syncthreads();
}
template <typename TParams>
VLLM_DSV4_DEVICE static void stage1_epilogue(TParams params,
uint32_t offset, void* _ws,
void* _smem) {
auto cluster = cooperative_groups::this_cluster();
const auto smem = static_cast<Smem*>(_smem);
const auto tx = threadIdx.x;
const auto local_above = smem->counter_gt;
const auto local_equal = smem->counter_eq;
const auto cluster_rank = blockIdx.y;
constexpr uint32_t kAboveMask = (1 << kAboveBits) - 1;
static_assert(kAboveMask >= K);
static_assert(kMaxTies <= kBlockSize);
const auto idx_above = tx < local_above ? params.indices_in[tx] : 0;
const auto tie_value = tx < local_equal ? smem->tie_buffer[tx] : Tie{0, 0.0f};
// Push counts to remote shared memory to reduce inter-block latency.
if (tx < kClusterSize) {
const auto value = (local_equal << kAboveBits) | local_above;
const auto dst_addr = cluster.map_shared_rank(smem->local_above_equal, tx);
dst_addr[cluster_rank] = value;
}
// After this final sync, every block can read only its own smem (peer
// ranks may have already exited), so we don't touch remote smem again.
cluster.sync();
if (tx < kClusterSize) {
const auto value = tx < cluster_rank ? smem->local_above_equal[tx] : 0;
const auto kActiveMask = (1u << kClusterSize) - 1;
smem->prefix_above_equal = warp_reduce_sum<kClusterSize>(value, kActiveMask);
}
__syncthreads();
const auto prefix_packed = smem->prefix_above_equal;
const auto prefix_above = prefix_packed & kAboveMask;
const auto prefix_equal = prefix_packed >> kAboveBits;
// Page-translate strictly-above entries.
if (tx < local_above) {
params.write(tx + prefix_above, idx_above + offset);
}
// Stage ties into the per-row workspace (regular global writes).
const auto ws = static_cast<WorkSpace*>(_ws);
if (tx < local_equal && tx + prefix_equal < kMaxTies) {
ws->ties[tx + prefix_equal] = {tie_value.idx + offset, tie_value.score};
}
// Last cluster rank publishes the sums into ws->metadata.
if (cluster_rank == kClusterSize - 1 && tx == 0) {
const auto sum_above = prefix_above + local_above;
const auto sum_equal = prefix_equal + local_equal;
ws->metadata = make_uint2(sum_above, sum_equal);
}
}
template <typename TParams>
VLLM_DSV4_DEVICE static void transform(TParams params, const void* _ws,
void* _smem) {
const auto ws = static_cast<const WorkSpace*>(_ws);
const auto meta = &ws->metadata;
const auto num_above = meta->x;
const auto num_equal = meta->y;
if (num_above >= K || num_equal == 0) return;
const auto clamped_ties = min(num_equal, kMaxTies);
tie_handle_transform(ws->ties, clamped_ties, num_above, K, params, _smem);
}
};
} // namespace vllm::dsv4_topk
-219
View File
@@ -1,219 +0,0 @@
// SPDX-License-Identifier: Apache-2.0
// SPDX-FileCopyrightText: Copyright contributors to the vLLM project
//
// Shared types/utilities for the three DeepSeek V4 top-k strategies
// (Register / Streaming / Cluster). Ported from sglang's
// jit_kernel/include/sgl_kernel/deepseek_v4/topk/common.cuh.
#pragma once
#include "utils.cuh"
#include <cuda_fp16.h>
#include <cstdint>
namespace vllm::dsv4_topk {
inline constexpr uint32_t kMaxTopK = 1024;
inline constexpr uint32_t kBlockSize = 1024;
inline constexpr uint32_t kNumWarps = kBlockSize / kWarpThreads;
// 1 element per thread in the tie-breaking pass.
inline constexpr uint32_t kMaxTies = 1024;
inline constexpr uint32_t kRadixBins = 256;
static_assert(kMaxTopK <= kBlockSize && kMaxTies <= kBlockSize);
// Always vectorize global loads as float4.
using Vec4 = AlignedVector<float, 4>;
// page_to_indices: convert a flat compressed-token index into a (block * page_size + offset)
// page-table-resolved index. page_size must be a power of 2; page_bits = log2(page_size).
VLLM_DSV4_DEVICE int32_t page_to_indices(const int32_t* __restrict__ page_table,
uint32_t i, uint32_t page_bits) {
const uint32_t mask = (1u << page_bits) - 1u;
return (page_table[i >> page_bits] << page_bits) | (i & mask);
}
// Output-side description of how each strategy commits its top-k output.
//
// Two modes, picked at compile time via ``kRawOutput``:
// - kRawOutput=false (paged): fold the page-table gather into the output
// store. ``write(dst, src)`` emits ``page_to_indices(table, src, bits)``;
// ``transform(idx)`` reads ``indices_in[idx]`` and re-emits via the
// page lookup. This is the original kernel behavior.
// - kRawOutput=true (raw): skip the page lookup entirely. The kernel
// just writes row-local raw indices, matching ``persistent_topk``'s
// output contract. ``page_table`` and ``page_bits`` are unused; the
// compiler eliminates the dead loads via ``if constexpr``.
template <bool kRawOutput>
struct TransformParamsT {
const int32_t* __restrict__ page_table;
const int32_t* __restrict__ indices_in;
int32_t* __restrict__ indices_out;
uint32_t page_bits;
VLLM_DSV4_DEVICE void transform(uint32_t idx) const {
if constexpr (kRawOutput) {
indices_out[idx] = static_cast<int32_t>(indices_in[idx]);
} else {
indices_out[idx] =
page_to_indices(page_table, indices_in[idx], page_bits);
}
}
VLLM_DSV4_DEVICE void write(uint32_t dst, uint32_t src) const {
if constexpr (kRawOutput) {
indices_out[dst] = static_cast<int32_t>(src);
} else {
indices_out[dst] = page_to_indices(page_table, src, page_bits);
}
}
};
// Back-compat alias. The four kernels in fast_topk_v2.cu instantiate both
// variants explicitly via templates.
using TransformParams = TransformParamsT<false>;
struct alignas(16) MatchBin {
uint32_t bin;
uint32_t above_count;
uint32_t equal_count;
};
struct alignas(8) Tie {
uint32_t idx;
float score;
};
// Shared-memory layout for the final tie-breaking radix pass. Reused by both
// the streaming kernel (overlapping `score_buffer`) and the cluster kernel.
struct TieHandleSmem {
alignas(128) uint32_t counter;
alignas(128) MatchBin match;
uint32_t histogram[kRadixBins];
uint32_t warp_sum[kNumWarps];
};
// Order-preserving fp32 -> uint key, truncated to the top kBits. Used for the
// coarse histogram pass.
template <uint32_t kBits>
VLLM_DSV4_DEVICE uint32_t extract_coarse_bin(float x) {
static_assert(0 < kBits && kBits < 15);
__half h = __float2half_rn(x);
uint16_t bits = __half_as_ushort(h);
uint16_t key = (bits & 0x8000) ? static_cast<uint16_t>(~bits)
: static_cast<uint16_t>(bits | 0x8000);
return key >> (16 - kBits);
}
// Full 32-bit order-preserving key, used in tie-breaking.
VLLM_DSV4_DEVICE uint32_t extract_exact_bin(float x) {
uint32_t bits = __float_as_uint(x);
return (bits & 0x80000000u) ? ~bits : (bits | 0x80000000u);
}
VLLM_DSV4_DEVICE uint32_t warp_inclusive_sum(uint32_t lane_id, uint32_t val) {
static_assert(kWarpThreads == 32);
#pragma unroll
for (uint32_t offset = 1; offset < 32; offset *= 2) {
uint32_t n = __shfl_up_sync(0xFFFFFFFF, val, offset);
if (lane_id >= offset) val += n;
}
return val;
}
// Fast path when seq_len <= K: identity mapping, padded to K with -1.
template <typename TParams>
VLLM_DSV4_DEVICE void trivial_transform(const TParams& params, uint32_t length,
uint32_t K) {
const auto tx = threadIdx.x;
if (tx < length) {
params.write(tx, tx);
} else if (tx < K) {
params.indices_out[tx] = -1;
}
}
// Tie-break the threshold-bin candidates that didn't fit in the strict-above
// region. One block-wide radix pass over the full 32-bit key (fp32 bit
// pattern, with idx as a secondary key). Writes at most `K - num_above`
// entries via params.write(...).
template <typename TParams>
VLLM_DSV4_DEVICE void tie_handle_transform(const Tie* __restrict__ ties,
uint32_t num_ties, uint32_t num_above,
uint32_t K, TParams params,
void* _smem) {
auto* smem = static_cast<TieHandleSmem*>(_smem);
const auto tx = threadIdx.x;
const auto lane_id = tx % kWarpThreads;
const auto warp_id = tx / kWarpThreads;
const bool has_elem = tx < num_ties;
const auto tie = has_elem ? ties[tx] : Tie{0, 0.0f};
const uint32_t key = extract_exact_bin(tie.score);
const uint32_t idx = tie.idx;
bool active = has_elem;
uint32_t topk_remain = K - num_above;
uint32_t write_pos = K;
smem->counter = 0;
__syncthreads();
// 256 bins / 32 lanes = 8 warps span the histogram inter-warp prefix.
constexpr uint32_t kRadixWarps = kRadixBins / kWarpThreads;
#pragma unroll
for (int round = 0; round < 4; round++) {
const uint32_t shift = 24 - round * 8;
const uint32_t bin = (key >> shift) & 0xFFu;
// 1. Histogram.
if (tx < kRadixBins) smem->histogram[tx] = 0;
__syncthreads();
if (active) atomicAdd(&smem->histogram[bin], 1);
__syncthreads();
// 2. Two-pass prefix sum across the 256 bins.
uint32_t hist_val = 0;
uint32_t warp_inc = 0;
if (tx < kRadixBins) {
hist_val = smem->histogram[tx];
warp_inc = warp_inclusive_sum(lane_id, hist_val);
if (lane_id == kWarpThreads - 1) smem->warp_sum[warp_id] = warp_inc;
}
__syncthreads();
if (tx < kRadixBins) {
const auto tmp = (lane_id < kRadixWarps) ? smem->warp_sum[lane_id] : 0;
const auto total = warp_reduce_sum(tmp);
const auto inter = warp_reduce_sum(lane_id < warp_id ? tmp : 0);
const auto prefix = inter + warp_inc;
const auto above = total - prefix;
// 3. Find threshold bin.
if (above < topk_remain && above + hist_val >= topk_remain) {
smem->match = {tx, above, topk_remain - above};
}
}
__syncthreads();
const auto thr = smem->match.bin;
const auto n_above = smem->match.above_count;
// 4. Scatter.
if (active) {
if (bin > thr) {
write_pos = num_above + atomicAdd(&smem->counter, 1);
active = false;
} else if (bin < thr) {
active = false;
} else if (round == 3) {
write_pos = K - atomicAdd(&smem->match.equal_count, -1u);
}
// bin == thr && round < 3: stay active for the next radix round.
}
topk_remain -= n_above;
if (topk_remain == 0) break;
}
if (write_pos < K) params.write(write_pos, idx);
}
} // namespace vllm::dsv4_topk
-66
View File
@@ -1,66 +0,0 @@
// SPDX-License-Identifier: Apache-2.0
// SPDX-FileCopyrightText: Copyright contributors to the vLLM project
//
// Thin wrappers around the CUDA PTX intrinsics used by the top-k pipeline.
// All of these require sm_90+. Ported from sglang's
// jit_kernel/include/sgl_kernel/deepseek_v4/topk/ptx.cuh.
#pragma once
#include "utils.cuh"
#include <cuda/ptx>
#include <cstdint>
namespace vllm::dsv4_topk::ptx {
VLLM_DSV4_DEVICE void mbarrier_init(uint64_t* addr, uint32_t arrives) {
cuda::ptx::mbarrier_init(addr, arrives);
}
VLLM_DSV4_DEVICE void mbarrier_arrive(uint64_t* addr) {
cuda::ptx::mbarrier_arrive(cuda::ptx::sem_relaxed, cuda::ptx::scope_cta,
cuda::ptx::space_shared, addr);
}
VLLM_DSV4_DEVICE void mbarrier_arrive_expect_tx(uint64_t* addr, uint32_t tx) {
cuda::ptx::mbarrier_arrive_expect_tx(cuda::ptx::sem_relaxed,
cuda::ptx::scope_cta,
cuda::ptx::space_shared, addr, tx);
}
VLLM_DSV4_DEVICE void mbarrier_wait(uint64_t* addr, uint32_t phase) {
while (!cuda::ptx::mbarrier_try_wait_parity(cuda::ptx::sem_relaxed,
cuda::ptx::scope_cta, addr,
phase))
;
}
VLLM_DSV4_DEVICE void tma_load(void* dst, const void* src, uint32_t num_bytes,
uint64_t* mbar) {
cuda::ptx::cp_async_bulk(cuda::ptx::space_shared, cuda::ptx::space_global,
dst, src, num_bytes, mbar);
}
// elect.sync: pick a single arbitrary thread out of an active mask. Used to
// fire a single TMA load per warp without the full ``if (tx == 0)`` cost.
VLLM_DSV4_DEVICE uint32_t elect_sync() {
uint32_t pred = 0;
asm volatile(
"{\n\t"
".reg .pred %%px;\n\t"
"elect.sync _|%%px, %1;\n\t"
"@%%px mov.s32 %0, 1;\n\t"
"}"
: "+r"(pred)
: "r"(0xFFFFFFFF));
return pred;
}
VLLM_DSV4_DEVICE bool elect_sync_cta(uint32_t tx) {
const auto warp_id = tx / 32;
const auto uniform_warp_id = __shfl_sync(0xFFFFFFFF, warp_id, 0);
return (uniform_warp_id == 0 && elect_sync());
}
} // namespace vllm::dsv4_topk::ptx
-314
View File
@@ -1,314 +0,0 @@
// SPDX-License-Identifier: Apache-2.0
// SPDX-FileCopyrightText: Copyright contributors to the vLLM project
//
// Register-resident top-k strategy for the DeepSeek V4 indexer (small N
// fast path). One block per row; up to ``kMax2PassLength`` scores per row
// streamed through registers, with a single 12-bit-coarse radix pass and
// a final tie-break round. Ported from
// jit_kernel/include/sgl_kernel/deepseek_v4/topk/register.cuh.
#pragma once
#include "common.cuh"
#include "ptx.cuh"
#include "utils.cuh"
#include <cfloat>
#include <cstdint>
namespace vllm::dsv4_topk {
template <uint32_t K>
struct RegisterTopK {
static constexpr uint32_t kHistBits = 12;
static constexpr uint32_t kHistBins = 1 << kHistBits;
static constexpr uint32_t kVecsPerThread = 4;
static constexpr uint32_t kMaxTolerance = 0;
// Length covered by registers in a single pass.
static constexpr uint32_t kMax1PassLength = kVecsPerThread * 4 * kBlockSize;
// Extra length staged through shared memory in the 2-pass path.
static constexpr uint32_t kMaxExtraLength = kMax1PassLength;
static constexpr uint32_t kMax2PassLength = kMax1PassLength + kMaxExtraLength;
struct Smem {
using HistVec = AlignedVector<uint32_t, kHistBins / kBlockSize>;
alignas(128) uint32_t counter_gt;
alignas(128) uint32_t counter_eq;
uint64_t mbarrier; // for the cp.async.bulk in the 2-pass path
MatchBin match;
uint32_t warp_sum[kNumWarps];
union {
uint32_t histogram[kHistBins];
HistVec histogram_vec[kBlockSize];
Tie tie_buffer[kMaxTies];
};
alignas(16) float score_buffer[kMaxExtraLength];
};
template <bool kIs2Pass = false>
VLLM_DSV4_DEVICE static void run(const float* scores, int32_t* indices,
uint32_t length, void* _smem,
bool use_pdl = false) {
const auto smem = static_cast<Smem*>(_smem);
const auto tx = threadIdx.x;
const auto lane_id = tx % kWarpThreads;
const auto warp_id = tx / kWarpThreads;
// Init histogram + counters.
{
typename Smem::HistVec hist_vec;
hist_vec.fill(0);
smem->histogram_vec[tx] = hist_vec;
if (tx == 0) {
smem->counter_gt = smem->counter_eq = 0;
if constexpr (kIs2Pass) {
ptx::mbarrier_init(&smem->mbarrier, 1);
}
}
__syncthreads();
}
if (use_pdl) pdl_wait_primary<true>();
// Stream the first `kMax1PassLength` scores into registers.
Vec4 local[kVecsPerThread];
#pragma unroll
for (uint32_t v = 0; v < kVecsPerThread; ++v) {
const uint32_t base = (tx + v * kBlockSize) * 4;
if (base >= length) break;
local[v].load(scores, tx + v * kBlockSize);
}
// Issue the 2-pass TMA prefetch (next chunk of scores into smem).
if constexpr (kIs2Pass) {
if (ptx::elect_sync_cta(tx)) {
const auto length_aligned = (length + 3u - kMax1PassLength) & ~3u;
const auto size_bytes = length_aligned * sizeof(float);
ptx::tma_load(smem->score_buffer, scores + kMax1PassLength, size_bytes,
&smem->mbarrier);
ptx::mbarrier_arrive_expect_tx(&smem->mbarrier, size_bytes);
}
__syncwarp();
}
// Phase 1: histogram via shared-memory atomics.
#pragma unroll
for (uint32_t v = 0; v < kVecsPerThread; ++v) {
#pragma unroll
for (uint32_t e = 0; e < 4; ++e) {
if constexpr (!kIs2Pass) {
const uint32_t idx = (tx + v * kBlockSize) * 4 + e;
if (idx >= length) goto LABEL_ACC_FINISH;
}
atomicAdd(&smem->histogram[extract_coarse_bin<kHistBits>(local[v][e])],
1);
}
}
if constexpr (kIs2Pass) {
if (lane_id == 0) ptx::mbarrier_wait(&smem->mbarrier, 0);
__syncwarp();
for (uint32_t i = tx; i + kMax1PassLength < length; i += kBlockSize) {
const auto val = smem->score_buffer[i];
atomicAdd(&smem->histogram[extract_coarse_bin<kHistBits>(val)], 1);
}
}
[[maybe_unused]] LABEL_ACC_FINISH:
__syncthreads();
// Phase 2: prefix scan over the histogram, locate the threshold bin.
{
constexpr uint32_t kItems = kHistBins / kBlockSize;
uint32_t orig[kItems];
const auto hist_vec = smem->histogram_vec[tx];
uint32_t tmp_local_sum = 0;
#pragma unroll
for (uint32_t i = 0; i < kItems; ++i) {
orig[i] = hist_vec[i];
tmp_local_sum += orig[i];
}
const auto warp_inc = warp_inclusive_sum(lane_id, tmp_local_sum);
const auto warp_exc = warp_inc - tmp_local_sum;
if (lane_id == kWarpThreads - 1) {
smem->warp_sum[warp_id] = warp_inc;
}
__syncthreads();
const auto tmp = smem->warp_sum[lane_id];
// Exactly one bin satisfies above < K && above + count >= K.
uint32_t prefix_sum = warp_reduce_sum(lane_id < warp_id ? tmp : 0);
prefix_sum += warp_exc;
#pragma unroll
for (uint32_t i = 0; i < kItems; ++i) {
prefix_sum += orig[i];
const auto above = length - prefix_sum;
if (above < K && above + orig[i] >= K) {
smem->match = {
.bin = tx * kItems + i,
.above_count = above,
.equal_count = orig[i],
};
}
}
__syncthreads();
}
const auto thr_bin = smem->match.bin;
const auto num_above = smem->match.above_count;
const auto num_equal = smem->match.equal_count;
// Phase 3: Scatter.
// - bin > thr -> write directly to output (strictly above).
// - bin == thr -> when no tie-break is needed, admit first-come;
// otherwise stash into tie_buffer for phase 4.
const bool need_tiebreak = (num_equal + num_above > K + kMaxTolerance);
const auto topk_indices = indices;
const auto tie_buffer = smem->tie_buffer;
#pragma unroll
for (uint32_t v = 0; v < kVecsPerThread; ++v) {
#pragma unroll
for (uint32_t e = 0; e < 4; ++e) {
const uint32_t idx = (tx + v * kBlockSize) * 4 + e;
if constexpr (!kIs2Pass) {
if (idx >= length) goto LABEL_SCATTER_DONE;
}
const uint32_t bin = extract_coarse_bin<kHistBits>(local[v][e]);
if (bin > thr_bin) {
topk_indices[atomicAdd(&smem->counter_gt, 1)] = idx;
} else if (bin == thr_bin) {
const auto pos = atomicAdd(&smem->counter_eq, 1);
if (need_tiebreak) {
if (pos < kMaxTies) {
tie_buffer[pos] = {.idx = idx, .score = local[v][e]};
}
} else {
if (const auto which = pos + num_above; which < K) {
topk_indices[which] = idx;
}
}
}
}
// 2-pass: pull the next chunk in from the staged smem buffer.
if constexpr (kIs2Pass) {
local[v].load(smem->score_buffer, tx + v * kBlockSize);
}
}
if constexpr (kIs2Pass) {
#pragma unroll
for (uint32_t v = 0; v < kVecsPerThread; ++v) {
#pragma unroll
for (uint32_t e = 0; e < 4; ++e) {
const uint32_t idx =
(tx + v * kBlockSize) * 4 + e + kMax1PassLength;
if (idx >= length) goto LABEL_SCATTER_DONE;
const uint32_t bin = extract_coarse_bin<kHistBits>(local[v][e]);
if (bin > thr_bin) {
topk_indices[atomicAdd(&smem->counter_gt, 1)] = idx;
} else if (bin == thr_bin) {
const auto pos = atomicAdd(&smem->counter_eq, 1);
if (need_tiebreak) {
if (pos < kMaxTies) {
tie_buffer[pos] = {.idx = idx, .score = local[v][e]};
}
} else {
if (const auto which = pos + num_above; which < K) {
topk_indices[which] = idx;
}
}
}
}
}
}
[[maybe_unused]] LABEL_SCATTER_DONE:
if (!need_tiebreak) return;
// Phase 4: tie-break within the threshold bin. We assume num_ties <=
// kBlockSize (one block of ties), so each thread takes one tied element,
// counts the number of tied elements with strictly higher (score, -idx),
// and writes to output if its rank is below the remaining quota.
__syncthreads();
static_assert(kMaxTies <= kBlockSize);
const uint32_t num_ties = min(num_equal, kMaxTies);
const uint32_t topk_remain = K - num_above;
const auto is_greater = [](const Tie& a, const Tie& b) {
return (a.score > b.score) || (a.score == b.score && a.idx < b.idx);
};
if (num_ties <= kWarpThreads) {
static_assert(kWarpThreads <= kNumWarps);
if (lane_id >= num_ties || warp_id >= num_ties) return;
const uint32_t mask = (1ull << num_ties) - 1u;
const auto tie = tie_buffer[lane_id];
const auto target_tie = tie_buffer[warp_id];
const bool pred = is_greater(tie, target_tie);
const auto rank =
static_cast<uint32_t>(__popc(__ballot_sync(mask, pred)));
if (lane_id == 0 && rank < topk_remain) {
topk_indices[num_above + rank] = target_tie.idx;
}
} else if (num_ties <= kWarpThreads * 2) {
// 64x64 case: each thread takes 2 elements.
const auto lane_id_1 = lane_id + kWarpThreads;
const auto warp_id_1 = warp_id + kWarpThreads;
const auto invalid = Tie{.idx = 0xFFFFFFFFu, .score = -FLT_MAX};
const auto tie_0 = tie_buffer[lane_id];
const auto tie_1 = lane_id_1 < num_ties ? tie_buffer[lane_id_1] : invalid;
{
const auto target = tie_buffer[warp_id];
const bool pred_0 = is_greater(tie_0, target);
const bool pred_1 = is_greater(tie_1, target);
const auto rank_0 =
static_cast<uint32_t>(__popc(__ballot_sync(0xFFFFFFFF, pred_0)));
const auto rank_1 =
static_cast<uint32_t>(__popc(__ballot_sync(0xFFFFFFFF, pred_1)));
const auto rank = rank_0 + rank_1;
if (lane_id == 0 && rank < topk_remain) {
topk_indices[num_above + rank] = target.idx;
}
}
if (warp_id_1 < num_ties) {
const auto target = tie_buffer[warp_id_1];
const bool pred_0 = is_greater(tie_0, target);
const bool pred_1 = is_greater(tie_1, target);
const auto rank_0 =
static_cast<uint32_t>(__popc(__ballot_sync(0xFFFFFFFF, pred_0)));
const auto rank_1 =
static_cast<uint32_t>(__popc(__ballot_sync(0xFFFFFFFF, pred_1)));
const auto rank = rank_0 + rank_1;
if (lane_id == 0 && rank < topk_remain) {
topk_indices[num_above + rank] = target.idx;
}
}
} else {
[[unlikely]];
// Block-wide fallback. Rarely reached.
for (auto i = warp_id; i < num_ties; i += kNumWarps) {
const auto target_tie = tie_buffer[i];
uint32_t local_rank = 0;
for (auto j = lane_id; j < num_ties; j += kWarpThreads) {
const auto tie = tie_buffer[j];
if (is_greater(tie, target_tie)) local_rank++;
}
const auto rank = warp_reduce_sum(local_rank);
if (lane_id == 0 && rank < topk_remain) {
topk_indices[num_above + rank] = target_tie.idx;
}
}
}
}
template <typename TParams>
VLLM_DSV4_DEVICE static void transform(TParams params) {
__syncthreads();
if (const auto tx = threadIdx.x; tx < K) params.transform(tx);
}
};
} // namespace vllm::dsv4_topk
-209
View File
@@ -1,209 +0,0 @@
// SPDX-License-Identifier: Apache-2.0
// SPDX-FileCopyrightText: Copyright contributors to the vLLM project
//
// Streaming top-k strategy for medium N. Uses a TMA-driven double-buffered
// histogram pass + scatter pass over chunks of `kSizePerStage` floats.
// Ported from
// jit_kernel/include/sgl_kernel/deepseek_v4/topk/streaming.cuh.
#pragma once
#include "common.cuh"
#include "ptx.cuh"
#include "utils.cuh"
#include <cfloat>
#include <cstdint>
namespace vllm::dsv4_topk {
template <uint32_t K>
struct StreamingTopK {
static constexpr uint32_t kHistBits = 12;
static constexpr uint32_t kHistBins = 1 << kHistBits;
static constexpr uint32_t kElemPerStage = 8;
static constexpr uint32_t kSizePerStage = kElemPerStage * kBlockSize;
static constexpr uint32_t kNumStages = 2; // double buffer
static constexpr uint32_t kHistItems = kHistBins / kBlockSize; // 4
static_assert(kHistItems * kBlockSize == kHistBins);
using HistVec = AlignedVector<uint32_t, kHistItems>;
struct Smem {
// [phase = 0 (histogram) | 1 (scatter)] x [buffer = 0 | 1]
uint64_t barrier[2][kNumStages];
alignas(128) uint32_t counter_gt;
alignas(128) uint32_t counter_eq;
alignas(128) MatchBin match;
alignas(128) uint32_t warp_sum[kNumWarps];
union {
uint32_t histogram[kHistBins];
HistVec histogram_vec[kBlockSize];
Tie tie_buffer[kMaxTies];
};
union {
float score_buffer[kNumStages][kSizePerStage];
TieHandleSmem stage2; // reused for the tie-handling phase
};
};
// length must be 4-aligned (caller rounds up); TMA wants 16-byte alignment.
template <bool kIsScatter>
VLLM_DSV4_DEVICE static void issue_tma(const float* scores, uint32_t stage,
uint32_t length, Smem* smem) {
const auto buf_idx = stage % kNumStages;
const auto offset = stage * kSizePerStage;
const auto size = min(kSizePerStage, length - offset);
const auto size_bytes = size * sizeof(float);
const auto bar = &smem->barrier[kIsScatter][buf_idx];
ptx::tma_load(smem->score_buffer[buf_idx], scores + offset, size_bytes,
bar);
ptx::mbarrier_arrive_expect_tx(bar, size_bytes);
}
// Unified streaming pass. kIsScatter=false: build histogram (phase A).
// kIsScatter=true: scatter using the threshold bin (phase C). Each barrier
// is reused across iterations via the reuse-arrive pattern.
template <bool kIsScatter>
VLLM_DSV4_DEVICE static void stream_pass(const float* scores, uint32_t length,
uint32_t thr_bin,
int32_t* s_topk_indices,
Smem* smem) {
const auto tx = threadIdx.x;
const auto num_iters = (length + kSizePerStage - 1) / kSizePerStage;
const auto lane_id = tx % kWarpThreads;
const auto length_aligned = (length + 3u) & ~3u;
if (tx == 0) {
#pragma unroll
for (uint32_t i = 0; i < kNumStages; i++) {
if (i >= num_iters) break;
issue_tma<kIsScatter>(scores, i, length_aligned, smem);
}
}
for (uint32_t iter = 0; iter < num_iters; iter++) {
const auto buf_idx = iter % kNumStages;
const auto offset = iter * kSizePerStage;
const auto this_size = min(kSizePerStage, length - offset);
if (lane_id == 1) {
const auto phase_bit = (iter / kNumStages) & 1;
ptx::mbarrier_wait(&smem->barrier[kIsScatter][buf_idx], phase_bit);
}
__syncwarp();
#pragma unroll
for (uint32_t i = 0; i < kElemPerStage; i++) {
const auto local_idx = tx + i * kBlockSize;
if (local_idx >= this_size) break;
const auto score = smem->score_buffer[buf_idx][local_idx];
const auto bin = extract_coarse_bin<kHistBits>(score);
if constexpr (kIsScatter) {
const auto global_idx = offset + local_idx;
if (bin > thr_bin) {
const auto pos = atomicAdd(&smem->counter_gt, 1);
if (pos < K) s_topk_indices[pos] = global_idx;
} else if (bin == thr_bin) {
const auto pos = atomicAdd(&smem->counter_eq, 1);
if (pos < kMaxTies) smem->tie_buffer[pos] = {global_idx, score};
}
} else {
atomicAdd(&smem->histogram[bin], 1);
}
}
__syncthreads();
if (tx == 0) {
if (const auto next_iter = iter + kNumStages; next_iter < num_iters) {
issue_tma<kIsScatter>(scores, next_iter, length_aligned, smem);
}
}
}
}
// Phase B: locate threshold bin via warp-level prefix scan.
VLLM_DSV4_DEVICE static void find_threshold(uint32_t length, Smem* smem) {
const auto tx = threadIdx.x;
const auto lane_id = tx % kWarpThreads;
const auto warp_id = tx / kWarpThreads;
uint32_t orig[kHistItems];
const auto hist_vec = smem->histogram_vec[tx];
uint32_t local_sum = 0;
#pragma unroll
for (uint32_t i = 0; i < kHistItems; ++i) {
orig[i] = hist_vec[i];
local_sum += orig[i];
}
const auto warp_inc = warp_inclusive_sum(lane_id, local_sum);
const auto warp_exc = warp_inc - local_sum;
if (lane_id == kWarpThreads - 1) smem->warp_sum[warp_id] = warp_inc;
__syncthreads();
const auto tmp = smem->warp_sum[lane_id];
uint32_t prefix_sum = warp_reduce_sum(lane_id < warp_id ? tmp : 0);
prefix_sum += warp_exc;
#pragma unroll
for (uint32_t i = 0; i < kHistItems; ++i) {
prefix_sum += orig[i];
const auto above = length - prefix_sum;
if (above < K && above + orig[i] >= K) {
smem->match = {
.bin = tx * kHistItems + i,
.above_count = above,
.equal_count = orig[i],
};
}
}
__syncthreads();
}
VLLM_DSV4_DEVICE static void run(const float* scores, uint32_t length,
int32_t* topk_indices, void* _smem) {
const auto smem = static_cast<Smem*>(_smem);
const auto tx = threadIdx.x;
__builtin_assume(tx < kBlockSize);
{
HistVec zero;
zero.fill(0);
smem->histogram_vec[tx] = zero;
if (tx < 2 * kNumStages) {
const auto base_barrier = &smem->barrier[0][0];
ptx::mbarrier_init(&base_barrier[tx], 1);
}
if (tx == 0) {
smem->counter_gt = 0;
smem->counter_eq = 0;
}
__syncthreads();
}
// Phase A: histogram.
stream_pass<false>(scores, length, 0, nullptr, smem);
// Phase B: threshold bin.
find_threshold(length, smem);
// Phase C: scatter.
stream_pass<true>(scores, length, smem->match.bin, topk_indices, smem);
}
template <typename TParams>
VLLM_DSV4_DEVICE static void transform(TParams params, void* _smem) {
// Phase D: page-translate above entries, then refine ties.
const auto smem = static_cast<Smem*>(_smem);
const auto tx = threadIdx.x;
const auto num_above = smem->match.above_count;
if (tx < num_above) params.transform(tx);
const auto num_equal = smem->counter_eq;
if (num_above >= K || num_equal == 0) return;
const auto clamped_ties = min(num_equal, kMaxTies);
tie_handle_transform(smem->tie_buffer, clamped_ties, num_above, K, params,
&smem->stage2);
}
};
} // namespace vllm::dsv4_topk
-75
View File
@@ -1,75 +0,0 @@
// SPDX-License-Identifier: Apache-2.0
// SPDX-FileCopyrightText: Copyright contributors to the vLLM project
//
// Minimal device-side utilities used by the DeepSeek V4 indexer top-k port.
// Replaces sgl_kernel/{utils,warp,vec,type}.cuh — we only need the bits the
// top-k kernels actually touch.
#pragma once
#include <cstddef>
#include <cstdint>
#include <cuda_fp16.h>
#include <cuda_runtime.h>
namespace vllm::dsv4_topk {
#define VLLM_DSV4_DEVICE __forceinline__ __device__
inline constexpr uint32_t kWarpThreads = 32u;
inline constexpr uint32_t kFullMask = 0xffffffffu;
// Programmatic Dependent Launch (sm_90+). When enabled, the kernel waits for
// the predecessor on the same stream to advance past its dependents-launch
// trigger before doing anything memory-dependent. Used to overlap the
// fp8_paged_mqa_logits epilogue with the first stage of top-k.
template <bool kUsePDL>
VLLM_DSV4_DEVICE void pdl_wait_primary() {
if constexpr (kUsePDL) {
asm volatile("griddepcontrol.wait;" ::: "memory");
}
}
template <bool kUsePDL>
VLLM_DSV4_DEVICE void pdl_trigger_secondary() {
if constexpr (kUsePDL) {
asm volatile("griddepcontrol.launch_dependents;" :::);
}
}
// Warp-level XOR-shuffle reduce. kThreads must be a power of 2 and <= 32.
template <uint32_t kThreads = kWarpThreads, typename T>
VLLM_DSV4_DEVICE T warp_reduce_sum(T value, uint32_t active_mask = kFullMask) {
#pragma unroll
for (auto offset = kThreads >> 1; offset > 0; offset >>= 1) {
value = value + __shfl_xor_sync(active_mask, value, offset, 32);
}
return value;
}
// 128-bit-aligned vector of N elements of T (N must be a power of 2, total
// size <= 16 bytes). Used for vectorized loads/stores into shared memory.
template <typename T, std::size_t N>
struct alignas(sizeof(T) * N) AlignedVector {
static_assert(N > 0 && (N & (N - 1)) == 0, "N must be a power of two");
static_assert(sizeof(T) * N <= 16,
"AlignedVector exceeds the 128-bit CUDA vector limit");
T data[N];
VLLM_DSV4_DEVICE void load(const void* ptr, std::size_t offset = 0) {
*reinterpret_cast<AlignedVector*>(this) =
reinterpret_cast<const AlignedVector*>(ptr)[offset];
}
VLLM_DSV4_DEVICE void store(void* ptr, std::size_t offset = 0) const {
reinterpret_cast<AlignedVector*>(ptr)[offset] = *this;
}
VLLM_DSV4_DEVICE void fill(T value) {
#pragma unroll
for (std::size_t i = 0; i < N; ++i) data[i] = value;
}
VLLM_DSV4_DEVICE T& operator[](std::size_t i) { return data[i]; }
VLLM_DSV4_DEVICE const T& operator[](std::size_t i) const { return data[i]; }
};
} // namespace vllm::dsv4_topk
+6 -3
View File
@@ -137,15 +137,18 @@ fused_add_rms_norm_static_fp8_quant_kernel(
_f16Vec<scalar_t, width> res = residual_v[id];
_f16Vec<scalar_t, width> w = weight_v[idx];
using Converter = _typeConvert<scalar_t>;
using HipT = typename Converter::hip_type;
#pragma unroll
for (int i = 0; i < width; ++i) {
float x = Converter::convert(res.data[i]);
float wf = Converter::convert(w.data[i]);
// See note in rms_norm_static_fp8_quant_kernel: round through scalar_t
// to match the unfused composite path at FP8 boundaries.
scalar_t out_norm = Converter::convert(x * s_variance * wf);
// to match the unfused composite path at FP8 boundaries. We use the
// backend's hip_type for the intermediate since c10::Half/BFloat16 has
// ambiguous conversions on CUDA and no implicit conversion on ROCm.
HipT out_norm_h = Converter::convert(x * s_variance * wf);
out[id * width + i] = scaled_fp8_conversion<true, fp8_type>(
static_cast<float>(out_norm), scale_inv);
Converter::convert(out_norm_h), scale_inv);
}
}
}
+2 -34
View File
@@ -125,40 +125,6 @@ void persistent_topk(const torch::Tensor& logits, const torch::Tensor& lengths,
torch::Tensor& output, torch::Tensor& workspace, int64_t k,
int64_t max_seq_len);
// DeepSeek V4 indexer top-k (k = 512). Hopper (sm_90a) and Blackwell
// datacenter (sm_100/sm_103) — needs thread-block clusters, TMA, and PDL.
// Two-step API:
// 1. fast_topk_v2_plan inspects the seq_lens distribution and writes a
// cluster_threshold + per-row Metadata into a (B+1, 4) int32 tensor. The
// plan is amortized when cudagraph-captured: once per shape, reused across
// layers.
// 2. fast_topk_v2 selects the top-512 indices per row, folds the page-table
// gather into the radix store, and writes (B, 512) int32 page indices.
// Dispatches per row to one of three strategies (Register / Streaming /
// Cluster) using the planned threshold.
//
// Returns the size in int32s of the per-row workspace required by
// fast_topk_v2 (allocate `(B, fast_topk_v2_workspace_ints())` int32 contig).
void fast_topk_v2_plan(const torch::Tensor& seq_lens, torch::Tensor& metadata,
int64_t static_cluster_threshold);
void fast_topk_v2(const torch::Tensor& scores, const torch::Tensor& seq_lens,
const torch::Tensor& page_table, torch::Tensor& page_indices,
int64_t page_size, const torch::Tensor& workspace,
const torch::Tensor& metadata, int64_t topk);
// Top-k only, no page-table fold-in. Same selection as fast_topk_v2 but
// emits raw row-local indices into ``topk_indices`` (drop-in for
// persistent_topk's output contract). topk must be one of {512, 1024}.
void fast_topk_v2_raw(const torch::Tensor& scores,
const torch::Tensor& seq_lens,
torch::Tensor& topk_indices,
const torch::Tensor& workspace,
const torch::Tensor& metadata,
int64_t topk);
int64_t fast_topk_v2_workspace_ints();
void rms_norm_static_fp8_quant(torch::Tensor& out, torch::Tensor& input,
torch::Tensor& weight, torch::Tensor& scale,
double epsilon);
@@ -197,6 +163,8 @@ void rotary_embedding(torch::Tensor& positions, torch::Tensor& query,
void silu_and_mul(torch::Tensor& out, torch::Tensor& input);
void silu_and_mul_clamp(torch::Tensor& out, torch::Tensor& input, double limit);
void silu_and_mul_quant(torch::Tensor& out, torch::Tensor& input,
torch::Tensor& scale);
+6 -20
View File
@@ -106,6 +106,12 @@ TORCH_LIBRARY_EXPAND(TORCH_EXTENSION_NAME, ops) {
ops.def("silu_and_mul(Tensor! result, Tensor input) -> ()");
ops.impl("silu_and_mul", torch::kCUDA, &silu_and_mul);
// SwiGLU activation with input clamping.
ops.def(
"silu_and_mul_with_clamp(Tensor! result, Tensor input, float limit) "
"-> ()");
ops.impl("silu_and_mul_with_clamp", torch::kCUDA, &silu_and_mul_clamp);
ops.def(
"silu_and_mul_quant(Tensor! result, Tensor input, Tensor scale) -> ()");
ops.impl("silu_and_mul_quant", torch::kCUDA, &silu_and_mul_quant);
@@ -215,26 +221,6 @@ TORCH_LIBRARY_EXPAND(TORCH_EXTENSION_NAME, ops) {
"Tensor workspace, int k, int max_seq_len) -> ()");
ops.impl("persistent_topk", torch::kCUDA, &persistent_topk);
// DeepSeek V4 indexer top-k (k=512), ported from sglang's topk_v2 family.
// Built for sm_90a (Hopper) + sm_100a/sm_103 (Blackwell datacenter).
// Schema only here; impl is registered in csrc/deepseek_v4/fast_topk_v2.cu
// so it's only present when CMake compiles the source for a supported arch.
ops.def(
"fast_topk_v2_plan(Tensor seq_lens, Tensor! metadata, "
"int static_cluster_threshold) -> ()");
ops.def(
"fast_topk_v2(Tensor scores, Tensor seq_lens, Tensor page_table, "
"Tensor! page_indices, int page_size, Tensor workspace, "
"Tensor metadata, int topk) -> ()");
ops.def(
"fast_topk_v2_raw(Tensor scores, Tensor seq_lens, "
"Tensor! topk_indices, Tensor workspace, Tensor metadata, int topk)"
" -> ()");
ops.def("fast_topk_v2_workspace_ints() -> int");
// Layernorm-quant
// Apply Root Mean Square (RMS) Normalization to the input tensor.
ops.def(
+1 -1
View File
@@ -538,7 +538,7 @@ RUN CUDA_VERSION_DASH=$(echo $CUDA_VERSION | cut -d. -f1,2 | tr '.' '-') && \
cuda-nvrtc-${CUDA_VERSION_DASH} \
cuda-cuobjdump-${CUDA_VERSION_DASH} \
libcurand-dev-${CUDA_VERSION_DASH} \
libcublas-${CUDA_VERSION_DASH} \
libcublas-dev-${CUDA_VERSION_DASH} \
# Required by fastsafetensors (fixes #20384)
libnuma-dev && \
# Fixes nccl_allocator requiring nccl.h at runtime
+7 -7
View File
@@ -68,7 +68,7 @@ You can pass a single image to the `'image'` field of the multi-modal dictionary
print(generated_text)
```
Full example: [examples/offline_inference/vision_language.py](../../examples/offline_inference/vision_language.py)
Full example: [examples/generate/multimodal/vision_language_offline.py](../../examples/generate/multimodal/vision_language_offline.py)
To substitute multiple images inside the same text prompt, you can pass in a list of images instead:
@@ -101,7 +101,7 @@ To substitute multiple images inside the same text prompt, you can pass in a lis
print(generated_text)
```
Full example: [examples/offline_inference/vision_language_multi_image.py](../../examples/offline_inference/vision_language_multi_image.py)
Full example: [examples/generate/multimodal/vision_language_multi_image_offline.py](../../examples/generate/multimodal/vision_language_multi_image_offline.py)
If using the [LLM.chat](../models/generative_models.md#llmchat) method, you can pass images directly in the message content using various formats: image URLs, PIL Image objects, or pre-computed embeddings:
@@ -287,13 +287,13 @@ Instead of NumPy arrays, you can also pass `'torch.Tensor'` instances, as shown
!!! note
'process_vision_info' is only applicable to Qwen2.5-VL and similar models.
Full example: [examples/offline_inference/vision_language.py](../../examples/offline_inference/vision_language.py)
Full example: [examples/generate/multimodal/vision_language_offline.py](../../examples/generate/multimodal/vision_language_offline.py)
### Audio Inputs
You can pass a tuple `(array, sampling_rate)` to the `'audio'` field of the multi-modal dictionary.
Full example: [examples/offline_inference/audio_language.py](../../examples/offline_inference/audio_language.py)
Full example: [examples/generate/multimodal/audio_language_offline.py](../../examples/generate/multimodal/audio_language_offline.py)
#### Chunking Long Audio for Transcription
@@ -674,7 +674,7 @@ Then, you can use the OpenAI client as follows:
print("Chat completion output:", chat_response.choices[0].message.content)
```
Full example: [examples/online_serving/openai_chat_completion_client_for_multimodal.py](../../examples/online_serving/openai_chat_completion_client_for_multimodal.py)
Full example: [examples/generate/multimodal/openai_chat_completion_client_for_multimodal.py](../../examples/generate/multimodal/openai_chat_completion_client_for_multimodal.py)
!!! tip
Loading from local file paths is also supported on vLLM: You can specify the allowed local media path via `--allowed-local-media-path` when launching the API server/engine,
@@ -745,7 +745,7 @@ Then, you can use the OpenAI client as follows:
print("Chat completion output from image url:", result)
```
Full example: [examples/online_serving/openai_chat_completion_client_for_multimodal.py](../../examples/online_serving/openai_chat_completion_client_for_multimodal.py)
Full example: [examples/generate/multimodal/openai_chat_completion_client_for_multimodal.py](../../examples/generate/multimodal/openai_chat_completion_client_for_multimodal.py)
!!! note
By default, the timeout for fetching videos through HTTP URL is `30` seconds.
@@ -958,7 +958,7 @@ Alternatively, you can pass `audio_url`, which is the audio counterpart of `imag
print("Chat completion output from audio url:", result)
```
Full example: [examples/online_serving/openai_chat_completion_client_for_multimodal.py](../../examples/online_serving/openai_chat_completion_client_for_multimodal.py)
Full example: [examples/generate/multimodal/openai_chat_completion_client_for_multimodal.py](../../examples/generate/multimodal/openai_chat_completion_client_for_multimodal.py)
!!! note
By default, the timeout for fetching audios through HTTP URL is `10` seconds.
+1
View File
@@ -20,6 +20,7 @@ The following are the supported quantization formats for vLLM:
- [AMD Quark](quark.md)
- [Quantized KV Cache](quantized_kvcache.md)
- [TorchAO](torchao.md)
- [FP8 ViT Encoder Attention](fp8_vit_attn.md)
## Supported Hardware
+109
View File
@@ -0,0 +1,109 @@
# FP8 ViT Encoder Attention
For visual understanding workloads with large images (e.g. QHD, 4K) and relatively
short text prompts/generation, the ViT encoder attention can become a significant
bottleneck, especially when the text model is quantized (e.g. NVFP4). vLLM
supports optional FP8 quantization for the ViT encoder attention via the
FlashInfer cuDNN backend. Q/K/V are quantized on-the-fly to FP8 before the
cuDNN attention call.
!!! note
- Currently supports Qwen3-VL family models only (`qwen3_vl`, `qwen3_vl_moe`,
`qwen3_5`, `qwen3_5_moe`, and other models using Qwen3 ViT).
- Dynamic scaling is not compatible with ViT full CUDA graphs.
- Performance gains are mostly visible at QHD/4K resolutions or multi-image
requests. Smaller images may see no speedup due to quantization overhead
(3 quantization kernel launches + un-padding).
- FP8 tensor-core speedup is more pronounced on GB300 than GB200.
## Requirements
- FlashInfer cuDNN backend with cuDNN >= 9.17.1.
## Usage
Enable FP8 ViT attention by passing `--mm-encoder-attn-dtype fp8` together
with `--mm-encoder-attn-backend FLASHINFER`:
```bash
vllm serve $MODEL \
--mm-encoder-attn-backend FLASHINFER \
--mm-encoder-attn-dtype fp8
```
By default (no scale file), **dynamic scaling** is used: a 16-entry circular
buffer of observed Q/K/V amax values drives per-forward scale updates. This
matches BF16 accuracy without any calibration but adds a small per-forward
overhead.
## Calibrate-Once, Reuse Workflow (Recommended)
For production, calibrate static scales on a representative dataset once and
reuse them to avoid the dynamic overhead:
```bash
# Step 1: calibrate and save scales (runs dynamic scaling for 16 passes,
# then dumps the learned scales to JSON).
vllm bench mm-processor \
--model $MODEL --mm-encoder-attn-backend FLASHINFER \
--mm-encoder-attn-dtype fp8 \
--mm-encoder-fp8-scale-save-path /path/to/scales.json \
--dataset-name hf --dataset-path lmarena-ai/VisionArena-Chat \
--num-prompts 100
# Step 2: serve with static scales (no dynamic overhead).
vllm serve $MODEL \
--mm-encoder-attn-backend FLASHINFER \
--mm-encoder-attn-dtype fp8 \
--mm-encoder-fp8-scale-path /path/to/scales.json
```
Saved scales are multiplied by `--mm-encoder-fp8-scale-save-margin` (default
`1.5`) to leave headroom against activation outliers not present in the
calibration set. The default has been validated to generalize across datasets
(e.g. VisionArena-Chat calibration maintains BF16 accuracy on ChartQA).
## Scale File Format
```json
{
"visual.blocks.0.attn.attn": {"q": 224.0, "k": 198.0, "v": 210.0},
"visual.blocks.1.attn.attn": {"q": 218.0, "k": 195.0, "v": 207.0}
}
```
Keys `q_scale` / `k_scale` / `v_scale` are accepted as aliases.
## Performance
**Core cuDNN attention kernel** (PyTorch profiler, `cudnn_generated_fort_native_sdpa_sm100_flash_fprop`, head_dim=128, seq_len=8192):
| Hardware | BF16 | FP8 | Speedup |
| -------- | ---- | ---- | ------- |
| GB200 | 350 us | 312 us | **1.12x** |
| GB300 | 300 us | 211 us | **1.42x** |
**End-to-end encoder forward time** (Qwen3-VL-30B-A3B-Instruct on GB200, 3 images/request):
| Resolution | BF16 median | FP8 median | Speedup |
| ---------- | ----------- | ---------- | ------- |
| HD (720x1280) | 31.77 ms | 36.39 ms | 0.87x |
| FullHD (1080x1920) | 57.99 ms | 58.73 ms | ~same |
| QHD (1440x2560) | 131.83 ms | 122.30 ms | **1.08x** |
| 4K (2160x3840) | 543.44 ms | 460.31 ms | **1.18x** |
Crossover is around FullHD with 3 images/request. At QHD and above, FP8 wins.
## Accuracy
ChartQA, Qwen3-VL-8B-Instruct, 500 samples. FP8 static uses scales calibrated
on VisionArena-Chat (with default 1.5x margin):
| Metric | BF16 | FP8 dynamic | FP8 static |
| ------ | ---- | ----------- | ---------- |
| relaxed_accuracy | 0.780 | 0.776 | 0.780 |
| anywhere_accuracy | 0.806 | 0.816 | 0.814 |
| exact_match | 0.584 | 0.582 | 0.578 |
All three configurations match within statistical noise, confirming that
static scales calibrated on one dataset generalize to another.
+1 -1
View File
@@ -202,7 +202,7 @@ The reasoning content is also available when both tool calling and the reasoning
print(f"Arguments: {tool_call.arguments}")
```
For more examples, please refer to [examples/online_serving/openai_chat_completion_tool_calls_with_reasoning.py](../../examples/online_serving/openai_chat_completion_tool_calls_with_reasoning.py).
For more examples, please refer to [examples/reasoning/openai_chat_completion_tool_calls_with_reasoning.py](../../examples/reasoning/openai_chat_completion_tool_calls_with_reasoning.py).
## Server-Level Default Chat Template Kwargs
+2
View File
@@ -439,6 +439,7 @@ th {
| `Mamba2ForCausalLM` | Mamba2 | `mistralai/Mamba-Codestral-7B-v0.1`, etc. | | ✅︎ |
| `MiMoForCausalLM` | MiMo | `XiaomiMiMo/MiMo-7B-RL`, etc. | ✅︎ | ✅︎ |
| `MiMoV2FlashForCausalLM` | MiMoV2Flash | `XiaomiMiMo/MiMo-V2-Flash`, etc. | | ✅︎ |
| `MiMoV2ProForCausalLM` | MiMoV2Pro | `XiaomiMiMo/MiMo-V2.5-Pro`, etc. | | ✅︎ |
| `MiniCPMForCausalLM` | MiniCPM | `openbmb/MiniCPM-2B-sft-bf16`, `openbmb/MiniCPM-2B-dpo-bf16`, `openbmb/MiniCPM-S-1B-sft`, etc. | ✅︎ | ✅︎ |
| `MiniCPM3ForCausalLM` | MiniCPM3 | `openbmb/MiniCPM3-4B`, etc. | ✅︎ | ✅︎ |
| `MiniMaxForCausalLM` | MiniMax-Text | `MiniMaxAI/MiniMax-Text-01-hf`, etc. | | |
@@ -590,6 +591,7 @@ These models primarily accept the [`LLM.generate`](./generative_models.md#llmgen
| `LlavaNextVideoForConditionalGeneration` | LLaVA-NeXT-Video | T + V | `llava-hf/LLaVA-NeXT-Video-7B-hf`, etc. | | ✅︎ |
| `LlavaOnevisionForConditionalGeneration` | LLaVA-Onevision | T + I<sup>+</sup> + V<sup>+</sup> | `llava-hf/llava-onevision-qwen2-7b-ov-hf`, `llava-hf/llava-onevision-qwen2-0.5b-ov-hf`, etc. | | ✅︎ |
| `MiDashengLMModel` | MiDashengLM | T + A<sup>+</sup> | `mispeech/midashenglm-7b` | | ✅︎ |
| `MiMoV2OmniForCausalLM` | MiMo-V2.5-Omni | T + I<sup>E+</sup> + V<sup>E+</sup> + A<sup>+</sup> | `XiaomiMiMo/MiMo-V2.5-Omni` | | ✅︎ |
| `MiniCPMO` | MiniCPM-O | T + I<sup>E+</sup> + V<sup>E+</sup> + A<sup>E+</sup> | `openbmb/MiniCPM-o-2_6`, etc. | ✅︎ | ✅︎ |
| `MiniCPMV` | MiniCPM-V | T + I<sup>E+</sup> + V<sup>E+</sup> | `openbmb/MiniCPM-V-2` (see note), `openbmb/MiniCPM-Llama3-V-2_5`, `openbmb/MiniCPM-V-2_6`, `openbmb/MiniCPM-V-4`, `openbmb/MiniCPM-V-4_5`, etc. | ✅︎ | |
| `MiniMaxVL01ForConditionalGeneration` | MiniMax-VL | T + I<sup>E+</sup> | `MiniMaxAI/MiniMax-VL-01`, etc. | | ✅︎ |
+3 -3
View File
@@ -251,7 +251,7 @@ The following extra parameters are supported:
Our Responses API is compatible with [OpenAI's Responses API](https://platform.openai.com/docs/api-reference/responses);
you can use the [official OpenAI Python client](https://github.com/openai/openai-python) to interact with it.
Code example: [examples/online_serving/openai_responses_client_with_tools.py](../../examples/online_serving/openai_responses_client_with_tools.py)
Code example: [examples/online_serving/openai_responses_client_with_tools.py](../../examples/tool_calling/openai_responses_client_with_tools.py)
#### Extra parameters
@@ -279,7 +279,7 @@ you can use the [official OpenAI Python client](https://github.com/openai/openai
!!! note
To use the Transcriptions API, please install with extra audio dependencies using `pip install vllm[audio]`.
Code example: [examples/online_serving/openai_transcription_client.py](../../examples/online_serving/openai_transcription_client.py)
Code example: [examples/speech_to_text/openai/openai_transcription_client.py](../../examples/speech_to_text/openai/openai_transcription_client.py)
NOTE: beam search is currently supported in the transcriptions endpoint for encoder-decoder multimodal models, e.g., whisper, but highly inefficient as work for handling the encoder/decoder cache is actively ongoing. This is an active point of ongoing optimization and will be handled properly in the very near future.
@@ -397,7 +397,7 @@ Please mind that the popular `openai/whisper-large-v3-turbo` model does not supp
!!! note
To use the Translation API, please install with extra audio dependencies using `pip install vllm[audio]`.
Code example: [examples/online_serving/openai_translation_client.py](../../examples/online_serving/openai_translation_client.py)
Code example: [examples/speech_to_text/openai/openai_translation_client.py](../../examples/speech_to_text/openai/openai_translation_client.py)
#### Extra Parameters
@@ -6,15 +6,15 @@ This folder provides several example scripts on how to inference Qwen2.5-Omni of
```bash
# Audio + image + video
python examples/offline_inference/qwen2_5_omni/only_thinker.py \
python examples/generate/multimodal/qwen2_5_omni/only_thinker.py \
-q mixed_modalities
# Read vision and audio inputs from a single video file
python examples/offline_inference/qwen2_5_omni/only_thinker.py \
python examples/generate/multimodal/qwen2_5_omni/only_thinker.py \
-q use_audio_in_video
# Multiple audios
python examples/offline_inference/qwen2_5_omni/only_thinker.py \
python examples/generate/multimodal/qwen2_5_omni/only_thinker.py \
-q multi_audios
```
@@ -24,16 +24,16 @@ You can also test Qwen2.5-Omni on a single modality:
```bash
# Process audio inputs
python examples/offline_inference/audio_language.py \
python examples/generate/multimodal/audio_language_offline.py \
--model-type qwen2_5_omni
# Process image inputs
python examples/offline_inference/vision_language.py \
python examples/generate/multimodal/vision_language_offline.py \
--modality image \
--model-type qwen2_5_omni
# Process video inputs
python examples/offline_inference/vision_language.py \
python examples/generate/multimodal/vision_language_offline.py \
--modality video \
--model-type qwen2_5_omni
```
@@ -1402,7 +1402,7 @@ def run_mantis(questions: list[str], modality: str) -> ModelRequestData:
# MiniCPM-V
def run_minicpmv_base(questions: list[str], modality: str, model_name):
assert modality in ["image", "video", "image+video"]
# If you want to use `MiniCPM-o-2_6` with audio inputs, check `audio_language.py` # noqa
# If you want to use `MiniCPM-o-2_6` with audio inputs, check `audio_language_offline.py` # noqa
# 2.0
# The official repo doesn't work yet, so we need to use a fork for now
+1 -1
View File
@@ -12,7 +12,7 @@ torchvision==0.26.0 # Required for phi3v processor. See https://github.com/pytor
flashinfer-python==0.6.8.post1
flashinfer-cubin==0.6.8.post1
apache-tvm-ffi==0.1.9
tilelang
tilelang==0.1.9
# Cap nvidia-cudnn-frontend (transitive dep of flashinfer) due to
# breaking changes in 1.19.0
nvidia-cudnn-frontend>=1.13.0,<1.19.0
@@ -261,6 +261,8 @@ def _compare_sp(
},
"use_inductor_graph_partition": use_inductor_graph_partition,
}
if not use_inductor_graph_partition:
compilation_config["splitting_ops"] = []
tp_sp_args = [
*common_args,
+5
View File
@@ -116,6 +116,11 @@ def run_e2e_fusion_test(monkeypatch, caplog_mp_spawn):
model_kwargs["attention_config"] = {"backend": attn_backend.backend.name}
model_kwargs["tensor_parallel_size"] = tp_size
# Cap warmup memory: tests use small max_model_len (1024) but the
# engine default max_num_batched_tokens is 16384. Warming up large
# models (e.g. Llama-4-Scout-FP8) at 16384 tokens may trigger OOM.
model_kwargs.setdefault("max_num_batched_tokens", 8192)
# Sparse MLA models (DSv3.2) hit an over-strict inductor assertion in
# decompose_auto_functionalized when +rotary_embedding is forced into
# the compile graph. Disable qk_norm+rope fusion (which auto-enables
@@ -19,6 +19,7 @@ from vllm.config import (
VllmConfig,
set_current_vllm_config,
)
from vllm.config.utils import Range
from vllm.distributed import (
tensor_model_parallel_all_gather,
tensor_model_parallel_reduce_scatter,
@@ -288,6 +289,22 @@ def test_async_tp_pass_replace(
run_torch_spawn(async_tp_pass_on_test_model, num_processes)
def test_async_tp_pass_requires_full_graph_compilation():
vllm_config = VllmConfig()
vllm_config.compilation_config.use_inductor_graph_partition = False
vllm_config.compilation_config.splitting_ops = [
"vllm::unified_attention_with_output"
]
async_tp_pass = object.__new__(AsyncTPPass)
async_tp_pass.compilation_config = vllm_config.compilation_config
with pytest.raises(
AssertionError, match="AsyncTPPass requires full-graph compilation"
):
async_tp_pass.is_applicable_for_range(Range(start=8, end=8))
def async_tp_pass_on_test_model(
local_rank: int,
world_size: int,
@@ -22,6 +22,7 @@ from vllm.config import (
get_current_vllm_config,
set_current_vllm_config,
)
from vllm.config.utils import Range
from vllm.distributed import tensor_model_parallel_all_reduce
from vllm.distributed.parallel_state import (
init_distributed_environment,
@@ -216,6 +217,24 @@ def test_sequence_parallelism_pass(
run_torch_spawn(sequence_parallelism_pass_on_test_model, num_processes)
def test_sequence_parallelism_pass_requires_full_graph_compilation():
vllm_config = VllmConfig()
vllm_config.compilation_config.use_inductor_graph_partition = False
vllm_config.compilation_config.splitting_ops = [
"vllm::unified_attention_with_output"
]
sequence_parallelism_pass = object.__new__(SequenceParallelismPass)
sequence_parallelism_pass.compilation_config = vllm_config.compilation_config
sequence_parallelism_pass.min_token_num = 1
with pytest.raises(
AssertionError,
match="SequenceParallelismPass requires full-graph compilation",
):
sequence_parallelism_pass.is_applicable_for_range(Range(start=8, end=8))
def sequence_parallelism_pass_on_test_model(
local_rank: int,
world_size: int,
+118 -1
View File
@@ -407,7 +407,7 @@ def test_should_split():
(None, 257, 1, False, 2048, CUDAGraphMode.FULL_AND_PIECEWISE, 256),
# max from list
([1, 2, 4, 15], None, 1, False, 2048, CUDAGraphMode.FULL_AND_PIECEWISE, 15),
# filtered out 15 due to SP
# SP forces full-graph compilation, sizes are filtered by TP
([1, 2, 4, 15], None, 2, True, 2048, CUDAGraphMode.FULL_AND_PIECEWISE, 4),
# limited by the max_tokens
([1, 2, 4, 15], None, 1, False, 8, CUDAGraphMode.FULL_AND_PIECEWISE, 4),
@@ -465,6 +465,123 @@ def test_cudagraph_sizes_post_init(
)
@pytest.mark.skipif(
not current_platform.support_static_graph_mode(),
reason="Skip if not cudagraph mode supported",
)
@pytest.mark.parametrize(
(
"cudagraph_mode",
"use_inductor_graph_partition",
"expected_enable_sp",
"expected_cudagraph_mode",
"expected_piecewise_compile",
"expected_capture_sizes",
"expected_max_size",
),
[
(CUDAGraphMode.PIECEWISE, False, True, CUDAGraphMode.FULL, False, [2, 4], 4),
(
CUDAGraphMode.FULL_DECODE_ONLY,
False,
True,
CUDAGraphMode.FULL_DECODE_ONLY,
False,
[2, 4],
4,
),
(
CUDAGraphMode.FULL_AND_PIECEWISE,
False,
True,
CUDAGraphMode.FULL,
False,
[2, 4],
4,
),
(
CUDAGraphMode.FULL_AND_PIECEWISE,
True,
True,
CUDAGraphMode.FULL_AND_PIECEWISE,
True,
[2, 4],
4,
),
],
)
def test_sequence_parallelism_requires_full_graph_compilation(
cudagraph_mode: CUDAGraphMode,
use_inductor_graph_partition: bool,
expected_enable_sp: bool,
expected_cudagraph_mode: CUDAGraphMode,
expected_piecewise_compile: bool,
expected_capture_sizes: list[int],
expected_max_size: int,
):
with patch.object(current_platform, "device_count", return_value=2):
vllm_config = VllmConfig(
parallel_config=ParallelConfig(tensor_parallel_size=2),
scheduler_config=SchedulerConfig(
max_num_seqs=128,
max_num_batched_tokens=2048,
max_model_len=2048,
is_encoder_decoder=False,
),
)
vllm_config.model_config = MagicMock(
dtype=torch.float16,
enforce_eager=False,
is_moe=False,
disable_cascade_attn=False,
get_hidden_size=MagicMock(return_value=4096),
)
vllm_config.compilation_config = CompilationConfig(
mode=CompilationMode.VLLM_COMPILE,
cudagraph_capture_sizes=[1, 2, 4, 15],
max_cudagraph_capture_size=None,
compile_sizes=["cudagraph_capture_sizes"],
use_inductor_graph_partition=use_inductor_graph_partition,
pass_config=PassConfig(
enable_sp=True,
fuse_gemm_comms=True,
fuse_norm_quant=True,
fuse_act_quant=True,
eliminate_noops=True,
sp_min_token_num=512,
),
cudagraph_mode=cudagraph_mode,
)
vllm_config.compilation_config.set_splitting_ops_for_v1(
all2all_backend=vllm_config.parallel_config.all2all_backend,
data_parallel_size=1,
)
vllm_config._set_compile_ranges()
vllm_config._set_cudagraph_sizes()
assert (
vllm_config.compilation_config.use_inductor_graph_partition
== use_inductor_graph_partition
)
assert (
bool(vllm_config.compilation_config.splitting_ops) == expected_piecewise_compile
)
assert vllm_config.compilation_config.pass_config.enable_sp == expected_enable_sp
assert (
vllm_config.compilation_config.pass_config.fuse_gemm_comms == expected_enable_sp
)
assert vllm_config.compilation_config.cudagraph_mode == expected_cudagraph_mode
assert (
vllm_config.compilation_config.cudagraph_capture_sizes == expected_capture_sizes
)
assert (
vllm_config.compilation_config.max_cudagraph_capture_size == expected_max_size
)
assert (
511 in vllm_config.compilation_config.compile_ranges_endpoints
) == expected_enable_sp
def test_cached_compilation_config(default_vllm_config):
import torch
from torch._inductor.utils import run_and_get_code
+18
View File
@@ -41,3 +41,21 @@ def test_language_model_only_affects_model_hash():
base_hash = ModelConfig(model).compute_hash()
lm_only_hash = ModelConfig(model, language_model_only=True).compute_hash()
assert base_hash != lm_only_hash
def test_mm_encoder_fp8_scale_path_requires_fp8():
with pytest.raises(ValueError, match="mm_encoder_attn_dtype"):
MultiModalConfig(mm_encoder_fp8_scale_path="/tmp/scales.json")
def test_mm_encoder_attn_dtype_hash_updates(tmp_path):
scale_file = tmp_path / "scales.json"
scale_file.write_text("{}")
base_hash = MultiModalConfig().compute_hash()
fp8_hash = MultiModalConfig(mm_encoder_attn_dtype="fp8").compute_hash()
fp8_static_hash = MultiModalConfig(
mm_encoder_attn_dtype="fp8",
mm_encoder_fp8_scale_path=str(scale_file),
).compute_hash()
assert base_hash != fp8_hash
assert fp8_hash != fp8_static_hash
@@ -0,0 +1,76 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
"""Unit tests for ``system_fingerprint`` construction."""
from types import SimpleNamespace
import pytest
from vllm.entrypoints.openai import fingerprint as fp
def _cfg(tp=1, pp=1, dp=1, ep=False, digest="a3b21f94deadbeef"):
c = SimpleNamespace(
parallel_config=SimpleNamespace(
tensor_parallel_size=tp,
pipeline_parallel_size=pp,
data_parallel_size=dp,
enable_expert_parallel=ep,
)
)
c.compute_hash = lambda: digest # type: ignore[attr-defined]
return c
@pytest.fixture(autouse=True)
def _reset():
fp.set_default_fingerprint_mode("full")
yield
fp.set_default_fingerprint_mode("full")
def test_four_modes_produce_expected_shapes():
from vllm import __version__ as v
cfg = _cfg(tp=8, ep=True)
assert fp.build_system_fingerprint(cfg, "full") == (f"vllm-{v}-tp8-ep-a3b21f94")
assert fp.build_system_fingerprint(cfg, "hash") == f"vllm-{v}-a3b21f94"
assert fp.build_system_fingerprint(cfg, "custom", "my-fp") == "my-fp"
assert fp.build_system_fingerprint(cfg, "none") is None
def test_full_mode_emits_only_non_trivial_parallelism():
from vllm import __version__ as v
# Single-GPU: nothing between version and hash.
assert fp.build_system_fingerprint(_cfg(), "full") == f"vllm-{v}-a3b21f94"
# All parallelism axes.
assert (
fp.build_system_fingerprint(_cfg(tp=8, pp=2, dp=4, ep=True), "full")
== f"vllm-{v}-tp8-pp2-dp4-ep-a3b21f94"
)
def test_get_respects_set_default():
cfg = _cfg(tp=8)
full = fp.get_system_fingerprint(cfg)
assert full == fp.get_system_fingerprint(cfg)
fp.set_default_fingerprint_mode("hash")
hashed = fp.get_system_fingerprint(cfg)
assert hashed != full
assert "tp8" not in hashed
fp.set_default_fingerprint_mode("custom", "deploy-42")
assert fp.get_system_fingerprint(cfg) == "deploy-42"
fp.set_default_fingerprint_mode("none")
assert fp.get_system_fingerprint(cfg) is None
def test_compute_hash_failure_does_not_raise():
cfg = _cfg()
cfg.compute_hash = lambda: (_ for _ in ()).throw(RuntimeError("boom"))
assert fp.build_system_fingerprint(cfg, "full").endswith("-nohash")
assert fp.build_system_fingerprint(cfg, "hash").endswith("-nohash")
@@ -128,7 +128,7 @@ def test_deepgemm_fp8_mqa_logits(clean_logits: bool):
q_fp8 = q.to(torch.float8_e4m3fn)
kv_fp8 = per_custom_dims_cast_to_fp8(kv, (0,), False)
logits = fp8_fp4_mqa_logits(
q_fp8, kv_fp8, weights, ks, ke, clean_logits=clean_logits
(q_fp8, None), kv_fp8, weights, ks, ke, clean_logits=clean_logits
)
ref_logits = _ref_fp8_mqa_logits(
+80
View File
@@ -16,6 +16,7 @@ from vllm.model_executor.layers.activation import (
NewGELU,
QuickGELU,
SiluAndMul,
SiluAndMulWithClamp,
SwigluOAIAndMul,
SwigluStepAndMul,
swiglustep_and_mul_triton,
@@ -116,6 +117,85 @@ def test_act_and_mul(
opcheck(fn, (out, x))
SWIGLU_LIMITS = [3.0, 7.0, 15.0]
@pytest.mark.parametrize("swiglu_limit", SWIGLU_LIMITS)
@pytest.mark.parametrize("num_tokens", NUM_TOKENS)
@pytest.mark.parametrize("d", D)
@pytest.mark.parametrize("dtype", DTYPES)
@pytest.mark.parametrize("seed", SEEDS)
@pytest.mark.parametrize("device", CUDA_DEVICES)
@torch.inference_mode()
def test_silu_and_mul_with_clamp(
default_vllm_config,
swiglu_limit: float,
num_tokens: int,
d: int,
dtype: torch.dtype,
seed: int,
device: str,
) -> None:
"""SiluAndMulWithClamp: cuda kernel must match native reference."""
set_random_seed(seed)
torch.set_default_device(device)
# Use large values to ensure clamping is exercised.
x = torch.randn(num_tokens, 2 * d, dtype=dtype) * swiglu_limit * 2
layer = SiluAndMulWithClamp(swiglu_limit, compile_native=False)
out = layer(x)
ref_out = layer.forward_native(x)
rtol = {
torch.float16: 2e-3,
torch.bfloat16: 2e-2,
torch.float: 1.3e-6,
}
torch.testing.assert_close(
out, ref_out, atol=get_default_atol(out), rtol=rtol[out.dtype]
)
# Verify clamping is actually being applied: the clamped output should
# differ from the unclamped SiluAndMul output when inputs are large.
unclamped_out = SiluAndMul.forward_native(x)
assert not torch.equal(ref_out.float(), unclamped_out.float()), (
"Input was not large enough to exercise the clamp; increase scale"
)
# Verify gate clamping semantics with a controlled scalar case.
# gate=large_val is clamped to limit first, then silu(limit) * 1.0.
x_gate = torch.tensor(
[[swiglu_limit * 20.0, 1.0]], dtype=torch.float32, device=device
)
out_gate = SiluAndMulWithClamp(swiglu_limit, compile_native=False)(x_gate)
expected_gate = torch.nn.functional.silu(
torch.tensor(swiglu_limit, dtype=torch.float32)
).item()
torch.testing.assert_close(
out_gate,
torch.tensor([[expected_gate]], dtype=torch.float32, device=device),
atol=1e-3,
rtol=1e-3,
)
# Verify up clamping semantics: up >> limit gets clamped to limit.
x_up = torch.tensor(
[[1.0, swiglu_limit * 20.0]], dtype=torch.float32, device=device
)
out_up = SiluAndMulWithClamp(swiglu_limit, compile_native=False)(x_up)
silu_1 = torch.nn.functional.silu(torch.tensor(1.0)).item()
torch.testing.assert_close(
out_up,
torch.tensor([[silu_1 * swiglu_limit]], dtype=torch.float32, device=device),
atol=1e-3,
rtol=1e-3,
)
# opcheck
out_buf = torch.empty(x.shape[:-1] + (d,), dtype=dtype, device=device)
opcheck(torch.ops._C.silu_and_mul_with_clamp, (out_buf, x, swiglu_limit))
@pytest.mark.parametrize(
"activation",
[
+279
View File
@@ -0,0 +1,279 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
"""Tests for the full FP8 ViT attention path (quantize -> cuDNN -> un-pad)."""
import contextlib
import pytest
import torch
from vllm.triton_utils import HAS_TRITON
from vllm.utils.flashinfer import (
is_flashinfer_cudnn_fp8_prefill_attn_supported,
)
from vllm.v1.attention.backends.registry import AttentionBackendEnum
def _has_flashinfer_cudnn() -> bool:
"""Check if FlashInfer cuDNN backend is available."""
try:
from flashinfer.prefill import (
cudnn_batch_prefill_with_kv_cache, # noqa: F401
)
return True
except ImportError:
return False
HEAD_DIMS = [72, 80]
SEQ_LENS = [256]
NUM_HEADS = [16]
@pytest.fixture
def _fp8_attention():
"""Create FP8-enabled MMEncoderAttention via config."""
from types import SimpleNamespace
from unittest.mock import patch
from vllm.config import VllmConfig, set_current_vllm_config
from vllm.config.multimodal import MultiModalConfig
if not is_flashinfer_cudnn_fp8_prefill_attn_supported():
pytest.skip("FlashInfer cuDNN FP8 prefill attention not supported")
mm_config = MultiModalConfig(mm_encoder_attn_dtype="fp8")
vllm_config = VllmConfig()
vllm_config.model_config = SimpleNamespace(multimodal_config=mm_config)
# MMEncoderAttention reads torch.get_default_dtype() during init
# to determine the output dtype. In real model loading this is bf16.
old_dtype = torch.get_default_dtype()
torch.set_default_dtype(torch.bfloat16)
with (
set_current_vllm_config(vllm_config),
patch(
"vllm.model_executor.layers.attention.mm_encoder_attention"
".get_vit_attn_backend",
return_value=AttentionBackendEnum.FLASHINFER,
),
):
yield
torch.set_default_dtype(old_dtype)
def _build_cu_seqlens_and_meta(
seq_len: int,
num_heads: int,
head_dim: int,
fp8_padded_hidden_size: int | None = None,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
"""Build cu_seqlens, max_seqlen, sequence_lengths for a single sequence."""
import numpy as np
from vllm.model_executor.layers.attention.mm_encoder_attention import (
MMEncoderAttention,
)
cu_seqlens_np = np.array([0, seq_len], dtype=np.int32)
sequence_lengths = MMEncoderAttention.maybe_compute_seq_lens(
AttentionBackendEnum.FLASHINFER,
cu_seqlens_np,
torch.device("cuda"),
)
max_seqlen = torch.tensor(
MMEncoderAttention.compute_max_seqlen(
AttentionBackendEnum.FLASHINFER, cu_seqlens_np
),
dtype=torch.int32,
)
cu_seqlens = MMEncoderAttention.maybe_recompute_cu_seqlens(
AttentionBackendEnum.FLASHINFER,
cu_seqlens_np,
num_heads * head_dim,
1, # tp_size
torch.device("cuda"),
fp8_padded_hidden_size=fp8_padded_hidden_size,
)
return cu_seqlens, max_seqlen, sequence_lengths
@pytest.mark.skipif(
not (HAS_TRITON and _has_flashinfer_cudnn()),
reason="Triton and FlashInfer cuDNN required",
)
@pytest.mark.parametrize("head_dim", HEAD_DIMS)
@pytest.mark.parametrize("seq_len", SEQ_LENS)
@pytest.mark.parametrize("num_heads", NUM_HEADS)
def test_fp8_attn_output_shape(
head_dim: int,
seq_len: int,
num_heads: int,
_fp8_attention,
) -> None:
"""Verify FP8 attention produces correct output shape after un-padding."""
from vllm.model_executor.layers.attention.mm_encoder_attention import (
MMEncoderAttention,
)
from vllm.utils.math_utils import round_up
attn = None
with contextlib.suppress(ValueError, ImportError):
attn = MMEncoderAttention(
num_heads=num_heads,
head_size=head_dim,
prefix="visual.blocks.0.attn",
).to("cuda")
if attn is None or not attn.fp8_enabled:
pytest.skip("FP8 MMEncoderAttention not available")
assert attn is not None # mypy narrowing
# FP8 always needs fp8_padded_hidden_size for correct cu_seqlens
fp8_padded_hidden_size = num_heads * round_up(head_dim, 16)
cu_seqlens, max_seqlen, sequence_lengths = _build_cu_seqlens_and_meta(
seq_len, num_heads, head_dim, fp8_padded_hidden_size=fp8_padded_hidden_size
)
q = torch.randn(
seq_len,
num_heads,
head_dim,
device="cuda",
dtype=torch.bfloat16,
)
k = torch.randn_like(q)
v = torch.randn_like(q)
output = attn._forward_flashinfer(q, k, v, cu_seqlens, max_seqlen, sequence_lengths)
# Output should have original head_dim (un-padded)
assert output.shape[-1] == head_dim
assert output.dtype == torch.bfloat16
@pytest.mark.skipif(
not (HAS_TRITON and _has_flashinfer_cudnn()),
reason="Triton and FlashInfer cuDNN required",
)
@pytest.mark.parametrize("head_dim", HEAD_DIMS)
@pytest.mark.parametrize("seq_len", SEQ_LENS)
@pytest.mark.parametrize("num_heads", NUM_HEADS)
def test_fp8_vs_bf16_close(
head_dim: int, seq_len: int, num_heads: int, _fp8_attention
) -> None:
"""FP8 attention output should be reasonably close to BF16 baseline."""
from vllm.model_executor.layers.attention.mm_encoder_attention import (
MMEncoderAttention,
)
from vllm.utils.math_utils import round_up
torch.manual_seed(42)
q = torch.randn(
1,
seq_len,
num_heads,
head_dim,
device="cuda",
dtype=torch.bfloat16,
)
k = torch.randn_like(q)
v = torch.randn_like(q)
# FP8 path
attn_fp8 = None
with contextlib.suppress(ValueError, ImportError):
attn_fp8 = MMEncoderAttention(
num_heads=num_heads,
head_size=head_dim,
prefix="visual.blocks.0.attn",
).to("cuda")
if attn_fp8 is None or not attn_fp8.fp8_enabled:
pytest.skip("FP8 MMEncoderAttention not available")
assert attn_fp8 is not None # mypy narrowing
fp8_padded_hidden_size = num_heads * round_up(head_dim, 16)
cu_seqlens, max_seqlen, seq_lengths = _build_cu_seqlens_and_meta(
seq_len,
num_heads,
head_dim,
fp8_padded_hidden_size=fp8_padded_hidden_size,
)
out_fp8 = attn_fp8._forward_flashinfer(
q.clone(),
k.clone(),
v.clone(),
cu_seqlens,
max_seqlen,
seq_lengths,
)
# BF16 baseline (create non-FP8 attention by using scale=attn_fp8.scale
# and calling the wrapper directly without FP8 quantization)
from vllm.model_executor.layers.attention.mm_encoder_attention import (
_get_flashinfer_workspace_buffer,
)
from vllm.v1.attention.ops.vit_attn_wrappers import (
vit_flashinfer_wrapper,
)
out_bf16 = vit_flashinfer_wrapper(
q=q.clone(),
k=k.clone(),
v=v.clone(),
scale=attn_fp8.scale,
workspace_buffer=_get_flashinfer_workspace_buffer(),
cu_seqlens=cu_seqlens,
max_seqlen=max_seqlen,
sequence_lengths=seq_lengths,
)
out_fp8_f = out_fp8.float()
out_bf16_f = out_bf16.float()
abs_diff = (out_fp8_f - out_bf16_f).abs()
abs_diff_flat = abs_diff.flatten()
# Relative diff (avoid division by zero)
denom = out_bf16_f.abs().clamp(min=1e-6)
rel_diff_flat = (abs_diff / denom).flatten()
cosine_sim = torch.nn.functional.cosine_similarity(
out_fp8_f.flatten().unsqueeze(0),
out_bf16_f.flatten().unsqueeze(0),
).item()
pcts = [50, 90, 95, 99, 99.9]
abs_pct = {p: torch.quantile(abs_diff_flat, p / 100).item() for p in pcts}
rel_pct = {p: torch.quantile(rel_diff_flat, p / 100).item() for p in pcts}
print(f"\nFP8 vs BF16 (head_dim={head_dim}, seq_len={seq_len}):")
print(f" cosine_sim={cosine_sim:.6f}")
print(
f" abs_diff: max={abs_diff_flat.max().item():.6f}, "
f"mean={abs_diff_flat.mean().item():.6f}, "
+ ", ".join(f"p{p}={abs_pct[p]:.6f}" for p in pcts)
)
print(
f" rel_diff: max={rel_diff_flat.max().item():.6f}, "
f"mean={rel_diff_flat.mean().item():.6f}, "
+ ", ".join(f"p{p}={rel_pct[p]:.6f}" for p in pcts)
)
assert abs_diff_flat.max().item() < 0.3, (
f"FP8 vs BF16 max abs diff too large: {abs_diff_flat.max().item()}"
)
assert abs_diff_flat.mean().item() < 0.03, (
f"FP8 vs BF16 mean abs diff too large: {abs_diff_flat.mean().item()}"
)
assert cosine_sim > 0.99, f"Cosine similarity too low: {cosine_sim:.6f}"
+124
View File
@@ -0,0 +1,124 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
"""Tests for the stride-aware FP8 quantization kernel with head_dim padding."""
import pytest
import torch
from vllm.platforms import current_platform
from vllm.triton_utils import HAS_TRITON
if HAS_TRITON:
from vllm.kernels.triton.qkv_padded_fp8_quant import (
quantize_fp8_pad_head_dim_triton,
)
HEAD_DIMS = [72, 80, 128]
SEQ_LENS = [64, 256]
NUM_HEADS = [16]
SCALES = [0.01, 0.1, 1.0]
def _naive_fp8_quantize(
tensor: torch.Tensor, scale: torch.Tensor, skip_scale: bool
) -> torch.Tensor:
"""Reference FP8 quantization in PyTorch."""
fp8_dtype = current_platform.fp8_dtype()
fp8_max = torch.finfo(fp8_dtype).max
fp8_min = -fp8_max
x = tensor.float()
if not skip_scale:
x = x / scale.item()
x = x.clamp(fp8_min, fp8_max)
return x.to(fp8_dtype)
@pytest.mark.skipif(not HAS_TRITON, reason="Triton not available")
@pytest.mark.parametrize("head_dim", HEAD_DIMS)
@pytest.mark.parametrize("seq_len", SEQ_LENS)
@pytest.mark.parametrize("num_heads", NUM_HEADS)
@pytest.mark.parametrize("scale_val", SCALES)
def test_quantize_contiguous(
head_dim: int, seq_len: int, num_heads: int, scale_val: float
) -> None:
"""Test quantization of contiguous 3D tensors."""
torch.manual_seed(42)
tensor = torch.randn(
seq_len, num_heads, head_dim, device="cuda", dtype=torch.bfloat16
)
scale = torch.tensor([scale_val], dtype=torch.float32, device="cuda").view(
1, 1, 1, 1
)
result = quantize_fp8_pad_head_dim_triton(tensor, scale)
padded_dim = (head_dim + 15) // 16 * 16
assert result.shape == (seq_len, num_heads, padded_dim)
assert result.is_contiguous()
assert result.dtype == current_platform.fp8_dtype()
# Compare unpadded portion against reference
ref = _naive_fp8_quantize(tensor, scale, skip_scale=False)
torch.testing.assert_close(result[:, :, :head_dim].float(), ref.float())
# Padded region should be zero
if padded_dim > head_dim:
assert (result[:, :, head_dim:].float() == 0).all()
@pytest.mark.skipif(not HAS_TRITON, reason="Triton not available")
@pytest.mark.parametrize("head_dim", [72, 80])
def test_quantize_non_contiguous(head_dim: int) -> None:
"""Test quantization from non-contiguous QKV views (interleaved buffer)."""
seq_len, num_heads = 64, 16
# Simulate interleaved QKV buffer: shape (seq_len, 3 * num_heads, head_dim)
qkv = torch.randn(
seq_len, 3 * num_heads, head_dim, device="cuda", dtype=torch.bfloat16
)
# Q is every 3rd head slice - non-contiguous view
q = qkv[:, 0::3, :]
assert not q.is_contiguous()
scale = torch.tensor([0.1], dtype=torch.float32, device="cuda").view(1, 1, 1, 1)
result = quantize_fp8_pad_head_dim_triton(q, scale)
padded_dim = (head_dim + 15) // 16 * 16
assert result.shape == (seq_len, num_heads, padded_dim)
assert result.is_contiguous()
# Compare against contiguous reference
ref = _naive_fp8_quantize(q.contiguous(), scale, skip_scale=False)
torch.testing.assert_close(result[:, :, :head_dim].float(), ref.float())
@pytest.mark.skipif(not HAS_TRITON, reason="Triton not available")
def test_skip_scale() -> None:
"""Test skip_scale=True produces cast-only output (no division)."""
seq_len, num_heads, head_dim = 32, 8, 80
tensor = torch.randn(
seq_len, num_heads, head_dim, device="cuda", dtype=torch.bfloat16
)
scale = torch.tensor([0.5], dtype=torch.float32, device="cuda").view(1, 1, 1, 1)
result_skip = quantize_fp8_pad_head_dim_triton(tensor, scale, skip_scale=True)
result_noskip = quantize_fp8_pad_head_dim_triton(tensor, scale, skip_scale=False)
# skip_scale should just cast, not divide
ref_cast = _naive_fp8_quantize(tensor, scale, skip_scale=True)
torch.testing.assert_close(result_skip[:, :, :head_dim].float(), ref_cast.float())
# With scale != 1.0, skip and no-skip should differ
assert not torch.equal(result_skip.float(), result_noskip.float())
@pytest.mark.skipif(not HAS_TRITON, reason="Triton not available")
def test_4d_input() -> None:
"""Test that 4D input (B, S, H, D) is handled correctly."""
B, S, H, D = 2, 32, 8, 72
tensor = torch.randn(B, S, H, D, device="cuda", dtype=torch.bfloat16)
scale = torch.tensor([0.1], dtype=torch.float32, device="cuda").view(1, 1, 1, 1)
result = quantize_fp8_pad_head_dim_triton(tensor, scale)
padded_dim = (D + 15) // 16 * 16
assert result.shape == (B, S, H, padded_dim)
+251
View File
@@ -0,0 +1,251 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
"""Tests for FP8 scaling (dynamic and static) in MMEncoderAttention."""
import contextlib
import json
from types import SimpleNamespace
from unittest.mock import patch
import pytest
import torch
from vllm.model_executor.layers.attention.mm_encoder_attention import (
_FP8_AMAX_HISTORY_LEN,
_FP8_MAX,
)
from vllm.utils.flashinfer import (
is_flashinfer_cudnn_fp8_prefill_attn_supported,
)
LAYER_0 = "visual.blocks.0.attn.attn"
LAYER_1 = "visual.blocks.1.attn.attn"
NUM_HEADS = 16
HEAD_DIM = 72
@contextlib.contextmanager
def _build_attention(mm_config):
"""Yield an MMEncoderAttention with the given multimodal config.
The VllmConfig context stays active while the test runs so that
``get_multimodal_config()`` calls during the forward path resolve. Also
invokes ``process_weights_after_loading`` to simulate the model loader's
auto-scan. Yields ``None`` if FlashInfer cuDNN is not available.
"""
from vllm.config import VllmConfig, set_current_vllm_config
from vllm.model_executor.layers.attention.mm_encoder_attention import (
MMEncoderAttention,
)
from vllm.v1.attention.backends.registry import AttentionBackendEnum
if not is_flashinfer_cudnn_fp8_prefill_attn_supported():
yield None
return
vllm_config = VllmConfig()
vllm_config.model_config = SimpleNamespace(multimodal_config=mm_config)
with (
set_current_vllm_config(vllm_config),
patch(
"vllm.model_executor.layers.attention.mm_encoder_attention"
".get_vit_attn_backend",
return_value=AttentionBackendEnum.FLASHINFER,
),
):
attn = MMEncoderAttention(
num_heads=NUM_HEADS,
head_size=HEAD_DIM,
prefix=LAYER_0,
)
attn.process_weights_after_loading(torch.bfloat16)
yield attn
@pytest.fixture
def _make_attention():
"""Create an MMEncoderAttention with dynamic FP8 scaling."""
from vllm.config.multimodal import MultiModalConfig
with _build_attention(MultiModalConfig(mm_encoder_attn_dtype="fp8")) as attn:
yield attn
@pytest.fixture
def _make_static_attention(tmp_path):
"""Create an MMEncoderAttention with static FP8 scales from a file."""
from vllm.config.multimodal import MultiModalConfig
scale_file = tmp_path / "scales.json"
scale_file.write_text(
json.dumps(
{
LAYER_0: {"q": 224.0, "k": 198.0, "v": 210.0},
LAYER_1: {"q": 100.0, "k": 110.0, "v": 120.0},
}
)
)
with _build_attention(
MultiModalConfig(
mm_encoder_attn_dtype="fp8",
mm_encoder_fp8_scale_path=str(scale_file),
)
) as attn:
yield attn
def test_dynamic_scaling_updates_scales(_make_attention) -> None:
"""Verify that _record_amax_and_update_scales updates scale buffers."""
attn = _make_attention
if attn is None or not attn.fp8_enabled:
pytest.skip("FP8 attention not available (FlashInfer backend required)")
attn = attn.to("cuda")
S, H, D = 32, NUM_HEADS, HEAD_DIM
q = torch.full((S, H, D), 2.0, device="cuda", dtype=torch.bfloat16)
k = torch.full((S, H, D), 3.0, device="cuda", dtype=torch.bfloat16)
v = torch.full((S, H, D), 4.0, device="cuda", dtype=torch.bfloat16)
attn._record_amax_and_update_scales(q, k, v)
expected_q_scale = 2.0 / _FP8_MAX
expected_k_scale = 3.0 / _FP8_MAX
expected_v_scale = 4.0 / _FP8_MAX
torch.testing.assert_close(attn._fp8_q_scale.item(), expected_q_scale)
torch.testing.assert_close(attn._fp8_k_scale.item(), expected_k_scale)
torch.testing.assert_close(attn._fp8_v_scale.item(), expected_v_scale)
def test_circular_buffer_wraps(_make_attention) -> None:
"""Verify the amax circular buffer wraps at HISTORY_LEN."""
attn = _make_attention
if attn is None or not attn.fp8_enabled:
pytest.skip("FP8 attention not available (FlashInfer backend required)")
attn = attn.to("cuda")
S, H, D = 16, NUM_HEADS, HEAD_DIM
for i in range(_FP8_AMAX_HISTORY_LEN + 2):
mag = float(i + 1)
q = torch.full((S, H, D), mag, device="cuda", dtype=torch.bfloat16)
k = torch.full((S, H, D), mag, device="cuda", dtype=torch.bfloat16)
v = torch.full((S, H, D), mag, device="cuda", dtype=torch.bfloat16)
attn._record_amax_and_update_scales(q, k, v)
assert attn._fp8_amax_pos == 2
expected_max = float(_FP8_AMAX_HISTORY_LEN + 2)
expected_scale = expected_max / _FP8_MAX
torch.testing.assert_close(attn._fp8_q_scale.item(), expected_scale)
def test_static_scales_loaded(_make_static_attention) -> None:
"""Verify static scales are loaded from the JSON file."""
attn = _make_static_attention
if attn is None or not attn.fp8_enabled:
pytest.skip("FP8 attention not available (FlashInfer backend required)")
assert attn.fp8_enabled
assert not attn._fp8_dynamic_scale
# Layer 0 scales (the layer this attention was created with).
assert attn._fp8_q_scale.item() == 224.0
assert attn._fp8_k_scale.item() == 198.0
assert attn._fp8_v_scale.item() == 210.0
assert not attn.skip_scale_q
assert not attn.skip_scale_k
assert not attn.skip_scale_v
# No amax history buffers for static scaling.
assert not hasattr(attn, "_fp8_q_amax")
def test_static_scales_missing_layer(tmp_path) -> None:
"""Verify error when requested layer is not in the scale file."""
from vllm.config import VllmConfig, set_current_vllm_config
from vllm.config.multimodal import MultiModalConfig
from vllm.v1.attention.backends.registry import AttentionBackendEnum
if not is_flashinfer_cudnn_fp8_prefill_attn_supported():
pytest.skip("FlashInfer cuDNN not available")
scale_file = tmp_path / "wrong_layer.json"
scale_file.write_text(
json.dumps({"visual.blocks.99.attn": {"q": 1.0, "k": 1.0, "v": 1.0}})
)
mm_config = MultiModalConfig(
mm_encoder_attn_dtype="fp8",
mm_encoder_fp8_scale_path=str(scale_file),
)
vllm_config = VllmConfig()
vllm_config.model_config = SimpleNamespace(multimodal_config=mm_config)
from vllm.model_executor.layers.attention.mm_encoder_attention import (
MMEncoderAttention,
)
with (
set_current_vllm_config(vllm_config),
patch(
"vllm.model_executor.layers.attention.mm_encoder_attention"
".get_vit_attn_backend",
return_value=AttentionBackendEnum.FLASHINFER,
),
):
attn = MMEncoderAttention(
num_heads=NUM_HEADS,
head_size=HEAD_DIM,
prefix=LAYER_0,
)
with pytest.raises(ValueError, match="scales not found for layer"):
attn.process_weights_after_loading(torch.bfloat16)
def test_dynamic_scales_auto_save(tmp_path) -> None:
"""Verify scales are saved to disk after the amax buffer fills."""
import vllm.model_executor.layers.attention.mm_encoder_attention as _mod
from vllm.config.multimodal import MultiModalConfig
if not is_flashinfer_cudnn_fp8_prefill_attn_supported():
pytest.skip("FlashInfer cuDNN not available")
# Reset module-level state between runs (other tests may have left
# state behind after triggering a save).
_mod._fp8_scale_save_path = None
_mod._fp8_saved_scale_refs.clear()
save_file = tmp_path / "auto_scales.json"
with _build_attention(
MultiModalConfig(
mm_encoder_attn_dtype="fp8",
mm_encoder_fp8_scale_save_path=str(save_file),
)
) as attn:
if attn is None or not attn.fp8_enabled:
pytest.skip("FP8 attention not available")
attn = attn.to("cuda")
S, H, D = 16, NUM_HEADS, HEAD_DIM
# Run exactly _FP8_AMAX_HISTORY_LEN forward passes.
for i in range(_FP8_AMAX_HISTORY_LEN):
mag = float(i + 1)
q = torch.full((S, H, D), mag, device="cuda", dtype=torch.bfloat16)
k = torch.full((S, H, D), mag * 0.5, device="cuda", dtype=torch.bfloat16)
v = torch.full((S, H, D), mag * 0.3, device="cuda", dtype=torch.bfloat16)
attn._record_amax_and_update_scales(q, k, v)
# File should have been written on the 16th call (buffer wrap).
assert save_file.is_file(), "Scale file was not saved"
scales = json.loads(save_file.read_text())
assert LAYER_0 in scales
assert set(scales[LAYER_0].keys()) == {"q", "k", "v"}
for val in scales[LAYER_0].values():
assert isinstance(val, float) and val > 0
# Path is cleared after the one-shot save fires.
assert _mod._fp8_scale_save_path is None
-667
View File
@@ -1,667 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
"""Correctness tests for fast_topk_v2 (DeepSeek V4 indexer top-k, k=512).
Run::
.venv/bin/python -m pytest tests/kernels/test_fast_topk_v2.py -v
Coverage:
- All four execution paths: trivial (sl<=512), Register (1- and 2-pass),
Streaming, and Cluster.
- Both launch shapes: fused (batch<=kNumClusters=15) and two-stage (>15).
- Mixed-length batches that exercise the per-row dispatch in the stage-2
combine kernel.
- Page-table fold-in: parametrised across page_size in {1, 32, 64}.
The kernel emits page-table-resolved indices. By using
``page_table[b, i] = i`` with ``page_size=1`` we can compare the kernel's
output 1:1 against ``torch.topk`` on the masked scores. For other page sizes
the test inverts the page resolution before comparing.
"""
from __future__ import annotations
import pytest
import torch
from vllm.platforms import current_platform
from vllm.v1.attention.ops.deepseek_v4_ops.fast_topk import (
fast_topk_v2,
fast_topk_v2_raw,
plan_topk_v2,
workspace_ints_per_batch,
)
# Match the kernel's compile-time constant.
TOPK = 512
# Thresholds inside the kernel (mirrors values in topk/register.cuh,
# topk_v2.cuh). Keep these in sync if the kernel changes.
SMALL_1PASS = 4 * 4 * 1024 # RegisterTopK::kMax1PassLength
SMALL_2PASS = 2 * SMALL_1PASS # RegisterTopK::kMax2PassLength = 32768
DEFAULT_CLUSTER_THRESHOLD = SMALL_2PASS # plan picks this for batch<=30
NUM_CLUSTERS = 15 # kNumClusters in fast_topk_v2.cu
# --------------------------------------------------------------------------
# Helpers
# --------------------------------------------------------------------------
def _max_blocks_for(seq_len: int, page_size: int) -> int:
return (seq_len + page_size - 1) // page_size
def _trivial_page_table(batch_size: int, max_blocks: int,
device: torch.device) -> torch.Tensor:
"""Identity page table: page_table[b, i] = i, so page_to_indices is a no-op
when ``page_size == 1`` (page_bits == 0)."""
return (
torch.arange(max_blocks, dtype=torch.int32, device=device)
.unsqueeze(0)
.expand(batch_size, -1)
.contiguous()
)
def _shuffled_page_table(batch_size: int, max_blocks: int, seed: int,
device: torch.device) -> torch.Tensor:
"""Per-row independent permutation of [0, max_blocks)."""
g = torch.Generator(device=device).manual_seed(seed)
rows = []
for _ in range(batch_size):
rows.append(torch.randperm(max_blocks, generator=g, device=device,
dtype=torch.int32))
return torch.stack(rows, dim=0)
def _resolve(raw_idx: int, b: int, page_table: torch.Tensor,
page_size: int) -> int:
"""Mirror of the device-side page_to_indices."""
block = raw_idx // page_size
offset = raw_idx % page_size
return int(page_table[b, block]) * page_size + offset
def _invert_resolved(resolved_idx: int, b: int, page_table: torch.Tensor,
page_size: int) -> int:
"""Find a raw_idx in [0, max_blocks*page_size) such that
_resolve(raw_idx, b) == resolved_idx. Used to translate kernel output
back to raw scores for comparison with torch.topk."""
block = resolved_idx // page_size
offset = resolved_idx % page_size
# Find the row in page_table[b] that holds `block`.
matches = (page_table[b] == block).nonzero(as_tuple=False)
assert matches.numel() == 1, (
f"page_table row {b} is not a permutation: block {block} appears "
f"{matches.numel()} times")
return int(matches.item()) * page_size + offset
def _reference_topk(scores: torch.Tensor, seq_lens: torch.Tensor,
page_table: torch.Tensor, page_size: int) -> list[set[int]]:
"""Per-row reference: page-resolved set of indices that fast_topk_v2
should emit (excluding -1 padding)."""
B, _ = scores.shape
out: list[set[int]] = []
for b in range(B):
sl = int(seq_lens[b])
if sl <= TOPK:
valid = list(range(sl))
else:
row = scores[b, :sl]
_, raw = torch.topk(row, TOPK)
valid = raw.tolist()
out.append({_resolve(i, b, page_table, page_size) for i in valid})
return out
def _check(scores: torch.Tensor, seq_lens: torch.Tensor,
page_table: torch.Tensor, page_size: int) -> None:
metadata = plan_topk_v2(seq_lens)
workspace = scores.new_empty(
(scores.shape[0], workspace_ints_per_batch()), dtype=torch.int32)
indices = fast_topk_v2(scores, seq_lens, page_table, page_size,
metadata=metadata, workspace=workspace)
torch.cuda.synchronize()
expected = _reference_topk(scores, seq_lens, page_table, page_size)
B = scores.shape[0]
for b in range(B):
sl = int(seq_lens[b])
valid_count = min(sl, TOPK)
row = indices[b].tolist()
# Padding region: -1 (only when sl < TOPK).
if sl < TOPK:
assert all(v == -1 for v in row[sl:]), (
f"row {b}: expected -1 padding after position {sl}, got "
f"{row[sl:sl + 8]}")
got = set(row[:valid_count])
assert -1 not in got, f"row {b}: -1 inside valid region (sl={sl})"
assert got == expected[b], (
f"row {b} (sl={sl}, page_size={page_size}): "
f"missing={len(expected[b] - got)} extra={len(got - expected[b])}")
# --------------------------------------------------------------------------
# Skip non-CUDA / non-Hopper-or-later
# --------------------------------------------------------------------------
def _supports_clusters() -> bool:
if not current_platform.is_cuda():
return False
major, _ = torch.cuda.get_device_capability()
# Thread-block clusters / TMA / PDL are sm_90+. sm_120 (consumer
# Blackwell) is missing some of these; skip when we detect it.
return major == 9 or major == 10
pytestmark = pytest.mark.skipif(
not _supports_clusters(),
reason="fast_topk_v2 requires sm_90 (Hopper) or sm_100 (Blackwell DC)",
)
# --------------------------------------------------------------------------
# Path coverage
# --------------------------------------------------------------------------
@pytest.mark.parametrize("seq_lens", [
pytest.param([1], id="trivial_1"),
pytest.param([300], id="trivial_300"),
pytest.param([512], id="trivial_boundary_512"),
pytest.param([513, 600, 100, 511, 512], id="trivial_mix"),
])
def test_trivial_path(seq_lens):
"""sl <= 512: identity-style fill, no radix, no tie-break."""
torch.manual_seed(0)
device = torch.device("cuda")
B = len(seq_lens)
L = max(max(seq_lens), 1024) # round up so stride is multiple of 4
L = (L + 3) & ~3
seq_lens_t = torch.tensor(seq_lens, dtype=torch.int32, device=device)
scores = torch.randn(B, L, dtype=torch.float32, device=device)
page_table = _trivial_page_table(B, _max_blocks_for(L, 1), device)
_check(scores, seq_lens_t, page_table, page_size=1)
@pytest.mark.parametrize("seq_len", [
pytest.param(513, id="just_above_topk"),
pytest.param(2048, id="2k"),
pytest.param(SMALL_1PASS - 1, id="register_1pass_max"),
pytest.param(SMALL_1PASS, id="register_1pass_boundary"),
pytest.param(SMALL_1PASS + 1, id="register_2pass_first"),
pytest.param(SMALL_2PASS - 1, id="register_2pass_max"),
])
def test_register_path(seq_len):
"""Register strategy (small N; both 1- and 2-pass)."""
torch.manual_seed(seq_len)
device = torch.device("cuda")
B = 4
L = (seq_len + 3) & ~3
scores = torch.randn(B, L, dtype=torch.float32, device=device)
seq_lens = torch.full((B,), seq_len, dtype=torch.int32, device=device)
page_table = _trivial_page_table(B, _max_blocks_for(L, 1), device)
_check(scores, seq_lens, page_table, page_size=1)
@pytest.mark.parametrize("seq_len", [
pytest.param(SMALL_2PASS, id="streaming_first"),
pytest.param(40000, id="streaming_40k"),
pytest.param(DEFAULT_CLUSTER_THRESHOLD, id="streaming_at_cluster_thresh"),
])
def test_streaming_path(seq_len):
"""Streaming strategy (medium N). With small batch and seq_len <=
auto-picked cluster_threshold (>= 32K when batch <= 30), the per-row
dispatch routes here."""
torch.manual_seed(seq_len)
device = torch.device("cuda")
B = 4
L = (seq_len + 3) & ~3
scores = torch.randn(B, L, dtype=torch.float32, device=device)
seq_lens = torch.full((B,), seq_len, dtype=torch.int32, device=device)
page_table = _trivial_page_table(B, _max_blocks_for(L, 1), device)
_check(scores, seq_lens, page_table, page_size=1)
@pytest.mark.parametrize("batch_size,seq_len", [
pytest.param(2, 65536, id="cluster_fused_64k"),
pytest.param(NUM_CLUSTERS, 131072, id="cluster_fused_max_batch"),
pytest.param(NUM_CLUSTERS + 1, 65536, id="cluster_two_stage_just_over"),
pytest.param(32, 96000, id="cluster_two_stage_32x96k"),
])
def test_cluster_path(batch_size, seq_len):
"""Large strategy (Hopper thread-block clusters). Force seq_len above
the auto threshold by passing static_cluster_threshold=SMALL_2PASS."""
torch.manual_seed(seq_len * batch_size)
device = torch.device("cuda")
L = (seq_len + 3) & ~3
scores = torch.randn(batch_size, L, dtype=torch.float32, device=device)
seq_lens = torch.full((batch_size,), seq_len, dtype=torch.int32,
device=device)
page_table = _trivial_page_table(batch_size, _max_blocks_for(L, 1), device)
metadata = plan_topk_v2(seq_lens, static_cluster_threshold=SMALL_2PASS)
indices = fast_topk_v2(scores, seq_lens, page_table, page_size=1,
metadata=metadata)
torch.cuda.synchronize()
expected = _reference_topk(scores, seq_lens, page_table, page_size=1)
for b in range(batch_size):
got = set(indices[b].tolist())
assert got == expected[b], (
f"row {b}: missing={len(expected[b] - got)} "
f"extra={len(got - expected[b])}")
@pytest.mark.parametrize("page_size", [1, 32, 64])
def test_page_table_fold_in(page_size):
"""page_to_indices: kernel-side fold of the page-table gather."""
torch.manual_seed(page_size)
device = torch.device("cuda")
B, seq_len = 4, 6000
L = (seq_len + 3) & ~3
max_blocks = (L + page_size - 1) // page_size
scores = torch.randn(B, L, dtype=torch.float32, device=device)
seq_lens = torch.full((B,), seq_len, dtype=torch.int32, device=device)
page_table = _shuffled_page_table(B, max_blocks, seed=page_size,
device=device)
_check(scores, seq_lens, page_table, page_size=page_size)
def test_mixed_lengths_route_per_row():
"""Per-row dispatch in topk_combine_transform: trivial / Register /
Streaming / Cluster all in one batch. Use static_cluster_threshold to
force a mix that includes the Large path."""
torch.manual_seed(7)
device = torch.device("cuda")
seq_lens = [
100, # trivial
SMALL_1PASS - 100, # 1-pass register
SMALL_2PASS - 100, # 2-pass register
50000, # streaming
40000, # streaming
80000, # cluster (above static_cluster_threshold)
]
B = len(seq_lens)
L = (max(seq_lens) + 3) & ~3
scores = torch.randn(B, L, dtype=torch.float32, device=device)
seq_lens_t = torch.tensor(seq_lens, dtype=torch.int32, device=device)
page_table = _trivial_page_table(B, _max_blocks_for(L, 1), device)
# Force seq_len > 49152 to take the Cluster path.
metadata = plan_topk_v2(seq_lens_t, static_cluster_threshold=49152)
indices = fast_topk_v2(scores, seq_lens_t, page_table, page_size=1,
metadata=metadata)
torch.cuda.synchronize()
expected = _reference_topk(scores, seq_lens_t, page_table, page_size=1)
for b, sl in enumerate(seq_lens):
valid = min(sl, TOPK)
row = indices[b].tolist()
if sl < TOPK:
assert all(v == -1 for v in row[sl:])
got = set(row[:valid])
assert got == expected[b], f"row {b} (sl={sl}) mismatched"
def test_metadata_can_be_reused_across_calls():
"""plan_topk_v2 is amortizable: same metadata reused across calls."""
torch.manual_seed(123)
device = torch.device("cuda")
B, seq_len = 8, 4096
L = (seq_len + 3) & ~3
seq_lens = torch.full((B,), seq_len, dtype=torch.int32, device=device)
page_table = _trivial_page_table(B, _max_blocks_for(L, 1), device)
metadata = plan_topk_v2(seq_lens)
# Two independent score buffers, same metadata.
scores_a = torch.randn(B, L, dtype=torch.float32, device=device)
scores_b = torch.randn(B, L, dtype=torch.float32, device=device)
out_a = fast_topk_v2(scores_a, seq_lens, page_table, page_size=1,
metadata=metadata)
out_b = fast_topk_v2(scores_b, seq_lens, page_table, page_size=1,
metadata=metadata)
torch.cuda.synchronize()
expected_a = _reference_topk(scores_a, seq_lens, page_table, page_size=1)
expected_b = _reference_topk(scores_b, seq_lens, page_table, page_size=1)
for b in range(B):
assert set(out_a[b].tolist()) == expected_a[b]
assert set(out_b[b].tolist()) == expected_b[b]
# --------------------------------------------------------------------------
# sparse_attn_indexer integration: parity with persistent_topk on the V4
# indexer decode shapes. This is the contract the wire-up depends on — the
# kernel must produce the same top-512 set as the existing path.
# --------------------------------------------------------------------------
@pytest.mark.parametrize("config", [
# (B, next_n, L, label). L is max compressed seq_len. Bounded above by
# max_model_len/compress_ratio: ~1024 for C128A, ~32768 for C4A.
pytest.param((1, 1, 1024), id="c128a_short"),
pytest.param((8, 1, 1024), id="c128a_b8"),
pytest.param((16, 1, 1024), id="c128a_b16"),
pytest.param((32, 1, 1024), id="c128a_b32"),
pytest.param((1, 1, 32768), id="c4a_long"),
pytest.param((8, 1, 32768), id="c4a_b8"),
pytest.param((4, 4, 4096), id="c4a_native_mtp"), # 2D seq_lens
])
def test_indexer_dispatch_matches_persistent_topk(config):
"""The dispatch path the indexer takes for V4 (plan once + raw kernel)
must produce the same top-512 set as the fallback persistent_topk on
every shape the V4 decode path actually feeds it."""
from vllm.model_executor.layers.sparse_attn_indexer import (
RADIX_TOPK_WORKSPACE_SIZE, _can_use_fast_topk_v2,
)
from vllm.v1.worker.workspace import (
current_workspace_manager,
init_workspace_manager,
is_workspace_manager_initialized,
)
if not _can_use_fast_topk_v2(512):
pytest.skip("fast_topk_v2 not callable in this environment")
device = torch.device("cuda")
if not is_workspace_manager_initialized():
init_workspace_manager(device=device, num_ubatches=1)
wsm = current_workspace_manager()
B, next_n, L = config
num_rows = B * next_n
L_aligned = (L + 3) & ~3
torch.manual_seed(B * next_n * L)
logits = torch.randn(num_rows, L_aligned, dtype=torch.float32,
device=device)
seq_lens_2d = torch.randint(1, L + 1, (B, next_n), dtype=torch.int32,
device=device)
# Mirror the production flow exactly: plan once into a per-call buffer
# (the indexer dispatch stashes this on attn_metadata), then call the
# raw kernel with that planned metadata.
out_v2 = torch.full((num_rows, TOPK), -1, dtype=torch.int32, device=device)
seq_lens_flat = seq_lens_2d.reshape(-1)
metadata = plan_topk_v2(seq_lens_flat)
(workspace,) = wsm.get_simultaneous(
((num_rows, workspace_ints_per_batch()), torch.int32),
)
fast_topk_v2_raw(
logits, seq_lens_flat,
metadata=metadata, workspace=workspace, topk_indices=out_v2,
)
out_ref = torch.full((num_rows, TOPK), -1, dtype=torch.int32, device=device)
(ref_workspace,) = wsm.get_simultaneous(
((RADIX_TOPK_WORKSPACE_SIZE,), torch.uint8))
torch.ops._C.persistent_topk(logits, seq_lens_2d, out_ref, ref_workspace,
TOPK, L_aligned)
torch.cuda.synchronize()
flat_seq_lens = seq_lens_2d.reshape(-1)
for r in range(num_rows):
sl = int(flat_seq_lens[r])
valid = min(sl, TOPK)
v2 = set(out_v2[r, :valid].tolist()) - {-1}
ref = set(out_ref[r, :valid].tolist()) - {-1}
assert v2 == ref, (
f"row {r} sl={sl}: v2 has {len(v2 - ref)} not in ref, "
f"ref has {len(ref - v2)} not in v2")
if sl < TOPK:
assert (out_v2[r, sl:] == -1).all(), f"row {r}: pad violated"
def test_workspace_can_be_preallocated():
"""Workspace passed in by the caller (cudagraph-friendly path)."""
torch.manual_seed(0)
device = torch.device("cuda")
B, seq_len = 16, 70000
L = (seq_len + 3) & ~3
seq_lens = torch.full((B,), seq_len, dtype=torch.int32, device=device)
page_table = _trivial_page_table(B, _max_blocks_for(L, 1), device)
scores = torch.randn(B, L, dtype=torch.float32, device=device)
metadata = plan_topk_v2(seq_lens, static_cluster_threshold=SMALL_2PASS)
workspace = scores.new_empty((B, workspace_ints_per_batch()),
dtype=torch.int32)
page_indices = scores.new_empty((B, TOPK), dtype=torch.int32)
out = fast_topk_v2(scores, seq_lens, page_table, page_size=1,
metadata=metadata, workspace=workspace,
page_indices=page_indices)
torch.cuda.synchronize()
assert out.data_ptr() == page_indices.data_ptr(), (
"kernel must write into the caller-supplied page_indices tensor")
expected = _reference_topk(scores, seq_lens, page_table, page_size=1)
for b in range(B):
assert set(out[b].tolist()) == expected[b]
# --------------------------------------------------------------------------
# Raw output path (no page-table fold-in). Same selection algorithm; just
# emits row-local raw indices straight to the output. Used by
# sparse_attn_indexer.py as a drop-in for persistent_topk.
# --------------------------------------------------------------------------
def _reference_topk_raw(scores, seq_lens):
"""Per-row reference: row-local raw top-k indices, no page resolution."""
B = scores.shape[0]
out = []
for b in range(B):
sl = int(seq_lens[b])
if sl <= TOPK:
out.append(set(range(sl)))
else:
_, raw = torch.topk(scores[b, :sl], TOPK)
out.append(set(raw.tolist()))
return out
@pytest.mark.parametrize("seq_len", [
pytest.param(300, id="trivial"),
pytest.param(2048, id="register_1p"),
pytest.param(SMALL_2PASS - 1, id="register_2p"),
pytest.param(40000, id="streaming"),
])
def test_raw_path_simple_shapes(seq_len):
"""fast_topk_v2_raw on simple paths."""
torch.manual_seed(seq_len)
device = torch.device("cuda")
B = 4
L = (seq_len + 3) & ~3
scores = torch.randn(B, L, dtype=torch.float32, device=device)
seq_lens = torch.full((B,), seq_len, dtype=torch.int32, device=device)
indices = fast_topk_v2_raw(scores, seq_lens)
torch.cuda.synchronize()
expected = _reference_topk_raw(scores, seq_lens)
for b in range(B):
sl = int(seq_lens[b])
valid = min(sl, TOPK)
row = indices[b].tolist()
if sl < TOPK:
assert all(v == -1 for v in row[sl:])
got = set(row[:valid]) - {-1}
assert got == expected[b], (
f"row {b} sl={sl}: missing={len(expected[b] - got)} "
f"extra={len(got - expected[b])}")
def test_raw_path_matches_paged_with_identity_table():
"""Cross-check: the kernel's two output modes (raw and paged) must
agree on the selected top-k set. With ``page_size=1`` and an identity
page_table, ``page_to_indices`` reduces to the identity, so
``fast_topk_v2_raw`` and ``fast_topk_v2`` should pick the same indices.
Guards against the ``if constexpr (kRawOutput)`` branch in the kernel
drifting from the paged code path."""
torch.manual_seed(0)
device = torch.device("cuda")
B, seq_len = 8, 8192
L = (seq_len + 3) & ~3
scores = torch.randn(B, L, dtype=torch.float32, device=device)
seq_lens = torch.full((B,), seq_len, dtype=torch.int32, device=device)
# Raw path
raw_out = fast_topk_v2_raw(scores, seq_lens)
# Paged path with page_size=1 + identity table
identity_pt = (torch.arange(L, dtype=torch.int32, device=device)
.unsqueeze(0).expand(B, L))
paged_out = fast_topk_v2(scores, seq_lens, identity_pt, page_size=1)
torch.cuda.synchronize()
# Per-row sets should match (top-k order may differ).
for b in range(B):
assert set(raw_out[b].tolist()) == set(paged_out[b].tolist()), (
f"row {b}: raw and paged-with-identity emitted different sets")
# --------------------------------------------------------------------------
# k=1024 (V4-Pro). The kernel templates K so all the same dispatch paths
# (Register / Streaming / Cluster) apply at this K too — these tests just
# repeat the trivial / register / streaming / cluster coverage with K=1024
# and verify parity against torch.topk and persistent_topk.
# --------------------------------------------------------------------------
K_PRO = 1024
def _reference_topk_raw_k(scores, seq_lens, k):
B = scores.shape[0]
out = []
for b in range(B):
sl = int(seq_lens[b])
if sl <= k:
out.append(set(range(sl)))
else:
_, raw = torch.topk(scores[b, :sl], k)
out.append(set(raw.tolist()))
return out
@pytest.mark.parametrize("seq_len", [
pytest.param(700, id="trivial"), # sl <= K=1024
pytest.param(1024, id="trivial_boundary"), # sl == K
pytest.param(1025, id="register_just_above_k"),
pytest.param(8192, id="register_1p"),
pytest.param(SMALL_2PASS - 1, id="register_2p"),
pytest.param(40000, id="streaming"),
])
def test_pro_simple_paths(seq_len):
"""k=1024 across trivial / register / streaming."""
torch.manual_seed(seq_len)
device = torch.device("cuda")
B = 4
L = (seq_len + 3) & ~3
scores = torch.randn(B, L, dtype=torch.float32, device=device)
seq_lens = torch.full((B,), seq_len, dtype=torch.int32, device=device)
indices = fast_topk_v2_raw(scores, seq_lens, topk=K_PRO)
torch.cuda.synchronize()
expected = _reference_topk_raw_k(scores, seq_lens, K_PRO)
for b in range(B):
sl = int(seq_lens[b])
valid = min(sl, K_PRO)
row = indices[b].tolist()
if sl < K_PRO:
assert all(v == -1 for v in row[sl:])
got = set(row[:valid]) - {-1}
assert got == expected[b], (
f"row {b} sl={sl}: missing={len(expected[b] - got)} "
f"extra={len(got - expected[b])}")
@pytest.mark.parametrize("batch_size,seq_len", [
pytest.param(2, 65536, id="cluster_fused_64k"),
pytest.param(NUM_CLUSTERS + 1, 65536, id="cluster_two_stage_just_over"),
pytest.param(32, 96000, id="cluster_two_stage_32x96k"),
])
def test_pro_cluster_path(batch_size, seq_len):
"""k=1024 across the cluster paths (fused and two-stage)."""
torch.manual_seed(seq_len * batch_size)
device = torch.device("cuda")
L = (seq_len + 3) & ~3
scores = torch.randn(batch_size, L, dtype=torch.float32, device=device)
seq_lens = torch.full((batch_size,), seq_len, dtype=torch.int32,
device=device)
metadata = plan_topk_v2(seq_lens, static_cluster_threshold=SMALL_2PASS)
indices = fast_topk_v2_raw(scores, seq_lens, topk=K_PRO,
metadata=metadata)
torch.cuda.synchronize()
expected = _reference_topk_raw_k(scores, seq_lens, K_PRO)
for b in range(batch_size):
got = set(indices[b].tolist()) - {-1}
assert got == expected[b], (
f"row {b}: missing={len(expected[b] - got)} "
f"extra={len(got - expected[b])}")
@pytest.mark.parametrize("config", [
pytest.param((1, 1, 1024), id="pro_short"),
pytest.param((8, 1, 8192), id="pro_register"),
pytest.param((4, 4, 4096), id="pro_native_mtp"), # 2D seq_lens
pytest.param((16, 1, 32768), id="pro_register_2pass"),
pytest.param((32, 1, 50000), id="pro_streaming"),
])
def test_pro_dispatch_matches_persistent_topk(config):
"""k=1024 parity against persistent_topk on V4-Pro decode shapes."""
from vllm.model_executor.layers.sparse_attn_indexer import (
RADIX_TOPK_WORKSPACE_SIZE, _can_use_fast_topk_v2,
)
from vllm.v1.worker.workspace import (
current_workspace_manager,
init_workspace_manager,
is_workspace_manager_initialized,
)
if not _can_use_fast_topk_v2(K_PRO):
pytest.skip("fast_topk_v2 not callable in this environment")
device = torch.device("cuda")
if not is_workspace_manager_initialized():
init_workspace_manager(device=device, num_ubatches=1)
wsm = current_workspace_manager()
B, next_n, L = config
num_rows = B * next_n
L_aligned = (L + 3) & ~3
torch.manual_seed(B * next_n * L)
logits = torch.randn(num_rows, L_aligned, dtype=torch.float32,
device=device)
seq_lens_2d = torch.randint(1, L + 1, (B, next_n), dtype=torch.int32,
device=device)
seq_lens_flat = seq_lens_2d.reshape(-1)
out_v2 = torch.full((num_rows, K_PRO), -1, dtype=torch.int32,
device=device)
metadata = plan_topk_v2(seq_lens_flat)
(workspace,) = wsm.get_simultaneous(
((num_rows, workspace_ints_per_batch()), torch.int32),
)
fast_topk_v2_raw(
logits, seq_lens_flat, topk=K_PRO,
metadata=metadata, workspace=workspace, topk_indices=out_v2,
)
out_ref = torch.full((num_rows, K_PRO), -1, dtype=torch.int32,
device=device)
(ref_workspace,) = wsm.get_simultaneous(
((RADIX_TOPK_WORKSPACE_SIZE,), torch.uint8))
torch.ops._C.persistent_topk(logits, seq_lens_2d, out_ref, ref_workspace,
K_PRO, L_aligned)
torch.cuda.synchronize()
for r in range(num_rows):
sl = int(seq_lens_flat[r])
valid = min(sl, K_PRO)
v2 = set(out_v2[r, :valid].tolist()) - {-1}
ref = set(out_ref[r, :valid].tolist()) - {-1}
assert v2 == ref, (
f"row {r} sl={sl}: v2 has {len(v2 - ref)} not in ref, "
f"ref has {len(ref - v2)} not in v2")
if sl < K_PRO:
assert (out_v2[r, sl:] == -1).all(), f"row {r}: pad violated"
+35 -2
View File
@@ -260,7 +260,9 @@ _TEXT_GENERATION_EXAMPLE_MODELS = {
trust_remote_code=True,
),
"DeepseekV32ForCausalLM": _HfExamplesInfo("deepseek-ai/DeepSeek-V3.2-Exp"),
"DeepseekV4ForCausalLM": _HfExamplesInfo("deepseek-ai/DeepSeek-V4-Flash"),
"DeepseekV4ForCausalLM": _HfExamplesInfo(
"deepseek-ai/DeepSeek-V4-Flash", is_available_online=False
),
"Ernie4_5ForCausalLM": _HfExamplesInfo("baidu/ERNIE-4.5-0.3B-PT"),
"Ernie4_5_MoeForCausalLM": _HfExamplesInfo("baidu/ERNIE-4.5-21B-A3B-PT"),
"ExaoneForCausalLM": _HfExamplesInfo(
@@ -592,6 +594,9 @@ _TEXT_GENERATION_EXAMPLE_MODELS = {
"MiMoV2FlashForCausalLM": _HfExamplesInfo(
"XiaomiMiMo/MiMo-V2-Flash", trust_remote_code=True
),
"MiMoV2ProForCausalLM": _HfExamplesInfo(
"XiaomiMiMo/MiMo-V2.5-Pro", trust_remote_code=True, is_available_online=False
),
"Dots1ForCausalLM": _HfExamplesInfo("rednote-hilab/dots.llm1.inst"),
}
@@ -959,6 +964,18 @@ _MULTIMODAL_EXAMPLE_MODELS = {
"PerceptronAI/Isaac-0.1",
trust_remote_code=True,
extras={"0.2-2B-Preview": "PerceptronAI/Isaac-0.2-2B-Preview"},
max_transformers_version="4.57",
transformers_version_reason={
"vllm": (
"Custom Isaac code is not compatible with Transformers v5. "
"The model should be upstreamed to Transformers for "
"long-term support."
),
"hf": (
"Isaac's remote model and processor code import or configure "
"APIs that changed in Transformers v5."
),
},
),
"InternS1ForConditionalGeneration": _HfExamplesInfo(
"internlm/Intern-S1",
@@ -1055,6 +1072,9 @@ _MULTIMODAL_EXAMPLE_MODELS = {
"MiDashengLMModel": _HfExamplesInfo(
"mispeech/midashenglm-7b", trust_remote_code=True
),
"MiMoV2OmniForCausalLM": _HfExamplesInfo(
"XiaomiMiMo/MiMo-V2.5-Omni", trust_remote_code=True, is_available_online=False
),
"MiniCPMO": _HfExamplesInfo(
"openbmb/MiniCPM-o-2_6",
trust_remote_code=True,
@@ -1483,10 +1503,11 @@ _SPECULATIVE_DECODING_EXAMPLE_MODELS = {
speculative_model="luccafong/deepseek_mtp_draft_random",
trust_remote_code=True,
),
"DeepSeekV4MTP": _HfExamplesInfo(
"DeepSeekV4MTPModel": _HfExamplesInfo(
"deepseek-ai/DeepSeek-V4-Flash",
speculative_model="deepseek-ai/DeepSeek-V4-Flash",
trust_remote_code=True,
is_available_online=False,
),
"ErnieMTPModel": _HfExamplesInfo(
"baidu/ERNIE-4.5-21B-A3B-PT",
@@ -1537,6 +1558,18 @@ _SPECULATIVE_DECODING_EXAMPLE_MODELS = {
trust_remote_code=True,
speculative_model="XiaomiMiMo/MiMo-7B-RL",
),
"MiMoV2MTPModel": _HfExamplesInfo(
"XiaomiMiMo/MiMo-V2.5-Pro",
trust_remote_code=True,
speculative_model="XiaomiMiMo/MiMo-V2.5-Pro",
is_available_online=False,
),
"MiMoV2OmniMTPModel": _HfExamplesInfo(
"XiaomiMiMo/MiMo-V2.5-Omni",
trust_remote_code=True,
speculative_model="XiaomiMiMo/MiMo-V2.5-Omni",
is_available_online=False,
),
"NemotronHMTPModel": _HfExamplesInfo(
"nvidia/Nemotron-Super-Placeholder",
speculative_model="nvidia/Nemotron-Super-Placeholder",
+2 -1
View File
@@ -5,7 +5,6 @@ from types import SimpleNamespace
import pytest
import torch
from vllm.third_party.deep_gemm.utils import per_token_cast_to_fp8
from vllm.model_executor.models.deepseek_v4 import (
DeepseekV4MegaMoEExperts,
@@ -112,6 +111,8 @@ def test_deepseek_v4_mega_moe_weight_loader_uses_ep_expert_ownership():
reason="DeepSeek V4 MegaMoE fused input staging requires CUDA.",
)
def test_deepseek_v4_mega_moe_fused_input_staging_is_bitwise_exact():
from vllm.third_party.deep_gemm.utils import per_token_cast_to_fp8
device = torch.device("cuda")
num_tokens = 7
hidden_size = 256
+16
View File
@@ -5,6 +5,8 @@ Unit tests for MultiModalRegistry.supports_multimodal_inputs and
Qwen2.5-VL visual component loading behavior.
"""
from types import SimpleNamespace
import pytest
from vllm.multimodal import MULTIMODAL_REGISTRY
@@ -32,3 +34,17 @@ def test_supports_multimodal_inputs(model_id, limit_mm_per_prompt, expected):
limit_mm_per_prompt=limit_mm_per_prompt,
)
assert MULTIMODAL_REGISTRY.supports_multimodal_inputs(ctx.model_config) is expected
def test_create_processor_error_uses_served_model_name():
model_config = SimpleNamespace(
is_multimodal_model=False,
model="/path/to/model/weights",
served_model_name="friendly-model-name",
)
with pytest.raises(
ValueError,
match="friendly-model-name is not a multimodal model",
):
MULTIMODAL_REGISTRY.create_processor(model_config)
+185
View File
@@ -0,0 +1,185 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
"""Tests for Cutlass W4A16 (Machete) kernel on Hopper.
Verifies that W4A16 quantized models loaded through vllm select the
MacheteLinearKernel on sm_90 GPUs, that weights are correctly repacked,
and that inference produces valid output.
Run `pytest tests/quantization/test_cutlass_w4a16.py`.
"""
import pytest
import torch
from vllm.platforms import current_platform
if not current_platform.has_device_capability(90):
pytest.skip(
"Machete W4A16 requires Hopper (sm_90).",
allow_module_level=True,
)
from vllm.model_executor.kernels.linear import (
MPLinearLayerConfig,
choose_mp_linear_kernel,
)
from vllm.model_executor.kernels.linear.mixed_precision import (
MacheteLinearKernel,
)
from vllm.model_executor.layers.quantization.compressed_tensors.compressed_tensors import ( # noqa: E501
CompressedTensorsLinearMethod,
CompressedTensorsWNA16,
)
from vllm.scalar_type import scalar_types
@pytest.fixture(scope="function", autouse=True)
def enable_pickle(monkeypatch):
"""`LLM.apply_model` requires pickling a function."""
monkeypatch.setenv("VLLM_ALLOW_INSECURE_SERIALIZATION", "1")
@pytest.mark.parametrize(
"act_type,weight_type,group_size,zero_points",
[
(torch.float16, scalar_types.uint4b8, 128, False),
(torch.bfloat16, scalar_types.uint4b8, 128, False),
(torch.float16, scalar_types.uint4, 128, True),
(torch.float16, scalar_types.uint4b8, -1, False),
],
ids=[
"fp16-gptq-g128",
"bf16-gptq-g128",
"fp16-awq-g128",
"fp16-channelwise",
],
)
def test_machete_kernel_selected(act_type, weight_type, group_size, zero_points):
"""Verify choose_mp_linear_kernel picks MacheteLinearKernel."""
config = MPLinearLayerConfig(
full_weight_shape=(4096, 4096),
partition_weight_shape=(4096, 4096),
act_type=act_type,
weight_type=weight_type,
group_size=group_size,
zero_points=zero_points,
has_g_idx=False,
)
kernel = choose_mp_linear_kernel(config)
assert kernel is MacheteLinearKernel, (
f"Expected MacheteLinearKernel, got {kernel.__name__}"
)
@pytest.mark.parametrize(
"full_shape,part_shape,weight_type,group_size,has_g_idx,expected_reason",
[
((4096, 4096), (2048, 4096), scalar_types.uint4b8, 128, True, "Act reordering"),
(
(4096, 4096),
(4096, 4096),
scalar_types.float6_e3m2f,
128,
False,
"Quant type",
),
((4096, 4096), (4096, 4096), scalar_types.uint4b8, 32, False, "Group size"),
],
ids=["partitioned-g_idx", "unsupported-quant-type", "unsupported-group-size"],
)
def test_machete_rejects_invalid_config(
full_shape, part_shape, weight_type, group_size, has_g_idx, expected_reason
):
"""Verify Machete rejects unsupported configurations."""
config = MPLinearLayerConfig(
full_weight_shape=full_shape,
partition_weight_shape=part_shape,
act_type=torch.float16,
weight_type=weight_type,
group_size=group_size,
zero_points=False,
has_g_idx=has_g_idx,
)
can_impl, reason = MacheteLinearKernel.can_implement(config)
assert not can_impl
assert expected_reason in reason
def test_kernel_selection_with_disabled_machete(monkeypatch):
"""Verify kernel selection falls back when Machete is disabled."""
monkeypatch.setattr("vllm.envs.VLLM_DISABLED_KERNELS", ["MacheteLinearKernel"])
config = MPLinearLayerConfig(
full_weight_shape=(4096, 4096),
partition_weight_shape=(4096, 4096),
act_type=torch.float16,
weight_type=scalar_types.uint4b8,
group_size=128,
zero_points=False,
has_g_idx=False,
)
kernel = choose_mp_linear_kernel(config)
assert kernel is not MacheteLinearKernel, "MacheteLinearKernel should be disabled"
@pytest.mark.parametrize(
"model_name",
[
"nm-testing/tinyllama-oneshot-w4a16-channel-v2",
"nm-testing/TinyLlama-1.1B-Chat-v1.0-W4A16-G128-Asym-Updated-ActOrder",
],
)
def test_w4a16_machete_e2e(vllm_runner, model_name):
"""Load a W4A16 model, verify Machete kernel is used, and generate."""
with vllm_runner(model_name, enforce_eager=True, gpu_memory_utilization=0.5) as llm:
def check_model(model):
layer = model.model.layers[0]
qkv_proj = layer.self_attn.qkv_proj
assert isinstance(qkv_proj.quant_method, CompressedTensorsLinearMethod)
assert isinstance(qkv_proj.scheme, CompressedTensorsWNA16)
assert isinstance(qkv_proj.scheme.kernel, MacheteLinearKernel), (
f"Expected MacheteLinearKernel on Hopper, "
f"got {type(qkv_proj.scheme.kernel).__name__}"
)
assert hasattr(qkv_proj, "weight_packed")
assert hasattr(qkv_proj, "weight_scale")
assert qkv_proj.weight_packed.dtype == torch.int32
llm.apply_model(check_model)
output = llm.generate_greedy("Hello my name is", max_tokens=10)
assert output
assert len(output[0][1]) > 0
def test_w4a16_machete_bfloat16_deterministic(vllm_runner):
"""Verify Machete works with bf16 activations and is deterministic."""
model_name = "nm-testing/tinyllama-oneshot-w4a16-channel-v2"
prompt = "The capital of France is"
with vllm_runner(
model_name,
enforce_eager=True,
dtype="bfloat16",
gpu_memory_utilization=0.5,
) as llm:
def check_kernel_type(model):
layer = model.model.layers[0]
scheme = layer.self_attn.qkv_proj.scheme
assert isinstance(scheme.kernel, MacheteLinearKernel), (
f"Expected MacheteLinearKernel with bf16, "
f"got {type(scheme.kernel).__name__}"
)
llm.apply_model(check_kernel_type)
out1 = llm.generate_greedy(prompt, max_tokens=10)
out2 = llm.generate_greedy(prompt, max_tokens=10)
assert out1[0][1] == out2[0][1], (
f"Non-deterministic: '{out1[0][1]}' vs '{out2[0][1]}'"
)
+108
View File
@@ -2512,3 +2512,111 @@ def test_block_lookup_cache_multi_blocks_per_key():
assert cache.pop(key1, 11) is block11
assert cache.get_one_block(key1) is None
assert cache.pop(key1, 12) is None
def test_can_fit_full_sequence_swa_cap_admits_long_prompt():
"""Hybrid full+SWA model with a pool sized at the startup minimum should
admit a prompt longer than the SWA cap, because SlidingWindowManager
recycles blocks during chunked prefill (issue #39734)."""
block_size = 16
sliding_window = 4 * block_size # 64 tokens
max_num_batched_tokens = 8 * block_size # 128 tokens
max_model_len = 64 * block_size # 1024 tokens — much larger than the SWA cap
# Startup pool sizing: full demands cdiv(max_model_len, bs) = 64 blocks,
# SWA demands cdiv(SW-1+max_batched, bs) + 1 = cdiv(191, 16) + 1 = 13.
# Pool minimum = 64 + 13 = 77; +1 for the null block.
num_blocks = 64 + 13 + 1
config = KVCacheConfig(
num_blocks=num_blocks,
kv_cache_tensors=[],
kv_cache_groups=[
KVCacheGroupSpec(
["layer_full"],
FullAttentionSpec(
block_size=block_size,
num_kv_heads=1,
head_size=1,
dtype=torch.float32,
),
),
KVCacheGroupSpec(
["layer_swa"],
SlidingWindowSpec(
block_size=block_size,
num_kv_heads=1,
head_size=1,
dtype=torch.float32,
sliding_window=sliding_window,
),
),
],
)
manager = KVCacheManager(
config,
max_model_len=max_model_len,
max_num_batched_tokens=max_num_batched_tokens,
enable_caching=True,
hash_block_size=block_size,
)
# A prompt that is shorter than max_model_len but longer than SW + chunk:
# cdiv(prompt_len, bs) = 32 blocks. Without the cap, admission would
# demand 32 (full) + 32 (SWA) = 64 blocks. With the cap, SWA contributes
# only 13, so total = 32 + 13 = 45 ≤ pool size.
prompt_len = 32 * block_size
req = make_request("long", list(range(prompt_len)), block_size, sha256)
assert manager.can_fit_full_sequence(req)
def test_can_fit_full_sequence_full_attention_still_gates_oversized():
"""The cap only loosens the SWA group; a prompt that exceeds the
full-attention pool capacity must still be rejected."""
block_size = 16
sliding_window = 4 * block_size
max_num_batched_tokens = 8 * block_size
max_model_len = 64 * block_size
# Provide a tiny pool — even a small prompt should be rejected.
num_blocks = 5
config = KVCacheConfig(
num_blocks=num_blocks,
kv_cache_tensors=[],
kv_cache_groups=[
KVCacheGroupSpec(
["layer_full"],
FullAttentionSpec(
block_size=block_size,
num_kv_heads=1,
head_size=1,
dtype=torch.float32,
),
),
KVCacheGroupSpec(
["layer_swa"],
SlidingWindowSpec(
block_size=block_size,
num_kv_heads=1,
head_size=1,
dtype=torch.float32,
sliding_window=sliding_window,
),
),
],
)
manager = KVCacheManager(
config,
max_model_len=max_model_len,
max_num_batched_tokens=max_num_batched_tokens,
enable_caching=True,
hash_block_size=block_size,
)
# 16 blocks of full attention demand alone exceeds the 5-block pool.
prompt_len = 16 * block_size
req = make_request("oversized", list(range(prompt_len)), block_size, sha256)
assert not manager.can_fit_full_sequence(req)
@@ -22,11 +22,13 @@ pytestmark = pytest.mark.cpu_test
def get_sliding_window_manager(sliding_window_spec, block_pool, enable_caching=True):
# Tests don't exercise admission gating; pass a large cap that is a no-op.
return SlidingWindowManager(
sliding_window_spec,
block_pool=block_pool,
enable_caching=enable_caching,
kv_cache_group_id=0,
max_admission_blocks_per_request=10**9,
)
@@ -38,6 +40,7 @@ def get_chunked_local_attention_manager(
block_pool=block_pool,
enable_caching=enable_caching,
kv_cache_group_id=0,
max_admission_blocks_per_request=10**9,
)
@@ -478,3 +478,59 @@ class TestSlidingWindowLookup:
sched._sliding_window_lookup(to_keys([1, 2, 3, 4]), 2, _EMPTY_REQ_CTX)
is None
)
@pytest.mark.parametrize("async_scheduling", [True, False])
def test_do_remote_decode_stores_all_blocks(request_runner, async_scheduling: bool):
"""With do_remote_decode=True, after loading prefix blocks from CPU,
all blocks must be re-stored not just the newly computed ones.
This supports P/D disaggregation where the prefill instance offloads the
complete KV cache so a remote decode node can consume it."""
offloaded_block_size = 12
gpu_block_size = 4
num_gpu_blocks = 100
runner = request_runner(
offloaded_block_size=offloaded_block_size,
gpu_block_size=gpu_block_size,
num_gpu_blocks=num_gpu_blocks,
async_scheduling=async_scheduling,
)
# Store 1 offloaded block (3 GPU blocks) via a normal request.
runner.new_request(token_ids=[0] * offloaded_block_size)
runner.manager.prepare_store.side_effect = (
lambda keys, req_context: generate_store_output(keys)
)
runner.run(
decoded_tokens=[EOS_TOKEN_ID],
expected_stored_gpu_block_indexes=(0, 1, 2),
)
# Reset GPU prefix cache so the next request must load from CPU.
runner.scheduler.reset_prefix_cache()
# New request with do_remote_decode=True and 2 offloaded blocks.
# The first offloaded block matches what we stored in CPU.
runner.new_request(
token_ids=[0] * offloaded_block_size * 2,
kv_transfer_params={"do_remote_decode": True},
)
runner.connector_scheduler._maximal_prefix_lookup = lambda key, req_context: 1
runner.manager.prepare_store.side_effect = (
lambda keys, req_context: generate_store_output(keys)
)
# Load the first offloaded block from CPU.
runner.run(
decoded_tokens=[0],
expected_loaded_gpu_block_indexes=(0, 1, 2),
)
# Store must include ALL 6 GPU blocks (both the loaded prefix and
# the newly computed block), not just the 3 new ones.
runner.run(
decoded_tokens=[EOS_TOKEN_ID],
expected_stored_gpu_block_indexes=(0, 1, 2, 3, 4, 5),
)
@@ -270,7 +270,11 @@ class RequestRunner:
slot_mapping={},
)
def new_request(self, token_ids: list[int]):
def new_request(
self,
token_ids: list[int],
kv_transfer_params: dict | None = None,
):
self.req_id += 1
sampling_params = SamplingParams(max_tokens=1000)
@@ -283,6 +287,8 @@ class RequestRunner:
pooling_params=None,
block_hasher=self._block_hasher,
)
if kv_transfer_params is not None:
req.kv_transfer_params = kv_transfer_params
self.scheduler.add_request(req)
@@ -406,16 +406,13 @@ class AsyncTPPass(VllmPatternMatcherPass):
self.dump_patterns(config, self.patterns)
def is_applicable_for_range(self, compile_range: Range) -> bool:
# This pass is applied on top of the sequence parallelism pass.
# It inherits the same applicability condition as `SequenceParallelismPass`.
# See `SequenceParallelismPass.is_applicable` for more details.
if (
not self.compilation_config.splitting_ops
or self.compilation_config.use_inductor_graph_partition
):
return True
tp_size = get_tensor_model_parallel_world_size()
return bool(compile_range.is_single_size() and compile_range.end % tp_size == 0)
# This pass is applied on top of the sequence parallelism pass,
# which is only supported in fullgraph compilation mode.
assert (
self.compilation_config.use_inductor_graph_partition
or not self.compilation_config.splitting_ops
), "AsyncTPPass requires full-graph compilation"
return True
@VllmInductorPass.time_and_log
def __call__(self, graph: fx.Graph) -> None:
@@ -341,22 +341,18 @@ class SequenceParallelismPass(VllmPatternMatcherPass):
significantly reduce communication overhead and improve overall model
performance.
This pass is only supported when compiling the whole graph (fullgraph
mode, i.e. using Inductor graph partition or empty splitting_ops).
Piecewise compilation is not supported because the residual tensor
gets split across TP ranks, causing size mismatches at subgraph
boundaries.
This pass splits up the residual tensor across TP ranks and hence divides its size.
Because the pattern matcher starts at the end of the graph, the replacement
contains a slice that temporarily conforms the input residual to the correct size.
After all patterns have been matched, we use a NoOpEliminationPass to clean up
what have now become no-op slices.
Note that an older version of the pass did not need this as it operated only on
custom rms_norm and fused_rms_norm_add custom ops which did not complain about
mismatched shapes during replacement. So this approach has the same assumption that
correctness is only maintained if all rms_norm operations are split across ranks.
Correctness-wise, this is approach strictly better than before - before,
the graph was incorrect semantically and shape-wise during the pass.
With this approach there's only semantic incorrectness during the pass.
Both approaches restore a correct graph once all patterns are matched.
This pass splits up the residual tensor across TP ranks and hence
divides its size. Because the pattern matcher starts at the end of
the graph, the replacement contains a slice that temporarily conforms
the input residual to the correct size. After all patterns have been
matched, we use a NoOpEliminationPass to clean up what have now
become no-op slices.
"""
@enable_fake_mode
@@ -419,19 +415,13 @@ class SequenceParallelismPass(VllmPatternMatcherPass):
and gathering tensors across TP ranks outweighs the benefits.
Returns False (SP disabled) when:
- Using piecewise compilation with non-concrete or TP-indivisible sizes
- min_token_num is None (SP disabled for this device/config)
- The compile range starts below the minimum token threshold
"""
# For piecewise compilation (not using inductor graph partition),
# we need concrete sizes that are divisible by TP for correct splitting
if (
not self.compilation_config.use_inductor_graph_partition
and self.compilation_config.splitting_ops
):
tp_size = get_tensor_model_parallel_world_size()
if not compile_range.is_single_size() or compile_range.end % tp_size != 0:
return False
assert (
self.compilation_config.use_inductor_graph_partition
or not self.compilation_config.splitting_ops
), "SequenceParallelismPass requires full-graph compilation"
# min_token_num is None when SP is disabled for this device/config
# (e.g., non-CUDA platform, unsupported GPU, or small hidden_size)
+19
View File
@@ -1149,6 +1149,25 @@ class CompilationConfig:
self.cudagraph_mode = CUDAGraphMode.FULL
self.splitting_ops = []
if (
not self.use_inductor_graph_partition
and (self.pass_config.enable_sp or self.pass_config.fuse_gemm_comms)
and self.splitting_ops
):
logger.warning_once(
"Sequence parallelism requires full-graph compilation when "
"use_inductor_graph_partition is off. Setting splitting_ops "
"to an empty list to preserve SP and async TP."
)
self.splitting_ops = []
if self.cudagraph_mode.has_piecewise_cudagraphs():
logger.warning_once(
"Sequence parallelism is incompatible with piecewise "
"cudagraph when use_inductor_graph_partition is off. "
"Setting cudagraph_mode to FULL."
)
self.cudagraph_mode = CUDAGraphMode.FULL
# Disable CUDA graphs for DeepEP high-throughput since its not CG compatible
if (
all2all_backend == "deepep_high_throughput"
+6 -6
View File
@@ -50,7 +50,7 @@ class IrOpPriorityConfig:
name: {
provider: IrOp.registry[name].impls[provider].uuid() for provider in p
}
for name, p in asdict(self).items()
for name, p in asdict(self).items() # type: ignore[call-overload]
}
return hash_factors(factors)
@@ -77,7 +77,7 @@ class IrOpPriorityConfig:
current_platform.import_ir_kernels()
with contextlib.ExitStack() as stack:
for field in fields(self):
for field in fields(self): # type: ignore[arg-type]
op_priority = getattr(self, field.name)
assert op_priority is not None, (
f"IR op priority for {field.name} must be set"
@@ -98,7 +98,7 @@ class IrOpPriorityConfig:
A helper to create an IrOpPriorityConfig where fields not specified in kwargs
use the given default list.
"""
for field in fields(cls):
for field in fields(cls): # type: ignore[arg-type]
if field.name not in kwargs:
kwargs[field.name] = list(default)
@@ -108,8 +108,8 @@ class IrOpPriorityConfig:
MoEBackend = Literal[
"auto",
"triton",
"triton_unfused",
"deep_gemm",
"deep_gemm_mega_moe",
"cutlass",
"flashinfer_trtllm",
"flashinfer_cutlass",
@@ -137,9 +137,9 @@ class KernelConfig:
"""Backend for MoE expert computation kernels. Available options:
- "auto": Automatically select the best backend based on model and hardware
- "triton": Use Triton-based fused MoE kernels (SWIGLUOAI activation only)
- "triton_unfused": Use Triton-based unfused MoE kernels (supports SILU/GELU)
- "triton": Use Triton-based fused MoE kernels
- "deep_gemm": Use DeepGEMM kernels (FP8 block-quantized only)
- "deep_gemm_mega_moe": Use DeepGEMM mega MoE kernels
- "cutlass": Use vLLM CUTLASS kernels
- "flashinfer_trtllm": Use FlashInfer with TRTLLM-GEN kernels
- "flashinfer_cutlass": Use FlashInfer with CUTLASS kernels
+13 -1
View File
@@ -326,6 +326,10 @@ class ModelConfig:
mm_encoder_only: InitVar[bool | None] = None
mm_encoder_tp_mode: InitVar[MMEncoderTPMode | None] = None
mm_encoder_attn_backend: InitVar[AttentionBackendEnum | str | None] = None
mm_encoder_attn_dtype: InitVar[str | None] = None
mm_encoder_fp8_scale_path: InitVar[str | None] = None
mm_encoder_fp8_scale_save_path: InitVar[str | None] = None
mm_encoder_fp8_scale_save_margin: InitVar[float | None] = None
interleave_mm_strings: InitVar[bool | None] = None
skip_mm_profiling: InitVar[bool | None] = None
video_pruning_rate: InitVar[float | None] = None
@@ -447,6 +451,10 @@ class ModelConfig:
mm_encoder_only: bool | None,
mm_encoder_tp_mode: MMEncoderTPMode | None,
mm_encoder_attn_backend: AttentionBackendEnum | str | None,
mm_encoder_attn_dtype: str | None,
mm_encoder_fp8_scale_path: str | None,
mm_encoder_fp8_scale_save_path: str | None,
mm_encoder_fp8_scale_save_margin: float | None,
interleave_mm_strings: bool | None,
skip_mm_profiling: bool | None,
video_pruning_rate: float | None,
@@ -513,6 +521,7 @@ class ModelConfig:
if dict_overrides:
self._apply_dict_overrides(hf_config, dict_overrides)
self.hf_text_config = get_hf_text_config(self.hf_config)
self.model_arch_config = self.get_model_arch_config()
self.attention_chunk_size = getattr(
self.hf_text_config, "attention_chunk_size", None
)
@@ -520,7 +529,6 @@ class ModelConfig:
self.hf_image_processor_config = get_hf_image_processor_config(
self.model, hf_token=self.hf_token, revision=self.revision
)
self.model_arch_config = self.get_model_arch_config()
architectures = self.architectures
registry = self.registry
@@ -643,6 +651,10 @@ class ModelConfig:
mm_encoder_only=mm_encoder_only,
mm_encoder_tp_mode=mm_encoder_tp_mode,
mm_encoder_attn_backend=mm_encoder_attn_backend,
mm_encoder_attn_dtype=mm_encoder_attn_dtype,
mm_encoder_fp8_scale_path=mm_encoder_fp8_scale_path,
mm_encoder_fp8_scale_save_path=mm_encoder_fp8_scale_save_path,
mm_encoder_fp8_scale_save_margin=mm_encoder_fp8_scale_save_margin,
interleave_mm_strings=interleave_mm_strings,
skip_mm_profiling=skip_mm_profiling,
video_pruning_rate=video_pruning_rate,
+51
View File
@@ -2,6 +2,7 @@
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
from collections.abc import Mapping
from pathlib import Path
from typing import Any, Literal, TypeAlias, TypedDict, final
from pydantic import ConfigDict, Field, field_validator, model_validator
@@ -158,6 +159,24 @@ class MultiModalConfig:
"""Optional override for the multi-modal encoder attention backend when
using vision transformers. Accepts any value from
`vllm.v1.attention.backends.registry.AttentionBackendEnum` (e.g. `FLASH_ATTN`)."""
mm_encoder_attn_dtype: Literal["fp8"] | None = None
"""Optional dtype override for ViT encoder attention. Set to `"fp8"` to
enable FP8 quantization via the FlashInfer cuDNN backend. When set to
`"fp8"` without a scale file, dynamic scaling is used automatically.
See docs/features/quantization/fp8_vit_attn.md for details."""
mm_encoder_fp8_scale_path: str | None = None
"""Path to a JSON file containing per-layer FP8 Q/K/V scales for ViT
encoder attention. When provided (with `mm_encoder_attn_dtype="fp8"`),
static scaling is used. When omitted, dynamic scaling is used."""
mm_encoder_fp8_scale_save_path: str | None = None
"""When set with dynamic FP8 scaling (`mm_encoder_attn_dtype="fp8"`
and no `mm_encoder_fp8_scale_path`), saves the calibrated scales to
this file after the amax history buffer is full. The saved file can
then be used as `mm_encoder_fp8_scale_path` in subsequent runs."""
mm_encoder_fp8_scale_save_margin: float = Field(default=1.5, gt=0.0)
"""Safety margin multiplied onto scales when auto-saving. A value > 1
leaves headroom so that inputs with larger activations than the
calibration set do not overflow FP8 range. Default 1.5."""
interleave_mm_strings: bool = False
"""Enable fully interleaved support for multimodal prompts, while using
--chat-template-content-format=string."""
@@ -233,6 +252,36 @@ class MultiModalConfig:
"'mm_shm_cache_max_object_size_mb' should only be set when "
"'mm_processor_cache_type' is 'shm'."
)
# Validate FP8 scale path combinations.
if self.mm_encoder_attn_dtype != "fp8" and (
self.mm_encoder_fp8_scale_path is not None
or self.mm_encoder_fp8_scale_save_path is not None
):
raise ValueError(
"'mm_encoder_fp8_scale_path' and "
"'mm_encoder_fp8_scale_save_path' require "
"'mm_encoder_attn_dtype' to be 'fp8'."
)
if (
self.mm_encoder_fp8_scale_path is not None
and self.mm_encoder_fp8_scale_save_path is not None
):
raise ValueError(
"'mm_encoder_fp8_scale_save_path' cannot be used with "
"'mm_encoder_fp8_scale_path' (saving requires dynamic scaling)."
)
# Validate file paths exist.
if self.mm_encoder_fp8_scale_path is not None:
scale_path = Path(self.mm_encoder_fp8_scale_path)
if not scale_path.is_file():
raise FileNotFoundError(f"FP8 scale file not found: {scale_path}")
if self.mm_encoder_fp8_scale_save_path is not None:
save_parent = Path(self.mm_encoder_fp8_scale_save_path).parent
if not save_parent.is_dir():
raise FileNotFoundError(
f"Parent directory for FP8 scale save path not found: {save_parent}"
)
return self
def compute_hash(self) -> str:
@@ -252,6 +301,8 @@ class MultiModalConfig:
if self.mm_encoder_attn_backend is not None
else None,
self.mm_encoder_tp_mode,
self.mm_encoder_attn_dtype,
self.mm_encoder_fp8_scale_path,
]
hash_str = safe_hash(str(factors).encode(), usedforsecurity=False).hexdigest()
return hash_str
+58 -6
View File
@@ -34,6 +34,7 @@ logger = init_logger(__name__)
MTPModelTypes = Literal[
"deepseek_mtp",
"mimo_mtp",
"mimo_v2_mtp",
"glm4_moe_mtp",
"glm4_moe_lite_mtp",
"glm_ocr_mtp",
@@ -63,7 +64,8 @@ SpeculativeMethod = Literal[
EagleModelTypes,
NgramGPUTypes,
]
RejectionSampleMethod = Literal["strict", "probabilistic", "synthetic"]
RejectionSampleMethod = Literal["standard", "synthetic"]
DraftSampleMethod = Literal["greedy", "gumbel"]
@config
@@ -183,11 +185,11 @@ class SpeculativeConfig:
"""Load config for the draft model. If not specified, will use the load
config from the target model."""
rejection_sample_method: RejectionSampleMethod = "strict"
"""Whether to use strict (target and draft sampled tokens match exactly)
or probabilistic rejection sampling. Both respect the target model
distribution, but the latter yields a higher acceptance rate at the cost
of more memory to cache draft logits."""
rejection_sample_method: RejectionSampleMethod = "standard"
"""The rejection sampling method to use. 'standard' uses probabilistic
rejection sampling (with or without cached draft logits, controlled by
draft_sample_method). 'synthetic' accepts draft tokens with a decaying
probability calibrated to synthetic_acceptance_rate."""
synthetic_acceptance_rates: list[float] | None = None
"""Per-position *unconditional* acceptance rates for synthetic rejection
@@ -248,6 +250,14 @@ class SpeculativeConfig:
)
return SpeculativeConfig._acceptance_length_to_rates(length, n)
draft_sample_method: DraftSampleMethod = "greedy"
"""How the draft model samples tokens. 'greedy' always picks the argmax
token, and the draft probabilities are treated as one-hot during rejection
sampling. 'gumbel' adds Gumbel noise for stochastic sampling, and the full
draft logits are used for the probability ratio test during rejection
sampling. This comes at the cost of additional GPU memory usage. This
parameter currently only applies to Model Runner V2."""
def compute_hash(self) -> str:
"""
WARNING: Whenever a new field is added to this config,
@@ -323,6 +333,48 @@ class SpeculativeConfig:
}
)
if (arch := hf_config.architectures[0]) in (
"MiMoV2ProForCausalLM",
"MiMoV2OmniForCausalLM",
):
from vllm.model_executor.models.mimo_v2_mtp import (
_MIMO_V2_PRO_NUM_MTP_LAYERS,
)
mtp_arch_maps = {
"MiMoV2ProForCausalLM": "MiMoV2MTPModel",
"MiMoV2OmniForCausalLM": "MiMoV2OmniMTPModel",
}
hf_config.model_type = "mimo_v2_mtp"
# vLLM currently supports only the first MiMo-V2 MTP layer.
n_predict = _MIMO_V2_PRO_NUM_MTP_LAYERS
hf_config.update(
{
"num_hidden_layers": 0,
"n_predict": n_predict,
"num_nextn_predict_layers": n_predict,
"architectures": [mtp_arch_maps[arch]],
}
)
if hf_config.architectures[0] == "MiMoV2FlashForCausalLM":
from vllm.model_executor.models.mimo_v2_mtp import (
_MIMO_V2_FLASH_NUM_MTP_LAYERS,
)
hf_config.model_type = "mimo_v2_mtp"
# vLLM currently supports only the first MiMo-V2 MTP layer.
n_predict = _MIMO_V2_FLASH_NUM_MTP_LAYERS
hf_config.update(
{
"num_hidden_layers": 0,
"n_predict": n_predict,
"num_nextn_predict_layers": n_predict,
"architectures": ["MiMoV2MTPModel"],
}
)
if hf_config.architectures[0] == "Glm4MoeForCausalLM":
hf_config.model_type = "glm4_moe_mtp"
n_predict = getattr(hf_config, "num_nextn_predict_layers", None)
+17 -28
View File
@@ -983,19 +983,16 @@ class VllmConfig:
)
self.compilation_config.cudagraph_mode = CUDAGraphMode.NONE
# async tp is built on top of sequence parallelism
# and requires it to be enabled.
if self.compilation_config.pass_config.fuse_gemm_comms:
self.compilation_config.pass_config.enable_sp = True
if self.compilation_config.pass_config.enable_sp:
# async tp is built on top of sequence parallelism and requires it.
pass_config = self.compilation_config.pass_config
if pass_config.fuse_gemm_comms:
pass_config.enable_sp = True
if pass_config.enable_sp:
if self.parallel_config.tensor_parallel_size == 1:
logger.warning("Sequence Parallelism requires TP>1, disabling")
self.compilation_config.pass_config.enable_sp = False
self.compilation_config.pass_config.fuse_gemm_comms = False
pass_config.enable_sp = False
pass_config.fuse_gemm_comms = False
else:
# Compute SP threshold early; disable if None (model too
# small for SP to be beneficial).
pass_config = self.compilation_config.pass_config
if pass_config.sp_min_token_num is None:
from vllm.compilation.passes.fusion.sequence_parallelism import (
get_sequence_parallelism_threshold,
@@ -1015,8 +1012,8 @@ class VllmConfig:
"threshold heuristic, disabling. To force SP, "
"set pass_config.sp_min_token_num manually."
)
self.compilation_config.pass_config.enable_sp = False
self.compilation_config.pass_config.fuse_gemm_comms = False
pass_config.enable_sp = False
pass_config.fuse_gemm_comms = False
from vllm.utils.torch_utils import HAS_OPAQUE_TYPE
@@ -1098,6 +1095,7 @@ class VllmConfig:
self.compilation_config.cudagraph_num_of_warmups = 1
self._set_cudagraph_sizes()
else:
self.compilation_config.cudagraph_mode = CUDAGraphMode.NONE
@@ -1171,8 +1169,8 @@ class VllmConfig:
)
if self.compilation_config.pass_config.enable_sp:
# With pipeline parallelism or dynamo partitioning,
# native rms norm tracing errors due to incorrect residual shape.
# With pipeline parallelism, native rms norm tracing errors due to
# incorrect residual shape.
# Use custom rms norm to unblock. In the future,
# the pass will operate on higher-level IR to avoid the issue.
# TODO: https://github.com/vllm-project/vllm/issues/27894
@@ -1183,24 +1181,15 @@ class VllmConfig:
self.compilation_config.mode,
)
is_fullgraph = (
self.compilation_config.use_inductor_graph_partition
or len(self.compilation_config.splitting_ops or []) == 0
)
if self.parallel_config.pipeline_parallel_size > 1 or not is_fullgraph:
if self.parallel_config.pipeline_parallel_size > 1:
if "-rms_norm" not in self.compilation_config.custom_ops:
self.compilation_config.custom_ops.append("+rms_norm")
else:
regime = (
"Dynamo partition"
if not is_fullgraph
else "pipeline parallelism"
)
logger.warning_once(
"Sequence parallelism not supported with "
"native rms_norm when using %s, "
"this will likely lead to an error.",
regime,
"pipeline parallelism",
)
# final check of cudagraph mode after all possible updates
@@ -1212,9 +1201,9 @@ class VllmConfig:
and not self.compilation_config.cudagraph_mode.has_piecewise_cudagraphs() # noqa: E501
):
logger.warning_once(
"No piecewise cudagraph for executing cascade attention."
" Will fall back to eager execution if a batch runs "
"into cascade attentions."
"No piecewise cudagraph for executing cascade attention. "
"Will fall back to eager execution if a batch runs into "
"cascade attentions."
)
if self.compilation_config.cudagraph_mode.requires_piecewise_compilation():
+40 -32
View File
@@ -128,13 +128,6 @@ class CuMemAllocator:
return CuMemAllocator.instance
def __init__(self):
conf = os.environ.get("PYTORCH_CUDA_ALLOC_CONF", "")
assert "expandable_segments:True" not in conf, (
"Expandable segments are not compatible with memory pool. "
"Please track https://github.com/pytorch/pytorch/issues/147851 "
"for the latest updates."
)
self.pointer_to_data: dict[int, AllocationData] = {}
self.current_tag: str = CuMemAllocator.default_tag
self.allocator_and_pools: dict[str, Any] = {}
@@ -264,34 +257,49 @@ class CuMemAllocator:
assert isinstance(tag, str)
# Expandable segments are incompatible with the memory pool used for
# sleep mode (see https://github.com/pytorch/pytorch/issues/147851).
# If the user has enabled expandable segments via
# PYTORCH_CUDA_ALLOC_CONF, temporarily disable them for the duration
# of the memory pool context and restore on exit.
conf = os.environ.get("PYTORCH_CUDA_ALLOC_CONF", "")
expandable_was_enabled = "expandable_segments:True" in conf
if expandable_was_enabled:
torch.cuda.memory._set_allocator_settings("expandable_segments:False")
old_tag = self.current_tag
self.current_tag = tag
with use_memory_pool_with_allocator(
self.python_malloc_callback, self.python_free_callback
) as data:
# start to hit another PyTorch bug in PyTorch 2.6,
# possibly because of gc-related issue w.r.t. the allocator and
# the memory pool.
# to avoid the issue, we keep a reference of the data.
# see https://github.com/pytorch/pytorch/issues/146431 .
self.allocator_and_pools[tag] = data
yield
# PyTorch's bug, calling torch.cuda.empty_cache() will error
# when using pluggable allocator, see
# https://github.com/pytorch/pytorch/issues/145168 .
# if we have some memory allocated and then freed,
# the memory will not be released, e.g. in online quantization,
# where the model is created in higher precision, and then
# quantized in lower precision.
# Find all unused allocations and manually release them.
# TODO: we should expose `empty_cache` method in the memory pool.
# TODO: ask for help from PyTorch team to expose this method.
allocations = data[0].snapshot()
for allocation in allocations:
if allocation["allocated_size"] == 0:
handle = self._python_free_callback(allocation["address"])
unmap_and_release(handle)
try:
with use_memory_pool_with_allocator(
self.python_malloc_callback, self.python_free_callback
) as data:
# start to hit another PyTorch bug in PyTorch 2.6,
# possibly because of gc-related issue w.r.t. the allocator
# and the memory pool.
# to avoid the issue, we keep a reference of the data.
# see https://github.com/pytorch/pytorch/issues/146431 .
self.allocator_and_pools[tag] = data
yield
# PyTorch's bug, calling torch.cuda.empty_cache() will error
# when using pluggable allocator, see
# https://github.com/pytorch/pytorch/issues/145168 .
# if we have some memory allocated and then freed,
# the memory will not be released, e.g. in online
# quantization, where the model is created in higher
# precision, and then quantized in lower precision.
# Find all unused allocations and manually release them.
# TODO: we should expose `empty_cache` method in the memory
# pool.
# TODO: ask for help from PyTorch team to expose this method.
allocations = data[0].snapshot()
for allocation in allocations:
if allocation["allocated_size"] == 0:
handle = self._python_free_callback(allocation["address"])
unmap_and_release(handle)
finally:
self.current_tag = old_tag
if expandable_was_enabled:
torch.cuda.memory._set_allocator_settings("expandable_segments:True")
def get_current_usage(self) -> int:
"""
@@ -492,15 +492,18 @@ class FlashInferNVLinkTwoSidedManager(All2AllManagerBase):
CustomCommunicator,
)
dp_config = MnnvlConfig(
comm_backend=CustomCommunicator(get_dp_group().cpu_group),
# MNNVL workspace is allocated per rank in the comm_backend's group; the
# flashinfer kernel asserts workspace.size(0) == moe_ep_size, so the backend
# must span the EP group (= DP*PCP*TP), not the DP group.
ep_config = MnnvlConfig(
comm_backend=CustomCommunicator(self.cpu_group),
fabric_page_size=1 << 29, # 512MB
allocation_granularity=0, # Auto-detect
)
self.workspace_tensor = MnnvlMoe.get_moe_workspaces(self.mapping, dp_config)
self.workspace_tensor = MnnvlMoe.get_moe_workspaces(self.mapping, ep_config)
self.prepare_workspace_tensor = MnnvlMoe.get_moe_prepare_workspace(
self.mapping, dp_config
self.mapping, ep_config
)
self.world_size = world_size
@@ -605,8 +608,11 @@ class FlashInferNVLinkOneSidedManager(All2AllManagerBase):
CustomCommunicator,
)
dp_config = MnnvlConfig(
comm_backend=CustomCommunicator(get_dp_group().cpu_group),
# MNNVL workspace is allocated per rank in the comm_backend's group; the
# flashinfer kernel asserts workspace.size(0) == moe_ep_size, so the backend
# must span the EP group (= DP*PCP*TP), not the DP group.
ep_config = MnnvlConfig(
comm_backend=CustomCommunicator(self.cpu_group),
)
total_dispatch_payload_size_per_token = (
hidden_size // 2 # nvfp4 hidden states
@@ -628,7 +634,7 @@ class FlashInferNVLinkOneSidedManager(All2AllManagerBase):
top_k=top_k,
num_experts=num_experts,
workspace_size_per_rank=self.workspace_size,
mnnvl_config=dp_config,
mnnvl_config=ep_config,
)
self.gpus_per_node = gpus_per_node
@@ -2,6 +2,7 @@
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
"""Scheduler-side logic for the NIXL connector."""
import os
import threading
import time
from typing import TYPE_CHECKING, Any
@@ -88,6 +89,12 @@ class NixlConnectorScheduler:
if vllm_config.scheduler_config.disable_hybrid_kv_cache_manager:
logger.info("Hybrid Memory Allocator is enabled with NIXL")
if os.environ.get("VLLM_NIXL_ABORT_REQUEST_TIMEOUT") is not None:
logger.warning(
"VLLM_NIXL_ABORT_REQUEST_TIMEOUT is deprecated and will be "
"removed in release 0.22.0."
)
# Background thread for handling new handshake requests.
self._nixl_handshake_listener_t: threading.Thread | None = None
self._stop_event = threading.Event()
@@ -314,6 +314,9 @@ class OffloadingConnectorScheduler:
num_locally_computed_tokens = req_status.num_locally_computed_tokens
num_cached_tokens = num_locally_computed_tokens + num_external_tokens
params = req_status.req_context.kv_transfer_params
do_remote_decode = params is not None and params.get("do_remote_decode")
keys_to_load: list[OffloadKey] = []
dst_block_ids: list[int] = []
# per group
@@ -360,7 +363,11 @@ class OffloadingConnectorScheduler:
group_sizes.append(num_pending_gpu_blocks)
block_indices.append(num_locally_computed_gpu_blocks)
group_state.next_stored_block_idx = num_blocks
if not do_remote_decode:
# For P/D prefill requests (do_remote_decode=True), we do
# NOT skip saving the hit prefix, as we need to stream the
# entire KV cache so a remote decode node can consume it.
group_state.next_stored_block_idx = num_blocks
src_spec = self.manager.prepare_load(keys_to_load, req_status.req_context)
dst_spec = GPULoadStoreSpec(
+28
View File
@@ -542,6 +542,14 @@ class EngineArgs:
mm_encoder_attn_backend: AttentionBackendEnum | str | None = (
MultiModalConfig.mm_encoder_attn_backend
)
mm_encoder_attn_dtype: str | None = MultiModalConfig.mm_encoder_attn_dtype
mm_encoder_fp8_scale_path: str | None = MultiModalConfig.mm_encoder_fp8_scale_path
mm_encoder_fp8_scale_save_path: str | None = (
MultiModalConfig.mm_encoder_fp8_scale_save_path
)
mm_encoder_fp8_scale_save_margin: float = (
MultiModalConfig.mm_encoder_fp8_scale_save_margin
)
io_processor_plugin: str | None = None
renderer_num_workers: int = 1
skip_mm_profiling: bool = MultiModalConfig.skip_mm_profiling
@@ -1179,6 +1187,22 @@ class EngineArgs:
"--mm-encoder-attn-backend",
**multimodal_kwargs["mm_encoder_attn_backend"],
)
multimodal_group.add_argument(
"--mm-encoder-attn-dtype",
**multimodal_kwargs["mm_encoder_attn_dtype"],
)
multimodal_group.add_argument(
"--mm-encoder-fp8-scale-path",
**multimodal_kwargs["mm_encoder_fp8_scale_path"],
)
multimodal_group.add_argument(
"--mm-encoder-fp8-scale-save-path",
**multimodal_kwargs["mm_encoder_fp8_scale_save_path"],
)
multimodal_group.add_argument(
"--mm-encoder-fp8-scale-save-margin",
**multimodal_kwargs["mm_encoder_fp8_scale_save_margin"],
)
multimodal_group.add_argument(
"--interleave-mm-strings", **multimodal_kwargs["interleave_mm_strings"]
)
@@ -1517,6 +1541,10 @@ class EngineArgs:
mm_encoder_only=self.mm_encoder_only,
mm_encoder_tp_mode=self.mm_encoder_tp_mode,
mm_encoder_attn_backend=self.mm_encoder_attn_backend,
mm_encoder_attn_dtype=self.mm_encoder_attn_dtype,
mm_encoder_fp8_scale_path=self.mm_encoder_fp8_scale_path,
mm_encoder_fp8_scale_save_path=self.mm_encoder_fp8_scale_save_path,
mm_encoder_fp8_scale_save_margin=self.mm_encoder_fp8_scale_save_margin,
pooler_config=self.pooler_config,
generation_config=self.generation_config,
override_generation_config=self.override_generation_config,
@@ -317,4 +317,5 @@ class OpenAIServingChatBatch(OpenAIServingChat):
model=model_name,
choices=choices,
usage=usage,
system_fingerprint=self.system_fingerprint,
)
@@ -129,6 +129,9 @@ class ChatCompletionStreamResponse(OpenAIBaseModel):
model: str
choices: list[ChatCompletionResponseStreamChoice]
usage: UsageInfo | None = Field(default=None)
# Set only on the final chunk of a stream to mirror non-streaming responses
# without the per-chunk serialization overhead.
system_fingerprint: str | None = None
# not part of the OpenAI spec but for tracing the tokens
prompt_token_ids: list[int] | None = None
@@ -1195,6 +1195,16 @@ class OpenAIServingChat(OpenAIServing):
choices=[choice_data],
model=model_name,
)
# Stamp the fingerprint on terminal chunks only (those with
# finish_reason set). When ``include_usage`` is on, the
# trailing usage chunk below overrides this as the true
# final message.
if (
not include_usage
and self.system_fingerprint is not None
and choice_data.finish_reason is not None
):
chunk.system_fingerprint = self.system_fingerprint
# handle usage stats if requested & if continuous
if include_continuous_usage:
@@ -1229,6 +1239,7 @@ class OpenAIServingChat(OpenAIServing):
choices=[],
model=model_name,
usage=final_usage,
system_fingerprint=self.system_fingerprint,
)
final_usage_data = final_usage_chunk.model_dump_json(
exclude_unset=True, exclude_none=True
@@ -1637,6 +1648,7 @@ class OpenAIServingChat(OpenAIServing):
model=model_name,
choices=choices,
usage=usage,
system_fingerprint=self.system_fingerprint,
prompt_logprobs=clamp_prompt_logprobs(final_res.prompt_logprobs),
prompt_token_ids=(
final_res.prompt_token_ids if request.return_token_ids else None
+13 -1
View File
@@ -153,9 +153,21 @@ class BaseFrontendArgs:
"""If set to True, log the stack trace of error responses"""
tokens_only: bool = False
"""
If set to True, only enable the Tokens In<>Out endpoint.
If set to True, only enable the Tokens In<>Out endpoint.
This is intended for use in a Disaggregated Everything setup.
"""
fingerprint_mode: Literal["full", "hash", "custom", "none"] = "full"
"""Controls the ``system_fingerprint`` field on responses.
- ``full`` (default): ``vllm-<version>[-<parallelism>]-<hash8>``. Encodes
server version, non-trivial parallelism degrees (tp/pp/dp/ep), and an
8-char config hash.
- ``hash``: ``vllm-<version>-<hash8>``. Parallelism stripped.
- ``custom``: emits the literal string from ``--fingerprint-value``.
- ``none``: the field is omitted (serialized as ``null``).
"""
fingerprint_value: str | None = None
"""Literal fingerprint string used when ``--fingerprint-mode=custom``."""
@classmethod
def _customize_cli_kwargs(
@@ -512,3 +512,6 @@ class CompletionStreamResponse(OpenAIBaseModel):
model: str
choices: list[CompletionResponseStreamChoice]
usage: UsageInfo | None = Field(default=None)
# Set only on the final chunk of a stream to mirror non-streaming responses
# without the per-chunk serialization overhead.
system_fingerprint: str | None = None
+12 -1
View File
@@ -383,6 +383,7 @@ class OpenAIServingCompletion(OpenAIServing):
chunk = CompletionStreamResponse(
id=request_id,
object="text_completion",
created=created_time,
model=model_name,
choices=[
@@ -401,6 +402,14 @@ class OpenAIServingCompletion(OpenAIServing):
)
],
)
# Stamp on terminal chunk only when no trailing usage chunk
# will follow (that one is the true final message).
if (
not include_usage
and self.system_fingerprint is not None
and finish_reason is not None
):
chunk.system_fingerprint = self.system_fingerprint
if include_continuous_usage:
prompt_tokens = num_prompt_tokens[prompt_idx]
completion_tokens = previous_num_tokens[i]
@@ -410,7 +419,7 @@ class OpenAIServingCompletion(OpenAIServing):
total_tokens=prompt_tokens + completion_tokens,
)
response_json = chunk.model_dump_json(exclude_unset=False)
response_json = chunk.model_dump_json(exclude_unset=True)
yield f"data: {response_json}\n\n"
total_prompt_tokens = sum(num_prompt_tokens)
@@ -433,6 +442,7 @@ class OpenAIServingCompletion(OpenAIServing):
model=model_name,
choices=[],
usage=final_usage_info,
system_fingerprint=self.system_fingerprint,
)
final_usage_data = final_usage_chunk.model_dump_json(
exclude_unset=False, exclude_none=True
@@ -562,6 +572,7 @@ class OpenAIServingCompletion(OpenAIServing):
model=model_name,
choices=choices,
usage=usage,
system_fingerprint=self.system_fingerprint,
kv_transfer_params=kv_transfer_params,
)
+13
View File
@@ -157,6 +157,19 @@ class OpenAIServing:
self.renderer = engine_client.renderer
self.input_processor = engine_client.input_processor
# Computed once at startup (cached by ``vllm_config`` identity) and
# stamped on non-streaming responses. Streaming chunks deliberately
# omit it to avoid per-chunk overhead.
from vllm.entrypoints.openai.fingerprint import get_system_fingerprint
try:
self.system_fingerprint: str | None = get_system_fingerprint(
engine_client.vllm_config
)
except Exception:
# Never fail server startup over the fingerprint.
self.system_fingerprint = None
async def beam_search(
self,
prompt: EngineInput,
+84
View File
@@ -0,0 +1,84 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
"""Build the ``system_fingerprint`` string returned by the OpenAI-compatible
server.
Four modes, configured via ``--fingerprint-mode``:
* ``full`` (default): ``vllm-<version>[-<parallelism>]-<hash8>`` encodes
server version, any non-trivial parallelism degree (tp/pp/dp/ep), and an
8-char prefix of ``vllm_config.compute_hash()`` (covers model identity,
quant config, speculative, attention backend, etc.).
* ``hash``: ``vllm-<version>-<hash8>`` parallelism stripped.
* ``custom``: user-provided literal via ``--fingerprint-value``.
* ``none``: the field is omitted (serialized as ``null``).
``get_system_fingerprint`` is only called at serving-class init (a handful
of times per server); each subclass caches the returned string on
``self.system_fingerprint``, so per-request cost is one attribute read.
"""
from __future__ import annotations
from typing import Any, Literal
FingerprintMode = Literal["full", "hash", "custom", "none"]
_DEFAULT_MODE: FingerprintMode = "full"
_CUSTOM_VALUE: str | None = None
def set_default_fingerprint_mode(
mode: FingerprintMode,
custom_value: str | None = None,
) -> None:
"""Configure the fingerprint mode for subsequent ``get_system_fingerprint``
calls. Called once at server startup."""
global _DEFAULT_MODE, _CUSTOM_VALUE
_DEFAULT_MODE = mode
_CUSTOM_VALUE = custom_value
def get_system_fingerprint(vllm_config: Any) -> str | None:
"""Return the fingerprint for ``vllm_config`` using the mode configured by
``set_default_fingerprint_mode``."""
return build_system_fingerprint(vllm_config, _DEFAULT_MODE, _CUSTOM_VALUE)
def build_system_fingerprint(
vllm_config: Any,
mode: FingerprintMode = "full",
custom_value: str | None = None,
) -> str | None:
if mode == "none":
return None
if mode == "custom":
return custom_value
from vllm import __version__ as vllm_version
try:
hash8 = vllm_config.compute_hash()[:8]
except Exception:
hash8 = "nohash"
if mode == "hash":
return f"vllm-{vllm_version}-{hash8}"
# mode == "full"
parts: list[str] = [f"vllm-{vllm_version}"]
pc = getattr(vllm_config, "parallel_config", None)
if pc is not None:
tp = getattr(pc, "tensor_parallel_size", 1)
if tp > 1:
parts.append(f"tp{tp}")
pp = getattr(pc, "pipeline_parallel_size", 1)
if pp > 1:
parts.append(f"pp{pp}")
dp = getattr(pc, "data_parallel_size", 1)
if dp > 1:
parts.append(f"dp{dp}")
if getattr(pc, "enable_expert_parallel", False):
parts.append("ep")
parts.append(hash8)
return "-".join(parts)
@@ -61,9 +61,17 @@ async def init_generate_state(
)
from vllm.entrypoints.openai.chat_completion.serving import OpenAIServingChat
from vllm.entrypoints.openai.completion.serving import OpenAIServingCompletion
from vllm.entrypoints.openai.fingerprint import set_default_fingerprint_mode
from vllm.entrypoints.openai.responses.serving import OpenAIServingResponses
from vllm.entrypoints.serve.disagg.serving import ServingTokens
# Applied before any serving class is constructed so that each one picks
# up the chosen mode on its first cache miss.
set_default_fingerprint_mode(
getattr(args, "fingerprint_mode", "full"),
getattr(args, "fingerprint_value", None),
)
if args.tool_server == "demo":
tool_server: ToolServer | None = DemoToolServer()
assert isinstance(tool_server, DemoToolServer)
-6
View File
@@ -247,7 +247,6 @@ if TYPE_CHECKING:
VLLM_SHARED_EXPERTS_STREAM_TOKEN_THRESHOLD: int = 256
VLLM_COMPILE_CACHE_SAVE_FORMAT: Literal["binary", "unpacked"] = "binary"
VLLM_USE_V2_MODEL_RUNNER: bool = False
VLLM_DEEPSEEK_V4_USE_MEGA_MOE: bool = False
VLLM_LOG_MODEL_INSPECTION: bool = False
VLLM_DEBUG_MFU_METRICS: bool = False
VLLM_WEIGHT_OFFLOADING_DISABLE_PIN_MEMORY: bool = False
@@ -1676,11 +1675,6 @@ environment_variables: dict[str, Callable[[], Any]] = {
"VLLM_USE_V2_MODEL_RUNNER": lambda: bool(
int(os.getenv("VLLM_USE_V2_MODEL_RUNNER", "0"))
),
# Use the DeepGEMM MegaMoE fused expert kernel for DeepSeek V4 routed
# experts. Set to 0 to fall back to the standard SharedFusedMoE path.
"VLLM_DEEPSEEK_V4_USE_MEGA_MOE": lambda: bool(
int(os.getenv("VLLM_DEEPSEEK_V4_USE_MEGA_MOE", "0"))
),
# Log model inspection after loading.
# If enabled, logs a transformers-style hierarchical view of the model
# with quantization methods and attention backends.

Some files were not shown because too many files have changed in this diff Show More