diff --git a/ggml/src/ggml-cuda/fattn-mma-f16.cuh b/ggml/src/ggml-cuda/fattn-mma-f16.cuh index 578f6cf79c25..1588912b77d3 100644 --- a/ggml/src/ggml-cuda/fattn-mma-f16.cuh +++ b/ggml/src/ggml-cuda/fattn-mma-f16.cuh @@ -1757,7 +1757,8 @@ static __device__ __forceinline__ void flash_attn_ext_f16_process_tile( static constexpr __host__ __device__ bool ggml_cuda_flash_attn_ext_mma_f16_may_use_sparse( const int DKQ, const int DV, const int ncols1, const int ncols2) { - return (DKQ == 512 && DV == 512 && ncols1 == 1 && ncols2 == 8) || + return (DKQ == 512 && DV == 512 && ncols1 == 1 && ncols2 == 8) || + (DKQ == 512 && DV == 512 && ncols1 == 1 && ncols2 == 32) || // Volta: >= 32 columns per launch (DKQ == 576 && DV == 512 && ncols1 == 1 && ncols2 == 16); } @@ -2118,6 +2119,7 @@ extern DECL_FATTN_MMA_F16_CASE(512, 512, 1, 8); extern DECL_FATTN_MMA_F16_CASE(512, 512, 2, 8); extern DECL_FATTN_MMA_F16_CASE(512, 512, 4, 8); extern DECL_FATTN_MMA_F16_CASE(512, 512, 8, 8); +extern DECL_FATTN_MMA_F16_CASE(512, 512, 1, 32); // sparse on Volta // The number of viable configurations for Deepseek is very limited: extern DECL_FATTN_MMA_F16_CASE(576, 512, 1, 16); diff --git a/ggml/src/ggml-cuda/fattn.cu b/ggml/src/ggml-cuda/fattn.cu index ceb4727931d4..5a3f495150de 100644 --- a/ggml/src/ggml-cuda/fattn.cu +++ b/ggml/src/ggml-cuda/fattn.cu @@ -106,15 +106,25 @@ void ggml_cuda_flash_attn_ext_compact_mask( #endif // !defined(GGML_USE_HIP) && !defined(GGML_USE_MUSA) } -bool ggml_cuda_flash_attn_ext_mma_f16_shall_use_sparse(ggml_backend_cuda_context & ctx, ggml_tensor * dst) { +// Volta can use the sparse kernels as well, but only via the <512, 512, 1, 32> instance: +// the sparse gather needs one index row per launch (ncols1 == 1) and the Volta build has no device code +// for launches with fewer than 32 columns, so all GQA heads of one query row go into a single launch. +static bool ggml_cuda_fattn_volta_sparse_ok(const int cc, const ggml_tensor * dst) { + const ggml_tensor * Q = dst->src[0]; + const ggml_tensor * K = dst->src[1]; + const ggml_tensor * V = dst->src[2]; + return cc == GGML_CUDA_CC_VOLTA && Q->ne[0] == 512 && V->ne[0] == 512 && + K->ne[2] > 0 && Q->ne[2] % K->ne[2] == 0 && (Q->ne[2] / K->ne[2]) % 32 == 0; +} + +static bool ggml_cuda_fattn_sparse_applicable(const int cc, const ggml_tensor * dst) { #if defined(GGML_USE_HIP) || defined(GGML_USE_MUSA) - GGML_UNUSED_VARS(ctx, dst); + GGML_UNUSED_VARS(cc, dst); return false; #else const ggml_tensor * Q = dst->src[0]; const ggml_tensor * K = dst->src[1]; const ggml_tensor * mask = dst->src[3]; - const int cc = ggml_cuda_info().devices[ctx.device].cc; float max_bias = 0.0f; float logit_softcap = 0.0f; @@ -122,13 +132,17 @@ bool ggml_cuda_flash_attn_ext_mma_f16_shall_use_sparse(ggml_backend_cuda_context memcpy(&logit_softcap, (const float *) dst->op_params + 2, sizeof(float)); const int32_t n_kv_max = ggml_get_op_params_i32(dst, 4); - return GGML_CUDA_CC_IS_NVIDIA(cc) && turing_mma_available(cc) && + return GGML_CUDA_CC_IS_NVIDIA(cc) && (turing_mma_available(cc) || ggml_cuda_fattn_volta_sparse_ok(cc, dst)) && mask != nullptr && n_kv_max > 0 && max_bias == 0.0f && logit_softcap == 0.0f && mask->ne[0] == K->ne[1] && mask->ne[1] >= Q->ne[1] && mask->ne[2] == 1 && K->ne[1] >= std::max(4096, 2LL*n_kv_max); #endif // !defined(GGML_USE_HIP) && !defined(GGML_USE_MUSA) } +bool ggml_cuda_flash_attn_ext_mma_f16_shall_use_sparse(ggml_backend_cuda_context & ctx, ggml_tensor * dst) { + return ggml_cuda_fattn_sparse_applicable(ggml_cuda_info().devices[ctx.device].cc, dst); +} + template static void ggml_cuda_flash_attn_ext_mma_f16_switch_ncols1(ggml_backend_cuda_context & ctx, ggml_tensor * dst) { const int cc = ggml_cuda_info().devices[ggml_cuda_get_device()].cc; @@ -136,7 +150,7 @@ static void ggml_cuda_flash_attn_ext_mma_f16_switch_ncols1(ggml_backend_cuda_con #if !defined(GGML_USE_HIP) && !defined(GGML_USE_MUSA) if constexpr (ggml_cuda_flash_attn_ext_mma_f16_may_use_sparse(DKQ, DV, 1, ncols2)) { - if (ggml_cuda_flash_attn_ext_mma_f16_shall_use_sparse(ctx, dst)) { + if ((cc != GGML_CUDA_CC_VOLTA || ncols2 >= 32) && ggml_cuda_flash_attn_ext_mma_f16_shall_use_sparse(ctx, dst)) { ggml_cuda_flash_attn_ext_mma_f16_case(ctx, dst); return; } @@ -198,6 +212,13 @@ static void ggml_cuda_flash_attn_ext_mma_f16_switch_ncols2(ggml_backend_cuda_con // On Volta the GQA optimizations aren't as impactful vs. minimizing wasted compute: if (cc == GGML_CUDA_CC_VOLTA) { + if constexpr (ggml_cuda_flash_attn_ext_mma_f16_may_use_sparse(DKQ, DV, 1, 32)) { + if (use_gqa_opt && gqa_ratio % 32 == 0 && ggml_cuda_flash_attn_ext_mma_f16_shall_use_sparse(ctx, dst)) { + ggml_cuda_flash_attn_ext_mma_f16_case(ctx, dst); + return; + } + } + if (use_gqa_opt && gqa_ratio % 8 == 0) { ggml_cuda_flash_attn_ext_mma_f16_switch_ncols1(ctx, dst); return; @@ -642,6 +663,11 @@ static best_fattn_kernel ggml_cuda_get_best_fattn_kernel(const int device, const } if (volta_mma_available(cc) && Q->ne[0] != 40 && Q->ne[0] != 72) { + // a sparse launch reads n_kv_max entries instead of all of K, which beats the dense tile kernel + // already for a single token as soon as the KV cache is much larger than n_kv_max + if (gqa_opt_applies && ggml_cuda_fattn_sparse_applicable(cc, dst)) { + return BEST_FATTN_KERNEL_MMA_F16; + } if (can_use_vector_kernel && Q->ne[1] * gqa_ratio_eff <= 2) { return BEST_FATTN_KERNEL_VEC; } diff --git a/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_1-ncols2_32.cu b/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_1-ncols2_32.cu index 8fc3b17976e7..6330ea14c04c 100644 --- a/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_1-ncols2_32.cu +++ b/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_1-ncols2_32.cu @@ -3,4 +3,5 @@ #include "../fattn-mma-f16.cuh" DECL_FATTN_MMA_F16_CASE(320, 256, 1, 32); +DECL_FATTN_MMA_F16_CASE(512, 512, 1, 32); DECL_FATTN_MMA_F16_CASE(576, 512, 1, 32); diff --git a/ggml/src/ggml-cuda/template-instances/generate_cu_files.py b/ggml/src/ggml-cuda/template-instances/generate_cu_files.py index d7cd271675e0..d7b6639d36bd 100755 --- a/ggml/src/ggml-cuda/template-instances/generate_cu_files.py +++ b/ggml/src/ggml-cuda/template-instances/generate_cu_files.py @@ -93,11 +93,13 @@ def get_short_name(long_quant_name): continue if head_size_kq == 320 and ncols2 != 32: # Mistral Small 4 continue - if head_size_kq == 512 and ncols2 not in (2, 4, 8): # Gemma 4 (+ MTP) + if head_size_kq == 512 and ncols2 not in (2, 4, 8, 32): # Gemma 4 (+ MTP); 32: sparse FA on Volta + continue + if head_size_kq == 512 and ncols2 == 32 and ncols1 != 1: # sparse gathers one index row per launch continue if head_size_kq == 576 and ncols2 not in (4, 16, 32): # Deepseek, GLM 4.7 Flash continue - if head_size_kq not in (192, 320, 576) and ncols2 in (16, 32): + if head_size_kq not in (192, 320, 512, 576) and ncols2 in (16, 32): continue head_size_v = HEAD_SIZES_V_OVERRIDE.get(head_size_kq, head_size_kq) f.write(SOURCE_FATTN_MMA_CASE.format(ncols1=ncols1, ncols2=ncols2, head_size_kq=head_size_kq, head_size_v=head_size_v)) diff --git a/tests/test-backend-ops.cpp b/tests/test-backend-ops.cpp index d02297cf518d..7879612f8ee4 100644 --- a/tests/test-backend-ops.cpp +++ b/tests/test-backend-ops.cpp @@ -10684,6 +10684,10 @@ static std::vector> make_test_cases_eval() { // Sparse mask hint: supported decode/prefill layouts and dense fallbacks. test_cases.emplace_back(new test_flash_attn_ext(512, 512, 1, { 8, 1}, 4096, 1, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16, {0, 1, 2, 3}, true, false, 512)); + // GQA 64 on 1 KV head + sinks: the <512, 512, 1, 32> sparse launch (used on Volta) + test_cases.emplace_back(new test_flash_attn_ext(512, 512, 1, {64, 1}, 8192, 32, true, true, 0, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16, {0, 1, 2, 3}, true, false, 640)); + test_cases.emplace_back(new test_flash_attn_ext(512, 512, 1, {64, 1}, 8192, 32, true, true, 0, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16, {0, 1, 2, 3}, true, true, 640)); + test_cases.emplace_back(new test_flash_attn_ext(512, 512, 1, {64, 1}, 8192, 1, true, true, 0, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16, {0, 1, 2, 3}, true, false, 640)); test_cases.emplace_back(new test_flash_attn_ext(512, 512, 1, { 8, 2}, 4096, 3, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16, {0, 1, 2, 3}, true, false, 768)); test_cases.emplace_back(new test_flash_attn_ext(576, 512, 1, {16, 1}, 4096, 1, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16, {0, 1, 2, 3}, true, true, 512)); test_cases.emplace_back(new test_flash_attn_ext(576, 512, 1, {16, 2}, 4096, 2, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16, {0, 1, 2, 3}, true, true, 768));