Commit 5b59b83f4 for llama.cpp
commit 5b59b83f4e2101ea173d4f853a0522d9971f48c6
Author: Georgi Gerganov <ggerganov@gmail.com>
Date: Sat Sep 19 13:14:44 2026 +0300
metal : add MoE and SSM_CONV fusion optimizations (#28948)
* metal : add top-k MoE fusion
Adds a Metal fusion for SOFT_MAX + ARGSORT + GET_ROWS with optional
routing-weight normalization and scale, matching the top-k MoE fusion
available in the CUDA and Vulkan backends. The fused kernel writes the
selected expert ids and routing weights directly, eliding the separate
softmax, argsort, get-rows, sum-rows, clamp, div and scale kernels.
Assisted-by: pi:llama.cpp/DeepSeek-V4-Flash-Vision-Exp
* metal : add MoE weighted reduction fusion
Fuses MUL(experts, weights) plus the expert VIEW/ADD chain into one kernel
that computes the weighted sum directly. The graph_optimize hook keeps the
expert and weight buffers alive until the fused output so the allocator cannot
reuse them while the kernel is still reading them.
Assisted-by: pi:llama.cpp/DeepSeek-V4-Flash-Vision-Exp
* tests : expose MoE weighted reduction in fusion baseline
Use 2 experts per token in the generated MoE test models so the Metal
MoE weighted reduction fusion (MUL + ADD) is exercised by test-fusion.
Assisted-by: pi:llama.cpp/DeepSeek-V4-Flash-Vision-Exp
* metal : fuse RMS_NORM + SCALE
Adds NORM/RMS_NORM + SCALE fusion to the Metal backend by reusing the
norm+mul kernel with a scalar scale flag. Adds test coverage for both
NORM+SCALE and RMS_NORM+SCALE and regenerates the fusion baseline.
Assisted-by: pi:llama.cpp/DeepSeek-V4-Flash-Vision-Exp
* metal : use function constant for RMS_NORM + SCALE
Replaces the runtime use_scale karg with a Metal function constant. The
norm+mul kernel is compiled with FC_norm_use_scale=false for MUL fusion and
FC_norm_use_scale=true for SCALE fusion, so the fused kernel has no runtime
branch.
Assisted-by: pi:llama.cpp/DeepSeek-V4-Flash-Vision-Exp
* metal : use function constant for top-k MoE with_norm
Replaces the runtime with_norm karg with a Metal function constant. The
top-k MoE kernel is compiled separately for the normalized and non-normalized
routing variants, removing the runtime branch.
Assisted-by: pi:llama.cpp/DeepSeek-V4-Flash-Vision-Exp
* metal : rename moe_weighted_reduction suffix to moe_reduce
Shortens the MoE weighted-reduction fusion identifiers, kernel, pipeline,
matcher, args struct, and test op name from moe_weighted_reduction to
moe_reduce.
Assisted-by: pi:llama.cpp/DeepSeek-V4-Flash-Vision-Exp
* metal : add MUL_MAT + UNARY and MUL_MAT + ADD + UNARY fusion
Adds dense mat-vec activation fusion for sigmoid/silu and bias+softplus.
The mat-vec kernels apply the activation/bias epilogue via function
constants, avoiding the separate unary/add passes.
Assisted-by: pi:llama.cpp/DeepSeek-V4-Flash-Vision-Exp
* metal : revert MUL_MAT + UNARY and MUL_MAT + ADD + UNARY fusion
The mat-vec activation fusion regressed decode throughput on Qwen3.6-35B-A3B
by ~8% (tg32 81.5 vs 88.5 t/s). The regression is caused by loss of
concurrency: the standalone unary kernels previously overlapped with other
mat-vec work, while fusing the activation into the mat-vec kernel serializes
it on the critical path.
Assisted-by: pi:llama.cpp/DeepSeek-V4-Flash-Vision-Exp
* metal : add SSM_CONV + UNARY (silu) fusion
The SSM_CONV kernels apply silu directly via a function constant, eliding
the separate unary pass. Regenerates the fusion baseline.
Assisted-by: pi:llama.cpp/DeepSeek-V4-Flash-Vision-Exp
* 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
* metal : fix unused parameter in top-k MoE fusion check
Assisted-by: pi:llama.cpp/DeepSeek-V4-Flash-Vision-Exp
* metal : guard SSM_CONV fusion lookup behind use_fusion
Assisted-by: pi:llama.cpp/DeepSeek-V4-Flash-Vision-Exp
* metal : track all fused outputs in graph reorder
Assisted-by: pi:llama.cpp/DeepSeek-V4-Flash-Vision-Exp
* metal : keep top-k MoE logits alive until fused output
Assisted-by: pi:llama.cpp/DeepSeek-V4-Flash-Vision-Exp
* metal : refactor alloc deps to pattern-driven approach
Assisted-by: pi:llama.cpp/DeepSeek-V4-Flash-Vision-Exp
* metal : check fused kernel destination in concurrency tracking
Assisted-by: pi:llama.cpp/DeepSeek-V4-Flash-Vision-Exp
* meta : forward graph_optimize to underlying backends
Assisted-by: pi:llama.cpp/DeepSeek-V4-Flash-Vision-Exp
* metal : use vector for fusion table
Assisted-by: pi:llama.cpp/DeepSeek-V4-Flash-Vision-Exp
* meta : keep graph_optimize unimplemented
Assisted-by: pi:llama.cpp/DeepSeek-V4-Flash-Vision-Exp
* parallel : fix non-deterministic prompt selection
Assisted-by: pi:llama.cpp/DeepSeek-V4-Flash-Vision-Exp
* parallel : support dummy models and add global logits run hash
Assisted-by: pi:llama.cpp/DeepSeek-V4-Flash-Vision-Exp
* metal : sync cross-device copies with destination completion event
Assisted-by: pi:llama.cpp/DeepSeek-V4-Flash-Vision-Exp
* metal : avoid const_cast in fusion alloc deps
Assisted-by: pi:llama.cpp/DeepSeek-V4-Flash-Vision-Exp
* metal : skip fusions with aliased sources
Assisted-by: pi:llama.cpp/DeepSeek-V4-Flash-Vision-Exp
* metal : hide fusion pattern definition
Assisted-by: pi:llama.cpp/DeepSeek-V4-Flash-Vision-Exp
* metal : use vector fusion op sequences
Assisted-by: pi:llama.cpp/DeepSeek-V4-Flash-Vision-Exp
* metal : drop redundant struct keywords
Assisted-by: pi:llama.cpp/DeepSeek-V4-Flash-Vision-Exp
* metal : add alloc deps comment separator
Assisted-by: pi:llama.cpp/DeepSeek-V4-Flash-Vision-Exp
* metal : generalize fusion output memory ranges
Assisted-by: pi:llama.cpp/DeepSeek-V4-Flash-Vision-Exp
* metal : rename fusion out_offsets to outs
Assisted-by: pi:llama.cpp/DeepSeek-V4-Flash-Vision-Exp
* metal : avoid dst vector in memory range check
Assisted-by: pi:llama.cpp/DeepSeek-V4-Flash-Vision-Exp
* metal : optimize fusion matching and multi-output handling
- use pointer arithmetic for fusion info count lookup
- avoid heap allocations in top-k MoE and MoE reduce pattern matchers
- use fusion outs for multi-output subgraph checks
Assisted-by: pi:llama.cpp/DeepSeek-V4-Flash-Vision-Exp
* Revert "parallel : support dummy models and add global logits run hash"
This reverts commit 57c7caf941c1b43c270fd5009c9f175063522e96.
* fusion : update MTL.csv
* metal : unroll constant loops in top-k MoE kernel
Assisted-by: pi:llama.cpp/DeepSeek-V4-Flash-Vision-Exp
* metal : use function constants for top-k MoE n_expert and top_k
Assisted-by: pi:llama.cpp/DeepSeek-V4-Flash-Vision-Exp
* metal : rename fusion kargs to scale and clamp
Assisted-by: pi:llama.cpp/DeepSeek-V4-Flash-Vision-Exp
* metal : use function constants for moe_reduce and ssm_conv
Assisted-by: pi:llama.cpp/DeepSeek-V4-Flash-Vision-Exp
* fusion : update MTL.csv
diff --git a/examples/parallel/parallel.cpp b/examples/parallel/parallel.cpp
index a46400c5b..4b74540f0 100644
--- a/examples/parallel/parallel.cpp
+++ b/examples/parallel/parallel.cpp
@@ -11,6 +11,7 @@
#include <clocale>
#include <cmath>
#include <cstdio>
+#include <random>
#include <string>
#include <vector>
#include <ctime>
@@ -156,7 +157,7 @@ static std::vector<std::string> split_string(const std::string& input, char deli
int main(int argc, char ** argv) {
std::setlocale(LC_NUMERIC, "C");
- srand(1234);
+ std::mt19937 rng(1234);
common_params params;
@@ -321,7 +322,7 @@ int main(int argc, char ** argv) {
client.t_start_prompt = ggml_time_us();
client.t_start_gen = 0;
- client.input = k_prompts[rand() % k_prompts.size()];
+ client.input = k_prompts[rng() % k_prompts.size()];
client.response = "";
// construct the prompt:
@@ -334,10 +335,10 @@ int main(int argc, char ** argv) {
client.prompt += k_system;
}
- const int n_junk_cur = rand() % n_junk;
+ const int n_junk_cur = rng() % n_junk;
for (int i = 0; i < n_junk_cur; ++i) {
- const int r = rand() % k_questions.size();
+ const int r = rng() % k_questions.size();
client.prompt += "User:\n" + k_questions[r] + "\nAssistant:\n " + k_answers[r] + "\n";
}
client.prompt += "User:\n" + client.input + "\nAssistant:\n";
diff --git a/ggml/src/ggml-metal/ggml-metal-common.cpp b/ggml/src/ggml-metal/ggml-metal-common.cpp
index 05755eb3b..388ac4185 100644
--- a/ggml/src/ggml-metal/ggml-metal-common.cpp
+++ b/ggml/src/ggml-metal/ggml-metal-common.cpp
@@ -222,38 +222,63 @@ struct node_info {
void add_fused(ggml_tensor * t) {
fused.push_back(t);
}
+
+ bool is_output(const ggml_tensor * t) const {
+ if (t == node) {
+ return true;
+ }
+ for (const auto * f : fused) {
+ if (t == f) {
+ return true;
+ }
+ }
+ return false;
+ }
};
static std::vector<int> ggml_metal_graph_optimize_reorder(const std::vector<node_info> & nodes) {
// helper to add node src and dst ranges
const auto & h_add = [](ggml_mem_ranges_t mrs, const node_info & node) {
+ // only external sources matter: sources produced by the fused group are internal
for (int i = 0; i < GGML_MAX_SRC; i++) {
- if (node.node->src[i]) {
- if (!ggml_mem_ranges_add_src(mrs, node.node->src[i])) {
+ const ggml_tensor * src = node.node->src[i];
+ if (src && !node.is_output(src)) {
+ if (!ggml_mem_ranges_add_src(mrs, src)) {
return false;
}
}
}
- // keep track of the sources of the fused nodes as well
for (const auto * fused : node.fused) {
for (int i = 0; i < GGML_MAX_SRC; i++) {
- if (fused->src[i]) {
- if (!ggml_mem_ranges_add_src(mrs, fused->src[i])) {
+ const ggml_tensor * src = fused->src[i];
+ if (src && !node.is_output(src)) {
+ if (!ggml_mem_ranges_add_src(mrs, src)) {
return false;
}
}
}
}
- return ggml_mem_ranges_add_dst(mrs, node.dst());
+ // all fused tensors are produced by the fused kernel
+ if (!ggml_mem_ranges_add_dst(mrs, node.node)) {
+ return false;
+ }
+ for (const auto * fused : node.fused) {
+ if (!ggml_mem_ranges_add_dst(mrs, fused)) {
+ return false;
+ }
+ }
+
+ return true;
};
// helper to check if a node can run concurrently with the existing set of nodes
const auto & h_check = [](ggml_mem_ranges_t mrs, const node_info & node) {
for (int i = 0; i < GGML_MAX_SRC; i++) {
- if (node.node->src[i]) {
- if (!ggml_mem_ranges_check_src(mrs, node.node->src[i])) {
+ const ggml_tensor * src = node.node->src[i];
+ if (src && !node.is_output(src)) {
+ if (!ggml_mem_ranges_check_src(mrs, src)) {
return false;
}
}
@@ -261,15 +286,25 @@ static std::vector<int> ggml_metal_graph_optimize_reorder(const std::vector<node
for (const auto * fused : node.fused) {
for (int i = 0; i < GGML_MAX_SRC; i++) {
- if (fused->src[i]) {
- if (!ggml_mem_ranges_check_src(mrs, fused->src[i])) {
+ const ggml_tensor * src = fused->src[i];
+ if (src && !node.is_output(src)) {
+ if (!ggml_mem_ranges_check_src(mrs, src)) {
return false;
}
}
}
}
- return ggml_mem_ranges_check_dst(mrs, node.dst());
+ if (!ggml_mem_ranges_check_dst(mrs, node.node)) {
+ return false;
+ }
+ for (const auto * fused : node.fused) {
+ if (!ggml_mem_ranges_check_dst(mrs, fused)) {
+ return false;
+ }
+ }
+
+ return true;
};
// perform reorders only across these types of ops
diff --git a/ggml/src/ggml-metal/ggml-metal-context.m b/ggml/src/ggml-metal/ggml-metal-context.m
index bf4fe2dcd..442ed2a07 100644
--- a/ggml/src/ggml-metal/ggml-metal-context.m
+++ b/ggml/src/ggml-metal/ggml-metal-context.m
@@ -30,7 +30,8 @@ struct ggml_metal {
ggml_metal_device_t dev;
ggml_metal_library_t lib;
- ggml_metal_event_t ev_cpy; // for async copies
+ ggml_metal_event_t ev_cpy; // for async copies
+ ggml_metal_event_t ev_sync; // destination completion signal
dispatch_queue_t d_queue;
@@ -129,7 +130,8 @@ ggml_metal_t ggml_metal_init(ggml_metal_device_t dev) {
}
}
- res->ev_cpy = ggml_metal_device_event_init(dev);
+ res->ev_cpy = ggml_metal_device_event_init(dev);
+ res->ev_sync = ggml_metal_device_event_init(dev);
const struct ggml_metal_device_props * props_dev = ggml_metal_device_get_props(dev);
@@ -240,6 +242,7 @@ void ggml_metal_free(ggml_metal_t ctx) {
dispatch_release(ctx->d_queue);
ggml_metal_device_event_free(ctx->dev, ctx->ev_cpy);
+ ggml_metal_device_event_free(ctx->dev, ctx->ev_sync);
free(ctx);
}
@@ -421,10 +424,23 @@ bool ggml_metal_cpy_tensor_async(ggml_metal_t ctx_src, ggml_metal_t ctx_dst, con
return false;
}
+ id<MTLCommandQueue> dst_queue = ggml_metal_device_get_queue(ctx_dst->dev);
+ id<MTLCommandBuffer> sync_cmd_buf = [dst_queue commandBuffer];
+
+ ggml_metal_event_encode_signal(ctx_dst->ev_sync, sync_cmd_buf);
+
+ [sync_cmd_buf commit];
+
+ [ctx_dst->cmd_bufs_ext addObject:sync_cmd_buf];
+ ctx_dst->cmd_buf_last = sync_cmd_buf;
+
+ [sync_cmd_buf retain];
+
// queue the copy operation into the Metal context
// this will be queued at the end, after any currently ongoing GPU operations
id<MTLCommandQueue> queue = ggml_metal_device_get_queue(ctx_src->dev);
id<MTLCommandBuffer> cmd_buf = [queue commandBuffer];
+ ggml_metal_event_encode_wait(ctx_dst->ev_sync, cmd_buf);
id<MTLBlitCommandEncoder> encoder = [cmd_buf blitCommandEncoder];
[encoder copyFromBuffer:bid_src.metal
diff --git a/ggml/src/ggml-metal/ggml-metal-device.cpp b/ggml/src/ggml-metal/ggml-metal-device.cpp
index 0dcfad3af..c08ec10b6 100644
--- a/ggml/src/ggml-metal/ggml-metal-device.cpp
+++ b/ggml/src/ggml-metal/ggml-metal-device.cpp
@@ -529,7 +529,8 @@ ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_dsv4_hc(ggml_met
return res;
}
-ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_ssm_conv(ggml_metal_library_t lib, const ggml_tensor * op) {
+ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_ssm_conv(
+ ggml_metal_library_t lib, const ggml_tensor * op, int32_t nc, bool use_silu) {
GGML_ASSERT(op->src[0]->type == GGML_TYPE_F32);
GGML_ASSERT(op->src[1]->type == GGML_TYPE_F32);
@@ -546,17 +547,24 @@ ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_ssm_conv(ggml_me
}
snprintf(base, 256, "kernel_ssm_conv_%s_%s%s", ggml_type_name(op->src[0]->type), ggml_type_name(op->src[1]->type), suffix);
- snprintf(name, 256, "%s", base);
+ snprintf(name, 256, "%s_nc=%d_silu=%d", base, nc, use_silu ? 1 : 0);
ggml_metal_pipeline_with_params res = ggml_metal_library_get_pipeline(lib, name);
if (!res.pipeline) {
- res = ggml_metal_library_compile_pipeline(lib, base, name, nullptr);
+ ggml_metal_cv_t cv = ggml_metal_cv_init();
+ ggml_metal_cv_set_bool(cv, use_silu, FC_SSM_CONV + 1);
+ ggml_metal_cv_set_int32(cv, nc, FC_SSM_CONV + 2);
+
+ res = ggml_metal_library_compile_pipeline(lib, base, name, cv);
+
+ ggml_metal_cv_free(cv);
}
return res;
}
-ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_ssm_conv_batched(ggml_metal_library_t lib, const ggml_tensor * op, int ssm_conv_bs) {
+ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_ssm_conv_batched(
+ ggml_metal_library_t lib, const ggml_tensor * op, int ssm_conv_bs, int32_t nc, bool use_silu) {
GGML_ASSERT(op->src[0]->type == GGML_TYPE_F32);
GGML_ASSERT(op->src[1]->type == GGML_TYPE_F32);
@@ -572,13 +580,15 @@ ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_ssm_conv_batched
}
snprintf(base, 256, "kernel_ssm_conv_%s_%s_batched%s", ggml_type_name(op->src[0]->type), ggml_type_name(op->src[1]->type), suffix);
- snprintf(name, 256, "%s_ssm_conv_bs=%d", base, ssm_conv_bs);
+ snprintf(name, 256, "%s_ssm_conv_bs=%d_nc=%d_silu=%d", base, ssm_conv_bs, nc, use_silu ? 1 : 0);
ggml_metal_pipeline_with_params res = ggml_metal_library_get_pipeline(lib, name);
if (!res.pipeline) {
ggml_metal_cv_t cv = ggml_metal_cv_init();
ggml_metal_cv_set_int16(cv, ssm_conv_bs, FC_SSM_CONV + 0);
+ ggml_metal_cv_set_bool(cv, use_silu, FC_SSM_CONV + 1);
+ ggml_metal_cv_set_int32(cv, nc, FC_SSM_CONV + 2);
res = ggml_metal_library_compile_pipeline(lib, base, name, cv);
@@ -1548,6 +1558,49 @@ ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_top_k_merge(ggml
return res;
}
+ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_topk_moe(
+ ggml_metal_library_t lib, int32_t n_expert, int32_t top_k, bool with_norm) {
+ char base[256];
+ char name[256];
+
+ snprintf(base, 256, "kernel_topk_moe_f32");
+ snprintf(name, 256, "%s_n_expert=%d_top_k=%d_with_norm=%d", base, n_expert, top_k, with_norm ? 1 : 0);
+
+ ggml_metal_pipeline_with_params res = ggml_metal_library_get_pipeline(lib, name);
+ if (!res.pipeline) {
+ ggml_metal_cv_t cv = ggml_metal_cv_init();
+ ggml_metal_cv_set_bool (cv, with_norm, FC_TOPK_MOE + 0);
+ ggml_metal_cv_set_int32(cv, n_expert, FC_TOPK_MOE + 1);
+ ggml_metal_cv_set_int32(cv, top_k, FC_TOPK_MOE + 2);
+
+ res = ggml_metal_library_compile_pipeline(lib, base, name, cv);
+
+ ggml_metal_cv_free(cv);
+ }
+
+ return res;
+}
+
+ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_moe_reduce(ggml_metal_library_t lib, int32_t n_expert_used) {
+ char base[256];
+ char name[256];
+
+ snprintf(base, 256, "kernel_moe_reduce_f32");
+ snprintf(name, 256, "%s_n_expert_used=%d", base, n_expert_used);
+
+ ggml_metal_pipeline_with_params res = ggml_metal_library_get_pipeline(lib, name);
+ if (!res.pipeline) {
+ ggml_metal_cv_t cv = ggml_metal_cv_init();
+ ggml_metal_cv_set_int32(cv, n_expert_used, FC_MOE_REDUCE + 0);
+
+ res = ggml_metal_library_compile_pipeline(lib, base, name, cv);
+
+ ggml_metal_cv_free(cv);
+ }
+
+ return res;
+}
+
ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_flash_attn_ext_pad(
ggml_metal_library_t lib,
const struct ggml_tensor * op,
@@ -2002,7 +2055,48 @@ ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_norm(ggml_metal_
ggml_metal_pipeline_with_params res = ggml_metal_library_get_pipeline(lib, name);
if (!res.pipeline) {
- res = ggml_metal_library_compile_pipeline(lib, base, name, nullptr);
+ ggml_metal_cv_t cv = ggml_metal_cv_init();
+ ggml_metal_cv_set_bool(cv, false, FC_NORM + 0);
+
+ res = ggml_metal_library_compile_pipeline(lib, base, name, cv);
+
+ ggml_metal_cv_free(cv);
+ }
+
+ res.smem = 32*sizeof(float);
+
+ return res;
+}
+
+ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_norm_scale(ggml_metal_library_t lib, const ggml_tensor * op) {
+ assert(op->op == GGML_OP_NORM || op->op == GGML_OP_RMS_NORM);
+
+ GGML_ASSERT(ggml_is_contiguous_rows(op->src[0]));
+
+ char base[256];
+ char name[256];
+
+ const char * suffix = "";
+ if (op->ne[0] % 4 == 0) {
+ suffix = "_4";
+ }
+
+ switch (op->op) {
+ case GGML_OP_NORM: snprintf(base, 256, "kernel_norm_mul_f32%s", suffix); break;
+ case GGML_OP_RMS_NORM: snprintf(base, 256, "kernel_rms_norm_mul_f32%s", suffix); break;
+ default: GGML_ABORT("fatal error");
+ }
+
+ snprintf(name, 256, "%s_use_scale", base);
+
+ ggml_metal_pipeline_with_params res = ggml_metal_library_get_pipeline(lib, name);
+ if (!res.pipeline) {
+ ggml_metal_cv_t cv = ggml_metal_cv_init();
+ ggml_metal_cv_set_bool(cv, true, FC_NORM + 0);
+
+ res = ggml_metal_library_compile_pipeline(lib, base, name, cv);
+
+ ggml_metal_cv_free(cv);
}
res.smem = 32*sizeof(float);
diff --git a/ggml/src/ggml-metal/ggml-metal-device.h b/ggml/src/ggml-metal/ggml-metal-device.h
index 0514f9ef0..2497e45c3 100644
--- a/ggml/src/ggml-metal/ggml-metal-device.h
+++ b/ggml/src/ggml-metal/ggml-metal-device.h
@@ -127,8 +127,8 @@ struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_tri
struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_soft_max (ggml_metal_library_t lib, const struct ggml_tensor * op);
struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_lightning_indexer (ggml_metal_library_t lib, const struct ggml_tensor * op);
struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_dsv4_hc (ggml_metal_library_t lib, const struct ggml_tensor * op);
-struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_ssm_conv (ggml_metal_library_t lib, const struct ggml_tensor * op);
-struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_ssm_conv_batched (ggml_metal_library_t lib, const struct ggml_tensor * op, int ssm_conv_bs);
+struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_ssm_conv (ggml_metal_library_t lib, const struct ggml_tensor * op, int32_t nc, bool use_silu);
+struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_ssm_conv_batched (ggml_metal_library_t lib, const struct ggml_tensor * op, int ssm_conv_bs, int32_t nc, bool use_silu);
struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_ssm_scan (ggml_metal_library_t lib, const struct ggml_tensor * op, bool tail);
struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_ssm_scan_ssd_mma (ggml_metal_library_t lib, const struct ggml_tensor * op);
struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_rwkv (ggml_metal_library_t lib, const struct ggml_tensor * op);
@@ -149,11 +149,14 @@ struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_fwht
struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_top_k (ggml_metal_library_t lib, const struct ggml_tensor * op);
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_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_topk_moe (ggml_metal_library_t lib, int32_t n_expert, int32_t top_k, bool with_norm);
+struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_moe_reduce (ggml_metal_library_t lib, int32_t n_expert_used);
+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);
struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_group_norm (ggml_metal_library_t lib, const struct ggml_tensor * op);
struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_norm (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_norm_scale (ggml_metal_library_t lib, const struct ggml_tensor * op);
struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_rope (ggml_metal_library_t lib, const struct ggml_tensor * op);
struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_im2col (ggml_metal_library_t lib, const struct ggml_tensor * op);
struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_conv_transpose_1d (ggml_metal_library_t lib, const struct ggml_tensor * op);
diff --git a/ggml/src/ggml-metal/ggml-metal-fusion.cpp b/ggml/src/ggml-metal/ggml-metal-fusion.cpp
index ac3ac0414..5a160a18d 100644
--- a/ggml/src/ggml-metal/ggml-metal-fusion.cpp
+++ b/ggml/src/ggml-metal/ggml-metal-fusion.cpp
@@ -4,9 +4,38 @@
#include "ggml-metal-device.h"
#include <algorithm>
+#include <cstddef>
+#include <cstring>
+#include <set>
#include <string>
#include <vector>
+struct ggml_metal_fusion {
+ ggml_metal_fusion_id id;
+
+ std::vector<ggml_op> ops; // op sequence (fixed length, non-empty nodes)
+ std::vector<ggml_op> ops_all; // full raw op sequence (may include empty RESHAPE/VIEW nodes)
+ std::vector<int> outs; // additional fused output nodes, relative to 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)
+ bool unsafe;
+
+ // extra backend constraints on top of ggml_can_fuse_subgraph
+ // nodes[j] is the j-th node of the pattern; node_idxs[idx + j] is its raw graph index
+ bool (*check)(const struct ggml_metal_fusion * fusion,
+ const struct ggml_tensor * const * nodes,
+ const struct ggml_cgraph * gf,
+ const int * node_idxs,
+ int idx,
+ ggml_metal_fusion_mode mode);
+};
+
+ggml_metal_fusion_id ggml_metal_fusion_get_id(const ggml_metal_fusion * fusion) {
+ return fusion->id;
+}
+
// ---- helpers -------------------------------------------------------------
// true if two tensors live in the same Metal buffer
@@ -31,12 +60,30 @@ static bool ggml_metal_fusion_same_buffer(const ggml_tensor * a, const ggml_tens
static bool ggml_metal_fusion_check_norm(
const ggml_metal_fusion * fusion,
const ggml_tensor * const * nodes,
+ const ggml_cgraph * gf,
+ const int * node_idxs,
+ int idx,
ggml_metal_fusion_mode mode) {
GGML_UNUSED(mode);
+ GGML_UNUSED(gf);
+ GGML_UNUSED(node_idxs);
+ GGML_UNUSED(idx);
+
+ GGML_ASSERT(fusion->ops.size() >= 2);
- GGML_ASSERT(fusion->n_ops >= 2);
+ if (fusion->id == GGML_METAL_FUSION_NORM_SCALE) {
+ GGML_ASSERT(fusion->ops.size() == 2);
- for (int j = 1; j < fusion->n_ops; j++) {
+ const ggml_tensor * scale = nodes[1];
+ if (scale->op != GGML_OP_SCALE || scale->src[0] != nodes[0] || scale->src[1] ||
+ scale->type != GGML_TYPE_F32) {
+ return false;
+ }
+
+ return true;
+ }
+
+ for (int j = 1; j < (int) fusion->ops.size(); j++) {
// the fused MUL/ADD must read the previous node as src0
if (nodes[j]->src[0] != nodes[j - 1]) {
return false;
@@ -59,15 +106,53 @@ static bool ggml_metal_fusion_check_norm(
return true;
}
+// SSM_CONV + UNARY (silu)
+static bool ggml_metal_fusion_check_ssm_conv_silu(
+ const ggml_metal_fusion * fusion,
+ const ggml_tensor * const * nodes,
+ const ggml_cgraph * gf,
+ const int * node_idxs,
+ int idx,
+ ggml_metal_fusion_mode mode) {
+ GGML_UNUSED(fusion);
+ GGML_UNUSED(gf);
+ GGML_UNUSED(node_idxs);
+ GGML_UNUSED(idx);
+ GGML_UNUSED(mode);
+
+ const ggml_tensor * conv = nodes[0];
+ const ggml_tensor * un = nodes[1];
+
+ if (conv->op != GGML_OP_SSM_CONV || un->op != GGML_OP_UNARY || un->src[0] != conv || un->src[1]) {
+ return false;
+ }
+
+ if (ggml_get_unary_op(un) != GGML_UNARY_OP_SILU) {
+ return false;
+ }
+
+ if (conv->type != GGML_TYPE_F32 || un->type != GGML_TYPE_F32 || !ggml_is_contiguous_rows(un)) {
+ return false;
+ }
+
+ return true;
+}
+
// 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_fusion_check_add_chain(
const ggml_metal_fusion * fusion,
const ggml_tensor * const * nodes,
+ const ggml_cgraph * gf,
+ const int * node_idxs,
+ int idx,
ggml_metal_fusion_mode mode) {
- GGML_ASSERT(fusion->n_ops >= 2);
+ GGML_UNUSED(gf);
+ GGML_UNUSED(node_idxs);
+ GGML_UNUSED(idx);
+ GGML_ASSERT(fusion->ops.size() >= 2);
- for (int j = 1; j < fusion->n_ops; j++) {
+ for (int j = 1; j < (int) fusion->ops.size(); j++) {
if (nodes[j]->src[0] != nodes[j - 1]) {
return false;
}
@@ -94,8 +179,14 @@ static bool ggml_metal_fusion_check_add_chain(
static bool ggml_metal_fusion_check_gdn_cache(
const ggml_metal_fusion * fusion,
const ggml_tensor * const * nodes,
+ const ggml_cgraph * gf,
+ const int * node_idxs,
+ int idx,
ggml_metal_fusion_mode mode) {
GGML_UNUSED(fusion);
+ GGML_UNUSED(gf);
+ GGML_UNUSED(node_idxs);
+ GGML_UNUSED(idx);
const ggml_tensor * gdn = nodes[0];
const ggml_tensor * cpy = nodes[1];
@@ -150,9 +241,15 @@ static bool ggml_metal_fusion_check_gdn_cache(
static bool ggml_metal_fusion_check_snake(
const ggml_metal_fusion * fusion,
const ggml_tensor * const * nodes,
+ const ggml_cgraph * gf,
+ const int * node_idxs,
+ int idx,
ggml_metal_fusion_mode mode) {
GGML_UNUSED(fusion);
GGML_UNUSED(mode);
+ GGML_UNUSED(gf);
+ GGML_UNUSED(node_idxs);
+ GGML_UNUSED(idx);
const ggml_tensor * mul0 = nodes[0];
const ggml_tensor * sin_node = nodes[1];
@@ -195,42 +292,465 @@ static bool ggml_metal_fusion_check_snake(
return types_ok && shape_ok && dim_ok && contig_ok && x_in_add == x;
}
+#define GGML_METAL_TOPK_MOE_MAX_EXPERTS 1024
+
+// 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 std::vector<ggml_op> ops_topk_moe_all = {
+ GGML_OP_SOFT_MAX, GGML_OP_RESHAPE, GGML_OP_ARGSORT, GGML_OP_VIEW, GGML_OP_GET_ROWS
+};
+static const std::vector<ggml_op> ops_topk_moe_scale_all = {
+ GGML_OP_SOFT_MAX, GGML_OP_RESHAPE, GGML_OP_ARGSORT, GGML_OP_VIEW, GGML_OP_GET_ROWS, GGML_OP_SCALE
+};
+static const std::vector<ggml_op> ops_topk_moe_norm_all = {
+ 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 std::vector<ggml_op> ops_topk_moe_norm_scale_all = {
+ 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,
+ const ggml_cgraph * gf,
+ const int * node_idxs,
+ int idx,
+ ggml_metal_fusion_mode mode) {
+ GGML_ASSERT(fusion->ops.size() >= 3);
+ GGML_UNUSED(nodes);
+
+ const int n_ops = (int) fusion->ops.size();
+
+ const bool with_norm = n_ops >= 6;
+ const bool with_scale = n_ops == 4 || n_ops == 7;
+
+ // the fusion table operates on the non-empty node sequence; the raw graph also
+ // contains the RESHAPE/VIEW nodes that the fused kernel elides.
+ const std::vector<ggml_op> & ops_all = fusion->ops_all;
+
+ const int raw_start = node_idxs[idx];
+ int raw_end = node_idxs[idx + n_ops - 1];
+
+ // the norm variant ends with a RESHAPE that the non-empty sequence filters out;
+ // include it so the output use-count check sees the real final routing tensor
+ if (with_norm && !with_scale) {
+ if (raw_end + 1 >= gf->n_nodes) {
+ return false;
+ }
+ const ggml_tensor * trailing_reshape = gf->nodes[raw_end + 1];
+ if (trailing_reshape->op != GGML_OP_RESHAPE || trailing_reshape->src[0] != gf->nodes[raw_end]) {
+ return false;
+ }
+ raw_end++;
+ }
+
+ const int raw_count = raw_end - raw_start + 1;
+ if (raw_count != (int) ops_all.size()) {
+ return false;
+ }
+
+ int raw_idxs[GGML_METAL_FUSION_MAX];
+ for (int i = 0; i < raw_count; ++i) {
+ raw_idxs[i] = raw_start + i;
+ if (gf->nodes[raw_start + i]->op != ops_all[i]) {
+ return false;
+ }
+ }
+
+ const ggml_tensor * softmax = gf->nodes[raw_start];
+ const ggml_tensor * probs_reshaped = gf->nodes[raw_start + 1];
+ const ggml_tensor * argsort = gf->nodes[raw_start + 2];
+ const ggml_tensor * ids = gf->nodes[raw_start + 3];
+ const ggml_tensor * get_rows = gf->nodes[raw_start + 4];
+ const ggml_tensor * out = gf->nodes[raw_end];
+ const ggml_tensor * logits = softmax->src[0];
+
+ // the fused kernel implements plain softmax only
+ float scale = 1.0f;
+ float max_bias = 0.0f;
+ memcpy(&scale, ((const int32_t *) softmax->op_params) + 0, sizeof(scale));
+ memcpy(&max_bias, ((const int32_t *) softmax->op_params) + 1, sizeof(max_bias));
+ if (scale != 1.0f || max_bias != 0.0f || softmax->src[1] || softmax->src[2]) {
+ return false;
+ }
+
+ if (logits->type != GGML_TYPE_F32 || softmax->type != GGML_TYPE_F32 ||
+ out->type != GGML_TYPE_F32 || ids->type != GGML_TYPE_I32) {
+ return false;
+ }
+
+ const int64_t n_expert = logits->ne[0];
+ const int64_t n_tokens = logits->ne[1];
+ const int64_t n_expert_used = ids->ne[0];
+
+ if (n_expert <= 0 || n_tokens <= 0 || n_expert_used <= 0 || n_expert_used > n_expert ||
+ n_expert > GGML_METAL_TOPK_MOE_MAX_EXPERTS || n_expert_used > GGML_METAL_TOPK_MOE_MAX_EXPERTS) {
+ return false;
+ }
+
+ if (logits->ne[2] != 1 || logits->ne[3] != 1 ||
+ ids->ne[1] != n_tokens || ids->ne[2] != 1 || ids->ne[3] != 1 ||
+ out->ne[0] != 1 || out->ne[1] != n_expert_used || out->ne[2] != n_tokens || out->ne[3] != 1) {
+ return false;
+ }
+
+ if (!ggml_is_contiguous(logits) || !ggml_is_contiguous(out) ||
+ ids->nb[0] != ggml_type_size(GGML_TYPE_I32) ||
+ ids->nb[1] != ggml_type_size(GGML_TYPE_I32) * n_expert) {
+ return false;
+ }
+
+ if (probs_reshaped->src[0] != softmax || argsort->src[0] != softmax ||
+ ids->src[0] != argsort || get_rows->src[0] != probs_reshaped || get_rows->src[1] != ids) {
+ return false;
+ }
+
+ if (with_norm) {
+ const ggml_tensor * weights_reshaped = gf->nodes[raw_start + 5];
+ const ggml_tensor * sum_rows = gf->nodes[raw_start + 6];
+ const ggml_tensor * clamp = gf->nodes[raw_start + 7];
+ const ggml_tensor * div = gf->nodes[raw_start + 8];
+ const ggml_tensor * out_reshaped = gf->nodes[raw_start + 9];
+
+ if (weights_reshaped->src[0] != get_rows || sum_rows->src[0] != weights_reshaped ||
+ clamp->src[0] != sum_rows || div->src[0] != weights_reshaped || div->src[1] != clamp ||
+ out_reshaped->src[0] != div) {
+ return false;
+ }
+
+ if (with_scale) {
+ const ggml_tensor * scale_node = gf->nodes[raw_start + 10];
+ if (scale_node->src[0] != out_reshaped) {
+ return false;
+ }
+ }
+ } else if (with_scale) {
+ const ggml_tensor * scale_node = gf->nodes[raw_start + 5];
+ if (scale_node->src[0] != get_rows) {
+ return false;
+ }
+ }
+
+ const int outputs[2] = { raw_start + 3, raw_end };
+ if (!ggml_can_fuse_subgraph_ext(gf, raw_idxs, raw_count, ops_all.data(), outputs, 2)) {
+ return false;
+ }
+
+ if (mode == GGML_METAL_FUSION_FULL) {
+ if (!logits->data || !out->data || !ids->data) {
+ return false;
+ }
+ }
+
+ return true;
+}
+
+#define GGML_METAL_MOE_REDUCE_MAX_EXPERTS 8
+
+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, const std::vector<ggml_op> & ops_all,
+ ggml_metal_moe_reduce_match * match) {
+ if (match == nullptr || node_idx < 0 || node_idx + (int) ops_all.size() > gf->n_nodes) {
+ return false;
+ }
+
+ const ggml_tensor * mul = gf->nodes[node_idx];
+ if (mul->op != GGML_OP_MUL || mul->type != GGML_TYPE_F32) {
+ return false;
+ }
+
+ // MUL, then one VIEW per expert, then one ADD per additional expert
+ const int raw_count = (int) ops_all.size();
+ const int n_expert_used = raw_count / 2;
+
+ if (n_expert_used < 2 || n_expert_used > GGML_METAL_MOE_REDUCE_MAX_EXPERTS ||
+ raw_count != 2 * n_expert_used) {
+ return false;
+ }
+
+ int n_views = 0;
+ while (node_idx + 1 + n_views < gf->n_nodes &&
+ gf->nodes[node_idx + 1 + n_views]->op == GGML_OP_VIEW) {
+ n_views++;
+ }
+
+ if (n_views != n_expert_used) {
+ return false;
+ }
+
+ for (int i = n_expert_used + 1; i < raw_count; ++i) {
+ if (gf->nodes[node_idx + i]->op != GGML_OP_ADD) {
+ return false;
+ }
+ }
+
+ int raw_idxs[GGML_METAL_FUSION_MAX];
+ for (int i = 0; i < raw_count; ++i) {
+ raw_idxs[i] = node_idx + i;
+ if (gf->nodes[node_idx + i]->op != ops_all[i]) {
+ return false;
+ }
+ }
+
+ const ggml_tensor * experts = mul->src[0];
+ const ggml_tensor * weights = mul->src[1];
+ const ggml_tensor * dst = gf->nodes[node_idx + raw_count - 1];
+
+ if (experts->type != GGML_TYPE_F32 || weights->type != GGML_TYPE_F32 || dst->type != GGML_TYPE_F32) {
+ return false;
+ }
+
+ const int64_t n_embd = experts->ne[0];
+ const int64_t n_tokens = experts->ne[2];
+
+ if (n_embd <= 0 || n_tokens <= 0 || experts->ne[1] != n_expert_used || experts->ne[3] != 1 ||
+ weights->ne[0] != 1 || weights->ne[1] != n_expert_used || weights->ne[2] != n_tokens || weights->ne[3] != 1 ||
+ dst->ne[0] != n_embd || dst->ne[1] != n_tokens || dst->ne[2] != 1 || dst->ne[3] != 1) {
+ return false;
+ }
+
+ if (!ggml_is_contiguous(experts) || !ggml_is_contiguous(weights) || !ggml_is_contiguous(dst)) {
+ return false;
+ }
+
+ for (int i = 1; i <= n_expert_used; ++i) {
+ const ggml_tensor * view = gf->nodes[node_idx + i];
+ if (view->view_src != mul || view->src[0] != mul ||
+ view->view_offs != (size_t) (i - 1) * mul->nb[1] ||
+ view->ne[0] != n_embd || view->ne[1] != n_tokens ||
+ view->nb[1] != mul->nb[2]) {
+ return false;
+ }
+ }
+
+ const ggml_tensor * prev_add = nullptr;
+ for (int j = 1; j < n_expert_used; ++j) {
+ const ggml_tensor * add = gf->nodes[node_idx + n_expert_used + j];
+ const ggml_tensor * rhs = gf->nodes[node_idx + j + 1];
+ const ggml_tensor * lhs = j == 1 ? gf->nodes[node_idx + 1] : prev_add;
+ if (add->src[0] != lhs || add->src[1] != rhs) {
+ return false;
+ }
+ prev_add = add;
+ }
+
+ const int outputs[1] = { node_idx + raw_count - 1 };
+ if (!ggml_can_fuse_subgraph_ext(gf, raw_idxs, raw_count, ops_all.data(), outputs, 1)) {
+ return false;
+ }
+
+ match->experts = experts;
+ match->weights = weights;
+ match->dst = dst;
+ match->node_count = raw_count;
+ return true;
+}
+
+static bool ggml_metal_fusion_check_moe_reduce(
+ const ggml_metal_fusion * fusion,
+ const ggml_tensor * const * nodes,
+ const ggml_cgraph * gf,
+ const int * node_idxs,
+ int idx,
+ ggml_metal_fusion_mode mode) {
+ GGML_UNUSED(nodes);
+
+ ggml_metal_moe_reduce_match match;
+ if (!ggml_metal_fusion_match_moe_reduce(gf, node_idxs[idx], fusion->ops_all, &match)) {
+ return false;
+ }
+
+ if ((int) fusion->ops.size() != match.experts->ne[1]) {
+ return false;
+ }
+
+ const int raw_end = node_idxs[idx] + match.node_count - 1;
+ if (node_idxs[idx + (int) fusion->ops.size() - 1] != raw_end) {
+ return false;
+ }
+
+ if (mode == GGML_METAL_FUSION_FULL) {
+ if (!match.experts->data || !match.weights->data || !match.dst->data) {
+ return false;
+ }
+ }
+
+ return true;
+}
+
// ---- patterns ------------------------------------------------------------
-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 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_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_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_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 },
+static const std::vector<ggml_op> ops_norm_mul = { GGML_OP_NORM, GGML_OP_MUL };
+static const std::vector<ggml_op> ops_norm_mul_add = { GGML_OP_NORM, GGML_OP_MUL, GGML_OP_ADD };
+static const std::vector<ggml_op> ops_norm_scale = { GGML_OP_NORM, GGML_OP_SCALE };
+static const std::vector<ggml_op> ops_rms_norm_mul = { GGML_OP_RMS_NORM, GGML_OP_MUL };
+static const std::vector<ggml_op> ops_rms_norm_mul_add = { GGML_OP_RMS_NORM, GGML_OP_MUL, GGML_OP_ADD };
+static const std::vector<ggml_op> ops_rms_norm_scale = { GGML_OP_RMS_NORM, GGML_OP_SCALE };
+
+static const std::vector<ggml_op> ops_add_2 = { GGML_OP_ADD, GGML_OP_ADD };
+static const std::vector<ggml_op> ops_add_3 = { GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD };
+static const std::vector<ggml_op> ops_add_4 = { GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD };
+static const std::vector<ggml_op> ops_add_5 = { GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD };
+static const std::vector<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 std::vector<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 std::vector<ggml_op> ops_snake = { GGML_OP_MUL, GGML_OP_SIN, GGML_OP_SQR, GGML_OP_MUL, GGML_OP_ADD };
+
+static const std::vector<ggml_op> ops_gdn_cache = { GGML_OP_GATED_DELTA_NET, GGML_OP_CPY };
+
+static const std::vector<ggml_op> ops_topk_moe = {
+ GGML_OP_SOFT_MAX, GGML_OP_ARGSORT, GGML_OP_GET_ROWS
+};
+static const std::vector<ggml_op> ops_topk_moe_scale = {
+ GGML_OP_SOFT_MAX, GGML_OP_ARGSORT, GGML_OP_GET_ROWS, GGML_OP_SCALE
+};
+static const std::vector<ggml_op> ops_topk_moe_norm = {
+ GGML_OP_SOFT_MAX, GGML_OP_ARGSORT, GGML_OP_GET_ROWS,
+ GGML_OP_SUM_ROWS, GGML_OP_CLAMP, GGML_OP_DIV
+};
+static const std::vector<ggml_op> ops_topk_moe_norm_scale = {
+ GGML_OP_SOFT_MAX, GGML_OP_ARGSORT, GGML_OP_GET_ROWS,
+ GGML_OP_SUM_ROWS, GGML_OP_CLAMP, GGML_OP_DIV, GGML_OP_SCALE
+};
+
+static const std::vector<ggml_op> ops_ssm_conv_silu = { GGML_OP_SSM_CONV, GGML_OP_UNARY };
+
+static const std::vector<ggml_op> ops_moe_reduce_2 = { GGML_OP_MUL, GGML_OP_ADD };
+static const std::vector<ggml_op> ops_moe_reduce_3 = { GGML_OP_MUL, GGML_OP_ADD, GGML_OP_ADD };
+static const std::vector<ggml_op> ops_moe_reduce_4 = { GGML_OP_MUL, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD };
+static const std::vector<ggml_op> ops_moe_reduce_5 = { GGML_OP_MUL, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD };
+static const std::vector<ggml_op> ops_moe_reduce_6 = { GGML_OP_MUL, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD };
+static const std::vector<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 std::vector<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 std::vector<ggml_op> ops_moe_reduce_all_2 = {
+ GGML_OP_MUL, GGML_OP_VIEW, GGML_OP_VIEW, GGML_OP_ADD
+};
+static const std::vector<ggml_op> ops_moe_reduce_all_3 = {
+ GGML_OP_MUL, GGML_OP_VIEW, GGML_OP_VIEW, GGML_OP_VIEW, GGML_OP_ADD, GGML_OP_ADD
+};
+static const std::vector<ggml_op> ops_moe_reduce_all_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 std::vector<ggml_op> ops_moe_reduce_all_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 std::vector<ggml_op> ops_moe_reduce_all_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 std::vector<ggml_op> ops_moe_reduce_all_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 std::vector<ggml_op> ops_moe_reduce_all_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 std::vector<ggml_metal_fusion> ggml_metal_fusions = {
+ { GGML_METAL_FUSION_NORM_MUL, ops_norm_mul, ops_norm_mul, {}, false, ggml_metal_fusion_check_norm },
+ { GGML_METAL_FUSION_NORM_MUL_ADD, ops_norm_mul_add, ops_norm_mul_add, {}, false, ggml_metal_fusion_check_norm },
+ { GGML_METAL_FUSION_NORM_SCALE, ops_norm_scale, ops_norm_scale, {}, false, ggml_metal_fusion_check_norm },
+ { GGML_METAL_FUSION_NORM_MUL, ops_rms_norm_mul, ops_rms_norm_mul, {}, false, ggml_metal_fusion_check_norm },
+ { GGML_METAL_FUSION_NORM_MUL_ADD, ops_rms_norm_mul_add, ops_rms_norm_mul_add, {}, false, ggml_metal_fusion_check_norm },
+ { GGML_METAL_FUSION_NORM_SCALE, ops_rms_norm_scale, ops_rms_norm_scale, {}, false, ggml_metal_fusion_check_norm },
+ { GGML_METAL_FUSION_ADD_CHAIN, ops_add_2, ops_add_2, {}, false, ggml_metal_fusion_check_add_chain },
+ { GGML_METAL_FUSION_ADD_CHAIN, ops_add_3, ops_add_3, {}, false, ggml_metal_fusion_check_add_chain },
+ { GGML_METAL_FUSION_ADD_CHAIN, ops_add_4, ops_add_4, {}, false, ggml_metal_fusion_check_add_chain },
+ { GGML_METAL_FUSION_ADD_CHAIN, ops_add_5, ops_add_5, {}, false, ggml_metal_fusion_check_add_chain },
+ { GGML_METAL_FUSION_ADD_CHAIN, ops_add_6, ops_add_6, {}, false, ggml_metal_fusion_check_add_chain },
+ { GGML_METAL_FUSION_ADD_CHAIN, ops_add_7, ops_add_7, {}, false, ggml_metal_fusion_check_add_chain },
+ { GGML_METAL_FUSION_SNAKE, ops_snake, ops_snake, {}, false, ggml_metal_fusion_check_snake },
+ { GGML_METAL_FUSION_GDN_CACHE, ops_gdn_cache, ops_gdn_cache, {}, true, ggml_metal_fusion_check_gdn_cache },
+ { GGML_METAL_FUSION_TOPK_MOE, ops_topk_moe, ops_topk_moe_all, {1}, true, ggml_metal_fusion_check_topk_moe },
+ { GGML_METAL_FUSION_TOPK_MOE, ops_topk_moe_scale, ops_topk_moe_scale_all, {1}, true, ggml_metal_fusion_check_topk_moe },
+ { GGML_METAL_FUSION_TOPK_MOE, ops_topk_moe_norm, ops_topk_moe_norm_all, {1}, true, ggml_metal_fusion_check_topk_moe },
+ { GGML_METAL_FUSION_TOPK_MOE, ops_topk_moe_norm_scale, ops_topk_moe_norm_scale_all, {1}, true, ggml_metal_fusion_check_topk_moe },
+ { GGML_METAL_FUSION_MOE_REDUCE, ops_moe_reduce_2, ops_moe_reduce_all_2, {}, true, ggml_metal_fusion_check_moe_reduce },
+ { GGML_METAL_FUSION_MOE_REDUCE, ops_moe_reduce_3, ops_moe_reduce_all_3, {}, true, ggml_metal_fusion_check_moe_reduce },
+ { GGML_METAL_FUSION_MOE_REDUCE, ops_moe_reduce_4, ops_moe_reduce_all_4, {}, true, ggml_metal_fusion_check_moe_reduce },
+ { GGML_METAL_FUSION_MOE_REDUCE, ops_moe_reduce_5, ops_moe_reduce_all_5, {}, true, ggml_metal_fusion_check_moe_reduce },
+ { GGML_METAL_FUSION_MOE_REDUCE, ops_moe_reduce_6, ops_moe_reduce_all_6, {}, true, ggml_metal_fusion_check_moe_reduce },
+ { GGML_METAL_FUSION_MOE_REDUCE, ops_moe_reduce_7, ops_moe_reduce_all_7, {}, true, ggml_metal_fusion_check_moe_reduce },
+ { GGML_METAL_FUSION_MOE_REDUCE, ops_moe_reduce_8, ops_moe_reduce_all_8, {}, true, ggml_metal_fusion_check_moe_reduce },
+ { GGML_METAL_FUSION_SSM_CONV_SILU, ops_ssm_conv_silu, ops_ssm_conv_silu, {}, false, ggml_metal_fusion_check_ssm_conv_silu },
};
-const ggml_metal_fusion * ggml_metal_fusion_all(int * n) {
- *n = (int) sizeof(ggml_metal_fusions) / sizeof(ggml_metal_fusions[0]);
+// ---- alloc deps -----------------------------------------------------------
+
+static bool ggml_metal_fusion_match_raw_pattern(
+ const ggml_cgraph * gf, int node_idx, const std::vector<ggml_op> & ops) {
+ if (node_idx < 0 || node_idx + (int) ops.size() > gf->n_nodes) {
+ return false;
+ }
+
+ for (int i = 0; i < (int) ops.size(); ++i) {
+ if (gf->nodes[node_idx + i]->op != ops[i]) {
+ return false;
+ }
+ }
- return ggml_metal_fusions;
+ return true;
+}
+
+static void ggml_metal_fusion_add_pattern_alloc_deps(
+ void * user_data,
+ void (*add_alloc_dep)(void *, ggml_tensor *, ggml_tensor *),
+ ggml_cgraph * gf,
+ const ggml_metal_fusion * fusion,
+ int node_idx) {
+ const int last_node = node_idx + (int) fusion->ops_all.size() - 1;
+
+ // keep all external inputs alive until the fused output
+ std::set<ggml_tensor *> seen;
+ for (int j = 0; j < (int) fusion->ops_all.size(); ++j) {
+ ggml_tensor * node = gf->nodes[node_idx + j];
+ for (int s = 0; s < GGML_MAX_SRC; ++s) {
+ ggml_tensor * src = node->src[s];
+ if (src && seen.insert(src).second) {
+ add_alloc_dep(user_data, src, 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 *),
+ ggml_cgraph * gf) {
+ for (int i = 0; i < gf->n_nodes; ++i) {
+ const ggml_metal_fusion * best = nullptr;
+ int best_raw = 0;
+
+ for (const ggml_metal_fusion & fusion : ggml_metal_fusions) {
+ if ((int) fusion.ops_all.size() <= best_raw) {
+ continue;
+ }
+ if (ggml_metal_fusion_match_raw_pattern(gf, i, fusion.ops_all)) {
+ best = &fusion;
+ best_raw = (int) fusion.ops_all.size();
+ }
+ }
+
+ if (best) {
+ ggml_metal_fusion_add_pattern_alloc_deps(user_data, add_alloc_dep, gf, best, i);
+ i += best_raw - 1;
+ }
+ }
}
// ---- shared fusion info ---------------------------------------------------
@@ -239,7 +759,7 @@ static std::string ggml_metal_fusion_label(const ggml_metal_fusion * fusion) {
GGML_ASSERT(fusion != nullptr);
std::string label;
- for (int j = 0; j < fusion->n_ops; j++) {
+ for (int j = 0; j < (int) fusion->ops.size(); j++) {
if (j > 0) {
label += '+';
}
@@ -271,90 +791,77 @@ struct ggml_metal_fusion_info * ggml_metal_fusion_info_init(bool enabled, int de
return finfo;
}
-void ggml_metal_fusion_info_free(struct ggml_metal_fusion_info * finfo) {
+void ggml_metal_fusion_info_free(ggml_metal_fusion_info * finfo) {
delete finfo;
}
-bool ggml_metal_fusion_info_enabled(const struct ggml_metal_fusion_info * finfo) {
+bool ggml_metal_fusion_info_enabled(const ggml_metal_fusion_info * finfo) {
return finfo->enabled;
}
-bool ggml_metal_fusion_info_stats(const struct ggml_metal_fusion_info * finfo) {
+bool ggml_metal_fusion_info_stats(const ggml_metal_fusion_info * finfo) {
return finfo->stats;
}
-int ggml_metal_fusion_info_debug(const struct ggml_metal_fusion_info * finfo) {
+int ggml_metal_fusion_info_debug(const ggml_metal_fusion_info * finfo) {
return finfo->debug;
}
-int ggml_metal_fusion_info_n_fusions(const struct ggml_metal_fusion_info * finfo) {
+int ggml_metal_fusion_info_n_fusions(const ggml_metal_fusion_info * finfo) {
return (int) finfo->labels.size();
}
-const char * ggml_metal_fusion_info_label(const struct ggml_metal_fusion_info * finfo, int idx) {
+const char * ggml_metal_fusion_info_label(const ggml_metal_fusion_info * finfo, int idx) {
GGML_ASSERT(idx >= 0 && idx < (int) finfo->labels.size());
return finfo->labels[idx].c_str();
}
-uint64_t ggml_metal_fusion_info_count(const struct ggml_metal_fusion_info * finfo, int idx) {
+uint64_t ggml_metal_fusion_info_count(const ggml_metal_fusion_info * finfo, int idx) {
GGML_ASSERT(idx >= 0 && idx < (int) finfo->counts.size());
return finfo->counts[idx];
}
-void ggml_metal_fusion_info_count_fusion(struct ggml_metal_fusion_info * finfo, const struct ggml_metal_fusion * fusion) {
+void ggml_metal_fusion_info_count_fusion(ggml_metal_fusion_info * finfo, const ggml_metal_fusion * fusion) {
if (!finfo->stats || fusion == nullptr) {
return;
}
- int n = 0;
- const ggml_metal_fusion * all = ggml_metal_fusion_all(&n);
-
- int idx = -1;
- for (int i = 0; i < n; i++) {
- if (&all[i] == fusion) {
- idx = i;
- break;
- }
- }
-
- if (idx >= 0 && idx < (int) finfo->counts.size()) {
+ const ptrdiff_t idx = fusion - ggml_metal_fusions.data();
+ if (idx >= 0 && idx < (ptrdiff_t) finfo->counts.size()) {
finfo->counts[idx]++;
}
}
-void ggml_metal_fusion_info_set_enabled(struct ggml_metal_fusion_info * finfo, bool enabled) {
+void ggml_metal_fusion_info_set_enabled(ggml_metal_fusion_info * finfo, bool enabled) {
finfo->enabled = enabled;
}
-void ggml_metal_fusion_info_labels_init(struct ggml_metal_fusion_info * finfo) {
+void ggml_metal_fusion_info_labels_init(ggml_metal_fusion_info * finfo) {
if (finfo->labels_set) {
return;
}
- int n = 0;
- const ggml_metal_fusion * all = ggml_metal_fusion_all(&n);
-
finfo->labels.clear();
- finfo->counts.assign(n, 0);
- finfo->labels.reserve(n);
+ finfo->counts.assign(ggml_metal_fusions.size(), 0);
+ finfo->labels.reserve(ggml_metal_fusions.size());
- for (int i = 0; i < n; i++) {
- finfo->labels.emplace_back(ggml_metal_fusion_label(&all[i]));
+ for (const ggml_metal_fusion & fusion : ggml_metal_fusions) {
+ finfo->labels.emplace_back(ggml_metal_fusion_label(&fusion));
}
finfo->labels_set = true;
}
-void ggml_metal_fusion_info_stats_init(struct ggml_metal_fusion_info * finfo) {
+void ggml_metal_fusion_info_stats_init(ggml_metal_fusion_info * finfo) {
finfo->stats = true;
ggml_metal_fusion_info_labels_init(finfo);
}
-void ggml_metal_fusion_info_stats_reset(struct ggml_metal_fusion_info * finfo) {
+void ggml_metal_fusion_info_stats_reset(ggml_metal_fusion_info * finfo) {
std::fill(finfo->counts.begin(), finfo->counts.end(), 0);
}
-int ggml_metal_fusion_info_stats_get(const struct ggml_metal_fusion_info * finfo, const char ** labels, uint64_t * counts, int n) {
+int ggml_metal_fusion_info_stats_get(const ggml_metal_fusion_info * finfo, const char ** labels, uint64_t * counts, int n) {
const int n_fusions = (int) finfo->labels.size();
if (labels == nullptr) {
@@ -372,6 +879,97 @@ int ggml_metal_fusion_info_stats_get(const struct ggml_metal_fusion_info * finfo
return n_fill;
}
+// ---- memory-range checks -------------------------------------------------
+
+// reject fusions where an external source overlaps any fused output. the fused
+// kernels elide intermediate nodes, so only sources that are not part of the
+// fused subgraph can cause read/write races with the output.
+static bool ggml_metal_fusion_check_memory_ranges(
+ const ggml_metal_fusion * fusion,
+ const ggml_tensor * const * nodes,
+ int node_count) {
+ // some fused kernels write through a tensor that also appears as a source (e.g. the gdn
+ // cache cpy), so a source that is the same memory as the output is not an external read
+ // source
+ auto same_memory = [](const ggml_tensor * a, const ggml_tensor * b) {
+ if (a->data && b->data && a->data == b->data) {
+ return true;
+ }
+ for (const ggml_tensor * v = a; v; v = v->view_src) {
+ if (v == b) {
+ return true;
+ }
+ }
+ for (const ggml_tensor * v = b; v; v = v->view_src) {
+ if (v == a) {
+ return true;
+ }
+ }
+ return false;
+ };
+
+ auto nodes_overlap = [](const ggml_tensor * a, const ggml_tensor * b) {
+ if (!a || !b || !a->data || !b->data || !a->buffer || !b->buffer) {
+ return false;
+ }
+
+ if (a->buffer != b->buffer) {
+ return false;
+ }
+
+ const int64_t a_start = (int64_t) a->data;
+ const int64_t a_end = a_start + ggml_backend_buft_get_alloc_size(a->buffer->buft, a);
+ const int64_t b_start = (int64_t) b->data;
+ const int64_t b_end = b_start + ggml_backend_buft_get_alloc_size(b->buffer->buft, b);
+
+ return (b_start <= a_start && a_start < b_end) ||
+ (a_start <= b_start && b_start < a_end);
+ };
+
+ auto is_intermediate = [](const ggml_tensor * src, const ggml_tensor * const * nodes, int j) {
+ for (int k = 0; k < j; ++k) {
+ if (src == nodes[k]) {
+ return true;
+ }
+ for (const ggml_tensor * view_src = src->view_src; view_src; view_src = view_src->view_src) {
+ if (view_src == nodes[k]) {
+ return true;
+ }
+ }
+ }
+ return false;
+ };
+
+ auto check_dst = [&](const ggml_tensor * dst) {
+ for (int j = 0; j < node_count; ++j) {
+ for (int s = 0; s < GGML_MAX_SRC; ++s) {
+ const ggml_tensor * src = nodes[j]->src[s];
+ if (!src || src->op == GGML_OP_NONE || same_memory(src, dst)) {
+ continue;
+ }
+
+ if (nodes_overlap(dst, src) && !is_intermediate(src, nodes, j)) {
+ return false;
+ }
+ }
+ }
+ return true;
+ };
+
+ if (!check_dst(nodes[node_count - 1])) {
+ return false;
+ }
+
+ for (int offset : fusion->outs) {
+ GGML_ASSERT(offset >= 0 && offset < node_count);
+ if (!check_dst(nodes[offset])) {
+ return false;
+ }
+ }
+
+ return true;
+}
+
// ---- queries -------------------------------------------------------------
// find the longest pattern matching the node sequence starting at idx
@@ -383,20 +981,17 @@ const ggml_metal_fusion * ggml_metal_fusion_next(
int idx,
ggml_metal_fusion_mode mode,
int * n_out) {
- int n = 0;
- const ggml_metal_fusion * all = ggml_metal_fusion_all(&n);
-
const ggml_metal_fusion * res = nullptr;
int best = 1;
- for (int i = 0; i < n; i++) {
- const ggml_metal_fusion * fusion = &all[i];
+ for (const ggml_metal_fusion & fusion : ggml_metal_fusions) {
+ const int n_ops = (int) fusion.ops.size();
// only look for a longer match than the current best
- if (fusion->n_ops <= best) {
+ if (n_ops <= best) {
continue;
}
- if (idx + fusion->n_ops > n_idxs) {
+ if (idx + n_ops > n_idxs) {
continue;
}
@@ -404,9 +999,9 @@ const ggml_metal_fusion * ggml_metal_fusion_next(
// the op sequence must match exactly
bool ok = true;
- for (int j = 0; j < fusion->n_ops; j++) {
+ for (int j = 0; j < n_ops; j++) {
nodes[j] = gf->nodes[node_idxs[idx + j]];
- if (nodes[j]->op != fusion->ops[j]) {
+ if (nodes[j]->op != fusion.ops[j]) {
ok = false;
break;
}
@@ -415,10 +1010,10 @@ const ggml_metal_fusion * ggml_metal_fusion_next(
continue;
}
- if (!fusion->unsafe) {
+ if (!fusion.unsafe) {
// common element-wise chain constraints: each node reads the previous one,
// and all nodes have the same shape
- for (int j = 1; j < fusion->n_ops && ok; j++) {
+ for (int j = 1; j < n_ops && ok; j++) {
if (nodes[j]->src[0] != nodes[j - 1] && nodes[j]->src[1] != nodes[j - 1]) {
ok = false;
break;
@@ -432,24 +1027,37 @@ const ggml_metal_fusion * ggml_metal_fusion_next(
continue;
}
- // all current fusions are single-output elision chains, so the last node is the only output
- // TODO: multi-output fusions: store pattern-relative offsets in the table and translate them here
- int outputs_buf[1];
- outputs_buf[0] = node_idxs[idx + fusion->n_ops - 1];
+ // primary output is the last node; additional outputs come from fusion.outs
+ int outputs_buf[GGML_METAL_FUSION_MAX];
+ outputs_buf[0] = node_idxs[idx + n_ops - 1];
+ for (size_t i = 0; i < fusion.outs.size(); ++i) {
+ const int out_offset = fusion.outs[i];
+ GGML_ASSERT(out_offset >= 0 && out_offset < n_ops);
+ outputs_buf[i + 1] = node_idxs[idx + out_offset];
+ }
+
+ const int n_outputs = 1 + (int) fusion.outs.size();
// structural subgraph checks (op sequence, elidable uses, view containment)
- if (!ggml_can_fuse_subgraph_ext(gf, node_idxs + idx, fusion->n_ops, fusion->ops, outputs_buf, 1)) {
+ if (!ggml_can_fuse_subgraph_ext(gf, node_idxs + idx, n_ops, fusion.ops.data(), outputs_buf, n_outputs)) {
continue;
}
}
// pattern-specific checks (the sole validator for unsafe patterns)
- if (fusion->check && !fusion->check(fusion, nodes, mode)) {
+ if (fusion.check && !fusion.check(&fusion, nodes, gf, node_idxs, idx, mode)) {
+ continue;
+ }
+
+ // the compute phase has allocated tensors and can detect aliasing between
+ // external sources and fused outputs; the optimizer phase cannot do this yet
+ if (mode == GGML_METAL_FUSION_FULL &&
+ !ggml_metal_fusion_check_memory_ranges(&fusion, nodes, n_ops)) {
continue;
}
- best = fusion->n_ops;
- res = fusion;
+ best = n_ops;
+ res = &fusion;
}
*n_out = best;
diff --git a/ggml/src/ggml-metal/ggml-metal-fusion.h b/ggml/src/ggml-metal/ggml-metal-fusion.h
index e8515bdec..6b139a69b 100644
--- a/ggml/src/ggml-metal/ggml-metal-fusion.h
+++ b/ggml/src/ggml-metal/ggml-metal-fusion.h
@@ -32,33 +32,27 @@ typedef enum ggml_metal_fusion_id {
GGML_METAL_FUSION_NONE = 0,
GGML_METAL_FUSION_NORM_MUL, // NORM/RMS_NORM + MUL
GGML_METAL_FUSION_NORM_MUL_ADD, // NORM/RMS_NORM + MUL + ADD
+ GGML_METAL_FUSION_NORM_SCALE, // NORM/RMS_NORM + SCALE
GGML_METAL_FUSION_ADD_CHAIN, // ADD x N (N in [2, 7])
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_SSM_CONV_SILU, // SSM_CONV + UNARY (silu)
} ggml_metal_fusion_id;
-struct ggml_metal_fusion {
- ggml_metal_fusion_id id;
-
- const enum ggml_op * ops; // op sequence (fixed length)
- int n_ops; // number of 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)
- 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_metal_fusion * fusion,
- const struct ggml_tensor * const * nodes,
- ggml_metal_fusion_mode mode);
-};
+struct ggml_metal_fusion; // defined in ggml-metal-fusion.cpp
typedef struct ggml_metal_fusion ggml_metal_fusion;
-// the single table of all fusions supported by the Metal backend
-const ggml_metal_fusion * ggml_metal_fusion_all(int * n);
+// access the fusion identifier without exposing the full pattern definition
+ggml_metal_fusion_id ggml_metal_fusion_get_id(const struct ggml_metal_fusion * fusion);
+
+// 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 *),
+ struct ggml_cgraph * gf);
// ---- shared fusion info ---------------------------------------------------
diff --git a/ggml/src/ggml-metal/ggml-metal-impl.h b/ggml/src/ggml-metal/ggml-metal-impl.h
index d84ca937b..eaa4278db 100644
--- a/ggml/src/ggml-metal/ggml-metal-impl.h
+++ b/ggml/src/ggml-metal/ggml-metal-impl.h
@@ -116,6 +116,9 @@
#define FC_SUM_ROWS 1400
#define FC_UPSCALE 1500
#define FC_GATED_DELTA_NET 1600
+#define FC_NORM 1700
+#define FC_TOPK_MOE 1800
+#define FC_MOE_REDUCE 1900
// op-specific constants
#define OP_FLASH_ATTN_EXT_NQPSG 8
@@ -622,6 +625,7 @@ typedef struct {
uint64_t nbf1[3];
uint64_t nbf2[3];
uint64_t nbf3[3];
+ float scale;
} ggml_metal_kargs_norm;
typedef struct {
@@ -910,7 +914,6 @@ typedef struct {
uint64_t nb00;
uint64_t nb01;
uint64_t nb02;
- int64_t ne10;
int64_t ne11;
uint64_t nb10;
uint64_t nb11;
@@ -1231,6 +1234,19 @@ typedef struct {
int32_t top_k; // k
} ggml_metal_kargs_top_k;
+typedef struct {
+ int32_t ne01; // n_tokens
+ uint64_t nb01; // logits row stride
+ uint64_t nb1_ids; // ids row stride
+ float clamp;
+ float scale;
+} ggml_metal_kargs_topk_moe;
+
+typedef struct {
+ int32_t ne00; // n_embd
+ int32_t ne02; // n_tokens
+} ggml_metal_kargs_moe_reduce;
+
typedef struct {
int32_t nrows;
} ggml_metal_kargs_fwht;
diff --git a/ggml/src/ggml-metal/ggml-metal-ops.cpp b/ggml/src/ggml-metal/ggml-metal-ops.cpp
index 77c399bdb..c86a74236 100644
--- a/ggml/src/ggml-metal/ggml-metal-ops.cpp
+++ b/ggml/src/ggml-metal/ggml-metal-ops.cpp
@@ -226,7 +226,16 @@ static int ggml_metal_op_encode_impl(ggml_metal_op_t ctx, int idx) {
// otherwise, we add the new ranges to the encoding context and process the node concurrently
//
{
- const bool is_concurrent = ggml_metal_op_concurrency_check(ctx, node);
+ bool is_concurrent = ggml_metal_op_concurrency_check(ctx, node);
+
+ if (is_concurrent && ctx->use_fusion()) {
+ int n_fuse = 1;
+ const ggml_metal_fusion * fusion = ctx->can_fuse(idx, GGML_METAL_FUSION_FULL, &n_fuse);
+ if (fusion) {
+ // fused kernels write to the last node of the group, not necessarily to the first node's dst
+ is_concurrent = ggml_mem_ranges_check(ctx->mem_ranges, ctx->node(idx + n_fuse - 1));
+ }
+ }
if (!is_concurrent) {
ggml_metal_op_concurrency_reset(ctx);
@@ -1540,6 +1549,14 @@ int ggml_metal_op_dsv4_hc(ggml_metal_op_t ctx, int idx) {
int ggml_metal_op_soft_max(ggml_metal_op_t ctx, int idx) {
ggml_tensor * op = ctx->node(idx);
+ if (ctx->use_fusion()) {
+ int n = 1;
+ const ggml_metal_fusion * fusion = ctx->can_fuse(idx, GGML_METAL_FUSION_FULL, &n);
+ if (fusion && ggml_metal_fusion_get_id(fusion) == GGML_METAL_FUSION_TOPK_MOE) {
+ return ggml_metal_op_topk_moe(ctx, idx);
+ }
+ }
+
ggml_metal_library_t lib = ctx->lib;
ggml_metal_encoder_t enc = ctx->enc;
@@ -1640,6 +1657,20 @@ int ggml_metal_op_ssm_conv(ggml_metal_op_t ctx, int idx) {
GGML_TENSOR_LOCALS( int32_t, ne, op, ne);
GGML_TENSOR_LOCALS(uint64_t, nb, op, nb);
+ int n_fuse = 1;
+ bool use_silu = false;
+
+ if (ctx->use_fusion()) {
+ int n = 1;
+ const ggml_metal_fusion * fusion = ctx->can_fuse(idx, GGML_METAL_FUSION_FULL, &n);
+ if (fusion && ggml_metal_fusion_get_id(fusion) == GGML_METAL_FUSION_SSM_CONV_SILU) {
+ n_fuse = n;
+ use_silu = true;
+
+ ctx->count_fusions(fusion);
+ }
+ }
+
ggml_metal_kargs_ssm_conv args = {
/*.ne00 =*/ ne00,
/*.ne01 =*/ ne01,
@@ -1647,7 +1678,6 @@ int ggml_metal_op_ssm_conv(ggml_metal_op_t ctx, int idx) {
/*.nb00 =*/ nb00,
/*.nb01 =*/ nb01,
/*.nb02 =*/ nb02,
- /*.ne10 =*/ ne10,
/*.ne11 =*/ ne11,
/*.nb10 =*/ nb10,
/*.nb11 =*/ nb11,
@@ -1659,6 +1689,8 @@ int ggml_metal_op_ssm_conv(ggml_metal_op_t ctx, int idx) {
/*.nb2 =*/ nb2,
};
+ const ggml_metal_buffer_id bid_dst = ggml_metal_get_buffer_id(n_fuse > 1 ? ctx->node(idx + n_fuse - 1) : op);
+
// Use batched kernel for prefill (ne1 > 1) to reduce threadgroup dispatch overhead
const bool use_batched = (ne1 > 1);
@@ -1673,31 +1705,35 @@ int ggml_metal_op_ssm_conv(ggml_metal_op_t ctx, int idx) {
else if (ne1 > 4 ) BATCH_SIZE = 8;
else BATCH_SIZE = 2;
- auto pipeline = ggml_metal_library_get_pipeline_ssm_conv_batched(lib, op, BATCH_SIZE);
+ auto pipeline = ggml_metal_library_get_pipeline_ssm_conv_batched(lib, op, BATCH_SIZE, (int32_t) ne10, use_silu);
ggml_metal_encoder_set_pipeline(enc, pipeline);
ggml_metal_encoder_set_bytes(enc, &args, sizeof(args), 0);
ggml_metal_encoder_set_buffer(enc, ggml_metal_get_buffer_id(op->src[0]), 1);
ggml_metal_encoder_set_buffer(enc, ggml_metal_get_buffer_id(op->src[1]), 2);
- ggml_metal_encoder_set_buffer(enc, ggml_metal_get_buffer_id(op), 3);
+ ggml_metal_encoder_set_buffer(enc, bid_dst, 3);
// Dispatch: ne01 rows, ceil(ne1/BATCH_SIZE) token batches, ne02 sequences
// Each threadgroup has BATCH_SIZE threads, each handling one token
const int n_token_batches = (ne1 + BATCH_SIZE - 1) / BATCH_SIZE;
ggml_metal_encoder_dispatch_threadgroups(enc, ne01, n_token_batches, ne02, BATCH_SIZE, 1, 1);
} else {
- auto pipeline = ggml_metal_library_get_pipeline_ssm_conv(lib, op);
+ auto pipeline = ggml_metal_library_get_pipeline_ssm_conv(lib, op, (int32_t) ne10, use_silu);
ggml_metal_encoder_set_pipeline(enc, pipeline);
ggml_metal_encoder_set_bytes(enc, &args, sizeof(args), 0);
ggml_metal_encoder_set_buffer(enc, ggml_metal_get_buffer_id(op->src[0]), 1);
ggml_metal_encoder_set_buffer(enc, ggml_metal_get_buffer_id(op->src[1]), 2);
- ggml_metal_encoder_set_buffer(enc, ggml_metal_get_buffer_id(op), 3);
+ ggml_metal_encoder_set_buffer(enc, bid_dst, 3);
ggml_metal_encoder_dispatch_threadgroups(enc, ne01, ne1, ne02, 1, 1, 1);
}
- return 1;
+ if (n_fuse > 1 && ggml_metal_fusion_info_debug(ctx->finfo) > 1) {
+ GGML_LOG_DEBUG("%s: fuse: SSM_CONV + UNARY\n", __func__);
+ }
+
+ return n_fuse;
}
int ggml_metal_op_ssm_scan(ggml_metal_op_t ctx, int idx) {
@@ -1904,7 +1940,7 @@ int ggml_metal_op_gated_delta_net(ggml_metal_op_t ctx, int idx) {
int n = 1;
const ggml_metal_fusion * fusion = ctx->can_fuse(idx, GGML_METAL_FUSION_FULL, &n);
- if (fusion && fusion->id == GGML_METAL_FUSION_GDN_CACHE) {
+ if (fusion && ggml_metal_fusion_get_id(fusion) == GGML_METAL_FUSION_GDN_CACHE) {
const ggml_tensor * dst_cache = ctx->node(idx + 1)->src[1]; // cache view
bid_out = ggml_metal_get_buffer_id(dst_cache);
@@ -3821,10 +3857,16 @@ int ggml_metal_op_bin(ggml_metal_op_t ctx, int idx) {
n_fuse = n;
// snake activation autofuse: mul -> sin -> sqr -> mul -> add
- if (fusion && fusion->id == GGML_METAL_FUSION_SNAKE) {
+ if (fusion && ggml_metal_fusion_get_id(fusion) == GGML_METAL_FUSION_SNAKE) {
ctx->count_fusions(fusion);
return ggml_metal_op_snake_fused(ctx, idx);
}
+
+ // MoE output reduction: experts * weights -> weighted sum
+ if (fusion && ggml_metal_fusion_get_id(fusion) == GGML_METAL_FUSION_MOE_REDUCE) {
+ ctx->count_fusions(fusion);
+ return ggml_metal_op_moe_reduce(ctx, idx);
+ }
}
ggml_tensor * op = ctx->node(idx);
@@ -3883,7 +3925,7 @@ int ggml_metal_op_bin(ggml_metal_op_t ctx, int idx) {
// c[1] = add(c[0], b[1])
// c[2] = add(c[1], b[2])
// ...
- if (use_fusion && fusion && fusion->id == GGML_METAL_FUSION_ADD_CHAIN) {
+ if (use_fusion && fusion && ggml_metal_fusion_get_id(fusion) == GGML_METAL_FUSION_ADD_CHAIN) {
// the offsets of the fused addends are relative to the start of the src1 buffer
for (int i = 1; i < n_fuse; i++) {
args.o1[i] = ggml_metal_get_buffer_id(ctx->node(idx + i)->src[1]).offs;
@@ -4127,20 +4169,22 @@ int ggml_metal_op_norm(ggml_metal_op_t ctx, int idx) {
/*.nbf1 =*/ { nb01 },
/*.nbf2 =*/ { nb02 },
/*.nbf3 =*/ { nb03 },
+ /*.scale =*/ 1.0f,
};
int n_fuse = 1;
+ bool fused_norm_scale = false;
ggml_metal_buffer_id bid_fuse[2] = { bid_src0, bid_src0 };
// d[0] = norm(a)
- // d[1] = mul(d[0], b)
+ // d[1] = mul(d[0], b) or scale(d[0])
// d[2] = add(d[1], c)
if (use_fusion) {
int n = 1;
const ggml_metal_fusion * fusion = ctx->can_fuse(idx, GGML_METAL_FUSION_FULL, &n);
- if (fusion && (fusion->id == GGML_METAL_FUSION_NORM_MUL || fusion->id == GGML_METAL_FUSION_NORM_MUL_ADD)) {
+ if (fusion && (ggml_metal_fusion_get_id(fusion) == GGML_METAL_FUSION_NORM_MUL || ggml_metal_fusion_get_id(fusion) == GGML_METAL_FUSION_NORM_MUL_ADD)) {
n_fuse = n;
ctx->count_fusions(fusion);
@@ -4168,6 +4212,20 @@ int ggml_metal_op_norm(ggml_metal_op_t ctx, int idx) {
}
}
}
+
+ if (fusion && ggml_metal_fusion_get_id(fusion) == GGML_METAL_FUSION_NORM_SCALE) {
+ n_fuse = n;
+ fused_norm_scale = true;
+
+ ctx->count_fusions(fusion);
+
+ const ggml_tensor * scale_node = ctx->node(idx + 1);
+ args.scale = ggml_get_op_params_f32(scale_node, 0);
+
+ if (debug_fusion > 1) {
+ GGML_LOG_DEBUG("%s: fuse: %s + SCALE\n", __func__, ggml_op_name(op->op));
+ }
+ }
}
if (n_fuse > 1) {
@@ -4182,7 +4240,9 @@ int ggml_metal_op_norm(ggml_metal_op_t ctx, int idx) {
}
}
- auto pipeline = ggml_metal_library_get_pipeline_norm(lib, op, n_fuse);
+ auto pipeline = fused_norm_scale ?
+ ggml_metal_library_get_pipeline_norm_scale(lib, op) :
+ ggml_metal_library_get_pipeline_norm(lib, op, n_fuse);
int nth = 32; // SIMD width
@@ -5403,6 +5463,110 @@ static void ggml_metal_op_top_k_radix(ggml_metal_op_t ctx, int idx) {
ggml_metal_encoder_dispatch_threadgroups(enc, ne01, ne02, ne03, nth, 1, 1);
}
+int ggml_metal_op_topk_moe(ggml_metal_op_t ctx, int idx) {
+ ggml_metal_library_t lib = ctx->lib;
+ ggml_metal_encoder_t enc = ctx->enc;
+
+ int n_fuse = 1;
+ const ggml_metal_fusion * fusion = ctx->can_fuse(idx, GGML_METAL_FUSION_FULL, &n_fuse);
+ if (!fusion || ggml_metal_fusion_get_id(fusion) != GGML_METAL_FUSION_TOPK_MOE) {
+ return 1;
+ }
+
+ ggml_tensor * softmax = ctx->node(idx);
+ ggml_tensor * logits = softmax->src[0];
+ ggml_tensor * get_rows = ctx->node(idx + 2);
+ ggml_tensor * ids = get_rows->src[1];
+ ggml_tensor * weights = ctx->node(idx + n_fuse - 1);
+
+ const int64_t n_expert = logits->ne[0];
+ const int64_t n_tokens = logits->ne[1];
+ const int64_t n_expert_used = ids->ne[0];
+
+ const bool with_norm = n_fuse >= 6;
+ const bool with_scale = n_fuse == 4 || n_fuse == 7;
+
+ float clamp = -INFINITY;
+ if (with_norm) {
+ ggml_tensor * clamp_node = ctx->node(idx + 4);
+ clamp = ggml_get_op_params_f32(clamp_node, 0);
+ }
+
+ float scale = 1.0f;
+ if (with_scale) {
+ ggml_tensor * scale_node = ctx->node(idx + n_fuse - 1);
+ scale = ggml_get_op_params_f32(scale_node, 0);
+ }
+
+ ggml_metal_kargs_topk_moe args = {
+ /*.ne01 =*/ (int32_t) n_tokens,
+ /*.nb01 =*/ logits->nb[1],
+ /*.nb1_ids =*/ ids->nb[1],
+ /*.clamp =*/ clamp,
+ /*.scale =*/ scale,
+ };
+
+ auto pipeline = ggml_metal_library_get_pipeline_topk_moe(lib, (int32_t) n_expert, (int32_t) n_expert_used, with_norm);
+
+ ggml_metal_encoder_set_pipeline(enc, pipeline);
+ ggml_metal_encoder_set_bytes (enc, &args, sizeof(args), 0);
+ ggml_metal_encoder_set_buffer (enc, ggml_metal_get_buffer_id(logits), 1);
+ ggml_metal_encoder_set_buffer (enc, ggml_metal_get_buffer_id(weights), 2);
+ ggml_metal_encoder_set_buffer (enc, ggml_metal_get_buffer_id(ids), 3);
+
+ ggml_metal_encoder_dispatch_threadgroups(enc, (uint32_t) n_tokens, 1, 1, 32, 1, 1);
+
+ ctx->count_fusions(fusion);
+
+ if (ggml_metal_fusion_info_debug(ctx->finfo) > 1) {
+ GGML_LOG_DEBUG("%s: fuse: SOFT_MAX + ARGSORT + GET_ROWS\n", __func__);
+ }
+
+ return n_fuse;
+}
+
+int ggml_metal_op_moe_reduce(ggml_metal_op_t ctx, int idx) {
+ ggml_metal_library_t lib = ctx->lib;
+ ggml_metal_encoder_t enc = ctx->enc;
+
+ int n_fuse = 1;
+ const ggml_metal_fusion * fusion = ctx->can_fuse(idx, GGML_METAL_FUSION_FULL, &n_fuse);
+ if (!fusion || ggml_metal_fusion_get_id(fusion) != GGML_METAL_FUSION_MOE_REDUCE) {
+ return 1;
+ }
+
+ ggml_tensor * mul = ctx->node(idx);
+ ggml_tensor * experts = mul->src[0];
+ ggml_tensor * weights = mul->src[1];
+ ggml_tensor * dst = ctx->node(idx + n_fuse - 1);
+
+ ggml_metal_kargs_moe_reduce args = {
+ /*.ne00 =*/ (int32_t) experts->ne[0],
+ /*.ne02 =*/ (int32_t) experts->ne[2],
+ };
+
+ auto pipeline = ggml_metal_library_get_pipeline_moe_reduce(lib, (int32_t) experts->ne[1]);
+
+ const int nth = std::min(256, ggml_metal_pipeline_max_theads_per_threadgroup(pipeline));
+ const int n_col_tiles = (args.ne00 + nth - 1) / nth;
+
+ ggml_metal_encoder_set_pipeline(enc, pipeline);
+ ggml_metal_encoder_set_bytes (enc, &args, sizeof(args), 0);
+ ggml_metal_encoder_set_buffer (enc, ggml_metal_get_buffer_id(experts), 1);
+ ggml_metal_encoder_set_buffer (enc, ggml_metal_get_buffer_id(weights), 2);
+ ggml_metal_encoder_set_buffer (enc, ggml_metal_get_buffer_id(dst), 3);
+
+ ggml_metal_encoder_dispatch_threadgroups(enc, (uint32_t) args.ne02, (uint32_t) n_col_tiles, 1, nth, 1, 1);
+
+ ctx->count_fusions(fusion);
+
+ if (ggml_metal_fusion_info_debug(ctx->finfo) > 1) {
+ GGML_LOG_DEBUG("%s: fuse: MOE_REDUCE\n", __func__);
+ }
+
+ return n_fuse;
+}
+
int ggml_metal_op_top_k(ggml_metal_op_t ctx, int idx) {
ggml_tensor * op = ctx->node(idx);
diff --git a/ggml/src/ggml-metal/ggml-metal-ops.h b/ggml/src/ggml-metal/ggml-metal-ops.h
index ae72e8820..583d1156b 100644
--- a/ggml/src/ggml-metal/ggml-metal-ops.h
+++ b/ggml/src/ggml-metal/ggml-metal-ops.h
@@ -98,6 +98,8 @@ int ggml_metal_op_timestep_embedding(ggml_metal_op_t ctx, int idx);
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_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);
diff --git a/ggml/src/ggml-metal/ggml-metal.cpp b/ggml/src/ggml-metal/ggml-metal.cpp
index 4f9440f9e..c6c8ce836 100644
--- a/ggml/src/ggml-metal/ggml-metal.cpp
+++ b/ggml/src/ggml-metal/ggml-metal.cpp
@@ -562,7 +562,11 @@ static void ggml_backend_metal_event_wait(ggml_backend_t backend, ggml_backend_e
}
static void ggml_backend_metal_graph_optimize(ggml_backend_t backend, ggml_cgraph * cgraph, ggml_backend_graph_optimize_params * params) {
- GGML_UNUSED(params);
+ GGML_ASSERT(params && params->add_alloc_dep);
+
+ // 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
+ ggml_metal_fusion_add_alloc_deps(params->user_data, params->add_alloc_dep, cgraph);
ggml_metal_t ctx = (ggml_metal_t)backend->context;
diff --git a/ggml/src/ggml-metal/kernels/argsort.metal b/ggml/src/ggml-metal/kernels/argsort.metal
index e81d194c3..5231b8395 100644
--- a/ggml/src/ggml-metal/kernels/argsort.metal
+++ b/ggml/src/ggml-metal/kernels/argsort.metal
@@ -1,5 +1,11 @@
#include "common.h"
+constant bool FC_topk_moe_with_norm [[function_constant(FC_TOPK_MOE + 0)]];
+constant int FC_topk_moe_n_expert [[function_constant(FC_TOPK_MOE + 1)]];
+constant int FC_topk_moe_top_k [[function_constant(FC_TOPK_MOE + 2)]];
+
+constant int FC_moe_reduce_n_expert_used [[function_constant(FC_MOE_REDUCE + 0)]];
+
// bitonic sort implementation following the CUDA kernels as reference
typedef void (argsort_t)(
constant ggml_metal_kargs_argsort & args,
@@ -335,3 +341,139 @@ kernel void kernel_top_k_f32_i32(
}
}
}
+
+// fused SOFT_MAX + top-k + GET_ROWS (+ optional norm/scale) for MoE routing.
+// One SIMDgroup handles one token row; n_expert is limited to 1024 by the host.
+kernel void kernel_topk_moe_f32(
+ constant ggml_metal_kargs_topk_moe & args,
+ device const char * src0,
+ device float * weights,
+ device int32_t * ids,
+ uint3 tgpig[[threadgroup_position_in_grid]],
+ ushort tiisg[[thread_index_in_simdgroup]]) {
+ const int row = (int) tgpig.x;
+ if (row >= args.ne01) {
+ return;
+ }
+
+ const int n_expert = FC_topk_moe_n_expert;
+ const int top_k = FC_topk_moe_top_k;
+ const int lane = (int) tiisg;
+ const int n_per_lane = (n_expert + 31) / 32;
+
+ device const float * logits_row = (device const float *) (src0 + row * args.nb01);
+ device float * weights_row = weights + row * top_k;
+ device int32_t * ids_row = ids + row * (args.nb1_ids / sizeof(int32_t));
+
+ float wt[32];
+ float output_weights[32];
+ FOR_UNROLL (int i = 0; i < 32; ++i) {
+ wt[i] = -INFINITY;
+ output_weights[i] = 0.0f;
+ }
+
+ for (int i = lane; i < n_expert; i += 32) {
+ const float v = logits_row[i];
+ wt[i / 32] = isnan(v) ? -FLT_MAX : v;
+ }
+
+ // softmax over the expert logits
+ float max_val = -INFINITY;
+ FOR_UNROLL (int i = 0; i < n_per_lane; ++i) {
+ max_val = max(max_val, wt[i]);
+ }
+ max_val = simd_max(max_val);
+
+ float sum_val = 0.0f;
+ FOR_UNROLL (int i = 0; i < n_per_lane; ++i) {
+ wt[i] = exp(wt[i] - max_val);
+ sum_val += wt[i];
+ }
+ sum_val = simd_sum(sum_val);
+
+ const float inv_sum = 1.0f / sum_val;
+ FOR_UNROLL (int i = 0; i < n_per_lane; ++i) {
+ wt[i] *= inv_sum;
+ }
+
+ float wt_sum = 0.0f;
+
+ for (int k = 0; k < top_k; ++k) {
+ float best_val = -INFINITY;
+ int best_expert = -1;
+
+ FOR_UNROLL (int i = 0; i < n_per_lane; ++i) {
+ const int expert = lane + i * 32;
+ if (expert < n_expert && (wt[i] > best_val || (wt[i] == best_val && expert < best_expert))) {
+ best_val = wt[i];
+ best_expert = expert;
+ }
+ }
+
+ FOR_UNROLL (int mask = 16; mask > 0; mask >>= 1) {
+ const float val = simd_shuffle_xor(best_val, mask);
+ const int expert = simd_shuffle_xor(best_expert, mask);
+ if (val > best_val || (val == best_val && expert < best_expert)) {
+ best_val = val;
+ best_expert = expert;
+ }
+ }
+
+ if ((best_expert & 31) == lane) {
+ wt[best_expert / 32] = -INFINITY;
+ }
+
+ if ((k & 31) == lane) {
+ output_weights[k / 32] = best_val;
+ }
+
+ if ((best_expert & 31) == lane) {
+ ids_row[k] = best_expert;
+ if (FC_topk_moe_with_norm) {
+ wt_sum += best_val;
+ }
+ }
+ }
+
+ if (FC_topk_moe_with_norm) {
+ wt_sum = simd_sum(wt_sum);
+ wt_sum = max(wt_sum, args.clamp);
+ const float inv = 1.0f / wt_sum;
+ FOR_UNROLL (int i = 0; i < n_per_lane; ++i) {
+ output_weights[i] *= inv;
+ }
+ }
+
+ FOR_UNROLL (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;
+ }
+ }
+}
+
+// fused MoE expert weighting + reduction: weighted = sum(experts[e] * weights[e]).
+// The host guarantees all tensors are contiguous F32.
+kernel void kernel_moe_reduce_f32(
+ constant ggml_metal_kargs_moe_reduce & args,
+ device const float * experts,
+ device const float * weights,
+ device float * dst,
+ uint3 tgpig[[threadgroup_position_in_grid]],
+ ushort3 tpitg[[thread_position_in_threadgroup]],
+ ushort3 ntg[[threads_per_threadgroup]]) {
+ const int64_t token = tgpig.x;
+ const int64_t col = (int64_t) tgpig.y * ntg.x + tpitg.x;
+ if (token >= args.ne02 || col >= args.ne00) {
+ return;
+ }
+
+ const int n_expert_used = FC_moe_reduce_n_expert_used;
+
+ const int64_t base = token * (int64_t) n_expert_used * args.ne00 + col;
+ float sum = 0.0f;
+ FOR_UNROLL (int e = 0; e < n_expert_used; ++e) {
+ sum += experts[base + e * args.ne00] * weights[token * n_expert_used + e];
+ }
+ dst[token * args.ne00 + col] = sum;
+}
diff --git a/ggml/src/ggml-metal/kernels/norm.metal b/ggml/src/ggml-metal/kernels/norm.metal
index 7e42389fe..c76d0c7ff 100644
--- a/ggml/src/ggml-metal/kernels/norm.metal
+++ b/ggml/src/ggml-metal/kernels/norm.metal
@@ -1,5 +1,7 @@
#include "common.h"
+constant bool FC_norm_use_scale [[function_constant(FC_NORM + 0)]];
+
// F == 1 : norm (no fuse)
// F == 2 : norm + mul
// F == 3 : norm + mul + add
@@ -80,7 +82,11 @@ kernel void kernel_norm_fuse_impl(
y[i00] = (y[i00]*scale);
}
if (F == 2) {
- y[i00] = (y[i00]*scale)*f0[i00];
+ if (FC_norm_use_scale) {
+ y[i00] = (y[i00]*scale) * args.scale;
+ } else {
+ y[i00] = (y[i00]*scale)*f0[i00];
+ }
}
if (F == 3) {
y[i00] = (y[i00]*scale)*f0[i00] + f1[i00];
@@ -155,7 +161,11 @@ kernel void kernel_rms_norm_fuse_impl(
y[i00] = (x[i00]*scale);
}
if (F == 2) {
- y[i00] = (x[i00]*scale)*f0[i00];
+ if (FC_norm_use_scale) {
+ y[i00] = (x[i00]*scale) * args.scale;
+ } else {
+ y[i00] = (x[i00]*scale)*f0[i00];
+ }
}
if (F == 3) {
y[i00] = (x[i00]*scale)*f0[i00] + f1[i00];
diff --git a/ggml/src/ggml-metal/kernels/ssm.metal b/ggml/src/ggml-metal/kernels/ssm.metal
index d3118a831..b21c53b74 100644
--- a/ggml/src/ggml-metal/kernels/ssm.metal
+++ b/ggml/src/ggml-metal/kernels/ssm.metal
@@ -1,5 +1,8 @@
#include "common.h"
+constant bool FC_ssm_conv_silu [[function_constant(FC_SSM_CONV + 1)]];
+constant int FC_ssm_conv_nc [[function_constant(FC_SSM_CONV + 2)]];
+
// ref: ggml.c:ggml_compute_forward_ssm_conv_f32
kernel void kernel_ssm_conv_f32_f32(
constant ggml_metal_kargs_ssm_conv & args,
@@ -13,7 +16,7 @@ kernel void kernel_ssm_conv_f32_f32(
const int64_t i2 = tgpig.y;
const int64_t i3 = tgpig.z;
- const int64_t nc = args.ne10;
+ const int64_t nc = FC_ssm_conv_nc;
//const int64_t ncs = args.ne00;
//const int64_t nr = args.ne01;
//const int64_t n_t = args.ne1;
@@ -25,11 +28,11 @@ kernel void kernel_ssm_conv_f32_f32(
float sumf = 0.0f;
- for (int64_t i0 = 0; i0 < nc; ++i0) {
+ FOR_UNROLL (int64_t i0 = 0; i0 < nc; ++i0) {
sumf += s[i0] * c[i0];
}
- x[0] = sumf;
+ x[0] = FC_ssm_conv_silu ? sumf/(1.0f + exp(-sumf)) : sumf;
}
kernel void kernel_ssm_conv_f32_f32_4(
@@ -44,7 +47,7 @@ kernel void kernel_ssm_conv_f32_f32_4(
const int64_t i2 = tgpig.y;
const int64_t i3 = tgpig.z;
- const int64_t nc = args.ne10;
+ const int64_t nc = FC_ssm_conv_nc;
//const int64_t ncs = args.ne00;
//const int64_t nr = args.ne01;
//const int64_t n_t = args.ne1;
@@ -56,11 +59,11 @@ kernel void kernel_ssm_conv_f32_f32_4(
float sumf = 0.0f;
- for (int64_t i0 = 0; i0 < nc/4; ++i0) {
+ FOR_UNROLL (int64_t i0 = 0; i0 < nc/4; ++i0) {
sumf += dot(s[i0], c[i0]);
}
- x[0] = sumf;
+ x[0] = FC_ssm_conv_silu ? sumf/(1.0f + exp(-sumf)) : sumf;
}
constant short FC_ssm_conv_bs [[function_constant(FC_SSM_CONV + 0)]];
@@ -87,7 +90,7 @@ kernel void kernel_ssm_conv_f32_f32_batched(
const int64_t i2_off = tpitg.x;
const int64_t i2 = i2_base + i2_off;
- const int64_t nc = args.ne10; // conv kernel size (typically 4)
+ const int64_t nc = FC_ssm_conv_nc; // conv kernel size (typically 4)
const int64_t n_t = args.ne1; // number of tokens
// Bounds check for partial batches at the end
@@ -105,11 +108,11 @@ kernel void kernel_ssm_conv_f32_f32_batched(
device float * x = (device float *) ((device char *) dst + ir*args.nb0 + i2*args.nb1 + i3*args.nb2);
float sumf = 0.0f;
- for (int64_t i0 = 0; i0 < nc; ++i0) {
+ FOR_UNROLL (int64_t i0 = 0; i0 < nc; ++i0) {
sumf += s[i0] * c[i0];
}
- x[0] = sumf;
+ x[0] = FC_ssm_conv_silu ? sumf/(1.0f + exp(-sumf)) : sumf;
}
kernel void kernel_ssm_conv_f32_f32_batched_4(
@@ -132,7 +135,7 @@ kernel void kernel_ssm_conv_f32_f32_batched_4(
const int64_t i2_off = tpitg.x;
const int64_t i2 = i2_base + i2_off;
- const int64_t nc = args.ne10; // conv kernel size (typically 4)
+ const int64_t nc = FC_ssm_conv_nc; // conv kernel size (typically 4)
const int64_t n_t = args.ne1; // number of tokens
// Bounds check for partial batches at the end
@@ -150,11 +153,11 @@ kernel void kernel_ssm_conv_f32_f32_batched_4(
device float * x = (device float *) ((device char *) dst + ir*args.nb0 + i2*args.nb1 + i3*args.nb2);
float sumf = 0.0f;
- for (int64_t i0 = 0; i0 < nc/4; ++i0) {
+ FOR_UNROLL (int64_t i0 = 0; i0 < nc/4; ++i0) {
sumf += dot(s[i0], c[i0]);
}
- x[0] = sumf;
+ x[0] = FC_ssm_conv_silu ? sumf/(1.0f + exp(-sumf)) : sumf;
}
// ref: ggml.c:ggml_compute_forward_ssm_scan_f32, Mamba-2 part
diff --git a/tests/fusion/MTL.csv b/tests/fusion/MTL.csv
index 067316abf..3970da970 100644
--- a/tests/fusion/MTL.csv
+++ b/tests/fusion/MTL.csv
@@ -1,37 +1,68 @@
# test-fusion baseline for device MTL
# arch ,moe ,mode ,label , count
+afmoe ,1 ,any ,MUL+ADD , 2
+afmoe ,1 ,any ,RMS_NORM+MUL , 10
+afmoe ,1 ,any ,RMS_NORM+MUL+ADD , 3
arcee ,0 ,any ,RMS_NORM+MUL , 5
+arctic ,0 ,any ,MUL+ADD , 4
arctic ,0 ,any ,RMS_NORM+MUL , 7
+arctic ,0 ,any ,SOFT_MAX+ARGSORT+GET_ROWS+SUM_ROWS+CLAMP+DIV, 2
baichuan ,0 ,any ,RMS_NORM+MUL , 5
bailingmoe ,1 ,any ,ADD+ADD , 2
+bailingmoe ,1 ,any ,MUL+ADD , 4
bailingmoe ,1 ,any ,RMS_NORM+MUL , 5
+bailingmoe ,1 ,any ,SOFT_MAX+ARGSORT+GET_ROWS , 2
bailingmoe2 ,1 ,any ,ADD+ADD , 1
+bailingmoe2 ,1 ,any ,MUL+ADD , 2
bailingmoe2 ,1 ,any ,RMS_NORM+MUL , 9
bailingmoe3 ,1 ,any ,ADD+ADD , 1
bailingmoe3 ,1 ,any ,GATED_DELTA_NET+CPY , 1
+bailingmoe3 ,1 ,any ,MUL+ADD , 2
bailingmoe3 ,1 ,any ,RMS_NORM+MUL , 8
+bailingmoe3 ,1 ,any ,RMS_NORM+SCALE , 2
bloom ,0 ,any ,NORM+MUL+ADD , 6
chatglm ,0 ,any ,RMS_NORM+MUL , 5
codeshell ,0 ,any ,NORM+MUL+ADD , 5
cogvlm ,0 ,any ,RMS_NORM+MUL , 5
+cohere2 ,0 ,any ,ADD+ADD , 2
+cohere2 ,0 ,any ,NORM+MUL , 3
+cohere2moe ,1 ,any ,ADD+ADD , 2
+cohere2moe ,1 ,any ,MUL+ADD , 2
+cohere2moe ,1 ,any ,RMS_NORM+MUL , 3
command-r ,0 ,any ,NORM+MUL , 3
+dbrx ,0 ,any ,MUL+ADD , 4
dbrx ,0 ,any ,NORM+MUL , 5
+dbrx ,0 ,any ,SOFT_MAX+ARGSORT+GET_ROWS+SUM_ROWS+CLAMP+DIV, 2
deci ,0 ,any ,RMS_NORM+MUL , 5
deepseek ,0 ,any ,ADD+ADD , 1
+deepseek ,0 ,any ,MUL+ADD , 2
deepseek ,0 ,any ,RMS_NORM+MUL , 5
+deepseek ,0 ,any ,SOFT_MAX+ARGSORT+GET_ROWS , 1
deepseek2 ,0 ,any ,ADD+ADD , 1
+deepseek2 ,0 ,any ,MUL+ADD , 2
deepseek2 ,0 ,any ,RMS_NORM+MUL , 9
deepseek32 ,0 ,any ,ADD+ADD , 1
+deepseek32 ,0 ,any ,MUL+ADD , 2
deepseek32 ,0 ,any ,NORM+MUL+ADD , 2
deepseek32 ,0 ,any ,RMS_NORM+MUL , 9
+deepseek4 ,0 ,any ,MUL+ADD , 8
deepseek4 ,0 ,any ,RMS_NORM+MUL , 20
dots1 ,0 ,any ,ADD+ADD , 1
+dots1 ,0 ,any ,MUL+ADD , 2
dots1 ,0 ,any ,RMS_NORM+MUL , 9
+dots3note ,0 ,any ,ADD+ADD , 1
+dots3note ,0 ,any ,MUL+ADD , 2
+dots3note ,0 ,any ,NORM+MUL+ADD , 1
+dots3note ,0 ,any ,RMS_NORM+MUL , 11
dream ,0 ,any ,RMS_NORM+MUL , 5
ernie4_5-moe ,1 ,any ,ADD+ADD , 1
+ernie4_5-moe ,1 ,any ,MUL+ADD , 2
ernie4_5-moe ,1 ,any ,RMS_NORM+MUL , 5
ernie4_5 ,0 ,any ,RMS_NORM+MUL , 5
exaone ,0 ,any ,RMS_NORM+MUL , 5
+exaone-moe ,1 ,any ,ADD+ADD , 1
+exaone-moe ,1 ,any ,MUL+ADD , 2
+exaone-moe ,1 ,any ,RMS_NORM+MUL , 9
exaone4 ,0 ,any ,RMS_NORM+MUL , 5
exaone4 ,0 ,any ,RMS_NORM+MUL+ADD , 4
falcon ,0 ,any ,ADD+ADD , 2
@@ -41,31 +72,48 @@ falcon-h1 ,0 ,any ,RMS_NORM+MUL , 9
gemma ,0 ,any ,RMS_NORM+MUL , 5
gemma2 ,0 ,any ,RMS_NORM+MUL , 5
gemma2 ,0 ,any ,RMS_NORM+MUL+ADD , 4
+gemma3 ,0 ,any ,RMS_NORM+MUL , 9
+gemma3 ,0 ,any ,RMS_NORM+MUL+ADD , 4
glm-dsa ,0 ,any ,ADD+ADD , 1
+glm-dsa ,0 ,any ,MUL+ADD , 2
glm-dsa ,0 ,any ,NORM+MUL+ADD , 2
glm-dsa ,0 ,any ,RMS_NORM+MUL , 9
glm4 ,0 ,any ,RMS_NORM+MUL , 5
glm4 ,0 ,any ,RMS_NORM+MUL+ADD , 4
glm4moe ,1 ,any ,ADD+ADD , 1
+glm4moe ,1 ,any ,MUL+ADD , 2
glm4moe ,1 ,any ,RMS_NORM+MUL , 9
+gpt-oss ,0 ,any ,MUL+ADD , 4
gpt-oss ,0 ,any ,RMS_NORM+MUL , 5
gpt2 ,0 ,any ,NORM+MUL+ADD , 5
gptneox ,0 ,any ,NORM+MUL+ADD , 5
granite ,0 ,any ,RMS_NORM+MUL , 5
+granite ,0 ,any ,MUL+ADD , 4
granite ,0 ,any ,RMS_NORM+MUL , 5
+granite ,0 ,any ,SOFT_MAX+ARGSORT+GET_ROWS+SUM_ROWS+CLAMP+DIV, 2
+granite_swa ,0 ,any ,RMS_NORM+MUL , 5
granitehybrid ,0 ,any ,RMS_NORM+MUL , 6
granitemoe ,1 ,any ,RMS_NORM+MUL , 5
+granitemoe ,1 ,any ,MUL+ADD , 4
granitemoe ,1 ,any ,RMS_NORM+MUL , 5
+granitemoe ,1 ,any ,SOFT_MAX+ARGSORT+GET_ROWS+SUM_ROWS+CLAMP+DIV, 2
+grok ,0 ,any ,MUL+ADD , 4
grok ,0 ,any ,RMS_NORM+MUL , 5
grok ,0 ,any ,RMS_NORM+MUL+ADD , 4
+grok ,0 ,any ,SOFT_MAX+ARGSORT+GET_ROWS+SUM_ROWS+CLAMP+DIV, 2
grovemoe ,1 ,any ,ADD+ADD , 2
+grovemoe ,1 ,any ,MUL+ADD , 8
grovemoe ,1 ,any ,RMS_NORM+MUL , 9
hunyuan-dense ,0 ,any ,RMS_NORM+MUL , 9
hunyuan-moe ,1 ,any ,ADD+ADD , 2
+hunyuan-moe ,1 ,any ,MUL+ADD , 4
hunyuan-moe ,1 ,any ,RMS_NORM+MUL , 9
+hunyuan-moe ,1 ,any ,SOFT_MAX+ARGSORT+GET_ROWS+SUM_ROWS+CLAMP+DIV, 2
hunyuan_vl ,0 ,any ,RMS_NORM+MUL , 9
hy_v3 ,0 ,any ,ADD+ADD , 2
+hy_v3 ,0 ,any ,MUL+ADD , 4
hy_v3 ,0 ,any ,RMS_NORM+MUL , 9
+hy_v4 ,0 ,any ,MUL+ADD , 2
hy_v4 ,0 ,any ,NORM+MUL+ADD , 1
hy_v4 ,0 ,any ,RMS_NORM+MUL , 9
internlm2 ,0 ,any ,RMS_NORM+MUL , 5
@@ -73,38 +121,73 @@ jais ,0 ,any ,NORM+MUL+ADD , 5
jais2 ,0 ,any ,NORM+MUL+ADD , 5
jamba ,0 ,any ,RMS_NORM+MUL , 8
kimi-k3 ,0 ,any ,GATED_DELTA_NET+CPY , 1
+kimi-k3 ,0 ,any ,MUL+ADD , 2
kimi-k3 ,0 ,any ,RMS_NORM+MUL , 17
+kimi-k3 ,0 ,any ,RMS_NORM+SCALE , 2
kimi-linear ,0 ,any ,ADD+ADD , 1
kimi-linear ,0 ,any ,GATED_DELTA_NET+CPY , 1
+kimi-linear ,0 ,any ,MUL+ADD , 2
kimi-linear ,0 ,any ,RMS_NORM+MUL , 7
+kimi-linear ,0 ,any ,RMS_NORM+SCALE , 2
+laguna ,0 ,any ,ADD+ADD , 1
+laguna ,0 ,any ,MUL+ADD , 2
+laguna ,0 ,any ,RMS_NORM+MUL , 9
lfm2 ,0 ,any ,RMS_NORM+MUL , 7
+lfm2moe ,1 ,any ,MUL+ADD , 2
lfm2moe ,1 ,any ,RMS_NORM+MUL , 7
llada ,0 ,any ,RMS_NORM+MUL , 5
+llada-moe ,1 ,any ,MUL+ADD , 4
llada-moe ,1 ,any ,RMS_NORM+MUL , 9
+llada-moe ,1 ,any ,SOFT_MAX+ARGSORT+GET_ROWS , 2
llama ,0 ,any ,RMS_NORM+MUL , 5
+llama ,0 ,any ,MUL+ADD , 4
llama ,0 ,any ,RMS_NORM+MUL , 5
+llama ,0 ,any ,SOFT_MAX+ARGSORT+GET_ROWS+SUM_ROWS+CLAMP+DIV, 2
llama4 ,0 ,any ,ADD+ADD , 2
llama4 ,0 ,any ,RMS_NORM+MUL , 9
maincoder ,0 ,any ,RMS_NORM+MUL , 9
mamba ,0 ,any ,RMS_NORM+MUL , 3
mamba2 ,0 ,any ,RMS_NORM+MUL , 5
+maple ,0 ,any ,MUL+ADD , 4
+maple ,0 ,any ,RMS_NORM+MUL , 9
+maple ,0 ,any ,SOFT_MAX+ARGSORT+GET_ROWS+SUM_ROWS+CLAMP+DIV, 2
+mellum ,0 ,any ,MUL+ADD , 4
+mellum ,0 ,any ,RMS_NORM+MUL , 9
+mellum ,0 ,any ,SOFT_MAX+ARGSORT+GET_ROWS+SUM_ROWS+CLAMP+DIV, 2
+mimo2 ,0 ,any ,MUL+ADD , 4
+mimo2 ,0 ,any ,RMS_NORM+MUL , 5
minicpm ,0 ,any ,RMS_NORM+MUL , 5
+minicpm ,0 ,any ,MUL+ADD , 4
minicpm ,0 ,any ,RMS_NORM+MUL , 5
+minicpm ,0 ,any ,SOFT_MAX+ARGSORT+GET_ROWS+SUM_ROWS+CLAMP+DIV, 2
minicpm3 ,0 ,any ,RMS_NORM+MUL , 9
+minimax-01 ,0 ,any ,MUL+ADD , 4
minimax-01 ,0 ,any ,RMS_NORM+MUL , 6
+minimax-01 ,0 ,any ,SOFT_MAX+ARGSORT+GET_ROWS+SUM_ROWS+CLAMP+DIV, 2
+minimax-m2 ,0 ,any ,MUL+ADD , 4
minimax-m2 ,0 ,any ,RMS_NORM+MUL , 9
minimax-m3 ,0 ,any ,ADD+ADD , 1
+minimax-m3 ,0 ,any ,MUL+ADD , 2
minimax-m3 ,0 ,any ,RMS_NORM+MUL , 11
mistral3 ,0 ,any ,RMS_NORM+MUL , 5
+mistral3 ,0 ,any ,MUL+ADD , 4
mistral3 ,0 ,any ,RMS_NORM+MUL , 5
+mistral3 ,0 ,any ,SOFT_MAX+ARGSORT+GET_ROWS+SUM_ROWS+CLAMP+DIV, 2
mistral4 ,0 ,any ,ADD+ADD , 1
+mistral4 ,0 ,any ,MUL+ADD , 2
mistral4 ,0 ,any ,RMS_NORM+MUL , 9
mpt ,0 ,any ,NORM+MUL+ADD , 5
+muse-glimmer ,0 ,any ,RMS_NORM+MUL , 10
+muse-glimmer ,0 ,any ,RMS_NORM+MUL+ADD , 3
nanbeige ,0 ,any ,RMS_NORM+MUL , 5
nemotron ,0 ,any ,NORM+MUL+ADD , 5
nemotron_h ,0 ,any ,RMS_NORM+MUL , 5
nemotron_h_moe ,1 ,any ,RMS_NORM+MUL , 5
+olmo2 ,0 ,any ,RMS_NORM+MUL , 5
+olmo2 ,0 ,any ,RMS_NORM+MUL+ADD , 4
+olmoe ,1 ,any ,MUL+ADD , 4
olmoe ,1 ,any ,RMS_NORM+MUL , 9
+olmoe ,1 ,any ,SOFT_MAX+ARGSORT+GET_ROWS , 2
openelm ,0 ,any ,RMS_NORM+MUL , 9
orion ,0 ,any ,NORM+MUL+ADD , 5
paddleocr ,0 ,any ,RMS_NORM+MUL , 5
@@ -112,40 +195,69 @@ pangu-embedded ,0 ,any ,RMS_NORM+MUL , 5
phi2 ,0 ,any ,ADD+ADD , 2
phi2 ,0 ,any ,NORM+MUL+ADD , 3
phi3 ,0 ,any ,RMS_NORM+MUL , 5
+phimoe ,1 ,any ,MUL+ADD , 4
phimoe ,1 ,any ,RMS_NORM+MUL+ADD , 5
+phimoe ,1 ,any ,SOFT_MAX+ARGSORT+GET_ROWS+SUM_ROWS+CLAMP+DIV, 2
plamo ,0 ,any ,ADD+ADD , 2
plamo ,0 ,any ,RMS_NORM+MUL , 3
plamo2 ,0 ,any ,RMS_NORM+MUL , 10
plamo2 ,0 ,any ,RMS_NORM+MUL+ADD , 4
+plamo2 ,0 ,any ,SSM_CONV+UNARY , 1
+plamo3 ,0 ,any ,RMS_NORM+MUL , 9
+plamo3 ,0 ,any ,RMS_NORM+MUL+ADD , 4
pockettts ,0 ,any ,NORM+MUL+ADD , 5
qwen ,0 ,any ,RMS_NORM+MUL , 5
qwen2 ,0 ,any ,RMS_NORM+MUL , 5
qwen2moe ,1 ,any ,ADD+ADD , 2
+qwen2moe ,1 ,any ,MUL+ADD , 4
qwen2moe ,1 ,any ,RMS_NORM+MUL , 5
+qwen2moe ,1 ,any ,SOFT_MAX+ARGSORT+GET_ROWS , 2
qwen2vl ,0 ,any ,RMS_NORM+MUL , 5
qwen3 ,0 ,any ,RMS_NORM+MUL , 9
qwen35 ,0 ,any ,GATED_DELTA_NET+CPY , 1
qwen35 ,0 ,any ,RMS_NORM+MUL , 8
+qwen35 ,0 ,any ,RMS_NORM+SCALE , 2
+qwen35 ,0 ,any ,SSM_CONV+UNARY , 1
qwen35moe ,1 ,any ,ADD+ADD , 2
qwen35moe ,1 ,any ,GATED_DELTA_NET+CPY , 1
+qwen35moe ,1 ,any ,MUL+ADD , 4
qwen35moe ,1 ,any ,RMS_NORM+MUL , 8
+qwen35moe ,1 ,any ,RMS_NORM+SCALE , 2
+qwen35moe ,1 ,any ,SOFT_MAX+ARGSORT+GET_ROWS+SUM_ROWS+CLAMP+DIV, 2
+qwen35moe ,1 ,any ,SSM_CONV+UNARY , 1
+qwen3moe ,1 ,any ,MUL+ADD , 4
qwen3moe ,1 ,any ,RMS_NORM+MUL , 9
+qwen3moe ,1 ,any ,SOFT_MAX+ARGSORT+GET_ROWS+SUM_ROWS+CLAMP+DIV, 2
qwen3next ,0 ,any ,ADD+ADD , 2
qwen3next ,0 ,any ,GATED_DELTA_NET+CPY , 1
+qwen3next ,0 ,any ,MUL+ADD , 4
qwen3next ,0 ,any ,RMS_NORM+MUL , 8
+qwen3next ,0 ,any ,RMS_NORM+SCALE , 2
+qwen3next ,0 ,any ,SOFT_MAX+ARGSORT+GET_ROWS+SUM_ROWS+CLAMP+DIV, 2
+qwen3next ,0 ,any ,SSM_CONV+UNARY , 1
qwen3tts ,0 ,any ,RMS_NORM+MUL , 9
qwen3vl ,0 ,any ,RMS_NORM+MUL , 9
+qwen3vlmoe ,1 ,any ,MUL+ADD , 4
qwen3vlmoe ,1 ,any ,RMS_NORM+MUL , 9
-qwen4exp ,0 ,any ,ADD+ADD+ADD , 5
+qwen3vlmoe ,1 ,any ,SOFT_MAX+ARGSORT+GET_ROWS+SUM_ROWS+CLAMP+DIV, 2
+qwen4exp ,0 ,any ,ADD+ADD+ADD , 1
qwen4exp ,0 ,any ,ADD+ADD+ADD+ADD+ADD+ADD+ADD , 9
qwen4exp ,0 ,any ,GATED_DELTA_NET+CPY , 1
-qwen4exp ,0 ,any ,RMS_NORM+MUL , 5
+qwen4exp ,0 ,any ,MUL+ADD , 4
+qwen4exp ,0 ,any ,RMS_NORM+MUL , 13
+qwen4exp ,0 ,any ,RMS_NORM+SCALE , 2
+qwen4exp ,0 ,any ,SOFT_MAX+ARGSORT+GET_ROWS+SUM_ROWS+CLAMP+DIV, 2
+qwen4exp ,0 ,any ,SSM_CONV+UNARY , 1
refact ,0 ,any ,RMS_NORM+MUL , 5
refact ,0 ,any ,RMS_NORM+MUL , 5
+rnd1 ,0 ,any ,MUL+ADD , 4
rnd1 ,0 ,any ,RMS_NORM+MUL , 9
+rnd1 ,0 ,any ,SOFT_MAX+ARGSORT+GET_ROWS+SUM_ROWS+CLAMP+DIV, 2
seed_oss ,0 ,any ,RMS_NORM+MUL , 5
+smallthinker ,0 ,any ,MUL+ADD , 4
smallthinker ,0 ,any ,RMS_NORM+MUL , 5
smollm3 ,0 ,any ,RMS_NORM+MUL , 5
+spark2_5 ,0 ,any ,RMS_NORM+MUL , 5
stablelm ,0 ,any ,NORM+MUL , 4
stablelm ,0 ,any ,NORM+MUL+ADD , 5
starcoder ,0 ,any ,NORM+MUL+ADD , 5
diff --git a/tests/test-backend-ops.cpp b/tests/test-backend-ops.cpp
index d8f4c3708..34f5e4587 100644
--- a/tests/test-backend-ops.cpp
+++ b/tests/test-backend-ops.cpp
@@ -3601,6 +3601,44 @@ struct test_norm_mul_add : public test_case {
return out;
}
};
+// GGML_OP_NORM/RMS_NORM + GGML_OP_SCALE
+struct test_norm_scale : public test_case {
+ const ggml_type type;
+ const std::array<int64_t, 4> ne;
+ const float eps;
+ const bool rms;
+ const float scale;
+
+ std::string vars() override {
+ return VARS_TO_STR5(type, ne, eps, rms, scale);
+ }
+
+ test_norm_scale(ggml_type type = GGML_TYPE_F32,
+ std::array<int64_t, 4> ne = {64, 5, 4, 3},
+ float eps = 1e-6f,
+ bool rms = false,
+ float scale = 1.5f)
+ : type(type), ne(ne), eps(eps), rms(rms), scale(scale) {}
+
+ std::string op_desc(ggml_tensor * t) override {
+ GGML_UNUSED(t);
+ return rms ? "RMS_NORM_SCALE" : "NORM_SCALE";
+ }
+
+ bool run_whole_graph() override { return true; }
+
+ ggml_tensor * build_graph(ggml_context * ctx) override {
+ ggml_tensor * a = ggml_new_tensor(ctx, type, 4, ne.data());
+ ggml_set_name(a, "a");
+
+ ggml_tensor * n = rms ? ggml_rms_norm(ctx, a, eps) : ggml_norm(ctx, a, eps);
+ ggml_tensor * out = ggml_scale(ctx, n, scale);
+ ggml_set_name(out, "out");
+
+ return out;
+ }
+};
+
// GGML_OP_RMS_NORM
struct test_rms_norm : public test_case {
const ggml_type type;
@@ -6802,7 +6840,7 @@ struct test_topk_moe : public test_case {
}
};
-struct test_moe_weighted_reduction : public test_case {
+struct test_moe_reduce : public test_case {
const int64_t n_embd;
const int64_t n_expert_used;
const int64_t n_tokens;
@@ -6810,7 +6848,7 @@ struct test_moe_weighted_reduction : public test_case {
const bool with_expert_scale;
const bool interleaved_views_adds;
- test_moe_weighted_reduction(
+ test_moe_reduce(
int64_t n_embd, int64_t n_expert_used, int64_t n_tokens,
bool unaligned_experts = false, bool with_expert_scale = false, bool interleaved_views_adds = false) :
n_embd(n_embd), n_expert_used(n_expert_used), n_tokens(n_tokens),
@@ -6823,7 +6861,7 @@ struct test_moe_weighted_reduction : public test_case {
std::string op_desc(ggml_tensor * t) override {
GGML_UNUSED(t);
- return "MOE_WEIGHTED_REDUCTION";
+ return "MOE_REDUCE";
}
bool run_whole_graph() override { return true; }
@@ -6870,7 +6908,7 @@ struct test_moe_weighted_reduction : public test_case {
ggml_build_forward_expand(gf, out);
}
}
- ggml_set_name(out, "moe_weighted_reduction");
+ ggml_set_name(out, "moe_reduce");
return out;
}
};
@@ -9656,6 +9694,8 @@ static std::vector<std::unique_ptr<test_case>> make_test_cases_eval() {
test_cases.emplace_back(new test_rms_norm(GGML_TYPE_F32, { n, 5, 4, 3 }, v, eps));
}
test_cases.emplace_back(new test_norm(GGML_TYPE_F32, { n, 5, 4, 3 }, false, eps, true));
+ test_cases.emplace_back(new test_norm_scale(GGML_TYPE_F32, { n, 5, 4, 3 }, eps, false, 1.5f));
+ test_cases.emplace_back(new test_norm_scale(GGML_TYPE_F32, { n, 5, 4, 3 }, eps, true, 1.5f));
test_cases.emplace_back(new test_rms_norm_back(GGML_TYPE_F32, { n, 5, 4, 3 }, eps));
test_cases.emplace_back(new test_l2_norm(GGML_TYPE_F32, { n, 5, 4, 3 }, eps, false));
test_cases.emplace_back(new test_l2_norm(GGML_TYPE_F32, { n, 5, 4, 3 }, eps, true));
@@ -10903,12 +10943,12 @@ static std::vector<std::unique_ptr<test_case>> make_test_cases_eval() {
}
// Cover the supported boundaries, common k = 8 shapes, interleaved views and adds, and k = 16 fallback.
- test_cases.emplace_back(new test_moe_weighted_reduction(63, 2, 17));
- test_cases.emplace_back(new test_moe_weighted_reduction(2048, 8, 128));
- test_cases.emplace_back(new test_moe_weighted_reduction(2048, 8, 128, false, true));
- test_cases.emplace_back(new test_moe_weighted_reduction(63, 12, 33, true, true, true));
- test_cases.emplace_back(new test_moe_weighted_reduction(2048, 15, 40, false, true));
- test_cases.emplace_back(new test_moe_weighted_reduction(2048, 16, 32, false, true));
+ test_cases.emplace_back(new test_moe_reduce(63, 2, 17));
+ test_cases.emplace_back(new test_moe_reduce(2048, 8, 128));
+ test_cases.emplace_back(new test_moe_reduce(2048, 8, 128, false, true));
+ test_cases.emplace_back(new test_moe_reduce(63, 12, 33, true, true, true));
+ test_cases.emplace_back(new test_moe_reduce(2048, 15, 40, false, true));
+ test_cases.emplace_back(new test_moe_reduce(2048, 16, 32, false, true));
test_cases.emplace_back(new test_gated_delta_net(GGML_TYPE_F32, 32, 128, 1, 1));
test_cases.emplace_back(new test_gated_delta_net(GGML_TYPE_F32, 32, 16, 1, 1));
diff --git a/tests/test-llama-archs.cpp b/tests/test-llama-archs.cpp
index f848fc139..80a045185 100644
--- a/tests/test-llama-archs.cpp
+++ b/tests/test-llama-archs.cpp
@@ -360,7 +360,7 @@ static gguf_context_ptr get_gguf_ctx(const llm_arch arch, const bool moe) {
ms.add_kv(LLM_KV_EXPERT_LATENT_LENGTH, n_ff);
ms.add_kv(LLM_KV_INTERLEAVE_MOE_LAYER_STEP, uint32_t(2));
ms.add_kv(LLM_KV_EXPERT_COUNT, uint32_t(2));
- ms.add_kv(LLM_KV_EXPERT_USED_COUNT, uint32_t(1));
+ ms.add_kv(LLM_KV_EXPERT_USED_COUNT, uint32_t(2));
ms.add_kv(LLM_KV_EXPERT_SHARED_COUNT, uint32_t(1));
ms.add_kv(LLM_KV_EXPERT_GATING_FUNC, arch == LLM_ARCH_DEEPSEEK4 ? uint32_t(4) : uint32_t(2)); // sqrtsoftplus : sigmoid
ms.add_kv(LLM_KV_EXPERT_GROUP_SCALE, 1.0f);