vllm.model_executor.layers.mamba.ops.ssu_dispatch
¶
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'.
Classes:
-
CPUSSUBackend–CPU SSU backend using the compiled C++ VSX/scalar kernel.
-
FlashInferSSUBackend–FlashInfer-based SSU backend.
-
MambaSSUBackend–Abstract base class for Mamba SSU backends.
-
TritonSSUBackend–Triton-based SSU backend (vLLM's default).
Functions:
-
commit_replayssm_ring_trackers–Commit the preceding speculative window and record the current one.
-
flashinfer_replayssm_autotune_supported–Return True when FlashInfer exposes ReplaySSM autotuning.
-
get_mamba_ssu_backend–Get the current Mamba SSU backend. Raises if not initialized.
-
initialize_mamba_ssu_backend–Initialize the Mamba SSU backend and optional FlashInfer ReplaySSM.
-
reset_replayssm_ring_trackers–Reset selected ReplaySSM ring trackers.
-
selective_state_update–Unified dispatch for Mamba selective state update.
-
selective_state_update_replayssm_flashinfer–Run FlashInfer checkpointing SSU and optionally advance shared trackers.
-
update_replayssm_ring_trackers–Reset selected trackers, or advance them when a window is provided.
CPUSSUBackend
¶
Bases: MambaSSUBackend
CPU SSU backend using the compiled C++ VSX/scalar kernel.
On CPU-only platforms (PowerPC, x86 without CUDA) this dispatches to
the vectorized C++ kernel registered as torch.ops._C.selective_state_update_cpu.
That kernel uses vec_op SIMD intrinsics (VSX on ppc64le, AVX2 on x86,
scalar fallback elsewhere) and is parallelised with OpenMP across heads.
Falls back to the pure-PyTorch implementation only if the C++ op is unavailable (e.g. a CPU-less build).
Source code in vllm/model_executor/layers/mamba/ops/ssu_dispatch.py
FlashInferSSUBackend
¶
Bases: MambaSSUBackend
FlashInfer-based SSU backend.
Source code in vllm/model_executor/layers/mamba/ops/ssu_dispatch.py
MambaSSUBackend
¶
Bases: ABC
Abstract base class for Mamba SSU backends.
Source code in vllm/model_executor/layers/mamba/ops/ssu_dispatch.py
TritonSSUBackend
¶
Bases: MambaSSUBackend
Triton-based SSU backend (vLLM's default).
Source code in vllm/model_executor/layers/mamba/ops/ssu_dispatch.py
commit_replayssm_ring_trackers(ring_start, prev_num_accepted, prev_query_len, state_batch_indices, num_accepted_tokens, query_start_loc, logical_window, ring_buffer_len, pad_slot_id=NULL_BLOCK_ID)
¶
Commit the preceding speculative window and record the current one.
MTP evaluates a target token and its draft tokens together, but the number
accepted from that query is available only on the next forward pass. This
function then advances the per-request ring by the accepted prefix and
records the current query length for the following pass. A zero
prev_query_len means that reset/prefill left no prior MTP query to
commit; slot validity is handled independently by the kernel mask.
Standard single-token decode needs no delayed commit: its ReplaySSM kernel advances the shared ring trackers directly after every token.
Source code in vllm/model_executor/layers/mamba/ops/ssu_dispatch.py
flashinfer_replayssm_autotune_supported()
cached
¶
Return True when FlashInfer exposes ReplaySSM autotuning.
Source code in vllm/model_executor/layers/mamba/ops/ssu_dispatch.py
get_mamba_ssu_backend()
¶
Get the current Mamba SSU backend. Raises if not initialized.
Source code in vllm/model_executor/layers/mamba/ops/ssu_dispatch.py
initialize_mamba_ssu_backend(mamba_config, kv_cache_config, *, use_replayssm=False)
¶
Initialize the Mamba SSU backend and optional FlashInfer ReplaySSM.
Source code in vllm/model_executor/layers/mamba/ops/ssu_dispatch.py
reset_replayssm_ring_trackers(ring_start, prev_num_accepted, prev_query_len, state_batch_indices, pad_slot_id=NULL_BLOCK_ID)
¶
Reset selected ReplaySSM ring trackers.
Source code in vllm/model_executor/layers/mamba/ops/ssu_dispatch.py
selective_state_update(state, x, dt, A, B, C, D, dt_bias, z=None, dt_softplus=False, state_batch_indices=None, dst_state_batch_indices=None, null_block_id=NULL_BLOCK_ID, out=None, num_accepted_tokens=None, cu_seqlens=None, is_blackwell=False)
¶
Unified dispatch for Mamba selective state update.
Delegates to the initialized backend (Triton or FlashInfer).
Source code in vllm/model_executor/layers/mamba/ops/ssu_dispatch.py
selective_state_update_replayssm_flashinfer(state, x, dt, A, B, C, out, x_cache, B_cache, dt_cache, ring_start, prev_num_accepted_tokens, prev_query_len, logical_window, D=None, dt_bias=None, dt_softplus=False, state_batch_indices=None, null_block_id=NULL_BLOCK_ID, scratch=None, update_trackers=True, enable_stochastic_rounding=False, stochastic_rounding_philox_rounds=0, cu_seqlens=None, max_seqlen=None, enable_pdl=False)
¶
Run FlashInfer checkpointing SSU and optionally advance shared trackers.
Source code in vllm/model_executor/layers/mamba/ops/ssu_dispatch.py
497 498 499 500 501 502 503 504 505 506 507 508 509 510 511 512 513 514 515 516 517 518 519 520 521 522 523 524 525 526 527 528 529 530 531 532 533 534 535 536 537 538 539 540 541 542 543 544 545 546 547 548 549 550 551 552 553 554 555 556 557 558 559 560 561 562 563 564 565 566 567 568 569 570 571 572 573 574 575 576 577 578 579 580 581 582 583 584 585 586 587 588 589 590 | |
update_replayssm_ring_trackers(ring_start, prev_num_accepted, prev_query_len, state_batch_indices, logical_window=None, ring_buffer_len=None, pad_slot_id=NULL_BLOCK_ID)
¶
Reset selected trackers, or advance them when a window is provided.