diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp index de8e2673067b..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 ] @@ -1058,6 +1073,7 @@ struct vk_device_struct { 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; @@ -1750,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; @@ -2440,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; @@ -5857,11 +5884,12 @@ 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 + // 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 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); + 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); @@ -14027,7 +14055,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,14 +14099,36 @@ 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]; + // 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) { @@ -14182,46 +14232,53 @@ 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]; +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]; - 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 }; + 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] - 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 }; + // 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]; + } - ggml_pipeline_request_descriptor_sets(ctx, pipeline, 1); - ggml_vk_dispatch_pipeline(ctx, subctx, pipeline, { src0_buf, dst_buf }, pc, elements); -} + 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; -// 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; -} + vk_pipeline pipeline = ctx->device->pipeline_topk_radix_qsa; + GGML_ASSERT(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); + // 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) { @@ -15885,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: @@ -17297,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); @@ -17676,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); @@ -17765,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)) { @@ -17881,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; } } @@ -18077,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); @@ -18095,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; @@ -18904,12 +19071,13 @@ 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; + // 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; } case GGML_OP_UPSCALE: if (op->op_params[0] & GGML_SCALE_FLAG_ANTIALIAS) { 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/topk_radix_select.comp b/ggml/src/ggml-vulkan/vulkan-shaders/topk_radix_select.comp new file mode 100644 index 000000000000..c8975391b30f --- /dev/null +++ b/ggml/src/ggml-vulkan/vulkan-shaders/topk_radix_select.comp @@ -0,0 +1,165 @@ +#version 450 + +#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" + +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 sg_cnt[64]; // per-subgroup hit counts for the slot scan + +// 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(); + } + + + // 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; + 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(); + } + } +} + +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 3624fb51a43f..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,8 +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_f32", "topk_radix.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/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; 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..168c5366b9e8 100644 --- a/src/models/models.h +++ b/src/models/models.h @@ -10,6 +10,13 @@ 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]; + + 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)) { diff --git a/tests/test-backend-ops.cpp b/tests/test-backend-ops.cpp index 516c8094739e..6f917f542db6 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,18 +6281,91 @@ 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)); + } + } + } +}; + +// 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(), r * t->nb[1], t->ne[0] * sizeof(float)); + 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 { @@ -9800,12 +9872,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)); @@ -9834,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));