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 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( def _mamba_chunk_scan_combined_fwd(
x, x,
dt, dt,
@@ -44,6 +71,8 @@ def _mamba_chunk_scan_combined_fwd(
dt_limit=(0.0, float("inf")), dt_limit=(0.0, float("inf")),
state_dtype=None, state_dtype=None,
out=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" assert is_int_pow_2(chunk_size), "chunk_size must be integer power of 2"
batch, seqlen, nheads, headdim = x.shape batch, seqlen, nheads, headdim = x.shape
@@ -175,7 +204,24 @@ def _mamba_chunk_scan_combined_fwd(
states.squeeze(0), states.squeeze(0),
initial_states=initial_states, 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( def mamba_chunk_scan_combined(
@@ -200,6 +246,9 @@ def mamba_chunk_scan_combined(
return_varlen_states=False, return_varlen_states=False,
return_intermediate_states=False, return_intermediate_states=False,
state_dtype=None, state_dtype=None,
return_track_states=False,
track_seq_idx=None,
track_end_locs=None,
): ):
""" """
Argument: Argument:
@@ -246,8 +295,17 @@ def mamba_chunk_scan_combined(
dt_limit=dt_limit, dt_limit=dt_limit,
out=out, out=out,
state_dtype=state_dtype, 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_intermediate_states:
if return_varlen_states: if return_varlen_states:
varlen_states = rest[0] varlen_states = rest[0]
@@ -122,6 +122,9 @@ class MambaAttnBackendBase(AttentionBackend):
track_ssm_h_dst = None track_ssm_h_dst = None
track_ssm_final_src = None track_ssm_final_src = None
track_ssm_final_dst = 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( mamba_cache_indices = self.req_to_token_pool.get_mamba_indices(
forward_batch.req_pool_indices forward_batch.req_pool_indices
@@ -252,6 +255,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_end_locs,
track_ssm_recompute_dst,
) = self._init_track_ssm_indices(mamba_cache_indices, forward_batch) ) = self._init_track_ssm_indices(mamba_cache_indices, forward_batch)
else: else:
raise ValueError(f"Invalid forward mode: {forward_batch.forward_mode=}") 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_h_dst=track_ssm_h_dst,
track_ssm_final_src=track_ssm_final_src, track_ssm_final_src=track_ssm_final_src,
track_ssm_final_dst=track_ssm_final_dst, 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, has_mamba_track_mask=has_mamba_track_mask,
replayssm_write_pos=replayssm_write_pos, replayssm_write_pos=replayssm_write_pos,
replayssm_force_flush=replayssm_force_flush, replayssm_force_flush=replayssm_force_flush,
@@ -358,17 +367,8 @@ class MambaAttnBackendBase(AttentionBackend):
mamba_track_seqlens = forward_batch.mamba_track_seqlens.cpu() mamba_track_seqlens = forward_batch.mamba_track_seqlens.cpu()
prefix_lens = forward_batch.extend_prefix_lens.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_to_track = mamba_track_seqlens - prefix_lens
lens_masked = lens_to_track[mamba_track_mask] 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] dst_masked = mamba_track_indices[mamba_track_mask]
is_aligned = (lens_masked % state_chunk_size) == 0 is_aligned = (lens_masked % state_chunk_size) == 0
@@ -379,16 +379,50 @@ class MambaAttnBackendBase(AttentionBackend):
# Unaligned: intermediate state from h. # Unaligned: intermediate state from h.
not_aligned = ~is_aligned not_aligned = ~is_aligned
track_ssm_h_src = offset_masked[not_aligned] + ( 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 lens_masked[not_aligned] // state_chunk_size
) )
track_ssm_h_dst = dst_masked[not_aligned]
def to_device(t):
return None if t is None else t.to(self.device, non_blocking=True)
return ( return (
track_ssm_h_src.to(self.device, non_blocking=True), to_device(track_ssm_h_src),
track_ssm_h_dst.to(self.device, non_blocking=True), to_device(track_ssm_h_dst),
track_ssm_final_src.to(self.device, non_blocking=True), to_device(track_ssm_final_src),
track_ssm_final_dst.to(self.device, non_blocking=True), 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( def init_forward_metadata_capture_cpu_graph(
@@ -847,6 +881,7 @@ class MambaAttnBackendBase(AttentionBackend):
h: Optional[torch.Tensor], h: Optional[torch.Tensor],
ssm_states: torch.Tensor, ssm_states: torch.Tensor,
forward_metadata: ForwardMetadata, forward_metadata: ForwardMetadata,
track_states: Optional[torch.Tensor] = None,
): ):
"""Copy extend SSM state at the last chunk boundary to track slots (source """Copy extend SSM state at the last chunk boundary to track slots (source
depends on chunk alignment; see `_init_track_ssm_indices`).""" 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[ ssm_states[forward_metadata.track_ssm_h_dst] = h[
forward_metadata.track_ssm_h_src forward_metadata.track_ssm_h_src
].to(ssm_states.dtype, copy=False) ].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: if forward_metadata.track_ssm_final_src.numel() > 0:
ssm_states[forward_metadata.track_ssm_final_dst] = ssm_states[ ssm_states[forward_metadata.track_ssm_final_dst] = ssm_states[
forward_metadata.track_ssm_final_src 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 use_triton_causal_conv or get_memory().enable_page_major_kv_layout
) )
layer_cache = self.req_to_token_pool.mamba2_layer_cache(layer_id) 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, hidden_states=hidden_states,
output=output, output=output,
layer_cache=layer_cache, layer_cache=layer_cache,
@@ -957,6 +1000,7 @@ class Mamba2AttnBackend(MambaAttnBackendBase):
intermediate_states, intermediate_states,
layer_cache.temporal, layer_cache.temporal,
self.forward_metadata, self.forward_metadata,
track_states=track_states,
) )
if self.forward_metadata.num_decodes > 0: if self.forward_metadata.num_decodes > 0:
@@ -461,6 +461,7 @@ class MambaMixer2(torch.nn.Module):
conv_state = layer_cache.conv[0] conv_state = layer_cache.conv[0]
ssm_state = layer_cache.temporal ssm_state = layer_cache.temporal
intermediate_states = None intermediate_states = None
track_states = None
query_start_loc = metadata.query_start_loc 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 # 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( hidden_states_p.view(
1, num_prefill_tokens, local_num_heads, self.head_dim 1, num_prefill_tokens, local_num_heads, self.head_dim
), ),
@@ -619,7 +620,9 @@ class MambaMixer2(torch.nn.Module):
initial_states=initial_states, initial_states=initial_states,
return_varlen_states=True, return_varlen_states=True,
return_final_states=False, 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_softplus=True,
dt_limit=(0.0, float("inf")), dt_limit=(0.0, float("inf")),
out=preallocated_ssm_out_p.view( out=preallocated_ssm_out_p.view(
@@ -763,7 +766,7 @@ class MambaMixer2(torch.nn.Module):
if output is not None: if output is not None:
output[:padded_num_tokens].copy_(mixer_out) output[:padded_num_tokens].copy_(mixer_out)
return mixer_out, intermediate_states return mixer_out, intermediate_states, track_states
@property @property
def mamba_type(self) -> str: def mamba_type(self) -> str:
@@ -58,6 +58,9 @@ class ForwardMetadata:
state_checkpoint_cu_starts: Optional[torch.Tensor] = None state_checkpoint_cu_starts: Optional[torch.Tensor] = None
num_state_checkpoints: int = 0 num_state_checkpoints: int = 0
state_checkpoint_every_n_tokens: 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 is_target_verify: bool = False
draft_token_num: int = 1 draft_token_num: int = 1
@@ -202,6 +205,9 @@ class Mamba2Metadata(ForwardMetadata):
track_ssm_h_dst=forward_metadata.track_ssm_h_dst, track_ssm_h_dst=forward_metadata.track_ssm_h_dst,
track_ssm_final_src=forward_metadata.track_ssm_final_src, track_ssm_final_src=forward_metadata.track_ssm_final_src,
track_ssm_final_dst=forward_metadata.track_ssm_final_dst, 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, has_mamba_track_mask=forward_metadata.has_mamba_track_mask,
num_decodes=len(seq_lens) if num_decodes is None else num_decodes, num_decodes=len(seq_lens) if num_decodes is None else num_decodes,
num_prefills=0, num_prefills=0,
@@ -303,6 +309,9 @@ class Mamba2Metadata(ForwardMetadata):
track_ssm_h_dst=forward_metadata.track_ssm_h_dst, track_ssm_h_dst=forward_metadata.track_ssm_h_dst,
track_ssm_final_src=forward_metadata.track_ssm_final_src, track_ssm_final_src=forward_metadata.track_ssm_final_src,
track_ssm_final_dst=forward_metadata.track_ssm_final_dst, 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, has_mamba_track_mask=forward_metadata.has_mamba_track_mask,
mamba_track_mask_indices=mamba_track_mask_indices, mamba_track_mask_indices=mamba_track_mask_indices,
conv_states_mask_indices=conv_states_mask_indices, conv_states_mask_indices=conv_states_mask_indices,
@@ -692,6 +692,94 @@ def test_mamba_chunk_scan_intermediate_states(
) )
@pytest.mark.parametrize("chunk_size", [64, 128])
@pytest.mark.parametrize(
"seqlens",
[
(500, 700, 900),
(256, 300, 400),
(130, 200, 300),
],
)
def test_mamba_chunk_scan_track_states_at_request_boundary(chunk_size, seqlens):
device = get_device()
if device not in ["cuda", "xpu"]:
pytest.skip("Test only supports CUDA and XPU devices")
n_heads, d_head, itype = 8, 64, torch.float32
A, dt, X, B, C = generate_random_inputs(
1, sum(seqlens), n_heads, d_head, itype, device
)
starts = [sum(seqlens[:i]) for i in range(len(seqlens) + 1)]
def run(x, d, b, c, lens, **kwargs):
cu_seqlens = torch.tensor(
[sum(lens[:i]) for i in range(len(lens) + 1)],
dtype=torch.int32,
device=device,
)
seq_idx = torch.repeat_interleave(
torch.arange(len(lens), dtype=torch.int32, device=device),
torch.tensor(lens, dtype=torch.int32, device=device),
).unsqueeze(0)
return mamba_chunk_scan_combined(
x,
d,
A,
b,
c,
chunk_size,
D=None,
cu_seqlens=cu_seqlens,
seq_idx=seq_idx,
initial_states=None,
out=torch.empty_like(x),
return_varlen_states=True,
**kwargs,
)
tracked = [
(i, (ln // chunk_size) * chunk_size)
for i, ln in enumerate(seqlens)
if ln >= chunk_size and ln % chunk_size != 0
]
assert tracked, "case must exercise the unaligned path"
_, _, track_states = run(
X,
dt,
B,
C,
list(seqlens),
return_track_states=True,
track_seq_idx=torch.tensor(
[i for i, _ in tracked], dtype=torch.int64, device=device
),
track_end_locs=torch.tensor(
[starts[i] + n for i, n in tracked], dtype=torch.int32, device=device
),
)
assert track_states.shape == (1, len(tracked), n_heads, d_head, d_head)
track_states = track_states.squeeze(0)
for j, (i, n) in enumerate(tracked):
sl = slice(starts[i], starts[i] + n)
expected = run(
X[:, sl].contiguous(),
dt[:, sl].contiguous(),
B[:, sl].contiguous(),
C[:, sl].contiguous(),
[n],
)
torch.testing.assert_close(
track_states[j],
expected[0],
atol=1e-2,
rtol=5e-3,
msg=lambda m, i=i, n=n: f"request {i} snapshot at {n} tokens: " + m,
)
if __name__ == "__main__": if __name__ == "__main__":
import sys import sys
@@ -0,0 +1,119 @@
"""A tracked SSM position that sits on the global chunk grid must be read from
`h`, not rebuilt: rebuilding it is numerically fine but not bit-exact, which
makes the cached state drift with the prefill batch composition.
"""
import unittest
from types import SimpleNamespace
import torch
from sglang.srt.layers.attention.hybrid_linear_attn_backend import Mamba2AttnBackend
from sglang.test.ci.ci_register import register_cpu_ci
register_cpu_ci(est_time=5, suite="base-a-test-cpu")
CHUNK = 128
def _backend():
backend = object.__new__(Mamba2AttnBackend)
backend.device = "cpu"
backend._mamba_chunk_size = CHUNK
return backend
def _forward_batch(extend_lens, prefix_lens, track_seqlens, track_mask):
return SimpleNamespace(
extend_seq_lens=torch.tensor(extend_lens),
extend_prefix_lens=torch.tensor(prefix_lens),
mamba_track_seqlens=torch.tensor(track_seqlens),
mamba_track_mask=torch.tensor(track_mask),
mamba_track_indices=torch.arange(100, 100 + len(extend_lens)),
)
def _split(extend_lens, prefix_lens, track_seqlens, track_mask):
backend = _backend()
cache_indices = torch.arange(len(extend_lens))
(
h_src,
h_dst,
_final_src,
_final_dst,
seq_idx,
end_locs,
recompute_dst,
) = backend._init_track_ssm_indices(
cache_indices,
_forward_batch(extend_lens, prefix_lens, track_seqlens, track_mask),
)
return h_src, h_dst, seq_idx, end_locs, recompute_dst
class TestMamba2TrackSsmIndices(unittest.TestCase):
def test_on_grid_request_reads_h(self):
# req0 is chunk-aligned, so req1 starts at flat 256 and its tracked
# position 256 + 128 is on the grid.
h_src, h_dst, seq_idx, _end_locs, recompute_dst = _split(
extend_lens=[256, 320],
prefix_lens=[0, 0],
track_seqlens=[0, 200],
track_mask=[False, True],
)
self.assertEqual(h_src.tolist(), [3])
self.assertEqual(h_dst.tolist(), [101])
self.assertEqual(seq_idx.numel(), 0)
self.assertEqual(recompute_dst.numel(), 0)
def test_off_grid_request_is_rebuilt(self):
# req0 is not chunk-aligned, so req1 starts at flat 192 and its tracked
# position 192 + 128 = 320 is not a multiple of 128.
h_src, _h_dst, seq_idx, end_locs, recompute_dst = _split(
extend_lens=[192, 320],
prefix_lens=[0, 0],
track_seqlens=[0, 200],
track_mask=[False, True],
)
self.assertEqual(h_src.numel(), 0)
self.assertEqual(seq_idx.tolist(), [1])
self.assertEqual(end_locs.tolist(), [320])
self.assertEqual(recompute_dst.tolist(), [101])
def test_mixed_batch_splits_without_overlap(self):
# req0 on grid (flat 0), req1 off grid (flat 192), req2 on grid (flat 512).
h_src, h_dst, seq_idx, end_locs, recompute_dst = _split(
extend_lens=[192, 320, 448],
prefix_lens=[0, 0, 0],
track_seqlens=[150, 200, 300],
track_mask=[True, True, True],
)
self.assertEqual(h_src.tolist(), [1, 6])
self.assertEqual(h_dst.tolist(), [100, 102])
self.assertEqual(seq_idx.tolist(), [1])
self.assertEqual(end_locs.tolist(), [320])
self.assertEqual(recompute_dst.tolist(), [101])
self.assertEqual(
sorted(h_dst.tolist() + recompute_dst.tolist()), [100, 101, 102]
)
def test_end_locs_are_chunk_boundaries_of_the_flat_batch(self):
# h is indexed by flat chunk, so a read index must satisfy
# h_src * CHUNK == the flat token the state belongs to.
starts = [0, 192, 512]
h_src, h_dst, seq_idx, end_locs, _recompute_dst = _split(
extend_lens=[192, 320, 448],
prefix_lens=[0, 0, 0],
track_seqlens=[150, 200, 300],
track_mask=[True, True, True],
)
want = {100: starts[0] + 128, 102: starts[2] + 256}
for dst, src in zip(h_dst.tolist(), h_src.tolist()):
self.assertEqual(src * CHUNK, want[dst])
for i, end in zip(seq_idx.tolist(), end_locs.tolist()):
self.assertNotEqual(end % CHUNK, 0)
self.assertGreaterEqual(end, starts[i])
if __name__ == "__main__":
unittest.main()