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