mirror of
https://github.com/ggml-org/llama.cpp.git
synced 2026-09-27 21:46:57 +02:00
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
This commit is contained in:
@@ -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";
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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);
|
||||
|
||||
@@ -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);
|
||||
|
||||
@@ -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->n_ops >= 2);
|
||||
GGML_ASSERT(fusion->ops.size() >= 2);
|
||||
|
||||
for (int j = 1; j < fusion->n_ops; j++) {
|
||||
if (fusion->id == GGML_METAL_FUSION_NORM_SCALE) {
|
||||
GGML_ASSERT(fusion->ops.size() == 2);
|
||||
|
||||
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;
|
||||
}
|
||||
|
||||
// ---- patterns ------------------------------------------------------------
|
||||
#define GGML_METAL_TOPK_MOE_MAX_EXPERTS 1024
|
||||
|
||||
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 },
|
||||
// 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
|
||||
};
|
||||
|
||||
const ggml_metal_fusion * ggml_metal_fusion_all(int * n) {
|
||||
*n = (int) sizeof(ggml_metal_fusions) / sizeof(ggml_metal_fusions[0]);
|
||||
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);
|
||||
|
||||
return ggml_metal_fusions;
|
||||
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 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 },
|
||||
};
|
||||
|
||||
// ---- 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 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;
|
||||
}
|
||||
|
||||
best = fusion->n_ops;
|
||||
res = fusion;
|
||||
// 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 = n_ops;
|
||||
res = &fusion;
|
||||
}
|
||||
|
||||
*n_out = best;
|
||||
|
||||
@@ -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 ---------------------------------------------------
|
||||
|
||||
|
||||
@@ -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;
|
||||
|
||||
@@ -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);
|
||||
|
||||
|
||||
@@ -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);
|
||||
|
||||
@@ -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;
|
||||
|
||||
|
||||
@@ -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;
|
||||
}
|
||||
|
||||
@@ -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];
|
||||
|
||||
@@ -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
|
||||
|
||||
+114
-2
@@ -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
|
||||
|
||||
|
+50
-10
@@ -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));
|
||||
|
||||
@@ -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);
|
||||
|
||||
Reference in New Issue
Block a user