use std::sync::Arc;
use cudarc::driver::{CudaStream, PushKernelArg};
use crate::mamba_ssm::gpu::blas::{TypedPtr, gpu_gemm_typed_forward_raw};
use crate::mamba_ssm::gpu::buffers::{DtypedBuf, GpuBuffer};
use crate::mamba_ssm::gpu::dtype::WeightDtype;
use crate::mamba_ssm::gpu::launch::{grid_1d, grid_norm};
use crate::mamba3_siso::config::Mamba3Config;
use crate::mamba3_siso::gpu::state::{CHUNK_SIZE, GpuMamba3StateBufs, M3Exec, Mamba3LayerPtrs};
use crate::mamba3_siso::gpu::weights_mixed_train::GpuMamba3TrainMixedWeights;
pub struct GpuMamba3LayerMixedActs {
pub residual: GpuBuffer,
pub rms_vals: GpuBuffer,
pub post_norm: DtypedBuf,
pub z: DtypedBuf,
pub x: DtypedBuf,
pub b_raw: DtypedBuf,
pub c_raw: DtypedBuf,
pub dd_dt_raw: GpuBuffer,
pub dd_a_raw: GpuBuffer,
pub trap_raw: GpuBuffer,
pub angles_raw: GpuBuffer,
pub dt: GpuBuffer,
pub a_val: GpuBuffer,
pub trap: GpuBuffer,
pub b_normed: DtypedBuf,
pub c_normed: DtypedBuf,
pub b_rms: GpuBuffer,
pub c_rms: GpuBuffer,
pub b_biased: DtypedBuf,
pub c_biased: DtypedBuf,
pub k: DtypedBuf,
pub q: DtypedBuf,
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: DtypedBuf,
pub da_cumsum_saved: GpuBuffer,
pub k_scaled_saved: DtypedBuf,
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: DtypedBuf,
}
pub struct GpuMamba3BackboneMixedActs {
pub input_proj_inputs: DtypedBuf,
pub input_proj_outputs: DtypedBuf,
pub layers: Vec<GpuMamba3LayerMixedActs>,
pub norm_f_input: GpuBuffer,
pub norm_f_rms: GpuBuffer,
pub dtype: WeightDtype,
}
impl GpuMamba3BackboneMixedActs {
pub fn new(
stream: &Arc<CudaStream>,
cfg: &Mamba3Config,
batch: usize,
seq_len: usize,
input_dim: usize,
dtype: WeightDtype,
) -> Result<Self, String> {
let dm = cfg.d_model;
let di = cfg.d_inner();
let ds = cfg.d_state;
let nh = cfg.nheads();
let ng = cfg.ngroups;
let na = cfg.num_rope_angles();
let hd = cfg.headdim;
let bt = batch * seq_len;
let nc_max = seq_len.div_ceil(CHUNK_SIZE);
let da_cs_len = batch * nc_max * nh * CHUNK_SIZE;
let layers = (0..cfg.n_layers)
.map(|_| {
Ok(GpuMamba3LayerMixedActs {
residual: GpuBuffer::zeros(stream, bt * dm)?,
rms_vals: GpuBuffer::zeros(stream, bt)?,
post_norm: DtypedBuf::zeros(stream, bt * dm, dtype)?,
z: DtypedBuf::zeros(stream, bt * di, dtype)?,
x: DtypedBuf::zeros(stream, bt * di, dtype)?,
b_raw: DtypedBuf::zeros(stream, bt * ng * ds, dtype)?,
c_raw: DtypedBuf::zeros(stream, bt * ng * ds, dtype)?,
dd_dt_raw: GpuBuffer::zeros(stream, bt * nh)?,
dd_a_raw: GpuBuffer::zeros(stream, bt * nh)?,
trap_raw: GpuBuffer::zeros(stream, bt * nh)?,
angles_raw: GpuBuffer::zeros(stream, bt * na.max(1))?,
dt: GpuBuffer::zeros(stream, bt * nh)?,
a_val: GpuBuffer::zeros(stream, bt * nh)?,
trap: GpuBuffer::zeros(stream, bt * nh)?,
b_normed: DtypedBuf::zeros(stream, bt * ng * ds, dtype)?,
c_normed: DtypedBuf::zeros(stream, bt * ng * ds, dtype)?,
b_rms: GpuBuffer::zeros(stream, bt * ng)?,
c_rms: GpuBuffer::zeros(stream, bt * ng)?,
b_biased: DtypedBuf::zeros(stream, bt * nh * ds, dtype)?,
c_biased: DtypedBuf::zeros(stream, bt * nh * ds, dtype)?,
k: DtypedBuf::zeros(stream, bt * nh * ds, dtype)?,
q: DtypedBuf::zeros(stream, bt * nh * ds, dtype)?,
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, batch * (seq_len + 1) * di * ds)?,
k_prev_saved: GpuBuffer::zeros(stream, bt * nh * ds)?,
v_prev_saved: GpuBuffer::zeros(stream, bt * nh * hd)?,
y: DtypedBuf::zeros(stream, bt * di, dtype)?,
da_cumsum_saved: GpuBuffer::zeros(stream, da_cs_len)?,
k_scaled_saved: DtypedBuf::zeros(stream, bt * nh * ds, dtype)?,
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, batch * nc_max * nh * hd * ds)?,
gated_rms_vals: GpuBuffer::zeros(stream, bt * nh)?,
gated: DtypedBuf::zeros(stream, bt * di, dtype)?,
})
})
.collect::<Result<Vec<_>, String>>()?;
let s = Self {
input_proj_inputs: DtypedBuf::zeros(stream, bt * input_dim, dtype)?,
input_proj_outputs: DtypedBuf::zeros(stream, bt * dm, dtype)?,
layers,
norm_f_input: GpuBuffer::zeros(stream, bt * dm)?,
norm_f_rms: GpuBuffer::zeros(stream, bt)?,
dtype,
};
stream
.synchronize()
.map_err(|e| format!("sync after m3 mixed acts alloc: {e:?}"))?;
Ok(s)
}
}
pub fn gpu_forward_mamba3_layer_mixed(
exec: &M3Exec<'_>,
temporal_f32: &mut GpuBuffer,
acts: &mut GpuMamba3LayerMixedActs,
w: &crate::mamba3_siso::gpu::weights::GpuMamba3MixedLayerWeights,
layer_ptrs: &Mamba3LayerPtrs,
scratch: &mut GpuMamba3MixedScratch,
dtype: WeightDtype,
) -> Result<(), String> {
let M3Exec {
ctx,
kernels: m3k,
dims,
} = *exec;
let Mamba3LayerPtrs {
ssm_state,
k_state,
v_state,
angle_state,
} = *layer_ptrs;
let GpuMamba3MixedScratch {
proj_flat: proj_flat_scratch,
out_flat: out_flat_scratch,
alpha: alpha_scratch,
beta: beta_scratch,
gamma: gamma_scratch,
adt_temp: adt_temp_scratch,
chunk_states: chunk_states_scratch,
final_states: final_states_scratch,
..
} = scratch;
let bt = dims.bt();
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 ip = dims.in_proj_dim;
let na = dims.n_angles;
acts.residual.copy_from(temporal_f32, &ctx.stream)?;
{
let bt_i = bt as i32;
let dm_i = dm as i32;
let eps: f32 = 1e-5;
let mut bld = ctx
.stream
.launch_builder(m3k.rmsnorm_fwd_f32in_typed.get(dtype));
let pn = acts.post_norm.cached_ptr();
let rms = acts.rms_vals.cached_ptr();
let x = acts.residual.cached_ptr();
let nw = w.norm_weight.ptr();
bld.arg(&pn);
bld.arg(&rms);
bld.arg(&x);
bld.arg(&nw);
bld.arg(&bt_i);
bld.arg(&dm_i);
bld.arg(&eps);
unsafe { bld.launch(grid_norm(bt, dm)) }
.map_err(|e| format!("m3_mixed F1 rmsnorm: {e:?}"))?;
}
gpu_gemm_typed_forward_raw(
ctx,
TypedPtr {
ptr: proj_flat_scratch.cached_ptr(),
dtype,
},
TypedPtr {
ptr: acts.post_norm.cached_ptr(),
dtype,
},
TypedPtr {
ptr: w.in_proj_w.ptr(),
dtype,
},
None,
(bt, dm, ip),
)?;
{
let n_i = bt as i32;
let di_i = di as i32;
let ng_i = ng as i32;
let ds_i = ds as i32;
let nh_i = nh as i32;
let na_i = na as i32;
let db_ptr = w.dt_bias.ptr();
let mut bld = ctx.stream.launch_builder(m3k.m3_split_typed.get(dtype));
let z = acts.z.cached_ptr();
let x = acts.x.cached_ptr();
let br = acts.b_raw.cached_ptr();
let cr = acts.c_raw.cached_ptr();
let dt = acts.dt.cached_ptr();
let av = acts.a_val.cached_ptr();
let tr = acts.trap.cached_ptr();
let ar = acts.angles_raw.cached_ptr();
let dd_dt = acts.dd_dt_raw.cached_ptr();
let dd_a = acts.dd_a_raw.cached_ptr();
let tr_r = acts.trap_raw.cached_ptr();
let proj = proj_flat_scratch.cached_ptr();
bld.arg(&z);
bld.arg(&x);
bld.arg(&br);
bld.arg(&cr);
bld.arg(&dt);
bld.arg(&av);
bld.arg(&tr);
bld.arg(&ar);
bld.arg(&dd_dt);
bld.arg(&dd_a);
bld.arg(&tr_r);
bld.arg(&proj);
bld.arg(&db_ptr);
bld.arg(&dims.a_floor);
bld.arg(&n_i);
bld.arg(&di_i);
bld.arg(&ng_i);
bld.arg(&ds_i);
bld.arg(&nh_i);
bld.arg(&na_i);
unsafe { bld.launch(grid_1d(bt * ip)) }.map_err(|e| format!("m3_mixed F3 split: {e:?}"))?;
}
{
let n_i = bt as i32;
let ng_i = ng as i32;
let ds_i = ds as i32;
let cfg = cudarc::driver::LaunchConfig {
grid_dim: ((bt * ng) as u32, 2, 1),
block_dim: (ds as u32, 1, 1),
shared_mem_bytes: ds as u32 * 4,
};
let mut bld = ctx
.stream
.launch_builder(m3k.bcnorm_fwd_bc_typed.get(dtype));
let bn = acts.b_normed.cached_ptr();
let cn = acts.c_normed.cached_ptr();
let brms = acts.b_rms.cached_ptr();
let crms = acts.c_rms.cached_ptr();
let br = acts.b_raw.cached_ptr();
let cr = acts.c_raw.cached_ptr();
let bnw = w.b_norm_weight.ptr();
let cnw = w.c_norm_weight.ptr();
bld.arg(&bn);
bld.arg(&cn);
bld.arg(&brms);
bld.arg(&crms);
bld.arg(&br);
bld.arg(&cr);
bld.arg(&bnw);
bld.arg(&cnw);
bld.arg(&n_i);
bld.arg(&ng_i);
bld.arg(&ds_i);
unsafe { bld.launch(cfg) }.map_err(|e| format!("m3_mixed F4ab bcnorm_bc: {e:?}"))?;
}
{
let n_i = bt as i32;
let nh_i = nh as i32;
let ng_i = ng as i32;
let ds_i = ds as i32;
let cfg = cudarc::driver::LaunchConfig {
grid_dim: ((bt * nh * ds).div_ceil(256) as u32, 2, 1),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let mut bld = ctx
.stream
.launch_builder(m3k.bc_bias_add_bc_typed.get(dtype));
let bb = acts.b_biased.cached_ptr();
let cb = acts.c_biased.cached_ptr();
let bn = acts.b_normed.cached_ptr();
let cn = acts.c_normed.cached_ptr();
let bb_p = w.b_bias.ptr();
let cb_p = w.c_bias.ptr();
bld.arg(&bb);
bld.arg(&cb);
bld.arg(&bn);
bld.arg(&cn);
bld.arg(&bb_p);
bld.arg(&cb_p);
bld.arg(&n_i);
bld.arg(&nh_i);
bld.arg(&ng_i);
bld.arg(&ds_i);
unsafe { bld.launch(cfg) }.map_err(|e| format!("m3_mixed F4cd bias: {e:?}"))?;
}
if na > 0 {
let b_i = (bt / dims.seq_len) as i32;
let t_i = dims.seq_len as i32;
let nh_i = nh as i32;
let na_i = na as i32;
let mut bld = ctx.stream.launch_builder(&m3k.m3_angle_dt_fwd_seq);
let ac = acts.angle_cumsum.cached_ptr();
let ar = acts.angles_raw.cached_ptr();
let dt = acts.dt.cached_ptr();
bld.arg(&ac);
bld.arg(&angle_state);
bld.arg(&ar);
bld.arg(&dt);
bld.arg(&b_i);
bld.arg(&t_i);
bld.arg(&nh_i);
bld.arg(&na_i);
let grid = cudarc::driver::LaunchConfig {
grid_dim: (
(bt / dims.seq_len) as u32,
(nh * na).div_ceil(256) as u32,
1,
),
block_dim: (256.min((nh * na) as u32), 1, 1),
shared_mem_bytes: 0,
};
unsafe { bld.launch(grid) }.map_err(|e| format!("m3_mixed F5 angle_dt: {e:?}"))?;
}
if na > 0 {
let n_i = bt as i32;
let nh_i = nh as i32;
let ds_i = ds as i32;
let na_i = na as i32;
let mut bld = ctx.stream.launch_builder(m3k.rope_fwd_typed.get(dtype));
let k = acts.k.cached_ptr();
let q = acts.q.cached_ptr();
let bb = acts.b_biased.cached_ptr();
let cb = acts.c_biased.cached_ptr();
let ac = acts.angle_cumsum.cached_ptr();
bld.arg(&k);
bld.arg(&q);
bld.arg(&bb);
bld.arg(&cb);
bld.arg(&ac);
bld.arg(&n_i);
bld.arg(&nh_i);
bld.arg(&ds_i);
bld.arg(&na_i);
unsafe { bld.launch(grid_1d(bt * nh * ds)) }
.map_err(|e| format!("m3_mixed F4ef rope: {e:?}"))?;
} else {
let bytes = bt * nh * ds * dtype.size_bytes();
let stream = ctx.stream.cu_stream();
unsafe {
let res = cudarc::driver::sys::cuMemcpyDtoDAsync_v2(
acts.k.cached_ptr(),
acts.b_biased.cached_ptr(),
bytes,
stream,
);
if res != cudarc::driver::sys::CUresult::CUDA_SUCCESS {
return Err(format!("m3_mixed F4ef k D2D copy: {res:?}"));
}
let res = cudarc::driver::sys::cuMemcpyDtoDAsync_v2(
acts.q.cached_ptr(),
acts.c_biased.cached_ptr(),
bytes,
stream,
);
if res != cudarc::driver::sys::CUresult::CUDA_SUCCESS {
return Err(format!("m3_mixed F4ef q D2D copy: {res:?}"));
}
}
}
{
let n_total = (bt * nh) as i32;
let mut bld = ctx.stream.launch_builder(&m3k.m3_compute_abg);
bld.arg(alpha_scratch.inner_mut());
bld.arg(beta_scratch.inner_mut());
bld.arg(gamma_scratch.inner_mut());
bld.arg(acts.dt.inner());
bld.arg(acts.a_val.inner());
bld.arg(acts.trap.inner());
bld.arg(&n_total);
unsafe { bld.launch(grid_1d(bt * nh)) }.map_err(|e| format!("m3_mixed F5b abg: {e:?}"))?;
}
acts.alpha.copy_from(alpha_scratch, &ctx.stream)?;
acts.beta.copy_from(beta_scratch, &ctx.stream)?;
acts.gamma.copy_from(gamma_scratch, &ctx.stream)?;
if dims.use_parallel_scan {
let dp_ptr = w.d_param.ptr();
let b_i = dims.batch as i32;
let t_i = dims.seq_len as i32;
let nh_i = nh as i32;
let hd_i = hd as i32;
let ds_i = ds as i32;
let cs = dims.chunk_size() as i32;
let nc = dims.n_chunks();
{
let n_total = (bt * nh) as i32;
let mut bld = ctx.stream.launch_builder(&m3k.elementwise_mul);
bld.arg(adt_temp_scratch.inner_mut());
bld.arg(acts.a_val.inner());
bld.arg(acts.dt.inner());
bld.arg(&n_total);
unsafe { bld.launch(grid_1d(bt * nh)) }
.map_err(|e| format!("m3_mixed F6 adt: {e:?}"))?;
}
{
let cfg = cudarc::driver::LaunchConfig {
grid_dim: ((dims.batch * nc) as u32, nh as u32, 1),
block_dim: (cs as u32, 1, 1),
shared_mem_bytes: 0,
};
let mut bld = ctx
.stream
.launch_builder(m3k.m3_preprocess_chunks_typed.get(dtype));
let ks = acts.k_scaled_saved.cached_ptr();
let kp = acts.k.cached_ptr();
let qp = acts.q.cached_ptr();
bld.arg(&ks);
bld.arg(acts.qk_dot_saved.inner_mut());
bld.arg(acts.scale_saved.inner_mut());
bld.arg(acts.gamma_saved.inner_mut());
bld.arg(&kp);
bld.arg(&qp);
bld.arg(acts.dt.inner());
bld.arg(acts.trap.inner());
bld.arg(&b_i);
bld.arg(&t_i);
bld.arg(&nh_i);
bld.arg(&ds_i);
bld.arg(&cs);
unsafe { bld.launch(cfg) }.map_err(|e| format!("m3_mixed F6 K1 preprocess: {e:?}"))?;
}
{
let block_x = nh.min(256) as u32;
let grid_z = nh.div_ceil(block_x as usize) as u32;
let cfg = cudarc::driver::LaunchConfig {
grid_dim: (dims.batch as u32, nc as u32, grid_z),
block_dim: (block_x, 1, 1),
shared_mem_bytes: 0,
};
let mut bld = ctx.stream.launch_builder(&m3k.m3_da_cumsum);
bld.arg(acts.da_cumsum_saved.inner_mut());
bld.arg(adt_temp_scratch.inner());
bld.arg(&b_i);
bld.arg(&t_i);
bld.arg(&nh_i);
bld.arg(&cs);
unsafe { bld.launch(cfg) }.map_err(|e| format!("m3_mixed F6 K2 da_cumsum: {e:?}"))?;
}
{
let cfg = cudarc::driver::LaunchConfig {
grid_dim: ((dims.batch * nc) as u32, nh as u32, 1),
block_dim: (hd as u32, 1, 1),
shared_mem_bytes: 0,
};
let mut bld = ctx
.stream
.launch_builder(m3k.m3_chunk_state_fwd_typed.get(dtype));
let xp = acts.x.cached_ptr();
let ks = acts.k_scaled_saved.cached_ptr();
bld.arg(chunk_states_scratch.inner_mut());
bld.arg(&xp);
bld.arg(&ks);
bld.arg(acts.da_cumsum_saved.inner());
bld.arg(&b_i);
bld.arg(&t_i);
bld.arg(&nh_i);
bld.arg(&hd_i);
bld.arg(&ds_i);
bld.arg(&cs);
unsafe { bld.launch(cfg) }
.map_err(|e| format!("m3_mixed F6 K3 chunk_state_fwd: {e:?}"))?;
}
{
let dim = hd * ds;
let block_x = dim.min(256) as u32;
let grid_z = dim.div_ceil(block_x as usize) as u32;
let nc_i = nc as i32;
let cfg = cudarc::driver::LaunchConfig {
grid_dim: (dims.batch as u32, nh as u32, grid_z),
block_dim: (block_x, 1, 1),
shared_mem_bytes: 0,
};
let mut bld = ctx.stream.launch_builder(&m3k.m3_state_passing_fwd);
bld.arg(chunk_states_scratch.inner_mut());
bld.arg(final_states_scratch.inner_mut());
bld.arg(acts.da_cumsum_saved.inner());
bld.arg(&b_i);
bld.arg(&nc_i);
bld.arg(&nh_i);
bld.arg(&hd_i);
bld.arg(&ds_i);
bld.arg(&cs);
bld.arg(&t_i);
unsafe { bld.launch(cfg) }
.map_err(|e| format!("m3_mixed F6 K4 state_passing: {e:?}"))?;
}
acts.chunk_states_saved
.copy_from(chunk_states_scratch, &ctx.stream)?;
{
let cfg = cudarc::driver::LaunchConfig {
grid_dim: ((dims.batch * nc) as u32, nh as u32, 1),
block_dim: (hd as u32, 1, 1),
shared_mem_bytes: 0,
};
let mut bld = ctx
.stream
.launch_builder(m3k.m3_chunk_scan_fwd_typed.get(dtype));
let yp = acts.y.cached_ptr();
let xp = acts.x.cached_ptr();
let qp = acts.q.cached_ptr();
let ks = acts.k_scaled_saved.cached_ptr();
bld.arg(&yp);
bld.arg(&xp);
bld.arg(&qp);
bld.arg(&ks);
bld.arg(acts.qk_dot_saved.inner());
bld.arg(acts.da_cumsum_saved.inner());
bld.arg(chunk_states_scratch.inner());
bld.arg(&dp_ptr);
bld.arg(&b_i);
bld.arg(&t_i);
bld.arg(&nh_i);
bld.arg(&hd_i);
bld.arg(&ds_i);
bld.arg(&cs);
unsafe { bld.launch(cfg) }
.map_err(|e| format!("m3_mixed F6 K5 chunk_scan_fwd: {e:?}"))?;
}
{
let block_x = hd.max(ds) as u32;
let cfg = cudarc::driver::LaunchConfig {
grid_dim: (dims.batch as u32, nh as u32, 1),
block_dim: (block_x, 1, 1),
shared_mem_bytes: 0,
};
let mut bld = ctx
.stream
.launch_builder(m3k.m3_writeback_parallel_states_typed.get(dtype));
let kp = acts.k.cached_ptr();
let xp = acts.x.cached_ptr();
bld.arg(&ssm_state);
bld.arg(&k_state);
bld.arg(&v_state);
bld.arg(final_states_scratch.inner());
bld.arg(&kp);
bld.arg(&xp);
bld.arg(&b_i);
bld.arg(&t_i);
bld.arg(&nh_i);
bld.arg(&hd_i);
bld.arg(&ds_i);
unsafe { bld.launch(cfg) }.map_err(|e| format!("m3_mixed F6 K6 writeback: {e:?}"))?;
}
} else {
let dp_ptr = w.d_param.ptr();
let b_i = dims.batch as i32;
let t_i = dims.seq_len as i32;
let nh_i = nh as i32;
let hd_i = hd as i32;
let ds_i = ds as i32;
let cfg = cudarc::driver::LaunchConfig {
grid_dim: (dims.batch as u32, nh as u32, 1),
block_dim: (hd as u32, 1, 1),
shared_mem_bytes: 0,
};
let kernel = match dtype {
WeightDtype::F32 => &m3k.m3_burnin_fwd,
WeightDtype::Bf16 => &m3k.m3_burnin_fwd_typed_bf16,
WeightDtype::F16 => &m3k.m3_burnin_fwd_typed_f16,
};
let mut bld = ctx.stream.launch_builder(kernel);
bld.arg(&ssm_state);
bld.arg(&k_state);
bld.arg(&v_state);
let y = acts.y.cached_ptr();
bld.arg(&y);
bld.arg(acts.h_saved.inner_mut());
bld.arg(acts.k_prev_saved.inner_mut());
bld.arg(acts.v_prev_saved.inner_mut());
let x = acts.x.cached_ptr();
let k = acts.k.cached_ptr();
let q = acts.q.cached_ptr();
bld.arg(&x);
bld.arg(&k);
bld.arg(&q);
bld.arg(alpha_scratch.inner());
bld.arg(beta_scratch.inner());
bld.arg(gamma_scratch.inner());
bld.arg(&dp_ptr);
bld.arg(&b_i);
bld.arg(&t_i);
bld.arg(&nh_i);
bld.arg(&hd_i);
bld.arg(&ds_i);
unsafe { bld.launch(cfg) }.map_err(|e| format!("m3_mixed F6 burnin: {e:?}"))?;
}
if dims.is_outproj_norm {
assert!(di <= 1024, "d_inner {di} > 1024 rmsnorm_gated shm limit");
let nw_ptr = w.norm_gate_weight.ptr();
let bt_i = bt as i32;
let di_i = di as i32;
let hd_i = hd as i32;
let grid = cudarc::driver::LaunchConfig {
grid_dim: (bt as u32, 1, 1),
block_dim: (di as u32, 1, 1),
shared_mem_bytes: (di * std::mem::size_of::<f32>()) as u32,
};
let mut bld = ctx
.stream
.launch_builder(m3k.rmsnorm_gated_fwd_typed.get(dtype));
let g = acts.gated.cached_ptr();
let gr = acts.gated_rms_vals.cached_ptr();
let y = acts.y.cached_ptr();
let z = acts.z.cached_ptr();
bld.arg(&g);
bld.arg(&gr);
bld.arg(&y);
bld.arg(&z);
bld.arg(&nw_ptr);
bld.arg(&bt_i);
bld.arg(&di_i);
bld.arg(&hd_i);
unsafe { bld.launch(grid) }.map_err(|e| format!("m3_mixed F7 rmsnorm_gated: {e:?}"))?;
} else {
let n = (bt * di) as i32;
let mut bld = ctx
.stream
.launch_builder(m3k.silu_gate_fwd_typed.get(dtype));
let g = acts.gated.cached_ptr();
let y = acts.y.cached_ptr();
let z = acts.z.cached_ptr();
bld.arg(&g);
bld.arg(&y);
bld.arg(&z);
bld.arg(&n);
unsafe { bld.launch(grid_1d(bt * di)) }
.map_err(|e| format!("m3_mixed F7 silu_gate: {e:?}"))?;
}
gpu_gemm_typed_forward_raw(
ctx,
TypedPtr {
ptr: out_flat_scratch.cached_ptr(),
dtype,
},
TypedPtr {
ptr: acts.gated.cached_ptr(),
dtype,
},
TypedPtr {
ptr: w.out_proj_w.ptr(),
dtype,
},
None,
(bt, di, dm),
)?;
{
let n = (bt * dm) as i32;
let mut bld = ctx
.stream
.launch_builder(m3k.residual_add_f32_typed.get(dtype));
let dst = temporal_f32.cached_ptr();
let a = acts.residual.cached_ptr();
let b = out_flat_scratch.cached_ptr();
bld.arg(&dst);
bld.arg(&a);
bld.arg(&b);
bld.arg(&n);
unsafe { bld.launch(grid_1d(bt * dm)) }
.map_err(|e| format!("m3_mixed F8 residual_add: {e:?}"))?;
}
Ok(())
}
pub struct GpuMamba3MixedScratch {
pub proj_flat: DtypedBuf,
pub out_flat: DtypedBuf,
pub alpha: GpuBuffer,
pub beta: GpuBuffer,
pub gamma: GpuBuffer,
pub d_temporal_typed: DtypedBuf,
pub d_gated_typed: DtypedBuf,
pub d_y_typed: DtypedBuf,
pub d_z_typed: DtypedBuf,
pub d_b_normed_typed: DtypedBuf,
pub d_c_normed_typed: DtypedBuf,
pub d_b_raw_typed: DtypedBuf,
pub d_c_raw_typed: DtypedBuf,
pub d_proj_typed: DtypedBuf,
pub d_post_norm_typed: DtypedBuf,
pub adt_temp: GpuBuffer,
pub da_cumsum: GpuBuffer,
pub chunk_states: GpuBuffer,
pub final_states: GpuBuffer,
pub dtype: WeightDtype,
}
impl GpuMamba3MixedScratch {
pub fn new(
stream: &Arc<CudaStream>,
cfg: &Mamba3Config,
batch: usize,
seq_len: usize,
dtype: WeightDtype,
) -> Result<Self, String> {
let bt = batch * seq_len;
let dm = cfg.d_model;
let di = cfg.d_inner();
let ng = cfg.ngroups;
let ds = cfg.d_state;
let nh = cfg.nheads();
let hd = cfg.headdim;
let ip = cfg.in_proj_out_dim();
let n_chunks_max = seq_len.div_ceil(CHUNK_SIZE);
let s = Self {
proj_flat: DtypedBuf::zeros(stream, bt * ip, dtype)?,
out_flat: DtypedBuf::zeros(stream, bt * dm, dtype)?,
alpha: GpuBuffer::zeros(stream, bt * nh)?,
beta: GpuBuffer::zeros(stream, bt * nh)?,
gamma: GpuBuffer::zeros(stream, bt * nh)?,
d_temporal_typed: DtypedBuf::zeros(stream, bt * dm, dtype)?,
d_gated_typed: DtypedBuf::zeros(stream, bt * di, dtype)?,
d_y_typed: DtypedBuf::zeros(stream, bt * di, dtype)?,
d_z_typed: DtypedBuf::zeros(stream, bt * di, dtype)?,
d_b_normed_typed: DtypedBuf::zeros(stream, bt * ng * ds, dtype)?,
d_c_normed_typed: DtypedBuf::zeros(stream, bt * ng * ds, dtype)?,
d_b_raw_typed: DtypedBuf::zeros(stream, bt * ng * ds, dtype)?,
d_c_raw_typed: DtypedBuf::zeros(stream, bt * ng * ds, dtype)?,
d_proj_typed: DtypedBuf::zeros(stream, bt * ip, dtype)?,
d_post_norm_typed: DtypedBuf::zeros(stream, bt * dm, dtype)?,
adt_temp: GpuBuffer::zeros(stream, bt * nh)?,
da_cumsum: GpuBuffer::zeros(stream, batch * n_chunks_max * nh * CHUNK_SIZE)?,
chunk_states: GpuBuffer::zeros(stream, batch * n_chunks_max * nh * hd * ds)?,
final_states: GpuBuffer::zeros(stream, batch * nh * hd * ds)?,
dtype,
};
stream
.synchronize()
.map_err(|e| format!("sync after m3 mixed scratch alloc: {e:?}"))?;
Ok(s)
}
}
pub fn gpu_forward_mamba3_backbone_mixed(
exec: &M3Exec<'_>,
temporal_f32: &mut GpuBuffer,
acts: &mut GpuMamba3BackboneMixedActs,
w: &GpuMamba3TrainMixedWeights,
mamba_input: &GpuBuffer,
states: GpuMamba3StateBufs<'_>,
scratch: &mut GpuMamba3MixedScratch,
) -> Result<(), String> {
let M3Exec {
ctx,
kernels: m3k,
dims,
} = *exec;
let dtype = acts.dtype;
let bt = dims.bt();
let dm = dims.d_model;
let ds = dims.d_state;
let nh = dims.nheads;
let hd = dims.headdim;
let na = dims.n_angles.max(1);
if dims.use_parallel_scan {
states.ssm.zero(&ctx.stream)?;
states.k.zero(&ctx.stream)?;
states.v.zero(&ctx.stream)?;
states.angle.zero(&ctx.stream)?;
}
if w.compute.input_proj_w.len_elems() == 0 {
let bytes = bt * dm * 4;
let stream = ctx.stream.cu_stream();
unsafe {
let res = cudarc::driver::sys::cuMemcpyDtoDAsync_v2(
temporal_f32.cached_ptr(),
mamba_input.cached_ptr(),
bytes,
stream,
);
if res != cudarc::driver::sys::CUresult::CUDA_SUCCESS {
return Err(format!("m3_mixed identity input_proj D2D: {res:?}"));
}
}
} else {
return Err(
"m3_mixed forward: non-identity input_proj not yet implemented — run \
with `cpu.input_proj_w.clear()` (identity branch)"
.into(),
);
}
let f32_sz = std::mem::size_of::<f32>() as u64;
let ssm_base = states.ssm.cached_ptr();
let k_base = states.k.cached_ptr();
let v_base = states.v.cached_ptr();
let a_base = states.angle.cached_ptr();
for l in 0..dims.n_layers {
let ssm_off = dims.batch * l * nh * hd * ds;
let k_off = dims.batch * l * nh * ds;
let v_off = dims.batch * l * nh * hd;
let a_off = dims.batch * l * nh * na;
let layer_ptrs = Mamba3LayerPtrs {
ssm_state: ssm_base + ssm_off as u64 * f32_sz,
k_state: k_base + k_off as u64 * f32_sz,
v_state: v_base + v_off as u64 * f32_sz,
angle_state: a_base + a_off as u64 * f32_sz,
};
gpu_forward_mamba3_layer_mixed(
exec,
temporal_f32,
&mut acts.layers[l],
&w.compute.layers[l],
&layer_ptrs,
scratch,
dtype,
)?;
}
acts.norm_f_input.copy_from(temporal_f32, &ctx.stream)?;
{
let bt_i = bt as i32;
let dm_i = dm as i32;
let eps: f32 = 1e-5;
let mut bld = ctx
.stream
.launch_builder(m3k.rmsnorm_fwd_f32in_typed.get(dtype));
let post = scratch.out_flat.cached_ptr();
let rms = acts.norm_f_rms.cached_ptr();
let x = acts.norm_f_input.cached_ptr();
let nw = w.compute.norm_f_weight.ptr();
bld.arg(&post);
bld.arg(&rms);
bld.arg(&x);
bld.arg(&nw);
bld.arg(&bt_i);
bld.arg(&dm_i);
bld.arg(&eps);
unsafe { bld.launch(grid_norm(bt, dm)) }.map_err(|e| format!("m3_mixed norm_f: {e:?}"))?;
}
Ok(())
}