use super::kernels::Mamba3Kernels;
use super::weights::{
GpuMamba3Grads, GpuMamba3LayerGrads, GpuMamba3LayerWeights, GpuMamba3Weights,
};
use crate::mamba_ssm::gpu::blas::{gpu_sgemm_backward_grad_raw, gpu_sgemm_forward_raw};
use crate::mamba_ssm::gpu::buffers::GpuBuffer;
use crate::mamba_ssm::gpu::context::GpuCtx;
use crate::mamba_ssm::gpu::launch::{grid_1d, grid_norm};
use cudarc::driver::PushKernelArg;
use std::sync::Arc;
type CUptr = cudarc::driver::sys::CUdeviceptr;
#[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 {
64
}
pub fn n_chunks(&self) -> usize {
self.seq_len.div_ceil(self.chunk_size())
}
}
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 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)?,
})
}
}
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)?,
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)?,
})
}
}
pub fn gpu_forward_mamba3_layer(
ctx: &GpuCtx,
m3k: &Mamba3Kernels,
temporal: &mut GpuBuffer, acts: &mut GpuMamba3LayerActs,
lw: &GpuMamba3LayerWeights,
layer_ptrs: &Mamba3LayerPtrs,
scratch: &mut GpuMamba3Scratch,
dims: &GpuMamba3Dims,
) -> Result<(), String> {
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, &ctx.stream)?;
{
let mut builder = ctx.stream.launch_builder(&m3k.rmsnorm_fwd);
builder.arg(acts.post_norm.inner_mut());
builder.arg(acts.rms_vals.inner_mut());
builder.arg(temporal.inner());
let nw_ptr = lw.norm_weight.raw_ptr(&ctx.stream);
builder.arg(&nw_ptr);
let bt_i = bt as i32;
let dm_i = dm as i32;
let eps: f32 = 1e-5;
builder.arg(&bt_i);
builder.arg(&dm_i);
builder.arg(&eps);
unsafe { builder.launch(grid_norm(bt, dm)) }
.map_err(|e| format!("rmsnorm_fwd m3 F1: {:?}", e))?;
}
gpu_sgemm_forward_raw(
ctx,
&mut scratch.proj_flat,
&acts.post_norm,
lw.in_proj_w.raw_ptr(&ctx.stream),
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 = lw.dt_bias.raw_ptr(&ctx.stream);
let mut builder = ctx.stream.launch_builder(&m3k.m3_split);
builder.arg(acts.z.inner_mut());
builder.arg(acts.x.inner_mut());
builder.arg(acts.b_raw.inner_mut());
builder.arg(acts.c_raw.inner_mut());
builder.arg(acts.dt.inner_mut());
builder.arg(acts.a_val.inner_mut());
builder.arg(acts.trap.inner_mut());
builder.arg(acts.angles_raw.inner_mut());
builder.arg(acts.dd_dt_raw.inner_mut());
builder.arg(acts.dd_a_raw.inner_mut());
builder.arg(acts.trap_raw.inner_mut());
builder.arg(scratch.proj_flat.inner());
builder.arg(&db_ptr);
builder.arg(&dims.a_floor);
builder.arg(&n_i);
builder.arg(&di_i);
builder.arg(&ng_i);
builder.arg(&ds_i);
builder.arg(&nh_i);
builder.arg(&na_i);
unsafe { builder.launch(grid_1d(bt * ip)) }.map_err(|e| format!("m3_split F3: {:?}", e))?;
}
{
let bn_ptr = lw.b_norm_weight.raw_ptr(&ctx.stream);
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, 1, 1),
block_dim: (ds as u32, 1, 1),
shared_mem_bytes: ds as u32 * 4,
};
let mut builder = ctx.stream.launch_builder(&m3k.bcnorm_fwd);
builder.arg(acts.b_normed.inner_mut());
builder.arg(acts.b_rms.inner_mut());
builder.arg(acts.b_raw.inner());
builder.arg(&bn_ptr);
builder.arg(&n_i);
builder.arg(&ng_i);
builder.arg(&ds_i);
unsafe { builder.launch(cfg) }.map_err(|e| format!("bcnorm_fwd B F4a: {:?}", e))?;
}
{
let cn_ptr = lw.c_norm_weight.raw_ptr(&ctx.stream);
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, 1, 1),
block_dim: (ds as u32, 1, 1),
shared_mem_bytes: ds as u32 * 4,
};
let mut builder = ctx.stream.launch_builder(&m3k.bcnorm_fwd);
builder.arg(acts.c_normed.inner_mut());
builder.arg(acts.c_rms.inner_mut());
builder.arg(acts.c_raw.inner());
builder.arg(&cn_ptr);
builder.arg(&n_i);
builder.arg(&ng_i);
builder.arg(&ds_i);
unsafe { builder.launch(cfg) }.map_err(|e| format!("bcnorm_fwd C F4b: {:?}", e))?;
}
{
let bb_ptr = lw.b_bias.raw_ptr(&ctx.stream);
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 mut builder = ctx.stream.launch_builder(&m3k.bc_bias_add);
builder.arg(acts.b_biased.inner_mut());
builder.arg(acts.b_normed.inner());
builder.arg(&bb_ptr);
builder.arg(&n_i);
builder.arg(&nh_i);
builder.arg(&ng_i);
builder.arg(&ds_i);
unsafe { builder.launch(grid_1d(bt * nh * ds)) }
.map_err(|e| format!("bc_bias_add B F4c: {:?}", e))?;
}
{
let cb_ptr = lw.c_bias.raw_ptr(&ctx.stream);
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 mut builder = ctx.stream.launch_builder(&m3k.bc_bias_add);
builder.arg(acts.c_biased.inner_mut());
builder.arg(acts.c_normed.inner());
builder.arg(&cb_ptr);
builder.arg(&n_i);
builder.arg(&nh_i);
builder.arg(&ng_i);
builder.arg(&ds_i);
unsafe { builder.launch(grid_1d(bt * nh * ds)) }
.map_err(|e| format!("bc_bias_add C F4d: {:?}", 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 angle_st = layer_ptrs.angle_state;
let mut builder = ctx.stream.launch_builder(&m3k.m3_angle_dt_fwd_seq);
builder.arg(acts.angle_cumsum.inner_mut());
builder.arg(&angle_st);
builder.arg(acts.angles_raw.inner());
builder.arg(acts.dt.inner());
builder.arg(&b_i);
builder.arg(&t_i);
builder.arg(&nh_i);
builder.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 { builder.launch(grid) }.map_err(|e| format!("angle_dt_fwd_seq F5: {:?}", 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 builder = ctx.stream.launch_builder(&m3k.rope_fwd);
builder.arg(acts.k.inner_mut());
builder.arg(acts.q.inner_mut());
builder.arg(acts.b_biased.inner());
builder.arg(acts.c_biased.inner());
builder.arg(acts.angle_cumsum.inner());
builder.arg(&n_i);
builder.arg(&nh_i);
builder.arg(&ds_i);
builder.arg(&na_i);
unsafe { builder.launch(grid_1d(bt * nh * ds)) }
.map_err(|e| format!("rope_fwd F4ef: {:?}", e))?;
} else {
acts.k.copy_from(&acts.b_biased, &ctx.stream)?;
acts.q.copy_from(&acts.c_biased, &ctx.stream)?;
}
{
let n_total = (bt * nh) as i32;
let mut builder = ctx.stream.launch_builder(&m3k.m3_compute_abg);
builder.arg(acts.alpha.inner_mut());
builder.arg(acts.beta.inner_mut());
builder.arg(acts.gamma.inner_mut());
builder.arg(acts.dt.inner());
builder.arg(acts.a_val.inner());
builder.arg(acts.trap.inner());
builder.arg(&n_total);
unsafe { builder.launch(grid_1d(bt * nh)) }
.map_err(|e| format!("m3_compute_abg F5b: {:?}", e))?;
}
if dims.use_parallel_scan {
let dp_ptr = lw.d_param.raw_ptr(&ctx.stream);
let nh_i = nh as i32;
let hd_i = hd as i32;
let ds_i = ds as i32;
let t_i = dims.seq_len as i32;
let cs = dims.chunk_size() as i32;
let nc = dims.n_chunks();
let b_i = dims.batch as i32;
{
let n_total = (bt * nh) as i32;
let mut builder = ctx.stream.launch_builder(&m3k.elementwise_mul);
builder.arg(scratch.d_alpha.inner_mut());
builder.arg(acts.a_val.inner());
builder.arg(acts.dt.inner());
builder.arg(&n_total);
unsafe { builder.launch(grid_1d(bt * nh)) }
.map_err(|e| format!("adt compute F6: {:?}", e))?;
}
{
let cfg = cudarc::driver::LaunchConfig {
grid_dim: ((dims.batch * nc) as u32, nh as u32, 1),
block_dim: (dims.chunk_size() as u32, 1, 1),
shared_mem_bytes: 0,
};
let mut builder = ctx.stream.launch_builder(&m3k.m3_preprocess_chunks);
builder.arg(scratch.d_q.inner_mut());
builder.arg(scratch.d_beta.inner_mut());
builder.arg(scratch.d_gamma.inner_mut());
builder.arg(scratch.d_dd_dt.inner_mut());
builder.arg(acts.k.inner());
builder.arg(acts.q.inner());
builder.arg(acts.dt.inner());
builder.arg(acts.trap.inner());
builder.arg(&b_i);
builder.arg(&t_i);
builder.arg(&nh_i);
builder.arg(&ds_i);
builder.arg(&cs);
unsafe { builder.launch(cfg) }
.map_err(|e| format!("m3_preprocess_chunks F6 K1: {:?}", 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 builder = ctx.stream.launch_builder(&m3k.m3_da_cumsum);
builder.arg(scratch.da_cumsum.inner_mut());
builder.arg(scratch.d_alpha.inner());
builder.arg(&b_i);
builder.arg(&t_i);
builder.arg(&nh_i);
builder.arg(&cs);
unsafe { builder.launch(cfg) }.map_err(|e| format!("m3_dA_cumsum F6 K2: {:?}", 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 builder = ctx.stream.launch_builder(&m3k.m3_chunk_state_fwd);
builder.arg(scratch.chunk_states.inner_mut());
builder.arg(acts.x.inner());
builder.arg(scratch.d_q.inner());
builder.arg(scratch.da_cumsum.inner());
builder.arg(&b_i);
builder.arg(&t_i);
builder.arg(&nh_i);
builder.arg(&hd_i);
builder.arg(&ds_i);
builder.arg(&cs);
unsafe { builder.launch(cfg) }
.map_err(|e| format!("m3_chunk_state_fwd F6 K3: {:?}", 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 builder = ctx.stream.launch_builder(&m3k.m3_state_passing_fwd);
builder.arg(scratch.chunk_states.inner_mut());
builder.arg(scratch.final_states.inner_mut());
builder.arg(scratch.da_cumsum.inner());
builder.arg(&b_i);
builder.arg(&nc_i);
builder.arg(&nh_i);
builder.arg(&hd_i);
builder.arg(&ds_i);
builder.arg(&cs);
builder.arg(&t_i);
unsafe { builder.launch(cfg) }
.map_err(|e| format!("m3_state_passing_fwd F6 K4: {:?}", 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 builder = ctx.stream.launch_builder(&m3k.m3_chunk_scan_fwd);
builder.arg(acts.y.inner_mut());
builder.arg(acts.x.inner());
builder.arg(acts.q.inner());
builder.arg(scratch.d_q.inner());
builder.arg(scratch.d_beta.inner());
builder.arg(scratch.da_cumsum.inner());
builder.arg(scratch.chunk_states.inner());
builder.arg(&dp_ptr);
builder.arg(&b_i);
builder.arg(&t_i);
builder.arg(&nh_i);
builder.arg(&hd_i);
builder.arg(&ds_i);
builder.arg(&cs);
unsafe { builder.launch(cfg) }
.map_err(|e| format!("m3_chunk_scan_fwd F6 K5: {:?}", e))?;
}
acts.da_cumsum_saved
.copy_from(&scratch.da_cumsum, &ctx.stream)?;
acts.k_scaled_saved.copy_from(&scratch.d_q, &ctx.stream)?; acts.scale_saved.copy_from(&scratch.d_gamma, &ctx.stream)?; acts.gamma_saved.copy_from(&scratch.d_dd_dt, &ctx.stream)?; acts.qk_dot_saved.copy_from(&scratch.d_beta, &ctx.stream)?; acts.chunk_states_saved
.copy_from(&scratch.chunk_states, &ctx.stream)?;
{
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 builder = ctx.stream.launch_builder(&m3k.m3_writeback_parallel_states);
builder.arg(&layer_ptrs.ssm_state);
builder.arg(&layer_ptrs.k_state);
builder.arg(&layer_ptrs.v_state);
builder.arg(scratch.final_states.inner());
builder.arg(acts.k.inner());
builder.arg(acts.x.inner());
builder.arg(&b_i);
builder.arg(&t_i);
builder.arg(&nh_i);
builder.arg(&hd_i);
builder.arg(&ds_i);
unsafe { builder.launch(cfg) }
.map_err(|e| format!("m3_writeback_parallel_states F6: {:?}", e))?;
}
} else {
let dp_ptr = lw.d_param.raw_ptr(&ctx.stream);
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 mut builder = ctx.stream.launch_builder(&m3k.m3_burnin_fwd);
builder.arg(&layer_ptrs.ssm_state);
builder.arg(&layer_ptrs.k_state);
builder.arg(&layer_ptrs.v_state);
builder.arg(acts.y.inner_mut());
builder.arg(acts.h_saved.inner_mut());
builder.arg(acts.k_prev_saved.inner_mut());
builder.arg(acts.v_prev_saved.inner_mut());
builder.arg(acts.x.inner());
builder.arg(acts.k.inner());
builder.arg(acts.q.inner());
builder.arg(acts.alpha.inner());
builder.arg(acts.beta.inner());
builder.arg(acts.gamma.inner());
builder.arg(&dp_ptr);
builder.arg(&b_i);
builder.arg(&t_i);
builder.arg(&nh_i);
builder.arg(&hd_i);
builder.arg(&ds_i);
unsafe { builder.launch(cfg) }.map_err(|e| format!("m3_burnin_fwd F6 seq: {:?}", e))?;
}
if dims.is_outproj_norm {
assert!(
di <= 1024,
"d_inner ({di}) exceeds rmsnorm_gated shared memory limit (1024)"
);
let nw_ptr = lw.norm_gate_weight.raw_ptr(&ctx.stream);
let bt_i = bt as i32;
let di_i = di as i32;
let hd_i = dims.headdim 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 builder = ctx.stream.launch_builder(&m3k.rmsnorm_gated_fwd);
builder.arg(acts.gated.inner_mut());
builder.arg(acts.gated_rms_vals.inner_mut());
builder.arg(acts.y.inner());
builder.arg(acts.z.inner());
builder.arg(&nw_ptr);
builder.arg(&bt_i);
builder.arg(&di_i);
builder.arg(&hd_i);
unsafe { builder.launch(grid) }.map_err(|e| format!("rmsnorm_gated_fwd m3 F7: {:?}", e))?;
} else {
let n = (bt * di) as i32;
let mut builder = ctx.stream.launch_builder(&m3k.silu_gate_fwd);
builder.arg(acts.gated.inner_mut());
builder.arg(acts.y.inner());
builder.arg(acts.z.inner());
builder.arg(&n);
unsafe { builder.launch(grid_1d(bt * di)) }
.map_err(|e| format!("silu_gate_fwd m3 F7: {:?}", e))?;
}
gpu_sgemm_forward_raw(
ctx,
&mut scratch.out_flat,
&acts.gated,
lw.out_proj_w.raw_ptr(&ctx.stream),
None,
(bt, di, dm),
)?;
{
let ne = (bt * dm) as i32;
let mut builder = ctx.stream.launch_builder(&m3k.residual_add);
builder.arg(temporal.inner_mut());
builder.arg(scratch.out_flat.inner());
builder.arg(acts.residual.inner());
builder.arg(&ne);
unsafe { builder.launch(grid_1d(bt * dm)) }
.map_err(|e| format!("residual_add m3 F8: {:?}", e))?;
}
Ok(())
}
pub fn gpu_forward_mamba3_backbone(
ctx: &GpuCtx,
m3k: &Mamba3Kernels,
temporal: &mut GpuBuffer,
acts: &mut GpuMamba3BackboneActs,
mamba_w: &GpuMamba3Weights,
mamba_input: &GpuBuffer,
ssm_states: &mut GpuBuffer,
k_states: &mut GpuBuffer,
v_states: &mut GpuBuffer,
angle_states: &mut GpuBuffer,
scratch: &mut GpuMamba3Scratch,
dims: &GpuMamba3Dims,
) -> Result<(), String> {
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);
acts.input_proj_inputs.copy_from(mamba_input, &ctx.stream)?;
gpu_sgemm_forward_raw(
ctx,
temporal,
mamba_input,
mamba_w.input_proj_w.raw_ptr(&ctx.stream),
Some(mamba_w.input_proj_b.raw_ptr(&ctx.stream)),
(bt, dims.mamba_input_dim, dm),
)?;
acts.input_proj_outputs.copy_from(temporal, &ctx.stream)?;
let f32_sz = std::mem::size_of::<f32>() as u64;
let ssm_base = ssm_states.raw_ptr(&ctx.stream);
let k_base = k_states.raw_ptr(&ctx.stream);
let v_base = v_states.raw_ptr(&ctx.stream);
let a_base = angle_states.raw_ptr(&ctx.stream);
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(
ctx,
m3k,
temporal,
&mut acts.layers[l],
&mamba_w.layers[l],
&layer_ptrs,
scratch,
dims,
)?;
}
acts.norm_f_input.copy_from(temporal, &ctx.stream)?;
{
let nf_ptr = mamba_w.norm_f_weight.raw_ptr(&ctx.stream);
let bt_i = bt as i32;
let dm_i = dm as i32;
let eps: f32 = 1e-5;
let mut builder = ctx.stream.launch_builder(&m3k.rmsnorm_fwd);
builder.arg(temporal.inner_mut());
builder.arg(acts.norm_f_rms.inner_mut());
builder.arg(acts.norm_f_input.inner());
builder.arg(&nf_ptr);
builder.arg(&bt_i);
builder.arg(&dm_i);
builder.arg(&eps);
unsafe { builder.launch(grid_norm(bt, dm)) }
.map_err(|e| format!("rmsnorm_fwd norm_f m3: {:?}", e))?;
}
Ok(())
}
pub fn gpu_backward_mamba3_layer(
ctx: &GpuCtx,
m3k: &Mamba3Kernels,
d_temporal: &mut GpuBuffer, acts: &GpuMamba3LayerActs,
lw: &GpuMamba3LayerWeights,
lg: &GpuMamba3LayerGrads,
scratch: &mut GpuMamba3Scratch,
dims: &GpuMamba3Dims,
) -> Result<(), String> {
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 b = dims.batch as i32;
let t = dims.seq_len as i32;
gpu_sgemm_backward_grad_raw(
ctx,
&mut scratch.d_gated,
(&lg.out_proj_w, None),
d_temporal,
&acts.gated,
lw.out_proj_w.raw_ptr(&ctx.stream),
(bt, di, dm),
)?;
if dims.is_outproj_norm {
assert!(
di <= 1024,
"d_inner ({di}) exceeds rmsnorm_gated shared memory limit"
);
let nw_ptr = lw.norm_gate_weight.raw_ptr(&ctx.stream);
let bt_i = bt as i32;
let di_i = di as i32;
let hd_i = dims.headdim 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 builder = ctx.stream.launch_builder(&m3k.rmsnorm_gated_bwd);
builder.arg(scratch.d_y.inner_mut());
builder.arg(scratch.d_z.inner_mut());
builder.arg(scratch.d_norm_gate_w.inner_mut());
builder.arg(scratch.d_gated.inner());
builder.arg(acts.y.inner());
builder.arg(acts.z.inner());
builder.arg(&nw_ptr);
builder.arg(acts.gated_rms_vals.inner());
builder.arg(&bt_i);
builder.arg(&di_i);
builder.arg(&hd_i);
unsafe { builder.launch(grid) }.map_err(|e| format!("rmsnorm_gated_bwd m3 B7: {:?}", e))?;
{
let n_i = di as i32;
let mut builder = ctx.stream.launch_builder(&m3k.colsum_accumulate);
let dst = lg.norm_gate_weight.ptr();
builder.arg(&dst);
builder.arg(scratch.d_norm_gate_w.inner());
let bt_i = bt as i32;
builder.arg(&bt_i);
builder.arg(&n_i);
unsafe { builder.launch(grid_1d(di)) }
.map_err(|e| format!("colsum d_norm_gate_w m3: {:?}", e))?;
}
} else {
let n = (bt * di) as i32;
let mut builder = ctx.stream.launch_builder(&m3k.silu_gate_bwd);
builder.arg(scratch.d_y.inner_mut());
builder.arg(scratch.d_z.inner_mut());
builder.arg(scratch.d_gated.inner());
builder.arg(acts.y.inner());
builder.arg(acts.z.inner());
builder.arg(&n);
unsafe { builder.launch(grid_1d(bt * di)) }
.map_err(|e| format!("silu_gate_bwd m3 B7: {:?}", e))?;
}
scratch.d_k.zero(&ctx.stream)?;
scratch.d_q.zero(&ctx.stream)?;
scratch.d_alpha.zero(&ctx.stream)?;
scratch.d_beta.zero(&ctx.stream)?;
scratch.d_gamma.zero(&ctx.stream)?;
scratch.d_d_local.zero(&ctx.stream)?;
if !dims.use_parallel_scan {
let dp_ptr = lw.d_param.raw_ptr(&ctx.stream);
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 mut builder = ctx.stream.launch_builder(&m3k.m3_backward_seq);
builder.arg(acts.h_saved.inner());
builder.arg(acts.k_prev_saved.inner());
builder.arg(acts.v_prev_saved.inner());
builder.arg(acts.x.inner());
builder.arg(acts.k.inner());
builder.arg(acts.q.inner());
builder.arg(acts.alpha.inner());
builder.arg(acts.beta.inner());
builder.arg(acts.gamma.inner());
builder.arg(&dp_ptr);
builder.arg(scratch.d_y.inner()); builder.arg(scratch.d_x.inner_mut());
builder.arg(scratch.d_k.inner_mut());
builder.arg(scratch.d_q.inner_mut());
builder.arg(scratch.d_alpha.inner_mut());
builder.arg(scratch.d_beta.inner_mut());
builder.arg(scratch.d_gamma.inner_mut());
builder.arg(scratch.d_d_local.inner_mut());
builder.arg(&b_i);
builder.arg(&t_i);
builder.arg(&nh_i);
builder.arg(&hd_i);
builder.arg(&ds_i);
unsafe { builder.launch(cfg) }.map_err(|e| format!("m3_backward_seq B6: {:?}", 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 builder = ctx.stream.launch_builder(&m3k.rope_bwd);
builder.arg(scratch.d_b_pre_rope.inner_mut()); builder.arg(scratch.d_c_pre_rope.inner_mut()); builder.arg(scratch.d_angle_cumsum.inner_mut());
builder.arg(scratch.d_k.inner()); builder.arg(scratch.d_q.inner()); builder.arg(acts.b_biased.inner()); builder.arg(acts.c_biased.inner()); builder.arg(acts.angle_cumsum.inner());
builder.arg(&n_i);
builder.arg(&nh_i);
builder.arg(&ds_i);
builder.arg(&na_i);
unsafe { builder.launch(grid_1d(bt * nh * ds)) }
.map_err(|e| format!("rope_bwd seq B6: {:?}", e))?;
} else {
scratch.d_b_pre_rope.copy_from(&scratch.d_k, &ctx.stream)?;
scratch.d_c_pre_rope.copy_from(&scratch.d_q, &ctx.stream)?;
}
{
let d_bb_ptr = lg.b_bias.ptr();
let bt_i = bt as i32;
let nhds_i = (nh * ds) as i32;
let mut builder = ctx.stream.launch_builder(&m3k.colsum_accumulate);
builder.arg(&d_bb_ptr);
builder.arg(scratch.d_b_pre_rope.inner());
builder.arg(&bt_i);
builder.arg(&nhds_i);
unsafe { builder.launch(grid_1d(nh * ds)) }
.map_err(|e| format!("colsum d_b_bias seq: {:?}", e))?;
}
{
let d_cb_ptr = lg.c_bias.ptr();
let bt_i = bt as i32;
let nhds_i = (nh * ds) as i32;
let mut builder = ctx.stream.launch_builder(&m3k.colsum_accumulate);
builder.arg(&d_cb_ptr);
builder.arg(scratch.d_c_pre_rope.inner());
builder.arg(&bt_i);
builder.arg(&nhds_i);
unsafe { builder.launch(grid_1d(nh * ds)) }
.map_err(|e| format!("colsum d_c_bias seq: {:?}", e))?;
}
{
let d_dp = lg.d_param.ptr();
let b_i2 = dims.batch as i32;
let nh_i = nh as i32;
let hd_i = hd as i32;
let mut builder = ctx.stream.launch_builder(&m3k.m3_reduce_d_d);
builder.arg(&d_dp);
builder.arg(scratch.d_d_local.inner());
builder.arg(&b_i2);
builder.arg(&nh_i);
builder.arg(&hd_i);
let cfg = cudarc::driver::LaunchConfig {
grid_dim: (nh as u32, 1, 1),
block_dim: (32, 1, 1),
shared_mem_bytes: 0,
};
unsafe { builder.launch(cfg) }.map_err(|e| format!("m3_reduce_d_D seq B6: {:?}", e))?;
}
} else {
let dp_ptr = lw.d_param.raw_ptr(&ctx.stream);
let nh_i = nh as i32;
let hd_i = hd as i32;
let ds_i = ds as i32;
let t_i = dims.seq_len as i32;
let cs = dims.chunk_size() as i32;
let nc = dims.n_chunks();
let b_i = dims.batch as i32;
let na = dims.n_angles;
scratch
.da_cumsum
.copy_from(&acts.da_cumsum_saved, &ctx.stream)?;
scratch
.d_b_pre_rope
.copy_from(&acts.k_scaled_saved, &ctx.stream)?; scratch.d_scale.copy_from(&acts.scale_saved, &ctx.stream)?;
scratch
.d_gamma_par
.copy_from(&acts.gamma_saved, &ctx.stream)?;
scratch
.d_qk_dot
.copy_from(&acts.qk_dot_saved, &ctx.stream)?;
scratch
.chunk_states
.copy_from(&acts.chunk_states_saved, &ctx.stream)?;
{
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 builder = ctx.stream.launch_builder(&m3k.m3_extract_da_cs_sum);
builder.arg(scratch.da_cs_sum.inner_mut());
builder.arg(scratch.da_cumsum.inner());
builder.arg(&b_i);
builder.arg(&t_i);
builder.arg(&nh_i);
builder.arg(&cs);
unsafe { builder.launch(cfg) }
.map_err(|e| format!("m3_extract_da_cs_sum B6: {:?}", e))?;
}
{
let zero: f32 = 0.0;
for (buf, sz, label) in [
(&mut scratch.d_x, bt * di, "d_x"),
(&mut scratch.d_k, bt * nh * ds, "d_k"),
(&mut scratch.d_q, bt * nh * ds, "d_q"),
] {
let ne = sz as i32;
let mut builder = ctx.stream.launch_builder(&m3k.fill_scalar);
builder.arg(buf.inner_mut());
builder.arg(&zero);
builder.arg(&ne);
unsafe { builder.launch(grid_1d(sz)) }
.map_err(|e| format!("zero {label} B6 par: {:?}", e))?;
}
}
{
let cs_u = dims.chunk_size();
let smem = (cs_u * ds + cs_u * ds + cs_u * hd + cs_u * hd + cs_u + cs_u + hd * ds) * 4;
let cfg = cudarc::driver::LaunchConfig {
grid_dim: (nh as u32, dims.batch as u32, 1),
block_dim: (hd as u32, 1, 1),
shared_mem_bytes: smem as u32,
};
let mut builder = ctx.stream.launch_builder(&m3k.m3_dqkv);
builder.arg(scratch.d_q.inner_mut()); builder.arg(scratch.d_k.inner_mut()); builder.arg(scratch.d_x.inner_mut()); builder.arg(scratch.d_alpha.inner_mut()); builder.arg(scratch.d_beta.inner_mut()); let d_dp_ptr = lg.d_param.ptr();
builder.arg(&d_dp_ptr); builder.arg(acts.q.inner()); builder.arg(scratch.d_b_pre_rope.inner()); builder.arg(acts.x.inner()); builder.arg(scratch.da_cumsum.inner()); builder.arg(scratch.da_cs_sum.inner()); builder.arg(scratch.d_qk_dot.inner()); builder.arg(scratch.chunk_states.inner()); builder.arg(scratch.d_y.inner()); builder.arg(&dp_ptr); builder.arg(&b_i);
builder.arg(&t_i);
builder.arg(&nh_i);
builder.arg(&hd_i);
builder.arg(&ds_i);
builder.arg(&cs);
unsafe { builder.launch(cfg) }.map_err(|e| format!("m3_dqkv B6 S1: {:?}", e))?;
}
if na > 0 {
let n_ac = (dims.batch * dims.seq_len * nh * na) as i32;
let mut builder = ctx.stream.launch_builder(&m3k.fill_scalar);
builder.arg(scratch.d_angle_cumsum.inner_mut());
let zero: f32 = 0.0;
builder.arg(&zero);
builder.arg(&n_ac);
unsafe { builder.launch(grid_1d(n_ac as usize)) }
.map_err(|e| format!("zero d_angle_cumsum: {:?}", e))?;
}
{
let na_i = na as i32;
let cfg = cudarc::driver::LaunchConfig {
grid_dim: ((dims.batch * nc) as u32, nh as u32, 1),
block_dim: (dims.chunk_size() as u32, 1, 1),
shared_mem_bytes: 0,
};
let d_cb_ptr = lg.c_bias.ptr();
let d_bb_ptr = lg.b_bias.ptr();
let mut builder = ctx.stream.launch_builder(&m3k.m3_dqktheta);
builder.arg(scratch.d_c_pre_rope.inner_mut()); builder.arg(scratch.d_b_pre_rope.inner_mut()); builder.arg(scratch.d_angle_cumsum.inner_mut()); let scale_in_ptr = scratch.d_scale.raw_ptr(&ctx.stream);
let gamma_in_ptr = scratch.d_gamma_par.raw_ptr(&ctx.stream);
builder.arg(scratch.d_scale.inner_mut()); builder.arg(scratch.d_gamma_par.inner_mut()); builder.arg(&d_cb_ptr); builder.arg(&d_bb_ptr); builder.arg(acts.c_biased.inner()); builder.arg(acts.b_biased.inner()); builder.arg(&scale_in_ptr); builder.arg(&gamma_in_ptr); builder.arg(acts.angle_cumsum.inner()); builder.arg(scratch.d_q.inner()); builder.arg(scratch.d_k.inner()); builder.arg(scratch.d_beta.inner()); builder.arg(&b_i);
builder.arg(&t_i);
builder.arg(&nh_i);
builder.arg(&ds_i);
builder.arg(&na_i);
builder.arg(&cs);
unsafe { builder.launch(cfg) }.map_err(|e| format!("m3_dqktheta B6 S2: {:?}", e))?;
}
scratch.d_q.copy_from(&scratch.d_c_pre_rope, &ctx.stream)?;
scratch.d_k.copy_from(&scratch.d_b_pre_rope, &ctx.stream)?;
{
let mut builder = ctx.stream.launch_builder(&m3k.m3_ddt_dtrap);
builder.arg(scratch.d_gamma.inner_mut()); builder.arg(scratch.d_trap_raw.inner_mut()); builder.arg(scratch.d_scale.inner()); builder.arg(scratch.d_gamma_par.inner()); builder.arg(acts.dt.inner()); builder.arg(acts.trap.inner()); builder.arg(&b_i);
builder.arg(&t_i);
builder.arg(&nh_i);
unsafe { builder.launch(grid_1d(bt * nh)) }
.map_err(|e| format!("m3_ddt_dtrap B6 S3: {:?}", e))?;
}
scratch.d_b_pre_rope.copy_from(&scratch.d_k, &ctx.stream)?;
scratch.d_c_pre_rope.copy_from(&scratch.d_q, &ctx.stream)?;
}
if na > 0 {
{
let ne = (bt * na) as i32;
let zero: f32 = 0.0;
let mut builder = ctx.stream.launch_builder(&m3k.fill_scalar);
builder.arg(scratch.d_angles_raw.inner_mut());
builder.arg(&zero);
builder.arg(&ne);
unsafe { builder.launch(grid_1d(bt * na)) }
.map_err(|e| format!("fill d_angles_raw B5a: {:?}", e))?;
}
{
let ne = (bt * nh) as i32;
let zero: f32 = 0.0;
let mut builder = ctx.stream.launch_builder(&m3k.fill_scalar);
builder.arg(scratch.d_dt_angle.inner_mut());
builder.arg(&zero);
builder.arg(&ne);
unsafe { builder.launch(grid_1d(bt * nh)) }
.map_err(|e| format!("fill d_dt_angle B5a: {:?}", e))?;
}
let t_i = t;
let nh_i = nh as i32;
let na_i = na as i32;
let b_i = b;
let mut builder = ctx.stream.launch_builder(&m3k.m3_angle_dt_bwd_seq);
builder.arg(scratch.d_angles_raw.inner_mut()); builder.arg(scratch.d_dt_angle.inner_mut()); builder.arg(scratch.d_angle_cumsum.inner()); builder.arg(acts.angles_raw.inner()); builder.arg(acts.dt.inner()); builder.arg(&b_i); builder.arg(&t_i); builder.arg(&nh_i); builder.arg(&na_i); let grid = cudarc::driver::LaunchConfig {
grid_dim: (dims.batch 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 { builder.launch(grid) }.map_err(|e| format!("m3_angle_dt_bwd_seq B5a: {:?}", e))?;
} else {
let ne = (bt * nh) as i32;
let zero: f32 = 0.0;
let mut builder = ctx.stream.launch_builder(&m3k.fill_scalar);
builder.arg(scratch.d_dt_angle.inner_mut());
builder.arg(&zero);
builder.arg(&ne);
unsafe { builder.launch(grid_1d(bt * nh)) }
.map_err(|e| format!("fill d_dt_angle no-angle B5a: {:?}", e))?;
}
if !dims.use_parallel_scan {
let n_total = (bt * nh) as i32;
let nh_i = nh as i32;
let dtb_ptr = lw.dt_bias.raw_ptr(&ctx.stream);
let mut builder = ctx.stream.launch_builder(&m3k.m3_abg_bwd);
builder.arg(scratch.d_dd_dt.inner_mut());
builder.arg(scratch.d_dd_a.inner_mut());
builder.arg(scratch.d_trap_raw.inner_mut());
builder.arg(scratch.d_alpha.inner());
builder.arg(scratch.d_beta.inner());
builder.arg(scratch.d_gamma.inner());
builder.arg(scratch.d_dt_angle.inner());
builder.arg(acts.dt.inner());
builder.arg(acts.a_val.inner());
builder.arg(acts.alpha.inner());
builder.arg(acts.dd_dt_raw.inner());
builder.arg(acts.dd_a_raw.inner());
builder.arg(acts.trap_raw.inner());
builder.arg(&dtb_ptr);
builder.arg(&dims.a_floor);
builder.arg(&n_total);
builder.arg(&nh_i);
unsafe { builder.launch(grid_1d(bt * nh)) }
.map_err(|e| format!("m3_abg_bwd seq B5b: {:?}", e))?;
} else {
let n_total = (bt * nh) as i32;
let nh_i = nh as i32;
let dtb_ptr = lw.dt_bias.raw_ptr(&ctx.stream);
let mut builder = ctx.stream.launch_builder(&m3k.m3_final_grads);
builder.arg(scratch.d_dd_dt.inner_mut());
builder.arg(scratch.d_dd_a.inner_mut());
builder.arg(scratch.d_alpha.inner()); builder.arg(scratch.d_gamma.inner()); builder.arg(scratch.d_dt_angle.inner()); builder.arg(acts.a_val.inner()); builder.arg(acts.dt.inner()); builder.arg(acts.dd_dt_raw.inner()); builder.arg(acts.dd_a_raw.inner()); builder.arg(&dtb_ptr); builder.arg(&dims.a_floor); builder.arg(&n_total); builder.arg(&nh_i); unsafe { builder.launch(grid_1d(bt * nh)) }
.map_err(|e| format!("m3_final_grads B5b: {:?}", e))?;
}
{
let d_dtb_ptr = lg.dt_bias.ptr();
let bt_i = bt as i32;
let nh_i = nh as i32;
let mut builder = ctx.stream.launch_builder(&m3k.colsum_accumulate);
builder.arg(&d_dtb_ptr);
builder.arg(scratch.d_dd_dt.inner());
builder.arg(&bt_i);
builder.arg(&nh_i);
unsafe { builder.launch(grid_1d(nh)) }
.map_err(|e| format!("colsum d_dt_bias B5b: {:?}", 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 mut builder = ctx.stream.launch_builder(&m3k.bc_bias_add_bwd);
builder.arg(scratch.d_b_normed.inner_mut());
builder.arg(scratch.d_b_pre_rope.inner());
builder.arg(&n_i);
builder.arg(&nh_i);
builder.arg(&ng_i);
builder.arg(&ds_i);
unsafe { builder.launch(grid_1d(bt * ng * ds)) }
.map_err(|e| format!("bc_bias_add_bwd B B4b: {:?}", 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 mut builder = ctx.stream.launch_builder(&m3k.bc_bias_add_bwd);
builder.arg(scratch.d_c_normed.inner_mut());
builder.arg(scratch.d_c_pre_rope.inner());
builder.arg(&n_i);
builder.arg(&nh_i);
builder.arg(&ng_i);
builder.arg(&ds_i);
unsafe { builder.launch(grid_1d(bt * ng * ds)) }
.map_err(|e| format!("bc_bias_add_bwd C B4b: {:?}", e))?;
}
{
let bn_ptr = lw.b_norm_weight.raw_ptr(&ctx.stream);
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, 1, 1),
block_dim: (ds as u32, 1, 1),
shared_mem_bytes: ds as u32 * 4,
};
let mut builder = ctx.stream.launch_builder(&m3k.bcnorm_bwd);
builder.arg(scratch.d_b_raw.inner_mut()); builder.arg(scratch.d_b_norm_w.inner_mut()); builder.arg(scratch.d_b_normed.inner()); builder.arg(acts.b_raw.inner()); builder.arg(acts.b_rms.inner()); builder.arg(&bn_ptr); builder.arg(&n_i); builder.arg(&ng_i); builder.arg(&ds_i); unsafe { builder.launch(cfg) }.map_err(|e| format!("bcnorm_bwd B B4c: {:?}", e))?;
}
{
let d_bnw_ptr = lg.b_norm_weight.ptr();
let rows = (bt * ng) as i32;
let ds_i = ds as i32;
let mut builder = ctx.stream.launch_builder(&m3k.colsum_accumulate);
builder.arg(&d_bnw_ptr);
builder.arg(scratch.d_b_norm_w.inner());
builder.arg(&rows);
builder.arg(&ds_i);
unsafe { builder.launch(grid_1d(ds)) }
.map_err(|e| format!("colsum d_b_norm_w B4c: {:?}", e))?;
}
{
let cn_ptr = lw.c_norm_weight.raw_ptr(&ctx.stream);
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, 1, 1),
block_dim: (ds as u32, 1, 1),
shared_mem_bytes: ds as u32 * 4,
};
let mut builder = ctx.stream.launch_builder(&m3k.bcnorm_bwd);
builder.arg(scratch.d_c_raw.inner_mut()); builder.arg(scratch.d_c_norm_w.inner_mut()); builder.arg(scratch.d_c_normed.inner()); builder.arg(acts.c_raw.inner()); builder.arg(acts.c_rms.inner()); builder.arg(&cn_ptr); builder.arg(&n_i); builder.arg(&ng_i); builder.arg(&ds_i); unsafe { builder.launch(cfg) }.map_err(|e| format!("bcnorm_bwd C B4c: {:?}", e))?;
}
{
let d_cnw_ptr = lg.c_norm_weight.ptr();
let rows = (bt * ng) as i32;
let ds_i = ds as i32;
let mut builder = ctx.stream.launch_builder(&m3k.colsum_accumulate);
builder.arg(&d_cnw_ptr);
builder.arg(scratch.d_c_norm_w.inner());
builder.arg(&rows);
builder.arg(&ds_i);
unsafe { builder.launch(grid_1d(ds)) }
.map_err(|e| format!("colsum d_c_norm_w B4c: {:?}", e))?;
}
{
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 mut builder = ctx.stream.launch_builder(&m3k.m3_split_bwd);
builder.arg(scratch.d_proj.inner_mut()); builder.arg(scratch.d_z.inner()); builder.arg(scratch.d_x.inner()); builder.arg(scratch.d_b_raw.inner()); builder.arg(scratch.d_c_raw.inner()); builder.arg(scratch.d_dd_dt.inner()); builder.arg(scratch.d_dd_a.inner()); builder.arg(scratch.d_trap_raw.inner()); builder.arg(scratch.d_angles_raw.inner()); builder.arg(&n_i); builder.arg(&di_i); builder.arg(&ng_i); builder.arg(&ds_i); builder.arg(&nh_i); builder.arg(&na_i); unsafe { builder.launch(grid_1d(bt * ip)) }
.map_err(|e| format!("m3_split_bwd B3: {:?}", e))?;
}
gpu_sgemm_backward_grad_raw(
ctx,
&mut scratch.d_norm,
(&lg.in_proj_w, None),
&scratch.d_proj,
&acts.post_norm,
lw.in_proj_w.raw_ptr(&ctx.stream),
(bt, dm, ip),
)?;
{
let nw_ptr = lw.norm_weight.raw_ptr(&ctx.stream);
let d_nw_ptr = lg.norm_weight.ptr();
let bt_i = bt as i32;
let dm_i = dm as i32;
let mut builder = ctx.stream.launch_builder(&m3k.rmsnorm_bwd);
builder.arg(scratch.d_pre_norm.inner_mut()); builder.arg(&d_nw_ptr); builder.arg(scratch.d_norm.inner()); builder.arg(acts.residual.inner()); builder.arg(&nw_ptr); builder.arg(acts.rms_vals.inner()); builder.arg(&bt_i);
builder.arg(&dm_i);
unsafe { builder.launch(grid_norm(bt, dm)) }
.map_err(|e| format!("rmsnorm_bwd m3 B1: {:?}", e))?;
}
{
let ne = (bt * dm) as i32;
let mut builder = ctx.stream.launch_builder(&m3k.vec_add_inplace);
builder.arg(d_temporal.inner_mut());
builder.arg(scratch.d_pre_norm.inner());
builder.arg(&ne);
unsafe { builder.launch(grid_1d(bt * dm)) }
.map_err(|e| format!("vec_add d_temporal m3: {:?}", e))?;
}
Ok(())
}
pub fn gpu_backward_mamba3_backbone(
ctx: &GpuCtx,
m3k: &Mamba3Kernels,
d_temporal: &mut GpuBuffer, acts: &GpuMamba3BackboneActs,
mamba_w: &GpuMamba3Weights,
grads: &GpuMamba3Grads,
scratch: &mut GpuMamba3Scratch,
dims: &GpuMamba3Dims,
) -> Result<(), String> {
let bt = dims.bt();
let dm = dims.d_model;
{
let nf_ptr = mamba_w.norm_f_weight.raw_ptr(&ctx.stream);
let d_nf_ptr = grads.norm_f_weight.ptr();
let bt_i = bt as i32;
let dm_i = dm as i32;
let mut builder = ctx.stream.launch_builder(&m3k.rmsnorm_bwd);
builder.arg(scratch.d_norm.inner_mut()); builder.arg(&d_nf_ptr); builder.arg(d_temporal.inner()); builder.arg(acts.norm_f_input.inner()); builder.arg(&nf_ptr); builder.arg(acts.norm_f_rms.inner()); builder.arg(&bt_i);
builder.arg(&dm_i);
unsafe { builder.launch(grid_norm(bt, dm)) }
.map_err(|e| format!("rmsnorm_bwd norm_f m3: {:?}", e))?;
}
d_temporal.copy_from(&scratch.d_norm, &ctx.stream)?;
for l in (0..dims.n_layers).rev() {
gpu_backward_mamba3_layer(
ctx,
m3k,
d_temporal,
&acts.layers[l],
&mamba_w.layers[l],
&grads.layers[l],
scratch,
dims,
)?;
}
gpu_sgemm_backward_grad_raw(
ctx,
&mut scratch.d_input_proj_dx,
(&grads.input_proj_w, Some(&grads.input_proj_b)),
d_temporal,
&acts.input_proj_inputs,
mamba_w.input_proj_w.raw_ptr(&ctx.stream),
(bt, dims.mamba_input_dim, dm),
)?;
Ok(())
}
pub fn gpu_forward_mamba3_target_burnin(
ctx: &GpuCtx,
m3k: &Mamba3Kernels,
temporal: &mut GpuBuffer, mamba_w: &GpuMamba3Weights,
mamba_input: &GpuBuffer,
tgt: &mut GpuMamba3TargetScratch,
dims: &GpuMamba3Dims,
) -> Result<(), String> {
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.max(1);
let b = dims.batch as i32;
let t = dims.seq_len as i32;
let f32_sz = std::mem::size_of::<f32>() as u64;
tgt.ssm_states.zero(&ctx.stream)?;
tgt.k_states.zero(&ctx.stream)?;
tgt.v_states.zero(&ctx.stream)?;
tgt.angle_states.zero(&ctx.stream)?;
gpu_sgemm_forward_raw(
ctx,
&mut tgt.temporal_work,
mamba_input,
mamba_w.input_proj_w.raw_ptr(&ctx.stream),
Some(mamba_w.input_proj_b.raw_ptr(&ctx.stream)),
(bt, dims.mamba_input_dim, dm),
)?;
for l in 0..dims.n_layers {
let lw = &mamba_w.layers[l];
tgt.residual.copy_from(&tgt.temporal_work, &ctx.stream)?;
{
let bt_i = bt as i32;
let dm_i = dm as i32;
let eps: f32 = 1e-5;
let nw_ptr = lw.norm_weight.raw_ptr(&ctx.stream);
let mut builder = ctx.stream.launch_builder(&m3k.rmsnorm_fwd);
builder.arg(tgt.out_flat.inner_mut()); builder.arg(tgt.rms_discard.inner_mut());
builder.arg(tgt.residual.inner());
builder.arg(&nw_ptr);
builder.arg(&bt_i);
builder.arg(&dm_i);
builder.arg(&eps);
unsafe { builder.launch(grid_norm(bt, dm)) }
.map_err(|e| format!("rmsnorm_fwd m3 tgt L{l}: {:?}", e))?;
}
gpu_sgemm_forward_raw(
ctx,
&mut tgt.proj_flat,
&tgt.out_flat,
lw.in_proj_w.raw_ptr(&ctx.stream),
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 = lw.dt_bias.raw_ptr(&ctx.stream);
let mut builder = ctx.stream.launch_builder(&m3k.m3_split);
builder.arg(tgt.z.inner_mut());
builder.arg(tgt.x.inner_mut());
builder.arg(tgt.b_raw.inner_mut());
builder.arg(tgt.c_raw.inner_mut());
builder.arg(tgt.dt.inner_mut());
builder.arg(tgt.a_val.inner_mut());
builder.arg(tgt.trap.inner_mut());
builder.arg(tgt.angles_raw.inner_mut());
builder.arg(tgt.dd_dt_raw.inner_mut());
builder.arg(tgt.dd_a_raw.inner_mut());
builder.arg(tgt.trap_raw.inner_mut());
builder.arg(tgt.proj_flat.inner());
builder.arg(&db_ptr);
builder.arg(&dims.a_floor);
builder.arg(&n_i);
builder.arg(&di_i);
builder.arg(&ng_i);
builder.arg(&ds_i);
builder.arg(&nh_i);
builder.arg(&na_i);
unsafe { builder.launch(grid_1d(bt * ip)) }
.map_err(|e| format!("m3_split tgt L{l}: {:?}", e))?;
}
{
let bn_ptr = lw.b_norm_weight.raw_ptr(&ctx.stream);
let cfg = cudarc::driver::LaunchConfig {
grid_dim: ((bt * ng) as u32, 1, 1),
block_dim: (ds as u32, 1, 1),
shared_mem_bytes: ds as u32 * 4,
};
let mut builder = ctx.stream.launch_builder(&m3k.bcnorm_fwd);
builder.arg(tgt.b_normed.inner_mut());
builder.arg(tgt.b_rms.inner_mut());
builder.arg(tgt.b_raw.inner());
builder.arg(&bn_ptr);
let n_i = bt as i32;
let ng_i = ng as i32;
let ds_i = ds as i32;
builder.arg(&n_i);
builder.arg(&ng_i);
builder.arg(&ds_i);
unsafe { builder.launch(cfg) }
.map_err(|e| format!("bcnorm_fwd B tgt L{l}: {:?}", e))?;
}
{
let cn_ptr = lw.c_norm_weight.raw_ptr(&ctx.stream);
let cfg = cudarc::driver::LaunchConfig {
grid_dim: ((bt * ng) as u32, 1, 1),
block_dim: (ds as u32, 1, 1),
shared_mem_bytes: ds as u32 * 4,
};
let mut builder = ctx.stream.launch_builder(&m3k.bcnorm_fwd);
builder.arg(tgt.c_normed.inner_mut());
builder.arg(tgt.c_rms.inner_mut());
builder.arg(tgt.c_raw.inner());
builder.arg(&cn_ptr);
let n_i = bt as i32;
let ng_i = ng as i32;
let ds_i = ds as i32;
builder.arg(&n_i);
builder.arg(&ng_i);
builder.arg(&ds_i);
unsafe { builder.launch(cfg) }
.map_err(|e| format!("bcnorm_fwd C tgt L{l}: {:?}", e))?;
}
{
let bb_ptr = lw.b_bias.raw_ptr(&ctx.stream);
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 mut builder = ctx.stream.launch_builder(&m3k.bc_bias_add);
builder.arg(tgt.b_biased.inner_mut());
builder.arg(tgt.b_normed.inner());
builder.arg(&bb_ptr);
builder.arg(&n_i);
builder.arg(&nh_i);
builder.arg(&ng_i);
builder.arg(&ds_i);
unsafe { builder.launch(grid_1d(bt * nh * ds)) }
.map_err(|e| format!("bc_bias_add B tgt L{l}: {:?}", e))?;
}
{
let cb_ptr = lw.c_bias.raw_ptr(&ctx.stream);
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 mut builder = ctx.stream.launch_builder(&m3k.bc_bias_add);
builder.arg(tgt.c_biased.inner_mut());
builder.arg(tgt.c_normed.inner());
builder.arg(&cb_ptr);
builder.arg(&n_i);
builder.arg(&nh_i);
builder.arg(&ng_i);
builder.arg(&ds_i);
unsafe { builder.launch(grid_1d(bt * nh * ds)) }
.map_err(|e| format!("bc_bias_add C tgt L{l}: {:?}", e))?;
}
if dims.n_angles > 0 {
{
let a_off = dims.batch * l * nh * na;
let angle_st = tgt.angle_states.raw_ptr(&ctx.stream) + a_off as u64 * f32_sz;
let b_i = dims.batch as i32;
let t_i = dims.seq_len as i32;
let nh_i = nh as i32;
let na_i = na as i32;
let mut builder = ctx.stream.launch_builder(&m3k.m3_angle_dt_fwd_seq);
builder.arg(tgt.angle_cumsum.inner_mut());
builder.arg(&angle_st);
builder.arg(tgt.angles_raw.inner());
builder.arg(tgt.dt.inner());
builder.arg(&b_i);
builder.arg(&t_i);
builder.arg(&nh_i);
builder.arg(&na_i);
let grid = cudarc::driver::LaunchConfig {
grid_dim: (dims.batch 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 { builder.launch(grid) }
.map_err(|e| format!("angle_dt_fwd_seq tgt L{l}: {:?}", e))?;
}
{
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 builder = ctx.stream.launch_builder(&m3k.rope_fwd);
builder.arg(tgt.k.inner_mut()); builder.arg(tgt.q.inner_mut()); builder.arg(tgt.b_biased.inner()); builder.arg(tgt.c_biased.inner()); builder.arg(tgt.angle_cumsum.inner()); builder.arg(&n_i);
builder.arg(&nh_i);
builder.arg(&ds_i);
builder.arg(&na_i);
unsafe { builder.launch(grid_1d(bt * nh * ds)) }
.map_err(|e| format!("rope_fwd tgt L{l}: {:?}", e))?;
}
} else {
tgt.k.copy_from(&tgt.b_biased, &ctx.stream)?;
tgt.q.copy_from(&tgt.c_biased, &ctx.stream)?;
}
{
let n_total = (bt * nh) as i32;
let mut builder = ctx.stream.launch_builder(&m3k.m3_compute_abg);
builder.arg(tgt.alpha.inner_mut());
builder.arg(tgt.beta.inner_mut());
builder.arg(tgt.gamma.inner_mut());
builder.arg(tgt.dt.inner());
builder.arg(tgt.a_val.inner());
builder.arg(tgt.trap.inner());
builder.arg(&n_total);
unsafe { builder.launch(grid_1d(bt * nh)) }
.map_err(|e| format!("m3_compute_abg tgt L{l}: {:?}", e))?;
}
{
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 ssm_ptr = tgt.ssm_states.raw_ptr(&ctx.stream) + ssm_off as u64 * f32_sz;
let k_ptr = tgt.k_states.raw_ptr(&ctx.stream) + k_off as u64 * f32_sz;
let v_ptr = tgt.v_states.raw_ptr(&ctx.stream) + v_off as u64 * f32_sz;
let dp_ptr = lw.d_param.raw_ptr(&ctx.stream);
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 mut builder = ctx.stream.launch_builder(&m3k.m3_burnin_fwd_nosave);
builder.arg(&ssm_ptr); builder.arg(&k_ptr); builder.arg(&v_ptr); builder.arg(tgt.y.inner_mut()); builder.arg(tgt.x.inner()); builder.arg(tgt.k.inner()); builder.arg(tgt.q.inner()); builder.arg(tgt.alpha.inner()); builder.arg(tgt.beta.inner()); builder.arg(tgt.gamma.inner()); builder.arg(&dp_ptr); builder.arg(&b); builder.arg(&t); builder.arg(&nh_i); builder.arg(&hd_i); builder.arg(&ds_i); unsafe { builder.launch(cfg) }
.map_err(|e| format!("m3_burnin_fwd_nosave tgt L{l}: {:?}", e))?;
}
if dims.is_outproj_norm {
assert!(
di <= 1024,
"d_inner ({di}) exceeds rmsnorm_gated shared memory limit"
);
let nw_ptr = lw.norm_gate_weight.raw_ptr(&ctx.stream);
let bt_i = bt as i32;
let di_i = di as i32;
let hd_i = dims.headdim 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 builder = ctx.stream.launch_builder(&m3k.rmsnorm_gated_fwd);
builder.arg(tgt.gated.inner_mut());
builder.arg(tgt.rms_discard.inner_mut());
builder.arg(tgt.y.inner());
builder.arg(tgt.z.inner());
builder.arg(&nw_ptr);
builder.arg(&bt_i);
builder.arg(&di_i);
builder.arg(&hd_i);
unsafe { builder.launch(grid) }
.map_err(|e| format!("rmsnorm_gated_fwd m3 tgt L{l}: {:?}", e))?;
} else {
let n = (bt * di) as i32;
let mut builder = ctx.stream.launch_builder(&m3k.silu_gate_fwd);
builder.arg(tgt.gated.inner_mut());
builder.arg(tgt.y.inner());
builder.arg(tgt.z.inner());
builder.arg(&n);
unsafe { builder.launch(grid_1d(bt * di)) }
.map_err(|e| format!("silu_gate_fwd m3 tgt L{l}: {:?}", e))?;
}
gpu_sgemm_forward_raw(
ctx,
&mut tgt.out_flat,
&tgt.gated,
lw.out_proj_w.raw_ptr(&ctx.stream),
None,
(bt, di, dm),
)?;
{
let ne = (bt * dm) as i32;
let mut builder = ctx.stream.launch_builder(&m3k.residual_add);
builder.arg(tgt.temporal_work.inner_mut());
builder.arg(tgt.out_flat.inner());
builder.arg(tgt.residual.inner());
builder.arg(&ne);
unsafe { builder.launch(grid_1d(bt * dm)) }
.map_err(|e| format!("residual_add m3 tgt L{l}: {:?}", e))?;
}
}
{
tgt.residual.copy_from(&tgt.temporal_work, &ctx.stream)?;
let nf_ptr = mamba_w.norm_f_weight.raw_ptr(&ctx.stream);
let bt_i = bt as i32;
let dm_i = dm as i32;
let eps: f32 = 1e-5;
let mut builder = ctx.stream.launch_builder(&m3k.rmsnorm_fwd);
builder.arg(tgt.temporal_work.inner_mut());
builder.arg(tgt.rms_discard.inner_mut());
builder.arg(tgt.residual.inner());
builder.arg(&nf_ptr);
builder.arg(&bt_i);
builder.arg(&dm_i);
builder.arg(&eps);
unsafe { builder.launch(grid_norm(bt, dm)) }
.map_err(|e| format!("rmsnorm_fwd norm_f m3 tgt: {:?}", e))?;
}
{
let b_i = dims.batch as i32;
let t_i = dims.seq_len as i32;
let dm_i = dm as i32;
let mut builder = ctx.stream.launch_builder(&m3k.gather_last_timestep);
builder.arg(temporal.inner_mut());
builder.arg(tgt.temporal_work.inner());
builder.arg(&b_i);
builder.arg(&t_i);
builder.arg(&dm_i);
unsafe { builder.launch(grid_1d(dims.batch * dm)) }
.map_err(|e| format!("gather_last_timestep m3 tgt: {:?}", e))?;
}
Ok(())
}