#[derive(Clone)]
pub struct MambaLayerState {
pub conv_state: Vec<f32>,
pub ssm_state: Vec<f32>,
}
impl MambaLayerState {
pub fn zeros(d_inner: usize, d_state: usize, d_conv: usize) -> Self {
assert!(d_conv > 0, "d_conv must be > 0");
Self {
conv_state: vec![0.0; (d_conv - 1) * d_inner],
ssm_state: vec![0.0; d_inner * d_state],
}
}
pub fn reset(&mut self) {
self.conv_state.fill(0.0);
self.ssm_state.fill(0.0);
}
}
#[derive(Clone)]
pub struct MambaState {
pub layers: Vec<MambaLayerState>,
}
impl MambaState {
pub fn zeros(n_layers: usize, d_inner: usize, d_state: usize, d_conv: usize) -> Self {
Self {
layers: (0..n_layers)
.map(|_| MambaLayerState::zeros(d_inner, d_state, d_conv))
.collect(),
}
}
pub fn reset(&mut self) {
for layer in &mut self.layers {
layer.reset();
}
}
}