Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
17 commits
Select commit Hold shift + click to select a range
20edd90
vulkan: sparse flash attention for qwen4exp top-k masks
LynxPDA Sep 7, 2026
7e56d92
vulkan: size the sparse-FA compaction per-subgroup tally to the workg…
LynxPDA Sep 10, 2026
c46ecbc
docs: recipe - shared per-subgroup tally sizing bug in the sparse FA …
LynxPDA Sep 10, 2026
949cc2c
docs: recipe - the sparse-FA shared-tile tiling contract
LynxPDA Sep 11, 2026
ec22659
vulkan: fix grouped union prefill host addressing for non-square head…
LynxPDA Sep 11, 2026
161ab59
docs: recipe - push constants carry the destination shape, not the di…
LynxPDA Sep 11, 2026
0b08831
feat: enable the grouped union prefill for GQA caches by default
LynxPDA Sep 11, 2026
0209047
vulkan: overlap the union scan of group g+1 with flash attention of g…
LynxPDA Sep 12, 2026
2b8bfbf
vulkan: key the union estimate and stat buffer by the real group size
LynxPDA Sep 13, 2026
a391db9
graph: skip an s_copy input that no node reads
LynxPDA Sep 13, 2026
29b5b84
qwen4exp: run the MTP draft block with its own indexer (sparse QSA)
LynxPDA Sep 13, 2026
c00d9b9
docs: recipe - sparse QSA for the qwen4exp MTP draft block
LynxPDA Sep 13, 2026
188988f
vulkan: widen the top-k radix pass and batch the emit scan
LynxPDA Sep 14, 2026
380b00a
vulkan: fold the QSA score transpose into the top-k gather
LynxPDA Sep 14, 2026
929aa8d
vulkan: clear the fusion label when a fusion is declined
LynxPDA Sep 14, 2026
cd81b81
docs: move PR-internal working notes out of the repo
LynxPDA Sep 15, 2026
1f4e257
vulkan: revert the CUDA sparse-FA consumer, keep the fork Vulkan-only
LynxPDA Sep 15, 2026
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
6 changes: 6 additions & 0 deletions ggml/include/ggml.h
Original file line number Diff line number Diff line change
Expand Up @@ -2433,6 +2433,12 @@ extern "C" {
GGML_API enum ggml_prec ggml_flash_attn_ext_get_prec(
const struct ggml_tensor * a);

// Use finite mask entries as a sparse K/V set. Set 0 to disable.
// n_kv_max must bound the number of finite entries in every mask row.
GGML_API void ggml_flash_attn_ext_set_sparse(
struct ggml_tensor * a,
int32_t n_kv_max);

GGML_API void ggml_flash_attn_ext_add_sinks(
struct ggml_tensor * a,
struct ggml_tensor * sinks);
Expand Down
696 changes: 611 additions & 85 deletions ggml/src/ggml-vulkan/ggml-vulkan.cpp

Large diffs are not rendered by default.

46 changes: 28 additions & 18 deletions ggml/src/ggml-vulkan/vulkan-shaders/flash_attn.comp
Original file line number Diff line number Diff line change
Expand Up @@ -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);
}
}
}
Expand Down Expand Up @@ -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]);
}
}

Expand Down Expand Up @@ -305,20 +308,22 @@ 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;
}

FLOAT_TYPEV4 K_Tf;
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]);
Expand All @@ -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;
}

Expand All @@ -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]);
Expand Down Expand Up @@ -493,14 +500,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]);
}
}

Expand All @@ -511,7 +519,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;
}

Expand All @@ -526,12 +536,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);
Expand Down
36 changes: 34 additions & 2 deletions ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_base.glsl
Original file line number Diff line number Diff line change
Expand Up @@ -28,7 +28,9 @@ const bool OLD_AMD_WINDOWS = (Flags & 8) != 0;
// GPU, where the row count is only known after a dedup pass and so cannot be pushed. The
// workgroup counts derive from neq1/neq2/neq3 and never from KV, so no indirect dispatch is
// needed: only this loop bound changes. Folds away for every other pipeline.
const bool DYNAMIC_KV = (Flags & 16) != 0;
const bool DYNAMIC_KV = (Flags & 32) != 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;
Expand Down Expand Up @@ -87,6 +89,9 @@ 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 KVB {uint32_t data_kv_dyn[];};
// The sparse index list shares this binding with the dynamic-KV row count (the two features
// are mutually exclusive), so the same 32-bit words are reinterpreted as signed indices.
#define data_sparse(i) (int(data_kv_dyn[(i)]))

#define MASK_OPT_ALL_NEG_INF 1
#define MASK_OPT_ALL_ZERO 2
Expand Down Expand Up @@ -199,7 +204,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, m_row_len, gqa_ratio, split_k_num, output_k_num;
q_stride, k_stride, v_stride, m_stride, m_row_len, gqa_ratio, split_k_num, output_k_num, sparse_base;
bool partial_output;

void init_indices()
Expand Down Expand Up @@ -281,6 +286,33 @@ void init_indices()
// under the sparse split path, where this dispatch only covers the raw prefix (KV) but
// the mask rows span the whole K range.
m_row_len = mask_stride_in_split_kv ? p.split_kv : 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.
Expand Down
49 changes: 33 additions & 16 deletions ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_cm1.comp
Original file line number Diff line number Diff line change
Expand Up @@ -179,9 +179,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)],
Expand Down Expand Up @@ -209,6 +216,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);
}
}
}
Expand All @@ -234,17 +243,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]);
}
}

Expand All @@ -269,25 +280,27 @@ 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) {
uint32_t col_vec = (idx + tid) % (MatBr / 4);
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]);
}
}

Expand Down Expand Up @@ -408,17 +421,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 vcol;
bool kv_active = fa_kv_index(j * Bc + c, vcol);
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 = 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
#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 + vcol * v_stride / 4 + d]);
}
}

Expand Down Expand Up @@ -455,21 +470,23 @@ 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 vcol;
bool kv_active = fa_kv_index(j * Bc + row, vcol);
const uint v_row = USE_SPARSE ? vcol : (j * Bc + 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);
Expand All @@ -491,7 +508,7 @@ void main() {
if (hsv_offset < HSV_pad) {
[[unroll]] for (uint32_t bc_chunk = 0; bc_chunk < Bc / MatBc; ++bc_chunk) {
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;
Expand Down
Loading