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
This commit is contained in:
@@ -657,6 +657,21 @@ static constexpr std::initializer_list<ggml_op> snake_pattern { GGM
|
|||||||
GGML_OP_SQR, GGML_OP_MUL,
|
GGML_OP_SQR, GGML_OP_MUL,
|
||||||
GGML_OP_ADD };
|
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<ggml_op> 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<std::array<int, 3>> 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 #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 #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 ]
|
//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_f32[num_argsort_pipelines];
|
||||||
vk_pipeline pipeline_argsort_large_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_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_sum_rows_f32;
|
||||||
vk_pipeline pipeline_cross_entropy_loss_f32, pipeline_cross_entropy_loss_f32_wg512;
|
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;
|
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;
|
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 {
|
struct vk_op_im2col_push_constants {
|
||||||
uint64_t dst_addr;
|
uint64_t dst_addr;
|
||||||
uint32_t batch_offset; uint32_t offset_delta;
|
uint32_t batch_offset; uint32_t offset_delta;
|
||||||
@@ -2439,6 +2465,8 @@ struct ggml_backend_vk_context {
|
|||||||
int fused_ops_write_mask {};
|
int fused_ops_write_mask {};
|
||||||
topk_moe_mode fused_topk_moe_mode {};
|
topk_moe_mode fused_topk_moe_mode {};
|
||||||
bool fused_topk_moe_scale {};
|
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
|
// for GGML_VK_PERF_LOGGER
|
||||||
std::unique_ptr<vk_perf_logger> perf_logger;
|
std::unique_ptr<vk_perf_logger> perf_logger;
|
||||||
@@ -5814,6 +5842,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_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);
|
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);
|
||||||
@@ -13940,6 +13976,31 @@ static void ggml_vk_topk(ggml_backend_vk_context * ctx, vk_context& subctx, cons
|
|||||||
uint32_t nrows = ggml_nrows(src0);
|
uint32_t nrows = ggml_nrows(src0);
|
||||||
uint32_t k = dst->ne[0];
|
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<uint32_t, 3> 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 };
|
vk_op_topk_push_constants pc { ncols, ncols, ncols, k, nrows, 0, 0 };
|
||||||
|
|
||||||
if (ctx->prealloc_x_need_sync) {
|
if (ctx->prealloc_x_need_sync) {
|
||||||
@@ -14043,6 +14104,55 @@ static void ggml_vk_topk(ggml_backend_vk_context * ctx, vk_context& subctx, cons
|
|||||||
ctx->prealloc_x_need_sync = true;
|
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<uint32_t, 3> 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) {
|
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));
|
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);
|
ggml_vk_op_f32(ctx, subctx, src0, nullptr, nullptr, nullptr, dst, GGML_OP_SUM, p);
|
||||||
@@ -15704,7 +15814,11 @@ static bool ggml_vk_build_graph(ggml_backend_vk_context * ctx, ggml_cgraph * cgr
|
|||||||
|
|
||||||
break;
|
break;
|
||||||
case GGML_OP_GET_ROWS:
|
case GGML_OP_GET_ROWS:
|
||||||
|
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);
|
ggml_vk_get_rows(ctx, compute_ctx, src0, src1, node);
|
||||||
|
}
|
||||||
|
|
||||||
break;
|
break;
|
||||||
case GGML_OP_GET_ROWS_BACK:
|
case GGML_OP_GET_ROWS_BACK:
|
||||||
@@ -17116,6 +17230,92 @@ static bool ggml_vk_can_fuse_topk_moe(ggml_backend_vk_context * ctx, const struc
|
|||||||
return true;
|
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<ggml_op> & 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,
|
static bool ggml_vk_can_fuse_rope_set_rows(ggml_backend_vk_context * ctx, const struct ggml_cgraph * cgraph,
|
||||||
int node_idx) {
|
int node_idx) {
|
||||||
GGML_UNUSED(ctx);
|
GGML_UNUSED(ctx);
|
||||||
@@ -17495,6 +17695,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_mode = TOPK_MOE_COUNT;
|
||||||
ctx->fused_topk_moe_scale = false;
|
ctx->fused_topk_moe_scale = false;
|
||||||
|
ctx->fused_topk_qsa = false;
|
||||||
const char *fusion_string {};
|
const char *fusion_string {};
|
||||||
if (!ctx->device->disable_fusion) {
|
if (!ctx->device->disable_fusion) {
|
||||||
uint32_t num_adds = ggml_vk_fuse_multi_add(ctx, cgraph, i);
|
uint32_t num_adds = ggml_vk_fuse_multi_add(ctx, cgraph, i);
|
||||||
@@ -17584,6 +17785,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
|
// with a data dependency on that register. The overlap check still
|
||||||
// rejects partial overlaps (different base or size).
|
// rejects partial overlaps (different base or size).
|
||||||
std::fill_n(op_srcs_fused_elementwise, 5, true);
|
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 }) &&
|
} 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_check_edges(cgraph, i, topk_moe_early_softmax_norm_edges) &&
|
||||||
ggml_vk_can_fuse_topk_moe(ctx, cgraph, i, TOPK_MOE_EARLY_SOFTMAX_NORM)) {
|
ggml_vk_can_fuse_topk_moe(ctx, cgraph, i, TOPK_MOE_EARLY_SOFTMAX_NORM)) {
|
||||||
@@ -17700,6 +17906,7 @@ static ggml_status ggml_backend_vk_graph_compute(ggml_backend_t backend, ggml_cg
|
|||||||
ctx->fused_ops_write_mask = 1;
|
ctx->fused_ops_write_mask = 1;
|
||||||
ctx->fused_topk_moe_mode = TOPK_MOE_COUNT;
|
ctx->fused_topk_moe_mode = TOPK_MOE_COUNT;
|
||||||
ctx->fused_topk_moe_scale = false;
|
ctx->fused_topk_moe_scale = false;
|
||||||
|
ctx->fused_topk_qsa = false;
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -17896,6 +18103,9 @@ static void ggml_vk_graph_optimize(ggml_backend_t backend, struct ggml_cgraph *
|
|||||||
if (keep_pattern(snake_pattern)) {
|
if (keep_pattern(snake_pattern)) {
|
||||||
continue;
|
continue;
|
||||||
}
|
}
|
||||||
|
if (keep_pattern(topk_qsa_pattern)) {
|
||||||
|
continue;
|
||||||
|
}
|
||||||
|
|
||||||
// First, grab the next unused node.
|
// First, grab the next unused node.
|
||||||
current_set.push_back(first_unused);
|
current_set.push_back(first_unused);
|
||||||
@@ -17914,13 +18124,23 @@ static void ggml_vk_graph_optimize(ggml_backend_t backend, struct ggml_cgraph *
|
|||||||
if (is_empty(graph->nodes[j])) {
|
if (is_empty(graph->nodes[j])) {
|
||||||
continue;
|
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) ||
|
if (match_pattern(topk_moe_early_softmax_norm, j) ||
|
||||||
match_pattern(topk_moe_sigmoid_norm_bias, j) ||
|
match_pattern(topk_moe_sigmoid_norm_bias, j) ||
|
||||||
match_pattern(topk_moe_sqrt_softplus_norm_bias, j) ||
|
match_pattern(topk_moe_sqrt_softplus_norm_bias, j) ||
|
||||||
match_pattern(topk_moe_early_softmax, j) ||
|
match_pattern(topk_moe_early_softmax, j) ||
|
||||||
match_pattern(topk_moe_late_softmax, j) ||
|
match_pattern(topk_moe_late_softmax, j) ||
|
||||||
match_pattern(snake_pattern, j)) {
|
match_pattern(snake_pattern, j) ||
|
||||||
|
in_qsa_pattern(j)) {
|
||||||
continue;
|
continue;
|
||||||
}
|
}
|
||||||
bool ok = true;
|
bool ok = true;
|
||||||
@@ -18723,15 +18943,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])) {
|
if (!ggml_is_contiguous(op) || !ggml_is_contiguous(op->src[0])) {
|
||||||
return false;
|
return false;
|
||||||
}
|
}
|
||||||
// We could potentially support larger, using argsort to sort the
|
// large k falls back to radix-select
|
||||||
// whole thing. Not clear if this is needed.
|
const uint32_t min_pipeline =
|
||||||
uint32_t min_pipeline = (uint32_t)log2f(float(op->ne[0])) + 1;
|
std::max((uint32_t) log2f(float(op->ne[0])) + 1, device->subgroup_size_log2);
|
||||||
if (min_pipeline >= num_topk_pipelines ||
|
if (min_pipeline < num_topk_pipelines && device->pipeline_topk_f32[min_pipeline]) {
|
||||||
!device->pipeline_topk_f32[min_pipeline]) {
|
|
||||||
return false;
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return true;
|
return true;
|
||||||
|
}
|
||||||
|
return device->pipeline_topk_radix_f32 != nullptr;
|
||||||
|
}
|
||||||
case GGML_OP_UPSCALE:
|
case GGML_OP_UPSCALE:
|
||||||
if (op->op_params[0] & GGML_SCALE_FLAG_ANTIALIAS) {
|
if (op->op_params[0] & GGML_SCALE_FLAG_ANTIALIAS) {
|
||||||
if ((op->op_params[0] & 0xFF) != GGML_SCALE_MODE_BILINEAR) {
|
if ((op->op_params[0] & 0xFF) != GGML_SCALE_MODE_BILINEAR) {
|
||||||
|
|||||||
@@ -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);
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -1028,6 +1028,7 @@ void process_shaders() {
|
|||||||
|
|
||||||
string_to_spv("topk_argsort_f32", "topk_argsort.comp", {{"A_TYPE", "float"}});
|
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_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("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("sum_rows_f32", "sum_rows.comp", merge_maps(base_dict, {{"A_TYPE", "float"}, {"D_TYPE", "float"}}));
|
||||||
|
|||||||
@@ -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<ggml_tensor *> 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<int32_t> 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<ggml_fp16_t> 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<int32_t> 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 {
|
enum MoeGatingFunc {
|
||||||
GATING_FUNC_SOFTMAX,
|
GATING_FUNC_SOFTMAX,
|
||||||
GATING_FUNC_SIGMOID,
|
GATING_FUNC_SIGMOID,
|
||||||
@@ -9813,6 +9894,22 @@ static std::vector<std::unique_ptr<test_case>> make_test_cases_eval() {
|
|||||||
test_cases.emplace_back(new test_top_k(GGML_TYPE_F32, {2049, 2, 1, 3}, k));
|
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
|
// exhaustive top_k tests
|
||||||
//for (int i = 1; i < 9999; ++i) {
|
//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));
|
// test_cases.emplace_back(new test_top_k(GGML_TYPE_F32, {i, 2, 1, 3}, rand() % i + 1));
|
||||||
|
|||||||
Reference in New Issue
Block a user