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
varlen_selective_scan_fwd_cpu
def varlen_selective_scan_fwd_cpu[kernel_dtype: DType, DSTATE: Int](dim: Int, ngroups: Int, batch: Int, pad_slot_id: Int32, delta_softplus: Int8, u: TileTensor[kernel_dtype, Engine=u.Engine, address_space=u.address_space, linear_idx_type=u.linear_idx_type], delta: TileTensor[kernel_dtype, Engine=delta.Engine, address_space=delta.address_space, linear_idx_type=delta.linear_idx_type], A: TileTensor[kernel_dtype, Engine=A.Engine, address_space=A.address_space, linear_idx_type=A.linear_idx_type], B: TileTensor[kernel_dtype, Engine=B.Engine, address_space=B.address_space, linear_idx_type=B.linear_idx_type], C: TileTensor[kernel_dtype, Engine=C.Engine, address_space=C.address_space, linear_idx_type=C.linear_idx_type], D: TileTensor[kernel_dtype, Engine=D.Engine, address_space=D.address_space, linear_idx_type=D.linear_idx_type], z: TileTensor[kernel_dtype, Engine=z.Engine, address_space=z.address_space, linear_idx_type=z.linear_idx_type], delta_bias: TileTensor[kernel_dtype, Engine=delta_bias.Engine, address_space=delta_bias.address_space, linear_idx_type=delta_bias.linear_idx_type], ssm_states: TileTensor[kernel_dtype, Engine=ssm_states.Engine, address_space=ssm_states.address_space, linear_idx_type=ssm_states.linear_idx_type], output: TileTensor[kernel_dtype, Engine=output.Engine, address_space=output.address_space, linear_idx_type=output.linear_idx_type], query_start_loc: TileTensor[.int32, Engine=query_start_loc.Engine, address_space=query_start_loc.address_space, linear_idx_type=query_start_loc.linear_idx_type], cache_indices: TileTensor[.int32, Engine=cache_indices.Engine, address_space=cache_indices.address_space, linear_idx_type=cache_indices.linear_idx_type], has_initial_state: TileTensor[.bool, Engine=has_initial_state.Engine, address_space=has_initial_state.address_space, linear_idx_type=has_initial_state.linear_idx_type], u_strides: IndexList[Int(2)], delta_strides: IndexList[Int(2)], A_strides: IndexList[Int(2)], B_strides: IndexList[Int(3)], C_strides: IndexList[Int(3)], D_strides: IndexList[Int(1)], z_strides: IndexList[Int(2)], delta_bias_strides: IndexList[Int(1)], ssm_states_strides: IndexList[Int(3)], out_strides: IndexList[Int(2)], ctx: Optional[DeviceContext] = None)
CPU kernel for variable-length selective scan.