From bdcc3b7d30fae0cd7eeec4387b4a322f759b3a70 Mon Sep 17 00:00:00 2001 From: jgrusewski Date: Sat, 11 Apr 2026 10:57:57 +0200 Subject: [PATCH] perf: cublasLt RELU_BIAS epilogue fusion for trunk + branch hidden GEMMs Fuses GEMM + bias-add + ReLU into a single cublasLtMatmul kernel via CUBLASLT_EPILOGUE_RELU_BIAS. Eliminates 7 separate add_bias_relu kernel launches per forward pass (3 trunk + 4 branch hidden layers). Creates separate cached descriptors with epilogue enabled at init time. Falls back to separate kernels if the epilogue heuristic isn't available. Bias pointer set dynamically per-call via set_matmul_desc_attribute. Co-Authored-By: Claude Opus 4.6 (1M context) --- .../ml/src/cuda_pipeline/batched_forward.rs | 243 +++++++++++++++++- 1 file changed, 234 insertions(+), 9 deletions(-) diff --git a/crates/ml/src/cuda_pipeline/batched_forward.rs b/crates/ml/src/cuda_pipeline/batched_forward.rs index ed266511f..06472e84e 100644 --- a/crates/ml/src/cuda_pipeline/batched_forward.rs +++ b/crates/ml/src/cuda_pipeline/batched_forward.rs @@ -221,6 +221,9 @@ pub struct CublasForward { /// Map from (n, batch, k, ldb) → pre-created descriptors + algo. /// All forward GEMMs share TRANSA=T, TRANSB=N. gemm_cache: HashMap, + /// Cached descriptors with RELU_BIAS epilogue (for hidden layers). + /// Bias pointer is set dynamically per-call. Same key as gemm_cache. + gemm_cache_relu_bias: HashMap, } impl CublasForward { @@ -371,6 +374,25 @@ impl CublasForward { gemm_cache.insert((n, b, k, ldb), desc); } + // Create RELU_BIAS epilogue variants for hidden layers. + // These fuse GEMM + bias-add + ReLU into one kernel, eliminating separate launches. + let relu_bias_shapes: Vec = vec![ + (shared_h1, batch_size, state_dim, state_dim_padded), // h_s1 + (shared_h2, batch_size, shared_h1, shared_h1), // h_s2 + (value_h, batch_size, shared_h2, shared_h2), // h_v + (adv_h, batch_size, shared_h2, shared_h2), // h_bd (×4) + ]; + let mut gemm_cache_relu_bias = HashMap::new(); + for &(n, b, k, ldb) in &relu_bias_shapes { + match create_cached_fwd_gemm_desc_relu_bias(lt_raw_handle, n, b, k, ldb, lt_ws_size) { + Ok(desc) => { gemm_cache_relu_bias.insert((n, b, k, ldb), desc); } + Err(e) => { + tracing::warn!("RELU_BIAS epilogue not available for ({n},{b},{k},{ldb}): {e}"); + // Fall back to separate kernels — no entry in cache + } + } + } + Ok(Self { handle: SendSyncCublasHandle(raw_handle), _workspace_buf: workspace_buf, @@ -401,6 +423,7 @@ impl CublasForward { trunk_done_event, branch_done_events, gemm_cache, + gemm_cache_relu_bias, }) } @@ -478,15 +501,29 @@ impl CublasForward { let b = self.batch_size; // First layer: ldb = state_dim_padded (CUTLASS K-tile alignment). - // States buffer is padded to [B, pad128(state_dim)] with zero columns. - self.sgemm_f32_ldb(stream, w_ptrs[0], states_ptr, h_s1_ptr, self.shared_h1, b, self.state_dim, self.state_dim_padded, "h_s1")?; - self.launch_add_bias_relu_f32_raw(stream, h_s1_ptr, w_ptrs[1], self.shared_h1, b)?; + // Try fused GEMM+bias+ReLU epilogue, fall back to separate kernels. + let ws = self.lt_workspace_ptr; + let wss = self.lt_workspace_size; + if self.sgemm_f32_fused_relu_bias(stream, w_ptrs[0], states_ptr, h_s1_ptr, w_ptrs[1], + self.shared_h1, b, self.state_dim, self.state_dim_padded, ws, wss, "h_s1").is_err() + { + self.sgemm_f32_ldb(stream, w_ptrs[0], states_ptr, h_s1_ptr, self.shared_h1, b, self.state_dim, self.state_dim_padded, "h_s1")?; + self.launch_add_bias_relu_f32_raw(stream, h_s1_ptr, w_ptrs[1], self.shared_h1, b)?; + } - self.sgemm_f32(stream, w_ptrs[2], h_s1_ptr, h_s2_ptr, self.shared_h2, b, self.shared_h1, "h_s2")?; - self.launch_add_bias_relu_f32_raw(stream, h_s2_ptr, w_ptrs[3], self.shared_h2, b)?; + if self.sgemm_f32_fused_relu_bias(stream, w_ptrs[2], h_s1_ptr, h_s2_ptr, w_ptrs[3], + self.shared_h2, b, self.shared_h1, self.shared_h1, ws, wss, "h_s2").is_err() + { + self.sgemm_f32(stream, w_ptrs[2], h_s1_ptr, h_s2_ptr, self.shared_h2, b, self.shared_h1, "h_s2")?; + self.launch_add_bias_relu_f32_raw(stream, h_s2_ptr, w_ptrs[3], self.shared_h2, b)?; + } - self.sgemm_f32(stream, w_ptrs[4], h_s2_ptr, h_v_ptr, self.value_h, b, self.shared_h2, "h_v")?; - self.launch_add_bias_relu_f32_raw(stream, h_v_ptr, w_ptrs[5], self.value_h, b)?; + if self.sgemm_f32_fused_relu_bias(stream, w_ptrs[4], h_s2_ptr, h_v_ptr, w_ptrs[5], + self.value_h, b, self.shared_h2, self.shared_h2, ws, wss, "h_v").is_err() + { + self.sgemm_f32(stream, w_ptrs[4], h_s2_ptr, h_v_ptr, self.value_h, b, self.shared_h2, "h_v")?; + self.launch_add_bias_relu_f32_raw(stream, h_v_ptr, w_ptrs[5], self.value_h, b)?; + } // Output layer: cublasLtMatmul writes f32 C-matrix (no f32 truncation overflow) self.sgemm_f32(stream, w_ptrs[6], h_v_ptr, v_logits_ptr, self.num_atoms, b, self.value_h, "v_logits")?; @@ -532,8 +569,15 @@ impl CublasForward { .map_err(|e| MLError::ModelError(format!("branch {d} wait trunk: {e}")))?; // Per-branch workspace: eliminates contention between parallel branch streams. - self.sgemm_f32_branch(bs, w_ptrs[w_fc_idx], h_s2_ptr, branch_h_ptrs[d], self.adv_h, b, self.shared_h2, d, "h_bd")?; - self.launch_add_bias_relu_f32_raw(bs, branch_h_ptrs[d], w_ptrs[w_fc_idx + 1], self.adv_h, b)?; + // Try fused GEMM+bias+ReLU epilogue, fall back to separate kernels. + let bws = self.branch_workspace_ptrs[d]; + let bwss = self.lt_workspace_size; + if self.sgemm_f32_fused_relu_bias(bs, w_ptrs[w_fc_idx], h_s2_ptr, branch_h_ptrs[d], w_ptrs[w_fc_idx + 1], + self.adv_h, b, self.shared_h2, self.shared_h2, bws, bwss, "h_bd").is_err() + { + self.sgemm_f32_branch(bs, w_ptrs[w_fc_idx], h_s2_ptr, branch_h_ptrs[d], self.adv_h, b, self.shared_h2, d, "h_bd")?; + self.launch_add_bias_relu_f32_raw(bs, branch_h_ptrs[d], w_ptrs[w_fc_idx + 1], self.adv_h, b)?; + } let adv_out_ptr = b_logits_ptr + logit_byte_offset; self.sgemm_f32_branch(bs, w_ptrs[w_fc_idx + 2], branch_h_ptrs[d], adv_out_ptr, n_d * na, b, self.adv_h, d, "adv_logits")?; @@ -1030,6 +1074,72 @@ impl CublasForward { Ok(()) } + /// Fused GEMM + bias + ReLU via cublasLt epilogue. + /// Uses cached descriptor with RELU_BIAS epilogue. Sets bias pointer per-call. + /// Falls back to separate GEMM + add_bias_relu if epilogue not available. + pub(crate) fn sgemm_f32_fused_relu_bias( + &self, + stream: &CudaStream, + w_ptr: u64, + input_ptr: u64, + output_ptr: u64, + bias_ptr: u64, + n: usize, // out_dim + b: usize, // batch + k: usize, // in_dim + ldb: usize, + ws_ptr: u64, + ws_size: usize, + label: &str, + ) -> Result<(), MLError> { + let key: FwdGemmKey = (n, b, k, ldb); + if let Some(cached) = self.gemm_cache_relu_bias.get(&key) { + // Set bias pointer on the descriptor (lightweight CPU write) + unsafe { + cublaslt_result::set_matmul_desc_attribute( + cached.matmul_desc, + cublaslt_sys::cublasLtMatmulDescAttributes_t::CUBLASLT_MATMUL_DESC_BIAS_POINTER, + &bias_ptr as *const u64 as *const std::ffi::c_void, + std::mem::size_of::(), + ).map_err(|e| MLError::ModelError(format!("set bias ptr {label}: {e:?}")))?; + } + + // Launch fused GEMM+bias+ReLU + let alpha = 1.0_f32; + let beta = 0.0_f32; + unsafe { + let cu_stream = stream.cu_stream() as cublaslt_sys::cudaStream_t; + let status = cublaslt_sys::cublasLtMatmul( + self.lt_handle.0, + cached.matmul_desc, + &alpha as *const f32 as *const std::ffi::c_void, + w_ptr as *const std::ffi::c_void, + cached.a_layout, + input_ptr as *const std::ffi::c_void, + cached.b_layout, + &beta as *const f32 as *const std::ffi::c_void, + output_ptr as *const std::ffi::c_void, + cached.c_layout, + output_ptr as *mut std::ffi::c_void, + cached.d_layout, + &cached.algo as *const cublaslt_sys::cublasLtMatmulAlgo_t, + ws_ptr as *mut std::ffi::c_void, + ws_size, + cu_stream, + ); + if status != cublaslt_sys::cublasStatus_t::CUBLAS_STATUS_SUCCESS { + let e = cublaslt_result::CublasError(status); + tracing::error!(m=n, n_batch=b, k=k, ?e, "cublasLtMatmul FUSED RELU_BIAS FAILED for {label}"); + return Err(MLError::ModelError(format!("cublasLtMatmul fused {label}: {e:?}"))); + } + } + Ok(()) + } else { + // No cached epilogue descriptor — caller should use separate GEMM + bias_relu + Err(MLError::ModelError(format!("no RELU_BIAS epilogue for {label} (n={n},b={b},k={k})"))) + } + } + /// Slow path: inline descriptor creation for uncached GEMM shapes. /// Used for diagnostic calls and any shape not pre-cached at init. #[allow(clippy::too_many_arguments)] @@ -1407,6 +1517,121 @@ fn create_cached_fwd_gemm_desc( } } +/// Create a cached GEMM descriptor with RELU_BIAS epilogue. +/// The bias pointer is set dynamically per-call via set_matmul_desc_attribute. +/// This fuses GEMM + bias-add + ReLU into a single cublasLtMatmul kernel, +/// eliminating separate add_bias_relu kernel launches. +fn create_cached_fwd_gemm_desc_relu_bias( + lt_handle: cublaslt_sys::cublasLtHandle_t, + n: usize, b: usize, k: usize, ldb: usize, ws_size: usize, +) -> 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!("cached+relu MatmulDescCreate (n={n},b={b},k={k}): {e:?}")))?; + + let transa: i32 = 1; // CUBLAS_OP_T + 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!("cached+relu set TRANSA: {e:?}")))?; + + let transb: i32 = 0; // CUBLAS_OP_N + 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!("cached+relu set TRANSB: {e:?}")))?; + + // Set RELU_BIAS epilogue — fuses bias add + ReLU into the matmul kernel + let epilogue: i32 = cublaslt_sys::cublasLtEpilogue_t::CUBLASLT_EPILOGUE_RELU_BIAS as i32; + cublaslt_result::set_matmul_desc_attribute( + matmul_desc, + cublaslt_sys::cublasLtMatmulDescAttributes_t::CUBLASLT_MATMUL_DESC_EPILOGUE, + &epilogue as *const i32 as *const std::ffi::c_void, + std::mem::size_of::(), + ).map_err(|e| MLError::ModelError(format!("cached+relu set EPILOGUE: {e:?}")))?; + + // Set bias data type to F32 + let bias_type = cublaslt_sys::cudaDataType_t::CUDA_R_32F; + cublaslt_result::set_matmul_desc_attribute( + matmul_desc, + cublaslt_sys::cublasLtMatmulDescAttributes_t::CUBLASLT_MATMUL_DESC_BIAS_DATA_TYPE, + &bias_type as *const _ as *const std::ffi::c_void, + std::mem::size_of::(), + ).map_err(|e| MLError::ModelError(format!("cached+relu set BIAS_DATA_TYPE: {e:?}")))?; + + let a_layout = cublaslt_result::create_matrix_layout(f32_type, k as u64, n as u64, k as i64) + .map_err(|e| MLError::ModelError(format!("cached+relu A layout: {e:?}")))?; + let b_layout = cublaslt_result::create_matrix_layout(f32_type, k as u64, b as u64, ldb as i64) + .map_err(|e| MLError::ModelError(format!("cached+relu B layout: {e:?}")))?; + let c_layout = cublaslt_result::create_matrix_layout(f32_type, n as u64, b as u64, n as i64) + .map_err(|e| MLError::ModelError(format!("cached+relu C layout: {e:?}")))?; + let d_layout = cublaslt_result::create_matrix_layout(f32_type, n as u64, b as u64, n as i64) + .map_err(|e| MLError::ModelError(format!("cached+relu D layout: {e:?}")))?; + + let matmul_pref = cublaslt_result::create_matmul_pref() + .map_err(|e| MLError::ModelError(format!("cached+relu pref: {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| MLError::ModelError(format!("cached+relu set pref ws: {e:?}")))?; + + let mut heuristics: [std::mem::MaybeUninit; 3] = + std::mem::MaybeUninit::uninit().assume_init(); + let mut algo_count: i32 = 0; + + let status = cublaslt_sys::cublasLtMatmulAlgoGetHeuristic( + lt_handle, matmul_desc, a_layout, b_layout, c_layout, d_layout, + matmul_pref, 3, heuristics[0].as_mut_ptr(), &mut algo_count, + ); + let _ = cublaslt_result::destroy_matmul_pref(matmul_pref); + + if status != cublaslt_sys::cublasStatus_t::CUBLAS_STATUS_SUCCESS || algo_count == 0 { + let _ = cublaslt_result::destroy_matrix_layout(d_layout); + let _ = cublaslt_result::destroy_matrix_layout(c_layout); + let _ = cublaslt_result::destroy_matrix_layout(b_layout); + let _ = cublaslt_result::destroy_matrix_layout(a_layout); + let _ = cublaslt_result::destroy_matmul_desc(matmul_desc); + tracing::warn!(n, b, k, ldb, "RELU_BIAS epilogue heuristic failed — falling back to separate kernels"); + return Err(MLError::ModelError(format!( + "cached+relu algo heuristic (n={n},b={b},k={k}): status={status:?}, count={algo_count}" + ))); + } + + let best = heuristics[0].assume_init(); + if best.state != cublaslt_sys::cublasStatus_t::CUBLAS_STATUS_SUCCESS { + let _ = cublaslt_result::destroy_matrix_layout(d_layout); + let _ = cublaslt_result::destroy_matrix_layout(c_layout); + let _ = cublaslt_result::destroy_matrix_layout(b_layout); + let _ = cublaslt_result::destroy_matrix_layout(a_layout); + let _ = cublaslt_result::destroy_matmul_desc(matmul_desc); + tracing::warn!(n, b, k, ldb, "RELU_BIAS epilogue algo invalid — falling back"); + return Err(MLError::ModelError(format!( + "cached+relu algo state invalid (n={n},b={b},k={k}): {:?}", best.state + ))); + } + + tracing::info!( + n, b, k, ldb, algo_count, + ws_needed = best.workspaceSize, + "cached fwd GEMM+RELU_BIAS desc created" + ); + + Ok(CachedGemmDesc { + matmul_desc, a_layout, b_layout, c_layout, d_layout, + algo: best.algo, + }) + } +} + // ── Compute BF16 weight pointers from flat params_buf ─────────────────────── /// Compute the 20 raw BF16 device pointers into a flat params_buf at GOFF_* offsets.