use rayon::prelude::*;
use crate::ops::blas::{sgemm_forward, sgemm_forward_par};
use crate::ops::dims::MambaDims;
use crate::ops::fast_math::{fast_exp_inplace, fast_exp_scalar};
use crate::state::MambaState;
use crate::weights::{MambaLayerWeights, MambaWeights};
const MAX_D_STATE: usize = 64;
type GemmFn = fn(&mut [f32], &[f32], &[f32], Option<&[f32]>, usize, usize, usize);
const TRANSPOSE_TILE: usize = 64;
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum PrefillMode {
Single,
Parallel,
}
pub(crate) fn zip_rows<F>(mode: PrefillMode, dst: (&mut [f32], usize), src: (&[f32], usize), f: F)
where
F: Fn(&mut [f32], &[f32]) + Sync + Send,
{
let (dst, dst_w) = dst;
let (src, src_w) = src;
match mode {
PrefillMode::Single => dst
.chunks_mut(dst_w)
.zip(src.chunks(src_w))
.for_each(|(d, s)| f(d, s)),
PrefillMode::Parallel => dst
.par_chunks_mut(dst_w)
.zip(src.par_chunks(src_w))
.for_each(|(d, s)| f(d, s)),
}
}
pub(crate) fn for_rows<F>(mode: PrefillMode, buf: &mut [f32], row_w: usize, f: F)
where
F: Fn(&mut [f32]) + Sync + Send,
{
match mode {
PrefillMode::Single => buf.chunks_mut(row_w).for_each(f),
PrefillMode::Parallel => buf.par_chunks_mut(row_w).for_each(f),
}
}
pub struct PrefillScratch {
fingerprint: (usize, usize, usize, usize, usize, usize, usize),
post_norm: Vec<f32>,
proj: Vec<f32>,
gate_silu: Vec<f32>,
x_cm: Vec<f32>,
u_cm: Vec<f32>,
delta_cm: Vec<f32>,
y_cm: Vec<f32>,
u_rm: Vec<f32>,
delta_rm: Vec<f32>,
y_rm: Vec<f32>,
xdbl: Vec<f32>,
dt_in: Vec<f32>,
out: Vec<f32>,
conv_reg: Vec<f32>,
}
fn fingerprint(dims: &MambaDims) -> (usize, usize, usize, usize, usize, usize, usize) {
(
dims.seq_len,
dims.d_model,
dims.d_inner,
dims.d_state,
dims.d_conv,
dims.dt_rank,
dims.mamba_input_dim,
)
}
impl PrefillScratch {
pub fn new(dims: &MambaDims) -> Self {
let t = dims.seq_len;
let di = dims.d_inner;
Self {
fingerprint: fingerprint(dims),
post_norm: vec![0.0; t * dims.d_model],
proj: vec![0.0; t * 2 * di],
gate_silu: vec![0.0; t * di],
x_cm: vec![0.0; di * t],
u_cm: vec![0.0; di * t],
delta_cm: vec![0.0; di * t],
y_cm: vec![0.0; di * t],
u_rm: vec![0.0; t * di],
delta_rm: vec![0.0; t * di],
y_rm: vec![0.0; t * di],
xdbl: vec![0.0; t * dims.xdbl_dim],
dt_in: vec![0.0; t * dims.dt_rank],
out: vec![0.0; t * dims.d_model],
conv_reg: vec![0.0; di * dims.d_conv],
}
}
pub fn ensure(&mut self, dims: &MambaDims) {
if fingerprint(dims) != self.fingerprint {
*self = Self::new(dims);
}
}
}
fn conv_channel(
d: usize,
u_col: &mut [f32],
reg: &mut [f32],
x_col: &[f32],
lw: &MambaLayerWeights,
dims: &MambaDims,
) {
let dc = dims.d_conv;
let hist = dc - 1;
let t_len = u_col.len();
let w_base = d * dc;
let bias = lw.conv1d_bias[d];
let head = hist.min(t_len);
for (u_td, &xb) in u_col[..head].iter_mut().zip(&x_col[..head]) {
for k in 0..dc - 1 {
reg[k] = reg[k + 1];
}
reg[dc - 1] = xb;
let mut val = bias;
for (r, w_k) in reg.iter().zip(&lw.conv1d_weight[w_base..w_base + dc]) {
val += r * w_k;
}
*u_td = val;
}
if t_len > hist {
let bulk = &mut u_col[hist..];
bulk.fill(bias);
for (k, &w_k) in lw.conv1d_weight[w_base..w_base + dc].iter().enumerate() {
for (u, &xv) in bulk.iter_mut().zip(&x_col[k..k + t_len - hist]) {
*u += xv * w_k;
}
}
for (r, &xv) in reg[1..dc].iter_mut().zip(&x_col[t_len - hist..]) {
*r = xv;
}
}
for u in u_col.iter_mut() {
*u = *u / (1.0 + fast_exp_scalar(-*u));
}
}
const SSM_BLOCK: usize = 8;
fn ssm_channels(
block: usize,
y_cols: &mut [f32],
ssm: &mut [f32],
io: (&[f32], &[f32], &[f32]),
lw: &MambaLayerWeights,
dims: &MambaDims,
) {
let (u_cols, delta_cols, xdbl) = io;
let ds = dims.d_state;
let dr = dims.dt_rank;
let xd = dims.xdbl_dim;
let t_len = dims.seq_len;
let d0 = block * SSM_BLOCK;
let nch = y_cols.len() / t_len;
let b_offset = dr;
let c_offset = dr + ds;
let mut da = [0.0f32; MAX_D_STATE * SSM_BLOCK];
let mut hloc = [0.0f32; MAX_D_STATE * SSM_BLOCK];
let mut du = [0.0f32; SSM_BLOCK];
let mut y_acc = [0.0f32; SSM_BLOCK];
for (c, ssm_ch) in ssm.chunks(ds).enumerate() {
for (n, &s) in ssm_ch.iter().enumerate() {
hloc[n * SSM_BLOCK + c] = s;
}
}
for t in 0..t_len {
let xdbl_row = t * xd;
for (c, du_c) in du[..nch].iter_mut().enumerate() {
let delta_d = delta_cols[c * t_len + t];
*du_c = delta_d * u_cols[c * t_len + t];
let a_base = (d0 + c) * ds;
for (n, a_n) in lw.a_neg[a_base..a_base + ds].iter().enumerate() {
da[n * SSM_BLOCK + c] = delta_d * a_n;
}
}
fast_exp_inplace(&mut da[..ds * SSM_BLOCK]);
y_acc.fill(0.0);
for (n, (da_l, h_l)) in da
.chunks_exact(SSM_BLOCK)
.zip(hloc.chunks_exact_mut(SSM_BLOCK))
.enumerate()
.take(ds)
{
let b_n = xdbl[xdbl_row + b_offset + n];
let c_n = xdbl[xdbl_row + c_offset + n];
for c in 0..SSM_BLOCK {
let h = da_l[c] * h_l[c] + du[c] * b_n;
h_l[c] = h;
y_acc[c] += h * c_n;
}
}
for (c, &y_c) in y_acc[..nch].iter().enumerate() {
let u_td = u_cols[c * t_len + t];
y_cols[c * t_len + t] = y_c + lw.d_param[d0 + c] * u_td;
}
}
for (c, ssm_ch) in ssm.chunks_mut(ds).enumerate() {
for (n, s) in ssm_ch.iter_mut().enumerate() {
*s = hloc[n * SSM_BLOCK + c];
}
}
}
pub(crate) fn transpose_cm_to_rm(
dst: &mut [f32],
src: &[f32],
di: usize,
t_len: usize,
mode: PrefillMode,
) {
let row_block = TRANSPOSE_TILE * di;
let work = |(tb, dst_tile): (usize, &mut [f32])| {
let t0 = tb * TRANSPOSE_TILE;
let rows = dst_tile.len() / di;
for d0 in (0..di).step_by(TRANSPOSE_TILE) {
let d1 = (d0 + TRANSPOSE_TILE).min(di);
for d in d0..d1 {
let col = &src[d * t_len + t0..d * t_len + t0 + rows];
for (r, &v) in col.iter().enumerate() {
dst_tile[r * di + d] = v;
}
}
}
};
match mode {
PrefillMode::Single => dst[..t_len * di]
.chunks_mut(row_block)
.enumerate()
.for_each(work),
PrefillMode::Parallel => dst[..t_len * di]
.par_chunks_mut(row_block)
.enumerate()
.for_each(work),
}
}
pub(crate) fn transpose_rows_to_cm(
dst: &mut [f32],
src: &[f32],
di: usize,
t_len: usize,
stride: (usize, usize),
mode: PrefillMode,
) {
let (src_stride, src_off) = stride;
let col_block = TRANSPOSE_TILE * t_len;
let work = |(db, dst_tile): (usize, &mut [f32])| {
let d0 = db * TRANSPOSE_TILE;
for t0 in (0..t_len).step_by(TRANSPOSE_TILE) {
let t1 = (t0 + TRANSPOSE_TILE).min(t_len);
for (c, col) in dst_tile.chunks_mut(t_len).enumerate() {
for (i, v) in col[t0..t1].iter_mut().enumerate() {
*v = src[(t0 + i) * src_stride + src_off + d0 + c];
}
}
}
};
match mode {
PrefillMode::Single => dst[..di * t_len]
.chunks_mut(col_block)
.enumerate()
.for_each(work),
PrefillMode::Parallel => dst[..di * t_len]
.par_chunks_mut(col_block)
.enumerate()
.for_each(work),
}
}
pub fn forward_mamba_backbone_prefill(
temporal_out: &mut [f32],
mamba_input_flat: &[f32],
w: &MambaWeights,
state: &mut MambaState,
scratch: &mut PrefillScratch,
dims: &MambaDims,
) {
forward_mamba_backbone_prefill_mode(
temporal_out,
mamba_input_flat,
w,
state,
scratch,
dims,
PrefillMode::Single,
);
}
pub fn forward_mamba_backbone_prefill_mode(
temporal_out: &mut [f32],
mamba_input_flat: &[f32],
w: &MambaWeights,
state: &mut MambaState,
scratch: &mut PrefillScratch,
dims: &MambaDims,
mode: PrefillMode,
) {
let dm = dims.d_model;
let di = dims.d_inner;
let ds = dims.d_state;
let dc = dims.d_conv;
let dr = dims.dt_rank;
let xd = dims.xdbl_dim;
let mid = dims.mamba_input_dim;
let t_len = dims.seq_len;
let hist = dc - 1;
let eps = dims.rms_norm_eps;
assert!(ds <= MAX_D_STATE, "d_state > {MAX_D_STATE} unsupported");
assert_eq!(temporal_out.len(), t_len * dm, "temporal_out shape");
assert_eq!(mamba_input_flat.len(), t_len * mid, "input shape");
assert_eq!(state.layers.len(), dims.n_layers, "state layer count");
scratch.ensure(dims);
let gemm: GemmFn = match mode {
PrefillMode::Single => sgemm_forward,
PrefillMode::Parallel => sgemm_forward_par,
};
if w.input_proj_w.is_empty() {
debug_assert_eq!(mid, dm, "identity input_proj requires input_dim == d_model");
temporal_out[..t_len * dm].copy_from_slice(&mamba_input_flat[..t_len * dm]);
} else {
gemm(
temporal_out,
mamba_input_flat,
&w.input_proj_w,
Some(&w.input_proj_b),
t_len,
mid,
dm,
);
}
for (layer_idx, lw) in w.layers.iter().enumerate() {
let lstate = &mut state.layers[layer_idx];
zip_rows(
mode,
(&mut scratch.post_norm[..t_len * dm], dm),
(&temporal_out[..t_len * dm], dm),
|dst, src| {
let mut sum_sq = 0.0_f32;
for &v in src {
sum_sq += v * v;
}
let inv_rms = 1.0 / (sum_sq / dm as f32 + eps).sqrt();
for ((y, &s), nw) in dst.iter_mut().zip(src).zip(&lw.norm_weight[..dm]) {
*y = s * inv_rms * nw;
}
},
);
gemm(
&mut scratch.proj,
&scratch.post_norm,
&lw.in_proj_w,
None,
t_len,
dm,
2 * di,
);
zip_rows(
mode,
(&mut scratch.gate_silu[..t_len * di], di),
(&scratch.proj[..t_len * 2 * di], 2 * di),
|gs, proj_row| {
for (g_out, &g) in gs.iter_mut().zip(&proj_row[di..2 * di]) {
*g_out = g * (1.0 / (1.0 + fast_exp_scalar(-g)));
}
},
);
transpose_rows_to_cm(
&mut scratch.x_cm,
&scratch.proj,
di,
t_len,
(2 * di, 0),
mode,
);
for d in 0..di {
scratch.conv_reg[d * dc] = 0.0;
scratch.conv_reg[d * dc + 1..d * dc + dc]
.copy_from_slice(&lstate.conv_state[d * hist..(d + 1) * hist]);
}
{
let x_cm = &scratch.x_cm[..];
match mode {
PrefillMode::Single => {
for (d, (u_col, reg)) in scratch
.u_cm
.chunks_mut(t_len)
.zip(scratch.conv_reg.chunks_mut(dc))
.enumerate()
{
conv_channel(d, u_col, reg, &x_cm[d * t_len..(d + 1) * t_len], lw, dims);
}
}
PrefillMode::Parallel => {
scratch
.u_cm
.par_chunks_mut(t_len)
.zip(scratch.conv_reg.par_chunks_mut(dc))
.enumerate()
.for_each(|(d, (u_col, reg))| {
conv_channel(
d,
u_col,
reg,
&x_cm[d * t_len..(d + 1) * t_len],
lw,
dims,
);
});
}
}
}
transpose_cm_to_rm(&mut scratch.u_rm, &scratch.u_cm, di, t_len, mode);
gemm(
&mut scratch.xdbl,
&scratch.u_rm,
&lw.x_proj_w,
None,
t_len,
di,
xd,
);
zip_rows(
mode,
(&mut scratch.dt_in[..t_len * dr], dr),
(&scratch.xdbl[..t_len * xd], xd),
|dst, src| dst.copy_from_slice(&src[..dr]),
);
gemm(
&mut scratch.delta_rm,
&scratch.dt_in,
&lw.dt_proj_w,
Some(&lw.dt_proj_b),
t_len,
dr,
di,
);
for_rows(mode, &mut scratch.delta_rm[..t_len * di], di, |row| {
for v in row {
if *v <= 20.0 {
*v = fast_exp_scalar(*v).ln_1p();
}
}
});
transpose_rows_to_cm(
&mut scratch.delta_cm,
&scratch.delta_rm,
di,
t_len,
(di, 0),
mode,
);
{
let u_cm = &scratch.u_cm[..];
let delta_cm = &scratch.delta_cm[..];
let xdbl = &scratch.xdbl[..];
let blk_t = SSM_BLOCK * t_len;
match mode {
PrefillMode::Single => {
for (blk, (y_cols, ssm)) in scratch
.y_cm
.chunks_mut(blk_t)
.zip(lstate.ssm_state.chunks_mut(SSM_BLOCK * ds))
.enumerate()
{
let span = blk * blk_t..blk * blk_t + y_cols.len();
let io = (&u_cm[span.clone()], &delta_cm[span], xdbl);
ssm_channels(blk, y_cols, ssm, io, lw, dims);
}
}
PrefillMode::Parallel => {
scratch
.y_cm
.par_chunks_mut(blk_t)
.zip(lstate.ssm_state.par_chunks_mut(SSM_BLOCK * ds))
.enumerate()
.for_each(|(blk, (y_cols, ssm))| {
let span = blk * blk_t..blk * blk_t + y_cols.len();
let io = (&u_cm[span.clone()], &delta_cm[span], xdbl);
ssm_channels(blk, y_cols, ssm, io, lw, dims);
});
}
}
}
transpose_cm_to_rm(&mut scratch.y_rm, &scratch.y_cm, di, t_len, mode);
zip_rows(
mode,
(&mut scratch.y_rm[..t_len * di], di),
(&scratch.gate_silu[..t_len * di], di),
|y_row, g_row| {
for (y, &g) in y_row.iter_mut().zip(g_row) {
*y *= g;
}
},
);
for d in 0..di {
lstate.conv_state[d * hist..(d + 1) * hist]
.copy_from_slice(&scratch.conv_reg[d * dc + 1..d * dc + dc]);
}
gemm(
&mut scratch.out,
&scratch.y_rm,
&lw.out_proj_w,
None,
t_len,
di,
dm,
);
zip_rows(
mode,
(&mut temporal_out[..t_len * dm], dm),
(&scratch.out[..t_len * dm], dm),
|y_row, o_row| {
for (y, &o) in y_row.iter_mut().zip(o_row) {
*y += o;
}
},
);
}
for_rows(mode, &mut temporal_out[..t_len * dm], dm, |row| {
let mut sum_sq = 0.0_f32;
for &v in row.iter() {
sum_sq += v * v;
}
let inv_rms = 1.0 / (sum_sq / dm as f32 + eps).sqrt();
for (y, nw) in row.iter_mut().zip(&w.norm_f_weight[..dm]) {
*y *= inv_rms * nw;
}
});
}
pub fn prefill_batch(
outputs: &mut [f32],
inputs: &[f32],
w: &MambaWeights,
states: &mut [MambaState],
scratches: &mut [PrefillScratch],
dims: &MambaDims,
) {
let t_len = dims.seq_len;
let dm = dims.d_model;
let mid = dims.mamba_input_dim;
let b = states.len();
assert_eq!(scratches.len(), b, "one scratch per sample");
assert_eq!(outputs.len(), b * t_len * dm, "outputs shape");
assert_eq!(inputs.len(), b * t_len * mid, "inputs shape");
outputs
.par_chunks_mut(t_len * dm)
.zip(states.par_iter_mut())
.zip(scratches.par_iter_mut())
.enumerate()
.for_each(|(i, ((out, state), scratch))| {
let inp = &inputs[i * t_len * mid..(i + 1) * t_len * mid];
forward_mamba_backbone_prefill(out, inp, w, state, scratch, dims);
});
}