diff --git a/crates/ml-ppo/src/cuda_nn/activations.rs b/crates/ml-ppo/src/cuda_nn/activations.rs index f88c3fcdb..a9356f861 100644 --- a/crates/ml-ppo/src/cuda_nn/activations.rs +++ b/crates/ml-ppo/src/cuda_nn/activations.rs @@ -1,142 +1,26 @@ -//! CUDA activation function kernels (`ReLU`, tanh, sigmoid). +//! CUDA activation function wrappers. //! -//! All activations operate in-place or to a pre-allocated output buffer, -//! avoiding Candle tensor overhead. +//! Thin wrappers around `ml_core::cuda_autograd::ActivationKernels` that +//! present the original `CudaVec`-returning API. The underlying CUDA kernels +//! are compiled and executed by ml-core -- no duplicate PTX source here. use std::sync::Arc; -use cudarc; -use cudarc::driver::{CudaSlice, CudaStream, LaunchConfig, PushKernelArg}; -use cudarc::nvrtc::Ptx; +use cudarc::driver::{CudaSlice, CudaStream}; +use ml_core::cuda_autograd::ActivationKernels; use ml_core::MLError; use super::tensor_util::CudaVec; -const ACTIVATION_KERNELS: &str = r#" -extern "C" __global__ void relu_forward( - float* __restrict__ output, - const float* __restrict__ input, - int n -) { - int idx = blockIdx.x * blockDim.x + threadIdx.x; - if (idx < n) { - output[idx] = fmaxf(input[idx], 0.0f); - } -} - -extern "C" __global__ void relu_backward( - float* __restrict__ grad_input, - const float* __restrict__ grad_output, - const float* __restrict__ input, - int n -) { - int idx = blockIdx.x * blockDim.x + threadIdx.x; - if (idx < n) { - grad_input[idx] = input[idx] > 0.0f ? grad_output[idx] : 0.0f; - } -} - -extern "C" __global__ void tanh_forward( - float* __restrict__ output, - const float* __restrict__ input, - int n -) { - int idx = blockIdx.x * blockDim.x + threadIdx.x; - if (idx < n) { - output[idx] = tanhf(input[idx]); - } -} - -extern "C" __global__ void sigmoid_forward( - float* __restrict__ output, - const float* __restrict__ input, - int n -) { - int idx = blockIdx.x * blockDim.x + threadIdx.x; - if (idx < n) { - output[idx] = 1.0f / (1.0f + expf(-input[idx])); - } -} -"#; - -/// Lazily compiled activation kernel handles. -struct ActivationKernels { - relu_fwd: cudarc::driver::CudaFunction, - relu_bwd: cudarc::driver::CudaFunction, - tanh_fwd: cudarc::driver::CudaFunction, - sigmoid_fwd: cudarc::driver::CudaFunction, -} - -fn compile_activation_kernels(stream: &Arc) -> Result { - let context = stream.context(); - let ptx: Ptx = cudarc::nvrtc::compile_ptx(ACTIVATION_KERNELS).map_err(|e| { - MLError::InitializationError { - component: "activations".to_owned(), - message: format!("Failed to compile activation kernels: {e}"), - } - })?; - let module = context.load_module(ptx).map_err(|e| { - MLError::InitializationError { - component: "activations".to_owned(), - message: format!("Failed to load activation module: {e}"), - } - })?; - Ok(ActivationKernels { - relu_fwd: module.load_function("relu_forward").map_err(|e| { - MLError::InitializationError { - component: "activations".to_owned(), - message: format!("Failed to load relu_forward: {e}"), - } - })?, - relu_bwd: module.load_function("relu_backward").map_err(|e| { - MLError::InitializationError { - component: "activations".to_owned(), - message: format!("Failed to load relu_backward: {e}"), - } - })?, - tanh_fwd: module.load_function("tanh_forward").map_err(|e| { - MLError::InitializationError { - component: "activations".to_owned(), - message: format!("Failed to load tanh_forward: {e}"), - } - })?, - sigmoid_fwd: module.load_function("sigmoid_forward").map_err(|e| { - MLError::InitializationError { - component: "activations".to_owned(), - message: format!("Failed to load sigmoid_forward: {e}"), - } - })?, - }) -} - -const fn launch_config(n: usize) -> LaunchConfig { - let block_size = 256_u32; - let grid_size = (n as u32).div_ceil(block_size); - LaunchConfig { - grid_dim: (grid_size, 1, 1), - block_dim: (block_size, 1, 1), - shared_mem_bytes: 0, - } -} - /// Apply `ReLU` activation: output = max(0, input) -pub fn cuda_relu(stream: &Arc, input: &CudaSlice, len: usize) -> Result { - let kernels = compile_activation_kernels(stream)?; - let mut output = stream.alloc_zeros::(len).map_err(|e| { - MLError::ModelError(format!("Failed to alloc relu output: {e}")) - })?; - let n = len as i32; - // SAFETY: relu_forward reads `len` f32 elements from input and writes - // `len` elements to output. Both are pre-allocated with matching sizes. - unsafe { - stream.launch_builder(&kernels.relu_fwd) - .arg(&mut output) - .arg(input) - .arg(&n) - .launch(launch_config(len)) - .map_err(|e| MLError::ModelError(format!("relu_forward launch failed: {e}")))?; - } +pub fn cuda_relu( + stream: &Arc, + input: &CudaSlice, + len: usize, +) -> Result { + let kernels = ActivationKernels::new(stream)?; + let output = kernels.relu_fwd_raw(input, len, stream)?; Ok(CudaVec::new(output, len)) } @@ -147,61 +31,29 @@ pub fn cuda_relu_backward( input: &CudaSlice, len: usize, ) -> Result { - let kernels = compile_activation_kernels(stream)?; - let mut grad_input = stream.alloc_zeros::(len).map_err(|e| { - MLError::ModelError(format!("Failed to alloc relu_backward output: {e}")) - })?; - let n = len as i32; - // SAFETY: relu_backward reads `len` elements from grad_output and input, - // writes `len` elements to grad_input. All buffers are pre-allocated. - unsafe { - stream.launch_builder(&kernels.relu_bwd) - .arg(&mut grad_input) - .arg(grad_output) - .arg(input) - .arg(&n) - .launch(launch_config(len)) - .map_err(|e| MLError::ModelError(format!("relu_backward launch failed: {e}")))?; - } - Ok(CudaVec::new(grad_input, len)) + let kernels = ActivationKernels::new(stream)?; + let dx = kernels.relu_bwd_raw(grad_output, input, len, stream)?; + Ok(CudaVec::new(dx, len)) } /// Apply tanh activation: output = tanh(input) -pub fn cuda_tanh(stream: &Arc, input: &CudaSlice, len: usize) -> Result { - let kernels = compile_activation_kernels(stream)?; - let mut output = stream.alloc_zeros::(len).map_err(|e| { - MLError::ModelError(format!("Failed to alloc tanh output: {e}")) - })?; - let n = len as i32; - // SAFETY: tanh_forward reads `len` f32 elements from input and writes - // `len` elements to output. Both buffers are pre-allocated with matching sizes. - unsafe { - stream.launch_builder(&kernels.tanh_fwd) - .arg(&mut output) - .arg(input) - .arg(&n) - .launch(launch_config(len)) - .map_err(|e| MLError::ModelError(format!("tanh_forward launch failed: {e}")))?; - } +pub fn cuda_tanh( + stream: &Arc, + input: &CudaSlice, + len: usize, +) -> Result { + let kernels = ActivationKernels::new(stream)?; + let output = kernels.tanh_fwd_raw(input, len, stream)?; Ok(CudaVec::new(output, len)) } /// Apply sigmoid activation: output = 1 / (1 + exp(-input)) -pub fn cuda_sigmoid(stream: &Arc, input: &CudaSlice, len: usize) -> Result { - let kernels = compile_activation_kernels(stream)?; - let mut output = stream.alloc_zeros::(len).map_err(|e| { - MLError::ModelError(format!("Failed to alloc sigmoid output: {e}")) - })?; - let n = len as i32; - // SAFETY: sigmoid_forward reads `len` f32 elements from input and writes - // `len` elements to output. Both buffers are pre-allocated with matching sizes. - unsafe { - stream.launch_builder(&kernels.sigmoid_fwd) - .arg(&mut output) - .arg(input) - .arg(&n) - .launch(launch_config(len)) - .map_err(|e| MLError::ModelError(format!("sigmoid_forward launch failed: {e}")))?; - } +pub fn cuda_sigmoid( + stream: &Arc, + input: &CudaSlice, + len: usize, +) -> Result { + let kernels = ActivationKernels::new(stream)?; + let output = kernels.sigmoid_fwd_raw(input, len, stream)?; Ok(CudaVec::new(output, len)) } diff --git a/crates/ml-ppo/src/cuda_nn/adam.rs b/crates/ml-ppo/src/cuda_nn/adam.rs index 7e7b63b4c..13cc7b406 100644 --- a/crates/ml-ppo/src/cuda_nn/adam.rs +++ b/crates/ml-ppo/src/cuda_nn/adam.rs @@ -1,72 +1,65 @@ -//! GPU-resident Adam optimizer using CUDA kernels. +//! GPU-resident AdamW optimizer using CUDA kernels. //! +//! Flat-param-group API over `ml_core::cuda_autograd`'s AdamW CUDA kernel. //! All optimizer state (m, v, parameters) stays on GPU. The update step -//! runs as a single CUDA kernel launch per parameter group — zero CPU +//! runs as a single CUDA kernel launch per parameter group -- zero CPU //! round-trips per optimization step. +//! +//! This module provides the `CudaAdam` struct with indexed parameter groups, +//! suitable for networks that manage weights as raw `CudaSlice` buffers +//! rather than named parameters in a `GpuVarStore`. use std::sync::Arc; -use cudarc; use cudarc::driver::{CudaSlice, CudaStream, LaunchConfig, PushKernelArg}; -use cudarc::nvrtc::Ptx; use ml_core::MLError; -const ADAM_KERNEL: &str = r#" -extern "C" __global__ void adam_step( - float* __restrict__ param, - const float* __restrict__ grad, - float* __restrict__ m, // first moment - float* __restrict__ v, // second moment - float lr, - float beta1, - float beta2, - float eps, - float weight_decay, - float bias_correction1, // 1 - beta1^t - float bias_correction2, // 1 - beta2^t - int n -) { - int idx = blockIdx.x * blockDim.x + threadIdx.x; - if (idx < n) { - float g = grad[idx]; +/// CUDA kernel source for the AdamW update -- identical to ml-core's +/// `adamw_update` kernel (decoupled weight decay, bias-corrected moments). +/// Kept as a string constant so compilation is self-contained per optimizer +/// instance. +const ADAMW_KERNEL: &str = r#" +extern "C" __global__ +void adamw_update(float* __restrict__ param, + const float* __restrict__ grad, + float* __restrict__ m, + float* __restrict__ v, + float lr, + float beta1, + float beta2, + float epsilon, + float weight_decay, + int t, + int n) { + int i = blockIdx.x * blockDim.x + threadIdx.x; + if (i < n) { + float g = grad[i]; + float p = param[i]; - // Weight decay (decoupled, AdamW-style) - if (weight_decay > 0.0f) { - g += weight_decay * param[idx]; - } + // Decoupled weight decay: p = p * (1 - lr * wd) + p *= (1.0f - lr * weight_decay); - // Update biased first moment estimate - float m_new = beta1 * m[idx] + (1.0f - beta1) * g; - m[idx] = m_new; + // Moment updates + float mi = beta1 * m[i] + (1.0f - beta1) * g; + float vi = beta2 * v[i] + (1.0f - beta2) * g * g; + m[i] = mi; + v[i] = vi; - // Update biased second raw moment estimate - float v_new = beta2 * v[idx] + (1.0f - beta2) * g * g; - v[idx] = v_new; + // Bias correction + float bc1 = 1.0f - powf(beta1, (float)t); + float bc2 = 1.0f - powf(beta2, (float)t); + float m_hat = mi / bc1; + float v_hat = vi / bc2; - // Compute bias-corrected estimates - float m_hat = m_new / bias_correction1; - float v_hat = v_new / bias_correction2; - - // Update parameter - param[idx] -= lr * m_hat / (sqrtf(v_hat) + eps); - } -} - -extern "C" __global__ void grad_clip_and_norm( - float* __restrict__ grad, - float max_norm, - float scale, // If total_norm > max_norm, scale = max_norm / total_norm, else 1.0 - int n -) { - int idx = blockIdx.x * blockDim.x + threadIdx.x; - if (idx < n) { - grad[idx] *= scale; + // Parameter update + p -= lr * m_hat / (sqrtf(v_hat) + epsilon); + param[i] = p; } } "#; -/// Configuration for the CUDA Adam optimizer. +/// Configuration for the CUDA AdamW optimizer. #[derive(Debug, Clone)] pub struct CudaAdamConfig { pub lr: f32, @@ -88,7 +81,7 @@ impl Default for CudaAdamConfig { } } -/// A single parameter group managed by the CUDA Adam optimizer. +/// A single parameter group managed by the CUDA AdamW optimizer. pub struct AdamParamGroup { /// First moment (m) buffer, same size as parameter pub m: CudaSlice, @@ -98,10 +91,10 @@ pub struct AdamParamGroup { pub len: usize, } -/// GPU-resident Adam optimizer. +/// GPU-resident AdamW optimizer with flat indexed parameter groups. /// /// Manages first and second moment buffers for each parameter group. -/// The `step()` method launches one CUDA kernel per parameter to update +/// The `step_group()` method launches one CUDA kernel per parameter to update /// weights in-place using the accumulated gradients. pub struct CudaAdam { config: CudaAdamConfig, @@ -124,27 +117,26 @@ impl std::fmt::Debug for CudaAdam { } impl CudaAdam { - /// Create a new Adam optimizer. + /// Create a new AdamW optimizer. /// - /// Call `add_param_group()` for each (weight, bias) pair before calling `step()`. + /// Call `add_param_group()` for each (weight, bias) pair before calling `step_group()`. pub fn new(stream: Arc, config: CudaAdamConfig) -> Result { let context = stream.context(); - let ptx: Ptx = cudarc::nvrtc::compile_ptx(ADAM_KERNEL).map_err(|e| { - MLError::InitializationError { + let ptx = ml_core::cuda_compile::compile_ptx_for_device(ADAMW_KERNEL, &context) + .map_err(|e| MLError::InitializationError { component: "CudaAdam".to_owned(), - message: format!("Failed to compile adam kernel: {e}"), - } - })?; + message: format!("Failed to compile adamw kernel: {e}"), + })?; let module = context.load_module(ptx).map_err(|e| { MLError::InitializationError { component: "CudaAdam".to_owned(), - message: format!("Failed to load adam module: {e}"), + message: format!("Failed to load adamw module: {e}"), } })?; - let adam_func = module.load_function("adam_step").map_err(|e| { + let adam_func = module.load_function("adamw_update").map_err(|e| { MLError::InitializationError { component: "CudaAdam".to_owned(), - message: format!("Failed to load adam_step: {e}"), + message: format!("Failed to load adamw_update: {e}"), } })?; @@ -178,7 +170,7 @@ impl CudaAdam { Ok(idx) } - /// Perform one Adam optimization step for a specific parameter group. + /// Perform one AdamW optimization step for a specific parameter group. /// /// Updates `param` in-place using `grad`. pub fn step_group( @@ -191,17 +183,15 @@ impl CudaAdam { MLError::ConfigError(format!("Adam group index {} out of range", group_idx)) })?; - let t = self.step_count + 1; - let bias_correction1 = 1.0 - self.config.beta1.powi(t as i32); - let bias_correction2 = 1.0 - self.config.beta2.powi(t as i32); + let t = (self.step_count + 1) as i32; let n = group.len as i32; let block_size = 256_u32; let grid_size = (group.len as u32).div_ceil(block_size); - // SAFETY: adam_step kernel reads `n` gradient elements and updates `n` + // SAFETY: adamw_update kernel reads `n` gradient elements and updates `n` // parameter, momentum (m), and velocity (v) elements in-place. All - // buffers are registered with matching sizes during `register_params`. + // buffers are registered with matching sizes during `add_param_group`. unsafe { self.stream .launch_builder(&self.adam_func) @@ -214,15 +204,14 @@ impl CudaAdam { .arg(&self.config.beta2) .arg(&self.config.eps) .arg(&self.config.weight_decay) - .arg(&bias_correction1) - .arg(&bias_correction2) + .arg(&t) .arg(&n) .launch(LaunchConfig { grid_dim: (grid_size, 1, 1), block_dim: (block_size, 1, 1), shared_mem_bytes: 0, }) - .map_err(|e| MLError::TrainingError(format!("adam_step launch failed: {e}")))?; + .map_err(|e| MLError::TrainingError(format!("adamw_update launch failed: {e}")))?; } Ok(())