Skip to content
Closed
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
4 changes: 3 additions & 1 deletion ggml/src/ggml-cuda/fattn-mma-f16.cuh
Original file line number Diff line number Diff line change
Expand Up @@ -1757,7 +1757,8 @@ static __device__ __forceinline__ void flash_attn_ext_f16_process_tile(

static constexpr __host__ __device__ bool ggml_cuda_flash_attn_ext_mma_f16_may_use_sparse(
const int DKQ, const int DV, const int ncols1, const int ncols2) {
return (DKQ == 512 && DV == 512 && ncols1 == 1 && ncols2 == 8) ||
return (DKQ == 512 && DV == 512 && ncols1 == 1 && ncols2 == 8) ||
(DKQ == 512 && DV == 512 && ncols1 == 1 && ncols2 == 32) || // Volta: >= 32 columns per launch
(DKQ == 576 && DV == 512 && ncols1 == 1 && ncols2 == 16);
}

Expand Down Expand Up @@ -2118,6 +2119,7 @@ extern DECL_FATTN_MMA_F16_CASE(512, 512, 1, 8);
extern DECL_FATTN_MMA_F16_CASE(512, 512, 2, 8);
extern DECL_FATTN_MMA_F16_CASE(512, 512, 4, 8);
extern DECL_FATTN_MMA_F16_CASE(512, 512, 8, 8);
extern DECL_FATTN_MMA_F16_CASE(512, 512, 1, 32); // sparse on Volta

// The number of viable configurations for Deepseek is very limited:
extern DECL_FATTN_MMA_F16_CASE(576, 512, 1, 16);
Expand Down
36 changes: 31 additions & 5 deletions ggml/src/ggml-cuda/fattn.cu
Original file line number Diff line number Diff line change
Expand Up @@ -106,37 +106,51 @@ void ggml_cuda_flash_attn_ext_compact_mask(
#endif // !defined(GGML_USE_HIP) && !defined(GGML_USE_MUSA)
}

bool ggml_cuda_flash_attn_ext_mma_f16_shall_use_sparse(ggml_backend_cuda_context & ctx, ggml_tensor * dst) {
// Volta can use the sparse kernels as well, but only via the <512, 512, 1, 32> instance:
// the sparse gather needs one index row per launch (ncols1 == 1) and the Volta build has no device code
// for launches with fewer than 32 columns, so all GQA heads of one query row go into a single launch.
static bool ggml_cuda_fattn_volta_sparse_ok(const int cc, const ggml_tensor * dst) {
const ggml_tensor * Q = dst->src[0];
const ggml_tensor * K = dst->src[1];
const ggml_tensor * V = dst->src[2];
return cc == GGML_CUDA_CC_VOLTA && Q->ne[0] == 512 && V->ne[0] == 512 &&
K->ne[2] > 0 && Q->ne[2] % K->ne[2] == 0 && (Q->ne[2] / K->ne[2]) % 32 == 0;
}

static bool ggml_cuda_fattn_sparse_applicable(const int cc, const ggml_tensor * dst) {
#if defined(GGML_USE_HIP) || defined(GGML_USE_MUSA)
GGML_UNUSED_VARS(ctx, dst);
GGML_UNUSED_VARS(cc, dst);
return false;
#else
const ggml_tensor * Q = dst->src[0];
const ggml_tensor * K = dst->src[1];
const ggml_tensor * mask = dst->src[3];
const int cc = ggml_cuda_info().devices[ctx.device].cc;

float max_bias = 0.0f;
float logit_softcap = 0.0f;
memcpy(&max_bias, (const float *) dst->op_params + 1, sizeof(float));
memcpy(&logit_softcap, (const float *) dst->op_params + 2, sizeof(float));

const int32_t n_kv_max = ggml_get_op_params_i32(dst, 4);
return GGML_CUDA_CC_IS_NVIDIA(cc) && turing_mma_available(cc) &&
return GGML_CUDA_CC_IS_NVIDIA(cc) && (turing_mma_available(cc) || ggml_cuda_fattn_volta_sparse_ok(cc, dst)) &&
mask != nullptr && n_kv_max > 0 && max_bias == 0.0f && logit_softcap == 0.0f &&
mask->ne[0] == K->ne[1] && mask->ne[1] >= Q->ne[1] && mask->ne[2] == 1 &&
K->ne[1] >= std::max<int64_t>(4096, 2LL*n_kv_max);
#endif // !defined(GGML_USE_HIP) && !defined(GGML_USE_MUSA)
}

bool ggml_cuda_flash_attn_ext_mma_f16_shall_use_sparse(ggml_backend_cuda_context & ctx, ggml_tensor * dst) {
return ggml_cuda_fattn_sparse_applicable(ggml_cuda_info().devices[ctx.device].cc, dst);
}

template <int DKQ, int DV, int ncols2>
static void ggml_cuda_flash_attn_ext_mma_f16_switch_ncols1(ggml_backend_cuda_context & ctx, ggml_tensor * dst) {
const int cc = ggml_cuda_info().devices[ggml_cuda_get_device()].cc;
const ggml_tensor * Q = dst->src[0];

#if !defined(GGML_USE_HIP) && !defined(GGML_USE_MUSA)
if constexpr (ggml_cuda_flash_attn_ext_mma_f16_may_use_sparse(DKQ, DV, 1, ncols2)) {
if (ggml_cuda_flash_attn_ext_mma_f16_shall_use_sparse(ctx, dst)) {
if ((cc != GGML_CUDA_CC_VOLTA || ncols2 >= 32) && ggml_cuda_flash_attn_ext_mma_f16_shall_use_sparse(ctx, dst)) {
ggml_cuda_flash_attn_ext_mma_f16_case<DKQ, DV, 1, ncols2>(ctx, dst);
return;
}
Expand Down Expand Up @@ -198,6 +212,13 @@ static void ggml_cuda_flash_attn_ext_mma_f16_switch_ncols2(ggml_backend_cuda_con

// On Volta the GQA optimizations aren't as impactful vs. minimizing wasted compute:
if (cc == GGML_CUDA_CC_VOLTA) {
if constexpr (ggml_cuda_flash_attn_ext_mma_f16_may_use_sparse(DKQ, DV, 1, 32)) {
if (use_gqa_opt && gqa_ratio % 32 == 0 && ggml_cuda_flash_attn_ext_mma_f16_shall_use_sparse(ctx, dst)) {
ggml_cuda_flash_attn_ext_mma_f16_case<DKQ, DV, 1, 32>(ctx, dst);
return;
}
}

if (use_gqa_opt && gqa_ratio % 8 == 0) {
ggml_cuda_flash_attn_ext_mma_f16_switch_ncols1<DKQ, DV, 8>(ctx, dst);
return;
Expand Down Expand Up @@ -642,6 +663,11 @@ static best_fattn_kernel ggml_cuda_get_best_fattn_kernel(const int device, const
}

if (volta_mma_available(cc) && Q->ne[0] != 40 && Q->ne[0] != 72) {
// a sparse launch reads n_kv_max entries instead of all of K, which beats the dense tile kernel
// already for a single token as soon as the KV cache is much larger than n_kv_max
if (gqa_opt_applies && ggml_cuda_fattn_sparse_applicable(cc, dst)) {
return BEST_FATTN_KERNEL_MMA_F16;
}
if (can_use_vector_kernel && Q->ne[1] * gqa_ratio_eff <= 2) {
return BEST_FATTN_KERNEL_VEC;
}
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -3,4 +3,5 @@
#include "../fattn-mma-f16.cuh"

DECL_FATTN_MMA_F16_CASE(320, 256, 1, 32);
DECL_FATTN_MMA_F16_CASE(512, 512, 1, 32);
DECL_FATTN_MMA_F16_CASE(576, 512, 1, 32);
6 changes: 4 additions & 2 deletions ggml/src/ggml-cuda/template-instances/generate_cu_files.py
Original file line number Diff line number Diff line change
Expand Up @@ -93,11 +93,13 @@ def get_short_name(long_quant_name):
continue
if head_size_kq == 320 and ncols2 != 32: # Mistral Small 4
continue
if head_size_kq == 512 and ncols2 not in (2, 4, 8): # Gemma 4 (+ MTP)
if head_size_kq == 512 and ncols2 not in (2, 4, 8, 32): # Gemma 4 (+ MTP); 32: sparse FA on Volta
continue
if head_size_kq == 512 and ncols2 == 32 and ncols1 != 1: # sparse gathers one index row per launch
continue
if head_size_kq == 576 and ncols2 not in (4, 16, 32): # Deepseek, GLM 4.7 Flash
continue
if head_size_kq not in (192, 320, 576) and ncols2 in (16, 32):
if head_size_kq not in (192, 320, 512, 576) and ncols2 in (16, 32):
continue
head_size_v = HEAD_SIZES_V_OVERRIDE.get(head_size_kq, head_size_kq)
f.write(SOURCE_FATTN_MMA_CASE.format(ncols1=ncols1, ncols2=ncols2, head_size_kq=head_size_kq, head_size_v=head_size_v))
Expand Down
4 changes: 4 additions & 0 deletions tests/test-backend-ops.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -10684,6 +10684,10 @@ static std::vector<std::unique_ptr<test_case>> make_test_cases_eval() {

// Sparse mask hint: supported decode/prefill layouts and dense fallbacks.
test_cases.emplace_back(new test_flash_attn_ext(512, 512, 1, { 8, 1}, 4096, 1, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16, {0, 1, 2, 3}, true, false, 512));
// GQA 64 on 1 KV head + sinks: the <512, 512, 1, 32> sparse launch (used on Volta)
test_cases.emplace_back(new test_flash_attn_ext(512, 512, 1, {64, 1}, 8192, 32, true, true, 0, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16, {0, 1, 2, 3}, true, false, 640));
test_cases.emplace_back(new test_flash_attn_ext(512, 512, 1, {64, 1}, 8192, 32, true, true, 0, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16, {0, 1, 2, 3}, true, true, 640));
test_cases.emplace_back(new test_flash_attn_ext(512, 512, 1, {64, 1}, 8192, 1, true, true, 0, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16, {0, 1, 2, 3}, true, false, 640));
test_cases.emplace_back(new test_flash_attn_ext(512, 512, 1, { 8, 2}, 4096, 3, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16, {0, 1, 2, 3}, true, false, 768));
test_cases.emplace_back(new test_flash_attn_ext(576, 512, 1, {16, 1}, 4096, 1, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16, {0, 1, 2, 3}, true, true, 512));
test_cases.emplace_back(new test_flash_attn_ext(576, 512, 1, {16, 2}, 4096, 2, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16, {0, 1, 2, 3}, true, true, 768));
Expand Down