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).

Mojo function

mamba2_ssd_chunk_scan_varlen_fwd_inplace_shape

def mamba2_ssd_chunk_scan_varlen_fwd_inplace_shape(x: T, dt: T, A: T, B: T, C: T, D: T, dt_bias: T, ssm_pool: T, query_start_loc: T, has_initial_state: T, cache_indices: T) -> IndexList[Int(3)]

Computes the output shape for the mamba2_ssd_chunk_scan_varlen_fwd_inplace graph op.

Args:

  • ​x (T): Packed input tensor of shape (total_len, nheads, head_dim).
  • ​dt (T): Per-head time deltas of shape (total_len, nheads).
  • ​A (T): Per-head scalar decay of shape (nheads,).
  • ​B (T): Grouped input projection of shape (total_len, ngroups, dstate).
  • ​C (T): Grouped output projection of shape (total_len, ngroups, dstate).
  • ​D (T): Per-head skip connection of shape (nheads,); may be empty when unused.
  • ​dt_bias (T): Per-head bias added to dt of shape (nheads,); may be empty when unused.
  • ​ssm_pool (T): Mutable SSM state pool of shape (max_slots, nheads, head_dim, dstate) in float32; final states are written in place at the slots indexed by cache_indices.
  • ​query_start_loc (T): Cumulative sequence lengths of shape (batch + 1,) in int32.
  • ​has_initial_state (T): Per-sequence flag of shape (batch,) in bool indicating whether to load the initial state from ssm_pool; may be empty when no initial states are used.
  • ​cache_indices (T): Per-sequence slot indices of shape (batch,) in uint32 selecting where in ssm_pool the final states are written.

Returns:

IndexList[Int(3)]

Was this page helpful?