From a8657d50194ac7564feea715d701c2d4cdc2f3fb Mon Sep 17 00:00:00 2001 From: Georgi Gerganov Date: Tue, 1 Sep 2026 14:41:18 +0300 Subject: [PATCH] metal : tidy fusion pattern checks and table - const-correct ggml_metal_fuse_outputs buffer - annotate unused check-callback parameters - drop a redundant size_t cast - align the ops/table initializers and add blank-line separation Assisted-by: pi:llama.cpp/DeepSeek-V4-Flash-0731 --- ggml/src/ggml-metal/ggml-metal-fuse.cpp | 129 +++++++++++++----------- ggml/src/ggml-metal/ggml-metal-fuse.h | 28 ++--- 2 files changed, 85 insertions(+), 72 deletions(-) diff --git a/ggml/src/ggml-metal/ggml-metal-fuse.cpp b/ggml/src/ggml-metal/ggml-metal-fuse.cpp index a89e375436..93fcb50f95 100644 --- a/ggml/src/ggml-metal/ggml-metal-fuse.cpp +++ b/ggml/src/ggml-metal/ggml-metal-fuse.cpp @@ -8,7 +8,7 @@ // ---- helpers ------------------------------------------------------------- // the pattern outputs (absolute graph node indices); the default is the last node -static const int * ggml_metal_fuse_outputs(const struct ggml_metal_fuse * fuse, int * buf) { +static const int * ggml_metal_fuse_outputs(const ggml_metal_fuse * fuse, const int * buf) { if (fuse->outputs) { return fuse->outputs; } @@ -17,7 +17,7 @@ static const int * ggml_metal_fuse_outputs(const struct ggml_metal_fuse * fuse, } // true if two tensors live in the same Metal buffer -static bool ggml_metal_fuse_same_buffer(const struct ggml_tensor * a, const struct ggml_tensor * b) { +static bool ggml_metal_fuse_same_buffer(const ggml_tensor * a, const ggml_tensor * b) { if (!a || !b) { return false; } @@ -35,9 +35,11 @@ static bool ggml_metal_fuse_same_buffer(const struct ggml_tensor * a, const stru // NORM/RMS_NORM + MUL + ADD: the weight/bias of each fused step must match the norm input // width, be contiguous rows, and the fused outputs must stay F32 -static bool ggml_metal_fuse_check_norm(const struct ggml_tensor * const * nodes, - const struct ggml_metal_fuse * fuse, - enum ggml_metal_fuse_mode mode) { +static bool ggml_metal_fuse_check_norm(const ggml_tensor * const * nodes, + const ggml_metal_fuse * fuse, + ggml_metal_fuse_mode mode) { + GGML_UNUSED(mode); + GGML_ASSERT(fuse->n_ops >= 2); for (int j = 1; j < fuse->n_ops; j++) { @@ -65,9 +67,9 @@ static bool ggml_metal_fuse_check_norm(const struct ggml_tensor * const * nodes, // ADD x N: each ADD reads the previous ADD as src0, and all addends must share layout // (and, in FULL mode, live in the same Metal buffer) -static bool ggml_metal_fuse_check_add_chain(const struct ggml_tensor * const * nodes, - const struct ggml_metal_fuse * fuse, - enum ggml_metal_fuse_mode mode) { +static bool ggml_metal_fuse_check_add_chain(const ggml_tensor * const * nodes, + const ggml_metal_fuse * fuse, + ggml_metal_fuse_mode mode) { GGML_ASSERT(fuse->n_ops >= 2); for (int j = 1; j < fuse->n_ops; j++) { @@ -94,11 +96,13 @@ static bool ggml_metal_fuse_check_add_chain(const struct ggml_tensor * const * n // mirrors ggml_metal_op_can_fuse_gdn_cache (PR #25788). the gdn output has other consumers (the // attn scores view), so unlike the other patterns this is not an elision chain: the structural // checks live entirely in this callback (unsafe = true). -static bool ggml_metal_fuse_check_gdn_cache(const struct ggml_tensor * const * nodes, - const struct ggml_metal_fuse * fuse, - enum ggml_metal_fuse_mode mode) { - const struct ggml_tensor * gdn = nodes[0]; - const struct ggml_tensor * cpy = nodes[1]; +static bool ggml_metal_fuse_check_gdn_cache(const ggml_tensor * const * nodes, + const ggml_metal_fuse * fuse, + ggml_metal_fuse_mode mode) { + GGML_UNUSED(fuse); + + const ggml_tensor * gdn = nodes[0]; + const ggml_tensor * cpy = nodes[1]; // the kernel skips the snapshot tail, so the gdn output must not be a graph output if (gdn->type != GGML_TYPE_F32 || (gdn->flags & GGML_TENSOR_FLAG_OUTPUT)) { @@ -119,8 +123,8 @@ static bool ggml_metal_fuse_check_gdn_cache(const struct ggml_tensor * const * n const int64_t D = S_v * S_v * H; const int64_t n_written = std::min(n_tokens, K); - const struct ggml_tensor * src = cpy->src[0]; // gdn snapshot tail view - const struct ggml_tensor * dst = cpy->src[1]; // cache view + const ggml_tensor * src = cpy->src[0]; // gdn snapshot tail view + const ggml_tensor * dst = cpy->src[1]; // cache view // src must be this gdn's snapshot tail (contiguous, at the tail offset) if (src->op != GGML_OP_VIEW || src->view_src != gdn || @@ -132,7 +136,7 @@ static bool ggml_metal_fuse_check_gdn_cache(const struct ggml_tensor * const * n if (dst->type != GGML_TYPE_F32 || !std::equal(expected_ne, expected_ne + GGML_MAX_DIMS, dst->ne) || dst->nb[0] != ggml_type_size(GGML_TYPE_F32) || - dst->nb[1] != (size_t) ggml_row_size(GGML_TYPE_F32, D)) { + dst->nb[1] != ggml_row_size(GGML_TYPE_F32, D)) { return false; } @@ -147,24 +151,27 @@ static bool ggml_metal_fuse_check_gdn_cache(const struct ggml_tensor * const * n } // MUL + SIN + SQR + MUL + ADD (snake activation) -static bool ggml_metal_fuse_check_snake(const struct ggml_tensor * const * nodes, - const struct ggml_metal_fuse * fuse, - enum ggml_metal_fuse_mode mode) { - const struct ggml_tensor * mul0 = nodes[0]; - const struct ggml_tensor * sin_node = nodes[1]; - const struct ggml_tensor * sqr = nodes[2]; - const struct ggml_tensor * mul1 = nodes[3]; - const struct ggml_tensor * add = nodes[4]; +static bool ggml_metal_fuse_check_snake(const ggml_tensor * const * nodes, + const ggml_metal_fuse * fuse, + ggml_metal_fuse_mode mode) { + GGML_UNUSED(fuse); + GGML_UNUSED(mode); + + const ggml_tensor * mul0 = nodes[0]; + const ggml_tensor * sin_node = nodes[1]; + const ggml_tensor * sqr = nodes[2]; + const ggml_tensor * mul1 = nodes[3]; + const ggml_tensor * add = nodes[4]; // x carries the full activation shape, a is the broadcast operand - const struct ggml_tensor * x = ggml_are_same_shape(mul0, mul0->src[0]) ? mul0->src[0] : mul0->src[1]; - const struct ggml_tensor * a = (x == mul0->src[0]) ? mul0->src[1] : mul0->src[0]; + const ggml_tensor * x = ggml_are_same_shape(mul0, mul0->src[0]) ? mul0->src[0] : mul0->src[1]; + const ggml_tensor * a = (x == mul0->src[0]) ? mul0->src[1] : mul0->src[0]; // mul1 reads sqr and inv_b in either operand order - const struct ggml_tensor * inv_b = (mul1->src[0] == sqr) ? mul1->src[1] : mul1->src[0]; + const ggml_tensor * inv_b = (mul1->src[0] == sqr) ? mul1->src[1] : mul1->src[0]; // closure check: the trailing add reads the same x as the leading mul - const struct ggml_tensor * x_in_add = (add->src[0] == mul1) ? add->src[1] : add->src[0]; + const ggml_tensor * x_in_add = (add->src[0] == mul1) ? add->src[1] : add->src[0]; // x is in the supported whitelist and every chain intermediate shares x's type. // a and inv_b bind as device const float * in the kernel, so they stay F32. @@ -174,6 +181,7 @@ static bool ggml_metal_fuse_check_snake(const struct ggml_tensor * const * nodes (mul0->type == x->type) && (sin_node->type == x->type) && (sqr->type == x->type) && (mul1->type == x->type) && (add->type == x->type); + // a / inv_b collapse to [1, C, 1, 1], x and add stay 2D const bool shape_ok = ggml_are_same_shape(a, inv_b) && a->ne[0] == 1 && a->ne[1] == x->ne[1]; const bool dim_ok = @@ -181,6 +189,7 @@ static bool ggml_metal_fuse_check_snake(const struct ggml_tensor * const * nodes (add->ne[2] == 1) && (add->ne[3] == 1) && (a->ne[2] == 1) && (a->ne[3] == 1) && (inv_b->ne[2] == 1) && (inv_b->ne[3] == 1); + // kernel reads x[idx] and a[c] / inv_b[c] linearly, so every operand is contiguous const bool contig_ok = ggml_is_contiguous(x) && ggml_is_contiguous(add) && @@ -191,35 +200,37 @@ static bool ggml_metal_fuse_check_snake(const struct ggml_tensor * const * nodes // ---- patterns ------------------------------------------------------------ -static const enum ggml_op ops_norm_mul[] = { GGML_OP_NORM, GGML_OP_MUL }; -static const enum ggml_op ops_norm_mul_add[] = { GGML_OP_NORM, GGML_OP_MUL, GGML_OP_ADD }; -static const enum ggml_op ops_rms_norm_mul[] = { GGML_OP_RMS_NORM, GGML_OP_MUL }; -static const enum ggml_op ops_rms_norm_mul_add[] = { GGML_OP_RMS_NORM, GGML_OP_MUL, GGML_OP_ADD }; -static const enum ggml_op ops_add_2[] = { GGML_OP_ADD, GGML_OP_ADD }; -static const enum ggml_op ops_add_3[] = { GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD }; -static const enum ggml_op ops_add_4[] = { GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD }; -static const enum ggml_op ops_add_5[] = { GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD }; -static const enum ggml_op ops_add_6[] = { GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD }; -static const enum ggml_op ops_add_7[] = { GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD }; -static const enum ggml_op ops_snake[] = { GGML_OP_MUL, GGML_OP_SIN, GGML_OP_SQR, GGML_OP_MUL, GGML_OP_ADD }; -static const enum ggml_op ops_gdn_cache[] = { GGML_OP_GATED_DELTA_NET, GGML_OP_CPY }; +static const ggml_op ops_norm_mul[] = { GGML_OP_NORM, GGML_OP_MUL }; +static const ggml_op ops_norm_mul_add[] = { GGML_OP_NORM, GGML_OP_MUL, GGML_OP_ADD }; +static const ggml_op ops_rms_norm_mul[] = { GGML_OP_RMS_NORM, GGML_OP_MUL }; +static const ggml_op ops_rms_norm_mul_add[] = { GGML_OP_RMS_NORM, GGML_OP_MUL, GGML_OP_ADD }; -static const struct ggml_metal_fuse ggml_metal_fuses[] = { +static const ggml_op ops_add_2[] = { GGML_OP_ADD, GGML_OP_ADD }; +static const ggml_op ops_add_3[] = { GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD }; +static const ggml_op ops_add_4[] = { GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD }; +static const ggml_op ops_add_5[] = { GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD }; +static const ggml_op ops_add_6[] = { GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD }; +static const ggml_op ops_add_7[] = { GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD }; +static const ggml_op ops_snake[] = { GGML_OP_MUL, GGML_OP_SIN, GGML_OP_SQR, GGML_OP_MUL, GGML_OP_ADD }; + +static const ggml_op ops_gdn_cache[] = { GGML_OP_GATED_DELTA_NET, GGML_OP_CPY }; + +static const ggml_metal_fuse ggml_metal_fuses[] = { { GGML_METAL_FUSE_NORM_MUL, ops_norm_mul, 2, nullptr, 0, false, ggml_metal_fuse_check_norm }, { GGML_METAL_FUSE_NORM_MUL_ADD, ops_norm_mul_add, 3, nullptr, 0, false, ggml_metal_fuse_check_norm }, { GGML_METAL_FUSE_NORM_MUL, ops_rms_norm_mul, 2, nullptr, 0, false, ggml_metal_fuse_check_norm }, { GGML_METAL_FUSE_NORM_MUL_ADD, ops_rms_norm_mul_add, 3, nullptr, 0, false, ggml_metal_fuse_check_norm }, - { GGML_METAL_FUSE_ADD_CHAIN, ops_add_2, 2, nullptr, 0, false, ggml_metal_fuse_check_add_chain }, - { GGML_METAL_FUSE_ADD_CHAIN, ops_add_3, 3, nullptr, 0, false, ggml_metal_fuse_check_add_chain }, - { GGML_METAL_FUSE_ADD_CHAIN, ops_add_4, 4, nullptr, 0, false, ggml_metal_fuse_check_add_chain }, - { GGML_METAL_FUSE_ADD_CHAIN, ops_add_5, 5, nullptr, 0, false, ggml_metal_fuse_check_add_chain }, - { GGML_METAL_FUSE_ADD_CHAIN, ops_add_6, 6, nullptr, 0, false, ggml_metal_fuse_check_add_chain }, - { GGML_METAL_FUSE_ADD_CHAIN, ops_add_7, 7, nullptr, 0, false, ggml_metal_fuse_check_add_chain }, - { GGML_METAL_FUSE_SNAKE, ops_snake, 5, nullptr, 0, false, ggml_metal_fuse_check_snake }, - { GGML_METAL_FUSE_GDN_CACHE, ops_gdn_cache, 2, nullptr, 0, true, ggml_metal_fuse_check_gdn_cache }, + { GGML_METAL_FUSE_ADD_CHAIN, ops_add_2, 2, nullptr, 0, false, ggml_metal_fuse_check_add_chain }, + { GGML_METAL_FUSE_ADD_CHAIN, ops_add_3, 3, nullptr, 0, false, ggml_metal_fuse_check_add_chain }, + { GGML_METAL_FUSE_ADD_CHAIN, ops_add_4, 4, nullptr, 0, false, ggml_metal_fuse_check_add_chain }, + { GGML_METAL_FUSE_ADD_CHAIN, ops_add_5, 5, nullptr, 0, false, ggml_metal_fuse_check_add_chain }, + { GGML_METAL_FUSE_ADD_CHAIN, ops_add_6, 6, nullptr, 0, false, ggml_metal_fuse_check_add_chain }, + { GGML_METAL_FUSE_ADD_CHAIN, ops_add_7, 7, nullptr, 0, false, ggml_metal_fuse_check_add_chain }, + { GGML_METAL_FUSE_SNAKE, ops_snake, 5, nullptr, 0, false, ggml_metal_fuse_check_snake }, + { GGML_METAL_FUSE_GDN_CACHE, ops_gdn_cache, 2, nullptr, 0, true, ggml_metal_fuse_check_gdn_cache }, }; -const struct ggml_metal_fuse * ggml_metal_fuse_all(int * n) { +const ggml_metal_fuse * ggml_metal_fuse_all(int * n) { *n = (int) sizeof(ggml_metal_fuses) / sizeof(ggml_metal_fuses[0]); return ggml_metal_fuses; @@ -229,21 +240,21 @@ const struct ggml_metal_fuse * ggml_metal_fuse_all(int * n) { // find the longest pattern matching the node sequence starting at idx // (idx is a position in node_idxs, which maps to graph node indices) -const struct ggml_metal_fuse * ggml_metal_fuse_next( - const struct ggml_cgraph * gf, +const ggml_metal_fuse * ggml_metal_fuse_next( + const ggml_cgraph * gf, const int * node_idxs, int n_idxs, int idx, - enum ggml_metal_fuse_mode mode, + ggml_metal_fuse_mode mode, int * n_out) { int n = 0; - const struct ggml_metal_fuse * all = ggml_metal_fuse_all(&n); + const ggml_metal_fuse * all = ggml_metal_fuse_all(&n); - const struct ggml_metal_fuse * res = nullptr; + const ggml_metal_fuse * res = nullptr; int best = 1; for (int i = 0; i < n; i++) { - const struct ggml_metal_fuse * fuse = &all[i]; + const ggml_metal_fuse * fuse = &all[i]; // only look for a longer match than the current best if (fuse->n_ops <= best) { @@ -253,7 +264,7 @@ const struct ggml_metal_fuse * ggml_metal_fuse_next( continue; } - const struct ggml_tensor * nodes[GGML_METAL_FUSE_MAX]; + const ggml_tensor * nodes[GGML_METAL_FUSE_MAX]; // the op sequence must match exactly bool ok = true; @@ -315,7 +326,7 @@ const struct ggml_metal_fuse * ggml_metal_fuse_next( // could be fused, chaining patterns back-to-back. matching runs on the same filtered (view // transparent) node sequence that the compute phase uses, so the returned count is the raw index // span from idx to the last matched node (intermediate views are packed along). -int ggml_metal_fuse_max(const struct ggml_cgraph * gf, int idx) { +int ggml_metal_fuse_max(const ggml_cgraph * gf, int idx) { // an empty/view node cannot start a pattern - pack it alone if (ggml_op_is_empty(gf->nodes[idx]->op) || ggml_is_empty(gf->nodes[idx])) { return 1; @@ -338,7 +349,7 @@ int ggml_metal_fuse_max(const struct ggml_cgraph * gf, int idx) { while (i_f < n_idxs && total < GGML_METAL_FUSE_MAX) { int len = 1; - const struct ggml_metal_fuse * fuse = ggml_metal_fuse_next(gf, idxs, n_idxs, i_f, GGML_METAL_FUSE_STRUCTURAL, &len); + const ggml_metal_fuse * fuse = ggml_metal_fuse_next(gf, idxs, n_idxs, i_f, GGML_METAL_FUSE_STRUCTURAL, &len); if (!fuse || total + len > GGML_METAL_FUSE_MAX) { break; } diff --git a/ggml/src/ggml-metal/ggml-metal-fuse.h b/ggml/src/ggml-metal/ggml-metal-fuse.h index 7316967cbb..edf6d0f87d 100644 --- a/ggml/src/ggml-metal/ggml-metal-fuse.h +++ b/ggml/src/ggml-metal/ggml-metal-fuse.h @@ -38,38 +38,40 @@ enum ggml_metal_fuse_id { }; struct ggml_metal_fuse { - enum ggml_metal_fuse_id id; - const enum ggml_op * ops; // op sequence (fixed length) - int n_ops; // number of ops - const int * outputs; // output node indices (absolute graph indices; nullptr => the last node) - int n_outputs;// number of outputs (0 => default last node) + ggml_metal_fuse_id id; + const ggml_op * ops; // op sequence (fixed length) + int n_ops; // number of ops + const int * outputs; // output node indices (absolute graph indices; nullptr => the last node) + int n_outputs; // number of outputs (0 => default last node) + // if unsafe: the generic chain/shape + ggml_can_fuse_subgraph checks are skipped and the // check callback below is the sole validator (used for patterns that are not elision chains, // e.g. the gdn + cache-cpy write-through fusion) bool unsafe; + // extra backend constraints on top of ggml_can_fuse_subgraph // nodes[j] is the j-th node of the pattern - bool (*check)(const struct ggml_tensor * const * nodes, - const struct ggml_metal_fuse * fuse, - enum ggml_metal_fuse_mode mode); + bool (*check)(const ggml_tensor * const * nodes, + const ggml_metal_fuse * fuse, + ggml_metal_fuse_mode mode); }; // the single table of all fusions supported by the Metal backend -const struct ggml_metal_fuse * ggml_metal_fuse_all(int * n); +const ggml_metal_fuse * ggml_metal_fuse_all(int * n); // compute phase: longest fusion starting at idx (a position in node_idxs) that matches in `mode`. // returns the matching pattern (nullptr if no fusion) and sets *n_out to the number of nodes consumed. -const struct ggml_metal_fuse * ggml_metal_fuse_next( - const struct ggml_cgraph * gf, +const ggml_metal_fuse * ggml_metal_fuse_next( + const ggml_cgraph * gf, const int * node_idxs, int n_idxs, int idx, - enum ggml_metal_fuse_mode mode, + ggml_metal_fuse_mode mode, int * n_out); // optimize phase: maximum number of nodes starting at idx (a raw sequential graph index) that // could be fused, chaining patterns back-to-back. returns at least 1. -int ggml_metal_fuse_max(const struct ggml_cgraph * gf, int idx); +int ggml_metal_fuse_max(const ggml_cgraph * gf, int idx); #ifdef __cplusplus }