use crate::config::MambaConfig;
#[derive(Debug, Clone, Copy, PartialEq)]
pub struct MambaDims {
pub d_model: usize,
pub d_inner: usize,
pub d_state: usize,
pub d_conv: usize,
pub dt_rank: usize,
pub xdbl_dim: usize,
pub seq_len: usize,
pub mamba_input_dim: usize,
pub n_layers: usize,
pub rms_norm_eps: f32,
}
impl MambaDims {
pub fn new(
(d_model, d_inner, d_state, d_conv, dt_rank, seq_len, mamba_input_dim, n_layers): (
usize,
usize,
usize,
usize,
usize,
usize,
usize,
usize,
),
) -> Self {
Self {
d_model,
d_inner,
d_state,
d_conv,
dt_rank,
xdbl_dim: dt_rank + 2 * d_state,
seq_len,
mamba_input_dim,
n_layers,
rms_norm_eps: crate::ops::fast_math::RMS_NORM_EPS,
}
}
pub fn from_config(config: &MambaConfig, seq_len: usize, input_dim: usize) -> Self {
let mut dims = Self::new((
config.d_model,
config.d_inner(),
config.d_state,
config.d_conv,
config.dt_rank(),
seq_len,
input_dim,
config.n_layers,
));
dims.rms_norm_eps = config.rms_norm_eps;
dims
}
}
pub struct MambaRecurrentState<'a> {
pub conv: &'a mut [f32],
pub ssm: &'a mut [f32],
pub a_neg: &'a [f32],
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_mamba_dims_defaults() {
let dims = MambaDims::new((128, 256, 16, 4, 8, 33, 346, 3));
assert_eq!(dims.d_model, 128);
assert_eq!(dims.d_inner, 256);
assert_eq!(dims.d_state, 16);
assert_eq!(dims.d_conv, 4);
assert_eq!(dims.dt_rank, 8);
assert_eq!(dims.xdbl_dim, 40); assert_eq!(dims.seq_len, 33);
assert_eq!(dims.mamba_input_dim, 346);
assert_eq!(dims.n_layers, 3);
}
#[test]
fn test_mamba_dims_from_config() {
let config = MambaConfig::default();
let dims = MambaDims::from_config(&config, 33, 346);
assert_eq!(dims.d_model, 128);
assert_eq!(dims.d_inner, 256); assert_eq!(dims.d_state, 16);
assert_eq!(dims.d_conv, 4);
assert_eq!(dims.dt_rank, 8); assert_eq!(dims.xdbl_dim, 40); assert_eq!(dims.seq_len, 33);
assert_eq!(dims.mamba_input_dim, 346);
assert_eq!(dims.n_layers, 3);
}
}