use super::backward::GpuMambaTargetScratch;
use super::blas::gpu_sgemm_forward_raw;
use super::buffers::GpuBuffer;
use super::context::GpuCtx;
use super::launch::{grid_1d, grid_norm, grid_parallel_scan};
use super::weights::{GpuMambaTrainLayerWeights, GpuMambaTrainWeights};
use cudarc::driver::PushKernelArg;
use std::sync::Arc;
pub const PARALLEL_SCAN_THRESHOLD: usize = crate::config::ScanMode::PARALLEL_SCAN_THRESHOLD;
#[derive(Debug, Clone, Copy)]
pub struct GpuMambaDims {
pub batch: usize,
pub d_model: usize,
pub d_inner: usize,
pub d_state: usize,
pub d_conv: usize,
pub dt_rank: usize,
pub xdbl_dim: usize,
pub seq_len: usize,
pub mamba_input_dim: usize,
pub n_layers: usize,
pub scan_mode: crate::config::ScanMode,
pub rms_norm_eps: f32,
}
impl GpuMambaDims {
pub fn bt(&self) -> usize {
self.batch * self.seq_len
}
}
pub struct GpuMambaLayerActs {
pub residual: GpuBuffer,
pub rms_vals: GpuBuffer,
pub post_norm: GpuBuffer,
pub gate_pre_silu: GpuBuffer,
pub gate_post_silu: GpuBuffer,
pub conv_states: GpuBuffer,
pub post_conv: GpuBuffer,
pub u: GpuBuffer,
pub xdbl: GpuBuffer,
pub delta_raw: GpuBuffer,
pub delta: GpuBuffer,
pub h_saved: GpuBuffer,
pub da_exp: GpuBuffer,
pub y: GpuBuffer,
pub gated: GpuBuffer,
}
pub struct GpuMambaBackboneActs {
pub input_proj_inputs: GpuBuffer,
pub input_proj_outputs: GpuBuffer,
pub layers: Vec<GpuMambaLayerActs>,
pub norm_f_input: GpuBuffer,
pub norm_f_rms: GpuBuffer,
}
impl GpuMambaBackboneActs {
pub fn new(
stream: &Arc<cudarc::driver::CudaStream>,
dims: &GpuMambaDims,
) -> 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(GpuMambaLayerActs {
residual: GpuBuffer::zeros(stream, bt * d_model)?,
rms_vals: GpuBuffer::zeros(stream, bt)?,
post_norm: GpuBuffer::zeros(stream, bt * d_model)?,
gate_pre_silu: GpuBuffer::zeros(stream, bt * d_inner)?,
gate_post_silu: GpuBuffer::zeros(stream, bt * d_inner)?,
conv_states: GpuBuffer::zeros(stream, bt * d_inner * d_conv)?,
post_conv: GpuBuffer::zeros(stream, bt * d_inner)?,
u: GpuBuffer::zeros(stream, bt * d_inner)?,
xdbl: GpuBuffer::zeros(stream, bt * xdbl_dim)?,
delta_raw: GpuBuffer::zeros(stream, bt * d_inner)?,
delta: GpuBuffer::zeros(stream, bt * d_inner)?,
h_saved: GpuBuffer::zeros(stream, batch * (seq_len + 1) * d_inner * d_state)?,
da_exp: GpuBuffer::zeros(stream, bt * d_inner * d_state)?,
y: GpuBuffer::zeros(stream, bt * d_inner)?,
gated: GpuBuffer::zeros(stream, bt * d_inner)?,
})
})
.collect::<Result<Vec<_>, String>>()?;
Ok(Self {
input_proj_inputs: GpuBuffer::zeros(stream, bt * mamba_input_dim)?,
input_proj_outputs: GpuBuffer::zeros(stream, bt * d_model)?,
layers,
norm_f_input: GpuBuffer::zeros(stream, bt * d_model)?,
norm_f_rms: GpuBuffer::zeros(stream, bt)?,
})
}
}
pub struct GpuMambaScratch {
pub dims: GpuMambaDims,
pub proj_flat: GpuBuffer,
pub x_branch: GpuBuffer,
pub out_flat: GpuBuffer,
pub dt_gather_buf: GpuBuffer,
pub d_gated: GpuBuffer,
pub d_y: GpuBuffer,
pub d_gate: GpuBuffer,
pub d_delta: GpuBuffer,
pub d_delta_raw: GpuBuffer,
pub d_u: GpuBuffer,
pub d_u_xproj: GpuBuffer,
pub d_xdbl: GpuBuffer,
pub d_x_branch: GpuBuffer,
pub d_proj: GpuBuffer,
pub d_norm: GpuBuffer,
pub d_pre_norm: GpuBuffer,
pub d_dt_input: GpuBuffer,
pub dt_xdbl_buf: GpuBuffer,
pub d_b_local: GpuBuffer,
pub d_c_local: GpuBuffer,
pub d_d_local: GpuBuffer,
pub d_a_log_local: GpuBuffer,
pub d_b_reduced: GpuBuffer,
pub d_c_reduced: GpuBuffer,
pub d_input_proj_dx: GpuBuffer,
pub axis0_partials: GpuBuffer,
}
impl GpuMambaScratch {
pub fn new(
stream: &Arc<cudarc::driver::CudaStream>,
dims: &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 dt_rank = dims.dt_rank;
let mamba_input_dim = dims.mamba_input_dim;
let bt = batch * dims.seq_len;
let xdbl_dim = dt_rank + 2 * d_state;
Ok(Self {
dims: *dims,
proj_flat: GpuBuffer::zeros(stream, bt * 2 * d_inner)?,
x_branch: GpuBuffer::zeros(stream, bt * d_inner)?,
out_flat: GpuBuffer::zeros(stream, bt * d_model)?,
dt_gather_buf: GpuBuffer::zeros(stream, bt * dt_rank)?,
d_gated: GpuBuffer::zeros(stream, bt * d_inner)?,
d_y: GpuBuffer::zeros(stream, bt * d_inner)?,
d_gate: GpuBuffer::zeros(stream, bt * d_inner)?,
d_delta: GpuBuffer::zeros(stream, bt * d_inner)?,
d_delta_raw: GpuBuffer::zeros(stream, bt * d_inner)?,
d_u: GpuBuffer::zeros(stream, bt * d_inner)?,
d_u_xproj: GpuBuffer::zeros(stream, bt * d_inner)?,
d_xdbl: GpuBuffer::zeros(stream, bt * xdbl_dim)?,
d_x_branch: GpuBuffer::zeros(stream, bt * d_inner)?,
d_proj: GpuBuffer::zeros(stream, bt * 2 * d_inner)?,
d_norm: GpuBuffer::zeros(stream, bt * d_model)?,
d_pre_norm: GpuBuffer::zeros(stream, bt * d_model)?,
d_dt_input: GpuBuffer::zeros(stream, bt * dt_rank)?,
dt_xdbl_buf: GpuBuffer::zeros(stream, bt * dt_rank)?,
d_b_local: GpuBuffer::zeros(stream, bt * d_inner * d_state)?,
d_c_local: GpuBuffer::zeros(stream, bt * d_inner * d_state)?,
d_d_local: GpuBuffer::zeros(stream, batch * d_inner)?,
d_a_log_local: GpuBuffer::zeros(stream, batch * d_inner * d_state)?,
d_b_reduced: GpuBuffer::zeros(stream, bt * d_state)?,
d_c_reduced: GpuBuffer::zeros(stream, bt * d_state)?,
d_input_proj_dx: GpuBuffer::zeros(stream, bt * mamba_input_dim)?,
axis0_partials: GpuBuffer::zeros(
stream,
std::cmp::max(bt * d_model, batch * d_inner * (dims.d_conv + 1)),
)?,
})
}
}
pub struct MambaLayerPtrs {
pub conv_state: cudarc::driver::sys::CUdeviceptr, pub ssm_state: cudarc::driver::sys::CUdeviceptr, pub a_neg: cudarc::driver::sys::CUdeviceptr, }
pub fn gpu_forward_mamba_layer(
ctx: &GpuCtx,
temporal: &mut GpuBuffer,
acts: &mut GpuMambaLayerActs,
lw: &GpuMambaTrainLayerWeights,
layer_ptrs: &MambaLayerPtrs,
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;
acts.residual.copy_from(temporal, &ctx.stream)?;
{
let batch_i = bt as i32;
let dim_i = dm as i32;
let eps: f32 = dims.rms_norm_eps;
let mut builder = ctx.stream.launch_builder(&ctx.kernels.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.cached_ptr();
builder.arg(&nw_ptr);
builder.arg(&batch_i);
builder.arg(&dim_i);
builder.arg(&eps);
unsafe { builder.launch(grid_norm(bt, dm)) }
.map_err(|e| format!("rmsnorm_fwd mamba: {:?}", e))?;
}
gpu_sgemm_forward_raw(
ctx,
&mut scratch.proj_flat,
&acts.post_norm,
lw.in_proj_w.cached_ptr(),
None,
(bt, dm, 2 * di),
)?;
{
let batch_i = bt as i32;
let di_i = di as i32;
let mut builder = ctx.stream.launch_builder(&ctx.kernels.split_gate_silu);
builder.arg(scratch.x_branch.inner_mut());
builder.arg(acts.gate_pre_silu.inner_mut());
builder.arg(acts.gate_post_silu.inner_mut());
builder.arg(scratch.proj_flat.inner());
builder.arg(&batch_i);
builder.arg(&di_i);
unsafe { builder.launch(grid_1d(bt * di)) }
.map_err(|e| format!("split_gate_silu 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 mut builder = ctx.stream.launch_builder(&ctx.kernels.conv1d_burnin_fwd);
builder.arg(acts.u.inner_mut());
builder.arg(acts.post_conv.inner_mut());
builder.arg(acts.conv_states.inner_mut());
builder.arg(&layer_ptrs.conv_state); builder.arg(scratch.x_branch.inner());
let cw_ptr = lw.conv1d_weight.cached_ptr();
let cb_ptr = lw.conv1d_bias.cached_ptr();
builder.arg(&cw_ptr);
builder.arg(&cb_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_fwd mamba: {:?}", e))?;
}
gpu_sgemm_forward_raw(
ctx,
&mut acts.xdbl,
&acts.u,
lw.x_proj_w.cached_ptr(),
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 builder = ctx.stream.launch_builder(&ctx.kernels.gather_cols);
builder.arg(scratch.dt_gather_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_cols dt mamba: {:?}", e))?;
}
gpu_sgemm_forward_raw(
ctx,
&mut acts.delta_raw,
&scratch.dt_gather_buf,
lw.dt_proj_w.cached_ptr(),
Some(lw.dt_proj_b.cached_ptr()),
(bt, dt_rank, di),
)?;
{
let n = (bt * di) as i32;
let mut builder = ctx.stream.launch_builder(&ctx.kernels.softplus_copy);
builder.arg(acts.delta.inner_mut());
builder.arg(acts.delta_raw.inner());
builder.arg(&n);
unsafe { builder.launch(grid_1d(bt * di)) }
.map_err(|e| format!("softplus_copy 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 fwd mamba: {:?}", 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 dims.scan_mode.use_parallel(t, ds) {
let mut builder = ctx.stream.launch_builder(&ctx.kernels.ssm_parallel_fwd);
builder.arg(&layer_ptrs.ssm_state);
builder.arg(acts.y.inner_mut());
builder.arg(acts.h_saved.inner_mut());
builder.arg(acts.da_exp.inner_mut());
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(&layer_ptrs.a_neg);
let dp_ptr = lw.d_param.cached_ptr();
builder.arg(&dp_ptr);
builder.arg(&b_i);
builder.arg(&t_i);
builder.arg(&di_i);
builder.arg(&ds_i);
unsafe { builder.launch(grid_parallel_scan(b, di)) }
.map_err(|e| format!("ssm_parallel_fwd mamba: {:?}", e))?;
} else {
let mut builder = ctx.stream.launch_builder(&ctx.kernels.ssm_burnin_fwd);
builder.arg(&layer_ptrs.ssm_state);
builder.arg(acts.y.inner_mut());
builder.arg(acts.h_saved.inner_mut());
builder.arg(acts.da_exp.inner_mut());
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(&layer_ptrs.a_neg);
let dp_ptr = lw.d_param.cached_ptr();
builder.arg(&dp_ptr);
builder.arg(&b_i);
builder.arg(&t_i);
builder.arg(&di_i);
builder.arg(&ds_i);
unsafe { builder.launch(grid_1d(b * di)) }
.map_err(|e| format!("ssm_burnin_fwd mamba: {:?}", e))?;
}
}
{
let n = (bt * di) as i32;
let mut builder = ctx.stream.launch_builder(&ctx.kernels.elementwise_mul);
builder.arg(acts.gated.inner_mut());
builder.arg(acts.y.inner());
builder.arg(acts.gate_post_silu.inner());
builder.arg(&n);
unsafe { builder.launch(grid_1d(bt * di)) }
.map_err(|e| format!("elementwise_mul gating mamba: {:?}", e))?;
}
gpu_sgemm_forward_raw(
ctx,
&mut scratch.out_flat,
&acts.gated,
lw.out_proj_w.cached_ptr(),
None,
(bt, di, dm),
)?;
{
let n = (bt * dm) as i32;
let mut builder = ctx.stream.launch_builder(&ctx.kernels.residual_add);
builder.arg(temporal.inner_mut());
builder.arg(acts.residual.inner());
builder.arg(scratch.out_flat.inner());
builder.arg(&n);
unsafe { builder.launch(grid_1d(bt * dm)) }
.map_err(|e| format!("residual_add mamba: {:?}", e))?;
}
Ok(())
}
pub struct GpuRecurrentState {
pub conv_states: GpuBuffer,
pub ssm_states: GpuBuffer,
pub a_neg_all: GpuBuffer,
}
pub fn gpu_forward_mamba_backbone(
ctx: &GpuCtx,
temporal: &mut GpuBuffer,
acts: &mut GpuMambaBackboneActs,
mamba_w: &GpuMambaTrainWeights,
mamba_input: &GpuBuffer,
state: &mut GpuRecurrentState,
scratch: &mut GpuMambaScratch,
) -> Result<(), String> {
let dims = scratch.dims;
let bt = dims.bt();
acts.input_proj_inputs.copy_from(mamba_input, &ctx.stream)?;
gpu_sgemm_forward_raw(
ctx,
temporal,
mamba_input,
mamba_w.input_proj_w.cached_ptr(),
Some(mamba_w.input_proj_b.cached_ptr()),
(bt, dims.mamba_input_dim, dims.d_model),
)?;
acts.input_proj_outputs.copy_from(temporal, &ctx.stream)?;
let conv_per_layer = dims.batch * dims.d_inner * dims.d_conv;
let ssm_per_layer = dims.batch * dims.d_inner * dims.d_state;
let a_neg_per_layer = dims.d_inner * dims.d_state;
for layer_idx in 0..dims.n_layers {
let conv_base = state.conv_states.raw_ptr(&ctx.stream);
let ssm_base = state.ssm_states.raw_ptr(&ctx.stream);
let aneg_base = state.a_neg_all.raw_ptr(&ctx.stream);
let f32_sz = std::mem::size_of::<f32>() as u64;
let layer_ptrs = MambaLayerPtrs {
conv_state: conv_base + (layer_idx * conv_per_layer) as u64 * f32_sz,
ssm_state: ssm_base + (layer_idx * ssm_per_layer) as u64 * f32_sz,
a_neg: aneg_base + (layer_idx * a_neg_per_layer) as u64 * f32_sz,
};
gpu_forward_mamba_layer(
ctx,
temporal,
&mut acts.layers[layer_idx],
&mamba_w.layers[layer_idx],
&layer_ptrs,
scratch,
)?;
}
{
let bt_i = bt as i32;
let dm_i = dims.d_model as i32;
let eps: f32 = dims.rms_norm_eps;
acts.norm_f_input.copy_from(temporal, &ctx.stream)?;
let mut builder = ctx.stream.launch_builder(&ctx.kernels.rmsnorm_fwd);
builder.arg(temporal.inner_mut()); builder.arg(acts.norm_f_rms.inner_mut()); builder.arg(acts.norm_f_input.inner()); let nf_ptr = mamba_w.norm_f_weight.cached_ptr();
builder.arg(&nf_ptr);
builder.arg(&bt_i);
builder.arg(&dm_i);
builder.arg(&eps);
unsafe { builder.launch(grid_norm(bt, dims.d_model)) }
.map_err(|e| format!("rmsnorm_fwd norm_f: {:?}", e))?;
}
Ok(())
}
pub fn gpu_forward_mamba_target_burnin(
ctx: &GpuCtx,
target_temporal: &mut GpuBuffer, ip_out_flat: &GpuBuffer, target_w: &GpuMambaTrainWeights,
a_neg_all: &GpuBuffer,
scratch: &mut GpuMambaTargetScratch,
) -> Result<(), String> {
let dims = &scratch.dims;
let seq_len = dims.seq_len;
let b = dims.batch;
let bt = b * 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 = dims.xdbl_dim;
let d_conv = dims.d_conv;
let t = seq_len;
scratch.conv_states.zero(&ctx.stream)?;
scratch.ssm_states.zero(&ctx.stream)?;
scratch.out_flat.copy_from(ip_out_flat, &ctx.stream)?;
let conv_per_layer = b * di * d_conv;
let ssm_per_layer = b * di * ds;
let a_neg_per_layer = di * ds;
for layer_idx in 0..dims.n_layers {
let lw = &target_w.layers[layer_idx];
let conv_base = scratch.conv_states.raw_ptr(&ctx.stream);
let ssm_base = scratch.ssm_states.raw_ptr(&ctx.stream);
let aneg_base = a_neg_all.raw_ptr(&ctx.stream);
let f32_sz = std::mem::size_of::<f32>() as u64;
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 a_neg_ptr = aneg_base + (layer_idx * a_neg_per_layer) as u64 * f32_sz;
scratch.residual.copy_from(&scratch.out_flat, &ctx.stream)?;
{
let bt_i = bt as i32;
let dm_i = dm as i32;
let eps: f32 = dims.rms_norm_eps;
let mut builder = ctx.stream.launch_builder(&ctx.kernels.rmsnorm_fwd);
builder.arg(scratch.out_flat.inner_mut()); builder.arg(scratch.rms_discard.inner_mut()); builder.arg(scratch.residual.inner()); let nw_ptr = lw.norm_weight.cached_ptr();
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 target L{layer_idx}: {:?}", e))?;
}
gpu_sgemm_forward_raw(
ctx,
&mut scratch.proj_flat,
&scratch.out_flat,
lw.in_proj_w.cached_ptr(),
None,
(bt, dm, 2 * di),
)?;
{
let bt_i = bt as i32;
let di_i = di as i32;
let gs_raw = scratch.gate_silu.raw_ptr(&ctx.stream);
let mut builder = ctx.stream.launch_builder(&ctx.kernels.split_gate_silu);
builder.arg(scratch.x_branch.inner_mut());
builder.arg(scratch.gate_silu.inner_mut()); builder.arg(&gs_raw); builder.arg(scratch.proj_flat.inner());
builder.arg(&bt_i);
builder.arg(&di_i);
unsafe { builder.launch(grid_1d(bt * di)) }
.map_err(|e| format!("split_gate target 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 mut builder = ctx
.stream
.launch_builder(&ctx.kernels.conv1d_burnin_fwd_nosave);
builder.arg(scratch.u.inner_mut()); builder.arg(&conv_ptr); builder.arg(scratch.x_branch.inner());
let cw_ptr = lw.conv1d_weight.cached_ptr();
let cb_ptr = lw.conv1d_bias.cached_ptr();
builder.arg(&cw_ptr);
builder.arg(&cb_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_nosave target L{layer_idx}: {:?}", e))?;
}
gpu_sgemm_forward_raw(
ctx,
&mut scratch.xdbl,
&scratch.u,
lw.x_proj_w.cached_ptr(),
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 builder = ctx.stream.launch_builder(&ctx.kernels.gather_cols);
builder.arg(scratch.dt_gather.inner_mut());
builder.arg(scratch.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 target L{layer_idx}: {:?}", e))?;
}
gpu_sgemm_forward_raw(
ctx,
&mut scratch.delta,
&scratch.dt_gather,
lw.dt_proj_w.cached_ptr(),
Some(lw.dt_proj_b.cached_ptr()),
(bt, dt_rank, di),
)?;
{
let n = (bt * di) as i32;
let mut builder = ctx.stream.launch_builder(&ctx.kernels.softplus_fwd);
builder.arg(scratch.delta.inner_mut());
builder.arg(&n);
unsafe { builder.launch(grid_1d(bt * di)) }
.map_err(|e| format!("softplus target 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_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.b_gathered.inner_mut());
builder.arg(scratch.c_gathered.inner_mut());
builder.arg(scratch.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 target 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 dims.scan_mode.use_parallel(t, ds) {
let mut builder = ctx
.stream
.launch_builder(&ctx.kernels.ssm_parallel_fwd_nosave);
builder.arg(&ssm_ptr);
builder.arg(scratch.y.inner_mut());
builder.arg(scratch.delta.inner());
builder.arg(scratch.u.inner());
builder.arg(scratch.b_gathered.inner());
builder.arg(scratch.c_gathered.inner());
builder.arg(&a_neg_ptr);
let dp_ptr = lw.d_param.cached_ptr();
builder.arg(&dp_ptr);
builder.arg(&b_i);
builder.arg(&t_i);
builder.arg(&di_i);
builder.arg(&ds_i);
unsafe { builder.launch(grid_parallel_scan(b, di)) }
.map_err(|e| format!("ssm_parallel_nosave target L{layer_idx}: {:?}", e))?;
} else {
let mut builder = ctx
.stream
.launch_builder(&ctx.kernels.ssm_burnin_fwd_nosave);
builder.arg(&ssm_ptr);
builder.arg(scratch.y.inner_mut());
builder.arg(scratch.delta.inner());
builder.arg(scratch.u.inner());
builder.arg(scratch.b_gathered.inner());
builder.arg(scratch.c_gathered.inner());
builder.arg(&a_neg_ptr);
let dp_ptr = lw.d_param.cached_ptr();
builder.arg(&dp_ptr);
builder.arg(&b_i);
builder.arg(&t_i);
builder.arg(&di_i);
builder.arg(&ds_i);
unsafe { builder.launch(grid_1d(b * di)) }
.map_err(|e| format!("ssm_nosave target L{layer_idx}: {:?}", e))?;
}
}
{
let n = (bt * di) as i32;
let mut builder = ctx.stream.launch_builder(&ctx.kernels.elementwise_mul);
builder.arg(scratch.gated.inner_mut());
builder.arg(scratch.y.inner());
builder.arg(scratch.gate_silu.inner());
builder.arg(&n);
unsafe { builder.launch(grid_1d(bt * di)) }
.map_err(|e| format!("gating target L{layer_idx}: {:?}", e))?;
}
gpu_sgemm_forward_raw(
ctx,
&mut scratch.out_flat,
&scratch.gated,
lw.out_proj_w.cached_ptr(),
None,
(bt, di, dm),
)?;
{
let n = (bt * dm) as i32;
let mut builder = ctx.stream.launch_builder(&ctx.kernels.vec_add_inplace);
builder.arg(scratch.out_flat.inner_mut());
builder.arg(scratch.residual.inner());
builder.arg(&n);
unsafe { builder.launch(grid_1d(bt * dm)) }
.map_err(|e| format!("residual target L{layer_idx}: {:?}", e))?;
}
}
{
let bt_i = bt as i32;
let dm_i = dm as i32;
let eps: f32 = dims.rms_norm_eps;
scratch.residual.copy_from(&scratch.out_flat, &ctx.stream)?;
let mut builder = ctx.stream.launch_builder(&ctx.kernels.rmsnorm_fwd);
builder.arg(scratch.out_flat.inner_mut()); builder.arg(scratch.rms_discard.inner_mut()); builder.arg(scratch.residual.inner()); let tnf_ptr = target_w.norm_f_weight.cached_ptr();
builder.arg(&tnf_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 target: {:?}", e))?;
}
{
let b_i = b as i32;
let t_i = t as i32;
let dm_i = dm as i32;
let mut builder = ctx.stream.launch_builder(&ctx.kernels.gather_last_timestep);
builder.arg(target_temporal.inner_mut());
builder.arg(scratch.out_flat.inner());
builder.arg(&b_i);
builder.arg(&t_i);
builder.arg(&dm_i);
unsafe { builder.launch(grid_1d(b * dm)) }
.map_err(|e| format!("gather_last target: {:?}", e))?;
}
Ok(())
}