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:
@@ -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.
|
||||
|
||||
Reference in New Issue
Block a user