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