use crate::config::MambaConfig;
use crate::inference::{MambaStepScratch, mamba_step};
use crate::state::MambaState;
use crate::weights::{MambaLayerWeights, MambaWeights};
pub struct MambaBackbone {
weights: MambaWeights,
cfg: MambaConfig,
input_dim: usize,
}
impl MambaBackbone {
pub fn init(cfg: MambaConfig, input_dim: usize, seed: u64) -> Self {
let weights = MambaWeights::init(&cfg, input_dim, seed);
Self {
weights,
cfg,
input_dim,
}
}
pub fn from_weights(cfg: MambaConfig, weights: MambaWeights) -> Result<Self, String> {
let input_dim = weights.input_proj_w.len() / cfg.d_model;
weights.validate(&cfg, input_dim)?;
Ok(Self {
weights,
cfg,
input_dim,
})
}
pub fn into_weights(self) -> MambaWeights {
self.weights
}
pub fn weights(&self) -> &MambaWeights {
&self.weights
}
pub fn weights_mut(&mut self) -> &mut MambaWeights {
&mut self.weights
}
pub fn layer(&self, index: usize) -> &MambaLayerWeights {
&self.weights.layers[index]
}
pub fn layer_mut(&mut self, index: usize) -> &mut MambaLayerWeights {
&mut self.weights.layers[index]
}
pub fn n_layers(&self) -> usize {
self.cfg.n_layers
}
pub fn param_count(&self) -> usize {
self.weights.param_count(self.input_dim, &self.cfg)
}
pub fn config(&self) -> &MambaConfig {
&self.cfg
}
pub fn input_dim(&self) -> usize {
self.input_dim
}
pub fn forward_step(
&self,
input: &[f32],
output: &mut [f32],
state: &mut MambaState,
scratch: &mut MambaStepScratch,
) {
mamba_step(
input,
output,
&self.weights,
&mut state.layers,
scratch,
&self.cfg,
self.input_dim,
);
}
pub fn forward_sequence(
&self,
inputs: &[f32],
outputs: &mut [f32],
state: &mut MambaState,
scratch: &mut MambaStepScratch,
seq_len: usize,
) {
let dm = self.cfg.d_model;
debug_assert_eq!(inputs.len(), seq_len * self.input_dim);
debug_assert_eq!(outputs.len(), seq_len * dm);
for t in 0..seq_len {
let inp = &inputs[t * self.input_dim..(t + 1) * self.input_dim];
let out = &mut outputs[t * dm..(t + 1) * dm];
self.forward_step(inp, out, state, scratch);
}
}
pub fn forward_step_batch(
&self,
inputs: &[f32],
outputs: &mut [f32],
states: &mut [MambaState],
scratches: &mut [MambaStepScratch],
) {
crate::inference::mamba_step_batch(
inputs,
outputs,
&self.weights,
states,
scratches,
&self.cfg,
self.input_dim,
);
}
pub fn alloc_state(&self) -> MambaState {
MambaState::zeros(
self.cfg.n_layers,
self.cfg.d_inner(),
self.cfg.d_state,
self.cfg.d_conv,
)
}
pub fn alloc_scratch(&self) -> MambaStepScratch {
MambaStepScratch::new(&self.cfg)
}
}