fix: backward branches reuse forward workspace — saves 128MB (OOM fix)

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) <noreply@anthropic.com>
This commit is contained in:
jgrusewski
2026-04-18 16:57:34 +02:00
parent 80d871ae1d
commit 32151ba8bc
3 changed files with 13 additions and 19 deletions

View File

@@ -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<u8>; 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<Self, MLError> {
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::<u8>(branch_ws_size)
.map_err(|e| MLError::ModelError(format!("bw branch workspace 0 alloc: {e}")))?;
let bw_ws_buf_1 = stream.alloc_zeros::<u8>(branch_ws_size)
.map_err(|e| MLError::ModelError(format!("bw branch workspace 1 alloc: {e}")))?;
let bw_ws_buf_2 = stream.alloc_zeros::<u8>(branch_ws_size)
.map_err(|e| MLError::ModelError(format!("bw branch workspace 2 alloc: {e}")))?;
let bw_ws_buf_3 = stream.alloc_zeros::<u8>(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,

View File

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

View File

@@ -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)");