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_shape

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

Computes the output shape for the mamba2_ssd_chunk_scan_varlen_fwd 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.
  • ​initial_states (T): Optional initial SSM states of shape (batch, nheads, head_dim, dstate) in float32; may be empty when has_initial_state is all false.
  • ​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 initial_states; may be empty when no initial states are used.

Returns:

IndexList[Int(3)]

Was this page helpful?