mamba2_ssm_prefill

Function mamba2_ssm_prefill 

Source
pub fn mamba2_ssm_prefill(
    gpu: &dyn GpuBackend,
    kernel: KernelHandle,
    h_state: DevicePtr,
    x: DevicePtr,
    b_proj: DevicePtr,
    c_proj: DevicePtr,
    dt_raw: DevicePtr,
    a_log: DevicePtr,
    d_param: DevicePtr,
    dt_bias: DevicePtr,
    output: DevicePtr,
    batch_size: u32,
    seq_len: u32,
    num_heads: u32,
    head_dim: u32,
    state_size: u32,
    n_groups: u32,
    dt_min: f32,
    dt_max: f32,
    x_stride: u32,
    bc_stride: u32,
    dt_stride: u32,
    y_stride: u32,
    stream: u64,
) -> Result<()>
Expand description

Mamba-2 SSM prefill: sequential recurrence across seq_len tokens in a single kernel.

Same algorithm as decode but loops over tokens, avoiding per-token launch overhead. Supports non-contiguous layouts via per-tensor strides (BF16 elements between tokens).

Grid: (num_heads, batch_size, 1) Block: (state_size, 1, 1)