use super::blas::{TypedPtr, gpu_gemm_ex_backward_dx_typed, gpu_sgemm_backward_dw_grad_typed};
use super::buffers::GpuBuffer;
use super::context::GpuCtx;
use super::dtype::WeightDtype;
use super::forward_mixed::{
GpuMambaBackboneMixedActs, GpuMambaLayerMixedActs, GpuMambaMixedTrainScratch,
};
use super::launch::{grid_1d, grid_norm};
use super::weights::{
GpuMambaGrads, GpuMambaLayerGrads, GpuMambaMixedLayerWeights, GpuMambaMixedWeights,
};
use cudarc::driver::PushKernelArg;
#[derive(Clone, Copy)]
pub struct MixedLayerBwd<'a> {
pub d_lw: &'a GpuMambaLayerGrads,
pub acts: &'a GpuMambaLayerMixedActs,
pub lw: &'a GpuMambaMixedLayerWeights,
}
pub fn gpu_backward_mamba_layer_mixed(
ctx: &GpuCtx,
d_temporal: &mut GpuBuffer,
layer: &MixedLayerBwd<'_>,
a_neg_ptr: cudarc::driver::sys::CUdeviceptr,
scratch: &mut GpuMambaMixedTrainScratch,
dtype: WeightDtype,
) -> Result<(), String> {
let MixedLayerBwd { d_lw, acts, lw } = *layer;
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;
let k = &ctx.kernels;
scratch.temporal_typed.zero(&ctx.stream)?;
{
let n = (bt * dm) as i32;
let mut bld = ctx
.stream
.launch_builder(k.vec_add_inplace_typed.get(dtype));
let a = scratch.temporal_typed.cached_ptr();
let b_p = d_temporal.cached_ptr();
bld.arg(&a);
bld.arg(&b_p);
bld.arg(&n);
unsafe { bld.launch(grid_1d(bt * dm)) }
.map_err(|e| format!("cast d_temporal→typed: {e:?}"))?;
}
gpu_sgemm_backward_dw_grad_typed(
ctx,
&d_lw.out_proj_w,
TypedPtr {
ptr: scratch.temporal_typed.cached_ptr(),
dtype,
},
TypedPtr {
ptr: acts.gated.cached_ptr(),
dtype,
},
bt,
di,
dm,
)?;
gpu_gemm_ex_backward_dx_typed(
ctx,
TypedPtr {
ptr: scratch.d_gated.cached_ptr(),
dtype,
},
TypedPtr {
ptr: scratch.temporal_typed.cached_ptr(),
dtype,
},
TypedPtr {
ptr: lw.out_proj_w.ptr(),
dtype,
},
bt,
di,
dm,
)?;
{
let n = (bt * di) as i32;
let mut bld = ctx.stream.launch_builder(k.gating_bwd_typed.get(dtype));
let dy = scratch.d_y.cached_ptr();
let dg = scratch.d_gate.cached_ptr();
let dgin = scratch.d_gated.cached_ptr();
let y = acts.y.cached_ptr();
let gp = acts.gate_pre_silu.cached_ptr();
let gs = acts.gate_post_silu.cached_ptr();
bld.arg(&dy);
bld.arg(&dg);
bld.arg(&dgin);
bld.arg(&y);
bld.arg(&gp);
bld.arg(&gs);
bld.arg(&n);
unsafe { bld.launch(grid_1d(bt * di)) }.map_err(|e| format!("gating_bwd_typed: {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(dtype));
let bb = scratch.b_buf.cached_ptr();
let cb = scratch.c_buf.cached_ptr();
let xd = 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 bwd: {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 h_p = acts.h_saved.cached_ptr();
let delta_p = acts.delta.cached_ptr();
let u_p = acts.u.cached_ptr();
let b_p = scratch.b_buf.cached_ptr();
let c_p = scratch.c_buf.cached_ptr();
let dp = lw.d_param.ptr();
let dy = scratch.d_y.cached_ptr();
let dd = scratch.d_delta.cached_ptr();
let du = scratch.d_u.cached_ptr();
let dbl = scratch.d_b_local.cached_ptr();
let dcl = scratch.d_c_local.cached_ptr();
let ddd = scratch.d_d_local.cached_ptr();
let da = scratch.d_a_log_local.cached_ptr();
if dims.scan_mode.use_parallel(t, ds) {
let mut bld = ctx
.stream
.launch_builder(k.ssm_parallel_bwd_typed.get(dtype));
bld.arg(&h_p);
bld.arg(&delta_p);
bld.arg(&u_p);
bld.arg(&b_p);
bld.arg(&c_p);
bld.arg(&a_neg_ptr);
bld.arg(&dp);
bld.arg(&dy);
bld.arg(&dd);
bld.arg(&du);
bld.arg(&dbl);
bld.arg(&dcl);
bld.arg(&ddd);
bld.arg(&da);
bld.arg(&b_i);
bld.arg(&t_i);
bld.arg(&di_i);
bld.arg(&ds_i);
unsafe { bld.launch(super::launch::grid_parallel_scan_bwd(b, di)) }
.map_err(|e| format!("ssm_parallel_bwd_typed: {e:?}"))?;
} else {
let mut bld = ctx
.stream
.launch_builder(k.ssm_backward_local_typed.get(dtype));
bld.arg(&h_p);
bld.arg(&delta_p);
bld.arg(&u_p);
bld.arg(&b_p);
bld.arg(&c_p);
bld.arg(&a_neg_ptr);
bld.arg(&dp);
bld.arg(&dy);
bld.arg(&dd);
bld.arg(&du);
bld.arg(&dbl);
bld.arg(&dcl);
bld.arg(&ddd);
bld.arg(&da);
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_backward_local_typed: {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 reduce_db = match dtype {
WeightDtype::Bf16 => &k.ssm_reduce_d_b_bf16,
WeightDtype::F16 => &k.ssm_reduce_d_b_f16,
WeightDtype::F32 => &k.ssm_reduce_d_b,
};
let mut bld = ctx.stream.launch_builder(reduce_db);
bld.arg(scratch.d_b_reduced.inner_mut());
let src = scratch.d_b_local.cached_ptr();
bld.arg(&src);
bld.arg(&b_i);
bld.arg(&t_i);
bld.arg(&di_i);
bld.arg(&ds_i);
unsafe { bld.launch(grid_1d(bt * ds)) }
.map_err(|e| format!("ssm_reduce_d_b typed: {e:?}"))?;
let reduce_dc = match dtype {
WeightDtype::Bf16 => &k.ssm_reduce_d_c_bf16,
WeightDtype::F16 => &k.ssm_reduce_d_c_f16,
WeightDtype::F32 => &k.ssm_reduce_d_c,
};
let mut bld = ctx.stream.launch_builder(reduce_dc);
bld.arg(scratch.d_c_reduced.inner_mut());
let src = scratch.d_c_local.cached_ptr();
bld.arg(&src);
bld.arg(&b_i);
bld.arg(&t_i);
bld.arg(&di_i);
bld.arg(&ds_i);
unsafe { bld.launch(grid_1d(bt * ds)) }
.map_err(|e| format!("ssm_reduce_d_c typed: {e:?}"))?;
let mut bld = ctx.stream.launch_builder(&k.ssm_reduce_d_d);
let p = d_lw.d_param.ptr();
bld.arg(&p);
bld.arg(scratch.d_d_local.inner());
bld.arg(&b_i);
bld.arg(&di_i);
unsafe { bld.launch(grid_1d(di)) }.map_err(|e| format!("ssm_reduce_d_d: {e:?}"))?;
let mut bld = ctx.stream.launch_builder(&k.ssm_reduce_d_a_log);
let p = d_lw.a_log.ptr();
bld.arg(&p);
bld.arg(scratch.d_a_log_local.inner());
bld.arg(&b_i);
bld.arg(&di_i);
bld.arg(&ds_i);
unsafe { bld.launch(grid_1d(di * ds)) }
.map_err(|e| format!("ssm_reduce_d_a_log: {e:?}"))?;
}
scratch.d_xdbl.zero(&ctx.stream)?;
{
let n_bc = (bt * ds) as i32;
scratch.b_buf.zero(&ctx.stream)?;
let mut bld = ctx
.stream
.launch_builder(k.vec_add_inplace_typed.get(dtype));
let dst = scratch.b_buf.cached_ptr();
let src = scratch.d_b_reduced.cached_ptr();
bld.arg(&dst);
bld.arg(&src);
bld.arg(&n_bc);
unsafe { bld.launch(grid_1d(bt * ds)) }
.map_err(|e| format!("cast d_b_reduced→typed: {e:?}"))?;
scratch.c_buf.zero(&ctx.stream)?;
let mut bld = ctx
.stream
.launch_builder(k.vec_add_inplace_typed.get(dtype));
let dst = scratch.c_buf.cached_ptr();
let src = scratch.d_c_reduced.cached_ptr();
bld.arg(&dst);
bld.arg(&src);
bld.arg(&n_bc);
unsafe { bld.launch(grid_1d(bt * ds)) }
.map_err(|e| format!("cast d_c_reduced→typed: {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.scatter_add_cols_typed.get(dtype));
let dst = scratch.d_xdbl.cached_ptr();
let src = scratch.b_buf.cached_ptr();
bld.arg(&dst);
bld.arg(&src);
bld.arg(&bt_i);
bld.arg(&xdbl_i);
bld.arg(&ds_i);
bld.arg(&b_off);
unsafe { bld.launch(grid_1d(bt * ds)) }.map_err(|e| format!("scatter d_b typed: {e:?}"))?;
let mut bld = ctx
.stream
.launch_builder(k.scatter_add_cols_typed.get(dtype));
let dst = scratch.d_xdbl.cached_ptr();
let src = scratch.c_buf.cached_ptr();
bld.arg(&dst);
bld.arg(&src);
bld.arg(&bt_i);
bld.arg(&xdbl_i);
bld.arg(&ds_i);
bld.arg(&c_off);
unsafe { bld.launch(grid_1d(bt * ds)) }.map_err(|e| format!("scatter d_c typed: {e:?}"))?;
}
{
let n = (bt * di) as i32;
let mut bld = ctx.stream.launch_builder(k.softplus_bwd_typed.get(dtype));
let dx = scratch.d_delta_raw.cached_ptr();
let xs = acts.delta_raw.cached_ptr();
let dy = scratch.d_delta.cached_ptr();
bld.arg(&dx);
bld.arg(&xs);
bld.arg(&dy);
bld.arg(&n);
unsafe { bld.launch(grid_1d(bt * di)) }
.map_err(|e| format!("softplus_bwd_typed: {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 bld = ctx.stream.launch_builder(k.gather_cols_typed.get(dtype));
let dst = scratch.dt_xdbl_buf.cached_ptr();
let src = acts.xdbl.cached_ptr();
bld.arg(&dst);
bld.arg(&src);
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 bwd typed: {e:?}"))?;
}
gpu_sgemm_backward_dw_grad_typed(
ctx,
&d_lw.dt_proj_w,
TypedPtr {
ptr: scratch.d_delta_raw.cached_ptr(),
dtype,
},
TypedPtr {
ptr: scratch.dt_xdbl_buf.cached_ptr(),
dtype,
},
bt,
dt_rank,
di,
)?;
gpu_gemm_ex_backward_dx_typed(
ctx,
TypedPtr {
ptr: scratch.d_dt_input.cached_ptr(),
dtype,
},
TypedPtr {
ptr: scratch.d_delta_raw.cached_ptr(),
dtype,
},
TypedPtr {
ptr: lw.dt_proj_w.ptr(),
dtype,
},
bt,
dt_rank,
di,
)?;
{
let bt_i = bt as i32;
let di_i = di as i32;
let mut bld = ctx.stream.launch_builder(k.reduce_bias_typed.get(dtype));
let db = d_lw.dt_proj_b.ptr();
let dy = scratch.d_delta_raw.cached_ptr();
bld.arg(&db);
bld.arg(&dy);
bld.arg(&bt_i);
bld.arg(&di_i);
let threads = 256u32;
let cfg = cudarc::driver::LaunchConfig {
grid_dim: (di as u32, 1, 1),
block_dim: (threads, 1, 1),
shared_mem_bytes: (threads as usize * std::mem::size_of::<f32>()) as u32,
};
unsafe { bld.launch(cfg) }.map_err(|e| format!("reduce_bias dt_proj: {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 bld = ctx
.stream
.launch_builder(k.scatter_add_cols_typed.get(dtype));
let dst = scratch.d_xdbl.cached_ptr();
let src = scratch.d_dt_input.cached_ptr();
bld.arg(&dst);
bld.arg(&src);
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!("scatter dt typed bwd: {e:?}"))?;
}
gpu_sgemm_backward_dw_grad_typed(
ctx,
&d_lw.x_proj_w,
TypedPtr {
ptr: scratch.d_xdbl.cached_ptr(),
dtype,
},
TypedPtr {
ptr: acts.u.cached_ptr(),
dtype,
},
bt,
di,
xdbl_dim,
)?;
gpu_gemm_ex_backward_dx_typed(
ctx,
TypedPtr {
ptr: scratch.d_u_xproj.cached_ptr(),
dtype,
},
TypedPtr {
ptr: scratch.d_xdbl.cached_ptr(),
dtype,
},
TypedPtr {
ptr: lw.x_proj_w.ptr(),
dtype,
},
bt,
di,
xdbl_dim,
)?;
{
let bt_i = bt as i32;
let di_i_dst = di as i32;
let di_i_src = di as i32;
let offset: i32 = 0;
let mut bld = ctx
.stream
.launch_builder(k.scatter_add_cols_typed.get(dtype));
let dst = scratch.d_u.cached_ptr();
let src = scratch.d_u_xproj.cached_ptr();
bld.arg(&dst);
bld.arg(&src);
bld.arg(&bt_i);
bld.arg(&di_i_dst);
bld.arg(&di_i_src);
bld.arg(&offset);
unsafe { bld.launch(grid_1d(bt * di)) }
.map_err(|e| format!("d_u += d_u_xproj typed: {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 bld = ctx
.stream
.launch_builder(k.conv1d_burnin_bwd_typed.get(dtype));
let dxb = scratch.d_x_branch.cached_ptr();
let du = scratch.d_u.cached_ptr();
let pc = acts.post_conv.cached_ptr();
let cs = acts.conv_states.cached_ptr();
let w = lw.conv1d_weight.ptr();
bld.arg(&dxb);
bld.arg(&wp_ptr); bld.arg(&bp_ptr); bld.arg(&du);
bld.arg(&pc);
bld.arg(&cs);
bld.arg(&w);
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_bwd_typed 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 bld = ctx.stream.launch_builder(&ctx.kernels.reduce_sum_axis0);
bld.arg(&p);
bld.arg(&wp_ptr);
bld.arg(&b_i);
bld.arg(&dim_w);
bld.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 { bld.launch(cfg) }
.map_err(|e| format!("conv1d_burnin_bwd_typed 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 bld = ctx.stream.launch_builder(&ctx.kernels.reduce_sum_axis0);
bld.arg(&p);
bld.arg(&bp_ptr);
bld.arg(&b_i);
bld.arg(&di_i);
bld.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 { bld.launch(cfg) }
.map_err(|e| format!("conv1d_burnin_bwd_typed bias final: {e:?}"))?;
}
}
{
let bt_i = bt as i32;
let di_i = di as i32;
let mut bld = ctx.stream.launch_builder(k.concat_halves_typed.get(dtype));
let dst = scratch.d_proj.cached_ptr();
let fh = scratch.d_x_branch.cached_ptr();
let sh = scratch.d_gate.cached_ptr();
bld.arg(&dst);
bld.arg(&fh);
bld.arg(&sh);
bld.arg(&bt_i);
bld.arg(&di_i);
unsafe { bld.launch(grid_1d(bt * di)) }
.map_err(|e| format!("concat_halves_typed bwd: {e:?}"))?;
}
gpu_sgemm_backward_dw_grad_typed(
ctx,
&d_lw.in_proj_w,
TypedPtr {
ptr: scratch.d_proj.cached_ptr(),
dtype,
},
TypedPtr {
ptr: acts.post_norm.cached_ptr(),
dtype,
},
bt,
dm,
2 * di,
)?;
gpu_gemm_ex_backward_dx_typed(
ctx,
TypedPtr {
ptr: scratch.d_norm.cached_ptr(),
dtype,
},
TypedPtr {
ptr: scratch.d_proj.cached_ptr(),
dtype,
},
TypedPtr {
ptr: lw.in_proj_w.ptr(),
dtype,
},
bt,
dm,
2 * di,
)?;
let rmsnorm_bwd = match dtype {
WeightDtype::F32 => &k.rmsnorm_bwd,
WeightDtype::Bf16 | WeightDtype::F16 => k.rmsnorm_bwd_f32in_typed.get(dtype),
};
{
let bt_i = bt as i32;
let dm_i = dm as i32;
let axis0_ptr = scratch.axis0_partials.cached_ptr();
{
let mut bld = ctx.stream.launch_builder(rmsnorm_bwd);
let dx = scratch.d_pre_norm.cached_ptr();
let dy = scratch.d_norm.cached_ptr();
let x = acts.residual.cached_ptr();
let sc = lw.norm_weight.ptr();
let rms = acts.rms_vals.cached_ptr();
bld.arg(&dx);
bld.arg(&axis0_ptr); bld.arg(&dy);
bld.arg(&x);
bld.arg(&sc);
bld.arg(&rms);
bld.arg(&bt_i);
bld.arg(&dm_i);
unsafe { bld.launch(grid_norm(bt, dm)) }
.map_err(|e| format!("rmsnorm_bwd_f32in_typed 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 bld = ctx.stream.launch_builder(&ctx.kernels.reduce_sum_axis0);
bld.arg(&p);
bld.arg(&axis0_ptr);
bld.arg(&bt_i);
bld.arg(&dm_i);
bld.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 { bld.launch(cfg) }
.map_err(|e| format!("rmsnorm_bwd_f32in_typed final: {e:?}"))?;
}
}
{
let n = (bt * dm) as i32;
let mut bld = ctx.stream.launch_builder(&k.vec_add_inplace);
bld.arg(d_temporal.inner_mut());
bld.arg(scratch.d_pre_norm.inner());
bld.arg(&n);
unsafe { bld.launch(grid_1d(bt * dm)) }
.map_err(|e| format!("vec_add residual bwd mixed: {e:?}"))?;
}
Ok(())
}
pub fn gpu_backward_mamba_backbone_mixed(
ctx: &GpuCtx,
d_temporal: &mut GpuBuffer,
d_mamba: &GpuMambaGrads,
acts: &GpuMambaBackboneMixedActs,
mamba_w: &GpuMambaMixedWeights,
a_neg_all: &GpuBuffer,
scratch: &mut GpuMambaMixedTrainScratch,
) -> Result<(), String> {
let dims = scratch.dims;
let dtype = acts.dtype;
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 bld = ctx.stream.launch_builder(&ctx.kernels.rmsnorm_bwd);
bld.arg(scratch.d_pre_norm.inner_mut());
bld.arg(&axis0_ptr); bld.arg(d_temporal.inner());
bld.arg(acts.norm_f_input.inner());
let nf = mamba_w.norm_f_weight.ptr();
bld.arg(&nf);
bld.arg(acts.norm_f_rms.inner());
bld.arg(&bt_i);
bld.arg(&dm_i);
unsafe { bld.launch(grid_norm(bt, dims.d_model)) }
.map_err(|e| format!("rmsnorm_bwd norm_f mixed 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 bld = ctx.stream.launch_builder(&ctx.kernels.reduce_sum_axis0);
bld.arg(&p);
bld.arg(&axis0_ptr);
bld.arg(&bt_i);
bld.arg(&dm_i);
bld.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 { bld.launch(cfg) }
.map_err(|e| format!("rmsnorm_bwd norm_f mixed final: {e:?}"))?;
}
d_temporal.copy_from(&scratch.d_pre_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_mixed(
ctx,
d_temporal,
&MixedLayerBwd {
d_lw: &d_mamba.layers[layer_idx],
acts: &acts.layers[layer_idx],
lw: &mamba_w.layers[layer_idx],
},
a_neg_ptr,
scratch,
dtype,
)?;
}
if mamba_w.input_proj_w.len_elems() != 0 {
let bt = dims.batch * dims.seq_len;
let dm = dims.d_model;
let mid = dims.mamba_input_dim;
let dy_ptr = acts.input_proj_outputs.cached_ptr();
{
let n = (bt * dm) as i32;
let cast = match dtype {
WeightDtype::Bf16 => &ctx.kernels.cast_f32_to_bf16,
WeightDtype::F16 => &ctx.kernels.cast_f32_to_f16,
WeightDtype::F32 => {
return Err("mixed backward: unexpected f32 compute dtype".into());
}
};
let src = d_temporal.cached_ptr();
let mut bld = ctx.stream.launch_builder(cast);
bld.arg(&dy_ptr);
bld.arg(&src);
bld.arg(&n);
unsafe { bld.launch(grid_1d(bt * dm)) }
.map_err(|e| format!("input_proj dY cast: {e:?}"))?;
}
{
let bt_i = bt as i32;
let dm_i = dm as i32;
let mut bld = ctx
.stream
.launch_builder(ctx.kernels.reduce_bias_typed.get(dtype));
let db = d_mamba.input_proj_b.ptr();
bld.arg(&db);
bld.arg(&dy_ptr);
bld.arg(&bt_i);
bld.arg(&dm_i);
let threads = 256u32;
let cfg = cudarc::driver::LaunchConfig {
grid_dim: (dm as u32, 1, 1),
block_dim: (threads, 1, 1),
shared_mem_bytes: (threads as usize * std::mem::size_of::<f32>()) as u32,
};
unsafe { bld.launch(cfg) }.map_err(|e| format!("reduce_bias input_proj: {e:?}"))?;
}
gpu_sgemm_backward_dw_grad_typed(
ctx,
&d_mamba.input_proj_w,
TypedPtr { ptr: dy_ptr, dtype },
TypedPtr {
ptr: acts.input_proj_inputs.cached_ptr(),
dtype,
},
bt,
mid,
dm,
)?;
}
Ok(())
}