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 {
const DT_INIT_FLOOR: f32 = 1e-4;
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();
linear_default_uniform(&mut w.input_proj_w, input_dim, &mut rng);
let residual_rescale = 1.0 / (cfg.n_layers as f64).sqrt() as f32;
for lw in &mut w.layers {
linear_default_uniform(&mut lw.in_proj_w, d, &mut rng);
linear_default_uniform(&mut lw.conv1d_weight, dc, &mut rng);
linear_default_uniform(&mut lw.conv1d_bias, dc, &mut rng);
linear_default_uniform(&mut lw.x_proj_w, di, &mut rng);
linear_default_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()
.max(DT_INIT_FLOOR);
*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();
}
}
linear_default_uniform(&mut lw.out_proj_w, di, &mut rng);
for v in &mut lw.out_proj_w {
*v *= residual_rescale;
}
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 linear_default_uniform(w: &mut [f32], fan_in: usize, rng: &mut SimpleRng) {
let bound = (1.0 / fan_in as f64).sqrt() as f32;
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_m1().ln() }
}
#[cfg(test)]
mod init_tests {
use super::*;
fn cfg(n_layers: usize) -> MambaConfig {
MambaConfig {
d_model: 64,
d_state: 16,
d_conv: 4,
expand: 2,
n_layers,
..MambaConfig::default()
}
}
fn bound_of(w: &[f32]) -> f32 {
w.iter().fold(0.0_f32, |m, v| m.max(v.abs()))
}
#[test]
fn linear_weights_draw_the_pytorch_default_bound() {
let c = cfg(2);
let w = MambaWeights::init(&c, c.d_model, 11);
let lw = &w.layers[0];
let expect = |fan_in: usize| 1.0 / (fan_in as f32).sqrt();
for (name, buf, fan_in) in [
("in_proj", &lw.in_proj_w, c.d_model),
("x_proj", &lw.x_proj_w, c.d_inner()),
("dt_proj", &lw.dt_proj_w, c.dt_rank()),
("conv1d", &lw.conv1d_weight, c.d_conv),
] {
let b = bound_of(buf);
let e = expect(fan_in);
assert!(
b <= e && b > 0.80 * e,
"{name}: max |w| = {b}, expected just under {e}"
);
}
}
#[test]
fn out_proj_carries_the_residual_rescale() {
let c = cfg(16);
let w = MambaWeights::init(&c, c.d_model, 12);
let b = bound_of(&w.layers[0].out_proj_w);
let e = (1.0 / (c.d_inner() as f32).sqrt()) / (c.n_layers as f32).sqrt();
assert!(
b <= e && b > 0.80 * e,
"out_proj: max |w| = {b}, expected just under {e}"
);
}
#[test]
fn conv1d_bias_is_drawn_not_zeroed() {
let c = cfg(2);
let w = MambaWeights::init(&c, c.d_model, 13);
let bias = &w.layers[0].conv1d_bias;
assert!(
bias.iter().any(|v| *v != 0.0),
"conv1d bias must not be all zeros"
);
let b = bound_of(bias);
let e = 1.0 / (c.d_conv as f32).sqrt();
assert!(b <= e, "conv1d bias: max |b| = {b}, bound {e}");
}
#[test]
fn a_log_is_the_s4d_real_ladder() {
let c = cfg(1);
let w = MambaWeights::init(&c, c.d_model, 14);
let lw = &w.layers[0];
for d in 0..c.d_inner() {
for n in 0..c.d_state {
let got = lw.a_log[d * c.d_state + n];
let want = ((n + 1) as f32).ln();
assert!(
(got - want).abs() < 1e-6,
"A_log[{d},{n}] = {got}, want {want}"
);
}
}
}
}