From ff141691a11901400f5bbe9f2c7a13a2bccbc9ab Mon Sep 17 00:00:00 2001 From: Kawrakow Date: Thu, 13 Aug 2026 15:22:01 +0200 Subject: [PATCH] Use f32 accumulation in CUDA DSA implementation (#2311) --- ggml/src/ggml-cuda/dsa_attn.cu | 22 +++++++++++++--------- 1 file changed, 13 insertions(+), 9 deletions(-) diff --git a/ggml/src/ggml-cuda/dsa_attn.cu b/ggml/src/ggml-cuda/dsa_attn.cu index ba20f801..6cebea2d 100644 --- a/ggml/src/ggml-cuda/dsa_attn.cu +++ b/ggml/src/ggml-cuda/dsa_attn.cu @@ -68,7 +68,8 @@ static __global__ void k_copy_dst(int nelem, const half * kqv16, float * dst) { } template -static __global__ void soft_max_f16_simple(half * x, const half * mask, const float * sinks, const int ncols_par, const int nrows_y, const float scale) { +static __global__ void soft_max_f16_simple(half * x, const half * mask, const float * sinks, const int ncols_par, const int nrows_y, + const float scale) { const int ncols = ncols_template == 0 ? ncols_par : ncols_template; const int tid = threadIdx.x; @@ -253,6 +254,7 @@ bool ggml_cuda_dsa_attn_ext(ggml_backend_cuda_context & ctx, ggml_tensor * dst) const half alpha = 1.0f; const half beta = 0.0f; + float alpha_32 = 1.0f, beta_32 = 0.0f; int max_rows = std::min(Q->ne[1], k_max_rows); bool is_k_view = v_is_k_view(K, V); @@ -331,17 +333,19 @@ bool ggml_cuda_dsa_attn_ext(ggml_backend_cuda_context & ctx, ggml_tensor * dst) CUDA_CHECK(cudaGetLastError()); if (is_k_view) { - CUBLAS_CHECK(cublasHgemmStridedBatched(ctx.cublas_handle(), CUBLAS_OP_N, CUBLAS_OP_N, + CUBLAS_CHECK(cublasGemmStridedBatchedEx(ctx.cublas_handle(), CUBLAS_OP_N, CUBLAS_OP_N, V->ne[0], Q->ne[2], indexer->ne[0], - &alpha, k16.get() + v_offset, K->ne[0], K->ne[0]*indexer->ne[0], - kq16.get(), indexer->ne[0], indexer->ne[0]*Q->ne[2], - &beta, kqv16.get(), V->ne[0], V->ne[0]*Q->ne[2], nrows)); + &alpha_32, k16.get() + v_offset, CUDA_R_16F, K->ne[0], K->ne[0]*indexer->ne[0], + kq16.get(), CUDA_R_16F, indexer->ne[0], indexer->ne[0]*Q->ne[2], + &beta_32, kqv16.get(), CUDA_R_16F, V->ne[0], V->ne[0]*Q->ne[2], nrows, + CUDA_R_32F, CUBLAS_GEMM_DEFAULT_TENSOR_OP)); } else { - CUBLAS_CHECK(cublasHgemmStridedBatched(ctx.cublas_handle(), CUBLAS_OP_N, CUBLAS_OP_N, + CUBLAS_CHECK(cublasGemmStridedBatchedEx(ctx.cublas_handle(), CUBLAS_OP_N, CUBLAS_OP_N, V->ne[0], Q->ne[2], indexer->ne[0], - &alpha, v16.get(), V->ne[0], V->ne[0]*indexer->ne[0], - kq16.get(), indexer->ne[0], indexer->ne[0]*Q->ne[2], - &beta, kqv16.get(), V->ne[0], V->ne[0]*Q->ne[2], nrows)); + &alpha_32, v16.get(), CUDA_R_16F, V->ne[0], V->ne[0]*indexer->ne[0], + kq16.get(), CUDA_R_16F, indexer->ne[0], indexer->ne[0]*Q->ne[2], + &beta_32, kqv16.get(), CUDA_R_16F, V->ne[0], V->ne[0]*Q->ne[2], nrows, + CUDA_R_32F, CUBLAS_GEMM_DEFAULT_TENSOR_OP)); } {