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
selective_scan_fwd_gpu
def selective_scan_fwd_gpu[kernel_dtype: DType, DSTATE: Int, output_LT: TensorLayout, x_LT: TensorLayout, out_z_LT: TensorLayout, u_LT: TensorLayout, delta_LT: TensorLayout, A_LT: TensorLayout, B_LT: TensorLayout, C_LT: TensorLayout, D_LT: TensorLayout, z_LT: TensorLayout, delta_bias_LT: TensorLayout](total_batch_dim: Int, batch: Int, dim: Int, seqlen: Int, group_size: Int, delta_softplus: Int8, output: TileTensor[kernel_dtype, output_LT, MutUntrackedOrigin], x: TileTensor[kernel_dtype, x_LT, MutUntrackedOrigin], out_z: TileTensor[kernel_dtype, out_z_LT, MutUntrackedOrigin], u: TileTensor[kernel_dtype, u_LT, MutUntrackedOrigin], delta: TileTensor[kernel_dtype, delta_LT, MutUntrackedOrigin], A: TileTensor[kernel_dtype, A_LT, MutUntrackedOrigin], B: TileTensor[kernel_dtype, B_LT, MutUntrackedOrigin], C: TileTensor[kernel_dtype, C_LT, MutUntrackedOrigin], D: TileTensor[kernel_dtype, D_LT, MutUntrackedOrigin], z: TileTensor[kernel_dtype, z_LT, MutUntrackedOrigin], delta_bias: TileTensor[kernel_dtype, delta_bias_LT, MutUntrackedOrigin], output_strides: IndexList[3], x_strides: IndexList[4], out_z_strides: IndexList[3], u_strides: IndexList[3], delta_strides: IndexList[3], A_strides: IndexList[2], B_strides: IndexList[4], C_strides: IndexList[4], D_strides: IndexList[1], z_strides: IndexList[3], delta_bias_strides: IndexList[1])
GPU kernel for selective scan forward pass.
Each thread processes one (batch, dim) pair and iterates through the sequence.