From 555330fbba462f9ebd1c835a4df0b9f2b88c98b0 Mon Sep 17 00:00:00 2001 From: Kawrakow Date: Fri, 28 Aug 2026 18:10:50 +0200 Subject: [PATCH] CUDA: handle GQA = 12 for head size = 256 via new MMA (#2372) * iCUDA: handle GQA = 12 for head size = 256 via new MMA * Cleanup --- ggml/src/ggml-cuda/fattn-new-mma.cu | 2 +- ggml/src/ggml-cuda/fattn.cu | 6 ++++-- 2 files changed, 5 insertions(+), 3 deletions(-) diff --git a/ggml/src/ggml-cuda/fattn-new-mma.cu b/ggml/src/ggml-cuda/fattn-new-mma.cu index 681d4fa4..97274fe7 100644 --- a/ggml/src/ggml-cuda/fattn-new-mma.cu +++ b/ggml/src/ggml-cuda/fattn-new-mma.cu @@ -2277,7 +2277,7 @@ void ggml_cuda_flash_attn_ext_mma_new(ggml_backend_cuda_context & ctx, ggml_tens } if (K->ne[0] == 256) { GGML_ASSERT(Q->ne[0] == 256 && V->ne[0] == 256); - if (gqa_ratio == 6) { + if (gqa_ratio % 6 == 0) { ggml_cuda_flash_attn_ext_mma_f16_case<256, 256, 1, 8>(ctx, dst); } else { GGML_ABORT("Not implemented"); diff --git a/ggml/src/ggml-cuda/fattn.cu b/ggml/src/ggml-cuda/fattn.cu index 4e487e7d..c1ab847d 100644 --- a/ggml/src/ggml-cuda/fattn.cu +++ b/ggml/src/ggml-cuda/fattn.cu @@ -117,13 +117,15 @@ void ggml_cuda_flash_attn_ext(ggml_backend_cuda_context & ctx, ggml_tensor * dst return; } + int gqa_ratio = Q->ne[2] / K->ne[2]; + if (new_mma_available(cc) && K->ne[0] == 128 && V->ne[0] == 128 && Q->ne[0] == 128 && Q->ne[1] == 1 && - (Q->ne[2] / K->ne[2] == 12 || Q->ne[2] / K->ne[2] == 6 || Q->ne[2] / K->ne[2] == 10)) { + (gqa_ratio == 12 || gqa_ratio == 6 || gqa_ratio == 10)) { ggml_cuda_flash_attn_ext_mma_new(ctx, dst); return; } - if (new_mma_available(cc) && K->ne[0] == 256 && V->ne[0] == 256 && Q->ne[0] == 256 && Q->ne[1] == 1 && Q->ne[2] / K->ne[2] == 6) { + if (new_mma_available(cc) && K->ne[0] == 256 && V->ne[0] == 256 && Q->ne[0] == 256 && Q->ne[1] == 1 && (gqa_ratio == 6 || gqa_ratio == 12)) { ggml_cuda_flash_attn_ext_mma_new(ctx, dst); return; }