From 03a667aa304f2a8e02a9a02b2e3fb45d64bcae7f Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Fran=C3=A7ois-Xavier=20Gsell?= Date: Mon, 28 Sep 2026 20:20:14 +0800 Subject: [PATCH] vulkan: fuse qwen4exp's SCALE -> SIGMOID -> SCALE -> hc_post chain (#29520) --- ggml/src/ggml-vulkan/ggml-vulkan-common.h | 2 +- .../ggml-vulkan/ggml-vulkan-push-constants.h | 4 ++ ggml/src/ggml-vulkan/ggml-vulkan-types.h | 10 ++++ ggml/src/ggml-vulkan/ggml-vulkan.cpp | 50 ++++++++++++++++--- .../vulkan-shaders/dsv4_hc_post.comp | 7 ++- tests/test-backend-ops.cpp | 17 +++++-- 6 files changed, 79 insertions(+), 11 deletions(-) diff --git a/ggml/src/ggml-vulkan/ggml-vulkan-common.h b/ggml/src/ggml-vulkan/ggml-vulkan-common.h index 4ae5fea7a8..5f95312a76 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan-common.h +++ b/ggml/src/ggml-vulkan/ggml-vulkan-common.h @@ -97,7 +97,7 @@ vk_pipeline ggml_vk_get_quantize_pipeline(ggml_backend_vk_context * ctx, ggml_ty void ggml_vk_quantize_q8_1(ggml_backend_vk_context * ctx, vk_context& subctx, const vk_subbuffer & in, const vk_subbuffer & out, uint32_t ne); void ggml_vk_dsv4_hc_comb(ggml_backend_vk_context * ctx, vk_context& subctx, const ggml_tensor * mixes, const ggml_tensor * scale, const ggml_tensor * base, ggml_tensor * dst); void ggml_vk_dsv4_hc_pre(ggml_backend_vk_context * ctx, vk_context& subctx, const ggml_tensor * x, const ggml_tensor * weights, ggml_tensor * dst); -void ggml_vk_dsv4_hc_post(ggml_backend_vk_context * ctx, vk_context& subctx, const ggml_tensor * x, const ggml_tensor * residual, const ggml_tensor * post, const ggml_tensor * comb, ggml_tensor * dst); +void ggml_vk_dsv4_hc_post(ggml_backend_vk_context * ctx, vk_context& subctx, const ggml_tensor * x, const ggml_tensor * residual, const ggml_tensor * post, const ggml_tensor * comb, ggml_tensor * dst, const ggml_tensor * gate_scale_in = nullptr); void ggml_vk_mul_mat(ggml_backend_vk_context * ctx, vk_context& subctx, const struct ggml_cgraph * cgraph, int node_idx); bool ggml_vk_use_mul_mat_vec_id(const struct ggml_cgraph * cgraph, int node_idx); void ggml_vk_mul_mat_id(ggml_backend_vk_context * ctx, vk_context& subctx, const struct ggml_cgraph * cgraph, int node_idx); diff --git a/ggml/src/ggml-vulkan/ggml-vulkan-push-constants.h b/ggml/src/ggml-vulkan/ggml-vulkan-push-constants.h index f066d10644..037bcee825 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan-push-constants.h +++ b/ggml/src/ggml-vulkan/ggml-vulkan-push-constants.h @@ -205,6 +205,10 @@ struct vk_op_dsv4_hc_post_push_constants { uint32_t p_offset; uint32_t c_offset; uint32_t d_offset; + + uint32_t gate; + float gate_scale_in; + float gate_scale_out; }; struct vk_op_count_experts_push_constants { diff --git a/ggml/src/ggml-vulkan/ggml-vulkan-types.h b/ggml/src/ggml-vulkan/ggml-vulkan-types.h index 252359bf89..84962be040 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan-types.h +++ b/ggml/src/ggml-vulkan/ggml-vulkan-types.h @@ -553,6 +553,15 @@ static constexpr std::initializer_list rms_norm_view_set_rows_pattern { static constexpr std::initializer_list rope_view_set_rows_pattern { GGML_OP_ROPE, GGML_OP_VIEW, GGML_OP_SET_ROWS }; +// scale_out*sigmoid(scale_in*x) as the hc_post weights (qwen4exp hc_combine) +static constexpr std::initializer_list hc_post_gate_pattern { GGML_OP_SCALE, GGML_OP_UNARY, GGML_OP_SCALE, GGML_OP_DSV4_HC_POST }; + +static constexpr std::initializer_list> hc_post_gate_edges { + { 1, 0, 0 }, // sigmoid->src[0] == scale + { 2, 0, 1 }, // scale->src[0] == sigmoid + { 3, 2, 2 }, // hc_post->src[2] == scale (post) +}; + static constexpr std::initializer_list> topk_moe_early_softmax_norm_edges { { 1, 0, 0 }, // reshape->src[0] == softmax { 2, 0, 0 }, // argsort->src[0] == softmax @@ -1284,6 +1293,7 @@ struct ggml_backend_vk_context { bool fused_topk_moe_scale {}; // QSA indexer gather+add+top_k fused into one radix-select bool fused_topk_qsa {}; + bool fused_hc_post_gate {}; rms_norm_mode fused_rms_norm_mode {RMS_NORM_COUNT}; // for GGML_VK_PERF_LOGGER diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp index 1521e508c9..29077dfc8e 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp +++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp @@ -7167,7 +7167,7 @@ void ggml_vk_dsv4_hc_pre(ggml_backend_vk_context * ctx, vk_context& subctx, cons ggml_vk_dispatch_pipeline(ctx, subctx, pipeline, { x_buf, w_buf, d_buf }, pc, { n_embd, n_tokens, 1 }); } -void ggml_vk_dsv4_hc_post(ggml_backend_vk_context * ctx, vk_context& subctx, const ggml_tensor * x, const ggml_tensor * residual, const ggml_tensor * post, const ggml_tensor * comb, ggml_tensor * dst) { +void ggml_vk_dsv4_hc_post(ggml_backend_vk_context * ctx, vk_context& subctx, const ggml_tensor * x, const ggml_tensor * residual, const ggml_tensor * post, const ggml_tensor * comb, ggml_tensor * dst, const ggml_tensor * gate_scale_in) { VK_LOG_DEBUG("ggml_vk_dsv4_hc_post(" << x << ", " << residual << ", " << post << ", " << comb << ", " << dst << ")"); vk_pipeline pipeline = comb ? ctx->device->pipeline_dsv4_hc_post_f32 : ctx->device->pipeline_dsv4_hc_post_nocomb_f32; @@ -7180,7 +7180,9 @@ void ggml_vk_dsv4_hc_post(ggml_backend_vk_context * ctx, vk_context& subctx, con const vk_subbuffer x_buf = ggml_vk_tensor_subbuffer(ctx, x, true); const vk_subbuffer r_buf = ggml_vk_tensor_subbuffer(ctx, residual, true); - const vk_subbuffer p_buf = ggml_vk_tensor_subbuffer(ctx, post, true); + // with a fused gate, post is scale(sigmoid(scale(p_src))) and the shader applies it to p_src + const ggml_tensor * p_src = gate_scale_in ? gate_scale_in->src[0] : post; + const vk_subbuffer p_buf = ggml_vk_tensor_subbuffer(ctx, p_src, true); const vk_subbuffer c_buf = comb ? ggml_vk_tensor_subbuffer(ctx, comb, true) : x_buf; const vk_subbuffer d_buf = ggml_vk_tensor_subbuffer(ctx, dst, true); @@ -7188,12 +7190,15 @@ void ggml_vk_dsv4_hc_post(ggml_backend_vk_context * ctx, vk_context& subctx, con n_embd, n_tokens, ggml_vk_nb_elem(x, 0), ggml_vk_nb_elem(x, 1), ggml_vk_nb_elem(residual, 0), ggml_vk_nb_elem(residual, 1), ggml_vk_nb_elem(residual, 2), - ggml_vk_nb_elem(post, 0), ggml_vk_nb_elem(post, 1), + ggml_vk_nb_elem(p_src, 0), ggml_vk_nb_elem(p_src, 1), comb ? ggml_vk_nb_elem(comb, 0) : 0, comb ? ggml_vk_nb_elem(comb, 1) : 0, comb ? ggml_vk_nb_elem(comb, 2) : 0, ggml_vk_nb_elem(dst, 0), ggml_vk_nb_elem(dst, 1), ggml_vk_nb_elem(dst, 2), 0, 0, 0, 0, 0, + gate_scale_in ? 1u : 0u, + gate_scale_in ? ggml_get_op_params_f32(gate_scale_in, 0) : 1.0f, + gate_scale_in ? ggml_get_op_params_f32(post, 0) : 1.0f, }; - init_pushconst_tensor_offsets(ctx, pc, x, residual, post, comb, dst); + init_pushconst_tensor_offsets(ctx, pc, x, residual, p_src, comb, dst); ggml_vk_dispatch_pipeline(ctx, subctx, pipeline, { x_buf, r_buf, p_buf, c_buf, d_buf }, pc, { n_embd, n_tokens, 1 }); } @@ -12356,7 +12361,12 @@ bool ggml_vk_build_graph(ggml_backend_vk_context * ctx, ggml_cgraph * cgraph, in break; case GGML_OP_SCALE: - ggml_vk_scale(ctx, compute_ctx, src0, node); + if (ctx->fused_hc_post_gate) { + ggml_tensor * hc_post = cgraph->nodes[node_idx + ctx->num_additional_fused_ops]; + ggml_vk_dsv4_hc_post(ctx, compute_ctx, hc_post->src[0], hc_post->src[1], hc_post->src[2], hc_post->src[3], hc_post, node); + } else { + ggml_vk_scale(ctx, compute_ctx, src0, node); + } break; case GGML_OP_SQR: @@ -13431,6 +13441,19 @@ static bool ggml_vk_can_fuse_unary_mul_pair(const struct ggml_cgraph * cgraph, i ggml_vk_can_fuse_unary_mul(cgraph, node_idx, node_idx + 1); } +static bool ggml_vk_can_fuse_hc_post_gate(const struct ggml_cgraph * cgraph, int node_idx) { + const ggml_tensor * scale_in = cgraph->nodes[node_idx]; + const ggml_tensor * sigmoid = cgraph->nodes[node_idx + 1]; + const ggml_tensor * scale_out = cgraph->nodes[node_idx + 2]; + + // the shader folds scale -> sigmoid -> scale; a bias on either scale is not handled + return ggml_get_unary_op(sigmoid) == GGML_UNARY_OP_SIGMOID && + ggml_get_op_params_f32(scale_in, 1) == 0.0f && + ggml_get_op_params_f32(scale_out, 1) == 0.0f && + scale_in->src[0]->type == GGML_TYPE_F32 && + ggml_are_same_shape(scale_in->src[0], scale_out); +} + bool ggml_vk_can_fuse(const ggml_backend_vk_context * ctx, const struct ggml_cgraph * cgraph, int node_idx, std::initializer_list ops) { if (ops.size() == 2 && ops.begin()[0] == GGML_OP_UNARY && ops.begin()[1] == GGML_OP_MUL) { return ggml_vk_can_fuse_unary_mul_pair(cgraph, node_idx); @@ -14278,6 +14301,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; + ctx->fused_hc_post_gate = false; ctx->fused_rms_norm_mode = RMS_NORM_COUNT; const char *fusion_string {}; if (!ctx->device->disable_fusion) { @@ -14334,6 +14358,13 @@ static ggml_status ggml_backend_vk_graph_compute(ggml_backend_t backend, ggml_cg op_srcs_fused_elementwise[0] = false; op_srcs_fused_elementwise[1] = true; op_srcs_fused_elementwise[2] = true; + } else if (ggml_can_fuse_subgraph(cgraph, i, hc_post_gate_pattern, { i + 3 }) && + ggml_check_edges(cgraph, i, hc_post_gate_edges) && + ggml_vk_can_fuse_hc_post_gate(cgraph, i)) { + ctx->num_additional_fused_ops = hc_post_gate_pattern.size() - 1; + ctx->fused_hc_post_gate = true; + fusion_string = "HC_POST_GATE"; + std::fill_n(op_srcs_fused_elementwise, ctx->num_additional_fused_ops + 1, false); } else if (ggml_vk_can_fuse(ctx, cgraph, i, rms_norm_mul_add_mul_pattern)) { ctx->num_additional_fused_ops = 3; ctx->fused_rms_norm_mode = RMS_NORM_MUL_ADD_MUL; @@ -14512,6 +14543,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; + ctx->fused_hc_post_gate = false; ctx->fused_rms_norm_mode = RMS_NORM_COUNT; fusion_string = nullptr; } @@ -14768,6 +14800,11 @@ void ggml_vk_graph_optimize(ggml_backend_t backend, struct ggml_cgraph * graph, if (keep_pattern(rope_view_set_rows_pattern)) { continue; } + if (match_pattern(hc_post_gate_pattern, first_unused)) { + add_pattern_alloc_deps(hc_post_gate_pattern, first_unused + (int) hc_post_gate_pattern.size() - 1); + keep_pattern(hc_post_gate_pattern); + continue; + } // First, grab the next unused node. current_set.push_back(first_unused); @@ -14807,7 +14844,8 @@ void ggml_vk_graph_optimize(ggml_backend_t backend, struct ggml_cgraph * graph, match_pattern(rms_norm_mul_add_pattern, j) || match_pattern(rms_norm_mul_rope_view_set_rows_pattern, j) || match_pattern(rms_norm_view_set_rows_pattern, j) || - match_pattern(rope_view_set_rows_pattern, j)) { + match_pattern(rope_view_set_rows_pattern, j) || + match_pattern(hc_post_gate_pattern, j)) { continue; } bool ok = true; diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/dsv4_hc_post.comp b/ggml/src/ggml-vulkan/vulkan-shaders/dsv4_hc_post.comp index e521fd9d45..b80c077259 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/dsv4_hc_post.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/dsv4_hc_post.comp @@ -33,6 +33,10 @@ layout(push_constant) uniform parameter uint p_offset; uint c_offset; uint d_offset; + + uint gate; // post = gate_scale_out*sigmoid(gate_scale_in*p) + float gate_scale_in; + float gate_scale_out; }; layout(binding = 0, std430) readonly buffer X { float data_x[]; }; @@ -51,7 +55,8 @@ void main() { const uint it = gl_WorkGroupID.y; if (tid < hc) { - post_s[tid] = data_p[p_offset + tid * nbp0 + it * nbp1]; + const float p = data_p[p_offset + tid * nbp0 + it * nbp1]; + post_s[tid] = gate != 0 ? (1.0f / (1.0f + exp(-(p * gate_scale_in)))) * gate_scale_out : p; } if (HAS_COMB == 1 && tid < hc * hc) { const uint idst = tid & 3; diff --git a/tests/test-backend-ops.cpp b/tests/test-backend-ops.cpp index ff5a832957..e11fb751a9 100644 --- a/tests/test-backend-ops.cpp +++ b/tests/test-backend-ops.cpp @@ -4342,6 +4342,7 @@ struct test_dsv4_hc_post : public test_dsv4_hc { const int64_t n_embd; const int64_t n_tokens; const bool identity; + const bool gated; std::string op_desc(ggml_tensor * t) override { GGML_UNUSED(t); @@ -4349,11 +4350,14 @@ struct test_dsv4_hc_post : public test_dsv4_hc { } std::string vars() override { - return VARS_TO_STR3(n_embd, n_tokens, identity); + return VARS_TO_STR4(n_embd, n_tokens, identity, gated); } - test_dsv4_hc_post(int64_t n_embd = 31, int64_t n_tokens = 17, bool identity = false) - : n_embd(n_embd), n_tokens(n_tokens), identity(identity) {} + // gated: post = 2*sigmoid(post/hc), as qwen4exp builds it, so backends can fuse the chain + bool run_whole_graph() override { return gated; } + + test_dsv4_hc_post(int64_t n_embd = 31, int64_t n_tokens = 17, bool identity = false, bool gated = false) + : n_embd(n_embd), n_tokens(n_tokens), identity(identity), gated(gated) {} ggml_tensor * build_graph(ggml_context * ctx) override { ggml_tensor * x = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, n_embd, n_tokens); @@ -4365,6 +4369,10 @@ struct test_dsv4_hc_post : public test_dsv4_hc { ggml_tensor * post = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, hc, n_tokens); ggml_set_name(post, "post"); + if (gated) { + post = ggml_scale(ctx, ggml_sigmoid(ctx, ggml_scale(ctx, post, 1.0f / (float) hc)), 2.0f); + } + ggml_tensor * comb = nullptr; if (!identity) { comb = ggml_new_tensor_3d(ctx, GGML_TYPE_F32, hc, hc, n_tokens); @@ -9198,6 +9206,9 @@ static std::vector> make_test_cases_eval() { test_cases.emplace_back(new test_dsv4_hc_post(4096, 21)); test_cases.emplace_back(new test_dsv4_hc_post(31, 17, true)); test_cases.emplace_back(new test_dsv4_hc_post(4096, 21, true)); + test_cases.emplace_back(new test_dsv4_hc_post(31, 17, true, true)); + test_cases.emplace_back(new test_dsv4_hc_post(2560, 21, true, true)); + test_cases.emplace_back(new test_dsv4_hc_post(31, 17, false, true)); // glu ops for (ggml_type type : {GGML_TYPE_F16, GGML_TYPE_F32}) {