use std::sync::Arc;
use cudarc::driver::{CudaSlice, CudaStream, PushKernelArg};
use crate::mamba_ssm::gpu::buffers::GpuBuffer;
use crate::mamba_ssm::gpu::context::GpuCtx;
use crate::mamba_ssm::gpu::kernels::MambaKernels;
use crate::mamba_ssm::gpu::launch::grid_1d;
#[derive(Clone, Debug)]
pub struct DynamicLossScaler {
scale: f32,
growth_factor: f32,
backoff_factor: f32,
growth_interval: u32,
growth_tracker: u32,
max_scale: f32,
min_scale: f32,
enabled: bool,
}
impl Default for DynamicLossScaler {
fn default() -> Self {
Self::new()
}
}
impl DynamicLossScaler {
pub fn new() -> Self {
Self {
scale: 65_536.0, growth_factor: 2.0,
backoff_factor: 0.5,
growth_interval: 2_000,
growth_tracker: 0,
max_scale: 16_777_216.0, min_scale: 1.0,
enabled: true,
}
}
pub fn disabled() -> Self {
Self {
enabled: false,
scale: 1.0,
..Self::new()
}
}
#[must_use]
pub fn with_init_scale(mut self, init_scale: f32) -> Self {
assert!(
init_scale.is_finite() && init_scale > 0.0,
"init_scale must be finite and positive, got {init_scale}"
);
self.scale = init_scale;
self
}
#[must_use]
pub fn with_growth_interval(mut self, n: u32) -> Self {
assert!(n > 0, "growth_interval must be > 0");
self.growth_interval = n;
self
}
#[must_use]
pub fn with_max_scale(mut self, s: f32) -> Self {
assert!(
s.is_finite() && s > 0.0,
"max_scale must be finite and positive, got {s}"
);
assert!(
s >= self.min_scale,
"max_scale ({s}) must be >= min_scale ({})",
self.min_scale
);
self.max_scale = s;
self
}
#[must_use]
pub fn with_min_scale(mut self, s: f32) -> Self {
assert!(
s.is_finite() && s > 0.0,
"min_scale must be finite and positive, got {s}"
);
assert!(
s <= self.max_scale,
"min_scale ({s}) must be <= max_scale ({})",
self.max_scale
);
self.min_scale = s;
self
}
pub fn state(&self) -> (f32, u32) {
(self.scale, self.growth_tracker)
}
pub fn load_state(&mut self, scale: f32, growth_tracker: u32) {
assert!(
scale.is_finite() && scale > 0.0,
"loaded scale must be finite and positive, got {scale}"
);
self.scale = scale.clamp(self.min_scale, self.max_scale);
self.growth_tracker = growth_tracker.min(self.growth_interval.saturating_sub(1));
}
pub fn scale(&self) -> f32 {
self.scale
}
pub fn clean_step_count(&self) -> u32 {
self.growth_tracker
}
pub fn enabled(&self) -> bool {
self.enabled
}
pub fn update(&mut self, had_overflow: bool) {
if !self.enabled {
return;
}
if had_overflow {
self.scale = (self.scale * self.backoff_factor).max(self.min_scale);
self.growth_tracker = 0;
} else {
self.growth_tracker = self.growth_tracker.saturating_add(1);
if self.growth_tracker >= self.growth_interval {
self.scale = (self.scale * self.growth_factor).min(self.max_scale);
self.growth_tracker = 0;
}
}
}
}
pub struct OverflowFlag {
data: CudaSlice<i32>,
}
impl OverflowFlag {
pub fn new(stream: &Arc<CudaStream>) -> Result<Self, String> {
let data = stream
.alloc_zeros::<i32>(1)
.map_err(|e| format!("OverflowFlag alloc: {:?}", e))?;
Ok(Self { data })
}
pub fn zero(&mut self, stream: &Arc<CudaStream>) -> Result<(), String> {
stream
.memset_zeros(&mut self.data)
.map_err(|e| format!("OverflowFlag zero: {:?}", e))
}
pub fn read(&self, stream: &Arc<CudaStream>) -> Result<i32, String> {
let host = stream
.clone_dtoh(&self.data)
.map_err(|e| format!("OverflowFlag read: {:?}", e))?;
Ok(host[0])
}
pub(crate) fn cuda_slice(&mut self) -> &mut CudaSlice<i32> {
&mut self.data
}
pub fn stable_ptr(&self, stream: &Arc<CudaStream>) -> cudarc::driver::sys::CUdeviceptr {
use cudarc::driver::DevicePtr;
let (p, _g) = self.data.device_ptr(stream);
p
}
}
pub fn check_inf_nan_gpu(
ctx: &GpuCtx,
kernels: &MambaKernels,
flag: &mut OverflowFlag,
grads: &GpuBuffer,
) -> Result<(), String> {
if grads.is_empty() {
return Ok(());
}
let n = grads.len() as i32;
let cfg = grid_1d(grads.len());
let mut bld = ctx.stream.launch_builder(&kernels.check_inf_nan_f32);
let flag_ptr = {
use cudarc::driver::DevicePtr;
let (p, _g) = flag.cuda_slice().device_ptr(&ctx.stream);
p
};
let grad_ptr = grads.cached_ptr();
bld.arg(&flag_ptr);
bld.arg(&grad_ptr);
bld.arg(&n);
unsafe { bld.launch(cfg) }.map_err(|e| format!("check_inf_nan_f32: {:?}", e))?;
Ok(())
}
pub struct UnscaleFactor {
buf: GpuBuffer,
}
impl UnscaleFactor {
pub fn new(stream: &Arc<CudaStream>) -> Result<Self, String> {
Ok(Self {
buf: GpuBuffer::zeros(stream, 1)?,
})
}
pub fn write(&mut self, stream: &Arc<CudaStream>, unscale: f32) -> Result<(), String> {
self.buf.upload(stream, &[unscale])
}
pub fn ptr(&self) -> cudarc::driver::sys::CUdeviceptr {
self.buf.cached_ptr()
}
}
pub fn scale_grads_skip_gpu(
ctx: &GpuCtx,
kernels: &MambaKernels,
flag: &mut OverflowFlag,
grads: &mut GpuBuffer,
unscale: &UnscaleFactor,
) -> Result<(), String> {
if grads.is_empty() {
return Ok(());
}
let n = grads.len() as i32;
let cfg = grid_1d(grads.len());
let mut bld = ctx.stream.launch_builder(&kernels.scale_grads_skip_f32);
let flag_ptr = {
use cudarc::driver::DevicePtr;
let (p, _g) = flag.cuda_slice().device_ptr(&ctx.stream);
p
};
let grad_ptr = grads.cached_ptr();
let unscale_ptr = unscale.ptr();
bld.arg(&grad_ptr);
bld.arg(&flag_ptr);
bld.arg(&unscale_ptr);
bld.arg(&n);
unsafe { bld.launch(cfg) }.map_err(|e| format!("scale_grads_skip_f32: {e:?}"))?;
Ok(())
}
pub fn scale_grads_gpu(
ctx: &GpuCtx,
kernels: &MambaKernels,
grads: &mut GpuBuffer,
scale: f32,
) -> Result<(), String> {
if grads.is_empty() {
return Ok(());
}
let n = grads.len() as i32;
let cfg = grid_1d(grads.len());
let mut bld = ctx.stream.launch_builder(&kernels.scale_grads_f32);
let grad_ptr = grads.cached_ptr();
bld.arg(&grad_ptr);
bld.arg(&scale);
bld.arg(&n);
unsafe { bld.launch(cfg) }.map_err(|e| format!("scale_grads_f32: {:?}", e))?;
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn cpu_state_machine_clean_growth() {
let mut s = DynamicLossScaler::new()
.with_init_scale(8.0)
.with_growth_interval(3);
assert_eq!(s.scale(), 8.0);
s.update(false); s.update(false); assert_eq!(s.scale(), 8.0);
s.update(false); assert_eq!(s.scale(), 16.0);
assert_eq!(s.clean_step_count(), 0);
}
#[test]
fn cpu_state_machine_overflow_backoff() {
let mut s = DynamicLossScaler::new().with_init_scale(8.0);
s.update(false);
s.update(false);
s.update(true); assert_eq!(s.scale(), 4.0);
assert_eq!(s.clean_step_count(), 0);
}
#[test]
fn cpu_state_machine_min_scale_floor() {
let mut s = DynamicLossScaler::new().with_init_scale(2.0);
s.update(true);
s.update(true);
s.update(true);
assert_eq!(s.scale(), 1.0);
}
#[test]
fn cpu_state_machine_max_scale_cap() {
let mut s = DynamicLossScaler::new()
.with_init_scale(8.0)
.with_growth_interval(1)
.with_max_scale(16.0);
s.update(false); assert_eq!(s.scale(), 16.0);
s.update(false); assert_eq!(s.scale(), 16.0);
}
#[test]
fn disabled_scaler_is_identity() {
let mut s = DynamicLossScaler::disabled();
assert_eq!(s.scale(), 1.0);
assert!(!s.enabled());
s.update(true);
assert_eq!(s.scale(), 1.0);
}
#[test]
fn state_dict_round_trip() {
let mut s = DynamicLossScaler::new()
.with_init_scale(8.0)
.with_growth_interval(10);
s.update(false);
s.update(false);
s.update(false);
let (saved_scale, saved_tracker) = s.state();
assert_eq!(saved_scale, 8.0);
assert_eq!(saved_tracker, 3);
let mut s2 = DynamicLossScaler::new().with_growth_interval(10);
s2.load_state(saved_scale, saved_tracker);
assert_eq!(s2.scale(), 8.0);
assert_eq!(s2.clean_step_count(), 3);
}
#[test]
fn load_state_clamps_to_bounds() {
let mut s = DynamicLossScaler::new()
.with_init_scale(8.0)
.with_max_scale(100.0);
s.load_state(1e9, 999_999); assert_eq!(s.scale(), 100.0); assert!(s.clean_step_count() < s.growth_interval);
}
#[test]
#[should_panic(expected = "init_scale must be finite and positive")]
fn init_scale_zero_panics() {
let _ = DynamicLossScaler::new().with_init_scale(0.0);
}
#[test]
#[should_panic(expected = "init_scale must be finite and positive")]
fn init_scale_neg_panics() {
let _ = DynamicLossScaler::new().with_init_scale(-1.0);
}
#[test]
#[should_panic(expected = "init_scale must be finite and positive")]
fn init_scale_inf_panics() {
let _ = DynamicLossScaler::new().with_init_scale(f32::INFINITY);
}
#[test]
#[should_panic(expected = "init_scale must be finite and positive")]
fn init_scale_nan_panics() {
let _ = DynamicLossScaler::new().with_init_scale(f32::NAN);
}
#[test]
#[should_panic(expected = "growth_interval must be > 0")]
fn growth_interval_zero_panics() {
let _ = DynamicLossScaler::new().with_growth_interval(0);
}
#[test]
#[should_panic(expected = "max_scale")]
fn max_below_min_panics() {
let _ = DynamicLossScaler::new().with_max_scale(0.5);
}
#[test]
#[should_panic(expected = "min_scale")]
fn min_above_max_panics() {
let _ = DynamicLossScaler::new().with_min_scale(1e10);
}
#[test]
fn min_scale_pytorch_equivalent() {
let mut s = DynamicLossScaler::new()
.with_init_scale(2.0)
.with_min_scale(f32::MIN_POSITIVE);
s.update(true); s.update(true); s.update(true); assert!((s.scale() - 0.25).abs() < 1e-6);
}
}