From 2d704eaa072cd4ae334c086ad287cd81f712c692 Mon Sep 17 00:00:00 2001 From: jgrusewski Date: Tue, 14 Apr 2026 00:01:34 +0200 Subject: [PATCH] =?UTF-8?q?fix:=20backward=20cuBLASLt=20uses=20heuristic?= =?UTF-8?q?=20algo=20=E2=80=94=20AlgoGetIds=20selects=20graph-incompatible?= =?UTF-8?q?=20algo=20on=20H100?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit AlgoGetIds picks algorithms by ID order. On H100 (SM90), the first valid algo uses split-K/stream-K with internal workspace allocation that silently produces zero output when replayed via CUDA Graph. The heuristic selects graph-safe algorithms by design. Also removes FOXHUNT_NO_GRAPH temp env var and raw DtoH diagnostic probe. Co-Authored-By: Claude Opus 4.6 (1M context) --- .../ml/src/cuda_pipeline/batched_backward.rs | 89 +++++++------------ .../ml/src/cuda_pipeline/gpu_dqn_trainer.rs | 19 ---- infra/k8s/argo/train-template.yaml | 4 - 3 files changed, 30 insertions(+), 82 deletions(-) diff --git a/crates/ml/src/cuda_pipeline/batched_backward.rs b/crates/ml/src/cuda_pipeline/batched_backward.rs index 732ab43f0..4268c1eef 100644 --- a/crates/ml/src/cuda_pipeline/batched_backward.rs +++ b/crates/ml/src/cuda_pipeline/batched_backward.rs @@ -1148,81 +1148,52 @@ 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 cublasLtMatmulAlgoGetIds ── - let max_ids = 64; - let mut algo_ids = vec![0i32; max_ids]; - let mut id_count: i32 = 0; + // ── Algorithm selection via heuristic ── + // AlgoGetIds picks the first valid algorithm by ID order, which on H100 + // (SM90) can select split-K or stream-K algorithms that use internal + // synchronization incompatible with CUDA Graph capture (produces zero + // output on replay). The heuristic selects the optimal graph-safe + // algorithm for the given shape and workspace. + let matmul_pref = cublaslt_result::create_matmul_pref() + .map_err(|e| MLError::ModelError(format!("cached bw MatmulPrefCreate (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| { + let _ = cublaslt_result::destroy_matmul_pref(matmul_pref); + MLError::ModelError(format!("cached bw set pref ws (m={m},n={n},k={k}): {e:?}")) + })?; - let status = cublaslt_sys::cublasLtMatmulAlgoGetIds( - lt_handle, - compute_type, - f32_type, f32_type, f32_type, f32_type, f32_type, - max_ids as i32, - algo_ids.as_mut_ptr(), - &mut id_count, + let heuristic = cublaslt_result::get_matmul_algo_heuristic( + lt_handle, matmul_desc, + a_layout, b_layout, c_layout, d_layout, + matmul_pref, ); - if status != cublaslt_sys::cublasStatus_t::CUBLAS_STATUS_SUCCESS || id_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); - return Err(MLError::ModelError(format!( - "cached bw AlgoGetIds (transa={transa_i32},transb={transb_i32},m={m},n={n},k={k}): status={status:?}, count={id_count}" - ))); - } + let _ = cublaslt_result::destroy_matmul_pref(matmul_pref); - // Find first algorithm that supports our shape via AlgoCheck - let mut found_algo: Option = None; - let mut found_id: i32 = -1; - for &id in &algo_ids[..id_count as usize] { - let mut algo: cublaslt_sys::cublasLtMatmulAlgo_t = std::mem::zeroed(); - let init_s = cublaslt_sys::cublasLtMatmulAlgoInit( - lt_handle, compute_type, - f32_type, f32_type, f32_type, f32_type, f32_type, - id, &mut algo, - ); - if init_s != cublaslt_sys::cublasStatus_t::CUBLAS_STATUS_SUCCESS { continue; } - - let mut result: std::mem::MaybeUninit = - std::mem::MaybeUninit::uninit(); - let check_s = cublaslt_sys::cublasLtMatmulAlgoCheck( - lt_handle, matmul_desc, - a_layout, b_layout, c_layout, d_layout, - &algo, result.as_mut_ptr(), - ); - if check_s == cublaslt_sys::cublasStatus_t::CUBLAS_STATUS_SUCCESS { - let r = result.assume_init(); - if r.state == cublaslt_sys::cublasStatus_t::CUBLAS_STATUS_SUCCESS - && r.workspaceSize <= ws_size - { - found_algo = Some(algo); - found_id = id; - break; - } - } - } - - let algo = match found_algo { - Some(a) => a, - None => { + let heuristic = match heuristic { + Ok(h) => h, + Err(e) => { 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 no valid algo (transa={transa_i32},transb={transb_i32},m={m},n={n},k={k}): checked {id_count} IDs" + "cached bw heuristic (transa={transa_i32},transb={transb_i32},m={m},n={n},k={k}): {e:?}" ))); } }; + let algo = heuristic.algo; + tracing::info!( transa = transa_i32, transb = transb_i32, m = m, n = n, k = k, lda = lda, ldb = ldb, ldc = ldc, - algo_id = found_id, - total_ids = id_count, - "cached bwd GEMM desc created (deterministic: AlgoGetIds)" + ws_needed = heuristic.workspaceSize, + "cached bwd GEMM desc created (heuristic algorithm)" ); Ok(CachedBwdGemmDesc { diff --git a/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs b/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs index bf8eff14b..bbb3bb9ba 100644 --- a/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs +++ b/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs @@ -2143,30 +2143,11 @@ impl GpuDqnTrainer { }; let grad_norm = self.debug_buffer_norm_f32(self.ptrs.grad_buf, self.total_params).unwrap_or(f32::NAN); - // Raw DtoH probe: read first 4 floats directly from grad_buf and grad_norm_buf. - // Bypasses all kernel infrastructure to verify the buffer actually has data. - let mut raw_grad = [0.0_f32; 4]; - let mut raw_norm = [0.0_f32; 1]; - unsafe { - cudarc::driver::sys::cuStreamSynchronize(self.stream.cu_stream()); - cudarc::driver::sys::cuMemcpyDtoH_v2( - raw_grad.as_mut_ptr().cast(), self.ptrs.grad_buf, 16, - ); - cudarc::driver::sys::cuMemcpyDtoH_v2( - raw_norm.as_mut_ptr().cast(), self.ptrs.grad_norm_buf, 4, - ); - } - tracing::warn!( step, d_val_norm, d_adv_norm, grad_norm, - raw_grad_0 = raw_grad[0], - raw_grad_1 = raw_grad[1], - raw_grad_2 = raw_grad[2], - raw_grad_3 = raw_grad[3], - raw_norm = raw_norm[0], c51_alpha = self.c51_alpha, "BUFFER_DIAG: per-step buffer norms after forward+backward" ); diff --git a/infra/k8s/argo/train-template.yaml b/infra/k8s/argo/train-template.yaml index bb02f1375..818fd64b3 100644 --- a/infra/k8s/argo/train-template.yaml +++ b/infra/k8s/argo/train-template.yaml @@ -400,8 +400,6 @@ spec: value: ":4096:8" - name: FOXHUNT_FEATURE_CACHE_DIR value: /feature-cache - - name: FOXHUNT_NO_GRAPH - value: "1" resources: requests: nvidia.com/gpu: "1" @@ -486,8 +484,6 @@ spec: value: ":4096:8" - name: FOXHUNT_FEATURE_CACHE_DIR value: /feature-cache - - name: FOXHUNT_NO_GRAPH - value: "1" resources: requests: nvidia.com/gpu: "1"