Fix mamba radix cache ssm state indexing (#37836)
Co-authored-by: adityakamat24 <adityakamat007@gmail.com>
This commit is contained in:
@@ -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
|
||||
|
||||
|
||||
@@ -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()
|
||||
Reference in New Issue
Block a user