use crate::config::MambaConfig;
#[derive(Clone)]
pub struct MambaLayerWeights {
pub norm_weight: Vec<f32>,
pub in_proj_w: Vec<f32>,
pub conv1d_weight: Vec<f32>,
pub conv1d_bias: Vec<f32>,
pub x_proj_w: Vec<f32>,
pub dt_proj_w: Vec<f32>,
pub dt_proj_b: Vec<f32>,
pub a_log: Vec<f32>,
pub a_neg: Vec<f32>,
pub d_param: Vec<f32>,
pub out_proj_w: Vec<f32>,
}
impl MambaLayerWeights {
pub fn compute_a_neg(&mut self) {
for i in 0..self.a_log.len() {
self.a_neg[i] = -self.a_log[i].exp();
}
}
}
#[derive(Clone)]
pub struct MambaWeights {
pub input_proj_w: Vec<f32>,
pub input_proj_b: Vec<f32>,
pub layers: Vec<MambaLayerWeights>,
pub norm_f_weight: Vec<f32>,
}
impl MambaWeights {
pub fn zeros(cfg: &MambaConfig, input_dim: usize) -> Self {
let d = cfg.d_model;
let di = cfg.d_inner();
let ds = cfg.d_state;
let dc = cfg.d_conv;
let dr = cfg.dt_rank();
let xd = cfg.xdbl_dim();
Self {
input_proj_w: vec![0.0; input_dim * d],
input_proj_b: vec![0.0; d],
layers: (0..cfg.n_layers)
.map(|_| MambaLayerWeights {
norm_weight: vec![1.0; d], in_proj_w: vec![0.0; d * 2 * di],
conv1d_weight: vec![0.0; di * dc],
conv1d_bias: vec![0.0; di],
x_proj_w: vec![0.0; di * xd],
dt_proj_w: vec![0.0; dr * di],
dt_proj_b: vec![0.0; di],
a_log: vec![0.0; di * ds],
a_neg: vec![0.0; di * ds], d_param: vec![1.0; di], out_proj_w: vec![0.0; di * d],
})
.collect(),
norm_f_weight: vec![1.0; d], }
}
pub fn init(cfg: &MambaConfig, input_dim: usize, seed: u64) -> Self {
let mut w = Self::zeros(cfg, input_dim);
let mut rng = SimpleRng::new(seed);
let d = cfg.d_model;
let di = cfg.d_inner();
let ds = cfg.d_state;
let dc = cfg.d_conv;
let dr = cfg.dt_rank();
kaiming_uniform(&mut w.input_proj_w, input_dim, &mut rng);
for lw in &mut w.layers {
kaiming_uniform(&mut lw.in_proj_w, d, &mut rng);
kaiming_uniform(&mut lw.conv1d_weight, dc, &mut rng);
kaiming_uniform(&mut lw.x_proj_w, di, &mut rng);
kaiming_uniform(&mut lw.dt_proj_w, dr, &mut rng);
let log_dt_min = 0.001_f32.ln();
let log_dt_max = 0.1_f32.ln();
for b in &mut lw.dt_proj_b {
let dt = (rng.next_f32() * (log_dt_max - log_dt_min) + log_dt_min).exp();
*b = inv_softplus(dt);
}
for d_idx in 0..di {
for n in 0..ds {
lw.a_log[d_idx * ds + n] = ((n + 1) as f32).ln();
}
}
kaiming_uniform(&mut lw.out_proj_w, di, &mut rng);
lw.compute_a_neg();
}
w
}
pub fn validate(&self, cfg: &MambaConfig, input_dim: usize) -> Result<(), String> {
let d = cfg.d_model;
let di = cfg.d_inner();
let ds = cfg.d_state;
let dc = cfg.d_conv;
let dr = cfg.dt_rank();
let xd = cfg.xdbl_dim();
let check = |name: &str, actual: usize, expected: usize| -> Result<(), String> {
if actual != expected {
return Err(format!("{name}: expected {expected}, got {actual}"));
}
Ok(())
};
check("input_proj_w", self.input_proj_w.len(), input_dim * d)?;
check("input_proj_b", self.input_proj_b.len(), d)?;
check("norm_f_weight", self.norm_f_weight.len(), d)?;
if self.layers.len() != cfg.n_layers {
return Err(format!(
"expected {} layers, got {}",
cfg.n_layers,
self.layers.len()
));
}
for (i, lw) in self.layers.iter().enumerate() {
let p = |n: &str| format!("layer[{i}].{n}");
check(&p("norm_weight"), lw.norm_weight.len(), d)?;
check(&p("in_proj_w"), lw.in_proj_w.len(), d * 2 * di)?;
check(&p("conv1d_weight"), lw.conv1d_weight.len(), di * dc)?;
check(&p("conv1d_bias"), lw.conv1d_bias.len(), di)?;
check(&p("x_proj_w"), lw.x_proj_w.len(), di * xd)?;
check(&p("dt_proj_w"), lw.dt_proj_w.len(), dr * di)?;
check(&p("dt_proj_b"), lw.dt_proj_b.len(), di)?;
check(&p("a_log"), lw.a_log.len(), di * ds)?;
check(&p("d_param"), lw.d_param.len(), di)?;
check(&p("out_proj_w"), lw.out_proj_w.len(), di * d)?;
}
Ok(())
}
pub fn param_count(&self, input_dim: usize, cfg: &MambaConfig) -> usize {
let d = cfg.d_model;
let di = cfg.d_inner();
let ds = cfg.d_state;
let dc = cfg.d_conv;
let dr = cfg.dt_rank();
let xd = cfg.xdbl_dim();
let per_layer =
d + d * 2 * di + di * dc + di + di * xd + dr * di + di + di * ds + di + di * d;
input_dim * d + d + cfg.n_layers * per_layer + d
}
}
struct SimpleRng(u64);
impl SimpleRng {
fn new(seed: u64) -> Self {
Self(seed)
}
fn next_u64(&mut self) -> u64 {
self.0 = self
.0
.wrapping_mul(6364136223846793005)
.wrapping_add(1442695040888963407);
self.0
}
fn next_f32(&mut self) -> f32 {
(self.next_u64() >> 40) as f32 / (1u64 << 24) as f32
}
}
fn kaiming_uniform(w: &mut [f32], fan_in: usize, rng: &mut SimpleRng) {
let bound = (3.0 / fan_in as f32).sqrt();
for v in w.iter_mut() {
*v = -bound + 2.0 * bound * rng.next_f32();
}
}
fn inv_softplus(y: f32) -> f32 {
if y > 20.0 { y } else { (y.exp() - 1.0).ln() }
}