Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
3 changes: 2 additions & 1 deletion docs/build.md
Original file line number Diff line number Diff line change
Expand Up @@ -323,7 +323,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_PEER_MAX_BATCH_SIZE | Positive integer | 128 | Maximum batch size for which to enable peer access between multiple GPUs. Peer access requires either Linux or NVLink. When using NVLink enabling peer access for larger batch sizes is potentially beneficial. |
| 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

Expand Down
2 changes: 2 additions & 0 deletions ggml/CMakeLists.txt
Original file line number Diff line number Diff line change
Expand Up @@ -209,6 +209,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
Expand Down
71 changes: 71 additions & 0 deletions ggml/cmake/common.cmake
Original file line number Diff line number Diff line change
Expand Up @@ -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()
13 changes: 2 additions & 11 deletions ggml/src/ggml-cuda/CMakeLists.txt
Original file line number Diff line number Diff line change
Expand Up @@ -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}
Expand Down
202 changes: 103 additions & 99 deletions ggml/src/ggml-cuda/fattn.cu
Original file line number Diff line number Diff line change
Expand Up @@ -257,90 +257,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

GGML_ABORT("fatal error");
const ggml_tensor * Q = dst->src[0];
const ggml_tensor * K = dst->src[1];
const ggml_tensor * V = dst->src[2];

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:
Expand All @@ -351,20 +362,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;
Expand Down Expand Up @@ -455,12 +463,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;
}
Expand Down Expand Up @@ -557,6 +559,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];

Expand All @@ -581,10 +584,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;
}
Expand Down
4 changes: 2 additions & 2 deletions ggml/src/ggml-cuda/ggml-cuda.cu
Original file line number Diff line number Diff line change
Expand Up @@ -5994,8 +5994,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

{
Expand Down
Loading