Commit a043d38a6 for llama.cpp

commit a043d38a62aeacde876dee02b9386cf6c3baf134
Author: Foad Abo Dahood <32059146+masterFoad@users.noreply.github.com>
Date:   Tue Oct 6 16:40:16 2026 +0300

    metal : fix excess threadgroup memory in quantized flash attention (#29340)

diff --git a/ggml/src/ggml-metal/ggml-metal-ops.cpp b/ggml/src/ggml-metal/ggml-metal-ops.cpp
index dfe46bac6..2a14cd0c2 100644
--- a/ggml/src/ggml-metal/ggml-metal-ops.cpp
+++ b/ggml/src/ggml-metal/ggml-metal-ops.cpp
@@ -3088,7 +3088,9 @@ static bool ggml_metal_op_flash_attn_ext_use_kv_f16(const ggml_tensor * op) {
     // depending on compute/bandwidth ratio, dequant to f16 kv is not always beneficial
     // ref: https://github.com/ggml-org/llama.cpp/pull/27390#issuecomment-5355152767
     // TODO: tune per device
-    if (op->src[0]->ne[1] < 32) {
+    // large heads need the upfront dequant to fit the non-vec threadgroup memory
+    if (op->src[0]->ne[1] < 32 &&
+        (op->src[0]->ne[0] < 512 || ggml_metal_op_flash_attn_ext_use_vec(op))) {
         return false;
     }

@@ -3827,6 +3829,8 @@ int ggml_metal_op_flash_attn_ext(ggml_metal_op_t ctx, int idx) {
         ggml_metal_encoder_set_buffer  (enc, bid_blk,  7);
         ggml_metal_encoder_set_buffer  (enc, bid_dst,  8);

+        GGML_ASSERT(smem <= props_dev->max_theadgroup_memory_size);
+
         ggml_metal_encoder_set_threadgroup_memory_size(enc, smem, 0);

         ggml_metal_encoder_dispatch_threadgroups(enc, (ne01 + nqptg - 1)/nqptg, ne02, ne03, 32, nsg, 1);
diff --git a/tests/test-backend-ops.cpp b/tests/test-backend-ops.cpp
index d8893025e..51d309b46 100644
--- a/tests/test-backend-ops.cpp
+++ b/tests/test-backend-ops.cpp
@@ -11327,6 +11327,8 @@ static std::vector<std::unique_ptr<test_case>> make_test_cases_eval() {
     test_cases.emplace_back(new test_flash_attn_ext(576, 512, 1, {8, 1},  113,   1, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_Q8_0, GGML_TYPE_Q8_0, {0, 1, 2, 3}, true, true));
     test_cases.emplace_back(new test_flash_attn_ext(576, 512, 1, {8, 1}, 1024,   1, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_Q8_0, GGML_TYPE_Q8_0, {0, 1, 2, 3}, true, true));
     test_cases.emplace_back(new test_flash_attn_ext(576, 512, 1, {8, 1}, 1024,  64, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_Q8_0, GGML_TYPE_Q8_0, {0, 1, 2, 3}, true, true));
+    test_cases.emplace_back(new test_flash_attn_ext(576, 512, 1, {8, 1}, 4096,  24, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_Q8_0, GGML_TYPE_Q8_0, {0, 1, 2, 3}, true, true));
+    test_cases.emplace_back(new test_flash_attn_ext(512, 512, 1, {8, 1}, 4096,  24, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_Q8_0, GGML_TYPE_Q8_0, {0, 1, 2, 3}, true, false));

     // 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));