From 32151ba8bca1d25aa25325623223634976f79cdc Mon Sep 17 00:00:00 2001 From: jgrusewski Date: Sat, 18 Apr 2026 16:57:34 +0200 Subject: [PATCH] =?UTF-8?q?fix:=20backward=20branches=20reuse=20forward=20?= =?UTF-8?q?workspace=20=E2=80=94=20saves=20128MB=20(OOM=20fix)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit The backward branch parallelism allocated 4 × 32MB = 128MB for per-branch cuBLAS workspaces, pushing H100 VRAM over 80GB → OOM. Forward and backward are in the same child graph (sequential) — workspaces never conflict. Now backward receives forward's branch_workspace_ptrs as a parameter. Co-Authored-By: Claude Opus 4.6 (1M context) --- .../ml/src/cuda_pipeline/batched_backward.rs | 26 +++++-------------- .../ml/src/cuda_pipeline/batched_forward.rs | 5 ++++ .../ml/src/cuda_pipeline/gpu_dqn_trainer.rs | 1 + 3 files changed, 13 insertions(+), 19 deletions(-) diff --git a/crates/ml/src/cuda_pipeline/batched_backward.rs b/crates/ml/src/cuda_pipeline/batched_backward.rs index 747fc84cc..9789c0766 100644 --- a/crates/ml/src/cuda_pipeline/batched_backward.rs +++ b/crates/ml/src/cuda_pipeline/batched_backward.rs @@ -205,7 +205,8 @@ pub struct CublasBackwardSet { /// Per-branch cublasLt workspace buffers — eliminates workspace contention /// when 4 branch streams run GEMMs in parallel. Each 32 MB. - _branch_workspace_bufs: [CudaSlice; 4], + // Branch workspaces reused from forward pass (same graph, sequential). + // No separate allocation — saves 128MB VRAM. branch_workspace_ptrs: [u64; 4], /// Reused CUDA events for fork-join synchronization. @@ -248,6 +249,7 @@ impl CublasBackwardSet { config: &GpuDqnTrainConfig, kan_gate_backward_kernel: CudaFunction, kan_grad_reduce_kernel: CudaFunction, + fwd_branch_workspace_ptrs: [u64; 4], ) -> Result { let stream = &shared.stream; let lt_raw_handle = shared.lt_handle.0; @@ -345,23 +347,10 @@ impl CublasBackwardSet { stream.fork().map_err(|e| MLError::DeviceError(format!("bw fork branch stream 3: {e}")))?, ]; - // ── Per-branch cublasLt workspace buffers (32 MB each) ────── - let branch_ws_size: usize = 32 * 1024 * 1024; - let bw_ws_buf_0 = stream.alloc_zeros::(branch_ws_size) - .map_err(|e| MLError::ModelError(format!("bw branch workspace 0 alloc: {e}")))?; - let bw_ws_buf_1 = stream.alloc_zeros::(branch_ws_size) - .map_err(|e| MLError::ModelError(format!("bw branch workspace 1 alloc: {e}")))?; - let bw_ws_buf_2 = stream.alloc_zeros::(branch_ws_size) - .map_err(|e| MLError::ModelError(format!("bw branch workspace 2 alloc: {e}")))?; - let bw_ws_buf_3 = stream.alloc_zeros::(branch_ws_size) - .map_err(|e| MLError::ModelError(format!("bw branch workspace 3 alloc: {e}")))?; - let branch_workspace_ptrs = [ - bw_ws_buf_0.raw_ptr(), - bw_ws_buf_1.raw_ptr(), - bw_ws_buf_2.raw_ptr(), - bw_ws_buf_3.raw_ptr(), - ]; - let branch_workspace_bufs = [bw_ws_buf_0, bw_ws_buf_1, bw_ws_buf_2, bw_ws_buf_3]; + // Reuse the FORWARD pass's branch workspace pointers — forward and backward + // are in the same child graph (sequential), so workspaces never conflict. + // Saves 128MB VRAM (4 × 32MB) that was causing OOM on H100. + let branch_workspace_ptrs = fwd_branch_workspace_ptrs; // ── Pre-allocate CUDA events for fork-join ────────────────── let ctx = stream.context(); @@ -440,7 +429,6 @@ impl CublasBackwardSet { branch_2_size: config.branch_2_size, branch_3_size: config.branch_3_size, branch_streams, - _branch_workspace_bufs: branch_workspace_bufs, branch_workspace_ptrs, trunk_done_event, branch_done_events, diff --git a/crates/ml/src/cuda_pipeline/batched_forward.rs b/crates/ml/src/cuda_pipeline/batched_forward.rs index acd7d0b5d..ed15f529b 100644 --- a/crates/ml/src/cuda_pipeline/batched_forward.rs +++ b/crates/ml/src/cuda_pipeline/batched_forward.rs @@ -417,6 +417,11 @@ impl CublasGemmSet { self.lt_matmul(stream, w_ptr, a_ptr, c_ptr, m, n, k, k, self.handle.lt_workspace_ptr, self.handle.lt_workspace_size, "DIAG_raw") } + /// Per-branch cuBLAS workspace pointers (shared with backward pass). + pub fn branch_workspace_ptrs(&self) -> [u64; 4] { + self.branch_workspace_ptrs + } + /// CRITICAL: cublasSetStream resets workspace to the default pool. /// Delegates to `SharedCublasHandle::set_stream` which atomically /// rebinds the handle and restores the user-owned workspace. diff --git a/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs b/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs index 7b3e7fd5d..1f123aafa 100644 --- a/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs +++ b/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs @@ -4902,6 +4902,7 @@ impl GpuDqnTrainer { // ── Initialize cuBLAS backward context (required) ────────── let cublas_backward = CublasBackwardSet::new( Arc::clone(&shared_cublas), &config, kan_gate_backward_kernel.clone(), kan_grad_reduce_kernel.clone(), + cublas_forward.branch_workspace_ptrs(), )?; info!("GpuDqnTrainer: cuBLAS batched backward initialized (KAN backward wired)");