From b1f86ee1bd8cb702576f432b8a6fcd563ad01a62 Mon Sep 17 00:00:00 2001 From: Nathan Wilson Date: Mon, 13 Jul 2026 11:09:53 +0000 Subject: [PATCH 01/68] experiment: dequant-once FA scratch for all KV quant types (q4_0/q4_1/q5_0/q5_1/iq4_nl) Evidence branch only - NOT for upstream. Extends the q8_0 dequant-once FA path to every KV-eligible quant type via per-type fused dequant+transpose shaders, plus a prefill-only fa_kv_ok gate for iq4_nl (no native coopmat1 path) and a GGML_VK_FA_DEQUANT env toggle. Correctness: dequant-once == native FA bit-exact for q4_0/q4_1/q5_0/q5_1; iq4_nl matches CPU. Finding: prefill is quant-type independent (all dequant to identical f16 scratch); iq4_nl is a poor KV type (ppl ~2x q4_0 at equal bits). Retained as gating evidence. Assisted-by: Claude Opus 4.8 --- ggml/src/ggml-vulkan/ggml-vulkan.cpp | 9 +++++++- .../vulkan-shaders/dequant_iq4_nl.comp | 21 +++++++++++++++---- .../vulkan-shaders/dequant_q4_0.comp | 12 +++++++++++ .../vulkan-shaders/dequant_q4_1.comp | 11 ++++++++++ .../vulkan-shaders/dequant_q5_0.comp | 11 ++++++++++ .../vulkan-shaders/dequant_q5_1.comp | 11 ++++++++++ .../vulkan-shaders/vulkan-shaders-gen.cpp | 2 +- 7 files changed, 71 insertions(+), 6 deletions(-) diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp index 8147894b0e4..0b24dff2d88 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp +++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp @@ -5618,6 +5618,10 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { ggml_vk_create_pipeline(device, device->pipeline_dequant[GGML_TYPE_Q4_1], "dequant_q4_1", dequant_q4_1_len, dequant_q4_1_data, "main", 2, 5 * sizeof(uint32_t), {256 * 16, 1, 1}, {}, 1); ggml_vk_create_pipeline(device, device->pipeline_dequant[GGML_TYPE_Q5_0], "dequant_q5_0", dequant_q5_0_len, dequant_q5_0_data, "main", 2, 5 * sizeof(uint32_t), {256 * 16, 1, 1}, {}, 1); ggml_vk_create_pipeline(device, device->pipeline_dequant[GGML_TYPE_Q5_1], "dequant_q5_1", dequant_q5_1_len, dequant_q5_1_data, "main", 2, 5 * sizeof(uint32_t), {256 * 16, 1, 1}, {}, 1); + ggml_vk_create_pipeline(device, device->pipeline_dequant_transpose[GGML_TYPE_Q4_0], "dequant_q4_0_transpose", dequant_q4_0_transpose_len, dequant_q4_0_transpose_data, "main", 2, 5 * sizeof(uint32_t), {256 * 16, 1, 1}, {}, 1); + ggml_vk_create_pipeline(device, device->pipeline_dequant_transpose[GGML_TYPE_Q4_1], "dequant_q4_1_transpose", dequant_q4_1_transpose_len, dequant_q4_1_transpose_data, "main", 2, 5 * sizeof(uint32_t), {256 * 16, 1, 1}, {}, 1); + ggml_vk_create_pipeline(device, device->pipeline_dequant_transpose[GGML_TYPE_Q5_0], "dequant_q5_0_transpose", dequant_q5_0_transpose_len, dequant_q5_0_transpose_data, "main", 2, 5 * sizeof(uint32_t), {256 * 16, 1, 1}, {}, 1); + ggml_vk_create_pipeline(device, device->pipeline_dequant_transpose[GGML_TYPE_Q5_1], "dequant_q5_1_transpose", dequant_q5_1_transpose_len, dequant_q5_1_transpose_data, "main", 2, 5 * sizeof(uint32_t), {256 * 16, 1, 1}, {}, 1); ggml_vk_create_pipeline(device, device->pipeline_dequant[GGML_TYPE_Q8_0], "dequant_q8_0", dequant_q8_0_len, dequant_q8_0_data, "main", 2, 5 * sizeof(uint32_t), {256 * 16, 1, 1}, {}, 1); ggml_vk_create_pipeline(device, device->pipeline_dequant_transpose[GGML_TYPE_Q8_0], "dequant_q8_0_transpose", dequant_q8_0_transpose_len, dequant_q8_0_transpose_data, "main", 2, 5 * sizeof(uint32_t), {256 * 16, 1, 1}, {}, 1); ggml_vk_create_pipeline(device, device->pipeline_dequant[GGML_TYPE_Q2_K], "dequant_q2_k", dequant_q2_k_len, dequant_q2_k_data, "main", 2, 5 * sizeof(uint32_t), {256 * 64, 1, 1}, {}, 1); @@ -5640,6 +5644,7 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { ggml_vk_create_pipeline(device, device->pipeline_dequant[GGML_TYPE_Q3_0_ROCMFPX], "dequant_rocmfpx_fp3", dequant_rocmfpx_fp3_len, dequant_rocmfpx_fp3_data, "main", 2, 5 * sizeof(uint32_t), {256 * 16, 1, 1}, {}, 1); ggml_vk_create_pipeline(device, device->pipeline_dequant[GGML_TYPE_Q6_0_ROCMFPX], "dequant_rocmfpx_fp6", dequant_rocmfpx_fp6_len, dequant_rocmfpx_fp6_data, "main", 2, 5 * sizeof(uint32_t), {256 * 16, 1, 1}, {}, 1); ggml_vk_create_pipeline(device, device->pipeline_dequant[GGML_TYPE_Q8_0_ROCMFPX], "dequant_rocmfpx_fp8", dequant_rocmfpx_fp8_len, dequant_rocmfpx_fp8_data, "main", 2, 5 * sizeof(uint32_t), {256 * 16, 1, 1}, {}, 1); + ggml_vk_create_pipeline(device, device->pipeline_dequant_transpose[GGML_TYPE_IQ4_NL], "dequant_iq4_nl_transpose", dequant_iq4_nl_transpose_len, dequant_iq4_nl_transpose_data, "main", 2, 5 * sizeof(uint32_t), {256 * 16, 1, 1}, {}, 1); ggml_vk_create_pipeline(device, device->pipeline_dequant[GGML_TYPE_MXFP4], "dequant_mxfp4", dequant_mxfp4_len, dequant_mxfp4_data, "main", 2, 5 * sizeof(uint32_t), {256 * 16, 1, 1}, {}, 1); ggml_vk_create_pipeline(device, device->pipeline_dequant[GGML_TYPE_NVFP4], "dequant_nvfp4", dequant_nvfp4_len, dequant_nvfp4_data, "main", 2, 5 * sizeof(uint32_t), {256 * 16, 1, 1}, {}, 1); @@ -11262,7 +11267,9 @@ static void ggml_vk_flash_attn(ggml_backend_vk_context * ctx, vk_context& subctx }; const bool k_quant = k->type != GGML_TYPE_F16 && k->type != GGML_TYPE_BF16 && k->type != GGML_TYPE_F32; const bool v_quant = v->type != GGML_TYPE_F16 && v->type != GGML_TYPE_BF16 && v->type != GGML_TYPE_F32; - const bool use_dequant_kv = k_quant && v_quant && neq1 >= 64 && + static const char * fa_dequant_env = getenv("GGML_VK_FA_DEQUANT"); + const bool fa_dequant_off = fa_dequant_env && fa_dequant_env[0] == '0'; + const bool use_dequant_kv = !fa_dequant_off && k_quant && v_quant && neq1 >= 64 && is_dense_kv_cache(k) && is_dense_kv_cache(v) && (uint64_t)ggml_nelements(k) * sizeof(ggml_fp16_t) <= ctx->device->properties.limits.maxStorageBufferRange && (uint64_t)ggml_nelements(v) * sizeof(ggml_fp16_t) <= ctx->device->properties.limits.maxStorageBufferRange && diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/dequant_iq4_nl.comp b/ggml/src/ggml-vulkan/vulkan-shaders/dequant_iq4_nl.comp index 8f7833eab2e..befcdc1c510 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/dequant_iq4_nl.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/dequant_iq4_nl.comp @@ -10,8 +10,6 @@ layout (binding = 1) writeonly buffer D {D_TYPE data_b[];}; void main() { const uint i = gl_WorkGroupID.x * 4 + gl_LocalInvocationID.x / 64; - init_iq_shmem(gl_WorkGroupSize); - const uint tid = gl_LocalInvocationID.x % 64; const uint il = tid/32; const uint ir = tid%32; @@ -21,12 +19,27 @@ void main() { } const uint q_idx = 8*il; + +#ifdef DEQUANT_TRANSPOSE + // Fused dequant+transpose for FA quant-KV: source physical is [HS, NH, KV, NS] + // (p.M=HS, p.K=NH, p.stride_a=KV); write to per-head-contiguous dest [HS, KV, NH, NS] so the + // f16 FA reads KV coalesced. An iq4_nl block = 32 consecutive HS elements at fixed (head,kv) -> + // 32 contiguous dest positions (intra-block nibble offsets unchanged). + const uint HS = p.M, NH = p.K, KVn = p.stride_a; + const uint e0 = ib * 32; + const uint b_idx = (e0 % HS) + + ((e0 / (HS * NH)) % KVn) * HS + + ((e0 / HS) % NH) * (HS * KVn) + + (e0 / (HS * NH * KVn)) * (HS * KVn * NH) + + q_idx; +#else const uint b_idx = 1024*i + 32*ir + q_idx; +#endif const float d = float(data_a[ib].d); [[unroll]] for (uint l = 0; l < 8; ++l) { - data_b[b_idx + l + 0] = D_TYPE(d * kvalues_iq4nl[data_a[ib].qs[q_idx + l] & 0xF]); - data_b[b_idx + l + 16] = D_TYPE(d * kvalues_iq4nl[data_a[ib].qs[q_idx + l] >> 4]); + data_b[b_idx + l + 0] = D_TYPE(d * float(kvalues_iq4nl_const[data_a[ib].qs[q_idx + l] & 0xF])); + data_b[b_idx + l + 16] = D_TYPE(d * float(kvalues_iq4nl_const[data_a[ib].qs[q_idx + l] >> 4])); } } diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/dequant_q4_0.comp b/ggml/src/ggml-vulkan/vulkan-shaders/dequant_q4_0.comp index b92b292135b..51a6c89a60c 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/dequant_q4_0.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/dequant_q4_0.comp @@ -19,7 +19,19 @@ void main() { } const uint q_idx = 8*il; + +#ifdef DEQUANT_TRANSPOSE + // fused dequant+transpose for FA quant-KV: per-head-contiguous f16 scratch (see dequant_q8_0.comp) + const uint HS = p.M, NH = p.K, KVn = p.stride_a; + const uint e0 = ib * 32; + const uint b_idx = (e0 % HS) + + ((e0 / (HS * NH)) % KVn) * HS + + ((e0 / HS) % NH) * (HS * KVn) + + (e0 / (HS * NH * KVn)) * (HS * KVn * NH) + + q_idx; +#else const uint b_idx = 1024*i + 32*ir + q_idx; +#endif const float d = float(data_a[ib].d); diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/dequant_q4_1.comp b/ggml/src/ggml-vulkan/vulkan-shaders/dequant_q4_1.comp index 6b63cbe5833..76f2d958cba 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/dequant_q4_1.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/dequant_q4_1.comp @@ -18,7 +18,18 @@ void main() { return; } +#ifdef DEQUANT_TRANSPOSE + // fused dequant+transpose for FA quant-KV: per-head-contiguous f16 scratch (see dequant_q8_0.comp) + const uint HS = p.M, NH = p.K, KVn = p.stride_a; + const uint e0 = ib * 32; + const uint b_idx = (e0 % HS) + + ((e0 / (HS * NH)) % KVn) * HS + + ((e0 / HS) % NH) * (HS * KVn) + + (e0 / (HS * NH * KVn)) * (HS * KVn * NH) + + 8*il; +#else const uint b_idx = 1024*i + 32*ir + 8*il; +#endif const float d = float(data_a[ib].d); const float m = float(data_a[ib].m); diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/dequant_q5_0.comp b/ggml/src/ggml-vulkan/vulkan-shaders/dequant_q5_0.comp index f1b0bac8727..1402fa4292d 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/dequant_q5_0.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/dequant_q5_0.comp @@ -18,7 +18,18 @@ void main() { return; } +#ifdef DEQUANT_TRANSPOSE + // fused dequant+transpose for FA quant-KV: per-head-contiguous f16 scratch (see dequant_q8_0.comp) + const uint HS = p.M, NH = p.K, KVn = p.stride_a; + const uint e0 = ib * 32; + const uint b_idx = (e0 % HS) + + ((e0 / (HS * NH)) % KVn) * HS + + ((e0 / HS) % NH) * (HS * KVn) + + (e0 / (HS * NH * KVn)) * (HS * KVn * NH) + + 8*il; +#else const uint b_idx = 1024*i + 32*ir + 8*il; +#endif const float d = float(data_a[ib].d); const uint qh = uint(data_a[ib].qh[1]) << 16 | data_a[ib].qh[0]; diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/dequant_q5_1.comp b/ggml/src/ggml-vulkan/vulkan-shaders/dequant_q5_1.comp index c495b31f175..1fd6e2552af 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/dequant_q5_1.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/dequant_q5_1.comp @@ -18,7 +18,18 @@ void main() { return; } +#ifdef DEQUANT_TRANSPOSE + // fused dequant+transpose for FA quant-KV: per-head-contiguous f16 scratch (see dequant_q8_0.comp) + const uint HS = p.M, NH = p.K, KVn = p.stride_a; + const uint e0 = ib * 32; + const uint b_idx = (e0 % HS) + + ((e0 / (HS * NH)) % KVn) * HS + + ((e0 / HS) % NH) * (HS * KVn) + + (e0 / (HS * NH * KVn)) * (HS * KVn * NH) + + 8*il; +#else const uint b_idx = 1024*i + 32*ir + 8*il; +#endif const float d = float(data_a[ib].d); const float m = float(data_a[ib].m); 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 1a57fdc24e8..8baf0f2ef4b 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp @@ -791,7 +791,7 @@ void process_shaders() { string_to_spv("dequant_" + tname, "dequant_" + tname + ".comp", merge_maps(base_dict, {{data_a_key, "1"}, {"D_TYPE", "float16_t"}})); } // Fused dequant+transpose variant for FA quant-KV (per-head-contiguous f16 scratch). - if (tname == "q8_0") { + if (tname == "q8_0" || tname == "iq4_nl" || tname == "q4_0" || tname == "q4_1" || tname == "q5_0" || tname == "q5_1") { string_to_spv("dequant_" + tname + "_transpose", "dequant_" + tname + ".comp", merge_maps(base_dict, {{data_a_key, "1"}, {"D_TYPE", "float16_t"}, {"DEQUANT_TRANSPOSE", "1"}})); } From 6dc9116a0528831bb5d24771cd823cb19bc4679f Mon Sep 17 00:00:00 2001 From: Nathan Wilson Date: Tue, 14 Jul 2026 14:09:02 +0000 Subject: [PATCH 02/68] vulkan : gate FA dequant-once scratch on device-local capacity The dequant-once path materializes a per-layer f16 K/V scratch (~2 KiB/token). On discrete devices that resident footprint can push the working set past free VRAM, at which point the driver silently pages device-local memory: measured ~15x prefill regression on an 8 GB card (RTX 3070, driver 591.86) at long context, with no error reported. Integrated/UMA devices have no separate device pool to overflow and are unaffected. Gate the path on this process's device-local usage against the physical heap size, keeping a conservative reserve for memory not observable in-process. heapBudget is deliberately not used as the signal: ggml_backend_vk_get_device_memory computes heapBudget - heapUsage in unsigned arithmetic, which wraps to a huge value exactly when the device is oversubscribed. The allocation cannot gate itself either - on WDDM vkAllocateMemory only fails at roughly physical heap size, which is above the free-VRAM level where paging begins, so a successful allocation is not evidence of a resident fit. Also fix the scratch size check: K and V are bound as one storage buffer, so their sum must fit maxStorageBufferRange, not each half independently. GGML_VK_FA_DEQUANT=0 forces the path off and =1 skips the capacity check; GGML_VK_FA_DEQUANT_RESERVE_MB overrides the reserve. Assisted-by: Claude (Opus 5) --- ggml/src/ggml-vulkan/ggml-vulkan.cpp | 87 ++++++++++++++++++++++++++-- 1 file changed, 83 insertions(+), 4 deletions(-) diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp index 0b24dff2d88..096aefa9730 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp +++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp @@ -2416,6 +2416,10 @@ struct ggml_backend_vk_context { ggml_vk_garbage_collector gc; size_t prealloc_size_x, prealloc_size_y, prealloc_size_split_k, prealloc_size_add_rms_partials, prealloc_size_add_rms_partials_offset; vk_buffer prealloc_x, prealloc_y, prealloc_split_k, prealloc_add_rms_partials, sync_staging; + // memoized capacity decision for the FA dequant-once scratch, see ggml_vk_fa_dequant_scratch_fits + uint64_t fa_dequant_gate_sz; + bool fa_dequant_gate_fits; + bool fa_dequant_gate_logged; vk::Fence fence, almost_ready_fence; bool submit_pending {}; bool almost_ready_fence_pending {}; @@ -7897,6 +7901,9 @@ static void ggml_vk_init(ggml_backend_vk_context * ctx, size_t idx) { ctx->prealloc_size_x = 0; ctx->prealloc_size_y = 0; ctx->prealloc_size_split_k = 0; + ctx->fa_dequant_gate_sz = 0; + ctx->fa_dequant_gate_fits = false; + ctx->fa_dequant_gate_logged = false; // Fixed size of 1KB, for deterministic behavior ctx->prealloc_size_add_rms_partials = 1024; @@ -11199,6 +11206,69 @@ static bool ggml_vk_flash_attn_coopmat_shmem_support(const vk_device& device, co return supported; } +// Capacity gate for the dequant-once f16 K/V scratch. On discrete devices the scratch can push the +// working set past free VRAM, and the driver then silently pages device-local memory (~15x prefill +// regression measured on an 8 GB card at long context). UMA has no separate pool to overflow. +// +// Gates our own device-local usage against the physical heap size, less a reserve for memory not +// visible in-process. heapBudget is deliberately not the signal: ggml_backend_vk_get_device_memory +// computes heapBudget - heapUsage unsigned, which wraps when the device is oversubscribed. Nor can +// the allocation gate itself: on WDDM vkAllocateMemory only fails near physical heap size, above +// the free-VRAM level where paging starts. The reserve is necessarily conservative because other +// processes' VRAM use is invisible to us. GGML_VK_FA_DEQUANT=0/1 forces the path off/on; +// GGML_VK_FA_DEQUANT_RESERVE_MB overrides the reserve. +static bool ggml_vk_fa_dequant_scratch_fits(ggml_backend_vk_context * ctx, uint64_t scratch_sz) { + const vk_device& device = ctx->device; + + if (device->uma) { + return true; + } + + // Decided once per scratch size, so the decision cannot flip between layers as usage grows. + if (ctx->fa_dequant_gate_sz == scratch_sz) { + return ctx->fa_dequant_gate_fits; + } + + static const uint64_t reserve = [] { + const char * env = getenv("GGML_VK_FA_DEQUANT_RESERVE_MB"); + return (uint64_t)(env ? atoi(env) : 1024) * 1024 * 1024; + }(); + + bool fits = false; + + // Without VK_EXT_memory_budget our usage is unknowable, so leave the path disabled. + if (vk_instance.device_supports_membudget[device->idx]) { + vk::PhysicalDeviceMemoryBudgetPropertiesEXT budgetprops; + vk::PhysicalDeviceMemoryProperties2 memprops = {}; + memprops.pNext = &budgetprops; + device->physical_device.getMemoryProperties2(&memprops); + + uint64_t heap_size = 0; + uint64_t heap_used = 0; + for (uint32_t i = 0; i < memprops.memoryProperties.memoryHeapCount; ++i) { + const vk::MemoryHeap & heap = memprops.memoryProperties.memoryHeaps[i]; + if (heap.flags & vk::MemoryHeapFlagBits::eDeviceLocal) { + heap_size += heap.size; + heap_used += budgetprops.heapUsage[i]; + } + } + // heap_used already covers scratch allocated on a previous ubatch, so counting scratch_sz + // in full is conservative by up to the current scratch size. + fits = heap_size > reserve && heap_used + scratch_sz + reserve <= heap_size; + } + + if (!fits && !ctx->fa_dequant_gate_logged) { + ctx->fa_dequant_gate_logged = true; + GGML_LOG_INFO("ggml_vulkan: flash attention dequant-once disabled: %llu MiB K/V scratch does not fit " + "device-local memory with a %llu MiB reserve. Set GGML_VK_FA_DEQUANT=1 to force it on.\n", + (unsigned long long)(scratch_sz >> 20), (unsigned long long)(reserve >> 20)); + } + + ctx->fa_dequant_gate_sz = scratch_sz; + ctx->fa_dequant_gate_fits = fits; + return fits; +} + static void ggml_vk_flash_attn(ggml_backend_vk_context * ctx, vk_context& subctx, const ggml_tensor * q, const ggml_tensor * k, const ggml_tensor * v, const ggml_tensor * mask, const ggml_tensor * sinks, ggml_tensor * dst) { VK_LOG_DEBUG("ggml_vk_flash_attn((" << q << ", name=" << q->name << ", type=" << q->type << ", ne0=" << q->ne[0] << ", ne1=" << q->ne[1] << ", ne2=" << q->ne[2] << ", ne3=" << q->ne[3] << ", nb0=" << q->nb[0] << ", nb1=" << q->nb[1] << ", nb2=" << q->nb[2] << ", nb3=" << q->nb[3]; std::cerr << "), (" << k << ", name=" << k->name << ", type=" << k->type << ", ne0=" << k->ne[0] << ", ne1=" << k->ne[1] << ", ne2=" << k->ne[2] << ", ne3=" << k->ne[3] << ", nb0=" << k->nb[0] << ", nb1=" << k->nb[1] << ", nb2=" << k->nb[2] << ", nb3=" << k->nb[3]; @@ -11258,7 +11328,14 @@ static void ggml_vk_flash_attn(ggml_backend_vk_context * ctx, vk_context& subctx const bool f32acc = !ctx->device->fp16 || dst->op_params[3] == GGML_PREC_F32 || k->type == GGML_TYPE_BF16; - // dequant K/V once into an f16 scratch, reordered KV layout so FA can read without a stride + // For prefill with quantized K/V, dequantize+transpose K/V once into a per-head-contiguous + // f16 scratch and run the f16 FA path, instead of the coopmat1 shader re-dequantizing the + // whole KV inside every Q-workgroup. The KV-cache view reaching FA is [0,2,1,3]-permuted but + // dense, so we require dense allocation (not ggml_is_contiguous) and block-contiguous dim0, + // and only engage where a fused dequant-transpose shader exists. Prefill only + // (n_rows >= 64); measured neutral at shallow depth and up to ~2x at long context. + // The scratch is bound as a single storage buffer holding K and V back to back, so the SUM of + // the two must fit maxStorageBufferRange, not each half independently. auto is_dense_kv_cache = [](const ggml_tensor * t) { return t->nb[0] == ggml_type_size(t->type) && t->nb[2] == ggml_row_size(t->type, t->ne[0]) && @@ -11267,19 +11344,21 @@ static void ggml_vk_flash_attn(ggml_backend_vk_context * ctx, vk_context& subctx }; const bool k_quant = k->type != GGML_TYPE_F16 && k->type != GGML_TYPE_BF16 && k->type != GGML_TYPE_F32; const bool v_quant = v->type != GGML_TYPE_F16 && v->type != GGML_TYPE_BF16 && v->type != GGML_TYPE_F32; + const uint64_t kv_f16_sz = ((uint64_t)ggml_nelements(k) + (uint64_t)ggml_nelements(v)) * sizeof(ggml_fp16_t); static const char * fa_dequant_env = getenv("GGML_VK_FA_DEQUANT"); const bool fa_dequant_off = fa_dequant_env && fa_dequant_env[0] == '0'; + const bool fa_dequant_on = fa_dequant_env && fa_dequant_env[0] == '1'; const bool use_dequant_kv = !fa_dequant_off && k_quant && v_quant && neq1 >= 64 && is_dense_kv_cache(k) && is_dense_kv_cache(v) && - (uint64_t)ggml_nelements(k) * sizeof(ggml_fp16_t) <= ctx->device->properties.limits.maxStorageBufferRange && - (uint64_t)ggml_nelements(v) * sizeof(ggml_fp16_t) <= ctx->device->properties.limits.maxStorageBufferRange && + kv_f16_sz <= ctx->device->properties.limits.maxStorageBufferRange && ctx->device->pipeline_dequant_transpose[k->type] != nullptr && ctx->device->pipeline_dequant_transpose[v->type] != nullptr && // coopmat2 path does not benefit from the f16 scratch !ctx->device->coopmat2 && // Intel Xe1 regresses, see PR 25494 (ctx->device->vendor_id != VK_VENDOR_ID_INTEL || - (ctx->device->coopmat_support && ctx->device->architecture != vk_device_architecture::INTEL_XE1)); + (ctx->device->coopmat_support && ctx->device->architecture != vk_device_architecture::INTEL_XE1)) && + (fa_dequant_on || ggml_vk_fa_dequant_scratch_fits(ctx, kv_f16_sz)); const ggml_type k_type_eff = use_dequant_kv ? GGML_TYPE_F16 : k->type; const ggml_type v_type_eff = use_dequant_kv ? GGML_TYPE_F16 : v->type; From e52cc7cee0f318fb13e1eb3006faa5fc7d9bf5ca Mon Sep 17 00:00:00 2001 From: Nathan Wilson Date: Tue, 28 Jul 2026 10:21:44 +0000 Subject: [PATCH 03/68] vulkan: contiguize strided f16 KV for FA prefill (GGML_VK_FA_KV_CONTIG, env-gated) The KV-cache view reaching FA has head-interleaved rows ([HS, NH, KV] physically), and the cm1 shader's direct-from-global coopMatLoads run ~2x slower on that stride than on per-head-contiguous K/V: a 16x16 tile touches 16 distant cache lines instead of 4. Measured on Strix Halo (RADV gfx1151), hd128/GQA8/kv10240/nb2048 f16: 29.8ms contiguous vs 63.1ms dense-permuted (the model layout; matches the in-model 59.9ms from the perf logger, where FLASH_ATTN_EXT was 72.6% of the graph at pp2048@d8192). GGML_VK_FA_KV_CONTIG=1 extends the dequant-once FA scratch to f16 K/V: dequant_f16_transpose.comp is a pure strided copy ([HS,NH,KV] -> [HS,KV,NH], same push-constant ABI and dispatch as the quant transpose shaders), engaged only when the rows are actually strided, prefill only (neq1 >= 64). FA op 63.1 -> 30.5ms (2.07x) incl. copy cost. Model-level (Qwen3-Coder-30B Q6_K_XL, ub2048, f16 KV, r=3, vs ROCm 571d0d5 nowmma): pp8192 877 -> 1199 t/s (ROCm 1216, was -28% now parity); pp2048@d4096/8192/16384: 776/513/300 -> 1120/850/580. Shallow prefill unchanged-to-better (pp2048 1542 -> 1633). Not the fix: shmem staging on AMD (loses on occupancy, 29.8 -> 54.4ms contiguous), bigger-tile/GQA-packed streaming (1-KV-head L2-resident probe runs identical -> kernel is issue-bound, not bandwidth-bound). test-backend-ops -o FLASH_ATTN_EXT green with the flag off and on (pre-existing iq4_nl+sinks failures unchanged). Co-Authored-By: Claude Fable 5 Assisted-by: Claude (Opus 5) --- ggml/src/ggml-vulkan/ggml-vulkan.cpp | 14 +++++++- .../vulkan-shaders/dequant_f16_transpose.comp | 32 +++++++++++++++++++ .../vulkan-shaders/vulkan-shaders-gen.cpp | 4 +++ 3 files changed, 49 insertions(+), 1 deletion(-) create mode 100644 ggml/src/ggml-vulkan/vulkan-shaders/dequant_f16_transpose.comp diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp index 096aefa9730..5f353e5630d 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp +++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp @@ -5628,6 +5628,7 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { ggml_vk_create_pipeline(device, device->pipeline_dequant_transpose[GGML_TYPE_Q5_1], "dequant_q5_1_transpose", dequant_q5_1_transpose_len, dequant_q5_1_transpose_data, "main", 2, 5 * sizeof(uint32_t), {256 * 16, 1, 1}, {}, 1); ggml_vk_create_pipeline(device, device->pipeline_dequant[GGML_TYPE_Q8_0], "dequant_q8_0", dequant_q8_0_len, dequant_q8_0_data, "main", 2, 5 * sizeof(uint32_t), {256 * 16, 1, 1}, {}, 1); ggml_vk_create_pipeline(device, device->pipeline_dequant_transpose[GGML_TYPE_Q8_0], "dequant_q8_0_transpose", dequant_q8_0_transpose_len, dequant_q8_0_transpose_data, "main", 2, 5 * sizeof(uint32_t), {256 * 16, 1, 1}, {}, 1); + ggml_vk_create_pipeline(device, device->pipeline_dequant_transpose[GGML_TYPE_F16], "dequant_f16_transpose", dequant_f16_transpose_len, dequant_f16_transpose_data, "main", 2, 5 * sizeof(uint32_t), {256 * 8, 1, 1}, {}, 1); ggml_vk_create_pipeline(device, device->pipeline_dequant[GGML_TYPE_Q2_K], "dequant_q2_k", dequant_q2_k_len, dequant_q2_k_data, "main", 2, 5 * sizeof(uint32_t), {256 * 64, 1, 1}, {}, 1); ggml_vk_create_pipeline(device, device->pipeline_dequant[GGML_TYPE_TQ2_0], "dequant_tq2_0", dequant_tq2_0_len, dequant_tq2_0_data, "main", 2, 5 * sizeof(uint32_t), {256 * 64, 1, 1}, {}, 1); ggml_vk_create_pipeline(device, device->pipeline_dequant[GGML_TYPE_Q3_K], "dequant_q3_k", dequant_q3_k_len, dequant_q3_k_data, "main", 2, 5 * sizeof(uint32_t), {256 * 64, 1, 1}, {}, 1); @@ -11348,7 +11349,18 @@ static void ggml_vk_flash_attn(ggml_backend_vk_context * ctx, vk_context& subctx static const char * fa_dequant_env = getenv("GGML_VK_FA_DEQUANT"); const bool fa_dequant_off = fa_dequant_env && fa_dequant_env[0] == '0'; const bool fa_dequant_on = fa_dequant_env && fa_dequant_env[0] == '1'; - const bool use_dequant_kv = !fa_dequant_off && k_quant && v_quant && neq1 >= 64 && + // EXPERIMENT (GGML_VK_FA_KV_CONTIG=1): run the same contiguize pass for f16 K/V. The + // KV-cache view reaching FA is head-interleaved ([HS, NH, KV] physically), and the cm1 + // direct-from-global coopMatLoads run ~2-5x slower on those strided rows than on + // per-head-contiguous K/V. Copy K/V once into the scratch instead (dequant_f16_transpose + // is a pure strided copy). Only engages when the rows are actually strided. + static const char * fa_kv_contig_env = getenv("GGML_VK_FA_KV_CONTIG"); + const bool fa_kv_contig = fa_kv_contig_env && fa_kv_contig_env[0] == '1'; + const bool kv_f16_strided = k->type == GGML_TYPE_F16 && v->type == GGML_TYPE_F16 && + (k->nb[1] != (uint64_t)HSK * sizeof(ggml_fp16_t) || + v->nb[1] != (uint64_t)HSV * sizeof(ggml_fp16_t)) && + (HSK % 8) == 0 && (HSV % 8) == 0; + const bool use_dequant_kv = !fa_dequant_off && ((k_quant && v_quant) || (fa_kv_contig && kv_f16_strided)) && neq1 >= 64 && is_dense_kv_cache(k) && is_dense_kv_cache(v) && kv_f16_sz <= ctx->device->properties.limits.maxStorageBufferRange && ctx->device->pipeline_dequant_transpose[k->type] != nullptr && diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/dequant_f16_transpose.comp b/ggml/src/ggml-vulkan/vulkan-shaders/dequant_f16_transpose.comp new file mode 100644 index 00000000000..dda53d749ac --- /dev/null +++ b/ggml/src/ggml-vulkan/vulkan-shaders/dequant_f16_transpose.comp @@ -0,0 +1,32 @@ +#version 450 + +#include "dequant_head.glsl" + +#extension GL_EXT_shader_explicit_arithmetic_types_float16 : require + +layout(local_size_x = 256, local_size_y = 1, local_size_z = 1) in; + +layout (binding = 0) readonly buffer A {f16vec4 data_a[];}; +layout (binding = 1) writeonly buffer D {f16vec4 data_b[];}; + +// Strided-copy counterpart of the fused dequant+transpose shaders for FA f16 KV: +// source physical is [HS, NH, KV, NS] (p.M=HS, p.K=NH, p.stride_a=KV); write to +// per-head-contiguous dest [HS, KV, NH, NS] so the f16 FA reads KV coalesced. +// HS stays innermost in both layouts, so each invocation moves 8 HS-contiguous +// elements (two f16vec4) requiring HS % 8 == 0 (enforced by the host gate). +void main() { + const uint i = gl_GlobalInvocationID.x; + const uint e0 = i * 8; + if (e0 >= p.nel) { + return; + } + + const uint HS = p.M, NH = p.K, KVn = p.stride_a; + const uint dst = (e0 % HS) + + ((e0 / (HS * NH)) % KVn) * HS + + ((e0 / HS) % NH) * (HS * KVn) + + (e0 / (HS * NH * KVn)) * (HS * KVn * NH); + + data_b[dst / 4 ] = data_a[e0 / 4 ]; + data_b[dst / 4 + 1] = data_a[e0 / 4 + 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 8baf0f2ef4b..81e556b0c78 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp @@ -794,6 +794,10 @@ void process_shaders() { if (tname == "q8_0" || tname == "iq4_nl" || tname == "q4_0" || tname == "q4_1" || tname == "q5_0" || tname == "q5_1") { string_to_spv("dequant_" + tname + "_transpose", "dequant_" + tname + ".comp", merge_maps(base_dict, {{data_a_key, "1"}, {"D_TYPE", "float16_t"}, {"DEQUANT_TRANSPOSE", "1"}})); } + // Strided-copy counterpart for f16 KV (contiguize the head-interleaved cache layout). + if (tname == "f16") { + string_to_spv("dequant_f16_transpose", "dequant_f16_transpose.comp", {}); + } shader = (tname == "f32" || tname == "f16" || tname == "bf16") ? "get_rows.comp" : "get_rows_quant.comp"; From d61a7fd6c4b9c9e4abe0041bba7d33039ef8bd68 Mon Sep 17 00:00:00 2001 From: Nathan Wilson Date: Tue, 28 Jul 2026 10:21:56 +0000 Subject: [PATCH 04/68] tests: dense-permuted K/V option + Strix FA prefill perf/probe cases test_flash_attn_ext always built K/V as sparse views (physical dim1 doubled, half viewed), which can never satisfy the contiguize path's ggml_is_contiguously_allocated gate - so no permuted correctness case exercised it. Add a kv_view parameter (default true = unchanged) and dense-permuted eval cases matching the real KV-cache layout, including ALiBi and logit-softcap variants; all pass vs CPU with GGML_VK_FA_KV_CONTIG=1. Perf additions: Qwen3-Coder-30B prefill-at-depth shapes (hd128, 4 KV heads, GQA 8, kv up to 10240, nb 512/2048), the dense-permuted variant (model layout), and the probe set used to establish that the contiguous cm1 kernel is issue-bound: 32-distinct-KV-head MALL-spill (flat), 1-KV-head L2-resident (flat), mask=0 (-5.5%), f16 acc (-3%). Co-Authored-By: Claude Fable 5 Assisted-by: Claude (Opus 5) --- tests/test-backend-ops.cpp | 25 ++++++++++++++++++++++++- 1 file changed, 24 insertions(+), 1 deletion(-) diff --git a/tests/test-backend-ops.cpp b/tests/test-backend-ops.cpp index 3919e555273..606ecc53b95 100644 --- a/tests/test-backend-ops.cpp +++ b/tests/test-backend-ops.cpp @@ -7915,7 +7915,7 @@ struct test_flash_attn_ext : public test_case { const ggml_type type_K; const ggml_type type_V; std::array permute; - const bool kv_view; // create K/V as views of a larger buffer (like a KV cache) + const bool kv_view; // create K/V as views of a larger buffer (like a KV cache); false = dense permuted like the model KV cache const bool v_is_view_of_k; const int64_t n_kv_max; @@ -10986,6 +10986,12 @@ static std::vector> make_test_cases_eval() { } } + // dense-permuted K/V (model KV-cache layout, engages the f16 contiguize path at nb>=64) + test_cases.emplace_back(new test_flash_attn_ext(128, 128, 4, {8, 1}, 1024, 128, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16, {0, 2, 1, 3}, false)); + test_cases.emplace_back(new test_flash_attn_ext(96, 96, 8, {4, 1}, 512, 80, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16, {0, 2, 1, 3}, false)); + test_cases.emplace_back(new test_flash_attn_ext(64, 64, 4, {8, 1}, 512, 75, true, false, 8.0f, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16, {0, 2, 1, 3}, false)); + test_cases.emplace_back(new test_flash_attn_ext(128, 128, 4, {8, 1}, 512, 96, true, false, 0, 30.0f, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16, {0, 2, 1, 3}, false)); + // mixed quant and Q1_0 test cases test_cases.emplace_back(new test_flash_attn_ext(64, 64, 4, {1, 1}, 128, 2, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_Q8_0, GGML_TYPE_Q4_0)); test_cases.emplace_back(new test_flash_attn_ext(64, 64, 4, {1, 1}, 128, 2, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_Q4_0, GGML_TYPE_F16)); @@ -11534,6 +11540,23 @@ static std::vector> 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)); + // Qwen3-Coder-30B-A3B prefill at depth: hd128, 4 KV heads, GQA 8, ub 2048 + test_cases.emplace_back(new test_flash_attn_ext(128, 128, 4, {8, 1}, 2048, 2048, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16)); + test_cases.emplace_back(new test_flash_attn_ext(128, 128, 4, {8, 1}, 6144, 2048, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16)); + test_cases.emplace_back(new test_flash_attn_ext(128, 128, 4, {8, 1}, 10240, 2048, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16)); + test_cases.emplace_back(new test_flash_attn_ext(128, 128, 4, {8, 1}, 10240, 512, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16)); + // MALL-spill probe: same FLOPs, 32 distinct KV heads (no GQA) -> K/V footprint 8x (168MB > 32MB MALL) + test_cases.emplace_back(new test_flash_attn_ext(128, 128, 32, {1, 1}, 10240, 2048, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16)); + // KV-cache layout probe: same shape, K/V strided token-major like the real cache (heads interleaved) + test_cases.emplace_back(new test_flash_attn_ext(128, 128, 4, {8, 1}, 10240, 2048, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16, {0, 2, 1, 3})); + // Same, dense-permuted (exact model KV-cache layout; eligible for the f16 contiguize path) + test_cases.emplace_back(new test_flash_attn_ext(128, 128, 4, {8, 1}, 10240, 2048, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16, {0, 2, 1, 3}, false)); + // L2-residence probe: 1 KV head x GQA 32 (K/V 5.2MB fits L2) - distinguishes cache-BW-bound from issue-bound + test_cases.emplace_back(new test_flash_attn_ext(128, 128, 1, {32, 1}, 10240, 2048, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16)); + // cost-partition probes: no mask; f16 accumulate + test_cases.emplace_back(new test_flash_attn_ext(128, 128, 4, {8, 1}, 10240, 2048, false, false, 0, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16)); + test_cases.emplace_back(new test_flash_attn_ext(128, 128, 4, {8, 1}, 10240, 2048, true, false, 0, 0, GGML_PREC_DEFAULT, 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_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)); From 6d056bf9db9854e6f1c31907f00c412d1ccbfd87 Mon Sep 17 00:00:00 2001 From: Nathan Wilson Date: Tue, 28 Jul 2026 13:25:49 +0000 Subject: [PATCH 05/68] vulkan: enable the f16 KV contiguize pass by default (GGML_VK_FA_KV_CONTIG=0 opts out) Flip 442d7df from opt-in to default-on, matching the quant dequant-once path (GGML_VK_FA_DEQUANT) convention. The pass still self-gates: f16 K/V only, prefill only (neq1 >= 64), only when rows are actually strided, dense allocation, and the shared scratch-capacity check. Validated on Strix Halo (RADV gfx1151): FLASH_ATTN_EXT suite green with default env (dense-permuted cases exercise the pass) and with the opt-out; model-level pp2048@d8192 with no FA env matches the explicit GGML_VK_FA_KV_CONTIG=1 validation run (846.6 vs 847.9 t/s). Co-Authored-By: Claude Fable 5 Assisted-by: Claude (Opus 5) --- ggml/src/ggml-vulkan/ggml-vulkan.cpp | 9 +++++---- 1 file changed, 5 insertions(+), 4 deletions(-) diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp index 5f353e5630d..6d24357bfb4 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp +++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp @@ -11349,13 +11349,14 @@ static void ggml_vk_flash_attn(ggml_backend_vk_context * ctx, vk_context& subctx static const char * fa_dequant_env = getenv("GGML_VK_FA_DEQUANT"); const bool fa_dequant_off = fa_dequant_env && fa_dequant_env[0] == '0'; const bool fa_dequant_on = fa_dequant_env && fa_dequant_env[0] == '1'; - // EXPERIMENT (GGML_VK_FA_KV_CONTIG=1): run the same contiguize pass for f16 K/V. The - // KV-cache view reaching FA is head-interleaved ([HS, NH, KV] physically), and the cm1 + // Contiguize pass for f16 K/V (GGML_VK_FA_KV_CONTIG=0 opts out). The KV-cache view + // reaching FA is head-interleaved ([HS, NH, KV] physically), and the cm1 // direct-from-global coopMatLoads run ~2-5x slower on those strided rows than on // per-head-contiguous K/V. Copy K/V once into the scratch instead (dequant_f16_transpose - // is a pure strided copy). Only engages when the rows are actually strided. + // is a pure strided copy). Only engages when the rows are actually strided, and shares + // the quant path's prefill/allocation/scratch-capacity gates below. static const char * fa_kv_contig_env = getenv("GGML_VK_FA_KV_CONTIG"); - const bool fa_kv_contig = fa_kv_contig_env && fa_kv_contig_env[0] == '1'; + const bool fa_kv_contig = !(fa_kv_contig_env && fa_kv_contig_env[0] == '0'); const bool kv_f16_strided = k->type == GGML_TYPE_F16 && v->type == GGML_TYPE_F16 && (k->nb[1] != (uint64_t)HSK * sizeof(ggml_fp16_t) || v->nb[1] != (uint64_t)HSV * sizeof(ggml_fp16_t)) && From 82c3fe972128ff3fc04225065957359eea9c5ec3 Mon Sep 17 00:00:00 2001 From: Nathan Wilson Date: Tue, 28 Jul 2026 23:40:27 +0000 Subject: [PATCH 06/68] vulkan: single source of truth for native FA K/V types + non-native hard-gate Rebased adaptation of the iq4_nl routing fix: upstream 8161641 made iq4_nl a native FA type, so the original motivation (iq4_nl had no native shader and silently read garbage outside the dequant-once path) no longer applies to any currently-admitted type. Keep the machinery as hardening: ggml_vk_fa_kv_native() is the one list, supports_op mirrors every hard condition of the dispatch-time dequant gate for any future non-native type, and dispatch asserts the invariant instead of falling back to a garbage-reading shader. Native list synced with upstream (iq4_nl in, q1_0 out to match current admission). Co-Authored-By: Claude Fable 5 Assisted-by: Claude (Opus 5) --- ggml/src/ggml-vulkan/ggml-vulkan.cpp | 80 +++++++++++++++++++++------- 1 file changed, 60 insertions(+), 20 deletions(-) diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp index 6d24357bfb4..752d1f3e9d3 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp +++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp @@ -11207,6 +11207,27 @@ static bool ggml_vk_flash_attn_coopmat_shmem_support(const vk_device& device, co return supported; } +// K/V types the FA shaders can read directly (scalar/coopmat1 select the dequant code via the +// FaTypeK/FaTypeV spec constants; a type outside this list silently reads garbage). Types that +// are FA-supported but not listed here (iq4_nl) are only correct through the dequant-once +// scratch path, so supports_op and the dispatch-time gate must agree on when that path runs. +static bool ggml_vk_fa_kv_native(ggml_type t, bool coopmat2) { + switch (t) { + case GGML_TYPE_F32: + case GGML_TYPE_F16: + case GGML_TYPE_BF16: + case GGML_TYPE_Q4_0: + case GGML_TYPE_Q4_1: + case GGML_TYPE_Q5_0: + case GGML_TYPE_Q5_1: + case GGML_TYPE_Q8_0: + case GGML_TYPE_IQ4_NL: // native FA support since upstream 8161641 + return true; + default: + return false; + } +} + // Capacity gate for the dequant-once f16 K/V scratch. On discrete devices the scratch can push the // working set past free VRAM, and the driver then silently pages device-local memory (~15x prefill // regression measured on an 8 GB card at long context). UMA has no separate pool to overflow. @@ -11361,17 +11382,27 @@ static void ggml_vk_flash_attn(ggml_backend_vk_context * ctx, vk_context& subctx (k->nb[1] != (uint64_t)HSK * sizeof(ggml_fp16_t) || v->nb[1] != (uint64_t)HSV * sizeof(ggml_fp16_t)) && (HSK % 8) == 0 && (HSV % 8) == 0; - const bool use_dequant_kv = !fa_dequant_off && ((k_quant && v_quant) || (fa_kv_contig && kv_f16_strided)) && neq1 >= 64 && + // A K/V type the FA shaders cannot read directly (iq4_nl) is only correct through the + // dequant path. supports_op only admits such types when the hard conditions below hold, + // and the VRAM heuristic must not veto them (there is no fallback), so force the path on. + const bool kv_needs_dequant = !ggml_vk_fa_kv_native(k->type, ctx->device->coopmat2) || + !ggml_vk_fa_kv_native(v->type, ctx->device->coopmat2); + const bool use_dequant_kv = !fa_dequant_off && + ((k_quant && v_quant) || kv_needs_dequant || (fa_kv_contig && kv_f16_strided)) && neq1 >= 64 && is_dense_kv_cache(k) && is_dense_kv_cache(v) && kv_f16_sz <= ctx->device->properties.limits.maxStorageBufferRange && ctx->device->pipeline_dequant_transpose[k->type] != nullptr && ctx->device->pipeline_dequant_transpose[v->type] != nullptr && - // coopmat2 path does not benefit from the f16 scratch - !ctx->device->coopmat2 && + // coopmat2 reads its native types directly; non-native still needs the scratch + (kv_needs_dequant || !ctx->device->coopmat2) && // Intel Xe1 regresses, see PR 25494 - (ctx->device->vendor_id != VK_VENDOR_ID_INTEL || + (kv_needs_dequant || + ctx->device->vendor_id != VK_VENDOR_ID_INTEL || (ctx->device->coopmat_support && ctx->device->architecture != vk_device_architecture::INTEL_XE1)) && - (fa_dequant_on || ggml_vk_fa_dequant_scratch_fits(ctx, kv_f16_sz)); + (fa_dequant_on || kv_needs_dequant || ggml_vk_fa_dequant_scratch_fits(ctx, kv_f16_sz)); + // If this fires, supports_op admitted a non-native K/V type the gate then rejected; the + // native shader would return garbage rather than fail, so abort instead. + GGML_ASSERT(use_dequant_kv || !kv_needs_dequant); const ggml_type k_type_eff = use_dequant_kv ? GGML_TYPE_F16 : k->type; const ggml_type v_type_eff = use_dequant_kv ? GGML_TYPE_F16 : v->type; @@ -19120,25 +19151,34 @@ static bool ggml_backend_vk_device_supports_op(ggml_backend_dev_t dev, const ggm if (op->src[3] && op->src[3]->type != GGML_TYPE_F16) { return false; } - auto fa_kv_ok = [](ggml_type t) { - switch (t) { - case GGML_TYPE_F32: - case GGML_TYPE_F16: - case GGML_TYPE_BF16: - case GGML_TYPE_Q8_0: - case GGML_TYPE_Q5_1: - case GGML_TYPE_Q5_0: - case GGML_TYPE_Q4_1: - case GGML_TYPE_Q4_0: - case GGML_TYPE_IQ4_NL: - return true; - default: - return false; - } + auto fa_kv_ok = [&](ggml_type t) { + // ggml_vk_fa_kv_native is the single source of truth; on this base every + // admitted type is native, and the non-native hard-gate below is dormant + // hardening for any future type routed through the dequant-once scratch. + return ggml_vk_fa_kv_native(t, coopmat2); }; if (!fa_kv_ok(op->src[1]->type) || !fa_kv_ok(op->src[2]->type)) { return false; } + if (!ggml_vk_fa_kv_native(op->src[1]->type, coopmat2) || !ggml_vk_fa_kv_native(op->src[2]->type, coopmat2)) { + // Only correct through the dequant-once scratch path; admit only when every + // hard condition of the dispatch-time gate holds, so dispatch can never fall + // back to the native shader (it reads garbage for these types, not an error). + const ggml_tensor * k = op->src[1]; + const ggml_tensor * v = op->src[2]; + static const char * fa_dequant_env = getenv("GGML_VK_FA_DEQUANT"); + const bool fa_dequant_off = fa_dequant_env && fa_dequant_env[0] == '0'; + const uint64_t kv_f16_sz = ((uint64_t)ggml_nelements(k) + (uint64_t)ggml_nelements(v)) * sizeof(ggml_fp16_t); + if (fa_dequant_off || + op->src[0]->ne[1] < 64 || + device->pipeline_dequant_transpose[k->type] == nullptr || + device->pipeline_dequant_transpose[v->type] == nullptr || + k->nb[0] != ggml_type_size(k->type) || v->nb[0] != ggml_type_size(v->type) || + !ggml_is_contiguously_allocated(k) || !ggml_is_contiguously_allocated(v) || + kv_f16_sz > device->properties.limits.maxStorageBufferRange) { + return false; + } + } if ((op->src[1]->type == GGML_TYPE_BF16) != (op->src[2]->type == GGML_TYPE_BF16)) { return false; } From fa117d1d055f20467a87844368b7a62f050ebf0b Mon Sep 17 00:00:00 2001 From: Nathan Wilson Date: Tue, 28 Jul 2026 23:51:32 +0000 Subject: [PATCH 07/68] tests: dense-permuted iq4_nl FA cases (dequant-once route coverage) iq4_nl has no native Vulkan FA shader; the dequant-once scratch path is its only route (1abdd92). The sweep's iq4_nl cases are all sparse-view, which that path correctly rejects, so iq4_nl had zero passing FA coverage. Add model-layout (dense [0,2,1,3]-permuted) cases at prefill batch size: iq4_nl/iq4_nl with sinks off and on, hd72+GQA with sinks, and mixed K=iq4_nl/V=f16. Validated on RADV gfx1151 at 146fb73: FLASH_ATTN_EXT 4765/4765 with default env and with GGML_VK_FA_KV_CONTIG=0/1; 4761/4761 with GGML_VK_FA_DEQUANT=0 (new cases correctly report unsupported); full test-backend-ops suite 15538/15538. Co-Authored-By: Claude Fable 5 Assisted-by: Claude (Opus 5) --- tests/test-backend-ops.cpp | 7 +++++++ 1 file changed, 7 insertions(+) diff --git a/tests/test-backend-ops.cpp b/tests/test-backend-ops.cpp index 606ecc53b95..5b7893d4512 100644 --- a/tests/test-backend-ops.cpp +++ b/tests/test-backend-ops.cpp @@ -10992,6 +10992,13 @@ static std::vector> make_test_cases_eval() { test_cases.emplace_back(new test_flash_attn_ext(64, 64, 4, {8, 1}, 512, 75, true, false, 8.0f, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16, {0, 2, 1, 3}, false)); test_cases.emplace_back(new test_flash_attn_ext(128, 128, 4, {8, 1}, 512, 96, true, false, 0, 30.0f, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16, {0, 2, 1, 3}, false)); + // dense-permuted iq4_nl K/V at prefill batch sizes: iq4_nl has no native FA shader, so these + // exercise the only supported route (the dequant-once path), incl. sinks and mixed-with-f16 + test_cases.emplace_back(new test_flash_attn_ext(64, 64, 4, {1, 1}, 512, 75, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_IQ4_NL, GGML_TYPE_IQ4_NL, {0, 2, 1, 3}, false)); + test_cases.emplace_back(new test_flash_attn_ext(64, 64, 4, {1, 1}, 512, 75, true, true, 0, 0, GGML_PREC_F32, GGML_TYPE_IQ4_NL, GGML_TYPE_IQ4_NL, {0, 2, 1, 3}, false)); + test_cases.emplace_back(new test_flash_attn_ext(72, 72, 4, {4, 1}, 113, 75, true, true, 0, 0, GGML_PREC_F32, GGML_TYPE_IQ4_NL, GGML_TYPE_IQ4_NL, {0, 2, 1, 3}, false)); + test_cases.emplace_back(new test_flash_attn_ext(64, 64, 4, {1, 1}, 512, 75, true, true, 0, 0, GGML_PREC_F32, GGML_TYPE_IQ4_NL, GGML_TYPE_F16, {0, 2, 1, 3}, false)); + // mixed quant and Q1_0 test cases test_cases.emplace_back(new test_flash_attn_ext(64, 64, 4, {1, 1}, 128, 2, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_Q8_0, GGML_TYPE_Q4_0)); test_cases.emplace_back(new test_flash_attn_ext(64, 64, 4, {1, 1}, 128, 2, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_Q4_0, GGML_TYPE_F16)); From 5d5759ac92f12a117987814b12a8533df54860e1 Mon Sep 17 00:00:00 2001 From: Nathan Wilson Date: Fri, 31 Jul 2026 01:25:46 +0000 Subject: [PATCH 08/68] vulkan: bill the FA K/V contiguize pass on its own perf-logger line The contiguize/dequant pass is dispatched inside the FLASH_ATTN_EXT node handler, and the perf logger writes one timestamp per graph node, so its cost was charged to FLASH_ATTN_EXT with no way to separate the two. That made the copy invisible and the kernel look correspondingly slower. Adds a sub-node timestamp: a handler can close an interval mid-node and have it logged under its own name. Measured on Coder-30B UD-Q6_K_XL, pp2048/ub2048 at d32768, f16 KV: graph total 4612.3 ms FLASH_ATTN_EXT 3490.6 ms 75.68% of graph FA_KV_CONTIGUIZE 33.1 ms 0.72% of graph (0.95% of FA) so the copy is under 1% of the graph and FA is 75.7% of it at that depth. Bench throughput is unchanged with the instrumentation compiled in (441.09 vs 442.14 t/s), since the marks are only emitted when the logger is on in per-op mode. Two latent bugs in the query-pool handling fall out of this and are fixed here: - The pool is created with n_nodes+100 slots but only the first n_nodes+1 were reset each graph, so anything using the headroom would read stale results. - The results buffer was sized n_nodes+1 while getQueryPoolResults was asked for query_idx entries. Equal today, but it is an overflow waiting for the first caller that writes an extra timestamp. Sub-op intervals log no flops, since they move bytes rather than doing math; attributing the node's flop count to them would corrupt the GFLOPS column for both halves. Assisted-by: Claude (Opus 5) --- ggml/src/ggml-vulkan/ggml-vulkan.cpp | 50 +++++++++++++++++++++++++--- 1 file changed, 46 insertions(+), 4 deletions(-) diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp index 752d1f3e9d3..990a90729d5 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp +++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp @@ -2383,6 +2383,13 @@ class vk_perf_logger { timings[name].push_back(time); } + // Log a sub-node interval under a caller-supplied name. Used for work a node's handler + // dispatches before the op itself (e.g. the FA K/V contiguize/dequant pass), which would + // otherwise be billed to the op and invisible. No flops: these move bytes, not math. + void log_timing_named(const char *name, uint64_t time) { + timings[std::string(name)].push_back(time); + } + void log_timing(const std::vector &nodes, const std::vector &names, uint64_t time) { uint64_t total_flops = 0; std::string name; @@ -2479,6 +2486,9 @@ struct ggml_backend_vk_context { std::vector query_fusion_node_count; std::vector query_nodes; std::vector query_node_idx; + // non-null => this query slot closes a sub-node interval logged under this literal name, + // not a graph node. See ggml_vk_perf_mark_subop. + std::vector query_sub_names; int32_t num_queries {}; int32_t query_idx {}; }; @@ -11291,6 +11301,22 @@ static bool ggml_vk_fa_dequant_scratch_fits(ggml_backend_vk_context * ctx, uint6 return fits; } +// Close a timestamp interval mid-node so work a handler dispatches before its op is billed +// separately instead of being folded into the op's own time. `name` must be a string literal +// (stored by pointer). No-op unless the perf logger is on in per-op mode. +static void ggml_vk_perf_mark_subop(ggml_backend_vk_context * ctx, vk_context& subctx, const char * name) { + if (!vk_perf_logger_enabled || vk_perf_logger_concurrent || ctx->query_pool == VK_NULL_HANDLE) { + return; + } + if (ctx->query_idx >= (int)ctx->num_queries) { + return; // pool headroom exhausted; drop the mark rather than overflow + } + ctx->query_nodes[ctx->query_idx] = nullptr; + ctx->query_fusion_names[ctx->query_idx] = nullptr; + ctx->query_sub_names[ctx->query_idx] = name; + subctx->s->buffer->buf.writeTimestamp(vk::PipelineStageFlagBits::eAllCommands, ctx->query_pool, ctx->query_idx++); +} + static void ggml_vk_flash_attn(ggml_backend_vk_context * ctx, vk_context& subctx, const ggml_tensor * q, const ggml_tensor * k, const ggml_tensor * v, const ggml_tensor * mask, const ggml_tensor * sinks, ggml_tensor * dst) { VK_LOG_DEBUG("ggml_vk_flash_attn((" << q << ", name=" << q->name << ", type=" << q->type << ", ne0=" << q->ne[0] << ", ne1=" << q->ne[1] << ", ne2=" << q->ne[2] << ", ne3=" << q->ne[3] << ", nb0=" << q->nb[0] << ", nb1=" << q->nb[1] << ", nb2=" << q->nb[2] << ", nb3=" << q->nb[3]; std::cerr << "), (" << k << ", name=" << k->name << ", type=" << k->type << ", ne0=" << k->ne[0] << ", ne1=" << k->ne[1] << ", ne2=" << k->ne[2] << ", ne3=" << k->ne[3] << ", nb0=" << k->nb[0] << ", nb1=" << k->nb[1] << ", nb2=" << k->nb[2] << ", nb3=" << k->nb[3]; @@ -11602,6 +11628,11 @@ static void ggml_vk_flash_attn(ggml_backend_vk_context * ctx, vk_context& subctx ggml_vk_sync_buffers(ctx, subctx); k_buf = k_dst; v_buf = v_dst; + // Bill the K/V contiguize/dequant pass on its own line. Without this it is charged to + // FLASH_ATTN_EXT, which makes the copy invisible and the kernel look slower than it is. + ggml_vk_perf_mark_subop(ctx, subctx, kv_needs_dequant || (k_quant && v_quant) + ? "FA_KV_DEQUANT (sub-op)" + : "FA_KV_CONTIGUIZE (sub-op)"); } uint32_t mask_n_head_log2 = ((sinks != nullptr) << 24) | n_head_log2; @@ -18014,9 +18045,11 @@ static ggml_status ggml_backend_vk_graph_compute(ggml_backend_t backend, ggml_cg ctx->query_fusion_node_count.resize(ctx->num_queries); ctx->query_nodes.resize(ctx->num_queries); ctx->query_node_idx.resize(ctx->num_queries); + ctx->query_sub_names.resize(ctx->num_queries); } - ctx->device->device.resetQueryPool(ctx->query_pool, 0, cgraph->n_nodes+1); + // Reset the whole pool, not just n_nodes+1: sub-op marks consume slots past that. + ctx->device->device.resetQueryPool(ctx->query_pool, 0, ctx->num_queries); std::fill(ctx->query_fusion_names.begin(), ctx->query_fusion_names.end(), nullptr); std::fill(ctx->query_fusion_node_count.begin(), ctx->query_fusion_node_count.end(), 0); std::fill(ctx->query_nodes.begin(), ctx->query_nodes.end(), nullptr); @@ -18352,6 +18385,7 @@ static ggml_status ggml_backend_vk_graph_compute(ggml_backend_t backend, ggml_cg // track a single node/fusion for the current query ctx->query_nodes[ctx->query_idx] = cgraph->nodes[i]; ctx->query_fusion_names[ctx->query_idx] = fusion_string; + ctx->query_sub_names[ctx->query_idx] = nullptr; compute_ctx->s->buffer->buf.writeTimestamp(vk::PipelineStageFlagBits::eAllCommands, ctx->query_pool, ctx->query_idx++); ggml_vk_sync_buffers(ctx, compute_ctx); } else { @@ -18393,14 +18427,22 @@ static ggml_status ggml_backend_vk_graph_compute(ggml_backend_t backend, ggml_cg ctx->compute_ctx.reset(); // Get the results and pass them to the logger - std::vector timestamps(cgraph->n_nodes + 1); - VK_CHECK(ctx->device->device.getQueryPoolResults(ctx->query_pool, 0, ctx->query_idx, (cgraph->n_nodes + 1)*sizeof(uint64_t), timestamps.data(), sizeof(uint64_t), vk::QueryResultFlagBits::e64 | vk::QueryResultFlagBits::eWait), "get timestamp results", ctx->device); + // Sized to the pool, not n_nodes+1: sub-op marks push query_idx past the node count. + std::vector timestamps(ctx->num_queries); + VK_CHECK(ctx->device->device.getQueryPoolResults(ctx->query_pool, 0, ctx->query_idx, ctx->num_queries*sizeof(uint64_t), timestamps.data(), sizeof(uint64_t), vk::QueryResultFlagBits::e64 | vk::QueryResultFlagBits::eWait), "get timestamp results", ctx->device); if (!vk_perf_logger_concurrent) { // Log each op separately for (int i = 1; i < ctx->query_idx; i++) { + const uint64_t dt = uint64_t((timestamps[i] - timestamps[i-1]) * ctx->device->properties.limits.timestampPeriod); + if (ctx->query_sub_names[i] != nullptr) { + // sub-node interval (e.g. the FA K/V contiguize pass) - billed separately so + // it is not silently folded into the op that dispatched it + ctx->perf_logger->log_timing_named(ctx->query_sub_names[i], dt); + continue; + } auto node = ctx->query_nodes[i]; auto name = ctx->query_fusion_names[i]; - ctx->perf_logger->log_timing(node, name, uint64_t((timestamps[i] - timestamps[i-1]) * ctx->device->properties.limits.timestampPeriod)); + ctx->perf_logger->log_timing(node, name, dt); } } else { // Log each group of nodes From 8f178b8477543fe2cebed12057d9387b4bfcdd5f Mon Sep 17 00:00:00 2001 From: Nathan Wilson Date: Sun, 9 Aug 2026 10:43:34 +0000 Subject: [PATCH 09/68] vulkan: scale the FA MMQ dot product in fp32 before narrowing acc is an int32 sum of dotPacked4x8EXT results. With q8_0 both operands are full int8, so its bound is d_per_step*4*127*127, which overflows f16 when ACC_TYPE is f16 (GGML_PREC_DEFAULT): the score goes +inf, softmax is destroyed and the output comes back -FLT_MAX. Nibble types bound at ~30480 and stay in range, which is why only q8_0 tripped it. Apply the scales in fp32, then narrow. Identical arithmetic when ACC_TYPE is float, so the f32acc path is untouched. This is an upstream bug, not a fork regression, and it was fixed here once before - 61e77f4 carried it and the rebase onto b10133 dropped it. Only the scalar shader has MMQ, so the failing shape is hsk=128 + q8_0 K + prec=def + nb=1 + nr23=[1,1]; GQA>1 and nb>1 route to coopmat1 and escape. test-backend-ops -o FLASH_ATTN_EXT on gfx1151, quiet box: 13295/13295 twice, up from 13257/13295. All 38 failures were type_K=q8_0 prec=def. Both ggml_flash_attn_ext call sites in the tree force GGML_PREC_F32, so no model path reaches this - it is a landmine for the next person touching precision, not a live bug. Co-Authored-By: Claude Opus 5 Assisted-by: Claude (Opus 5) --- ggml/src/ggml-vulkan/vulkan-shaders/flash_attn.comp | 6 +++++- 1 file changed, 5 insertions(+), 1 deletion(-) diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn.comp b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn.comp index 0c1b6d0673e..18a37add9bf 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn.comp @@ -420,7 +420,11 @@ void main() { acc += dotPacked4x8EXT(Qf[qib].qs[qiqs + d], k_quants[d]); } - Sf[r][c] += ACC_TYPE(acc) * ACC_TYPE(Qf[qib].ds.x) * k_dm.x; + // scale in fp32 before narrowing: acc is an int32 sum of dotPacked4x8EXT + // results, bounded by d_per_step*4*127*127 with q8_0 on both sides, which + // overflows f16 when ACC_TYPE is f16 (GGML_PREC_DEFAULT). Identical + // arithmetic when ACC_TYPE is float. + Sf[r][c] += ACC_TYPE(float(acc) * float(Qf[qib].ds.x) * float(k_dm.x)); if ((d_tid * (HSK_per_thread / 4) + d_block) % 8 == 0) { Sf[r][c] += k_dot_correction(qib, k_dm); } From e72ffec16eb59e9b024288c9f669963effb096f2 Mon Sep 17 00:00:00 2001 From: Gaetan Puleo <12990773+gaetan-puleo@users.noreply.github.com> Date: Sat, 1 Aug 2026 14:57:00 +0200 Subject: [PATCH 10/68] vulkan: DeepSeek V4 lightning indexer kernels + indexed sparse FA Implements GGML_OP_LIGHTNING_INDEXER on Vulkan (scalar subgroup shader for small batches, coopmat 16x16 tiles for prefill, dedicated decode variant) and an indexed sparse flash-attention path that consumes the indexer's top-k selection directly via a new ggml_flash_attn_ext_add_top_k() API (FA src[5] + op_param[4] = n_kv_raw dense prefix), instead of attending densely over the full compressed KV. The sparse path engages only for V4's CSA shape (hd 512, 64 heads, MQA, f16 K==V latent) when dense_kv >= 3x active_kv; everything else falls through to the dense path, which stays correct because the kq_mask still carries the top-k selection. Dropped from the original: the mul_mat_id tokens-per-expert pipeline selection, which duplicates GGML_VK_MMID_SMALLN already on this branch. Originally by Gaetan Puleo (llama-cpp-nathan-toolbox-deepseek-v4-poc, branch deepseek-v4-flash-strix-halo); cherry-picked with the mmid hunk dropped. --- ggml/include/ggml.h | 5 + ggml/src/ggml-vulkan/ggml-vulkan.cpp | 170 ++++++++++++++++++ .../vulkan-shaders/flash_attn_top_k.comp | 144 +++++++++++++++ .../vulkan-shaders/lightning_indexer_cm.comp | 125 +++++++++++++ .../lightning_indexer_decode_cm.comp | 110 ++++++++++++ .../lightning_indexer_scalar64.comp | 92 ++++++++++ .../vulkan-shaders/vulkan-shaders-gen.cpp | 7 + ggml/src/ggml.c | 14 ++ src/llama-graph.cpp | 7 +- src/llama-graph.h | 4 +- src/models/deepseek4.cpp | 2 +- tests/test-backend-ops.cpp | 10 +- 12 files changed, 686 insertions(+), 4 deletions(-) create mode 100644 ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_top_k.comp create mode 100644 ggml/src/ggml-vulkan/vulkan-shaders/lightning_indexer_cm.comp create mode 100644 ggml/src/ggml-vulkan/vulkan-shaders/lightning_indexer_decode_cm.comp create mode 100644 ggml/src/ggml-vulkan/vulkan-shaders/lightning_indexer_scalar64.comp diff --git a/ggml/include/ggml.h b/ggml/include/ggml.h index f0eda79a1ae..2ff12d00295 100644 --- a/ggml/include/ggml.h +++ b/ggml/include/ggml.h @@ -2486,6 +2486,11 @@ extern "C" { struct ggml_tensor * a, struct ggml_tensor * sinks); + GGML_API void ggml_flash_attn_ext_add_top_k( + struct ggml_tensor * a, + struct ggml_tensor * top_k, + int64_t n_kv_raw); + // TODO: needs to be adapted to ggml_flash_attn_ext GGML_API struct ggml_tensor * ggml_flash_attn_back( struct ggml_context * ctx, diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp index 990a90729d5..44efa5c4897 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp +++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp @@ -1104,6 +1104,10 @@ struct vk_device_struct { vk_pipeline pipeline_lightning_indexer_f32[GGML_TYPE_COUNT]; // [size_idx][kda] where size_idx: 0=d16, 1=d32, 2=d64, 3=d128 vk_pipeline pipeline_gated_delta_net[4][2]; + vk_pipeline pipeline_lightning_indexer_f16; + vk_pipeline pipeline_lightning_indexer_cm_f16; + vk_pipeline pipeline_lightning_indexer_decode_cm_f16; + vk_pipeline pipeline_flash_attn_top_k_f16; vk_pipeline pipeline_ssm_scan_f32_d128; vk_pipeline pipeline_ssm_scan_f32_d256; vk_pipeline pipeline_ssm_conv_f32; @@ -1933,6 +1937,29 @@ struct vk_op_gated_delta_net_push_constants { uint32_t K; }; +// push constants for the fork's wave64 f16 lightning-indexer kernels (scalar-64 + CM family) +struct vk_op_lightning_indexer_cm_push_constants { + uint32_t n_kv, n_batch, n_stream, nem3; + uint32_t nb1, nb3; + uint32_t nbq1, nbq2, nbq3; + uint32_t nbk2, nbk3; + uint32_t nbw1, nbw3; + uint32_t nbm1, nbm3; +}; +static_assert(sizeof(vk_op_lightning_indexer_cm_push_constants) <= 128); + +struct vk_op_flash_attn_top_k_push_constants { + uint32_t n_batch, n_kv, n_kv_raw, n_top_k, n_head; + uint32_t nbq1, nbq2, nbq3; + uint32_t nbk1, nbk3; + uint32_t nbm1, nbm3; + uint32_t nbt1, nbt3; + uint32_t nb1, nb2, nb3; + float scale; + uint32_t has_sinks; +}; +static_assert(sizeof(vk_op_flash_attn_top_k_push_constants) <= 128); + struct vk_op_ssm_scan_push_constants { uint32_t nb02, nb03, nb12, nb13; uint32_t nb21, nb22, nb31; @@ -6205,6 +6232,29 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { } } + if (device->subgroup_arithmetic && device->subgroup_size == 64) { + ggml_vk_create_pipeline(device, device->pipeline_lightning_indexer_f16, + "lightning_indexer_f16", lightning_indexer_f16_len, lightning_indexer_f16_data, "main", 5, + sizeof(vk_op_lightning_indexer_cm_push_constants), {8, 1, 1}, {device->subgroup_size}, 1, true, true, + device->subgroup_size); +#if defined(VK_KHR_cooperative_matrix) && defined(GGML_VULKAN_COOPMAT_GLSLC_SUPPORT) + if (device->coopmat_support && device->coopmat_support_16x16x16_f32acc && device->subgroup_size_control) { + ggml_vk_create_pipeline(device, device->pipeline_lightning_indexer_cm_f16, + "lightning_indexer_cm_f16", lightning_indexer_cm_f16_len, lightning_indexer_cm_f16_data, "main", 5, + sizeof(vk_op_lightning_indexer_cm_push_constants), {16, 16, 1}, {device->subgroup_size}, 1, true, true, + device->subgroup_size); + ggml_vk_create_pipeline(device, device->pipeline_lightning_indexer_decode_cm_f16, + "lightning_indexer_decode_cm_f16", lightning_indexer_decode_cm_f16_len, lightning_indexer_decode_cm_f16_data, "main", 5, + sizeof(vk_op_lightning_indexer_cm_push_constants), {16, 1, 1}, {device->subgroup_size}, 1, true, true, + device->subgroup_size); + } +#endif + ggml_vk_create_pipeline(device, device->pipeline_flash_attn_top_k_f16, + "flash_attn_top_k_f16", flash_attn_top_k_f16_len, flash_attn_top_k_f16_data, "main", 6, + sizeof(vk_op_flash_attn_top_k_push_constants), {1, 1, 1}, {512, device->subgroup_size}, 1, true, true, + device->subgroup_size); + } + if (device->subgroup_arithmetic && device->subgroup_require_full_support) { ggml_vk_create_pipeline(device, device->pipeline_ssm_scan_f32_d128, "ssm_scan_128_f32", ssm_scan_subgroup_f32_len, ssm_scan_subgroup_f32_data, "main", 8, sizeof(vk_op_ssm_scan_push_constants), {1, 1, 1}, {128, device->subgroup_size}, 1, true, true); ggml_vk_create_pipeline(device, device->pipeline_ssm_scan_f32_d256, "ssm_scan_256_f32", ssm_scan_subgroup_f32_len, ssm_scan_subgroup_f32_data, "main", 8, sizeof(vk_op_ssm_scan_push_constants), {1, 1, 1}, {256, device->subgroup_size}, 1, true, true); @@ -11317,6 +11367,69 @@ static void ggml_vk_perf_mark_subop(ggml_backend_vk_context * ctx, vk_context& s subctx->s->buffer->buf.writeTimestamp(vk::PipelineStageFlagBits::eAllCommands, ctx->query_pool, ctx->query_idx++); } +static bool ggml_vk_flash_attn_top_k(ggml_backend_vk_context * ctx, vk_context & subctx, + const ggml_tensor * q, const ggml_tensor * k, const ggml_tensor * v, + const ggml_tensor * mask, const ggml_tensor * sinks, ggml_tensor * dst) { + const ggml_tensor * top_k = dst->src[5]; + if (!top_k || !ctx->device->pipeline_flash_attn_top_k_f16 || + q->type != GGML_TYPE_F32 || k->type != GGML_TYPE_F16 || v->type != GGML_TYPE_F16 || + !mask || mask->type != GGML_TYPE_F16 || top_k->type != GGML_TYPE_I32 || + q->ne[0] != 512 || q->ne[1] < 64 || k->ne[0] != 512 || v->ne[0] != 512 || + q->ne[2] != 64 || k->ne[2] != 1 || v->ne[2] != 1 || + q->ne[1] != top_k->ne[1] || q->ne[3] != top_k->ne[3] || + k->ne[1] != v->ne[1] || k->buffer != v->buffer || k->data != v->data || + !ggml_is_contiguous(mask) || !ggml_is_contiguous(top_k)) { + return false; + } + + float scale = 0.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 (max_bias != 0.0f || logit_softcap != 0.0f) { + return false; + } + + const int32_t n_kv_raw = ggml_get_op_params_i32(dst, 4); + if (n_kv_raw < 0 || n_kv_raw > k->ne[1] || top_k->ne[0] > k->ne[1] - n_kv_raw) { + return false; + } + const int64_t n_kv_active = n_kv_raw + top_k->ne[0]; + if (k->ne[1] < 3 * n_kv_active) { + return false; + } + + const vk_op_flash_attn_top_k_push_constants pc = { + (uint32_t) q->ne[1], (uint32_t) k->ne[1], (uint32_t) n_kv_raw, + (uint32_t) top_k->ne[0], (uint32_t) q->ne[2], + (uint32_t) (q->nb[1] / sizeof(float)), + (uint32_t) (q->nb[2] / sizeof(float)), + (uint32_t) (q->nb[3] / sizeof(float)), + (uint32_t) (k->nb[1] / sizeof(ggml_fp16_t)), + (uint32_t) (k->nb[3] / sizeof(ggml_fp16_t)), + (uint32_t) (mask->nb[1] / sizeof(ggml_fp16_t)), + (uint32_t) (mask->nb[3] / sizeof(ggml_fp16_t)), + (uint32_t) (top_k->nb[1] / sizeof(int32_t)), + (uint32_t) (top_k->nb[3] / sizeof(int32_t)), + (uint32_t) (dst->nb[1] / sizeof(float)), + (uint32_t) (dst->nb[2] / sizeof(float)), + (uint32_t) (dst->nb[3] / sizeof(float)), + scale, sinks != nullptr, + }; + + const vk_subbuffer q_buf = ggml_vk_tensor_subbuffer(ctx, q); + const vk_subbuffer sinks_buf = sinks ? ggml_vk_tensor_subbuffer(ctx, sinks) : q_buf; + vk_pipeline pipeline = ctx->device->pipeline_flash_attn_top_k_f16; + ggml_pipeline_request_descriptor_sets(ctx, pipeline, 1); + ggml_vk_dispatch_pipeline(ctx, subctx, pipeline, + {q_buf, ggml_vk_tensor_subbuffer(ctx, k), ggml_vk_tensor_subbuffer(ctx, mask), sinks_buf, + ggml_vk_tensor_subbuffer(ctx, top_k), ggml_vk_tensor_subbuffer(ctx, dst)}, + pc, {(uint32_t) q->ne[1], (uint32_t) CEIL_DIV(q->ne[2], 8), (uint32_t) q->ne[3]}); + return true; +} + static void ggml_vk_flash_attn(ggml_backend_vk_context * ctx, vk_context& subctx, const ggml_tensor * q, const ggml_tensor * k, const ggml_tensor * v, const ggml_tensor * mask, const ggml_tensor * sinks, ggml_tensor * dst) { VK_LOG_DEBUG("ggml_vk_flash_attn((" << q << ", name=" << q->name << ", type=" << q->type << ", ne0=" << q->ne[0] << ", ne1=" << q->ne[1] << ", ne2=" << q->ne[2] << ", ne3=" << q->ne[3] << ", nb0=" << q->nb[0] << ", nb1=" << q->nb[1] << ", nb2=" << q->nb[2] << ", nb3=" << q->nb[3]; std::cerr << "), (" << k << ", name=" << k->name << ", type=" << k->type << ", ne0=" << k->ne[0] << ", ne1=" << k->ne[1] << ", ne2=" << k->ne[2] << ", ne3=" << k->ne[3] << ", nb0=" << k->nb[0] << ", nb1=" << k->nb[1] << ", nb2=" << k->nb[2] << ", nb3=" << k->nb[3]; @@ -11368,6 +11481,9 @@ static void ggml_vk_flash_attn(ggml_backend_vk_context * ctx, vk_context& subctx assert(dst->type == GGML_TYPE_F32); assert(q->type == GGML_TYPE_F32); + if (ggml_vk_flash_attn_top_k(ctx, subctx, q, k, v, mask, sinks, dst)) { + return; + } uint32_t gqa_ratio = 1; uint32_t qk_ratio = neq2 / nek2; uint32_t workgroups_x = (uint32_t)neq1; @@ -12275,6 +12391,17 @@ static vk_pipeline ggml_vk_op_get_pipeline(ggml_backend_vk_context * ctx, const } return nullptr; case GGML_OP_LIGHTNING_INDEXER: + // fork fast path: f16 K on wave64 subgroup-arithmetic devices routes to the tuned + // scalar-64/CM kernels; anything else falls through to the generic pipeline table + if (src0->type == GGML_TYPE_F32 && src1->type == GGML_TYPE_F16 && dst->type == GGML_TYPE_F32 && + src0->ne[0] == 128 && src0->ne[1] == 64 && src1->ne[1] == 1 && + ctx->device->pipeline_lightning_indexer_f16) { + if (ctx->device->pipeline_lightning_indexer_decode_cm_f16 && src0->ne[2] == 1) { + return ctx->device->pipeline_lightning_indexer_decode_cm_f16; + } + return ctx->device->pipeline_lightning_indexer_cm_f16 && src0->ne[2] >= 16 ? + ctx->device->pipeline_lightning_indexer_cm_f16 : ctx->device->pipeline_lightning_indexer_f16; + } // only the k type selects a pipeline, the other types are fixed by ggml_lightning_indexer() if (ggml_vk_lightning_indexer_k_type_supported(src1->type)) { return ctx->device->pipeline_lightning_indexer_f32[src1->type]; @@ -13355,6 +13482,8 @@ static void ggml_vk_gated_linear_attn(ggml_backend_vk_context * ctx, vk_context& pc, { (uint32_t)(n_seqs * n_heads), 1, 1 }); } +static void ggml_vk_lightning_indexer_cm(ggml_backend_vk_context * ctx, vk_context& subctx, ggml_tensor * dst, vk_pipeline pipeline); + static void ggml_vk_lightning_indexer(ggml_backend_vk_context * ctx, vk_context& subctx, ggml_tensor * dst) { const ggml_tensor * q = dst->src[0]; const ggml_tensor * k = dst->src[1]; @@ -13364,6 +13493,14 @@ static void ggml_vk_lightning_indexer(ggml_backend_vk_context * ctx, vk_context& vk_pipeline pipeline = ggml_vk_op_get_pipeline(ctx, q, k, w, dst, dst->op); GGML_ASSERT(pipeline != nullptr); + // the fork's wave64 f16 kernels take their own push-constant layout + if (pipeline == ctx->device->pipeline_lightning_indexer_f16 || + pipeline == ctx->device->pipeline_lightning_indexer_cm_f16 || + pipeline == ctx->device->pipeline_lightning_indexer_decode_cm_f16) { + ggml_vk_lightning_indexer_cm(ctx, subctx, dst, pipeline); + return; + } + ggml_pipeline_request_descriptor_sets(ctx, pipeline, 1); const uint32_t n_kv = k->ne[2]; @@ -13461,6 +13598,39 @@ static void ggml_vk_gated_delta_net(ggml_backend_vk_context * ctx, vk_context& s pc, { H, n_seqs, S_v }); } +static void ggml_vk_lightning_indexer_cm(ggml_backend_vk_context * ctx, vk_context& subctx, ggml_tensor * dst, vk_pipeline pipeline) { + const ggml_tensor * q = dst->src[0]; + const ggml_tensor * k = dst->src[1]; + const ggml_tensor * w = dst->src[2]; + const ggml_tensor * m = dst->src[3]; + + ggml_pipeline_request_descriptor_sets(ctx, pipeline, 1); + + const vk_op_lightning_indexer_cm_push_constants pc = { + (uint32_t) k->ne[2], + (uint32_t) q->ne[2], + (uint32_t) q->ne[3], + (uint32_t) m->ne[3], + (uint32_t) (dst->nb[1] / sizeof(float)), + (uint32_t) (dst->nb[3] / sizeof(float)), + (uint32_t) (q->nb[1] / sizeof(float)), + (uint32_t) (q->nb[2] / sizeof(float)), + (uint32_t) (q->nb[3] / sizeof(float)), + (uint32_t) (k->nb[2] / sizeof(ggml_fp16_t)), + (uint32_t) (k->nb[3] / sizeof(ggml_fp16_t)), + (uint32_t) (w->nb[1] / sizeof(float)), + (uint32_t) (w->nb[3] / sizeof(float)), + (uint32_t) (m->nb[1] / sizeof(ggml_fp16_t)), + (uint32_t) (m->nb[3] / sizeof(ggml_fp16_t)), + }; + + ggml_vk_dispatch_pipeline(ctx, subctx, pipeline, + {ggml_vk_tensor_subbuffer(ctx, q), ggml_vk_tensor_subbuffer(ctx, k), + ggml_vk_tensor_subbuffer(ctx, w), ggml_vk_tensor_subbuffer(ctx, m), + ggml_vk_tensor_subbuffer(ctx, dst)}, + pc, {(uint32_t) k->ne[2], (uint32_t) q->ne[2], (uint32_t) q->ne[3]}); +} + static void ggml_vk_ssm_scan(ggml_backend_vk_context * ctx, vk_context& subctx, ggml_tensor * dst) { const ggml_tensor * src0 = dst->src[0]; const ggml_tensor * src1 = dst->src[1]; diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_top_k.comp b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_top_k.comp new file mode 100644 index 00000000000..7f82266781f --- /dev/null +++ b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_top_k.comp @@ -0,0 +1,144 @@ +#version 450 + +#extension GL_EXT_control_flow_attributes : require +#extension GL_EXT_shader_16bit_storage : require +#extension GL_EXT_shader_explicit_arithmetic_types_float16 : require +#extension GL_KHR_shader_subgroup_arithmetic : require +#extension GL_KHR_shader_subgroup_basic : require + +layout(constant_id = 0) const uint WORKGROUP_SIZE = 512; +layout(constant_id = 1) const uint SUBGROUP_SIZE = 64; +layout(local_size_x_id = 0, local_size_y = 1, local_size_z = 1) in; + +layout(binding = 0) readonly buffer QBuf { float data_q[]; }; +layout(binding = 1) readonly buffer KBuf { float16_t data_k[]; }; +layout(binding = 2) readonly buffer MaskBuf { float16_t data_m[]; }; +layout(binding = 3) readonly buffer SinkBuf { float data_s[]; }; +layout(binding = 4) readonly buffer TopBuf { int data_top[]; }; +layout(binding = 5) writeonly buffer DstBuf { float data_dst[]; }; + +layout(push_constant) uniform Parameters { + uint n_batch; + uint n_kv; + uint n_kv_raw; + uint n_top_k; + uint n_head; + uint nbq1; + uint nbq2; + uint nbq3; + uint nbk1; + uint nbk3; + uint nbm1; + uint nbm3; + uint nbt1; + uint nbt3; + uint nb1; + uint nb2; + uint nb3; + float scale; + uint has_sinks; +} p; + +const uint HEAD_SIZE = 512; +const uint HEADS_PER_GROUP = 8; +const uint KEYS_PER_BLOCK = 16; + +shared float16_t key_sh[KEYS_PER_BLOCK * HEAD_SIZE]; +shared uint key_idx[KEYS_PER_BLOCK]; + +void main() { + const uint tid = gl_LocalInvocationIndex; + const uint lane = gl_SubgroupInvocationID; + const uint head = gl_WorkGroupID.y * HEADS_PER_GROUP + gl_SubgroupID; + const uint token = gl_WorkGroupID.x; + const uint stream = gl_WorkGroupID.z; + + if (token >= p.n_batch || head >= p.n_head) { + return; + } + + float accum[HEAD_SIZE / SUBGROUP_SIZE]; + [[unroll]] for (uint i = 0; i < HEAD_SIZE / SUBGROUP_SIZE; ++i) { + accum[i] = 0.0; + } + + float row_max = uintBitsToFloat(0xff800000); + float row_sum = 0.0; + const uint q_base = stream * p.nbq3 + head * p.nbq2 + token * p.nbq1; + const uint mask_base = stream * p.nbm3 + token * p.nbm1; + const uint top_base = stream * p.nbt3 + token * p.nbt1; + const uint total_keys = p.n_kv_raw + p.n_top_k; + + for (uint kb = 0; kb < total_keys; kb += KEYS_PER_BLOCK) { + if (tid < KEYS_PER_BLOCK) { + const uint selected = kb + tid; + uint key = p.n_kv; + if (selected < p.n_kv_raw) { + key = selected; + } else if (selected < total_keys) { + const int compressed = data_top[top_base + selected - p.n_kv_raw]; + if (compressed >= 0 && uint(compressed) < p.n_kv - p.n_kv_raw) { + key = p.n_kv_raw + uint(compressed); + } + } + key_idx[tid] = key; + } + barrier(); + + for (uint idx = tid; idx < KEYS_PER_BLOCK * HEAD_SIZE; idx += WORKGROUP_SIZE) { + const uint col = idx / HEAD_SIZE; + const uint dim = idx % HEAD_SIZE; + const uint key = key_idx[col]; + key_sh[idx] = key < p.n_kv ? data_k[stream * p.nbk3 + key * p.nbk1 + dim] : float16_t(0.0); + } + barrier(); + + [[unroll]] for (uint col = 0; col < KEYS_PER_BLOCK; ++col) { + const uint selected = kb + col; + const uint key = key_idx[col]; + if (selected >= total_keys || key >= p.n_kv) { + continue; + } + + float partial = 0.0; + [[unroll]] for (uint i = 0; i < HEAD_SIZE / SUBGROUP_SIZE; ++i) { + const uint dim = lane + i * SUBGROUP_SIZE; + partial += data_q[q_base + dim] * float(key_sh[col * HEAD_SIZE + dim]); + } + const float mask = float(data_m[mask_base + key]); + const float score = subgroupAdd(partial) * p.scale + mask; + if (mask < -65500.0) { + continue; + } + + const float new_max = max(row_max, score); + const float old_scale = row_sum == 0.0 ? 0.0 : exp(row_max - new_max); + const float value_scale = exp(score - new_max); + row_sum = row_sum * old_scale + value_scale; + row_max = new_max; + + [[unroll]] for (uint i = 0; i < HEAD_SIZE / SUBGROUP_SIZE; ++i) { + const uint dim = lane + i * SUBGROUP_SIZE; + accum[i] = accum[i] * old_scale + value_scale * float(key_sh[col * HEAD_SIZE + dim]); + } + } + barrier(); + } + + if (p.has_sinks != 0) { + const float sink = data_s[head]; + const float new_max = max(row_max, sink); + const float old_scale = row_sum == 0.0 ? 0.0 : exp(row_max - new_max); + const float sink_scale = exp(sink - new_max); + row_sum = row_sum * old_scale + sink_scale; + [[unroll]] for (uint i = 0; i < HEAD_SIZE / SUBGROUP_SIZE; ++i) { + accum[i] *= old_scale; + } + } + + const uint dst_base = stream * p.nb3 + token * p.nb2 + head * p.nb1; + const float inv_sum = row_sum == 0.0 ? 0.0 : 1.0 / row_sum; + [[unroll]] for (uint i = 0; i < HEAD_SIZE / SUBGROUP_SIZE; ++i) { + data_dst[dst_base + lane + i * SUBGROUP_SIZE] = accum[i] * inv_sum; + } +} diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/lightning_indexer_cm.comp b/ggml/src/ggml-vulkan/vulkan-shaders/lightning_indexer_cm.comp new file mode 100644 index 00000000000..a0a3639d254 --- /dev/null +++ b/ggml/src/ggml-vulkan/vulkan-shaders/lightning_indexer_cm.comp @@ -0,0 +1,125 @@ +#version 450 + +#extension GL_EXT_control_flow_attributes : require +#extension GL_EXT_shader_16bit_storage : require +#extension GL_EXT_shader_explicit_arithmetic_types_float16 : require +#extension GL_KHR_cooperative_matrix : require +#extension GL_KHR_memory_scope_semantics : require + +layout(constant_id = 0) const uint SUBGROUP_SIZE = 64; +layout(local_size_x_id = 0, local_size_y = 1, local_size_z = 1) in; + +layout(binding = 0) readonly buffer QBuf { float data_q[]; }; +layout(binding = 1) readonly buffer KBuf { float16_t data_k[]; }; +layout(binding = 2) readonly buffer WBuf { float data_w[]; }; +layout(binding = 3) readonly buffer MBuf { float16_t data_m[]; }; +layout(binding = 4) writeonly buffer DstBuf { float data_dst[]; }; + +layout(push_constant) uniform Parameters { + uint n_kv; + uint n_batch; + uint n_stream; + uint nem3; + uint nb1; + uint nb3; + uint nbq1; + uint nbq2; + uint nbq3; + uint nbk2; + uint nbk3; + uint nbw1; + uint nbw3; + uint nbm1; + uint nbm3; +} p; + +const uint TILE = 16; +const uint HEAD_SIZE = 128; +const uint N_HEAD = 64; +const uint VEC_PER_HEAD = HEAD_SIZE / 4; +const uint TILE_STRIDE = VEC_PER_HEAD + 2; +const uint SCORE_STRIDE = TILE / 4 + 1; + +shared f16vec4 q_sh[TILE * TILE_STRIDE]; +shared f16vec4 k_sh[TILE * TILE_STRIDE]; +shared vec4 score_sh[TILE * SCORE_STRIDE]; + +void main() { + const uint tid = gl_LocalInvocationIndex; + const uint kv_base = gl_WorkGroupID.x * TILE; + const uint token_base = gl_WorkGroupID.y * TILE; + const uint stream = gl_WorkGroupID.z; + + float totals[4]; + [[unroll]] for (uint i = 0; i < 4; ++i) { + totals[i] = 0.0; + } + + for (uint idx = tid; idx < TILE * VEC_PER_HEAD; idx += SUBGROUP_SIZE) { + const uint key = idx / VEC_PER_HEAD; + const uint d4 = idx % VEC_PER_HEAD; + const uint kv = kv_base + key; + f16vec4 value = f16vec4(0.0); + if (kv < p.n_kv) { + const uint offset = stream * p.nbk3 + kv * p.nbk2 + d4 * 4; + value = f16vec4(data_k[offset], data_k[offset + 1], data_k[offset + 2], data_k[offset + 3]); + } + k_sh[key * TILE_STRIDE + d4] = value; + } + barrier(); + + for (uint head = 0; head < N_HEAD; ++head) { + for (uint idx = tid; idx < TILE * VEC_PER_HEAD; idx += SUBGROUP_SIZE) { + const uint token_local = idx / VEC_PER_HEAD; + const uint d4 = idx % VEC_PER_HEAD; + const uint token = token_base + token_local; + f16vec4 value = f16vec4(0.0); + if (token < p.n_batch) { + const uint offset = stream * p.nbq3 + token * p.nbq2 + head * p.nbq1 + d4 * 4; + value = f16vec4(data_q[offset], data_q[offset + 1], data_q[offset + 2], data_q[offset + 3]); + } + q_sh[token_local * TILE_STRIDE + d4] = value; + } + barrier(); + + coopmat scores = + coopmat(0.0); + coopmat kmat; + coopmat qmat; + + [[unroll]] for (uint d = 0; d < HEAD_SIZE; d += TILE) { + coopMatLoad(kmat, k_sh, d / 4, TILE_STRIDE, gl_CooperativeMatrixLayoutRowMajor); + coopMatLoad(qmat, q_sh, d / 4, TILE_STRIDE, gl_CooperativeMatrixLayoutColumnMajor); + scores = coopMatMulAdd(kmat, qmat, scores); + } + + coopMatStore(scores, score_sh, 0, SCORE_STRIDE, gl_CooperativeMatrixLayoutRowMajor); + barrier(); + + [[unroll]] for (uint i = 0; i < 4; ++i) { + const uint idx = tid + i * SUBGROUP_SIZE; + const uint key = idx / TILE; + const uint token_local = idx % TILE; + const uint token = token_base + token_local; + if (token < p.n_batch && kv_base + key < p.n_kv) { + const float score = score_sh[key * SCORE_STRIDE + token_local / 4][token_local % 4]; + const float weight = data_w[stream * p.nbw3 + token * p.nbw1 + head]; + totals[i] += max(score, 0.0) * weight; + } + } + barrier(); + } + + [[unroll]] for (uint i = 0; i < 4; ++i) { + const uint idx = tid + i * SUBGROUP_SIZE; + const uint key = idx / TILE; + const uint token_local = idx % TILE; + const uint kv = kv_base + key; + const uint token = token_base + token_local; + if (kv < p.n_kv && token < p.n_batch) { + const uint mask_base = (stream % p.nem3) * p.nbm3 + token * p.nbm1; + const uint dst_base = stream * p.nb3 + token * p.nb1; + data_dst[dst_base + kv] = totals[i] + float(data_m[mask_base + kv]); + } + } +} diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/lightning_indexer_decode_cm.comp b/ggml/src/ggml-vulkan/vulkan-shaders/lightning_indexer_decode_cm.comp new file mode 100644 index 00000000000..fd555c76806 --- /dev/null +++ b/ggml/src/ggml-vulkan/vulkan-shaders/lightning_indexer_decode_cm.comp @@ -0,0 +1,110 @@ +#version 450 + +#extension GL_EXT_control_flow_attributes : require +#extension GL_EXT_shader_16bit_storage : require +#extension GL_EXT_shader_explicit_arithmetic_types_float16 : require +#extension GL_KHR_cooperative_matrix : require +#extension GL_KHR_memory_scope_semantics : require + +layout(constant_id = 0) const uint SUBGROUP_SIZE = 64; +layout(local_size_x_id = 0, local_size_y = 1, local_size_z = 1) in; + +layout(binding = 0) readonly buffer QBuf { float data_q[]; }; +layout(binding = 1) readonly buffer KBuf { float16_t data_k[]; }; +layout(binding = 2) readonly buffer WBuf { float data_w[]; }; +layout(binding = 3) readonly buffer MBuf { float16_t data_m[]; }; +layout(binding = 4) writeonly buffer DstBuf { float data_dst[]; }; + +layout(push_constant) uniform Parameters { + uint n_kv; + uint n_batch; + uint n_stream; + uint nem3; + uint nb1; + uint nb3; + uint nbq1; + uint nbq2; + uint nbq3; + uint nbk2; + uint nbk3; + uint nbw1; + uint nbw3; + uint nbm1; + uint nbm3; +} p; + +const uint TILE = 16; +const uint HEAD_SIZE = 128; +const uint N_HEAD = 64; +const uint VEC_PER_HEAD = HEAD_SIZE / 4; +const uint TILE_STRIDE = VEC_PER_HEAD + 2; +const uint SCORE_STRIDE = TILE / 4 + 1; + +shared f16vec4 q_sh[TILE * TILE_STRIDE]; +shared f16vec4 k_sh[TILE * TILE_STRIDE]; +shared vec4 score_sh[TILE * SCORE_STRIDE]; + +void main() { + const uint tid = gl_LocalInvocationIndex; + const uint kv_base = gl_WorkGroupID.x * TILE; + const uint token = gl_WorkGroupID.y; + const uint stream = gl_WorkGroupID.z; + + for (uint idx = tid; idx < TILE * VEC_PER_HEAD; idx += SUBGROUP_SIZE) { + const uint key = idx / VEC_PER_HEAD; + const uint d4 = idx % VEC_PER_HEAD; + const uint kv = kv_base + key; + f16vec4 value = f16vec4(0.0); + if (kv < p.n_kv) { + const uint offset = stream * p.nbk3 + kv * p.nbk2 + d4 * 4; + value = f16vec4(data_k[offset], data_k[offset + 1], data_k[offset + 2], data_k[offset + 3]); + } + k_sh[key * TILE_STRIDE + d4] = value; + } + barrier(); + + float total = 0.0; + for (uint head_base = 0; head_base < N_HEAD; head_base += TILE) { + for (uint idx = tid; idx < TILE * VEC_PER_HEAD; idx += SUBGROUP_SIZE) { + const uint head_local = idx / VEC_PER_HEAD; + const uint d4 = idx % VEC_PER_HEAD; + const uint head = head_base + head_local; + const uint offset = stream * p.nbq3 + token * p.nbq2 + head * p.nbq1 + d4 * 4; + q_sh[head_local * TILE_STRIDE + d4] = + f16vec4(data_q[offset], data_q[offset + 1], data_q[offset + 2], data_q[offset + 3]); + } + barrier(); + + coopmat scores = + coopmat(0.0); + coopmat kmat; + coopmat qmat; + + [[unroll]] for (uint d = 0; d < HEAD_SIZE; d += TILE) { + coopMatLoad(kmat, k_sh, d / 4, TILE_STRIDE, gl_CooperativeMatrixLayoutRowMajor); + coopMatLoad(qmat, q_sh, d / 4, TILE_STRIDE, gl_CooperativeMatrixLayoutColumnMajor); + scores = coopMatMulAdd(kmat, qmat, scores); + } + + coopMatStore(scores, score_sh, 0, SCORE_STRIDE, gl_CooperativeMatrixLayoutRowMajor); + barrier(); + + if (tid < TILE && kv_base + tid < p.n_kv) { + [[unroll]] for (uint head_local = 0; head_local < TILE; ++head_local) { + const float score = score_sh[tid * SCORE_STRIDE + head_local / 4][head_local % 4]; + const float weight = data_w[stream * p.nbw3 + token * p.nbw1 + head_base + head_local]; + total += max(score, 0.0) * weight; + } + } + barrier(); + } + + if (tid < TILE) { + const uint kv = kv_base + tid; + if (kv < p.n_kv) { + const uint mask_base = (stream % p.nem3) * p.nbm3 + token * p.nbm1; + const uint dst_base = stream * p.nb3 + token * p.nb1; + data_dst[dst_base + kv] = total + float(data_m[mask_base + kv]); + } + } +} diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/lightning_indexer_scalar64.comp b/ggml/src/ggml-vulkan/vulkan-shaders/lightning_indexer_scalar64.comp new file mode 100644 index 00000000000..693bb3ece8e --- /dev/null +++ b/ggml/src/ggml-vulkan/vulkan-shaders/lightning_indexer_scalar64.comp @@ -0,0 +1,92 @@ +#version 450 + +#extension GL_EXT_control_flow_attributes : require +#extension GL_EXT_shader_16bit_storage : require +#extension GL_EXT_shader_explicit_arithmetic_types_float16 : require +#extension GL_KHR_shader_subgroup_arithmetic : require +#extension GL_KHR_shader_subgroup_basic : require + +layout(constant_id = 0) const uint SUBGROUP_SIZE = 64; +layout(local_size_x_id = 0, local_size_y = 1, local_size_z = 1) in; + +layout(binding = 0) readonly buffer QBuf { float data_q[]; }; +layout(binding = 1) readonly buffer KBuf { float16_t data_k[]; }; +layout(binding = 2) readonly buffer WBuf { float data_w[]; }; +layout(binding = 3) readonly buffer MaskBuf { float16_t data_m[]; }; +layout(binding = 4) writeonly buffer DstBuf { float data_dst[]; }; + +layout(push_constant) uniform Parameters { + uint n_kv; + uint n_batch; + uint n_stream; + uint nem3; + uint nb1; + uint nb3; + uint nbq1; + uint nbq2; + uint nbq3; + uint nbk2; + uint nbk3; + uint nbw1; + uint nbw3; + uint nbm1; + uint nbm3; +} p; + +const uint K_PER_GROUP = 8; +const uint N_HEAD = 64; + +void main() { + const uint lane = gl_SubgroupInvocationID; + const uint token = gl_WorkGroupID.y; + const uint stream = gl_WorkGroupID.z; + const uint kv_base = gl_WorkGroupID.x * K_PER_GROUP; + + if (token >= p.n_batch || stream >= p.n_stream) { + return; + } + + float k0[K_PER_GROUP]; + float k1[K_PER_GROUP]; + [[unroll]] for (uint j = 0; j < K_PER_GROUP; ++j) { + const uint kv = kv_base + j; + if (kv < p.n_kv) { + const uint k_base = stream * p.nbk3 + kv * p.nbk2; + k0[j] = float(data_k[k_base + lane]); + k1[j] = float(data_k[k_base + lane + SUBGROUP_SIZE]); + } else { + k0[j] = 0.0; + k1[j] = 0.0; + } + } + + float score[K_PER_GROUP]; + [[unroll]] for (uint j = 0; j < K_PER_GROUP; ++j) { + score[j] = 0.0; + } + + for (uint head = 0; head < N_HEAD; ++head) { + const uint q_base = stream * p.nbq3 + token * p.nbq2 + head * p.nbq1; + const float q0 = data_q[q_base + lane]; + const float q1 = data_q[q_base + lane + SUBGROUP_SIZE]; + const float weight = data_w[stream * p.nbw3 + token * p.nbw1 + head]; + + [[unroll]] for (uint j = 0; j < K_PER_GROUP; ++j) { + const float qk = subgroupAdd(q0 * k0[j] + q1 * k1[j]); + if (lane == 0) { + score[j] += max(qk, 0.0) * weight; + } + } + } + + if (lane == 0) { + const uint mask_base = (stream % p.nem3) * p.nbm3 + token * p.nbm1; + const uint dst_base = stream * p.nb3 + token * p.nb1; + [[unroll]] for (uint j = 0; j < K_PER_GROUP; ++j) { + const uint kv = kv_base + j; + if (kv < p.n_kv) { + data_dst[dst_base + kv] = score[j] + float(data_m[mask_base + kv]); + } + } + } +} 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 81e556b0c78..9b485311baa 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp @@ -811,6 +811,13 @@ void process_shaders() { string_to_spv("get_rows_i32", "get_rows.comp", {{"TEMP_TYPE", "uint"}, {"A_TYPE", "uint"}, {"B_TYPE", "int"}, {"D_TYPE", "uint"}}); + string_to_spv("lightning_indexer_f16", "lightning_indexer_scalar64.comp", {}); +#if defined(GGML_VULKAN_COOPMAT_GLSLC_SUPPORT) + string_to_spv("lightning_indexer_cm_f16", "lightning_indexer_cm.comp", {}); + string_to_spv("lightning_indexer_decode_cm_f16", "lightning_indexer_decode_cm.comp", {}); +#endif + string_to_spv("flash_attn_top_k_f16", "flash_attn_top_k.comp", {}); + string_to_spv("mul_mat_vec_p021_f16_f32_subgroup_add", "mul_mat_vec_p021.comp", {{"A_TYPE", "float16_t"}, {"A_TYPEV4", "f16vec4"}, {"B_TYPE", "float"}, {"B_TYPEV4", "vec4"}, {"D_TYPE", "float"}, {"USE_SUBGROUP_ADD", "1"}}); string_to_spv("mul_mat_vec_p021_f16_f32", "mul_mat_vec_p021.comp", {{"A_TYPE", "float16_t"}, {"A_TYPEV4", "f16vec4"}, {"B_TYPE", "float"}, {"B_TYPEV4", "vec4"}, {"D_TYPE", "float"}}); string_to_spv("mul_mat_vec_nc_f16_f32", "mul_mat_vec_nc.comp", {{"A_TYPE", "float16_t"}, {"A_TYPEV4", "f16vec4"}, {"B_TYPE", "float"}, {"B_TYPEV4", "vec4"}, {"D_TYPE", "float"}}); diff --git a/ggml/src/ggml.c b/ggml/src/ggml.c index 5fcb6c2f300..c5093a039aa 100644 --- a/ggml/src/ggml.c +++ b/ggml/src/ggml.c @@ -5597,6 +5597,20 @@ void ggml_flash_attn_ext_add_sinks( a->src[4] = sinks; } +void ggml_flash_attn_ext_add_top_k( + struct ggml_tensor * a, + struct ggml_tensor * top_k, + int64_t n_kv_raw) { + GGML_ASSERT(a->op == GGML_OP_FLASH_ATTN_EXT); + GGML_ASSERT(a->src[5] == NULL); + GGML_ASSERT(top_k->type == GGML_TYPE_I32); + GGML_ASSERT(top_k->ne[1] == a->src[0]->ne[1]); + GGML_ASSERT(n_kv_raw >= 0 && n_kv_raw <= a->src[1]->ne[1]); + + a->src[5] = top_k; + ggml_set_op_params_i32(a, 4, (int32_t) n_kv_raw); +} + // ggml_flash_attn_back struct ggml_tensor * ggml_flash_attn_back( diff --git a/src/llama-graph.cpp b/src/llama-graph.cpp index b137c6cfe8a..ed56fdbb104 100644 --- a/src/llama-graph.cpp +++ b/src/llama-graph.cpp @@ -2554,7 +2554,9 @@ ggml_tensor * llm_graph_context::build_attn_mha( ggml_tensor * v_mla, int64_t n_kv_max, float kq_scale, - int il) const { + int il, + ggml_tensor * top_k, + int64_t n_kv_raw) const { const bool v_trans = v->nb[1] > v->nb[2]; // split the batch into streams if needed @@ -2592,6 +2594,9 @@ ggml_tensor * llm_graph_context::build_attn_mha( ggml_flash_attn_ext_add_sinks(cur, sinks); GGML_ASSERT(n_kv_max >= 0 && n_kv_max <= INT32_MAX); ggml_flash_attn_ext_set_n_kv_max(cur, static_cast(n_kv_max)); + if (top_k) { + ggml_flash_attn_ext_add_top_k(cur, top_k, n_kv_raw); + } ggml_flash_attn_ext_set_prec (cur, GGML_PREC_F32); if (v_mla) { diff --git a/src/llama-graph.h b/src/llama-graph.h index dddfdac7b51..862142b4fd7 100644 --- a/src/llama-graph.h +++ b/src/llama-graph.h @@ -1173,7 +1173,9 @@ struct llm_graph_context { ggml_tensor * v_mla, // [n_embd_head_v_mla, n_embd_head_v, n_head_v] int64_t n_kv_max, float kq_scale, - int il) const; + int il, + ggml_tensor * top_k = nullptr, + int64_t n_kv_raw = 0) const; llm_graph_input_attn_no_cache * build_attn_inp_no_cache() const; diff --git a/src/models/deepseek4.cpp b/src/models/deepseek4.cpp index 6bf9d344494..8588d069eed 100644 --- a/src/models/deepseek4.cpp +++ b/src/models/deepseek4.cpp @@ -758,7 +758,7 @@ ggml_tensor * llama_model_deepseek4::graph::build_csa_lid_attention( cb(kq_mask, "csa_lid_kq_mask", il); const int64_t n_kv_max = std::min(raw_mask->ne[0], hparams.n_swa) + top_k->ne[0]; - ggml_tensor * out = build_attn_mha(q, k_all, k_all, nullptr, kq_mask, sinks, nullptr, n_kv_max, kq_scale, il); + ggml_tensor * out = build_attn_mha(q, k_all, k_all, nullptr, kq_mask, sinks, nullptr, n_kv_max, kq_scale, il, top_k, raw_k->ne[2]); if (k_rot) { out = llama_mul_mat_hadamard(ctx0, out, k_rot); } diff --git a/tests/test-backend-ops.cpp b/tests/test-backend-ops.cpp index 5b7893d4512..13d8ae928fd 100644 --- a/tests/test-backend-ops.cpp +++ b/tests/test-backend-ops.cpp @@ -11260,6 +11260,9 @@ static std::vector> make_test_cases_eval() { } } } + test_cases.emplace_back(new test_lightning_indexer(128, 64, 257, 1, 1, 1, GGML_TYPE_F16)); + test_cases.emplace_back(new test_lightning_indexer(128, 64, 257, 17, 1, 1, GGML_TYPE_F16)); + test_cases.emplace_back(new test_lightning_indexer(128, 64, 512, 512, 1, 1, GGML_TYPE_F16)); for (int kv : { 1, 7, 8, 63, 64, 65 }) { for (ggml_type type_K : {GGML_TYPE_F32, GGML_TYPE_F16, GGML_TYPE_BF16, GGML_TYPE_Q8_0, GGML_TYPE_Q5_1, GGML_TYPE_Q5_0, GGML_TYPE_Q4_1, GGML_TYPE_Q4_0}) { @@ -11302,6 +11305,11 @@ static std::vector> make_test_cases_perf() { GGML_TYPE_F32, {n_kv, 512, 64, 1}, false, {2, 1, 0, 3})); } + for (ggml_type type_a : { GGML_TYPE_IQ2_XS, GGML_TYPE_IQ3_XXS }) { + test_cases.emplace_back(new test_mul_mat_id(type_a, GGML_TYPE_F32, 256, 6, false, 2048, 512, 4096)); + test_cases.emplace_back(new test_mul_mat_id(type_a, GGML_TYPE_F32, 256, 6, false, 4096, 512, 2048)); + } + // Conv2d: K=CRS=NPQ=4096 matmul performance uint32_t iwh_idx = 0; uint32_t kwh_idx = 1; @@ -11740,7 +11748,7 @@ static std::vector> make_test_cases_perf() { test_cases.emplace_back(new test_gated_delta_net(GGML_TYPE_F32, 32, 128, 2048, 1, 1, false, true)); // lightning_indexer - for (int kv : { 256, 4096, 65536 }) { + for (int kv : { 256, 512, 4096, 65536 }) { for (int bs : { 1, 512, 2048 }) { for (int nh : { 32, 64 }) { for (int ns : { 1, 4 }) { From 56a1843fd9b8f4737ec4ac162a6f1a76a1f5affa Mon Sep 17 00:00:00 2001 From: Nathan Wilson Date: Sat, 1 Aug 2026 13:30:31 +0000 Subject: [PATCH 11/68] vulkan: harden the sparse-FA shader and document the top-k API - flash_attn_top_k.comp: remove the dead bounds check that would skip barriers for part of the workgroup if it ever fired (barrier divergence is UB; the dispatch gate sizes the grid exactly), pin per-lane sizing to LANES=64 instead of the SUBGROUP_SIZE spec constant (array bounds fold from the spec default at compile time), name the f16 mask threshold constant - ggml.h: document ggml_flash_attn_ext_add_top_k semantics (index base, dense prefix, backends-may-ignore contract) - tests: cover the scalar indexer variant (batch 4 and boundary 15), previously only the cm and decode-cm variants had eval parity cases Co-Authored-By: Claude Fable 5 Assisted-by: Claude (Opus 5) --- ggml/include/ggml.h | 5 +++ .../vulkan-shaders/flash_attn_top_k.comp | 36 ++++++++++++------- tests/test-backend-ops.cpp | 4 +++ 3 files changed, 32 insertions(+), 13 deletions(-) diff --git a/ggml/include/ggml.h b/ggml/include/ggml.h index 2ff12d00295..fb83a874e87 100644 --- a/ggml/include/ggml.h +++ b/ggml/include/ggml.h @@ -2486,6 +2486,11 @@ extern "C" { struct ggml_tensor * a, struct ggml_tensor * sinks); + // sparse attention hint: attend only to the first n_kv_raw keys (dense prefix) plus the + // keys selected by top_k. top_k is I32 [n_top_k, n_tokens, 1, n_streams]; each index i + // selects absolute key n_kv_raw + i. Negative or out-of-range indices are ignored. + // Backends may ignore the hint: the kq_mask must still encode the same selection, so a + // dense fallback computes the identical result. GGML_API void ggml_flash_attn_ext_add_top_k( struct ggml_tensor * a, struct ggml_tensor * top_k, diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_top_k.comp b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_top_k.comp index 7f82266781f..b8b49c677bd 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_top_k.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_top_k.comp @@ -39,9 +39,18 @@ layout(push_constant) uniform Parameters { uint has_sinks; } p; +// Shape constants pinned by the dispatch gate in ggml_vk_flash_attn_top_k: DeepSeek V4 +// CSA attention only (hd 512, 64 heads MQA over an f16 K==V latent). The gate also pins +// the pipeline to subgroup size 64, so per-lane sizing uses LANES rather than the +// SUBGROUP_SIZE specialization constant (a spec-constant array bound would be folded at +// compile time from the default and silently break under a different runtime subgroup). const uint HEAD_SIZE = 512; const uint HEADS_PER_GROUP = 8; const uint KEYS_PER_BLOCK = 16; +const uint LANES = 64; +// f16 -inf (or the lowest f16 normal some mask writers use in its place) marks a +// masked-out key. +const float MASK_NEG_INF = -65500.0; shared float16_t key_sh[KEYS_PER_BLOCK * HEAD_SIZE]; shared uint key_idx[KEYS_PER_BLOCK]; @@ -53,12 +62,13 @@ void main() { const uint token = gl_WorkGroupID.x; const uint stream = gl_WorkGroupID.z; - if (token >= p.n_batch || head >= p.n_head) { - return; - } + // No bounds check: the grid is sized exactly (x = n_batch, y * HEADS_PER_GROUP covers + // n_head == 64). An early return here would skip the barriers below for part of the + // workgroup, which is undefined behavior — do not reintroduce one without restructuring + // the barrier flow. - float accum[HEAD_SIZE / SUBGROUP_SIZE]; - [[unroll]] for (uint i = 0; i < HEAD_SIZE / SUBGROUP_SIZE; ++i) { + float accum[HEAD_SIZE / LANES]; + [[unroll]] for (uint i = 0; i < HEAD_SIZE / LANES; ++i) { accum[i] = 0.0; } @@ -101,13 +111,13 @@ void main() { } float partial = 0.0; - [[unroll]] for (uint i = 0; i < HEAD_SIZE / SUBGROUP_SIZE; ++i) { - const uint dim = lane + i * SUBGROUP_SIZE; + [[unroll]] for (uint i = 0; i < HEAD_SIZE / LANES; ++i) { + const uint dim = lane + i * LANES; partial += data_q[q_base + dim] * float(key_sh[col * HEAD_SIZE + dim]); } const float mask = float(data_m[mask_base + key]); const float score = subgroupAdd(partial) * p.scale + mask; - if (mask < -65500.0) { + if (mask < MASK_NEG_INF) { continue; } @@ -117,8 +127,8 @@ void main() { row_sum = row_sum * old_scale + value_scale; row_max = new_max; - [[unroll]] for (uint i = 0; i < HEAD_SIZE / SUBGROUP_SIZE; ++i) { - const uint dim = lane + i * SUBGROUP_SIZE; + [[unroll]] for (uint i = 0; i < HEAD_SIZE / LANES; ++i) { + const uint dim = lane + i * LANES; accum[i] = accum[i] * old_scale + value_scale * float(key_sh[col * HEAD_SIZE + dim]); } } @@ -131,14 +141,14 @@ void main() { const float old_scale = row_sum == 0.0 ? 0.0 : exp(row_max - new_max); const float sink_scale = exp(sink - new_max); row_sum = row_sum * old_scale + sink_scale; - [[unroll]] for (uint i = 0; i < HEAD_SIZE / SUBGROUP_SIZE; ++i) { + [[unroll]] for (uint i = 0; i < HEAD_SIZE / LANES; ++i) { accum[i] *= old_scale; } } const uint dst_base = stream * p.nb3 + token * p.nb2 + head * p.nb1; const float inv_sum = row_sum == 0.0 ? 0.0 : 1.0 / row_sum; - [[unroll]] for (uint i = 0; i < HEAD_SIZE / SUBGROUP_SIZE; ++i) { - data_dst[dst_base + lane + i * SUBGROUP_SIZE] = accum[i] * inv_sum; + [[unroll]] for (uint i = 0; i < HEAD_SIZE / LANES; ++i) { + data_dst[dst_base + lane + i * LANES] = accum[i] * inv_sum; } } diff --git a/tests/test-backend-ops.cpp b/tests/test-backend-ops.cpp index 13d8ae928fd..16a8e382356 100644 --- a/tests/test-backend-ops.cpp +++ b/tests/test-backend-ops.cpp @@ -11260,7 +11260,11 @@ static std::vector> make_test_cases_eval() { } } } + // batch 1 = Vulkan decode-cm variant, 4/15 = scalar subgroup variant (below the cm + // threshold of 16, 15 is the boundary), 17/512 = cm prefill variant test_cases.emplace_back(new test_lightning_indexer(128, 64, 257, 1, 1, 1, GGML_TYPE_F16)); + test_cases.emplace_back(new test_lightning_indexer(128, 64, 257, 4, 1, 1, GGML_TYPE_F16)); + test_cases.emplace_back(new test_lightning_indexer(128, 64, 257, 15, 1, 1, GGML_TYPE_F16)); test_cases.emplace_back(new test_lightning_indexer(128, 64, 257, 17, 1, 1, GGML_TYPE_F16)); test_cases.emplace_back(new test_lightning_indexer(128, 64, 512, 512, 1, 1, GGML_TYPE_F16)); From 79149d1bacf0ae127a6ee5e1d2249228f1f80e3e Mon Sep 17 00:00:00 2001 From: Nathan Wilson Date: Sat, 1 Aug 2026 14:55:05 +0000 Subject: [PATCH 12/68] tests: sparse top-k FA parity + perf coverage (V4 CSA shape) Adds test_flash_attn_ext_top_k: builds the DeepSeek V4 CSA attention shape (hd 512, 64-head MQA, V as a view of K) with a consistent per-token top-k/mask pair, one deliberately invalid index, and cases on both sides of the Vulkan engagement gates. nb >= 64 cases are the first numerical parity coverage the sparse prefill shader has had; nb < 64 and sub-3x-kv cases pin the dense-fallback contract. Perf cases sweep kv 8k/32k/64k at nb 1/8/64/512 with a fixed active set. Measured on gfx1151: the sparse shader is flat vs kv at prefill (~2.2 TFLOPS active-only) while nb < 64 falls back to dense and scales with kv (1326 us at 64k, nb=1) - the gap a sparse decode path needs to close. Co-Authored-By: Claude Fable 5 Assisted-by: Claude (Opus 5) --- tests/test-backend-ops.cpp | 134 +++++++++++++++++++++++++++++++++++++ 1 file changed, 134 insertions(+) diff --git a/tests/test-backend-ops.cpp b/tests/test-backend-ops.cpp index 16a8e382356..248fbd0555f 100644 --- a/tests/test-backend-ops.cpp +++ b/tests/test-backend-ops.cpp @@ -8028,6 +8028,117 @@ struct test_flash_attn_ext : public test_case { } }; +// GGML_OP_FLASH_ATTN_EXT with a top-k sparse selection hint (DeepSeek V4 CSA shape). +// The kq_mask encodes the same selection as the top_k indices, so a backend that ignores +// the hint (CPU) computes the identical result densely — this is exactly the contract +// that keeps sparse and dense paths interchangeable, and what this test verifies. +struct test_flash_attn_ext_top_k : public test_case { + const int64_t kv; // total KV size (compressed region + dense prefix) + const int64_t nb; // batch size (query tokens) + const int64_t n_kv_raw; // dense prefix always attended + const int64_t n_top_k; // selected keys per query token + const bool sinks; + + static constexpr int64_t hs = 512; // V4 CSA head size, K == V latent + static constexpr int64_t nh = 64; // V4 CSA query heads (MQA) + + std::string vars() override { + return VARS_TO_STR5(kv, nb, n_kv_raw, n_top_k, sinks); + } + + double max_nmse_err() override { + return 5e-4; + } + + uint64_t op_flops(ggml_tensor * t) override { + GGML_UNUSED(t); + // only the active keys contribute compute on a sparse backend; count those so + // perf mode reports the useful-work rate + return 2 * nh * nb * (hs + hs) * (n_kv_raw + n_top_k); + } + + test_flash_attn_ext_top_k(int64_t kv = 768, int64_t nb = 8, int64_t n_kv_raw = 64, int64_t n_top_k = 128, bool sinks = false) + : kv(kv), nb(nb), n_kv_raw(n_kv_raw), n_top_k(n_top_k), sinks(sinks) {} + + ggml_tensor * build_graph(ggml_context * ctx) override { + ggml_tensor * q = ggml_new_tensor_4d(ctx, GGML_TYPE_F32, hs, nb, nh, 1); + ggml_set_name(q, "q"); + + ggml_tensor * k = ggml_new_tensor_4d(ctx, GGML_TYPE_F16, hs, kv, 1, 1); + ggml_set_name(k, "k"); + + // V4 CSA attends over the K latent itself: V is the same cache tensor + ggml_tensor * v = ggml_view_4d(ctx, k, hs, kv, 1, 1, k->nb[1], k->nb[2], k->nb[3], 0); + ggml_set_name(v, "v"); + + ggml_tensor * m = ggml_new_tensor_4d(ctx, GGML_TYPE_F16, kv, nb, 1, 1); + ggml_set_name(m, "m"); + + ggml_tensor * t = ggml_new_tensor_4d(ctx, GGML_TYPE_I32, n_top_k, nb, 1, 1); + ggml_set_name(t, "top_k"); + + ggml_tensor * s = nullptr; + if (sinks) { + s = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, nh); + ggml_set_name(s, "s"); + } + + ggml_tensor * out = ggml_flash_attn_ext(ctx, q, k, v, m, 1.0f/sqrtf(hs), 0.0f, 0.0f); + ggml_flash_attn_ext_add_sinks(out, s); + ggml_flash_attn_ext_add_top_k(out, t, n_kv_raw); + ggml_flash_attn_ext_set_prec (out, GGML_PREC_F32); + ggml_set_name(out, "out"); + + return out; + } + + void initialize_tensors(ggml_context * ctx) override { + const int64_t range = kv - n_kv_raw; // size of the selectable compressed region + + for (ggml_tensor * t = ggml_get_first_tensor(ctx); t != NULL; t = ggml_get_next_tensor(ctx, t)) { + if (strcmp(t->name, "top_k") == 0 || strcmp(t->name, "m") == 0) { + continue; // filled together below + } + if (strcmp(t->name, "s") == 0) { + init_tensor_uniform(t, -10.0f, 10.0f); + } else { + init_tensor_uniform(t); + } + } + + // build a consistent (top_k, mask) pair: a deterministic per-token selection, + // strided so adjacent tokens select overlapping-but-different keys, with one + // deliberately invalid index (-1) whose mask slot stays -inf + std::vector top(n_top_k * nb); + std::vector mask(kv * nb); + const ggml_fp16_t minus_inf = ggml_fp32_to_fp16(-INFINITY); + const ggml_fp16_t zero = ggml_fp32_to_fp16(0.0f); + + for (int64_t b = 0; b < nb; ++b) { + for (int64_t i = 0; i < kv; ++i) { + mask[b * kv + i] = i < n_kv_raw ? zero : minus_inf; + } + for (int64_t j = 0; j < n_top_k; ++j) { + int32_t idx = (int32_t) ((j * range) / n_top_k + b) % (int32_t) range; + if (j == n_top_k - 1 && b == 0) { + idx = -1; // exercise the ignore-invalid-index path + } else { + mask[b * kv + n_kv_raw + idx] = zero; + } + top[b * n_top_k + j] = idx; + } + } + + for (ggml_tensor * t = ggml_get_first_tensor(ctx); t != NULL; t = ggml_get_next_tensor(ctx, t)) { + if (strcmp(t->name, "top_k") == 0) { + ggml_backend_tensor_set(t, top.data(), 0, top.size() * sizeof(int32_t)); + } else if (strcmp(t->name, "m") == 0) { + ggml_backend_tensor_set(t, mask.data(), 0, mask.size() * sizeof(ggml_fp16_t)); + } + } + } +}; + // GGML_OP_CROSS_ENTROPY_LOSS struct test_cross_entropy_loss : public test_case { const ggml_type type; @@ -11274,6 +11385,18 @@ static std::vector> make_test_cases_eval() { } } + // sparse top-k FA: (kv, nb, n_kv_raw, n_top_k, sinks). The Vulkan sparse path engages + // when kv >= 3*(n_kv_raw + n_top_k) AND nb >= 64 (prefill-only); the nb < 64 cases + // and the kv=512 case verify dense-fallback parity with the hint attached, the + // nb=64/128 cases exercise the sparse shader itself. + test_cases.emplace_back(new test_flash_attn_ext_top_k(4096, 1, 256, 512, false)); + test_cases.emplace_back(new test_flash_attn_ext_top_k( 768, 8, 64, 128, false)); + test_cases.emplace_back(new test_flash_attn_ext_top_k( 768, 17, 64, 128, false)); + test_cases.emplace_back(new test_flash_attn_ext_top_k( 512, 4, 64, 128, false)); + test_cases.emplace_back(new test_flash_attn_ext_top_k( 768, 64, 64, 128, false)); + test_cases.emplace_back(new test_flash_attn_ext_top_k( 768, 64, 64, 128, true)); + test_cases.emplace_back(new test_flash_attn_ext_top_k(4096, 128, 256, 512, false)); + return test_cases; } #ifdef _MSC_VER @@ -11764,6 +11887,17 @@ static std::vector> make_test_cases_perf() { } } + // sparse top-k FA at V4 decode/prefill shapes — the A/B instrument for the + // gather-to-compact work (n_active = n_kv_raw + n_top_k stays fixed as kv grows). + // nb 1/8 currently takes the DENSE path (the sparse shader gates on nb >= 64): + // those rows measure the decode cost gather-to-compact must beat. nb 64/512 + // measures the existing sparse prefill shader. + for (int kv : { 8192, 32768, 65536 }) { + for (int nb : { 1, 8, 64, 512 }) { + test_cases.emplace_back(new test_flash_attn_ext_top_k(kv, nb, 1024, 512, false)); + } + } + return test_cases; } From 5d1180a63265cc1394d31aaed6d2062a25f65a51 Mon Sep 17 00:00:00 2001 From: Nathan Wilson Date: Sat, 1 Aug 2026 15:08:58 +0000 Subject: [PATCH 13/68] vulkan: gather-to-compact sparse decode FA for DeepSeek V4 top-k selection The sparse prefill shader gates on q->ne[1] >= 64, so single-token decode attends densely over the whole compressed KV and its cost grows with context. This adds a gather pass (flash_attn_gather.comp): copy the active rows (dense prefix + top-k selection; MQA, so all 64 query heads share one set) plus their mask values into a compact contiguous scratch in prealloc_y, then run the ordinary dense FA over the compacted K/V/mask. V is the K latent, so one gather serves both. Invalid indices and padding get zeroed K and -inf mask. The FA function itself only has its inputs swapped: KV, mask geometry, strides and the K/V/mask bindings are overridden up front and every downstream decision (pipeline choice, split-k, workgroup sizing, use_mask_opt) sizes itself to the compact KV unchanged. Engages for the V4 CSA decode shape when kv >= 2x the padded active set; GGML_VK_FA_TOPK_GATHER=0 disables. Measured (gfx1151, test-backend-ops perf, active set 1536): kv=8192 nb=1: 255.6 us -> 55.8 us (4.6x) kv=32768 nb=1: 986.4 us -> 56.1 us (17.6x) kv=65536 nb=1: 1333.8 us -> 58.7 us (22.7x) Time is flat vs context. nb>1 still falls back to dense pending a union gather. FLASH_ATTN_EXT eval suite green incl. the top-k parity cases (the kv=4096 nb=1 case exercises this path end-to-end vs CPU). Co-Authored-By: Claude Fable 5 Fold-in note for the toolbox branch: use_dequant_kv is additionally gated on !fa_compact.active (the compact scratch is already contiguous f16, and the two scratch layers must not stack), and the compact stride overrides chain through the toolbox's nb*_eff values so the contiguize/dequant path keeps its strides when the gather is inactive. Co-Authored-By: Claude Fable 5 Assisted-by: Claude (Opus 5) --- ggml/src/ggml-vulkan/ggml-vulkan.cpp | 143 +++++++++++++++++- .../vulkan-shaders/flash_attn_gather.comp | 70 +++++++++ .../vulkan-shaders/vulkan-shaders-gen.cpp | 1 + 3 files changed, 207 insertions(+), 7 deletions(-) create mode 100644 ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_gather.comp diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp index 44efa5c4897..012b18a44ba 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp +++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp @@ -1108,6 +1108,7 @@ struct vk_device_struct { vk_pipeline pipeline_lightning_indexer_cm_f16; vk_pipeline pipeline_lightning_indexer_decode_cm_f16; vk_pipeline pipeline_flash_attn_top_k_f16; + vk_pipeline pipeline_flash_attn_gather_f16; vk_pipeline pipeline_ssm_scan_f32_d128; vk_pipeline pipeline_ssm_scan_f32_d256; vk_pipeline pipeline_ssm_conv_f32; @@ -1948,6 +1949,12 @@ struct vk_op_lightning_indexer_cm_push_constants { }; static_assert(sizeof(vk_op_lightning_indexer_cm_push_constants) <= 128); +struct vk_op_flash_attn_gather_push_constants { + uint32_t n_kv, n_kv_raw, n_top_k, kv_c; + uint32_t nbk1, nbk3, nbt3, nbm3, nem3; +}; +static_assert(sizeof(vk_op_flash_attn_gather_push_constants) <= 128); + struct vk_op_flash_attn_top_k_push_constants { uint32_t n_batch, n_kv, n_kv_raw, n_top_k, n_head; uint32_t nbq1, nbq2, nbq3; @@ -6253,6 +6260,10 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { "flash_attn_top_k_f16", flash_attn_top_k_f16_len, flash_attn_top_k_f16_data, "main", 6, sizeof(vk_op_flash_attn_top_k_push_constants), {1, 1, 1}, {512, device->subgroup_size}, 1, true, true, device->subgroup_size); + ggml_vk_create_pipeline(device, device->pipeline_flash_attn_gather_f16, + "flash_attn_gather_f16", flash_attn_gather_f16_len, flash_attn_gather_f16_data, "main", 5, + sizeof(vk_op_flash_attn_gather_push_constants), {1, 1, 1}, {}, 1, true, true, + device->subgroup_size); } if (device->subgroup_arithmetic && device->subgroup_require_full_support) { @@ -11430,6 +11441,95 @@ static bool ggml_vk_flash_attn_top_k(ggml_backend_vk_context * ctx, vk_context & return true; } +struct vk_fa_compact_state { + bool active = false; + uint32_t kv_c = 0; + vk_subbuffer kc_buf, mc_buf; +}; + +// V4 sparse decode (gather-to-compact): the sparse prefill shader above gates on +// q->ne[1] >= 64, so single-token decode otherwise attends densely over the whole +// compressed KV, at a cost that grows with context. Instead, gather the active rows +// (dense prefix + top-k selection; MQA, so all query heads share one set) into a +// compact contiguous scratch in prealloc_y, and let the ordinary dense FA below run +// over the compacted K/V/mask. Correct by the same contract as the sparse shader: +// the source mask carries the selection, and the gathered mask preserves it. +static bool ggml_vk_flash_attn_gather_compact(ggml_backend_vk_context * ctx, vk_context & subctx, + const ggml_tensor * q, const ggml_tensor * k, const ggml_tensor * v, + const ggml_tensor * mask, ggml_tensor * dst, vk_fa_compact_state & st) { + const ggml_tensor * top_k = dst->src[5]; + static const char * gather_env = getenv("GGML_VK_FA_TOPK_GATHER"); + if ((gather_env && gather_env[0] == '0') || + !top_k || !ctx->device->pipeline_flash_attn_gather_f16 || + q->ne[1] != 1 || // single-token decode only; batched queries need a union gather + q->type != GGML_TYPE_F32 || k->type != GGML_TYPE_F16 || v->type != GGML_TYPE_F16 || + !mask || mask->type != GGML_TYPE_F16 || top_k->type != GGML_TYPE_I32 || + q->ne[0] != 512 || k->ne[0] != 512 || v->ne[0] != 512 || q->ne[2] != 64 || + k->ne[2] != 1 || v->ne[2] != 1 || + q->ne[1] != top_k->ne[1] || q->ne[3] != top_k->ne[3] || + k->ne[1] != v->ne[1] || k->buffer != v->buffer || k->data != v->data || + !ggml_is_contiguous(mask) || !ggml_is_contiguous(top_k)) { + return false; + } + + 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)); + if (max_bias != 0.0f || logit_softcap != 0.0f) { + return false; + } + + const int32_t n_kv_raw = ggml_get_op_params_i32(dst, 4); + if (n_kv_raw < 0 || n_kv_raw > k->ne[1] || top_k->ne[0] > k->ne[1] - n_kv_raw) { + return false; + } + + const uint32_t kv_c = GGML_PAD((uint32_t)(n_kv_raw + top_k->ne[0]), 256u); + // the gather writes then re-reads ~the active bytes; dense reads the source KV once, + // so compaction only pays when the source is comfortably larger than the active set + if ((uint64_t) k->ne[1] < 2ull * kv_c) { + return false; + } + + const uint32_t ns = (uint32_t) q->ne[3]; + const size_t kc_sz = (size_t) ns * kv_c * 512 * sizeof(ggml_fp16_t); + const size_t mc_sz = (size_t) ns * kv_c * sizeof(ggml_fp16_t); + + if (ctx->prealloc_size_y < kc_sz + mc_sz) { + ctx->prealloc_size_y = kc_sz + mc_sz; + ggml_vk_preallocate_buffers(ctx, subctx); + } + if (ctx->prealloc_y_need_sync) { + ggml_vk_sync_buffers(ctx, subctx); + } + + vk_pipeline pipeline = ctx->device->pipeline_flash_attn_gather_f16; + ggml_pipeline_request_descriptor_sets(ctx, pipeline, 1); + + const vk_op_flash_attn_gather_push_constants pc = { + (uint32_t) k->ne[1], (uint32_t) n_kv_raw, (uint32_t) top_k->ne[0], kv_c, + (uint32_t) (k->nb[1] / sizeof(ggml_fp16_t)), + (uint32_t) (k->nb[3] / sizeof(ggml_fp16_t)), + (uint32_t) (top_k->nb[3] / sizeof(int32_t)), + (uint32_t) (mask->nb[3] / sizeof(ggml_fp16_t)), + (uint32_t) mask->ne[3], + }; + + st.kc_buf = ggml_vk_subbuffer(ctx, ctx->prealloc_y, 0); + st.mc_buf = ggml_vk_subbuffer(ctx, ctx->prealloc_y, kc_sz); + ggml_vk_dispatch_pipeline(ctx, subctx, pipeline, + { ggml_vk_tensor_subbuffer(ctx, k), ggml_vk_tensor_subbuffer(ctx, top_k), + ggml_vk_tensor_subbuffer(ctx, mask), st.kc_buf, st.mc_buf }, + pc, { kv_c, 1, ns }); + ggml_vk_sync_buffers(ctx, subctx); + ctx->prealloc_y_need_sync = true; + + st.active = true; + st.kv_c = kv_c; + return true; +} + static void ggml_vk_flash_attn(ggml_backend_vk_context * ctx, vk_context& subctx, const ggml_tensor * q, const ggml_tensor * k, const ggml_tensor * v, const ggml_tensor * mask, const ggml_tensor * sinks, ggml_tensor * dst) { VK_LOG_DEBUG("ggml_vk_flash_attn((" << q << ", name=" << q->name << ", type=" << q->type << ", ne0=" << q->ne[0] << ", ne1=" << q->ne[1] << ", ne2=" << q->ne[2] << ", ne3=" << q->ne[3] << ", nb0=" << q->nb[0] << ", nb1=" << q->nb[1] << ", nb2=" << q->nb[2] << ", nb3=" << q->nb[3]; std::cerr << "), (" << k << ", name=" << k->name << ", type=" << k->type << ", ne0=" << k->ne[0] << ", ne1=" << k->ne[1] << ", ne2=" << k->ne[2] << ", ne3=" << k->ne[3] << ", nb0=" << k->nb[0] << ", nb1=" << k->nb[1] << ", nb2=" << k->nb[2] << ", nb3=" << k->nb[3]; @@ -11449,15 +11549,15 @@ static void ggml_vk_flash_attn(ggml_backend_vk_context * ctx, vk_context& subctx GGML_TENSOR_LOCALS(int64_t, ne, dst, ne) GGML_TENSOR_LOCALS(size_t, nb, dst, nb) - const uint32_t nem0 = mask ? mask->ne[0] : 0; - const uint32_t nem1 = mask ? mask->ne[1] : 0; - const uint32_t nem2 = mask ? mask->ne[2] : 0; - const uint32_t nem3 = mask ? mask->ne[3] : 0; + uint32_t nem0 = mask ? mask->ne[0] : 0; + uint32_t nem1 = mask ? mask->ne[1] : 0; + uint32_t nem2 = mask ? mask->ne[2] : 0; + uint32_t nem3 = mask ? mask->ne[3] : 0; const uint32_t HSK = nek0; const uint32_t HSV = nev0; uint32_t N = neq1; - const uint32_t KV = nek1; + uint32_t KV = nek1; GGML_ASSERT(ne0 == HSV); GGML_ASSERT(ne2 == N); @@ -11484,6 +11584,17 @@ static void ggml_vk_flash_attn(ggml_backend_vk_context * ctx, vk_context& subctx if (ggml_vk_flash_attn_top_k(ctx, subctx, q, k, v, mask, sinks, dst)) { return; } + // V4 sparse decode: gather the active set into a compact scratch and run the dense + // FA below on it. Overrides KV, the mask geometry, and (further down) the K/V/mask + // bindings and strides; every other decision then sizes itself to the compact KV. + vk_fa_compact_state fa_compact; + if (ggml_vk_flash_attn_gather_compact(ctx, subctx, q, k, v, mask, dst, fa_compact)) { + KV = fa_compact.kv_c; + nem0 = fa_compact.kv_c; + nem1 = N; + nem2 = 1; + nem3 = (uint32_t) q->ne[3]; + } uint32_t gqa_ratio = 1; uint32_t qk_ratio = neq2 / nek2; uint32_t workgroups_x = (uint32_t)neq1; @@ -11530,6 +11641,9 @@ static void ggml_vk_flash_attn(ggml_backend_vk_context * ctx, vk_context& subctx const bool kv_needs_dequant = !ggml_vk_fa_kv_native(k->type, ctx->device->coopmat2) || !ggml_vk_fa_kv_native(v->type, ctx->device->coopmat2); const bool use_dequant_kv = !fa_dequant_off && + // the gather-to-compact scratch is already contiguous f16; the + // dequant/contiguize pass must not run on top of it + !fa_compact.active && ((k_quant && v_quant) || kv_needs_dequant || (fa_kv_contig && kv_f16_strided)) && neq1 >= 64 && is_dense_kv_cache(k) && is_dense_kv_cache(v) && kv_f16_sz <= ctx->device->properties.limits.maxStorageBufferRange && @@ -11568,6 +11682,10 @@ static void ggml_vk_flash_attn(ggml_backend_vk_context * ctx, vk_context& subctx 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)); + if (fa_compact.active) { + k_stride = 512; + v_stride = 512; + } // For F32, the shader treats it as a block of size 4 (for vec4 loads) if (k->type == GGML_TYPE_F32) { @@ -11715,6 +11833,11 @@ static void ggml_vk_flash_attn(ggml_backend_vk_context * ctx, vk_context& subctx vk_subbuffer v_buf = ggml_vk_tensor_subbuffer(ctx, v); vk_subbuffer dst_buf = ggml_vk_tensor_subbuffer(ctx, dst); vk_subbuffer mask_buf = mask ? ggml_vk_tensor_subbuffer(ctx, mask) : q_buf; + if (fa_compact.active) { + k_buf = fa_compact.kc_buf; + v_buf = fa_compact.kc_buf; // V is the K latent; one gather serves both + mask_buf = fa_compact.mc_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; @@ -11773,6 +11896,12 @@ static void ggml_vk_flash_attn(ggml_backend_vk_context * ctx, vk_context& subctx ggml_vk_sync_buffers(ctx, subctx); } + // compact scratch layout: [512, kv_c, 1, ns] f16, tightly packed + const uint32_t eff_nbk2 = fa_compact.active ? fa_compact.kv_c * 512 * (uint32_t)sizeof(ggml_fp16_t) : nbk2_eff; + const uint32_t eff_nbk3 = fa_compact.active ? fa_compact.kv_c * 512 * (uint32_t)sizeof(ggml_fp16_t) : nbk3_eff; + const uint32_t eff_nbv2 = fa_compact.active ? eff_nbk2 : nbv2_eff; + const uint32_t eff_nbv3 = fa_compact.active ? eff_nbk3 : nbv3_eff; + const vk_flash_attn_push_constants pc = { N, KV, (uint32_t)ne1, (uint32_t)ne2, (uint32_t)ne3, (uint32_t)neq2, (uint32_t)neq3, @@ -11780,8 +11909,8 @@ static void ggml_vk_flash_attn(ggml_backend_vk_context * ctx, vk_context& subctx (uint32_t)nev2, (uint32_t)nev3, nem1, nem2, nem3, q_stride, (uint32_t)nbq2, (uint32_t)nbq3, - k_stride, nbk2_eff, nbk3_eff, - v_stride, nbv2_eff, nbv3_eff, + k_stride, eff_nbk2, eff_nbk3, + v_stride, eff_nbv2, eff_nbv3, scale, max_bias, logit_softcap, mask_n_head_log2, m0, m1, gqa_ratio, split_kv, split_k }; diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_gather.comp b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_gather.comp new file mode 100644 index 00000000000..4e3dc1c4262 --- /dev/null +++ b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_gather.comp @@ -0,0 +1,70 @@ +#version 450 + +#extension GL_EXT_control_flow_attributes : require +#extension GL_EXT_shader_16bit_storage : require +#extension GL_EXT_shader_explicit_arithmetic_types_float16 : require + +// Gathers the active KV rows of a top-k sparse attention (DeepSeek V4 CSA decode) into a +// compact contiguous scratch: rows [0, n_kv_raw) of the source (the dense prefix), then the +// n_top_k selected rows, then zero padding up to kv_c. The gathered mask row keeps the +// per-key mask values so causality/validity survive compaction; invalid top-k indices and +// padding get -inf mask and zeroed K (softmax-neutral either way, zeroed so no NaN*0). +// One workgroup per compact row; V is the K latent (V==K), so a single gather serves both. + +layout(local_size_x = 64, local_size_y = 1, local_size_z = 1) in; + +layout(binding = 0) readonly buffer KBuf { float16_t data_k[]; }; +layout(binding = 1) readonly buffer TopBuf { int data_top[]; }; +layout(binding = 2) readonly buffer MBuf { float16_t data_m[]; }; +layout(binding = 3) writeonly buffer KcBuf { float16_t data_kc[]; }; +layout(binding = 4) writeonly buffer McBuf { float16_t data_mc[]; }; + +layout(push_constant) uniform Parameters { + uint n_kv; // total source KV rows + uint n_kv_raw; // dense prefix length + uint n_top_k; // selected rows for the (single) query token + uint kv_c; // padded compact row count == dispatch row range + uint nbk1; // K source row stride, elements + uint nbk3; // K source stream stride, elements + uint nbt3; // top_k stream stride, elements + uint nbm3; // mask source stream stride, elements + uint nem3; // mask ne[3], for stream broadcast +} p; + +const uint HEAD_SIZE = 512; +const uint LANES = 64; + +void main() { + const uint row = gl_WorkGroupID.x; + const uint stream = gl_WorkGroupID.z; + const uint tid = gl_LocalInvocationIndex; + + // map compact row -> source row; p.n_kv is the invalid sentinel + uint src = p.n_kv; + if (row < p.n_kv_raw) { + src = row; + } else if (row < p.n_kv_raw + p.n_top_k) { + const int idx = data_top[stream * p.nbt3 + (row - p.n_kv_raw)]; + if (idx >= 0 && uint(idx) < p.n_kv - p.n_kv_raw) { + src = p.n_kv_raw + uint(idx); + } + } + + const uint dst_base = (stream * p.kv_c + row) * HEAD_SIZE; + if (src < p.n_kv) { + const uint src_base = stream * p.nbk3 + src * p.nbk1; + [[unroll]] for (uint i = 0; i < HEAD_SIZE / LANES; ++i) { + data_kc[dst_base + tid + i * LANES] = data_k[src_base + tid + i * LANES]; + } + if (tid == 0) { + data_mc[stream * p.kv_c + row] = data_m[(stream % p.nem3) * p.nbm3 + src]; + } + } else { + [[unroll]] for (uint i = 0; i < HEAD_SIZE / LANES; ++i) { + data_kc[dst_base + tid + i * LANES] = float16_t(0.0); + } + if (tid == 0) { + data_mc[stream * p.kv_c + row] = float16_t(uintBitsToFloat(0xff800000)); + } + } +} 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 9b485311baa..5c06f822292 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp @@ -817,6 +817,7 @@ void process_shaders() { string_to_spv("lightning_indexer_decode_cm_f16", "lightning_indexer_decode_cm.comp", {}); #endif string_to_spv("flash_attn_top_k_f16", "flash_attn_top_k.comp", {}); + string_to_spv("flash_attn_gather_f16", "flash_attn_gather.comp", {}); string_to_spv("mul_mat_vec_p021_f16_f32_subgroup_add", "mul_mat_vec_p021.comp", {{"A_TYPE", "float16_t"}, {"A_TYPEV4", "f16vec4"}, {"B_TYPE", "float"}, {"B_TYPEV4", "vec4"}, {"D_TYPE", "float"}, {"USE_SUBGROUP_ADD", "1"}}); string_to_spv("mul_mat_vec_p021_f16_f32", "mul_mat_vec_p021.comp", {{"A_TYPE", "float16_t"}, {"A_TYPEV4", "f16vec4"}, {"B_TYPE", "float"}, {"B_TYPEV4", "vec4"}, {"D_TYPE", "float"}}); From 554519bbfe3251dc7fe7023a9e8bbf82ae10a65b Mon Sep 17 00:00:00 2001 From: Nathan Wilson Date: Sun, 2 Aug 2026 00:04:55 +0000 Subject: [PATCH 14/68] vulkan: fused DeepSeek V4 hyper-connection ops (HC pre / comb / post) Ports ggml-cuda/dsv4-hc.cu to Vulkan: hc_pre (mix the HC input streams down to one embedding), hc_comb (per-token 4x4 stream-mixing matrix - per-source softmax then eps-stabilized alternating column/row sinkhorn normalization, whole matrix in registers, one thread per token), and hc_post (redistribute the layer output back into the streams with the mixed residual). Same launch geometry as the CUDA kernels; plain f32 compute, no subgroup or coopmat requirements, so the pipelines are created unconditionally. The value at decode is dispatch-count collapse: the unfused fallback runs the decomposed graph (measured on gfx1151 config-a partial offload: SUM_ROWS alone 79 dispatches x 39.7us = 3.1ms per graph, plus DIV/MUL/ ADD shares at 4x4 shapes) where the fused form is 3 dispatches per layer. resolve_fused_ops now keeps all three fusions enabled on Vulkan instead of printing 'not supported, set to disabled'. Parity: upstream test-backend-ops DSV4_HC_PRE/COMB/POST cases green vs CPU on first build. Co-Authored-By: Claude Fable 5 Assisted-by: Claude (Opus 5) --- ggml/src/ggml-vulkan/ggml-vulkan.cpp | 153 ++++++++++++++++++ .../vulkan-shaders/dsv4_hc_comb.comp | 101 ++++++++++++ .../vulkan-shaders/dsv4_hc_post.comp | 44 +++++ .../vulkan-shaders/dsv4_hc_pre.comp | 37 +++++ .../vulkan-shaders/vulkan-shaders-gen.cpp | 3 + 5 files changed, 338 insertions(+) create mode 100644 ggml/src/ggml-vulkan/vulkan-shaders/dsv4_hc_comb.comp create mode 100644 ggml/src/ggml-vulkan/vulkan-shaders/dsv4_hc_post.comp create mode 100644 ggml/src/ggml-vulkan/vulkan-shaders/dsv4_hc_pre.comp diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp index 012b18a44ba..8a2d6fe9954 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp +++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp @@ -1109,6 +1109,9 @@ struct vk_device_struct { vk_pipeline pipeline_lightning_indexer_decode_cm_f16; vk_pipeline pipeline_flash_attn_top_k_f16; vk_pipeline pipeline_flash_attn_gather_f16; + vk_pipeline pipeline_dsv4_hc_pre_f32; + vk_pipeline pipeline_dsv4_hc_comb_f32; + vk_pipeline pipeline_dsv4_hc_post_f32; vk_pipeline pipeline_ssm_scan_f32_d128; vk_pipeline pipeline_ssm_scan_f32_d256; vk_pipeline pipeline_ssm_conv_f32; @@ -1949,6 +1952,35 @@ struct vk_op_lightning_indexer_cm_push_constants { }; static_assert(sizeof(vk_op_lightning_indexer_cm_push_constants) <= 128); +struct vk_op_dsv4_hc_pre_push_constants { + uint32_t n_embd, hc, nr; + uint32_t sx0, sx1, sx2; + uint32_t sw0, sw1; + uint32_t sd0, sd1; +}; +static_assert(sizeof(vk_op_dsv4_hc_pre_push_constants) <= 128); + +struct vk_op_dsv4_hc_comb_push_constants { + uint32_t n_tokens; + uint32_t sm0, sm1; + uint32_t ss0; + uint32_t sb0; + uint32_t sd0, sd1, sd2; + float eps; + int32_t n_iter; +}; +static_assert(sizeof(vk_op_dsv4_hc_comb_push_constants) <= 128); + +struct vk_op_dsv4_hc_post_push_constants { + uint32_t n_embd, hc, nr; + uint32_t sx0, sx1; + uint32_t sr0, sr1, sr2; + uint32_t sp0, sp1; + uint32_t sc0, sc1, sc2; + uint32_t sd0, sd1, sd2; +}; +static_assert(sizeof(vk_op_dsv4_hc_post_push_constants) <= 128); + struct vk_op_flash_attn_gather_push_constants { uint32_t n_kv, n_kv_raw, n_top_k, kv_c; uint32_t nbk1, nbk3, nbt3, nbm3, nem3; @@ -6266,6 +6298,17 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { device->subgroup_size); } + // DSv4 fused hyper-connection ops: plain f32 compute, no subgroup/coopmat requirements + ggml_vk_create_pipeline(device, device->pipeline_dsv4_hc_pre_f32, + "dsv4_hc_pre_f32", dsv4_hc_pre_f32_len, dsv4_hc_pre_f32_data, "main", 3, + sizeof(vk_op_dsv4_hc_pre_push_constants), {256, 1, 1}, {}, 1); + ggml_vk_create_pipeline(device, device->pipeline_dsv4_hc_comb_f32, + "dsv4_hc_comb_f32", dsv4_hc_comb_f32_len, dsv4_hc_comb_f32_data, "main", 4, + sizeof(vk_op_dsv4_hc_comb_push_constants), {256, 1, 1}, {}, 1); + ggml_vk_create_pipeline(device, device->pipeline_dsv4_hc_post_f32, + "dsv4_hc_post_f32", dsv4_hc_post_f32_len, dsv4_hc_post_f32_data, "main", 5, + sizeof(vk_op_dsv4_hc_post_push_constants), {256, 1, 1}, {}, 1); + if (device->subgroup_arithmetic && device->subgroup_require_full_support) { ggml_vk_create_pipeline(device, device->pipeline_ssm_scan_f32_d128, "ssm_scan_128_f32", ssm_scan_subgroup_f32_len, ssm_scan_subgroup_f32_data, "main", 8, sizeof(vk_op_ssm_scan_push_constants), {1, 1, 1}, {128, device->subgroup_size}, 1, true, true); ggml_vk_create_pipeline(device, device->pipeline_ssm_scan_f32_d256, "ssm_scan_256_f32", ssm_scan_subgroup_f32_len, ssm_scan_subgroup_f32_data, "main", 8, sizeof(vk_op_ssm_scan_push_constants), {1, 1, 1}, {256, device->subgroup_size}, 1, true, true); @@ -13760,6 +13803,91 @@ static void ggml_vk_lightning_indexer_cm(ggml_backend_vk_context * ctx, vk_conte pc, {(uint32_t) k->ne[2], (uint32_t) q->ne[2], (uint32_t) q->ne[3]}); } +// DSv4 fused hyper-connection ops — ports of ggml-cuda/dsv4-hc.cu. Strides are passed in +// f32 elements; grids mirror the CUDA launch geometry (flat 1D, 256 threads per workgroup). + +static void ggml_vk_dsv4_hc_pre(ggml_backend_vk_context * ctx, vk_context& subctx, ggml_tensor * dst) { + const ggml_tensor * x = dst->src[0]; + const ggml_tensor * w = dst->src[1]; + + const uint32_t n_embd = (uint32_t) x->ne[0]; + const uint32_t hc = (uint32_t) x->ne[1]; + const uint32_t n_tokens = (uint32_t) x->ne[2]; + const uint32_t nr = n_embd * n_tokens; + + vk_pipeline pipeline = ctx->device->pipeline_dsv4_hc_pre_f32; + ggml_pipeline_request_descriptor_sets(ctx, pipeline, 1); + + const vk_op_dsv4_hc_pre_push_constants pc = { + n_embd, hc, nr, + (uint32_t)(x->nb[0] / sizeof(float)), (uint32_t)(x->nb[1] / sizeof(float)), (uint32_t)(x->nb[2] / sizeof(float)), + (uint32_t)(w->nb[0] / sizeof(float)), (uint32_t)(w->nb[1] / sizeof(float)), + (uint32_t)(dst->nb[0] / sizeof(float)), (uint32_t)(dst->nb[1] / sizeof(float)), + }; + + ggml_vk_dispatch_pipeline(ctx, subctx, pipeline, + {ggml_vk_tensor_subbuffer(ctx, x), ggml_vk_tensor_subbuffer(ctx, w), + ggml_vk_tensor_subbuffer(ctx, dst)}, + pc, {nr, 1, 1}); +} + +static void ggml_vk_dsv4_hc_comb(ggml_backend_vk_context * ctx, vk_context& subctx, ggml_tensor * dst) { + const ggml_tensor * mixes = dst->src[0]; + const ggml_tensor * scale = dst->src[1]; + const ggml_tensor * base = dst->src[2]; + + const uint32_t n_tokens = (uint32_t) mixes->ne[1]; + const float eps = ggml_get_op_params_f32(dst, 0); + const int32_t n_iter = ggml_get_op_params_i32(dst, 1); + + vk_pipeline pipeline = ctx->device->pipeline_dsv4_hc_comb_f32; + ggml_pipeline_request_descriptor_sets(ctx, pipeline, 1); + + const vk_op_dsv4_hc_comb_push_constants pc = { + n_tokens, + (uint32_t)(mixes->nb[0] / sizeof(float)), (uint32_t)(mixes->nb[1] / sizeof(float)), + (uint32_t)(scale->nb[0] / sizeof(float)), + (uint32_t)(base->nb[0] / sizeof(float)), + (uint32_t)(dst->nb[0] / sizeof(float)), (uint32_t)(dst->nb[1] / sizeof(float)), (uint32_t)(dst->nb[2] / sizeof(float)), + eps, n_iter, + }; + + ggml_vk_dispatch_pipeline(ctx, subctx, pipeline, + {ggml_vk_tensor_subbuffer(ctx, mixes), ggml_vk_tensor_subbuffer(ctx, scale), + ggml_vk_tensor_subbuffer(ctx, base), ggml_vk_tensor_subbuffer(ctx, dst)}, + pc, {n_tokens, 1, 1}); +} + +static void ggml_vk_dsv4_hc_post(ggml_backend_vk_context * ctx, vk_context& subctx, ggml_tensor * dst) { + const ggml_tensor * x = dst->src[0]; + const ggml_tensor * residual = dst->src[1]; + const ggml_tensor * post = dst->src[2]; + const ggml_tensor * comb = dst->src[3]; + + const uint32_t n_embd = (uint32_t) x->ne[0]; + const uint32_t n_tokens = (uint32_t) x->ne[1]; + const uint32_t hc = (uint32_t) residual->ne[1]; + const uint32_t nr = n_embd * hc * n_tokens; + + vk_pipeline pipeline = ctx->device->pipeline_dsv4_hc_post_f32; + ggml_pipeline_request_descriptor_sets(ctx, pipeline, 1); + + const vk_op_dsv4_hc_post_push_constants pc = { + n_embd, hc, nr, + (uint32_t)(x->nb[0] / sizeof(float)), (uint32_t)(x->nb[1] / sizeof(float)), + (uint32_t)(residual->nb[0] / sizeof(float)), (uint32_t)(residual->nb[1] / sizeof(float)), (uint32_t)(residual->nb[2] / sizeof(float)), + (uint32_t)(post->nb[0] / sizeof(float)), (uint32_t)(post->nb[1] / sizeof(float)), + (uint32_t)(comb->nb[0] / sizeof(float)), (uint32_t)(comb->nb[1] / sizeof(float)), (uint32_t)(comb->nb[2] / sizeof(float)), + (uint32_t)(dst->nb[0] / sizeof(float)), (uint32_t)(dst->nb[1] / sizeof(float)), (uint32_t)(dst->nb[2] / sizeof(float)), + }; + + ggml_vk_dispatch_pipeline(ctx, subctx, pipeline, + {ggml_vk_tensor_subbuffer(ctx, x), ggml_vk_tensor_subbuffer(ctx, residual), + ggml_vk_tensor_subbuffer(ctx, post), ggml_vk_tensor_subbuffer(ctx, comb), + ggml_vk_tensor_subbuffer(ctx, dst)}, + pc, {nr, 1, 1}); +} + static void ggml_vk_ssm_scan(ggml_backend_vk_context * ctx, vk_context& subctx, ggml_tensor * dst) { const ggml_tensor * src0 = dst->src[0]; const ggml_tensor * src1 = dst->src[1]; @@ -16921,6 +17049,21 @@ static bool ggml_vk_build_graph(ggml_backend_vk_context * ctx, ggml_cgraph * cgr break; + case GGML_OP_DSV4_HC_PRE: + ggml_vk_dsv4_hc_pre(ctx, compute_ctx, node); + + break; + + case GGML_OP_DSV4_HC_COMB: + ggml_vk_dsv4_hc_comb(ctx, compute_ctx, node); + + break; + + case GGML_OP_DSV4_HC_POST: + ggml_vk_dsv4_hc_post(ctx, compute_ctx, node); + + break; + case GGML_OP_SSM_SCAN: ggml_vk_ssm_scan(ctx, compute_ctx, node); @@ -19858,6 +20001,16 @@ static bool ggml_backend_vk_device_supports_op(ggml_backend_dev_t dev, const ggm case GGML_OP_GATED_LINEAR_ATTN: // the shader block size is hardcoded to head_size 64 return op->src[0]->type == GGML_TYPE_F32 && op->type == GGML_TYPE_F32 && op->src[0]->ne[0] == 64; + case GGML_OP_DSV4_HC_PRE: + return op->src[0]->type == GGML_TYPE_F32 && op->src[1]->type == GGML_TYPE_F32 && + op->type == GGML_TYPE_F32; + case GGML_OP_DSV4_HC_COMB: + return op->src[0]->type == GGML_TYPE_F32 && op->src[1]->type == GGML_TYPE_F32 && + op->src[2]->type == GGML_TYPE_F32 && op->type == GGML_TYPE_F32; + case GGML_OP_DSV4_HC_POST: + return op->src[0]->type == GGML_TYPE_F32 && op->src[1]->type == GGML_TYPE_F32 && + op->src[2]->type == GGML_TYPE_F32 && op->src[3]->type == GGML_TYPE_F32 && + op->type == GGML_TYPE_F32; case GGML_OP_LIGHTNING_INDEXER: { const ggml_tensor * q = op->src[0]; diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/dsv4_hc_comb.comp b/ggml/src/ggml-vulkan/vulkan-shaders/dsv4_hc_comb.comp new file mode 100644 index 00000000000..0449715bf45 --- /dev/null +++ b/ggml/src/ggml-vulkan/vulkan-shaders/dsv4_hc_comb.comp @@ -0,0 +1,101 @@ +#version 450 + +// DeepSeek V4 fused hyper-connection "comb": build the 4x4 stream-mixing matrix per token — +// per-source softmax over scaled+biased logits, then eps-stabilized alternating column/row +// (sinkhorn) normalization. Port of ggml-cuda/dsv4-hc.cu (hc_comb): one THREAD per token, +// the whole 4x4 lives in registers. At decode this is a single active thread by design — +// the fusion's value is collapsing the ~dozen decomposed graph ops (and their intermediate +// tensors) into one dispatch, not throughput. + +layout(local_size_x = 256, local_size_y = 1, local_size_z = 1) in; + +layout(binding = 0) readonly buffer MBuf { float data_mixes[]; }; +layout(binding = 1) readonly buffer SBuf { float data_scale[]; }; +layout(binding = 2) readonly buffer BBuf { float data_base[]; }; +layout(binding = 3) writeonly buffer DBuf { float data_d[]; }; + +layout(push_constant) uniform Parameters { + uint n_tokens; + uint sm0, sm1; + uint ss0; + uint sb0; + uint sd0, sd1, sd2; + float eps; + int n_iter; +} p; + +const uint HC = 4; +const uint COMB_OFFSET = 2 * HC; // comb logits start after the pre/post blocks of the mix vector + +void norm_cols(inout float comb[HC * HC]) { + for (uint idst = 0; idst < HC; ++idst) { + float sum = p.eps; + for (uint isrc = 0; isrc < HC; ++isrc) { + sum += comb[idst + HC * isrc]; + } + const float inv_sum = 1.0 / sum; + for (uint isrc = 0; isrc < HC; ++isrc) { + comb[idst + HC * isrc] *= inv_sum; + } + } +} + +void norm_rows(inout float comb[HC * HC]) { + for (uint isrc = 0; isrc < HC; ++isrc) { + float sum = p.eps; + for (uint idst = 0; idst < HC; ++idst) { + sum += comb[idst + HC * isrc]; + } + const float inv_sum = 1.0 / sum; + for (uint idst = 0; idst < HC; ++idst) { + comb[idst + HC * isrc] *= inv_sum; + } + } +} + +void main() { + const uint it = gl_GlobalInvocationID.x; + if (it >= p.n_tokens) { + return; + } + + const float scale_comb = data_scale[2 * p.ss0]; + float comb[HC * HC]; + + for (uint isrc = 0; isrc < HC; ++isrc) { + float vmax = uintBitsToFloat(0xff800000); // -inf + for (uint idst = 0; idst < HC; ++idst) { + const uint idx = idst + HC * isrc; + const float v = data_mixes[(COMB_OFFSET + idx) * p.sm0 + it * p.sm1] * scale_comb + + data_base[(COMB_OFFSET + idx) * p.sb0]; + comb[idx] = v; + vmax = max(vmax, v); + } + + float sum = 0.0; + for (uint idst = 0; idst < HC; ++idst) { + const uint idx = idst + HC * isrc; + const float v = exp(comb[idx] - vmax); + comb[idx] = v; + sum += v; + } + + const float inv_sum = 1.0 / sum; + for (uint idst = 0; idst < HC; ++idst) { + const uint idx = idst + HC * isrc; + comb[idx] = comb[idx] * inv_sum + p.eps; + } + } + + norm_cols(comb); + for (int i = 1; i < p.n_iter; ++i) { + norm_rows(comb); + norm_cols(comb); + } + + for (uint isrc = 0; isrc < HC; ++isrc) { + for (uint idst = 0; idst < HC; ++idst) { + data_d[idst * p.sd0 + isrc * p.sd1 + it * p.sd2] = comb[idst + HC * isrc]; + } + } +} diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/dsv4_hc_post.comp b/ggml/src/ggml-vulkan/vulkan-shaders/dsv4_hc_post.comp new file mode 100644 index 00000000000..212109dfd43 --- /dev/null +++ b/ggml/src/ggml-vulkan/vulkan-shaders/dsv4_hc_post.comp @@ -0,0 +1,44 @@ +#version 450 + +// DeepSeek V4 fused hyper-connection "post": redistribute the layer output back into the HC +// streams with the sinkhorn-mixed residual, +// dst[i0, idst, it] = x[i0, it] * post[idst, it] + sum_isrc residual[i0, isrc, it] * comb[idst, isrc, it]. +// Port of ggml-cuda/dsv4-hc.cu (hc_post). Flat elementwise over n_embd * hc * n_tokens. + +layout(local_size_x = 256, local_size_y = 1, local_size_z = 1) in; + +layout(binding = 0) readonly buffer XBuf { float data_x[]; }; +layout(binding = 1) readonly buffer RBuf { float data_r[]; }; +layout(binding = 2) readonly buffer PBuf { float data_p[]; }; +layout(binding = 3) readonly buffer CBuf { float data_c[]; }; +layout(binding = 4) writeonly buffer DBuf { float data_d[]; }; + +layout(push_constant) uniform Parameters { + uint n_embd; + uint hc; + uint nr; // n_embd * hc * n_tokens + uint sx0, sx1; + uint sr0, sr1, sr2; + uint sp0, sp1; + uint sc0, sc1, sc2; + uint sd0, sd1, sd2; +} p; + +void main() { + const uint ir = gl_GlobalInvocationID.x; + if (ir >= p.nr) { + return; + } + + const uint i0 = ir % p.n_embd; + const uint idst = (ir / p.n_embd) % p.hc; + const uint it = ir / (p.n_embd * p.hc); + + float sum = data_x[i0 * p.sx0 + it * p.sx1] * data_p[idst * p.sp0 + it * p.sp1]; + for (uint isrc = 0; isrc < p.hc; ++isrc) { + sum += data_r[i0 * p.sr0 + isrc * p.sr1 + it * p.sr2] + * data_c[idst * p.sc0 + isrc * p.sc1 + it * p.sc2]; + } + + data_d[i0 * p.sd0 + idst * p.sd1 + it * p.sd2] = sum; +} diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/dsv4_hc_pre.comp b/ggml/src/ggml-vulkan/vulkan-shaders/dsv4_hc_pre.comp new file mode 100644 index 00000000000..b6cc09fe2cf --- /dev/null +++ b/ggml/src/ggml-vulkan/vulkan-shaders/dsv4_hc_pre.comp @@ -0,0 +1,37 @@ +#version 450 + +// DeepSeek V4 fused hyper-connection "pre": mix the HC input streams down to one embedding, +// dst[i0, it] = sum_ih x[i0, ih, it] * w[ih, it]. Port of ggml-cuda/dsv4-hc.cu (hc_pre). +// Flat elementwise kernel over n_embd * n_tokens; strides are in f32 elements. + +layout(local_size_x = 256, local_size_y = 1, local_size_z = 1) in; + +layout(binding = 0) readonly buffer XBuf { float data_x[]; }; +layout(binding = 1) readonly buffer WBuf { float data_w[]; }; +layout(binding = 2) writeonly buffer DBuf { float data_d[]; }; + +layout(push_constant) uniform Parameters { + uint n_embd; + uint hc; + uint nr; // n_embd * n_tokens + uint sx0, sx1, sx2; + uint sw0, sw1; + uint sd0, sd1; +} p; + +void main() { + const uint ir = gl_GlobalInvocationID.x; + if (ir >= p.nr) { + return; + } + + const uint i0 = ir % p.n_embd; + const uint it = ir / p.n_embd; + + float sum = data_x[i0 * p.sx0 + it * p.sx2] * data_w[it * p.sw1]; + for (uint ih = 1; ih < p.hc; ++ih) { + sum += data_x[i0 * p.sx0 + ih * p.sx1 + it * p.sx2] * data_w[ih * p.sw0 + it * p.sw1]; + } + + data_d[i0 * p.sd0 + it * p.sd1] = sum; +} 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 5c06f822292..3232278e588 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp @@ -818,6 +818,9 @@ void process_shaders() { #endif string_to_spv("flash_attn_top_k_f16", "flash_attn_top_k.comp", {}); string_to_spv("flash_attn_gather_f16", "flash_attn_gather.comp", {}); + string_to_spv("dsv4_hc_pre_f32", "dsv4_hc_pre.comp", {}); + string_to_spv("dsv4_hc_comb_f32", "dsv4_hc_comb.comp", {}); + string_to_spv("dsv4_hc_post_f32", "dsv4_hc_post.comp", {}); string_to_spv("mul_mat_vec_p021_f16_f32_subgroup_add", "mul_mat_vec_p021.comp", {{"A_TYPE", "float16_t"}, {"A_TYPEV4", "f16vec4"}, {"B_TYPE", "float"}, {"B_TYPEV4", "vec4"}, {"D_TYPE", "float"}, {"USE_SUBGROUP_ADD", "1"}}); string_to_spv("mul_mat_vec_p021_f16_f32", "mul_mat_vec_p021.comp", {{"A_TYPE", "float16_t"}, {"A_TYPEV4", "f16vec4"}, {"B_TYPE", "float"}, {"B_TYPEV4", "vec4"}, {"D_TYPE", "float"}}); From b8d2503264dfd6d13d12c5c56b3ffdd979d01c19 Mon Sep 17 00:00:00 2001 From: Nathan Wilson Date: Sun, 2 Aug 2026 16:07:19 +0000 Subject: [PATCH 15/68] llama: keep DeepSeek lightning-indexer key cache f16 under quantized -ctk The fused indexer kernels read f16 keys only, so quantizing the small (128-dim) indexer cache silently disables them and falls back to the decomposed full-KV indexer path, whose contiguize cost grows superlinearly with depth (measured 0.8ms -> 97ms per dispatch by 12k context on Vulkan). Pin the indexer key cache to f16 in both DSA cache variants; the memory cost vs q8_0 is ~120 bytes per token per layer. Assisted-by: Claude Fable 5 --- src/llama-kv-cache-dsa.cpp | 5 ++++- src/llama-kv-cache-dsv4.cpp | 5 ++++- 2 files changed, 8 insertions(+), 2 deletions(-) diff --git a/src/llama-kv-cache-dsa.cpp b/src/llama-kv-cache-dsa.cpp index 96cb045d2e5..e926e34314b 100644 --- a/src/llama-kv-cache-dsa.cpp +++ b/src/llama-kv-cache-dsa.cpp @@ -47,8 +47,11 @@ llama_kv_cache_dsa::llama_kv_cache_dsa( LLAMA_LOG_INFO("%s: creating indexer KV cache, size = %u cells\n", __func__, kv_size); + // keep indexer keys f16 regardless of type_k: the fused indexer kernels read + // f16 only, and quantizing this small cache (128 dims) saves little while + // forcing the much slower decomposed indexer path kv_lid = std::make_unique( - model, hparams_lid, type_k, type_v, + model, hparams_lid, GGML_TYPE_F16, type_v, v_trans, offload, unified, kv_size, n_seq_max, n_pad, n_swa, swa_type, nullptr, filter_lid, reuse, nullptr); } diff --git a/src/llama-kv-cache-dsv4.cpp b/src/llama-kv-cache-dsv4.cpp index 9e82a6198a0..2097271a0f1 100644 --- a/src/llama-kv-cache-dsv4.cpp +++ b/src/llama-kv-cache-dsv4.cpp @@ -1305,8 +1305,11 @@ llama_kv_cache_dsv4::llama_kv_cache_dsv4( LLAMA_LOG_INFO("%s: creating DSV4 lightning-indexer KV cache, size = %u cells\n", __func__, dsv4_comp_size(kv_size, DSV4_CSA_RATIO)); + // keep indexer keys f16 regardless of type_k: the fused indexer kernels read + // f16 only, and quantizing this small cache (128 dims) saves little while + // forcing the much slower decomposed indexer path kv_lid = std::make_unique( - model, hparams_lid, type_k, type_v, + model, hparams_lid, GGML_TYPE_F16, type_v, v_trans, offload, unified_compressed, GGML_PAD(dsv4_comp_size(kv_size, DSV4_CSA_RATIO), 256u), n_seq_max, n_pad, 0, LLAMA_SWA_TYPE_NONE, nullptr, filter_csa, nullptr, nullptr); From 9cb8c288889ac32765211bf3026b262e36010459 Mon Sep 17 00:00:00 2001 From: Nathan Wilson Date: Mon, 3 Aug 2026 01:48:02 +0000 Subject: [PATCH 16/68] llama: contiguize grouped o-proj input for small multi-token batches (DSv4) The permuted view feeding the grouped wo_a matmul gives B a 128 KiB power-of-2 token stride. Single-token decode never touches it, but small multi-token batches (speculative verify, n=2-4) hit a strided-B matmul path that runs ~11x slower than contiguous (40 vs 460 GFLOPS measured on gfx1151, ~54% of GPU time in a draft-verify window). A contiguous copy for 2..8 tokens is far cheaper than the stride tax; n=1 and large prefill batches are unaffected. Assisted-by: Claude Fable 5 --- src/models/deepseek4.cpp | 5 +++++ 1 file changed, 5 insertions(+) diff --git a/src/models/deepseek4.cpp b/src/models/deepseek4.cpp index 8588d069eed..494e643276c 100644 --- a/src/models/deepseek4.cpp +++ b/src/models/deepseek4.cpp @@ -1209,6 +1209,11 @@ ggml_tensor * llama_model_deepseek4::graph::build_attention_impl( out = ggml_reshape_3d(ctx0, out, o_group_dim, n_groups, nt); out = ggml_permute(ctx0, out, 0, 2, 1, 3); + // small multi-token batches (speculative verify) hit a pathological strided-B + // path in the grouped matmul below; a contiguous copy is much cheaper + if (nt > 1 && nt <= 8) { + out = ggml_cont(ctx0, out); + } ggml_tensor * oa = ggml_mul_mat(ctx0, layer.wo_a, out); cb(oa, "attn_wo_a", il); oa = ggml_permute(ctx0, oa, 0, 2, 1, 3); From 78e31af07176d46c6792b08bdba3d0b4f4b574ec Mon Sep 17 00:00:00 2001 From: Jaap Buurman Date: Wed, 12 Aug 2026 22:25:40 +0200 Subject: [PATCH 17/68] vulkan: accelerate DeepSeek V4 sparse prefill FA Assisted-by: OpenAI Codex --- .../DSV4-vulkan-sparse-prefill-progress.md | 259 +++++++++++++++++ ggml/src/ggml-vulkan/ggml-vulkan.cpp | 16 +- .../vulkan-shaders/flash_attn_top_k_cm.comp | 269 ++++++++++++++++++ .../vulkan-shaders/vulkan-shaders-gen.cpp | 3 + tests/test-backend-ops.cpp | 1 + 5 files changed, 545 insertions(+), 3 deletions(-) create mode 100644 docs/development/DSV4-vulkan-sparse-prefill-progress.md create mode 100644 ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_top_k_cm.comp diff --git a/docs/development/DSV4-vulkan-sparse-prefill-progress.md b/docs/development/DSV4-vulkan-sparse-prefill-progress.md new file mode 100644 index 00000000000..80bf562a062 --- /dev/null +++ b/docs/development/DSV4-vulkan-sparse-prefill-progress.md @@ -0,0 +1,259 @@ +# DeepSeek V4 Vulkan sparse prefill progress + +This file is a self-contained handoff for the DeepSeek V4 sparse-attention prompt-processing optimization on AMD Strix Halo. Read the repository `AGENTS.md` and `CONTRIBUTING.md` before continuing. + +## Repository state + +- Repository: `https://github.com/Nathanw1014/llama.cpp` +- Branch: `strix-halo-vulkan` +- Starting commit: `baf0025de861c6f6ea3720fa81c52ae1b2e6c078` +- Target GPU: AMD Radeon 8060S / gfx1151, RADV, Vulkan, wave64 +- The device reports `GL_KHR_cooperative_matrix`, f16 inputs with f32 accumulation, 64 KiB shared memory, and a maximum 512-thread workgroup used by this path. + +The implementation is not upstream `ggml-org/llama.cpp`. It builds on this branch's DeepSeek V4 graph, Lightning Indexer, sparse top-K hint, decode gather path, fused HC kernels, and Vulkan profiler changes. + +## Important execution constraints + +The system is an APU. CPU compilation and GPU benchmarking share power and memory bandwidth. Never build and benchmark at the same time. Serialize all builds, correctness tests, and performance tests. + +GPU commands must run with host GPU access. In an agent sandbox, request elevated/out-of-sandbox execution. A sandboxed benchmark showed only CPU activity and is invalid. + +Redirect the full model benchmark to a log. Inspect only the last profiler block with `tail`; do not load the full log into agent context. + +## Build and ccache + +The build directory is `build`, configured as Release with Vulkan enabled. The default ccache directory was read-only in the agent environment, so use a writable directory: + +```bash +cmake -S . -B build \ + -DGGML_VULKAN=ON \ + -DCMAKE_BUILD_TYPE=Release \ + -DCMAKE_C_COMPILER_LAUNCHER=ccache \ + -DCMAKE_CXX_COMPILER_LAUNCHER=ccache + +CCACHE_DIR=/tmp/llama-cpp-ccache cmake --build build --config Release \ + --target llama-bench test-backend-ops -j "$(nproc)" + +CCACHE_DIR=/tmp/llama-cpp-ccache ccache -s +``` + +ccache was verified active. The final build reported direct hits. Note that an incremental change to `ggml-vulkan.cpp` is one large C++ translation unit and therefore uses one compiler core even with `-j`. Shader object regeneration can run in parallel. + +## Canonical benchmark command + +Do not change or omit switches for the 32k acceptance run: + +```bash +GGML_VK_PERF_LOGGER=1 ./build/bin/llama-bench \ + -m ~/Projects/docker/localLLaMA/models/models--unsloth--DeepSeek-V4-Flash-0731-GGUF/snapshots/109848da2469efe1f1aab9e11acea08a065ccd4f/UD-IQ3_XXS/DeepSeek-V4-Flash-0731-UD-IQ3_XXS-00001-of-00004.gguf \ + -r 1 -d 32768 -p 2048 -ub 2048 -fa 1 -n 0 \ + > /tmp/dsv4-vulkan-32k.log 2>&1 + +tail -n 180 /tmp/dsv4-vulkan-32k.log +``` + +The known model path exists on the target system. + +## Graph and dispatch findings + +DeepSeek V4 builds the Lightning Indexer and sparse attention in `src/models/deepseek4.cpp`: + +1. `build_lid_top_k()` creates indexer Q/K/weights and calls `ggml_lightning_indexer()`. +2. `ggml_top_k()` selects up to `hparams.indexer_top_k` compressed-cache indices for every query token. +3. `build_csa_lid_attention()` concatenates the raw SWA K prefix with compressed CSA K, builds a dense mask carrying the same sparse selection, and calls `build_attn_mha(..., top_k, raw_k->ne[2])`. +4. `build_attn_mha()` attaches `top_k` and `n_kv_raw` to `GGML_OP_FLASH_ATTN_EXT` through `ggml_flash_attn_ext_add_top_k()`. + +At the final 32k PP2048 batch: + +- total K rows: 11,008 +- `n_kv_raw`: 2,304 raw SWA rows, always attended subject to the mask +- `n_top_k`: 512 selected compressed rows per query token +- active rows: 2,816 +- selectable compressed region: 8,704 + +The custom Vulkan path is selected by `ggml_vk_flash_attn_top_k()` before ordinary FA. Its gate requires the DeepSeek V4 shape and `total_k >= 3 * (n_kv_raw + n_top_k)`. The final shape satisfies `11008 >= 3 * 2816`. + +The old `flash_attn_top_k.comp` shader is scalar/subgroup code. One 512-thread workgroup covers eight heads for one query token. It stages 16 selected 512-wide K/V rows, computes QK with scalar FMAs and `subgroupAdd`, updates online softmax one key at a time, and accumulates PV manually. It does not use cooperative matrices. + +The top-K set differs by query token but is shared by all 64 query heads for that token. This makes the attention for one token a regular matrix problem across heads and selected keys despite sparse per-token indexing. + +## Root cause evidence + +The existing Vulkan timestamp infrastructure was extended with `ggml_vk_perf_mark_subop()` after the sparse dispatch. This reports the sparse kernel separately as `FA_TOP_K_SPARSE (sub-op)` or `FA_TOP_K_CM (sub-op)`. + +Focused test shape: + +```bash +GGML_VK_PERF_LOGGER=1 ./build/bin/test-backend-ops perf \ + -b Vulkan0 -o FLASH_ATTN_EXT \ + -p 'kv=32768,nb=512,n_kv_raw=1024,n_top_k=512,sinks=0' +``` + +Results: + +- old scalar sparse kernel: 61.42 ms, 1.68 TFLOPS of useful active-set work +- ordinary dense FA diagnostic (`GGML_VK_FA_TOPK=0`): 255.97 ms, about 8.7 TFLOPS over the full dense work +- final cooperative sparse kernel: 32.55 ms, 3.17 TFLOPS of useful active-set work + +The residual `FLASH_ATTN_EXT` interval after the sparse timestamp is only about 4-7 us. The cost is inside the shader, not dispatch or surrounding synchronization. + +A temporary uniform stage-profiling mode was used and removed. For the final 32-head tile: + +- selected K gather + cooperative QK: 14.30 ms +- gather + QK + serial softmax: 26.34 ms +- gather + QK + parallel softmax: 15.53 ms +- full kernel: 32.55 ms +- the remaining cooperative PV/output portion is about 17.0 ms + +The old scalar shader was compute/issue inefficient. Dense FA proved matrix hardware is much faster but was still too expensive because it processes all K rows. The final implementation preserves sparsity and uses the matrix hardware for both QK and PV. + +## Implementation + +Files changed: + +- `ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_top_k_cm.comp` + - New cooperative-matrix sparse prefill shader. + - One 512-thread workgroup covers 32 query heads for one token. + - Eight wave64 subgroups cover two 16-head tiles by four 16-key or 16-output-dimension tiles. + - Processes 64 selected keys per online-softmax block. + - Stages only indexed selected K/V tiles, never the full K range. + - Uses f16 cooperative-matrix inputs and f32 accumulation for QK and PV. + - Uses a 16-lane segmented softmax per head. XOR subgroup shuffles reduce max and sum for four independent heads per wave without workgroup barriers. + - Keeps f32 output accumulators and normalizes after all active blocks. + - Preserves the raw prefix, top-K index validation, mask, sinks, stream strides, and K == V latent behavior. +- `ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp` + - Embeds the new shader when cooperative-matrix shader support is available. +- `ggml/src/ggml-vulkan/ggml-vulkan.cpp` + - Adds the cooperative sparse pipeline when the device supports the required 16x16x16 f16/f32 cooperative matrix shape. + - Selects it by capability and keeps the scalar shader as fallback. + - Adds `GGML_VK_FA_TOPK=0` to force ordinary dense FA for diagnostics. + - Adds `GGML_VK_FA_TOPK_CM=0` to force the old scalar sparse shader for A/B tests. + - Adds sparse sub-operation timestamps through the existing profiler. + +The path is capability-based, not hardcoded to Strix Halo. The current sparse shape gate remains DeepSeek V4-specific. Devices without the required cooperative matrix support keep the correct scalar sparse or dense fallback. + +Two discarded prototypes are useful context: + +- A 16-head, 64-key streaming cooperative tile was correct and reduced the focused test from 61.4 to 44.6 ms. +- Keeping 32 complete 512-wide K/V rows in LDS grew shared memory to about 41 KiB, reduced residency, doubled block/barrier count, and regressed to 73.8 ms. Do not retry full-row LDS staging without solving occupancy. +- A 32-head, 64-key streaming tile halved irregular row loads but initially stayed near 44.7 ms because serial softmax cost about 11.5 ms. Parallel segmented softmax produced the final 32.55 ms result. + +## Correctness validation + +Run: + +```bash +./build/bin/test-backend-ops test -b Vulkan0 -o FLASH_ATTN_EXT -p 'n_top_k=' +``` + +Final result: 8/8 sparse top-K FA cases passed against the CPU reference. Cases include: + +- decode and short batches that use dense/gather fallback +- sparse prefill batch sizes 64 and 128 +- `n_kv_raw` plus top-K selection +- invalid top-K index handling from the test fixture +- sinks enabled and disabled +- sparse threshold transitions +- an active-key count of 193, which exercises a partial final 64-key block + +The test uses the existing FA tolerance of NMSE <= `5e-4`. No NaN or Inf failure occurred. The implementation changes Q and probability inputs to f16 cooperative-matrix operands with f32 accumulation, matching the precision strategy of ordinary Vulkan cooperative FA. + +Still desirable before broader submission: + +- compare model logits on controlled prompts between `GGML_VK_FA_TOPK_CM=0` and the default cooperative path + +## Canonical 32k results + +Exact clean runs, same command and machine, no concurrent build: + +```text +32k context, PP 2048, ub 2048 + +Before (commit baf0025de): +112.29 tok/s +Total Vulkan: 18.1957 s +Sparse FA: 8.84496 s, 421.189 ms/layer +Lightning Indexer: 1.19438 s +TOP_K: 0.076948 s + +After: +152.32 tok/s +Total Vulkan: 13.4032 s +Sparse FA: 4.44701 s, 211.762 ms/layer +Lightning Indexer: 1.13156 s +TOP_K: 0.073945 s + +Change: +Throughput: +35.65% +Total Vulkan time: -26.34% +Sparse FA time: -49.72% +Sparse FA saved: 4.398 s +Total GPU time saved: 4.793 s +``` + +The profiler now lists the optimized dispatch as `FA_TOP_K_CM (sub-op)`. The following residual `FLASH_ATTN_EXT` line is only the post-mark interval and must not be interpreted as the kernel time. + +## Context-depth measurements + +All points use PP2048, ub2048, FA enabled, one repetition, and no token generation. They were run sequentially with no compiler active: + +```text +Existing depth tok/s Total Vulkan Final large FA Lightning Indexer TOP_K +0 253.44 8.041 s 0.294 s 0.095 s 0.001 s +8192 211.01 9.666 s 1.625 s 0.379 s 0.022 s +16384 177.19 11.518 s 3.031 s 0.678 s 0.044 s +32768 152.32 13.403 s 4.447 s 1.132 s 0.074 s +``` + +At 0, 8k, and 16k, total K is below the existing sparse-path gate `total_k >= 3 * (n_kv_raw + n_top_k)`. These points use the unchanged ordinary dense FA implementation, so the cooperative sparse change does not affect or regress them. At 32k, total K is 11,008 and the cooperative sparse path engages. The 32k `Final large FA` value is the `FA_TOP_K_CM (sub-op)` total; the lower-depth values are the large ordinary `FLASH_ATTN_EXT` totals. + +Logs: + +- `/tmp/dsv4-vulkan-cm-0k.log` +- `/tmp/dsv4-vulkan-cm-8k.log` +- `/tmp/dsv4-vulkan-cm-16k.log` +- `/tmp/dsv4-vulkan-cm-32k.log` +- `/tmp/dsv4-vulkan-baseline-clean.log` +- `/tmp/dsv4-fa-cm-correctness-final.log` + +## Next optimization target + +The cooperative sparse FA remains the largest context-dependent cost at about 4.45 s total. Stage profiling indicates approximately 14.3 ms of focused-test time in gather/QK, about 1.2 ms in parallel softmax, and about 17 ms in PV/output. + +The next useful work is PV and output accumulation, not TOP_K. Investigate: + +- reducing repeated selected V staging across the two 32-head workgroups per token without increasing LDS enough to lose occupancy +- reducing the eight output-dimension passes or retaining more PV state in cooperative fragments/registers +- checking register count and spills for the 32 f32 output accumulators per invocation using RADV shader statistics +- alternate 32-head layouts that keep the same eight-wave occupancy but improve PV scheduling +- query-tile overlap/union gathering only if measured top-K overlap is high enough; a whole-2048-query union is unlikely to help + +Do not optimize TOP_K first. At 32k it is only about 74 ms total. Lightning Indexer is about 1.13 s and is the next context-dependent target only after sparse FA improves further. + +## Useful diagnostics + +Force old scalar sparse path: + +```bash +GGML_VK_FA_TOPK_CM=0 GGML_VK_PERF_LOGGER=1 ./build/bin/test-backend-ops perf \ + -b Vulkan0 -o FLASH_ATTN_EXT \ + -p 'kv=32768,nb=512,n_kv_raw=1024,n_top_k=512,sinks=0' +``` + +Force ordinary dense FA: + +```bash +GGML_VK_FA_TOPK=0 GGML_VK_PERF_LOGGER=1 ./build/bin/test-backend-ops perf \ + -b Vulkan0 -o FLASH_ATTN_EXT \ + -p 'kv=32768,nb=512,n_kv_raw=1024,n_top_k=512,sinks=0' +``` + +Default cooperative sparse path: + +```bash +GGML_VK_PERF_LOGGER=1 ./build/bin/test-backend-ops perf \ + -b Vulkan0 -o FLASH_ATTN_EXT \ + -p 'kv=32768,nb=512,n_kv_raw=1024,n_top_k=512,sinks=0' +``` + +Always run these sequentially. Do not run a compiler concurrently on this APU. diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp index 8a2d6fe9954..6c2dcef824b 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp +++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp @@ -1108,6 +1108,7 @@ struct vk_device_struct { vk_pipeline pipeline_lightning_indexer_cm_f16; vk_pipeline pipeline_lightning_indexer_decode_cm_f16; vk_pipeline pipeline_flash_attn_top_k_f16; + vk_pipeline pipeline_flash_attn_top_k_cm_f16; vk_pipeline pipeline_flash_attn_gather_f16; vk_pipeline pipeline_dsv4_hc_pre_f32; vk_pipeline pipeline_dsv4_hc_comb_f32; @@ -6286,6 +6287,10 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { "lightning_indexer_decode_cm_f16", lightning_indexer_decode_cm_f16_len, lightning_indexer_decode_cm_f16_data, "main", 5, sizeof(vk_op_lightning_indexer_cm_push_constants), {16, 1, 1}, {device->subgroup_size}, 1, true, true, device->subgroup_size); + ggml_vk_create_pipeline(device, device->pipeline_flash_attn_top_k_cm_f16, + "flash_attn_top_k_cm_f16", flash_attn_top_k_cm_f16_len, flash_attn_top_k_cm_f16_data, "main", 6, + sizeof(vk_op_flash_attn_top_k_push_constants), {1, 1, 1}, {512, device->subgroup_size}, 1, true, true, + device->subgroup_size); } #endif ggml_vk_create_pipeline(device, device->pipeline_flash_attn_top_k_f16, @@ -11425,7 +11430,9 @@ static bool ggml_vk_flash_attn_top_k(ggml_backend_vk_context * ctx, vk_context & const ggml_tensor * q, const ggml_tensor * k, const ggml_tensor * v, const ggml_tensor * mask, const ggml_tensor * sinks, ggml_tensor * dst) { const ggml_tensor * top_k = dst->src[5]; - if (!top_k || !ctx->device->pipeline_flash_attn_top_k_f16 || + static const char * top_k_env = getenv("GGML_VK_FA_TOPK"); + if ((top_k_env && top_k_env[0] == '0') || + !top_k || (!ctx->device->pipeline_flash_attn_top_k_f16 && !ctx->device->pipeline_flash_attn_top_k_cm_f16) || q->type != GGML_TYPE_F32 || k->type != GGML_TYPE_F16 || v->type != GGML_TYPE_F16 || !mask || mask->type != GGML_TYPE_F16 || top_k->type != GGML_TYPE_I32 || q->ne[0] != 512 || q->ne[1] < 64 || k->ne[0] != 512 || v->ne[0] != 512 || @@ -11475,12 +11482,15 @@ static bool ggml_vk_flash_attn_top_k(ggml_backend_vk_context * ctx, vk_context & const vk_subbuffer q_buf = ggml_vk_tensor_subbuffer(ctx, q); const vk_subbuffer sinks_buf = sinks ? ggml_vk_tensor_subbuffer(ctx, sinks) : q_buf; - vk_pipeline pipeline = ctx->device->pipeline_flash_attn_top_k_f16; + static const char * top_k_cm_env = getenv("GGML_VK_FA_TOPK_CM"); + const bool use_cm = (!top_k_cm_env || top_k_cm_env[0] != '0') && ctx->device->pipeline_flash_attn_top_k_cm_f16; + vk_pipeline pipeline = use_cm ? ctx->device->pipeline_flash_attn_top_k_cm_f16 : ctx->device->pipeline_flash_attn_top_k_f16; ggml_pipeline_request_descriptor_sets(ctx, pipeline, 1); ggml_vk_dispatch_pipeline(ctx, subctx, pipeline, {q_buf, ggml_vk_tensor_subbuffer(ctx, k), ggml_vk_tensor_subbuffer(ctx, mask), sinks_buf, ggml_vk_tensor_subbuffer(ctx, top_k), ggml_vk_tensor_subbuffer(ctx, dst)}, - pc, {(uint32_t) q->ne[1], (uint32_t) CEIL_DIV(q->ne[2], 8), (uint32_t) q->ne[3]}); + pc, {(uint32_t) q->ne[1], (uint32_t) CEIL_DIV(q->ne[2], use_cm ? 32 : 8), (uint32_t) q->ne[3]}); + ggml_vk_perf_mark_subop(ctx, subctx, use_cm ? "FA_TOP_K_CM (sub-op)" : "FA_TOP_K_SPARSE (sub-op)"); return true; } diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_top_k_cm.comp b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_top_k_cm.comp new file mode 100644 index 00000000000..420a05375a7 --- /dev/null +++ b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_top_k_cm.comp @@ -0,0 +1,269 @@ +#version 450 + +#extension GL_EXT_control_flow_attributes : require +#extension GL_EXT_shader_16bit_storage : require +#extension GL_EXT_shader_explicit_arithmetic_types_float16 : require +#extension GL_KHR_cooperative_matrix : require +#extension GL_KHR_memory_scope_semantics : require +#extension GL_KHR_shader_subgroup_basic : require +#extension GL_KHR_shader_subgroup_shuffle : require + +layout(constant_id = 0) const uint WORKGROUP_SIZE = 512; +layout(constant_id = 1) const uint SUBGROUP_SIZE = 64; +layout(local_size_x_id = 0, local_size_y = 1, local_size_z = 1) in; + +layout(binding = 0) readonly buffer QBuf { float data_q[]; }; +layout(binding = 1) readonly buffer KBuf { float16_t data_k[]; }; +layout(binding = 2) readonly buffer MaskBuf { float16_t data_m[]; }; +layout(binding = 3) readonly buffer SinkBuf { float data_s[]; }; +layout(binding = 4) readonly buffer TopBuf { int data_top[]; }; +layout(binding = 5) writeonly buffer DstBuf { float data_dst[]; }; + +layout(push_constant) uniform Parameters { + uint n_batch; + uint n_kv; + uint n_kv_raw; + uint n_top_k; + uint n_head; + uint nbq1; + uint nbq2; + uint nbq3; + uint nbk1; + uint nbk3; + uint nbm1; + uint nbm3; + uint nbt1; + uint nbt3; + uint nb1; + uint nb2; + uint nb3; + float scale; + uint has_sinks; +} p; + +const uint TILE = 16; +const uint HEAD_SIZE = 512; +const uint HEADS_PER_GROUP = 32; +const uint KEYS_PER_BLOCK = 64; +const uint DIMS_PER_BLOCK = 64; +const uint QK_STRIDE = TILE / 4 + 2; +const uint SCORE_STRIDE = HEADS_PER_GROUP / 4 + 1; +const uint P_STRIDE = KEYS_PER_BLOCK / 4 + 2; +const uint V_STRIDE = DIMS_PER_BLOCK / 4 + 2; +const uint PV_STRIDE = DIMS_PER_BLOCK / 4; +const float MASK_NEG_INF = -65500.0; + +shared uint key_idx[KEYS_PER_BLOCK]; +shared f16vec4 q_sh[HEADS_PER_GROUP * QK_STRIDE]; +shared f16vec4 k_sh[KEYS_PER_BLOCK * QK_STRIDE]; +shared vec4 score_sh[KEYS_PER_BLOCK * SCORE_STRIDE]; +shared f16vec4 p_sh[HEADS_PER_GROUP * P_STRIDE]; +shared f16vec4 v_sh[KEYS_PER_BLOCK * V_STRIDE]; +shared vec4 pv_sh[HEADS_PER_GROUP * PV_STRIDE]; +shared float old_scale_sh[HEADS_PER_GROUP]; +shared float row_max_sh[HEADS_PER_GROUP]; +shared float row_sum_sh[HEADS_PER_GROUP]; + +void main() { + const uint tid = gl_LocalInvocationIndex; + const uint token = gl_WorkGroupID.x; + const uint head_base = gl_WorkGroupID.y * HEADS_PER_GROUP; + const uint stream = gl_WorkGroupID.z; + const uint mask_base = stream * p.nbm3 + token * p.nbm1; + const uint top_base = stream * p.nbt3 + token * p.nbt1; + const uint total_keys = p.n_kv_raw + p.n_top_k; + + float accum[HEADS_PER_GROUP * HEAD_SIZE / WORKGROUP_SIZE]; + [[unroll]] for (uint i = 0; i < accum.length(); ++i) { + accum[i] = 0.0; + } + + if (tid < HEADS_PER_GROUP) { + row_max_sh[tid] = uintBitsToFloat(0xff800000); + row_sum_sh[tid] = 0.0; + } + barrier(); + + for (uint kb = 0; kb < total_keys; kb += KEYS_PER_BLOCK) { + if (tid < KEYS_PER_BLOCK) { + const uint selected = kb + tid; + uint key = p.n_kv; + if (selected < p.n_kv_raw) { + key = selected; + } else if (selected < total_keys) { + const int compressed = data_top[top_base + selected - p.n_kv_raw]; + if (compressed >= 0 && uint(compressed) < p.n_kv - p.n_kv_raw) { + key = p.n_kv_raw + uint(compressed); + } + } + key_idx[tid] = key; + } + barrier(); + + coopmat scores = + coopmat(0.0); + coopmat kmat; + coopmat qmat; + + [[unroll]] for (uint d = 0; d < HEAD_SIZE; d += TILE) { + if (tid < KEYS_PER_BLOCK * (TILE / 4)) { + const uint key_local = tid / (TILE / 4); + const uint d4 = tid % (TILE / 4); + const uint key = key_idx[key_local]; + f16vec4 value = f16vec4(0.0); + if (key < p.n_kv) { + const uint offset = stream * p.nbk3 + key * p.nbk1 + d + d4 * 4; + value = f16vec4(data_k[offset], data_k[offset + 1], data_k[offset + 2], data_k[offset + 3]); + } + k_sh[key_local * QK_STRIDE + d4] = value; + } + if (tid < HEADS_PER_GROUP * (TILE / 4)) { + const uint head_local = tid / (TILE / 4); + const uint d4 = tid % (TILE / 4); + const uint head = head_base + head_local; + const uint offset = stream * p.nbq3 + head * p.nbq2 + token * p.nbq1 + d + d4 * 4; + q_sh[head_local * QK_STRIDE + d4] = f16vec4( + data_q[offset], data_q[offset + 1], data_q[offset + 2], data_q[offset + 3]); + } + barrier(); + + const uint key_chunk = gl_SubgroupID % (KEYS_PER_BLOCK / TILE); + const uint head_tile = gl_SubgroupID / (KEYS_PER_BLOCK / TILE); + coopMatLoad(kmat, k_sh, key_chunk * TILE * QK_STRIDE, + QK_STRIDE, gl_CooperativeMatrixLayoutRowMajor); + coopMatLoad(qmat, q_sh, head_tile * TILE * QK_STRIDE, + QK_STRIDE, gl_CooperativeMatrixLayoutColumnMajor); + scores = coopMatMulAdd(kmat, qmat, scores); + barrier(); + } + + const uint score_key_chunk = gl_SubgroupID % (KEYS_PER_BLOCK / TILE); + const uint score_head_tile = gl_SubgroupID / (KEYS_PER_BLOCK / TILE); + coopMatStore(scores, score_sh, + score_key_chunk * TILE * SCORE_STRIDE + score_head_tile * (TILE / 4), + SCORE_STRIDE, gl_CooperativeMatrixLayoutRowMajor); + barrier(); + + { + const uint head_local = tid / (SUBGROUP_SIZE / 4); + const uint softmax_lane = tid % (SUBGROUP_SIZE / 4); + float block_max = uintBitsToFloat(0xff800000); + [[unroll]] for (uint i = 0; i < KEYS_PER_BLOCK / (SUBGROUP_SIZE / 4); ++i) { + const uint key_local = softmax_lane + i * (SUBGROUP_SIZE / 4); + const uint key = key_idx[key_local]; + const float mask = key < p.n_kv ? float(data_m[mask_base + key]) : uintBitsToFloat(0xff800000); + const float score = float(score_sh[key_local * SCORE_STRIDE + head_local / 4][head_local % 4]) * p.scale + mask; + block_max = mask < MASK_NEG_INF ? block_max : max(block_max, score); + } + [[unroll]] for (uint delta = 1; delta < SUBGROUP_SIZE / 4; delta *= 2) { + block_max = max(block_max, subgroupShuffleXor(block_max, delta)); + } + + const float old_row_max = row_max_sh[head_local]; + const float old_row_sum = row_sum_sh[head_local]; + const float new_max = max(old_row_max, block_max); + const float old_scale = old_row_sum == 0.0 ? 0.0 : exp(old_row_max - new_max); + float block_sum = 0.0; + [[unroll]] for (uint i = 0; i < KEYS_PER_BLOCK / (SUBGROUP_SIZE / 4); ++i) { + const uint key_local = softmax_lane + i * (SUBGROUP_SIZE / 4); + const uint key = key_idx[key_local]; + const float mask = key < p.n_kv ? float(data_m[mask_base + key]) : uintBitsToFloat(0xff800000); + float weight = 0.0; + if (mask >= MASK_NEG_INF) { + const float score = float(score_sh[key_local * SCORE_STRIDE + head_local / 4][head_local % 4]) * p.scale + mask; + weight = exp(score - new_max); + block_sum += weight; + } + p_sh[head_local * P_STRIDE + key_local / 4][key_local % 4] = float16_t(weight); + } + [[unroll]] for (uint delta = 1; delta < SUBGROUP_SIZE / 4; delta *= 2) { + block_sum += subgroupShuffleXor(block_sum, delta); + } + + if (softmax_lane == 0) { + row_sum_sh[head_local] = old_row_sum * old_scale + block_sum; + row_max_sh[head_local] = new_max; + old_scale_sh[head_local] = old_scale; + } + } + barrier(); + + [[unroll]] for (uint i = 0; i < accum.length(); ++i) { + const uint out_idx = tid + i * WORKGROUP_SIZE; + const uint head_local = out_idx / HEAD_SIZE; + accum[i] *= old_scale_sh[head_local]; + } + + [[unroll]] for (uint dim_base = 0; dim_base < HEAD_SIZE; dim_base += DIMS_PER_BLOCK) { + [[unroll]] for (uint idx = tid; idx < KEYS_PER_BLOCK * (DIMS_PER_BLOCK / 4); idx += WORKGROUP_SIZE) { + const uint key_local = idx / (DIMS_PER_BLOCK / 4); + const uint d4 = idx % (DIMS_PER_BLOCK / 4); + const uint key = key_idx[key_local]; + f16vec4 value = f16vec4(0.0); + if (key < p.n_kv) { + const uint offset = stream * p.nbk3 + key * p.nbk1 + dim_base + d4 * 4; + value = f16vec4(data_k[offset], data_k[offset + 1], data_k[offset + 2], data_k[offset + 3]); + } + v_sh[key_local * V_STRIDE + d4] = value; + } + barrier(); + + coopmat pv = + coopmat(0.0); + coopmat pmat; + coopmat vmat; + + const uint pv_head_tile = gl_SubgroupID / (DIMS_PER_BLOCK / TILE); + const uint pv_dim_tile = gl_SubgroupID % (DIMS_PER_BLOCK / TILE); + [[unroll]] for (uint key_chunk = 0; key_chunk < KEYS_PER_BLOCK / TILE; ++key_chunk) { + coopMatLoad(pmat, p_sh, + pv_head_tile * TILE * P_STRIDE + key_chunk * (TILE / 4), + P_STRIDE, gl_CooperativeMatrixLayoutRowMajor); + coopMatLoad(vmat, v_sh, + key_chunk * TILE * V_STRIDE + pv_dim_tile * (TILE / 4), + V_STRIDE, gl_CooperativeMatrixLayoutRowMajor); + pv = coopMatMulAdd(pmat, vmat, pv); + } + + coopMatStore(pv, pv_sh, + pv_head_tile * TILE * PV_STRIDE + pv_dim_tile * (TILE / 4), PV_STRIDE, + gl_CooperativeMatrixLayoutRowMajor); + barrier(); + + [[unroll]] for (uint i = 0; i < accum.length(); ++i) { + const uint out_idx = tid + i * WORKGROUP_SIZE; + const uint head_local = out_idx / HEAD_SIZE; + const uint dim = out_idx % HEAD_SIZE; + if (dim >= dim_base && dim < dim_base + DIMS_PER_BLOCK) { + accum[i] += pv_sh[head_local * PV_STRIDE + (dim - dim_base) / 4][dim % 4]; + } + } + barrier(); + } + } + + if (p.has_sinks != 0 && tid < HEADS_PER_GROUP) { + const float sink = data_s[head_base + tid]; + const float new_max = max(row_max_sh[tid], sink); + const float old_scale = row_sum_sh[tid] == 0.0 ? 0.0 : exp(row_max_sh[tid] - new_max); + row_sum_sh[tid] = row_sum_sh[tid] * old_scale + exp(sink - new_max); + row_max_sh[tid] = new_max; + old_scale_sh[tid] = old_scale; + } + barrier(); + + if (p.has_sinks != 0) { + [[unroll]] for (uint i = 0; i < accum.length(); ++i) { + const uint out_idx = tid + i * WORKGROUP_SIZE; + accum[i] *= old_scale_sh[out_idx / HEAD_SIZE]; + } + } + [[unroll]] for (uint i = 0; i < accum.length(); ++i) { + const uint out_idx = tid + i * WORKGROUP_SIZE; + const uint head_local = out_idx / HEAD_SIZE; + const uint dim = out_idx % HEAD_SIZE; + const uint dst_base = stream * p.nb3 + token * p.nb2 + (head_base + head_local) * p.nb1; + const float inv_sum = row_sum_sh[head_local] == 0.0 ? 0.0 : 1.0 / row_sum_sh[head_local]; + data_dst[dst_base + dim] = accum[i] * inv_sum; + } +} 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 3232278e588..db65e1d49ae 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp @@ -817,6 +817,9 @@ void process_shaders() { string_to_spv("lightning_indexer_decode_cm_f16", "lightning_indexer_decode_cm.comp", {}); #endif string_to_spv("flash_attn_top_k_f16", "flash_attn_top_k.comp", {}); +#if defined(GGML_VULKAN_COOPMAT_GLSLC_SUPPORT) + string_to_spv("flash_attn_top_k_cm_f16", "flash_attn_top_k_cm.comp", {}); +#endif string_to_spv("flash_attn_gather_f16", "flash_attn_gather.comp", {}); string_to_spv("dsv4_hc_pre_f32", "dsv4_hc_pre.comp", {}); string_to_spv("dsv4_hc_comb_f32", "dsv4_hc_comb.comp", {}); diff --git a/tests/test-backend-ops.cpp b/tests/test-backend-ops.cpp index 248fbd0555f..5534f798edf 100644 --- a/tests/test-backend-ops.cpp +++ b/tests/test-backend-ops.cpp @@ -11395,6 +11395,7 @@ static std::vector> make_test_cases_eval() { test_cases.emplace_back(new test_flash_attn_ext_top_k( 512, 4, 64, 128, false)); test_cases.emplace_back(new test_flash_attn_ext_top_k( 768, 64, 64, 128, false)); test_cases.emplace_back(new test_flash_attn_ext_top_k( 768, 64, 64, 128, true)); + test_cases.emplace_back(new test_flash_attn_ext_top_k(1024, 64, 65, 128, false)); test_cases.emplace_back(new test_flash_attn_ext_top_k(4096, 128, 256, 512, false)); return test_cases; From d4cf90797fc52952393afe07c8a730619a0d1cfe Mon Sep 17 00:00:00 2001 From: Jaap Buurman Date: Thu, 13 Aug 2026 11:12:21 +0200 Subject: [PATCH 18/68] vulkan: split sparse prefill attention Assisted-by: Codex --- .../DSV4-vulkan-sparse-prefill-progress.md | 92 +++++++++++++++++ ggml/src/ggml-vulkan/ggml-vulkan.cpp | 99 ++++++++++++++++++- .../vulkan-shaders/flash_attn_base.glsl | 39 ++++---- .../vulkan-shaders/flash_attn_cm1.comp | 16 +-- .../vulkan-shaders/flash_attn_top_k.comp | 2 + .../vulkan-shaders/flash_attn_top_k_cm.comp | 38 +++++-- tests/test-backend-ops.cpp | 1 + 7 files changed, 251 insertions(+), 36 deletions(-) diff --git a/docs/development/DSV4-vulkan-sparse-prefill-progress.md b/docs/development/DSV4-vulkan-sparse-prefill-progress.md index 80bf562a062..93c42476cfc 100644 --- a/docs/development/DSV4-vulkan-sparse-prefill-progress.md +++ b/docs/development/DSV4-vulkan-sparse-prefill-progress.md @@ -257,3 +257,95 @@ GGML_VK_PERF_LOGGER=1 ./build/bin/test-backend-ops perf \ ``` Always run these sequentially. Do not run a compiler concurrently on this APU. + +## Experimental raw-prefix split prototype + +An uncommitted follow-up prototype was tested after commit `24c7ead76cc1a9631fa1b42b0bfa53a15169e1ba`. Check `git status` before continuing. It was initially tested with `GGML_VK_FA_TOPK_SPLIT=1`. After the successful 32k run, the split path was changed to default-on for a llama-server coherence test. Set `GGML_VK_FA_TOPK_SPLIT=0` to restore the single cooperative sparse kernel. + +Stage profiling on the exact production sparse shape (`kv=19200,nb=2048,n_kv_raw=2304,n_top_k=512,sinks=0`) showed: + +```text +cooperative QK and selected-K gather: 88.885 ms +softmax increment: 4.909 ms +cooperative PV and output increment: 128.399 ms +full cooperative sparse kernel: 222.193 ms +``` + +PV and output were 57.8% of the kernel. More importantly, most active keys are not sparse: all 2,304 raw-prefix rows are contiguous, while only 512 compressed rows use top-K indices. The prototype therefore makes two attention partitions: + +1. Ordinary optimized Vulkan cooperative FA processes the contiguous raw prefix. +2. The cooperative top-K shader processes only the 512 selected compressed rows. +3. The existing split-K reduction combines both online-softmax partitions and applies sinks. + +This preserves the exact raw-prefix, sparse top-K, causal mask, and softmax semantics. It reuses the existing ordinary FA and split-K reduction rather than adding a new subsystem. + +The exact-shape microbenchmark improved from 222.19 ms to 111.33 ms. Its steady split stages were about 62-64 ms raw-prefix FA, 43-46 ms selected sparse FA, and 3.6-3.9 ms reduction. + +Correctness validation: + +```bash +GGML_VK_FA_TOPK_SPLIT=1 ./build/bin/test-backend-ops test \ + -b Vulkan0 -o FLASH_ATTN_EXT -p 'n_top_k=' + +./build/bin/test-backend-ops test -b Vulkan0 -o FLASH_ATTN_EXT +``` + +Results were 8/8 sparse top-K cases and 13,296/13,296 complete Vulkan FA cases against the CPU reference. No NaN or Inf failure occurred. + +Canonical 32k prototype command: + +```bash +GGML_VK_FA_TOPK_SPLIT=1 GGML_VK_PERF_LOGGER=1 ./build/bin/llama-bench \ + -m ~/Projects/docker/localLLaMA/models/models--unsloth--DeepSeek-V4-Flash-0731-GGUF/snapshots/109848da2469efe1f1aab9e11acea08a065ccd4f/UD-IQ3_XXS/DeepSeek-V4-Flash-0731-UD-IQ3_XXS-00001-of-00004.gguf \ + -r 1 -d 32768 -p 2048 -ub 2048 -fa 1 -n 0 \ + > /tmp/dsv4-vulkan-split-32k.log 2>&1 +``` + +Final 32k result: + +```text +32k context, PP 2048, ub 2048 + +Old scalar: 112.29 tok/s, 18.1957 s total, 8.84496 s sparse FA +Committed coopmat: 152.32 tok/s, 13.4032 s total, 4.44701 s sparse FA +Experimental split: 208.70 tok/s, 9.7704 s total, 1.20170 s split sparse FA + +Experimental split stages: +raw-prefix FA: 0.210775 s total, 10.037 ms/layer +selected sparse FA: 0.908190 s total, 43.247 ms/layer +split reduction: 0.082732 s total, 3.940 ms/layer +Lightning Indexer: 1.110410 s total, 52.877 ms/layer +TOP_K: 0.077394 s total, 3.685 ms/layer +``` + +Relative to the committed cooperative path, throughput improved 37.0%, total Vulkan time fell 27.1%, and sparse FA time fell 73.0%. Relative to the old scalar path, throughput improved 85.9%, total Vulkan time fell 46.3%, and sparse FA time fell 86.4%. + +The main unresolved tradeoff is scratch memory. At PP2048, the two output partitions use about 539 MiB because each stores an f32 partial output for 512 dimensions x 64 heads x 2048 queries, plus L/M data. The device maximum storage-buffer range is checked and the code falls back when the allocation is unavailable. The path is temporarily default-on for a llama-server coherence test but should not be finalized without discussing this footprint. A likely next step is to avoid materializing both full output partitions, for example by directly merging the selected partition into the raw result or processing query tiles, while retaining the measured split-path speed. + +The old sparse selection heuristic required `total_k >= 3 * active_k`. For the coherence test it now selects sparse attention whenever `total_k > active_k`; equality still uses dense FA because no keys are pruned. The shape, capability, and allocation gates remain unchanged. + +Prototype logs: + +- `/tmp/dsv4-fa-exact-qk.log` +- `/tmp/dsv4-fa-exact-softmax.log` +- `/tmp/dsv4-fa-exact-full.log` +- `/tmp/dsv4-fa-exact-split.log` +- `/tmp/dsv4-fa-split-correctness.log` +- `/tmp/dsv4-fa-all-correctness.log` +- `/tmp/dsv4-vulkan-split-32k.log` + +## Llama-server coherence check + +After making the split path default-on and changing the sparse crossover to `total_k > active_k`, `llama-server` produced a coherent response from a 2,044-token prompt at significant context depth. The final profiler block confirmed that all 21 sparse-attention layers used the split path: + +```text +FA_TOP_K_RAW: 0.179223 s total, 8.534 ms/layer +FA_TOP_K_SELECTED: 0.886629 s total, 42.220 ms/layer +FA_TOP_K_REDUCE: 0.083240 s total, 3.964 ms/layer +Split sparse FA: 1.149092 s total, 54.719 ms/layer +Lightning Indexer: 1.054610 s total, 50.220 ms/layer +TOP_K: 0.044540 s total, 2.121 ms/layer +Total Vulkan: 9.797560 s +``` + +This closely matches the canonical PP2048 llama-bench result of 9.770 s total and 1.202 s split sparse FA. The focused sparse CPU-reference test was rerun after the crossover change and passed 8/8 cases. diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp index 6c2dcef824b..bf5f889faa8 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp +++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp @@ -1997,6 +1997,8 @@ struct vk_op_flash_attn_top_k_push_constants { uint32_t nb1, nb2, nb3; float scale; uint32_t has_sinks; + uint32_t profile_stage; + uint32_t split_mode; }; static_assert(sizeof(vk_op_flash_attn_top_k_push_constants) <= 128); @@ -11458,11 +11460,11 @@ static bool ggml_vk_flash_attn_top_k(ggml_backend_vk_context * ctx, vk_context & return false; } const int64_t n_kv_active = n_kv_raw + top_k->ne[0]; - if (k->ne[1] < 3 * n_kv_active) { + if (k->ne[1] <= n_kv_active) { return false; } - const vk_op_flash_attn_top_k_push_constants pc = { + vk_op_flash_attn_top_k_push_constants pc = { (uint32_t) q->ne[1], (uint32_t) k->ne[1], (uint32_t) n_kv_raw, (uint32_t) top_k->ne[0], (uint32_t) q->ne[2], (uint32_t) (q->nb[1] / sizeof(float)), @@ -11477,14 +11479,105 @@ static bool ggml_vk_flash_attn_top_k(ggml_backend_vk_context * ctx, vk_context & (uint32_t) (dst->nb[1] / sizeof(float)), (uint32_t) (dst->nb[2] / sizeof(float)), (uint32_t) (dst->nb[3] / sizeof(float)), - scale, sinks != nullptr, + scale, sinks != nullptr, 0, 0, }; const vk_subbuffer q_buf = ggml_vk_tensor_subbuffer(ctx, q); const vk_subbuffer sinks_buf = sinks ? ggml_vk_tensor_subbuffer(ctx, sinks) : q_buf; static const char * top_k_cm_env = getenv("GGML_VK_FA_TOPK_CM"); const bool use_cm = (!top_k_cm_env || top_k_cm_env[0] != '0') && ctx->device->pipeline_flash_attn_top_k_cm_f16; + static const char * top_k_profile_env = getenv("GGML_VK_FA_TOPK_PROFILE"); + pc.profile_stage = use_cm && top_k_profile_env ? atoi(top_k_profile_env) : 0; vk_pipeline pipeline = use_cm ? ctx->device->pipeline_flash_attn_top_k_cm_f16 : ctx->device->pipeline_flash_attn_top_k_f16; + + static const char * top_k_split_env = getenv("GGML_VK_FA_TOPK_SPLIT"); + const uint32_t mask_stride = (uint32_t) (mask->nb[1] / sizeof(ggml_fp16_t)); + const bool try_split = use_cm && (!top_k_split_env || top_k_split_env[0] != '0') && n_kv_raw > 0 && top_k->ne[0] > 0 && mask_stride <= 0xffff; + if (try_split) { + const uint32_t N = (uint32_t) q->ne[1]; + const uint32_t D = 512; + const uint32_t NH = 64; + const uint32_t NS = (uint32_t) q->ne[3]; + const uint32_t raw_kv = (uint32_t) n_kv_raw; + const uint32_t partitions = 2; + const bool f32acc = true; + vk_fa_tuning_params tuning = get_fa_tuning_params(ctx->device, D, D, N, raw_kv, GGML_TYPE_F16, GGML_TYPE_F16, f32acc); + + const uint32_t q_stride = (uint32_t) (q->nb[1] / sizeof(float)); + const uint32_t k_stride = (uint32_t) (k->nb[1] / sizeof(ggml_fp16_t)); + const bool aligned = raw_kv % tuning.block_cols == 0 && (q_stride & 7) == 0 && (k_stride & 7) == 0; + const vk_fa_pipeline_state raw_state = get_fa_pipeline_state(ctx->device, tuning, D, D, aligned, f32acc, + true, false, false, GGML_TYPE_F16, GGML_TYPE_F16); + if (raw_state.path == FA_COOPMAT1 && ctx->device->pipeline_flash_attn_split_k_reduce) { + vk_pipeline raw_pipeline; + { + std::lock_guard guard(ctx->device->compile_mutex); + auto & pipelines = ctx->device->pipeline_flash_attn_f32_f16; + auto it = pipelines.find(raw_state); + if (it != pipelines.end()) { + raw_pipeline = it->second; + } else { + pipelines[raw_state] = raw_pipeline = std::make_shared(); + } + } + ggml_pipeline_request_descriptor_sets(ctx, raw_pipeline, 1); + ggml_pipeline_request_descriptor_sets(ctx, pipeline, 1); + ggml_pipeline_request_descriptor_sets(ctx, ctx->device->pipeline_flash_attn_split_k_reduce, 1); + + const uint64_t split_size = ((uint64_t) D * NH * sizeof(float) + NH * 2 * sizeof(float)) * partitions * N * NS; + if (split_size <= ctx->device->properties.limits.maxStorageBufferRange) { + if (ctx->prealloc_size_split_k < split_size) { + ctx->prealloc_size_split_k = split_size; + ggml_vk_preallocate_buffers(ctx, subctx); + } + if (ctx->prealloc_split_k_need_sync) { + ggml_vk_sync_buffers(ctx, subctx); + } + + const vk_subbuffer split_buf = ggml_vk_subbuffer(ctx, ctx->prealloc_split_k, 0); + const uint32_t n_head_log2 = 64; + const uint32_t packed_gqa = (mask_stride << 16) | 1; + const uint32_t packed_partitions = (partitions << 16) | 1; + const vk_flash_attn_push_constants raw_pc = { + N, raw_kv, + NH, N, NS, + NH, NS, + 1, NS, + 1, NS, + (uint32_t) mask->ne[1], (uint32_t) mask->ne[2], (uint32_t) mask->ne[3], + q_stride, (uint32_t) q->nb[2], (uint32_t) q->nb[3], + k_stride, (uint32_t) k->nb[2], (uint32_t) k->nb[3], + k_stride, (uint32_t) k->nb[2], (uint32_t) k->nb[3], + scale, 0.0f, 0.0f, + n_head_log2, 1.0f, 1.0f, + packed_gqa, raw_kv, packed_partitions, + }; + + ggml_vk_dispatch_pipeline(ctx, subctx, raw_pipeline, + {q_buf, ggml_vk_tensor_subbuffer(ctx, k), ggml_vk_tensor_subbuffer(ctx, k), + ggml_vk_tensor_subbuffer(ctx, mask), q_buf, split_buf, q_buf}, + raw_pc, {N, NH, NS}); + ggml_vk_perf_mark_subop(ctx, subctx, "FA_TOP_K_RAW (sub-op)"); + + pc.split_mode = 1; + ggml_vk_dispatch_pipeline(ctx, subctx, pipeline, + {q_buf, ggml_vk_tensor_subbuffer(ctx, k), ggml_vk_tensor_subbuffer(ctx, mask), sinks_buf, + ggml_vk_tensor_subbuffer(ctx, top_k), split_buf}, + pc, {N, (uint32_t) CEIL_DIV(q->ne[2], 32), NS}); + ggml_vk_perf_mark_subop(ctx, subctx, "FA_TOP_K_SELECTED (sub-op)"); + + ggml_vk_sync_buffers(ctx, subctx); + const vk_op_flash_attn_split_k_reduce_push_constants reduce_pc = {D, NH, N, NS, partitions, sinks != nullptr}; + ggml_vk_dispatch_pipeline(ctx, subctx, ctx->device->pipeline_flash_attn_split_k_reduce, + {split_buf, sinks_buf, ggml_vk_tensor_subbuffer(ctx, dst)}, + reduce_pc, {NH, D, N * NS}); + ctx->prealloc_split_k_need_sync = true; + ggml_vk_perf_mark_subop(ctx, subctx, "FA_TOP_K_REDUCE (sub-op)"); + return true; + } + } + } + ggml_pipeline_request_descriptor_sets(ctx, pipeline, 1); ggml_vk_dispatch_pipeline(ctx, subctx, pipeline, {q_buf, ggml_vk_tensor_subbuffer(ctx, k), ggml_vk_tensor_subbuffer(ctx, mask), sinks_buf, 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 0ce4503a884..800a79a97d1 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_base.glsl +++ b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_base.glsl @@ -124,7 +124,7 @@ ACC_TYPE perElemOpStoreCol0(const in uint32_t r, const in uint32_t c, const in A // Load the slope matrix, indexed by Q's dimension 2. ACC_TYPE perElemOpComputeSlope(const in uint32_t r, const in uint32_t c, const in ACC_TYPE elem, const in uint32_t iq2) { - const uint32_t h = iq2 + (r % p.gqa_ratio); + const uint32_t h = iq2 + (r % (p.gqa_ratio & 0xffff)); uint32_t n_head_log2 = p.mask_n_head_log2 & N_LOG2_MASK; @@ -137,32 +137,40 @@ ACC_TYPE perElemOpComputeSlope(const in uint32_t r, const in uint32_t c, const i // Load the sink value, indexed by Q's dimension 2. ACC_TYPE perElemOpGetSink(const in uint32_t r, const in uint32_t c, const in ACC_TYPE elem, const in uint32_t iq2) { - const uint32_t h = iq2 + (r % p.gqa_ratio); + const uint32_t h = iq2 + (r % (p.gqa_ratio & 0xffff)); return ACC_TYPE(data_s[h]); } 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, gqa_ratio, split_k_num, output_k_num; +bool partial_output; void init_indices() { N = p.N; KV = p.KV; + gqa_ratio = p.gqa_ratio & 0xffff; + split_k_num = p.k_num & 0xffff; + output_k_num = p.k_num >> 16; + partial_output = output_k_num != 0; + if (!partial_output) { + output_k_num = split_k_num; + } - if (p.k_num > 1) { - if (p.gqa_ratio > 1) { + if (split_k_num > 1) { + if (gqa_ratio > 1) { i = 0; // batch and split_k share gl_WorkGroupID.x - gqa_iq1 = gl_WorkGroupID.x / p.k_num; - split_k_index = gl_WorkGroupID.x % p.k_num; + gqa_iq1 = gl_WorkGroupID.x / split_k_num; + split_k_index = gl_WorkGroupID.x % split_k_num; } else { gqa_iq1 = 0; - split_k_index = gl_WorkGroupID.x % p.k_num; - i = gl_WorkGroupID.x / p.k_num; + split_k_index = gl_WorkGroupID.x % split_k_num; + i = gl_WorkGroupID.x / split_k_num; } - } else if (p.gqa_ratio > 1) { + } else if (gqa_ratio > 1) { i = 0; gqa_iq1 = gl_WorkGroupID.x; split_k_index = 0; @@ -179,7 +187,7 @@ void init_indices() // When not using grouped query attention, all rows share the same iq2, equal to gl_WorkGroupID.y. // When using grouped query attention, each workgroup does gqa_ratio consecutive values of iq2. - iq2 = gl_WorkGroupID.y * p.gqa_ratio; + iq2 = gl_WorkGroupID.y * gqa_ratio; iq3 = gl_WorkGroupID.z; // broadcast factors @@ -200,14 +208,11 @@ void init_indices() // nb?1 are already divided by the type size and are in units of elements. // When using grouped query attention, Q is indexed by iq2, so the stride // should be nb02 (which is in bytes). - q_stride = p.gqa_ratio > 1 ? (p.nb02 / 4) : p.nb01; + q_stride = gqa_ratio > 1 ? (p.nb02 / 4) : p.nb01; k_stride = p.nb11; v_stride = p.nb21; - // When using grouped query attention, all rows use the same mask (stride 0). - // "p.gqa_ratio >> 16" is just a roundabout way of writing zero - // 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; + const uint32_t mask_stride_override = p.gqa_ratio >> 16; + m_stride = mask_stride_override != 0 ? mask_stride_override : (gqa_ratio > 1 ? 0 : 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 057ed739aa8..e3ea909a59c 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_cm1.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_cm1.comp @@ -169,7 +169,7 @@ void main() { } // Only load if the block is not all zeros if (mask_opt_bits != MASK_OPT_ALL_ZERO) { - bool nem1_bounds_check = !(p.gqa_ratio > 1) && (p.nem1 % Br) != 0; + bool nem1_bounds_check = !(gqa_ratio > 1) && (p.nem1 % Br) != 0; float max_mask = NEG_FLT_MAX_OVER_2; [[unroll]] for (uint32_t idx = 0; idx < Bc * Br / 4; idx += gl_WorkGroupSize.x) { @@ -533,10 +533,10 @@ void main() { // If there is split_k, then the split_k resolve shader does the final // division by L. Store the intermediate O value and per-row m and L values. - if (p.k_num > 1) { - if (p.gqa_ratio > 1) { + if (partial_output || split_k_num > 1) { + if (gqa_ratio > 1) { // note: O and Q have swapped coord 1,2. - uint32_t o_offset = HSV * p.ne1 * (split_k_index + p.k_num * (gqa_iq1 + p.ne2 * iq3)) / 4; + uint32_t o_offset = HSV * p.ne1 * (split_k_index + output_k_num * (gqa_iq1 + p.ne2 * iq3)) / 4; [[unroll]] for (uint32_t r = 0; r < rows_per_thread; ++r) { if (tile_row(r) < N) { @@ -549,7 +549,7 @@ void main() { } } - o_offset = HSV * p.ne1 * p.k_num * p.ne2 * p.ne3 + p.ne1 * 2 * (split_k_index + p.k_num * (gqa_iq1 + p.ne2 * iq3)); + o_offset = HSV * p.ne1 * output_k_num * p.ne2 * p.ne3 + p.ne1 * 2 * (split_k_index + output_k_num * (gqa_iq1 + p.ne2 * iq3)); [[unroll]] for (uint32_t r = 0; r < rows_per_thread; ++r) { if (tile_row(r) < N) { perElemOpStoreCol0(tile_row(r), 0u, ACC_TYPE(Lf[r]), o_offset, iq2, N); @@ -562,7 +562,7 @@ void main() { const uint global_row = i * Br + row; if (global_row < N) { - uint32_t o_offset = HSV * p.ne1 * (split_k_index + p.k_num * (global_row + p.ne2 * iq3)) / 4; + uint32_t o_offset = HSV * p.ne1 * (split_k_index + output_k_num * (global_row + p.ne2 * iq3)) / 4; [[unroll]] for (uint32_t d0 = 0; d0 < HSV / 4; d0 += threads_per_rowgroup) { const uint d = d0 + col_tid; @@ -572,7 +572,7 @@ void main() { } if (global_row < N && col_tid == 0) { - uint32_t lm_offset = HSV * p.ne1 * p.k_num * p.ne2 * p.ne3 + p.ne1 * 2 * (split_k_index + p.k_num * (global_row + p.ne2 * iq3)); + uint32_t lm_offset = HSV * p.ne1 * output_k_num * p.ne2 * p.ne3 + p.ne1 * 2 * (split_k_index + output_k_num * (global_row + p.ne2 * iq3)); data_o[lm_offset + iq2] = D_TYPE(Lf[r]); data_o[lm_offset + p.ne1 + iq2] = D_TYPE(Mf[r]); } @@ -621,7 +621,7 @@ void main() { uint32_t o_offset = (gqa_iq1*p.ne1*HSV + iq3*p.ne2*p.ne1*HSV) / 4; - if (p.gqa_ratio > 1) { + if (gqa_ratio > 1) { [[unroll]] for (uint32_t r = 0; r < rows_per_thread; ++r) { if (tile_row(r) < N) { [[unroll]] for (uint32_t d0 = 0; d0 < HSV / 4; d0 += threads_per_rowgroup) { diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_top_k.comp b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_top_k.comp index b8b49c677bd..62807471bdd 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_top_k.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_top_k.comp @@ -37,6 +37,8 @@ layout(push_constant) uniform Parameters { uint nb3; float scale; uint has_sinks; + uint profile_stage; + uint split_mode; } p; // Shape constants pinned by the dispatch gate in ggml_vk_flash_attn_top_k: DeepSeek V4 diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_top_k_cm.comp b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_top_k_cm.comp index 420a05375a7..f3dfbe59ac0 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_top_k_cm.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_top_k_cm.comp @@ -39,6 +39,8 @@ layout(push_constant) uniform Parameters { uint nb3; float scale; uint has_sinks; + uint profile_stage; + uint split_mode; } p; const uint TILE = 16; @@ -71,7 +73,7 @@ void main() { const uint stream = gl_WorkGroupID.z; const uint mask_base = stream * p.nbm3 + token * p.nbm1; const uint top_base = stream * p.nbt3 + token * p.nbt1; - const uint total_keys = p.n_kv_raw + p.n_top_k; + const uint total_keys = p.split_mode != 0 ? p.n_top_k : p.n_kv_raw + p.n_top_k; float accum[HEADS_PER_GROUP * HEAD_SIZE / WORKGROUP_SIZE]; [[unroll]] for (uint i = 0; i < accum.length(); ++i) { @@ -88,10 +90,11 @@ void main() { if (tid < KEYS_PER_BLOCK) { const uint selected = kb + tid; uint key = p.n_kv; - if (selected < p.n_kv_raw) { + if (p.split_mode == 0 && selected < p.n_kv_raw) { key = selected; } else if (selected < total_keys) { - const int compressed = data_top[top_base + selected - p.n_kv_raw]; + const uint top_pos = p.split_mode != 0 ? selected : selected - p.n_kv_raw; + const int compressed = data_top[top_base + top_pos]; if (compressed >= 0 && uint(compressed) < p.n_kv - p.n_kv_raw) { key = p.n_kv_raw + uint(compressed); } @@ -144,6 +147,10 @@ void main() { SCORE_STRIDE, gl_CooperativeMatrixLayoutRowMajor); barrier(); + if (p.profile_stage == 1) { + continue; + } + { const uint head_local = tid / (SUBGROUP_SIZE / 4); const uint softmax_lane = tid % (SUBGROUP_SIZE / 4); @@ -194,6 +201,10 @@ void main() { accum[i] *= old_scale_sh[head_local]; } + if (p.profile_stage == 2) { + continue; + } + [[unroll]] for (uint dim_base = 0; dim_base < HEAD_SIZE; dim_base += DIMS_PER_BLOCK) { [[unroll]] for (uint idx = tid; idx < KEYS_PER_BLOCK * (DIMS_PER_BLOCK / 4); idx += WORKGROUP_SIZE) { const uint key_local = idx / (DIMS_PER_BLOCK / 4); @@ -242,7 +253,7 @@ void main() { } } - if (p.has_sinks != 0 && tid < HEADS_PER_GROUP) { + if (p.split_mode == 0 && p.has_sinks != 0 && tid < HEADS_PER_GROUP) { const float sink = data_s[head_base + tid]; const float new_max = max(row_max_sh[tid], sink); const float old_scale = row_sum_sh[tid] == 0.0 ? 0.0 : exp(row_max_sh[tid] - new_max); @@ -252,7 +263,7 @@ void main() { } barrier(); - if (p.has_sinks != 0) { + if (p.split_mode == 0 && p.has_sinks != 0) { [[unroll]] for (uint i = 0; i < accum.length(); ++i) { const uint out_idx = tid + i * WORKGROUP_SIZE; accum[i] *= old_scale_sh[out_idx / HEAD_SIZE]; @@ -262,8 +273,19 @@ void main() { const uint out_idx = tid + i * WORKGROUP_SIZE; const uint head_local = out_idx / HEAD_SIZE; const uint dim = out_idx % HEAD_SIZE; - const uint dst_base = stream * p.nb3 + token * p.nb2 + (head_base + head_local) * p.nb1; - const float inv_sum = row_sum_sh[head_local] == 0.0 ? 0.0 : 1.0 / row_sum_sh[head_local]; - data_dst[dst_base + dim] = accum[i] * inv_sum; + if (p.split_mode != 0) { + const uint part_idx = 1; + const uint matrix_base = HEAD_SIZE * p.n_head * (part_idx + 2 * (token + p.n_batch * stream)); + const uint lm_base = HEAD_SIZE * p.n_head * p.n_batch * 2 + p.n_head * 2 * (part_idx + 2 * (token + p.n_batch * stream)); + data_dst[matrix_base + (head_base + head_local) * HEAD_SIZE + dim] = accum[i]; + if (dim == 0) { + data_dst[lm_base + head_base + head_local] = row_sum_sh[head_local]; + data_dst[lm_base + p.n_head + head_base + head_local] = row_max_sh[head_local]; + } + } else { + const uint dst_base = stream * p.nb3 + token * p.nb2 + (head_base + head_local) * p.nb1; + const float inv_sum = row_sum_sh[head_local] == 0.0 ? 0.0 : 1.0 / row_sum_sh[head_local]; + data_dst[dst_base + dim] = accum[i] * inv_sum; + } } } diff --git a/tests/test-backend-ops.cpp b/tests/test-backend-ops.cpp index 5534f798edf..c1629865ce4 100644 --- a/tests/test-backend-ops.cpp +++ b/tests/test-backend-ops.cpp @@ -11898,6 +11898,7 @@ static std::vector> make_test_cases_perf() { test_cases.emplace_back(new test_flash_attn_ext_top_k(kv, nb, 1024, 512, false)); } } + test_cases.emplace_back(new test_flash_attn_ext_top_k(19200, 2048, 2304, 512, false)); return test_cases; } From 57a64cb26af0928106cd239bf98668198bb40a8c Mon Sep 17 00:00:00 2001 From: Jaap Buurman Date: Thu, 13 Aug 2026 11:50:12 +0200 Subject: [PATCH 19/68] vulkan: tile sparse prefill scratch Assisted-by: Codex --- .../DSV4-vulkan-sparse-prefill-progress.md | 58 ++++++++++- ggml/src/ggml-vulkan/ggml-vulkan.cpp | 97 +++++++++++-------- tests/test-backend-ops.cpp | 1 + 3 files changed, 117 insertions(+), 39 deletions(-) diff --git a/docs/development/DSV4-vulkan-sparse-prefill-progress.md b/docs/development/DSV4-vulkan-sparse-prefill-progress.md index 93c42476cfc..5ace979a08d 100644 --- a/docs/development/DSV4-vulkan-sparse-prefill-progress.md +++ b/docs/development/DSV4-vulkan-sparse-prefill-progress.md @@ -320,7 +320,7 @@ TOP_K: 0.077394 s total, 3.685 ms/layer Relative to the committed cooperative path, throughput improved 37.0%, total Vulkan time fell 27.1%, and sparse FA time fell 73.0%. Relative to the old scalar path, throughput improved 85.9%, total Vulkan time fell 46.3%, and sparse FA time fell 86.4%. -The main unresolved tradeoff is scratch memory. At PP2048, the two output partitions use about 539 MiB because each stores an f32 partial output for 512 dimensions x 64 heads x 2048 queries, plus L/M data. The device maximum storage-buffer range is checked and the code falls back when the allocation is unavailable. The path is temporarily default-on for a llama-server coherence test but should not be finalized without discussing this footprint. A likely next step is to avoid materializing both full output partitions, for example by directly merging the selected partition into the raw result or processing query tiles, while retaining the measured split-path speed. +The initial split implementation used 538,968,064 bytes (514 MiB) of scratch at PP2048 because each of two partitions stored an f32 partial output for 512 dimensions x 64 heads x 2048 queries, plus L/M data. This was subsequently reduced by query tiling as described below. The old sparse selection heuristic required `total_k >= 3 * active_k`. For the coherence test it now selects sparse attention whenever `total_k > active_k`; equality still uses dense FA because no keys are pruned. The shape, capability, and allocation gates remain unchanged. @@ -349,3 +349,59 @@ Total Vulkan: 9.797560 s ``` This closely matches the canonical PP2048 llama-bench result of 9.770 s total and 1.202 s split sparse FA. The focused sparse CPU-reference test was rerun after the crossover change and passed 8/8 cases. + +## Tiled split scratch optimization + +Commit `4bbe53e4775f0707de8158e4977a16ab770829da` used two full PP2048 output partitions. The same two-partition algorithm now processes at most 256 query tokens per tile and reuses the split scratch between tiles. Q, mask, top-K, and destination descriptors are offset to the tile while K/V remain shared. A Vulkan pipeline barrier separates reuse of each scratch tile. + +Scratch at PP2048 changed from: + +```text +Before: 538,968,064 bytes (514 MiB) +After: 67,371,008 bytes (64.25 MiB) +Change: 8x reduction +``` + +The exact production-shape microbenchmark (`kv=19200,nb=2048,n_kv_raw=2304,n_top_k=512,sinks=0`) measured: + +```text +Full-batch two partitions: 114.65 ms +256-query tiled path: 113.70 ms +``` + +Two one-partition alternatives were tested and discarded. Merging the raw partial directly inside the cooperative selected shader measured 119.17 ms. Writing selected output separately and using a lightweight merge kernel measured 120.45 ms. Both cut scratch in half but regressed because writing selected output outside the contiguous split layout increased the selected stage from about 44 ms to about 51 ms. Query tiling preserves the faster memory layout. + +Canonical 32k result after tiling: + +```text +32k context, PP 2048, ub 2048 + +Full-batch split: 208.70 tok/s, 9.77039 s total, 1.20170 s sparse FA +Tiled split: 209.62 tok/s, 9.72897 s total, 1.18721 s sparse FA + +Tiled split stages: +raw-prefix FA: 0.192879 s total, 168 tile dispatches +selected sparse FA: 0.910153 s total, 168 tile dispatches +split reduction: 0.084179 s total, 168 tile dispatches +Lightning Indexer: 1.100230 s total +TOP_K: 0.075158 s total +``` + +The 168 dispatch count is eight tiles x 21 sparse-attention layers. Relative to the full-batch split path, throughput improved 0.44%, total Vulkan time fell 0.42%, and sparse FA time fell 1.21%. The main result is the 8x scratch reduction without a performance regression. + +Final correctness after removing the discarded merge prototypes: + +- focused sparse top-K suite: 9/9 passed against the CPU reference, including a 257-query tile-boundary case +- complete Vulkan Flash Attention suite: 13,296/13,296 passed +- no NaN or Inf failure + +Logs: + +- `/tmp/dsv4-fa-exact-fused.log` +- `/tmp/dsv4-fa-exact-merge.log` +- `/tmp/dsv4-fa-exact-two-part-current.log` +- `/tmp/dsv4-fa-exact-tiled.log` +- `/tmp/dsv4-fa-tiled-final-correctness.log` +- `/tmp/dsv4-fa-tiled-boundary-correctness.log` +- `/tmp/dsv4-fa-tiled-all-correctness.log` +- `/tmp/dsv4-vulkan-tiled-32k.log` diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp index bf5f889faa8..0cd804993d7 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp +++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp @@ -11500,6 +11500,8 @@ static bool ggml_vk_flash_attn_top_k(ggml_backend_vk_context * ctx, vk_context & const uint32_t NS = (uint32_t) q->ne[3]; const uint32_t raw_kv = (uint32_t) n_kv_raw; const uint32_t partitions = 2; + const uint32_t tile_size = std::min(N, 256u); + const uint32_t n_tiles = CEIL_DIV(N, tile_size); const bool f32acc = true; vk_fa_tuning_params tuning = get_fa_tuning_params(ctx->device, D, D, N, raw_kv, GGML_TYPE_F16, GGML_TYPE_F16, f32acc); @@ -11520,11 +11522,12 @@ static bool ggml_vk_flash_attn_top_k(ggml_backend_vk_context * ctx, vk_context & pipelines[raw_state] = raw_pipeline = std::make_shared(); } } - ggml_pipeline_request_descriptor_sets(ctx, raw_pipeline, 1); - ggml_pipeline_request_descriptor_sets(ctx, pipeline, 1); - ggml_pipeline_request_descriptor_sets(ctx, ctx->device->pipeline_flash_attn_split_k_reduce, 1); + ggml_pipeline_request_descriptor_sets(ctx, raw_pipeline, n_tiles); + ggml_pipeline_request_descriptor_sets(ctx, pipeline, n_tiles); + ggml_pipeline_request_descriptor_sets(ctx, ctx->device->pipeline_flash_attn_split_k_reduce, n_tiles); - const uint64_t split_size = ((uint64_t) D * NH * sizeof(float) + NH * 2 * sizeof(float)) * partitions * N * NS; + const uint64_t partition_size = ((uint64_t) D * NH + NH * 2) * sizeof(float) * tile_size * NS; + const uint64_t split_size = partition_size * partitions; if (split_size <= ctx->device->properties.limits.maxStorageBufferRange) { if (ctx->prealloc_size_split_k < split_size) { ctx->prealloc_size_split_k = split_size; @@ -11534,45 +11537,63 @@ static bool ggml_vk_flash_attn_top_k(ggml_backend_vk_context * ctx, vk_context & ggml_vk_sync_buffers(ctx, subctx); } - const vk_subbuffer split_buf = ggml_vk_subbuffer(ctx, ctx->prealloc_split_k, 0); const uint32_t n_head_log2 = 64; const uint32_t packed_gqa = (mask_stride << 16) | 1; const uint32_t packed_partitions = (partitions << 16) | 1; - const vk_flash_attn_push_constants raw_pc = { - N, raw_kv, - NH, N, NS, - NH, NS, - 1, NS, - 1, NS, - (uint32_t) mask->ne[1], (uint32_t) mask->ne[2], (uint32_t) mask->ne[3], - q_stride, (uint32_t) q->nb[2], (uint32_t) q->nb[3], - k_stride, (uint32_t) k->nb[2], (uint32_t) k->nb[3], - k_stride, (uint32_t) k->nb[2], (uint32_t) k->nb[3], - scale, 0.0f, 0.0f, - n_head_log2, 1.0f, 1.0f, - packed_gqa, raw_kv, packed_partitions, + const vk_subbuffer split_buf = ggml_vk_subbuffer(ctx, ctx->prealloc_split_k, 0); + const vk_subbuffer k_buf = ggml_vk_tensor_subbuffer(ctx, k); + const vk_subbuffer mask_buf = ggml_vk_tensor_subbuffer(ctx, mask); + const vk_subbuffer top_buf = ggml_vk_tensor_subbuffer(ctx, top_k); + const vk_subbuffer dst_buf = ggml_vk_tensor_subbuffer(ctx, dst); + const auto sliced = [](const vk_subbuffer & buf, uint64_t offset) { + return vk_subbuffer{buf.buffer, buf.offset + offset, buf.size - offset}; }; - ggml_vk_dispatch_pipeline(ctx, subctx, raw_pipeline, - {q_buf, ggml_vk_tensor_subbuffer(ctx, k), ggml_vk_tensor_subbuffer(ctx, k), - ggml_vk_tensor_subbuffer(ctx, mask), q_buf, split_buf, q_buf}, - raw_pc, {N, NH, NS}); - ggml_vk_perf_mark_subop(ctx, subctx, "FA_TOP_K_RAW (sub-op)"); - - pc.split_mode = 1; - ggml_vk_dispatch_pipeline(ctx, subctx, pipeline, - {q_buf, ggml_vk_tensor_subbuffer(ctx, k), ggml_vk_tensor_subbuffer(ctx, mask), sinks_buf, - ggml_vk_tensor_subbuffer(ctx, top_k), split_buf}, - pc, {N, (uint32_t) CEIL_DIV(q->ne[2], 32), NS}); - ggml_vk_perf_mark_subop(ctx, subctx, "FA_TOP_K_SELECTED (sub-op)"); - - ggml_vk_sync_buffers(ctx, subctx); - const vk_op_flash_attn_split_k_reduce_push_constants reduce_pc = {D, NH, N, NS, partitions, sinks != nullptr}; - ggml_vk_dispatch_pipeline(ctx, subctx, ctx->device->pipeline_flash_attn_split_k_reduce, - {split_buf, sinks_buf, ggml_vk_tensor_subbuffer(ctx, dst)}, - reduce_pc, {NH, D, N * NS}); - ctx->prealloc_split_k_need_sync = true; - ggml_vk_perf_mark_subop(ctx, subctx, "FA_TOP_K_REDUCE (sub-op)"); + for (uint32_t tile = 0; tile < n_tiles; ++tile) { + if (tile != 0) { + ggml_vk_sync_buffers(ctx, subctx); + } + const uint32_t token_offset = tile * tile_size; + const uint32_t tile_n = std::min(tile_size, N - token_offset); + const vk_subbuffer tile_q = sliced(q_buf, (uint64_t) token_offset * q->nb[1]); + const vk_subbuffer tile_mask = sliced(mask_buf, (uint64_t) token_offset * mask->nb[1]); + const vk_subbuffer tile_top = sliced(top_buf, (uint64_t) token_offset * top_k->nb[1]); + const vk_subbuffer tile_dst = sliced(dst_buf, (uint64_t) token_offset * dst->nb[2]); + const vk_flash_attn_push_constants raw_pc = { + tile_n, raw_kv, + NH, tile_n, NS, + NH, NS, + 1, NS, + 1, NS, + (uint32_t) mask->ne[1], (uint32_t) mask->ne[2], (uint32_t) mask->ne[3], + q_stride, (uint32_t) q->nb[2], (uint32_t) q->nb[3], + k_stride, (uint32_t) k->nb[2], (uint32_t) k->nb[3], + k_stride, (uint32_t) k->nb[2], (uint32_t) k->nb[3], + scale, 0.0f, 0.0f, + n_head_log2, 1.0f, 1.0f, + packed_gqa, raw_kv, packed_partitions, + }; + + ggml_vk_dispatch_pipeline(ctx, subctx, raw_pipeline, + {tile_q, k_buf, k_buf, tile_mask, tile_q, split_buf, tile_q}, + raw_pc, {tile_n, NH, NS}); + ggml_vk_perf_mark_subop(ctx, subctx, "FA_TOP_K_RAW (sub-op)"); + + pc.n_batch = tile_n; + pc.split_mode = 1; + ggml_vk_dispatch_pipeline(ctx, subctx, pipeline, + {tile_q, k_buf, tile_mask, sinks_buf, tile_top, split_buf}, + pc, {tile_n, (uint32_t) CEIL_DIV(q->ne[2], 32), NS}); + ggml_vk_perf_mark_subop(ctx, subctx, "FA_TOP_K_SELECTED (sub-op)"); + + ctx->prealloc_split_k_need_sync = true; + ggml_vk_sync_buffers(ctx, subctx); + const vk_op_flash_attn_split_k_reduce_push_constants reduce_pc = {D, NH, tile_n, NS, partitions, sinks != nullptr}; + ggml_vk_dispatch_pipeline(ctx, subctx, ctx->device->pipeline_flash_attn_split_k_reduce, + {split_buf, sinks_buf, tile_dst}, reduce_pc, {NH, D, tile_n * NS}); + ctx->prealloc_split_k_need_sync = true; + ggml_vk_perf_mark_subop(ctx, subctx, "FA_TOP_K_REDUCE (sub-op)"); + } return true; } } diff --git a/tests/test-backend-ops.cpp b/tests/test-backend-ops.cpp index c1629865ce4..5d455155085 100644 --- a/tests/test-backend-ops.cpp +++ b/tests/test-backend-ops.cpp @@ -11397,6 +11397,7 @@ static std::vector> make_test_cases_eval() { test_cases.emplace_back(new test_flash_attn_ext_top_k( 768, 64, 64, 128, true)); test_cases.emplace_back(new test_flash_attn_ext_top_k(1024, 64, 65, 128, false)); test_cases.emplace_back(new test_flash_attn_ext_top_k(4096, 128, 256, 512, false)); + test_cases.emplace_back(new test_flash_attn_ext_top_k(4096, 257, 256, 512, false)); return test_cases; } From 989c21b1c5cb1426ab3fdbf4314408ff48e4c8c9 Mon Sep 17 00:00:00 2001 From: Jaap Buurman Date: Thu, 13 Aug 2026 12:19:10 +0200 Subject: [PATCH 20/68] vulkan: reuse sparse FA probability fragments Assisted-by: Codex --- .../DSV4-vulkan-sparse-prefill-progress.md | 85 +++++++++++++++++++ .../vulkan-shaders/flash_attn_top_k_cm.comp | 27 +++--- 2 files changed, 100 insertions(+), 12 deletions(-) diff --git a/docs/development/DSV4-vulkan-sparse-prefill-progress.md b/docs/development/DSV4-vulkan-sparse-prefill-progress.md index 5ace979a08d..82843486bc0 100644 --- a/docs/development/DSV4-vulkan-sparse-prefill-progress.md +++ b/docs/development/DSV4-vulkan-sparse-prefill-progress.md @@ -405,3 +405,88 @@ Logs: - `/tmp/dsv4-fa-tiled-boundary-correctness.log` - `/tmp/dsv4-fa-tiled-all-correctness.log` - `/tmp/dsv4-vulkan-tiled-32k.log` + +## Selected PV probability reuse + +The selected cooperative-matrix stage remained the largest split sparse-FA component. Profiling the exact tiled production shape showed: + +```text +selected K gather and cooperative QK: about 20.0 ms +QK plus softmax: about 22.3 ms +full selected stage: about 45.5 ms +cooperative PV and output increment: about 23.2 ms +``` + +PV and output were about 51% of the selected stage. The shader previously loaded each of four 16x16 probability cooperative-matrix fragments again for every one of the eight 64-dimension PV passes. The optimized shader loads these four fragments once per 64-key block and retains them across all PV dimension passes. This follows the probability-fragment lifetime used by ordinary cooperative Vulkan FA. + +The shader also aliases the shared score and PV-output matrices because their lifetimes do not overlap. Shader resource statistics changed as follows: + +```text + Before After +VGPRs 192 192 +VGPR spills 0 0 +LDS 36,864 28,672 bytes +static instructions 11,524 11,358 +``` + +The exact production-shape microbenchmark changed from 113.70 ms for the committed tiled path to 110.11 ms. The selected stage fell from about 45.5 ms to about 43.7 ms. The raw-prefix and reduction implementations are unchanged. + +Two canonical 32k runs after this change measured: + +```text +32k context, PP 2048, ub 2048 + +Committed tiled reference: +209.62 tok/s +Total Vulkan: 9.72897 s +Split sparse FA: 1.18721 s + raw-prefix FA: 0.192879 s + selected sparse FA: 0.910153 s + split reduction: 0.084179 s +Lightning Indexer: 1.100230 s +TOP_K: 0.075158 s + +Probability reuse, run 1: +211.03 tok/s +Total Vulkan: 9.66363 s +Split sparse FA: 1.12776 s + raw-prefix FA: 0.192078 s + selected sparse FA: 0.849862 s + split reduction: 0.085821 s +Lightning Indexer: 1.098980 s +TOP_K: 0.071036 s + +Probability reuse, run 2: +209.57 tok/s +Total Vulkan: 9.72968 s +Split sparse FA: 1.13553 s + raw-prefix FA: 0.192987 s + selected sparse FA: 0.859890 s + split reduction: 0.082653 s +Lightning Indexer: 1.102550 s +TOP_K: 0.071474 s +``` + +The two-run selected-stage improvement is 5.5-6.6%, and the complete split sparse-FA improvement is 4.4-5.0%. The two-run throughput mean is 210.30 tok/s, 0.32% above the 209.62 tok/s reference. End-to-end noise in unrelated model kernels is larger than this small total-throughput change, but both profiler runs isolate a consistent gain in the modified selected stage. + +Correctness after the change: + +- focused sparse top-K suite: 9/9 passed against the CPU reference +- complete Vulkan Flash Attention suite: 13,297/13,297 passed +- pipeline statistics probe: 1/1 passed, with no SGPR or VGPR spills +- no NaN or Inf failure + +Logs: + +- `/tmp/dsv4-fa-selected-qk-tiled.log` +- `/tmp/dsv4-fa-selected-softmax-tiled.log` +- `/tmp/dsv4-fa-lds-alias-stats.log` +- `/tmp/dsv4-fa-pmat-stats.log` +- `/tmp/dsv4-fa-exact-lds-alias.log` +- `/tmp/dsv4-fa-exact-pmat.log` +- `/tmp/dsv4-fa-pmat-correctness.log` +- `/tmp/dsv4-fa-pmat-all-correctness.log` +- `/tmp/dsv4-vulkan-32k-pmat.log` +- `/tmp/dsv4-vulkan-32k-pmat-repeat.log` + +The next sparse-FA optimization should continue to target the selected PV/output half. The retained probability fragments remove redundant cooperative loads without increasing reported VGPR allocation. More invasive changes such as doubling the PV dimension tile can reduce barriers but must be designed around the eight available wave64 subgroups, 64 KiB LDS limit, and already high 192-VGPR allocation. Do not build a crossover matrix until the next kernel layout is settled because an optimization can shift the crossover. diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_top_k_cm.comp b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_top_k_cm.comp index f3dfbe59ac0..d4bc4011d07 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_top_k_cm.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_top_k_cm.comp @@ -58,10 +58,9 @@ const float MASK_NEG_INF = -65500.0; shared uint key_idx[KEYS_PER_BLOCK]; shared f16vec4 q_sh[HEADS_PER_GROUP * QK_STRIDE]; shared f16vec4 k_sh[KEYS_PER_BLOCK * QK_STRIDE]; -shared vec4 score_sh[KEYS_PER_BLOCK * SCORE_STRIDE]; +shared vec4 matrix_sh[KEYS_PER_BLOCK * SCORE_STRIDE]; shared f16vec4 p_sh[HEADS_PER_GROUP * P_STRIDE]; shared f16vec4 v_sh[KEYS_PER_BLOCK * V_STRIDE]; -shared vec4 pv_sh[HEADS_PER_GROUP * PV_STRIDE]; shared float old_scale_sh[HEADS_PER_GROUP]; shared float row_max_sh[HEADS_PER_GROUP]; shared float row_sum_sh[HEADS_PER_GROUP]; @@ -142,7 +141,7 @@ void main() { const uint score_key_chunk = gl_SubgroupID % (KEYS_PER_BLOCK / TILE); const uint score_head_tile = gl_SubgroupID / (KEYS_PER_BLOCK / TILE); - coopMatStore(scores, score_sh, + coopMatStore(scores, matrix_sh, score_key_chunk * TILE * SCORE_STRIDE + score_head_tile * (TILE / 4), SCORE_STRIDE, gl_CooperativeMatrixLayoutRowMajor); barrier(); @@ -159,7 +158,7 @@ void main() { const uint key_local = softmax_lane + i * (SUBGROUP_SIZE / 4); const uint key = key_idx[key_local]; const float mask = key < p.n_kv ? float(data_m[mask_base + key]) : uintBitsToFloat(0xff800000); - const float score = float(score_sh[key_local * SCORE_STRIDE + head_local / 4][head_local % 4]) * p.scale + mask; + const float score = float(matrix_sh[key_local * SCORE_STRIDE + head_local / 4][head_local % 4]) * p.scale + mask; block_max = mask < MASK_NEG_INF ? block_max : max(block_max, score); } [[unroll]] for (uint delta = 1; delta < SUBGROUP_SIZE / 4; delta *= 2) { @@ -177,7 +176,7 @@ void main() { const float mask = key < p.n_kv ? float(data_m[mask_base + key]) : uintBitsToFloat(0xff800000); float weight = 0.0; if (mask >= MASK_NEG_INF) { - const float score = float(score_sh[key_local * SCORE_STRIDE + head_local / 4][head_local % 4]) * p.scale + mask; + const float score = float(matrix_sh[key_local * SCORE_STRIDE + head_local / 4][head_local % 4]) * p.scale + mask; weight = exp(score - new_max); block_sum += weight; } @@ -205,6 +204,14 @@ void main() { continue; } + coopmat pmats[KEYS_PER_BLOCK / TILE]; + [[unroll]] for (uint key_chunk = 0; key_chunk < KEYS_PER_BLOCK / TILE; ++key_chunk) { + const uint pv_head_tile = gl_SubgroupID / (DIMS_PER_BLOCK / TILE); + coopMatLoad(pmats[key_chunk], p_sh, + pv_head_tile * TILE * P_STRIDE + key_chunk * (TILE / 4), + P_STRIDE, gl_CooperativeMatrixLayoutRowMajor); + } + [[unroll]] for (uint dim_base = 0; dim_base < HEAD_SIZE; dim_base += DIMS_PER_BLOCK) { [[unroll]] for (uint idx = tid; idx < KEYS_PER_BLOCK * (DIMS_PER_BLOCK / 4); idx += WORKGROUP_SIZE) { const uint key_local = idx / (DIMS_PER_BLOCK / 4); @@ -221,22 +228,18 @@ void main() { coopmat pv = coopmat(0.0); - coopmat pmat; coopmat vmat; const uint pv_head_tile = gl_SubgroupID / (DIMS_PER_BLOCK / TILE); const uint pv_dim_tile = gl_SubgroupID % (DIMS_PER_BLOCK / TILE); [[unroll]] for (uint key_chunk = 0; key_chunk < KEYS_PER_BLOCK / TILE; ++key_chunk) { - coopMatLoad(pmat, p_sh, - pv_head_tile * TILE * P_STRIDE + key_chunk * (TILE / 4), - P_STRIDE, gl_CooperativeMatrixLayoutRowMajor); coopMatLoad(vmat, v_sh, key_chunk * TILE * V_STRIDE + pv_dim_tile * (TILE / 4), V_STRIDE, gl_CooperativeMatrixLayoutRowMajor); - pv = coopMatMulAdd(pmat, vmat, pv); + pv = coopMatMulAdd(pmats[key_chunk], vmat, pv); } - coopMatStore(pv, pv_sh, + coopMatStore(pv, matrix_sh, pv_head_tile * TILE * PV_STRIDE + pv_dim_tile * (TILE / 4), PV_STRIDE, gl_CooperativeMatrixLayoutRowMajor); barrier(); @@ -246,7 +249,7 @@ void main() { const uint head_local = out_idx / HEAD_SIZE; const uint dim = out_idx % HEAD_SIZE; if (dim >= dim_base && dim < dim_base + DIMS_PER_BLOCK) { - accum[i] += pv_sh[head_local * PV_STRIDE + (dim - dim_base) / 4][dim % 4]; + accum[i] += matrix_sh[head_local * PV_STRIDE + (dim - dim_base) / 4][dim % 4]; } } barrier(); From 94ecd38a20042244c10f668c4474b16449987580 Mon Sep 17 00:00:00 2001 From: Jaap Buurman Date: Thu, 13 Aug 2026 12:59:50 +0200 Subject: [PATCH 21/68] vulkan: cache sparse FA masks per key block Assisted-by: Codex --- .../DSV4-vulkan-sparse-prefill-progress.md | 103 +++++++++++++++++- ggml/src/ggml-vulkan/ggml-vulkan.cpp | 7 +- .../vulkan-shaders/flash_attn_base.glsl | 3 +- .../vulkan-shaders/flash_attn_top_k_cm.comp | 30 +++-- tests/test-backend-ops.cpp | 5 +- 5 files changed, 123 insertions(+), 25 deletions(-) diff --git a/docs/development/DSV4-vulkan-sparse-prefill-progress.md b/docs/development/DSV4-vulkan-sparse-prefill-progress.md index 82843486bc0..e880ae63e99 100644 --- a/docs/development/DSV4-vulkan-sparse-prefill-progress.md +++ b/docs/development/DSV4-vulkan-sparse-prefill-progress.md @@ -262,7 +262,7 @@ Always run these sequentially. Do not run a compiler concurrently on this APU. An uncommitted follow-up prototype was tested after commit `24c7ead76cc1a9631fa1b42b0bfa53a15169e1ba`. Check `git status` before continuing. It was initially tested with `GGML_VK_FA_TOPK_SPLIT=1`. After the successful 32k run, the split path was changed to default-on for a llama-server coherence test. Set `GGML_VK_FA_TOPK_SPLIT=0` to restore the single cooperative sparse kernel. -Stage profiling on the exact production sparse shape (`kv=19200,nb=2048,n_kv_raw=2304,n_top_k=512,sinks=0`) showed: +Stage profiling on the previously used 19,200-row sparse shape (`kv=19200,nb=2048,n_kv_raw=2304,n_top_k=512,sinks=0`) showed: ```text cooperative QK and selected-K gather: 88.885 ms @@ -362,7 +362,7 @@ After: 67,371,008 bytes (64.25 MiB) Change: 8x reduction ``` -The exact production-shape microbenchmark (`kv=19200,nb=2048,n_kv_raw=2304,n_top_k=512,sinks=0`) measured: +The previously used 19,200-row microbenchmark (`kv=19200,nb=2048,n_kv_raw=2304,n_top_k=512,sinks=0`) measured: ```text Full-batch two partitions: 114.65 ms @@ -408,7 +408,7 @@ Logs: ## Selected PV probability reuse -The selected cooperative-matrix stage remained the largest split sparse-FA component. Profiling the exact tiled production shape showed: +The selected cooperative-matrix stage remained the largest split sparse-FA component. Profiling the tiled 19,200-row shape showed: ```text selected K gather and cooperative QK: about 20.0 ms @@ -429,7 +429,7 @@ LDS 36,864 28,672 bytes static instructions 11,524 11,358 ``` -The exact production-shape microbenchmark changed from 113.70 ms for the committed tiled path to 110.11 ms. The selected stage fell from about 45.5 ms to about 43.7 ms. The raw-prefix and reduction implementations are unchanged. +The 19,200-row microbenchmark changed from 113.70 ms for the committed tiled path to 110.11 ms. The selected stage fell from about 45.5 ms to about 43.7 ms. The raw-prefix and reduction implementations are unchanged. Two canonical 32k runs after this change measured: @@ -490,3 +490,98 @@ Logs: - `/tmp/dsv4-vulkan-32k-pmat-repeat.log` The next sparse-FA optimization should continue to target the selected PV/output half. The retained probability fragments remove redundant cooperative loads without increasing reported VGPR allocation. More invasive changes such as doubling the PV dimension tile can reduce barriers but must be designed around the eight available wave64 subgroups, 64 KiB LDS limit, and already high 192-VGPR allocation. Do not build a crossover matrix until the next kernel layout is settled because an optimization can shift the crossover. + +## Selected mask caching and deep-context micro matrix + +The sparse FA `kv` dimension is compressed K/V rows, not source-token context depth. For the tested DeepSeek V4 graph, a PP2048 batch uses 2,304 raw rows and approximately one compressed row per four source tokens. The useful synthetic mapping is: + +```text +Source context Sparse FA kv rows +32k 11,008 +64k 19,200 +128k 35,584 +256k 68,352 +512k 133,888 +``` + +The performance test registry now contains PP2048 cases at all five K extents with `n_kv_raw=2304` and `n_top_k=512`. Use this command template and replace `KV` with a value from the table: + +```bash +GGML_VK_PERF_LOGGER=1 ./build/bin/test-backend-ops perf \ + -b Vulkan0 -o FLASH_ATTN_EXT \ + -p 'kv=KV,nb=2048,n_kv_raw=2304,n_top_k=512,sinks=0' +``` + +These are synthetic sparse-FA depth tests. They do not include the Lightning Indexer or the rest of the model graph and do not replace the canonical 32k llama-bench. + +The selected shader previously loaded the same selected-key mask value independently for all 32 heads during both softmax passes. The shader now loads each mask value once per 64-key block and reuses it from LDS. Q and probability storage also share one LDS allocation, while K and V staging share another because both pairs have disjoint lifetimes. + +Shader resources changed from the probability-reuse commit: + +```text + Before After +VGPRs 192 192 +VGPR spills 0 0 +LDS 28,672 24,576 bytes +static instructions 11,358 11,233 +``` + +A 128-dimension V staging experiment was correct but rejected. With operand aliasing it used exactly 32 KiB LDS. It was 2.3% slower at simulated 32k, approximately equal at 64k, and 1.6% faster at 128k. The retained 64-dimension layout is better for the canonical 32k target and is simpler. + +Selected-mask caching improved every synthetic depth relative to the same 64-dimension operand-alias layout: + +```text +Depth Selected before Selected after Split total before Split total after +32k 43.16 ms 41.23 ms 109.45 ms 108.07 ms +64k 43.22 ms 41.42 ms 110.45 ms 109.11 ms +128k 52.26 ms 48.74 ms 119.15 ms 116.92 ms +256k 51.64 ms 49.02 ms 118.06 ms 116.93 ms +512k 51.33 ms 49.20 ms 117.87 ms 116.27 ms +``` + +The 256k case exposed a separate path-selection cutoff. The raw-prefix ordinary FA dispatch encoded its mask row stride in 16 bits and disabled split sparse FA above 65,535 K rows. It now uses a flagged convention that carries the full 32-bit mask stride in the existing `split_kv` push constant for this one-partition partial-output dispatch. No push-constant structure was enlarged. At simulated 256k, this changes the selected path from unsplit cooperative sparse FA at 206.22 ms to split sparse FA at 118.06 ms before mask caching, a 42.8% reduction. The 512k case also uses the split path successfully. + +Canonical 32k PP2048 after operand aliasing and selected-mask caching: + +```text +Previous commit, two runs: +209.57-211.03 tok/s +Total Vulkan: 9.664-9.730 s +Split sparse FA: 1.128-1.136 s + raw-prefix FA: 0.192-0.193 s + selected sparse FA: 0.850-0.860 s + split reduction: 0.083-0.086 s + +Current result: +210.47 tok/s +Total Vulkan: 9.68921 s +Split sparse FA: 1.10160 s + raw-prefix FA: 0.192860 s + selected sparse FA: 0.821381 s + split reduction: 0.087354 s +Lightning Indexer: 1.102230 s +TOP_K: 0.072488 s +``` + +The selected stage improves 3.4-4.5% and complete split sparse FA improves 2.3-3.0% relative to the previous commit's two canonical runs. End-to-end throughput remains within model-wide benchmark noise. + +Correctness: + +- focused sparse top-K suite: 9/9 passed +- complete Vulkan Flash Attention suite: 13,297/13,297 passed +- no NaN or Inf failure +- simulated 256k and 512k performance cases both selected the split path and completed successfully + +Benchmark policy for this laptop APU: do not repeat llama-bench when the exact-shape microbenchmarks and the first canonical profiler block agree. Sustained load can lower GPU clocks and bias a repeat. Repeat only when the first result is anomalous or contradicts the microbenchmarks. Never build and benchmark concurrently. + +Logs: + +- `/tmp/dsv4-fa-mask-cache-32k.log` +- `/tmp/dsv4-fa-mask-cache-64k.log` +- `/tmp/dsv4-fa-mask-cache-128k.log` +- `/tmp/dsv4-fa-mask-cache-256k.log` +- `/tmp/dsv4-fa-mask-cache-512k.log` +- `/tmp/dsv4-fa-final-micro-32k.log` +- `/tmp/dsv4-fa-final-focused-correctness.log` +- `/tmp/dsv4-fa-mask-cache-all-correctness.log` +- `/tmp/dsv4-vulkan-32k-mask-cache.log` diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp index 0cd804993d7..5ba53b06a5d 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp +++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp @@ -11492,7 +11492,7 @@ static bool ggml_vk_flash_attn_top_k(ggml_backend_vk_context * ctx, vk_context & static const char * top_k_split_env = getenv("GGML_VK_FA_TOPK_SPLIT"); const uint32_t mask_stride = (uint32_t) (mask->nb[1] / sizeof(ggml_fp16_t)); - const bool try_split = use_cm && (!top_k_split_env || top_k_split_env[0] != '0') && n_kv_raw > 0 && top_k->ne[0] > 0 && mask_stride <= 0xffff; + const bool try_split = use_cm && (!top_k_split_env || top_k_split_env[0] != '0') && n_kv_raw > 0 && top_k->ne[0] > 0; if (try_split) { const uint32_t N = (uint32_t) q->ne[1]; const uint32_t D = 512; @@ -11538,7 +11538,8 @@ static bool ggml_vk_flash_attn_top_k(ggml_backend_vk_context * ctx, vk_context & } const uint32_t n_head_log2 = 64; - const uint32_t packed_gqa = (mask_stride << 16) | 1; + const uint32_t mask_stride_in_split_kv = 1u << 31; + const uint32_t packed_gqa = mask_stride_in_split_kv | 1u; const uint32_t packed_partitions = (partitions << 16) | 1; const vk_subbuffer split_buf = ggml_vk_subbuffer(ctx, ctx->prealloc_split_k, 0); const vk_subbuffer k_buf = ggml_vk_tensor_subbuffer(ctx, k); @@ -11571,7 +11572,7 @@ static bool ggml_vk_flash_attn_top_k(ggml_backend_vk_context * ctx, vk_context & k_stride, (uint32_t) k->nb[2], (uint32_t) k->nb[3], scale, 0.0f, 0.0f, n_head_log2, 1.0f, 1.0f, - packed_gqa, raw_kv, packed_partitions, + packed_gqa, mask_stride, packed_partitions, }; ggml_vk_dispatch_pipeline(ctx, subctx, raw_pipeline, 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 800a79a97d1..8fba7a450f7 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_base.glsl +++ b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_base.glsl @@ -211,8 +211,9 @@ void init_indices() q_stride = gqa_ratio > 1 ? (p.nb02 / 4) : p.nb01; k_stride = p.nb11; v_stride = p.nb21; + const bool mask_stride_in_split_kv = (p.gqa_ratio & 0x80000000u) != 0; const uint32_t mask_stride_override = p.gqa_ratio >> 16; - m_stride = mask_stride_override != 0 ? mask_stride_override : (gqa_ratio > 1 ? 0 : KV); + m_stride = mask_stride_in_split_kv ? p.split_kv : (mask_stride_override != 0 ? mask_stride_override : (gqa_ratio > 1 ? 0 : KV)); } // Bias applied to softmax to stay in fp16 range. diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_top_k_cm.comp b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_top_k_cm.comp index d4bc4011d07..3d0024553e7 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_top_k_cm.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_top_k_cm.comp @@ -56,11 +56,10 @@ const uint PV_STRIDE = DIMS_PER_BLOCK / 4; const float MASK_NEG_INF = -65500.0; shared uint key_idx[KEYS_PER_BLOCK]; -shared f16vec4 q_sh[HEADS_PER_GROUP * QK_STRIDE]; -shared f16vec4 k_sh[KEYS_PER_BLOCK * QK_STRIDE]; +shared float key_mask[KEYS_PER_BLOCK]; +shared f16vec4 q_p_sh[HEADS_PER_GROUP * P_STRIDE]; +shared f16vec4 k_v_sh[KEYS_PER_BLOCK * V_STRIDE]; shared vec4 matrix_sh[KEYS_PER_BLOCK * SCORE_STRIDE]; -shared f16vec4 p_sh[HEADS_PER_GROUP * P_STRIDE]; -shared f16vec4 v_sh[KEYS_PER_BLOCK * V_STRIDE]; shared float old_scale_sh[HEADS_PER_GROUP]; shared float row_max_sh[HEADS_PER_GROUP]; shared float row_sum_sh[HEADS_PER_GROUP]; @@ -99,6 +98,7 @@ void main() { } } key_idx[tid] = key; + key_mask[tid] = key < p.n_kv ? float(data_m[mask_base + key]) : uintBitsToFloat(0xff800000); } barrier(); @@ -117,23 +117,23 @@ void main() { const uint offset = stream * p.nbk3 + key * p.nbk1 + d + d4 * 4; value = f16vec4(data_k[offset], data_k[offset + 1], data_k[offset + 2], data_k[offset + 3]); } - k_sh[key_local * QK_STRIDE + d4] = value; + k_v_sh[key_local * QK_STRIDE + d4] = value; } if (tid < HEADS_PER_GROUP * (TILE / 4)) { const uint head_local = tid / (TILE / 4); const uint d4 = tid % (TILE / 4); const uint head = head_base + head_local; const uint offset = stream * p.nbq3 + head * p.nbq2 + token * p.nbq1 + d + d4 * 4; - q_sh[head_local * QK_STRIDE + d4] = f16vec4( + q_p_sh[head_local * QK_STRIDE + d4] = f16vec4( data_q[offset], data_q[offset + 1], data_q[offset + 2], data_q[offset + 3]); } barrier(); const uint key_chunk = gl_SubgroupID % (KEYS_PER_BLOCK / TILE); const uint head_tile = gl_SubgroupID / (KEYS_PER_BLOCK / TILE); - coopMatLoad(kmat, k_sh, key_chunk * TILE * QK_STRIDE, + coopMatLoad(kmat, k_v_sh, key_chunk * TILE * QK_STRIDE, QK_STRIDE, gl_CooperativeMatrixLayoutRowMajor); - coopMatLoad(qmat, q_sh, head_tile * TILE * QK_STRIDE, + coopMatLoad(qmat, q_p_sh, head_tile * TILE * QK_STRIDE, QK_STRIDE, gl_CooperativeMatrixLayoutColumnMajor); scores = coopMatMulAdd(kmat, qmat, scores); barrier(); @@ -156,8 +156,7 @@ void main() { float block_max = uintBitsToFloat(0xff800000); [[unroll]] for (uint i = 0; i < KEYS_PER_BLOCK / (SUBGROUP_SIZE / 4); ++i) { const uint key_local = softmax_lane + i * (SUBGROUP_SIZE / 4); - const uint key = key_idx[key_local]; - const float mask = key < p.n_kv ? float(data_m[mask_base + key]) : uintBitsToFloat(0xff800000); + const float mask = key_mask[key_local]; const float score = float(matrix_sh[key_local * SCORE_STRIDE + head_local / 4][head_local % 4]) * p.scale + mask; block_max = mask < MASK_NEG_INF ? block_max : max(block_max, score); } @@ -172,15 +171,14 @@ void main() { float block_sum = 0.0; [[unroll]] for (uint i = 0; i < KEYS_PER_BLOCK / (SUBGROUP_SIZE / 4); ++i) { const uint key_local = softmax_lane + i * (SUBGROUP_SIZE / 4); - const uint key = key_idx[key_local]; - const float mask = key < p.n_kv ? float(data_m[mask_base + key]) : uintBitsToFloat(0xff800000); + const float mask = key_mask[key_local]; float weight = 0.0; if (mask >= MASK_NEG_INF) { const float score = float(matrix_sh[key_local * SCORE_STRIDE + head_local / 4][head_local % 4]) * p.scale + mask; weight = exp(score - new_max); block_sum += weight; } - p_sh[head_local * P_STRIDE + key_local / 4][key_local % 4] = float16_t(weight); + q_p_sh[head_local * P_STRIDE + key_local / 4][key_local % 4] = float16_t(weight); } [[unroll]] for (uint delta = 1; delta < SUBGROUP_SIZE / 4; delta *= 2) { block_sum += subgroupShuffleXor(block_sum, delta); @@ -207,7 +205,7 @@ void main() { coopmat pmats[KEYS_PER_BLOCK / TILE]; [[unroll]] for (uint key_chunk = 0; key_chunk < KEYS_PER_BLOCK / TILE; ++key_chunk) { const uint pv_head_tile = gl_SubgroupID / (DIMS_PER_BLOCK / TILE); - coopMatLoad(pmats[key_chunk], p_sh, + coopMatLoad(pmats[key_chunk], q_p_sh, pv_head_tile * TILE * P_STRIDE + key_chunk * (TILE / 4), P_STRIDE, gl_CooperativeMatrixLayoutRowMajor); } @@ -222,7 +220,7 @@ void main() { const uint offset = stream * p.nbk3 + key * p.nbk1 + dim_base + d4 * 4; value = f16vec4(data_k[offset], data_k[offset + 1], data_k[offset + 2], data_k[offset + 3]); } - v_sh[key_local * V_STRIDE + d4] = value; + k_v_sh[key_local * V_STRIDE + d4] = value; } barrier(); @@ -233,7 +231,7 @@ void main() { const uint pv_head_tile = gl_SubgroupID / (DIMS_PER_BLOCK / TILE); const uint pv_dim_tile = gl_SubgroupID % (DIMS_PER_BLOCK / TILE); [[unroll]] for (uint key_chunk = 0; key_chunk < KEYS_PER_BLOCK / TILE; ++key_chunk) { - coopMatLoad(vmat, v_sh, + coopMatLoad(vmat, k_v_sh, key_chunk * TILE * V_STRIDE + pv_dim_tile * (TILE / 4), V_STRIDE, gl_CooperativeMatrixLayoutRowMajor); pv = coopMatMulAdd(pmats[key_chunk], vmat, pv); diff --git a/tests/test-backend-ops.cpp b/tests/test-backend-ops.cpp index 5d455155085..23913f62700 100644 --- a/tests/test-backend-ops.cpp +++ b/tests/test-backend-ops.cpp @@ -11899,7 +11899,10 @@ static std::vector> make_test_cases_perf() { test_cases.emplace_back(new test_flash_attn_ext_top_k(kv, nb, 1024, 512, false)); } } - test_cases.emplace_back(new test_flash_attn_ext_top_k(19200, 2048, 2304, 512, false)); + // PP2048 compressed-K rows for source context depths 32k through 512k. + for (int kv : { 11008, 19200, 35584, 68352, 133888 }) { + test_cases.emplace_back(new test_flash_attn_ext_top_k(kv, 2048, 2304, 512, false)); + } return test_cases; } From 5f60911fe26e6ff2a537a5697c020e37951ba079 Mon Sep 17 00:00:00 2001 From: Nathan Wilson Date: Thu, 13 Aug 2026 14:32:06 +0000 Subject: [PATCH 22/68] vulkan: fix DeepSeek V4 sparse split attention with multiple sequences The raw/selected split path mis-indexed three things once q->ne[3] > 1. All three are inert at a single sequence, so the existing tests never reached them. - flash_attn_top_k_cm.comp wrote the L/M rows at a base that omitted the stream count. The split buffer is [O matrices][L/M rows] with both regions spanning ne3, so with more than one sequence the selected partition's L/M landed inside the O region and corrupted partition 0. - The mask's per-stream offset stepped by nem1 * KV. That assumes the mask row length equals KV, which the split path breaks: it overrides the mask stride to the full K range while KV covers only the raw prefix. Separate the two quantities - m_stride is the row-to-row step inside the tile (0 under GQA), m_row_len is the mask's real row length used to step between tokens and streams. They coincide everywhere except GQA and this path. - Query tiling conflicts with multiple streams: flash_attn_split_k_reduce takes the tile height as ne2 and derives the destination row from it, so with several tiles it writes stream s at s*tile_size instead of s*N. Tile only when there is one stream, which is the prefill case the tiling was added for. Co-Authored-By: Claude Opus 5 Assisted-by: Claude (Opus 5) --- ggml/src/ggml-vulkan/ggml-vulkan.cpp | 7 ++++++- .../vulkan-shaders/flash_attn_base.glsl | 18 +++++++++++++++--- .../vulkan-shaders/flash_attn_cm1.comp | 4 ++-- .../vulkan-shaders/flash_attn_top_k_cm.comp | 9 ++++++++- 4 files changed, 31 insertions(+), 7 deletions(-) diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp index 5ba53b06a5d..4342e78f2a6 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp +++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp @@ -11500,7 +11500,12 @@ static bool ggml_vk_flash_attn_top_k(ggml_backend_vk_context * ctx, vk_context & const uint32_t NS = (uint32_t) q->ne[3]; const uint32_t raw_kv = (uint32_t) n_kv_raw; const uint32_t partitions = 2; - const uint32_t tile_size = std::min(N, 256u); + // Query tiling keeps the split scratch small, but flash_attn_split_k_reduce derives the + // destination row from ne2, which it is handed as the TILE height. With one tile that + // equals N and the stream stride is right; with several tiles and more than one stream + // it would write stream s at s*tile_size instead of s*N. Only tile when there is a + // single stream, which is the prefill case the tiling exists for. + const uint32_t tile_size = (NS == 1) ? std::min(N, 256u) : N; const uint32_t n_tiles = CEIL_DIV(N, tile_size); const bool f32acc = true; vk_fa_tuning_params tuning = get_fa_tuning_params(ctx->device, D, D, N, raw_kv, GGML_TYPE_F16, GGML_TYPE_F16, f32acc); 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 8fba7a450f7..e308c7214dd 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_base.glsl +++ b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_base.glsl @@ -144,7 +144,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, 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; bool partial_output; void init_indices() @@ -211,9 +211,21 @@ void init_indices() q_stride = gqa_ratio > 1 ? (p.nb02 / 4) : p.nb01; k_stride = p.nb11; v_stride = p.nb21; + // Bit 31 of gqa_ratio means "the mask row stride is in split_kv", used by the DeepSeek V4 + // sparse split path where the mask spans the full K range but this dispatch only covers the + // raw prefix, so m_stride != KV. That path always sets split_k_num == 1, which is what keeps + // split_kv free to carry the stride; ggml_vk_flash_attn_top_k asserts the invariant. + // Otherwise: when using grouped query attention all rows share the same mask (stride 0). + // "p.gqa_ratio >> 16" is just a roundabout way of writing zero that prevents the compiler + // from folding the "&" through the select and breaking the alignment detection. const bool mask_stride_in_split_kv = (p.gqa_ratio & 0x80000000u) != 0; - const uint32_t mask_stride_override = p.gqa_ratio >> 16; - m_stride = mask_stride_in_split_kv ? p.split_kv : (mask_stride_override != 0 ? mask_stride_override : (gqa_ratio > 1 ? 0 : KV)); + m_stride = mask_stride_in_split_kv ? p.split_kv : ((gqa_ratio > 1) ? (p.gqa_ratio >> 16) : KV); + // Distinct from m_stride: m_stride is the row-to-row step INSIDE this tile (0 under GQA, + // where every row shares one mask row), while m_row_len is the mask tensor's actual row + // length, used to step between tokens and between streams. They differ under GQA and + // 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; } // 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 e3ea909a59c..38ae42c81f0 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_cm1.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_cm1.comp @@ -139,9 +139,9 @@ void main() { // FaBlockBytesK/V == 2 for f16 (sizeof f16) and == 16 for f32 (vec4) and == ggml block size for quants. uint32_t k_offset = (ik2*p.nb12 + ik3*p.nb13) / FaBlockBytesK; uint32_t v_offset = (iv2*p.nb22 + iv3*p.nb23) / FaBlockBytesV; - uint32_t m_offset = gqa_iq1*KV; + uint32_t m_offset = gqa_iq1*m_row_len; if (p.nem2 != 1 || p.nem3 != 1) { - m_offset += ((iq3 % p.nem3) * p.nem2 + (iq2 % p.nem2)) * p.nem1 * KV; + m_offset += ((iq3 % p.nem3) * p.nem2 + (iq2 % p.nem2)) * p.nem1 * m_row_len; mo_offset += ((iq3 % p.nem3) * p.nem2 + (iq2 % p.nem2)) * CEIL_DIV(p.nem1, Br) * mo_stride; } diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_top_k_cm.comp b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_top_k_cm.comp index 3d0024553e7..e0628c65705 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_top_k_cm.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_top_k_cm.comp @@ -276,8 +276,15 @@ void main() { const uint dim = out_idx % HEAD_SIZE; if (p.split_mode != 0) { const uint part_idx = 1; + // Split-K buffer layout, shared with flash_attn_cm1.comp and + // flash_attn_split_k_reduce.comp: all O matrices [HSV, ne1, k_num, ne2, ne3] + // first, then the L/M rows [ne1, k_num, ne2, ne3]. Here ne1 = n_head, + // ne2 = n_batch and ne3 = the stream count, so the L/M base must span every + // stream -- omitting gl_NumWorkGroups.z lands L/M inside the O region as soon + // as there is more than one sequence. + const uint n_streams = gl_NumWorkGroups.z; const uint matrix_base = HEAD_SIZE * p.n_head * (part_idx + 2 * (token + p.n_batch * stream)); - const uint lm_base = HEAD_SIZE * p.n_head * p.n_batch * 2 + p.n_head * 2 * (part_idx + 2 * (token + p.n_batch * stream)); + const uint lm_base = HEAD_SIZE * p.n_head * p.n_batch * n_streams * 2 + p.n_head * 2 * (part_idx + 2 * (token + p.n_batch * stream)); data_dst[matrix_base + (head_base + head_local) * HEAD_SIZE + dim] = accum[i]; if (dim == 0) { data_dst[lm_base + head_base + head_local] = row_sum_sh[head_local]; From 76b18c3a9bf461ad98671f860a6574546b32b358 Mon Sep 17 00:00:00 2001 From: Nathan Wilson Date: Thu, 13 Aug 2026 14:32:23 +0000 Subject: [PATCH 23/68] vulkan: harden the sparse FA split path and drop its debug scaffolding No behaviour change on any path that runs today; this removes footguns the split path left in shared code and clears the leftovers. - flash_attn.comp and flash_attn_cm2.comp still read p.k_num and p.gqa_ratio raw, while the split path packs a partition count into the high half of k_num and a flag into bit 31 of gqa_ratio. Only the coopmat1 gate keeps them from ever seeing those values, and nothing said so. Use the decoded globals init_indices() already computes, so routing the split path at a sibling shader fails loudly instead of writing at wild offsets. - Drop profile_stage and GGML_VK_FA_TOPK_PROFILE. This also removes two continues from the sparse kernel's key loop. - Remove the unreachable mask-stride override branch (nothing sets bits 16..30 without bit 31) and restore the comment explaining why the GQA case writes zero the roundabout way: the compiler must not fold it, or stride alignment detection breaks. - Assert that the mask stride covers the raw KV range. The raw dispatch smuggles that stride through split_kv, which the shader also uses to derive its KV range, so a stride below KV would silently clip it. - Request descriptor sets only after the split scratch is known to fit, so the fallback path does not inherit n_tiles of unused requests. - n_head_log2 is unread with max_bias == 0; pass 0 rather than a number that looks computed. - Make the crossover depend on the path that will run. The coopmat sparse path wins as soon as any key is pruned, but the scalar fallback does not and regresses against dense FA until the pruned fraction is large, so it keeps its original 3x margin. - Record that the coopmat sparse shader requires 64-wide subgroups and eight of them, at the constants that encode it. Co-Authored-By: Claude Opus 5 Assisted-by: Claude (Opus 5) --- ggml/src/ggml-vulkan/ggml-vulkan.cpp | 30 ++++++++++++------- .../vulkan-shaders/flash_attn.comp | 20 ++++++------- .../vulkan-shaders/flash_attn_cm2.comp | 24 +++++++-------- .../vulkan-shaders/flash_attn_top_k.comp | 1 - .../vulkan-shaders/flash_attn_top_k_cm.comp | 15 ++++------ 5 files changed, 48 insertions(+), 42 deletions(-) diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp index 4342e78f2a6..6e4d10f69f2 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp +++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp @@ -1997,7 +1997,6 @@ struct vk_op_flash_attn_top_k_push_constants { uint32_t nb1, nb2, nb3; float scale; uint32_t has_sinks; - uint32_t profile_stage; uint32_t split_mode; }; static_assert(sizeof(vk_op_flash_attn_top_k_push_constants) <= 128); @@ -11460,7 +11459,14 @@ static bool ggml_vk_flash_attn_top_k(ggml_backend_vk_context * ctx, vk_context & return false; } const int64_t n_kv_active = n_kv_raw + top_k->ne[0]; - if (k->ne[1] <= n_kv_active) { + // Crossover against ordinary dense FA. The coopmat sparse path (and especially the + // raw/selected split below) is cheap enough to win as soon as any key is pruned; the + // scalar fallback shader is not, and regresses against dense FA until the pruned + // fraction is large, so it keeps its original 3x margin. Equality always uses dense + // FA because nothing is pruned. + const bool have_cm_sparse = ctx->device->pipeline_flash_attn_top_k_cm_f16 != nullptr; + const int64_t min_total_k = have_cm_sparse ? n_kv_active + 1 : 3 * n_kv_active; + if (k->ne[1] < min_total_k) { return false; } @@ -11479,15 +11485,13 @@ static bool ggml_vk_flash_attn_top_k(ggml_backend_vk_context * ctx, vk_context & (uint32_t) (dst->nb[1] / sizeof(float)), (uint32_t) (dst->nb[2] / sizeof(float)), (uint32_t) (dst->nb[3] / sizeof(float)), - scale, sinks != nullptr, 0, 0, + scale, sinks != nullptr, 0, }; const vk_subbuffer q_buf = ggml_vk_tensor_subbuffer(ctx, q); const vk_subbuffer sinks_buf = sinks ? ggml_vk_tensor_subbuffer(ctx, sinks) : q_buf; static const char * top_k_cm_env = getenv("GGML_VK_FA_TOPK_CM"); const bool use_cm = (!top_k_cm_env || top_k_cm_env[0] != '0') && ctx->device->pipeline_flash_attn_top_k_cm_f16; - static const char * top_k_profile_env = getenv("GGML_VK_FA_TOPK_PROFILE"); - pc.profile_stage = use_cm && top_k_profile_env ? atoi(top_k_profile_env) : 0; vk_pipeline pipeline = use_cm ? ctx->device->pipeline_flash_attn_top_k_cm_f16 : ctx->device->pipeline_flash_attn_top_k_f16; static const char * top_k_split_env = getenv("GGML_VK_FA_TOPK_SPLIT"); @@ -11527,13 +11531,12 @@ static bool ggml_vk_flash_attn_top_k(ggml_backend_vk_context * ctx, vk_context & pipelines[raw_state] = raw_pipeline = std::make_shared(); } } - ggml_pipeline_request_descriptor_sets(ctx, raw_pipeline, n_tiles); - ggml_pipeline_request_descriptor_sets(ctx, pipeline, n_tiles); - ggml_pipeline_request_descriptor_sets(ctx, ctx->device->pipeline_flash_attn_split_k_reduce, n_tiles); - const uint64_t partition_size = ((uint64_t) D * NH + NH * 2) * sizeof(float) * tile_size * NS; const uint64_t split_size = partition_size * partitions; if (split_size <= ctx->device->properties.limits.maxStorageBufferRange) { + ggml_pipeline_request_descriptor_sets(ctx, raw_pipeline, n_tiles); + ggml_pipeline_request_descriptor_sets(ctx, pipeline, n_tiles); + ggml_pipeline_request_descriptor_sets(ctx, ctx->device->pipeline_flash_attn_split_k_reduce, n_tiles); if (ctx->prealloc_size_split_k < split_size) { ctx->prealloc_size_split_k = split_size; ggml_vk_preallocate_buffers(ctx, subctx); @@ -11542,7 +11545,14 @@ static bool ggml_vk_flash_attn_top_k(ggml_backend_vk_context * ctx, vk_context & ggml_vk_sync_buffers(ctx, subctx); } - const uint32_t n_head_log2 = 64; + // ALiBi is disabled (max_bias == 0), so n_head_log2 is never read; keep it 0 + // rather than a value that looks computed. + const uint32_t n_head_log2 = 0; + // The raw dispatch smuggles the mask row stride through split_kv, which the FA + // shader also uses to derive its KV range as min(KV, (split_k_index+1)*split_kv). + // That is only safe while split_k_index == 0 and the stride covers the whole raw + // prefix -- both hold here, but assert rather than rely on it silently. + GGML_ASSERT(mask_stride >= raw_kv && "split_kv carries the mask stride; it must not clip the raw KV range"); const uint32_t mask_stride_in_split_kv = 1u << 31; const uint32_t packed_gqa = mask_stride_in_split_kv | 1u; const uint32_t packed_partitions = (partitions << 16) | 1; diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn.comp b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn.comp index 18a37add9bf..5b19ff61093 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn.comp @@ -186,9 +186,9 @@ void main() { // FaBlockBytesK/V == 2 for f16, 16 for f32, ggml block byte size for quants. uint32_t k_offset = (ik2*p.nb12 + ik3*p.nb13) / FaBlockBytesK; uint32_t v_offset = (iv2*p.nb22 + iv3*p.nb23) / FaBlockBytesV; - uint32_t m_offset = gqa_iq1*KV; + uint32_t m_offset = gqa_iq1*m_row_len; if (p.nem2 != 1 || p.nem3 != 1) { - m_offset += ((iq3 % p.nem3) * p.nem2 + (iq2 % p.nem2)) * p.nem1 * KV; + m_offset += ((iq3 % p.nem3) * p.nem2 + (iq2 % p.nem2)) * p.nem1 * m_row_len; mo_offset += ((iq3 % p.nem3) * p.nem2 + (iq2 % p.nem2)) * CEIL_DIV(p.nem1, Br) * mo_stride; } @@ -210,7 +210,7 @@ void main() { } // Only load if the block is not all zeros if (mask_opt_bits != MASK_OPT_ALL_ZERO) { - bool nem1_bounds_check = !(p.gqa_ratio > 1) && (p.nem1 % Br) != 0; + bool nem1_bounds_check = !(gqa_ratio > 1) && (p.nem1 % Br) != 0; float max_mask = NEG_FLT_MAX_OVER_2; barrier(); @@ -660,10 +660,10 @@ void main() { // If there is split_k, then the split_k resolve shader does the final // division by L. Store the intermediate O value and per-row m and L values. - if (p.k_num > 1) { - if (p.gqa_ratio > 1) { + if (partial_output || split_k_num > 1) { + if (gqa_ratio > 1) { // note: O and Q have swapped coord 1,2. - uint32_t o_offset = HSV * p.ne1 * (split_k_index + p.k_num * (gqa_iq1 + p.ne2 * iq3)) / 4; + uint32_t o_offset = HSV * p.ne1 * (split_k_index + output_k_num * (gqa_iq1 + p.ne2 * iq3)) / 4; [[unroll]] for (uint32_t r = 0; r < rows_per_thread; ++r) { const uint row = tile_row(r); @@ -674,7 +674,7 @@ void main() { } } - o_offset = HSV * p.ne1 * p.k_num * p.ne2 * p.ne3 + p.ne1 * 2 * (split_k_index + p.k_num * (gqa_iq1 + p.ne2 * iq3)); + o_offset = HSV * p.ne1 * output_k_num * p.ne2 * p.ne3 + p.ne1 * 2 * (split_k_index + output_k_num * (gqa_iq1 + p.ne2 * iq3)); [[unroll]] for (uint32_t r = 0; r < rows_per_thread; ++r) { const uint row = tile_row(r); if (row < N) { @@ -688,7 +688,7 @@ void main() { const uint global_row = i * Br + row; if (global_row < N) { - uint32_t o_offset = HSV * p.ne1 * (split_k_index + p.k_num * (global_row + p.ne2 * iq3)) / 4; + uint32_t o_offset = HSV * p.ne1 * (split_k_index + output_k_num * (global_row + p.ne2 * iq3)) / 4; [[unroll]] for (uint32_t d = 0; d < HSV_per_thread / 4; ++d) { data_ov4[o_offset + iq2 * HSV/4 + d * D_split + d_tid] = D_TYPEV4(Of[r][d]); @@ -696,7 +696,7 @@ void main() { } if (global_row < N && d_tid == 0 && col_tid == 0) { - uint32_t lm_offset = HSV * p.ne1 * p.k_num * p.ne2 * p.ne3 + p.ne1 * 2 * (split_k_index + p.k_num * (global_row + p.ne2 * iq3)); + uint32_t lm_offset = HSV * p.ne1 * output_k_num * p.ne2 * p.ne3 + p.ne1 * 2 * (split_k_index + output_k_num * (global_row + p.ne2 * iq3)); data_o[lm_offset + iq2] = D_TYPE(Lf[r]); data_o[lm_offset + p.ne1 + iq2] = D_TYPE(Mf[r]); } @@ -742,7 +742,7 @@ void main() { uint32_t o_offset = (gqa_iq1*p.ne1*HSV + iq3*p.ne2*p.ne1*HSV) / 4; - if (p.gqa_ratio > 1) { + if (gqa_ratio > 1) { [[unroll]] for (uint32_t r = 0; r < rows_per_thread; ++r) { const uint row = tile_row(r); if (row < N) { 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 31741115308..54be1e6daa9 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_cm2.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_cm2.comp @@ -151,7 +151,7 @@ D_TYPE perElemOpNonGqaSplitKStore(const in uint32_t r, const in uint32_t c, cons uint32_t global_row = i * Br + r; if (global_row < N && c < HSV) { uint32_t o_off = HSV * p.ne1 - * (split_k_index + p.k_num * (global_row + p.ne2 * iq3)); + * (split_k_index + output_k_num * (global_row + p.ne2 * iq3)); data_o[o_off + iq2 * HSV + c] = D_TYPE(elem); } return elem; @@ -161,8 +161,8 @@ D_TYPE perElemOpNonGqaSplitKStore(const in uint32_t r, const in uint32_t c, cons ACC_TYPE perElemOpNonGqaSplitKStoreCol0(const in uint32_t r, const in uint32_t c, const in ACC_TYPE elem, const in uint32_t lm_base, const in uint32_t iq2, const in uint32_t N) { uint32_t global_row = i * Br + r; if (global_row < N && c == 0) { - uint32_t lm_off = HSV * p.ne1 * p.k_num * p.ne2 * p.ne3 - + p.ne1 * 2 * (split_k_index + p.k_num * (global_row + p.ne2 * iq3)); + uint32_t lm_off = HSV * p.ne1 * output_k_num * p.ne2 * p.ne3 + + p.ne1 * 2 * (split_k_index + output_k_num * (global_row + p.ne2 * iq3)); data_o[lm_off + lm_base + iq2] = D_TYPE(elem); } return elem; @@ -242,9 +242,9 @@ void main() { // mo_offset will point to the tile starting at row i*Br and col 0 uint32_t mo_offset = mo_stride * i; - uint32_t m_offset = gqa_iq1*KV * 2 /*sizeof(float16_t)*/; + uint32_t m_offset = gqa_iq1*m_row_len * 2 /*sizeof(float16_t)*/; if (p.nem2 != 1 || p.nem3 != 1) { - m_offset += ((iq3 % p.nem3) * p.nem2 + (iq2 % p.nem2)) * p.nem1 * KV * 2 /*sizeof(float16_t)*/; + m_offset += ((iq3 % p.nem3) * p.nem2 + (iq2 % p.nem2)) * p.nem1 * m_row_len * 2 /*sizeof(float16_t)*/; mo_offset += ((iq3 % p.nem3) * p.nem2 + (iq2 % p.nem2)) * CEIL_DIV(p.nem1, Br) * mo_stride; } @@ -268,7 +268,7 @@ void main() { } // Only load if the block is not all zeros if (mask_opt_bits != MASK_OPT_ALL_ZERO) { - bool nem1_bounds_check = !(p.gqa_ratio > 1) && (p.nem1 % Br) != 0; + bool nem1_bounds_check = !(gqa_ratio > 1) && (p.nem1 % Br) != 0; if (nem1_bounds_check) { tensorLayoutNV<2, gl_CooperativeMatrixClampModeConstantNV> tensorLayoutM = createTensorLayoutNV(2, gl_CooperativeMatrixClampModeConstantNV); @@ -287,7 +287,7 @@ void main() { } else { tensorLayoutNV<2, Clamp> tensorLayoutM = createTensorLayoutNV(2, Clamp); // Don't clamp against nem1 when GQA is enabled - uint32_t m_height = p.gqa_ratio > 1 ? ~0 : p.nem1; + uint32_t m_height = gqa_ratio > 1 ? ~0 : p.nem1; tensorLayoutM = setTensorLayoutDimensionNV(tensorLayoutM, m_height, KV); tensorLayoutM = setTensorLayoutStrideNV(tensorLayoutM, m_stride, 1); @@ -406,15 +406,15 @@ void main() { // If there is split_k, then the split_k resolve shader does the final // division by L. Store the intermediate O value and per-row m and L values. - if (p.k_num > 1) { + if (partial_output || split_k_num > 1) { coopmat O_D = coopmat(O); - if (p.gqa_ratio > 1) { + if (gqa_ratio > 1) { // note: O and Q have swapped coord 1,2. - uint32_t o_offset = HSV * p.ne1 * (split_k_index + p.k_num * (gqa_iq1 + p.ne2 * iq3)); + uint32_t o_offset = HSV * p.ne1 * (split_k_index + output_k_num * (gqa_iq1 + p.ne2 * iq3)); coopMatPerElementNV(O_D, O_D, perElemOpGqaStore, o_offset, iq2, N); - o_offset = HSV * p.ne1 * p.k_num * p.ne2 * p.ne3 + p.ne1 * 2 * (split_k_index + p.k_num * (gqa_iq1 + p.ne2 * iq3)); + o_offset = HSV * p.ne1 * output_k_num * p.ne2 * p.ne3 + p.ne1 * 2 * (split_k_index + output_k_num * (gqa_iq1 + p.ne2 * iq3)); coopMatPerElementNV(L, L, perElemOpStoreCol0, o_offset, iq2, N); coopMatPerElementNV(M, M, perElemOpStoreCol0, o_offset + p.ne1, iq2, N); } else { @@ -473,7 +473,7 @@ void main() { uint32_t o_offset = gqa_iq1*p.ne1*HSV + iq3*p.ne2*p.ne1*HSV; - if (p.gqa_ratio > 1) { + if (gqa_ratio > 1) { coopMatPerElementNV(O_D, O_D, perElemOpGqaStore, o_offset, iq2, N); } else { tensorLayoutNV<3, gl_CooperativeMatrixClampModeConstantNV> tensorLayoutD = createTensorLayoutNV(3, gl_CooperativeMatrixClampModeConstantNV); diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_top_k.comp b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_top_k.comp index 62807471bdd..930777b5f0e 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_top_k.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_top_k.comp @@ -37,7 +37,6 @@ layout(push_constant) uniform Parameters { uint nb3; float scale; uint has_sinks; - uint profile_stage; uint split_mode; } p; diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_top_k_cm.comp b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_top_k_cm.comp index e0628c65705..7a328938777 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_top_k_cm.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_top_k_cm.comp @@ -39,10 +39,15 @@ layout(push_constant) uniform Parameters { uint nb3; float scale; uint has_sinks; - uint profile_stage; uint split_mode; } p; +// This shader is hard-wired to a 512-thread workgroup of eight 64-wide subgroups: +// the softmax segments SUBGROUP_SIZE/4 = 16 lanes per head (so head_local stays < 32), +// and the QK/PV tiling assumes gl_NumSubgroups == 8 via gl_SubgroupID % 4 / gl_SubgroupID / 4. +// A different subgroup size indexes past row_max_sh/q_p_sh and corrupts the score tiles, so +// the pipeline is only created under `device->subgroup_size == 64` in ggml_vk_load_shaders. +// Keep that gate in sync with these constants. const uint TILE = 16; const uint HEAD_SIZE = 512; const uint HEADS_PER_GROUP = 32; @@ -146,10 +151,6 @@ void main() { SCORE_STRIDE, gl_CooperativeMatrixLayoutRowMajor); barrier(); - if (p.profile_stage == 1) { - continue; - } - { const uint head_local = tid / (SUBGROUP_SIZE / 4); const uint softmax_lane = tid % (SUBGROUP_SIZE / 4); @@ -198,10 +199,6 @@ void main() { accum[i] *= old_scale_sh[head_local]; } - if (p.profile_stage == 2) { - continue; - } - coopmat pmats[KEYS_PER_BLOCK / TILE]; [[unroll]] for (uint key_chunk = 0; key_chunk < KEYS_PER_BLOCK / TILE; ++key_chunk) { const uint pv_head_tile = gl_SubgroupID / (DIMS_PER_BLOCK / TILE); From 75e195fcdd37357ba0dc96b616d5fa00d5c48b78 Mon Sep 17 00:00:00 2001 From: Nathan Wilson Date: Thu, 13 Aug 2026 14:32:59 +0000 Subject: [PATCH 24/68] test-backend-ops: cover sparse top-k FA with more than one sequence The sparse top-k fixture pinned every tensor's fourth dimension to 1, so no case reached the split path's stream indexing and three separate mis-indexings passed the suite. Add an ns parameter and three cases: two at ns=2 (with and without sinks, single tile) and one at ns=3 with 300 query tokens, which also crosses the 256-token tile boundary. The per-token selection is offset by the stream so a dropped stream stride reads another sequence's keys rather than the same ones. All three fail on the unfixed shaders and pass with them fixed. Co-Authored-By: Claude Opus 5 Assisted-by: Claude (Opus 5) --- tests/test-backend-ops.cpp | 56 +++++++++++++++++++++++--------------- 1 file changed, 34 insertions(+), 22 deletions(-) diff --git a/tests/test-backend-ops.cpp b/tests/test-backend-ops.cpp index 23913f62700..23b30958db3 100644 --- a/tests/test-backend-ops.cpp +++ b/tests/test-backend-ops.cpp @@ -8038,12 +8038,13 @@ struct test_flash_attn_ext_top_k : public test_case { const int64_t n_kv_raw; // dense prefix always attended const int64_t n_top_k; // selected keys per query token const bool sinks; + const int64_t ns; // sequences (ne3); >1 exercises the split-K stream stride static constexpr int64_t hs = 512; // V4 CSA head size, K == V latent static constexpr int64_t nh = 64; // V4 CSA query heads (MQA) std::string vars() override { - return VARS_TO_STR5(kv, nb, n_kv_raw, n_top_k, sinks); + return VARS_TO_STR6(kv, nb, n_kv_raw, n_top_k, sinks, ns); } double max_nmse_err() override { @@ -8054,27 +8055,27 @@ struct test_flash_attn_ext_top_k : public test_case { GGML_UNUSED(t); // only the active keys contribute compute on a sparse backend; count those so // perf mode reports the useful-work rate - return 2 * nh * nb * (hs + hs) * (n_kv_raw + n_top_k); + return 2 * nh * nb * ns * (hs + hs) * (n_kv_raw + n_top_k); } - test_flash_attn_ext_top_k(int64_t kv = 768, int64_t nb = 8, int64_t n_kv_raw = 64, int64_t n_top_k = 128, bool sinks = false) - : kv(kv), nb(nb), n_kv_raw(n_kv_raw), n_top_k(n_top_k), sinks(sinks) {} + test_flash_attn_ext_top_k(int64_t kv = 768, int64_t nb = 8, int64_t n_kv_raw = 64, int64_t n_top_k = 128, bool sinks = false, int64_t ns = 1) + : kv(kv), nb(nb), n_kv_raw(n_kv_raw), n_top_k(n_top_k), sinks(sinks), ns(ns) {} ggml_tensor * build_graph(ggml_context * ctx) override { - ggml_tensor * q = ggml_new_tensor_4d(ctx, GGML_TYPE_F32, hs, nb, nh, 1); + ggml_tensor * q = ggml_new_tensor_4d(ctx, GGML_TYPE_F32, hs, nb, nh, ns); ggml_set_name(q, "q"); - ggml_tensor * k = ggml_new_tensor_4d(ctx, GGML_TYPE_F16, hs, kv, 1, 1); + ggml_tensor * k = ggml_new_tensor_4d(ctx, GGML_TYPE_F16, hs, kv, 1, ns); ggml_set_name(k, "k"); // V4 CSA attends over the K latent itself: V is the same cache tensor - ggml_tensor * v = ggml_view_4d(ctx, k, hs, kv, 1, 1, k->nb[1], k->nb[2], k->nb[3], 0); + ggml_tensor * v = ggml_view_4d(ctx, k, hs, kv, 1, ns, k->nb[1], k->nb[2], k->nb[3], 0); ggml_set_name(v, "v"); - ggml_tensor * m = ggml_new_tensor_4d(ctx, GGML_TYPE_F16, kv, nb, 1, 1); + ggml_tensor * m = ggml_new_tensor_4d(ctx, GGML_TYPE_F16, kv, nb, 1, ns); ggml_set_name(m, "m"); - ggml_tensor * t = ggml_new_tensor_4d(ctx, GGML_TYPE_I32, n_top_k, nb, 1, 1); + ggml_tensor * t = ggml_new_tensor_4d(ctx, GGML_TYPE_I32, n_top_k, nb, 1, ns); ggml_set_name(t, "top_k"); ggml_tensor * s = nullptr; @@ -8109,23 +8110,29 @@ struct test_flash_attn_ext_top_k : public test_case { // build a consistent (top_k, mask) pair: a deterministic per-token selection, // strided so adjacent tokens select overlapping-but-different keys, with one // deliberately invalid index (-1) whose mask slot stays -inf - std::vector top(n_top_k * nb); - std::vector mask(kv * nb); + std::vector top(n_top_k * nb * ns); + std::vector mask(kv * nb * ns); const ggml_fp16_t minus_inf = ggml_fp32_to_fp16(-INFINITY); const ggml_fp16_t zero = ggml_fp32_to_fp16(0.0f); - for (int64_t b = 0; b < nb; ++b) { - for (int64_t i = 0; i < kv; ++i) { - mask[b * kv + i] = i < n_kv_raw ? zero : minus_inf; - } - for (int64_t j = 0; j < n_top_k; ++j) { - int32_t idx = (int32_t) ((j * range) / n_top_k + b) % (int32_t) range; - if (j == n_top_k - 1 && b == 0) { - idx = -1; // exercise the ignore-invalid-index path - } else { - mask[b * kv + n_kv_raw + idx] = zero; + for (int64_t s = 0; s < ns; ++s) { + for (int64_t b = 0; b < nb; ++b) { + const int64_t mrow = (s * nb + b) * kv; + const int64_t trow = (s * nb + b) * n_top_k; + for (int64_t i = 0; i < kv; ++i) { + mask[mrow + i] = i < n_kv_raw ? zero : minus_inf; + } + for (int64_t j = 0; j < n_top_k; ++j) { + // offset the selection by the stream too, so a dropped stream stride + // reads another sequence's keys and shows up as a mismatch + int32_t idx = (int32_t) ((j * range) / n_top_k + b + s * 7) % (int32_t) range; + if (j == n_top_k - 1 && b == 0 && s == 0) { + idx = -1; // exercise the ignore-invalid-index path + } else { + mask[mrow + n_kv_raw + idx] = zero; + } + top[trow + j] = idx; } - top[b * n_top_k + j] = idx; } } @@ -11398,6 +11405,11 @@ static std::vector> make_test_cases_eval() { test_cases.emplace_back(new test_flash_attn_ext_top_k(1024, 64, 65, 128, false)); test_cases.emplace_back(new test_flash_attn_ext_top_k(4096, 128, 256, 512, false)); test_cases.emplace_back(new test_flash_attn_ext_top_k(4096, 257, 256, 512, false)); + // ns > 1: the split-K partial-output path indexes O and L/M by stream, so these cover + // the stream stride in both regions (single tile and multi-tile). + test_cases.emplace_back(new test_flash_attn_ext_top_k( 768, 64, 64, 128, false, 2)); + test_cases.emplace_back(new test_flash_attn_ext_top_k( 768, 64, 64, 128, true, 2)); + test_cases.emplace_back(new test_flash_attn_ext_top_k(4096, 300, 256, 512, false, 3)); return test_cases; } From 04311a08e5acdc0152ba9c3bcc624fdd1a764899 Mon Sep 17 00:00:00 2001 From: Nathan Wilson Date: Thu, 13 Aug 2026 15:34:07 +0000 Subject: [PATCH 25/68] vulkan: let sparse FA query tiling and multiple sequences coexist flash_attn_split_k_reduce was using one value for two different things: ne2 is the split buffer's query count, which the sparse path hands the TILE height, but the destination's stream stride needs the FULL query count. With several tiles and more than one sequence it wrote stream s at s*tile_size instead of s*N, so the previous fix simply refused to tile whenever there was more than one stream. Give the reduce a separate dst_ne2 and use it for the destination index only. Ordinary FA split-K passes ne2 for it and is bit-identical. The sparse path passes the full batch, so tiling is unconditional again. This matters for memory, which is the binding constraint on this model. The scratch is capped at 256 queries per tile rather than scaling with the batch, so at ub2048 a 2-sequence context drops from 1028 MB to 128.5 MB and a 4-sequence one from 2056 MB to 257 MB. It also removes a cliff: 8 sequences at ub2048 would have exceeded maxStorageBufferRange and silently fallen back to the slower non-split kernel. Sparse FA timings are unchanged (every shape within 0.45%, inside run-to-run spread). 13,308 FLASH_ATTN_EXT cases pass, which covers the ordinary split-K path that shares this shader. Co-Authored-By: Claude Opus 5 Assisted-by: Claude (Opus 5) --- ggml/src/ggml-vulkan/ggml-vulkan.cpp | 18 ++++++++++-------- .../flash_attn_split_k_reduce.comp | 3 ++- 2 files changed, 12 insertions(+), 9 deletions(-) diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp index 6e4d10f69f2..a4abb75d56b 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp +++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp @@ -2152,7 +2152,11 @@ struct vk_quantize_q8_1_push_constants { struct vk_op_flash_attn_split_k_reduce_push_constants { uint32_t D; uint32_t ne1; + // ne2 describes the SPLIT BUFFER (which may cover only a tile of queries); dst_ne2 is the + // destination's query count. They differ only when the caller tiles the split buffer, and + // the destination stride between streams must always use the full count. uint32_t ne2; + uint32_t dst_ne2; uint32_t ne3; uint32_t k_num; uint32_t sinks; @@ -11504,12 +11508,10 @@ static bool ggml_vk_flash_attn_top_k(ggml_backend_vk_context * ctx, vk_context & const uint32_t NS = (uint32_t) q->ne[3]; const uint32_t raw_kv = (uint32_t) n_kv_raw; const uint32_t partitions = 2; - // Query tiling keeps the split scratch small, but flash_attn_split_k_reduce derives the - // destination row from ne2, which it is handed as the TILE height. With one tile that - // equals N and the stream stride is right; with several tiles and more than one stream - // it would write stream s at s*tile_size instead of s*N. Only tile when there is a - // single stream, which is the prefill case the tiling exists for. - const uint32_t tile_size = (NS == 1) ? std::min(N, 256u) : N; + // Query tiling caps the split scratch at 256 queries regardless of batch or stream + // count. The reduce takes the tile height as ne2 and the full query count as dst_ne2, + // so the destination stream stride stays correct across tiles. + const uint32_t tile_size = std::min(N, 256u); const uint32_t n_tiles = CEIL_DIV(N, tile_size); const bool f32acc = true; vk_fa_tuning_params tuning = get_fa_tuning_params(ctx->device, D, D, N, raw_kv, GGML_TYPE_F16, GGML_TYPE_F16, f32acc); @@ -11604,7 +11606,7 @@ static bool ggml_vk_flash_attn_top_k(ggml_backend_vk_context * ctx, vk_context & ctx->prealloc_split_k_need_sync = true; ggml_vk_sync_buffers(ctx, subctx); - const vk_op_flash_attn_split_k_reduce_push_constants reduce_pc = {D, NH, tile_n, NS, partitions, sinks != nullptr}; + const vk_op_flash_attn_split_k_reduce_push_constants reduce_pc = {D, NH, tile_n, N, NS, partitions, sinks != nullptr}; ggml_vk_dispatch_pipeline(ctx, subctx, ctx->device->pipeline_flash_attn_split_k_reduce, {split_buf, sinks_buf, tile_dst}, reduce_pc, {NH, D, tile_n * NS}); ctx->prealloc_split_k_need_sync = true; @@ -12121,7 +12123,7 @@ static void ggml_vk_flash_attn(ggml_backend_vk_context * ctx, vk_context& subctx pc, { dispatch_x, workgroups_y, workgroups_z }); ggml_vk_sync_buffers(ctx, subctx); - const vk_op_flash_attn_split_k_reduce_push_constants pc2 = { HSV, (uint32_t)ne1, (uint32_t)ne2, (uint32_t)ne3, split_k, (sinks != nullptr) }; + const vk_op_flash_attn_split_k_reduce_push_constants pc2 = { HSV, (uint32_t)ne1, (uint32_t)ne2, (uint32_t)ne2, (uint32_t)ne3, split_k, (sinks != nullptr) }; ggml_vk_dispatch_pipeline(ctx, subctx, ctx->device->pipeline_flash_attn_split_k_reduce, {split_k_buf, sinks_buf, dst_buf}, pc2, { (uint32_t)ne1, HSV, (uint32_t)(ne2 * ne3) }); diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_split_k_reduce.comp b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_split_k_reduce.comp index 68917fc0bb0..69342ba3cee 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_split_k_reduce.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_split_k_reduce.comp @@ -14,6 +14,7 @@ layout (push_constant) uniform parameter { uint D; uint ne1; uint ne2; + uint dst_ne2; uint ne3; uint k_num; uint sinks; @@ -116,6 +117,6 @@ void main() { const float FLT_MAX = uintBitsToFloat(0x7F7FFFFF); O = clamp(O, -FLT_MAX, FLT_MAX); - data_d[(i3 * p.ne2 + i2) * p.ne1 * D + D * n + d] = O; + data_d[(i3 * p.dst_ne2 + i2) * p.ne1 * D + D * n + d] = O; } } From feff28ebf178472cb7ab044e3ff869c1197de4a2 Mon Sep 17 00:00:00 2001 From: Jaap Buurman Date: Thu, 13 Aug 2026 20:00:25 +0200 Subject: [PATCH 26/68] vulkan: parallelize DSV4 Lightning Indexer prefill Assisted-by: Codex --- .../DSV4-vulkan-lightning-indexer-progress.md | 158 ++++++++++++++++++ ggml/src/ggml-vulkan/ggml-vulkan.cpp | 18 +- .../vulkan-shaders/lightning_indexer_cm.comp | 94 +++++++---- .../vulkan-shaders/vulkan-shaders-gen.cpp | 3 +- tests/test-backend-ops.cpp | 9 + 5 files changed, 241 insertions(+), 41 deletions(-) create mode 100644 docs/development/DSV4-vulkan-lightning-indexer-progress.md diff --git a/docs/development/DSV4-vulkan-lightning-indexer-progress.md b/docs/development/DSV4-vulkan-lightning-indexer-progress.md new file mode 100644 index 00000000000..d19fbd6903a --- /dev/null +++ b/docs/development/DSV4-vulkan-lightning-indexer-progress.md @@ -0,0 +1,158 @@ +# DeepSeek V4 Vulkan Lightning Indexer progress + +This is a restart note for the Strix Halo Lightning Indexer optimization. It is a development scratch pad and can be removed before the final PR. + +## Repository state + +- Main repository: `/home/jaap/Projects/git/llama.cpp` +- Optimization worktree: `/tmp/llama-strix-beta-bench` +- Branch: `strix-halo-vulkan-lightning-indexer` +- Base commit: `316c72ee9eab590f5891089d3b6bfc0d01d00d19` +- Base branch: Nathan's `strix-halo-vulkan-beta` +- Decode microbench work is stored in the main worktree as `stash@{0}: On strix-halo-vulkan: wip: DSV4 decode microbench depth matrix`. +- The Indexer changes are uncommitted. Do not commit without explicit user approval. An assisted commit needs an `Assisted-by:` trailer. +- Do not run builds and GPU benchmarks together. The APU shares its power and memory-bandwidth budget. +- GPU commands need sandbox escalation. + +## Objective and result + +After sparse prefill attention was flattened, the context-dependent Lightning Indexer became the next prefill bottleneck. The old cooperative-matrix shader used one wave64 subgroup per workgroup, processed one 16-key tile, and loaded one query head at a time. + +The new wide pipeline uses eight wave64 subgroups per workgroup. Each subgroup processes a separate 16-key tile, so one workgroup covers 128 keys. It stages four query heads and their weights together, reuses them across all eight subgroups, and uses subgroup-scoped synchronization between cooperative-matrix result stores. A workgroup barrier remains between four-head groups because all subgroups reuse the shared query storage. + +The optimized shader requires 512 workgroup invocations and 64 KiB shared memory. Pipeline creation is capability-based. Devices without those limits use a one-wave, one-head cooperative-matrix specialization. The scalar implementation remains the fallback when cooperative matrices are unavailable. The decode-specific cooperative-matrix pipeline is unchanged. + +At the 32k-equivalent prefill microbench shape: + +| Version | Time per layer | Throughput | +| --- | ---: | ---: | +| Baseline | 50.51 ms | 5.83 TFLOPS | +| Optimized | 31.81 ms | 9.25 TFLOPS | + +This is a 37.0% reduction in Lightning Indexer kernel time. + +The canonical 32k llama-bench improved from 209.45 to 216.32 tokens/s. Total Vulkan time fell from 9.73729 to 9.42448 seconds. Total Lightning Indexer time fell from 1.11860 to 0.697276 seconds. Sparse attention and top-K were effectively unchanged. + +## Changed files + +- `ggml/src/ggml-vulkan/vulkan-shaders/lightning_indexer_cm.comp`: parameterizes the shader, adds the eight-wave four-head implementation, and remains usable for the small fallback. +- `ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp`: generates wide `N_WAVES=8`, `HEADS_PER_TILE=4` and small `N_WAVES=1`, `HEADS_PER_TILE=1` variants. +- `ggml/src/ggml-vulkan/ggml-vulkan.cpp`: creates and selects the capability-gated wide pipeline and the small cooperative-matrix fallback. +- `tests/test-backend-ops.cpp`: adds 127, 128, and 129-key correctness boundaries and PP2048 performance shapes through 512k simulated source context. + +## Performance data + +The performance rows model the actual PP2048 Indexer shapes after the source context is filled in 2048-token batches. `kv=8704` is the measured shape near 32k source context. It differs from 32768 because the Indexer compresses source tokens into rows. + +| Source depth | `kv` | Baseline | Optimized | Reduction | +| ---: | ---: | ---: | ---: | ---: | +| 0 | 512 | 4.23 ms | 2.14 ms | 49.5% | +| 8k | 2560 | 17.10 ms | 10.09 ms | 41.0% | +| 16k | 4608 | 29.97 ms | 17.62 ms | 41.2% | +| 32k | 8704 | 50.51 ms | 32.94 ms | 34.8% | +| 64k | 16896 | 94.36 ms | 64.87 ms | 31.3% | +| 128k | 33280 | 174.96 ms | 128.76 ms | 26.4% | +| 256k | 66048 | 352.73 ms | 250.88 ms | 28.9% | +| 512k | 131584 | 696.79 ms | 495.90 ms | 28.8% | + +The final isolated 32k run after cleanup measured 31.81159 ms. Small matrix differences are normal laptop GPU clock variation. + +| Canonical 32k metric | Baseline | Optimized | +| --- | ---: | ---: | +| PP2048 | 209.45 tokens/s | 216.32 tokens/s | +| Total Vulkan | 9.73729 s | 9.42448 s | +| Lightning Indexer | 1.11860 s | 0.697276 s | +| Sparse FA raw | 0.198438 s | 0.195367 s | +| Sparse FA selected | 0.835110 s | 0.843385 s | +| Sparse FA reduce | 0.088463 s | 0.086678 s | +| TOP_K | 0.075473 s | 0.075350 s | + +Logs: + +- Baseline matrix: `/tmp/dsv4-lightning-prefill-baseline.log` +- Optimized matrix: `/tmp/dsv4-lightning-four-head-matrix.log` +- Final selected-pipeline 32k microbench: `/tmp/dsv4-lightning-final-selected-32k.log` +- Baseline canonical 32k llama-bench: `/tmp/dsv4-nathan-beta-32k-rerun-new-first.log` +- Optimized canonical 32k llama-bench: `/tmp/dsv4-lightning-final-32k-llama-bench.log` + +## Correctness and resources + +- Wide pipeline: all 20 focused F16 cases passed, including 127, 128, and 129-key boundaries. +- Small cooperative-matrix fallback: temporarily forced and all the same 20 cases passed. +- Wide shader on gfx1151: 168 VGPRs, 63,488 bytes LDS, no spills, eight subgroups per SIMD. +- The final pipeline-statistics run confirmed `lightning_indexer_cm_f16` was selected. +- No NaN, Inf, or comparison failures were reported. + +Correctness logs: + +- Wide: `/tmp/dsv4-lightning-consolidated-wide-correctness.log` +- Small fallback: `/tmp/dsv4-lightning-consolidated-small-correctness.log` + +## Experiments and decisions + +- Four waves improved the 32k shape about 3% and became worse at deep simulated contexts. +- Eight waves improved it about 8% before the other changes. +- Subgroup-scoped synchronization after cooperative-matrix stores improved the eight-wave version. +- Staging four query heads and weights produced the large gain by reducing redundant loads and barriers. +- Using only a subgroup barrier between head groups failed four boundary tests. One subgroup could overwrite shared query data while another still read it. A workgroup barrier is required there. +- A separate fallback shader source was avoided. Generator definitions create both variants from one file. + +## Commands + +Build only, with no GPU benchmark running: + +```sh +cd /tmp/llama-strix-beta-bench +git diff --check +cmake --build build --config Release --target test-backend-ops llama-bench -j "$(nproc)" +``` + +Focused correctness: + +```sh +cd /tmp/llama-strix-beta-bench +./build/bin/test-backend-ops test -b Vulkan0 -o LIGHTNING_INDEXER -p 'type_K=f16' > /tmp/dsv4-lightning-correctness.log 2>&1 +tail -n 30 /tmp/dsv4-lightning-correctness.log +``` + +Final 32k-equivalent microbench and pipeline selection: + +```sh +cd /tmp/llama-strix-beta-bench +GGML_VK_PIPELINE_STATS=lightning_indexer_cm_f16 ./build/bin/test-backend-ops perf -b Vulkan0 -o LIGHTNING_INDEXER -p 'kv=8704' > /tmp/dsv4-lightning-final-selected-32k.log 2>&1 +tail -n 16 /tmp/dsv4-lightning-final-selected-32k.log +``` + +Full Indexer depth matrix: + +```sh +cd /tmp/llama-strix-beta-bench +./build/bin/test-backend-ops perf -b Vulkan0 -o LIGHTNING_INDEXER -p 'nb=2048,nh=64,ns=1,nm=1,type_K=f16' > /tmp/dsv4-lightning-matrix.log 2>&1 +rg 'kv=(512|2560|4608|8704|16896|33280|66048|131584),nb=2048' /tmp/dsv4-lightning-matrix.log +``` + +Canonical 32k llama-bench. Run it only when needed, never while compiling, and inspect only the final block: + +```sh +cd /tmp/llama-strix-beta-bench +GGML_VK_PERF_LOGGER=1 ./build/bin/llama-bench -m /home/jaap/Projects/docker/localLLaMA/models/models--unsloth--DeepSeek-V4-Flash-0731-GGUF/snapshots/109848da2469efe1f1aab9e11acea08a065ccd4f/UD-IQ3_XXS/DeepSeek-V4-Flash-0731-UD-IQ3_XXS-00001-of-00004.gguf -r 1 -d 32768 -p 2048 -ub 2048 -fa 1 -n 0 > /tmp/dsv4-lightning-32k-llama-bench.log 2>&1 +last=$(grep -n 'Vulkan Timings:' /tmp/dsv4-lightning-32k-llama-bench.log | tail -n 1 | cut -d: -f1) +sed -n "${last},\$p" /tmp/dsv4-lightning-32k-llama-bench.log | tail -n 180 +``` + +Patch inspection: + +```sh +cd /tmp/llama-strix-beta-bench +git diff --check +git diff --stat +git diff +git status --short +``` + +## Next actions + +1. The user reviews and understands the four-file implementation and result summary. +2. Commit only after explicit user approval for that commit action. +3. Remove this scratch pad before a PR if it is not useful as permanent documentation. +4. Restore the separate decode microbench stash from the main worktree only if that work resumes. diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp index a4abb75d56b..81021b6204c 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp +++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp @@ -1106,6 +1106,7 @@ struct vk_device_struct { vk_pipeline pipeline_gated_delta_net[4][2]; vk_pipeline pipeline_lightning_indexer_f16; vk_pipeline pipeline_lightning_indexer_cm_f16; + vk_pipeline pipeline_lightning_indexer_cm_small_f16; vk_pipeline pipeline_lightning_indexer_decode_cm_f16; vk_pipeline pipeline_flash_attn_top_k_f16; vk_pipeline pipeline_flash_attn_top_k_cm_f16; @@ -6284,10 +6285,18 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { device->subgroup_size); #if defined(VK_KHR_cooperative_matrix) && defined(GGML_VULKAN_COOPMAT_GLSLC_SUPPORT) if (device->coopmat_support && device->coopmat_support_16x16x16_f32acc && device->subgroup_size_control) { - ggml_vk_create_pipeline(device, device->pipeline_lightning_indexer_cm_f16, - "lightning_indexer_cm_f16", lightning_indexer_cm_f16_len, lightning_indexer_cm_f16_data, "main", 5, + ggml_vk_create_pipeline(device, device->pipeline_lightning_indexer_cm_small_f16, + "lightning_indexer_cm_small_f16", lightning_indexer_cm_small_f16_len, lightning_indexer_cm_small_f16_data, "main", 5, sizeof(vk_op_lightning_indexer_cm_push_constants), {16, 16, 1}, {device->subgroup_size}, 1, true, true, device->subgroup_size); + if (device->properties.limits.maxComputeWorkGroupInvocations >= 512 && + device->properties.limits.maxComputeWorkGroupSize[0] >= 512 && + device->properties.limits.maxComputeSharedMemorySize >= 64 * 1024) { + ggml_vk_create_pipeline(device, device->pipeline_lightning_indexer_cm_f16, + "lightning_indexer_cm_f16", lightning_indexer_cm_f16_len, lightning_indexer_cm_f16_data, "main", 5, + sizeof(vk_op_lightning_indexer_cm_push_constants), {128, 16, 1}, {device->subgroup_size}, 1, true, true, + device->subgroup_size); + } ggml_vk_create_pipeline(device, device->pipeline_lightning_indexer_decode_cm_f16, "lightning_indexer_decode_cm_f16", lightning_indexer_decode_cm_f16_len, lightning_indexer_decode_cm_f16_data, "main", 5, sizeof(vk_op_lightning_indexer_cm_push_constants), {16, 1, 1}, {device->subgroup_size}, 1, true, true, @@ -12713,8 +12722,9 @@ static vk_pipeline ggml_vk_op_get_pipeline(ggml_backend_vk_context * ctx, const if (ctx->device->pipeline_lightning_indexer_decode_cm_f16 && src0->ne[2] == 1) { return ctx->device->pipeline_lightning_indexer_decode_cm_f16; } - return ctx->device->pipeline_lightning_indexer_cm_f16 && src0->ne[2] >= 16 ? - ctx->device->pipeline_lightning_indexer_cm_f16 : ctx->device->pipeline_lightning_indexer_f16; + vk_pipeline cm = ctx->device->pipeline_lightning_indexer_cm_f16 ? + ctx->device->pipeline_lightning_indexer_cm_f16 : ctx->device->pipeline_lightning_indexer_cm_small_f16; + return cm && src0->ne[2] >= 16 ? cm : ctx->device->pipeline_lightning_indexer_f16; } // only the k type selects a pipeline, the other types are fixed by ggml_lightning_indexer() if (ggml_vk_lightning_indexer_k_type_supported(src1->type)) { diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/lightning_indexer_cm.comp b/ggml/src/ggml-vulkan/vulkan-shaders/lightning_indexer_cm.comp index a0a3639d254..c53eb74f8d4 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/lightning_indexer_cm.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/lightning_indexer_cm.comp @@ -5,9 +5,14 @@ #extension GL_EXT_shader_explicit_arithmetic_types_float16 : require #extension GL_KHR_cooperative_matrix : require #extension GL_KHR_memory_scope_semantics : require +#extension GL_KHR_shader_subgroup_basic : require layout(constant_id = 0) const uint SUBGROUP_SIZE = 64; +#if N_WAVES == 1 layout(local_size_x_id = 0, local_size_y = 1, local_size_z = 1) in; +#else +layout(local_size_x = 512, local_size_y = 1, local_size_z = 1) in; +#endif layout(binding = 0) readonly buffer QBuf { float data_q[]; }; layout(binding = 1) readonly buffer KBuf { float16_t data_k[]; }; @@ -40,13 +45,16 @@ const uint VEC_PER_HEAD = HEAD_SIZE / 4; const uint TILE_STRIDE = VEC_PER_HEAD + 2; const uint SCORE_STRIDE = TILE / 4 + 1; -shared f16vec4 q_sh[TILE * TILE_STRIDE]; -shared f16vec4 k_sh[TILE * TILE_STRIDE]; -shared vec4 score_sh[TILE * SCORE_STRIDE]; +shared f16vec4 q_sh[HEADS_PER_TILE][TILE * TILE_STRIDE]; +shared f16vec4 k_sh[N_WAVES][TILE * TILE_STRIDE]; +shared vec4 score_sh[N_WAVES][TILE * SCORE_STRIDE]; +shared float weight_sh[HEADS_PER_TILE][TILE]; void main() { const uint tid = gl_LocalInvocationIndex; - const uint kv_base = gl_WorkGroupID.x * TILE; + const uint lane = gl_SubgroupInvocationID; + const uint wave = gl_SubgroupID; + const uint kv_base = (gl_WorkGroupID.x * N_WAVES + wave) * TILE; const uint token_base = gl_WorkGroupID.y * TILE; const uint stream = gl_WorkGroupID.z; @@ -55,63 +63,77 @@ void main() { totals[i] = 0.0; } - for (uint idx = tid; idx < TILE * VEC_PER_HEAD; idx += SUBGROUP_SIZE) { - const uint key = idx / VEC_PER_HEAD; - const uint d4 = idx % VEC_PER_HEAD; - const uint kv = kv_base + key; + for (uint idx = tid; idx < N_WAVES * TILE * VEC_PER_HEAD; idx += gl_WorkGroupSize.x) { + const uint load_wave = idx / (TILE * VEC_PER_HEAD); + const uint wave_idx = idx % (TILE * VEC_PER_HEAD); + const uint key = wave_idx / VEC_PER_HEAD; + const uint d4 = wave_idx % VEC_PER_HEAD; + const uint kv = (gl_WorkGroupID.x * N_WAVES + load_wave) * TILE + key; f16vec4 value = f16vec4(0.0); if (kv < p.n_kv) { const uint offset = stream * p.nbk3 + kv * p.nbk2 + d4 * 4; value = f16vec4(data_k[offset], data_k[offset + 1], data_k[offset + 2], data_k[offset + 3]); } - k_sh[key * TILE_STRIDE + d4] = value; + k_sh[load_wave][key * TILE_STRIDE + d4] = value; } barrier(); - for (uint head = 0; head < N_HEAD; ++head) { - for (uint idx = tid; idx < TILE * VEC_PER_HEAD; idx += SUBGROUP_SIZE) { - const uint token_local = idx / VEC_PER_HEAD; - const uint d4 = idx % VEC_PER_HEAD; + for (uint head_base = 0; head_base < N_HEAD; head_base += HEADS_PER_TILE) { + for (uint idx = tid; idx < HEADS_PER_TILE * TILE * VEC_PER_HEAD; idx += gl_WorkGroupSize.x) { + const uint head_local = idx / (TILE * VEC_PER_HEAD); + const uint head_idx = idx % (TILE * VEC_PER_HEAD); + const uint token_local = head_idx / VEC_PER_HEAD; + const uint d4 = head_idx % VEC_PER_HEAD; const uint token = token_base + token_local; f16vec4 value = f16vec4(0.0); if (token < p.n_batch) { - const uint offset = stream * p.nbq3 + token * p.nbq2 + head * p.nbq1 + d4 * 4; + const uint offset = stream * p.nbq3 + token * p.nbq2 + (head_base + head_local) * p.nbq1 + d4 * 4; value = f16vec4(data_q[offset], data_q[offset + 1], data_q[offset + 2], data_q[offset + 3]); } - q_sh[token_local * TILE_STRIDE + d4] = value; + q_sh[head_local][token_local * TILE_STRIDE + d4] = value; + } + if (tid < HEADS_PER_TILE * TILE) { + const uint head_local = tid / TILE; + const uint token_local = tid % TILE; + const uint token = token_base + token_local; + weight_sh[head_local][token_local] = token < p.n_batch ? data_w[stream * p.nbw3 + token * p.nbw1 + head_base + head_local] : 0.0; } barrier(); - coopmat scores = - coopmat(0.0); - coopmat kmat; - coopmat qmat; + [[unroll]] for (uint head_local = 0; head_local < HEADS_PER_TILE; ++head_local) { + coopmat scores = + coopmat(0.0); + coopmat kmat; + coopmat qmat; - [[unroll]] for (uint d = 0; d < HEAD_SIZE; d += TILE) { - coopMatLoad(kmat, k_sh, d / 4, TILE_STRIDE, gl_CooperativeMatrixLayoutRowMajor); - coopMatLoad(qmat, q_sh, d / 4, TILE_STRIDE, gl_CooperativeMatrixLayoutColumnMajor); - scores = coopMatMulAdd(kmat, qmat, scores); - } + [[unroll]] for (uint d = 0; d < HEAD_SIZE; d += TILE) { + coopMatLoad(kmat, k_sh[wave], d / 4, TILE_STRIDE, gl_CooperativeMatrixLayoutRowMajor); + coopMatLoad(qmat, q_sh[head_local], d / 4, TILE_STRIDE, gl_CooperativeMatrixLayoutColumnMajor); + scores = coopMatMulAdd(kmat, qmat, scores); + } - coopMatStore(scores, score_sh, 0, SCORE_STRIDE, gl_CooperativeMatrixLayoutRowMajor); - barrier(); + coopMatStore(scores, score_sh[wave], 0, SCORE_STRIDE, gl_CooperativeMatrixLayoutRowMajor); + controlBarrier(gl_ScopeSubgroup, gl_ScopeSubgroup, gl_StorageSemanticsShared, gl_SemanticsAcquireRelease); - [[unroll]] for (uint i = 0; i < 4; ++i) { - const uint idx = tid + i * SUBGROUP_SIZE; - const uint key = idx / TILE; - const uint token_local = idx % TILE; - const uint token = token_base + token_local; - if (token < p.n_batch && kv_base + key < p.n_kv) { - const float score = score_sh[key * SCORE_STRIDE + token_local / 4][token_local % 4]; - const float weight = data_w[stream * p.nbw3 + token * p.nbw1 + head]; - totals[i] += max(score, 0.0) * weight; + [[unroll]] for (uint i = 0; i < 4; ++i) { + const uint idx = lane + i * SUBGROUP_SIZE; + const uint key = idx / TILE; + const uint token_local = idx % TILE; + const uint token = token_base + token_local; + if (token < p.n_batch && kv_base + key < p.n_kv) { + const float score = score_sh[wave][key * SCORE_STRIDE + token_local / 4][token_local % 4]; + totals[i] += max(score, 0.0) * weight_sh[head_local][token_local]; + } + } + if (head_local + 1 < HEADS_PER_TILE) { + controlBarrier(gl_ScopeSubgroup, gl_ScopeSubgroup, gl_StorageSemanticsShared, gl_SemanticsAcquireRelease); } } barrier(); } [[unroll]] for (uint i = 0; i < 4; ++i) { - const uint idx = tid + i * SUBGROUP_SIZE; + const uint idx = lane + i * SUBGROUP_SIZE; const uint key = idx / TILE; const uint token_local = idx % TILE; const uint kv = kv_base + key; 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 db65e1d49ae..2b8642e45f8 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp @@ -813,7 +813,8 @@ void process_shaders() { string_to_spv("lightning_indexer_f16", "lightning_indexer_scalar64.comp", {}); #if defined(GGML_VULKAN_COOPMAT_GLSLC_SUPPORT) - string_to_spv("lightning_indexer_cm_f16", "lightning_indexer_cm.comp", {}); + string_to_spv("lightning_indexer_cm_f16", "lightning_indexer_cm.comp", {{"N_WAVES", "8"}, {"HEADS_PER_TILE", "4"}}); + string_to_spv("lightning_indexer_cm_small_f16", "lightning_indexer_cm.comp", {{"N_WAVES", "1"}, {"HEADS_PER_TILE", "1"}}); string_to_spv("lightning_indexer_decode_cm_f16", "lightning_indexer_decode_cm.comp", {}); #endif string_to_spv("flash_attn_top_k_f16", "flash_attn_top_k.comp", {}); diff --git a/tests/test-backend-ops.cpp b/tests/test-backend-ops.cpp index 23b30958db3..86e4713303c 100644 --- a/tests/test-backend-ops.cpp +++ b/tests/test-backend-ops.cpp @@ -11411,6 +11411,10 @@ static std::vector> make_test_cases_eval() { test_cases.emplace_back(new test_flash_attn_ext_top_k( 768, 64, 64, 128, true, 2)); test_cases.emplace_back(new test_flash_attn_ext_top_k(4096, 300, 256, 512, false, 3)); + for (int kv : { 127, 128, 129 }) { + test_cases.emplace_back(new test_lightning_indexer(128, 64, kv, 32, 4, 1, GGML_TYPE_F16)); + } + return test_cases; } #ifdef _MSC_VER @@ -11900,6 +11904,11 @@ static std::vector> make_test_cases_perf() { } } } + // DSV4 PP2048 indexer rows after filling source contexts from 8k through 512k. + // The zero-depth kv=512 shape is covered above. + for (int kv : { 2560, 4608, 8704, 16896, 33280, 66048, 131584 }) { + test_cases.emplace_back(new test_lightning_indexer(128, 64, kv, 2048, 1, 1, GGML_TYPE_F16)); + } // sparse top-k FA at V4 decode/prefill shapes — the A/B instrument for the // gather-to-compact work (n_active = n_kv_raw + n_top_k stays fixed as kv grows). From 2a6d1de4376204749f8cfd4a80eb50c57c2fdb83 Mon Sep 17 00:00:00 2001 From: Nathan Wilson Date: Thu, 13 Aug 2026 16:00:20 +0000 Subject: [PATCH 27/68] vulkan: extend DeepSeek V4 gather-to-compact to small batches The gap and its consequence were diagnosed by Jaap Buurman: the sparse prefill path gates on batch >= 64 and gather-to-compact gated on batch == 1, so batch 2..63 fell through to dense attention over the whole compressed KV, at a cost that grows with context. That is where a speculative draft lands (n_max 3-5), which is why token generation dropped off sharply with DSpark enabled. This implements the fix for that diagnosis. Give each query token its own gathered top-k block rather than deduplicating into a union. No dedup pass, no atomics, and the size is bounded by n_kv_raw + n_batch*n_top_k regardless of depth. Cross-token rows are neutralised through the mask, which already encodes each token's selection: token t reads -inf on any block that is not its own, so the softmax cannot double count. The compact mask is token-major [n_batch][kv_c], which is what the GQA mask path already expects (m_stride 0, rows stepped by gqa_iq1 * m_row_len). Measured on gfx1151, test-backend-ops perf, medians of 2 launches, n_kv_raw=2304 n_top_k=512. Gather cost is flat in depth (894 us at batch 4 at every depth); dense is not: kv rows batch dense gathered speedup 11008 2 2194 us 680 us 3.23x 11008 4 2204 us 895 us 2.46x 11008 8 2217 us 2217 us 1.00x (gate declines: kv < 2*kv_c) 35584 2 7065 us 680 us 10.38x 35584 4 7099 us 895 us 7.93x 35584 8 7114 us 1316 us 5.41x 133888 2 7552 us 680 us 11.10x 133888 4 7971 us 894 us 8.92x 133888 8 8348 us 1320 us 6.32x Batch 1 is unchanged in behaviour and stays on the same code path. Not validated end to end. A deduplicated union would shrink the gathered set further wherever adjacent draft tokens select overlapping keys, and would lift the batch ceiling documented in the following commit. Suggested-by: Jaap Buurman (@Mushoz) Co-Authored-By: Claude Opus 5 Assisted-by: Claude (Opus 5) --- ggml/src/ggml-vulkan/ggml-vulkan.cpp | 19 +++++-- .../vulkan-shaders/flash_attn_gather.comp | 52 ++++++++++++++----- tests/test-backend-ops.cpp | 17 ++++++ 3 files changed, 70 insertions(+), 18 deletions(-) diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp index 81021b6204c..840d9e376be 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp +++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp @@ -1985,7 +1985,7 @@ static_assert(sizeof(vk_op_dsv4_hc_post_push_constants) <= 128); struct vk_op_flash_attn_gather_push_constants { uint32_t n_kv, n_kv_raw, n_top_k, kv_c; - uint32_t nbk1, nbk3, nbt3, nbm3, nem3; + uint32_t nbk1, nbk3, nbt1, nbt3, nbm1, nbm3, nem3, n_batch; }; static_assert(sizeof(vk_op_flash_attn_gather_push_constants) <= 128); @@ -11638,6 +11638,7 @@ static bool ggml_vk_flash_attn_top_k(ggml_backend_vk_context * ctx, vk_context & struct vk_fa_compact_state { bool active = false; uint32_t kv_c = 0; + uint32_t n_batch = 1; vk_subbuffer kc_buf, mc_buf; }; @@ -11655,7 +11656,7 @@ static bool ggml_vk_flash_attn_gather_compact(ggml_backend_vk_context * ctx, vk_ static const char * gather_env = getenv("GGML_VK_FA_TOPK_GATHER"); if ((gather_env && gather_env[0] == '0') || !top_k || !ctx->device->pipeline_flash_attn_gather_f16 || - q->ne[1] != 1 || // single-token decode only; batched queries need a union gather + q->ne[1] < 1 || q->ne[1] >= 64 || // 1..63: >=64 goes to the sparse prefill path q->type != GGML_TYPE_F32 || k->type != GGML_TYPE_F16 || v->type != GGML_TYPE_F16 || !mask || mask->type != GGML_TYPE_F16 || top_k->type != GGML_TYPE_I32 || q->ne[0] != 512 || k->ne[0] != 512 || v->ne[0] != 512 || q->ne[2] != 64 || @@ -11679,7 +11680,10 @@ static bool ggml_vk_flash_attn_gather_compact(ggml_backend_vk_context * ctx, vk_ return false; } - const uint32_t kv_c = GGML_PAD((uint32_t)(n_kv_raw + top_k->ne[0]), 256u); + const uint32_t n_batch = (uint32_t) q->ne[1]; + // every token gets its own top-k block; bounded by n_kv_raw + n_batch*n_top_k regardless + // of context depth, which is the whole point at decode + const uint32_t kv_c = GGML_PAD((uint32_t)(n_kv_raw + (int64_t) n_batch * top_k->ne[0]), 256u); // the gather writes then re-reads ~the active bytes; dense reads the source KV once, // so compaction only pays when the source is comfortably larger than the active set if ((uint64_t) k->ne[1] < 2ull * kv_c) { @@ -11688,7 +11692,7 @@ static bool ggml_vk_flash_attn_gather_compact(ggml_backend_vk_context * ctx, vk_ const uint32_t ns = (uint32_t) q->ne[3]; const size_t kc_sz = (size_t) ns * kv_c * 512 * sizeof(ggml_fp16_t); - const size_t mc_sz = (size_t) ns * kv_c * sizeof(ggml_fp16_t); + const size_t mc_sz = (size_t) ns * n_batch * kv_c * sizeof(ggml_fp16_t); if (ctx->prealloc_size_y < kc_sz + mc_sz) { ctx->prealloc_size_y = kc_sz + mc_sz; @@ -11705,9 +11709,12 @@ static bool ggml_vk_flash_attn_gather_compact(ggml_backend_vk_context * ctx, vk_ (uint32_t) k->ne[1], (uint32_t) n_kv_raw, (uint32_t) top_k->ne[0], kv_c, (uint32_t) (k->nb[1] / sizeof(ggml_fp16_t)), (uint32_t) (k->nb[3] / sizeof(ggml_fp16_t)), + (uint32_t) (top_k->nb[1] / sizeof(int32_t)), (uint32_t) (top_k->nb[3] / sizeof(int32_t)), + (uint32_t) (mask->nb[1] / sizeof(ggml_fp16_t)), (uint32_t) (mask->nb[3] / sizeof(ggml_fp16_t)), (uint32_t) mask->ne[3], + n_batch, }; st.kc_buf = ggml_vk_subbuffer(ctx, ctx->prealloc_y, 0); @@ -11721,6 +11728,7 @@ static bool ggml_vk_flash_attn_gather_compact(ggml_backend_vk_context * ctx, vk_ st.active = true; st.kv_c = kv_c; + st.n_batch = n_batch; return true; } @@ -11785,6 +11793,9 @@ static void ggml_vk_flash_attn(ggml_backend_vk_context * ctx, vk_context& subctx if (ggml_vk_flash_attn_gather_compact(ctx, subctx, q, k, v, mask, dst, fa_compact)) { KV = fa_compact.kv_c; nem0 = fa_compact.kv_c; + nem1 = fa_compact.n_batch; + nem2 = 1; + nem3 = (uint32_t) q->ne[3]; nem1 = N; nem2 = 1; nem3 = (uint32_t) q->ne[3]; diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_gather.comp b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_gather.comp index 4e3dc1c4262..3c9ac68e2c6 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_gather.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_gather.comp @@ -4,9 +4,15 @@ #extension GL_EXT_shader_16bit_storage : require #extension GL_EXT_shader_explicit_arithmetic_types_float16 : require -// Gathers the active KV rows of a top-k sparse attention (DeepSeek V4 CSA decode) into a -// compact contiguous scratch: rows [0, n_kv_raw) of the source (the dense prefix), then the -// n_top_k selected rows, then zero padding up to kv_c. The gathered mask row keeps the +// Gathers the active KV rows of a top-k sparse attention (DeepSeek V4 CSA) into a compact +// contiguous scratch: rows [0, n_kv_raw) of the source (the dense prefix), then, for each of +// the n_batch query tokens, that token's n_top_k selected rows, then zero padding up to kv_c. +// +// n_batch > 1 (small-batch decode, e.g. speculative drafts) gives every token its own block +// rather than deduplicating into a union. That costs n_batch*n_top_k gathered rows instead of +// |union|, but needs no dedup pass and no atomics, and the size is bounded by the worst case +// the caller already allocates for. Cross-token rows are neutralised through the mask: token t +// sees -inf on any block that is not its own, so nothing is double counted in the softmax. The gathered mask row keeps the // per-key mask values so causality/validity survive compaction; invalid top-k indices and // padding get -inf mask and zeroed K (softmax-neutral either way, zeroed so no NaN*0). // One workgroup per compact row; V is the K latent (V==K), so a single gather serves both. @@ -26,9 +32,12 @@ layout(push_constant) uniform Parameters { uint kv_c; // padded compact row count == dispatch row range uint nbk1; // K source row stride, elements uint nbk3; // K source stream stride, elements + uint nbt1; // top_k row (per query token) stride, elements uint nbt3; // top_k stream stride, elements + uint nbm1; // mask source row (per query token) stride, elements uint nbm3; // mask source stream stride, elements uint nem3; // mask ne[3], for stream broadcast + uint n_batch; // query tokens sharing this gather; <= LANES } p; const uint HEAD_SIZE = 512; @@ -39,14 +48,23 @@ void main() { const uint stream = gl_WorkGroupID.z; const uint tid = gl_LocalInvocationIndex; - // map compact row -> source row; p.n_kv is the invalid sentinel - uint src = p.n_kv; + // map compact row -> source row; p.n_kv is the invalid sentinel. + // owner is the token whose block this row belongs to, or ALL_TOKENS for the shared prefix. + const uint ALL_TOKENS = 0xffffffffu; + uint src = p.n_kv; + uint owner = ALL_TOKENS; if (row < p.n_kv_raw) { src = row; - } else if (row < p.n_kv_raw + p.n_top_k) { - const int idx = data_top[stream * p.nbt3 + (row - p.n_kv_raw)]; - if (idx >= 0 && uint(idx) < p.n_kv - p.n_kv_raw) { - src = p.n_kv_raw + uint(idx); + } else { + const uint off = row - p.n_kv_raw; + const uint tok = off / p.n_top_k; + const uint slot = off - tok * p.n_top_k; + if (tok < p.n_batch) { + owner = tok; + const int idx = data_top[stream * p.nbt3 + tok * p.nbt1 + slot]; + if (idx >= 0 && uint(idx) < p.n_kv - p.n_kv_raw) { + src = p.n_kv_raw + uint(idx); + } } } @@ -56,15 +74,21 @@ void main() { [[unroll]] for (uint i = 0; i < HEAD_SIZE / LANES; ++i) { data_kc[dst_base + tid + i * LANES] = data_k[src_base + tid + i * LANES]; } - if (tid == 0) { - data_mc[stream * p.kv_c + row] = data_m[(stream % p.nem3) * p.nbm3 + src]; - } } else { [[unroll]] for (uint i = 0; i < HEAD_SIZE / LANES; ++i) { data_kc[dst_base + tid + i * LANES] = float16_t(0.0); } - if (tid == 0) { - data_mc[stream * p.kv_c + row] = float16_t(uintBitsToFloat(0xff800000)); + } + + // Compact mask is token-major [n_batch][kv_c], which is what the GQA mask path expects + // (m_stride 0, rows stepped by gqa_iq1 * m_row_len). One lane per token; n_batch <= LANES. + const float NEG_INF = uintBitsToFloat(0xff800000); + if (tid < p.n_batch) { + const uint mc_idx = (stream * p.n_batch + tid) * p.kv_c + row; + float mv = NEG_INF; + if (src < p.n_kv && (owner == ALL_TOKENS || owner == tid)) { + mv = float(data_m[(stream % p.nem3) * p.nbm3 + tid * p.nbm1 + src]); } + data_mc[mc_idx] = float16_t(mv); } } diff --git a/tests/test-backend-ops.cpp b/tests/test-backend-ops.cpp index 86e4713303c..a3fd3562dbc 100644 --- a/tests/test-backend-ops.cpp +++ b/tests/test-backend-ops.cpp @@ -11407,6 +11407,16 @@ static std::vector> make_test_cases_eval() { test_cases.emplace_back(new test_flash_attn_ext_top_k(4096, 257, 256, 512, false)); // ns > 1: the split-K partial-output path indexes O and L/M by stream, so these cover // the stream stride in both regions (single tile and multi-tile). + // small-batch decode (speculative drafts): each token gets its own gathered top-k block, + // so cross-token rows must be masked out or the softmax double counts. kv must be large + // enough that compaction is worth it (the gather gates on kv >= 2*kv_c). + test_cases.emplace_back(new test_flash_attn_ext_top_k(8192, 2, 1024, 512, false)); + test_cases.emplace_back(new test_flash_attn_ext_top_k(8192, 3, 1024, 512, false)); + test_cases.emplace_back(new test_flash_attn_ext_top_k(8192, 4, 1024, 512, false)); + test_cases.emplace_back(new test_flash_attn_ext_top_k(8192, 8, 1024, 512, true)); + test_cases.emplace_back(new test_flash_attn_ext_top_k(32768, 16, 2304, 512, false)); + test_cases.emplace_back(new test_flash_attn_ext_top_k(65536, 63, 2304, 512, false)); + test_cases.emplace_back(new test_flash_attn_ext_top_k(8192, 4, 1024, 512, false, 2)); test_cases.emplace_back(new test_flash_attn_ext_top_k( 768, 64, 64, 128, false, 2)); test_cases.emplace_back(new test_flash_attn_ext_top_k( 768, 64, 64, 128, true, 2)); test_cases.emplace_back(new test_flash_attn_ext_top_k(4096, 300, 256, 512, false, 3)); @@ -11924,6 +11934,13 @@ static std::vector> make_test_cases_perf() { for (int kv : { 11008, 19200, 35584, 68352, 133888 }) { test_cases.emplace_back(new test_flash_attn_ext_top_k(kv, 2048, 2304, 512, false)); } + // small-batch decode at depth: the speculative-draft regime (batch 2-8), where the old + // path fell through to dense attention over the whole compressed KV. + for (int kv : { 11008, 35584, 133888 }) { + for (int nb : { 1, 2, 4, 8 }) { + test_cases.emplace_back(new test_flash_attn_ext_top_k(kv, nb, 2304, 512, false)); + } + } return test_cases; } From 6a31bf0adbc787510f9c4c0a465019dc135e99f0 Mon Sep 17 00:00:00 2001 From: Nathan Wilson Date: Thu, 13 Aug 2026 22:20:56 +0000 Subject: [PATCH 28/68] vulkan: record the resource limits behind the two DSV4 prefill kernels Both additions sit close to a limit that is invisible at the call site. The parallel Lightning Indexer uses 62720 of 65536 bytes of shared memory at N_WAVES=8 / HEADS_PER_TILE=4, so exactly one workgroup fits per CU. That is the intended trade, but raising either constant overruns the budget and the pipeline then fails to create and silently falls back to the small variant. Write the arithmetic down next to the arrays. The small-batch gather's compact set is independent of context depth but grows with batch, so its attention work is quadratic in batch against dense's linear. The kv >= 2*kv_c gate already caps this at about batch 6 at 32k depth and ~30 at 128k, and beyond the cap dense runs instead, so it is never slower. But the measured speedups do not show that ceiling and it is the main argument for building the deduplicated union later. Co-Authored-By: Claude Opus 5 Assisted-by: Claude (Opus 5) --- ggml/src/ggml-vulkan/ggml-vulkan.cpp | 10 ++++++++-- .../vulkan-shaders/lightning_indexer_cm.comp | 10 ++++++++++ 2 files changed, 18 insertions(+), 2 deletions(-) diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp index 840d9e376be..b84c1512541 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp +++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp @@ -11681,8 +11681,14 @@ static bool ggml_vk_flash_attn_gather_compact(ggml_backend_vk_context * ctx, vk_ } const uint32_t n_batch = (uint32_t) q->ne[1]; - // every token gets its own top-k block; bounded by n_kv_raw + n_batch*n_top_k regardless - // of context depth, which is the whole point at decode + // Every token gets its own top-k block, so the compact set is n_kv_raw + n_batch*n_top_k: + // independent of context depth, which is the point, but GROWING WITH BATCH. Attention work + // is then n_batch * (n_kv_raw + n_batch*n_top_k), i.e. quadratic in batch, against dense's + // n_batch * n_kv. Break-even is n_batch = (n_kv - n_kv_raw) / n_top_k, and the + // kv >= 2*kv_c gate below caps the useful batch at (n_kv/2 - n_kv_raw) / n_top_k -- + // about 6 at 32k depth, ~30 at 128k, batch-capped at 512k. Beyond that the gate declines + // and dense runs, so this can never be slower; it just stops helping. A deduplicated union + // would lift that ceiling wherever draft tokens select overlapping keys. const uint32_t kv_c = GGML_PAD((uint32_t)(n_kv_raw + (int64_t) n_batch * top_k->ne[0]), 256u); // the gather writes then re-reads ~the active bytes; dense reads the source KV once, // so compaction only pays when the source is comfortably larger than the active set diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/lightning_indexer_cm.comp b/ggml/src/ggml-vulkan/vulkan-shaders/lightning_indexer_cm.comp index c53eb74f8d4..d4379e8899d 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/lightning_indexer_cm.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/lightning_indexer_cm.comp @@ -45,6 +45,16 @@ const uint VEC_PER_HEAD = HEAD_SIZE / 4; const uint TILE_STRIDE = VEC_PER_HEAD + 2; const uint SCORE_STRIDE = TILE / 4 + 1; +// Shared-memory budget, N_WAVES=8 / HEADS_PER_TILE=4 (TILE 16, HEAD_SIZE 128): +// k_sh 8 * 16 * 34 * 8 B = 34816 +// q_sh 4 * 16 * 34 * 8 B = 17408 +// score_sh 8 * 16 * 5 * 16 B = 10240 +// weight_sh 4 * 16 * 4 B = 256 +// total = 62720 of 65536 (95.7%) +// Only one workgroup fits per CU at that size, which is the intended trade. Raising either +// constant overruns: HEADS_PER_TILE=8 needs 80128 B and the pipeline then fails to create, +// silently falling back to the small variant. The host gates on +// maxComputeSharedMemorySize >= 64 KiB; keep that gate and this arithmetic in sync. shared f16vec4 q_sh[HEADS_PER_TILE][TILE * TILE_STRIDE]; shared f16vec4 k_sh[N_WAVES][TILE * TILE_STRIDE]; shared vec4 score_sh[N_WAVES][TILE * SCORE_STRIDE]; From 964a218fc92ede4bb18a15dae5fea89f1fbbfd9c Mon Sep 17 00:00:00 2001 From: Nathan Wilson Date: Fri, 14 Aug 2026 00:01:20 +0000 Subject: [PATCH 29/68] vulkan: deduplicated union for DeepSeek V4 small-batch decode Replaces the per-token top-k blocks with one row per DISTINCT selected key, so the compact set stops growing linearly with batch. Measured on the real model, adjacent tokens share 60% of their selections over 4 tokens and 76% over 8, so the union is materially smaller than n_batch*n_top_k. The size is only known on the GPU, so flash-attention now reads its KV bound from a buffer instead of the push constant, behind a new DYNAMIC_KV pipeline flag. Every other pipeline folds the flag away at compile time and is unchanged (13308 FLASH_ATTN_EXT cases pass identically with the union off). No indirect dispatch is needed: FA workgroup counts come from neq1/neq2/neq3 and never from KV, so only the loop bound moves. Padding the count to 256 keeps KV % Bc == 0, which lets the aligned pipeline variant still apply. Dedup marks a bitmap from the top-k lists and compacts by scanning bitmap WORDS. An earlier version scanned the mask row by row: simpler, but it made dedup cost scale with depth and measured 0.71x at kv=133888, ie a loss. The bitmap form is O(n_batch*n_top_k) to mark and R/32 to compact, and is depth-independent. test-backend-ops perf, medians of 2, n_kv_raw=2304 n_top_k=512, fixture overlap 60% (the default generator produces near-zero overlap and would make the union look worthless by construction): kv batch per-token union speedup 35584 2 680 us 632 us 1.08x 35584 4 892 us 720 us 1.24x 35584 8 1315 us 900 us 1.46x 133888 2 679 us 638 us 1.06x 133888 4 892 us 726 us 1.23x 133888 8 1314 us 909 us 1.45x Union cost is flat in depth (720 vs 726 us at batch 4 across a 3.8x depth range), which the per-token form was not. Opt-in via GGML_VK_FA_TOPK_UNION=1, single stream only, and it falls back to the per-token blocks when the compressed region exceeds the shared bitmap. The kv >= 2*kv_c gate still uses the worst case, so the batch ceiling is NOT yet lifted; doing that needs a GPU-side fallback for an oversized union. Suggested-by: Jaap Buurman (@Mushoz) Co-Authored-By: Claude Opus 5 Assisted-by: Claude (Opus 5) --- ggml/src/ggml-vulkan/ggml-vulkan.cpp | 209 +++++++++++++++++- .../vulkan-shaders/flash_attn_base.glsl | 8 +- .../flash_attn_gather_union.comp | 82 +++++++ .../vulkan-shaders/flash_attn_union.comp | 122 ++++++++++ .../vulkan-shaders/vulkan-shaders-gen.cpp | 2 + tests/test-backend-ops.cpp | 20 +- 6 files changed, 428 insertions(+), 15 deletions(-) create mode 100644 ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_gather_union.comp create mode 100644 ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_union.comp diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp index b84c1512541..1e22e8be459 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp +++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp @@ -64,6 +64,7 @@ typedef struct VkPhysicalDeviceCooperativeMatrixDecodeVectorFeaturesNV { #include #include #include +#include #include #include #include @@ -1111,6 +1112,8 @@ struct vk_device_struct { vk_pipeline pipeline_flash_attn_top_k_f16; vk_pipeline pipeline_flash_attn_top_k_cm_f16; vk_pipeline pipeline_flash_attn_gather_f16; + vk_pipeline pipeline_flash_attn_union_f16; + vk_pipeline pipeline_flash_attn_gather_union_f16; vk_pipeline pipeline_dsv4_hc_pre_f32; vk_pipeline pipeline_dsv4_hc_comb_f32; vk_pipeline pipeline_dsv4_hc_post_f32; @@ -1983,6 +1986,12 @@ struct vk_op_dsv4_hc_post_push_constants { }; static_assert(sizeof(vk_op_dsv4_hc_post_push_constants) <= 128); +struct vk_op_flash_attn_union_push_constants { + uint32_t n_kv, n_kv_raw, n_batch, n_top_k, max_union, nbt1, max_words, pad_to; +}; +struct vk_op_flash_attn_gather_union_push_constants { + uint32_t n_kv, n_kv_raw, kv_c_max, nbk1, nbm1, n_batch; +}; struct vk_op_flash_attn_gather_push_constants { uint32_t n_kv, n_kv_raw, n_top_k, kv_c; uint32_t nbk1, nbk3, nbt1, nbt3, nbm1, nbm3, nem3, n_batch; @@ -4050,14 +4059,16 @@ 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, ggml_type k_type, ggml_type v_type, + bool use_dynamic_kv = false) { 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_dynamic_kv ? 16 : 0); const uint32_t subgroup_size = params.disable_subgroups ? 0 : params.subgroup_size; @@ -4732,7 +4743,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); @@ -4768,7 +4779,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); @@ -4805,7 +4816,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); } @@ -6315,6 +6326,16 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { "flash_attn_gather_f16", flash_attn_gather_f16_len, flash_attn_gather_f16_data, "main", 5, sizeof(vk_op_flash_attn_gather_push_constants), {1, 1, 1}, {}, 1, true, true, device->subgroup_size); + if (device->subgroup_arithmetic) { + ggml_vk_create_pipeline(device, device->pipeline_flash_attn_union_f16, + "flash_attn_union_f16", flash_attn_union_f16_len, flash_attn_union_f16_data, "main", 3, + sizeof(vk_op_flash_attn_union_push_constants), {1, 1, 1}, {}, 1, true, true, + device->subgroup_size); + ggml_vk_create_pipeline(device, device->pipeline_flash_attn_gather_union_f16, + "flash_attn_gather_union_f16", flash_attn_gather_union_f16_len, flash_attn_gather_union_f16_data, "main", 6, + sizeof(vk_op_flash_attn_gather_union_push_constants), {1, 1, 1}, {}, 1, true, true, + device->subgroup_size); + } } // DSv4 fused hyper-connection ops: plain f32 compute, no subgroup/coopmat requirements @@ -11483,6 +11504,56 @@ static bool ggml_vk_flash_attn_top_k(ggml_backend_vk_context * ctx, vk_context & return false; } + // ---- diagnostic: adjacent-token top-k overlap (GGML_VK_TOPK_OVERLAP=1) --------------- + // Decides whether a deduplicated union is worth building for the small-batch path. A + // speculative draft attends at adjacent POSITIONS, so overlap between adjacent prefill + // tokens is the same quantity and a plain deep prefill samples it thousands of times + // without needing a draft model. Reports |union| / (W * n_top_k) for window sizes W. + // Sampled, host-side, off by default; the readback would be far too costly otherwise. + { + static const char * ov_env = getenv("GGML_VK_TOPK_OVERLAP"); + if (ov_env && ov_env[0] == '1') { + static std::mutex ov_mu; + static uint64_t ov_seen = 0; + static std::map> ov_acc; // W -> {sum ratio, n} + std::lock_guard lock(ov_mu); + if ((ov_seen++ % 32) == 0) { // sample: this is a multi-MB readback + const uint32_t nb = (uint32_t) q->ne[1]; + const uint32_t tk = (uint32_t) top_k->ne[0]; + const int32_t rng = (int32_t) (k->ne[1] - n_kv_raw); + std::vector idx((size_t) nb * tk); + vk_subbuffer sb = ggml_vk_tensor_subbuffer(ctx, top_k); + ggml_vk_buffer_read(sb.buffer, sb.offset, idx.data(), idx.size() * sizeof(int32_t)); + for (uint32_t W : {2u, 4u, 8u}) { + if (nb < W) continue; + double sum = 0.0; uint64_t n = 0; + for (uint32_t t0 = 0; t0 + W <= nb; t0 += W) { // disjoint windows + std::unordered_set u; + uint64_t valid = 0; + for (uint32_t t = t0; t < t0 + W; ++t) { + for (uint32_t j = 0; j < tk; ++j) { + const int32_t v = idx[(size_t) t * tk + j]; + if (v >= 0 && v < rng) { u.insert(v); ++valid; } + } + } + if (valid) { sum += (double) u.size() / (double) valid; ++n; } + } + if (n) { auto & a = ov_acc[W]; a.first += sum; a.second += n; } + } + fprintf(stderr, "[topk-overlap] sample %llu (n_batch=%u, kv=%lld)\n", + (unsigned long long) ov_seen, nb, (long long) k->ne[1]); + for (const auto & e : ov_acc) { + const double r = e.second.first / (double) e.second.second; + const double now = (double) n_kv_raw + (double) e.first * tk; + const double dedup = (double) n_kv_raw + r * (double) e.first * tk; + fprintf(stderr, "[topk-overlap] W=%u union/selected=%.3f (overlap %.1f%%) " + "kv_c %.0f -> %.0f projected small-batch gain %.1f%%\n", + e.first, r, 100.0 * (1.0 - r), now, dedup, 100.0 * (1.0 - dedup / now)); + } + } + } + } + vk_op_flash_attn_top_k_push_constants pc = { (uint32_t) q->ne[1], (uint32_t) k->ne[1], (uint32_t) n_kv_raw, (uint32_t) top_k->ne[0], (uint32_t) q->ne[2], @@ -11602,7 +11673,7 @@ static bool ggml_vk_flash_attn_top_k(ggml_backend_vk_context * ctx, vk_context & }; ggml_vk_dispatch_pipeline(ctx, subctx, raw_pipeline, - {tile_q, k_buf, k_buf, tile_mask, tile_q, split_buf, tile_q}, + {tile_q, k_buf, k_buf, tile_mask, tile_q, split_buf, tile_q, tile_q /* dyn-KV: unused */}, raw_pc, {tile_n, NH, NS}); ggml_vk_perf_mark_subop(ctx, subctx, "FA_TOP_K_RAW (sub-op)"); @@ -11637,8 +11708,10 @@ static bool ggml_vk_flash_attn_top_k(ggml_backend_vk_context * ctx, vk_context & struct vk_fa_compact_state { bool active = false; - uint32_t kv_c = 0; + bool dynamic_kv = false; // KV row count lives in kv_buf, not the push constant + uint32_t kv_c = 0; // upper bound; the real count is runtime when dynamic_kv uint32_t n_batch = 1; + vk_subbuffer kv_buf; vk_subbuffer kc_buf, mc_buf; }; @@ -11696,6 +11769,74 @@ static bool ggml_vk_flash_attn_gather_compact(ggml_backend_vk_context * ctx, vk_ return false; } + // ---- deduplicated union (GGML_VK_FA_TOPK_UNION=1) ---------------------------------- + // Same compact layout, but one row per DISTINCT selected key instead of one block per + // token. Measured adjacent-token overlap on the real model is 60% at 4 tokens and 76% at + // 8, so the union is materially smaller. Its size is only known on the GPU, so the FA + // reads its KV bound from a buffer (the DYNAMIC_KV pipeline flag) rather than a push + // constant; padding the count to 256 keeps KV % Bc == 0 so the aligned variant still + // applies. Single stream only: the FA takes one KV for all streams. + static const char * union_env = getenv("GGML_VK_FA_TOPK_UNION"); + if (union_env && union_env[0] == '1' && q->ne[3] == 1 && n_batch > 1 && + ctx->device->pipeline_flash_attn_union_f16 && ctx->device->pipeline_flash_attn_gather_union_f16) { + const uint32_t max_union = (uint32_t) ((int64_t) n_batch * top_k->ne[0]); + // shared bitmap capacity in flash_attn_union.comp + const uint32_t max_words = 12288; + const uint32_t need_words = (uint32_t) (((k->ne[1] - n_kv_raw) + 31) / 32); + if (need_words > max_words) { + goto union_unavailable; // fall through to the per-token block form + } + const size_t ukc_sz = (size_t) kv_c * 512 * sizeof(ggml_fp16_t); + const size_t umc_sz = (size_t) n_batch * kv_c * sizeof(ggml_fp16_t); + const size_t ul_sz = (size_t) max_union * sizeof(uint32_t); + const size_t uc_sz = 2 * sizeof(uint32_t); + const size_t need = ukc_sz + umc_sz + ul_sz + uc_sz; + if (ctx->prealloc_size_y < need) { + ctx->prealloc_size_y = need; + ggml_vk_preallocate_buffers(ctx, subctx); + } + if (ctx->prealloc_y_need_sync) { + ggml_vk_sync_buffers(ctx, subctx); + } + const vk_subbuffer kc_buf = ggml_vk_subbuffer(ctx, ctx->prealloc_y, 0); + const vk_subbuffer mc_buf = ggml_vk_subbuffer(ctx, ctx->prealloc_y, ukc_sz); + const vk_subbuffer ul_buf = ggml_vk_subbuffer(ctx, ctx->prealloc_y, ukc_sz + umc_sz); + const vk_subbuffer uc_buf = ggml_vk_subbuffer(ctx, ctx->prealloc_y, ukc_sz + umc_sz + ul_sz); + + ggml_pipeline_request_descriptor_sets(ctx, ctx->device->pipeline_flash_attn_union_f16, 1); + ggml_pipeline_request_descriptor_sets(ctx, ctx->device->pipeline_flash_attn_gather_union_f16, 1); + + const vk_op_flash_attn_union_push_constants upc = { + (uint32_t) k->ne[1], (uint32_t) n_kv_raw, n_batch, (uint32_t) top_k->ne[0], max_union, + (uint32_t) (top_k->nb[1] / sizeof(int32_t)), max_words, 256u, + }; + ggml_vk_dispatch_pipeline(ctx, subctx, ctx->device->pipeline_flash_attn_union_f16, + { ggml_vk_tensor_subbuffer(ctx, top_k), ul_buf, uc_buf }, upc, { 1, 1, 1 }); + ggml_vk_sync_buffers(ctx, subctx); + + const vk_op_flash_attn_gather_union_push_constants gpc = { + (uint32_t) k->ne[1], (uint32_t) n_kv_raw, kv_c, + (uint32_t) (k->nb[1] / sizeof(ggml_fp16_t)), + (uint32_t) (mask->nb[1] / sizeof(ggml_fp16_t)), + n_batch, + }; + ggml_vk_dispatch_pipeline(ctx, subctx, ctx->device->pipeline_flash_attn_gather_union_f16, + { ggml_vk_tensor_subbuffer(ctx, k), ul_buf, ggml_vk_tensor_subbuffer(ctx, mask), + kc_buf, mc_buf, uc_buf }, gpc, { kv_c, 1, 1 }); + ggml_vk_sync_buffers(ctx, subctx); + ctx->prealloc_y_need_sync = true; + + st.active = true; + st.dynamic_kv = true; + st.kv_c = kv_c; + st.n_batch = n_batch; + st.kc_buf = kc_buf; + st.mc_buf = mc_buf; + st.kv_buf = uc_buf; + return true; + } +union_unavailable:; + const uint32_t ns = (uint32_t) q->ne[3]; const size_t kc_sz = (size_t) ns * kv_c * 512 * sizeof(ggml_fp16_t); const size_t mc_sz = (size_t) ns * n_batch * kv_c * sizeof(ggml_fp16_t); @@ -11732,6 +11873,53 @@ static bool ggml_vk_flash_attn_gather_compact(ggml_backend_vk_context * ctx, vk_ ggml_vk_sync_buffers(ctx, subctx); ctx->prealloc_y_need_sync = true; + // ---- diagnostic: measure real top-k overlap between the query tokens ---------------- + // GGML_VK_TOPK_OVERLAP=1. Off by default and never touched on the hot path. Answers the + // one question that decides whether a deduplicated union is worth building: this path + // gathers n_batch*n_top_k rows, a union would gather |union| rows, and the op's cost is + // linear in that count (measured: 204.4 / 205.7 / 205.5 us per row at batch 2 / 4 / 8). + // Needs a real model - synthetic top-k selections say nothing about real overlap, which + // is exactly how the mul_mat_id large-tile probe misled once before. + static const char * overlap_env = getenv("GGML_VK_TOPK_OVERLAP"); + if (overlap_env && overlap_env[0] == '1' && n_batch > 1) { + static std::mutex ov_mutex; + static uint64_t ov_calls = 0, ov_selected = 0, ov_union = 0; + static std::map> ov_by_batch; // n_batch -> {selected, union} + const size_t n_idx = (size_t) n_batch * top_k->ne[0]; + std::vector idx(n_idx); + vk_subbuffer top_sb = ggml_vk_tensor_subbuffer(ctx, top_k); + ggml_vk_buffer_read(top_sb.buffer, top_sb.offset, idx.data(), n_idx * sizeof(int32_t)); + + const int32_t range = (int32_t) (k->ne[1] - n_kv_raw); + std::unordered_set uni; + uint64_t valid = 0; + for (size_t i = 0; i < n_idx; ++i) { + const int32_t v = idx[i]; + if (v >= 0 && v < range) { uni.insert(v); ++valid; } + } + std::lock_guard lock(ov_mutex); + ov_calls++; ov_selected += valid; ov_union += uni.size(); + auto & e = ov_by_batch[n_batch]; + e.first += valid; e.second += uni.size(); + if ((ov_calls % 256) == 0) { + fprintf(stderr, "[topk-overlap] calls=%llu selected=%llu union=%llu " + "union/selected=%.3f => a dedup union would gather %.1f%% fewer compressed rows\n", + (unsigned long long) ov_calls, (unsigned long long) ov_selected, + (unsigned long long) ov_union, + ov_selected ? (double) ov_union / (double) ov_selected : 0.0, + ov_selected ? 100.0 * (1.0 - (double) ov_union / (double) ov_selected) : 0.0); + for (const auto & kv : ov_by_batch) { + const double ratio = kv.second.first ? (double) kv.second.second / (double) kv.second.first : 0.0; + // projected op speedup uses the measured linear cost model on kv_c + const double now = (double) n_kv_raw + (double) kv.first * (double) top_k->ne[0]; + const double dedup = (double) n_kv_raw + ratio * (double) kv.first * (double) top_k->ne[0]; + fprintf(stderr, "[topk-overlap] n_batch=%u union/selected=%.3f " + "kv_c %.0f -> %.0f projected op gain %.1f%%\n", + kv.first, ratio, now, dedup, 100.0 * (1.0 - dedup / now)); + } + } + } + st.active = true; st.kv_c = kv_c; st.n_batch = n_batch; @@ -11943,7 +12131,8 @@ static void ggml_vk_flash_attn(ggml_backend_vk_context * ctx, vk_context& subctx bool use_mask_opt = mask && 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, k_type_eff, v_type_eff, + fa_compact.dynamic_kv); vk_pipeline pipeline = nullptr; @@ -12145,7 +12334,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, fa_compact.dynamic_kv ? fa_compact.kv_buf : q_buf}, pc, { dispatch_x, workgroups_y, workgroups_z }); ggml_vk_sync_buffers(ctx, subctx); @@ -12160,7 +12349,7 @@ 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, fa_compact.dynamic_kv ? fa_compact.kv_buf : q_buf}, pc, { workgroups_x, workgroups_y, workgroups_z }); } 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 e308c7214dd..b562c5d7874 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,11 @@ 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; +// KV comes from a buffer instead of the push constant. Used by paths that compact K/V on the +// 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; // Round up head sizes to a multiple of 16, for coopmat1/coopmat2 paths const uint32_t HSK_pad = (HSK + 15) & ~15; @@ -81,6 +86,7 @@ layout (binding = 5) writeonly buffer O {D_TYPE data_o[];}; 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[];}; #define MASK_OPT_ALL_NEG_INF 1 #define MASK_OPT_ALL_ZERO 2 @@ -150,7 +156,7 @@ bool partial_output; void init_indices() { N = p.N; - KV = p.KV; + KV = DYNAMIC_KV ? data_kv_dyn[0] : p.KV; gqa_ratio = p.gqa_ratio & 0xffff; split_k_num = p.k_num & 0xffff; output_k_num = p.k_num >> 16; diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_gather_union.comp b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_gather_union.comp new file mode 100644 index 00000000000..775b1bc96c6 --- /dev/null +++ b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_gather_union.comp @@ -0,0 +1,82 @@ +#version 450 + +#extension GL_EXT_control_flow_attributes : require +#extension GL_EXT_shader_16bit_storage : require +#extension GL_EXT_shader_explicit_arithmetic_types_float16 : require + +// Gathers K/V and the mask for the DeepSeek V4 small-batch decode path, using the deduplicated +// union produced by flash_attn_union.comp: rows [0, n_kv_raw) of the source, then one row per +// distinct selected compressed row, then padding. +// +// Simpler than the block-per-token form it replaces, because a union row appears exactly once: +// there is no owner to track and no cross-token masking to apply. Each token just reads its own +// mask value for the gathered source row, which is already -inf where that token did not select +// it, so the softmax still cannot double count. +// +// Row count is a runtime value in data_c[0]; rows past the union are zeroed K and -inf mask, +// which is softmax-neutral and keeps the padded tail harmless. + +layout(local_size_x = 64, local_size_y = 1, local_size_z = 1) in; + +layout(binding = 0) readonly buffer KBuf { float16_t data_k[]; }; +layout(binding = 1) readonly buffer UBuf { uint data_u[]; }; +layout(binding = 2) readonly buffer MBuf { float16_t data_m[]; }; +layout(binding = 3) writeonly buffer KcBuf { float16_t data_kc[]; }; +layout(binding = 4) writeonly buffer McBuf { float16_t data_mc[]; }; +layout(binding = 5) readonly buffer CBuf { uint data_c[]; }; + +layout(push_constant) uniform Parameters { + uint n_kv; + uint n_kv_raw; + uint kv_c_max; + uint nbk1; + uint nbm1; + uint n_batch; +} p; + +const uint HEAD_SIZE = 512; +const uint LANES = 64; + +void main() { + const uint row = gl_WorkGroupID.x; + const uint tid = gl_LocalInvocationIndex; + + const uint kv_c = data_c[0]; // padded compact rows, the FA's runtime KV + const uint n_uni = data_c[1]; // unpadded union size + + if (row >= kv_c) { + return; + } + + uint src = p.n_kv; // sentinel: invalid + if (row < p.n_kv_raw) { + src = row; + } else if (row - p.n_kv_raw < n_uni) { + src = p.n_kv_raw + data_u[row - p.n_kv_raw]; + } + + const uint dst_base = row * HEAD_SIZE; + if (src < p.n_kv) { + const uint src_base = src * p.nbk1; + [[unroll]] for (uint i = 0; i < HEAD_SIZE / LANES; ++i) { + data_kc[dst_base + tid + i * LANES] = data_k[src_base + tid + i * LANES]; + } + } else { + [[unroll]] for (uint i = 0; i < HEAD_SIZE / LANES; ++i) { + data_kc[dst_base + tid + i * LANES] = float16_t(0.0); + } + } + + // Compact mask is token-major [n_batch][kv_c]. The stride must be the RUNTIME kv_c, not + // kv_c_max: the FA derives m_row_len from KV, which is now that same runtime value, and a + // mismatch would step the mask by the wrong amount for every token past the first. The + // buffer is allocated for kv_c_max, so a smaller stride simply leaves a tail unused. + const float NEG_INF = uintBitsToFloat(0xff800000); + if (tid < p.n_batch) { + float mv = NEG_INF; + if (src < p.n_kv) { + mv = float(data_m[tid * p.nbm1 + src]); + } + data_mc[tid * kv_c + row] = float16_t(mv); + } +} diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_union.comp b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_union.comp new file mode 100644 index 00000000000..cd779e0a0f9 --- /dev/null +++ b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_union.comp @@ -0,0 +1,122 @@ +#version 450 + +#extension GL_EXT_control_flow_attributes : require +#extension GL_KHR_shader_subgroup_arithmetic : require +#extension GL_KHR_shader_subgroup_basic : require + +// Builds the deduplicated union of the compressed rows selected by any of the n_batch query +// tokens, for the DeepSeek V4 small-batch decode path. +// +// Marks a bitmap from the top-k index lists, then compacts by scanning bitmap WORDS. The +// marking is O(n_batch * n_top_k), independent of context depth, and the compaction touches +// R/32 words instead of R rows. An earlier version scanned the mask row by row instead: that +// is simpler, but it made dedup cost scale with depth and measured 0.71x (ie a loss) at +// kv=133888, which defeats the purpose. Do not go back to it. +// +// Ascending source order falls out of the bitmap scan, so the result is deterministic. +// +// Emits the index list plus the PADDED compact row count. Padding to a multiple of pad_to +// (a multiple of every FA block width) is what lets flash-attention keep its "aligned" +// pipeline variant even though the row count is now a runtime value. +// +// One workgroup: the bitmap and the running offset both live in shared memory. + +layout(local_size_x = 256, local_size_y = 1, local_size_z = 1) in; + +layout(binding = 0) readonly buffer TBuf { int data_top[]; }; +layout(binding = 1) writeonly buffer UBuf { uint data_u[]; }; +layout(binding = 2) writeonly buffer CBuf { uint data_c[]; }; + +layout(push_constant) uniform Parameters { + uint n_kv; + uint n_kv_raw; + uint n_batch; + uint n_top_k; + uint max_union; + uint nbt1; + uint max_words; // capacity of the shared bitmap, host-checked + uint pad_to; +} p; + +// 12288 words = 393216 compressed rows, ~48 KiB of shared memory. +const uint MAX_WORDS = 12288; + +shared uint bitmap[MAX_WORDS]; +shared uint wave_totals[16]; +shared uint base_sh; + +void main() { + const uint tid = gl_LocalInvocationIndex; + const uint lane = gl_SubgroupInvocationID; + const uint wave = gl_SubgroupID; + const uint nwave = gl_NumSubgroups; + const uint R = p.n_kv - p.n_kv_raw; + const uint words = min((R + 31) / 32, p.max_words); + + // phase 1: clear + for (uint w = tid; w < words; w += gl_WorkGroupSize.x) { + bitmap[w] = 0; + } + if (tid == 0) { + base_sh = 0; + } + barrier(); + + // phase 2: mark. Depth-independent: one pass over the top-k lists. + const uint n_cand = p.n_batch * p.n_top_k; + for (uint c = tid; c < n_cand; c += gl_WorkGroupSize.x) { + const uint t = c / p.n_top_k; + const uint j = c - t * p.n_top_k; + const int idx = data_top[t * p.nbt1 + j]; + if (idx >= 0 && uint(idx) < R) { + atomicOr(bitmap[uint(idx) >> 5], 1u << (uint(idx) & 31u)); + } + } + barrier(); + + // phase 3: compact. R/32 iterations, ascending, exact running offset in shared memory. + for (uint chunk = 0; chunk < words; chunk += gl_WorkGroupSize.x) { + const uint w = chunk + tid; + const uint bits = w < words ? bitmap[w] : 0u; + const uint cnt = bitCount(bits); + + const uint wave_off = subgroupExclusiveAdd(cnt); + const uint wave_tot = subgroupAdd(cnt); + if (lane == 0) { + wave_totals[wave] = wave_tot; + } + barrier(); + + uint prefix = 0; + for (uint i = 0; i < wave; ++i) { + prefix += wave_totals[i]; + } + uint total = 0; + for (uint i = 0; i < nwave; ++i) { + total += wave_totals[i]; + } + + uint slot = base_sh + prefix + wave_off; + uint rem = bits; + while (rem != 0) { + const uint b = findLSB(rem); + rem &= rem - 1; + if (slot < p.max_union) { + data_u[slot] = w * 32 + b; + } + ++slot; + } + barrier(); + if (tid == 0) { + base_sh += total; + } + barrier(); + } + + if (tid == 0) { + const uint u = min(base_sh, p.max_union); + const uint rows = p.n_kv_raw + u; + data_c[0] = ((rows + p.pad_to - 1) / p.pad_to) * p.pad_to; + data_c[1] = u; + } +} 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 2b8642e45f8..38e9e4e0af5 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp @@ -822,6 +822,8 @@ void process_shaders() { string_to_spv("flash_attn_top_k_cm_f16", "flash_attn_top_k_cm.comp", {}); #endif string_to_spv("flash_attn_gather_f16", "flash_attn_gather.comp", {}); + string_to_spv("flash_attn_union_f16", "flash_attn_union.comp", {}); + string_to_spv("flash_attn_gather_union_f16", "flash_attn_gather_union.comp", {}); string_to_spv("dsv4_hc_pre_f32", "dsv4_hc_pre.comp", {}); string_to_spv("dsv4_hc_comb_f32", "dsv4_hc_comb.comp", {}); string_to_spv("dsv4_hc_post_f32", "dsv4_hc_post.comp", {}); diff --git a/tests/test-backend-ops.cpp b/tests/test-backend-ops.cpp index a3fd3562dbc..c6c4796c26c 100644 --- a/tests/test-backend-ops.cpp +++ b/tests/test-backend-ops.cpp @@ -8039,12 +8039,13 @@ struct test_flash_attn_ext_top_k : public test_case { const int64_t n_top_k; // selected keys per query token const bool sinks; const int64_t ns; // sequences (ne3); >1 exercises the split-K stream stride + const int64_t ov; // % of each token's picks shared with its neighbours (dedup-union realism) static constexpr int64_t hs = 512; // V4 CSA head size, K == V latent static constexpr int64_t nh = 64; // V4 CSA query heads (MQA) std::string vars() override { - return VARS_TO_STR6(kv, nb, n_kv_raw, n_top_k, sinks, ns); + return VARS_TO_STR7(kv, nb, n_kv_raw, n_top_k, sinks, ns, ov); } double max_nmse_err() override { @@ -8058,8 +8059,8 @@ struct test_flash_attn_ext_top_k : public test_case { return 2 * nh * nb * ns * (hs + hs) * (n_kv_raw + n_top_k); } - test_flash_attn_ext_top_k(int64_t kv = 768, int64_t nb = 8, int64_t n_kv_raw = 64, int64_t n_top_k = 128, bool sinks = false, int64_t ns = 1) - : kv(kv), nb(nb), n_kv_raw(n_kv_raw), n_top_k(n_top_k), sinks(sinks), ns(ns) {} + test_flash_attn_ext_top_k(int64_t kv = 768, int64_t nb = 8, int64_t n_kv_raw = 64, int64_t n_top_k = 128, bool sinks = false, int64_t ns = 1, int64_t ov = 0) + : kv(kv), nb(nb), n_kv_raw(n_kv_raw), n_top_k(n_top_k), sinks(sinks), ns(ns), ov(ov) {} ggml_tensor * build_graph(ggml_context * ctx) override { ggml_tensor * q = ggml_new_tensor_4d(ctx, GGML_TYPE_F32, hs, nb, nh, ns); @@ -8125,7 +8126,10 @@ struct test_flash_attn_ext_top_k : public test_case { for (int64_t j = 0; j < n_top_k; ++j) { // offset the selection by the stream too, so a dropped stream stride // reads another sequence's keys and shows up as a mismatch - int32_t idx = (int32_t) ((j * range) / n_top_k + b + s * 7) % (int32_t) range; + const bool shared = (int64_t) j * 100 < n_top_k * ov; + int32_t idx = shared + ? (int32_t) ((j * range) / n_top_k + s * 7) % (int32_t) range + : (int32_t) ((j * range) / n_top_k + b + s * 7) % (int32_t) range; if (j == n_top_k - 1 && b == 0 && s == 0) { idx = -1; // exercise the ignore-invalid-index path } else { @@ -11941,6 +11945,14 @@ static std::vector> make_test_cases_perf() { test_cases.emplace_back(new test_flash_attn_ext_top_k(kv, nb, 2304, 512, false)); } } + // Same shapes with realistic adjacent-token overlap. Measured on DeepSeek-V4-Flash the + // real overlap is 60% over 4 adjacent tokens and 76% over 8; the default generator is + // near 0%, which would make a deduplicated union look worthless by construction. + for (int kv : { 35584, 133888 }) { + for (int nb : { 2, 4, 8 }) { + test_cases.emplace_back(new test_flash_attn_ext_top_k(kv, nb, 2304, 512, false, 1, 60)); + } + } return test_cases; } From 628788d27a8de55f53b5b95d2272319e3645899b Mon Sep 17 00:00:00 2001 From: Nathan Wilson Date: Fri, 14 Aug 2026 01:21:31 +0000 Subject: [PATCH 30/68] vulkan: gate the DeepSeek V4 small-batch union on the measured union size The union shrinks the compact set but the gate could not see it. kv_c is the worst case n_kv_raw + n_batch*n_top_k, so at batch 8 and 32k depth the gate priced 6400 rows against 11008 source rows, declined, and dense attention ran over the whole compressed KV -- even though those selections deduplicate to about 3300 rows. That is the batch ceiling the previous commit documented and did not lift. The union size is only known on the device, and the note on the previous commit assumed lifting the ceiling therefore needed a GPU-side fallback for an oversized union. It does not. The union is bounded by the source, so there is no dispatch to recover from; the missing piece was only a measurement. The union shader now writes its count into a small host-visible buffer and the host prices the next step from it. The read is deliberately unsynchronised and one graph stale -- overlap is a property of the model and the draft, not of one op -- and a wrong estimate costs part of one step, never correctness, because every allocation and dispatch bound is still the worst case. Estimates are held per batch size. A speculative decode varies the batch with the accept count, and the overlap itself varies with the batch (0.64 at 2 tokens, 0.40 at 4, 0.24 at 8), so a single slot would be invalidated on nearly every step. Each estimate tracks the latest measurement, lightly smoothed, and a declined step dispatches the scan in a new count_only mode that writes no index list. Both of those are deliberate and were measured the other way round first. Holding a decaying peak to stay conservative, and sampling the probe every 256 declines to stay cheap, together produced 1992 us on a shape the union runs in 898: the asymmetry actually runs the other way, because an estimate that is too high declines compaction and forgoes 2-3x for as long as it stays high, while one that is too low costs a single step and is corrected by the count that step produces. Sampling makes it worse still, since a decline is exactly the state in which the compact path stops refreshing the estimate, so the stale value stays latched for the whole sampling period. The probe is one workgroup against the ~2.2 ms dense op it rides along with, and it costs 0.2% of a declined step. The same measurement settles the reverse case. Where selections do not overlap, the union is the same size as the per-token blocks and its scan is pure cost, so those now stay on the per-token path rather than being admitted whenever the worst-case gate happened to allow them. test-backend-ops perf, medians of 3 launches counterbalanced A B B A A B, spreads at or under 0.4%, n_kv_raw=2304 n_top_k=512, at kv=11008 (~32k source) which is the shape where the gate used to decline. The fixture's ov is a per-token share, not the union/selected ratio the model was measured by: at nb tokens it yields (ov + (1-ov)*nb)/nb of the selections, so ov=86 is the setting that reproduces the 0.243 measured on DeepSeek-V4-Flash over 8 adjacent tokens, and ov=60 is deliberately more pessimistic than the model. batch overlap before after 8 ov=86 2215.2 us 704.8 us 3.14x 16 ov=86 2227.3 us 841.3 us 2.65x 8 ov=60 2214.5 us 897.9 us 2.47x 16 ov=60 2226.2 us 2231.2 us 1.00x union does not fit, declined 8 ov=0 2213.5 us 2219.0 us 1.00x nothing to deduplicate, declined The union cost is depth-independent as before, so the lifted cells now sit alongside the deeper ones: 897.9 / 900.7 / 906.6 us at batch 8 ov=60 across kv 11008 / 35584 / 133888. Every cell at kv 35584 and 133888 is 1.00x: this changes which shapes are admitted, not how the union performs once admitted. Both declines above are the estimate working rather than failing. At ov=60 and batch 16 the union really is 5888 rows against 11008 source rows and compaction would not pay; at ov=0 there is no overlap to exploit at all. Verified from the executed graph with GGML_VK_FA_UNION_STATS=1, added here, which reports the measured union/candidate ratio and the resulting decision rather than leaving engagement to be inferred from a timing. It reads 0.246 at ov=86 and batch 8, against the 0.243 measured on the model. 13310 FLASH_ATTN_EXT cases pass with the union off and on, including three new overlap cases at the shapes the gate now admits. End to end on a model that has no top-k tensor at all, so the compaction path is never reached and the change should be structurally inert: Qwen3-Coder-30B-A3B UD-Q4_K_XL, same counterbalanced order, medians of 3. pp512 1578.20 -> 1591.67 t/s and tg32 96.89 -> 97.28 at d0; pp512 700.64 -> 699.32 and tg32 53.92 -> 54.45 at d16384. All within 1%. The union path itself cannot be checked this way: it needs batch 2..63 decode, which only a speculative draft produces. Still measured against a synthetic fixture rather than real draft tokens: that needs a runnable target and draft pair, which this box cannot host. Suggested-by: Jaap Buurman (@Mushoz) Co-Authored-By: Claude Opus 5 Assisted-by: Claude (Opus 5) --- ggml/src/ggml-vulkan/ggml-vulkan.cpp | 214 +++++++++++++++--- .../vulkan-shaders/flash_attn_union.comp | 14 +- tests/test-backend-ops.cpp | 21 +- 3 files changed, 216 insertions(+), 33 deletions(-) diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp index 1e22e8be459..182f51741cd 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp +++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp @@ -1987,7 +1987,7 @@ struct vk_op_dsv4_hc_post_push_constants { static_assert(sizeof(vk_op_dsv4_hc_post_push_constants) <= 128); struct vk_op_flash_attn_union_push_constants { - uint32_t n_kv, n_kv_raw, n_batch, n_top_k, max_union, nbt1, max_words, pad_to; + uint32_t n_kv, n_kv_raw, n_batch, n_top_k, max_union, nbt1, max_words, pad_to, count_only; }; struct vk_op_flash_attn_gather_union_push_constants { uint32_t n_kv, n_kv_raw, kv_c_max, nbk1, nbm1, n_batch; @@ -2509,6 +2509,14 @@ struct ggml_backend_vk_context { uint64_t fa_dequant_gate_sz; bool fa_dequant_gate_fits; bool fa_dequant_gate_logged; + // DeepSeek V4 small-batch union: the compact row count is produced on the device, so the + // host prices compaction from the last count it wrote. See ggml_vk_fa_union_estimate. + // Per batch size, because the overlap depends on it (measured 0.64 at 2 tokens, 0.40 at 4, + // 0.24 at 8) and because a speculative decode varies the batch with the accept count, so a + // single slot would be invalidated on nearly every step. This path caps the batch at 64. + vk_buffer fa_union_stat; + float fa_union_est_ratio[64]; // union / candidates, decaying peak; 0 = unseeded + uint64_t fa_union_declines; vk::Fence fence, almost_ready_fence; bool submit_pending {}; bool almost_ready_fence_pending {}; @@ -8059,6 +8067,8 @@ static void ggml_vk_init(ggml_backend_vk_context * ctx, size_t idx) { ctx->fa_dequant_gate_sz = 0; ctx->fa_dequant_gate_fits = false; ctx->fa_dequant_gate_logged = false; + memset(ctx->fa_union_est_ratio, 0, sizeof(ctx->fa_union_est_ratio)); + ctx->fa_union_declines = 0; // Fixed size of 1KB, for deterministic behavior ctx->prealloc_size_add_rms_partials = 1024; @@ -11715,6 +11725,80 @@ struct vk_fa_compact_state { vk_subbuffer kc_buf, mc_buf; }; +// Small host-visible buffer holding the last union count the device produced: +// [0] padded compact rows (also read by the gather and by the FA), [1] raw union size, +// [2] the candidate count it came from, [3] the batch it came from. Host-visible is a +// requirement rather than a preference here - the whole point is that the host can read it +// without submitting. +static bool ggml_vk_fa_union_stat_init(ggml_backend_vk_context * ctx) { + if (ctx->fa_union_stat) { + return ctx->fa_union_stat->ptr != nullptr; + } + try { + ctx->fa_union_stat = ggml_vk_create_buffer(ctx->device, 64, + {vk::MemoryPropertyFlagBits::eDeviceLocal | vk::MemoryPropertyFlagBits::eHostVisible | vk::MemoryPropertyFlagBits::eHostCoherent, + vk::MemoryPropertyFlagBits::eHostVisible | vk::MemoryPropertyFlagBits::eHostCoherent}); + } catch (const vk::SystemError &) { + return false; + } + if (ctx->fa_union_stat->ptr == nullptr) { + return false; + } + memset(ctx->fa_union_stat->ptr, 0, 64); + return true; +} + +// Host-coherent memory still needs the write made available to the host domain. This is the +// only barrier in the path that names eHost, and it is cheap: nothing waits on it, it just +// lets the next graph's host-side read see the count this one produced. +static void ggml_vk_fa_union_stat_host_barrier(vk_context & subctx) { + subctx->s->buffer->buf.pipelineBarrier( + vk::PipelineStageFlagBits::eComputeShader, + vk::PipelineStageFlagBits::eHost, + {}, + { { vk::AccessFlagBits::eShaderWrite, vk::AccessFlagBits::eHostRead } }, + {}, {}); +} + +// Price the union without synchronising for it. The count read here belongs to whichever call +// last wrote the slot, which under a decode is the final flash-attention op of the previous +// graph - one token stale, and that is fine: selection overlap is a property of the model and +// the draft, not of an individual op, and a wrong estimate costs part of one step rather than +// correctness (the compact buffers are still sized for the worst case). +// +// Tracks the latest measurement, lightly smoothed. The asymmetry runs the other way from what +// a conservative estimator would assume: an estimate that is too HIGH declines compaction and +// forgoes 2-3x for as long as it stays high, while one that is too low costs a single step at +// roughly dense cost and is corrected by the count that step produces. An earlier version held +// a decaying peak instead and measured 1992 us where the union delivers 900, because a spell of +// genuinely low overlap pinned the estimate and 0.999 per read took hundreds of steps to relax. +// +// The words are read without ordering against the device write, so they can come from +// different calls. The sample is filed under the batch the device reported rather than the +// batch being priced, so a torn read costs one mispriced step for that batch and then +// corrects, which is the same failure the estimate already tolerates. +// +// Returns the padded compact row count to gate on, or 0 when this batch is unseeded. +static uint32_t ggml_vk_fa_union_estimate(ggml_backend_vk_context * ctx, uint32_t n_kv_raw, + uint32_t n_batch, uint32_t n_cand) { + const volatile uint32_t * stat = (const volatile uint32_t *) ctx->fa_union_stat->ptr; + const uint32_t u = stat[1]; + const uint32_t cand = stat[2]; + const uint32_t nb_obs = stat[3]; + + if (u > 0 && cand > 0 && nb_obs > 0 && nb_obs < 64) { + const float r = std::min(1.0f, (float) u / (float) cand); + float & e = ctx->fa_union_est_ratio[nb_obs]; + e = e > 0.0f ? 0.5f * r + 0.5f * e : r; + } + + const float ratio = ctx->fa_union_est_ratio[n_batch]; + if (ratio <= 0.0f) { + return 0; + } + return GGML_PAD(n_kv_raw + (uint32_t) ceilf(ratio * (float) n_cand), 256u); +} + // V4 sparse decode (gather-to-compact): the sparse prefill shader above gates on // q->ne[1] >= 64, so single-token decode otherwise attends densely over the whole // compressed KV, at a cost that grows with context. Instead, gather the active rows @@ -11754,20 +11838,17 @@ static bool ggml_vk_flash_attn_gather_compact(ggml_backend_vk_context * ctx, vk_ } const uint32_t n_batch = (uint32_t) q->ne[1]; - // Every token gets its own top-k block, so the compact set is n_kv_raw + n_batch*n_top_k: - // independent of context depth, which is the point, but GROWING WITH BATCH. Attention work - // is then n_batch * (n_kv_raw + n_batch*n_top_k), i.e. quadratic in batch, against dense's - // n_batch * n_kv. Break-even is n_batch = (n_kv - n_kv_raw) / n_top_k, and the - // kv >= 2*kv_c gate below caps the useful batch at (n_kv/2 - n_kv_raw) / n_top_k -- - // about 6 at 32k depth, ~30 at 128k, batch-capped at 512k. Beyond that the gate declines - // and dense runs, so this can never be slower; it just stops helping. A deduplicated union - // would lift that ceiling wherever draft tokens select overlapping keys. - const uint32_t kv_c = GGML_PAD((uint32_t)(n_kv_raw + (int64_t) n_batch * top_k->ne[0]), 256u); - // the gather writes then re-reads ~the active bytes; dense reads the source KV once, - // so compaction only pays when the source is comfortably larger than the active set - if ((uint64_t) k->ne[1] < 2ull * kv_c) { - return false; - } + const uint32_t n_cand = (uint32_t) ((int64_t) n_batch * top_k->ne[0]); + // Worst case: every token gets its own top-k block, so the compact set is + // n_kv_raw + n_batch*n_top_k -- independent of context depth, which is the point, but + // GROWING WITH BATCH. Attention work is then n_batch * (n_kv_raw + n_batch*n_top_k), i.e. + // quadratic in batch, against dense's n_batch * n_kv. Break-even is + // n_batch = (n_kv - n_kv_raw) / n_top_k, and a kv >= 2*kv_c gate on this worst case caps + // the useful batch at (n_kv/2 - n_kv_raw) / n_top_k -- about 6 at 32k depth, ~30 at 128k. + // Beyond that the gate declines and dense runs, so this can never be slower; it just stops + // helping. The union below lifts that ceiling by gating on what the selections actually + // deduplicate to; kv_c remains the bound for every allocation and dispatch count. + const uint32_t kv_c = GGML_PAD((uint32_t) (n_kv_raw + (int64_t) n_cand), 256u); // ---- deduplicated union (GGML_VK_FA_TOPK_UNION=1) ---------------------------------- // Same compact layout, but one row per DISTINCT selected key instead of one block per @@ -11776,39 +11857,101 @@ static bool ggml_vk_flash_attn_gather_compact(ggml_backend_vk_context * ctx, vk_ // reads its KV bound from a buffer (the DYNAMIC_KV pipeline flag) rather than a push // constant; padding the count to 256 keeps KV % Bc == 0 so the aligned variant still // applies. Single stream only: the FA takes one KV for all streams. + // + // Gated on the ESTIMATED union rather than on kv_c, which is why this sits above the + // worst-case gate: at batch 8 and 32k depth the worst case is 6400 rows against 11008 + // source rows and would decline, while the union measures around 3300 and is well worth + // compacting. kv_c stays the worst case for every allocation and dispatch bound, so a + // wrong estimate is a slow step, never a wrong answer. + const uint32_t max_words = 12288; // shared bitmap capacity in flash_attn_union.comp + const bool bitmap_fits = (uint64_t) ((k->ne[1] - n_kv_raw) + 31) / 32 <= max_words; + static const char * union_env = getenv("GGML_VK_FA_TOPK_UNION"); - if (union_env && union_env[0] == '1' && q->ne[3] == 1 && n_batch > 1 && - ctx->device->pipeline_flash_attn_union_f16 && ctx->device->pipeline_flash_attn_gather_union_f16) { - const uint32_t max_union = (uint32_t) ((int64_t) n_batch * top_k->ne[0]); - // shared bitmap capacity in flash_attn_union.comp - const uint32_t max_words = 12288; - const uint32_t need_words = (uint32_t) (((k->ne[1] - n_kv_raw) + 31) / 32); - if (need_words > max_words) { - goto union_unavailable; // fall through to the per-token block form + if (union_env && union_env[0] == '1' && q->ne[3] == 1 && n_batch > 1 && bitmap_fits && + ctx->device->pipeline_flash_attn_union_f16 && ctx->device->pipeline_flash_attn_gather_union_f16 && + ggml_vk_fa_union_stat_init(ctx)) { + const uint32_t max_union = n_cand; + const uint32_t kv_c_est = ggml_vk_fa_union_estimate(ctx, (uint32_t) n_kv_raw, n_batch, n_cand); + // Two separate questions. Does the compact set fit under the gate at all, and does + // deduplicating actually shrink it: with no overlap to exploit the union is the same + // size as the per-token blocks and the scan is pure cost, measured at 1.2% of the op + // at 512k depth. The worst-case bound on the source keeps a collapse in overlap to + // roughly dense cost for the one step it takes the estimate to catch up. + const bool worth_it = kv_c_est != 0 && kv_c_est < kv_c && + (uint64_t) k->ne[1] >= 2ull * kv_c_est && + (uint64_t) k->ne[1] >= (uint64_t) kv_c; + + // GGML_VK_FA_UNION_STATS=1: what the gate actually decided and on what measurement. + // The alternative is inferring engagement from a timing, which is how a sparse path + // gets credited for a run it never took. + static const char * stats_env = getenv("GGML_VK_FA_UNION_STATS"); + if (stats_env && stats_env[0] == '1') { + static uint64_t calls = 0; + if ((calls++ % 256) == 0) { + fprintf(stderr, "[fa-union] n_kv=%lld n_kv_raw=%d n_batch=%u cand=%u " + "union/cand=%.3f kv_c %u -> est %u %s\n", + (long long) k->ne[1], n_kv_raw, n_batch, n_cand, (double) ctx->fa_union_est_ratio[n_batch], + kv_c, kv_c_est, worth_it ? "UNION" : "declined"); + } + } + + if (!worth_it) { + // Nothing is known about this batch shape, or the last count says compaction does + // not pay. Either way the count is what settles it, so produce one: the scan is a + // single workgroup and depth-independent, and count_only needs no index list. Dense + // runs this step; the next one decides on a measurement instead of a bound. + // + // Every declined step, not a sample: a decline is exactly the state in which the + // estimate stops being refreshed by the compact path, so sampling it leaves a stale + // estimate latched for as many steps as the sampling period. The dispatch is one + // workgroup against the ~2.2 ms dense op it is riding along with. + ctx->fa_union_declines++; + { + const vk_op_flash_attn_union_push_constants ppc = { + (uint32_t) k->ne[1], (uint32_t) n_kv_raw, n_batch, (uint32_t) top_k->ne[0], max_union, + (uint32_t) (top_k->nb[1] / sizeof(int32_t)), max_words, 256u, 1u, + }; + const vk_subbuffer stat_buf = ggml_vk_subbuffer(ctx, ctx->fa_union_stat); + ggml_pipeline_request_descriptor_sets(ctx, ctx->device->pipeline_flash_attn_union_f16, 1); + // The probe writes the same slot a taken union path reads back as its row + // count, and a graph can contain both: the estimate is refreshed from the + // device as the graph is recorded, so an op late in the graph can be admitted + // after an earlier one was declined. Order it explicitly rather than rely on + // the compact path's own sync, which this path does not go through. + ggml_vk_sync_buffers(ctx, subctx); + ggml_vk_dispatch_pipeline(ctx, subctx, ctx->device->pipeline_flash_attn_union_f16, + { ggml_vk_tensor_subbuffer(ctx, top_k), stat_buf, stat_buf }, ppc, { 1, 1, 1 }); + ggml_vk_fa_union_stat_host_barrier(subctx); + } + goto union_unavailable; } + const size_t ukc_sz = (size_t) kv_c * 512 * sizeof(ggml_fp16_t); const size_t umc_sz = (size_t) n_batch * kv_c * sizeof(ggml_fp16_t); const size_t ul_sz = (size_t) max_union * sizeof(uint32_t); - const size_t uc_sz = 2 * sizeof(uint32_t); - const size_t need = ukc_sz + umc_sz + ul_sz + uc_sz; + const size_t need = ukc_sz + umc_sz + ul_sz; if (ctx->prealloc_size_y < need) { ctx->prealloc_size_y = need; ggml_vk_preallocate_buffers(ctx, subctx); } - if (ctx->prealloc_y_need_sync) { - ggml_vk_sync_buffers(ctx, subctx); - } + // Unconditional, not gated on prealloc_y_need_sync: this also orders the count slot + // against a probe dispatched earlier in the same graph, which leaves that flag clear. + ggml_vk_sync_buffers(ctx, subctx); const vk_subbuffer kc_buf = ggml_vk_subbuffer(ctx, ctx->prealloc_y, 0); const vk_subbuffer mc_buf = ggml_vk_subbuffer(ctx, ctx->prealloc_y, ukc_sz); const vk_subbuffer ul_buf = ggml_vk_subbuffer(ctx, ctx->prealloc_y, ukc_sz + umc_sz); - const vk_subbuffer uc_buf = ggml_vk_subbuffer(ctx, ctx->prealloc_y, ukc_sz + umc_sz + ul_sz); + // The count lives in the stat buffer rather than in prealloc_y so that the same write + // that the gather and the FA consume is also the one the host prices the next step from. + // Both are single-slot and rewritten by every layer, so the WAR hazard is unchanged: the + // prealloc_y sync above is a global barrier and orders the previous layer's read. + const vk_subbuffer uc_buf = ggml_vk_subbuffer(ctx, ctx->fa_union_stat); ggml_pipeline_request_descriptor_sets(ctx, ctx->device->pipeline_flash_attn_union_f16, 1); ggml_pipeline_request_descriptor_sets(ctx, ctx->device->pipeline_flash_attn_gather_union_f16, 1); const vk_op_flash_attn_union_push_constants upc = { (uint32_t) k->ne[1], (uint32_t) n_kv_raw, n_batch, (uint32_t) top_k->ne[0], max_union, - (uint32_t) (top_k->nb[1] / sizeof(int32_t)), max_words, 256u, + (uint32_t) (top_k->nb[1] / sizeof(int32_t)), max_words, 256u, 0u, }; ggml_vk_dispatch_pipeline(ctx, subctx, ctx->device->pipeline_flash_attn_union_f16, { ggml_vk_tensor_subbuffer(ctx, top_k), ul_buf, uc_buf }, upc, { 1, 1, 1 }); @@ -11824,6 +11967,7 @@ static bool ggml_vk_flash_attn_gather_compact(ggml_backend_vk_context * ctx, vk_ { ggml_vk_tensor_subbuffer(ctx, k), ul_buf, ggml_vk_tensor_subbuffer(ctx, mask), kc_buf, mc_buf, uc_buf }, gpc, { kv_c, 1, 1 }); ggml_vk_sync_buffers(ctx, subctx); + ggml_vk_fa_union_stat_host_barrier(subctx); ctx->prealloc_y_need_sync = true; st.active = true; @@ -11837,6 +11981,13 @@ static bool ggml_vk_flash_attn_gather_compact(ggml_backend_vk_context * ctx, vk_ } union_unavailable:; + // Per-token blocks have no dedup, so this form really does cost kv_c: the gather writes + // then re-reads ~the active bytes while dense reads the source KV once, so compaction only + // pays when the source is comfortably larger than the active set. + if ((uint64_t) k->ne[1] < 2ull * kv_c) { + return false; + } + const uint32_t ns = (uint32_t) q->ne[3]; const size_t kc_sz = (size_t) ns * kv_c * 512 * sizeof(ggml_fp16_t); const size_t mc_sz = (size_t) ns * n_batch * kv_c * sizeof(ggml_fp16_t); @@ -17574,8 +17725,11 @@ static void ggml_vk_cleanup(ggml_backend_vk_context * ctx) { ggml_vk_destroy_buffer(ctx->prealloc_y); ggml_vk_destroy_buffer(ctx->prealloc_split_k); ggml_vk_destroy_buffer(ctx->prealloc_add_rms_partials); + ggml_vk_destroy_buffer(ctx->fa_union_stat); ggml_vk_destroy_buffer(ctx->sync_staging); + memset(ctx->fa_union_est_ratio, 0, sizeof(ctx->fa_union_est_ratio)); + ctx->prealloc_y_last_pipeline_used = nullptr; ctx->prealloc_y_last_tensor_used = nullptr; ctx->prealloc_y_last_k_padded = false; diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_union.comp b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_union.comp index cd779e0a0f9..51337570930 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_union.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_union.comp @@ -19,6 +19,15 @@ // (a multiple of every FA block width) is what lets flash-attention keep its "aligned" // pipeline variant even though the row count is now a runtime value. // +// Also emits the raw union size, the candidate count and the batch it came from. The host +// reads those back to decide whether compaction is worth it at all, and files the sample under +// the batch reported here rather than the one it is pricing: it cannot tell which call last +// wrote the slot, and the overlap it is measuring depends on the batch. +// +// count_only skips the list writes. The scan still has to run to produce the count, but with +// nothing to write the index buffer need not exist, which is what lets the host price the +// union on a call it is not going to compact. +// // One workgroup: the bitmap and the running offset both live in shared memory. layout(local_size_x = 256, local_size_y = 1, local_size_z = 1) in; @@ -36,6 +45,7 @@ layout(push_constant) uniform Parameters { uint nbt1; uint max_words; // capacity of the shared bitmap, host-checked uint pad_to; + uint count_only; // 1: produce the count, write no index list } p; // 12288 words = 393216 compressed rows, ~48 KiB of shared memory. @@ -97,7 +107,7 @@ void main() { } uint slot = base_sh + prefix + wave_off; - uint rem = bits; + uint rem = p.count_only != 0 ? 0u : bits; while (rem != 0) { const uint b = findLSB(rem); rem &= rem - 1; @@ -118,5 +128,7 @@ void main() { const uint rows = p.n_kv_raw + u; data_c[0] = ((rows + p.pad_to - 1) / p.pad_to) * p.pad_to; data_c[1] = u; + data_c[2] = n_cand; + data_c[3] = p.n_batch; } } diff --git a/tests/test-backend-ops.cpp b/tests/test-backend-ops.cpp index c6c4796c26c..6ab027c7ea1 100644 --- a/tests/test-backend-ops.cpp +++ b/tests/test-backend-ops.cpp @@ -11420,6 +11420,12 @@ static std::vector> make_test_cases_eval() { test_cases.emplace_back(new test_flash_attn_ext_top_k(8192, 8, 1024, 512, true)); test_cases.emplace_back(new test_flash_attn_ext_top_k(32768, 16, 2304, 512, false)); test_cases.emplace_back(new test_flash_attn_ext_top_k(65536, 63, 2304, 512, false)); + // overlapping selections at the shapes where the compaction gate is tightest: with a + // deduplicated union these are admitted on the estimated union size rather than the + // worst case, so they cover the estimator's gate as well as the union itself. + test_cases.emplace_back(new test_flash_attn_ext_top_k(11008, 8, 2304, 512, false, 1, 60)); + test_cases.emplace_back(new test_flash_attn_ext_top_k(11008, 16, 2304, 512, false, 1, 60)); + test_cases.emplace_back(new test_flash_attn_ext_top_k(11008, 16, 2304, 512, false, 1, 86)); test_cases.emplace_back(new test_flash_attn_ext_top_k(8192, 4, 1024, 512, false, 2)); test_cases.emplace_back(new test_flash_attn_ext_top_k( 768, 64, 64, 128, false, 2)); test_cases.emplace_back(new test_flash_attn_ext_top_k( 768, 64, 64, 128, true, 2)); @@ -11948,11 +11954,22 @@ static std::vector> make_test_cases_perf() { // Same shapes with realistic adjacent-token overlap. Measured on DeepSeek-V4-Flash the // real overlap is 60% over 4 adjacent tokens and 76% over 8; the default generator is // near 0%, which would make a deduplicated union look worthless by construction. - for (int kv : { 35584, 133888 }) { - for (int nb : { 2, 4, 8 }) { + // kv=11008 (~32k source) and nb=16 are where the compaction gate is tightest: the + // worst-case compact set 2304 + nb*512 crosses kv/2 at nb=6, so those cells measure + // whether the gate can be opened by the union rather than by the worst case. + for (int kv : { 11008, 35584, 133888 }) { + for (int nb : { 2, 4, 8, 16 }) { test_cases.emplace_back(new test_flash_attn_ext_top_k(kv, nb, 2304, 512, false, 1, 60)); } } + // ov is a per-token share, not the union/selected ratio the model was measured by: at nb + // tokens it gives a union of (ov + (1-ov)*nb)/nb of the selections, so ov=60 is 0.475 at + // nb=8 where the model measured 0.243. ov=86 is the setting that reproduces the model, and + // at kv=11008 it is the difference between a union that fits under the gate and one that + // does not. + for (int nb : { 8, 16 }) { + test_cases.emplace_back(new test_flash_attn_ext_top_k(11008, nb, 2304, 512, false, 1, 86)); + } return test_cases; } From 857f81cf66bb0c8240a3ec52ae8ba3d263a38726 Mon Sep 17 00:00:00 2001 From: Nathan Wilson Date: Fri, 14 Aug 2026 02:27:30 +0000 Subject: [PATCH 31/68] vulkan: measure the V4 union on real draft tokens, and correct the fixture calibration The previous commit calibrated its fixture against a PROXY, because a proxy was all that was reachable at the time: adjacent PREFILL tokens standing in for draft tokens. With the draft now drivable, the proxy turns out to have been optimistic, and the correction runs against that commit's own message. DeepSeek-V4-Flash UD-IQ3_XXS with the DSpark draft, Vulkan attention and CPU experts, 25538 tokens of context. The gate reports union/candidate 0.560 and 0.609 on the target layers and 0.484 to 0.651 on a second shape (n_kv_raw 256, most consistent with the draft's own sparse attention, also GPU-resident). Mean about 0.56 at batch 4, against 0.397 from the proxy at a window of 4. Draft tokens diverge from each other more than adjacent prefill tokens do, which is what a draft exploring rather than committed text should look like. The fixture's ov maps to (ov + (1-ov)*nb)/nb, so ov=60 predicts 0.55 at batch 4 against 0.56 measured: ov=60 reproduces real drafting almost exactly. The previous commit says ov=86 is the setting that reproduces the model and ov=60 is deliberately pessimistic. That is backwards for real draft tokens, and the realistic headline at batch 8 is 2.47x rather than 3.14x. The kernel numbers in that table are unchanged and still correct for the overlap each column states; only which column describes reality has moved. Batch 8 overlap itself is still not measured on real drafts, for the reason below, so 2.47x is an extrapolation from batch 4. The path does engage: 56 of 57 sampled gate decisions took the union, and the single decline is the unseeded first call, which then probes. That is the designed bootstrap, observed on the real model rather than a fixture. Scope, stated plainly because the previous commit implies more: DSpark drafts 3 tokens, so the batch is 4 and --spec-draft-n-max does not raise it, an MTP head having a fixed width. At batch 4 and this depth kv_c is 4352 against n_kv 8704, so the OLD worst-case gate admits too and the gate change makes no difference there; the smaller compact set comes from the union commit. For this target and draft the gate change buys a lower depth threshold instead: roughly 21.0k tokens rather than 25.5k on the target layers, and 11.8k rather than 17.7k on the second shape. The 2.47x and 2.65x cells need batch 8 or 16, which this pair does not generate. GGML_VK_FA_UNION_STATS takes a period now rather than being fixed at 256, and reports running totals. The fixed period is why this took two runs: a 64-token speculative decode over 25k of context emitted exactly ONE line, which reported only that the first call was unseeded. A diagnostic whose sampling rate can hide the thing it exists to measure is not a diagnostic. Measurements and provenance: ~/strix-results/derived-dsv4-union-gate-20260814.md section 9. Co-Authored-By: Claude Opus 5 Assisted-by: Claude (Opus 5) --- ggml/src/ggml-vulkan/ggml-vulkan.cpp | 24 ++++++++++++++++-------- 1 file changed, 16 insertions(+), 8 deletions(-) diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp index 182f51741cd..07912c4cf7a 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp +++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp @@ -11881,17 +11881,25 @@ static bool ggml_vk_flash_attn_gather_compact(ggml_backend_vk_context * ctx, vk_ (uint64_t) k->ne[1] >= 2ull * kv_c_est && (uint64_t) k->ne[1] >= (uint64_t) kv_c; - // GGML_VK_FA_UNION_STATS=1: what the gate actually decided and on what measurement. - // The alternative is inferring engagement from a timing, which is how a sparse path - // gets credited for a run it never took. + // GGML_VK_FA_UNION_STATS=N: report every Nth call (N=1 means every call) what the gate + // decided and on what measurement. The alternative is inferring engagement from a + // timing, which is how a sparse path gets credited for a run it never took. + // + // The period is a parameter because a fixed one is a way to miss the answer: a 64-token + // speculative decode over 25k of context produced ONE line at a period of 256, which + // said only that the first call was unseeded. Running totals rather than instants, so a + // single late line still reports whether the path engaged. static const char * stats_env = getenv("GGML_VK_FA_UNION_STATS"); - if (stats_env && stats_env[0] == '1') { - static uint64_t calls = 0; - if ((calls++ % 256) == 0) { + if (stats_env && stats_env[0] != '\0' && stats_env[0] != '0') { + static uint64_t calls = 0, taken = 0; + const uint64_t period = std::max(1ull, (unsigned long long) atoll(stats_env)); + taken += worth_it ? 1 : 0; + if ((calls++ % period) == 0) { fprintf(stderr, "[fa-union] n_kv=%lld n_kv_raw=%d n_batch=%u cand=%u " - "union/cand=%.3f kv_c %u -> est %u %s\n", + "union/cand=%.3f kv_c %u -> est %u %s (%llu/%llu taken)\n", (long long) k->ne[1], n_kv_raw, n_batch, n_cand, (double) ctx->fa_union_est_ratio[n_batch], - kv_c, kv_c_est, worth_it ? "UNION" : "declined"); + kv_c, kv_c_est, worth_it ? "UNION" : "declined", + (unsigned long long) taken, (unsigned long long) calls); } } From fa825766d42c596f8f622d8a256dbcf463846c92 Mon Sep 17 00:00:00 2001 From: Nathan Wilson Date: Fri, 14 Aug 2026 03:11:13 +0000 Subject: [PATCH 32/68] vulkan: default the DeepSeek V4 small-batch union on Every other gate in this family is on by default with =0 to disable: GGML_VK_FA_TOPK, GGML_VK_FA_TOPK_CM, GGML_VK_FA_TOPK_SPLIT. GGML_VK_FA_TOPK_UNION=1 was the odd one out because it began as a prototype, and it is no longer one. What makes it safe to default is that the path declines itself rather than needing a user to know when to avoid it. Where the selections do not deduplicate, the estimate says so and the per-token form runs instead; where the compact set would not fit under the gate, dense runs. Both were measured, at 1.00x, rather than assumed. Kept as one commit of its own rather than folded into the gate change, because it is a policy change and not a mechanism one: reverting these four lines restores opt-in behaviour without touching the gate logic, which is what a bisect would want if a regression turns up. 13310 FLASH_ATTN_EXT cases pass with the variable UNSET, which is the configuration this changes and the one that ships. That is the same code path as the previously validated GGML_VK_FA_TOPK_UNION=1, but it was run again rather than argued from equivalence. Co-Authored-By: Claude Opus 5 Assisted-by: Claude (Opus 5) --- ggml/src/ggml-vulkan/ggml-vulkan.cpp | 8 ++++---- 1 file changed, 4 insertions(+), 4 deletions(-) diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp index 07912c4cf7a..f72612c91f9 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp +++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp @@ -11850,10 +11850,10 @@ static bool ggml_vk_flash_attn_gather_compact(ggml_backend_vk_context * ctx, vk_ // deduplicate to; kv_c remains the bound for every allocation and dispatch count. const uint32_t kv_c = GGML_PAD((uint32_t) (n_kv_raw + (int64_t) n_cand), 256u); - // ---- deduplicated union (GGML_VK_FA_TOPK_UNION=1) ---------------------------------- + // ---- deduplicated union (default on, GGML_VK_FA_TOPK_UNION=0 disables) -------------- // Same compact layout, but one row per DISTINCT selected key instead of one block per - // token. Measured adjacent-token overlap on the real model is 60% at 4 tokens and 76% at - // 8, so the union is materially smaller. Its size is only known on the GPU, so the FA + // token. Measured on real draft tokens the union is 0.56 of the selections at batch 4, + // so it is materially smaller. Its size is only known on the GPU, so the FA // reads its KV bound from a buffer (the DYNAMIC_KV pipeline flag) rather than a push // constant; padding the count to 256 keeps KV % Bc == 0 so the aligned variant still // applies. Single stream only: the FA takes one KV for all streams. @@ -11867,7 +11867,7 @@ static bool ggml_vk_flash_attn_gather_compact(ggml_backend_vk_context * ctx, vk_ const bool bitmap_fits = (uint64_t) ((k->ne[1] - n_kv_raw) + 31) / 32 <= max_words; static const char * union_env = getenv("GGML_VK_FA_TOPK_UNION"); - if (union_env && union_env[0] == '1' && q->ne[3] == 1 && n_batch > 1 && bitmap_fits && + if ((!union_env || union_env[0] != '0') && q->ne[3] == 1 && n_batch > 1 && bitmap_fits && ctx->device->pipeline_flash_attn_union_f16 && ctx->device->pipeline_flash_attn_gather_union_f16 && ggml_vk_fa_union_stat_init(ctx)) { const uint32_t max_union = n_cand; From 1481bbeff75160788bb05b4629dfd3c10979638e Mon Sep 17 00:00:00 2001 From: Nathan Wilson Date: Fri, 14 Aug 2026 06:27:04 +0000 Subject: [PATCH 33/68] vulkan: let the DeepSeek V4 small-batch gather serve quantised K/V The gather rejected anything but f16 K/V, so a DSv4 run with -ctk q8_0 -ctv q8_0 took the dense fallback for every batch 2..63 decode - the whole small-batch path, silently off. Two community test sets were run that way before anyone noticed, which is what prompted this. The restriction was incidental rather than essential. The gather RELOCATES rows; it never reads a value out of one. Addressing K as raw 4-byte words instead of f16 elements makes the same shader serve any type whose row is a whole number of words, which every block-quantised type here satisfies at head size 512. The union shader needed no change at all - it only ever touched top_k indices. Zeroing an unused row now writes zero BYTES. That decodes to zero for the block-quantised types as well as for f16, because a zero scale zeroes the block; those rows are -inf in the mask regardless, so zeroing only keeps a garbage dot product out of the softmax as a NaN. The type gate asks ggml_vk_fa_kv_native rather than just "is it quantised". While the compact scratch is active the dequant/contiguize pass is disabled by construction, and the assert guarding that combination aborts rather than falling back, so admitting a type flash-attention has no native shader for would crash instead of degrade. test-backend-ops perf, kv=11008 n_kv_raw=2304 n_top_k=512 ov=60, gather off vs on: batch f16 dense -> gathered q8_0 dense -> gathered 2 2192 -> 645 us 3.40x 3917 -> 1110 us 3.53x 4 2203 -> 726 us 3.03x 3924 -> 1257 us 3.12x 8 2213 -> 910 us 2.43x 3936 -> 1567 us 2.51x 16 2219 -> 2232 us declined 3930 -> 3952 us declined Two things worth stating because they contradict what I expected going in. Compaction does NOT pay more on quantised K/V. The reasoning was that a gathered set is contiguous and would dodge the strided-read and channel-aliasing taxes that hurt quantised reads more; the measured ratios differ by 3 to 4%, inside noise. It pays the same proportion from a worse starting point. That worse starting point is the real result: q8_0 attention is about 1.78x SLOWER than f16 here, dense (3917 vs 2192) and gathered (1110 vs 645) alike. Halving the bytes does not pay for the dequant work in the inner loop at this head size. So this change makes q8_0 K/V survivable for DSv4 rather than advisable - f16 is still the faster cache by a wide margin, and that inverts the usual q8_0-KV-buys-speed rule of thumb for this model. 13318 FLASH_ATTN_EXT cases pass, including 8 new q8_0/q4_0 top-k cases at the shapes a DSv4 decode actually hits. test_flash_attn_ext_top_k takes a type_K parameter for them. Co-Authored-By: Claude Opus 5 Assisted-by: Claude (Opus 5) --- ggml/src/ggml-vulkan/ggml-vulkan.cpp | 55 ++++++++++++++----- .../vulkan-shaders/flash_attn_gather.comp | 25 +++++---- .../flash_attn_gather_union.comp | 25 ++++++--- tests/test-backend-ops.cpp | 24 ++++++-- 4 files changed, 91 insertions(+), 38 deletions(-) diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp index f72612c91f9..de2cb074542 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp +++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp @@ -1989,12 +1989,14 @@ static_assert(sizeof(vk_op_dsv4_hc_post_push_constants) <= 128); struct vk_op_flash_attn_union_push_constants { uint32_t n_kv, n_kv_raw, n_batch, n_top_k, max_union, nbt1, max_words, pad_to, count_only; }; +// nbk1/nbk3 are in 4-byte WORDS, not elements: the gather relocates K rows verbatim and never +// interprets what is in them, so it works for any type whose row is a whole number of words. struct vk_op_flash_attn_gather_union_push_constants { - uint32_t n_kv, n_kv_raw, kv_c_max, nbk1, nbm1, n_batch; + uint32_t n_kv, n_kv_raw, kv_c_max, nbk1, nbm1, n_batch, row_words; }; struct vk_op_flash_attn_gather_push_constants { uint32_t n_kv, n_kv_raw, n_top_k, kv_c; - uint32_t nbk1, nbk3, nbt1, nbt3, nbm1, nbm3, nem3, n_batch; + uint32_t nbk1, nbk3, nbt1, nbt3, nbm1, nbm3, nem3, n_batch, row_words; }; static_assert(sizeof(vk_op_flash_attn_gather_push_constants) <= 128); @@ -11721,6 +11723,8 @@ struct vk_fa_compact_state { bool dynamic_kv = false; // KV row count lives in kv_buf, not the push constant uint32_t kv_c = 0; // upper bound; the real count is runtime when dynamic_kv uint32_t n_batch = 1; + uint32_t row_bytes = 0; // bytes per compact K row; K may be quantised + uint32_t row_elems = 0; // K row stride in ELEMENTS/blocks, for the FA push constant vk_subbuffer kv_buf; vk_subbuffer kc_buf, mc_buf; }; @@ -11811,10 +11815,22 @@ static bool ggml_vk_flash_attn_gather_compact(ggml_backend_vk_context * ctx, vk_ const ggml_tensor * mask, ggml_tensor * dst, vk_fa_compact_state & st) { const ggml_tensor * top_k = dst->src[5]; static const char * gather_env = getenv("GGML_VK_FA_TOPK_GATHER"); + // K/V type: the gather relocates rows verbatim and never reads a value out of them, so it + // does not have to be f16 - it has to be a type whose row is a whole number of 4-byte words, + // and one flash-attention has a NATIVE shader for. The dequant/contiguize pass is disabled + // while this scratch is active and asserts if it turns out to be needed, so admitting a + // non-native type here would abort rather than fall back. + const bool kv_word_addressable = + k->type == v->type && + ggml_vk_fa_kv_native(k->type, ctx->device->coopmat2) && + k->ne[0] % ggml_blck_size(k->type) == 0 && + ggml_row_size(k->type, k->ne[0]) % 4 == 0 && + k->nb[1] % 4 == 0 && k->nb[3] % 4 == 0; + if ((gather_env && gather_env[0] == '0') || !top_k || !ctx->device->pipeline_flash_attn_gather_f16 || q->ne[1] < 1 || q->ne[1] >= 64 || // 1..63: >=64 goes to the sparse prefill path - q->type != GGML_TYPE_F32 || k->type != GGML_TYPE_F16 || v->type != GGML_TYPE_F16 || + q->type != GGML_TYPE_F32 || !kv_word_addressable || !mask || mask->type != GGML_TYPE_F16 || top_k->type != GGML_TYPE_I32 || q->ne[0] != 512 || k->ne[0] != 512 || v->ne[0] != 512 || q->ne[2] != 64 || k->ne[2] != 1 || v->ne[2] != 1 || @@ -11824,6 +11840,9 @@ static bool ggml_vk_flash_attn_gather_compact(ggml_backend_vk_context * ctx, vk_ return false; } + const uint32_t k_row_bytes = (uint32_t) ggml_row_size(k->type, k->ne[0]); + const uint32_t k_row_words = k_row_bytes / 4; + float max_bias = 0.0f; float logit_softcap = 0.0f; memcpy(&max_bias, (const float *) dst->op_params + 1, sizeof(float)); @@ -11934,7 +11953,7 @@ static bool ggml_vk_flash_attn_gather_compact(ggml_backend_vk_context * ctx, vk_ goto union_unavailable; } - const size_t ukc_sz = (size_t) kv_c * 512 * sizeof(ggml_fp16_t); + const size_t ukc_sz = (size_t) kv_c * k_row_bytes; const size_t umc_sz = (size_t) n_batch * kv_c * sizeof(ggml_fp16_t); const size_t ul_sz = (size_t) max_union * sizeof(uint32_t); const size_t need = ukc_sz + umc_sz + ul_sz; @@ -11967,9 +11986,9 @@ static bool ggml_vk_flash_attn_gather_compact(ggml_backend_vk_context * ctx, vk_ const vk_op_flash_attn_gather_union_push_constants gpc = { (uint32_t) k->ne[1], (uint32_t) n_kv_raw, kv_c, - (uint32_t) (k->nb[1] / sizeof(ggml_fp16_t)), + (uint32_t) (k->nb[1] / 4), (uint32_t) (mask->nb[1] / sizeof(ggml_fp16_t)), - n_batch, + n_batch, k_row_words, }; ggml_vk_dispatch_pipeline(ctx, subctx, ctx->device->pipeline_flash_attn_gather_union_f16, { ggml_vk_tensor_subbuffer(ctx, k), ul_buf, ggml_vk_tensor_subbuffer(ctx, mask), @@ -11982,6 +12001,8 @@ static bool ggml_vk_flash_attn_gather_compact(ggml_backend_vk_context * ctx, vk_ st.dynamic_kv = true; st.kv_c = kv_c; st.n_batch = n_batch; + st.row_bytes = k_row_bytes; + st.row_elems = (uint32_t) (k->ne[0] / ggml_blck_size(k->type)); st.kc_buf = kc_buf; st.mc_buf = mc_buf; st.kv_buf = uc_buf; @@ -11997,7 +12018,7 @@ union_unavailable:; } const uint32_t ns = (uint32_t) q->ne[3]; - const size_t kc_sz = (size_t) ns * kv_c * 512 * sizeof(ggml_fp16_t); + const size_t kc_sz = (size_t) ns * kv_c * k_row_bytes; const size_t mc_sz = (size_t) ns * n_batch * kv_c * sizeof(ggml_fp16_t); if (ctx->prealloc_size_y < kc_sz + mc_sz) { @@ -12013,14 +12034,14 @@ union_unavailable:; const vk_op_flash_attn_gather_push_constants pc = { (uint32_t) k->ne[1], (uint32_t) n_kv_raw, (uint32_t) top_k->ne[0], kv_c, - (uint32_t) (k->nb[1] / sizeof(ggml_fp16_t)), - (uint32_t) (k->nb[3] / sizeof(ggml_fp16_t)), + (uint32_t) (k->nb[1] / 4), + (uint32_t) (k->nb[3] / 4), (uint32_t) (top_k->nb[1] / sizeof(int32_t)), (uint32_t) (top_k->nb[3] / sizeof(int32_t)), (uint32_t) (mask->nb[1] / sizeof(ggml_fp16_t)), (uint32_t) (mask->nb[3] / sizeof(ggml_fp16_t)), (uint32_t) mask->ne[3], - n_batch, + n_batch, k_row_words, }; st.kc_buf = ggml_vk_subbuffer(ctx, ctx->prealloc_y, 0); @@ -12082,6 +12103,8 @@ union_unavailable:; st.active = true; st.kv_c = kv_c; st.n_batch = n_batch; + st.row_bytes = k_row_bytes; + st.row_elems = (uint32_t) (k->ne[0] / ggml_blck_size(k->type)); return true; } @@ -12241,8 +12264,10 @@ static void ggml_vk_flash_attn(ggml_backend_vk_context * ctx, vk_context& subctx uint32_t k_stride = (uint32_t)(nbk1 / ggml_type_size(k->type)); uint32_t v_stride = (uint32_t)(nbv1 / ggml_type_size(v->type)); if (fa_compact.active) { - k_stride = 512; - v_stride = 512; + // rows are tightly packed in the compact scratch; for a quantised K this is the block + // count per row, which is what nbk1 / ggml_type_size would have given for the source + k_stride = fa_compact.row_elems; + v_stride = fa_compact.row_elems; } // For F32, the shader treats it as a block of size 4 (for vec4 loads) @@ -12455,9 +12480,9 @@ static void ggml_vk_flash_attn(ggml_backend_vk_context * ctx, vk_context& subctx ggml_vk_sync_buffers(ctx, subctx); } - // compact scratch layout: [512, kv_c, 1, ns] f16, tightly packed - const uint32_t eff_nbk2 = fa_compact.active ? fa_compact.kv_c * 512 * (uint32_t)sizeof(ggml_fp16_t) : nbk2_eff; - const uint32_t eff_nbk3 = fa_compact.active ? fa_compact.kv_c * 512 * (uint32_t)sizeof(ggml_fp16_t) : nbk3_eff; + // compact scratch layout: [512, kv_c, 1, ns] tightly packed, in K's own type + const uint32_t eff_nbk2 = fa_compact.active ? fa_compact.kv_c * fa_compact.row_bytes : nbk2_eff; + const uint32_t eff_nbk3 = fa_compact.active ? fa_compact.kv_c * fa_compact.row_bytes : nbk3_eff; const uint32_t eff_nbv2 = fa_compact.active ? eff_nbk2 : nbv2_eff; const uint32_t eff_nbv3 = fa_compact.active ? eff_nbk3 : nbv3_eff; diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_gather.comp b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_gather.comp index 3c9ac68e2c6..6dd24d31b51 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_gather.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_gather.comp @@ -19,10 +19,12 @@ layout(local_size_x = 64, local_size_y = 1, local_size_z = 1) in; -layout(binding = 0) readonly buffer KBuf { float16_t data_k[]; }; +// K is addressed as raw 4-byte WORDS, never as values: a gather relocates rows and does not +// need to know whether they hold f16 elements or quantised blocks. Only the MASK is typed. +layout(binding = 0) readonly buffer KBuf { uint data_k[]; }; layout(binding = 1) readonly buffer TopBuf { int data_top[]; }; layout(binding = 2) readonly buffer MBuf { float16_t data_m[]; }; -layout(binding = 3) writeonly buffer KcBuf { float16_t data_kc[]; }; +layout(binding = 3) writeonly buffer KcBuf { uint data_kc[]; }; layout(binding = 4) writeonly buffer McBuf { float16_t data_mc[]; }; layout(push_constant) uniform Parameters { @@ -30,17 +32,17 @@ layout(push_constant) uniform Parameters { uint n_kv_raw; // dense prefix length uint n_top_k; // selected rows for the (single) query token uint kv_c; // padded compact row count == dispatch row range - uint nbk1; // K source row stride, elements - uint nbk3; // K source stream stride, elements + uint nbk1; // K source row stride, WORDS + uint nbk3; // K source stream stride, WORDS uint nbt1; // top_k row (per query token) stride, elements uint nbt3; // top_k stream stride, elements uint nbm1; // mask source row (per query token) stride, elements uint nbm3; // mask source stream stride, elements uint nem3; // mask ne[3], for stream broadcast uint n_batch; // query tokens sharing this gather; <= LANES + uint row_words; // bytes per K row / 4 } p; -const uint HEAD_SIZE = 512; const uint LANES = 64; void main() { @@ -68,15 +70,18 @@ void main() { } } - const uint dst_base = (stream * p.kv_c + row) * HEAD_SIZE; + // Zeroing writes zero BYTES, which decode to zero for every block-quantised type here (a + // zero scale zeroes the block) as well as for f16. The row is -inf in the mask either way; + // zeroing only keeps a garbage dot product from reaching the softmax as a NaN. + const uint dst_base = (stream * p.kv_c + row) * p.row_words; if (src < p.n_kv) { const uint src_base = stream * p.nbk3 + src * p.nbk1; - [[unroll]] for (uint i = 0; i < HEAD_SIZE / LANES; ++i) { - data_kc[dst_base + tid + i * LANES] = data_k[src_base + tid + i * LANES]; + for (uint i = tid; i < p.row_words; i += LANES) { + data_kc[dst_base + i] = data_k[src_base + i]; } } else { - [[unroll]] for (uint i = 0; i < HEAD_SIZE / LANES; ++i) { - data_kc[dst_base + tid + i * LANES] = float16_t(0.0); + for (uint i = tid; i < p.row_words; i += LANES) { + data_kc[dst_base + i] = 0u; } } diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_gather_union.comp b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_gather_union.comp index 775b1bc96c6..4669e0d1c75 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_gather_union.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_gather_union.comp @@ -18,10 +18,14 @@ layout(local_size_x = 64, local_size_y = 1, local_size_z = 1) in; -layout(binding = 0) readonly buffer KBuf { float16_t data_k[]; }; +// K is addressed as raw 4-byte WORDS, never as values. A gather relocates rows; it does not +// need to know whether they hold f16 elements or quantised blocks, so the same shader serves +// every K type whose row is a whole number of words. Only the MASK is typed, and a mask is +// always f16. +layout(binding = 0) readonly buffer KBuf { uint data_k[]; }; layout(binding = 1) readonly buffer UBuf { uint data_u[]; }; layout(binding = 2) readonly buffer MBuf { float16_t data_m[]; }; -layout(binding = 3) writeonly buffer KcBuf { float16_t data_kc[]; }; +layout(binding = 3) writeonly buffer KcBuf { uint data_kc[]; }; layout(binding = 4) writeonly buffer McBuf { float16_t data_mc[]; }; layout(binding = 5) readonly buffer CBuf { uint data_c[]; }; @@ -29,12 +33,12 @@ layout(push_constant) uniform Parameters { uint n_kv; uint n_kv_raw; uint kv_c_max; - uint nbk1; + uint nbk1; // K source row stride, WORDS uint nbm1; uint n_batch; + uint row_words; // bytes per K row / 4 } p; -const uint HEAD_SIZE = 512; const uint LANES = 64; void main() { @@ -55,15 +59,18 @@ void main() { src = p.n_kv_raw + data_u[row - p.n_kv_raw]; } - const uint dst_base = row * HEAD_SIZE; + // Zeroing an unused row writes zero BYTES, which decode to zero for every block-quantised + // type here (a zero scale zeroes the block) as well as for f16. The row is -inf in the mask + // either way; zeroing only keeps a garbage dot product from reaching the softmax as a NaN. + const uint dst_base = row * p.row_words; if (src < p.n_kv) { const uint src_base = src * p.nbk1; - [[unroll]] for (uint i = 0; i < HEAD_SIZE / LANES; ++i) { - data_kc[dst_base + tid + i * LANES] = data_k[src_base + tid + i * LANES]; + for (uint i = tid; i < p.row_words; i += LANES) { + data_kc[dst_base + i] = data_k[src_base + i]; } } else { - [[unroll]] for (uint i = 0; i < HEAD_SIZE / LANES; ++i) { - data_kc[dst_base + tid + i * LANES] = float16_t(0.0); + for (uint i = tid; i < p.row_words; i += LANES) { + data_kc[dst_base + i] = 0u; } } diff --git a/tests/test-backend-ops.cpp b/tests/test-backend-ops.cpp index 6ab027c7ea1..749369be2ab 100644 --- a/tests/test-backend-ops.cpp +++ b/tests/test-backend-ops.cpp @@ -8040,12 +8040,13 @@ struct test_flash_attn_ext_top_k : public test_case { const bool sinks; const int64_t ns; // sequences (ne3); >1 exercises the split-K stream stride const int64_t ov; // % of each token's picks shared with its neighbours (dedup-union realism) + const ggml_type type_K; // K/V cache type; V is the same tensor, so one type covers both static constexpr int64_t hs = 512; // V4 CSA head size, K == V latent static constexpr int64_t nh = 64; // V4 CSA query heads (MQA) std::string vars() override { - return VARS_TO_STR7(kv, nb, n_kv_raw, n_top_k, sinks, ns, ov); + return VARS_TO_STR8(kv, nb, n_kv_raw, n_top_k, sinks, ns, ov, type_K); } double max_nmse_err() override { @@ -8059,14 +8060,15 @@ struct test_flash_attn_ext_top_k : public test_case { return 2 * nh * nb * ns * (hs + hs) * (n_kv_raw + n_top_k); } - test_flash_attn_ext_top_k(int64_t kv = 768, int64_t nb = 8, int64_t n_kv_raw = 64, int64_t n_top_k = 128, bool sinks = false, int64_t ns = 1, int64_t ov = 0) - : kv(kv), nb(nb), n_kv_raw(n_kv_raw), n_top_k(n_top_k), sinks(sinks), ns(ns), ov(ov) {} + test_flash_attn_ext_top_k(int64_t kv = 768, int64_t nb = 8, int64_t n_kv_raw = 64, int64_t n_top_k = 128, bool sinks = false, int64_t ns = 1, int64_t ov = 0, + ggml_type type_K = GGML_TYPE_F16) + : kv(kv), nb(nb), n_kv_raw(n_kv_raw), n_top_k(n_top_k), sinks(sinks), ns(ns), ov(ov), type_K(type_K) {} ggml_tensor * build_graph(ggml_context * ctx) override { ggml_tensor * q = ggml_new_tensor_4d(ctx, GGML_TYPE_F32, hs, nb, nh, ns); ggml_set_name(q, "q"); - ggml_tensor * k = ggml_new_tensor_4d(ctx, GGML_TYPE_F16, hs, kv, 1, ns); + ggml_tensor * k = ggml_new_tensor_4d(ctx, type_K, hs, kv, 1, ns); ggml_set_name(k, "k"); // V4 CSA attends over the K latent itself: V is the same cache tensor @@ -11426,6 +11428,15 @@ static std::vector> make_test_cases_eval() { test_cases.emplace_back(new test_flash_attn_ext_top_k(11008, 8, 2304, 512, false, 1, 60)); test_cases.emplace_back(new test_flash_attn_ext_top_k(11008, 16, 2304, 512, false, 1, 60)); test_cases.emplace_back(new test_flash_attn_ext_top_k(11008, 16, 2304, 512, false, 1, 86)); + // quantised K/V: the gather relocates rows verbatim, so it should serve any type whose row + // is a whole number of 4-byte words. These are the shapes a DSv4 decode with -ctk q8_0 hits, + // which took the dense fallback entirely before the gather learned to address rows as bytes. + for (ggml_type tk : { GGML_TYPE_Q8_0, GGML_TYPE_Q4_0 }) { + test_cases.emplace_back(new test_flash_attn_ext_top_k(8192, 4, 1024, 512, false, 1, 0, tk)); + test_cases.emplace_back(new test_flash_attn_ext_top_k(11008, 8, 2304, 512, false, 1, 60, tk)); + test_cases.emplace_back(new test_flash_attn_ext_top_k(11008, 16, 2304, 512, false, 1, 86, tk)); + test_cases.emplace_back(new test_flash_attn_ext_top_k(8192, 1, 1024, 512, false, 1, 0, tk)); + } test_cases.emplace_back(new test_flash_attn_ext_top_k(8192, 4, 1024, 512, false, 2)); test_cases.emplace_back(new test_flash_attn_ext_top_k( 768, 64, 64, 128, false, 2)); test_cases.emplace_back(new test_flash_attn_ext_top_k( 768, 64, 64, 128, true, 2)); @@ -11970,6 +11981,11 @@ static std::vector> make_test_cases_perf() { for (int nb : { 8, 16 }) { test_cases.emplace_back(new test_flash_attn_ext_top_k(11008, nb, 2304, 512, false, 1, 86)); } + // q8_0 K/V at the same shapes: DSv4 with -ctk q8_0 took the dense fallback before the + // gather became type-agnostic, so this is the cell that says whether it now pays there. + for (int nb : { 2, 4, 8, 16 }) { + test_cases.emplace_back(new test_flash_attn_ext_top_k(11008, nb, 2304, 512, false, 1, 60, GGML_TYPE_Q8_0)); + } return test_cases; } From 35b75782ffe19c8c01eea1755299c664da5b37e2 Mon Sep 17 00:00:00 2001 From: Nathan Wilson Date: Fri, 14 Aug 2026 06:52:40 +0000 Subject: [PATCH 34/68] vulkan: dequantise q8_0 K/V inside the DeepSeek V4 small-batch gather Relocating quantised rows verbatim leaves flash attention to decode them in its inner loop, which it does once per query block that reads a row rather than once per row. Measured, that costs 0.15 us per KV row attended - 0.1481 / 0.1510 / 0.1512 us/row at batch 8 / 4 / 2 and 0.1565 dense, so a constant across a 3.6x span of row counts and across both regimes. It is what made q8_0 attention 1.78x slower than f16 while reading half the bytes. The gather is the natural place to pay it instead: it already touches every selected row exactly once. Decoding there makes the compact scratch f16, flash attention takes its f16 path, and the redundancy disappears. No extra pass - the pass already existed. test-backend-ops perf, kv=11008 n_kv_raw=2304 n_top_k=512 ov=60, q8_0 K/V: batch dense verbatim gather decoded gather f16 gather 2 3917 us 1110 us 630 us 640 us 4 3924 us 1257 us 715 us 727 us 8 3936 us 1567 us 891 us 908 us q8_0 now edges f16 by 1.5 to 1.8% rather than trailing it by 78%: the FA does identical f16 work either way, and the gather's scattered read side moves half the bytes. Against the dense fallback a q8_0 cache took before any of this work, batch 8 is 4.41x. Flash attention picks its pipeline from what the SCRATCH holds, not from the source tensor, so k_type_eff/v_type_eff now also key off the compact state - the same mechanism use_dequant_kv already used, which is why that variable existed in this form. q8_0 only, deliberately. It is the type both community test sets used and the one worth having; every other quantised type keeps the verbatim path, which is correct for them and validated. Generalising means one shader variant per type through dequant_funcs.glsl, which is only worth the shader-permutation cost if this shape of win reproduces elsewhere. Not covered: the per-token gather still relocates verbatim, so a batch where the union declines but the per-token form fits keeps paying the inline decode. Same fix applies, and the same measurement will say whether it is worth a second shader. 13318 FLASH_ATTN_EXT cases pass. Co-Authored-By: Claude Opus 5 Assisted-by: Claude (Opus 5) --- ggml/src/ggml-vulkan/ggml-vulkan.cpp | 33 +++++-- .../flash_attn_gather_union_dq.comp | 99 +++++++++++++++++++ .../vulkan-shaders/vulkan-shaders-gen.cpp | 1 + 3 files changed, 124 insertions(+), 9 deletions(-) create mode 100644 ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_gather_union_dq.comp diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp index de2cb074542..b599012284e 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp +++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp @@ -1114,6 +1114,7 @@ struct vk_device_struct { vk_pipeline pipeline_flash_attn_gather_f16; vk_pipeline pipeline_flash_attn_union_f16; vk_pipeline pipeline_flash_attn_gather_union_f16; + vk_pipeline pipeline_flash_attn_gather_union_dq_q8_0; vk_pipeline pipeline_dsv4_hc_pre_f32; vk_pipeline pipeline_dsv4_hc_comb_f32; vk_pipeline pipeline_dsv4_hc_post_f32; @@ -6345,6 +6346,10 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { "flash_attn_gather_union_f16", flash_attn_gather_union_f16_len, flash_attn_gather_union_f16_data, "main", 6, sizeof(vk_op_flash_attn_gather_union_push_constants), {1, 1, 1}, {}, 1, true, true, device->subgroup_size); + ggml_vk_create_pipeline(device, device->pipeline_flash_attn_gather_union_dq_q8_0, + "flash_attn_gather_union_dq_q8_0", flash_attn_gather_union_dq_q8_0_len, flash_attn_gather_union_dq_q8_0_data, "main", 6, + sizeof(vk_op_flash_attn_gather_union_push_constants), {1, 1, 1}, {}, 1, true, true, + device->subgroup_size); } } @@ -11725,6 +11730,7 @@ struct vk_fa_compact_state { uint32_t n_batch = 1; uint32_t row_bytes = 0; // bytes per compact K row; K may be quantised uint32_t row_elems = 0; // K row stride in ELEMENTS/blocks, for the FA push constant + bool dequantized = false; // scratch holds f16 because the gather decoded on the way in vk_subbuffer kv_buf; vk_subbuffer kc_buf, mc_buf; }; @@ -11953,7 +11959,12 @@ static bool ggml_vk_flash_attn_gather_compact(ggml_backend_vk_context * ctx, vk_ goto union_unavailable; } - const size_t ukc_sz = (size_t) kv_c * k_row_bytes; + // Decoding on the way in makes the scratch f16 and hands flash attention its f16 path, + // which is worth far more than the extra scratch bytes: the inline decode it replaces + // costs a measured 0.15 us per KV row attended, every step. + const bool dq = k->type == GGML_TYPE_Q8_0 && ctx->device->pipeline_flash_attn_gather_union_dq_q8_0; + const uint32_t u_row_by = dq ? (uint32_t) (k->ne[0] * sizeof(ggml_fp16_t)) : k_row_bytes; + const size_t ukc_sz = (size_t) kv_c * u_row_by; const size_t umc_sz = (size_t) n_batch * kv_c * sizeof(ggml_fp16_t); const size_t ul_sz = (size_t) max_union * sizeof(uint32_t); const size_t need = ukc_sz + umc_sz + ul_sz; @@ -11973,8 +11984,10 @@ static bool ggml_vk_flash_attn_gather_compact(ggml_backend_vk_context * ctx, vk_ // prealloc_y sync above is a global barrier and orders the previous layer's read. const vk_subbuffer uc_buf = ggml_vk_subbuffer(ctx, ctx->fa_union_stat); + vk_pipeline gather_pipe = dq ? ctx->device->pipeline_flash_attn_gather_union_dq_q8_0 + : ctx->device->pipeline_flash_attn_gather_union_f16; ggml_pipeline_request_descriptor_sets(ctx, ctx->device->pipeline_flash_attn_union_f16, 1); - ggml_pipeline_request_descriptor_sets(ctx, ctx->device->pipeline_flash_attn_gather_union_f16, 1); + ggml_pipeline_request_descriptor_sets(ctx, gather_pipe, 1); const vk_op_flash_attn_union_push_constants upc = { (uint32_t) k->ne[1], (uint32_t) n_kv_raw, n_batch, (uint32_t) top_k->ne[0], max_union, @@ -11984,13 +11997,14 @@ static bool ggml_vk_flash_attn_gather_compact(ggml_backend_vk_context * ctx, vk_ { ggml_vk_tensor_subbuffer(ctx, top_k), ul_buf, uc_buf }, upc, { 1, 1, 1 }); ggml_vk_sync_buffers(ctx, subctx); + // the fused decoder steps K in BLOCKS and writes elements; the verbatim one does words const vk_op_flash_attn_gather_union_push_constants gpc = { (uint32_t) k->ne[1], (uint32_t) n_kv_raw, kv_c, - (uint32_t) (k->nb[1] / 4), + dq ? (uint32_t) (k->nb[1] / ggml_type_size(k->type)) : (uint32_t) (k->nb[1] / 4), (uint32_t) (mask->nb[1] / sizeof(ggml_fp16_t)), - n_batch, k_row_words, + n_batch, dq ? (uint32_t) k->ne[0] : k_row_words, }; - ggml_vk_dispatch_pipeline(ctx, subctx, ctx->device->pipeline_flash_attn_gather_union_f16, + ggml_vk_dispatch_pipeline(ctx, subctx, gather_pipe, { ggml_vk_tensor_subbuffer(ctx, k), ul_buf, ggml_vk_tensor_subbuffer(ctx, mask), kc_buf, mc_buf, uc_buf }, gpc, { kv_c, 1, 1 }); ggml_vk_sync_buffers(ctx, subctx); @@ -12001,8 +12015,9 @@ static bool ggml_vk_flash_attn_gather_compact(ggml_backend_vk_context * ctx, vk_ st.dynamic_kv = true; st.kv_c = kv_c; st.n_batch = n_batch; - st.row_bytes = k_row_bytes; - st.row_elems = (uint32_t) (k->ne[0] / ggml_blck_size(k->type)); + st.row_bytes = u_row_by; + st.row_elems = dq ? (uint32_t) k->ne[0] : (uint32_t) (k->ne[0] / ggml_blck_size(k->type)); + st.dequantized = dq; st.kc_buf = kc_buf; st.mc_buf = mc_buf; st.kv_buf = uc_buf; @@ -12240,8 +12255,8 @@ static void ggml_vk_flash_attn(ggml_backend_vk_context * ctx, vk_context& subctx // If this fires, supports_op admitted a non-native K/V type the gate then rejected; the // native shader would return garbage rather than fail, so abort instead. GGML_ASSERT(use_dequant_kv || !kv_needs_dequant); - const ggml_type k_type_eff = use_dequant_kv ? GGML_TYPE_F16 : k->type; - const ggml_type v_type_eff = use_dequant_kv ? GGML_TYPE_F16 : v->type; + const ggml_type k_type_eff = (use_dequant_kv || fa_compact.dequantized) ? GGML_TYPE_F16 : k->type; + const ggml_type v_type_eff = (use_dequant_kv || fa_compact.dequantized) ? GGML_TYPE_F16 : v->type; // For scalar/coopmat1 FA, we can use the "large" size to accommodate qga. // For coopmat2 FA, we always use the small size (which is still pretty large for gqa). diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_gather_union_dq.comp b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_gather_union_dq.comp new file mode 100644 index 00000000000..fc2a0f4541a --- /dev/null +++ b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_gather_union_dq.comp @@ -0,0 +1,99 @@ +#version 450 + +#extension GL_EXT_control_flow_attributes : require +#extension GL_EXT_shader_16bit_storage : require +#extension GL_EXT_shader_explicit_arithmetic_types_float16 : require +#extension GL_EXT_shader_explicit_arithmetic_types_int8 : require +#extension GL_EXT_shader_explicit_arithmetic_types_int16 : require +#extension GL_EXT_shader_explicit_arithmetic_types_int32 : require + +#include "types.glsl" + +// Gather + DEQUANTISE in one pass, for the DeepSeek V4 small-batch decode path with a +// quantised KV cache. +// +// The sibling flash_attn_gather_union.comp relocates rows verbatim, which leaves flash +// attention to dequantise them in its inner loop. That costs a MEASURED 0.15 us per KV row +// attended - constant across batch 2/4/8 and across the dense and gathered regimes alike, and +// large enough to make q8_0 attention 1.78x slower than f16 despite reading half the bytes. +// It is redundant work: the FA re-decodes the same row for every query block that reads it, +// whereas the gather already touches each selected row exactly once. +// +// So dequantise here instead. The compact scratch becomes f16, flash attention takes its f16 +// path, and the decode is paid once per row rather than once per use. The extra cost is +// writing f16 instead of q8_0 into the scratch; the pass itself already existed. +// +// q8_0 only. Other quantised types keep the verbatim path, which is correct for them and was +// validated; generalising means one variant per type through dequant_funcs.glsl, which is only +// worth doing if this shape of win holds up. + +layout(local_size_x = 64, local_size_y = 1, local_size_z = 1) in; + +layout(binding = 0) readonly buffer KBuf { block_q8_0_packed16 data_k[]; }; +layout(binding = 1) readonly buffer UBuf { uint data_u[]; }; +layout(binding = 2) readonly buffer MBuf { float16_t data_m[]; }; +layout(binding = 3) writeonly buffer KcBuf { float16_t data_kc[]; }; +layout(binding = 4) writeonly buffer McBuf { float16_t data_mc[]; }; +layout(binding = 5) readonly buffer CBuf { uint data_c[]; }; + +layout(push_constant) uniform Parameters { + uint n_kv; + uint n_kv_raw; + uint kv_c_max; + uint nbk1; // K source row stride, BLOCKS + uint nbm1; + uint n_batch; + uint row_elems; // elements per K row (== head size) +} p; + +const uint LANES = 64; +const uint QK8_0 = 32; + +void main() { + const uint row = gl_WorkGroupID.x; + const uint tid = gl_LocalInvocationIndex; + + const uint kv_c = data_c[0]; // padded compact rows, the FA's runtime KV + const uint n_uni = data_c[1]; // unpadded union size + + if (row >= kv_c) { + return; + } + + uint src = p.n_kv; // sentinel: invalid + if (row < p.n_kv_raw) { + src = row; + } else if (row - p.n_kv_raw < n_uni) { + src = p.n_kv_raw + data_u[row - p.n_kv_raw]; + } + + // 64 lanes x 8 elements covers a 512-element row exactly; a q8_0 block is 32 elements, so + // each lane stays inside one block and reads four int16 pairs from it. + const uint dst_base = row * p.row_elems; + const uint e0 = tid * 8; + if (src < p.n_kv) { + const uint ib = src * p.nbk1 + e0 / QK8_0; + const uint iqs = e0 % QK8_0; + const float d = float(data_k[ib].d); + [[unroll]] for (uint j = 0; j < 4; ++j) { + const i8vec2 v = unpack8(int32_t(data_k[ib].qs[iqs / 2 + j])).xy; + data_kc[dst_base + e0 + 2 * j] = float16_t(d * float(v.x)); + data_kc[dst_base + e0 + 2 * j + 1] = float16_t(d * float(v.y)); + } + } else { + [[unroll]] for (uint j = 0; j < 8; ++j) { + data_kc[dst_base + e0 + j] = float16_t(0.0); + } + } + + // Compact mask is token-major [n_batch][kv_c], stride the RUNTIME kv_c: the FA derives + // m_row_len from KV, which is that same runtime value. + const float NEG_INF = uintBitsToFloat(0xff800000); + if (tid < p.n_batch) { + float mv = NEG_INF; + if (src < p.n_kv) { + mv = float(data_m[tid * p.nbm1 + src]); + } + data_mc[tid * kv_c + row] = float16_t(mv); + } +} 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 38e9e4e0af5..44c00f4acac 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp @@ -824,6 +824,7 @@ void process_shaders() { string_to_spv("flash_attn_gather_f16", "flash_attn_gather.comp", {}); string_to_spv("flash_attn_union_f16", "flash_attn_union.comp", {}); string_to_spv("flash_attn_gather_union_f16", "flash_attn_gather_union.comp", {}); + string_to_spv("flash_attn_gather_union_dq_q8_0", "flash_attn_gather_union_dq.comp", {}); string_to_spv("dsv4_hc_pre_f32", "dsv4_hc_pre.comp", {}); string_to_spv("dsv4_hc_comb_f32", "dsv4_hc_comb.comp", {}); string_to_spv("dsv4_hc_post_f32", "dsv4_hc_post.comp", {}); From 987f66222055ae2945597d416f3d38da2a03df5d Mon Sep 17 00:00:00 2001 From: Nathan Wilson Date: Fri, 14 Aug 2026 07:29:56 +0000 Subject: [PATCH 35/68] vulkan: decode q4_0 in the V4 gather too, with the real element mapping q4_0 turns out to want this more than q8_0 did. Its inline-decode penalty is the same shape - a constant per KV row attended, flat across batch and across dense vs gathered - but larger: 0.172 to 0.179 us/row against q8_0's 0.148 to 0.153. Unpacking nibbles costs more ALU than the halved bytes save, which is also why q4_0 was the SLOWEST of the three types before this. kv=11008 n_kv_raw=2304 n_top_k=512 ov=60, batch 2 / 4 / 8 type dense verbatim gather decoded gather f16 2192/2203/2213 - 640/726/902 q8_0 3917/3924/3936 1110/1257/1567 628/712/888 q4_0 4134/4141/4145 1189/1342/1671 626/711/886 All three converge, because after the gather they run the same f16 attention; what is left is the scattered-read side, where fewer bytes now wins. Against the dense fallback each type took before any of this work, batch 8 is 2.45x for f16, 4.43x for q8_0 and 4.68x for q4_0. The interesting part is what the first attempt got wrong. Generalising via dequantize4() from dequant_funcs.glsl looked obvious and passed on q8_0, then failed q4_0 with NMSE 1.17. That helper is permutation AGNOSTIC by design: mul_mat_vec only ever feeds it into a dot product, where any consistent element order gives the same sum, so for q4_0 it returns nibbles in packed order rather than element order. q8_0 passed only because its layout happens to be contiguous. Materialising a row into memory is precisely the use where the order matters, so each type's true mapping is now written out - for q4_0, byte j carries element j in its low nibble and element j+16 in its high one, exactly as dequant_q4_0.comp materialises it. That is also why this stops at two types rather than looping over the whole native list: every type needs its mapping stated and tested, an #error guards anything else reaching the shader, and q8_0 and q4_0 are the two actually used as a KV cache. q4_1/q5_0/q5_1/iq4_nl keep the verbatim gather, which is correct for them. 13318 FLASH_ATTN_EXT cases pass, including q4_0 and q8_0 top-k cases at the shapes a DSv4 decode hits, plus q4_0 perf cells alongside the q8_0 ones. Co-Authored-By: Claude Opus 5 Assisted-by: Claude (Opus 5) --- ggml/src/ggml-vulkan/ggml-vulkan.cpp | 17 +++-- .../flash_attn_gather_union_dq.comp | 70 +++++++++++++------ .../vulkan-shaders/vulkan-shaders-gen.cpp | 9 ++- tests/test-backend-ops.cpp | 6 +- 4 files changed, 70 insertions(+), 32 deletions(-) diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp index b599012284e..7de378a8f48 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp +++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp @@ -1114,7 +1114,7 @@ struct vk_device_struct { vk_pipeline pipeline_flash_attn_gather_f16; vk_pipeline pipeline_flash_attn_union_f16; vk_pipeline pipeline_flash_attn_gather_union_f16; - vk_pipeline pipeline_flash_attn_gather_union_dq_q8_0; + vk_pipeline pipeline_flash_attn_gather_union_dq[GGML_TYPE_COUNT]; vk_pipeline pipeline_dsv4_hc_pre_f32; vk_pipeline pipeline_dsv4_hc_comb_f32; vk_pipeline pipeline_dsv4_hc_post_f32; @@ -6346,10 +6346,15 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { "flash_attn_gather_union_f16", flash_attn_gather_union_f16_len, flash_attn_gather_union_f16_data, "main", 6, sizeof(vk_op_flash_attn_gather_union_push_constants), {1, 1, 1}, {}, 1, true, true, device->subgroup_size); - ggml_vk_create_pipeline(device, device->pipeline_flash_attn_gather_union_dq_q8_0, - "flash_attn_gather_union_dq_q8_0", flash_attn_gather_union_dq_q8_0_len, flash_attn_gather_union_dq_q8_0_data, "main", 6, - sizeof(vk_op_flash_attn_gather_union_push_constants), {1, 1, 1}, {}, 1, true, true, +#define CREATE_FA_GATHER_DQ(TYPE, NAMED) \ + ggml_vk_create_pipeline(device, device->pipeline_flash_attn_gather_union_dq[TYPE], \ + "flash_attn_gather_union_dq_" #NAMED, flash_attn_gather_union_dq_ ## NAMED ## _len, \ + flash_attn_gather_union_dq_ ## NAMED ## _data, "main", 6, \ + sizeof(vk_op_flash_attn_gather_union_push_constants), {1, 1, 1}, {}, 1, true, true, \ device->subgroup_size); + CREATE_FA_GATHER_DQ(GGML_TYPE_Q4_0, q4_0) + CREATE_FA_GATHER_DQ(GGML_TYPE_Q8_0, q8_0) +#undef CREATE_FA_GATHER_DQ } } @@ -11962,7 +11967,7 @@ static bool ggml_vk_flash_attn_gather_compact(ggml_backend_vk_context * ctx, vk_ // Decoding on the way in makes the scratch f16 and hands flash attention its f16 path, // which is worth far more than the extra scratch bytes: the inline decode it replaces // costs a measured 0.15 us per KV row attended, every step. - const bool dq = k->type == GGML_TYPE_Q8_0 && ctx->device->pipeline_flash_attn_gather_union_dq_q8_0; + const bool dq = ggml_is_quantized(k->type) && ctx->device->pipeline_flash_attn_gather_union_dq[k->type]; const uint32_t u_row_by = dq ? (uint32_t) (k->ne[0] * sizeof(ggml_fp16_t)) : k_row_bytes; const size_t ukc_sz = (size_t) kv_c * u_row_by; const size_t umc_sz = (size_t) n_batch * kv_c * sizeof(ggml_fp16_t); @@ -11984,7 +11989,7 @@ static bool ggml_vk_flash_attn_gather_compact(ggml_backend_vk_context * ctx, vk_ // prealloc_y sync above is a global barrier and orders the previous layer's read. const vk_subbuffer uc_buf = ggml_vk_subbuffer(ctx, ctx->fa_union_stat); - vk_pipeline gather_pipe = dq ? ctx->device->pipeline_flash_attn_gather_union_dq_q8_0 + vk_pipeline gather_pipe = dq ? ctx->device->pipeline_flash_attn_gather_union_dq[k->type] : ctx->device->pipeline_flash_attn_gather_union_f16; ggml_pipeline_request_descriptor_sets(ctx, ctx->device->pipeline_flash_attn_union_f16, 1); ggml_pipeline_request_descriptor_sets(ctx, gather_pipe, 1); diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_gather_union_dq.comp b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_gather_union_dq.comp index fc2a0f4541a..fd3d0708b2f 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_gather_union_dq.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_gather_union_dq.comp @@ -7,29 +7,32 @@ #extension GL_EXT_shader_explicit_arithmetic_types_int16 : require #extension GL_EXT_shader_explicit_arithmetic_types_int32 : require +// iq4_nl's value table lives in shared memory and is typed FLOAT_TYPE +#define FLOAT_TYPE float #include "types.glsl" // Gather + DEQUANTISE in one pass, for the DeepSeek V4 small-batch decode path with a -// quantised KV cache. +// quantised KV cache. One variant per K type, via DATA_A_*. // // The sibling flash_attn_gather_union.comp relocates rows verbatim, which leaves flash -// attention to dequantise them in its inner loop. That costs a MEASURED 0.15 us per KV row -// attended - constant across batch 2/4/8 and across the dense and gathered regimes alike, and -// large enough to make q8_0 attention 1.78x slower than f16 despite reading half the bytes. -// It is redundant work: the FA re-decodes the same row for every query block that reads it, -// whereas the gather already touches each selected row exactly once. +// attention to decode them in its inner loop - once per query block that reads a row, rather +// than once per row. Measured, that is a constant per KV row attended: // -// So dequantise here instead. The compact scratch becomes f16, flash attention takes its f16 -// path, and the decode is paid once per row rather than once per use. The extra cost is -// writing f16 instead of q8_0 into the scratch; the pass itself already existed. +// q8_0 0.148 - 0.153 us/row q4_0 0.172 - 0.179 us/row // -// q8_0 only. Other quantised types keep the verbatim path, which is correct for them and was -// validated; generalising means one variant per type through dequant_funcs.glsl, which is only -// worth doing if this shape of win holds up. +// flat across batch 2/4/8 and across the dense and gathered regimes alike. It is what made +// quantised attention slower than f16 despite reading fewer bytes, and q4_0 worse than q8_0 +// despite reading fewer bytes still: unpacking nibbles costs more ALU than the bytes save. +// +// The gather already touches each selected row exactly once, so decoding here converts per-use +// work into per-row work. The compact scratch becomes f16 and flash attention takes its f16 +// path. The pass itself already existed, so the only added cost is writing f16 rather than +// blocks into the scratch. layout(local_size_x = 64, local_size_y = 1, local_size_z = 1) in; -layout(binding = 0) readonly buffer KBuf { block_q8_0_packed16 data_k[]; }; +layout (binding = 0) readonly buffer A {A_TYPE data_a[];}; + layout(binding = 1) readonly buffer UBuf { uint data_u[]; }; layout(binding = 2) readonly buffer MBuf { float16_t data_m[]; }; layout(binding = 3) writeonly buffer KcBuf { float16_t data_kc[]; }; @@ -47,9 +50,12 @@ layout(push_constant) uniform Parameters { } p; const uint LANES = 64; -const uint QK8_0 = 32; void main() { +#ifdef NEEDS_INIT_IQ_SHMEM + // barrier inside; must run before any divergent return below + init_iq_shmem(gl_WorkGroupSize); +#endif const uint row = gl_WorkGroupID.x; const uint tid = gl_LocalInvocationIndex; @@ -67,19 +73,37 @@ void main() { src = p.n_kv_raw + data_u[row - p.n_kv_raw]; } - // 64 lanes x 8 elements covers a 512-element row exactly; a q8_0 block is 32 elements, so - // each lane stays inside one block and reads four int16 pairs from it. + // 64 lanes x 8 elements covers a 512-element row exactly, and QUANT_K is 32, so a lane's 8 + // elements sit wholly inside one block AND wholly inside one nibble half. + // + // Deliberately NOT dequantize4() from dequant_funcs.glsl. That helper is permutation + // AGNOSTIC: mul_mat_vec only ever feeds it into a dot product, where any consistent element + // order gives the same answer, so for q4_0 it returns nibbles in packed order rather than + // element order. Materialising a row to memory is the one use where the order matters, and + // using it here produced NMSE 1.17 on q4_0 while q8_0 passed - q8_0's layout just happens to + // be contiguous. Each type's true element mapping is written out below instead. const uint dst_base = row * p.row_elems; const uint e0 = tid * 8; if (src < p.n_kv) { - const uint ib = src * p.nbk1 + e0 / QK8_0; - const uint iqs = e0 % QK8_0; - const float d = float(data_k[ib].d); - [[unroll]] for (uint j = 0; j < 4; ++j) { - const i8vec2 v = unpack8(int32_t(data_k[ib].qs[iqs / 2 + j])).xy; - data_kc[dst_base + e0 + 2 * j] = float16_t(d * float(v.x)); - data_kc[dst_base + e0 + 2 * j + 1] = float16_t(d * float(v.y)); + const uint ib = src * p.nbk1 + e0 / QUANT_K; + const uint iqs = e0 % QUANT_K; + const float d = float(data_a[ib].d); +#if defined(DATA_A_Q8_0) + // element e is qs[e] + [[unroll]] for (uint l = 0; l < 8; ++l) { + data_kc[dst_base + e0 + l] = float16_t(d * float(data_a[ib].qs[iqs + l])); + } +#elif defined(DATA_A_Q4_0) + // byte j carries element j in its low nibble and element j+16 in its high one + const uint shift = (iqs >> 4) * 4; // 0 for elements 0..15, 4 for 16..31 + const uint byte0 = iqs & 0xF; + [[unroll]] for (uint l = 0; l < 8; ++l) { + const float q = float((data_a[ib].qs[byte0 + l] >> shift) & 0xF) - 8.0f; + data_kc[dst_base + e0 + l] = float16_t(d * q); } +#else +#error "no element mapping written for this K type" +#endif } else { [[unroll]] for (uint j = 0; j < 8; ++j) { data_kc[dst_base + e0 + j] = float16_t(0.0); 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 44c00f4acac..efef6542a77 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp @@ -824,7 +824,14 @@ void process_shaders() { string_to_spv("flash_attn_gather_f16", "flash_attn_gather.comp", {}); string_to_spv("flash_attn_union_f16", "flash_attn_union.comp", {}); string_to_spv("flash_attn_gather_union_f16", "flash_attn_gather_union.comp", {}); - string_to_spv("flash_attn_gather_union_dq_q8_0", "flash_attn_gather_union_dq.comp", {}); + // one decoder per quantised K type flash-attention supports natively; f16/bf16/f32 need no + // decode and take the verbatim gather + // q8_0 and q4_0 only: each needs its true element mapping written out (see the shader), and + // these are the two types actually used as a KV cache. The rest take the verbatim gather. + for (const auto& tname : {"q4_0", "q8_0"}) { + string_to_spv("flash_attn_gather_union_dq_" + std::string(tname), "flash_attn_gather_union_dq.comp", + {{"DATA_A_" + to_uppercase(tname), "1"}}); + } string_to_spv("dsv4_hc_pre_f32", "dsv4_hc_pre.comp", {}); string_to_spv("dsv4_hc_comb_f32", "dsv4_hc_comb.comp", {}); string_to_spv("dsv4_hc_post_f32", "dsv4_hc_post.comp", {}); diff --git a/tests/test-backend-ops.cpp b/tests/test-backend-ops.cpp index 749369be2ab..c3cb4ef8e99 100644 --- a/tests/test-backend-ops.cpp +++ b/tests/test-backend-ops.cpp @@ -11983,8 +11983,10 @@ static std::vector> make_test_cases_perf() { } // q8_0 K/V at the same shapes: DSv4 with -ctk q8_0 took the dense fallback before the // gather became type-agnostic, so this is the cell that says whether it now pays there. - for (int nb : { 2, 4, 8, 16 }) { - test_cases.emplace_back(new test_flash_attn_ext_top_k(11008, nb, 2304, 512, false, 1, 60, GGML_TYPE_Q8_0)); + for (ggml_type tk : { GGML_TYPE_Q8_0, GGML_TYPE_Q4_0 }) { + for (int nb : { 2, 4, 8, 16 }) { + test_cases.emplace_back(new test_flash_attn_ext_top_k(11008, nb, 2304, 512, false, 1, 60, tk)); + } } return test_cases; From fa5f2517f19297d5ae4c32d03b52faa26db81891 Mon Sep 17 00:00:00 2001 From: Nathan Wilson Date: Mon, 17 Aug 2026 09:31:21 +0000 Subject: [PATCH 36/68] vulkan : restore the unrolled row copy in the DSV4 gather shaders Generalising the gather from f16 elements to raw words (6b2cade31) replaced [[unroll]] for (uint i = 0; i < HEAD_SIZE / LANES; ++i) with a loop bounded by the row_words push constant, which cannot be unrolled and pays a bounds check per iteration. Half as many, twice as wide accesses were still 6% slower on the smallest op the gather serves. Bisected over the 19 beta3 commits with the tip's test file overlaid at every point, so the instrument stayed fixed while beta3 grew its own perf cases. nb=1 at kv=8192: 54.58 us at the pre-beta3 base, flat through commit 16 (54.84), 58.24 at 6b2cade31, 58.23 at the tip. Step 4*LANES instead. Each access stays contiguous across the lanes, so coalescing is unchanged, and an f16 row at head size 512 is 256 words == 4*LANES: one iteration, no loop overhead. The tail loop carries the quantised row sizes, which are not multiples of 4*LANES (q8_0 is 136 words, q4_0 is 72). nb=1 against the pre-beta3 base, f16 K/V: -0.07% at kv=8192, -2.83% at kv=35584, mean -1.29% across the six shapes, so the regression is gone rather than reduced. Against beta3 as merged: -4.8% to -6.2%. The nb>=2 gains are untouched, worst cell -0.51%. 13318 FLASH_ATTN_EXT cases pass. Worth noting the gather itself was never the cost at nb=1: with GGML_VK_FA_TOPK_GATHER=0 that cell is 249.84 us against 58.14 us with it on. Assisted-by: Claude Opus 5 --- .../vulkan-shaders/flash_attn_gather.comp | 22 +++++++++++++++++-- .../flash_attn_gather_union.comp | 20 +++++++++++++++-- 2 files changed, 38 insertions(+), 4 deletions(-) diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_gather.comp b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_gather.comp index 6dd24d31b51..4e16b2a9cec 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_gather.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_gather.comp @@ -74,13 +74,31 @@ void main() { // zero scale zeroes the block) as well as for f16. The row is -inf in the mask either way; // zeroing only keeps a garbage dot product from reaching the softmax as a NaN. const uint dst_base = (stream * p.kv_c + row) * p.row_words; + // row_words is a push constant, so this cannot be [[unroll]]ed the way it was when the row + // was a fixed 512 f16 elements. Stepping 4*LANES recovers that: each access stays contiguous + // across the lanes, and an f16 row (256 words == 4*LANES) is one iteration with no loop + // overhead. The tail carries the quantised row sizes, which are not multiples of 4*LANES. if (src < p.n_kv) { const uint src_base = stream * p.nbk3 + src * p.nbk1; - for (uint i = tid; i < p.row_words; i += LANES) { + uint i = tid; + for (; i + 3 * LANES < p.row_words; i += 4 * LANES) { + data_kc[dst_base + i] = data_k[src_base + i]; + data_kc[dst_base + i + LANES] = data_k[src_base + i + LANES]; + data_kc[dst_base + i + 2 * LANES] = data_k[src_base + i + 2 * LANES]; + data_kc[dst_base + i + 3 * LANES] = data_k[src_base + i + 3 * LANES]; + } + for (; i < p.row_words; i += LANES) { data_kc[dst_base + i] = data_k[src_base + i]; } } else { - for (uint i = tid; i < p.row_words; i += LANES) { + uint i = tid; + for (; i + 3 * LANES < p.row_words; i += 4 * LANES) { + data_kc[dst_base + i] = 0u; + data_kc[dst_base + i + LANES] = 0u; + data_kc[dst_base + i + 2 * LANES] = 0u; + data_kc[dst_base + i + 3 * LANES] = 0u; + } + for (; i < p.row_words; i += LANES) { data_kc[dst_base + i] = 0u; } } diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_gather_union.comp b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_gather_union.comp index 4669e0d1c75..306372ba8af 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_gather_union.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_gather_union.comp @@ -63,13 +63,29 @@ void main() { // type here (a zero scale zeroes the block) as well as for f16. The row is -inf in the mask // either way; zeroing only keeps a garbage dot product from reaching the softmax as a NaN. const uint dst_base = row * p.row_words; + // Stepped 4*LANES for the same reason as flash_attn_gather.comp: row_words is a push + // constant, so the copy cannot be [[unroll]]ed, and an f16 row is one iteration this way. if (src < p.n_kv) { const uint src_base = src * p.nbk1; - for (uint i = tid; i < p.row_words; i += LANES) { + uint i = tid; + for (; i + 3 * LANES < p.row_words; i += 4 * LANES) { + data_kc[dst_base + i] = data_k[src_base + i]; + data_kc[dst_base + i + LANES] = data_k[src_base + i + LANES]; + data_kc[dst_base + i + 2 * LANES] = data_k[src_base + i + 2 * LANES]; + data_kc[dst_base + i + 3 * LANES] = data_k[src_base + i + 3 * LANES]; + } + for (; i < p.row_words; i += LANES) { data_kc[dst_base + i] = data_k[src_base + i]; } } else { - for (uint i = tid; i < p.row_words; i += LANES) { + uint i = tid; + for (; i + 3 * LANES < p.row_words; i += 4 * LANES) { + data_kc[dst_base + i] = 0u; + data_kc[dst_base + i + LANES] = 0u; + data_kc[dst_base + i + 2 * LANES] = 0u; + data_kc[dst_base + i + 3 * LANES] = 0u; + } + for (; i < p.row_words; i += LANES) { data_kc[dst_base + i] = 0u; } } From 7ceebb6b5261ce565d1bc192a59ab619a7c9c521 Mon Sep 17 00:00:00 2001 From: Nathan Wilson Date: Thu, 20 Aug 2026 00:38:13 +0000 Subject: [PATCH 37/68] vulkan: decode quantised K/V inside the DeepSeek V4 per-token gather The union gather decodes q8_0/q4_0 on the way in, but it exists only for n_batch > 1: a single token has nothing to deduplicate, so plain autoregressive decode takes the per-token gather, which relocated rows verbatim and left flash attention paying the same inline decode the union work measured at ~0.15 us per KV row attended. flash_attn_gather_dq.comp closes that: same row mapping as flash_attn_gather.comp, same per-type element mapping as flash_attn_gather_union_dq.comp (q4_0 materialised in element order, not packed order). The compact scratch becomes f16 and flash attention takes its f16 path, keyed off the compact state exactly as the union path already is. test-backend-ops perf, kv=11008 n_kv_raw=2304 n_top_k=512 ov=60, nb=1: type verbatim gather decoded gather q8_0 226.9 us 102.9 us q4_0 258.4 us 101.8 us f16 114.5 us 105.5 us (f16 never decodes; its two runs bound the harness spread) Both quantised types land on the f16 op time instead of trailing it 2.0-2.3x: after the gather the attention work is identical f16, and the gather's scattered read side moves fewer bytes. q8_0 and q4_0 only, as in the union: each type's element mapping must be stated and tested, an #error guards anything else reaching the shader, and every other type keeps the verbatim gather. nb=1 rows are added to the top-k perf grids so this regime stays measured. Co-Authored-By: Claude Fable 5 Assisted-by: Claude (Opus 5) --- ggml/src/ggml-vulkan/ggml-vulkan.cpp | 33 ++++-- .../vulkan-shaders/flash_attn_gather_dq.comp | 106 ++++++++++++++++++ .../vulkan-shaders/vulkan-shaders-gen.cpp | 2 + tests/test-backend-ops.cpp | 6 +- 4 files changed, 138 insertions(+), 9 deletions(-) create mode 100644 ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_gather_dq.comp diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp index 7de378a8f48..92e2c58282b 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp +++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp @@ -1112,6 +1112,7 @@ struct vk_device_struct { vk_pipeline pipeline_flash_attn_top_k_f16; vk_pipeline pipeline_flash_attn_top_k_cm_f16; vk_pipeline pipeline_flash_attn_gather_f16; + vk_pipeline pipeline_flash_attn_gather_dq[GGML_TYPE_COUNT]; vk_pipeline pipeline_flash_attn_union_f16; vk_pipeline pipeline_flash_attn_gather_union_f16; vk_pipeline pipeline_flash_attn_gather_union_dq[GGML_TYPE_COUNT]; @@ -6337,6 +6338,15 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { "flash_attn_gather_f16", flash_attn_gather_f16_len, flash_attn_gather_f16_data, "main", 5, sizeof(vk_op_flash_attn_gather_push_constants), {1, 1, 1}, {}, 1, true, true, device->subgroup_size); +#define CREATE_FA_GATHER_TOK_DQ(TYPE, NAMED) \ + ggml_vk_create_pipeline(device, device->pipeline_flash_attn_gather_dq[TYPE], \ + "flash_attn_gather_dq_" #NAMED, flash_attn_gather_dq_ ## NAMED ## _len, \ + flash_attn_gather_dq_ ## NAMED ## _data, "main", 5, \ + sizeof(vk_op_flash_attn_gather_push_constants), {1, 1, 1}, {}, 1, true, true, \ + device->subgroup_size); + CREATE_FA_GATHER_TOK_DQ(GGML_TYPE_Q4_0, q4_0) + CREATE_FA_GATHER_TOK_DQ(GGML_TYPE_Q8_0, q8_0) +#undef CREATE_FA_GATHER_TOK_DQ if (device->subgroup_arithmetic) { ggml_vk_create_pipeline(device, device->pipeline_flash_attn_union_f16, "flash_attn_union_f16", flash_attn_union_f16_len, flash_attn_union_f16_data, "main", 3, @@ -12037,8 +12047,14 @@ union_unavailable:; return false; } + // Decode on the way in, as the union path does: the gather touches each row once, while + // flash attention decodes once per query block that reads it. The union only covers + // n_batch > 1, so single-token decode lands here. + const bool tok_dq = ggml_is_quantized(k->type) && ctx->device->pipeline_flash_attn_gather_dq[k->type]; + const uint32_t row_by = tok_dq ? (uint32_t) (k->ne[0] * sizeof(ggml_fp16_t)) : k_row_bytes; + const uint32_t ns = (uint32_t) q->ne[3]; - const size_t kc_sz = (size_t) ns * kv_c * k_row_bytes; + const size_t kc_sz = (size_t) ns * kv_c * row_by; const size_t mc_sz = (size_t) ns * n_batch * kv_c * sizeof(ggml_fp16_t); if (ctx->prealloc_size_y < kc_sz + mc_sz) { @@ -12049,19 +12065,21 @@ union_unavailable:; ggml_vk_sync_buffers(ctx, subctx); } - vk_pipeline pipeline = ctx->device->pipeline_flash_attn_gather_f16; + vk_pipeline pipeline = tok_dq ? ctx->device->pipeline_flash_attn_gather_dq[k->type] + : ctx->device->pipeline_flash_attn_gather_f16; ggml_pipeline_request_descriptor_sets(ctx, pipeline, 1); + // the fused decoder steps K in BLOCKS and writes elements; the verbatim one does words const vk_op_flash_attn_gather_push_constants pc = { (uint32_t) k->ne[1], (uint32_t) n_kv_raw, (uint32_t) top_k->ne[0], kv_c, - (uint32_t) (k->nb[1] / 4), - (uint32_t) (k->nb[3] / 4), + tok_dq ? (uint32_t) (k->nb[1] / ggml_type_size(k->type)) : (uint32_t) (k->nb[1] / 4), + tok_dq ? (uint32_t) (k->nb[3] / ggml_type_size(k->type)) : (uint32_t) (k->nb[3] / 4), (uint32_t) (top_k->nb[1] / sizeof(int32_t)), (uint32_t) (top_k->nb[3] / sizeof(int32_t)), (uint32_t) (mask->nb[1] / sizeof(ggml_fp16_t)), (uint32_t) (mask->nb[3] / sizeof(ggml_fp16_t)), (uint32_t) mask->ne[3], - n_batch, k_row_words, + n_batch, tok_dq ? (uint32_t) k->ne[0] : k_row_words, }; st.kc_buf = ggml_vk_subbuffer(ctx, ctx->prealloc_y, 0); @@ -12123,8 +12141,9 @@ union_unavailable:; st.active = true; st.kv_c = kv_c; st.n_batch = n_batch; - st.row_bytes = k_row_bytes; - st.row_elems = (uint32_t) (k->ne[0] / ggml_blck_size(k->type)); + st.row_bytes = row_by; + st.row_elems = tok_dq ? (uint32_t) k->ne[0] : (uint32_t) (k->ne[0] / ggml_blck_size(k->type)); + st.dequantized = tok_dq; return true; } diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_gather_dq.comp b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_gather_dq.comp new file mode 100644 index 00000000000..f2711539641 --- /dev/null +++ b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_gather_dq.comp @@ -0,0 +1,106 @@ +#version 450 + +#extension GL_EXT_control_flow_attributes : require +#extension GL_EXT_shader_16bit_storage : require +#extension GL_EXT_shader_explicit_arithmetic_types_float16 : require +#extension GL_EXT_shader_explicit_arithmetic_types_int8 : require +#extension GL_EXT_shader_explicit_arithmetic_types_int16 : require +#extension GL_EXT_shader_explicit_arithmetic_types_int32 : require + +#define FLOAT_TYPE float +#include "types.glsl" + +// Gather + DEQUANTISE for the per-token compact form: same row mapping as the verbatim +// flash_attn_gather.comp, same decode as flash_attn_gather_union_dq.comp. One variant per K +// type, via DATA_A_*. +// +// The union carries the decode only for n_batch > 1, because it needs more than one token to +// deduplicate. Single-token decode is the common case and it lands here; a verbatim gather +// leaves flash attention to decode inline, measured 1.98x the f16 op at q8_0 and 2.26x at +// q4_0, kv 11008. + +layout(local_size_x = 64, local_size_y = 1, local_size_z = 1) in; + +layout(binding = 0) readonly buffer A { A_TYPE data_a[]; }; +layout(binding = 1) readonly buffer TopBuf { int data_top[]; }; +layout(binding = 2) readonly buffer MBuf { float16_t data_m[]; }; +layout(binding = 3) writeonly buffer KcBuf { float16_t data_kc[]; }; +layout(binding = 4) writeonly buffer McBuf { float16_t data_mc[]; }; + +layout(push_constant) uniform Parameters { + uint n_kv; + uint n_kv_raw; + uint n_top_k; + uint kv_c; + uint nbk1; // K source row stride, BLOCKS + uint nbk3; // K source stream stride, BLOCKS + uint nbt1; + uint nbt3; + uint nbm1; + uint nbm3; + uint nem3; + uint n_batch; + uint row_elems; // elements per K row; 64 lanes x 8 covers the 512 the gate pins +} p; + +void main() { + const uint row = gl_WorkGroupID.x; + const uint stream = gl_WorkGroupID.z; + const uint tid = gl_LocalInvocationIndex; + + const uint ALL_TOKENS = 0xffffffffu; + uint src = p.n_kv; + uint owner = ALL_TOKENS; + if (row < p.n_kv_raw) { + src = row; + } else { + const uint off = row - p.n_kv_raw; + const uint tok = off / p.n_top_k; + const uint slot = off - tok * p.n_top_k; + if (tok < p.n_batch) { + owner = tok; + const int idx = data_top[stream * p.nbt3 + tok * p.nbt1 + slot]; + if (idx >= 0 && uint(idx) < p.n_kv - p.n_kv_raw) { + src = p.n_kv_raw + uint(idx); + } + } + } + + const uint dst_base = (stream * p.kv_c + row) * p.row_elems; + const uint e0 = tid * 8; + if (src < p.n_kv) { + const uint ib = stream * p.nbk3 + src * p.nbk1 + e0 / QUANT_K; + const uint iqs = e0 % QUANT_K; + const float d = float(data_a[ib].d); +#if defined(DATA_A_Q8_0) + // element e is qs[e] + [[unroll]] for (uint l = 0; l < 8; ++l) { + data_kc[dst_base + e0 + l] = float16_t(d * float(data_a[ib].qs[iqs + l])); + } +#elif defined(DATA_A_Q4_0) + // byte j carries element j in its low nibble and element j+16 in its high one + const uint shift = (iqs >> 4) * 4; + const uint byte0 = iqs & 0xF; + [[unroll]] for (uint l = 0; l < 8; ++l) { + const float q = float((data_a[ib].qs[byte0 + l] >> shift) & 0xF) - 8.0f; + data_kc[dst_base + e0 + l] = float16_t(d * q); + } +#else +#error "no element mapping written for this K type" +#endif + } else { + [[unroll]] for (uint l = 0; l < 8; ++l) { + data_kc[dst_base + e0 + l] = float16_t(0.0); + } + } + + const float NEG_INF = uintBitsToFloat(0xff800000); + if (tid < p.n_batch) { + const uint mc_idx = (stream * p.n_batch + tid) * p.kv_c + row; + float mv = NEG_INF; + if (src < p.n_kv && (owner == ALL_TOKENS || owner == tid)) { + mv = float(data_m[(stream % p.nem3) * p.nbm3 + tid * p.nbm1 + src]); + } + data_mc[mc_idx] = float16_t(mv); + } +} 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 efef6542a77..fe82b8fb6cb 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp @@ -831,6 +831,8 @@ void process_shaders() { for (const auto& tname : {"q4_0", "q8_0"}) { string_to_spv("flash_attn_gather_union_dq_" + std::string(tname), "flash_attn_gather_union_dq.comp", {{"DATA_A_" + to_uppercase(tname), "1"}}); + string_to_spv("flash_attn_gather_dq_" + std::string(tname), "flash_attn_gather_dq.comp", + {{"DATA_A_" + to_uppercase(tname), "1"}}); } string_to_spv("dsv4_hc_pre_f32", "dsv4_hc_pre.comp", {}); string_to_spv("dsv4_hc_comb_f32", "dsv4_hc_comb.comp", {}); diff --git a/tests/test-backend-ops.cpp b/tests/test-backend-ops.cpp index c3cb4ef8e99..7c943a503ae 100644 --- a/tests/test-backend-ops.cpp +++ b/tests/test-backend-ops.cpp @@ -11969,7 +11969,7 @@ static std::vector> make_test_cases_perf() { // worst-case compact set 2304 + nb*512 crosses kv/2 at nb=6, so those cells measure // whether the gate can be opened by the union rather than by the worst case. for (int kv : { 11008, 35584, 133888 }) { - for (int nb : { 2, 4, 8, 16 }) { + for (int nb : { 1, 2, 4, 8, 16 }) { test_cases.emplace_back(new test_flash_attn_ext_top_k(kv, nb, 2304, 512, false, 1, 60)); } } @@ -11983,8 +11983,10 @@ static std::vector> make_test_cases_perf() { } // q8_0 K/V at the same shapes: DSv4 with -ctk q8_0 took the dense fallback before the // gather became type-agnostic, so this is the cell that says whether it now pays there. + // nb=1 is plain autoregressive decode: the union needs nb > 1 to dedup, so this width + // takes the per-token gather. Its f16 row is in the grid above. for (ggml_type tk : { GGML_TYPE_Q8_0, GGML_TYPE_Q4_0 }) { - for (int nb : { 2, 4, 8, 16 }) { + for (int nb : { 1, 2, 4, 8, 16 }) { test_cases.emplace_back(new test_flash_attn_ext_top_k(11008, nb, 2304, 512, false, 1, 60, tk)); } } From f9471c66bbc54b92b5207d56ddd119c7cc3b5a8b Mon Sep 17 00:00:00 2001 From: Nathan Wilson Date: Thu, 20 Aug 2026 00:25:10 +0000 Subject: [PATCH 38/68] vulkan: dequantise the cache for the DeepSeek V4 sparse prefill ggml_vk_flash_attn_top_k declined any K/V but f16, so a DSv4 run with a quantised cache took dense FA for every prefill batch - O(kv) against the sparse path's O(n_kv_raw + n_top_k), a penalty that grows with depth. A 128GB community run measured it end to end: prefill -5.4/-22.2/-39.4/-52.9% against f16 at source depths 17k/33k/67k/134k with -ctk q8_0. The sparse shaders stay f16-only. A quantised cache is dequantised once per op into the prealloc_x f16 scratch by the dense path's own fused dequant+transpose pipeline, and the shaders read the scratch. K and V are the same tensor in this path, so one pass covers both at half the dense path's scratch footprint. Types without a dequant-transpose pipeline, oversized caches, and GGML_VK_FA_DEQUANT=0 keep the dense fallback. test-backend-ops perf at nb=1024, kv 5504/11008/19200/35584 (the community depths in compressed-K rows): q8_0 was 60.9 -> 427.6 ms dense (1.30x -> 8.27x of f16, growing with kv) and lands within +0.8% of f16's flat 46.9-51.7 ms with the scratch; the dequant pass itself costs under 1%. 36/36 FLASH_ATTN_EXT top-k cases pass, including 6 new quantised prefill cases (nb=64 ns=1/2, nb=128 split form). Co-Authored-By: Claude Fable 5 Assisted-by: Claude (Opus 5) --- ggml/src/ggml-vulkan/ggml-vulkan.cpp | 56 ++++++++++++++++++++++++---- tests/test-backend-ops.cpp | 17 +++++++++ 2 files changed, 65 insertions(+), 8 deletions(-) diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp index 92e2c58282b..7fbc651d97f 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp +++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp @@ -11498,9 +11498,22 @@ static bool ggml_vk_flash_attn_top_k(ggml_backend_vk_context * ctx, vk_context & const ggml_tensor * mask, const ggml_tensor * sinks, ggml_tensor * dst) { const ggml_tensor * top_k = dst->src[5]; static const char * top_k_env = getenv("GGML_VK_FA_TOPK"); + // The sparse shaders read f16 only. Serve a quantised cache by dequantising it once into + // the f16 scratch, the same pass the dense path uses. K and V are the same tensor here, so + // one pass covers both. Without this a quantised cache declines to dense FA, which costs + // O(kv) where the sparse path costs O(n_kv_raw + n_top_k). + static const char * fa_dequant_env = getenv("GGML_VK_FA_DEQUANT"); + const uint64_t kv_f16_sz = (uint64_t) ggml_nelements(k) * sizeof(ggml_fp16_t); + const bool dequant_kv = top_k && k->type != GGML_TYPE_F16 && + !(fa_dequant_env && fa_dequant_env[0] == '0') && + ctx->device->pipeline_dequant_transpose[k->type] != nullptr && + k->nb[0] == ggml_type_size(k->type) && + ggml_is_contiguously_allocated(k) && + kv_f16_sz <= ctx->device->properties.limits.maxStorageBufferRange && + ggml_vk_fa_dequant_scratch_fits(ctx, kv_f16_sz); if ((top_k_env && top_k_env[0] == '0') || !top_k || (!ctx->device->pipeline_flash_attn_top_k_f16 && !ctx->device->pipeline_flash_attn_top_k_cm_f16) || - q->type != GGML_TYPE_F32 || k->type != GGML_TYPE_F16 || v->type != GGML_TYPE_F16 || + q->type != GGML_TYPE_F32 || (k->type != GGML_TYPE_F16 && !dequant_kv) || v->type != k->type || !mask || mask->type != GGML_TYPE_F16 || top_k->type != GGML_TYPE_I32 || q->ne[0] != 512 || q->ne[1] < 64 || k->ne[0] != 512 || v->ne[0] != 512 || q->ne[2] != 64 || k->ne[2] != 1 || v->ne[2] != 1 || @@ -11586,14 +11599,43 @@ static bool ggml_vk_flash_attn_top_k(ggml_backend_vk_context * ctx, vk_context & } } + vk_subbuffer k_buf = ggml_vk_tensor_subbuffer(ctx, k); + if (dequant_kv) { + if (ctx->prealloc_size_x < kv_f16_sz) { + ctx->prealloc_size_x = kv_f16_sz; + ggml_vk_preallocate_buffers(ctx, subctx); + } + vk_pipeline tr_k = ctx->device->pipeline_dequant_transpose[k->type]; + ggml_pipeline_request_descriptor_sets(ctx, tr_k, 1); + if (ctx->prealloc_x_need_sync) { + ggml_vk_sync_buffers(ctx, subctx); + } + const vk_subbuffer k_dst = vk_subbuffer{ ctx->prealloc_x, 0, kv_f16_sz }; + const uint32_t k_nel = (uint32_t) ggml_nelements(k); + const std::vector tr_pc = { (uint32_t) k->ne[0], (uint32_t) k->ne[2], (uint32_t) k->ne[1], 0, k_nel }; + ggml_vk_dispatch_pipeline(ctx, subctx, tr_k, { k_buf, k_dst }, tr_pc, { k_nel, 1, 1 }); + ggml_vk_sync_buffers(ctx, subctx); + ctx->prealloc_x_need_sync = true; + k_buf = k_dst; + ggml_vk_perf_mark_subop(ctx, subctx, "FA_KV_DEQUANT (sub-op)"); + } + // strides of what the shaders actually read: the source cache, or the contiguous + // [HS, KV, n_head_kv, ns] f16 scratch the dequant just wrote + const uint32_t k_stride = dequant_kv ? (uint32_t) k->ne[0] : (uint32_t) (k->nb[1] / sizeof(ggml_fp16_t)); + const uint32_t k_nb3_el = dequant_kv ? (uint32_t) ((uint64_t) k->ne[0] * k->ne[1] * k->ne[2]) + : (uint32_t) (k->nb[3] / sizeof(ggml_fp16_t)); + const uint32_t k_nb2_byte = dequant_kv ? (uint32_t) ((uint64_t) k->ne[0] * k->ne[1] * sizeof(ggml_fp16_t)) + : (uint32_t) k->nb[2]; + const uint32_t k_nb3_byte = dequant_kv ? (uint32_t) (k_nb3_el * sizeof(ggml_fp16_t)) : (uint32_t) k->nb[3]; + vk_op_flash_attn_top_k_push_constants pc = { (uint32_t) q->ne[1], (uint32_t) k->ne[1], (uint32_t) n_kv_raw, (uint32_t) top_k->ne[0], (uint32_t) q->ne[2], (uint32_t) (q->nb[1] / sizeof(float)), (uint32_t) (q->nb[2] / sizeof(float)), (uint32_t) (q->nb[3] / sizeof(float)), - (uint32_t) (k->nb[1] / sizeof(ggml_fp16_t)), - (uint32_t) (k->nb[3] / sizeof(ggml_fp16_t)), + k_stride, + k_nb3_el, (uint32_t) (mask->nb[1] / sizeof(ggml_fp16_t)), (uint32_t) (mask->nb[3] / sizeof(ggml_fp16_t)), (uint32_t) (top_k->nb[1] / sizeof(int32_t)), @@ -11629,7 +11671,6 @@ static bool ggml_vk_flash_attn_top_k(ggml_backend_vk_context * ctx, vk_context & vk_fa_tuning_params tuning = get_fa_tuning_params(ctx->device, D, D, N, raw_kv, GGML_TYPE_F16, GGML_TYPE_F16, f32acc); const uint32_t q_stride = (uint32_t) (q->nb[1] / sizeof(float)); - const uint32_t k_stride = (uint32_t) (k->nb[1] / sizeof(ggml_fp16_t)); const bool aligned = raw_kv % tuning.block_cols == 0 && (q_stride & 7) == 0 && (k_stride & 7) == 0; const vk_fa_pipeline_state raw_state = get_fa_pipeline_state(ctx->device, tuning, D, D, aligned, f32acc, true, false, false, GGML_TYPE_F16, GGML_TYPE_F16); @@ -11671,7 +11712,6 @@ static bool ggml_vk_flash_attn_top_k(ggml_backend_vk_context * ctx, vk_context & const uint32_t packed_gqa = mask_stride_in_split_kv | 1u; const uint32_t packed_partitions = (partitions << 16) | 1; const vk_subbuffer split_buf = ggml_vk_subbuffer(ctx, ctx->prealloc_split_k, 0); - const vk_subbuffer k_buf = ggml_vk_tensor_subbuffer(ctx, k); const vk_subbuffer mask_buf = ggml_vk_tensor_subbuffer(ctx, mask); const vk_subbuffer top_buf = ggml_vk_tensor_subbuffer(ctx, top_k); const vk_subbuffer dst_buf = ggml_vk_tensor_subbuffer(ctx, dst); @@ -11697,8 +11737,8 @@ static bool ggml_vk_flash_attn_top_k(ggml_backend_vk_context * ctx, vk_context & 1, NS, (uint32_t) mask->ne[1], (uint32_t) mask->ne[2], (uint32_t) mask->ne[3], q_stride, (uint32_t) q->nb[2], (uint32_t) q->nb[3], - k_stride, (uint32_t) k->nb[2], (uint32_t) k->nb[3], - k_stride, (uint32_t) k->nb[2], (uint32_t) k->nb[3], + k_stride, k_nb2_byte, k_nb3_byte, + k_stride, k_nb2_byte, k_nb3_byte, scale, 0.0f, 0.0f, n_head_log2, 1.0f, 1.0f, packed_gqa, mask_stride, packed_partitions, @@ -11731,7 +11771,7 @@ static bool ggml_vk_flash_attn_top_k(ggml_backend_vk_context * ctx, vk_context & ggml_pipeline_request_descriptor_sets(ctx, pipeline, 1); ggml_vk_dispatch_pipeline(ctx, subctx, pipeline, - {q_buf, ggml_vk_tensor_subbuffer(ctx, k), ggml_vk_tensor_subbuffer(ctx, mask), sinks_buf, + {q_buf, k_buf, ggml_vk_tensor_subbuffer(ctx, mask), sinks_buf, ggml_vk_tensor_subbuffer(ctx, top_k), ggml_vk_tensor_subbuffer(ctx, dst)}, pc, {(uint32_t) q->ne[1], (uint32_t) CEIL_DIV(q->ne[2], use_cm ? 32 : 8), (uint32_t) q->ne[3]}); ggml_vk_perf_mark_subop(ctx, subctx, use_cm ? "FA_TOP_K_CM (sub-op)" : "FA_TOP_K_SPARSE (sub-op)"); diff --git a/tests/test-backend-ops.cpp b/tests/test-backend-ops.cpp index 7c943a503ae..f00b69e2496 100644 --- a/tests/test-backend-ops.cpp +++ b/tests/test-backend-ops.cpp @@ -11436,6 +11436,12 @@ static std::vector> make_test_cases_eval() { test_cases.emplace_back(new test_flash_attn_ext_top_k(11008, 8, 2304, 512, false, 1, 60, tk)); test_cases.emplace_back(new test_flash_attn_ext_top_k(11008, 16, 2304, 512, false, 1, 86, tk)); test_cases.emplace_back(new test_flash_attn_ext_top_k(8192, 1, 1024, 512, false, 1, 0, tk)); + // prefill widths (nb >= 64), where the sparse shaders run on a dequantised f16 scratch + // instead of the cache. ns=2 covers the scratch's stream stride, and the 4096 case is + // wide enough for the raw/selected split form. + test_cases.emplace_back(new test_flash_attn_ext_top_k( 768, 64, 64, 128, false, 1, 0, tk)); + test_cases.emplace_back(new test_flash_attn_ext_top_k( 768, 64, 64, 128, false, 2, 0, tk)); + test_cases.emplace_back(new test_flash_attn_ext_top_k(4096, 128, 256, 512, false, 1, 0, tk)); } test_cases.emplace_back(new test_flash_attn_ext_top_k(8192, 4, 1024, 512, false, 2)); test_cases.emplace_back(new test_flash_attn_ext_top_k( 768, 64, 64, 128, false, 2)); @@ -11990,6 +11996,17 @@ static std::vector> make_test_cases_perf() { test_cases.emplace_back(new test_flash_attn_ext_top_k(11008, nb, 2304, 512, false, 1, 60, tk)); } } + // PREFILL widths with quantised K/V, which the sparse path now serves through a one-shot + // dequant into the f16 scratch; quantised should sit within ~1% of f16 at every kv here. + // nb=1024 is the reporting user's --ubatch-size; the kv list is their four source depths + // (17k/33k/67k/134k) in compressed-K rows. GGML_VK_FA_DEQUANT=0 reproduces the old dense + // fallback, whose gap GROWS with kv (dense is O(kv), sparse O(n_kv_raw + n_top_k)); f16 + // under GGML_VK_FA_TOPK=0 is the falsification arm for attributing that gap to the gate. + for (ggml_type tk : { GGML_TYPE_F16, GGML_TYPE_Q8_0, GGML_TYPE_Q4_0 }) { + for (int kv : { 5504, 11008, 19200, 35584 }) { + test_cases.emplace_back(new test_flash_attn_ext_top_k(kv, 1024, 2304, 512, false, 1, 0, tk)); + } + } return test_cases; } From f7ffc13272cb4fb6b29e3cbcd3f34d6f17a409e7 Mon Sep 17 00:00:00 2001 From: pepuscz Date: Mon, 24 Aug 2026 22:17:58 +0200 Subject: [PATCH 39/68] vulkan: use small Lightning Indexer CM for batches 4-15 --- ggml/src/ggml-vulkan/ggml-vulkan.cpp | 5 +++++ 1 file changed, 5 insertions(+) diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp index 7fbc651d97f..8969f3328d4 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp +++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp @@ -13191,6 +13191,11 @@ static vk_pipeline ggml_vk_op_get_pipeline(ggml_backend_vk_context * ctx, const if (ctx->device->pipeline_lightning_indexer_decode_cm_f16 && src0->ne[2] == 1) { return ctx->device->pipeline_lightning_indexer_decode_cm_f16; } + static const char * small_cm_env = getenv("GGML_VK_LIGHTNING_INDEXER_SMALL_CM"); + if (ctx->device->pipeline_lightning_indexer_cm_small_f16 && + src0->ne[2] >= 4 && src0->ne[2] < 16 && small_cm_env && small_cm_env[0] == '1') { + return ctx->device->pipeline_lightning_indexer_cm_small_f16; + } vk_pipeline cm = ctx->device->pipeline_lightning_indexer_cm_f16 ? ctx->device->pipeline_lightning_indexer_cm_f16 : ctx->device->pipeline_lightning_indexer_cm_small_f16; return cm && src0->ne[2] >= 16 ? cm : ctx->device->pipeline_lightning_indexer_f16; From ef56966fa024f002d94d9416f79ffec204935df9 Mon Sep 17 00:00:00 2001 From: Nathan Wilson Date: Wed, 26 Aug 2026 05:01:45 +0000 Subject: [PATCH 40/68] vulkan: hoist the Lightning Indexer K fragments out of the head loop The K tile is invariant across the head loop but coopMatLoad sits inside it, so the shader re-read the same 8 MatrixA fragments from shared memory once per head tile - 4x at N_HEAD=64. Each wave64 fragment load moves 1 KiB of LDS traffic because the 16x16 f16 fragment is replicated 4x across the subgroup. Same loads of the same data feeding the same coopMatMulAdd sequence, so the output is bit-identical; the fragments just live in registers (8 of them, 16 f16 per lane) instead of being re-fetched. gfx1151, test-backend-ops perf, kv 8704/33280/131584 x batch 1-15: 16-20% faster at every shape, no spills, no regression. Batch 1 is ordinary DSv4 decode and gets 12-18% of that. Co-Authored-By: Claude Opus 5 Assisted-by: Claude (Opus 5) --- .../lightning_indexer_decode_cm.comp | 15 ++++++++++++--- 1 file changed, 12 insertions(+), 3 deletions(-) diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/lightning_indexer_decode_cm.comp b/ggml/src/ggml-vulkan/vulkan-shaders/lightning_indexer_decode_cm.comp index fd555c76806..24d1b98d571 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/lightning_indexer_decode_cm.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/lightning_indexer_decode_cm.comp @@ -63,6 +63,17 @@ void main() { } barrier(); + // The K tile is invariant across the head loop, but coopMatLoad is inside it, so the shipped + // shader re-reads the same 8 A fragments from shared memory once per head tile (4x at + // N_HEAD=64). Each wave64 fragment load moves 1 KiB of LDS traffic because the 16x16 f16 + // fragment is replicated 4x across the subgroup, so those re-reads are the largest single + // item of per-tile cost after the multiplies themselves. Hoisting costs 8 A fragments of + // register state (16 f16 per lane each). + coopmat kmat[HEAD_SIZE / TILE]; + [[unroll]] for (uint d = 0; d < HEAD_SIZE; d += TILE) { + coopMatLoad(kmat[d / TILE], k_sh, d / 4, TILE_STRIDE, gl_CooperativeMatrixLayoutRowMajor); + } + float total = 0.0; for (uint head_base = 0; head_base < N_HEAD; head_base += TILE) { for (uint idx = tid; idx < TILE * VEC_PER_HEAD; idx += SUBGROUP_SIZE) { @@ -77,13 +88,11 @@ void main() { coopmat scores = coopmat(0.0); - coopmat kmat; coopmat qmat; [[unroll]] for (uint d = 0; d < HEAD_SIZE; d += TILE) { - coopMatLoad(kmat, k_sh, d / 4, TILE_STRIDE, gl_CooperativeMatrixLayoutRowMajor); coopMatLoad(qmat, q_sh, d / 4, TILE_STRIDE, gl_CooperativeMatrixLayoutColumnMajor); - scores = coopMatMulAdd(kmat, qmat, scores); + scores = coopMatMulAdd(kmat[d / TILE], qmat, scores); } coopMatStore(scores, score_sh, 0, SCORE_STRIDE, gl_CooperativeMatrixLayoutRowMajor); From 24d18f710e83e12f42233d4727165843bead22e7 Mon Sep 17 00:00:00 2001 From: Nathan Wilson Date: Wed, 26 Aug 2026 05:01:45 +0000 Subject: [PATCH 41/68] vulkan: route the whole small-batch Lightning Indexer window to the decode CM shader Builds directly on pepuscz's PR #6 and issue #10. Their per-kernel table - 43.8 us scalar against 27.9 us small CM per 1k scanned rows at batch 5, and the same ratio at every depth - is what showed the indexer's cost is a per-tile constant rather than anything to do with tokens or bytes. Once the cost is per tile, the thing to minimise is the tile count, and that is what this changes. The finding is downstream of their work; only the choice of shader differs. The decode CM shader puts 16 HEADS in the coopmat N dimension and dispatches one workgroup per token, so it issues 4*n_batch tiles per 16 KV rows. That is the arithmetic minimum - (64 heads x n_batch tokens) / 16 columns - because 64 heads fill its 16 columns exactly, with no remainder at any batch. The small CM shader puts 16 TOKENS in N and pays a flat 64 tiles however small the batch is, so at batch 5 eleven of its sixteen columns are padding. Measured cost tracks tile count: at kv=131584 batch 1 costs 1.95 us per 1k scanned rows and batch 5 costs 26.28, i.e. 13.5x for 5x the work, and both shaders sit at 6.6-7.8 ns per tile. The decode CM shader body has no n_batch == 1 assumption - token is gl_WorkGroupID.y and it indexes q/w/mask/dst by it - so the old gate was an artefact of where it was written. gfx1151, kv=131584 (526k source tokens), us/run, shipped vs this: batch 2 2335.6 -> 419.0 (5.6x, was on the scalar path) batch 3 3492.3 -> 631.8 (5.5x, was on the scalar path) batch 4 3267.9 -> 840.8 (3.9x) batch 5 3465.6 -> 1119.4 (3.1x) <- DSpark n-max 4 verify shape batch 8 3704.9 -> 1764.1 (2.1x) batch 15 4467.9 -> 3231.6 (1.4x) Batch 16 and 32 are unchanged, which confirms the arms are isolated. PR #6's route is kept as the opt-out arm rather than deleted, and is promoted from opt-in to default-on so that one variable is enough to reach it: default decode CM for the whole 2-15 window ..._DECODE_CM_BATCH=0 small CM for 4-15, scalar for 2-3 (PR #6) ..._DECODE_CM_BATCH=0 SMALL_CM=0 scalar for 2-15 (pre-PR #6 baseline) Verified all three route as documented and pass 29/29. Keeping it is not just courtesy: the tile-count argument is hardware independent, but decode CM re-reads the K tile once per token and that part is bandwidth dependent, so the crossover need not sit in the same place on other devices, and these numbers are from one gfx1151 box. A (head, token) packing shader was also tried and is strictly worse - it reaches 4b tiles only when the batch divides 16, and it divides the workgroup count by the batch. Kept out of tree. Numerics: decode CM sums the 64 heads in tiles of 16 rather than one at a time, so f32 accumulation order differs from the small CM path. Cleared by KLD A/B on trunc10 (-c 8192, -b/-ub small so every indexer dispatch routes through the window, 8 chunks wikitext-2, same binary, only the env var differing): -ub 8 mean KLD 0.000000, max 6.0e-5, RMS dp 0.000%, same-top 99.994%, PPL 10507558.3156 identical in both arms -ub 5 mean KLD 0.000000, max 8.8e-4 but 99.9% 4.9e-5 (one tail event, no argmax flip), RMS dp 0.000%, same-top 100.000% Same class as PR #6's own measured numbers (max 5.5e-5, same-top 100.000%) and 45x tighter than the FA_WAVE32 change already shipped. Adds the eval coverage the 2-15 window never had (batches 4/8/15 at kv=256) and the perf grid these numbers came from. Co-Authored-By: Claude Opus 5 Assisted-by: Claude (Opus 5) --- ggml/src/ggml-vulkan/ggml-vulkan.cpp | 33 +++++++++++++++++++++++++--- tests/test-backend-ops.cpp | 10 ++++++++- 2 files changed, 39 insertions(+), 4 deletions(-) diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp index 8969f3328d4..a47792968ec 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp +++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp @@ -13188,12 +13188,39 @@ static vk_pipeline ggml_vk_op_get_pipeline(ggml_backend_vk_context * ctx, const if (src0->type == GGML_TYPE_F32 && src1->type == GGML_TYPE_F16 && dst->type == GGML_TYPE_F32 && src0->ne[0] == 128 && src0->ne[1] == 64 && src1->ne[1] == 1 && ctx->device->pipeline_lightning_indexer_f16) { - if (ctx->device->pipeline_lightning_indexer_decode_cm_f16 && src0->ne[2] == 1) { + // Small-batch routing. Three arms, one env var each, most specific first: + // + // default decode CM for the whole 2-15 window + // ..._DECODE_CM_BATCH=0 small CM for 4-15, scalar for 2-3 (PR #6) + // ..._DECODE_CM_BATCH=0 SMALL_CM=0 scalar for 2-15 (pre-PR #6 baseline) + // + // The decode CM shader puts 16 HEADS in the coopmat N dimension and dispatches one + // workgroup per token, so it issues 4*n_batch tiles per 16 KV rows - the arithmetic + // minimum, since 64 heads fill its 16 columns exactly. The small CM shader puts 16 + // TOKENS in N and pays a flat 64 tiles no matter how small the batch. Routing the + // whole window to decode CM is 2.7x at batch 5 and 4.8x at batch 2-3 (526k source + // tokens, gfx1151); the shader body has no n_batch == 1 assumption, token is + // gl_WorkGroupID.y, so the old ne[2] == 1 gate was an artefact. + // + // That measurement only exists because of pepuscz's PR #6 and issue #10: their + // per-kernel table (43.8 us scalar vs 27.9 us small CM per 1k scanned rows at batch + // 5, the same ratio at every depth) is what showed the cost is a per-tile constant, + // which is what makes the tile count the thing to minimise. Their small CM route is + // kept as the opt-out arm rather than deleted: the tile-count argument is hardware + // independent, but decode CM re-reads the K tile once per token, and that part is + // bandwidth dependent, so the crossover need not sit here on other devices. + static const char * decode_cm_batch_env = getenv("GGML_VK_LIGHTNING_INDEXER_DECODE_CM_BATCH"); + static const bool decode_cm_batch_on = !decode_cm_batch_env || decode_cm_batch_env[0] != '0'; + const int64_t decode_cm_max = decode_cm_batch_on ? 15 : 1; + if (ctx->device->pipeline_lightning_indexer_decode_cm_f16 && src0->ne[2] <= decode_cm_max) { return ctx->device->pipeline_lightning_indexer_decode_cm_f16; } + // PR #6 as its author proposed it, now default-on so that the kill switch above is + // enough on its own to reach it. static const char * small_cm_env = getenv("GGML_VK_LIGHTNING_INDEXER_SMALL_CM"); - if (ctx->device->pipeline_lightning_indexer_cm_small_f16 && - src0->ne[2] >= 4 && src0->ne[2] < 16 && small_cm_env && small_cm_env[0] == '1') { + static const bool small_cm_on = !small_cm_env || small_cm_env[0] != '0'; + if (ctx->device->pipeline_lightning_indexer_cm_small_f16 && small_cm_on && + src0->ne[2] >= 4 && src0->ne[2] < 16) { return ctx->device->pipeline_lightning_indexer_cm_small_f16; } vk_pipeline cm = ctx->device->pipeline_lightning_indexer_cm_f16 ? diff --git a/tests/test-backend-ops.cpp b/tests/test-backend-ops.cpp index f00b69e2496..55d25ccf9c3 100644 --- a/tests/test-backend-ops.cpp +++ b/tests/test-backend-ops.cpp @@ -11374,7 +11374,7 @@ static std::vector> make_test_cases_eval() { // lightning_indexer for (int kv : { 256 }) { - for (int bs : { 1, 512 }) { + for (int bs : { 1, 4, 8, 15, 512 }) { for (int nh : { 32, 64 }) { for (auto [ns, nm] : { std::pair{1, 1}, std::pair{4, 4}, std::pair{4, 1} }) { for (ggml_type type_K : {GGML_TYPE_F32, GGML_TYPE_F16, GGML_TYPE_BF16, GGML_TYPE_Q8_0, GGML_TYPE_Q5_1, GGML_TYPE_Q5_0, GGML_TYPE_Q4_1, GGML_TYPE_Q4_0, GGML_TYPE_IQ4_NL}) { @@ -11947,6 +11947,14 @@ static std::vector> make_test_cases_perf() { test_cases.emplace_back(new test_lightning_indexer(128, 64, kv, 2048, 1, 1, GGML_TYPE_F16)); } + // DSpark verify-step indexer rows: batch 1-32 across the 4-15 small-CM routing window, + // at the compressed key counts a 128k/491k source context produces (source/4). + for (int kv : { 8704, 33280, 131584 }) { + for (int bs : { 1, 2, 3, 4, 5, 6, 8, 12, 15, 16, 32 }) { + test_cases.emplace_back(new test_lightning_indexer(128, 64, kv, bs, 1, 1, GGML_TYPE_F16)); + } + } + // sparse top-k FA at V4 decode/prefill shapes — the A/B instrument for the // gather-to-compact work (n_active = n_kv_raw + n_top_k stays fixed as kv grows). // nb 1/8 currently takes the DENSE path (the sparse shader gates on nb >= 64): From 4d83452a433a5da46b1faa0bda25d273b5cf1bc0 Mon Sep 17 00:00:00 2001 From: Nathan Wilson Date: Wed, 26 Aug 2026 07:15:56 +0000 Subject: [PATCH 42/68] vulkan: hoist the Lightning Indexer K fragments out of the CM head loop Same defect as the decode CM shader, larger blast radius. `k_sh[wave]` is written once above the head loop and never touched again, but `coopMatLoad` sits inside both the head_base and head_local loops, so each wave re-reads the same 8 MatrixA fragments N_HEAD times: 512 loads where 8 would do. Each wave64 fragment load moves 1 KiB of shared-memory traffic because the 16x16 f16 fragment is replicated 4x across the subgroup. Same loads of the same data feeding the same coopMatMulAdd order, so the output is bit-identical; the fragments live in registers (8 of them, 16 f16 per lane) instead of being re-fetched. This source compiles to both the prefill pipeline and the small-batch one, so both gain. gfx1151, test-backend-ops perf, us/run: prefill (lightning_indexer_cm_f16) kv 8704 nb 2048 27492.5 -> 18207.0 1.51x kv 33280 nb 2048 104101.9 -> 68947.8 1.51x kv 131584 nb 2048 421239.8 -> 281540.5 1.50x kv 131584 nb 16 3456.2 -> 2329.2 1.48x kv 131584 nb 32 6858.6 -> 4623.6 1.48x small batch (lightning_indexer_cm_small_f16) kv 131584 nb 5 3433.1 -> 2783.3 1.23x kv 131584 nb 8 3682.3 -> 3105.1 1.19x kv 131584 nb 15 4474.4 -> 4209.5 1.06x 1.48-1.54x on prefill at every depth and batch measured, flat. The small-batch variant gains less, which says the 1-wave configuration is bound by its Q staging and epilogue rather than by fragment traffic. No regression at any shape, so the 8 held fragments (64 VGPRs per lane) neither spill nor cost occupancy even in the prefill pipeline, which runs 512 threads at 62.7 KiB of shared memory and was already pinned to one workgroup per CU. LIGHTNING_INDEXER 29/29. Co-Authored-By: Claude Opus 5 Assisted-by: Claude (Opus 5) --- .../vulkan-shaders/lightning_indexer_cm.comp | 13 ++++++++++--- 1 file changed, 10 insertions(+), 3 deletions(-) diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/lightning_indexer_cm.comp b/ggml/src/ggml-vulkan/vulkan-shaders/lightning_indexer_cm.comp index d4379e8899d..6747feba505 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/lightning_indexer_cm.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/lightning_indexer_cm.comp @@ -88,6 +88,15 @@ void main() { } barrier(); + // k_sh[wave] is written once above and never touched again, but coopMatLoad sits inside both + // the head_base and head_local loops, so the same 8 A fragments are re-read from shared + // memory N_HEAD times per wave (512 loads where 8 would do). Each wave64 fragment load moves + // 1 KiB of LDS traffic because the 16x16 f16 fragment is replicated 4x across the subgroup. + coopmat kmat_h[HEAD_SIZE / TILE]; + [[unroll]] for (uint d = 0; d < HEAD_SIZE; d += TILE) { + coopMatLoad(kmat_h[d / TILE], k_sh[wave], d / 4, TILE_STRIDE, gl_CooperativeMatrixLayoutRowMajor); + } + for (uint head_base = 0; head_base < N_HEAD; head_base += HEADS_PER_TILE) { for (uint idx = tid; idx < HEADS_PER_TILE * TILE * VEC_PER_HEAD; idx += gl_WorkGroupSize.x) { const uint head_local = idx / (TILE * VEC_PER_HEAD); @@ -113,13 +122,11 @@ void main() { [[unroll]] for (uint head_local = 0; head_local < HEADS_PER_TILE; ++head_local) { coopmat scores = coopmat(0.0); - coopmat kmat; coopmat qmat; [[unroll]] for (uint d = 0; d < HEAD_SIZE; d += TILE) { - coopMatLoad(kmat, k_sh[wave], d / 4, TILE_STRIDE, gl_CooperativeMatrixLayoutRowMajor); coopMatLoad(qmat, q_sh[head_local], d / 4, TILE_STRIDE, gl_CooperativeMatrixLayoutColumnMajor); - scores = coopMatMulAdd(kmat, qmat, scores); + scores = coopMatMulAdd(kmat_h[d / TILE], qmat, scores); } coopMatStore(scores, score_sh[wave], 0, SCORE_STRIDE, gl_CooperativeMatrixLayoutRowMajor); From 883b4d8fafa92676e2567e221e42c266f302ff53 Mon Sep 17 00:00:00 2001 From: Nathan Wilson Date: Wed, 26 Aug 2026 13:29:50 +0000 Subject: [PATCH 43/68] vulkan: support arbitrary Lightning Indexer head counts via specialization constant The indexer pipelines hardcoded nh=64 (DeepSeek-V4's geometry) in supports_op and as a shader constant, so any other head count fell to the CPU backend. Qwen3.8-Flash-Next (qwen4_exp, released today) carries the same indexer at nh=4. N_HEAD becomes specialization constant 1, one pipeline per supported head count (LI_NH_VALUES = {4, 32, 64}), selected by q->ne[1] at dispatch. It must be a specialization constant, not a push constant: the head loop's trip count has to stay visible to the compiler - as a push constant the loop cannot unroll and decode costs ~55% (measured, gfx1151, kv=131584, and why this commit is not the simpler design). nh=4 does not fill a 16-wide head tile, so the tail tile stages q as zeros (relu(0)*w = 0, padded heads drop out of the sum) and the weight read clamps its index. Both guards are gated on N_HEAD % TILE != 0, a specialization-constant expression, so at nh=64 they fold away at pipeline compile and the codegen is equivalent to the previous shader - measured, because relying on the compiler to range-prove head < N_HEAD instead still cost ~45%: kv=131584, us/run before after decode b1 226.8 219.6 decode b5 1058.8 1007.4 decode b15 3174.0 2966.8 prefill b16 2329.2 2230.2 prefill b2048 281540.5 281169.1 No regression at nh=64; the small decode improvement is within a cautious reading of run-to-run variance and is not claimed. test-backend-ops LIGHTNING_INDEXER: 44/44, up from 29/29 - the suite's existing nh=32 cases had been silently reporting "not supported" under the old gate and now run against the CPU reference. Head size stays pinned at 128: the coopmat tiling and f16vec4 staging both assume it, and every known indexer model (DSv4, GLM-DSA, Qwen4-exp) uses 128. Co-Authored-By: Claude Opus 5 Assisted-by: Claude (Opus 5) --- ggml/src/ggml-vulkan/ggml-vulkan.cpp | 96 ++++++++++++------- .../vulkan-shaders/lightning_indexer_cm.comp | 10 +- .../lightning_indexer_decode_cm.comp | 28 +++++- 3 files changed, 94 insertions(+), 40 deletions(-) diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp index a47792968ec..acd8c03bf62 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp +++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp @@ -799,6 +799,22 @@ static bool ggml_vk_lightning_indexer_k_type_supported(ggml_type type) { return std::find(lightning_indexer_k_types.begin(), lightning_indexer_k_types.end(), type) != lightning_indexer_k_types.end(); } +// Indexer head counts we build kernels for: 4 = Qwen4-exp (qwen4_exp), 64 = DeepSeek-V4 and +// GLM-DSA, 32 kept because test-backend-ops exercises it. N_HEAD is a specialization constant, +// so each entry is a separately compiled kernel with the head loop fully unrolled - as a push +// constant the loop cannot unroll and decode costs ~55% more (measured, gfx1151, kv=131584). +static constexpr uint32_t LI_NH_VALUES[] = { 4, 32, 64 }; +#define LI_NH_COUNT (sizeof(LI_NH_VALUES) / sizeof(LI_NH_VALUES[0])) + +static int ggml_vk_li_nh_index(int64_t nh) { + for (size_t i = 0; i < LI_NH_COUNT; ++i) { + if ((int64_t) LI_NH_VALUES[i] == nh) { + return (int) i; + } + } + return -1; +} + struct vk_device_struct { std::recursive_mutex mutex; mutable std::shared_mutex pinned_memory_mutex; @@ -1105,10 +1121,11 @@ struct vk_device_struct { vk_pipeline pipeline_lightning_indexer_f32[GGML_TYPE_COUNT]; // [size_idx][kda] where size_idx: 0=d16, 1=d32, 2=d64, 3=d128 vk_pipeline pipeline_gated_delta_net[4][2]; - vk_pipeline pipeline_lightning_indexer_f16; - vk_pipeline pipeline_lightning_indexer_cm_f16; - vk_pipeline pipeline_lightning_indexer_cm_small_f16; - vk_pipeline pipeline_lightning_indexer_decode_cm_f16; + // One pipeline per supported indexer head count; see LI_NH_VALUES. + vk_pipeline pipeline_lightning_indexer_f16[LI_NH_COUNT]; + vk_pipeline pipeline_lightning_indexer_cm_f16[LI_NH_COUNT]; + vk_pipeline pipeline_lightning_indexer_cm_small_f16[LI_NH_COUNT]; + vk_pipeline pipeline_lightning_indexer_decode_cm_f16[LI_NH_COUNT]; vk_pipeline pipeline_flash_attn_top_k_f16; vk_pipeline pipeline_flash_attn_top_k_cm_f16; vk_pipeline pipeline_flash_attn_gather_f16; @@ -6302,28 +6319,36 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { } if (device->subgroup_arithmetic && device->subgroup_size == 64) { - ggml_vk_create_pipeline(device, device->pipeline_lightning_indexer_f16, - "lightning_indexer_f16", lightning_indexer_f16_len, lightning_indexer_f16_data, "main", 5, - sizeof(vk_op_lightning_indexer_cm_push_constants), {8, 1, 1}, {device->subgroup_size}, 1, true, true, - device->subgroup_size); + for (size_t nhi = 0; nhi < LI_NH_COUNT; ++nhi) { + ggml_vk_create_pipeline(device, device->pipeline_lightning_indexer_f16[nhi], + "lightning_indexer_f16", lightning_indexer_f16_len, lightning_indexer_f16_data, "main", 5, + sizeof(vk_op_lightning_indexer_cm_push_constants), {8, 1, 1}, {device->subgroup_size, LI_NH_VALUES[nhi]}, 1, true, true, + device->subgroup_size); + } #if defined(VK_KHR_cooperative_matrix) && defined(GGML_VULKAN_COOPMAT_GLSLC_SUPPORT) if (device->coopmat_support && device->coopmat_support_16x16x16_f32acc && device->subgroup_size_control) { - ggml_vk_create_pipeline(device, device->pipeline_lightning_indexer_cm_small_f16, - "lightning_indexer_cm_small_f16", lightning_indexer_cm_small_f16_len, lightning_indexer_cm_small_f16_data, "main", 5, - sizeof(vk_op_lightning_indexer_cm_push_constants), {16, 16, 1}, {device->subgroup_size}, 1, true, true, - device->subgroup_size); + for (size_t nhi = 0; nhi < LI_NH_COUNT; ++nhi) { + ggml_vk_create_pipeline(device, device->pipeline_lightning_indexer_cm_small_f16[nhi], + "lightning_indexer_cm_small_f16", lightning_indexer_cm_small_f16_len, lightning_indexer_cm_small_f16_data, "main", 5, + sizeof(vk_op_lightning_indexer_cm_push_constants), {16, 16, 1}, {device->subgroup_size, LI_NH_VALUES[nhi]}, 1, true, true, + device->subgroup_size); + } if (device->properties.limits.maxComputeWorkGroupInvocations >= 512 && device->properties.limits.maxComputeWorkGroupSize[0] >= 512 && device->properties.limits.maxComputeSharedMemorySize >= 64 * 1024) { - ggml_vk_create_pipeline(device, device->pipeline_lightning_indexer_cm_f16, - "lightning_indexer_cm_f16", lightning_indexer_cm_f16_len, lightning_indexer_cm_f16_data, "main", 5, - sizeof(vk_op_lightning_indexer_cm_push_constants), {128, 16, 1}, {device->subgroup_size}, 1, true, true, + for (size_t nhi = 0; nhi < LI_NH_COUNT; ++nhi) { + ggml_vk_create_pipeline(device, device->pipeline_lightning_indexer_cm_f16[nhi], + "lightning_indexer_cm_f16", lightning_indexer_cm_f16_len, lightning_indexer_cm_f16_data, "main", 5, + sizeof(vk_op_lightning_indexer_cm_push_constants), {128, 16, 1}, {device->subgroup_size, LI_NH_VALUES[nhi]}, 1, true, true, + device->subgroup_size); + } + } + for (size_t nhi = 0; nhi < LI_NH_COUNT; ++nhi) { + ggml_vk_create_pipeline(device, device->pipeline_lightning_indexer_decode_cm_f16[nhi], + "lightning_indexer_decode_cm_f16", lightning_indexer_decode_cm_f16_len, lightning_indexer_decode_cm_f16_data, "main", 5, + sizeof(vk_op_lightning_indexer_cm_push_constants), {16, 1, 1}, {device->subgroup_size, LI_NH_VALUES[nhi]}, 1, true, true, device->subgroup_size); } - ggml_vk_create_pipeline(device, device->pipeline_lightning_indexer_decode_cm_f16, - "lightning_indexer_decode_cm_f16", lightning_indexer_decode_cm_f16_len, lightning_indexer_decode_cm_f16_data, "main", 5, - sizeof(vk_op_lightning_indexer_cm_push_constants), {16, 1, 1}, {device->subgroup_size}, 1, true, true, - device->subgroup_size); ggml_vk_create_pipeline(device, device->pipeline_flash_attn_top_k_cm_f16, "flash_attn_top_k_cm_f16", flash_attn_top_k_cm_f16_len, flash_attn_top_k_cm_f16_data, "main", 6, sizeof(vk_op_flash_attn_top_k_push_constants), {1, 1, 1}, {512, device->subgroup_size}, 1, true, true, @@ -13184,10 +13209,12 @@ static vk_pipeline ggml_vk_op_get_pipeline(ggml_backend_vk_context * ctx, const return nullptr; case GGML_OP_LIGHTNING_INDEXER: // fork fast path: f16 K on wave64 subgroup-arithmetic devices routes to the tuned - // scalar-64/CM kernels; anything else falls through to the generic pipeline table + // scalar-64/CM kernels (head counts in LI_NH_VALUES); anything else falls through to + // the generic pipeline table if (src0->type == GGML_TYPE_F32 && src1->type == GGML_TYPE_F16 && dst->type == GGML_TYPE_F32 && - src0->ne[0] == 128 && src0->ne[1] == 64 && src1->ne[1] == 1 && - ctx->device->pipeline_lightning_indexer_f16) { + src0->ne[0] == 128 && src1->ne[1] == 1 && + ggml_vk_li_nh_index(src0->ne[1]) >= 0 && ctx->device->pipeline_lightning_indexer_f16[0]) { + const int nhi = ggml_vk_li_nh_index(src0->ne[1]); // Small-batch routing. Three arms, one env var each, most specific first: // // default decode CM for the whole 2-15 window @@ -13212,20 +13239,20 @@ static vk_pipeline ggml_vk_op_get_pipeline(ggml_backend_vk_context * ctx, const static const char * decode_cm_batch_env = getenv("GGML_VK_LIGHTNING_INDEXER_DECODE_CM_BATCH"); static const bool decode_cm_batch_on = !decode_cm_batch_env || decode_cm_batch_env[0] != '0'; const int64_t decode_cm_max = decode_cm_batch_on ? 15 : 1; - if (ctx->device->pipeline_lightning_indexer_decode_cm_f16 && src0->ne[2] <= decode_cm_max) { - return ctx->device->pipeline_lightning_indexer_decode_cm_f16; + if (ctx->device->pipeline_lightning_indexer_decode_cm_f16[nhi] && src0->ne[2] <= decode_cm_max) { + return ctx->device->pipeline_lightning_indexer_decode_cm_f16[nhi]; } // PR #6 as its author proposed it, now default-on so that the kill switch above is // enough on its own to reach it. static const char * small_cm_env = getenv("GGML_VK_LIGHTNING_INDEXER_SMALL_CM"); static const bool small_cm_on = !small_cm_env || small_cm_env[0] != '0'; - if (ctx->device->pipeline_lightning_indexer_cm_small_f16 && small_cm_on && + if (ctx->device->pipeline_lightning_indexer_cm_small_f16[nhi] && small_cm_on && src0->ne[2] >= 4 && src0->ne[2] < 16) { - return ctx->device->pipeline_lightning_indexer_cm_small_f16; + return ctx->device->pipeline_lightning_indexer_cm_small_f16[nhi]; } - vk_pipeline cm = ctx->device->pipeline_lightning_indexer_cm_f16 ? - ctx->device->pipeline_lightning_indexer_cm_f16 : ctx->device->pipeline_lightning_indexer_cm_small_f16; - return cm && src0->ne[2] >= 16 ? cm : ctx->device->pipeline_lightning_indexer_f16; + vk_pipeline cm = ctx->device->pipeline_lightning_indexer_cm_f16[nhi] ? + ctx->device->pipeline_lightning_indexer_cm_f16[nhi] : ctx->device->pipeline_lightning_indexer_cm_small_f16[nhi]; + return cm && src0->ne[2] >= 16 ? cm : ctx->device->pipeline_lightning_indexer_f16[nhi]; } // only the k type selects a pipeline, the other types are fixed by ggml_lightning_indexer() if (ggml_vk_lightning_indexer_k_type_supported(src1->type)) { @@ -14319,9 +14346,14 @@ static void ggml_vk_lightning_indexer(ggml_backend_vk_context * ctx, vk_context& GGML_ASSERT(pipeline != nullptr); // the fork's wave64 f16 kernels take their own push-constant layout - if (pipeline == ctx->device->pipeline_lightning_indexer_f16 || - pipeline == ctx->device->pipeline_lightning_indexer_cm_f16 || - pipeline == ctx->device->pipeline_lightning_indexer_decode_cm_f16) { + bool fork_li = false; + for (size_t nhi = 0; nhi < LI_NH_COUNT && !fork_li; ++nhi) { + fork_li = pipeline == ctx->device->pipeline_lightning_indexer_f16[nhi] || + pipeline == ctx->device->pipeline_lightning_indexer_cm_f16[nhi] || + pipeline == ctx->device->pipeline_lightning_indexer_cm_small_f16[nhi] || + pipeline == ctx->device->pipeline_lightning_indexer_decode_cm_f16[nhi]; + } + if (fork_li) { ggml_vk_lightning_indexer_cm(ctx, subctx, dst, pipeline); return; } diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/lightning_indexer_cm.comp b/ggml/src/ggml-vulkan/vulkan-shaders/lightning_indexer_cm.comp index 6747feba505..a2004bd93d7 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/lightning_indexer_cm.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/lightning_indexer_cm.comp @@ -8,6 +8,10 @@ #extension GL_KHR_shader_subgroup_basic : require layout(constant_id = 0) const uint SUBGROUP_SIZE = 64; +// Head count is a SPECIALIZATION constant, not a push constant: the head loop's trip +// count must stay visible to the compiler. As a push constant it cannot unroll and +// decode costs ~55% more (measured, gfx1151, kv=131584). +layout(constant_id = 1) const uint N_HEAD = 64; #if N_WAVES == 1 layout(local_size_x_id = 0, local_size_y = 1, local_size_z = 1) in; #else @@ -40,7 +44,6 @@ layout(push_constant) uniform Parameters { const uint TILE = 16; const uint HEAD_SIZE = 128; -const uint N_HEAD = 64; const uint VEC_PER_HEAD = HEAD_SIZE / 4; const uint TILE_STRIDE = VEC_PER_HEAD + 2; const uint SCORE_STRIDE = TILE / 4 + 1; @@ -105,7 +108,7 @@ void main() { const uint d4 = head_idx % VEC_PER_HEAD; const uint token = token_base + token_local; f16vec4 value = f16vec4(0.0); - if (token < p.n_batch) { + if (token < p.n_batch && (N_HEAD % HEADS_PER_TILE == 0 || head_base + head_local < N_HEAD)) { const uint offset = stream * p.nbq3 + token * p.nbq2 + (head_base + head_local) * p.nbq1 + d4 * 4; value = f16vec4(data_q[offset], data_q[offset + 1], data_q[offset + 2], data_q[offset + 3]); } @@ -115,7 +118,8 @@ void main() { const uint head_local = tid / TILE; const uint token_local = tid % TILE; const uint token = token_base + token_local; - weight_sh[head_local][token_local] = token < p.n_batch ? data_w[stream * p.nbw3 + token * p.nbw1 + head_base + head_local] : 0.0; + weight_sh[head_local][token_local] = (token < p.n_batch && (N_HEAD % HEADS_PER_TILE == 0 || head_base + head_local < N_HEAD)) + ? data_w[stream * p.nbw3 + token * p.nbw1 + head_base + head_local] : 0.0; } barrier(); diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/lightning_indexer_decode_cm.comp b/ggml/src/ggml-vulkan/vulkan-shaders/lightning_indexer_decode_cm.comp index 24d1b98d571..5d32e51361b 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/lightning_indexer_decode_cm.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/lightning_indexer_decode_cm.comp @@ -7,6 +7,10 @@ #extension GL_KHR_memory_scope_semantics : require layout(constant_id = 0) const uint SUBGROUP_SIZE = 64; +// Head count is a SPECIALIZATION constant, not a push constant: the head loop's trip +// count must stay visible to the compiler. As a push constant it cannot unroll and +// decode costs ~55% more (measured, gfx1151, kv=131584). +layout(constant_id = 1) const uint N_HEAD = 64; layout(local_size_x_id = 0, local_size_y = 1, local_size_z = 1) in; layout(binding = 0) readonly buffer QBuf { float data_q[]; }; @@ -35,7 +39,6 @@ layout(push_constant) uniform Parameters { const uint TILE = 16; const uint HEAD_SIZE = 128; -const uint N_HEAD = 64; const uint VEC_PER_HEAD = HEAD_SIZE / 4; const uint TILE_STRIDE = VEC_PER_HEAD + 2; const uint SCORE_STRIDE = TILE / 4 + 1; @@ -80,9 +83,19 @@ void main() { const uint head_local = idx / VEC_PER_HEAD; const uint d4 = idx % VEC_PER_HEAD; const uint head = head_base + head_local; - const uint offset = stream * p.nbq3 + token * p.nbq2 + head * p.nbq1 + d4 * 4; - q_sh[head_local * TILE_STRIDE + d4] = - f16vec4(data_q[offset], data_q[offset + 1], data_q[offset + 2], data_q[offset + 3]); + // n_head need not be a multiple of TILE (Qwen4-exp indexers are nh=4), so the tail + // tile stages zeros rather than reading past the end of Q. Zero q gives score 0, + // and relu(0) * w = 0, so the padded heads drop out of the sum on their own. + // The guard exists only when N_HEAD is not tile-aligned (nh=4). The condition is a + // specialization-constant expression, so at nh=64 it folds to true and this compiles + // to the original unguarded load - measured: relying on the compiler to range-prove + // head < N_HEAD instead costs ~45% at nh=64. + f16vec4 value = f16vec4(0.0); + if (N_HEAD % TILE == 0 || head < N_HEAD) { + const uint offset = stream * p.nbq3 + token * p.nbq2 + head * p.nbq1 + d4 * 4; + value = f16vec4(data_q[offset], data_q[offset + 1], data_q[offset + 2], data_q[offset + 3]); + } + q_sh[head_local * TILE_STRIDE + d4] = value; } barrier(); @@ -100,8 +113,13 @@ void main() { if (tid < TILE && kv_base + tid < p.n_kv) { [[unroll]] for (uint head_local = 0; head_local < TILE; ++head_local) { + // Padded heads staged q as zero, so their score is 0 and the term dies in the + // multiply. Only the weight READ needs bounding, and clamping the index does that + // without a branch - a dynamic break here forces the loop rolled and costs ~60%. + const uint head = N_HEAD % TILE == 0 ? head_base + head_local + : min(head_base + head_local, N_HEAD - 1); const float score = score_sh[tid * SCORE_STRIDE + head_local / 4][head_local % 4]; - const float weight = data_w[stream * p.nbw3 + token * p.nbw1 + head_base + head_local]; + const float weight = data_w[stream * p.nbw3 + token * p.nbw1 + head]; total += max(score, 0.0) * weight; } } From 808ce1c87b3fa1ec942233cbc3d204a48e36a3f4 Mon Sep 17 00:00:00 2001 From: Nathan Wilson Date: Sun, 30 Aug 2026 09:02:36 +0000 Subject: [PATCH 44/68] vulkan: N_HEAD spec constant for the scalar-64 lightning indexer The scalar64 shader kept its hardcoded 64-head loop when the pipelines went per-head-count; nh 4 and 32 would have run with the wrong trip count. Assisted-by: Claude Fable 5 --- .../vulkan-shaders/lightning_indexer_scalar64.comp | 5 ++++- 1 file changed, 4 insertions(+), 1 deletion(-) diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/lightning_indexer_scalar64.comp b/ggml/src/ggml-vulkan/vulkan-shaders/lightning_indexer_scalar64.comp index 693bb3ece8e..83fd923fe73 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/lightning_indexer_scalar64.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/lightning_indexer_scalar64.comp @@ -7,6 +7,10 @@ #extension GL_KHR_shader_subgroup_basic : require layout(constant_id = 0) const uint SUBGROUP_SIZE = 64; +// Head count is a SPECIALIZATION constant, not a push constant: the head loop's trip +// count must stay visible to the compiler. As a push constant it cannot unroll and +// decode costs ~55% more (measured, gfx1151, kv=131584). +layout(constant_id = 1) const uint N_HEAD = 64; layout(local_size_x_id = 0, local_size_y = 1, local_size_z = 1) in; layout(binding = 0) readonly buffer QBuf { float data_q[]; }; @@ -34,7 +38,6 @@ layout(push_constant) uniform Parameters { } p; const uint K_PER_GROUP = 8; -const uint N_HEAD = 64; void main() { const uint lane = gl_SubgroupInvocationID; From 1007fdc4ada16480c06fc747df3aa28f714d35f3 Mon Sep 17 00:00:00 2001 From: Nathan Wilson Date: Thu, 30 Jul 2026 01:21:49 +0000 Subject: [PATCH 45/68] vulkan: hoist the coopmat1 FA P-fragment load out of the hsv_tile loop The GEMM2 P load depends only on bc_chunk, but sat inside the hsv_tile loop, so all four fragments were re-read from shared memory once per tile (2 tiles at HSV=128, 4 at HSV=256). Load them once into a coopmat array before the loop. Psh is not written again until the next KV block, so the fragments stay valid across the barriers inside it. Measured on gfx1151 (RADV), Qwen3-Coder-30B-A3B, pp2048 f16 KV: d8192 +6.9%, d16384 +8.1%, d32768 +9.2%. VGPRs and LDS unchanged, 0 spilled. FLASH_ATTN_EXT suite 5105/5105. Assisted-by: Claude Fable 5 --- .../ggml-vulkan/vulkan-shaders/flash_attn_cm1.comp | 11 ++++++++--- 1 file changed, 8 insertions(+), 3 deletions(-) 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 38ae42c81f0..695d459d0a4 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_cm1.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_cm1.comp @@ -423,6 +423,13 @@ void main() { const uint num_hsv_tiles = (HSV + MatBc * row_split - 1) / (MatBc * row_split); // round up + // Psh is not written again until the next KV block, so the P fragments are the same + // for every hsv_tile. Load them once instead of re-reading LDS per tile. + coopmat PMat[Bc / MatBc]; + [[unroll]] for (uint32_t bc_chunk = 0; bc_chunk < Bc / MatBc; ++bc_chunk) { + coopMatLoad(PMat[bc_chunk], Psh, bc_chunk * MatBc * psh_stride, psh_stride, gl_CooperativeMatrixLayoutColumnMajor); + } + // Each subgroup handles HSV/4 columns [[unroll]] for (uint32_t hsv_tile = 0; hsv_tile < num_hsv_tiles; ++hsv_tile) { const uint hsv_offset = (hsv_tile * row_split + gl_SubgroupID) * 16; @@ -476,8 +483,6 @@ void main() { if (hsv_offset < HSV_pad) { [[unroll]] for (uint32_t bc_chunk = 0; bc_chunk < Bc / MatBc; ++bc_chunk) { - coopMatLoad(KMat, Psh, bc_chunk * MatBc * psh_stride, psh_stride, gl_CooperativeMatrixLayoutColumnMajor); - if (SHMEM_STAGING == 0) { if (!USE_DECODE_V && !KV_bounds_check) { // F16/BF16 values can be loaded directly from global memory @@ -493,7 +498,7 @@ void main() { coopMatLoad(QMat, kvsh, v_tile_offset, kvsh_stride, gl_CooperativeMatrixLayoutRowMajor); } - PVMat = coopMatMulAdd(KMat, QMat, PVMat); + PVMat = coopMatMulAdd(PMat[bc_chunk], QMat, PVMat); } // Store PVMat to pvsh and load into Of From 11eaefa454042843b86d739fa8c090c74ad83ee3 Mon Sep 17 00:00:00 2001 From: Nathan Wilson Date: Thu, 30 Jul 2026 01:23:34 +0000 Subject: [PATCH 46/68] vulkan: store coopmat1 FA Psh query-major so the GEMM2 A load vectorizes Psh held P as [kv][query]. The GEMM2 UseA load therefore had to request ColumnMajor, and RADV only attaches an alignment hint to the internal column-major case, which for UseA means the RowMajor request. The load was emitted as 16 separate 16-bit shared reads per fragment. Store P as [query][kv] instead and request RowMajor. The producer now writes four scalar components rather than one vec4; those go to disjoint bytes, so there is no read-modify-write and no partially written vec4. Also updates the host shared-memory estimator, which mirrored the old stride and would otherwise disagree with the shader. Perf-neutral on gfx1151 (within 1% at hd128 and hd256), but it shrinks Psh: LDS 16384 -> 15360 B and code size 13896 -> 13596 at hd128, VGPRs unchanged at 96 with 0 spilled. Kept because it is free and removes a scalar shared-memory access pattern. FLASH_ATTN_EXT suite 5105/5105. Assisted-by: Claude Fable 5 --- ggml/src/ggml-vulkan/ggml-vulkan.cpp | 4 ++-- .../vulkan-shaders/flash_attn_cm1.comp | 17 ++++++++++++----- 2 files changed, 14 insertions(+), 7 deletions(-) diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp index acd8c03bf62..869b9f655b7 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp +++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp @@ -11393,8 +11393,8 @@ static bool ggml_vk_flash_attn_coopmat_shmem_support(const vk_device& device, co const uint32_t qstride = hsk_pad / 4 + 2; const uint32_t Qf = Br * qstride * f16vec4; - const uint32_t psh_stride = Br / 4 + 2; - const uint32_t Psh = Bc * psh_stride * f16vec4; + const uint32_t psh_stride = Bc / 4 + 2; + const uint32_t Psh = Br * psh_stride * f16vec4; const uint32_t sfshstride = (hsk <= 128) ? (Br + 8) : Br; const uint32_t sfsh = Bc * sfshstride * acctype; 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 695d459d0a4..195901a7d90 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_cm1.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_cm1.comp @@ -44,8 +44,11 @@ shared float tmpsh[row_split]; const uint32_t qstride = HSK_pad / 4 + 2; shared FLOAT_TYPEV4 Qf[Br * qstride]; -const uint psh_stride = Br / 4 + 2; -shared FLOAT_TYPEV4 Psh[Bc * psh_stride]; +// P is stored query-major with KV contiguous so the GEMM2 UseA load can be RowMajor: +// RADV flips the requested layout for UseA, and only the resulting internal column-major +// case gets an alignment hint, which is what lets the 16 element loads vectorize. +const uint psh_stride = Bc / 4 + 2; +shared FLOAT_TYPEV4 Psh[Br * psh_stride]; // Avoid padding for hsk==256 to make it fit in 48KB shmem. const uint32_t sfshstride = (HSK <= 128) ? (Br / 4 + 2) : Br / 4; @@ -382,15 +385,19 @@ void main() { [[unroll]] for (uint32_t r = 0; r < rows_per_thread; r += 4) { const uint row = tile_row(r); + const uint pcol_vec = col / 4; + const uint pcol_comp = col % 4; if (KV_bounds_check && j * Bc + col >= KV) { - Psh[col * psh_stride + row / 4] = FLOAT_TYPEV4(0.0f); + [[unroll]] for (uint32_t vec_idx = 0; vec_idx < 4; ++vec_idx) { + Psh[(row + vec_idx) * psh_stride + pcol_vec][pcol_comp] = FLOAT_TYPE(0.0f); + } } else { const vec4 mfvec = vec4(Mf[r], Mf[r + 1], Mf[r + 2], Mf[r + 3]); const FLOAT_TYPEV4 Pf = FLOAT_TYPEV4(exp(vec4(sfsh[row / 4 + col * sfshstride]) - mfvec)); [[unroll]] for (uint32_t vec_idx = 0; vec_idx < 4; ++vec_idx) { Lf[r + vec_idx] += Pf[vec_idx]; + Psh[(row + vec_idx) * psh_stride + pcol_vec][pcol_comp] = Pf[vec_idx]; } - Psh[col * psh_stride + row / 4] = Pf; } } } @@ -427,7 +434,7 @@ void main() { // for every hsv_tile. Load them once instead of re-reading LDS per tile. coopmat PMat[Bc / MatBc]; [[unroll]] for (uint32_t bc_chunk = 0; bc_chunk < Bc / MatBc; ++bc_chunk) { - coopMatLoad(PMat[bc_chunk], Psh, bc_chunk * MatBc * psh_stride, psh_stride, gl_CooperativeMatrixLayoutColumnMajor); + coopMatLoad(PMat[bc_chunk], Psh, bc_chunk * (MatBc / 4), psh_stride, gl_CooperativeMatrixLayoutRowMajor); } // Each subgroup handles HSV/4 columns From bc2b9ad55236fde185f5e9a2f3e25cecdf37c50f Mon Sep 17 00:00:00 2001 From: Nathan Wilson Date: Thu, 30 Jul 2026 01:25:28 +0000 Subject: [PATCH 47/68] vulkan: pin a 32-wide subgroup for coopmat1 FA where narrowing is free get_fa_tuning_params_coopmat1 took device->subgroup_size unconditionally, so on a 64-wide device the coopmat1 FA path always ran wave64. The sibling scalar path already does AMD-specific wave selection; the coopmat1 path never got the equivalent. Narrowing is free exactly when it does not add an iteration to the O-accumulation loop, which runs ceil((HSV/4) / threads_per_rowgroup) per row. On a 64-wide device the test reduces to hsv <= 128. Above it the narrow subgroup issues 1.5x to 1.8x the instructions for the same SIMD passes and hd256 measures 6 to 18 percent slower, so the rule declines. Gated behind GGML_VK_FA_WAVE32 (=1 rule, =2 forces the pin regardless of head size, diagnostic only). Off by default. Measured on gfx1151 (RADV), model-level pp2048, Qwen3-Coder-30B-A3B: d0 +2.5%, d8192 +8.4%, d16384 +10.1%, d32768 +11.3%. Op-level across head sizes, largest where wave64 wastes the most lanes: hsv=64 +12.3%, hsv=96 +7.4%, hsv=128 +6.5%, hsv=256 correctly declined. Full test-backend-ops suite 15884/15884 at =0, =1 and =2; the pin was confirmed to engage in-band via pipeline VGPR statistics rather than inferred from the env var being set. Assisted-by: Claude Fable 5 --- ggml/src/ggml-vulkan/ggml-vulkan.cpp | 33 ++++++++++++++++++++++++++++ 1 file changed, 33 insertions(+) diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp index 869b9f655b7..e4796f12a78 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp +++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp @@ -4005,8 +4005,41 @@ static vk_fa_tuning_params get_fa_tuning_params_coopmat1(const vk_device& device result.block_cols = coopmat_block_cols * num_subgroups; result.row_split = num_subgroups; result.subgroup_size = device->subgroup_size; + + // Pin a 32-wide subgroup only where narrowing is free. The shader derives cols_per_iter, + // threads_per_rowgroup and every strided load loop from gl_WorkGroupSize.x, and + // workgroup_size is num_subgroups * subgroup_size, so threads_per_rowgroup always equals the + // real subgroup size. Halving the subgroup halves the workgroup, and the per-lane O state + // grows as d_per_thread = ceil((HSV/4) / threads_per_rowgroup). Pin only when that count is + // unchanged. Above that point the narrow subgroup issues roughly 1.5x to 1.8x the + // instructions for the same number of SIMD passes, which loses on an issue-bound kernel: + // hd256 measures 6 to 18 percent slower. The test depends on HSV only; HSK does not enter + // d_per_thread. On a 64-wide device it reduces exactly to hsv <= 128. + // =1 applies the rule; =2 forces the pin regardless of head size, for measuring the + // configurations the rule rejects. Diagnostic only. + static const int fa_wave32 = [] { + const char * e = getenv("GGML_VK_FA_WAVE32"); + return e ? atoi(e) : 0; + }(); + if (fa_wave32 != 0 && + device->subgroup_size_control && + 32 < device->subgroup_size && // narrow only, never widen + device->subgroup_min_size <= 32 && 32 <= device->subgroup_max_size && + (result.block_cols % 32) == 0 && // cols_per_thread stays >= 1 + (result.block_cols * result.block_rows / 4) >= num_subgroups * 32 && // mask_cache != 0 + (fa_wave32 == 2 || + CEIL_DIV(hsv / 4, 32u) == CEIL_DIV(hsv / 4, device->subgroup_size))) { + result.subgroup_size = 32; + } + result.workgroup_size = num_subgroups * result.subgroup_size; + // threads_per_rowgroup == the real subgroup size is load-bearing in three places: + // the subgroupMax row reduction, the subgroupAdd of Lf, and tmpsh[gl_SubgroupID], which is + // sized by row_split and would be written out of bounds if gl_NumSubgroups exceeded it. + GGML_ASSERT(result.workgroup_size == result.row_split * result.subgroup_size); + GGML_ASSERT(result.block_cols % result.subgroup_size == 0); + const uint32_t D_lsb = D ^ (D & (D-1)); // extract lowest set bit result.d_split = std::min(std::min(result.subgroup_size, 8u), D_lsb / 4); From 1556cd12dc1d2e148700e01d9f494b949608f1da Mon Sep 17 00:00:00 2001 From: Nathan Wilson Date: Sun, 30 Aug 2026 10:17:20 +0000 Subject: [PATCH 48/68] vulkan: enable the coopmat1 FA wave32 narrowing rule by default Split from the fork's flag-flip commit (f7d804ee7): GGML_VK_FA_WAVE32 now defaults to 1 (apply the HSV-based narrowing rule); =0 opts out and =2 keeps its diagnostic force meaning. The subgroup_size_control device guard is unchanged, so devices without a 32-wide subgroup mode are unaffected. Co-Authored-By: Claude Fable 5 Assisted-by: Claude (Opus 5) --- ggml/src/ggml-vulkan/ggml-vulkan.cpp | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp index e4796f12a78..ee35ba7c5a9 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp +++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp @@ -4015,11 +4015,11 @@ static vk_fa_tuning_params get_fa_tuning_params_coopmat1(const vk_device& device // instructions for the same number of SIMD passes, which loses on an issue-bound kernel: // hd256 measures 6 to 18 percent slower. The test depends on HSV only; HSK does not enter // d_per_thread. On a 64-wide device it reduces exactly to hsv <= 128. - // =1 applies the rule; =2 forces the pin regardless of head size, for measuring the - // configurations the rule rejects. Diagnostic only. + // On by default (=1, applies the rule); =0 disables. =2 forces the pin regardless of + // head size, for measuring the configurations the rule rejects. Diagnostic only. static const int fa_wave32 = [] { const char * e = getenv("GGML_VK_FA_WAVE32"); - return e ? atoi(e) : 0; + return e ? atoi(e) : 1; }(); if (fa_wave32 != 0 && device->subgroup_size_control && From e75a22bf1bca22f64b458e0b05dd8eb9a0cbf25e Mon Sep 17 00:00:00 2001 From: Nathan Wilson Date: Tue, 14 Jul 2026 05:53:12 +0000 Subject: [PATCH 49/68] vulkan: mul_mat_id per-expert-n tile selection (Stage 2a, env-gated) Select the matmul_id tile by expected per-expert token count (nei1*nei0/n_expert) instead of aggregate nei1, which always picked the widest tile and left most N-lanes empty at MoE prefill (~16 rows/expert vs BN=64). Gated on GGML_VK_MMID_SMALLN=1, default off pending cross-model validation. Was null (+0.7%) standalone pre-row-lists: smaller tiles meant more workgroups each re-paying the per-WG id scan. With the Stage 1 row-list prepass the scan is gone and the occupancy win materializes. Strix Halo (RADV gfx1151), Qwen3.6-35B-A3B UD-Q5_K_XL, clean window, pp512 4-way (rowlists x smalln): 914.7 / 937.9 / 1005.3 / 1063.6 t/s (combined +16.3% vs baseline). MUL_MAT_ID q5_K 2.33->4.43 TFLOPS (1.91x), q6_K 1.99->3.57 (1.79x) cumulative. Hot pipelines verified via probe: matmul_id_*_f32 (y_f32 path, quantize_y does not engage for these quants); smalln flips tile _m -> _s. 790/790 MUL_MAT_ID both rowlists configs. Co-Authored-By: Claude Fable 5 Assisted-by: Claude (Opus 5) --- ggml/src/ggml-vulkan/ggml-vulkan.cpp | 16 +++++++++++++--- 1 file changed, 13 insertions(+), 3 deletions(-) diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp index ee35ba7c5a9..74372580bc3 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp +++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp @@ -10844,7 +10844,18 @@ static void ggml_vk_mul_mat_id_q_f16(ggml_backend_vk_context * ctx, vk_context& const ggml_type effective_src1_type = quantize_y ? GGML_TYPE_Q8_1 : (y_f32_kernel ? GGML_TYPE_F32 : src1->type); - const uint32_t kpad = quantize_y ? 0 : ggml_vk_align_size(ne10, ggml_vk_guess_matmul_id_pipeline_align(ctx, mmp, ne01, nei1, qx_needs_dequant ? f16_type : src0->type, effective_src1_type)); + // EXPERIMENT (GGML_VK_MMID_SMALLN=1): select the matmul tile by the EXPECTED PER-EXPERT + // token count rather than the whole batch. With E experts and nei0 active per token, each + // expert sees ~nei1*nei0/E rows; selecting by aggregate nei1 always picks the widest tile + // and leaves most N-lanes empty at MoE prefill (measured: MUL_MAT_ID at ~20% of dense + // matmul efficiency, ~78% of MoE prefill time on Qwen3.6-35B-A3B). + uint32_t n_for_tile = (uint32_t)nei1; + static const char * mmid_smalln_env = getenv("GGML_VK_MMID_SMALLN"); + if (mmid_smalln_env && atoi(mmid_smalln_env) != 0 && ne02 > 1) { + n_for_tile = std::max(1u, (uint32_t)((nei1 * nei0 + ne02 - 1) / ne02)); + } + + const uint32_t kpad = quantize_y ? 0 : ggml_vk_align_size(ne10, ggml_vk_guess_matmul_id_pipeline_align(ctx, mmp, ne01, n_for_tile, qx_needs_dequant ? f16_type : src0->type, effective_src1_type)); // Coopmat2 MUL_MAT_ID BK specialization constants in ggml_vk_load_shaders are at most 64. const uint32_t y_staged_row_stride = ctx->device->coopmat2 && !quantize_y ? ggml_vk_align_size(ne10, 64) : ne10; const bool y_needs_k_padding = ne10 != y_staged_row_stride; @@ -10853,10 +10864,9 @@ static void ggml_vk_mul_mat_id_q_f16(ggml_backend_vk_context * ctx, vk_context& // Not implemented GGML_ASSERT(y_needs_reformat || !qy_needs_dequant); // NOLINT - const bool aligned = !quantize_y && ne10 == kpad && ne01 > 8 && nei1 > 8; - vk_pipeline pipeline = ggml_vk_guess_matmul_id_pipeline(ctx, mmp, ne01, nei1, aligned, qx_needs_dequant ? f16_type : src0->type, effective_src1_type); + vk_pipeline pipeline = ggml_vk_guess_matmul_id_pipeline(ctx, mmp, ne01, n_for_tile, aligned, qx_needs_dequant ? f16_type : src0->type, effective_src1_type); if (ggml_nbytes(src0) > ctx->device->properties.limits.maxStorageBufferRange) { pipeline = ggml_vk_get_64b_indexing_pipeline(ctx, pipeline); From 33a08e893d996ffd1c75883a37edb4ecda13a9e9 Mon Sep 17 00:00:00 2001 From: Nathan Wilson Date: Tue, 14 Jul 2026 06:04:18 +0000 Subject: [PATCH 50/68] vulkan: mul_mat_id small-tile shape probes (env-gated) Two env-gated overrides of the mmid small-tile config (KHR coopmat branch only, scoped to the mul_mat_id quant pipelines; dense untouched): - GGML_VK_MMID_TILE16: BN/WN 32->16. NEGATIVE (-3.8% e2e): experts with n_e>16 split into two column tiles and re-stream their full weight matrix; A-traffic scales with sum(ceil(n_e/BN)), so BN must not drop below the mean per-expert n. Kept as a documented dead end. - GGML_VK_MMID_BM64: BM 32->64, BLOCK_SIZE 64->128 (two warps). +1.3% e2e on top of rowlists+smalln: same A-traffic, half the ir-tiles so half the per-expert B re-reads, larger WGs hide latency. Strix Halo, Qwen3.6-35B-A3B UD-Q5_K_XL pp512, clean window, all on top of RL+SN control 1065.4 +/- 3.3: BM64 1078.9 +/- 3.5, TILE16 1024.6, BM64+TILE16 1007.5. Cumulative vs pre-Stage-1 baseline: 914.7 -> 1078.9 (+17.9%). 790/790 MUL_MAT_ID for all configs. Co-Authored-By: Claude Fable 5 Assisted-by: Claude (Opus 5) --- ggml/src/ggml-vulkan/ggml-vulkan.cpp | 28 ++++++++++++++++++++++++++++ 1 file changed, 28 insertions(+) diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp index 74372580bc3..f1211d67a25 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp +++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp @@ -5107,6 +5107,34 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { } #endif + // EXPERIMENT (GGML_VK_MMID_TILE16=1): narrow the matmul_id small tile to BN=16. + // MoE prefill leaves ~nei1*nei0/n_expert rows per expert (~16 at ub512/256E/8a), + // so even the 32-wide small tile runs half empty. Only meaningful stacked on + // GGML_VK_MMID_SMALLN=1 (routes mmid to the small tile) + the row-list prepass. + // Shadows the s-tile config for the mmid quant pipelines only; dense unaffected. + auto s_warptile_mmq_id16 = s_warptile_mmq; + auto s_mmq_wg_denoms_id16 = s_mmq_wg_denoms; + { + const char * tile16_env = getenv("GGML_VK_MMID_TILE16"); + if (tile16_env && atoi(tile16_env) != 0) { + s_warptile_mmq_id16[2] = 16; // BN + s_warptile_mmq_id16[5] = 16; // WN + s_mmq_wg_denoms_id16[1] = 16; + } + // GGML_VK_MMID_BM64=1: taller small tile (BM 32->64, two warps). Same + // A-traffic (ic-tile count unchanged), halves per-expert B re-reads + // (ir-tile count), larger WGs to hide latency. + const char * bm64_env = getenv("GGML_VK_MMID_BM64"); + if (bm64_env && atoi(bm64_env) != 0) { + s_warptile_mmq_id16[0] = 2 * mul_mat_subgroup_size; // BLOCK_SIZE + s_warptile_mmq_id16[1] = 64; // BM + s_mmq_wg_denoms_id16[0] = 64; + } + } + { + const auto &s_warptile_mmq = s_warptile_mmq_id16; + const auto &s_mmq_wg_denoms = s_mmq_wg_denoms_id16; + CREATE_MM2(GGML_TYPE_Q1_0, pipeline_dequant_mul_mat_mat_id[GGML_TYPE_Q1_0], matmul_id_subgroup_q1_0_f32, mmq_wg_denoms, warptile_mmq, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); CREATE_MM2(GGML_TYPE_Q2_0, pipeline_dequant_mul_mat_mat_id[GGML_TYPE_Q2_0], matmul_id_subgroup_q2_0_f32, mmq_wg_denoms, warptile_mmq, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); CREATE_MM2(GGML_TYPE_Q4_0, pipeline_dequant_mul_mat_mat_id[GGML_TYPE_Q4_0], matmul_id_subgroup_q4_0_f32, mmq_wg_denoms, warptile_mmq, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); From fe10c7f2ae9f6a7fff98ae38e456b226fedcf28a Mon Sep 17 00:00:00 2001 From: Nathan Wilson Date: Tue, 14 Jul 2026 08:42:07 +0000 Subject: [PATCH 51/68] vulkan: mul_mat_id taller medium tile probe (GGML_VK_MMID_M128, env-gated) Extends the BM64 idea to the medium tile: GGML_VK_MMID_M128=1 raises the mmid m-tile to BM=128 / BLOCK_SIZE=4*subgroup (four warps), same A-traffic, half the ir-tiles and per-expert B re-reads. The medium tile is what the per-expert-n heuristic selects at n~64 (e.g. ub2048 on 256-expert/top-8 models, or ub1024 at 128 experts). Strix Halo, Qwen3.6-35B-A3B UD-Q5_K_XL pp2048, drain-verified window: ub2048 stack 1039.9 +/- 1.1 -> +M128 1091.6 +/- 2.6 (+5.0%); inert in the s-tile regime (ub1024 1137.2 vs 1147.9 ref). ub1024 remains the throughput sweet spot for this model. 790/790 MUL_MAT_ID. Also measured this session, NOT kept: caching counts+row lists across the three expert matmuls of a layer (they share one ids tensor) was correctness-clean but perf-null (1066.3 vs 1070.9 baseline) - the prepass dispatches and barriers are already free on this queue; reverted rather than carry the invalidation surface. Co-Authored-By: Claude Fable 5 Assisted-by: Claude (Opus 5) --- ggml/src/ggml-vulkan/ggml-vulkan.cpp | 12 ++++++++++++ 1 file changed, 12 insertions(+) diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp index f1211d67a25..79797c69525 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp +++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp @@ -5114,6 +5114,8 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { // Shadows the s-tile config for the mmid quant pipelines only; dense unaffected. auto s_warptile_mmq_id16 = s_warptile_mmq; auto s_mmq_wg_denoms_id16 = s_mmq_wg_denoms; + auto m_warptile_mmq_id128 = m_warptile_mmq; + auto m_mmq_wg_denoms_id128 = m_mmq_wg_denoms; { const char * tile16_env = getenv("GGML_VK_MMID_TILE16"); if (tile16_env && atoi(tile16_env) != 0) { @@ -5130,10 +5132,20 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { s_warptile_mmq_id16[1] = 64; // BM s_mmq_wg_denoms_id16[0] = 64; } + // GGML_VK_MMID_M128=1: same idea for the medium tile (BM 64->128, four + // warps) — the tile the per-expert-n heuristic picks at n~64 (e.g. ub2048). + const char * m128_env = getenv("GGML_VK_MMID_M128"); + if (m128_env && atoi(m128_env) != 0) { + m_warptile_mmq_id128[0] = 4 * mul_mat_subgroup_size; // BLOCK_SIZE + m_warptile_mmq_id128[1] = 128; // BM + m_mmq_wg_denoms_id128[0] = 128; + } } { const auto &s_warptile_mmq = s_warptile_mmq_id16; const auto &s_mmq_wg_denoms = s_mmq_wg_denoms_id16; + const auto &m_warptile_mmq = m_warptile_mmq_id128; + const auto &m_mmq_wg_denoms = m_mmq_wg_denoms_id128; CREATE_MM2(GGML_TYPE_Q1_0, pipeline_dequant_mul_mat_mat_id[GGML_TYPE_Q1_0], matmul_id_subgroup_q1_0_f32, mmq_wg_denoms, warptile_mmq, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); CREATE_MM2(GGML_TYPE_Q2_0, pipeline_dequant_mul_mat_mat_id[GGML_TYPE_Q2_0], matmul_id_subgroup_q2_0_f32, mmq_wg_denoms, warptile_mmq, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); From 07e35febaf3182611c849e980e96348d5297b1d8 Mon Sep 17 00:00:00 2001 From: Nathan Wilson Date: Tue, 14 Jul 2026 11:20:31 +0000 Subject: [PATCH 52/68] vulkan: mmid wave32 probe (GGML_VK_MMID_WAVE32, env-gated) Force required subgroup size 32 on the KHR-coopmat mmid quant pipelines (pipeline_dequant_mul_mat_mat_id[*] only; dense untouched). Hypothesis: RDNA3.5 WMMA is wave32-native, so RADV may lower KHR_coopmat better at wave32 than at the reported default of 64. Mechanics: the cm1 path in mul_mm.comp derives its warp grid from the real subgroup (warp_i = gl_SubgroupID, tiw = gl_SubgroupInvocationID) while NUM_WARPS = BLOCK_SIZE/WARP (spec constants) sizes coopmat_stage[] and ballots_sh[] and warp_r/warp_c assume NUM_WARPS == (BM/WM)*(BN/WN). Forcing sg32 with WARP=64 would over-run those shared arrays and leave warp_c outside the tile, so the gate (a) sets WARP=32 in shadowed copies of the s/m/l mmid warptiles (composes after the BM64/M128 shadows) and (b) halves WM (or WN, keeping WM>=TM, WN>=TN) until the doubled subgroup count exactly tiles BM x BN again, asserting both invariants. BLOCK_SIZE is kept, so workgroup shape, load loops and shmem match the wave64 stack; per-lane accumulator footprint is also unchanged (half the lanes per subgroup, half the (WM/TM)*(WN/TN) fragments). The required size is passed via a scoped CREATE_MM redefinition adding a trailing required_subgroup_size arg (fp16-branch pattern), gated on subgroup_size_control covering 32 since ggml_vk_create_pipeline_func silently drops the required size otherwise. Correctness (test-backend-ops -o MUL_MAT_ID -b Vulkan0): 790/790 plain, 790/790 WAVE32=1, 790/790 WAVE32+SMALLN+BM64. Perf, Qwen3.6-35B-A3B-UD-Q5_K_XL, fa=1 b/ub=512 ctk/ctv=q8_0 pp512 r=3, atomic window, canary clean: stack (SMALLN+BM64), wave64 default: 1076.61 +/- 3.71 t/s stack + WAVE32: 1106.73 +/- 2.52 t/s (+2.8%) Co-Authored-By: Claude Fable 5 Assisted-by: Claude (Opus 5) --- ggml/src/ggml-vulkan/ggml-vulkan.cpp | 58 ++++++++++++++++++++++++++++ 1 file changed, 58 insertions(+) diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp index 79797c69525..e8bb9555ed0 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp +++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp @@ -5116,6 +5116,8 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { auto s_mmq_wg_denoms_id16 = s_mmq_wg_denoms; auto m_warptile_mmq_id128 = m_warptile_mmq; auto m_mmq_wg_denoms_id128 = m_mmq_wg_denoms; + auto l_warptile_mmq_idw = l_warptile_mmq; + uint32_t mmid_req_sgs = 0; { const char * tile16_env = getenv("GGML_VK_MMID_TILE16"); if (tile16_env && atoi(tile16_env) != 0) { @@ -5140,12 +5142,68 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { m_warptile_mmq_id128[1] = 128; // BM m_mmq_wg_denoms_id128[0] = 128; } + // GGML_VK_MMID_WAVE32=1: force required subgroup size 32 on the mmid + // quant coopmat pipelines (RDNA3.x WMMA is wave32-native; probe whether + // RADV lowers KHR_coopmat better at wave32). The cm1 path derives its + // warp grid from the real subgroup (warp_i = gl_SubgroupID, tiw = + // gl_SubgroupInvocationID) and sizes shared arrays (coopmat_stage, + // ballots_sh) with NUM_WARPS = BLOCK_SIZE / WARP, so the WARP spec + // constant must equal the forced size, and the coverage invariant + // NUM_WARPS == (BM/WM)*(BN/WN) must be restored. BLOCK_SIZE is kept + // (same workgroup shape, load loops and shmem as the wave64 stack), so + // the subgroup count doubles and WM (or WN) is halved until the warp + // grid exactly tiles BM x BN again. Per-lane coopmat accumulator + // footprint is unchanged: half the lanes per subgroup, half the + // (WM/TM)*(WN/TN) fragments per subgroup. Runs after the BM64/M128 + // gates so it composes with the probe stack. Applies only when the + // driver honors a required size (subgroup_size_control covering 32); + // otherwise WARP=32 with a real subgroup of 64 would corrupt tiling. + const char * wave32_env = getenv("GGML_VK_MMID_WAVE32"); + if (wave32_env && atoi(wave32_env) != 0 && device->subgroup_size_control && + device->subgroup_min_size <= 32 && 32 <= device->subgroup_max_size) { + mmid_req_sgs = 32; + auto wave32_tile = [](std::vector &w) { + // {BLOCK_SIZE, BM, BN, BK, WM, WN, WMITER, TM, TN, TK, WARP} + w[10] = 32; // WARP: must match the forced subgroup size + for (int guard = 0; guard < 4 && w[0] / w[10] != (w[1] / w[4]) * (w[2] / w[5]); ++guard) { + if (w[4] >= w[5] && w[4] > w[7]) { + w[4] /= 2; // halve WM, keeping WM >= TM + } else { + w[5] /= 2; // halve WN + } + } + GGML_ASSERT(w[0] / w[10] == (w[1] / w[4]) * (w[2] / w[5])); // NUM_WARPS == (BM/WM)*(BN/WN) + GGML_ASSERT(w[4] >= w[7] && w[5] >= w[8]); // WM >= TM, WN >= TN + }; + wave32_tile(s_warptile_mmq_id16); + wave32_tile(m_warptile_mmq_id128); + wave32_tile(l_warptile_mmq_idw); + } } { const auto &s_warptile_mmq = s_warptile_mmq_id16; const auto &s_mmq_wg_denoms = s_mmq_wg_denoms_id16; const auto &m_warptile_mmq = m_warptile_mmq_id128; const auto &m_mmq_wg_denoms = m_mmq_wg_denoms_id128; + const auto &l_warptile_mmq = l_warptile_mmq_idw; + + // Same expansion as CREATE_MM above, plus a trailing required subgroup + // size (0 = driver default) for the GGML_VK_MMID_WAVE32 probe. Scoped to + // the mmid quant pipelines below; dense pipelines are untouched. +#undef CREATE_MM +#define CREATE_MM(TYPE, PIPELINE_NAME, NAMELC, F16ACC, WG_DENOMS, WARPTILE, PUSHCONST, PARAMCOUNT, ID) \ + if (device->mul_mat ## ID ## _l[TYPE]) \ + ggml_vk_create_pipeline(device, device-> PIPELINE_NAME ->l, #NAMELC #F16ACC "_l", NAMELC ## F16ACC ## _cm1_len, NAMELC ## F16ACC ## _cm1_data, "main", PARAMCOUNT, sizeof(PUSHCONST), l_ ## WG_DENOMS, ggml_vk_mul_mm_spec(l_ ## WARPTILE, false), 1, false, true, mmid_req_sgs); \ + if (device->mul_mat ## ID ## _m[TYPE]) \ + ggml_vk_create_pipeline(device, device-> PIPELINE_NAME ->m, #NAMELC #F16ACC "_m", NAMELC ## F16ACC ## _cm1_len, NAMELC ## F16ACC ## _cm1_data, "main", PARAMCOUNT, sizeof(PUSHCONST), m_ ## WG_DENOMS, ggml_vk_mul_mm_spec(m_ ## WARPTILE, false), 1, false, true, mmid_req_sgs); \ + if (device->mul_mat ## ID ## _s[TYPE]) \ + ggml_vk_create_pipeline(device, device-> PIPELINE_NAME ->s, #NAMELC #F16ACC "_s", NAMELC ## F16ACC ## _cm1_len, NAMELC ## F16ACC ## _cm1_data, "main", PARAMCOUNT, sizeof(PUSHCONST), s_ ## WG_DENOMS, ggml_vk_mul_mm_spec(s_ ## WARPTILE, false), 1, false, true, mmid_req_sgs); \ + if (device->mul_mat ## ID ## _l[TYPE]) \ + ggml_vk_create_pipeline(device, device-> PIPELINE_NAME ->a_l, #NAMELC #F16ACC "_aligned_l", NAMELC ## F16ACC ## _cm1_len, NAMELC ## F16ACC ## _cm1_data, "main", PARAMCOUNT, sizeof(PUSHCONST), l_ ## WG_DENOMS, ggml_vk_mul_mm_spec(l_ ## WARPTILE, true), l_align, false, true, mmid_req_sgs); \ + if (device->mul_mat ## ID ## _m[TYPE]) \ + ggml_vk_create_pipeline(device, device-> PIPELINE_NAME ->a_m, #NAMELC #F16ACC "_aligned_m", NAMELC ## F16ACC ## _cm1_len, NAMELC ## F16ACC ## _cm1_data, "main", PARAMCOUNT, sizeof(PUSHCONST), m_ ## WG_DENOMS, ggml_vk_mul_mm_spec(m_ ## WARPTILE, true), m_align, false, true, mmid_req_sgs); \ + if (device->mul_mat ## ID ## _s[TYPE]) \ + ggml_vk_create_pipeline(device, device-> PIPELINE_NAME ->a_s, #NAMELC #F16ACC "_aligned_s", NAMELC ## F16ACC ## _cm1_len, NAMELC ## F16ACC ## _cm1_data, "main", PARAMCOUNT, sizeof(PUSHCONST), s_ ## WG_DENOMS, ggml_vk_mul_mm_spec(s_ ## WARPTILE, true), s_align, false, true, mmid_req_sgs); \ CREATE_MM2(GGML_TYPE_Q1_0, pipeline_dequant_mul_mat_mat_id[GGML_TYPE_Q1_0], matmul_id_subgroup_q1_0_f32, mmq_wg_denoms, warptile_mmq, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); CREATE_MM2(GGML_TYPE_Q2_0, pipeline_dequant_mul_mat_mat_id[GGML_TYPE_Q2_0], matmul_id_subgroup_q2_0_f32, mmq_wg_denoms, warptile_mmq, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); From 18569a90308efad8e597ae15721eb9de704f692c Mon Sep 17 00:00:00 2001 From: Nathan Wilson Date: Tue, 14 Jul 2026 11:29:58 +0000 Subject: [PATCH 53/68] vulkan: mul_mat_id f16-B probe (GGML_VK_MMID_F16B, env-gated) Convert contiguous f32 activations (B) to f16 for quantized MUL_MAT_ID on KHR_coopmat devices and run the matmul_id_subgroup__f16 kernels instead of the f32-B ones. Halves B bytes and buf_b shared memory. The _f16 SPIR-V already exists for every quant type; this adds a parallel env-gated pipeline array (pipeline_dequant_mul_mat_mat_id_f16b), extends the getter to return it for src1=F16 on non-coopmat2 devices, relaxes the src1-type assert, and forces the existing y_non_contig convert-to- prealloc_y plumbing (same as coopmat2). Default OFF, zero behavior change when unset. Measured on Radeon 8060S (RADV gfx1151), Qwen3.6-35B-A3B-UD-Q5_K_XL, -fa 1 -b 512 -ub 512 -ctk q8_0 -ctv q8_0 -p 512 -r 3, stacked on GGML_VK_MMID_SMALLN=1 GGML_VK_MMID_BM64=1: stack (f32 B, canary): pp512 1075.04 +/- 8.66 t/s stack + F16B: pp512 1100.66 +/- 2.37 t/s (+2.4%) test-backend-ops test -o MUL_MAT_ID -b Vulkan0: 790/790 with and without GGML_VK_MMID_F16B=1. Co-Authored-By: Claude Fable 5 Assisted-by: Claude (Opus 5) --- ggml/src/ggml-vulkan/ggml-vulkan.cpp | 72 +++++++++++++++++++++++++++- 1 file changed, 70 insertions(+), 2 deletions(-) diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp index e8bb9555ed0..f7bf2ac664a 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp +++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp @@ -955,6 +955,10 @@ struct vk_device_struct { vk_matmul_pipeline2 pipeline_dequant_mul_mat_mat_id[GGML_TYPE_COUNT]; vk_matmul_pipeline2 pipeline_dequant_mul_mat_mat_id_q8_1[GGML_TYPE_COUNT]; + // EXPERIMENT (GGML_VK_MMID_F16B=1): f16-B mul_mat_id pipelines on KHR_coopmat + // devices (upstream only builds f32-B there). Populated only when the env flag + // is set; empty otherwise. + vk_matmul_pipeline2 pipeline_dequant_mul_mat_mat_id_f16b[GGML_TYPE_COUNT]; vk_pipeline pipeline_matmul_split_k_reduce; vk_pipeline pipeline_quantize_q8_1_x4; @@ -4454,6 +4458,18 @@ struct CompileTask { uint32_t required_subgroup_size; }; +// EXPERIMENT (GGML_VK_MMID_F16B=1): convert the contiguous f32 activations (B) of +// quantized MUL_MAT_ID to f16 and run the f16-B matmul_id kernels instead of the +// f32-B ones. Halves B bytes and buf_b shared memory (better occupancy) at the +// f32->f16 rounding cost upstream already accepts on the coopmat2 path. +static bool ggml_vk_mmid_f16b_enabled() { + static const bool enabled = [] { + const char * env = getenv("GGML_VK_MMID_F16B"); + return env != nullptr && atoi(env) != 0; + }(); + return enabled; +} + static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { VK_LOG_DEBUG("ggml_vk_load_shaders(" << device->name << ")"); @@ -5242,6 +5258,37 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { CREATE_MM2(GGML_TYPE_MXFP4, pipeline_dequant_mul_mat_mat_id[GGML_TYPE_MXFP4], matmul_id_subgroup_mxfp4_f32, mmq_wg_denoms, warptile_mmq, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); CREATE_MM2(GGML_TYPE_NVFP4, pipeline_dequant_mul_mat_mat_id[GGML_TYPE_NVFP4], matmul_id_subgroup_nvfp4_f32, mmq_wg_denoms, warptile_mmq, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); } + + // EXPERIMENT (GGML_VK_MMID_F16B=1): f16-B variants of the quant matmul_id + // pipelines. The _f16 SPIR-V exists for every quant type; upstream just never + // instantiates it in the KHR_coopmat branch. Same warptiles as the f32-B lines + // (including the SMALLN/BM64/M128 tile experiments shadowed above). + if (ggml_vk_mmid_f16b_enabled()) { + CREATE_MM2(GGML_TYPE_Q1_0, pipeline_dequant_mul_mat_mat_id_f16b[GGML_TYPE_Q1_0], matmul_id_subgroup_q1_0_f16, mmq_wg_denoms, warptile_mmq, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); + CREATE_MM2(GGML_TYPE_Q4_0, pipeline_dequant_mul_mat_mat_id_f16b[GGML_TYPE_Q4_0], matmul_id_subgroup_q4_0_f16, mmq_wg_denoms, warptile_mmq, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); + CREATE_MM2(GGML_TYPE_Q4_1, pipeline_dequant_mul_mat_mat_id_f16b[GGML_TYPE_Q4_1], matmul_id_subgroup_q4_1_f16, mmq_wg_denoms, warptile_mmq, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); + CREATE_MM2(GGML_TYPE_Q5_0, pipeline_dequant_mul_mat_mat_id_f16b[GGML_TYPE_Q5_0], matmul_id_subgroup_q5_0_f16, mmq_wg_denoms, warptile_mmq, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); + CREATE_MM2(GGML_TYPE_Q5_1, pipeline_dequant_mul_mat_mat_id_f16b[GGML_TYPE_Q5_1], matmul_id_subgroup_q5_1_f16, mmq_wg_denoms, warptile_mmq, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); + CREATE_MM2(GGML_TYPE_Q8_0, pipeline_dequant_mul_mat_mat_id_f16b[GGML_TYPE_Q8_0], matmul_id_subgroup_q8_0_f16, mmq_wg_denoms, warptile_mmq, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); + CREATE_MM2(GGML_TYPE_Q2_K, pipeline_dequant_mul_mat_mat_id_f16b[GGML_TYPE_Q2_K], matmul_id_subgroup_q2_k_f16, mmq_wg_denoms, warptile_mmq, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); + CREATE_MM2(GGML_TYPE_Q3_K, pipeline_dequant_mul_mat_mat_id_f16b[GGML_TYPE_Q3_K], matmul_id_subgroup_q3_k_f16, mmq_wg_denoms, warptile_mmq, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); + CREATE_MM2(GGML_TYPE_Q4_K, pipeline_dequant_mul_mat_mat_id_f16b[GGML_TYPE_Q4_K], matmul_id_subgroup_q4_k_f16, mmq_wg_denoms, warptile_mmq, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); + CREATE_MM2(GGML_TYPE_Q5_K, pipeline_dequant_mul_mat_mat_id_f16b[GGML_TYPE_Q5_K], matmul_id_subgroup_q5_k_f16, mmq_wg_denoms, warptile_mmq, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); + CREATE_MM2(GGML_TYPE_Q6_K, pipeline_dequant_mul_mat_mat_id_f16b[GGML_TYPE_Q6_K], matmul_id_subgroup_q6_k_f16, mmq_wg_denoms, warptile_mmq, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); + CREATE_MM2(GGML_TYPE_IQ1_S, pipeline_dequant_mul_mat_mat_id_f16b[GGML_TYPE_IQ1_S], matmul_id_subgroup_iq1_s_f16, mmq_wg_denoms, warptile_mmq, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); + CREATE_MM2(GGML_TYPE_IQ1_M, pipeline_dequant_mul_mat_mat_id_f16b[GGML_TYPE_IQ1_M], matmul_id_subgroup_iq1_m_f16, mmq_wg_denoms, warptile_mmq, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); + CREATE_MM2(GGML_TYPE_IQ2_XXS, pipeline_dequant_mul_mat_mat_id_f16b[GGML_TYPE_IQ2_XXS], matmul_id_subgroup_iq2_xxs_f16, mmq_wg_denoms, warptile_mmq, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); + CREATE_MM2(GGML_TYPE_IQ2_XS, pipeline_dequant_mul_mat_mat_id_f16b[GGML_TYPE_IQ2_XS], matmul_id_subgroup_iq2_xs_f16, mmq_wg_denoms, warptile_mmq, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); + CREATE_MM2(GGML_TYPE_IQ2_S, pipeline_dequant_mul_mat_mat_id_f16b[GGML_TYPE_IQ2_S], matmul_id_subgroup_iq2_s_f16, mmq_wg_denoms, warptile_mmq, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); + CREATE_MM2(GGML_TYPE_IQ3_XXS, pipeline_dequant_mul_mat_mat_id_f16b[GGML_TYPE_IQ3_XXS], matmul_id_subgroup_iq3_xxs_f16, mmq_wg_denoms, warptile_mmq, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); + CREATE_MM2(GGML_TYPE_IQ3_S, pipeline_dequant_mul_mat_mat_id_f16b[GGML_TYPE_IQ3_S], matmul_id_subgroup_iq3_s_f16, mmq_wg_denoms, warptile_mmq, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); + CREATE_MM2(GGML_TYPE_IQ4_XS, pipeline_dequant_mul_mat_mat_id_f16b[GGML_TYPE_IQ4_XS], matmul_id_subgroup_iq4_xs_f16, mmq_wg_denoms, warptile_mmq, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); + CREATE_MM2(GGML_TYPE_IQ4_NL, pipeline_dequant_mul_mat_mat_id_f16b[GGML_TYPE_IQ4_NL], matmul_id_subgroup_iq4_nl_f16, mmq_wg_denoms, warptile_mmq, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); + CREATE_MM2(GGML_TYPE_MXFP4, pipeline_dequant_mul_mat_mat_id_f16b[GGML_TYPE_MXFP4], matmul_id_subgroup_mxfp4_f16, mmq_wg_denoms, warptile_mmq, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); + CREATE_MM2(GGML_TYPE_NVFP4, pipeline_dequant_mul_mat_mat_id_f16b[GGML_TYPE_NVFP4], matmul_id_subgroup_nvfp4_f16, mmq_wg_denoms, warptile_mmq, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); + } + } + #undef CREATE_MM2 #undef CREATE_MM } else @@ -8534,7 +8581,8 @@ static vk_matmul_pipeline ggml_vk_get_mul_mat_mat_id_pipeline(ggml_backend_vk_co return pipelines; } - GGML_ASSERT(src1_type == GGML_TYPE_F32 || (ctx->device->coopmat2 && src1_type == GGML_TYPE_F16)); + GGML_ASSERT(src1_type == GGML_TYPE_F32 || + ((ctx->device->coopmat2 || ggml_vk_mmid_f16b_enabled()) && src1_type == GGML_TYPE_F16)); switch (src0_type) { case GGML_TYPE_Q1_0: @@ -8571,7 +8619,11 @@ static vk_matmul_pipeline ggml_vk_get_mul_mat_mat_id_pipeline(ggml_backend_vk_co return nullptr; } - vk_matmul_pipeline2& mmp = ctx->device->pipeline_dequant_mul_mat_mat_id[src0_type]; + // GGML_VK_MMID_F16B: on KHR_coopmat devices the f16-B mmid pipelines live in a + // parallel array (coopmat2's main array already holds f16-B pipelines). + vk_matmul_pipeline2& mmp = (src1_type == GGML_TYPE_F16 && !ctx->device->coopmat2) + ? ctx->device->pipeline_dequant_mul_mat_mat_id_f16b[src0_type] + : ctx->device->pipeline_dequant_mul_mat_mat_id[src0_type]; // XXX TODO 'prec' is not actually allowed in mul_mat_id. bool prefer_fp16acc = ctx->device->fp16 /*&& prec == GGML_PREC_DEFAULT*/; bool support_fp16acc = !mmp.f16acc->is_empty(); @@ -10914,7 +10966,23 @@ static void ggml_vk_mul_mat_id_q_f16(ggml_backend_vk_context * ctx, vk_context& #else const bool y_decode_vector_staging = false; #endif + // EXPERIMENT (GGML_VK_MMID_F16B=1): route quantized MUL_MAT_ID through the f16-B + // kernels. Treating contiguous f32 B as y_non_contig reuses the existing + // convert-to-prealloc_y plumbing (to_fp16_vk_1), exactly like coopmat2 does. + // Gated on coopmat_support because the f16b pipelines are only created there. + const bool mmid_f16b = ggml_vk_mmid_f16b_enabled() && + ctx->device->coopmat_support && !ctx->device->coopmat2 && + ggml_is_quantized(src0->type) && src1->type == GGML_TYPE_F32; + if (mmid_f16b) { + static bool mmid_f16b_logged = false; + if (!mmid_f16b_logged) { + mmid_f16b_logged = true; + fprintf(stderr, "ggml_vulkan: MUL_MAT_ID f16-B path engaged (GGML_VK_MMID_F16B)\n"); + } + } + const bool y_non_contig = y_decode_vector_staging || + mmid_f16b || (ctx->device->coopmat2 && src1->type == GGML_TYPE_F32) || (src0->type == GGML_TYPE_BF16 && src1->type != GGML_TYPE_BF16) || !ggml_vk_dim01_contiguous(src1); From 8ad46e197db665b0204bf6dc36ddab43a354e5e5 Mon Sep 17 00:00:00 2001 From: Nathan Wilson Date: Sun, 26 Jul 2026 10:53:57 +0000 Subject: [PATCH 54/68] vulkan: guard mmid f16-B path on pipeline existence (Q2_0 fallback) Assisted-by: Claude (Opus 5) --- ggml/src/ggml-vulkan/ggml-vulkan.cpp | 6 +++++- 1 file changed, 5 insertions(+), 1 deletion(-) diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp index f7bf2ac664a..a7c7f15b4ea 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp +++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp @@ -10972,7 +10972,11 @@ static void ggml_vk_mul_mat_id_q_f16(ggml_backend_vk_context * ctx, vk_context& // Gated on coopmat_support because the f16b pipelines are only created there. const bool mmid_f16b = ggml_vk_mmid_f16b_enabled() && ctx->device->coopmat_support && !ctx->device->coopmat2 && - ggml_is_quantized(src0->type) && src1->type == GGML_TYPE_F32; + ggml_is_quantized(src0->type) && src1->type == GGML_TYPE_F32 && + // only take the f16-B path if a pipeline exists for this src0 type + // (e.g. Q2_0 has none); otherwise fall through to the normal f32-B path. + !(ctx->device->pipeline_dequant_mul_mat_mat_id_f16b[src0->type].f16acc->is_empty() && + ctx->device->pipeline_dequant_mul_mat_mat_id_f16b[src0->type].f32acc->is_empty()); if (mmid_f16b) { static bool mmid_f16b_logged = false; if (!mmid_f16b_logged) { From 68df0714b84f0ea041c1c00a2bc6a7ff2896ae57 Mon Sep 17 00:00:00 2001 From: Nathan Wilson Date: Thu, 6 Aug 2026 22:30:38 +0000 Subject: [PATCH 55/68] tests: cover MMQ tile boundaries in MUL_MAT and MUL_MAT_ID The MMQ J config is chosen from n, so a wave-partitioning error in one J tile only shows up when n sits on that tile's boundary. The neighbouring n picks a different J and passes, which hides it. The existing quantized cases stop at n=129 and the general MUL_MAT set jumps 64 -> 4096, so nothing lands on 256 or 512 and the whole class went untested. Sweep n over 255/256/257/511/512/513 for q8_0, q4_0, q4_K, q5_K and q6_K, in both MUL_MAT and MUL_MAT_ID. On an RDNA3.5 build that runs the J128 kernel with 16 wave32 waves over a 128-row tile these fail for every n that selects J128 and pass at n=513, which selects J112. A build predating that config passes the whole sweep. Co-Authored-By: Claude Opus 5 (cherry picked from commit b54cd8a8add4ddb2710eaf256b9fb2ba2be4b383) Assisted-by: Claude (Opus 5) --- tests/test-backend-ops.cpp | 13 +++++++++++++ 1 file changed, 13 insertions(+) diff --git a/tests/test-backend-ops.cpp b/tests/test-backend-ops.cpp index 55d25ccf9c3..964c5172866 100644 --- a/tests/test-backend-ops.cpp +++ b/tests/test-backend-ops.cpp @@ -10515,6 +10515,19 @@ static std::vector> make_test_cases_eval() { int k = 256; test_cases.emplace_back(new test_mul_mat_id(type_a, type_b, n_mats, n_used, b, m, n, k)); } + // MMQ tile-boundary cases. The MMQ J config is picked from n, and a wave-partitioning error in + // one J tile only shows up when n sits on that tile's boundary: the neighbouring n selects a + // different J and passes, which hides it. The general cases above stop at 129 and the MUL_MAT + // set jumps 64 -> 4096, so no existing case lands on 256 or 512. + // MMQ tile-boundary sweep: n on and either side of the 256 / 512 J boundaries. + for (ggml_type type_a : {GGML_TYPE_Q8_0, GGML_TYPE_Q4_0, GGML_TYPE_Q4_K, GGML_TYPE_Q5_K, GGML_TYPE_Q6_K}) { + for (int64_t n : {255, 256, 257, 511, 512, 513}) { + test_cases.emplace_back(new test_mul_mat(type_a, GGML_TYPE_F32, 1024, n, 256, {1, 1}, {1, 1})); + test_cases.emplace_back(new test_mul_mat_id(type_a, GGML_TYPE_F32, 8, 2, false, 1024, n, 256)); + } + } + + } } } From 77214b50c05fe5a21db8144992051b910a9bf2d9 Mon Sep 17 00:00:00 2001 From: Nathan Wilson Date: Sat, 8 Aug 2026 04:24:38 +0000 Subject: [PATCH 56/68] vulkan: create coopmat2 mul_mat_id pipelines with the real param count 30d8bb02b raised mul_mat_id_param_count to 6 for the fused MUL epilogue and gave every mul_mat_id shader a binding 5 (FusedScale), including the coopmat2 one, which binds it purely for descriptor-layout parity. The coopmat2 pipeline creation block was left passing a literal 5, so two things go wrong there: - the pipelines are created with 5 descriptors while the shader declares 6 - PARAMCOUNT == mul_mat_id_param_count doubles as the "this is mul_mat_id" argument to ggml_vk_mul_mm_cm2_spec, so it went false and every coopmat2 mul_mat_id pipeline was specialized as a plain matmul, dropping the trailing spec constant Only reachable where device->coopmat2 is true. gfx1151 does not take that path and the v0.5 release predates the constant bump, so neither is affected. Not validated on hardware - no coopmat2 device here. The block does compile: built with the pinned shaderc (GL_NV_cooperative_matrix2 supported). The two OCP FP4 sites stay behind GL_EXT_float_e2m1, which that glslc does not support. Co-Authored-By: Claude Opus 5 Assisted-by: Claude (Opus 5) --- ggml/src/ggml-vulkan/ggml-vulkan.cpp | 66 ++++++++++++++-------------- 1 file changed, 33 insertions(+), 33 deletions(-) diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp index a7c7f15b4ea..fc7c41b15b2 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp +++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp @@ -4989,48 +4989,48 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { GGML_ASSERT(device->subgroup_ballot); - CREATE_MM2(pipeline_matmul_id_f16, matmul_id_subgroup_f16, wg_denoms, warptile, vk_mat_mat_id_push_constants, 5) + CREATE_MM2(pipeline_matmul_id_f16, matmul_id_subgroup_f16, wg_denoms, warptile, vk_mat_mat_id_push_constants, mul_mat_id_param_count) #if defined(GGML_VULKAN_BFLOAT16_GLSLC_SUPPORT) if (device->coopmat_bf16_support) { - CREATE_MM(pipeline_matmul_id_bf16, matmul_id_subgroup_bf16, , wg_denoms, warptile, vk_mat_mat_id_push_constants, 5) + CREATE_MM(pipeline_matmul_id_bf16, matmul_id_subgroup_bf16, , wg_denoms, warptile, vk_mat_mat_id_push_constants, mul_mat_id_param_count) } #endif - CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_Q1_0], matmul_id_subgroup_q1_0_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, 5) - CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_Q2_0], matmul_id_subgroup_q2_0_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, 5) - CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_Q4_0], matmul_id_subgroup_q4_0_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, 5) - CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_Q4_1], matmul_id_subgroup_q4_1_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, 5) - CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_Q5_0], matmul_id_subgroup_q5_0_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, 5) - CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_Q5_1], matmul_id_subgroup_q5_1_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, 5) - CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_Q8_0], matmul_id_subgroup_q8_0_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, 5) - CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_Q2_K], matmul_id_subgroup_q2_k_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, 5) - CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_TQ2_0], matmul_id_subgroup_tq2_0_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, 5) - CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_Q3_K], matmul_id_subgroup_q3_k_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, 5) - CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_Q4_K], matmul_id_subgroup_q4_k_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, 5) - CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_Q5_K], matmul_id_subgroup_q5_k_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, 5) - CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_Q6_K], matmul_id_subgroup_q6_k_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, 5) - CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_IQ1_S], matmul_id_subgroup_iq1_s_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, 5) - CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_IQ1_M], matmul_id_subgroup_iq1_m_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, 5) - CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_IQ2_XXS], matmul_id_subgroup_iq2_xxs_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, 5) - CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_IQ2_XS], matmul_id_subgroup_iq2_xs_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, 5) - CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_IQ2_S], matmul_id_subgroup_iq2_s_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, 5) - CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_IQ3_XXS], matmul_id_subgroup_iq3_xxs_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, 5) - CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_IQ3_S], matmul_id_subgroup_iq3_s_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, 5) - CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_IQ4_XS], matmul_id_subgroup_iq4_xs_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, 5) - CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_IQ4_NL], matmul_id_subgroup_iq4_nl_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, 5) - CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_Q4_0_ROCMFP4], matmul_id_subgroup_rocmfp4_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, 5) - CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_Q4_0_ROCMFP4_FAST], matmul_id_subgroup_rocmfp4_fast_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, 5) - CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_Q3_0_ROCMFPX], matmul_id_subgroup_rocmfpx_fp3_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, 5) - CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_Q6_0_ROCMFPX], matmul_id_subgroup_rocmfpx_fp6_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, 5) - CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_Q8_0_ROCMFPX], matmul_id_subgroup_rocmfpx_fp8_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, 5) + CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_Q1_0], matmul_id_subgroup_q1_0_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, mul_mat_id_param_count) + CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_Q2_0], matmul_id_subgroup_q2_0_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, mul_mat_id_param_count) + CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_Q4_0], matmul_id_subgroup_q4_0_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, mul_mat_id_param_count) + CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_Q4_1], matmul_id_subgroup_q4_1_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, mul_mat_id_param_count) + CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_Q5_0], matmul_id_subgroup_q5_0_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, mul_mat_id_param_count) + CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_Q5_1], matmul_id_subgroup_q5_1_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, mul_mat_id_param_count) + CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_Q8_0], matmul_id_subgroup_q8_0_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, mul_mat_id_param_count) + CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_Q2_K], matmul_id_subgroup_q2_k_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, mul_mat_id_param_count) + CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_TQ2_0], matmul_id_subgroup_tq2_0_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, mul_mat_id_param_count) + CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_Q3_K], matmul_id_subgroup_q3_k_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, mul_mat_id_param_count) + CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_Q4_K], matmul_id_subgroup_q4_k_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, mul_mat_id_param_count) + CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_Q5_K], matmul_id_subgroup_q5_k_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, mul_mat_id_param_count) + CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_Q6_K], matmul_id_subgroup_q6_k_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, mul_mat_id_param_count) + CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_IQ1_S], matmul_id_subgroup_iq1_s_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, mul_mat_id_param_count) + CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_IQ1_M], matmul_id_subgroup_iq1_m_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, mul_mat_id_param_count) + CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_IQ2_XXS], matmul_id_subgroup_iq2_xxs_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, mul_mat_id_param_count) + CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_IQ2_XS], matmul_id_subgroup_iq2_xs_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, mul_mat_id_param_count) + CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_IQ2_S], matmul_id_subgroup_iq2_s_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, mul_mat_id_param_count) + CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_IQ3_XXS], matmul_id_subgroup_iq3_xxs_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, mul_mat_id_param_count) + CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_IQ3_S], matmul_id_subgroup_iq3_s_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, mul_mat_id_param_count) + CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_IQ4_XS], matmul_id_subgroup_iq4_xs_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, mul_mat_id_param_count) + CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_IQ4_NL], matmul_id_subgroup_iq4_nl_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, mul_mat_id_param_count) + CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_Q4_0_ROCMFP4], matmul_id_subgroup_rocmfp4_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, mul_mat_id_param_count) + CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_Q4_0_ROCMFP4_FAST], matmul_id_subgroup_rocmfp4_fast_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, mul_mat_id_param_count) + CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_Q3_0_ROCMFPX], matmul_id_subgroup_rocmfpx_fp3_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, mul_mat_id_param_count) + CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_Q6_0_ROCMFPX], matmul_id_subgroup_rocmfpx_fp6_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, mul_mat_id_param_count) + CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_Q8_0_ROCMFPX], matmul_id_subgroup_rocmfpx_fp8_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, mul_mat_id_param_count) #if defined(GGML_VULKAN_FLOAT_E2M1_GLSLC_SUPPORT) && defined(GGML_VULKAN_FLOAT_E4M3_GLSLC_SUPPORT) if (device->ocp_fp4) { - CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_MXFP4], matmul_id_subgroup_mxfp4_f16_ocp, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, 5) - CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_NVFP4], matmul_id_subgroup_nvfp4_f16_ocp, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, 5) + CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_MXFP4], matmul_id_subgroup_mxfp4_f16_ocp, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, mul_mat_id_param_count) + CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_NVFP4], matmul_id_subgroup_nvfp4_f16_ocp, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, mul_mat_id_param_count) } else #endif { - CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_MXFP4], matmul_id_subgroup_mxfp4_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, 5) - CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_NVFP4], matmul_id_subgroup_nvfp4_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, 5) + CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_MXFP4], matmul_id_subgroup_mxfp4_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, mul_mat_id_param_count) + CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_NVFP4], matmul_id_subgroup_nvfp4_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, mul_mat_id_param_count) } #undef CREATE_MM #undef CREATE_MM2 From 6192a050dc3f751a74134f524828f621e8bdc5d1 Mon Sep 17 00:00:00 2001 From: Nathan Wilson Date: Sun, 30 Aug 2026 10:20:41 +0000 Subject: [PATCH 57/68] vulkan: enable the Strix mmid tile gates by default Split from the fork's flag-flip commit (f7d804ee7): GGML_VK_MMID_F16B, BM64, M128, WAVE32 and SMALLN now default on with =0 opt-out. The wave32 gate keeps its subgroup_size_control device guard, so devices without a 32-wide subgroup mode are unaffected. GGML_VK_MMID_TILE16 stays opt-in (documented negative on gfx1151). Co-Authored-By: Claude Fable 5 Assisted-by: Claude (Opus 5) --- ggml/src/ggml-vulkan/ggml-vulkan.cpp | 10 +++++----- 1 file changed, 5 insertions(+), 5 deletions(-) diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp index fc7c41b15b2..04bcdec5606 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp +++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp @@ -4465,7 +4465,7 @@ struct CompileTask { static bool ggml_vk_mmid_f16b_enabled() { static const bool enabled = [] { const char * env = getenv("GGML_VK_MMID_F16B"); - return env != nullptr && atoi(env) != 0; + return env == nullptr || atoi(env) != 0; // on by default; =0 disables }(); return enabled; } @@ -5145,7 +5145,7 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { // A-traffic (ic-tile count unchanged), halves per-expert B re-reads // (ir-tile count), larger WGs to hide latency. const char * bm64_env = getenv("GGML_VK_MMID_BM64"); - if (bm64_env && atoi(bm64_env) != 0) { + if (!bm64_env || atoi(bm64_env) != 0) { // on by default; =0 disables s_warptile_mmq_id16[0] = 2 * mul_mat_subgroup_size; // BLOCK_SIZE s_warptile_mmq_id16[1] = 64; // BM s_mmq_wg_denoms_id16[0] = 64; @@ -5153,7 +5153,7 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { // GGML_VK_MMID_M128=1: same idea for the medium tile (BM 64->128, four // warps) — the tile the per-expert-n heuristic picks at n~64 (e.g. ub2048). const char * m128_env = getenv("GGML_VK_MMID_M128"); - if (m128_env && atoi(m128_env) != 0) { + if (!m128_env || atoi(m128_env) != 0) { // on by default; =0 disables m_warptile_mmq_id128[0] = 4 * mul_mat_subgroup_size; // BLOCK_SIZE m_warptile_mmq_id128[1] = 128; // BM m_mmq_wg_denoms_id128[0] = 128; @@ -5175,7 +5175,7 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { // driver honors a required size (subgroup_size_control covering 32); // otherwise WARP=32 with a real subgroup of 64 would corrupt tiling. const char * wave32_env = getenv("GGML_VK_MMID_WAVE32"); - if (wave32_env && atoi(wave32_env) != 0 && device->subgroup_size_control && + if ((!wave32_env || atoi(wave32_env) != 0) && device->subgroup_size_control && // on by default; =0 disables device->subgroup_min_size <= 32 && 32 <= device->subgroup_max_size) { mmid_req_sgs = 32; auto wave32_tile = [](std::vector &w) { @@ -11021,7 +11021,7 @@ static void ggml_vk_mul_mat_id_q_f16(ggml_backend_vk_context * ctx, vk_context& // matmul efficiency, ~78% of MoE prefill time on Qwen3.6-35B-A3B). uint32_t n_for_tile = (uint32_t)nei1; static const char * mmid_smalln_env = getenv("GGML_VK_MMID_SMALLN"); - if (mmid_smalln_env && atoi(mmid_smalln_env) != 0 && ne02 > 1) { + if (!(mmid_smalln_env && atoi(mmid_smalln_env) == 0) && ne02 > 1) { // on by default; =0 disables n_for_tile = std::max(1u, (uint32_t)((nei1 * nei0 + ne02 - 1) / ne02)); } From 891cfff9ae640eda8956957452a396947fac3212 Mon Sep 17 00:00:00 2001 From: Nathan Wilson Date: Tue, 18 Aug 2026 01:12:00 +0000 Subject: [PATCH 58/68] vulkan: run the quantised dense coopmat pipelines at wave32 RDNA3.x WMMA is wave32-native, so a wave64 subgroup issues each coopmat op as two halves. GGML_VK_MMID_WAVE32 already exploits this for mul_mat_id and leaves the dense pipelines at the driver default; this gives dense the same treatment, gated on measurement rather than a flag. BLOCK_SIZE is kept, so the subgroup count doubles and WM (or WN) halves until the warp grid tiles BM x BN again. Only the quantised tiles are retiled: on gfx1151 a standalone MUL_MAT microbench at both dense FFN shapes reads q6_K +5.2..+10.8%, q8_0 +5.4..+8.4%, q4_K +0.7..+9.1%, q4_0 -1.5..+1.8%, while f16 reads -6.7..+6.4% and bf16 ~0 - the float paths are bandwidth-bound on the weight stream, not issue-bound. The win tracks inline dequant instruction count (q6_K 3907 -> 3433 instructions, identical 192 VGPRs and 8 subgroups/SIMD). The required subgroup size is now the tile's own WARP for every dense coopmat pipeline. The cm1 shaders derive their warp grid from gl_SubgroupID and size shared arrays as NUM_WARPS = BLOCK_SIZE / WARP, so WARP and the real subgroup must agree; leaving that to the driver made the agreement incidental. Scoped to AMD coopmat1 on a wave64 default; other vendors keep the driver default. GGML_VK_DENSE_WAVE32=0 disables, =2 also retiles the float tiles. Qwen3-32B Q6_K_XL pp2048 +7.2% at ub256 / +3.9% at ub2048, Qwen3.8-27B +5.3% / +4.8%. PPL unchanged: 6.9496 +/- 0.24246 in both arms, all 20 per-chunk values identical, since the retile changes which warp owns an output sub-tile and not the K-reduction order within an element. Assisted-by: Claude Opus 5 (cherry picked from commit 448994e9637405610ced3e0ead02c7b6fa688314) --- ggml/src/ggml-vulkan/ggml-vulkan.cpp | 83 ++++++++++++++++++++++++++-- 1 file changed, 77 insertions(+), 6 deletions(-) diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp index 04bcdec5606..59acd2dfb8f 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp +++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp @@ -5038,20 +5038,91 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { #endif // defined(VK_NV_cooperative_matrix2) && defined(GGML_VULKAN_COOPMAT2_GLSLC_SUPPORT) #if defined(VK_KHR_cooperative_matrix) && defined(GGML_VULKAN_COOPMAT_GLSLC_SUPPORT) if (device->coopmat_support) { + // Deterministic subgroup sizing for the dense coopmat pipelines. Two parts: + // + // (1) The required subgroup size is the tile's own WARP, not the driver's choice. The cm1 + // shaders derive their warp grid from the real subgroup (warp_i = gl_SubgroupID) and + // size shared arrays with NUM_WARPS = BLOCK_SIZE / WARP, so WARP and the actual + // subgroup must agree - leaving that to the driver makes the agreement incidental. + // + // (2) The QUANTISED dense tiles run at wave32. RDNA3.x WMMA is wave32-native, so a wave64 + // subgroup issues each coopmat op as two halves. BLOCK_SIZE is kept, so the subgroup + // count doubles and WM (or WN) halves until the warp grid tiles BM x BN again: + // NUM_WARPS == (BM/WM)*(BN/WN). A tile that cannot be retiled is left at wave64. + // + // The FLOAT tiles are deliberately excluded. Measured on gfx1151 with a standalone + // MUL_MAT microbench at both dense FFN shapes (m=25600 k=5120, m=5120 k=25600): + // q6_K +5.2..+10.8%, q8_0 +5.4..+8.4%, q4_K +0.7..+9.1%, q4_0 -1.5..+1.8%, but + // f16 -6.7..+6.4% and bf16 ~0. The win tracks inline dequant instruction count + // (q6_K 3907 -> 3433 instructions at wave32, identical 192 VGPRs and 8 subgroups/SIMD), + // so it lands on the issue-bound quantised kernels and not on the float ones, which are + // bandwidth-bound on the weight stream. + // + // Scoped to AMD coopmat1 on a wave64 default; other vendors keep the driver default, + // since none of the above is validated there. + // GGML_VK_DENSE_WAVE32=0 disables, =2 additionally retiles the float tiles (probe). + const bool dense_sgs_scope = + device->vendor_id == VK_VENDOR_ID_AMD && + device->driver_id != vk::DriverId::eAmdProprietary && + device->subgroup_size_control; + const bool dense_wave32_possible = + dense_sgs_scope && + device->subgroup_min_size <= 32 && 32 <= device->subgroup_max_size && + device->subgroup_size > 32; + const char * dense_wave32_env = getenv("GGML_VK_DENSE_WAVE32"); + const int dense_wave32 = dense_wave32_env ? atoi(dense_wave32_env) : 1; + + if (dense_wave32_possible && dense_wave32 != 0) { + auto wave32_tile = [](std::vector & w) -> bool { + // {BLOCK_SIZE, BM, BN, BK, WM, WN, WMITER, TM, TN, TK, WARP} + std::vector t = w; + t[10] = 32; + for (int guard = 0; guard < 4 && t[0] / t[10] != (t[1] / t[4]) * (t[2] / t[5]); ++guard) { + if (t[4] >= t[5] && t[4] > t[7]) { + t[4] /= 2; // halve WM, keeping WM >= TM + } else { + t[5] /= 2; // halve WN + } + } + if (t[0] / t[10] != (t[1] / t[4]) * (t[2] / t[5]) || t[4] < t[7] || t[5] < t[8]) { + return false; + } + w = t; + return true; + }; + wave32_tile(l_warptile_mmq); + wave32_tile(m_warptile_mmq); + wave32_tile(s_warptile_mmq); + if (dense_wave32 >= 2) { + wave32_tile(l_warptile); + wave32_tile(m_warptile); + wave32_tile(s_warptile); + } + } + + // WARP -> required subgroup size, or 0 where the device cannot honor one. + auto dense_req_sgs = [dense_sgs_scope, &device](const std::vector & w) -> uint32_t { + const uint32_t warp = w[10]; + if (!dense_sgs_scope || warp < device->subgroup_min_size || warp > device->subgroup_max_size) { + return 0; + } + return warp; + }; + // Create 6 variants, {s,m,l}x{unaligned,aligned} #define CREATE_MM(TYPE, PIPELINE_NAME, NAMELC, F16ACC, WG_DENOMS, WARPTILE, PUSHCONST, PARAMCOUNT, ID) \ if (device->mul_mat ## ID ## _l[TYPE]) \ - ggml_vk_create_pipeline(device, device-> PIPELINE_NAME ->l, #NAMELC #F16ACC "_l", NAMELC ## F16ACC ## _cm1_len, NAMELC ## F16ACC ## _cm1_data, "main", PARAMCOUNT, sizeof(PUSHCONST), l_ ## WG_DENOMS, ggml_vk_mul_mm_spec(l_ ## WARPTILE, false), 1, false, true); \ + ggml_vk_create_pipeline(device, device-> PIPELINE_NAME ->l, #NAMELC #F16ACC "_l", NAMELC ## F16ACC ## _cm1_len, NAMELC ## F16ACC ## _cm1_data, "main", PARAMCOUNT, sizeof(PUSHCONST), l_ ## WG_DENOMS, ggml_vk_mul_mm_spec(l_ ## WARPTILE, false), 1, false, true, dense_req_sgs(l_ ## WARPTILE)); \ if (device->mul_mat ## ID ## _m[TYPE]) \ - ggml_vk_create_pipeline(device, device-> PIPELINE_NAME ->m, #NAMELC #F16ACC "_m", NAMELC ## F16ACC ## _cm1_len, NAMELC ## F16ACC ## _cm1_data, "main", PARAMCOUNT, sizeof(PUSHCONST), m_ ## WG_DENOMS, ggml_vk_mul_mm_spec(m_ ## WARPTILE, false), 1, false, true); \ + ggml_vk_create_pipeline(device, device-> PIPELINE_NAME ->m, #NAMELC #F16ACC "_m", NAMELC ## F16ACC ## _cm1_len, NAMELC ## F16ACC ## _cm1_data, "main", PARAMCOUNT, sizeof(PUSHCONST), m_ ## WG_DENOMS, ggml_vk_mul_mm_spec(m_ ## WARPTILE, false), 1, false, true, dense_req_sgs(m_ ## WARPTILE)); \ if (device->mul_mat ## ID ## _s[TYPE]) \ - ggml_vk_create_pipeline(device, device-> PIPELINE_NAME ->s, #NAMELC #F16ACC "_s", NAMELC ## F16ACC ## _cm1_len, NAMELC ## F16ACC ## _cm1_data, "main", PARAMCOUNT, sizeof(PUSHCONST), s_ ## WG_DENOMS, ggml_vk_mul_mm_spec(s_ ## WARPTILE, false), 1, false, true); \ + ggml_vk_create_pipeline(device, device-> PIPELINE_NAME ->s, #NAMELC #F16ACC "_s", NAMELC ## F16ACC ## _cm1_len, NAMELC ## F16ACC ## _cm1_data, "main", PARAMCOUNT, sizeof(PUSHCONST), s_ ## WG_DENOMS, ggml_vk_mul_mm_spec(s_ ## WARPTILE, false), 1, false, true, dense_req_sgs(s_ ## WARPTILE)); \ if (device->mul_mat ## ID ## _l[TYPE]) \ - ggml_vk_create_pipeline(device, device-> PIPELINE_NAME ->a_l, #NAMELC #F16ACC "_aligned_l", NAMELC ## F16ACC ## _cm1_len, NAMELC ## F16ACC ## _cm1_data, "main", PARAMCOUNT, sizeof(PUSHCONST), l_ ## WG_DENOMS, ggml_vk_mul_mm_spec(l_ ## WARPTILE, true), l_align, false, true); \ + ggml_vk_create_pipeline(device, device-> PIPELINE_NAME ->a_l, #NAMELC #F16ACC "_aligned_l", NAMELC ## F16ACC ## _cm1_len, NAMELC ## F16ACC ## _cm1_data, "main", PARAMCOUNT, sizeof(PUSHCONST), l_ ## WG_DENOMS, ggml_vk_mul_mm_spec(l_ ## WARPTILE, true), l_align, false, true, dense_req_sgs(l_ ## WARPTILE)); \ if (device->mul_mat ## ID ## _m[TYPE]) \ - ggml_vk_create_pipeline(device, device-> PIPELINE_NAME ->a_m, #NAMELC #F16ACC "_aligned_m", NAMELC ## F16ACC ## _cm1_len, NAMELC ## F16ACC ## _cm1_data, "main", PARAMCOUNT, sizeof(PUSHCONST), m_ ## WG_DENOMS, ggml_vk_mul_mm_spec(m_ ## WARPTILE, true), m_align, false, true); \ + ggml_vk_create_pipeline(device, device-> PIPELINE_NAME ->a_m, #NAMELC #F16ACC "_aligned_m", NAMELC ## F16ACC ## _cm1_len, NAMELC ## F16ACC ## _cm1_data, "main", PARAMCOUNT, sizeof(PUSHCONST), m_ ## WG_DENOMS, ggml_vk_mul_mm_spec(m_ ## WARPTILE, true), m_align, false, true, dense_req_sgs(m_ ## WARPTILE)); \ if (device->mul_mat ## ID ## _s[TYPE]) \ - ggml_vk_create_pipeline(device, device-> PIPELINE_NAME ->a_s, #NAMELC #F16ACC "_aligned_s", NAMELC ## F16ACC ## _cm1_len, NAMELC ## F16ACC ## _cm1_data, "main", PARAMCOUNT, sizeof(PUSHCONST), s_ ## WG_DENOMS, ggml_vk_mul_mm_spec(s_ ## WARPTILE, true), s_align, false, true); \ + ggml_vk_create_pipeline(device, device-> PIPELINE_NAME ->a_s, #NAMELC #F16ACC "_aligned_s", NAMELC ## F16ACC ## _cm1_len, NAMELC ## F16ACC ## _cm1_data, "main", PARAMCOUNT, sizeof(PUSHCONST), s_ ## WG_DENOMS, ggml_vk_mul_mm_spec(s_ ## WARPTILE, true), s_align, false, true, dense_req_sgs(s_ ## WARPTILE)); \ // Create 2 variants, {f16,f32} accumulator #define CREATE_MM2(TYPE, PIPELINE_NAME, NAMELC, WG_DENOMS, WARPTILE, PUSHCONST, PARAMCOUNT, ID) \ From f4d6f9e8da7886598083167416865cdcfee5da18 Mon Sep 17 00:00:00 2001 From: Nathan Wilson Date: Thu, 6 Aug 2026 01:34:10 +0000 Subject: [PATCH 59/68] vulkan: four env-gated Strix Halo prefill fixes for delta-net MoE All default OFF, so the same binary A/Bs each change. GGML_VK_CONCAT_TRANSPOSE: delta-net does ggml_transpose() into a dim-0 ggml_concat(), which the generic concat reads fully de-coalesced. Route that shape through a 32x32 shared-memory tile transpose. CONCAT 11877 -> 957 us/op. GGML_VK_MMID_SCALE_EPILOGUE: apply the following MUL's per-(expert,token) broadcast scale as mul_mat_id writes out, removing a 134 MB write plus read back. Prefill only; the existing fusion is gated to mat-vec. Not implemented in the coopmat2 shader, so it is refused there. GGML_VK_FUSE_UNARY_MUL: silu(x)*y is two nodes in the delta-net path; run it as the existing swiglu split. 750 -> 443 us/op. GGML_VK_MMID_WG256: the RADV tuning gives the dense large tile 256 threads on a 128x128 tile but left the mul_mat_id variants at 128. Qwen3.6-35B-A3B UD-Q4_K_XL, gfx1151, pp2048 at ub2048: 1223.75 -> 1733.64 t/s. test-backend-ops CONCAT/MUL_MAT_ID/UNARY/MUL pass; generated text is unchanged with each flag on, and the disabled paths are bit-identical to before. Co-Authored-By: Claude Opus 5 (cherry picked from commit 6f7a49e661d93fa42df528bdc88bc788eaf0ad2c) Assisted-by: Claude (Opus 5) --- ggml/src/ggml-vulkan/ggml-vulkan.cpp | 182 ++++++++++++++++-- .../vulkan-shaders/concat_transpose.comp | 43 +++++ .../ggml-vulkan/vulkan-shaders/mul_mm.comp | 20 +- .../vulkan-shaders/mul_mm_cm2.comp | 3 + .../ggml-vulkan/vulkan-shaders/mul_mmq.comp | 8 +- .../vulkan-shaders/vulkan-shaders-gen.cpp | 1 + 6 files changed, 241 insertions(+), 16 deletions(-) create mode 100644 ggml/src/ggml-vulkan/vulkan-shaders/concat_transpose.comp diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp index 59acd2dfb8f..8f297cb6696 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp +++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp @@ -999,6 +999,7 @@ struct vk_device_struct { vk_pipeline pipeline_add_id_f32; vk_pipeline pipeline_concat_i8, pipeline_concat_i16, pipeline_concat_i32, pipeline_concat_i64; + vk_pipeline pipeline_concat_transpose_i32; vk_pipeline pipeline_upscale_nearest_f32, pipeline_upscale_bilinear_f32, pipeline_upscale_bicubic_f32, pipeline_upscale_bilinear_antialias_f32; vk_pipeline pipeline_scale_f32; vk_pipeline pipeline_log[2]; @@ -1403,6 +1404,7 @@ struct vk_mat_mat_id_push_constants { uint32_t nei0; uint32_t nei1; uint32_t nbi1; uint32_t ne11; uint32_t n_experts; uint32_t hoist_row_ids; + uint32_t fusion_flags; }; struct vk_mat_vec_id_push_constants { uint32_t ncols; @@ -2418,6 +2420,20 @@ class vk_perf_logger { if (node->op == GGML_OP_UNARY) { return fusion_str + ggml_unary_op_name(ggml_get_unary_op(node)); } + if (node->op == GGML_OP_MUL && getenv("GGML_VK_PERF_SHAPES")) { + std::string name = "MUL "; + name += "dst(" + std::to_string(node->ne[0]) + "," + std::to_string(node->ne[1]) + "," + + std::to_string(node->ne[2]) + ") b(" + std::to_string(node->src[1]->ne[0]) + "," + + std::to_string(node->src[1]->ne[1]) + "," + std::to_string(node->src[1]->ne[2]) + ")"; + name += std::string(" a=") + ggml_op_name(node->src[0]->op); + if (node->src[0]->op == GGML_OP_UNARY) { name += std::string(":") + ggml_unary_op_name(ggml_get_unary_op(node->src[0])); } + name += std::string(" b=") + ggml_op_name(node->src[1]->op); + if (node->src[1]->op == GGML_OP_UNARY) { name += std::string(":") + ggml_unary_op_name(ggml_get_unary_op(node->src[1])); } + if (node->src[1]->op == GGML_OP_RESHAPE && node->src[1]->src[0]) { + name += std::string("(") + ggml_op_name(node->src[1]->src[0]->op) + ")"; + } + return fusion_str + name; + } if (node->op == GGML_OP_MUL_MAT || node->op == GGML_OP_MUL_MAT_ID) { const uint64_t m = node->ne[0]; const uint64_t n = node->ne[1]; @@ -4611,6 +4627,23 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { l_warptile = { 256, 128, 128, 16, mm_warp_8, 64, 2, tm_m, tn_m, tk_m, mm_warp_8 }; l_warptile_mmq = l_warptile_mmq_int = { 256, 128, 128, 32, mm_warp_8, 64, 2, tm_m, tn_m, tk_m, mm_warp_8 }; l_warptile_mmq_int_k = { 256, 128, 128, 32, mm_warp_16, 64, 1, 4, 2, 1, mm_warp_16 }; + + // EXPERIMENT (GGML_VK_MMID_WG256=1): the dense large tile above runs 256 threads on a + // 128x128 tile, but the mul_mat_id variants still run 128. Give MoE the same thread + // count per tile: same BM/BN/BK (so same shared memory), twice the threads sharing each + // A/B tile load, half the accumulators per thread. Warp split stays legal: + // (BM/WM)*(BN/WN) == wg/subgroup == 4, WNITER == (WM*WN)/(WARP*TM*TN*WMITER) == 2. + // A 256-expert MoE at ub=2048 sees only ~64 rows per expert, so the tile that actually + // runs is the medium one, not the large one. Override both. + static const char * mmid_wg256_env = getenv("GGML_VK_MMID_WG256"); + if (mmid_wg256_env && atoi(mmid_wg256_env) != 0) { + l_warptile_mmqid = { 256, 128, 128, 32, mul_mat_subgroup_size_8, 64, 2, tm_m, tn_m, tk_m, mul_mat_subgroup_size_8 }; + l_warptile_mmqid_int = { 256, 128, 128, 32, mul_mat_subgroup_size_8, 64, 2, 4, 4, 1, mul_mat_subgroup_size_8 }; + // BM=BN=64 at 4 warps needs WM=WN=32: (BM/WM)*(BN/WN) == 4, cms_per_row/col == 2. + m_warptile_mmqid = { 256, 64, 64, 32, 32, 32, 2, tm_m, tn_m, tk_m, mul_mat_subgroup_size_8 }; + m_warptile_mmqid_int = { 256, 64, 64, 32, 32, 32, 2, 2, 2, 1, mul_mat_subgroup_size_8 }; + fprintf(stderr, "ggml_vulkan: MUL_MAT_ID medium+large tiles at 256 threads (GGML_VK_MMID_WG256)\n"); + } } else if (device->vendor_id == VK_VENDOR_ID_INTEL && device->coopmat_support) { // Xe2/Xe3 with coopmat enabled - warptile performance tuning l_warptile = { 512, 128, 128, 16, mm_warp_8, 32, 2, tm_m, tn_m, tk_m, mm_warp_8 }; @@ -4913,7 +4946,7 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { return spec; }; - const int mul_mat_id_param_count = 5; + const int mul_mat_id_param_count = 6; // a, b, d, ids, expert_counts, fused scale #if defined(VK_NV_cooperative_matrix2) && defined(GGML_VULKAN_COOPMAT2_GLSLC_SUPPORT) if (device->coopmat2) { @@ -6250,6 +6283,8 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { ggml_vk_create_pipeline(device, device->pipeline_concat_i8, "concat_i8", concat_i8_len, concat_i8_data, "main", 3, sizeof(vk_op_binary_push_constants), {512, 1, 1}, {}, 1); ggml_vk_create_pipeline(device, device->pipeline_concat_i16, "concat_i16", concat_i16_len, concat_i16_data, "main", 3, sizeof(vk_op_binary_push_constants), {512, 1, 1}, {}, 1); ggml_vk_create_pipeline(device, device->pipeline_concat_i32, "concat_i32", concat_i32_len, concat_i32_data, "main", 3, sizeof(vk_op_binary_push_constants), {512, 1, 1}, {}, 1); + // One workgroup per 32x32 tile: elements are passed as (rows, cols, 1). + ggml_vk_create_pipeline(device, device->pipeline_concat_transpose_i32, "concat_transpose_i32", concat_transpose_i32_len, concat_transpose_i32_data, "main", 3, sizeof(vk_op_binary_push_constants), {32, 32, 1}, {}, 1); ggml_vk_create_pipeline(device, device->pipeline_concat_i64, "concat_i64", concat_i64_len, concat_i64_data, "main", 3, sizeof(vk_op_binary_push_constants), {512, 1, 1}, {}, 1); ggml_vk_create_pipeline(device, device->pipeline_upscale_nearest_f32, "upscale_f32", upscale_f32_len, upscale_f32_data, "main", 2, sizeof(vk_op_upscale_push_constants), {512, 1, 1}, {GGML_SCALE_MODE_NEAREST}, 1); @@ -9699,14 +9734,14 @@ static void ggml_vk_matmul_id( uint32_t m, uint32_t n, uint32_t k, uint32_t stride_a, uint32_t stride_b, uint32_t stride_d, uint32_t batch_stride_a, uint32_t batch_stride_b, uint32_t batch_stride_d, uint32_t n_as, uint32_t nei0, uint32_t nei1, uint32_t nbi1, uint32_t ne11, - bool hoist_row_ids) { + bool hoist_row_ids, const vk_subbuffer & fused_scale, uint32_t fusion_flags) { VK_LOG_DEBUG("ggml_vk_matmul_id(a: (" << a.buffer->buffer << ", " << a.offset << ", " << a.size << "), b: (" << b.buffer->buffer << ", " << b.offset << ", " << b.size << "), d: (" << d.buffer->buffer << ", " << d.offset << ", " << d.size << "), ids: (" << ids.buffer->buffer << ", " << ids.offset << ", " << ids.size << "), expert_count: (" << expert_count_buf.buffer->buffer << ", " << expert_count_buf.offset << ", " << expert_count_buf.size << "), " << "m: " << m << ", n: " << n << ", k: " << k << ", stride_a: " << stride_a << ", stride_b: " << stride_b << ", stride_d: " << stride_d << ", " << "batch_stride_a: " << batch_stride_a << ", batch_stride_b: " << batch_stride_b << ", batch_stride_d: " << batch_stride_d << ", " << "n_as: " << n_as << ", nei0: " << nei0 << ", nei1: " << nei1 << ", nbi1: " << nbi1 << ", ne11: " << ne11 << ")"); const vk_mat_mat_id_push_constants pc = { m, n, k, stride_a, stride_b, stride_d, batch_stride_a, batch_stride_b, batch_stride_d, - nei0, nei1, nbi1, ne11, n_as, uint32_t(hoist_row_ids) }; - ggml_vk_dispatch_pipeline(ctx, subctx, pipeline, { a, b, d, ids, expert_count_buf }, pc, { m, nei1, n_as }); + nei0, nei1, nbi1, ne11, n_as, uint32_t(hoist_row_ids), fusion_flags }; + ggml_vk_dispatch_pipeline(ctx, subctx, pipeline, { a, b, d, ids, expert_count_buf, fused_scale }, pc, { m, nei1, n_as }); } static bool ggml_vk_dim01_contiguous(const ggml_tensor * tensor) { @@ -10953,7 +10988,7 @@ static void ggml_vk_mul_mat(ggml_backend_vk_context * ctx, vk_context& subctx, c } } -static void ggml_vk_mul_mat_id_q_f16(ggml_backend_vk_context * ctx, vk_context& subctx, const ggml_tensor * src0, const ggml_tensor * src1, const ggml_tensor * ids, ggml_tensor * dst) { +static void ggml_vk_mul_mat_id_q_f16(ggml_backend_vk_context * ctx, vk_context& subctx, const ggml_tensor * src0, const ggml_tensor * src1, const ggml_tensor * ids, ggml_tensor * dst, const ggml_tensor * fused_scale = nullptr, ggml_tensor * fused_dst = nullptr) { VK_LOG_DEBUG("ggml_vk_mul_mat_id_q_f16((" << src0 << ", name=" << src0->name << ", type=" << src0->type << ", ne0=" << src0->ne[0] << ", ne1=" << src0->ne[1] << ", ne2=" << src0->ne[2] << ", ne3=" << src0->ne[3] << ", nb0=" << src0->nb[0] << ", nb1=" << src0->nb[1] << ", nb2=" << src0->nb[2] << ", nb3=" << src0->nb[3]; std::cerr << "), (" << src1 << ", name=" << src1->name << ", type=" << src1->type << ", ne0=" << src1->ne[0] << ", ne1=" << src1->ne[1] << ", ne2=" << src1->ne[2] << ", ne3=" << src1->ne[3] << ", nb0=" << src1->nb[0] << ", nb1=" << src1->nb[1] << ", nb2=" << src1->nb[2] << ", nb3=" << src1->nb[3]; std::cerr << "), (" << ids << ", name=" << ids->name << ", type=" << ids->type << ", ne0=" << ids->ne[0] << ", ne1=" << ids->ne[1] << ", ne2=" << ids->ne[2] << ", ne3=" << ids->ne[3] << ", nb0=" << ids->nb[0] << ", nb1=" << ids->nb[1] << ", nb2=" << ids->nb[2] << ", nb3=" << ids->nb[3]; @@ -10991,7 +11026,9 @@ static void ggml_vk_mul_mat_id_q_f16(ggml_backend_vk_context * ctx, vk_context& hoisted_row_id_words * sizeof(uint32_t) <= ctx->device->properties.limits.maxStorageBufferRange; - ggml_backend_vk_buffer_context * dst_buf_ctx = (ggml_backend_vk_buffer_context *)dst->buffer->context; + // When the following MUL is fused in, write the scaled result straight to its destination. + const ggml_tensor * out_dst = fused_dst ? fused_dst : dst; + ggml_backend_vk_buffer_context * dst_buf_ctx = (ggml_backend_vk_buffer_context *)out_dst->buffer->context; ggml_backend_vk_buffer_context * src0_buf_ctx = (ggml_backend_vk_buffer_context *)src0->buffer->context; ggml_backend_vk_buffer_context * src1_buf_ctx = (ggml_backend_vk_buffer_context *)src1->buffer->context; ggml_backend_vk_buffer_context * ids_buf_ctx = (ggml_backend_vk_buffer_context *)ids->buffer->context; @@ -11109,6 +11146,18 @@ static void ggml_vk_mul_mat_id_q_f16(ggml_backend_vk_context * ctx, vk_context& vk_pipeline pipeline = ggml_vk_guess_matmul_id_pipeline(ctx, mmp, ne01, n_for_tile, aligned, qx_needs_dequant ? f16_type : src0->type, effective_src1_type); + // PROBE (GGML_VK_MMID_PROBE=1): which mmid tile actually runs, and with how many threads. + static const char * mmid_probe_env = getenv("GGML_VK_MMID_PROBE"); + if (mmid_probe_env && atoi(mmid_probe_env) != 0) { + static std::set seen; + std::string key = pipeline->name + ":" + std::to_string(n_for_tile); + if (seen.insert(key).second) { + fprintf(stderr, "ggml_vulkan: mmid pipeline=%s n_for_tile=%u m=%u wg=(%u,%u,%u)\n", + pipeline->name.c_str(), n_for_tile, (uint32_t)ne01, + pipeline->wg_denoms[0], pipeline->wg_denoms[1], pipeline->wg_denoms[2]); + } + } + if (ggml_nbytes(src0) > ctx->device->properties.limits.maxStorageBufferRange) { pipeline = ggml_vk_get_64b_indexing_pipeline(ctx, pipeline); } @@ -11199,7 +11248,7 @@ static void ggml_vk_mul_mat_id_q_f16(ggml_backend_vk_context * ctx, vk_context& } vk_buffer d_D = dst_buf_ctx->dev_buffer; - const uint64_t d_buf_offset = vk_tensor_offset(dst) + dst->view_offs; + const uint64_t d_buf_offset = vk_tensor_offset(out_dst) + out_dst->view_offs; GGML_ASSERT(d_D != nullptr); vk_buffer d_X; uint64_t x_buf_offset = 0; @@ -11335,7 +11384,9 @@ static void ggml_vk_mul_mat_id_q_f16(ggml_backend_vk_context * ctx, vk_context& { d_D, d_buf_offset, d_sz }, { d_ids, ids_buf_offset, ids_sz }, expert_count_buf, ne01, ne21, ne10, ne10, stride_b_y, ne01, stride_batch_x, stride_batch_y, ne20*ne21, - n_as, nei0, nei1, nbi1 / ggml_type_size(ids->type), ne11, hoist_row_ids + n_as, nei0, nei1, nbi1 / ggml_type_size(ids->type), ne11, hoist_row_ids, + fused_scale ? ggml_vk_tensor_subbuffer(ctx, fused_scale) : vk_subbuffer{ d_D, d_buf_offset, d_sz }, + fused_scale ? 1u : 0u ); // NOLINT if (x_non_contig || qx_needs_dequant) { @@ -11599,7 +11650,16 @@ static void ggml_vk_mul_mat_id(ggml_backend_vk_context * ctx, vk_context& subctx if (ggml_vk_use_mul_mat_vec_id(cgraph, node_idx)) { ggml_vk_mul_mat_vec_id_q_f16(ctx, subctx, cgraph, node_idx); } else { - ggml_vk_mul_mat_id_q_f16(ctx, subctx, src0, src1, src2, dst); + // Fused scale epilogue: the MUL's other operand is applied as the matmul writes out, + // and the result goes straight to the MUL's destination. + const ggml_tensor * fused_scale = nullptr; + ggml_tensor * fused_dst = nullptr; + if (ctx->num_additional_fused_ops == 1) { + ggml_tensor * mul = cgraph->nodes[node_idx + 1]; + fused_scale = (mul->src[0] == dst) ? mul->src[1] : mul->src[0]; + fused_dst = mul; + } + ggml_vk_mul_mat_id_q_f16(ctx, subctx, src0, src1, src2, dst, fused_scale, fused_dst); } } @@ -12962,6 +13022,31 @@ static vk_conv_shapes ggml_vk_conv_select_shape(ggml_backend_vk_context * ctx, u } } +// EXPERIMENT (GGML_VK_CONCAT_TRANSPOSE=1): the delta-net conv-state path does +// ggml_transpose() straight into a dim-0 ggml_concat(), so the generic concat kernel reads +// src1 fully de-coalesced. Measured on Qwen3.6-35B-A3B: CONCAT is ~22% of pp2048 at ub=2048 +// and grows 3.1x for a 2x ubatch. Route that exact shape to a tiled-transpose kernel. +static bool ggml_vk_concat_is_transposed(const ggml_tensor * src0, const ggml_tensor * src1, const ggml_tensor * dst) { + static const char * env = getenv("GGML_VK_CONCAT_TRANSPOSE"); + if (!(env && atoi(env) != 0)) { + return false; + } + if (ggml_get_op_params_i32(dst, 0) != 0) { // dim 0 only + return false; + } + if (src0->ne[2] != 1 || src0->ne[3] != 1 || src1->ne[2] != 1 || src1->ne[3] != 1) { + return false; + } + const size_t ts = ggml_type_size(src0->type); + if (src0->nb[0] != ts || dst->nb[0] != ts) { // src0 and dst rows must be contiguous + return false; + } + if (src1->nb[0] <= src1->nb[1]) { // src1 must actually be transposed + return false; + } + return src0->ne[1] == src1->ne[1] && dst->ne[1] == src1->ne[1]; +} + static vk_pipeline ggml_vk_op_get_pipeline(ggml_backend_vk_context * ctx, const ggml_tensor * src0, const ggml_tensor * src1, const ggml_tensor * src2, const ggml_tensor * dst, ggml_op op) { switch (op) { case GGML_OP_GET_ROWS: @@ -13054,6 +13139,11 @@ static vk_pipeline ggml_vk_op_get_pipeline(ggml_backend_vk_context * ctx, const if (!ggml_vk_concat_supported(src0, src1, dst)) { return nullptr; } + // Tiled-transpose path handles unquantized 4-byte elements only. + if (!ggml_is_quantized(src0->type) && ggml_vk_concat_unit_size(src0->type) == 4 && + ggml_vk_concat_is_transposed(src0, src1, dst)) { + return ctx->device->pipeline_concat_transpose_i32; + } switch (ggml_vk_concat_unit_size(src0->type)) { case 1: return ctx->device->pipeline_concat_i8; @@ -14126,6 +14216,11 @@ static void ggml_vk_op_f32(ggml_backend_vk_context * ctx, vk_context& subctx, co case GGML_OP_GLU: case GGML_OP_CONV_2D_DW: { + // The tiled concat kernel is dispatched per 32x32 tile, not per element. + if (op == GGML_OP_CONCAT && pipeline == ctx->device->pipeline_concat_transpose_i32) { + elements = { (uint32_t)src1->ne[1], (uint32_t)src1->ne[0], 1 }; + break; + } uint32_t ne = ggml_nelements(dst); if (op == GGML_OP_CPY && ggml_is_quantized(src0->type) && ggml_is_quantized(dst->type)) { // Convert from number of logical elements to 2- or 4-byte units. @@ -17816,6 +17911,22 @@ static bool ggml_vk_build_graph(ggml_backend_vk_context * ctx, ggml_cgraph * cgr break; } + // Fused silu(x)*y: run it as a swiglu split, writing straight to the MUL's destination. + if (ctx->num_additional_fused_ops == 1) { + ggml_tensor * mul = cgraph->nodes[node_idx + 1]; + ggml_tensor * other = (mul->src[0] == node) ? mul->src[1] : mul->src[0]; + + ggml_tensor fused = *mul; + fused.op = GGML_OP_GLU; + memset(fused.op_params, 0, sizeof(fused.op_params)); + ggml_set_op_params_i32(&fused, 0, (int32_t) GGML_GLU_OP_SWIGLU); + fused.src[0] = node->src[0]; + fused.src[1] = other; + + ggml_vk_glu(ctx, compute_ctx, fused.src[0], fused.src[1], &fused); + break; + } + switch (ggml_get_unary_op(node)) { case GGML_UNARY_OP_ELU: case GGML_UNARY_OP_EXP: @@ -18851,15 +18962,61 @@ static bool ggml_vk_can_fuse(const ggml_backend_vk_context * ctx, const struct g } } + // EXPERIMENT (GGML_VK_FUSE_UNARY_MUL=1): silu(x)*y is emitted as two nodes by the delta-net + // path, so the silu result makes a full round trip through memory. That is the same shape + // swiglu-split already computes in one pass, so route the pair to the existing GLU pipeline. + if (ops.size() == 2 && ops.begin()[0] == GGML_OP_UNARY && ops.begin()[1] == GGML_OP_MUL) { + static const char * env = getenv("GGML_VK_FUSE_UNARY_MUL"); + if (!(env && atoi(env) != 0)) { + return false; + } + const ggml_tensor * unary = cgraph->nodes[node_idx]; + const ggml_tensor * mul = cgraph->nodes[node_idx + 1]; + + if (ggml_get_unary_op(unary) != GGML_UNARY_OP_SILU) { + return false; + } + if (mul->src[0] != unary && mul->src[1] != unary) { + return false; + } + const ggml_tensor * other = (mul->src[0] == unary) ? mul->src[1] : mul->src[0]; + // The GLU split shader walks both inputs and the output with the same element count. + if (unary->type != GGML_TYPE_F32 || other->type != GGML_TYPE_F32 || mul->type != GGML_TYPE_F32) { + return false; + } + if (!ggml_are_same_shape(unary, other) || !ggml_are_same_shape(unary, mul)) { + return false; + } + if (!ggml_is_contiguous(unary->src[0]) || !ggml_is_contiguous(other) || !ggml_is_contiguous(mul)) { + return false; + } + return true; + } + auto const &mmid_mul_ok = [&](const ggml_tensor *mmid, const ggml_tensor *mul) { const ggml_tensor *scale = mul->src[1]; if (mmid != mul->src[0]) { return false; } - // mat-vec only + // EXPERIMENT (GGML_VK_MMID_SCALE_EPILOGUE=1): the tile shader can apply the scale as it + // writes out, which removes a full write+read of the matmul result at prefill. The + // coopmat2 shader has the binding but not the epilogue, so it stays on the old path. if (!ggml_vk_use_mul_mat_vec_id(cgraph, node_idx)) { - return false; + static const char * env = getenv("GGML_VK_MMID_SCALE_EPILOGUE"); + if (!(env && atoi(env) != 0) || ctx->device->coopmat2) { + return false; + } + // Shader indexes the scale as [token * nei0 + expert_slot]. + if (scale->type != GGML_TYPE_F32 || mul->type != GGML_TYPE_F32 || !ggml_is_contiguous(scale)) { + return false; + } + if (get_misalign_bytes(ctx, scale) != 0) { + return false; + } + return scale->ne[0] == 1 && + scale->ne[1] == mmid->ne[1] && scale->ne[2] == mmid->ne[2] && scale->ne[3] == mmid->ne[3] && + ggml_are_same_shape(mul, mmid); } // shaders assume the types match if (mmid->type != scale->type) { @@ -19604,6 +19761,9 @@ static ggml_status ggml_backend_vk_graph_compute(ggml_backend_t backend, ggml_cg fusion_string = "MUL_MAT_ID_MUL"; op_srcs_fused_elementwise[0] = false; op_srcs_fused_elementwise[1] = true; + } else if (ggml_vk_can_fuse(ctx, cgraph, i, { GGML_OP_UNARY, GGML_OP_MUL })) { + ctx->num_additional_fused_ops = 1; + fusion_string = "SILU_MUL"; } else if (ggml_can_fuse_subgraph(cgraph, i, { GGML_OP_RMS_NORM, GGML_OP_MUL, GGML_OP_ROPE, GGML_OP_VIEW, GGML_OP_SET_ROWS }, { i + 4 }) && ggml_check_edges(cgraph, i, rms_norm_mul_rope_view_set_rows_edges) && ggml_vk_can_fuse_rms_norm_mul_rope(ctx, cgraph, i) && diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/concat_transpose.comp b/ggml/src/ggml-vulkan/vulkan-shaders/concat_transpose.comp new file mode 100644 index 00000000000..653aeaa011b --- /dev/null +++ b/ggml/src/ggml-vulkan/vulkan-shaders/concat_transpose.comp @@ -0,0 +1,43 @@ +#version 450 + +#include "types.glsl" +#include "generic_binary_head.glsl" + +layout(local_size_x = 32, local_size_y = 8, local_size_z = 1) in; + +// dim-0 concat whose src1 is transposed. The generic kernel reads src1 with the transposed +// stride, so neighbouring lanes touch different cache lines. Stage a 32x32 tile in shared +// memory instead, which keeps both the load and the store coalesced. 33 columns pads away +// the shared-memory bank conflicts. +shared A_TYPE tmp[32][33]; + +void main() { + const uint tx = gl_LocalInvocationID.x; + const uint ty = gl_LocalInvocationID.y; + + const uint row = gl_WorkGroupID.x * 32 + tx; + + // src0 is already contiguous, copy it straight through. + if (gl_WorkGroupID.y == 0 && row < p.ne01) { + for (uint i0 = ty; i0 < p.ne00; i0 += 8) { + data_d[get_doffset() + row*p.nb21 + i0*p.nb20] = D_TYPE(data_a[get_aoffset() + row*p.nb01 + i0*p.nb00]); + } + } + + [[unroll]] for (uint j = 0; j < 32; j += 8) { + const uint c = gl_WorkGroupID.y * 32 + ty + j; + if (c < p.ne10 && row < p.ne11) { + tmp[ty + j][tx] = A_TYPE(data_b[get_boffset() + c*p.nb10 + row*p.nb11]); + } + } + + barrier(); + + const uint col = gl_WorkGroupID.y * 32 + tx; + [[unroll]] for (uint j = 0; j < 32; j += 8) { + const uint r = gl_WorkGroupID.x * 32 + ty + j; + if (col < p.ne10 && r < p.ne11) { + data_d[get_doffset() + r*p.nb21 + (p.ne00 + col)*p.nb20] = D_TYPE(tmp[tx][ty + j]); + } + } +} diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm.comp b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm.comp index 63c4aaebcb1..7d32124b969 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm.comp @@ -68,6 +68,9 @@ layout (binding = 2) writeonly buffer D {D_TYPE data_d[];}; #ifdef MUL_MAT_ID layout (binding = 3) readonly buffer IDS {int data_ids[];}; layout (binding = 4) readonly buffer Counts {int data_expert_count[];}; +// Fused MUL epilogue: one scale per (expert slot, token), broadcast down the M dimension. +// Always bound; p.fusion_flags == 0 means ignore it. +layout (binding = 5) readonly buffer FusedScale {float data_fscale[];}; #endif layout (push_constant) uniform parameter @@ -90,6 +93,7 @@ layout (push_constant) uniform parameter uint ne11; uint n_experts; uint hoist_row_ids; + uint fusion_flags; #else uint base_work_group_z; uint num_batches; @@ -394,9 +398,13 @@ void main() { if (row_i >= _ne1) break; const u16vec2 row_idx = row_ids[row_i - ic * BN]; - if (dr + cm_row * TM + store_r < p.M) { - data_d[row_idx.y * p.batch_stride_d + row_idx.x * p.stride_d + dr + cm_row * TM + store_r] = D_TYPE(coopmat_stage[warp_i * TM * TN + (col + store_c) * TM + store_r]); + const uint didx = row_idx.y * p.batch_stride_d + row_idx.x * p.stride_d + dr + cm_row * TM + store_r; + if (p.fusion_flags != 0) { + data_d[didx] = D_TYPE(float(coopmat_stage[warp_i * TM * TN + (col + store_c) * TM + store_r]) * data_fscale[row_idx.y * p.nei0 + row_idx.x]); + } else { + data_d[didx] = D_TYPE(coopmat_stage[warp_i * TM * TN + (col + store_c) * TM + store_r]); + } } } barrier(); @@ -449,15 +457,19 @@ void main() { if (row_i >= _ne1) break; const u16vec2 row_idx = row_ids[row_i - ic * BN]; + const bool do_scale = p.fusion_flags != 0; + const float fscale = do_scale ? data_fscale[row_idx.y * p.nei0 + row_idx.x] : 1.0f; #endif // MUL_MAT_ID [[unroll]] for (uint cr = 0; cr < TM / 2; cr++) { const uint sums_idx = (wsic * TN + cc) * WMITER * (TM / 2) + wsir * (TM / 2) + cr; #ifdef MUL_MAT_ID if (dr_warp + 2 * cr < p.M) { - data_d[row_idx.y * p.batch_stride_d + row_idx.x * p.stride_d + dr_warp + 2 * cr] = D_TYPE(sums[sums_idx].x); + const uint didx = row_idx.y * p.batch_stride_d + row_idx.x * p.stride_d + dr_warp + 2 * cr; + data_d[didx] = do_scale ? D_TYPE(float(sums[sums_idx].x) * fscale) : D_TYPE(sums[sums_idx].x); } if (dr_warp + 2 * cr + 1 < p.M) { - data_d[row_idx.y * p.batch_stride_d + row_idx.x * p.stride_d + dr_warp + 2 * cr + 1] = D_TYPE(sums[sums_idx].y); + const uint didx = row_idx.y * p.batch_stride_d + row_idx.x * p.stride_d + dr_warp + 2 * cr + 1; + data_d[didx] = do_scale ? D_TYPE(float(sums[sums_idx].y) * fscale) : D_TYPE(sums[sums_idx].y); } #else if (dr_warp + 2 * cr < p.M && dc_warp + cc < p.N) { diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm_cm2.comp b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm_cm2.comp index 27f3178e7f2..9bfad031d1b 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm_cm2.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm_cm2.comp @@ -106,6 +106,9 @@ layout (binding = 1) readonly buffer B4 {B_TYPEV4 data_b_v4[];}; #ifdef MUL_MAT_ID layout (binding = 3) readonly buffer IDS {int data_ids[];}; layout (binding = 4) readonly buffer Counts {int data_expert_count[];}; +// Bound for descriptor-layout parity with the other mul_mat_id shaders. The fused MUL +// epilogue is not implemented here, so the host never enables it on coopmat2 devices. +layout (binding = 5) readonly buffer FusedScale {float data_fscale[];}; shared u16vec4 row_ids[BN]; diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq.comp b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq.comp index 1fbcbf6c933..b36d056c1c9 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq.comp @@ -36,6 +36,8 @@ layout (binding = 2) writeonly buffer D {D_TYPE data_d[];}; #ifdef MUL_MAT_ID layout (binding = 3) readonly buffer IDS {int data_ids[];}; layout (binding = 4) readonly buffer Counts {int data_expert_count[];}; +// Fused MUL epilogue, see mul_mm.comp. +layout (binding = 5) readonly buffer FusedScale {float data_fscale[];}; #endif layout (push_constant) uniform parameter @@ -58,6 +60,7 @@ layout (push_constant) uniform parameter uint ne11; uint n_experts; uint hoist_row_ids; + uint fusion_flags; #else uint base_work_group_z; uint num_batches; @@ -303,7 +306,10 @@ void main() { const uint sums_idx = (wsic * TN + cc) * WMITER * TM + wsir * TM + cr; #ifdef MUL_MAT_ID if (dr_warp + cr < p.M) { - data_d[row_idx.y * p.batch_stride_d + row_idx.x * p.stride_d + dr_warp + cr] = D_TYPE(sums[sums_idx].x); + const uint didx = row_idx.y * p.batch_stride_d + row_idx.x * p.stride_d + dr_warp + cr; + data_d[didx] = p.fusion_flags != 0 + ? D_TYPE(float(sums[sums_idx].x) * data_fscale[row_idx.y * p.nei0 + row_idx.x]) + : D_TYPE(sums[sums_idx].x); } #else if (dr_warp + cr < p.M && dc_warp + cc < p.N) { 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 fe82b8fb6cb..295e17b7b08 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp @@ -951,6 +951,7 @@ void process_shaders() { string_to_spv("concat_i16", "concat.comp", {{"A_TYPE", "uint16_t"}, {"B_TYPE", "uint16_t"}, {"D_TYPE", "uint16_t"}}); string_to_spv("concat_i32", "concat.comp", {{"A_TYPE", "uint"}, {"B_TYPE", "uint"}, {"D_TYPE", "uint"}}); string_to_spv("concat_i64", "concat.comp", {{"A_TYPE", "uvec2"}, {"B_TYPE", "uvec2"}, {"D_TYPE", "uvec2"}}); + string_to_spv("concat_transpose_i32", "concat_transpose.comp", {{"A_TYPE", "uint"}, {"B_TYPE", "uint"}, {"D_TYPE", "uint"}}); string_to_spv("upscale_f32", "upscale.comp", {{"A_TYPE", "float"}, {"B_TYPE", "float"}, {"D_TYPE", "float"}}); From 874b9043d1b4e1761219bfcd2ec7f0c1e37703f1 Mon Sep 17 00:00:00 2001 From: Nathan Wilson Date: Thu, 6 Aug 2026 08:48:36 +0000 Subject: [PATCH 60/68] vulkan: reject ne[3] > 1 in the mul_mat_id scale epilogue The fused epilogue derives the scale index from row_ids as [token * nei0 + expert_slot], which carries no 4th dimension, but the gate admitted any ne[3] as long as the scale and the matmul agreed. A tensor with ne[3] > 1 would read the wrong scale for every batch past the first and return quietly wrong results. test-backend-ops never generates such a case, so the suite passed throughout; found by reading the gate against the shader. Co-Authored-By: Claude Opus 5 (cherry picked from commit 016e906788057b9734ab7727de645d58f8080716) Assisted-by: Claude (Opus 5) --- ggml/src/ggml-vulkan/ggml-vulkan.cpp | 6 ++++-- 1 file changed, 4 insertions(+), 2 deletions(-) diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp index 8f297cb6696..c51440149d7 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp +++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp @@ -19014,8 +19014,10 @@ static bool ggml_vk_can_fuse(const ggml_backend_vk_context * ctx, const struct g if (get_misalign_bytes(ctx, scale) != 0) { return false; } - return scale->ne[0] == 1 && - scale->ne[1] == mmid->ne[1] && scale->ne[2] == mmid->ne[2] && scale->ne[3] == mmid->ne[3] && + // The shader indexes the scale as [token * nei0 + expert_slot] from row_ids, which + // carries no 4th dimension, so ne[3] must be 1 or later batches read the wrong scale. + return scale->ne[0] == 1 && mmid->ne[3] == 1 && scale->ne[3] == 1 && + scale->ne[1] == mmid->ne[1] && scale->ne[2] == mmid->ne[2] && ggml_are_same_shape(mul, mmid); } // shaders assume the types match From e249248565e8fbd736f9ce9e2db5352734ad5844 Mon Sep 17 00:00:00 2001 From: Nathan Wilson Date: Mon, 17 Aug 2026 06:13:27 +0000 Subject: [PATCH 61/68] vulkan : default the transposed-concat path on The delta-net conv-state path transposes straight into a dim-0 concat, so the generic concat kernel walks src1 with a conv_channels * 4 byte stride. On qwen35 that is 40960 B, which is 160 * 256 B with 160 % 16 == 0, so every read lands on the same one of the 16 memory channels: 13.7 GB/s against 138.9 GB/s for the tiled path. The tiled-transpose route has been behind GGML_VK_CONCAT_TRANSPOSE=1 since it landed. Turn it on by default and keep GGML_VK_CONCAT_TRANSPOSE=0 as the opt-out. Qwen3.8-27B pp2048: +0.4% at ub 256, +4.7% at ub 1024, +7.2% at ub 2048. Assisted-by: Claude Opus 5 --- ggml/src/ggml-vulkan/ggml-vulkan.cpp | 12 +++++++----- 1 file changed, 7 insertions(+), 5 deletions(-) diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp index c51440149d7..52f986fb92e 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp +++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp @@ -13022,13 +13022,15 @@ static vk_conv_shapes ggml_vk_conv_select_shape(ggml_backend_vk_context * ctx, u } } -// EXPERIMENT (GGML_VK_CONCAT_TRANSPOSE=1): the delta-net conv-state path does -// ggml_transpose() straight into a dim-0 ggml_concat(), so the generic concat kernel reads -// src1 fully de-coalesced. Measured on Qwen3.6-35B-A3B: CONCAT is ~22% of pp2048 at ub=2048 -// and grows 3.1x for a 2x ubatch. Route that exact shape to a tiled-transpose kernel. +// The delta-net conv-state path does ggml_transpose() straight into a dim-0 ggml_concat(), so +// the generic concat kernel walks src1 with a conv_channels * 4 byte stride. On qwen35 that is +// 40960 B = 160 * 256 B and 160 % 16 == 0, so every read lands on one of the 16 memory channels: +// 13.7 GB/s against 138.9 GB/s for the tiled path. Route that exact shape to a tiled-transpose +// kernel. On by default; GGML_VK_CONCAT_TRANSPOSE=0 opts out. +// Qwen3.8-27B pp2048: +0.4% at ub 256, +4.7% at ub 1024, +7.2% at ub 2048. static bool ggml_vk_concat_is_transposed(const ggml_tensor * src0, const ggml_tensor * src1, const ggml_tensor * dst) { static const char * env = getenv("GGML_VK_CONCAT_TRANSPOSE"); - if (!(env && atoi(env) != 0)) { + if (env && env[0] == '0') { return false; } if (ggml_get_op_params_i32(dst, 0) != 0) { // dim 0 only From 7db0111d53cbcc6750f61c4d2fdcb10bc01ce7d0 Mon Sep 17 00:00:00 2001 From: Nathan Wilson Date: Sun, 2 Aug 2026 00:13:47 +0000 Subject: [PATCH 62/68] vulkan: flush pending compute ctx before perf logger timestamps The scheduler's async input copies between graph splits land in the compute ctx on devices without a separate transfer queue, so the perf logger's fresh-ctx assert fired under partial offload (--n-cpu-moe). Assisted-by: Claude Fable 5 --- ggml/src/ggml-vulkan/ggml-vulkan.cpp | 6 ++++++ 1 file changed, 6 insertions(+) diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp index 52f986fb92e..a4f619b2061 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp +++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp @@ -19630,6 +19630,12 @@ static ggml_status ggml_backend_vk_graph_compute(ggml_backend_t backend, ggml_cg std::fill(ctx->query_nodes.begin(), ctx->query_nodes.end(), nullptr); std::fill(ctx->query_node_idx.begin(), ctx->query_node_idx.end(), 0); + // Under partial offload the scheduler's async input copies between graph + // splits can leave commands in a pending compute ctx. Flush it so the + // timestamp stream starts on a fresh command buffer. + if (!ctx->compute_ctx.expired()) { + ggml_vk_synchronize(ctx); + } GGML_ASSERT(ctx->compute_ctx.expired()); compute_ctx = ggml_vk_get_compute_ctx(ctx); ctx->query_idx = 0; From 6ab7cb6b49316488600461dab94748b73290a6bb Mon Sep 17 00:00:00 2001 From: Nathan Wilson Date: Sun, 2 Aug 2026 11:56:37 +0000 Subject: [PATCH 63/68] vulkan: bound command buffers by memory traffic, not just flops Nodes with no flops estimate (large copies, set_rows, mask fills) can pack a command buffer whose execution time grows with context length until it exceeds the amdgpu ring timeout (10s on the compute ring), causing the ring resets and DeviceLost reported at long context. Add a bytes-per-submit cap (default 8 GiB, GGML_VK_MAX_MB_PER_SUBMIT to override, 0 disables) alongside the existing flops and node-count gates. Assisted-by: Claude Fable 5 --- ggml/src/ggml-vulkan/ggml-vulkan.cpp | 34 ++++++++++++++++++++++++++-- 1 file changed, 32 insertions(+), 2 deletions(-) diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp index a4f619b2061..3e954312372 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp +++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp @@ -869,6 +869,7 @@ struct vk_device_struct { bool add_rms_fusion; uint32_t partials_binding_alignment; uint32_t max_nodes_per_submit; + uint64_t max_bytes_per_submit; bool shader_64b_indexing; @@ -2302,6 +2303,20 @@ static bool vk_enable_sync_logger = false; static uint32_t vk_perf_logger_frequency = 1; static std::string vk_pipeline_stats_filter; +// Total memory traffic of a node (dst + srcs). Used to bound command buffer +// execution time for bandwidth-bound ops with no flops estimate (large copies, +// set_rows, mask fills at long context) - packing too many of them into one +// submission can exceed the driver timeout. +static uint64_t ggml_vk_get_node_bytes(const ggml_tensor * node) { + uint64_t bytes = ggml_nbytes(node); + for (int i = 0; i < GGML_MAX_SRC; i++) { + if (node->src[i]) { + bytes += ggml_nbytes(node->src[i]); + } + } + return bytes; +} + static uint64_t ggml_vk_get_node_flops(const ggml_tensor * node) { if (node->op == GGML_OP_MUL_MAT || node->op == GGML_OP_MUL_MAT_ID) { const uint64_t m = node->ne[0]; @@ -7225,6 +7240,15 @@ static vk_device ggml_vk_get_device(size_t idx) { device->max_nodes_per_submit = std::max(max_nodes_per_submit, 1u); } + // Also submit once a batch has accumulated enough memory traffic, so that + // bandwidth-bound nodes with no flops estimate cannot grow a command buffer + // past the driver timeout. 0 disables the limit. + device->max_bytes_per_submit = 8ull * 1024 * 1024 * 1024; + const char* GGML_VK_MAX_MB_PER_SUBMIT = getenv("GGML_VK_MAX_MB_PER_SUBMIT"); + if (GGML_VK_MAX_MB_PER_SUBMIT != nullptr) { + device->max_bytes_per_submit = std::stoull(GGML_VK_MAX_MB_PER_SUBMIT) * 1024 * 1024; + } + const bool force_disable_f16 = getenv("GGML_VK_DISABLE_F16") != nullptr; device->fp16 = !force_disable_f16 && fp16_storage && fp16_compute; @@ -19662,6 +19686,7 @@ static ggml_status ggml_backend_vk_graph_compute(ggml_backend_t backend, ggml_cg uint32_t submitted_nodes = 0; uint32_t submit_count = 0; uint64_t batch_flops = 0; + uint64_t batch_bytes = 0; uint64_t total_flops = 0; uint64_t flops_cap = 200'000'000'000ULL; @@ -19699,6 +19724,7 @@ static ggml_status ggml_backend_vk_graph_compute(ggml_backend_t backend, ggml_cg first_node_in_batch = true; submitted_nodes = 0; batch_flops = 0; + batch_bytes = 0; if (submit_count < 3) { flops_per_submit *= 2; } @@ -19713,9 +19739,11 @@ static ggml_status ggml_backend_vk_graph_compute(ggml_backend_t backend, ggml_cg { auto node_flops = ggml_vk_get_node_flops(cgraph->nodes[i]); total_flops += node_flops; + auto node_bytes = ggml_vk_get_node_bytes(cgraph->nodes[i]); - // Flush the current batch before recording a node that would push it over the flop threshold - if (flops_per_submit != 0 && submitted_nodes > 0 && batch_flops + node_flops >= flops_per_submit) { + // Flush the current batch before recording a node that would push it over the flop or byte threshold + if ((flops_per_submit != 0 && submitted_nodes > 0 && batch_flops + node_flops >= flops_per_submit) || + (ctx->device->max_bytes_per_submit != 0 && submitted_nodes > 0 && batch_bytes + node_bytes >= ctx->device->max_bytes_per_submit)) { vk_context flush_ctx = ggml_vk_get_compute_ctx(ctx); ggml_vk_ctx_end(flush_ctx); flush_ctx->exit_tensor_idx = -1; @@ -19726,6 +19754,7 @@ static ggml_status ggml_backend_vk_graph_compute(ggml_backend_t backend, ggml_cg } batch_flops += node_flops; + batch_bytes += node_bytes; } // op_srcs_fused_elementwise indicates whether an op's srcs all contribute to @@ -19958,6 +19987,7 @@ static ggml_status ggml_backend_vk_graph_compute(ggml_backend_t backend, ggml_cg bool almost_ready = (cgraph->n_nodes - i) < cgraph->n_nodes / 5; bool submit = (submitted_nodes >= ctx->device->max_nodes_per_submit) || (flops_per_submit != 0 && batch_flops >= flops_per_submit) || + (ctx->device->max_bytes_per_submit != 0 && batch_bytes >= ctx->device->max_bytes_per_submit) || (i + ctx->num_additional_fused_ops >= last_node) || (almost_ready && !ctx->almost_ready_fence_pending); From e38725201c560ab333e561c830ce960060a1fcce Mon Sep 17 00:00:00 2001 From: Nathan Wilson Date: Sun, 9 Aug 2026 02:48:56 +0000 Subject: [PATCH 64/68] ggml: cut backend splits on the input constant, not the grown capacity #22789 replaced the fixed 30-entry split input array with a growable one and, in the same edit, changed the split-cutting heuristic from the constant to split->inputs_capacity: - if (split->n_inputs == GGML_SCHED_MAX_SPLIT_INPUTS) { + if (split->n_inputs >= split->inputs_capacity) { inputs_capacity starts at GGML_SCHED_MAX_SPLIT_INPUTS but doubles on demand and is never reset for the life of the sched, so once a split slot grows, the scheduler stops cutting there and the cut point ratchets up for every later graph build. Longer splits mean every cross-backend input copy is materialised at the split's start and stays live to its last use inside it, which raises the peak the compute-buffer allocator has to cover - n_copies times over under pipeline parallelism. Only multi-backend configurations can reach this. Keep the growable array, which is what fixes the original >30-input assert, and cut on the constant again as before #22789. >= rather than == so the check keeps firing for splits that did have to grow. DeepSeek-V4-Flash UD-IQ3_XXS, gfx1151, -c 400000 -ub 2048 -fa 1 --fit off: Vulkan0 compute buffer 4714.00 MiB and 9157 graph nodes, byte-identical to the unpatched tree, and neither run grows a split past 30 inputs. Expected - one Vulkan device plus the CPU backend cannot exercise the path on this box. The reported case is 3 devices with pipeline parallelism. test-backend-ops -o FLASH_ATTN_EXT, run alone on gfx1151: 13257/13295 on both this and the unpatched tree, with the same 38 failing cases (identical case list, all type_K=q8_0 prec=def kv_view=1). Pre-existing on the branch, not touched by this change. Co-Authored-By: Claude Opus 5 Assisted-by: Claude (Opus 5) --- ggml/src/ggml-backend.cpp | 5 ++++- 1 file changed, 4 insertions(+), 1 deletion(-) diff --git a/ggml/src/ggml-backend.cpp b/ggml/src/ggml-backend.cpp index dc6b98946e2..0955a431b17 100644 --- a/ggml/src/ggml-backend.cpp +++ b/ggml/src/ggml-backend.cpp @@ -1446,7 +1446,10 @@ void ggml_backend_sched_split_graph(ggml_backend_sched_t sched, struct ggml_cgra } // check if the split has too many inputs // FIXME: count the number of inputs instead of only checking when full - if (split->n_inputs >= split->inputs_capacity) { + // cut on the constant, not on inputs_capacity: capacity doubles on demand and + // is never reset, so using it lets the cut point drift up and keeps every input + // copy of an ever-longer split live at once + if (split->n_inputs >= GGML_SCHED_MAX_SPLIT_INPUTS) { const size_t id = hash_id(src); int src_backend_id = sched->hv_tensor_backend_ids[id]; bool supported = ggml_backend_sched_buffer_supported(sched, src, cur_backend_id); From 492e443eac327a170104ac37147b8bd20e364c5f Mon Sep 17 00:00:00 2001 From: Nathan Wilson Date: Sun, 16 Aug 2026 11:49:48 +0000 Subject: [PATCH 65/68] vulkan : optional f16 B operand for quantized MUL_MAT on coopmat1 The _f16 SPIR-V and pipeline_dequant_mul_mat_mat_f16 already exist, but are only populated in the coopmat2 branch. GGML_VK_DENSE_F16B populates them for coopmat1 too and routes B through the existing convert-to-prealloc_y path. Off by default: it helps large dense models and costs a little elsewhere. gfx1151, pp2048: Qwen3.8-27B and Qwen3-32B (hidden 5120) +5 to +7% for both q6_K and q8_0 weights, Qwen2.5-7B -1.2%, Qwen3-Coder-30B MoE -0.5%. Decode is untouched, ne1==1 does not reach this path. Numerically identical: mul_mm stages B into shared FLOAT_TYPE either way, so the f32-B kernel already rounds B to f16. Wikitext PPL matches to 4 dp. Assisted-by: Claude Opus 5 --- ggml/src/ggml-vulkan/ggml-vulkan.cpp | 45 ++++++++++++++++++++++++++-- 1 file changed, 42 insertions(+), 3 deletions(-) diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp index 3e954312372..8eee6436d7f 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp +++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp @@ -4498,6 +4498,11 @@ static bool ggml_vk_mmid_f16b_enabled() { const char * env = getenv("GGML_VK_MMID_F16B"); return env == nullptr || atoi(env) != 0; // on by default; =0 disables }(); +// GGML_VK_DENSE_F16B=1: same idea as GGML_VK_MMID_F16B but for plain MUL_MAT. Halves the B +// bytes, which keeps the activations inside the LLC at large ubatch and moves their row stride +// off the 1-of-16 channel pattern. Off by default, it is a loss when B already fits. +static bool ggml_vk_dense_f16b_enabled() { + static const bool enabled = getenv("GGML_VK_DENSE_F16B") != nullptr; return enabled; } @@ -5205,6 +5210,21 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { CREATE_MM2(GGML_TYPE_Q4_K, pipeline_dequant_mul_mat_mat[GGML_TYPE_Q4_K], matmul_q4_k_f32, mmq_wg_denoms, warptile_mmq, vk_mat_mat_push_constants, 3, ); CREATE_MM2(GGML_TYPE_Q5_K, pipeline_dequant_mul_mat_mat[GGML_TYPE_Q5_K], matmul_q5_k_f32, mmq_wg_denoms, warptile_mmq, vk_mat_mat_push_constants, 3, ); CREATE_MM2(GGML_TYPE_Q6_K, pipeline_dequant_mul_mat_mat[GGML_TYPE_Q6_K], matmul_q6_k_f32, mmq_wg_denoms, warptile_mmq, vk_mat_mat_push_constants, 3, ); + + // f16-B variants of the same quant pipelines. The _f16 SPIR-V is already built for + // every type; upstream only instantiates it in the coopmat2 branch. + if (ggml_vk_dense_f16b_enabled()) { + CREATE_MM2(GGML_TYPE_Q4_0, pipeline_dequant_mul_mat_mat_f16[GGML_TYPE_Q4_0], matmul_q4_0_f16, mmq_wg_denoms, warptile_mmq, vk_mat_mat_push_constants, 3, ); + CREATE_MM2(GGML_TYPE_Q4_1, pipeline_dequant_mul_mat_mat_f16[GGML_TYPE_Q4_1], matmul_q4_1_f16, mmq_wg_denoms, warptile_mmq, vk_mat_mat_push_constants, 3, ); + CREATE_MM2(GGML_TYPE_Q5_0, pipeline_dequant_mul_mat_mat_f16[GGML_TYPE_Q5_0], matmul_q5_0_f16, mmq_wg_denoms, warptile_mmq, vk_mat_mat_push_constants, 3, ); + CREATE_MM2(GGML_TYPE_Q5_1, pipeline_dequant_mul_mat_mat_f16[GGML_TYPE_Q5_1], matmul_q5_1_f16, mmq_wg_denoms, warptile_mmq, vk_mat_mat_push_constants, 3, ); + CREATE_MM2(GGML_TYPE_Q8_0, pipeline_dequant_mul_mat_mat_f16[GGML_TYPE_Q8_0], matmul_q8_0_f16, mmq_wg_denoms, warptile_mmq, vk_mat_mat_push_constants, 3, ); + CREATE_MM2(GGML_TYPE_Q2_K, pipeline_dequant_mul_mat_mat_f16[GGML_TYPE_Q2_K], matmul_q2_k_f16, mmq_wg_denoms, warptile_mmq, vk_mat_mat_push_constants, 3, ); + CREATE_MM2(GGML_TYPE_Q3_K, pipeline_dequant_mul_mat_mat_f16[GGML_TYPE_Q3_K], matmul_q3_k_f16, mmq_wg_denoms, warptile_mmq, vk_mat_mat_push_constants, 3, ); + CREATE_MM2(GGML_TYPE_Q4_K, pipeline_dequant_mul_mat_mat_f16[GGML_TYPE_Q4_K], matmul_q4_k_f16, mmq_wg_denoms, warptile_mmq, vk_mat_mat_push_constants, 3, ); + CREATE_MM2(GGML_TYPE_Q5_K, pipeline_dequant_mul_mat_mat_f16[GGML_TYPE_Q5_K], matmul_q5_k_f16, mmq_wg_denoms, warptile_mmq, vk_mat_mat_push_constants, 3, ); + CREATE_MM2(GGML_TYPE_Q6_K, pipeline_dequant_mul_mat_mat_f16[GGML_TYPE_Q6_K], matmul_q6_k_f16, mmq_wg_denoms, warptile_mmq, vk_mat_mat_push_constants, 3, ); + } CREATE_MM2(GGML_TYPE_IQ1_S, pipeline_dequant_mul_mat_mat[GGML_TYPE_IQ1_S], matmul_iq1_s_f32, mmq_wg_denoms, warptile_mmq, vk_mat_mat_push_constants, 3, ); CREATE_MM2(GGML_TYPE_IQ1_M, pipeline_dequant_mul_mat_mat[GGML_TYPE_IQ1_M], matmul_iq1_m_f32, mmq_wg_denoms, warptile_mmq, vk_mat_mat_push_constants, 3, ); CREATE_MM2(GGML_TYPE_IQ2_XXS, pipeline_dequant_mul_mat_mat[GGML_TYPE_IQ2_XXS], matmul_iq2_xxs_f32, mmq_wg_denoms, warptile_mmq, vk_mat_mat_push_constants, 3, ); @@ -8532,7 +8552,8 @@ static vk_matmul_pipeline ggml_vk_get_mul_mat_mat_pipeline(ggml_backend_vk_conte return pipelines; } - if (src1_type != GGML_TYPE_F32 && !ctx->device->coopmat2) { + if (src1_type != GGML_TYPE_F32 && !ctx->device->coopmat2 && + !(src1_type == GGML_TYPE_F16 && ggml_vk_dense_f16b_enabled())) { return nullptr; } @@ -8576,7 +8597,9 @@ static vk_matmul_pipeline ggml_vk_get_mul_mat_mat_pipeline(ggml_backend_vk_conte return prec == GGML_PREC_DEFAULT ? ctx->device->pipeline_dequant_mul_mat_mat_f16[src0_type].f16acc : ctx->device->pipeline_dequant_mul_mat_mat_f16[src0_type].f32acc; } if (ctx->device->coopmat_support) { - return (ctx->device->fp16 && ctx->device->coopmat_acc_f16_support && prec == GGML_PREC_DEFAULT) ? ctx->device->pipeline_dequant_mul_mat_mat[src0_type].f16acc : ctx->device->pipeline_dequant_mul_mat_mat[src0_type].f32acc; + vk_matmul_pipeline2 & p = (src1_type == GGML_TYPE_F16) ? ctx->device->pipeline_dequant_mul_mat_mat_f16[src0_type] + : ctx->device->pipeline_dequant_mul_mat_mat[src0_type]; + return (ctx->device->fp16 && ctx->device->coopmat_acc_f16_support && prec == GGML_PREC_DEFAULT) ? p.f16acc : p.f32acc; } return (ctx->device->fp16 && prec == GGML_PREC_DEFAULT) ? ctx->device->pipeline_dequant_mul_mat_mat[src0_type].f16acc : ctx->device->pipeline_dequant_mul_mat_mat[src0_type].f32acc; } @@ -10071,7 +10094,23 @@ static void ggml_vk_mul_mat_q_f16(ggml_backend_vk_context * ctx, vk_context& sub // Reformat and convert to fp16 if non-contiguous, or for coopmat2 for better perf const bool x_non_contig = (ctx->device->coopmat2 && src0->type == GGML_TYPE_F32) || !ggml_vk_dim01_contiguous(src0); - const bool y_non_contig = (ctx->device->coopmat2 && src1->type == GGML_TYPE_F32) || + // Route quantized MUL_MAT through the f16-B kernels. Treating contiguous f32 B as + // y_non_contig reuses the convert-to-prealloc_y plumbing, like coopmat2 does. + const bool dense_f16b = ggml_vk_dense_f16b_enabled() && + ctx->device->coopmat_support && !ctx->device->coopmat2 && + ggml_is_quantized(src0->type) && src1->type == GGML_TYPE_F32 && + !(ctx->device->pipeline_dequant_mul_mat_mat_f16[src0->type].f16acc->is_empty() && + ctx->device->pipeline_dequant_mul_mat_mat_f16[src0->type].f32acc->is_empty()); + if (dense_f16b) { + static bool dense_f16b_logged = false; + if (!dense_f16b_logged) { + dense_f16b_logged = true; + fprintf(stderr, "ggml_vulkan: MUL_MAT f16-B path engaged (GGML_VK_DENSE_F16B)\n"); + } + } + + const bool y_non_contig = dense_f16b || + (ctx->device->coopmat2 && src1->type == GGML_TYPE_F32) || (src0->type == GGML_TYPE_BF16 && src1->type != GGML_TYPE_BF16) || !ggml_vk_dim01_contiguous(src1); From d737bd5f72b30c920ee398136971161c4b4be49b Mon Sep 17 00:00:00 2001 From: Nathan Wilson Date: Sun, 16 Aug 2026 14:34:15 +0000 Subject: [PATCH 66/68] vulkan : add auto mode to GGML_VK_DENSE_F16B =auto restricts the f16-B path to ne10 == 5120, the only reduction width measured to gain, so it cannot fire on the widths that lose. =1 keeps the old all-shapes behaviour as a manual override. gfx1151, Qwen3.8-27B UD-Q6_K_XL pp2048, auto vs off: +5.8 / +6.0 / +5.9 / +5.3% at ub 256/512/1024/2048, which is 97% of the all-shapes win at ub256 and 82-86% above it. Qwen3-Coder-30B MoE is untouched, the gate never fires. The width equality is a stopgap until a per-shape predicate is derived; it will silently do nothing for a dense model of another width. Assisted-by: Claude Opus 5 --- ggml/src/ggml-vulkan/ggml-vulkan.cpp | 24 +++++++++++++++++++----- 1 file changed, 19 insertions(+), 5 deletions(-) diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp index 8eee6436d7f..6c97e1b4d46 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp +++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp @@ -4498,14 +4498,25 @@ static bool ggml_vk_mmid_f16b_enabled() { const char * env = getenv("GGML_VK_MMID_F16B"); return env == nullptr || atoi(env) != 0; // on by default; =0 disables }(); -// GGML_VK_DENSE_F16B=1: same idea as GGML_VK_MMID_F16B but for plain MUL_MAT. Halves the B -// bytes, which keeps the activations inside the LLC at large ubatch and moves their row stride -// off the 1-of-16 channel pattern. Off by default, it is a loss when B already fits. -static bool ggml_vk_dense_f16b_enabled() { - static const bool enabled = getenv("GGML_VK_DENSE_F16B") != nullptr; return enabled; } +// GGML_VK_DENSE_F16B: same idea as GGML_VK_MMID_F16B but for plain MUL_MAT. Halves the B bytes +// moved. Numerically identical: mul_mm stages B into shared FLOAT_TYPE either way, so the f32-B +// kernel already rounds B to f16. Helps wide dense models, costs ~1% on narrow ones. +// 0 = off, 1 = all quantized dense matmuls, 2 = auto (only the K we have positive data for) +static int ggml_vk_dense_f16b_mode() { + static const int mode = [] { + const char * e = getenv("GGML_VK_DENSE_F16B"); + if (e == nullptr) return 0; + if (e[0] == 'a') return 2; + return atoi(e) != 0 ? 1 : 0; + }(); + return mode; +} + +static bool ggml_vk_dense_f16b_enabled() { return ggml_vk_dense_f16b_mode() != 0; } + static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { VK_LOG_DEBUG("ggml_vk_load_shaders(" << device->name << ")"); @@ -10096,7 +10107,10 @@ static void ggml_vk_mul_mat_q_f16(ggml_backend_vk_context * ctx, vk_context& sub !ggml_vk_dim01_contiguous(src0); // Route quantized MUL_MAT through the f16-B kernels. Treating contiguous f32 B as // y_non_contig reuses the convert-to-prealloc_y plumbing, like coopmat2 does. + // auto mode restricts to ne10 == 5120: the only width measured to gain. Narrow on purpose, + // so widths measured as losses (3584 dense, 2048 MoE) cannot trigger it. const bool dense_f16b = ggml_vk_dense_f16b_enabled() && + (ggml_vk_dense_f16b_mode() == 1 || ne10 == 5120) && ctx->device->coopmat_support && !ctx->device->coopmat2 && ggml_is_quantized(src0->type) && src1->type == GGML_TYPE_F32 && !(ctx->device->pipeline_dequant_mul_mat_mat_f16[src0->type].f16acc->is_empty() && From 9f5ece3a5ae1a3417be0d45cc18ac4797a6b7a68 Mon Sep 17 00:00:00 2001 From: Nathan Wilson Date: Tue, 8 Sep 2026 09:46:19 +0000 Subject: [PATCH 67/68] vulkan: derive the mul_mat_id tiles from the wave64 dense tiles The dense wave32 change rewrites s/m/l_warptile_mmq in place (WARP=32, WM/WN halved). The mul_mat_id tile block copied those rewritten vectors, so with GGML_VK_MMID_WAVE32=0 the mmid pipelines carried a WARP=32 spec constant but no required subgroup size and ran at wave64 with the wrong warp grid: test-backend-ops MUL_MAT_ID 700/7432 for every quantised type at n >= 16 (garbage output, not an error). Default settings were not affected; the opt-out was. Snapshot the dense tiles before the wave32 rewrite and build the mmid tiles from the snapshot. When the mmid wave32 pin is not applied, assert the coverage invariant and require the tile's own WARP as the subgroup size where the driver honours it, so the two can never disagree. Assisted-by: Claude (Opus 5) --- ggml/src/ggml-vulkan/ggml-vulkan.cpp | 25 ++++++++++++++++++++++--- 1 file changed, 22 insertions(+), 3 deletions(-) diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp index 6c97e1b4d46..d22e355aac1 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp +++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp @@ -5136,6 +5136,12 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { const char * dense_wave32_env = getenv("GGML_VK_DENSE_WAVE32"); const int dense_wave32 = dense_wave32_env ? atoi(dense_wave32_env) : 1; + // The mul_mat_id tiles below start from these. Snapshot them before the dense + // wave32 rewrite so an mmid pipeline created without a required subgroup size + // (GGML_VK_MMID_WAVE32=0) keeps a WARP that matches the real wave64 subgroup. + const auto l_warptile_mmq_w64 = l_warptile_mmq; + const auto m_warptile_mmq_w64 = m_warptile_mmq; + const auto s_warptile_mmq_w64 = s_warptile_mmq; if (dense_wave32_possible && dense_wave32 != 0) { auto wave32_tile = [](std::vector & w) -> bool { // {BLOCK_SIZE, BM, BN, BK, WM, WN, WMITER, TM, TN, TK, WARP} @@ -5278,11 +5284,11 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { // so even the 32-wide small tile runs half empty. Only meaningful stacked on // GGML_VK_MMID_SMALLN=1 (routes mmid to the small tile) + the row-list prepass. // Shadows the s-tile config for the mmid quant pipelines only; dense unaffected. - auto s_warptile_mmq_id16 = s_warptile_mmq; + auto s_warptile_mmq_id16 = s_warptile_mmq_w64; auto s_mmq_wg_denoms_id16 = s_mmq_wg_denoms; - auto m_warptile_mmq_id128 = m_warptile_mmq; + auto m_warptile_mmq_id128 = m_warptile_mmq_w64; auto m_mmq_wg_denoms_id128 = m_mmq_wg_denoms; - auto l_warptile_mmq_idw = l_warptile_mmq; + auto l_warptile_mmq_idw = l_warptile_mmq_w64; uint32_t mmid_req_sgs = 0; { const char * tile16_env = getenv("GGML_VK_MMID_TILE16"); @@ -5344,6 +5350,19 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { wave32_tile(s_warptile_mmq_id16); wave32_tile(m_warptile_mmq_id128); wave32_tile(l_warptile_mmq_idw); + } else { + // Wave64 stack: the tiles are the dense wave64 tiles plus the BM64/M128 reshapes, + // which keep NUM_WARPS == (BM/WM)*(BN/WN). Pin the subgroup to the tile's WARP + // where the driver honours it, so the spec constant can never disagree with + // the real subgroup size. + for (const auto * w : { &s_warptile_mmq_id16, &m_warptile_mmq_id128, &l_warptile_mmq_idw }) { + GGML_ASSERT((*w)[0] / (*w)[10] == ((*w)[1] / (*w)[4]) * ((*w)[2] / (*w)[5])); + } + const uint32_t warp = s_warptile_mmq_id16[10]; + if (device->subgroup_size_control && warp == m_warptile_mmq_id128[10] && warp == l_warptile_mmq_idw[10] && + device->subgroup_min_size <= warp && warp <= device->subgroup_max_size) { + mmid_req_sgs = warp; + } } } { From 1debd524b05de9502a318386e737c76366986793 Mon Sep 17 00:00:00 2001 From: Nathan Wilson Date: Tue, 8 Sep 2026 10:56:08 +0000 Subject: [PATCH 68/68] vulkan: keep the coopmat1 FA wave32 pin to multi-row dispatches Decode dispatches carry N = gqa_ratio query rows and the narrow subgroup loses there: Qwen3-Coder-30B q8_0 KV tg64 at d8192/d32768 measured 3.5 to 4 percent slower with the pin than without, while the prefill gain it was introduced for (up to 10 percent at d8192) needs the multi-row shapes. Apply the rule only when n_rows >= 32. Assisted-by: Claude (Opus 5) --- ggml/src/ggml-vulkan/ggml-vulkan.cpp | 5 +++++ 1 file changed, 5 insertions(+) diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp index d22e355aac1..353edb87c46 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp +++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp @@ -4056,7 +4056,12 @@ static vk_fa_tuning_params get_fa_tuning_params_coopmat1(const vk_device& device const char * e = getenv("GGML_VK_FA_WAVE32"); return e ? atoi(e) : 1; }(); + // Multi-row dispatches only. Decode dispatches carry N = gqa_ratio rows (8 on + // Qwen3-Coder-30B) and there the narrow subgroup loses: tg64 at d8192/d32768 measured + // 3.5 to 4 percent slower with the pin (q8_0 KV), while prefill gains up to 10 percent + // at d8192. Keep the pin to the prefill shapes it was measured on. if (fa_wave32 != 0 && + n_rows >= 32 && device->subgroup_size_control && 32 < device->subgroup_size && // narrow only, never widen device->subgroup_min_size <= 32 && 32 <= device->subgroup_max_size &&