perf: multi-stream branch dispatch — 3 advantage heads in parallel
Fork 3 CUDA streams at CublasForward construction. In forward_online_raw and forward_online_f32, the trunk records an event, each branch stream waits on it, submits its GEMM+bias ops on its own stream (via cublasSetStream), then the main stream joins all three. This overlaps the 3 independent advantage head computations (exposure, order, urgency) that previously executed sequentially. Safety: a distinct_branches guard checks h_b0!=h_b1!=h_b2 at runtime; callers that alias branch hidden buffers (Pass 3 Double DQN scratch reuse) fall back to the sequential loop. CUDA Graph capture (CUDA 12+) captures the fork/join pattern as graph dependencies. Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
This commit is contained in:
@@ -126,6 +126,12 @@ pub struct CublasForward {
|
||||
branch_0_size: usize,
|
||||
branch_1_size: usize,
|
||||
branch_2_size: usize,
|
||||
|
||||
// ── Multi-stream branch dispatch ──
|
||||
/// 3 forked CUDA streams for parallel advantage branch execution.
|
||||
/// Each branch (exposure, order, urgency) submits GEMMs to its own stream,
|
||||
/// then the main stream joins all three before consuming the logits.
|
||||
branch_streams: [Arc<CudaStream>; 3],
|
||||
}
|
||||
|
||||
impl CublasForward {
|
||||
@@ -180,6 +186,13 @@ impl CublasForward {
|
||||
let (add_bias_relu_bf16_kernel, add_bias_bf16_kernel, add_bias_f32_kernel, add_bias_relu_f32_kernel, add_bias_f32_f32bias_kernel) =
|
||||
compile_bias_kernels(stream)?;
|
||||
|
||||
// ── Fork 3 branch streams for parallel advantage head dispatch ──
|
||||
let branch_streams = [
|
||||
stream.fork().map_err(|e| MLError::DeviceError(format!("fork branch stream 0: {e}")))?,
|
||||
stream.fork().map_err(|e| MLError::DeviceError(format!("fork branch stream 1: {e}")))?,
|
||||
stream.fork().map_err(|e| MLError::DeviceError(format!("fork branch stream 2: {e}")))?,
|
||||
];
|
||||
|
||||
Ok(Self {
|
||||
handle: SendSyncCublasHandle(raw_handle),
|
||||
_workspace_buf: workspace_buf,
|
||||
@@ -200,6 +213,7 @@ impl CublasForward {
|
||||
branch_0_size,
|
||||
branch_1_size,
|
||||
branch_2_size,
|
||||
branch_streams,
|
||||
})
|
||||
}
|
||||
|
||||
@@ -266,22 +280,87 @@ impl CublasForward {
|
||||
let branch_h_ptrs = [h_b0_ptr, h_b1_ptr, h_b2_ptr];
|
||||
let branch_w_base = [8_usize, 12, 16];
|
||||
let na = self.num_atoms;
|
||||
let mut logit_byte_offset: u64 = 0;
|
||||
for d in 0..3 {
|
||||
let n_d = branch_sizes[d];
|
||||
let w_fc_idx = branch_w_base[d];
|
||||
|
||||
// Hidden layer: stays bf16
|
||||
self.gemmex_bf16(w_ptrs[w_fc_idx], h_s2_ptr, branch_h_ptrs[d], self.adv_h, b, self.shared_h2, "h_bd")?;
|
||||
self.launch_add_bias_relu_bf16_raw(stream, branch_h_ptrs[d], w_ptrs[w_fc_idx + 1], self.adv_h, b)?;
|
||||
// Multi-stream only when branch hidden buffers are distinct (no aliasing).
|
||||
// Pass 3 (Double DQN on next_states) reuses a single scratch for all 3
|
||||
// branches — must stay sequential to avoid concurrent writes.
|
||||
let distinct_branches = h_b0_ptr != h_b1_ptr && h_b1_ptr != h_b2_ptr && h_b0_ptr != h_b2_ptr;
|
||||
|
||||
// Output layer: GemmEx writes f32 C-matrix
|
||||
let adv_out_ptr = b_logits_ptr + logit_byte_offset;
|
||||
self.gemmex_bf16_to_f32(w_ptrs[w_fc_idx + 2], branch_h_ptrs[d], adv_out_ptr, n_d * na, b, self.adv_h, "adv_logits")?;
|
||||
self.launch_add_bias_f32_raw(stream, adv_out_ptr, w_ptrs[w_fc_idx + 3], n_d * na, b)?;
|
||||
if distinct_branches {
|
||||
// Multi-stream branch dispatch: submit each branch's GEMMs to a
|
||||
// dedicated CUDA stream so all 3 advantage heads execute in parallel.
|
||||
// Safe because branches read h_s2 (shared trunk, read-only) and write
|
||||
// to non-overlapping regions of branch_h_ptrs and b_logits_buf.
|
||||
//
|
||||
// Works inside CUDA Graph capture (CUDA 12+): cuStreamWaitEvent
|
||||
// on a non-captured stream makes it join the capture graph, and the
|
||||
// fork/join pattern is recorded as graph dependencies.
|
||||
|
||||
logit_byte_offset += (b * n_d * na * std::mem::size_of::<f32>()) as u64;
|
||||
// Record trunk completion on the main stream.
|
||||
let trunk_done = stream.record_event(None)
|
||||
.map_err(|e| MLError::ModelError(format!("trunk event record: {e}")))?;
|
||||
|
||||
let mut logit_byte_offset: u64 = 0;
|
||||
for d in 0..3 {
|
||||
let bs = &self.branch_streams[d];
|
||||
let n_d = branch_sizes[d];
|
||||
let w_fc_idx = branch_w_base[d];
|
||||
|
||||
// Branch stream waits for trunk completion.
|
||||
bs.wait(&trunk_done)
|
||||
.map_err(|e| MLError::ModelError(format!("branch {d} wait trunk: {e}")))?;
|
||||
|
||||
// Redirect cuBLAS handle to this branch stream.
|
||||
unsafe {
|
||||
let cu_stream = bs.cu_stream() as *mut cublas_sys::CUstream_st;
|
||||
cublas_result::set_stream(self.handle.0, cu_stream)
|
||||
.map_err(|e| MLError::ModelError(format!("cublasSetStream branch {d}: {e:?}")))?;
|
||||
}
|
||||
|
||||
// Hidden layer: stays bf16
|
||||
self.gemmex_bf16(w_ptrs[w_fc_idx], h_s2_ptr, branch_h_ptrs[d], self.adv_h, b, self.shared_h2, "h_bd")?;
|
||||
self.launch_add_bias_relu_bf16_raw(bs, branch_h_ptrs[d], w_ptrs[w_fc_idx + 1], self.adv_h, b)?;
|
||||
|
||||
// Output layer: GemmEx writes f32 C-matrix
|
||||
let adv_out_ptr = b_logits_ptr + logit_byte_offset;
|
||||
self.gemmex_bf16_to_f32(w_ptrs[w_fc_idx + 2], branch_h_ptrs[d], adv_out_ptr, n_d * na, b, self.adv_h, "adv_logits")?;
|
||||
self.launch_add_bias_f32_raw(bs, adv_out_ptr, w_ptrs[w_fc_idx + 3], n_d * na, b)?;
|
||||
|
||||
logit_byte_offset += (b * n_d * na * std::mem::size_of::<f32>()) as u64;
|
||||
}
|
||||
|
||||
// Restore cuBLAS handle to the main stream.
|
||||
unsafe {
|
||||
let cu_stream = stream.cu_stream() as *mut cublas_sys::CUstream_st;
|
||||
cublas_result::set_stream(self.handle.0, cu_stream)
|
||||
.map_err(|e| MLError::ModelError(format!("cublasSetStream restore: {e:?}")))?;
|
||||
}
|
||||
|
||||
// Join: main stream waits for all 3 branches to complete.
|
||||
for d in 0..3 {
|
||||
let branch_done = self.branch_streams[d].record_event(None)
|
||||
.map_err(|e| MLError::ModelError(format!("branch {d} done event: {e}")))?;
|
||||
stream.wait(&branch_done)
|
||||
.map_err(|e| MLError::ModelError(format!("main wait branch {d}: {e}")))?;
|
||||
}
|
||||
} else {
|
||||
// Sequential fallback: branch hidden buffers alias (scratch reuse).
|
||||
let mut logit_byte_offset: u64 = 0;
|
||||
for d in 0..3 {
|
||||
let n_d = branch_sizes[d];
|
||||
let w_fc_idx = branch_w_base[d];
|
||||
|
||||
self.gemmex_bf16(w_ptrs[w_fc_idx], h_s2_ptr, branch_h_ptrs[d], self.adv_h, b, self.shared_h2, "h_bd")?;
|
||||
self.launch_add_bias_relu_bf16_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;
|
||||
self.gemmex_bf16_to_f32(w_ptrs[w_fc_idx + 2], branch_h_ptrs[d], adv_out_ptr, n_d * na, b, self.adv_h, "adv_logits")?;
|
||||
self.launch_add_bias_f32_raw(stream, adv_out_ptr, w_ptrs[w_fc_idx + 3], n_d * na, b)?;
|
||||
|
||||
logit_byte_offset += (b * n_d * na * std::mem::size_of::<f32>()) as u64;
|
||||
}
|
||||
}
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
@@ -321,20 +400,69 @@ impl CublasForward {
|
||||
let branch_h_ptrs = [h_b0_ptr, h_b1_ptr, h_b2_ptr];
|
||||
let branch_w_base = [8_usize, 12, 16];
|
||||
let na = self.num_atoms;
|
||||
let mut logit_byte_offset: u64 = 0;
|
||||
for d in 0..3 {
|
||||
let n_d = branch_sizes[d];
|
||||
let w_fc_idx = branch_w_base[d];
|
||||
|
||||
self.sgemm_f32(w_ptrs[w_fc_idx], h_s2_ptr, branch_h_ptrs[d], self.adv_h, b, self.shared_h2, "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 distinct_branches = h_b0_ptr != h_b1_ptr && h_b1_ptr != h_b2_ptr && h_b0_ptr != h_b2_ptr;
|
||||
|
||||
let adv_out_ptr = b_logits_ptr + logit_byte_offset;
|
||||
self.sgemm_f32(w_ptrs[w_fc_idx + 2], branch_h_ptrs[d], adv_out_ptr, n_d * na, b, self.adv_h, "f32_adv_logits")?;
|
||||
self.launch_add_bias_f32_raw(stream, adv_out_ptr, w_ptrs[w_fc_idx + 3], n_d * na, b)?;
|
||||
if distinct_branches {
|
||||
// Multi-stream branch dispatch (f32 path, experience collection).
|
||||
let trunk_done = stream.record_event(None)
|
||||
.map_err(|e| MLError::ModelError(format!("f32 trunk event record: {e}")))?;
|
||||
|
||||
logit_byte_offset += (b * n_d * na * std::mem::size_of::<f32>()) as u64;
|
||||
let mut logit_byte_offset: u64 = 0;
|
||||
for d in 0..3 {
|
||||
let bs = &self.branch_streams[d];
|
||||
let n_d = branch_sizes[d];
|
||||
let w_fc_idx = branch_w_base[d];
|
||||
|
||||
bs.wait(&trunk_done)
|
||||
.map_err(|e| MLError::ModelError(format!("f32 branch {d} wait trunk: {e}")))?;
|
||||
|
||||
unsafe {
|
||||
let cu_stream = bs.cu_stream() as *mut cublas_sys::CUstream_st;
|
||||
cublas_result::set_stream(self.handle.0, cu_stream)
|
||||
.map_err(|e| MLError::ModelError(format!("cublasSetStream f32 branch {d}: {e:?}")))?;
|
||||
}
|
||||
|
||||
self.sgemm_f32(w_ptrs[w_fc_idx], h_s2_ptr, branch_h_ptrs[d], self.adv_h, b, self.shared_h2, "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;
|
||||
self.sgemm_f32(w_ptrs[w_fc_idx + 2], branch_h_ptrs[d], adv_out_ptr, n_d * na, b, self.adv_h, "f32_adv_logits")?;
|
||||
self.launch_add_bias_f32_raw(bs, adv_out_ptr, w_ptrs[w_fc_idx + 3], n_d * na, b)?;
|
||||
|
||||
logit_byte_offset += (b * n_d * na * std::mem::size_of::<f32>()) as u64;
|
||||
}
|
||||
|
||||
unsafe {
|
||||
let cu_stream = stream.cu_stream() as *mut cublas_sys::CUstream_st;
|
||||
cublas_result::set_stream(self.handle.0, cu_stream)
|
||||
.map_err(|e| MLError::ModelError(format!("cublasSetStream f32 restore: {e:?}")))?;
|
||||
}
|
||||
|
||||
for d in 0..3 {
|
||||
let branch_done = self.branch_streams[d].record_event(None)
|
||||
.map_err(|e| MLError::ModelError(format!("f32 branch {d} done event: {e}")))?;
|
||||
stream.wait(&branch_done)
|
||||
.map_err(|e| MLError::ModelError(format!("f32 main wait branch {d}: {e}")))?;
|
||||
}
|
||||
} else {
|
||||
// Sequential fallback: branch hidden buffers alias.
|
||||
let mut logit_byte_offset: u64 = 0;
|
||||
for d in 0..3 {
|
||||
let n_d = branch_sizes[d];
|
||||
let w_fc_idx = branch_w_base[d];
|
||||
|
||||
self.sgemm_f32(w_ptrs[w_fc_idx], h_s2_ptr, branch_h_ptrs[d], self.adv_h, b, self.shared_h2, "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;
|
||||
self.sgemm_f32(w_ptrs[w_fc_idx + 2], branch_h_ptrs[d], adv_out_ptr, n_d * na, b, self.adv_h, "f32_adv_logits")?;
|
||||
self.launch_add_bias_f32_raw(stream, adv_out_ptr, w_ptrs[w_fc_idx + 3], n_d * na, b)?;
|
||||
|
||||
logit_byte_offset += (b * n_d * na * std::mem::size_of::<f32>()) as u64;
|
||||
}
|
||||
}
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user