Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion ggml/src/ggml-vulkan/ggml-vulkan-common.h
Original file line number Diff line number Diff line change
Expand Up @@ -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);
Expand Down
4 changes: 4 additions & 0 deletions ggml/src/ggml-vulkan/ggml-vulkan-push-constants.h
Original file line number Diff line number Diff line change
Expand Up @@ -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 {
Expand Down
10 changes: 10 additions & 0 deletions ggml/src/ggml-vulkan/ggml-vulkan-types.h
Original file line number Diff line number Diff line change
Expand Up @@ -553,6 +553,15 @@ static constexpr std::initializer_list<ggml_op> rms_norm_view_set_rows_pattern {

static constexpr std::initializer_list<ggml_op> 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<ggml_op> hc_post_gate_pattern { GGML_OP_SCALE, GGML_OP_UNARY, GGML_OP_SCALE, GGML_OP_DSV4_HC_POST };

static constexpr std::initializer_list<std::array<int, 3>> 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<std::array<int, 3>> topk_moe_early_softmax_norm_edges {
{ 1, 0, 0 }, // reshape->src[0] == softmax
{ 2, 0, 0 }, // argsort->src[0] == softmax
Expand Down Expand Up @@ -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
Expand Down
50 changes: 44 additions & 6 deletions ggml/src/ggml-vulkan/ggml-vulkan.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -7157,7 +7157,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;
Expand All @@ -7170,20 +7170,25 @@ 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);

vk_op_dsv4_hc_post_push_constants pc = {
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 });
}
Expand Down Expand Up @@ -12340,7 +12345,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:
Expand Down Expand Up @@ -13415,6 +13425,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<enum ggml_op> 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);
Expand Down Expand Up @@ -14262,6 +14285,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) {
Expand Down Expand Up @@ -14318,6 +14342,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;
Expand Down Expand Up @@ -14496,6 +14527,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;
}
Expand Down Expand Up @@ -14752,6 +14784,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);
Expand Down Expand Up @@ -14791,7 +14828,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;
Expand Down
7 changes: 6 additions & 1 deletion ggml/src/ggml-vulkan/vulkan-shaders/dsv4_hc_post.comp
Original file line number Diff line number Diff line change
Expand Up @@ -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[]; };
Expand All @@ -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;
Expand Down
17 changes: 14 additions & 3 deletions tests/test-backend-ops.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -4342,18 +4342,22 @@ 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);
return "DSV4_HC_POST";
}

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);
Expand All @@ -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);
Expand Down Expand Up @@ -9183,6 +9191,9 @@ static std::vector<std::unique_ptr<test_case>> 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}) {
Expand Down