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 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,
}
impl GpuMambaTrainingStepGraph {
#[allow(clippy::too_many_arguments)]
pub fn capture(
ctx: &GpuCtx,
cfg: &crate::config::MambaConfig,
train_w: &mut GpuMambaTrainMixedWeights,
adam: &GpuAdamW,
bias: &AdamWBiasFactors,
grads: &mut GpuMambaGrads,
acts: &mut GpuMambaBackboneMixedActs,
scratch: &mut GpuMambaMixedTrainScratch,
a_neg_all: &GpuBuffer,
mamba_input: &GpuBuffer,
d_temporal: &mut GpuBuffer,
state: &mut GpuRecurrentState,
batch: usize,
seq_len: usize,
) -> Result<Self, String> {
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 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 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,
})
}
#[allow(clippy::too_many_arguments)]
pub fn replay(
&self,
ctx: &GpuCtx,
train_w: &GpuMambaTrainMixedWeights,
adam: &GpuAdamW,
bias: &AdamWBiasFactors,
grads: &GpuMambaGrads,
a_neg_all: &GpuBuffer,
mamba_input: &GpuBuffer,
d_temporal: &GpuBuffer,
state: &GpuRecurrentState,
) -> Result<(), String> {
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?)"
);
self.graph
.launch()
.map_err(|e| format!("training_graph launch: {e:?}"))
}
}
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 {
#[allow(clippy::too_many_arguments)]
pub fn capture(
ctx: &GpuCtx,
cfg: &crate::config::MambaConfig,
weights: &mut GpuMambaTrainWeights,
adam: &GpuAdamW,
bias: &AdamWBiasFactors,
grads: &mut GpuMambaGrads,
acts: &mut GpuMambaBackboneActs,
scratch: &mut GpuMambaScratch,
a_neg_all: &GpuBuffer,
temporal: &mut GpuBuffer,
mamba_input: &GpuBuffer,
d_temporal: &mut GpuBuffer,
state: &mut GpuRecurrentState,
batch: usize,
seq_len: usize,
) -> Result<Self, String> {
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,
})
}
#[allow(clippy::too_many_arguments)]
pub fn replay(
&self,
weights: &GpuMambaTrainWeights,
adam: &GpuAdamW,
bias: &AdamWBiasFactors,
grads: &GpuMambaGrads,
temporal: &GpuBuffer,
a_neg_all: &GpuBuffer,
mamba_input: &GpuBuffer,
d_temporal: &GpuBuffer,
state: &GpuRecurrentState,
) -> Result<(), String> {
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:?}"))
}
}