metal : dequant kv cache only for large batches (llama/27438)
This commit is contained in:
parent
292241818a
commit
60f14a5ecc
|
|
@ -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);
|
||||||
|
|
||||||
|
|
|
||||||
Loading…
Reference in New Issue