From 4fa9da9fc1277d954cae04e0cf9d9b420245dc5d Mon Sep 17 00:00:00 2001 From: jgrusewski Date: Wed, 27 May 2026 20:19:45 +0200 Subject: [PATCH] perf(cuda): enable TF32 Tensor Core math on all cuBLAS handles MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit DQN head (dqn.rs) and Mamba2 encoder (mamba2_block.rs) were forcing CUBLAS_COMPUTE_32F — pure FP32 scalar, bypassing Tensor Cores entirely. TF32 (19-bit mantissa) gives 2-3× SGEMM throughput on L40S/H100 with negligible precision loss (well within RL gradient noise). 8+ SGEMMs per step now use Tensor Cores: DQN forward/backward (4) + Mamba2 W_in/W_a/W_b/W_out projections (4+). Co-Authored-By: Claude Opus 4.7 --- crates/ml-alpha/src/mamba2_block.rs | 6 ++++++ crates/ml-alpha/src/rl/dqn.rs | 6 ++++++ 2 files changed, 12 insertions(+) diff --git a/crates/ml-alpha/src/mamba2_block.rs b/crates/ml-alpha/src/mamba2_block.rs index 79275b8cd..bde8edfba 100644 --- a/crates/ml-alpha/src/mamba2_block.rs +++ b/crates/ml-alpha/src/mamba2_block.rs @@ -547,6 +547,12 @@ impl Mamba2Block { ) .result() .map_err(|e| anyhow!("Mamba2Block: cublasSetWorkspace_v2: {e:?}"))?; + cudarc::cublas::sys::cublasSetMathMode( + *cublas.handle(), + cudarc::cublas::sys::cublasMath_t::CUBLAS_TF32_TENSOR_OP_MATH, + ) + .result() + .map_err(|e| anyhow!("Mamba2Block: cublasSetMathMode TF32: {e:?}"))?; } // ── Parameter allocation + initialisation. ────────────────────── diff --git a/crates/ml-alpha/src/rl/dqn.rs b/crates/ml-alpha/src/rl/dqn.rs index 0d7d40500..c725b6a69 100644 --- a/crates/ml-alpha/src/rl/dqn.rs +++ b/crates/ml-alpha/src/rl/dqn.rs @@ -311,6 +311,12 @@ impl DqnHead { ) .result() .map_err(|e| anyhow::anyhow!("DqnHead: cublasSetWorkspace_v2: {e:?}"))?; + cudarc::cublas::sys::cublasSetMathMode( + *cublas.handle(), + cudarc::cublas::sys::cublasMath_t::CUBLAS_TF32_TENSOR_OP_MATH, + ) + .result() + .map_err(|e| anyhow::anyhow!("DqnHead: cublasSetMathMode TF32: {e:?}"))?; } // Bias-add and reduce-sum-axis0 kernel handles from ml-core.