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:
Piotr Wilkin (ilintar)
2026-09-14 20:45:06 +03:00
committed by Georgi Gerganov
co-authored by Johannes Gäßler
parent dca2df4211
commit 4d506f58e4
7 changed files with 182 additions and 132 deletions
+2
View File
@@ -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
+71
View File
@@ -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()
+2 -11
View File
@@ -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
View File
@@ -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;
}
+2 -2
View File
@@ -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
{
+2 -11
View File
@@ -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}
+2 -11
View File
@@ -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})