diff --git a/crates/ml/src/cuda_pipeline/batched_forward.rs b/crates/ml/src/cuda_pipeline/batched_forward.rs index 882691d8a..507bc80db 100644 --- a/crates/ml/src/cuda_pipeline/batched_forward.rs +++ b/crates/ml/src/cuda_pipeline/batched_forward.rs @@ -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; 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::()) 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::()) 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::()) 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::()) 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::()) 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::()) as u64; + } } + Ok(()) }