fix: remove machine-specific cuBLAS algo cache
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) <noreply@anthropic.com>
This commit is contained in:
@@ -1 +1 @@
|
||||
{"sessionId":"4d4aa47f-4eb8-44d0-9d38-840da6e33fc0","pid":4104094,"acquiredAt":1775984198815}
|
||||
{"sessionId":"4d4aa47f-4eb8-44d0-9d38-840da6e33fc0","pid":173970,"acquiredAt":1776019129035}
|
||||
@@ -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"
|
||||
}
|
||||
}
|
||||
@@ -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::<usize>(),
|
||||
).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<cublaslt_sys::cublasLtMatmulHeuristicResult_t>; 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<cublaslt_sys::cublasLtMatmulAlgo_t, MLError> {
|
||||
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::<usize>(),
|
||||
).map_err(|e| MLError::ModelError(format!("bwd set pref ws (m={m},n={n},k={k}): {e:?}")))?;
|
||||
|
||||
let mut heuristics: [std::mem::MaybeUninit<cublaslt_sys::cublasLtMatmulHeuristicResult_t>; 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.
|
||||
|
||||
@@ -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<cublaslt_sys::cublasLtMatmulAlgo_t, MLError> {
|
||||
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::<usize>(),
|
||||
).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<cublaslt_sys::cublasLtMatmulHeuristicResult_t> =
|
||||
std::mem::MaybeUninit::uninit();
|
||||
// Request top 3 algorithms, pick the best (first valid).
|
||||
let mut heuristics: [std::mem::MaybeUninit<cublaslt_sys::cublasLtMatmulHeuristicResult_t>; 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<cublaslt_sys::cublasLtMatmulAlgo_t, MLError> {
|
||||
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::<usize>(),
|
||||
).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<cublaslt_sys::cublasLtMatmulHeuristicResult_t>; 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,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -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<Option<AlgoCache>> = 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::<cublaslt_sys::cublasLtMatmulAlgo_t>();
|
||||
|
||||
/// 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<String, String>, // 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<AlgoCache> {
|
||||
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<cublaslt_sys::cublasLtMatmulAlgo_t> {
|
||||
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<cublaslt_sys::cublasLtMatmulAlgo_t> {
|
||||
let bytes: Vec<u8> = (0..hex.len())
|
||||
.step_by(2)
|
||||
.map(|i| u8::from_str_radix(&hex[i..i + 2], 16))
|
||||
.collect::<Result<Vec<_>, _>>()
|
||||
.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)
|
||||
}
|
||||
@@ -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;
|
||||
|
||||
@@ -313,9 +313,6 @@ impl FusedTrainingCtx {
|
||||
batch_size: usize,
|
||||
stream: Arc<cudarc::driver::CudaStream>,
|
||||
) -> Result<Self> {
|
||||
// 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(|| {
|
||||
|
||||
@@ -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];
|
||||
|
||||
Reference in New Issue
Block a user