[Perf] Optimize DSv4 prefill chunk planning, 4.0% E2E Throughput Improvement (#45061)

Signed-off-by: yewentao256 <[email protected]>
This commit is contained in:
Wentao Ye
2026-06-15 19:50:21 +00:00
committed by GitHub
parent 51ec5cf08f
commit e18fe932ca
3 changed files with 133 additions and 23 deletions
@@ -122,3 +122,24 @@ def test_sparse_flashmla_prefill_smoke():
assert out.shape == (s_q, h_q, d_v)
assert max_logits.shape == (s_q, h_q)
assert lse.shape == (s_q, h_q)
def test_deepseek_v4_prefill_chunk_planning_expands_for_short_sequences():
from vllm.v1.attention.backends.mla.sparse_swa import DeepseekSparseSWAMetadata
metadata = DeepseekSparseSWAMetadata(
block_table=torch.empty(0, dtype=torch.int32),
slot_mapping=torch.empty(0, dtype=torch.int32),
block_size=64,
num_prefills=5,
prefill_seq_lens_cpu=torch.tensor([80, 96, 112, 128, 144], dtype=torch.int32),
prefill_query_lens_cpu=torch.tensor([4, 4, 4, 4, 4], dtype=torch.int32),
prefill_window_size=64,
prefill_max_model_len=1024,
prefill_max_num_batched_tokens=128,
)
chunk_plan = metadata.get_prefill_chunk_plan(compress_ratio=4, prefill_chunk_size=4)
# the adaptive plan keeps all 5 in one chunk
assert chunk_plan == [(0, 5, 36, 103)]
+12 -20
View File
@@ -246,7 +246,6 @@ class DeepseekV4FlashMLAAttention(DeepseekV4Attention):
) -> None:
swa_only = attn_metadata is None
num_prefills = swa_metadata.num_prefills
num_prefill_tokens = swa_metadata.num_prefill_tokens
num_decodes = swa_metadata.num_decodes
num_decode_tokens = swa_metadata.num_decode_tokens
@@ -274,29 +273,22 @@ class DeepseekV4FlashMLAAttention(DeepseekV4Attention):
assert attn_metadata is not None
topk_indices = attn_metadata.c128a_prefill_topk_indices
top_k = topk_indices.shape[-1]
# Compressed region must fit the full compressed pool (seq_len //
# compress_ratio), not just top_k. top_k bounds how many indices
# the indexer selects, not the pool size it indexes into.
N = (self.max_model_len + self.compress_ratio - 1) // self.compress_ratio
else:
# NOTE(woosuk): topk_indices will not be used for SWA-only layers.
assert self.topk_indices_buffer is not None
topk_indices = self.topk_indices_buffer[num_decode_tokens:]
top_k = 0
N = 0
M = N + self.window_size + self.max_num_batched_tokens
chunk_size_const = self.PREFILL_CHUNK_SIZE
num_chunks = (num_prefills + chunk_size_const - 1) // chunk_size_const
chunk_plan = swa_metadata.get_prefill_chunk_plan(
compress_ratio=self.compress_ratio,
prefill_chunk_size=self.PREFILL_CHUNK_SIZE,
)
assert chunk_plan, "prefill chunk plan must be non-empty when num_prefills > 0"
workspace_manager = current_workspace_manager()
kv = workspace_manager.get_simultaneous(
((chunk_size_const, M, q.shape[-1]), torch.bfloat16),
)[0]
for chunk_idx in range(num_chunks):
chunk_start = chunk_idx * chunk_size_const
chunk_end = min(chunk_start + chunk_size_const, num_prefills)
for chunk_start, chunk_end, chunk_N, chunk_M in chunk_plan:
chunk_size = chunk_end - chunk_start
kv = workspace_manager.get_simultaneous(
((chunk_size, chunk_M, q.shape[-1]), torch.bfloat16),
)[0]
if not swa_only:
# Gather compressed KV
assert attn_metadata is not None
@@ -320,7 +312,7 @@ class DeepseekV4FlashMLAAttention(DeepseekV4Attention):
gather_lens=gather_lens[chunk_start:chunk_end],
block_table=swa_block_table[chunk_start:chunk_end],
block_size=swa_metadata.block_size,
offset=N,
offset=chunk_N,
)
# Combine the topk indices and SWA indices for gathered KV cache
@@ -341,8 +333,8 @@ class DeepseekV4FlashMLAAttention(DeepseekV4Attention):
self.window_size,
self.compress_ratio,
top_k,
M,
N,
chunk_M,
chunk_N,
)
flash_mla_sparse_fwd(
q=q[query_start:query_end],
+100 -3
View File
@@ -9,6 +9,7 @@ from vllm.config import CacheConfig, VllmConfig, get_current_vllm_config
from vllm.model_executor.layers.attention_layer_base import AttentionLayerBase
from vllm.platforms import current_platform
from vllm.triton_utils import tl, triton
from vllm.utils.math_utils import cdiv
from vllm.v1.attention.backend import (
AttentionBackend,
AttentionCGSupport,
@@ -172,7 +173,12 @@ class DeepseekSparseSWAMetadata:
# Pre-computed prefill metadata shared across all DeepseekV4 attention layers.
prefill_seq_lens: torch.Tensor | None = None
prefill_seq_lens_cpu: torch.Tensor | None = None
prefill_gather_lens: torch.Tensor | None = None
prefill_query_lens_cpu: torch.Tensor | None = None
prefill_window_size: int = 0
prefill_max_model_len: int = 0
prefill_max_num_batched_tokens: int = 0
# Per-layer-type FlashMLA tile-scheduler metadata. One FlashMLASchedMeta
# per present DeepseekV4 layer type, shared across all ~60 layers of that type
@@ -188,6 +194,79 @@ class DeepseekSparseSWAMetadata:
tile_sched_c4a: "FlashMLASchedMeta | None" = None
tile_sched_c128a: "FlashMLASchedMeta | None" = None
def get_prefill_chunk_plan(
self, compress_ratio: int, prefill_chunk_size: int
) -> list[tuple[int, int, int, int]]:
if self.num_prefills == 0:
return []
assert self.prefill_seq_lens_cpu is not None
assert self.prefill_query_lens_cpu is not None
# query_len <= max_num_batched_tokens and
# gather_len = query_len + min(prefix_len, window_size - 1), so the
# worst-case gathered width is bounded by
# max_num_batched_tokens + window_size - 1. The compressed prefix pool
# is bounded by ceil(max_model_len / compress_ratio).
max_workspace_area = prefill_chunk_size * (
(
0
if compress_ratio <= 1
else cdiv(self.prefill_max_model_len, compress_ratio)
)
+ self.prefill_window_size
+ self.prefill_max_num_batched_tokens
)
prefix_lens_cpu = self.prefill_seq_lens_cpu - self.prefill_query_lens_cpu
gather_lens_cpu = self.prefill_query_lens_cpu + torch.clamp(
prefix_lens_cpu, min=0, max=self.prefill_window_size - 1
)
compressed_lens_cpu = (
torch.zeros_like(self.prefill_seq_lens_cpu)
if compress_ratio <= 1
else torch.div(
self.prefill_seq_lens_cpu,
compress_ratio,
rounding_mode="floor",
)
)
chunk_plan: list[tuple[int, int, int, int]] = []
chunk_start = 0
while chunk_start < self.num_prefills:
chunk_max_compressed = int(compressed_lens_cpu[chunk_start].item())
chunk_max_gather = int(gather_lens_cpu[chunk_start].item())
chunk_end = chunk_start + 1
while chunk_end < self.num_prefills:
candidate_max_compressed = max(
chunk_max_compressed,
int(compressed_lens_cpu[chunk_end].item()),
)
candidate_max_gather = max(
chunk_max_gather,
int(gather_lens_cpu[chunk_end].item()),
)
candidate_width = candidate_max_compressed + candidate_max_gather
candidate_area = (chunk_end - chunk_start + 1) * candidate_width
if candidate_area > max_workspace_area:
break
chunk_max_compressed = candidate_max_compressed
chunk_max_gather = candidate_max_gather
chunk_end += 1
chunk_plan.append(
(
chunk_start,
chunk_end,
chunk_max_compressed,
chunk_max_compressed + chunk_max_gather,
)
)
chunk_start = chunk_end
return chunk_plan
class DeepseekSparseSWAMetadataBuilder(AttentionMetadataBuilder):
"""Builds metadata for DeepseekV4 SWA cache.
@@ -213,6 +292,10 @@ class DeepseekSparseSWAMetadataBuilder(AttentionMetadataBuilder):
self.head_size = mla_spec.head_size # Already considered quantization.
self.compress_ratio = mla_spec.compress_ratio
self.block_size = mla_spec.block_size
self.max_model_len = self.vllm_config.model_config.max_model_len
self.max_num_batched_tokens = (
self.vllm_config.scheduler_config.max_num_batched_tokens
)
# Handle MTP: adjust decode_threshold like the indexer does
self.num_speculative_tokens = (
@@ -279,6 +362,7 @@ class DeepseekSparseSWAMetadataBuilder(AttentionMetadataBuilder):
"""
num_reqs = common_attn_metadata.num_reqs
seq_lens = common_attn_metadata.seq_lens
seq_lens_cpu = common_attn_metadata.seq_lens_cpu_upper_bound
query_start_loc = common_attn_metadata.query_start_loc
query_start_loc_cpu = common_attn_metadata.query_start_loc_cpu
block_table = common_attn_metadata.block_table_tensor
@@ -323,7 +407,9 @@ class DeepseekSparseSWAMetadataBuilder(AttentionMetadataBuilder):
num_decodes,
num_prefills,
seq_lens,
seq_lens_cpu,
query_start_loc,
query_start_loc_cpu,
)
# Per-layer-type tile-scheduler plan holders. Empty FlashMLASchedMeta
@@ -350,7 +436,7 @@ class DeepseekSparseSWAMetadataBuilder(AttentionMetadataBuilder):
tile_sched_swaonly=tile_sched[_LAYER_TYPE_SWAONLY],
tile_sched_c4a=tile_sched[_LAYER_TYPE_C4A],
tile_sched_c128a=tile_sched[_LAYER_TYPE_C128A],
**deepseek_v4_fields,
**deepseek_v4_fields, # type: ignore[arg-type]
)
def build_tile_scheduler(
@@ -391,8 +477,10 @@ class DeepseekSparseSWAMetadataBuilder(AttentionMetadataBuilder):
num_decodes: int,
num_prefills: int,
seq_lens: torch.Tensor,
seq_lens_cpu: torch.Tensor | None,
query_start_loc: torch.Tensor,
) -> dict[str, torch.Tensor | None]:
query_start_loc_cpu: torch.Tensor,
) -> dict[str, torch.Tensor | int | None]:
"""Pre-compute DeepseekV4 prefill metadata during the metadata build phase.
Returns a dict of keyword arguments to pass to the
@@ -401,10 +489,11 @@ class DeepseekSparseSWAMetadataBuilder(AttentionMetadataBuilder):
Note: C128A topk indices are computed by the FlashMLASparse builder
(which owns the C128A block_table), not here.
"""
result: dict[str, torch.Tensor | None] = {}
result: dict[str, torch.Tensor | int | None] = {}
# --- Prefill query metadata (single Triton kernel + CPU slicing) ---
if num_prefills > 0:
assert seq_lens_cpu is not None
pfx_gather_lens = torch.empty(
num_prefills, dtype=torch.int32, device=seq_lens.device
)
@@ -419,7 +508,15 @@ class DeepseekSparseSWAMetadataBuilder(AttentionMetadataBuilder):
)
result["prefill_seq_lens"] = seq_lens[num_decodes:]
result["prefill_seq_lens_cpu"] = seq_lens_cpu[num_decodes:]
result["prefill_gather_lens"] = pfx_gather_lens
result["prefill_query_lens_cpu"] = (
query_start_loc_cpu[num_decodes + 1 : num_decodes + num_prefills + 1]
- query_start_loc_cpu[num_decodes : num_decodes + num_prefills]
).to(dtype=torch.int32)
result["prefill_window_size"] = self.window_size
result["prefill_max_model_len"] = self.max_model_len
result["prefill_max_num_batched_tokens"] = self.max_num_batched_tokens
return result