Commit dac308739 for llama.cpp
commit dac308739429d60ccf99c46ecc1fbb1dff20ca77
Author: Jeff Bolz <jbolz@nvidia.com>
Date: Thu Oct 8 04:36:20 2026 -0500
vulkan: extend sparse FA support to coopmat2 (#30003)
diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp
index 740de5b5a..998cc693c 100644
--- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp
+++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp
@@ -8154,10 +8154,14 @@ void ggml_vk_flash_attn(ggml_backend_vk_context * ctx, vk_context& subctx, const
// cm2 dense is fast, so it needs a larger reduction to win.
// With quantized K/V, sparse only breaks even around 16x (measured on RDNA3/RDNA4).
const int64_t min_ratio = tuning_params.path == FA_COOPMAT2 ? 4 : (kv_f16 ? 2 : 16);
+ // coopmat2 vector decode requires 8B strides.
+ auto sparse_gather_aligned = [](const ggml_tensor * t) {
+ return (t->type != GGML_TYPE_F16 && t->type != GGML_TYPE_BF16) ||
+ (t->nb[1] | t->nb[2] | t->nb[3]) % (4 * sizeof(ggml_fp16_t)) == 0;
+ };
const bool use_sparse = !disable_sparse && n_kv_max > 0 && mask &&
max_bias == 0.0f && logit_softcap == 0.0f &&
- // the cm2 sparse gather only reads f16
- (kv_f16 || tuning_params.path != FA_COOPMAT2) &&
+ (tuning_params.path != FA_COOPMAT2 || (sparse_gather_aligned(k) && sparse_gather_aligned(v))) &&
nem0 == KV &&
(int64_t)KV >= std::max<int64_t>(4096, min_ratio * (int64_t)n_kv_max) &&
(gqa_ratio > 1 || (tuning_params.path == FA_SCALAR && N == 1));
diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_cm2.comp b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_cm2.comp
index c6ed63dd4..c7253784c 100644
--- a/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_cm2.comp
+++ b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_cm2.comp
@@ -18,7 +18,8 @@
#ifdef GL_NV_cooperative_matrix_decode_vector
#extension GL_NV_cooperative_matrix_decode_vector : enable
#endif
-#extension GL_EXT_buffer_reference : enable
+#extension GL_EXT_buffer_reference2 : enable
+#extension GL_EXT_shader_explicit_arithmetic_types_int64 : enable
#extension GL_KHR_shader_subgroup_ballot : enable
#extension GL_KHR_shader_subgroup_vote : enable
#extension GL_EXT_null_initializer : enable
@@ -35,6 +36,10 @@
#define FA_GATHER_BS 1u
#endif
+layout(buffer_reference, std430, buffer_reference_align = 1) buffer decodeBufFA_Byte {
+ uint8_t raw;
+};
+
// buffer_reference stride = sizeof(struct) = FaBlockBytesK/V.
layout(buffer_reference, std430, buffer_reference_align = 1) buffer decodeBufFA_K {
uint8_t raw[FaBlockBytesK];
@@ -113,48 +118,71 @@ layout (binding = 1) readonly buffer K {uint8_t data_k[];};
layout (binding = 2) readonly buffer V {uint8_t data_v[];};
layout (binding = 3) readonly buffer M {uint8_t data_m[];};
-// f16 aliases for the sparse gather callbacks.
-layout (binding = 1) readonly buffer KF16 {float16_t data_kf16[];};
-layout (binding = 2) readonly buffer VF16 {float16_t data_vf16[];};
+// Native 16-bit aliases for the sparse gather callbacks.
+layout (binding = 1) readonly buffer K16 {FLOAT_TYPE data_k16[];};
+layout (binding = 2) readonly buffer V16 {FLOAT_TYPE data_v16[];};
layout (binding = 3) readonly buffer MF16 {float16_t data_mf16[];};
#ifdef GL_NV_cooperative_matrix_decode_vector
-layout (binding = 1) readonly buffer KF16V4 {f16vec4 data_kf16v4[];};
-layout (binding = 2) readonly buffer VF16V4 {f16vec4 data_vf16v4[];};
+layout (binding = 1) readonly buffer K16V4 {FLOAT_TYPEV4 data_k16v4[];};
+layout (binding = 2) readonly buffer V16V4 {FLOAT_TYPEV4 data_v16v4[];};
#endif
-// K/V/mask f16-element offsets for the current head/batch, set in main().
+// K/V/mask offsets in 16-bit elements for the current head/batch, set in main().
uint32_t g_k_off_elem, g_v_off_elem, g_m_off_elem;
-#if !defined(BFLOAT16)
-// blockCoords are in block units: KV slot = blockCoords[0],
-// head dim = blockCoords[1]*FA_GATHER_BS + coordInBlock[1].
-float16_t faGatherK(const decodeBufFA_K unused, const uint32_t blockCoords[2], const uint32_t coordInBlock[2]) {
- if (blockCoords[0] >= p.split_kv) { return float16_t(0); }
+FLOAT_TYPE faGatherK(const decodeBufFA_K bl_in, const uint32_t blockCoords[2], const uint32_t coordInBlock[2]) {
+ if (blockCoords[0] >= p.split_kv) { return FLOAT_TYPE(0.0); }
const int r = data_sparse[sparse_base + blockCoords[0]];
- return r < 0 ? float16_t(0) : data_kf16[g_k_off_elem + uint(r) * k_stride + blockCoords[1] * FA_GATHER_BS + coordInBlock[1]];
+ if (r < 0) { return FLOAT_TYPE(0.0); }
+#if !defined(BFLOAT16)
+ if (USE_DECODE_K) {
+ decodeBufFA_K block = decodeBufFA_K(decodeBufFA_Byte(bl_in) + uint64_t(uint(r) - blockCoords[0]) * k_stride * FaBlockBytesK);
+ return faDecodeK(block, blockCoords, coordInBlock);
+ }
+#endif
+ return data_k16[g_k_off_elem + uint(r) * k_stride + blockCoords[1] * FA_GATHER_BS + coordInBlock[1]];
}
-float16_t faGatherV(const decodeBufFA_V unused, const uint32_t blockCoords[2], const uint32_t coordInBlock[2]) {
- if (blockCoords[0] >= p.split_kv) { return float16_t(0); }
+FLOAT_TYPE faGatherV(const decodeBufFA_V bl_in, const uint32_t blockCoords[2], const uint32_t coordInBlock[2]) {
+ if (blockCoords[0] >= p.split_kv) { return FLOAT_TYPE(0.0); }
const int r = data_sparse[sparse_base + blockCoords[0]];
- return r < 0 ? float16_t(0) : data_vf16[g_v_off_elem + uint(r) * v_stride + blockCoords[1] * FA_GATHER_BS + coordInBlock[1]];
+ if (r < 0) { return FLOAT_TYPE(0.0); }
+#if !defined(BFLOAT16)
+ if (USE_DECODE_V) {
+ decodeBufFA_V block = decodeBufFA_V(decodeBufFA_Byte(bl_in) + uint64_t(uint(r) - blockCoords[0]) * v_stride * FaBlockBytesV);
+ return faDecodeV(block, blockCoords, coordInBlock);
+ }
+#endif
+ return data_v16[g_v_off_elem + uint(r) * v_stride + blockCoords[1] * FA_GATHER_BS + coordInBlock[1]];
}
#ifdef GL_NV_cooperative_matrix_decode_vector
-f16vec4 faGatherKVector(const decodeBufFA_K unused, const uint32_t blockCoords[2], const uint32_t coordInBlock[2]) {
- if (blockCoords[0] >= p.split_kv) { return f16vec4(0); }
+FLOAT_TYPEV4 faGatherKVector(const decodeBufFA_K bl_in, const uint32_t blockCoords[2], const uint32_t coordInBlock[2]) {
+ if (blockCoords[0] >= p.split_kv) { return FLOAT_TYPEV4(0.0); }
const int r = data_sparse[sparse_base + blockCoords[0]];
- if (r < 0) { return f16vec4(0); }
+ if (r < 0) { return FLOAT_TYPEV4(0.0); }
+#if !defined(BFLOAT16)
+ if (USE_DECODE_K) {
+ decodeBufFA_K block = decodeBufFA_K(decodeBufFA_Byte(bl_in) + uint64_t(uint(r) - blockCoords[0]) * k_stride * FaBlockBytesK);
+ return faDecodeKVector(block, blockCoords, coordInBlock);
+ }
+#endif
const uint32_t o = g_k_off_elem + uint(r) * k_stride + blockCoords[1] * FA_GATHER_BS + coordInBlock[1];
- return data_kf16v4[o / 4];
+ return data_k16v4[o / 4];
}
-f16vec4 faGatherVVector(const decodeBufFA_V unused, const uint32_t blockCoords[2], const uint32_t coordInBlock[2]) {
- if (blockCoords[0] >= p.split_kv) { return f16vec4(0); }
+FLOAT_TYPEV4 faGatherVVector(const decodeBufFA_V bl_in, const uint32_t blockCoords[2], const uint32_t coordInBlock[2]) {
+ if (blockCoords[0] >= p.split_kv) { return FLOAT_TYPEV4(0.0); }
const int r = data_sparse[sparse_base + blockCoords[0]];
- if (r < 0) { return f16vec4(0); }
+ if (r < 0) { return FLOAT_TYPEV4(0.0); }
+#if !defined(BFLOAT16)
+ if (USE_DECODE_V) {
+ decodeBufFA_V block = decodeBufFA_V(decodeBufFA_Byte(bl_in) + uint64_t(uint(r) - blockCoords[0]) * v_stride * FaBlockBytesV);
+ return faDecodeVVector(block, blockCoords, coordInBlock);
+ }
+#endif
const uint32_t o = g_v_off_elem + uint(r) * v_stride + blockCoords[1] * FA_GATHER_BS + coordInBlock[1];
- return data_vf16v4[o / 4];
+ return data_v16v4[o / 4];
}
#define FAGATHERK , faGatherK, faGatherKVector
@@ -163,7 +191,6 @@ f16vec4 faGatherVVector(const decodeBufFA_V unused, const uint32_t blockCoords[2
#define FAGATHERK , faGatherK
#define FAGATHERV , faGatherV
#endif
-#endif
// Add gathered mask to S (slope==1 since sparse requires max_bias==0). col = slot in block jblk.
ACC_TYPE faAddSparseMask(const uint32_t row, const uint32_t col, const ACC_TYPE elem, const uint32_t jblk) {
@@ -252,8 +279,8 @@ void main() {
tensorViewNV<2, false, 1, 0> tensorViewTranspose = createTensorViewNV(2, false, 1, 0);
- const uint bs_k = USE_SPARSE ? FA_GATHER_BS : fa_block_elems(FaTypeK);
- const uint bs_v = USE_SPARSE ? FA_GATHER_BS : fa_block_elems(FaTypeV);
+ const uint bs_k = USE_SPARSE ? max(FA_GATHER_BS, BLOCK_SIZE_K) : BLOCK_SIZE_K;
+ const uint bs_v = USE_SPARSE ? max(FA_GATHER_BS, BLOCK_SIZE_V) : BLOCK_SIZE_V;
tensorLayoutK = setTensorLayoutBlockSizeNV(tensorLayoutK, 1, bs_k);
tensorLayoutV = setTensorLayoutBlockSizeNV(tensorLayoutV, 1, bs_v);
@@ -384,18 +411,15 @@ void main() {
uint32_t k_offset = ik2*p.nb12 + ik3*p.nb13;
// F16: bs_k==1 (direct load). F32: bs_k==4 (vec4 / dequantFuncF32). Quantized types: bs_k==32.
-#if defined(BFLOAT16)
- coopMatLoadTensorNV(K_T, data_k, k_offset, sliceTensorLayoutNV(tensorLayoutK, j * Bc, Bc, 0, HSK_pad), tensorViewTranspose);
-#else
- const bool k_use_decode = (bs_k > 1u);
if (USE_SPARSE) {
coopMatLoadTensorNV(K_T, data_k, k_offset, sliceTensorLayoutNV(tensorLayoutK, j * Bc, Bc, 0, HSK_pad), tensorViewTranspose FAGATHERK);
- } else if (k_use_decode) {
+#if !defined(BFLOAT16)
+ } else if (USE_DECODE_K) {
coopMatLoadTensorNV(K_T, data_k, k_offset, sliceTensorLayoutNV(tensorLayoutK, j * Bc, Bc, 0, HSK_pad), tensorViewTranspose FADECODEK);
+#endif
} else {
coopMatLoadTensorNV(K_T, data_k, k_offset, sliceTensorLayoutNV(tensorLayoutK, j * Bc, Bc, 0, HSK_pad), tensorViewTranspose);
}
-#endif
S = coopMatMulAdd(Qf16, K_T, S);
if (LOGIT_SOFTCAP) {
@@ -458,18 +482,15 @@ void main() {
coopmat<FLOAT_TYPE, gl_ScopeWorkgroup, Bc, HSV_pad, gl_MatrixUseB> V;
uint32_t v_offset = iv2*p.nb22 + iv3*p.nb23;
-#if defined(BFLOAT16)
- coopMatLoadTensorNV(V, data_v, v_offset, sliceTensorLayoutNV(tensorLayoutV, j * Bc, Bc, 0, HSV_pad));
-#else
- const bool v_use_decode = (bs_v > 1u);
if (USE_SPARSE) {
coopMatLoadTensorNV(V, data_v, v_offset, sliceTensorLayoutNV(tensorLayoutV, j * Bc, Bc, 0, HSV_pad) FAGATHERV);
- } else if (v_use_decode) {
+#if !defined(BFLOAT16)
+ } else if (USE_DECODE_V) {
coopMatLoadTensorNV(V, data_v, v_offset, sliceTensorLayoutNV(tensorLayoutV, j * Bc, Bc, 0, HSV_pad) FADECODEV);
+#endif
} else {
coopMatLoadTensorNV(V, data_v, v_offset, sliceTensorLayoutNV(tensorLayoutV, j * Bc, Bc, 0, HSV_pad));
}
-#endif
L = eM*L + rowsum;
diff --git a/tests/test-backend-ops.cpp b/tests/test-backend-ops.cpp
index 453fb78d7..709997680 100644
--- a/tests/test-backend-ops.cpp
+++ b/tests/test-backend-ops.cpp
@@ -11526,6 +11526,15 @@ static std::vector<std::unique_ptr<test_case>> make_test_cases_eval() {
// KV not a multiple of the compaction workgroup size.
test_cases.emplace_back(new test_flash_attn_ext(128, 128, 1, { 8, 1}, 5003, 4, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16, {0, 1, 2, 3}, true, false, 512));
+ // Sparse gather: native block sizes, padded slots, and head/batch strides.
+ for (ggml_type type : {GGML_TYPE_F32, GGML_TYPE_F16, GGML_TYPE_BF16, GGML_TYPE_Q4_0, GGML_TYPE_Q4_1, GGML_TYPE_Q5_0, GGML_TYPE_Q5_1, GGML_TYPE_Q8_0, GGML_TYPE_IQ4_NL}) {
+ test_cases.emplace_back(new test_flash_attn_ext(128, 96, 2, {8, 2}, 5003, 3, true, false, 0, 0, GGML_PREC_F32, type, type, {0, 1, 2, 3}, true, false, 257));
+ }
+ test_cases.emplace_back(new test_flash_attn_ext(128, 96, 2, {8, 2}, 5003, 3, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_Q8_0, GGML_TYPE_F16, {0, 2, 1, 3}, true, false, 257));
+ test_cases.emplace_back(new test_flash_attn_ext(128, 96, 2, {8, 2}, 5003, 3, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_Q4_0, {0, 2, 1, 3}, true, false, 257));
+ test_cases.emplace_back(new test_flash_attn_ext(128, 96, 2, {8, 2}, 5003, 3, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_Q8_0, GGML_TYPE_F32, {0, 1, 2, 3}, true, false, 257));
+ test_cases.emplace_back(new test_flash_attn_ext(128, 96, 2, {8, 2}, 5003, 3, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_F32, GGML_TYPE_Q8_0, {0, 1, 2, 3}, true, false, 257));
+
// more V-is-sub-view-of-K cases: other head shapes, and full views with equal head sizes
test_cases.emplace_back(new test_flash_attn_ext(320, 256, 1, {32, 1}, 512, 1, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16, {0, 1, 2, 3}, true, true));
test_cases.emplace_back(new test_flash_attn_ext(192, 128, 4, {8, 1}, 512, 1, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16, {0, 1, 2, 3}, true, true));
@@ -12017,6 +12026,8 @@ static std::vector<std::unique_ptr<test_case>> make_test_cases_perf() {
test_cases.emplace_back(new test_flash_attn_ext(576, 512, 1, {16, 1}, kv, 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(256, 256, 2, {12, 1}, kv, 1, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16, {0, 1, 2, 3}, true, false, 2048));
test_cases.emplace_back(new test_flash_attn_ext(256, 256, 2, {12, 1}, kv, 1, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_Q8_0, GGML_TYPE_Q8_0, {0, 1, 2, 3}, true, false, 2048));
+ test_cases.emplace_back(new test_flash_attn_ext(256, 256, 2, {12, 1}, kv, 1, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_Q4_0, GGML_TYPE_Q4_0, {0, 1, 2, 3}, true, false, 2048));
+ test_cases.emplace_back(new test_flash_attn_ext(256, 256, 2, {12, 1}, kv, 1, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_BF16, GGML_TYPE_BF16, {0, 1, 2, 3}, true, false, 2048));
test_cases.emplace_back(new test_flash_attn_ext(256, 256, 2, {12, 1}, kv, 1, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_Q8_0, GGML_TYPE_Q8_0, {0, 1, 2, 3}, true, false, 0));
}