Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
36 commits
Select commit Hold shift + click to select a range
eab5d53
add replayssm configuration
Aug 14, 2026
a7e36e9
Refactor ReplaySSM backend integration to support FlashInfer and Triton
Aug 14, 2026
becbedd
improve replayssm performance by allowing for split kernel runs
Aug 14, 2026
b526ca5
add metadata building for FI ReplaySSM
Aug 14, 2026
27bcd0b
add proper ring flushing and advancement
Aug 14, 2026
10e9d18
validation for FlashInfer ReplaySSM cache mode to prevent unsupporte…
Aug 14, 2026
165f5a4
add FlashInfer ReplaySSM ring tracker updates and lifecycle tests
Aug 14, 2026
6f27481
remove useless flags
Aug 14, 2026
6ae4546
Add startup autotuning for FlashInfer ReplaySSM
Aug 15, 2026
0ec9060
Batch ReplaySSM tracker updates across layers
Aug 16, 2026
f93b18d
Autotune ReplaySSM precompute launch geometry
Aug 16, 2026
3247272
Share ReplaySSM trackers by KV cache group
Aug 16, 2026
7c0bfd7
Use native FlashInfer ReplaySSM autotuning
Aug 16, 2026
1cab249
Simplify FlashInfer ReplaySSM autotune warmup
Aug 16, 2026
9a88661
refactor(mamba): simplify FlashInfer ReplaySSM wiring
Aug 16, 2026
e8ea5d4
fix(mamba): support ReplaySSM autotuning on runner v2
Aug 17, 2026
67807dd
fix(mamba): enforce ReplaySSM integration contracts
Aug 17, 2026
fc8aff3
fix(mamba): warm ReplaySSM tracker kernels
Aug 17, 2026
06aaf92
fix(mamba): cover V1 tracker warmup layouts
Aug 17, 2026
e5e9466
refactor(mamba): simplify ReplaySSM integration
Aug 17, 2026
5e7d1b1
refactor(mamba): remove redundant ReplaySSM assignments
Aug 17, 2026
1f18705
refactor(mamba): name ReplaySSM tracker resets
Aug 17, 2026
200b7c7
refactor(mamba): drop unreachable ReplaySSM ring sizing
Aug 17, 2026
7a98740
Clean up ReplaySSM tracker grouping and add V2 e2e coverage.
Aug 27, 2026
a53b2fe
Refactor FlashInfer checkpointing SSU integration in tests
Aug 27, 2026
d024660
Remove excessive FlashInfer ReplaySSM tests
Aug 28, 2026
1523e79
remove `raise` if replayssm autotuning is not available
Aug 31, 2026
e64eab6
Merge remote-tracking branch 'upstream/main' into feat/add_flashinfer…
Aug 31, 2026
e3b819f
clean up tests; remove unnecessary fixtures
Aug 31, 2026
d455455
move replayssm warmup to a separate file
Aug 31, 2026
b464da0
undo removing comments
Aug 31, 2026
c89f2db
test(mamba): skip FlashInfer ReplaySSM warmup tests without CUDA
Sep 1, 2026
51c7a68
Merge branch 'main' into feat/add_flashinfer_replayssm_kernel
askliar Sep 1, 2026
6ce9d6e
Merge branch 'main' of https://github.com/vllm-project/vllm into feat…
Sep 1, 2026
00a1d2e
fix pre-commit
Sep 1, 2026
dc9fc45
fix test args
Sep 1, 2026
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
72 changes: 70 additions & 2 deletions tests/kernels/mamba/test_ssu_dispatch.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,12 +7,15 @@
import torch

from vllm.config.mamba import MambaBackendEnum, MambaConfig, MambaSSUAlgorithm
from vllm.model_executor.layers.mamba.mamba_utils import MambaStateShapeCalculator
from vllm.model_executor.layers.mamba.ops.ssu_dispatch import (
FlashInferSSUBackend,
TritonSSUBackend,
get_mamba_ssu_backend,
initialize_mamba_ssu_backend,
reset_replayssm_ring_trackers,
selective_state_update,
update_replayssm_ring_trackers,
)
from vllm.utils.torch_utils import set_random_seed
from vllm.v1.attention.backends.registry import MambaAttentionBackendEnum
Expand All @@ -30,6 +33,47 @@
HAS_FLASHINFER = False


@pytest.fixture(autouse=True)
def restore_backend_state():
import vllm.model_executor.layers.mamba.ops.ssu_dispatch as mod

old_backend = mod._mamba_ssu_backend
old_replayssm_kernel = mod._flashinfer_replayssm_kernel
yield
mod._mamba_ssu_backend = old_backend
mod._flashinfer_replayssm_kernel = old_replayssm_kernel


def test_flashinfer_replayssm_ring_tracker_lifecycle():
ring_start = torch.zeros(2, dtype=torch.int32, device="cuda")
prev_num_accepted = torch.zeros(2, dtype=torch.int32, device="cuda")
state_batch_indices = torch.tensor([1], dtype=torch.int32, device="cuda")

observed = []
for _ in range(33):
update_replayssm_ring_trackers(
ring_start,
prev_num_accepted,
state_batch_indices,
logical_window=16,
ring_buffer_len=17,
)
observed.append((int(ring_start[1]), int(prev_num_accepted[1])))

assert observed[4] == (0, 5)
assert observed[15] == (0, 16)
assert observed[16] == (16, 1)
assert observed[31] == (16, 16)
assert observed[32] == (15, 1)

reset_replayssm_ring_trackers(
ring_start,
prev_num_accepted,
state_batch_indices,
)
assert (ring_start[1].item(), prev_num_accepted[1].item()) == (0, 0)


def _kv_cache_config_with_ssu(
mamba_type: MambaAttentionBackendEnum = MambaAttentionBackendEnum.MAMBA2,
) -> KVCacheConfig:
Expand Down Expand Up @@ -116,11 +160,10 @@ def test_flashinfer_forwards_ssu_algorithm(
def test_uninitialized_backend_raises():
import vllm.model_executor.layers.mamba.ops.ssu_dispatch as mod

old = mod._mamba_ssu_backend
# restore_backend_state (autouse) puts the global back afterwards.
mod._mamba_ssu_backend = None
with pytest.raises(RuntimeError, match="not been initialized"):
get_mamba_ssu_backend()
mod._mamba_ssu_backend = old


@pytest.mark.parametrize(
Expand Down Expand Up @@ -186,3 +229,28 @@ def test_triton_basic_call():
out=out,
)
assert not torch.isnan(out).any()


@pytest.mark.parametrize(
("backend", "expected_ring_len"),
[
(MambaBackendEnum.TRITON, 16),
(MambaBackendEnum.FLASHINFER, 17),
],
)
def test_replayssm_physical_ring_shape(backend, expected_ring_len):
base_shapes = ((64, 3), (8, 4, 16))

shapes = MambaStateShapeCalculator.append_replayssm_ring(
base_shapes,
n_groups=4,
tp_world_size=2,
logical_window=16,
backend=backend,
)

assert shapes[2:] == (
(8, expected_ring_len, 4),
(8, expected_ring_len),
(2, expected_ring_len, 16),
)
150 changes: 150 additions & 0 deletions tests/model_executor/test_replayssm_warmup.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,150 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project

from types import SimpleNamespace
from unittest.mock import Mock, patch

import numpy as np
import pytest
import torch

from vllm.config.mamba import MambaBackendEnum
from vllm.model_executor.layers.mamba.mamba_mixer2 import MambaMixer2
from vllm.model_executor.warmup import replayssm_warmup as warmup
from vllm.platforms import current_platform
from vllm.utils.flashinfer import has_flashinfer

pytestmark = pytest.mark.skipif(
not current_platform.is_cuda() or not has_flashinfer(),
reason="FlashInfer ReplaySSM warmup tests require CUDA and FlashInfer",
)

PREFILL_KWARGS = {
"num_tokens": 128,
"skip_eplb": True,
"is_profile": True,
"randomize_inputs": True,
}


def _autotune_runner(
*,
use_v2_model_runner: bool = False,
query_len: int = 6,
max_num_seqs: int = 32,
num_blocks: int = 17,
max_num_tokens: int = 100,
use_replayssm: bool = True,
backend: MambaBackendEnum = MambaBackendEnum.FLASHINFER,
) -> SimpleNamespace:
return SimpleNamespace(
vllm_config=SimpleNamespace(
cache_config=SimpleNamespace(use_replayssm=use_replayssm),
mamba_config=SimpleNamespace(backend=backend),
use_v2_model_runner=use_v2_model_runner,
),
uniform_decode_query_len=query_len,
decode_query_len=query_len,
max_num_tokens=max_num_tokens,
scheduler_config=SimpleNamespace(max_num_seqs=max_num_seqs),
kv_cache_config=SimpleNamespace(num_blocks=num_blocks),
)


@pytest.mark.parametrize(
("runner_kwargs", "expected_num_reqs"),
[
# max_num_seqs (32) vs max_num_tokens // query_len (16) vs blocks-1 (16).
(dict(query_len=6, use_v2_model_runner=False), 16),
(dict(query_len=6, use_v2_model_runner=True), 16),
# num_blocks - 1 is the binding constraint.
(dict(query_len=1, max_num_tokens=128, max_num_seqs=64, num_blocks=5), 4),
],
ids=["v1", "v2", "clamped_to_state_capacity"],
)
def test_replayssm_autotune_decode_kwargs(runner_kwargs, expected_num_reqs):
query_len = runner_kwargs["query_len"]
with patch.object(
warmup, "flashinfer_replayssm_autotune_supported", return_value=True
):
result = warmup._replayssm_autotune_kwargs(_autotune_runner(**runner_kwargs))

expected_kwargs = {
**PREFILL_KWARGS,
"num_tokens": expected_num_reqs * query_len,
"uniform_decode": True,
}
if runner_kwargs.get("use_v2_model_runner"):
expected_kwargs["valid_dummy_state_slots"] = True
else:
expected_kwargs.update(
allow_microbatching=False,
force_attention=True,
profile_seq_lens=query_len + 1,
)
assert result == (expected_num_reqs, expected_kwargs)


@pytest.mark.parametrize(
("runner_kwargs", "flashinfer_supported"),
[
(dict(use_replayssm=False), True),
(dict(backend=MambaBackendEnum.TRITON), True),
({}, False),
],
ids=["replayssm_disabled", "non_flashinfer_backend", "kernel_unavailable"],
)
def test_replayssm_autotune_kwargs_skipped(runner_kwargs, flashinfer_supported):
with patch.object(
warmup,
"flashinfer_replayssm_autotune_supported",
return_value=flashinfer_supported,
):
result = warmup._replayssm_autotune_kwargs(_autotune_runner(**runner_kwargs))
assert result is None


def test_replayssm_autotune_slots_restore_state_and_trackers():
mixer = MambaMixer2.__new__(MambaMixer2)
torch.nn.Module.__init__(mixer)
mixer.use_replayssm = True
mixer.replayssm_buffer_len = 16
mixer.kv_cache = (
torch.full((4, 2), 3.0),
torch.full((4, 2), 3.0),
*(torch.full((4, 2, 17), 3.0) for _ in range(3)),
)
mixer._replayssm_ring_start = torch.full((4,), 3, dtype=torch.int32)
mixer._replayssm_prev_num_accepted = torch.full((4,), 3, dtype=torch.int32)
tracked = (
*mixer.kv_cache,
mixer._replayssm_ring_start,
mixer._replayssm_prev_num_accepted,
)

block_ids = np.arange(10, 14, dtype=np.int32).reshape(4, 1)
original_block_ids = block_ids.copy()
block_table = SimpleNamespace(block_table=SimpleNamespace(np=block_ids))
multi_group_block_table = SimpleNamespace(
block_tables=[block_table], commit_block_table=Mock()
)
runner = SimpleNamespace(
vllm_config=SimpleNamespace(use_v2_model_runner=False),
input_batch=SimpleNamespace(block_table=multi_group_block_table),
get_model=lambda: SimpleNamespace(modules=lambda: (mixer,)),
)

with warmup._temporary_replayssm_autotune_state(runner, 2):
assert block_ids[:2, 0].tolist() == [1, 2]
for tensor in tracked:
tensor[1:3].fill_(9)

assert np.array_equal(block_ids, original_block_ids)
assert multi_group_block_table.commit_block_table.call_args_list == [
((2,), {}),
((2,), {}),
]
for tensor in tracked:
assert torch.count_nonzero(tensor[1:3]) == 0
assert torch.all(tensor[0] == 3)
assert torch.all(tensor[3] == 3)
58 changes: 46 additions & 12 deletions tests/v1/attention/test_replayssm_metadata_builder.py
Original file line number Diff line number Diff line change
Expand Up @@ -16,6 +16,7 @@
create_common_attn_metadata,
create_vllm_config,
)
from vllm.config.mamba import MambaBackendEnum
from vllm.v1.kv_cache_interface import MambaSpec

BLOCK_SIZE = 16
Expand Down Expand Up @@ -187,34 +188,47 @@ class ReplaySSMBuildCase:
}


def _make_mamba_spec(buffer_len: int) -> MambaSpec:
# Five-tensor ReplaySSM page; the builder only reads shapes[4][0] (bc groups).
def _make_mamba_spec(
buffer_len: int,
mamba_backend: MambaBackendEnum,
) -> MambaSpec:
ring_buffer_len = buffer_len + (
1 if mamba_backend == MambaBackendEnum.FLASHINFER else 0
)
shapes = (
(1, 1),
(1, 1, 1),
(1, ring_buffer_len, 1),
(1, ring_buffer_len),
(1, ring_buffer_len, 1),
)
return MambaSpec(
block_size=BLOCK_SIZE,
shapes=(
(1, 1),
(1, 1, 1),
(1, buffer_len, 1),
(1, buffer_len),
(1, buffer_len, 1),
),
shapes=shapes,
dtypes=(torch.float32,),
)


def _create_replayssm_builder(
buffer_len: int, mamba_cache_mode: str = "none"
buffer_len: int,
mamba_cache_mode: str = "none",
*,
mamba_backend: MambaBackendEnum = MambaBackendEnum.TRITON,
) -> MockMambaBuilder:
vllm_config = create_vllm_config(
model_name="Qwen/Qwen3.5-0.8B", block_size=BLOCK_SIZE
)
# Set the flags after construction to skip validate_mamba_cached_kernel
# (it requires a Triton backend) on the mock model.
# (it requires a real SupportsReplaySSM model) on the mock model.
vllm_config.cache_config.use_replayssm = True
vllm_config.cache_config.replayssm_buffer_len = buffer_len
vllm_config.cache_config.mamba_cache_mode = mamba_cache_mode
vllm_config.mamba_config.backend = mamba_backend
return MockMambaBuilder(
_make_mamba_spec(buffer_len), ["layer0"], vllm_config, DEVICE
_make_mamba_spec(buffer_len, mamba_backend),
["layer0"],
vllm_config,
DEVICE,
)


Expand Down Expand Up @@ -254,3 +268,23 @@ def test_resumed_request_differs_from_fresh():

assert meta.write_pos_d.tolist()[:2] == [5, 0]
assert meta.is_flush_d.tolist()[:2] == [0, 0]


def test_flashinfer_replayssm_scratch_metadata_fresh_decode():
checkpointing_ssu = pytest.importorskip("flashinfer.mamba.checkpointing_ssu")
if not hasattr(checkpointing_ssu, "allocate_checkpointing_ssu_scratch"):
pytest.skip("FlashInfer does not expose ReplaySSM scratch allocation")

builder = _create_replayssm_builder(16, mamba_backend=MambaBackendEnum.FLASHINFER)
case = REPLAYSSM_BUILD_CASES["fresh_decode"]
meta = _build(builder, case)

assert meta.write_pos_d is None
assert meta.is_flush_d is None
assert meta.bc_pre_scratch is None
assert meta.replayssm_scratch is not None
assert [tensor.shape for tensor in meta.replayssm_scratch] == [
(1, 1, 32, 8),
(1, 1, 16),
(1, 1, 32, 8),
]
Loading
Loading