use std::sync::Arc;
use cudarc::driver::{LaunchConfig, PushKernelArg};
use super::buffers::{GpuBuffer, GpuByteBuffer};
use super::context::GpuCtx;
use super::launch::grid_1d;
pub const GRAD_CLIP_PARTIALS: usize = 512;
const GRAD_CLIP_THREADS: u32 = 256;
pub fn alloc_partials(stream: &Arc<cudarc::driver::CudaStream>) -> Result<GpuByteBuffer, String> {
GpuByteBuffer::zeros(stream, GRAD_CLIP_PARTIALS * std::mem::size_of::<f64>())
}
pub fn global_grad_norm(
ctx: &GpuCtx,
grads_flat: &GpuBuffer,
partials: &mut GpuByteBuffer,
partials_host: &mut [f64],
) -> Result<f64, String> {
assert_eq!(
partials.len_bytes(),
GRAD_CLIP_PARTIALS * std::mem::size_of::<f64>(),
"grad-clip partials buffer has the wrong size"
);
assert_eq!(
partials_host.len(),
GRAD_CLIP_PARTIALS,
"grad-clip host partials slice has the wrong length"
);
let n = grads_flat.len() as i32;
let dst = partials.cached_ptr();
let src = grads_flat.cached_ptr();
let mut b = ctx
.stream
.launch_builder(&ctx.kernels.grad_sumsq_partial_f32);
b.arg(&dst);
b.arg(&src);
b.arg(&n);
let launch_cfg = LaunchConfig {
grid_dim: (GRAD_CLIP_PARTIALS as u32, 1, 1),
block_dim: (GRAD_CLIP_THREADS, 1, 1),
shared_mem_bytes: 0,
};
unsafe { b.launch(launch_cfg) }.map_err(|e| format!("grad_sumsq_partial_f32: {e:?}"))?;
ctx.stream
.synchronize()
.map_err(|e| format!("grad norm sync: {e:?}"))?;
partials.download_f64(&ctx.stream, partials_host)?;
let mut sum = 0.0f64;
for &p in partials_host.iter() {
sum += p;
}
Ok(sum.sqrt())
}
pub fn clip_grads_device(
ctx: &GpuCtx,
grads_flat: &mut GpuBuffer,
partials: &mut GpuByteBuffer,
scratch: &mut GpuBuffer,
max_norm: f32,
) -> Result<f32, String> {
assert_eq!(
partials.len_bytes(),
GRAD_CLIP_PARTIALS * std::mem::size_of::<f64>(),
"grad-clip partials buffer has the wrong size"
);
assert_eq!(
scratch.len(),
2,
"grad-clip scratch must hold exactly [coef, norm]"
);
let n_elems = grads_flat.len();
let n = n_elems as i32;
{
let dst = partials.cached_ptr();
let src = grads_flat.cached_ptr();
let mut b = ctx
.stream
.launch_builder(&ctx.kernels.grad_sumsq_partial_f32);
b.arg(&dst);
b.arg(&src);
b.arg(&n);
let cfg = LaunchConfig {
grid_dim: (GRAD_CLIP_PARTIALS as u32, 1, 1),
block_dim: (GRAD_CLIP_THREADS, 1, 1),
shared_mem_bytes: 0,
};
unsafe { b.launch(cfg) }.map_err(|e| format!("grad_sumsq_partial_f32: {e:?}"))?;
}
{
let coef_ptr = scratch.cached_ptr();
let norm_ptr = coef_ptr + std::mem::size_of::<f32>() as u64;
let part_ptr = partials.cached_ptr();
let np = GRAD_CLIP_PARTIALS as i32;
let mut b = ctx.stream.launch_builder(&ctx.kernels.grad_clip_coef_f32);
b.arg(&coef_ptr);
b.arg(&norm_ptr);
b.arg(&part_ptr);
b.arg(&np);
b.arg(&max_norm);
let cfg = LaunchConfig {
grid_dim: (1, 1, 1),
block_dim: (1, 1, 1),
shared_mem_bytes: 0,
};
unsafe { b.launch(cfg) }.map_err(|e| format!("grad_clip_coef_f32: {e:?}"))?;
}
{
let coef_ptr = scratch.cached_ptr();
let mut b = ctx.stream.launch_builder(&ctx.kernels.scale_grads_dev_f32);
b.arg(grads_flat.inner_mut());
b.arg(&coef_ptr);
b.arg(&n);
unsafe { b.launch(grid_1d(n_elems)) }.map_err(|e| format!("scale_grads_dev_f32: {e:?}"))?;
}
let mut host = [0.0f32; 2];
scratch.download(&ctx.stream, &mut host)?;
Ok(host[1])
}
pub fn scale_grads(ctx: &GpuCtx, grads_flat: &mut GpuBuffer, factor: f32) -> Result<(), String> {
let n_elems = grads_flat.len();
let n = n_elems as i32;
let mut b = ctx.stream.launch_builder(&ctx.kernels.scale_grads_f32);
b.arg(grads_flat.inner_mut());
b.arg(&factor);
b.arg(&n);
unsafe { b.launch(grid_1d(n_elems)) }
.map(|_| ())
.map_err(|e| format!("scale_grads (clip): {e:?}"))
}