metal : dequant kv cache only for large batches (llama/27438)

This commit is contained in:
Georgi Gerganov 2026-08-20 17:00:54 +03:00
parent 292241818a
commit 60f14a5ecc
1 changed files with 11 additions and 3 deletions

View File

@ -2806,6 +2806,13 @@ bool ggml_metal_op_flash_attn_ext_use_vec(const ggml_tensor * op) {
static bool ggml_metal_op_flash_attn_ext_use_kv_f16(const ggml_tensor * op) { static bool ggml_metal_op_flash_attn_ext_use_kv_f16(const ggml_tensor * op) {
assert(op->op == GGML_OP_FLASH_ATTN_EXT); assert(op->op == GGML_OP_FLASH_ATTN_EXT);
// 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) {
return false;
}
switch (op->src[1]->type) { switch (op->src[1]->type) {
case GGML_TYPE_Q4_0: case GGML_TYPE_Q4_0:
case GGML_TYPE_Q4_1: case GGML_TYPE_Q4_1:
@ -2968,9 +2975,10 @@ size_t ggml_metal_op_flash_attn_ext_extra_tmp(const ggml_tensor * op) {
size_t ggml_metal_op_flash_attn_ext_extra_kv_f16(const ggml_tensor * op) { size_t ggml_metal_op_flash_attn_ext_extra_kv_f16(const ggml_tensor * op) {
assert(op->op == GGML_OP_FLASH_ATTN_EXT); assert(op->op == GGML_OP_FLASH_ATTN_EXT);
if (!ggml_metal_op_flash_attn_ext_use_kv_f16(op)) { // note: always reserve the temp buffer to avoid graph reallocations
return 0; //if (!ggml_metal_op_flash_attn_ext_use_kv_f16(op)) {
} // return 0;
//}
GGML_TENSOR_LOCALS( int32_t, ne2, op->src[2], ne); GGML_TENSOR_LOCALS( int32_t, ne2, op->src[2], ne);