Fix mamba radix cache ssm state indexing (#37836)
Co-authored-by: adityakamat24 <adityakamat007@gmail.com>
This commit is contained in:
@@ -25,6 +25,33 @@ def is_int_pow_2(n):
|
||||
return isinstance(n, int) and n > 0 and (n & (n - 1)) == 0
|
||||
|
||||
|
||||
def _track_states_at(
|
||||
B,
|
||||
x,
|
||||
dt,
|
||||
dA_cumsum,
|
||||
states,
|
||||
cu_seqlens,
|
||||
initial_states,
|
||||
track_seq_idx,
|
||||
track_end_locs,
|
||||
):
|
||||
starts = cu_seqlens[track_seq_idx]
|
||||
ends = track_end_locs.to(device=starts.device, dtype=starts.dtype)
|
||||
bounds = torch.stack([starts, ends], dim=1).flatten().contiguous()
|
||||
if initial_states is not None:
|
||||
initial_states = initial_states[track_seq_idx].repeat_interleave(2, dim=0)[:-1]
|
||||
return chunk_state_varlen(
|
||||
B.squeeze(0),
|
||||
x.squeeze(0),
|
||||
dt.squeeze(0),
|
||||
dA_cumsum.squeeze(0),
|
||||
bounds,
|
||||
states.squeeze(0),
|
||||
initial_states=initial_states,
|
||||
)[::2].unsqueeze(0)
|
||||
|
||||
|
||||
def _mamba_chunk_scan_combined_fwd(
|
||||
x,
|
||||
dt,
|
||||
@@ -44,6 +71,8 @@ def _mamba_chunk_scan_combined_fwd(
|
||||
dt_limit=(0.0, float("inf")),
|
||||
state_dtype=None,
|
||||
out=None,
|
||||
track_seq_idx=None,
|
||||
track_end_locs=None,
|
||||
):
|
||||
assert is_int_pow_2(chunk_size), "chunk_size must be integer power of 2"
|
||||
batch, seqlen, nheads, headdim = x.shape
|
||||
@@ -175,7 +204,24 @@ def _mamba_chunk_scan_combined_fwd(
|
||||
states.squeeze(0),
|
||||
initial_states=initial_states,
|
||||
)
|
||||
return out_x, dt, dA_cumsum, states, final_states, varlen_states
|
||||
track_states = None
|
||||
if (
|
||||
track_seq_idx is not None
|
||||
and track_end_locs is not None
|
||||
and track_end_locs.numel() > 0
|
||||
):
|
||||
track_states = _track_states_at(
|
||||
B,
|
||||
x,
|
||||
dt,
|
||||
dA_cumsum,
|
||||
states,
|
||||
cu_seqlens,
|
||||
initial_states,
|
||||
track_seq_idx,
|
||||
track_end_locs,
|
||||
)
|
||||
return out_x, dt, dA_cumsum, states, final_states, varlen_states, track_states
|
||||
|
||||
|
||||
def mamba_chunk_scan_combined(
|
||||
@@ -200,6 +246,9 @@ def mamba_chunk_scan_combined(
|
||||
return_varlen_states=False,
|
||||
return_intermediate_states=False,
|
||||
state_dtype=None,
|
||||
return_track_states=False,
|
||||
track_seq_idx=None,
|
||||
track_end_locs=None,
|
||||
):
|
||||
"""
|
||||
Argument:
|
||||
@@ -246,8 +295,17 @@ def mamba_chunk_scan_combined(
|
||||
dt_limit=dt_limit,
|
||||
out=out,
|
||||
state_dtype=state_dtype,
|
||||
track_seq_idx=track_seq_idx,
|
||||
track_end_locs=track_end_locs,
|
||||
)
|
||||
)
|
||||
if return_track_states:
|
||||
assert return_varlen_states, (
|
||||
"return_track_states requires return_varlen_states (cu_seqlens mode)"
|
||||
)
|
||||
# `states` rides along: a tracked position on the grid is read from it
|
||||
# directly, which stays bit-exact against a recompute-free run.
|
||||
return states, rest[0], rest[1]
|
||||
if return_intermediate_states:
|
||||
if return_varlen_states:
|
||||
varlen_states = rest[0]
|
||||
|
||||
@@ -122,6 +122,9 @@ class MambaAttnBackendBase(AttentionBackend):
|
||||
track_ssm_h_dst = None
|
||||
track_ssm_final_src = None
|
||||
track_ssm_final_dst = None
|
||||
track_ssm_seq_idx = None
|
||||
track_ssm_end_locs = None
|
||||
track_ssm_recompute_dst = None
|
||||
|
||||
mamba_cache_indices = self.req_to_token_pool.get_mamba_indices(
|
||||
forward_batch.req_pool_indices
|
||||
@@ -252,6 +255,9 @@ class MambaAttnBackendBase(AttentionBackend):
|
||||
track_ssm_h_dst,
|
||||
track_ssm_final_src,
|
||||
track_ssm_final_dst,
|
||||
track_ssm_seq_idx,
|
||||
track_ssm_end_locs,
|
||||
track_ssm_recompute_dst,
|
||||
) = self._init_track_ssm_indices(mamba_cache_indices, forward_batch)
|
||||
else:
|
||||
raise ValueError(f"Invalid forward mode: {forward_batch.forward_mode=}")
|
||||
@@ -270,6 +276,9 @@ class MambaAttnBackendBase(AttentionBackend):
|
||||
track_ssm_h_dst=track_ssm_h_dst,
|
||||
track_ssm_final_src=track_ssm_final_src,
|
||||
track_ssm_final_dst=track_ssm_final_dst,
|
||||
track_ssm_seq_idx=track_ssm_seq_idx,
|
||||
track_ssm_end_locs=track_ssm_end_locs,
|
||||
track_ssm_recompute_dst=track_ssm_recompute_dst,
|
||||
has_mamba_track_mask=has_mamba_track_mask,
|
||||
replayssm_write_pos=replayssm_write_pos,
|
||||
replayssm_force_flush=replayssm_force_flush,
|
||||
@@ -358,17 +367,8 @@ class MambaAttnBackendBase(AttentionBackend):
|
||||
mamba_track_seqlens = forward_batch.mamba_track_seqlens.cpu()
|
||||
prefix_lens = forward_batch.extend_prefix_lens.cpu()
|
||||
|
||||
if isinstance(self, Mamba2AttnBackend):
|
||||
num_h_states = extend_seq_lens // state_chunk_size
|
||||
else:
|
||||
num_h_states = (extend_seq_lens - 1) // state_chunk_size + 1
|
||||
|
||||
track_ssm_src_offset = torch.zeros_like(num_h_states)
|
||||
track_ssm_src_offset[1:] = torch.cumsum(num_h_states[:-1], dim=0)
|
||||
|
||||
lens_to_track = mamba_track_seqlens - prefix_lens
|
||||
lens_masked = lens_to_track[mamba_track_mask]
|
||||
offset_masked = track_ssm_src_offset[mamba_track_mask]
|
||||
dst_masked = mamba_track_indices[mamba_track_mask]
|
||||
|
||||
is_aligned = (lens_masked % state_chunk_size) == 0
|
||||
@@ -379,16 +379,50 @@ class MambaAttnBackendBase(AttentionBackend):
|
||||
|
||||
# Unaligned: intermediate state from h.
|
||||
not_aligned = ~is_aligned
|
||||
track_ssm_h_src = offset_masked[not_aligned] + (
|
||||
lens_masked[not_aligned] // state_chunk_size
|
||||
)
|
||||
track_ssm_h_dst = dst_masked[not_aligned]
|
||||
track_ssm_seq_idx = None
|
||||
track_ssm_end_locs = None
|
||||
track_ssm_recompute_dst = None
|
||||
|
||||
if isinstance(self, Mamba2AttnBackend):
|
||||
# `h` is one chunk grid over the whole flattened batch, so the wanted
|
||||
# state is in it only when query_start_loc[i] is chunk-aligned.
|
||||
query_start_loc = torch.zeros_like(extend_seq_lens)
|
||||
query_start_loc[1:] = torch.cumsum(extend_seq_lens[:-1], dim=0)
|
||||
aligned_len = (
|
||||
lens_masked[not_aligned] // state_chunk_size
|
||||
) * state_chunk_size
|
||||
seq_idx_unaligned = torch.nonzero(mamba_track_mask, as_tuple=True)[0][
|
||||
not_aligned
|
||||
]
|
||||
end_locs = query_start_loc[seq_idx_unaligned] + aligned_len
|
||||
on_grid = (end_locs % state_chunk_size) == 0
|
||||
|
||||
track_ssm_h_src = end_locs[on_grid] // state_chunk_size
|
||||
track_ssm_h_dst = track_ssm_h_dst[on_grid]
|
||||
|
||||
track_ssm_seq_idx = seq_idx_unaligned[~on_grid]
|
||||
track_ssm_end_locs = end_locs[~on_grid]
|
||||
track_ssm_recompute_dst = dst_masked[not_aligned][~on_grid]
|
||||
else:
|
||||
num_h_states = (extend_seq_lens - 1) // state_chunk_size + 1
|
||||
track_ssm_src_offset = torch.zeros_like(num_h_states)
|
||||
track_ssm_src_offset[1:] = torch.cumsum(num_h_states[:-1], dim=0)
|
||||
track_ssm_h_src = track_ssm_src_offset[mamba_track_mask][not_aligned] + (
|
||||
lens_masked[not_aligned] // state_chunk_size
|
||||
)
|
||||
|
||||
def to_device(t):
|
||||
return None if t is None else t.to(self.device, non_blocking=True)
|
||||
|
||||
return (
|
||||
track_ssm_h_src.to(self.device, non_blocking=True),
|
||||
track_ssm_h_dst.to(self.device, non_blocking=True),
|
||||
track_ssm_final_src.to(self.device, non_blocking=True),
|
||||
track_ssm_final_dst.to(self.device, non_blocking=True),
|
||||
to_device(track_ssm_h_src),
|
||||
to_device(track_ssm_h_dst),
|
||||
to_device(track_ssm_final_src),
|
||||
to_device(track_ssm_final_dst),
|
||||
to_device(track_ssm_seq_idx),
|
||||
to_device(track_ssm_end_locs),
|
||||
to_device(track_ssm_recompute_dst),
|
||||
)
|
||||
|
||||
def init_forward_metadata_capture_cpu_graph(
|
||||
@@ -847,6 +881,7 @@ class MambaAttnBackendBase(AttentionBackend):
|
||||
h: Optional[torch.Tensor],
|
||||
ssm_states: torch.Tensor,
|
||||
forward_metadata: ForwardMetadata,
|
||||
track_states: Optional[torch.Tensor] = None,
|
||||
):
|
||||
"""Copy extend SSM state at the last chunk boundary to track slots (source
|
||||
depends on chunk alignment; see `_init_track_ssm_indices`)."""
|
||||
@@ -859,6 +894,14 @@ class MambaAttnBackendBase(AttentionBackend):
|
||||
ssm_states[forward_metadata.track_ssm_h_dst] = h[
|
||||
forward_metadata.track_ssm_h_src
|
||||
].to(ssm_states.dtype, copy=False)
|
||||
if (
|
||||
forward_metadata.track_ssm_recompute_dst is not None
|
||||
and forward_metadata.track_ssm_recompute_dst.numel() > 0
|
||||
):
|
||||
assert track_states is not None
|
||||
ssm_states[forward_metadata.track_ssm_recompute_dst] = (
|
||||
track_states.squeeze(0).to(ssm_states.dtype, copy=False)
|
||||
)
|
||||
if forward_metadata.track_ssm_final_src.numel() > 0:
|
||||
ssm_states[forward_metadata.track_ssm_final_dst] = ssm_states[
|
||||
forward_metadata.track_ssm_final_src
|
||||
@@ -941,7 +984,7 @@ class Mamba2AttnBackend(MambaAttnBackendBase):
|
||||
use_triton_causal_conv or get_memory().enable_page_major_kv_layout
|
||||
)
|
||||
layer_cache = self.req_to_token_pool.mamba2_layer_cache(layer_id)
|
||||
mixer_out, intermediate_states = mixer.forward(
|
||||
mixer_out, intermediate_states, track_states = mixer.forward(
|
||||
hidden_states=hidden_states,
|
||||
output=output,
|
||||
layer_cache=layer_cache,
|
||||
@@ -957,6 +1000,7 @@ class Mamba2AttnBackend(MambaAttnBackendBase):
|
||||
intermediate_states,
|
||||
layer_cache.temporal,
|
||||
self.forward_metadata,
|
||||
track_states=track_states,
|
||||
)
|
||||
|
||||
if self.forward_metadata.num_decodes > 0:
|
||||
|
||||
@@ -461,6 +461,7 @@ class MambaMixer2(torch.nn.Module):
|
||||
conv_state = layer_cache.conv[0]
|
||||
ssm_state = layer_cache.temporal
|
||||
intermediate_states = None
|
||||
track_states = None
|
||||
|
||||
query_start_loc = metadata.query_start_loc
|
||||
|
||||
@@ -600,7 +601,7 @@ class MambaMixer2(torch.nn.Module):
|
||||
)
|
||||
|
||||
# NOTE: final output is an in-place update of out tensor
|
||||
intermediate_states, varlen_state = mamba_chunk_scan_combined(
|
||||
intermediate_states, varlen_state, track_states = mamba_chunk_scan_combined(
|
||||
hidden_states_p.view(
|
||||
1, num_prefill_tokens, local_num_heads, self.head_dim
|
||||
),
|
||||
@@ -619,7 +620,9 @@ class MambaMixer2(torch.nn.Module):
|
||||
initial_states=initial_states,
|
||||
return_varlen_states=True,
|
||||
return_final_states=False,
|
||||
return_intermediate_states=True,
|
||||
return_track_states=True,
|
||||
track_seq_idx=metadata.track_ssm_seq_idx,
|
||||
track_end_locs=metadata.track_ssm_end_locs,
|
||||
dt_softplus=True,
|
||||
dt_limit=(0.0, float("inf")),
|
||||
out=preallocated_ssm_out_p.view(
|
||||
@@ -763,7 +766,7 @@ class MambaMixer2(torch.nn.Module):
|
||||
if output is not None:
|
||||
output[:padded_num_tokens].copy_(mixer_out)
|
||||
|
||||
return mixer_out, intermediate_states
|
||||
return mixer_out, intermediate_states, track_states
|
||||
|
||||
@property
|
||||
def mamba_type(self) -> str:
|
||||
|
||||
@@ -58,6 +58,9 @@ class ForwardMetadata:
|
||||
state_checkpoint_cu_starts: Optional[torch.Tensor] = None
|
||||
num_state_checkpoints: int = 0
|
||||
state_checkpoint_every_n_tokens: int = 0
|
||||
track_ssm_seq_idx: Optional[torch.Tensor] = None
|
||||
track_ssm_end_locs: Optional[torch.Tensor] = None
|
||||
track_ssm_recompute_dst: Optional[torch.Tensor] = None
|
||||
|
||||
is_target_verify: bool = False
|
||||
draft_token_num: int = 1
|
||||
@@ -202,6 +205,9 @@ class Mamba2Metadata(ForwardMetadata):
|
||||
track_ssm_h_dst=forward_metadata.track_ssm_h_dst,
|
||||
track_ssm_final_src=forward_metadata.track_ssm_final_src,
|
||||
track_ssm_final_dst=forward_metadata.track_ssm_final_dst,
|
||||
track_ssm_seq_idx=forward_metadata.track_ssm_seq_idx,
|
||||
track_ssm_end_locs=forward_metadata.track_ssm_end_locs,
|
||||
track_ssm_recompute_dst=forward_metadata.track_ssm_recompute_dst,
|
||||
has_mamba_track_mask=forward_metadata.has_mamba_track_mask,
|
||||
num_decodes=len(seq_lens) if num_decodes is None else num_decodes,
|
||||
num_prefills=0,
|
||||
@@ -303,6 +309,9 @@ class Mamba2Metadata(ForwardMetadata):
|
||||
track_ssm_h_dst=forward_metadata.track_ssm_h_dst,
|
||||
track_ssm_final_src=forward_metadata.track_ssm_final_src,
|
||||
track_ssm_final_dst=forward_metadata.track_ssm_final_dst,
|
||||
track_ssm_seq_idx=forward_metadata.track_ssm_seq_idx,
|
||||
track_ssm_end_locs=forward_metadata.track_ssm_end_locs,
|
||||
track_ssm_recompute_dst=forward_metadata.track_ssm_recompute_dst,
|
||||
has_mamba_track_mask=forward_metadata.has_mamba_track_mask,
|
||||
mamba_track_mask_indices=mamba_track_mask_indices,
|
||||
conv_states_mask_indices=conv_states_mask_indices,
|
||||
|
||||
@@ -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