diff --git a/crates/ml/src/cuda_pipeline/dqn_utility_kernels.cu b/crates/ml/src/cuda_pipeline/dqn_utility_kernels.cu index 6b1ae5655..355990683 100644 --- a/crates/ml/src/cuda_pipeline/dqn_utility_kernels.cu +++ b/crates/ml/src/cuda_pipeline/dqn_utility_kernels.cu @@ -411,45 +411,49 @@ extern "C" __global__ void dqn_relu_mask_kernel( } /* ══════════════════════════════════════════════════════════════════════ - * SPECTRAL NORM POWER ITERATION KERNEL + * BATCHED SPECTRAL NORM — all weight matrices in a single launch * - * One step of power iteration for spectral normalization: - * v_new = W^T u / ||W^T u|| - * u_new = W v_new / ||W v_new|| - * sigma = u_new^T (W v_new) + * Each block processes one weight matrix independently. 10 blocks = 10 + * matrices, all running in parallel on H100's 132 SMs. * - * Then scales the weight matrix: W[i] /= max(1.0, sigma / sigma_max) + * Descriptor layout (per matrix, 6 values): + * [0] = W pointer (device), [1] = u pointer, [2] = v pointer, + * [3] = out_dim, [4] = in_dim, [5] = (unused, alignment pad) * - * For small matrices (256x256), this runs efficiently with a single block. - * Uses shared memory for the matmul + reduction. - * - * Launch: grid=(1,1,1), block=(256,1,1) + * Launch: grid=(num_matrices,1,1), block=(256,1,1), shmem=2048 * ══════════════════════════════════════════════════════════════════════ */ -extern "C" __global__ void spectral_norm_kernel( - __nv_bfloat16* __restrict__ W, /* [out_dim, in_dim] weight matrix — scaled in-place */ - __nv_bfloat16* __restrict__ u, /* [out_dim] left singular vector (persistent) */ - __nv_bfloat16* __restrict__ v, /* [in_dim] right singular vector (persistent) */ - int out_dim, - int in_dim, - float sigma_max /* clip sigma above this (typically 1.0) */ -) { - __shared__ __nv_bfloat16 shmem[256]; /* scratch for reductions */ +extern "C" __global__ void spectral_norm_batched( + const unsigned long long* __restrict__ descriptors, /* [num_matrices, 6] */ + float sigma_max, + int num_matrices) +{ + int mat_idx = blockIdx.x; + if (mat_idx >= num_matrices) return; + + /* Unpack descriptor for this matrix */ + const unsigned long long* desc = &descriptors[mat_idx * 6]; + __nv_bfloat16* W = (__nv_bfloat16*)desc[0]; + __nv_bfloat16* u = (__nv_bfloat16*)desc[1]; + __nv_bfloat16* v = (__nv_bfloat16*)desc[2]; + int out_dim = (int)desc[3]; + int in_dim = (int)desc[4]; + + __shared__ __nv_bfloat16 shmem[256]; int tid = threadIdx.x; - int bd = blockDim.x; /* typically 256 */ + int bd = blockDim.x; int n_total = out_dim * in_dim; - /* ── Step 1: v_new = W^T u ── (strided for dims > blockDim) */ - /* Each thread handles multiple v elements via stride loop. */ + /* ── Step 1: v_new = W^T u ── */ for (int col = tid; col < in_dim; col += bd) { __nv_bfloat16 val = bf16_zero(); for (int row = 0; row < out_dim; row++) val = val + W[row * in_dim + col] * u[row]; - v[col] = val; /* unnormalized — normalize below */ + v[col] = val; } __syncthreads(); - /* ── Normalize v_new: ||v||₂ via parallel reduction ── */ + /* ── Normalize v_new ── */ __nv_bfloat16 local_v2 = bf16_zero(); for (int col = tid; col < in_dim; col += bd) { __nv_bfloat16 vc = v[col]; local_v2 = local_v2 + vc * vc; } @@ -464,16 +468,16 @@ extern "C" __global__ void spectral_norm_kernel( v[col] = v[col] / v_norm; __syncthreads(); - /* ── Step 2: u_new = W v_new ── (strided for dims > blockDim) */ + /* ── Step 2: u_new = W v_new ── */ for (int row = tid; row < out_dim; row += bd) { __nv_bfloat16 val = bf16_zero(); for (int col = 0; col < in_dim; col++) val = val + W[row * in_dim + col] * v[col]; - u[row] = val; /* unnormalized — sigma = ||u_new|| */ + u[row] = val; } __syncthreads(); - /* ── Sigma = ||W v_new|| = ||u_new|| (before normalizing) ── */ + /* ── Sigma = ||u_new|| ── */ __nv_bfloat16 local_u2 = bf16_zero(); for (int row = tid; row < out_dim; row += bd) { __nv_bfloat16 ur = u[row]; local_u2 = local_u2 + ur * ur; } diff --git a/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs b/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs index 22ee4800c..6bc0cd2a4 100644 --- a/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs +++ b/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs @@ -501,7 +501,11 @@ pub struct GpuDqnTrainer { regime_scale_kernel: CudaFunction, shrink_perturb_kernel: CudaFunction, relu_mask_kernel: CudaFunction, - spectral_norm_kernel: CudaFunction, + spectral_norm_batched_kernel: CudaFunction, + /// Descriptor buffer for batched spectral norm: 10 matrices x 6 u64 = 60 elements. + /// Each entry: [W_ptr, u_ptr, v_ptr, out_dim, in_dim, pad=0]. + /// Built once at construction (pointers into params_bf16 are stable). + spectral_norm_descriptors: CudaSlice, clipped_saxpy_kernel: CudaFunction, clip_grad_kernel: CudaFunction, pad_states_kernel: CudaFunction, @@ -1858,71 +1862,26 @@ impl GpuDqnTrainer { self.params_initialized = true; } - let sh1 = self.config.shared_h1 as i32; - let _eg = EventTrackingGuard::new(self.stream.context()); - let sh2 = self.config.shared_h2 as i32; - let sd = self.config.state_dim as i32; - let sigma_max = self.config.spectral_norm_sigma_max; - let _evt_guard = EventTrackingGuard::new(self.stream.context()); + let sigma_max = self.config.spectral_norm_sigma_max; + let num_matrices = 10_i32; - // Compute raw u64 pointers into params_bf16 at GOFF_* offsets. - // These point directly into the flat bf16 shadow buffer — no separate - // allocations, no D2D sync needed after kernel writes. - let param_sizes = compute_param_sizes(&self.config); - let w_ptrs = bf16_weight_ptrs_from_base(self.ptrs.params_buf, ¶m_sizes); - - // ── Macro for launching spectral norm on a params_bf16 offset ──── - macro_rules! spec_norm { - ($goff_idx:expr, $u_slice:expr, $v_slice:expr, $out_dim:expr, $in_dim:expr, $label:literal) => {{ - let w_ptr = w_ptrs[$goff_idx]; - let u_ptr = $u_slice.raw_ptr(); - let v_ptr = $v_slice.raw_ptr(); - unsafe { - self.stream - .launch_builder(&self.spectral_norm_kernel) - .arg(&w_ptr).arg(&u_ptr).arg(&v_ptr) - .arg(&$out_dim).arg(&$in_dim).arg(&sigma_max) - .launch(LaunchConfig { - grid_dim: (1, 1, 1), - block_dim: (256, 1, 1), - shared_mem_bytes: 512 * 4, - }) - .map_err(|e| MLError::ModelError(format!("spectral_norm {}: {e}", $label)))?; - } - }}; + // Single batched launch: 10 blocks × 256 threads, one matrix per block. + // Descriptor buffer was built at construction with stable pointers. + unsafe { + self.stream + .launch_builder(&self.spectral_norm_batched_kernel) + .arg(&self.spectral_norm_descriptors) + .arg(&sigma_max) + .arg(&num_matrices) + .launch(LaunchConfig { + grid_dim: (num_matrices as u32, 1, 1), + block_dim: (256, 1, 1), + shared_mem_bytes: 0, // shmem[256] is static in the kernel + }) + .map_err(|e| MLError::ModelError(format!("spectral_norm_batched: {e}")))?; } - // W_s1 [shared_h1, state_dim] — GOFF index 0 - spec_norm!(0, self.spec_u_s1, self.spec_v_s1, sh1, sd, "W_s1"); - // W_s2 [shared_h2, shared_h1] — GOFF index 2 - spec_norm!(2, self.spec_u_s2, self.spec_v_s2, sh2, sh1, "W_s2"); - - // ── Head spectral norm kernel launches ─────────────────────────── - let vh = self.config.value_h as i32; - let ah = self.config.adv_h as i32; - let na = self.config.num_atoms as i32; - let b0 = self.config.branch_0_size as i32; - let b1 = self.config.branch_1_size as i32; - let b2 = self.config.branch_2_size as i32; - - // W_v1 [value_h, shared_h2] — GOFF index 4 - spec_norm!(4, self.spec_u_v1, self.spec_v_v1, vh, sh2, "W_v1"); - // W_v2 [num_atoms, value_h] — GOFF index 6 - spec_norm!(6, self.spec_u_v2, self.spec_v_v2, na, vh, "W_v2"); - // W_a1 [adv_h, shared_h2] — GOFF index 8 - spec_norm!(8, self.spec_u_a1, self.spec_v_a1, ah, sh2, "W_a1"); - // W_a2 [b0*num_atoms, adv_h] — GOFF index 10 - spec_norm!(10, self.spec_u_a2, self.spec_v_a2, b0 * na, ah, "W_a2"); - // W_bo1 [adv_h, shared_h2] — GOFF index 12 - spec_norm!(12, self.spec_u_bo1, self.spec_v_bo1, ah, sh2, "W_bo1"); - // W_bo2 [b1*num_atoms, adv_h] — GOFF index 14 - spec_norm!(14, self.spec_u_bo2, self.spec_v_bo2, b1 * na, ah, "W_bo2"); - // W_bu1 [adv_h, shared_h2] — GOFF index 16 - spec_norm!(16, self.spec_u_bu1, self.spec_v_bu1, ah, sh2, "W_bu1"); - // W_bu2 [b2*num_atoms, adv_h] — GOFF index 18 - spec_norm!(18, self.spec_u_bu2, self.spec_v_bu2, b2 * na, ah, "W_bu2"); - // No sync_w copies needed — spectral norm wrote directly to params_bf16. // Sync bf16 shadow → f32 master so Adam sees spectrally normalized weights. @@ -2433,6 +2392,50 @@ impl GpuDqnTrainer { // W_bu2 [b2*num_atoms, adv_h] alloc_spec_pair!(spec_u_bu2, spec_v_bu2, b2 * na, ah, "spec_u_bu2", "spec_v_bu2"); + // ── Build batched spectral norm descriptor buffer (10 matrices × 6 u64) ── + // Pointers are stable: params_bf16 base and spec_u/v allocations never move. + let spectral_norm_descriptors = { + let param_sizes = compute_param_sizes(&config); + let w_ptrs = bf16_weight_ptrs_from_base(params_bf16.raw_ptr(), ¶m_sizes); + let sd_u64 = config.state_dim as u64; + let sh1_u64 = config.shared_h1 as u64; + let sh2_u64 = config.shared_h2 as u64; + let vh_u64 = config.value_h as u64; + let ah_u64 = config.adv_h as u64; + let na_u64 = config.num_atoms as u64; + let b0_u64 = config.branch_0_size as u64; + let b1_u64 = config.branch_1_size as u64; + let b2_u64 = config.branch_2_size as u64; + + // 10 matrices, 6 u64 each: [W_ptr, u_ptr, v_ptr, out_dim, in_dim, pad] + let host_desc: [u64; 60] = [ + // [0] W_s1 [shared_h1, state_dim] — GOFF 0 + w_ptrs[0], spec_u_s1.raw_ptr(), spec_v_s1.raw_ptr(), sh1_u64, sd_u64, 0, + // [1] W_s2 [shared_h2, shared_h1] — GOFF 2 + w_ptrs[2], spec_u_s2.raw_ptr(), spec_v_s2.raw_ptr(), sh2_u64, sh1_u64, 0, + // [2] W_v1 [value_h, shared_h2] — GOFF 4 + w_ptrs[4], spec_u_v1.raw_ptr(), spec_v_v1.raw_ptr(), vh_u64, sh2_u64, 0, + // [3] W_v2 [num_atoms, value_h] — GOFF 6 + w_ptrs[6], spec_u_v2.raw_ptr(), spec_v_v2.raw_ptr(), na_u64, vh_u64, 0, + // [4] W_a1 [adv_h, shared_h2] — GOFF 8 + w_ptrs[8], spec_u_a1.raw_ptr(), spec_v_a1.raw_ptr(), ah_u64, sh2_u64, 0, + // [5] W_a2 [b0*num_atoms, adv_h] — GOFF 10 + w_ptrs[10], spec_u_a2.raw_ptr(), spec_v_a2.raw_ptr(), b0_u64 * na_u64, ah_u64, 0, + // [6] W_bo1 [adv_h, shared_h2] — GOFF 12 + w_ptrs[12], spec_u_bo1.raw_ptr(), spec_v_bo1.raw_ptr(), ah_u64, sh2_u64, 0, + // [7] W_bo2 [b1*num_atoms, adv_h] — GOFF 14 + w_ptrs[14], spec_u_bo2.raw_ptr(), spec_v_bo2.raw_ptr(), b1_u64 * na_u64, ah_u64, 0, + // [8] W_bu1 [adv_h, shared_h2] — GOFF 16 + w_ptrs[16], spec_u_bu1.raw_ptr(), spec_v_bu1.raw_ptr(), ah_u64, sh2_u64, 0, + // [9] W_bu2 [b2*num_atoms, adv_h] — GOFF 18 + w_ptrs[18], spec_u_bu2.raw_ptr(), spec_v_bu2.raw_ptr(), b2_u64 * na_u64, ah_u64, 0, + ]; + let mut desc_buf = stream.alloc_zeros::(60) + .map_err(|e| MLError::ModelError(format!("alloc spectral_norm_descriptors: {e}")))?; + stream.memcpy_htod(&host_desc, &mut desc_buf) + .map_err(|e| MLError::ModelError(format!("upload spectral_norm_descriptors: {e}")))?; + desc_buf + }; // ── Initialize cuBLAS forward context (required) ───────────── let cublas_forward = CublasForward::new( @@ -2667,7 +2670,8 @@ impl GpuDqnTrainer { regime_scale_kernel, shrink_perturb_kernel: shrink_perturb, relu_mask_kernel: relu_mask_standalone, - spectral_norm_kernel, + spectral_norm_batched_kernel: spectral_norm_kernel, + spectral_norm_descriptors, clipped_saxpy_kernel, clip_grad_kernel, pad_states_kernel, @@ -5796,10 +5800,10 @@ impl GpuDqnTrainer { /// /// The cubin is produced by build.rs from `common_device_functions.cuh` + /// `dqn_utility_kernels.cu`. Contains: grad_norm, adam_update, f32_to_bf16, -/// bf16_to_f32, saxpy, zero, regime_scale, shrink_perturb, relu_mask, spectral_norm, -/// clipped_saxpy, clip_grad. +/// bf16_to_f32, saxpy, zero, regime_scale, shrink_perturb, relu_mask, +/// spectral_norm_batched, clipped_saxpy, clip_grad. /// -/// Returns `(grad_norm, grad_norm_finalize, adam_update, f32_to_bf16, bf16_to_f32, saxpy, zero, regime_scale, shrink_perturb, relu_mask, spectral_norm, clipped_saxpy, clip_grad)`. +/// Returns `(grad_norm, grad_norm_finalize, adam_update, f32_to_bf16, bf16_to_f32, saxpy, zero, regime_scale, shrink_perturb, relu_mask, spectral_norm_batched, clipped_saxpy, clip_grad)`. fn compile_training_kernels( stream: &Arc, config: &GpuDqnTrainConfig, @@ -5839,8 +5843,8 @@ fn compile_training_kernels( .map_err(|e| MLError::ModelError(format!("dqn_shrink_perturb_kernel load: {e}")))?; let _relu_mask_from_module = module.load_function("dqn_relu_mask_kernel") .map_err(|e| MLError::ModelError(format!("dqn_relu_mask_kernel load: {e}")))?; - let spectral_norm = module.load_function("spectral_norm_kernel") - .map_err(|e| MLError::ModelError(format!("spectral_norm_kernel load: {e}")))?; + let spectral_norm = module.load_function("spectral_norm_batched") + .map_err(|e| MLError::ModelError(format!("spectral_norm_batched load: {e}")))?; let clipped_saxpy = module.load_function("dqn_clipped_saxpy_kernel") .map_err(|e| MLError::ModelError(format!("dqn_clipped_saxpy_kernel load: {e}")))?; let clip_grad = module.load_function("dqn_clip_grad_kernel")