use super::kernels::Mamba3Kernels;
use crate::mamba_ssm::gpu::buffers::GpuBuffer;
use crate::mamba_ssm::gpu::context::GpuCtx;
use std::sync::Arc;
pub(crate) type CUptr = cudarc::driver::sys::CUdeviceptr;
pub struct GpuMamba3StateBufs<'a> {
pub ssm: &'a mut GpuBuffer,
pub k: &'a mut GpuBuffer,
pub v: &'a mut GpuBuffer,
pub angle: &'a mut GpuBuffer,
}
impl GpuMamba3StateBufs<'_> {
pub fn reborrow(&mut self) -> GpuMamba3StateBufs<'_> {
GpuMamba3StateBufs {
ssm: self.ssm,
k: self.k,
v: self.v,
angle: self.angle,
}
}
}
#[derive(Clone, Copy)]
pub struct M3Exec<'a> {
pub ctx: &'a GpuCtx,
pub kernels: &'a Mamba3Kernels,
pub dims: &'a GpuMamba3Dims,
}
#[derive(Debug, Clone, Copy)]
pub struct GpuMamba3Dims {
pub batch: usize,
pub d_model: usize,
pub d_inner: usize,
pub d_state: usize,
pub nheads: usize,
pub headdim: usize,
pub ngroups: usize,
pub in_proj_dim: usize,
pub seq_len: usize,
pub mamba_input_dim: usize,
pub n_layers: usize,
pub n_angles: usize,
pub a_floor: f32,
pub is_outproj_norm: bool,
pub use_parallel_scan: bool,
}
impl GpuMamba3Dims {
pub fn bt(&self) -> usize {
self.batch * self.seq_len
}
pub fn chunk_size(&self) -> usize {
CHUNK_SIZE
}
pub fn n_chunks(&self) -> usize {
self.seq_len.div_ceil(self.chunk_size())
}
}
pub const CHUNK_SIZE: usize = 64;
pub struct GpuMamba3LayerActs {
pub residual: GpuBuffer, pub rms_vals: GpuBuffer, pub post_norm: GpuBuffer, pub z: GpuBuffer, pub x: GpuBuffer, pub b_raw: GpuBuffer, pub c_raw: GpuBuffer, pub dd_dt_raw: GpuBuffer, pub dd_a_raw: GpuBuffer, pub trap_raw: GpuBuffer, pub dt: GpuBuffer, pub a_val: GpuBuffer, pub trap: GpuBuffer, pub angles_raw: GpuBuffer, pub b_normed: GpuBuffer, pub c_normed: GpuBuffer, pub b_rms: GpuBuffer, pub c_rms: GpuBuffer, pub b_biased: GpuBuffer, pub c_biased: GpuBuffer, pub k: GpuBuffer, pub q: GpuBuffer, pub angle_cumsum: GpuBuffer, pub alpha: GpuBuffer, pub beta: GpuBuffer, pub gamma: GpuBuffer, pub h_saved: GpuBuffer, pub k_prev_saved: GpuBuffer, pub v_prev_saved: GpuBuffer, pub y: GpuBuffer, pub da_cumsum_saved: GpuBuffer, pub k_scaled_saved: GpuBuffer, pub scale_saved: GpuBuffer, pub gamma_saved: GpuBuffer, pub qk_dot_saved: GpuBuffer, pub chunk_states_saved: GpuBuffer, pub gated_rms_vals: GpuBuffer, pub gated: GpuBuffer, }
pub struct GpuMamba3BackboneActs {
pub input_proj_inputs: GpuBuffer, pub input_proj_outputs: GpuBuffer, pub layers: Vec<GpuMamba3LayerActs>,
pub norm_f_input: GpuBuffer, pub norm_f_rms: GpuBuffer, }
pub struct GpuMamba3Scratch {
pub proj_flat: GpuBuffer, pub out_flat: GpuBuffer,
pub d_gated: GpuBuffer, pub d_y: GpuBuffer, pub d_z: GpuBuffer, pub d_norm_gate_w: GpuBuffer, pub d_x: GpuBuffer, pub d_k: GpuBuffer, pub d_q: GpuBuffer, pub d_proj: GpuBuffer, pub d_norm: GpuBuffer, pub d_pre_norm: GpuBuffer,
pub d_d_local: GpuBuffer, pub d_alpha: GpuBuffer, pub d_beta: GpuBuffer, pub d_gamma: GpuBuffer,
pub d_dt_angle: GpuBuffer, pub d_angles_raw: GpuBuffer, pub d_dd_dt: GpuBuffer, pub d_dd_a: GpuBuffer, pub d_trap_raw: GpuBuffer,
pub d_b_pre_rope: GpuBuffer, pub d_c_pre_rope: GpuBuffer, pub d_angle_cumsum: GpuBuffer, pub d_b_normed: GpuBuffer, pub d_c_normed: GpuBuffer, pub d_b_raw: GpuBuffer, pub d_c_raw: GpuBuffer, pub d_b_norm_w: GpuBuffer, pub d_c_norm_w: GpuBuffer,
pub d_input_proj_dx: GpuBuffer,
pub da_cumsum: GpuBuffer, pub da_cs_sum: GpuBuffer, pub chunk_states: GpuBuffer, pub final_states: GpuBuffer, pub d_da_cumsum: GpuBuffer, pub d_prev_states: GpuBuffer, pub d_scale: GpuBuffer, pub d_gamma_par: GpuBuffer, pub d_qk_dot: GpuBuffer,
pub axis0_partials: GpuBuffer,
}
#[derive(Clone, Copy)]
pub struct Mamba3LayerPtrs {
pub ssm_state: CUptr, pub k_state: CUptr, pub v_state: CUptr, pub angle_state: CUptr, }
pub struct GpuMamba3TargetScratch {
pub proj_flat: GpuBuffer, pub z: GpuBuffer, pub x: GpuBuffer, pub b_raw: GpuBuffer, pub c_raw: GpuBuffer, pub dd_dt_raw: GpuBuffer, pub dd_a_raw: GpuBuffer, pub trap_raw: GpuBuffer, pub dt: GpuBuffer, pub a_val: GpuBuffer, pub trap: GpuBuffer, pub angles_raw: GpuBuffer, pub b_normed: GpuBuffer, pub c_normed: GpuBuffer, pub b_rms: GpuBuffer, pub c_rms: GpuBuffer, pub b_biased: GpuBuffer, pub c_biased: GpuBuffer, pub k: GpuBuffer, pub q: GpuBuffer, pub angle_cumsum: GpuBuffer, pub alpha: GpuBuffer, pub beta: GpuBuffer, pub gamma: GpuBuffer, pub y: GpuBuffer, pub gated: GpuBuffer, pub temporal_work: GpuBuffer, pub out_flat: GpuBuffer, pub residual: GpuBuffer, pub rms_discard: GpuBuffer, pub ssm_states: GpuBuffer, pub k_states: GpuBuffer, pub v_states: GpuBuffer, pub angle_states: GpuBuffer, }
impl GpuMamba3LayerActs {
pub fn new(
stream: &Arc<cudarc::driver::CudaStream>,
bt: usize,
dims: &GpuMamba3Dims,
) -> Result<Self, String> {
let dm = dims.d_model;
let di = dims.d_inner;
let ds = dims.d_state;
let nh = dims.nheads;
let hd = dims.headdim;
let ng = dims.ngroups;
let b = dims.batch;
let t = dims.seq_len;
let na = dims.n_angles;
let nc = dims.n_chunks();
let cs = dims.chunk_size();
Ok(Self {
residual: GpuBuffer::zeros(stream, bt * dm)?,
rms_vals: GpuBuffer::zeros(stream, bt)?,
post_norm: GpuBuffer::zeros(stream, bt * dm)?,
z: GpuBuffer::zeros(stream, bt * di)?,
x: GpuBuffer::zeros(stream, bt * di)?,
b_raw: GpuBuffer::zeros(stream, bt * ng * ds)?,
c_raw: GpuBuffer::zeros(stream, bt * ng * ds)?,
dd_dt_raw: GpuBuffer::zeros(stream, bt * nh)?,
dd_a_raw: GpuBuffer::zeros(stream, bt * nh)?,
trap_raw: GpuBuffer::zeros(stream, bt * nh)?,
dt: GpuBuffer::zeros(stream, bt * nh)?,
a_val: GpuBuffer::zeros(stream, bt * nh)?,
trap: GpuBuffer::zeros(stream, bt * nh)?,
angles_raw: GpuBuffer::zeros(stream, bt * na.max(1))?,
b_normed: GpuBuffer::zeros(stream, bt * ng * ds)?,
c_normed: GpuBuffer::zeros(stream, bt * ng * ds)?,
b_rms: GpuBuffer::zeros(stream, bt * ng)?,
c_rms: GpuBuffer::zeros(stream, bt * ng)?,
b_biased: GpuBuffer::zeros(stream, bt * nh * ds)?,
c_biased: GpuBuffer::zeros(stream, bt * nh * ds)?,
k: GpuBuffer::zeros(stream, bt * nh * ds)?,
q: GpuBuffer::zeros(stream, bt * nh * ds)?,
angle_cumsum: GpuBuffer::zeros(stream, bt * nh * na.max(1))?,
alpha: GpuBuffer::zeros(stream, bt * nh)?,
beta: GpuBuffer::zeros(stream, bt * nh)?,
gamma: GpuBuffer::zeros(stream, bt * nh)?,
h_saved: GpuBuffer::zeros(stream, b * (t + 1) * di * ds)?,
k_prev_saved: GpuBuffer::zeros(stream, bt * nh * ds)?,
v_prev_saved: GpuBuffer::zeros(stream, bt * nh * hd)?,
y: GpuBuffer::zeros(stream, bt * di)?,
da_cumsum_saved: GpuBuffer::zeros(stream, b * nc * nh * cs)?,
k_scaled_saved: GpuBuffer::zeros(stream, bt * nh * ds)?,
scale_saved: GpuBuffer::zeros(stream, bt * nh)?,
gamma_saved: GpuBuffer::zeros(stream, bt * nh)?,
qk_dot_saved: GpuBuffer::zeros(stream, bt * nh)?,
chunk_states_saved: GpuBuffer::zeros(stream, b * nc * nh * hd * ds)?,
gated_rms_vals: GpuBuffer::zeros(stream, bt * nh)?,
gated: GpuBuffer::zeros(stream, bt * di)?,
})
}
}
impl GpuMamba3BackboneActs {
pub fn new(
stream: &Arc<cudarc::driver::CudaStream>,
dims: &GpuMamba3Dims,
) -> Result<Self, String> {
let bt = dims.bt();
Ok(Self {
input_proj_inputs: GpuBuffer::zeros(stream, bt * dims.mamba_input_dim)?,
input_proj_outputs: GpuBuffer::zeros(stream, bt * dims.d_model)?,
layers: (0..dims.n_layers)
.map(|_| GpuMamba3LayerActs::new(stream, bt, dims))
.collect::<Result<_, _>>()?,
norm_f_input: GpuBuffer::zeros(stream, bt * dims.d_model)?,
norm_f_rms: GpuBuffer::zeros(stream, bt)?,
})
}
}
impl GpuMamba3Scratch {
pub fn new(
stream: &Arc<cudarc::driver::CudaStream>,
dims: &GpuMamba3Dims,
) -> Result<Self, String> {
let bt = dims.bt();
let di = dims.d_inner;
let dm = dims.d_model;
let ds = dims.d_state;
let nh = dims.nheads;
let hd = dims.headdim;
let ip = dims.in_proj_dim;
let b = dims.batch;
Ok(Self {
proj_flat: GpuBuffer::zeros(stream, bt * ip)?,
out_flat: GpuBuffer::zeros(stream, bt * dm)?,
d_gated: GpuBuffer::zeros(stream, bt * di)?,
d_y: GpuBuffer::zeros(stream, bt * di)?,
d_z: GpuBuffer::zeros(stream, bt * di)?,
d_norm_gate_w: GpuBuffer::zeros(stream, bt * di)?,
d_x: GpuBuffer::zeros(stream, bt * di)?,
d_k: GpuBuffer::zeros(stream, bt * nh * ds)?,
d_q: GpuBuffer::zeros(stream, bt * nh * ds)?,
d_proj: GpuBuffer::zeros(stream, bt * ip)?,
d_norm: GpuBuffer::zeros(stream, bt * dm)?,
d_pre_norm: GpuBuffer::zeros(stream, bt * dm)?,
d_d_local: GpuBuffer::zeros(stream, b * di)?,
d_alpha: GpuBuffer::zeros(stream, bt * nh)?,
d_beta: GpuBuffer::zeros(stream, bt * nh)?,
d_gamma: GpuBuffer::zeros(stream, bt * nh)?,
d_dt_angle: GpuBuffer::zeros(stream, bt * nh)?,
d_angles_raw: GpuBuffer::zeros(stream, bt * dims.n_angles.max(1))?,
d_dd_dt: GpuBuffer::zeros(stream, bt * nh)?,
d_dd_a: GpuBuffer::zeros(stream, bt * nh)?,
d_trap_raw: GpuBuffer::zeros(stream, bt * nh)?,
d_b_pre_rope: GpuBuffer::zeros(stream, bt * nh * ds)?,
d_c_pre_rope: GpuBuffer::zeros(stream, bt * nh * ds)?,
d_angle_cumsum: GpuBuffer::zeros(stream, bt * nh * dims.n_angles.max(1))?,
d_b_normed: GpuBuffer::zeros(stream, bt * dims.ngroups * ds)?,
d_c_normed: GpuBuffer::zeros(stream, bt * dims.ngroups * ds)?,
d_b_raw: GpuBuffer::zeros(stream, bt * dims.ngroups * ds)?,
d_c_raw: GpuBuffer::zeros(stream, bt * dims.ngroups * ds)?,
d_b_norm_w: GpuBuffer::zeros(stream, bt * dims.ngroups * ds)?,
d_c_norm_w: GpuBuffer::zeros(stream, bt * dims.ngroups * ds)?,
d_input_proj_dx: GpuBuffer::zeros(stream, bt * dims.mamba_input_dim)?,
da_cumsum: {
let cs = dims.chunk_size();
let nc = dims.n_chunks();
GpuBuffer::zeros(stream, b * nc * nh * cs)?
},
da_cs_sum: {
let nc = dims.n_chunks();
GpuBuffer::zeros(stream, b * nh * nc)?
},
chunk_states: {
let nc = dims.n_chunks();
GpuBuffer::zeros(stream, b * nc * nh * hd * ds)?
},
final_states: GpuBuffer::zeros(stream, b * nh * hd * ds)?,
d_da_cumsum: {
let cs = dims.chunk_size();
let nc = dims.n_chunks();
GpuBuffer::zeros(stream, b * nc * nh * cs)?
},
d_prev_states: {
let nc = dims.n_chunks();
GpuBuffer::zeros(stream, b * nc * nh * hd * ds)?
},
d_scale: GpuBuffer::zeros(stream, bt * nh)?,
d_gamma_par: GpuBuffer::zeros(stream, bt * nh)?,
d_qk_dot: GpuBuffer::zeros(stream, bt * nh)?,
axis0_partials: {
let na = dims.n_angles.max(1);
let angle_dt_sz = 2 * nh * bt * na;
let rmsnorm_sz = bt * dm;
let d_d_sz = b * nh;
GpuBuffer::zeros(stream, angle_dt_sz.max(rmsnorm_sz).max(d_d_sz))?
},
})
}
}
impl GpuMamba3TargetScratch {
pub fn new(
stream: &Arc<cudarc::driver::CudaStream>,
dims: &GpuMamba3Dims,
) -> Result<Self, String> {
let bt = dims.bt();
let di = dims.d_inner;
let dm = dims.d_model;
let ds = dims.d_state;
let nh = dims.nheads;
let hd = dims.headdim;
let ng = dims.ngroups;
let ip = dims.in_proj_dim;
let b = dims.batch;
let nl = dims.n_layers;
let na = dims.n_angles.max(1);
Ok(Self {
proj_flat: GpuBuffer::zeros(stream, bt * ip)?,
z: GpuBuffer::zeros(stream, bt * di)?,
x: GpuBuffer::zeros(stream, bt * di)?,
b_raw: GpuBuffer::zeros(stream, bt * ng * ds)?,
c_raw: GpuBuffer::zeros(stream, bt * ng * ds)?,
dd_dt_raw: GpuBuffer::zeros(stream, bt * nh)?,
dd_a_raw: GpuBuffer::zeros(stream, bt * nh)?,
trap_raw: GpuBuffer::zeros(stream, bt * nh)?,
dt: GpuBuffer::zeros(stream, bt * nh)?,
a_val: GpuBuffer::zeros(stream, bt * nh)?,
trap: GpuBuffer::zeros(stream, bt * nh)?,
angles_raw: GpuBuffer::zeros(stream, bt * na)?,
b_normed: GpuBuffer::zeros(stream, bt * ng * ds)?,
c_normed: GpuBuffer::zeros(stream, bt * ng * ds)?,
b_rms: GpuBuffer::zeros(stream, bt * ng)?,
c_rms: GpuBuffer::zeros(stream, bt * ng)?,
b_biased: GpuBuffer::zeros(stream, bt * nh * ds)?,
c_biased: GpuBuffer::zeros(stream, bt * nh * ds)?,
k: GpuBuffer::zeros(stream, bt * nh * ds)?,
q: GpuBuffer::zeros(stream, bt * nh * ds)?,
angle_cumsum: GpuBuffer::zeros(stream, bt * nh * na.max(1))?,
alpha: GpuBuffer::zeros(stream, bt * nh)?,
beta: GpuBuffer::zeros(stream, bt * nh)?,
gamma: GpuBuffer::zeros(stream, bt * nh)?,
y: GpuBuffer::zeros(stream, bt * di)?,
gated: GpuBuffer::zeros(stream, bt * di)?,
temporal_work: GpuBuffer::zeros(stream, bt * dm)?,
out_flat: GpuBuffer::zeros(stream, bt * dm)?,
residual: GpuBuffer::zeros(stream, bt * dm)?,
rms_discard: GpuBuffer::zeros(stream, bt * nh)?,
ssm_states: GpuBuffer::zeros(stream, b * nl * nh * hd * ds)?,
k_states: GpuBuffer::zeros(stream, b * nl * nh * ds)?,
v_states: GpuBuffer::zeros(stream, b * nl * nh * hd)?,
angle_states: GpuBuffer::zeros(stream, b * nl * nh * na.max(1))?,
})
}
}