use super::blas::gpu_sgemm_forward_raw;
use super::buffers::GpuBuffer;
use super::context::GpuCtx;
use super::device::GpuDevice;
use super::launch::{grid_1d, grid_norm};
use super::weights::GpuMambaWeights;
use crate::config::MambaConfig;
use crate::weights::MambaWeights;
use cudarc::driver::PushKernelArg;
use std::sync::Arc;
pub struct GpuInferenceState {
pub conv: GpuBuffer,
pub ssm: GpuBuffer,
batch: usize,
d_inner: usize,
d_conv: usize,
d_state: usize,
}
impl GpuInferenceState {
pub fn zeros(
stream: &Arc<cudarc::driver::CudaStream>,
batch: usize,
cfg: &MambaConfig,
) -> Result<Self, String> {
let di = cfg.d_inner();
let conv_len = cfg.n_layers * batch * di * cfg.d_conv;
let ssm_len = cfg.n_layers * batch * di * cfg.d_state;
Ok(Self {
conv: GpuBuffer::zeros(stream, conv_len)?,
ssm: GpuBuffer::zeros(stream, ssm_len)?,
batch,
d_inner: di,
d_conv: cfg.d_conv,
d_state: cfg.d_state,
})
}
pub fn reset(&mut self, stream: &Arc<cudarc::driver::CudaStream>) -> Result<(), String> {
self.conv.zero(stream)?;
self.ssm.zero(stream)
}
pub fn conv_offset(&self, layer: usize) -> usize {
layer * self.batch * self.d_inner * self.d_conv
}
pub fn ssm_offset(&self, layer: usize) -> usize {
layer * self.batch * self.d_inner * self.d_state
}
pub fn batch(&self) -> usize {
self.batch
}
}
pub struct GpuInferenceScratch {
pub gpu_input: GpuBuffer,
pub temporal: GpuBuffer,
pub residual: GpuBuffer,
pub proj: GpuBuffer,
pub x_branch: GpuBuffer,
pub gate_silu: GpuBuffer,
pub u: GpuBuffer,
pub xdbl: GpuBuffer,
pub dt_gather: GpuBuffer,
pub delta: GpuBuffer,
pub b_buf: GpuBuffer,
pub c_buf: GpuBuffer,
pub y: GpuBuffer,
pub rms_buf: GpuBuffer,
}
impl GpuInferenceScratch {
pub fn new(
stream: &Arc<cudarc::driver::CudaStream>,
batch: usize,
cfg: &MambaConfig,
input_dim: usize,
) -> Result<Self, String> {
let dm = cfg.d_model;
let di = cfg.d_inner();
let ds = cfg.d_state;
let dt_rank = cfg.dt_rank();
let xdbl_dim = cfg.xdbl_dim();
Ok(Self {
gpu_input: GpuBuffer::zeros(stream, batch * input_dim)?,
temporal: GpuBuffer::zeros(stream, batch * dm)?,
residual: GpuBuffer::zeros(stream, batch * dm)?,
proj: GpuBuffer::zeros(stream, batch * 2 * di)?,
x_branch: GpuBuffer::zeros(stream, batch * di)?,
gate_silu: GpuBuffer::zeros(stream, batch * di)?,
u: GpuBuffer::zeros(stream, batch * di)?,
xdbl: GpuBuffer::zeros(stream, batch * xdbl_dim)?,
dt_gather: GpuBuffer::zeros(stream, batch * dt_rank)?,
delta: GpuBuffer::zeros(stream, batch * di)?,
b_buf: GpuBuffer::zeros(stream, batch * ds)?,
c_buf: GpuBuffer::zeros(stream, batch * ds)?,
y: GpuBuffer::zeros(stream, batch * di)?,
rms_buf: GpuBuffer::zeros(stream, batch)?,
})
}
}
pub struct GpuMambaInference {
ctx: GpuCtx,
weights: GpuMambaWeights,
a_neg_all: GpuBuffer,
cfg: MambaConfig,
input_dim: usize,
batch: usize,
graph: Option<cudarc::driver::CudaGraph>,
captured_state_ptr: u64,
captured_scratch_ptr: u64,
}
impl GpuMambaInference {
pub fn new(
device: &GpuDevice,
cpu_weights: &MambaWeights,
cfg: MambaConfig,
input_dim: usize,
batch: usize,
) -> Result<Self, String> {
cfg.validate()?;
let ctx = GpuCtx::new(device)?;
let weights = GpuMambaWeights::from_cpu(&ctx.stream, cpu_weights, &cfg)?;
let di = cfg.d_inner();
let ds = cfg.d_state;
let total_aneg = cfg.n_layers * di * ds;
let a_neg_all = GpuBuffer::zeros(&ctx.stream, total_aneg)?;
for (layer_idx, lw) in weights.layers.iter().enumerate() {
let offset = layer_idx * di * ds;
let dst_ptr = a_neg_all.raw_ptr_at(&ctx.stream, offset);
let src_ptr = lw.a_log.ptr();
let n_i = (di * ds) as i32;
let mut builder = ctx.stream.launch_builder(&ctx.kernels.exp_negate);
builder.arg(&dst_ptr);
builder.arg(&src_ptr);
builder.arg(&n_i);
unsafe { builder.launch(grid_1d(di * ds)) }
.map_err(|e| format!("exp_negate layer {layer_idx}: {e:?}"))?;
}
Ok(Self {
ctx,
weights,
a_neg_all,
cfg,
input_dim,
batch,
graph: None,
captured_state_ptr: 0,
captured_scratch_ptr: 0,
})
}
pub fn capture_graph(
&mut self,
state: &mut GpuInferenceState,
scratch: &mut GpuInferenceScratch,
) -> Result<(), String> {
self.ctx
.stream
.synchronize()
.map_err(|e| format!("pre-capture sync: {e:?}"))?;
self.ctx
.stream
.begin_capture(
cudarc::driver::sys::CUstreamCaptureMode::CU_STREAM_CAPTURE_MODE_THREAD_LOCAL,
)
.map_err(|e| format!("begin_capture: {e:?}"))?;
let capture_result = self.step_kernels(state, scratch);
if capture_result.is_err() {
let _ = self.ctx.stream.end_capture(
cudarc::driver::sys::CUgraphInstantiate_flags::CUDA_GRAPH_INSTANTIATE_FLAG_AUTO_FREE_ON_LAUNCH,
);
return capture_result.map(|_| ());
}
let graph = self.ctx.stream
.end_capture(
cudarc::driver::sys::CUgraphInstantiate_flags::CUDA_GRAPH_INSTANTIATE_FLAG_AUTO_FREE_ON_LAUNCH,
)
.map_err(|e| format!("end_capture: {e:?}"))?;
self.graph = graph;
self.captured_state_ptr = state.conv.cached_ptr();
self.captured_scratch_ptr = scratch.gpu_input.cached_ptr();
Ok(())
}
pub fn has_graph(&self) -> bool {
self.graph.is_some()
}
pub fn alloc_state(&self) -> Result<GpuInferenceState, String> {
GpuInferenceState::zeros(&self.ctx.stream, self.batch, &self.cfg)
}
pub fn alloc_scratch(&self) -> Result<GpuInferenceScratch, String> {
GpuInferenceScratch::new(&self.ctx.stream, self.batch, &self.cfg, self.input_dim)
}
pub fn step(
&self,
input: &[f32],
output: &mut [f32],
state: &mut GpuInferenceState,
scratch: &mut GpuInferenceScratch,
) -> Result<(), String> {
scratch.gpu_input.upload(&self.ctx.stream, input)?;
if let Some(ref g) = self.graph {
assert_eq!(
state.conv.cached_ptr(),
self.captured_state_ptr,
"CUDA Graph replay requires the same state buffers used during capture"
);
assert_eq!(
scratch.gpu_input.cached_ptr(),
self.captured_scratch_ptr,
"CUDA Graph replay requires the same scratch buffers used during capture"
);
g.launch().map_err(|e| format!("graph launch: {e:?}"))?;
} else {
self.step_kernels(state, scratch)?;
}
self.ctx
.stream
.synchronize()
.map_err(|e| format!("sync: {e:?}"))?;
scratch.temporal.download(&self.ctx.stream, output)?;
Ok(())
}
fn step_kernels(
&self,
state: &mut GpuInferenceState,
scratch: &mut GpuInferenceScratch,
) -> Result<(), String> {
let b = self.batch;
let cfg = &self.cfg;
let dm = cfg.d_model;
let di = cfg.d_inner();
let ds = cfg.d_state;
let dt_rank = cfg.dt_rank();
let xdbl_dim = cfg.xdbl_dim();
let d_conv = cfg.d_conv;
let k = &self.ctx.kernels;
gpu_sgemm_forward_raw(
&self.ctx,
&mut scratch.temporal, &scratch.gpu_input,
self.weights.input_proj_w.ptr(),
Some(self.weights.input_proj_b.ptr()),
(b, self.input_dim, dm),
)?;
let f32_sz = std::mem::size_of::<f32>() as u64;
for layer_idx in 0..cfg.n_layers {
let lw = &self.weights.layers[layer_idx];
let conv_ptr = state.conv.cached_ptr() + (state.conv_offset(layer_idx) as u64) * f32_sz;
let ssm_ptr = state.ssm.cached_ptr() + (state.ssm_offset(layer_idx) as u64) * f32_sz;
let aneg_ptr = self.a_neg_all.cached_ptr() + (layer_idx * di * ds) as u64 * f32_sz;
scratch
.residual
.copy_from_raw(&scratch.temporal, &self.ctx.stream)?;
{
let b_i = b as i32;
let dm_i = dm as i32;
let eps: f32 = 1e-5;
let mut bld = self.ctx.stream.launch_builder(&k.rmsnorm_fwd);
let t_ptr = scratch.temporal.cached_ptr();
let rms_ptr = scratch.rms_buf.cached_ptr();
let res_ptr = scratch.residual.cached_ptr();
bld.arg(&t_ptr); bld.arg(&rms_ptr);
bld.arg(&res_ptr); let nw = lw.norm_weight.ptr();
bld.arg(&nw);
bld.arg(&b_i);
bld.arg(&dm_i);
bld.arg(&eps);
unsafe { bld.launch(grid_norm(b, dm)) }
.map_err(|e| format!("rmsnorm_fwd L{layer_idx}: {e:?}"))?;
}
gpu_sgemm_forward_raw(
&self.ctx,
&mut scratch.proj,
&scratch.temporal,
lw.in_proj_w.ptr(),
None,
(b, dm, 2 * di),
)?;
{
let b_i = b as i32;
let di_i = di as i32;
let mut bld = self.ctx.stream.launch_builder(&k.split_gate_silu);
let xb_ptr = scratch.x_branch.cached_ptr();
bld.arg(&xb_ptr);
let g_ptr = scratch.gate_silu.cached_ptr();
let p_ptr = scratch.proj.cached_ptr();
bld.arg(&g_ptr); bld.arg(&g_ptr); bld.arg(&p_ptr);
bld.arg(&b_i);
bld.arg(&di_i);
unsafe { bld.launch(grid_1d(b * di)) }
.map_err(|e| format!("split_gate_silu L{layer_idx}: {e:?}"))?;
}
{
let b_i = b as i32;
let di_i = di as i32;
let dc_i = d_conv as i32;
let mut bld = self.ctx.stream.launch_builder(&k.conv1d_step_fwd);
let u_ptr = scratch.u.cached_ptr();
let xb_ptr2 = scratch.x_branch.cached_ptr();
bld.arg(&u_ptr);
bld.arg(&conv_ptr); bld.arg(&xb_ptr2);
let cw = lw.conv1d_weight.ptr();
let cb = lw.conv1d_bias.ptr();
bld.arg(&cw);
bld.arg(&cb);
bld.arg(&b_i);
bld.arg(&di_i);
bld.arg(&dc_i);
unsafe { bld.launch(grid_1d(b * di)) }
.map_err(|e| format!("conv1d_step L{layer_idx}: {e:?}"))?;
}
{
let n = (b * di) as i32;
let mut bld = self.ctx.stream.launch_builder(&k.silu_fwd);
let u_silu_ptr = scratch.u.cached_ptr();
bld.arg(&u_silu_ptr); bld.arg(&n);
unsafe { bld.launch(grid_1d(b * di)) }
.map_err(|e| format!("silu_fwd L{layer_idx}: {e:?}"))?;
}
gpu_sgemm_forward_raw(
&self.ctx,
&mut scratch.xdbl,
&scratch.u,
lw.x_proj_w.ptr(),
None,
(b, di, xdbl_dim),
)?;
{
let b_i = b as i32;
let xdbl_i = xdbl_dim as i32;
let dt_i = dt_rank as i32;
let offset: i32 = 0;
let mut bld = self.ctx.stream.launch_builder(&k.gather_cols);
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(&b_i);
bld.arg(&xdbl_i);
bld.arg(&dt_i);
bld.arg(&offset);
unsafe { bld.launch(grid_1d(b * dt_rank)) }
.map_err(|e| format!("gather_cols dt L{layer_idx}: {e:?}"))?;
}
gpu_sgemm_forward_raw(
&self.ctx,
&mut scratch.delta,
&scratch.dt_gather,
lw.dt_proj_w.ptr(),
Some(lw.dt_proj_b.ptr()),
(b, dt_rank, di),
)?;
{
let n = (b * di) as i32;
let mut bld = self.ctx.stream.launch_builder(&k.softplus_copy);
let d_ptr = scratch.delta.cached_ptr();
bld.arg(&d_ptr); bld.arg(&d_ptr); bld.arg(&n);
unsafe { bld.launch(grid_1d(b * di)) }
.map_err(|e| format!("softplus L{layer_idx}: {e:?}"))?;
}
{
let b_i = b 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 = self.ctx.stream.launch_builder(&k.gather_bc_cols);
let bb_ptr = scratch.b_buf.cached_ptr();
let cb_ptr = scratch.c_buf.cached_ptr();
let xdbl_bc_ptr = scratch.xdbl.cached_ptr();
bld.arg(&bb_ptr);
bld.arg(&cb_ptr);
bld.arg(&xdbl_bc_ptr);
bld.arg(&b_i);
bld.arg(&xdbl_i);
bld.arg(&ds_i);
bld.arg(&b_off);
bld.arg(&c_off);
unsafe { bld.launch(grid_1d(b * ds)) }
.map_err(|e| format!("gather_bc L{layer_idx}: {e:?}"))?;
}
{
let b_i = b as i32;
let di_i = di as i32;
let ds_i = ds as i32;
let dp = lw.d_param.ptr();
let mut bld = self.ctx.stream.launch_builder(&k.ssm_step_fwd);
let y_ssm_ptr = scratch.y.cached_ptr();
let delta_ssm_ptr = scratch.delta.cached_ptr();
let u_ssm_ptr = scratch.u.cached_ptr();
let b_ssm_ptr = scratch.b_buf.cached_ptr();
let c_ssm_ptr = scratch.c_buf.cached_ptr();
bld.arg(&ssm_ptr);
bld.arg(&y_ssm_ptr);
bld.arg(&delta_ssm_ptr);
bld.arg(&u_ssm_ptr);
bld.arg(&b_ssm_ptr);
bld.arg(&c_ssm_ptr);
bld.arg(&aneg_ptr);
bld.arg(&dp);
bld.arg(&b_i);
bld.arg(&di_i);
bld.arg(&ds_i);
unsafe { bld.launch(grid_1d(b * di)) }
.map_err(|e| format!("ssm_step L{layer_idx}: {e:?}"))?;
}
{
let n = (b * di) as i32;
let mut bld = self.ctx.stream.launch_builder(&k.elementwise_mul);
let y_ptr = scratch.y.cached_ptr();
let gs_ptr = scratch.gate_silu.cached_ptr();
bld.arg(&y_ptr);
bld.arg(&y_ptr);
bld.arg(&gs_ptr);
bld.arg(&n);
unsafe { bld.launch(grid_1d(b * di)) }
.map_err(|e| format!("gating L{layer_idx}: {e:?}"))?;
}
gpu_sgemm_forward_raw(
&self.ctx,
&mut scratch.temporal,
&scratch.y,
lw.out_proj_w.ptr(),
None,
(b, di, dm),
)?;
{
let n = (b * dm) as i32;
let mut bld = self.ctx.stream.launch_builder(&k.residual_add);
let t_ptr = scratch.temporal.cached_ptr();
let r_ptr = scratch.residual.cached_ptr();
bld.arg(&t_ptr);
bld.arg(&r_ptr);
bld.arg(&t_ptr); bld.arg(&n);
unsafe { bld.launch(grid_1d(b * dm)) }
.map_err(|e| format!("residual L{layer_idx}: {e:?}"))?;
}
}
{
let b_i = b as i32;
let dm_i = dm as i32;
let eps: f32 = 1e-5;
let mut bld = self.ctx.stream.launch_builder(&k.rmsnorm_fwd);
let t_ptr = scratch.temporal.cached_ptr();
let rms_ptr = scratch.rms_buf.cached_ptr();
bld.arg(&t_ptr);
bld.arg(&rms_ptr);
bld.arg(&t_ptr);
let nfw = self.weights.norm_f_weight.ptr();
bld.arg(&nfw);
bld.arg(&b_i);
bld.arg(&dm_i);
bld.arg(&eps);
unsafe { bld.launch(grid_norm(b, dm)) }.map_err(|e| format!("norm_f: {e:?}"))?;
}
Ok(())
}
pub fn config(&self) -> &MambaConfig {
&self.cfg
}
pub fn batch(&self) -> usize {
self.batch
}
}
pub struct GpuMambaBackbone {
engine: GpuMambaInference,
state: GpuInferenceState,
scratch: GpuInferenceScratch,
}
impl GpuMambaBackbone {
pub fn new(
gpu_ordinal: usize,
cpu_weights: &MambaWeights,
cfg: MambaConfig,
input_dim: usize,
batch: usize,
) -> Result<Self, String> {
let device = GpuDevice::new(gpu_ordinal)?;
let engine = GpuMambaInference::new(&device, cpu_weights, cfg, input_dim, batch)?;
let state = engine.alloc_state()?;
let scratch = engine.alloc_scratch()?;
Ok(Self {
engine,
state,
scratch,
})
}
pub fn step(&mut self, input: &[f32], output: &mut [f32]) -> Result<(), String> {
self.engine
.step(input, output, &mut self.state, &mut self.scratch)
}
pub fn reset(&mut self) -> Result<(), String> {
self.state.reset(&self.engine.ctx.stream)
}
pub fn capture_graph(&mut self) -> Result<(), String> {
let input = vec![0.0f32; self.engine.batch * self.engine.input_dim];
let mut output = vec![0.0f32; self.engine.batch * self.engine.cfg.d_model];
self.engine
.step(&input, &mut output, &mut self.state, &mut self.scratch)?;
self.state.reset(&self.engine.ctx.stream)?;
self.engine
.capture_graph(&mut self.state, &mut self.scratch)
}
pub fn config(&self) -> &MambaConfig {
self.engine.config()
}
pub fn batch(&self) -> usize {
self.engine.batch()
}
pub fn has_graph(&self) -> bool {
self.engine.has_graph()
}
}