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 x_branch: 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 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)?,
x_branch: GpuBuffer::zeros(stream, bt * d_inner)?,
conv_states: GpuBuffer::zeros(stream, batch * 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,
if dims.scan_mode.use_parallel(seq_len, d_state)
&& super::launch::scan_tape_slim()
{
super::launch::scan_tape_len(batch, seq_len, d_inner, d_state)
} else {
batch * (seq_len + 1) * 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_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;
let bc_rows = if dims.scan_mode.use_parallel(dims.seq_len, d_state)
&& d_inner.is_multiple_of(super::launch::SCAN_BWD_DGROUP)
{
d_inner / super::launch::SCAN_BWD_DGROUP
} else {
d_inner
};
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_dt_input: GpuBuffer::zeros(stream, bt * dt_rank)?,
dt_xdbl_buf: GpuBuffer::zeros(stream, bt * dt_rank)?,
d_b_local: GpuBuffer::zeros(stream, bt * bc_rows * d_state)?,
d_c_local: GpuBuffer::zeros(stream, bt * bc_rows * d_state)?,
d_d_local: GpuBuffer::zeros(stream, batch * d_inner)?,
d_a_log_local: GpuBuffer::zeros(
stream,
batch * dims.seq_len.div_ceil(super::launch::SCAN_CHUNK).max(1) * 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 * dims.seq_len.div_ceil(128) * 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,
stream_out: cudarc::driver::sys::CUdeviceptr,
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;
{
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(acts.residual.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);
builder.arg(acts.x_branch.inner_mut());
builder.arg(acts.gate_pre_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 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_tiled_typed
.get(super::dtype::WeightDtype::F32),
);
builder.arg(acts.u.inner_mut());
builder.arg(&layer_ptrs.conv_state); builder.arg(acts.conv_states.inner_mut());
builder.arg(acts.post_conv.inner_mut());
builder.arg(acts.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(super::launch::grid_conv_tiled(b, di, t)) }
.map_err(|e| format!("conv1d_burnin_fwd_tiled 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 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 tmajor = dims.scan_mode.use_parallel(t, ds);
let kernel = if tmajor {
&ctx.kernels.gather_bc_cols_tmajor
} else {
&ctx.kernels.gather_bc_cols
};
let t_i = t as i32;
let mut builder = ctx.stream.launch_builder(kernel);
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);
if tmajor {
builder.arg(&t_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 tape_p = acts.h_saved.cached_ptr();
let slim_i: i32 = i32::from(super::launch::scan_tape_slim());
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.delta_raw.inner());
builder.arg(acts.delta.inner_mut());
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);
builder.arg(&tape_p);
builder.arg(&slim_i);
unsafe { builder.launch(grid_parallel_scan(b, di, ds)) }
.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.delta_raw.inner());
builder.arg(acts.delta.inner_mut());
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.gate_mul_silu);
builder.arg(acts.gated.inner_mut());
builder.arg(acts.y.inner());
builder.arg(acts.gate_pre_silu.inner());
builder.arg(&n);
unsafe { builder.launch(grid_1d(bt * di)) }
.map_err(|e| format!("gate_mul_silu 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(&stream_out);
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,
}
#[derive(Clone, Debug, PartialEq)]
pub struct RecurrentStateBlob {
pub conv_states: Vec<f32>,
pub ssm_states: Vec<f32>,
}
impl GpuRecurrentState {
pub fn export_state(
&self,
stream: &std::sync::Arc<cudarc::driver::CudaStream>,
) -> Result<RecurrentStateBlob, String> {
Ok(RecurrentStateBlob {
conv_states: self.conv_states.to_cpu(stream)?,
ssm_states: self.ssm_states.to_cpu(stream)?,
})
}
pub fn import_state(
&mut self,
stream: &std::sync::Arc<cudarc::driver::CudaStream>,
blob: &RecurrentStateBlob,
) -> Result<(), String> {
if blob.conv_states.len() != self.conv_states.len()
|| blob.ssm_states.len() != self.ssm_states.len()
{
return Err(format!(
"recurrent state mismatch: blob conv/ssm = {}/{} elements, \
state = {}/{} — the blob belongs to a different shape",
blob.conv_states.len(),
blob.ssm_states.len(),
self.conv_states.len(),
self.ssm_states.len()
));
}
self.conv_states.upload(stream, &blob.conv_states)?;
self.ssm_states.upload(stream, &blob.ssm_states)?;
Ok(())
}
}
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)?;
acts.layers[0].residual.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,
};
let stream_out = if layer_idx + 1 < dims.n_layers {
acts.layers[layer_idx + 1].residual.cached_ptr()
} else {
acts.norm_f_input.cached_ptr()
};
gpu_forward_mamba_layer(
ctx,
stream_out,
&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;
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 tmajor = dims.scan_mode.use_parallel(t, ds);
let kernel = if tmajor {
&ctx.kernels.gather_bc_cols_tmajor
} else {
&ctx.kernels.gather_bc_cols
};
let t_i = t as i32;
let mut builder = ctx.stream.launch_builder(kernel);
builder.arg(scratch.b_gathered.inner_mut());
builder.arg(scratch.c_gathered.inner_mut());
builder.arg(scratch.xdbl.inner());
builder.arg(&bt_i);
if tmajor {
builder.arg(&t_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);
let gate_stride = 0i32;
builder.arg(&dp_ptr);
builder.arg(&gate_stride);
builder.arg(&b_i);
builder.arg(&t_i);
builder.arg(&di_i);
builder.arg(&ds_i);
unsafe { builder.launch(grid_parallel_scan(b, di, ds)) }
.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);
let gate_stride = 0i32;
builder.arg(&dp_ptr);
builder.arg(&gate_stride);
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(())
}