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)
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 loadinitial_statesfor each sequence (empty to disable).
- x (TensorValue) – The
-
Returns:
-
(y, final_states)whereyis[total_len, nheads, head_dim](model dtype) andfinal_statesis[batch, nheads, head_dim, dstate]fp32. -
Return type: