Commit 5a4d0feca for llama.cpp
commit 5a4d0fecae272c9caf0b32eb384fa6a58dddb560
Author: Piotr Wilkin (ilintar) <piotr.wilkin@syndatis.com>
Date: Wed Sep 9 12:50:08 2026 +0200
CUDA: replace GGML_FA_ALL_QUANTS with GGML_FA_QUANTS, more control over what is compiled (#28079)
* CUDA: add configurable FA quant combinations
Assisted-by: Codex
* remove all flags but , add runtime fallback with warning for uncompiled combination
* Update docs/build.md
Co-authored-by: Johannes Gäßler <johannesg@5d6.de>
* apply code review comments
---------
Co-authored-by: Johannes Gäßler <johannesg@5d6.de>
diff --git a/docs/build.md b/docs/build.md
index f794d490b..28dcbc2e5 100644
--- a/docs/build.md
+++ b/docs/build.md
@@ -300,7 +300,8 @@ The following compilation options are also available to tweak performance:
|-------------------------------|------------------------|---------|----------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------|
| GGML_CUDA_FORCE_MMQ | Boolean | false | Force the use of custom matrix multiplication kernels for quantized models instead of FP16 cuBLAS even if there is no int8 tensor core implementation available (affects V100, CDNA and RDNA3+). MMQ kernels are enabled by default on GPUs with int8 tensor core support. With MMQ force enabled, speed for large batch sizes will be worse but VRAM consumption will be lower. |
| GGML_CUDA_FORCE_CUBLAS | Boolean | false | Force the use of FP16 cuBLAS instead of custom matrix multiplication kernels for quantized models. There may be issues with numerical overflows (except for V100, CDNA and RDNA4 which use FP32 compute type by default) and memory use will be higher. Prompt processing may become faster on recent datacenter GPUs (the custom kernels were tuned primarily for RTX 3000/4000). |
-| GGML_CUDA_FA_ALL_QUANTS | Boolean | false | Compile support for all KV cache quantization type (combinations) for the FlashAttention CUDA kernels. More fine-grained control over KV cache size but compilation takes much longer. |
+| GGML_CUDA_FA_QUANTS | `all` or `type_K-type_V` list | q4_0-q4_0;q8_0-q8_0;f16-f16;bf16-bf16 | Select which K/V type combinations to compile the FlashAttention CUDA kernels for. `all` compiles every combination, but compilation takes much longer. Otherwise a `;`-separated list of `type_K-type_V` pairs; f16-f16 is always compiled. Combinations that were not compiled fall back to f16-f16 kernel with a warning. Legal types: f16, bf16, q4_0, q4_1, q5_0, q5_1, q8_0. |
+| GGML_CUDA_FA_ALL_QUANTS | Boolean | false | Deprecated alias for `GGML_CUDA_FA_QUANTS=all`. |
## MUSA
diff --git a/ggml/CMakeLists.txt b/ggml/CMakeLists.txt
index d76ed8ab0..ba9bc83b9 100644
--- a/ggml/CMakeLists.txt
+++ b/ggml/CMakeLists.txt
@@ -204,6 +204,8 @@ option(GGML_CUDA_NO_PEER_COPY "ggml: do not use peer to peer copie
option(GGML_CUDA_NO_VMM "ggml: do not try to use CUDA VMM" OFF)
option(GGML_CUDA_FA "ggml: compile ggml FlashAttention CUDA kernels" ON)
option(GGML_CUDA_FA_ALL_QUANTS "ggml: compile all quants for FlashAttention" OFF)
+set (GGML_CUDA_FA_QUANTS "q4_0-q4_0;q8_0-q8_0;f16-f16;bf16-bf16" CACHE STRING
+ "ggml: FlashAttention K-V type combinations to compile, \"all\" or a list such as \"q8_0-q8_0;q8_0-q4_0\"")
option(GGML_CUDA_GRAPHS "ggml: use CUDA graphs (llama.cpp only)" ${GGML_CUDA_GRAPHS_DEFAULT})
option(GGML_CUDA_NCCL "ggml: use NVIDIA Collective Comm. Library" ON)
set (GGML_CUDA_COMPRESSION_MODE "size" CACHE STRING
diff --git a/ggml/cmake/common.cmake b/ggml/cmake/common.cmake
index cb6638833..25eff7a5e 100644
--- a/ggml/cmake/common.cmake
+++ b/ggml/cmake/common.cmake
@@ -48,3 +48,74 @@ function(ggml_get_system_arch)
set(GGML_SYSTEM_ARCH "UNKNOWN" PARENT_SCOPE)
endif()
endfunction()
+
+# Determines which FlashAttention vector kernel template instances to compile, returns them in OUT_SRCS.
+function(ggml_cuda_fattn_vec_instances DIR OUT_SRCS)
+ set(FA_TYPES q4_0 q4_1 q5_0 q5_1 q8_0 bf16 f16)
+
+ string(TOLOWER "${GGML_CUDA_FA_QUANTS}" FA_QUANTS)
+ string(STRIP "${FA_QUANTS}" FA_QUANTS)
+ if (GGML_CUDA_FA_ALL_QUANTS)
+ message(WARNING "GGML_CUDA_FA_ALL_QUANTS is deprecated, use GGML_CUDA_FA_QUANTS=all instead")
+ set(FA_QUANTS all)
+ endif()
+ if (NOT FA_QUANTS)
+ message(FATAL_ERROR "GGML_CUDA_FA_QUANTS must not be empty")
+ endif()
+
+ if (FA_QUANTS STREQUAL "all")
+ set(FA_COMBINATIONS "")
+ foreach (TYPE_V IN LISTS FA_TYPES)
+ foreach (TYPE_K IN LISTS FA_TYPES)
+ list(APPEND FA_COMBINATIONS ${TYPE_K}-${TYPE_V})
+ endforeach()
+ endforeach()
+ else()
+ set(FA_COMBINATIONS f16-f16)
+
+ string(REPLACE "," ";" FA_SELECTED "${FA_QUANTS}")
+ foreach (COMBINATION IN LISTS FA_SELECTED)
+ string(STRIP "${COMBINATION}" COMBINATION)
+ if (NOT COMBINATION MATCHES "^([a-z0-9_]+)-([a-z0-9_]+)$")
+ message(FATAL_ERROR "GGML_CUDA_FA_QUANTS: \"${COMBINATION}\" is not \"all\" or a <type_K>-<type_V> combination")
+ endif()
+ set(TYPE_K ${CMAKE_MATCH_1})
+ set(TYPE_V ${CMAKE_MATCH_2})
+ foreach (TYPE ${TYPE_K} ${TYPE_V})
+ if (NOT TYPE IN_LIST FA_TYPES)
+ message(FATAL_ERROR
+ "GGML_CUDA_FA_QUANTS: unknown type \"${TYPE}\" in \"${COMBINATION}\", must be one of: ${FA_TYPES}")
+ endif()
+ endforeach()
+ list(APPEND FA_COMBINATIONS ${TYPE_K}-${TYPE_V})
+ endforeach()
+ endif()
+ list(REMOVE_DUPLICATES FA_COMBINATIONS)
+
+ string(REPLACE ";" "," FA_QUANTS_DEFINE "${FA_QUANTS}")
+ add_compile_definitions(GGML_CUDA_FA_QUANTS="${FA_QUANTS_DEFINE}")
+ foreach (TYPE_V IN LISTS FA_TYPES)
+ foreach (TYPE_K IN LISTS FA_TYPES)
+ if ("${TYPE_K}-${TYPE_V}" IN_LIST FA_COMBINATIONS)
+ set(COMPILED 1)
+ else()
+ set(COMPILED 0)
+ endif()
+ string(TOUPPER "GGML_CUDA_FA_${TYPE_K}_${TYPE_V}" COMBINATION_DEF)
+ add_compile_definitions(${COMBINATION_DEF}=${COMPILED})
+ endforeach()
+ endforeach()
+
+ message(STATUS "FlashAttention K-V type combinations: ${FA_COMBINATIONS}")
+
+ set(SRCS "")
+ foreach (COMBINATION IN LISTS FA_COMBINATIONS)
+ set(SRC "${DIR}/template-instances/fattn-vec-instance-${COMBINATION}.cu")
+ if (NOT EXISTS "${SRC}")
+ message(FATAL_ERROR "FlashAttention template instance \"${SRC}\" does not exist")
+ endif()
+ list(APPEND SRCS "${SRC}")
+ endforeach()
+
+ set(${OUT_SRCS} ${SRCS} PARENT_SCOPE)
+endfunction()
diff --git a/ggml/src/ggml-cuda/CMakeLists.txt b/ggml/src/ggml-cuda/CMakeLists.txt
index 10828ad81..2254090cb 100644
--- a/ggml/src/ggml-cuda/CMakeLists.txt
+++ b/ggml/src/ggml-cuda/CMakeLists.txt
@@ -112,17 +112,8 @@ if (CUDAToolkit_FOUND)
file(GLOB SRCS "template-instances/mmf*.cu")
list(APPEND GGML_SOURCES_CUDA ${SRCS})
- if (GGML_CUDA_FA_ALL_QUANTS)
- file(GLOB SRCS "template-instances/fattn-vec*.cu")
- list(APPEND GGML_SOURCES_CUDA ${SRCS})
- add_compile_definitions(GGML_CUDA_FA_ALL_QUANTS)
- else()
- list(APPEND GGML_SOURCES_CUDA
- template-instances/fattn-vec-instance-f16-f16.cu
- template-instances/fattn-vec-instance-q4_0-q4_0.cu
- template-instances/fattn-vec-instance-q8_0-q8_0.cu
- template-instances/fattn-vec-instance-bf16-bf16.cu)
- endif()
+ ggml_cuda_fattn_vec_instances(${CMAKE_CURRENT_SOURCE_DIR} SRCS)
+ list(APPEND GGML_SOURCES_CUDA ${SRCS})
ggml_add_backend_library(ggml-cuda
${GGML_HEADERS_CUDA}
diff --git a/ggml/src/ggml-cuda/fattn.cu b/ggml/src/ggml-cuda/fattn.cu
index ae217fbd9..d11a964d5 100644
--- a/ggml/src/ggml-cuda/fattn.cu
+++ b/ggml/src/ggml-cuda/fattn.cu
@@ -374,90 +374,101 @@ static void ggml_cuda_flash_attn_ext_mma_f16(ggml_backend_cuda_context & ctx, gg
}
}
-#define FATTN_VEC_CASE(D, type_K, type_V) \
- { \
- const bool type_K_okay = K->type == (type_K) || (K->type == GGML_TYPE_F32 && (type_K) == GGML_TYPE_F16); \
- const bool type_V_okay = V->type == (type_V) || (V->type == GGML_TYPE_F32 && (type_V) == GGML_TYPE_F16); \
- if (Q->ne[0] == (D) && type_K_okay && type_V_okay) { \
- ggml_cuda_flash_attn_ext_vec_case<D, type_K, type_V>(ctx, dst); \
- return; \
- } \
- } \
-
-#define FATTN_VEC_CASES_ALL_D(type_K, type_V) \
- FATTN_VEC_CASE( 64, type_K, type_V) \
- FATTN_VEC_CASE(128, type_K, type_V) \
- FATTN_VEC_CASE(256, type_K, type_V) \
+#define FATTN_VEC_CASE(D, type_K_case, type_V_case) \
+ if constexpr (GGML_CUDA_FA_##type_K_case##_##type_V_case) { \
+ const bool type_K_okay = type_K == GGML_TYPE_##type_K_case || (type_K == GGML_TYPE_F32 && GGML_TYPE_##type_K_case == GGML_TYPE_F16); \
+ const bool type_V_okay = type_V == GGML_TYPE_##type_V_case || (type_V == GGML_TYPE_F32 && GGML_TYPE_##type_V_case == GGML_TYPE_F16); \
+ if (head_size == (D) && type_K_okay && type_V_okay) { \
+ return ggml_cuda_flash_attn_ext_vec_case<D, GGML_TYPE_##type_K_case, GGML_TYPE_##type_V_case>; \
+ } \
+ } \
+
+#define FATTN_VEC_CASES_ALL_D(type_K_case, type_V_case) \
+ FATTN_VEC_CASE( 64, type_K_case, type_V_case) \
+ FATTN_VEC_CASE(128, type_K_case, type_V_case) \
+ FATTN_VEC_CASE(256, type_K_case, type_V_case) \
+
+typedef void (* fattn_vec_case_t)(ggml_backend_cuda_context & ctx, ggml_tensor * dst);
+
+// Vector kernel for the given head size and K/V types, nullptr if its template instance was not compiled:
+static fattn_vec_case_t ggml_cuda_get_fattn_vec_case(const int64_t head_size, const ggml_type type_K, const ggml_type type_V) {
+ FATTN_VEC_CASES_ALL_D(F16, F16)
+ FATTN_VEC_CASES_ALL_D(Q4_0, F16)
+ FATTN_VEC_CASES_ALL_D(Q4_1, F16)
+ FATTN_VEC_CASES_ALL_D(Q5_0, F16)
+ FATTN_VEC_CASES_ALL_D(Q5_1, F16)
+ FATTN_VEC_CASES_ALL_D(Q8_0, F16)
+ FATTN_VEC_CASES_ALL_D(BF16, F16)
+
+ FATTN_VEC_CASES_ALL_D(F16, Q4_0)
+ FATTN_VEC_CASES_ALL_D(Q4_0, Q4_0)
+ FATTN_VEC_CASES_ALL_D(Q4_1, Q4_0)
+ FATTN_VEC_CASES_ALL_D(Q5_0, Q4_0)
+ FATTN_VEC_CASES_ALL_D(Q5_1, Q4_0)
+ FATTN_VEC_CASES_ALL_D(Q8_0, Q4_0)
+ FATTN_VEC_CASES_ALL_D(BF16, Q4_0)
+
+ FATTN_VEC_CASES_ALL_D(F16, Q4_1)
+ FATTN_VEC_CASES_ALL_D(Q4_0, Q4_1)
+ FATTN_VEC_CASES_ALL_D(Q4_1, Q4_1)
+ FATTN_VEC_CASES_ALL_D(Q5_0, Q4_1)
+ FATTN_VEC_CASES_ALL_D(Q5_1, Q4_1)
+ FATTN_VEC_CASES_ALL_D(Q8_0, Q4_1)
+ FATTN_VEC_CASES_ALL_D(BF16, Q4_1)
+
+ FATTN_VEC_CASES_ALL_D(F16, Q5_0)
+ FATTN_VEC_CASES_ALL_D(Q4_0, Q5_0)
+ FATTN_VEC_CASES_ALL_D(Q4_1, Q5_0)
+ FATTN_VEC_CASES_ALL_D(Q5_0, Q5_0)
+ FATTN_VEC_CASES_ALL_D(Q5_1, Q5_0)
+ FATTN_VEC_CASES_ALL_D(Q8_0, Q5_0)
+ FATTN_VEC_CASES_ALL_D(BF16, Q5_0)
+
+ FATTN_VEC_CASES_ALL_D(F16, Q5_1)
+ FATTN_VEC_CASES_ALL_D(Q4_0, Q5_1)
+ FATTN_VEC_CASES_ALL_D(Q4_1, Q5_1)
+ FATTN_VEC_CASES_ALL_D(Q5_0, Q5_1)
+ FATTN_VEC_CASES_ALL_D(Q5_1, Q5_1)
+ FATTN_VEC_CASES_ALL_D(Q8_0, Q5_1)
+ FATTN_VEC_CASES_ALL_D(BF16, Q5_1)
+
+ FATTN_VEC_CASES_ALL_D(F16, Q8_0)
+ FATTN_VEC_CASES_ALL_D(Q4_0, Q8_0)
+ FATTN_VEC_CASES_ALL_D(Q4_1, Q8_0)
+ FATTN_VEC_CASES_ALL_D(Q5_0, Q8_0)
+ FATTN_VEC_CASES_ALL_D(Q5_1, Q8_0)
+ FATTN_VEC_CASES_ALL_D(Q8_0, Q8_0)
+ FATTN_VEC_CASES_ALL_D(BF16, Q8_0)
+
+ FATTN_VEC_CASES_ALL_D(F16, BF16)
+ FATTN_VEC_CASES_ALL_D(Q4_0, BF16)
+ FATTN_VEC_CASES_ALL_D(Q4_1, BF16)
+ FATTN_VEC_CASES_ALL_D(Q5_0, BF16)
+ FATTN_VEC_CASES_ALL_D(Q5_1, BF16)
+ FATTN_VEC_CASES_ALL_D(Q8_0, BF16)
+ FATTN_VEC_CASES_ALL_D(BF16, BF16)
+
+ return nullptr;
+}
static void ggml_cuda_flash_attn_ext_vec(ggml_backend_cuda_context & ctx, ggml_tensor * dst) {
- ggml_tensor * Q = dst->src[0];
- ggml_tensor * K = dst->src[1];
- ggml_tensor * V = dst->src[2];
-
-#ifdef GGML_CUDA_FA_ALL_QUANTS
- FATTN_VEC_CASES_ALL_D(GGML_TYPE_F16, GGML_TYPE_F16)
- FATTN_VEC_CASES_ALL_D(GGML_TYPE_Q4_0, GGML_TYPE_F16)
- FATTN_VEC_CASES_ALL_D(GGML_TYPE_Q4_1, GGML_TYPE_F16)
- FATTN_VEC_CASES_ALL_D(GGML_TYPE_Q5_0, GGML_TYPE_F16)
- FATTN_VEC_CASES_ALL_D(GGML_TYPE_Q5_1, GGML_TYPE_F16)
- FATTN_VEC_CASES_ALL_D(GGML_TYPE_Q8_0, GGML_TYPE_F16)
- FATTN_VEC_CASES_ALL_D(GGML_TYPE_BF16, GGML_TYPE_F16)
-
- FATTN_VEC_CASES_ALL_D(GGML_TYPE_F16, GGML_TYPE_Q4_0)
- FATTN_VEC_CASES_ALL_D(GGML_TYPE_Q4_0, GGML_TYPE_Q4_0)
- FATTN_VEC_CASES_ALL_D(GGML_TYPE_Q4_1, GGML_TYPE_Q4_0)
- FATTN_VEC_CASES_ALL_D(GGML_TYPE_Q5_0, GGML_TYPE_Q4_0)
- FATTN_VEC_CASES_ALL_D(GGML_TYPE_Q5_1, GGML_TYPE_Q4_0)
- FATTN_VEC_CASES_ALL_D(GGML_TYPE_Q8_0, GGML_TYPE_Q4_0)
- FATTN_VEC_CASES_ALL_D(GGML_TYPE_BF16, GGML_TYPE_Q4_0)
-
- FATTN_VEC_CASES_ALL_D(GGML_TYPE_F16, GGML_TYPE_Q4_1)
- FATTN_VEC_CASES_ALL_D(GGML_TYPE_Q4_0, GGML_TYPE_Q4_1)
- FATTN_VEC_CASES_ALL_D(GGML_TYPE_Q4_1, GGML_TYPE_Q4_1)
- FATTN_VEC_CASES_ALL_D(GGML_TYPE_Q5_0, GGML_TYPE_Q4_1)
- FATTN_VEC_CASES_ALL_D(GGML_TYPE_Q5_1, GGML_TYPE_Q4_1)
- FATTN_VEC_CASES_ALL_D(GGML_TYPE_Q8_0, GGML_TYPE_Q4_1)
- FATTN_VEC_CASES_ALL_D(GGML_TYPE_BF16, GGML_TYPE_Q4_1)
-
- FATTN_VEC_CASES_ALL_D(GGML_TYPE_F16, GGML_TYPE_Q5_0)
- FATTN_VEC_CASES_ALL_D(GGML_TYPE_Q4_0, GGML_TYPE_Q5_0)
- FATTN_VEC_CASES_ALL_D(GGML_TYPE_Q4_1, GGML_TYPE_Q5_0)
- FATTN_VEC_CASES_ALL_D(GGML_TYPE_Q5_0, GGML_TYPE_Q5_0)
- FATTN_VEC_CASES_ALL_D(GGML_TYPE_Q5_1, GGML_TYPE_Q5_0)
- FATTN_VEC_CASES_ALL_D(GGML_TYPE_Q8_0, GGML_TYPE_Q5_0)
- FATTN_VEC_CASES_ALL_D(GGML_TYPE_BF16, GGML_TYPE_Q5_0)
-
- FATTN_VEC_CASES_ALL_D(GGML_TYPE_F16, GGML_TYPE_Q5_1)
- FATTN_VEC_CASES_ALL_D(GGML_TYPE_Q4_0, GGML_TYPE_Q5_1)
- FATTN_VEC_CASES_ALL_D(GGML_TYPE_Q4_1, GGML_TYPE_Q5_1)
- FATTN_VEC_CASES_ALL_D(GGML_TYPE_Q5_0, GGML_TYPE_Q5_1)
- FATTN_VEC_CASES_ALL_D(GGML_TYPE_Q5_1, GGML_TYPE_Q5_1)
- FATTN_VEC_CASES_ALL_D(GGML_TYPE_Q8_0, GGML_TYPE_Q5_1)
- FATTN_VEC_CASES_ALL_D(GGML_TYPE_BF16, GGML_TYPE_Q5_1)
-
- FATTN_VEC_CASES_ALL_D(GGML_TYPE_F16, GGML_TYPE_Q8_0)
- FATTN_VEC_CASES_ALL_D(GGML_TYPE_Q4_0, GGML_TYPE_Q8_0)
- FATTN_VEC_CASES_ALL_D(GGML_TYPE_Q4_1, GGML_TYPE_Q8_0)
- FATTN_VEC_CASES_ALL_D(GGML_TYPE_Q5_0, GGML_TYPE_Q8_0)
- FATTN_VEC_CASES_ALL_D(GGML_TYPE_Q5_1, GGML_TYPE_Q8_0)
- FATTN_VEC_CASES_ALL_D(GGML_TYPE_Q8_0, GGML_TYPE_Q8_0)
- FATTN_VEC_CASES_ALL_D(GGML_TYPE_BF16, GGML_TYPE_Q8_0)
-
- FATTN_VEC_CASES_ALL_D(GGML_TYPE_F16, GGML_TYPE_BF16)
- FATTN_VEC_CASES_ALL_D(GGML_TYPE_Q4_0, GGML_TYPE_BF16)
- FATTN_VEC_CASES_ALL_D(GGML_TYPE_Q4_1, GGML_TYPE_BF16)
- FATTN_VEC_CASES_ALL_D(GGML_TYPE_Q5_0, GGML_TYPE_BF16)
- FATTN_VEC_CASES_ALL_D(GGML_TYPE_Q5_1, GGML_TYPE_BF16)
- FATTN_VEC_CASES_ALL_D(GGML_TYPE_Q8_0, GGML_TYPE_BF16)
- FATTN_VEC_CASES_ALL_D(GGML_TYPE_BF16, GGML_TYPE_BF16)
-#else
- FATTN_VEC_CASES_ALL_D(GGML_TYPE_F16, GGML_TYPE_F16)
- FATTN_VEC_CASES_ALL_D(GGML_TYPE_Q4_0, GGML_TYPE_Q4_0)
- FATTN_VEC_CASES_ALL_D(GGML_TYPE_Q8_0, GGML_TYPE_Q8_0)
- FATTN_VEC_CASES_ALL_D(GGML_TYPE_BF16, GGML_TYPE_BF16)
-#endif // GGML_CUDA_FA_ALL_QUANTS
+ const ggml_tensor * Q = dst->src[0];
+ const ggml_tensor * K = dst->src[1];
+ const ggml_tensor * V = dst->src[2];
- GGML_ABORT("fatal error");
+ fattn_vec_case_t vec_case = ggml_cuda_get_fattn_vec_case(Q->ne[0], K->type, V->type);
+ if (vec_case == nullptr) {
+ static bool warned = false;
+ if (!warned) {
+ GGML_LOG_WARN("%s: no FlashAttention vector kernel compiled for K/V types %s-%s, converting K and V to f16 instead (slow). "
+ "Add \"%s-%s\" to GGML_CUDA_FA_QUANTS to compile it.\n",
+ __func__, ggml_type_name(K->type), ggml_type_name(V->type), ggml_type_name(K->type), ggml_type_name(V->type));
+ warned = true;
+ }
+ vec_case = ggml_cuda_get_fattn_vec_case(Q->ne[0], GGML_TYPE_F16, GGML_TYPE_F16);
+ }
+ GGML_ASSERT(vec_case != nullptr);
+ vec_case(ctx, dst);
}
// Best FlashAttention kernel for a specific GPU:
@@ -468,20 +479,17 @@ enum best_fattn_kernel {
BEST_FATTN_KERNEL_MMA_F16 = 400,
};
-static bool ggml_cuda_fattn_kv_type_supported(ggml_type type) {
+// K/V types for which there is a vector kernel template instance, other kernels convert these to f16:
+static bool ggml_cuda_fattn_kv_type_supported(const ggml_type type) {
switch (type) {
case GGML_TYPE_F32:
case GGML_TYPE_F16:
- return true;
+ case GGML_TYPE_BF16:
+ case GGML_TYPE_Q4_0:
case GGML_TYPE_Q4_1:
case GGML_TYPE_Q5_0:
case GGML_TYPE_Q5_1:
-#ifndef GGML_CUDA_FA_ALL_QUANTS
- return false;
-#endif // GGML_CUDA_FA_ALL_QUANTS
- case GGML_TYPE_Q4_0:
case GGML_TYPE_Q8_0:
- case GGML_TYPE_BF16:
return true;
default:
return false;
@@ -572,12 +580,6 @@ static best_fattn_kernel ggml_cuda_get_best_fattn_kernel(const int device, const
return BEST_FATTN_KERNEL_NONE;
}
-#ifndef GGML_CUDA_FA_ALL_QUANTS
- if (K->type != V->type) {
- return BEST_FATTN_KERNEL_NONE;
- }
-#endif // GGML_CUDA_FA_ALL_QUANTS
-
if (!ggml_cuda_fattn_kv_type_supported(K->type) || !ggml_cuda_fattn_kv_type_supported(V->type)) {
return BEST_FATTN_KERNEL_NONE;
}
@@ -669,6 +671,7 @@ static best_fattn_kernel ggml_cuda_get_best_fattn_kernel(const int device, const
size_t ggml_cuda_flash_attn_ext_get_alloc_size(int device, const ggml_tensor * dst) {
GGML_ASSERT(dst->op == GGML_OP_FLASH_ATTN_EXT);
+ const ggml_tensor * Q = dst->src[0];
const ggml_tensor * K = dst->src[1];
const ggml_tensor * V = dst->src[2];
@@ -686,10 +689,11 @@ size_t ggml_cuda_flash_attn_ext_get_alloc_size(int device, const ggml_tensor * d
need_f16_K = true;
need_f16_V = true;
break;
- case BEST_FATTN_KERNEL_VEC:
- need_f16_K = K->type == GGML_TYPE_F32;
- need_f16_V = V->type == GGML_TYPE_F32;
- break;
+ case BEST_FATTN_KERNEL_VEC: {
+ const bool f16_fallback = ggml_cuda_get_fattn_vec_case(Q->ne[0], K->type, V->type) == nullptr;
+ need_f16_K = K->type == GGML_TYPE_F32 || f16_fallback;
+ need_f16_V = V->type == GGML_TYPE_F32 || f16_fallback;
+ } break;
case BEST_FATTN_KERNEL_NONE:
break;
}
diff --git a/ggml/src/ggml-cuda/ggml-cuda.cu b/ggml/src/ggml-cuda/ggml-cuda.cu
index 38bd4c9a0..5ae3b8d22 100644
--- a/ggml/src/ggml-cuda/ggml-cuda.cu
+++ b/ggml/src/ggml-cuda/ggml-cuda.cu
@@ -5640,8 +5640,8 @@ static ggml_backend_feature * ggml_backend_cuda_get_features(ggml_backend_reg_t
features.push_back({ "USE_GRAPHS", "1" });
#endif
- #ifdef GGML_CUDA_FA_ALL_QUANTS
- features.push_back({ "FA_ALL_QUANTS", "1" });
+ #ifdef GGML_CUDA_FA_QUANTS
+ features.push_back({ "FA_QUANTS", GGML_CUDA_FA_QUANTS });
#endif
{
diff --git a/ggml/src/ggml-hip/CMakeLists.txt b/ggml/src/ggml-hip/CMakeLists.txt
index 47f16f56c..a6a6b7271 100644
--- a/ggml/src/ggml-hip/CMakeLists.txt
+++ b/ggml/src/ggml-hip/CMakeLists.txt
@@ -70,17 +70,8 @@ list(APPEND GGML_SOURCES_ROCM ${SRCS})
file(GLOB SRCS "../ggml-cuda/template-instances/mmf*.cu")
list(APPEND GGML_SOURCES_ROCM ${SRCS})
-if (GGML_CUDA_FA_ALL_QUANTS)
- file(GLOB SRCS "../ggml-cuda/template-instances/fattn-vec*.cu")
- list(APPEND GGML_SOURCES_ROCM ${SRCS})
- add_compile_definitions(GGML_CUDA_FA_ALL_QUANTS)
-else()
- list(APPEND GGML_SOURCES_ROCM
- ../ggml-cuda/template-instances/fattn-vec-instance-f16-f16.cu
- ../ggml-cuda/template-instances/fattn-vec-instance-q4_0-q4_0.cu
- ../ggml-cuda/template-instances/fattn-vec-instance-q8_0-q8_0.cu
- ../ggml-cuda/template-instances/fattn-vec-instance-bf16-bf16.cu)
-endif()
+ggml_cuda_fattn_vec_instances(${CMAKE_CURRENT_SOURCE_DIR}/../ggml-cuda SRCS)
+list(APPEND GGML_SOURCES_ROCM ${SRCS})
ggml_add_backend_library(ggml-hip
${GGML_HEADERS_ROCM}
diff --git a/ggml/src/ggml-musa/CMakeLists.txt b/ggml/src/ggml-musa/CMakeLists.txt
index faf979033..82b754f41 100644
--- a/ggml/src/ggml-musa/CMakeLists.txt
+++ b/ggml/src/ggml-musa/CMakeLists.txt
@@ -43,17 +43,8 @@ if (MUSAToolkit_FOUND)
add_compile_definitions(GGML_MUSA_MUDNN_COPY)
endif()
- if (GGML_CUDA_FA_ALL_QUANTS)
- file(GLOB SRCS "../ggml-cuda/template-instances/fattn-vec*.cu")
- list(APPEND GGML_SOURCES_MUSA ${SRCS})
- add_compile_definitions(GGML_CUDA_FA_ALL_QUANTS)
- else()
- list(APPEND GGML_SOURCES_MUSA
- ../ggml-cuda/template-instances/fattn-vec-instance-f16-f16.cu
- ../ggml-cuda/template-instances/fattn-vec-instance-q4_0-q4_0.cu
- ../ggml-cuda/template-instances/fattn-vec-instance-q8_0-q8_0.cu
- ../ggml-cuda/template-instances/fattn-vec-instance-bf16-bf16.cu)
- endif()
+ ggml_cuda_fattn_vec_instances(${CMAKE_CURRENT_SOURCE_DIR}/../ggml-cuda SRCS)
+ list(APPEND GGML_SOURCES_MUSA ${SRCS})
set_source_files_properties(${GGML_SOURCES_MUSA} PROPERTIES LANGUAGE CXX)
foreach(SOURCE ${GGML_SOURCES_MUSA})