class BaseMambaAttentionMetadataBuilder(AttentionMetadataBuilder[M], abc.ABC):
kv_cache_spec: MambaSpec
metadata_cls: type[M]
reorder_batch_threshold: int = 1
_cudagraph_support: ClassVar[AttentionCGSupport] = AttentionCGSupport.UNIFORM_BATCH
# Will be disabled if speculative decoding is used
supports_update_block_table: bool = True
needs_causal_conv1d_metadata: bool = True
def __init__(
self,
kv_cache_spec: MambaSpec,
layer_names: list[str],
vllm_config: VllmConfig,
device: torch.device,
):
super().__init__(kv_cache_spec, layer_names, vllm_config, device)
# Enable speculative decoding support
self.speculative_config = vllm_config.speculative_config
self.compilation_config = vllm_config.compilation_config
self.num_spec_tokens: int = vllm_config.num_speculative_tokens
self.use_spec_decode = self.num_spec_tokens > 0
self.use_replayssm = vllm_config.cache_config.use_replayssm
self.replayssm_buffer_len = vllm_config.cache_config.replayssm_buffer_len
self.use_flashinfer_replayssm = (
self.use_replayssm
and vllm_config.mamba_config.backend == MambaBackendEnum.FLASHINFER
)
scheduler_config = vllm_config.scheduler_config
self.decode_cudagraph_max_bs: int = scheduler_config.max_num_seqs
if self.compilation_config.max_cudagraph_capture_size is not None:
self.decode_cudagraph_max_bs = min(
self.decode_cudagraph_max_bs,
self.compilation_config.max_cudagraph_capture_size,
)
self.state_indices_tensor_d = torch.empty(
(self.decode_cudagraph_max_bs, 1 + self.num_spec_tokens),
dtype=torch.int32,
device=device,
)
# For speculative decoding, we need to store the following buffers
# for CUDA graph capture during decode
if self.num_spec_tokens > 0:
self.decode_num_accepted_tokens: torch.Tensor = torch.empty(
(self.decode_cudagraph_max_bs,),
dtype=torch.int32,
device=device,
)
self.decode_bc_pre_scratch: torch.Tensor | None = None
self.decode_replayssm_scratch: (
tuple[torch.Tensor, torch.Tensor, torch.Tensor] | None
) = None
self.decode_replayssm_state_indices_d: torch.Tensor | None = None
# ReplaySSM CUDA-graph buffers for the selected backend.
if self.use_replayssm and not self.use_flashinfer_replayssm:
self.decode_write_pos_d: torch.Tensor = torch.empty(
(self.decode_cudagraph_max_bs,),
dtype=torch.int32,
device=device,
)
self.decode_is_flush_d: torch.Tensor = torch.empty(
(self.decode_cudagraph_max_bs,),
dtype=torch.int8,
device=device,
)
# B_cache shape = (ngroups, replayssm_buffer_len, dstate); the page
# layout is (conv_state, ssm_state, x_cache, dt_cache, B_cache).
bc_ngroups = kv_cache_spec.shapes[4][0]
bc_scratch_bs = max(
self.decode_cudagraph_max_bs, scheduler_config.max_num_seqs
)
self.decode_bc_pre_scratch = torch.empty(
(
bc_scratch_bs,
bc_ngroups,
self.replayssm_buffer_len,
),
dtype=torch.float32,
device=device,
)
elif self.use_flashinfer_replayssm:
from flashinfer.mamba.checkpointing_ssu import (
allocate_checkpointing_ssu_scratch,
)
nheads = kv_cache_spec.shapes[2][0]
self.decode_replayssm_scratch = allocate_checkpointing_ssu_scratch(
batch_size=scheduler_config.max_num_seqs,
num_heads=nheads,
num_predicted_tokens=1 + self.num_spec_tokens,
max_window=self.replayssm_buffer_len,
dtype=vllm_config.model_config.dtype,
device=device,
)
# Full CUDA graphs retain capture-time tensor addresses. Keep the
# contiguous first-column view used by FlashInfer in a persistent
# buffer and refresh its contents before each replay.
self.decode_replayssm_state_indices_d = torch.empty(
(self.decode_cudagraph_max_bs,), dtype=torch.int32, device=device
)
self._init_reorder_batch_threshold(1, self.use_spec_decode)
if self.use_spec_decode:
self.supports_update_block_table = False
def build_for_cudagraph_capture(
self, common_attn_metadata: CommonAttentionMetadata
) -> M:
"""This method builds the metadata for full cudagraph capture.
Currently, only decode is supported for full cudagraphs with Mamba.
"""
m = common_attn_metadata
assert (
m.max_query_len <= 1 + self.num_spec_tokens
and m.num_reqs <= self.decode_cudagraph_max_bs
), (
"Mamba only supports decode-only full CUDAGraph capture. "
"Make sure all cudagraph capture sizes <= max_num_seq."
)
assert m.max_query_len == 1 + self.num_spec_tokens # decode-only
num_accepted_tokens = None
if self.num_spec_tokens > 0:
num_accepted_tokens = torch.diff(m.query_start_loc)
return self.build(0, m, num_accepted_tokens=num_accepted_tokens)
def build(
self,
common_prefix_len: int,
common_attn_metadata: CommonAttentionMetadata,
fast_build: bool = False,
*,
num_accepted_tokens: torch.Tensor | None = None,
num_decode_draft_tokens_cpu: torch.Tensor | None = None,
**kwargs: Any,
) -> M:
"""Default build implementation for Mamba-like attention backends.
Subclasses (e.g., Mamba2) can override to add additional metadata.
"""
return self._compute_common_metadata(
common_attn_metadata,
num_accepted_tokens=num_accepted_tokens,
num_decode_draft_tokens_cpu=num_decode_draft_tokens_cpu,
)
def _compute_chunk_metadata(
self,
chunk_size: int,
num_prefills: int,
num_computed_tokens_p_cpu: torch.Tensor,
query_start_loc_p_cpu: torch.Tensor,
) -> tuple[list[int], list[int], list[int]]:
"""Compute chunk-specific metadata for Mamba models.
The code below carefully constructs the chunks such that:
1. Chunks contain tokens from a *single* sequence only.
2. For every sequence, we are guaranteed that we can
retrieve the mamba state *every* chunk_size tokens.
Constraint (1) dramatically simplifies the mamba kernels.
Constraint (2) dramatically simplifies the implementation
of prefix caching for mamba (wip). We need to take care
of the interaction with chunked prefill in order to
satisfy constraint (2).
"""
# TODO (tdoublep): This code could probably be optimized.
cu_chunk_seqlen = []
seq_idx = []
last_chunk_indices = []
seqlen_pos = 0
for req_idx in range(num_prefills):
this_num_computed = num_computed_tokens_p_cpu[req_idx].item()
this_new_tokens = (
query_start_loc_p_cpu[req_idx + 1].item()
- query_start_loc_p_cpu[req_idx].item()
)
# if computed tokens are not chunk-aligned, use the first
# chunk to finish it off
if this_num_computed % chunk_size != 0:
seq_idx.append(req_idx)
cu_chunk_seqlen.append(seqlen_pos)
# how many tokens to finish the chunk?
chunk_len = (
cdiv(this_num_computed, chunk_size) * chunk_size - this_num_computed
)
# we can only use at most this_new_tokens
chunk_len = min(chunk_len, this_new_tokens)
seqlen_pos += chunk_len
this_new_tokens -= chunk_len
n_chunks = cdiv(this_new_tokens, chunk_size)
for chunk in range(n_chunks):
seq_idx.append(req_idx)
cu_chunk_seqlen.append(seqlen_pos)
chunk_len = min(chunk_size, this_new_tokens)
seqlen_pos += chunk_len
this_new_tokens -= chunk_len
assert this_new_tokens == 0
last_chunk_indices.append(len(cu_chunk_seqlen) - 1)
cu_chunk_seqlen.append(seqlen_pos)
return cu_chunk_seqlen, seq_idx, last_chunk_indices
def _prefill_cpu_metadata(
self,
common_attn_metadata: CommonAttentionMetadata,
num_reqs: int,
num_prefills: int,
num_decode_tokens: int,
) -> tuple[torch.Tensor, torch.Tensor]:
"""Prefill context lengths and query offsets, from CPU data only.
`seq_lens_cpu_upper_bound` is precise for prefill rows in all modes
(including async spec decode), so this avoids the D2H sync that
`compute_num_computed_tokens().cpu()` would force.
Returns (num_computed_tokens_p_cpu, query_start_loc_p_cpu).
"""
seq_lens_cpu = common_attn_metadata.seq_lens_cpu_upper_bound
assert seq_lens_cpu is not None
query_start_loc_p_cpu = (
common_attn_metadata.query_start_loc_cpu[-num_prefills - 1 :]
- num_decode_tokens
)
prefill_query_lens_cpu = query_start_loc_p_cpu[1:] - query_start_loc_p_cpu[:-1]
num_computed_tokens_p_cpu = (
seq_lens_cpu[num_reqs - num_prefills : num_reqs] - prefill_query_lens_cpu
)
return num_computed_tokens_p_cpu, query_start_loc_p_cpu
def _build_chunk_metadata_tensors(
self,
chunk_size: int,
common: M,
common_attn_metadata: CommonAttentionMetadata,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
"""Compute chunk metadata and return as device tensors.
Returns (cu_chunk_seqlen_p, seq_idx_p, last_chunk_indices_p).
"""
num_prefills = common.num_prefills
num_computed_tokens_p_cpu, query_start_loc_p_cpu = self._prefill_cpu_metadata(
common_attn_metadata,
common.num_reqs,
num_prefills,
common.num_decode_tokens,
)
cu_chunk_seqlen, seq_idx, last_chunk_indices = self._compute_chunk_metadata(
chunk_size,
num_prefills,
num_computed_tokens_p_cpu,
query_start_loc_p_cpu,
)
device = common_attn_metadata.query_start_loc.device
# Build on pinned CPU and upload non-blocking to avoid the synchronous
# H2D copy that `torch.as_tensor(list, device=cuda)` would force.
cu_chunk_seqlen_p = async_tensor_h2d(
cu_chunk_seqlen, dtype=torch.int32, device=device
)
seq_idx_p = async_tensor_h2d(seq_idx, dtype=torch.int32, device=device)
last_chunk_indices_p = async_tensor_h2d(
last_chunk_indices, dtype=torch.int32, device=device
)
return cu_chunk_seqlen_p, seq_idx_p, last_chunk_indices_p
def _compute_common_metadata(
self,
common_attn_metadata: CommonAttentionMetadata,
*,
num_accepted_tokens: torch.Tensor | None = None,
num_decode_draft_tokens_cpu: torch.Tensor | None = None,
) -> M:
"""Compute metadata common to both Mamba1 and Mamba2."""
num_reqs = common_attn_metadata.num_reqs
# Treat multi-token queries as decode requests when
# speculative decoding is enabled. Otherwise, use the
# default decode threshold to prevent misclassification
# of prefill queries as decode requests.
decode_threshold = (
self.reorder_batch_threshold if num_accepted_tokens is not None else 1
)
# FULL-CG dispatch is shape-based, so one-token prefills with
# prior Mamba state can replay a decode graph while `is_prefilling`
# is still true. Treat them as decode/update rows. This is required
# for NIXL disagg's h(N-1)->N recompute path and for sporadic
# final single-token prefill chunks that land in a `uniform` FULL-CG
# batch. Relies on `reorder` putting short extends before pure prefills.
is_prefilling = common_attn_metadata.is_prefilling
assert is_prefilling is not None
seq_lens_cpu = common_attn_metadata.seq_lens_cpu_upper_bound
assert seq_lens_cpu is not None
query_lens_cpu = torch.diff(common_attn_metadata.query_start_loc_cpu)
# First prompt chunks have no prior Mamba state and must stay prefills.
has_prior_state = seq_lens_cpu > query_lens_cpu
stateful_prefill_rows = is_prefilling & has_prior_state
# One-token prefills with prior state can use the decode/update path.
prefill_to_decode = stateful_prefill_rows & (query_lens_cpu == 1)
# The scheduler may pad a one-token remote prompt tail with placeholder
# drafts to retain the uniform K+1 decode graph. This is a speculative
# decode transaction even though the real token is still in the prompt:
# the decode kernels keep h(N) in the running slot and h(N+i) in scratch
# slots, so normal acceptance rollback remains valid. The prefill kernels
# only return h(N+K) and cannot roll the placeholders back.
if num_decode_draft_tokens_cpu is not None:
padded_prompt_tail_rows = (
stateful_prefill_rows
& (num_decode_draft_tokens_cpu >= 0)
& (query_lens_cpu == num_decode_draft_tokens_cpu + 1)
)
prefill_to_decode |= padded_prompt_tail_rows
if torch.any(prefill_to_decode).item():
# ReplaySSM handles these rows as single-token flushes (see the
# write-position derivation below), same as the baseline decode path.
is_prefilling = is_prefilling.clone()
is_prefilling[prefill_to_decode] = False
common_attn_metadata = common_attn_metadata.replace(
is_prefilling=is_prefilling
)
num_decodes, num_prefills, num_decode_tokens, num_prefill_tokens = (
split_decodes_and_prefills(
common_attn_metadata,
decode_threshold=decode_threshold,
treat_short_extends_as_decodes=False,
)
)
# Need flags to indicate if there are initial states
has_initial_states_p = None
query_start_loc_p = None
query_start_loc_d = None
# for causal_conv1d
nums_dict, batch_ptr, token_chunk_offset_ptr = None, None, None
write_pos_d = None
is_flush_d = None
replayssm_scratch = None
state_indices_tensor = mamba_get_block_table_tensor(
common_attn_metadata.block_table_tensor,
common_attn_metadata.seq_lens,
self.kv_cache_spec,
self.vllm_config.cache_config.mamba_cache_mode,
)
if state_indices_tensor.dim() == 1:
state_indices_tensor = state_indices_tensor.unsqueeze(-1)
state_indices_tensor_d, state_indices_tensor_p = torch.split(
state_indices_tensor,
[num_decodes, num_prefills],
dim=0,
)
state_indices_tensor_d = state_indices_tensor_d[:, : 1 + self.num_spec_tokens]
state_indices_tensor_p = state_indices_tensor_p[:, 0]
if num_decodes > 0 and self.use_spec_decode:
query_start_loc_d = common_attn_metadata.query_start_loc[: num_decodes + 1]
if num_accepted_tokens is None:
# Single-token prefill chunks can be reclassified as decodes before
# speculative decoding has produced acceptance counts. Treat each
# token as accepted so recurrent state and ReplaySSM trackers follow
# the normal speculative-decode path.
num_accepted_tokens = torch.diff(query_start_loc_d)
else:
num_accepted_tokens = num_accepted_tokens[:num_decodes]
if num_prefills > 0:
num_computed_tokens = common_attn_metadata.compute_num_computed_tokens()
query_start_loc_p = (
common_attn_metadata.query_start_loc[-num_prefills - 1 :]
- num_decode_tokens
)
has_initial_states_p = (
num_computed_tokens[num_reqs - num_prefills : num_reqs] > 0
)
if self.needs_causal_conv1d_metadata:
query_start_loc_p_cpu = (
common_attn_metadata.query_start_loc_cpu[-num_prefills - 1 :]
- num_decode_tokens
)
nums_dict, batch_ptr, token_chunk_offset_ptr = (
compute_causal_conv1d_metadata(
query_start_loc_p_cpu,
device=common_attn_metadata.query_start_loc.device,
)
)
if self.use_replayssm and not self.use_flashinfer_replayssm and num_decodes > 0:
decode_base_cpu = common_attn_metadata.replayssm_decode_base_cpu
seq_lens_cpu = common_attn_metadata.seq_lens_cpu_upper_bound
async_spec_decode = (
self.vllm_config.scheduler_config.async_scheduling
and self.vllm_config.speculative_config is not None
)
if decode_base_cpu is None or seq_lens_cpu is None or async_spec_decode:
raise ValueError(
"--use-replayssm requires exact CPU sequence lengths and "
"decode-base counts to derive decode write positions"
)
query_lens_cpu = (
common_attn_metadata.query_start_loc_cpu[1 : num_decodes + 1]
- common_attn_metadata.query_start_loc_cpu[:num_decodes]
)
num_computed_d = seq_lens_cpu[:num_decodes] - query_lens_cpu
decode_base_d = decode_base_cpu[:num_decodes]
align_mode = self.vllm_config.cache_config.mamba_cache_mode == "align"
block_size = self.kv_cache_spec.block_size
if align_mode:
# After a boundary the align copy leaves an exact checkpoint at
# the block start and the new block's ring restarts empty, so
# re-anchor there; max() keeps the prompt-end anchor for the
# first (partial) block.
effective_base = torch.maximum(
decode_base_d, (num_computed_d // block_size) * block_size
)
else:
effective_base = decode_base_d
# write_pos counts decode steps since the ring's last full-state
# write (the anchor), so a resumed request re-anchors correctly.
decode_steps_cpu = num_computed_d - effective_base
valid_decode_rows = query_lens_cpu > 0
# A single-token prefill row replayed as decode (query_len==1 with
# prior state) has decode_steps < 0; force it to a one-token flush
# (write_pos=0, is_flush=1). The flush branch reads an empty history
# window, so it applies exactly one recurrence step off the checkpoint
# -- identical to the baseline decode kernel for that row. The split
# (treat_short_extends_as_decodes=False) admits only such rows here.
leftover_prompt = valid_decode_rows & (decode_steps_cpu < 0)
decode_steps_cpu = torch.where(
valid_decode_rows & ~leftover_prompt,
decode_steps_cpu,
torch.zeros_like(decode_steps_cpu),
)
write_pos_cpu = torch.remainder(decode_steps_cpu, self.replayssm_buffer_len)
is_flush_cpu = (
write_pos_cpu == self.replayssm_buffer_len - 1
) | leftover_prompt
if align_mode:
# Force a flush on the step completing a mamba block so the exact
# boundary state is materialized for prefix caching.
is_flush_cpu = is_flush_cpu | (
valid_decode_rows
& ((num_computed_d + query_lens_cpu) % block_size == 0)
)
is_flush_cpu = is_flush_cpu.to(torch.int8)
write_pos_d = async_tensor_h2d(
write_pos_cpu.to(torch.int32).tolist(),
dtype=torch.int32,
device=common_attn_metadata.query_start_loc.device,
)
is_flush_d = async_tensor_h2d(
is_flush_cpu.tolist(),
dtype=torch.int8,
device=common_attn_metadata.query_start_loc.device,
)
if self.use_flashinfer_replayssm and num_decodes > 0:
assert self.decode_replayssm_scratch is not None
cb_scaled, cumAdt_vec, cb_old = self.decode_replayssm_scratch
replayssm_scratch = (
cb_scaled[:num_decodes],
cumAdt_vec[:num_decodes],
cb_old[:num_decodes],
)
bc_pre_scratch = None
if (
self.use_replayssm
and self.decode_bc_pre_scratch is not None
and num_decodes > 0
):
bc_pre_scratch = self.decode_bc_pre_scratch[:num_decodes]
metadata = self.metadata_cls(
num_prefills=num_prefills,
num_prefill_tokens=num_prefill_tokens,
num_decodes=num_decodes,
num_decode_tokens=num_decode_tokens,
query_start_loc_p=query_start_loc_p,
has_initial_states_p=has_initial_states_p,
state_indices_tensor_p=state_indices_tensor_p,
state_indices_tensor_d=state_indices_tensor_d,
write_pos_d=write_pos_d,
is_flush_d=is_flush_d,
bc_pre_scratch=bc_pre_scratch,
replayssm_scratch=replayssm_scratch,
num_accepted_tokens=num_accepted_tokens,
query_start_loc_d=query_start_loc_d,
num_reqs=num_reqs,
seq_lens=common_attn_metadata.seq_lens,
nums_dict=nums_dict,
batch_ptr=batch_ptr,
token_chunk_offset_ptr=token_chunk_offset_ptr,
)
return self._update_metadata_for_cudagraph_capture(metadata)
def _update_metadata_for_cudagraph_capture(
self,
metadata: M,
) -> M:
"""Update the metadata for cudagraph capture.
Currently, only decode is supported for full cudagraphs with Mamba.
"""
state_indices_tensor_d = metadata.state_indices_tensor_d
query_start_loc_d = metadata.query_start_loc_d
num_accepted_tokens = metadata.num_accepted_tokens
write_pos_d = metadata.write_pos_d
is_flush_d = metadata.is_flush_d
bc_pre_scratch = metadata.bc_pre_scratch
replayssm_scratch = metadata.replayssm_scratch
replayssm_state_indices_d = None
if (
metadata.num_prefills == 0
and metadata.num_decodes <= self.decode_cudagraph_max_bs
and self.compilation_config.cudagraph_mode.has_full_cudagraphs()
):
padded_bs = metadata.num_reqs
self.state_indices_tensor_d[: metadata.num_decodes].copy_(
state_indices_tensor_d, non_blocking=True
)
state_indices_tensor_d = self.state_indices_tensor_d[:padded_bs]
state_indices_tensor_d[metadata.num_decodes :] = NULL_BLOCK_ID
if self.use_spec_decode and num_accepted_tokens is not None:
assert query_start_loc_d is not None
query_start_loc_d = query_start_loc_d[: padded_bs + 1]
self.decode_num_accepted_tokens[: metadata.num_decodes].copy_(
num_accepted_tokens, non_blocking=True
)
num_accepted_tokens = self.decode_num_accepted_tokens[:padded_bs]
num_accepted_tokens[metadata.num_decodes :] = (
1 # pad with 1st slot index
)
if self.use_replayssm and not self.use_flashinfer_replayssm:
assert write_pos_d is not None
assert is_flush_d is not None
self.decode_write_pos_d[: metadata.num_decodes].copy_(
write_pos_d[: metadata.num_decodes],
non_blocking=True,
)
write_pos_d = self.decode_write_pos_d[:padded_bs]
write_pos_d[metadata.num_decodes :] = 0
self.decode_is_flush_d[: metadata.num_decodes].copy_(
is_flush_d[: metadata.num_decodes],
non_blocking=True,
)
is_flush_d = self.decode_is_flush_d[:padded_bs]
is_flush_d[metadata.num_decodes :] = 0
if self.decode_bc_pre_scratch is not None:
bc_pre_scratch = self.decode_bc_pre_scratch[:padded_bs]
elif self.use_flashinfer_replayssm:
assert self.decode_replayssm_scratch is not None
cb_scaled, cumAdt_vec, cb_old = self.decode_replayssm_scratch
replayssm_scratch = (
cb_scaled[:padded_bs],
cumAdt_vec[:padded_bs],
cb_old[:padded_bs],
)
assert self.decode_replayssm_state_indices_d is not None
self.decode_replayssm_state_indices_d[:padded_bs].copy_(
state_indices_tensor_d[:, 0], non_blocking=True
)
replayssm_state_indices_d = self.decode_replayssm_state_indices_d[
:padded_bs
]
if (
self.use_flashinfer_replayssm
and state_indices_tensor_d is not None
and replayssm_state_indices_d is None
):
replayssm_state_indices_d = state_indices_tensor_d[:, 0].contiguous()
return replace(
metadata,
state_indices_tensor_d=state_indices_tensor_d,
query_start_loc_d=query_start_loc_d,
num_accepted_tokens=num_accepted_tokens,
write_pos_d=write_pos_d,
is_flush_d=is_flush_d,
bc_pre_scratch=bc_pre_scratch,
replayssm_scratch=replayssm_scratch,
replayssm_state_indices_d=replayssm_state_indices_d,
)
def update_block_table(
self,
metadata: M,
blk_table: torch.Tensor,
slot_mapping: torch.Tensor,
) -> M:
state_indices_tensor = mamba_get_block_table_tensor(
blk_table,
metadata.seq_lens,
self.kv_cache_spec,
self.vllm_config.cache_config.mamba_cache_mode,
)
if state_indices_tensor.dim() == 1:
state_indices_tensor = state_indices_tensor.unsqueeze(-1)
assert (
metadata.num_prefills + metadata.num_decodes
== state_indices_tensor.shape[0]
), (
"Mismatch in number of requests when updating block table."
f" Expected {metadata.num_prefills + metadata.num_decodes}, "
f"got {state_indices_tensor.shape[0]}."
)
state_indices_tensor_d, state_indices_tensor_p = torch.split(
state_indices_tensor,
[metadata.num_decodes, metadata.num_prefills],
dim=0,
)
state_indices_tensor_d = state_indices_tensor_d[:, : 1 + self.num_spec_tokens]
state_indices_tensor_p = state_indices_tensor_p[:, 0]
new_metadata = replace(
metadata,
state_indices_tensor_d=state_indices_tensor_d,
state_indices_tensor_p=state_indices_tensor_p,
)
return self._update_metadata_for_cudagraph_capture(new_metadata)