Return intermediate Mamba states (#19716)
This commit is contained in:
@@ -198,6 +198,7 @@ def mamba_chunk_scan_combined(
|
|||||||
out=None,
|
out=None,
|
||||||
return_final_states=False,
|
return_final_states=False,
|
||||||
return_varlen_states=False,
|
return_varlen_states=False,
|
||||||
|
return_intermediate_states=False,
|
||||||
state_dtype=None,
|
state_dtype=None,
|
||||||
):
|
):
|
||||||
"""
|
"""
|
||||||
@@ -247,6 +248,19 @@ def mamba_chunk_scan_combined(
|
|||||||
state_dtype=state_dtype,
|
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_varlen_states:
|
||||||
if not return_final_states:
|
if not return_final_states:
|
||||||
return
|
return
|
||||||
|
|||||||
@@ -1,7 +1,7 @@
|
|||||||
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
|
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_cuda_ci(est_time=15, suite="stage-b-test-small-1-gpu")
|
||||||
register_amd_ci(est_time=30, suite="stage-b-test-small-1-gpu-amd")
|
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
|
# 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
|
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:
|
Arguments:
|
||||||
X: (batch, length, n_heads, d_head)
|
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
|
# Add output of intra-chunk and inter-chunk terms
|
||||||
# (diagonal and off-diagonal blocks)
|
# (diagonal and off-diagonal blocks)
|
||||||
Y = rearrange(Y_diag + Y_off, "b c l h p -> b (c l) h p")
|
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
|
return Y, final_state
|
||||||
|
|
||||||
|
|
||||||
@@ -612,6 +622,68 @@ def test_mamba_chunk_scan_cont_batch_prefill_chunking(chunk_size, seqlens):
|
|||||||
) # noqa: B023
|
) # 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__":
|
if __name__ == "__main__":
|
||||||
import sys
|
import sys
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user