Skip to content
Open
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
23 changes: 23 additions & 0 deletions tests/renderers/test_hf.py
Original file line number Diff line number Diff line change
Expand Up @@ -13,6 +13,7 @@
_convert_developer_to_system,
_detect_content_format,
_detect_developer_role_support,
_ensure_mamba_checkpoint_token,
_get_hf_base_chat_template_params,
_template_error_reason,
_try_extract_ast,
Expand All @@ -31,6 +32,28 @@
chatml_jinja_path = VLLM_PATH / "examples/template_chatml.jinja"
assert chatml_jinja_path.exists()


class _CheckpointTokenizer:
def __init__(self):
self.special_tokens = []

def add_special_tokens(self, tokens):
self.special_tokens.extend(tokens["additional_special_tokens"])

def encode(self, token, add_special_tokens=False):
assert not add_special_tokens
return [123] if token in self.special_tokens else [1, 2]


def test_ensure_mamba_checkpoint_token_registers_single_token():
tokenizer = _CheckpointTokenizer()

token_id = _ensure_mamba_checkpoint_token(tokenizer, "<|mamba_checkpoint|>")

assert token_id == 123
assert tokenizer.special_tokens == ["<|mamba_checkpoint|>"]


# Define models, templates, and their corresponding expected outputs
MODEL_TEMPLATE_GENERATION_OUTPUT = [
(
Expand Down
7 changes: 7 additions & 0 deletions tests/v1/core/test_mamba_align_chunk_split.py
Original file line number Diff line number Diff line change
Expand Up @@ -392,3 +392,10 @@ def test_unaligned_resume_never_runs_past_its_block(
f"intermediate chunk end {end} is neither block-aligned nor the "
f"partial-tail stop ({tail_stop})"
)


def test_explicit_checkpoint_stops_prefill_chunk() -> None:
(request,) = create_requests(1, num_tokens=3602, block_size=ATTN_BLOCK_SIZE)
request.mamba_checkpoint_position = 800

assert _split(request, 1600, use_eagle=False) == 800
194 changes: 194 additions & 0 deletions tests/v1/core/test_mamba_checkpoint_scheduler.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,194 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
"""Unit tests for Mamba prefix checkpointing in the V1 Scheduler.

Covers:
1. Same-step producer/consumer pairing: a consumer request arriving in the same
scheduling step as its checkpoint producer inherits the producer's request ID,
gets mamba_checkpoint_source_block_ids, and schedules alongside it.
2. Cross-step dependency state machine: when a checkpoint is not yet ready, consumer
requests are skipped (waiting_for_mamba_checkpoint=True) without blocking unrelated
requests, and are resumed once the checkpoint is marked ready.
"""

import pytest
import torch

from vllm.config import (
CacheConfig,
ModelConfig,
SchedulerConfig,
VllmConfig,
)
from vllm.sampling_params import SamplingParams
from vllm.utils.hashing import sha256 as vllm_sha256
from vllm.v1.core.kv_cache_utils import get_request_block_hasher, init_none_hash
from vllm.v1.core.sched.scheduler import Scheduler
from vllm.v1.core.single_type_kv_cache_manager import register_all_kvcache_specs
from vllm.v1.kv_cache_interface import (
FullAttentionSpec,
KVCacheConfig,
KVCacheGroupSpec,
MambaSpec,
)
from vllm.v1.request import Request
from vllm.v1.structured_output import StructuredOutputManager

pytestmark = pytest.mark.cpu_test

BLOCK_SIZE = 16


def _create_hybrid_mamba_scheduler(
num_blocks: int = 1000,
block_size: int = BLOCK_SIZE,
num_prefill_checkpoint_blocks: int = 0,
) -> Scheduler:
from unittest.mock import patch
from transformers import OPTConfig

mock_cfg = OPTConfig(
vocab_size=1000,
hidden_size=64,
num_hidden_layers=1,
num_attention_heads=1,
)
mock_cfg.architectures = ["OPTForCausalLM"]

with patch("vllm.config.model.get_config", return_value=mock_cfg):
model_config = ModelConfig(
model="facebook/opt-125m",
tokenizer="facebook/opt-125m",
seed=42,
skip_tokenizer_init=True,
)
vllm_config = VllmConfig(
scheduler_config=SchedulerConfig(
max_num_seqs=8,
max_num_batched_tokens=8192,
max_model_len=8192,
enable_chunked_prefill=True,
is_encoder_decoder=False,
watermark=0.0,
),
model_config=model_config,
cache_config=CacheConfig(
block_size=block_size,
enable_prefix_caching=True,
mamba_cache_mode="align",
mamba_checkpoint_token="<|mamba_checkpoint|>",
),
)
vllm_config.cache_config.num_gpu_blocks = num_blocks
kv_cache_config = KVCacheConfig(
num_blocks=num_blocks,
kv_cache_tensors=[],
kv_cache_groups=[
KVCacheGroupSpec(
["fa"],
FullAttentionSpec(
block_size=block_size,
num_kv_heads=1,
head_size=1,
dtype=torch.float32,
),
),
KVCacheGroupSpec(
["mamba"],
MambaSpec(
block_size=block_size,
shapes=((1, 1),),
dtypes=(torch.float32,),
mamba_cache_mode="align",
num_speculative_blocks=0,
num_prefill_checkpoint_blocks=num_prefill_checkpoint_blocks,
),
),
],
)
register_all_kvcache_specs(vllm_config)
return Scheduler(
vllm_config=vllm_config,
kv_cache_config=kv_cache_config,
structured_output_manager=StructuredOutputManager(vllm_config),
block_size=block_size,
hash_block_size=block_size,
log_stats=True,
)


def test_scheduler_producer_checkpoint_stops_at_boundary():
"""Producer prefill chunk stops exactly at the mamba_checkpoint_position."""
scheduler = _create_hybrid_mamba_scheduler()
init_none_hash(vllm_sha256)
block_hasher = get_request_block_hasher(BLOCK_SIZE, vllm_sha256)

checkpoint_pos = 48
tokens_producer = [10] * 100

req_producer = Request(
request_id="producer_0",
prompt_token_ids=tokens_producer,
sampling_params=SamplingParams(max_tokens=5),
pooling_params=None,
block_hasher=block_hasher,
mamba_checkpoint_position=checkpoint_pos,
)

scheduler.add_request(req_producer)
sched_out = scheduler.schedule()

# Producer chunk stops at checkpoint boundary (48 tokens).
assert [r.req_id for r in sched_out.scheduled_new_reqs] == ["producer_0"]
assert sched_out.num_scheduled_tokens["producer_0"] == checkpoint_pos
assert scheduler.kv_cache_manager.has_unready_checkpoint(req_producer)


def test_scheduler_cross_step_checkpoint_pending_and_wakeup():
"""Consumer is skipped while checkpoint is pending, and woken up once ready."""
scheduler = _create_hybrid_mamba_scheduler()
init_none_hash(vllm_sha256)
block_hasher = get_request_block_hasher(BLOCK_SIZE, vllm_sha256)

checkpoint_pos = 48
tokens_producer = [10] * 100
tokens_consumer = [10] * checkpoint_pos + [20] * 50

req_producer = Request(
request_id="producer_0",
prompt_token_ids=tokens_producer,
sampling_params=SamplingParams(max_tokens=5),
pooling_params=None,
block_hasher=block_hasher,
mamba_checkpoint_position=checkpoint_pos,
)
req_consumer = Request(
request_id="consumer_0",
prompt_token_ids=tokens_consumer,
sampling_params=SamplingParams(max_tokens=5),
pooling_params=None,
block_hasher=block_hasher,
mamba_checkpoint_position=checkpoint_pos,
)

# Step 1: Producer is scheduled.
scheduler.add_request(req_producer)
out1 = scheduler.schedule()
assert [r.req_id for r in out1.scheduled_new_reqs] == ["producer_0"]
assert scheduler.kv_cache_manager.has_unready_checkpoint(req_producer)

# Step 2: Consumer arrives while producer's checkpoint is still unready.
scheduler.add_request(req_consumer)
out2 = scheduler.schedule()
assert len(out2.scheduled_new_reqs) == 0
assert req_consumer.waiting_for_mamba_checkpoint is True

# Step 3: Producer completes and marks checkpoint ready.
scheduler.kv_cache_manager.mark_checkpoint_ready("producer_0")
assert not scheduler.kv_cache_manager.has_unready_checkpoint(req_producer)

# Step 4: Next schedule step resumes consumer, hitting the prefix cache.
out3 = scheduler.schedule()
assert [r.req_id for r in out3.scheduled_new_reqs] == ["consumer_0"]
assert req_consumer.waiting_for_mamba_checkpoint is False
assert req_consumer.mamba_prefix_producer_id is None
138 changes: 61 additions & 77 deletions tests/v1/core/test_prefix_caching.py
Original file line number Diff line number Diff line change
Expand Up @@ -1248,83 +1248,6 @@ def make_kv_cache_config_three_types(
)


def test_prefix_cache_hit_uses_per_group_dcp_geometry():
"""Prefix lookup must use each group's DCP size, not the process-wide one.

Target and draft MLA are both sharded (DCP=8); Mamba stays replicated
(DCP=1). Hits then align to the sharded full-attention block, not the
unsharded page size.
"""
block_size = 16
dcp = 8
sharded_block = block_size * dcp
kv_cache_config = KVCacheConfig(
num_blocks=64,
kv_cache_tensors=[],
kv_cache_groups=[
KVCacheGroupSpec(
["target_mla"],
FullAttentionSpec(
block_size=block_size,
num_kv_heads=1,
head_size=1,
dtype=torch.float32,
),
),
KVCacheGroupSpec(
["draft_mla"],
MLAAttentionSpec(
block_size=block_size,
num_kv_heads=1,
head_size=1,
dtype=torch.float32,
),
),
KVCacheGroupSpec(
["mamba"],
MambaSpec(
block_size=block_size,
shapes=(1, 1),
dtypes=(torch.float32,),
),
),
],
)
manager = KVCacheManager(
kv_cache_config,
max_model_len=8192,
enable_caching=True,
hash_block_size=block_size,
scheduler_block_size=sharded_block,
dcp_world_size=dcp,
)
target_mgr, draft_mgr, mamba_mgr = manager.coordinator.single_type_managers
assert target_mgr.dcp_world_size == dcp
assert draft_mgr.dcp_world_size == dcp
assert mamba_mgr.dcp_world_size == 1
assert target_mgr.block_size == sharded_block
assert draft_mgr.block_size == sharded_block
assert mamba_mgr.block_size == block_size

hash_fn = sha256
common_token_ids = [i for i in range(2) for _ in range(sharded_block)]
req0 = make_request("0", common_token_ids + [99] * 7, block_size, hash_fn)
computed_blocks, num_computed_tokens, _ = manager.get_computed_blocks(req0)
assert num_computed_tokens == 0
blocks = manager.allocate_slots(
req0, len(req0.prompt_token_ids), 0, computed_blocks
)
assert blocks is not None

req1 = make_request("1", common_token_ids + [100] * 5, block_size, hash_fn)
computed_blocks, num_computed_tokens, _ = manager.get_computed_blocks(req1)
assert num_computed_tokens == 2 * sharded_block
assert [len(group) for group in computed_blocks.blocks] == [2, 2, 16]

manager.free(req0)
manager.free(req1)


@pytest.mark.parametrize("hash_fn", [sha256, sha256_cbor])
def test_prefill(hash_fn):
block_size = 16
Expand Down Expand Up @@ -2201,6 +2124,67 @@ def test_hybrid_cache_mamba_align_shared_prefix_detection():
manager.free(req_2)


def test_mamba_checkpoint_only_caches_explicit_boundary():
config = KVCacheConfig(
num_blocks=100,
kv_cache_tensors=[],
kv_cache_groups=[
KVCacheGroupSpec(
["full"],
FullAttentionSpec(
block_size=16,
num_kv_heads=1,
head_size=1,
dtype=torch.float32,
),
),
KVCacheGroupSpec(
["mamba"],
MambaSpec(
block_size=32,
shapes=((1, 1),),
dtypes=(torch.float32,),
mamba_cache_mode="align",
),
),
],
)
manager = make_kv_cache_manager(
config,
max_model_len=128,
enable_caching=True,
hash_block_size=16,
)
request = make_request("checkpoint", list(range(48)), 16, sha256)
request.mamba_checkpoint_position = 16

assert manager.allocate_slots(request, 16) is not None
mamba_group_id = 1
assert manager.has_unready_checkpoint(request)
assert (
manager.block_pool.get_cached_block(request.block_hashes[0], [mamba_group_id])
is None
)
manager.mark_checkpoint_ready(request.request_id)
assert not manager.has_unready_checkpoint(request)
assert (
manager.block_pool.get_cached_block(request.block_hashes[0], [mamba_group_id])
is not None
)

request.num_computed_tokens = 16
assert manager.allocate_slots(request, 32) is not None
assert (
manager.block_pool.get_cached_block(request.block_hashes[1], [mamba_group_id])
is None
)
assert (
manager.block_pool.get_cached_block(request.block_hashes[2], [mamba_group_id])
is None
)
manager.free(request)


def test_hybrid_model_mamba_align_with_dynamic_draft_tokens():
"""Regression test for https://github.com/vllm-project/vllm/issues/39271.

Expand Down
Loading
Loading