Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
49 changes: 9 additions & 40 deletions tests/v1/simple_kv_offload/test_kv_events.py
Original file line number Diff line number Diff line change
Expand Up @@ -714,17 +714,8 @@ def test_disk_mode_block_removed_relabeled_storage() -> None:
assert ev.medium == MEDIUM_STORAGE


def test_mamba_dcp_metadata_guard() -> None:
"""Mamba+DCP: token_ids guard emits [] when capture and emission sizes differ.

The metadata capture at _prepare_eager_store_specs uses
g_block_size = spec.block_size * cp_world_size = 32 * 2 = 64 to slice
tokens. The event emission uses block_size = spec.block_size * 1 = 32
for Mamba (SSM state is not CP-sharded). The guard in
_process_store_completion detects len(token_ids) != block_size and emits
token_ids=[] instead of misleading data. When #49962 fixes the capture
path, the guard stops triggering and correct token_ids flow through.
"""
def test_mamba_align_skips_positional_event_metadata() -> None:
"""Mamba align emits events only from explicit boundary handoffs."""
vbs = BLOCK_SIZE * 2
kv_cfg = _make_mixed_kv_cache_config(
num_blocks=16, with_mamba=True, mamba_block_size=32
Expand All @@ -739,8 +730,6 @@ def test_mamba_dcp_metadata_guard() -> None:
sched = fix.scheduler
gpu_pool = fix.gpu_block_pool

# Mamba g_block_size = 32 * 2 = 64, so 2 Mamba blocks need 128 computed
# tokens to be "ready"; FA g_block_size = 16 * 2 = 32 needs only 64.
req = _make_cp_request(num_blocks=4, virtual_block_size=vbs)
fa_blocks = _allocate_cp_gpu_blocks(gpu_pool, req, 2, vbs, group_id=0)
mamba_blocks = _allocate_cp_gpu_blocks(gpu_pool, req, 2, vbs, group_id=1)
Expand All @@ -759,32 +748,12 @@ def test_mamba_dcp_metadata_guard() -> None:

events = list(sched.take_events())
stored = [e for e in events if isinstance(e, BlockStored)]
assert len(stored) == 4, (
f"expected 4 BlockStored (2 FA + 2 Mamba), got {len(stored)}"
)

by_group: dict[int, list[BlockStored]] = {}
for ev in stored:
by_group.setdefault(ev.group_idx, []).append(ev)

# FA group: capture size (64) == emission size (64) → token_ids present.
fa_events = by_group[0]
assert len(fa_events) == 2
for ev in fa_events:
assert len(stored) == 2
assert {ev.group_idx for ev in stored} == {0}
for idx, ev in enumerate(stored):
assert ev.block_size == vbs
assert len(ev.token_ids) == vbs, (
f"FA token_ids should be {vbs}, got {len(ev.token_ids)}"
)

# Mamba group: capture size (64) != emission size (32) → guard emits [].
mamba_events = by_group[1]
assert len(mamba_events) == 2
for ev in mamba_events:
assert ev.block_size == 32
assert ev.token_ids == [], (
f"Mamba token_ids should be [] (guard), got {ev.token_ids}"
)
assert ev.parent_block_hash is None, (
"Mamba parent_block_hash should be None (guard)"
assert ev.token_ids == req.prompt_token_ids[idx * vbs : (idx + 1) * vbs]
expected_hash = maybe_convert_block_hash(
get_block_hash(make_block_hash_with_group_id(req.block_hashes[idx], 0))
)
assert ev.extra_keys is None, "Mamba extra_keys should be None (guard)"
assert ev.block_hashes == [expected_hash]
Loading
Loading