class MambaHybridModelState(DefaultModelState):
"""Model state for hybrid attention + Mamba / linear-attention models."""
def __init__(
self,
vllm_config: VllmConfig,
model: nn.Module,
encoder_cache: EncoderCache | None,
device: torch.device,
) -> None:
super().__init__(vllm_config, model, encoder_cache, device)
self.cache_config = vllm_config.cache_config
self.num_accepted_tokens_gpu = torch.ones(
self.max_num_reqs, dtype=torch.int32, device=self.device
)
# Pre-copy "align" prefix-cache state (V2). The migration of each
# request's mamba state across block boundaries runs as a fused GPU
# kernel reusing the postprocess copy machinery, so the per-step src
# columns and the running state_idx are kept GPU-resident.
self._align_mode = self.cache_config.mamba_cache_mode == "align"
self.recoverssm = (
RecoverSSMState() if self.cache_config.use_kda_recoverssm else None
)
if self._align_mode:
self._mamba_state_idx_gpu = torch.zeros(
self.max_num_reqs, dtype=torch.int32, device=self.device
)
self._mamba_src_col_gpu = torch.full(
(self.max_num_reqs,), -1, dtype=torch.int32, device=self.device
)
self._mamba_src_off_gpu = torch.zeros(
self.max_num_reqs, dtype=torch.int32, device=self.device
)
self._mamba_ctx: MambaSpecDecodeGPUContext | None = None
self._mamba_group_ids: list[int] = []
self._mamba_spec: MambaSpec | None = None
self._mamba_state_copy_funcs: MambaStateCopyFuncsByType | None = None
def add_request(self, req_index: int, new_req_data: NewRequestData) -> None:
super().add_request(req_index, new_req_data)
# Must reset the speculative acceptance count in this idx which could be stale.
self.num_accepted_tokens_gpu[req_index].fill_(1)
if self._align_mode:
# Seed the running state block from the resumed/prefilled position.
self._mamba_state_idx_gpu[req_index].fill_(
(new_req_data.num_computed_tokens - 1) // self.cache_config.block_size
)
def _get_mamba_group_info(
self, kv_cache_config: KVCacheConfig
) -> tuple[list[int], MambaSpec]:
if self._mamba_spec is None:
mamba_groups = get_mamba_groups(kv_cache_config)
mamba_spec = next(iter(mamba_groups))
assert all(
spec.block_size == mamba_spec.block_size
and spec.num_speculative_blocks == mamba_spec.num_speculative_blocks
and spec.mamba_cache_mode == mamba_spec.mamba_cache_mode
for spec in mamba_groups
), "all mamba groups must share cache scheduling parameters"
self._mamba_group_ids = get_mamba_group_ids(mamba_groups)
self._mamba_spec = mamba_spec
return self._mamba_group_ids, self._mamba_spec
def _ensure_align_ctx(
self,
kv_cache_config: KVCacheConfig,
mamba_group_ids: list[int],
block_tables: tuple[torch.Tensor, ...],
) -> MambaSpecDecodeGPUContext:
if self._mamba_state_copy_funcs is None:
mamba_groups = get_mamba_groups(kv_cache_config)
mamba_types = {spec.mamba_type for spec in mamba_groups}
copy_funcs = self.model.get_mamba_state_copy_funcs(mamba_types)
validate_mamba_state_copy_funcs(mamba_groups, copy_funcs)
self._mamba_state_copy_funcs = copy_funcs
copy_funcs = self._mamba_state_copy_funcs
if self._mamba_ctx is None:
# Both SD and DS conv layouts support a >0 spec-decode shift: the
# fused pre-copy kernel (``_copy_mamba_state_block``) applies the
# ``token_bias = num_accepted - 1`` window shift per conv layout
# (SD: contiguous slice; DS: per-dim-row strided slice), matching
# the V1 ``get_conv_copy_spec`` semantics.
self._mamba_ctx = MambaSpecDecodeGPUContext.create(
max_num_reqs=self.max_num_reqs,
kv_cache_config=kv_cache_config,
copy_funcs=copy_funcs,
device=self.device,
make_buffer=lambda n, dtype: CpuGpuBuffer(
n, dtype=dtype, device=self.device
),
)
ctx = self._mamba_ctx
if not ctx.is_initialized:
forward_context = self.vllm_config.compilation_config.static_forward_context
# block_tables are batch-order slices of the persistent
# input_block_tables (stable data_ptr), so the metadata is captured
# once here and reused across steps.
ctx.initialize_from_forward_context(
kv_cache_config,
forward_context,
copy_funcs,
[block_tables[gid] for gid in mamba_group_ids],
)
return ctx
def preprocess_state(
self,
input_batch: InputBatch,
block_tables: tuple[torch.Tensor, ...],
kv_cache_config: KVCacheConfig,
num_computed_tokens: torch.Tensor,
) -> None:
"""Migrate each request's mamba state across block boundaries before the
forward (V1 align semantics, done on GPU). Runs on real batches only
(dummy DP/profiling runs skip preprocess_state), and before
``prepare_attn`` gathers ``num_accepted_tokens``, so the boundary reset
is visible to the forward kernels.
"""
if not self._align_mode:
return
num_reqs = input_batch.num_reqs
if num_reqs == 0:
return
mamba_group_ids, mamba_spec = self._get_mamba_group_info(kv_cache_config)
ctx = self._ensure_align_ctx(kv_cache_config, mamba_group_ids, block_tables)
# The state-advance + pre-copy kernels run every step; they fast-exit per
# request when src_col < 0 or src_col == dst_col, so no copy happens on
# steps that don't cross a block boundary. (Skipping the launch entirely
# would need a V1-style async-D2H of the actual num_computed, since
# num_computed_tokens_np is an optimistic mirror under async scheduling;
# the launch cost is ~0.3% of TPOT, so the GPU fast-exit suffices.)
block = 256
grid = (triton.cdiv(num_reqs, block),)
preprocess_mamba_align_fused_kernel[grid](
input_batch.idx_mapping,
self._mamba_state_idx_gpu,
num_computed_tokens,
input_batch.query_start_loc,
self.num_accepted_tokens_gpu,
self._mamba_src_col_gpu,
self._mamba_src_off_gpu,
num_reqs,
BLOCK_SIZE=block,
MAMBA_BLOCK_SIZE=mamba_spec.block_size,
)
ctx.run_fused_precopy(
num_reqs,
self._mamba_state_idx_gpu,
self._mamba_src_col_gpu,
self._mamba_src_off_gpu,
input_batch.idx_mapping,
)
def prepare_attn(
self,
input_batch: InputBatch,
cudagraph_mode: CUDAGraphMode,
block_tables: tuple[torch.Tensor, ...],
slot_mappings: torch.Tensor,
attn_groups: list[list[AttentionGroup]],
kv_cache_config: KVCacheConfig,
for_capture: bool = False,
ubatch_idx: int = 0,
model_specific_attn_metadata: ModelSpecificAttnMetadata | None = None,
) -> dict[str, Any]:
assert ubatch_idx == 0, "DBO is not supported"
assert model_specific_attn_metadata is None
if cudagraph_mode == CUDAGraphMode.FULL:
num_reqs = input_batch.num_reqs_after_padding
num_tokens = input_batch.num_tokens_after_padding
else:
num_reqs = input_batch.num_reqs
num_tokens = input_batch.num_tokens
query_start_loc_cpu = torch.from_numpy(input_batch.query_start_loc_np)
# Prefer the promised bound: a capture dummy's measured max is its even
# split, not the length the graph must replay.
max_query_len = input_batch.max_query_len
if max_query_len is None:
max_query_len = input_batch.num_scheduled_tokens.max().item()
seq_lens_cpu_upper_bound = input_batch.seq_lens_cpu_upper_bound
if for_capture:
# Capture with worst-case max_seq_len so the graph is valid at any replay.
max_seq_len = self.max_model_len
else:
max_seq_len = seq_lens_cpu_upper_bound[:num_reqs].max().item()
is_prefilling_np = input_batch.is_prefilling_np
if input_batch.prefill_runs_as_decode_np is not None:
# A prompt tail the scheduler padded with placeholder drafts must run
# as a spec-decode row: the prefill kernels can't roll them back.
is_prefilling_np = is_prefilling_np & ~input_batch.prefill_runs_as_decode_np
is_prefilling = torch.zeros(num_reqs, dtype=torch.bool, device="cpu")
is_prefilling[: input_batch.num_reqs] = torch.from_numpy(is_prefilling_np)
# During CUDAGraph capture, num_decode_draft_tokens_cpu and num_accepted_tokens
# are created by attn_metadata_builder.build_for_cudagraph_capture, so we only
# compute them during actual (non-capture) forward execution.
num_accepted_tokens = None
num_decode_draft_tokens_cpu = None
if not for_capture and self.vllm_config.num_speculative_tokens > 0:
num_accepted_tokens = self.num_accepted_tokens_gpu.new_ones(num_reqs)
num_accepted_tokens[: input_batch.num_reqs] = self.num_accepted_tokens_gpu[
input_batch.idx_mapping
]
# GDN uses >= 0 to select spec-decode rows, so non-decode rows
# need the -1 sentinel rather than a raw zero draft count.
num_decode_draft_tokens_np = np.full(num_reqs, -1, dtype=np.int32)
num_draft_tokens_per_req = input_batch.num_draft_tokens_per_req
if num_draft_tokens_per_req is not None:
# Test request state, not num_scheduled_tokens == draft_count+1:
# adaptive rewrites num_scheduled_tokens to an even split, so that
# equality rarely holds and would demote every verify row to decode.
is_decode = ~is_prefilling_np & (input_batch.num_scheduled_tokens > 0)
spec_decode_mask = (num_draft_tokens_per_req > 0) & is_decode
num_decode_draft_tokens_np[: input_batch.num_reqs] = np.where(
spec_decode_mask, num_draft_tokens_per_req, -1
)
num_decode_draft_tokens_cpu = torch.from_numpy(num_decode_draft_tokens_np)
if self._align_mode:
mamba_group_ids, _ = self._get_mamba_group_info(kv_cache_config)
aligned_index_builders = []
for group_idx, group_id in enumerate(mamba_group_ids):
for group in attn_groups[group_id]:
builder = group.get_metadata_builder(0)
if hasattr(builder, "mamba_aligned_state_indices"):
aligned_index_builders.append((group_idx, builder))
if aligned_index_builders:
ctx = self._ensure_align_ctx(
kv_cache_config, mamba_group_ids, block_tables
)
all_group_indices = ctx.compute_aligned_state_indices(
input_batch.seq_lens, num_reqs
)
for group_idx, builder in aligned_index_builders:
builder.mamba_aligned_state_indices = all_group_indices[group_idx]
mamba_attn_metadata = MambaHybridAttnMetadata(
is_prefilling=is_prefilling,
num_accepted_tokens=num_accepted_tokens,
num_decode_draft_tokens_cpu=num_decode_draft_tokens_cpu,
)
attn_metadata = build_attn_metadata(
attn_groups=attn_groups,
num_reqs=num_reqs,
num_tokens=num_tokens,
query_start_loc_gpu=input_batch.query_start_loc,
query_start_loc_cpu=query_start_loc_cpu,
max_query_len=max_query_len,
seq_lens=input_batch.seq_lens,
max_seq_len=max_seq_len,
block_tables=block_tables,
slot_mappings=slot_mappings,
kv_cache_config=kv_cache_config,
seq_lens_cpu_upper_bound=seq_lens_cpu_upper_bound,
dcp_local_seq_lens=input_batch.dcp_local_seq_lens,
positions=input_batch.positions,
model_specific_attn_metadata=mamba_attn_metadata,
for_cudagraph_capture=for_capture,
rswa_prefix_lens=input_batch.prompt_lens,
)
if self.recoverssm is not None:
self.recoverssm.record_step(
attn_metadata, attn_groups, for_capture=for_capture
)
return attn_metadata
def postprocess_state(
self,
idx_mapping: torch.Tensor,
num_sampled: torch.Tensor | int,
num_computed_tokens: torch.Tensor | None = None,
) -> None:
# Chunked prefill does not sample a token, so num_sampled can be 0.
# Mamba treats num_accepted_tokens=1 as the neutral non-spec value.
num_reqs = idx_mapping.shape[0]
if num_reqs:
if not isinstance(num_sampled, int):
# idx_mapping may contain -1 sentinels (filtered rows) under PP; the
# kernel skips them rather than scattering with a host-side gather.
_scatter_num_accepted_kernel[(num_reqs,)](
idx_mapping,
num_sampled,
self.num_accepted_tokens_gpu,
)
else:
# Fill with single value.
_fill_num_accepted_kernel[(num_reqs,)](
idx_mapping,
self.num_accepted_tokens_gpu,
max(num_sampled, 1),
)
if self.recoverssm is not None:
self.recoverssm.commit_step(
num_sampled,
idx_mapping,
state_indices=(self._mamba_state_idx_gpu if self._align_mode else None),
num_accepted_tokens=self.num_accepted_tokens_gpu,
)
if not num_reqs:
return
# Align: save the running state to the block-aligned position when
# spec-decode acceptance leaves the sequence non-block-aligned (mirrors
# the V1 align postprocess). num_computed_tokens already holds the
# post-step advanced count.
if (
self._align_mode
and num_computed_tokens is not None
and self._mamba_ctx is not None
):
self._mamba_ctx.run_fused_postprocess_align(
num_reqs,
self.num_accepted_tokens_gpu,
self._mamba_state_idx_gpu,
num_computed_tokens,
idx_mapping,
)