perf: fused RELU_BIAS epilogue for target, collector, and ensemble forward paths

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) <noreply@anthropic.com>
This commit is contained in:
jgrusewski
2026-04-18 22:40:46 +02:00
parent 76cfc0d185
commit 4b6d9d2e82

View File

@@ -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.