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_inplace

mamba2_ssd_chunk_scan_varlen_fwd_inplace()​

max.nn.state_space.mamba2_ssd_chunk_scan_varlen_fwd_inplace(x, dt, A, B, C, D, dt_bias, ssm_pool, query_start_loc, has_initial_state, cache_indices)

source

Performs the Mamba-2 SSD chunked-scan forward, writing final states back into the SSM pool in place.

Identical to mamba2_ssd_chunk_scan_varlen_fwd() but writes final states directly into ssm_pool[cache_indices[b], ...] in place instead of returning a separate final_states output tensor. This eliminates the graph-side buffer_load -> gather -> scatter_nd -> buffer_store whole-pool RMW that otherwise dominates decode GPU time.

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).
  • ssm_pool (BufferValue) – The [max_slots, nheads, head_dim, dstate] mutable state pool (fp32; bf16 on Apple GPUs — storage only, the scan accumulates in fp32). Read at ssm_pool[cache_indices[b]] when has_initial_state[b] is true; written in-place with final state.
  • query_start_loc (TensorValue) – The [batch + 1] int32 cumulative sequence lengths.
  • has_initial_state (TensorValue) – The [batch] bool, whether to load initial state for each sequence (empty to disable).
  • cache_indices (TensorValue) – The [batch] uint32 slot indices into ssm_pool.

Returns:

y, the [total_len, nheads, head_dim] output (model dtype). ssm_pool is mutated in place.

Return type:

TensorValue