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
9 changes: 5 additions & 4 deletions megatron/core/inference/contexts/dynamic_context.py
Original file line number Diff line number Diff line change
Expand Up @@ -3223,9 +3223,10 @@ def _compute_prefix_match(
else:
prefix_skip_tokens = 0

# Hybrid models with Mamba caching: skip based on Mamba match count.
# Only applies to the first chunk (finished == 0); continuation chunks
# already had Mamba state restored during the first chunk.
# Hybrid models can skip tokens only when the first chunk restores a
# matching Mamba state. Continuation chunks carry live state from the
# previous chunk and must compute every subsequent recurrent transition,
# even when their attention KV blocks are already cached.
if self.is_hybrid_model and self.mamba_slot_allocator is not None and finished == 0:
num_mamba_matched = self._find_mamba_match_count(
req=req,
Expand All @@ -3249,7 +3250,7 @@ def _compute_prefix_match(
prefix_skip_tokens = raw_skip
else:
prefix_skip_tokens = 0
elif self.is_hybrid_model and finished == 0:
elif self.is_hybrid_model:
if record_mamba_match:
req._mamba_num_matched_blocks = 0
prefix_skip_tokens = 0
Expand Down
73 changes: 73 additions & 0 deletions tests/unit_tests/inference/contexts/test_dynamic_prefix_caching.py
Original file line number Diff line number Diff line change
Expand Up @@ -1747,6 +1747,79 @@ def test_unaligned_continuation_redirects_inherited_partial_block(self):
assert ctx.kv_block_allocator.block_ref_counts[block_id].item() == 2


@pytest.mark.internal
@pytest.mark.parametrize("record_mamba_match", [False, True])
@pytest.mark.parametrize(
"mamba_boundaries",
[None, [], [2], [2, 6, 9]],
ids=["memory_only", "snapshots_evicted", "earlier_snapshot", "later_snapshots"],
)
def test_hybrid_cached_continuation_preserves_recurrent_position(
mamba_boundaries, record_mamba_match
):
"""KV reuse must not skip transitions in a continuation's live Mamba state."""
ctx = object.__new__(DynamicInferenceContext)
ctx.block_size_tokens = 256
ctx.enable_prefix_caching = True
ctx.enable_mtp_kv_cache = False
ctx.is_hybrid_model = True
hashes = list(range(1, 11))
ctx.kv_block_allocator = SimpleNamespace(kv_hash_to_block_id={h: h for h in hashes})
ctx.mamba_slot_allocator = (
None
if mamba_boundaries is None
else SimpleNamespace(hash_to_block_id={h: h for h in mamba_boundaries})
)
req = SimpleNamespace(finished_chunk_token_count=0, precomputed_block_hashes=hashes)
bs = ctx.block_size_tokens
recurrent_position = 0

for chunk_length in (4 * bs, 4 * bs, 2 * bs + 2):
finished = req.finished_chunk_token_count
match = ctx._compute_prefix_match(req, chunk_length, record_mamba_match=record_mamba_match)
allocated = match.already_allocated_blocks
required = match.overall_required_blocks
assert match.matched_block_ids == hashes[allocated:required]
assert match.num_blocks_from_pool == required - allocated - len(match.matched_block_ids)
if finished == 0:
# A valid first-chunk snapshot still permits its matching prefix skip.
assert match.prefix_skip_tokens == (2 * bs if mamba_boundaries else 0)
recurrent_position = match.prefix_skip_tokens
else:
# No Mamba restore occurs on continuation admission. Sharing KV must
# not advance the logical position past uncomputed recurrent state.
assert match.prefix_skip_tokens == 0
assert match.effective_prefill_chunk_length == chunk_length
recurrent_position += match.effective_prefill_chunk_length
req.finished_chunk_token_count += chunk_length
assert recurrent_position == req.finished_chunk_token_count


@pytest.mark.internal
def test_attention_only_cached_continuation_still_skips():
"""Attention-only continuations can still skip matching KV blocks."""
ctx = object.__new__(DynamicInferenceContext)
ctx.block_size_tokens = 256
ctx.enable_prefix_caching = True
ctx.enable_mtp_kv_cache = False
ctx.is_hybrid_model = False
ctx.mamba_slot_allocator = None
hashes = list(range(1, 11))
ctx.kv_block_allocator = SimpleNamespace(kv_hash_to_block_id={h: h for h in hashes})
req = SimpleNamespace(
finished_chunk_token_count=4 * ctx.block_size_tokens, precomputed_block_hashes=hashes
)

match = ctx._compute_prefix_match(req, 6 * ctx.block_size_tokens + 2)

assert match.matched_block_ids == hashes[4:]
assert match.num_blocks_from_pool == 1
assert match.already_allocated_blocks == 4
assert match.overall_required_blocks == 11
assert match.prefix_skip_tokens == 6 * ctx.block_size_tokens
assert match.effective_prefill_chunk_length == 2


def _make_cpu_mamba_slot_allocator(
monkeypatch, *, total_blocks: int, max_slots: int
) -> MambaSlotAllocator:
Expand Down
Loading