From 0b57847ebfd4e28a2c371892b9d4ac778e34b3fa Mon Sep 17 00:00:00 2001 From: Ke Bao Date: Fri, 4 Sep 2026 15:46:33 +0800 Subject: [PATCH] Fix mamba radix cache ssm state indexing (#37836) Co-authored-by: adityakamat24 --- .../ops/mamba/triton_ops/ssd_combined.py | 60 ++++++++- .../attention/hybrid_linear_attn_backend.py | 78 +++++++++--- .../srt/layers/attention/mamba/mamba.py | 9 +- .../layers/attention/mamba/mamba2_metadata.py | 9 ++ .../layers/mamba/test_mamba_ssm_ssd.py | 88 +++++++++++++ .../layers/test_mamba2_track_ssm_indices.py | 119 ++++++++++++++++++ 6 files changed, 342 insertions(+), 21 deletions(-) create mode 100644 test/registered/unit/layers/test_mamba2_track_ssm_indices.py diff --git a/python/sglang/kernels/ops/mamba/triton_ops/ssd_combined.py b/python/sglang/kernels/ops/mamba/triton_ops/ssd_combined.py index b79ff2365..c12e87166 100644 --- a/python/sglang/kernels/ops/mamba/triton_ops/ssd_combined.py +++ b/python/sglang/kernels/ops/mamba/triton_ops/ssd_combined.py @@ -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] diff --git a/python/sglang/srt/layers/attention/hybrid_linear_attn_backend.py b/python/sglang/srt/layers/attention/hybrid_linear_attn_backend.py index ce625166b..772ccd142 100644 --- a/python/sglang/srt/layers/attention/hybrid_linear_attn_backend.py +++ b/python/sglang/srt/layers/attention/hybrid_linear_attn_backend.py @@ -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: diff --git a/python/sglang/srt/layers/attention/mamba/mamba.py b/python/sglang/srt/layers/attention/mamba/mamba.py index 13c0f35fc..0778c3ab6 100644 --- a/python/sglang/srt/layers/attention/mamba/mamba.py +++ b/python/sglang/srt/layers/attention/mamba/mamba.py @@ -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: diff --git a/python/sglang/srt/layers/attention/mamba/mamba2_metadata.py b/python/sglang/srt/layers/attention/mamba/mamba2_metadata.py index 60a9da352..0aaa01ffd 100644 --- a/python/sglang/srt/layers/attention/mamba/mamba2_metadata.py +++ b/python/sglang/srt/layers/attention/mamba/mamba2_metadata.py @@ -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, diff --git a/test/registered/layers/mamba/test_mamba_ssm_ssd.py b/test/registered/layers/mamba/test_mamba_ssm_ssd.py index d19eaf391..6933ea72e 100644 --- a/test/registered/layers/mamba/test_mamba_ssm_ssd.py +++ b/test/registered/layers/mamba/test_mamba_ssm_ssd.py @@ -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__": import sys diff --git a/test/registered/unit/layers/test_mamba2_track_ssm_indices.py b/test/registered/unit/layers/test_mamba2_track_ssm_indices.py new file mode 100644 index 000000000..b081f5425 --- /dev/null +++ b/test/registered/unit/layers/test_mamba2_track_ssm_indices.py @@ -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()