use cudarc::driver::{CudaGraph, PushKernelArg};
use crate::mamba_ssm::gpu::adamw::{AdamWBiasFactors, GpuAdamW, step_m1_capturable};
use crate::mamba_ssm::gpu::backward::gpu_backward_mamba_backbone;
use crate::mamba_ssm::gpu::backward_mixed::gpu_backward_mamba_backbone_mixed;
use crate::mamba_ssm::gpu::buffers::GpuBuffer;
use crate::mamba_ssm::gpu::context::GpuCtx;
use crate::mamba_ssm::gpu::dtype::WeightDtype;
use crate::mamba_ssm::gpu::forward::{
GpuMambaBackboneActs, GpuMambaScratch, GpuRecurrentState, gpu_forward_mamba_backbone,
};
use crate::mamba_ssm::gpu::forward_mixed::{
GpuMambaBackboneMixedActs, GpuMambaMixedTrainScratch, gpu_forward_mamba_backbone_mixed,
};
use crate::mamba_ssm::gpu::graph_capture::capture_into_graph;
use crate::mamba_ssm::gpu::launch::grid_1d;
use crate::mamba_ssm::gpu::weights::{
GpuMambaGrads, GpuMambaTrainLayerWeights, GpuMambaTrainWeights,
};
use crate::mamba_ssm::gpu::weights_mixed_train::GpuMambaTrainMixedWeights;
fn recompute_a_neg_captured(
ctx: &GpuCtx,
master_layers: &[GpuMambaTrainLayerWeights],
a_neg_all: &GpuBuffer,
state_a_neg_all: &GpuBuffer,
d_inner: usize,
d_state: usize,
) -> Result<(), String> {
let per_layer = d_inner * d_state;
if per_layer == 0 {
return Ok(());
}
let n_i32 = per_layer as i32;
for (li, mw) in master_layers.iter().enumerate() {
let src = mw.a_log.cached_ptr();
let dst_a = a_neg_all.inner_at(li * per_layer);
let mut b1 = ctx.stream.launch_builder(&ctx.kernels.exp_negate);
b1.arg(&dst_a);
b1.arg(&src);
b1.arg(&n_i32);
unsafe { b1.launch(grid_1d(per_layer)) }
.map_err(|e| format!("exp_negate captured a_neg_all L{li}: {e:?}"))?;
let dst_s = state_a_neg_all.inner_at(li * per_layer);
let mut b2 = ctx.stream.launch_builder(&ctx.kernels.exp_negate);
b2.arg(&dst_s);
b2.arg(&src);
b2.arg(&n_i32);
unsafe { b2.launch(grid_1d(per_layer)) }
.map_err(|e| format!("exp_negate captured state.a_neg_all L{li}: {e:?}"))?;
}
Ok(())
}
pub struct MambaMixedCapture<'a> {
pub train_w: &'a mut GpuMambaTrainMixedWeights,
pub adam: &'a GpuAdamW,
pub bias: &'a AdamWBiasFactors,
pub grads: &'a mut GpuMambaGrads,
pub acts: &'a mut GpuMambaBackboneMixedActs,
pub scratch: &'a mut GpuMambaMixedTrainScratch,
pub a_neg_all: &'a GpuBuffer,
pub mamba_input: &'a GpuBuffer,
pub d_temporal: &'a mut GpuBuffer,
pub state: &'a mut GpuRecurrentState,
}
#[derive(Clone, Copy)]
pub struct MambaMixedReplay<'a> {
pub train_w: &'a GpuMambaTrainMixedWeights,
pub adam: &'a GpuAdamW,
pub bias: &'a AdamWBiasFactors,
pub grads: &'a GpuMambaGrads,
pub a_neg_all: &'a GpuBuffer,
pub mamba_input: &'a GpuBuffer,
pub d_temporal: &'a GpuBuffer,
pub state: &'a GpuRecurrentState,
}
pub struct GpuMambaTrainingStepGraph {
pub graph: CudaGraph,
pub batch: usize,
pub seq_len: usize,
pub dtype: WeightDtype,
captured_input_ptr: u64,
captured_d_temporal_ptr: u64,
captured_grads_flat_ptr: u64,
captured_adam_m_ptr: u64,
captured_adam_v_ptr: u64,
captured_bias_factors_ptr: u64,
captured_state_ssm_states_ptr: u64,
captured_state_conv_states_ptr: u64,
captured_state_a_neg_all_ptr: u64,
captured_a_neg_all_ptr: u64,
captured_master_input_proj_w_ptr: u64,
captured_master_norm_f_ptr: u64,
captured_compute_input_proj_w_ptr: u64,
captured_compute_norm_f_ptr: u64,
captured_half_staging_ptr: u64,
captured_bi_upcast_ptrs: [u64; 3],
}
impl GpuMambaTrainingStepGraph {
pub fn capture(
ctx: &GpuCtx,
cfg: &crate::config::MambaConfig,
cap: MambaMixedCapture<'_>,
batch: usize,
seq_len: usize,
) -> Result<Self, String> {
let MambaMixedCapture {
train_w,
adam,
bias,
grads,
acts,
scratch,
a_neg_all,
mamba_input,
d_temporal,
state,
} = cap;
assert!(
matches!(train_w.dtype, WeightDtype::Bf16),
"Step 14 graph capture supports bf16 only (f16 needs in-graph overflow check)"
);
assert_eq!(acts.dtype, WeightDtype::Bf16);
ctx.presize_half_staging_for_train(cfg, batch, seq_len, train_w.dtype)?;
let input_dim = mamba_input.len() / (batch * seq_len);
ctx.presize_bi_upcast_scratch_for_train_with_input(
cfg,
batch,
seq_len,
input_dim,
train_w.dtype,
)?;
let snap_input = mamba_input.cached_ptr();
let snap_d_temporal = d_temporal.cached_ptr();
let snap_grads_flat = grads.flat.cached_ptr();
let snap_adam_m = adam.m.cached_ptr();
let snap_adam_v = adam.v.cached_ptr();
let snap_bias = bias.ptr();
let snap_state_ssm = state.ssm_states.cached_ptr();
let snap_state_conv = state.conv_states.cached_ptr();
let snap_state_a_neg = state.a_neg_all.cached_ptr();
let snap_a_neg = a_neg_all.cached_ptr();
let snap_master_input = train_w.master.input_proj_w.cached_ptr();
let snap_master_norm_f = train_w.master.norm_f_weight.cached_ptr();
let snap_compute_input = train_w.compute.input_proj_w.ptr();
let snap_compute_norm_f = train_w.compute.norm_f_weight.ptr();
let snap_half_staging = ctx.half_staging_ptr();
let snap_bi_upcast = ctx.bi_upcast_scratch_ptrs();
let graph = capture_into_graph(&ctx.stream, || {
grads.zero(&ctx.stream)?;
gpu_forward_mamba_backbone_mixed(
ctx,
acts,
&train_w.compute,
mamba_input,
state,
scratch,
)?;
gpu_backward_mamba_backbone_mixed(
ctx,
d_temporal,
grads,
acts,
&train_w.compute,
a_neg_all,
scratch,
)?;
step_m1_capturable(
ctx,
&ctx.kernels.adamw_step_f32_capturable,
adam,
bias.ptr(),
&mut train_w.master,
grads,
)?;
train_w.sync_master_to_compute(ctx)?;
recompute_a_neg_captured(
ctx,
&train_w.master.layers,
a_neg_all,
&state.a_neg_all,
cfg.d_inner(),
cfg.d_state,
)?;
Ok(())
})?;
Ok(Self {
graph,
batch,
seq_len,
dtype: train_w.dtype,
captured_input_ptr: snap_input,
captured_d_temporal_ptr: snap_d_temporal,
captured_grads_flat_ptr: snap_grads_flat,
captured_adam_m_ptr: snap_adam_m,
captured_adam_v_ptr: snap_adam_v,
captured_bias_factors_ptr: snap_bias,
captured_state_ssm_states_ptr: snap_state_ssm,
captured_state_conv_states_ptr: snap_state_conv,
captured_state_a_neg_all_ptr: snap_state_a_neg,
captured_a_neg_all_ptr: snap_a_neg,
captured_master_input_proj_w_ptr: snap_master_input,
captured_master_norm_f_ptr: snap_master_norm_f,
captured_compute_input_proj_w_ptr: snap_compute_input,
captured_compute_norm_f_ptr: snap_compute_norm_f,
captured_half_staging_ptr: snap_half_staging,
captured_bi_upcast_ptrs: snap_bi_upcast,
})
}
pub fn replay(&self, ctx: &GpuCtx, rp: &MambaMixedReplay<'_>) -> Result<(), String> {
let MambaMixedReplay {
train_w,
adam,
bias,
grads,
a_neg_all,
mamba_input,
d_temporal,
state,
} = *rp;
assert_eq!(
mamba_input.cached_ptr(),
self.captured_input_ptr,
"training_graph replay: mamba_input pointer changed since capture"
);
assert_eq!(
d_temporal.cached_ptr(),
self.captured_d_temporal_ptr,
"training_graph replay: d_temporal pointer changed since capture"
);
assert_eq!(
grads.flat.cached_ptr(),
self.captured_grads_flat_ptr,
"training_graph replay: grads.flat pointer changed since capture"
);
assert_eq!(
adam.m.cached_ptr(),
self.captured_adam_m_ptr,
"training_graph replay: adam.m pointer changed since capture"
);
assert_eq!(
adam.v.cached_ptr(),
self.captured_adam_v_ptr,
"training_graph replay: adam.v pointer changed since capture"
);
assert_eq!(
bias.ptr(),
self.captured_bias_factors_ptr,
"training_graph replay: bias_factors pointer changed since capture"
);
assert_eq!(
state.ssm_states.cached_ptr(),
self.captured_state_ssm_states_ptr,
"training_graph replay: state.ssm_states pointer changed since capture"
);
assert_eq!(
state.conv_states.cached_ptr(),
self.captured_state_conv_states_ptr,
"training_graph replay: state.conv_states pointer changed since capture"
);
assert_eq!(
state.a_neg_all.cached_ptr(),
self.captured_state_a_neg_all_ptr,
"training_graph replay: state.a_neg_all pointer changed since capture"
);
assert_eq!(
a_neg_all.cached_ptr(),
self.captured_a_neg_all_ptr,
"training_graph replay: standalone a_neg_all pointer changed since capture"
);
assert_eq!(
train_w.master.input_proj_w.cached_ptr(),
self.captured_master_input_proj_w_ptr,
"training_graph replay: master input_proj_w pointer changed since capture"
);
assert_eq!(
train_w.master.norm_f_weight.cached_ptr(),
self.captured_master_norm_f_ptr,
"training_graph replay: master norm_f_weight pointer changed since capture"
);
assert_eq!(
train_w.compute.input_proj_w.ptr(),
self.captured_compute_input_proj_w_ptr,
"training_graph replay: compute input_proj_w pointer changed since capture"
);
assert_eq!(
train_w.compute.norm_f_weight.ptr(),
self.captured_compute_norm_f_ptr,
"training_graph replay: compute norm_f_weight pointer changed since capture"
);
assert_eq!(
ctx.half_staging_ptr(),
self.captured_half_staging_ptr,
"training_graph replay: half_staging pointer changed since capture \
(lazy grow during a previous step?)"
);
assert_eq!(
ctx.bi_upcast_scratch_ptrs(),
self.captured_bi_upcast_ptrs,
"training_graph replay: bi_upcast_scratch pointer changed since \
capture (a larger typed bi GEMM regrew the scratch after this \
graph was captured — re-capture or presize for the larger shape)"
);
self.graph
.launch()
.map_err(|e| format!("training_graph launch: {e:?}"))
}
}
pub struct MambaF32Capture<'a> {
pub weights: &'a mut GpuMambaTrainWeights,
pub adam: &'a GpuAdamW,
pub bias: &'a AdamWBiasFactors,
pub grads: &'a mut GpuMambaGrads,
pub acts: &'a mut GpuMambaBackboneActs,
pub scratch: &'a mut GpuMambaScratch,
pub a_neg_all: &'a GpuBuffer,
pub temporal: &'a mut GpuBuffer,
pub mamba_input: &'a GpuBuffer,
pub d_temporal: &'a mut GpuBuffer,
pub state: &'a mut GpuRecurrentState,
}
#[derive(Clone, Copy)]
pub struct MambaF32Replay<'a> {
pub weights: &'a GpuMambaTrainWeights,
pub adam: &'a GpuAdamW,
pub bias: &'a AdamWBiasFactors,
pub grads: &'a GpuMambaGrads,
pub temporal: &'a GpuBuffer,
pub a_neg_all: &'a GpuBuffer,
pub mamba_input: &'a GpuBuffer,
pub d_temporal: &'a GpuBuffer,
pub state: &'a GpuRecurrentState,
}
pub struct GpuMambaF32TrainingStepGraph {
pub graph: CudaGraph,
pub batch: usize,
pub seq_len: usize,
captured_input_ptr: u64,
captured_d_temporal_ptr: u64,
captured_grads_flat_ptr: u64,
captured_adam_m_ptr: u64,
captured_adam_v_ptr: u64,
captured_bias_factors_ptr: u64,
captured_state_ssm_states_ptr: u64,
captured_state_conv_states_ptr: u64,
captured_state_a_neg_all_ptr: u64,
captured_temporal_ptr: u64,
captured_a_neg_all_ptr: u64,
captured_weights_input_proj_w_ptr: u64,
captured_weights_norm_f_ptr: u64,
}
impl GpuMambaF32TrainingStepGraph {
pub fn capture(
ctx: &GpuCtx,
cfg: &crate::config::MambaConfig,
cap: MambaF32Capture<'_>,
batch: usize,
seq_len: usize,
) -> Result<Self, String> {
let MambaF32Capture {
weights,
adam,
bias,
grads,
acts,
scratch,
a_neg_all,
temporal,
mamba_input,
d_temporal,
state,
} = cap;
let snap_input = mamba_input.cached_ptr();
let snap_d_temporal = d_temporal.cached_ptr();
let snap_grads_flat = grads.flat.cached_ptr();
let snap_adam_m = adam.m.cached_ptr();
let snap_adam_v = adam.v.cached_ptr();
let snap_bias = bias.ptr();
let snap_state_ssm = state.ssm_states.cached_ptr();
let snap_state_conv = state.conv_states.cached_ptr();
let snap_state_a_neg = state.a_neg_all.cached_ptr();
let snap_temporal = temporal.cached_ptr();
let snap_a_neg = a_neg_all.cached_ptr();
let snap_input_proj = weights.input_proj_w.cached_ptr();
let snap_norm_f = weights.norm_f_weight.cached_ptr();
let cfg_local = *cfg;
let graph = capture_into_graph(&ctx.stream, || {
grads.zero(&ctx.stream)?;
gpu_forward_mamba_backbone(ctx, temporal, acts, weights, mamba_input, state, scratch)?;
gpu_backward_mamba_backbone(ctx, d_temporal, grads, acts, weights, a_neg_all, scratch)?;
crate::mamba_ssm::gpu::adamw::step_m1_capturable(
ctx,
&ctx.kernels.adamw_step_f32_capturable,
adam,
bias.ptr(),
weights,
grads,
)?;
recompute_a_neg_captured(
ctx,
&weights.layers,
a_neg_all,
&state.a_neg_all,
cfg_local.d_inner(),
cfg_local.d_state,
)?;
Ok(())
})?;
Ok(Self {
graph,
batch,
seq_len,
captured_input_ptr: snap_input,
captured_d_temporal_ptr: snap_d_temporal,
captured_grads_flat_ptr: snap_grads_flat,
captured_adam_m_ptr: snap_adam_m,
captured_adam_v_ptr: snap_adam_v,
captured_bias_factors_ptr: snap_bias,
captured_state_ssm_states_ptr: snap_state_ssm,
captured_state_conv_states_ptr: snap_state_conv,
captured_state_a_neg_all_ptr: snap_state_a_neg,
captured_temporal_ptr: snap_temporal,
captured_a_neg_all_ptr: snap_a_neg,
captured_weights_input_proj_w_ptr: snap_input_proj,
captured_weights_norm_f_ptr: snap_norm_f,
})
}
pub fn replay(&self, rp: &MambaF32Replay<'_>) -> Result<(), String> {
let MambaF32Replay {
weights,
adam,
bias,
grads,
temporal,
a_neg_all,
mamba_input,
d_temporal,
state,
} = *rp;
assert_eq!(
mamba_input.cached_ptr(),
self.captured_input_ptr,
"f32 training_graph replay: mamba_input pointer changed since capture"
);
assert_eq!(
d_temporal.cached_ptr(),
self.captured_d_temporal_ptr,
"f32 training_graph replay: d_temporal pointer changed since capture"
);
assert_eq!(
grads.flat.cached_ptr(),
self.captured_grads_flat_ptr,
"f32 training_graph replay: grads.flat pointer changed since capture"
);
assert_eq!(
adam.m.cached_ptr(),
self.captured_adam_m_ptr,
"f32 training_graph replay: adam.m pointer changed since capture"
);
assert_eq!(
adam.v.cached_ptr(),
self.captured_adam_v_ptr,
"f32 training_graph replay: adam.v pointer changed since capture"
);
assert_eq!(
bias.ptr(),
self.captured_bias_factors_ptr,
"f32 training_graph replay: bias_factors pointer changed since capture"
);
assert_eq!(
state.ssm_states.cached_ptr(),
self.captured_state_ssm_states_ptr,
"f32 training_graph replay: state.ssm_states pointer changed since capture"
);
assert_eq!(
state.conv_states.cached_ptr(),
self.captured_state_conv_states_ptr,
"f32 training_graph replay: state.conv_states pointer changed since capture"
);
assert_eq!(
state.a_neg_all.cached_ptr(),
self.captured_state_a_neg_all_ptr,
"f32 training_graph replay: state.a_neg_all pointer changed since capture"
);
assert_eq!(
temporal.cached_ptr(),
self.captured_temporal_ptr,
"f32 training_graph replay: temporal pointer changed since capture"
);
assert_eq!(
a_neg_all.cached_ptr(),
self.captured_a_neg_all_ptr,
"f32 training_graph replay: standalone a_neg_all pointer changed since capture"
);
assert_eq!(
weights.input_proj_w.cached_ptr(),
self.captured_weights_input_proj_w_ptr,
"f32 training_graph replay: input_proj_w pointer changed since capture"
);
assert_eq!(
weights.norm_f_weight.cached_ptr(),
self.captured_weights_norm_f_ptr,
"f32 training_graph replay: norm_f_weight pointer changed since capture"
);
self.graph
.launch()
.map_err(|e| format!("f32 training_graph launch: {e:?}"))
}
}