diff --git a/crates/ml/src/cuda_pipeline/curiosity_training_kernel.cu b/crates/ml/src/cuda_pipeline/curiosity_training_kernel.cu index 6a6661511..c51f20647 100644 --- a/crates/ml/src/cuda_pipeline/curiosity_training_kernel.cu +++ b/crates/ml/src/cuda_pipeline/curiosity_training_kernel.cu @@ -1,16 +1,31 @@ /** - * GPU-resident curiosity forward model training kernel. + * GPU-resident curiosity forward model training kernels. * - * Trains the curiosity forward model (2-layer MLP) entirely on GPU with zero - * CPU involvement. Experience data (states, actions, next_states) is already - * on GPU from the experience collector. Weights are updated in-place via Adam. + * cuBLAS GEMM pipeline (replaces the 1085ms curiosity_fwd_bwd_per_block kernel): + * + * Forward: + * 1. curiosity_shift_states — build next_states from states[+1] + * 2. curiosity_prepare_input — build input[N, CUR_INPUT] (from inference cubin) + * 3. cuBLAS GEMM1 — hidden[N, CUR_HIDDEN] = input @ W1^T + * 4. curiosity_bias_leaky_relu — in-place +b1, LeakyReLU (from inference cubin) + * 5. cuBLAS GEMM2 — pred[N, CUR_OUTPUT] = hidden @ W2^T + * 6. curiosity_mse_fwd_grad — +b2, compute d_pred = 2/N*(pred+b2-target) + * + * Backward: + * 7. cuBLAS dW2 — grad_w2 += d_pred^T @ hidden + * 8. curiosity_bias_grad_reduce — grad_b2 = sum_batch(d_pred) + * 9. cuBLAS d_hidden — d_hidden = d_pred @ W2 + * 10. curiosity_leaky_relu_bwd — mask d_hidden by post-activation sign + * 11. cuBLAS dW1 — grad_w1 += d_hidden^T @ input + * 12. curiosity_bias_grad_reduce — grad_b1 = sum_batch(d_hidden) + * 13. curiosity_adam_step — Adam update per parameter group * * Architecture: [CUR_INPUT=MARKET_DIM+3] -> [CUR_HIDDEN=128] LeakyReLU(0.01) -> [CUR_OUTPUT=MARKET_DIM] - * Total params: 128*(MARKET_DIM+3) + 128 + MARKET_DIM*128 + MARKET_DIM + * Total params: 128*(MARKET_DIM+3) + 128 + MARKET_DIM*128 + MARKET_DIM = 11306 * * Requires common_device_functions.cuh to be prepended for CUR_INPUT/CUR_HIDDEN/CUR_OUTPUT. * - * Native float everywhere. Fully deterministic (no atomicAdd). + * Native float everywhere. */ /* Total trainable parameters */ @@ -327,3 +342,97 @@ extern "C" __global__ void curiosity_adam_step( float v_hat = v[i] / (1.0f - powf(beta2, (float)step)); params[i] = params[i] - lr_bf * m_hat / (sqrtf(v_hat) + eps_bf); } + +/* ------------------------------------------------------------------ */ +/* Kernel 5: MSE forward + d_pred computation (cuBLAS training path) */ +/* ------------------------------------------------------------------ */ + +/** + * One thread per element in pred[N, CUR_OUTPUT]. + * + * 1. Adds b2 in-place to pred (broadcast along batch). + * 2. Computes d_pred[n, o] = (2.0 / CUR_OUTPUT) * (pred[n,o] - target[n,o]). + * The 1/N mean is applied later by curiosity_adam_step (divides by batch_size). + * + * cuBLAS GEMM 2 writes pred[N, CUR_OUTPUT] without bias. We add bias here + * and compute the MSE gradient for the backward pass. + * + * Launch: grid=(ceil(N*CUR_OUTPUT/256)), block=(256) + */ +extern "C" __global__ void curiosity_mse_fwd_grad( + float* __restrict__ pred, /* [N, CUR_OUTPUT] — in-place +b2, then d_pred out */ + const float* __restrict__ b2, /* [CUR_OUTPUT] */ + const float* __restrict__ target, /* [N, state_dim] (uses first CUR_OUTPUT elements) */ + int N, + int state_dim +) { + int total = N * CUR_OUTPUT; + int i = blockIdx.x * blockDim.x + threadIdx.x; + if (i >= total) return; + + int n = i / CUR_OUTPUT; /* sample index */ + int o = i % CUR_OUTPUT; /* output index */ + + float p = pred[i] + b2[o]; + float t = target[n * state_dim + o]; + float d = (2.0f / (float)CUR_OUTPUT) * (p - t); + + /* Write back d_pred; pred buffer is reused as gradient buffer */ + pred[i] = d; +} + +/* ------------------------------------------------------------------ */ +/* Kernel 6: LeakyReLU backward (in-place on d_hidden) */ +/* ------------------------------------------------------------------ */ + +/** + * One thread per element in d_hidden[N, CUR_HIDDEN]. + * + * Gates d_hidden by the sign of the post-activation hidden values: + * d_hidden[i] *= (hidden_post[i] > 0) ? 1.0 : 0.01 + * + * This is valid because LeakyReLU preserves the sign of the input: + * f(x) > 0 iff x > 0, so the mask is equivalent to (pre_act > 0). + * + * Launch: grid=(ceil(N*CUR_HIDDEN/256)), block=(256) + */ +extern "C" __global__ void curiosity_leaky_relu_bwd( + float* __restrict__ d_hidden, /* [N, CUR_HIDDEN] — modified in-place */ + const float* __restrict__ hidden_post, /* [N, CUR_HIDDEN] — post-activation */ + int total_elements /* N * CUR_HIDDEN */ +) { + int i = blockIdx.x * blockDim.x + threadIdx.x; + if (i >= total_elements) return; + float mask = (hidden_post[i] > 0.0f) ? 1.0f : 0.01f; + d_hidden[i] *= mask; +} + +/* ------------------------------------------------------------------ */ +/* Kernel 7: Bias gradient reduction (sum over batch) */ +/* ------------------------------------------------------------------ */ + +/** + * One thread per output dimension d. + * + * Computes grad_b[d] = sum_{n=0}^{N-1} dy[n, d] + * + * Used for both grad_b1 (d=CUR_HIDDEN, dy=d_hidden) and + * grad_b2 (d=CUR_OUTPUT, dy=d_pred). + * + * Launch: grid=(ceil(out_dim/256)), block=(256) + */ +extern "C" __global__ void curiosity_bias_grad_reduce( + const float* __restrict__ dy, /* [N, out_dim] */ + float* __restrict__ grad_b, /* [out_dim] */ + int N, + int out_dim +) { + int d = blockIdx.x * blockDim.x + threadIdx.x; + if (d >= out_dim) return; + + float sum = 0.0f; + for (int n = 0; n < N; n++) { + sum += dy[n * out_dim + d]; + } + grad_b[d] = sum; +} diff --git a/crates/ml/src/cuda_pipeline/gpu_curiosity_trainer.rs b/crates/ml/src/cuda_pipeline/gpu_curiosity_trainer.rs index 57084ce56..d22c3c673 100644 --- a/crates/ml/src/cuda_pipeline/gpu_curiosity_trainer.rs +++ b/crates/ml/src/cuda_pipeline/gpu_curiosity_trainer.rs @@ -1,31 +1,52 @@ #![allow(unsafe_code)] // Required for CUDA kernel launches -//! GPU-resident curiosity forward model training. +//! GPU-resident curiosity forward model training — cuBLAS GEMM pipeline. //! //! Trains the curiosity forward model (2-layer MLP) entirely on GPU with zero //! CPU involvement. Experience data (states, actions, next_states) is already //! on GPU from the experience collector. Weights in [`CuriosityWeightSet`] are //! updated in-place via Adam optimizer -- no CPU roundtrips. //! -//! Architecture: `[MARKET_DIM+3] -> [128] LeakyReLU -> [MARKET_DIM]` (MARKET_DIM=42: 11_954 params) +//! Architecture: `[MARKET_DIM+3] -> [128] LeakyReLU -> [MARKET_DIM]` (MARKET_DIM=42: 11_306 params) //! -//! Kernels (deterministic path -- 4 launches per step): -//! - `curiosity_shift_states`: builds shifted next_states from states buffer -//! - `curiosity_fwd_bwd_per_block`: forward + backward, block-level reduce to partial gradients -//! - `curiosity_grad_reduce`: deterministic reduction of per-block partials into final gradients -//! - `curiosity_adam_step`: per-param-group Adam optimizer update +//! Pipeline (13 GPU ops per step, replaces the 1085ms curiosity_fwd_bwd_per_block): +//! +//! Forward: +//! 1. `curiosity_shift_states` — build next_states[N, state_dim] from states[+1] +//! 2. `curiosity_prepare_input` — input[N, CUR_INPUT] from states + action one-hot +//! 3. cuBLAS GEMM1 — hidden[N, CUR_HIDDEN] = input @ W1^T +//! 4. `curiosity_bias_leaky_relu` — in-place +b1, LeakyReLU(0.01) +//! 5. cuBLAS GEMM2 — pred[N, CUR_OUTPUT] = hidden @ W2^T +//! 6. `curiosity_mse_fwd_grad` — +b2 in-place, d_pred = 2/CUR_OUTPUT*(pred-target) +//! +//! Backward: +//! 7. cuBLAS dW2 — grad_w2 = d_pred^T @ hidden (GEMM TRANSA=T, TRANSB=N) +//! 8. `curiosity_bias_grad_reduce`— grad_b2 = sum_batch(d_pred) +//! 9. cuBLAS d_hidden — d_hidden = d_pred @ W2 (GEMM TRANSA=N, TRANSB=N) +//! 10. `curiosity_leaky_relu_bwd` — d_hidden *= mask(hidden_post) +//! 11. cuBLAS dW1 — grad_w1 = d_hidden^T @ input (GEMM TRANSA=T, TRANSB=N) +//! 12. `curiosity_bias_grad_reduce`— grad_b1 = sum_batch(d_hidden) +//! 13. `curiosity_adam_step` ×4 — Adam update per parameter group use std::sync::Arc; -use cudarc::driver::{CudaFunction, CudaSlice, CudaStream, LaunchConfig, PushKernelArg}; +use cudarc::cublas::result as cublas_result; +use cudarc::cublas::sys as cublas_sys; +use cudarc::cublaslt::result as cublaslt_result; +use cudarc::cublaslt::sys as cublaslt_sys; +use cudarc::driver::{CudaFunction, CudaSlice, CudaStream, DevicePtr, LaunchConfig, PushKernelArg}; use tracing::debug; use crate::MLError; use super::gpu_weights::CuriosityWeightSet; +use super::shared_cublas_handle::{SendSyncCublasHandle, SendSyncCublasLtHandle}; /// Precompiled curiosity training kernel cubin, embedded at compile time by build.rs. static CURIOSITY_TRAINING_CUBIN: &[u8] = include_bytes!(concat!(env!("OUT_DIR"), "/curiosity_training_kernel.cubin")); +/// Precompiled curiosity inference kernel cubin (provides prepare_input + bias_leaky_relu). +static CURIOSITY_INFERENCE_CUBIN: &[u8] = include_bytes!(concat!(env!("OUT_DIR"), "/curiosity_inference_kernel.cubin")); + // --------------------------------------------------------------------------- // Constants (must match CUDA kernel defines) // --------------------------------------------------------------------------- @@ -40,49 +61,421 @@ const CUR_B1_LEN: usize = CUR_HIDDEN; // [128] const CUR_W2_LEN: usize = CUR_OUTPUT * CUR_HIDDEN; // [42, 128] = 5376 const CUR_B2_LEN: usize = CUR_OUTPUT; // [42] -const CUR_TOTAL_PARAMS: usize = CUR_W1_LEN + CUR_B1_LEN + CUR_W2_LEN + CUR_B2_LEN; // 11306 - -/// Block size for the forward+backward kernel. -const FWD_BWD_BLOCK_SIZE: u32 = 256; - /// Adam optimizer hyperparameters. const ADAM_LR: f32 = 0.001; const ADAM_BETA1: f32 = 0.9; const ADAM_BETA2: f32 = 0.999; const ADAM_EPS: f32 = 1e-8; +// --------------------------------------------------------------------------- +// Minimal cuBLAS context for the curiosity trainer +// --------------------------------------------------------------------------- + +/// Minimal cuBLAS + cublasLt context owned by the curiosity trainer. +/// +/// The curiosity trainer runs on its own CUDA stream (the experience collector +/// stream), separate from the main DQN trainer. Creating a dedicated handle +/// avoids handle-sharing conflicts during CUDA Graph capture. +#[allow(missing_debug_implementations)] +struct CuriosityGemm { + handle: SendSyncCublasHandle, + lt_handle: SendSyncCublasLtHandle, + _workspace_buf: CudaSlice, + workspace_ptr: u64, + workspace_size: usize, +} + +impl CuriosityGemm { + fn new(stream: &Arc) -> Result { + let raw_handle = cublas_result::create_handle() + .map_err(|e| MLError::ModelError(format!("curiosity cublasCreate: {e:?}")))?; + + unsafe { + let cu_stream = stream.cu_stream() as *mut cublas_sys::CUstream_st; + cublas_result::set_stream(raw_handle, cu_stream) + .map_err(|e| MLError::ModelError(format!("curiosity cublasSetStream: {e:?}")))?; + // TF32 tensor cores for the curiosity MLP — fast, no determinism requirement + // (curiosity is an auxiliary model; ~1e-4 per-GEMM variance is acceptable) + cublas_sys::cublasSetMathMode( + raw_handle, + cublas_sys::cublasMath_t::CUBLAS_TF32_TENSOR_OP_MATH, + ); + } + + // 16 MB workspace: curiosity GEMMs are small (N≤8192, CUR_INPUT=45, CUR_HIDDEN=128) + let workspace_size: usize = 16 * 1024 * 1024; + let workspace_buf = stream + .alloc_zeros::(workspace_size) + .map_err(|e| MLError::ModelError(format!("curiosity cuBLAS workspace alloc: {e}")))?; + let workspace_ptr = { + let (dev_ptr, _guard) = workspace_buf.device_ptr(stream); + dev_ptr + }; + unsafe { + let status = cublas_sys::cublasSetWorkspace_v2( + raw_handle, + workspace_ptr as *mut std::ffi::c_void, + workspace_size, + ); + if status != cublas_sys::cublasStatus_t::CUBLAS_STATUS_SUCCESS { + return Err(MLError::ModelError(format!( + "curiosity cublasSetWorkspace_v2 failed: {status:?}" + ))); + } + } + + let lt_raw_handle = cublaslt_result::create_handle() + .map_err(|e| MLError::ModelError(format!("curiosity cublasLtCreate: {e:?}")))?; + + Ok(Self { + handle: SendSyncCublasHandle(raw_handle), + lt_handle: SendSyncCublasLtHandle(lt_raw_handle), + _workspace_buf: workspace_buf, + workspace_ptr, + workspace_size, + }) + } + + /// cublasLtMatmul: `C = alpha * op(A) @ op(B) + beta * C`. + /// Used only by `gemm_fwd` (forward GEMM with TF32 tensor cores). + #[allow(clippy::too_many_arguments)] + fn sgemm( + &self, + stream: &CudaStream, + transa: cublas_sys::cublasOperation_t, + transb: cublas_sys::cublasOperation_t, + m: i32, n: i32, k: i32, + alpha: f32, + a_ptr: u64, lda: i32, + b_ptr: u64, ldb: i32, + beta: f32, + c_ptr: u64, ldc: i32, + label: &str, + ) -> Result<(), MLError> { + let f32_type = cublaslt_sys::cudaDataType_t::CUDA_R_32F; + let compute_type = cublaslt_sys::cublasComputeType_t::CUBLAS_COMPUTE_32F_FAST_TF32; + + unsafe { + let matmul_desc = cublaslt_result::create_matmul_desc(compute_type, f32_type) + .map_err(|e| MLError::ModelError(format!("cublasLtMatmulDescCreate {label}: {e:?}")))?; + + let transa_i32: i32 = transa as i32; + cublaslt_result::set_matmul_desc_attribute( + matmul_desc, + cublaslt_sys::cublasLtMatmulDescAttributes_t::CUBLASLT_MATMUL_DESC_TRANSA, + &transa_i32 as *const i32 as *const std::ffi::c_void, + std::mem::size_of::(), + ).map_err(|e| MLError::ModelError(format!("set TRANSA {label}: {e:?}")))?; + + let transb_i32: i32 = transb as i32; + cublaslt_result::set_matmul_desc_attribute( + matmul_desc, + cublaslt_sys::cublasLtMatmulDescAttributes_t::CUBLASLT_MATMUL_DESC_TRANSB, + &transb_i32 as *const i32 as *const std::ffi::c_void, + std::mem::size_of::(), + ).map_err(|e| MLError::ModelError(format!("set TRANSB {label}: {e:?}")))?; + + let a_layout = if transa == cublas_sys::cublasOperation_t::CUBLAS_OP_N { + cublaslt_result::create_matrix_layout(f32_type, m as u64, k as u64, lda as i64) + } else { + cublaslt_result::create_matrix_layout(f32_type, k as u64, m as u64, lda as i64) + }.map_err(|e| MLError::ModelError(format!("A layout {label}: {e:?}")))?; + + let b_layout = if transb == cublas_sys::cublasOperation_t::CUBLAS_OP_N { + cublaslt_result::create_matrix_layout(f32_type, k as u64, n as u64, ldb as i64) + } else { + cublaslt_result::create_matrix_layout(f32_type, n as u64, k as u64, ldb as i64) + }.map_err(|e| MLError::ModelError(format!("B layout {label}: {e:?}")))?; + + let c_layout = cublaslt_result::create_matrix_layout(f32_type, m as u64, n as u64, ldc as i64) + .map_err(|e| MLError::ModelError(format!("C layout {label}: {e:?}")))?; + let d_layout = cublaslt_result::create_matrix_layout(f32_type, m as u64, n as u64, ldc as i64) + .map_err(|e| MLError::ModelError(format!("D layout {label}: {e:?}")))?; + + let matmul_pref = cublaslt_result::create_matmul_pref() + .map_err(|e| MLError::ModelError(format!("matmul pref {label}: {e:?}")))?; + let ws_size_u64 = self.workspace_size as u64; + cublaslt_result::set_matmul_pref_attribute( + matmul_pref, + cublaslt_sys::cublasLtMatmulPreferenceAttributes_t::CUBLASLT_MATMUL_PREF_MAX_WORKSPACE_BYTES, + &ws_size_u64 as *const u64 as *const std::ffi::c_void, + std::mem::size_of::(), + ).map_err(|e| MLError::ModelError(format!("set workspace pref {label}: {e:?}")))?; + + let mut heuristic_result = std::mem::MaybeUninit::::uninit(); + let mut num_results: i32 = 0; + let heuristic_status = cublaslt_sys::cublasLtMatmulAlgoGetHeuristic( + self.lt_handle.0, + matmul_desc, + a_layout, + b_layout, + c_layout, + d_layout, + matmul_pref, + 1, + heuristic_result.as_mut_ptr(), + &mut num_results, + ); + if heuristic_status != cublaslt_sys::cublasStatus_t::CUBLAS_STATUS_SUCCESS || num_results == 0 { + return Err(MLError::ModelError(format!( + "cublasLtMatmulAlgoGetHeuristic {label}: status={heuristic_status:?}, results={num_results}" + ))); + } + let chosen_algo = &(*heuristic_result.as_ptr()).algo as *const cublaslt_sys::cublasLtMatmulAlgo_t; + + let cu_stream = stream.cu_stream() as cublaslt_sys::cudaStream_t; + let status = cublaslt_sys::cublasLtMatmul( + self.lt_handle.0, + matmul_desc, + &alpha as *const f32 as *const std::ffi::c_void, + a_ptr as *const std::ffi::c_void, + a_layout, + b_ptr as *const std::ffi::c_void, + b_layout, + &beta as *const f32 as *const std::ffi::c_void, + c_ptr as *const std::ffi::c_void, + c_layout, + c_ptr as *mut std::ffi::c_void, + d_layout, + chosen_algo, + self.workspace_ptr as *mut std::ffi::c_void, + self.workspace_size, + cu_stream, + ); + if status != cublaslt_sys::cublasStatus_t::CUBLAS_STATUS_SUCCESS { + return Err(MLError::ModelError(format!( + "cublasLtMatmul {label} (m={m},n={n},k={k}): {status:?}" + ))); + } + + // Clean up descriptors + cublaslt_result::destroy_matmul_desc(matmul_desc).ok(); + cublaslt_result::destroy_matrix_layout(a_layout).ok(); + cublaslt_result::destroy_matrix_layout(b_layout).ok(); + cublaslt_result::destroy_matrix_layout(c_layout).ok(); + cublaslt_result::destroy_matrix_layout(d_layout).ok(); + cublaslt_result::destroy_matmul_pref(matmul_pref).ok(); + } + Ok(()) + } + + /// Forward GEMM for the curiosity MLP: + /// C[N, out_dim] = A[N, in_dim] @ W^T[out_dim, in_dim] + /// + /// Row-major GEMM using the column-major transpose trick: + /// cuBLAS col-major: C^T = W @ A^T → TRANSA=N (W), TRANSB=T (A) + /// BUT cublasLt row-major representation: we swap A↔B and set TRANSA=T, TRANSB=N. + /// + /// This matches `CublasGemmSet::sgemm_f32` which uses the same convention. + fn gemm_fwd( + &self, + stream: &CudaStream, + w_ptr: u64, // W[out_dim, in_dim] (row-major) + a_ptr: u64, // A[N, in_dim] (row-major) + c_ptr: u64, // C[N, out_dim] (row-major, output) + out_dim: usize, + n: usize, + in_dim: usize, + label: &str, + ) -> Result<(), MLError> { + // Same call convention as CublasGemmSet::lt_matmul with ldb==k: + // m=out_dim, n=batch, k=in_dim (column-major args) + // A(W) is [K=in_dim, M=out_dim] col-major → row-major [out_dim, in_dim] → TRANSA=T + // B(input) is [K=in_dim, N=batch] col-major → row-major [N, in_dim] → TRANSB=N + let m = out_dim as i32; + let batch = n as i32; + let k = in_dim as i32; + self.sgemm( + stream, + cublas_sys::cublasOperation_t::CUBLAS_OP_T, // W is transposed (row→col convention) + cublas_sys::cublasOperation_t::CUBLAS_OP_N, // A (input) not transposed + m, batch, k, + 1.0, + w_ptr, k, // lda = K (W is [out_dim, in_dim] row-major, treated as [K, M] col-major) + a_ptr, k, // ldb = K (A is [N, in_dim] row-major, treated as [K, N] col-major) + 0.0, + c_ptr, m, // ldc = M = out_dim + label, + ) + } + + /// Compute weight gradient: dW[out_dim, in_dim] row-major = dy^T @ x + /// + /// dy is `[N, out_dim]` row-major (upstream gradient). + /// x is `[N, in_dim]` row-major (saved activation or input). + fn gemm_dw( + &self, + stream: &CudaStream, + dy_ptr: u64, // dy[N, out_dim] row-major + x_ptr: u64, // x[N, in_dim] row-major + dw_ptr: u64, // dW[out_dim, in_dim] row-major (output) + out_dim: usize, + n_batch: usize, + in_dim: usize, + label: &str, + ) -> Result<(), MLError> { + self.sgemm_raw(stream, out_dim, in_dim, n_batch, dy_ptr, x_ptr, dw_ptr, label) + } + + /// Raw cublasSgemm call for weight gradient: dW[out_dim, in_dim] = dy^T @ x + /// + /// cublasSgemm: dW[out_dim, in_dim] row-major = dy^T @ x + /// + /// Computes C^T[in_dim, out_dim] col-major = x^T[in_dim, N] @ dy[N, out_dim] + /// which reads back as dW[out_dim, in_dim] row-major. + /// m=in_dim, n=out_dim, k=N + /// A=x col-major [in_dim, N], TRANSA=N → op(A)=[in_dim, N], lda=in_dim + /// B=dy col-major [out_dim, N], TRANSB=T → op(B)=[N, out_dim], ldb=out_dim + /// C = [in_dim, out_dim] col-major = dW[out_dim, in_dim] row-major ✓ + #[allow(clippy::too_many_arguments)] + fn sgemm_raw( + &self, + stream: &CudaStream, + out_dim: usize, + in_dim: usize, + n_batch: usize, + dy_ptr: u64, // dy[n_batch, out_dim] row-major + x_ptr: u64, // x[n_batch, in_dim] row-major + dw_ptr: u64, // dW[out_dim, in_dim] row-major output + label: &str, + ) -> Result<(), MLError> { + let m = in_dim as i32; + let n = out_dim as i32; + let k = n_batch as i32; + + unsafe { + let cu_stream = stream.cu_stream() as *mut cublas_sys::CUstream_st; + cublas_result::set_stream(self.handle.0, cu_stream) + .map_err(|e| MLError::ModelError(format!("set_stream {label}: {e:?}")))?; + cublas_sys::cublasSetWorkspace_v2( + self.handle.0, + self.workspace_ptr as *mut std::ffi::c_void, + self.workspace_size, + ); + + let alpha: f32 = 1.0; + let beta: f32 = 0.0; + let status = cublas_sys::cublasSgemm_v2( + self.handle.0, + cublas_sys::cublasOperation_t::CUBLAS_OP_N, // A=x: no transpose (col-major [in_dim, N]) + cublas_sys::cublasOperation_t::CUBLAS_OP_T, // B=dy: transpose (col-major [out_dim, N] → [N, out_dim]) + m, // M = in_dim + n, // N = out_dim + k, // K = batch + &alpha, + x_ptr as *const f32, + m, // lda = in_dim (leading dim of x col-major [in_dim, N]) + dy_ptr as *const f32, + n, // ldb = out_dim (leading dim of dy col-major [out_dim, N]) + &beta, + dw_ptr as *mut f32, + m, // ldc = in_dim (leading dim of dW col-major [in_dim, out_dim]) + ); + if status != cublas_sys::cublasStatus_t::CUBLAS_STATUS_SUCCESS { + return Err(MLError::ModelError(format!( + "cublasSgemm dW {label} (m={m},n={n},k={k}): {status:?}" + ))); + } + } + Ok(()) + } + + /// Backward GEMM: dx[N, in_dim] = dy[N, out_dim] @ W[out_dim, in_dim] (row-major). + /// + /// Computes dx^T col-major = W^T col-major @ dy^T col-major: + /// m=in_dim, n=N, k=out_dim + /// A=W col-major [in_dim, out_dim], TRANSA=N, lda=in_dim + /// B=dy col-major [out_dim, N], TRANSB=N, ldb=out_dim + /// C = [in_dim, N] col-major = dx^T → dx[N, in_dim] row-major ✓ + fn gemm_dx( + &self, + stream: &CudaStream, + dy_ptr: u64, // dy[N, out_dim] row-major + w_ptr: u64, // W[out_dim, in_dim] row-major + dx_ptr: u64, // dx[N, in_dim] row-major (output) + n_batch: usize, + out_dim: usize, + in_dim: usize, + label: &str, + ) -> Result<(), MLError> { + let m = in_dim as i32; + let n = n_batch as i32; + let k = out_dim as i32; + + unsafe { + let cu_stream = stream.cu_stream() as *mut cublas_sys::CUstream_st; + cublas_result::set_stream(self.handle.0, cu_stream) + .map_err(|e| MLError::ModelError(format!("set_stream {label}: {e:?}")))?; + cublas_sys::cublasSetWorkspace_v2( + self.handle.0, + self.workspace_ptr as *mut std::ffi::c_void, + self.workspace_size, + ); + + let alpha: f32 = 1.0; + let beta: f32 = 0.0; + let status = cublas_sys::cublasSgemm_v2( + self.handle.0, + cublas_sys::cublasOperation_t::CUBLAS_OP_N, // W: no transpose (col-major [in_dim, out_dim]) + cublas_sys::cublasOperation_t::CUBLAS_OP_N, // dy: no transpose (col-major [out_dim, N]) + m, // M = in_dim + n, // N = batch + k, // K = out_dim + &alpha, + w_ptr as *const f32, + m, // lda = in_dim (leading dim of W col-major [in_dim, out_dim]) + dy_ptr as *const f32, + k, // ldb = out_dim (leading dim of dy col-major [out_dim, N]) + &beta, + dx_ptr as *mut f32, + m, // ldc = in_dim (leading dim of dx^T col-major [in_dim, N]) + ); + if status != cublas_sys::cublasStatus_t::CUBLAS_STATUS_SUCCESS { + return Err(MLError::ModelError(format!( + "cublasSgemm dx {label} (m={m},n={n},k={k}): {status:?}" + ))); + } + } + Ok(()) + } +} + // --------------------------------------------------------------------------- // GpuCuriosityTrainer // --------------------------------------------------------------------------- -/// GPU-resident curiosity forward model trainer. +/// GPU-resident curiosity forward model trainer — cuBLAS GEMM pipeline. /// /// Trains the curiosity MLP entirely on GPU using experience data that is /// already device-resident. Maintains gradient buffers and Adam optimizer /// state. Modifies [`CuriosityWeightSet`] in-place -- zero CPU traffic. /// -/// Gradient accumulation is fully deterministic: per-block shared-memory -/// reduction followed by a sequential sum over block partials. No atomicAdd. +/// The 1085ms `curiosity_fwd_bwd_per_block` kernel is replaced by cuBLAS +/// GEMMs (same approach as `launch_curiosity_inference` in the DQN trainer). #[allow(missing_debug_implementations)] // CudaSlice does not implement Debug pub struct GpuCuriosityTrainer { stream: Arc, - // Kernel functions + // cuBLAS context for curiosity GEMMs + gemm: CuriosityGemm, + + // Kernel functions (from training cubin) shift_func: CudaFunction, - fwd_bwd_per_block_func: CudaFunction, - grad_reduce_func: CudaFunction, + mse_fwd_grad_func: CudaFunction, + leaky_relu_bwd_func: CudaFunction, + bias_grad_reduce_func: CudaFunction, adam_func: CudaFunction, + // Kernel functions (from inference cubin — reused for forward pass) + prepare_input_func: CudaFunction, + bias_leaky_relu_func: CudaFunction, + // Gradient buffers grad_w1: CudaSlice, // [CUR_W1_LEN] grad_b1: CudaSlice, // [CUR_B1_LEN] grad_w2: CudaSlice, // [CUR_W2_LEN] grad_b2: CudaSlice, // [CUR_B2_LEN] - // Per-block partial gradient buffer for deterministic reduction - partial_grads: CudaSlice, // [max_blocks * CUR_TOTAL_PARAMS] - max_blocks: usize, - // Adam first moment (per-param-group) adam_m_w1: CudaSlice, // [CUR_W1_LEN] adam_m_b1: CudaSlice, // [CUR_B1_LEN] @@ -95,8 +488,12 @@ pub struct GpuCuriosityTrainer { adam_v_w2: CudaSlice, // [CUR_W2_LEN] adam_v_b2: CudaSlice, // [CUR_B2_LEN] - // Shifted next_states buffer - next_states_buf: CudaSlice, + // Intermediate buffers for forward/backward pass + next_states_buf: CudaSlice, // [max_samples * state_dim] + input_buf: CudaSlice, // [max_samples * CUR_INPUT] — prepare_input output + hidden_buf: CudaSlice, // [max_samples * CUR_HIDDEN] — post-activation hidden + pred_buf: CudaSlice, // [max_samples * CUR_OUTPUT] — pred + d_pred (reused) + d_hidden_buf: CudaSlice, // [max_samples * CUR_HIDDEN] — d_hidden backward // Adam step counter (1-based) step: i32, @@ -104,14 +501,11 @@ pub struct GpuCuriosityTrainer { // State dimension of the experience data state_dim: usize, - // Maximum number of samples the next_states buffer was allocated for + // Maximum number of samples the buffers were allocated for buf_capacity: usize, } /// Launch Adam optimizer step for one parameter group. -/// -/// Free function to avoid borrow conflicts when calling from -/// `train_on_collector_buffers` (which mutably borrows multiple fields). fn launch_adam_step( stream: &CudaStream, adam_func: &CudaFunction, @@ -155,117 +549,111 @@ fn launch_adam_step( impl GpuCuriosityTrainer { /// Create a new GPU curiosity trainer. /// - /// Loads the precompiled CUDA training cubin, allocates gradient and - /// Adam optimizer state buffers as zeros on GPU. + /// Loads the precompiled CUDA training + inference cubins, creates a + /// dedicated cuBLAS handle, and allocates all intermediate buffers on GPU. /// /// # Arguments /// * `stream` - CUDA stream for all GPU operations /// * `state_dim` - Dimensionality of state vectors in the experience buffer - /// * `max_samples` - Maximum number of training samples (for next_states buffer sizing) + /// * `max_samples` - Maximum number of training samples (for buffer sizing) pub fn new( stream: Arc, state_dim: usize, max_samples: usize, ) -> Result { - // ---- Load precompiled cubin ---- + // ---- Load training cubin ---- let context = stream.context(); - let module = context.load_cubin(CURIOSITY_TRAINING_CUBIN.to_vec()).map_err(|e| { + let train_module = context.load_cubin(CURIOSITY_TRAINING_CUBIN.to_vec()).map_err(|e| { MLError::ModelError(format!("curiosity training module load: {e}")) })?; - let shift_func = module.load_function("curiosity_shift_states").map_err(|e| { + let shift_func = train_module.load_function("curiosity_shift_states").map_err(|e| { MLError::ModelError(format!("curiosity_shift_states load: {e}")) })?; - let fwd_bwd_per_block_func = module.load_function("curiosity_fwd_bwd_per_block").map_err(|e| { - MLError::ModelError(format!("curiosity_fwd_bwd_per_block load: {e}")) + let mse_fwd_grad_func = train_module.load_function("curiosity_mse_fwd_grad").map_err(|e| { + MLError::ModelError(format!("curiosity_mse_fwd_grad load: {e}")) })?; - let grad_reduce_func = module.load_function("curiosity_grad_reduce").map_err(|e| { - MLError::ModelError(format!("curiosity_grad_reduce load: {e}")) + let leaky_relu_bwd_func = train_module.load_function("curiosity_leaky_relu_bwd").map_err(|e| { + MLError::ModelError(format!("curiosity_leaky_relu_bwd load: {e}")) })?; - let adam_func = module.load_function("curiosity_adam_step").map_err(|e| { + let bias_grad_reduce_func = train_module.load_function("curiosity_bias_grad_reduce").map_err(|e| { + MLError::ModelError(format!("curiosity_bias_grad_reduce load: {e}")) + })?; + let adam_func = train_module.load_function("curiosity_adam_step").map_err(|e| { MLError::ModelError(format!("curiosity_adam_step load: {e}")) })?; + // ---- Load inference cubin (for prepare_input + bias_leaky_relu) ---- + let infer_module = context.load_cubin(CURIOSITY_INFERENCE_CUBIN.to_vec()).map_err(|e| { + MLError::ModelError(format!("curiosity inference module load: {e}")) + })?; + let prepare_input_func = infer_module.load_function("curiosity_prepare_input").map_err(|e| { + MLError::ModelError(format!("curiosity_prepare_input load: {e}")) + })?; + let bias_leaky_relu_func = infer_module.load_function("curiosity_bias_leaky_relu").map_err(|e| { + MLError::ModelError(format!("curiosity_bias_leaky_relu load: {e}")) + })?; + + // ---- Create cuBLAS context ---- + let gemm = CuriosityGemm::new(&stream)?; + // ---- Allocate gradient buffers ---- - let grad_w1 = stream.alloc_zeros::(CUR_W1_LEN).map_err(|e| { - MLError::ModelError(format!("alloc grad_w1: {e}")) - })?; - let grad_b1 = stream.alloc_zeros::(CUR_B1_LEN).map_err(|e| { - MLError::ModelError(format!("alloc grad_b1: {e}")) - })?; - let grad_w2 = stream.alloc_zeros::(CUR_W2_LEN).map_err(|e| { - MLError::ModelError(format!("alloc grad_w2: {e}")) - })?; - let grad_b2 = stream.alloc_zeros::(CUR_B2_LEN).map_err(|e| { - MLError::ModelError(format!("alloc grad_b2: {e}")) - })?; + let grad_w1 = stream.alloc_zeros::(CUR_W1_LEN).map_err(|e| MLError::ModelError(format!("alloc grad_w1: {e}")))?; + let grad_b1 = stream.alloc_zeros::(CUR_B1_LEN).map_err(|e| MLError::ModelError(format!("alloc grad_b1: {e}")))?; + let grad_w2 = stream.alloc_zeros::(CUR_W2_LEN).map_err(|e| MLError::ModelError(format!("alloc grad_w2: {e}")))?; + let grad_b2 = stream.alloc_zeros::(CUR_B2_LEN).map_err(|e| MLError::ModelError(format!("alloc grad_b2: {e}")))?; - // ---- Allocate per-block partial gradient buffer ---- - // With FWD_BWD_BLOCK_SIZE=256 threads/block and max_samples samples: - // max_blocks = ceil(max_samples / 256). Buffer = max_blocks * CUR_TOTAL_PARAMS floats. - let max_blocks = (max_samples + FWD_BWD_BLOCK_SIZE as usize - 1) / FWD_BWD_BLOCK_SIZE as usize; - let partial_grads = stream - .alloc_zeros::(max_blocks * CUR_TOTAL_PARAMS) - .map_err(|e| { - MLError::ModelError(format!("alloc partial_grads: {e}")) - })?; + // ---- Allocate Adam moment buffers ---- + let adam_m_w1 = stream.alloc_zeros::(CUR_W1_LEN).map_err(|e| MLError::ModelError(format!("alloc adam_m_w1: {e}")))?; + let adam_m_b1 = stream.alloc_zeros::(CUR_B1_LEN).map_err(|e| MLError::ModelError(format!("alloc adam_m_b1: {e}")))?; + let adam_m_w2 = stream.alloc_zeros::(CUR_W2_LEN).map_err(|e| MLError::ModelError(format!("alloc adam_m_w2: {e}")))?; + let adam_m_b2 = stream.alloc_zeros::(CUR_B2_LEN).map_err(|e| MLError::ModelError(format!("alloc adam_m_b2: {e}")))?; - // ---- Allocate Adam first moment buffers ---- - let adam_m_w1 = stream.alloc_zeros::(CUR_W1_LEN).map_err(|e| { - MLError::ModelError(format!("alloc adam_m_w1: {e}")) - })?; - let adam_m_b1 = stream.alloc_zeros::(CUR_B1_LEN).map_err(|e| { - MLError::ModelError(format!("alloc adam_m_b1: {e}")) - })?; - let adam_m_w2 = stream.alloc_zeros::(CUR_W2_LEN).map_err(|e| { - MLError::ModelError(format!("alloc adam_m_w2: {e}")) - })?; - let adam_m_b2 = stream.alloc_zeros::(CUR_B2_LEN).map_err(|e| { - MLError::ModelError(format!("alloc adam_m_b2: {e}")) - })?; + let adam_v_w1 = stream.alloc_zeros::(CUR_W1_LEN).map_err(|e| MLError::ModelError(format!("alloc adam_v_w1: {e}")))?; + let adam_v_b1 = stream.alloc_zeros::(CUR_B1_LEN).map_err(|e| MLError::ModelError(format!("alloc adam_v_b1: {e}")))?; + let adam_v_w2 = stream.alloc_zeros::(CUR_W2_LEN).map_err(|e| MLError::ModelError(format!("alloc adam_v_w2: {e}")))?; + let adam_v_b2 = stream.alloc_zeros::(CUR_B2_LEN).map_err(|e| MLError::ModelError(format!("alloc adam_v_b2: {e}")))?; - // ---- Allocate Adam second moment buffers ---- - let adam_v_w1 = stream.alloc_zeros::(CUR_W1_LEN).map_err(|e| { - MLError::ModelError(format!("alloc adam_v_w1: {e}")) - })?; - let adam_v_b1 = stream.alloc_zeros::(CUR_B1_LEN).map_err(|e| { - MLError::ModelError(format!("alloc adam_v_b1: {e}")) - })?; - let adam_v_w2 = stream.alloc_zeros::(CUR_W2_LEN).map_err(|e| { - MLError::ModelError(format!("alloc adam_v_w2: {e}")) - })?; - let adam_v_b2 = stream.alloc_zeros::(CUR_B2_LEN).map_err(|e| { - MLError::ModelError(format!("alloc adam_v_b2: {e}")) - })?; - - // ---- Allocate shifted next_states buffer ---- + // ---- Allocate intermediate buffers ---- let next_states_buf = stream .alloc_zeros::(max_samples * state_dim) - .map_err(|e| { - MLError::ModelError(format!("alloc next_states_buf: {e}")) - })?; + .map_err(|e| MLError::ModelError(format!("alloc next_states_buf: {e}")))?; + let input_buf = stream + .alloc_zeros::(max_samples * CUR_INPUT) + .map_err(|e| MLError::ModelError(format!("alloc input_buf: {e}")))?; + let hidden_buf = stream + .alloc_zeros::(max_samples * CUR_HIDDEN) + .map_err(|e| MLError::ModelError(format!("alloc hidden_buf: {e}")))?; + let pred_buf = stream + .alloc_zeros::(max_samples * CUR_OUTPUT) + .map_err(|e| MLError::ModelError(format!("alloc pred_buf: {e}")))?; + let d_hidden_buf = stream + .alloc_zeros::(max_samples * CUR_HIDDEN) + .map_err(|e| MLError::ModelError(format!("alloc d_hidden_buf: {e}")))?; debug!( state_dim, max_samples, - max_blocks, - partial_grads_bytes = max_blocks * CUR_TOTAL_PARAMS * 4, - total_params = CUR_TOTAL_PARAMS, - "GPU curiosity trainer initialized (deterministic per-block reduce)" + cur_input = CUR_INPUT, + cur_hidden = CUR_HIDDEN, + cur_output = CUR_OUTPUT, + "GPU curiosity trainer initialized (cuBLAS GEMM pipeline)" ); Ok(Self { stream, + gemm, shift_func, - fwd_bwd_per_block_func, - grad_reduce_func, + mse_fwd_grad_func, + leaky_relu_bwd_func, + bias_grad_reduce_func, adam_func, + prepare_input_func, + bias_leaky_relu_func, grad_w1, grad_b1, grad_w2, grad_b2, - partial_grads, - max_blocks, adam_m_w1, adam_m_b1, adam_m_w2, @@ -275,6 +663,10 @@ impl GpuCuriosityTrainer { adam_v_w2, adam_v_b2, next_states_buf, + input_buf, + hidden_buf, + pred_buf, + d_hidden_buf, step: 0, state_dim, buf_capacity: max_samples, @@ -283,37 +675,26 @@ impl GpuCuriosityTrainer { /// Train the curiosity forward model on GPU-resident experience data. /// - /// Performs one training step (forward + backward + Adam update) entirely - /// on GPU. The states buffer is shifted by one timestep to produce - /// next_states -- episode boundary noise is negligible for this tiny - /// auxiliary model. - /// - /// Gradient accumulation is fully deterministic: - /// 1. `curiosity_fwd_bwd_per_block` -- block-level shared-memory reduce - /// 2. `curiosity_grad_reduce` -- sequential sum over block partials + /// Performs one training step via cuBLAS GEMM pipeline (forward + backward + Adam). + /// All operations run on GPU with zero CPU roundtrips. /// /// # Arguments /// * `weights` - Curiosity model weights to update in-place on GPU /// * `states` - State observations `[n_samples * state_dim]` on GPU - /// * `actions` - Action indices `[n_samples]` on GPU (i32, 0-4 for DQN) - /// * `n_samples` - Number of experience samples (will use n_samples-1 for training, - /// since the last sample has no valid next_state) - /// - /// # Errors - /// Returns `MLError::ModelError` on kernel launch or buffer size mismatch. + /// * `actions` - Action indices `[n_samples]` on GPU + /// * `n_samples` - Number of experience samples pub fn train_on_collector_buffers( &mut self, weights: &mut CuriosityWeightSet, - states: &CudaSlice, // #30: f32 states from experience collector + states: &CudaSlice, actions: &CudaSlice, n_samples: usize, ) -> Result<(), MLError> { - // Need at least 2 samples (one for state, one for next_state via shift) if n_samples < 2 { return Ok(()); } - // Effective training samples: n_samples - 1 (last has no valid next_state) + // n_train = n_samples - 1 (last sample has no valid next_state via shift) let n_train = n_samples - 1; if n_train > self.buf_capacity { @@ -326,11 +707,12 @@ impl GpuCuriosityTrainer { let sd = self.state_dim; let sd_i32 = sd as i32; let n_i32 = n_train as i32; + let stream = &*self.stream; self.stream.synchronize() .map_err(|e| MLError::ModelError(format!("curiosity PRE-SHIFT sync FAILED: {e}")))?; - // ---- Launch 1/4: Build shifted next_states buffer ---- + // ── Step 1: Build shifted next_states buffer ───────────────────────── let shift_total = n_train * sd; let shift_cfg = LaunchConfig { grid_dim: (((shift_total as u32) + 255) / 256, 1, 1), @@ -338,108 +720,149 @@ impl GpuCuriosityTrainer { shared_mem_bytes: 0, }; unsafe { - self.stream + stream .launch_builder(&self.shift_func) .arg(states) .arg(&mut self.next_states_buf) .arg(&n_i32) .arg(&sd_i32) .launch(shift_cfg) - .map_err(|e| { - MLError::ModelError(format!("curiosity_shift_states launch: {e}")) - })?; + .map_err(|e| MLError::ModelError(format!("curiosity_shift_states: {e}")))?; } - self.stream.synchronize() - .map_err(|e| MLError::ModelError(format!("curiosity shift kernel CRASHED: {e}")))?; - - // ---- Launch 2/4: Forward + backward per-block ---- - let num_blocks = ((n_train as u32) + FWD_BWD_BLOCK_SIZE - 1) / FWD_BWD_BLOCK_SIZE; - - if (num_blocks as usize) > self.max_blocks { - return Err(MLError::ModelError(format!( - "curiosity trainer: num_blocks={num_blocks} exceeds max_blocks={}", - self.max_blocks - ))); - } - - let fwd_bwd_cfg = LaunchConfig { - grid_dim: (num_blocks, 1, 1), - block_dim: (FWD_BWD_BLOCK_SIZE, 1, 1), - shared_mem_bytes: FWD_BWD_BLOCK_SIZE * std::mem::size_of::() as u32, - }; - unsafe { - self.stream - .launch_builder(&self.fwd_bwd_per_block_func) - .arg(states) - .arg(actions) - .arg(&self.next_states_buf) - .arg(&weights.w1) - .arg(&weights.b1) - .arg(&weights.w2) - .arg(&weights.b2) - .arg(&mut self.partial_grads) - .arg(&n_i32) - .arg(&sd_i32) - .launch(fwd_bwd_cfg) - .map_err(|e| MLError::ModelError(format!("curiosity fwd_bwd_per_block launch: {e}")))?; - } - self.stream.synchronize() - .map_err(|e| MLError::ModelError(format!("curiosity fwd_bwd_per_block CRASHED: {e}")))?; - - // ---- Launch 3/4: Deterministic gradient reduction ---- - let total_params_i32 = CUR_TOTAL_PARAMS as i32; - let num_blocks_i32 = num_blocks as i32; - let reduce_cfg = LaunchConfig { - grid_dim: (((CUR_TOTAL_PARAMS as u32) + 255) / 256, 1, 1), + // ── Step 2: Build input[N, CUR_INPUT] from states + action one-hot ─── + let blocks_n = ((n_train as u32) + 255) / 256; + let elem_cfg = |total: usize| LaunchConfig { + grid_dim: (((total as u32) + 255) / 256, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0, }; - // Suppress unused variable warning — total_params_i32 used only for clarity - let _ = total_params_i32; unsafe { - self.stream - .launch_builder(&self.grad_reduce_func) - .arg(&self.partial_grads) - .arg(&mut self.grad_w1) - .arg(&mut self.grad_b1) - .arg(&mut self.grad_w2) - .arg(&mut self.grad_b2) - .arg(&num_blocks_i32) - .launch(reduce_cfg) - .map_err(|e| MLError::ModelError(format!("curiosity grad_reduce launch: {e}")))?; + stream + .launch_builder(&self.prepare_input_func) + .arg(states) + .arg(actions) + .arg(&mut self.input_buf) + .arg(&n_i32) + .arg(&sd_i32) + .launch(LaunchConfig { + grid_dim: (blocks_n, 1, 1), + block_dim: (256, 1, 1), + shared_mem_bytes: 0, + }) + .map_err(|e| MLError::ModelError(format!("curiosity_prepare_input: {e}")))?; } - // ---- Launch 4/4: Adam optimizer update (4 separate launches) ---- + // ── Step 3: cuBLAS GEMM1 — hidden[N, CUR_HIDDEN] = input @ W1^T ───── + let w1_ptr = weights.w1.raw_ptr(); + let input_ptr = self.input_buf.raw_ptr(); + let hidden_ptr = self.hidden_buf.raw_ptr(); + self.gemm.gemm_fwd(stream, w1_ptr, input_ptr, hidden_ptr, CUR_HIDDEN, n_train, CUR_INPUT, "cur_train_gemm1")?; + + // ── Step 4: bias_leaky_relu — hidden += b1, LeakyReLU(0.01) in-place ─ + let total_hidden = (n_train * CUR_HIDDEN) as i32; + let cur_hidden_i32 = CUR_HIDDEN as i32; + let b1_ptr = weights.b1.raw_ptr(); + unsafe { + stream + .launch_builder(&self.bias_leaky_relu_func) + .arg(&hidden_ptr) + .arg(&b1_ptr) + .arg(&cur_hidden_i32) + .arg(&total_hidden) + .launch(elem_cfg(n_train * CUR_HIDDEN)) + .map_err(|e| MLError::ModelError(format!("curiosity_bias_leaky_relu: {e}")))?; + } + + // ── Step 5: cuBLAS GEMM2 — pred[N, CUR_OUTPUT] = hidden @ W2^T ────── + let w2_ptr = weights.w2.raw_ptr(); + let pred_ptr = self.pred_buf.raw_ptr(); + self.gemm.gemm_fwd(stream, w2_ptr, hidden_ptr, pred_ptr, CUR_OUTPUT, n_train, CUR_HIDDEN, "cur_train_gemm2")?; + + // ── Step 6: mse_fwd_grad — +b2 in-place, d_pred = 2/CUR_OUTPUT*(pred-target) ─ + // pred_buf is reused: after this kernel it contains d_pred[N, CUR_OUTPUT] + let b2_ptr = weights.b2.raw_ptr(); + let next_states_ptr = self.next_states_buf.raw_ptr(); + unsafe { + stream + .launch_builder(&self.mse_fwd_grad_func) + .arg(&pred_ptr) + .arg(&b2_ptr) + .arg(&next_states_ptr) + .arg(&n_i32) + .arg(&sd_i32) + .launch(elem_cfg(n_train * CUR_OUTPUT)) + .map_err(|e| MLError::ModelError(format!("curiosity_mse_fwd_grad: {e}")))?; + } + + // ── Step 7: cuBLAS dW2 — grad_w2 = d_pred^T @ hidden ──────────────── + // d_pred = pred_buf[N, CUR_OUTPUT], hidden = hidden_buf[N, CUR_HIDDEN] + // grad_w2[CUR_OUTPUT, CUR_HIDDEN] = d_pred^T @ hidden + let grad_w2_ptr = self.grad_w2.raw_ptr(); + self.gemm.gemm_dw(stream, pred_ptr, hidden_ptr, grad_w2_ptr, CUR_OUTPUT, n_train, CUR_HIDDEN, "cur_dW2")?; + + // ── Step 8: bias_grad_reduce — grad_b2 = sum_batch(d_pred) ─────────── + let grad_b2_ptr = self.grad_b2.raw_ptr(); + let cur_output_i32 = CUR_OUTPUT as i32; + unsafe { + stream + .launch_builder(&self.bias_grad_reduce_func) + .arg(&pred_ptr) // d_pred[N, CUR_OUTPUT] + .arg(&grad_b2_ptr) + .arg(&n_i32) + .arg(&cur_output_i32) + .launch(elem_cfg(CUR_OUTPUT)) + .map_err(|e| MLError::ModelError(format!("curiosity_bias_grad_reduce b2: {e}")))?; + } + + // ── Step 9: cuBLAS d_hidden — d_hidden = d_pred @ W2 ───────────────── + // d_pred[N, CUR_OUTPUT] @ W2[CUR_OUTPUT, CUR_HIDDEN] = d_hidden[N, CUR_HIDDEN] + let d_hidden_ptr = self.d_hidden_buf.raw_ptr(); + self.gemm.gemm_dx(stream, pred_ptr, w2_ptr, d_hidden_ptr, n_train, CUR_OUTPUT, CUR_HIDDEN, "cur_d_hidden")?; + + // ── Step 10: leaky_relu_bwd — mask d_hidden by post-activation sign ── + let total_hidden_usize = n_train * CUR_HIDDEN; + let total_hidden_i32 = total_hidden_usize as i32; + unsafe { + stream + .launch_builder(&self.leaky_relu_bwd_func) + .arg(&d_hidden_ptr) + .arg(&hidden_ptr) // post-activation hidden (sign = sign of pre-activation) + .arg(&total_hidden_i32) + .launch(elem_cfg(total_hidden_usize)) + .map_err(|e| MLError::ModelError(format!("curiosity_leaky_relu_bwd: {e}")))?; + } + + // ── Step 11: cuBLAS dW1 — grad_w1 = d_hidden^T @ input ─────────────── + // d_hidden[N, CUR_HIDDEN], input[N, CUR_INPUT] + // grad_w1[CUR_HIDDEN, CUR_INPUT] = d_hidden^T @ input + let grad_w1_ptr = self.grad_w1.raw_ptr(); + self.gemm.gemm_dw(stream, d_hidden_ptr, input_ptr, grad_w1_ptr, CUR_HIDDEN, n_train, CUR_INPUT, "cur_dW1")?; + + // ── Step 12: bias_grad_reduce — grad_b1 = sum_batch(d_hidden) ───────── + let grad_b1_ptr = self.grad_b1.raw_ptr(); + let cur_hidden_i32_2 = CUR_HIDDEN as i32; + unsafe { + stream + .launch_builder(&self.bias_grad_reduce_func) + .arg(&d_hidden_ptr) // d_hidden[N, CUR_HIDDEN] + .arg(&grad_b1_ptr) + .arg(&n_i32) + .arg(&cur_hidden_i32_2) + .launch(elem_cfg(CUR_HIDDEN)) + .map_err(|e| MLError::ModelError(format!("curiosity_bias_grad_reduce b1: {e}")))?; + } + + // ── Step 13: Adam update ────────────────────────────────────────────── self.step += 1; let step = self.step; - launch_adam_step( - &self.stream, &self.adam_func, - &mut weights.w1, &self.grad_w1, CUR_W1_LEN, - &mut self.adam_m_w1, &mut self.adam_v_w1, - n_train, step, - )?; - launch_adam_step( - &self.stream, &self.adam_func, - &mut weights.b1, &self.grad_b1, CUR_B1_LEN, - &mut self.adam_m_b1, &mut self.adam_v_b1, - n_train, step, - )?; - launch_adam_step( - &self.stream, &self.adam_func, - &mut weights.w2, &self.grad_w2, CUR_W2_LEN, - &mut self.adam_m_w2, &mut self.adam_v_w2, - n_train, step, - )?; - launch_adam_step( - &self.stream, &self.adam_func, - &mut weights.b2, &self.grad_b2, CUR_B2_LEN, - &mut self.adam_m_b2, &mut self.adam_v_b2, - n_train, step, - )?; + launch_adam_step(stream, &self.adam_func, &mut weights.w1, &self.grad_w1, CUR_W1_LEN, &mut self.adam_m_w1, &mut self.adam_v_w1, n_train, step)?; + launch_adam_step(stream, &self.adam_func, &mut weights.b1, &self.grad_b1, CUR_B1_LEN, &mut self.adam_m_b1, &mut self.adam_v_b1, n_train, step)?; + launch_adam_step(stream, &self.adam_func, &mut weights.w2, &self.grad_w2, CUR_W2_LEN, &mut self.adam_m_w2, &mut self.adam_v_w2, n_train, step)?; + launch_adam_step(stream, &self.adam_func, &mut weights.b2, &self.grad_b2, CUR_B2_LEN, &mut self.adam_m_b2, &mut self.adam_v_b2, n_train, step)?; - debug!(step, n_train, num_blocks, "curiosity GPU training step complete (deterministic)"); + debug!(step, n_train, "curiosity GPU training step complete (cuBLAS GEMM pipeline)"); Ok(()) }