use std::sync::Arc;
use cudarc::driver::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::context::GpuCtx;
use crate::mamba_ssm::gpu::dtype::WeightDtype;
use crate::mamba_ssm::gpu::forward::{GpuMambaDims, GpuRecurrentState};
use crate::mamba_ssm::gpu::launch::{grid_1d, grid_norm};
use crate::mamba_ssm::gpu::weights::GpuMambaMixedWeights;
use crate::mamba_ssm::gpu::weights_mixed_train::GpuMambaTrainMixedWeights;
pub struct GpuMambaLayerMixedActs {
pub residual: GpuBuffer,
pub rms_vals: GpuBuffer,
pub post_norm: DtypedBuf,
pub gate_pre_silu: DtypedBuf,
pub gate_post_silu: DtypedBuf,
pub conv_states: GpuBuffer,
pub post_conv: DtypedBuf,
pub u: DtypedBuf,
pub xdbl: DtypedBuf,
pub delta_raw: DtypedBuf,
pub delta: DtypedBuf,
pub h_saved: GpuBuffer,
pub da_exp: GpuBuffer,
pub y: DtypedBuf,
pub gated: DtypedBuf,
}
pub struct GpuMambaBackboneMixedActs {
pub input_proj_inputs: DtypedBuf,
pub input_proj_outputs: DtypedBuf,
pub layers: Vec<GpuMambaLayerMixedActs>,
pub norm_f_input: GpuBuffer,
pub norm_f_rms: GpuBuffer,
pub dtype: WeightDtype,
}
impl GpuMambaBackboneMixedActs {
pub fn new(
stream: &Arc<cudarc::driver::CudaStream>,
dims: &GpuMambaDims,
dtype: WeightDtype,
) -> Result<Self, String> {
let batch = dims.batch;
let seq_len = dims.seq_len;
let d_model = dims.d_model;
let d_inner = dims.d_inner;
let d_state = dims.d_state;
let d_conv = dims.d_conv;
let dt_rank = dims.dt_rank;
let n_layers = dims.n_layers;
let mamba_input_dim = dims.mamba_input_dim;
let bt = batch * seq_len;
let xdbl_dim = dt_rank + 2 * d_state;
let layers = (0..n_layers)
.map(|_| {
Ok(GpuMambaLayerMixedActs {
residual: GpuBuffer::zeros(stream, bt * d_model)?,
rms_vals: GpuBuffer::zeros(stream, bt)?,
conv_states: GpuBuffer::zeros(stream, bt * d_inner * d_conv)?,
h_saved: GpuBuffer::zeros(stream, batch * (seq_len + 1) * d_inner * d_state)?,
da_exp: GpuBuffer::zeros(stream, bt * d_inner * d_state)?,
post_norm: DtypedBuf::zeros(stream, bt * d_model, dtype)?,
gate_pre_silu: DtypedBuf::zeros(stream, bt * d_inner, dtype)?,
gate_post_silu: DtypedBuf::zeros(stream, bt * d_inner, dtype)?,
post_conv: DtypedBuf::zeros(stream, bt * d_inner, dtype)?,
u: DtypedBuf::zeros(stream, bt * d_inner, dtype)?,
xdbl: DtypedBuf::zeros(stream, bt * xdbl_dim, dtype)?,
delta_raw: DtypedBuf::zeros(stream, bt * d_inner, dtype)?,
delta: DtypedBuf::zeros(stream, bt * d_inner, dtype)?,
y: DtypedBuf::zeros(stream, bt * d_inner, dtype)?,
gated: DtypedBuf::zeros(stream, bt * d_inner, dtype)?,
})
})
.collect::<Result<Vec<_>, String>>()?;
let s = Self {
input_proj_inputs: DtypedBuf::zeros(stream, bt * mamba_input_dim, dtype)?,
input_proj_outputs: DtypedBuf::zeros(stream, bt * d_model, dtype)?,
layers,
norm_f_input: GpuBuffer::zeros(stream, bt * d_model)?,
norm_f_rms: GpuBuffer::zeros(stream, bt)?,
dtype,
};
stream
.synchronize()
.map_err(|e| format!("sync after mixed acts alloc: {e:?}"))?;
Ok(s)
}
}
pub struct GpuMambaMixedTrainScratch {
pub dims: GpuMambaDims,
pub dtype: WeightDtype,
pub proj_flat: DtypedBuf,
pub x_branch: DtypedBuf,
pub dt_gather: DtypedBuf,
pub b_buf: DtypedBuf,
pub c_buf: DtypedBuf,
pub out_flat: DtypedBuf,
pub temporal_typed: DtypedBuf,
pub d_gated: DtypedBuf,
pub d_y: DtypedBuf,
pub d_gate: DtypedBuf,
pub d_b_local: DtypedBuf,
pub d_c_local: DtypedBuf,
pub d_delta: DtypedBuf,
pub d_u: DtypedBuf,
pub d_u_xproj: DtypedBuf,
pub d_delta_raw: DtypedBuf,
pub dt_xdbl_buf: DtypedBuf,
pub d_dt_input: DtypedBuf,
pub d_xdbl: DtypedBuf,
pub d_x_branch: DtypedBuf,
pub d_proj: DtypedBuf,
pub d_norm: DtypedBuf,
pub d_b_reduced: GpuBuffer,
pub d_c_reduced: GpuBuffer,
pub d_d_local: GpuBuffer,
pub d_a_log_local: GpuBuffer,
pub d_pre_norm: GpuBuffer,
pub d_input_proj_dx: GpuBuffer,
pub axis0_partials: GpuBuffer,
}
impl GpuMambaMixedTrainScratch {
pub fn new(
stream: &Arc<cudarc::driver::CudaStream>,
dims: &GpuMambaDims,
dtype: WeightDtype,
) -> Result<Self, String> {
let bt = dims.batch * dims.seq_len;
let xdbl_dim = dims.dt_rank + 2 * dims.d_state;
let di = dims.d_inner;
let ds = dims.d_state;
let dm = dims.d_model;
let b = dims.batch;
let s = Self {
dims: *dims,
dtype,
proj_flat: DtypedBuf::zeros(stream, bt * 2 * di, dtype)?,
x_branch: DtypedBuf::zeros(stream, bt * di, dtype)?,
dt_gather: DtypedBuf::zeros(stream, bt * dims.dt_rank, dtype)?,
b_buf: DtypedBuf::zeros(stream, bt * ds, dtype)?,
c_buf: DtypedBuf::zeros(stream, bt * ds, dtype)?,
out_flat: DtypedBuf::zeros(stream, bt * dm, dtype)?,
temporal_typed: DtypedBuf::zeros(stream, bt * dm, dtype)?,
d_gated: DtypedBuf::zeros(stream, bt * di, dtype)?,
d_y: DtypedBuf::zeros(stream, bt * di, dtype)?,
d_gate: DtypedBuf::zeros(stream, bt * di, dtype)?,
d_b_local: DtypedBuf::zeros(stream, bt * di * ds, dtype)?,
d_c_local: DtypedBuf::zeros(stream, bt * di * ds, dtype)?,
d_delta: DtypedBuf::zeros(stream, bt * di, dtype)?,
d_u: DtypedBuf::zeros(stream, bt * di, dtype)?,
d_u_xproj: DtypedBuf::zeros(stream, bt * di, dtype)?,
d_delta_raw: DtypedBuf::zeros(stream, bt * di, dtype)?,
dt_xdbl_buf: DtypedBuf::zeros(stream, bt * dims.dt_rank, dtype)?,
d_dt_input: DtypedBuf::zeros(stream, bt * dims.dt_rank, dtype)?,
d_xdbl: DtypedBuf::zeros(stream, bt * xdbl_dim, dtype)?,
d_x_branch: DtypedBuf::zeros(stream, bt * di, dtype)?,
d_proj: DtypedBuf::zeros(stream, bt * 2 * di, dtype)?,
d_norm: DtypedBuf::zeros(stream, bt * dm, dtype)?,
d_b_reduced: GpuBuffer::zeros(stream, bt * ds)?,
d_c_reduced: GpuBuffer::zeros(stream, bt * ds)?,
d_d_local: GpuBuffer::zeros(stream, b * di)?,
d_a_log_local: GpuBuffer::zeros(stream, b * di * ds)?,
d_pre_norm: GpuBuffer::zeros(stream, bt * dm)?,
d_input_proj_dx: GpuBuffer::zeros(stream, bt * dims.mamba_input_dim)?,
axis0_partials: GpuBuffer::zeros(
stream,
std::cmp::max(bt * dm, b * di * (dims.d_conv + 1)),
)?,
};
stream
.synchronize()
.map_err(|e| format!("sync after mixed train scratch alloc: {e:?}"))?;
Ok(s)
}
}
pub fn gpu_forward_mamba_backbone_mixed(
ctx: &GpuCtx,
acts: &mut GpuMambaBackboneMixedActs,
mamba_w: &GpuMambaMixedWeights,
mamba_input: &GpuBuffer, state: &mut GpuRecurrentState,
scratch: &mut GpuMambaMixedTrainScratch,
) -> Result<(), String> {
assert_eq!(acts.dtype, scratch.dtype, "acts/scratch dtype mismatch");
assert_eq!(
acts.dtype, mamba_w.bulk_dtype,
"acts dtype must match weights bulk_dtype"
);
let dims = scratch.dims;
let bt = dims.batch * dims.seq_len;
let dm = dims.d_model;
let di = dims.d_inner;
let ds = dims.d_state;
let dt_rank = dims.dt_rank;
let xdbl_dim = dt_rank + 2 * ds;
let b = dims.batch;
let t = dims.seq_len;
let d_conv = dims.d_conv;
let dt = acts.dtype;
let k = &ctx.kernels;
if mamba_w.input_proj_w.len_elems() == 0 {
let bytes = bt * dm * 4;
let res = unsafe {
cudarc::driver::sys::cuMemcpyDtoDAsync_v2(
acts.layers[0].residual.cached_ptr(),
mamba_input.cached_ptr(),
bytes,
ctx.stream.cu_stream(),
)
};
if res != cudarc::driver::sys::CUresult::CUDA_SUCCESS {
return Err(format!("identity proj D2D failed: {res:?}"));
}
} else {
return Err(
"non-identity input_proj for mixed training not yet implemented \
(HF Mamba uses identity_proj — use from_hf path)"
.to_string(),
);
}
let conv_per_layer = b * di * d_conv;
let ssm_per_layer = b * di * ds;
let a_neg_per_layer = di * ds;
let f32_sz = std::mem::size_of::<f32>() as u64;
let conv_base = state.conv_states.cached_ptr();
let ssm_base = state.ssm_states.cached_ptr();
let aneg_base = state.a_neg_all.cached_ptr();
for layer_idx in 0..dims.n_layers {
let conv_ptr = conv_base + (layer_idx * conv_per_layer) as u64 * f32_sz;
let ssm_ptr = ssm_base + (layer_idx * ssm_per_layer) as u64 * f32_sz;
let aneg_ptr = aneg_base + (layer_idx * a_neg_per_layer) as u64 * f32_sz;
let lw = &mamba_w.layers[layer_idx];
let layer_acts = &mut acts.layers[layer_idx];
{
let bt_i = bt as i32;
let dm_i = dm as i32;
let eps: f32 = 1e-5;
let mut bld = ctx.stream.launch_builder(k.rmsnorm_fwd_f32in_typed.get(dt));
let pn_ptr = layer_acts.post_norm.cached_ptr();
let rms_ptr = layer_acts.rms_vals.cached_ptr();
let res_ptr = layer_acts.residual.cached_ptr();
bld.arg(&pn_ptr);
bld.arg(&rms_ptr);
bld.arg(&res_ptr);
let nw = lw.norm_weight.ptr();
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!("rmsnorm_f32in_typed L{layer_idx}: {e:?}"))?;
}
gpu_gemm_typed_forward_raw(
ctx,
TypedPtr {
ptr: scratch.proj_flat.cached_ptr(),
dtype: dt,
},
TypedPtr {
ptr: layer_acts.post_norm.cached_ptr(),
dtype: dt,
},
TypedPtr {
ptr: lw.in_proj_w.ptr(),
dtype: dt,
},
None,
(bt, dm, 2 * di),
)?;
{
let bt_i = bt as i32;
let di_i = di as i32;
let mut bld = ctx.stream.launch_builder(k.split_gate_silu_typed.get(dt));
let xb = scratch.x_branch.cached_ptr();
let gp = layer_acts.gate_pre_silu.cached_ptr();
let gs = layer_acts.gate_post_silu.cached_ptr();
let pf = scratch.proj_flat.cached_ptr();
bld.arg(&xb);
bld.arg(&gp);
bld.arg(&gs);
bld.arg(&pf);
bld.arg(&bt_i);
bld.arg(&di_i);
unsafe { bld.launch(grid_1d(bt * di)) }
.map_err(|e| format!("split_gate_silu_typed L{layer_idx}: {e:?}"))?;
}
{
let b_i = b as i32;
let t_i = t as i32;
let di_i = di as i32;
let dc_i = d_conv as i32;
let kernel = match dt {
WeightDtype::F32 => &k.conv1d_burnin_fwd_f32_typed,
WeightDtype::Bf16 => &k.conv1d_burnin_fwd_bf16,
WeightDtype::F16 => &k.conv1d_burnin_fwd_f16,
};
let mut bld = ctx.stream.launch_builder(kernel);
let u = layer_acts.u.cached_ptr();
let cs = layer_acts.conv_states.cached_ptr();
let pc = layer_acts.post_conv.cached_ptr();
let xb = scratch.x_branch.cached_ptr();
bld.arg(&u);
bld.arg(&conv_ptr); bld.arg(&cs);
bld.arg(&pc);
bld.arg(&xb);
let cw = lw.conv1d_weight.ptr();
let cb = lw.conv1d_bias.ptr();
bld.arg(&cw);
bld.arg(&cb);
bld.arg(&b_i);
bld.arg(&t_i);
bld.arg(&di_i);
bld.arg(&dc_i);
unsafe { bld.launch(grid_1d(b * di)) }
.map_err(|e| format!("conv1d_burnin_typed L{layer_idx}: {e:?}"))?;
}
gpu_gemm_typed_forward_raw(
ctx,
TypedPtr {
ptr: layer_acts.xdbl.cached_ptr(),
dtype: dt,
},
TypedPtr {
ptr: layer_acts.u.cached_ptr(),
dtype: dt,
},
TypedPtr {
ptr: lw.x_proj_w.ptr(),
dtype: dt,
},
None,
(bt, di, xdbl_dim),
)?;
{
let bt_i = bt as i32;
let xdbl_i = xdbl_dim as i32;
let dt_i = dt_rank as i32;
let offset: i32 = 0;
let mut bld = ctx.stream.launch_builder(k.gather_cols_typed.get(dt));
let dg = scratch.dt_gather.cached_ptr();
let xd = layer_acts.xdbl.cached_ptr();
bld.arg(&dg);
bld.arg(&xd);
bld.arg(&bt_i);
bld.arg(&xdbl_i);
bld.arg(&dt_i);
bld.arg(&offset);
unsafe { bld.launch(grid_1d(bt * dt_rank)) }
.map_err(|e| format!("gather_cols dt typed L{layer_idx}: {e:?}"))?;
}
gpu_gemm_typed_forward_raw(
ctx,
TypedPtr {
ptr: layer_acts.delta_raw.cached_ptr(),
dtype: dt,
},
TypedPtr {
ptr: scratch.dt_gather.cached_ptr(),
dtype: dt,
},
TypedPtr {
ptr: lw.dt_proj_w.ptr(),
dtype: dt,
},
Some(lw.dt_proj_b.ptr()),
(bt, dt_rank, di),
)?;
{
let n = (bt * di) as i32;
let mut bld = ctx.stream.launch_builder(k.softplus_copy_typed.get(dt));
let dl = layer_acts.delta.cached_ptr();
let dr = layer_acts.delta_raw.cached_ptr();
bld.arg(&dl);
bld.arg(&dr);
bld.arg(&n);
unsafe { bld.launch(grid_1d(bt * di)) }
.map_err(|e| format!("softplus_copy_typed L{layer_idx}: {e:?}"))?;
}
{
let bt_i = bt as i32;
let xdbl_i = xdbl_dim as i32;
let ds_i = ds as i32;
let b_off = dt_rank as i32;
let c_off = (dt_rank + ds) as i32;
let mut bld = ctx.stream.launch_builder(k.gather_bc_cols_typed.get(dt));
let bb = scratch.b_buf.cached_ptr();
let cb = scratch.c_buf.cached_ptr();
let xd = layer_acts.xdbl.cached_ptr();
bld.arg(&bb);
bld.arg(&cb);
bld.arg(&xd);
bld.arg(&bt_i);
bld.arg(&xdbl_i);
bld.arg(&ds_i);
bld.arg(&b_off);
bld.arg(&c_off);
unsafe { bld.launch(grid_1d(bt * ds)) }
.map_err(|e| format!("gather_bc_cols typed L{layer_idx}: {e:?}"))?;
}
{
let b_i = b as i32;
let t_i = t as i32;
let di_i = di as i32;
let ds_i = ds as i32;
if t > super::forward::PARALLEL_SCAN_THRESHOLD || ds > 64 {
let kernel = k.ssm_parallel_fwd_typed.get(dt);
let mut bld = ctx.stream.launch_builder(kernel);
let y = layer_acts.y.cached_ptr();
let hs = layer_acts.h_saved.cached_ptr();
let dae = layer_acts.da_exp.cached_ptr();
let dl = layer_acts.delta.cached_ptr();
let u = layer_acts.u.cached_ptr();
let bb = scratch.b_buf.cached_ptr();
let cb = scratch.c_buf.cached_ptr();
let dp = lw.d_param.ptr();
bld.arg(&ssm_ptr);
bld.arg(&y);
bld.arg(&hs);
bld.arg(&dae);
bld.arg(&dl);
bld.arg(&u);
bld.arg(&bb);
bld.arg(&cb);
bld.arg(&aneg_ptr);
bld.arg(&dp);
bld.arg(&b_i);
bld.arg(&t_i);
bld.arg(&di_i);
bld.arg(&ds_i);
unsafe {
bld.launch(super::launch::grid_parallel_scan_typed(
b,
di,
dt.size_bytes(),
))
}
.map_err(|e| format!("ssm_parallel_fwd typed L{layer_idx}: {e:?}"))?;
} else {
assert!(
ds <= 64,
"ssm_burnin_forward_typed requires d_state <= 64 (got {ds})"
);
let kernel = match dt {
WeightDtype::F32 => &k.ssm_burnin_fwd,
WeightDtype::Bf16 => &k.ssm_burnin_fwd_bf16,
WeightDtype::F16 => &k.ssm_burnin_fwd_f16,
};
let mut bld = ctx.stream.launch_builder(kernel);
let y = layer_acts.y.cached_ptr();
let hs = layer_acts.h_saved.cached_ptr();
let dae = layer_acts.da_exp.cached_ptr();
let dl = layer_acts.delta.cached_ptr();
let u = layer_acts.u.cached_ptr();
let bb = scratch.b_buf.cached_ptr();
let cb = scratch.c_buf.cached_ptr();
let dp = lw.d_param.ptr();
bld.arg(&ssm_ptr);
bld.arg(&y);
bld.arg(&hs);
bld.arg(&dae);
bld.arg(&dl);
bld.arg(&u);
bld.arg(&bb);
bld.arg(&cb);
bld.arg(&aneg_ptr);
bld.arg(&dp);
bld.arg(&b_i);
bld.arg(&t_i);
bld.arg(&di_i);
bld.arg(&ds_i);
unsafe { bld.launch(grid_1d(b * di)) }
.map_err(|e| format!("ssm_burnin_forward typed L{layer_idx}: {e:?}"))?;
}
}
{
let n = (bt * di) as i32;
let mut bld = ctx.stream.launch_builder(k.elementwise_mul_typed.get(dt));
let g = layer_acts.gated.cached_ptr();
let y = layer_acts.y.cached_ptr();
let gs = layer_acts.gate_post_silu.cached_ptr();
bld.arg(&g);
bld.arg(&y);
bld.arg(&gs);
bld.arg(&n);
unsafe { bld.launch(grid_1d(bt * di)) }
.map_err(|e| format!("elementwise_mul_typed L{layer_idx}: {e:?}"))?;
}
gpu_gemm_typed_forward_raw(
ctx,
TypedPtr {
ptr: scratch.out_flat.cached_ptr(),
dtype: dt,
},
TypedPtr {
ptr: layer_acts.gated.cached_ptr(),
dtype: dt,
},
TypedPtr {
ptr: lw.out_proj_w.ptr(),
dtype: dt,
},
None,
(bt, di, dm),
)?;
let next_res_ptr = if layer_idx + 1 < dims.n_layers {
acts.layers[layer_idx + 1].residual.cached_ptr()
} else {
acts.norm_f_input.cached_ptr()
};
{
let n = (bt * dm) as i32;
let mut bld = ctx.stream.launch_builder(k.residual_add_f32_typed.get(dt));
let cur_res = acts.layers[layer_idx].residual.cached_ptr();
let of = scratch.out_flat.cached_ptr();
bld.arg(&next_res_ptr); bld.arg(&cur_res); bld.arg(&of); bld.arg(&n);
unsafe { bld.launch(grid_1d(bt * dm)) }
.map_err(|e| format!("residual_add_f32_typed L{layer_idx}: {e:?}"))?;
}
}
{
let bt_i = bt as i32;
let dm_i = dm as i32;
let eps: f32 = 1e-5;
let mut bld = ctx.stream.launch_builder(k.rmsnorm_fwd_f32in_typed.get(dt));
let tt = scratch.temporal_typed.cached_ptr();
let nfr = acts.norm_f_rms.cached_ptr();
let nfi = acts.norm_f_input.cached_ptr();
bld.arg(&tt);
bld.arg(&nfr);
bld.arg(&nfi);
let nfw = mamba_w.norm_f_weight.ptr();
bld.arg(&nfw);
bld.arg(&bt_i);
bld.arg(&dm_i);
bld.arg(&eps);
unsafe { bld.launch(grid_norm(bt, dm)) }
.map_err(|e| format!("rmsnorm_f32in_typed norm_f: {e:?}"))?;
}
Ok(())
}
pub fn gpu_forward_mamba_backbone_train_mixed(
ctx: &GpuCtx,
acts: &mut GpuMambaBackboneMixedActs,
train_w: &GpuMambaTrainMixedWeights,
mamba_input: &GpuBuffer,
state: &mut GpuRecurrentState,
scratch: &mut GpuMambaMixedTrainScratch,
) -> Result<(), String> {
gpu_forward_mamba_backbone_mixed(ctx, acts, &train_w.compute, mamba_input, state, scratch)
}