Commit fc82583e6 for llama.cpp
commit fc82583e65ad753710fbd69a9244d9a35dca667a
Author: Ruben Ortlam <rortlam@redhat.com>
Date: Tue Sep 15 11:30:27 2026 +0200
vulkan: support sparse Flash Attention (#28105)
* vulkan: add sparse Flash Attention support for DSV4/GLM
* tune implementation
* add tests
* avoid nondeterministic atomicAdd
* add cm2 decode vector support
* simplify logic and make variable names more consistent
* add cm2 f16vec4 binding for decode vector
diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp
index 0dfa44dbf..f936127a6 100644
--- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp
+++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp
@@ -1171,6 +1171,10 @@ struct vk_device_struct {
std::map<std::pair<uint32_t, uint32_t>, vk_pipeline> pipeline_fa_mask_opt;
+ vk_pipeline pipeline_fa_sparse_compact;
+ vk_pipeline pipeline_fa_sparse_compact_subgroup;
+ bool fa_sparse_compact_use_subgroups;
+
vk_pipeline pipeline_flash_attn_split_k_reduce;
vk_pipeline pipeline_count_experts;
@@ -2196,6 +2200,16 @@ struct vk_op_flash_attn_mask_opt_push_constants {
uint32_t nbd3;
};
+struct vk_op_flash_attn_sparse_compact_push_constants {
+ uint32_t KV;
+ uint32_t nem1;
+ uint32_t nem2;
+ uint32_t nbm1;
+ uint32_t nbm2;
+ uint32_t nbm3;
+ uint32_t n_kv_max;
+};
+
// Allow pre-recording command buffers
struct vk_staging_memcpy {
vk_staging_memcpy(void * _dst, const void * _src, size_t _n) : dst(_dst), src(_src), n(_n) {}
@@ -4119,14 +4133,15 @@ static vk_fa_tuning_params get_fa_tuning_params(const vk_device& device, uint32_
}
static vk_fa_pipeline_state get_fa_pipeline_state(const vk_device& device, const vk_fa_tuning_params& params, uint32_t hsk, uint32_t hsv, bool aligned, bool f32acc,
- bool use_mask, bool use_mask_opt, bool use_logit_softcap, ggml_type k_type, ggml_type v_type) {
+ bool use_mask, bool use_mask_opt, bool use_logit_softcap, bool use_sparse, ggml_type k_type, ggml_type v_type) {
const bool old_amd_windows = device->vendor_id == VK_VENDOR_ID_AMD && device->driver_id == vk::DriverId::eAmdProprietary &&
(device->architecture == AMD_GCN || device->architecture == AMD_RDNA1 || device->architecture == AMD_RDNA2);
uint32_t flags = (use_mask_opt ? 1 : 0) |
(use_mask ? 2 : 0) |
(use_logit_softcap ? 4 : 0) |
- (old_amd_windows ? 8 : 0);
+ (old_amd_windows ? 8 : 0) |
+ (use_sparse ? 16 : 0);
const uint32_t subgroup_size = params.disable_subgroups ? 0 : params.subgroup_size;
@@ -4746,7 +4761,7 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) {
}
name = aligned ? "flash_attn_f32_f16_aligned" : "flash_attn_f32_f16";
}
- ggml_vk_create_pipeline(device, fa.second, name, spv_size, spv_data, "main", 7,
+ ggml_vk_create_pipeline(device, fa.second, name, spv_size, spv_data, "main", 8,
sizeof(vk_flash_attn_push_constants), {Br, 1, 1},
get_fa_spec_constants(fa.first), aligned ? Bc : 1, true,
!fa_ds, !fa_ds ? fa_sgs : 0);
@@ -4782,7 +4797,7 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) {
else { spv_data = flash_attn_f32_f16_f16acc_cm1_data; spv_size = flash_attn_f32_f16_f16acc_cm1_len; }
name = aligned ? "flash_attn_f32_f16_aligned_cm1" : "flash_attn_f32_f16_cm1";
}
- ggml_vk_create_pipeline(device, fa.second, name, spv_size, spv_data, "main", 7,
+ ggml_vk_create_pipeline(device, fa.second, name, spv_size, spv_data, "main", 8,
sizeof(vk_flash_attn_push_constants), {Br, 1, 1},
get_fa_spec_constants(fa.first), aligned ? Bc : 1, true,
!fa_ds, !fa_ds ? fa_sgs : 0);
@@ -4819,7 +4834,7 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) {
if (f32acc) { spv_data = flash_attn_f32_f16_cm2_data; spv_size = flash_attn_f32_f16_cm2_len; name = "flash_attn_f32_f16_f32acc_cm2"; }
else { spv_data = flash_attn_f32_f16_f16acc_cm2_data; spv_size = flash_attn_f32_f16_f16acc_cm2_len; name = "flash_attn_f32_f16_f16acc_cm2"; }
}
- ggml_vk_create_pipeline(device, fa.second, name, spv_size, spv_data, "main", 7,
+ ggml_vk_create_pipeline(device, fa.second, name, spv_size, spv_data, "main", 8,
sizeof(vk_flash_attn_push_constants), {Br, 1, 1},
get_fa_spec_constants(fa.first), aligned ? Bc : 1, true, false, 0);
}
@@ -5783,6 +5798,22 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) {
ggml_vk_create_pipeline(device, it.second, "fa_mask_opt", fa_mask_opt_len, fa_mask_opt_data, "main", 2, sizeof(vk_op_flash_attn_mask_opt_push_constants), {1, 1, 1}, {128, 128 / device->subgroup_size, BrBc.first, BrBc.second}, 1, true, true, device->subgroup_size);
}
+ {
+ // Large workgroup so the per-row KV scan parallelizes; capped to device limits.
+ const uint32_t compact_max = std::min({1024u, device->properties.limits.maxComputeWorkGroupInvocations, device->properties.limits.maxComputeWorkGroupSize[0]});
+
+ // Fast ballot prefix-sum path when the device supports full subgroups; otherwise
+ // a shared-memory prefix-sum fallback. Both emit a deterministic ascending list.
+ device->fa_sparse_compact_use_subgroups = device->subgroup_ballot && device->subgroup_require_full_support;
+ if (device->fa_sparse_compact_use_subgroups) {
+ const uint32_t compact_wg = std::max(device->subgroup_size, (compact_max / device->subgroup_size) * device->subgroup_size);
+ const uint32_t compact_num_sg = compact_wg / device->subgroup_size;
+ ggml_vk_create_pipeline(device, device->pipeline_fa_sparse_compact_subgroup, "fa_sparse_compact_subgroup", fa_sparse_compact_subgroup_len, fa_sparse_compact_subgroup_data, "main", 2, sizeof(vk_op_flash_attn_sparse_compact_push_constants), {1, 1, 1}, {compact_wg, compact_num_sg}, 1, true, true, device->subgroup_size);
+ } else {
+ ggml_vk_create_pipeline(device, device->pipeline_fa_sparse_compact, "fa_sparse_compact", fa_sparse_compact_len, fa_sparse_compact_data, "main", 2, sizeof(vk_op_flash_attn_sparse_compact_push_constants), {1, 1, 1}, {compact_max}, 1, true);
+ }
+ }
+
if (device->subgroup_clustered && device->subgroup_require_full_support) {
ggml_vk_create_pipeline(device, device->pipeline_quantize_q8_1_x4, "quantize_q8_1_x4", quantize_q8_1_x4_subgroup_len, quantize_q8_1_x4_subgroup_data, "main", 2, sizeof(vk_quantize_q8_1_push_constants), {32 * device->subgroup_size / 8, 1, 1}, { device->subgroup_size }, 1, true, true);
} else {
@@ -11276,6 +11307,30 @@ static void ggml_vk_flash_attn(ggml_backend_vk_context * ctx, vk_context& subctx
tuning_params = get_fa_tuning_params(ctx->device, HSK, HSV, N, KV, k_type_eff, v_type_eff, f32acc);
+ float scale = 1.0f;
+ float max_bias = 0.0f;
+ float logit_softcap = 0.0f;
+
+ memcpy(&scale, (const float *) dst->op_params + 0, sizeof(float));
+ memcpy(&max_bias, (const float *) dst->op_params + 1, sizeof(float));
+ memcpy(&logit_softcap, (const float *) dst->op_params + 2, sizeof(float));
+
+ if (logit_softcap != 0) {
+ scale /= logit_softcap;
+ }
+
+ // Sparse mask hint (op_params[4]): compact the <= n_kv_max finite positions and gather only those.
+ const int32_t n_kv_max = mask ? ggml_get_op_params_i32(dst, 4) : 0;
+ static const bool disable_sparse = getenv("GGML_VK_FA_SPARSE_DISABLE") != nullptr;
+ // cm2 dense is fast, so it needs a larger reduction to win.
+ const int64_t min_ratio = tuning_params.path == FA_COOPMAT2 ? 4 : 2;
+ const bool use_sparse = !disable_sparse && n_kv_max > 0 && mask &&
+ max_bias == 0.0f && logit_softcap == 0.0f &&
+ k_type_eff == GGML_TYPE_F16 && v_type_eff == GGML_TYPE_F16 &&
+ 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));
+
const uint32_t q_stride = (uint32_t)(nbq1 / ggml_type_size(q->type));
uint32_t k_stride = (uint32_t)(nbk1 / ggml_type_size(k->type));
uint32_t v_stride = (uint32_t)(nbv1 / ggml_type_size(v->type));
@@ -11298,7 +11353,6 @@ static void ggml_vk_flash_attn(ggml_backend_vk_context * ctx, vk_context& subctx
nbv2_eff = (uint32_t)((uint64_t)HSV * KV * sizeof(ggml_fp16_t));
nbv3_eff = (uint32_t)((uint64_t)HSV * KV * nev2 * sizeof(ggml_fp16_t));
}
-
const uint32_t alignment = tuning_params.block_cols;
bool aligned = (KV % alignment) == 0 &&
// the "aligned" shader variant will forcibly align strides, for performance
@@ -11309,23 +11363,11 @@ static void ggml_vk_flash_attn(ggml_backend_vk_context * ctx, vk_context& subctx
aligned = false;
}
- float scale = 1.0f;
- float max_bias = 0.0f;
- float logit_softcap = 0.0f;
-
- memcpy(&scale, (const float *) dst->op_params + 0, sizeof(float));
- memcpy(&max_bias, (const float *) dst->op_params + 1, sizeof(float));
- memcpy(&logit_softcap, (const float *) dst->op_params + 2, sizeof(float));
-
- if (logit_softcap != 0) {
- scale /= logit_softcap;
- }
-
// Only use mask opt when the mask is fairly large. This hasn't been tuned extensively.
- bool use_mask_opt = mask && nem1 >= 32 && nem0 * nem1 > 32768 && nem0 >= tuning_params.block_cols * 16
+ bool use_mask_opt = mask && !use_sparse && nem1 >= 32 && nem0 * nem1 > 32768 && nem0 >= tuning_params.block_cols * 16
&& (ctx->device->architecture != vk_device_architecture::AMD_GCN || HSK > 256 || HSV > 256);
vk_fa_pipeline_state fa_pipeline_state = get_fa_pipeline_state(ctx->device, tuning_params, HSK, HSV, aligned, f32acc,
- mask != nullptr, use_mask_opt, logit_softcap != 0, k_type_eff, v_type_eff);
+ mask != nullptr, use_mask_opt, logit_softcap != 0, use_sparse, k_type_eff, v_type_eff);
vk_pipeline pipeline = nullptr;
@@ -11360,7 +11402,19 @@ static void ggml_vk_flash_attn(ggml_backend_vk_context * ctx, vk_context& subctx
const uint32_t Tr = CEIL_DIV(N, Br);
// Try to use split_k when KV is large enough to be worth the overhead.
- if (gqa_ratio > 1 && workgroups_x <= Br) {
+ // Sparse: split_kv carries n_kv_max, split_k partitions its blocks for occupancy.
+ if (use_sparse) {
+ split_kv = (uint32_t)n_kv_max;
+ const uint32_t total_blocks = CEIL_DIV((uint32_t)n_kv_max, Bc);
+ const uint32_t base_wgs = (gqa_ratio > 1 ? workgroups_x : Tr) * workgroups_y * workgroups_z;
+ if (base_wgs < shader_core_count * 2) {
+ split_k = shader_core_count * 2 / base_wgs;
+ }
+ split_k = std::max(1u, std::min(split_k, total_blocks));
+ // Match the shader's per-split block count so no split is empty.
+ const uint32_t per_blocks = CEIL_DIV(total_blocks, split_k);
+ split_k = CEIL_DIV(total_blocks, per_blocks);
+ } else if (gqa_ratio > 1 && workgroups_x <= Br) {
split_k = shader_core_count * 2 / (workgroups_x * workgroups_y * workgroups_z);
} else if (gqa_ratio <= 1) {
uint32_t total_wgs_no_split = Tr * workgroups_y * workgroups_z;
@@ -11369,7 +11423,7 @@ static void ggml_vk_flash_attn(ggml_backend_vk_context * ctx, vk_context& subctx
}
}
- if (split_k > 1) {
+ if (!use_sparse && split_k > 1) {
// Try to evenly split KV into split_k chunks, but it needs to be a multiple
// of "align", so recompute split_k based on that.
split_kv = ROUNDUP_POW2(std::max(1u, KV / split_k), alignment);
@@ -11416,6 +11470,24 @@ static void ggml_vk_flash_attn(ggml_backend_vk_context * ctx, vk_context& subctx
}
}
+ // Sparse index scratch reuses prealloc_y (mutually exclusive with mask opt).
+ const uint64_t sparse_idx_size = use_sparse
+ ? sizeof(int32_t) * (uint64_t)n_kv_max * nem1 * nem2 * nem3
+ : 0;
+ vk_pipeline sparse_compact_pipeline = ctx->device->fa_sparse_compact_use_subgroups
+ ? ctx->device->pipeline_fa_sparse_compact_subgroup
+ : ctx->device->pipeline_fa_sparse_compact;
+ if (use_sparse) {
+ ggml_pipeline_request_descriptor_sets(ctx, sparse_compact_pipeline, 1);
+ if (ctx->prealloc_size_y < sparse_idx_size) {
+ ctx->prealloc_size_y = sparse_idx_size;
+ ggml_vk_preallocate_buffers(ctx, subctx);
+ }
+ if (ctx->prealloc_y_need_sync) {
+ ggml_vk_sync_buffers(ctx, subctx);
+ }
+ }
+
const uint32_t n_head_kv = neq2;
const uint32_t n_head_log2 = 1u << (uint32_t) floorf(log2f((float) n_head_kv));
const float m0 = powf(2.0f, -(max_bias ) / n_head_log2);
@@ -11428,6 +11500,7 @@ static void ggml_vk_flash_attn(ggml_backend_vk_context * ctx, vk_context& subctx
vk_subbuffer mask_buf = mask ? ggml_vk_tensor_subbuffer(ctx, mask) : q_buf;
vk_subbuffer sinks_buf = sinks ? ggml_vk_tensor_subbuffer(ctx, sinks) : q_buf;
vk_subbuffer mask_opt_buf = use_mask_opt ? ggml_vk_subbuffer(ctx, ctx->prealloc_y, 0) : q_buf;
+ vk_subbuffer sparse_buf = use_sparse ? ggml_vk_subbuffer(ctx, ctx->prealloc_y, 0) : q_buf;
if (use_dequant_kv) {
const uint64_t fp = sizeof(ggml_fp16_t);
@@ -11479,6 +11552,24 @@ static void ggml_vk_flash_attn(ggml_backend_vk_context * ctx, vk_context& subctx
ggml_vk_sync_buffers(ctx, subctx);
}
+ if (use_sparse)
+ {
+ const vk_op_flash_attn_sparse_compact_push_constants sc_pc = {
+ KV,
+ nem1,
+ nem2,
+ (uint32_t)(mask->nb[1] / sizeof(ggml_fp16_t)),
+ (uint32_t)(mask->nb[2] / sizeof(ggml_fp16_t)),
+ (uint32_t)(mask->nb[3] / sizeof(ggml_fp16_t)),
+ (uint32_t)n_kv_max,
+ };
+
+ ggml_vk_dispatch_pipeline(ctx, subctx, sparse_compact_pipeline,
+ { mask_buf, sparse_buf }, sc_pc,
+ { nem1, nem2, nem3 });
+ ggml_vk_sync_buffers(ctx, subctx);
+ }
+
const vk_flash_attn_push_constants pc = { N, KV,
(uint32_t)ne1, (uint32_t)ne2, (uint32_t)ne3,
(uint32_t)neq2, (uint32_t)neq3,
@@ -11511,7 +11602,7 @@ static void ggml_vk_flash_attn(ggml_backend_vk_context * ctx, vk_context& subctx
vk_subbuffer split_k_buf = ggml_vk_subbuffer(ctx, ctx->prealloc_split_k, 0);
ggml_vk_dispatch_pipeline(ctx, subctx, pipeline,
- {q_buf, k_buf, v_buf, mask_buf, sinks_buf, split_k_buf, mask_opt_buf},
+ {q_buf, k_buf, v_buf, mask_buf, sinks_buf, split_k_buf, mask_opt_buf, sparse_buf},
pc, { dispatch_x, workgroups_y, workgroups_z });
ggml_vk_sync_buffers(ctx, subctx);
@@ -11526,13 +11617,16 @@ static void ggml_vk_flash_attn(ggml_backend_vk_context * ctx, vk_context& subctx
workgroups_x *= pipeline->wg_denoms[0];
}
ggml_vk_dispatch_pipeline(ctx, subctx, pipeline,
- {q_buf, k_buf, v_buf, mask_buf, sinks_buf, dst_buf, mask_opt_buf},
+ {q_buf, k_buf, v_buf, mask_buf, sinks_buf, dst_buf, mask_opt_buf, sparse_buf},
pc, { workgroups_x, workgroups_y, workgroups_z });
}
if (use_dequant_kv) {
ctx->prealloc_x_need_sync = true;
}
+ if (use_mask_opt || use_sparse) {
+ ctx->prealloc_y_need_sync = true;
+ }
}
static vk_conv_shapes ggml_vk_conv_select_shape(ggml_backend_vk_context * ctx, uint32_t K, uint32_t NPQ) {
diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn.comp b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn.comp
index 9a12cdfb8..107d44aaa 100644
--- a/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn.comp
+++ b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn.comp
@@ -218,12 +218,14 @@ void main() {
uint32_t c = (idx + tid) % Bc;
uint32_t r = (idx + tid) / Bc;
if (idx + tid < Bc * Br) {
- if ((!KV_bounds_check || j * Bc + c < KV) && (!nem1_bounds_check || i * Br + r < p.nem1)) {
- FLOAT_TYPE m = FLOAT_TYPE(data_m[m_offset + (i * Br + r) * m_stride + (j * Bc + c)]);
+ uint32_t kcol;
+ bool kv_active = fa_kv_index(j * Bc + c, kcol);
+ if (kv_active && (!nem1_bounds_check || i * Br + r < p.nem1)) {
+ FLOAT_TYPE m = FLOAT_TYPE(data_m[m_offset + (i * Br + r) * m_stride + kcol]);
masksh[c * masksh_stride + r] = m;
max_mask = max(max_mask, float(m));
} else {
- masksh[c * masksh_stride + r] = FLOAT_TYPE(0);
+ masksh[c * masksh_stride + r] = USE_SPARSE ? FLOAT_TYPE(NEG_FLT_MAX_OVER_2) : FLOAT_TYPE(0);
}
}
}
@@ -258,14 +260,15 @@ void main() {
uint32_t c = (idx + tid) / (HSK / 4);
if (idx + gl_WorkGroupSize.x <= Bc * HSK / 4 || c < Bc) {
FLOAT_TYPEV4 K_Tf = FLOAT_TYPEV4(0);
- if (!KV_bounds_check || j * Bc + c < KV) {
+ uint32_t kcol;
+ if (fa_kv_index(j * Bc + c, kcol)) {
if (USE_DECODE_K) {
- uint coord = (j * Bc + c) * k_stride * BLOCK_SIZE_K + 4 * d;
+ uint coord = kcol * k_stride * BLOCK_SIZE_K + 4 * d;
uint ib = coord / BLOCK_SIZE_K;
uint iqs = (coord % BLOCK_SIZE_K);
K_Tf = dequantize4(ib, iqs, k_offset, BINDING_IDX_K);
} else {
- K_Tf = FLOAT_TYPEV4(data_kv4[k_offset / 4 + (j * Bc + c) * k_stride / 4 + d]);
+ K_Tf = FLOAT_TYPEV4(data_kv4[k_offset / 4 + kcol * k_stride / 4 + d]);
}
}
@@ -305,7 +308,9 @@ void main() {
}
[[unroll]] for (uint32_t c = 0; c < cols_per_thread; ++c) {
- if (KV_bounds_check && j * Bc + c * cols_per_iter + col_tid >= KV) {
+ uint32_t kcol;
+ bool kv_active = fa_kv_index(j * Bc + c * cols_per_iter + col_tid, kcol);
+ if (!kv_active) {
continue;
}
@@ -313,12 +318,12 @@ void main() {
if (SHMEM_STAGING != 0) {
K_Tf = kvsh[(c * cols_per_iter + col_tid) * kvsh_stride + (d * D_split + d_tid)];
} else if (USE_DECODE_K) {
- uint coord = (j * Bc + c * cols_per_iter + col_tid) * k_stride * BLOCK_SIZE_K + 4 * (d * D_split + d_tid);
+ uint coord = kcol * k_stride * BLOCK_SIZE_K + 4 * (d * D_split + d_tid);
uint ib = coord / BLOCK_SIZE_K;
uint iqs = (coord % BLOCK_SIZE_K);
K_Tf = dequantize4(ib, iqs, k_offset, BINDING_IDX_K);
} else {
- K_Tf = FLOAT_TYPEV4(data_kv4[k_offset / 4 + (j * Bc + c * cols_per_iter + col_tid) * k_stride / 4 + d * D_split + d_tid]);
+ K_Tf = FLOAT_TYPEV4(data_kv4[k_offset / 4 + kcol * k_stride / 4 + d * D_split + d_tid]);
}
[[unroll]] for (uint32_t r = 0; r < rows_per_thread; ++r) {
Sf[r][c] = dot_product(Q_cache[r], K_Tf, Sf[r][c]);
@@ -327,7 +332,9 @@ void main() {
}
} else {
[[unroll]] for (uint32_t c = 0; c < cols_per_thread; ++c) {
- if (KV_bounds_check && j * Bc + c * cols_per_iter + col_tid >= KV) {
+ uint32_t kcol;
+ bool kv_active = fa_kv_index(j * Bc + c * cols_per_iter + col_tid, kcol);
+ if (!kv_active) {
continue;
}
@@ -336,12 +343,12 @@ void main() {
if (SHMEM_STAGING != 0) {
K_Tf = kvsh[(c * cols_per_iter + col_tid) * kvsh_stride + (d * D_split + d_tid)];
} else if (USE_DECODE_K) {
- uint coord = (j * Bc + c * cols_per_iter + col_tid) * k_stride * BLOCK_SIZE_K + 4 * (d * D_split + d_tid);
+ uint coord = kcol * k_stride * BLOCK_SIZE_K + 4 * (d * D_split + d_tid);
uint ib = coord / BLOCK_SIZE_K;
uint iqs = (coord % BLOCK_SIZE_K);
K_Tf = dequantize4(ib, iqs, k_offset, BINDING_IDX_K);
} else {
- K_Tf = FLOAT_TYPEV4(data_kv4[k_offset / 4 + (j * Bc + c * cols_per_iter + col_tid) * k_stride / 4 + d * D_split + d_tid]);
+ K_Tf = FLOAT_TYPEV4(data_kv4[k_offset / 4 + kcol * k_stride / 4 + d * D_split + d_tid]);
}
[[unroll]] for (uint32_t r = 0; r < rows_per_thread; ++r) {
Sf[r][c] = dot_product(Qf[tile_row(r) * qf_stride + d * D_split + d_tid], K_Tf, Sf[r][c]);
@@ -489,14 +496,15 @@ void main() {
uint32_t c = (idx + tid) / (HSV / 4);
if (idx + gl_WorkGroupSize.x <= Bc * HSV / 4 || c < Bc) {
FLOAT_TYPEV4 V_Tf = FLOAT_TYPEV4(0);
- if (!KV_bounds_check || j * Bc + c < KV) {
+ uint32_t vcol;
+ if (fa_kv_index(j * Bc + c, vcol)) {
if (USE_DECODE_V) {
- uint coord = (j * Bc + c) * v_stride * BLOCK_SIZE_V + 4 * d;
+ uint coord = vcol * v_stride * BLOCK_SIZE_V + 4 * d;
uint ib = coord / BLOCK_SIZE_V;
uint iqs = (coord % BLOCK_SIZE_V);
V_Tf = dequantize4(ib, iqs, v_offset, BINDING_IDX_V);
} else {
- V_Tf = FLOAT_TYPEV4(data_vv4[v_offset / 4 + (j * Bc + c) * v_stride / 4 + d]);
+ V_Tf = FLOAT_TYPEV4(data_vv4[v_offset / 4 + vcol * v_stride / 4 + d]);
}
}
@@ -507,7 +515,9 @@ void main() {
}
[[unroll]] for (uint32_t c = 0; c < cols_per_thread; ++c) {
- if (KV_bounds_check && j * Bc + c * cols_per_iter + col_tid >= KV) {
+ uint32_t vcol;
+ bool kv_active = fa_kv_index(j * Bc + c * cols_per_iter + col_tid, vcol);
+ if (!kv_active) {
continue;
}
@@ -522,12 +532,12 @@ void main() {
if (SHMEM_STAGING != 0) {
Vf = kvsh[(c * cols_per_iter + col_tid) * kvsh_stride + (d * D_split + d_tid)];
} else if (USE_DECODE_V) {
- uint coord = (j * Bc + c * cols_per_iter + col_tid) * v_stride * BLOCK_SIZE_V + 4 * (d * D_split + d_tid);
+ uint coord = vcol * v_stride * BLOCK_SIZE_V + 4 * (d * D_split + d_tid);
uint ib = coord / BLOCK_SIZE_V;
uint iqs = (coord % BLOCK_SIZE_V);
Vf = dequantize4(ib, iqs, v_offset, BINDING_IDX_V);
} else {
- Vf = FLOAT_TYPEV4(data_vv4[v_offset / 4 + (j * Bc + c * cols_per_iter + col_tid) * v_stride / 4 + d * D_split + d_tid]);
+ Vf = FLOAT_TYPEV4(data_vv4[v_offset / 4 + vcol * v_stride / 4 + d * D_split + d_tid]);
}
[[unroll]] for (uint32_t r = 0; r < rows_per_thread; ++r) {
Of[r][d] += FLOAT_TYPEV4(Pf[r] * Vf);
diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_base.glsl b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_base.glsl
index a4be1ebf9..2e0e23bc1 100644
--- a/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_base.glsl
+++ b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_base.glsl
@@ -24,6 +24,8 @@ const bool USE_MASK_OPT = (Flags & 1) != 0;
const bool MASK_ENABLE = (Flags & 2) != 0;
const bool LOGIT_SOFTCAP = (Flags & 4) != 0;
const bool OLD_AMD_WINDOWS = (Flags & 8) != 0;
+// Sparse: gather binding-7 indices instead of scanning [0,KV); p.split_kv = n_kv_max.
+const bool USE_SPARSE = (Flags & 16) != 0;
// Round up head sizes to a multiple of 16, for coopmat1/coopmat2 paths
const uint32_t HSK_pad = (HSK + 15) & ~15;
@@ -82,6 +84,8 @@ layout (binding = 5) writeonly buffer OV4 {D_TYPEV4 data_ov4[];};
layout (binding = 6) readonly buffer MO {uint32_t data_mask_opt[];};
+layout (binding = 7) readonly buffer SP {int32_t data_sparse[];};
+
#define MASK_OPT_ALL_NEG_INF 1
#define MASK_OPT_ALL_ZERO 2
@@ -144,7 +148,7 @@ ACC_TYPE perElemOpGetSink(const in uint32_t r, const in uint32_t c, const in ACC
uint32_t i, N, KV, split_k_index, Tr, start_j, end_j,
gqa_iq1, iq2, iq3, rk2, rk3, rv2, rv3, ik2, ik3, iv2, iv3,
- q_stride, k_stride, v_stride, m_stride;
+ q_stride, k_stride, v_stride, m_stride, sparse_base;
void init_indices()
{
@@ -208,6 +212,33 @@ void init_indices()
// that prevents the compiler from folding the "&" through the select
// and breaking the alignment detection.
m_stride = (p.gqa_ratio > 1) ? (p.gqa_ratio >> 16) : KV;
+
+ // Sparse: the tile shares one mask row (gqa heads, or Br==1). split_k
+ // partitions the n_kv_max blocks.
+ if (USE_SPARSE) {
+ uint32_t qrow = (p.gqa_ratio > 1) ? gqa_iq1 : (i * Br);
+ sparse_base = (((iq3 % p.nem3) * p.nem2 + (iq2 % p.nem2)) * p.nem1 + qrow) * p.split_kv;
+
+ uint32_t total_blocks = CEIL_DIV(p.split_kv, Bc);
+ uint32_t per_blocks = CEIL_DIV(total_blocks, p.k_num);
+ start_j = min(split_k_index * per_blocks, total_blocks);
+ end_j = min((split_k_index + 1) * per_blocks, total_blocks);
+ }
+}
+
+// Resolve a linear KV slot to a real column; false for inactive (sparse padding/-1, or dense OOB).
+bool fa_kv_index(uint lin, out uint kv_col) {
+ if (USE_SPARSE) {
+ if (lin >= p.split_kv) {
+ kv_col = 0;
+ return false;
+ }
+ int idx = data_sparse[sparse_base + lin];
+ kv_col = idx >= 0 ? uint(idx) : 0;
+ return idx >= 0;
+ }
+ kv_col = lin;
+ return !KV_bounds_check || lin < KV;
}
// Bias applied to softmax to stay in fp16 range.
diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_cm1.comp b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_cm1.comp
index 057ed739a..aa9dd624b 100644
--- a/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_cm1.comp
+++ b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_cm1.comp
@@ -176,9 +176,16 @@ void main() {
uint32_t c = (idx + tid) / (Br / 4);
uint32_t r = (idx + tid) % (Br / 4);
if (idx + tid < Bc * Br / 4 || idx + gl_WorkGroupSize.x <= Bc * Br / 4) {
- if ((!KV_bounds_check || j * Bc + c < KV)) {
+ uint32_t kcol;
+ bool kv_active = fa_kv_index(j * Bc + c, kcol);
+ if (kv_active) {
f16vec4 m;
- if (!nem1_bounds_check || i * Br + r * 4 + 3 < p.nem1) {
+ if (USE_SPARSE) {
+ // sparse is gqa-gated (m_stride == 0): all four rows share the value
+ FLOAT_TYPE mv = FLOAT_TYPE(data_m[m_offset + kcol]);
+ m = f16vec4(mv);
+ max_mask = max(max_mask, float(mv));
+ } else if (!nem1_bounds_check || i * Br + r * 4 + 3 < p.nem1) {
m = f16vec4(data_m[m_offset + (i * Br + r * 4 ) * m_stride + (j * Bc + c)],
data_m[m_offset + (i * Br + r * 4 + 1) * m_stride + (j * Bc + c)],
data_m[m_offset + (i * Br + r * 4 + 2) * m_stride + (j * Bc + c)],
@@ -206,6 +213,8 @@ void main() {
m = f16vec4(0.0);
}
mask_cache[idx / WorkGroupSize] = m;
+ } else if (USE_SPARSE) {
+ mask_cache[idx / WorkGroupSize] = f16vec4(NEG_FLT_MAX_OVER_2);
}
}
}
@@ -231,17 +240,19 @@ void main() {
uint32_t c = (idx + tid) / (HSK_pad / 4);
if (idx + gl_WorkGroupSize.x <= Bc * HSK_pad / 4 || c < Bc) {
FLOAT_TYPEV4 K_Tf = FLOAT_TYPEV4(0);
- if ((!KV_bounds_check || j * Bc + c < KV) && (HSK == HSK_pad || d < HSK / 4)) {
+ uint32_t kcol;
+ bool kv_active = fa_kv_index(j * Bc + c, kcol);
+ if (kv_active && (HSK == HSK_pad || d < HSK / 4)) {
#if !defined(BFLOAT16)
if (USE_DECODE_K) {
- uint coord = (j * Bc + c) * k_stride * BLOCK_SIZE_K + 4 * d;
+ uint coord = kcol * k_stride * BLOCK_SIZE_K + 4 * d;
uint ib = coord / BLOCK_SIZE_K;
uint iqs = (coord % BLOCK_SIZE_K);
K_Tf = dequantize4(ib, iqs, k_offset, BINDING_IDX_K);
} else
#endif
{
- K_Tf = FLOAT_TYPEV4(data_kv4[k_offset / 4 + (j * Bc + c) * k_stride / 4 + d]);
+ K_Tf = FLOAT_TYPEV4(data_kv4[k_offset / 4 + kcol * k_stride / 4 + d]);
}
}
@@ -266,7 +277,7 @@ void main() {
if (SHMEM_STAGING == 0) {
// For quants we always need to dequant into kvsh; for f16/bf16 we can load
// directly from global memory when alignment / bounds allow it.
- const bool stage_k = USE_DECODE_K || KV_bounds_check || d * 16 + 16 > HSK;
+ const bool stage_k = USE_DECODE_K || KV_bounds_check || USE_SPARSE || d * 16 + 16 > HSK;
if (stage_k) {
barrier();
[[unroll]] for (uint32_t idx = 0; idx < Bc * MatBr / 4; idx += gl_WorkGroupSize.x) {
@@ -274,17 +285,19 @@ void main() {
uint32_t row = (idx + tid) / (MatBr / 4);
if (idx + tid < Bc * MatBr / 4) {
FLOAT_TYPEV4 K_Tf = FLOAT_TYPEV4(0);
- if ((!KV_bounds_check || j * Bc + row < KV) && (HSK == HSK_pad || d * 16 + col_vec * 4 < HSK)) {
+ uint32_t kcol;
+ bool kv_active = fa_kv_index(j * Bc + row, kcol);
+ if (kv_active && (HSK == HSK_pad || d * 16 + col_vec * 4 < HSK)) {
#if !defined(BFLOAT16)
if (USE_DECODE_K) {
- uint coord = (j * Bc + row) * k_stride * BLOCK_SIZE_K + d * 16 + col_vec * 4;
+ uint coord = kcol * k_stride * BLOCK_SIZE_K + d * 16 + col_vec * 4;
uint ib = coord / BLOCK_SIZE_K;
uint iqs = (coord % BLOCK_SIZE_K);
K_Tf = dequantize4(ib, iqs, k_offset, BINDING_IDX_K);
} else
#endif
{
- K_Tf = FLOAT_TYPEV4(data_kv4[k_offset / 4 + (j * Bc + row) * k_stride / 4 + d * 16 / 4 + col_vec]);
+ K_Tf = FLOAT_TYPEV4(data_kv4[k_offset / 4 + kcol * k_stride / 4 + d * 16 / 4 + col_vec]);
}
}
@@ -401,17 +414,19 @@ void main() {
uint32_t c = (idx + tid) / (HSV_pad / 4);
if (idx + gl_WorkGroupSize.x <= Bc * HSV_pad / 4 || c < Bc) {
FLOAT_TYPEV4 V_Tf = FLOAT_TYPEV4(0);
- if ((!KV_bounds_check || j * Bc + c < KV) && (HSV == HSV_pad || d < HSV / 4)) {
+ uint32_t v_row;
+ bool kv_active = fa_kv_index(j * Bc + c, v_row);
+ if (kv_active && (HSV == HSV_pad || d < HSV / 4)) {
#if !defined(BFLOAT16)
if (USE_DECODE_V) {
- uint coord = (j * Bc + c) * v_stride * BLOCK_SIZE_V + 4 * d;
+ uint coord = v_row * v_stride * BLOCK_SIZE_V + 4 * d;
uint ib = coord / BLOCK_SIZE_V;
uint iqs = (coord % BLOCK_SIZE_V);
V_Tf = dequantize4(ib, iqs, v_offset, BINDING_IDX_V);
} else
#endif
{
- V_Tf = FLOAT_TYPEV4(data_vv4[v_offset / 4 + (j * Bc + c) * v_stride / 4 + d]);
+ V_Tf = FLOAT_TYPEV4(data_vv4[v_offset / 4 + v_row * v_stride / 4 + d]);
}
}
@@ -441,21 +456,22 @@ void main() {
if (SHMEM_STAGING == 0) {
// For quants we always preload via kvsh. For f16/bf16 we only preload when
// alignment / bounds force it (otherwise we coopMatLoad direct from data_vv4).
- const bool stage_v = USE_DECODE_V || KV_bounds_check;
+ const bool stage_v = USE_DECODE_V || KV_bounds_check || USE_SPARSE;
if (stage_v) {
[[unroll]] for (uint32_t i = 0; i < v_loads_per_thread; ++i) {
const uint idx = i * gl_WorkGroupSize.x + tid;
const uint row = idx / v_cols;
const uint col = idx % v_cols;
- const uint v_row = j * Bc + row;
+ uint32_t v_row;
+ bool kv_active = fa_kv_index(j * Bc + row, v_row);
const uint v_col = hsv_tile * MatBc * row_split + col * 4;
const uint coord = v_row * v_stride * BLOCK_SIZE_V + v_col;
const uint ib = coord / BLOCK_SIZE_V;
const uint iqs = coord % BLOCK_SIZE_V;
- if (!KV_bounds_check || (v_row < KV && v_col < HSV)) {
+ if (USE_SPARSE ? (kv_active && v_col < HSV) : (!KV_bounds_check || (v_row < KV && v_col < HSV))) {
#if !defined(BFLOAT16)
if (USE_DECODE_V) {
kvsh[row * vsh_stride + col] = dequantize4(ib, iqs, v_offset, BINDING_IDX_V);
@@ -479,7 +495,7 @@ void main() {
coopMatLoad(KMat, Psh, bc_chunk * MatBc * psh_stride, psh_stride, gl_CooperativeMatrixLayoutColumnMajor);
if (SHMEM_STAGING == 0) {
- if (!USE_DECODE_V && !KV_bounds_check) {
+ if (!USE_DECODE_V && !KV_bounds_check && !USE_SPARSE) {
// F16/BF16 values can be loaded directly from global memory
const uint v_tile_row = j * Bc + bc_chunk * MatBc;
const uint v_tile_offset = v_offset / 4 + v_tile_row * v_stride / 4 + hsv_offset / 4;
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 5a9abe226..c6ed63dd4 100644
--- a/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_cm2.comp
+++ b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_cm2.comp
@@ -29,6 +29,12 @@
#include "dequant_funcs_cm2.glsl"
#endif
+#ifdef GL_NV_cooperative_matrix_decode_vector
+#define FA_GATHER_BS 4u
+#else
+#define FA_GATHER_BS 1u
+#endif
+
// buffer_reference stride = sizeof(struct) = FaBlockBytesK/V.
layout(buffer_reference, std430, buffer_reference_align = 1) buffer decodeBufFA_K {
uint8_t raw[FaBlockBytesK];
@@ -107,6 +113,67 @@ 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[];};
+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[];};
+#endif
+
+// K/V/mask f16-element offsets 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); }
+ 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]];
+}
+
+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); }
+ 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]];
+}
+
+#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); }
+ const int r = data_sparse[sparse_base + blockCoords[0]];
+ if (r < 0) { return f16vec4(0); }
+ const uint32_t o = g_k_off_elem + uint(r) * k_stride + blockCoords[1] * FA_GATHER_BS + coordInBlock[1];
+ return data_kf16v4[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); }
+ const int r = data_sparse[sparse_base + blockCoords[0]];
+ if (r < 0) { return f16vec4(0); }
+ const uint32_t o = g_v_off_elem + uint(r) * v_stride + blockCoords[1] * FA_GATHER_BS + coordInBlock[1];
+ return data_vf16v4[o / 4];
+}
+
+#define FAGATHERK , faGatherK, faGatherKVector
+#define FAGATHERV , faGatherV, faGatherVVector
+#else
+#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) {
+ const float NEG = uintBitsToFloat(0xFEFFFFFF);
+ const uint32_t kvslot = jblk * Bc + col;
+ if (kvslot >= p.split_kv) { return ACC_TYPE(NEG); }
+ const int r = data_sparse[sparse_base + kvslot];
+ return r < 0 ? ACC_TYPE(NEG) : elem + ACC_TYPE(data_mf16[g_m_off_elem + row * m_stride + uint(r)]);
+}
+
ACC_TYPE maxReduce(const in ACC_TYPE x, const in ACC_TYPE y) {
return max(x, y);
}
@@ -185,14 +252,16 @@ void main() {
tensorViewNV<2, false, 1, 0> tensorViewTranspose = createTensorViewNV(2, false, 1, 0);
- const uint bs_k = fa_block_elems(FaTypeK);
- const uint bs_v = fa_block_elems(FaTypeV);
+ 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);
tensorLayoutK = setTensorLayoutBlockSizeNV(tensorLayoutK, 1, bs_k);
tensorLayoutV = setTensorLayoutBlockSizeNV(tensorLayoutV, 1, bs_v);
+ // Sparse iterates n_kv_max (in split_kv); the decode callbacks remap each slot.
+ const uint32_t KV_iter = USE_SPARSE ? p.split_kv : KV;
tensorLayoutQ = setTensorLayoutDimensionNV(tensorLayoutQ, N, HSK);
- tensorLayoutK = setTensorLayoutDimensionNV(tensorLayoutK, KV, HSK);
- tensorLayoutV = setTensorLayoutDimensionNV(tensorLayoutV, KV, HSV);
+ tensorLayoutK = setTensorLayoutDimensionNV(tensorLayoutK, KV_iter, HSK);
+ tensorLayoutV = setTensorLayoutDimensionNV(tensorLayoutV, KV_iter, HSV);
// hint to the compiler that strides are aligned for the aligned variant of the shader
if (Clamp != gl_CooperativeMatrixClampModeConstantNV)
@@ -250,6 +319,10 @@ void main() {
mo_offset += ((iq3 % p.nem3) * p.nem2 + (iq2 % p.nem2)) * CEIL_DIV(p.nem1, Br) * mo_stride;
}
+ g_k_off_elem = (ik2*p.nb12 + ik3*p.nb13) / 2;
+ g_v_off_elem = (iv2*p.nb22 + iv3*p.nb23) / 2;
+ g_m_off_elem = m_offset / 2;
+
uint32_t mask_opt = 0;
uint32_t mask_opt_idx = ~0;
@@ -257,7 +330,7 @@ void main() {
for (uint32_t j = start_j; j < end_j; ++j) {
coopmat<float16_t, gl_ScopeWorkgroup, Br, Bc, gl_MatrixUseAccumulator> mv = coopmat<float16_t, gl_ScopeWorkgroup, Br, Bc, gl_MatrixUseAccumulator>(0);
- if (MASK_ENABLE) {
+ if (MASK_ENABLE && !USE_SPARSE) {
if (USE_MASK_OPT && mask_opt_idx != j / 16) {
mask_opt_idx = j / 16;
@@ -315,7 +388,9 @@ void main() {
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 (k_use_decode) {
+ 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) {
coopMatLoadTensorNV(K_T, data_k, k_offset, sliceTensorLayoutNV(tensorLayoutK, j * Bc, Bc, 0, HSK_pad), tensorViewTranspose FADECODEK);
} else {
coopMatLoadTensorNV(K_T, data_k, k_offset, sliceTensorLayoutNV(tensorLayoutK, j * Bc, Bc, 0, HSK_pad), tensorViewTranspose);
@@ -330,7 +405,9 @@ void main() {
}
}
- if (MASK_ENABLE) {
+ if (MASK_ENABLE && USE_SPARSE) {
+ coopMatPerElementNV(S, S, faAddSparseMask, j);
+ } else if (MASK_ENABLE) {
S += slopeMat*coopmat<ACC_TYPE, gl_ScopeWorkgroup, Br, Bc, gl_MatrixUseAccumulator>(mv);
}
@@ -385,7 +462,9 @@ void main() {
coopMatLoadTensorNV(V, data_v, v_offset, sliceTensorLayoutNV(tensorLayoutV, j * Bc, Bc, 0, HSV_pad));
#else
const bool v_use_decode = (bs_v > 1u);
- if (v_use_decode) {
+ if (USE_SPARSE) {
+ coopMatLoadTensorNV(V, data_v, v_offset, sliceTensorLayoutNV(tensorLayoutV, j * Bc, Bc, 0, HSV_pad) FAGATHERV);
+ } else if (v_use_decode) {
coopMatLoadTensorNV(V, data_v, v_offset, sliceTensorLayoutNV(tensorLayoutV, j * Bc, Bc, 0, HSV_pad) FADECODEV);
} else {
coopMatLoadTensorNV(V, data_v, v_offset, sliceTensorLayoutNV(tensorLayoutV, j * Bc, Bc, 0, HSV_pad));
diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_sparse_compact.comp b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_sparse_compact.comp
new file mode 100644
index 000000000..3d3136266
--- /dev/null
+++ b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_sparse_compact.comp
@@ -0,0 +1,102 @@
+#version 450
+
+#extension GL_EXT_control_flow_attributes : enable
+#extension GL_EXT_shader_16bit_storage : require
+#extension GL_EXT_shader_explicit_arithmetic_types_int32 : require
+#ifdef USE_SUBGROUPS
+#extension GL_KHR_shader_subgroup_basic : require
+#extension GL_KHR_shader_subgroup_ballot : require
+#endif
+
+layout(local_size_x_id = 0, local_size_y = 1, local_size_z = 1) in;
+layout(constant_id = 0) const uint BLOCK_SIZE = 128;
+layout(constant_id = 1) const uint NUM_SUBGROUPS = 1;
+
+layout (binding = 0) readonly buffer M {float16_t data_m[];};
+layout (binding = 1) writeonly buffer I {int32_t data_i[];};
+
+layout (push_constant) uniform parameter {
+ uint KV;
+ uint nem1;
+ uint nem2;
+ uint nbm1;
+ uint nbm2;
+ uint nbm3;
+ uint n_kv_max;
+} p;
+
+#ifdef USE_SUBGROUPS
+shared uvec4 ballots_sh[NUM_SUBGROUPS];
+#else
+shared uint scan[BLOCK_SIZE];
+#endif
+
+// One workgroup per mask row: compact the finite-mask KV positions into a
+// per-row index list of length n_kv_max, -1 padded. Emitted in ascending KV
+// order so the downstream attention accumulation is deterministic.
+void main() {
+ const uint i1 = gl_WorkGroupID.x;
+ const uint i2 = gl_WorkGroupID.y;
+ const uint i3 = gl_WorkGroupID.z;
+ const uint tid = gl_LocalInvocationIndex;
+
+ const uint m_base = i3 * p.nbm3 + i2 * p.nbm2 + i1 * p.nbm1;
+ const uint out_base = ((i3 * p.nem2 + i2) * p.nem1 + i1) * p.n_kv_max;
+
+ uint base = 0;
+ for (uint chunk = 0; chunk < p.KV; chunk += BLOCK_SIZE) {
+ const uint k = chunk + tid;
+ bool selected = false;
+ if (k < p.KV) {
+ const float v = float(data_m[m_base + k]);
+ selected = !isinf(v) && !isnan(v);
+ }
+
+#ifdef USE_SUBGROUPS
+ const uvec4 ballot = subgroupBallot(selected);
+ if (subgroupElect()) {
+ ballots_sh[gl_SubgroupID] = ballot;
+ }
+ barrier();
+
+ uint subgroup_base = 0;
+ uint total = 0;
+ [[unroll]] for (uint s = 0; s < gl_NumSubgroups; ++s) {
+ if (s == gl_SubgroupID) {
+ subgroup_base = total;
+ }
+ total += subgroupBallotBitCount(ballots_sh[s]);
+ }
+ barrier();
+
+ const uint slot = base + subgroup_base + subgroupBallotExclusiveBitCount(ballot);
+#else
+ // Hillis-Steele inclusive prefix sum over the workgroup.
+ scan[tid] = selected ? 1u : 0u;
+ barrier();
+ for (uint off = 1; off < BLOCK_SIZE; off <<= 1) {
+ uint add = 0;
+ if (tid >= off) {
+ add = scan[tid - off];
+ }
+ barrier();
+ scan[tid] += add;
+ barrier();
+ }
+
+ const uint inclusive = scan[tid];
+ const uint total = scan[BLOCK_SIZE - 1];
+ const uint slot = base + inclusive - 1u;
+#endif
+
+ if (selected && slot < p.n_kv_max) {
+ data_i[out_base + slot] = int32_t(k);
+ }
+ base += total;
+ barrier();
+ }
+
+ for (uint s = min(base, p.n_kv_max) + tid; s < p.n_kv_max; s += BLOCK_SIZE) {
+ data_i[out_base + s] = int32_t(-1);
+ }
+}
diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp b/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp
index 30fe0884e..d3f425968 100644
--- a/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp
+++ b/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp
@@ -922,6 +922,8 @@ void process_shaders() {
string_to_spv("fa_split_k_reduce", "flash_attn_split_k_reduce.comp", {});
string_to_spv("fa_mask_opt", "flash_attn_mask_opt.comp", {});
+ string_to_spv("fa_sparse_compact", "flash_attn_sparse_compact.comp", {});
+ string_to_spv("fa_sparse_compact_subgroup", "flash_attn_sparse_compact.comp", {{"USE_SUBGROUPS", "1"}});
string_to_spv("quantize_q8_1", "quantize_q8_1.comp", {});
string_to_spv("quantize_q8_1_subgroup", "quantize_q8_1.comp", {{"USE_SUBGROUPS", "1"}});
diff --git a/tests/test-backend-ops.cpp b/tests/test-backend-ops.cpp
index 4f2665497..0e074770d 100644
--- a/tests/test-backend-ops.cpp
+++ b/tests/test-backend-ops.cpp
@@ -10707,6 +10707,9 @@ static std::vector<std::unique_ptr<test_case>> make_test_cases_eval() {
test_cases.emplace_back(new test_flash_attn_ext(128, 128, 1, { 8, 1}, 4096, 1, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_Q8_0, GGML_TYPE_Q8_0, {0, 1, 2, 3}, true, false, 512));
test_cases.emplace_back(new test_flash_attn_ext(128, 128, 1, { 8, 1}, 4096, 64, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_Q8_0, GGML_TYPE_Q8_0, {0, 1, 2, 3}, true, false, 512));
+ // Qwen QSA: 256/256, gqa 12, budget 2048.
+ test_cases.emplace_back(new test_flash_attn_ext(256, 256, 2, {12, 1}, 8192, 1, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16, {0, 1, 2, 3}, true, false, 2048));
+
// 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));
@@ -11165,6 +11168,14 @@ static std::vector<std::unique_ptr<test_case>> make_test_cases_perf() {
// Qwen3-VL-8B https://github.com/ggml-org/llama.cpp/issues/17012
test_cases.emplace_back(new test_flash_attn_ext(72, 72, 16, {1, 1}, 5776, 5776, false, false, 0, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16));
+ // Sparse flash attention (n_kv_max hint) decode across KV depths.
+ // Shapes: 576/512 DeepSeek MLA, 512/512 DeepSeek-V4/GLM-5.2, 256/256 gqa12 Qwen QSA.
+ for (int64_t kv : {4096, 16384, 32768}) {
+ test_cases.emplace_back(new test_flash_attn_ext(512, 512, 1, { 8, 1}, kv, 1, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16, {0, 1, 2, 3}, true, false, 512));
+ 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(64, 64, 8, {8, 1}, 7680, 1, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16));
test_cases.emplace_back(new test_flash_attn_ext(64, 64, 8, {8, 1}, 7680, 4, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16));
test_cases.emplace_back(new test_flash_attn_ext(64, 64, 8, {8, 1}, 7680, 1, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_Q4_0, GGML_TYPE_Q4_0));