use super::config::Mamba3Config;
#[derive(Clone)]
pub struct Mamba3LayerState {
pub ssm_state: Vec<f32>,
pub k_state: Vec<f32>,
pub v_state: Vec<f32>,
pub angle_state: Vec<f32>,
}
impl Mamba3LayerState {
pub fn zeros(nheads: usize, headdim: usize, d_state: usize, num_rope_angles: usize) -> Self {
Self {
ssm_state: vec![0.0; nheads * headdim * d_state],
k_state: vec![0.0; nheads * d_state],
v_state: vec![0.0; nheads * headdim],
angle_state: vec![0.0; nheads * num_rope_angles.max(1)],
}
}
pub fn reset(&mut self) {
self.ssm_state.fill(0.0);
self.k_state.fill(0.0);
self.v_state.fill(0.0);
self.angle_state.fill(0.0);
}
}
#[derive(Clone)]
pub struct Mamba3State {
pub layers: Vec<Mamba3LayerState>,
}
impl Mamba3State {
pub fn zeros(cfg: &Mamba3Config) -> Self {
let nh = cfg.nheads();
let hd = cfg.headdim;
let ds = cfg.d_state;
let na = cfg.num_rope_angles();
Self {
layers: (0..cfg.n_layers)
.map(|_| Mamba3LayerState::zeros(nh, hd, ds, na))
.collect(),
}
}
pub fn reset(&mut self) {
for layer in &mut self.layers {
layer.reset();
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_state_zeros() {
let cfg = Mamba3Config::default();
let state = Mamba3State::zeros(&cfg);
assert_eq!(state.layers.len(), cfg.n_layers);
let l = &state.layers[0];
assert_eq!(l.ssm_state.len(), cfg.nheads() * cfg.headdim * cfg.d_state);
assert_eq!(l.k_state.len(), cfg.nheads() * cfg.d_state);
assert_eq!(l.v_state.len(), cfg.nheads() * cfg.headdim);
assert_eq!(
l.angle_state.len(),
cfg.nheads() * cfg.num_rope_angles().max(1)
);
assert!(l.ssm_state.iter().all(|&v| v == 0.0));
}
#[test]
fn test_state_reset() {
let cfg = Mamba3Config::default();
let mut state = Mamba3State::zeros(&cfg);
state.layers[0].ssm_state[0] = 42.0;
state.layers[0].angle_state[0] = std::f32::consts::PI;
state.reset();
assert_eq!(state.layers[0].ssm_state[0], 0.0);
assert_eq!(state.layers[0].angle_state[0], 0.0);
}
}