class DeepseekV32IndexerMetadataBuilder(AttentionMetadataBuilder):
# The indexer opts out of the shared reorder-threshold vote (see __init__),
# so this is None; its own split uses self.decode_threshold.
reorder_batch_threshold: int | None = None
requires_block_table_width = True
@classmethod
def get_cudagraph_support(
cls,
vllm_config: VllmConfig,
kv_cache_spec: KVCacheSpec,
) -> AttentionCGSupport:
if _supports_varlen_paged_mqa_logits() or _use_flattening(vllm_config):
return AttentionCGSupport.ALWAYS
return AttentionCGSupport.UNIFORM_BATCH
def __init__(self, *args, block_table_width: int, **kwargs) -> None:
super().__init__(*args, **kwargs)
scheduler_config = self.vllm_config.scheduler_config
parallel_config = self.vllm_config.parallel_config
self.dcp_world_size = parallel_config.decode_context_parallel_size
self.dcp_rank = get_dcp_group().rank_in_group if self.dcp_world_size > 1 else 0
self.pcp_world_size = parallel_config.prefill_context_parallel_size
self.use_pcp = self.pcp_world_size > 1
self.pcp_rank = get_pcp_group().rank_in_group if self.use_pcp else 0
self.cp_kv_cache_interleave_size = parallel_config.cp_kv_cache_interleave_size
# KV compression (DeepseekV4). Default to 1 for no compression.
self.compress_ratio = 1
if isinstance(self.kv_cache_spec, MLAAttentionSpec):
assert isinstance(self.kv_cache_spec.tokens_per_state, int)
self.compress_ratio = self.kv_cache_spec.tokens_per_state
# NOTE(Chen):an estimated max size of flattened_kv. Need to double check.
# Counted in compressed rows, like the chunker's seq_lens and the
# workspace.
self.max_prefill_buffer_size = (
get_max_prefill_buffer_size(self.vllm_config) // self.compress_ratio
)
self.num_speculative_tokens = (
self.vllm_config.speculative_config.num_speculative_tokens
if self.vllm_config.speculative_config
else 0
)
self.indexer_uses_fp4 = dsa_indexer_uses_fp4(self.vllm_config)
next_n = self.num_speculative_tokens + 1
self.decode_threshold = next_n
self.reorder_batch_threshold = None
self.use_flattening = _use_flattening(self.vllm_config)
self.supports_varlen = _supports_varlen_paged_mqa_logits()
logger.info_once(
"DSA indexer decode path: use_flattening=%s supports_varlen=%s "
"(next_n=%d, use_fp4_cache=%s)",
self.use_flattening,
self.supports_varlen,
next_n,
self.indexer_uses_fp4,
)
sm_count = num_compute_units(self.device.index)
self.num_sms = sm_count
self.offsets_buffer = torch.arange(
next_n, device=self.device, dtype=torch.int32
)
self.decode_lens_buffer = torch.zeros(
(scheduler_config.max_num_batched_tokens,),
dtype=torch.int32,
device=self.device,
)
self.per_req_decode_lens_buffer = torch.zeros(
(scheduler_config.max_num_batched_tokens,),
dtype=torch.int32,
device=self.device,
)
# Shared workspace for decode seq_lens. Native MTP views this as
# (B, max_decode_len) at runtime, keeping context_lens contiguous even
# when max_decode_len is smaller than next_n.
self.decode_seq_lens_buffer = torch.zeros(
(scheduler_config.max_num_batched_tokens,),
dtype=torch.int32,
device=self.device,
)
self.global_decode_seq_lens_buffer = torch.zeros(
(scheduler_config.max_num_batched_tokens,),
dtype=torch.int32,
device=self.device,
)
self.decode_indices_buffer = torch.zeros(
(scheduler_config.max_num_batched_tokens,),
dtype=torch.int32,
device=self.device,
)
self.arange_buffer = torch.arange(
max(
scheduler_config.max_num_seqs * next_n,
scheduler_config.max_num_batched_tokens,
),
dtype=torch.int32,
device=self.device,
)
# Materialize the rank on device during builder initialization. Creating
# this scalar in build() would introduce a GPU<->CPU sync in the decode
# hot path.
self.dcp_rank_tensor = torch.tensor(
self.dcp_rank, dtype=torch.int32, device=self.device
)
self.expanded_block_table_buffer = torch.zeros(
(scheduler_config.max_num_batched_tokens, block_table_width),
dtype=torch.int32,
device=self.device,
)
# See: DeepGMM/csrc/apis/attention.hpp. Sized for one slot per SM;
# build() narrows it to whatever the kernel actually schedules.
self.scheduler_metadata_buffer = torch.empty(
(self.num_sms + 1, 2), dtype=torch.int32, device=self.device
)
if self.dcp_world_size > 1 and self.compress_ratio > 1:
raise NotImplementedError(
"DCP is not supported with sparse indexer KV compression "
f"(compress_ratio={self.compress_ratio})."
)
# Pre-allocate buffers for CUDA graph compatibility when
if self.compress_ratio > 1:
# compress_ratio > 1 (DeepseekV4)
# Compressed slot mapping output buffer
self.compressed_slot_mapping_buffer = torch.zeros(
(scheduler_config.max_num_batched_tokens,),
dtype=torch.int64,
device=self.device,
)
# Buffer for compressed seq_lens in decode path
self.expanded_seq_lens_buffer = torch.zeros(
(scheduler_config.max_num_batched_tokens,),
dtype=torch.int32,
device=self.device,
)
self.indexer_decode_block_table_buffer: torch.Tensor | None = None
self._max_num_batched_tokens = scheduler_config.max_num_batched_tokens
def _dcp_localize_decode_seq_lens(
self,
seq_lens: torch.Tensor,
num_decodes: int,
seq_lens_is_buffer_view: bool,
) -> torch.Tensor:
local_seq_lens = get_dcp_local_seq_lens(
seq_lens,
self.dcp_world_size,
self.dcp_rank_tensor,
self.cp_kv_cache_interleave_size,
)
if seq_lens_is_buffer_view:
seq_lens.copy_(local_seq_lens)
return seq_lens
out = self.decode_seq_lens_buffer[:num_decodes]
out.copy_(local_seq_lens)
return out
def _prepare_decode_tensors(
self,
seq_lens: torch.Tensor,
block_table: torch.Tensor,
decode_lens: torch.Tensor,
decode_lens_cpu: torch.Tensor,
query_start_loc: torch.Tensor,
num_decodes: int,
num_decode_tokens: int,
use_native: bool,
next_n: int,
max_decode_len: int,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, int, bool]:
"""Prepare native or per-token flattened decode tensors."""
spec_config = self.vllm_config.speculative_config
adaptive = bool(spec_config and spec_config.enable_adaptive_verification)
min_decode_len = int(decode_lens_cpu.min().item())
if not use_native:
assert self.decode_seq_lens_buffer.dim() == 1
if (
not self.supports_varlen
and (num_decodes == 1 or not adaptive)
and min_decode_len == max_decode_len
and num_decodes * max_decode_len == num_decode_tokens
):
# Uniform decode lengths with no cudagraph token padding.
_PREPARE_UNIFORM_DECODE_KERNEL(
seq_lens,
self.decode_seq_lens_buffer,
block_table,
self.expanded_block_table_buffer,
self.decode_lens_buffer,
num_decode_tokens,
max_decode_len,
)
self.decode_seq_lens_buffer[num_decode_tokens:] = 0
seq_lens = self.decode_seq_lens_buffer[:num_decode_tokens]
block_table = self.expanded_block_table_buffer[:num_decode_tokens]
decode_lens = self.decode_lens_buffer[:num_decode_tokens]
return seq_lens, block_table, decode_lens, num_decode_tokens, False
else:
# Variable decode lengths.
# Assume 4 requests with seq_lens [10, 7, 12, 0] (the final req is
# padding) and decode_lens [3, 1, 4, 0] in the below example comments.
# The context lengths are therefore
# [10-3, 7-1, 12-4, 0-0] = [7, 6, 8, 0].
# 3 + 1 + 4 + 0 = 8
actual_expanded = int(decode_lens_cpu.sum().item())
# Fuse expanded_base and expanded_starts into a single
# repeat_interleave:
# seq_len_i = (context_start[b] - query_start_loc[b]) + arange[i] + 1
# where context_start[b] = seq_lens[b] - decode_lens[b].
# Example: offsets = [7-0, 6-3, 8-4, 0-8] = [7, 3, 4, -8]
# expanded_offsets = [7, 7, 7, 3, 4, 4, 4, 4]
# result = [8, 9, 10, 7, 9, 10, 11, 12]
expanded_offsets = torch.repeat_interleave(
seq_lens - decode_lens - query_start_loc,
decode_lens,
output_size=actual_expanded,
)
# [8, 9, 10, 7, 9, 10, 11, 12, ...] where ... is unused buffer space
self.decode_seq_lens_buffer[:actual_expanded] = (
expanded_offsets + self.arange_buffer[:actual_expanded] + 1
)
self.decode_seq_lens_buffer[actual_expanded:] = 0
seq_lens = self.decode_seq_lens_buffer[:num_decode_tokens]
# Give each of the flattened entries the same block table row as the
# original request.
self.expanded_block_table_buffer[:actual_expanded] = (
torch.repeat_interleave(
block_table, decode_lens, dim=0, output_size=actual_expanded
)
)
if actual_expanded < num_decode_tokens:
self.expanded_block_table_buffer[
actual_expanded:num_decode_tokens, 0
] = 0
block_table = self.expanded_block_table_buffer[:num_decode_tokens]
# All reqs now have decode_len=1
self.decode_lens_buffer[:num_decode_tokens] = 1
decode_lens = self.decode_lens_buffer[:num_decode_tokens]
return seq_lens, block_table, decode_lens, num_decode_tokens, False
else:
# Native path: plain decode (next_n==1) or spec decode
# with 2D per-token context lengths (next_n > 1).
#
# When decode_lens are not truly uniform (e.g. some requests have
# decode_len < next_n due to padding or short prefills), the simple
# reshape in sparse_attn_indexer won't work. Use pack_seq_triton
# (requires_padding) instead.
requires_padding = min_decode_len != max_decode_len
if use_native and next_n > 1:
assert self.decode_seq_lens_buffer.dim() == 1
# (B, max_decode_len): token j attends to
# L - max_decode_len + j + 1 KV tokens.
seq_lens_buffer = self.decode_seq_lens_buffer[
: num_decodes * max_decode_len
].view(num_decodes, max_decode_len)
# Clamp at 0: padding requests have seq_len == 0, which would
# otherwise make token 0 negative (next_n=2 gives 0-2+1+0 = -1).
# Downstream kernels read these as uint32, turning -1 into ~4e9.
seq_lens_buffer[:] = (
seq_lens.unsqueeze(1)
- max_decode_len
+ 1
+ self.offsets_buffer[:max_decode_len]
).clamp_(min=0)
seq_lens = seq_lens_buffer
return seq_lens, block_table, decode_lens, num_decodes, requires_padding
def _prepare_global_decode_seq_lens(
self,
global_seq_lens: torch.Tensor | None,
decode_lens: torch.Tensor,
decode_lens_cpu: torch.Tensor,
query_start_loc: torch.Tensor,
num_decode_tokens: int,
use_native: bool,
max_decode_len: int,
) -> torch.Tensor | None:
if global_seq_lens is None:
return None
if use_native or max_decode_len <= 1:
return global_seq_lens
actual_expanded = int(decode_lens_cpu.sum().item())
if actual_expanded > 0:
expanded_offsets = torch.repeat_interleave(
global_seq_lens - decode_lens - query_start_loc,
decode_lens,
output_size=actual_expanded,
)
self.global_decode_seq_lens_buffer[:actual_expanded] = (
expanded_offsets + self.arange_buffer[:actual_expanded] + 1
)
self.global_decode_seq_lens_buffer[actual_expanded:num_decode_tokens] = 0
return self.global_decode_seq_lens_buffer[:num_decode_tokens]
def _split_pcp_dcp_prefill_chunks(
self,
row_req_idx: np.ndarray,
row_shard_rows: np.ndarray,
row_query_lens_cpu: torch.Tensor,
max_logits_bytes: int,
request_offset: int,
) -> list[tuple[slice, slice]]:
"""Chunk by request rather than by row, so a split prefill's two rows
charge their shared context once, then widen each chunk back to rows.
``row_shard_rows`` holds the whole request's largest DCP shard on
every row, so the plan is identical on every PCP rank.
"""
row_bounds = request_row_bounds(row_req_idx)
first_rows = row_bounds[:-1]
# Each request's context, padded to whole DCP shards.
seq_lens = row_shard_rows[first_rows] * self.dcp_world_size
# Every rank holds a full-size chunk of a split request (shorter ones
# are replicated), so rows x longest row is the same on every rank,
# whichever row holds the short tail.
row_query_lens = row_query_lens_cpu.numpy()
query_lens = np.diff(row_bounds) * np.maximum.reduceat(
row_query_lens, first_rows
)
chunk_specs = self._split_indexer_prefill_chunks(
torch.from_numpy(seq_lens.astype(np.int32)),
torch.from_numpy(query_lens.astype(np.int32)),
self.max_prefill_buffer_size,
max_logits_bytes,
)
return [
(
slice(
request_offset + int(row_bounds[request_slice.start]),
request_offset + int(row_bounds[request_slice.stop]),
),
query_slice,
)
for request_slice, query_slice in chunk_specs
]
def _prefill_split_seq_lens(self, seq_lens_cpu: torch.Tensor) -> torch.Tensor:
"""Per-request KV lengths the prefill chunker budgets logits with;
subclasses whose logits rows are wider than the context override."""
return seq_lens_cpu
@staticmethod
def _split_indexer_prefill_chunks(
compressed_seq_lens_cpu: torch.Tensor,
prefill_query_lens_cpu: torch.Tensor,
workspace_size: int,
max_logits_bytes: int,
request_offset: int = 0,
) -> list[tuple[slice, slice]]:
"""Split this step's prefill requests into chunks, respecting:
- N constraint: total_seq_lens <= workspace_size (existing O(N)
workspace)
- Logits constraint: M * N * 4 <= max_logits_bytes
When a single request-level chunk still exceeds the logits budget,
sub-chunks on the query dimension (M) to bound peak memory.
Returns list of (req_slice, query_slice) tuples.
"""
chunks: list[tuple[slice, slice]] = []
n = len(compressed_seq_lens_cpu)
max_logits_elems = max_logits_bytes // 4
end = 0
while end < n:
start, chunk_m, chunk_n = end, 0, 0
while end < n:
q, s = (
prefill_query_lens_cpu[end].item(),
compressed_seq_lens_cpu[end].item(),
)
new_m, new_n = chunk_m + q, chunk_n + s
if new_n <= workspace_size and new_m * new_n <= max_logits_elems:
chunk_m, chunk_n = new_m, new_n
end += 1
else:
break
# A single request can exceed the budget, requiring sub-chunking
# on the query dimension.
if end == start:
chunk_m, chunk_n = (
prefill_query_lens_cpu[end].item(),
compressed_seq_lens_cpu[end].item(),
)
end += 1
req_slice = slice(start + request_offset, end + request_offset)
max_q = (
max(1, max_logits_elems // chunk_n) if chunk_n > 0 else max(1, chunk_m)
)
for q_off in range(0, chunk_m, max_q):
sub_m = min(max_q, chunk_m - q_off)
chunks.append((req_slice, slice(q_off, q_off + sub_m)))
return chunks
def build(
self,
common_prefix_len: int,
common_attn_metadata: CommonAttentionMetadata,
fast_build: bool = False,
) -> DeepseekV32IndexerMetadata:
num_reqs = common_attn_metadata.num_reqs
num_tokens = common_attn_metadata.num_actual_tokens
query_start_loc = common_attn_metadata.query_start_loc
query_start_loc_cpu = common_attn_metadata.query_start_loc_cpu
seq_lens = common_attn_metadata.seq_lens
slot_mapping = common_attn_metadata.slot_mapping
block_table = common_attn_metadata.block_table_tensor
dcp_local_seq_lens = common_attn_metadata.dcp_local_seq_lens
compressed_slot_mapping = slot_mapping
indexer_block_table = block_table
if self.compress_ratio > 1:
kernel_block_size = self.kernel_block_size
if (
kernel_block_size is not None
and self.kv_cache_spec.block_size != kernel_block_size
and self.kv_cache_spec.block_size % kernel_block_size == 0
):
factor = self.kv_cache_spec.block_size // kernel_block_size
indexer_block_table = (block_table[:, ::factor] // factor).contiguous()
padded_num_tokens = num_tokens
local_slot_mapping = slot_mapping
if self.use_pcp:
# The gathered layout holds each rank's local tokens, padded, in
# rank order, so this rank's segment lines up with query_start_loc.
padded_num_tokens = slot_mapping.shape[0] // self.pcp_world_size
local_slot_mapping = slot_mapping[
self.pcp_rank * padded_num_tokens : (self.pcp_rank + 1)
* padded_num_tokens
]
compressed_slot_mapping = get_compressed_slot_mapping(
num_tokens,
local_slot_mapping,
query_start_loc,
seq_lens,
indexer_block_table,
self.kv_cache_spec.num_states,
self.compress_ratio,
out=self.compressed_slot_mapping_buffer,
)
if self.pcp_world_size > 1:
compressed_slot_mapping = get_pcp_group().all_gather(
self.compressed_slot_mapping_buffer[:padded_num_tokens],
dim=0,
)
# PCP decode sharding keeps a zero-token placeholder on ranks that own
# no request in a step so collectives retain a uniform rank shape. Do
# not turn that placeholder into a zero-length indexer decode request.
if num_tokens == 0:
return DeepseekV32IndexerMetadata(
seq_lens=seq_lens,
max_seq_len=common_attn_metadata.max_seq_len,
slot_mapping=compressed_slot_mapping,
num_decodes=0,
num_decode_tokens=0,
num_prefills=0,
num_prefill_tokens=0,
prefill=None,
decode=None,
)
num_decodes, num_prefills, num_decode_tokens, num_prefill_tokens = (
split_decodes_and_prefills(
common_attn_metadata,
decode_threshold=self.decode_threshold,
require_uniform=not (self.use_flattening or self.supports_varlen),
treat_short_extends_as_decodes=not self.use_pcp,
)
)
assert num_decodes + num_prefills == num_reqs
assert num_decode_tokens + num_prefill_tokens == num_tokens
prefill_metadata = None
if num_prefills > 0:
compressed_seq_lens = (
seq_lens // self.compress_ratio if self.compress_ratio > 1 else seq_lens
)
# This CPU value is an upper bound for async-spec extend rows. It
# is safe for chunking/allocation because CUDA metadata below is
# built from exact device seq_lens and gather ignores the tail.
assert common_attn_metadata.seq_lens_cpu_upper_bound is not None
seq_lens_cpu = common_attn_metadata.seq_lens_cpu_upper_bound
compressed_seq_lens_cpu = (
seq_lens_cpu // self.compress_ratio
if self.compress_ratio > 1
else seq_lens_cpu
)
prefill_query_lens_cpu = torch.diff(
query_start_loc_cpu[num_decodes : num_decodes + num_prefills + 1]
)
max_logits_bytes = envs.VLLM_SPARSE_INDEXER_MAX_LOGITS_MB * 1024 * 1024
# Upper bound is exact for prefill rows (the `[num_decodes:]`
# slice below).
assert common_attn_metadata.seq_lens_cpu_upper_bound is not None
seq_lens_cpu = common_attn_metadata.seq_lens_cpu_upper_bound
req_idx = None
shard_rows = None
if self.use_pcp and self.dcp_world_size > 1:
# The gathered KV must be packed identically on every PCP rank:
# chunk by request from its DCP shard rows, which every rank
# holds. A dummy batch bypasses the PCP manager and has one
# row per request, so its own extent is the request's.
req_idx = common_attn_metadata.req_idx
if req_idx is None:
req_idx = np.arange(num_reqs)
shard_rows_cpu = common_attn_metadata.dcp_local_seq_lens_cpu_upper_bound
if shard_rows_cpu is None:
shard_rows_cpu = get_dcp_local_seq_lens(
seq_lens_cpu,
self.dcp_world_size,
0,
self.cp_kv_cache_interleave_size,
)
shard_rows = shard_rows_cpu.numpy()
chunk_specs = self._split_pcp_dcp_prefill_chunks(
req_idx[num_decodes:],
shard_rows[num_decodes:],
prefill_query_lens_cpu,
max_logits_bytes,
request_offset=num_decodes,
)
else:
chunk_specs = self._split_indexer_prefill_chunks(
self._prefill_split_seq_lens(compressed_seq_lens_cpu[num_decodes:]),
prefill_query_lens_cpu,
self.max_prefill_buffer_size,
max_logits_bytes,
request_offset=num_decodes,
)
chunks = []
for req_slice, query_slice in chunk_specs:
pcp_plan = None
if req_idx is not None:
assert shard_rows is not None
pcp_plan = build_pcp_global_chunk_plan(
req_idx[req_slice],
shard_rows[req_slice],
self.dcp_world_size,
self.device,
self.cp_kv_cache_interleave_size,
)
metadata = build_prefill_chunk_metadata(
req_slice.start,
req_slice.stop,
query_start_loc,
query_start_loc_cpu,
seq_lens,
compressed_seq_lens,
compressed_seq_lens_cpu,
indexer_block_table,
self.compress_ratio,
query_slice=query_slice,
skip_kv_gather=query_slice.start > 0,
dcp_rank=self.dcp_rank,
dcp_world_size=self.dcp_world_size,
cp_kv_cache_interleave_size=self.cp_kv_cache_interleave_size,
pcp_plan=pcp_plan,
)
# Skip when total_seq_lens is 0 (i.e., no compressed token).
if metadata is not None:
chunks.append(metadata)
prefill_metadata = DeepseekV32IndexerPrefillMetadata(
chunks,
max_prefill_seq_len=(
int(seq_lens_cpu[num_decodes:].max().item())
if num_prefills > 0
else 0
),
)
decode_metadata = None
if num_decodes > 0:
if not self.supports_varlen:
torch.diff(
common_attn_metadata.query_start_loc[: num_decodes + 1],
out=self.decode_lens_buffer[:num_decodes],
)
self.per_req_decode_lens_buffer[:num_decodes].copy_(
self.decode_lens_buffer[:num_decodes]
)
decode_lens = self.decode_lens_buffer[:num_decodes]
decode_lens_cpu = torch.diff(
common_attn_metadata.query_start_loc_cpu[: num_decodes + 1]
)
# Under DCP the per-token decode bounds must be localized AFTER the
# per-token expansion below, not before. Expanding from a
# request-level localized length subtracts decode offsets in local
# space and yields too-short bounds (e.g. world=2, rank=1, global
# per-token bounds [8, 9, 10] -> [3, 4, 5] instead of [4, 4, 5]), so
# the first decode token would run top-k against too short a local KV
# range and miss valid tokens. Keep the global seq_lens here and
# localize the expanded bounds further down.
global_seq_lens_for_decode: torch.Tensor | None = None
if dcp_local_seq_lens is not None:
global_seq_lens_for_decode = common_attn_metadata.seq_lens[:num_decodes]
seq_lens = common_attn_metadata.seq_lens[:num_decodes]
block_table = common_attn_metadata.block_table_tensor[:num_decodes, ...]
max_decode_len = int(decode_lens_cpu.max().item())
min_decode_len = int(decode_lens_cpu.min().item())
write_is_uniform = min_decode_len == max_decode_len
next_n = 1 + self.num_speculative_tokens
# The kernel sees max_decode_len Q rows, not the configured next_n,
# so legality is per-step: on SM90 a uniformly 3-deep batch has no
# native kernel. max_decode_len <= 1 always has one.
step_next_n_ok = max_decode_len <= 1 or _supports_native_decode(
max_decode_len
)
use_native = (
not (self.use_flattening or self.supports_varlen)
and max_decode_len <= next_n
and step_next_n_ok
)
if not self.supports_varlen:
global_seq_lens_for_decode = self._prepare_global_decode_seq_lens(
global_seq_lens=global_seq_lens_for_decode,
decode_lens=decode_lens,
decode_lens_cpu=decode_lens_cpu,
query_start_loc=common_attn_metadata.query_start_loc[:num_decodes],
num_decode_tokens=num_decode_tokens,
use_native=use_native,
max_decode_len=max_decode_len,
)
decode_indices = None
if self.supports_varlen:
from vllm.v1.attention.ops.metadata import (
_indexer_decode_metadata_kernel,
)
capacity = self.decode_seq_lens_buffer.numel()
grid = max(
num_decodes,
num_decode_tokens + triton.cdiv(capacity - num_decode_tokens, 256),
)
_indexer_decode_metadata_kernel[(grid,)](
query_start_loc,
seq_lens,
block_table,
self.decode_seq_lens_buffer,
self.expanded_block_table_buffer,
self.decode_lens_buffer,
self.decode_indices_buffer,
self.per_req_decode_lens_buffer,
num_decodes,
num_decode_tokens,
capacity,
block_table.stride(0),
self.expanded_block_table_buffer.stride(0),
BLOCK_COLS=block_table.shape[1],
num_warps=4,
)
seq_lens = self.decode_seq_lens_buffer[:num_decode_tokens]
block_table = self.expanded_block_table_buffer[:num_decode_tokens]
decode_lens = self.decode_lens_buffer[:num_decode_tokens]
decode_indices = self.decode_indices_buffer[:num_decode_tokens]
requires_padding = False
if global_seq_lens_for_decode is not None and max_decode_len > 1:
self.global_decode_seq_lens_buffer[:num_decode_tokens].copy_(
seq_lens
)
global_seq_lens_for_decode = self.global_decode_seq_lens_buffer[
:num_decode_tokens
]
else:
seq_lens, block_table, decode_lens, batch_size, requires_padding = (
self._prepare_decode_tensors(
seq_lens=seq_lens,
block_table=block_table,
decode_lens=decode_lens,
decode_lens_cpu=decode_lens_cpu,
query_start_loc=common_attn_metadata.query_start_loc[
:num_decodes
],
num_decodes=num_decodes,
num_decode_tokens=num_decode_tokens,
use_native=use_native,
next_n=next_n,
max_decode_len=max_decode_len,
)
)
if self.compress_ratio > 1:
kernel_block_size = self.kernel_block_size
if (
kernel_block_size is not None
and self.kv_cache_spec.block_size != kernel_block_size
and self.kv_cache_spec.block_size % kernel_block_size == 0
):
factor = self.kv_cache_spec.block_size // kernel_block_size
compressed = block_table[:, ::factor] // factor
rows, cols = compressed.shape
if self.indexer_decode_block_table_buffer is None:
self.indexer_decode_block_table_buffer = torch.zeros(
(self._max_num_batched_tokens, cols),
dtype=torch.int32,
device=self.device,
)
self.indexer_decode_block_table_buffer[:rows, :cols].copy_(
compressed
)
block_table = self.indexer_decode_block_table_buffer[:rows, :cols]
# Flattening always returns a buffer view, including single-token
# batches. Keep its address stable across varlen graph replays.
seq_lens_is_buffer_view = not use_native or next_n > 1
# DCP: localize the now-expanded per-token global bounds to this
# rank's owned KV. Done here (after expansion) so each token's global
# causal length is localized individually; see the comment above.
if dcp_local_seq_lens is not None:
seq_lens = self._dcp_localize_decode_seq_lens(
seq_lens, num_decodes, seq_lens_is_buffer_view
)
# For DeepseekV4 (compress_ratio > 1), the indexer KV cache stores
# compressed tokens. Convert uncompressed seq_lens to compressed.
if self.compress_ratio > 1:
if seq_lens_is_buffer_view:
seq_lens //= self.compress_ratio
else:
# Copy to avoid mutating shared state; keeps CG address stable.
self.expanded_seq_lens_buffer[:num_decodes] = (
seq_lens // self.compress_ratio
)
self.expanded_seq_lens_buffer[num_decodes:num_decode_tokens] = 0
seq_lens = self.expanded_seq_lens_buffer[:num_decode_tokens]
# Non-MTP: deep_gemm paged MQA logits requires 2D context_lens
# (csrc/apis/attention.hpp). Unsqueeze to (B, 1) so downstream
# kernels see the same (B, next_n) layout as the MTP path.
if seq_lens.dim() == 1:
seq_lens = seq_lens.unsqueeze(-1)
# DeepGEMM is required for the paged MQA logits on CUDA devices
schedule_metadata = self.scheduler_metadata_buffer
if current_platform.is_cuda() and has_deep_gemm():
metadata = get_paged_mqa_logits_metadata(
seq_lens,
self.kv_cache_spec.num_states,
self.num_sms,
indices=decode_indices,
)
schedule_metadata = self.scheduler_metadata_buffer[: metadata.shape[0]]
schedule_metadata[:] = metadata
decode_metadata = DeepSeekV32IndexerDecodeMetadata(
block_table=block_table,
seq_lens=seq_lens,
decode_lens=decode_lens,
requires_padding=requires_padding,
schedule_metadata=schedule_metadata,
indices=decode_indices,
global_seq_lens=global_seq_lens_for_decode,
per_req_decode_lens=self.per_req_decode_lens_buffer[:num_decodes],
decode_is_uniform=write_is_uniform,
write_max_decode_len=max_decode_len,
)
attn_metadata = DeepseekV32IndexerMetadata(
seq_lens=common_attn_metadata.seq_lens,
max_seq_len=common_attn_metadata.max_seq_len,
slot_mapping=compressed_slot_mapping,
num_decodes=num_decodes,
num_decode_tokens=num_decode_tokens,
num_prefills=num_prefills,
num_prefill_tokens=num_prefill_tokens,
prefill=prefill_metadata,
decode=decode_metadata,
)
return attn_metadata