From 88cf7a321ed71ea0d78005f36707d8acb45c28f4 Mon Sep 17 00:00:00 2001 From: jgrusewski Date: Mon, 13 Apr 2026 08:34:39 +0200 Subject: [PATCH] fix: remove machine-specific cuBLAS algo cache MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit The algo cache serialized heuristic-selected algo bytes to disk with GPU name + CUDA version validation. This created a divergent code path between dev (RTX 3050) and prod (H100) — different machines would get different cached algos, producing different but "frozen" results. Determinism should come from the algorithm itself being deterministic, not from caching one machine's non-deterministic output. The proper fix is routing evaluation through the trainer's CUDA-Graphed forward pass (which IS deterministic by design). Kept: COMPUTE_32F_PEDANTIC on all cublasLt matmul descriptors (disables TF32, universal across GPUs). Kept: IQN + attention backward determinism, single-stream eval, evaluator reuse. Co-Authored-By: Claude Opus 4.6 (1M context) --- .claude/scheduled_tasks.lock | 2 +- crates/ml/config/cublas_algo_cache.json | 48 ----- .../ml/src/cuda_pipeline/batched_backward.rs | 135 ++++++------- .../ml/src/cuda_pipeline/batched_forward.rs | 184 ++++++++--------- .../ml/src/cuda_pipeline/cublas_algo_cache.rs | 188 ------------------ crates/ml/src/cuda_pipeline/mod.rs | 1 - crates/ml/src/trainers/dqn/fused_training.rs | 3 - crates/ml/src/trainers/dqn/trainer/mod.rs | 4 - 8 files changed, 144 insertions(+), 421 deletions(-) delete mode 100644 crates/ml/config/cublas_algo_cache.json delete mode 100644 crates/ml/src/cuda_pipeline/cublas_algo_cache.rs diff --git a/.claude/scheduled_tasks.lock b/.claude/scheduled_tasks.lock index 62cee8a51..65ca2349c 100644 --- a/.claude/scheduled_tasks.lock +++ b/.claude/scheduled_tasks.lock @@ -1 +1 @@ -{"sessionId":"4d4aa47f-4eb8-44d0-9d38-840da6e33fc0","pid":4104094,"acquiredAt":1775984198815} \ No newline at end of file +{"sessionId":"4d4aa47f-4eb8-44d0-9d38-840da6e33fc0","pid":173970,"acquiredAt":1776019129035} \ No newline at end of file diff --git a/crates/ml/config/cublas_algo_cache.json b/crates/ml/config/cublas_algo_cache.json deleted file mode 100644 index 23708f36b..000000000 --- a/crates/ml/config/cublas_algo_cache.json +++ /dev/null @@ -1,48 +0,0 @@ -{ - "cuda_version": 13000, - "device_name": "NVIDIA GeForce RTX 3050 Ti Laptop GPU", - "algos": { - "128_32_64_64_64_128_1_0_0": "010000000b0000000000000001000000000000000100000000000000000000000100000045d80100000000000000000000000000450000000000000000000000", - "64_64_64_64_64_64_0_1_2": "010000000b0000000000000001000000000000000100000000000000000000000100000045d80100000000000000000000000000450000000000000000000000", - "52_512_128_128_128_52_1_0_0": "010000000b0000000000000001000000000000000100000000000000000000000100c4a245d80100000000000000000000000000450000000000000000000000", - "64_1_64_64_64_64_1_0_1": "0d000000000000000000000001000000000000000000000058000000000000000100000045d80100000000000000000000000000450000000000000000000000", - "128_512_64_64_64_128_1_0_1": "000000000e0000000000000001000000000000000100000000000000000000000100000045d80100000000000000000000000000450000000000000000000000", - "128_64_64_64_64_128_1_0_1": "010000000b0000000000000001000000000000000100000000000000000000000100000045d80100000000000000000000000000450000000000000000000000", - "64_32_64_64_64_64_1_0_0": "010000000b0000000000000001000000000000000100000000000000000000000100000045d80100000000000000000000000000450000000000000000000000", - "156_512_128_128_128_156_1_0_0": "010000000d0000000000000001000000000000000100000000000000000000000100c4a245d80100000000000000000000000000450000000000000000000000", - "156_64_128_128_128_156_1_0_0": "010000000b0000000000000001000000000000000100000000000000000000000100000045d80100000000000000000000000000450000000000000000000000", - "128_1_64_64_64_128_1_0_0": "0d000000000000000000000001000000000000000000000058000000000000000100000045d80100000000000000000000000000450000000000000000000000", - "64_512_64_64_64_64_1_0_0": "010000000b0000000000000001000000000000000100000000000000000000000100c4a245d80100000000000000000000000000450000000000000000000000", - "64_512_80_80_128_64_1_0_0": "010000000d0000000000000001000000000000000100000000000000000000000100c4a245d80100000000000000000000000000450000000000000000000000", - "128_156_64_128_156_128_0_1_2": "010000000d0000000000000001000000000000000000000000000000000000000100000045d80100000000000000000000000000450000000000000000000000", - "64_64_64_64_64_64_0_0_2": "000000000e0000000000000001000000000000000000000000000000000000000100000045d80100000000000000000000000000450000000000000000000000", - "64_64_128_64_128_64_0_0_2": "010000000b0000000000000001000000000000000100000000000000000000000100000045d80100000000000000000000000000450000000000000000000000", - "128_64_156_128_156_128_0_0_2": "010000000b0000000000000001000000000000000000000000000000000000000100000045d80100000000000000000000000000450000000000000000000000", - "52_1_128_128_128_52_1_0_0": "0d00000000000000000000000100000000000000000000004f000000000000000100000045d80100000000000000000000000000450000000000000000000000", - "128_64_52_128_52_128_0_0_2": "000000000e0000000000000001000000000000000000000000000000000000000100000045d80100000000000000000000000000450000000000000000000000", - "156_1_128_128_128_156_1_0_0": "0d000000000000000000000001000000000000000000000057000000000000000100000045d80100000000000000000000000000450000000000000000000000", - "64_512_80_80_128_64_1_0_1": "010000000d0000000000000001000000000000000100000000000000000000000100000045d80100000000000000000000000000450000000000000000000000", - "64_32_54_54_54_64_1_0_1": "010000000b0000000000000001000000000000000100000000000000000000000100000045d80100000000000000000000000000450000000000000000000000", - "128_32_64_64_64_128_1_0_1": "010000000b0000000000000001000000000000000100000000000000000000000100000045d80100000000000000000000000000450000000000000000000000", - "64_512_64_64_64_64_1_0_1": "010000000b0000000000000001000000000000000100000000000000000000000100000045d80100000000000000000000000000450000000000000000000000", - "64_1_64_64_64_64_1_0_0": "0d000000000000000000000001000000000000000000000058000000000000000100000045d80100000000000000000000000000450000000000000000000000", - "64_1_80_80_128_64_1_0_0": "0d000000000000000000000001000000000000000000000058000000000000000100000045d80100000000000000000000000000450000000000000000000000", - "64_1_80_80_128_64_1_0_1": "0d000000000000000000000001000000000000000000000058000000000000000100000045d80100000000000000000000000000450000000000000000000000", - "64_32_64_64_64_64_1_0_1": "010000000b0000000000000001000000000000000100000000000000000000000100000045d80100000000000000000000000000450000000000000000000000", - "64_128_64_64_128_64_0_1_2": "010000000d0000000000000001000000000000000000000000000000000000000100000045d80100000000000000000000000000450000000000000000000000", - "52_32_128_128_128_52_1_0_0": "010000000b0000000000000001000000000000000100000000000000000000000100000045d80100000000000000000000000000450000000000000000000000", - "64_32_54_54_54_64_1_0_0": "010000000b0000000000000001000000000000000100000000000000000000000100000045d80100000000000000000000000000450000000000000000000000", - "156_32_128_128_128_156_1_0_0": "010000000b0000000000000001000000000000000100000000000000000000000100000045d80100000000000000000000000000450000000000000000000000", - "52_64_128_128_128_52_1_0_0": "010000000b0000000000000001000000000000000100000000000000000000000100000045d80100000000000000000000000000450000000000000000000000", - "64_64_64_64_64_64_1_0_0": "010000000b0000000000000001000000000000000100000000000000000000000100000045d80100000000000000000000000000450000000000000000000000", - "128_64_64_64_64_128_1_0_0": "010000000b0000000000000001000000000000000100000000000000000000000100000045d80100000000000000000000000000450000000000000000000000", - "64_64_64_64_64_64_1_0_1": "010000000b0000000000000001000000000000000100000000000000000000000100000045d80100000000000000000000000000450000000000000000000000", - "54_64_64_54_64_54_0_1_2": "14000000100000001c00000001000000000000000000000000000000000000000100000045d80100000000000000000000000000450000000000000000000000", - "128_512_64_64_64_128_1_0_0": "000000000e0000000000000001000000000000000100000000000000000000000100c4a245d80100000000000000000000000000450000000000000000000000", - "128_1_64_64_64_128_1_0_1": "0d000000000000000000000001000000000000000000000058000000000000000100000045d80100000000000000000000000000450000000000000000000000", - "54_64_64_54_64_54_0_0_2": "000000000e0000000000000001000000000000000000000000000000000000000100000045d80100000000000000000000000000450000000000000000000000", - "64_64_54_54_54_64_1_0_0": "010000000b0000000000000001000000000000000100000000000000000000000100000045d80100000000000000000000000000450000000000000000000000", - "64_64_54_54_54_64_1_0_1": "010000000b0000000000000001000000000000000100000000000000000000000100000045d80100000000000000000000000000450000000000000000000000", - "128_52_64_128_52_128_0_1_2": "010000000b0000000000000001000000000000000100000000000000000000000100000045d80100000000000000000000000000450000000000000000000000" - } -} \ No newline at end of file diff --git a/crates/ml/src/cuda_pipeline/batched_backward.rs b/crates/ml/src/cuda_pipeline/batched_backward.rs index 32b729d2d..0781efd1f 100644 --- a/crates/ml/src/cuda_pipeline/batched_backward.rs +++ b/crates/ml/src/cuda_pipeline/batched_backward.rs @@ -1245,21 +1245,67 @@ fn create_cached_bwd_gemm_desc( let d_layout = cublaslt_result::create_matrix_layout(f32_type, m as u64, n as u64, ldc as i64) .map_err(|e| MLError::ModelError(format!("cached bw D layout (m={m},n={n},k={k}): {e:?}")))?; - // ── Deterministic algorithm selection via algo cache ── - use super::cublas_algo_cache::{self, GemmShapeKey}; - let shape_key = GemmShapeKey { - m, n, k, - lda, ldb, ldc, - transa: transa_i32, transb: transb_i32, - variant: 2, // 2 = backward GEMM - }; + // Algorithm heuristic (requestedAlgoCount=3) + let matmul_pref = cublaslt_result::create_matmul_pref() + .map_err(|e| MLError::ModelError(format!("cached bw matmul pref (m={m},n={n},k={k}): {e:?}")))?; + cublaslt_result::set_matmul_pref_attribute( + matmul_pref, + cublaslt_sys::cublasLtMatmulPreferenceAttributes_t::CUBLASLT_MATMUL_PREF_MAX_WORKSPACE_BYTES, + &ws_size as *const usize as *const std::ffi::c_void, + std::mem::size_of::(), + ).map_err(|e| MLError::ModelError(format!("cached bw set pref ws (m={m},n={n},k={k}): {e:?}")))?; - let algo = if let Some(cached) = cublas_algo_cache::get_cached_algo(&shape_key) { - tracing::info!(transa = transa_i32, transb = transb_i32, m, n, k, lda, ldb, ldc, "bwd GEMM: using cached algo"); - cached - } else { - select_algo_heuristic_bwd(lt_handle, matmul_desc, a_layout, b_layout, c_layout, d_layout, ws_size, &shape_key, compute_type, transa_i32, transb_i32, m, n, k, lda, ldb, ldc)? - }; + let mut heuristics: [std::mem::MaybeUninit; 3] = + std::mem::MaybeUninit::uninit().assume_init(); + let mut algo_count: i32 = 0; + + let status = cublaslt_sys::cublasLtMatmulAlgoGetHeuristic( + lt_handle, + matmul_desc, + a_layout, + b_layout, + c_layout, + d_layout, + matmul_pref, + 3, + heuristics[0].as_mut_ptr(), + &mut algo_count, + ); + if status != cublaslt_sys::cublasStatus_t::CUBLAS_STATUS_SUCCESS || algo_count == 0 { + let _ = cublaslt_result::destroy_matmul_pref(matmul_pref); + let _ = cublaslt_result::destroy_matrix_layout(d_layout); + let _ = cublaslt_result::destroy_matrix_layout(c_layout); + let _ = cublaslt_result::destroy_matrix_layout(b_layout); + let _ = cublaslt_result::destroy_matrix_layout(a_layout); + let _ = cublaslt_result::destroy_matmul_desc(matmul_desc); + return Err(MLError::ModelError(format!( + "cached bw algo heuristic (transa={transa_i32},transb={transb_i32},m={m},n={n},k={k}): status={status:?}, count={algo_count}" + ))); + } + + let best = heuristics[0].assume_init(); + if best.state != cublaslt_sys::cublasStatus_t::CUBLAS_STATUS_SUCCESS { + let _ = cublaslt_result::destroy_matmul_pref(matmul_pref); + let _ = cublaslt_result::destroy_matrix_layout(d_layout); + let _ = cublaslt_result::destroy_matrix_layout(c_layout); + let _ = cublaslt_result::destroy_matrix_layout(b_layout); + let _ = cublaslt_result::destroy_matrix_layout(a_layout); + let _ = cublaslt_result::destroy_matmul_desc(matmul_desc); + return Err(MLError::ModelError(format!( + "cached bw algo heuristic state invalid (transa={transa_i32},transb={transb_i32},m={m},n={n},k={k}): {:?}", + best.state + ))); + } + + tracing::info!( + transa = transa_i32, transb = transb_i32, + m = m, n = n, k = k, lda = lda, ldb = ldb, ldc = ldc, + algo_count = algo_count, + ws_needed = best.workspaceSize, + "cached bwd GEMM desc created" + ); + + let _ = cublaslt_result::destroy_matmul_pref(matmul_pref); Ok(CachedBwdGemmDesc { matmul_desc, @@ -1267,70 +1313,11 @@ fn create_cached_bwd_gemm_desc( b_layout, c_layout, d_layout, - algo, + algo: best.algo, }) } } -/// Run heuristic for backward GEMM, cache the full algo struct, return the algo. -#[allow(clippy::too_many_arguments)] -fn select_algo_heuristic_bwd( - lt_handle: cublaslt_sys::cublasLtHandle_t, - matmul_desc: cublaslt_sys::cublasLtMatmulDesc_t, - a_layout: cublaslt_sys::cublasLtMatrixLayout_t, - b_layout: cublaslt_sys::cublasLtMatrixLayout_t, - c_layout: cublaslt_sys::cublasLtMatrixLayout_t, - d_layout: cublaslt_sys::cublasLtMatrixLayout_t, - ws_size: usize, - shape_key: &super::cublas_algo_cache::GemmShapeKey, - _compute_type: cublaslt_sys::cublasComputeType_t, - transa_i32: i32, transb_i32: i32, - m: i32, n: i32, k: i32, lda: i32, ldb: i32, ldc: i32, -) -> Result { - use super::cublas_algo_cache; - unsafe { - let matmul_pref = cublaslt_result::create_matmul_pref() - .map_err(|e| MLError::ModelError(format!("bwd pref (m={m},n={n},k={k}): {e:?}")))?; - cublaslt_result::set_matmul_pref_attribute( - matmul_pref, - cublaslt_sys::cublasLtMatmulPreferenceAttributes_t::CUBLASLT_MATMUL_PREF_MAX_WORKSPACE_BYTES, - &ws_size as *const usize as *const std::ffi::c_void, - std::mem::size_of::(), - ).map_err(|e| MLError::ModelError(format!("bwd set pref ws (m={m},n={n},k={k}): {e:?}")))?; - - let mut heuristics: [std::mem::MaybeUninit; 3] = - std::mem::MaybeUninit::uninit().assume_init(); - let mut algo_count: i32 = 0; - - let status = cublaslt_sys::cublasLtMatmulAlgoGetHeuristic( - lt_handle, matmul_desc, a_layout, b_layout, c_layout, d_layout, - matmul_pref, 3, heuristics[0].as_mut_ptr(), &mut algo_count, - ); - let _ = cublaslt_result::destroy_matmul_pref(matmul_pref); - - if status != cublaslt_sys::cublasStatus_t::CUBLAS_STATUS_SUCCESS || algo_count == 0 { - return Err(MLError::ModelError(format!( - "bwd algo heuristic (transa={transa_i32},transb={transb_i32},m={m},n={n},k={k}): status={status:?}, count={algo_count}" - ))); - } - - let best = heuristics[0].assume_init(); - if best.state != cublaslt_sys::cublasStatus_t::CUBLAS_STATUS_SUCCESS { - return Err(MLError::ModelError(format!( - "bwd algo state invalid (transa={transa_i32},transb={transb_i32},m={m},n={n},k={k}): {:?}", best.state - ))); - } - - let algo = best.algo; - - // Cache the full algo struct for cross-run determinism - cublas_algo_cache::store_algo(shape_key, &algo); - tracing::info!(transa = transa_i32, transb = transb_i32, m, n, k, lda, ldb, ldc, "bwd GEMM: caching algo"); - - Ok(algo) - } -} - // ── Kernel compilation ──────────────────────────────────────────────────────── /// Compile `relu_mask_kernel` and `bias_grad_reduce_kernel` from inline NVRTC. diff --git a/crates/ml/src/cuda_pipeline/batched_forward.rs b/crates/ml/src/cuda_pipeline/batched_forward.rs index 3aa957ea0..981f6d9e9 100644 --- a/crates/ml/src/cuda_pipeline/batched_forward.rs +++ b/crates/ml/src/cuda_pipeline/batched_forward.rs @@ -1457,76 +1457,79 @@ fn create_cached_fwd_gemm_desc( let d_layout = cublaslt_result::create_matrix_layout(f32_type, n as u64, b as u64, n as i64) .map_err(|e| MLError::ModelError(format!("cached D layout (n={n},b={b},k={k}): {e:?}")))?; - // ── Deterministic algorithm selection via algo cache ── - // The cublasLt heuristic returns different algorithms between process invocations - // (GPU clock/thermal state affects performance estimates). We cache the full algo - // struct on first run and reuse it directly for determinism. - use super::cublas_algo_cache::{self, GemmShapeKey}; - let shape_key = GemmShapeKey { - m: n as i32, n: b as i32, k: k as i32, - lda: k as i32, ldb: ldb as i32, ldc: n as i32, - transa: 1, transb: 0, variant: 0, // 0 = plain forward GEMM - }; - - let algo = if let Some(cached) = cublas_algo_cache::get_cached_algo(&shape_key) { - tracing::info!(n, b, k, ldb, "fwd GEMM: using cached algo"); - cached - } else { - select_algo_heuristic_fwd(lt_handle, matmul_desc, a_layout, b_layout, c_layout, d_layout, ws_size, &shape_key, compute_type, n, b, k, ldb)? - }; - - Ok(CachedGemmDesc { matmul_desc, a_layout, b_layout, c_layout, d_layout, algo }) - } -} - -/// Run heuristic for forward GEMM, cache the full algo struct, return the algo. -#[allow(clippy::too_many_arguments)] -fn select_algo_heuristic_fwd( - lt_handle: cublaslt_sys::cublasLtHandle_t, - matmul_desc: cublaslt_sys::cublasLtMatmulDesc_t, - a_layout: cublaslt_sys::cublasLtMatrixLayout_t, - b_layout: cublaslt_sys::cublasLtMatrixLayout_t, - c_layout: cublaslt_sys::cublasLtMatrixLayout_t, - d_layout: cublaslt_sys::cublasLtMatrixLayout_t, - ws_size: usize, - shape_key: &super::cublas_algo_cache::GemmShapeKey, - _compute_type: cublaslt_sys::cublasComputeType_t, - n: usize, b: usize, k: usize, ldb: usize, -) -> Result { - use super::cublas_algo_cache; - unsafe { + // ── Algorithm heuristic (requestedAlgoCount=3 for better selection) ── let matmul_pref = cublaslt_result::create_matmul_pref() - .map_err(|e| MLError::ModelError(format!("fwd pref (n={n},b={b},k={k}): {e:?}")))?; + .map_err(|e| MLError::ModelError(format!("cached matmul pref (n={n},b={b},k={k}): {e:?}")))?; cublaslt_result::set_matmul_pref_attribute( matmul_pref, cublaslt_sys::cublasLtMatmulPreferenceAttributes_t::CUBLASLT_MATMUL_PREF_MAX_WORKSPACE_BYTES, &ws_size as *const usize as *const std::ffi::c_void, std::mem::size_of::(), - ).map_err(|e| MLError::ModelError(format!("fwd set pref ws: {e:?}")))?; + ).map_err(|e| MLError::ModelError(format!("cached set pref ws (n={n},b={b},k={k}): {e:?}")))?; - let mut heuristic: std::mem::MaybeUninit = - std::mem::MaybeUninit::uninit(); + // Request top 3 algorithms, pick the best (first valid). + let mut heuristics: [std::mem::MaybeUninit; 3] = + std::mem::MaybeUninit::uninit().assume_init(); let mut algo_count: i32 = 0; - let status = cublaslt_sys::cublasLtMatmulAlgoGetHeuristic( - lt_handle, matmul_desc, a_layout, b_layout, c_layout, d_layout, - matmul_pref, 1, heuristic.as_mut_ptr(), &mut algo_count, - ); - let _ = cublaslt_result::destroy_matmul_pref(matmul_pref); + let status = cublaslt_sys::cublasLtMatmulAlgoGetHeuristic( + lt_handle, + matmul_desc, + a_layout, + b_layout, + c_layout, + d_layout, + matmul_pref, + 3, // requestedAlgoCount — try 3 for better selection + heuristics[0].as_mut_ptr(), + &mut algo_count, + ); if status != cublaslt_sys::cublasStatus_t::CUBLAS_STATUS_SUCCESS || algo_count == 0 { + // Cleanup on failure + let _ = cublaslt_result::destroy_matmul_pref(matmul_pref); + let _ = cublaslt_result::destroy_matrix_layout(d_layout); + let _ = cublaslt_result::destroy_matrix_layout(c_layout); + let _ = cublaslt_result::destroy_matrix_layout(b_layout); + let _ = cublaslt_result::destroy_matrix_layout(a_layout); + let _ = cublaslt_result::destroy_matmul_desc(matmul_desc); return Err(MLError::ModelError(format!( - "fwd algo heuristic (n={n},b={b},k={k},ldb={ldb}): status={status:?}, count={algo_count}" + "cached algo heuristic (n={n},b={b},k={k},ldb={ldb}): status={status:?}, count={algo_count}" ))); } - let best = heuristic.assume_init(); - let algo = best.algo; + // Pick the first valid heuristic result. + let best = heuristics[0].assume_init(); + if best.state != cublaslt_sys::cublasStatus_t::CUBLAS_STATUS_SUCCESS { + let _ = cublaslt_result::destroy_matmul_pref(matmul_pref); + let _ = cublaslt_result::destroy_matrix_layout(d_layout); + let _ = cublaslt_result::destroy_matrix_layout(c_layout); + let _ = cublaslt_result::destroy_matrix_layout(b_layout); + let _ = cublaslt_result::destroy_matrix_layout(a_layout); + let _ = cublaslt_result::destroy_matmul_desc(matmul_desc); + return Err(MLError::ModelError(format!( + "cached algo heuristic state invalid (n={n},b={b},k={k},ldb={ldb}): {:?}", + best.state + ))); + } - // Cache the full algo struct for cross-run determinism - cublas_algo_cache::store_algo(shape_key, &algo); - tracing::info!(n, b, k, ldb, "fwd GEMM: caching algo"); + tracing::info!( + n = n, b = b, k = k, ldb = ldb, + algo_count = algo_count, + ws_needed = best.workspaceSize, + "cached fwd GEMM desc created" + ); - Ok(algo) + // Preference is no longer needed after heuristic query. + let _ = cublaslt_result::destroy_matmul_pref(matmul_pref); + + Ok(CachedGemmDesc { + matmul_desc, + a_layout, + b_layout, + c_layout, + d_layout, + algo: best.algo, + }) } } @@ -1588,52 +1591,14 @@ fn create_cached_fwd_gemm_desc_relu_bias( let d_layout = cublaslt_result::create_matrix_layout(f32_type, n as u64, b as u64, n as i64) .map_err(|e| MLError::ModelError(format!("cached+relu D layout: {e:?}")))?; - // ── Deterministic algorithm selection via algo cache ── - use super::cublas_algo_cache::{self, GemmShapeKey}; - let shape_key = GemmShapeKey { - m: n as i32, n: b as i32, k: k as i32, - lda: k as i32, ldb: ldb as i32, ldc: n as i32, - transa: 1, transb: 0, variant: 1, // 1 = fused RELU_BIAS epilogue - }; - - let algo = if let Some(cached) = cublas_algo_cache::get_cached_algo(&shape_key) { - tracing::info!(n, b, k, ldb, "fwd GEMM+RELU_BIAS: using cached algo"); - cached - } else { - select_algo_heuristic_relu_bias(lt_handle, matmul_desc, a_layout, b_layout, c_layout, d_layout, ws_size, &shape_key, compute_type, n, b, k, ldb)? - }; - - Ok(CachedGemmDesc { - matmul_desc, a_layout, b_layout, c_layout, d_layout, - algo, - }) - } -} - -/// Run heuristic for forward GEMM+RELU_BIAS, cache the full algo struct, return the algo. -#[allow(clippy::too_many_arguments)] -fn select_algo_heuristic_relu_bias( - lt_handle: cublaslt_sys::cublasLtHandle_t, - matmul_desc: cublaslt_sys::cublasLtMatmulDesc_t, - a_layout: cublaslt_sys::cublasLtMatrixLayout_t, - b_layout: cublaslt_sys::cublasLtMatrixLayout_t, - c_layout: cublaslt_sys::cublasLtMatrixLayout_t, - d_layout: cublaslt_sys::cublasLtMatrixLayout_t, - ws_size: usize, - shape_key: &super::cublas_algo_cache::GemmShapeKey, - _compute_type: cublaslt_sys::cublasComputeType_t, - n: usize, b: usize, k: usize, ldb: usize, -) -> Result { - use super::cublas_algo_cache; - unsafe { let matmul_pref = cublaslt_result::create_matmul_pref() - .map_err(|e| MLError::ModelError(format!("relu_bias pref (n={n},b={b},k={k}): {e:?}")))?; + .map_err(|e| MLError::ModelError(format!("cached+relu pref: {e:?}")))?; cublaslt_result::set_matmul_pref_attribute( matmul_pref, cublaslt_sys::cublasLtMatmulPreferenceAttributes_t::CUBLASLT_MATMUL_PREF_MAX_WORKSPACE_BYTES, &ws_size as *const usize as *const std::ffi::c_void, std::mem::size_of::(), - ).map_err(|e| MLError::ModelError(format!("relu_bias set pref ws: {e:?}")))?; + ).map_err(|e| MLError::ModelError(format!("cached+relu set pref ws: {e:?}")))?; let mut heuristics: [std::mem::MaybeUninit; 3] = std::mem::MaybeUninit::uninit().assume_init(); @@ -1646,25 +1611,40 @@ fn select_algo_heuristic_relu_bias( let _ = cublaslt_result::destroy_matmul_pref(matmul_pref); if status != cublaslt_sys::cublasStatus_t::CUBLAS_STATUS_SUCCESS || algo_count == 0 { + let _ = cublaslt_result::destroy_matrix_layout(d_layout); + let _ = cublaslt_result::destroy_matrix_layout(c_layout); + let _ = cublaslt_result::destroy_matrix_layout(b_layout); + let _ = cublaslt_result::destroy_matrix_layout(a_layout); + let _ = cublaslt_result::destroy_matmul_desc(matmul_desc); + tracing::warn!(n, b, k, ldb, "RELU_BIAS epilogue heuristic failed — falling back to separate kernels"); return Err(MLError::ModelError(format!( - "relu_bias algo heuristic (n={n},b={b},k={k},ldb={ldb}): status={status:?}, count={algo_count}" + "cached+relu algo heuristic (n={n},b={b},k={k}): status={status:?}, count={algo_count}" ))); } let best = heuristics[0].assume_init(); if best.state != cublaslt_sys::cublasStatus_t::CUBLAS_STATUS_SUCCESS { + let _ = cublaslt_result::destroy_matrix_layout(d_layout); + let _ = cublaslt_result::destroy_matrix_layout(c_layout); + let _ = cublaslt_result::destroy_matrix_layout(b_layout); + let _ = cublaslt_result::destroy_matrix_layout(a_layout); + let _ = cublaslt_result::destroy_matmul_desc(matmul_desc); + tracing::warn!(n, b, k, ldb, "RELU_BIAS epilogue algo invalid — falling back"); return Err(MLError::ModelError(format!( - "relu_bias algo state invalid (n={n},b={b},k={k}): {:?}", best.state + "cached+relu algo state invalid (n={n},b={b},k={k}): {:?}", best.state ))); } - let algo = best.algo; + tracing::info!( + n, b, k, ldb, algo_count, + ws_needed = best.workspaceSize, + "cached fwd GEMM+RELU_BIAS desc created" + ); - // Cache the full algo struct for cross-run determinism - cublas_algo_cache::store_algo(shape_key, &algo); - tracing::info!(n, b, k, ldb, "fwd GEMM+RELU_BIAS: caching algo"); - - Ok(algo) + Ok(CachedGemmDesc { + matmul_desc, a_layout, b_layout, c_layout, d_layout, + algo: best.algo, + }) } } diff --git a/crates/ml/src/cuda_pipeline/cublas_algo_cache.rs b/crates/ml/src/cuda_pipeline/cublas_algo_cache.rs deleted file mode 100644 index 991767df1..000000000 --- a/crates/ml/src/cuda_pipeline/cublas_algo_cache.rs +++ /dev/null @@ -1,188 +0,0 @@ -//! Deterministic cublasLt algorithm cache. -//! -//! cublasLtMatmulAlgoGetHeuristic returns different algorithms between process -//! invocations (GPU clock/thermal state affects performance estimates). This -//! causes non-reproducible training — different GEMM algorithms produce different -//! FP results due to internal tiling/reduction order. -//! -//! Solution: cache the heuristic-selected algo ID per GEMM shape. On subsequent -//! runs, reconstruct the exact same algo via cublasLtMatmulAlgoInit. The cache -//! is serialized to disk with metadata (CUDA version + GPU name) for validation. -//! -//! When the same algo ID is used, cublasLtMatmul produces bit-identical results. - -#![allow(unsafe_code)] - -use std::collections::HashMap; -use std::sync::Mutex; - -use cudarc::cublaslt::sys as cublaslt_sys; - -static ALGO_CACHE: Mutex> = Mutex::new(None); - -const CACHE_FILE: &str = "config/cublas_algo_cache.json"; - -/// Key uniquely identifying a GEMM problem shape. -#[derive(Debug, Clone, Hash, Eq, PartialEq, serde::Serialize, serde::Deserialize)] -pub struct GemmShapeKey { - pub m: i32, - pub n: i32, - pub k: i32, - pub lda: i32, - pub ldb: i32, - pub ldc: i32, - pub transa: i32, - pub transb: i32, - /// Distinguishes plain GEMM (0) from fused GEMM+RELU_BIAS epilogue (1), backward (2), etc. - pub variant: i32, -} - -/// Size of cublasLtMatmulAlgo_t opaque struct (64 bytes). -const ALGO_BYTES: usize = std::mem::size_of::(); - -/// Serializable algo cache with environment metadata. -/// Stores the FULL opaque algo struct (base64-encoded) — not just the algo ID. -/// This preserves epilogue support, tiling config, and all internal state. -#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)] -struct AlgoCache { - cuda_version: i32, - device_name: String, - algos: HashMap, // shape_key → base64-encoded algo bytes -} - -impl AlgoCache { - fn key_string(key: &GemmShapeKey) -> String { - format!( - "{}_{}_{}_{}_{}_{}_{}_{}_{}", - key.m, key.n, key.k, key.lda, key.ldb, key.ldc, key.transa, key.transb, key.variant - ) - } -} - -/// Get the current CUDA runtime version. -fn cuda_runtime_version() -> i32 { - let mut version: i32 = 0; - unsafe { - cudarc::driver::sys::cuDriverGetVersion(&mut version); - } - version -} - -/// Get the GPU device name. -fn gpu_device_name() -> String { - let mut name = [0u8; 256]; - unsafe { - let _ = cudarc::driver::sys::cuDeviceGetName( - name.as_mut_ptr() as *mut i8, - 256, - 0, // device 0 - ); - } - let len = name.iter().position(|&b| b == 0).unwrap_or(256); - String::from_utf8_lossy(&name[..len]).to_string() -} - -/// Load the algo cache from disk. Returns None if file doesn't exist or metadata mismatches. -fn load_cache() -> Option { - let data = std::fs::read_to_string(CACHE_FILE).ok()?; - let cache: AlgoCache = serde_json::from_str(&data).ok()?; - - let current_cuda = cuda_runtime_version(); - let current_gpu = gpu_device_name(); - - if cache.cuda_version != current_cuda || cache.device_name != current_gpu { - tracing::info!( - cached_cuda = cache.cuda_version, - current_cuda, - cached_gpu = %cache.device_name, - current_gpu = %current_gpu, - "cuBLAS algo cache invalidated (environment changed)" - ); - return None; - } - - tracing::info!( - entries = cache.algos.len(), - "cuBLAS algo cache loaded from {CACHE_FILE}" - ); - Some(cache) -} - -/// Save the algo cache to disk. -fn save_cache(cache: &AlgoCache) { - if let Some(parent) = std::path::Path::new(CACHE_FILE).parent() { - let _ = std::fs::create_dir_all(parent); - } - match serde_json::to_string_pretty(cache) { - Ok(json) => { - if let Err(e) = std::fs::write(CACHE_FILE, json) { - tracing::warn!("Failed to save cuBLAS algo cache: {e}"); - } else { - tracing::info!(entries = cache.algos.len(), "cuBLAS algo cache saved to {CACHE_FILE}"); - } - } - Err(e) => tracing::warn!("Failed to serialize cuBLAS algo cache: {e}"), - } -} - -/// Look up a cached algo for the given GEMM shape. -/// Returns Some(algo) with the full opaque struct if found, None otherwise. -pub fn get_cached_algo(key: &GemmShapeKey) -> Option { - let guard = ALGO_CACHE.lock().ok()?; - let cache = guard.as_ref()?; - let key_str = AlgoCache::key_string(key); - let hex = cache.algos.get(&key_str)?; - algo_from_hex(hex) -} - -/// Store a full algo struct in the cache and persist to disk. -pub fn store_algo(key: &GemmShapeKey, algo: &cublaslt_sys::cublasLtMatmulAlgo_t) { - let mut guard = match ALGO_CACHE.lock() { - Ok(g) => g, - Err(_) => return, - }; - - let cache = guard.get_or_insert_with(|| AlgoCache { - cuda_version: cuda_runtime_version(), - device_name: gpu_device_name(), - algos: HashMap::new(), - }); - - let key_str = AlgoCache::key_string(key); - cache.algos.insert(key_str, algo_to_hex(algo)); - save_cache(cache); -} - -/// Initialize the global algo cache from disk (call once at startup). -pub fn init_cache() { - let mut guard = match ALGO_CACHE.lock() { - Ok(g) => g, - Err(_) => return, - }; - if guard.is_none() { - *guard = load_cache(); - } -} - -/// Serialize a cublasLtMatmulAlgo_t to hex string (no external deps). -fn algo_to_hex(algo: &cublaslt_sys::cublasLtMatmulAlgo_t) -> String { - let bytes: &[u8] = unsafe { - std::slice::from_raw_parts(algo as *const _ as *const u8, ALGO_BYTES) - }; - bytes.iter().map(|b| format!("{b:02x}")).collect() -} - -/// Deserialize a cublasLtMatmulAlgo_t from hex string. -fn algo_from_hex(hex: &str) -> Option { - let bytes: Vec = (0..hex.len()) - .step_by(2) - .map(|i| u8::from_str_radix(&hex[i..i + 2], 16)) - .collect::, _>>() - .ok()?; - if bytes.len() != ALGO_BYTES { return None; } - let mut algo: cublaslt_sys::cublasLtMatmulAlgo_t = unsafe { std::mem::zeroed() }; - unsafe { - std::ptr::copy_nonoverlapping(bytes.as_ptr(), &mut algo as *mut _ as *mut u8, ALGO_BYTES); - } - Some(algo) -} diff --git a/crates/ml/src/cuda_pipeline/mod.rs b/crates/ml/src/cuda_pipeline/mod.rs index f2512d4b3..75f417134 100644 --- a/crates/ml/src/cuda_pipeline/mod.rs +++ b/crates/ml/src/cuda_pipeline/mod.rs @@ -32,7 +32,6 @@ pub mod batched_backward; pub mod gpu_her; #[cfg(test)] mod cublaslt_debug; -pub mod cublas_algo_cache; pub mod gpu_iql_trainer; pub mod gpu_iqn_head; pub mod gpu_attention; diff --git a/crates/ml/src/trainers/dqn/fused_training.rs b/crates/ml/src/trainers/dqn/fused_training.rs index ff844b664..e184d5850 100644 --- a/crates/ml/src/trainers/dqn/fused_training.rs +++ b/crates/ml/src/trainers/dqn/fused_training.rs @@ -313,9 +313,6 @@ impl FusedTrainingCtx { batch_size: usize, stream: Arc, ) -> Result { - // Initialize deterministic cuBLAS algo cache from disk (if available). - crate::cuda_pipeline::cublas_algo_cache::init_cache(); - let dqn = agent.primary_dqn(); dqn.branching_q_network.as_ref().ok_or_else(|| { diff --git a/crates/ml/src/trainers/dqn/trainer/mod.rs b/crates/ml/src/trainers/dqn/trainer/mod.rs index 98d31c0da..5c3aa20b4 100644 --- a/crates/ml/src/trainers/dqn/trainer/mod.rs +++ b/crates/ml/src/trainers/dqn/trainer/mod.rs @@ -563,10 +563,6 @@ impl DQNTrainer { { use crate::cuda_pipeline::gpu_walk_forward::GpuWalkForwardConfig; - // Initialize cuBLAS algo cache BEFORE any CublasForward construction. - // Must be the first CUDA-related call to ensure all instances use cached algos. - crate::cuda_pipeline::cublas_algo_cache::init_cache(); - // Convert all_data to fixed-size arrays ONCE let features: Vec<[f64; 42]> = training_data.iter().map(|(fv, _)| { let mut f = [0.0_f64; 42];