From d799bfafbe2403ee03dd6adfebdca4f1cf0e9129 Mon Sep 17 00:00:00 2001 From: jgrusewski Date: Fri, 17 Apr 2026 23:38:00 +0200 Subject: [PATCH] feat: rewrite IQL trainer to batched cuBLAS GEMMs MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Replace per-sample CUDA kernels (16,384 launches per GEMM) with batched cublasLtMatmul — ONE call per GEMM layer. Forward pass uses 3 GEMMs + SiLU + bias_add + expectile_loss. Backward pass uses 5 GEMMs + SiLU backward + bias_grad_reduce. Eliminates grads_per_sample buffer (was 27MB), tiled backward loop, and weight_grad_reduce kernel. Constructor now accepts Arc instead of bare stream, sharing the cuBLAS handle with the trunk forward/backward pipeline. Eight IqlGemmDesc descriptors are cached at init with heuristic algo selection for CUDA Graph compatibility. Co-Authored-By: Claude Opus 4.6 (1M context) --- .../ml/src/cuda_pipeline/gpu_iql_trainer.rs | 952 +++++++++++++----- crates/ml/src/trainers/dqn/fused_training.rs | 5 +- 2 files changed, 683 insertions(+), 274 deletions(-) diff --git a/crates/ml/src/cuda_pipeline/gpu_iql_trainer.rs b/crates/ml/src/cuda_pipeline/gpu_iql_trainer.rs index ca5ca0506..b4a716b20 100644 --- a/crates/ml/src/cuda_pipeline/gpu_iql_trainer.rs +++ b/crates/ml/src/cuda_pipeline/gpu_iql_trainer.rs @@ -3,19 +3,25 @@ //! GPU-accelerated IQL (Implicit Q-Learning) value network trainer. //! //! Implements the value network component of IQL (Kostrikov et al., 2021) -//! as a fused CUDA kernel pipeline — zero atomicAdd, fully deterministic: +//! using batched cuBLAS GEMMs for matrix multiplication — zero atomicAdd, +//! fully deterministic: //! -//! 1. **Forward + Expectile Loss** — V(s) MLP forward + asymmetric loss -//! 2. **Backward (per-sample)** — backprop into per-sample gradient buffer -//! 3. **Weight Grad Reduce** — deterministic sum across samples -//! 4. **Grad Norm Phase 1+2** — two-phase L2 norm (no atomicAdd) -//! 5. **Adam** — AdamW update with gradient clipping +//! 1. **Forward** — 3 cublasLtMatmul GEMMs + SiLU + bias_add + expectile loss +//! 2. **Backward** — 5 cublasLtMatmul GEMMs + SiLU backward + bias_grad_reduce +//! 3. **Grad Norm Phase 1+2** — two-phase L2 norm (no atomicAdd) +//! 4. **Adam** — AdamW update with gradient clipping //! //! ## Architecture //! //! V(s) is a 2-hidden-layer MLP with SiLU activation: //! state -> Linear(STATE_DIM, H) -> SiLU -> Linear(H, H) -> SiLU -> Linear(H, 1) //! +//! ## cuBLAS integration +//! +//! All matrix multiplications use `cublasLtMatmul` with `CUBLAS_COMPUTE_32F_FAST_TF32` +//! via pre-cached GEMM descriptors (IqlGemmDesc). Element-wise operations (SiLU, +//! expectile loss, bias add/grad) remain as precompiled cubin kernels. +//! //! ## Integration //! //! Called after the main DQN training step in `FusedTrainingCtx::submit_aux_ops()`. @@ -24,10 +30,46 @@ use std::sync::Arc; +use cudarc::cublaslt::result as cublaslt_result; +use cudarc::cublaslt::sys as cublaslt_sys; use cudarc::driver::{CudaFunction, CudaSlice, CudaStream, LaunchConfig, PushKernelArg}; use tracing::info; use crate::MLError; +use super::shared_cublas_handle::SharedCublasHandle; + +// --------------------------------------------------------------------------- +// cuBLAS GEMM descriptor +// --------------------------------------------------------------------------- + +/// Cached cuBLAS GEMM descriptor for IQL forward/backward. +/// Same structure as CachedGemmDesc in batched_forward.rs. +struct IqlGemmDesc { + matmul_desc: cublaslt_sys::cublasLtMatmulDesc_t, + a_layout: cublaslt_sys::cublasLtMatrixLayout_t, + b_layout: cublaslt_sys::cublasLtMatrixLayout_t, + c_layout: cublaslt_sys::cublasLtMatrixLayout_t, + d_layout: cublaslt_sys::cublasLtMatrixLayout_t, + algo: cublaslt_sys::cublasLtMatmulAlgo_t, +} + +/// SAFETY: cublasLt descriptors are opaque pointers owned exclusively by +/// `GpuIqlTrainer`, which is single-threaded in practice (owned by `FusedTrainingCtx` +/// on the GPU trainer thread). +unsafe impl Send for IqlGemmDesc {} +unsafe impl Sync for IqlGemmDesc {} + +impl Drop for IqlGemmDesc { + fn drop(&mut self) { + unsafe { + let _ = cublaslt_result::destroy_matrix_layout(self.d_layout); + let _ = cublaslt_result::destroy_matrix_layout(self.c_layout); + let _ = cublaslt_result::destroy_matrix_layout(self.b_layout); + let _ = cublaslt_result::destroy_matrix_layout(self.a_layout); + let _ = cublaslt_result::destroy_matmul_desc(self.matmul_desc); + } + } +} // --------------------------------------------------------------------------- // Configuration @@ -107,24 +149,47 @@ impl GpuIqlConfig { /// GPU-accelerated IQL value network trainer. /// /// Owns pre-allocated GPU buffers for the V(s) network weights, Adam state, -/// per-sample gradients, and compiled CUDA kernels. Operates entirely on GPU +/// cuBLAS GEMM descriptors, and compiled CUDA kernels. Operates entirely on GPU /// with no per-step CPU-GPU data transfers. /// -/// Zero atomicAdd — all gradient accumulation uses per-sample buffers with -/// deterministic cross-sample reduction. +/// Zero atomicAdd — all matrix multiplications use batched cuBLAS GEMMs. #[allow(missing_debug_implementations)] pub struct GpuIqlTrainer { config: GpuIqlConfig, stream: Arc, - // ── Compiled kernels ──────────────────────────────────────────── - forward_loss_kernel: CudaFunction, - backward_per_sample_kernel: CudaFunction, - weight_grad_reduce_kernel: CudaFunction, + // ── cuBLAS handle (shared with trunk) ────────────────────────── + shared_handle: Arc, + + // ── cuBLAS GEMM descriptors (cached at init) ────────────────── + /// h1_pre = W1^T @ states: M=H, N=B, K=state_dim, ldb=state_dim_padded + gemm_fwd_h1: IqlGemmDesc, + /// h2_pre = W2^T @ h1: M=H, N=B, K=H, ldb=H + gemm_fwd_h2: IqlGemmDesc, + /// v = W3^T @ h2: M=1, N=B, K=H, ldb=H + gemm_fwd_v: IqlGemmDesc, + /// dW3 = dv @ h2^T: M=1, N=H, K=B (TRANSA=N, TRANSB=T) + gemm_bwd_dw3: IqlGemmDesc, + /// dh2 = W3 @ dv: M=H, N=B, K=1 (TRANSA=N, TRANSB=N) + gemm_bwd_dh2: IqlGemmDesc, + /// dW2 = dh2_pre @ h1^T: M=H, N=H, K=B (TRANSA=N, TRANSB=T) + gemm_bwd_dw2: IqlGemmDesc, + /// dh1 = W2 @ dh2_pre: M=H, N=B, K=H (TRANSA=N, TRANSB=N) + gemm_bwd_dh1: IqlGemmDesc, + /// dW1 = dh1_pre @ states^T: M=H, N=state_dim, K=B (TRANSA=N, TRANSB=T) + gemm_bwd_dw1: IqlGemmDesc, + + // ── Element-wise cubin kernels ───────────────────────────────── + silu_fwd_kernel: CudaFunction, + silu_bwd_kernel: CudaFunction, + expectile_loss_kernel: CudaFunction, + bias_add_kernel: CudaFunction, + bias_grad_reduce_kernel: CudaFunction, loss_reduce_kernel: CudaFunction, grad_norm_phase1_kernel: CudaFunction, grad_norm_phase2_kernel: CudaFunction, adam_kernel: CudaFunction, + // Keep all non-training kernels unchanged: forward_kernel: CudaFunction, gather_q_taken_kernel: CudaFunction, advantage_weight_kernel: CudaFunction, @@ -135,67 +200,147 @@ pub struct GpuIqlTrainer { expectile_gap_kernel: CudaFunction, gap_mean_kernel: CudaFunction, per_sample_epsilon_kernel: CudaFunction, + support_floor_kernel: CudaFunction, + adv_sigma_ema_kernel: CudaFunction, // ── V network parameters (flat f32 on GPU) ───────────────────── + // Layout: W1[H*SD] + b1[H] + W2[H*H] + b2[H] + W3[H] + b3[1] params_buf: CudaSlice, // ── Adam optimizer state ──────────────────────────────────────── - m_buf: CudaSlice, // first moment - v_buf: CudaSlice, // second moment - grad_buf: CudaSlice, // reduced gradient [total_params] - grads_per_sample: CudaSlice, // per-sample gradients [B * total_params] - grad_norm_buf: CudaSlice, // [1] gradient L2 norm (sum-of-squares) - grad_norm_partials: CudaSlice, // [num_blocks] partial sums for phase1 + m_buf: CudaSlice, + v_buf: CudaSlice, + grad_buf: CudaSlice, // [total_params] reduced gradients + grad_norm_buf: CudaSlice, + grad_norm_partials: CudaSlice, - // ── Activation save buffers (forward -> backward) ─────────────── - save_pre1: CudaSlice, // [B, H] pre-activation layer 1 - save_pre2: CudaSlice, // [B, H] pre-activation layer 2 - save_h1: CudaSlice, // [B, H] post-activation layer 1 - save_h2: CudaSlice, // [B, H] post-activation layer 2 + // ── cuBLAS forward intermediate buffers ───────────────────────── + h1_pre_buf: CudaSlice, // [H, B] pre-activation layer 1 (GEMM output) + h1_buf: CudaSlice, // [H, B] post-SiLU layer 1 + h2_pre_buf: CudaSlice, // [H, B] pre-activation layer 2 + h2_buf: CudaSlice, // [H, B] post-SiLU layer 2 - // ── Output buffers ────────────────────────────────────────────── - v_out_buf: CudaSlice, // [B] V(s) predictions - loss_buf: CudaSlice, // [B] per-sample loss - total_loss_buf: CudaSlice, // [1] batch-mean loss - q_taken_buf: CudaSlice, // [B] Q(s, a_taken) gathered from q_out - advantage_weights_buf: CudaSlice, // [B] advantage weights + // ── cuBLAS backward intermediate buffers ──────────────────────── + dv_buf: CudaSlice, // [1, B] d_expectile_loss output + dh2_buf: CudaSlice, // [H, B] dh2 (pre silu_bwd) + dh2_pre_buf: CudaSlice, // [H, B] dh2_pre (post silu_bwd) + dh1_buf: CudaSlice, // [H, B] dh1 + dh1_pre_buf: CudaSlice, // [H, B] dh1_pre - // ── New integration buffers ───────────────────────────────────── - adv_stats_buf: CudaSlice, // [2] mean, variance - adv_sigma_ema_buf: CudaSlice, // [1] GPU-side EMA of advantage std - readiness_buf: CudaSlice, // [4]: readiness, cv_initial, cv_max, reserved - p5_state_buf: CudaSlice, // [2]: [0]=p5_estimate, [1]=step_count - support_floor_kernel: CudaFunction, - adv_sigma_ema_kernel: CudaFunction, - per_sample_support_buf: CudaSlice, // [B, 3] - branch_scales_buf: CudaSlice, // [B, 4] - expectile_gap_buf: CudaSlice, // [B] - gap_mean_buf: CudaSlice, // [1] - per_sample_epsilon_buf: CudaSlice, // [B] + // ── Output buffers (same as before) ──────────────────────────── + v_out_buf: CudaSlice, + loss_buf: CudaSlice, + total_loss_buf: CudaSlice, + q_taken_buf: CudaSlice, + advantage_weights_buf: CudaSlice, - // ── Training state ────────────────────────────────────────────── + // ── Integration buffers (unchanged) ──────────────────────────── + adv_stats_buf: CudaSlice, + adv_sigma_ema_buf: CudaSlice, + readiness_buf: CudaSlice, + p5_state_buf: CudaSlice, + per_sample_support_buf: CudaSlice, + branch_scales_buf: CudaSlice, + expectile_gap_buf: CudaSlice, + gap_mean_buf: CudaSlice, + per_sample_epsilon_buf: CudaSlice, + + // ── Training state ───────────────────────────────────────────── adam_step: i32, t_buf: CudaSlice, total_params: usize, grad_norm_blocks: usize, - grad_tile_size: usize, // per-sample grad tile (min(B, 256)) + #[allow(dead_code)] + state_dim_padded: usize, } impl GpuIqlTrainer { - /// Create a new GPU IQL trainer. + /// Create a new GPU IQL trainer with batched cuBLAS GEMMs. /// - /// Compiles 9 CUDA kernels, initializes V(s) weights with Xavier/Glorot, - /// and pre-allocates all GPU buffers including per-sample gradient storage. + /// Creates 8 cached GEMM descriptors, loads element-wise cubin kernels, + /// initializes V(s) weights with Xavier/Glorot, and pre-allocates all GPU + /// buffers including cuBLAS intermediate storage. pub fn new( - stream: Arc, + shared_handle: Arc, config: GpuIqlConfig, ) -> Result { + let stream = Arc::clone(&shared_handle.stream); let total_params = config.total_params(); let b = config.batch_size; let h = config.value_hidden_dim; + let sd = config.state_dim; + let state_dim_padded = (sd + 127) & !127; + let lt_handle = shared_handle.lt_handle.0; + let lt_ws_size = shared_handle.lt_workspace_size; - // Compile all 9 kernels - let kernels = compile_iql_kernels(&stream, &config)?; + // Create 3 forward GEMM descriptors + let gemm_fwd_h1 = create_iql_gemm_desc(lt_handle, h, b, sd, state_dim_padded, lt_ws_size, + 1, 0, "iql_fwd_h1")?; // TRANSA=T, TRANSB=N + let gemm_fwd_h2 = create_iql_gemm_desc(lt_handle, h, b, h, h, lt_ws_size, + 1, 0, "iql_fwd_h2")?; + let gemm_fwd_v = create_iql_gemm_desc(lt_handle, 1, b, h, h, lt_ws_size, + 1, 0, "iql_fwd_v")?; + + // Create 5 backward GEMM descriptors + // dW3 = dv[1,B] @ h2^T[B,H] -> [1,H] TRANSA=N, TRANSB=T + let gemm_bwd_dw3 = create_iql_gemm_desc_nt(lt_handle, 1, h, b, 1, h, lt_ws_size, "iql_bwd_dw3")?; + // dh2 = W3[H,1] @ dv[1,B] -> [H,B] TRANSA=N, TRANSB=N + let gemm_bwd_dh2 = create_iql_gemm_desc(lt_handle, h, b, 1, 1, lt_ws_size, + 0, 0, "iql_bwd_dh2")?; + // dW2 = dh2_pre[H,B] @ h1^T[B,H] -> [H,H] + let gemm_bwd_dw2 = create_iql_gemm_desc_nt(lt_handle, h, h, b, h, h, lt_ws_size, "iql_bwd_dw2")?; + // dh1 = W2[H,H] @ dh2_pre[H,B] -> [H,B] + let gemm_bwd_dh1 = create_iql_gemm_desc(lt_handle, h, b, h, h, lt_ws_size, + 0, 0, "iql_bwd_dh1")?; + // dW1 = dh1_pre[H,B] @ states^T[B,SD] -> [H,SD] + let gemm_bwd_dw1 = create_iql_gemm_desc_nt(lt_handle, h, sd, b, h, state_dim_padded, lt_ws_size, "iql_bwd_dw1")?; + + // Load element-wise cubin kernels + let iql_cubin = include_bytes!(concat!(env!("OUT_DIR"), "/iql_value_kernel.cubin")); + let context = stream.context(); + let module = context.load_cubin(iql_cubin.to_vec()) + .map_err(|e| MLError::ModelError(format!("iql cubin: {e}")))?; + + let load = |name: &str| -> Result { + module.load_function(name) + .map_err(|e| MLError::ModelError(format!("{name} load: {e}"))) + }; + + // New cuBLAS-era element-wise kernels + let silu_fwd_kernel = load("iql_silu_fwd")?; + let silu_bwd_kernel = load("iql_silu_bwd")?; + let expectile_loss_kernel = load("iql_expectile_loss")?; + let bias_add_kernel = load("iql_bias_add")?; + let bias_grad_reduce_kernel = load("iql_bias_grad_reduce")?; + + // Existing kernels (unchanged) + let loss_reduce_kernel = load("iql_loss_reduce")?; + let grad_norm_phase1_kernel = load("iql_grad_norm_phase1")?; + let grad_norm_phase2_kernel = load("iql_grad_norm_phase2")?; + let adam_kernel = load("iql_adam_kernel")?; + let forward_kernel = load("iql_forward_kernel")?; + let gather_q_taken_kernel = load("iql_gather_q_taken")?; + let advantage_weight_kernel = load("iql_compute_advantage_weights")?; + let modulate_td_kernel = load("iql_modulate_td_errors")?; + let adv_variance_kernel = load("iql_adv_variance_reduce")?; + let per_sample_support_kernel = load("iql_compute_per_sample_support")?; + let branch_advantage_kernel = load("iql_per_branch_advantage")?; + let expectile_gap_kernel = load("iql_expectile_gap")?; + let gap_mean_kernel = load("iql_gap_mean_reduce")?; + let per_sample_epsilon_kernel = load("iql_compute_per_sample_epsilon")?; + let adv_sigma_ema_kernel = load("iql_adv_sigma_ema_update")?; + let support_floor_kernel = load("iql_support_floor")?; + + // Allocate cuBLAS intermediate buffers (col-major [out_dim, B]) + let h1_pre_buf = alloc_f32(&stream, h * b, "iql_h1_pre")?; + let h1_buf = alloc_f32(&stream, h * b, "iql_h1")?; + let h2_pre_buf = alloc_f32(&stream, h * b, "iql_h2_pre")?; + let h2_buf = alloc_f32(&stream, h * b, "iql_h2")?; + let dv_buf = alloc_f32(&stream, b, "iql_dv")?; + let dh2_buf = alloc_f32(&stream, h * b, "iql_dh2")?; + let dh2_pre_buf = alloc_f32(&stream, h * b, "iql_dh2_pre")?; + let dh1_buf = alloc_f32(&stream, h * b, "iql_dh1")?; + let dh1_pre_buf = alloc_f32(&stream, h * b, "iql_dh1_pre")?; // Allocate parameter buffer and initialize with Xavier/Glorot let params_buf = init_xavier_weights(&stream, &config)?; @@ -204,22 +349,12 @@ impl GpuIqlTrainer { let m_buf = alloc_f32(&stream, total_params, "iql_m")?; let v_buf = alloc_f32(&stream, total_params, "iql_v")?; let grad_buf = alloc_f32(&stream, total_params, "iql_grad")?; - // Tiled per-sample gradients: process TILE samples at a time instead of all B. - // At B=16384, P=27009: full buffer = 1.7GB. Tiled at 256: 27MB. 64x smaller. - let grad_tile_size = b.min(256); - let grads_per_sample = alloc_f32(&stream, grad_tile_size * total_params, "iql_grads_tile")?; // Grad norm buffers (two-phase) let grad_norm_blocks = (total_params + 255) / 256; let grad_norm_buf = alloc_f32(&stream, 1, "iql_grad_norm")?; let grad_norm_partials = alloc_f32(&stream, grad_norm_blocks, "iql_grad_norm_partials")?; - // Allocate activation save buffers - let save_pre1 = alloc_f32(&stream, b * h, "iql_save_pre1")?; - let save_pre2 = alloc_f32(&stream, b * h, "iql_save_pre2")?; - let save_h1 = alloc_f32(&stream, b * h, "iql_save_h1")?; - let save_h2 = alloc_f32(&stream, b * h, "iql_save_h2")?; - // Allocate output buffers let v_out_buf = alloc_f32(&stream, b, "iql_v_out")?; let loss_buf = alloc_f32(&stream, b, "iql_loss")?; @@ -227,11 +362,11 @@ impl GpuIqlTrainer { let q_taken_buf = alloc_f32(&stream, b, "iql_q_taken")?; let advantage_weights_buf = alloc_f32(&stream, b, "iql_adv_weights")?; - // New integration buffers + // Integration buffers let adv_stats_buf = alloc_f32(&stream, 2, "iql_adv_stats")?; let adv_sigma_ema_buf = alloc_f32(&stream, 1, "iql_adv_sigma_ema")?; - let readiness_buf = alloc_f32(&stream, 4, "iql_readiness")?; // CV readiness + cv_max - let p5_state_buf = alloc_f32(&stream, 2, "iql_p5_state")?; // zeros: triggers first-batch init + let readiness_buf = alloc_f32(&stream, 4, "iql_readiness")?; + let p5_state_buf = alloc_f32(&stream, 2, "iql_p5_state")?; let mut per_sample_support_buf = alloc_f32(&stream, b * 3, "iql_per_sample_support")?; let branch_scales_buf = alloc_f32(&stream, b * 4, "iql_branch_scales")?; let expectile_gap_buf = alloc_f32(&stream, b, "iql_expectile_gap")?; @@ -249,11 +384,11 @@ impl GpuIqlTrainer { } super::htod_f32(&stream, &default_support, &mut per_sample_support_buf)?; - let per_sample_vram = grad_tile_size * total_params * 4; - let vram_bytes = total_params * 4 * 4 // params + m + v + grad - + per_sample_vram // per-sample gradients - + b * h * 4 * 4 // 4 activation buffers - + (b * 3 + 2) * 4; // v_out + loss + total_loss + adv_weights + grad_norm + let cublas_fwd_vram = h * b * 4 * 4 + b * 4; // h1_pre+h1+h2_pre+h2 + dv + let cublas_bwd_vram = h * b * 4 * 4; // dh2+dh2_pre+dh1+dh1_pre + let vram_bytes = total_params * 4 * 4 // params + m + v + grad + + cublas_fwd_vram + cublas_bwd_vram + + (b * 3 + 2) * 4; // v_out + loss + total_loss + adv_weights + grad_norm info!( state_dim = config.state_dim, @@ -262,9 +397,9 @@ impl GpuIqlTrainer { total_params, expectile_tau = config.expectile_tau, advantage_temperature = config.advantage_temperature, - per_sample_grad_mb = per_sample_vram / (1024 * 1024), + state_dim_padded, vram_kb = vram_bytes / 1024, - "GpuIqlTrainer initialized: 17 kernels, zero atomicAdd, V(s) Xavier-initialized" + "GpuIqlTrainer initialized: 8 cuBLAS GEMMs + element-wise kernels, zero atomicAdd" ); let t_buf = stream.alloc_zeros::(1) @@ -272,34 +407,51 @@ impl GpuIqlTrainer { Ok(Self { config, stream, - forward_loss_kernel: kernels.forward_loss, - backward_per_sample_kernel: kernels.backward_per_sample, - weight_grad_reduce_kernel: kernels.weight_grad_reduce, - loss_reduce_kernel: kernels.loss_reduce, - grad_norm_phase1_kernel: kernels.grad_norm_phase1, - grad_norm_phase2_kernel: kernels.grad_norm_phase2, - adam_kernel: kernels.adam, - forward_kernel: kernels.forward, - gather_q_taken_kernel: kernels.gather_q_taken, - advantage_weight_kernel: kernels.advantage_weight, - modulate_td_kernel: kernels.modulate_td, - adv_variance_kernel: kernels.adv_variance, - per_sample_support_kernel: kernels.per_sample_support, - branch_advantage_kernel: kernels.branch_advantage, - expectile_gap_kernel: kernels.expectile_gap, - gap_mean_kernel: kernels.gap_mean, - per_sample_epsilon_kernel: kernels.per_sample_epsilon, + shared_handle, + gemm_fwd_h1, + gemm_fwd_h2, + gemm_fwd_v, + gemm_bwd_dw3, + gemm_bwd_dh2, + gemm_bwd_dw2, + gemm_bwd_dh1, + gemm_bwd_dw1, + silu_fwd_kernel, + silu_bwd_kernel, + expectile_loss_kernel, + bias_add_kernel, + bias_grad_reduce_kernel, + loss_reduce_kernel, + grad_norm_phase1_kernel, + grad_norm_phase2_kernel, + adam_kernel, + forward_kernel, + gather_q_taken_kernel, + advantage_weight_kernel, + modulate_td_kernel, + adv_variance_kernel, + per_sample_support_kernel, + branch_advantage_kernel, + expectile_gap_kernel, + gap_mean_kernel, + per_sample_epsilon_kernel, + support_floor_kernel, + adv_sigma_ema_kernel, params_buf, m_buf, v_buf, grad_buf, - grads_per_sample, grad_norm_buf, grad_norm_partials, - save_pre1, - save_pre2, - save_h1, - save_h2, + h1_pre_buf, + h1_buf, + h2_pre_buf, + h2_buf, + dv_buf, + dh2_buf, + dh2_pre_buf, + dh1_buf, + dh1_pre_buf, v_out_buf, loss_buf, total_loss_buf, @@ -309,8 +461,6 @@ impl GpuIqlTrainer { adv_sigma_ema_buf, readiness_buf, p5_state_buf, - support_floor_kernel: kernels.support_floor, - adv_sigma_ema_kernel: kernels.adv_sigma_ema_update, per_sample_support_buf, branch_scales_buf, expectile_gap_buf, @@ -320,16 +470,17 @@ impl GpuIqlTrainer { t_buf, total_params, grad_norm_blocks, - grad_tile_size, + state_dim_padded, }) } /// Run one IQL value network training step — fully deterministic. /// - /// Pipeline: forward+loss → backward_per_sample → weight_grad_reduce - /// → loss_reduce → grad_norm phase1+2 → Adam. + /// Pipeline: 3 forward GEMMs + SiLU + bias_add + expectile_loss + /// -> 5 backward GEMMs + SiLU backward + bias_grad_reduce + /// -> loss_reduce -> grad_norm phase1+2 -> Adam. /// - /// Zero atomicAdd. All cross-sample reduction is deterministic. + /// Zero atomicAdd. All matrix multiplications use batched cuBLAS GEMMs. /// /// `gather_q_taken()` MUST be called first to populate `q_taken_buf`. /// The expectile regression target is Q(s, a_taken) from the DQN. @@ -338,107 +489,360 @@ impl GpuIqlTrainer { states_f32: &CudaSlice, ) -> Result<(), MLError> { let b = self.config.batch_size; - let state_dim_i32 = self.config.state_dim as i32; - let batch_size_i32 = b as i32; + let h = self.config.value_hidden_dim; + let sd = self.config.state_dim; let total_params_i32 = self.total_params as i32; + let batch_size_i32 = b as i32; let expectile_tau = self.config.expectile_tau; - // 1. Forward + loss kernel (256 threads per sample) - let fwd_shmem = 8 * std::mem::size_of::(); + let lt_handle = self.shared_handle.lt_handle.0; + let lt_ws_ptr = self.shared_handle.lt_workspace_ptr; + let lt_ws_size = self.shared_handle.lt_workspace_size; + let cu_stream = self.stream.cu_stream() as cublaslt_sys::cudaStream_t; + + let alpha: f32 = 1.0; + let beta: f32 = 0.0; + + // Compute param offsets (same layout as current flat buffer) + let w1_off = 0; + let b1_off = h * sd; + let w2_off = b1_off + h; + let b2_off = w2_off + h * h; + let w3_off = b2_off + h; + let b3_off = w3_off + h; + let f32_sz = std::mem::size_of::(); + + // Raw pointers into params_buf + let w1_ptr = self.params_buf.raw_ptr() + (w1_off * f32_sz) as u64; + let b1_ptr = self.params_buf.raw_ptr() + (b1_off * f32_sz) as u64; + let w2_ptr = self.params_buf.raw_ptr() + (w2_off * f32_sz) as u64; + let b2_ptr = self.params_buf.raw_ptr() + (b2_off * f32_sz) as u64; + let w3_ptr = self.params_buf.raw_ptr() + (w3_off * f32_sz) as u64; + let b3_ptr = self.params_buf.raw_ptr() + (b3_off * f32_sz) as u64; + + // ── Forward pass (3 GEMMs + 2 SiLU + bias_add) ── + + // 1. h1_pre[H,B] = W1^T @ states: cublasLtMatmul unsafe { - self.stream - .launch_builder(&self.forward_loss_kernel) - .arg(states_f32) - .arg(&self.q_taken_buf) - .arg(&self.params_buf) - .arg(&mut self.v_out_buf) - .arg(&mut self.loss_buf) - .arg(&mut self.save_pre1) - .arg(&mut self.save_pre2) - .arg(&mut self.save_h1) - .arg(&mut self.save_h2) - .arg(&batch_size_i32) - .arg(&state_dim_i32) - .arg(&expectile_tau) - .launch(LaunchConfig { - grid_dim: (b as u32, 1, 1), - block_dim: (256, 1, 1), - shared_mem_bytes: fwd_shmem as u32, - }) - .map_err(|e| MLError::ModelError(format!("IQL forward+loss: {e}")))?; + cublaslt_sys::cublasLtMatmul( + lt_handle, + self.gemm_fwd_h1.matmul_desc, + &alpha as *const f32 as *const std::ffi::c_void, + w1_ptr as *const std::ffi::c_void, // A = W1 + self.gemm_fwd_h1.a_layout, + states_f32.raw_ptr() as *const std::ffi::c_void, // B = states + self.gemm_fwd_h1.b_layout, + &beta as *const f32 as *const std::ffi::c_void, + self.h1_pre_buf.raw_ptr() as *mut std::ffi::c_void, // C + self.gemm_fwd_h1.c_layout, + self.h1_pre_buf.raw_ptr() as *mut std::ffi::c_void, // D = C + self.gemm_fwd_h1.d_layout, + &self.gemm_fwd_h1.algo as *const _, + lt_ws_ptr as *mut std::ffi::c_void, + lt_ws_size, + cu_stream, + ); + } + // bias_add: h1_pre += b1 (broadcast along batch dimension) + let n_h1 = (h * b) as i32; + let h_i32 = h as i32; + let blocks_h1 = ((h * b + 255) / 256) as u32; + unsafe { + self.stream.launch_builder(&self.bias_add_kernel) + .arg(&self.h1_pre_buf) // in-place + .arg(&b1_ptr) + .arg(&h_i32) // out_dim (repeating pattern) + .arg(&n_h1) // total elements + .launch(LaunchConfig { grid_dim: (blocks_h1, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 }) + .map_err(|e| MLError::ModelError(format!("IQL bias_add h1: {e}")))?; + } + // SiLU: h1 = silu(h1_pre), saves h1_pre for backward + unsafe { + self.stream.launch_builder(&self.silu_fwd_kernel) + .arg(&self.h1_pre_buf) + .arg(&mut self.h1_buf) + .arg(&n_h1) + .launch(LaunchConfig { grid_dim: (blocks_h1, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 }) + .map_err(|e| MLError::ModelError(format!("IQL silu_fwd h1: {e}")))?; } - // 2+3. Tiled backward + reduce: process TILE samples at a time. - // backward_per_sample writes to grads_tile [TILE * P] - // weight_grad_reduce sums TILE samples → accumulates into grad_buf [P] - // All launches are async on the same stream — no host sync between tiles. - let tile = self.grad_tile_size; - let reduce_blocks = (self.total_params + 255) / 256; - let grads_tile_ptr = self.grads_per_sample.raw_ptr(); - let grad_buf_ptr = self.grad_buf.raw_ptr(); + // 2. h2_pre[H,B] = W2^T @ h1 + unsafe { + cublaslt_sys::cublasLtMatmul( + lt_handle, + self.gemm_fwd_h2.matmul_desc, + &alpha as *const f32 as *const std::ffi::c_void, + w2_ptr as *const std::ffi::c_void, + self.gemm_fwd_h2.a_layout, + self.h1_buf.raw_ptr() as *const std::ffi::c_void, + self.gemm_fwd_h2.b_layout, + &beta as *const f32 as *const std::ffi::c_void, + self.h2_pre_buf.raw_ptr() as *mut std::ffi::c_void, + self.gemm_fwd_h2.c_layout, + self.h2_pre_buf.raw_ptr() as *mut std::ffi::c_void, + self.gemm_fwd_h2.d_layout, + &self.gemm_fwd_h2.algo as *const _, + lt_ws_ptr as *mut std::ffi::c_void, + lt_ws_size, + cu_stream, + ); + } + // bias_add + SiLU for layer 2 (same pattern as layer 1) + unsafe { + self.stream.launch_builder(&self.bias_add_kernel) + .arg(&self.h2_pre_buf) + .arg(&b2_ptr) + .arg(&h_i32) + .arg(&n_h1) + .launch(LaunchConfig { grid_dim: (blocks_h1, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 }) + .map_err(|e| MLError::ModelError(format!("IQL bias_add h2: {e}")))?; + } + unsafe { + self.stream.launch_builder(&self.silu_fwd_kernel) + .arg(&self.h2_pre_buf) + .arg(&mut self.h2_buf) + .arg(&n_h1) + .launch(LaunchConfig { grid_dim: (blocks_h1, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 }) + .map_err(|e| MLError::ModelError(format!("IQL silu_fwd h2: {e}")))?; + } - // Zero grad_buf before tiled accumulation + // 3. v_pre[1,B] = W3^T @ h2 -> v_out_buf + unsafe { + cublaslt_sys::cublasLtMatmul( + lt_handle, + self.gemm_fwd_v.matmul_desc, + &alpha as *const f32 as *const std::ffi::c_void, + w3_ptr as *const std::ffi::c_void, + self.gemm_fwd_v.a_layout, + self.h2_buf.raw_ptr() as *const std::ffi::c_void, + self.gemm_fwd_v.b_layout, + &beta as *const f32 as *const std::ffi::c_void, + self.v_out_buf.raw_ptr() as *mut std::ffi::c_void, + self.gemm_fwd_v.c_layout, + self.v_out_buf.raw_ptr() as *mut std::ffi::c_void, + self.gemm_fwd_v.d_layout, + &self.gemm_fwd_v.algo as *const _, + lt_ws_ptr as *mut std::ffi::c_void, + lt_ws_size, + cu_stream, + ); + } + // Add output bias (b3, scalar broadcast) + let n_v = b as i32; + let one_i32 = 1_i32; + let v_blocks = ((b + 255) / 256) as u32; + unsafe { + self.stream.launch_builder(&self.bias_add_kernel) + .arg(&self.v_out_buf) + .arg(&b3_ptr) + .arg(&one_i32) + .arg(&n_v) + .launch(LaunchConfig { grid_dim: (v_blocks, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 }) + .map_err(|e| MLError::ModelError(format!("IQL bias_add v: {e}")))?; + } + + // 4. Expectile loss + dv: combined kernel + // loss[B] = |tau - 1(u<0)| * u^2 where u = q_taken - v_out + // dv[B] = d(loss)/d(v_out) + unsafe { + self.stream.launch_builder(&self.expectile_loss_kernel) + .arg(&self.v_out_buf) + .arg(&self.q_taken_buf) + .arg(&expectile_tau) + .arg(&mut self.loss_buf) + .arg(&mut self.dv_buf) + .arg(&batch_size_i32) + .launch(LaunchConfig { grid_dim: (v_blocks, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 }) + .map_err(|e| MLError::ModelError(format!("IQL expectile_loss: {e}")))?; + } + + // ── Backward pass (5 GEMMs + 2 SiLU backward + bias grad reduce) ── + + // Zero grad_buf before accumulating weight gradients self.stream.memset_zeros(&mut self.grad_buf) .map_err(|e| MLError::ModelError(format!("IQL zero grad_buf: {e}")))?; - for tile_start in (0..b).step_by(tile) { - let tile_b = (b - tile_start).min(tile); - let tile_b_i32 = tile_b as i32; - let f32_sz = std::mem::size_of::(); + // inv_batch for 1/N mean reduction + let inv_batch: f32 = 1.0 / b as f32; + let beta_zero: f32 = 0.0; - // Backward: process tile_b samples starting at tile_start - // The kernel reads from states/q_taken/v_out/saves at [tile_start..] offsets. - // We pass offset pointers so the kernel sees sample indices 0..tile_b. - let states_off = states_f32.raw_ptr() + (tile_start * self.config.state_dim * f32_sz) as u64; - let qt_off = self.q_taken_buf.raw_ptr() + (tile_start * f32_sz) as u64; - let vo_off = self.v_out_buf.raw_ptr() + (tile_start * f32_sz) as u64; - let sp1_off = self.save_pre1.raw_ptr() + (tile_start * self.config.value_hidden_dim * f32_sz) as u64; - let sp2_off = self.save_pre2.raw_ptr() + (tile_start * self.config.value_hidden_dim * f32_sz) as u64; - let sh1_off = self.save_h1.raw_ptr() + (tile_start * self.config.value_hidden_dim * f32_sz) as u64; - let sh2_off = self.save_h2.raw_ptr() + (tile_start * self.config.value_hidden_dim * f32_sz) as u64; - - unsafe { - self.stream - .launch_builder(&self.backward_per_sample_kernel) - .arg(&states_off) - .arg(&qt_off) - .arg(&vo_off) - .arg(&self.params_buf) - .arg(&sp1_off) - .arg(&sp2_off) - .arg(&sh1_off) - .arg(&sh2_off) - .arg(&grads_tile_ptr) - .arg(&tile_b_i32) // grid size (samples in this tile) - .arg(&state_dim_i32) - .arg(&total_params_i32) - .arg(&expectile_tau) - .arg(&batch_size_i32) // FULL batch_size for 1/N mean reduction - .launch(LaunchConfig { - grid_dim: (tile_b as u32, 1, 1), - block_dim: (256, 1, 1), - shared_mem_bytes: 0, - }) - .map_err(|e| MLError::ModelError(format!("IQL backward tile {tile_start}: {e}")))?; - } - - // Reduce tile_b samples into grad_buf (accumulates via +=) - unsafe { - self.stream - .launch_builder(&self.weight_grad_reduce_kernel) - .arg(&grads_tile_ptr) - .arg(&grad_buf_ptr) - .arg(&tile_b_i32) - .arg(&total_params_i32) - .launch(LaunchConfig { - grid_dim: (reduce_blocks as u32, 1, 1), - block_dim: (256, 1, 1), - shared_mem_bytes: 0, - }) - .map_err(|e| MLError::ModelError(format!("IQL reduce tile {tile_start}: {e}")))?; - } + // dW3[1,H] = dv[1,B] @ h2^T[B,H] (alpha=1/B for mean reduction) + let dw3_ptr = self.grad_buf.raw_ptr() + (w3_off * f32_sz) as u64; + unsafe { + cublaslt_sys::cublasLtMatmul( + lt_handle, + self.gemm_bwd_dw3.matmul_desc, + &inv_batch as *const f32 as *const std::ffi::c_void, + self.dv_buf.raw_ptr() as *const std::ffi::c_void, + self.gemm_bwd_dw3.a_layout, + self.h2_buf.raw_ptr() as *const std::ffi::c_void, + self.gemm_bwd_dw3.b_layout, + &beta_zero as *const f32 as *const std::ffi::c_void, + dw3_ptr as *mut std::ffi::c_void, + self.gemm_bwd_dw3.c_layout, + dw3_ptr as *mut std::ffi::c_void, + self.gemm_bwd_dw3.d_layout, + &self.gemm_bwd_dw3.algo as *const _, + lt_ws_ptr as *mut std::ffi::c_void, + lt_ws_size, + cu_stream, + ); } + // db3: sum(dv) / B -- scalar bias gradient via bias_grad_reduce + let db3_ptr = self.grad_buf.raw_ptr() + (b3_off * f32_sz) as u64; + unsafe { + self.stream.launch_builder(&self.bias_grad_reduce_kernel) + .arg(&self.dv_buf) + .arg(&db3_ptr) + .arg(&one_i32) // out_dim = 1 (scalar bias) + .arg(&batch_size_i32) + .arg(&inv_batch) + .launch(LaunchConfig { grid_dim: (1, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 }) + .map_err(|e| MLError::ModelError(format!("IQL bias_grad_reduce b3: {e}")))?; + } + + // dh2[H,B] = W3[H,1] @ dv[1,B] + unsafe { + cublaslt_sys::cublasLtMatmul( + lt_handle, + self.gemm_bwd_dh2.matmul_desc, + &alpha as *const f32 as *const std::ffi::c_void, + w3_ptr as *const std::ffi::c_void, + self.gemm_bwd_dh2.a_layout, + self.dv_buf.raw_ptr() as *const std::ffi::c_void, + self.gemm_bwd_dh2.b_layout, + &beta_zero as *const f32 as *const std::ffi::c_void, + self.dh2_buf.raw_ptr() as *mut std::ffi::c_void, + self.gemm_bwd_dh2.c_layout, + self.dh2_buf.raw_ptr() as *mut std::ffi::c_void, + self.gemm_bwd_dh2.d_layout, + &self.gemm_bwd_dh2.algo as *const _, + lt_ws_ptr as *mut std::ffi::c_void, + lt_ws_size, + cu_stream, + ); + } + + // dh2_pre = dh2 * d_silu(h2_pre) -- element-wise + unsafe { + self.stream.launch_builder(&self.silu_bwd_kernel) + .arg(&self.h2_pre_buf) // pre-activation (saved from forward) + .arg(&self.dh2_buf) // upstream gradient + .arg(&mut self.dh2_pre_buf) // output + .arg(&n_h1) + .launch(LaunchConfig { grid_dim: (blocks_h1, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 }) + .map_err(|e| MLError::ModelError(format!("IQL silu_bwd h2: {e}")))?; + } + + // dW2[H,H] = dh2_pre[H,B] @ h1^T[B,H] (alpha=1/B) + let dw2_ptr = self.grad_buf.raw_ptr() + (w2_off * f32_sz) as u64; + unsafe { + cublaslt_sys::cublasLtMatmul( + lt_handle, + self.gemm_bwd_dw2.matmul_desc, + &inv_batch as *const f32 as *const std::ffi::c_void, + self.dh2_pre_buf.raw_ptr() as *const std::ffi::c_void, + self.gemm_bwd_dw2.a_layout, + self.h1_buf.raw_ptr() as *const std::ffi::c_void, + self.gemm_bwd_dw2.b_layout, + &beta_zero as *const f32 as *const std::ffi::c_void, + dw2_ptr as *mut std::ffi::c_void, + self.gemm_bwd_dw2.c_layout, + dw2_ptr as *mut std::ffi::c_void, + self.gemm_bwd_dw2.d_layout, + &self.gemm_bwd_dw2.algo as *const _, + lt_ws_ptr as *mut std::ffi::c_void, + lt_ws_size, + cu_stream, + ); + } + + // db2: sum columns of dh2_pre -> db2[H] + let db2_ptr = self.grad_buf.raw_ptr() + (b2_off * f32_sz) as u64; + let bias_blocks_h = ((h + 255) / 256) as u32; + unsafe { + self.stream.launch_builder(&self.bias_grad_reduce_kernel) + .arg(&self.dh2_pre_buf) + .arg(&db2_ptr) + .arg(&h_i32) // out_dim = H + .arg(&batch_size_i32) + .arg(&inv_batch) + .launch(LaunchConfig { grid_dim: (bias_blocks_h, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 }) + .map_err(|e| MLError::ModelError(format!("IQL bias_grad_reduce b2: {e}")))?; + } + + // dh1[H,B] = W2[H,H] @ dh2_pre[H,B] + unsafe { + cublaslt_sys::cublasLtMatmul( + lt_handle, + self.gemm_bwd_dh1.matmul_desc, + &alpha as *const f32 as *const std::ffi::c_void, + w2_ptr as *const std::ffi::c_void, + self.gemm_bwd_dh1.a_layout, + self.dh2_pre_buf.raw_ptr() as *const std::ffi::c_void, + self.gemm_bwd_dh1.b_layout, + &beta_zero as *const f32 as *const std::ffi::c_void, + self.dh1_buf.raw_ptr() as *mut std::ffi::c_void, + self.gemm_bwd_dh1.c_layout, + self.dh1_buf.raw_ptr() as *mut std::ffi::c_void, + self.gemm_bwd_dh1.d_layout, + &self.gemm_bwd_dh1.algo as *const _, + lt_ws_ptr as *mut std::ffi::c_void, + lt_ws_size, + cu_stream, + ); + } + + // dh1_pre = dh1 * d_silu(h1_pre) + unsafe { + self.stream.launch_builder(&self.silu_bwd_kernel) + .arg(&self.h1_pre_buf) + .arg(&self.dh1_buf) + .arg(&mut self.dh1_pre_buf) + .arg(&n_h1) + .launch(LaunchConfig { grid_dim: (blocks_h1, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 }) + .map_err(|e| MLError::ModelError(format!("IQL silu_bwd h1: {e}")))?; + } + + // dW1[H,SD] = dh1_pre[H,B] @ states^T[B,SD] (alpha=1/B) + let dw1_ptr = self.grad_buf.raw_ptr() + (w1_off * f32_sz) as u64; + unsafe { + cublaslt_sys::cublasLtMatmul( + lt_handle, + self.gemm_bwd_dw1.matmul_desc, + &inv_batch as *const f32 as *const std::ffi::c_void, + self.dh1_pre_buf.raw_ptr() as *const std::ffi::c_void, + self.gemm_bwd_dw1.a_layout, + states_f32.raw_ptr() as *const std::ffi::c_void, + self.gemm_bwd_dw1.b_layout, + &beta_zero as *const f32 as *const std::ffi::c_void, + dw1_ptr as *mut std::ffi::c_void, + self.gemm_bwd_dw1.c_layout, + dw1_ptr as *mut std::ffi::c_void, + self.gemm_bwd_dw1.d_layout, + &self.gemm_bwd_dw1.algo as *const _, + lt_ws_ptr as *mut std::ffi::c_void, + lt_ws_size, + cu_stream, + ); + } + + // db1: sum columns of dh1_pre -> db1[H] + let db1_ptr = self.grad_buf.raw_ptr() + (b1_off * f32_sz) as u64; + unsafe { + self.stream.launch_builder(&self.bias_grad_reduce_kernel) + .arg(&self.dh1_pre_buf) + .arg(&db1_ptr) + .arg(&h_i32) // out_dim = H + .arg(&batch_size_i32) + .arg(&inv_batch) + .launch(LaunchConfig { grid_dim: (bias_blocks_h, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 }) + .map_err(|e| MLError::ModelError(format!("IQL bias_grad_reduce b1: {e}")))?; + } + + // ── Loss reduce + grad norm + Adam (unchanged) ── + // 4. Loss reduce (deterministic sequential sum) unsafe { self.stream @@ -486,7 +890,7 @@ impl GpuIqlTrainer { .map_err(|e| MLError::ModelError(format!("IQL grad_norm_phase2: {e}")))?; } - // 7. Adam update — adam_step on GPU (async HtoD) + // 7. Adam update -- adam_step on GPU (async HtoD) self.adam_step += 1; unsafe { cudarc::driver::sys::cuMemcpyHtoDAsync_v2( @@ -652,7 +1056,7 @@ impl GpuIqlTrainer { self.adam_step += 1; } - // ── New integration launchers ──────────────────────────────────── + // ── Integration launchers ──────────────────────────────────────── /// Modulate td_errors in-place with advantage weights and staleness decay. pub fn modulate_td_errors( @@ -689,7 +1093,7 @@ impl GpuIqlTrainer { Ok(()) } - /// Update the EMA of advantage standard deviation entirely on GPU — zero DtoH. + /// Update the EMA of advantage standard deviation entirely on GPU -- zero DtoH. pub fn update_adv_sigma(&mut self) -> Result<(), MLError> { let batch_i32 = self.config.batch_size as i32; @@ -708,7 +1112,7 @@ impl GpuIqlTrainer { .map_err(|e| MLError::ModelError(format!("IQL adv_variance_reduce: {e}")))?; } - // Step 2: GPU-side EMA update — reads adv_stats[1], updates adv_sigma_ema_buf[0] + // Step 2: GPU-side EMA update -- reads adv_stats[1], updates adv_sigma_ema_buf[0] let ema_beta = 0.99_f32; unsafe { self.stream @@ -910,88 +1314,6 @@ impl GpuIqlTrainer { } } -// --------------------------------------------------------------------------- -// Compiled kernel set -// --------------------------------------------------------------------------- - -struct IqlKernels { - forward_loss: CudaFunction, - backward_per_sample: CudaFunction, - weight_grad_reduce: CudaFunction, - loss_reduce: CudaFunction, - grad_norm_phase1: CudaFunction, - grad_norm_phase2: CudaFunction, - adam: CudaFunction, - forward: CudaFunction, - gather_q_taken: CudaFunction, - advantage_weight: CudaFunction, - modulate_td: CudaFunction, - adv_variance: CudaFunction, - per_sample_support: CudaFunction, - branch_advantage: CudaFunction, - expectile_gap: CudaFunction, - gap_mean: CudaFunction, - per_sample_epsilon: CudaFunction, - adv_sigma_ema_update: CudaFunction, - support_floor: CudaFunction, -} - -// --------------------------------------------------------------------------- -// Kernel compilation -// --------------------------------------------------------------------------- - -/// Precompiled IQL value kernel cubin, embedded at compile time by build.rs. -static IQL_VALUE_CUBIN: &[u8] = include_bytes!(concat!(env!("OUT_DIR"), "/iql_value_kernel.cubin")); - -/// Load all 17 IQL CUDA kernels from precompiled cubin. -fn compile_iql_kernels( - stream: &Arc, - config: &GpuIqlConfig, -) -> Result { - info!( - state_dim = config.state_dim, - hidden_dim = config.value_hidden_dim, - expectile_tau = config.expectile_tau, - total_params = config.total_params(), - "GpuIqlTrainer: loading 17 deterministic IQL kernels (zero atomicAdd)" - ); - - let context = stream.context(); - let module = context.load_cubin(IQL_VALUE_CUBIN.to_vec()).map_err(|e| { - MLError::ModelError(format!("iql_value module load: {e}")) - })?; - - let load = |name: &str| -> Result { - module.load_function(name) - .map_err(|e| MLError::ModelError(format!("{name} load: {e}"))) - }; - - let kernels = IqlKernels { - forward_loss: load("iql_forward_loss_kernel")?, - backward_per_sample: load("iql_backward_per_sample")?, - weight_grad_reduce: load("iql_weight_grad_reduce")?, - loss_reduce: load("iql_loss_reduce")?, - grad_norm_phase1: load("iql_grad_norm_phase1")?, - grad_norm_phase2: load("iql_grad_norm_phase2")?, - adam: load("iql_adam_kernel")?, - forward: load("iql_forward_kernel")?, - gather_q_taken: load("iql_gather_q_taken")?, - advantage_weight: load("iql_compute_advantage_weights")?, - modulate_td: load("iql_modulate_td_errors")?, - adv_variance: load("iql_adv_variance_reduce")?, - per_sample_support: load("iql_compute_per_sample_support")?, - branch_advantage: load("iql_per_branch_advantage")?, - expectile_gap: load("iql_expectile_gap")?, - gap_mean: load("iql_gap_mean_reduce")?, - per_sample_epsilon: load("iql_compute_per_sample_epsilon")?, - adv_sigma_ema_update: load("iql_adv_sigma_ema_update")?, - support_floor: load("iql_support_floor")?, - }; - - info!("GpuIqlTrainer: 17 kernels loaded"); - Ok(kernels) -} - // --------------------------------------------------------------------------- // Weight initialization // --------------------------------------------------------------------------- @@ -1048,6 +1370,92 @@ fn init_xavier_weights( Ok(params_buf) } +// --------------------------------------------------------------------------- +// GEMM descriptor creation +// --------------------------------------------------------------------------- + +/// Create a GEMM descriptor: C = alpha * op(A) @ op(B) + beta * C +/// transa/transb: 0=N, 1=T +fn create_iql_gemm_desc( + lt_handle: cublaslt_sys::cublasLtHandle_t, + m: usize, n: usize, k: usize, ldb: usize, ws_size: usize, + transa: i32, transb: i32, label: &str, +) -> Result { + 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!("{label} MatmulDescCreate: {e:?}")))?; + + cublaslt_result::set_matmul_desc_attribute( + matmul_desc, + cublaslt_sys::cublasLtMatmulDescAttributes_t::CUBLASLT_MATMUL_DESC_TRANSA, + &transa as *const i32 as *const std::ffi::c_void, + std::mem::size_of::(), + ).map_err(|e| MLError::ModelError(format!("{label} set TRANSA: {e:?}")))?; + + cublaslt_result::set_matmul_desc_attribute( + matmul_desc, + cublaslt_sys::cublasLtMatmulDescAttributes_t::CUBLASLT_MATMUL_DESC_TRANSB, + &transb as *const i32 as *const std::ffi::c_void, + std::mem::size_of::(), + ).map_err(|e| MLError::ModelError(format!("{label} set TRANSB: {e:?}")))?; + + // Physical A layout: [K, M] col-major when TRANSA=T, [M, K] when TRANSA=N + let (a_rows, a_cols, a_ld) = if transa == 1 { (k, m, k) } else { (m, k, m) }; + let a_layout = cublaslt_result::create_matrix_layout( + f32_type, a_rows as u64, a_cols as u64, a_ld as i64, + ).map_err(|e| MLError::ModelError(format!("{label} A layout: {e:?}")))?; + + // Physical B layout: [K, N] col-major when TRANSB=N, [N, K] when TRANSB=T + let (b_rows, b_cols, b_ld) = if transb == 1 { (n, k, ldb) } else { (k, n, ldb) }; + let b_layout = cublaslt_result::create_matrix_layout( + f32_type, b_rows as u64, b_cols as u64, b_ld as i64, + ).map_err(|e| MLError::ModelError(format!("{label} B layout: {e:?}")))?; + + let c_layout = cublaslt_result::create_matrix_layout( + f32_type, m as u64, n as u64, m as i64, + ).map_err(|e| MLError::ModelError(format!("{label} C layout: {e:?}")))?; + + let d_layout = cublaslt_result::create_matrix_layout( + f32_type, m as u64, n as u64, m as i64, + ).map_err(|e| MLError::ModelError(format!("{label} D layout: {e:?}")))?; + + let matmul_pref = cublaslt_result::create_matmul_pref() + .map_err(|e| MLError::ModelError(format!("{label} MatmulPrefCreate: {e:?}")))?; + cublaslt_result::set_matmul_pref_attribute( + matmul_pref, + cublaslt_sys::cublasLtMatmulPreferenceAttributes_t::CUBLASLT_MATMUL_PREF_MAX_WORKSPACE_BYTES, + &ws_size as *const usize as *const std::ffi::c_void, + std::mem::size_of::(), + ).map_err(|e| { + let _ = cublaslt_result::destroy_matmul_pref(matmul_pref); + MLError::ModelError(format!("{label} set pref ws: {e:?}")) + })?; + + let heuristic = cublaslt_result::get_matmul_algo_heuristic( + lt_handle, matmul_desc, a_layout, b_layout, c_layout, d_layout, matmul_pref, + ); + let _ = cublaslt_result::destroy_matmul_pref(matmul_pref); + let algo = heuristic.map_err(|e| MLError::ModelError(format!("{label} heuristic: {e:?}")))?.algo; + + tracing::info!(%label, m, n, k, ldb, transa, transb, "IQL GEMM desc created"); + Ok(IqlGemmDesc { matmul_desc, a_layout, b_layout, c_layout, d_layout, algo }) + } +} + +/// Create GEMM desc for dW = A[M,K] @ B^T[K,N] where B has ldb leading dim. +/// Wrapper for TRANSA=N, TRANSB=T (used in all dW backward GEMMs). +fn create_iql_gemm_desc_nt( + lt_handle: cublaslt_sys::cublasLtHandle_t, + m: usize, n: usize, k: usize, _lda: usize, ldb: usize, ws_size: usize, label: &str, +) -> Result { + // dW = A @ B^T: TRANSA=N, TRANSB=T + // A's ld is M, set automatically in create_iql_gemm_desc (m, k, m) branch when transa=0. + create_iql_gemm_desc(lt_handle, m, n, k, ldb, ws_size, 0, 1, label) +} + // --------------------------------------------------------------------------- // Helpers // --------------------------------------------------------------------------- diff --git a/crates/ml/src/trainers/dqn/fused_training.rs b/crates/ml/src/trainers/dqn/fused_training.rs index 575c3a7c4..f22dc8180 100644 --- a/crates/ml/src/trainers/dqn/fused_training.rs +++ b/crates/ml/src/trainers/dqn/fused_training.rs @@ -486,14 +486,15 @@ impl FusedTrainingCtx { gamma: hyperparams.gamma as f32, ..GpuIqlConfig::default() }; - let gpu_iql = GpuIqlTrainer::new(stream.clone(), iql_config.clone()) + let shared_cublas_for_iql = Arc::clone(trainer.shared_cublas()); + let gpu_iql = GpuIqlTrainer::new(Arc::clone(&shared_cublas_for_iql), iql_config.clone()) .map_err(|e| anyhow::anyhow!("GPU IQL (high-tau) init: {e}"))?; let iql_low_config = GpuIqlConfig { expectile_tau: 1.0 - hyperparams.iql_expectile_tau, ..iql_config }; - let gpu_iql_low = GpuIqlTrainer::new(stream.clone(), iql_low_config) + let gpu_iql_low = GpuIqlTrainer::new(shared_cublas_for_iql, iql_low_config) .map_err(|e| anyhow::anyhow!("GPU IQL (low-tau) init: {e}"))?; // Initialize GPU IQN dual-head when iqn_lambda > 0 (CVaR risk sizing)