use std::sync::Arc;
use cudarc::driver::{CudaFunction, CudaStream, PushKernelArg};
use crate::mamba_ssm::gpu::buffers::{GpuBuffer, GradSlice};
use crate::mamba_ssm::gpu::context::GpuCtx;
use crate::mamba_ssm::gpu::launch::grid_1d;
pub struct AdamWBiasFactors {
pub buf: GpuBuffer,
}
impl AdamWBiasFactors {
pub fn new(stream: &Arc<CudaStream>) -> Result<Self, String> {
let buf = GpuBuffer::zeros(stream, 2)?;
let mut this = Self { buf };
this.write(stream, 1.0, 1.0)?;
Ok(this)
}
pub fn write(&mut self, stream: &Arc<CudaStream>, bc1: f32, bc2: f32) -> Result<(), String> {
debug_assert!(
bc1.is_finite() && bc2.is_finite() && bc1 > 0.0 && bc2 > 0.0,
"AdamWBiasFactors::write got non-finite or non-positive values: bc1={bc1} bc2={bc2}"
);
self.buf.upload(stream, &[bc1, bc2])
}
pub fn ptr(&self) -> cudarc::driver::sys::CUdeviceptr {
self.buf.cached_ptr()
}
}
#[derive(Clone, Copy, Debug)]
pub struct AdamWParamPtrs {
pub weight: cudarc::driver::sys::CUdeviceptr,
pub grad: cudarc::driver::sys::CUdeviceptr,
pub m: cudarc::driver::sys::CUdeviceptr,
pub v: cudarc::driver::sys::CUdeviceptr,
}
pub struct GpuAdamW {
pub m: GpuBuffer,
pub v: GpuBuffer,
pub step: u64,
pub lr: f32,
pub beta1: f32,
pub beta2: f32,
pub eps: f32,
pub weight_decay: f32,
pub reference_no_decay: bool,
}
impl GpuAdamW {
pub fn new(stream: &Arc<CudaStream>, n_params: usize) -> Result<Self, String> {
Ok(Self {
m: GpuBuffer::zeros(stream, n_params)?,
v: GpuBuffer::zeros(stream, n_params)?,
step: 0,
lr: 1e-3,
beta1: 0.9,
beta2: 0.999,
eps: 1e-8,
weight_decay: 1e-2,
reference_no_decay: false,
})
}
#[must_use]
pub fn with_reference_no_decay(mut self, on: bool) -> Self {
self.reference_no_decay = on;
self
}
#[must_use]
pub fn with_lr(mut self, lr: f32) -> Self {
assert!(lr.is_finite() && lr >= 0.0, "lr must be finite and >= 0");
self.lr = lr;
self
}
#[must_use]
pub fn with_betas(mut self, beta1: f32, beta2: f32) -> Self {
assert!(
(0.0..1.0).contains(&beta1) && (0.0..1.0).contains(&beta2),
"betas must be in [0, 1), got beta1={beta1} beta2={beta2}"
);
self.beta1 = beta1;
self.beta2 = beta2;
self
}
#[must_use]
pub fn with_eps(mut self, eps: f32) -> Self {
assert!(eps.is_finite() && eps > 0.0, "eps must be finite and > 0");
self.eps = eps;
self
}
#[must_use]
pub fn with_weight_decay(mut self, wd: f32) -> Self {
assert!(
wd.is_finite() && wd >= 0.0,
"weight_decay must be finite and >= 0"
);
self.weight_decay = wd;
self
}
pub fn zero_state(&mut self, stream: &Arc<CudaStream>) -> Result<(), String> {
self.m.zero(stream)?;
self.v.zero(stream)?;
self.step = 0;
Ok(())
}
pub fn state(&self) -> (u64, f32) {
(self.step, self.lr)
}
pub fn step_one(
&self,
ctx: &GpuCtx,
adamw_kernel: &CudaFunction,
ptrs: AdamWParamPtrs,
len: usize,
bias_c1: f32,
bias_c2: f32,
) -> Result<(), String> {
if len == 0 {
return Ok(());
}
let n = len as i32;
let cfg = grid_1d(len);
let mut bld = ctx.stream.launch_builder(adamw_kernel);
bld.arg(&ptrs.weight);
bld.arg(&ptrs.grad);
bld.arg(&ptrs.m);
bld.arg(&ptrs.v);
bld.arg(&self.lr);
bld.arg(&self.beta1);
bld.arg(&self.beta2);
bld.arg(&self.eps);
bld.arg(&self.weight_decay);
bld.arg(&bias_c1);
bld.arg(&bias_c2);
bld.arg(&n);
unsafe { bld.launch(cfg) }.map_err(|e| format!("adamw_step_f32: {e:?}"))?;
Ok(())
}
pub fn step_one_capturable(
&self,
ctx: &GpuCtx,
adamw_kernel: &CudaFunction,
ptrs: AdamWParamPtrs,
bias_factors_ptr: cudarc::driver::sys::CUdeviceptr,
len: usize,
) -> Result<(), String> {
self.step_one_capturable_wd(
ctx,
adamw_kernel,
ptrs,
bias_factors_ptr,
len,
self.weight_decay,
)
}
pub fn step_one_capturable_wd(
&self,
ctx: &GpuCtx,
adamw_kernel: &CudaFunction,
ptrs: AdamWParamPtrs,
bias_factors_ptr: cudarc::driver::sys::CUdeviceptr,
len: usize,
weight_decay: f32,
) -> Result<(), String> {
if len == 0 {
return Ok(());
}
let n = len as i32;
let cfg = grid_1d(len);
let mut bld = ctx.stream.launch_builder(adamw_kernel);
bld.arg(&ptrs.weight);
bld.arg(&ptrs.grad);
bld.arg(&ptrs.m);
bld.arg(&ptrs.v);
bld.arg(&self.lr);
bld.arg(&self.beta1);
bld.arg(&self.beta2);
bld.arg(&self.eps);
bld.arg(&weight_decay);
bld.arg(&bias_factors_ptr);
bld.arg(&n);
unsafe { bld.launch(cfg) }.map_err(|e| format!("adamw_step_f32_capturable: {e:?}"))?;
Ok(())
}
pub fn advance(&mut self) -> (u64, f32, f32) {
self.step += 1;
let t = self.step.min(1 << 30) as i32;
let denom1 = 1.0 - (self.beta1 as f64).powi(t);
let denom2 = 1.0 - (self.beta2 as f64).powi(t);
let bias_c1 = (1.0 / denom1.max(1e-30)) as f32;
let bias_c2 = (1.0 / denom2.max(1e-30)) as f32;
(self.step, bias_c1, bias_c2)
}
}
pub fn run_pairs(
ctx: &GpuCtx,
adamw_kernel: &CudaFunction,
adam: &GpuAdamW,
bias_c1: f32,
bias_c2: f32,
flat_grad_base: cudarc::driver::sys::CUdeviceptr,
pairs: &[(&GpuBuffer, &GradSlice)],
) -> Result<(), String> {
let m_base = adam.m.cached_ptr();
let v_base = adam.v.cached_ptr();
for (w, g) in pairs {
if w.is_empty() {
continue;
}
if g.len() != w.len() {
return Err(format!(
"adamw: weight/grad len mismatch: w={} g={}",
w.len(),
g.len()
));
}
let g_ptr = g.ptr();
let off_bytes = g_ptr - flat_grad_base;
let off_elems = off_bytes / 4;
let m_ptr = m_base + off_bytes;
let v_ptr = v_base + off_bytes;
debug_assert!(
off_elems as usize + g.len() <= adam.m.len(),
"adamw m/v slice OOB: off_elems={off_elems} len={} m.len={}",
g.len(),
adam.m.len()
);
adam.step_one(
ctx,
adamw_kernel,
AdamWParamPtrs {
weight: w.cached_ptr(),
grad: g_ptr,
m: m_ptr,
v: v_ptr,
},
g.len(),
bias_c1,
bias_c2,
)?;
}
Ok(())
}
pub fn run_pairs_capturable(
ctx: &GpuCtx,
adamw_kernel: &CudaFunction,
adam: &GpuAdamW,
bias_factors_ptr: cudarc::driver::sys::CUdeviceptr,
flat_grad_base: cudarc::driver::sys::CUdeviceptr,
pairs: &[(&GpuBuffer, &GradSlice)],
) -> Result<(), String> {
for pair in pairs {
run_one_capturable_wd(
ctx,
adamw_kernel,
adam,
bias_factors_ptr,
flat_grad_base,
*pair,
adam.weight_decay,
)?;
}
Ok(())
}
fn run_one_capturable_wd(
ctx: &GpuCtx,
adamw_kernel: &CudaFunction,
adam: &GpuAdamW,
bias_factors_ptr: cudarc::driver::sys::CUdeviceptr,
flat_grad_base: cudarc::driver::sys::CUdeviceptr,
(w, g): (&GpuBuffer, &GradSlice),
weight_decay: f32,
) -> Result<(), String> {
if w.is_empty() {
return Ok(());
}
if g.len() != w.len() {
return Err(format!(
"adamw: weight/grad len mismatch: w={} g={}",
w.len(),
g.len()
));
}
let g_ptr = g.ptr();
let off_bytes = g_ptr - flat_grad_base;
let off_elems = off_bytes / 4;
let m_ptr = adam.m.cached_ptr() + off_bytes;
let v_ptr = adam.v.cached_ptr() + off_bytes;
debug_assert!(
off_elems as usize + g.len() <= adam.m.len(),
"adamw m/v slice OOB"
);
adam.step_one_capturable_wd(
ctx,
adamw_kernel,
AdamWParamPtrs {
weight: w.cached_ptr(),
grad: g_ptr,
m: m_ptr,
v: v_ptr,
},
bias_factors_ptr,
g.len(),
weight_decay,
)
}
pub fn step_m1(
ctx: &GpuCtx,
adamw_kernel: &CudaFunction,
adam: &mut GpuAdamW,
weights: &mut crate::mamba_ssm::gpu::weights::GpuMambaTrainWeights,
grads: &crate::mamba_ssm::gpu::weights::GpuMambaGrads,
) -> Result<(), String> {
let (_, bc1, bc2) = adam.advance();
let flat_base = grads.flat.cached_ptr();
let mut pairs: Vec<(&GpuBuffer, &GradSlice)> =
Vec::with_capacity(3 + 10 * weights.layers.len());
pairs.push((&weights.input_proj_w, &grads.input_proj_w));
pairs.push((&weights.input_proj_b, &grads.input_proj_b));
for (lw, lg) in weights.layers.iter().zip(&grads.layers) {
pairs.push((&lw.norm_weight, &lg.norm_weight));
pairs.push((&lw.in_proj_w, &lg.in_proj_w));
pairs.push((&lw.conv1d_weight, &lg.conv1d_weight));
pairs.push((&lw.conv1d_bias, &lg.conv1d_bias));
pairs.push((&lw.x_proj_w, &lg.x_proj_w));
pairs.push((&lw.dt_proj_w, &lg.dt_proj_w));
pairs.push((&lw.dt_proj_b, &lg.dt_proj_b));
pairs.push((&lw.a_log, &lg.a_log));
pairs.push((&lw.d_param, &lg.d_param));
pairs.push((&lw.out_proj_w, &lg.out_proj_w));
}
pairs.push((&weights.norm_f_weight, &grads.norm_f_weight));
run_pairs(ctx, adamw_kernel, adam, bc1, bc2, flat_base, &pairs)
}
pub fn step_m1_capturable(
ctx: &GpuCtx,
adamw_kernel: &CudaFunction,
adam: &GpuAdamW,
bias_factors_ptr: cudarc::driver::sys::CUdeviceptr,
weights: &mut crate::mamba_ssm::gpu::weights::GpuMambaTrainWeights,
grads: &crate::mamba_ssm::gpu::weights::GpuMambaGrads,
) -> Result<(), String> {
let flat_base = grads.flat.cached_ptr();
let mut pairs: Vec<(&GpuBuffer, &GradSlice, bool)> =
Vec::with_capacity(3 + 10 * weights.layers.len());
pairs.push((&weights.input_proj_w, &grads.input_proj_w, false));
pairs.push((&weights.input_proj_b, &grads.input_proj_b, false));
for (lw, lg) in weights.layers.iter().zip(&grads.layers) {
pairs.push((&lw.norm_weight, &lg.norm_weight, true));
pairs.push((&lw.in_proj_w, &lg.in_proj_w, false));
pairs.push((&lw.conv1d_weight, &lg.conv1d_weight, false));
pairs.push((&lw.conv1d_bias, &lg.conv1d_bias, false));
pairs.push((&lw.x_proj_w, &lg.x_proj_w, false));
pairs.push((&lw.dt_proj_w, &lg.dt_proj_w, false));
pairs.push((&lw.dt_proj_b, &lg.dt_proj_b, true));
pairs.push((&lw.a_log, &lg.a_log, true));
pairs.push((&lw.d_param, &lg.d_param, true));
pairs.push((&lw.out_proj_w, &lg.out_proj_w, false));
}
pairs.push((&weights.norm_f_weight, &grads.norm_f_weight, true));
for (w, g, no_decay) in pairs {
let wd = if no_decay && adam.reference_no_decay {
0.0
} else {
adam.weight_decay
};
run_one_capturable_wd(
ctx,
adamw_kernel,
adam,
bias_factors_ptr,
flat_base,
(w, g),
wd,
)?;
}
Ok(())
}
pub fn step_m3_capturable(
ctx: &GpuCtx,
adamw_kernel: &CudaFunction,
adam: &GpuAdamW,
bias_factors_ptr: cudarc::driver::sys::CUdeviceptr,
weights: &mut crate::mamba3_siso::gpu::weights::GpuMamba3Weights,
grads: &crate::mamba3_siso::gpu::weights::GpuMamba3Grads,
) -> Result<(), String> {
let flat_base = grads.flat.cached_ptr();
let mut pairs: Vec<(&GpuBuffer, &GradSlice, bool)> =
Vec::with_capacity(3 + 10 * weights.layers.len());
pairs.push((&weights.input_proj_w, &grads.input_proj_w, false));
pairs.push((&weights.input_proj_b, &grads.input_proj_b, false));
for (lw, lg) in weights.layers.iter().zip(&grads.layers) {
pairs.push((&lw.norm_weight, &lg.norm_weight, true));
pairs.push((&lw.in_proj_w, &lg.in_proj_w, false));
pairs.push((&lw.dt_bias, &lg.dt_bias, true));
pairs.push((&lw.b_norm_weight, &lg.b_norm_weight, true));
pairs.push((&lw.c_norm_weight, &lg.c_norm_weight, true));
pairs.push((&lw.b_bias, &lg.b_bias, false));
pairs.push((&lw.c_bias, &lg.c_bias, false));
pairs.push((&lw.d_param, &lg.d_param, true));
pairs.push((&lw.norm_gate_weight, &lg.norm_gate_weight, true));
pairs.push((&lw.out_proj_w, &lg.out_proj_w, false));
}
pairs.push((&weights.norm_f_weight, &grads.norm_f_weight, true));
for (w, g, no_decay) in pairs {
let wd = if no_decay && adam.reference_no_decay {
0.0
} else {
adam.weight_decay
};
run_one_capturable_wd(
ctx,
adamw_kernel,
adam,
bias_factors_ptr,
flat_base,
(w, g),
wd,
)?;
}
Ok(())
}
pub fn step_m3(
ctx: &GpuCtx,
adamw_kernel: &CudaFunction,
adam: &mut GpuAdamW,
weights: &mut crate::mamba3_siso::gpu::weights::GpuMamba3Weights,
grads: &crate::mamba3_siso::gpu::weights::GpuMamba3Grads,
) -> Result<(), String> {
let (_, bc1, bc2) = adam.advance();
let flat_base = grads.flat.cached_ptr();
let mut pairs: Vec<(&GpuBuffer, &GradSlice)> =
Vec::with_capacity(3 + 10 * weights.layers.len());
pairs.push((&weights.input_proj_w, &grads.input_proj_w));
pairs.push((&weights.input_proj_b, &grads.input_proj_b));
for (lw, lg) in weights.layers.iter().zip(&grads.layers) {
pairs.push((&lw.norm_weight, &lg.norm_weight));
pairs.push((&lw.in_proj_w, &lg.in_proj_w));
pairs.push((&lw.dt_bias, &lg.dt_bias));
pairs.push((&lw.b_norm_weight, &lg.b_norm_weight));
pairs.push((&lw.c_norm_weight, &lg.c_norm_weight));
pairs.push((&lw.b_bias, &lg.b_bias));
pairs.push((&lw.c_bias, &lg.c_bias));
pairs.push((&lw.d_param, &lg.d_param));
pairs.push((&lw.norm_gate_weight, &lg.norm_gate_weight));
pairs.push((&lw.out_proj_w, &lg.out_proj_w));
}
pairs.push((&weights.norm_f_weight, &grads.norm_f_weight));
run_pairs(ctx, adamw_kernel, adam, bc1, bc2, flat_base, &pairs)
}
#[cfg(test)]
mod cpu_state_tests {
#[test]
fn bias_correction_step_one() {
let beta1 = 0.9_f64;
let beta2 = 0.999_f64;
let bc1 = (1.0 / (1.0 - beta1.powi(1))) as f32;
let bc2 = (1.0 / (1.0 - beta2.powi(1))) as f32;
assert!((bc1 - 10.0).abs() < 1e-3, "bc1={bc1}");
assert!((bc2 - 1000.0).abs() < 1e-1, "bc2={bc2}");
}
#[test]
fn bias_correction_step_large() {
let beta1 = 0.9_f64;
let beta2 = 0.999_f64;
let bc1 = (1.0 / (1.0 - beta1.powi(2000))) as f32;
let bc2 = (1.0 / (1.0 - beta2.powi(2000))) as f32;
assert!(bc1 < 1.001, "bc1={bc1}");
assert!(bc2 < 1.2, "bc2={bc2}");
}
}