use super::flat::{MambaBackboneFlat, MambaLayerFlat};
use super::scratch::PhaseScratch;
use super::weights::{TrainMambaLayerWeights, TrainMambaWeights};
use crate::ops::blas::{matvec_forward, sgemm_forward};
use crate::ops::dims::{MambaDims, MambaRecurrentState};
use crate::ops::fast_math::{RMS_NORM_EPS, fast_exp_inplace, fast_exp_scalar};
pub fn forward_mamba_layer_batched(
temporal_flat: &mut [f32],
acts: &mut MambaLayerFlat,
layer_w: &TrainMambaLayerWeights,
state: &mut MambaRecurrentState<'_>,
scratch: &mut PhaseScratch,
dims: &MambaDims,
) {
let conv_state = &mut *state.conv;
let ssm_state = &mut *state.ssm;
let a_neg = state.a_neg;
let dm = dims.d_model;
let di = dims.d_inner;
let ds = dims.d_state;
let dc = dims.d_conv;
let dt_rank = dims.dt_rank;
let xdbl_dim = dims.xdbl_dim;
let seq_len = dims.seq_len;
for t in 0..seq_len {
let off = t * dm;
let src = &temporal_flat[off..off + dm];
acts.residual_mut(t).copy_from_slice(src);
let mut sum_sq = 0.0_f32;
for &v in &src[..dm] {
sum_sq += v * v;
}
let mean_sq = sum_sq / dm as f32;
let rms = (mean_sq + RMS_NORM_EPS).sqrt();
acts.set_rms_val(t, rms);
let inv_rms = 1.0 / rms;
let pn = &mut scratch.post_norm_flat[off..off + dm];
for d in 0..dm {
pn[d] = src[d] * inv_rms * layer_w.norm_weight[d];
}
acts.post_norm_mut(t).copy_from_slice(pn);
}
sgemm_forward(
&mut scratch.proj_flat,
&scratch.post_norm_flat,
&layer_w.in_proj_w,
None,
seq_len,
dm,
2 * di,
);
for t in 0..seq_len {
let proj_off = t * 2 * di;
acts.x_branch_mut(t)
.copy_from_slice(&scratch.proj_flat[proj_off..proj_off + di]);
let gate = &scratch.proj_flat[proj_off + di..proj_off + 2 * di];
acts.gate_pre_silu_mut(t).copy_from_slice(gate);
let gs_off = t * di;
for (d, &g) in gate.iter().enumerate().take(di) {
let sig = 1.0 / (1.0 + fast_exp_scalar(-g));
scratch.gate_silu_flat[gs_off + d] = g * sig;
}
acts.gate_post_silu_mut(t)
.copy_from_slice(&scratch.gate_silu_flat[gs_off..gs_off + di]);
}
let b_offset = dt_rank;
let c_offset = dt_rank + ds;
for t in 0..seq_len {
{
let x_branch = acts.x_branch(t);
for (d, &xb) in x_branch.iter().enumerate().take(di) {
let base = d * dc;
for k in 0..dc - 1 {
conv_state[base + k] = conv_state[base + k + 1];
}
conv_state[base + dc - 1] = xb;
}
}
acts.conv_state_mut(t)
.copy_from_slice(&conv_state[..di * dc]);
{
let step_base = t * acts.offsets.step_stride;
let pc_start = step_base + acts.offsets.post_conv;
let u_start = step_base + acts.offsets.u;
for d in 0..di {
let base = d * dc;
let mut val = layer_w.conv1d_bias[d];
for k in 0..dc {
val += conv_state[base + k] * layer_w.conv1d_weight[base + k];
}
acts.data[pc_start + d] = val;
acts.data[u_start + d] = val / (1.0 + fast_exp_scalar(-val));
}
}
let gs_off = t * di;
scratch.gate_silu_flat[..di].copy_from_slice(acts.u(t));
matvec_forward(
acts.xdbl_mut(t),
&scratch.gate_silu_flat[..di],
&layer_w.x_proj_w,
None,
di,
xdbl_dim,
);
scratch.gate_silu_flat[..dt_rank].copy_from_slice(&acts.xdbl(t)[..dt_rank]);
matvec_forward(
acts.delta_raw_mut(t),
&scratch.gate_silu_flat[..dt_rank],
&layer_w.dt_proj_w,
Some(&layer_w.dt_proj_b),
dt_rank,
di,
);
{
let step_base = t * acts.offsets.step_stride;
let dr_start = step_base + acts.offsets.delta_raw;
let d_start = step_base + acts.offsets.delta;
for d in 0..di {
let raw = acts.data[dr_start + d];
acts.data[d_start + d] = if raw > 20.0 {
raw
} else {
(1.0_f32 + fast_exp_scalar(raw)).ln()
};
}
}
acts.h_prev_mut(t).copy_from_slice(&ssm_state[..di * ds]);
{
let step_base = t * acts.offsets.step_stride;
let delta_start = step_base + acts.offsets.delta;
let u_start = step_base + acts.offsets.u;
let xdbl_start = step_base + acts.offsets.xdbl;
let da_start = step_base + acts.offsets.da_exp;
let y_start = step_base + acts.offsets.y;
for d in 0..di {
let delta_d = acts.data[delta_start + d];
let u_d = acts.data[u_start + d];
let delta_u_d = delta_d * u_d; let a_base = d * ds;
for n in 0..ds {
scratch.da_buf[n] = delta_d * a_neg[a_base + n];
}
fast_exp_inplace(&mut scratch.da_buf[..ds]);
let mut y_d = 0.0_f32;
for n in 0..ds {
let idx = a_base + n;
let b_n = acts.data[xdbl_start + b_offset + n];
let c_n = acts.data[xdbl_start + c_offset + n];
let da = scratch.da_buf[n];
acts.data[da_start + idx] = da;
let h_prev = ssm_state[idx];
ssm_state[idx] = da * h_prev + delta_u_d * b_n;
y_d += ssm_state[idx] * c_n;
}
y_d += layer_w.d_param[d] * u_d;
acts.data[y_start + d] = y_d;
}
}
acts.h_curr_mut(t).copy_from_slice(&ssm_state[..di * ds]);
{
let step_base = t * acts.offsets.step_stride;
let y_start = step_base + acts.offsets.y;
let gpost_start = step_base + acts.offsets.gate_post_silu;
let gated_start = step_base + acts.offsets.gated;
for d in 0..di {
scratch.gated_flat[gs_off + d] =
acts.data[y_start + d] * acts.data[gpost_start + d];
acts.data[gated_start + d] = scratch.gated_flat[gs_off + d];
}
}
}
sgemm_forward(
&mut scratch.out_flat,
&scratch.gated_flat,
&layer_w.out_proj_w,
None,
seq_len,
di,
dm,
);
for t in 0..seq_len {
let off = t * dm;
let residual = acts.residual(t);
for d in 0..dm {
temporal_flat[off + d] = residual[d] + scratch.out_flat[off + d];
}
}
}
pub fn forward_mamba_backbone_batched(
temporal_flat: &mut [f32],
acts: &mut MambaBackboneFlat,
mamba_w: &TrainMambaWeights,
mamba_input_flat: &[f32],
state: &mut MambaRecurrentState<'_>,
scratch: &mut PhaseScratch,
dims: &MambaDims,
) {
let conv_states = &mut *state.conv;
let ssm_states = &mut *state.ssm;
let a_neg_all = state.a_neg;
let dm = dims.d_model;
let di = dims.d_inner;
let ds = dims.d_state;
let dc = dims.d_conv;
let mid = dims.mamba_input_dim;
let seq_len = dims.seq_len;
let n_layers = dims.n_layers;
acts.input_proj_inputs
.copy_from_slice(&mamba_input_flat[..seq_len * mid]);
sgemm_forward(
&mut temporal_flat[..seq_len * dm],
&mamba_input_flat[..seq_len * mid],
&mamba_w.input_proj_w,
Some(&mamba_w.input_proj_b),
seq_len,
mid,
dm,
);
acts.input_proj_outputs[..seq_len * dm].copy_from_slice(&temporal_flat[..seq_len * dm]);
let conv_per_layer = di * dc;
let ssm_per_layer = di * ds;
let a_neg_per_layer = di * ds;
for layer_idx in 0..n_layers {
let conv_start = layer_idx * conv_per_layer;
let ssm_start = layer_idx * ssm_per_layer;
let a_neg_start = layer_idx * a_neg_per_layer;
forward_mamba_layer_batched(
temporal_flat,
&mut acts.layers[layer_idx],
&mamba_w.layers[layer_idx],
&mut MambaRecurrentState {
conv: &mut conv_states[conv_start..conv_start + conv_per_layer],
ssm: &mut ssm_states[ssm_start..ssm_start + ssm_per_layer],
a_neg: &a_neg_all[a_neg_start..a_neg_start + a_neg_per_layer],
},
scratch,
dims,
);
}
acts.norm_f_input[..seq_len * dm].copy_from_slice(&temporal_flat[..seq_len * dm]);
for t in 0..seq_len {
let off = t * dm;
let mean_sq: f32 = temporal_flat[off..off + dm]
.iter()
.map(|v| v * v)
.sum::<f32>()
/ dm as f32;
let rms = (mean_sq + RMS_NORM_EPS).sqrt();
acts.norm_f_rms[t] = rms;
let inv_rms = 1.0 / rms;
for d in 0..dm {
temporal_flat[off + d] = temporal_flat[off + d] * inv_rms * mamba_w.norm_f_weight[d];
}
}
}