BEST_FATTN_KERNEL_MMA_F16 = 400,
};
+static bool ggml_cuda_fattn_kv_type_supported(ggml_type type) {
+ switch (type) {
+ case GGML_TYPE_F32:
+ case GGML_TYPE_F16:
+ return true;
+ 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;
+ }
+}
+
static best_fattn_kernel ggml_cuda_get_best_fattn_kernel(const int device, const ggml_tensor * dst) {
#ifndef FLASH_ATTN_AVAILABLE
GGML_UNUSED(device); GGML_UNUSED(dst);
}
#endif // GGML_CUDA_FA_ALL_QUANTS
- switch (K->type) {
- case GGML_TYPE_F32:
- case GGML_TYPE_F16:
- break;
- case GGML_TYPE_Q4_1:
- case GGML_TYPE_Q5_0:
- case GGML_TYPE_Q5_1:
-#ifndef GGML_CUDA_FA_ALL_QUANTS
- return BEST_FATTN_KERNEL_NONE;
-#endif // GGML_CUDA_FA_ALL_QUANTS
- case GGML_TYPE_Q4_0:
- case GGML_TYPE_Q8_0:
- case GGML_TYPE_BF16:
- break;
- default:
- return BEST_FATTN_KERNEL_NONE;
+ if (!ggml_cuda_fattn_kv_type_supported(K->type) || !ggml_cuda_fattn_kv_type_supported(V->type)) {
+ return BEST_FATTN_KERNEL_NONE;
}
if (mask && mask->ne[2] != 1) {