mirror of
https://github.com/vllm-project/vllm.git
synced 2026-08-03 12:28:06 +00:00
[Perf] Optimize DSv4 prefill chunk planning, 4.0% E2E Throughput Improvement (#45061)
Signed-off-by: yewentao256 <[email protected]>
This commit is contained in:
@@ -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)]
|
||||
|
||||
@@ -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],
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
Reference in New Issue
Block a user