From c76251f70c68060cc452e15b211dc580600cb6b7 Mon Sep 17 00:00:00 2001 From: roikoren755 <26850796+roikoren755@users.noreply.github.com> Date: Mon, 9 Mar 2026 10:04:36 +0200 Subject: [PATCH] Return intermediate Mamba states (#19716) --- .../attention/mamba/ops/ssd_combined.py | 14 ++++ .../layers/mamba/test_mamba_ssm_ssd.py | 78 ++++++++++++++++++- 2 files changed, 89 insertions(+), 3 deletions(-) diff --git a/python/sglang/srt/layers/attention/mamba/ops/ssd_combined.py b/python/sglang/srt/layers/attention/mamba/ops/ssd_combined.py index 6e2e74752..c7f16e70e 100644 --- a/python/sglang/srt/layers/attention/mamba/ops/ssd_combined.py +++ b/python/sglang/srt/layers/attention/mamba/ops/ssd_combined.py @@ -198,6 +198,7 @@ def mamba_chunk_scan_combined( out=None, return_final_states=False, return_varlen_states=False, + return_intermediate_states=False, state_dtype=None, ): """ @@ -247,6 +248,19 @@ def mamba_chunk_scan_combined( state_dtype=state_dtype, ) ) + if return_intermediate_states: + if return_varlen_states: + varlen_states = rest[0] + if return_final_states: + return states, final_states, varlen_states + else: + return states, varlen_states + else: + if return_final_states: + return states, final_states + else: + return states + if not return_varlen_states: if not return_final_states: return diff --git a/test/registered/layers/mamba/test_mamba_ssm_ssd.py b/test/registered/layers/mamba/test_mamba_ssm_ssd.py index f6191d0bf..02b5f9f1e 100644 --- a/test/registered/layers/mamba/test_mamba_ssm_ssd.py +++ b/test/registered/layers/mamba/test_mamba_ssm_ssd.py @@ -1,7 +1,7 @@ from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci -register_cuda_ci(est_time=13, suite="stage-b-test-small-1-gpu") -register_amd_ci(est_time=30, suite="stage-b-test-small-1-gpu-amd") +register_cuda_ci(est_time=15, suite="stage-b-test-small-1-gpu") +register_amd_ci(est_time=34, suite="stage-b-test-small-1-gpu-amd") # Adapted from https://github.com/vllm-project/vllm/blob/633f943e30a4444d890d26b81850f7217736f840/tests/kernels/mamba/test_mamba_ssm_ssd.py @@ -38,7 +38,15 @@ def segsum(x): return x_segsum -def ssd_minimal_discrete(X, A, B, C, block_len, initial_states=None): +def ssd_minimal_discrete( + X, + A, + B, + C, + block_len, + initial_states=None, + return_intermediate_states=False, +): """ Arguments: X: (batch, length, n_heads, d_head) @@ -86,6 +94,8 @@ def ssd_minimal_discrete(X, A, B, C, block_len, initial_states=None): # Add output of intra-chunk and inter-chunk terms # (diagonal and off-diagonal blocks) Y = rearrange(Y_diag + Y_off, "b c l h p -> b (c l) h p") + if return_intermediate_states: + return Y, final_state, states return Y, final_state @@ -612,6 +622,68 @@ def test_mamba_chunk_scan_cont_batch_prefill_chunking(chunk_size, seqlens): ) # noqa: B023 +@pytest.mark.parametrize("itype", [torch.float32, torch.bfloat16]) +@pytest.mark.parametrize("n_heads", [4, 16]) +@pytest.mark.parametrize("d_head", [32, 64]) +@pytest.mark.parametrize("seq_len_chunk_size", [(128, 32), (256, 64)]) +def test_mamba_chunk_scan_intermediate_states( + d_head, + n_heads, + seq_len_chunk_size, + itype, +): + if not torch.cuda.is_available(): + pytest.skip("CUDA device not available") + + if itype == torch.bfloat16: + atol, rtol = 5e-2, 5e-2 + else: + atol, rtol = 8e-3, 5e-3 + + batch_size = 1 + seqlen, chunk_size = seq_len_chunk_size + + A, dt, X, B, C = generate_random_inputs(batch_size, seqlen, n_heads, d_head, itype) + + _, ref_final_state, ref_states = ssd_minimal_discrete( + X * dt.unsqueeze(-1), A * dt, B, C, chunk_size, return_intermediate_states=True + ) + + Y = torch.empty_like(X) + states, final_state = mamba_chunk_scan_combined( + X, + dt, + A, + B, + C, + chunk_size, + D=None, + return_intermediate_states=True, + return_final_states=True, + out=Y, + ) + + num_chunks = seqlen // chunk_size + assert states.shape == (batch_size, num_chunks, n_heads, d_head, d_head) + assert ref_states.shape == states.shape + + torch.testing.assert_close( + final_state[:, -1], + ref_final_state[:, -1].to(torch.float32), + atol=atol, + rtol=rtol, + ) + + for chunk_idx in range(num_chunks): + torch.testing.assert_close( + states[:, chunk_idx, -1], + ref_states[:, chunk_idx, -1].to(states.dtype), + atol=atol, + rtol=rtol, + msg=lambda x: f"chunk {chunk_idx} " + x, + ) + + if __name__ == "__main__": import sys