IMPORTANT: To view this page as Markdown, append `.md` to the URL (e.g. /get-started.md). For the complete documentation index, see llms.txt.
Skip to main content
For the complete documentation index, see llms.txt. Markdown versions of all pages are available by appending .md to any URL (e.g. /get-started.md).

Python function

mamba2_ssd_chunk_scan_varlen_fwd

mamba2_ssd_chunk_scan_varlen_fwd()​

max.nn.state_space.mamba2_ssd_chunk_scan_varlen_fwd(x, dt, A, B, C, D, dt_bias, initial_states, query_start_loc, has_initial_state)

source

Performs the Mamba-2 SSD chunked-scan forward for prefill and decode.

Parameters:

  • x (TensorValue) – The [total_len, nheads, head_dim] SSM input (model dtype).
  • dt (TensorValue) – The [total_len, nheads] per-head time deltas (model dtype).
  • A (TensorValue) – The [nheads] per-head scalar (model dtype; already -exp(A_log)).
  • B (TensorValue) – The [total_len, ngroups, dstate] grouped input proj (model dtype).
  • C (TensorValue) – The [total_len, ngroups, dstate] grouped output proj (model dtype).
  • D (TensorValue) – The [nheads] skip connection (model dtype; empty to disable).
  • dt_bias (TensorValue) – The [nheads] dt bias (model dtype; empty to disable softplus bias).
  • initial_states (TensorValue) – The [batch, nheads, head_dim, dstate] fp32 initial SSM state (empty [0, ...] for a fresh prefill).
  • query_start_loc (TensorValue) – The [batch + 1] int32 cumulative sequence lengths.
  • has_initial_state (TensorValue) – The [batch] bool, whether to load initial_states for each sequence (empty to disable).

Returns:

(y, final_states) where y is [total_len, nheads, head_dim] (model dtype) and final_states is [batch, nheads, head_dim, dstate] fp32.

Return type:

tuple[TensorValue, TensorValue]