Return intermediate Mamba states (#19716)

This commit is contained in:
roikoren755
2026-03-09 10:04:36 +02:00
committed by GitHub
parent 484f53c40e
commit c76251f70c
2 changed files with 89 additions and 3 deletions

View File

@@ -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

View File

@@ -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