From 04c8909061988a261a01f3456873db228ad7463a Mon Sep 17 00:00:00 2001 From: Nathan Wilson Date: Fri, 4 Sep 2026 00:17:26 +0000 Subject: [PATCH 1/6] vulkan: replace the EngramHalo radix top-k with upstream's radix select This reverts the Vulkan part of "Halo/vulkan topk radix (#17)" (ca062e9f3), which ported the EngramHalo HIP radix kernel as topk_radix.comp, so that the next commit can carry upstream's top_k radix select (ggml-org/llama.cpp#28032) unchanged. One implementation of the op, the one upstream settled on, keeps this tree closest to upstream and lets the deterministic slot assignment that follows apply as-is. The two shaders solve the same problem; the difference that matters for Qwen 3.8 Flash-Next is order: topk_radix.comp assigns output slots with atomicAdd, so the selected indices come out in a scheduling-dependent order, and the sparse attention sums them in that order (see the deterministic-scan commit for the measured effect). Upstream's variant also carries the QSA indexer fusion. Co-authored-by: Claude (Opus 5) --- ggml/src/ggml-vulkan/ggml-vulkan.cpp | 71 ++--------- .../vulkan-shaders/topk_radix.comp | 120 ------------------ .../vulkan-shaders/vulkan-shaders-gen.cpp | 2 - tests/test-backend-ops.cpp | 23 +--- 4 files changed, 13 insertions(+), 203 deletions(-) delete mode 100644 ggml/src/ggml-vulkan/vulkan-shaders/topk_radix.comp diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp index de8e2673067b..a795d3595a1e 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp +++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp @@ -1057,7 +1057,6 @@ struct vk_device_struct { vk_pipeline pipeline_argsort_f32[num_argsort_pipelines]; vk_pipeline pipeline_argsort_large_f32[num_argsort_pipelines]; vk_pipeline pipeline_topk_f32[num_topk_pipelines]; - vk_pipeline pipeline_topk_radix_f32; vk_pipeline pipeline_sum_rows_f32; vk_pipeline pipeline_cross_entropy_loss_f32, pipeline_cross_entropy_loss_f32_wg512; vk_pipeline pipeline_cross_entropy_loss_back_f32, pipeline_cross_entropy_loss_back_f32_wg512; @@ -5857,13 +5856,6 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { } } - // radix selection for k beyond the shared-memory shaders' single-workgroup reach: - // one workgroup per row, single dispatch, no scratch buffers - { - const uint32_t RADIX_WG = std::min(1024, 1u << device->max_workgroup_size_log2); - ggml_vk_create_pipeline2(device, device->pipeline_topk_radix_f32, "topk_radix_f32", topk_radix_f32_len, topk_radix_f32_data, "main", 2, sizeof(vk_op_topk_push_constants), {RADIX_WG, 1, 1}, {RADIX_WG}, 1, true); - } - ggml_vk_create_pipeline(device, device->pipeline_argmax_f32, "argmax_f32", argmax_f32_len, argmax_f32_data, "main", 2, sizeof(vk_op_push_constants), {1, 1, 1}, { device->subgroup_size }, 1); ggml_vk_create_pipeline(device, device->pipeline_sum_rows_f32, "sum_rows_f32", sum_rows_f32_len, sum_rows_f32_data, "main", 2, sizeof(vk_op_sum_rows_push_constants), {1, 1, 1}, { device->subgroup_size }, 1); @@ -14027,7 +14019,7 @@ static void ggml_vk_argsort(ggml_backend_vk_context * ctx, vk_context& subctx, c std::array elements; elements[0] = ncolsp2; - elements[1] = std::min(nrows, ctx->device->properties.limits.maxComputeWorkGroupCount[1]); + elements[1] = std::min((uint32_t)ggml_nrows(src0), ctx->device->properties.limits.maxComputeWorkGroupCount[1]); elements[2] = 1; // First dispatch initializes tmp_idx and does the first N passes where @@ -14071,10 +14063,7 @@ static void ggml_vk_argsort(ggml_backend_vk_context * ctx, vk_context& subctx, c } } -// Tournament-reduction top_k, only correct while k fits a single workgroup's -// tournament pipeline (see ggml_vk_topk's dispatcher and the large-k path below -// for why this can't just be given a bigger workgroup). -static void ggml_vk_topk_small(ggml_backend_vk_context * ctx, vk_context& subctx, const ggml_tensor * src0, ggml_tensor * dst) { +static void ggml_vk_topk(ggml_backend_vk_context * ctx, vk_context& subctx, const ggml_tensor * src0, ggml_tensor * dst) { uint32_t ncols = src0->ne[0]; uint32_t nrows = ggml_nrows(src0); uint32_t k = dst->ne[0]; @@ -14182,48 +14171,6 @@ static void ggml_vk_topk_small(ggml_backend_vk_context * ctx, vk_context& subctx ctx->prealloc_x_need_sync = true; } -// Large-k top_k: the tournament reduction above can only discard elements a -// workgroup can see all of at once, so it can never make progress once k is -// >= the max deployable workgroup size (~1024 on real hardware) - any element -// in a smaller chunk could belong to the true global top-k. Radix selection -// discards on a *global* bucket boundary instead, so it converges regardless -// of the workgroup-size-vs-k relationship: one workgroup per row, a single -// dispatch, no scratch buffers. -static void ggml_vk_topk_large(ggml_backend_vk_context * ctx, vk_context& subctx, const ggml_tensor * src0, ggml_tensor * dst) { - const uint32_t ncols = (uint32_t) src0->ne[0]; - const uint32_t nrows = (uint32_t) ggml_nrows(src0); - const uint32_t k = (uint32_t) dst->ne[0]; - - vk_pipeline pipeline = ctx->device->pipeline_topk_radix_f32; - GGML_ASSERT(pipeline != nullptr); - - vk_op_topk_push_constants pc { ncols, ncols, ncols, k, nrows, 0, 0 }; - - vk_subbuffer src0_buf = ggml_vk_tensor_subbuffer(ctx, src0); - vk_subbuffer dst_buf = ggml_vk_tensor_subbuffer(ctx, dst); - - // one workgroup per row - std::array elements = { nrows*pipeline->wg_denoms[0], 1, 1 }; - - ggml_pipeline_request_descriptor_sets(ctx, pipeline, 1); - ggml_vk_dispatch_pipeline(ctx, subctx, pipeline, { src0_buf, dst_buf }, pc, elements); -} - -// true when k fits a single-workgroup tournament pipeline; the index is only -// valid to probe after the range check, hence the explicit ordering here -static bool ggml_vk_topk_k_fits_pipeline(const vk_device_struct * device, uint32_t k) { - const uint32_t min_pipeline = (uint32_t) log2f(float(k)) + 1; - return min_pipeline < num_topk_pipelines && device->pipeline_topk_f32[min_pipeline] != nullptr; -} - -static void ggml_vk_topk(ggml_backend_vk_context * ctx, vk_context& subctx, const ggml_tensor * src0, ggml_tensor * dst) { - if (ggml_vk_topk_k_fits_pipeline(ctx->device.get(), (uint32_t) dst->ne[0])) { - ggml_vk_topk_small(ctx, subctx, src0, dst); - } else { - ggml_vk_topk_large(ctx, subctx, src0, dst); - } -} - static void ggml_vk_sum(ggml_backend_vk_context * ctx, vk_context& subctx, const ggml_tensor * src0, ggml_tensor * dst) { vk_op_sum_rows_push_constants p = vk_op_sum_rows_push_constants_init(src0, dst, ggml_nelements(src0)); ggml_vk_op_f32(ctx, subctx, src0, nullptr, nullptr, nullptr, dst, GGML_OP_SUM, p); @@ -18904,13 +18851,15 @@ static bool ggml_backend_vk_device_supports_op(ggml_backend_dev_t dev, const ggm if (!ggml_is_contiguous(op) || !ggml_is_contiguous(op->src[0])) { return false; } - // op is the dst tensor here, so op->ne[0] == k (ggml_top_k builds dst - // as ggml_new_tensor_4d(..., k, a->ne[1], ...)), not the candidate pool - // width - ggml_vk_topk_small's multi-pass loop already handles arbitrarily - // large ncols correctly as long as k fits a single tournament pipeline. - // Beyond that, the radix path takes over, which has no extra requirements. - return true; + // We could potentially support larger, using argsort to sort the + // whole thing. Not clear if this is needed. + uint32_t min_pipeline = (uint32_t)log2f(float(op->ne[0])) + 1; + if (min_pipeline >= num_topk_pipelines || + !device->pipeline_topk_f32[min_pipeline]) { + return false; + } } + return true; case GGML_OP_UPSCALE: if (op->op_params[0] & GGML_SCALE_FLAG_ANTIALIAS) { if ((op->op_params[0] & 0xFF) != GGML_SCALE_MODE_BILINEAR) { diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/topk_radix.comp b/ggml/src/ggml-vulkan/vulkan-shaders/topk_radix.comp deleted file mode 100644 index abb88fe4e1e5..000000000000 --- a/ggml/src/ggml-vulkan/vulkan-shaders/topk_radix.comp +++ /dev/null @@ -1,120 +0,0 @@ -#version 450 - -#extension GL_EXT_control_flow_attributes : enable - -#include "types.glsl" - -// Exact top-k selection for rows too wide (or k too large) for the shared-memory -// nary-search shader: iterative 8-bit radix selection on the order-preserving -// unsigned transform of the keys, one workgroup per row, no temporary buffers. -// The resulting indices are in no particular order (the ggml_top_k contract). -// Ported to Vulkan from the HIP kernel in EngramHalo.cpp (Aristo94). - -layout(constant_id = 0) const int BLOCK_SIZE = 1024; - -layout(local_size_x_id = 0, local_size_y = 1, local_size_z = 1) in; - -layout (binding = 0) readonly buffer A {float data_a[];}; -layout (binding = 1) writeonly buffer D {int data_d[];}; - -layout (push_constant) uniform parameter { - uint orig_ncols; - uint ncols_input; - uint ncols_output; - uint k; - uint nrows; - uint first_pass; - uint last_pass; -} p; - -shared uint hist[256]; -shared uint sh_prefix; -shared uint sh_rank; -shared uint sh_greater; -shared uint sh_equal; - -// map float bits to uint such that unsigned comparison matches float comparison -uint f2ui(float x) { - uint y = floatBitsToUint(x); - if ((y & 0x80000000u) != 0u) { - y ^= ~0u; - } else { - y |= 0x80000000u; - } - return y; -} - -void main() { - const uint row = gl_WorkGroupID.x; - const uint tid = gl_LocalInvocationID.x; - const uint wgs = gl_WorkGroupSize.x; - const uint base = row*p.ncols_input; - - if (row >= p.nrows) { - return; - } - - // find the exact key of the k-th largest element, 8 bits per pass - uint prefix = 0u; - uint pmask = 0u; - uint rank = p.k; - - for (int shift = 24; shift >= 0; shift -= 8) { - // strided: the workgroup can be narrower than the 256 bins (Vulkan only - // guarantees maxComputeWorkGroupInvocations >= 128) - for (uint bin = tid; bin < 256u; bin += wgs) { - hist[bin] = 0u; - } - barrier(); - - // count only the elements that still match the prefix found so far - for (uint col = tid; col < p.ncols_input; col += wgs) { - const uint key = f2ui(data_a[base + col]); - if ((key & pmask) == prefix) { - atomicAdd(hist[(key >> shift) & 255u], 1u); - } - } - barrier(); - - // walk the bins from the largest keys down to the one holding the k-th largest - if (tid == 0u) { - uint r = rank; - int bin = 255; - while (bin > 0 && hist[bin] < r) { - r -= hist[bin]; - bin--; - } - sh_prefix = prefix | (uint(bin) << shift); - sh_rank = r; - } - barrier(); - - prefix = sh_prefix; - rank = sh_rank; - pmask |= 255u << shift; - } - - // gather: all elements strictly greater than the threshold key are in the top k, - // ties on the exact threshold fill the remaining `rank` slots - if (tid == 0u) { - sh_greater = 0u; - sh_equal = 0u; - } - barrier(); - - const uint threshold = prefix; - const uint out_base = row*p.k; - - for (uint col = tid; col < p.ncols_input; col += wgs) { - const uint key = f2ui(data_a[base + col]); - if (key > threshold) { - const uint pos = atomicAdd(sh_greater, 1u); - data_d[out_base + pos] = int(col); - } else if (key == threshold) { - const uint pos = atomicAdd(sh_equal, 1u); - if (pos < rank) { - data_d[out_base + (p.k - rank) + pos] = int(col); - } - } - } -} 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 3624fb51a43f..f0610f4cd82d 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp @@ -1029,8 +1029,6 @@ void process_shaders() { string_to_spv("topk_argsort_f32", "topk_argsort.comp", {{"A_TYPE", "float"}}); string_to_spv("topk_nary_search_f32", "topk_nary_search.comp", {{"A_TYPE", "float"}}); - string_to_spv("topk_radix_f32", "topk_radix.comp", {{"A_TYPE", "float"}}); - string_to_spv("argmax_f32", "argmax.comp", merge_maps(base_dict, {{"A_TYPE", "float"}, {"D_TYPE", "int"}})); string_to_spv("sum_rows_f32", "sum_rows.comp", merge_maps(base_dict, {{"A_TYPE", "float"}, {"D_TYPE", "float"}})); string_to_spv("cross_entropy_loss_f32", "cross_entropy_loss.comp", merge_maps(base_dict, {{"A_TYPE", "float"}, {"B_TYPE", "float"}, {"D_TYPE", "float"}})); diff --git a/tests/test-backend-ops.cpp b/tests/test-backend-ops.cpp index 516c8094739e..a9d91acf6ee7 100644 --- a/tests/test-backend-ops.cpp +++ b/tests/test-backend-ops.cpp @@ -6172,17 +6172,16 @@ struct test_top_k : public test_case { const std::array ne; const int k; const bool ties; - const bool neg_inf; ggml_tensor * input {}; std::string vars() override { - return VARS_TO_STR5(type, ne, k, ties, neg_inf); + return VARS_TO_STR4(type, ne, k, ties); } test_top_k(ggml_type type = GGML_TYPE_F32, std::array ne = {16, 10, 10, 10}, - int k = 4, bool ties = false, bool neg_inf = false) - : type(type), ne(ne), k(k), ties(ties), neg_inf(neg_inf) {} + int k = 4, bool ties = false) + : type(type), ne(ne), k(k), ties(ties) {} double max_err() override { return 0.0; @@ -6282,14 +6281,6 @@ struct test_top_k : public test_case { } } std::shuffle(data.begin(), data.end(), rng); - if (neg_inf) { - // keep exactly k finite candidates; the rest are masked (-INFINITY), - // mirroring qwen4exp's indexer score tensor where invisible cells are - // -INFINITY and must never be selected regardless of magnitude - for (int64_t i = k; i < t->ne[0]; i++) { - data[i] = -INFINITY; - } - } ggml_backend_tensor_set(t, data.data(), r * t->nb[1], t->ne[0] * sizeof(float)); } } @@ -9800,12 +9791,6 @@ static std::vector> make_test_cases_eval() { test_cases.emplace_back(new test_top_k(GGML_TYPE_F32, {(1<> make_test_cases_eval() { test_cases.emplace_back(new test_top_k(GGML_TYPE_F32, {202048, nrows, 1, 1}, k, true)); } } - // sparse-attention indexer shape: ncols tracks a growing KV cache, k is far past one workgroup - test_cases.emplace_back(new test_top_k(GGML_TYPE_F32, {32768, 1, 1, 1}, 2051, false, true)); for (int k : {1, 2, 3, 7, 15}) { test_cases.emplace_back(new test_top_k(GGML_TYPE_F32, {16, 10, 10, 10}, k)); From 55a6adcc27e38cbe89e728ff1311b23d265f3148 Mon Sep 17 00:00:00 2001 From: Ruben Ortlam Date: Mon, 31 Aug 2026 07:04:34 +0200 Subject: [PATCH 2/6] vulkan: top_k radix select for k >= 1024 for Qwen 3.8 Flash Next (#28032) * vulkan: add top-k radix sort shader for k >= 1024 * add Qwen 3.8 Flash Next top-k tests * add top-k qsa fusion * clean up code (cherry picked from commit daef7b6874397a5a7c3d7e38b55e2ee0adf7da38) --- ggml/src/ggml-vulkan/ggml-vulkan.cpp | 239 +++++++++++++++++- .../vulkan-shaders/topk_radix_select.comp | 144 +++++++++++ .../vulkan-shaders/vulkan-shaders-gen.cpp | 1 + tests/test-backend-ops.cpp | 97 +++++++ 4 files changed, 471 insertions(+), 10 deletions(-) create mode 100644 ggml/src/ggml-vulkan/vulkan-shaders/topk_radix_select.comp diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp index a795d3595a1e..05be6e93712a 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp +++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp @@ -657,6 +657,21 @@ static constexpr std::initializer_list snake_pattern { GGM GGML_OP_SQR, GGML_OP_MUL, GGML_OP_ADD }; +// qwen4 QSA indexer: gather per-block scores to cells + add f16 mask (cast+reshape) + top-k, +// fused into one radix-select. The cast/reshape are elided; the raw f16 mask is read in-shader. +static constexpr std::initializer_list topk_qsa_pattern { GGML_OP_GET_ROWS, GGML_OP_PERMUTE, + GGML_OP_CONT, GGML_OP_CPY, + GGML_OP_RESHAPE, GGML_OP_ADD, + GGML_OP_TOP_K }; +static constexpr std::initializer_list> topk_qsa_edges { + { 1, 0, 0 }, // permute->src[0] == get_rows + { 2, 0, 1 }, // cont->src[0] == permute + { 4, 0, 3 }, // reshape->src[0] == cpy (mask cast) + { 5, 0, 2 }, // add->src[0] == cont + { 5, 1, 4 }, // add->src[1] == reshape + { 6, 0, 5 }, // top_k->src[0] == add +}; + //node #978 ( SOFT_MAX): ffn_moe_probs-15 ( 0K) [Vulka ] use=2: ffn_moe_logits-15 ( 0K) [Vulka ] //node #979 ( RESHAPE): ffn_moe_probs-15 (re ( 0K) [Vulka ] use=1: ffn_moe_probs-15 ( 0K) [Vulka ] //node #980 ( ARGSORT): ffn_moe_argsort-15 ( 0K) [Vulka ] use=1: ffn_moe_probs-15 ( 0K) [Vulka ] @@ -1057,6 +1072,8 @@ struct vk_device_struct { vk_pipeline pipeline_argsort_f32[num_argsort_pipelines]; vk_pipeline pipeline_argsort_large_f32[num_argsort_pipelines]; vk_pipeline pipeline_topk_f32[num_topk_pipelines]; + vk_pipeline pipeline_topk_radix_f32; + vk_pipeline pipeline_topk_radix_qsa; // qwen4 QSA indexer fusion (f16 mask) vk_pipeline pipeline_sum_rows_f32; vk_pipeline pipeline_cross_entropy_loss_f32, pipeline_cross_entropy_loss_f32_wg512; vk_pipeline pipeline_cross_entropy_loss_back_f32, pipeline_cross_entropy_loss_back_f32_wg512; @@ -1749,6 +1766,15 @@ struct vk_op_topk_push_constants { uint32_t last_pass; }; +struct vk_op_topk_radix_push_constants { + uint32_t ncols; + uint32_t k; + uint32_t nrows; + uint32_t n_tps; // QSA only + uint32_t n_blocks; // QSA only + uint32_t n_stream; // QSA only +}; + struct vk_op_im2col_push_constants { uint64_t dst_addr; uint32_t batch_offset; uint32_t offset_delta; @@ -2439,6 +2465,8 @@ struct ggml_backend_vk_context { int fused_ops_write_mask {}; topk_moe_mode fused_topk_moe_mode {}; bool fused_topk_moe_scale {}; + // QSA indexer gather+add+top_k fused into one radix-select + bool fused_topk_qsa {}; // for GGML_VK_PERF_LOGGER std::unique_ptr perf_logger; @@ -5856,6 +5884,14 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { } } + // large-k fallback: one workgroup per row, radix-select instead of a full sort. The QSA + // variant (spec constant 1) additionally gathers the qwen4 indexer input on the fly. + { + const uint32_t BLOCK_SIZE = 1u << std::min(10u, device->max_workgroup_size_log2); + ggml_vk_create_pipeline2(device, device->pipeline_topk_radix_f32, "topk_radix_f32", topk_radix_select_f32_len, topk_radix_select_f32_data, "main", 5, sizeof(vk_op_topk_radix_push_constants), {BLOCK_SIZE, 1, 1}, {BLOCK_SIZE, 0}, 1, true); + ggml_vk_create_pipeline2(device, device->pipeline_topk_radix_qsa, "topk_radix_qsa", topk_radix_select_f32_len, topk_radix_select_f32_data, "main", 5, sizeof(vk_op_topk_radix_push_constants), {BLOCK_SIZE, 1, 1}, {BLOCK_SIZE, 1}, 1, true); + } + ggml_vk_create_pipeline(device, device->pipeline_argmax_f32, "argmax_f32", argmax_f32_len, argmax_f32_data, "main", 2, sizeof(vk_op_push_constants), {1, 1, 1}, { device->subgroup_size }, 1); ggml_vk_create_pipeline(device, device->pipeline_sum_rows_f32, "sum_rows_f32", sum_rows_f32_len, sum_rows_f32_data, "main", 2, sizeof(vk_op_sum_rows_push_constants), {1, 1, 1}, { device->subgroup_size }, 1); @@ -14068,6 +14104,31 @@ static void ggml_vk_topk(ggml_backend_vk_context * ctx, vk_context& subctx, cons uint32_t nrows = ggml_nrows(src0); uint32_t k = dst->ne[0]; + // tournament path is faster where it fits; use radix-select only past its k limit + const uint32_t k_min_pipeline = std::max((uint32_t) log2f(float(k)) + 1, ctx->device->subgroup_size_log2); + if (k_min_pipeline >= num_topk_pipelines || ctx->device->pipeline_topk_f32[k_min_pipeline] == nullptr) { + vk_pipeline pipeline = ctx->device->pipeline_topk_radix_f32; + GGML_ASSERT(pipeline != nullptr); + + if (ctx->prealloc_x_need_sync) { + ggml_vk_sync_buffers(ctx, subctx); + } + + vk_op_topk_radix_push_constants pc { ncols, k, nrows, 0, 0, 0 }; + std::array elements { + pipeline->wg_denoms[0], + std::min(nrows, ctx->device->properties.limits.maxComputeWorkGroupCount[1]), + 1, + }; + // the non-QSA path only uses bindings 0/1; bind valid buffers for the unused QSA slots + vk_subbuffer src0_buf = ggml_vk_tensor_subbuffer(ctx, src0); + vk_subbuffer dst_buf = ggml_vk_tensor_subbuffer(ctx, dst); + ggml_pipeline_request_descriptor_sets(ctx, pipeline, 1); + ggml_vk_dispatch_pipeline(ctx, subctx, pipeline, + { src0_buf, dst_buf, src0_buf, src0_buf, src0_buf }, pc, elements); + return; + } + vk_op_topk_push_constants pc { ncols, ncols, ncols, k, nrows, 0, 0 }; if (ctx->prealloc_x_need_sync) { @@ -14171,6 +14232,55 @@ static void ggml_vk_topk(ggml_backend_vk_context * ctx, vk_context& subctx, cons ctx->prealloc_x_need_sync = true; } +static void ggml_vk_topk_qsa(ggml_backend_vk_context * ctx, vk_context& subctx, const ggml_cgraph * cgraph, int node_idx) { + const ggml_tensor * get_rows = cgraph->nodes[node_idx + 0]; + const ggml_tensor * add = cgraph->nodes[node_idx + ctx->num_additional_fused_ops - 1]; + ggml_tensor * top_k = cgraph->nodes[node_idx + ctx->num_additional_fused_ops]; + + const ggml_tensor * scores = get_rows->src[0]; // [n_tps, n_blocks, n_stream] + const ggml_tensor * cell_blk = get_rows->src[1]; // [n_kv, n_stream] + + // raw f16 mask: follow the reshape/cpy chain back to the materialized input + const ggml_tensor * mask = add->src[1]; + while (mask->op == GGML_OP_RESHAPE || mask->op == GGML_OP_CPY) { + mask = mask->src[0]; + } + + const uint32_t n_tps = scores->ne[0]; + const uint32_t n_blocks = scores->ne[1]; + const uint32_t n_stream = scores->ne[2]; + const uint32_t n_kv = cell_blk->ne[0]; + const uint32_t width = top_k->ne[0]; + const uint32_t nrows = n_tps * n_stream; + + vk_pipeline pipeline = ctx->device->pipeline_topk_radix_qsa; + GGML_ASSERT(pipeline != nullptr); + + // scratch holds the gathered+masked input, materialized once and reused across passes + const size_t scratch_size = size_t{ n_kv } * nrows * sizeof(float); + if (ctx->prealloc_size_x < scratch_size) { + ctx->prealloc_size_x = scratch_size; + ggml_vk_preallocate_buffers(ctx, subctx); + } + if (ctx->prealloc_x_need_sync) { + ggml_vk_sync_buffers(ctx, subctx); + } + + vk_op_topk_radix_push_constants pc { n_kv, width, nrows, n_tps, n_blocks, n_stream }; + std::array elements { + pipeline->wg_denoms[0], + std::min(nrows, ctx->device->properties.limits.maxComputeWorkGroupCount[1]), + 1, + }; + vk_subbuffer scratch_buf { ctx->prealloc_x, 0, ctx->prealloc_x->size }; + ggml_pipeline_request_descriptor_sets(ctx, pipeline, 1); + ggml_vk_dispatch_pipeline(ctx, subctx, pipeline, + { ggml_vk_tensor_subbuffer(ctx, scores), ggml_vk_tensor_subbuffer(ctx, top_k), + ggml_vk_tensor_subbuffer(ctx, cell_blk), ggml_vk_tensor_subbuffer(ctx, mask), + scratch_buf }, pc, elements); + ctx->prealloc_x_need_sync = true; +} + static void ggml_vk_sum(ggml_backend_vk_context * ctx, vk_context& subctx, const ggml_tensor * src0, ggml_tensor * dst) { vk_op_sum_rows_push_constants p = vk_op_sum_rows_push_constants_init(src0, dst, ggml_nelements(src0)); ggml_vk_op_f32(ctx, subctx, src0, nullptr, nullptr, nullptr, dst, GGML_OP_SUM, p); @@ -15832,7 +15942,11 @@ static bool ggml_vk_build_graph(ggml_backend_vk_context * ctx, ggml_cgraph * cgr break; case GGML_OP_GET_ROWS: - ggml_vk_get_rows(ctx, compute_ctx, src0, src1, node); + if (ctx->fused_topk_qsa) { + ggml_vk_topk_qsa(ctx, compute_ctx, cgraph, node_idx); + } else { + ggml_vk_get_rows(ctx, compute_ctx, src0, src1, node); + } break; case GGML_OP_GET_ROWS_BACK: @@ -17244,6 +17358,92 @@ static bool ggml_vk_can_fuse_topk_moe(ggml_backend_vk_context * ctx, const struc return true; } +// Manual op-sequence match (ggml_can_fuse_subgraph rejects the mask's external reshape/cpy). +static bool ggml_vk_match_ops(const struct ggml_cgraph * cgraph, int node_idx, + const std::initializer_list & ops) { + if (node_idx + (int) ops.size() > cgraph->n_nodes) { + return false; + } + for (size_t j = 0; j < ops.size(); ++j) { + const ggml_tensor * node = cgraph->nodes[node_idx + j]; + if (node->op != ops.begin()[j] || + (node->flags & GGML_TENSOR_FLAG_COMPUTE) == 0 || + (node->flags & GGML_TENSOR_FLAG_OUTPUT) != 0) { + return false; + } + } + return true; +} + +// True if the qwen4 QSA indexer top-k can be fused at node_idx (the get_rows). +static bool ggml_vk_can_fuse_topk_qsa(ggml_backend_vk_context * ctx, const struct ggml_cgraph * cgraph, int node_idx) { + if (ctx->device->disable_fusion || !ctx->device->pipeline_topk_radix_qsa) { + return false; + } + + const int n_ops = topk_qsa_pattern.size(); + if (!ggml_vk_match_ops(cgraph, node_idx, topk_qsa_pattern) || + !ggml_check_edges(cgraph, node_idx, topk_qsa_edges)) { + return false; + } + + // elided nodes must be single-use (cpy counts its own src[1] self-reference) + for (int j = 0; j < n_ops - 1; ++j) { + const ggml_tensor * node = cgraph->nodes[node_idx + j]; + const int32_t want = node->op == GGML_OP_CPY ? 2 : 1; + if (ggml_node_get_use_count(cgraph, node_idx + j) != want) { + return false; + } + } + + const ggml_tensor * get_rows = cgraph->nodes[node_idx + 0]; + const ggml_tensor * add = cgraph->nodes[node_idx + n_ops - 2]; + const ggml_tensor * top_k = cgraph->nodes[node_idx + n_ops - 1]; + + const ggml_tensor * scores = get_rows->src[0]; // [n_tps, n_blocks, n_stream] + const ggml_tensor * cell_blk = get_rows->src[1]; // [n_kv, n_stream] + const ggml_tensor * expanded = add->src[0]; // [n_kv, n_tps, n_stream] + + // raw mask: follow the reshape/cpy chain back to the materialized f16 input + const ggml_tensor * mask = add->src[1]; + while (mask && (mask->op == GGML_OP_RESHAPE || mask->op == GGML_OP_CPY)) { + mask = mask->src[0]; + } + if (!mask || mask->type != GGML_TYPE_F16) { + return false; + } + + if (scores->type != GGML_TYPE_F32 || cell_blk->type != GGML_TYPE_I32 || top_k->type != GGML_TYPE_I32) { + return false; + } + if (!ggml_is_contiguous(scores) || !ggml_is_contiguous(cell_blk) || !ggml_is_contiguous(mask) || + !ggml_is_contiguous(expanded) || !ggml_is_contiguous(top_k)) { + return false; + } + + const int64_t n_tps = scores->ne[0]; + const int64_t n_blocks = scores->ne[1]; + const int64_t n_stream = scores->ne[2]; + const int64_t n_kv = cell_blk->ne[0]; + const int64_t width = top_k->ne[0]; + + // pin the indexer layout the shader's addressing assumes + if (scores->ne[3] != 1 || cell_blk->ne[1] != n_stream || ggml_nrows(cell_blk) != n_stream || + ggml_nelements(mask) != n_kv * n_tps * n_stream || + expanded->ne[0] != n_kv || expanded->ne[1] != n_tps || expanded->ne[2] != n_stream || + top_k->ne[1] != n_tps || top_k->ne[2] != n_stream || top_k->ne[3] != 1 || + n_blocks <= 0 || n_kv <= 0 || width <= 0 || width > n_kv) { + return false; + } + + // only worth it in the radix regime; small k uses the faster tournament unfused + const uint32_t k_min_pipeline = std::max((uint32_t) log2f(float(width)) + 1, ctx->device->subgroup_size_log2); + if (k_min_pipeline < num_topk_pipelines && ctx->device->pipeline_topk_f32[k_min_pipeline]) { + return false; + } + return true; +} + static bool ggml_vk_can_fuse_rope_set_rows(ggml_backend_vk_context * ctx, const struct ggml_cgraph * cgraph, int node_idx) { GGML_UNUSED(ctx); @@ -17623,6 +17823,7 @@ static ggml_status ggml_backend_vk_graph_compute(ggml_backend_t backend, ggml_cg ctx->fused_topk_moe_mode = TOPK_MOE_COUNT; ctx->fused_topk_moe_scale = false; + ctx->fused_topk_qsa = false; const char *fusion_string {}; if (!ctx->device->disable_fusion) { uint32_t num_adds = ggml_vk_fuse_multi_add(ctx, cgraph, i); @@ -17712,6 +17913,11 @@ static ggml_status ggml_backend_vk_graph_compute(ggml_backend_t backend, ggml_cg // with a data dependency on that register. The overlap check still // rejects partial overlaps (different base or size). std::fill_n(op_srcs_fused_elementwise, 5, true); + } else if (ggml_vk_can_fuse_topk_qsa(ctx, cgraph, i)) { + ctx->num_additional_fused_ops = topk_qsa_pattern.size() - 1; + ctx->fused_topk_qsa = true; + fusion_string = "TOPK_QSA"; + std::fill_n(op_srcs_fused_elementwise, ctx->num_additional_fused_ops + 1, false); } else if (ggml_can_fuse_subgraph(cgraph, i, topk_moe_early_softmax_norm, { i + 3, i + 9 }) && ggml_check_edges(cgraph, i, topk_moe_early_softmax_norm_edges) && ggml_vk_can_fuse_topk_moe(ctx, cgraph, i, TOPK_MOE_EARLY_SOFTMAX_NORM)) { @@ -17828,6 +18034,7 @@ static ggml_status ggml_backend_vk_graph_compute(ggml_backend_t backend, ggml_cg ctx->fused_ops_write_mask = 1; ctx->fused_topk_moe_mode = TOPK_MOE_COUNT; ctx->fused_topk_moe_scale = false; + ctx->fused_topk_qsa = false; } } @@ -18024,6 +18231,9 @@ static void ggml_vk_graph_optimize(ggml_backend_t backend, struct ggml_cgraph * if (keep_pattern(snake_pattern)) { continue; } + if (keep_pattern(topk_qsa_pattern)) { + continue; + } // First, grab the next unused node. current_set.push_back(first_unused); @@ -18042,13 +18252,23 @@ static void ggml_vk_graph_optimize(ggml_backend_t backend, struct ggml_cgraph * if (is_empty(graph->nodes[j])) { continue; } - // Don't pull forward nodes from fusion patterns + // Protect every interior QSA node (not just the start): the mask branch is + // independent, so it gets pulled out and breaks keep_pattern otherwise. + auto const &in_qsa_pattern = [&](int n) -> bool { + for (int o = 0; o < (int) topk_qsa_pattern.size(); ++o) { + if (n - o >= 0 && match_pattern(topk_qsa_pattern, n - o)) { + return true; + } + } + return false; + }; if (match_pattern(topk_moe_early_softmax_norm, j) || match_pattern(topk_moe_sigmoid_norm_bias, j) || match_pattern(topk_moe_sqrt_softplus_norm_bias, j) || match_pattern(topk_moe_early_softmax, j) || match_pattern(topk_moe_late_softmax, j) || - match_pattern(snake_pattern, j)) { + match_pattern(snake_pattern, j) || + in_qsa_pattern(j)) { continue; } bool ok = true; @@ -18851,15 +19071,14 @@ static bool ggml_backend_vk_device_supports_op(ggml_backend_dev_t dev, const ggm if (!ggml_is_contiguous(op) || !ggml_is_contiguous(op->src[0])) { return false; } - // We could potentially support larger, using argsort to sort the - // whole thing. Not clear if this is needed. - uint32_t min_pipeline = (uint32_t)log2f(float(op->ne[0])) + 1; - if (min_pipeline >= num_topk_pipelines || - !device->pipeline_topk_f32[min_pipeline]) { - return false; + // large k falls back to radix-select + const uint32_t min_pipeline = + std::max((uint32_t) log2f(float(op->ne[0])) + 1, device->subgroup_size_log2); + if (min_pipeline < num_topk_pipelines && device->pipeline_topk_f32[min_pipeline]) { + return true; } + return device->pipeline_topk_radix_f32 != nullptr; } - return true; case GGML_OP_UPSCALE: if (op->op_params[0] & GGML_SCALE_FLAG_ANTIALIAS) { if ((op->op_params[0] & 0xFF) != GGML_SCALE_MODE_BILINEAR) { diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/topk_radix_select.comp b/ggml/src/ggml-vulkan/vulkan-shaders/topk_radix_select.comp new file mode 100644 index 000000000000..8e14b2e99253 --- /dev/null +++ b/ggml/src/ggml-vulkan/vulkan-shaders/topk_radix_select.comp @@ -0,0 +1,144 @@ +#version 450 + +#extension GL_EXT_control_flow_attributes : enable +#extension GL_EXT_shader_16bit_storage : require + +#include "types.glsl" + +layout(constant_id = 0) const int BLOCK_SIZE = 1024; +layout(constant_id = 1) const int QSA = 0; // 1: fuse the qwen4 QSA indexer gather + f16 mask + +layout(local_size_x_id = 0, local_size_y = 1, local_size_z = 1) in; + +layout (binding = 0) readonly buffer A {float data_a[];}; // input values, or QSA block scores [n_tps, n_blocks, n_stream] +layout (binding = 1) writeonly buffer D {int data_d[];}; // [k, ...] +layout (binding = 2) readonly buffer CB {int cell_blk[];}; // QSA: cell->block map [n_kv, n_stream] +layout (binding = 3) readonly buffer M {float16_t mask[];}; // QSA: raw f16 kq_mask [n_kv, n_tps, n_stream] +layout (binding = 4) buffer S {float scratch[];}; // QSA: [nrows, n_kv] gathered inputs + +layout (push_constant) uniform parameter { + uint ncols; + uint k; + uint nrows; + uint n_tps; // QSA only + uint n_blocks; // QSA only + uint n_stream; // QSA only +} p; + +#define RADIX_BITS 8 +#define RADIX_SIZE (1 << RADIX_BITS) + +shared uint histo[RADIX_SIZE]; +shared uint sh_bucket; +shared uint sh_above; +shared uint out_count; + +// order-preserving float -> uint mapping +uint f2ui(float x) { + uint y = floatBitsToUint(x); + if ((y & 0x80000000u) != 0u) { + y ^= 0xFFFFFFFFu; + } else { + y |= 0x80000000u; + } + return y; +} + +// QSA element i of row (t,s): score[cell_blk[i,s], t, s] + mask[i,t,s] +float gather(uint row, uint i) { + const uint t = row % p.n_tps; + const uint s = row / p.n_tps; + const uint block = uint(cell_blk[s * p.ncols + i]); + const float a = data_a[(s * p.n_blocks + block) * p.n_tps + t]; + const float m = float(mask[(s * p.n_tps + t) * p.ncols + i]); + return a + m; +} + +float load(uint row, uint i, bool first) { + if (QSA == 0) { + return data_a[row * p.ncols + i]; + } + // materialize the scattered gather on the first pass and reuse it after; each + // invocation only touches its own scratch entries, so no barrier is needed + const uint off = row * p.ncols + i; + if (first) { + const float v = gather(row, i); + scratch[off] = v; + return v; + } + return scratch[off]; +} + +// one workgroup per row: radix-select the K-th largest, then compact it plus enough ties +void topk(const uint row) { + const uint tid = gl_LocalInvocationID.x; + const uint ncols = p.ncols; + const uint row_out = row * p.k; + + uint prefix = 0; // fixed high bits of the threshold key + uint desired = p.k; // count still needed from the candidate range + + [[unroll]] for (int shift = 32 - RADIX_BITS; shift >= 0; shift -= RADIX_BITS) { + for (uint i = tid; i < RADIX_SIZE; i += BLOCK_SIZE) { + histo[i] = 0; + } + barrier(); + + const bool first = (shift == 32 - RADIX_BITS); + const uint hi_mask = (shift + RADIX_BITS >= 32) ? 0u : (0xFFFFFFFFu << uint(shift + RADIX_BITS)); + const uint prefix_hi = prefix & hi_mask; + for (uint i = tid; i < ncols; i += BLOCK_SIZE) { + const uint key = f2ui(load(row, i, first)); + if ((key & hi_mask) == prefix_hi) { + atomicAdd(histo[(key >> uint(shift)) & (RADIX_SIZE - 1)], 1u); + } + } + barrier(); + + // top-down scan for the bucket holding the K-th value + if (tid == 0) { + uint acc = 0; + uint b = 0; + for (int bb = RADIX_SIZE - 1; bb >= 0; --bb) { + const uint c = histo[bb]; + if (acc + c >= desired) { b = uint(bb); break; } + acc += c; + } + sh_bucket = b; + sh_above = acc; + } + barrier(); + + prefix |= sh_bucket << uint(shift); + desired -= sh_above; + barrier(); + } + + if (tid == 0) { + out_count = 0; + } + barrier(); + + // emit everything above the threshold, then fill the rest from ties + const uint threshold = prefix; + for (uint i = tid; i < ncols; i += BLOCK_SIZE) { + if (f2ui(load(row, i, false)) > threshold) { + data_d[row_out + atomicAdd(out_count, 1u)] = int(i); + } + } + barrier(); + for (uint i = tid; i < ncols; i += BLOCK_SIZE) { + if (f2ui(load(row, i, false)) == threshold) { + const uint pos = atomicAdd(out_count, 1u); + if (pos < p.k) { + data_d[row_out + pos] = int(i); + } + } + } +} + +void main() { + for (uint row = gl_WorkGroupID.y; row < p.nrows; row += gl_NumWorkGroups.y) { + topk(row); + } +} 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 f0610f4cd82d..27ff68c10d5b 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp @@ -1028,6 +1028,7 @@ void process_shaders() { string_to_spv("topk_argsort_f32", "topk_argsort.comp", {{"A_TYPE", "float"}}); string_to_spv("topk_nary_search_f32", "topk_nary_search.comp", {{"A_TYPE", "float"}}); + string_to_spv("topk_radix_select_f32", "topk_radix_select.comp", {{"A_TYPE", "float"}}); string_to_spv("argmax_f32", "argmax.comp", merge_maps(base_dict, {{"A_TYPE", "float"}, {"D_TYPE", "int"}})); string_to_spv("sum_rows_f32", "sum_rows.comp", merge_maps(base_dict, {{"A_TYPE", "float"}, {"D_TYPE", "float"}})); diff --git a/tests/test-backend-ops.cpp b/tests/test-backend-ops.cpp index a9d91acf6ee7..6f917f542db6 100644 --- a/tests/test-backend-ops.cpp +++ b/tests/test-backend-ops.cpp @@ -6287,6 +6287,87 @@ struct test_top_k : public test_case { } }; +// qwen4exp QSA indexer top-k fusion: expand per-block scores to cells, add the f16 mask, top-k. +struct test_topk_qsa : public test_case { + const int64_t n_blocks; + const int64_t n_kv; + const int64_t n_tps; + const int64_t n_stream; + const int width; + ggml_tensor * out {}; + + std::string op_desc(ggml_tensor * t) override { + GGML_UNUSED(t); + return "TOPK_QSA"; + } + + std::string vars() override { + return VARS_TO_STR5(n_blocks, n_kv, n_tps, n_stream, width); + } + + test_topk_qsa(int64_t n_blocks = 512, int64_t n_kv = 2048, int64_t n_tps = 2, int64_t n_stream = 1, int width = 1500) + : n_blocks(n_blocks), n_kv(n_kv), n_tps(n_tps), n_stream(n_stream), width(width) {} + + double max_err() override { return 0.0; } + bool run_whole_graph() override { return true; } + + ggml_tensor * build_graph(ggml_context * ctx) override { + ggml_tensor * score = ggml_new_tensor_3d(ctx, GGML_TYPE_F32, n_blocks, n_tps, n_stream); + ggml_set_name(score, "score"); + ggml_tensor * cell_blk = ggml_new_tensor_2d(ctx, GGML_TYPE_I32, n_kv, n_stream); + ggml_set_name(cell_blk, "cell_blk"); + ggml_tensor * kq_mask = ggml_new_tensor_3d(ctx, GGML_TYPE_F16, n_kv, n_tps, n_stream); + ggml_set_name(kq_mask, "kq_mask"); + + ggml_tensor * a = ggml_cont(ctx, ggml_permute(ctx, score, 1, 0, 2, 3)); + ggml_tensor * e = ggml_get_rows(ctx, a, cell_blk); + e = ggml_cont(ctx, ggml_permute(ctx, e, 1, 0, 2, 3)); + ggml_tensor * m = ggml_cast(ctx, kq_mask, GGML_TYPE_F32); + e = ggml_add(ctx, e, ggml_reshape_3d(ctx, m, n_kv, n_tps, n_stream)); + out = ggml_top_k(ctx, e, width); + ggml_set_name(out, "out"); + return out; + } + + std::vector fusion_test_nodes() override { return { out }; } + + // distinct mask ramp + small scores keep every cell value unique, so no top-k ties + void initialize_tensors(ggml_context * ctx) override { + for (ggml_tensor * t = ggml_get_first_tensor(ctx); t != NULL; t = ggml_get_next_tensor(ctx, t)) { + if (t->op != GGML_OP_NONE) { + continue; + } + if (t->type == GGML_TYPE_I32) { + std::vector data(ggml_nelements(t)); + for (auto & v : data) { v = rand() % n_blocks; } + ggml_backend_tensor_set(t, data.data(), 0, data.size() * sizeof(int32_t)); + } else if (t->type == GGML_TYPE_F16) { + std::vector data(ggml_nelements(t)); + for (int64_t r = 0; r < ggml_nrows(t); r++) { + for (int64_t i = 0; i < n_kv; i++) { + data[r * n_kv + i] = ggml_fp32_to_fp16((float) i); + } + } + ggml_backend_tensor_set(t, data.data(), 0, data.size() * sizeof(ggml_fp16_t)); + } else { + init_tensor_uniform(t, 0.0f, 0.5f); + } + } + } + + // top-k output order is unspecified; compare as a set of indices + double err(const float * a, const float * b, size_t n) override { + std::vector ia(n), ib(n); + double diff = 0.0; + for (size_t i = 0; i < n; i++) { + ia[i] = (int32_t) a[i]; + ib[i] = (int32_t) b[i]; + diff += std::fabs(a[i] - ia[i]) + std::fabs(b[i] - ib[i]); + } + return diff + jdst(ia.data(), ib.data(), n); + } +}; + enum MoeGatingFunc { GATING_FUNC_SOFTMAX, GATING_FUNC_SIGMOID, @@ -9817,6 +9898,22 @@ static std::vector> make_test_cases_eval() { test_cases.emplace_back(new test_top_k(GGML_TYPE_F32, {2049, 2, 1, 3}, k)); } + // Large-k, including multi-row and ties (qwen4exp) + test_cases.emplace_back(new test_top_k(GGML_TYPE_F32, { 1024, 1, 1, 1 }, 1024)); + test_cases.emplace_back(new test_top_k(GGML_TYPE_F32, { 2048, 2, 1, 1 }, 1024)); + test_cases.emplace_back(new test_top_k(GGML_TYPE_F32, { 4096, 1, 1, 1 }, 2048)); + test_cases.emplace_back(new test_top_k(GGML_TYPE_F32, { 8192, 2, 1, 1 }, 2051)); + test_cases.emplace_back(new test_top_k(GGML_TYPE_F32, { 33024, 1, 1, 1 }, 2051)); + test_cases.emplace_back(new test_top_k(GGML_TYPE_F32, { 33024, 4, 1, 1 }, 2051)); + test_cases.emplace_back(new test_top_k(GGML_TYPE_F32, { 8192, 2, 1, 1 }, 2051, true)); + test_cases.emplace_back(new test_top_k(GGML_TYPE_F32, { 33024, 4, 1, 1 }, 2051, true)); + + // qwen4exp QSA indexer top-k fusion (get_rows + f16 mask + top_k) + test_cases.emplace_back(new test_topk_qsa(512, 2048, 1, 1, 1500)); + test_cases.emplace_back(new test_topk_qsa(512, 2048, 2, 1, 1500)); + test_cases.emplace_back(new test_topk_qsa(256, 2048, 4, 2, 2000)); + test_cases.emplace_back(new test_topk_qsa(64, 256, 2, 1, 200)); // small k: unfused fallback + // exhaustive top_k tests //for (int i = 1; i < 9999; ++i) { // test_cases.emplace_back(new test_top_k(GGML_TYPE_F32, {i, 2, 1, 3}, rand() % i + 1)); From 36faf5415e794f1e362ec1774d2b7c8f88f7a6d5 Mon Sep 17 00:00:00 2001 From: Nathan Wilson Date: Thu, 3 Sep 2026 02:36:12 +0000 Subject: [PATCH 3/6] vulkan: deterministic slot assignment in top_k radix select The radix select (#28032) assigned output slots with atomicAdd, so identical runs emitted the selected indices in a scheduling-dependent order and, at the tie boundary, could drop a different tied element. On Qwen 3.8 Flash-Next the sparse attention sums the selection in that order, the ULP noise cascades through the QSA layers, and identical greedy requests above 2051 prompt tokens diverge (32k prompt x4: 2 to 3 distinct on every unfixed build, upstream included). Replace both atomicAdd slot counters with a per-chunk subgroup-ballot exclusive scan, walking the columns in ascending index order: the output order is fixed (ascending, which also helps gather locality) and the tie fill keeps the lowest-indexed tied elements. The histogram atomicAdd is unchanged. Verified on gfx1151: test-backend-ops TOP_K 453/453; 32k x4 byte-identical (was 2 to 3 distinct); 1024 x16 byte-identical; layer 3 and layer 47 top-k dumps identical in set and order across runs; prefill 337.5 vs 339.6 t/s and decode 26.3 vs 26.2 t/s at 16k depth (within noise). Co-authored-by: Claude (Opus 5) (cherry picked from commit fe6620c88c599b2f3a6e5ff4a6b99bb4773be38e) --- .../vulkan-shaders/topk_radix_select.comp | 55 +++++++++++++------ 1 file changed, 38 insertions(+), 17 deletions(-) diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/topk_radix_select.comp b/ggml/src/ggml-vulkan/vulkan-shaders/topk_radix_select.comp index 8e14b2e99253..c8975391b30f 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/topk_radix_select.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/topk_radix_select.comp @@ -2,6 +2,8 @@ #extension GL_EXT_control_flow_attributes : enable #extension GL_EXT_shader_16bit_storage : require +#extension GL_KHR_shader_subgroup_basic : enable +#extension GL_KHR_shader_subgroup_ballot : enable #include "types.glsl" @@ -31,7 +33,7 @@ layout (push_constant) uniform parameter { shared uint histo[RADIX_SIZE]; shared uint sh_bucket; shared uint sh_above; -shared uint out_count; +shared uint sg_cnt[64]; // per-subgroup hit counts for the slot scan // order-preserving float -> uint mapping uint f2ui(float x) { @@ -114,25 +116,44 @@ void topk(const uint row) { barrier(); } - if (tid == 0) { - out_count = 0; - } - barrier(); - // emit everything above the threshold, then fill the rest from ties + // Emit everything above the threshold, then fill the rest from ties. Slots come from + // an exclusive scan over the candidate flags, one BLOCK_SIZE chunk at a time in ascending + // index order, so the output is identical on every run. The previous atomicAdd slot + // counter made the ORDER scheduling-dependent, and at the tie boundary the SET as well: + // the QSA width is top_k + ratio - 1, so the boundary block's ratio tied cells race for + // ratio - 1 slots and a different cell lost each run (non-repeatable output at depth). + // With the scan the lowest-indexed tied cells win. const uint threshold = prefix; - for (uint i = tid; i < ncols; i += BLOCK_SIZE) { - if (f2ui(load(row, i, false)) > threshold) { - data_d[row_out + atomicAdd(out_count, 1u)] = int(i); - } - } - barrier(); - for (uint i = tid; i < ncols; i += BLOCK_SIZE) { - if (f2ui(load(row, i, false)) == threshold) { - const uint pos = atomicAdd(out_count, 1u); - if (pos < p.k) { - data_d[row_out + pos] = int(i); + uint base = 0; + [[dont_unroll]] for (uint pass = 0; pass < 2; ++pass) { + for (uint c = 0; c < ncols; c += BLOCK_SIZE) { + const uint i = c + tid; + bool hit = false; + if (i < ncols) { + const uint key = f2ui(load(row, i, false)); + hit = (pass == 0) ? (key > threshold) : (key == threshold); + } + const uvec4 ballot = subgroupBallot(hit); + const uint rank_sg = subgroupBallotExclusiveBitCount(ballot); + const uint cnt_sg = subgroupBallotBitCount(ballot); + if (subgroupElect()) { + sg_cnt[gl_SubgroupID] = cnt_sg; + } + barrier(); + uint sg_base = 0; + uint total = 0; + for (uint sg = 0; sg < gl_NumSubgroups; ++sg) { + const uint v = sg_cnt[sg]; + sg_base += (sg < gl_SubgroupID) ? v : 0; + total += v; + } + const uint slot = base + sg_base + rank_sg; + if (hit && slot < p.k) { + data_d[row_out + slot] = int(i); } + base += total; + barrier(); } } } From 50b83f81d89f7b7a2c36c64d1a509e869755ab48 Mon Sep 17 00:00:00 2001 From: danielhanchen Date: Mon, 31 Aug 2026 02:58:30 +0000 Subject: [PATCH 4/6] models: use flash-linear-attention's l2norm for gated delta net q/k The GDN q/k normalization is defined by flash-linear-attention as l2norm(x) = x * rsqrt(sum(x*x) + eps) with eps inside the root. Every GDN call site in the tree uses ggml_l2_norm instead, which is x / max(sqrt(sum(x*x)), eps), i.e. torch.nn.functional.normalize - its CUDA kernel cites that page. The clamp never engages at these magnitudes, so in practice llama.cpp normalizes with no epsilon at all where the reference has one inside the root. transformers made the same substitution when it first added Qwen3-Next and corrected it three days later in huggingface/transformers#40842, 'Fix the misalignment between the l2norm in GDN of Qwen3-Next and the implementation in the FLA library'. vLLM and SGLang vendor FLA rather than reimplementing it, so neither ever had the clamp. eps keeps coming from the checkpoint, exactly as every call site already passed it. The references hardcode 1e-6 for this norm; that is a separate question and the two agree on every GDN checkpoint in the wild. ggml_l2_norm itself is correct and unchanged, as is rwkv7-base, its original caller, which passes normalize's own default eps of 1e-12. No new ggml op: rms_norm already carries eps inside the root, so rms_norm(x, eps/n) * (1/sqrt(n)) is exactly x * rsqrt(sum(x*x) + eps). (cherry picked from commit 0dd72fd59ec4062daea5ac3c16095732c3ba5ac7) --- src/models/bailingmoe3.cpp | 4 ++-- src/models/kimi-k3.cpp | 6 +++--- src/models/kimi-linear.cpp | 5 +++-- src/models/models.h | 6 ++++++ src/models/qwen35.cpp | 5 +++-- src/models/qwen35moe.cpp | 5 +++-- src/models/qwen3next.cpp | 5 +++-- src/models/qwen4exp.cpp | 5 +++-- 8 files changed, 26 insertions(+), 15 deletions(-) diff --git a/src/models/bailingmoe3.cpp b/src/models/bailingmoe3.cpp index 0637931cc0c9..abc4440a2fea 100644 --- a/src/models/bailingmoe3.cpp +++ b/src/models/bailingmoe3.cpp @@ -281,8 +281,8 @@ llama_model_bailingmoe3::graph::graph(const llama_model & model, const llm_graph ggml_tensor * beta = ggml_mul_mat(ctx0, layer.ssm_beta, cur); beta = ggml_sigmoid(ctx0, ggml_reshape_4d(ctx0, beta, 1, n_head, n_seq_tokens, n_seqs)); - q = ggml_l2_norm(ctx0, q, hparams.f_norm_rms_eps); - k = ggml_l2_norm(ctx0, k, hparams.f_norm_rms_eps); + q = build_gdn_l2_norm(ctx0, q, hparams.f_norm_rms_eps); + k = build_gdn_l2_norm(ctx0, k, hparams.f_norm_rms_eps); ggml_tensor * states_all = mctx_cur->get_s_l(il); ggml_tensor * state = build_rs(inp_rs, states_all, hparams.n_embd_s(), n_seqs); diff --git a/src/models/kimi-k3.cpp b/src/models/kimi-k3.cpp index d952d72cdf13..6d09ed00e062 100644 --- a/src/models/kimi-k3.cpp +++ b/src/models/kimi-k3.cpp @@ -441,9 +441,9 @@ ggml_tensor * llama_model_kimi_k3::graph::build_kda_layer( ggml_tensor * state = build_rs(inp_rs, ssm_states_all, hparams.n_embd_s(), n_seqs); state = ggml_reshape_4d(ctx0, state, head_dim, head_dim, n_head_kda, n_seqs); - const float eps = hparams.f_norm_rms_eps; - Qcur = ggml_l2_norm(ctx0, Qcur, eps); - Kcur = ggml_l2_norm(ctx0, Kcur, eps); + const float eps_norm = hparams.f_norm_rms_eps; + Qcur = build_gdn_l2_norm(ctx0, Qcur, eps_norm); + Kcur = build_gdn_l2_norm(ctx0, Kcur, eps_norm); auto attn_out = build_delta_net(Qcur, Kcur, Vcur, g1, beta, state, il); diff --git a/src/models/kimi-linear.cpp b/src/models/kimi-linear.cpp index 367f6990d1fb..69962f220f81 100644 --- a/src/models/kimi-linear.cpp +++ b/src/models/kimi-linear.cpp @@ -331,10 +331,11 @@ llama_model_kimi_linear::graph::graph(const llama_model & model, const llm_graph ggml_tensor * state = build_rs(inp_rs, ssm_states_all, hparams.n_embd_s(), n_seqs); state = ggml_reshape_4d(ctx0, state, head_dim, head_dim, n_head, n_seqs); + const float eps_norm = hparams.f_norm_rms_eps; - Qcur = ggml_l2_norm(ctx0, Qcur, eps_norm); - Kcur = ggml_l2_norm(ctx0, Kcur, eps_norm); + Qcur = build_gdn_l2_norm(ctx0, Qcur, eps_norm); + Kcur = build_gdn_l2_norm(ctx0, Kcur, eps_norm); // Choose between build_delta_net_chunking and build_delta_net_recurrent based on n_tokens auto attn_out = build_delta_net(Qcur, Kcur, Vcur, g1, beta, state, il); diff --git a/src/models/models.h b/src/models/models.h index f285d3d315ac..f2d90cd45d3a 100644 --- a/src/models/models.h +++ b/src/models/models.h @@ -10,6 +10,12 @@ class llama_memory_hybrid_idx_context; +static inline ggml_tensor * build_gdn_l2_norm(ggml_context * ctx, ggml_tensor * x, float eps) { + const float n = x->ne[0]; + + return ggml_scale(ctx, ggml_rms_norm(ctx, x, eps/n), 1.0f/sqrtf(n)); +} + // // base classes // diff --git a/src/models/qwen35.cpp b/src/models/qwen35.cpp index 309dd432447c..f6a71672c4d6 100644 --- a/src/models/qwen35.cpp +++ b/src/models/qwen35.cpp @@ -427,10 +427,11 @@ ggml_tensor * llama_model_qwen35::graph::build_layer_attn_linear( cb(k_conv, "k_conv", il); cb(v_conv, "v_conv", il); + const float eps_norm = hparams.f_norm_rms_eps; - q_conv = ggml_l2_norm(ctx0, q_conv, eps_norm); - k_conv = ggml_l2_norm(ctx0, k_conv, eps_norm); + q_conv = build_gdn_l2_norm(ctx0, q_conv, eps_norm); + k_conv = build_gdn_l2_norm(ctx0, k_conv, eps_norm); //q_conv = ggml_cont_4d(ctx0, q_conv, head_k_dim, num_k_heads, n_seq_tokens, n_seqs); //k_conv = ggml_cont_4d(ctx0, k_conv, head_k_dim, num_k_heads, n_seq_tokens, n_seqs); diff --git a/src/models/qwen35moe.cpp b/src/models/qwen35moe.cpp index 38f2a57985a9..988cbd6ccd7a 100644 --- a/src/models/qwen35moe.cpp +++ b/src/models/qwen35moe.cpp @@ -451,10 +451,11 @@ ggml_tensor * llama_model_qwen35moe::graph::build_layer_attn_linear( cb(k_conv, "k_conv", il); cb(v_conv, "v_conv", il); + const float eps_norm = hparams.f_norm_rms_eps; - q_conv = ggml_l2_norm(ctx0, q_conv, eps_norm); - k_conv = ggml_l2_norm(ctx0, k_conv, eps_norm); + q_conv = build_gdn_l2_norm(ctx0, q_conv, eps_norm); + k_conv = build_gdn_l2_norm(ctx0, k_conv, eps_norm); //q_conv = ggml_cont_4d(ctx0, q_conv, head_k_dim, num_k_heads, n_seq_tokens, n_seqs); //k_conv = ggml_cont_4d(ctx0, k_conv, head_k_dim, num_k_heads, n_seq_tokens, n_seqs); diff --git a/src/models/qwen3next.cpp b/src/models/qwen3next.cpp index 0808fd87aa0e..783289f4d1de 100644 --- a/src/models/qwen3next.cpp +++ b/src/models/qwen3next.cpp @@ -507,10 +507,11 @@ ggml_tensor * llama_model_qwen3next::graph::build_layer_attn_linear( cb(k_conv, "k_conv", il); cb(v_conv, "v_conv", il); + const float eps_norm = hparams.f_norm_rms_eps; - q_conv = ggml_l2_norm(ctx0, q_conv, eps_norm); - k_conv = ggml_l2_norm(ctx0, k_conv, eps_norm); + q_conv = build_gdn_l2_norm(ctx0, q_conv, eps_norm); + k_conv = build_gdn_l2_norm(ctx0, k_conv, eps_norm); //q_conv = ggml_cont_4d(ctx0, q_conv, head_k_dim, num_k_heads, n_seq_tokens, n_seqs); //k_conv = ggml_cont_4d(ctx0, k_conv, head_k_dim, num_k_heads, n_seq_tokens, n_seqs); diff --git a/src/models/qwen4exp.cpp b/src/models/qwen4exp.cpp index ca47c3706cae..9d40e8683b0d 100644 --- a/src/models/qwen4exp.cpp +++ b/src/models/qwen4exp.cpp @@ -1084,10 +1084,11 @@ ggml_tensor * llama_model_qwen4exp::graph::build_layer_attn_linear( cb(k_conv, "k_conv", il); cb(v_conv, "v_conv", il); + const float eps_norm = hparams.f_norm_rms_eps; - q_conv = ggml_l2_norm(ctx0, q_conv, eps_norm); - k_conv = ggml_l2_norm(ctx0, k_conv, eps_norm); + q_conv = build_gdn_l2_norm(ctx0, q_conv, eps_norm); + k_conv = build_gdn_l2_norm(ctx0, k_conv, eps_norm); // repeat to match shapes when head keys != value keys; unneeded with the fused GDN if (num_k_heads != num_v_heads && (!cparams.fused_gdn_ar || !cparams.fused_gdn_ch)) { From 8d37b7c0960b13c59a0cf099d3c567f2ff3017f9 Mon Sep 17 00:00:00 2001 From: Daniel Han Date: Mon, 31 Aug 2026 00:12:52 -0700 Subject: [PATCH 5/6] Update src/models/models.h Co-authored-by: Georgi Gerganov (cherry picked from commit fe4632898773e1fe58acbc7fcf1505a0c56c898f) --- src/models/models.h | 1 + 1 file changed, 1 insertion(+) diff --git a/src/models/models.h b/src/models/models.h index f2d90cd45d3a..168c5366b9e8 100644 --- a/src/models/models.h +++ b/src/models/models.h @@ -10,6 +10,7 @@ class llama_memory_hybrid_idx_context; +// ref: https://github.com/ggml-org/llama.cpp/pull/28068 static inline ggml_tensor * build_gdn_l2_norm(ggml_context * ctx, ggml_tensor * x, float eps) { const float n = x->ne[0]; From aad5adb08f5d59925990b28cd0d8cebc0ffc5880 Mon Sep 17 00:00:00 2001 From: Nathan Wilson Date: Thu, 3 Sep 2026 09:54:56 +0000 Subject: [PATCH 6/6] kv-cache: zero freed cells so masked-out rows never carry stale K/V On gfx11 (RDNA3) WMMA, x + (-0.0) is not exact in f16, so a fully masked flash-attention column with P == +0.0 still leaks the sign of whatever V the cell last held into the accumulator: the output of a request depends on what the previous request left in the cache. On a 1024-token prompt with 16 identical greedy requests, Qwen2.5-7B (f16 KV) alternates between two continuations and Qwen3.8 Flash-Next gives 5-6 distinct ones. Rather than patching the shader (which costs 8-18% dense prefill at depth on gfx1151 by its mere presence in flash_attn_cm1.comp), keep the invariant that a free cell is always zero: the KV buffers are cleared at construction, and every cell that becomes free is memset again, in every layer, when it is freed (seq_rm, seq_keep; clear() wipes up to a per-stream high-water mark, and stream copies mirror the mark). Caches that mirror another cache's cells (the Flash-Next indexer) register with the owner and are zeroed with it. On UMA the Vulkan backend's tensor memset is a host memset; it happens once per free, never per token, and never as a graph node, so graph reuse cannot replay it. Not covered: masked-out cells that belong to another sequence in a unified multi-sequence cache; only a shader-side fix covers that case. Gate (16 identical greedy requests, 1024-token prompt, 129 tokens, per-position top-8 logprob streams compared): 16/16 identical; the fresh-server output is unchanged byte for byte. Co-authored-by: Claude (Opus 5) --- src/llama-kv-cache.cpp | 113 +++++++++++++++++++++++++++++++++++++++++ src/llama-kv-cache.h | 18 ++++++- 2 files changed, 130 insertions(+), 1 deletion(-) diff --git a/src/llama-kv-cache.cpp b/src/llama-kv-cache.cpp index 65afbd8c3778..0542a4b863c5 100644 --- a/src/llama-kv-cache.cpp +++ b/src/llama-kv-cache.cpp @@ -100,6 +100,11 @@ llama_kv_cache::llama_kv_cache( GGML_ASSERT(kv_size % n_pad == 0); + if (other) { + // our rows go stale exactly when the owner frees a cell: let it zero them with its own + other->sharers.push_back(this); + } + const uint32_t n_layer = hparams.n_layer_all; // define a comparator for the buft -> ctx map to ensure that the order is well-defined: @@ -368,6 +373,15 @@ llama_kv_cache::llama_kv_cache( } void llama_kv_cache::clear(bool data) { + if (!other) { + // every cell becomes free: wipe the rows that were ever written (the sharers' rows too, + // their own clear() may not run, and a buffer clear below only covers our buffers) + rows_hw_init(); + for (uint32_t s = 0; s < n_stream; ++s) { + zero_rows(s, 0, rows_hw[s]); + rows_hw[s] = 0; + } + } for (uint32_t s = 0; s < n_stream; ++s) { v_cells[s].reset(); v_heads[s] = 0; @@ -403,6 +417,8 @@ bool llama_kv_cache::seq_rm(llama_seq_id seq_id, llama_pos p0, llama_pos p1) { uint32_t new_head = cells.size(); + std::vector freed; + for (uint32_t i = 0; i < cells.size(); ++i) { if (!cells.pos_in(i, p0, p1)) { continue; @@ -412,9 +428,13 @@ bool llama_kv_cache::seq_rm(llama_seq_id seq_id, llama_pos p0, llama_pos p1) { if (new_head == cells.size()) { new_head = i; } + + freed.push_back(i); } } + zero_idxs(seq_to_stream[seq_id], freed); + // If we freed up a slot, set head to it so searching can start there. if (new_head != cells.size() && new_head < head) { head = new_head; @@ -427,6 +447,8 @@ bool llama_kv_cache::seq_rm(llama_seq_id seq_id, llama_pos p0, llama_pos p1) { uint32_t new_head = cells.size(); + std::vector freed; + for (uint32_t i = 0; i < cells.size(); ++i) { if (!cells.pos_in(i, p0, p1)) { continue; @@ -437,8 +459,12 @@ bool llama_kv_cache::seq_rm(llama_seq_id seq_id, llama_pos p0, llama_pos p1) { if (new_head == cells.size()) { new_head = i; } + + freed.push_back(i); } + zero_idxs(s, freed); + // If we freed up a slot, set head to it so searching can start there. if (new_head != cells.size() && new_head < head) { head = new_head; @@ -554,14 +580,20 @@ void llama_kv_cache::seq_keep(llama_seq_id seq_id) { uint32_t new_head = cells.size(); + std::vector freed; + for (uint32_t i = 0; i < cells.size(); ++i) { if (cells.seq_keep(i, seq_id)) { if (new_head == cells.size()) { new_head = i; } + + freed.push_back(i); } } + zero_idxs(seq_to_stream[seq_id], freed); + // If we freed up a slot, set head to it so searching can start there. if (new_head != cells.size() && new_head < head) { head = new_head; @@ -841,6 +873,10 @@ bool llama_kv_cache::update(llama_context * lctx, bool do_shift, const stream_co LLAMA_LOG_DEBUG("%s: copying KV buffer: stream %d to stream %d\n", __func__, ssrc, sdst); + // the whole stream is copied, rows above the source high-water mark included + rows_hw_init(); + rows_hw[sdst] = rows_hw[ssrc]; + assert(ssrc != sdst); for (uint32_t il = 0; il < layers.size(); ++il) { @@ -1184,6 +1220,15 @@ void llama_kv_cache::apply_ubatch(const slot_info & sinfo, const llama_ubatch & head = sinfo.idxs[s].back() + 1; } + + // rows at or above the highest used cell are zero (buffer clear at construction, zero_rows + // on every free); remember the written extent so clear() only wipes what it has to + rows_hw_init(); + for (uint32_t s = 0; s < sinfo.n_stream(); ++s) { + const uint32_t strm = sinfo.strm[s]; + + rows_hw[strm] = std::max(rows_hw[strm], v_cells[strm].used_max_p1()); + } } bool llama_kv_cache::get_can_shift() const { @@ -1316,6 +1361,74 @@ ggml_tensor * llama_kv_cache::get_v(ggml_context * ctx, int32_t il, uint32_t n_k ggml_row_size(v->type, kv_size*n_embd_v_gqa)*sinfo.s0); } +llama_kv_cache::~llama_kv_cache() { + if (other) { + auto & v = other->sharers; + v.erase(std::remove(v.begin(), v.end(), this), v.end()); + } + for (auto * c : sharers) { + c->other = nullptr; + } +} + +void llama_kv_cache::rows_hw_init() { + if (rows_hw.size() != n_stream) { + rows_hw.assign(n_stream, 0); + } +} + +void llama_kv_cache::zero_rows(uint32_t strm, uint32_t r0, uint32_t r1) { + r1 = std::min(r1, get_size()); + + if (r0 >= r1) { + return; + } + + const auto zero = [strm, r0, r1](const llama_kv_cache * c) { + for (const auto & layer : c->layers) { + // V stored column-major (no flash attention) has no contiguous rows to wipe; that + // path does not go through the WMMA accumulate the zeroing is for + for (ggml_tensor * t : { layer.k, c->v_trans ? nullptr : layer.v }) { + if (!t || !t->buffer) { + continue; + } + + const size_t offset = (size_t) strm*t->nb[2] + (size_t) r0*t->nb[1]; + const size_t size = (size_t) (r1 - r0)*t->nb[1]; + + if (offset + size > ggml_nbytes(t)) { + continue; + } + + // a host memset on UMA, a fill command elsewhere; never a graph node, so graph + // reuse cannot replay it with stale offsets + ggml_backend_tensor_memset(t, 0, offset, size); + } + } + }; + + zero(this); + + for (const auto * c : sharers) { + zero(c); + } +} + +void llama_kv_cache::zero_idxs(uint32_t strm, const std::vector & idxs) { + // coalesce runs of consecutive cells into one memset each + size_t i = 0; + while (i < idxs.size()) { + size_t j = i; + while (j + 1 < idxs.size() && idxs[j + 1] == idxs[j] + 1) { + ++j; + } + + zero_rows(strm, idxs[i], idxs[j] + 1); + + i = j + 1; + } +} + ggml_tensor * llama_kv_cache::cpy_k(ggml_context * ctx, ggml_tensor * k_cur, ggml_tensor * k_idxs, int32_t il, const slot_info & sinfo) const { GGML_UNUSED(sinfo); diff --git a/src/llama-kv-cache.h b/src/llama-kv-cache.h index c4d8699def12..53b896b2e463 100644 --- a/src/llama-kv-cache.h +++ b/src/llama-kv-cache.h @@ -116,7 +116,7 @@ class llama_kv_cache : public llama_memory_i { // a model can hold more than one cache, so the tensor names have to stay unique const char * name_tag = ""); - ~llama_kv_cache() = default; + ~llama_kv_cache() override; // // llama_memory_i @@ -189,6 +189,14 @@ class llama_kv_cache : public llama_memory_i { ggml_tensor * get_k(ggml_context * ctx, int32_t il, uint32_t n_kv, const slot_info & sinfo) const; ggml_tensor * get_v(ggml_context * ctx, int32_t il, uint32_t n_kv, const slot_info & sinfo) const; + // zero the K rows (and the V rows when V is stored row-major) [r0, r1) of stream strm in + // every layer, including the layers of the caches that mirror these cells (sharers). + // Called for every cell that becomes free, so a masked-out cell never holds stale content: + // on gfx11 WMMA, P*V with P == +0.0 is not exact for V != +0.0, so the FA output would + // otherwise depend on whatever the masked-out cells last held. + void zero_rows(uint32_t strm, uint32_t r0, uint32_t r1); + void zero_idxs(uint32_t strm, const std::vector & idxs); // ascending cell indices + // store k_cur and v_cur in the cache based on the provided head location ggml_tensor * cpy_k(ggml_context * ctx, ggml_tensor * k_cur, ggml_tensor * k_idxs, int32_t il, const slot_info & sinfo) const; ggml_tensor * cpy_v(ggml_context * ctx, ggml_tensor * v_cur, ggml_tensor * v_idxs, int32_t il, const slot_info & sinfo) const; @@ -295,6 +303,14 @@ class llama_kv_cache : public llama_memory_i { // note: this is not part of the KV state and it's only used to speed-up the find_slot() method std::vector v_heads; + // one past the highest cell row ever written, per stream. Rows above it are still zero + // from the buffer clear at construction, so clear() only has to wipe [0, rows_hw). + std::vector rows_hw; + void rows_hw_init(); + + // caches that mirror our cells (see `other`): their rows go stale exactly when ours do + std::vector sharers; + // TODO: temporary until we refactor to be able to share the same cells between 2 kv caches [TAG_KV_CACHE_SHARE_CELLS] llama_kv_cache * other;