use std::cell::{Cell, RefCell};
use std::sync::atomic::{AtomicU64, Ordering};
use rayon::prelude::*;
use super::backward::backward_mamba3_layer_batched;
use super::dims::Mamba3Dims;
use super::flat::Mamba3LayerFlat;
use super::forward::forward_mamba3_layer_batched;
use super::scratch::Mamba3Scratch;
use super::weights::TrainMamba3Weights;
struct Mamba3ThreadScratch {
phase: Mamba3Scratch,
temporal_flat: Vec<f32>,
d_model: usize,
d_inner: usize,
seq_len: usize,
}
impl Mamba3ThreadScratch {
fn new(dims: &Mamba3Dims) -> Self {
Self {
phase: Mamba3Scratch::zeros(dims),
temporal_flat: vec![0.0; dims.seq_len * dims.d_model],
d_model: dims.d_model,
d_inner: dims.d_inner,
seq_len: dims.seq_len,
}
}
}
static M3_BWD_EPOCH: AtomicU64 = AtomicU64::new(0);
thread_local! {
static M3_THREAD_SCRATCH: RefCell<Option<Mamba3ThreadScratch>> = const { RefCell::new(None) };
static M3_THREAD_GRADS: RefCell<Option<TrainMamba3Weights>> = const { RefCell::new(None) };
static M3_THREAD_GRADS_EPOCH: Cell<u64> = const { Cell::new(0) };
}
fn ensure_grads_zeroed(dims: &Mamba3Dims, input_dim: usize, epoch: u64) {
M3_THREAD_GRADS_EPOCH.with(|ep| {
if ep.get() != epoch {
M3_THREAD_GRADS.with(|cell| {
let mut opt = cell.borrow_mut();
let g = opt.get_or_insert_with(|| TrainMamba3Weights::zeros(dims, input_dim));
g.zero();
});
ep.set(epoch);
}
});
}
fn with_scratch<F, R>(dims: &Mamba3Dims, f: F) -> R
where
F: FnOnce(&mut Mamba3ThreadScratch) -> R,
{
M3_THREAD_SCRATCH.with(|cell| {
let mut opt = cell.borrow_mut();
let needs_reinit = opt.as_ref().is_some_and(|s| {
s.d_model != dims.d_model || s.d_inner != dims.d_inner || s.seq_len != dims.seq_len
});
if needs_reinit {
*opt = None;
}
let scratch = opt.get_or_insert_with(|| Mamba3ThreadScratch::new(dims));
f(scratch)
})
}
pub fn invalidate_mamba3_scratch() {
rayon::broadcast(|_| {
M3_THREAD_SCRATCH.with(|cell| *cell.borrow_mut() = None);
M3_THREAD_GRADS.with(|cell| *cell.borrow_mut() = None);
});
}
struct BatchPtrs {
temporal: *mut f32,
ssm: *mut f32,
k: *mut f32,
v: *mut f32,
angle: *mut f32,
acts: *mut Vec<Mamba3LayerFlat>,
}
unsafe impl Send for BatchPtrs {}
unsafe impl Sync for BatchPtrs {}
#[allow(clippy::too_many_arguments)]
pub fn parallel_mamba3_forward(
temporal_out: &mut [f32],
batch_acts: &mut [Vec<Mamba3LayerFlat>],
ssm_states: &mut [f32],
k_states: &mut [f32],
v_states: &mut [f32],
angle_states: &mut [f32],
weights: &TrainMamba3Weights,
dims: &Mamba3Dims,
batch_size: usize,
) {
let dm = dims.d_model;
let nl = dims.n_layers;
let nh = dims.nheads;
let hd = dims.headdim;
let ds = dims.d_state;
let seq_len = dims.seq_len;
let na = dims.num_rope_angles.max(1);
let ssm_per = nl * nh * hd * ds;
let k_per = nl * nh * ds;
let v_per = nl * nh * hd;
let a_per = nl * nh * na;
let t_per = seq_len * dm;
let ptrs = &BatchPtrs {
temporal: temporal_out.as_mut_ptr(),
ssm: ssm_states.as_mut_ptr(),
k: k_states.as_mut_ptr(),
v: v_states.as_mut_ptr(),
angle: angle_states.as_mut_ptr(),
acts: batch_acts.as_mut_ptr(),
};
(0..batch_size).into_par_iter().for_each(|b| {
with_scratch(dims, |tls| {
let temporal =
unsafe { std::slice::from_raw_parts_mut(ptrs.temporal.add(b * t_per), t_per) };
let acts = unsafe { &mut *ptrs.acts.add(b) };
for (layer_idx, lw) in weights.layers.iter().enumerate() {
let ssm_off = b * ssm_per + layer_idx * nh * hd * ds;
let k_off = b * k_per + layer_idx * nh * ds;
let v_off = b * v_per + layer_idx * nh * hd;
let a_off = b * a_per + layer_idx * nh * na;
let ssm =
unsafe { std::slice::from_raw_parts_mut(ptrs.ssm.add(ssm_off), nh * hd * ds) };
let k = unsafe { std::slice::from_raw_parts_mut(ptrs.k.add(k_off), nh * ds) };
let v = unsafe { std::slice::from_raw_parts_mut(ptrs.v.add(v_off), nh * hd) };
let angle =
unsafe { std::slice::from_raw_parts_mut(ptrs.angle.add(a_off), nh * na) };
forward_mamba3_layer_batched(
temporal,
&mut acts[layer_idx],
lw,
ssm,
k,
v,
angle,
&mut tls.phase,
dims,
);
}
});
});
}
#[allow(clippy::too_many_arguments)]
pub fn parallel_mamba3_backward(
d_temporal_out: &mut [f32],
batch_acts: &[Vec<Mamba3LayerFlat>],
weights: &TrainMamba3Weights,
d_weights: &mut TrainMamba3Weights,
dims: &Mamba3Dims,
batch_size: usize,
_input_dim: usize,
) {
let dm = dims.d_model;
let seq_len = dims.seq_len;
let t_per = seq_len * dm;
let _epoch = M3_BWD_EPOCH.fetch_add(1, Ordering::Relaxed) + 1;
for b in 0..batch_size {
let d_temporal = &mut d_temporal_out[b * t_per..(b + 1) * t_per];
for (layer_idx, lw) in weights.layers.iter().enumerate().rev() {
let mut scratch = Mamba3Scratch::zeros(dims);
backward_mamba3_layer_batched(
d_temporal,
&batch_acts[b][layer_idx],
lw,
&mut d_weights.layers[layer_idx],
&mut scratch,
dims,
None, );
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::mamba3_siso::config::Mamba3Config;
fn test_cfg() -> Mamba3Config {
Mamba3Config {
d_model: 16,
d_state: 4,
expand: 2,
headdim: 4,
ngroups: 1,
n_layers: 2,
rope_fraction: 0.5,
a_floor: 0.0625,
is_outproj_norm: false,
}
}
#[test]
fn test_parallel_forward_no_panic() {
let cfg = test_cfg();
let dims = Mamba3Dims::from_config(&cfg, 4);
let batch = 3;
let mut w = TrainMamba3Weights::zeros(&dims, 16);
for l in &mut w.layers {
for v in &mut l.norm_weight {
*v = 1.0;
}
for v in &mut l.d_param {
*v = 1.0;
}
for v in &mut l.b_norm_weight {
*v = 1.0;
}
for v in &mut l.c_norm_weight {
*v = 1.0;
}
}
let nh = dims.nheads;
let hd = dims.headdim;
let ds = dims.d_state;
let nl = dims.n_layers;
let na = dims.num_rope_angles.max(1);
let mut temporal = vec![1.0_f32; batch * dims.seq_len * dims.d_model];
let mut acts: Vec<Vec<Mamba3LayerFlat>> = (0..batch)
.map(|_| (0..nl).map(|_| Mamba3LayerFlat::zeros(dims)).collect())
.collect();
let mut ssm = vec![0.0; batch * nl * nh * hd * ds];
let mut k = vec![0.0; batch * nl * nh * ds];
let mut v = vec![0.0; batch * nl * nh * hd];
let mut angle = vec![0.0; batch * nl * nh * na];
parallel_mamba3_forward(
&mut temporal,
&mut acts,
&mut ssm,
&mut k,
&mut v,
&mut angle,
&w,
&dims,
batch,
);
assert!(temporal.iter().all(|v| v.is_finite()));
}
}