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)