From 4b6d9d2e82ba78ac4a0a239dd3de2cbaaa373908 Mon Sep 17 00:00:00 2001 From: jgrusewski Date: Sat, 18 Apr 2026 22:40:46 +0200 Subject: [PATCH] perf: fused RELU_BIAS epilogue for target, collector, and ensemble forward paths MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Adds try-fused-fallback-to-separate pattern to forward_target_raw, forward_online_f32, and forward_value_head — matching the existing pattern in forward_online_raw. Eliminates ~10-14 separate bias+ReLU kernel launches by fusing them into the preceding cuBLAS GEMM epilogue. Co-Authored-By: Claude Opus 4.6 (1M context) --- .../ml/src/cuda_pipeline/batched_forward.rs | 119 +++++++++++++----- 1 file changed, 85 insertions(+), 34 deletions(-) diff --git a/crates/ml/src/cuda_pipeline/batched_forward.rs b/crates/ml/src/cuda_pipeline/batched_forward.rs index ed15f529b..d99857ca7 100644 --- a/crates/ml/src/cuda_pipeline/batched_forward.rs +++ b/crates/ml/src/cuda_pipeline/batched_forward.rs @@ -627,18 +627,32 @@ impl CublasGemmSet { mag_concat_ptr: u64, ) -> Result<(), MLError> { let b = self.batch_size; + let ws = self.handle.lt_workspace_ptr; + let wss = self.handle.lt_workspace_size; - // All layers: pure F32 SGEMM + F32 bias+relu - self.sgemm_f32_ldb(stream, w_ptrs[0], states_f32_ptr, h_s1_ptr, self.shared_h1, b, self.s1_input_dim, self.s1_ldb, "f32_h_s1")?; - self.launch_add_bias_relu_f32_raw(stream, h_s1_ptr, w_ptrs[1], self.shared_h1, b)?; + // Trunk layers: fused GEMM+bias+ReLU epilogue when available + if self.sgemm_f32_fused_relu_bias(stream, w_ptrs[0], states_f32_ptr, h_s1_ptr, w_ptrs[1], + self.shared_h1, b, self.s1_input_dim, self.s1_ldb, ws, wss, "f32_h_s1").is_err() + { + self.sgemm_f32_ldb(stream, w_ptrs[0], states_f32_ptr, h_s1_ptr, self.shared_h1, b, self.s1_input_dim, self.s1_ldb, "f32_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, "f32_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, "f32_h_s2").is_err() + { + self.sgemm_f32(stream, w_ptrs[2], h_s1_ptr, h_s2_ptr, self.shared_h2, b, self.shared_h1, "f32_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, "f32_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, "f32_h_v").is_err() + { + self.sgemm_f32(stream, w_ptrs[4], h_s2_ptr, h_v_ptr, self.value_h, b, self.shared_h2, "f32_h_v")?; + self.launch_add_bias_relu_f32_raw(stream, h_v_ptr, w_ptrs[5], self.value_h, b)?; + } - // Value output: F32 cublasLtMatmul → F32 logits (same as f32 path output type) + // Value output: linear (no ReLU) self.sgemm_f32(stream, w_ptrs[6], h_v_ptr, v_logits_ptr, self.num_atoms, b, self.value_h, "f32_v_logits")?; self.launch_add_bias_f32_raw(stream, v_logits_ptr, w_ptrs[7], self.num_atoms, b)?; @@ -674,13 +688,19 @@ impl CublasGemmSet { true, "f32", )?; } else { + let bws = self.branch_workspace_ptrs[d]; + let bwss = self.handle.lt_workspace_size; let (fc_input, fc_k) = if d == 1 && mag_concat_ptr != 0 { (mag_concat_ptr, self.mag_concat_dim) } else { (h_s2_ptr, self.shared_h2) }; - self.sgemm_f32_branch(bs, w_ptrs[w_fc_idx], fc_input, branch_h_ptrs[d], self.adv_h, b, fc_k, d, "f32_h_bd")?; - self.launch_add_bias_relu_f32_raw(bs, branch_h_ptrs[d], w_ptrs[w_fc_idx + 1], self.adv_h, b)?; + if self.sgemm_f32_fused_relu_bias(bs, w_ptrs[w_fc_idx], fc_input, branch_h_ptrs[d], w_ptrs[w_fc_idx + 1], + self.adv_h, b, fc_k, fc_k, bws, bwss, "f32_h_bd").is_err() + { + self.sgemm_f32_branch(bs, w_ptrs[w_fc_idx], fc_input, branch_h_ptrs[d], self.adv_h, b, fc_k, d, "f32_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; @@ -716,8 +736,12 @@ impl CublasGemmSet { } else { (h_s2_ptr, self.shared_h2) }; - self.sgemm_f32(stream, w_ptrs[w_fc_idx], fc_input, branch_h_ptrs[d], self.adv_h, b, fc_k, "f32_h_bd")?; - self.launch_add_bias_relu_f32_raw(stream, branch_h_ptrs[d], w_ptrs[w_fc_idx + 1], self.adv_h, b)?; + if self.sgemm_f32_fused_relu_bias(stream, w_ptrs[w_fc_idx], fc_input, branch_h_ptrs[d], w_ptrs[w_fc_idx + 1], + self.adv_h, b, fc_k, fc_k, ws, wss, "f32_h_bd").is_err() + { + self.sgemm_f32(stream, w_ptrs[w_fc_idx], fc_input, branch_h_ptrs[d], self.adv_h, b, fc_k, "f32_h_bd")?; + self.launch_add_bias_relu_f32_raw(stream, branch_h_ptrs[d], w_ptrs[w_fc_idx + 1], self.adv_h, b)?; + } } let adv_out_ptr = b_logits_ptr + logit_byte_offset; @@ -785,14 +809,21 @@ impl CublasGemmSet { v_logits_out: u64, // [B, NA] output logits batch_size: usize, ) -> Result<(), MLError> { - // h_v = ReLU(h_s2 @ W_v1^T + b_v1) - self.sgemm_f32( - stream, w_v1, h_s2_ptr, h_v_scratch, - self.value_h, batch_size, self.shared_h2, "ens_h_v", - )?; - self.launch_add_bias_relu_f32_raw(stream, h_v_scratch, b_v1, self.value_h, batch_size)?; + let ws = self.handle.lt_workspace_ptr; + let wss = self.handle.lt_workspace_size; - // v_logits = h_v @ W_v2^T + b_v2 (f32 output — no f32 truncation) + // h_v = ReLU(h_s2 @ W_v1^T + b_v1) — fused epilogue when available + if self.sgemm_f32_fused_relu_bias(stream, w_v1, h_s2_ptr, h_v_scratch, b_v1, + self.value_h, batch_size, self.shared_h2, self.shared_h2, ws, wss, "ens_h_v").is_err() + { + self.sgemm_f32( + stream, w_v1, h_s2_ptr, h_v_scratch, + self.value_h, batch_size, self.shared_h2, "ens_h_v", + )?; + self.launch_add_bias_relu_f32_raw(stream, h_v_scratch, b_v1, self.value_h, batch_size)?; + } + + // v_logits = h_v @ W_v2^T + b_v2 (linear output — no ReLU) self.sgemm_f32( stream, w_v2, h_v_scratch, v_logits_out, self.num_atoms, batch_size, self.value_h, "ens_v_logits", @@ -832,23 +863,37 @@ impl CublasGemmSet { mag_concat_ptr: u64, ) -> Result<(), MLError> { let b = self.batch_size; + let ws = self.handle.lt_workspace_ptr; + let wss = self.handle.lt_workspace_size; - // h_s1[B, SH1] — first layer: ldb = state_dim_padded (CUTLASS K-tile alignment) - self.sgemm_f32_ldb(stream, tg_w_ptrs[0], states_ptr, h_s1_ptr, - self.shared_h1, b, self.s1_input_dim, self.s1_ldb, "tg_h_s1")?; - self.launch_add_bias_relu_f32_raw(stream, h_s1_ptr, tg_w_ptrs[1], self.shared_h1, b)?; + // h_s1[B, SH1] — fused GEMM+bias+ReLU epilogue, fall back to separate kernels + if self.sgemm_f32_fused_relu_bias(stream, tg_w_ptrs[0], states_ptr, h_s1_ptr, tg_w_ptrs[1], + self.shared_h1, b, self.s1_input_dim, self.s1_ldb, ws, wss, "tg_h_s1").is_err() + { + self.sgemm_f32_ldb(stream, tg_w_ptrs[0], states_ptr, h_s1_ptr, + self.shared_h1, b, self.s1_input_dim, self.s1_ldb, "tg_h_s1")?; + self.launch_add_bias_relu_f32_raw(stream, h_s1_ptr, tg_w_ptrs[1], self.shared_h1, b)?; + } // h_s2[B, SH2] - self.sgemm_f32(stream, tg_w_ptrs[2], h_s1_ptr, h_s2_ptr, - self.shared_h2, b, self.shared_h1, "tg_h_s2")?; - self.launch_add_bias_relu_f32_raw(stream, h_s2_ptr, tg_w_ptrs[3], self.shared_h2, b)?; + if self.sgemm_f32_fused_relu_bias(stream, tg_w_ptrs[2], h_s1_ptr, h_s2_ptr, tg_w_ptrs[3], + self.shared_h2, b, self.shared_h1, self.shared_h1, ws, wss, "tg_h_s2").is_err() + { + self.sgemm_f32(stream, tg_w_ptrs[2], h_s1_ptr, h_s2_ptr, + self.shared_h2, b, self.shared_h1, "tg_h_s2")?; + self.launch_add_bias_relu_f32_raw(stream, h_s2_ptr, tg_w_ptrs[3], self.shared_h2, b)?; + } // h_v[B, VH] - self.sgemm_f32(stream, tg_w_ptrs[4], h_s2_ptr, h_v_ptr, - self.value_h, b, self.shared_h2, "tg_h_v")?; - self.launch_add_bias_relu_f32_raw(stream, h_v_ptr, tg_w_ptrs[5], self.value_h, b)?; + if self.sgemm_f32_fused_relu_bias(stream, tg_w_ptrs[4], h_s2_ptr, h_v_ptr, tg_w_ptrs[5], + self.value_h, b, self.shared_h2, self.shared_h2, ws, wss, "tg_h_v").is_err() + { + self.sgemm_f32(stream, tg_w_ptrs[4], h_s2_ptr, h_v_ptr, + self.value_h, b, self.shared_h2, "tg_h_v")?; + self.launch_add_bias_relu_f32_raw(stream, h_v_ptr, tg_w_ptrs[5], self.value_h, b)?; + } - // v_logits[B, NA] — output layer: f32 cublasLtMatmul + f32 bias + // v_logits[B, NA] — output layer: no ReLU (linear output) self.sgemm_f32(stream, tg_w_ptrs[6], h_v_ptr, v_logits_ptr, self.num_atoms, b, self.value_h, "tg_v_logits")?; self.launch_add_bias_f32_raw(stream, v_logits_ptr, tg_w_ptrs[7], self.num_atoms, b)?; @@ -884,15 +929,21 @@ impl CublasGemmSet { true, "tg", )?; } else { - // Legacy GEMM+bias+ReLU path + // GEMM+bias+ReLU — try fused epilogue, fall back to separate kernels + let bws = self.branch_workspace_ptrs[d]; + let bwss = self.handle.lt_workspace_size; let (fc_input, fc_k) = if d == 1 && mag_concat_ptr != 0 { (mag_concat_ptr, self.mag_concat_dim) } else { (h_s2_ptr, self.shared_h2) }; - self.sgemm_f32_branch(bs, tg_w_ptrs[w_fc_idx], fc_input, branch_h_ptrs[d], - self.adv_h, b, fc_k, d, "tg_h_bd")?; - self.launch_add_bias_relu_f32_raw(bs, branch_h_ptrs[d], tg_w_ptrs[w_fc_idx + 1], self.adv_h, b)?; + if self.sgemm_f32_fused_relu_bias(bs, tg_w_ptrs[w_fc_idx], fc_input, branch_h_ptrs[d], tg_w_ptrs[w_fc_idx + 1], + self.adv_h, b, fc_k, fc_k, bws, bwss, "tg_h_bd").is_err() + { + self.sgemm_f32_branch(bs, tg_w_ptrs[w_fc_idx], fc_input, branch_h_ptrs[d], + self.adv_h, b, fc_k, d, "tg_h_bd")?; + self.launch_add_bias_relu_f32_raw(bs, branch_h_ptrs[d], tg_w_ptrs[w_fc_idx + 1], self.adv_h, b)?; + } } // Output layer: f32 cublasLtMatmul + f32 bias — per-branch workspace.