mirror of
https://github.com/ggml-org/llama.cpp.git
synced 2026-09-17 20:31:47 +02:00
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:
@@ -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);
|
||||
|
||||
@@ -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) {
|
||||
|
||||
@@ -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);
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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);
|
||||
|
||||
@@ -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);
|
||||
|
||||
@@ -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;
|
||||
|
||||
|
||||
@@ -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;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user