Fix mamba radix cache ssm state indexing (#37836)

Co-authored-by: adityakamat24 <adityakamat007@gmail.com>
This commit is contained in:
Ke Bao
2026-09-04 15:46:33 +08:00
committed by GitHub
co-authored by adityakamat24
parent 01e66a62db
commit 0b57847ebf
6 changed files with 342 additions and 21 deletions
@@ -25,6 +25,33 @@ def is_int_pow_2(n):
return isinstance(n, int) and n > 0 and (n & (n - 1)) == 0
def _track_states_at(
B,
x,
dt,
dA_cumsum,
states,
cu_seqlens,
initial_states,
track_seq_idx,
track_end_locs,
):
starts = cu_seqlens[track_seq_idx]
ends = track_end_locs.to(device=starts.device, dtype=starts.dtype)
bounds = torch.stack([starts, ends], dim=1).flatten().contiguous()
if initial_states is not None:
initial_states = initial_states[track_seq_idx].repeat_interleave(2, dim=0)[:-1]
return chunk_state_varlen(
B.squeeze(0),
x.squeeze(0),
dt.squeeze(0),
dA_cumsum.squeeze(0),
bounds,
states.squeeze(0),
initial_states=initial_states,
)[::2].unsqueeze(0)
def _mamba_chunk_scan_combined_fwd(
x,
dt,
@@ -44,6 +71,8 @@ def _mamba_chunk_scan_combined_fwd(
dt_limit=(0.0, float("inf")),
state_dtype=None,
out=None,
track_seq_idx=None,
track_end_locs=None,
):
assert is_int_pow_2(chunk_size), "chunk_size must be integer power of 2"
batch, seqlen, nheads, headdim = x.shape
@@ -175,7 +204,24 @@ def _mamba_chunk_scan_combined_fwd(
states.squeeze(0),
initial_states=initial_states,
)
return out_x, dt, dA_cumsum, states, final_states, varlen_states
track_states = None
if (
track_seq_idx is not None
and track_end_locs is not None
and track_end_locs.numel() > 0
):
track_states = _track_states_at(
B,
x,
dt,
dA_cumsum,
states,
cu_seqlens,
initial_states,
track_seq_idx,
track_end_locs,
)
return out_x, dt, dA_cumsum, states, final_states, varlen_states, track_states
def mamba_chunk_scan_combined(
@@ -200,6 +246,9 @@ def mamba_chunk_scan_combined(
return_varlen_states=False,
return_intermediate_states=False,
state_dtype=None,
return_track_states=False,
track_seq_idx=None,
track_end_locs=None,
):
"""
Argument:
@@ -246,8 +295,17 @@ def mamba_chunk_scan_combined(
dt_limit=dt_limit,
out=out,
state_dtype=state_dtype,
track_seq_idx=track_seq_idx,
track_end_locs=track_end_locs,
)
)
if return_track_states:
assert return_varlen_states, (
"return_track_states requires return_varlen_states (cu_seqlens mode)"
)
# `states` rides along: a tracked position on the grid is read from it
# directly, which stays bit-exact against a recompute-free run.
return states, rest[0], rest[1]
if return_intermediate_states:
if return_varlen_states:
varlen_states = rest[0]
@@ -122,6 +122,9 @@ class MambaAttnBackendBase(AttentionBackend):
track_ssm_h_dst = None
track_ssm_final_src = None
track_ssm_final_dst = None
track_ssm_seq_idx = None
track_ssm_end_locs = None
track_ssm_recompute_dst = None
mamba_cache_indices = self.req_to_token_pool.get_mamba_indices(
forward_batch.req_pool_indices
@@ -252,6 +255,9 @@ class MambaAttnBackendBase(AttentionBackend):
track_ssm_h_dst,
track_ssm_final_src,
track_ssm_final_dst,
track_ssm_seq_idx,
track_ssm_end_locs,
track_ssm_recompute_dst,
) = self._init_track_ssm_indices(mamba_cache_indices, forward_batch)
else:
raise ValueError(f"Invalid forward mode: {forward_batch.forward_mode=}")
@@ -270,6 +276,9 @@ class MambaAttnBackendBase(AttentionBackend):
track_ssm_h_dst=track_ssm_h_dst,
track_ssm_final_src=track_ssm_final_src,
track_ssm_final_dst=track_ssm_final_dst,
track_ssm_seq_idx=track_ssm_seq_idx,
track_ssm_end_locs=track_ssm_end_locs,
track_ssm_recompute_dst=track_ssm_recompute_dst,
has_mamba_track_mask=has_mamba_track_mask,
replayssm_write_pos=replayssm_write_pos,
replayssm_force_flush=replayssm_force_flush,
@@ -358,17 +367,8 @@ class MambaAttnBackendBase(AttentionBackend):
mamba_track_seqlens = forward_batch.mamba_track_seqlens.cpu()
prefix_lens = forward_batch.extend_prefix_lens.cpu()
if isinstance(self, Mamba2AttnBackend):
num_h_states = extend_seq_lens // state_chunk_size
else:
num_h_states = (extend_seq_lens - 1) // state_chunk_size + 1
track_ssm_src_offset = torch.zeros_like(num_h_states)
track_ssm_src_offset[1:] = torch.cumsum(num_h_states[:-1], dim=0)
lens_to_track = mamba_track_seqlens - prefix_lens
lens_masked = lens_to_track[mamba_track_mask]
offset_masked = track_ssm_src_offset[mamba_track_mask]
dst_masked = mamba_track_indices[mamba_track_mask]
is_aligned = (lens_masked % state_chunk_size) == 0
@@ -379,16 +379,50 @@ class MambaAttnBackendBase(AttentionBackend):
# Unaligned: intermediate state from h.
not_aligned = ~is_aligned
track_ssm_h_src = offset_masked[not_aligned] + (
lens_masked[not_aligned] // state_chunk_size
)
track_ssm_h_dst = dst_masked[not_aligned]
track_ssm_seq_idx = None
track_ssm_end_locs = None
track_ssm_recompute_dst = None
if isinstance(self, Mamba2AttnBackend):
# `h` is one chunk grid over the whole flattened batch, so the wanted
# state is in it only when query_start_loc[i] is chunk-aligned.
query_start_loc = torch.zeros_like(extend_seq_lens)
query_start_loc[1:] = torch.cumsum(extend_seq_lens[:-1], dim=0)
aligned_len = (
lens_masked[not_aligned] // state_chunk_size
) * state_chunk_size
seq_idx_unaligned = torch.nonzero(mamba_track_mask, as_tuple=True)[0][
not_aligned
]
end_locs = query_start_loc[seq_idx_unaligned] + aligned_len
on_grid = (end_locs % state_chunk_size) == 0
track_ssm_h_src = end_locs[on_grid] // state_chunk_size
track_ssm_h_dst = track_ssm_h_dst[on_grid]
track_ssm_seq_idx = seq_idx_unaligned[~on_grid]
track_ssm_end_locs = end_locs[~on_grid]
track_ssm_recompute_dst = dst_masked[not_aligned][~on_grid]
else:
num_h_states = (extend_seq_lens - 1) // state_chunk_size + 1
track_ssm_src_offset = torch.zeros_like(num_h_states)
track_ssm_src_offset[1:] = torch.cumsum(num_h_states[:-1], dim=0)
track_ssm_h_src = track_ssm_src_offset[mamba_track_mask][not_aligned] + (
lens_masked[not_aligned] // state_chunk_size
)
def to_device(t):
return None if t is None else t.to(self.device, non_blocking=True)
return (
track_ssm_h_src.to(self.device, non_blocking=True),
track_ssm_h_dst.to(self.device, non_blocking=True),
track_ssm_final_src.to(self.device, non_blocking=True),
track_ssm_final_dst.to(self.device, non_blocking=True),
to_device(track_ssm_h_src),
to_device(track_ssm_h_dst),
to_device(track_ssm_final_src),
to_device(track_ssm_final_dst),
to_device(track_ssm_seq_idx),
to_device(track_ssm_end_locs),
to_device(track_ssm_recompute_dst),
)
def init_forward_metadata_capture_cpu_graph(
@@ -847,6 +881,7 @@ class MambaAttnBackendBase(AttentionBackend):
h: Optional[torch.Tensor],
ssm_states: torch.Tensor,
forward_metadata: ForwardMetadata,
track_states: Optional[torch.Tensor] = None,
):
"""Copy extend SSM state at the last chunk boundary to track slots (source
depends on chunk alignment; see `_init_track_ssm_indices`)."""
@@ -859,6 +894,14 @@ class MambaAttnBackendBase(AttentionBackend):
ssm_states[forward_metadata.track_ssm_h_dst] = h[
forward_metadata.track_ssm_h_src
].to(ssm_states.dtype, copy=False)
if (
forward_metadata.track_ssm_recompute_dst is not None
and forward_metadata.track_ssm_recompute_dst.numel() > 0
):
assert track_states is not None
ssm_states[forward_metadata.track_ssm_recompute_dst] = (
track_states.squeeze(0).to(ssm_states.dtype, copy=False)
)
if forward_metadata.track_ssm_final_src.numel() > 0:
ssm_states[forward_metadata.track_ssm_final_dst] = ssm_states[
forward_metadata.track_ssm_final_src
@@ -941,7 +984,7 @@ class Mamba2AttnBackend(MambaAttnBackendBase):
use_triton_causal_conv or get_memory().enable_page_major_kv_layout
)
layer_cache = self.req_to_token_pool.mamba2_layer_cache(layer_id)
mixer_out, intermediate_states = mixer.forward(
mixer_out, intermediate_states, track_states = mixer.forward(
hidden_states=hidden_states,
output=output,
layer_cache=layer_cache,
@@ -957,6 +1000,7 @@ class Mamba2AttnBackend(MambaAttnBackendBase):
intermediate_states,
layer_cache.temporal,
self.forward_metadata,
track_states=track_states,
)
if self.forward_metadata.num_decodes > 0:
@@ -461,6 +461,7 @@ class MambaMixer2(torch.nn.Module):
conv_state = layer_cache.conv[0]
ssm_state = layer_cache.temporal
intermediate_states = None
track_states = None
query_start_loc = metadata.query_start_loc
@@ -600,7 +601,7 @@ class MambaMixer2(torch.nn.Module):
)
# NOTE: final output is an in-place update of out tensor
intermediate_states, varlen_state = mamba_chunk_scan_combined(
intermediate_states, varlen_state, track_states = mamba_chunk_scan_combined(
hidden_states_p.view(
1, num_prefill_tokens, local_num_heads, self.head_dim
),
@@ -619,7 +620,9 @@ class MambaMixer2(torch.nn.Module):
initial_states=initial_states,
return_varlen_states=True,
return_final_states=False,
return_intermediate_states=True,
return_track_states=True,
track_seq_idx=metadata.track_ssm_seq_idx,
track_end_locs=metadata.track_ssm_end_locs,
dt_softplus=True,
dt_limit=(0.0, float("inf")),
out=preallocated_ssm_out_p.view(
@@ -763,7 +766,7 @@ class MambaMixer2(torch.nn.Module):
if output is not None:
output[:padded_num_tokens].copy_(mixer_out)
return mixer_out, intermediate_states
return mixer_out, intermediate_states, track_states
@property
def mamba_type(self) -> str:
@@ -58,6 +58,9 @@ class ForwardMetadata:
state_checkpoint_cu_starts: Optional[torch.Tensor] = None
num_state_checkpoints: int = 0
state_checkpoint_every_n_tokens: int = 0
track_ssm_seq_idx: Optional[torch.Tensor] = None
track_ssm_end_locs: Optional[torch.Tensor] = None
track_ssm_recompute_dst: Optional[torch.Tensor] = None
is_target_verify: bool = False
draft_token_num: int = 1
@@ -202,6 +205,9 @@ class Mamba2Metadata(ForwardMetadata):
track_ssm_h_dst=forward_metadata.track_ssm_h_dst,
track_ssm_final_src=forward_metadata.track_ssm_final_src,
track_ssm_final_dst=forward_metadata.track_ssm_final_dst,
track_ssm_seq_idx=forward_metadata.track_ssm_seq_idx,
track_ssm_end_locs=forward_metadata.track_ssm_end_locs,
track_ssm_recompute_dst=forward_metadata.track_ssm_recompute_dst,
has_mamba_track_mask=forward_metadata.has_mamba_track_mask,
num_decodes=len(seq_lens) if num_decodes is None else num_decodes,
num_prefills=0,
@@ -303,6 +309,9 @@ class Mamba2Metadata(ForwardMetadata):
track_ssm_h_dst=forward_metadata.track_ssm_h_dst,
track_ssm_final_src=forward_metadata.track_ssm_final_src,
track_ssm_final_dst=forward_metadata.track_ssm_final_dst,
track_ssm_seq_idx=forward_metadata.track_ssm_seq_idx,
track_ssm_end_locs=forward_metadata.track_ssm_end_locs,
track_ssm_recompute_dst=forward_metadata.track_ssm_recompute_dst,
has_mamba_track_mask=forward_metadata.has_mamba_track_mask,
mamba_track_mask_indices=mamba_track_mask_indices,
conv_states_mask_indices=conv_states_mask_indices,