use crate::ops::dims::MambaDims;
#[derive(Debug, Clone, Copy)]
pub struct FieldOffsets {
pub residual: usize,
pub rms_val: usize,
pub post_norm: usize,
pub x_branch: usize,
pub conv_state: usize,
pub post_conv: usize,
pub u: usize,
pub xdbl: usize,
pub delta_raw: usize,
pub delta: usize,
pub h_prev: usize,
pub h_curr: usize,
pub da_exp: usize,
pub y: usize,
pub gate_pre_silu: usize,
pub gate_post_silu: usize,
pub gated: usize,
pub step_stride: usize,
}
impl FieldOffsets {
pub fn new(dims: &MambaDims) -> Self {
let dm = dims.d_model;
let di = dims.d_inner;
let ds = dims.d_state;
let dc = dims.d_conv;
let xdbl = dims.xdbl_dim;
let mut off = 0usize;
let residual = off;
off += dm;
let rms_val = off;
off += 1;
let post_norm = off;
off += dm;
let x_branch = off;
off += di;
let conv_state = off;
off += di * dc;
let post_conv = off;
off += di;
let u = off;
off += di;
let xdbl_off = off;
off += xdbl;
let delta_raw = off;
off += di;
let delta = off;
off += di;
let h_prev = off;
off += di * ds;
let h_curr = off;
off += di * ds;
let da_exp = off;
off += di * ds;
let y = off;
off += di;
let gate_pre_silu = off;
off += di;
let gate_post_silu = off;
off += di;
let gated = off;
off += di;
Self {
residual,
rms_val,
post_norm,
x_branch,
conv_state,
post_conv,
u,
xdbl: xdbl_off,
delta_raw,
delta,
h_prev,
h_curr,
da_exp,
y,
gate_pre_silu,
gate_post_silu,
gated,
step_stride: off,
}
}
}
pub struct MambaLayerFlat {
pub data: Vec<f32>,
pub offsets: FieldOffsets,
pub dims: MambaDims,
}
impl MambaLayerFlat {
pub fn zeros(dims: MambaDims) -> Self {
let offsets = FieldOffsets::new(&dims);
let total = dims.seq_len * offsets.step_stride;
Self {
data: vec![0.0; total],
offsets,
dims,
}
}
#[inline(always)]
fn base(&self, t: usize) -> usize {
t * self.offsets.step_stride
}
#[inline]
pub fn residual(&self, t: usize) -> &[f32] {
let b = self.base(t) + self.offsets.residual;
&self.data[b..b + self.dims.d_model]
}
#[inline]
pub fn rms_val(&self, t: usize) -> f32 {
self.data[self.base(t) + self.offsets.rms_val]
}
#[inline]
pub fn post_norm(&self, t: usize) -> &[f32] {
let b = self.base(t) + self.offsets.post_norm;
&self.data[b..b + self.dims.d_model]
}
#[inline]
pub fn x_branch(&self, t: usize) -> &[f32] {
let b = self.base(t) + self.offsets.x_branch;
&self.data[b..b + self.dims.d_inner]
}
#[inline]
pub fn conv_state(&self, t: usize) -> &[f32] {
let b = self.base(t) + self.offsets.conv_state;
&self.data[b..b + self.dims.d_inner * self.dims.d_conv]
}
#[inline]
pub fn post_conv(&self, t: usize) -> &[f32] {
let b = self.base(t) + self.offsets.post_conv;
&self.data[b..b + self.dims.d_inner]
}
#[inline]
pub fn u(&self, t: usize) -> &[f32] {
let b = self.base(t) + self.offsets.u;
&self.data[b..b + self.dims.d_inner]
}
#[inline]
pub fn xdbl(&self, t: usize) -> &[f32] {
let b = self.base(t) + self.offsets.xdbl;
&self.data[b..b + self.dims.xdbl_dim]
}
#[inline]
pub fn delta_raw(&self, t: usize) -> &[f32] {
let b = self.base(t) + self.offsets.delta_raw;
&self.data[b..b + self.dims.d_inner]
}
#[inline]
pub fn delta(&self, t: usize) -> &[f32] {
let b = self.base(t) + self.offsets.delta;
&self.data[b..b + self.dims.d_inner]
}
#[inline]
pub fn h_prev(&self, t: usize) -> &[f32] {
let b = self.base(t) + self.offsets.h_prev;
&self.data[b..b + self.dims.d_inner * self.dims.d_state]
}
#[inline]
pub fn h_curr(&self, t: usize) -> &[f32] {
let b = self.base(t) + self.offsets.h_curr;
&self.data[b..b + self.dims.d_inner * self.dims.d_state]
}
#[inline]
pub fn da_exp(&self, t: usize) -> &[f32] {
let b = self.base(t) + self.offsets.da_exp;
&self.data[b..b + self.dims.d_inner * self.dims.d_state]
}
#[inline]
pub fn y(&self, t: usize) -> &[f32] {
let b = self.base(t) + self.offsets.y;
&self.data[b..b + self.dims.d_inner]
}
#[inline]
pub fn gate_pre_silu(&self, t: usize) -> &[f32] {
let b = self.base(t) + self.offsets.gate_pre_silu;
&self.data[b..b + self.dims.d_inner]
}
#[inline]
pub fn gate_post_silu(&self, t: usize) -> &[f32] {
let b = self.base(t) + self.offsets.gate_post_silu;
&self.data[b..b + self.dims.d_inner]
}
#[inline]
pub fn gated(&self, t: usize) -> &[f32] {
let b = self.base(t) + self.offsets.gated;
&self.data[b..b + self.dims.d_inner]
}
#[inline]
pub fn residual_mut(&mut self, t: usize) -> &mut [f32] {
let b = self.base(t) + self.offsets.residual;
let dm = self.dims.d_model;
&mut self.data[b..b + dm]
}
#[inline]
pub fn set_rms_val(&mut self, t: usize, val: f32) {
let idx = self.base(t) + self.offsets.rms_val;
self.data[idx] = val;
}
#[inline]
pub fn post_norm_mut(&mut self, t: usize) -> &mut [f32] {
let b = self.base(t) + self.offsets.post_norm;
let dm = self.dims.d_model;
&mut self.data[b..b + dm]
}
#[inline]
pub fn x_branch_mut(&mut self, t: usize) -> &mut [f32] {
let b = self.base(t) + self.offsets.x_branch;
let di = self.dims.d_inner;
&mut self.data[b..b + di]
}
#[inline]
pub fn conv_state_mut(&mut self, t: usize) -> &mut [f32] {
let b = self.base(t) + self.offsets.conv_state;
let len = self.dims.d_inner * self.dims.d_conv;
&mut self.data[b..b + len]
}
#[inline]
pub fn post_conv_mut(&mut self, t: usize) -> &mut [f32] {
let b = self.base(t) + self.offsets.post_conv;
let di = self.dims.d_inner;
&mut self.data[b..b + di]
}
#[inline]
pub fn u_mut(&mut self, t: usize) -> &mut [f32] {
let b = self.base(t) + self.offsets.u;
let di = self.dims.d_inner;
&mut self.data[b..b + di]
}
#[inline]
pub fn xdbl_mut(&mut self, t: usize) -> &mut [f32] {
let b = self.base(t) + self.offsets.xdbl;
let xd = self.dims.xdbl_dim;
&mut self.data[b..b + xd]
}
#[inline]
pub fn delta_raw_mut(&mut self, t: usize) -> &mut [f32] {
let b = self.base(t) + self.offsets.delta_raw;
let di = self.dims.d_inner;
&mut self.data[b..b + di]
}
#[inline]
pub fn delta_mut(&mut self, t: usize) -> &mut [f32] {
let b = self.base(t) + self.offsets.delta;
let di = self.dims.d_inner;
&mut self.data[b..b + di]
}
#[inline]
pub fn h_prev_mut(&mut self, t: usize) -> &mut [f32] {
let b = self.base(t) + self.offsets.h_prev;
let len = self.dims.d_inner * self.dims.d_state;
&mut self.data[b..b + len]
}
#[inline]
pub fn h_curr_mut(&mut self, t: usize) -> &mut [f32] {
let b = self.base(t) + self.offsets.h_curr;
let len = self.dims.d_inner * self.dims.d_state;
&mut self.data[b..b + len]
}
#[inline]
pub fn da_exp_mut(&mut self, t: usize) -> &mut [f32] {
let b = self.base(t) + self.offsets.da_exp;
let len = self.dims.d_inner * self.dims.d_state;
&mut self.data[b..b + len]
}
#[inline]
pub fn y_mut(&mut self, t: usize) -> &mut [f32] {
let b = self.base(t) + self.offsets.y;
let di = self.dims.d_inner;
&mut self.data[b..b + di]
}
#[inline]
pub fn gate_pre_silu_mut(&mut self, t: usize) -> &mut [f32] {
let b = self.base(t) + self.offsets.gate_pre_silu;
let di = self.dims.d_inner;
&mut self.data[b..b + di]
}
#[inline]
pub fn gate_post_silu_mut(&mut self, t: usize) -> &mut [f32] {
let b = self.base(t) + self.offsets.gate_post_silu;
let di = self.dims.d_inner;
&mut self.data[b..b + di]
}
#[inline]
pub fn gated_mut(&mut self, t: usize) -> &mut [f32] {
let b = self.base(t) + self.offsets.gated;
let di = self.dims.d_inner;
&mut self.data[b..b + di]
}
pub fn copy_post_norm_all(&self, dst: &mut [f32]) {
let dm = self.dims.d_model;
for t in 0..self.dims.seq_len {
let src = self.post_norm(t);
dst[t * dm..(t + 1) * dm].copy_from_slice(src);
}
}
pub fn copy_u_all(&self, dst: &mut [f32]) {
let di = self.dims.d_inner;
for t in 0..self.dims.seq_len {
let src = self.u(t);
dst[t * di..(t + 1) * di].copy_from_slice(src);
}
}
pub fn copy_gated_all(&self, dst: &mut [f32]) {
let di = self.dims.d_inner;
for t in 0..self.dims.seq_len {
let src = self.gated(t);
dst[t * di..(t + 1) * di].copy_from_slice(src);
}
}
pub fn copy_xdbl_dt_all(&self, dst: &mut [f32]) {
let dr = self.dims.dt_rank;
for t in 0..self.dims.seq_len {
let src = self.xdbl(t);
dst[t * dr..(t + 1) * dr].copy_from_slice(&src[..dr]);
}
}
}
pub struct MambaBackboneFlat {
pub input_proj_inputs: Vec<f32>,
pub input_proj_outputs: Vec<f32>,
pub layers: Vec<MambaLayerFlat>,
pub norm_f_input: Vec<f32>,
pub norm_f_rms: Vec<f32>,
}
impl MambaBackboneFlat {
pub fn zeros(dims: MambaDims) -> Self {
Self {
input_proj_inputs: vec![0.0; dims.seq_len * dims.mamba_input_dim],
input_proj_outputs: vec![0.0; dims.seq_len * dims.d_model],
layers: (0..dims.n_layers)
.map(|_| MambaLayerFlat::zeros(dims))
.collect(),
norm_f_input: vec![0.0; dims.seq_len * dims.d_model],
norm_f_rms: vec![0.0; dims.seq_len],
}
}
}
#[cfg(test)]
mod tests {
use super::*;
fn default_dims() -> MambaDims {
MambaDims::new((128, 256, 16, 4, 8, 33, 346, 3))
}
#[test]
fn test_field_offsets_no_overlap() {
let dims = default_dims();
let o = FieldOffsets::new(&dims);
let dm = dims.d_model;
let di = dims.d_inner;
let ds = dims.d_state;
let dc = dims.d_conv;
let xd = dims.xdbl_dim;
let fields: Vec<(usize, usize)> = vec![
(o.residual, dm),
(o.rms_val, 1),
(o.post_norm, dm),
(o.x_branch, di),
(o.conv_state, di * dc),
(o.post_conv, di),
(o.u, di),
(o.xdbl, xd),
(o.delta_raw, di),
(o.delta, di),
(o.h_prev, di * ds),
(o.h_curr, di * ds),
(o.da_exp, di * ds),
(o.y, di),
(o.gate_pre_silu, di),
(o.gate_post_silu, di),
(o.gated, di),
];
for i in 0..fields.len() - 1 {
let end_i = fields[i].0 + fields[i].1;
let start_next = fields[i + 1].0;
assert_eq!(
end_i,
start_next,
"field {} ends at {end_i} but field {} starts at {start_next}",
i,
i + 1,
);
}
let last = fields.last().unwrap();
assert_eq!(
last.0 + last.1,
o.step_stride,
"last field end ({}) != step_stride ({})",
last.0 + last.1,
o.step_stride,
);
assert_eq!(o.step_stride, 15913);
}
#[test]
fn test_field_accessor_roundtrip() {
let dims = default_dims();
let mut layer = MambaLayerFlat::zeros(dims);
for &t in &[0usize, dims.seq_len - 1] {
layer.residual_mut(t).iter_mut().for_each(|v| *v = 1.0);
assert!(layer.residual(t).iter().all(|&v| v == 1.0));
assert_eq!(layer.residual(t).len(), dims.d_model);
layer.set_rms_val(t, 42.0);
assert_eq!(layer.rms_val(t), 42.0);
layer.post_norm_mut(t).iter_mut().for_each(|v| *v = 2.0);
assert!(layer.post_norm(t).iter().all(|&v| v == 2.0));
assert_eq!(layer.post_norm(t).len(), dims.d_model);
layer.x_branch_mut(t).iter_mut().for_each(|v| *v = 3.0);
assert!(layer.x_branch(t).iter().all(|&v| v == 3.0));
assert_eq!(layer.x_branch(t).len(), dims.d_inner);
layer.conv_state_mut(t).iter_mut().for_each(|v| *v = 4.0);
assert!(layer.conv_state(t).iter().all(|&v| v == 4.0));
assert_eq!(layer.conv_state(t).len(), dims.d_inner * dims.d_conv);
layer.post_conv_mut(t).iter_mut().for_each(|v| *v = 5.0);
assert!(layer.post_conv(t).iter().all(|&v| v == 5.0));
assert_eq!(layer.post_conv(t).len(), dims.d_inner);
layer.u_mut(t).iter_mut().for_each(|v| *v = 6.0);
assert!(layer.u(t).iter().all(|&v| v == 6.0));
assert_eq!(layer.u(t).len(), dims.d_inner);
layer.xdbl_mut(t).iter_mut().for_each(|v| *v = 7.0);
assert!(layer.xdbl(t).iter().all(|&v| v == 7.0));
assert_eq!(layer.xdbl(t).len(), dims.xdbl_dim);
layer.delta_raw_mut(t).iter_mut().for_each(|v| *v = 8.0);
assert!(layer.delta_raw(t).iter().all(|&v| v == 8.0));
assert_eq!(layer.delta_raw(t).len(), dims.d_inner);
layer.delta_mut(t).iter_mut().for_each(|v| *v = 9.0);
assert!(layer.delta(t).iter().all(|&v| v == 9.0));
assert_eq!(layer.delta(t).len(), dims.d_inner);
layer.h_prev_mut(t).iter_mut().for_each(|v| *v = 10.0);
assert!(layer.h_prev(t).iter().all(|&v| v == 10.0));
assert_eq!(layer.h_prev(t).len(), dims.d_inner * dims.d_state);
layer.h_curr_mut(t).iter_mut().for_each(|v| *v = 11.0);
assert!(layer.h_curr(t).iter().all(|&v| v == 11.0));
assert_eq!(layer.h_curr(t).len(), dims.d_inner * dims.d_state);
layer.da_exp_mut(t).iter_mut().for_each(|v| *v = 12.0);
assert!(layer.da_exp(t).iter().all(|&v| v == 12.0));
assert_eq!(layer.da_exp(t).len(), dims.d_inner * dims.d_state);
layer.y_mut(t).iter_mut().for_each(|v| *v = 13.0);
assert!(layer.y(t).iter().all(|&v| v == 13.0));
assert_eq!(layer.y(t).len(), dims.d_inner);
layer
.gate_pre_silu_mut(t)
.iter_mut()
.for_each(|v| *v = 14.0);
assert!(layer.gate_pre_silu(t).iter().all(|&v| v == 14.0));
assert_eq!(layer.gate_pre_silu(t).len(), dims.d_inner);
layer
.gate_post_silu_mut(t)
.iter_mut()
.for_each(|v| *v = 15.0);
assert!(layer.gate_post_silu(t).iter().all(|&v| v == 15.0));
assert_eq!(layer.gate_post_silu(t).len(), dims.d_inner);
layer.gated_mut(t).iter_mut().for_each(|v| *v = 16.0);
assert!(layer.gated(t).iter().all(|&v| v == 16.0));
assert_eq!(layer.gated(t).len(), dims.d_inner);
}
if dims.seq_len > 2 {
assert!(layer.residual(1).iter().all(|&v| v == 0.0));
assert_eq!(layer.rms_val(1), 0.0);
assert!(layer.gated(1).iter().all(|&v| v == 0.0));
}
}
#[test]
fn test_bulk_copy_post_norm_all() {
let dims = MambaDims::new((128, 256, 16, 4, 8, 4, 346, 3));
let mut layer = MambaLayerFlat::zeros(dims);
for t in 0..dims.seq_len {
let val = (t + 1) as f32;
layer.post_norm_mut(t).iter_mut().for_each(|v| *v = val);
}
let mut dst = vec![0.0f32; dims.seq_len * dims.d_model];
layer.copy_post_norm_all(&mut dst);
for t in 0..dims.seq_len {
let expected = (t + 1) as f32;
let chunk = &dst[t * dims.d_model..(t + 1) * dims.d_model];
assert!(
chunk.iter().all(|&v| v == expected),
"post_norm_all mismatch at t={t}",
);
}
}
#[test]
fn test_bulk_copy_u_all() {
let dims = MambaDims::new((128, 256, 16, 4, 8, 4, 346, 3));
let mut layer = MambaLayerFlat::zeros(dims);
for t in 0..dims.seq_len {
let val = (t + 10) as f32;
layer.u_mut(t).iter_mut().for_each(|v| *v = val);
}
let mut dst = vec![0.0f32; dims.seq_len * dims.d_inner];
layer.copy_u_all(&mut dst);
for t in 0..dims.seq_len {
let expected = (t + 10) as f32;
let chunk = &dst[t * dims.d_inner..(t + 1) * dims.d_inner];
assert!(
chunk.iter().all(|&v| v == expected),
"u_all mismatch at t={t}",
);
}
}
#[test]
fn test_bulk_copy_gated_all() {
let dims = MambaDims::new((128, 256, 16, 4, 8, 4, 346, 3));
let mut layer = MambaLayerFlat::zeros(dims);
for t in 0..dims.seq_len {
let val = (t + 20) as f32;
layer.gated_mut(t).iter_mut().for_each(|v| *v = val);
}
let mut dst = vec![0.0f32; dims.seq_len * dims.d_inner];
layer.copy_gated_all(&mut dst);
for t in 0..dims.seq_len {
let expected = (t + 20) as f32;
let chunk = &dst[t * dims.d_inner..(t + 1) * dims.d_inner];
assert!(
chunk.iter().all(|&v| v == expected),
"gated_all mismatch at t={t}",
);
}
}
#[test]
fn test_bulk_copy_xdbl_dt_all() {
let dims = MambaDims::new((128, 256, 16, 4, 8, 4, 346, 3));
let mut layer = MambaLayerFlat::zeros(dims);
for t in 0..dims.seq_len {
let xdbl = layer.xdbl_mut(t);
for (i, v) in xdbl.iter_mut().enumerate() {
if i < dims.dt_rank {
*v = (t * 100 + i) as f32;
} else {
*v = -1.0;
}
}
}
let mut dst = vec![0.0f32; dims.seq_len * dims.dt_rank];
layer.copy_xdbl_dt_all(&mut dst);
for t in 0..dims.seq_len {
for i in 0..dims.dt_rank {
let expected = (t * 100 + i) as f32;
assert_eq!(
dst[t * dims.dt_rank + i],
expected,
"xdbl_dt_all mismatch at t={t}, i={i}",
);
}
}
}
#[test]
fn test_backbone_flat_allocation() {
let dims = default_dims();
let backbone = MambaBackboneFlat::zeros(dims);
assert_eq!(
backbone.input_proj_inputs.len(),
dims.seq_len * dims.mamba_input_dim
);
assert_eq!(
backbone.input_proj_outputs.len(),
dims.seq_len * dims.d_model
);
assert_eq!(backbone.layers.len(), dims.n_layers);
for layer in &backbone.layers {
let expected_len = dims.seq_len * layer.offsets.step_stride;
assert_eq!(layer.data.len(), expected_len);
}
}
}