Commit a3a1c4747 for llama.cpp
commit a3a1c4747fdc0dcad40b3946108b89375d9a7d0e
Author: pratiknarola-t <pratik.narola@tether.io>
Date: Mon Oct 5 10:59:13 2026 +0530
metal : few-row MMA mat-mul (#29869)
* metal : few-row MMA mat-mul and batched copies for speculative decoding
Speculative decoding verifies a few draft tokens per step. Without the tensor API, Metal ran these mat-muls with the mat-vec kernels, whose time grows with every src1 row, so DFlash2 decoding on an M3 Ultra was slower than serial decoding.
- add mat-mul kernels for 2..16 src1 rows on 8x8 simdgroup matrices: each weight is dequantized once for all rows, and the simdgroups of a threadgroup split K. Q4_0, Q8_0 and Q5_K have their own kernels, F32, F16, Q4_1, Q5_0, Q5_1, Q4_K and Q6_K use a generic path over the 16-weight dequantizers, and Q4_0 at 2 rows uses a 2-row variant of the mat-vec kernel
- use them only on MTLGPUFamilyApple7+ without the tensor API, from the row count at which they beat the mat-vec kernels on an M3 Ultra (F32: 6, F16, Q4_K, Q5_0, Q5_1: 3, other types: 2)
- fusion table: MUL_MAT + ADD adds a same-shape residual in the MMA store, and up to 16 adjacent same-layout f32 copies between the same two tensors run as one dispatch
- the fusion checks and ggml_graph_optimize take the device props, so the reorder packs MUL_MAT + ADD only on devices that can fuse it, at every src1 row count
- views do not count toward GGML_METAL_FUSION_MAX when the reorder packs a group, so 16 recurrent state snapshot copies with views between them stay one group
- the encoder checks the inner nodes of a fused group for concurrency, tracks written views by their extent, and does not count the destination of a CPY as a read
- CONCAT splits long rows across threadgroups when there are few rows
- tests: few-row MUL_MAT, MUL_MAT_ADD, CPY_BATCH and CONCAT cases in test-backend-ops (with a prepare_graph hook for the copy order), test-metal-graph-optimize, test-metal-cpy-batch-alias
* metal : remove the CPY_BATCH fusion and the memory range changes
Remove the batched copy fusion with its kernel and tests, and revert the
memory range changes, as suggested in review. The memory ranges, the
graph reorder and the CPY encoder are again the same as on master.
* cont : clean-up
* cont : drop has_tensor gate
* cont : clean-up operand/residual logic
* cont : drop Q4_0 ne11=2 special-case
* cont : add kernels/mul_mv_mma.metal
* cont : consolidate mma pipeline selection logic
* cont : decouple fusion logic from device props
---------
Co-authored-by: Georgi Gerganov <ggerganov@gmail.com>
diff --git a/ggml/src/ggml-metal/CMakeLists.txt b/ggml/src/ggml-metal/CMakeLists.txt
index 08408a2d4..05cdea156 100644
--- a/ggml/src/ggml-metal/CMakeLists.txt
+++ b/ggml/src/ggml-metal/CMakeLists.txt
@@ -53,6 +53,7 @@ set(METALLIB_KERNEL_SOURCES
kernels/fa_vec_q5_1.metal
kernels/fa_vec_q8_0.metal
kernels/mul_mv.metal
+ kernels/mul_mv_mma.metal
kernels/mul_mm.metal
kernels/quantize.metal
kernels/softmax.metal
diff --git a/ggml/src/ggml-metal/ggml-metal-common.cpp b/ggml/src/ggml-metal/ggml-metal-common.cpp
index e43023cc6..8b065f54a 100644
--- a/ggml/src/ggml-metal/ggml-metal-common.cpp
+++ b/ggml/src/ggml-metal/ggml-metal-common.cpp
@@ -48,6 +48,90 @@ bool ggml_metal_op_mul_mat_id_use_mm(const struct ggml_tensor * op, bool has_sim
return has_simdgroup_mm && ne00 >= 64 && ne21 >= 32;
}
+// the most src1 rows of the few-row MMA kernels
+static constexpr int64_t GGML_METAL_MMA_ROWS_MAX = 16;
+
+// src1 rows per 8x8 simdgroup matrix tile of the few-row MMA kernels
+static constexpr int64_t GGML_METAL_MMA_TILE_ROWS = 8;
+
+// weights per K step of the q5_K and generic few-row MMA kernels
+static constexpr int64_t GGML_METAL_MMA_K_CHUNK = 64;
+
+enum ggml_metal_mma_kind ggml_metal_mul_mv_mma_kind(enum ggml_type type, int rt) {
+ if (type == GGML_TYPE_Q4_0 || (type == GGML_TYPE_Q8_0 && rt == 1)) {
+ return GGML_METAL_MMA_KIND_BLK;
+ }
+ return type == GGML_TYPE_Q5_K ? GGML_METAL_MMA_KIND_Q5_K : GGML_METAL_MMA_KIND_GEN;
+}
+
+int ggml_metal_mul_mv_mma_rt(const struct ggml_tensor * op) {
+ return op->src[1]->ne[1] > GGML_METAL_MMA_TILE_ROWS ? 2 : 1;
+}
+
+static bool ggml_metal_mul_mv_mma_type_supported(enum ggml_type type) {
+ switch (type) {
+ case GGML_TYPE_F32:
+ case GGML_TYPE_F16:
+ case GGML_TYPE_Q4_0:
+ case GGML_TYPE_Q4_1:
+ case GGML_TYPE_Q5_0:
+ case GGML_TYPE_Q5_1:
+ case GGML_TYPE_Q8_0:
+ case GGML_TYPE_Q4_K:
+ case GGML_TYPE_Q5_K:
+ case GGML_TYPE_Q6_K:
+ return true;
+ default:
+ return false;
+ }
+}
+
+int64_t ggml_metal_mul_mv_mma_k_step(enum ggml_type type, int rt) {
+ if (!ggml_metal_mul_mv_mma_type_supported(type)) {
+ return 0;
+ }
+ return ggml_metal_mul_mv_mma_kind(type, rt) == GGML_METAL_MMA_KIND_BLK ? ggml_blck_size(type) : GGML_METAL_MMA_K_CHUNK;
+}
+
+static bool ggml_metal_mul_mat_mma_type_ok(const struct ggml_tensor * op) {
+ const ggml_tensor * src0 = op->src[0];
+ const int64_t step = ggml_metal_mul_mv_mma_k_step(src0->type, ggml_metal_mul_mv_mma_rt(op));
+
+ return step > 0 && src0->ne[0] % step == 0 && src0->nb[0] == ggml_type_size(src0->type);
+}
+
+// the fewest src1 rows at which the few-row MMA kernels beat the mat-vec kernels (measured on an M3 Ultra)
+static int64_t ggml_metal_mul_mv_mma_rows_min(enum ggml_type type) {
+ switch (type) {
+ case GGML_TYPE_F32:
+ return 6;
+ case GGML_TYPE_F16:
+ case GGML_TYPE_Q4_K:
+ case GGML_TYPE_Q5_0:
+ case GGML_TYPE_Q5_1:
+ return 3;
+ default:
+ return 2;
+ }
+}
+
+bool ggml_metal_op_mul_mat_use_mma(const struct ggml_tensor * op) {
+ const ggml_tensor * src0 = op->src[0];
+ const ggml_tensor * src1 = op->src[1];
+
+ // the batch shape goes into int16 function constants
+ const bool batch_ok = src1->ne[2] <= INT16_MAX && src1->ne[2]/src0->ne[2] <= INT16_MAX && src1->ne[3]/src0->ne[3] <= INT16_MAX;
+
+ return ggml_metal_mul_mat_mma_type_ok(op) && batch_ok &&
+ src1->type == GGML_TYPE_F32 && src1->ne[1] >= ggml_metal_mul_mv_mma_rows_min(src0->type) && src1->ne[1] <= GGML_METAL_MMA_ROWS_MAX &&
+ !ggml_is_transposed(src0) && !ggml_is_transposed(src1) &&
+ src1->nb[0] == sizeof(float) && src1->nb[1] % 16 == 0 && src1->nb[2] % 16 == 0 && src1->nb[3] % 16 == 0;
+}
+
+bool ggml_metal_op_mul_mat_may_use_mma(const struct ggml_tensor * op) {
+ return ggml_metal_mul_mv_mma_type_supported(op->src[0]->type) && op->src[1]->type == GGML_TYPE_F32;
+}
+
// represents a memory range (i.e. an interval from a starting address p0 to an ending address p1 in a given buffer pb)
// the type indicates whether it is a source range (i.e. ops read data from it) or a destination range (i.e. ops write data to it)
struct ggml_mem_range {
diff --git a/ggml/src/ggml-metal/ggml-metal-common.h b/ggml/src/ggml-metal/ggml-metal-common.h
index 6b5a1883f..c86d92b9d 100644
--- a/ggml/src/ggml-metal/ggml-metal-common.h
+++ b/ggml/src/ggml-metal/ggml-metal-common.h
@@ -2,6 +2,8 @@
#pragma once
+#include "ggml.h"
+
#include <stdbool.h>
#include <stddef.h>
@@ -42,17 +44,30 @@ bool ggml_mem_ranges_add(ggml_mem_ranges_t mrs, const struct ggml_tensor * tenso
// - new dst range overlaps with any existing range (src or dst)
bool ggml_mem_ranges_check(ggml_mem_ranges_t mrs, const struct ggml_tensor * tensor);
-// reorder the nodes in the graph to improve concurrency, while respecting fusion
+// reorder the nodes in the graph to improve concurrency, while respecting the fusions of a device with props
//
// note: this implementation is generic and not specific to metal
// if it proves to work well, we can start using it for other backends in the future
void ggml_graph_optimize(struct ggml_cgraph * gf);
// mat-mat vs mat-vec dispatch; used by both supports_op and ggml_metal_op_mul_mat*
-bool ggml_metal_op_mul_mat_use_fwht(const struct ggml_tensor * op, size_t max_tg_mem);
+bool ggml_metal_op_mul_mat_use_fwht (const struct ggml_tensor * op, size_t max_tg_mem);
bool ggml_metal_op_mul_mat_use_mm (const struct ggml_tensor * op, bool has_simdgroup_mm);
bool ggml_metal_op_mul_mat_id_use_mm(const struct ggml_tensor * op, bool has_simdgroup_mm);
+bool ggml_metal_op_mul_mat_use_mma (const struct ggml_tensor * op);
+bool ggml_metal_op_mul_mat_may_use_mma(const struct ggml_tensor * op); // graph structure only
+
+// the few-row MMA kernel for a src0 type and rt src1 tiles: per 32-weight block (q4_0, q8_0 with one tile), q5_K, or the generic 64-weight chunk kernel
+enum ggml_metal_mma_kind { GGML_METAL_MMA_KIND_BLK, GGML_METAL_MMA_KIND_Q5_K, GGML_METAL_MMA_KIND_GEN };
+enum ggml_metal_mma_kind ggml_metal_mul_mv_mma_kind(enum ggml_type type, int rt);
+
+// the src1 tiles of the few-row MMA kernels for mat-mul op: one 8-row tile, or two above 8 rows
+int ggml_metal_mul_mv_mma_rt(const struct ggml_tensor * op);
+
+// the weights of K per simdgroup step of the few-row MMA kernel for a src0 type and rt src1 tiles, 0 if none takes the type
+int64_t ggml_metal_mul_mv_mma_k_step(enum ggml_type type, int rt);
+
#ifdef __cplusplus
}
#endif
diff --git a/ggml/src/ggml-metal/ggml-metal-device.cpp b/ggml/src/ggml-metal/ggml-metal-device.cpp
index 91b6ef1dc..8ac6e0830 100644
--- a/ggml/src/ggml-metal/ggml-metal-device.cpp
+++ b/ggml/src/ggml-metal/ggml-metal-device.cpp
@@ -1,4 +1,5 @@
#include "ggml-metal-device.h"
+#include "ggml-metal-common.h"
#include "ggml-metal-impl.h"
#include "ggml-metal-tuning.h"
@@ -801,6 +802,125 @@ ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_mul_mv_ext(ggml_
return res;
}
+// threadgroups needed to fill the GPU
+// TODO: dedup
+static constexpr int64_t GGML_METAL_MIN_THREADGROUPS = 128;
+
+// few-row MMA mat-mul: 8x8 simdgroup matrices for the src1 rows of ggml_metal_mul_mat_use_mma
+static constexpr int GGML_METAL_MMA_TILE = 8;
+static constexpr int GGML_METAL_MMA_NT_MAX = 4;
+static constexpr int GGML_METAL_MMA_SMEM_MAX = 16384;
+
+// simdgroups per threadgroup for at most FEW_ROWS, at most MANY_ROWS, and more src0 rows
+static constexpr int64_t GGML_METAL_MMA_FEW_ROWS = 64;
+static constexpr int64_t GGML_METAL_MMA_MANY_ROWS = 6144;
+static constexpr int GGML_METAL_MMA_NSG_FEW_ROWS = 32;
+static constexpr int GGML_METAL_MMA_NSG_MID_ROWS = 16;
+static constexpr int GGML_METAL_MMA_NSG_MANY_ROWS = 8;
+
+struct ggml_metal_mma_tiling {
+ int nsg; // simdgroups per threadgroup, each over a slice of K
+ int nt; // 8-row src0 tiles per threadgroup
+ int rt; // 8-row src1 tiles per threadgroup
+};
+
+// halves n while it is above limit
+static int ggml_metal_halve_to_limit(int n, int64_t limit) {
+ while (n > 1 && n > limit) {
+ n /= 2;
+ }
+ return n;
+}
+
+static size_t ggml_metal_mul_mv_mma_smem(int nsg, int nt, int rt) {
+ // one 8x8 float simdgroup matrix per output tile
+ constexpr size_t tile_bytes = 8*8*sizeof(float);
+ return (size_t) nsg*nt*rt*tile_bytes;
+}
+
+// fewer src0 rows need more simdgroups per threadgroup (a finer K split) to fill the GPU.
+// the tiles become narrower when the K-slice reduction buffer is too large or when too few threadgroups fill the GPU.
+static ggml_metal_mma_tiling ggml_metal_op_mul_mat_mma_tiling(const ggml_tensor * op) {
+ const ggml_type type = op->src[0]->type;
+ const int64_t ne00 = op->src[0]->ne[0];
+ const int64_t ne01 = op->src[0]->ne[1];
+
+ ggml_metal_mma_tiling res;
+ res.rt = ggml_metal_mul_mv_mma_rt(op);
+
+ const int64_t n_steps = ne00/ggml_metal_mul_mv_mma_k_step(type, res.rt);
+
+ int nsg = GGML_METAL_MMA_NSG_MANY_ROWS;
+ if (ne01 <= GGML_METAL_MMA_FEW_ROWS) {
+ nsg = GGML_METAL_MMA_NSG_FEW_ROWS;
+ } else if (ne01 <= GGML_METAL_MMA_MANY_ROWS) {
+ nsg = GGML_METAL_MMA_NSG_MID_ROWS;
+ }
+ res.nsg = ggml_metal_halve_to_limit(nsg, n_steps);
+
+ const int64_t nt_smem = GGML_METAL_MMA_SMEM_MAX/(int64_t) ggml_metal_mul_mv_mma_smem(res.nsg, 1, res.rt);
+ const int64_t nt_rows = ne01/(GGML_METAL_MIN_THREADGROUPS*GGML_METAL_MMA_TILE);
+ res.nt = ggml_metal_halve_to_limit(GGML_METAL_MMA_NT_MAX, std::min(nt_smem, nt_rows));
+
+ return res;
+}
+
+// the pipeline for the tiling, with fewer simdgroups when the device cannot run that many threads
+ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_mul_mv_mma_auto(ggml_metal_library_t lib, const ggml_tensor * op, bool add) {
+ ggml_metal_mma_tiling tiling = ggml_metal_op_mul_mat_mma_tiling(op);
+
+ auto pipeline = ggml_metal_library_get_pipeline_mul_mv_mma(lib, op, tiling.nsg, tiling.nt, tiling.rt, add);
+ while (tiling.nsg > 1 && ggml_metal_pipeline_max_theads_per_threadgroup(pipeline) < tiling.nsg*32) {
+ tiling.nsg /= 2;
+ pipeline = ggml_metal_library_get_pipeline_mul_mv_mma(lib, op, tiling.nsg, tiling.nt, tiling.rt, add);
+ }
+
+ return pipeline;
+}
+
+ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_mul_mv_mma(ggml_metal_library_t lib, const ggml_tensor * op, int nsg, int nt, int rt, bool add) {
+ char base[256];
+ char name[256];
+
+ const ggml_type tsrc0 = op->src[0]->type;
+ const ggml_type tsrc1 = op->src[1]->type;
+ const int ne12 = op->src[1]->ne[2];
+ const int r2 = ne12 / op->src[0]->ne[2];
+ const int r3 = op->src[1]->ne[3] / op->src[0]->ne[3];
+
+ GGML_ASSERT(ne12 <= INT16_MAX && r2 <= INT16_MAX && r3 <= INT16_MAX);
+
+ // the specialized kernels unroll over a compile-time row length
+ const int ne00 = ggml_metal_mul_mv_mma_kind(tsrc0, rt) != GGML_METAL_MMA_KIND_GEN ? op->src[0]->ne[0] : 0;
+
+ snprintf(base, 256, "kernel_mul_mv_mma_%s_%s_nt%d_rt%d", ggml_type_name(tsrc0), ggml_type_name(tsrc1), nt, rt);
+ snprintf(name, 256, "%s_nsg=%d_ne12=%d_r2=%d_r3=%d_ne00=%d_add=%d", base, nsg, ne12, r2, r3, ne00, add);
+
+ 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, ne00, FC_MUL_MV_MMA + 4);
+ ggml_metal_cv_set_int16(cv, nsg, FC_MUL_MV_MMA + 0);
+ ggml_metal_cv_set_int16(cv, (int16_t) ne12, FC_MUL_MV_MMA + 1);
+ ggml_metal_cv_set_int16(cv, (int16_t) r2, FC_MUL_MV_MMA + 2);
+ ggml_metal_cv_set_int16(cv, (int16_t) r3, FC_MUL_MV_MMA + 3);
+ ggml_metal_cv_set_bool (cv, add, FC_MUL_MV_MMA + 5);
+
+ res = ggml_metal_library_compile_pipeline(lib, base, name, cv);
+
+ ggml_metal_cv_free(cv);
+ }
+
+ res.nsg = nsg;
+ res.smem = ggml_metal_mul_mv_mma_smem(nsg, nt, rt);
+
+ res.nr0 = GGML_METAL_MMA_TILE*nt;
+ res.nr1 = GGML_METAL_MMA_TILE*rt;
+
+ return res;
+}
+
ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_mul_mm(ggml_metal_library_t lib, const ggml_tensor * op) {
char base[256];
char name[256];
diff --git a/ggml/src/ggml-metal/ggml-metal-device.h b/ggml/src/ggml-metal/ggml-metal-device.h
index 794fe979e..4012b9d5f 100644
--- a/ggml/src/ggml-metal/ggml-metal-device.h
+++ b/ggml/src/ggml-metal/ggml-metal-device.h
@@ -135,6 +135,8 @@ struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_rwkv
struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_gated_delta_net (ggml_metal_library_t lib, const struct ggml_tensor * op);
struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_solve_tri (ggml_metal_library_t lib, const struct ggml_tensor * op);
struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_mul_mv_ext (ggml_metal_library_t lib, const struct ggml_tensor * op, int nsg, int nxpsg, int r1ptg);
+struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_mul_mv_mma (ggml_metal_library_t lib, const struct ggml_tensor * op, int nsg, int nt, int rt, bool add);
+struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_mul_mv_mma_auto (ggml_metal_library_t lib, const struct ggml_tensor * op, bool add);
struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_mul_mm (ggml_metal_library_t lib, const struct ggml_tensor * op);
struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_mul_mv (ggml_metal_library_t lib, const struct ggml_tensor * op);
struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_mul_mm_id_map0 (ggml_metal_library_t lib, int ne02, int ne20);
@@ -302,8 +304,6 @@ struct ggml_metal_device_props {
bool use_residency_sets;
bool use_shared_buffers;
- bool supports_gpu_family_apple7;
-
enum ggml_metal_device_id device_id;
int gpu_family;
diff --git a/ggml/src/ggml-metal/ggml-metal-device.m b/ggml/src/ggml-metal/ggml-metal-device.m
index 951cb802a..827d92622 100644
--- a/ggml/src/ggml-metal/ggml-metal-device.m
+++ b/ggml/src/ggml-metal/ggml-metal-device.m
@@ -126,6 +126,7 @@ int ggml_metal_pipeline_max_theads_per_threadgroup(struct ggml_metal_pipeline_wi
X(FA_VEC_Q5_1, fa_vec_q5_1) \
X(FA_VEC_Q8_0, fa_vec_q8_0) \
X(MUL_MV, mul_mv) \
+ X(MUL_MV_MMA, mul_mv_mma) \
X(MUL_MM, mul_mm) \
X(QUANTIZE, quantize) \
X(SOFTMAX, softmax) \
@@ -1278,8 +1279,6 @@ ggml_metal_device_t ggml_metal_device_init(int device, int n_devices) {
dev->props.use_shared_buffers = true;
}
- dev->props.supports_gpu_family_apple7 = [dev->mtl_device supportsFamily:MTLGPUFamilyApple7];
-
dev->props.device_id = ggml_metal_device_id_parse([[dev->mtl_device name] UTF8String]);
dev->props.op_offload_min_batch_size = getenv("GGML_OP_OFFLOAD_MIN_BATCH") ? atoi(getenv("GGML_OP_OFFLOAD_MIN_BATCH")) : 32;
diff --git a/ggml/src/ggml-metal/ggml-metal-fusion.cpp b/ggml/src/ggml-metal/ggml-metal-fusion.cpp
index e55b01503..d78ca7c4b 100644
--- a/ggml/src/ggml-metal/ggml-metal-fusion.cpp
+++ b/ggml/src/ggml-metal/ggml-metal-fusion.cpp
@@ -1,8 +1,10 @@
#include "ggml-metal-fusion.h"
-#include "ggml-backend-impl.h"
+#include "ggml-metal-common.h"
#include "ggml-metal-device.h"
+#include "ggml-backend-impl.h"
+
#include <algorithm>
#include <cstddef>
#include <cstring>
@@ -85,17 +87,34 @@ static bool ggml_metal_fusion_same_buffer(const ggml_tensor * a, const ggml_tens
return ggml_metal_buffer_get_id(ca, a).metal == ggml_metal_buffer_get_id(cb, b).metal;
}
+// true if the memory of two tensors overlaps in the same Metal buffer
+static bool ggml_metal_fusion_overlap(const ggml_tensor * a, const ggml_tensor * b) {
+ ggml_backend_buffer_t ba = a->view_src ? a->view_src->buffer : a->buffer;
+ ggml_backend_buffer_t bb = b->view_src ? b->view_src->buffer : b->buffer;
+
+ const ggml_metal_buffer_id bid_a = ggml_metal_buffer_get_id((ggml_metal_buffer_t) ba->context, a);
+ const ggml_metal_buffer_id bid_b = ggml_metal_buffer_get_id((ggml_metal_buffer_t) bb->context, b);
+
+ if (bid_a.metal == nullptr || bid_a.metal != bid_b.metal) {
+ return false;
+ }
+
+ return bid_a.offs <= bid_b.offs
+ ? bid_b.offs - bid_a.offs < ggml_nbytes(a)
+ : bid_a.offs - bid_b.offs < ggml_nbytes(b);
+}
+
// ---- pattern checks ------------------------------------------------------
// NORM/RMS_NORM + MUL + ADD: the weight/bias of each fused step must match the norm input
// width, be contiguous rows, and the fused outputs must stay F32
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) {
+ 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);
@@ -140,12 +159,12 @@ static bool ggml_metal_fusion_check_norm(
// 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) {
+ 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);
@@ -173,12 +192,12 @@ static bool ggml_metal_fusion_check_ssm_conv_silu(
// 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) {
+ 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(gf);
GGML_UNUSED(node_idxs);
GGML_UNUSED(idx);
@@ -209,11 +228,11 @@ static bool ggml_metal_fusion_check_add_chain(
// attn scores view), so unlike the other patterns this is not an elision chain: the structural
// checks live entirely in this callback (unsafe = true).
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,
+ 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);
@@ -271,12 +290,12 @@ static bool ggml_metal_fusion_check_gdn_cache(
// MUL + SIN + SQR + MUL + ADD (snake activation)
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) {
+ 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);
@@ -345,12 +364,12 @@ static const std::vector<ggml_op> ops_topk_moe_norm_scale = {
};
static bool ggml_metal_fusion_check_topk_moe(
- const ggml_metal_fusion * fusion,
- const ggml_tensor * const * nodes,
- const ggml_cgraph * gf,
- const int * node_idxs,
- int idx,
- ggml_metal_fusion_mode mode) {
+ 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);
@@ -591,12 +610,12 @@ static bool ggml_metal_fusion_match_moe_reduce(
}
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) {
+ 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;
@@ -622,6 +641,71 @@ static bool ggml_metal_fusion_check_moe_reduce(
return true;
}
+// true if t is or views a tensor in a buffer marked as weights, such as a bias; the model loader marks its buffers before
+// any graph is optimized, and tensors in unmarked or not yet allocated buffers count as non-weights in both phases
+static bool ggml_metal_tensor_is_weight(const struct ggml_tensor * t) {
+ const ggml_tensor * base = t->view_src != NULL ? t->view_src : t;
+
+ return base->buffer != NULL && ggml_backend_buffer_get_usage(base->buffer) == GGML_BACKEND_BUFFER_USAGE_WEIGHTS;
+}
+
+static const struct ggml_tensor * ggml_metal_mul_mat_add_operand(const struct ggml_tensor * mm, const struct ggml_tensor * add) {
+ if (add->op != GGML_OP_ADD || (add->src[0] == mm) == (add->src[1] == mm)) {
+ return NULL;
+ }
+
+ const ggml_tensor * other = add->src[0] == mm ? add->src[1] : add->src[0];
+
+ const bool ok = other->type == GGML_TYPE_F32 && add->type == GGML_TYPE_F32 && !ggml_metal_tensor_is_weight(other);
+
+ return ok ? other : NULL;
+}
+
+static const struct ggml_tensor * ggml_metal_mul_mat_add_residual(const struct ggml_tensor * mm, const struct ggml_tensor * add) {
+ const ggml_tensor * res = ggml_metal_mul_mat_add_operand(mm, add);
+
+ const bool ok = res != NULL && ggml_are_same_shape(res, mm) &&
+ ggml_is_contiguous(res) && ggml_is_contiguous(mm) && ggml_is_contiguous(add);
+
+ return ok ? res : NULL;
+}
+
+// MUL_MAT + ADD of an f32 non-weight: the reorder packs it without reading row counts, so ubatch sizes share one order;
+// the encoder fuses only a same-shape residual in the few-row MMA store, which the sum may overlap only in place
+static bool ggml_metal_fusion_check_mul_mat_add(
+ 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(gf);
+ GGML_UNUSED(node_idxs);
+ GGML_UNUSED(idx);
+ GGML_UNUSED(fusion);
+
+ const ggml_tensor * mm = nodes[0];
+ const ggml_tensor * add = nodes[1];
+
+ if (ggml_metal_mul_mat_add_operand(mm, add) == nullptr ||
+ !ggml_metal_op_mul_mat_may_use_mma(mm)) {
+ return false;
+ }
+
+ if (mode == GGML_METAL_FUSION_STRUCTURAL) {
+ return true;
+ }
+
+ const ggml_tensor * res = ggml_metal_mul_mat_add_residual(mm, add);
+
+ if (res == nullptr || !ggml_metal_op_mul_mat_use_mma(mm)) {
+ return false;
+ }
+
+ return !ggml_metal_fusion_overlap(add, mm->src[0]) && !ggml_metal_fusion_overlap(add, mm->src[1]) &&
+ (add->data == res->data || !ggml_metal_fusion_overlap(add, res));
+}
+
// ---- patterns ------------------------------------------------------------
static const std::vector<ggml_op> ops_norm_mul = { GGML_OP_NORM, GGML_OP_MUL };
@@ -671,6 +755,8 @@ static const std::vector<ggml_op> ops_moe_reduce_8 = {
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_mul_mat_add = { GGML_OP_MUL_MAT, GGML_OP_ADD };
+
static const std::vector<ggml_metal_fusion> ggml_metal_fusions = {
{ GGML_METAL_FUSION_NORM_MUL, ops_norm_mul, {}, false, ggml_metal_fusion_check_norm },
{ GGML_METAL_FUSION_NORM_MUL_ADD, ops_norm_mul_add, {}, false, ggml_metal_fusion_check_norm },
@@ -698,6 +784,7 @@ static const std::vector<ggml_metal_fusion> ggml_metal_fusions = {
{ GGML_METAL_FUSION_MOE_REDUCE, ops_moe_reduce_7, {}, true, ggml_metal_fusion_check_moe_reduce },
{ GGML_METAL_FUSION_MOE_REDUCE, ops_moe_reduce_8, {}, true, ggml_metal_fusion_check_moe_reduce },
{ GGML_METAL_FUSION_SSM_CONV_SILU, ops_ssm_conv_silu, {}, false, ggml_metal_fusion_check_ssm_conv_silu },
+ { GGML_METAL_FUSION_MUL_MAT_ADD, ops_mul_mat_add, {}, false, ggml_metal_fusion_check_mul_mat_add },
};
// ---- alloc deps -----------------------------------------------------------
diff --git a/ggml/src/ggml-metal/ggml-metal-fusion.h b/ggml/src/ggml-metal/ggml-metal-fusion.h
index 6b139a69b..5d994c249 100644
--- a/ggml/src/ggml-metal/ggml-metal-fusion.h
+++ b/ggml/src/ggml-metal/ggml-metal-fusion.h
@@ -2,8 +2,9 @@
//
// every fusable subgraph is declared exactly once as a ggml_metal_fusion entry in
// the table in ggml-metal-fusion.cpp. both the graph optimizer (ggml_metal_fusion_max)
-// and the op encoders (ggml_metal_fusion_next) consult this same table, so the two
-// phases can never disagree about what can be fused.
+// and the op encoders (ggml_metal_fusion_next) consult this same table with the same device
+// properties. a check that passes in FULL mode also passes in STRUCTURAL mode, so the encoders
+// fuse only groups that the optimizer may pack.
#pragma once
@@ -15,13 +16,15 @@
extern "C" {
#endif
+struct ggml_metal_device_props;
+
// the maximum number of nodes that can be fused in a single kernel
// (also the maximum length of a packed fusion group during graph optimization)
#define GGML_METAL_FUSION_MAX 16
typedef enum ggml_metal_fusion_mode {
- // structural checks only; used by the graph optimizer, at which point the graph
- // tensors are not allocated yet, so buffer placement cannot be verified
+ // structural checks for the graph optimizer, which runs before the graph tensors are allocated (weights
+ // already are); a check may skip conditions that differ between batch sizes, and so accept more than FULL
GGML_METAL_FUSION_STRUCTURAL = 0,
// full checks, including buffer placement; used by the op encoders
GGML_METAL_FUSION_FULL,
@@ -39,6 +42,7 @@ typedef enum ggml_metal_fusion_id {
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_MUL_MAT_ADD, // MUL_MAT + ADD (residual added in the few-row MMA store)
} ggml_metal_fusion_id;
struct ggml_metal_fusion; // defined in ggml-metal-fusion.cpp
@@ -79,8 +83,8 @@ void ggml_metal_fusion_info_stats_reset( struct ggml_metal_fusion_info * fi
int ggml_metal_fusion_info_stats_get (const struct ggml_metal_fusion_info * finfo, const char ** labels, uint64_t * counts, int n);
void ggml_metal_fusion_info_labels_init( struct ggml_metal_fusion_info * finfo);
-// compute phase: longest fusion starting at idx (a position in node_idxs) that matches in `mode`.
-// returns the matching pattern (nullptr if no fusion) and sets *n_out to the number of nodes consumed.
+// compute phase: longest fusion starting at idx (a position in node_idxs) that matches in `mode` on a device with
+// props. returns the matching pattern (nullptr if no fusion) and sets *n_out to the number of nodes consumed.
const ggml_metal_fusion * ggml_metal_fusion_next(
const struct ggml_cgraph * gf,
const int * node_idxs,
@@ -90,7 +94,7 @@ const ggml_metal_fusion * ggml_metal_fusion_next(
int * n_out);
// optimize phase: maximum number of nodes starting at idx (a raw sequential graph index) that
-// could be fused, chaining patterns back-to-back. returns at least 1.
+// could be fused on a device with props, chaining patterns back-to-back. returns at least 1.
int ggml_metal_fusion_max(const struct ggml_cgraph * gf, int idx);
#ifdef __cplusplus
diff --git a/ggml/src/ggml-metal/ggml-metal-impl.h b/ggml/src/ggml-metal/ggml-metal-impl.h
index a5bc79f57..e7c895631 100644
--- a/ggml/src/ggml-metal/ggml-metal-impl.h
+++ b/ggml/src/ggml-metal/ggml-metal-impl.h
@@ -123,6 +123,7 @@
#define FC_PAD 2100
#define FC_FLASH_ATTN_EXT_TENSOR 2200
#define FC_LIGHTNING_INDEXER 2200
+#define FC_MUL_MV_MMA 2300
// op-specific constants
#define OP_FLASH_ATTN_EXT_NQPSG 8
@@ -215,6 +216,7 @@ typedef struct {
uint64_t nb2;
uint64_t nb3;
int32_t dim;
+ int32_t nc0;
} ggml_metal_kargs_concat;
typedef struct {
diff --git a/ggml/src/ggml-metal/ggml-metal-ops.cpp b/ggml/src/ggml-metal/ggml-metal-ops.cpp
index 4a7da2d07..dfe46bac6 100644
--- a/ggml/src/ggml-metal/ggml-metal-ops.cpp
+++ b/ggml/src/ggml-metal/ggml-metal-ops.cpp
@@ -15,6 +15,10 @@
#include <limits>
#include <cmath>
+// threadgroups needed to fill the GPU
+// TODO: dedup
+static constexpr int64_t GGML_METAL_MIN_THREADGROUPS = 128;
+
static ggml_metal_buffer_id ggml_metal_get_buffer_id(const ggml_tensor * t) {
if (!t) {
return { nullptr, 0 };
@@ -584,6 +588,27 @@ int ggml_metal_op_concat(ggml_metal_op_t ctx, int idx) {
ne0_arg = ne0/blck;
}
+ int nth = std::min(256, ne0_arg);
+
+ // when rows are small, we can batch them together in a single threadgroup
+ int nrptg = 1;
+ if (nth < 256) {
+ nrptg = std::min((256 + nth - 1) / nth, ne1);
+ if (nrptg * nth > 256) {
+ nrptg = 256 / nth;
+ }
+ }
+
+ const int nw0 = (ne1 + nrptg - 1) / nrptg;
+
+ // split long rows across threadgroups when there are too few rows to fill the GPU
+ const int64_t n_rows = (int64_t) nw0*ne2*ne3;
+
+ int nc0 = 1;
+ if (nrptg == 1) {
+ nc0 = (int) std::max<int64_t>(1, std::min<int64_t>((ne0_arg + nth - 1)/nth, (GGML_METAL_MIN_THREADGROUPS + n_rows - 1)/n_rows));
+ }
+
ggml_metal_kargs_concat args = {
/*.ne00 =*/ ne00_arg,
/*.ne01 =*/ ne01,
@@ -610,6 +635,7 @@ int ggml_metal_op_concat(ggml_metal_op_t ctx, int idx) {
/*.nb2 =*/ nb2,
/*.nb3 =*/ nb3,
/*.dim =*/ dim,
+ /*.nc0 =*/ nc0,
};
auto pipeline = ggml_metal_library_get_pipeline_concat(lib, op->type);
@@ -620,20 +646,7 @@ int ggml_metal_op_concat(ggml_metal_op_t ctx, int idx) {
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);
- int nth = std::min(256, ne0_arg);
-
- // when rows are small, we can batch them together in a single threadgroup
- int nrptg = 1;
- if (nth < 256) {
- nrptg = std::min((256 + nth - 1) / nth, ne1);
- if (nrptg * nth > 256) {
- nrptg = 256 / nth;
- }
- }
-
- const int nw0 = (ne1 + nrptg - 1) / nrptg;
-
- ggml_metal_encoder_dispatch_threadgroups(enc, nw0, ne2, ne3, nth, nrptg, 1);
+ ggml_metal_encoder_dispatch_threadgroups(enc, nw0*nc0, ne2, ne3, nth, nrptg, 1);
return 1;
}
@@ -2421,6 +2434,157 @@ int ggml_metal_op_pool_2d(ggml_metal_op_t ctx, int idx) {
return 1;
}
+// the number of nodes from idx on that the fusion table fuses as pattern id into one dispatch, 1 if it fuses none
+static int ggml_metal_op_try_fusion(ggml_metal_op_t ctx, int idx, ggml_metal_fusion_id id) {
+ if (!ctx->use_fusion()) {
+ return 1;
+ }
+
+ int n = 1;
+ const ggml_metal_fusion * fusion = ctx->can_fuse(idx, GGML_METAL_FUSION_FULL, &n);
+ if (fusion == nullptr || ggml_metal_fusion_get_id(fusion) != id) {
+ return 1;
+ }
+
+ ctx->count_fusions(fusion);
+
+ if (ggml_metal_fusion_info_debug(ctx->finfo) > 1) {
+ GGML_LOG_DEBUG("%s: fuse: %s to %s, %d nodes\n", __func__, ggml_op_name(ctx->node(idx)->op), ggml_op_name(ctx->node(idx + n - 1)->op), n);
+ }
+
+ return n;
+}
+
+static int ggml_metal_op_mul_mat_mma(ggml_metal_op_t ctx, int idx) {
+ ggml_tensor * op = ctx->node(idx);
+
+ ggml_metal_library_t lib = ctx->lib;
+ ggml_metal_encoder_t enc = ctx->enc;
+
+ GGML_TENSOR_LOCALS( int32_t, ne0, op->src[0], ne);
+ GGML_TENSOR_LOCALS(uint64_t, nb0, op->src[0], nb);
+ GGML_TENSOR_LOCALS( int32_t, ne1, op->src[1], ne);
+ GGML_TENSOR_LOCALS(uint64_t, nb1, op->src[1], nb);
+ GGML_TENSOR_LOCALS( int32_t, ne, op, ne);
+
+ // the MMA store adds the residual when the table fuses the ADD after this mat-mul
+ const int n_fuse = ggml_metal_op_try_fusion(ctx, idx, GGML_METAL_FUSION_MUL_MAT_ADD);
+ const bool fuse_add = n_fuse > 1;
+
+ const ggml_tensor * dst = op;
+ const ggml_tensor * res = dst;
+
+ if (fuse_add) {
+ dst = ctx->node(idx + n_fuse - 1);
+ res = dst->src[0]->op == GGML_OP_MUL_MAT ? dst->src[1] : dst->src[0];
+ }
+
+ auto pipeline = ggml_metal_library_get_pipeline_mul_mv_mma_auto(lib, op, fuse_add);
+
+ ggml_metal_kargs_mul_mv_ext args = {
+ /*.ne00 =*/ ne00,
+ /*.ne01 =*/ ne01,
+ /*.ne02 =*/ ne02,
+ /*.nb00 =*/ nb00,
+ /*.nb01 =*/ nb01,
+ /*.nb02 =*/ nb02,
+ /*.nb03 =*/ nb03,
+ /*.ne10 =*/ ne10,
+ /*.ne11 =*/ ne11,
+ /*.ne12 =*/ ne12,
+ /*.nb10 =*/ nb10,
+ /*.nb11 =*/ nb11,
+ /*.nb12 =*/ nb12,
+ /*.nb13 =*/ nb13,
+ /*.ne0 =*/ ne0,
+ /*.ne1 =*/ ne1,
+ /*.r2 =*/ (int16_t) (ne12/ne02),
+ /*.r3 =*/ (int16_t) (ne13/ne03),
+ };
+
+ 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(dst), 3);
+ ggml_metal_encoder_set_buffer (enc, ggml_metal_get_buffer_id(res), 4);
+
+ ggml_metal_encoder_set_threadgroup_memory_size(enc, pipeline.smem, 0);
+
+ const int rows0 = pipeline.nr0;
+ const int rows1 = pipeline.nr1;
+
+ ggml_metal_encoder_dispatch_threadgroups(enc, (ne01 + rows0 - 1)/rows0, (ne11 + rows1 - 1)/rows1, ne12*ne13, 32, pipeline.nsg, 1);
+
+ return n_fuse;
+}
+
+// the generic mat-vec kernel, or its 2-row Q4_0 variant if nc
+static int ggml_metal_op_mul_mat_mv(ggml_metal_op_t ctx, int idx) {
+ ggml_tensor * op = ctx->node(idx);
+
+ ggml_metal_library_t lib = ctx->lib;
+ ggml_metal_encoder_t enc = ctx->enc;
+
+ GGML_TENSOR_LOCALS( int32_t, ne0, op->src[0], ne);
+ GGML_TENSOR_LOCALS(uint64_t, nb0, op->src[0], nb);
+ GGML_TENSOR_LOCALS( int32_t, ne1, op->src[1], ne);
+ GGML_TENSOR_LOCALS(uint64_t, nb1, op->src[1], nb);
+ GGML_TENSOR_LOCALS( int32_t, ne, op, ne);
+
+ const int16_t r2 = ne12/ne02;
+ const int16_t r3 = ne13/ne03;
+
+ auto pipeline = ggml_metal_library_get_pipeline_mul_mv(lib, op);
+
+ const int nr0 = pipeline.nr0;
+ const int nr1 = pipeline.nr1;
+ const int nsg = pipeline.nsg;
+
+ const size_t smem = pipeline.smem;
+
+ ggml_metal_kargs_mul_mv args = {
+ /*.ne00 =*/ ne00,
+ /*.ne01 =*/ ne01,
+ /*.ne02 =*/ ne02,
+ /*.nb00 =*/ nb00,
+ /*.nb01 =*/ nb01,
+ /*.nb02 =*/ nb02,
+ /*.nb03 =*/ nb03,
+ /*.ne10 =*/ ne10,
+ /*.ne11 =*/ ne11,
+ /*.ne12 =*/ ne12,
+ /*.nb10 =*/ nb10,
+ /*.nb11 =*/ nb11,
+ /*.nb12 =*/ nb12,
+ /*.nb13 =*/ nb13,
+ /*.ne0 =*/ ne0,
+ /*.ne1 =*/ ne1,
+ /*.nr0 =*/ nr0,
+ /*.r2 =*/ r2,
+ /*.r3 =*/ r3,
+ };
+
+ 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_threadgroup_memory_size(enc, smem, 0);
+
+ if (op->src[0]->type == GGML_TYPE_F32 ||
+ op->src[0]->type == GGML_TYPE_F16 ||
+ op->src[0]->type == GGML_TYPE_BF16 ||
+ op->src[0]->type == GGML_TYPE_Q8_0) {
+ ggml_metal_encoder_dispatch_threadgroups(enc, ((ne01 + nr0 - 1)/(nr0)), ((ne11 + nr1 - 1)/nr1), ne12*ne13, 32, nsg, 1);
+ } else {
+ ggml_metal_encoder_dispatch_threadgroups(enc, ((ne01 + nr0*nsg - 1)/(nr0*nsg)), ((ne11 + nr1 - 1)/nr1), ne12*ne13, 32, nsg, 1);
+ }
+
+ return 1;
+}
+
int ggml_metal_op_mul_mat(ggml_metal_op_t ctx, int idx) {
ggml_tensor * op = ctx->node(idx);
@@ -2433,6 +2597,10 @@ int ggml_metal_op_mul_mat(ggml_metal_op_t ctx, int idx) {
return ggml_metal_op_fwht(ctx, idx);
}
+ if (props_dev->has_simdgroup_mm && ggml_metal_op_mul_mat_use_mma(op)) {
+ return ggml_metal_op_mul_mat_mma(ctx, idx);
+ }
+
GGML_TENSOR_LOCALS( int32_t, ne0, op->src[0], ne);
GGML_TENSOR_LOCALS(uint64_t, nb0, op->src[0], nb);
GGML_TENSOR_LOCALS( int32_t, ne1, op->src[1], ne);
@@ -2597,52 +2765,7 @@ int ggml_metal_op_mul_mat(ggml_metal_op_t ctx, int idx) {
ggml_metal_encoder_dispatch_threadgroups(enc, ((ne11 + nr1 - 1) / nr1), ((ne01 + nr0 - 1) / nr0), ne12 * ne13, 32, nsg, 1);
} else {
- auto pipeline = ggml_metal_library_get_pipeline_mul_mv(lib, op);
-
- const int nr0 = pipeline.nr0;
- const int nr1 = pipeline.nr1;
- const int nsg = pipeline.nsg;
-
- const size_t smem = pipeline.smem;
-
- ggml_metal_kargs_mul_mv args = {
- /*.ne00 =*/ ne00,
- /*.ne01 =*/ ne01,
- /*.ne02 =*/ ne02,
- /*.nb00 =*/ nb00,
- /*.nb01 =*/ nb01,
- /*.nb02 =*/ nb02,
- /*.nb03 =*/ nb03,
- /*.ne10 =*/ ne10,
- /*.ne11 =*/ ne11,
- /*.ne12 =*/ ne12,
- /*.nb10 =*/ nb10,
- /*.nb11 =*/ nb11,
- /*.nb12 =*/ nb12,
- /*.nb13 =*/ nb13,
- /*.ne0 =*/ ne0,
- /*.ne1 =*/ ne1,
- /*.nr0 =*/ nr0,
- /*.r2 =*/ r2,
- /*.r3 =*/ r3,
- };
-
- 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_threadgroup_memory_size(enc, smem, 0);
-
- if (op->src[0]->type == GGML_TYPE_F32 ||
- op->src[0]->type == GGML_TYPE_F16 ||
- op->src[0]->type == GGML_TYPE_BF16 ||
- op->src[0]->type == GGML_TYPE_Q8_0) {
- ggml_metal_encoder_dispatch_threadgroups(enc, ((ne01 + nr0 - 1)/(nr0)), ((ne11 + nr1 - 1)/nr1), ne12*ne13, 32, nsg, 1);
- } else {
- ggml_metal_encoder_dispatch_threadgroups(enc, ((ne01 + nr0*nsg - 1)/(nr0*nsg)), ((ne11 + nr1 - 1)/nr1), ne12*ne13, 32, nsg, 1);
- }
+ return ggml_metal_op_mul_mat_mv(ctx, idx);
}
return 1;
@@ -3968,7 +4091,9 @@ int ggml_metal_op_bin(ggml_metal_op_t ctx, int idx) {
int n_fuse = 1;
const ggml_metal_fusion * fusion = nullptr;
- if (ctx->use_fusion()) {
+ const bool use_fusion = ctx->use_fusion();
+
+ if (use_fusion) {
int n = 1;
fusion = ctx->can_fuse(idx, GGML_METAL_FUSION_FULL, &n);
n_fuse = n;
@@ -3991,8 +4116,6 @@ int ggml_metal_op_bin(ggml_metal_op_t ctx, int idx) {
ggml_metal_library_t lib = ctx->lib;
ggml_metal_encoder_t enc = ctx->enc;
- const bool use_fusion = ctx->use_fusion();
-
const int debug_fusion = ggml_metal_fusion_info_debug(ctx->finfo);
GGML_TENSOR_LOCALS( int32_t, ne0, op->src[0], ne);
diff --git a/ggml/src/ggml-metal/kernels/mul_mv_mma.metal b/ggml/src/ggml-metal/kernels/mul_mv_mma.metal
new file mode 100644
index 000000000..fcf981a93
--- /dev/null
+++ b/ggml/src/ggml-metal/kernels/mul_mv_mma.metal
@@ -0,0 +1,608 @@
+#include "common.h"
+#include "dequantize.h"
+
+constant short FC_mul_mv_mma_nsg [[function_constant(FC_MUL_MV_MMA + 0)]];
+constant short FC_mul_mv_mma_ne12 [[function_constant(FC_MUL_MV_MMA + 1)]];
+constant short FC_mul_mv_mma_r2 [[function_constant(FC_MUL_MV_MMA + 2)]];
+constant short FC_mul_mv_mma_r3 [[function_constant(FC_MUL_MV_MMA + 3)]];
+constant int FC_mul_mv_mma_ne00 [[function_constant(FC_MUL_MV_MMA + 4)]];
+constant bool FC_mul_mv_mma_add [[function_constant(FC_MUL_MV_MMA + 5)]];
+
+// a lane of a few-row MMA tile: A fragment row fm, B fragment columns fn and fn + 1
+struct mul_mv_mma_tile {
+ short fm;
+ short fn;
+ int i01;
+ int i11;
+ int i1m;
+ uint64_t offset0;
+ uint64_t offset1;
+};
+
+// the A fragment row and the first B fragment column that lane l holds in an 8x8 simdgroup matrix
+inline short mul_mv_mma_lane_fm(ushort l) { return ((l/4) & 4) + ((l/2) % 4); }
+inline short mul_mv_mma_lane_fn(ushort l) { return ((l/4) & 2)*2 + (l % 2)*2; }
+
+template<short NT, short RT>
+inline mul_mv_mma_tile mul_mv_mma_tile_init(constant ggml_metal_kargs_mul_mv_ext & args, uint3 tgpig, ushort tiisg) {
+ mul_mv_mma_tile tile;
+ tile.fm = mul_mv_mma_lane_fm(tiisg);
+ tile.fn = mul_mv_mma_lane_fn(tiisg);
+ tile.i01 = tgpig.x*(8*NT);
+ tile.i11 = tgpig.y*(8*RT);
+ tile.i1m = tgpig.z;
+
+ const int i12 = tile.i1m%FC_mul_mv_mma_ne12;
+ const int i13 = tile.i1m/FC_mul_mv_mma_ne12;
+
+ tile.offset0 = (i12/FC_mul_mv_mma_r2)*args.nb02 + (i13/FC_mul_mv_mma_r3)*args.nb03;
+ tile.offset1 = i12*args.nb12 + i13*args.nb13;
+
+ return tile;
+}
+
+// the src0 row of A fragment row fm in 8-row tile t, clamped to the last row
+inline device const char * mul_mv_mma_src0_row(
+ thread const mul_mv_mma_tile & tile, constant ggml_metal_kargs_mul_mv_ext & args, device const char * src0, short t) {
+ const int r = min(tile.i01 + 8*t + tile.fm, args.ne01 - 1);
+ return src0 + tile.offset0 + (uint64_t) r*args.nb01;
+}
+
+// the src1 row of B fragment column fn + e in 8-row tile rt, clamped to the last row
+inline device const float * mul_mv_mma_src1_row(
+ thread const mul_mv_mma_tile & tile, constant ggml_metal_kargs_mul_mv_ext & args, device const char * src1, short rt, short e) {
+ const int r = min(tile.i11 + 8*rt + tile.fn + e, args.ne11 - 1);
+ return (device const float *) (src1 + tile.offset1 + (uint64_t) r*args.nb11);
+}
+
+constexpr constant static ushort mma_f16_1024_bits = 0x6400;
+constexpr constant static half mma_f16_1024 = 1024.0h;
+
+// the halves 1024 + q for integers q < 1024: exact normal values, unlike the subnormal q*2^-24 that Metal may flush to zero
+inline half2 mul_mv_mma_1024_plus(ushort2 q) {
+ return as_type<half2>(q | mma_f16_1024_bits);
+}
+
+// adds up the K slices of the NSG simdgroups for an 8*NT x 8*RT output tile and writes it, plus the residual in src2 if FC_mul_mv_mma_add
+template<short NT, short RT>
+inline void mul_mv_mma_store(
+ thread float (&acc)[RT][NT][2],
+ constant ggml_metal_kargs_mul_mv_ext & args,
+ device const char * src2,
+ device char * dst,
+ threadgroup char * shmem,
+ thread const mul_mv_mma_tile & tile, ushort tiisg, ushort sgitg) {
+ const short NSG = FC_mul_mv_mma_nsg;
+
+ threadgroup float * red = (threadgroup float *) shmem;
+
+ FOR_UNROLL (short rt = 0; rt < RT; ++rt) {
+ FOR_UNROLL (short t = 0; t < NT; ++t) {
+ red[((sgitg*RT + rt)*NT + t)*64 + 2*tiisg + 0] = acc[rt][t][0];
+ red[((sgitg*RT + rt)*NT + t)*64 + 2*tiisg + 1] = acc[rt][t][1];
+ }
+ }
+
+ threadgroup_barrier(mem_flags::mem_threadgroup);
+
+ device float * dst_f32 = (device float *) dst + (uint64_t) tile.i1m*args.ne0*args.ne1;
+ device const float * res_f32 = (device const float *) src2 + (uint64_t) tile.i1m*args.ne0*args.ne1;
+
+ for (short idx = sgitg*32 + tiisg; idx < RT*NT*64; idx += NSG*32) {
+ float sum = 0.0f;
+ for (short sg = 0; sg < NSG; ++sg) {
+ sum += red[sg*(RT*NT*64) + idx];
+ }
+
+ const short rt = idx/(NT*64);
+ const short t = (idx/64) % NT;
+ const short l = (idx % 64)/2;
+ const short e = idx % 2;
+
+ const int r0 = tile.i01 + 8*t + mul_mv_mma_lane_fm(l);
+ const int r1 = tile.i11 + 8*rt + mul_mv_mma_lane_fn(l) + e;
+
+ if (r0 < args.ne01 && r1 < args.ne11) {
+ const uint64_t i = (uint64_t) r1*args.ne0 + r0;
+ dst_f32[i] = FC_mul_mv_mma_add ? sum + res_f32[i] : sum;
+ }
+ }
+}
+
+// the per-type parts of kernel_mul_mv_mma_blk for 32-weight blocks. a src1 block splits into halves b0 and b1,
+// and MMA step s uses half b1 when b1_step(s) and the .y value of a pair when y_step(s)
+// q4_0: A lane (m, j) holds qs ushorts j and j + 1 (j even); a high nibble stays in place as 16*q, so b1 is divided by 16.
+// B lane k = fm holds src1 values 2*k, 2*k + 1 (b0, low nibbles) and 2*k + 16, 2*k + 17 (b1, high nibbles) of a block
+struct mul_mv_mma_q4_0 {
+ typedef block_q4_0 block;
+ typedef ushort2 quants;
+
+ // weights per block, and the float2 offset of b1 in a src1 block
+ enum { qk = QK4_0, b1 = QK4_0/4 };
+
+ static short a_off(short fn) { return 1 + fn; }
+ static short b_off(short fm) { return 2*fm; }
+
+ static quants load(device const ushort * qs) { return ushort2(qs[0], qs[1]); }
+ static quants prep(quants q) { return q; }
+ static float2 prep_b1(float2 v) {
+ constexpr float hi_scale = 1.0f/16;
+ return v*hi_scale;
+ }
+
+ static bool b1_step(short s) { return s % 2 != 0; }
+ static bool y_step (short s) { return s >= 2; }
+
+ static half2 frag(quants q, short s) {
+ constexpr ushort lo_mask = 0x000F;
+ constexpr ushort hi_mask = 0x00F0;
+ constexpr half lo_zero = 8.0h;
+ constexpr half hi_zero = 16*lo_zero;
+
+ const ushort2 qq = s < 2 ? q : q >> 8;
+ return s % 2 == 0 ? mul_mv_mma_1024_plus(qq & lo_mask) - (mma_f16_1024 + lo_zero) : mul_mv_mma_1024_plus(qq & hi_mask) - (mma_f16_1024 + hi_zero);
+ }
+};
+
+// q8_0: A lane (m, j) holds qs bytes 4*j .. 4*j + 7 (j even); flipping the sign bit of a quant byte gives the unsigned q + 128.
+// B lane k = fm holds src1 values b, b + 1 (b0) and b + 4, b + 5 (b1) of a block, b = 8*(k/2) + 2*(k%2)
+struct mul_mv_mma_q8_0 {
+ typedef block_q8_0 block;
+ typedef ushort4 quants;
+
+ // weights per block, and the float2 offset of b1 in a src1 block
+ enum { qk = QK8_0, b1 = 2 };
+
+ static short a_off(short fn) { return 1 + 2*fn; }
+ static short b_off(short fm) { return 8*(fm/2) + 2*(fm%2); }
+
+ static quants load(device const ushort * qs) { return ushort4(qs[0], qs[1], qs[2], qs[3]); }
+ static quants prep(quants q) {
+ constexpr ushort sign_bits = 0x8080;
+ return q ^ ushort4(sign_bits);
+ }
+ static float2 prep_b1(float2 v) { return v; }
+
+ static bool b1_step(short s) { return s >= 2; }
+ static bool y_step (short s) { return s % 2 != 0; }
+
+ static half2 frag(quants q, short s) {
+ constexpr ushort byte_mask = 0x00FF;
+ constexpr half q_bias = mma_f16_1024 + 128.0h;
+
+ const ushort2 w = s < 2 ? q.xy : q.zw;
+ const ushort2 qq = s % 2 == 0 ? w : w >> 8;
+ return mul_mv_mma_1024_plus(qq & byte_mask) - q_bias;
+ }
+};
+
+template<typename Q, short NT>
+inline void load_mma_blk_a(device const ushort * const x[NT], int off, short fn, thread typename Q::quants * q, thread float * d) {
+ FOR_UNROLL (short t = 0; t < NT; ++t) {
+ device const ushort * qs = x[t] + off;
+ q[t] = Q::load(qs);
+ d[t] = as_type<half>(*(qs - Q::a_off(fn)));
+ }
+}
+
+template<typename Q, short RT>
+inline void load_mma_blk_b(device const float2 * const y[RT][2], int ib, thread float2 (*b0)[2], thread float2 (*b1)[2]) {
+ FOR_UNROLL (short rt = 0; rt < RT; ++rt) {
+ FOR_UNROLL (short e = 0; e < 2; ++e) {
+ b0[rt][e] = y[rt][e][ib*(Q::qk/2)];
+ b1[rt][e] = y[rt][e][ib*(Q::qk/2) + Q::b1];
+ }
+ }
+}
+
+// few-row mat-mat (2..16 src1 rows) on 8x8 simdgroup matrices for 32-weight block types: a threadgroup reads each weight once
+// for 8*NT src0 rows x 8*RT src1 rows, and its NSG simdgroups split K
+template<short NT, short RT, typename Q>
+kernel void kernel_mul_mv_mma_blk(
+ constant ggml_metal_kargs_mul_mv_ext & args,
+ device const char * src0,
+ device const char * src1,
+ device char * dst,
+ device const char * src2,
+ threadgroup char * shmem [[threadgroup(0)]],
+ uint3 tgpig[[threadgroup_position_in_grid]],
+ ushort tiisg[[thread_index_in_simdgroup]],
+ ushort sgitg[[simdgroup_index_in_threadgroup]]) {
+ const short NSG = FC_mul_mv_mma_nsg;
+
+ const mul_mv_mma_tile tile = mul_mv_mma_tile_init<NT, RT>(args, tgpig, tiisg);
+
+ device const ushort * x[NT];
+ FOR_UNROLL (short t = 0; t < NT; ++t) {
+ x[t] = (device const ushort *) mul_mv_mma_src0_row(tile, args, src0, t) + Q::a_off(tile.fn);
+ }
+
+ device const float2 * y[RT][2];
+ FOR_UNROLL (short rt = 0; rt < RT; ++rt) {
+ FOR_UNROLL (short e = 0; e < 2; ++e) {
+ y[rt][e] = (device const float2 *) (mul_mv_mma_src1_row(tile, args, src1, rt, e) + Q::b_off(tile.fm));
+ }
+ }
+
+ float acc[RT][NT][2] = {};
+
+ const int nb = FC_mul_mv_mma_ne00/Q::qk;
+
+ // a block is d, then the quants
+ constexpr short us_blk = sizeof(typename Q::block)/2;
+
+ typename Q::quants q[NT];
+ float d[NT];
+ float2 b0[RT][2];
+ float2 b1[RT][2];
+
+ const int ib0 = min((int) sgitg, nb - 1);
+ load_mma_blk_a<Q, NT>(x, ib0*us_blk, tile.fn, q, d);
+ load_mma_blk_b<Q, RT>(y, ib0, b0, b1);
+
+ for (int ib = sgitg; ib < nb; ib += NSG) {
+ typename Q::quants qc[NT];
+ float dc[NT];
+ float2 b0c[RT][2];
+ float2 b1c[RT][2];
+
+ FOR_UNROLL (short t = 0; t < NT; ++t) {
+ qc[t] = Q::prep(q[t]);
+ dc[t] = d[t];
+ }
+ FOR_UNROLL (short rt = 0; rt < RT; ++rt) {
+ FOR_UNROLL (short e = 0; e < 2; ++e) {
+ b0c[rt][e] = b0[rt][e];
+ b1c[rt][e] = Q::prep_b1(b1[rt][e]);
+ }
+ }
+
+ const int ibn = min(ib + NSG, nb - 1);
+ load_mma_blk_a<Q, NT>(x, ibn*us_blk, tile.fn, q, d);
+ load_mma_blk_b<Q, RT>(y, ibn, b0, b1);
+
+ simdgroup_float8x8 mp[RT][NT];
+ FOR_UNROLL (short rt = 0; rt < RT; ++rt) {
+ FOR_UNROLL (short t = 0; t < NT; ++t) {
+ mp[rt][t] = make_filled_simdgroup_matrix<float, 8>(0.0f);
+ }
+ }
+
+ FOR_UNROLL (short s = 0; s < 4; ++s) {
+ simdgroup_float8x8 mb[RT];
+ FOR_UNROLL (short rt = 0; rt < RT; ++rt) {
+ const float2 v0 = Q::b1_step(s) ? b1c[rt][0] : b0c[rt][0];
+ const float2 v1 = Q::b1_step(s) ? b1c[rt][1] : b0c[rt][1];
+ mb[rt].thread_elements()[0] = Q::y_step(s) ? v0.y : v0.x;
+ mb[rt].thread_elements()[1] = Q::y_step(s) ? v1.y : v1.x;
+ }
+
+ FOR_UNROLL (short t = 0; t < NT; ++t) {
+ const half2 h = Q::frag(qc[t], s);
+
+ simdgroup_half8x8 ma;
+ ma.thread_elements()[0] = h.x;
+ ma.thread_elements()[1] = h.y;
+
+ FOR_UNROLL (short rt = 0; rt < RT; ++rt) {
+ simdgroup_multiply_accumulate(mp[rt][t], ma, mb[rt], mp[rt][t]);
+ }
+ }
+ }
+
+ FOR_UNROLL (short t = 0; t < NT; ++t) {
+ FOR_UNROLL (short rt = 0; rt < RT; ++rt) {
+ acc[rt][t][0] = fma(dc[t], mp[rt][t].thread_elements()[0], acc[rt][t][0]);
+ acc[rt][t][1] = fma(dc[t], mp[rt][t].thread_elements()[1], acc[rt][t][1]);
+ }
+ }
+ }
+
+ mul_mv_mma_store<NT, RT>(acc, args, src2, dst, shmem, tile, tiisg, sgitg);
+}
+
+typedef decltype(kernel_mul_mv_mma_blk<4, 1, mul_mv_mma_q4_0>) mul_mv_mma_t;
+
+template [[host_name("kernel_mul_mv_mma_q4_0_f32_nt1_rt1")]] kernel mul_mv_mma_t kernel_mul_mv_mma_blk<1, 1, mul_mv_mma_q4_0>;
+template [[host_name("kernel_mul_mv_mma_q4_0_f32_nt2_rt1")]] kernel mul_mv_mma_t kernel_mul_mv_mma_blk<2, 1, mul_mv_mma_q4_0>;
+template [[host_name("kernel_mul_mv_mma_q4_0_f32_nt4_rt1")]] kernel mul_mv_mma_t kernel_mul_mv_mma_blk<4, 1, mul_mv_mma_q4_0>;
+template [[host_name("kernel_mul_mv_mma_q4_0_f32_nt1_rt2")]] kernel mul_mv_mma_t kernel_mul_mv_mma_blk<1, 2, mul_mv_mma_q4_0>;
+template [[host_name("kernel_mul_mv_mma_q4_0_f32_nt2_rt2")]] kernel mul_mv_mma_t kernel_mul_mv_mma_blk<2, 2, mul_mv_mma_q4_0>;
+template [[host_name("kernel_mul_mv_mma_q4_0_f32_nt4_rt2")]] kernel mul_mv_mma_t kernel_mul_mv_mma_blk<4, 2, mul_mv_mma_q4_0>;
+
+template [[host_name("kernel_mul_mv_mma_q8_0_f32_nt1_rt1")]] kernel mul_mv_mma_t kernel_mul_mv_mma_blk<1, 1, mul_mv_mma_q8_0>;
+template [[host_name("kernel_mul_mv_mma_q8_0_f32_nt2_rt1")]] kernel mul_mv_mma_t kernel_mul_mv_mma_blk<2, 1, mul_mv_mma_q8_0>;
+template [[host_name("kernel_mul_mv_mma_q8_0_f32_nt4_rt1")]] kernel mul_mv_mma_t kernel_mul_mv_mma_blk<4, 1, mul_mv_mma_q8_0>;
+
+// q5_K scale and min of sub-block j from the 12 packed bytes, held as 3 words
+inline float2 mul_mv_mma_q5_K_scale_min(thread const uint * w, short j) {
+ if (j < 4) {
+ return float2((w[0] >> 8*j) & 63, (w[1] >> 8*j) & 63);
+ }
+ const short k = 8*(j - 4);
+ return float2(((w[2] >> k) & 0xF) | (((w[0] >> (k + 6)) & 3) << 4), ((w[2] >> (k + 4)) & 0xF) | (((w[1] >> (k + 6)) & 3) << 4));
+}
+
+// the A fragments of one qs/qh word: lo holds the 5-bit quants q of the low-nibble sub-block of a pair, hi holds 16*q of the high-nibble sub-block.
+// step e0 takes bytes 0 and 2 of the word, step e1 takes bytes 1 and 3; hs holds the qh bits of the pair.
+inline void mul_mv_mma_q5_K_frags(uint q, uint hs, thread half2 * lo, thread half2 * hi) {
+ const ushort2 qw = as_type<ushort2>(q);
+ const ushort2 hw = as_type<ushort2>(hs);
+
+ lo[0] = mul_mv_mma_1024_plus((qw & 0x000F) | ((hw << 4) & 0x0010)) - mma_f16_1024;
+ lo[1] = mul_mv_mma_1024_plus(((qw >> 8) & 0x000F) | ((hw >> 4) & 0x0010)) - mma_f16_1024;
+ hi[0] = mul_mv_mma_1024_plus((qw & 0x00F0) | ((hw << 7) & 0x0100)) - mma_f16_1024;
+ hi[1] = mul_mv_mma_1024_plus(((qw >> 8) & 0x00F0) | ((hw >> 1) & 0x0100)) - mma_f16_1024;
+}
+
+struct mul_mv_mma_q5_K_a {
+ uint2 q;
+ uint2 h;
+ uint sc[3];
+ uint dm;
+};
+
+template<short NT>
+inline void load_q5_K_mma_a(device const block_q5_K * const x[NT], int ip, short fn, thread mul_mv_mma_q5_K_a * a) {
+ constexpr short pairs = QK_K/64;
+
+ FOR_UNROLL (short t = 0; t < NT; ++t) {
+ device const block_q5_K * xb = x[t] + ip/pairs;
+ device const uint * sp = (device const uint *) xb->scales;
+
+ a[t].q = *((device const uint2 *) (xb->qs + 32*(ip%pairs)) + fn/2);
+ a[t].h = *((device const uint2 *) xb->qh + fn/2);
+ a[t].sc[0] = sp[0];
+ a[t].sc[1] = sp[1];
+ a[t].sc[2] = sp[2];
+ a[t].dm = *((device const uint *) xb);
+ }
+}
+
+// xs[rt][e][h]: src1 values of sub-block h of the pair, for the steps (b, b + 1) and (b + 4, b + 5)
+template<short RT>
+inline void load_q5_K_mma_b(device const float2 * const y[RT][2], int ip, thread float2 (*xs)[2][2][2]) {
+ FOR_UNROLL (short rt = 0; rt < RT; ++rt) {
+ FOR_UNROLL (short e = 0; e < 2; ++e) {
+ FOR_UNROLL (short h = 0; h < 2; ++h) {
+ xs[rt][e][h][0] = y[rt][e][ip*32 + h*16 + 0];
+ xs[rt][e][h][1] = y[rt][e][ip*32 + h*16 + 2];
+ }
+ }
+ }
+}
+
+// few-row mat-mat for q5_K over pairs of 32-weight sub-blocks, laid out like mul_mv_mma_q4_0: A lane (m, j) holds qs and qh bytes 4*j .. 4*j + 7 of a pair (j even).
+// the 1/16 of the in-place high nibbles and the sub-block scales go into the accumulation; the mins are removed with the src1 sums.
+template<short NT, short RT>
+kernel void kernel_mul_mv_mma_q5_K_f32(
+ constant ggml_metal_kargs_mul_mv_ext & args,
+ device const char * src0,
+ device const char * src1,
+ device char * dst,
+ device const char * src2,
+ threadgroup char * shmem [[threadgroup(0)]],
+ uint3 tgpig[[threadgroup_position_in_grid]],
+ ushort tiisg[[thread_index_in_simdgroup]],
+ ushort sgitg[[simdgroup_index_in_threadgroup]]) {
+ const short NSG = FC_mul_mv_mma_nsg;
+
+ constexpr float hi_scale = 1.0f/16;
+ constexpr short pairs = QK_K/64;
+
+ const mul_mv_mma_tile tile = mul_mv_mma_tile_init<NT, RT>(args, tgpig, tiisg);
+
+ device const block_q5_K * x[NT];
+ FOR_UNROLL (short t = 0; t < NT; ++t) {
+ x[t] = (device const block_q5_K *) mul_mv_mma_src0_row(tile, args, src0, t);
+ }
+
+ // B lane k = fm holds src1 values b, b + 1, b + 4, b + 5 of each sub-block, b = 8*(k/2) + 2*(k%2)
+ device const float2 * y[RT][2];
+ FOR_UNROLL (short rt = 0; rt < RT; ++rt) {
+ FOR_UNROLL (short e = 0; e < 2; ++e) {
+ y[rt][e] = (device const float2 *) (mul_mv_mma_src1_row(tile, args, src1, rt, e) + 8*(tile.fm/2) + 2*(tile.fm%2));
+ }
+ }
+
+ float acc[RT][NT][2] = {};
+
+ const int np = FC_mul_mv_mma_ne00/64;
+
+ mul_mv_mma_q5_K_a an[NT];
+ float2 xn[RT][2][2][2];
+
+ const int ip0 = min((int) sgitg, np - 1);
+ load_q5_K_mma_a<NT>(x, ip0, tile.fn, an);
+ load_q5_K_mma_b<RT>(y, ip0, xn);
+
+ for (int ip = sgitg; ip < np; ip += NSG) {
+ const short p = ip%pairs;
+
+ mul_mv_mma_q5_K_a ac[NT];
+ float2 xs[RT][2][2][2];
+ FOR_UNROLL (short t = 0; t < NT; ++t) {
+ ac[t] = an[t];
+ }
+ FOR_UNROLL (short rt = 0; rt < RT; ++rt) {
+ FOR_UNROLL (short e = 0; e < 2; ++e) {
+ FOR_UNROLL (short h = 0; h < 2; ++h) {
+ xs[rt][e][h][0] = xn[rt][e][h][0];
+ xs[rt][e][h][1] = xn[rt][e][h][1];
+ }
+ }
+ }
+
+ const int ipn = min(ip + NSG, np - 1);
+ load_q5_K_mma_a<NT>(x, ipn, tile.fn, an);
+ load_q5_K_mma_b<RT>(y, ipn, xn);
+
+ float c[RT][2][2];
+ FOR_UNROLL (short rt = 0; rt < RT; ++rt) {
+ FOR_UNROLL (short e = 0; e < 2; ++e) {
+ FOR_UNROLL (short h = 0; h < 2; ++h) {
+ float u = xs[rt][e][h][0].x + xs[rt][e][h][0].y + xs[rt][e][h][1].x + xs[rt][e][h][1].y;
+ u += simd_shuffle_xor(u, 2);
+ u += simd_shuffle_xor(u, 4);
+ u += simd_shuffle_xor(u, 16);
+ c[rt][e][h] = u;
+ }
+ }
+ }
+
+ FOR_UNROLL (short t = 0; t < NT; ++t) {
+ const float2 dm = float2(as_type<half2>(ac[t].dm));
+ const float2 sm0 = mul_mv_mma_q5_K_scale_min(ac[t].sc, 2*p + 0);
+ const float2 sm1 = mul_mv_mma_q5_K_scale_min(ac[t].sc, 2*p + 1);
+
+ half2 a[2][2][2];
+ mul_mv_mma_q5_K_frags(ac[t].q.x, ac[t].h.x >> 2*p, a[0][0], a[1][0]);
+ mul_mv_mma_q5_K_frags(ac[t].q.y, ac[t].h.y >> 2*p, a[0][1], a[1][1]);
+
+ FOR_UNROLL (short hh = 0; hh < 2; ++hh) {
+ simdgroup_float8x8 mp[RT];
+ FOR_UNROLL (short rt = 0; rt < RT; ++rt) {
+ mp[rt] = make_filled_simdgroup_matrix<float, 8>(0.0f);
+ }
+
+ FOR_UNROLL (short s = 0; s < 4; ++s) {
+ simdgroup_half8x8 ma;
+ ma.thread_elements()[0] = a[hh][s/2][s%2].x;
+ ma.thread_elements()[1] = a[hh][s/2][s%2].y;
+
+ FOR_UNROLL (short rt = 0; rt < RT; ++rt) {
+ simdgroup_float8x8 mb;
+ mb.thread_elements()[0] = s % 2 == 0 ? xs[rt][0][hh][s/2].x : xs[rt][0][hh][s/2].y;
+ mb.thread_elements()[1] = s % 2 == 0 ? xs[rt][1][hh][s/2].x : xs[rt][1][hh][s/2].y;
+
+ simdgroup_multiply_accumulate(mp[rt], ma, mb, mp[rt]);
+ }
+ }
+
+ const float2 smh = hh == 0 ? sm0 : sm1;
+ const float dsc = dm.x*smh.x*(hh == 0 ? 1.0f : hi_scale);
+ const float dmn = dm.y*smh.y;
+
+ FOR_UNROLL (short rt = 0; rt < RT; ++rt) {
+ acc[rt][t][0] = fma(dsc, mp[rt].thread_elements()[0], fma(-dmn, c[rt][0][hh], acc[rt][t][0]));
+ acc[rt][t][1] = fma(dsc, mp[rt].thread_elements()[1], fma(-dmn, c[rt][1][hh], acc[rt][t][1]));
+ }
+ }
+ }
+ }
+
+ mul_mv_mma_store<NT, RT>(acc, args, src2, dst, shmem, tile, tiisg, sgitg);
+}
+
+template [[host_name("kernel_mul_mv_mma_q5_K_f32_nt1_rt1")]] kernel mul_mv_mma_t kernel_mul_mv_mma_q5_K_f32<1, 1>;
+template [[host_name("kernel_mul_mv_mma_q5_K_f32_nt2_rt1")]] kernel mul_mv_mma_t kernel_mul_mv_mma_q5_K_f32<2, 1>;
+template [[host_name("kernel_mul_mv_mma_q5_K_f32_nt4_rt1")]] kernel mul_mv_mma_t kernel_mul_mv_mma_q5_K_f32<4, 1>;
+template [[host_name("kernel_mul_mv_mma_q5_K_f32_nt1_rt2")]] kernel mul_mv_mma_t kernel_mul_mv_mma_q5_K_f32<1, 2>;
+template [[host_name("kernel_mul_mv_mma_q5_K_f32_nt2_rt2")]] kernel mul_mv_mma_t kernel_mul_mv_mma_q5_K_f32<2, 2>;
+template [[host_name("kernel_mul_mv_mma_q5_K_f32_nt4_rt2")]] kernel mul_mv_mma_t kernel_mul_mv_mma_q5_K_f32<4, 2>;
+
+// few-row mat-mat for any type with a 16-weight dequantizer: a lane dequantizes 16 consecutive weights of a 64-weight chunk once for all src1 rows.
+// MMA step s at MMA-k index j reads chunk weight 16*(j/2) + 8*(j%2) + s.
+template<short NT, short RT, typename block_q, short nl, void (*dequantize_func)(device const block_q *, short, thread float4x4 &)>
+kernel void kernel_mul_mv_mma_gen(
+ constant ggml_metal_kargs_mul_mv_ext & args,
+ device const char * src0,
+ device const char * src1,
+ device char * dst,
+ device const char * src2,
+ threadgroup char * shmem [[threadgroup(0)]],
+ uint3 tgpig[[threadgroup_position_in_grid]],
+ ushort tiisg[[thread_index_in_simdgroup]],
+ ushort sgitg[[simdgroup_index_in_threadgroup]]) {
+ const short NSG = FC_mul_mv_mma_nsg;
+
+ const mul_mv_mma_tile tile = mul_mv_mma_tile_init<NT, RT>(args, tgpig, tiisg);
+
+ device const block_q * x[NT];
+ FOR_UNROLL (short t = 0; t < NT; ++t) {
+ x[t] = (device const block_q *) mul_mv_mma_src0_row(tile, args, src0, t);
+ }
+
+ // B lane j = fm reads chunk values 16*(fm/2) + 8*(fm%2) .. +7
+ device const float4 * y[RT][2];
+ FOR_UNROLL (short rt = 0; rt < RT; ++rt) {
+ FOR_UNROLL (short e = 0; e < 2; ++e) {
+ y[rt][e] = (device const float4 *) mul_mv_mma_src1_row(tile, args, src1, rt, e) + 4*(tile.fm/2) + 2*(tile.fm%2);
+ }
+ }
+
+ simdgroup_float8x8 mc[RT][NT];
+ FOR_UNROLL (short rt = 0; rt < RT; ++rt) {
+ FOR_UNROLL (short t = 0; t < NT; ++t) {
+ mc[rt][t] = make_filled_simdgroup_matrix<float, 8>(0.0f);
+ }
+ }
+
+ const int nch = args.ne00/64;
+
+ for (int g = sgitg; g < nch; g += NSG) {
+ simdgroup_float8x8 mb[RT][8];
+ FOR_UNROLL (short rt = 0; rt < RT; ++rt) {
+ const float4 a0 = y[rt][0][16*g + 0];
+ const float4 a1 = y[rt][0][16*g + 1];
+ const float4 b0 = y[rt][1][16*g + 0];
+ const float4 b1 = y[rt][1][16*g + 1];
+ FOR_UNROLL (short s = 0; s < 4; ++s) {
+ mb[rt][s ].thread_elements()[0] = a0[s];
+ mb[rt][s ].thread_elements()[1] = b0[s];
+ mb[rt][s + 4].thread_elements()[0] = a1[s];
+ mb[rt][s + 4].thread_elements()[1] = b1[s];
+ }
+ }
+
+ const int ci = 4*g + tile.fn/2;
+
+ FOR_UNROLL (short t = 0; t < NT; ++t) {
+ float4x4 w;
+ dequantize_func(x[t] + ci/nl, ci%nl, w);
+
+ FOR_UNROLL (short s = 0; s < 8; ++s) {
+ simdgroup_float8x8 ma;
+ ma.thread_elements()[0] = w[s/4 ][s%4];
+ ma.thread_elements()[1] = w[s/4 + 2][s%4];
+
+ FOR_UNROLL (short rt = 0; rt < RT; ++rt) {
+ simdgroup_multiply_accumulate(mc[rt][t], ma, mb[rt][s], mc[rt][t]);
+ }
+ }
+ }
+ }
+
+ float acc[RT][NT][2];
+ FOR_UNROLL (short rt = 0; rt < RT; ++rt) {
+ FOR_UNROLL (short t = 0; t < NT; ++t) {
+ acc[rt][t][0] = mc[rt][t].thread_elements()[0];
+ acc[rt][t][1] = mc[rt][t].thread_elements()[1];
+ }
+ }
+
+ mul_mv_mma_store<NT, RT>(acc, args, src2, dst, shmem, tile, tiisg, sgitg);
+}
+
+#define MUL_MV_MMA_GEN(tname, bq, nl, deq) \
+template [[host_name("kernel_mul_mv_mma_" tname "_f32_nt1_rt1")]] kernel mul_mv_mma_t kernel_mul_mv_mma_gen<1, 1, bq, nl, deq>; \
+template [[host_name("kernel_mul_mv_mma_" tname "_f32_nt2_rt1")]] kernel mul_mv_mma_t kernel_mul_mv_mma_gen<2, 1, bq, nl, deq>; \
+template [[host_name("kernel_mul_mv_mma_" tname "_f32_nt4_rt1")]] kernel mul_mv_mma_t kernel_mul_mv_mma_gen<4, 1, bq, nl, deq>; \
+template [[host_name("kernel_mul_mv_mma_" tname "_f32_nt1_rt2")]] kernel mul_mv_mma_t kernel_mul_mv_mma_gen<1, 2, bq, nl, deq>; \
+template [[host_name("kernel_mul_mv_mma_" tname "_f32_nt2_rt2")]] kernel mul_mv_mma_t kernel_mul_mv_mma_gen<2, 2, bq, nl, deq>; \
+template [[host_name("kernel_mul_mv_mma_" tname "_f32_nt4_rt2")]] kernel mul_mv_mma_t kernel_mul_mv_mma_gen<4, 2, bq, nl, deq>;
+
+// q8_0 with 9..16 src1 rows: the per-block scaling above is slower than dequantizing to f32
+template [[host_name("kernel_mul_mv_mma_q8_0_f32_nt1_rt2")]] kernel mul_mv_mma_t kernel_mul_mv_mma_gen<1, 2, block_q8_0, 2, dequantize_q8_0>;
+template [[host_name("kernel_mul_mv_mma_q8_0_f32_nt2_rt2")]] kernel mul_mv_mma_t kernel_mul_mv_mma_gen<2, 2, block_q8_0, 2, dequantize_q8_0>;
+template [[host_name("kernel_mul_mv_mma_q8_0_f32_nt4_rt2")]] kernel mul_mv_mma_t kernel_mul_mv_mma_gen<4, 2, block_q8_0, 2, dequantize_q8_0>;
+
+MUL_MV_MMA_GEN("f32", float4x4, 1, dequantize_f32)
+MUL_MV_MMA_GEN("f16", half4x4, 1, dequantize_f16)
+MUL_MV_MMA_GEN("q4_1", block_q4_1, 2, dequantize_q4_1)
+MUL_MV_MMA_GEN("q5_0", block_q5_0, 2, dequantize_q5_0)
+MUL_MV_MMA_GEN("q5_1", block_q5_1, 2, dequantize_q5_1)
+MUL_MV_MMA_GEN("q4_K", block_q4_K, QK_NL, dequantize_q4_K)
+MUL_MV_MMA_GEN("q6_K", block_q6_K, QK_NL, dequantize_q6_K)
+
+#undef MUL_MV_MMA_GEN
diff --git a/ggml/src/ggml-metal/kernels/quantize.metal b/ggml/src/ggml-metal/kernels/quantize.metal
index 42ca6d74a..73679b9ab 100644
--- a/ggml/src/ggml-metal/kernels/quantize.metal
+++ b/ggml/src/ggml-metal/kernels/quantize.metal
@@ -171,7 +171,8 @@ kernel void kernel_concat(
const int i3 = tgpig.z;
const int i2 = tgpig.y;
- const int i1 = ntg.y == 1 ? tgpig.x : tgpig.x*ntg.y + tpitg.y;
+ const int i1 = ntg.y == 1 ? tgpig.x/args.nc0 : tgpig.x*ntg.y + tpitg.y;
+ const int ic = ntg.y == 1 ? tgpig.x%args.nc0 : 0;
if (i1 >= args.ne1) {
return;
@@ -180,7 +181,11 @@ kernel void kernel_concat(
int o[4] = {0, 0, 0, 0};
o[args.dim] = args.dim == 0 ? args.ne00 : (args.dim == 1 ? args.ne01 : (args.dim == 2 ? args.ne02 : args.ne03));
- for (int i0 = tpitg.x; i0 < args.ne0; i0 += ntg.x) {
+ // chunk ic of nc0 along the row
+ const int n0 = (args.ne0 + args.nc0 - 1)/args.nc0;
+ const int i0e = min(args.ne0, (ic + 1)*n0);
+
+ for (int i0 = ic*n0 + tpitg.x; i0 < i0e; i0 += ntg.x) {
device const T * x;
if (i0 < args.ne00 && i1 < args.ne01 && i2 < args.ne02 && i3 < args.ne03) {
@@ -220,7 +225,8 @@ kernel void kernel_concat_q(
// note: for quantized types, the args are in units of blocks (nb0 == type_size)
const int i3 = tgpig.z;
const int i2 = tgpig.y;
- const int i1 = ntg.y == 1 ? tgpig.x : tgpig.x*ntg.y + tpitg.y;
+ const int i1 = ntg.y == 1 ? tgpig.x/args.nc0 : tgpig.x*ntg.y + tpitg.y;
+ const int ic = ntg.y == 1 ? tgpig.x%args.nc0 : 0;
if (i1 >= args.ne1) {
return;
@@ -229,7 +235,11 @@ kernel void kernel_concat_q(
int o[4] = {0, 0, 0, 0};
o[args.dim] = args.dim == 0 ? args.ne00 : (args.dim == 1 ? args.ne01 : (args.dim == 2 ? args.ne02 : args.ne03));
- for (int i0 = tpitg.x; i0 < args.ne0; i0 += ntg.x) {
+ // chunk ic of nc0 along the row
+ const int n0 = (args.ne0 + args.nc0 - 1)/args.nc0;
+ const int i0e = min(args.ne0, (ic + 1)*n0);
+
+ for (int i0 = ic*n0 + tpitg.x; i0 < i0e; i0 += ntg.x) {
device const block_q * x;
if (i0 < args.ne00 && i1 < args.ne01 && i2 < args.ne02 && i3 < args.ne03) {
diff --git a/tests/CMakeLists.txt b/tests/CMakeLists.txt
index 96d910a75..b5ef02c02 100644
--- a/tests/CMakeLists.txt
+++ b/tests/CMakeLists.txt
@@ -276,6 +276,12 @@ if (NOT LLAMA_SANITIZE_ADDRESS AND NOT GGML_SCHED_NO_REALLOC)
# TODO: repair known memory leaks
llama_build_and_test(test-opt.cpp)
endif()
+if (GGML_METAL AND NOT GGML_BACKEND_DL)
+ # calls the Metal backend's graph optimizer and fusion table directly
+ llama_build_and_test(test-metal-graph-optimize.cpp)
+ target_include_directories(test-metal-graph-optimize PRIVATE ${PROJECT_SOURCE_DIR}/ggml/src ${PROJECT_SOURCE_DIR}/ggml/src/ggml-metal)
+ target_link_libraries(test-metal-graph-optimize PRIVATE ggml-metal)
+endif()
# TODO: make this test (and others) not link `libllama` as it is not needed [TAG_TESTS_LLAMA_LINK]
llama_build(test-backend-ops.cpp)
diff --git a/tests/test-backend-ops.cpp b/tests/test-backend-ops.cpp
index 070d378f4..c82af18ff 100644
--- a/tests/test-backend-ops.cpp
+++ b/tests/test-backend-ops.cpp
@@ -7193,6 +7193,92 @@ struct test_moe_reduce : public test_case {
}
};
+// mul_mat with src1 in [0, 1]: a zero-mean src1 hides errors in the zero point or the min of a quantized src0
+struct test_mul_mat_pos : public test_mul_mat {
+ using test_mul_mat::test_mul_mat;
+
+ std::string vars() override {
+ return test_mul_mat::vars() + ",src1=[0,1]";
+ }
+
+ void initialize_tensors(ggml_context * ctx) override {
+ for (ggml_tensor * t = ggml_get_first_tensor(ctx); t != nullptr; t = ggml_get_next_tensor(ctx, t)) {
+ if (t->type == GGML_TYPE_F32) {
+ init_tensor_uniform(t, 0.0f, 1.0f);
+ } else {
+ init_tensor_uniform(t);
+ }
+ }
+ }
+};
+
+enum mul_mat_add_mode {
+ MUL_MAT_ADD_MM_RES, // mm + res
+ MUL_MAT_ADD_RES_MM, // res + mm
+ MUL_MAT_ADD_ROW, // mm + a one-row res, which a same-shape fusion must leave alone
+ MUL_MAT_ADD_RES_INPLACE, // res += mm
+ MUL_MAT_ADD_B_INPLACE, // b += mm: the sum overwrites the mat-mul input (m == k)
+};
+
+static std::string var_to_str(mul_mat_add_mode mode) {
+ switch (mode) {
+ case MUL_MAT_ADD_MM_RES: return "mm+res";
+ case MUL_MAT_ADD_RES_MM: return "res+mm";
+ case MUL_MAT_ADD_ROW: return "mm+row";
+ case MUL_MAT_ADD_RES_INPLACE: return "res+=mm";
+ case MUL_MAT_ADD_B_INPLACE: return "b+=mm";
+ }
+ return "unknown";
+}
+
+// mul_mat followed by an add of a residual, which backends may fuse into the mat-mul
+struct test_mul_mat_add : public test_case {
+ const ggml_type type_a;
+ const int64_t m;
+ const int64_t n;
+ const int64_t k;
+ const mul_mat_add_mode mode;
+
+ test_mul_mat_add(ggml_type type_a, int64_t m, int64_t n, int64_t k, mul_mat_add_mode mode = MUL_MAT_ADD_MM_RES)
+ : type_a(type_a), m(m), n(n), k(k), mode(mode) {
+ GGML_ASSERT(mode != MUL_MAT_ADD_B_INPLACE || m == k);
+ }
+
+ std::string vars() override {
+ return VARS_TO_STR5(type_a, m, n, k, mode);
+ }
+
+ std::string op_desc(ggml_tensor * t) override {
+ GGML_UNUSED(t);
+ return "MUL_MAT_ADD";
+ }
+
+ bool run_whole_graph() override { return true; }
+
+ double max_nmse_err() override {
+ return 5e-4;
+ }
+
+ ggml_tensor * build_graph(ggml_context * ctx) override {
+ ggml_tensor * a = ggml_new_tensor_2d(ctx, type_a, k, m);
+ ggml_tensor * b = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, k, n);
+ ggml_tensor * res = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, m, mode == MUL_MAT_ADD_ROW ? 1 : n);
+
+ ggml_tensor * mm = ggml_mul_mat(ctx, a, b);
+ ggml_tensor * out = nullptr;
+ switch (mode) {
+ case MUL_MAT_ADD_MM_RES:
+ case MUL_MAT_ADD_ROW: out = ggml_add(ctx, mm, res); break;
+ case MUL_MAT_ADD_RES_MM: out = ggml_add(ctx, res, mm); break;
+ case MUL_MAT_ADD_RES_INPLACE: out = ggml_add_inplace(ctx, res, mm); break;
+ case MUL_MAT_ADD_B_INPLACE: out = ggml_add_inplace(ctx, b, mm); break;
+ }
+ ggml_set_name(out, "out");
+
+ return out;
+ }
+};
+
struct test_mul_mat_vec_fusion : public test_case {
const ggml_type type;
const ggml_glu_op glu_op;
@@ -10289,6 +10375,41 @@ static std::vector<std::unique_ptr<test_case>> make_test_cases_eval() {
test_cases.emplace_back(new test_mul_mat(GGML_TYPE_F32, GGML_TYPE_F32, 32, 509, 2112, {1, 1}, {1, 1}));
test_cases.emplace_back(new test_mul_mat(GGML_TYPE_Q8_0, GGML_TYPE_F32, 32, 509, 2112, {1, 1}, {1, 1}));
+ // few src1 rows (speculative verify): odd m, a single K block, long K, every tile width, broadcast batches
+ for (int64_t n : {2, 3, 5, 8, 9, 13, 16}) {
+ for (auto [m, k] : std::vector<std::pair<int64_t, int64_t>>{{40, 32}, {100, 96}, {1000, 5120}, {3000, 1024}, {6144, 5120}, {17408, 512}}) {
+ test_cases.emplace_back(new test_mul_mat(GGML_TYPE_Q4_0, GGML_TYPE_F32, m, n, k, {1, 1}, {1, 1}));
+ }
+ test_cases.emplace_back(new test_mul_mat(GGML_TYPE_Q4_0, GGML_TYPE_F32, 64, n, 256, {3, 2}, {2, 1}));
+ for (ggml_type type_a : {GGML_TYPE_F32, GGML_TYPE_F16, GGML_TYPE_Q4_1, GGML_TYPE_Q5_0, GGML_TYPE_Q5_1, GGML_TYPE_Q8_0,
+ GGML_TYPE_Q4_K, GGML_TYPE_Q5_K, GGML_TYPE_Q6_K}) {
+ test_cases.emplace_back(new test_mul_mat(type_a, GGML_TYPE_F32, 48, n, 2560, {1, 1}, {1, 1}));
+ test_cases.emplace_back(new test_mul_mat(type_a, GGML_TYPE_F32, 1000, n, 1024, {1, 1}, {1, 1}));
+ test_cases.emplace_back(new test_mul_mat(type_a, GGML_TYPE_F32, 3000, n, 512, {1, 1}, {1, 1}));
+ test_cases.emplace_back(new test_mul_mat(type_a, GGML_TYPE_F32, 4100, n, 256, {1, 1}, {1, 1}));
+ test_cases.emplace_back(new test_mul_mat(type_a, GGML_TYPE_F32, 64, n, 256, {3, 2}, {2, 1}));
+ // a K that is not a multiple of the 64-weight chunk takes the other kernels
+ if (96 % ggml_blck_size(type_a) == 0) {
+ test_cases.emplace_back(new test_mul_mat(type_a, GGML_TYPE_F32, 100, n, 96, {1, 1}, {1, 1}));
+ }
+ }
+ for (ggml_type type_a : {GGML_TYPE_Q4_0, GGML_TYPE_Q8_0, GGML_TYPE_Q5_K, GGML_TYPE_F16}) {
+ for (mul_mat_add_mode mode : {MUL_MAT_ADD_MM_RES, MUL_MAT_ADD_RES_MM, MUL_MAT_ADD_ROW, MUL_MAT_ADD_RES_INPLACE}) {
+ test_cases.emplace_back(new test_mul_mat_add(type_a, 1000, n, 1024, mode));
+ }
+ test_cases.emplace_back(new test_mul_mat_add(type_a, 2048, n, 2048, MUL_MAT_ADD_B_INPLACE));
+ }
+ // the widest tiles with two src1 tiles, and src1 rows padded in memory
+ for (ggml_type type_a : {GGML_TYPE_F16, GGML_TYPE_Q4_0, GGML_TYPE_Q8_0, GGML_TYPE_Q4_K, GGML_TYPE_Q5_K, GGML_TYPE_Q6_K}) {
+ test_cases.emplace_back(new test_mul_mat(type_a, GGML_TYPE_F32, 8192, n, 512, {1, 1}, {1, 1}));
+ test_cases.emplace_back(new test_mul_mat(type_a, GGML_TYPE_F32, 1000, n, 1024, {1, 1}, {1, 1}, {0, 1, 2, 3}, 1280));
+ }
+ // a src1 with a nonzero mean, for the zero points and mins of quantized src0 types
+ for (ggml_type type_a : {GGML_TYPE_Q4_0, GGML_TYPE_Q4_1, GGML_TYPE_Q5_0, GGML_TYPE_Q8_0, GGML_TYPE_Q4_K, GGML_TYPE_Q5_K}) {
+ test_cases.emplace_back(new test_mul_mat_pos(type_a, GGML_TYPE_F32, 1000, n, 1024, {1, 1}, {1, 1}));
+ }
+ }
+
#if 0
{
// Test paths in OpenCL
@@ -10874,6 +10995,13 @@ static std::vector<std::unique_ptr<test_case>> make_test_cases_eval() {
}
}
+ // few long rows, which backends may split across workgroups
+ for (int dim : { 0, 1 }) {
+ test_cases.emplace_back(new test_concat(GGML_TYPE_F32, {98304, 8, 1, 1}, dim == 0 ? 288 : 2, dim, 0));
+ test_cases.emplace_back(new test_concat(GGML_TYPE_F32, {20000, 3, 2, 1}, dim == 0 ? 333 : 1, dim, 1));
+ test_cases.emplace_back(new test_concat(GGML_TYPE_Q8_0, {8192, 2, 1, 1}, dim == 0 ? 4096 : 1, dim, 0));
+ }
+
for (ggml_sort_order order : {GGML_SORT_ORDER_ASC, GGML_SORT_ORDER_DESC}) {
for (uint32_t i = 4; i <= 1024*1024; i *= 2) {
test_cases.emplace_back(new test_argsort(GGML_TYPE_F32, {i-1, 1, 1, 1}));
diff --git a/tests/test-metal-graph-optimize.cpp b/tests/test-metal-graph-optimize.cpp
new file mode 100644
index 000000000..c46987b12
--- /dev/null
+++ b/tests/test-metal-graph-optimize.cpp
@@ -0,0 +1,285 @@
+// checks the MUL_MAT+ADD packs of the Metal graph reorder (devices, pack readers, src1 rows), and that the encoder fuses
+// no MUL_MAT+ADD the reorder leaves unpacked
+#include "ggml.h"
+#include "ggml-alloc.h"
+#include "ggml-backend.h"
+#include "ggml-impl.h"
+#include "ggml-metal-common.h"
+#include "ggml-metal-device.h"
+#include "ggml-metal-fusion.h"
+
+#include <cstdio>
+#include <vector>
+
+static constexpr int n_tensors = 64;
+static constexpr float scale = 2.0f;
+static constexpr float norm_eps = 1e-6f;
+// a mat-mul the few-row MMA kernels take: K a multiple of their 64-weight step, 2..16 src1 rows
+static constexpr int64_t n_k = 64;
+static constexpr int64_t n_m = 16;
+static constexpr int64_t n_rows = 8;
+// src1 row counts below, inside and above the 2..16 rows of the few-row MMA kernels
+static const std::vector<int64_t> batch_rows = { 1, 4, 9, 64, 512 };
+
+struct device_case {
+ const char * name;
+
+ bool has_native_simdgroup_mm;
+ bool packed;
+};
+
+// a device with probed simdgroup matrices that are native (MTLGPUFamilyApple7+) only if has_native_simdgroup_mm
+static ggml_metal_device_props device_props(bool has_native_simdgroup_mm) {
+ ggml_metal_device_props props = {};
+ props.has_simdgroup_mm = has_native_simdgroup_mm;
+ return props;
+}
+
+// the index of t in the nodes of graph, -1 if absent
+static int node_index(ggml_cgraph * graph, const ggml_tensor * t) {
+ for (int i = 0; i < ggml_graph_n_nodes(graph); ++i) {
+ if (ggml_graph_node(graph, i) == t) {
+ return i;
+ }
+ }
+ return -1;
+}
+
+// the reordered positions of the tracked nodes
+static std::vector<int> node_positions(ggml_cgraph * graph, const std::vector<ggml_tensor *> & tracked) {
+ std::vector<int> res;
+ for (const ggml_tensor * t : tracked) {
+ res.push_back(node_index(graph, t));
+ }
+ return res;
+}
+
+// allocates the tensors of ctx in a CPU buffer marked as weights, as the model loader does
+static ggml_backend_buffer_t alloc_weights(ggml_context * ctx) {
+ ggml_backend_buffer_type_t buft = ggml_backend_dev_buffer_type(ggml_backend_dev_by_type(GGML_BACKEND_DEVICE_TYPE_CPU));
+ ggml_backend_buffer_t buffer = ggml_backend_alloc_ctx_tensors_from_buft(ctx, buft);
+ ggml_backend_buffer_set_usage(buffer, GGML_BACKEND_BUFFER_USAGE_WEIGHTS);
+ return buffer;
+}
+
+// what the add sums with the mat-mul
+enum add_operand_kind {
+ ADD_RESIDUAL, // a same-shape activation
+ ADD_BIAS, // a one-row bias weight
+ ADD_BIAS_VIEW, // a view of a bias weight
+};
+
+static const char * add_operand_name(add_operand_kind kind) {
+ switch (kind) {
+ case ADD_RESIDUAL: return "residual";
+ case ADD_BIAS: return "bias";
+ case ADD_BIAS_VIEW: return "bias view";
+ }
+ return "?";
+}
+
+static ggml_tensor * new_add_operand(ggml_context * ctx, ggml_context * ctx_w, add_operand_kind kind, int64_t rows) {
+ switch (kind) {
+ case ADD_RESIDUAL: return ggml_new_tensor_2d(ctx, GGML_TYPE_F32, n_m, rows);
+ case ADD_BIAS: return ggml_new_tensor_2d(ctx_w, GGML_TYPE_F32, n_m, 1);
+ case ADD_BIAS_VIEW: return ggml_reshape_2d(ctx, ggml_new_tensor_1d(ctx_w, GGML_TYPE_F32, n_m), n_m, 1);
+ }
+ return nullptr;
+}
+
+// the reordered positions of a mat-mul with rows src1 rows, the add of an operand of kind to it and an independent
+// mat-mul, which the reorder runs between the first mat-mul and the add unless the pair is packed
+static std::vector<int> reorder_mul_mat_add(int64_t rows, add_operand_kind kind) {
+ ggml_init_params params = { n_tensors*ggml_tensor_overhead() + ggml_graph_overhead(), nullptr, true };
+ ggml_context * ctx = ggml_init(params);
+ ggml_context * ctx_w = ggml_init(params);
+ ggml_tensor * x = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, n_k, rows);
+ ggml_tensor * mm = ggml_mul_mat(ctx, ggml_new_tensor_2d(ctx_w, GGML_TYPE_F32, n_k, n_m), x);
+ ggml_tensor * add = ggml_add(ctx, mm, new_add_operand(ctx, ctx_w, kind, rows));
+ ggml_tensor * other = ggml_mul_mat(ctx, ggml_new_tensor_2d(ctx_w, GGML_TYPE_F32, n_k, n_m), x);
+ ggml_backend_buffer_t weights = alloc_weights(ctx_w);
+
+ ggml_cgraph * graph = ggml_new_graph(ctx);
+ ggml_build_forward_expand(graph, add);
+ ggml_build_forward_expand(graph, other);
+ ggml_graph_optimize(graph);
+ const std::vector<int> res = node_positions(graph, { mm, add, other });
+
+ ggml_backend_buffer_free(weights);
+ ggml_free(ctx_w);
+ ggml_free(ctx);
+ return res;
+}
+
+// true if the independent mat-mul of reorder_mul_mat_add does not run between the mat-mul and the add
+static bool add_packed(const std::vector<int> & pos) {
+ return !(pos[0] < pos[2] && pos[2] < pos[1]);
+}
+
+static bool check_pack(const device_case & c) {
+ const std::vector<int> pos = reorder_mul_mat_add(n_rows, ADD_RESIDUAL);
+ const bool packed = add_packed(pos);
+
+ const bool ok = packed == c.packed;
+ std::printf("%s: mat-mul and add packed %d (expected %d): %s\n", c.name, packed, c.packed, ok ? "OK" : "FAIL");
+ return ok;
+}
+
+static int run_pack_cases(const std::vector<device_case> & cases) {
+ int failures = 0;
+ for (const device_case & c : cases) {
+ failures += check_pack(c) ? 0 : 1;
+ }
+ return failures;
+}
+
+// true if every row count in rows gives the reorder of the first one, on a device that fuses MUL_MAT+ADD
+static bool same_reorder_for_rows(const std::vector<int64_t> & rows, add_operand_kind kind) {
+ const std::vector<int> first = reorder_mul_mat_add(rows.front(), kind);
+ for (auto r = rows.begin() + 1; r != rows.end(); ++r) {
+ if (reorder_mul_mat_add(*r, kind) != first) {
+ return false;
+ }
+ }
+ return true;
+}
+
+// the pack must not depend on the batch size, or graphs with the same nodes get another allocation per ubatch size,
+// and a bias add, which the encoder never fuses, must stay free to run next to independent nodes
+static bool check_row_independent_pack(add_operand_kind kind) {
+ const bool same = same_reorder_for_rows(batch_rows, kind);
+ const bool packed = add_packed(reorder_mul_mat_add(batch_rows.front(), kind));
+ const bool expected = kind == ADD_RESIDUAL;
+ const bool ok = same && packed == expected;
+ std::printf("MUL_MAT+ADD of a %s, reorder independent of src1 rows %d, packed %d (expected %d): %s\n",
+ add_operand_name(kind), same, packed, expected, ok ? "OK" : "FAIL");
+ return ok;
+}
+
+static int run_row_independent_pack_cases() {
+ int failures = 0;
+ for (add_operand_kind kind : { ADD_RESIDUAL, ADD_BIAS, ADD_BIAS_VIEW }) {
+ failures += check_row_independent_pack(kind) ? 0 : 1;
+ }
+ return failures;
+}
+
+// true if the encoder's check fuses a few-row MUL_MAT with a same-shape residual, which is a weight if weight_res;
+// the check compares Metal buffer ranges, so the tensors live in Metal buffers
+static bool encoder_fuses_mul_mat_add(bool weight_res) {
+ ggml_backend_buffer_type_t buft = ggml_backend_dev_buffer_type(ggml_backend_dev_by_type(GGML_BACKEND_DEVICE_TYPE_GPU));
+ ggml_init_params params = { n_tensors*ggml_tensor_overhead(), nullptr, true };
+ ggml_context * ctx = ggml_init(params);
+ ggml_context * ctx_w = ggml_init(params);
+ ggml_tensor * x = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, n_k, n_rows);
+ ggml_tensor * mm = ggml_mul_mat(ctx, ggml_new_tensor_2d(ctx_w, GGML_TYPE_F32, n_k, n_m), x);
+ ggml_tensor * add = ggml_add(ctx, mm, ggml_new_tensor_2d(weight_res ? ctx_w : ctx, GGML_TYPE_F32, n_m, n_rows));
+ ggml_backend_buffer_t buffer = ggml_backend_alloc_ctx_tensors_from_buft(ctx, buft);
+ ggml_backend_buffer_t weights = ggml_backend_alloc_ctx_tensors_from_buft(ctx_w, buft);
+ ggml_backend_buffer_set_usage(weights, GGML_BACKEND_BUFFER_USAGE_WEIGHTS);
+
+ ggml_init_params params_gf = { ggml_graph_overhead(), nullptr, true };
+ ggml_context * ctx_gf = ggml_init(params_gf);
+ ggml_cgraph * gf = ggml_new_graph(ctx_gf);
+ ggml_build_forward_expand(gf, add);
+
+ const int idxs[] = { 0, 1 };
+ int n_fused = 1;
+ const ggml_metal_fusion * fusion = ggml_metal_fusion_next(gf, idxs, 2, 0, GGML_METAL_FUSION_FULL, &n_fused);
+ const bool fused = fusion != nullptr && ggml_metal_fusion_get_id(fusion) == GGML_METAL_FUSION_MUL_MAT_ADD;
+ ggml_free(ctx_gf);
+
+ ggml_backend_buffer_free(weights);
+ ggml_backend_buffer_free(buffer);
+ ggml_free(ctx_w);
+ ggml_free(ctx);
+ return fused;
+}
+
+// the encoder may fuse only what the reorder packs, so it must leave a same-shape residual alone when it is a weight
+static bool check_encoder_skips_weight_residual() {
+ const bool fuses_residual = encoder_fuses_mul_mat_add(false);
+ const bool fuses_weight = encoder_fuses_mul_mat_add(true);
+ const bool ok = fuses_residual && !fuses_weight;
+ std::printf("encoder fuses MUL_MAT+ADD of a residual %d (expected 1), of a weight %d (expected 0): %s\n",
+ fuses_residual, fuses_weight, ok ? "OK" : "FAIL");
+ return ok;
+}
+
+static ggml_context * graph_ctx() {
+ ggml_init_params params = { n_tensors*ggml_tensor_overhead() + ggml_graph_overhead(), nullptr, true };
+ return ggml_init(params);
+}
+
+static void expand_all(ggml_cgraph * graph, const std::vector<ggml_tensor *> & outputs) {
+ for (ggml_tensor * t : outputs) {
+ ggml_build_forward_expand(graph, t);
+ }
+}
+
+// the graph of outputs in build order, reordered for a device that fuses MUL_MAT+ADD
+static ggml_cgraph * optimized_graph(ggml_context * ctx, const std::vector<ggml_tensor *> & outputs) {
+ ggml_cgraph * graph = ggml_new_graph(ctx);
+ expand_all(graph, outputs);
+ ggml_graph_optimize(graph);
+ return graph;
+}
+
+// the position of t among the nodes the encoder runs (views are skipped), -1 if absent
+static int encoded_index(ggml_cgraph * graph, const ggml_tensor * t) {
+ int res = 0;
+ for (int i = 0; i < ggml_graph_n_nodes(graph); ++i) {
+ const ggml_tensor * node = ggml_graph_node(graph, i);
+ if (node == t) {
+ return res;
+ }
+ res += ggml_op_is_empty(node->op) ? 0 : 1;
+ }
+ return -1;
+}
+
+// true if the encoder runs the nodes back to back in this order, so they fuse
+static bool encoded_in_a_row(ggml_cgraph * graph, const std::vector<ggml_tensor *> & nodes) {
+ const int first = encoded_index(graph, nodes[0]);
+ for (size_t j = 1; j < nodes.size(); ++j) {
+ if (encoded_index(graph, nodes[j]) != first + (int) j) {
+ return false;
+ }
+ }
+ return true;
+}
+
+static bool report_reader(const char * name, bool packed, bool reader_after) {
+ const bool ok = packed && reader_after;
+ std::printf("%s: packed %d, reader after the write %d (expected 1, 1): %s\n", name, packed, reader_after, ok ? "OK" : "FAIL");
+ return ok;
+}
+
+// x feeds a MUL_MAT+ADD pack chained with RMS_NORM+MUL, and c reads the residual sum h, an output inside the pack
+static bool check_chained_pack_reader() {
+ ggml_context * ctx = graph_ctx();
+ ggml_tensor * x = ggml_scale(ctx, ggml_new_tensor_2d(ctx, GGML_TYPE_F32, n_k, n_rows), scale);
+ ggml_tensor * mm = ggml_mul_mat(ctx, ggml_new_tensor_2d(ctx, GGML_TYPE_F32, n_k, n_m), x);
+ ggml_tensor * h = ggml_add(ctx, mm, ggml_new_tensor_2d(ctx, GGML_TYPE_F32, n_m, n_rows));
+ ggml_tensor * n = ggml_rms_norm(ctx, h, norm_eps);
+ ggml_tensor * m = ggml_mul(ctx, n, ggml_new_tensor_1d(ctx, GGML_TYPE_F32, n_m));
+ ggml_tensor * c = ggml_scale(ctx, h, scale);
+
+ ggml_cgraph * graph = optimized_graph(ctx, { m, c });
+ const bool packed = encoded_in_a_row(graph, { mm, h, n, m });
+ const bool after = encoded_index(graph, c) > encoded_index(graph, h);
+ ggml_free(ctx);
+
+ return report_reader("reader of a chained MUL_MAT+ADD sum", packed, after);
+}
+
+int main() {
+ const std::vector<device_case> devices = {
+ { "native simdgroup matrices", true, true },
+ { "probed or no simdgroup matrices", false, true },
+ };
+ const int failures = run_pack_cases(devices) + (check_chained_pack_reader() ? 0 : 1) + run_row_independent_pack_cases() +
+ (check_encoder_skips_weight_residual() ? 0 : 1);
+
+ return failures == 0 ? 0 : 1;
+}