diff --git a/ggml/src/ggml-metal/ggml-metal-fusion.cpp b/ggml/src/ggml-metal/ggml-metal-fusion.cpp index c37364fbf5..605a3e38a1 100644 --- a/ggml/src/ggml-metal/ggml-metal-fusion.cpp +++ b/ggml/src/ggml-metal/ggml-metal-fusion.cpp @@ -5,6 +5,7 @@ #include #include +#include #include #include @@ -269,6 +270,21 @@ static bool ggml_metal_fusion_check_snake( // SOFT_MAX + ARGSORT + GET_ROWS (plus optional norm/scale) for MoE routing. // This is a multi-output elision chain: the fused kernel writes both the selected // expert ids and the gathered/normalized routing weights. +static const ggml_op ops_topk_moe_raw[] = { + GGML_OP_SOFT_MAX, GGML_OP_RESHAPE, GGML_OP_ARGSORT, GGML_OP_VIEW, GGML_OP_GET_ROWS +}; +static const ggml_op ops_topk_moe_scale_raw[] = { + GGML_OP_SOFT_MAX, GGML_OP_RESHAPE, GGML_OP_ARGSORT, GGML_OP_VIEW, GGML_OP_GET_ROWS, GGML_OP_SCALE +}; +static const ggml_op ops_topk_moe_norm_raw[] = { + GGML_OP_SOFT_MAX, GGML_OP_RESHAPE, GGML_OP_ARGSORT, GGML_OP_VIEW, GGML_OP_GET_ROWS, + GGML_OP_RESHAPE, GGML_OP_SUM_ROWS, GGML_OP_CLAMP, GGML_OP_DIV, GGML_OP_RESHAPE +}; +static const ggml_op ops_topk_moe_norm_scale_raw[] = { + GGML_OP_SOFT_MAX, GGML_OP_RESHAPE, GGML_OP_ARGSORT, GGML_OP_VIEW, GGML_OP_GET_ROWS, + GGML_OP_RESHAPE, GGML_OP_SUM_ROWS, GGML_OP_CLAMP, GGML_OP_DIV, GGML_OP_RESHAPE, GGML_OP_SCALE +}; + static bool ggml_metal_fusion_check_topk_moe( const ggml_metal_fusion * fusion, const ggml_tensor * const * nodes, @@ -286,21 +302,6 @@ static bool ggml_metal_fusion_check_topk_moe( // the fusion table operates on the non-empty node sequence; the raw graph also // contains the RESHAPE/VIEW nodes that the fused kernel elides. - static const ggml_op ops_topk_moe_raw[] = { - GGML_OP_SOFT_MAX, GGML_OP_RESHAPE, GGML_OP_ARGSORT, GGML_OP_VIEW, GGML_OP_GET_ROWS - }; - static const ggml_op ops_topk_moe_scale_raw[] = { - GGML_OP_SOFT_MAX, GGML_OP_RESHAPE, GGML_OP_ARGSORT, GGML_OP_VIEW, GGML_OP_GET_ROWS, GGML_OP_SCALE - }; - static const ggml_op ops_topk_moe_norm_raw[] = { - GGML_OP_SOFT_MAX, GGML_OP_RESHAPE, GGML_OP_ARGSORT, GGML_OP_VIEW, GGML_OP_GET_ROWS, - GGML_OP_RESHAPE, GGML_OP_SUM_ROWS, GGML_OP_CLAMP, GGML_OP_DIV, GGML_OP_RESHAPE - }; - static const ggml_op ops_topk_moe_norm_scale_raw[] = { - GGML_OP_SOFT_MAX, GGML_OP_RESHAPE, GGML_OP_ARGSORT, GGML_OP_VIEW, GGML_OP_GET_ROWS, - GGML_OP_RESHAPE, GGML_OP_SUM_ROWS, GGML_OP_CLAMP, GGML_OP_DIV, GGML_OP_RESHAPE, GGML_OP_SCALE - }; - const ggml_op * expected = nullptr; int expected_n = 0; @@ -446,71 +447,6 @@ struct ggml_metal_moe_reduce_match { int node_count; }; -struct ggml_metal_topk_moe_match { - const ggml_tensor * logits; - const ggml_tensor * dst; - int node_count; -}; - -static bool ggml_metal_fusion_match_topk_moe(const ggml_cgraph * gf, int node_idx, ggml_metal_topk_moe_match * match) { - if (match == nullptr || node_idx < 0 || node_idx + 5 > gf->n_nodes) { - return false; - } - - const ggml_tensor * softmax = gf->nodes[node_idx]; - if (softmax->op != GGML_OP_SOFT_MAX || softmax->src[0] == nullptr) { - return false; - } - - static const ggml_op ops_topk_moe_raw[] = { - GGML_OP_SOFT_MAX, GGML_OP_RESHAPE, GGML_OP_ARGSORT, GGML_OP_VIEW, GGML_OP_GET_ROWS - }; - static const ggml_op ops_topk_moe_scale_raw[] = { - GGML_OP_SOFT_MAX, GGML_OP_RESHAPE, GGML_OP_ARGSORT, GGML_OP_VIEW, GGML_OP_GET_ROWS, GGML_OP_SCALE - }; - static const ggml_op ops_topk_moe_norm_raw[] = { - GGML_OP_SOFT_MAX, GGML_OP_RESHAPE, GGML_OP_ARGSORT, GGML_OP_VIEW, GGML_OP_GET_ROWS, - GGML_OP_RESHAPE, GGML_OP_SUM_ROWS, GGML_OP_CLAMP, GGML_OP_DIV, GGML_OP_RESHAPE - }; - static const ggml_op ops_topk_moe_norm_scale_raw[] = { - GGML_OP_SOFT_MAX, GGML_OP_RESHAPE, GGML_OP_ARGSORT, GGML_OP_VIEW, GGML_OP_GET_ROWS, - GGML_OP_RESHAPE, GGML_OP_SUM_ROWS, GGML_OP_CLAMP, GGML_OP_DIV, GGML_OP_RESHAPE, GGML_OP_SCALE - }; - - const ggml_op * expected[] = { - ops_topk_moe_raw, - ops_topk_moe_scale_raw, - ops_topk_moe_norm_raw, - ops_topk_moe_norm_scale_raw, - }; - const int expected_n[] = { 5, 6, 10, 11 }; - - for (int v = 0; v < 4; ++v) { - const int n = expected_n[v]; - if (node_idx + n > gf->n_nodes) { - continue; - } - - bool ok = true; - for (int j = 0; j < n; ++j) { - if (gf->nodes[node_idx + j]->op != expected[v][j]) { - ok = false; - break; - } - } - if (!ok) { - continue; - } - - match->logits = softmax->src[0]; - match->dst = gf->nodes[node_idx + n - 1]; - match->node_count = n; - return true; - } - - return false; -} - static bool ggml_metal_fusion_match_moe_reduce( const ggml_cgraph * gf, int node_idx, ggml_metal_moe_reduce_match * match) { if (match == nullptr || node_idx < 0 || node_idx + 3 > gf->n_nodes) { @@ -683,33 +619,61 @@ static const ggml_op ops_moe_reduce_6[] = { GGML_OP_MUL, GGML_OP_ADD, GGML_OP_AD static const ggml_op ops_moe_reduce_7[] = { GGML_OP_MUL, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD }; static const ggml_op ops_moe_reduce_8[] = { GGML_OP_MUL, 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_moe_reduce_raw_2[] = { + GGML_OP_MUL, GGML_OP_VIEW, GGML_OP_VIEW, GGML_OP_ADD +}; +static const ggml_op ops_moe_reduce_raw_3[] = { + GGML_OP_MUL, GGML_OP_VIEW, GGML_OP_VIEW, GGML_OP_VIEW, GGML_OP_ADD, GGML_OP_ADD +}; +static const ggml_op ops_moe_reduce_raw_4[] = { + GGML_OP_MUL, GGML_OP_VIEW, GGML_OP_VIEW, GGML_OP_VIEW, GGML_OP_VIEW, + GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD +}; +static const ggml_op ops_moe_reduce_raw_5[] = { + GGML_OP_MUL, GGML_OP_VIEW, GGML_OP_VIEW, GGML_OP_VIEW, GGML_OP_VIEW, GGML_OP_VIEW, + GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD +}; +static const ggml_op ops_moe_reduce_raw_6[] = { + GGML_OP_MUL, GGML_OP_VIEW, GGML_OP_VIEW, GGML_OP_VIEW, GGML_OP_VIEW, GGML_OP_VIEW, GGML_OP_VIEW, + GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD +}; +static const ggml_op ops_moe_reduce_raw_7[] = { + GGML_OP_MUL, GGML_OP_VIEW, GGML_OP_VIEW, GGML_OP_VIEW, GGML_OP_VIEW, GGML_OP_VIEW, GGML_OP_VIEW, GGML_OP_VIEW, + GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD +}; +static const ggml_op ops_moe_reduce_raw_8[] = { + GGML_OP_MUL, + GGML_OP_VIEW, GGML_OP_VIEW, GGML_OP_VIEW, GGML_OP_VIEW, GGML_OP_VIEW, GGML_OP_VIEW, GGML_OP_VIEW, GGML_OP_VIEW, + GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD +}; + static const ggml_metal_fusion ggml_metal_fusions[] = { - { GGML_METAL_FUSION_NORM_MUL, ops_norm_mul, 2, false, ggml_metal_fusion_check_norm }, - { GGML_METAL_FUSION_NORM_MUL_ADD, ops_norm_mul_add, 3, false, ggml_metal_fusion_check_norm }, - { GGML_METAL_FUSION_NORM_SCALE, ops_norm_scale, 2, false, ggml_metal_fusion_check_norm }, - { GGML_METAL_FUSION_NORM_MUL, ops_rms_norm_mul, 2, false, ggml_metal_fusion_check_norm }, - { GGML_METAL_FUSION_NORM_MUL_ADD, ops_rms_norm_mul_add, 3, false, ggml_metal_fusion_check_norm }, - { GGML_METAL_FUSION_NORM_SCALE, ops_rms_norm_scale, 2, false, ggml_metal_fusion_check_norm }, - { GGML_METAL_FUSION_ADD_CHAIN, ops_add_2, 2, false, ggml_metal_fusion_check_add_chain }, - { GGML_METAL_FUSION_ADD_CHAIN, ops_add_3, 3, false, ggml_metal_fusion_check_add_chain }, - { GGML_METAL_FUSION_ADD_CHAIN, ops_add_4, 4, false, ggml_metal_fusion_check_add_chain }, - { GGML_METAL_FUSION_ADD_CHAIN, ops_add_5, 5, false, ggml_metal_fusion_check_add_chain }, - { GGML_METAL_FUSION_ADD_CHAIN, ops_add_6, 6, false, ggml_metal_fusion_check_add_chain }, - { GGML_METAL_FUSION_ADD_CHAIN, ops_add_7, 7, false, ggml_metal_fusion_check_add_chain }, - { GGML_METAL_FUSION_SNAKE, ops_snake, 5, false, ggml_metal_fusion_check_snake }, - { GGML_METAL_FUSION_GDN_CACHE, ops_gdn_cache, 2, true, ggml_metal_fusion_check_gdn_cache }, - { GGML_METAL_FUSION_TOPK_MOE, ops_topk_moe, 3, true, ggml_metal_fusion_check_topk_moe }, - { GGML_METAL_FUSION_TOPK_MOE, ops_topk_moe_scale, 4, true, ggml_metal_fusion_check_topk_moe }, - { GGML_METAL_FUSION_TOPK_MOE, ops_topk_moe_norm, 6, true, ggml_metal_fusion_check_topk_moe }, - { GGML_METAL_FUSION_TOPK_MOE, ops_topk_moe_norm_scale, 7, true, ggml_metal_fusion_check_topk_moe }, - { GGML_METAL_FUSION_MOE_REDUCE, ops_moe_reduce_2, 2, true, ggml_metal_fusion_check_moe_reduce }, - { GGML_METAL_FUSION_MOE_REDUCE, ops_moe_reduce_3, 3, true, ggml_metal_fusion_check_moe_reduce }, - { GGML_METAL_FUSION_MOE_REDUCE, ops_moe_reduce_4, 4, true, ggml_metal_fusion_check_moe_reduce }, - { GGML_METAL_FUSION_MOE_REDUCE, ops_moe_reduce_5, 5, true, ggml_metal_fusion_check_moe_reduce }, - { GGML_METAL_FUSION_MOE_REDUCE, ops_moe_reduce_6, 6, true, ggml_metal_fusion_check_moe_reduce }, - { GGML_METAL_FUSION_MOE_REDUCE, ops_moe_reduce_7, 7, true, ggml_metal_fusion_check_moe_reduce }, - { GGML_METAL_FUSION_MOE_REDUCE, ops_moe_reduce_8, 8, true, ggml_metal_fusion_check_moe_reduce }, - { GGML_METAL_FUSION_SSM_CONV_SILU, ops_ssm_conv_silu, 2, false, ggml_metal_fusion_check_ssm_conv_silu }, + { GGML_METAL_FUSION_NORM_MUL, ops_norm_mul, 2, ops_norm_mul, 2, false, ggml_metal_fusion_check_norm }, + { GGML_METAL_FUSION_NORM_MUL_ADD, ops_norm_mul_add, 3, ops_norm_mul_add, 3, false, ggml_metal_fusion_check_norm }, + { GGML_METAL_FUSION_NORM_SCALE, ops_norm_scale, 2, ops_norm_scale, 2, false, ggml_metal_fusion_check_norm }, + { GGML_METAL_FUSION_NORM_MUL, ops_rms_norm_mul, 2, ops_rms_norm_mul, 2, false, ggml_metal_fusion_check_norm }, + { GGML_METAL_FUSION_NORM_MUL_ADD, ops_rms_norm_mul_add, 3, ops_rms_norm_mul_add, 3, false, ggml_metal_fusion_check_norm }, + { GGML_METAL_FUSION_NORM_SCALE, ops_rms_norm_scale, 2, ops_rms_norm_scale, 2, false, ggml_metal_fusion_check_norm }, + { GGML_METAL_FUSION_ADD_CHAIN, ops_add_2, 2, ops_add_2, 2, false, ggml_metal_fusion_check_add_chain }, + { GGML_METAL_FUSION_ADD_CHAIN, ops_add_3, 3, ops_add_3, 3, false, ggml_metal_fusion_check_add_chain }, + { GGML_METAL_FUSION_ADD_CHAIN, ops_add_4, 4, ops_add_4, 4, false, ggml_metal_fusion_check_add_chain }, + { GGML_METAL_FUSION_ADD_CHAIN, ops_add_5, 5, ops_add_5, 5, false, ggml_metal_fusion_check_add_chain }, + { GGML_METAL_FUSION_ADD_CHAIN, ops_add_6, 6, ops_add_6, 6, false, ggml_metal_fusion_check_add_chain }, + { GGML_METAL_FUSION_ADD_CHAIN, ops_add_7, 7, ops_add_7, 7, false, ggml_metal_fusion_check_add_chain }, + { GGML_METAL_FUSION_SNAKE, ops_snake, 5, ops_snake, 5, false, ggml_metal_fusion_check_snake }, + { GGML_METAL_FUSION_GDN_CACHE, ops_gdn_cache, 2, ops_gdn_cache, 2, true, ggml_metal_fusion_check_gdn_cache }, + { GGML_METAL_FUSION_TOPK_MOE, ops_topk_moe, 3, ops_topk_moe_raw, 5, true, ggml_metal_fusion_check_topk_moe }, + { GGML_METAL_FUSION_TOPK_MOE, ops_topk_moe_scale, 4, ops_topk_moe_scale_raw, 6, true, ggml_metal_fusion_check_topk_moe }, + { GGML_METAL_FUSION_TOPK_MOE, ops_topk_moe_norm, 6, ops_topk_moe_norm_raw, 10, true, ggml_metal_fusion_check_topk_moe }, + { GGML_METAL_FUSION_TOPK_MOE, ops_topk_moe_norm_scale, 7, ops_topk_moe_norm_scale_raw, 11, true, ggml_metal_fusion_check_topk_moe }, + { GGML_METAL_FUSION_MOE_REDUCE, ops_moe_reduce_2, 2, ops_moe_reduce_raw_2, 4, true, ggml_metal_fusion_check_moe_reduce }, + { GGML_METAL_FUSION_MOE_REDUCE, ops_moe_reduce_3, 3, ops_moe_reduce_raw_3, 6, true, ggml_metal_fusion_check_moe_reduce }, + { GGML_METAL_FUSION_MOE_REDUCE, ops_moe_reduce_4, 4, ops_moe_reduce_raw_4, 8, true, ggml_metal_fusion_check_moe_reduce }, + { GGML_METAL_FUSION_MOE_REDUCE, ops_moe_reduce_5, 5, ops_moe_reduce_raw_5, 10, true, ggml_metal_fusion_check_moe_reduce }, + { GGML_METAL_FUSION_MOE_REDUCE, ops_moe_reduce_6, 6, ops_moe_reduce_raw_6, 12, true, ggml_metal_fusion_check_moe_reduce }, + { GGML_METAL_FUSION_MOE_REDUCE, ops_moe_reduce_7, 7, ops_moe_reduce_raw_7, 14, true, ggml_metal_fusion_check_moe_reduce }, + { GGML_METAL_FUSION_MOE_REDUCE, ops_moe_reduce_8, 8, ops_moe_reduce_raw_8, 16, true, ggml_metal_fusion_check_moe_reduce }, + { GGML_METAL_FUSION_SSM_CONV_SILU, ops_ssm_conv_silu, 2, ops_ssm_conv_silu, 2, false, ggml_metal_fusion_check_ssm_conv_silu }, }; const ggml_metal_fusion * ggml_metal_fusion_all(int * n) { @@ -718,34 +682,69 @@ const ggml_metal_fusion * ggml_metal_fusion_all(int * n) { return ggml_metal_fusions; } +static bool ggml_metal_fusion_match_raw_pattern( + const ggml_cgraph * gf, int node_idx, const ggml_op * ops, int n_ops) { + if (node_idx < 0 || node_idx + n_ops > gf->n_nodes) { + return false; + } + + for (int i = 0; i < n_ops; ++i) { + if (gf->nodes[node_idx + i]->op != ops[i]) { + return false; + } + } + + return true; +} + +static void ggml_metal_fusion_add_pattern_alloc_deps( + void * user_data, + void (*add_alloc_dep)(void *, ggml_tensor *, ggml_tensor *), + const ggml_cgraph * gf, + const ggml_metal_fusion * fusion, + int node_idx) { + const int last_node = node_idx + fusion->n_raw_ops - 1; + + // keep all external inputs alive until the fused output + std::set seen; + for (int j = 0; j < fusion->n_raw_ops; ++j) { + const ggml_tensor * node = gf->nodes[node_idx + j]; + for (int s = 0; s < GGML_MAX_SRC; ++s) { + const ggml_tensor * src = node->src[s]; + if (src && seen.insert(src).second) { + add_alloc_dep(user_data, const_cast(src), const_cast(gf->nodes[last_node])); + } + } + seen.insert(node); + } +} + void ggml_metal_fusion_add_alloc_deps( void * user_data, void (*add_alloc_dep)(void *, ggml_tensor *, ggml_tensor *), const ggml_cgraph * gf) { - for (int i = 0; i < gf->n_nodes; ++i) { - if (gf->nodes[i]->op == GGML_OP_SOFT_MAX) { - ggml_metal_topk_moe_match match; - if (ggml_metal_fusion_match_topk_moe(gf, i, &match)) { - add_alloc_dep(user_data, const_cast(match.logits), const_cast(match.dst)); + int n_fusions = 0; + const ggml_metal_fusion * all = ggml_metal_fusion_all(&n_fusions); - i += match.node_count - 1; + for (int i = 0; i < gf->n_nodes; ++i) { + const ggml_metal_fusion * best = nullptr; + int best_raw = 0; + + for (int f = 0; f < n_fusions; ++f) { + const ggml_metal_fusion * fusion = &all[f]; + if (fusion->n_raw_ops <= best_raw) { continue; } + if (ggml_metal_fusion_match_raw_pattern(gf, i, fusion->raw_ops, fusion->n_raw_ops)) { + best = fusion; + best_raw = fusion->n_raw_ops; + } } - if (gf->nodes[i]->op != GGML_OP_MUL) { - continue; + if (best) { + ggml_metal_fusion_add_pattern_alloc_deps(user_data, add_alloc_dep, gf, best, i); + i += best_raw - 1; } - - ggml_metal_moe_reduce_match match; - if (!ggml_metal_fusion_match_moe_reduce(gf, i, &match)) { - continue; - } - - add_alloc_dep(user_data, const_cast(match.experts), const_cast(match.dst)); - add_alloc_dep(user_data, const_cast(match.weights), const_cast(match.dst)); - - i += match.node_count - 1; } } diff --git a/ggml/src/ggml-metal/ggml-metal-fusion.h b/ggml/src/ggml-metal/ggml-metal-fusion.h index 7e1284a79f..8df0cd6bfa 100644 --- a/ggml/src/ggml-metal/ggml-metal-fusion.h +++ b/ggml/src/ggml-metal/ggml-metal-fusion.h @@ -44,9 +44,12 @@ typedef enum ggml_metal_fusion_id { struct ggml_metal_fusion { ggml_metal_fusion_id id; - const enum ggml_op * ops; // op sequence (fixed length) + const enum ggml_op * ops; // op sequence (fixed length, non-empty nodes) int n_ops; // number of ops + const enum ggml_op * raw_ops; // full raw op sequence (may include empty RESHAPE/VIEW nodes) + int n_raw_ops; // number of raw ops + // 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)