use crate::ops::dims::MambaDims;
#[derive(Clone)]
pub struct TrainMambaLayerWeights {
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 d_param: Vec<f32>,
pub out_proj_w: Vec<f32>,
}
#[derive(Clone)]
pub struct TrainMambaWeights {
pub input_proj_w: Vec<f32>,
pub input_proj_b: Vec<f32>,
pub layers: Vec<TrainMambaLayerWeights>,
pub norm_f_weight: Vec<f32>,
}
impl TrainMambaWeights {
pub fn zeros_from_dims(dims: &MambaDims) -> Self {
let dm = dims.d_model;
let di = dims.d_inner;
let ds = dims.d_state;
let dc = dims.d_conv;
let dr = dims.dt_rank;
let xdbl = dims.xdbl_dim;
let mid = dims.mamba_input_dim;
Self {
input_proj_w: vec![0.0; mid * dm],
input_proj_b: vec![0.0; dm],
layers: (0..dims.n_layers)
.map(|_| TrainMambaLayerWeights {
norm_weight: vec![0.0; dm],
in_proj_w: vec![0.0; dm * 2 * di],
conv1d_weight: vec![0.0; di * dc],
conv1d_bias: vec![0.0; di],
x_proj_w: vec![0.0; di * xdbl],
dt_proj_w: vec![0.0; dr * di],
dt_proj_b: vec![0.0; di],
a_log: vec![0.0; di * ds],
d_param: vec![0.0; di],
out_proj_w: vec![0.0; di * dm],
})
.collect(),
norm_f_weight: vec![0.0; dm],
}
}
pub fn zero(&mut self) {
self.input_proj_w.fill(0.0);
self.input_proj_b.fill(0.0);
for l in &mut self.layers {
l.norm_weight.fill(0.0);
l.in_proj_w.fill(0.0);
l.conv1d_weight.fill(0.0);
l.conv1d_bias.fill(0.0);
l.x_proj_w.fill(0.0);
l.dt_proj_w.fill(0.0);
l.dt_proj_b.fill(0.0);
l.a_log.fill(0.0);
l.d_param.fill(0.0);
l.out_proj_w.fill(0.0);
}
self.norm_f_weight.fill(0.0);
}
pub fn add_inplace(&mut self, other: &Self) {
add_vecs(&mut self.input_proj_w, &other.input_proj_w);
add_vecs(&mut self.input_proj_b, &other.input_proj_b);
for (sl, ol) in self.layers.iter_mut().zip(other.layers.iter()) {
add_vecs(&mut sl.norm_weight, &ol.norm_weight);
add_vecs(&mut sl.in_proj_w, &ol.in_proj_w);
add_vecs(&mut sl.conv1d_weight, &ol.conv1d_weight);
add_vecs(&mut sl.conv1d_bias, &ol.conv1d_bias);
add_vecs(&mut sl.x_proj_w, &ol.x_proj_w);
add_vecs(&mut sl.dt_proj_w, &ol.dt_proj_w);
add_vecs(&mut sl.dt_proj_b, &ol.dt_proj_b);
add_vecs(&mut sl.a_log, &ol.a_log);
add_vecs(&mut sl.d_param, &ol.d_param);
add_vecs(&mut sl.out_proj_w, &ol.out_proj_w);
}
add_vecs(&mut self.norm_f_weight, &other.norm_f_weight);
}
}
#[inline]
fn add_vecs(a: &mut [f32], b: &[f32]) {
for (ai, &bi) in a.iter_mut().zip(b.iter()) {
*ai += bi;
}
}