use super::backward::{GpuMambaTargetMixedScratch, GpuMambaTargetScratch};
use super::blas::{
TypedPtr, gpu_gemm_forward_dispatch, gpu_gemm_typed_forward_raw, gpu_sgemm_forward_raw,
};
use super::buffers::GpuBuffer;
use super::context::GpuCtx;
use super::forward::{GpuMambaDims, PARALLEL_SCAN_THRESHOLD};
use super::inference::GpuInferenceState;
use super::launch::{grid_1d, grid_norm, grid_parallel_scan};
use super::weights::{MambaLayerWeightsView, MambaWeightsView};
use cudarc::driver::PushKernelArg;
pub struct PrefillInputs<'a, W: MambaWeightsView> {
pub ip_out_flat: &'a GpuBuffer,
pub weights: &'a W,
pub a_neg_all: &'a GpuBuffer,
}
pub fn gpu_forward_inference_prefill<W: MambaWeightsView>(
ctx: &GpuCtx,
target_temporal: &mut GpuBuffer,
inputs: PrefillInputs<'_, W>,
state: &mut GpuInferenceState,
scratch: &mut GpuMambaTargetScratch,
) -> Result<(), String> {
let PrefillInputs {
ip_out_flat,
weights,
a_neg_all,
} = inputs;
let dims: GpuMambaDims = 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.out_flat.copy_from(ip_out_flat, &ctx.stream)?;
let f32_sz = std::mem::size_of::<f32>() as u64;
for layer_idx in 0..weights.n_layers() {
let lw = weights.layer(layer_idx);
let conv_per_layer = b * di * d_conv;
let ssm_per_layer = b * di * ds;
let conv_ptr = state.conv.cached_ptr() + (layer_idx * conv_per_layer) as u64 * f32_sz;
let ssm_ptr = state.ssm.cached_ptr() + (layer_idx * ssm_per_layer) as u64 * f32_sz;
let a_neg_ptr = a_neg_all.cached_ptr() + (layer_idx * di * ds) 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 = lw.norm_weight();
builder.arg(&nw);
builder.arg(&bt_i);
builder.arg(&dm_i);
builder.arg(&eps);
unsafe { builder.launch(grid_norm(bt, dm)) }
.map_err(|e| format!("rmsnorm prefill L{layer_idx}: {e:?}"))?;
}
let (ipw, ipw_dt) = lw.in_proj_w();
gpu_gemm_forward_dispatch(
ctx,
&mut scratch.proj_flat,
&scratch.out_flat,
ipw,
ipw_dt,
None,
(bt, dm, 2 * di),
)?;
{
let bt_i = bt as i32;
let di_i = di as i32;
let gs_raw = scratch.gate_silu.cached_ptr();
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 prefill 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 = lw.conv1d_weight();
let cb = lw.conv1d_bias();
builder.arg(&cw);
builder.arg(&cb);
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 prefill L{layer_idx}: {e:?}"))?;
}
let (xpw, xpw_dt) = lw.x_proj_w();
gpu_gemm_forward_dispatch(
ctx,
&mut scratch.xdbl,
&scratch.u,
xpw,
xpw_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 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 prefill L{layer_idx}: {e:?}"))?;
}
let (dpw, dpw_dt) = lw.dt_proj_w();
gpu_gemm_forward_dispatch(
ctx,
&mut scratch.delta,
&scratch.dt_gather,
dpw,
dpw_dt,
Some(lw.dt_proj_b()),
(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 prefill 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 prefill 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 = lw.d_param();
builder.arg(&dp);
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 prefill 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 = lw.d_param();
builder.arg(&dp);
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 prefill 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 prefill L{layer_idx}: {e:?}"))?;
}
let (opw, opw_dt) = lw.out_proj_w();
gpu_gemm_forward_dispatch(
ctx,
&mut scratch.out_flat,
&scratch.gated,
opw,
opw_dt,
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 prefill 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 nfw = weights.norm_f_weight();
builder.arg(&nfw);
builder.arg(&bt_i);
builder.arg(&dm_i);
builder.arg(&eps);
unsafe { builder.launch(grid_norm(bt, dm)) }
.map_err(|e| format!("norm_f prefill: {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 prefill: {e:?}"))?;
}
let _ = gpu_sgemm_forward_raw;
Ok(())
}
pub fn gpu_forward_inference_prefill_mixed<W: MambaWeightsView>(
ctx: &GpuCtx,
target_temporal: &super::buffers::DtypedBuf,
inputs: PrefillInputs<'_, W>,
state: &mut GpuInferenceState,
scratch: &mut GpuMambaTargetMixedScratch,
) -> Result<(), String> {
let PrefillInputs {
ip_out_flat,
weights,
a_neg_all,
} = inputs;
let dt = scratch.dtype;
assert_eq!(
target_temporal.dtype(),
dt,
"target_temporal dtype must match mixed scratch dtype"
);
let dims: GpuMambaDims = 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;
let k = &ctx.kernels;
scratch.residual.copy_from(ip_out_flat, &ctx.stream)?;
let f32_sz = std::mem::size_of::<f32>() as u64;
for layer_idx in 0..weights.n_layers() {
let lw = weights.layer(layer_idx);
let conv_per_layer = b * di * d_conv;
let ssm_per_layer = b * di * ds;
let conv_ptr = state.conv.cached_ptr() + (layer_idx * conv_per_layer) as u64 * f32_sz;
let ssm_ptr = state.ssm.cached_ptr() + (layer_idx * ssm_per_layer) as u64 * f32_sz;
let a_neg_ptr = a_neg_all.cached_ptr() + (layer_idx * di * ds) as u64 * f32_sz;
{
let bt_i = bt as i32;
let dm_i = dm as i32;
let eps: f32 = dims.rms_norm_eps;
let mut bld = ctx.stream.launch_builder(k.rmsnorm_fwd_f32in_typed.get(dt));
let out_ptr = scratch.out_flat.cached_ptr();
let rms_ptr = scratch.rms_discard.cached_ptr();
let res_ptr = scratch.residual.cached_ptr();
bld.arg(&out_ptr);
bld.arg(&rms_ptr);
bld.arg(&res_ptr);
let nw = lw.norm_weight();
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 prefill L{layer_idx}: {e:?}"))?;
}
let (ipw, ipw_dt) = lw.in_proj_w();
gpu_gemm_typed_forward_raw(
ctx,
TypedPtr {
ptr: scratch.proj_flat.cached_ptr(),
dtype: dt,
},
TypedPtr {
ptr: scratch.out_flat.cached_ptr(),
dtype: dt,
},
TypedPtr {
ptr: ipw,
dtype: ipw_dt,
},
None,
(bt, dm, 2 * di),
)?;
{
let bt_i = bt as i32;
let di_i = di as i32;
let gs_raw = scratch.gate_silu.cached_ptr();
let mut bld = ctx.stream.launch_builder(k.split_gate_silu_typed.get(dt));
let xb_ptr = scratch.x_branch.cached_ptr();
let proj_ptr = scratch.proj_flat.cached_ptr();
bld.arg(&xb_ptr);
bld.arg(&gs_raw);
bld.arg(&gs_raw);
bld.arg(&proj_ptr);
bld.arg(&bt_i);
bld.arg(&di_i);
unsafe { bld.launch(grid_1d(bt * di)) }
.map_err(|e| format!("split_gate prefill 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 bld = ctx
.stream
.launch_builder(k.conv1d_burnin_nosave_typed.get(dt));
let u_ptr = scratch.u.cached_ptr();
let xb_ptr = scratch.x_branch.cached_ptr();
bld.arg(&u_ptr);
bld.arg(&conv_ptr);
bld.arg(&xb_ptr);
let cw = lw.conv1d_weight();
let cb = lw.conv1d_bias();
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_nosave prefill L{layer_idx}: {e:?}"))?;
}
let (xpw, xpw_dt) = lw.x_proj_w();
gpu_gemm_typed_forward_raw(
ctx,
TypedPtr {
ptr: scratch.xdbl.cached_ptr(),
dtype: dt,
},
TypedPtr {
ptr: scratch.u.cached_ptr(),
dtype: dt,
},
TypedPtr {
ptr: xpw,
dtype: xpw_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 dtg_ptr = scratch.dt_gather.cached_ptr();
let xdbl_ptr = scratch.xdbl.cached_ptr();
bld.arg(&dtg_ptr);
bld.arg(&xdbl_ptr);
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 dt prefill L{layer_idx}: {e:?}"))?;
}
let (dpw, dpw_dt) = lw.dt_proj_w();
gpu_gemm_typed_forward_raw(
ctx,
TypedPtr {
ptr: scratch.delta.cached_ptr(),
dtype: dt,
},
TypedPtr {
ptr: scratch.dt_gather.cached_ptr(),
dtype: dt,
},
TypedPtr {
ptr: dpw,
dtype: dpw_dt,
},
Some(lw.dt_proj_b()),
(bt, dt_rank, di),
)?;
{
let n = (bt * di) as i32;
let mut bld = ctx.stream.launch_builder(k.softplus_fwd_typed.get(dt));
let d_ptr = scratch.delta.cached_ptr();
bld.arg(&d_ptr);
bld.arg(&n);
unsafe { bld.launch(grid_1d(bt * di)) }
.map_err(|e| format!("softplus prefill 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 bld = ctx.stream.launch_builder(k.gather_bc_cols_typed.get(dt));
let bb_ptr = scratch.b_gathered.cached_ptr();
let cb_ptr = scratch.c_gathered.cached_ptr();
let xdbl_ptr = scratch.xdbl.cached_ptr();
bld.arg(&bb_ptr);
bld.arg(&cb_ptr);
bld.arg(&xdbl_ptr);
bld.arg(&bt_i);
bld.arg(&xdbl_i);
bld.arg(&ds_i);
bld.arg(&b_offset);
bld.arg(&c_offset);
unsafe { bld.launch(grid_1d(bt * ds)) }
.map_err(|e| format!("gather_bc prefill 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;
let _ = (PARALLEL_SCAN_THRESHOLD, grid_parallel_scan);
let mut bld = ctx.stream.launch_builder(k.ssm_burnin_nosave_typed.get(dt));
let y_ptr = scratch.y.cached_ptr();
let delta_ptr = scratch.delta.cached_ptr();
let u_ptr = scratch.u.cached_ptr();
let bb_ptr = scratch.b_gathered.cached_ptr();
let cb_ptr = scratch.c_gathered.cached_ptr();
bld.arg(&ssm_ptr);
bld.arg(&y_ptr);
bld.arg(&delta_ptr);
bld.arg(&u_ptr);
bld.arg(&bb_ptr);
bld.arg(&cb_ptr);
bld.arg(&a_neg_ptr);
let dp = lw.d_param();
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_nosave prefill L{layer_idx}: {e:?}"))?;
}
{
let n = (bt * di) as i32;
let mut bld = ctx.stream.launch_builder(k.elementwise_mul_typed.get(dt));
let gated_ptr = scratch.gated.cached_ptr();
let y_ptr = scratch.y.cached_ptr();
let gs_ptr = scratch.gate_silu.cached_ptr();
bld.arg(&gated_ptr);
bld.arg(&y_ptr);
bld.arg(&gs_ptr);
bld.arg(&n);
unsafe { bld.launch(grid_1d(bt * di)) }
.map_err(|e| format!("gating prefill L{layer_idx}: {e:?}"))?;
}
let (opw, opw_dt) = lw.out_proj_w();
gpu_gemm_typed_forward_raw(
ctx,
TypedPtr {
ptr: scratch.out_flat.cached_ptr(),
dtype: dt,
},
TypedPtr {
ptr: scratch.gated.cached_ptr(),
dtype: dt,
},
TypedPtr {
ptr: opw,
dtype: opw_dt,
},
None,
(bt, di, dm),
)?;
{
let n = (bt * dm) as i32;
let mut bld = ctx.stream.launch_builder(k.residual_add_f32_typed.get(dt));
let r_ptr = scratch.residual.cached_ptr();
let t_ptr = scratch.out_flat.cached_ptr();
bld.arg(&r_ptr);
bld.arg(&r_ptr);
bld.arg(&t_ptr);
bld.arg(&n);
unsafe { bld.launch(grid_1d(bt * dm)) }
.map_err(|e| format!("residual_add_f32 prefill L{layer_idx}: {e:?}"))?;
}
}
{
let bt_i = bt as i32;
let dm_i = dm as i32;
let eps: f32 = dims.rms_norm_eps;
let mut bld = ctx.stream.launch_builder(k.rmsnorm_fwd_f32in_typed.get(dt));
let out_ptr = scratch.out_flat.cached_ptr();
let rms_ptr = scratch.rms_discard.cached_ptr();
let res_ptr = scratch.residual.cached_ptr();
bld.arg(&out_ptr);
bld.arg(&rms_ptr);
bld.arg(&res_ptr);
let nfw = weights.norm_f_weight();
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!("norm_f prefill mixed: {e:?}"))?;
}
{
let b_i = b as i32;
let t_i = t as i32;
let dm_i = dm as i32;
let mut bld = ctx
.stream
.launch_builder(k.gather_last_timestep_typed.get(dt));
let dst_ptr = target_temporal.cached_ptr();
let src_ptr = scratch.out_flat.cached_ptr();
bld.arg(&dst_ptr);
bld.arg(&src_ptr);
bld.arg(&b_i);
bld.arg(&t_i);
bld.arg(&dm_i);
unsafe { bld.launch(grid_1d(b * dm)) }
.map_err(|e| format!("gather_last_typed prefill: {e:?}"))?;
}
Ok(())
}