metal : address fusion review comments

- Fix declaration/table alignment
- Rename top-k MoE kargs fields to val_clamp / val_scale
- Move moe-reduce alloc-deps handling into a general fusion helper
- Remove the public moe-reduce matcher API

Assisted-by: pi:llama.cpp/DeepSeek-V4-Flash-Vision-Exp
This commit is contained in:
Georgi Gerganov
2026-09-15 16:17:57 +03:00
parent 5fbcd5c0f3
commit c1820afb84
8 changed files with 74 additions and 66 deletions
+1 -1
View File
@@ -148,7 +148,7 @@ struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_top_k
struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_top_k_radix (ggml_metal_library_t lib, const struct ggml_tensor * op);
struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_top_k_merge (ggml_metal_library_t lib, const struct ggml_tensor * op);
struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_topk_moe (ggml_metal_library_t lib, bool with_norm);
struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_moe_reduce (ggml_metal_library_t lib);
struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_moe_reduce (ggml_metal_library_t lib);
struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_bin (ggml_metal_library_t lib, const struct ggml_tensor * op, int32_t n_fuse );
struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_bin_one (ggml_metal_library_t lib, enum ggml_op op);
struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_l2_norm (ggml_metal_library_t lib, const struct ggml_tensor * op);
+55 -27
View File
@@ -438,7 +438,14 @@ static bool ggml_metal_fusion_check_topk_moe(
#define GGML_METAL_MOE_REDUCE_MAX_EXPERTS 8
bool ggml_metal_fusion_match_moe_reduce(
struct ggml_metal_moe_reduce_match {
const ggml_tensor * experts;
const ggml_tensor * weights;
const ggml_tensor * dst;
int node_count;
};
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) {
return false;
@@ -611,32 +618,32 @@ static const ggml_op ops_moe_reduce_7[] = { GGML_OP_MUL, GGML_OP_ADD, GGML_OP_AD
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_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, 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 },
};
const ggml_metal_fusion * ggml_metal_fusion_all(int * n) {
@@ -645,6 +652,27 @@ const ggml_metal_fusion * ggml_metal_fusion_all(int * n) {
return ggml_metal_fusions;
}
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_MUL) {
continue;
}
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;
}
}
// ---- shared fusion info ---------------------------------------------------
static std::string ggml_metal_fusion_label(const ggml_metal_fusion * fusion) {
+6 -12
View File
@@ -37,7 +37,7 @@ typedef enum ggml_metal_fusion_id {
GGML_METAL_FUSION_SNAKE, // MUL + SIN + SQR + MUL + ADD
GGML_METAL_FUSION_GDN_CACHE, // GATED_DELTA_NET + CPY (write snapshots into the recurrent cache)
GGML_METAL_FUSION_TOPK_MOE, // SOFT_MAX + ARGSORT + GET_ROWS + norm/scale (MoE routing)
GGML_METAL_FUSION_MOE_REDUCE, // MUL + expert VIEWs + ADD chain (MoE output reduction)
GGML_METAL_FUSION_MOE_REDUCE, // MUL + expert VIEWs + ADD chain (MoE output reduction)
GGML_METAL_FUSION_SSM_CONV_SILU, // SSM_CONV + UNARY (silu)
} ggml_metal_fusion_id;
@@ -64,17 +64,11 @@ struct ggml_metal_fusion {
typedef struct ggml_metal_fusion ggml_metal_fusion;
struct ggml_metal_moe_reduce_match {
const struct ggml_tensor * experts;
const struct ggml_tensor * weights;
const struct ggml_tensor * dst;
int node_count;
};
// match MUL(experts, weights) + expert VIEWs + ADD chain; used by both the fusion
// validator and the graph-optimize alloc-dependency hook
bool ggml_metal_fusion_match_moe_reduce(
const struct ggml_cgraph * gf, int node_idx, struct ggml_metal_moe_reduce_match * match);
// apply any alloc-dependencies required by the fused kernels during graph optimize
void ggml_metal_fusion_add_alloc_deps(
void * user_data,
void (*add_alloc_dep)(void *, struct ggml_tensor *, struct ggml_tensor *),
const struct ggml_cgraph * gf);
// the single table of all fusions supported by the Metal backend
const ggml_metal_fusion * ggml_metal_fusion_all(int * n);
+2 -2
View File
@@ -1230,8 +1230,8 @@ typedef struct {
uint64_t nb01; // logits row stride
uint64_t nb1_ids; // ids row stride
int32_t top_k; // n_expert_used
float clamp_val;
float scale_val;
float val_clamp;
float val_scale;
} ggml_metal_kargs_topk_moe;
typedef struct {
+6 -6
View File
@@ -5422,16 +5422,16 @@ int ggml_metal_op_topk_moe(ggml_metal_op_t ctx, int idx) {
const bool with_norm = n_fuse >= 6;
const bool with_scale = n_fuse == 4 || n_fuse == 7;
float clamp_val = -INFINITY;
float val_clamp = -INFINITY;
if (with_norm) {
ggml_tensor * clamp = ctx->node(idx + 4);
clamp_val = ggml_get_op_params_f32(clamp, 0);
val_clamp = ggml_get_op_params_f32(clamp, 0);
}
float scale_val = 1.0f;
float val_scale = 1.0f;
if (with_scale) {
ggml_tensor * scale = ctx->node(idx + n_fuse - 1);
scale_val = ggml_get_op_params_f32(scale, 0);
val_scale = ggml_get_op_params_f32(scale, 0);
}
ggml_metal_kargs_topk_moe args = {
@@ -5440,8 +5440,8 @@ int ggml_metal_op_topk_moe(ggml_metal_op_t ctx, int idx) {
/*.nb01 =*/ logits->nb[1],
/*.nb1_ids =*/ ids->nb[1],
/*.top_k =*/ (int32_t) n_expert_used,
/*.clamp_val =*/ clamp_val,
/*.scale_val =*/ scale_val,
/*.val_clamp =*/ val_clamp,
/*.val_scale =*/ val_scale,
};
auto pipeline = ggml_metal_library_get_pipeline_topk_moe(lib, with_norm);
+1 -1
View File
@@ -98,7 +98,7 @@ int ggml_metal_op_argmax (ggml_metal_op_t ctx, int idx);
int ggml_metal_op_argsort (ggml_metal_op_t ctx, int idx);
int ggml_metal_op_top_k (ggml_metal_op_t ctx, int idx);
int ggml_metal_op_topk_moe (ggml_metal_op_t ctx, int idx);
int ggml_metal_op_moe_reduce (ggml_metal_op_t ctx, int idx);
int ggml_metal_op_moe_reduce (ggml_metal_op_t ctx, int idx);
int ggml_metal_op_tri (ggml_metal_op_t ctx, int idx);
int ggml_metal_op_opt_step_adamw (ggml_metal_op_t ctx, int idx);
int ggml_metal_op_opt_step_sgd (ggml_metal_op_t ctx, int idx);
+1 -15
View File
@@ -565,21 +565,7 @@ static void ggml_backend_metal_graph_optimize(ggml_backend_t backend, ggml_cgrap
// keep the MoE weighted-reduction inputs alive until the fused output so the
// allocator cannot reuse them while the fused kernel is still reading them
for (int i = 0; i < cgraph->n_nodes; ++i) {
if (cgraph->nodes[i]->op != GGML_OP_MUL) {
continue;
}
ggml_metal_moe_reduce_match match;
if (!ggml_metal_fusion_match_moe_reduce(cgraph, i, &match)) {
continue;
}
params->add_alloc_dep(params->user_data, (ggml_tensor *) match.experts, (ggml_tensor *) match.dst);
params->add_alloc_dep(params->user_data, (ggml_tensor *) match.weights, (ggml_tensor *) match.dst);
i += match.node_count - 1;
}
ggml_metal_fusion_add_alloc_deps(params->user_data, params->add_alloc_dep, cgraph);
ggml_metal_t ctx = (ggml_metal_t)backend->context;
+2 -2
View File
@@ -433,7 +433,7 @@ kernel void kernel_topk_moe_f32(
if (FC_topk_moe_with_norm) {
wt_sum = simd_sum(wt_sum);
wt_sum = max(wt_sum, args.clamp_val);
wt_sum = max(wt_sum, args.val_clamp);
const float inv = 1.0f / wt_sum;
for (int i = 0; i < n_per_lane; ++i) {
output_weights[i] *= inv;
@@ -443,7 +443,7 @@ kernel void kernel_topk_moe_f32(
for (int i = 0; i < n_per_lane; ++i) {
const int idx = i * 32 + lane;
if (idx < top_k) {
weights_row[idx] = output_weights[i] * args.scale_val;
weights_row[idx] = output_weights[i] * args.val_scale;
}
}
}