Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
21 changes: 21 additions & 0 deletions tests/kernels/attention/test_flashmla_sparse.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)]
32 changes: 12 additions & 20 deletions vllm/models/deepseek_v4/nvidia/flashmla.py
Original file line number Diff line number Diff line change
Expand Up @@ -246,7 +246,6 @@ def _forward_prefill(
) -> 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
Expand Down Expand Up @@ -274,29 +273,22 @@ def _forward_prefill(
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
Expand All @@ -320,7 +312,7 @@ def _forward_prefill(
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
Expand All @@ -341,8 +333,8 @@ def _forward_prefill(
self.window_size,
self.compress_ratio,
top_k,
M,
N,
chunk_M,
chunk_N,
)
flash_mla_sparse_fwd(
q=q[query_start:query_end],
Expand Down
103 changes: 100 additions & 3 deletions vllm/v1/attention/backends/mla/sparse_swa.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,6 +9,7 @@
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,
Expand Down Expand Up @@ -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
Expand All @@ -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.
Expand All @@ -213,6 +292,10 @@ def __init__(self, *args, **kwargs):
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 = (
Expand Down Expand Up @@ -279,6 +362,7 @@ def build(
"""
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
Expand Down Expand Up @@ -323,7 +407,9 @@ def build(
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
Expand All @@ -350,7 +436,7 @@ def build(
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(
Expand Down Expand Up @@ -391,8 +477,10 @@ def _build_deepseek_v4_metadata(
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
Expand All @@ -401,10 +489,11 @@ def _build_deepseek_v4_metadata(
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
)
Expand All @@ -419,7 +508,15 @@ def _build_deepseek_v4_metadata(
)

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

Expand Down
Loading