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) <noreply@anthropic.com>
This commit is contained in:
jgrusewski
2026-04-11 10:57:57 +02:00
parent 7c34003f03
commit bdcc3b7d30

View File

@@ -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<FwdGemmKey, CachedGemmDesc>,
/// 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<FwdGemmKey, CachedGemmDesc>,
}
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<FwdGemmKey> = 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::<u64>(),
).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<CachedGemmDesc, 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!("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::<i32>(),
).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::<i32>(),
).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::<i32>(),
).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::<cublaslt_sys::cudaDataType_t>(),
).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::<usize>(),
).map_err(|e| MLError::ModelError(format!("cached+relu set pref ws: {e:?}")))?;
let mut heuristics: [std::mem::MaybeUninit<cublaslt_sys::cublasLtMatmulHeuristicResult_t>; 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.