From eab5d5330104bc3051905b7c7c292082dcebe17b Mon Sep 17 00:00:00 2001 From: Andrii Skliar Date: Fri, 14 Aug 2026 15:14:42 +0200 Subject: [PATCH 01/33] add replayssm configuration Signed-off-by: Andrii Skliar --- tests/kernels/mamba/test_ssu_dispatch.py | 160 +++++++ vllm/config/cache.py | 6 +- vllm/config/vllm.py | 12 +- .../layers/mamba/mamba_mixer2.py | 19 +- .../layers/mamba/ops/ssu_dispatch.py | 412 ++++++++++++++++-- vllm/v1/worker/gpu/model_runner.py | 4 +- vllm/v1/worker/gpu_model_runner.py | 4 +- 7 files changed, 560 insertions(+), 57 deletions(-) diff --git a/tests/kernels/mamba/test_ssu_dispatch.py b/tests/kernels/mamba/test_ssu_dispatch.py index c37f74eb5079..082784d3126a 100644 --- a/tests/kernels/mamba/test_ssu_dispatch.py +++ b/tests/kernels/mamba/test_ssu_dispatch.py @@ -8,11 +8,17 @@ from vllm.config.mamba import MambaBackendEnum, MambaConfig, MambaSSUAlgorithm from vllm.model_executor.layers.mamba.ops.ssu_dispatch import ( + FlashInferReplaySSMBackend, FlashInferSSUBackend, + TritonReplaySSMBackend, TritonSSUBackend, get_mamba_ssu_backend, + get_replayssm_backend, initialize_mamba_ssu_backend, + initialize_replayssm_backend, selective_state_update, + selective_state_update_replayssm, + translate_vllm_replayssm_bookkeeping, ) from vllm.utils.torch_utils import set_random_seed from vllm.v1.attention.backends.registry import MambaAttentionBackendEnum @@ -29,6 +35,13 @@ except ImportError: HAS_FLASHINFER = False +try: + from flashinfer.mamba import checkpointing_ssu # noqa: F401 + + HAS_FLASHINFER_CHECKPOINTING_SSU = True +except ImportError: + HAS_FLASHINFER_CHECKPOINTING_SSU = False + def _kv_cache_config_with_ssu( mamba_type: MambaAttentionBackendEnum = MambaAttentionBackendEnum.MAMBA2, @@ -186,3 +199,150 @@ def test_triton_basic_call(): out=out, ) assert not torch.isnan(out).any() + + +def test_replayssm_default_backend_is_triton(): + initialize_replayssm_backend(MambaConfig(), use_replayssm=True) + backend = get_replayssm_backend() + assert isinstance(backend, TritonReplaySSMBackend) + assert backend.name == "triton" + + +def test_replayssm_explicit_triton_backend(): + initialize_replayssm_backend( + MambaConfig(backend=MambaBackendEnum.TRITON), use_replayssm=True + ) + backend = get_replayssm_backend() + assert isinstance(backend, TritonReplaySSMBackend) + + +@pytest.mark.skipif( + not HAS_FLASHINFER_CHECKPOINTING_SSU, + reason="flashinfer.mamba.checkpointing_ssu not available", +) +def test_replayssm_flashinfer_backend_init(): + initialize_replayssm_backend( + MambaConfig(backend=MambaBackendEnum.FLASHINFER), use_replayssm=True + ) + backend = get_replayssm_backend() + assert isinstance(backend, FlashInferReplaySSMBackend) + assert backend.name == "flashinfer" + + +def test_replayssm_disabled_clears_backend(): + initialize_replayssm_backend(MambaConfig(), use_replayssm=True) + assert get_replayssm_backend() is not None + initialize_replayssm_backend(MambaConfig(), use_replayssm=False) + with pytest.raises(RuntimeError, match="not been initialized"): + get_replayssm_backend() + + +def test_replayssm_cpu_backend_rejected(): + with pytest.raises(ValueError, match="does not support mamba backend"): + initialize_replayssm_backend( + MambaConfig(backend=MambaBackendEnum.CPU), use_replayssm=True + ) + + +def test_replayssm_uninitialized_backend_raises(): + import vllm.model_executor.layers.mamba.ops.ssu_dispatch as mod + + old = mod._replayssm_backend + mod._replayssm_backend = None + try: + with pytest.raises(RuntimeError, match="not been initialized"): + get_replayssm_backend() + finally: + mod._replayssm_backend = old + + +def test_replayssm_bookkeeping_adapter_not_implemented(): + with pytest.raises(NotImplementedError, match="bookkeeping adapter"): + translate_vllm_replayssm_bookkeeping( + write_pos=torch.zeros(1, dtype=torch.int32), + is_flush=torch.zeros(1, dtype=torch.int8), + state_batch_indices=None, + max_cache_len=16, + batch=1, + ) + + +@pytest.mark.skipif( + not HAS_FLASHINFER_CHECKPOINTING_SSU, + reason="flashinfer.mamba.checkpointing_ssu not available", +) +def test_replayssm_flashinfer_call_hits_bookkeeping_gap(monkeypatch): + import flashinfer.mamba + + kernel = Mock() + monkeypatch.setattr(flashinfer.mamba, "checkpointing_ssu", kernel) + backend = FlashInferReplaySSMBackend( + MambaConfig(backend=MambaBackendEnum.FLASHINFER) + ) + + batch, nheads, dim, dstate, ngroups, L = 1, 2, 4, 8, 1, 16 + state = torch.empty(1, nheads, dim, dstate) + x = torch.empty(batch, nheads, dim) + dt = torch.empty(batch, nheads, dim) + A = torch.empty(nheads, dim, dstate) + B = torch.empty(batch, ngroups, dstate) + C = torch.empty(batch, ngroups, dstate) + D = torch.empty(nheads, dim) + dt_bias = torch.empty(nheads, dim) + out = torch.empty_like(x) + x_cache = torch.empty(1, nheads, L, dim) + dt_cache = torch.empty(1, nheads, L) + B_cache = torch.empty(1, ngroups, L, dstate) + write_pos = torch.zeros(batch, dtype=torch.int32) + is_flush = torch.zeros(batch, dtype=torch.int8) + + with pytest.raises(NotImplementedError, match="bookkeeping adapter"): + backend( + state, + x, + dt, + A, + B, + C, + D=D, + dt_bias=dt_bias, + x_cache=x_cache, + dt_cache=dt_cache, + B_cache=B_cache, + write_pos=write_pos, + is_flush=is_flush, + max_cache_len=L, + out=out, + ) + kernel.assert_not_called() + + +@pytest.mark.skipif( + HAS_FLASHINFER_CHECKPOINTING_SSU, + reason="flashinfer checkpointing_ssu is installed", +) +def test_replayssm_flashinfer_import_error(): + with pytest.raises(ImportError, match="FlashInfer is required"): + FlashInferReplaySSMBackend(MambaConfig(backend=MambaBackendEnum.FLASHINFER)) + + +def test_replayssm_dispatch_fn_uses_initialized_backend(): + import vllm.model_executor.layers.mamba.ops.ssu_dispatch as mod + + called = Mock(return_value=torch.empty(1)) + old = mod._replayssm_backend + mod._replayssm_backend = called + try: + tensor = torch.empty(1) + selective_state_update_replayssm( + tensor, + tensor, + tensor, + tensor, + tensor, + tensor, + out=tensor, + ) + assert called.call_count == 1 + finally: + mod._replayssm_backend = old diff --git a/vllm/config/cache.py b/vllm/config/cache.py index dc4e6faf3cec..40741e5eaddf 100644 --- a/vllm/config/cache.py +++ b/vllm/config/cache.py @@ -201,9 +201,9 @@ class CacheConfig: """Use the ReplaySSM Mamba2 decode kernel: cache recent SSM inputs and skip the per-step full-state store, writing the checkpoint back only on flush. Requires mamba_cache_mode 'none' or 'align' (prefix caching) and the Triton - mamba backend; standard (non-speculative) decode only. In align mode flushes - are most efficient when mamba_block_size is a multiple of replayssm_buffer_len, - but this is not required.""" + or FlashInfer mamba backend; standard (non-speculative) decode only. In align + mode flushes are most efficient when mamba_block_size is a multiple of + replayssm_buffer_len, but this is not required.""" use_kda_recoverssm: bool = field(default=False, init=False) """Whether Kimi-K3 KDA uses RecoverSSM speculative decode.""" diff --git a/vllm/config/vllm.py b/vllm/config/vllm.py index a7fa4e1dca86..f575c89e44ce 100644 --- a/vllm/config/vllm.py +++ b/vllm/config/vllm.py @@ -2708,8 +2708,16 @@ def validate_mamba_cached_kernel(self) -> "VllmConfig": "--use-replayssm supports prefix caching only in align mode; " "pass --mamba-cache-mode align" ) - if self.mamba_config.backend != MambaBackendEnum.TRITON: - raise ValueError("--use-replayssm requires --mamba-backend triton") + if self.num_speculative_tokens > 0: + raise ValueError("--use-replayssm does not support speculative decoding") + if self.mamba_config.backend not in ( + MambaBackendEnum.TRITON, + MambaBackendEnum.FLASHINFER, + ): + raise ValueError( + "--use-replayssm requires --mamba-backend triton or flashinfer " + f"(got {self.mamba_config.backend.value!r})" + ) if ( self.kv_transfer_config is not None and self.kv_transfer_config.is_kv_transfer_instance diff --git a/vllm/model_executor/layers/mamba/mamba_mixer2.py b/vllm/model_executor/layers/mamba/mamba_mixer2.py index 513402a0b149..7adcbc0f2d2e 100644 --- a/vllm/model_executor/layers/mamba/mamba_mixer2.py +++ b/vllm/model_executor/layers/mamba/mamba_mixer2.py @@ -32,13 +32,13 @@ causal_conv1d_update, ) from vllm.model_executor.layers.mamba.ops.layernorm_gated import rms_norm_gated -from vllm.model_executor.layers.mamba.ops.selective_state_update_replayssm_output_only import ( # noqa: E501 - selective_state_update_replayssm_output_only, -) from vllm.model_executor.layers.mamba.ops.ssd_combined import ( mamba_chunk_scan_combined_varlen, ) -from vllm.model_executor.layers.mamba.ops.ssu_dispatch import selective_state_update +from vllm.model_executor.layers.mamba.ops.ssu_dispatch import ( + selective_state_update, + selective_state_update_replayssm, +) from vllm.model_executor.layers.quantization import QuantizationConfig from vllm.model_executor.model_loader.weight_utils import ( LoaderFunction, @@ -1060,7 +1060,7 @@ def conv_ssm_forward( ) if self.use_replayssm: assert self.replayssm_buffer_len is not None - selective_state_update_replayssm_output_only( + selective_state_update_replayssm( ssm_state, hidden_states_d, dt_d, @@ -1079,15 +1079,6 @@ def conv_ssm_forward( max_cache_len=self.replayssm_buffer_len, state_batch_indices=state_indices_tensor_d_input, out=preallocated_ssm_out_d, - # Stochastic Rounding for the vanilla decode path is read - # from mamba_config inside ssu_dispatch; the replay kernel - # isn't a dispatch backend, so pass it here. - enable_stochastic_rounding=( - self.mamba_config.enable_stochastic_rounding - ), - cache_philox_rounds=( - self.mamba_config.stochastic_rounding_philox_rounds - ), ) else: selective_state_update( diff --git a/vllm/model_executor/layers/mamba/ops/ssu_dispatch.py b/vllm/model_executor/layers/mamba/ops/ssu_dispatch.py index 24bdba688e70..1083104ea594 100644 --- a/vllm/model_executor/layers/mamba/ops/ssu_dispatch.py +++ b/vllm/model_executor/layers/mamba/ops/ssu_dispatch.py @@ -3,10 +3,17 @@ """ Dispatch module for Mamba selective state update (SSU) backends. -Provides a unified `selective_state_update` function that dispatches to -the Triton, FlashInfer, or CPU backend based on the configured -`MambaBackendEnum`. On CPU-only platforms (PowerPC, x86 without CUDA) -the backend defaults to 'cpu'. +Provides unified ``selective_state_update`` (baseline decode) and +``selective_state_update_replayssm`` (cached-input ReplaySSM decode) that +dispatch to Triton / FlashInfer / CPU based on ``MambaBackendEnum``. On +CPU-only platforms (PowerPC, x86 without CUDA) the baseline SSU backend +defaults to ``cpu``. + +The FlashInfer ReplaySSM path imports ``flashinfer.mamba.checkpointing_ssu`` +and reshapes T=1 AR tensors, but the vLLM ``write_pos`` / ``is_flush`` / +``bc_pre`` → FlashInfer ``ring_start`` / ``prev_num_accepted_tokens`` / +scratch contract is intentionally unfinished (see +``translate_vllm_replayssm_bookkeeping``). """ from abc import ABC, abstractmethod @@ -262,52 +269,385 @@ def __call__( _mamba_ssu_backend: MambaSSUBackend | None = None +class ReplaySSMBackend(ABC): + """Abstract base class for ReplaySSM decode backends.""" + + def __init__(self, mamba_config: MambaConfig): + self._mamba_config = mamba_config + + @property + @abstractmethod + def name(self) -> str: ... + + @abstractmethod + def __call__( + self, + state: torch.Tensor, + x: torch.Tensor, + dt: torch.Tensor, + A: torch.Tensor, + B: torch.Tensor, + C: torch.Tensor, + D: torch.Tensor | None = None, + dt_bias: torch.Tensor | None = None, + z: torch.Tensor | None = None, + dt_softplus: bool = False, + x_cache: torch.Tensor | None = None, + dt_cache: torch.Tensor | None = None, + B_cache: torch.Tensor | None = None, + bc_pre: torch.Tensor | None = None, + write_pos: torch.Tensor | None = None, + is_flush: torch.Tensor | None = None, + max_cache_len: int = 16, + state_batch_indices: torch.Tensor | None = None, + null_block_id: int = NULL_BLOCK_ID, + out: torch.Tensor | None = None, + ) -> torch.Tensor: ... + + +class TritonReplaySSMBackend(ReplaySSMBackend): + """vLLM's in-tree Triton ReplaySSM output_only kernel.""" + + def __init__(self, mamba_config: MambaConfig): + super().__init__(mamba_config) + from vllm.model_executor.layers.mamba.ops.selective_state_update_replayssm_output_only import ( # noqa: E501 + selective_state_update_replayssm_output_only as _triton_replayssm, + ) + + self._kernel = _triton_replayssm + + @property + def name(self) -> str: + return "triton" + + def __call__( + self, + state: torch.Tensor, + x: torch.Tensor, + dt: torch.Tensor, + A: torch.Tensor, + B: torch.Tensor, + C: torch.Tensor, + D: torch.Tensor | None = None, + dt_bias: torch.Tensor | None = None, + z: torch.Tensor | None = None, + dt_softplus: bool = False, + x_cache: torch.Tensor | None = None, + dt_cache: torch.Tensor | None = None, + B_cache: torch.Tensor | None = None, + bc_pre: torch.Tensor | None = None, + write_pos: torch.Tensor | None = None, + is_flush: torch.Tensor | None = None, + max_cache_len: int = 16, + state_batch_indices: torch.Tensor | None = None, + null_block_id: int = NULL_BLOCK_ID, + out: torch.Tensor | None = None, + ) -> torch.Tensor: + return self._kernel( + state, + x, + dt, + A, + B, + C, + D=D, + dt_bias=dt_bias, + z=z, + dt_softplus=dt_softplus, + x_cache=x_cache, + dt_cache=dt_cache, + B_cache=B_cache, + bc_pre=bc_pre, + write_pos=write_pos, + is_flush=is_flush, + max_cache_len=max_cache_len, + state_batch_indices=state_batch_indices, + null_block_id=null_block_id, + out=out, + enable_stochastic_rounding=self._mamba_config.enable_stochastic_rounding, + cache_philox_rounds=self._mamba_config.stochastic_rounding_philox_rounds, + ) + + +def translate_vllm_replayssm_bookkeeping( + *, + write_pos: torch.Tensor, + is_flush: torch.Tensor, + state_batch_indices: torch.Tensor | None, + max_cache_len: int, + batch: int, +) -> tuple[torch.Tensor, torch.Tensor]: + """Map vLLM ReplaySSM host metadata to FlashInfer checkpointing_ssu args. + + vLLM (Triton ReplaySSM) provides per-decode-row: + - ``write_pos``: ring append cursor in ``[0, max_cache_len)`` + - ``is_flush``: materialize full SSM state this step + + FlashInfer ``checkpointing_ssu`` expects per-cache-slot: + - ``ring_start``: oldest live ring index + - ``prev_num_accepted_tokens``: live history length to replay + + Returns: + ``(ring_start, prev_num_accepted_tokens)`` shaped for the FlashInfer + call (typically indexed by cache slot, not batch row). + + Raises: + NotImplementedError: contract adapter not wired yet. + """ + raise NotImplementedError( + "FlashInfer ReplaySSM bookkeeping adapter is not implemented yet. " + "Map vLLM write_pos/is_flush (and state_batch_indices) to FlashInfer " + "ring_start/prev_num_accepted_tokens for max_cache_len=" + f"{max_cache_len}, batch={batch}." + ) + + +class FlashInferReplaySSMBackend(ReplaySSMBackend): + """FlashInfer ``checkpointing_ssu`` ReplaySSM backend (contract pending).""" + + def __init__(self, mamba_config: MambaConfig): + super().__init__(mamba_config) + try: + from flashinfer.mamba import checkpointing_ssu as _fi_checkpointing_ssu + except ImportError as e: + raise ImportError( + "FlashInfer is required for the flashinfer ReplaySSM backend. " + "Please install flashinfer with mamba.checkpointing_ssu support: " + "pip install flashinfer-python" + ) from e + self._kernel = _fi_checkpointing_ssu + + @property + def name(self) -> str: + return "flashinfer" + + def __call__( + self, + state: torch.Tensor, + x: torch.Tensor, + dt: torch.Tensor, + A: torch.Tensor, + B: torch.Tensor, + C: torch.Tensor, + D: torch.Tensor | None = None, + dt_bias: torch.Tensor | None = None, + z: torch.Tensor | None = None, + dt_softplus: bool = False, + x_cache: torch.Tensor | None = None, + dt_cache: torch.Tensor | None = None, + B_cache: torch.Tensor | None = None, + bc_pre: torch.Tensor | None = None, + write_pos: torch.Tensor | None = None, + is_flush: torch.Tensor | None = None, + max_cache_len: int = 16, + state_batch_indices: torch.Tensor | None = None, + null_block_id: int = NULL_BLOCK_ID, + out: torch.Tensor | None = None, + ) -> torch.Tensor: + del bc_pre # Triton-only scratch; FI uses its own precompute buffers. + if write_pos is None or is_flush is None: + raise ValueError( + "FlashInfer ReplaySSM requires write_pos and is_flush metadata" + ) + if out is None: + raise ValueError("FlashInfer ReplaySSM requires a preallocated out tensor") + if x_cache is None or dt_cache is None or B_cache is None: + raise ValueError( + "FlashInfer ReplaySSM requires x_cache, dt_cache, and B_cache" + ) + + # Mechanical T=1 reshape for AR decode. MTP (T>1) is out of scope until + # ReplaySSM speculative decode is enabled. + if x.dim() == 3: + x_t = x.unsqueeze(1) + dt_t = dt.unsqueeze(1) + B_t = B.unsqueeze(1) + C_t = C.unsqueeze(1) + out_t = out.unsqueeze(1) + z_t = z.unsqueeze(1) if z is not None else None + else: + x_t, dt_t, B_t, C_t, out_t, z_t = x, dt, B, C, out, z + + batch = x_t.shape[0] + ring_start, prev_num_accepted = translate_vllm_replayssm_bookkeeping( + write_pos=write_pos, + is_flush=is_flush, + state_batch_indices=state_batch_indices, + max_cache_len=max_cache_len, + batch=batch, + ) + + rand_seed = ( + torch.randint(0, 2**32, (1,), device=state.device, dtype=torch.int64) + if self._mamba_config.enable_stochastic_rounding + else None + ) + indices = state_batch_indices + if indices is not None and indices.dim() > 1: + indices = indices[:, 0] + + return self._kernel( + state, + x_cache, + B_cache, + dt_cache, + ring_start, + prev_num_accepted, + x_t, + dt_t, + A, + B_t, + C_t, + out_t, + D=D, + z=z_t, + dt_bias=dt_bias, + dt_softplus=dt_softplus, + state_batch_indices=indices, + pad_slot_id=null_block_id, + rand_seed=rand_seed, + philox_rounds=self._mamba_config.stochastic_rounding_philox_rounds or 10, + ) + + +_REPLAYSSM_BACKEND_REGISTRY: dict[MambaBackendEnum, type[ReplaySSMBackend]] = { + MambaBackendEnum.TRITON: TritonReplaySSMBackend, + MambaBackendEnum.FLASHINFER: FlashInferReplaySSMBackend, +} + +_replayssm_backend: ReplaySSMBackend | None = None + + +def initialize_replayssm_backend( + mamba_config: MambaConfig, + *, + use_replayssm: bool, +) -> None: + """Initialize the global ReplaySSM backend when ``--use-replayssm`` is set.""" + global _replayssm_backend + if not use_replayssm: + _replayssm_backend = None + return + + backend = mamba_config.backend + if backend not in _REPLAYSSM_BACKEND_REGISTRY: + raise ValueError( + f"--use-replayssm does not support mamba backend {backend.value!r}. " + f"Valid options: {[b.value for b in _REPLAYSSM_BACKEND_REGISTRY]}" + ) + + backend_cls = _REPLAYSSM_BACKEND_REGISTRY[backend] + if isinstance(_replayssm_backend, backend_cls): + return + + _replayssm_backend = backend_cls(mamba_config) + logger.info("Using %s ReplaySSM backend.", _replayssm_backend.name) + + +def get_replayssm_backend() -> ReplaySSMBackend: + """Get the current ReplaySSM backend. Raises if not initialized.""" + if _replayssm_backend is None: + raise RuntimeError( + "ReplaySSM backend has not been initialized. " + "Call initialize_mamba_ssu_backend() with use_replayssm=True first." + ) + return _replayssm_backend + + +def selective_state_update_replayssm( + state: torch.Tensor, + x: torch.Tensor, + dt: torch.Tensor, + A: torch.Tensor, + B: torch.Tensor, + C: torch.Tensor, + D: torch.Tensor | None = None, + dt_bias: torch.Tensor | None = None, + z: torch.Tensor | None = None, + dt_softplus: bool = False, + x_cache: torch.Tensor | None = None, + dt_cache: torch.Tensor | None = None, + B_cache: torch.Tensor | None = None, + bc_pre: torch.Tensor | None = None, + write_pos: torch.Tensor | None = None, + is_flush: torch.Tensor | None = None, + max_cache_len: int = 16, + state_batch_indices: torch.Tensor | None = None, + null_block_id: int = NULL_BLOCK_ID, + out: torch.Tensor | None = None, +) -> torch.Tensor: + """Unified dispatch for ReplaySSM selective state update.""" + return get_replayssm_backend()( + state, + x, + dt, + A, + B, + C, + D=D, + dt_bias=dt_bias, + z=z, + dt_softplus=dt_softplus, + x_cache=x_cache, + dt_cache=dt_cache, + B_cache=B_cache, + bc_pre=bc_pre, + write_pos=write_pos, + is_flush=is_flush, + max_cache_len=max_cache_len, + state_batch_indices=state_batch_indices, + null_block_id=null_block_id, + out=out, + ) + + def initialize_mamba_ssu_backend( mamba_config: MambaConfig, kv_cache_config: KVCacheConfig, + *, + use_replayssm: bool = False, ) -> None: - """Initialize the global Mamba SSU backend. + """Initialize the global Mamba SSU backend (and ReplaySSM when enabled). - No-op if `kv_cache_config` contains no specs that call - selective_state_update. + No-op for baseline SSU if `kv_cache_config` contains no specs that call + selective_state_update. Always (re)considers ReplaySSM when + ``use_replayssm`` is set. """ - if not any( + if any( isinstance(g.kv_cache_spec, MambaSpec) and g.kv_cache_spec.mamba_type in (MambaAttentionBackendEnum.MAMBA1, MambaAttentionBackendEnum.MAMBA2) for g in kv_cache_config.kv_cache_groups ): - return - - global _mamba_ssu_backend - - backend = mamba_config.backend - - # On CPU-only platforms (PowerPC, x86 without CUDA) Triton JIT is - # unstable or unavailable. Silently fall back to the CPU - # backend unless the user explicitly chose something other than "triton". - if backend == MambaBackendEnum.TRITON: - from vllm.platforms import current_platform - - if current_platform.is_cpu(): - logger.info( - "CPU platform detected: overriding Mamba SSU backend " - "from 'triton' to 'cpu'." + global _mamba_ssu_backend + + backend = mamba_config.backend + + # On CPU-only platforms (PowerPC, x86 without CUDA) Triton JIT is + # unstable or unavailable. Silently fall back to the CPU + # backend unless the user explicitly chose something other than "triton". + if backend == MambaBackendEnum.TRITON: + from vllm.platforms import current_platform + + if current_platform.is_cpu(): + logger.info( + "CPU platform detected: overriding Mamba SSU backend " + "from 'triton' to 'cpu'." + ) + backend = MambaBackendEnum.CPU + + if backend not in _BACKEND_REGISTRY: + raise ValueError( + f"Unknown Mamba SSU backend: {backend}. " + f"Valid options: {list(_BACKEND_REGISTRY.keys())}" ) - backend = MambaBackendEnum.CPU - if backend not in _BACKEND_REGISTRY: - raise ValueError( - f"Unknown Mamba SSU backend: {backend}. " - f"Valid options: {list(_BACKEND_REGISTRY.keys())}" - ) - - backend_cls = _BACKEND_REGISTRY[backend] - if isinstance(_mamba_ssu_backend, backend_cls): - return + backend_cls = _BACKEND_REGISTRY[backend] + if not isinstance(_mamba_ssu_backend, backend_cls): + _mamba_ssu_backend = backend_cls(mamba_config) + logger.info("Using %s Mamba SSU backend.", _mamba_ssu_backend.name) - _mamba_ssu_backend = backend_cls(mamba_config) - logger.info("Using %s Mamba SSU backend.", _mamba_ssu_backend.name) + initialize_replayssm_backend(mamba_config, use_replayssm=use_replayssm) def get_mamba_ssu_backend() -> MambaSSUBackend: diff --git a/vllm/v1/worker/gpu/model_runner.py b/vllm/v1/worker/gpu/model_runner.py index 5ba4b97f3c5d..b25cdb4d033e 100644 --- a/vllm/v1/worker/gpu/model_runner.py +++ b/vllm/v1/worker/gpu/model_runner.py @@ -615,7 +615,9 @@ def initialize_kv_cache( cls=self.pcp_manager_cls, ) initialize_mamba_ssu_backend( - self.vllm_config.mamba_config, self.kv_cache_config + self.vllm_config.mamba_config, + self.kv_cache_config, + use_replayssm=self.vllm_config.cache_config.use_replayssm, ) if self.adaptive_verification is not None: self.compilation_config.cudagraph_mode = CUDAGraphMode.FULL_AND_PIECEWISE diff --git a/vllm/v1/worker/gpu_model_runner.py b/vllm/v1/worker/gpu_model_runner.py index a1cab24fbef5..78933bbc6920 100644 --- a/vllm/v1/worker/gpu_model_runner.py +++ b/vllm/v1/worker/gpu_model_runner.py @@ -7477,7 +7477,9 @@ def initialize_kv_cache( self.maybe_add_kv_sharing_layers_to_kv_cache_groups(kv_cache_config) self.initialize_attn_backend(kv_cache_config, is_profiling=is_profiling) initialize_mamba_ssu_backend( - self.vllm_config.mamba_config, self.kv_cache_config + self.vllm_config.mamba_config, + self.kv_cache_config, + use_replayssm=self.vllm_config.cache_config.use_replayssm, ) # The kernel block size for all KV cache groups. For example, if # kv_cache_manager uses block_size 256 for a given group, but the attention From a7e36e959fb2c687b98f3b1b7a54404dabb524a0 Mon Sep 17 00:00:00 2001 From: Andrii Skliar Date: Fri, 14 Aug 2026 15:27:44 +0200 Subject: [PATCH 02/33] Refactor ReplaySSM backend integration to support FlashInfer and Triton Signed-off-by: Andrii Skliar --- tests/kernels/mamba/test_ssu_dispatch.py | 92 ++++--- .../layers/mamba/mamba_mixer2.py | 70 ++++-- .../layers/mamba/ops/ssu_dispatch.py | 231 +++++++++--------- vllm/v1/attention/backends/mamba_attn.py | 35 ++- 4 files changed, 241 insertions(+), 187 deletions(-) diff --git a/tests/kernels/mamba/test_ssu_dispatch.py b/tests/kernels/mamba/test_ssu_dispatch.py index 082784d3126a..c748c6424309 100644 --- a/tests/kernels/mamba/test_ssu_dispatch.py +++ b/tests/kernels/mamba/test_ssu_dispatch.py @@ -17,8 +17,8 @@ initialize_mamba_ssu_backend, initialize_replayssm_backend, selective_state_update, - selective_state_update_replayssm, - translate_vllm_replayssm_bookkeeping, + selective_state_update_replayssm_flashinfer, + selective_state_update_replayssm_triton, ) from vllm.utils.torch_utils import set_random_seed from vllm.v1.attention.backends.registry import MambaAttentionBackendEnum @@ -256,14 +256,24 @@ def test_replayssm_uninitialized_backend_raises(): mod._replayssm_backend = old -def test_replayssm_bookkeeping_adapter_not_implemented(): - with pytest.raises(NotImplementedError, match="bookkeeping adapter"): - translate_vllm_replayssm_bookkeeping( - write_pos=torch.zeros(1, dtype=torch.int32), - is_flush=torch.zeros(1, dtype=torch.int8), - state_batch_indices=None, - max_cache_len=16, - batch=1, +@pytest.mark.skipif( + not HAS_FLASHINFER_CHECKPOINTING_SSU, + reason="flashinfer.mamba.checkpointing_ssu not available", +) +def test_replayssm_triton_entry_rejects_flashinfer_backend(): + initialize_replayssm_backend( + MambaConfig(backend=MambaBackendEnum.FLASHINFER), use_replayssm=True + ) + tensor = torch.empty(1) + with pytest.raises(RuntimeError, match="Triton ReplaySSM"): + selective_state_update_replayssm_triton( + tensor, + tensor, + tensor, + tensor, + tensor, + tensor, + out=tensor, ) @@ -271,13 +281,13 @@ def test_replayssm_bookkeeping_adapter_not_implemented(): not HAS_FLASHINFER_CHECKPOINTING_SSU, reason="flashinfer.mamba.checkpointing_ssu not available", ) -def test_replayssm_flashinfer_call_hits_bookkeeping_gap(monkeypatch): +def test_replayssm_flashinfer_call(monkeypatch): import flashinfer.mamba - kernel = Mock() + kernel = Mock(return_value=torch.empty(1, 1, 2, 4)) monkeypatch.setattr(flashinfer.mamba, "checkpointing_ssu", kernel) - backend = FlashInferReplaySSMBackend( - MambaConfig(backend=MambaBackendEnum.FLASHINFER) + initialize_replayssm_backend( + MambaConfig(backend=MambaBackendEnum.FLASHINFER), use_replayssm=True ) batch, nheads, dim, dstate, ngroups, L = 1, 2, 4, 8, 1, 16 @@ -293,28 +303,33 @@ def test_replayssm_flashinfer_call_hits_bookkeeping_gap(monkeypatch): x_cache = torch.empty(1, nheads, L, dim) dt_cache = torch.empty(1, nheads, L) B_cache = torch.empty(1, ngroups, L, dstate) - write_pos = torch.zeros(batch, dtype=torch.int32) - is_flush = torch.zeros(batch, dtype=torch.int8) - - with pytest.raises(NotImplementedError, match="bookkeeping adapter"): - backend( - state, - x, - dt, - A, - B, - C, - D=D, - dt_bias=dt_bias, - x_cache=x_cache, - dt_cache=dt_cache, - B_cache=B_cache, - write_pos=write_pos, - is_flush=is_flush, - max_cache_len=L, - out=out, - ) - kernel.assert_not_called() + ring_start = torch.zeros(1, dtype=torch.int32) + prev_num_accepted = torch.zeros(1, dtype=torch.int32) + + selective_state_update_replayssm_flashinfer( + state, + x, + dt, + A, + B, + C, + out, + x_cache, + B_cache, + dt_cache, + ring_start, + prev_num_accepted, + D=D, + dt_bias=dt_bias, + dt_softplus=True, + ) + assert kernel.call_count == 1 + kwargs = kernel.call_args.kwargs + assert kwargs["dt_softplus"] is True + # ring_start / prev_num_accepted are positional after the caches. + args = kernel.call_args.args + assert args[4] is ring_start + assert args[5] is prev_num_accepted @pytest.mark.skipif( @@ -331,10 +346,11 @@ def test_replayssm_dispatch_fn_uses_initialized_backend(): called = Mock(return_value=torch.empty(1)) old = mod._replayssm_backend - mod._replayssm_backend = called + mod._replayssm_backend = TritonReplaySSMBackend(MambaConfig()) + mod._replayssm_backend._kernel = called try: tensor = torch.empty(1) - selective_state_update_replayssm( + selective_state_update_replayssm_triton( tensor, tensor, tensor, diff --git a/vllm/model_executor/layers/mamba/mamba_mixer2.py b/vllm/model_executor/layers/mamba/mamba_mixer2.py index 7adcbc0f2d2e..2d8b6d5fcec2 100644 --- a/vllm/model_executor/layers/mamba/mamba_mixer2.py +++ b/vllm/model_executor/layers/mamba/mamba_mixer2.py @@ -36,8 +36,10 @@ mamba_chunk_scan_combined_varlen, ) from vllm.model_executor.layers.mamba.ops.ssu_dispatch import ( + get_replayssm_backend, selective_state_update, - selective_state_update_replayssm, + selective_state_update_replayssm_flashinfer, + selective_state_update_replayssm_triton, ) from vllm.model_executor.layers.quantization import QuantizationConfig from vllm.model_executor.model_loader.weight_utils import ( @@ -1060,26 +1062,52 @@ def conv_ssm_forward( ) if self.use_replayssm: assert self.replayssm_buffer_len is not None - selective_state_update_replayssm( - ssm_state, - hidden_states_d, - dt_d, - A_d, - B_d, - C_d, - D_d, - dt_bias, - dt_softplus=True, - x_cache=x_cache, - dt_cache=dt_cache, - B_cache=B_cache, - bc_pre=attn_metadata.bc_pre_scratch, - write_pos=attn_metadata.write_pos_d, - is_flush=attn_metadata.is_flush_d, - max_cache_len=self.replayssm_buffer_len, - state_batch_indices=state_indices_tensor_d_input, - out=preallocated_ssm_out_d, - ) + replayssm_backend = get_replayssm_backend() + if replayssm_backend.name == "flashinfer": + assert attn_metadata.ring_start_d is not None + assert attn_metadata.prev_num_accepted_d is not None + selective_state_update_replayssm_flashinfer( + ssm_state, + hidden_states_d, + dt_d, + A_d, + B_d, + C_d, + preallocated_ssm_out_d, + x_cache, + B_cache, + dt_cache, + attn_metadata.ring_start_d, + attn_metadata.prev_num_accepted_d, + D=D_d, + dt_bias=dt_bias, + dt_softplus=True, + state_batch_indices=state_indices_tensor_d_input, + cb_scaled=attn_metadata.fi_cb_scaled_scratch, + cumAdt_vec=attn_metadata.fi_cumAdt_vec_scratch, + cb_old=attn_metadata.fi_cb_old_scratch, + ) + else: + selective_state_update_replayssm_triton( + ssm_state, + hidden_states_d, + dt_d, + A_d, + B_d, + C_d, + D_d, + dt_bias, + dt_softplus=True, + x_cache=x_cache, + dt_cache=dt_cache, + B_cache=B_cache, + bc_pre=attn_metadata.bc_pre_scratch, + write_pos=attn_metadata.write_pos_d, + is_flush=attn_metadata.is_flush_d, + max_cache_len=self.replayssm_buffer_len, + state_batch_indices=state_indices_tensor_d_input, + out=preallocated_ssm_out_d, + ) else: selective_state_update( ssm_state, diff --git a/vllm/model_executor/layers/mamba/ops/ssu_dispatch.py b/vllm/model_executor/layers/mamba/ops/ssu_dispatch.py index 1083104ea594..d85d5711a1d7 100644 --- a/vllm/model_executor/layers/mamba/ops/ssu_dispatch.py +++ b/vllm/model_executor/layers/mamba/ops/ssu_dispatch.py @@ -3,17 +3,17 @@ """ Dispatch module for Mamba selective state update (SSU) backends. -Provides unified ``selective_state_update`` (baseline decode) and -``selective_state_update_replayssm`` (cached-input ReplaySSM decode) that -dispatch to Triton / FlashInfer / CPU based on ``MambaBackendEnum``. On -CPU-only platforms (PowerPC, x86 without CUDA) the baseline SSU backend -defaults to ``cpu``. - -The FlashInfer ReplaySSM path imports ``flashinfer.mamba.checkpointing_ssu`` -and reshapes T=1 AR tensors, but the vLLM ``write_pos`` / ``is_flush`` / -``bc_pre`` → FlashInfer ``ring_start`` / ``prev_num_accepted_tokens`` / -scratch contract is intentionally unfinished (see -``translate_vllm_replayssm_bookkeeping``). +Provides unified ``selective_state_update`` (baseline decode) and ReplaySSM +decode entry points that dispatch to Triton / FlashInfer / CPU based on +``MambaBackendEnum``. On CPU-only platforms (PowerPC, x86 without CUDA) the +baseline SSU backend defaults to ``cpu``. + +ReplaySSM backends: + - Triton: ``write_pos`` / ``is_flush`` / ``bc_pre`` + (``selective_state_update_replayssm_triton``) + - FlashInfer: ``ring_start`` / ``prev_num_accepted_tokens`` (+ optional + two-kernel scratch), matching ``flashinfer.mamba.checkpointing_ssu`` + (``selective_state_update_replayssm_flashinfer``) """ from abc import ABC, abstractmethod @@ -270,7 +270,7 @@ def __call__( class ReplaySSMBackend(ABC): - """Abstract base class for ReplaySSM decode backends.""" + """Marker base for ReplaySSM decode backends.""" def __init__(self, mamba_config: MambaConfig): self._mamba_config = mamba_config @@ -279,34 +279,9 @@ def __init__(self, mamba_config: MambaConfig): @abstractmethod def name(self) -> str: ... - @abstractmethod - def __call__( - self, - state: torch.Tensor, - x: torch.Tensor, - dt: torch.Tensor, - A: torch.Tensor, - B: torch.Tensor, - C: torch.Tensor, - D: torch.Tensor | None = None, - dt_bias: torch.Tensor | None = None, - z: torch.Tensor | None = None, - dt_softplus: bool = False, - x_cache: torch.Tensor | None = None, - dt_cache: torch.Tensor | None = None, - B_cache: torch.Tensor | None = None, - bc_pre: torch.Tensor | None = None, - write_pos: torch.Tensor | None = None, - is_flush: torch.Tensor | None = None, - max_cache_len: int = 16, - state_batch_indices: torch.Tensor | None = None, - null_block_id: int = NULL_BLOCK_ID, - out: torch.Tensor | None = None, - ) -> torch.Tensor: ... - class TritonReplaySSMBackend(ReplaySSMBackend): - """vLLM's in-tree Triton ReplaySSM output_only kernel.""" + """vLLM Triton ReplaySSM (``write_pos`` / ``is_flush`` / ``bc_pre``).""" def __init__(self, mamba_config: MambaConfig): super().__init__(mamba_config) @@ -369,41 +344,8 @@ def __call__( ) -def translate_vllm_replayssm_bookkeeping( - *, - write_pos: torch.Tensor, - is_flush: torch.Tensor, - state_batch_indices: torch.Tensor | None, - max_cache_len: int, - batch: int, -) -> tuple[torch.Tensor, torch.Tensor]: - """Map vLLM ReplaySSM host metadata to FlashInfer checkpointing_ssu args. - - vLLM (Triton ReplaySSM) provides per-decode-row: - - ``write_pos``: ring append cursor in ``[0, max_cache_len)`` - - ``is_flush``: materialize full SSM state this step - - FlashInfer ``checkpointing_ssu`` expects per-cache-slot: - - ``ring_start``: oldest live ring index - - ``prev_num_accepted_tokens``: live history length to replay - - Returns: - ``(ring_start, prev_num_accepted_tokens)`` shaped for the FlashInfer - call (typically indexed by cache slot, not batch row). - - Raises: - NotImplementedError: contract adapter not wired yet. - """ - raise NotImplementedError( - "FlashInfer ReplaySSM bookkeeping adapter is not implemented yet. " - "Map vLLM write_pos/is_flush (and state_batch_indices) to FlashInfer " - "ring_start/prev_num_accepted_tokens for max_cache_len=" - f"{max_cache_len}, batch={batch}." - ) - - class FlashInferReplaySSMBackend(ReplaySSMBackend): - """FlashInfer ``checkpointing_ssu`` ReplaySSM backend (contract pending).""" + """FlashInfer ``checkpointing_ssu`` ReplaySSM backend.""" def __init__(self, mamba_config: MambaConfig): super().__init__(mamba_config) @@ -429,53 +371,32 @@ def __call__( A: torch.Tensor, B: torch.Tensor, C: torch.Tensor, + out: torch.Tensor, + x_cache: torch.Tensor, + B_cache: torch.Tensor, + dt_cache: torch.Tensor, + ring_start: torch.Tensor, + prev_num_accepted_tokens: torch.Tensor, D: torch.Tensor | None = None, dt_bias: torch.Tensor | None = None, z: torch.Tensor | None = None, dt_softplus: bool = False, - x_cache: torch.Tensor | None = None, - dt_cache: torch.Tensor | None = None, - B_cache: torch.Tensor | None = None, - bc_pre: torch.Tensor | None = None, - write_pos: torch.Tensor | None = None, - is_flush: torch.Tensor | None = None, - max_cache_len: int = 16, state_batch_indices: torch.Tensor | None = None, null_block_id: int = NULL_BLOCK_ID, - out: torch.Tensor | None = None, + cb_scaled: torch.Tensor | None = None, + cumAdt_vec: torch.Tensor | None = None, + cb_old: torch.Tensor | None = None, + algorithm: str = "auto", ) -> torch.Tensor: - del bc_pre # Triton-only scratch; FI uses its own precompute buffers. - if write_pos is None or is_flush is None: - raise ValueError( - "FlashInfer ReplaySSM requires write_pos and is_flush metadata" - ) - if out is None: - raise ValueError("FlashInfer ReplaySSM requires a preallocated out tensor") - if x_cache is None or dt_cache is None or B_cache is None: - raise ValueError( - "FlashInfer ReplaySSM requires x_cache, dt_cache, and B_cache" - ) - - # Mechanical T=1 reshape for AR decode. MTP (T>1) is out of scope until - # ReplaySSM speculative decode is enabled. + # AR decode currently passes (batch, nheads, dim); checkpointing_ssu + # expects a predicted-token axis T. Unsqueeze T=1 here. if x.dim() == 3: - x_t = x.unsqueeze(1) - dt_t = dt.unsqueeze(1) - B_t = B.unsqueeze(1) - C_t = C.unsqueeze(1) - out_t = out.unsqueeze(1) - z_t = z.unsqueeze(1) if z is not None else None - else: - x_t, dt_t, B_t, C_t, out_t, z_t = x, dt, B, C, out, z - - batch = x_t.shape[0] - ring_start, prev_num_accepted = translate_vllm_replayssm_bookkeeping( - write_pos=write_pos, - is_flush=is_flush, - state_batch_indices=state_batch_indices, - max_cache_len=max_cache_len, - batch=batch, - ) + x = x.unsqueeze(1) + dt = dt.unsqueeze(1) + B = B.unsqueeze(1) + C = C.unsqueeze(1) + out = out.unsqueeze(1) + z = z.unsqueeze(1) if z is not None else None rand_seed = ( torch.randint(0, 2**32, (1,), device=state.device, dtype=torch.int64) @@ -492,21 +413,25 @@ def __call__( B_cache, dt_cache, ring_start, - prev_num_accepted, - x_t, - dt_t, + prev_num_accepted_tokens, + x, + dt, A, - B_t, - C_t, - out_t, + B, + C, + out, D=D, - z=z_t, + z=z, dt_bias=dt_bias, dt_softplus=dt_softplus, state_batch_indices=indices, pad_slot_id=null_block_id, rand_seed=rand_seed, philox_rounds=self._mamba_config.stochastic_rounding_philox_rounds or 10, + cb_scaled=cb_scaled, + cumAdt_vec=cumAdt_vec, + cb_old=cb_old, + algorithm=algorithm, ) @@ -554,7 +479,7 @@ def get_replayssm_backend() -> ReplaySSMBackend: return _replayssm_backend -def selective_state_update_replayssm( +def selective_state_update_replayssm_triton( state: torch.Tensor, x: torch.Tensor, dt: torch.Tensor, @@ -576,8 +501,15 @@ def selective_state_update_replayssm( null_block_id: int = NULL_BLOCK_ID, out: torch.Tensor | None = None, ) -> torch.Tensor: - """Unified dispatch for ReplaySSM selective state update.""" - return get_replayssm_backend()( + """Triton ReplaySSM decode (``write_pos`` / ``is_flush`` / ``bc_pre``).""" + backend = get_replayssm_backend() + if not isinstance(backend, TritonReplaySSMBackend): + raise RuntimeError( + "selective_state_update_replayssm_triton is the Triton ReplaySSM " + f"entry point; current backend is {backend.name!r}. Use " + "selective_state_update_replayssm_flashinfer for FlashInfer." + ) + return backend( state, x, dt, @@ -601,6 +533,63 @@ def selective_state_update_replayssm( ) +def selective_state_update_replayssm_flashinfer( + state: torch.Tensor, + x: torch.Tensor, + dt: torch.Tensor, + A: torch.Tensor, + B: torch.Tensor, + C: torch.Tensor, + out: torch.Tensor, + x_cache: torch.Tensor, + B_cache: torch.Tensor, + dt_cache: torch.Tensor, + ring_start: torch.Tensor, + prev_num_accepted_tokens: torch.Tensor, + D: torch.Tensor | None = None, + dt_bias: torch.Tensor | None = None, + z: torch.Tensor | None = None, + dt_softplus: bool = False, + state_batch_indices: torch.Tensor | None = None, + null_block_id: int = NULL_BLOCK_ID, + cb_scaled: torch.Tensor | None = None, + cumAdt_vec: torch.Tensor | None = None, + cb_old: torch.Tensor | None = None, + algorithm: str = "auto", +) -> torch.Tensor: + """FlashInfer ReplaySSM decode (``checkpointing_ssu``).""" + backend = get_replayssm_backend() + if not isinstance(backend, FlashInferReplaySSMBackend): + raise RuntimeError( + "selective_state_update_replayssm_flashinfer requires the " + f"flashinfer ReplaySSM backend; current backend is {backend.name!r}." + ) + return backend( + state, + x, + dt, + A, + B, + C, + out, + x_cache, + B_cache, + dt_cache, + ring_start, + prev_num_accepted_tokens, + D=D, + dt_bias=dt_bias, + z=z, + dt_softplus=dt_softplus, + state_batch_indices=state_batch_indices, + null_block_id=null_block_id, + cb_scaled=cb_scaled, + cumAdt_vec=cumAdt_vec, + cb_old=cb_old, + algorithm=algorithm, + ) + + def initialize_mamba_ssu_backend( mamba_config: MambaConfig, kv_cache_config: KVCacheConfig, diff --git a/vllm/v1/attention/backends/mamba_attn.py b/vllm/v1/attention/backends/mamba_attn.py index 6b7ac71ae8c3..f7da100cfd8b 100644 --- a/vllm/v1/attention/backends/mamba_attn.py +++ b/vllm/v1/attention/backends/mamba_attn.py @@ -8,6 +8,7 @@ import torch from vllm.config import VllmConfig +from vllm.config.mamba import MambaBackendEnum from vllm.utils.math_utils import cdiv from vllm.utils.torch_utils import async_tensor_h2d from vllm.v1.attention.backend import ( @@ -74,12 +75,20 @@ class BaseMambaAttentionMetadata: nums_dict: dict | None = None batch_ptr: torch.Tensor | None = None token_chunk_offset_ptr: torch.Tensor | None = None - # ReplaySSM standard decode: per-row ring cursor and flush flag, plus the - # per-step (decode_rows, ngroups, replayssm_buffer_len) fp32 scratch for the - # precomputed k^T q products. All None when use_replayssm is disabled. + # ReplaySSM standard decode — Triton: per-row ring cursor and flush flag, + # plus (decode_rows, ngroups, replayssm_buffer_len) fp32 scratch for + # precomputed B·C products. All None when use_replayssm is disabled or + # the FlashInfer ReplaySSM backend is selected. write_pos_d: torch.Tensor | None = None is_flush_d: torch.Tensor | None = None bc_pre_scratch: torch.Tensor | None = None + # ReplaySSM — FlashInfer checkpointing_ssu bookkeeping / scratch. + # All None unless that backend is on. + ring_start_d: torch.Tensor | None = None + prev_num_accepted_d: torch.Tensor | None = None + fi_cb_scaled_scratch: torch.Tensor | None = None + fi_cumAdt_vec_scratch: torch.Tensor | None = None + fi_cb_old_scratch: torch.Tensor | None = None class BaseMambaAttentionMetadataBuilder(AttentionMetadataBuilder[M], abc.ABC): @@ -107,6 +116,10 @@ def __init__( self.use_spec_decode = self.num_spec_tokens > 0 self.use_replayssm = vllm_config.cache_config.use_replayssm self.replayssm_buffer_len = vllm_config.cache_config.replayssm_buffer_len + self.use_flashinfer_replayssm = ( + self.use_replayssm + and vllm_config.mamba_config.backend == MambaBackendEnum.FLASHINFER + ) scheduler_config = vllm_config.scheduler_config self.decode_cudagraph_max_bs: int = scheduler_config.max_num_seqs @@ -166,9 +179,10 @@ def __init__( dtype=torch.int32, device=device, ) - # ReplaySSM standard-decode CUDA-graph buffers: per-row ring cursor, - # flush flag, and the k^T q precompute scratch. - if self.use_replayssm: + # ReplaySSM CUDA-graph buffers (Triton write_pos/is_flush/bc_pre). + # FlashInfer ring_start/pnat buffers are allocated when that path is + # wired. + if self.use_replayssm and not self.use_flashinfer_replayssm: self.decode_write_pos_d: torch.Tensor = torch.empty( (self.decode_cudagraph_max_bs,), dtype=torch.int32, @@ -587,6 +601,13 @@ def _compute_common_metadata( ] if self.use_replayssm and num_decodes > 0: + if self.use_flashinfer_replayssm: + raise NotImplementedError( + "FlashInfer ReplaySSM metadata is not implemented yet. " + "Build ring_start / prev_num_accepted_tokens (and optional " + "checkpointing_ssu two-kernel scratch) for " + "flashinfer.mamba.checkpointing_ssu." + ) decode_base_cpu = common_attn_metadata.replayssm_decode_base_cpu num_computed_tokens_cpu = common_attn_metadata._num_computed_tokens_cpu if decode_base_cpu is None or num_computed_tokens_cpu is None: @@ -769,7 +790,7 @@ def _update_metadata_for_cudagraph_capture( ) block_idx_last_scheduled_token_prev_step[metadata.num_decodes :] = 0 - if self.use_replayssm: + if self.use_replayssm and not self.use_flashinfer_replayssm: assert write_pos_d is not None assert is_flush_d is not None self.decode_write_pos_d[: metadata.num_decodes].copy_( From becbedd64fe138cd1ef3430a18e09d76a444bc9d Mon Sep 17 00:00:00 2001 From: Andrii Skliar Date: Fri, 14 Aug 2026 16:00:26 +0200 Subject: [PATCH 03/33] improve replayssm performance by allowing for split kernel runs Signed-off-by: Andrii Skliar --- .../test_replayssm_metadata_builder.py | 44 +++++++++- tests/v1/e2e/test_replayssm_decode.py | 73 ++++++++++++++++- .../layers/mamba/mamba_mixer2.py | 10 ++- vllm/v1/attention/backends/mamba_attn.py | 81 ++++++++++++++++--- 4 files changed, 190 insertions(+), 18 deletions(-) diff --git a/tests/v1/attention/test_replayssm_metadata_builder.py b/tests/v1/attention/test_replayssm_metadata_builder.py index cbbfedee964a..667a26623286 100644 --- a/tests/v1/attention/test_replayssm_metadata_builder.py +++ b/tests/v1/attention/test_replayssm_metadata_builder.py @@ -16,11 +16,16 @@ create_common_attn_metadata, create_vllm_config, ) +from vllm.config.mamba import MambaBackendEnum from vllm.v1.kv_cache_interface import MambaSpec BLOCK_SIZE = 16 DEVICE = torch.device("cpu") +# Flip when FlashInfer ReplaySSM metadata (ring_start / prev_num_accepted) is +# implemented in BaseMambaAttentionMetadataBuilder. +FLASHINFER_REPLAYSSM_METADATA_READY = False + @dataclass class ReplaySSMBuildCase: @@ -203,16 +208,20 @@ def _make_mamba_spec(buffer_len: int) -> MambaSpec: 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 ) @@ -254,3 +263,34 @@ 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_metadata_pending(): + """FlashInfer path must not silently reuse Triton write_pos metadata.""" + builder = _create_replayssm_builder( + 16, mamba_backend=MambaBackendEnum.FLASHINFER + ) + case = REPLAYSSM_BUILD_CASES["fresh_decode"] + with pytest.raises(NotImplementedError, match="FlashInfer ReplaySSM metadata"): + _build(builder, case) + + +@pytest.mark.skipif( + not FLASHINFER_REPLAYSSM_METADATA_READY, + reason="FlashInfer ReplaySSM metadata not implemented yet", +) +def test_flashinfer_replayssm_ring_metadata_fresh_decode(): + """Fresh decode: ring_start / prev_num_accepted for checkpointing_ssu.""" + builder = _create_replayssm_builder( + 16, mamba_backend=MambaBackendEnum.FLASHINFER + ) + meta = _build(builder, REPLAYSSM_BUILD_CASES["fresh_decode"]) + + assert meta.ring_start_d is not None + assert meta.prev_num_accepted_d is not None + assert meta.write_pos_d is None + assert meta.is_flush_d is None + assert meta.bc_pre_scratch is None + # Fill in expected ring_start / prev_num_accepted once the FI schedule is + # defined; until then this test stays skipped via the flag above. + raise NotImplementedError("set expected ring_start / prev_num_accepted") diff --git a/tests/v1/e2e/test_replayssm_decode.py b/tests/v1/e2e/test_replayssm_decode.py index 4fc6768e7446..d044223fccec 100644 --- a/tests/v1/e2e/test_replayssm_decode.py +++ b/tests/v1/e2e/test_replayssm_decode.py @@ -9,6 +9,16 @@ from ...models.utils import check_logprobs_close from ...utils import large_gpu_mark, multi_gpu_test +try: + from flashinfer.mamba import checkpointing_ssu # noqa: F401 + + HAS_FLASHINFER_CHECKPOINTING_SSU = True +except ImportError: + HAS_FLASHINFER_CHECKPOINTING_SSU = False + +# Flip when FlashInfer ReplaySSM metadata is implemented in mamba_attn. +FLASHINFER_REPLAYSSM_METADATA_READY = False + # Mamba2 (Nemotron-3) hybrid. MAMBA2_MODEL = "nvidia/NVIDIA-Nemotron-3-Nano-4B-BF16" MODELS = [ @@ -21,7 +31,14 @@ ] -def _check_replayssm_parity(vllm_runner, model_name, *, tensor_parallel_size=1): +def _check_replayssm_parity( + vllm_runner, + model_name, + *, + tensor_parallel_size=1, + mamba_backend: str = "triton", + name_1: str = "replayssm", +): # Compare logprobs, not greedy ids: ReplaySSM's fp arithmetic can flip a # near-tie. Baseline and ReplaySSM run at the same TP, so TP numerics are # common-mode and only ReplaySSM varies. @@ -31,6 +48,7 @@ def _check_replayssm_parity(vllm_runner, model_name, *, tensor_parallel_size=1): enable_prefix_caching=False, mamba_cache_mode="none", tensor_parallel_size=tensor_parallel_size, + mamba_backend=mamba_backend, ) with vllm_runner(model_name, **common) as llm: baseline = llm.generate_greedy_logprobs(PROMPTS, max_tokens=32, num_logprobs=5) @@ -43,7 +61,7 @@ def _check_replayssm_parity(vllm_runner, model_name, *, tensor_parallel_size=1): outputs_0_lst=baseline, outputs_1_lst=replay, name_0="baseline", - name_1="replayssm", + name_1=name_1, ) @@ -60,6 +78,57 @@ def test_replayssm_decode_matches_baseline_tp2(vllm_runner, model_name): _check_replayssm_parity(vllm_runner, model_name, tensor_parallel_size=2) +@pytest.mark.skipif( + not HAS_FLASHINFER_CHECKPOINTING_SSU, + reason="flashinfer.mamba.checkpointing_ssu not available", +) +@pytest.mark.skipif( + not FLASHINFER_REPLAYSSM_METADATA_READY, + reason="FlashInfer ReplaySSM metadata not implemented yet", +) +@pytest.mark.parametrize("model_name", MODELS) +def test_replayssm_flashinfer_decode_matches_baseline(vllm_runner, model_name): + _check_replayssm_parity( + vllm_runner, + model_name, + mamba_backend="flashinfer", + name_1="replayssm_flashinfer", + ) + + +@pytest.mark.skipif( + not HAS_FLASHINFER_CHECKPOINTING_SSU, + reason="flashinfer.mamba.checkpointing_ssu not available", +) +@pytest.mark.skipif( + not FLASHINFER_REPLAYSSM_METADATA_READY, + reason="FlashInfer ReplaySSM metadata not implemented yet", +) +@pytest.mark.parametrize("model_name", MODELS) +def test_replayssm_flashinfer_matches_triton_replayssm(vllm_runner, model_name): + common = dict( + max_model_len=1024, + trust_remote_code=True, + enable_prefix_caching=False, + mamba_cache_mode="none", + use_replayssm=True, + replayssm_buffer_len=16, + ) + with vllm_runner(model_name, mamba_backend="triton", **common) as llm: + triton = llm.generate_greedy_logprobs(PROMPTS, max_tokens=32, num_logprobs=5) + with vllm_runner(model_name, mamba_backend="flashinfer", **common) as llm: + flashinfer = llm.generate_greedy_logprobs( + PROMPTS, max_tokens=32, num_logprobs=5 + ) + + check_logprobs_close( + outputs_0_lst=triton, + outputs_1_lst=flashinfer, + name_0="replayssm_triton", + name_1="replayssm_flashinfer", + ) + + # Prefix spans several mamba blocks; prefix caching only reuses full blocks. _PC_SENTENCE = ( "In a detailed survey of state space models, the authors compared many " diff --git a/vllm/model_executor/layers/mamba/mamba_mixer2.py b/vllm/model_executor/layers/mamba/mamba_mixer2.py index 2d8b6d5fcec2..e9228d5f9cde 100644 --- a/vllm/model_executor/layers/mamba/mamba_mixer2.py +++ b/vllm/model_executor/layers/mamba/mamba_mixer2.py @@ -1066,6 +1066,9 @@ def conv_ssm_forward( if replayssm_backend.name == "flashinfer": assert attn_metadata.ring_start_d is not None assert attn_metadata.prev_num_accepted_d is not None + assert attn_metadata.cb_scaled is not None + assert attn_metadata.cumAdt_vec is not None + assert attn_metadata.cb_old is not None selective_state_update_replayssm_flashinfer( ssm_state, hidden_states_d, @@ -1083,9 +1086,10 @@ def conv_ssm_forward( dt_bias=dt_bias, dt_softplus=True, state_batch_indices=state_indices_tensor_d_input, - cb_scaled=attn_metadata.fi_cb_scaled_scratch, - cumAdt_vec=attn_metadata.fi_cumAdt_vec_scratch, - cb_old=attn_metadata.fi_cb_old_scratch, + cb_scaled=attn_metadata.cb_scaled, + cumAdt_vec=attn_metadata.cumAdt_vec, + cb_old=attn_metadata.cb_old, + algorithm="auto", ) else: selective_state_update_replayssm_triton( diff --git a/vllm/v1/attention/backends/mamba_attn.py b/vllm/v1/attention/backends/mamba_attn.py index f7da100cfd8b..77c836302bcd 100644 --- a/vllm/v1/attention/backends/mamba_attn.py +++ b/vllm/v1/attention/backends/mamba_attn.py @@ -82,13 +82,15 @@ class BaseMambaAttentionMetadata: write_pos_d: torch.Tensor | None = None is_flush_d: torch.Tensor | None = None bc_pre_scratch: torch.Tensor | None = None - # ReplaySSM — FlashInfer checkpointing_ssu bookkeeping / scratch. - # All None unless that backend is on. + # ReplaySSM — FlashInfer checkpointing_ssu: + # ring_start / prev_num_accepted (always), plus two-kernel scratch + # (cb_scaled / cumAdt_vec / cb_old) so monolith and two-kernel are both + # available via algorithm="auto". All None unless that backend is on. ring_start_d: torch.Tensor | None = None prev_num_accepted_d: torch.Tensor | None = None - fi_cb_scaled_scratch: torch.Tensor | None = None - fi_cumAdt_vec_scratch: torch.Tensor | None = None - fi_cb_old_scratch: torch.Tensor | None = None + cb_scaled: torch.Tensor | None = None + cumAdt_vec: torch.Tensor | None = None + cb_old: torch.Tensor | None = None class BaseMambaAttentionMetadataBuilder(AttentionMetadataBuilder[M], abc.ABC): @@ -179,9 +181,11 @@ def __init__( dtype=torch.int32, device=device, ) - # ReplaySSM CUDA-graph buffers (Triton write_pos/is_flush/bc_pre). - # FlashInfer ring_start/pnat buffers are allocated when that path is - # wired. + # ReplaySSM CUDA-graph buffers. + # Triton: write_pos / is_flush / bc_pre. + # FlashInfer: ring_start / prev_num_accepted + two-kernel scratch + # (cb_scaled / cumAdt_vec / cb_old) so algorithm="auto" can pick + # monolith or two-kernel. if self.use_replayssm and not self.use_flashinfer_replayssm: self.decode_write_pos_d: torch.Tensor = torch.empty( (self.decode_cudagraph_max_bs,), @@ -208,8 +212,64 @@ def __init__( dtype=torch.float32, device=device, ) + self.decode_ring_start_d = None + self.decode_prev_num_accepted_d = None + self.decode_cb_scaled = None + self.decode_cumAdt_vec = None + self.decode_cb_old = None + elif self.use_flashinfer_replayssm: + self.decode_bc_pre_scratch = None + scratch_bs = max( + self.decode_cudagraph_max_bs, scheduler_config.max_num_seqs + ) + # checkpointing_ssu ring_start / prev_num_accepted are per cache + # slot; for decode metadata we keep per-row views sized to the + # CG batch (indexed via state_batch_indices inside the kernel). + self.decode_ring_start_d = torch.empty( + (scratch_bs,), + dtype=torch.int32, + device=device, + ) + self.decode_prev_num_accepted_d = torch.empty( + (scratch_bs,), + dtype=torch.int32, + device=device, + ) + # x_cache page: (nheads, L, head_dim); AR decode uses T=1. + # Scratch shapes follow flashinfer's bench_checkpointing_ssu.py + # (WARP_SIZE / MMA_FRAG_SIZE are not exported from Python). + nheads = kv_cache_spec.shapes[2][0] + npredicted = 1 # AR decode + max_window = self.replayssm_buffer_len - 1 + k_old = (max_window + 7) // 8 * 8 + # cumAdt_vec: next_multiple_of_16(T) + t_pad = ((npredicted + 15) // 16) * 16 + warp_size = 32 + # cb_scaled: (..., 32, 8) = fragA for m16n8k16 + mma_frag_size = t_pad // 2 + act_dtype = vllm_config.model_config.dtype + self.decode_cb_scaled = torch.empty( + (scratch_bs, nheads, warp_size, mma_frag_size), + dtype=act_dtype, + device=device, + ) + self.decode_cumAdt_vec = torch.empty( + (scratch_bs, nheads, t_pad), + dtype=torch.float32, + device=device, + ) + self.decode_cb_old = torch.empty( + (scratch_bs, nheads, warp_size, k_old // 2), + dtype=act_dtype, + device=device, + ) else: self.decode_bc_pre_scratch = None + self.decode_ring_start_d = None + self.decode_prev_num_accepted_d = None + self.decode_cb_scaled = None + self.decode_cumAdt_vec = None + self.decode_cb_old = None self._init_reorder_batch_threshold(1, self.use_spec_decode) if self.use_spec_decode: @@ -604,9 +664,8 @@ def _compute_common_metadata( if self.use_flashinfer_replayssm: raise NotImplementedError( "FlashInfer ReplaySSM metadata is not implemented yet. " - "Build ring_start / prev_num_accepted_tokens (and optional " - "checkpointing_ssu two-kernel scratch) for " - "flashinfer.mamba.checkpointing_ssu." + "Fill ring_start_d / prev_num_accepted_d (CG buffers and " + "two-kernel scratch are already allocated)." ) decode_base_cpu = common_attn_metadata.replayssm_decode_base_cpu num_computed_tokens_cpu = common_attn_metadata._num_computed_tokens_cpu From b526ca5c462990c7b4700dc441f7aca0992b4f6c Mon Sep 17 00:00:00 2001 From: Andrii Skliar Date: Fri, 14 Aug 2026 16:22:52 +0200 Subject: [PATCH 04/33] add metadata building for FI ReplaySSM Signed-off-by: Andrii Skliar --- .../test_replayssm_metadata_builder.py | 36 +++----- vllm/v1/attention/backends/mamba_attn.py | 91 ++++++++++++++----- 2 files changed, 80 insertions(+), 47 deletions(-) diff --git a/tests/v1/attention/test_replayssm_metadata_builder.py b/tests/v1/attention/test_replayssm_metadata_builder.py index 667a26623286..368d6cfc3e5d 100644 --- a/tests/v1/attention/test_replayssm_metadata_builder.py +++ b/tests/v1/attention/test_replayssm_metadata_builder.py @@ -22,11 +22,6 @@ BLOCK_SIZE = 16 DEVICE = torch.device("cpu") -# Flip when FlashInfer ReplaySSM metadata (ring_start / prev_num_accepted) is -# implemented in BaseMambaAttentionMetadataBuilder. -FLASHINFER_REPLAYSSM_METADATA_READY = False - - @dataclass class ReplaySSMBuildCase: """A decode batch and its expected per-row write_pos / is_flush. @@ -265,32 +260,25 @@ def test_resumed_request_differs_from_fresh(): assert meta.is_flush_d.tolist()[:2] == [0, 0] -def test_flashinfer_replayssm_metadata_pending(): - """FlashInfer path must not silently reuse Triton write_pos metadata.""" - builder = _create_replayssm_builder( - 16, mamba_backend=MambaBackendEnum.FLASHINFER - ) - case = REPLAYSSM_BUILD_CASES["fresh_decode"] - with pytest.raises(NotImplementedError, match="FlashInfer ReplaySSM metadata"): - _build(builder, case) - - -@pytest.mark.skipif( - not FLASHINFER_REPLAYSSM_METADATA_READY, - reason="FlashInfer ReplaySSM metadata not implemented yet", -) def test_flashinfer_replayssm_ring_metadata_fresh_decode(): - """Fresh decode: ring_start / prev_num_accepted for checkpointing_ssu.""" + """FlashInfer receives per-slot ring metadata and per-row scratch.""" builder = _create_replayssm_builder( 16, mamba_backend=MambaBackendEnum.FLASHINFER ) - meta = _build(builder, REPLAYSSM_BUILD_CASES["fresh_decode"]) + case = REPLAYSSM_BUILD_CASES["fresh_decode"] + meta = _build(builder, case) assert meta.ring_start_d is not None assert meta.prev_num_accepted_d is not None assert meta.write_pos_d is None assert meta.is_flush_d is None assert meta.bc_pre_scratch is None - # Fill in expected ring_start / prev_num_accepted once the FI schedule is - # defined; until then this test stays skipped via the flag above. - raise NotImplementedError("set expected ring_start / prev_num_accepted") + state_slot = int(meta.state_indices_tensor_d[0, 0]) + assert torch.count_nonzero(meta.ring_start_d) == 0 + assert int(meta.prev_num_accepted_d[state_slot]) == case.expected_write_pos[0] + assert meta.cb_scaled is not None + assert meta.cb_scaled.shape == (1, 1, 32, 8) + assert meta.cumAdt_vec is not None + assert meta.cumAdt_vec.shape == (1, 1, 16) + assert meta.cb_old is not None + assert meta.cb_old.shape == (1, 1, 32, 8) diff --git a/vllm/v1/attention/backends/mamba_attn.py b/vllm/v1/attention/backends/mamba_attn.py index 77c836302bcd..0d626a3d057e 100644 --- a/vllm/v1/attention/backends/mamba_attn.py +++ b/vllm/v1/attention/backends/mamba_attn.py @@ -222,16 +222,19 @@ def __init__( scratch_bs = max( self.decode_cudagraph_max_bs, scheduler_config.max_num_seqs ) - # checkpointing_ssu ring_start / prev_num_accepted are per cache - # slot; for decode metadata we keep per-row views sized to the - # CG batch (indexed via state_batch_indices inside the kernel). - self.decode_ring_start_d = torch.empty( - (scratch_bs,), + # checkpointing_ssu indexes these tensors by state_batch_indices, + # so they cover cache slots rather than decode rows. + num_cache_slots = vllm_config.cache_config.num_gpu_blocks + if num_cache_slots is None: + # Unit-test builders run before cache sizing. + num_cache_slots = scratch_bs + self.decode_ring_start_d = torch.zeros( + (num_cache_slots,), dtype=torch.int32, device=device, ) - self.decode_prev_num_accepted_d = torch.empty( - (scratch_bs,), + self.decode_prev_num_accepted_d = torch.zeros( + (num_cache_slots,), dtype=torch.int32, device=device, ) @@ -574,6 +577,11 @@ def _compute_common_metadata( nums_dict, batch_ptr, token_chunk_offset_ptr = None, None, None write_pos_d = None is_flush_d = None + ring_start_d = None + prev_num_accepted_d = None + cb_scaled = None + cumAdt_vec = None + cb_old = None if self.vllm_config.cache_config.mamba_cache_mode == "all": num_computed_tokens = common_attn_metadata.compute_num_computed_tokens() @@ -661,12 +669,6 @@ def _compute_common_metadata( ] if self.use_replayssm and num_decodes > 0: - if self.use_flashinfer_replayssm: - raise NotImplementedError( - "FlashInfer ReplaySSM metadata is not implemented yet. " - "Fill ring_start_d / prev_num_accepted_d (CG buffers and " - "two-kernel scratch are already allocated)." - ) decode_base_cpu = common_attn_metadata.replayssm_decode_base_cpu num_computed_tokens_cpu = common_attn_metadata._num_computed_tokens_cpu if decode_base_cpu is None or num_computed_tokens_cpu is None: @@ -720,16 +722,35 @@ def _compute_common_metadata( & ((num_computed_d + query_lens_cpu) % block_size == 0) ) is_flush_cpu = is_flush_cpu.to(torch.int8) - write_pos_d = async_tensor_h2d( - write_pos_cpu.to(torch.int32).tolist(), - dtype=torch.int32, - device=common_attn_metadata.query_start_loc.device, - ) - is_flush_d = async_tensor_h2d( - is_flush_cpu.tolist(), - dtype=torch.int8, - device=common_attn_metadata.query_start_loc.device, - ) + if self.use_flashinfer_replayssm: + assert self.decode_ring_start_d is not None + assert self.decode_prev_num_accepted_d is not None + assert self.decode_cb_scaled is not None + assert self.decode_cumAdt_vec is not None + assert self.decode_cb_old is not None + ring_start_d = self.decode_ring_start_d + prev_num_accepted_d = self.decode_prev_num_accepted_d + state_slots = state_indices_tensor_d[:num_decodes, 0].long() + prev_num_accepted = async_tensor_h2d( + write_pos_cpu.to(torch.int32).tolist(), + dtype=torch.int32, + device=common_attn_metadata.query_start_loc.device, + ) + prev_num_accepted_d.index_copy_(0, state_slots, prev_num_accepted) + cb_scaled = self.decode_cb_scaled[:num_decodes] + cumAdt_vec = self.decode_cumAdt_vec[:num_decodes] + cb_old = self.decode_cb_old[:num_decodes] + else: + write_pos_d = async_tensor_h2d( + write_pos_cpu.to(torch.int32).tolist(), + dtype=torch.int32, + device=common_attn_metadata.query_start_loc.device, + ) + is_flush_d = async_tensor_h2d( + is_flush_cpu.tolist(), + dtype=torch.int8, + device=common_attn_metadata.query_start_loc.device, + ) bc_pre_scratch = None if ( @@ -751,6 +772,11 @@ def _compute_common_metadata( write_pos_d=write_pos_d, is_flush_d=is_flush_d, bc_pre_scratch=bc_pre_scratch, + ring_start_d=ring_start_d, + prev_num_accepted_d=prev_num_accepted_d, + cb_scaled=cb_scaled, + cumAdt_vec=cumAdt_vec, + cb_old=cb_old, num_accepted_tokens=num_accepted_tokens, query_start_loc_d=query_start_loc_d, block_idx_last_scheduled_token=block_idx_last_scheduled_token, @@ -788,6 +814,11 @@ def _update_metadata_for_cudagraph_capture( write_pos_d = metadata.write_pos_d is_flush_d = metadata.is_flush_d bc_pre_scratch = metadata.bc_pre_scratch + ring_start_d = metadata.ring_start_d + prev_num_accepted_d = metadata.prev_num_accepted_d + cb_scaled = metadata.cb_scaled + cumAdt_vec = metadata.cumAdt_vec + cb_old = metadata.cb_old if ( metadata.num_prefills == 0 and metadata.num_decodes <= self.decode_cudagraph_max_bs @@ -868,6 +899,15 @@ def _update_metadata_for_cudagraph_capture( if self.decode_bc_pre_scratch is not None: bc_pre_scratch = self.decode_bc_pre_scratch[:padded_bs] + elif self.use_flashinfer_replayssm: + assert ring_start_d is not None + assert prev_num_accepted_d is not None + assert self.decode_cb_scaled is not None + assert self.decode_cumAdt_vec is not None + assert self.decode_cb_old is not None + cb_scaled = self.decode_cb_scaled[:padded_bs] + cumAdt_vec = self.decode_cumAdt_vec[:padded_bs] + cb_old = self.decode_cb_old[:padded_bs] return replace( metadata, @@ -877,6 +917,11 @@ def _update_metadata_for_cudagraph_capture( write_pos_d=write_pos_d, is_flush_d=is_flush_d, bc_pre_scratch=bc_pre_scratch, + ring_start_d=ring_start_d, + prev_num_accepted_d=prev_num_accepted_d, + cb_scaled=cb_scaled, + cumAdt_vec=cumAdt_vec, + cb_old=cb_old, block_idx_last_scheduled_token=block_idx_last_scheduled_token, block_idx_last_computed_token=block_idx_last_computed_token, block_idx_last_scheduled_token_prev_step=( From 27bcd0bad82a5021ab530386cddef05ae1c4ffe9 Mon Sep 17 00:00:00 2001 From: Andrii Skliar Date: Fri, 14 Aug 2026 16:50:16 +0200 Subject: [PATCH 05/33] add proper ring flushing and advancement Signed-off-by: Andrii Skliar --- .../test_replayssm_metadata_builder.py | 32 +++++++++++--- vllm/config/cache.py | 5 ++- .../layers/mamba/mamba_mixer2.py | 6 ++- .../layers/mamba/mamba_utils.py | 16 +++---- vllm/model_executor/models/nemotron_h.py | 8 +++- vllm/v1/attention/backends/mamba_attn.py | 43 ++++++++++++++++++- 6 files changed, 91 insertions(+), 19 deletions(-) diff --git a/tests/v1/attention/test_replayssm_metadata_builder.py b/tests/v1/attention/test_replayssm_metadata_builder.py index 368d6cfc3e5d..3055e48e4111 100644 --- a/tests/v1/attention/test_replayssm_metadata_builder.py +++ b/tests/v1/attention/test_replayssm_metadata_builder.py @@ -17,6 +17,9 @@ create_vllm_config, ) from vllm.config.mamba import MambaBackendEnum +from vllm.v1.attention.backends.mamba_attn import ( + _derive_flashinfer_replayssm_ring_state, +) from vllm.v1.kv_cache_interface import MambaSpec BLOCK_SIZE = 16 @@ -187,16 +190,22 @@ class ReplaySSMBuildCase: } -def _make_mamba_spec(buffer_len: int) -> MambaSpec: +def _make_mamba_spec( + buffer_len: int, + mamba_backend: MambaBackendEnum, +) -> MambaSpec: # Five-tensor ReplaySSM page; the builder only reads shapes[4][0] (bc groups). + ring_buffer_len = buffer_len + ( + 1 if mamba_backend == MambaBackendEnum.FLASHINFER else 0 + ) return MambaSpec( block_size=BLOCK_SIZE, shapes=( (1, 1), (1, 1, 1), - (1, buffer_len, 1), - (1, buffer_len), - (1, buffer_len, 1), + (1, ring_buffer_len, 1), + (1, ring_buffer_len), + (1, ring_buffer_len, 1), ), dtypes=(torch.float32,), ) @@ -218,7 +227,10 @@ def _create_replayssm_builder( 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, ) @@ -282,3 +294,13 @@ def test_flashinfer_replayssm_ring_metadata_fresh_decode(): assert meta.cumAdt_vec.shape == (1, 1, 16) assert meta.cb_old is not None assert meta.cb_old.shape == (1, 1, 32, 8) + + +def test_flashinfer_replayssm_ring_lifecycle_across_flushes(): + ring_start, prev_num_accepted = _derive_flashinfer_replayssm_ring_state( + torch.tensor([0, 5, 16, 17, 32, 33], dtype=torch.int32), + logical_window=16, + ) + + assert ring_start.tolist() == [0, 0, 0, 16, 16, 15] + assert prev_num_accepted.tolist() == [0, 5, 16, 1, 16, 1] diff --git a/vllm/config/cache.py b/vllm/config/cache.py index 40741e5eaddf..7ff08d992c43 100644 --- a/vllm/config/cache.py +++ b/vllm/config/cache.py @@ -195,8 +195,9 @@ class CacheConfig: caching is enabled. """ replayssm_buffer_len: int = Field(default=16, gt=0) - """ReplaySSM history buffer length B for standard Mamba2 decode. Kimi-K3 - speculative decoding does not use B. Default 16.""" + """ReplaySSM logical history length B for Mamba2. Triton uses B physical + rows and FlashInfer uses B+1. Kimi-K3 speculative decode does not use B. + Default 16.""" use_replayssm: bool = False """Use the ReplaySSM Mamba2 decode kernel: cache recent SSM inputs and skip the per-step full-state store, writing the checkpoint back only on flush. diff --git a/vllm/model_executor/layers/mamba/mamba_mixer2.py b/vllm/model_executor/layers/mamba/mamba_mixer2.py index e9228d5f9cde..debdb3c3cb0f 100644 --- a/vllm/model_executor/layers/mamba/mamba_mixer2.py +++ b/vllm/model_executor/layers/mamba/mamba_mixer2.py @@ -6,6 +6,7 @@ from torch import nn from vllm.config import CacheConfig, ModelConfig, get_current_vllm_config +from vllm.config.mamba import MambaBackendEnum from vllm.distributed import ( divide, get_tensor_model_parallel_rank, @@ -1159,11 +1160,14 @@ def get_state_shape(self) -> tuple[tuple[int, ...], ...]: ) if self.use_replayssm: assert self.replayssm_buffer_len is not None + ring_buffer_len = self.replayssm_buffer_len + if self.mamba_config.backend == MambaBackendEnum.FLASHINFER: + ring_buffer_len += 1 return MambaStateShapeCalculator.append_replayssm_ring( base_shape, self.n_groups, tp_world_size, - self.replayssm_buffer_len, + ring_buffer_len, ) return base_shape diff --git a/vllm/model_executor/layers/mamba/mamba_utils.py b/vllm/model_executor/layers/mamba/mamba_utils.py index 71872de47c0b..ab7d99d34334 100644 --- a/vllm/model_executor/layers/mamba/mamba_utils.py +++ b/vllm/model_executor/layers/mamba/mamba_utils.py @@ -213,20 +213,20 @@ def append_replayssm_ring( base_shapes: tuple[tuple[int, ...], ...], n_groups: int, tp_world_size: int, - replayssm_buffer_len: int, + ring_buffer_len: int, ) -> tuple[tuple[int, ...], ...]: - """Append the ReplaySSM ring shapes (x_cache, dt_cache, B_cache) to a - base ``(conv, ssm)`` tuple. ``base_shapes[1]`` is the ssm shape - ``(nheads // tp, head_dim, state_size)``; B_cache uses the un-extended - ``n_groups``. + """Append the physical ReplaySSM ring shapes to ``(conv, ssm)``. + + ``base_shapes[1]`` is ``(nheads // tp, head_dim, state_size)``; + B_cache uses the un-extended ``n_groups``. """ local_nheads, head_dim, state_size = base_shapes[1] local_ngroups = divide(n_groups, tp_world_size) return ( *base_shapes, - (local_nheads, replayssm_buffer_len, head_dim), - (local_nheads, replayssm_buffer_len), - (local_ngroups, replayssm_buffer_len, state_size), + (local_nheads, ring_buffer_len, head_dim), + (local_nheads, ring_buffer_len), + (local_ngroups, ring_buffer_len, state_size), ) @classmethod diff --git a/vllm/model_executor/models/nemotron_h.py b/vllm/model_executor/models/nemotron_h.py index b62e46d44af4..83e74bb6c54b 100644 --- a/vllm/model_executor/models/nemotron_h.py +++ b/vllm/model_executor/models/nemotron_h.py @@ -26,6 +26,7 @@ from vllm.compilation.decorators import support_torch_compile from vllm.config import CacheConfig, ModelConfig, VllmConfig +from vllm.config.mamba import MambaBackendEnum from vllm.config.parallel import ParallelConfig from vllm.distributed import get_ep_group, get_tensor_model_parallel_world_size from vllm.distributed.communication_op import tensor_model_parallel_all_gather @@ -783,11 +784,16 @@ def get_mamba_state_shape_from_config( num_spec=vllm_config.num_speculative_tokens, ) if cache_config.use_replayssm: + ring_buffer_len = cache_config.replayssm_buffer_len + if vllm_config.mamba_config.backend == MambaBackendEnum.FLASHINFER: + # FlashInfer's physical ring includes room for the new token + # while replaying the B live old tokens. + ring_buffer_len += 1 return MambaStateShapeCalculator.append_replayssm_ring( base_shape, hf_config.n_groups, parallel_config.tensor_parallel_size, - cache_config.replayssm_buffer_len, + ring_buffer_len, ) return base_shape diff --git a/vllm/v1/attention/backends/mamba_attn.py b/vllm/v1/attention/backends/mamba_attn.py index 0d626a3d057e..12e0e63e39dc 100644 --- a/vllm/v1/attention/backends/mamba_attn.py +++ b/vllm/v1/attention/backends/mamba_attn.py @@ -27,6 +27,29 @@ M = TypeVar("M", bound="BaseMambaAttentionMetadata") +def _derive_flashinfer_replayssm_ring_state( + decode_steps: torch.Tensor, + logical_window: int, +) -> tuple[torch.Tensor, torch.Tensor]: + """Return the ring cursor and live-token count before each T=1 call.""" + positive_steps = torch.clamp(decode_steps - 1, min=0) + flushes_before = torch.div( + positive_steps, + logical_window, + rounding_mode="floor", + ) + prev_num_accepted = torch.where( + decode_steps == 0, + torch.zeros_like(decode_steps), + torch.remainder(positive_steps, logical_window) + 1, + ) + ring_start = torch.remainder( + flushes_before * logical_window, + logical_window + 1, + ) + return ring_start, prev_num_accepted + + @dataclass class BaseMambaAttentionMetadata: num_prefills: int @@ -243,7 +266,7 @@ def __init__( # (WARP_SIZE / MMA_FRAG_SIZE are not exported from Python). nheads = kv_cache_spec.shapes[2][0] npredicted = 1 # AR decode - max_window = self.replayssm_buffer_len - 1 + max_window = self.replayssm_buffer_len k_old = (max_window + 7) // 8 * 8 # cumAdt_vec: next_multiple_of_16(T) t_pad = ((npredicted + 15) // 16) * 16 @@ -731,11 +754,27 @@ def _compute_common_metadata( ring_start_d = self.decode_ring_start_d prev_num_accepted_d = self.decode_prev_num_accepted_d state_slots = state_indices_tensor_d[:num_decodes, 0].long() + # For T=1, a checkpoint replays B old tokens and leaves the + # fresh token live. Derive the host-owned ring state before + # this call from the number of completed decode steps. + logical_window = self.replayssm_buffer_len + ring_start_cpu, prev_num_accepted_cpu = ( + _derive_flashinfer_replayssm_ring_state( + decode_steps_cpu, + logical_window, + ) + ) + ring_start = async_tensor_h2d( + ring_start_cpu.to(torch.int32).tolist(), + dtype=torch.int32, + device=common_attn_metadata.query_start_loc.device, + ) prev_num_accepted = async_tensor_h2d( - write_pos_cpu.to(torch.int32).tolist(), + prev_num_accepted_cpu.to(torch.int32).tolist(), dtype=torch.int32, device=common_attn_metadata.query_start_loc.device, ) + ring_start_d.index_copy_(0, state_slots, ring_start) prev_num_accepted_d.index_copy_(0, state_slots, prev_num_accepted) cb_scaled = self.decode_cb_scaled[:num_decodes] cumAdt_vec = self.decode_cumAdt_vec[:num_decodes] From 10e9d18408fb0ea909da6b3fa2a4da4575cb23e9 Mon Sep 17 00:00:00 2001 From: Andrii Skliar Date: Fri, 14 Aug 2026 17:07:07 +0200 Subject: [PATCH 06/33] validation for FlashInfer ReplaySSM cache mode to prevent unsupported configurations Signed-off-by: Andrii Skliar --- vllm/config/vllm.py | 8 ++++++++ 1 file changed, 8 insertions(+) diff --git a/vllm/config/vllm.py b/vllm/config/vllm.py index f575c89e44ce..bcba17222948 100644 --- a/vllm/config/vllm.py +++ b/vllm/config/vllm.py @@ -2718,6 +2718,14 @@ def validate_mamba_cached_kernel(self) -> "VllmConfig": "--use-replayssm requires --mamba-backend triton or flashinfer " f"(got {self.mamba_config.backend.value!r})" ) + if ( + self.mamba_config.backend == MambaBackendEnum.FLASHINFER + and self.cache_config.mamba_cache_mode == "align" + ): + raise ValueError( + "FlashInfer ReplaySSM does not support " + "--mamba-cache-mode align yet; use none" + ) if ( self.kv_transfer_config is not None and self.kv_transfer_config.is_kv_transfer_instance From 165f5a427cf5466a3bdad63230b490080c3f8b33 Mon Sep 17 00:00:00 2001 From: Andrii Skliar Date: Fri, 14 Aug 2026 17:20:53 +0200 Subject: [PATCH 07/33] add FlashInfer ReplaySSM ring tracker updates and lifecycle tests Signed-off-by: Andrii Skliar --- tests/kernels/mamba/test_ssu_dispatch.py | 24 ++++ .../test_replayssm_metadata_builder.py | 42 ++---- vllm/config/vllm.py | 1 - .../layers/mamba/mamba_mixer2.py | 88 ++++++++++-- .../layers/mamba/mamba_utils.py | 15 +- .../layers/mamba/ops/ssu_dispatch.py | 106 +++++++++++++- vllm/model_executor/models/nemotron_h.py | 9 +- vllm/v1/attention/backends/mamba_attn.py | 130 ++++-------------- 8 files changed, 268 insertions(+), 147 deletions(-) diff --git a/tests/kernels/mamba/test_ssu_dispatch.py b/tests/kernels/mamba/test_ssu_dispatch.py index c748c6424309..1893c70070c6 100644 --- a/tests/kernels/mamba/test_ssu_dispatch.py +++ b/tests/kernels/mamba/test_ssu_dispatch.py @@ -19,6 +19,7 @@ selective_state_update, selective_state_update_replayssm_flashinfer, selective_state_update_replayssm_triton, + update_replayssm_ring_trackers, ) from vllm.utils.torch_utils import set_random_seed from vllm.v1.attention.backends.registry import MambaAttentionBackendEnum @@ -43,6 +44,29 @@ HAS_FLASHINFER_CHECKPOINTING_SSU = False +@pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") +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, + ) + 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) + + def _kv_cache_config_with_ssu( mamba_type: MambaAttentionBackendEnum = MambaAttentionBackendEnum.MAMBA2, ) -> KVCacheConfig: diff --git a/tests/v1/attention/test_replayssm_metadata_builder.py b/tests/v1/attention/test_replayssm_metadata_builder.py index 3055e48e4111..2ba86c08dad8 100644 --- a/tests/v1/attention/test_replayssm_metadata_builder.py +++ b/tests/v1/attention/test_replayssm_metadata_builder.py @@ -17,9 +17,6 @@ create_vllm_config, ) from vllm.config.mamba import MambaBackendEnum -from vllm.v1.attention.backends.mamba_attn import ( - _derive_flashinfer_replayssm_ring_state, -) from vllm.v1.kv_cache_interface import MambaSpec BLOCK_SIZE = 16 @@ -194,19 +191,23 @@ def _make_mamba_spec( buffer_len: int, mamba_backend: MambaBackendEnum, ) -> MambaSpec: - # Five-tensor ReplaySSM page; the builder only reads shapes[4][0] (bc groups). + # The builder only reads the x/B ring shapes; include FlashInfer's trackers + # so the mock page matches the production cache layout. 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), + ) + if mamba_backend == MambaBackendEnum.FLASHINFER: + shapes = (*shapes, (), ()) return MambaSpec( block_size=BLOCK_SIZE, - shapes=( - (1, 1), - (1, 1, 1), - (1, ring_buffer_len, 1), - (1, ring_buffer_len), - (1, ring_buffer_len, 1), - ), + shapes=shapes, dtypes=(torch.float32,), ) @@ -272,35 +273,20 @@ def test_resumed_request_differs_from_fresh(): assert meta.is_flush_d.tolist()[:2] == [0, 0] -def test_flashinfer_replayssm_ring_metadata_fresh_decode(): - """FlashInfer receives per-slot ring metadata and per-row scratch.""" +def test_flashinfer_replayssm_scratch_metadata_fresh_decode(): + """FlashInfer receives per-row scratch; ring state is layer-local.""" builder = _create_replayssm_builder( 16, mamba_backend=MambaBackendEnum.FLASHINFER ) case = REPLAYSSM_BUILD_CASES["fresh_decode"] meta = _build(builder, case) - assert meta.ring_start_d is not None - assert meta.prev_num_accepted_d is not None assert meta.write_pos_d is None assert meta.is_flush_d is None assert meta.bc_pre_scratch is None - state_slot = int(meta.state_indices_tensor_d[0, 0]) - assert torch.count_nonzero(meta.ring_start_d) == 0 - assert int(meta.prev_num_accepted_d[state_slot]) == case.expected_write_pos[0] assert meta.cb_scaled is not None assert meta.cb_scaled.shape == (1, 1, 32, 8) assert meta.cumAdt_vec is not None assert meta.cumAdt_vec.shape == (1, 1, 16) assert meta.cb_old is not None assert meta.cb_old.shape == (1, 1, 32, 8) - - -def test_flashinfer_replayssm_ring_lifecycle_across_flushes(): - ring_start, prev_num_accepted = _derive_flashinfer_replayssm_ring_state( - torch.tensor([0, 5, 16, 17, 32, 33], dtype=torch.int32), - logical_window=16, - ) - - assert ring_start.tolist() == [0, 0, 0, 16, 16, 15] - assert prev_num_accepted.tolist() == [0, 5, 16, 1, 16, 1] diff --git a/vllm/config/vllm.py b/vllm/config/vllm.py index bcba17222948..d06555b05e8a 100644 --- a/vllm/config/vllm.py +++ b/vllm/config/vllm.py @@ -2670,7 +2670,6 @@ def validate_mamba_cached_kernel(self) -> "VllmConfig": self.cache_config.use_kda_recoverssm = False return self self.cache_config.use_kda_recoverssm = self.num_speculative_tokens > 0 - if self.model_config is not None and not self.model_config.supports_replayssm: raise ValueError( "--use-replayssm is not supported for architecture " diff --git a/vllm/model_executor/layers/mamba/mamba_mixer2.py b/vllm/model_executor/layers/mamba/mamba_mixer2.py index debdb3c3cb0f..7475b22a120c 100644 --- a/vllm/model_executor/layers/mamba/mamba_mixer2.py +++ b/vllm/model_executor/layers/mamba/mamba_mixer2.py @@ -38,6 +38,7 @@ ) from vllm.model_executor.layers.mamba.ops.ssu_dispatch import ( get_replayssm_backend, + reset_replayssm_ring_trackers, selective_state_update, selective_state_update_replayssm_flashinfer, selective_state_update_replayssm_triton, @@ -519,9 +520,19 @@ def __init__( "--use-replayssm requires tensor-parallel heads to divide evenly" ) # The tuple is (conv_state, ssm_state); with the cached (ReplaySSM) decode - # kernel enabled it is (conv_state, ssm_state, x_cache, dt_cache, B_cache). - _n_state = 5 if self.use_replayssm else 2 + # kernel enabled it also has x/dt/B rings and, for FlashInfer, two + # per-slot ring trackers. + if self.use_replayssm: + _n_state = ( + 7 + if self.mamba_config.backend == MambaBackendEnum.FLASHINFER + else 5 + ) + else: + _n_state = 2 self.kv_cache = tuple(torch.tensor([]) for _ in range(_n_state)) + self._replayssm_ring_start = torch.tensor([]) + self._replayssm_prev_num_accepted = torch.tensor([]) self.num_spec = vllm_config.num_speculative_tokens if self.num_spec > 0: @@ -549,6 +560,22 @@ def __init__( # Check if running on Blackwell (SM100+) for kernel tuning self.is_blackwell = current_platform.is_device_capability_family(100) + def _get_contiguous_replayssm_tracker( + self, + source: torch.Tensor, + attr_name: str, + ) -> torch.Tensor: + buffer = getattr(self, attr_name) + if ( + buffer.shape != source.shape + or buffer.device != source.device + or buffer.dtype != source.dtype + ): + buffer = torch.zeros_like(source, memory_format=torch.contiguous_format) + setattr(self, attr_name, buffer) + buffer.copy_(source) + return buffer + def forward( self, hidden_states: torch.Tensor, @@ -709,6 +736,8 @@ def conv_ssm_forward( assert self.cache_config is not None mamba_block_size = self.cache_config.mamba_block_size is_mamba_cache_all = self.cache_config.mamba_cache_mode == "all" + ring_start = prev_num_accepted = None + ring_start_src = prev_num_accepted_src = None attn_metadata: AttentionMetadata | None = None if attn_metadata_raw is not None: @@ -726,9 +755,28 @@ def conv_ssm_forward( ) ssm_state = self.kv_cache[1] if self.use_replayssm: - x_cache, dt_cache, B_cache = self.kv_cache[2:] + x_cache, dt_cache, B_cache = self.kv_cache[2:5] + if self.mamba_config.backend == MambaBackendEnum.FLASHINFER: + ring_start, prev_num_accepted = self.kv_cache[5:] + ring_start_src = ring_start + prev_num_accepted_src = prev_num_accepted + # Scalar states are strided across packed KV pages, while + # FlashInfer requires contiguous tracker arrays. + if not ring_start.is_contiguous(): + ring_start = self._get_contiguous_replayssm_tracker( + ring_start, + "_replayssm_ring_start", + ) + if not prev_num_accepted.is_contiguous(): + prev_num_accepted = self._get_contiguous_replayssm_tracker( + prev_num_accepted, + "_replayssm_prev_num_accepted", + ) + else: + ring_start = prev_num_accepted = None else: x_cache = dt_cache = B_cache = None + ring_start = prev_num_accepted = None has_initial_states_p = attn_metadata.has_initial_states_p prep_initial_states = attn_metadata.prep_initial_states chunk_size = attn_metadata.chunk_size @@ -985,6 +1033,13 @@ def conv_ssm_forward( # tensor assert state_indices_tensor_p is not None ssm_state[state_indices_tensor_p] = varlen_states + if ring_start is not None: + assert prev_num_accepted is not None + reset_replayssm_ring_trackers( + ring_start, + prev_num_accepted, + state_indices_tensor_p, + ) # Process decode requests if has_decode: @@ -1065,8 +1120,8 @@ def conv_ssm_forward( assert self.replayssm_buffer_len is not None replayssm_backend = get_replayssm_backend() if replayssm_backend.name == "flashinfer": - assert attn_metadata.ring_start_d is not None - assert attn_metadata.prev_num_accepted_d is not None + assert ring_start is not None + assert prev_num_accepted is not None assert attn_metadata.cb_scaled is not None assert attn_metadata.cumAdt_vec is not None assert attn_metadata.cb_old is not None @@ -1081,8 +1136,8 @@ def conv_ssm_forward( x_cache, B_cache, dt_cache, - attn_metadata.ring_start_d, - attn_metadata.prev_num_accepted_d, + ring_start, + prev_num_accepted, D=D_d, dt_bias=dt_bias, dt_softplus=True, @@ -1132,6 +1187,16 @@ def conv_ssm_forward( is_blackwell=self.is_blackwell, ) + if ring_start_src is not None and ring_start is not ring_start_src: + assert ring_start is not None + ring_start_src.copy_(ring_start) + if ( + prev_num_accepted_src is not None + and prev_num_accepted is not prev_num_accepted_src + ): + assert prev_num_accepted is not None + prev_num_accepted_src.copy_(prev_num_accepted) + def get_state_dtype(self) -> tuple[torch.dtype, ...]: assert self.model_config is not None assert self.cache_config is not None @@ -1142,7 +1207,11 @@ def get_state_dtype(self) -> tuple[torch.dtype, ...]: ) if self.use_replayssm: return MambaStateDtypeCalculator.append_replayssm_ring( - base_dtype, self.model_config.dtype + base_dtype, + self.model_config.dtype, + include_trackers=( + self.mamba_config.backend == MambaBackendEnum.FLASHINFER + ), ) return base_dtype @@ -1168,6 +1237,9 @@ def get_state_shape(self) -> tuple[tuple[int, ...], ...]: self.n_groups, tp_world_size, ring_buffer_len, + include_trackers=( + self.mamba_config.backend == MambaBackendEnum.FLASHINFER + ), ) return base_shape diff --git a/vllm/model_executor/layers/mamba/mamba_utils.py b/vllm/model_executor/layers/mamba/mamba_utils.py index ab7d99d34334..6ab18a81bc8f 100644 --- a/vllm/model_executor/layers/mamba/mamba_utils.py +++ b/vllm/model_executor/layers/mamba/mamba_utils.py @@ -85,12 +85,17 @@ def append_replayssm_ring( cls, base_dtypes: tuple[torch.dtype, ...], model_dtype: ModelDType | torch.dtype, + include_trackers: bool = False, ) -> tuple[torch.dtype, ...]: """Append the ReplaySSM ring dtypes to a base ``(conv, ssm)`` tuple: ``(x_cache, dt_cache, B_cache)`` = ``(activation, fp32, activation)``. + FlashInfer also appends int32 ``(ring_start, prev_num_accepted)``. """ activation_dtype = get_kv_cache_torch_dtype("auto", model_dtype) - return (*base_dtypes, activation_dtype, torch.float32, activation_dtype) + ring_dtypes = (activation_dtype, torch.float32, activation_dtype) + if include_trackers: + return (*base_dtypes, *ring_dtypes, torch.int32, torch.int32) + return (*base_dtypes, *ring_dtypes) @classmethod def _mamba_state_dtype( @@ -214,20 +219,24 @@ def append_replayssm_ring( n_groups: int, tp_world_size: int, ring_buffer_len: int, + include_trackers: bool = False, ) -> tuple[tuple[int, ...], ...]: - """Append the physical ReplaySSM ring shapes to ``(conv, ssm)``. + """Append the physical ReplaySSM ring and optional tracker shapes. ``base_shapes[1]`` is ``(nheads // tp, head_dim, state_size)``; B_cache uses the un-extended ``n_groups``. """ local_nheads, head_dim, state_size = base_shapes[1] local_ngroups = divide(n_groups, tp_world_size) - return ( + shapes = ( *base_shapes, (local_nheads, ring_buffer_len, head_dim), (local_nheads, ring_buffer_len), (local_ngroups, ring_buffer_len, state_size), ) + if include_trackers: + return (*shapes, (), ()) + return shapes @classmethod def short_conv_state_shape( diff --git a/vllm/model_executor/layers/mamba/ops/ssu_dispatch.py b/vllm/model_executor/layers/mamba/ops/ssu_dispatch.py index d85d5711a1d7..5fb064cbf2d7 100644 --- a/vllm/model_executor/layers/mamba/ops/ssu_dispatch.py +++ b/vllm/model_executor/layers/mamba/ops/ssu_dispatch.py @@ -22,6 +22,7 @@ from vllm.config.mamba import MambaBackendEnum, MambaConfig, MambaSSUAlgorithm from vllm.logger import init_logger +from vllm.triton_utils import tl, triton from vllm.v1.attention.backends.registry import MambaAttentionBackendEnum from vllm.v1.attention.backends.utils import NULL_BLOCK_ID from vllm.v1.kv_cache_interface import KVCacheConfig, MambaSpec @@ -29,6 +30,100 @@ logger = init_logger(__name__) +@triton.jit +def _update_replayssm_ring_trackers_kernel( + ring_start, + prev_num_accepted, + state_batch_indices, + n_slots, + logical_window: tl.constexpr, + ring_buffer_len: tl.constexpr, + pad_slot_id: tl.constexpr, + BLOCK: tl.constexpr, +) -> None: + offsets = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK) + mask = offsets < n_slots + slots = tl.load( + state_batch_indices + offsets, mask=mask, other=pad_slot_id + ) + valid = mask & (slots != pad_slot_id) + prev = tl.load(prev_num_accepted + slots, mask=valid, other=0) + start = tl.load(ring_start + slots, mask=valid, other=0) + must_checkpoint = prev + 1 > logical_window + next_start = tl.where( + must_checkpoint, + (start + prev) % ring_buffer_len, + start, + ) + next_prev = tl.where(must_checkpoint, 1, prev + 1) + tl.store(ring_start + slots, next_start, mask=valid) + tl.store(prev_num_accepted + slots, next_prev, mask=valid) + + +@triton.jit +def _reset_replayssm_ring_trackers_kernel( + ring_start, + prev_num_accepted, + state_batch_indices, + n_slots, + pad_slot_id: tl.constexpr, + BLOCK: tl.constexpr, +) -> None: + offsets = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK) + mask = offsets < n_slots + slots = tl.load( + state_batch_indices + offsets, mask=mask, other=pad_slot_id + ) + valid = mask & (slots != pad_slot_id) + tl.store(ring_start + slots, 0, mask=valid) + tl.store(prev_num_accepted + slots, 0, mask=valid) + + +def update_replayssm_ring_trackers( + ring_start: torch.Tensor, + prev_num_accepted: torch.Tensor, + state_batch_indices: torch.Tensor, + logical_window: int, + pad_slot_id: int = NULL_BLOCK_ID, +) -> None: + state_batch_indices = state_batch_indices.reshape(-1) + n_slots = state_batch_indices.numel() + if n_slots == 0: + return + block = 128 + _update_replayssm_ring_trackers_kernel[(triton.cdiv(n_slots, block),)]( + ring_start, + prev_num_accepted, + state_batch_indices, + n_slots, + logical_window, + logical_window + 1, + pad_slot_id, + BLOCK=block, + ) + + +def reset_replayssm_ring_trackers( + ring_start: torch.Tensor, + prev_num_accepted: torch.Tensor, + state_batch_indices: torch.Tensor, + pad_slot_id: int = NULL_BLOCK_ID, +) -> None: + state_batch_indices = state_batch_indices.reshape(-1) + n_slots = state_batch_indices.numel() + if n_slots == 0: + return + block = 128 + _reset_replayssm_ring_trackers_kernel[(triton.cdiv(n_slots, block),)]( + ring_start, + prev_num_accepted, + state_batch_indices, + n_slots, + pad_slot_id, + BLOCK=block, + ) + + class MambaSSUBackend(ABC): """Abstract base class for Mamba SSU backends.""" @@ -407,7 +502,7 @@ def __call__( if indices is not None and indices.dim() > 1: indices = indices[:, 0] - return self._kernel( + result = self._kernel( state, x_cache, B_cache, @@ -433,6 +528,15 @@ def __call__( cb_old=cb_old, algorithm=algorithm, ) + if indices is not None: + update_replayssm_ring_trackers( + ring_start, + prev_num_accepted_tokens, + indices, + logical_window=x_cache.size(2) - 1, + pad_slot_id=null_block_id, + ) + return result _REPLAYSSM_BACKEND_REGISTRY: dict[MambaBackendEnum, type[ReplaySSMBackend]] = { diff --git a/vllm/model_executor/models/nemotron_h.py b/vllm/model_executor/models/nemotron_h.py index 83e74bb6c54b..5ab8afd62f7b 100644 --- a/vllm/model_executor/models/nemotron_h.py +++ b/vllm/model_executor/models/nemotron_h.py @@ -748,7 +748,11 @@ def get_mamba_state_dtype_from_config( ) if cache_config.use_replayssm: return MambaStateDtypeCalculator.append_replayssm_ring( - base_dtype, vllm_config.model_config.dtype + base_dtype, + vllm_config.model_config.dtype, + include_trackers=( + vllm_config.mamba_config.backend == MambaBackendEnum.FLASHINFER + ), ) return base_dtype @@ -794,6 +798,9 @@ def get_mamba_state_shape_from_config( hf_config.n_groups, parallel_config.tensor_parallel_size, ring_buffer_len, + include_trackers=( + vllm_config.mamba_config.backend == MambaBackendEnum.FLASHINFER + ), ) return base_shape diff --git a/vllm/v1/attention/backends/mamba_attn.py b/vllm/v1/attention/backends/mamba_attn.py index 12e0e63e39dc..bbbb39305081 100644 --- a/vllm/v1/attention/backends/mamba_attn.py +++ b/vllm/v1/attention/backends/mamba_attn.py @@ -27,29 +27,6 @@ M = TypeVar("M", bound="BaseMambaAttentionMetadata") -def _derive_flashinfer_replayssm_ring_state( - decode_steps: torch.Tensor, - logical_window: int, -) -> tuple[torch.Tensor, torch.Tensor]: - """Return the ring cursor and live-token count before each T=1 call.""" - positive_steps = torch.clamp(decode_steps - 1, min=0) - flushes_before = torch.div( - positive_steps, - logical_window, - rounding_mode="floor", - ) - prev_num_accepted = torch.where( - decode_steps == 0, - torch.zeros_like(decode_steps), - torch.remainder(positive_steps, logical_window) + 1, - ) - ring_start = torch.remainder( - flushes_before * logical_window, - logical_window + 1, - ) - return ring_start, prev_num_accepted - - @dataclass class BaseMambaAttentionMetadata: num_prefills: int @@ -105,12 +82,8 @@ class BaseMambaAttentionMetadata: write_pos_d: torch.Tensor | None = None is_flush_d: torch.Tensor | None = None bc_pre_scratch: torch.Tensor | None = None - # ReplaySSM — FlashInfer checkpointing_ssu: - # ring_start / prev_num_accepted (always), plus two-kernel scratch - # (cb_scaled / cumAdt_vec / cb_old) so monolith and two-kernel are both - # available via algorithm="auto". All None unless that backend is on. - ring_start_d: torch.Tensor | None = None - prev_num_accepted_d: torch.Tensor | None = None + # ReplaySSM — FlashInfer checkpointing_ssu two-kernel scratch. + # The per-layer ring trackers live in the Mamba KV cache. cb_scaled: torch.Tensor | None = None cumAdt_vec: torch.Tensor | None = None cb_old: torch.Tensor | None = None @@ -235,8 +208,6 @@ def __init__( dtype=torch.float32, device=device, ) - self.decode_ring_start_d = None - self.decode_prev_num_accepted_d = None self.decode_cb_scaled = None self.decode_cumAdt_vec = None self.decode_cb_old = None @@ -245,22 +216,6 @@ def __init__( scratch_bs = max( self.decode_cudagraph_max_bs, scheduler_config.max_num_seqs ) - # checkpointing_ssu indexes these tensors by state_batch_indices, - # so they cover cache slots rather than decode rows. - num_cache_slots = vllm_config.cache_config.num_gpu_blocks - if num_cache_slots is None: - # Unit-test builders run before cache sizing. - num_cache_slots = scratch_bs - self.decode_ring_start_d = torch.zeros( - (num_cache_slots,), - dtype=torch.int32, - device=device, - ) - self.decode_prev_num_accepted_d = torch.zeros( - (num_cache_slots,), - dtype=torch.int32, - device=device, - ) # x_cache page: (nheads, L, head_dim); AR decode uses T=1. # Scratch shapes follow flashinfer's bench_checkpointing_ssu.py # (WARP_SIZE / MMA_FRAG_SIZE are not exported from Python). @@ -291,8 +246,6 @@ def __init__( ) else: self.decode_bc_pre_scratch = None - self.decode_ring_start_d = None - self.decode_prev_num_accepted_d = None self.decode_cb_scaled = None self.decode_cumAdt_vec = None self.decode_cb_old = None @@ -600,8 +553,6 @@ def _compute_common_metadata( nums_dict, batch_ptr, token_chunk_offset_ptr = None, None, None write_pos_d = None is_flush_d = None - ring_start_d = None - prev_num_accepted_d = None cb_scaled = None cumAdt_vec = None cb_old = None @@ -691,7 +642,11 @@ def _compute_common_metadata( num_reqs - num_prefills : num_reqs ] - if self.use_replayssm and num_decodes > 0: + if ( + self.use_replayssm + and not self.use_flashinfer_replayssm + and num_decodes > 0 + ): decode_base_cpu = common_attn_metadata.replayssm_decode_base_cpu num_computed_tokens_cpu = common_attn_metadata._num_computed_tokens_cpu if decode_base_cpu is None or num_computed_tokens_cpu is None: @@ -745,51 +700,24 @@ def _compute_common_metadata( & ((num_computed_d + query_lens_cpu) % block_size == 0) ) is_flush_cpu = is_flush_cpu.to(torch.int8) - if self.use_flashinfer_replayssm: - assert self.decode_ring_start_d is not None - assert self.decode_prev_num_accepted_d is not None - assert self.decode_cb_scaled is not None - assert self.decode_cumAdt_vec is not None - assert self.decode_cb_old is not None - ring_start_d = self.decode_ring_start_d - prev_num_accepted_d = self.decode_prev_num_accepted_d - state_slots = state_indices_tensor_d[:num_decodes, 0].long() - # For T=1, a checkpoint replays B old tokens and leaves the - # fresh token live. Derive the host-owned ring state before - # this call from the number of completed decode steps. - logical_window = self.replayssm_buffer_len - ring_start_cpu, prev_num_accepted_cpu = ( - _derive_flashinfer_replayssm_ring_state( - decode_steps_cpu, - logical_window, - ) - ) - ring_start = async_tensor_h2d( - ring_start_cpu.to(torch.int32).tolist(), - dtype=torch.int32, - device=common_attn_metadata.query_start_loc.device, - ) - prev_num_accepted = async_tensor_h2d( - prev_num_accepted_cpu.to(torch.int32).tolist(), - dtype=torch.int32, - device=common_attn_metadata.query_start_loc.device, - ) - ring_start_d.index_copy_(0, state_slots, ring_start) - prev_num_accepted_d.index_copy_(0, state_slots, prev_num_accepted) - cb_scaled = self.decode_cb_scaled[:num_decodes] - cumAdt_vec = self.decode_cumAdt_vec[:num_decodes] - cb_old = self.decode_cb_old[:num_decodes] - else: - write_pos_d = async_tensor_h2d( - write_pos_cpu.to(torch.int32).tolist(), - dtype=torch.int32, - device=common_attn_metadata.query_start_loc.device, - ) - is_flush_d = async_tensor_h2d( - is_flush_cpu.tolist(), - dtype=torch.int8, - device=common_attn_metadata.query_start_loc.device, - ) + write_pos_d = async_tensor_h2d( + write_pos_cpu.to(torch.int32).tolist(), + dtype=torch.int32, + device=common_attn_metadata.query_start_loc.device, + ) + is_flush_d = async_tensor_h2d( + is_flush_cpu.tolist(), + dtype=torch.int8, + device=common_attn_metadata.query_start_loc.device, + ) + + if self.use_flashinfer_replayssm and num_decodes > 0: + assert self.decode_cb_scaled is not None + assert self.decode_cumAdt_vec is not None + assert self.decode_cb_old is not None + cb_scaled = self.decode_cb_scaled[:num_decodes] + cumAdt_vec = self.decode_cumAdt_vec[:num_decodes] + cb_old = self.decode_cb_old[:num_decodes] bc_pre_scratch = None if ( @@ -811,8 +739,6 @@ def _compute_common_metadata( write_pos_d=write_pos_d, is_flush_d=is_flush_d, bc_pre_scratch=bc_pre_scratch, - ring_start_d=ring_start_d, - prev_num_accepted_d=prev_num_accepted_d, cb_scaled=cb_scaled, cumAdt_vec=cumAdt_vec, cb_old=cb_old, @@ -853,8 +779,6 @@ def _update_metadata_for_cudagraph_capture( write_pos_d = metadata.write_pos_d is_flush_d = metadata.is_flush_d bc_pre_scratch = metadata.bc_pre_scratch - ring_start_d = metadata.ring_start_d - prev_num_accepted_d = metadata.prev_num_accepted_d cb_scaled = metadata.cb_scaled cumAdt_vec = metadata.cumAdt_vec cb_old = metadata.cb_old @@ -939,8 +863,6 @@ def _update_metadata_for_cudagraph_capture( if self.decode_bc_pre_scratch is not None: bc_pre_scratch = self.decode_bc_pre_scratch[:padded_bs] elif self.use_flashinfer_replayssm: - assert ring_start_d is not None - assert prev_num_accepted_d is not None assert self.decode_cb_scaled is not None assert self.decode_cumAdt_vec is not None assert self.decode_cb_old is not None @@ -956,8 +878,6 @@ def _update_metadata_for_cudagraph_capture( write_pos_d=write_pos_d, is_flush_d=is_flush_d, bc_pre_scratch=bc_pre_scratch, - ring_start_d=ring_start_d, - prev_num_accepted_d=prev_num_accepted_d, cb_scaled=cb_scaled, cumAdt_vec=cumAdt_vec, cb_old=cb_old, From 6f2748103683019f26780cac66ff6a4160d2f477 Mon Sep 17 00:00:00 2001 From: Andrii Skliar Date: Fri, 14 Aug 2026 17:30:56 +0200 Subject: [PATCH 08/33] remove useless flags Signed-off-by: Andrii Skliar --- tests/v1/e2e/test_replayssm_decode.py | 11 ----------- 1 file changed, 11 deletions(-) diff --git a/tests/v1/e2e/test_replayssm_decode.py b/tests/v1/e2e/test_replayssm_decode.py index d044223fccec..c328341c93fa 100644 --- a/tests/v1/e2e/test_replayssm_decode.py +++ b/tests/v1/e2e/test_replayssm_decode.py @@ -16,9 +16,6 @@ except ImportError: HAS_FLASHINFER_CHECKPOINTING_SSU = False -# Flip when FlashInfer ReplaySSM metadata is implemented in mamba_attn. -FLASHINFER_REPLAYSSM_METADATA_READY = False - # Mamba2 (Nemotron-3) hybrid. MAMBA2_MODEL = "nvidia/NVIDIA-Nemotron-3-Nano-4B-BF16" MODELS = [ @@ -82,10 +79,6 @@ def test_replayssm_decode_matches_baseline_tp2(vllm_runner, model_name): not HAS_FLASHINFER_CHECKPOINTING_SSU, reason="flashinfer.mamba.checkpointing_ssu not available", ) -@pytest.mark.skipif( - not FLASHINFER_REPLAYSSM_METADATA_READY, - reason="FlashInfer ReplaySSM metadata not implemented yet", -) @pytest.mark.parametrize("model_name", MODELS) def test_replayssm_flashinfer_decode_matches_baseline(vllm_runner, model_name): _check_replayssm_parity( @@ -100,10 +93,6 @@ def test_replayssm_flashinfer_decode_matches_baseline(vllm_runner, model_name): not HAS_FLASHINFER_CHECKPOINTING_SSU, reason="flashinfer.mamba.checkpointing_ssu not available", ) -@pytest.mark.skipif( - not FLASHINFER_REPLAYSSM_METADATA_READY, - reason="FlashInfer ReplaySSM metadata not implemented yet", -) @pytest.mark.parametrize("model_name", MODELS) def test_replayssm_flashinfer_matches_triton_replayssm(vllm_runner, model_name): common = dict( From 6ae4546510495faaf718dfbda5719bb1034171fa Mon Sep 17 00:00:00 2001 From: Andrii Skliar Date: Sat, 15 Aug 2026 15:32:53 -0700 Subject: [PATCH 09/33] Add startup autotuning for FlashInfer ReplaySSM Signed-off-by: Andrii Skliar --- .../test_flashinfer_replayssm_warmup.py | 213 ++++++ tests/v1/worker/test_utils.py | 78 +++ .../layers/mamba/mamba_mixer2.py | 71 +- .../layers/mamba/ops/ssu_dispatch.py | 92 ++- .../warmup/flashinfer_replayssm_warmup.py | 650 ++++++++++++++++++ vllm/model_executor/warmup/kernel_warmup.py | 4 + vllm/v1/worker/gpu_model_runner.py | 18 +- 7 files changed, 1066 insertions(+), 60 deletions(-) create mode 100644 tests/model_executor/test_flashinfer_replayssm_warmup.py create mode 100644 vllm/model_executor/warmup/flashinfer_replayssm_warmup.py diff --git a/tests/model_executor/test_flashinfer_replayssm_warmup.py b/tests/model_executor/test_flashinfer_replayssm_warmup.py new file mode 100644 index 000000000000..64330800bc92 --- /dev/null +++ b/tests/model_executor/test_flashinfer_replayssm_warmup.py @@ -0,0 +1,213 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project + +from types import SimpleNamespace +from unittest.mock import Mock + +import pytest +import torch + +from vllm.config.mamba import MambaBackendEnum, MambaConfig +from vllm.forward_context import BatchDescriptor +from vllm.model_executor.layers.mamba.ops import ssu_dispatch +from vllm.model_executor.layers.mamba.ops.ssu_dispatch import ( + FLASHINFER_REPLAYSSM_AUTO_TACTIC, + FlashInferReplaySSMBackend, + FlashInferReplaySSMTactic, + use_flashinfer_replayssm_tactic, +) +from vllm.model_executor.warmup.flashinfer_replayssm_warmup import ( + FLASHINFER_REPLAYSSM_TUNING_CANDIDATES, + FlashInferReplaySSMAutotuneResult, + _load_cache, + _make_cache_key, + _ReplaySSMBenchmark, + _save_cache, + _select_fastest, + flashinfer_replayssm_autotune_warmup, +) + +_STAGES_ENV = "FLASHINFER_SSU_MAIN_PIPELINE_STAGES" +_CTAS_ENV = "FLASHINFER_SSU_MAIN_CTA_PER_SM" + + +def _fake_backend() -> FlashInferReplaySSMBackend: + backend = FlashInferReplaySSMBackend.__new__(FlashInferReplaySSMBackend) + backend._mamba_config = MambaConfig(backend=MambaBackendEnum.FLASHINFER) + backend._kernel = Mock(return_value=torch.empty(1)) + backend._algorithm = "auto" + return backend + + +def test_replayssm_tuning_candidates_and_deterministic_selection(): + assert [tactic.name for tactic in FLASHINFER_REPLAYSSM_TUNING_CANDIDATES] == [ + "auto", + "monolith", + "two_kernel_s1_c1", + "two_kernel_s1_c2", + "two_kernel_s1_c4", + "two_kernel_s1_c8", + "two_kernel_s1_c16", + "two_kernel_s2_c1", + "two_kernel_s2_c2", + "two_kernel_s2_c4", + "two_kernel_s2_c8", + "two_kernel_s2_c16", + ] + timings = [3.0, 2.0, 2.0, 4.0, 5.0, 6.0, 7.0, 8.0, 9.0, 10.0, 11.0, 12.0] + assert _select_fastest(timings).name == "monolith" + assert _select_fastest([float("inf")] * len(timings)) is None + + +def test_replayssm_tuning_key_distinguishes_batch_and_T(): + fingerprint = {"geometry": "h64_d64_n128"} + assert _make_cache_key(fingerprint, 32, 8) != _make_cache_key(fingerprint, 256, 1) + assert _make_cache_key(fingerprint, 32, 8) == _make_cache_key(fingerprint, 32, 8) + + +def test_replayssm_autotune_cache_round_trip(tmp_path): + path = tmp_path / "replayssm.json" + _save_cache(path, {"key": "two_kernel_s2_c16", "bad": "unknown"}) + assert _load_cache(path) == {"key": "two_kernel_s2_c16"} + + +@pytest.mark.parametrize("payload", ["null", "[]", "1"]) +def test_replayssm_autotune_ignores_non_mapping_cache(tmp_path, payload): + path = tmp_path / "replayssm.json" + path.write_text(payload) + assert _load_cache(path) == {} + + +def test_replayssm_benchmark_reserves_null_cache_slot(): + pytest.importorskip("flashinfer.mamba") + cache_slots, nheads, headdim, dstate, ngroups = 5, 2, 4, 8, 1 + layer = SimpleNamespace( + kv_cache=( + torch.empty(0), + torch.empty(cache_slots, nheads, headdim, dstate), + torch.empty(cache_slots, nheads, 17, headdim), + torch.empty(cache_slots, nheads, 17), + torch.empty(cache_slots, ngroups, 17, dstate), + ), + A=torch.empty(nheads), + D=torch.empty(nheads), + dt_bias=torch.empty(nheads), + mamba_config=SimpleNamespace( + enable_stochastic_rounding=False, + stochastic_rounding_philox_rounds=None, + ), + ) + + benchmark = _ReplaySSMBenchmark(layer, 3) + assert benchmark.indices.tolist() == [1, 2, 3] + assert benchmark.ring_start.shape == (cache_slots,) + assert benchmark.initial_ring_start.tolist() == [0, 0, 1, 2, 0] + assert benchmark.initial_prev_num_accepted.tolist() == [0, 1, 2, 3, 0] + _ReplaySSMBenchmark(layer, cache_slots - 1) + with pytest.raises(ValueError, match="needs 6 cache slots"): + _ReplaySSMBenchmark(layer, cache_slots) + + +@pytest.mark.parametrize( + ("use_v2_model_runner", "use_ubatching"), + [(True, False), (False, True)], +) +def test_replayssm_autotune_safely_skips_unsupported_runners( + use_v2_model_runner, use_ubatching +): + runner = SimpleNamespace( + parallel_config=SimpleNamespace(use_ubatching=use_ubatching) + ) + worker = SimpleNamespace( + model_runner=runner, + vllm_config=SimpleNamespace( + kernel_config=SimpleNamespace(enable_flashinfer_autotune=True) + ), + model_config=SimpleNamespace(enforce_eager=False), + use_v2_model_runner=use_v2_model_runner, + ) + + flashinfer_replayssm_autotune_warmup(worker) + assert runner.flashinfer_replayssm_autotune_result is None + + +def test_replayssm_tactic_scope_restores_algorithm_and_environment( + monkeypatch, +): + backend = _fake_backend() + monkeypatch.setattr(ssu_dispatch, "_replayssm_backend", backend) + monkeypatch.setenv(_STAGES_ENV, "7") + monkeypatch.setenv(_CTAS_ENV, "11") + + tactic = FlashInferReplaySSMTactic("two-kernel", 2, 16) + with ( + pytest.raises(RuntimeError, match="sentinel"), + use_flashinfer_replayssm_tactic(tactic), + ): + assert backend._algorithm == "two-kernel" + assert ssu_dispatch.os.environ[_STAGES_ENV] == "2" + assert ssu_dispatch.os.environ[_CTAS_ENV] == "16" + raise RuntimeError("sentinel") + + assert backend._algorithm == "auto" + assert ssu_dispatch.os.environ[_STAGES_ENV] == "7" + assert ssu_dispatch.os.environ[_CTAS_ENV] == "11" + + with use_flashinfer_replayssm_tactic(FLASHINFER_REPLAYSSM_AUTO_TACTIC): + assert _STAGES_ENV not in ssu_dispatch.os.environ + assert _CTAS_ENV not in ssu_dispatch.os.environ + + +def test_replayssm_backend_uses_scoped_algorithm(monkeypatch): + backend = _fake_backend() + monkeypatch.setattr(ssu_dispatch, "_replayssm_backend", backend) + tensor = torch.empty(1) + + with use_flashinfer_replayssm_tactic(FlashInferReplaySSMTactic("monolith")): + backend( + tensor, + tensor, + tensor, + tensor, + tensor, + tensor, + tensor, + tensor, + tensor, + tensor, + tensor, + tensor, + ) + + assert backend._kernel.call_args.kwargs["algorithm"] == "monolith" + + +def test_replayssm_capture_tactic_uses_request_batch_not_tokens(): + result = FlashInferReplaySSMAutotuneResult( + spec_query_len=1, + tactics={32: FlashInferReplaySSMTactic("two-kernel", 2, 8)}, + ) + assert ( + result.tactic_for( + BatchDescriptor(num_tokens=32, num_reqs=32, uniform=True) + ).name + == "two_kernel_s2_c8" + ) + assert ( + result.tactic_for(BatchDescriptor(num_tokens=256, num_reqs=32, uniform=True)) + is None + ) + mtp_result = FlashInferReplaySSMAutotuneResult( + spec_query_len=8, + tactics={32: FlashInferReplaySSMTactic("two-kernel", 2, 8)}, + ) + assert ( + mtp_result.tactic_for( + BatchDescriptor(num_tokens=256, num_reqs=32, uniform=True) + ).name + == "two_kernel_s2_c8" + ) + assert ( + result.tactic_for(BatchDescriptor(num_tokens=32, num_reqs=None, uniform=False)) + is None + ) diff --git a/tests/v1/worker/test_utils.py b/tests/v1/worker/test_utils.py index 016b8aa5a635..a760f1dbd1ba 100644 --- a/tests/v1/worker/test_utils.py +++ b/tests/v1/worker/test_utils.py @@ -1,11 +1,89 @@ # SPDX-License-Identifier: Apache-2.0 # SPDX-FileCopyrightText: Copyright contributors to the vLLM project +from types import SimpleNamespace + import torch +from vllm.config.mamba import MambaBackendEnum, MambaConfig +from vllm.model_executor.layers.mamba.mamba_mixer2 import MambaMixer2 from vllm.v1.worker.utils import bind_kv_cache +class _TestReplaySSMMixer(MambaMixer2): + _state_shapes = ((2,), (3,), (4,), (5,), (6,), (), ()) + _state_dtypes = ( + torch.float32, + torch.float32, + torch.float32, + torch.float32, + torch.float32, + torch.int32, + torch.int32, + ) + + def __init__(self): + torch.nn.Module.__init__(self) + self.use_replayssm = True + self.mamba_config = MambaConfig(backend=MambaBackendEnum.FLASHINFER) + self.cache_config = SimpleNamespace(mamba_cache_mode="none") + self._replayssm_ring_start = torch.empty(0, dtype=torch.int32) + self._replayssm_prev_num_accepted = torch.empty(0, dtype=torch.int32) + + def get_state_shape(self) -> tuple[tuple[int, ...], ...]: + return self._state_shapes + + def get_state_dtype(self) -> tuple[torch.dtype, ...]: + return self._state_dtypes + + +def _packed_replayssm_cache(num_blocks: int, fill_value: int = 0) -> torch.Tensor: + return torch.full((num_blocks, 1, 1, 88), fill_value, dtype=torch.int8) + + +def test_bind_kv_cache_uses_contiguous_replayssm_tracker_sidecars(): + mixer = _TestReplaySSMMixer() + mixer.bind_kv_cache(_packed_replayssm_cache(3, fill_value=1)) + + packed_ring_start, packed_prev_num_accepted = mixer.kv_cache[5:] + assert not packed_ring_start.is_contiguous() + assert not packed_prev_num_accepted.is_contiguous() + + for tracker in ( + mixer._replayssm_ring_start, + mixer._replayssm_prev_num_accepted, + ): + assert tracker.shape == (3,) + assert tracker.dtype == torch.int32 + assert tracker.is_contiguous() + assert torch.count_nonzero(tracker) == 0 + + assert torch.count_nonzero(packed_ring_start) == 3 + assert torch.count_nonzero(packed_prev_num_accepted) == 3 + assert not dict(mixer.named_buffers()) + + +def test_bind_kv_cache_recreates_replayssm_tracker_sidecars(): + mixer = _TestReplaySSMMixer() + mixer.bind_kv_cache(_packed_replayssm_cache(2)) + old_ring_start = mixer._replayssm_ring_start + old_prev_num_accepted = mixer._replayssm_prev_num_accepted + old_ring_start.fill_(7) + old_prev_num_accepted.fill_(9) + + mixer.bind_kv_cache(_packed_replayssm_cache(4)) + + assert mixer._replayssm_ring_start.shape == (4,) + assert mixer._replayssm_prev_num_accepted.shape == (4,) + assert torch.count_nonzero(mixer._replayssm_ring_start) == 0 + assert torch.count_nonzero(mixer._replayssm_prev_num_accepted) == 0 + assert mixer._replayssm_ring_start.data_ptr() != old_ring_start.data_ptr() + assert ( + mixer._replayssm_prev_num_accepted.data_ptr() + != old_prev_num_accepted.data_ptr() + ) + + def test_bind_kv_cache(default_vllm_config): from vllm.model_executor.layers.attention import Attention diff --git a/vllm/model_executor/layers/mamba/mamba_mixer2.py b/vllm/model_executor/layers/mamba/mamba_mixer2.py index 7475b22a120c..8478da96c9d1 100644 --- a/vllm/model_executor/layers/mamba/mamba_mixer2.py +++ b/vllm/model_executor/layers/mamba/mamba_mixer2.py @@ -524,15 +524,13 @@ def __init__( # per-slot ring trackers. if self.use_replayssm: _n_state = ( - 7 - if self.mamba_config.backend == MambaBackendEnum.FLASHINFER - else 5 + 7 if self.mamba_config.backend == MambaBackendEnum.FLASHINFER else 5 ) else: _n_state = 2 self.kv_cache = tuple(torch.tensor([]) for _ in range(_n_state)) - self._replayssm_ring_start = torch.tensor([]) - self._replayssm_prev_num_accepted = torch.tensor([]) + self._replayssm_ring_start = torch.empty(0, dtype=torch.int32) + self._replayssm_prev_num_accepted = torch.empty(0, dtype=torch.int32) self.num_spec = vllm_config.num_speculative_tokens if self.num_spec > 0: @@ -560,21 +558,29 @@ def __init__( # Check if running on Blackwell (SM100+) for kernel tuning self.is_blackwell = current_platform.is_device_capability_family(100) - def _get_contiguous_replayssm_tracker( - self, - source: torch.Tensor, - attr_name: str, - ) -> torch.Tensor: - buffer = getattr(self, attr_name) + def bind_kv_cache(self, kv_cache: torch.Tensor) -> None: + super().bind_kv_cache(kv_cache) if ( - buffer.shape != source.shape - or buffer.device != source.device - or buffer.dtype != source.dtype + self.use_replayssm + and self.mamba_config.backend == MambaBackendEnum.FLASHINFER ): - buffer = torch.zeros_like(source, memory_format=torch.contiguous_format) - setattr(self, attr_name, buffer) - buffer.copy_(source) - return buffer + assert self.cache_config is not None + assert self.cache_config.mamba_cache_mode == "none" + # FI ReplaySSM is restricted to cache mode "none", so these + # sidecars are authoritative; packed tracker fields remain reserved. + ring_start, prev_num_accepted = self.kv_cache[5:] + assert ring_start.dtype == torch.int32 + assert prev_num_accepted.dtype == torch.int32 + self._replayssm_ring_start = torch.zeros( + ring_start.shape, + dtype=torch.int32, + device=ring_start.device, + ) + self._replayssm_prev_num_accepted = torch.zeros( + prev_num_accepted.shape, + dtype=torch.int32, + device=prev_num_accepted.device, + ) def forward( self, @@ -737,7 +743,6 @@ def conv_ssm_forward( mamba_block_size = self.cache_config.mamba_block_size is_mamba_cache_all = self.cache_config.mamba_cache_mode == "all" ring_start = prev_num_accepted = None - ring_start_src = prev_num_accepted_src = None attn_metadata: AttentionMetadata | None = None if attn_metadata_raw is not None: @@ -757,21 +762,8 @@ def conv_ssm_forward( if self.use_replayssm: x_cache, dt_cache, B_cache = self.kv_cache[2:5] if self.mamba_config.backend == MambaBackendEnum.FLASHINFER: - ring_start, prev_num_accepted = self.kv_cache[5:] - ring_start_src = ring_start - prev_num_accepted_src = prev_num_accepted - # Scalar states are strided across packed KV pages, while - # FlashInfer requires contiguous tracker arrays. - if not ring_start.is_contiguous(): - ring_start = self._get_contiguous_replayssm_tracker( - ring_start, - "_replayssm_ring_start", - ) - if not prev_num_accepted.is_contiguous(): - prev_num_accepted = self._get_contiguous_replayssm_tracker( - prev_num_accepted, - "_replayssm_prev_num_accepted", - ) + ring_start = self._replayssm_ring_start + prev_num_accepted = self._replayssm_prev_num_accepted else: ring_start = prev_num_accepted = None else: @@ -1145,7 +1137,6 @@ def conv_ssm_forward( cb_scaled=attn_metadata.cb_scaled, cumAdt_vec=attn_metadata.cumAdt_vec, cb_old=attn_metadata.cb_old, - algorithm="auto", ) else: selective_state_update_replayssm_triton( @@ -1187,16 +1178,6 @@ def conv_ssm_forward( is_blackwell=self.is_blackwell, ) - if ring_start_src is not None and ring_start is not ring_start_src: - assert ring_start is not None - ring_start_src.copy_(ring_start) - if ( - prev_num_accepted_src is not None - and prev_num_accepted is not prev_num_accepted_src - ): - assert prev_num_accepted is not None - prev_num_accepted_src.copy_(prev_num_accepted) - def get_state_dtype(self) -> tuple[torch.dtype, ...]: assert self.model_config is not None assert self.cache_config is not None diff --git a/vllm/model_executor/layers/mamba/ops/ssu_dispatch.py b/vllm/model_executor/layers/mamba/ops/ssu_dispatch.py index 5fb064cbf2d7..95d89b6283e8 100644 --- a/vllm/model_executor/layers/mamba/ops/ssu_dispatch.py +++ b/vllm/model_executor/layers/mamba/ops/ssu_dispatch.py @@ -16,7 +16,11 @@ (``selective_state_update_replayssm_flashinfer``) """ +import os from abc import ABC, abstractmethod +from collections.abc import Iterator +from contextlib import contextmanager +from dataclasses import dataclass import torch @@ -29,6 +33,41 @@ logger = init_logger(__name__) +_FLASHINFER_SSU_PIPELINE_STAGES_ENV = "FLASHINFER_SSU_MAIN_PIPELINE_STAGES" +_FLASHINFER_SSU_CTA_PER_SM_ENV = "FLASHINFER_SSU_MAIN_CTA_PER_SM" + + +@dataclass(frozen=True) +class FlashInferReplaySSMTactic: + algorithm: str + pipeline_stages: int | None = None + ctas_per_sm: int | None = None + + def __post_init__(self) -> None: + if self.algorithm not in {"auto", "monolith", "two-kernel"}: + raise ValueError(f"Unsupported ReplaySSM algorithm: {self.algorithm}") + has_launch_config = ( + self.pipeline_stages is not None or self.ctas_per_sm is not None + ) + if self.algorithm == "two-kernel": + if self.pipeline_stages not in {1, 2}: + raise ValueError("two-kernel requires pipeline_stages in {1, 2}") + if self.ctas_per_sm is None or self.ctas_per_sm <= 0: + raise ValueError("two-kernel requires a positive ctas_per_sm") + elif has_launch_config: + raise ValueError( + f"{self.algorithm} does not accept pipeline or CTA settings" + ) + + @property + def name(self) -> str: + if self.algorithm != "two-kernel": + return self.algorithm + return f"two_kernel_s{self.pipeline_stages}_c{self.ctas_per_sm}" + + +FLASHINFER_REPLAYSSM_AUTO_TACTIC = FlashInferReplaySSMTactic("auto") + @triton.jit def _update_replayssm_ring_trackers_kernel( @@ -43,9 +82,7 @@ def _update_replayssm_ring_trackers_kernel( ) -> None: offsets = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK) mask = offsets < n_slots - slots = tl.load( - state_batch_indices + offsets, mask=mask, other=pad_slot_id - ) + slots = tl.load(state_batch_indices + offsets, mask=mask, other=pad_slot_id) valid = mask & (slots != pad_slot_id) prev = tl.load(prev_num_accepted + slots, mask=valid, other=0) start = tl.load(ring_start + slots, mask=valid, other=0) @@ -71,9 +108,7 @@ def _reset_replayssm_ring_trackers_kernel( ) -> None: offsets = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK) mask = offsets < n_slots - slots = tl.load( - state_batch_indices + offsets, mask=mask, other=pad_slot_id - ) + slots = tl.load(state_batch_indices + offsets, mask=mask, other=pad_slot_id) valid = mask & (slots != pad_slot_id) tl.store(ring_start + slots, 0, mask=valid) tl.store(prev_num_accepted + slots, 0, mask=valid) @@ -453,6 +488,7 @@ def __init__(self, mamba_config: MambaConfig): "pip install flashinfer-python" ) from e self._kernel = _fi_checkpointing_ssu + self._algorithm = FLASHINFER_REPLAYSSM_AUTO_TACTIC.algorithm @property def name(self) -> str: @@ -481,7 +517,7 @@ def __call__( cb_scaled: torch.Tensor | None = None, cumAdt_vec: torch.Tensor | None = None, cb_old: torch.Tensor | None = None, - algorithm: str = "auto", + algorithm: str | None = None, ) -> torch.Tensor: # AR decode currently passes (batch, nheads, dim); checkpointing_ssu # expects a predicted-token axis T. Unsqueeze T=1 here. @@ -526,7 +562,7 @@ def __call__( cb_scaled=cb_scaled, cumAdt_vec=cumAdt_vec, cb_old=cb_old, - algorithm=algorithm, + algorithm=self._algorithm if algorithm is None else algorithm, ) if indices is not None: update_replayssm_ring_trackers( @@ -539,6 +575,44 @@ def __call__( return result +@contextmanager +def use_flashinfer_replayssm_tactic( + tactic: FlashInferReplaySSMTactic, +) -> Iterator[None]: + """Apply a ReplaySSM launch tactic during serial warmup or graph capture.""" + backend = get_replayssm_backend() + if not isinstance(backend, FlashInferReplaySSMBackend): + yield + return + + old_algorithm = backend._algorithm + old_stages = os.environ.get(_FLASHINFER_SSU_PIPELINE_STAGES_ENV) + old_ctas = os.environ.get(_FLASHINFER_SSU_CTA_PER_SM_ENV) + backend._algorithm = tactic.algorithm + try: + if tactic.algorithm == "two-kernel": + assert tactic.pipeline_stages is not None + assert tactic.ctas_per_sm is not None + os.environ[_FLASHINFER_SSU_PIPELINE_STAGES_ENV] = str( + tactic.pipeline_stages + ) + os.environ[_FLASHINFER_SSU_CTA_PER_SM_ENV] = str(tactic.ctas_per_sm) + else: + os.environ.pop(_FLASHINFER_SSU_PIPELINE_STAGES_ENV, None) + os.environ.pop(_FLASHINFER_SSU_CTA_PER_SM_ENV, None) + yield + finally: + backend._algorithm = old_algorithm + if old_stages is None: + os.environ.pop(_FLASHINFER_SSU_PIPELINE_STAGES_ENV, None) + else: + os.environ[_FLASHINFER_SSU_PIPELINE_STAGES_ENV] = old_stages + if old_ctas is None: + os.environ.pop(_FLASHINFER_SSU_CTA_PER_SM_ENV, None) + else: + os.environ[_FLASHINFER_SSU_CTA_PER_SM_ENV] = old_ctas + + _REPLAYSSM_BACKEND_REGISTRY: dict[MambaBackendEnum, type[ReplaySSMBackend]] = { MambaBackendEnum.TRITON: TritonReplaySSMBackend, MambaBackendEnum.FLASHINFER: FlashInferReplaySSMBackend, @@ -659,7 +733,7 @@ def selective_state_update_replayssm_flashinfer( cb_scaled: torch.Tensor | None = None, cumAdt_vec: torch.Tensor | None = None, cb_old: torch.Tensor | None = None, - algorithm: str = "auto", + algorithm: str | None = None, ) -> torch.Tensor: """FlashInfer ReplaySSM decode (``checkpointing_ssu``).""" backend = get_replayssm_backend() diff --git a/vllm/model_executor/warmup/flashinfer_replayssm_warmup.py b/vllm/model_executor/warmup/flashinfer_replayssm_warmup.py new file mode 100644 index 000000000000..c8a9bd276440 --- /dev/null +++ b/vllm/model_executor/warmup/flashinfer_replayssm_warmup.py @@ -0,0 +1,650 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""Startup autotuning for FlashInfer ReplaySSM CUDA-graph launches.""" + +from __future__ import annotations + +import gc +import json +import math +import statistics +import time +from contextlib import contextmanager, nullcontext +from dataclasses import dataclass +from pathlib import Path +from typing import TYPE_CHECKING, Any + +import torch + +from vllm.config.mamba import MambaBackendEnum +from vllm.distributed.parallel_state import get_world_group +from vllm.forward_context import BatchDescriptor +from vllm.logger import init_logger +from vllm.model_executor.layers.mamba.ops.ssu_dispatch import ( + FLASHINFER_REPLAYSSM_AUTO_TACTIC, + FlashInferReplaySSMTactic, + update_replayssm_ring_trackers, + use_flashinfer_replayssm_tactic, +) +from vllm.model_executor.warmup.flashinfer_autotune_cache import ( + resolve_flashinfer_autotune_file, + write_flashinfer_autotune_cache, +) + +if TYPE_CHECKING: + from vllm.model_executor.layers.mamba.mamba_mixer2 import MambaMixer2 + from vllm.v1.worker.gpu_model_runner import GPUModelRunner + from vllm.v1.worker.gpu_worker import Worker + +logger = init_logger(__name__) + +_CACHE_SCHEMA_VERSION = 1 +_TUNING_FIXTURE = "t1_mixed_history_cycle_v2" +_CACHE_FILE_NAME = "replayssm_autotune_configs.json" + +FLASHINFER_REPLAYSSM_TUNING_CANDIDATES = ( + FLASHINFER_REPLAYSSM_AUTO_TACTIC, + FlashInferReplaySSMTactic("monolith"), + FlashInferReplaySSMTactic("two-kernel", 1, 1), + FlashInferReplaySSMTactic("two-kernel", 1, 2), + FlashInferReplaySSMTactic("two-kernel", 1, 4), + FlashInferReplaySSMTactic("two-kernel", 1, 8), + FlashInferReplaySSMTactic("two-kernel", 1, 16), + FlashInferReplaySSMTactic("two-kernel", 2, 1), + FlashInferReplaySSMTactic("two-kernel", 2, 2), + FlashInferReplaySSMTactic("two-kernel", 2, 4), + FlashInferReplaySSMTactic("two-kernel", 2, 8), + FlashInferReplaySSMTactic("two-kernel", 2, 16), +) +_TACTICS_BY_NAME = { + tactic.name: tactic for tactic in FLASHINFER_REPLAYSSM_TUNING_CANDIDATES +} + + +@dataclass +class FlashInferReplaySSMAutotuneResult: + spec_query_len: int + tactics: dict[int, FlashInferReplaySSMTactic] + + def tactic_for( + self, batch_descriptor: BatchDescriptor + ) -> FlashInferReplaySSMTactic | None: + if ( + not batch_descriptor.uniform + or batch_descriptor.num_reqs is None + or batch_descriptor.num_tokens + != batch_descriptor.num_reqs * self.spec_query_len + ): + return None + return self.tactics.get(batch_descriptor.num_reqs) + + +def _make_cache_key(fingerprint: dict[str, Any], batch: int, T: int) -> str: + return json.dumps( + { + **fingerprint, + "batch_sequences": batch, + "spec_query_len": T, + "fixture": _TUNING_FIXTURE, + "candidate_schema": [ + tactic.name for tactic in FLASHINFER_REPLAYSSM_TUNING_CANDIDATES + ], + }, + sort_keys=True, + separators=(",", ":"), + ) + + +def _select_fastest(timings: list[float]) -> FlashInferReplaySSMTactic | None: + if len(timings) != len(FLASHINFER_REPLAYSSM_TUNING_CANDIDATES): + raise ValueError("one timing is required for each ReplaySSM tactic") + finite = [i for i, timing in enumerate(timings) if math.isfinite(timing)] + if not finite: + return None + winner = min(finite, key=lambda i: (timings[i], i)) + return FLASHINFER_REPLAYSSM_TUNING_CANDIDATES[winner] + + +def _load_cache(path: Path) -> dict[str, str]: + try: + payload = json.loads(path.read_text()) + except (OSError, ValueError, TypeError): + return {} + if not isinstance(payload, dict): + return {} + if payload.get("schema_version") != _CACHE_SCHEMA_VERSION: + return {} + entries = payload.get("entries") + if not isinstance(entries, dict): + return {} + return { + key: value + for key, value in entries.items() + if isinstance(key, str) and isinstance(value, str) and value in _TACTICS_BY_NAME + } + + +def _save_cache(path: Path, entries: dict[str, str]) -> None: + payload = { + "schema_version": _CACHE_SCHEMA_VERSION, + "entries": dict(sorted(entries.items())), + } + write_flashinfer_autotune_cache( + path, + (json.dumps(payload, sort_keys=True, indent=2) + "\n").encode(), + ) + + +def _find_replayssm_layers(runner: GPUModelRunner) -> tuple[MambaMixer2, ...]: + from vllm.model_executor.layers.mamba.mamba_mixer2 import MambaMixer2 + + return tuple( + module + for module in runner.get_model().modules() + if ( + isinstance(module, MambaMixer2) + and module.use_replayssm + and module.mamba_config.backend == MambaBackendEnum.FLASHINFER + ) + ) + + +def _uniform_capture_batches(runner: GPUModelRunner) -> tuple[int, ...]: + return tuple( + sorted( + { + desc.num_reqs + for _, descs in runner.cudagraph_dispatcher.get_capture_descs() + for desc in descs + if desc.uniform and desc.num_reqs is not None + } + ) + ) + + +def _layer_fingerprint(layer: MambaMixer2) -> dict[str, Any]: + import flashinfer + from flashinfer.jit import env as flashinfer_jit_env + + _, state, x_cache, dt_cache, B_cache, *_ = layer.kv_cache + props = torch.cuda.get_device_properties(state.device) + return { + "flashinfer_version": flashinfer.__version__, + "flashinfer_workspace": [ + flashinfer_jit_env.FLASHINFER_WORKSPACE_DIR.parent.name, + flashinfer_jit_env.FLASHINFER_WORKSPACE_DIR.name, + ], + "gpu_name": props.name, + "gpu_capability": list(torch.cuda.get_device_capability(state.device)), + "gpu_sm_count": props.multi_processor_count, + "nheads": state.shape[1], + "headdim": state.shape[2], + "dstate": state.shape[3], + "ngroups": B_cache.shape[1], + "cache_slots": state.shape[0], + "physical_ring_len": x_cache.shape[2], + "state_dtype": str(state.dtype), + "activation_dtype": str(x_cache.dtype), + "dt_cache_dtype": str(dt_cache.dtype), + "state_stride": list(state.stride()), + "x_cache_stride": list(x_cache.stride()), + "B_cache_stride": list(B_cache.stride()), + "dt_cache_stride": list(dt_cache.stride()), + "A_dtype": str(layer.A.dtype), + "A_stride": list(layer.A.stride()), + "D_dtype": str(layer.D.dtype), + "D_stride": list(layer.D.stride()), + "dt_bias_dtype": str(layer.dt_bias.dtype), + "dt_bias_stride": list(layer.dt_bias.stride()), + "stochastic_rounding": layer.mamba_config.enable_stochastic_rounding, + "stochastic_rounding_philox_rounds": ( + layer.mamba_config.stochastic_rounding_philox_rounds + ), + "tp_size": layer.tp_size, + } + + +class _ReplaySSMBenchmark: + def __init__(self, layer: MambaMixer2, batch: int): + from flashinfer.mamba import checkpointing_ssu + + _, self.state, self.x_cache, self.dt_cache, self.B_cache, *_ = layer.kv_cache + if self.state.shape[0] <= batch: + raise ValueError( + f"ReplaySSM autotune batch {batch} needs {batch + 1} cache " + f"slots, but only {self.state.shape[0]} are available" + ) + + self._kernel = checkpointing_ssu + self.batch = batch + self.logical_window = self.x_cache.shape[2] - 1 + if self.logical_window <= 0: + raise ValueError("ReplaySSM history window must be positive") + + device = self.state.device + activation_dtype = self.x_cache.dtype + nheads = self.state.shape[1] + headdim = self.state.shape[2] + dstate = self.state.shape[3] + ngroups = self.B_cache.shape[1] + generator = torch.Generator(device=device) + generator.manual_seed(0x5253534D + batch) + + self.x = torch.randn( + batch, + 1, + nheads, + headdim, + dtype=activation_dtype, + device=device, + generator=generator, + ) + dt_base = torch.randn( + batch, + 1, + nheads, + dtype=activation_dtype, + device=device, + generator=generator, + ) + self.dt = dt_base.unsqueeze(-1).expand(batch, 1, nheads, headdim) + self.B = torch.randn( + batch, + 1, + ngroups, + dstate, + dtype=activation_dtype, + device=device, + generator=generator, + ) + self.C = torch.randn( + self.B.shape, + dtype=activation_dtype, + device=device, + generator=generator, + ) + self.out = torch.empty_like(self.x) + self.indices = torch.arange(1, batch + 1, dtype=torch.int32, device=device) + self.ring_start = torch.zeros( + self.state.shape[0], dtype=torch.int32, device=device + ) + self.prev_num_accepted = torch.zeros_like(self.ring_start) + rows = torch.arange(batch, dtype=torch.int32, device=device) + self.initial_prev_num_accepted = torch.zeros_like(self.prev_num_accepted) + self.initial_prev_num_accepted[1 : batch + 1] = rows.remainder( + self.logical_window + ).add_(1) + self.initial_ring_start = torch.zeros_like(self.ring_start) + self.initial_ring_start[1 : batch + 1] = rows.remainder(self.logical_window + 1) + + self.A = ( + layer.A[:, None, ...][:, :, None] + .expand(-1, headdim, dstate) + .to(dtype=torch.float32) + ) + self.D = layer.D[:, None, ...].expand(-1, headdim) + self.dt_bias = layer.dt_bias[:, None, ...].expand(-1, headdim) + self.rand_seed = ( + torch.zeros(1, dtype=torch.int64, device=device) + if layer.mamba_config.enable_stochastic_rounding + else None + ) + self.philox_rounds = layer.mamba_config.stochastic_rounding_philox_rounds or 10 + + k_old = ((self.logical_window + 7) // 8) * 8 + self.cb_scaled = torch.empty( + batch, nheads, 32, 8, dtype=activation_dtype, device=device + ) + self.cumAdt_vec = torch.empty( + batch, nheads, 16, dtype=torch.float32, device=device + ) + self.cb_old = torch.empty( + batch, + nheads, + 32, + k_old // 2, + dtype=activation_dtype, + device=device, + ) + + def reset(self) -> None: + for tensor in (self.state, self.x_cache, self.dt_cache, self.B_cache): + tensor[: self.batch + 1].zero_() + self.out.zero_() + self.ring_start.copy_(self.initial_ring_start) + self.prev_num_accepted.copy_(self.initial_prev_num_accepted) + + def call(self, tactic: FlashInferReplaySSMTactic) -> None: + self._kernel( + self.state, + self.x_cache, + self.B_cache, + self.dt_cache, + self.ring_start, + self.prev_num_accepted, + self.x, + self.dt, + self.A, + self.B, + self.C, + self.out, + D=self.D, + dt_bias=self.dt_bias, + dt_softplus=True, + state_batch_indices=self.indices, + pad_slot_id=0, + rand_seed=self.rand_seed, + philox_rounds=self.philox_rounds, + cb_scaled=self.cb_scaled, + cumAdt_vec=self.cumAdt_vec, + cb_old=self.cb_old, + precompute_heads_per_cta=0, + algorithm=tactic.algorithm, + ) + update_replayssm_ring_trackers( + self.ring_start, + self.prev_num_accepted, + self.indices, + logical_window=self.logical_window, + pad_slot_id=0, + ) + + def benchmark(self, tactic: FlashInferReplaySSMTactic) -> float: + calls_per_graph = self.logical_window + graph: torch.cuda.CUDAGraph | None = None + samples = [] + try: + with use_flashinfer_replayssm_tactic(tactic): + self.reset() + side_stream = torch.cuda.Stream() + side_stream.wait_stream(torch.cuda.current_stream()) + with torch.cuda.stream(side_stream): + for _ in range(3): + self.call(tactic) + torch.cuda.current_stream().wait_stream(side_stream) + torch.cuda.synchronize() + + self.reset() + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph): + for _ in range(calls_per_graph): + self.call(tactic) + torch.cuda.synchronize() + + for _ in range(3): + self.reset() + for _ in range(3): + graph.replay() + torch.cuda.synchronize() + start = torch.cuda.Event(enable_timing=True) + end = torch.cuda.Event(enable_timing=True) + start.record() + for _ in range(10): + graph.replay() + end.record() + end.synchronize() + samples.append(start.elapsed_time(end) / (10 * calls_per_graph)) + finally: + self.reset() + del graph + return statistics.median(samples) + + +def _aggregate_timings(timings: list[float]) -> list[float]: + world = get_world_group() + if world.world_size == 1: + return timings + values = torch.tensor(timings, dtype=torch.float64) + torch.distributed.all_reduce( + values, + op=torch.distributed.ReduceOp.MAX, + group=world.cpu_group, + ) + return values.tolist() + + +def _all_ranks_support_tuning(supported: bool) -> bool: + world = get_world_group() + if world.world_size == 1: + return supported + flag = torch.tensor([int(supported)], dtype=torch.int32) + torch.distributed.all_reduce( + flag, + op=torch.distributed.ReduceOp.MIN, + group=world.cpu_group, + ) + return bool(flag.item()) + + +@torch.inference_mode() +def flashinfer_replayssm_autotune_warmup(worker: Worker) -> None: + runner = worker.model_runner + runner.flashinfer_replayssm_autotune_result = None + if worker.vllm_config.kernel_config.enable_flashinfer_autotune is not True: + return + if worker.model_config.enforce_eager: + logger.info_once("Skipping FlashInfer ReplaySSM autotune without CUDA graphs.") + return + if getattr(worker, "use_v2_model_runner", False): + logger.info_once( + "Skipping FlashInfer ReplaySSM autotune with the V2 model runner." + ) + return + if runner.parallel_config.use_ubatching: + logger.info_once( + "Skipping FlashInfer ReplaySSM autotune with uniform microbatching." + ) + return + + layers = _find_replayssm_layers(runner) + if not _all_ranks_support_tuning(bool(layers)): + return + assert layers + local_layer_fingerprints = None + try: + local_layer_fingerprints = tuple(_layer_fingerprint(layer) for layer in layers) + except Exception: + logger.warning( + "Could not fingerprint the FlashInfer ReplaySSM layers.", + exc_info=True, + ) + if not _all_ranks_support_tuning(local_layer_fingerprints is not None): + logger.warning_once( + "Skipping FlashInfer ReplaySSM autotune because layer " + "fingerprinting failed on at least one rank." + ) + return + assert local_layer_fingerprints is not None + layers_are_homogeneous = all( + fingerprint == local_layer_fingerprints[0] + for fingerprint in local_layer_fingerprints[1:] + ) + if not _all_ranks_support_tuning(layers_are_homogeneous): + logger.warning_once( + "Skipping FlashInfer ReplaySSM autotune because ReplaySSM layer " + "geometries differ within a rank." + ) + return + layer = layers[0] + T = runner.uniform_decode_query_len + if T != 1: + logger.warning_once( + "Skipping FlashInfer ReplaySSM autotune for T=%d; this vLLM " + "integration currently supports only T=1.", + T, + ) + return + + local_batches = _uniform_capture_batches(runner) + world = get_world_group() + batches = world.broadcast_object( + local_batches if world.rank_in_group == 0 else None, src=0 + ) + if not _all_ranks_support_tuning(local_batches == batches): + logger.warning_once( + "Skipping FlashInfer ReplaySSM autotune because CUDA-graph " + "descriptors differ across ranks." + ) + return + if not batches: + logger.info_once( + "Skipping FlashInfer ReplaySSM autotune because there are no " + "uniform FULL CUDA-graph capture descriptors." + ) + return + + local_fingerprint = local_layer_fingerprints[0] + fingerprint = world.broadcast_object( + local_fingerprint if world.rank_in_group == 0 else None, src=0 + ) + if not _all_ranks_support_tuning(local_fingerprint == fingerprint): + logger.warning_once( + "Skipping FlashInfer ReplaySSM autotune because Mamba geometry " + "differs across ranks." + ) + return + + cache_path = None + try: + cache_path = resolve_flashinfer_autotune_file(runner).with_name( + _CACHE_FILE_NAME + ) + except Exception: + logger.warning( + "Could not resolve the FlashInfer ReplaySSM autotune cache path.", + exc_info=True, + ) + if not _all_ranks_support_tuning(cache_path is not None): + logger.warning_once( + "Skipping FlashInfer ReplaySSM autotune because the cache path " + "could not be resolved on every rank." + ) + return + assert cache_path is not None + cached_entries = _load_cache(cache_path) if world.rank_in_group == 0 else None + cached_entries = world.broadcast_object(cached_entries, src=0) + assert cached_entries is not None + + selected: dict[int, FlashInferReplaySSMTactic] = {} + cache_changed = False + tuning_started = time.perf_counter() + for batch in batches: + cache_key = _make_cache_key(fingerprint, batch, T) + cached_name = cached_entries.get(cache_key) + if cached_name in _TACTICS_BY_NAME: + tactic = _TACTICS_BY_NAME[cached_name] + selected[batch] = tactic + if world.rank_in_group == 0: + logger.info( + "FlashInfer ReplaySSM autotune cache hit for batch %d: %s", + batch, + tactic.name, + ) + continue + + benchmark = None + try: + benchmark = _ReplaySSMBenchmark(layer, batch) + except Exception: + logger.warning( + "Could not construct the FlashInfer ReplaySSM benchmark for batch %d.", + batch, + exc_info=True, + ) + if not _all_ranks_support_tuning(benchmark is not None): + logger.warning_once( + "Skipping FlashInfer ReplaySSM autotune for batch %d because " + "the benchmark could not be constructed on every rank.", + batch, + ) + del benchmark + gc.collect() + torch.cuda.empty_cache() + continue + assert benchmark is not None + + candidate_count = len(FLASHINFER_REPLAYSSM_TUNING_CANDIDATES) + shift = batch % candidate_count + candidate_indices = tuple(range(candidate_count)) + candidate_order = candidate_indices[shift:] + candidate_indices[:shift] + if (batch // candidate_count) % 2: + candidate_order = tuple(reversed(candidate_order)) + local_timings = [float("inf")] * candidate_count + for candidate_index in candidate_order: + tactic = FLASHINFER_REPLAYSSM_TUNING_CANDIDATES[candidate_index] + try: + local_timings[candidate_index] = benchmark.benchmark(tactic) + except Exception: + logger.warning( + "FlashInfer ReplaySSM tactic %s failed for batch %d.", + tactic.name, + batch, + exc_info=True, + ) + timings = _aggregate_timings(local_timings) + tactic = _select_fastest(timings) + if tactic is None: + logger.warning_once( + "Every FlashInfer ReplaySSM tactic failed for batch %d; " + "leaving the default launch policy unchanged.", + batch, + ) + del benchmark + gc.collect() + torch.cuda.empty_cache() + continue + selected[batch] = tactic + cached_entries[cache_key] = tactic.name + cache_changed = True + if world.rank_in_group == 0: + timing_log = ", ".join( + f"{candidate.name}={timing:.6f}ms" + for candidate, timing in zip( + FLASHINFER_REPLAYSSM_TUNING_CANDIDATES, + timings, + strict=True, + ) + ) + logger.info( + "FlashInfer ReplaySSM autotune selected batch %d: %s (%s)", + batch, + tactic.name, + timing_log, + ) + del benchmark + gc.collect() + torch.cuda.empty_cache() + + if cache_changed and world.rank_in_group == 0: + try: + _save_cache(cache_path, cached_entries) + except Exception: + logger.warning( + "Could not save the FlashInfer ReplaySSM autotune cache to %s.", + cache_path, + exc_info=True, + ) + runner.flashinfer_replayssm_autotune_result = FlashInferReplaySSMAutotuneResult( + T, selected + ) + if world.rank_in_group == 0: + logger.info( + "FlashInfer ReplaySSM autotune prepared %d CUDA-graph batch " + "tactics in %.2f seconds.", + len(selected), + time.perf_counter() - tuning_started, + ) + + +@contextmanager +def use_flashinfer_replayssm_tactic_for_capture( + runner: GPUModelRunner, + batch_descriptor: BatchDescriptor, +): + result = runner.flashinfer_replayssm_autotune_result + tactic = result.tactic_for(batch_descriptor) if result is not None else None + scope = ( + use_flashinfer_replayssm_tactic(tactic) if tactic is not None else nullcontext() + ) + with scope: + yield diff --git a/vllm/model_executor/warmup/kernel_warmup.py b/vllm/model_executor/warmup/kernel_warmup.py index f4dc08844769..bdc075340502 100644 --- a/vllm/model_executor/warmup/kernel_warmup.py +++ b/vllm/model_executor/warmup/kernel_warmup.py @@ -27,6 +27,9 @@ resolve_flashinfer_autotune_file, write_flashinfer_autotune_cache, ) +from vllm.model_executor.warmup.flashinfer_replayssm_warmup import ( + flashinfer_replayssm_autotune_warmup, +) from vllm.model_executor.warmup.flashinfer_sparse_mla_warmup import ( deepseek_v4_sparse_mla_attention_warmup, flashinfer_sparse_mla_decode_autotune_warmup, @@ -194,6 +197,7 @@ def kernel_warmup(worker: "Worker", *, process_local_only: bool = False): logger.info_once("Skipping FlashInfer autotune because it is disabled.") elif has_flashinfer() and current_platform.has_device_capability(90): flashinfer_autotune(worker.model_runner) + flashinfer_replayssm_autotune_warmup(worker) # FlashInfer attention warmup # Only warmup if the model has FlashInfer attention groups diff --git a/vllm/v1/worker/gpu_model_runner.py b/vllm/v1/worker/gpu_model_runner.py index 78933bbc6920..d13029a2e513 100644 --- a/vllm/v1/worker/gpu_model_runner.py +++ b/vllm/v1/worker/gpu_model_runner.py @@ -912,6 +912,7 @@ def __init__( # Cudagraph dispatcher for runtime cudagraph dispatching. self.cudagraph_dispatcher = CudagraphDispatcher(self.vllm_config) + self.flashinfer_replayssm_autotune_result: Any | None = None self.mm_budget = ( MultiModalBudget(self.vllm_config, self.mm_registry) @@ -7046,6 +7047,10 @@ def _capture_cudagraphs( if not batch_descriptors: return + from vllm.model_executor.warmup.flashinfer_replayssm_warmup import ( + use_flashinfer_replayssm_tactic_for_capture, + ) + uniform_decode = batch_descriptors[0].uniform # Only rank 0 should print progress bar during capture @@ -7075,12 +7080,13 @@ def _capture_cudagraphs( uniform_decode=uniform_decode, ) ) - self._warmup_and_capture( - batch_desc, - cudagraph_runtime_mode=cudagraph_runtime_mode, - allow_microbatching=allow_microbatching, - profiler=profiler, - ) + with use_flashinfer_replayssm_tactic_for_capture(self, batch_desc): + self._warmup_and_capture( + batch_desc, + cudagraph_runtime_mode=cudagraph_runtime_mode, + allow_microbatching=allow_microbatching, + profiler=profiler, + ) torch.accelerator.synchronize() self.maybe_remove_all_loras(self.lora_config) From 0ec9060e791a42053f18d01c8a19ece8ad30d305 Mon Sep 17 00:00:00 2001 From: Andrii Skliar Date: Sun, 16 Aug 2026 01:01:04 -0700 Subject: [PATCH 10/33] Batch ReplaySSM tracker updates across layers Signed-off-by: Andrii Skliar --- tests/kernels/mamba/test_ssu_dispatch.py | 34 ++++++++++++ tests/v1/worker/test_utils.py | 31 +++++++++++ .../layers/mamba/mamba_mixer2.py | 53 +++++++++++++++++++ .../layers/mamba/ops/ssu_dispatch.py | 47 +++++++++++++--- vllm/v1/worker/utils.py | 17 ++++++ 5 files changed, 174 insertions(+), 8 deletions(-) diff --git a/tests/kernels/mamba/test_ssu_dispatch.py b/tests/kernels/mamba/test_ssu_dispatch.py index 1893c70070c6..757385115697 100644 --- a/tests/kernels/mamba/test_ssu_dispatch.py +++ b/tests/kernels/mamba/test_ssu_dispatch.py @@ -67,6 +67,40 @@ def test_flashinfer_replayssm_ring_tracker_lifecycle(): assert observed[32] == (15, 1) +@pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") +def test_flashinfer_replayssm_batched_ring_tracker_update(): + ring_start = torch.tensor( + [[0, 2, 4, 6], [1, 3, 5, 7], [8, 10, 12, 14]], + dtype=torch.int32, + device="cuda", + ) + prev_num_accepted = torch.tensor( + [[0, 4, 8, 12], [1, 5, 9, 13], [2, 6, 10, 14]], + dtype=torch.int32, + device="cuda", + ) + expected_ring_start = ring_start.clone() + expected_prev_num_accepted = prev_num_accepted.clone() + state_batch_indices = torch.tensor([1, 3], dtype=torch.int32, device="cuda") + + update_replayssm_ring_trackers( + ring_start, + prev_num_accepted, + state_batch_indices, + logical_window=16, + ) + for layer_index in range(ring_start.shape[0]): + update_replayssm_ring_trackers( + expected_ring_start[layer_index], + expected_prev_num_accepted[layer_index], + state_batch_indices, + logical_window=16, + ) + + torch.testing.assert_close(ring_start, expected_ring_start) + torch.testing.assert_close(prev_num_accepted, expected_prev_num_accepted) + + def _kv_cache_config_with_ssu( mamba_type: MambaAttentionBackendEnum = MambaAttentionBackendEnum.MAMBA2, ) -> KVCacheConfig: diff --git a/tests/v1/worker/test_utils.py b/tests/v1/worker/test_utils.py index a760f1dbd1ba..b53f677848f2 100644 --- a/tests/v1/worker/test_utils.py +++ b/tests/v1/worker/test_utils.py @@ -29,6 +29,9 @@ def __init__(self): self.cache_config = SimpleNamespace(mamba_cache_mode="none") self._replayssm_ring_start = torch.empty(0, dtype=torch.int32) self._replayssm_prev_num_accepted = torch.empty(0, dtype=torch.int32) + self._updates_replayssm_trackers = True + self._replayssm_ring_starts_to_update = None + self._replayssm_prev_num_accepted_to_update = None def get_state_shape(self) -> tuple[tuple[int, ...], ...]: return self._state_shapes @@ -84,6 +87,34 @@ def test_bind_kv_cache_recreates_replayssm_tracker_sidecars(): ) +def test_bind_kv_cache_batches_replayssm_tracker_updates(): + mixers = [_TestReplaySSMMixer() for _ in range(3)] + layer_names = [f"layers.{i}.mixer" for i in range(3)] + ctx = dict(zip(layer_names, mixers)) + kv_cache = { + layer_name: _packed_replayssm_cache(4) for layer_name in layer_names + } + + bind_kv_cache(kv_cache, ctx, []) + + assert len({m._replayssm_ring_start.data_ptr() for m in mixers}) == 3 + assert len({m._replayssm_prev_num_accepted.data_ptr() for m in mixers}) == 3 + ring_starts = mixers[-1]._replayssm_ring_starts_to_update + prev_num_accepted = mixers[-1]._replayssm_prev_num_accepted_to_update + assert ring_starts is not None + assert prev_num_accepted is not None + assert ring_starts.shape == (3, 4) + assert prev_num_accepted.shape == (3, 4) + for layer_index, mixer in enumerate(mixers): + assert mixer._replayssm_ring_start.data_ptr() == ring_starts[ + layer_index + ].data_ptr() + assert mixer._replayssm_prev_num_accepted.data_ptr() == prev_num_accepted[ + layer_index + ].data_ptr() + assert [m._updates_replayssm_trackers for m in mixers] == [False, False, True] + + def test_bind_kv_cache(default_vllm_config): from vllm.model_executor.layers.attention import Attention diff --git a/vllm/model_executor/layers/mamba/mamba_mixer2.py b/vllm/model_executor/layers/mamba/mamba_mixer2.py index 8478da96c9d1..f708b0d20a50 100644 --- a/vllm/model_executor/layers/mamba/mamba_mixer2.py +++ b/vllm/model_executor/layers/mamba/mamba_mixer2.py @@ -531,6 +531,9 @@ def __init__( self.kv_cache = tuple(torch.tensor([]) for _ in range(_n_state)) self._replayssm_ring_start = torch.empty(0, dtype=torch.int32) self._replayssm_prev_num_accepted = torch.empty(0, dtype=torch.int32) + self._updates_replayssm_trackers = True + self._replayssm_ring_starts_to_update: torch.Tensor | None = None + self._replayssm_prev_num_accepted_to_update: torch.Tensor | None = None self.num_spec = vllm_config.num_speculative_tokens if self.num_spec > 0: @@ -1137,6 +1140,11 @@ def conv_ssm_forward( cb_scaled=attn_metadata.cb_scaled, cumAdt_vec=attn_metadata.cumAdt_vec, cb_old=attn_metadata.cb_old, + update_trackers=self._updates_replayssm_trackers, + ring_starts_to_update=self._replayssm_ring_starts_to_update, + prev_num_accepted_to_update=( + self._replayssm_prev_num_accepted_to_update + ), ) else: selective_state_update_replayssm_triton( @@ -1229,6 +1237,51 @@ def mamba_type(self) -> MambaAttentionBackendEnum: return MambaAttentionBackendEnum.MAMBA2 +def batch_replayssm_ring_tracker_updates(mixers: list[MambaMixer2]) -> None: + """Keep per-layer cursors and update them together after the final layer.""" + if not mixers: + return + + first_ring_start = mixers[0]._replayssm_ring_start + first_prev_num_accepted = mixers[0]._replayssm_prev_num_accepted + expected = ( + first_ring_start.shape, + first_ring_start.device, + first_ring_start.dtype, + first_prev_num_accepted.shape, + first_prev_num_accepted.device, + first_prev_num_accepted.dtype, + ) + for mixer in mixers: + actual = ( + mixer._replayssm_ring_start.shape, + mixer._replayssm_ring_start.device, + mixer._replayssm_ring_start.dtype, + mixer._replayssm_prev_num_accepted.shape, + mixer._replayssm_prev_num_accepted.device, + mixer._replayssm_prev_num_accepted.dtype, + ) + if actual != expected: + raise ValueError("ReplaySSM tracker sidecars must have matching layouts") + + ring_starts = torch.zeros( + (len(mixers), *first_ring_start.shape), + dtype=first_ring_start.dtype, + device=first_ring_start.device, + ) + prev_num_accepted = torch.zeros_like(ring_starts) + for layer_index, mixer in enumerate(mixers): + mixer._replayssm_ring_start = ring_starts[layer_index] + mixer._replayssm_prev_num_accepted = prev_num_accepted[layer_index] + mixer._updates_replayssm_trackers = False + mixer._replayssm_ring_starts_to_update = None + mixer._replayssm_prev_num_accepted_to_update = None + + mixers[-1]._updates_replayssm_trackers = True + mixers[-1]._replayssm_ring_starts_to_update = ring_starts + mixers[-1]._replayssm_prev_num_accepted_to_update = prev_num_accepted + + def mamba_mixer2( projected_states: torch.Tensor, output: torch.Tensor, diff --git a/vllm/model_executor/layers/mamba/ops/ssu_dispatch.py b/vllm/model_executor/layers/mamba/ops/ssu_dispatch.py index 95d89b6283e8..9e216441e07d 100644 --- a/vllm/model_executor/layers/mamba/ops/ssu_dispatch.py +++ b/vllm/model_executor/layers/mamba/ops/ssu_dispatch.py @@ -75,17 +75,20 @@ def _update_replayssm_ring_trackers_kernel( prev_num_accepted, state_batch_indices, n_slots, + tracker_stride, logical_window: tl.constexpr, ring_buffer_len: tl.constexpr, pad_slot_id: tl.constexpr, BLOCK: tl.constexpr, ) -> None: offsets = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK) + tracker_offset = tl.program_id(1) * tracker_stride mask = offsets < n_slots slots = tl.load(state_batch_indices + offsets, mask=mask, other=pad_slot_id) valid = mask & (slots != pad_slot_id) - prev = tl.load(prev_num_accepted + slots, mask=valid, other=0) - start = tl.load(ring_start + slots, mask=valid, other=0) + tracker_slots = tracker_offset + slots + prev = tl.load(prev_num_accepted + tracker_slots, mask=valid, other=0) + start = tl.load(ring_start + tracker_slots, mask=valid, other=0) must_checkpoint = prev + 1 > logical_window next_start = tl.where( must_checkpoint, @@ -93,8 +96,8 @@ def _update_replayssm_ring_trackers_kernel( start, ) next_prev = tl.where(must_checkpoint, 1, prev + 1) - tl.store(ring_start + slots, next_start, mask=valid) - tl.store(prev_num_accepted + slots, next_prev, mask=valid) + tl.store(ring_start + tracker_slots, next_start, mask=valid) + tl.store(prev_num_accepted + tracker_slots, next_prev, mask=valid) @triton.jit @@ -121,16 +124,27 @@ def update_replayssm_ring_trackers( logical_window: int, pad_slot_id: int = NULL_BLOCK_ID, ) -> None: + if ring_start.shape != prev_num_accepted.shape: + raise ValueError("ReplaySSM tracker tensors must have matching shapes") + if ring_start.dim() not in (1, 2): + raise ValueError("ReplaySSM tracker tensors must be one- or two-dimensional") + if not ring_start.is_contiguous() or not prev_num_accepted.is_contiguous(): + raise ValueError("ReplaySSM tracker tensors must be contiguous") state_batch_indices = state_batch_indices.reshape(-1) n_slots = state_batch_indices.numel() if n_slots == 0: return + num_trackers = 1 if ring_start.dim() == 1 else ring_start.shape[0] + tracker_stride = ring_start.shape[-1] block = 128 - _update_replayssm_ring_trackers_kernel[(triton.cdiv(n_slots, block),)]( + _update_replayssm_ring_trackers_kernel[ + (triton.cdiv(n_slots, block), num_trackers) + ]( ring_start, prev_num_accepted, state_batch_indices, n_slots, + tracker_stride, logical_window, logical_window + 1, pad_slot_id, @@ -518,6 +532,9 @@ def __call__( cumAdt_vec: torch.Tensor | None = None, cb_old: torch.Tensor | None = None, algorithm: str | None = None, + update_trackers: bool = True, + ring_starts_to_update: torch.Tensor | None = None, + prev_num_accepted_to_update: torch.Tensor | None = None, ) -> torch.Tensor: # AR decode currently passes (batch, nheads, dim); checkpointing_ssu # expects a predicted-token axis T. Unsqueeze T=1 here. @@ -564,10 +581,18 @@ def __call__( cb_old=cb_old, algorithm=self._algorithm if algorithm is None else algorithm, ) - if indices is not None: + if update_trackers and indices is not None: + if (ring_starts_to_update is None) != ( + prev_num_accepted_to_update is None + ): + raise ValueError("ReplaySSM tracker update tensors must be paired") update_replayssm_ring_trackers( - ring_start, - prev_num_accepted_tokens, + ring_start + if ring_starts_to_update is None + else ring_starts_to_update, + prev_num_accepted_tokens + if prev_num_accepted_to_update is None + else prev_num_accepted_to_update, indices, logical_window=x_cache.size(2) - 1, pad_slot_id=null_block_id, @@ -734,6 +759,9 @@ def selective_state_update_replayssm_flashinfer( cumAdt_vec: torch.Tensor | None = None, cb_old: torch.Tensor | None = None, algorithm: str | None = None, + update_trackers: bool = True, + ring_starts_to_update: torch.Tensor | None = None, + prev_num_accepted_to_update: torch.Tensor | None = None, ) -> torch.Tensor: """FlashInfer ReplaySSM decode (``checkpointing_ssu``).""" backend = get_replayssm_backend() @@ -765,6 +793,9 @@ def selective_state_update_replayssm_flashinfer( cumAdt_vec=cumAdt_vec, cb_old=cb_old, algorithm=algorithm, + update_trackers=update_trackers, + ring_starts_to_update=ring_starts_to_update, + prev_num_accepted_to_update=prev_num_accepted_to_update, ) diff --git a/vllm/v1/worker/utils.py b/vllm/v1/worker/utils.py index 8ee389df3c1f..8811b2e4bfdc 100644 --- a/vllm/v1/worker/utils.py +++ b/vllm/v1/worker/utils.py @@ -11,6 +11,7 @@ import torch from vllm.config import CacheConfig, VllmConfig +from vllm.config.mamba import MambaBackendEnum from vllm.logger import init_logger from vllm.model_executor.layers.attention import Attention from vllm.model_executor.models.interfaces import MultiModalEmbeddings @@ -590,6 +591,7 @@ def bind_kv_cache( for layer_name in kv_caches: index2name[extract_layer_index(layer_name, num_attn_module)].append(layer_name) + ordered_layer_names: list[str] = [] for layer_index in sorted(index2name.keys()): layer_names = index2name[layer_index] if len(layer_names) > 1: @@ -603,6 +605,7 @@ def bind_kv_cache( current_platform.check_runner_kv_caches_multi_layer() for layer_name in layer_names: runner_kv_caches.append(kv_caches[layer_name]) + ordered_layer_names.append(layer_name) # Bind kv_caches to forward context. Each layer's bind_kv_cache unpacks # its raw allocation into the per-layer view(s) it needs (e.g. Mamba @@ -611,6 +614,20 @@ def bind_kv_cache( for layer_name, kv_cache in kv_caches.items(): forward_context[layer_name].bind_kv_cache(kv_cache) + from vllm.model_executor.layers.mamba.mamba_mixer2 import ( + MambaMixer2, + batch_replayssm_ring_tracker_updates, + ) + + replayssm_mixers = [ + layer + for layer_name in ordered_layer_names + if isinstance((layer := forward_context[layer_name]), MambaMixer2) + and layer.use_replayssm + and layer.mamba_config.backend == MambaBackendEnum.FLASHINFER + ] + batch_replayssm_ring_tracker_updates(replayssm_mixers) + def copy_kv_cache_blocks_inplace( kv_caches: Iterable[torch.Tensor], From f93b18d6c813c3361f093a2144039cf3cac5e825 Mon Sep 17 00:00:00 2001 From: Andrii Skliar Date: Sun, 16 Aug 2026 01:01:39 -0700 Subject: [PATCH 11/33] Autotune ReplaySSM precompute launch geometry Signed-off-by: Andrii Skliar --- .../test_flashinfer_replayssm_warmup.py | 55 ++++++- .../layers/mamba/ops/ssu_dispatch.py | 19 ++- .../warmup/flashinfer_replayssm_warmup.py | 139 +++++++++++++++--- 3 files changed, 190 insertions(+), 23 deletions(-) diff --git a/tests/model_executor/test_flashinfer_replayssm_warmup.py b/tests/model_executor/test_flashinfer_replayssm_warmup.py index 64330800bc92..af6bc40d50b6 100644 --- a/tests/model_executor/test_flashinfer_replayssm_warmup.py +++ b/tests/model_executor/test_flashinfer_replayssm_warmup.py @@ -19,8 +19,10 @@ from vllm.model_executor.warmup.flashinfer_replayssm_warmup import ( FLASHINFER_REPLAYSSM_TUNING_CANDIDATES, FlashInferReplaySSMAutotuneResult, + _expanded_precompute_tactics, _load_cache, _make_cache_key, + _precompute_heads_per_cta_candidates, _ReplaySSMBenchmark, _save_cache, _select_fastest, @@ -36,6 +38,7 @@ def _fake_backend() -> FlashInferReplaySSMBackend: backend._mamba_config = MambaConfig(backend=MambaBackendEnum.FLASHINFER) backend._kernel = Mock(return_value=torch.empty(1)) backend._algorithm = "auto" + backend._precompute_heads_per_cta = 0 return backend @@ -57,6 +60,39 @@ def test_replayssm_tuning_candidates_and_deterministic_selection(): timings = [3.0, 2.0, 2.0, 4.0, 5.0, 6.0, 7.0, 8.0, 9.0, 10.0, 11.0, 12.0] assert _select_fastest(timings).name == "monolith" assert _select_fastest([float("inf")] * len(timings)) is None + near_tie = [1.0, 0.9995] + [float("inf")] * (len(timings) - 2) + assert _select_fastest(near_tie).name == "auto" + + +def test_replayssm_precompute_candidates_expand_top_main_tactics(): + assert _precompute_heads_per_cta_candidates(16) == (1, 2, 4, 8, 16) + timings = [float("inf")] * len(FLASHINFER_REPLAYSSM_TUNING_CANDIDATES) + timings[4] = 1.0 + timings[10] = 2.0 + timings[2] = 3.0 + + assert [ + tactic.name for tactic in _expanded_precompute_tactics(timings, 16) + ] == [ + "two_kernel_s1_c4_h1", + "two_kernel_s1_c4_h2", + "two_kernel_s1_c4_h4", + "two_kernel_s1_c4_h8", + "two_kernel_s1_c4_h16", + "two_kernel_s2_c8_h1", + "two_kernel_s2_c8_h2", + "two_kernel_s2_c8_h4", + "two_kernel_s2_c8_h8", + "two_kernel_s2_c8_h16", + ] + + +def test_replayssm_tactic_validates_precompute_geometry(): + assert FlashInferReplaySSMTactic( + "two-kernel", 1, 4, precompute_heads_per_cta=8 + ).name == "two_kernel_s1_c4_h8" + with pytest.raises(ValueError, match="does not accept precompute"): + FlashInferReplaySSMTactic("monolith", precompute_heads_per_cta=8) def test_replayssm_tuning_key_distinguishes_batch_and_T(): @@ -67,8 +103,8 @@ def test_replayssm_tuning_key_distinguishes_batch_and_T(): def test_replayssm_autotune_cache_round_trip(tmp_path): path = tmp_path / "replayssm.json" - _save_cache(path, {"key": "two_kernel_s2_c16", "bad": "unknown"}) - assert _load_cache(path) == {"key": "two_kernel_s2_c16"} + _save_cache(path, {"key": "two_kernel_s2_c16_h8", "bad": "unknown"}) + assert _load_cache(path) == {"key": "two_kernel_s2_c16_h8"} @pytest.mark.parametrize("payload", ["null", "[]", "1"]) @@ -139,17 +175,21 @@ def test_replayssm_tactic_scope_restores_algorithm_and_environment( monkeypatch.setenv(_STAGES_ENV, "7") monkeypatch.setenv(_CTAS_ENV, "11") - tactic = FlashInferReplaySSMTactic("two-kernel", 2, 16) + tactic = FlashInferReplaySSMTactic( + "two-kernel", 2, 16, precompute_heads_per_cta=8 + ) with ( pytest.raises(RuntimeError, match="sentinel"), use_flashinfer_replayssm_tactic(tactic), ): assert backend._algorithm == "two-kernel" + assert backend._precompute_heads_per_cta == 8 assert ssu_dispatch.os.environ[_STAGES_ENV] == "2" assert ssu_dispatch.os.environ[_CTAS_ENV] == "16" raise RuntimeError("sentinel") assert backend._algorithm == "auto" + assert backend._precompute_heads_per_cta == 0 assert ssu_dispatch.os.environ[_STAGES_ENV] == "7" assert ssu_dispatch.os.environ[_CTAS_ENV] == "11" @@ -163,7 +203,11 @@ def test_replayssm_backend_uses_scoped_algorithm(monkeypatch): monkeypatch.setattr(ssu_dispatch, "_replayssm_backend", backend) tensor = torch.empty(1) - with use_flashinfer_replayssm_tactic(FlashInferReplaySSMTactic("monolith")): + with use_flashinfer_replayssm_tactic( + FlashInferReplaySSMTactic( + "two-kernel", 1, 4, precompute_heads_per_cta=8 + ) + ): backend( tensor, tensor, @@ -179,7 +223,8 @@ def test_replayssm_backend_uses_scoped_algorithm(monkeypatch): tensor, ) - assert backend._kernel.call_args.kwargs["algorithm"] == "monolith" + assert backend._kernel.call_args.kwargs["algorithm"] == "two-kernel" + assert backend._kernel.call_args.kwargs["precompute_heads_per_cta"] == 8 def test_replayssm_capture_tactic_uses_request_batch_not_tokens(): diff --git a/vllm/model_executor/layers/mamba/ops/ssu_dispatch.py b/vllm/model_executor/layers/mamba/ops/ssu_dispatch.py index 9e216441e07d..a2c197a3191d 100644 --- a/vllm/model_executor/layers/mamba/ops/ssu_dispatch.py +++ b/vllm/model_executor/layers/mamba/ops/ssu_dispatch.py @@ -42,6 +42,7 @@ class FlashInferReplaySSMTactic: algorithm: str pipeline_stages: int | None = None ctas_per_sm: int | None = None + precompute_heads_per_cta: int = 0 def __post_init__(self) -> None: if self.algorithm not in {"auto", "monolith", "two-kernel"}: @@ -54,16 +55,27 @@ def __post_init__(self) -> None: raise ValueError("two-kernel requires pipeline_stages in {1, 2}") if self.ctas_per_sm is None or self.ctas_per_sm <= 0: raise ValueError("two-kernel requires a positive ctas_per_sm") + if self.precompute_heads_per_cta < 0: + raise ValueError( + "two-kernel requires non-negative precompute_heads_per_cta" + ) elif has_launch_config: raise ValueError( f"{self.algorithm} does not accept pipeline or CTA settings" ) + elif self.precompute_heads_per_cta != 0: + raise ValueError( + f"{self.algorithm} does not accept precompute_heads_per_cta" + ) @property def name(self) -> str: if self.algorithm != "two-kernel": return self.algorithm - return f"two_kernel_s{self.pipeline_stages}_c{self.ctas_per_sm}" + name = f"two_kernel_s{self.pipeline_stages}_c{self.ctas_per_sm}" + if self.precompute_heads_per_cta: + name += f"_h{self.precompute_heads_per_cta}" + return name FLASHINFER_REPLAYSSM_AUTO_TACTIC = FlashInferReplaySSMTactic("auto") @@ -503,6 +515,7 @@ def __init__(self, mamba_config: MambaConfig): ) from e self._kernel = _fi_checkpointing_ssu self._algorithm = FLASHINFER_REPLAYSSM_AUTO_TACTIC.algorithm + self._precompute_heads_per_cta = 0 @property def name(self) -> str: @@ -579,6 +592,7 @@ def __call__( cb_scaled=cb_scaled, cumAdt_vec=cumAdt_vec, cb_old=cb_old, + precompute_heads_per_cta=self._precompute_heads_per_cta, algorithm=self._algorithm if algorithm is None else algorithm, ) if update_trackers and indices is not None: @@ -611,9 +625,11 @@ def use_flashinfer_replayssm_tactic( return old_algorithm = backend._algorithm + old_precompute_heads_per_cta = backend._precompute_heads_per_cta old_stages = os.environ.get(_FLASHINFER_SSU_PIPELINE_STAGES_ENV) old_ctas = os.environ.get(_FLASHINFER_SSU_CTA_PER_SM_ENV) backend._algorithm = tactic.algorithm + backend._precompute_heads_per_cta = tactic.precompute_heads_per_cta try: if tactic.algorithm == "two-kernel": assert tactic.pipeline_stages is not None @@ -628,6 +644,7 @@ def use_flashinfer_replayssm_tactic( yield finally: backend._algorithm = old_algorithm + backend._precompute_heads_per_cta = old_precompute_heads_per_cta if old_stages is None: os.environ.pop(_FLASHINFER_SSU_PIPELINE_STAGES_ENV, None) else: diff --git a/vllm/model_executor/warmup/flashinfer_replayssm_warmup.py b/vllm/model_executor/warmup/flashinfer_replayssm_warmup.py index c8a9bd276440..8e39b95f15f8 100644 --- a/vllm/model_executor/warmup/flashinfer_replayssm_warmup.py +++ b/vllm/model_executor/warmup/flashinfer_replayssm_warmup.py @@ -10,7 +10,7 @@ import statistics import time from contextlib import contextmanager, nullcontext -from dataclasses import dataclass +from dataclasses import dataclass, replace from pathlib import Path from typing import TYPE_CHECKING, Any @@ -38,9 +38,11 @@ logger = init_logger(__name__) -_CACHE_SCHEMA_VERSION = 1 -_TUNING_FIXTURE = "t1_mixed_history_cycle_v2" +_CACHE_SCHEMA_VERSION = 3 +_TUNING_FIXTURE = "t1_mixed_history_cycle_v4" _CACHE_FILE_NAME = "replayssm_autotune_configs.json" +_PRECOMPUTE_MAIN_TACTIC_COUNT = 2 +_RELATIVE_TIE_TOLERANCE = 0.001 FLASHINFER_REPLAYSSM_TUNING_CANDIDATES = ( FLASHINFER_REPLAYSSM_AUTO_TACTIC, @@ -61,6 +63,56 @@ } +def _tactic_from_name(name: str) -> FlashInferReplaySSMTactic | None: + if tactic := _TACTICS_BY_NAME.get(name): + return tactic + for base in FLASHINFER_REPLAYSSM_TUNING_CANDIDATES: + if base.algorithm != "two-kernel": + continue + prefix = f"{base.name}_h" + if name.startswith(prefix): + try: + heads_per_cta = int(name.removeprefix(prefix)) + return replace(base, precompute_heads_per_cta=heads_per_cta) + except ValueError: + return None + return None + + +def _precompute_heads_per_cta_candidates(heads_per_group: int) -> tuple[int, ...]: + """Return the distinct HEADS_PER_GROUP >> k launch geometries.""" + if heads_per_group <= 0: + raise ValueError("heads_per_group must be positive") + candidates = set() + value = heads_per_group + while value: + candidates.add(value) + value >>= 1 + return tuple(sorted(candidates)) + + +def _expanded_precompute_tactics( + main_timings: list[float], heads_per_group: int +) -> tuple[FlashInferReplaySSMTactic, ...]: + ranked = sorted( + ( + index + for index, tactic in enumerate( + FLASHINFER_REPLAYSSM_TUNING_CANDIDATES + ) + if tactic.algorithm == "two-kernel" + and math.isfinite(main_timings[index]) + ), + key=lambda index: (main_timings[index], index), + )[:_PRECOMPUTE_MAIN_TACTIC_COUNT] + return tuple( + replace(base, precompute_heads_per_cta=heads_per_cta) + for index in ranked + for base in (FLASHINFER_REPLAYSSM_TUNING_CANDIDATES[index],) + for heads_per_cta in _precompute_heads_per_cta_candidates(heads_per_group) + ) + + @dataclass class FlashInferReplaySSMAutotuneResult: spec_query_len: int @@ -86,23 +138,39 @@ def _make_cache_key(fingerprint: dict[str, Any], batch: int, T: int) -> str: "batch_sequences": batch, "spec_query_len": T, "fixture": _TUNING_FIXTURE, - "candidate_schema": [ - tactic.name for tactic in FLASHINFER_REPLAYSSM_TUNING_CANDIDATES - ], + "candidate_schema": { + "main": [ + tactic.name + for tactic in FLASHINFER_REPLAYSSM_TUNING_CANDIDATES + ], + "precompute": "top2_main_x_heads_per_group_shift_chain", + "relative_tie_tolerance": _RELATIVE_TIE_TOLERANCE, + }, }, sort_keys=True, separators=(",", ":"), ) -def _select_fastest(timings: list[float]) -> FlashInferReplaySSMTactic | None: - if len(timings) != len(FLASHINFER_REPLAYSSM_TUNING_CANDIDATES): +def _select_fastest( + timings: list[float], + candidates: tuple[FlashInferReplaySSMTactic, ...] = ( + FLASHINFER_REPLAYSSM_TUNING_CANDIDATES + ), +) -> FlashInferReplaySSMTactic | None: + if len(timings) != len(candidates): raise ValueError("one timing is required for each ReplaySSM tactic") finite = [i for i, timing in enumerate(timings) if math.isfinite(timing)] if not finite: return None - winner = min(finite, key=lambda i: (timings[i], i)) - return FLASHINFER_REPLAYSSM_TUNING_CANDIDATES[winner] + best_timing = min(timings[i] for i in finite) + tied = [ + i + for i in finite + if timings[i] <= best_timing * (1 + _RELATIVE_TIE_TOLERANCE) + ] + winner = min(tied) + return candidates[winner] def _load_cache(path: Path) -> dict[str, str]: @@ -120,7 +188,9 @@ def _load_cache(path: Path) -> dict[str, str]: return { key: value for key, value in entries.items() - if isinstance(key, str) and isinstance(value, str) and value in _TACTICS_BY_NAME + if isinstance(key, str) + and isinstance(value, str) + and _tactic_from_name(value) is not None } @@ -338,7 +408,7 @@ def call(self, tactic: FlashInferReplaySSMTactic) -> None: cb_scaled=self.cb_scaled, cumAdt_vec=self.cumAdt_vec, cb_old=self.cb_old, - precompute_heads_per_cta=0, + precompute_heads_per_cta=tactic.precompute_heads_per_cta, algorithm=tactic.algorithm, ) update_replayssm_ring_trackers( @@ -531,8 +601,11 @@ def flashinfer_replayssm_autotune_warmup(worker: Worker) -> None: for batch in batches: cache_key = _make_cache_key(fingerprint, batch, T) cached_name = cached_entries.get(cache_key) - if cached_name in _TACTICS_BY_NAME: - tactic = _TACTICS_BY_NAME[cached_name] + cached_tactic = ( + _tactic_from_name(cached_name) if cached_name is not None else None + ) + if cached_tactic is not None: + tactic = cached_tactic selected[batch] = tactic if world.rank_in_group == 0: logger.info( @@ -582,7 +655,39 @@ def flashinfer_replayssm_autotune_warmup(worker: Worker) -> None: exc_info=True, ) timings = _aggregate_timings(local_timings) - tactic = _select_fastest(timings) + _, _, _, _, B_cache, *_ = layer.kv_cache + heads_per_group = layer.kv_cache[1].shape[1] // B_cache.shape[1] + expanded_candidates = _expanded_precompute_tactics( + timings, heads_per_group + ) + local_expanded_timings = [float("inf")] * len(expanded_candidates) + expanded_count = len(expanded_candidates) + if expanded_count: + expanded_shift = batch % expanded_count + expanded_order = tuple(range(expanded_count)) + expanded_order = ( + expanded_order[expanded_shift:] + + expanded_order[:expanded_shift] + ) + for candidate_index in expanded_order: + expanded_tactic = expanded_candidates[candidate_index] + try: + local_expanded_timings[candidate_index] = benchmark.benchmark( + expanded_tactic + ) + except Exception: + logger.warning( + "FlashInfer ReplaySSM tactic %s failed for batch %d.", + expanded_tactic.name, + batch, + exc_info=True, + ) + expanded_timings = _aggregate_timings(local_expanded_timings) + all_candidates = ( + FLASHINFER_REPLAYSSM_TUNING_CANDIDATES + expanded_candidates + ) + all_timings = timings + expanded_timings + tactic = _select_fastest(all_timings, all_candidates) if tactic is None: logger.warning_once( "Every FlashInfer ReplaySSM tactic failed for batch %d; " @@ -600,8 +705,8 @@ def flashinfer_replayssm_autotune_warmup(worker: Worker) -> None: timing_log = ", ".join( f"{candidate.name}={timing:.6f}ms" for candidate, timing in zip( - FLASHINFER_REPLAYSSM_TUNING_CANDIDATES, - timings, + all_candidates, + all_timings, strict=True, ) ) From 3247272aeac248d8018429fffad5fc63fd0c0a2e Mon Sep 17 00:00:00 2001 From: Andrii Skliar Date: Sun, 16 Aug 2026 03:10:18 -0700 Subject: [PATCH 12/33] Share ReplaySSM trackers by KV cache group Signed-off-by: Andrii Skliar --- tests/kernels/mamba/test_ssu_dispatch.py | 34 ------- tests/v1/worker/test_utils.py | 52 ++++++----- .../layers/mamba/mamba_mixer2.py | 93 +++++++++---------- .../layers/mamba/ops/ssu_dispatch.py | 40 ++------ vllm/v1/worker/gpu/attn_utils.py | 8 +- vllm/v1/worker/gpu_model_runner.py | 1 + vllm/v1/worker/utils.py | 33 +++++-- 7 files changed, 116 insertions(+), 145 deletions(-) diff --git a/tests/kernels/mamba/test_ssu_dispatch.py b/tests/kernels/mamba/test_ssu_dispatch.py index 757385115697..1893c70070c6 100644 --- a/tests/kernels/mamba/test_ssu_dispatch.py +++ b/tests/kernels/mamba/test_ssu_dispatch.py @@ -67,40 +67,6 @@ def test_flashinfer_replayssm_ring_tracker_lifecycle(): assert observed[32] == (15, 1) -@pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") -def test_flashinfer_replayssm_batched_ring_tracker_update(): - ring_start = torch.tensor( - [[0, 2, 4, 6], [1, 3, 5, 7], [8, 10, 12, 14]], - dtype=torch.int32, - device="cuda", - ) - prev_num_accepted = torch.tensor( - [[0, 4, 8, 12], [1, 5, 9, 13], [2, 6, 10, 14]], - dtype=torch.int32, - device="cuda", - ) - expected_ring_start = ring_start.clone() - expected_prev_num_accepted = prev_num_accepted.clone() - state_batch_indices = torch.tensor([1, 3], dtype=torch.int32, device="cuda") - - update_replayssm_ring_trackers( - ring_start, - prev_num_accepted, - state_batch_indices, - logical_window=16, - ) - for layer_index in range(ring_start.shape[0]): - update_replayssm_ring_trackers( - expected_ring_start[layer_index], - expected_prev_num_accepted[layer_index], - state_batch_indices, - logical_window=16, - ) - - torch.testing.assert_close(ring_start, expected_ring_start) - torch.testing.assert_close(prev_num_accepted, expected_prev_num_accepted) - - def _kv_cache_config_with_ssu( mamba_type: MambaAttentionBackendEnum = MambaAttentionBackendEnum.MAMBA2, ) -> KVCacheConfig: diff --git a/tests/v1/worker/test_utils.py b/tests/v1/worker/test_utils.py index b53f677848f2..02417c70c2c5 100644 --- a/tests/v1/worker/test_utils.py +++ b/tests/v1/worker/test_utils.py @@ -27,11 +27,10 @@ def __init__(self): self.use_replayssm = True self.mamba_config = MambaConfig(backend=MambaBackendEnum.FLASHINFER) self.cache_config = SimpleNamespace(mamba_cache_mode="none") + self.replayssm_buffer_len = 16 self._replayssm_ring_start = torch.empty(0, dtype=torch.int32) self._replayssm_prev_num_accepted = torch.empty(0, dtype=torch.int32) self._updates_replayssm_trackers = True - self._replayssm_ring_starts_to_update = None - self._replayssm_prev_num_accepted_to_update = None def get_state_shape(self) -> tuple[tuple[int, ...], ...]: return self._state_shapes @@ -87,32 +86,41 @@ def test_bind_kv_cache_recreates_replayssm_tracker_sidecars(): ) -def test_bind_kv_cache_batches_replayssm_tracker_updates(): +def test_bind_kv_cache_shares_replayssm_trackers_by_cache_group(): mixers = [_TestReplaySSMMixer() for _ in range(3)] layer_names = [f"layers.{i}.mixer" for i in range(3)] ctx = dict(zip(layer_names, mixers)) kv_cache = { - layer_name: _packed_replayssm_cache(4) for layer_name in layer_names + layer_names[0]: _packed_replayssm_cache(4), + layer_names[1]: _packed_replayssm_cache(4), + layer_names[2]: _packed_replayssm_cache(4), } + kv_cache_groups = [ + SimpleNamespace(layer_names=[layer_names[0], layer_names[2]]), + SimpleNamespace(layer_names=[layer_names[1]]), + ] + + bind_kv_cache(kv_cache, ctx, [], kv_cache_groups=kv_cache_groups) - bind_kv_cache(kv_cache, ctx, []) - - assert len({m._replayssm_ring_start.data_ptr() for m in mixers}) == 3 - assert len({m._replayssm_prev_num_accepted.data_ptr() for m in mixers}) == 3 - ring_starts = mixers[-1]._replayssm_ring_starts_to_update - prev_num_accepted = mixers[-1]._replayssm_prev_num_accepted_to_update - assert ring_starts is not None - assert prev_num_accepted is not None - assert ring_starts.shape == (3, 4) - assert prev_num_accepted.shape == (3, 4) - for layer_index, mixer in enumerate(mixers): - assert mixer._replayssm_ring_start.data_ptr() == ring_starts[ - layer_index - ].data_ptr() - assert mixer._replayssm_prev_num_accepted.data_ptr() == prev_num_accepted[ - layer_index - ].data_ptr() - assert [m._updates_replayssm_trackers for m in mixers] == [False, False, True] + assert ( + mixers[0]._replayssm_ring_start.data_ptr() + == mixers[2]._replayssm_ring_start.data_ptr() + ) + assert ( + mixers[0]._replayssm_prev_num_accepted.data_ptr() + == mixers[2]._replayssm_prev_num_accepted.data_ptr() + ) + assert ( + mixers[1]._replayssm_ring_start.data_ptr() + != mixers[0]._replayssm_ring_start.data_ptr() + ) + assert ( + mixers[1]._replayssm_prev_num_accepted.data_ptr() + != mixers[0]._replayssm_prev_num_accepted.data_ptr() + ) + assert mixers[0]._replayssm_ring_start.shape == (4,) + assert mixers[0]._replayssm_prev_num_accepted.shape == (4,) + assert [m._updates_replayssm_trackers for m in mixers] == [False, True, True] def test_bind_kv_cache(default_vllm_config): diff --git a/vllm/model_executor/layers/mamba/mamba_mixer2.py b/vllm/model_executor/layers/mamba/mamba_mixer2.py index f708b0d20a50..8201b3f69f34 100644 --- a/vllm/model_executor/layers/mamba/mamba_mixer2.py +++ b/vllm/model_executor/layers/mamba/mamba_mixer2.py @@ -532,8 +532,6 @@ def __init__( self._replayssm_ring_start = torch.empty(0, dtype=torch.int32) self._replayssm_prev_num_accepted = torch.empty(0, dtype=torch.int32) self._updates_replayssm_trackers = True - self._replayssm_ring_starts_to_update: torch.Tensor | None = None - self._replayssm_prev_num_accepted_to_update: torch.Tensor | None = None self.num_spec = vllm_config.num_speculative_tokens if self.num_spec > 0: @@ -1141,10 +1139,6 @@ def conv_ssm_forward( cumAdt_vec=attn_metadata.cumAdt_vec, cb_old=attn_metadata.cb_old, update_trackers=self._updates_replayssm_trackers, - ring_starts_to_update=self._replayssm_ring_starts_to_update, - prev_num_accepted_to_update=( - self._replayssm_prev_num_accepted_to_update - ), ) else: selective_state_update_replayssm_triton( @@ -1237,49 +1231,52 @@ def mamba_type(self) -> MambaAttentionBackendEnum: return MambaAttentionBackendEnum.MAMBA2 -def batch_replayssm_ring_tracker_updates(mixers: list[MambaMixer2]) -> None: - """Keep per-layer cursors and update them together after the final layer.""" - if not mixers: - return - - first_ring_start = mixers[0]._replayssm_ring_start - first_prev_num_accepted = mixers[0]._replayssm_prev_num_accepted - expected = ( - first_ring_start.shape, - first_ring_start.device, - first_ring_start.dtype, - first_prev_num_accepted.shape, - first_prev_num_accepted.device, - first_prev_num_accepted.dtype, - ) - for mixer in mixers: - actual = ( - mixer._replayssm_ring_start.shape, - mixer._replayssm_ring_start.device, - mixer._replayssm_ring_start.dtype, - mixer._replayssm_prev_num_accepted.shape, - mixer._replayssm_prev_num_accepted.device, - mixer._replayssm_prev_num_accepted.dtype, +def share_replayssm_ring_trackers( + mixer_groups: list[list[MambaMixer2]], +) -> None: + """Share ring cursors within each cache-slot index namespace. + + Layers backed by one KV-cache group use the same physical block indices and + can therefore share cursors. Different KV-cache groups may assign different + block indices to the same request and must keep separate cursor tensors. + The final local layer in each group advances its cursors after every layer + in that group has consumed the previous values. + """ + for mixers in mixer_groups: + if not mixers: + continue + + first_ring_start = mixers[0]._replayssm_ring_start + first_prev_num_accepted = mixers[0]._replayssm_prev_num_accepted + expected = ( + first_ring_start.shape, + first_ring_start.device, + first_ring_start.dtype, + first_prev_num_accepted.shape, + first_prev_num_accepted.device, + first_prev_num_accepted.dtype, + mixers[0].replayssm_buffer_len, ) - if actual != expected: - raise ValueError("ReplaySSM tracker sidecars must have matching layouts") - - ring_starts = torch.zeros( - (len(mixers), *first_ring_start.shape), - dtype=first_ring_start.dtype, - device=first_ring_start.device, - ) - prev_num_accepted = torch.zeros_like(ring_starts) - for layer_index, mixer in enumerate(mixers): - mixer._replayssm_ring_start = ring_starts[layer_index] - mixer._replayssm_prev_num_accepted = prev_num_accepted[layer_index] - mixer._updates_replayssm_trackers = False - mixer._replayssm_ring_starts_to_update = None - mixer._replayssm_prev_num_accepted_to_update = None - - mixers[-1]._updates_replayssm_trackers = True - mixers[-1]._replayssm_ring_starts_to_update = ring_starts - mixers[-1]._replayssm_prev_num_accepted_to_update = prev_num_accepted + for mixer in mixers: + actual = ( + mixer._replayssm_ring_start.shape, + mixer._replayssm_ring_start.device, + mixer._replayssm_ring_start.dtype, + mixer._replayssm_prev_num_accepted.shape, + mixer._replayssm_prev_num_accepted.device, + mixer._replayssm_prev_num_accepted.dtype, + mixer.replayssm_buffer_len, + ) + if actual != expected: + raise ValueError( + "ReplaySSM tracker sidecars must have matching layouts" + ) + + mixer._replayssm_ring_start = first_ring_start + mixer._replayssm_prev_num_accepted = first_prev_num_accepted + mixer._updates_replayssm_trackers = False + + mixers[-1]._updates_replayssm_trackers = True def mamba_mixer2( diff --git a/vllm/model_executor/layers/mamba/ops/ssu_dispatch.py b/vllm/model_executor/layers/mamba/ops/ssu_dispatch.py index a2c197a3191d..34f09115466e 100644 --- a/vllm/model_executor/layers/mamba/ops/ssu_dispatch.py +++ b/vllm/model_executor/layers/mamba/ops/ssu_dispatch.py @@ -87,20 +87,17 @@ def _update_replayssm_ring_trackers_kernel( prev_num_accepted, state_batch_indices, n_slots, - tracker_stride, logical_window: tl.constexpr, ring_buffer_len: tl.constexpr, pad_slot_id: tl.constexpr, BLOCK: tl.constexpr, ) -> None: offsets = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK) - tracker_offset = tl.program_id(1) * tracker_stride mask = offsets < n_slots slots = tl.load(state_batch_indices + offsets, mask=mask, other=pad_slot_id) valid = mask & (slots != pad_slot_id) - tracker_slots = tracker_offset + slots - prev = tl.load(prev_num_accepted + tracker_slots, mask=valid, other=0) - start = tl.load(ring_start + tracker_slots, mask=valid, other=0) + prev = tl.load(prev_num_accepted + slots, mask=valid, other=0) + start = tl.load(ring_start + slots, mask=valid, other=0) must_checkpoint = prev + 1 > logical_window next_start = tl.where( must_checkpoint, @@ -108,8 +105,8 @@ def _update_replayssm_ring_trackers_kernel( start, ) next_prev = tl.where(must_checkpoint, 1, prev + 1) - tl.store(ring_start + tracker_slots, next_start, mask=valid) - tl.store(prev_num_accepted + tracker_slots, next_prev, mask=valid) + tl.store(ring_start + slots, next_start, mask=valid) + tl.store(prev_num_accepted + slots, next_prev, mask=valid) @triton.jit @@ -138,25 +135,20 @@ def update_replayssm_ring_trackers( ) -> None: if ring_start.shape != prev_num_accepted.shape: raise ValueError("ReplaySSM tracker tensors must have matching shapes") - if ring_start.dim() not in (1, 2): - raise ValueError("ReplaySSM tracker tensors must be one- or two-dimensional") + if ring_start.dim() != 1: + raise ValueError("ReplaySSM tracker tensors must be one-dimensional") if not ring_start.is_contiguous() or not prev_num_accepted.is_contiguous(): raise ValueError("ReplaySSM tracker tensors must be contiguous") state_batch_indices = state_batch_indices.reshape(-1) n_slots = state_batch_indices.numel() if n_slots == 0: return - num_trackers = 1 if ring_start.dim() == 1 else ring_start.shape[0] - tracker_stride = ring_start.shape[-1] block = 128 - _update_replayssm_ring_trackers_kernel[ - (triton.cdiv(n_slots, block), num_trackers) - ]( + _update_replayssm_ring_trackers_kernel[(triton.cdiv(n_slots, block),)]( ring_start, prev_num_accepted, state_batch_indices, n_slots, - tracker_stride, logical_window, logical_window + 1, pad_slot_id, @@ -546,8 +538,6 @@ def __call__( cb_old: torch.Tensor | None = None, algorithm: str | None = None, update_trackers: bool = True, - ring_starts_to_update: torch.Tensor | None = None, - prev_num_accepted_to_update: torch.Tensor | None = None, ) -> torch.Tensor: # AR decode currently passes (batch, nheads, dim); checkpointing_ssu # expects a predicted-token axis T. Unsqueeze T=1 here. @@ -596,17 +586,9 @@ def __call__( algorithm=self._algorithm if algorithm is None else algorithm, ) if update_trackers and indices is not None: - if (ring_starts_to_update is None) != ( - prev_num_accepted_to_update is None - ): - raise ValueError("ReplaySSM tracker update tensors must be paired") update_replayssm_ring_trackers( - ring_start - if ring_starts_to_update is None - else ring_starts_to_update, - prev_num_accepted_tokens - if prev_num_accepted_to_update is None - else prev_num_accepted_to_update, + ring_start, + prev_num_accepted_tokens, indices, logical_window=x_cache.size(2) - 1, pad_slot_id=null_block_id, @@ -777,8 +759,6 @@ def selective_state_update_replayssm_flashinfer( cb_old: torch.Tensor | None = None, algorithm: str | None = None, update_trackers: bool = True, - ring_starts_to_update: torch.Tensor | None = None, - prev_num_accepted_to_update: torch.Tensor | None = None, ) -> torch.Tensor: """FlashInfer ReplaySSM decode (``checkpointing_ssu``).""" backend = get_replayssm_backend() @@ -811,8 +791,6 @@ def selective_state_update_replayssm_flashinfer( cb_old=cb_old, algorithm=algorithm, update_trackers=update_trackers, - ring_starts_to_update=ring_starts_to_update, - prev_num_accepted_to_update=prev_num_accepted_to_update, ) diff --git a/vllm/v1/worker/gpu/attn_utils.py b/vllm/v1/worker/gpu/attn_utils.py index 05febec328d4..8a7043e54a14 100644 --- a/vllm/v1/worker/gpu/attn_utils.py +++ b/vllm/v1/worker/gpu/attn_utils.py @@ -229,7 +229,13 @@ def init_kv_cache( in ("longcat_flash", "longcat_flash_ngram") else 1 ) - bind_kv_cache(kv_caches, forward_context, runner_kv_caches, num_attn_module) + bind_kv_cache( + kv_caches, + forward_context, + runner_kv_caches, + num_attn_module, + kv_cache_groups=kv_cache_config.kv_cache_groups, + ) return kv_caches diff --git a/vllm/v1/worker/gpu_model_runner.py b/vllm/v1/worker/gpu_model_runner.py index d13029a2e513..60f6fab81f28 100644 --- a/vllm/v1/worker/gpu_model_runner.py +++ b/vllm/v1/worker/gpu_model_runner.py @@ -7433,6 +7433,7 @@ def initialize_kv_cache_tensors( self.compilation_config.static_forward_context, self.kv_caches, num_attn_module, + kv_cache_groups=kv_cache_config.kv_cache_groups, ) return kv_caches diff --git a/vllm/v1/worker/utils.py b/vllm/v1/worker/utils.py index 8811b2e4bfdc..ed8abb502fbe 100644 --- a/vllm/v1/worker/utils.py +++ b/vllm/v1/worker/utils.py @@ -566,6 +566,7 @@ def bind_kv_cache( forward_context: dict[str, Attention], runner_kv_caches: list[torch.Tensor], num_attn_module: int = 1, + kv_cache_groups: Sequence[KVCacheGroupSpec] | None = None, ) -> None: """ Bind the allocated KV cache to both ModelRunner and forward context so @@ -616,17 +617,31 @@ def bind_kv_cache( from vllm.model_executor.layers.mamba.mamba_mixer2 import ( MambaMixer2, - batch_replayssm_ring_tracker_updates, + share_replayssm_ring_trackers, ) - replayssm_mixers = [ - layer - for layer_name in ordered_layer_names - if isinstance((layer := forward_context[layer_name]), MambaMixer2) - and layer.use_replayssm - and layer.mamba_config.backend == MambaBackendEnum.FLASHINFER - ] - batch_replayssm_ring_tracker_updates(replayssm_mixers) + layer_to_cache_group = { + layer_name: group_index + for group_index, group in enumerate(kv_cache_groups or ()) + for layer_name in group.layer_names + } + replayssm_mixer_groups = defaultdict(list) + for layer_name in ordered_layer_names: + layer = forward_context[layer_name] + if not ( + isinstance(layer, MambaMixer2) + and layer.use_replayssm + and layer.mamba_config.backend == MambaBackendEnum.FLASHINFER + ): + continue + group_index = layer_to_cache_group.get(layer_name) + if group_index is None: + # Callers without a KV-cache configuration cannot prove that two + # layers use the same block-index namespace, so keep them separate. + group_index = layer_name + replayssm_mixer_groups[group_index].append(layer) + + share_replayssm_ring_trackers(list(replayssm_mixer_groups.values())) def copy_kv_cache_blocks_inplace( From 7c0bfd753f1801ac9119f492a0bceaa5bce6c144 Mon Sep 17 00:00:00 2001 From: Andrii Skliar Date: Sun, 16 Aug 2026 04:58:59 -0700 Subject: [PATCH 13/33] Use native FlashInfer ReplaySSM autotuning Signed-off-by: Andrii Skliar --- tests/kernels/mamba/test_ssu_dispatch.py | 48 +- .../test_flashinfer_replayssm_warmup.py | 248 ++---- .../layers/mamba/ops/ssu_dispatch.py | 75 +- .../warmup/flashinfer_replayssm_warmup.py | 759 +++--------------- vllm/model_executor/warmup/kernel_warmup.py | 8 +- vllm/v1/worker/gpu_model_runner.py | 18 +- 6 files changed, 294 insertions(+), 862 deletions(-) diff --git a/tests/kernels/mamba/test_ssu_dispatch.py b/tests/kernels/mamba/test_ssu_dispatch.py index 1893c70070c6..d852908a5a2f 100644 --- a/tests/kernels/mamba/test_ssu_dispatch.py +++ b/tests/kernels/mamba/test_ssu_dispatch.py @@ -1,6 +1,8 @@ # SPDX-License-Identifier: Apache-2.0 # SPDX-FileCopyrightText: Copyright contributors to the vLLM project +import importlib +from types import SimpleNamespace from unittest.mock import Mock import pytest @@ -37,9 +39,18 @@ HAS_FLASHINFER = False try: - from flashinfer.mamba import checkpointing_ssu # noqa: F401 - - HAS_FLASHINFER_CHECKPOINTING_SSU = True + checkpointing_ssu_module = importlib.import_module( + "flashinfer.mamba.checkpointing_ssu" + ) + HAS_FLASHINFER_CHECKPOINTING_SSU = ( + hasattr(checkpointing_ssu_module, "CheckpointingSSURunner") + and getattr( + checkpointing_ssu_module, + "CHECKPOINTING_SSU_AUTOTUNE_ABI_VERSION", + 0, + ) + >= 1 + ) except ImportError: HAS_FLASHINFER_CHECKPOINTING_SSU = False @@ -306,10 +317,15 @@ def test_replayssm_triton_entry_rejects_flashinfer_backend(): reason="flashinfer.mamba.checkpointing_ssu not available", ) def test_replayssm_flashinfer_call(monkeypatch): - import flashinfer.mamba - kernel = Mock(return_value=torch.empty(1, 1, 2, 4)) - monkeypatch.setattr(flashinfer.mamba, "checkpointing_ssu", kernel) + checkpointing_ssu_module = importlib.import_module( + "flashinfer.mamba.checkpointing_ssu" + ) + monkeypatch.setattr(checkpointing_ssu_module, "checkpointing_ssu", kernel) + monkeypatch.setattr( + "vllm.model_executor.layers.mamba.ops.ssu_dispatch._replayssm_backend", + None, + ) initialize_replayssm_backend( MambaConfig(backend=MambaBackendEnum.FLASHINFER), use_replayssm=True ) @@ -354,6 +370,21 @@ def test_replayssm_flashinfer_call(monkeypatch): args = kernel.call_args.args assert args[4] is ring_start assert args[5] is prev_num_accepted + assert kwargs["algorithm"] == "auto" + assert kwargs["precompute_heads_per_cta"] == 0 + assert kwargs["main_pipeline_stages"] == 0 + assert kwargs["main_ctas_per_sm"] == 0 + + +def test_replayssm_requires_native_flashinfer_autotuning(monkeypatch): + import vllm.model_executor.layers.mamba.ops.ssu_dispatch as mod + + old_module = SimpleNamespace( + checkpointing_ssu=Mock(), CheckpointingSSURunner=object + ) + monkeypatch.setattr(mod.importlib, "import_module", lambda _: old_module) + with pytest.raises(ImportError, match="native checkpointing_ssu autotuning"): + FlashInferReplaySSMBackend(MambaConfig(backend=MambaBackendEnum.FLASHINFER)) @pytest.mark.skipif( @@ -361,7 +392,10 @@ def test_replayssm_flashinfer_call(monkeypatch): reason="flashinfer checkpointing_ssu is installed", ) def test_replayssm_flashinfer_import_error(): - with pytest.raises(ImportError, match="FlashInfer is required"): + with pytest.raises( + ImportError, + match="FlashInfer is required|native checkpointing_ssu autotuning", + ): FlashInferReplaySSMBackend(MambaConfig(backend=MambaBackendEnum.FLASHINFER)) diff --git a/tests/model_executor/test_flashinfer_replayssm_warmup.py b/tests/model_executor/test_flashinfer_replayssm_warmup.py index af6bc40d50b6..4a5c9e3d0da3 100644 --- a/tests/model_executor/test_flashinfer_replayssm_warmup.py +++ b/tests/model_executor/test_flashinfer_replayssm_warmup.py @@ -16,108 +16,20 @@ FlashInferReplaySSMTactic, use_flashinfer_replayssm_tactic, ) -from vllm.model_executor.warmup.flashinfer_replayssm_warmup import ( - FLASHINFER_REPLAYSSM_TUNING_CANDIDATES, - FlashInferReplaySSMAutotuneResult, - _expanded_precompute_tactics, - _load_cache, - _make_cache_key, - _precompute_heads_per_cta_candidates, - _ReplaySSMBenchmark, - _save_cache, - _select_fastest, - flashinfer_replayssm_autotune_warmup, -) - -_STAGES_ENV = "FLASHINFER_SSU_MAIN_PIPELINE_STAGES" -_CTAS_ENV = "FLASHINFER_SSU_MAIN_CTA_PER_SM" +from vllm.model_executor.warmup import flashinfer_replayssm_warmup as warmup def _fake_backend() -> FlashInferReplaySSMBackend: backend = FlashInferReplaySSMBackend.__new__(FlashInferReplaySSMBackend) backend._mamba_config = MambaConfig(backend=MambaBackendEnum.FLASHINFER) backend._kernel = Mock(return_value=torch.empty(1)) - backend._algorithm = "auto" - backend._precompute_heads_per_cta = 0 + backend._tactic = FLASHINFER_REPLAYSSM_AUTO_TACTIC return backend -def test_replayssm_tuning_candidates_and_deterministic_selection(): - assert [tactic.name for tactic in FLASHINFER_REPLAYSSM_TUNING_CANDIDATES] == [ - "auto", - "monolith", - "two_kernel_s1_c1", - "two_kernel_s1_c2", - "two_kernel_s1_c4", - "two_kernel_s1_c8", - "two_kernel_s1_c16", - "two_kernel_s2_c1", - "two_kernel_s2_c2", - "two_kernel_s2_c4", - "two_kernel_s2_c8", - "two_kernel_s2_c16", - ] - timings = [3.0, 2.0, 2.0, 4.0, 5.0, 6.0, 7.0, 8.0, 9.0, 10.0, 11.0, 12.0] - assert _select_fastest(timings).name == "monolith" - assert _select_fastest([float("inf")] * len(timings)) is None - near_tie = [1.0, 0.9995] + [float("inf")] * (len(timings) - 2) - assert _select_fastest(near_tie).name == "auto" - - -def test_replayssm_precompute_candidates_expand_top_main_tactics(): - assert _precompute_heads_per_cta_candidates(16) == (1, 2, 4, 8, 16) - timings = [float("inf")] * len(FLASHINFER_REPLAYSSM_TUNING_CANDIDATES) - timings[4] = 1.0 - timings[10] = 2.0 - timings[2] = 3.0 - - assert [ - tactic.name for tactic in _expanded_precompute_tactics(timings, 16) - ] == [ - "two_kernel_s1_c4_h1", - "two_kernel_s1_c4_h2", - "two_kernel_s1_c4_h4", - "two_kernel_s1_c4_h8", - "two_kernel_s1_c4_h16", - "two_kernel_s2_c8_h1", - "two_kernel_s2_c8_h2", - "two_kernel_s2_c8_h4", - "two_kernel_s2_c8_h8", - "two_kernel_s2_c8_h16", - ] - - -def test_replayssm_tactic_validates_precompute_geometry(): - assert FlashInferReplaySSMTactic( - "two-kernel", 1, 4, precompute_heads_per_cta=8 - ).name == "two_kernel_s1_c4_h8" - with pytest.raises(ValueError, match="does not accept precompute"): - FlashInferReplaySSMTactic("monolith", precompute_heads_per_cta=8) - - -def test_replayssm_tuning_key_distinguishes_batch_and_T(): - fingerprint = {"geometry": "h64_d64_n128"} - assert _make_cache_key(fingerprint, 32, 8) != _make_cache_key(fingerprint, 256, 1) - assert _make_cache_key(fingerprint, 32, 8) == _make_cache_key(fingerprint, 32, 8) - - -def test_replayssm_autotune_cache_round_trip(tmp_path): - path = tmp_path / "replayssm.json" - _save_cache(path, {"key": "two_kernel_s2_c16_h8", "bad": "unknown"}) - assert _load_cache(path) == {"key": "two_kernel_s2_c16_h8"} - - -@pytest.mark.parametrize("payload", ["null", "[]", "1"]) -def test_replayssm_autotune_ignores_non_mapping_cache(tmp_path, payload): - path = tmp_path / "replayssm.json" - path.write_text(payload) - assert _load_cache(path) == {} - - -def test_replayssm_benchmark_reserves_null_cache_slot(): - pytest.importorskip("flashinfer.mamba") - cache_slots, nheads, headdim, dstate, ngroups = 5, 2, 4, 8, 1 - layer = SimpleNamespace( +def _fake_layer(cache_slots: int = 5): + nheads, headdim, dstate, ngroups = 2, 4, 8, 1 + return SimpleNamespace( kv_cache=( torch.empty(0), torch.empty(cache_slots, nheads, headdim, dstate), @@ -134,79 +46,36 @@ def test_replayssm_benchmark_reserves_null_cache_slot(): ), ) - benchmark = _ReplaySSMBenchmark(layer, 3) - assert benchmark.indices.tolist() == [1, 2, 3] - assert benchmark.ring_start.shape == (cache_slots,) - assert benchmark.initial_ring_start.tolist() == [0, 0, 1, 2, 0] - assert benchmark.initial_prev_num_accepted.tolist() == [0, 1, 2, 3, 0] - _ReplaySSMBenchmark(layer, cache_slots - 1) - with pytest.raises(ValueError, match="needs 6 cache slots"): - _ReplaySSMBenchmark(layer, cache_slots) - -@pytest.mark.parametrize( - ("use_v2_model_runner", "use_ubatching"), - [(True, False), (False, True)], -) -def test_replayssm_autotune_safely_skips_unsupported_runners( - use_v2_model_runner, use_ubatching -): - runner = SimpleNamespace( - parallel_config=SimpleNamespace(use_ubatching=use_ubatching) - ) - worker = SimpleNamespace( - model_runner=runner, - vllm_config=SimpleNamespace( - kernel_config=SimpleNamespace(enable_flashinfer_autotune=True) - ), - model_config=SimpleNamespace(enforce_eager=False), - use_v2_model_runner=use_v2_model_runner, - ) - - flashinfer_replayssm_autotune_warmup(worker) - assert runner.flashinfer_replayssm_autotune_result is None +def test_replayssm_explicit_tactic_validation(): + tactic = FlashInferReplaySSMTactic("two-kernel", 1, 4, precompute_heads_per_cta=8) + assert tactic.name == "two_kernel_s1_c4_h8" + with pytest.raises(ValueError, match="does not accept precompute"): + FlashInferReplaySSMTactic("monolith", precompute_heads_per_cta=8) -def test_replayssm_tactic_scope_restores_algorithm_and_environment( - monkeypatch, -): +def test_replayssm_tactic_scope_restores_direct_launch_controls(monkeypatch): backend = _fake_backend() monkeypatch.setattr(ssu_dispatch, "_replayssm_backend", backend) - monkeypatch.setenv(_STAGES_ENV, "7") - monkeypatch.setenv(_CTAS_ENV, "11") + tactic = FlashInferReplaySSMTactic("two-kernel", 2, 16, precompute_heads_per_cta=8) - tactic = FlashInferReplaySSMTactic( - "two-kernel", 2, 16, precompute_heads_per_cta=8 - ) with ( pytest.raises(RuntimeError, match="sentinel"), use_flashinfer_replayssm_tactic(tactic), ): - assert backend._algorithm == "two-kernel" - assert backend._precompute_heads_per_cta == 8 - assert ssu_dispatch.os.environ[_STAGES_ENV] == "2" - assert ssu_dispatch.os.environ[_CTAS_ENV] == "16" + assert backend._tactic is tactic raise RuntimeError("sentinel") - assert backend._algorithm == "auto" - assert backend._precompute_heads_per_cta == 0 - assert ssu_dispatch.os.environ[_STAGES_ENV] == "7" - assert ssu_dispatch.os.environ[_CTAS_ENV] == "11" - - with use_flashinfer_replayssm_tactic(FLASHINFER_REPLAYSSM_AUTO_TACTIC): - assert _STAGES_ENV not in ssu_dispatch.os.environ - assert _CTAS_ENV not in ssu_dispatch.os.environ + assert backend._tactic is FLASHINFER_REPLAYSSM_AUTO_TACTIC -def test_replayssm_backend_uses_scoped_algorithm(monkeypatch): +def test_replayssm_backend_forwards_explicit_tactic(monkeypatch): backend = _fake_backend() monkeypatch.setattr(ssu_dispatch, "_replayssm_backend", backend) tensor = torch.empty(1) with use_flashinfer_replayssm_tactic( - FlashInferReplaySSMTactic( - "two-kernel", 1, 4, precompute_heads_per_cta=8 - ) + FlashInferReplaySSMTactic("two-kernel", 1, 4, precompute_heads_per_cta=8) ): backend( tensor, @@ -223,36 +92,63 @@ def test_replayssm_backend_uses_scoped_algorithm(monkeypatch): tensor, ) - assert backend._kernel.call_args.kwargs["algorithm"] == "two-kernel" - assert backend._kernel.call_args.kwargs["precompute_heads_per_cta"] == 8 + kwargs = backend._kernel.call_args.kwargs + assert kwargs["algorithm"] == "two-kernel" + assert kwargs["precompute_heads_per_cta"] == 8 + assert kwargs["main_pipeline_stages"] == 1 + assert kwargs["main_ctas_per_sm"] == 4 -def test_replayssm_capture_tactic_uses_request_batch_not_tokens(): - result = FlashInferReplaySSMAutotuneResult( - spec_query_len=1, - tactics={32: FlashInferReplaySSMTactic("two-kernel", 2, 8)}, - ) - assert ( - result.tactic_for( - BatchDescriptor(num_tokens=32, num_reqs=32, uniform=True) - ).name - == "two_kernel_s2_c8" - ) - assert ( - result.tactic_for(BatchDescriptor(num_tokens=256, num_reqs=32, uniform=True)) - is None - ) - mtp_result = FlashInferReplaySSMAutotuneResult( - spec_query_len=8, - tactics={32: FlashInferReplaySSMTactic("two-kernel", 2, 8)}, - ) - assert ( - mtp_result.tactic_for( - BatchDescriptor(num_tokens=256, num_reqs=32, uniform=True) - ).name - == "two_kernel_s2_c8" +def test_replayssm_tuning_call_uses_private_production_layout(): + layer = _fake_layer() + call = warmup._ReplaySSMTuningCall(layer, 3) + + assert call.state.shape == (4, *layer.kv_cache[1].shape[1:]) + assert call.state.stride() == layer.kv_cache[1].stride() + assert call.state.data_ptr() != layer.kv_cache[1].data_ptr() + assert call.x_cache.data_ptr() != layer.kv_cache[2].data_ptr() + assert call.indices.tolist() == [1, 2, 3] + assert call.ring_start.shape == (4,) + assert call.prev_num_accepted.tolist() == [0, 1, 2, 3] + assert call.dt.shape == (3, 2, 4) + assert call.dt.stride(-1) == 0 + + +def test_replayssm_tuning_trigger_uses_largest_supported_decode_batch(monkeypatch): + layer = _fake_layer(cache_slots=5) + runner = SimpleNamespace( + scheduler_config=SimpleNamespace(max_num_seqs=8), + cudagraph_dispatcher=SimpleNamespace( + get_capture_descs=lambda: [ + ( + None, + [ + BatchDescriptor(num_tokens=2, num_reqs=2, uniform=True), + BatchDescriptor(num_tokens=8, num_reqs=8, uniform=True), + ], + ) + ] + ), ) - assert ( - result.tactic_for(BatchDescriptor(num_tokens=32, num_reqs=None, uniform=False)) - is None + observed = SimpleNamespace(batch=None, ran=False) + + class FakeCall: + def __init__(self, _layer, batch): + assert _layer is layer + observed.batch = batch + + def run(self): + observed.ran = True + + monkeypatch.setattr(warmup, "_find_replayssm_layers", lambda _runner: (layer,)) + monkeypatch.setattr( + warmup, "_distributed_layers_are_compatible", lambda _layers: True ) + monkeypatch.setattr(warmup, "_distributed_min", lambda value: value) + monkeypatch.setattr(warmup, "_ReplaySSMTuningCall", FakeCall) + monkeypatch.setattr(torch.cuda, "synchronize", lambda: None) + + warmup.trigger_flashinfer_replayssm_autotune(runner) + + assert observed.batch == 4 + assert observed.ran diff --git a/vllm/model_executor/layers/mamba/ops/ssu_dispatch.py b/vllm/model_executor/layers/mamba/ops/ssu_dispatch.py index 34f09115466e..fba515b2b375 100644 --- a/vllm/model_executor/layers/mamba/ops/ssu_dispatch.py +++ b/vllm/model_executor/layers/mamba/ops/ssu_dispatch.py @@ -16,7 +16,7 @@ (``selective_state_update_replayssm_flashinfer``) """ -import os +import importlib from abc import ABC, abstractmethod from collections.abc import Iterator from contextlib import contextmanager @@ -33,9 +33,6 @@ logger = init_logger(__name__) -_FLASHINFER_SSU_PIPELINE_STAGES_ENV = "FLASHINFER_SSU_MAIN_PIPELINE_STAGES" -_FLASHINFER_SSU_CTA_PER_SM_ENV = "FLASHINFER_SSU_MAIN_CTA_PER_SM" - @dataclass(frozen=True) class FlashInferReplaySSMTactic: @@ -498,16 +495,32 @@ class FlashInferReplaySSMBackend(ReplaySSMBackend): def __init__(self, mamba_config: MambaConfig): super().__init__(mamba_config) try: - from flashinfer.mamba import checkpointing_ssu as _fi_checkpointing_ssu - except ImportError as e: + checkpointing_ssu_module = importlib.import_module( + "flashinfer.mamba.checkpointing_ssu" + ) + except (ImportError, ModuleNotFoundError) as e: raise ImportError( "FlashInfer is required for the flashinfer ReplaySSM backend. " "Please install flashinfer with mamba.checkpointing_ssu support: " "pip install flashinfer-python" ) from e - self._kernel = _fi_checkpointing_ssu - self._algorithm = FLASHINFER_REPLAYSSM_AUTO_TACTIC.algorithm - self._precompute_heads_per_cta = 0 + autotune_abi = getattr( + checkpointing_ssu_module, + "CHECKPOINTING_SSU_AUTOTUNE_ABI_VERSION", + 0, + ) + if ( + not hasattr(checkpointing_ssu_module, "CheckpointingSSURunner") + or autotune_abi < 1 + ): + raise ImportError( + "FlashInfer ReplaySSM requires native checkpointing_ssu " + "autotuning ABI version 1 or newer. Install a compatible " + "FlashInfer revision exposing " + "CHECKPOINTING_SSU_AUTOTUNE_ABI_VERSION >= 1." + ) + self._kernel = checkpointing_ssu_module.checkpointing_ssu + self._tactic = FLASHINFER_REPLAYSSM_AUTO_TACTIC @property def name(self) -> str: @@ -558,6 +571,9 @@ def __call__( if indices is not None and indices.dim() > 1: indices = indices[:, 0] + tactic = self._tactic + requested_algorithm = tactic.algorithm if algorithm is None else algorithm + use_explicit_tactic = algorithm is None result = self._kernel( state, x_cache, @@ -582,8 +598,14 @@ def __call__( cb_scaled=cb_scaled, cumAdt_vec=cumAdt_vec, cb_old=cb_old, - precompute_heads_per_cta=self._precompute_heads_per_cta, - algorithm=self._algorithm if algorithm is None else algorithm, + precompute_heads_per_cta=( + tactic.precompute_heads_per_cta if use_explicit_tactic else 0 + ), + main_pipeline_stages=( + tactic.pipeline_stages or 0 if use_explicit_tactic else 0 + ), + main_ctas_per_sm=(tactic.ctas_per_sm or 0 if use_explicit_tactic else 0), + algorithm=requested_algorithm, ) if update_trackers and indices is not None: update_replayssm_ring_trackers( @@ -600,41 +622,18 @@ def __call__( def use_flashinfer_replayssm_tactic( tactic: FlashInferReplaySSMTactic, ) -> Iterator[None]: - """Apply a ReplaySSM launch tactic during serial warmup or graph capture.""" + """Apply an explicit ReplaySSM tactic for tests or debugging.""" backend = get_replayssm_backend() if not isinstance(backend, FlashInferReplaySSMBackend): yield return - old_algorithm = backend._algorithm - old_precompute_heads_per_cta = backend._precompute_heads_per_cta - old_stages = os.environ.get(_FLASHINFER_SSU_PIPELINE_STAGES_ENV) - old_ctas = os.environ.get(_FLASHINFER_SSU_CTA_PER_SM_ENV) - backend._algorithm = tactic.algorithm - backend._precompute_heads_per_cta = tactic.precompute_heads_per_cta + old_tactic = backend._tactic + backend._tactic = tactic try: - if tactic.algorithm == "two-kernel": - assert tactic.pipeline_stages is not None - assert tactic.ctas_per_sm is not None - os.environ[_FLASHINFER_SSU_PIPELINE_STAGES_ENV] = str( - tactic.pipeline_stages - ) - os.environ[_FLASHINFER_SSU_CTA_PER_SM_ENV] = str(tactic.ctas_per_sm) - else: - os.environ.pop(_FLASHINFER_SSU_PIPELINE_STAGES_ENV, None) - os.environ.pop(_FLASHINFER_SSU_CTA_PER_SM_ENV, None) yield finally: - backend._algorithm = old_algorithm - backend._precompute_heads_per_cta = old_precompute_heads_per_cta - if old_stages is None: - os.environ.pop(_FLASHINFER_SSU_PIPELINE_STAGES_ENV, None) - else: - os.environ[_FLASHINFER_SSU_PIPELINE_STAGES_ENV] = old_stages - if old_ctas is None: - os.environ.pop(_FLASHINFER_SSU_CTA_PER_SM_ENV, None) - else: - os.environ[_FLASHINFER_SSU_CTA_PER_SM_ENV] = old_ctas + backend._tactic = old_tactic _REPLAYSSM_BACKEND_REGISTRY: dict[MambaBackendEnum, type[ReplaySSMBackend]] = { diff --git a/vllm/model_executor/warmup/flashinfer_replayssm_warmup.py b/vllm/model_executor/warmup/flashinfer_replayssm_warmup.py index 8e39b95f15f8..2160cad8055d 100644 --- a/vllm/model_executor/warmup/flashinfer_replayssm_warmup.py +++ b/vllm/model_executor/warmup/flashinfer_replayssm_warmup.py @@ -1,209 +1,26 @@ # SPDX-License-Identifier: Apache-2.0 # SPDX-FileCopyrightText: Copyright contributors to the vLLM project -"""Startup autotuning for FlashInfer ReplaySSM CUDA-graph launches.""" +"""Trigger native FlashInfer ReplaySSM tuning before CUDA graph capture.""" from __future__ import annotations -import gc -import json -import math -import statistics -import time -from contextlib import contextmanager, nullcontext -from dataclasses import dataclass, replace -from pathlib import Path from typing import TYPE_CHECKING, Any import torch from vllm.config.mamba import MambaBackendEnum from vllm.distributed.parallel_state import get_world_group -from vllm.forward_context import BatchDescriptor from vllm.logger import init_logger from vllm.model_executor.layers.mamba.ops.ssu_dispatch import ( - FLASHINFER_REPLAYSSM_AUTO_TACTIC, - FlashInferReplaySSMTactic, - update_replayssm_ring_trackers, - use_flashinfer_replayssm_tactic, -) -from vllm.model_executor.warmup.flashinfer_autotune_cache import ( - resolve_flashinfer_autotune_file, - write_flashinfer_autotune_cache, + selective_state_update_replayssm_flashinfer, ) if TYPE_CHECKING: from vllm.model_executor.layers.mamba.mamba_mixer2 import MambaMixer2 from vllm.v1.worker.gpu_model_runner import GPUModelRunner - from vllm.v1.worker.gpu_worker import Worker logger = init_logger(__name__) -_CACHE_SCHEMA_VERSION = 3 -_TUNING_FIXTURE = "t1_mixed_history_cycle_v4" -_CACHE_FILE_NAME = "replayssm_autotune_configs.json" -_PRECOMPUTE_MAIN_TACTIC_COUNT = 2 -_RELATIVE_TIE_TOLERANCE = 0.001 - -FLASHINFER_REPLAYSSM_TUNING_CANDIDATES = ( - FLASHINFER_REPLAYSSM_AUTO_TACTIC, - FlashInferReplaySSMTactic("monolith"), - FlashInferReplaySSMTactic("two-kernel", 1, 1), - FlashInferReplaySSMTactic("two-kernel", 1, 2), - FlashInferReplaySSMTactic("two-kernel", 1, 4), - FlashInferReplaySSMTactic("two-kernel", 1, 8), - FlashInferReplaySSMTactic("two-kernel", 1, 16), - FlashInferReplaySSMTactic("two-kernel", 2, 1), - FlashInferReplaySSMTactic("two-kernel", 2, 2), - FlashInferReplaySSMTactic("two-kernel", 2, 4), - FlashInferReplaySSMTactic("two-kernel", 2, 8), - FlashInferReplaySSMTactic("two-kernel", 2, 16), -) -_TACTICS_BY_NAME = { - tactic.name: tactic for tactic in FLASHINFER_REPLAYSSM_TUNING_CANDIDATES -} - - -def _tactic_from_name(name: str) -> FlashInferReplaySSMTactic | None: - if tactic := _TACTICS_BY_NAME.get(name): - return tactic - for base in FLASHINFER_REPLAYSSM_TUNING_CANDIDATES: - if base.algorithm != "two-kernel": - continue - prefix = f"{base.name}_h" - if name.startswith(prefix): - try: - heads_per_cta = int(name.removeprefix(prefix)) - return replace(base, precompute_heads_per_cta=heads_per_cta) - except ValueError: - return None - return None - - -def _precompute_heads_per_cta_candidates(heads_per_group: int) -> tuple[int, ...]: - """Return the distinct HEADS_PER_GROUP >> k launch geometries.""" - if heads_per_group <= 0: - raise ValueError("heads_per_group must be positive") - candidates = set() - value = heads_per_group - while value: - candidates.add(value) - value >>= 1 - return tuple(sorted(candidates)) - - -def _expanded_precompute_tactics( - main_timings: list[float], heads_per_group: int -) -> tuple[FlashInferReplaySSMTactic, ...]: - ranked = sorted( - ( - index - for index, tactic in enumerate( - FLASHINFER_REPLAYSSM_TUNING_CANDIDATES - ) - if tactic.algorithm == "two-kernel" - and math.isfinite(main_timings[index]) - ), - key=lambda index: (main_timings[index], index), - )[:_PRECOMPUTE_MAIN_TACTIC_COUNT] - return tuple( - replace(base, precompute_heads_per_cta=heads_per_cta) - for index in ranked - for base in (FLASHINFER_REPLAYSSM_TUNING_CANDIDATES[index],) - for heads_per_cta in _precompute_heads_per_cta_candidates(heads_per_group) - ) - - -@dataclass -class FlashInferReplaySSMAutotuneResult: - spec_query_len: int - tactics: dict[int, FlashInferReplaySSMTactic] - - def tactic_for( - self, batch_descriptor: BatchDescriptor - ) -> FlashInferReplaySSMTactic | None: - if ( - not batch_descriptor.uniform - or batch_descriptor.num_reqs is None - or batch_descriptor.num_tokens - != batch_descriptor.num_reqs * self.spec_query_len - ): - return None - return self.tactics.get(batch_descriptor.num_reqs) - - -def _make_cache_key(fingerprint: dict[str, Any], batch: int, T: int) -> str: - return json.dumps( - { - **fingerprint, - "batch_sequences": batch, - "spec_query_len": T, - "fixture": _TUNING_FIXTURE, - "candidate_schema": { - "main": [ - tactic.name - for tactic in FLASHINFER_REPLAYSSM_TUNING_CANDIDATES - ], - "precompute": "top2_main_x_heads_per_group_shift_chain", - "relative_tie_tolerance": _RELATIVE_TIE_TOLERANCE, - }, - }, - sort_keys=True, - separators=(",", ":"), - ) - - -def _select_fastest( - timings: list[float], - candidates: tuple[FlashInferReplaySSMTactic, ...] = ( - FLASHINFER_REPLAYSSM_TUNING_CANDIDATES - ), -) -> FlashInferReplaySSMTactic | None: - if len(timings) != len(candidates): - raise ValueError("one timing is required for each ReplaySSM tactic") - finite = [i for i, timing in enumerate(timings) if math.isfinite(timing)] - if not finite: - return None - best_timing = min(timings[i] for i in finite) - tied = [ - i - for i in finite - if timings[i] <= best_timing * (1 + _RELATIVE_TIE_TOLERANCE) - ] - winner = min(tied) - return candidates[winner] - - -def _load_cache(path: Path) -> dict[str, str]: - try: - payload = json.loads(path.read_text()) - except (OSError, ValueError, TypeError): - return {} - if not isinstance(payload, dict): - return {} - if payload.get("schema_version") != _CACHE_SCHEMA_VERSION: - return {} - entries = payload.get("entries") - if not isinstance(entries, dict): - return {} - return { - key: value - for key, value in entries.items() - if isinstance(key, str) - and isinstance(value, str) - and _tactic_from_name(value) is not None - } - - -def _save_cache(path: Path, entries: dict[str, str]) -> None: - payload = { - "schema_version": _CACHE_SCHEMA_VERSION, - "entries": dict(sorted(entries.items())), - } - write_flashinfer_autotune_cache( - path, - (json.dumps(payload, sort_keys=True, indent=2) + "\n").encode(), - ) - def _find_replayssm_layers(runner: GPUModelRunner) -> tuple[MambaMixer2, ...]: from vllm.model_executor.layers.mamba.mamba_mixer2 import MambaMixer2 @@ -219,7 +36,19 @@ def _find_replayssm_layers(runner: GPUModelRunner) -> tuple[MambaMixer2, ...]: ) -def _uniform_capture_batches(runner: GPUModelRunner) -> tuple[int, ...]: +def _layer_signature(layer: MambaMixer2) -> tuple[Any, ...]: + _, state, x_cache, dt_cache, B_cache, *_ = layer.kv_cache + tensors = (state, x_cache, dt_cache, B_cache, layer.A, layer.D, layer.dt_bias) + return tuple( + (tuple(tensor.shape), tuple(tensor.stride()), tensor.dtype) + for tensor in tensors + ) + ( + layer.mamba_config.enable_stochastic_rounding, + layer.mamba_config.stochastic_rounding_philox_rounds, + ) + + +def _uniform_decode_batches(runner: GPUModelRunner) -> tuple[int, ...]: return tuple( sorted( { @@ -232,64 +61,88 @@ def _uniform_capture_batches(runner: GPUModelRunner) -> tuple[int, ...]: ) -def _layer_fingerprint(layer: MambaMixer2) -> dict[str, Any]: - import flashinfer - from flashinfer.jit import env as flashinfer_jit_env - - _, state, x_cache, dt_cache, B_cache, *_ = layer.kv_cache - props = torch.cuda.get_device_properties(state.device) - return { - "flashinfer_version": flashinfer.__version__, - "flashinfer_workspace": [ - flashinfer_jit_env.FLASHINFER_WORKSPACE_DIR.parent.name, - flashinfer_jit_env.FLASHINFER_WORKSPACE_DIR.name, - ], - "gpu_name": props.name, - "gpu_capability": list(torch.cuda.get_device_capability(state.device)), - "gpu_sm_count": props.multi_processor_count, - "nheads": state.shape[1], - "headdim": state.shape[2], - "dstate": state.shape[3], - "ngroups": B_cache.shape[1], - "cache_slots": state.shape[0], - "physical_ring_len": x_cache.shape[2], - "state_dtype": str(state.dtype), - "activation_dtype": str(x_cache.dtype), - "dt_cache_dtype": str(dt_cache.dtype), - "state_stride": list(state.stride()), - "x_cache_stride": list(x_cache.stride()), - "B_cache_stride": list(B_cache.stride()), - "dt_cache_stride": list(dt_cache.stride()), - "A_dtype": str(layer.A.dtype), - "A_stride": list(layer.A.stride()), - "D_dtype": str(layer.D.dtype), - "D_stride": list(layer.D.stride()), - "dt_bias_dtype": str(layer.dt_bias.dtype), - "dt_bias_stride": list(layer.dt_bias.stride()), - "stochastic_rounding": layer.mamba_config.enable_stochastic_rounding, - "stochastic_rounding_philox_rounds": ( - layer.mamba_config.stochastic_rounding_philox_rounds +def _local_max_tuning_batch(runner: GPUModelRunner, state_capacity: int) -> int: + capture_batches = _uniform_decode_batches(runner) + requested = ( + capture_batches[-1] if capture_batches else runner.scheduler_config.max_num_seqs + ) + return max( + 0, + min( + requested, + runner.scheduler_config.max_num_seqs, + state_capacity - 1, ), - "tp_size": layer.tp_size, - } + ) -class _ReplaySSMBenchmark: - def __init__(self, layer: MambaMixer2, batch: int): - from flashinfer.mamba import checkpointing_ssu +def _distributed_min(value: int) -> int: + world = get_world_group() + if world.world_size == 1: + return value + tensor = torch.tensor([value], dtype=torch.int64) + torch.distributed.all_reduce( + tensor, + op=torch.distributed.ReduceOp.MIN, + group=world.cpu_group, + ) + return int(tensor.item()) + - _, self.state, self.x_cache, self.dt_cache, self.B_cache, *_ = layer.kv_cache - if self.state.shape[0] <= batch: +def _distributed_layers_are_compatible( + layers: tuple[MambaMixer2, ...], +) -> bool: + world = get_world_group() + local_signature = _layer_signature(layers[0]) if layers else None + local_homogeneous = bool(layers) and all( + _layer_signature(layer) == local_signature for layer in layers[1:] + ) + reference = world.broadcast_object( + local_signature if world.rank_in_group == 0 else None, + src=0, + ) + local_ok = local_homogeneous and local_signature == reference + if world.world_size == 1: + return local_ok + flag = torch.tensor([int(local_ok)], dtype=torch.int32) + torch.distributed.all_reduce( + flag, + op=torch.distributed.ReduceOp.MIN, + group=world.cpu_group, + ) + return bool(flag.item()) + + +def _empty_preserve_strides(tensor: torch.Tensor, cache_capacity: int) -> torch.Tensor: + return torch.empty_strided( + (cache_capacity, *tensor.shape[1:]), + tensor.stride(), + dtype=tensor.dtype, + device=tensor.device, + ) + + +class _ReplaySSMTuningCall: + """A private, production-layout ReplaySSM invocation for native tuning.""" + + def __init__(self, layer: MambaMixer2, batch: int): + _, live_state, live_x_cache, live_dt_cache, live_B_cache, *_ = layer.kv_cache + if batch <= 0 or batch >= live_state.shape[0]: raise ValueError( - f"ReplaySSM autotune batch {batch} needs {batch + 1} cache " - f"slots, but only {self.state.shape[0]} are available" + f"ReplaySSM tuning batch {batch} needs {batch + 1} state slots, " + f"but only {live_state.shape[0]} are available" ) - self._kernel = checkpointing_ssu - self.batch = batch - self.logical_window = self.x_cache.shape[2] - 1 - if self.logical_window <= 0: - raise ValueError("ReplaySSM history window must be positive") + # Preserve production inner shapes and every stride. Native FlashInfer + # treats cache capacity as a constrained dimension, so only the active + # slots plus the reserved padding slot need private storage. + private_capacity = batch + 1 + self.state = _empty_preserve_strides(live_state, private_capacity) + self.x_cache = _empty_preserve_strides(live_x_cache, private_capacity) + self.dt_cache = _empty_preserve_strides(live_dt_cache, private_capacity) + self.B_cache = _empty_preserve_strides(live_B_cache, private_capacity) + for tensor in (self.state, self.x_cache, self.dt_cache, self.B_cache): + tensor[: batch + 1].zero_() device = self.state.device activation_dtype = self.x_cache.dtype @@ -297,56 +150,32 @@ def __init__(self, layer: MambaMixer2, batch: int): headdim = self.state.shape[2] dstate = self.state.shape[3] ngroups = self.B_cache.shape[1] - generator = torch.Generator(device=device) - generator.manual_seed(0x5253534D + batch) + self.batch = batch + self.logical_window = self.x_cache.shape[2] - 1 + if self.logical_window <= 0: + raise ValueError("ReplaySSM history window must be positive") - self.x = torch.randn( - batch, - 1, - nheads, - headdim, - dtype=activation_dtype, - device=device, - generator=generator, - ) - dt_base = torch.randn( - batch, - 1, - nheads, - dtype=activation_dtype, - device=device, - generator=generator, - ) - self.dt = dt_base.unsqueeze(-1).expand(batch, 1, nheads, headdim) - self.B = torch.randn( - batch, - 1, - ngroups, - dstate, - dtype=activation_dtype, - device=device, - generator=generator, - ) - self.C = torch.randn( - self.B.shape, - dtype=activation_dtype, - device=device, - generator=generator, - ) - self.out = torch.empty_like(self.x) - self.indices = torch.arange(1, batch + 1, dtype=torch.int32, device=device) self.ring_start = torch.zeros( - self.state.shape[0], dtype=torch.int32, device=device + private_capacity, dtype=torch.int32, device=device ) self.prev_num_accepted = torch.zeros_like(self.ring_start) rows = torch.arange(batch, dtype=torch.int32, device=device) - self.initial_prev_num_accepted = torch.zeros_like(self.prev_num_accepted) - self.initial_prev_num_accepted[1 : batch + 1] = rows.remainder( + self.ring_start[1 : batch + 1] = rows.remainder(self.logical_window + 1) + self.prev_num_accepted[1 : batch + 1] = rows.remainder( self.logical_window ).add_(1) - self.initial_ring_start = torch.zeros_like(self.ring_start) - self.initial_ring_start[1 : batch + 1] = rows.remainder(self.logical_window + 1) + self.indices = torch.arange(1, batch + 1, dtype=torch.int32, device=device) + self.x = torch.zeros( + batch, nheads, headdim, dtype=activation_dtype, device=device + ) + dt_base = torch.zeros(batch, nheads, dtype=activation_dtype, device=device) + self.dt = dt_base.unsqueeze(-1).expand(batch, nheads, headdim) + self.B = torch.zeros( + batch, ngroups, dstate, dtype=activation_dtype, device=device + ) + self.C = torch.zeros_like(self.B) + self.out = torch.empty_like(self.x) self.A = ( layer.A[:, None, ...][:, :, None] .expand(-1, headdim, dstate) @@ -377,379 +206,55 @@ def __init__(self, layer: MambaMixer2, batch: int): device=device, ) - def reset(self) -> None: - for tensor in (self.state, self.x_cache, self.dt_cache, self.B_cache): - tensor[: self.batch + 1].zero_() - self.out.zero_() - self.ring_start.copy_(self.initial_ring_start) - self.prev_num_accepted.copy_(self.initial_prev_num_accepted) - - def call(self, tactic: FlashInferReplaySSMTactic) -> None: - self._kernel( + def run(self) -> None: + selective_state_update_replayssm_flashinfer( self.state, - self.x_cache, - self.B_cache, - self.dt_cache, - self.ring_start, - self.prev_num_accepted, self.x, self.dt, self.A, self.B, self.C, self.out, + self.x_cache, + self.B_cache, + self.dt_cache, + self.ring_start, + self.prev_num_accepted, D=self.D, dt_bias=self.dt_bias, dt_softplus=True, state_batch_indices=self.indices, - pad_slot_id=0, - rand_seed=self.rand_seed, - philox_rounds=self.philox_rounds, cb_scaled=self.cb_scaled, cumAdt_vec=self.cumAdt_vec, cb_old=self.cb_old, - precompute_heads_per_cta=tactic.precompute_heads_per_cta, - algorithm=tactic.algorithm, - ) - update_replayssm_ring_trackers( - self.ring_start, - self.prev_num_accepted, - self.indices, - logical_window=self.logical_window, - pad_slot_id=0, + algorithm="auto", ) - def benchmark(self, tactic: FlashInferReplaySSMTactic) -> float: - calls_per_graph = self.logical_window - graph: torch.cuda.CUDAGraph | None = None - samples = [] - try: - with use_flashinfer_replayssm_tactic(tactic): - self.reset() - side_stream = torch.cuda.Stream() - side_stream.wait_stream(torch.cuda.current_stream()) - with torch.cuda.stream(side_stream): - for _ in range(3): - self.call(tactic) - torch.cuda.current_stream().wait_stream(side_stream) - torch.cuda.synchronize() - - self.reset() - graph = torch.cuda.CUDAGraph() - with torch.cuda.graph(graph): - for _ in range(calls_per_graph): - self.call(tactic) - torch.cuda.synchronize() - - for _ in range(3): - self.reset() - for _ in range(3): - graph.replay() - torch.cuda.synchronize() - start = torch.cuda.Event(enable_timing=True) - end = torch.cuda.Event(enable_timing=True) - start.record() - for _ in range(10): - graph.replay() - end.record() - end.synchronize() - samples.append(start.elapsed_time(end) / (10 * calls_per_graph)) - finally: - self.reset() - del graph - return statistics.median(samples) - - -def _aggregate_timings(timings: list[float]) -> list[float]: - world = get_world_group() - if world.world_size == 1: - return timings - values = torch.tensor(timings, dtype=torch.float64) - torch.distributed.all_reduce( - values, - op=torch.distributed.ReduceOp.MAX, - group=world.cpu_group, - ) - return values.tolist() - - -def _all_ranks_support_tuning(supported: bool) -> bool: - world = get_world_group() - if world.world_size == 1: - return supported - flag = torch.tensor([int(supported)], dtype=torch.int32) - torch.distributed.all_reduce( - flag, - op=torch.distributed.ReduceOp.MIN, - group=world.cpu_group, - ) - return bool(flag.item()) - @torch.inference_mode() -def flashinfer_replayssm_autotune_warmup(worker: Worker) -> None: - runner = worker.model_runner - runner.flashinfer_replayssm_autotune_result = None - if worker.vllm_config.kernel_config.enable_flashinfer_autotune is not True: - return - if worker.model_config.enforce_eager: - logger.info_once("Skipping FlashInfer ReplaySSM autotune without CUDA graphs.") - return - if getattr(worker, "use_v2_model_runner", False): - logger.info_once( - "Skipping FlashInfer ReplaySSM autotune with the V2 model runner." - ) - return - if runner.parallel_config.use_ubatching: - logger.info_once( - "Skipping FlashInfer ReplaySSM autotune with uniform microbatching." - ) - return - +def trigger_flashinfer_replayssm_autotune(runner: GPUModelRunner) -> None: + """Make one maximum-batch call so FlashInfer tunes all decode buckets.""" layers = _find_replayssm_layers(runner) - if not _all_ranks_support_tuning(bool(layers)): - return - assert layers - local_layer_fingerprints = None - try: - local_layer_fingerprints = tuple(_layer_fingerprint(layer) for layer in layers) - except Exception: - logger.warning( - "Could not fingerprint the FlashInfer ReplaySSM layers.", - exc_info=True, - ) - if not _all_ranks_support_tuning(local_layer_fingerprints is not None): - logger.warning_once( - "Skipping FlashInfer ReplaySSM autotune because layer " - "fingerprinting failed on at least one rank." - ) - return - assert local_layer_fingerprints is not None - layers_are_homogeneous = all( - fingerprint == local_layer_fingerprints[0] - for fingerprint in local_layer_fingerprints[1:] - ) - if not _all_ranks_support_tuning(layers_are_homogeneous): - logger.warning_once( - "Skipping FlashInfer ReplaySSM autotune because ReplaySSM layer " - "geometries differ within a rank." - ) - return - layer = layers[0] - T = runner.uniform_decode_query_len - if T != 1: + if not _distributed_layers_are_compatible(layers): logger.warning_once( - "Skipping FlashInfer ReplaySSM autotune for T=%d; this vLLM " - "integration currently supports only T=1.", - T, + "Skipping native FlashInfer ReplaySSM autotuning because ReplaySSM " + "layers are absent or have incompatible rank-local geometries." ) return - local_batches = _uniform_capture_batches(runner) - world = get_world_group() - batches = world.broadcast_object( - local_batches if world.rank_in_group == 0 else None, src=0 - ) - if not _all_ranks_support_tuning(local_batches == batches): - logger.warning_once( - "Skipping FlashInfer ReplaySSM autotune because CUDA-graph " - "descriptors differ across ranks." - ) - return - if not batches: - logger.info_once( - "Skipping FlashInfer ReplaySSM autotune because there are no " - "uniform FULL CUDA-graph capture descriptors." - ) - return - - local_fingerprint = local_layer_fingerprints[0] - fingerprint = world.broadcast_object( - local_fingerprint if world.rank_in_group == 0 else None, src=0 - ) - if not _all_ranks_support_tuning(local_fingerprint == fingerprint): - logger.warning_once( - "Skipping FlashInfer ReplaySSM autotune because Mamba geometry " - "differs across ranks." - ) - return - - cache_path = None - try: - cache_path = resolve_flashinfer_autotune_file(runner).with_name( - _CACHE_FILE_NAME - ) - except Exception: - logger.warning( - "Could not resolve the FlashInfer ReplaySSM autotune cache path.", - exc_info=True, - ) - if not _all_ranks_support_tuning(cache_path is not None): + layer = layers[0] + state_capacity = layer.kv_cache[1].shape[0] + batch = _distributed_min(_local_max_tuning_batch(runner, state_capacity)) + if batch <= 0: logger.warning_once( - "Skipping FlashInfer ReplaySSM autotune because the cache path " - "could not be resolved on every rank." + "Skipping native FlashInfer ReplaySSM autotuning because no valid " + "decode batch fits in the state cache." ) return - assert cache_path is not None - cached_entries = _load_cache(cache_path) if world.rank_in_group == 0 else None - cached_entries = world.broadcast_object(cached_entries, src=0) - assert cached_entries is not None - - selected: dict[int, FlashInferReplaySSMTactic] = {} - cache_changed = False - tuning_started = time.perf_counter() - for batch in batches: - cache_key = _make_cache_key(fingerprint, batch, T) - cached_name = cached_entries.get(cache_key) - cached_tactic = ( - _tactic_from_name(cached_name) if cached_name is not None else None - ) - if cached_tactic is not None: - tactic = cached_tactic - selected[batch] = tactic - if world.rank_in_group == 0: - logger.info( - "FlashInfer ReplaySSM autotune cache hit for batch %d: %s", - batch, - tactic.name, - ) - continue - - benchmark = None - try: - benchmark = _ReplaySSMBenchmark(layer, batch) - except Exception: - logger.warning( - "Could not construct the FlashInfer ReplaySSM benchmark for batch %d.", - batch, - exc_info=True, - ) - if not _all_ranks_support_tuning(benchmark is not None): - logger.warning_once( - "Skipping FlashInfer ReplaySSM autotune for batch %d because " - "the benchmark could not be constructed on every rank.", - batch, - ) - del benchmark - gc.collect() - torch.cuda.empty_cache() - continue - assert benchmark is not None - - candidate_count = len(FLASHINFER_REPLAYSSM_TUNING_CANDIDATES) - shift = batch % candidate_count - candidate_indices = tuple(range(candidate_count)) - candidate_order = candidate_indices[shift:] + candidate_indices[:shift] - if (batch // candidate_count) % 2: - candidate_order = tuple(reversed(candidate_order)) - local_timings = [float("inf")] * candidate_count - for candidate_index in candidate_order: - tactic = FLASHINFER_REPLAYSSM_TUNING_CANDIDATES[candidate_index] - try: - local_timings[candidate_index] = benchmark.benchmark(tactic) - except Exception: - logger.warning( - "FlashInfer ReplaySSM tactic %s failed for batch %d.", - tactic.name, - batch, - exc_info=True, - ) - timings = _aggregate_timings(local_timings) - _, _, _, _, B_cache, *_ = layer.kv_cache - heads_per_group = layer.kv_cache[1].shape[1] // B_cache.shape[1] - expanded_candidates = _expanded_precompute_tactics( - timings, heads_per_group - ) - local_expanded_timings = [float("inf")] * len(expanded_candidates) - expanded_count = len(expanded_candidates) - if expanded_count: - expanded_shift = batch % expanded_count - expanded_order = tuple(range(expanded_count)) - expanded_order = ( - expanded_order[expanded_shift:] - + expanded_order[:expanded_shift] - ) - for candidate_index in expanded_order: - expanded_tactic = expanded_candidates[candidate_index] - try: - local_expanded_timings[candidate_index] = benchmark.benchmark( - expanded_tactic - ) - except Exception: - logger.warning( - "FlashInfer ReplaySSM tactic %s failed for batch %d.", - expanded_tactic.name, - batch, - exc_info=True, - ) - expanded_timings = _aggregate_timings(local_expanded_timings) - all_candidates = ( - FLASHINFER_REPLAYSSM_TUNING_CANDIDATES + expanded_candidates - ) - all_timings = timings + expanded_timings - tactic = _select_fastest(all_timings, all_candidates) - if tactic is None: - logger.warning_once( - "Every FlashInfer ReplaySSM tactic failed for batch %d; " - "leaving the default launch policy unchanged.", - batch, - ) - del benchmark - gc.collect() - torch.cuda.empty_cache() - continue - selected[batch] = tactic - cached_entries[cache_key] = tactic.name - cache_changed = True - if world.rank_in_group == 0: - timing_log = ", ".join( - f"{candidate.name}={timing:.6f}ms" - for candidate, timing in zip( - all_candidates, - all_timings, - strict=True, - ) - ) - logger.info( - "FlashInfer ReplaySSM autotune selected batch %d: %s (%s)", - batch, - tactic.name, - timing_log, - ) - del benchmark - gc.collect() - torch.cuda.empty_cache() - - if cache_changed and world.rank_in_group == 0: - try: - _save_cache(cache_path, cached_entries) - except Exception: - logger.warning( - "Could not save the FlashInfer ReplaySSM autotune cache to %s.", - cache_path, - exc_info=True, - ) - runner.flashinfer_replayssm_autotune_result = FlashInferReplaySSMAutotuneResult( - T, selected - ) - if world.rank_in_group == 0: - logger.info( - "FlashInfer ReplaySSM autotune prepared %d CUDA-graph batch " - "tactics in %.2f seconds.", - len(selected), - time.perf_counter() - tuning_started, - ) - -@contextmanager -def use_flashinfer_replayssm_tactic_for_capture( - runner: GPUModelRunner, - batch_descriptor: BatchDescriptor, -): - result = runner.flashinfer_replayssm_autotune_result - tactic = result.tactic_for(batch_descriptor) if result is not None else None - scope = ( - use_flashinfer_replayssm_tactic(tactic) if tactic is not None else nullcontext() + tuning_call = _ReplaySSMTuningCall(layer, batch) + tuning_call.run() + torch.cuda.synchronize() + logger.info_once( + "Triggered native FlashInfer ReplaySSM autotuning through batch %d.", batch ) - with scope: - yield diff --git a/vllm/model_executor/warmup/kernel_warmup.py b/vllm/model_executor/warmup/kernel_warmup.py index bdc075340502..0a8f163752c7 100644 --- a/vllm/model_executor/warmup/kernel_warmup.py +++ b/vllm/model_executor/warmup/kernel_warmup.py @@ -28,7 +28,7 @@ write_flashinfer_autotune_cache, ) from vllm.model_executor.warmup.flashinfer_replayssm_warmup import ( - flashinfer_replayssm_autotune_warmup, + trigger_flashinfer_replayssm_autotune, ) from vllm.model_executor.warmup.flashinfer_sparse_mla_warmup import ( deepseek_v4_sparse_mla_attention_warmup, @@ -197,7 +197,6 @@ def kernel_warmup(worker: "Worker", *, process_local_only: bool = False): logger.info_once("Skipping FlashInfer autotune because it is disabled.") elif has_flashinfer() and current_platform.has_device_capability(90): flashinfer_autotune(worker.model_runner) - flashinfer_replayssm_autotune_warmup(worker) # FlashInfer attention warmup # Only warmup if the model has FlashInfer attention groups @@ -336,6 +335,11 @@ def flashinfer_autotune(runner: "GPUModelRunner") -> None: fi_utils.autotune(tune_mode=True, **autotune_kwargs), ): _run_flashinfer_autotune_dummy_runs(runner) + # The generic dummy run is predominantly prefill-shaped and may + # not execute ReplaySSM decode. One private maximum-batch decode + # call lets FlashInfer populate every native dynamic batch bucket + # before CUDA graph capture. + trigger_flashinfer_replayssm_autotune(runner) finally: set_autotune_process_group(None) diff --git a/vllm/v1/worker/gpu_model_runner.py b/vllm/v1/worker/gpu_model_runner.py index 60f6fab81f28..b5e591dd067c 100644 --- a/vllm/v1/worker/gpu_model_runner.py +++ b/vllm/v1/worker/gpu_model_runner.py @@ -912,7 +912,6 @@ def __init__( # Cudagraph dispatcher for runtime cudagraph dispatching. self.cudagraph_dispatcher = CudagraphDispatcher(self.vllm_config) - self.flashinfer_replayssm_autotune_result: Any | None = None self.mm_budget = ( MultiModalBudget(self.vllm_config, self.mm_registry) @@ -7047,10 +7046,6 @@ def _capture_cudagraphs( if not batch_descriptors: return - from vllm.model_executor.warmup.flashinfer_replayssm_warmup import ( - use_flashinfer_replayssm_tactic_for_capture, - ) - uniform_decode = batch_descriptors[0].uniform # Only rank 0 should print progress bar during capture @@ -7080,13 +7075,12 @@ def _capture_cudagraphs( uniform_decode=uniform_decode, ) ) - with use_flashinfer_replayssm_tactic_for_capture(self, batch_desc): - self._warmup_and_capture( - batch_desc, - cudagraph_runtime_mode=cudagraph_runtime_mode, - allow_microbatching=allow_microbatching, - profiler=profiler, - ) + self._warmup_and_capture( + batch_desc, + cudagraph_runtime_mode=cudagraph_runtime_mode, + allow_microbatching=allow_microbatching, + profiler=profiler, + ) torch.accelerator.synchronize() self.maybe_remove_all_loras(self.lora_config) From 1cab2492ae4e49fcfa1e478b7196b430b0fde847 Mon Sep 17 00:00:00 2001 From: Andrii Skliar Date: Sun, 16 Aug 2026 07:05:12 -0700 Subject: [PATCH 14/33] Simplify FlashInfer ReplaySSM autotune warmup Signed-off-by: Andrii Skliar --- tests/kernels/mamba/test_ssu_dispatch.py | 78 +++++- .../test_flashinfer_replayssm_warmup.py | 154 ----------- tests/model_executor/test_kernel_warmup.py | 72 +++++ .../layers/mamba/ops/ssu_dispatch.py | 51 +--- .../warmup/flashinfer_replayssm_warmup.py | 260 ------------------ vllm/model_executor/warmup/kernel_warmup.py | 67 ++++- 6 files changed, 213 insertions(+), 469 deletions(-) delete mode 100644 tests/model_executor/test_flashinfer_replayssm_warmup.py create mode 100644 tests/model_executor/test_kernel_warmup.py delete mode 100644 vllm/model_executor/warmup/flashinfer_replayssm_warmup.py diff --git a/tests/kernels/mamba/test_ssu_dispatch.py b/tests/kernels/mamba/test_ssu_dispatch.py index d852908a5a2f..c826f55fbb69 100644 --- a/tests/kernels/mamba/test_ssu_dispatch.py +++ b/tests/kernels/mamba/test_ssu_dispatch.py @@ -10,7 +10,9 @@ from vllm.config.mamba import MambaBackendEnum, MambaConfig, MambaSSUAlgorithm from vllm.model_executor.layers.mamba.ops.ssu_dispatch import ( + FLASHINFER_REPLAYSSM_AUTO_TACTIC, FlashInferReplaySSMBackend, + FlashInferReplaySSMTactic, FlashInferSSUBackend, TritonReplaySSMBackend, TritonSSUBackend, @@ -22,6 +24,7 @@ selective_state_update_replayssm_flashinfer, selective_state_update_replayssm_triton, update_replayssm_ring_trackers, + use_flashinfer_replayssm_tactic, ) from vllm.utils.torch_utils import set_random_seed from vllm.v1.attention.backends.registry import MambaAttentionBackendEnum @@ -42,19 +45,69 @@ checkpointing_ssu_module = importlib.import_module( "flashinfer.mamba.checkpointing_ssu" ) - HAS_FLASHINFER_CHECKPOINTING_SSU = ( - hasattr(checkpointing_ssu_module, "CheckpointingSSURunner") - and getattr( - checkpointing_ssu_module, - "CHECKPOINTING_SSU_AUTOTUNE_ABI_VERSION", - 0, - ) - >= 1 + HAS_FLASHINFER_CHECKPOINTING_SSU = hasattr( + checkpointing_ssu_module, "CheckpointingSSURunner" ) except ImportError: HAS_FLASHINFER_CHECKPOINTING_SSU = False +def _fake_flashinfer_replayssm_backend() -> FlashInferReplaySSMBackend: + backend = FlashInferReplaySSMBackend.__new__(FlashInferReplaySSMBackend) + backend._mamba_config = MambaConfig(backend=MambaBackendEnum.FLASHINFER) + backend._kernel = Mock(return_value=torch.empty(1)) + backend._tactic = FLASHINFER_REPLAYSSM_AUTO_TACTIC + return backend + + +def test_replayssm_explicit_tactic_validation(): + tactic = FlashInferReplaySSMTactic( + "two-kernel", d_split=2, precompute_heads_per_cta=8 + ) + assert tactic.name == "two_kernel_d2_h8" + with pytest.raises(ValueError, match="does not accept precompute"): + FlashInferReplaySSMTactic("monolith", precompute_heads_per_cta=8) + + +def test_replayssm_tactic_scope_restores_direct_launch_controls(monkeypatch): + import vllm.model_executor.layers.mamba.ops.ssu_dispatch as mod + + backend = _fake_flashinfer_replayssm_backend() + monkeypatch.setattr(mod, "_replayssm_backend", backend) + tactic = FlashInferReplaySSMTactic( + "two-kernel", d_split=2, precompute_heads_per_cta=8 + ) + + with ( + pytest.raises(RuntimeError, match="sentinel"), + use_flashinfer_replayssm_tactic(tactic), + ): + assert backend._tactic is tactic + raise RuntimeError("sentinel") + + assert backend._tactic is FLASHINFER_REPLAYSSM_AUTO_TACTIC + + +def test_replayssm_backend_forwards_explicit_tactic(monkeypatch): + import vllm.model_executor.layers.mamba.ops.ssu_dispatch as mod + + backend = _fake_flashinfer_replayssm_backend() + monkeypatch.setattr(mod, "_replayssm_backend", backend) + tensor = torch.empty(1) + + with use_flashinfer_replayssm_tactic( + FlashInferReplaySSMTactic("two-kernel", d_split=2, precompute_heads_per_cta=8) + ): + backend(*(tensor,) * 12) + + kwargs = backend._kernel.call_args.kwargs + assert kwargs["algorithm"] == "two-kernel" + assert kwargs["d_split"] == 2 + assert kwargs["precompute_heads_per_cta"] == 8 + assert "main_pipeline_stages" not in kwargs + assert "main_ctas_per_sm" not in kwargs + + @pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") def test_flashinfer_replayssm_ring_tracker_lifecycle(): ring_start = torch.zeros(2, dtype=torch.int32, device="cuda") @@ -371,17 +424,16 @@ def test_replayssm_flashinfer_call(monkeypatch): assert args[4] is ring_start assert args[5] is prev_num_accepted assert kwargs["algorithm"] == "auto" + assert kwargs["d_split"] is None assert kwargs["precompute_heads_per_cta"] == 0 - assert kwargs["main_pipeline_stages"] == 0 - assert kwargs["main_ctas_per_sm"] == 0 + assert "main_pipeline_stages" not in kwargs + assert "main_ctas_per_sm" not in kwargs def test_replayssm_requires_native_flashinfer_autotuning(monkeypatch): import vllm.model_executor.layers.mamba.ops.ssu_dispatch as mod - old_module = SimpleNamespace( - checkpointing_ssu=Mock(), CheckpointingSSURunner=object - ) + old_module = SimpleNamespace(checkpointing_ssu=Mock()) monkeypatch.setattr(mod.importlib, "import_module", lambda _: old_module) with pytest.raises(ImportError, match="native checkpointing_ssu autotuning"): FlashInferReplaySSMBackend(MambaConfig(backend=MambaBackendEnum.FLASHINFER)) diff --git a/tests/model_executor/test_flashinfer_replayssm_warmup.py b/tests/model_executor/test_flashinfer_replayssm_warmup.py deleted file mode 100644 index 4a5c9e3d0da3..000000000000 --- a/tests/model_executor/test_flashinfer_replayssm_warmup.py +++ /dev/null @@ -1,154 +0,0 @@ -# SPDX-License-Identifier: Apache-2.0 -# SPDX-FileCopyrightText: Copyright contributors to the vLLM project - -from types import SimpleNamespace -from unittest.mock import Mock - -import pytest -import torch - -from vllm.config.mamba import MambaBackendEnum, MambaConfig -from vllm.forward_context import BatchDescriptor -from vllm.model_executor.layers.mamba.ops import ssu_dispatch -from vllm.model_executor.layers.mamba.ops.ssu_dispatch import ( - FLASHINFER_REPLAYSSM_AUTO_TACTIC, - FlashInferReplaySSMBackend, - FlashInferReplaySSMTactic, - use_flashinfer_replayssm_tactic, -) -from vllm.model_executor.warmup import flashinfer_replayssm_warmup as warmup - - -def _fake_backend() -> FlashInferReplaySSMBackend: - backend = FlashInferReplaySSMBackend.__new__(FlashInferReplaySSMBackend) - backend._mamba_config = MambaConfig(backend=MambaBackendEnum.FLASHINFER) - backend._kernel = Mock(return_value=torch.empty(1)) - backend._tactic = FLASHINFER_REPLAYSSM_AUTO_TACTIC - return backend - - -def _fake_layer(cache_slots: int = 5): - nheads, headdim, dstate, ngroups = 2, 4, 8, 1 - return SimpleNamespace( - kv_cache=( - torch.empty(0), - torch.empty(cache_slots, nheads, headdim, dstate), - torch.empty(cache_slots, nheads, 17, headdim), - torch.empty(cache_slots, nheads, 17), - torch.empty(cache_slots, ngroups, 17, dstate), - ), - A=torch.empty(nheads), - D=torch.empty(nheads), - dt_bias=torch.empty(nheads), - mamba_config=SimpleNamespace( - enable_stochastic_rounding=False, - stochastic_rounding_philox_rounds=None, - ), - ) - - -def test_replayssm_explicit_tactic_validation(): - tactic = FlashInferReplaySSMTactic("two-kernel", 1, 4, precompute_heads_per_cta=8) - assert tactic.name == "two_kernel_s1_c4_h8" - with pytest.raises(ValueError, match="does not accept precompute"): - FlashInferReplaySSMTactic("monolith", precompute_heads_per_cta=8) - - -def test_replayssm_tactic_scope_restores_direct_launch_controls(monkeypatch): - backend = _fake_backend() - monkeypatch.setattr(ssu_dispatch, "_replayssm_backend", backend) - tactic = FlashInferReplaySSMTactic("two-kernel", 2, 16, precompute_heads_per_cta=8) - - with ( - pytest.raises(RuntimeError, match="sentinel"), - use_flashinfer_replayssm_tactic(tactic), - ): - assert backend._tactic is tactic - raise RuntimeError("sentinel") - - assert backend._tactic is FLASHINFER_REPLAYSSM_AUTO_TACTIC - - -def test_replayssm_backend_forwards_explicit_tactic(monkeypatch): - backend = _fake_backend() - monkeypatch.setattr(ssu_dispatch, "_replayssm_backend", backend) - tensor = torch.empty(1) - - with use_flashinfer_replayssm_tactic( - FlashInferReplaySSMTactic("two-kernel", 1, 4, precompute_heads_per_cta=8) - ): - backend( - tensor, - tensor, - tensor, - tensor, - tensor, - tensor, - tensor, - tensor, - tensor, - tensor, - tensor, - tensor, - ) - - kwargs = backend._kernel.call_args.kwargs - assert kwargs["algorithm"] == "two-kernel" - assert kwargs["precompute_heads_per_cta"] == 8 - assert kwargs["main_pipeline_stages"] == 1 - assert kwargs["main_ctas_per_sm"] == 4 - - -def test_replayssm_tuning_call_uses_private_production_layout(): - layer = _fake_layer() - call = warmup._ReplaySSMTuningCall(layer, 3) - - assert call.state.shape == (4, *layer.kv_cache[1].shape[1:]) - assert call.state.stride() == layer.kv_cache[1].stride() - assert call.state.data_ptr() != layer.kv_cache[1].data_ptr() - assert call.x_cache.data_ptr() != layer.kv_cache[2].data_ptr() - assert call.indices.tolist() == [1, 2, 3] - assert call.ring_start.shape == (4,) - assert call.prev_num_accepted.tolist() == [0, 1, 2, 3] - assert call.dt.shape == (3, 2, 4) - assert call.dt.stride(-1) == 0 - - -def test_replayssm_tuning_trigger_uses_largest_supported_decode_batch(monkeypatch): - layer = _fake_layer(cache_slots=5) - runner = SimpleNamespace( - scheduler_config=SimpleNamespace(max_num_seqs=8), - cudagraph_dispatcher=SimpleNamespace( - get_capture_descs=lambda: [ - ( - None, - [ - BatchDescriptor(num_tokens=2, num_reqs=2, uniform=True), - BatchDescriptor(num_tokens=8, num_reqs=8, uniform=True), - ], - ) - ] - ), - ) - observed = SimpleNamespace(batch=None, ran=False) - - class FakeCall: - def __init__(self, _layer, batch): - assert _layer is layer - observed.batch = batch - - def run(self): - observed.ran = True - - monkeypatch.setattr(warmup, "_find_replayssm_layers", lambda _runner: (layer,)) - monkeypatch.setattr( - warmup, "_distributed_layers_are_compatible", lambda _layers: True - ) - monkeypatch.setattr(warmup, "_distributed_min", lambda value: value) - monkeypatch.setattr(warmup, "_ReplaySSMTuningCall", FakeCall) - monkeypatch.setattr(torch.cuda, "synchronize", lambda: None) - - warmup.trigger_flashinfer_replayssm_autotune(runner) - - assert observed.batch == 4 - assert observed.ran diff --git a/tests/model_executor/test_kernel_warmup.py b/tests/model_executor/test_kernel_warmup.py new file mode 100644 index 000000000000..c25d9cfcb285 --- /dev/null +++ b/tests/model_executor/test_kernel_warmup.py @@ -0,0 +1,72 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project + +from types import SimpleNamespace +from unittest.mock import Mock + +import numpy as np +import pytest + +from vllm.model_executor.layers.mamba.ops import ssu_dispatch +from vllm.model_executor.warmup import kernel_warmup as warmup + + +@pytest.mark.parametrize( + "backend_name, expected_calls", [("flashinfer", 1), ("triton", 0)] +) +def test_replayssm_autotune_uses_uniform_decode( + monkeypatch, backend_name, expected_calls +): + block_ids = np.zeros((32, 1), dtype=np.int32) + block_table = SimpleNamespace(block_table=SimpleNamespace(np=block_ids)) + multi_group_block_table = SimpleNamespace( + block_tables=[block_table], commit_block_table=Mock() + ) + + def dummy_run(**kwargs): + assert block_ids[:16, 0].tolist() == list(range(1, 17)) + + dummy_run = Mock(side_effect=dummy_run) + runner = SimpleNamespace( + uniform_decode_query_len=6, + max_num_tokens=100, + scheduler_config=SimpleNamespace(max_num_seqs=32), + input_batch=SimpleNamespace(block_table=multi_group_block_table), + get_model=lambda: SimpleNamespace(modules=lambda: ()), + _dummy_run=dummy_run, + ) + monkeypatch.setattr( + ssu_dispatch, + "get_replayssm_backend", + lambda: SimpleNamespace(name=backend_name), + ) + + warmup._flashinfer_replayssm_autotune_dummy_run(runner) + + assert dummy_run.call_count == expected_calls + if expected_calls: + dummy_run.assert_called_once_with( + num_tokens=96, + uniform_decode=True, + allow_microbatching=False, + skip_eplb=True, + is_profile=True, + randomize_inputs=True, + force_attention=True, + profile_seq_lens=7, + ) + assert not block_ids.any() + multi_group_block_table.commit_block_table.assert_called_once_with(16) + + +def test_replayssm_autotune_skips_uninitialized_backend(monkeypatch): + runner = SimpleNamespace(_dummy_run=Mock()) + + def get_backend(): + raise RuntimeError + + monkeypatch.setattr(ssu_dispatch, "get_replayssm_backend", get_backend) + + warmup._flashinfer_replayssm_autotune_dummy_run(runner) + + runner._dummy_run.assert_not_called() diff --git a/vllm/model_executor/layers/mamba/ops/ssu_dispatch.py b/vllm/model_executor/layers/mamba/ops/ssu_dispatch.py index fba515b2b375..373beb2c7b5c 100644 --- a/vllm/model_executor/layers/mamba/ops/ssu_dispatch.py +++ b/vllm/model_executor/layers/mamba/ops/ssu_dispatch.py @@ -37,39 +37,26 @@ @dataclass(frozen=True) class FlashInferReplaySSMTactic: algorithm: str - pipeline_stages: int | None = None - ctas_per_sm: int | None = None + d_split: int | None = None precompute_heads_per_cta: int = 0 def __post_init__(self) -> None: if self.algorithm not in {"auto", "monolith", "two-kernel"}: raise ValueError(f"Unsupported ReplaySSM algorithm: {self.algorithm}") - has_launch_config = ( - self.pipeline_stages is not None or self.ctas_per_sm is not None - ) - if self.algorithm == "two-kernel": - if self.pipeline_stages not in {1, 2}: - raise ValueError("two-kernel requires pipeline_stages in {1, 2}") - if self.ctas_per_sm is None or self.ctas_per_sm <= 0: - raise ValueError("two-kernel requires a positive ctas_per_sm") - if self.precompute_heads_per_cta < 0: - raise ValueError( - "two-kernel requires non-negative precompute_heads_per_cta" - ) - elif has_launch_config: - raise ValueError( - f"{self.algorithm} does not accept pipeline or CTA settings" - ) - elif self.precompute_heads_per_cta != 0: + if self.d_split is not None and self.d_split <= 0: + raise ValueError("d_split must be positive when specified") + if self.precompute_heads_per_cta < 0: + raise ValueError("precompute_heads_per_cta must be non-negative") + if self.algorithm == "monolith" and self.precompute_heads_per_cta != 0: raise ValueError( f"{self.algorithm} does not accept precompute_heads_per_cta" ) @property def name(self) -> str: - if self.algorithm != "two-kernel": - return self.algorithm - name = f"two_kernel_s{self.pipeline_stages}_c{self.ctas_per_sm}" + name = self.algorithm.replace("-", "_") + if self.d_split is not None: + name += f"_d{self.d_split}" if self.precompute_heads_per_cta: name += f"_h{self.precompute_heads_per_cta}" return name @@ -504,20 +491,11 @@ def __init__(self, mamba_config: MambaConfig): "Please install flashinfer with mamba.checkpointing_ssu support: " "pip install flashinfer-python" ) from e - autotune_abi = getattr( - checkpointing_ssu_module, - "CHECKPOINTING_SSU_AUTOTUNE_ABI_VERSION", - 0, - ) - if ( - not hasattr(checkpointing_ssu_module, "CheckpointingSSURunner") - or autotune_abi < 1 - ): + if not hasattr(checkpointing_ssu_module, "CheckpointingSSURunner"): raise ImportError( "FlashInfer ReplaySSM requires native checkpointing_ssu " - "autotuning ABI version 1 or newer. Install a compatible " - "FlashInfer revision exposing " - "CHECKPOINTING_SSU_AUTOTUNE_ABI_VERSION >= 1." + "autotuning support. Install a compatible FlashInfer revision " + "exposing CheckpointingSSURunner." ) self._kernel = checkpointing_ssu_module.checkpointing_ssu self._tactic = FLASHINFER_REPLAYSSM_AUTO_TACTIC @@ -598,13 +576,10 @@ def __call__( cb_scaled=cb_scaled, cumAdt_vec=cumAdt_vec, cb_old=cb_old, + d_split=tactic.d_split if use_explicit_tactic else None, precompute_heads_per_cta=( tactic.precompute_heads_per_cta if use_explicit_tactic else 0 ), - main_pipeline_stages=( - tactic.pipeline_stages or 0 if use_explicit_tactic else 0 - ), - main_ctas_per_sm=(tactic.ctas_per_sm or 0 if use_explicit_tactic else 0), algorithm=requested_algorithm, ) if update_trackers and indices is not None: diff --git a/vllm/model_executor/warmup/flashinfer_replayssm_warmup.py b/vllm/model_executor/warmup/flashinfer_replayssm_warmup.py deleted file mode 100644 index 2160cad8055d..000000000000 --- a/vllm/model_executor/warmup/flashinfer_replayssm_warmup.py +++ /dev/null @@ -1,260 +0,0 @@ -# SPDX-License-Identifier: Apache-2.0 -# SPDX-FileCopyrightText: Copyright contributors to the vLLM project -"""Trigger native FlashInfer ReplaySSM tuning before CUDA graph capture.""" - -from __future__ import annotations - -from typing import TYPE_CHECKING, Any - -import torch - -from vllm.config.mamba import MambaBackendEnum -from vllm.distributed.parallel_state import get_world_group -from vllm.logger import init_logger -from vllm.model_executor.layers.mamba.ops.ssu_dispatch import ( - selective_state_update_replayssm_flashinfer, -) - -if TYPE_CHECKING: - from vllm.model_executor.layers.mamba.mamba_mixer2 import MambaMixer2 - from vllm.v1.worker.gpu_model_runner import GPUModelRunner - -logger = init_logger(__name__) - - -def _find_replayssm_layers(runner: GPUModelRunner) -> tuple[MambaMixer2, ...]: - from vllm.model_executor.layers.mamba.mamba_mixer2 import MambaMixer2 - - return tuple( - module - for module in runner.get_model().modules() - if ( - isinstance(module, MambaMixer2) - and module.use_replayssm - and module.mamba_config.backend == MambaBackendEnum.FLASHINFER - ) - ) - - -def _layer_signature(layer: MambaMixer2) -> tuple[Any, ...]: - _, state, x_cache, dt_cache, B_cache, *_ = layer.kv_cache - tensors = (state, x_cache, dt_cache, B_cache, layer.A, layer.D, layer.dt_bias) - return tuple( - (tuple(tensor.shape), tuple(tensor.stride()), tensor.dtype) - for tensor in tensors - ) + ( - layer.mamba_config.enable_stochastic_rounding, - layer.mamba_config.stochastic_rounding_philox_rounds, - ) - - -def _uniform_decode_batches(runner: GPUModelRunner) -> tuple[int, ...]: - return tuple( - sorted( - { - desc.num_reqs - for _, descs in runner.cudagraph_dispatcher.get_capture_descs() - for desc in descs - if desc.uniform and desc.num_reqs is not None - } - ) - ) - - -def _local_max_tuning_batch(runner: GPUModelRunner, state_capacity: int) -> int: - capture_batches = _uniform_decode_batches(runner) - requested = ( - capture_batches[-1] if capture_batches else runner.scheduler_config.max_num_seqs - ) - return max( - 0, - min( - requested, - runner.scheduler_config.max_num_seqs, - state_capacity - 1, - ), - ) - - -def _distributed_min(value: int) -> int: - world = get_world_group() - if world.world_size == 1: - return value - tensor = torch.tensor([value], dtype=torch.int64) - torch.distributed.all_reduce( - tensor, - op=torch.distributed.ReduceOp.MIN, - group=world.cpu_group, - ) - return int(tensor.item()) - - -def _distributed_layers_are_compatible( - layers: tuple[MambaMixer2, ...], -) -> bool: - world = get_world_group() - local_signature = _layer_signature(layers[0]) if layers else None - local_homogeneous = bool(layers) and all( - _layer_signature(layer) == local_signature for layer in layers[1:] - ) - reference = world.broadcast_object( - local_signature if world.rank_in_group == 0 else None, - src=0, - ) - local_ok = local_homogeneous and local_signature == reference - if world.world_size == 1: - return local_ok - flag = torch.tensor([int(local_ok)], dtype=torch.int32) - torch.distributed.all_reduce( - flag, - op=torch.distributed.ReduceOp.MIN, - group=world.cpu_group, - ) - return bool(flag.item()) - - -def _empty_preserve_strides(tensor: torch.Tensor, cache_capacity: int) -> torch.Tensor: - return torch.empty_strided( - (cache_capacity, *tensor.shape[1:]), - tensor.stride(), - dtype=tensor.dtype, - device=tensor.device, - ) - - -class _ReplaySSMTuningCall: - """A private, production-layout ReplaySSM invocation for native tuning.""" - - def __init__(self, layer: MambaMixer2, batch: int): - _, live_state, live_x_cache, live_dt_cache, live_B_cache, *_ = layer.kv_cache - if batch <= 0 or batch >= live_state.shape[0]: - raise ValueError( - f"ReplaySSM tuning batch {batch} needs {batch + 1} state slots, " - f"but only {live_state.shape[0]} are available" - ) - - # Preserve production inner shapes and every stride. Native FlashInfer - # treats cache capacity as a constrained dimension, so only the active - # slots plus the reserved padding slot need private storage. - private_capacity = batch + 1 - self.state = _empty_preserve_strides(live_state, private_capacity) - self.x_cache = _empty_preserve_strides(live_x_cache, private_capacity) - self.dt_cache = _empty_preserve_strides(live_dt_cache, private_capacity) - self.B_cache = _empty_preserve_strides(live_B_cache, private_capacity) - for tensor in (self.state, self.x_cache, self.dt_cache, self.B_cache): - tensor[: batch + 1].zero_() - - device = self.state.device - activation_dtype = self.x_cache.dtype - nheads = self.state.shape[1] - headdim = self.state.shape[2] - dstate = self.state.shape[3] - ngroups = self.B_cache.shape[1] - self.batch = batch - self.logical_window = self.x_cache.shape[2] - 1 - if self.logical_window <= 0: - raise ValueError("ReplaySSM history window must be positive") - - self.ring_start = torch.zeros( - private_capacity, dtype=torch.int32, device=device - ) - self.prev_num_accepted = torch.zeros_like(self.ring_start) - rows = torch.arange(batch, dtype=torch.int32, device=device) - self.ring_start[1 : batch + 1] = rows.remainder(self.logical_window + 1) - self.prev_num_accepted[1 : batch + 1] = rows.remainder( - self.logical_window - ).add_(1) - self.indices = torch.arange(1, batch + 1, dtype=torch.int32, device=device) - - self.x = torch.zeros( - batch, nheads, headdim, dtype=activation_dtype, device=device - ) - dt_base = torch.zeros(batch, nheads, dtype=activation_dtype, device=device) - self.dt = dt_base.unsqueeze(-1).expand(batch, nheads, headdim) - self.B = torch.zeros( - batch, ngroups, dstate, dtype=activation_dtype, device=device - ) - self.C = torch.zeros_like(self.B) - self.out = torch.empty_like(self.x) - self.A = ( - layer.A[:, None, ...][:, :, None] - .expand(-1, headdim, dstate) - .to(dtype=torch.float32) - ) - self.D = layer.D[:, None, ...].expand(-1, headdim) - self.dt_bias = layer.dt_bias[:, None, ...].expand(-1, headdim) - self.rand_seed = ( - torch.zeros(1, dtype=torch.int64, device=device) - if layer.mamba_config.enable_stochastic_rounding - else None - ) - self.philox_rounds = layer.mamba_config.stochastic_rounding_philox_rounds or 10 - - k_old = ((self.logical_window + 7) // 8) * 8 - self.cb_scaled = torch.empty( - batch, nheads, 32, 8, dtype=activation_dtype, device=device - ) - self.cumAdt_vec = torch.empty( - batch, nheads, 16, dtype=torch.float32, device=device - ) - self.cb_old = torch.empty( - batch, - nheads, - 32, - k_old // 2, - dtype=activation_dtype, - device=device, - ) - - def run(self) -> None: - selective_state_update_replayssm_flashinfer( - self.state, - self.x, - self.dt, - self.A, - self.B, - self.C, - self.out, - self.x_cache, - self.B_cache, - self.dt_cache, - self.ring_start, - self.prev_num_accepted, - D=self.D, - dt_bias=self.dt_bias, - dt_softplus=True, - state_batch_indices=self.indices, - cb_scaled=self.cb_scaled, - cumAdt_vec=self.cumAdt_vec, - cb_old=self.cb_old, - algorithm="auto", - ) - - -@torch.inference_mode() -def trigger_flashinfer_replayssm_autotune(runner: GPUModelRunner) -> None: - """Make one maximum-batch call so FlashInfer tunes all decode buckets.""" - layers = _find_replayssm_layers(runner) - if not _distributed_layers_are_compatible(layers): - logger.warning_once( - "Skipping native FlashInfer ReplaySSM autotuning because ReplaySSM " - "layers are absent or have incompatible rank-local geometries." - ) - return - - layer = layers[0] - state_capacity = layer.kv_cache[1].shape[0] - batch = _distributed_min(_local_max_tuning_batch(runner, state_capacity)) - if batch <= 0: - logger.warning_once( - "Skipping native FlashInfer ReplaySSM autotuning because no valid " - "decode batch fits in the state cache." - ) - return - - tuning_call = _ReplaySSMTuningCall(layer, batch) - tuning_call.run() - torch.cuda.synchronize() - logger.info_once( - "Triggered native FlashInfer ReplaySSM autotuning through batch %d.", batch - ) diff --git a/vllm/model_executor/warmup/kernel_warmup.py b/vllm/model_executor/warmup/kernel_warmup.py index 0a8f163752c7..f45b73fe39d0 100644 --- a/vllm/model_executor/warmup/kernel_warmup.py +++ b/vllm/model_executor/warmup/kernel_warmup.py @@ -27,9 +27,6 @@ resolve_flashinfer_autotune_file, write_flashinfer_autotune_cache, ) -from vllm.model_executor.warmup.flashinfer_replayssm_warmup import ( - trigger_flashinfer_replayssm_autotune, -) from vllm.model_executor.warmup.flashinfer_sparse_mla_warmup import ( deepseek_v4_sparse_mla_attention_warmup, flashinfer_sparse_mla_decode_autotune_warmup, @@ -275,6 +272,68 @@ def _run_flashinfer_autotune_dummy_runs(runner: "GPUModelRunner") -> None: ) +def _flashinfer_replayssm_autotune_dummy_run(runner: "GPUModelRunner") -> None: + from vllm.model_executor.layers.mamba.ops.ssu_dispatch import ( + get_replayssm_backend, + ) + + try: + backend = get_replayssm_backend() + except RuntimeError: + return + if backend.name != "flashinfer": + return + + query_len = runner.uniform_decode_query_len + max_num_reqs = min( + runner.scheduler_config.max_num_seqs, + runner.max_num_tokens // query_len, + ) + if max_num_reqs == 0: + raise RuntimeError( + "FlashInfer ReplaySSM autotuning needs room for one decode request." + ) + + block_tables = runner.input_batch.block_table.block_tables + saved_block_ids = tuple( + block_table.block_table.np[:max_num_reqs, 0].copy() + for block_table in block_tables + ) + dummy_block_ids = range(1, max_num_reqs + 1) + for block_table in block_tables: + block_table.block_table.np[:max_num_reqs, 0] = dummy_block_ids + + try: + runner._dummy_run( + num_tokens=max_num_reqs * query_len, + uniform_decode=True, + allow_microbatching=False, + skip_eplb=True, + is_profile=True, + randomize_inputs=True, + force_attention=True, + profile_seq_lens=query_len + 1, + ) + finally: + for block_table, block_ids in zip(block_tables, saved_block_ids): + block_table.block_table.np[:max_num_reqs, 0] = block_ids + runner.input_batch.block_table.commit_block_table(max_num_reqs) + + from vllm.model_executor.layers.mamba.mamba_mixer2 import MambaMixer2 + + reset_tensors: set[int] = set() + for module in runner.get_model().modules(): + if not isinstance(module, MambaMixer2) or not module.use_replayssm: + continue + for tensor in module.kv_cache: + if not tensor.numel() or tensor.shape[0] <= max_num_reqs: + continue + data_ptr = tensor.data_ptr() + if data_ptr not in reset_tensors: + tensor[1 : max_num_reqs + 1].zero_() + reset_tensors.add(data_ptr) + + def flashinfer_autotune(runner: "GPUModelRunner") -> None: """ Autotune FlashInfer operations. @@ -339,7 +398,7 @@ def flashinfer_autotune(runner: "GPUModelRunner") -> None: # not execute ReplaySSM decode. One private maximum-batch decode # call lets FlashInfer populate every native dynamic batch bucket # before CUDA graph capture. - trigger_flashinfer_replayssm_autotune(runner) + _flashinfer_replayssm_autotune_dummy_run(runner) finally: set_autotune_process_group(None) From 9a88661f480c2cae8d3f0893cb52ea4c5cb17d38 Mon Sep 17 00:00:00 2001 From: Andrii Skliar Date: Sun, 16 Aug 2026 23:45:25 +0200 Subject: [PATCH 15/33] refactor(mamba): simplify FlashInfer ReplaySSM wiring Signed-off-by: Andrii Skliar --- tests/kernels/mamba/test_ssu_dispatch.py | 309 +++-------- tests/model_executor/test_kernel_warmup.py | 119 ++-- .../test_replayssm_metadata_builder.py | 26 +- tests/v1/e2e/test_replayssm_decode.py | 33 +- tests/v1/worker/test_utils.py | 53 +- .../layers/mamba/mamba_mixer2.py | 107 ++-- .../layers/mamba/mamba_utils.py | 15 +- .../layers/mamba/ops/ssu_dispatch.py | 514 ++++-------------- vllm/model_executor/models/nemotron_h.py | 6 - vllm/model_executor/warmup/kernel_warmup.py | 110 ++-- vllm/v1/attention/backends/mamba_attn.py | 104 ++-- vllm/v1/worker/utils.py | 42 +- 12 files changed, 442 insertions(+), 996 deletions(-) diff --git a/tests/kernels/mamba/test_ssu_dispatch.py b/tests/kernels/mamba/test_ssu_dispatch.py index c826f55fbb69..96e662014cc0 100644 --- a/tests/kernels/mamba/test_ssu_dispatch.py +++ b/tests/kernels/mamba/test_ssu_dispatch.py @@ -10,29 +10,17 @@ from vllm.config.mamba import MambaBackendEnum, MambaConfig, MambaSSUAlgorithm from vllm.model_executor.layers.mamba.ops.ssu_dispatch import ( - FLASHINFER_REPLAYSSM_AUTO_TACTIC, - FlashInferReplaySSMBackend, - FlashInferReplaySSMTactic, FlashInferSSUBackend, - TritonReplaySSMBackend, TritonSSUBackend, get_mamba_ssu_backend, - get_replayssm_backend, initialize_mamba_ssu_backend, - initialize_replayssm_backend, selective_state_update, selective_state_update_replayssm_flashinfer, - selective_state_update_replayssm_triton, update_replayssm_ring_trackers, - use_flashinfer_replayssm_tactic, ) from vllm.utils.torch_utils import set_random_seed from vllm.v1.attention.backends.registry import MambaAttentionBackendEnum -from vllm.v1.kv_cache_interface import ( - KVCacheConfig, - KVCacheGroupSpec, - MambaSpec, -) +from vllm.v1.kv_cache_interface import KVCacheConfig, KVCacheGroupSpec, MambaSpec try: import flashinfer.mamba # noqa: F401 @@ -45,69 +33,17 @@ checkpointing_ssu_module = importlib.import_module( "flashinfer.mamba.checkpointing_ssu" ) - HAS_FLASHINFER_CHECKPOINTING_SSU = hasattr( - checkpointing_ssu_module, "CheckpointingSSURunner" + HAS_FLASHINFER_CHECKPOINTING_SSU = all( + hasattr(checkpointing_ssu_module, name) + for name in ( + "CheckpointingSSURunner", + "allocate_checkpointing_ssu_scratch", + ) ) except ImportError: HAS_FLASHINFER_CHECKPOINTING_SSU = False -def _fake_flashinfer_replayssm_backend() -> FlashInferReplaySSMBackend: - backend = FlashInferReplaySSMBackend.__new__(FlashInferReplaySSMBackend) - backend._mamba_config = MambaConfig(backend=MambaBackendEnum.FLASHINFER) - backend._kernel = Mock(return_value=torch.empty(1)) - backend._tactic = FLASHINFER_REPLAYSSM_AUTO_TACTIC - return backend - - -def test_replayssm_explicit_tactic_validation(): - tactic = FlashInferReplaySSMTactic( - "two-kernel", d_split=2, precompute_heads_per_cta=8 - ) - assert tactic.name == "two_kernel_d2_h8" - with pytest.raises(ValueError, match="does not accept precompute"): - FlashInferReplaySSMTactic("monolith", precompute_heads_per_cta=8) - - -def test_replayssm_tactic_scope_restores_direct_launch_controls(monkeypatch): - import vllm.model_executor.layers.mamba.ops.ssu_dispatch as mod - - backend = _fake_flashinfer_replayssm_backend() - monkeypatch.setattr(mod, "_replayssm_backend", backend) - tactic = FlashInferReplaySSMTactic( - "two-kernel", d_split=2, precompute_heads_per_cta=8 - ) - - with ( - pytest.raises(RuntimeError, match="sentinel"), - use_flashinfer_replayssm_tactic(tactic), - ): - assert backend._tactic is tactic - raise RuntimeError("sentinel") - - assert backend._tactic is FLASHINFER_REPLAYSSM_AUTO_TACTIC - - -def test_replayssm_backend_forwards_explicit_tactic(monkeypatch): - import vllm.model_executor.layers.mamba.ops.ssu_dispatch as mod - - backend = _fake_flashinfer_replayssm_backend() - monkeypatch.setattr(mod, "_replayssm_backend", backend) - tensor = torch.empty(1) - - with use_flashinfer_replayssm_tactic( - FlashInferReplaySSMTactic("two-kernel", d_split=2, precompute_heads_per_cta=8) - ): - backend(*(tensor,) * 12) - - kwargs = backend._kernel.call_args.kwargs - assert kwargs["algorithm"] == "two-kernel" - assert kwargs["d_split"] == 2 - assert kwargs["precompute_heads_per_cta"] == 8 - assert "main_pipeline_stages" not in kwargs - assert "main_ctas_per_sm" not in kwargs - - @pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") def test_flashinfer_replayssm_ring_tracker_lifecycle(): ring_start = torch.zeros(2, dtype=torch.int32, device="cuda") @@ -158,8 +94,7 @@ def test_explicit_triton_backend(): initialize_mamba_ssu_backend( MambaConfig(backend=MambaBackendEnum.TRITON), _kv_cache_config_with_ssu() ) - backend = get_mamba_ssu_backend() - assert isinstance(backend, TritonSSUBackend) + assert isinstance(get_mamba_ssu_backend(), TritonSSUBackend) @pytest.mark.skipif(not HAS_FLASHINFER, reason="flashinfer not installed") @@ -198,18 +133,9 @@ def test_flashinfer_forwards_ssu_algorithm( ssu_algorithm=algorithm, ) ) - tensor = torch.empty(1) - backend( - tensor, - tensor, - tensor, - tensor, - tensor, - tensor, - tensor, - tensor, - ) + + backend(*(tensor,) * 8) assert kernel.call_args.kwargs["algorithm"] == expected @@ -219,9 +145,11 @@ def test_uninitialized_backend_raises(): old = mod._mamba_ssu_backend mod._mamba_ssu_backend = None - with pytest.raises(RuntimeError, match="not been initialized"): - get_mamba_ssu_backend() - mod._mamba_ssu_backend = old + try: + with pytest.raises(RuntimeError, match="not been initialized"): + get_mamba_ssu_backend() + finally: + mod._mamba_ssu_backend = old @pytest.mark.parametrize( @@ -242,8 +170,6 @@ def test_init_is_noop_for_non_ssu_mamba_type(mamba_type): MambaConfig(), _kv_cache_config_with_ssu(mamba_type) ) assert mod._mamba_ssu_backend is None - with pytest.raises(RuntimeError, match="not been initialized"): - get_mamba_ssu_backend() finally: mod._mamba_ssu_backend = old @@ -254,25 +180,24 @@ def test_flashinfer_import_error(): FlashInferSSUBackend(MambaConfig()) +@pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") def test_triton_basic_call(): set_random_seed(0) initialize_mamba_ssu_backend( MambaConfig(backend=MambaBackendEnum.TRITON), _kv_cache_config_with_ssu() ) - device = "cuda" batch_size = 2 dim = 64 dstate = 16 - - state = torch.randn(batch_size, dim, dstate, device=device) - x = torch.randn(batch_size, dim, device=device) + state = torch.randn(batch_size, dim, dstate, device="cuda") + x = torch.randn(batch_size, dim, device="cuda") out = torch.empty_like(x) - dt = torch.randn(batch_size, dim, device=device) - dt_bias = torch.rand(dim, device=device) - 4.0 - A = -torch.rand(dim, dstate, device=device) - B = torch.randn(batch_size, dstate, device=device) - C = torch.randn(batch_size, dstate, device=device) - D = torch.randn(dim, device=device) + dt = torch.randn(batch_size, dim, device="cuda") + dt_bias = torch.rand(dim, device="cuda") - 4.0 + A = -torch.rand(dim, dstate, device="cuda") + B = torch.randn(batch_size, dstate, device="cuda") + C = torch.randn(batch_size, dstate, device="cuda") + D = torch.randn(dim, device="cuda") selective_state_update( state, @@ -289,115 +214,26 @@ def test_triton_basic_call(): assert not torch.isnan(out).any() -def test_replayssm_default_backend_is_triton(): - initialize_replayssm_backend(MambaConfig(), use_replayssm=True) - backend = get_replayssm_backend() - assert isinstance(backend, TritonReplaySSMBackend) - assert backend.name == "triton" - - -def test_replayssm_explicit_triton_backend(): - initialize_replayssm_backend( - MambaConfig(backend=MambaBackendEnum.TRITON), use_replayssm=True - ) - backend = get_replayssm_backend() - assert isinstance(backend, TritonReplaySSMBackend) - - -@pytest.mark.skipif( - not HAS_FLASHINFER_CHECKPOINTING_SSU, - reason="flashinfer.mamba.checkpointing_ssu not available", -) -def test_replayssm_flashinfer_backend_init(): - initialize_replayssm_backend( - MambaConfig(backend=MambaBackendEnum.FLASHINFER), use_replayssm=True - ) - backend = get_replayssm_backend() - assert isinstance(backend, FlashInferReplaySSMBackend) - assert backend.name == "flashinfer" - - -def test_replayssm_disabled_clears_backend(): - initialize_replayssm_backend(MambaConfig(), use_replayssm=True) - assert get_replayssm_backend() is not None - initialize_replayssm_backend(MambaConfig(), use_replayssm=False) - with pytest.raises(RuntimeError, match="not been initialized"): - get_replayssm_backend() - - -def test_replayssm_cpu_backend_rejected(): - with pytest.raises(ValueError, match="does not support mamba backend"): - initialize_replayssm_backend( - MambaConfig(backend=MambaBackendEnum.CPU), use_replayssm=True - ) - - -def test_replayssm_uninitialized_backend_raises(): +def test_replayssm_flashinfer_call_forwards_explicit_controls(monkeypatch): import vllm.model_executor.layers.mamba.ops.ssu_dispatch as mod - old = mod._replayssm_backend - mod._replayssm_backend = None - try: - with pytest.raises(RuntimeError, match="not been initialized"): - get_replayssm_backend() - finally: - mod._replayssm_backend = old - - -@pytest.mark.skipif( - not HAS_FLASHINFER_CHECKPOINTING_SSU, - reason="flashinfer.mamba.checkpointing_ssu not available", -) -def test_replayssm_triton_entry_rejects_flashinfer_backend(): - initialize_replayssm_backend( - MambaConfig(backend=MambaBackendEnum.FLASHINFER), use_replayssm=True - ) - tensor = torch.empty(1) - with pytest.raises(RuntimeError, match="Triton ReplaySSM"): - selective_state_update_replayssm_triton( - tensor, - tensor, - tensor, - tensor, - tensor, - tensor, - out=tensor, - ) - - -@pytest.mark.skipif( - not HAS_FLASHINFER_CHECKPOINTING_SSU, - reason="flashinfer.mamba.checkpointing_ssu not available", -) -def test_replayssm_flashinfer_call(monkeypatch): kernel = Mock(return_value=torch.empty(1, 1, 2, 4)) - checkpointing_ssu_module = importlib.import_module( - "flashinfer.mamba.checkpointing_ssu" - ) - monkeypatch.setattr(checkpointing_ssu_module, "checkpointing_ssu", kernel) - monkeypatch.setattr( - "vllm.model_executor.layers.mamba.ops.ssu_dispatch._replayssm_backend", - None, - ) - initialize_replayssm_backend( - MambaConfig(backend=MambaBackendEnum.FLASHINFER), use_replayssm=True - ) + monkeypatch.setattr(mod, "_flashinfer_replayssm_kernel", kernel) - batch, nheads, dim, dstate, ngroups, L = 1, 2, 4, 8, 1, 16 + batch, nheads, dim, dstate, ngroups, window = 1, 2, 4, 8, 1, 16 state = torch.empty(1, nheads, dim, dstate) x = torch.empty(batch, nheads, dim) dt = torch.empty(batch, nheads, dim) A = torch.empty(nheads, dim, dstate) B = torch.empty(batch, ngroups, dstate) C = torch.empty(batch, ngroups, dstate) - D = torch.empty(nheads, dim) - dt_bias = torch.empty(nheads, dim) out = torch.empty_like(x) - x_cache = torch.empty(1, nheads, L, dim) - dt_cache = torch.empty(1, nheads, L) - B_cache = torch.empty(1, ngroups, L, dstate) + x_cache = torch.empty(1, nheads, window, dim) + dt_cache = torch.empty(1, nheads, window) + B_cache = torch.empty(1, ngroups, window, dstate) ring_start = torch.zeros(1, dtype=torch.int32) prev_num_accepted = torch.zeros(1, dtype=torch.int32) + scratch = (torch.empty(1), torch.empty(1), torch.empty(1)) selective_state_update_replayssm_flashinfer( state, @@ -412,63 +248,54 @@ def test_replayssm_flashinfer_call(monkeypatch): dt_cache, ring_start, prev_num_accepted, - D=D, - dt_bias=dt_bias, - dt_softplus=True, + scratch=scratch, + algorithm="two-kernel", + d_split=2, + precompute_heads_per_cta=8, + update_trackers=False, ) - assert kernel.call_count == 1 - kwargs = kernel.call_args.kwargs - assert kwargs["dt_softplus"] is True - # ring_start / prev_num_accepted are positional after the caches. + args = kernel.call_args.args + kwargs = kernel.call_args.kwargs assert args[4] is ring_start assert args[5] is prev_num_accepted - assert kwargs["algorithm"] == "auto" - assert kwargs["d_split"] is None - assert kwargs["precompute_heads_per_cta"] == 0 - assert "main_pipeline_stages" not in kwargs - assert "main_ctas_per_sm" not in kwargs + assert kwargs["algorithm"] == "two-kernel" + assert kwargs["d_split"] == 2 + assert kwargs["precompute_heads_per_cta"] == 8 + assert kwargs["cb_scaled"] is scratch[0] + assert kwargs["cumAdt_vec"] is scratch[1] + assert kwargs["cb_old"] is scratch[2] -def test_replayssm_requires_native_flashinfer_autotuning(monkeypatch): +def test_replayssm_requires_native_flashinfer_support(monkeypatch): import vllm.model_executor.layers.mamba.ops.ssu_dispatch as mod - old_module = SimpleNamespace(checkpointing_ssu=Mock()) + old_module = SimpleNamespace( + checkpointing_ssu=Mock(), CheckpointingSSURunner=object + ) monkeypatch.setattr(mod.importlib, "import_module", lambda _: old_module) - with pytest.raises(ImportError, match="native checkpointing_ssu autotuning"): - FlashInferReplaySSMBackend(MambaConfig(backend=MambaBackendEnum.FLASHINFER)) + with pytest.raises(ImportError, match="scratch allocation support"): + mod._initialize_flashinfer_replayssm(True) -@pytest.mark.skipif( - HAS_FLASHINFER_CHECKPOINTING_SSU, - reason="flashinfer checkpointing_ssu is installed", -) -def test_replayssm_flashinfer_import_error(): - with pytest.raises( - ImportError, - match="FlashInfer is required|native checkpointing_ssu autotuning", - ): - FlashInferReplaySSMBackend(MambaConfig(backend=MambaBackendEnum.FLASHINFER)) +def test_replayssm_flashinfer_import_error(monkeypatch): + import vllm.model_executor.layers.mamba.ops.ssu_dispatch as mod + def raise_import_error(_): + raise ImportError -def test_replayssm_dispatch_fn_uses_initialized_backend(): - import vllm.model_executor.layers.mamba.ops.ssu_dispatch as mod + monkeypatch.setattr(mod.importlib, "import_module", raise_import_error) + with pytest.raises(ImportError, match="FlashInfer is required"): + mod._initialize_flashinfer_replayssm(True) - called = Mock(return_value=torch.empty(1)) - old = mod._replayssm_backend - mod._replayssm_backend = TritonReplaySSMBackend(MambaConfig()) - mod._replayssm_backend._kernel = called - try: - tensor = torch.empty(1) - selective_state_update_replayssm_triton( - tensor, - tensor, - tensor, - tensor, - tensor, - tensor, - out=tensor, - ) - assert called.call_count == 1 - finally: - mod._replayssm_backend = old + +@pytest.mark.skipif( + not HAS_FLASHINFER_CHECKPOINTING_SSU, + reason="compatible flashinfer checkpointing_ssu not available", +) +def test_replayssm_flashinfer_backend_init(): + initialize_mamba_ssu_backend( + MambaConfig(backend=MambaBackendEnum.FLASHINFER), + _kv_cache_config_with_ssu(), + use_replayssm=True, + ) diff --git a/tests/model_executor/test_kernel_warmup.py b/tests/model_executor/test_kernel_warmup.py index c25d9cfcb285..08c1ce31bdd3 100644 --- a/tests/model_executor/test_kernel_warmup.py +++ b/tests/model_executor/test_kernel_warmup.py @@ -6,67 +6,94 @@ import numpy as np import pytest +import torch -from vllm.model_executor.layers.mamba.ops import ssu_dispatch +from vllm.config.mamba import MambaBackendEnum +from vllm.model_executor.layers.mamba.mamba_mixer2 import MambaMixer2 from vllm.model_executor.warmup import kernel_warmup as warmup @pytest.mark.parametrize( - "backend_name, expected_calls", [("flashinfer", 1), ("triton", 0)] + ("backend", "use_replayssm", "expected"), + [ + (MambaBackendEnum.FLASHINFER, True, True), + (MambaBackendEnum.TRITON, True, False), + (MambaBackendEnum.FLASHINFER, False, False), + ], ) -def test_replayssm_autotune_uses_uniform_decode( - monkeypatch, backend_name, expected_calls -): - block_ids = np.zeros((32, 1), dtype=np.int32) - block_table = SimpleNamespace(block_table=SimpleNamespace(np=block_ids)) - multi_group_block_table = SimpleNamespace( - block_tables=[block_table], commit_block_table=Mock() - ) - - def dummy_run(**kwargs): - assert block_ids[:16, 0].tolist() == list(range(1, 17)) - - dummy_run = Mock(side_effect=dummy_run) +def test_replayssm_autotune_decode_kwargs(backend, use_replayssm, expected): runner = SimpleNamespace( + vllm_config=SimpleNamespace( + cache_config=SimpleNamespace(use_replayssm=use_replayssm), + mamba_config=SimpleNamespace(backend=backend), + ), uniform_decode_query_len=6, max_num_tokens=100, scheduler_config=SimpleNamespace(max_num_seqs=32), - input_batch=SimpleNamespace(block_table=multi_group_block_table), - get_model=lambda: SimpleNamespace(modules=lambda: ()), - _dummy_run=dummy_run, ) - monkeypatch.setattr( - ssu_dispatch, - "get_replayssm_backend", - lambda: SimpleNamespace(name=backend_name), - ) - - warmup._flashinfer_replayssm_autotune_dummy_run(runner) + prefill_kwargs = { + "num_tokens": 128, + "skip_eplb": True, + "is_profile": True, + "randomize_inputs": True, + } - assert dummy_run.call_count == expected_calls - if expected_calls: - dummy_run.assert_called_once_with( - num_tokens=96, - uniform_decode=True, - allow_microbatching=False, - skip_eplb=True, - is_profile=True, - randomize_inputs=True, - force_attention=True, - profile_seq_lens=7, - ) - assert not block_ids.any() - multi_group_block_table.commit_block_table.assert_called_once_with(16) + result = warmup._flashinfer_replayssm_autotune_kwargs(runner, prefill_kwargs) + if not expected: + assert result is None + return + assert result == ( + 16, + { + **prefill_kwargs, + "num_tokens": 96, + "uniform_decode": True, + "allow_microbatching": False, + "force_attention": True, + "profile_seq_lens": 7, + }, + ) -def test_replayssm_autotune_skips_uninitialized_backend(monkeypatch): - runner = SimpleNamespace(_dummy_run=Mock()) - def get_backend(): - raise RuntimeError +def test_replayssm_autotune_slots_restore_state_and_trackers(): + mixer = MambaMixer2.__new__(MambaMixer2) + torch.nn.Module.__init__(mixer) + mixer.use_replayssm = True + mixer.kv_cache = ( + torch.full((4, 2), 3.0), + torch.full((4, 2), 3.0), + ) + mixer._replayssm_ring_start = torch.full((4,), 3, dtype=torch.int32) + mixer._replayssm_prev_num_accepted = torch.full((4,), 3, dtype=torch.int32) - monkeypatch.setattr(ssu_dispatch, "get_replayssm_backend", get_backend) + 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( + input_batch=SimpleNamespace(block_table=multi_group_block_table), + get_model=lambda: SimpleNamespace(modules=lambda: (mixer,)), + ) - warmup._flashinfer_replayssm_autotune_dummy_run(runner) + with warmup._temporary_replayssm_autotune_slots(runner, 2): + assert block_ids[:2, 0].tolist() == [1, 2] + for tensor in ( + *mixer.kv_cache, + mixer._replayssm_ring_start, + mixer._replayssm_prev_num_accepted, + ): + tensor[1:3].fill_(9) - runner._dummy_run.assert_not_called() + assert np.array_equal(block_ids, original_block_ids) + multi_group_block_table.commit_block_table.assert_called_once_with(2) + for tensor in ( + *mixer.kv_cache, + mixer._replayssm_ring_start, + mixer._replayssm_prev_num_accepted, + ): + assert torch.count_nonzero(tensor[1:3]) == 0 + assert torch.all(tensor[0] == 3) + assert torch.all(tensor[3] == 3) diff --git a/tests/v1/attention/test_replayssm_metadata_builder.py b/tests/v1/attention/test_replayssm_metadata_builder.py index 2ba86c08dad8..89cb0f566467 100644 --- a/tests/v1/attention/test_replayssm_metadata_builder.py +++ b/tests/v1/attention/test_replayssm_metadata_builder.py @@ -22,6 +22,7 @@ BLOCK_SIZE = 16 DEVICE = torch.device("cpu") + @dataclass class ReplaySSMBuildCase: """A decode batch and its expected per-row write_pos / is_flush. @@ -191,8 +192,6 @@ def _make_mamba_spec( buffer_len: int, mamba_backend: MambaBackendEnum, ) -> MambaSpec: - # The builder only reads the x/B ring shapes; include FlashInfer's trackers - # so the mock page matches the production cache layout. ring_buffer_len = buffer_len + ( 1 if mamba_backend == MambaBackendEnum.FLASHINFER else 0 ) @@ -203,8 +202,6 @@ def _make_mamba_spec( (1, ring_buffer_len), (1, ring_buffer_len, 1), ) - if mamba_backend == MambaBackendEnum.FLASHINFER: - shapes = (*shapes, (), ()) return MambaSpec( block_size=BLOCK_SIZE, shapes=shapes, @@ -274,19 +271,20 @@ def test_resumed_request_differs_from_fresh(): def test_flashinfer_replayssm_scratch_metadata_fresh_decode(): - """FlashInfer receives per-row scratch; ring state is layer-local.""" - builder = _create_replayssm_builder( - 16, mamba_backend=MambaBackendEnum.FLASHINFER - ) + 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.cb_scaled is not None - assert meta.cb_scaled.shape == (1, 1, 32, 8) - assert meta.cumAdt_vec is not None - assert meta.cumAdt_vec.shape == (1, 1, 16) - assert meta.cb_old is not None - assert meta.cb_old.shape == (1, 1, 32, 8) + 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), + ] diff --git a/tests/v1/e2e/test_replayssm_decode.py b/tests/v1/e2e/test_replayssm_decode.py index c328341c93fa..ea7ea91b5549 100644 --- a/tests/v1/e2e/test_replayssm_decode.py +++ b/tests/v1/e2e/test_replayssm_decode.py @@ -10,7 +10,9 @@ from ...utils import large_gpu_mark, multi_gpu_test try: - from flashinfer.mamba import checkpointing_ssu # noqa: F401 + from flashinfer.mamba.checkpointing_ssu import ( + allocate_checkpointing_ssu_scratch, # noqa: F401 + ) HAS_FLASHINFER_CHECKPOINTING_SSU = True except ImportError: @@ -89,35 +91,6 @@ def test_replayssm_flashinfer_decode_matches_baseline(vllm_runner, model_name): ) -@pytest.mark.skipif( - not HAS_FLASHINFER_CHECKPOINTING_SSU, - reason="flashinfer.mamba.checkpointing_ssu not available", -) -@pytest.mark.parametrize("model_name", MODELS) -def test_replayssm_flashinfer_matches_triton_replayssm(vllm_runner, model_name): - common = dict( - max_model_len=1024, - trust_remote_code=True, - enable_prefix_caching=False, - mamba_cache_mode="none", - use_replayssm=True, - replayssm_buffer_len=16, - ) - with vllm_runner(model_name, mamba_backend="triton", **common) as llm: - triton = llm.generate_greedy_logprobs(PROMPTS, max_tokens=32, num_logprobs=5) - with vllm_runner(model_name, mamba_backend="flashinfer", **common) as llm: - flashinfer = llm.generate_greedy_logprobs( - PROMPTS, max_tokens=32, num_logprobs=5 - ) - - check_logprobs_close( - outputs_0_lst=triton, - outputs_1_lst=flashinfer, - name_0="replayssm_triton", - name_1="replayssm_flashinfer", - ) - - # Prefix spans several mamba blocks; prefix caching only reuses full blocks. _PC_SENTENCE = ( "In a detailed survey of state space models, the authors compared many " diff --git a/tests/v1/worker/test_utils.py b/tests/v1/worker/test_utils.py index 02417c70c2c5..a49c76fe24b0 100644 --- a/tests/v1/worker/test_utils.py +++ b/tests/v1/worker/test_utils.py @@ -11,15 +11,13 @@ class _TestReplaySSMMixer(MambaMixer2): - _state_shapes = ((2,), (3,), (4,), (5,), (6,), (), ()) + _state_shapes = ((2,), (3,), (4,), (5,), (6,)) _state_dtypes = ( torch.float32, torch.float32, torch.float32, torch.float32, torch.float32, - torch.int32, - torch.int32, ) def __init__(self): @@ -40,50 +38,7 @@ def get_state_dtype(self) -> tuple[torch.dtype, ...]: def _packed_replayssm_cache(num_blocks: int, fill_value: int = 0) -> torch.Tensor: - return torch.full((num_blocks, 1, 1, 88), fill_value, dtype=torch.int8) - - -def test_bind_kv_cache_uses_contiguous_replayssm_tracker_sidecars(): - mixer = _TestReplaySSMMixer() - mixer.bind_kv_cache(_packed_replayssm_cache(3, fill_value=1)) - - packed_ring_start, packed_prev_num_accepted = mixer.kv_cache[5:] - assert not packed_ring_start.is_contiguous() - assert not packed_prev_num_accepted.is_contiguous() - - for tracker in ( - mixer._replayssm_ring_start, - mixer._replayssm_prev_num_accepted, - ): - assert tracker.shape == (3,) - assert tracker.dtype == torch.int32 - assert tracker.is_contiguous() - assert torch.count_nonzero(tracker) == 0 - - assert torch.count_nonzero(packed_ring_start) == 3 - assert torch.count_nonzero(packed_prev_num_accepted) == 3 - assert not dict(mixer.named_buffers()) - - -def test_bind_kv_cache_recreates_replayssm_tracker_sidecars(): - mixer = _TestReplaySSMMixer() - mixer.bind_kv_cache(_packed_replayssm_cache(2)) - old_ring_start = mixer._replayssm_ring_start - old_prev_num_accepted = mixer._replayssm_prev_num_accepted - old_ring_start.fill_(7) - old_prev_num_accepted.fill_(9) - - mixer.bind_kv_cache(_packed_replayssm_cache(4)) - - assert mixer._replayssm_ring_start.shape == (4,) - assert mixer._replayssm_prev_num_accepted.shape == (4,) - assert torch.count_nonzero(mixer._replayssm_ring_start) == 0 - assert torch.count_nonzero(mixer._replayssm_prev_num_accepted) == 0 - assert mixer._replayssm_ring_start.data_ptr() != old_ring_start.data_ptr() - assert ( - mixer._replayssm_prev_num_accepted.data_ptr() - != old_prev_num_accepted.data_ptr() - ) + return torch.full((num_blocks, 1, 1, 80), fill_value, dtype=torch.int8) def test_bind_kv_cache_shares_replayssm_trackers_by_cache_group(): @@ -120,6 +75,10 @@ def test_bind_kv_cache_shares_replayssm_trackers_by_cache_group(): ) assert mixers[0]._replayssm_ring_start.shape == (4,) assert mixers[0]._replayssm_prev_num_accepted.shape == (4,) + assert mixers[0]._replayssm_ring_start.dtype == torch.int32 + assert mixers[0]._replayssm_ring_start.is_contiguous() + assert torch.count_nonzero(mixers[0]._replayssm_ring_start) == 0 + assert torch.count_nonzero(mixers[0]._replayssm_prev_num_accepted) == 0 assert [m._updates_replayssm_trackers for m in mixers] == [False, True, True] diff --git a/vllm/model_executor/layers/mamba/mamba_mixer2.py b/vllm/model_executor/layers/mamba/mamba_mixer2.py index 8201b3f69f34..dfdbcbab4f8f 100644 --- a/vllm/model_executor/layers/mamba/mamba_mixer2.py +++ b/vllm/model_executor/layers/mamba/mamba_mixer2.py @@ -33,15 +33,16 @@ causal_conv1d_update, ) from vllm.model_executor.layers.mamba.ops.layernorm_gated import rms_norm_gated +from vllm.model_executor.layers.mamba.ops.selective_state_update_replayssm_output_only import ( # noqa: E501 + selective_state_update_replayssm_output_only, +) from vllm.model_executor.layers.mamba.ops.ssd_combined import ( mamba_chunk_scan_combined_varlen, ) from vllm.model_executor.layers.mamba.ops.ssu_dispatch import ( - get_replayssm_backend, reset_replayssm_ring_trackers, selective_state_update, selective_state_update_replayssm_flashinfer, - selective_state_update_replayssm_triton, ) from vllm.model_executor.layers.quantization import QuantizationConfig from vllm.model_executor.model_loader.weight_utils import ( @@ -519,15 +520,8 @@ def __init__( raise ValueError( "--use-replayssm requires tensor-parallel heads to divide evenly" ) - # The tuple is (conv_state, ssm_state); with the cached (ReplaySSM) decode - # kernel enabled it also has x/dt/B rings and, for FlashInfer, two - # per-slot ring trackers. - if self.use_replayssm: - _n_state = ( - 7 if self.mamba_config.backend == MambaBackendEnum.FLASHINFER else 5 - ) - else: - _n_state = 2 + # ReplaySSM appends x/dt/B rings to (conv_state, ssm_state). + _n_state = 5 if self.use_replayssm else 2 self.kv_cache = tuple(torch.tensor([]) for _ in range(_n_state)) self._replayssm_ring_start = torch.empty(0, dtype=torch.int32) self._replayssm_prev_num_accepted = torch.empty(0, dtype=torch.int32) @@ -559,30 +553,6 @@ def __init__( # Check if running on Blackwell (SM100+) for kernel tuning self.is_blackwell = current_platform.is_device_capability_family(100) - def bind_kv_cache(self, kv_cache: torch.Tensor) -> None: - super().bind_kv_cache(kv_cache) - if ( - self.use_replayssm - and self.mamba_config.backend == MambaBackendEnum.FLASHINFER - ): - assert self.cache_config is not None - assert self.cache_config.mamba_cache_mode == "none" - # FI ReplaySSM is restricted to cache mode "none", so these - # sidecars are authoritative; packed tracker fields remain reserved. - ring_start, prev_num_accepted = self.kv_cache[5:] - assert ring_start.dtype == torch.int32 - assert prev_num_accepted.dtype == torch.int32 - self._replayssm_ring_start = torch.zeros( - ring_start.shape, - dtype=torch.int32, - device=ring_start.device, - ) - self._replayssm_prev_num_accepted = torch.zeros( - prev_num_accepted.shape, - dtype=torch.int32, - device=prev_num_accepted.device, - ) - def forward( self, hidden_states: torch.Tensor, @@ -1111,13 +1081,10 @@ def conv_ssm_forward( ) if self.use_replayssm: assert self.replayssm_buffer_len is not None - replayssm_backend = get_replayssm_backend() - if replayssm_backend.name == "flashinfer": + if self.mamba_config.backend == MambaBackendEnum.FLASHINFER: assert ring_start is not None assert prev_num_accepted is not None - assert attn_metadata.cb_scaled is not None - assert attn_metadata.cumAdt_vec is not None - assert attn_metadata.cb_old is not None + assert attn_metadata.replayssm_scratch is not None selective_state_update_replayssm_flashinfer( ssm_state, hidden_states_d, @@ -1135,13 +1102,17 @@ def conv_ssm_forward( dt_bias=dt_bias, dt_softplus=True, state_batch_indices=state_indices_tensor_d_input, - cb_scaled=attn_metadata.cb_scaled, - cumAdt_vec=attn_metadata.cumAdt_vec, - cb_old=attn_metadata.cb_old, + scratch=attn_metadata.replayssm_scratch, update_trackers=self._updates_replayssm_trackers, + enable_stochastic_rounding=( + self.mamba_config.enable_stochastic_rounding + ), + stochastic_rounding_philox_rounds=( + self.mamba_config.stochastic_rounding_philox_rounds + ), ) else: - selective_state_update_replayssm_triton( + selective_state_update_replayssm_output_only( ssm_state, hidden_states_d, dt_d, @@ -1160,6 +1131,12 @@ def conv_ssm_forward( max_cache_len=self.replayssm_buffer_len, state_batch_indices=state_indices_tensor_d_input, out=preallocated_ssm_out_d, + enable_stochastic_rounding=( + self.mamba_config.enable_stochastic_rounding + ), + cache_philox_rounds=( + self.mamba_config.stochastic_rounding_philox_rounds + ), ) else: selective_state_update( @@ -1190,11 +1167,7 @@ def get_state_dtype(self) -> tuple[torch.dtype, ...]: ) if self.use_replayssm: return MambaStateDtypeCalculator.append_replayssm_ring( - base_dtype, - self.model_config.dtype, - include_trackers=( - self.mamba_config.backend == MambaBackendEnum.FLASHINFER - ), + base_dtype, self.model_config.dtype ) return base_dtype @@ -1220,9 +1193,6 @@ def get_state_shape(self) -> tuple[tuple[int, ...], ...]: self.n_groups, tp_world_size, ring_buffer_len, - include_trackers=( - self.mamba_config.backend == MambaBackendEnum.FLASHINFER - ), ) return base_shape @@ -1246,34 +1216,21 @@ def share_replayssm_ring_trackers( if not mixers: continue - first_ring_start = mixers[0]._replayssm_ring_start - first_prev_num_accepted = mixers[0]._replayssm_prev_num_accepted - expected = ( - first_ring_start.shape, - first_ring_start.device, - first_ring_start.dtype, - first_prev_num_accepted.shape, - first_prev_num_accepted.device, - first_prev_num_accepted.dtype, - mixers[0].replayssm_buffer_len, - ) + first_state = mixers[0].kv_cache[1] + expected = (first_state.shape[0], first_state.device) for mixer in mixers: - actual = ( - mixer._replayssm_ring_start.shape, - mixer._replayssm_ring_start.device, - mixer._replayssm_ring_start.dtype, - mixer._replayssm_prev_num_accepted.shape, - mixer._replayssm_prev_num_accepted.device, - mixer._replayssm_prev_num_accepted.dtype, - mixer.replayssm_buffer_len, - ) + state = mixer.kv_cache[1] + actual = (state.shape[0], state.device) if actual != expected: raise ValueError( - "ReplaySSM tracker sidecars must have matching layouts" + "ReplaySSM layers in one cache group must share cache capacity" ) - mixer._replayssm_ring_start = first_ring_start - mixer._replayssm_prev_num_accepted = first_prev_num_accepted + ring_start = torch.zeros(expected[0], dtype=torch.int32, device=expected[1]) + prev_num_accepted = torch.zeros_like(ring_start) + for mixer in mixers: + mixer._replayssm_ring_start = ring_start + mixer._replayssm_prev_num_accepted = prev_num_accepted mixer._updates_replayssm_trackers = False mixers[-1]._updates_replayssm_trackers = True diff --git a/vllm/model_executor/layers/mamba/mamba_utils.py b/vllm/model_executor/layers/mamba/mamba_utils.py index 6ab18a81bc8f..9a94576007bc 100644 --- a/vllm/model_executor/layers/mamba/mamba_utils.py +++ b/vllm/model_executor/layers/mamba/mamba_utils.py @@ -85,17 +85,12 @@ def append_replayssm_ring( cls, base_dtypes: tuple[torch.dtype, ...], model_dtype: ModelDType | torch.dtype, - include_trackers: bool = False, ) -> tuple[torch.dtype, ...]: """Append the ReplaySSM ring dtypes to a base ``(conv, ssm)`` tuple: ``(x_cache, dt_cache, B_cache)`` = ``(activation, fp32, activation)``. - FlashInfer also appends int32 ``(ring_start, prev_num_accepted)``. """ activation_dtype = get_kv_cache_torch_dtype("auto", model_dtype) - ring_dtypes = (activation_dtype, torch.float32, activation_dtype) - if include_trackers: - return (*base_dtypes, *ring_dtypes, torch.int32, torch.int32) - return (*base_dtypes, *ring_dtypes) + return (*base_dtypes, activation_dtype, torch.float32, activation_dtype) @classmethod def _mamba_state_dtype( @@ -219,24 +214,20 @@ def append_replayssm_ring( n_groups: int, tp_world_size: int, ring_buffer_len: int, - include_trackers: bool = False, ) -> tuple[tuple[int, ...], ...]: - """Append the physical ReplaySSM ring and optional tracker shapes. + """Append the physical ReplaySSM ring shapes. ``base_shapes[1]`` is ``(nheads // tp, head_dim, state_size)``; B_cache uses the un-extended ``n_groups``. """ local_nheads, head_dim, state_size = base_shapes[1] local_ngroups = divide(n_groups, tp_world_size) - shapes = ( + return ( *base_shapes, (local_nheads, ring_buffer_len, head_dim), (local_nheads, ring_buffer_len), (local_ngroups, ring_buffer_len, state_size), ) - if include_trackers: - return (*shapes, (), ()) - return shapes @classmethod def short_conv_state_shape( diff --git a/vllm/model_executor/layers/mamba/ops/ssu_dispatch.py b/vllm/model_executor/layers/mamba/ops/ssu_dispatch.py index 373beb2c7b5c..d3706609f57d 100644 --- a/vllm/model_executor/layers/mamba/ops/ssu_dispatch.py +++ b/vllm/model_executor/layers/mamba/ops/ssu_dispatch.py @@ -3,24 +3,15 @@ """ Dispatch module for Mamba selective state update (SSU) backends. -Provides unified ``selective_state_update`` (baseline decode) and ReplaySSM -decode entry points that dispatch to Triton / FlashInfer / CPU based on -``MambaBackendEnum``. On CPU-only platforms (PowerPC, x86 without CUDA) the -baseline SSU backend defaults to ``cpu``. - -ReplaySSM backends: - - Triton: ``write_pos`` / ``is_flush`` / ``bc_pre`` - (``selective_state_update_replayssm_triton``) - - FlashInfer: ``ring_start`` / ``prev_num_accepted_tokens`` (+ optional - two-kernel scratch), matching ``flashinfer.mamba.checkpointing_ssu`` - (``selective_state_update_replayssm_flashinfer``) +Provides a unified ``selective_state_update`` function that dispatches to +Triton, FlashInfer, or CPU based on ``MambaBackendEnum``. It also contains the +FlashInfer ReplaySSM adapter and shared ring-tracker kernels. On CPU-only +platforms the baseline SSU backend defaults to CPU. """ import importlib from abc import ABC, abstractmethod -from collections.abc import Iterator -from contextlib import contextmanager -from dataclasses import dataclass +from collections.abc import Callable import torch @@ -34,37 +25,6 @@ logger = init_logger(__name__) -@dataclass(frozen=True) -class FlashInferReplaySSMTactic: - algorithm: str - d_split: int | None = None - precompute_heads_per_cta: int = 0 - - def __post_init__(self) -> None: - if self.algorithm not in {"auto", "monolith", "two-kernel"}: - raise ValueError(f"Unsupported ReplaySSM algorithm: {self.algorithm}") - if self.d_split is not None and self.d_split <= 0: - raise ValueError("d_split must be positive when specified") - if self.precompute_heads_per_cta < 0: - raise ValueError("precompute_heads_per_cta must be non-negative") - if self.algorithm == "monolith" and self.precompute_heads_per_cta != 0: - raise ValueError( - f"{self.algorithm} does not accept precompute_heads_per_cta" - ) - - @property - def name(self) -> str: - name = self.algorithm.replace("-", "_") - if self.d_split is not None: - name += f"_d{self.d_split}" - if self.precompute_heads_per_cta: - name += f"_h{self.precompute_heads_per_cta}" - return name - - -FLASHINFER_REPLAYSSM_AUTO_TACTIC = FlashInferReplaySSMTactic("auto") - - @triton.jit def _update_replayssm_ring_trackers_kernel( ring_start, @@ -401,312 +361,35 @@ def __call__( _mamba_ssu_backend: MambaSSUBackend | None = None -class ReplaySSMBackend(ABC): - """Marker base for ReplaySSM decode backends.""" - - def __init__(self, mamba_config: MambaConfig): - self._mamba_config = mamba_config +_flashinfer_replayssm_kernel: Callable[..., torch.Tensor] | None = None - @property - @abstractmethod - def name(self) -> str: ... - - -class TritonReplaySSMBackend(ReplaySSMBackend): - """vLLM Triton ReplaySSM (``write_pos`` / ``is_flush`` / ``bc_pre``).""" - - def __init__(self, mamba_config: MambaConfig): - super().__init__(mamba_config) - from vllm.model_executor.layers.mamba.ops.selective_state_update_replayssm_output_only import ( # noqa: E501 - selective_state_update_replayssm_output_only as _triton_replayssm, - ) - - self._kernel = _triton_replayssm - - @property - def name(self) -> str: - return "triton" - def __call__( - self, - state: torch.Tensor, - x: torch.Tensor, - dt: torch.Tensor, - A: torch.Tensor, - B: torch.Tensor, - C: torch.Tensor, - D: torch.Tensor | None = None, - dt_bias: torch.Tensor | None = None, - z: torch.Tensor | None = None, - dt_softplus: bool = False, - x_cache: torch.Tensor | None = None, - dt_cache: torch.Tensor | None = None, - B_cache: torch.Tensor | None = None, - bc_pre: torch.Tensor | None = None, - write_pos: torch.Tensor | None = None, - is_flush: torch.Tensor | None = None, - max_cache_len: int = 16, - state_batch_indices: torch.Tensor | None = None, - null_block_id: int = NULL_BLOCK_ID, - out: torch.Tensor | None = None, - ) -> torch.Tensor: - return self._kernel( - state, - x, - dt, - A, - B, - C, - D=D, - dt_bias=dt_bias, - z=z, - dt_softplus=dt_softplus, - x_cache=x_cache, - dt_cache=dt_cache, - B_cache=B_cache, - bc_pre=bc_pre, - write_pos=write_pos, - is_flush=is_flush, - max_cache_len=max_cache_len, - state_batch_indices=state_batch_indices, - null_block_id=null_block_id, - out=out, - enable_stochastic_rounding=self._mamba_config.enable_stochastic_rounding, - cache_philox_rounds=self._mamba_config.stochastic_rounding_philox_rounds, - ) - - -class FlashInferReplaySSMBackend(ReplaySSMBackend): - """FlashInfer ``checkpointing_ssu`` ReplaySSM backend.""" - - def __init__(self, mamba_config: MambaConfig): - super().__init__(mamba_config) - try: - checkpointing_ssu_module = importlib.import_module( - "flashinfer.mamba.checkpointing_ssu" - ) - except (ImportError, ModuleNotFoundError) as e: - raise ImportError( - "FlashInfer is required for the flashinfer ReplaySSM backend. " - "Please install flashinfer with mamba.checkpointing_ssu support: " - "pip install flashinfer-python" - ) from e - if not hasattr(checkpointing_ssu_module, "CheckpointingSSURunner"): - raise ImportError( - "FlashInfer ReplaySSM requires native checkpointing_ssu " - "autotuning support. Install a compatible FlashInfer revision " - "exposing CheckpointingSSURunner." - ) - self._kernel = checkpointing_ssu_module.checkpointing_ssu - self._tactic = FLASHINFER_REPLAYSSM_AUTO_TACTIC - - @property - def name(self) -> str: - return "flashinfer" - - def __call__( - self, - state: torch.Tensor, - x: torch.Tensor, - dt: torch.Tensor, - A: torch.Tensor, - B: torch.Tensor, - C: torch.Tensor, - out: torch.Tensor, - x_cache: torch.Tensor, - B_cache: torch.Tensor, - dt_cache: torch.Tensor, - ring_start: torch.Tensor, - prev_num_accepted_tokens: torch.Tensor, - D: torch.Tensor | None = None, - dt_bias: torch.Tensor | None = None, - z: torch.Tensor | None = None, - dt_softplus: bool = False, - state_batch_indices: torch.Tensor | None = None, - null_block_id: int = NULL_BLOCK_ID, - cb_scaled: torch.Tensor | None = None, - cumAdt_vec: torch.Tensor | None = None, - cb_old: torch.Tensor | None = None, - algorithm: str | None = None, - update_trackers: bool = True, - ) -> torch.Tensor: - # AR decode currently passes (batch, nheads, dim); checkpointing_ssu - # expects a predicted-token axis T. Unsqueeze T=1 here. - if x.dim() == 3: - x = x.unsqueeze(1) - dt = dt.unsqueeze(1) - B = B.unsqueeze(1) - C = C.unsqueeze(1) - out = out.unsqueeze(1) - z = z.unsqueeze(1) if z is not None else None - - rand_seed = ( - torch.randint(0, 2**32, (1,), device=state.device, dtype=torch.int64) - if self._mamba_config.enable_stochastic_rounding - else None - ) - indices = state_batch_indices - if indices is not None and indices.dim() > 1: - indices = indices[:, 0] - - tactic = self._tactic - requested_algorithm = tactic.algorithm if algorithm is None else algorithm - use_explicit_tactic = algorithm is None - result = self._kernel( - state, - x_cache, - B_cache, - dt_cache, - ring_start, - prev_num_accepted_tokens, - x, - dt, - A, - B, - C, - out, - D=D, - z=z, - dt_bias=dt_bias, - dt_softplus=dt_softplus, - state_batch_indices=indices, - pad_slot_id=null_block_id, - rand_seed=rand_seed, - philox_rounds=self._mamba_config.stochastic_rounding_philox_rounds or 10, - cb_scaled=cb_scaled, - cumAdt_vec=cumAdt_vec, - cb_old=cb_old, - d_split=tactic.d_split if use_explicit_tactic else None, - precompute_heads_per_cta=( - tactic.precompute_heads_per_cta if use_explicit_tactic else 0 - ), - algorithm=requested_algorithm, - ) - if update_trackers and indices is not None: - update_replayssm_ring_trackers( - ring_start, - prev_num_accepted_tokens, - indices, - logical_window=x_cache.size(2) - 1, - pad_slot_id=null_block_id, - ) - return result - - -@contextmanager -def use_flashinfer_replayssm_tactic( - tactic: FlashInferReplaySSMTactic, -) -> Iterator[None]: - """Apply an explicit ReplaySSM tactic for tests or debugging.""" - backend = get_replayssm_backend() - if not isinstance(backend, FlashInferReplaySSMBackend): - yield +def _initialize_flashinfer_replayssm(enabled: bool) -> None: + global _flashinfer_replayssm_kernel + _flashinfer_replayssm_kernel = None + if not enabled: return - old_tactic = backend._tactic - backend._tactic = tactic try: - yield - finally: - backend._tactic = old_tactic - - -_REPLAYSSM_BACKEND_REGISTRY: dict[MambaBackendEnum, type[ReplaySSMBackend]] = { - MambaBackendEnum.TRITON: TritonReplaySSMBackend, - MambaBackendEnum.FLASHINFER: FlashInferReplaySSMBackend, -} - -_replayssm_backend: ReplaySSMBackend | None = None - - -def initialize_replayssm_backend( - mamba_config: MambaConfig, - *, - use_replayssm: bool, -) -> None: - """Initialize the global ReplaySSM backend when ``--use-replayssm`` is set.""" - global _replayssm_backend - if not use_replayssm: - _replayssm_backend = None - return - - backend = mamba_config.backend - if backend not in _REPLAYSSM_BACKEND_REGISTRY: - raise ValueError( - f"--use-replayssm does not support mamba backend {backend.value!r}. " - f"Valid options: {[b.value for b in _REPLAYSSM_BACKEND_REGISTRY]}" - ) - - backend_cls = _REPLAYSSM_BACKEND_REGISTRY[backend] - if isinstance(_replayssm_backend, backend_cls): - return - - _replayssm_backend = backend_cls(mamba_config) - logger.info("Using %s ReplaySSM backend.", _replayssm_backend.name) - - -def get_replayssm_backend() -> ReplaySSMBackend: - """Get the current ReplaySSM backend. Raises if not initialized.""" - if _replayssm_backend is None: - raise RuntimeError( - "ReplaySSM backend has not been initialized. " - "Call initialize_mamba_ssu_backend() with use_replayssm=True first." - ) - return _replayssm_backend - - -def selective_state_update_replayssm_triton( - state: torch.Tensor, - x: torch.Tensor, - dt: torch.Tensor, - A: torch.Tensor, - B: torch.Tensor, - C: torch.Tensor, - D: torch.Tensor | None = None, - dt_bias: torch.Tensor | None = None, - z: torch.Tensor | None = None, - dt_softplus: bool = False, - x_cache: torch.Tensor | None = None, - dt_cache: torch.Tensor | None = None, - B_cache: torch.Tensor | None = None, - bc_pre: torch.Tensor | None = None, - write_pos: torch.Tensor | None = None, - is_flush: torch.Tensor | None = None, - max_cache_len: int = 16, - state_batch_indices: torch.Tensor | None = None, - null_block_id: int = NULL_BLOCK_ID, - out: torch.Tensor | None = None, -) -> torch.Tensor: - """Triton ReplaySSM decode (``write_pos`` / ``is_flush`` / ``bc_pre``).""" - backend = get_replayssm_backend() - if not isinstance(backend, TritonReplaySSMBackend): - raise RuntimeError( - "selective_state_update_replayssm_triton is the Triton ReplaySSM " - f"entry point; current backend is {backend.name!r}. Use " - "selective_state_update_replayssm_flashinfer for FlashInfer." - ) - return backend( - state, - x, - dt, - A, - B, - C, - D=D, - dt_bias=dt_bias, - z=z, - dt_softplus=dt_softplus, - x_cache=x_cache, - dt_cache=dt_cache, - B_cache=B_cache, - bc_pre=bc_pre, - write_pos=write_pos, - is_flush=is_flush, - max_cache_len=max_cache_len, - state_batch_indices=state_batch_indices, - null_block_id=null_block_id, - out=out, + module = importlib.import_module("flashinfer.mamba.checkpointing_ssu") + except (ImportError, ModuleNotFoundError) as e: + raise ImportError( + "FlashInfer is required for the flashinfer ReplaySSM backend. " + "Install a compatible flashinfer-python package." + ) from e + + required = ( + "checkpointing_ssu", + "CheckpointingSSURunner", + "allocate_checkpointing_ssu_scratch", ) + missing = [name for name in required if not hasattr(module, name)] + if missing: + raise ImportError( + "FlashInfer ReplaySSM requires native autotuning and scratch " + f"allocation support; missing {missing}." + ) + _flashinfer_replayssm_kernel = module.checkpointing_ssu def selective_state_update_replayssm_flashinfer( @@ -728,44 +411,79 @@ def selective_state_update_replayssm_flashinfer( dt_softplus: bool = False, state_batch_indices: torch.Tensor | None = None, null_block_id: int = NULL_BLOCK_ID, - cb_scaled: torch.Tensor | None = None, - cumAdt_vec: torch.Tensor | None = None, - cb_old: torch.Tensor | None = None, - algorithm: str | None = None, + scratch: tuple[torch.Tensor, torch.Tensor, torch.Tensor] | None = None, + algorithm: str = "auto", + d_split: int | None = None, + precompute_heads_per_cta: int = 0, update_trackers: bool = True, + enable_stochastic_rounding: bool = False, + stochastic_rounding_philox_rounds: int | None = None, ) -> torch.Tensor: - """FlashInfer ReplaySSM decode (``checkpointing_ssu``).""" - backend = get_replayssm_backend() - if not isinstance(backend, FlashInferReplaySSMBackend): + """Run FlashInfer checkpointing SSU and optionally advance shared trackers.""" + if _flashinfer_replayssm_kernel is None: raise RuntimeError( - "selective_state_update_replayssm_flashinfer requires the " - f"flashinfer ReplaySSM backend; current backend is {backend.name!r}." + "FlashInfer ReplaySSM has not been initialized. " + "Call initialize_mamba_ssu_backend() with use_replayssm=True." ) - return backend( + + if x.dim() == 3: + x = x.unsqueeze(1) + dt = dt.unsqueeze(1) + B = B.unsqueeze(1) + C = C.unsqueeze(1) + out = out.unsqueeze(1) + z = z.unsqueeze(1) if z is not None else None + + indices = state_batch_indices + if indices is not None and indices.dim() > 1: + indices = indices[:, 0] + + cb_scaled = cumAdt_vec = cb_old = None + if scratch is not None: + cb_scaled, cumAdt_vec, cb_old = scratch + + rand_seed = ( + torch.randint(0, 2**32, (1,), device=state.device, dtype=torch.int64) + if enable_stochastic_rounding + else None + ) + result = _flashinfer_replayssm_kernel( state, + x_cache, + B_cache, + dt_cache, + ring_start, + prev_num_accepted_tokens, x, dt, A, B, C, out, - x_cache, - B_cache, - dt_cache, - ring_start, - prev_num_accepted_tokens, D=D, - dt_bias=dt_bias, z=z, + dt_bias=dt_bias, dt_softplus=dt_softplus, - state_batch_indices=state_batch_indices, - null_block_id=null_block_id, + state_batch_indices=indices, + pad_slot_id=null_block_id, + rand_seed=rand_seed, + philox_rounds=stochastic_rounding_philox_rounds or 10, cb_scaled=cb_scaled, cumAdt_vec=cumAdt_vec, cb_old=cb_old, + d_split=d_split, + precompute_heads_per_cta=precompute_heads_per_cta, algorithm=algorithm, - update_trackers=update_trackers, ) + if update_trackers and indices is not None: + update_replayssm_ring_trackers( + ring_start, + prev_num_accepted_tokens, + indices, + logical_window=x_cache.size(2) - x.size(1), + pad_slot_id=null_block_id, + ) + return result def initialize_mamba_ssu_backend( @@ -774,47 +492,49 @@ def initialize_mamba_ssu_backend( *, use_replayssm: bool = False, ) -> None: - """Initialize the global Mamba SSU backend (and ReplaySSM when enabled). - - No-op for baseline SSU if `kv_cache_config` contains no specs that call - selective_state_update. Always (re)considers ReplaySSM when - ``use_replayssm`` is set. - """ - if any( + """Initialize the Mamba SSU backend and optional FlashInfer ReplaySSM.""" + if not any( isinstance(g.kv_cache_spec, MambaSpec) and g.kv_cache_spec.mamba_type in (MambaAttentionBackendEnum.MAMBA1, MambaAttentionBackendEnum.MAMBA2) for g in kv_cache_config.kv_cache_groups ): - global _mamba_ssu_backend - - backend = mamba_config.backend - - # On CPU-only platforms (PowerPC, x86 without CUDA) Triton JIT is - # unstable or unavailable. Silently fall back to the CPU - # backend unless the user explicitly chose something other than "triton". - if backend == MambaBackendEnum.TRITON: - from vllm.platforms import current_platform - - if current_platform.is_cpu(): - logger.info( - "CPU platform detected: overriding Mamba SSU backend " - "from 'triton' to 'cpu'." - ) - backend = MambaBackendEnum.CPU - - if backend not in _BACKEND_REGISTRY: - raise ValueError( - f"Unknown Mamba SSU backend: {backend}. " - f"Valid options: {list(_BACKEND_REGISTRY.keys())}" + return + + global _mamba_ssu_backend + backend = mamba_config.backend + + if backend == MambaBackendEnum.TRITON: + from vllm.platforms import current_platform + + if current_platform.is_cpu(): + logger.info( + "CPU platform detected: overriding Mamba SSU backend " + "from 'triton' to 'cpu'." ) + backend = MambaBackendEnum.CPU + + if backend not in _BACKEND_REGISTRY: + raise ValueError( + f"Unknown Mamba SSU backend: {backend}. " + f"Valid options: {list(_BACKEND_REGISTRY.keys())}" + ) + if use_replayssm and backend not in ( + MambaBackendEnum.TRITON, + MambaBackendEnum.FLASHINFER, + ): + raise ValueError(f"ReplaySSM does not support mamba backend {backend.value!r}") - backend_cls = _BACKEND_REGISTRY[backend] - if not isinstance(_mamba_ssu_backend, backend_cls): - _mamba_ssu_backend = backend_cls(mamba_config) - logger.info("Using %s Mamba SSU backend.", _mamba_ssu_backend.name) + backend_cls = _BACKEND_REGISTRY[backend] + if not isinstance(_mamba_ssu_backend, backend_cls): + _mamba_ssu_backend = backend_cls(mamba_config) + logger.info("Using %s Mamba SSU backend.", _mamba_ssu_backend.name) - initialize_replayssm_backend(mamba_config, use_replayssm=use_replayssm) + _initialize_flashinfer_replayssm( + use_replayssm and backend == MambaBackendEnum.FLASHINFER + ) + if use_replayssm: + logger.info("Using %s ReplaySSM backend.", backend.value) def get_mamba_ssu_backend() -> MambaSSUBackend: diff --git a/vllm/model_executor/models/nemotron_h.py b/vllm/model_executor/models/nemotron_h.py index 5ab8afd62f7b..65f478cdbdb2 100644 --- a/vllm/model_executor/models/nemotron_h.py +++ b/vllm/model_executor/models/nemotron_h.py @@ -750,9 +750,6 @@ def get_mamba_state_dtype_from_config( return MambaStateDtypeCalculator.append_replayssm_ring( base_dtype, vllm_config.model_config.dtype, - include_trackers=( - vllm_config.mamba_config.backend == MambaBackendEnum.FLASHINFER - ), ) return base_dtype @@ -798,9 +795,6 @@ def get_mamba_state_shape_from_config( hf_config.n_groups, parallel_config.tensor_parallel_size, ring_buffer_len, - include_trackers=( - vllm_config.mamba_config.backend == MambaBackendEnum.FLASHINFER - ), ) return base_shape diff --git a/vllm/model_executor/warmup/kernel_warmup.py b/vllm/model_executor/warmup/kernel_warmup.py index f45b73fe39d0..789181b207d5 100644 --- a/vllm/model_executor/warmup/kernel_warmup.py +++ b/vllm/model_executor/warmup/kernel_warmup.py @@ -8,11 +8,14 @@ import sys import time -from typing import TYPE_CHECKING +from collections.abc import Iterator +from contextlib import contextmanager +from typing import TYPE_CHECKING, Any import torch import vllm.envs as envs +from vllm.config.mamba import MambaBackendEnum from vllm.logger import init_logger from vllm.model_executor.warmup.b12x_warmup import b12x_warmup from vllm.model_executor.warmup.cutedsl_warmup import cutedsl_warmup @@ -272,18 +275,15 @@ def _run_flashinfer_autotune_dummy_runs(runner: "GPUModelRunner") -> None: ) -def _flashinfer_replayssm_autotune_dummy_run(runner: "GPUModelRunner") -> None: - from vllm.model_executor.layers.mamba.ops.ssu_dispatch import ( - get_replayssm_backend, - ) - - try: - backend = get_replayssm_backend() - except RuntimeError: - return - if backend.name != "flashinfer": - return - +def _flashinfer_replayssm_autotune_kwargs( + runner: "GPUModelRunner", max_token_prefill_kwargs: dict[str, Any] +) -> tuple[int, dict[str, Any]] | None: + config = runner.vllm_config + if not ( + config.cache_config.use_replayssm + and config.mamba_config.backend == MambaBackendEnum.FLASHINFER + ): + return None query_len = runner.uniform_decode_query_len max_num_reqs = min( runner.scheduler_config.max_num_seqs, @@ -294,6 +294,45 @@ def _flashinfer_replayssm_autotune_dummy_run(runner: "GPUModelRunner") -> None: "FlashInfer ReplaySSM autotuning needs room for one decode request." ) + return max_num_reqs, { + **max_token_prefill_kwargs, + "num_tokens": max_num_reqs * query_len, + "uniform_decode": True, + "allow_microbatching": False, + "force_attention": True, + "profile_seq_lens": query_len + 1, + } + + +@contextmanager +def _temporary_replayssm_autotune_slots( + runner: "GPUModelRunner", max_num_reqs: int +) -> Iterator[None]: + from vllm.model_executor.layers.mamba.mamba_mixer2 import MambaMixer2 + + reset_tensors: list[torch.Tensor] = [] + seen: set[int] = set() + for module in runner.get_model().modules(): + if not isinstance(module, MambaMixer2) or not module.use_replayssm: + continue + tensors = ( + *module.kv_cache, + module._replayssm_ring_start, + module._replayssm_prev_num_accepted, + ) + for tensor in tensors: + if not tensor.numel(): + continue + if tensor.shape[0] <= max_num_reqs: + raise RuntimeError( + "FlashInfer ReplaySSM autotuning needs max_num_reqs + 1 " + "state slots." + ) + data_ptr = tensor.data_ptr() + if data_ptr not in seen: + reset_tensors.append(tensor) + seen.add(data_ptr) + block_tables = runner.input_batch.block_table.block_tables saved_block_ids = tuple( block_table.block_table.np[:max_num_reqs, 0].copy() @@ -304,34 +343,13 @@ def _flashinfer_replayssm_autotune_dummy_run(runner: "GPUModelRunner") -> None: block_table.block_table.np[:max_num_reqs, 0] = dummy_block_ids try: - runner._dummy_run( - num_tokens=max_num_reqs * query_len, - uniform_decode=True, - allow_microbatching=False, - skip_eplb=True, - is_profile=True, - randomize_inputs=True, - force_attention=True, - profile_seq_lens=query_len + 1, - ) + yield finally: for block_table, block_ids in zip(block_tables, saved_block_ids): block_table.block_table.np[:max_num_reqs, 0] = block_ids runner.input_batch.block_table.commit_block_table(max_num_reqs) - - from vllm.model_executor.layers.mamba.mamba_mixer2 import MambaMixer2 - - reset_tensors: set[int] = set() - for module in runner.get_model().modules(): - if not isinstance(module, MambaMixer2) or not module.use_replayssm: - continue - for tensor in module.kv_cache: - if not tensor.numel() or tensor.shape[0] <= max_num_reqs: - continue - data_ptr = tensor.data_ptr() - if data_ptr not in reset_tensors: - tensor[1 : max_num_reqs + 1].zero_() - reset_tensors.add(data_ptr) + for tensor in reset_tensors: + tensor[1 : max_num_reqs + 1].zero_() def flashinfer_autotune(runner: "GPUModelRunner") -> None: @@ -375,6 +393,15 @@ def flashinfer_autotune(runner: "GPUModelRunner") -> None: # which lead to some EP ranks receiving no tokens and skipping their # MoE kernel entirely, and cause hang due to all-reduce collective # during synchronized autotuning. + max_token_prefill_kwargs = dict( + num_tokens=runner.scheduler_config.max_num_batched_tokens, + skip_eplb=True, + is_profile=True, + randomize_inputs=True, + ) + replayssm_autotune = _flashinfer_replayssm_autotune_kwargs( + runner, max_token_prefill_kwargs + ) # Read cached autotune results and broadcast to all ranks. cached_results: bytes | None = None if is_leader and cache_path.exists(): @@ -394,11 +421,10 @@ def flashinfer_autotune(runner: "GPUModelRunner") -> None: fi_utils.autotune(tune_mode=True, **autotune_kwargs), ): _run_flashinfer_autotune_dummy_runs(runner) - # The generic dummy run is predominantly prefill-shaped and may - # not execute ReplaySSM decode. One private maximum-batch decode - # call lets FlashInfer populate every native dynamic batch bucket - # before CUDA graph capture. - _flashinfer_replayssm_autotune_dummy_run(runner) + if replayssm_autotune is not None: + max_num_reqs, max_batch_decode_kwargs = replayssm_autotune + with _temporary_replayssm_autotune_slots(runner, max_num_reqs): + runner._dummy_run(**max_batch_decode_kwargs) finally: set_autotune_process_group(None) diff --git a/vllm/v1/attention/backends/mamba_attn.py b/vllm/v1/attention/backends/mamba_attn.py index bbbb39305081..ce27e32ef55d 100644 --- a/vllm/v1/attention/backends/mamba_attn.py +++ b/vllm/v1/attention/backends/mamba_attn.py @@ -83,10 +83,7 @@ class BaseMambaAttentionMetadata: is_flush_d: torch.Tensor | None = None bc_pre_scratch: torch.Tensor | None = None # ReplaySSM — FlashInfer checkpointing_ssu two-kernel scratch. - # The per-layer ring trackers live in the Mamba KV cache. - cb_scaled: torch.Tensor | None = None - cumAdt_vec: torch.Tensor | None = None - cb_old: torch.Tensor | None = None + replayssm_scratch: tuple[torch.Tensor, torch.Tensor, torch.Tensor] | None = None class BaseMambaAttentionMetadataBuilder(AttentionMetadataBuilder[M], abc.ABC): @@ -179,9 +176,8 @@ def __init__( ) # ReplaySSM CUDA-graph buffers. # Triton: write_pos / is_flush / bc_pre. - # FlashInfer: ring_start / prev_num_accepted + two-kernel scratch - # (cb_scaled / cumAdt_vec / cb_old) so algorithm="auto" can pick - # monolith or two-kernel. + # FlashInfer: two-kernel scratch, so algorithm="auto" can pick the + # monolith or two-kernel implementation. if self.use_replayssm and not self.use_flashinfer_replayssm: self.decode_write_pos_d: torch.Tensor = torch.empty( (self.decode_cudagraph_max_bs,), @@ -208,47 +204,25 @@ def __init__( dtype=torch.float32, device=device, ) - self.decode_cb_scaled = None - self.decode_cumAdt_vec = None - self.decode_cb_old = None + self.decode_replayssm_scratch = None elif self.use_flashinfer_replayssm: - self.decode_bc_pre_scratch = None - scratch_bs = max( - self.decode_cudagraph_max_bs, scheduler_config.max_num_seqs + from flashinfer.mamba.checkpointing_ssu import ( + allocate_checkpointing_ssu_scratch, ) - # x_cache page: (nheads, L, head_dim); AR decode uses T=1. - # Scratch shapes follow flashinfer's bench_checkpointing_ssu.py - # (WARP_SIZE / MMA_FRAG_SIZE are not exported from Python). + + self.decode_bc_pre_scratch = None nheads = kv_cache_spec.shapes[2][0] - npredicted = 1 # AR decode - max_window = self.replayssm_buffer_len - k_old = (max_window + 7) // 8 * 8 - # cumAdt_vec: next_multiple_of_16(T) - t_pad = ((npredicted + 15) // 16) * 16 - warp_size = 32 - # cb_scaled: (..., 32, 8) = fragA for m16n8k16 - mma_frag_size = t_pad // 2 - act_dtype = vllm_config.model_config.dtype - self.decode_cb_scaled = torch.empty( - (scratch_bs, nheads, warp_size, mma_frag_size), - dtype=act_dtype, - device=device, - ) - self.decode_cumAdt_vec = torch.empty( - (scratch_bs, nheads, t_pad), - dtype=torch.float32, - device=device, - ) - self.decode_cb_old = torch.empty( - (scratch_bs, nheads, warp_size, k_old // 2), - dtype=act_dtype, + self.decode_replayssm_scratch = allocate_checkpointing_ssu_scratch( + batch_size=scheduler_config.max_num_seqs, + num_heads=nheads, + num_predicted_tokens=1, + max_window=self.replayssm_buffer_len, + dtype=vllm_config.model_config.dtype, device=device, ) else: self.decode_bc_pre_scratch = None - self.decode_cb_scaled = None - self.decode_cumAdt_vec = None - self.decode_cb_old = None + self.decode_replayssm_scratch = None self._init_reorder_batch_threshold(1, self.use_spec_decode) if self.use_spec_decode: @@ -553,9 +527,7 @@ def _compute_common_metadata( nums_dict, batch_ptr, token_chunk_offset_ptr = None, None, None write_pos_d = None is_flush_d = None - cb_scaled = None - cumAdt_vec = None - cb_old = None + replayssm_scratch = None if self.vllm_config.cache_config.mamba_cache_mode == "all": num_computed_tokens = common_attn_metadata.compute_num_computed_tokens() @@ -642,11 +614,7 @@ def _compute_common_metadata( num_reqs - num_prefills : num_reqs ] - if ( - self.use_replayssm - and not self.use_flashinfer_replayssm - and num_decodes > 0 - ): + if self.use_replayssm and not self.use_flashinfer_replayssm and num_decodes > 0: decode_base_cpu = common_attn_metadata.replayssm_decode_base_cpu num_computed_tokens_cpu = common_attn_metadata._num_computed_tokens_cpu if decode_base_cpu is None or num_computed_tokens_cpu is None: @@ -712,12 +680,13 @@ def _compute_common_metadata( ) if self.use_flashinfer_replayssm and num_decodes > 0: - assert self.decode_cb_scaled is not None - assert self.decode_cumAdt_vec is not None - assert self.decode_cb_old is not None - cb_scaled = self.decode_cb_scaled[:num_decodes] - cumAdt_vec = self.decode_cumAdt_vec[:num_decodes] - cb_old = self.decode_cb_old[:num_decodes] + assert self.decode_replayssm_scratch is not None + cb_scaled, cumAdt_vec, cb_old = self.decode_replayssm_scratch + replayssm_scratch = ( + cb_scaled[:num_decodes], + cumAdt_vec[:num_decodes], + cb_old[:num_decodes], + ) bc_pre_scratch = None if ( @@ -739,9 +708,7 @@ def _compute_common_metadata( write_pos_d=write_pos_d, is_flush_d=is_flush_d, bc_pre_scratch=bc_pre_scratch, - cb_scaled=cb_scaled, - cumAdt_vec=cumAdt_vec, - cb_old=cb_old, + replayssm_scratch=replayssm_scratch, num_accepted_tokens=num_accepted_tokens, query_start_loc_d=query_start_loc_d, block_idx_last_scheduled_token=block_idx_last_scheduled_token, @@ -779,9 +746,7 @@ def _update_metadata_for_cudagraph_capture( write_pos_d = metadata.write_pos_d is_flush_d = metadata.is_flush_d bc_pre_scratch = metadata.bc_pre_scratch - cb_scaled = metadata.cb_scaled - cumAdt_vec = metadata.cumAdt_vec - cb_old = metadata.cb_old + replayssm_scratch = metadata.replayssm_scratch if ( metadata.num_prefills == 0 and metadata.num_decodes <= self.decode_cudagraph_max_bs @@ -863,12 +828,13 @@ def _update_metadata_for_cudagraph_capture( if self.decode_bc_pre_scratch is not None: bc_pre_scratch = self.decode_bc_pre_scratch[:padded_bs] elif self.use_flashinfer_replayssm: - assert self.decode_cb_scaled is not None - assert self.decode_cumAdt_vec is not None - assert self.decode_cb_old is not None - cb_scaled = self.decode_cb_scaled[:padded_bs] - cumAdt_vec = self.decode_cumAdt_vec[:padded_bs] - cb_old = self.decode_cb_old[:padded_bs] + assert self.decode_replayssm_scratch is not None + cb_scaled, cumAdt_vec, cb_old = self.decode_replayssm_scratch + replayssm_scratch = ( + cb_scaled[:padded_bs], + cumAdt_vec[:padded_bs], + cb_old[:padded_bs], + ) return replace( metadata, @@ -878,9 +844,7 @@ def _update_metadata_for_cudagraph_capture( write_pos_d=write_pos_d, is_flush_d=is_flush_d, bc_pre_scratch=bc_pre_scratch, - cb_scaled=cb_scaled, - cumAdt_vec=cumAdt_vec, - cb_old=cb_old, + replayssm_scratch=replayssm_scratch, block_idx_last_scheduled_token=block_idx_last_scheduled_token, block_idx_last_computed_token=block_idx_last_computed_token, block_idx_last_scheduled_token_prev_step=( diff --git a/vllm/v1/worker/utils.py b/vllm/v1/worker/utils.py index ed8abb502fbe..7707fcb520c7 100644 --- a/vllm/v1/worker/utils.py +++ b/vllm/v1/worker/utils.py @@ -620,28 +620,38 @@ def bind_kv_cache( share_replayssm_ring_trackers, ) - layer_to_cache_group = { - layer_name: group_index - for group_index, group in enumerate(kv_cache_groups or ()) - for layer_name in group.layer_names - } - replayssm_mixer_groups = defaultdict(list) + replayssm_mixers: dict[str, MambaMixer2] = {} for layer_name in ordered_layer_names: layer = forward_context[layer_name] - if not ( + if ( isinstance(layer, MambaMixer2) and layer.use_replayssm and layer.mamba_config.backend == MambaBackendEnum.FLASHINFER ): - continue - group_index = layer_to_cache_group.get(layer_name) - if group_index is None: - # Callers without a KV-cache configuration cannot prove that two - # layers use the same block-index namespace, so keep them separate. - group_index = layer_name - replayssm_mixer_groups[group_index].append(layer) - - share_replayssm_ring_trackers(list(replayssm_mixer_groups.values())) + replayssm_mixers[layer_name] = layer + if kv_cache_groups: + mixer_groups = [] + grouped_names = { + layer_name for group in kv_cache_groups for layer_name in group.layer_names + } + for group in kv_cache_groups: + group_names = set(group.layer_names) + mixer_groups.append( + [ + replayssm_mixers[layer_name] + for layer_name in ordered_layer_names + if layer_name in group_names and layer_name in replayssm_mixers + ] + ) + mixer_groups.extend( + [mixer] + for layer_name, mixer in replayssm_mixers.items() + if layer_name not in grouped_names + ) + else: + # Without cache groups, block-index namespaces cannot be proven equal. + mixer_groups = [[mixer] for mixer in replayssm_mixers.values()] + share_replayssm_ring_trackers(mixer_groups) def copy_kv_cache_blocks_inplace( From e8ea5d4a588655f7deea3851f2ddfab417acfcf5 Mon Sep 17 00:00:00 2001 From: Andrii Skliar Date: Mon, 17 Aug 2026 10:14:44 +0200 Subject: [PATCH 16/33] fix(mamba): support ReplaySSM autotuning on runner v2 Signed-off-by: Andrii Skliar --- tests/model_executor/test_kernel_warmup.py | 107 +++++++++++++++++--- tests/v1/worker/test_gpu_block_table.py | 17 ++++ vllm/model_executor/warmup/kernel_warmup.py | 104 ++++++++++++++----- vllm/v1/worker/gpu/block_table.py | 16 ++- vllm/v1/worker/gpu/model_runner.py | 13 ++- 5 files changed, 211 insertions(+), 46 deletions(-) diff --git a/tests/model_executor/test_kernel_warmup.py b/tests/model_executor/test_kernel_warmup.py index 08c1ce31bdd3..8696b784f93d 100644 --- a/tests/model_executor/test_kernel_warmup.py +++ b/tests/model_executor/test_kernel_warmup.py @@ -14,22 +14,28 @@ @pytest.mark.parametrize( - ("backend", "use_replayssm", "expected"), + ("backend", "use_replayssm", "use_v2_model_runner", "expected"), [ - (MambaBackendEnum.FLASHINFER, True, True), - (MambaBackendEnum.TRITON, True, False), - (MambaBackendEnum.FLASHINFER, False, False), + (MambaBackendEnum.FLASHINFER, True, False, True), + (MambaBackendEnum.FLASHINFER, True, True, True), + (MambaBackendEnum.TRITON, True, False, False), + (MambaBackendEnum.FLASHINFER, False, False, False), ], ) -def test_replayssm_autotune_decode_kwargs(backend, use_replayssm, expected): +def test_replayssm_autotune_decode_kwargs( + backend, use_replayssm, use_v2_model_runner, expected +): runner = 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=6, + decode_query_len=6, max_num_tokens=100, scheduler_config=SimpleNamespace(max_num_seqs=32), + kv_cache_config=SimpleNamespace(num_blocks=17), ) prefill_kwargs = { "num_tokens": 128, @@ -43,18 +49,57 @@ def test_replayssm_autotune_decode_kwargs(backend, use_replayssm, expected): if not expected: assert result is None return - assert result == ( - 16, - { - **prefill_kwargs, - "num_tokens": 96, - "uniform_decode": True, - "allow_microbatching": False, - "force_attention": True, - "profile_seq_lens": 7, - }, + expected_kwargs = { + **prefill_kwargs, + "num_tokens": 96, + "uniform_decode": True, + } + if use_v2_model_runner: + expected_kwargs["dummy_first_block_id"] = 1 + else: + expected_kwargs.update( + allow_microbatching=False, + force_attention=True, + profile_seq_lens=7, + ) + assert result == (16, expected_kwargs) + + +def test_replayssm_autotune_decode_kwargs_clamps_to_state_capacity(): + runner = SimpleNamespace( + vllm_config=SimpleNamespace( + cache_config=SimpleNamespace(use_replayssm=True), + mamba_config=SimpleNamespace(backend=MambaBackendEnum.FLASHINFER), + use_v2_model_runner=False, + ), + uniform_decode_query_len=1, + max_num_tokens=128, + scheduler_config=SimpleNamespace(max_num_seqs=64), + kv_cache_config=SimpleNamespace(num_blocks=5), + ) + + result = warmup._flashinfer_replayssm_autotune_kwargs(runner, {}) + + assert result is not None + assert result[0] == 4 + assert result[1]["num_tokens"] == 4 + + +def test_replayssm_autotune_decode_kwargs_skips_without_state_slot(): + runner = SimpleNamespace( + vllm_config=SimpleNamespace( + cache_config=SimpleNamespace(use_replayssm=True), + mamba_config=SimpleNamespace(backend=MambaBackendEnum.FLASHINFER), + use_v2_model_runner=False, + ), + uniform_decode_query_len=1, + max_num_tokens=128, + scheduler_config=SimpleNamespace(max_num_seqs=64), + kv_cache_config=SimpleNamespace(num_blocks=1), ) + assert warmup._flashinfer_replayssm_autotune_kwargs(runner, {}) is None + def test_replayssm_autotune_slots_restore_state_and_trackers(): mixer = MambaMixer2.__new__(MambaMixer2) @@ -74,6 +119,7 @@ def test_replayssm_autotune_slots_restore_state_and_trackers(): 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,)), ) @@ -97,3 +143,34 @@ def test_replayssm_autotune_slots_restore_state_and_trackers(): assert torch.count_nonzero(tensor[1:3]) == 0 assert torch.all(tensor[0] == 3) assert torch.all(tensor[3] == 3) + + +def test_replayssm_autotune_slots_reset_v2_dummy_tables_and_state(): + mixer = MambaMixer2.__new__(MambaMixer2) + torch.nn.Module.__init__(mixer) + mixer.use_replayssm = True + mixer.kv_cache = (torch.full((4, 2), 3.0),) + mixer._replayssm_ring_start = torch.full((4,), 3, dtype=torch.int32) + mixer._replayssm_prev_num_accepted = torch.full((4,), 3, dtype=torch.int32) + block_tables = SimpleNamespace(get_dummy_block_tables=Mock()) + runner = SimpleNamespace( + vllm_config=SimpleNamespace(use_v2_model_runner=True), + block_tables=block_tables, + get_model=lambda: SimpleNamespace(modules=lambda: (mixer,)), + ) + + with warmup._temporary_replayssm_autotune_slots(runner, 2): + for tensor in ( + *mixer.kv_cache, + mixer._replayssm_ring_start, + mixer._replayssm_prev_num_accepted, + ): + tensor[1:3].fill_(9) + + block_tables.get_dummy_block_tables.assert_called_once_with(2) + for tensor in ( + *mixer.kv_cache, + mixer._replayssm_ring_start, + mixer._replayssm_prev_num_accepted, + ): + assert torch.count_nonzero(tensor[1:3]) == 0 diff --git a/tests/v1/worker/test_gpu_block_table.py b/tests/v1/worker/test_gpu_block_table.py index ee44ff24d581..7ebde608aa73 100644 --- a/tests/v1/worker/test_gpu_block_table.py +++ b/tests/v1/worker/test_gpu_block_table.py @@ -237,3 +237,20 @@ def test_get_dummy_block_tables_returns_zeroed_rows(): assert (dummy[0] == 0).all() # CUDA graph invariant: same persistent tensor, not a fresh allocation. assert dummy[0].data_ptr() == block_tables.input_block_tables[0].data_ptr() + + +def test_get_dummy_block_tables_can_assign_non_null_state_slots(): + block_tables = BlockTables( + block_sizes=[16], + max_num_reqs=4, + max_num_batched_tokens=64, + max_num_blocks_per_group=[8], + device=torch.device("cuda"), + kernel_block_sizes=[16], + ) + + (dummy,) = block_tables.get_dummy_block_tables(3, first_block_id=1) + + assert dummy[:, 0].tolist() == [1, 2, 3] + assert torch.count_nonzero(dummy[:, 1:]) == 0 + assert dummy.data_ptr() == block_tables.input_block_tables[0].data_ptr() diff --git a/vllm/model_executor/warmup/kernel_warmup.py b/vllm/model_executor/warmup/kernel_warmup.py index 789181b207d5..ed47d40369a6 100644 --- a/vllm/model_executor/warmup/kernel_warmup.py +++ b/vllm/model_executor/warmup/kernel_warmup.py @@ -284,24 +284,39 @@ def _flashinfer_replayssm_autotune_kwargs( and config.mamba_config.backend == MambaBackendEnum.FLASHINFER ): return None - query_len = runner.uniform_decode_query_len + use_v2_model_runner = config.use_v2_model_runner + v2_runner: Any = runner + query_len = ( + v2_runner.decode_query_len + if use_v2_model_runner + else runner.uniform_decode_query_len + ) max_num_reqs = min( runner.scheduler_config.max_num_seqs, runner.max_num_tokens // query_len, + runner.kv_cache_config.num_blocks - 1, ) - if max_num_reqs == 0: - raise RuntimeError( - "FlashInfer ReplaySSM autotuning needs room for one decode request." + if max_num_reqs <= 0: + logger.warning_once( + "Skipping FlashInfer ReplaySSM autotuning because no non-padding " + "state slot is available." ) + return None - return max_num_reqs, { + decode_kwargs = { **max_token_prefill_kwargs, "num_tokens": max_num_reqs * query_len, "uniform_decode": True, - "allow_microbatching": False, - "force_attention": True, - "profile_seq_lens": query_len + 1, } + if use_v2_model_runner: + decode_kwargs["dummy_first_block_id"] = 1 + else: + decode_kwargs.update( + allow_microbatching=False, + force_attention=True, + profile_seq_lens=query_len + 1, + ) + return max_num_reqs, decode_kwargs @contextmanager @@ -309,45 +324,82 @@ def _temporary_replayssm_autotune_slots( runner: "GPUModelRunner", max_num_reqs: int ) -> Iterator[None]: from vllm.model_executor.layers.mamba.mamba_mixer2 import MambaMixer2 + from vllm.model_executor.layers.mamba.ops.ssu_dispatch import ( + reset_replayssm_ring_trackers, + ) reset_tensors: list[torch.Tensor] = [] seen: set[int] = set() + tracker_pairs: list[tuple[torch.Tensor, torch.Tensor]] = [] + seen_tracker_pairs: set[tuple[int, int]] = set() for module in runner.get_model().modules(): if not isinstance(module, MambaMixer2) or not module.use_replayssm: continue - tensors = ( - *module.kv_cache, + tracker_pair = ( module._replayssm_ring_start, module._replayssm_prev_num_accepted, ) + tracker_ptrs = (tracker_pair[0].data_ptr(), tracker_pair[1].data_ptr()) + if tracker_ptrs not in seen_tracker_pairs: + tracker_pairs.append(tracker_pair) + seen_tracker_pairs.add(tracker_ptrs) + tensors = ( + *module.kv_cache, + *tracker_pair, + ) for tensor in tensors: if not tensor.numel(): continue - if tensor.shape[0] <= max_num_reqs: - raise RuntimeError( - "FlashInfer ReplaySSM autotuning needs max_num_reqs + 1 " - "state slots." - ) data_ptr = tensor.data_ptr() if data_ptr not in seen: reset_tensors.append(tensor) seen.add(data_ptr) - block_tables = runner.input_batch.block_table.block_tables - saved_block_ids = tuple( - block_table.block_table.np[:max_num_reqs, 0].copy() - for block_table in block_tables - ) - dummy_block_ids = range(1, max_num_reqs + 1) - for block_table in block_tables: - block_table.block_table.np[:max_num_reqs, 0] = dummy_block_ids + use_v2_model_runner = runner.vllm_config.use_v2_model_runner + v2_runner: Any = runner + block_tables = saved_block_ids = None + if not use_v2_model_runner: + block_tables = runner.input_batch.block_table.block_tables + saved_block_ids = tuple( + block_table.block_table.np[:max_num_reqs, 0].copy() + for block_table in block_tables + ) + dummy_block_ids = range(1, max_num_reqs + 1) + for block_table in block_tables: + block_table.block_table.np[:max_num_reqs, 0] = dummy_block_ids + + if tracker_pairs and tracker_pairs[0][0].is_cuda: + if use_v2_model_runner: + # Match the strided first-column view used by production prefill; + # Triton specializes this tracker reset separately from a fresh, + # contiguous arange tensor. + state_batch_indices = v2_runner.block_tables.get_dummy_block_tables( + max_num_reqs, first_block_id=1 + )[0][:, 0] + else: + state_batch_indices = torch.arange( + 1, + max_num_reqs + 1, + dtype=torch.int32, + device=tracker_pairs[0][0].device, + ) + for ring_start, prev_num_accepted in tracker_pairs: + reset_replayssm_ring_trackers( + ring_start, + prev_num_accepted, + state_batch_indices, + ) try: yield finally: - for block_table, block_ids in zip(block_tables, saved_block_ids): - block_table.block_table.np[:max_num_reqs, 0] = block_ids - runner.input_batch.block_table.commit_block_table(max_num_reqs) + if use_v2_model_runner: + v2_runner.block_tables.get_dummy_block_tables(max_num_reqs) + else: + assert block_tables is not None and saved_block_ids is not None + for block_table, block_ids in zip(block_tables, saved_block_ids): + block_table.block_table.np[:max_num_reqs, 0] = block_ids + runner.input_batch.block_table.commit_block_table(max_num_reqs) for tensor in reset_tensors: tensor[1 : max_num_reqs + 1].zero_() diff --git a/vllm/v1/worker/gpu/block_table.py b/vllm/v1/worker/gpu/block_table.py index 5e09383ee6ab..17bb172a9c17 100644 --- a/vllm/v1/worker/gpu/block_table.py +++ b/vllm/v1/worker/gpu/block_table.py @@ -167,7 +167,9 @@ def gather_block_tables( ) return tuple(bt[:num_reqs_padded] for bt in out) - def get_dummy_block_tables(self, num_reqs: int) -> tuple[torch.Tensor, ...]: + def get_dummy_block_tables( + self, num_reqs: int, first_block_id: int | None = None + ) -> tuple[torch.Tensor, ...]: # NOTE(woosuk): The output may be used for CUDA graph capture. # Therefore, this method must return the persistent tensor # with the same memory address as that used during the model's forward pass, @@ -176,9 +178,19 @@ def get_dummy_block_tables(self, num_reqs: int) -> tuple[torch.Tensor, ...]: # Zero the rows so dummy runs write mamba state to the reserved null # block rather than through the previous real step's (stale) block # ids, which may point at blocks since freed and reallocated. - return tuple( + result = tuple( block_table[:num_reqs].zero_() for block_table in self.input_block_tables ) + if first_block_id is not None: + block_ids = torch.arange( + first_block_id, + first_block_id + num_reqs, + dtype=torch.int32, + device=self.device, + ) + for block_table in result: + block_table[:, 0].copy_(block_ids) + return result def compute_slot_mappings( self, diff --git a/vllm/v1/worker/gpu/model_runner.py b/vllm/v1/worker/gpu/model_runner.py index b25cdb4d033e..796752485fab 100644 --- a/vllm/v1/worker/gpu/model_runner.py +++ b/vllm/v1/worker/gpu/model_runner.py @@ -689,6 +689,7 @@ def _dummy_run( context_len: int = 0, skip_eplb: bool = False, is_profile: bool = False, + dummy_first_block_id: int | None = None, **kwargs, ) -> tuple[torch.Tensor | None, torch.Tensor | None]: if skip_attn and not is_profile: @@ -746,6 +747,7 @@ def _dummy_run( skip_attn_for_dummy_run=skip_attn, is_profile=is_profile, context_len=context_len, + dummy_first_block_id=dummy_first_block_id, ) self.kv_connector.set_disabled(False) @@ -1371,9 +1373,11 @@ def prepare_attn( return block_tables, slot_mappings def prepare_dummy_attn( - self, input_batch: InputBatch + self, input_batch: InputBatch, first_block_id: int | None = None ) -> tuple[tuple[torch.Tensor, ...], torch.Tensor]: - block_tables = self.block_tables.get_dummy_block_tables(input_batch.num_reqs) + block_tables = self.block_tables.get_dummy_block_tables( + input_batch.num_reqs, first_block_id + ) slot_mappings = pcp.maybe_get_pcp_dummy_slot_mappings( self.pcp_manager, self.block_tables, input_batch.num_tokens ) @@ -1506,6 +1510,7 @@ def execute_model( skip_attn_for_dummy_run: bool = False, is_profile: bool = False, context_len: int = 0, + dummy_first_block_id: int | None = None, ) -> ModelRunnerOutput | IntermediateTensors | None: if not dummy_run: # Update the request states. @@ -1605,7 +1610,9 @@ def execute_model( max_query_len=batch_desc.max_query_len, ) if not skip_attn_for_dummy_run: - block_tables, slot_mappings = self.prepare_dummy_attn(input_batch) + block_tables, slot_mappings = self.prepare_dummy_attn( + input_batch, dummy_first_block_id + ) if context_len: set_dummy_context( input_batch, From 67807dd8d74b405777c6fed17d1287962f936432 Mon Sep 17 00:00:00 2001 From: Andrii Skliar Date: Mon, 17 Aug 2026 10:14:55 +0200 Subject: [PATCH 17/33] fix(mamba): enforce ReplaySSM integration contracts Signed-off-by: Andrii Skliar --- tests/kernels/mamba/test_ssu_dispatch.py | 89 ++++++++++++++++++- tests/test_config.py | 51 +++++++++++ vllm/config/mamba.py | 4 +- vllm/config/vllm.py | 14 +++ .../layers/mamba/mamba_mixer2.py | 16 ++-- .../layers/mamba/mamba_utils.py | 10 ++- .../layers/mamba/ops/ssu_dispatch.py | 75 +++++++++++----- vllm/model_executor/models/nemotron_h.py | 16 ++-- vllm/v1/attention/backends/mamba_attn.py | 6 +- 9 files changed, 235 insertions(+), 46 deletions(-) diff --git a/tests/kernels/mamba/test_ssu_dispatch.py b/tests/kernels/mamba/test_ssu_dispatch.py index 96e662014cc0..ee3edafd31c1 100644 --- a/tests/kernels/mamba/test_ssu_dispatch.py +++ b/tests/kernels/mamba/test_ssu_dispatch.py @@ -9,6 +9,7 @@ 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, @@ -44,6 +45,17 @@ HAS_FLASHINFER_CHECKPOINTING_SSU = 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 + + @pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") def test_flashinfer_replayssm_ring_tracker_lifecycle(): ring_start = torch.zeros(2, dtype=torch.int32, device="cuda") @@ -57,6 +69,7 @@ def test_flashinfer_replayssm_ring_tracker_lifecycle(): prev_num_accepted, state_batch_indices, logical_window=16, + ring_buffer_len=17, ) observed.append((int(ring_start[1]), int(prev_num_accepted[1]))) @@ -170,6 +183,8 @@ def test_init_is_noop_for_non_ssu_mamba_type(mamba_type): MambaConfig(), _kv_cache_config_with_ssu(mamba_type) ) assert mod._mamba_ssu_backend is None + with pytest.raises(RuntimeError, match="not been initialized"): + get_mamba_ssu_backend() finally: mod._mamba_ssu_backend = old @@ -248,6 +263,7 @@ def test_replayssm_flashinfer_call_forwards_explicit_controls(monkeypatch): dt_cache, ring_start, prev_num_accepted, + logical_window=window, scratch=scratch, algorithm="two-kernel", d_split=2, @@ -265,6 +281,40 @@ def test_replayssm_flashinfer_call_forwards_explicit_controls(monkeypatch): assert kwargs["cb_scaled"] is scratch[0] assert kwargs["cumAdt_vec"] is scratch[1] assert kwargs["cb_old"] is scratch[2] + assert kwargs["philox_rounds"] == 10 + + +def test_replayssm_flashinfer_forwards_explicit_philox_rounds(monkeypatch): + import vllm.model_executor.layers.mamba.ops.ssu_dispatch as mod + + kernel = Mock(return_value=torch.empty(1, 1, 2, 4)) + monkeypatch.setattr(mod, "_flashinfer_replayssm_kernel", kernel) + tensor = torch.empty(1, 2, 4) + state = torch.empty(1, 2, 4, 8) + group = torch.empty(1, 1, 8) + + selective_state_update_replayssm_flashinfer( + state, + tensor, + tensor, + torch.empty(2, 4, 8), + group, + group, + tensor, + torch.empty(1, 2, 17, 4), + torch.empty(1, 1, 17, 8), + torch.empty(1, 2, 17), + torch.zeros(1, dtype=torch.int32), + torch.zeros(1, dtype=torch.int32), + logical_window=16, + enable_stochastic_rounding=True, + stochastic_rounding_philox_rounds=6, + update_trackers=False, + ) + + assert kernel.call_args.kwargs["philox_rounds"] == 6 + assert kernel.call_args.kwargs["rand_seed"].shape == (1,) + assert kernel.call_args.kwargs["rand_seed"].dtype == torch.int64 def test_replayssm_requires_native_flashinfer_support(monkeypatch): @@ -273,7 +323,7 @@ def test_replayssm_requires_native_flashinfer_support(monkeypatch): old_module = SimpleNamespace( checkpointing_ssu=Mock(), CheckpointingSSURunner=object ) - monkeypatch.setattr(mod.importlib, "import_module", lambda _: old_module) + monkeypatch.setattr(mod, "import_module", lambda _: old_module) with pytest.raises(ImportError, match="scratch allocation support"): mod._initialize_flashinfer_replayssm(True) @@ -284,7 +334,7 @@ def test_replayssm_flashinfer_import_error(monkeypatch): def raise_import_error(_): raise ImportError - monkeypatch.setattr(mod.importlib, "import_module", raise_import_error) + monkeypatch.setattr(mod, "import_module", raise_import_error) with pytest.raises(ImportError, match="FlashInfer is required"): mod._initialize_flashinfer_replayssm(True) @@ -294,8 +344,43 @@ def raise_import_error(_): reason="compatible flashinfer checkpointing_ssu not available", ) def test_replayssm_flashinfer_backend_init(): + import vllm.model_executor.layers.mamba.ops.ssu_dispatch as mod + initialize_mamba_ssu_backend( MambaConfig(backend=MambaBackendEnum.FLASHINFER), _kv_cache_config_with_ssu(), use_replayssm=True, ) + assert isinstance(get_mamba_ssu_backend(), FlashInferSSUBackend) + assert ( + mod._flashinfer_replayssm_kernel is checkpointing_ssu_module.checkpointing_ssu + ) + + +@pytest.mark.parametrize( + ("backend", "num_speculative_tokens", "expected_ring_len"), + [ + (MambaBackendEnum.TRITON, 0, 16), + (MambaBackendEnum.FLASHINFER, 0, 17), + (MambaBackendEnum.FLASHINFER, 3, 20), + ], +) +def test_replayssm_physical_ring_shape( + backend, num_speculative_tokens, 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, + num_speculative_tokens=num_speculative_tokens, + ) + + assert shapes[2:] == ( + (8, expected_ring_len, 4), + (8, expected_ring_len), + (2, expected_ring_len, 16), + ) diff --git a/tests/test_config.py b/tests/test_config.py index df77d33e0cee..5cfd49c7367c 100644 --- a/tests/test_config.py +++ b/tests/test_config.py @@ -129,6 +129,57 @@ def test_v2_model_runner_env_tri_state(monkeypatch, env_value, expected): assert envs.VLLM_USE_V2_MODEL_RUNNER is expected +def _replayssm_config( + *, + backend: MambaBackendEnum, + use_v2_model_runner: bool = False, + ssu_algorithm: str | None = None, +) -> SimpleNamespace: + return SimpleNamespace( + cache_config=SimpleNamespace( + use_replayssm=True, + mamba_cache_mode="none", + ), + model_config=None, + num_speculative_tokens=0, + mamba_config=SimpleNamespace( + backend=backend, + ssu_algorithm=ssu_algorithm, + ), + use_v2_model_runner=use_v2_model_runner, + kv_transfer_config=None, + ) + + +def test_v2_replayssm_requires_flashinfer(): + config = _replayssm_config( + backend=MambaBackendEnum.TRITON, + use_v2_model_runner=True, + ) + + with pytest.raises(ValueError, match="Triton ReplaySSM does not support"): + VllmConfig.validate_mamba_cached_kernel(config) + + +def test_v2_flashinfer_replayssm_is_supported(): + config = _replayssm_config( + backend=MambaBackendEnum.FLASHINFER, + use_v2_model_runner=True, + ) + + assert VllmConfig.validate_mamba_cached_kernel(config) is config + + +def test_replayssm_rejects_plain_ssu_algorithm_override(): + config = _replayssm_config( + backend=MambaBackendEnum.FLASHINFER, + ssu_algorithm="auto", + ) + + with pytest.raises(ValueError, match="plain FlashInfer SSU tactic"): + VllmConfig.validate_mamba_cached_kernel(config) + + def test_rocm_keeps_compiled_deepseek_defaults(monkeypatch): """ROCm keeps DeepSeek V3.2 and V4 on their compiled MRV1 paths.""" from vllm.config.vllm import ( diff --git a/vllm/config/mamba.py b/vllm/config/mamba.py index e8988aab1940..07a674983bcb 100644 --- a/vllm/config/mamba.py +++ b/vllm/config/mamba.py @@ -46,8 +46,8 @@ class MambaConfig: numerical stability for long sequences.""" stochastic_rounding_philox_rounds: int = 0 """Number of Philox PRNG rounds for stochastic rounding random number - generation. 0 uses the Triton default. Higher values improve randomness - quality at the cost of compute.""" + generation. 0 uses the backend default (10 rounds for FlashInfer). Higher + values improve randomness quality at the cost of compute.""" ssu_algorithm: MambaSSUAlgorithm | None = None """Selective state update algorithm to use with the FlashInfer backend. diff --git a/vllm/config/vllm.py b/vllm/config/vllm.py index d06555b05e8a..06efc64acdd3 100644 --- a/vllm/config/vllm.py +++ b/vllm/config/vllm.py @@ -2709,6 +2709,11 @@ def validate_mamba_cached_kernel(self) -> "VllmConfig": ) if self.num_speculative_tokens > 0: raise ValueError("--use-replayssm does not support speculative decoding") + if self.mamba_config.ssu_algorithm is not None: + raise ValueError( + "--mamba-ssu-algorithm selects a plain FlashInfer SSU tactic " + "and is not compatible with ReplaySSM" + ) if self.mamba_config.backend not in ( MambaBackendEnum.TRITON, MambaBackendEnum.FLASHINFER, @@ -2725,6 +2730,15 @@ def validate_mamba_cached_kernel(self) -> "VllmConfig": "FlashInfer ReplaySSM does not support " "--mamba-cache-mode align yet; use none" ) + if ( + self.use_v2_model_runner + and self.mamba_config.backend == MambaBackendEnum.TRITON + ): + raise ValueError( + "Triton ReplaySSM does not support Model Runner V2 because it " + "requires V1 CPU decode-position metadata; use " + "--mamba-backend flashinfer or Model Runner V1" + ) if ( self.kv_transfer_config is not None and self.kv_transfer_config.is_kv_transfer_instance diff --git a/vllm/model_executor/layers/mamba/mamba_mixer2.py b/vllm/model_executor/layers/mamba/mamba_mixer2.py index dfdbcbab4f8f..8b4382a7ed1c 100644 --- a/vllm/model_executor/layers/mamba/mamba_mixer2.py +++ b/vllm/model_executor/layers/mamba/mamba_mixer2.py @@ -996,7 +996,7 @@ def conv_ssm_forward( # tensor assert state_indices_tensor_p is not None ssm_state[state_indices_tensor_p] = varlen_states - if ring_start is not None: + if ring_start is not None and self._updates_replayssm_trackers: assert prev_num_accepted is not None reset_replayssm_ring_trackers( ring_start, @@ -1098,6 +1098,7 @@ def conv_ssm_forward( dt_cache, ring_start, prev_num_accepted, + logical_window=self.replayssm_buffer_len, D=D_d, dt_bias=dt_bias, dt_softplus=True, @@ -1185,14 +1186,13 @@ def get_state_shape(self) -> tuple[tuple[int, ...], ...]: ) if self.use_replayssm: assert self.replayssm_buffer_len is not None - ring_buffer_len = self.replayssm_buffer_len - if self.mamba_config.backend == MambaBackendEnum.FLASHINFER: - ring_buffer_len += 1 return MambaStateShapeCalculator.append_replayssm_ring( - base_shape, - self.n_groups, - tp_world_size, - ring_buffer_len, + base_shapes=base_shape, + n_groups=self.n_groups, + tp_world_size=tp_world_size, + logical_window=self.replayssm_buffer_len, + backend=self.mamba_config.backend, + num_speculative_tokens=self.num_spec, ) return base_shape diff --git a/vllm/model_executor/layers/mamba/mamba_utils.py b/vllm/model_executor/layers/mamba/mamba_utils.py index 9a94576007bc..1b725d08a8a4 100644 --- a/vllm/model_executor/layers/mamba/mamba_utils.py +++ b/vllm/model_executor/layers/mamba/mamba_utils.py @@ -10,6 +10,7 @@ import vllm.envs as envs from vllm.config.cache import MambaDType +from vllm.config.mamba import MambaBackendEnum from vllm.config.model import ModelDType from vllm.distributed import divide from vllm.logger import init_logger @@ -213,13 +214,20 @@ def append_replayssm_ring( base_shapes: tuple[tuple[int, ...], ...], n_groups: int, tp_world_size: int, - ring_buffer_len: int, + logical_window: int, + backend: MambaBackendEnum, + num_speculative_tokens: int = 0, ) -> tuple[tuple[int, ...], ...]: """Append the physical ReplaySSM ring shapes. ``base_shapes[1]`` is ``(nheads // tp, head_dim, state_size)``; B_cache uses the un-extended ``n_groups``. """ + ring_buffer_len = logical_window + if backend == MambaBackendEnum.FLASHINFER: + # FlashInfer keeps the live window, the token being appended, and + # any predicted tokens in the physical ring at the same time. + ring_buffer_len += 1 + num_speculative_tokens local_nheads, head_dim, state_size = base_shapes[1] local_ngroups = divide(n_groups, tp_world_size) return ( diff --git a/vllm/model_executor/layers/mamba/ops/ssu_dispatch.py b/vllm/model_executor/layers/mamba/ops/ssu_dispatch.py index d3706609f57d..a4d53cc05082 100644 --- a/vllm/model_executor/layers/mamba/ops/ssu_dispatch.py +++ b/vllm/model_executor/layers/mamba/ops/ssu_dispatch.py @@ -9,9 +9,10 @@ platforms the baseline SSU backend defaults to CPU. """ -import importlib from abc import ABC, abstractmethod from collections.abc import Callable +from importlib import import_module +from types import ModuleType import torch @@ -75,14 +76,10 @@ def update_replayssm_ring_trackers( prev_num_accepted: torch.Tensor, state_batch_indices: torch.Tensor, logical_window: int, + ring_buffer_len: int, pad_slot_id: int = NULL_BLOCK_ID, ) -> None: - if ring_start.shape != prev_num_accepted.shape: - raise ValueError("ReplaySSM tracker tensors must have matching shapes") - if ring_start.dim() != 1: - raise ValueError("ReplaySSM tracker tensors must be one-dimensional") - if not ring_start.is_contiguous() or not prev_num_accepted.is_contiguous(): - raise ValueError("ReplaySSM tracker tensors must be contiguous") + _validate_replayssm_ring_trackers(ring_start, prev_num_accepted) state_batch_indices = state_batch_indices.reshape(-1) n_slots = state_batch_indices.numel() if n_slots == 0: @@ -94,18 +91,31 @@ def update_replayssm_ring_trackers( state_batch_indices, n_slots, logical_window, - logical_window + 1, + ring_buffer_len, pad_slot_id, BLOCK=block, ) +def _validate_replayssm_ring_trackers( + ring_start: torch.Tensor, + prev_num_accepted: torch.Tensor, +) -> None: + if ring_start.shape != prev_num_accepted.shape: + raise ValueError("ReplaySSM tracker tensors must have matching shapes") + if ring_start.dim() != 1: + raise ValueError("ReplaySSM tracker tensors must be one-dimensional") + if not ring_start.is_contiguous() or not prev_num_accepted.is_contiguous(): + raise ValueError("ReplaySSM tracker tensors must be contiguous") + + def reset_replayssm_ring_trackers( ring_start: torch.Tensor, prev_num_accepted: torch.Tensor, state_batch_indices: torch.Tensor, pad_slot_id: int = NULL_BLOCK_ID, ) -> None: + _validate_replayssm_ring_trackers(ring_start, prev_num_accepted) state_batch_indices = state_batch_indices.reshape(-1) n_slots = state_batch_indices.numel() if n_slots == 0: @@ -364,14 +374,9 @@ def __call__( _flashinfer_replayssm_kernel: Callable[..., torch.Tensor] | None = None -def _initialize_flashinfer_replayssm(enabled: bool) -> None: - global _flashinfer_replayssm_kernel - _flashinfer_replayssm_kernel = None - if not enabled: - return - +def _load_flashinfer_replayssm_module() -> ModuleType: try: - module = importlib.import_module("flashinfer.mamba.checkpointing_ssu") + module = import_module("flashinfer.mamba.checkpointing_ssu") except (ImportError, ModuleNotFoundError) as e: raise ImportError( "FlashInfer is required for the flashinfer ReplaySSM backend. " @@ -389,6 +394,37 @@ def _initialize_flashinfer_replayssm(enabled: bool) -> None: "FlashInfer ReplaySSM requires native autotuning and scratch " f"allocation support; missing {missing}." ) + return module + + +def allocate_flashinfer_replayssm_scratch( + *, + batch_size: int, + num_heads: int, + num_predicted_tokens: int, + max_window: int, + dtype: torch.dtype, + device: torch.device, +) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: + """Allocate native FlashInfer ReplaySSM scratch after capability checks.""" + module = _load_flashinfer_replayssm_module() + return module.allocate_checkpointing_ssu_scratch( + batch_size=batch_size, + num_heads=num_heads, + num_predicted_tokens=num_predicted_tokens, + max_window=max_window, + dtype=dtype, + device=device, + ) + + +def _initialize_flashinfer_replayssm(enabled: bool) -> None: + global _flashinfer_replayssm_kernel + _flashinfer_replayssm_kernel = None + if not enabled: + return + + module = _load_flashinfer_replayssm_module() _flashinfer_replayssm_kernel = module.checkpointing_ssu @@ -405,9 +441,9 @@ def selective_state_update_replayssm_flashinfer( dt_cache: torch.Tensor, ring_start: torch.Tensor, prev_num_accepted_tokens: torch.Tensor, + logical_window: int, D: torch.Tensor | None = None, dt_bias: torch.Tensor | None = None, - z: torch.Tensor | None = None, dt_softplus: bool = False, state_batch_indices: torch.Tensor | None = None, null_block_id: int = NULL_BLOCK_ID, @@ -417,7 +453,7 @@ def selective_state_update_replayssm_flashinfer( precompute_heads_per_cta: int = 0, update_trackers: bool = True, enable_stochastic_rounding: bool = False, - stochastic_rounding_philox_rounds: int | None = None, + stochastic_rounding_philox_rounds: int = 0, ) -> torch.Tensor: """Run FlashInfer checkpointing SSU and optionally advance shared trackers.""" if _flashinfer_replayssm_kernel is None: @@ -432,7 +468,6 @@ def selective_state_update_replayssm_flashinfer( B = B.unsqueeze(1) C = C.unsqueeze(1) out = out.unsqueeze(1) - z = z.unsqueeze(1) if z is not None else None indices = state_batch_indices if indices is not None and indices.dim() > 1: @@ -461,7 +496,6 @@ def selective_state_update_replayssm_flashinfer( C, out, D=D, - z=z, dt_bias=dt_bias, dt_softplus=dt_softplus, state_batch_indices=indices, @@ -480,7 +514,8 @@ def selective_state_update_replayssm_flashinfer( ring_start, prev_num_accepted_tokens, indices, - logical_window=x_cache.size(2) - x.size(1), + logical_window=logical_window, + ring_buffer_len=x_cache.size(2), pad_slot_id=null_block_id, ) return result diff --git a/vllm/model_executor/models/nemotron_h.py b/vllm/model_executor/models/nemotron_h.py index 65f478cdbdb2..a3401dfcf522 100644 --- a/vllm/model_executor/models/nemotron_h.py +++ b/vllm/model_executor/models/nemotron_h.py @@ -26,7 +26,6 @@ from vllm.compilation.decorators import support_torch_compile from vllm.config import CacheConfig, ModelConfig, VllmConfig -from vllm.config.mamba import MambaBackendEnum from vllm.config.parallel import ParallelConfig from vllm.distributed import get_ep_group, get_tensor_model_parallel_world_size from vllm.distributed.communication_op import tensor_model_parallel_all_gather @@ -785,16 +784,13 @@ def get_mamba_state_shape_from_config( num_spec=vllm_config.num_speculative_tokens, ) if cache_config.use_replayssm: - ring_buffer_len = cache_config.replayssm_buffer_len - if vllm_config.mamba_config.backend == MambaBackendEnum.FLASHINFER: - # FlashInfer's physical ring includes room for the new token - # while replaying the B live old tokens. - ring_buffer_len += 1 return MambaStateShapeCalculator.append_replayssm_ring( - base_shape, - hf_config.n_groups, - parallel_config.tensor_parallel_size, - ring_buffer_len, + base_shapes=base_shape, + n_groups=hf_config.n_groups, + tp_world_size=parallel_config.tensor_parallel_size, + logical_window=cache_config.replayssm_buffer_len, + backend=vllm_config.mamba_config.backend, + num_speculative_tokens=vllm_config.num_speculative_tokens, ) return base_shape diff --git a/vllm/v1/attention/backends/mamba_attn.py b/vllm/v1/attention/backends/mamba_attn.py index ce27e32ef55d..4293d18c1d44 100644 --- a/vllm/v1/attention/backends/mamba_attn.py +++ b/vllm/v1/attention/backends/mamba_attn.py @@ -206,13 +206,13 @@ def __init__( ) self.decode_replayssm_scratch = None elif self.use_flashinfer_replayssm: - from flashinfer.mamba.checkpointing_ssu import ( - allocate_checkpointing_ssu_scratch, + from vllm.model_executor.layers.mamba.ops.ssu_dispatch import ( + allocate_flashinfer_replayssm_scratch, ) self.decode_bc_pre_scratch = None nheads = kv_cache_spec.shapes[2][0] - self.decode_replayssm_scratch = allocate_checkpointing_ssu_scratch( + self.decode_replayssm_scratch = allocate_flashinfer_replayssm_scratch( batch_size=scheduler_config.max_num_seqs, num_heads=nheads, num_predicted_tokens=1, From fc8aff3862259979f4f3b058ab6c7839ff4fe682 Mon Sep 17 00:00:00 2001 From: Andrii Skliar Date: Mon, 17 Aug 2026 10:36:19 +0200 Subject: [PATCH 18/33] fix(mamba): warm ReplaySSM tracker kernels Signed-off-by: Andrii Skliar --- tests/model_executor/test_kernel_warmup.py | 13 +++- vllm/model_executor/warmup/kernel_warmup.py | 72 +++++++++++++++------ 2 files changed, 66 insertions(+), 19 deletions(-) diff --git a/tests/model_executor/test_kernel_warmup.py b/tests/model_executor/test_kernel_warmup.py index 8696b784f93d..5e5a5da4dc2f 100644 --- a/tests/model_executor/test_kernel_warmup.py +++ b/tests/model_executor/test_kernel_warmup.py @@ -105,9 +105,13 @@ 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), + torch.full((4, 2, 17), 3.0), + torch.full((4, 2, 17), 3.0), ) mixer._replayssm_ring_start = torch.full((4,), 3, dtype=torch.int32) mixer._replayssm_prev_num_accepted = torch.full((4,), 3, dtype=torch.int32) @@ -149,7 +153,14 @@ def test_replayssm_autotune_slots_reset_v2_dummy_tables_and_state(): mixer = MambaMixer2.__new__(MambaMixer2) torch.nn.Module.__init__(mixer) mixer.use_replayssm = True - mixer.kv_cache = (torch.full((4, 2), 3.0),) + 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), + torch.full((4, 2, 17), 3.0), + torch.full((4, 2, 17), 3.0), + ) mixer._replayssm_ring_start = torch.full((4,), 3, dtype=torch.int32) mixer._replayssm_prev_num_accepted = torch.full((4,), 3, dtype=torch.int32) block_tables = SimpleNamespace(get_dummy_block_tables=Mock()) diff --git a/vllm/model_executor/warmup/kernel_warmup.py b/vllm/model_executor/warmup/kernel_warmup.py index ed47d40369a6..a22f4cf25fb2 100644 --- a/vllm/model_executor/warmup/kernel_warmup.py +++ b/vllm/model_executor/warmup/kernel_warmup.py @@ -326,11 +326,12 @@ def _temporary_replayssm_autotune_slots( from vllm.model_executor.layers.mamba.mamba_mixer2 import MambaMixer2 from vllm.model_executor.layers.mamba.ops.ssu_dispatch import ( reset_replayssm_ring_trackers, + update_replayssm_ring_trackers, ) reset_tensors: list[torch.Tensor] = [] seen: set[int] = set() - tracker_pairs: list[tuple[torch.Tensor, torch.Tensor]] = [] + tracker_specs: list[tuple[torch.Tensor, torch.Tensor, int, int]] = [] seen_tracker_pairs: set[tuple[int, int]] = set() for module in runner.get_model().modules(): if not isinstance(module, MambaMixer2) or not module.use_replayssm: @@ -341,7 +342,14 @@ def _temporary_replayssm_autotune_slots( ) tracker_ptrs = (tracker_pair[0].data_ptr(), tracker_pair[1].data_ptr()) if tracker_ptrs not in seen_tracker_pairs: - tracker_pairs.append(tracker_pair) + assert module.replayssm_buffer_len is not None + tracker_specs.append( + ( + *tracker_pair, + module.replayssm_buffer_len, + module.kv_cache[2].size(2), + ) + ) seen_tracker_pairs.add(tracker_ptrs) tensors = ( *module.kv_cache, @@ -368,27 +376,55 @@ def _temporary_replayssm_autotune_slots( for block_table in block_tables: block_table.block_table.np[:max_num_reqs, 0] = dummy_block_ids - if tracker_pairs and tracker_pairs[0][0].is_cuda: + if tracker_specs and tracker_specs[0][0].is_cuda: if use_v2_model_runner: - # Match the strided first-column view used by production prefill; - # Triton specializes this tracker reset separately from a fresh, - # contiguous arange tensor. - state_batch_indices = v2_runner.block_tables.get_dummy_block_tables( - max_num_reqs, first_block_id=1 - )[0][:, 0] + # Match every group-specific first-column stride and pointer + # alignment used by production mixed decode/prefill batches. + # Triton specializes the tracker kernels for these views. Four + # row offsets cover every int32 pointer alignment class; retain a + # one-element view as well because scalar value 1 is specialized. + state_batch_indices_variants = tuple( + block_table[offset:, 0] + for block_table in v2_runner.block_tables.get_dummy_block_tables( + max_num_reqs, first_block_id=1 + ) + for offset in sorted( + {0, *range(1, min(4, max_num_reqs)), max_num_reqs - 1} + ) + ) else: - state_batch_indices = torch.arange( - 1, - max_num_reqs + 1, - dtype=torch.int32, - device=tracker_pairs[0][0].device, + state_batch_indices_variants = ( + torch.arange( + 1, + max_num_reqs + 1, + dtype=torch.int32, + device=tracker_specs[0][0].device, + ), ) - for ring_start, prev_num_accepted in tracker_pairs: - reset_replayssm_ring_trackers( + for state_batch_indices in state_batch_indices_variants: + for ( ring_start, prev_num_accepted, - state_batch_indices, - ) + logical_window, + ring_buffer_len, + ) in tracker_specs: + reset_replayssm_ring_trackers( + ring_start, + prev_num_accepted, + state_batch_indices, + ) + update_replayssm_ring_trackers( + ring_start, + prev_num_accepted, + state_batch_indices, + logical_window, + ring_buffer_len, + ) + reset_replayssm_ring_trackers( + ring_start, + prev_num_accepted, + state_batch_indices, + ) try: yield From 06aaf924f246e24923aca15f46ac473c5fbb1b30 Mon Sep 17 00:00:00 2001 From: Andrii Skliar Date: Mon, 17 Aug 2026 10:43:28 +0200 Subject: [PATCH 19/33] fix(mamba): cover V1 tracker warmup layouts Signed-off-by: Andrii Skliar --- vllm/model_executor/warmup/kernel_warmup.py | 37 ++++++++++----------- 1 file changed, 17 insertions(+), 20 deletions(-) diff --git a/vllm/model_executor/warmup/kernel_warmup.py b/vllm/model_executor/warmup/kernel_warmup.py index a22f4cf25fb2..3ddf8fb5fdeb 100644 --- a/vllm/model_executor/warmup/kernel_warmup.py +++ b/vllm/model_executor/warmup/kernel_warmup.py @@ -378,29 +378,26 @@ def _temporary_replayssm_autotune_slots( if tracker_specs and tracker_specs[0][0].is_cuda: if use_v2_model_runner: - # Match every group-specific first-column stride and pointer - # alignment used by production mixed decode/prefill batches. - # Triton specializes the tracker kernels for these views. Four - # row offsets cover every int32 pointer alignment class; retain a - # one-element view as well because scalar value 1 is specialized. - state_batch_indices_variants = tuple( - block_table[offset:, 0] - for block_table in v2_runner.block_tables.get_dummy_block_tables( - max_num_reqs, first_block_id=1 - ) - for offset in sorted( - {0, *range(1, min(4, max_num_reqs)), max_num_reqs - 1} - ) + index_block_tables = v2_runner.block_tables.get_dummy_block_tables( + max_num_reqs, first_block_id=1 ) else: - state_batch_indices_variants = ( - torch.arange( - 1, - max_num_reqs + 1, - dtype=torch.int32, - device=tracker_specs[0][0].device, - ), + assert block_tables is not None + runner.input_batch.block_table.commit_block_table(max_num_reqs) + index_block_tables = tuple( + block_table.block_table.gpu[:max_num_reqs] + for block_table in block_tables ) + # Match every group-specific first-column stride and pointer alignment + # used by production mixed decode/prefill batches. Triton specializes + # the tracker kernels for these views. Four row offsets cover every + # int32 pointer alignment class; retain a one-element view as well + # because scalar value 1 is specialized. + state_batch_indices_variants = tuple( + block_table[offset:, 0] + for block_table in index_block_tables + for offset in sorted({0, *range(1, min(4, max_num_reqs)), max_num_reqs - 1}) + ) for state_batch_indices in state_batch_indices_variants: for ( ring_start, From e5e9466eacd5d11ad62689f5ba21c060277f4f5f Mon Sep 17 00:00:00 2001 From: Andrii Skliar Date: Mon, 17 Aug 2026 15:46:09 +0200 Subject: [PATCH 20/33] refactor(mamba): simplify ReplaySSM integration Signed-off-by: Andrii Skliar --- tests/kernels/mamba/test_ssu_dispatch.py | 108 ++++------ tests/model_executor/test_kernel_warmup.py | 60 +++--- tests/test_config.py | 18 +- tests/v1/e2e/test_replayssm_decode.py | 19 +- tests/v1/worker/test_gpu_block_table.py | 20 +- vllm/config/mamba.py | 4 +- vllm/config/vllm.py | 40 ++-- .../layers/mamba/mamba_mixer2.py | 4 +- .../layers/mamba/ops/ssu_dispatch.py | 192 ++++++------------ vllm/model_executor/warmup/kernel_warmup.py | 114 ++++------- vllm/v1/attention/backends/mamba_attn.py | 11 +- vllm/v1/worker/gpu/block_table.py | 16 +- vllm/v1/worker/gpu/model_runner.py | 23 ++- 13 files changed, 237 insertions(+), 392 deletions(-) diff --git a/tests/kernels/mamba/test_ssu_dispatch.py b/tests/kernels/mamba/test_ssu_dispatch.py index ee3edafd31c1..7f85da1e666f 100644 --- a/tests/kernels/mamba/test_ssu_dispatch.py +++ b/tests/kernels/mamba/test_ssu_dispatch.py @@ -1,8 +1,6 @@ # SPDX-License-Identifier: Apache-2.0 # SPDX-FileCopyrightText: Copyright contributors to the vLLM project -import importlib -from types import SimpleNamespace from unittest.mock import Mock import pytest @@ -31,15 +29,17 @@ HAS_FLASHINFER = False try: - checkpointing_ssu_module = importlib.import_module( - "flashinfer.mamba.checkpointing_ssu" + from flashinfer.mamba.checkpointing_ssu import ( + CheckpointingSSURunner, + allocate_checkpointing_ssu_scratch, ) + from flashinfer.mamba.checkpointing_ssu import ( + checkpointing_ssu as checkpointing_ssu_kernel, + ) + HAS_FLASHINFER_CHECKPOINTING_SSU = all( - hasattr(checkpointing_ssu_module, name) - for name in ( - "CheckpointingSSURunner", - "allocate_checkpointing_ssu_scratch", - ) + callable(symbol) + for symbol in (CheckpointingSSURunner, allocate_checkpointing_ssu_scratch) ) except ImportError: HAS_FLASHINFER_CHECKPOINTING_SSU = False @@ -79,6 +79,31 @@ def test_flashinfer_replayssm_ring_tracker_lifecycle(): assert observed[31] == (16, 16) assert observed[32] == (15, 1) + update_replayssm_ring_trackers( + ring_start, + prev_num_accepted, + state_batch_indices, + ) + assert (ring_start[1].item(), prev_num_accepted[1].item()) == (0, 0) + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") +def test_flashinfer_replayssm_ring_tracker_ignores_invalid_slots(): + 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, 2, 1, 0], dtype=torch.int32, device="cuda") + + update_replayssm_ring_trackers( + ring_start, + prev_num_accepted, + state_batch_indices, + logical_window=16, + ring_buffer_len=17, + ) + + assert ring_start.tolist() == [0, 0] + assert prev_num_accepted.tolist() == [0, 1] + def _kv_cache_config_with_ssu( mamba_type: MambaAttentionBackendEnum = MambaAttentionBackendEnum.MAMBA2, @@ -268,6 +293,8 @@ def test_replayssm_flashinfer_call_forwards_explicit_controls(monkeypatch): algorithm="two-kernel", d_split=2, precompute_heads_per_cta=8, + enable_stochastic_rounding=True, + stochastic_rounding_philox_rounds=6, update_trackers=False, ) @@ -281,62 +308,9 @@ def test_replayssm_flashinfer_call_forwards_explicit_controls(monkeypatch): assert kwargs["cb_scaled"] is scratch[0] assert kwargs["cumAdt_vec"] is scratch[1] assert kwargs["cb_old"] is scratch[2] - assert kwargs["philox_rounds"] == 10 - - -def test_replayssm_flashinfer_forwards_explicit_philox_rounds(monkeypatch): - import vllm.model_executor.layers.mamba.ops.ssu_dispatch as mod - - kernel = Mock(return_value=torch.empty(1, 1, 2, 4)) - monkeypatch.setattr(mod, "_flashinfer_replayssm_kernel", kernel) - tensor = torch.empty(1, 2, 4) - state = torch.empty(1, 2, 4, 8) - group = torch.empty(1, 1, 8) - - selective_state_update_replayssm_flashinfer( - state, - tensor, - tensor, - torch.empty(2, 4, 8), - group, - group, - tensor, - torch.empty(1, 2, 17, 4), - torch.empty(1, 1, 17, 8), - torch.empty(1, 2, 17), - torch.zeros(1, dtype=torch.int32), - torch.zeros(1, dtype=torch.int32), - logical_window=16, - enable_stochastic_rounding=True, - stochastic_rounding_philox_rounds=6, - update_trackers=False, - ) - - assert kernel.call_args.kwargs["philox_rounds"] == 6 - assert kernel.call_args.kwargs["rand_seed"].shape == (1,) - assert kernel.call_args.kwargs["rand_seed"].dtype == torch.int64 - - -def test_replayssm_requires_native_flashinfer_support(monkeypatch): - import vllm.model_executor.layers.mamba.ops.ssu_dispatch as mod - - old_module = SimpleNamespace( - checkpointing_ssu=Mock(), CheckpointingSSURunner=object - ) - monkeypatch.setattr(mod, "import_module", lambda _: old_module) - with pytest.raises(ImportError, match="scratch allocation support"): - mod._initialize_flashinfer_replayssm(True) - - -def test_replayssm_flashinfer_import_error(monkeypatch): - import vllm.model_executor.layers.mamba.ops.ssu_dispatch as mod - - def raise_import_error(_): - raise ImportError - - monkeypatch.setattr(mod, "import_module", raise_import_error) - with pytest.raises(ImportError, match="FlashInfer is required"): - mod._initialize_flashinfer_replayssm(True) + assert kwargs["philox_rounds"] == 6 + assert kwargs["rand_seed"].shape == (1,) + assert kwargs["rand_seed"].dtype == torch.int64 @pytest.mark.skipif( @@ -352,9 +326,7 @@ def test_replayssm_flashinfer_backend_init(): use_replayssm=True, ) assert isinstance(get_mamba_ssu_backend(), FlashInferSSUBackend) - assert ( - mod._flashinfer_replayssm_kernel is checkpointing_ssu_module.checkpointing_ssu - ) + assert mod._flashinfer_replayssm_kernel is checkpointing_ssu_kernel @pytest.mark.parametrize( diff --git a/tests/model_executor/test_kernel_warmup.py b/tests/model_executor/test_kernel_warmup.py index 5e5a5da4dc2f..6477118d2d93 100644 --- a/tests/model_executor/test_kernel_warmup.py +++ b/tests/model_executor/test_kernel_warmup.py @@ -13,6 +13,21 @@ from vllm.model_executor.warmup import kernel_warmup as warmup +def _replayssm_mixer() -> MambaMixer2: + 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) + return mixer + + @pytest.mark.parametrize( ("backend", "use_replayssm", "use_v2_model_runner", "expected"), [ @@ -44,7 +59,7 @@ def test_replayssm_autotune_decode_kwargs( "randomize_inputs": True, } - result = warmup._flashinfer_replayssm_autotune_kwargs(runner, prefill_kwargs) + result = warmup._replayssm_autotune_kwargs(runner, prefill_kwargs) if not expected: assert result is None @@ -55,7 +70,7 @@ def test_replayssm_autotune_decode_kwargs( "uniform_decode": True, } if use_v2_model_runner: - expected_kwargs["dummy_first_block_id"] = 1 + expected_kwargs["valid_dummy_state_slots"] = True else: expected_kwargs.update( allow_microbatching=False, @@ -78,7 +93,7 @@ def test_replayssm_autotune_decode_kwargs_clamps_to_state_capacity(): kv_cache_config=SimpleNamespace(num_blocks=5), ) - result = warmup._flashinfer_replayssm_autotune_kwargs(runner, {}) + result = warmup._replayssm_autotune_kwargs(runner, {}) assert result is not None assert result[0] == 4 @@ -98,23 +113,11 @@ def test_replayssm_autotune_decode_kwargs_skips_without_state_slot(): kv_cache_config=SimpleNamespace(num_blocks=1), ) - assert warmup._flashinfer_replayssm_autotune_kwargs(runner, {}) is None + assert warmup._replayssm_autotune_kwargs(runner, {}) 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), - torch.full((4, 2, 17), 3.0), - torch.full((4, 2, 17), 3.0), - ) - mixer._replayssm_ring_start = torch.full((4,), 3, dtype=torch.int32) - mixer._replayssm_prev_num_accepted = torch.full((4,), 3, dtype=torch.int32) + mixer = _replayssm_mixer() block_ids = np.arange(10, 14, dtype=np.int32).reshape(4, 1) original_block_ids = block_ids.copy() @@ -128,7 +131,7 @@ def test_replayssm_autotune_slots_restore_state_and_trackers(): get_model=lambda: SimpleNamespace(modules=lambda: (mixer,)), ) - with warmup._temporary_replayssm_autotune_slots(runner, 2): + with warmup._temporary_replayssm_autotune_state(runner, 2): assert block_ids[:2, 0].tolist() == [1, 2] for tensor in ( *mixer.kv_cache, @@ -138,7 +141,10 @@ def test_replayssm_autotune_slots_restore_state_and_trackers(): tensor[1:3].fill_(9) assert np.array_equal(block_ids, original_block_ids) - multi_group_block_table.commit_block_table.assert_called_once_with(2) + assert multi_group_block_table.commit_block_table.call_args_list == [ + ((2,), {}), + ((2,), {}), + ] for tensor in ( *mixer.kv_cache, mixer._replayssm_ring_start, @@ -150,19 +156,7 @@ def test_replayssm_autotune_slots_restore_state_and_trackers(): def test_replayssm_autotune_slots_reset_v2_dummy_tables_and_state(): - 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), - torch.full((4, 2, 17), 3.0), - torch.full((4, 2, 17), 3.0), - ) - mixer._replayssm_ring_start = torch.full((4,), 3, dtype=torch.int32) - mixer._replayssm_prev_num_accepted = torch.full((4,), 3, dtype=torch.int32) + mixer = _replayssm_mixer() block_tables = SimpleNamespace(get_dummy_block_tables=Mock()) runner = SimpleNamespace( vllm_config=SimpleNamespace(use_v2_model_runner=True), @@ -170,7 +164,7 @@ def test_replayssm_autotune_slots_reset_v2_dummy_tables_and_state(): get_model=lambda: SimpleNamespace(modules=lambda: (mixer,)), ) - with warmup._temporary_replayssm_autotune_slots(runner, 2): + with warmup._temporary_replayssm_autotune_state(runner, 2): for tensor in ( *mixer.kv_cache, mixer._replayssm_ring_start, diff --git a/tests/test_config.py b/tests/test_config.py index 5cfd49c7367c..9eb8bddef457 100644 --- a/tests/test_config.py +++ b/tests/test_config.py @@ -133,7 +133,6 @@ def _replayssm_config( *, backend: MambaBackendEnum, use_v2_model_runner: bool = False, - ssu_algorithm: str | None = None, ) -> SimpleNamespace: return SimpleNamespace( cache_config=SimpleNamespace( @@ -142,10 +141,7 @@ def _replayssm_config( ), model_config=None, num_speculative_tokens=0, - mamba_config=SimpleNamespace( - backend=backend, - ssu_algorithm=ssu_algorithm, - ), + mamba_config=SimpleNamespace(backend=backend), use_v2_model_runner=use_v2_model_runner, kv_transfer_config=None, ) @@ -157,7 +153,7 @@ def test_v2_replayssm_requires_flashinfer(): use_v2_model_runner=True, ) - with pytest.raises(ValueError, match="Triton ReplaySSM does not support"): + with pytest.raises(ValueError, match="requires Model Runner V1"): VllmConfig.validate_mamba_cached_kernel(config) @@ -170,16 +166,6 @@ def test_v2_flashinfer_replayssm_is_supported(): assert VllmConfig.validate_mamba_cached_kernel(config) is config -def test_replayssm_rejects_plain_ssu_algorithm_override(): - config = _replayssm_config( - backend=MambaBackendEnum.FLASHINFER, - ssu_algorithm="auto", - ) - - with pytest.raises(ValueError, match="plain FlashInfer SSU tactic"): - VllmConfig.validate_mamba_cached_kernel(config) - - def test_rocm_keeps_compiled_deepseek_defaults(monkeypatch): """ROCm keeps DeepSeek V3.2 and V4 on their compiled MRV1 paths.""" from vllm.config.vllm import ( diff --git a/tests/v1/e2e/test_replayssm_decode.py b/tests/v1/e2e/test_replayssm_decode.py index ea7ea91b5549..a9036c984cb9 100644 --- a/tests/v1/e2e/test_replayssm_decode.py +++ b/tests/v1/e2e/test_replayssm_decode.py @@ -11,10 +11,11 @@ try: from flashinfer.mamba.checkpointing_ssu import ( + CheckpointingSSURunner, allocate_checkpointing_ssu_scratch, # noqa: F401 ) - HAS_FLASHINFER_CHECKPOINTING_SSU = True + HAS_FLASHINFER_CHECKPOINTING_SSU = callable(CheckpointingSSURunner) except ImportError: HAS_FLASHINFER_CHECKPOINTING_SSU = False @@ -91,6 +92,22 @@ def test_replayssm_flashinfer_decode_matches_baseline(vllm_runner, model_name): ) +@multi_gpu_test(num_gpus=2) +@pytest.mark.skipif( + not HAS_FLASHINFER_CHECKPOINTING_SSU, + reason="flashinfer.mamba.checkpointing_ssu not available", +) +@pytest.mark.parametrize("model_name", [MAMBA2_MODEL]) +def test_replayssm_flashinfer_decode_matches_baseline_tp2(vllm_runner, model_name): + _check_replayssm_parity( + vllm_runner, + model_name, + tensor_parallel_size=2, + mamba_backend="flashinfer", + name_1="replayssm_flashinfer_tp2", + ) + + # Prefix spans several mamba blocks; prefix caching only reuses full blocks. _PC_SENTENCE = ( "In a detailed survey of state space models, the authors compared many " diff --git a/tests/v1/worker/test_gpu_block_table.py b/tests/v1/worker/test_gpu_block_table.py index 7ebde608aa73..76cf61171825 100644 --- a/tests/v1/worker/test_gpu_block_table.py +++ b/tests/v1/worker/test_gpu_block_table.py @@ -1,11 +1,14 @@ # SPDX-License-Identifier: Apache-2.0 # SPDX-FileCopyrightText: Copyright contributors to the vLLM project +from types import SimpleNamespace + import pytest import torch from vllm.platforms import current_platform from vllm.v1.worker.gpu.block_table import BlockTables +from vllm.v1.worker.gpu.model_runner import GPUModelRunner pytestmark = pytest.mark.skipif( not current_platform.is_cuda(), @@ -239,18 +242,21 @@ def test_get_dummy_block_tables_returns_zeroed_rows(): assert dummy[0].data_ptr() == block_tables.input_block_tables[0].data_ptr() -def test_get_dummy_block_tables_can_assign_non_null_state_slots(): - block_tables = BlockTables( +def test_prepare_dummy_attn_can_assign_valid_state_slots(): + runner = object.__new__(GPUModelRunner) + runner.device = torch.device("cuda") + runner.pcp_manager = None + runner.block_tables = BlockTables( block_sizes=[16], max_num_reqs=4, max_num_batched_tokens=64, max_num_blocks_per_group=[8], - device=torch.device("cuda"), + device=runner.device, kernel_block_sizes=[16], ) + input_batch = SimpleNamespace(num_reqs=3, num_tokens=3) - (dummy,) = block_tables.get_dummy_block_tables(3, first_block_id=1) + block_tables, _ = runner.prepare_dummy_attn(input_batch, valid_state_slots=True) - assert dummy[:, 0].tolist() == [1, 2, 3] - assert torch.count_nonzero(dummy[:, 1:]) == 0 - assert dummy.data_ptr() == block_tables.input_block_tables[0].data_ptr() + assert block_tables[0][:, 0].tolist() == [1, 2, 3] + assert torch.count_nonzero(block_tables[0][:, 1:]) == 0 diff --git a/vllm/config/mamba.py b/vllm/config/mamba.py index 07a674983bcb..e8988aab1940 100644 --- a/vllm/config/mamba.py +++ b/vllm/config/mamba.py @@ -46,8 +46,8 @@ class MambaConfig: numerical stability for long sequences.""" stochastic_rounding_philox_rounds: int = 0 """Number of Philox PRNG rounds for stochastic rounding random number - generation. 0 uses the backend default (10 rounds for FlashInfer). Higher - values improve randomness quality at the cost of compute.""" + generation. 0 uses the Triton default. Higher values improve randomness + quality at the cost of compute.""" ssu_algorithm: MambaSSUAlgorithm | None = None """Selective state update algorithm to use with the FlashInfer backend. diff --git a/vllm/config/vllm.py b/vllm/config/vllm.py index 06efc64acdd3..f6306c3fd6f1 100644 --- a/vllm/config/vllm.py +++ b/vllm/config/vllm.py @@ -2670,6 +2670,7 @@ def validate_mamba_cached_kernel(self) -> "VllmConfig": self.cache_config.use_kda_recoverssm = False return self self.cache_config.use_kda_recoverssm = self.num_speculative_tokens > 0 + if self.model_config is not None and not self.model_config.supports_replayssm: raise ValueError( "--use-replayssm is not supported for architecture " @@ -2702,41 +2703,26 @@ def validate_mamba_cached_kernel(self) -> "VllmConfig": raise ValueError( "RecoverSSM currently requires pipeline_parallel_size=1" ) + if self.mamba_config.backend != MambaBackendEnum.TRITON: + raise ValueError("RecoverSSM requires --mamba-backend triton") elif self.cache_config.mamba_cache_mode == "all": raise ValueError( "--use-replayssm supports prefix caching only in align mode; " "pass --mamba-cache-mode align" ) - if self.num_speculative_tokens > 0: - raise ValueError("--use-replayssm does not support speculative decoding") - if self.mamba_config.ssu_algorithm is not None: - raise ValueError( - "--mamba-ssu-algorithm selects a plain FlashInfer SSU tactic " - "and is not compatible with ReplaySSM" - ) - if self.mamba_config.backend not in ( - MambaBackendEnum.TRITON, - MambaBackendEnum.FLASHINFER, - ): - raise ValueError( - "--use-replayssm requires --mamba-backend triton or flashinfer " - f"(got {self.mamba_config.backend.value!r})" - ) - if ( - self.mamba_config.backend == MambaBackendEnum.FLASHINFER - and self.cache_config.mamba_cache_mode == "align" - ): + elif self.mamba_config.backend == MambaBackendEnum.FLASHINFER: + if self.cache_config.mamba_cache_mode == "align": + raise ValueError( + "FlashInfer ReplaySSM does not support " + "--mamba-cache-mode align yet; use none" + ) + elif self.mamba_config.backend != MambaBackendEnum.TRITON: raise ValueError( - "FlashInfer ReplaySSM does not support " - "--mamba-cache-mode align yet; use none" + "--use-replayssm requires --mamba-backend triton or flashinfer" ) - if ( - self.use_v2_model_runner - and self.mamba_config.backend == MambaBackendEnum.TRITON - ): + elif self.use_v2_model_runner: raise ValueError( - "Triton ReplaySSM does not support Model Runner V2 because it " - "requires V1 CPU decode-position metadata; use " + "Triton ReplaySSM requires Model Runner V1; use " "--mamba-backend flashinfer or Model Runner V1" ) if ( diff --git a/vllm/model_executor/layers/mamba/mamba_mixer2.py b/vllm/model_executor/layers/mamba/mamba_mixer2.py index 8b4382a7ed1c..ceb345fb5913 100644 --- a/vllm/model_executor/layers/mamba/mamba_mixer2.py +++ b/vllm/model_executor/layers/mamba/mamba_mixer2.py @@ -40,9 +40,9 @@ mamba_chunk_scan_combined_varlen, ) from vllm.model_executor.layers.mamba.ops.ssu_dispatch import ( - reset_replayssm_ring_trackers, selective_state_update, selective_state_update_replayssm_flashinfer, + update_replayssm_ring_trackers, ) from vllm.model_executor.layers.quantization import QuantizationConfig from vllm.model_executor.model_loader.weight_utils import ( @@ -998,7 +998,7 @@ def conv_ssm_forward( ssm_state[state_indices_tensor_p] = varlen_states if ring_start is not None and self._updates_replayssm_trackers: assert prev_num_accepted is not None - reset_replayssm_ring_trackers( + update_replayssm_ring_trackers( ring_start, prev_num_accepted, state_indices_tensor_p, diff --git a/vllm/model_executor/layers/mamba/ops/ssu_dispatch.py b/vllm/model_executor/layers/mamba/ops/ssu_dispatch.py index a4d53cc05082..1deffb09c408 100644 --- a/vllm/model_executor/layers/mamba/ops/ssu_dispatch.py +++ b/vllm/model_executor/layers/mamba/ops/ssu_dispatch.py @@ -3,16 +3,14 @@ """ Dispatch module for Mamba selective state update (SSU) backends. -Provides a unified ``selective_state_update`` function that dispatches to -Triton, FlashInfer, or CPU based on ``MambaBackendEnum``. It also contains the -FlashInfer ReplaySSM adapter and shared ring-tracker kernels. On CPU-only -platforms the baseline SSU backend defaults to CPU. +Provides a unified `selective_state_update` function that dispatches to +the Triton, FlashInfer, or CPU backend based on the configured +`MambaBackendEnum`. On CPU-only platforms (PowerPC, x86 without CUDA) +the backend defaults to 'cpu'. """ from abc import ABC, abstractmethod from collections.abc import Callable -from importlib import import_module -from types import ModuleType import torch @@ -26,107 +24,80 @@ logger = init_logger(__name__) -@triton.jit +@triton.jit( + do_not_specialize=["n_slots", "state_batch_indices_stride"], + do_not_specialize_on_alignment=["state_batch_indices"], +) def _update_replayssm_ring_trackers_kernel( ring_start, prev_num_accepted, state_batch_indices, + state_batch_indices_stride, n_slots, + num_states, logical_window: tl.constexpr, ring_buffer_len: tl.constexpr, pad_slot_id: tl.constexpr, + RESET: tl.constexpr, BLOCK: tl.constexpr, ) -> None: offsets = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK) mask = offsets < n_slots - slots = tl.load(state_batch_indices + offsets, mask=mask, other=pad_slot_id) - valid = mask & (slots != pad_slot_id) - prev = tl.load(prev_num_accepted + slots, mask=valid, other=0) - start = tl.load(ring_start + slots, mask=valid, other=0) - must_checkpoint = prev + 1 > logical_window - next_start = tl.where( - must_checkpoint, - (start + prev) % ring_buffer_len, - start, + slots = tl.load( + state_batch_indices + offsets * state_batch_indices_stride, + mask=mask, + other=pad_slot_id, ) - next_prev = tl.where(must_checkpoint, 1, prev + 1) - tl.store(ring_start + slots, next_start, mask=valid) - tl.store(prev_num_accepted + slots, next_prev, mask=valid) - - -@triton.jit -def _reset_replayssm_ring_trackers_kernel( - ring_start, - prev_num_accepted, - state_batch_indices, - n_slots, - pad_slot_id: tl.constexpr, - BLOCK: tl.constexpr, -) -> None: - offsets = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK) - mask = offsets < n_slots - slots = tl.load(state_batch_indices + offsets, mask=mask, other=pad_slot_id) - valid = mask & (slots != pad_slot_id) - tl.store(ring_start + slots, 0, mask=valid) - tl.store(prev_num_accepted + slots, 0, mask=valid) + valid = mask & (slots != pad_slot_id) & (slots >= 0) & (slots < num_states) + if RESET: + tl.store(ring_start + slots, 0, mask=valid) + tl.store(prev_num_accepted + slots, 0, mask=valid) + else: + prev = tl.load(prev_num_accepted + slots, mask=valid, other=0) + start = tl.load(ring_start + slots, mask=valid, other=0) + must_checkpoint = prev + 1 > logical_window + next_start = tl.where( + must_checkpoint, + (start + prev) % ring_buffer_len, + start, + ) + next_prev = tl.where(must_checkpoint, 1, prev + 1) + tl.store(ring_start + slots, next_start, mask=valid) + tl.store(prev_num_accepted + slots, next_prev, mask=valid) def update_replayssm_ring_trackers( ring_start: torch.Tensor, prev_num_accepted: torch.Tensor, state_batch_indices: torch.Tensor, - logical_window: int, - ring_buffer_len: int, + logical_window: int | None = None, + ring_buffer_len: int | None = None, pad_slot_id: int = NULL_BLOCK_ID, ) -> None: - _validate_replayssm_ring_trackers(ring_start, prev_num_accepted) - state_batch_indices = state_batch_indices.reshape(-1) + """Reset selected trackers, or advance them when a window is provided.""" + if state_batch_indices.dim() > 1: + state_batch_indices = state_batch_indices[:, 0] n_slots = state_batch_indices.numel() if n_slots == 0: return + reset = logical_window is None + if reset: + logical_window = 0 + ring_buffer_len = 1 + else: + assert ring_buffer_len is not None block = 128 _update_replayssm_ring_trackers_kernel[(triton.cdiv(n_slots, block),)]( ring_start, prev_num_accepted, state_batch_indices, + state_batch_indices.stride(0), n_slots, + min(ring_start.numel(), prev_num_accepted.numel()), logical_window, ring_buffer_len, pad_slot_id, - BLOCK=block, - ) - - -def _validate_replayssm_ring_trackers( - ring_start: torch.Tensor, - prev_num_accepted: torch.Tensor, -) -> None: - if ring_start.shape != prev_num_accepted.shape: - raise ValueError("ReplaySSM tracker tensors must have matching shapes") - if ring_start.dim() != 1: - raise ValueError("ReplaySSM tracker tensors must be one-dimensional") - if not ring_start.is_contiguous() or not prev_num_accepted.is_contiguous(): - raise ValueError("ReplaySSM tracker tensors must be contiguous") - - -def reset_replayssm_ring_trackers( - ring_start: torch.Tensor, - prev_num_accepted: torch.Tensor, - state_batch_indices: torch.Tensor, - pad_slot_id: int = NULL_BLOCK_ID, -) -> None: - _validate_replayssm_ring_trackers(ring_start, prev_num_accepted) - state_batch_indices = state_batch_indices.reshape(-1) - n_slots = state_batch_indices.numel() - if n_slots == 0: - return - block = 128 - _reset_replayssm_ring_trackers_kernel[(triton.cdiv(n_slots, block),)]( - ring_start, - prev_num_accepted, - state_batch_indices, - n_slots, - pad_slot_id, + RESET=reset, BLOCK=block, ) @@ -374,60 +345,6 @@ def __call__( _flashinfer_replayssm_kernel: Callable[..., torch.Tensor] | None = None -def _load_flashinfer_replayssm_module() -> ModuleType: - try: - module = import_module("flashinfer.mamba.checkpointing_ssu") - except (ImportError, ModuleNotFoundError) as e: - raise ImportError( - "FlashInfer is required for the flashinfer ReplaySSM backend. " - "Install a compatible flashinfer-python package." - ) from e - - required = ( - "checkpointing_ssu", - "CheckpointingSSURunner", - "allocate_checkpointing_ssu_scratch", - ) - missing = [name for name in required if not hasattr(module, name)] - if missing: - raise ImportError( - "FlashInfer ReplaySSM requires native autotuning and scratch " - f"allocation support; missing {missing}." - ) - return module - - -def allocate_flashinfer_replayssm_scratch( - *, - batch_size: int, - num_heads: int, - num_predicted_tokens: int, - max_window: int, - dtype: torch.dtype, - device: torch.device, -) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: - """Allocate native FlashInfer ReplaySSM scratch after capability checks.""" - module = _load_flashinfer_replayssm_module() - return module.allocate_checkpointing_ssu_scratch( - batch_size=batch_size, - num_heads=num_heads, - num_predicted_tokens=num_predicted_tokens, - max_window=max_window, - dtype=dtype, - device=device, - ) - - -def _initialize_flashinfer_replayssm(enabled: bool) -> None: - global _flashinfer_replayssm_kernel - _flashinfer_replayssm_kernel = None - if not enabled: - return - - module = _load_flashinfer_replayssm_module() - _flashinfer_replayssm_kernel = module.checkpointing_ssu - - def selective_state_update_replayssm_flashinfer( state: torch.Tensor, x: torch.Tensor, @@ -536,7 +453,7 @@ def initialize_mamba_ssu_backend( ): return - global _mamba_ssu_backend + global _flashinfer_replayssm_kernel, _mamba_ssu_backend backend = mamba_config.backend if backend == MambaBackendEnum.TRITON: @@ -565,9 +482,20 @@ def initialize_mamba_ssu_backend( _mamba_ssu_backend = backend_cls(mamba_config) logger.info("Using %s Mamba SSU backend.", _mamba_ssu_backend.name) - _initialize_flashinfer_replayssm( - use_replayssm and backend == MambaBackendEnum.FLASHINFER - ) + _flashinfer_replayssm_kernel = None + if use_replayssm and backend == MambaBackendEnum.FLASHINFER: + try: + from flashinfer.mamba.checkpointing_ssu import ( + CheckpointingSSURunner, + checkpointing_ssu, + ) + except ImportError as e: + raise ImportError( + "FlashInfer ReplaySSM requires a compatible flashinfer-python package" + ) from e + if not callable(CheckpointingSSURunner): + raise ImportError("FlashInfer ReplaySSM requires native autotuning support") + _flashinfer_replayssm_kernel = checkpointing_ssu if use_replayssm: logger.info("Using %s ReplaySSM backend.", backend.value) diff --git a/vllm/model_executor/warmup/kernel_warmup.py b/vllm/model_executor/warmup/kernel_warmup.py index 3ddf8fb5fdeb..0e897a2c1d3d 100644 --- a/vllm/model_executor/warmup/kernel_warmup.py +++ b/vllm/model_executor/warmup/kernel_warmup.py @@ -275,7 +275,7 @@ def _run_flashinfer_autotune_dummy_runs(runner: "GPUModelRunner") -> None: ) -def _flashinfer_replayssm_autotune_kwargs( +def _replayssm_autotune_kwargs( runner: "GPUModelRunner", max_token_prefill_kwargs: dict[str, Any] ) -> tuple[int, dict[str, Any]] | None: config = runner.vllm_config @@ -309,7 +309,7 @@ def _flashinfer_replayssm_autotune_kwargs( "uniform_decode": True, } if use_v2_model_runner: - decode_kwargs["dummy_first_block_id"] = 1 + decode_kwargs["valid_dummy_state_slots"] = True else: decode_kwargs.update( allow_microbatching=False, @@ -320,48 +320,39 @@ def _flashinfer_replayssm_autotune_kwargs( @contextmanager -def _temporary_replayssm_autotune_slots( +def _temporary_replayssm_autotune_state( runner: "GPUModelRunner", max_num_reqs: int ) -> Iterator[None]: from vllm.model_executor.layers.mamba.mamba_mixer2 import MambaMixer2 from vllm.model_executor.layers.mamba.ops.ssu_dispatch import ( - reset_replayssm_ring_trackers, update_replayssm_ring_trackers, ) - reset_tensors: list[torch.Tensor] = [] - seen: set[int] = set() - tracker_specs: list[tuple[torch.Tensor, torch.Tensor, int, int]] = [] - seen_tracker_pairs: set[tuple[int, int]] = set() + reset_tensors: dict[int, torch.Tensor] = {} + tracker_specs: dict[int, tuple[torch.Tensor, torch.Tensor, int, int]] = {} for module in runner.get_model().modules(): if not isinstance(module, MambaMixer2) or not module.use_replayssm: continue - tracker_pair = ( - module._replayssm_ring_start, - module._replayssm_prev_num_accepted, + assert module.replayssm_buffer_len is not None + ring_start = module._replayssm_ring_start + prev_num_accepted = module._replayssm_prev_num_accepted + tracker_specs.setdefault( + ring_start.data_ptr(), + ( + ring_start, + prev_num_accepted, + module.replayssm_buffer_len, + module.kv_cache[2].size(2), + ), ) - tracker_ptrs = (tracker_pair[0].data_ptr(), tracker_pair[1].data_ptr()) - if tracker_ptrs not in seen_tracker_pairs: - assert module.replayssm_buffer_len is not None - tracker_specs.append( - ( - *tracker_pair, - module.replayssm_buffer_len, - module.kv_cache[2].size(2), - ) - ) - seen_tracker_pairs.add(tracker_ptrs) tensors = ( *module.kv_cache, - *tracker_pair, + ring_start, + prev_num_accepted, ) for tensor in tensors: - if not tensor.numel(): - continue - data_ptr = tensor.data_ptr() - if data_ptr not in seen: - reset_tensors.append(tensor) - seen.add(data_ptr) + if tensor.numel(): + reset_tensors.setdefault(tensor.data_ptr(), tensor) use_v2_model_runner = runner.vllm_config.use_v2_model_runner v2_runner: Any = runner @@ -375,53 +366,28 @@ def _temporary_replayssm_autotune_slots( dummy_block_ids = range(1, max_num_reqs + 1) for block_table in block_tables: block_table.block_table.np[:max_num_reqs, 0] = dummy_block_ids + runner.input_batch.block_table.commit_block_table(max_num_reqs) - if tracker_specs and tracker_specs[0][0].is_cuda: - if use_v2_model_runner: - index_block_tables = v2_runner.block_tables.get_dummy_block_tables( - max_num_reqs, first_block_id=1 - ) - else: - assert block_tables is not None - runner.input_batch.block_table.commit_block_table(max_num_reqs) - index_block_tables = tuple( - block_table.block_table.gpu[:max_num_reqs] - for block_table in block_tables - ) - # Match every group-specific first-column stride and pointer alignment - # used by production mixed decode/prefill batches. Triton specializes - # the tracker kernels for these views. Four row offsets cover every - # int32 pointer alignment class; retain a one-element view as well - # because scalar value 1 is specialized. - state_batch_indices_variants = tuple( - block_table[offset:, 0] - for block_table in index_block_tables - for offset in sorted({0, *range(1, min(4, max_num_reqs)), max_num_reqs - 1}) + first_tracker = next(iter(tracker_specs.values()), None) + if first_tracker is not None and first_tracker[0].is_cuda: + state_slots = torch.arange( + 1, max_num_reqs + 1, dtype=torch.int32, device=first_tracker[0].device ) - for state_batch_indices in state_batch_indices_variants: - for ( + for ( + ring_start, + prev_num_accepted, + logical_window, + ring_buffer_len, + ) in tracker_specs.values(): + update_replayssm_ring_trackers(ring_start, prev_num_accepted, state_slots) + update_replayssm_ring_trackers( ring_start, prev_num_accepted, + state_slots, logical_window, ring_buffer_len, - ) in tracker_specs: - reset_replayssm_ring_trackers( - ring_start, - prev_num_accepted, - state_batch_indices, - ) - update_replayssm_ring_trackers( - ring_start, - prev_num_accepted, - state_batch_indices, - logical_window, - ring_buffer_len, - ) - reset_replayssm_ring_trackers( - ring_start, - prev_num_accepted, - state_batch_indices, - ) + ) + update_replayssm_ring_trackers(ring_start, prev_num_accepted, state_slots) try: yield @@ -433,7 +399,7 @@ def _temporary_replayssm_autotune_slots( for block_table, block_ids in zip(block_tables, saved_block_ids): block_table.block_table.np[:max_num_reqs, 0] = block_ids runner.input_batch.block_table.commit_block_table(max_num_reqs) - for tensor in reset_tensors: + for tensor in reset_tensors.values(): tensor[1 : max_num_reqs + 1].zero_() @@ -484,9 +450,7 @@ def flashinfer_autotune(runner: "GPUModelRunner") -> None: is_profile=True, randomize_inputs=True, ) - replayssm_autotune = _flashinfer_replayssm_autotune_kwargs( - runner, max_token_prefill_kwargs - ) + replayssm_autotune = _replayssm_autotune_kwargs(runner, max_token_prefill_kwargs) # Read cached autotune results and broadcast to all ranks. cached_results: bytes | None = None if is_leader and cache_path.exists(): @@ -508,7 +472,7 @@ def flashinfer_autotune(runner: "GPUModelRunner") -> None: _run_flashinfer_autotune_dummy_runs(runner) if replayssm_autotune is not None: max_num_reqs, max_batch_decode_kwargs = replayssm_autotune - with _temporary_replayssm_autotune_slots(runner, max_num_reqs): + with _temporary_replayssm_autotune_state(runner, max_num_reqs): runner._dummy_run(**max_batch_decode_kwargs) finally: set_autotune_process_group(None) diff --git a/vllm/v1/attention/backends/mamba_attn.py b/vllm/v1/attention/backends/mamba_attn.py index 4293d18c1d44..5f2a6597bf5c 100644 --- a/vllm/v1/attention/backends/mamba_attn.py +++ b/vllm/v1/attention/backends/mamba_attn.py @@ -174,10 +174,7 @@ def __init__( dtype=torch.int32, device=device, ) - # ReplaySSM CUDA-graph buffers. - # Triton: write_pos / is_flush / bc_pre. - # FlashInfer: two-kernel scratch, so algorithm="auto" can pick the - # monolith or two-kernel implementation. + # ReplaySSM CUDA-graph buffers for the selected backend. if self.use_replayssm and not self.use_flashinfer_replayssm: self.decode_write_pos_d: torch.Tensor = torch.empty( (self.decode_cudagraph_max_bs,), @@ -206,13 +203,13 @@ def __init__( ) self.decode_replayssm_scratch = None elif self.use_flashinfer_replayssm: - from vllm.model_executor.layers.mamba.ops.ssu_dispatch import ( - allocate_flashinfer_replayssm_scratch, + from flashinfer.mamba.checkpointing_ssu import ( + allocate_checkpointing_ssu_scratch, ) self.decode_bc_pre_scratch = None nheads = kv_cache_spec.shapes[2][0] - self.decode_replayssm_scratch = allocate_flashinfer_replayssm_scratch( + self.decode_replayssm_scratch = allocate_checkpointing_ssu_scratch( batch_size=scheduler_config.max_num_seqs, num_heads=nheads, num_predicted_tokens=1, diff --git a/vllm/v1/worker/gpu/block_table.py b/vllm/v1/worker/gpu/block_table.py index 17bb172a9c17..5e09383ee6ab 100644 --- a/vllm/v1/worker/gpu/block_table.py +++ b/vllm/v1/worker/gpu/block_table.py @@ -167,9 +167,7 @@ def gather_block_tables( ) return tuple(bt[:num_reqs_padded] for bt in out) - def get_dummy_block_tables( - self, num_reqs: int, first_block_id: int | None = None - ) -> tuple[torch.Tensor, ...]: + def get_dummy_block_tables(self, num_reqs: int) -> tuple[torch.Tensor, ...]: # NOTE(woosuk): The output may be used for CUDA graph capture. # Therefore, this method must return the persistent tensor # with the same memory address as that used during the model's forward pass, @@ -178,19 +176,9 @@ def get_dummy_block_tables( # Zero the rows so dummy runs write mamba state to the reserved null # block rather than through the previous real step's (stale) block # ids, which may point at blocks since freed and reallocated. - result = tuple( + return tuple( block_table[:num_reqs].zero_() for block_table in self.input_block_tables ) - if first_block_id is not None: - block_ids = torch.arange( - first_block_id, - first_block_id + num_reqs, - dtype=torch.int32, - device=self.device, - ) - for block_table in result: - block_table[:, 0].copy_(block_ids) - return result def compute_slot_mappings( self, diff --git a/vllm/v1/worker/gpu/model_runner.py b/vllm/v1/worker/gpu/model_runner.py index 796752485fab..dcb18073a37f 100644 --- a/vllm/v1/worker/gpu/model_runner.py +++ b/vllm/v1/worker/gpu/model_runner.py @@ -689,7 +689,7 @@ def _dummy_run( context_len: int = 0, skip_eplb: bool = False, is_profile: bool = False, - dummy_first_block_id: int | None = None, + valid_dummy_state_slots: bool = False, **kwargs, ) -> tuple[torch.Tensor | None, torch.Tensor | None]: if skip_attn and not is_profile: @@ -747,7 +747,7 @@ def _dummy_run( skip_attn_for_dummy_run=skip_attn, is_profile=is_profile, context_len=context_len, - dummy_first_block_id=dummy_first_block_id, + valid_dummy_state_slots=valid_dummy_state_slots, ) self.kv_connector.set_disabled(False) @@ -1373,11 +1373,18 @@ def prepare_attn( return block_tables, slot_mappings def prepare_dummy_attn( - self, input_batch: InputBatch, first_block_id: int | None = None + self, input_batch: InputBatch, valid_state_slots: bool = False ) -> tuple[tuple[torch.Tensor, ...], torch.Tensor]: - block_tables = self.block_tables.get_dummy_block_tables( - input_batch.num_reqs, first_block_id - ) + block_tables = self.block_tables.get_dummy_block_tables(input_batch.num_reqs) + if valid_state_slots: + state_slots = torch.arange( + 1, + input_batch.num_reqs + 1, + dtype=torch.int32, + device=self.device, + ) + for block_table in block_tables: + block_table[:, 0].copy_(state_slots) slot_mappings = pcp.maybe_get_pcp_dummy_slot_mappings( self.pcp_manager, self.block_tables, input_batch.num_tokens ) @@ -1510,7 +1517,7 @@ def execute_model( skip_attn_for_dummy_run: bool = False, is_profile: bool = False, context_len: int = 0, - dummy_first_block_id: int | None = None, + valid_dummy_state_slots: bool = False, ) -> ModelRunnerOutput | IntermediateTensors | None: if not dummy_run: # Update the request states. @@ -1611,7 +1618,7 @@ def execute_model( ) if not skip_attn_for_dummy_run: block_tables, slot_mappings = self.prepare_dummy_attn( - input_batch, dummy_first_block_id + input_batch, valid_dummy_state_slots ) if context_len: set_dummy_context( From 5e7d1b19c4298e59e81662fbf5b348d376c3fdc9 Mon Sep 17 00:00:00 2001 From: Andrii Skliar Date: Mon, 17 Aug 2026 18:54:53 +0200 Subject: [PATCH 21/33] refactor(mamba): remove redundant ReplaySSM assignments Signed-off-by: Andrii Skliar --- vllm/model_executor/layers/mamba/mamba_mixer2.py | 3 --- vllm/v1/attention/backends/mamba_attn.py | 11 +++++------ 2 files changed, 5 insertions(+), 9 deletions(-) diff --git a/vllm/model_executor/layers/mamba/mamba_mixer2.py b/vllm/model_executor/layers/mamba/mamba_mixer2.py index ceb345fb5913..f526e00eb278 100644 --- a/vllm/model_executor/layers/mamba/mamba_mixer2.py +++ b/vllm/model_executor/layers/mamba/mamba_mixer2.py @@ -735,11 +735,8 @@ def conv_ssm_forward( if self.mamba_config.backend == MambaBackendEnum.FLASHINFER: ring_start = self._replayssm_ring_start prev_num_accepted = self._replayssm_prev_num_accepted - else: - ring_start = prev_num_accepted = None else: x_cache = dt_cache = B_cache = None - ring_start = prev_num_accepted = None has_initial_states_p = attn_metadata.has_initial_states_p prep_initial_states = attn_metadata.prep_initial_states chunk_size = attn_metadata.chunk_size diff --git a/vllm/v1/attention/backends/mamba_attn.py b/vllm/v1/attention/backends/mamba_attn.py index 5f2a6597bf5c..b1a0c2d0efbf 100644 --- a/vllm/v1/attention/backends/mamba_attn.py +++ b/vllm/v1/attention/backends/mamba_attn.py @@ -174,6 +174,10 @@ def __init__( dtype=torch.int32, device=device, ) + self.decode_bc_pre_scratch: torch.Tensor | None = None + self.decode_replayssm_scratch: ( + tuple[torch.Tensor, torch.Tensor, torch.Tensor] | None + ) = None # ReplaySSM CUDA-graph buffers for the selected backend. if self.use_replayssm and not self.use_flashinfer_replayssm: self.decode_write_pos_d: torch.Tensor = torch.empty( @@ -192,7 +196,7 @@ def __init__( bc_scratch_bs = max( self.decode_cudagraph_max_bs, scheduler_config.max_num_seqs ) - self.decode_bc_pre_scratch: torch.Tensor = torch.empty( + self.decode_bc_pre_scratch = torch.empty( ( bc_scratch_bs, bc_ngroups, @@ -201,13 +205,11 @@ def __init__( dtype=torch.float32, device=device, ) - self.decode_replayssm_scratch = None elif self.use_flashinfer_replayssm: from flashinfer.mamba.checkpointing_ssu import ( allocate_checkpointing_ssu_scratch, ) - self.decode_bc_pre_scratch = None nheads = kv_cache_spec.shapes[2][0] self.decode_replayssm_scratch = allocate_checkpointing_ssu_scratch( batch_size=scheduler_config.max_num_seqs, @@ -217,9 +219,6 @@ def __init__( dtype=vllm_config.model_config.dtype, device=device, ) - else: - self.decode_bc_pre_scratch = None - self.decode_replayssm_scratch = None self._init_reorder_batch_threshold(1, self.use_spec_decode) if self.use_spec_decode: From 1f187055ab3a79e949d55988031796597d48093a Mon Sep 17 00:00:00 2001 From: Andrii Skliar Date: Mon, 17 Aug 2026 18:55:36 +0200 Subject: [PATCH 22/33] refactor(mamba): name ReplaySSM tracker resets Signed-off-by: Andrii Skliar --- tests/kernels/mamba/test_ssu_dispatch.py | 3 ++- vllm/model_executor/layers/mamba/mamba_mixer2.py | 4 ++-- .../layers/mamba/ops/ssu_dispatch.py | 15 +++++++++++++++ vllm/model_executor/warmup/kernel_warmup.py | 7 +++++-- 4 files changed, 24 insertions(+), 5 deletions(-) diff --git a/tests/kernels/mamba/test_ssu_dispatch.py b/tests/kernels/mamba/test_ssu_dispatch.py index 7f85da1e666f..1d4bff4cacb7 100644 --- a/tests/kernels/mamba/test_ssu_dispatch.py +++ b/tests/kernels/mamba/test_ssu_dispatch.py @@ -13,6 +13,7 @@ TritonSSUBackend, get_mamba_ssu_backend, initialize_mamba_ssu_backend, + reset_replayssm_ring_trackers, selective_state_update, selective_state_update_replayssm_flashinfer, update_replayssm_ring_trackers, @@ -79,7 +80,7 @@ def test_flashinfer_replayssm_ring_tracker_lifecycle(): assert observed[31] == (16, 16) assert observed[32] == (15, 1) - update_replayssm_ring_trackers( + reset_replayssm_ring_trackers( ring_start, prev_num_accepted, state_batch_indices, diff --git a/vllm/model_executor/layers/mamba/mamba_mixer2.py b/vllm/model_executor/layers/mamba/mamba_mixer2.py index f526e00eb278..835544954029 100644 --- a/vllm/model_executor/layers/mamba/mamba_mixer2.py +++ b/vllm/model_executor/layers/mamba/mamba_mixer2.py @@ -40,9 +40,9 @@ mamba_chunk_scan_combined_varlen, ) from vllm.model_executor.layers.mamba.ops.ssu_dispatch import ( + reset_replayssm_ring_trackers, selective_state_update, selective_state_update_replayssm_flashinfer, - update_replayssm_ring_trackers, ) from vllm.model_executor.layers.quantization import QuantizationConfig from vllm.model_executor.model_loader.weight_utils import ( @@ -995,7 +995,7 @@ def conv_ssm_forward( ssm_state[state_indices_tensor_p] = varlen_states if ring_start is not None and self._updates_replayssm_trackers: assert prev_num_accepted is not None - update_replayssm_ring_trackers( + reset_replayssm_ring_trackers( ring_start, prev_num_accepted, state_indices_tensor_p, diff --git a/vllm/model_executor/layers/mamba/ops/ssu_dispatch.py b/vllm/model_executor/layers/mamba/ops/ssu_dispatch.py index 1deffb09c408..cc7575e9af4a 100644 --- a/vllm/model_executor/layers/mamba/ops/ssu_dispatch.py +++ b/vllm/model_executor/layers/mamba/ops/ssu_dispatch.py @@ -102,6 +102,21 @@ def update_replayssm_ring_trackers( ) +def reset_replayssm_ring_trackers( + ring_start: torch.Tensor, + prev_num_accepted: torch.Tensor, + state_batch_indices: torch.Tensor, + pad_slot_id: int = NULL_BLOCK_ID, +) -> None: + """Reset selected ReplaySSM ring trackers.""" + update_replayssm_ring_trackers( + ring_start, + prev_num_accepted, + state_batch_indices, + pad_slot_id=pad_slot_id, + ) + + class MambaSSUBackend(ABC): """Abstract base class for Mamba SSU backends.""" diff --git a/vllm/model_executor/warmup/kernel_warmup.py b/vllm/model_executor/warmup/kernel_warmup.py index 0e897a2c1d3d..92c17f1c93cd 100644 --- a/vllm/model_executor/warmup/kernel_warmup.py +++ b/vllm/model_executor/warmup/kernel_warmup.py @@ -325,6 +325,7 @@ def _temporary_replayssm_autotune_state( ) -> Iterator[None]: from vllm.model_executor.layers.mamba.mamba_mixer2 import MambaMixer2 from vllm.model_executor.layers.mamba.ops.ssu_dispatch import ( + reset_replayssm_ring_trackers, update_replayssm_ring_trackers, ) @@ -379,7 +380,9 @@ def _temporary_replayssm_autotune_state( logical_window, ring_buffer_len, ) in tracker_specs.values(): - update_replayssm_ring_trackers(ring_start, prev_num_accepted, state_slots) + # Compile reset (prefill) and advance (decode) before inference. + # The final reset leaves the decode tuning run in a clean state. + reset_replayssm_ring_trackers(ring_start, prev_num_accepted, state_slots) update_replayssm_ring_trackers( ring_start, prev_num_accepted, @@ -387,7 +390,7 @@ def _temporary_replayssm_autotune_state( logical_window, ring_buffer_len, ) - update_replayssm_ring_trackers(ring_start, prev_num_accepted, state_slots) + reset_replayssm_ring_trackers(ring_start, prev_num_accepted, state_slots) try: yield From 200b7c700a84463e9154d41e50264bd9fa3eff7a Mon Sep 17 00:00:00 2001 From: Andrii Skliar Date: Mon, 17 Aug 2026 18:55:49 +0200 Subject: [PATCH 23/33] refactor(mamba): drop unreachable ReplaySSM ring sizing Signed-off-by: Andrii Skliar --- tests/kernels/mamba/test_ssu_dispatch.py | 12 ++++-------- vllm/model_executor/layers/mamba/mamba_mixer2.py | 1 - vllm/model_executor/layers/mamba/mamba_utils.py | 6 ++---- vllm/model_executor/models/nemotron_h.py | 1 - 4 files changed, 6 insertions(+), 14 deletions(-) diff --git a/tests/kernels/mamba/test_ssu_dispatch.py b/tests/kernels/mamba/test_ssu_dispatch.py index 1d4bff4cacb7..fa26968f186f 100644 --- a/tests/kernels/mamba/test_ssu_dispatch.py +++ b/tests/kernels/mamba/test_ssu_dispatch.py @@ -331,16 +331,13 @@ def test_replayssm_flashinfer_backend_init(): @pytest.mark.parametrize( - ("backend", "num_speculative_tokens", "expected_ring_len"), + ("backend", "expected_ring_len"), [ - (MambaBackendEnum.TRITON, 0, 16), - (MambaBackendEnum.FLASHINFER, 0, 17), - (MambaBackendEnum.FLASHINFER, 3, 20), + (MambaBackendEnum.TRITON, 16), + (MambaBackendEnum.FLASHINFER, 17), ], ) -def test_replayssm_physical_ring_shape( - backend, num_speculative_tokens, expected_ring_len -): +def test_replayssm_physical_ring_shape(backend, expected_ring_len): base_shapes = ((64, 3), (8, 4, 16)) shapes = MambaStateShapeCalculator.append_replayssm_ring( @@ -349,7 +346,6 @@ def test_replayssm_physical_ring_shape( tp_world_size=2, logical_window=16, backend=backend, - num_speculative_tokens=num_speculative_tokens, ) assert shapes[2:] == ( diff --git a/vllm/model_executor/layers/mamba/mamba_mixer2.py b/vllm/model_executor/layers/mamba/mamba_mixer2.py index 835544954029..d0b9dc59cc35 100644 --- a/vllm/model_executor/layers/mamba/mamba_mixer2.py +++ b/vllm/model_executor/layers/mamba/mamba_mixer2.py @@ -1189,7 +1189,6 @@ def get_state_shape(self) -> tuple[tuple[int, ...], ...]: tp_world_size=tp_world_size, logical_window=self.replayssm_buffer_len, backend=self.mamba_config.backend, - num_speculative_tokens=self.num_spec, ) return base_shape diff --git a/vllm/model_executor/layers/mamba/mamba_utils.py b/vllm/model_executor/layers/mamba/mamba_utils.py index 1b725d08a8a4..bc579868f3db 100644 --- a/vllm/model_executor/layers/mamba/mamba_utils.py +++ b/vllm/model_executor/layers/mamba/mamba_utils.py @@ -216,7 +216,6 @@ def append_replayssm_ring( tp_world_size: int, logical_window: int, backend: MambaBackendEnum, - num_speculative_tokens: int = 0, ) -> tuple[tuple[int, ...], ...]: """Append the physical ReplaySSM ring shapes. @@ -225,9 +224,8 @@ def append_replayssm_ring( """ ring_buffer_len = logical_window if backend == MambaBackendEnum.FLASHINFER: - # FlashInfer keeps the live window, the token being appended, and - # any predicted tokens in the physical ring at the same time. - ring_buffer_len += 1 + num_speculative_tokens + # FlashInfer keeps the live window and appended token together. + ring_buffer_len += 1 local_nheads, head_dim, state_size = base_shapes[1] local_ngroups = divide(n_groups, tp_world_size) return ( diff --git a/vllm/model_executor/models/nemotron_h.py b/vllm/model_executor/models/nemotron_h.py index a3401dfcf522..e65e51f3bb9d 100644 --- a/vllm/model_executor/models/nemotron_h.py +++ b/vllm/model_executor/models/nemotron_h.py @@ -790,7 +790,6 @@ def get_mamba_state_shape_from_config( tp_world_size=parallel_config.tensor_parallel_size, logical_window=cache_config.replayssm_buffer_len, backend=vllm_config.mamba_config.backend, - num_speculative_tokens=vllm_config.num_speculative_tokens, ) return base_shape From 7a987408f1a0d95fd79a57b385c7e608375d39fc Mon Sep 17 00:00:00 2001 From: Andrii Skliar Date: Thu, 27 Aug 2026 19:57:10 +0200 Subject: [PATCH 24/33] Clean up ReplaySSM tracker grouping and add V2 e2e coverage. Move tracker sharing into share_replayssm_ring_trackers keyed by execution order, drop test-only FlashInfer ReplaySSM tactic kwargs so checkpointing_ssu uses native defaults, and add ModelRunnerV2 FlashInfer parity tests. Co-authored-by: Composer Signed-off-by: Andrii Skliar --- tests/kernels/mamba/test_ssu_dispatch.py | 45 +++++++------- tests/v1/e2e/test_replayssm_decode.py | 50 ++++++++++++++++ tests/v1/worker/test_utils.py | 6 +- .../layers/mamba/mamba_mixer2.py | 58 ++++++++++++++----- .../layers/mamba/ops/ssu_dispatch.py | 6 -- vllm/model_executor/warmup/kernel_warmup.py | 10 ++-- vllm/v1/worker/utils.py | 40 +------------ 7 files changed, 128 insertions(+), 87 deletions(-) diff --git a/tests/kernels/mamba/test_ssu_dispatch.py b/tests/kernels/mamba/test_ssu_dispatch.py index fa26968f186f..a9f56a1e25d6 100644 --- a/tests/kernels/mamba/test_ssu_dispatch.py +++ b/tests/kernels/mamba/test_ssu_dispatch.py @@ -255,26 +255,34 @@ def test_triton_basic_call(): assert not torch.isnan(out).any() -def test_replayssm_flashinfer_call_forwards_explicit_controls(monkeypatch): +@pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") +def test_replayssm_flashinfer_call_forwards_scratch_and_rounding(monkeypatch): import vllm.model_executor.layers.mamba.ops.ssu_dispatch as mod - kernel = Mock(return_value=torch.empty(1, 1, 2, 4)) + kernel = Mock(return_value=torch.empty(1, 1, 2, 4, device="cuda")) + tracker = Mock() monkeypatch.setattr(mod, "_flashinfer_replayssm_kernel", kernel) + monkeypatch.setattr(mod, "update_replayssm_ring_trackers", tracker) batch, nheads, dim, dstate, ngroups, window = 1, 2, 4, 8, 1, 16 - state = torch.empty(1, nheads, dim, dstate) - x = torch.empty(batch, nheads, dim) - dt = torch.empty(batch, nheads, dim) - A = torch.empty(nheads, dim, dstate) - B = torch.empty(batch, ngroups, dstate) - C = torch.empty(batch, ngroups, dstate) + state = torch.empty(1, nheads, dim, dstate, device="cuda") + x = torch.empty(batch, nheads, dim, device="cuda") + dt = torch.empty(batch, nheads, dim, device="cuda") + A = torch.empty(nheads, dim, dstate, device="cuda") + B = torch.empty(batch, ngroups, dstate, device="cuda") + C = torch.empty(batch, ngroups, dstate, device="cuda") out = torch.empty_like(x) - x_cache = torch.empty(1, nheads, window, dim) - dt_cache = torch.empty(1, nheads, window) - B_cache = torch.empty(1, ngroups, window, dstate) - ring_start = torch.zeros(1, dtype=torch.int32) - prev_num_accepted = torch.zeros(1, dtype=torch.int32) - scratch = (torch.empty(1), torch.empty(1), torch.empty(1)) + x_cache = torch.empty(1, nheads, window, dim, device="cuda") + dt_cache = torch.empty(1, nheads, window, device="cuda") + B_cache = torch.empty(1, ngroups, window, dstate, device="cuda") + ring_start = torch.zeros(1, dtype=torch.int32, device="cuda") + prev_num_accepted = torch.zeros(1, dtype=torch.int32, device="cuda") + state_batch_indices = torch.zeros(1, dtype=torch.int32, device="cuda") + scratch = ( + torch.empty(1, device="cuda"), + torch.empty(1, device="cuda"), + torch.empty(1, device="cuda"), + ) selective_state_update_replayssm_flashinfer( state, @@ -290,10 +298,8 @@ def test_replayssm_flashinfer_call_forwards_explicit_controls(monkeypatch): ring_start, prev_num_accepted, logical_window=window, + state_batch_indices=state_batch_indices, scratch=scratch, - algorithm="two-kernel", - d_split=2, - precompute_heads_per_cta=8, enable_stochastic_rounding=True, stochastic_rounding_philox_rounds=6, update_trackers=False, @@ -303,15 +309,14 @@ def test_replayssm_flashinfer_call_forwards_explicit_controls(monkeypatch): kwargs = kernel.call_args.kwargs assert args[4] is ring_start assert args[5] is prev_num_accepted - assert kwargs["algorithm"] == "two-kernel" - assert kwargs["d_split"] == 2 - assert kwargs["precompute_heads_per_cta"] == 8 assert kwargs["cb_scaled"] is scratch[0] assert kwargs["cumAdt_vec"] is scratch[1] assert kwargs["cb_old"] is scratch[2] assert kwargs["philox_rounds"] == 6 assert kwargs["rand_seed"].shape == (1,) assert kwargs["rand_seed"].dtype == torch.int64 + assert kwargs["rand_seed"].device.type == "cuda" + tracker.assert_not_called() @pytest.mark.skipif( diff --git a/tests/v1/e2e/test_replayssm_decode.py b/tests/v1/e2e/test_replayssm_decode.py index a9036c984cb9..bd62f54172a4 100644 --- a/tests/v1/e2e/test_replayssm_decode.py +++ b/tests/v1/e2e/test_replayssm_decode.py @@ -4,6 +4,7 @@ import pytest +import vllm.envs as envs from vllm.v1.metrics.reader import Counter from ...models.utils import check_logprobs_close @@ -38,10 +39,17 @@ def _check_replayssm_parity( tensor_parallel_size=1, mamba_backend: str = "triton", name_1: str = "replayssm", + require_v2: bool = False, + monkeypatch: pytest.MonkeyPatch | None = None, ): # Compare logprobs, not greedy ids: ReplaySSM's fp arithmetic can flip a # near-tie. Baseline and ReplaySSM run at the same TP, so TP numerics are # common-mode and only ReplaySSM varies. + if require_v2: + assert monkeypatch is not None + monkeypatch.setenv("VLLM_USE_V2_MODEL_RUNNER", "1") + envs.disable_envs_cache() + common = dict( max_model_len=1024, trust_remote_code=True, @@ -51,10 +59,14 @@ def _check_replayssm_parity( mamba_backend=mamba_backend, ) with vllm_runner(model_name, **common) as llm: + if require_v2: + assert llm.llm.llm_engine.vllm_config.use_v2_model_runner baseline = llm.generate_greedy_logprobs(PROMPTS, max_tokens=32, num_logprobs=5) with vllm_runner( model_name, use_replayssm=True, replayssm_buffer_len=16, **common ) as llm: + if require_v2: + assert llm.llm.llm_engine.vllm_config.use_v2_model_runner replay = llm.generate_greedy_logprobs(PROMPTS, max_tokens=32, num_logprobs=5) check_logprobs_close( @@ -92,6 +104,24 @@ def test_replayssm_flashinfer_decode_matches_baseline(vllm_runner, model_name): ) +@pytest.mark.skipif( + not HAS_FLASHINFER_CHECKPOINTING_SSU, + reason="flashinfer.mamba.checkpointing_ssu not available", +) +@pytest.mark.parametrize("model_name", MODELS) +def test_replayssm_flashinfer_decode_matches_baseline_v2( + vllm_runner, model_name, monkeypatch +): + _check_replayssm_parity( + vllm_runner, + model_name, + mamba_backend="flashinfer", + name_1="replayssm_flashinfer_v2", + require_v2=True, + monkeypatch=monkeypatch, + ) + + @multi_gpu_test(num_gpus=2) @pytest.mark.skipif( not HAS_FLASHINFER_CHECKPOINTING_SSU, @@ -108,6 +138,26 @@ def test_replayssm_flashinfer_decode_matches_baseline_tp2(vllm_runner, model_nam ) +@multi_gpu_test(num_gpus=2) +@pytest.mark.skipif( + not HAS_FLASHINFER_CHECKPOINTING_SSU, + reason="flashinfer.mamba.checkpointing_ssu not available", +) +@pytest.mark.parametrize("model_name", [MAMBA2_MODEL]) +def test_replayssm_flashinfer_decode_matches_baseline_v2_tp2( + vllm_runner, model_name, monkeypatch +): + _check_replayssm_parity( + vllm_runner, + model_name, + tensor_parallel_size=2, + mamba_backend="flashinfer", + name_1="replayssm_flashinfer_v2_tp2", + require_v2=True, + monkeypatch=monkeypatch, + ) + + # Prefix spans several mamba blocks; prefix caching only reuses full blocks. _PC_SENTENCE = ( "In a detailed survey of state space models, the authors compared many " diff --git a/tests/v1/worker/test_utils.py b/tests/v1/worker/test_utils.py index a49c76fe24b0..39bacb4680d7 100644 --- a/tests/v1/worker/test_utils.py +++ b/tests/v1/worker/test_utils.py @@ -45,10 +45,11 @@ def test_bind_kv_cache_shares_replayssm_trackers_by_cache_group(): mixers = [_TestReplaySSMMixer() for _ in range(3)] layer_names = [f"layers.{i}.mixer" for i in range(3)] ctx = dict(zip(layer_names, mixers)) + # Reverse insertion order: updater must follow layer index, not dict order. kv_cache = { - layer_names[0]: _packed_replayssm_cache(4), - layer_names[1]: _packed_replayssm_cache(4), layer_names[2]: _packed_replayssm_cache(4), + layer_names[1]: _packed_replayssm_cache(4), + layer_names[0]: _packed_replayssm_cache(4), } kv_cache_groups = [ SimpleNamespace(layer_names=[layer_names[0], layer_names[2]]), @@ -79,6 +80,7 @@ def test_bind_kv_cache_shares_replayssm_trackers_by_cache_group(): assert mixers[0]._replayssm_ring_start.is_contiguous() assert torch.count_nonzero(mixers[0]._replayssm_ring_start) == 0 assert torch.count_nonzero(mixers[0]._replayssm_prev_num_accepted) == 0 + # Group {0, 2} shares trackers; layer 2 (not 0) updates after both run. assert [m._updates_replayssm_trackers for m in mixers] == [False, True, True] diff --git a/vllm/model_executor/layers/mamba/mamba_mixer2.py b/vllm/model_executor/layers/mamba/mamba_mixer2.py index d0b9dc59cc35..0de22ead8339 100644 --- a/vllm/model_executor/layers/mamba/mamba_mixer2.py +++ b/vllm/model_executor/layers/mamba/mamba_mixer2.py @@ -2,6 +2,8 @@ # SPDX-FileCopyrightText: Copyright contributors to the vLLM project +from collections.abc import Sequence + import torch from torch import nn @@ -17,6 +19,7 @@ from vllm.forward_context import ForwardContext, get_forward_context from vllm.logger import init_logger from vllm.model_executor.custom_op import CustomOp, PluggableLayer +from vllm.model_executor.layers.attention import Attention from vllm.model_executor.layers.linear import ( ColumnParallelLinear, MergedColumnParallelLinear, @@ -63,6 +66,7 @@ from vllm.v1.attention.backend import AttentionMetadata from vllm.v1.attention.backends.mamba2_attn import Mamba2AttentionMetadata from vllm.v1.attention.backends.registry import MambaAttentionBackendEnum +from vllm.v1.kv_cache_interface import KVCacheGroupSpec logger = init_logger(__name__) @@ -1198,7 +1202,9 @@ def mamba_type(self) -> MambaAttentionBackendEnum: def share_replayssm_ring_trackers( - mixer_groups: list[list[MambaMixer2]], + ordered_layer_names: list[str], + forward_context: dict[str, Attention], + kv_cache_groups: Sequence[KVCacheGroupSpec] | None = None, ) -> None: """Share ring cursors within each cache-slot index namespace. @@ -1208,28 +1214,50 @@ def share_replayssm_ring_trackers( The final local layer in each group advances its cursors after every layer in that group has consumed the previous values. """ - for mixers in mixer_groups: - if not mixers: - continue - first_state = mixers[0].kv_cache[1] - expected = (first_state.shape[0], first_state.device) - for mixer in mixers: - state = mixer.kv_cache[1] - actual = (state.shape[0], state.device) - if actual != expected: + replayssm_mixers: dict[str, MambaMixer2] = {} + for layer_name in ordered_layer_names: + layer = forward_context[layer_name] + if ( + isinstance(layer, MambaMixer2) + and layer.use_replayssm + and layer.mamba_config.backend == MambaBackendEnum.FLASHINFER + ): + replayssm_mixers[layer_name] = layer + + layer_to_group: dict[str, int] = {} + if kv_cache_groups: + for group_idx, group in enumerate(kv_cache_groups): + for layer_name in group.layer_names: + layer_to_group[layer_name] = group_idx + + groups_by_namespace: dict[int | str, list[str]] = {} + for layer_name in ordered_layer_names: + if layer_name not in replayssm_mixers: + continue + namespace = layer_to_group.get(layer_name, layer_name) + groups_by_namespace.setdefault(namespace, []).append(layer_name) + + for group_layer_names in groups_by_namespace.values(): + last_layer_name = group_layer_names[-1] + + first_mixer = replayssm_mixers[group_layer_names[0]] + first_state = first_mixer.kv_cache[1] + num_blocks, device = first_state.shape[0], first_state.device + for layer_name in group_layer_names: + state = replayssm_mixers[layer_name].kv_cache[1] + if (state.shape[0], state.device) != (num_blocks, device): raise ValueError( "ReplaySSM layers in one cache group must share cache capacity" ) - ring_start = torch.zeros(expected[0], dtype=torch.int32, device=expected[1]) + ring_start = torch.zeros(num_blocks, dtype=torch.int32, device=device) prev_num_accepted = torch.zeros_like(ring_start) - for mixer in mixers: + for layer_name in group_layer_names: + mixer = replayssm_mixers[layer_name] mixer._replayssm_ring_start = ring_start mixer._replayssm_prev_num_accepted = prev_num_accepted - mixer._updates_replayssm_trackers = False - - mixers[-1]._updates_replayssm_trackers = True + mixer._updates_replayssm_trackers = layer_name == last_layer_name def mamba_mixer2( diff --git a/vllm/model_executor/layers/mamba/ops/ssu_dispatch.py b/vllm/model_executor/layers/mamba/ops/ssu_dispatch.py index cc7575e9af4a..a66756142b60 100644 --- a/vllm/model_executor/layers/mamba/ops/ssu_dispatch.py +++ b/vllm/model_executor/layers/mamba/ops/ssu_dispatch.py @@ -380,9 +380,6 @@ def selective_state_update_replayssm_flashinfer( state_batch_indices: torch.Tensor | None = None, null_block_id: int = NULL_BLOCK_ID, scratch: tuple[torch.Tensor, torch.Tensor, torch.Tensor] | None = None, - algorithm: str = "auto", - d_split: int | None = None, - precompute_heads_per_cta: int = 0, update_trackers: bool = True, enable_stochastic_rounding: bool = False, stochastic_rounding_philox_rounds: int = 0, @@ -437,9 +434,6 @@ def selective_state_update_replayssm_flashinfer( cb_scaled=cb_scaled, cumAdt_vec=cumAdt_vec, cb_old=cb_old, - d_split=d_split, - precompute_heads_per_cta=precompute_heads_per_cta, - algorithm=algorithm, ) if update_trackers and indices is not None: update_replayssm_ring_trackers( diff --git a/vllm/model_executor/warmup/kernel_warmup.py b/vllm/model_executor/warmup/kernel_warmup.py index 92c17f1c93cd..42e3532de835 100644 --- a/vllm/model_executor/warmup/kernel_warmup.py +++ b/vllm/model_executor/warmup/kernel_warmup.py @@ -284,11 +284,10 @@ def _replayssm_autotune_kwargs( and config.mamba_config.backend == MambaBackendEnum.FLASHINFER ): return None - use_v2_model_runner = config.use_v2_model_runner v2_runner: Any = runner query_len = ( v2_runner.decode_query_len - if use_v2_model_runner + if config.use_v2_model_runner else runner.uniform_decode_query_len ) max_num_reqs = min( @@ -308,7 +307,7 @@ def _replayssm_autotune_kwargs( "num_tokens": max_num_reqs * query_len, "uniform_decode": True, } - if use_v2_model_runner: + if config.use_v2_model_runner: decode_kwargs["valid_dummy_state_slots"] = True else: decode_kwargs.update( @@ -355,10 +354,9 @@ def _temporary_replayssm_autotune_state( if tensor.numel(): reset_tensors.setdefault(tensor.data_ptr(), tensor) - use_v2_model_runner = runner.vllm_config.use_v2_model_runner v2_runner: Any = runner block_tables = saved_block_ids = None - if not use_v2_model_runner: + if not runner.vllm_config.use_v2_model_runner: block_tables = runner.input_batch.block_table.block_tables saved_block_ids = tuple( block_table.block_table.np[:max_num_reqs, 0].copy() @@ -395,7 +393,7 @@ def _temporary_replayssm_autotune_state( try: yield finally: - if use_v2_model_runner: + if runner.vllm_config.use_v2_model_runner: v2_runner.block_tables.get_dummy_block_tables(max_num_reqs) else: assert block_tables is not None and saved_block_ids is not None diff --git a/vllm/v1/worker/utils.py b/vllm/v1/worker/utils.py index 7707fcb520c7..7e2bef04ae10 100644 --- a/vllm/v1/worker/utils.py +++ b/vllm/v1/worker/utils.py @@ -11,9 +11,9 @@ import torch from vllm.config import CacheConfig, VllmConfig -from vllm.config.mamba import MambaBackendEnum from vllm.logger import init_logger from vllm.model_executor.layers.attention import Attention +from vllm.model_executor.layers.mamba.mamba_mixer2 import share_replayssm_ring_trackers from vllm.model_executor.models.interfaces import MultiModalEmbeddings from vllm.model_executor.models.utils import extract_layer_index from vllm.platforms import current_platform @@ -615,43 +615,7 @@ def bind_kv_cache( for layer_name, kv_cache in kv_caches.items(): forward_context[layer_name].bind_kv_cache(kv_cache) - from vllm.model_executor.layers.mamba.mamba_mixer2 import ( - MambaMixer2, - share_replayssm_ring_trackers, - ) - - replayssm_mixers: dict[str, MambaMixer2] = {} - for layer_name in ordered_layer_names: - layer = forward_context[layer_name] - if ( - isinstance(layer, MambaMixer2) - and layer.use_replayssm - and layer.mamba_config.backend == MambaBackendEnum.FLASHINFER - ): - replayssm_mixers[layer_name] = layer - if kv_cache_groups: - mixer_groups = [] - grouped_names = { - layer_name for group in kv_cache_groups for layer_name in group.layer_names - } - for group in kv_cache_groups: - group_names = set(group.layer_names) - mixer_groups.append( - [ - replayssm_mixers[layer_name] - for layer_name in ordered_layer_names - if layer_name in group_names and layer_name in replayssm_mixers - ] - ) - mixer_groups.extend( - [mixer] - for layer_name, mixer in replayssm_mixers.items() - if layer_name not in grouped_names - ) - else: - # Without cache groups, block-index namespaces cannot be proven equal. - mixer_groups = [[mixer] for mixer in replayssm_mixers.values()] - share_replayssm_ring_trackers(mixer_groups) + share_replayssm_ring_trackers(ordered_layer_names, forward_context, kv_cache_groups) def copy_kv_cache_blocks_inplace( From a53b2fe0ce661d8feb0e01ce5310bae997099ad7 Mon Sep 17 00:00:00 2001 From: Andrii Skliar Date: Thu, 27 Aug 2026 20:01:07 +0200 Subject: [PATCH 25/33] Refactor FlashInfer checkpointing SSU integration in tests Signed-off-by: Andrii Skliar --- tests/kernels/mamba/test_ssu_dispatch.py | 35 +++++------------------- tests/v1/e2e/test_replayssm_decode.py | 35 +----------------------- tests/v1/worker/test_utils.py | 29 +++----------------- 3 files changed, 12 insertions(+), 87 deletions(-) diff --git a/tests/kernels/mamba/test_ssu_dispatch.py b/tests/kernels/mamba/test_ssu_dispatch.py index a9f56a1e25d6..1a339a83fe5e 100644 --- a/tests/kernels/mamba/test_ssu_dispatch.py +++ b/tests/kernels/mamba/test_ssu_dispatch.py @@ -30,18 +30,9 @@ HAS_FLASHINFER = False try: - from flashinfer.mamba.checkpointing_ssu import ( - CheckpointingSSURunner, - allocate_checkpointing_ssu_scratch, - ) - from flashinfer.mamba.checkpointing_ssu import ( - checkpointing_ssu as checkpointing_ssu_kernel, - ) + from flashinfer.mamba.checkpointing_ssu import CheckpointingSSURunner - HAS_FLASHINFER_CHECKPOINTING_SSU = all( - callable(symbol) - for symbol in (CheckpointingSSURunner, allocate_checkpointing_ssu_scratch) - ) + HAS_FLASHINFER_CHECKPOINTING_SSU = callable(CheckpointingSSURunner) except ImportError: HAS_FLASHINFER_CHECKPOINTING_SSU = False @@ -256,7 +247,7 @@ def test_triton_basic_call(): @pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") -def test_replayssm_flashinfer_call_forwards_scratch_and_rounding(monkeypatch): +def test_replayssm_flashinfer_wrapper_forwards_vllm_owned_args(monkeypatch): import vllm.model_executor.layers.mamba.ops.ssu_dispatch as mod kernel = Mock(return_value=torch.empty(1, 1, 2, 4, device="cuda")) @@ -309,6 +300,7 @@ def test_replayssm_flashinfer_call_forwards_scratch_and_rounding(monkeypatch): kwargs = kernel.call_args.kwargs assert args[4] is ring_start assert args[5] is prev_num_accepted + assert kwargs["state_batch_indices"] is state_batch_indices assert kwargs["cb_scaled"] is scratch[0] assert kwargs["cumAdt_vec"] is scratch[1] assert kwargs["cb_old"] is scratch[2] @@ -316,25 +308,12 @@ def test_replayssm_flashinfer_call_forwards_scratch_and_rounding(monkeypatch): assert kwargs["rand_seed"].shape == (1,) assert kwargs["rand_seed"].dtype == torch.int64 assert kwargs["rand_seed"].device.type == "cuda" + assert "algorithm" not in kwargs + assert "d_split" not in kwargs + assert "precompute_heads_per_cta" not in kwargs tracker.assert_not_called() -@pytest.mark.skipif( - not HAS_FLASHINFER_CHECKPOINTING_SSU, - reason="compatible flashinfer checkpointing_ssu not available", -) -def test_replayssm_flashinfer_backend_init(): - import vllm.model_executor.layers.mamba.ops.ssu_dispatch as mod - - initialize_mamba_ssu_backend( - MambaConfig(backend=MambaBackendEnum.FLASHINFER), - _kv_cache_config_with_ssu(), - use_replayssm=True, - ) - assert isinstance(get_mamba_ssu_backend(), FlashInferSSUBackend) - assert mod._flashinfer_replayssm_kernel is checkpointing_ssu_kernel - - @pytest.mark.parametrize( ("backend", "expected_ring_len"), [ diff --git a/tests/v1/e2e/test_replayssm_decode.py b/tests/v1/e2e/test_replayssm_decode.py index bd62f54172a4..27a4973f8976 100644 --- a/tests/v1/e2e/test_replayssm_decode.py +++ b/tests/v1/e2e/test_replayssm_decode.py @@ -11,10 +11,7 @@ from ...utils import large_gpu_mark, multi_gpu_test try: - from flashinfer.mamba.checkpointing_ssu import ( - CheckpointingSSURunner, - allocate_checkpointing_ssu_scratch, # noqa: F401 - ) + from flashinfer.mamba.checkpointing_ssu import CheckpointingSSURunner HAS_FLASHINFER_CHECKPOINTING_SSU = callable(CheckpointingSSURunner) except ImportError: @@ -90,20 +87,6 @@ def test_replayssm_decode_matches_baseline_tp2(vllm_runner, model_name): _check_replayssm_parity(vllm_runner, model_name, tensor_parallel_size=2) -@pytest.mark.skipif( - not HAS_FLASHINFER_CHECKPOINTING_SSU, - reason="flashinfer.mamba.checkpointing_ssu not available", -) -@pytest.mark.parametrize("model_name", MODELS) -def test_replayssm_flashinfer_decode_matches_baseline(vllm_runner, model_name): - _check_replayssm_parity( - vllm_runner, - model_name, - mamba_backend="flashinfer", - name_1="replayssm_flashinfer", - ) - - @pytest.mark.skipif( not HAS_FLASHINFER_CHECKPOINTING_SSU, reason="flashinfer.mamba.checkpointing_ssu not available", @@ -122,22 +105,6 @@ def test_replayssm_flashinfer_decode_matches_baseline_v2( ) -@multi_gpu_test(num_gpus=2) -@pytest.mark.skipif( - not HAS_FLASHINFER_CHECKPOINTING_SSU, - reason="flashinfer.mamba.checkpointing_ssu not available", -) -@pytest.mark.parametrize("model_name", [MAMBA2_MODEL]) -def test_replayssm_flashinfer_decode_matches_baseline_tp2(vllm_runner, model_name): - _check_replayssm_parity( - vllm_runner, - model_name, - tensor_parallel_size=2, - mamba_backend="flashinfer", - name_1="replayssm_flashinfer_tp2", - ) - - @multi_gpu_test(num_gpus=2) @pytest.mark.skipif( not HAS_FLASHINFER_CHECKPOINTING_SSU, diff --git a/tests/v1/worker/test_utils.py b/tests/v1/worker/test_utils.py index 39bacb4680d7..a4474c8a5843 100644 --- a/tests/v1/worker/test_utils.py +++ b/tests/v1/worker/test_utils.py @@ -11,34 +11,23 @@ class _TestReplaySSMMixer(MambaMixer2): - _state_shapes = ((2,), (3,), (4,), (5,), (6,)) - _state_dtypes = ( - torch.float32, - torch.float32, - torch.float32, - torch.float32, - torch.float32, - ) - def __init__(self): torch.nn.Module.__init__(self) self.use_replayssm = True self.mamba_config = MambaConfig(backend=MambaBackendEnum.FLASHINFER) - self.cache_config = SimpleNamespace(mamba_cache_mode="none") - self.replayssm_buffer_len = 16 self._replayssm_ring_start = torch.empty(0, dtype=torch.int32) self._replayssm_prev_num_accepted = torch.empty(0, dtype=torch.int32) self._updates_replayssm_trackers = True def get_state_shape(self) -> tuple[tuple[int, ...], ...]: - return self._state_shapes + return ((2,), (3,), (4,), (5,), (6,)) def get_state_dtype(self) -> tuple[torch.dtype, ...]: - return self._state_dtypes + return (torch.float32,) * 5 -def _packed_replayssm_cache(num_blocks: int, fill_value: int = 0) -> torch.Tensor: - return torch.full((num_blocks, 1, 1, 80), fill_value, dtype=torch.int8) +def _packed_replayssm_cache(num_blocks: int) -> torch.Tensor: + return torch.full((num_blocks, 1, 1, 80), 0, dtype=torch.int8) def test_bind_kv_cache_shares_replayssm_trackers_by_cache_group(): @@ -70,16 +59,6 @@ def test_bind_kv_cache_shares_replayssm_trackers_by_cache_group(): mixers[1]._replayssm_ring_start.data_ptr() != mixers[0]._replayssm_ring_start.data_ptr() ) - assert ( - mixers[1]._replayssm_prev_num_accepted.data_ptr() - != mixers[0]._replayssm_prev_num_accepted.data_ptr() - ) - assert mixers[0]._replayssm_ring_start.shape == (4,) - assert mixers[0]._replayssm_prev_num_accepted.shape == (4,) - assert mixers[0]._replayssm_ring_start.dtype == torch.int32 - assert mixers[0]._replayssm_ring_start.is_contiguous() - assert torch.count_nonzero(mixers[0]._replayssm_ring_start) == 0 - assert torch.count_nonzero(mixers[0]._replayssm_prev_num_accepted) == 0 # Group {0, 2} shares trackers; layer 2 (not 0) updates after both run. assert [m._updates_replayssm_trackers for m in mixers] == [False, True, True] From d0246605b20caaae7b09233adb4585b8f303716c Mon Sep 17 00:00:00 2001 From: Andrii Skliar Date: Fri, 28 Aug 2026 17:33:34 +0200 Subject: [PATCH 26/33] Remove excessive FlashInfer ReplaySSM tests Signed-off-by: Andrii Skliar --- tests/kernels/mamba/test_ssu_dispatch.py | 87 --------------------- tests/model_executor/test_kernel_warmup.py | 63 +-------------- tests/test_config.py | 9 --- tests/v1/e2e/test_replayssm_decode.py | 37 ++++++--- tests/v1/worker/test_gpu_block_table.py | 23 ------ vllm/model_executor/warmup/kernel_warmup.py | 7 -- 6 files changed, 28 insertions(+), 198 deletions(-) diff --git a/tests/kernels/mamba/test_ssu_dispatch.py b/tests/kernels/mamba/test_ssu_dispatch.py index 1a339a83fe5e..5a6943857fe6 100644 --- a/tests/kernels/mamba/test_ssu_dispatch.py +++ b/tests/kernels/mamba/test_ssu_dispatch.py @@ -15,7 +15,6 @@ initialize_mamba_ssu_backend, reset_replayssm_ring_trackers, selective_state_update, - selective_state_update_replayssm_flashinfer, update_replayssm_ring_trackers, ) from vllm.utils.torch_utils import set_random_seed @@ -79,24 +78,6 @@ def test_flashinfer_replayssm_ring_tracker_lifecycle(): assert (ring_start[1].item(), prev_num_accepted[1].item()) == (0, 0) -@pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") -def test_flashinfer_replayssm_ring_tracker_ignores_invalid_slots(): - 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, 2, 1, 0], dtype=torch.int32, device="cuda") - - update_replayssm_ring_trackers( - ring_start, - prev_num_accepted, - state_batch_indices, - logical_window=16, - ring_buffer_len=17, - ) - - assert ring_start.tolist() == [0, 0] - assert prev_num_accepted.tolist() == [0, 1] - - def _kv_cache_config_with_ssu( mamba_type: MambaAttentionBackendEnum = MambaAttentionBackendEnum.MAMBA2, ) -> KVCacheConfig: @@ -246,74 +227,6 @@ def test_triton_basic_call(): assert not torch.isnan(out).any() -@pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") -def test_replayssm_flashinfer_wrapper_forwards_vllm_owned_args(monkeypatch): - import vllm.model_executor.layers.mamba.ops.ssu_dispatch as mod - - kernel = Mock(return_value=torch.empty(1, 1, 2, 4, device="cuda")) - tracker = Mock() - monkeypatch.setattr(mod, "_flashinfer_replayssm_kernel", kernel) - monkeypatch.setattr(mod, "update_replayssm_ring_trackers", tracker) - - batch, nheads, dim, dstate, ngroups, window = 1, 2, 4, 8, 1, 16 - state = torch.empty(1, nheads, dim, dstate, device="cuda") - x = torch.empty(batch, nheads, dim, device="cuda") - dt = torch.empty(batch, nheads, dim, device="cuda") - A = torch.empty(nheads, dim, dstate, device="cuda") - B = torch.empty(batch, ngroups, dstate, device="cuda") - C = torch.empty(batch, ngroups, dstate, device="cuda") - out = torch.empty_like(x) - x_cache = torch.empty(1, nheads, window, dim, device="cuda") - dt_cache = torch.empty(1, nheads, window, device="cuda") - B_cache = torch.empty(1, ngroups, window, dstate, device="cuda") - ring_start = torch.zeros(1, dtype=torch.int32, device="cuda") - prev_num_accepted = torch.zeros(1, dtype=torch.int32, device="cuda") - state_batch_indices = torch.zeros(1, dtype=torch.int32, device="cuda") - scratch = ( - torch.empty(1, device="cuda"), - torch.empty(1, device="cuda"), - torch.empty(1, device="cuda"), - ) - - selective_state_update_replayssm_flashinfer( - state, - x, - dt, - A, - B, - C, - out, - x_cache, - B_cache, - dt_cache, - ring_start, - prev_num_accepted, - logical_window=window, - state_batch_indices=state_batch_indices, - scratch=scratch, - enable_stochastic_rounding=True, - stochastic_rounding_philox_rounds=6, - update_trackers=False, - ) - - args = kernel.call_args.args - kwargs = kernel.call_args.kwargs - assert args[4] is ring_start - assert args[5] is prev_num_accepted - assert kwargs["state_batch_indices"] is state_batch_indices - assert kwargs["cb_scaled"] is scratch[0] - assert kwargs["cumAdt_vec"] is scratch[1] - assert kwargs["cb_old"] is scratch[2] - assert kwargs["philox_rounds"] == 6 - assert kwargs["rand_seed"].shape == (1,) - assert kwargs["rand_seed"].dtype == torch.int64 - assert kwargs["rand_seed"].device.type == "cuda" - assert "algorithm" not in kwargs - assert "d_split" not in kwargs - assert "precompute_heads_per_cta" not in kwargs - tracker.assert_not_called() - - @pytest.mark.parametrize( ("backend", "expected_ring_len"), [ diff --git a/tests/model_executor/test_kernel_warmup.py b/tests/model_executor/test_kernel_warmup.py index 6477118d2d93..912adcc7d9b2 100644 --- a/tests/model_executor/test_kernel_warmup.py +++ b/tests/model_executor/test_kernel_warmup.py @@ -28,22 +28,12 @@ def _replayssm_mixer() -> MambaMixer2: return mixer -@pytest.mark.parametrize( - ("backend", "use_replayssm", "use_v2_model_runner", "expected"), - [ - (MambaBackendEnum.FLASHINFER, True, False, True), - (MambaBackendEnum.FLASHINFER, True, True, True), - (MambaBackendEnum.TRITON, True, False, False), - (MambaBackendEnum.FLASHINFER, False, False, False), - ], -) -def test_replayssm_autotune_decode_kwargs( - backend, use_replayssm, use_v2_model_runner, expected -): +@pytest.mark.parametrize("use_v2_model_runner", [False, True]) +def test_replayssm_autotune_decode_kwargs(use_v2_model_runner): runner = SimpleNamespace( vllm_config=SimpleNamespace( - cache_config=SimpleNamespace(use_replayssm=use_replayssm), - mamba_config=SimpleNamespace(backend=backend), + cache_config=SimpleNamespace(use_replayssm=True), + mamba_config=SimpleNamespace(backend=MambaBackendEnum.FLASHINFER), use_v2_model_runner=use_v2_model_runner, ), uniform_decode_query_len=6, @@ -61,9 +51,6 @@ def test_replayssm_autotune_decode_kwargs( result = warmup._replayssm_autotune_kwargs(runner, prefill_kwargs) - if not expected: - assert result is None - return expected_kwargs = { **prefill_kwargs, "num_tokens": 96, @@ -100,22 +87,6 @@ def test_replayssm_autotune_decode_kwargs_clamps_to_state_capacity(): assert result[1]["num_tokens"] == 4 -def test_replayssm_autotune_decode_kwargs_skips_without_state_slot(): - runner = SimpleNamespace( - vllm_config=SimpleNamespace( - cache_config=SimpleNamespace(use_replayssm=True), - mamba_config=SimpleNamespace(backend=MambaBackendEnum.FLASHINFER), - use_v2_model_runner=False, - ), - uniform_decode_query_len=1, - max_num_tokens=128, - scheduler_config=SimpleNamespace(max_num_seqs=64), - kv_cache_config=SimpleNamespace(num_blocks=1), - ) - - assert warmup._replayssm_autotune_kwargs(runner, {}) is None - - def test_replayssm_autotune_slots_restore_state_and_trackers(): mixer = _replayssm_mixer() @@ -153,29 +124,3 @@ def test_replayssm_autotune_slots_restore_state_and_trackers(): assert torch.count_nonzero(tensor[1:3]) == 0 assert torch.all(tensor[0] == 3) assert torch.all(tensor[3] == 3) - - -def test_replayssm_autotune_slots_reset_v2_dummy_tables_and_state(): - mixer = _replayssm_mixer() - block_tables = SimpleNamespace(get_dummy_block_tables=Mock()) - runner = SimpleNamespace( - vllm_config=SimpleNamespace(use_v2_model_runner=True), - block_tables=block_tables, - get_model=lambda: SimpleNamespace(modules=lambda: (mixer,)), - ) - - with warmup._temporary_replayssm_autotune_state(runner, 2): - for tensor in ( - *mixer.kv_cache, - mixer._replayssm_ring_start, - mixer._replayssm_prev_num_accepted, - ): - tensor[1:3].fill_(9) - - block_tables.get_dummy_block_tables.assert_called_once_with(2) - for tensor in ( - *mixer.kv_cache, - mixer._replayssm_ring_start, - mixer._replayssm_prev_num_accepted, - ): - assert torch.count_nonzero(tensor[1:3]) == 0 diff --git a/tests/test_config.py b/tests/test_config.py index 9eb8bddef457..1e8acaa87728 100644 --- a/tests/test_config.py +++ b/tests/test_config.py @@ -157,15 +157,6 @@ def test_v2_replayssm_requires_flashinfer(): VllmConfig.validate_mamba_cached_kernel(config) -def test_v2_flashinfer_replayssm_is_supported(): - config = _replayssm_config( - backend=MambaBackendEnum.FLASHINFER, - use_v2_model_runner=True, - ) - - assert VllmConfig.validate_mamba_cached_kernel(config) is config - - def test_rocm_keeps_compiled_deepseek_defaults(monkeypatch): """ROCm keeps DeepSeek V3.2 and V4 on their compiled MRV1 paths.""" from vllm.config.vllm import ( diff --git a/tests/v1/e2e/test_replayssm_decode.py b/tests/v1/e2e/test_replayssm_decode.py index 27a4973f8976..38fd1eec65d2 100644 --- a/tests/v1/e2e/test_replayssm_decode.py +++ b/tests/v1/e2e/test_replayssm_decode.py @@ -105,23 +105,34 @@ def test_replayssm_flashinfer_decode_matches_baseline_v2( ) -@multi_gpu_test(num_gpus=2) @pytest.mark.skipif( not HAS_FLASHINFER_CHECKPOINTING_SSU, reason="flashinfer.mamba.checkpointing_ssu not available", ) -@pytest.mark.parametrize("model_name", [MAMBA2_MODEL]) -def test_replayssm_flashinfer_decode_matches_baseline_v2_tp2( - vllm_runner, model_name, monkeypatch -): - _check_replayssm_parity( - vllm_runner, - model_name, - tensor_parallel_size=2, - mamba_backend="flashinfer", - name_1="replayssm_flashinfer_v2_tp2", - require_v2=True, - monkeypatch=monkeypatch, +@pytest.mark.parametrize("model_name", MODELS) +def test_replayssm_flashinfer_matches_triton_replayssm(vllm_runner, model_name): + # Both backends implement ReplaySSM; compare them directly on V1 because + # Triton ReplaySSM is not supported on Model Runner V2. + common = dict( + max_model_len=1024, + trust_remote_code=True, + enable_prefix_caching=False, + mamba_cache_mode="none", + use_replayssm=True, + replayssm_buffer_len=16, + ) + with vllm_runner(model_name, mamba_backend="triton", **common) as llm: + triton = llm.generate_greedy_logprobs(PROMPTS, max_tokens=32, num_logprobs=5) + with vllm_runner(model_name, mamba_backend="flashinfer", **common) as llm: + flashinfer = llm.generate_greedy_logprobs( + PROMPTS, max_tokens=32, num_logprobs=5 + ) + + check_logprobs_close( + outputs_0_lst=triton, + outputs_1_lst=flashinfer, + name_0="replayssm_triton", + name_1="replayssm_flashinfer", ) diff --git a/tests/v1/worker/test_gpu_block_table.py b/tests/v1/worker/test_gpu_block_table.py index 76cf61171825..ee44ff24d581 100644 --- a/tests/v1/worker/test_gpu_block_table.py +++ b/tests/v1/worker/test_gpu_block_table.py @@ -1,14 +1,11 @@ # SPDX-License-Identifier: Apache-2.0 # SPDX-FileCopyrightText: Copyright contributors to the vLLM project -from types import SimpleNamespace - import pytest import torch from vllm.platforms import current_platform from vllm.v1.worker.gpu.block_table import BlockTables -from vllm.v1.worker.gpu.model_runner import GPUModelRunner pytestmark = pytest.mark.skipif( not current_platform.is_cuda(), @@ -240,23 +237,3 @@ def test_get_dummy_block_tables_returns_zeroed_rows(): assert (dummy[0] == 0).all() # CUDA graph invariant: same persistent tensor, not a fresh allocation. assert dummy[0].data_ptr() == block_tables.input_block_tables[0].data_ptr() - - -def test_prepare_dummy_attn_can_assign_valid_state_slots(): - runner = object.__new__(GPUModelRunner) - runner.device = torch.device("cuda") - runner.pcp_manager = None - runner.block_tables = BlockTables( - block_sizes=[16], - max_num_reqs=4, - max_num_batched_tokens=64, - max_num_blocks_per_group=[8], - device=runner.device, - kernel_block_sizes=[16], - ) - input_batch = SimpleNamespace(num_reqs=3, num_tokens=3) - - block_tables, _ = runner.prepare_dummy_attn(input_batch, valid_state_slots=True) - - assert block_tables[0][:, 0].tolist() == [1, 2, 3] - assert torch.count_nonzero(block_tables[0][:, 1:]) == 0 diff --git a/vllm/model_executor/warmup/kernel_warmup.py b/vllm/model_executor/warmup/kernel_warmup.py index 42e3532de835..52c8a2eb3e67 100644 --- a/vllm/model_executor/warmup/kernel_warmup.py +++ b/vllm/model_executor/warmup/kernel_warmup.py @@ -295,13 +295,6 @@ def _replayssm_autotune_kwargs( runner.max_num_tokens // query_len, runner.kv_cache_config.num_blocks - 1, ) - if max_num_reqs <= 0: - logger.warning_once( - "Skipping FlashInfer ReplaySSM autotuning because no non-padding " - "state slot is available." - ) - return None - decode_kwargs = { **max_token_prefill_kwargs, "num_tokens": max_num_reqs * query_len, From 1523e793eafa1082dfad8db1cd395a292119d7e8 Mon Sep 17 00:00:00 2001 From: Andrii Skliar Date: Mon, 31 Aug 2026 11:39:23 +0200 Subject: [PATCH 27/33] remove `raise` if replayssm autotuning is not available Signed-off-by: Andrii Skliar --- tests/model_executor/test_kernel_warmup.py | 32 ++++++++++++++++++- .../layers/mamba/ops/ssu_dispatch.py | 18 +++++++---- vllm/model_executor/warmup/kernel_warmup.py | 10 ++++++ 3 files changed, 53 insertions(+), 7 deletions(-) diff --git a/tests/model_executor/test_kernel_warmup.py b/tests/model_executor/test_kernel_warmup.py index 912adcc7d9b2..54ffabb22066 100644 --- a/tests/model_executor/test_kernel_warmup.py +++ b/tests/model_executor/test_kernel_warmup.py @@ -2,7 +2,7 @@ # SPDX-FileCopyrightText: Copyright contributors to the vLLM project from types import SimpleNamespace -from unittest.mock import Mock +from unittest.mock import Mock, patch import numpy as np import pytest @@ -13,6 +13,16 @@ from vllm.model_executor.warmup import kernel_warmup as warmup +@pytest.fixture(autouse=True) +def _replayssm_autotune_supported(): + with patch.object( + warmup, + "flashinfer_replayssm_autotune_supported", + return_value=True, + ): + yield + + def _replayssm_mixer() -> MambaMixer2: mixer = MambaMixer2.__new__(MambaMixer2) torch.nn.Module.__init__(mixer) @@ -87,6 +97,26 @@ def test_replayssm_autotune_decode_kwargs_clamps_to_state_capacity(): assert result[1]["num_tokens"] == 4 +def test_replayssm_autotune_decode_kwargs_skips_without_runner(): + runner = SimpleNamespace( + vllm_config=SimpleNamespace( + cache_config=SimpleNamespace(use_replayssm=True), + mamba_config=SimpleNamespace(backend=MambaBackendEnum.FLASHINFER), + use_v2_model_runner=False, + ), + uniform_decode_query_len=1, + max_num_tokens=128, + scheduler_config=SimpleNamespace(max_num_seqs=64), + kv_cache_config=SimpleNamespace(num_blocks=5), + ) + with patch.object( + warmup, + "flashinfer_replayssm_autotune_supported", + return_value=False, + ): + assert warmup._replayssm_autotune_kwargs(runner, {}) is None + + def test_replayssm_autotune_slots_restore_state_and_trackers(): mixer = _replayssm_mixer() diff --git a/vllm/model_executor/layers/mamba/ops/ssu_dispatch.py b/vllm/model_executor/layers/mamba/ops/ssu_dispatch.py index a66756142b60..8c21a53e9518 100644 --- a/vllm/model_executor/layers/mamba/ops/ssu_dispatch.py +++ b/vllm/model_executor/layers/mamba/ops/ssu_dispatch.py @@ -11,6 +11,7 @@ from abc import ABC, abstractmethod from collections.abc import Callable +from functools import cache import torch @@ -360,6 +361,16 @@ def __call__( _flashinfer_replayssm_kernel: Callable[..., torch.Tensor] | None = None +@cache +def flashinfer_replayssm_autotune_supported() -> bool: + """Return True when FlashInfer exposes ReplaySSM autotuning.""" + try: + from flashinfer.mamba.checkpointing_ssu import CheckpointingSSURunner + except ImportError: + return False + return callable(CheckpointingSSURunner) + + def selective_state_update_replayssm_flashinfer( state: torch.Tensor, x: torch.Tensor, @@ -494,16 +505,11 @@ def initialize_mamba_ssu_backend( _flashinfer_replayssm_kernel = None if use_replayssm and backend == MambaBackendEnum.FLASHINFER: try: - from flashinfer.mamba.checkpointing_ssu import ( - CheckpointingSSURunner, - checkpointing_ssu, - ) + from flashinfer.mamba.checkpointing_ssu import checkpointing_ssu except ImportError as e: raise ImportError( "FlashInfer ReplaySSM requires a compatible flashinfer-python package" ) from e - if not callable(CheckpointingSSURunner): - raise ImportError("FlashInfer ReplaySSM requires native autotuning support") _flashinfer_replayssm_kernel = checkpointing_ssu if use_replayssm: logger.info("Using %s ReplaySSM backend.", backend.value) diff --git a/vllm/model_executor/warmup/kernel_warmup.py b/vllm/model_executor/warmup/kernel_warmup.py index 52c8a2eb3e67..2ce806d7dfa0 100644 --- a/vllm/model_executor/warmup/kernel_warmup.py +++ b/vllm/model_executor/warmup/kernel_warmup.py @@ -17,6 +17,9 @@ import vllm.envs as envs from vllm.config.mamba import MambaBackendEnum from vllm.logger import init_logger +from vllm.model_executor.layers.mamba.ops.ssu_dispatch import ( + flashinfer_replayssm_autotune_supported, +) from vllm.model_executor.warmup.b12x_warmup import b12x_warmup from vllm.model_executor.warmup.cutedsl_warmup import cutedsl_warmup from vllm.model_executor.warmup.deep_gemm_warmup import deep_gemm_warmup @@ -284,6 +287,13 @@ def _replayssm_autotune_kwargs( and config.mamba_config.backend == MambaBackendEnum.FLASHINFER ): return None + if not flashinfer_replayssm_autotune_supported(): + logger.info_once( + "Skipping FlashInfer ReplaySSM autotuning because " + "flashinfer.mamba.checkpointing_ssu.CheckpointingSSURunner " + "is unavailable." + ) + return None v2_runner: Any = runner query_len = ( v2_runner.decode_query_len From e3b819fca6f63c2ba688be689c203f5d8f9e47f8 Mon Sep 17 00:00:00 2001 From: Andrii Skliar Date: Mon, 31 Aug 2026 20:56:27 +0200 Subject: [PATCH 28/33] clean up tests; remove unnecessary fixtures Signed-off-by: Andrii Skliar --- tests/kernels/mamba/test_ssu_dispatch.py | 58 +++--- tests/model_executor/test_kernel_warmup.py | 181 +++++++++--------- tests/test_config.py | 28 --- tests/v1/e2e/test_replayssm_decode.py | 43 +---- .../layers/mamba/ops/ssu_dispatch.py | 6 +- 5 files changed, 122 insertions(+), 194 deletions(-) diff --git a/tests/kernels/mamba/test_ssu_dispatch.py b/tests/kernels/mamba/test_ssu_dispatch.py index 5a6943857fe6..6226340a9936 100644 --- a/tests/kernels/mamba/test_ssu_dispatch.py +++ b/tests/kernels/mamba/test_ssu_dispatch.py @@ -19,7 +19,11 @@ ) from vllm.utils.torch_utils import set_random_seed from vllm.v1.attention.backends.registry import MambaAttentionBackendEnum -from vllm.v1.kv_cache_interface import KVCacheConfig, KVCacheGroupSpec, MambaSpec +from vllm.v1.kv_cache_interface import ( + KVCacheConfig, + KVCacheGroupSpec, + MambaSpec, +) try: import flashinfer.mamba # noqa: F401 @@ -28,13 +32,6 @@ except ImportError: HAS_FLASHINFER = False -try: - from flashinfer.mamba.checkpointing_ssu import CheckpointingSSURunner - - HAS_FLASHINFER_CHECKPOINTING_SSU = callable(CheckpointingSSURunner) -except ImportError: - HAS_FLASHINFER_CHECKPOINTING_SSU = False - @pytest.fixture(autouse=True) def restore_backend_state(): @@ -47,7 +44,6 @@ def restore_backend_state(): mod._flashinfer_replayssm_kernel = old_replayssm_kernel -@pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") 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") @@ -105,7 +101,8 @@ def test_explicit_triton_backend(): initialize_mamba_ssu_backend( MambaConfig(backend=MambaBackendEnum.TRITON), _kv_cache_config_with_ssu() ) - assert isinstance(get_mamba_ssu_backend(), TritonSSUBackend) + backend = get_mamba_ssu_backend() + assert isinstance(backend, TritonSSUBackend) @pytest.mark.skipif(not HAS_FLASHINFER, reason="flashinfer not installed") @@ -144,9 +141,18 @@ def test_flashinfer_forwards_ssu_algorithm( ssu_algorithm=algorithm, ) ) - tensor = torch.empty(1) - backend(*(tensor,) * 8) + tensor = torch.empty(1) + backend( + tensor, + tensor, + tensor, + tensor, + tensor, + tensor, + tensor, + tensor, + ) assert kernel.call_args.kwargs["algorithm"] == expected @@ -154,13 +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 - try: - with pytest.raises(RuntimeError, match="not been initialized"): - get_mamba_ssu_backend() - finally: - mod._mamba_ssu_backend = old + with pytest.raises(RuntimeError, match="not been initialized"): + get_mamba_ssu_backend() @pytest.mark.parametrize( @@ -193,24 +196,25 @@ def test_flashinfer_import_error(): FlashInferSSUBackend(MambaConfig()) -@pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") def test_triton_basic_call(): set_random_seed(0) initialize_mamba_ssu_backend( MambaConfig(backend=MambaBackendEnum.TRITON), _kv_cache_config_with_ssu() ) + device = "cuda" batch_size = 2 dim = 64 dstate = 16 - state = torch.randn(batch_size, dim, dstate, device="cuda") - x = torch.randn(batch_size, dim, device="cuda") + + state = torch.randn(batch_size, dim, dstate, device=device) + x = torch.randn(batch_size, dim, device=device) out = torch.empty_like(x) - dt = torch.randn(batch_size, dim, device="cuda") - dt_bias = torch.rand(dim, device="cuda") - 4.0 - A = -torch.rand(dim, dstate, device="cuda") - B = torch.randn(batch_size, dstate, device="cuda") - C = torch.randn(batch_size, dstate, device="cuda") - D = torch.randn(dim, device="cuda") + dt = torch.randn(batch_size, dim, device=device) + dt_bias = torch.rand(dim, device=device) - 4.0 + A = -torch.rand(dim, dstate, device=device) + B = torch.randn(batch_size, dstate, device=device) + C = torch.randn(batch_size, dstate, device=device) + D = torch.randn(dim, device=device) selective_state_update( state, diff --git a/tests/model_executor/test_kernel_warmup.py b/tests/model_executor/test_kernel_warmup.py index 54ffabb22066..4b5527fd0f09 100644 --- a/tests/model_executor/test_kernel_warmup.py +++ b/tests/model_executor/test_kernel_warmup.py @@ -12,113 +12,112 @@ from vllm.model_executor.layers.mamba.mamba_mixer2 import MambaMixer2 from vllm.model_executor.warmup import kernel_warmup as warmup - -@pytest.fixture(autouse=True) -def _replayssm_autotune_supported(): - with patch.object( - warmup, - "flashinfer_replayssm_autotune_supported", - return_value=True, - ): - yield - - -def _replayssm_mixer() -> MambaMixer2: - 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) - return mixer - - -@pytest.mark.parametrize("use_v2_model_runner", [False, True]) -def test_replayssm_autotune_decode_kwargs(use_v2_model_runner): - runner = SimpleNamespace( +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=True), - mamba_config=SimpleNamespace(backend=MambaBackendEnum.FLASHINFER), + cache_config=SimpleNamespace(use_replayssm=use_replayssm), + mamba_config=SimpleNamespace(backend=backend), use_v2_model_runner=use_v2_model_runner, ), - uniform_decode_query_len=6, - decode_query_len=6, - max_num_tokens=100, - scheduler_config=SimpleNamespace(max_num_seqs=32), - kv_cache_config=SimpleNamespace(num_blocks=17), + 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), ) - prefill_kwargs = { - "num_tokens": 128, - "skip_eplb": True, - "is_profile": True, - "randomize_inputs": True, - } - result = warmup._replayssm_autotune_kwargs(runner, prefill_kwargs) + +@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), PREFILL_KWARGS + ) expected_kwargs = { - **prefill_kwargs, - "num_tokens": 96, + **PREFILL_KWARGS, + "num_tokens": expected_num_reqs * query_len, "uniform_decode": True, } - if use_v2_model_runner: + 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=7, + profile_seq_lens=query_len + 1, ) - assert result == (16, expected_kwargs) - - -def test_replayssm_autotune_decode_kwargs_clamps_to_state_capacity(): - runner = SimpleNamespace( - vllm_config=SimpleNamespace( - cache_config=SimpleNamespace(use_replayssm=True), - mamba_config=SimpleNamespace(backend=MambaBackendEnum.FLASHINFER), - use_v2_model_runner=False, - ), - uniform_decode_query_len=1, - max_num_tokens=128, - scheduler_config=SimpleNamespace(max_num_seqs=64), - kv_cache_config=SimpleNamespace(num_blocks=5), - ) - - result = warmup._replayssm_autotune_kwargs(runner, {}) - - assert result is not None - assert result[0] == 4 - assert result[1]["num_tokens"] == 4 - - -def test_replayssm_autotune_decode_kwargs_skips_without_runner(): - runner = SimpleNamespace( - vllm_config=SimpleNamespace( - cache_config=SimpleNamespace(use_replayssm=True), - mamba_config=SimpleNamespace(backend=MambaBackendEnum.FLASHINFER), - use_v2_model_runner=False, - ), - uniform_decode_query_len=1, - max_num_tokens=128, - scheduler_config=SimpleNamespace(max_num_seqs=64), - kv_cache_config=SimpleNamespace(num_blocks=5), - ) + 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=False, + return_value=flashinfer_supported, ): - assert warmup._replayssm_autotune_kwargs(runner, {}) is None + result = warmup._replayssm_autotune_kwargs( + _autotune_runner(**runner_kwargs), PREFILL_KWARGS + ) + assert result is None def test_replayssm_autotune_slots_restore_state_and_trackers(): - mixer = _replayssm_mixer() + 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() @@ -134,11 +133,7 @@ def test_replayssm_autotune_slots_restore_state_and_trackers(): with warmup._temporary_replayssm_autotune_state(runner, 2): assert block_ids[:2, 0].tolist() == [1, 2] - for tensor in ( - *mixer.kv_cache, - mixer._replayssm_ring_start, - mixer._replayssm_prev_num_accepted, - ): + for tensor in tracked: tensor[1:3].fill_(9) assert np.array_equal(block_ids, original_block_ids) @@ -146,11 +141,7 @@ def test_replayssm_autotune_slots_restore_state_and_trackers(): ((2,), {}), ((2,), {}), ] - for tensor in ( - *mixer.kv_cache, - mixer._replayssm_ring_start, - mixer._replayssm_prev_num_accepted, - ): + for tensor in tracked: assert torch.count_nonzero(tensor[1:3]) == 0 assert torch.all(tensor[0] == 3) assert torch.all(tensor[3] == 3) diff --git a/tests/test_config.py b/tests/test_config.py index e8654170c827..4c27a6fe699c 100644 --- a/tests/test_config.py +++ b/tests/test_config.py @@ -158,34 +158,6 @@ def test_v2_model_runner_env_tri_state(monkeypatch, env_value, expected): assert envs.VLLM_USE_V2_MODEL_RUNNER is expected -def _replayssm_config( - *, - backend: MambaBackendEnum, - use_v2_model_runner: bool = False, -) -> SimpleNamespace: - return SimpleNamespace( - cache_config=SimpleNamespace( - use_replayssm=True, - mamba_cache_mode="none", - ), - model_config=None, - num_speculative_tokens=0, - mamba_config=SimpleNamespace(backend=backend), - use_v2_model_runner=use_v2_model_runner, - kv_transfer_config=None, - ) - - -def test_v2_replayssm_requires_flashinfer(): - config = _replayssm_config( - backend=MambaBackendEnum.TRITON, - use_v2_model_runner=True, - ) - - with pytest.raises(ValueError, match="requires Model Runner V1"): - VllmConfig.validate_mamba_cached_kernel(config) - - def test_rocm_keeps_compiled_deepseek_defaults(monkeypatch): """ROCm keeps DeepSeek V3.2 and V4 on their compiled MRV1 paths.""" from vllm.config.vllm import ( diff --git a/tests/v1/e2e/test_replayssm_decode.py b/tests/v1/e2e/test_replayssm_decode.py index 38fd1eec65d2..2e66a4668924 100644 --- a/tests/v1/e2e/test_replayssm_decode.py +++ b/tests/v1/e2e/test_replayssm_decode.py @@ -10,13 +10,6 @@ from ...models.utils import check_logprobs_close from ...utils import large_gpu_mark, multi_gpu_test -try: - from flashinfer.mamba.checkpointing_ssu import CheckpointingSSURunner - - HAS_FLASHINFER_CHECKPOINTING_SSU = callable(CheckpointingSSURunner) -except ImportError: - HAS_FLASHINFER_CHECKPOINTING_SSU = False - # Mamba2 (Nemotron-3) hybrid. MAMBA2_MODEL = "nvidia/NVIDIA-Nemotron-3-Nano-4B-BF16" MODELS = [ @@ -87,14 +80,11 @@ def test_replayssm_decode_matches_baseline_tp2(vllm_runner, model_name): _check_replayssm_parity(vllm_runner, model_name, tensor_parallel_size=2) -@pytest.mark.skipif( - not HAS_FLASHINFER_CHECKPOINTING_SSU, - reason="flashinfer.mamba.checkpointing_ssu not available", -) @pytest.mark.parametrize("model_name", MODELS) def test_replayssm_flashinfer_decode_matches_baseline_v2( vllm_runner, model_name, monkeypatch ): + pytest.importorskip("flashinfer.mamba.checkpointing_ssu") _check_replayssm_parity( vllm_runner, model_name, @@ -105,37 +95,6 @@ def test_replayssm_flashinfer_decode_matches_baseline_v2( ) -@pytest.mark.skipif( - not HAS_FLASHINFER_CHECKPOINTING_SSU, - reason="flashinfer.mamba.checkpointing_ssu not available", -) -@pytest.mark.parametrize("model_name", MODELS) -def test_replayssm_flashinfer_matches_triton_replayssm(vllm_runner, model_name): - # Both backends implement ReplaySSM; compare them directly on V1 because - # Triton ReplaySSM is not supported on Model Runner V2. - common = dict( - max_model_len=1024, - trust_remote_code=True, - enable_prefix_caching=False, - mamba_cache_mode="none", - use_replayssm=True, - replayssm_buffer_len=16, - ) - with vllm_runner(model_name, mamba_backend="triton", **common) as llm: - triton = llm.generate_greedy_logprobs(PROMPTS, max_tokens=32, num_logprobs=5) - with vllm_runner(model_name, mamba_backend="flashinfer", **common) as llm: - flashinfer = llm.generate_greedy_logprobs( - PROMPTS, max_tokens=32, num_logprobs=5 - ) - - check_logprobs_close( - outputs_0_lst=triton, - outputs_1_lst=flashinfer, - name_0="replayssm_triton", - name_1="replayssm_flashinfer", - ) - - # Prefix spans several mamba blocks; prefix caching only reuses full blocks. _PC_SENTENCE = ( "In a detailed survey of state space models, the authors compared many " diff --git a/vllm/model_executor/layers/mamba/ops/ssu_dispatch.py b/vllm/model_executor/layers/mamba/ops/ssu_dispatch.py index 8c21a53e9518..62e4307b37b7 100644 --- a/vllm/model_executor/layers/mamba/ops/ssu_dispatch.py +++ b/vllm/model_executor/layers/mamba/ops/ssu_dispatch.py @@ -365,10 +365,12 @@ def __call__( def flashinfer_replayssm_autotune_supported() -> bool: """Return True when FlashInfer exposes ReplaySSM autotuning.""" try: - from flashinfer.mamba.checkpointing_ssu import CheckpointingSSURunner + from flashinfer.mamba.checkpointing_ssu import ( # noqa: F401 + CheckpointingSSURunner, + ) except ImportError: return False - return callable(CheckpointingSSURunner) + return True def selective_state_update_replayssm_flashinfer( From d45545506b5e24d86cb09b647ba85766883ac7ed Mon Sep 17 00:00:00 2001 From: Andrii Skliar Date: Tue, 1 Sep 2026 00:13:22 +0200 Subject: [PATCH 29/33] move replayssm warmup to a separate file Signed-off-by: Andrii Skliar --- ...nel_warmup.py => test_replayssm_warmup.py} | 10 +- vllm/model_executor/warmup/kernel_warmup.py | 157 +---------------- .../model_executor/warmup/replayssm_warmup.py | 159 ++++++++++++++++++ 3 files changed, 167 insertions(+), 159 deletions(-) rename tests/model_executor/{test_kernel_warmup.py => test_replayssm_warmup.py} (93%) create mode 100644 vllm/model_executor/warmup/replayssm_warmup.py diff --git a/tests/model_executor/test_kernel_warmup.py b/tests/model_executor/test_replayssm_warmup.py similarity index 93% rename from tests/model_executor/test_kernel_warmup.py rename to tests/model_executor/test_replayssm_warmup.py index 4b5527fd0f09..bb59318b7cd6 100644 --- a/tests/model_executor/test_kernel_warmup.py +++ b/tests/model_executor/test_replayssm_warmup.py @@ -10,7 +10,7 @@ from vllm.config.mamba import MambaBackendEnum from vllm.model_executor.layers.mamba.mamba_mixer2 import MambaMixer2 -from vllm.model_executor.warmup import kernel_warmup as warmup +from vllm.model_executor.warmup import replayssm_warmup as warmup PREFILL_KWARGS = { "num_tokens": 128, @@ -60,9 +60,7 @@ def test_replayssm_autotune_decode_kwargs(runner_kwargs, expected_num_reqs): with patch.object( warmup, "flashinfer_replayssm_autotune_supported", return_value=True ): - result = warmup._replayssm_autotune_kwargs( - _autotune_runner(**runner_kwargs), PREFILL_KWARGS - ) + result = warmup._replayssm_autotune_kwargs(_autotune_runner(**runner_kwargs)) expected_kwargs = { **PREFILL_KWARGS, @@ -95,9 +93,7 @@ def test_replayssm_autotune_kwargs_skipped(runner_kwargs, flashinfer_supported): "flashinfer_replayssm_autotune_supported", return_value=flashinfer_supported, ): - result = warmup._replayssm_autotune_kwargs( - _autotune_runner(**runner_kwargs), PREFILL_KWARGS - ) + result = warmup._replayssm_autotune_kwargs(_autotune_runner(**runner_kwargs)) assert result is None diff --git a/vllm/model_executor/warmup/kernel_warmup.py b/vllm/model_executor/warmup/kernel_warmup.py index 2ce806d7dfa0..eb1b30273428 100644 --- a/vllm/model_executor/warmup/kernel_warmup.py +++ b/vllm/model_executor/warmup/kernel_warmup.py @@ -8,18 +8,12 @@ import sys import time -from collections.abc import Iterator -from contextlib import contextmanager -from typing import TYPE_CHECKING, Any +from typing import TYPE_CHECKING import torch import vllm.envs as envs -from vllm.config.mamba import MambaBackendEnum from vllm.logger import init_logger -from vllm.model_executor.layers.mamba.ops.ssu_dispatch import ( - flashinfer_replayssm_autotune_supported, -) from vllm.model_executor.warmup.b12x_warmup import b12x_warmup from vllm.model_executor.warmup.cutedsl_warmup import cutedsl_warmup from vllm.model_executor.warmup.deep_gemm_warmup import deep_gemm_warmup @@ -41,6 +35,9 @@ kimi_k3_triton_warmup, ) from vllm.model_executor.warmup.qwen_triton_warmup import qwen_triton_warmup +from vllm.model_executor.warmup.replayssm_warmup import ( + replayssm_autotune_warmup, +) from vllm.model_executor.warmup.sparse_mla_triton_warmup import ( sparse_mla_triton_warmup, ) @@ -278,135 +275,6 @@ def _run_flashinfer_autotune_dummy_runs(runner: "GPUModelRunner") -> None: ) -def _replayssm_autotune_kwargs( - runner: "GPUModelRunner", max_token_prefill_kwargs: dict[str, Any] -) -> tuple[int, dict[str, Any]] | None: - config = runner.vllm_config - if not ( - config.cache_config.use_replayssm - and config.mamba_config.backend == MambaBackendEnum.FLASHINFER - ): - return None - if not flashinfer_replayssm_autotune_supported(): - logger.info_once( - "Skipping FlashInfer ReplaySSM autotuning because " - "flashinfer.mamba.checkpointing_ssu.CheckpointingSSURunner " - "is unavailable." - ) - return None - v2_runner: Any = runner - query_len = ( - v2_runner.decode_query_len - if config.use_v2_model_runner - else runner.uniform_decode_query_len - ) - max_num_reqs = min( - runner.scheduler_config.max_num_seqs, - runner.max_num_tokens // query_len, - runner.kv_cache_config.num_blocks - 1, - ) - decode_kwargs = { - **max_token_prefill_kwargs, - "num_tokens": max_num_reqs * query_len, - "uniform_decode": True, - } - if config.use_v2_model_runner: - decode_kwargs["valid_dummy_state_slots"] = True - else: - decode_kwargs.update( - allow_microbatching=False, - force_attention=True, - profile_seq_lens=query_len + 1, - ) - return max_num_reqs, decode_kwargs - - -@contextmanager -def _temporary_replayssm_autotune_state( - runner: "GPUModelRunner", max_num_reqs: int -) -> Iterator[None]: - from vllm.model_executor.layers.mamba.mamba_mixer2 import MambaMixer2 - from vllm.model_executor.layers.mamba.ops.ssu_dispatch import ( - reset_replayssm_ring_trackers, - update_replayssm_ring_trackers, - ) - - reset_tensors: dict[int, torch.Tensor] = {} - tracker_specs: dict[int, tuple[torch.Tensor, torch.Tensor, int, int]] = {} - for module in runner.get_model().modules(): - if not isinstance(module, MambaMixer2) or not module.use_replayssm: - continue - assert module.replayssm_buffer_len is not None - ring_start = module._replayssm_ring_start - prev_num_accepted = module._replayssm_prev_num_accepted - tracker_specs.setdefault( - ring_start.data_ptr(), - ( - ring_start, - prev_num_accepted, - module.replayssm_buffer_len, - module.kv_cache[2].size(2), - ), - ) - tensors = ( - *module.kv_cache, - ring_start, - prev_num_accepted, - ) - for tensor in tensors: - if tensor.numel(): - reset_tensors.setdefault(tensor.data_ptr(), tensor) - - v2_runner: Any = runner - block_tables = saved_block_ids = None - if not runner.vllm_config.use_v2_model_runner: - block_tables = runner.input_batch.block_table.block_tables - saved_block_ids = tuple( - block_table.block_table.np[:max_num_reqs, 0].copy() - for block_table in block_tables - ) - dummy_block_ids = range(1, max_num_reqs + 1) - for block_table in block_tables: - block_table.block_table.np[:max_num_reqs, 0] = dummy_block_ids - runner.input_batch.block_table.commit_block_table(max_num_reqs) - - first_tracker = next(iter(tracker_specs.values()), None) - if first_tracker is not None and first_tracker[0].is_cuda: - state_slots = torch.arange( - 1, max_num_reqs + 1, dtype=torch.int32, device=first_tracker[0].device - ) - for ( - ring_start, - prev_num_accepted, - logical_window, - ring_buffer_len, - ) in tracker_specs.values(): - # Compile reset (prefill) and advance (decode) before inference. - # The final reset leaves the decode tuning run in a clean state. - reset_replayssm_ring_trackers(ring_start, prev_num_accepted, state_slots) - update_replayssm_ring_trackers( - ring_start, - prev_num_accepted, - state_slots, - logical_window, - ring_buffer_len, - ) - reset_replayssm_ring_trackers(ring_start, prev_num_accepted, state_slots) - - try: - yield - finally: - if runner.vllm_config.use_v2_model_runner: - v2_runner.block_tables.get_dummy_block_tables(max_num_reqs) - else: - assert block_tables is not None and saved_block_ids is not None - for block_table, block_ids in zip(block_tables, saved_block_ids): - block_table.block_table.np[:max_num_reqs, 0] = block_ids - runner.input_batch.block_table.commit_block_table(max_num_reqs) - for tensor in reset_tensors.values(): - tensor[1 : max_num_reqs + 1].zero_() - - def flashinfer_autotune(runner: "GPUModelRunner") -> None: """ Autotune FlashInfer operations. @@ -443,18 +311,6 @@ def flashinfer_autotune(runner: "GPUModelRunner") -> None: if is_leader: logger.info_once("Using FlashInfer autotune cache file: %s", cache_path) - # We skip EPLB here since we don't want to record dummy metrics. - # Randomize inputs to avoid every token pick the same experts, - # which lead to some EP ranks receiving no tokens and skipping their - # MoE kernel entirely, and cause hang due to all-reduce collective - # during synchronized autotuning. - max_token_prefill_kwargs = dict( - num_tokens=runner.scheduler_config.max_num_batched_tokens, - skip_eplb=True, - is_profile=True, - randomize_inputs=True, - ) - replayssm_autotune = _replayssm_autotune_kwargs(runner, max_token_prefill_kwargs) # Read cached autotune results and broadcast to all ranks. cached_results: bytes | None = None if is_leader and cache_path.exists(): @@ -474,10 +330,7 @@ def flashinfer_autotune(runner: "GPUModelRunner") -> None: fi_utils.autotune(tune_mode=True, **autotune_kwargs), ): _run_flashinfer_autotune_dummy_runs(runner) - if replayssm_autotune is not None: - max_num_reqs, max_batch_decode_kwargs = replayssm_autotune - with _temporary_replayssm_autotune_state(runner, max_num_reqs): - runner._dummy_run(**max_batch_decode_kwargs) + replayssm_autotune_warmup(runner) finally: set_autotune_process_group(None) diff --git a/vllm/model_executor/warmup/replayssm_warmup.py b/vllm/model_executor/warmup/replayssm_warmup.py new file mode 100644 index 000000000000..60b3eced2016 --- /dev/null +++ b/vllm/model_executor/warmup/replayssm_warmup.py @@ -0,0 +1,159 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project + +from collections.abc import Iterator +from contextlib import contextmanager +from typing import TYPE_CHECKING, Any + +import torch + +from vllm.config.mamba import MambaBackendEnum +from vllm.logger import init_logger +from vllm.model_executor.layers.mamba.ops.ssu_dispatch import ( + flashinfer_replayssm_autotune_supported, +) + +if TYPE_CHECKING: + from vllm.v1.worker.gpu_model_runner import GPUModelRunner + +logger = init_logger(__name__) + + +def _replayssm_autotune_kwargs( + runner: "GPUModelRunner", +) -> tuple[int, dict[str, Any]] | None: + config = runner.vllm_config + if not ( + config.cache_config.use_replayssm + and config.mamba_config.backend == MambaBackendEnum.FLASHINFER + ): + return None + if not flashinfer_replayssm_autotune_supported(): + logger.info_once( + "Skipping FlashInfer ReplaySSM autotuning because " + "flashinfer.mamba.checkpointing_ssu.CheckpointingSSURunner " + "is unavailable." + ) + return None + v2_runner: Any = runner + query_len = ( + v2_runner.decode_query_len + if config.use_v2_model_runner + else runner.uniform_decode_query_len + ) + max_num_reqs = min( + runner.scheduler_config.max_num_seqs, + runner.max_num_tokens // query_len, + runner.kv_cache_config.num_blocks - 1, + ) + decode_kwargs = { + "num_tokens": max_num_reqs * query_len, + "skip_eplb": True, + "is_profile": True, + "randomize_inputs": True, + "uniform_decode": True, + } + if config.use_v2_model_runner: + decode_kwargs["valid_dummy_state_slots"] = True + else: + decode_kwargs.update( + allow_microbatching=False, + force_attention=True, + profile_seq_lens=query_len + 1, + ) + return max_num_reqs, decode_kwargs + + +@contextmanager +def _temporary_replayssm_autotune_state( + runner: "GPUModelRunner", max_num_reqs: int +) -> Iterator[None]: + from vllm.model_executor.layers.mamba.mamba_mixer2 import MambaMixer2 + from vllm.model_executor.layers.mamba.ops.ssu_dispatch import ( + reset_replayssm_ring_trackers, + update_replayssm_ring_trackers, + ) + + reset_tensors: dict[int, torch.Tensor] = {} + tracker_specs: dict[int, tuple[torch.Tensor, torch.Tensor, int, int]] = {} + for module in runner.get_model().modules(): + if not isinstance(module, MambaMixer2) or not module.use_replayssm: + continue + assert module.replayssm_buffer_len is not None + ring_start = module._replayssm_ring_start + prev_num_accepted = module._replayssm_prev_num_accepted + tracker_specs.setdefault( + ring_start.data_ptr(), + ( + ring_start, + prev_num_accepted, + module.replayssm_buffer_len, + module.kv_cache[2].size(2), + ), + ) + tensors = ( + *module.kv_cache, + ring_start, + prev_num_accepted, + ) + for tensor in tensors: + if tensor.numel(): + reset_tensors.setdefault(tensor.data_ptr(), tensor) + + v2_runner: Any = runner + block_tables = saved_block_ids = None + if not runner.vllm_config.use_v2_model_runner: + block_tables = runner.input_batch.block_table.block_tables + saved_block_ids = tuple( + block_table.block_table.np[:max_num_reqs, 0].copy() + for block_table in block_tables + ) + dummy_block_ids = range(1, max_num_reqs + 1) + for block_table in block_tables: + block_table.block_table.np[:max_num_reqs, 0] = dummy_block_ids + runner.input_batch.block_table.commit_block_table(max_num_reqs) + + first_tracker = next(iter(tracker_specs.values()), None) + if first_tracker is not None and first_tracker[0].is_cuda: + state_slots = torch.arange( + 1, max_num_reqs + 1, dtype=torch.int32, device=first_tracker[0].device + ) + for ( + ring_start, + prev_num_accepted, + logical_window, + ring_buffer_len, + ) in tracker_specs.values(): + # Compile reset (prefill) and advance (decode) before inference. + # The final reset leaves the decode tuning run in a clean state. + reset_replayssm_ring_trackers(ring_start, prev_num_accepted, state_slots) + update_replayssm_ring_trackers( + ring_start, + prev_num_accepted, + state_slots, + logical_window, + ring_buffer_len, + ) + reset_replayssm_ring_trackers(ring_start, prev_num_accepted, state_slots) + + try: + yield + finally: + if runner.vllm_config.use_v2_model_runner: + v2_runner.block_tables.get_dummy_block_tables(max_num_reqs) + else: + assert block_tables is not None and saved_block_ids is not None + for block_table, block_ids in zip(block_tables, saved_block_ids): + block_table.block_table.np[:max_num_reqs, 0] = block_ids + runner.input_batch.block_table.commit_block_table(max_num_reqs) + for tensor in reset_tensors.values(): + tensor[1 : max_num_reqs + 1].zero_() + + +def replayssm_autotune_warmup(runner: "GPUModelRunner") -> None: + autotune = _replayssm_autotune_kwargs(runner) + if autotune is None: + return + max_num_reqs, decode_kwargs = autotune + with _temporary_replayssm_autotune_state(runner, max_num_reqs): + runner._dummy_run(**decode_kwargs) From b464da0dd1370f86329a1a5dd91bb82adc685ca0 Mon Sep 17 00:00:00 2001 From: Andrii Skliar Date: Tue, 1 Sep 2026 00:14:15 +0200 Subject: [PATCH 30/33] undo removing comments Signed-off-by: Andrii Skliar --- vllm/model_executor/warmup/kernel_warmup.py | 5 +++++ 1 file changed, 5 insertions(+) diff --git a/vllm/model_executor/warmup/kernel_warmup.py b/vllm/model_executor/warmup/kernel_warmup.py index eb1b30273428..03c65fc95da3 100644 --- a/vllm/model_executor/warmup/kernel_warmup.py +++ b/vllm/model_executor/warmup/kernel_warmup.py @@ -311,6 +311,11 @@ def flashinfer_autotune(runner: "GPUModelRunner") -> None: if is_leader: logger.info_once("Using FlashInfer autotune cache file: %s", cache_path) + # We skip EPLB here since we don't want to record dummy metrics. + # Randomize inputs to avoid every token pick the same experts, + # which lead to some EP ranks receiving no tokens and skipping their + # MoE kernel entirely, and cause hang due to all-reduce collective + # during synchronized autotuning. # Read cached autotune results and broadcast to all ranks. cached_results: bytes | None = None if is_leader and cache_path.exists(): From c89f2dbccb010a0180fc3be20403f8cd68382db1 Mon Sep 17 00:00:00 2001 From: Andrii Skliar Date: Tue, 1 Sep 2026 08:20:05 +0200 Subject: [PATCH 31/33] test(mamba): skip FlashInfer ReplaySSM warmup tests without CUDA Signed-off-by: Andrii Skliar --- tests/model_executor/test_replayssm_warmup.py | 7 +++++++ 1 file changed, 7 insertions(+) diff --git a/tests/model_executor/test_replayssm_warmup.py b/tests/model_executor/test_replayssm_warmup.py index bb59318b7cd6..272fa2cf8528 100644 --- a/tests/model_executor/test_replayssm_warmup.py +++ b/tests/model_executor/test_replayssm_warmup.py @@ -11,6 +11,13 @@ 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, From 00a1d2ec9f96eb84e7c692f7573c21b4ee173804 Mon Sep 17 00:00:00 2001 From: Andrii Skliar Date: Tue, 1 Sep 2026 12:00:51 +0200 Subject: [PATCH 32/33] fix pre-commit Signed-off-by: Andrii Skliar --- vllm/model_executor/warmup/kernel_warmup.py | 3 --- 1 file changed, 3 deletions(-) diff --git a/vllm/model_executor/warmup/kernel_warmup.py b/vllm/model_executor/warmup/kernel_warmup.py index c7ba593df751..ca816cdb44e8 100644 --- a/vllm/model_executor/warmup/kernel_warmup.py +++ b/vllm/model_executor/warmup/kernel_warmup.py @@ -41,9 +41,6 @@ from vllm.model_executor.warmup.replayssm_warmup import ( replayssm_autotune_warmup, ) -from vllm.model_executor.warmup.sparse_mla_triton_warmup import ( - sparse_mla_triton_warmup, -) from vllm.platforms import current_platform from vllm.utils.deep_gemm import is_deep_gemm_supported from vllm.utils.flashinfer import has_flashinfer From dc9fc45e05748c435536aa549fd95db838f68036 Mon Sep 17 00:00:00 2001 From: Andrii Skliar Date: Tue, 1 Sep 2026 14:43:36 +0200 Subject: [PATCH 33/33] fix test args Signed-off-by: Andrii Skliar --- tests/v1/worker/test_kv_cache_allocation_scope.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/tests/v1/worker/test_kv_cache_allocation_scope.py b/tests/v1/worker/test_kv_cache_allocation_scope.py index 1818f173d978..a9579f1af69b 100644 --- a/tests/v1/worker/test_kv_cache_allocation_scope.py +++ b/tests/v1/worker/test_kv_cache_allocation_scope.py @@ -48,7 +48,7 @@ def bind(*args, **kwargs): result = attn_utils.init_kv_cache( [], {}, - object(), + SimpleNamespace(kv_cache_groups=[]), torch.device("cpu"), [], config, @@ -83,7 +83,7 @@ def bind(*args, **kwargs): ) result = gpu_model_runner.GPUModelRunner.initialize_kv_cache_tensors( runner, - object(), + SimpleNamespace(kv_cache_groups=[]), [], kv_cache_allocation_context=scope, )