metal : refactor alloc deps to pattern-driven approach

Assisted-by: pi:llama.cpp/DeepSeek-V4-Flash-Vision-Exp
This commit is contained in:
Georgi Gerganov
2026-09-15 23:46:21 +03:00
parent e6f001a696
commit 8b9a664a54
2 changed files with 127 additions and 125 deletions
+123 -124
View File
@@ -5,6 +5,7 @@
#include <algorithm>
#include <cstring>
#include <set>
#include <string>
#include <vector>
@@ -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<const ggml_tensor *> 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<ggml_tensor *>(src), const_cast<ggml_tensor *>(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<ggml_tensor *>(match.logits), const_cast<ggml_tensor *>(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<ggml_tensor *>(match.experts), const_cast<ggml_tensor *>(match.dst));
add_alloc_dep(user_data, const_cast<ggml_tensor *>(match.weights), const_cast<ggml_tensor *>(match.dst));
i += match.node_count - 1;
}
}
+4 -1
View File
@@ -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)