mirror of
https://github.com/ggml-org/whisper.cpp.git
synced 2026-09-21 13:37:35 -05:00
CUDA: replace GGML_FA_ALL_QUANTS with GGML_FA_QUANTS, more control over what is compiled (llama/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>
This commit is contained in:
committed by
Georgi Gerganov
co-authored by
Johannes Gäßler
parent
dca2df4211
commit
4d506f58e4
@@ -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
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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}
|
||||
|
||||
+101
-97
@@ -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_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, 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_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];
|
||||
const ggml_tensor * Q = dst->src[0];
|
||||
const ggml_tensor * K = dst->src[1];
|
||||
const 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
|
||||
|
||||
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;
|
||||
}
|
||||
|
||||
@@ -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
|
||||
|
||||
{
|
||||
|
||||
@@ -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}
|
||||
|
||||
@@ -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})
|
||||
|
||||
Reference in New Issue
Block a user