use super::blas::gpu_sgemm_backward_grad_raw;
use super::buffers::GpuBuffer;
use super::context::GpuCtx;
use super::forward::{GpuMambaBackboneActs, GpuMambaLayerActs, GpuMambaScratch};
use super::launch::{grid_1d, grid_norm};
use super::weights::{
GpuMambaGrads, GpuMambaLayerGrads, GpuMambaTrainLayerWeights, GpuMambaTrainWeights,
};
use cudarc::driver::PushKernelArg;
use std::sync::Arc;
pub fn gpu_backward_mamba_layer(
ctx: &GpuCtx,
d_temporal: &mut GpuBuffer,
d_lw: &GpuMambaLayerGrads,
acts: &GpuMambaLayerActs,
lw: &GpuMambaTrainLayerWeights,
a_neg_ptr: cudarc::driver::sys::CUdeviceptr,
scratch: &mut GpuMambaScratch,
) -> Result<(), String> {
let dims = scratch.dims; let bt = dims.bt();
let dm = dims.d_model;
let di = dims.d_inner;
let ds = dims.d_state;
let dt_rank = dims.dt_rank;
let xdbl_dim = dims.xdbl_dim;
let b = dims.batch;
let t = dims.seq_len;
let d_conv = dims.d_conv;
gpu_sgemm_backward_grad_raw(
ctx,
&mut scratch.d_gated,
(&d_lw.out_proj_w, None),
d_temporal,
&acts.gated,
lw.out_proj_w.cached_ptr(),
(bt, di, dm),
)?;
{
let n = (bt * di) as i32;
let mut builder = ctx.stream.launch_builder(&ctx.kernels.gating_backward);
builder.arg(scratch.d_y.inner_mut());
builder.arg(scratch.d_gate.inner_mut());
builder.arg(scratch.d_gated.inner());
builder.arg(acts.y.inner());
builder.arg(acts.gate_pre_silu.inner());
builder.arg(acts.gate_post_silu.inner());
builder.arg(&n);
unsafe { builder.launch(grid_1d(bt * di)) }
.map_err(|e| format!("gating_backward mamba: {:?}", e))?;
}
{
let bt_i = bt as i32;
let xdbl_i = xdbl_dim as i32;
let ds_i = ds as i32;
let b_offset = dt_rank as i32;
let c_offset = (dt_rank + ds) as i32;
let mut builder = ctx.stream.launch_builder(&ctx.kernels.gather_bc_cols);
builder.arg(scratch.d_b_reduced.inner_mut());
builder.arg(scratch.d_c_reduced.inner_mut());
builder.arg(acts.xdbl.inner());
builder.arg(&bt_i);
builder.arg(&xdbl_i);
builder.arg(&ds_i);
builder.arg(&b_offset);
builder.arg(&c_offset);
unsafe { builder.launch(grid_1d(bt * ds)) }
.map_err(|e| format!("gather_bc_cols bwd mamba: {:?}", e))?;
}
scratch.d_a_log_local.zero(&ctx.stream)?;
{
let b_i = b as i32;
let t_i = t as i32;
let di_i = di as i32;
let ds_i = ds as i32;
let use_parallel = dims.scan_mode.use_parallel(t, ds);
let kernel = if use_parallel {
ctx.kernels
.ssm_parallel_bwd_typed
.get(super::dtype::WeightDtype::F32)
} else {
&ctx.kernels.ssm_backward_local
};
let mut builder = ctx.stream.launch_builder(kernel);
builder.arg(acts.h_saved.inner());
builder.arg(acts.delta.inner());
builder.arg(acts.u.inner());
builder.arg(scratch.d_b_reduced.inner()); builder.arg(scratch.d_c_reduced.inner()); builder.arg(&a_neg_ptr); let dp_ptr = lw.d_param.cached_ptr();
builder.arg(&dp_ptr);
builder.arg(scratch.d_y.inner()); builder.arg(scratch.d_delta.inner_mut());
builder.arg(scratch.d_u.inner_mut());
builder.arg(scratch.d_b_local.inner_mut());
builder.arg(scratch.d_c_local.inner_mut());
builder.arg(scratch.d_d_local.inner_mut());
builder.arg(scratch.d_a_log_local.inner_mut());
builder.arg(&b_i);
builder.arg(&t_i);
builder.arg(&di_i);
builder.arg(&ds_i);
let cfg = if use_parallel {
super::launch::grid_parallel_scan_bwd(b, di)
} else {
grid_1d(b * di)
};
unsafe { builder.launch(cfg) }
.map_err(|e| format!("ssm bwd (parallel={use_parallel}): {e:?}"))?;
}
scratch.d_b_reduced.zero(&ctx.stream)?;
scratch.d_c_reduced.zero(&ctx.stream)?;
{
let b_i = b as i32;
let t_i = t as i32;
let di_i = di as i32;
let ds_i = ds as i32;
let mut builder = ctx.stream.launch_builder(&ctx.kernels.ssm_reduce_d_b);
builder.arg(scratch.d_b_reduced.inner_mut());
builder.arg(scratch.d_b_local.inner());
builder.arg(&b_i);
builder.arg(&t_i);
builder.arg(&di_i);
builder.arg(&ds_i);
unsafe { builder.launch(grid_1d(bt * ds)) }
.map_err(|e| format!("ssm_reduce_d_B mamba: {:?}", e))?;
let mut builder = ctx.stream.launch_builder(&ctx.kernels.ssm_reduce_d_c);
builder.arg(scratch.d_c_reduced.inner_mut());
builder.arg(scratch.d_c_local.inner());
builder.arg(&b_i);
builder.arg(&t_i);
builder.arg(&di_i);
builder.arg(&ds_i);
unsafe { builder.launch(grid_1d(bt * ds)) }
.map_err(|e| format!("ssm_reduce_d_C mamba: {:?}", e))?;
let mut builder = ctx.stream.launch_builder(&ctx.kernels.ssm_reduce_d_d);
let _p = d_lw.d_param.ptr();
builder.arg(&_p);
builder.arg(scratch.d_d_local.inner());
builder.arg(&b_i);
builder.arg(&di_i);
unsafe { builder.launch(grid_1d(di)) }
.map_err(|e| format!("ssm_reduce_d_D mamba: {:?}", e))?;
let mut builder = ctx.stream.launch_builder(&ctx.kernels.ssm_reduce_d_a_log);
let _p = d_lw.a_log.ptr();
builder.arg(&_p);
builder.arg(scratch.d_a_log_local.inner());
builder.arg(&b_i);
builder.arg(&di_i);
builder.arg(&ds_i);
unsafe { builder.launch(grid_1d(di * ds)) }
.map_err(|e| format!("ssm_reduce_d_a_log mamba: {:?}", e))?;
}
scratch.d_xdbl.zero(&ctx.stream)?;
{
let bt_i = bt as i32;
let xdbl_i = xdbl_dim as i32;
let ds_i = ds as i32;
let b_offset = dt_rank as i32;
let c_offset = (dt_rank + ds) as i32;
let mut builder = ctx.stream.launch_builder(&ctx.kernels.scatter_add_cols);
builder.arg(scratch.d_xdbl.inner_mut());
builder.arg(scratch.d_b_reduced.inner());
builder.arg(&bt_i);
builder.arg(&xdbl_i);
builder.arg(&ds_i);
builder.arg(&b_offset);
unsafe { builder.launch(grid_1d(bt * ds)) }
.map_err(|e| format!("scatter d_B mamba: {:?}", e))?;
let mut builder = ctx.stream.launch_builder(&ctx.kernels.scatter_add_cols);
builder.arg(scratch.d_xdbl.inner_mut());
builder.arg(scratch.d_c_reduced.inner());
builder.arg(&bt_i);
builder.arg(&xdbl_i);
builder.arg(&ds_i);
builder.arg(&c_offset);
unsafe { builder.launch(grid_1d(bt * ds)) }
.map_err(|e| format!("scatter d_C mamba: {:?}", e))?;
}
{
let n = (bt * di) as i32;
let mut builder = ctx.stream.launch_builder(&ctx.kernels.softplus_bwd);
builder.arg(scratch.d_delta_raw.inner_mut()); builder.arg(acts.delta_raw.inner()); builder.arg(scratch.d_delta.inner()); builder.arg(&n);
unsafe { builder.launch(grid_1d(bt * di)) }
.map_err(|e| format!("softplus_bwd mamba: {:?}", e))?;
}
{
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 builder = ctx.stream.launch_builder(&ctx.kernels.gather_cols);
builder.arg(scratch.dt_xdbl_buf.inner_mut());
builder.arg(acts.xdbl.inner());
builder.arg(&bt_i);
builder.arg(&xdbl_i);
builder.arg(&dt_i);
builder.arg(&offset);
unsafe { builder.launch(grid_1d(bt * dt_rank)) }
.map_err(|e| format!("gather dt x_saved mamba: {:?}", e))?;
}
gpu_sgemm_backward_grad_raw(
ctx,
&mut scratch.d_dt_input, (&d_lw.dt_proj_w, Some(&d_lw.dt_proj_b)),
&scratch.d_delta_raw, &scratch.dt_xdbl_buf, lw.dt_proj_w.cached_ptr(),
(bt, dt_rank, di),
)?;
{
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 builder = ctx.stream.launch_builder(&ctx.kernels.scatter_add_cols);
builder.arg(scratch.d_xdbl.inner_mut());
builder.arg(scratch.d_dt_input.inner());
builder.arg(&bt_i);
builder.arg(&xdbl_i);
builder.arg(&dt_i);
builder.arg(&offset);
unsafe { builder.launch(grid_1d(bt * dt_rank)) }
.map_err(|e| format!("scatter dt bwd mamba: {:?}", e))?;
}
gpu_sgemm_backward_grad_raw(
ctx,
&mut scratch.d_u_xproj,
(&d_lw.x_proj_w, None),
&scratch.d_xdbl,
&acts.u,
lw.x_proj_w.cached_ptr(),
(bt, di, xdbl_dim),
)?;
{
let n = (bt * di) as i32;
let mut builder = ctx.stream.launch_builder(&ctx.kernels.vec_add_inplace);
builder.arg(scratch.d_u.inner_mut());
builder.arg(scratch.d_u_xproj.inner());
builder.arg(&n);
unsafe { builder.launch(grid_1d(bt * di)) }
.map_err(|e| format!("vec_add d_u mamba: {:?}", 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 weight_partials_elems = b * di * d_conv;
let bias_offset_bytes = (weight_partials_elems * std::mem::size_of::<f32>()) as u64;
let axis0_base = scratch.axis0_partials.cached_ptr();
let wp_ptr = axis0_base;
let bp_ptr = axis0_base + bias_offset_bytes;
{
let mut builder = ctx.stream.launch_builder(&ctx.kernels.conv1d_burnin_bwd);
builder.arg(scratch.d_x_branch.inner_mut());
builder.arg(&wp_ptr); builder.arg(&bp_ptr); builder.arg(scratch.d_u.inner());
builder.arg(acts.post_conv.inner());
builder.arg(acts.conv_states.inner());
let cw_ptr = lw.conv1d_weight.cached_ptr();
builder.arg(&cw_ptr);
builder.arg(&b_i);
builder.arg(&t_i);
builder.arg(&di_i);
builder.arg(&dc_i);
unsafe { builder.launch(grid_1d(b * di)) }
.map_err(|e| format!("conv1d_burnin_bwd mamba partial: {:?}", e))?;
}
{
let block_dim = (b as u32).next_power_of_two().clamp(32, 256);
let accumulate_i: i32 = 1;
let dim_w = (di * d_conv) as i32;
let p = d_lw.conv1d_weight.ptr();
let mut builder = ctx.stream.launch_builder(&ctx.kernels.reduce_sum_axis0);
builder.arg(&p);
builder.arg(&wp_ptr);
builder.arg(&b_i);
builder.arg(&dim_w);
builder.arg(&accumulate_i);
let cfg = cudarc::driver::LaunchConfig {
grid_dim: ((di * d_conv) as u32, 1, 1),
block_dim: (block_dim, 1, 1),
shared_mem_bytes: (block_dim as usize * std::mem::size_of::<f32>()) as u32,
};
unsafe { builder.launch(cfg) }
.map_err(|e| format!("conv1d_burnin_bwd mamba weight final: {:?}", e))?;
}
{
let block_dim = (b as u32).next_power_of_two().clamp(32, 256);
let accumulate_i: i32 = 1;
let p = d_lw.conv1d_bias.ptr();
let mut builder = ctx.stream.launch_builder(&ctx.kernels.reduce_sum_axis0);
builder.arg(&p);
builder.arg(&bp_ptr);
builder.arg(&b_i);
builder.arg(&di_i);
builder.arg(&accumulate_i);
let cfg = cudarc::driver::LaunchConfig {
grid_dim: (di as u32, 1, 1),
block_dim: (block_dim, 1, 1),
shared_mem_bytes: (block_dim as usize * std::mem::size_of::<f32>()) as u32,
};
unsafe { builder.launch(cfg) }
.map_err(|e| format!("conv1d_burnin_bwd mamba bias final: {:?}", e))?;
}
}
{
let bt_i = bt as i32;
let di_i = di as i32;
let mut builder = ctx.stream.launch_builder(&ctx.kernels.concat_halves);
builder.arg(scratch.d_proj.inner_mut());
builder.arg(scratch.d_x_branch.inner());
builder.arg(scratch.d_gate.inner());
builder.arg(&bt_i);
builder.arg(&di_i);
unsafe { builder.launch(grid_1d(bt * di)) }
.map_err(|e| format!("concat_halves bwd mamba: {:?}", e))?;
}
gpu_sgemm_backward_grad_raw(
ctx,
&mut scratch.d_norm,
(&d_lw.in_proj_w, None),
&scratch.d_proj,
&acts.post_norm,
lw.in_proj_w.cached_ptr(),
(bt, dm, 2 * di),
)?;
{
let bt_i = bt as i32;
let dm_i = dm as i32;
let nw_ptr = lw.norm_weight.cached_ptr();
let axis0_ptr = scratch.axis0_partials.cached_ptr();
{
let mut builder = ctx.stream.launch_builder(&ctx.kernels.rmsnorm_bwd);
builder.arg(scratch.d_pre_norm.inner_mut());
builder.arg(&axis0_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 mamba partial: {:?}", e))?;
}
{
let block_dim = (bt as u32).next_power_of_two().clamp(32, 256);
let accumulate_i: i32 = 1;
let p = d_lw.norm_weight.ptr();
let mut builder = ctx.stream.launch_builder(&ctx.kernels.reduce_sum_axis0);
builder.arg(&p);
builder.arg(&axis0_ptr);
builder.arg(&bt_i);
builder.arg(&dm_i);
builder.arg(&accumulate_i);
let cfg = cudarc::driver::LaunchConfig {
grid_dim: (dm as u32, 1, 1),
block_dim: (block_dim, 1, 1),
shared_mem_bytes: (block_dim as usize * std::mem::size_of::<f32>()) as u32,
};
unsafe { builder.launch(cfg) }
.map_err(|e| format!("rmsnorm_bwd mamba final: {:?}", e))?;
}
}
{
let n = (bt * dm) as i32;
let mut builder = ctx.stream.launch_builder(&ctx.kernels.vec_add_inplace);
builder.arg(d_temporal.inner_mut());
builder.arg(scratch.d_pre_norm.inner());
builder.arg(&n);
unsafe { builder.launch(grid_1d(bt * dm)) }
.map_err(|e| format!("vec_add residual bwd mamba: {:?}", e))?;
}
Ok(())
}
pub fn gpu_backward_mamba_backbone(
ctx: &GpuCtx,
d_temporal: &mut GpuBuffer,
d_mamba: &GpuMambaGrads,
acts: &GpuMambaBackboneActs,
mamba_w: &GpuMambaTrainWeights,
a_neg_all: &GpuBuffer,
scratch: &mut GpuMambaScratch,
) -> Result<(), String> {
let dims = scratch.dims; let bt = dims.bt();
{
let bt_i = bt as i32;
let dm_i = dims.d_model as i32;
let axis0_ptr = scratch.axis0_partials.cached_ptr();
{
let mut builder = ctx.stream.launch_builder(&ctx.kernels.rmsnorm_bwd);
builder.arg(scratch.d_norm.inner_mut()); builder.arg(&axis0_ptr); builder.arg(d_temporal.inner()); builder.arg(acts.norm_f_input.inner()); let nf_ptr = mamba_w.norm_f_weight.cached_ptr();
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, dims.d_model)) }
.map_err(|e| format!("rmsnorm_bwd norm_f partial: {:?}", e))?;
}
{
let block_dim = (bt as u32).next_power_of_two().clamp(32, 256);
let accumulate_i: i32 = 1;
let p = d_mamba.norm_f_weight.ptr();
let mut builder = ctx.stream.launch_builder(&ctx.kernels.reduce_sum_axis0);
builder.arg(&p);
builder.arg(&axis0_ptr);
builder.arg(&bt_i);
builder.arg(&dm_i);
builder.arg(&accumulate_i);
let cfg = cudarc::driver::LaunchConfig {
grid_dim: (dims.d_model as u32, 1, 1),
block_dim: (block_dim, 1, 1),
shared_mem_bytes: (block_dim as usize * std::mem::size_of::<f32>()) as u32,
};
unsafe { builder.launch(cfg) }
.map_err(|e| format!("rmsnorm_bwd norm_f final: {:?}", e))?;
}
d_temporal.copy_from(&scratch.d_norm, &ctx.stream)?;
}
let a_neg_per_layer = dims.d_inner * dims.d_state;
for layer_idx in (0..dims.n_layers).rev() {
let base = a_neg_all.raw_ptr(&ctx.stream);
let a_neg_ptr = base + (layer_idx * a_neg_per_layer * std::mem::size_of::<f32>()) as u64;
gpu_backward_mamba_layer(
ctx,
d_temporal,
&d_mamba.layers[layer_idx],
&acts.layers[layer_idx],
&mamba_w.layers[layer_idx],
a_neg_ptr,
scratch,
)?;
}
gpu_sgemm_backward_grad_raw(
ctx,
&mut scratch.d_input_proj_dx,
(&d_mamba.input_proj_w, Some(&d_mamba.input_proj_b)),
d_temporal,
&acts.input_proj_inputs,
mamba_w.input_proj_w.cached_ptr(),
(bt, dims.mamba_input_dim, dims.d_model),
)?;
Ok(())
}
pub struct GpuMambaTargetScratch {
pub proj_flat: GpuBuffer, pub x_branch: GpuBuffer, pub gate_silu: GpuBuffer, pub u: GpuBuffer, pub xdbl: GpuBuffer, pub dt_gather: GpuBuffer, pub delta: GpuBuffer, pub y: GpuBuffer, pub gated: GpuBuffer, pub out_flat: GpuBuffer, pub residual: GpuBuffer, pub rms_discard: GpuBuffer, pub b_gathered: GpuBuffer, pub c_gathered: GpuBuffer, pub conv_states: GpuBuffer, pub ssm_states: GpuBuffer, pub dims: super::forward::GpuMambaDims,
}
impl GpuMambaTargetScratch {
pub fn new(
stream: &Arc<cudarc::driver::CudaStream>,
dims: &super::forward::GpuMambaDims,
) -> Result<Self, String> {
let batch = dims.batch;
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 bt = batch * dims.seq_len;
let xdbl_dim = dt_rank + 2 * d_state;
Ok(Self {
proj_flat: GpuBuffer::zeros(stream, bt * 2 * d_inner)?,
x_branch: GpuBuffer::zeros(stream, bt * d_inner)?,
gate_silu: GpuBuffer::zeros(stream, bt * d_inner)?,
u: GpuBuffer::zeros(stream, bt * d_inner)?,
xdbl: GpuBuffer::zeros(stream, bt * xdbl_dim)?,
dt_gather: GpuBuffer::zeros(stream, bt * dt_rank)?,
delta: GpuBuffer::zeros(stream, bt * d_inner)?,
y: GpuBuffer::zeros(stream, bt * d_inner)?,
gated: GpuBuffer::zeros(stream, bt * d_inner)?,
out_flat: GpuBuffer::zeros(stream, bt * d_model)?,
residual: GpuBuffer::zeros(stream, bt * d_model)?,
rms_discard: GpuBuffer::zeros(stream, bt)?,
b_gathered: GpuBuffer::zeros(stream, bt * d_state)?,
c_gathered: GpuBuffer::zeros(stream, bt * d_state)?,
conv_states: GpuBuffer::zeros(stream, batch * n_layers * d_inner * d_conv)?,
ssm_states: GpuBuffer::zeros(stream, batch * n_layers * d_inner * d_state)?,
dims: *dims,
})
}
}
use super::buffers::DtypedBuf;
use super::dtype::WeightDtype;
pub struct GpuMambaTargetMixedScratch {
pub proj_flat: DtypedBuf,
pub x_branch: DtypedBuf,
pub gate_silu: DtypedBuf,
pub u: DtypedBuf,
pub xdbl: DtypedBuf,
pub dt_gather: DtypedBuf,
pub delta: DtypedBuf,
pub y: DtypedBuf,
pub gated: DtypedBuf,
pub out_flat: DtypedBuf,
pub residual: GpuBuffer,
pub rms_discard: GpuBuffer,
pub b_gathered: DtypedBuf,
pub c_gathered: DtypedBuf,
pub dims: super::forward::GpuMambaDims,
pub dtype: WeightDtype,
}
impl GpuMambaTargetMixedScratch {
pub fn new(
stream: &Arc<cudarc::driver::CudaStream>,
dims: &super::forward::GpuMambaDims,
dtype: WeightDtype,
) -> Result<Self, String> {
if matches!(dtype, WeightDtype::F32) {
return Err("GpuMambaTargetMixedScratch requires bf16 or f16 dtype".to_string());
}
let batch = dims.batch;
let d_model = dims.d_model;
let d_inner = dims.d_inner;
let d_state = dims.d_state;
let dt_rank = dims.dt_rank;
let bt = batch * dims.seq_len;
let xdbl_dim = dt_rank + 2 * d_state;
Ok(Self {
proj_flat: DtypedBuf::zeros(stream, bt * 2 * d_inner, dtype)?,
x_branch: DtypedBuf::zeros(stream, bt * d_inner, dtype)?,
gate_silu: DtypedBuf::zeros(stream, bt * d_inner, dtype)?,
u: DtypedBuf::zeros(stream, bt * d_inner, dtype)?,
xdbl: DtypedBuf::zeros(stream, bt * xdbl_dim, dtype)?,
dt_gather: DtypedBuf::zeros(stream, bt * dt_rank, 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)?,
out_flat: DtypedBuf::zeros(stream, bt * d_model, dtype)?,
residual: GpuBuffer::zeros(stream, bt * d_model)?,
rms_discard: GpuBuffer::zeros(stream, bt)?,
b_gathered: DtypedBuf::zeros(stream, bt * d_state, dtype)?,
c_gathered: DtypedBuf::zeros(stream, bt * d_state, dtype)?,
dims: *dims,
dtype,
})
}
}