use crate::ops::dims::MambaDims;
pub struct PhaseScratch {
pub post_norm_flat: Vec<f32>,
pub proj_flat: Vec<f32>,
pub gate_silu_flat: Vec<f32>,
pub gated_flat: Vec<f32>,
pub out_flat: Vec<f32>,
pub da_buf: Vec<f32>,
pub u_flat: Vec<f32>,
pub xdbl_flat: Vec<f32>,
pub dt_in_flat: Vec<f32>,
pub delta_raw_flat: Vec<f32>,
}
impl PhaseScratch {
pub fn zeros(dims: &MambaDims) -> Self {
let t = dims.seq_len;
Self {
post_norm_flat: vec![0.0; t * dims.d_model],
proj_flat: vec![0.0; t * 2 * dims.d_inner],
gate_silu_flat: vec![0.0; t * dims.d_inner],
gated_flat: vec![0.0; t * dims.d_inner],
out_flat: vec![0.0; t * dims.d_model],
da_buf: vec![0.0; dims.d_state],
u_flat: vec![0.0; t * dims.d_inner],
xdbl_flat: vec![0.0; t * (dims.dt_rank + 2 * dims.d_state)],
dt_in_flat: vec![0.0; t * dims.dt_rank],
delta_raw_flat: vec![0.0; t * dims.d_inner],
}
}
}
pub struct BackwardPhaseScratch {
pub d_gated_flat: Vec<f32>,
pub d_y_flat: Vec<f32>,
pub d_gate_flat: Vec<f32>,
pub d_delta_flat: Vec<f32>,
pub d_delta_raw_flat: Vec<f32>,
pub d_u_flat: Vec<f32>,
pub d_u_xproj_flat: Vec<f32>,
pub d_xdbl_flat: Vec<f32>,
pub d_conv_out_flat: Vec<f32>,
pub d_x_branch_flat: Vec<f32>,
pub d_proj_flat: Vec<f32>,
pub d_norm_flat: Vec<f32>,
pub d_pre_norm_flat: Vec<f32>,
pub d_dt_input_flat: Vec<f32>,
pub d_h: Vec<f32>,
pub d_conv_carry: Vec<f32>,
pub xdbl_dt_buf: Vec<f32>,
pub u_buf: Vec<f32>,
pub gated_buf: Vec<f32>,
pub post_norm_buf: Vec<f32>,
pub d_input_proj_scratch: Vec<f32>,
pub d_norm_f_weight_local: Vec<f32>,
}
impl BackwardPhaseScratch {
pub fn zeros(dims: &MambaDims) -> Self {
let t = dims.seq_len;
let di = dims.d_inner;
let dm = dims.d_model;
Self {
d_gated_flat: vec![0.0; t * di],
d_y_flat: vec![0.0; t * di],
d_gate_flat: vec![0.0; t * di],
d_delta_flat: vec![0.0; t * di],
d_delta_raw_flat: vec![0.0; t * di],
d_u_flat: vec![0.0; t * di],
d_u_xproj_flat: vec![0.0; t * di],
d_xdbl_flat: vec![0.0; t * dims.xdbl_dim],
d_conv_out_flat: vec![0.0; t * di],
d_x_branch_flat: vec![0.0; t * di],
d_proj_flat: vec![0.0; t * 2 * di],
d_norm_flat: vec![0.0; t * dm],
d_pre_norm_flat: vec![0.0; t * dm],
d_dt_input_flat: vec![0.0; t * dims.dt_rank],
d_h: vec![0.0; di * dims.d_state],
d_conv_carry: vec![0.0; di * (dims.d_conv - 1)],
xdbl_dt_buf: vec![0.0; t * dims.dt_rank],
u_buf: vec![0.0; t * di],
gated_buf: vec![0.0; t * di],
post_norm_buf: vec![0.0; t * dm],
d_input_proj_scratch: vec![0.0; t * dims.mamba_input_dim.max(dm)],
d_norm_f_weight_local: vec![0.0; dm],
}
}
}