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,
use_parallel_scan: bool,
) -> 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,
if use_parallel_scan {
1
} else {
batch * (seq_len + 1) * di * ds
},
)?,
k_prev_saved: GpuBuffer::zeros(
stream,
if use_parallel_scan { 1 } else { bt * nh * ds },
)?,
v_prev_saved: GpuBuffer::zeros(
stream,
if use_parallel_scan { 1 } else { 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<'_>,
stream_out: cudarc::driver::sys::CUdeviceptr,
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;
{
let bt_i = bt as i32;
let dm_i = dm as i32;
let eps: f32 = dims.rms_norm_eps;
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);
let eps_g5: f32 = dims.rms_norm_eps;
bld.arg(&eps_g5);
unsafe { bld.launch(cfg) }.map_err(|e| format!("m3_mixed F4ab bcnorm_bc: {e:?}"))?;
}
if na > 0 {
crate::mamba3_siso::gpu::forward::gpu_angle_chunked_fwd(
ctx,
m3k,
&mut acts.angle_cumsum,
angle_state,
&acts.angles_raw,
&acts.dt,
&scratch.angle_chunk_sums,
&scratch.angle_chunk_carries,
bt / dims.seq_len,
dims.seq_len,
nh,
na,
)?;
}
{
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 na_i = na as i32;
let mut bld = ctx
.stream
.launch_builder(m3k.m3_bias_rope_fwd_typed.get(dtype));
let bb = acts.b_biased.cached_ptr();
let cb = acts.c_biased.cached_ptr();
let k = acts.k.cached_ptr();
let q = acts.q.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();
let ac = acts.angle_cumsum.cached_ptr();
bld.arg(&bb);
bld.arg(&cb);
bld.arg(&k);
bld.arg(&q);
bld.arg(&bn);
bld.arg(&cn);
bld.arg(&bb_p);
bld.arg(&cb_p);
bld.arg(&ac);
bld.arg(&n_i);
bld.arg(&nh_i);
bld.arg(&ng_i);
bld.arg(&ds_i);
bld.arg(&na_i);
unsafe { bld.launch(grid_1d(bt * nh * ds)) }
.map_err(|e| format!("m3_mixed F4 bias_rope: {e:?}"))?;
}
{
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_raw(alpha_scratch, &ctx.stream)?;
acts.beta.copy_from_raw(beta_scratch, &ctx.stream)?;
acts.gamma.copy_from_raw(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.div_ceil(2) as u32, 1),
block_dim: (hd as u32, 2, 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());
let init_states_null: crate::mamba3_siso::gpu::state::CUptr = 0;
bld.arg(&init_states_null);
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_raw(chunk_states_scratch, &ctx.stream)?;
{
let (coop, cfg) =
super::kernels::chunk_scan_cfg(dims.batch, nc, nh, hd, ds, dims.chunk_size());
let kern = if coop {
m3k.m3_chunk_scan_fwd_coop_typed.get(dtype)
} else {
m3k.m3_chunk_scan_fwd_typed.get(dtype)
};
let mut bld = ctx.stream.launch_builder(kern);
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);
let eps_g5: f32 = dims.rms_norm_eps;
bld.arg(&eps_g5);
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 = stream_out;
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 angle_chunk_sums: crate::mamba_ssm::gpu::buffers::GpuByteBuffer,
pub angle_chunk_carries: crate::mamba_ssm::gpu::buffers::GpuByteBuffer,
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)?,
angle_chunk_sums: crate::mamba_ssm::gpu::buffers::GpuByteBuffer::zeros(
stream,
batch
* n_chunks_max
* nh
* cfg.num_rope_angles().max(1)
* std::mem::size_of::<f64>(),
)?,
angle_chunk_carries: crate::mamba_ssm::gpu::buffers::GpuByteBuffer::zeros(
stream,
batch
* n_chunks_max
* nh
* cfg.num_rope_angles().max(1)
* std::mem::size_of::<f64>(),
)?,
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 {
let need = dims.batch * (dims.seq_len + 1) * dims.d_inner * ds;
if acts.layers.first().is_some_and(|l| l.h_saved.len() < need) {
return Err(
"m3_mixed forward: dims request the sequential scan but the acts \
were allocated for the chunked path (sentinel tapes) — construct \
the acts with the same use_parallel_scan flag the forward runs with"
.into(),
);
}
}
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 {
let mid = dims.mamba_input_dim;
{
let n = (bt * mid) as i32;
let cast = match dtype {
WeightDtype::Bf16 => &m3k.cast_f32_to_bf16,
WeightDtype::F16 => &m3k.cast_f32_to_f16,
WeightDtype::F32 => {
return Err("m3_mixed forward: unexpected f32 compute dtype".into());
}
};
let dst = acts.input_proj_inputs.cached_ptr();
let src = mamba_input.cached_ptr();
let mut bld = ctx.stream.launch_builder(cast);
bld.arg(&dst);
bld.arg(&src);
bld.arg(&n);
unsafe { bld.launch(grid_1d(bt * mid)) }
.map_err(|e| format!("m3_mixed input_proj input cast: {e:?}"))?;
}
gpu_gemm_typed_forward_raw(
ctx,
TypedPtr {
ptr: acts.input_proj_outputs.cached_ptr(),
dtype,
},
TypedPtr {
ptr: acts.input_proj_inputs.cached_ptr(),
dtype,
},
TypedPtr {
ptr: w.compute.input_proj_w.ptr(),
dtype,
},
Some(w.compute.input_proj_b.ptr()),
(bt, mid, dm),
)?;
{
let n = (bt * dm) as i32;
let cast = match dtype {
WeightDtype::Bf16 => &m3k.cast_bf16_to_f32,
WeightDtype::F16 => &m3k.cast_f16_to_f32,
WeightDtype::F32 => {
return Err("m3_mixed forward: unexpected f32 compute dtype".into());
}
};
let dst = temporal_f32.cached_ptr();
let src = acts.input_proj_outputs.cached_ptr();
let mut bld = ctx.stream.launch_builder(cast);
bld.arg(&dst);
bld.arg(&src);
bld.arg(&n);
unsafe { bld.launch(grid_1d(bt * dm)) }
.map_err(|e| format!("m3_mixed input_proj residual upcast: {e:?}"))?;
}
}
acts.layers[0]
.residual
.copy_from_raw(temporal_f32, &ctx.stream)?;
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,
};
let stream_out = if l + 1 < dims.n_layers {
acts.layers[l + 1].residual.cached_ptr()
} else {
acts.norm_f_input.cached_ptr()
};
gpu_forward_mamba3_layer_mixed(
exec,
stream_out,
&mut acts.layers[l],
&w.compute.layers[l],
&layer_ptrs,
scratch,
dtype,
)?;
}
{
let bt_i = bt as i32;
let dm_i = dm as i32;
let eps: f32 = dims.rms_norm_eps;
let mut bld = ctx.stream.launch_builder(&m3k.rmsnorm_fwd);
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(temporal_f32.inner_mut());
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(())
}