Return intermediate Mamba states (#19716)
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
Reference in New Issue
Block a user