use super::kernels::Mamba3Kernels;
use crate::mamba_ssm::gpu::buffers::{GpuBuffer, GpuByteBuffer};
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, PartialEq)]
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 rms_norm_eps: f32,
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 fn validate_index_budget(&self) -> Result<(), String> {
let b = self.batch;
let t = self.seq_len;
let di = self.d_inner;
let ds = self.d_state;
let nh = self.nheads;
let candidates: [(&str, usize); 6] = [
(
"sequential tape B*(T+1)*d_inner*d_state",
b * (t + 1) * di * ds,
),
("chunk-state tape B*n_chunks*nheads*headdim*d_state", {
b * self.n_chunks() * nh * self.headdim * ds
}),
("da cumsum B*n_chunks*nheads*chunk_size", {
b * self.n_chunks() * nh * CHUNK_SIZE
}),
("activation B*T*in_proj_dim", b * t * self.in_proj_dim),
("activation B*T*d_inner", b * t * di),
("bc lanes B*T*nheads*d_state", b * t * nh * ds),
];
for (name, len) in candidates {
if len > i32::MAX as usize {
return Err(format!(
"index budget overflow: {name} = {len} exceeds i32::MAX — the CUDA \
kernels index this range in 32-bit arithmetic and would silently \
corrupt; shrink batch/seq_len or split the run"
));
}
}
Ok(())
}
}
pub const CHUNK_SIZE: usize = 64;
#[cfg(test)]
mod index_budget_tests {
use super::GpuMamba3Dims;
fn dims(batch: usize, seq_len: usize) -> GpuMamba3Dims {
GpuMamba3Dims {
batch,
d_model: 384,
d_inner: 768,
d_state: 64,
nheads: 48,
headdim: 16,
ngroups: 1,
in_proj_dim: 1700,
seq_len,
mamba_input_dim: 384,
n_layers: 24,
n_angles: 16,
a_floor: 1e-4,
is_outproj_norm: false,
rms_norm_eps: 1e-5,
use_parallel_scan: true,
}
}
#[test]
fn index_budget_guards_long_context() {
dims(2, 4621).validate_index_budget().unwrap();
let err = dims(2, 25_000_000).validate_index_budget().unwrap_err();
assert!(err.contains("index budget overflow"), "{err}");
}
}
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 angle_chunk_sums: GpuByteBuffer,
pub angle_chunk_carries: GpuByteBuffer,
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,
if dims.use_parallel_scan {
1
} else {
b * (t + 1) * di * ds
},
)?,
k_prev_saved: GpuBuffer::zeros(
stream,
if dims.use_parallel_scan {
1
} else {
bt * nh * ds
},
)?,
v_prev_saved: GpuBuffer::zeros(
stream,
if dims.use_parallel_scan {
1
} else {
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;
let angle_stage_bytes =
b * dims.n_chunks() * nh * dims.n_angles.max(1) * std::mem::size_of::<f64>();
Ok(Self {
proj_flat: GpuBuffer::zeros(stream, bt * ip)?,
out_flat: GpuBuffer::zeros(stream, bt * dm)?,
angle_chunk_sums: GpuByteBuffer::zeros(stream, angle_stage_bytes)?,
angle_chunk_carries: GpuByteBuffer::zeros(stream, angle_stage_bytes)?,
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))?,
})
}
}