From be12d39a17de113be24f3433c802d9c9ac044eb3 Mon Sep 17 00:00:00 2001 From: Yuri Khrustalev Date: Tue, 25 Aug 2026 07:35:39 -0400 Subject: [PATCH 001/104] metal : null-check buffer alloc to fix OOM crash (llama/25371) * metal : null-check ggml_metal_buffer_init result to avoid OOM crash ggml_backend_metal_buffer_type_alloc_buffer used the result of ggml_metal_buffer_init without checking for NULL. ggml_metal_buffer_init returns NULL when the underlying Metal allocation fails (e.g. an out-of-memory condition), and the following ggml_metal_buffer_is_shared(res) call dereferences it, turning a recoverable allocation failure into a hard crash (EXC_BAD_ACCESS). This is easy to hit on memory-constrained devices such as iOS when a model/context exceeds the available Metal budget. Log the failure using the existing GGML_LOG_ERROR convention and return NULL so the allocator surfaces a diagnosable error up the stack instead of crashing. * cont : fix log --------- Co-authored-by: Georgi Gerganov --- ggml/src/ggml-metal/ggml-metal.cpp | 5 +++++ 1 file changed, 5 insertions(+) diff --git a/ggml/src/ggml-metal/ggml-metal.cpp b/ggml/src/ggml-metal/ggml-metal.cpp index e5b2ee8a55b..9756d47050c 100644 --- a/ggml/src/ggml-metal/ggml-metal.cpp +++ b/ggml/src/ggml-metal/ggml-metal.cpp @@ -204,6 +204,11 @@ static ggml_backend_buffer_t ggml_backend_metal_buffer_type_alloc_buffer(ggml_ba ggml_metal_device_t ctx_dev = (ggml_metal_device_t)buft->device->context; ggml_metal_buffer_t res = ggml_metal_buffer_init(ctx_dev, size, shared); + if (res == NULL) { + GGML_LOG_ERROR("%s: failed to allocate Metal buffer of %zu bytes (out of memory)\n", __func__, size); + return NULL; + } + ggml_backend_buffer_i buf_i = ggml_metal_buffer_is_shared(res) ? ggml_backend_metal_buffer_shared_i : ggml_backend_metal_buffer_private_i; From e820c280a64b37c14e9c3aa4879dda51451ecb72 Mon Sep 17 00:00:00 2001 From: Ryan C Date: Tue, 25 Aug 2026 17:12:15 +0000 Subject: [PATCH 002/104] rpc: support apple RDMA as an RPC transport (llama/26421) * rpc: support apple RDMA as an RPC transport * remove set_tensor micro optimization, rpc socket pinning per CR * remove transparent reconnect * trigger apple builds on RPC changes --------- Co-authored-by: Ryan Churaman --- ggml/src/ggml-rpc/CMakeLists.txt | 28 +- ggml/src/ggml-rpc/ggml-rpc.cpp | 7 +- ggml/src/ggml-rpc/transport-apple.cpp | 470 ++++++++++++++++++++++++++ ggml/src/ggml-rpc/transport-apple.h | 27 ++ ggml/src/ggml-rpc/transport.cpp | 106 ++++-- ggml/src/ggml-rpc/transport.h | 4 + 6 files changed, 607 insertions(+), 35 deletions(-) create mode 100644 ggml/src/ggml-rpc/transport-apple.cpp create mode 100644 ggml/src/ggml-rpc/transport-apple.h diff --git a/ggml/src/ggml-rpc/CMakeLists.txt b/ggml/src/ggml-rpc/CMakeLists.txt index 40e11fead63..b2f086380d5 100644 --- a/ggml/src/ggml-rpc/CMakeLists.txt +++ b/ggml/src/ggml-rpc/CMakeLists.txt @@ -9,10 +9,18 @@ if (WIN32) target_link_libraries(ggml-rpc PRIVATE ws2_32) endif() -# RDMA auto-detection (Linux only, requires libibverbs) -if (NOT WIN32 AND NOT APPLE) - find_library(IBVERBS_LIB ibverbs) - if (IBVERBS_LIB) +# RDMA auto-detection: Linux RoCE/IB via libibverbs, Apple RDMA-over-Thunderbolt via librdma +if (APPLE) + set(RDMA_LIB_NAME rdma) + set(RDMA_DESC "Apple RDMA-over-Thunderbolt, UC") +elseif (NOT WIN32) + set(RDMA_LIB_NAME ibverbs) + set(RDMA_DESC "auto-detected") +endif() + +if (RDMA_LIB_NAME) + find_library(RDMA_LIB ${RDMA_LIB_NAME}) + if (RDMA_LIB) option(GGML_RPC_RDMA "ggml: enable RDMA transport for RPC" ON) else() option(GGML_RPC_RDMA "ggml: enable RDMA transport for RPC" OFF) @@ -22,12 +30,16 @@ else() endif() if (GGML_RPC_RDMA) - if (NOT IBVERBS_LIB) - find_library(IBVERBS_LIB ibverbs REQUIRED) + if (NOT RDMA_LIB) + find_library(RDMA_LIB ${RDMA_LIB_NAME} REQUIRED) endif() target_compile_definitions(ggml-rpc PRIVATE GGML_RPC_RDMA) - target_link_libraries(ggml-rpc PRIVATE ${IBVERBS_LIB}) - message(STATUS " RDMA transport enabled (auto-detected)") + target_link_libraries(ggml-rpc PRIVATE ${RDMA_LIB}) + if (APPLE) + target_compile_definitions(ggml-rpc PRIVATE GGML_RPC_RDMA_APPLE) + target_sources(ggml-rpc PRIVATE transport-apple.cpp) + endif() + message(STATUS " RDMA transport enabled (${RDMA_DESC})") else() message(STATUS " RDMA transport disabled") endif() diff --git a/ggml/src/ggml-rpc/ggml-rpc.cpp b/ggml/src/ggml-rpc/ggml-rpc.cpp index 9d480226600..69a8a08ae17 100644 --- a/ggml/src/ggml-rpc/ggml-rpc.cpp +++ b/ggml/src/ggml-rpc/ggml-rpc.cpp @@ -253,7 +253,10 @@ static bool send_msg(socket_ptr sock, const void * msg, size_t msg_size) { if (!sock->send_data(&msg_size, sizeof(msg_size))) { return false; } - return sock->send_data(msg, msg_size); + if (!sock->send_data(msg, msg_size)) { + return false; + } + return sock->flush(); } static bool recv_msg(socket_ptr sock, void * msg, size_t msg_size) { @@ -308,7 +311,7 @@ static bool send_rpc_cmd(socket_ptr sock, enum rpc_cmd cmd, const void * input, if (!sock->send_data(input, input_size)) { return false; } - return true; + return sock->flush(); } // RPC request : | rpc_cmd (1 byte) | request_size (8 bytes) | request_data (request_size bytes) | diff --git a/ggml/src/ggml-rpc/transport-apple.cpp b/ggml/src/ggml-rpc/transport-apple.cpp new file mode 100644 index 00000000000..c8be77a6dce --- /dev/null +++ b/ggml/src/ggml-rpc/transport-apple.cpp @@ -0,0 +1,470 @@ +#include "transport-apple.h" +#include "transport.h" +#include "ggml-impl.h" + +#include + +#include +#include +#include +#include +#include +#include +#include + +// Apple RDMA-over-Thunderbolt (see Apple TN3205). +// +// Apple's RDMA is quite different from what's supported in Linux - deserving of its own transport implementation. +// see https://developer.apple.com/documentation/technotes/tn3205-low-latency-communication-with-rdma-over-thunderbolt for details +// at a high level the main differences are: +// UC(unreliable connection) on Apple vs RC(reliable connection) QP transport types on Linux (though in practice UC on Apple is still lossless) +// fixed 128KiB stride on Apple vs variable chunk size on Linux +// relying on Apple's hardware credit based flow control vs RNR NAKs + retries on Linux +// +// on Apple a SEND and its corresponding RECV must cover the same number of 4 KiB Thunderbolt frames, +// so every SEND posts a whole 128KiB stride over the wire, even when partially filled. +// (In testing 128KiB was the best performing among 32, 64, 128, 256) + +static constexpr uint32_t RDMA_SEG_MAGIC = 0x52534547u; // "RSEG" +static constexpr int RDMA_NBUF = 16; // ring depth (frames per direction) +static constexpr size_t RDMA_FRAME = 4096; // Thunderbolt frame (fixed on Apple) +static constexpr size_t RDMA_STRIDE = 128 * 1024; // 32 Thunderbolt frames; NBUF x this = 2 MiB pinned per direction +static constexpr uint32_t RDMA_PSN = 0; // any value works if both sides match: UC has no retransmit +static constexpr size_t RDMA_GID_SIZE = 16; + +static_assert(RDMA_STRIDE % RDMA_FRAME == 0, "RDMA_STRIDE must be a whole number of frames"); +// TN3205 counts queue depth in Thunderbolt frames, not work requests. +static constexpr uint32_t RDMA_QP_WR = (uint32_t)RDMA_NBUF * (RDMA_STRIDE / RDMA_FRAME); +static constexpr uint64_t RDMA_RECV_WR = 1ull << 20; // wr_id bit tagging recv completions +static constexpr uint64_t RDMA_WR_IDX_MASK = 0xffff; // buffer index in the low bits of wr_id +static constexpr uint8_t RDMA_SYNC_READY = 0x2A; // readiness-handshake byte (peer activated) + +struct rdma_seg_hdr { + uint32_t magic; // RDMA_SEG_MAGIC; a mismatch means the stream desynced + uint32_t len; // payload bytes in this frame; the rest of the stride is padding +}; +static constexpr size_t RDMA_PAYLOAD = RDMA_STRIDE - sizeof(rdma_seg_hdr); + +struct apple_rdma_caps { + uint32_t qpn; + uint16_t lid; + uint16_t reserved; + uint8_t gid[RDMA_GID_SIZE]; +}; + +static_assert(sizeof(apple_rdma_caps) == RPC_CONN_CAPS_SIZE, "apple_rdma_caps must match conn_caps size"); + +struct apple_rdma::impl { + int fd = -1; // bootstrap TCP socket, kept as the liveness anchor + + struct ibv_context * ctx = nullptr; + struct ibv_pd * pd = nullptr; + struct ibv_cq * cq = nullptr; // one CQ for both directions; RDMA_RECV_WR tags recv completions + struct ibv_qp * qp = nullptr; + + uint8_t * send_mem = nullptr; + struct ibv_mr * send_mr = nullptr; + uint8_t * recv_mem = nullptr; + struct ibv_mr * recv_mr = nullptr; + + int send_busy[RDMA_NBUF] = {}; // 1 while this buffer has a send in flight + // completed recv frames, oldest first: ring index, bytes already handed to + // the reader, and total payload length + struct { int buf; uint32_t off; uint32_t len; } inq[RDMA_NBUF] = {}; + int inq_head = 0; + int inq_count = 0; + int pend_buf = -1; + uint32_t pend_len = 0; + bool broken = false; + + uint32_t qpn = 0; + uint8_t port = 0; + int gid_idx = 0; + enum ibv_mtu path_mtu = IBV_MTU_1024; + + int progress(); + bool acquire_pending(); + bool post_pending(); + + bool post_recv(int i) { + struct ibv_sge sge = {}; + sge.addr = (uintptr_t)(recv_mem + (size_t)i * RDMA_STRIDE); + sge.length = (uint32_t)RDMA_STRIDE; + sge.lkey = recv_mr->lkey; + struct ibv_recv_wr wr = {}, * bad = nullptr; + wr.wr_id = RDMA_RECV_WR | (uint64_t)i; + wr.sg_list = &sge; + wr.num_sge = 1; + return ibv_post_recv(qp, &wr, &bad) == 0; + } + + bool post_send(int i, size_t len) { + struct ibv_sge sge = {}; + sge.addr = (uintptr_t)(send_mem + (size_t)i * RDMA_STRIDE); + sge.length = (uint32_t)len; + sge.lkey = send_mr->lkey; + struct ibv_send_wr wr = {}, * bad = nullptr; + wr.wr_id = (uint64_t)i; + wr.sg_list = &sge; + wr.num_sge = 1; + wr.opcode = IBV_WR_SEND; + wr.send_flags = IBV_SEND_SIGNALED; + return ibv_post_send(qp, &wr, &bad) == 0; + } + + ~impl() { + broken = true; + // the QP must be destroyed before the memory it can still write to is + // deregistered and freed: ERR only starts flushing the posted WQEs + if (qp) { + struct ibv_qp_attr a = {}; + a.qp_state = IBV_QPS_ERR; + ibv_modify_qp(qp, &a, IBV_QP_STATE); + struct ibv_wc wc[RDMA_NBUF * 2]; + while (ibv_poll_cq(cq, RDMA_NBUF * 2, wc) > 0) {} + ibv_destroy_qp(qp); + } + if (send_mr) ibv_dereg_mr(send_mr); + if (recv_mr) ibv_dereg_mr(recv_mr); + free(send_mem); + free(recv_mem); + if (cq) ibv_destroy_cq(cq); + if (pd) ibv_dealloc_pd(pd); + if (ctx) ibv_close_device(ctx); + } +}; + +apple_rdma::apple_rdma(std::unique_ptr p) : pimpl(std::move(p)) {} + +apple_rdma::~apple_rdma() = default; + +bool apple_rdma::broken() const { + return pimpl->broken; +} + +// The readiness handshake below still runs over the bootstrap socket, one byte +// each way, before the transport is declared live. +static bool tcp_send_byte(int fd, uint8_t b) { + ssize_t n; + do { n = ::send(fd, &b, sizeof(b), 0); } while (n < 0 && errno == EINTR); + return n == sizeof(b); +} + +static bool tcp_recv_byte(int fd, uint8_t * b) { + ssize_t n; + do { n = ::recv(fd, b, sizeof(*b), 0); } while (n < 0 && errno == EINTR); + return n == (ssize_t)sizeof(*b); +} + +// Index of the GID on this port equal to the target, or -1. Thunderbolt GIDs are +// RoCEv2 IPv4-mapped (::ffff:a.b.c.d), so this matches the local TCP address. +static int rdma_match_gid(struct ibv_context * ctx, uint8_t port, int gid_tbl_len, + const uint8_t * target, union ibv_gid * out) { + for (int i = 0; i < gid_tbl_len; i++) { + union ibv_gid g; + if (ibv_query_gid(ctx, port, i, &g) != 0) continue; + if (memcmp(g.raw, target, RDMA_GID_SIZE) != 0) continue; + if (out) *out = g; + return i; + } + return -1; +} + +// First ACTIVE port on the device. Only a cabled, up Thunderbolt link reports +// ACTIVE, and it is not always port 1, so the port cannot be hardcoded the way +// the Linux path does. Returns 0 if none. +static uint8_t rdma_first_active_port(struct ibv_context * ctx, struct ibv_port_attr * out) { + struct ibv_device_attr da; + if (ibv_query_device(ctx, &da) != 0) return 0; + for (uint8_t p = 1; p <= da.phys_port_cnt; p++) { + struct ibv_port_attr pa; + if (ibv_query_port(ctx, p, &pa) != 0) continue; + if (pa.state == IBV_PORT_ACTIVE) { if (out) *out = pa; return p; } + } + return 0; +} + +// Called before the endpoints are exchanged: pick the local device facing this +// peer, create a UC QP and register the frame rings. RDMA is point-to-point, so +// the device is the one whose GID equals the bootstrap connection's local +// address, i.e. the one cabled to the peer. +std::unique_ptr apple_rdma::probe(int fd, const uint8_t * target_gid, uint8_t * caps) { + int ndev = 0; + ibv_device ** devs = ibv_get_device_list(&ndev); + if (!devs) return nullptr; + + ibv_context * ctx = nullptr; + uint8_t port = 0; + struct ibv_port_attr pa = {}; + union ibv_gid gid = {}; + int gid_idx = -1; + std::string matched; + for (int d = 0; d < ndev; d++) { + ibv_context * c = ibv_open_device(devs[d]); + if (!c) continue; + struct ibv_port_attr p = {}; + uint8_t pt = rdma_first_active_port(c, &p); + int gi = pt ? rdma_match_gid(c, pt, p.gid_tbl_len, target_gid, &gid) : -1; + if (gi < 0) { ibv_close_device(c); continue; } + ctx = c; port = pt; pa = p; gid_idx = gi; + const char * name = ibv_get_device_name(devs[d]); + matched = name ? name : ""; + break; + } + ibv_free_device_list(devs); + if (!ctx) return nullptr; + + std::unique_ptr c(new impl()); + c->fd = fd; + c->ctx = ctx; + c->port = port; + c->gid_idx = gid_idx; + c->path_mtu = pa.active_mtu; + + c->pd = ibv_alloc_pd(ctx); + if (!c->pd) return nullptr; + + c->cq = ibv_create_cq(ctx, 2 * RDMA_QP_WR + 1, nullptr, nullptr, 0); + if (!c->cq) return nullptr; + + ibv_qp_init_attr qia = {}; + qia.send_cq = c->cq; + qia.recv_cq = c->cq; + qia.qp_type = IBV_QPT_UC; + qia.cap.max_send_wr = RDMA_QP_WR; + qia.cap.max_recv_wr = RDMA_QP_WR; + qia.cap.max_send_sge = 1; + qia.cap.max_recv_sge = 1; + c->qp = ibv_create_qp(c->pd, &qia); + if (!c->qp) return nullptr; + + { + ibv_qp_attr a = {}; + a.qp_state = IBV_QPS_INIT; + a.pkey_index = 0; + a.port_num = port; + a.qp_access_flags = IBV_ACCESS_LOCAL_WRITE | IBV_ACCESS_REMOTE_READ | IBV_ACCESS_REMOTE_WRITE; + if (ibv_modify_qp(c->qp, &a, + IBV_QP_STATE | IBV_QP_PKEY_INDEX | IBV_QP_PORT | IBV_QP_ACCESS_FLAGS) != 0) { + return nullptr; + } + } + + long page = sysconf(_SC_PAGESIZE); + if (page <= 0) page = 4096; + const size_t ring_bytes = (size_t)RDMA_NBUF * RDMA_STRIDE; + if (posix_memalign((void **)&c->send_mem, (size_t)page, ring_bytes) != 0) c->send_mem = nullptr; + if (posix_memalign((void **)&c->recv_mem, (size_t)page, ring_bytes) != 0) c->recv_mem = nullptr; + if (!c->send_mem || !c->recv_mem) return nullptr; + + // Apple's provider rejects LOCAL_WRITE-only MRs even for two-sided SEND/RECV. + const int mr_flags = IBV_ACCESS_LOCAL_WRITE | IBV_ACCESS_REMOTE_READ | IBV_ACCESS_REMOTE_WRITE; + c->send_mr = ibv_reg_mr(c->pd, c->send_mem, ring_bytes, mr_flags); + c->recv_mr = ibv_reg_mr(c->pd, c->recv_mem, ring_bytes, mr_flags); + if (!c->send_mr || !c->recv_mr) return nullptr; + + // Recvs are posted in activate() after the RTS transition, not here: Apple's + // provider rejects ibv_post_recv on a QP that has not reached RTS. + + c->qpn = c->qp->qp_num; + + apple_rdma_caps rc = {}; + rc.qpn = c->qpn; + rc.lid = pa.lid; + memcpy(rc.gid, gid.raw, RDMA_GID_SIZE); + memcpy(caps, &rc, sizeof(rc)); + + GGML_LOG_INFO("RDMA(Apple/UC) probed: dev=%s port=%u gid=%d qpn=%u lid=%u mtu=%d ring=%d x %zu KiB\n", + matched.c_str(), port, gid_idx, c->qpn, (unsigned)pa.lid, 128 << c->path_mtu, + RDMA_NBUF, RDMA_STRIDE / 1024); + return std::unique_ptr(new apple_rdma(std::move(c))); +} + +// Called once the peer's endpoint has arrived: INIT -> RTR -> RTS (UC: GID/GRH +// addressing, no timeout/retry/rnr/rd_atomic), then the readiness handshake. +bool apple_rdma::activate(const uint8_t * caps) { + impl * c = pimpl.get(); + + apple_rdma_caps rc = {}; + memcpy(&rc, caps, sizeof(rc)); + + bool ok = true; + { + ibv_qp_attr a = {}; + a.qp_state = IBV_QPS_RTR; + a.path_mtu = c->path_mtu; + a.rq_psn = RDMA_PSN; + a.dest_qp_num = rc.qpn; + a.ah_attr.is_global = 1; + a.ah_attr.port_num = c->port; + a.ah_attr.sl = 0; + a.ah_attr.src_path_bits = 0; + a.ah_attr.dlid = rc.lid; + a.ah_attr.grh.hop_limit = 1; + a.ah_attr.grh.sgid_index = (uint8_t)c->gid_idx; + memcpy(&a.ah_attr.grh.dgid, rc.gid, RDMA_GID_SIZE); + if (ibv_modify_qp(c->qp, &a, + IBV_QP_STATE | IBV_QP_AV | IBV_QP_PATH_MTU | IBV_QP_DEST_QPN | IBV_QP_RQ_PSN) != 0) { + GGML_LOG_ERROR("RDMA(Apple/UC) RTR failed: %s\n", strerror(errno)); + ok = false; + } + } + if (ok) { + ibv_qp_attr a = {}; + a.qp_state = IBV_QPS_RTS; + a.sq_psn = RDMA_PSN; + if (ibv_modify_qp(c->qp, &a, IBV_QP_STATE | IBV_QP_SQ_PSN) != 0) { + GGML_LOG_ERROR("RDMA(Apple/UC) RTS failed: %s\n", strerror(errno)); + ok = false; + } + } + + // Recvs are posted only now: the controller starts processing them at RTR. + for (int i = 0; ok && i < RDMA_NBUF; i++) { + if (!c->post_recv(i)) { + GGML_LOG_ERROR("RDMA(Apple/UC) post_recv %d/%d failed\n", i, RDMA_NBUF); + ok = false; + } + } + + // A queue pair processes receives only after RTR and the transitions above can + // fail on one side alone, so neither peer sends a frame until both report their + // recvs posted. + uint8_t peer_ready = 0; + if (!tcp_send_byte(c->fd, ok ? RDMA_SYNC_READY : 0) || !tcp_recv_byte(c->fd, &peer_ready)) { + return false; + } + if (!ok || peer_ready != RDMA_SYNC_READY) { + return false; + } + + GGML_LOG_INFO("RDMA(Apple/UC) activated: qpn=%u->%u mtu=%d rx_depth=%d\n", + c->qpn, rc.qpn, 128 << c->path_mtu, RDMA_NBUF); + return true; +} + +// Drain the CQ: release completed send buffers, queue completed recv frames for +// the reader. Returns the number of completions reaped, or -1 on error. +int apple_rdma::impl::progress() { + struct ibv_wc wc[RDMA_NBUF * 2]; + int n = ibv_poll_cq(cq, RDMA_NBUF * 2, wc); + if (n < 0) { GGML_LOG_ERROR("RDMA(Apple/UC) poll_cq failed\n"); broken = true; return -1; } + for (int j = 0; j < n; j++) { + uint64_t id = wc[j].wr_id; + bool is_recv = (id & RDMA_RECV_WR) != 0; + if (wc[j].status != IBV_WC_SUCCESS) { + GGML_LOG_ERROR("RDMA(Apple/UC) %s wc error: status=%d\n", is_recv ? "recv" : "send", wc[j].status); + broken = true; + return -1; + } + if (is_recv) { + int b = (int)(id & RDMA_WR_IDX_MASK); + const rdma_seg_hdr * h = (const rdma_seg_hdr *)(recv_mem + (size_t)b * RDMA_STRIDE); + if (h->magic != RDMA_SEG_MAGIC) { GGML_LOG_ERROR("RDMA(Apple/UC) bad frame magic\n"); broken = true; return -1; } + if (h->len > RDMA_PAYLOAD) { GGML_LOG_ERROR("RDMA(Apple/UC) frame len %u exceeds payload\n", h->len); broken = true; return -1; } + int slot = (inq_head + inq_count) % RDMA_NBUF; + inq[slot].buf = b; + inq[slot].off = 0; + inq[slot].len = h->len; + inq_count++; + } else { + send_busy[(int)(id & RDMA_WR_IDX_MASK)] = 0; + } + } + return n; +} + +// Reserve a free send buffer to coalesce into, waiting on progress if none free. +bool apple_rdma::impl::acquire_pending() { + if (pend_buf >= 0) return true; + for (;;) { + if (broken) return false; + for (int k = 0; k < RDMA_NBUF; k++) if (!send_busy[k]) { pend_buf = k; pend_len = 0; return true; } + if (progress() < 0) return false; + } +} + +// Post the pending frame. The whole STRIDE goes out even when only partly filled: +// TN3205 requires a SEND and its matching RECV to cover the same number of +// Thunderbolt frames, so a short send would fail the peer's receive. +bool apple_rdma::impl::post_pending() { + if (pend_buf < 0) return true; + int i = pend_buf; + rdma_seg_hdr * h = (rdma_seg_hdr *)(send_mem + (size_t)i * RDMA_STRIDE); + h->magic = RDMA_SEG_MAGIC; + h->len = pend_len; + if (!post_send(i, RDMA_STRIDE)) { broken = true; return false; } + send_busy[i] = 1; + pend_buf = -1; + pend_len = 0; + return true; +} + +// Coalescing write: append into the pending frame, posting a full frame when it +// fills. The trailing partial is posted by flush() at each message boundary. +bool apple_rdma::send(const void * data, size_t size) { + impl * c = pimpl.get(); + const uint8_t * p = (const uint8_t *)data; + while (size > 0) { + if (c->broken) return false; + if (!c->acquire_pending()) return false; + uint8_t * sb = c->send_mem + (size_t)c->pend_buf * RDMA_STRIDE; + size_t space = RDMA_PAYLOAD - c->pend_len; + size_t chunk = size < space ? size : space; + memcpy(sb + sizeof(rdma_seg_hdr) + c->pend_len, p, chunk); + c->pend_len += (uint32_t)chunk; + p += chunk; + size -= chunk; + if (c->pend_len == RDMA_PAYLOAD) { if (!c->post_pending()) return false; } + } + return true; +} + +bool apple_rdma::recv(void * data, size_t size) { + impl * c = pimpl.get(); + uint8_t * p = (uint8_t *)data; + if (!c->post_pending()) return false; // turnaround: flush the coalesced request + unsigned idle = 0; + while (size > 0) { + if (c->inq_count == 0) { + if (c->broken) return false; + int n = c->progress(); + if (n < 0) return false; + if (n == 0) { + // UC gives no disconnect notification, so the bootstrap TCP fd is + // the liveness anchor: nothing crosses it once RDMA is up, so any + // readability means the peer's FIN (macOS has no POLLRDHUP). + // Same idle interval as the Linux path. + if ((++idle & 0xFFFFF) == 0) { + struct pollfd pfd = { c->fd, POLLIN, 0 }; + if (poll(&pfd, 1, 0) > 0 && + (pfd.revents & (POLLIN | POLLHUP | POLLERR | POLLNVAL))) { + return false; + } + } + } else { + idle = 0; + } + continue; + } + idle = 0; + int slot = c->inq_head; + int b = c->inq[slot].buf; + uint32_t avail = c->inq[slot].len - c->inq[slot].off; + uint32_t take = (size < (size_t)avail) ? (uint32_t)size : avail; + memcpy(p, c->recv_mem + (size_t)b * RDMA_STRIDE + sizeof(rdma_seg_hdr) + c->inq[slot].off, take); + p += take; + size -= take; + c->inq[slot].off += take; + if (c->inq[slot].off == c->inq[slot].len) { + if (!c->post_recv(b)) { c->broken = true; return false; } + c->inq_head = (c->inq_head + 1) % RDMA_NBUF; + c->inq_count--; + } + } + return true; +} + +bool apple_rdma::flush() { + return pimpl->post_pending(); +} diff --git a/ggml/src/ggml-rpc/transport-apple.h b/ggml/src/ggml-rpc/transport-apple.h new file mode 100644 index 00000000000..7968d38a17a --- /dev/null +++ b/ggml/src/ggml-rpc/transport-apple.h @@ -0,0 +1,27 @@ +#pragma once + +#include +#include +#include + +struct apple_rdma { + // target_gid is 16 bytes in, caps is RPC_CONN_CAPS_SIZE bytes out. + static std::unique_ptr probe(int fd, const uint8_t * target_gid, uint8_t * caps); + ~apple_rdma(); + + // Peer endpoint from its caps, which must be non-zero: this blocks on a + // readiness handshake over fd that the peer only joins if it also has RDMA. + bool activate(const uint8_t * caps); + + bool send(const void * data, size_t size); + bool recv(void * data, size_t size); + // Post the trailing partial frame; must be called at every message boundary. + bool flush(); + // True once the connection has failed; the caller should drop the socket. + bool broken() const; + +private: + struct impl; + explicit apple_rdma(std::unique_ptr p); + std::unique_ptr pimpl; +}; diff --git a/ggml/src/ggml-rpc/transport.cpp b/ggml/src/ggml-rpc/transport.cpp index a728152421f..5ec15dc80c0 100644 --- a/ggml/src/ggml-rpc/transport.cpp +++ b/ggml/src/ggml-rpc/transport.cpp @@ -18,15 +18,20 @@ # include #endif #include +#include #include #include #ifdef GGML_RPC_RDMA # include +# include # include # ifndef _WIN32 # include # endif +# ifdef GGML_RPC_RDMA_APPLE +# include "transport-apple.h" +# endif #endif // GGML_RPC_RDMA #ifdef _WIN32 @@ -42,10 +47,13 @@ static const char * RPC_DEBUG = std::getenv("GGML_RPC_DEBUG"); do { if (RPC_DEBUG) GGML_LOG_DEBUG(__VA_ARGS__); } while (0) #ifdef GGML_RPC_RDMA -static constexpr size_t RDMA_CHUNK = 256 * 1024; // 256 KiB per send/recv (fits default 8 MiB memlock) -static constexpr int RDMA_RX_DEPTH = 24; // pre-posted recv ring: 24 × 256 KiB = 6 MiB static constexpr size_t RDMA_GID_SIZE = 16; // RoCE GID / IB GID is always 16 bytes using rdma_gid_t = std::array; +#endif // GGML_RPC_RDMA + +#if defined(GGML_RPC_RDMA) && !defined(GGML_RPC_RDMA_APPLE) +static constexpr size_t RDMA_CHUNK = 256 * 1024; // 256 KiB per send/recv (fits default 8 MiB memlock) +static constexpr int RDMA_RX_DEPTH = 24; // pre-posted recv ring: 24 × 256 KiB = 6 MiB struct rdma_conn { struct ibv_context * ctx = nullptr; @@ -111,27 +119,33 @@ struct rdma_caps { static_assert(sizeof(rdma_caps) == RPC_CONN_CAPS_SIZE, "rdma_caps must match conn_caps size"); -#endif // GGML_RPC_RDMA +#endif // GGML_RPC_RDMA && !GGML_RPC_RDMA_APPLE struct socket_t::impl { impl(sockfd_t fd) : use_rdma(false), fd(fd) {} ~impl(); bool send_data(const void * data, size_t size); bool recv_data(void * data, size_t size); + bool flush(); void get_caps(uint8_t * local_caps); void update_caps(const uint8_t * remote_caps); #ifdef GGML_RPC_RDMA - bool tcp_peer_closed(); std::optional rdma_build_target_gid(); + +# ifdef GGML_RPC_RDMA_APPLE + std::unique_ptr rdma; +# else bool rdma_probe(); - bool rdma_activate(uint32_t remote_qpn, uint32_t remote_psn, const uint8_t * remote_gid); - bool rdma_poll(struct ibv_cq * cq, struct ibv_wc * wc); bool rdma_send(const void * data, size_t size); bool rdma_recv(void * data, size_t size); + bool tcp_peer_closed(); + bool rdma_activate(uint32_t remote_qpn, uint32_t remote_psn, const uint8_t * remote_gid); + bool rdma_poll(struct ibv_cq * cq, struct ibv_wc * wc); std::unique_ptr rdma; rdma_local_info rdma_local = {}; +# endif #endif // GGML_RPC_RDMA bool use_rdma; sockfd_t fd; @@ -151,17 +165,6 @@ socket_t::impl::~impl() { #ifdef GGML_RPC_RDMA -bool socket_t::impl::tcp_peer_closed() { - if (fd < 0) return false; -#ifndef _WIN32 - struct pollfd pfd = { fd, POLLIN | POLLRDHUP, 0 }; - int r = poll(&pfd, 1, 0); - return r > 0 && (pfd.revents & (POLLHUP | POLLERR | POLLRDHUP)); -#else - return false; -#endif -} - // Build a RoCE GID-shaped 16-byte target from a TCP socket's local address. // Used to match the socket's local IP against the kernel's GID table so that // a single memcmp handles IPv4, IPv4-mapped IPv6, and native IPv6 uniformly: @@ -191,6 +194,19 @@ std::optional socket_t::impl::rdma_build_target_gid() { return std::nullopt; } +#ifndef GGML_RPC_RDMA_APPLE + +bool socket_t::impl::tcp_peer_closed() { + if (fd < 0) return false; +#ifndef _WIN32 + struct pollfd pfd = { fd, POLLIN | POLLRDHUP, 0 }; + int r = poll(&pfd, 1, 0); + return r > 0 && (pfd.revents & (POLLHUP | POLLERR | POLLRDHUP)); +#else + return false; +#endif +} + bool socket_t::impl::rdma_probe() { const char * dev_env = std::getenv("GGML_RDMA_DEV"); const char * gid_env = std::getenv("GGML_RDMA_GID"); @@ -457,10 +473,16 @@ bool socket_t::impl::rdma_recv(void * data, size_t size) { return true; } +#endif // !GGML_RPC_RDMA_APPLE (Linux RC transport) + #endif // GGML_RPC_RDMA bool socket_t::impl::send_data(const void * data, size_t size) { -#ifdef GGML_RPC_RDMA +#ifdef GGML_RPC_RDMA_APPLE + if (use_rdma) { + return rdma->send(data, size); + } +#elif defined(GGML_RPC_RDMA) if (use_rdma) { return rdma_send(data, size); } @@ -480,7 +502,11 @@ bool socket_t::impl::send_data(const void * data, size_t size) { } bool socket_t::impl::recv_data(void * data, size_t size) { -#ifdef GGML_RPC_RDMA +#ifdef GGML_RPC_RDMA_APPLE + if (use_rdma) { + return rdma->recv(data, size); + } +#elif defined(GGML_RPC_RDMA) if (use_rdma) { return rdma_recv(data, size); } @@ -506,6 +532,15 @@ bool socket_t::impl::recv_data(void * data, size_t size) { void socket_t::impl::get_caps(uint8_t * local_caps) { memset(local_caps, 0, RPC_CONN_CAPS_SIZE); #ifdef GGML_RPC_RDMA + if (std::getenv("GGML_RPC_NO_RDMA")) { + return; + } +# ifdef GGML_RPC_RDMA_APPLE + auto target_gid = rdma_build_target_gid(); + if (target_gid) { + rdma = apple_rdma::probe(fd, target_gid->data(), local_caps); + } +# else rdma_local = {}; if (rdma_probe()) { rdma_caps rc = {}; @@ -516,21 +551,30 @@ void socket_t::impl::get_caps(uint8_t * local_caps) { } else { rdma.reset(); } +# endif #endif // GGML_RPC_RDMA } void socket_t::impl::update_caps(const uint8_t * remote_caps) { #ifdef GGML_RPC_RDMA - if (!rdma) { - return; + // a peer that has no RDMA advertises all-zero caps and takes no further part + // in the negotiation, so drop to TCP without reporting a failure + bool remote_rdma = false; + for (size_t i = 0; i < RPC_CONN_CAPS_SIZE; i++) { + remote_rdma |= remote_caps[i] != 0; } - rdma_caps rc = {}; - memcpy(&rc, remote_caps, sizeof(rc)); - if (rc.qpn == 0) { + if (!rdma || !remote_rdma) { rdma.reset(); return; } - if (rdma_activate(rc.qpn, rc.psn, rc.gid)) { +# ifdef GGML_RPC_RDMA_APPLE + bool activated = rdma->activate(remote_caps); +# else + rdma_caps rc = {}; + memcpy(&rc, remote_caps, sizeof(rc)); + bool activated = rdma_activate(rc.qpn, rc.psn, rc.gid); +# endif + if (activated) { use_rdma = true; } else { GGML_LOG_ERROR("RDMA activate failed, staying on TCP\n"); @@ -541,6 +585,14 @@ void socket_t::impl::update_caps(const uint8_t * remote_caps) { #endif // GGML_RPC_RDMA } +bool socket_t::impl::flush() { +#ifdef GGML_RPC_RDMA_APPLE + if (use_rdma) { + return rdma->flush(); + } +#endif + return true; +} ///////////////////////////////////////////////////////////////////////////// @@ -556,6 +608,10 @@ bool socket_t::recv_data(void * data, size_t size) { return pimpl->recv_data(data, size); } +bool socket_t::flush() { + return pimpl->flush(); +} + void socket_t::get_caps(uint8_t * local_caps) { return pimpl->get_caps(local_caps); } diff --git a/ggml/src/ggml-rpc/transport.h b/ggml/src/ggml-rpc/transport.h index 73b85cc530a..3f747ecffd9 100644 --- a/ggml/src/ggml-rpc/transport.h +++ b/ggml/src/ggml-rpc/transport.h @@ -15,6 +15,10 @@ struct socket_t { bool send_data(const void * data, size_t size); bool recv_data(void * data, size_t size); + // Must be called at every message boundary: the RDMA transport coalesces + // writes into fixed-size frames and posts the trailing partial frame only + // here. No-op on TCP. + bool flush(); socket_ptr accept(); From 482956e744ea8484256985592d2018b2940dc4bc Mon Sep 17 00:00:00 2001 From: Jonathan Clohessy Date: Tue, 25 Aug 2026 22:07:29 +0100 Subject: [PATCH 003/104] kleidiai: Rework KleidiAI Build System/Integration (llama/26077) * Rework KleidiAI Build System/Integration Signed-off-by: Jonathan Clohessy * Add fp16 guard, and fix cmake caching issue Signed-off-by: Jonathan Clohessy * Fix formatting, and rebase issue Signed-off-by: Jonathan Clohessy --------- Signed-off-by: Jonathan Clohessy --- ggml/cmake/ggml-config.cmake.in | 10 ++ ggml/src/ggml-cpu/CMakeLists.txt | 175 +++++++--------------- ggml/src/ggml-cpu/kleidiai/CMakeLists.txt | 14 ++ ggml/src/ggml-cpu/kleidiai/kernels.cpp | 149 +++++++----------- ggml/src/ggml-cpu/kleidiai/kernels.h | 5 +- ggml/src/ggml-cpu/kleidiai/kleidiai.cpp | 3 +- 6 files changed, 134 insertions(+), 222 deletions(-) create mode 100644 ggml/src/ggml-cpu/kleidiai/CMakeLists.txt diff --git a/ggml/cmake/ggml-config.cmake.in b/ggml/cmake/ggml-config.cmake.in index abe17804a5a..a28e49e8342 100644 --- a/ggml/cmake/ggml-config.cmake.in +++ b/ggml/cmake/ggml-config.cmake.in @@ -110,6 +110,16 @@ set_and_check(GGML_INCLUDE_DIR "@PACKAGE_GGML_INCLUDE_INSTALL_DIR@") set_and_check(GGML_LIB_DIR "@PACKAGE_GGML_LIB_INSTALL_DIR@") #set_and_check(GGML_BIN_DIR "@PACKAGE_GGML_BIN_INSTALL_DIR@") +if (NOT GGML_SHARED_LIB AND GGML_CPU_KLEIDIAI) + unset(KLEIDIAI_LIBRARY CACHE) + unset(KLEIDIAI_LIBRARY) + find_library(KLEIDIAI_LIBRARY kleidiai + REQUIRED + HINTS ${GGML_LIB_DIR} + NO_CMAKE_FIND_ROOT_PATH) + list(APPEND GGML_CPU_INTERFACE_LINK_LIBRARIES ${KLEIDIAI_LIBRARY}) +endif() + if(NOT TARGET ggml::ggml) find_package(Threads REQUIRED) diff --git a/ggml/src/ggml-cpu/CMakeLists.txt b/ggml/src/ggml-cpu/CMakeLists.txt index e16ac996a4a..3c6343fb2a9 100644 --- a/ggml/src/ggml-cpu/CMakeLists.txt +++ b/ggml/src/ggml-cpu/CMakeLists.txt @@ -576,10 +576,25 @@ function(ggml_add_cpu_backend_variant_impl tag_name) endif() if (GGML_CPU_KLEIDIAI) - message(STATUS "Using KleidiAI optimized kernels if applicable") + # upstream repo requires at least cmake 3.16 + if (CMAKE_VERSION VERSION_LESS 3.16) + message(FATAL_ERROR "GGML_CPU_KLEIDIAI requires CMake >= 3.16") + endif() + + set(GGML_CPU_KLEIDIAI_AARCH64 OFF) + if (GGML_SYSTEM_ARCH STREQUAL "ARM" AND + (APPLE OR WIN32 OR CMAKE_SYSTEM_NAME MATCHES "^(Linux|Android)$") AND + (CMAKE_SYSTEM_PROCESSOR MATCHES "^(aarch64|arm64|ARM64|arm64-v8a)$" OR + CMAKE_OSX_ARCHITECTURES MATCHES "arm64" OR + CMAKE_GENERATOR_PLATFORM_LWR STREQUAL "arm64" OR + CMAKE_ANDROID_ARCH_ABI STREQUAL "arm64-v8a")) + set(GGML_CPU_KLEIDIAI_AARCH64 ON) + endif() + if (NOT GGML_CPU_KLEIDIAI_AARCH64) + message(FATAL_ERROR "GGML_CPU_KLEIDIAI requires a Linux, Android, Apple, or Windows AArch64/arm64 target") + endif() - # Disable the KleidiAI tests - set(KLEIDIAI_BUILD_TESTS OFF) + message(STATUS "Using KleidiAI optimized kernels if applicable") # Fetch KleidiAI sources: include(FetchContent) @@ -595,31 +610,49 @@ function(ggml_add_cpu_backend_variant_impl tag_name) list(APPEND KLEIDIAI_FETCH_ARGS DOWNLOAD_EXTRACT_TIMESTAMP NEW) endif() - if (CMAKE_VERSION VERSION_GREATER_EQUAL "3.28") - FetchContent_Declare(KleidiAI_Download - ${KLEIDIAI_FETCH_ARGS} - EXCLUDE_FROM_ALL - ) + FetchContent_Declare(kleidiai + ${KLEIDIAI_FETCH_ARGS} + ) - FetchContent_MakeAvailable(KleidiAI_Download) - FetchContent_GetProperties(KleidiAI_Download SOURCE_DIR KLEIDIAI_SRC) - else() - FetchContent_Declare(KleidiAI_Download - ${KLEIDIAI_FETCH_ARGS} - ) + # Disable tests and benchmark building + set(KLEIDIAI_BUILD_TESTS OFF CACHE BOOL "" FORCE) + set(KLEIDIAI_BUILD_BENCHMARK OFF CACHE BOOL "" FORCE) - FetchContent_GetProperties(KleidiAI_Download + # Use the Populate/add_subdirectory flow for compatibility with CMake 3.16. + FetchContent_GetProperties(kleidiai + SOURCE_DIR KLEIDIAI_SRC + BINARY_DIR KLEIDIAI_BIN + POPULATED KLEIDIAI_POPULATED + ) + if (NOT KLEIDIAI_POPULATED) + FetchContent_Populate(kleidiai) + FetchContent_GetProperties(kleidiai SOURCE_DIR KLEIDIAI_SRC - POPULATED KLEIDIAI_POPULATED + BINARY_DIR KLEIDIAI_BIN ) + endif() - if (NOT KLEIDIAI_POPULATED) - FetchContent_Populate(KleidiAI_Download) - FetchContent_GetProperties(KleidiAI_Download SOURCE_DIR KLEIDIAI_SRC) + if (NOT TARGET kleidiai) + add_subdirectory( + "${CMAKE_CURRENT_SOURCE_DIR}/ggml-cpu/kleidiai" + "${CMAKE_CURRENT_BINARY_DIR}/kleidiai-wrapper" + EXCLUDE_FROM_ALL + ) + if (NOT CMAKE_SKIP_INSTALL_RULES AND + (NOT DEFINED BUILD_SHARED_LIBS OR NOT BUILD_SHARED_LIBS)) + install(TARGETS kleidiai ARCHIVE) endif() endif() - add_compile_definitions(GGML_USE_CPU_KLEIDIAI) + if (NOT TARGET kleidiai) + message(FATAL_ERROR "KleidiAI target was not created") + endif() + + set_target_properties(kleidiai PROPERTIES POSITION_INDEPENDENT_CODE ON) + + target_link_libraries(${GGML_CPU_NAME} PRIVATE kleidiai) + + target_compile_definitions(${GGML_CPU_NAME} PRIVATE GGML_USE_CPU_KLEIDIAI) list(APPEND GGML_CPU_SOURCES ggml-cpu/kleidiai/kleidiai.cpp @@ -627,108 +660,6 @@ function(ggml_add_cpu_backend_variant_impl tag_name) ggml-cpu/kleidiai/kleidiai.h ggml-cpu/kleidiai/kernels.h ) - - # KleidiAI - include_directories( - ${KLEIDIAI_SRC}/ - ${KLEIDIAI_SRC}/kai/ - ${KLEIDIAI_SRC}/kai/ukernels/ - ${KLEIDIAI_SRC}/kai/ukernels/matmul/ - ${KLEIDIAI_SRC}/kai/ukernels/matmul/matmul_clamp_f32_qsi8d32p_qsi4c32p/ - ${KLEIDIAI_SRC}/kai/ukernels/matmul/matmul_clamp_f32_qai8dxp_qsi8cxp/ - ${KLEIDIAI_SRC}/kai/ukernels/matmul/matmul_clamp_fp32_bf16p_bf16p/ - ${KLEIDIAI_SRC}/kai/ukernels/matmul/matmul_clamp_f32_f16p_qsi4c32p/ - ${KLEIDIAI_SRC}/kai/ukernels/matmul/matmul_clamp_f32_f32p_f32p/ - ${KLEIDIAI_SRC}/kai/ukernels/matmul/matmul_clamp_f32_f32_f32p/ - ${KLEIDIAI_SRC}/kai/ukernels/matmul/pack/) - - set(ARCH_FLAGS_TEMP "${ARCH_FLAGS}") - if (NOT ARCH_FLAGS_TEMP) - string(REGEX MATCH "-march=[^ ]+" ARCH_FLAGS_TEMP "${CMAKE_C_FLAGS}") - endif() - string(FIND "${ARCH_FLAGS_TEMP}" "+dotprod" DOTPROD_ENABLED) - string(FIND "${ARCH_FLAGS_TEMP}" "+i8mm" I8MM_ENABLED) - string(FIND "${ARCH_FLAGS_TEMP}" "+sme" SME_ENABLED) - string(FIND "${ARCH_FLAGS_TEMP}" "+sve" SVE_ENABLED) - - set(PRIVATE_ARCH_FLAGS ${ARCH_FLAGS_TEMP}) - - list(APPEND GGML_KLEIDIAI_SOURCES - ${KLEIDIAI_SRC}/kai/ukernels/matmul/pack/kai_lhs_quant_pack_qsi8d32p_f32.c - ${KLEIDIAI_SRC}/kai/ukernels/matmul/pack/kai_lhs_quant_pack_qsi8d32p4x8sb_f32_neon.c - ${KLEIDIAI_SRC}/kai/ukernels/matmul/pack/kai_rhs_pack_nxk_qsi4c32ps1s0scalef16_qsu4c32s16s0_neon.c - ${KLEIDIAI_SRC}/kai/ukernels/matmul/pack/kai_lhs_quant_pack_qsi8d32p_f32_neon.c - ${KLEIDIAI_SRC}/kai/ukernels/matmul/pack/kai_rhs_pack_nxk_qsi4c32pscalef16_qsu4c32s16s0.c - ${KLEIDIAI_SRC}/kai/ukernels/matmul/pack/kai_lhs_quant_pack_qai8dxp_f32.c - ${KLEIDIAI_SRC}/kai/ukernels/matmul/pack/kai_rhs_pack_nxk_qsi8cxp_qsi8cx_neon.c) - - if (NOT DOTPROD_ENABLED MATCHES -1) - list(APPEND GGML_KLEIDIAI_SOURCES - ${KLEIDIAI_SRC}/kai/ukernels/matmul/matmul_clamp_f32_qsi8d32p_qsi4c32p/kai_matmul_clamp_f32_qsi8d32p1x8_qsi4c32p4x8_1x4x32_neon_dotprod.c - ${KLEIDIAI_SRC}/kai/ukernels/matmul/matmul_clamp_f32_qsi8d32p_qsi4c32p/kai_matmul_clamp_f32_qsi8d32p1x4_qsi4c32p4x4_1x4_neon_dotprod.c - ${KLEIDIAI_SRC}/kai/ukernels/matmul/matmul_clamp_f32_qsi8d32p_qsi4c32p/kai_matmul_clamp_f32_qsi8d32p4x4_qsi4c32p4x4_16x4_neon_dotprod.c - ${KLEIDIAI_SRC}/kai/ukernels/matmul/matmul_clamp_f32_qai8dxp_qsi8cxp/kai_matmul_clamp_f32_qai8dxp4x4_qsi8cxp4x4_16x4_neon_dotprod.c - ${KLEIDIAI_SRC}/kai/ukernels/matmul/matmul_clamp_f32_qai8dxp_qsi8cxp/kai_matmul_clamp_f32_qai8dxp1x4_qsi8cxp4x4_1x4_neon_dotprod.c - ${KLEIDIAI_SRC}/kai/ukernels/matmul/matmul_clamp_f32_qai8dxp_qsi8cxp/kai_matmul_clamp_f32_qai8dxp1x8_qsi8cxp4x8_1x4_neon_dotprod.c) - endif() - - if (NOT I8MM_ENABLED MATCHES -1) - list(APPEND GGML_KLEIDIAI_SOURCES - ${KLEIDIAI_SRC}/kai/ukernels/matmul/matmul_clamp_f32_qsi8d32p_qsi4c32p/kai_matmul_clamp_f32_qsi8d32p4x8_qsi4c32p4x8_16x4_neon_i8mm.c - ${KLEIDIAI_SRC}/kai/ukernels/matmul/matmul_clamp_f32_qai8dxp_qsi8cxp/kai_matmul_clamp_f32_qai8dxp4x8_qsi8cxp4x8_16x4_neon_i8mm.c) - endif() - - if (NOT SME_ENABLED MATCHES -1) - list(APPEND GGML_KLEIDIAI_SME_SOURCES - ${KLEIDIAI_SRC}/kai/ukernels/matmul/matmul_clamp_f32_qai8dxp_qsi8cxp/kai_matmul_clamp_f32_qai8dxp1vlx4_qsi8cxp4vlx4_1vlx4vl_sme_mopa.c - ${KLEIDIAI_SRC}/kai/ukernels/matmul/matmul_clamp_f32_qai8dxp_qsi8cxp/kai_matmul_clamp_f32_qai8dxp1vlx4_qsi8cxp4vlx4_1vlx4vl_sme_mopa_asm.S - ${KLEIDIAI_SRC}/kai/ukernels/matmul/matmul_clamp_f32_qai8dxp_qsi8cxp/kai_matmul_clamp_f32_qai8dxp1x4_qsi8cxp4vlx4_1x4vl_sme_dot.c - ${KLEIDIAI_SRC}/kai/ukernels/matmul/matmul_clamp_f32_qai8dxp_qsi8cxp/kai_matmul_clamp_f32_qai8dxp1x4_qsi8cxp4vlx4_1x4vl_sme_dot_asm.S - ${KLEIDIAI_SRC}/kai/ukernels/matmul/matmul_clamp_f32_f32p_f32p/kai_matmul_clamp_f32_f32p2vlx1_f32p2vlx1b_2vlx2vl_sme_mopa.c - ${KLEIDIAI_SRC}/kai/ukernels/matmul/matmul_clamp_f32_f32p_f32p/kai_matmul_clamp_f32_f32p2vlx1_f32p2vlx1b_2vlx2vl_sme_mopa_asm.S) - set_source_files_properties(${GGML_KLEIDIAI_SME_SOURCES} - PROPERTIES COMPILE_OPTIONS "-fno-tree-vectorize;${ARCH_FLAGS_TEMP}+sve+sve2+sme") - list(APPEND GGML_CPU_SOURCES ${GGML_KLEIDIAI_SME_SOURCES}) - - list(APPEND GGML_KLEIDIAI_SME2_SOURCES - ${KLEIDIAI_SRC}/kai/ukernels/matmul/matmul_clamp_f32_qsi8d32p_qsi4c32p/kai_matmul_clamp_f32_qsi8d32p1x4_qsi4c32p4vlx4_1x4vl_sme2_sdot.c - ${KLEIDIAI_SRC}/kai/ukernels/matmul/matmul_clamp_f32_qai8dxp_qsi8cxp/kai_matmul_clamp_f32_qai8dxp1vlx4_qsi8cxp4vlx4_1vlx4vl_sme2_mopa.c - ${KLEIDIAI_SRC}/kai/ukernels/matmul/matmul_clamp_f32_qai8dxp_qsi8cxp/kai_matmul_clamp_f32_qai8dxp1vlx4_qsi8cxp4vlx4_1vlx4vl_sme2_mopa_asm.S - ${KLEIDIAI_SRC}/kai/ukernels/matmul/matmul_clamp_f32_qai8dxp_qsi8cxp/kai_matmul_clamp_f32_qai8dxp1x4_qsi8cxp4vlx4_1x4vl_sme2_dot.c - ${KLEIDIAI_SRC}/kai/ukernels/matmul/matmul_clamp_f32_qai8dxp_qsi8cxp/kai_matmul_clamp_f32_qai8dxp1x4_qsi8cxp4vlx4_1x4vl_sme2_dot_asm.S - ${KLEIDIAI_SRC}/kai/ukernels/matmul/matmul_clamp_fp32_bf16p_bf16p/kai_matmul_clamp_f32_bf16p2vlx2_bf16p2vlx2_2vlx2vl_sme2_mopa.c - ${KLEIDIAI_SRC}/kai/ukernels/matmul/matmul_clamp_fp32_bf16p_bf16p/kai_matmul_clamp_f32_bf16p2vlx2_bf16p2vlx2_2vlx2vl_sme2_mopa_asm.S - ${KLEIDIAI_SRC}/kai/ukernels/matmul/matmul_clamp_f32_f16p_qsi4c32p/kai_matmul_clamp_f32_f16p1vlx2_qsi4c32p4vlx2_1vlx4vl_sme2_mopa.c - ${KLEIDIAI_SRC}/kai/ukernels/matmul/matmul_clamp_f32_f16p_qsi4c32p/kai_matmul_clamp_f32_f16p1vlx2_qsi4c32p4vlx2_1vlx4vl_sme2_mopa_asm.S - ${KLEIDIAI_SRC}/kai/ukernels/matmul/matmul_clamp_f32_f32p_f32p/kai_matmul_clamp_f32_f32p2vlx1_f32p2vlx1biasf32_sme2_mopa.c - ${KLEIDIAI_SRC}/kai/ukernels/matmul/matmul_clamp_f32_f32p_f32p/kai_matmul_clamp_f32_f32p2vlx1_f32p2vlx1biasf32_sme2_mopa_asm.S - ${KLEIDIAI_SRC}/kai/ukernels/matmul/matmul_clamp_f32_f32_f32p/kai_matmul_clamp_f32_f32_f32p2vlx1b_1x16vl_sme2_mla.c - ${KLEIDIAI_SRC}/kai/ukernels/matmul/matmul_clamp_f32_f32_f32p/kai_matmul_clamp_f32_f32_f32p2vlx1b_1x16vl_sme2_mla_asm.S - ${KLEIDIAI_SRC}/kai/ukernels/matmul/pack/kai_lhs_pack_bf16p2vlx2_f32_sme.c - ${KLEIDIAI_SRC}/kai/ukernels/matmul/pack/kai_rhs_pack_kxn_bf16p2vlx2b_f32_x32_sme.c - ${KLEIDIAI_SRC}/kai/ukernels/matmul/pack/kai_lhs_pack_f16pmrx2_f32_neon.c - ${KLEIDIAI_SRC}/kai/ukernels/matmul/pack/kai_lhs_pack_f32p2vlx1_f32_sme.c - ${KLEIDIAI_SRC}/kai/ukernels/matmul/pack/kai_lhs_pack_f32p2vlx1_f32_sme_asm.S - ${KLEIDIAI_SRC}/kai/ukernels/matmul/pack/kai_rhs_pack_nxk_f32p2vlx1biasf32_f32_f32_sme.c - ${KLEIDIAI_SRC}/kai/ukernels/matmul/pack/kai_rhs_pack_nxk_f32p2vlx1biasf32_f32_f32_sme_asm.S - ${KLEIDIAI_SRC}/kai/kai_common_sme_asm.S) - set_source_files_properties(${GGML_KLEIDIAI_SME2_SOURCES} - PROPERTIES COMPILE_OPTIONS "-fno-tree-vectorize;${ARCH_FLAGS_TEMP}+sve+sve2+sme2+fp16") - list(APPEND GGML_CPU_SOURCES ${GGML_KLEIDIAI_SME2_SOURCES}) - set(PRIVATE_ARCH_FLAGS "-fno-tree-vectorize;${PRIVATE_ARCH_FLAGS}") - endif() - - if (NOT SVE_ENABLED MATCHES -1) - list(APPEND GGML_KLEIDIAI_SOURCES - ${KLEIDIAI_SRC}/kai/kai_common_sve_asm.S - ${KLEIDIAI_SRC}/kai/ukernels/matmul/matmul_clamp_f32_qsi8d32p_qsi4c32p/kai_matmul_clamp_f32_qsi8d32p1x8_qsi4c32p8x8_1x8_sve_dotprod_asm.S - ${KLEIDIAI_SRC}/kai/ukernels/matmul/matmul_clamp_f32_qsi8d32p_qsi4c32p/kai_matmul_clamp_f32_qsi8d32p1x8_qsi4c32p8x8_1x8_sve_dotprod.c - ${KLEIDIAI_SRC}/kai/ukernels/matmul/matmul_clamp_f32_qsi8d32p_qsi4c32p/kai_matmul_clamp_f32_qsi8d32p4x8_qsi4c32p8x8_16x8_sve_i8mm_asm.S - ${KLEIDIAI_SRC}/kai/ukernels/matmul/matmul_clamp_f32_qsi8d32p_qsi4c32p/kai_matmul_clamp_f32_qsi8d32p4x8_qsi4c32p8x8_16x8_sve_i8mm.c) - endif() - - set_source_files_properties(${GGML_KLEIDIAI_SOURCES} PROPERTIES COMPILE_OPTIONS "${PRIVATE_ARCH_FLAGS}") - list(APPEND GGML_CPU_SOURCES ${GGML_KLEIDIAI_SOURCES}) endif() message(STATUS "Adding CPU backend variant ${GGML_CPU_NAME}: ${ARCH_FLAGS} ${ARCH_DEFINITIONS}") diff --git a/ggml/src/ggml-cpu/kleidiai/CMakeLists.txt b/ggml/src/ggml-cpu/kleidiai/CMakeLists.txt new file mode 100644 index 00000000000..b36cb6d3a9d --- /dev/null +++ b/ggml/src/ggml-cpu/kleidiai/CMakeLists.txt @@ -0,0 +1,14 @@ +set(BUILD_SHARED_LIBS OFF) +set(CMAKE_SKIP_INSTALL_RULES TRUE) + +add_subdirectory("${KLEIDIAI_SRC}" "${KLEIDIAI_BIN}" EXCLUDE_FROM_ALL) + +if (NOT TARGET kleidiai) + message(FATAL_ERROR "KleidiAI target was not created") +endif() + +if (MSVC) + target_compile_options(kleidiai PRIVATE $<$:/WX->) +else() + target_compile_options(kleidiai PRIVATE $<$:-Wno-error>) +endif() diff --git a/ggml/src/ggml-cpu/kleidiai/kernels.cpp b/ggml/src/ggml-cpu/kleidiai/kernels.cpp index 70b519f29ce..d4551298f86 100644 --- a/ggml/src/ggml-cpu/kleidiai/kernels.cpp +++ b/ggml/src/ggml-cpu/kleidiai/kernels.cpp @@ -3,44 +3,44 @@ // // KleidiAI micro-kernels -#include "kai_matmul_clamp_f32_qsi8d32p_qsi4c32p_interface.h" -#include "kai_matmul_clamp_f32_qai8dxp_qsi8cxp_interface.h" -#include "kai_matmul_clamp_f32_qsi8d32p1x8_qsi4c32p4x8_1x4x32_neon_dotprod.h" -#include "kai_matmul_clamp_f32_qsi8d32p1x4_qsi4c32p4x4_1x4_neon_dotprod.h" -#include "kai_matmul_clamp_f32_qsi8d32p4x4_qsi4c32p4x4_16x4_neon_dotprod.h" -#include "kai_matmul_clamp_f32_qsi8d32p4x8_qsi4c32p4x8_16x4_neon_i8mm.h" -#include "kai_matmul_clamp_f32_qsi8d32p1x4_qsi4c32p4vlx4_1x4vl_sme2_sdot.h" -#include "kai_matmul_clamp_f32_bf16p2vlx2_bf16p2vlx2_2vlx2vl_sme2_mopa.h" -#include "kai_matmul_clamp_f32_qai8dxp1vlx4_qsi8cxp4vlx4_1vlx4vl_sme2_mopa.h" -#include "kai_matmul_clamp_f32_qai8dxp1x4_qsi8cxp4vlx4_1x4vl_sme2_dot.h" -#include "kai_matmul_clamp_f32_qai8dxp1vlx4_qsi8cxp4vlx4_1vlx4vl_sme_mopa.h" -#include "kai_matmul_clamp_f32_qai8dxp1x4_qsi8cxp4vlx4_1x4vl_sme_dot.h" -#include "kai_matmul_clamp_f32_qai8dxp1x8_qsi8cxp4x8_1x4_neon_dotprod.h" -#include "kai_matmul_clamp_f32_qai8dxp1x4_qsi8cxp4x4_1x4_neon_dotprod.h" -#include "kai_matmul_clamp_f32_qai8dxp4x4_qsi8cxp4x4_16x4_neon_dotprod.h" -#include "kai_matmul_clamp_f32_qai8dxp4x8_qsi8cxp4x8_16x4_neon_i8mm.h" -#include "kai_matmul_clamp_f32_qsi8d32p4x8_qsi4c32p8x8_16x8_sve_i8mm.h" -#include "kai_matmul_clamp_f32_qsi8d32p1x8_qsi4c32p8x8_1x8_sve_dotprod.h" -#include "kai_matmul_clamp_f32_f16p1vlx2_qsi4c32p4vlx2_1vlx4vl_sme2_mopa.h" -#include "kai_matmul_clamp_f32_f32p2vlx1_f32p2vlx1biasf32_sme2_mopa.h" -#include "kai_matmul_clamp_f32_f32_f32p2vlx1b_1x16vl_sme2_mla.h" -#include "kai_matmul_clamp_f32_f32p2vlx1_f32p2vlx1b_2vlx2vl_sme_mopa.h" - -#include "kai_lhs_pack_bf16p2vlx2_f32_sme.h" -#include "kai_lhs_pack_f32p2vlx1_f32_sme.h" -#include "kai_lhs_quant_pack_qsi8d32p_f32.h" -#include "kai_lhs_quant_pack_qsi8d32p4x8sb_f32_neon.h" -#include "kai_lhs_quant_pack_qsi8d32p_f32_neon.h" -#include "kai_lhs_quant_pack_qai8dxp_f32.h" - -#include "kai_rhs_pack_kxn_bf16p2vlx2b_f32_x32_sme.h" -#include "kai_rhs_pack_nxk_f32p2vlx1biasf32_f32_f32_sme.h" -#include "kai_rhs_pack_nxk_qsi4c32pscalef16_qsu4c32s16s0.h" -#include "kai_rhs_pack_nxk_qsi4c32ps1s0scalef16_qsu4c32s16s0_neon.h" -#include "kai_rhs_pack_nxk_qsi8cxp_qsi8cx_neon.h" -#include "kai_lhs_pack_f16pmrx2_f32_neon.h" - -#include "kai_common.h" +#include "kai/ukernels/matmul/matmul_clamp_f32_qsi8d32p_qsi4c32p/kai_matmul_clamp_f32_qsi8d32p_qsi4c32p_interface.h" +#include "kai/ukernels/matmul/matmul_clamp_f32_qai8dxp_qsi8cxp/kai_matmul_clamp_f32_qai8dxp_qsi8cxp_interface.h" +#include "kai/ukernels/matmul/matmul_clamp_f32_qsi8d32p_qsi4c32p/kai_matmul_clamp_f32_qsi8d32p1x8_qsi4c32p4x8_1x4x32_neon_dotprod.h" +#include "kai/ukernels/matmul/matmul_clamp_f32_qsi8d32p_qsi4c32p/kai_matmul_clamp_f32_qsi8d32p1x4_qsi4c32p4x4_1x4_neon_dotprod.h" +#include "kai/ukernels/matmul/matmul_clamp_f32_qsi8d32p_qsi4c32p/kai_matmul_clamp_f32_qsi8d32p4x4_qsi4c32p4x4_16x4_neon_dotprod.h" +#include "kai/ukernels/matmul/matmul_clamp_f32_qsi8d32p_qsi4c32p/kai_matmul_clamp_f32_qsi8d32p4x8_qsi4c32p4x8_16x4_neon_i8mm.h" +#include "kai/ukernels/matmul/matmul_clamp_f32_qsi8d32p_qsi4c32p/kai_matmul_clamp_f32_qsi8d32p1x4_qsi4c32p4vlx4_1x4vl_sme2_sdot.h" +#include "kai/ukernels/matmul/matmul_clamp_fp32_bf16p_bf16p/kai_matmul_clamp_f32_bf16p2vlx2_bf16p2vlx2_2vlx2vl_sme2_mopa.h" +#include "kai/ukernels/matmul/matmul_clamp_f32_qai8dxp_qsi8cxp/kai_matmul_clamp_f32_qai8dxp1vlx4_qsi8cxp4vlx4_1vlx4vl_sme2_mopa.h" +#include "kai/ukernels/matmul/matmul_clamp_f32_qai8dxp_qsi8cxp/kai_matmul_clamp_f32_qai8dxp1x4_qsi8cxp4vlx4_1x4vl_sme2_dot.h" +#include "kai/ukernels/matmul/matmul_clamp_f32_qai8dxp_qsi8cxp/kai_matmul_clamp_f32_qai8dxp1vlx4_qsi8cxp4vlx4_1vlx4vl_sme_mopa.h" +#include "kai/ukernels/matmul/matmul_clamp_f32_qai8dxp_qsi8cxp/kai_matmul_clamp_f32_qai8dxp1x4_qsi8cxp4vlx4_1x4vl_sme_dot.h" +#include "kai/ukernels/matmul/matmul_clamp_f32_qai8dxp_qsi8cxp/kai_matmul_clamp_f32_qai8dxp1x8_qsi8cxp4x8_1x4_neon_dotprod.h" +#include "kai/ukernels/matmul/matmul_clamp_f32_qai8dxp_qsi8cxp/kai_matmul_clamp_f32_qai8dxp1x4_qsi8cxp4x4_1x4_neon_dotprod.h" +#include "kai/ukernels/matmul/matmul_clamp_f32_qai8dxp_qsi8cxp/kai_matmul_clamp_f32_qai8dxp4x4_qsi8cxp4x4_16x4_neon_dotprod.h" +#include "kai/ukernels/matmul/matmul_clamp_f32_qai8dxp_qsi8cxp/kai_matmul_clamp_f32_qai8dxp4x8_qsi8cxp4x8_16x4_neon_i8mm.h" +#include "kai/ukernels/matmul/matmul_clamp_f32_qsi8d32p_qsi4c32p/kai_matmul_clamp_f32_qsi8d32p4x8_qsi4c32p8x8_16x8_sve_i8mm.h" +#include "kai/ukernels/matmul/matmul_clamp_f32_qsi8d32p_qsi4c32p/kai_matmul_clamp_f32_qsi8d32p1x8_qsi4c32p8x8_1x8_sve_dotprod.h" +#include "kai/ukernels/matmul/matmul_clamp_f32_f16p_qsi4c32p/kai_matmul_clamp_f32_f16p1vlx2_qsi4c32p4vlx2_1vlx4vl_sme2_mopa.h" +#include "kai/ukernels/matmul/matmul_clamp_f32_f32p_f32p/kai_matmul_clamp_f32_f32p2vlx1_f32p2vlx1biasf32_sme2_mopa.h" +#include "kai/ukernels/matmul/matmul_clamp_f32_f32_f32p/kai_matmul_clamp_f32_f32_f32p2vlx1b_1x16vl_sme2_mla.h" +#include "kai/ukernels/matmul/matmul_clamp_f32_f32p_f32p/kai_matmul_clamp_f32_f32p2vlx1_f32p2vlx1b_2vlx2vl_sme_mopa.h" + +#include "kai/ukernels/matmul/pack/kai_lhs_pack_bf16p2vlx2_f32_sme.h" +#include "kai/ukernels/matmul/pack/kai_lhs_pack_f32p2vlx1_f32_sme.h" +#include "kai/ukernels/matmul/pack/kai_lhs_quant_pack_qsi8d32p_f32.h" +#include "kai/ukernels/matmul/pack/kai_lhs_quant_pack_qsi8d32p4x8sb_f32_neon.h" +#include "kai/ukernels/matmul/pack/kai_lhs_quant_pack_qsi8d32p_f32_neon.h" +#include "kai/ukernels/matmul/pack/kai_lhs_quant_pack_qai8dxp_f32.h" + +#include "kai/ukernels/matmul/pack/kai_rhs_pack_kxn_bf16p2vlx2b_f32_x32_sme.h" +#include "kai/ukernels/matmul/pack/kai_rhs_pack_nxk_f32p2vlx1biasf32_f32_f32_sme.h" +#include "kai/ukernels/matmul/pack/kai_rhs_pack_nxk_qsi4c32pscalef16_qsu4c32s16s0.h" +#include "kai/ukernels/matmul/pack/kai_rhs_pack_nxk_qsi4c32ps1s0scalef16_qsu4c32s16s0_neon.h" +#include "kai/ukernels/matmul/pack/kai_rhs_pack_nxk_qsi8cxp_qsi8cx_neon.h" +#include "kai/ukernels/matmul/pack/kai_lhs_pack_f16pmrx2_f32_neon.h" + +#include "kai/kai_common.h" #include "simd-mappings.h" @@ -328,9 +328,8 @@ static void dequantize_row_qsi8cxp( } static ggml_kleidiai_kernels gemm_gemv_kernels[] = { -#if defined(__ARM_FEATURE_SME) { - /* SME GEMM */ + /* SME2 GEMM */ /* .kern_info = */ { /* .get_m_step = */ kai_get_m_step_matmul_clamp_f32_f16p1vlx2_qsi4c32p4vlx2_1vlx4vl_sme2_mopa, /* .get_n_step = */ kai_get_n_step_matmul_clamp_f32_f16p1vlx2_qsi4c32p4vlx2_1vlx4vl_sme2_mopa, @@ -351,7 +350,7 @@ static ggml_kleidiai_kernels gemm_gemv_kernels[] = { /* .packed_size_ex = */ &lhs_ps_fn6, /* .pack_func_ex = */ &lhs_pack_void_fn10, }, - /* SME GEMV */ + /* SME2 GEMV */ /* .kern_info = */ { /* .get_m_step = */ kai_get_m_step_matmul_clamp_f32_qsi8d32p1x4_qsi4c32p4vlx4_1x4vl_sme2_sdot, /* .get_n_step = */ kai_get_n_step_matmul_clamp_f32_qsi8d32p1x4_qsi4c32p4vlx4_1x4vl_sme2_sdot, @@ -378,13 +377,13 @@ static ggml_kleidiai_kernels gemm_gemv_kernels[] = { /* .packed_stride_ex = */ &rhs_stride_fn4, /* .pack_func_ex = */ &rhs_pack_fn12, }, - /* .required_cpu = */ CPU_FEATURE_SME2, + /* .required_cpu = */ CPU_FEATURE_SME2 | CPU_FEATURE_FP16, /* .lhs_type = */ GGML_TYPE_F32, /* .rhs_type = */ GGML_TYPE_Q4_0, /* .op_type = */ GGML_TYPE_F32, }, { - /* SME GEMM */ + /* SME2 GEMM */ /* .kern_info = */ { /* .get_m_step = */ kai_get_m_step_matmul_clamp_f32_bf16p2vlx2_bf16p2vlx2_2vlx2vl_sme2_mopa, /* .get_n_step = */ kai_get_n_step_matmul_clamp_f32_bf16p2vlx2_bf16p2vlx2_2vlx2vl_sme2_mopa, @@ -404,7 +403,7 @@ static ggml_kleidiai_kernels gemm_gemv_kernels[] = { /* .packed_size_ex = */ &lhs_ps_fn5, /* .pack_func_ex = */ &lhs_pack_void_fn9, }, - /* SME GEMV */ + /* SME2 GEMV */ /* .kern_info = */ { /* .get_m_step = */ kai_get_m_step_matmul_clamp_f32_bf16p2vlx2_bf16p2vlx2_2vlx2vl_sme2_mopa, /* .get_n_step = */ kai_get_n_step_matmul_clamp_f32_bf16p2vlx2_bf16p2vlx2_2vlx2vl_sme2_mopa, @@ -436,9 +435,7 @@ static ggml_kleidiai_kernels gemm_gemv_kernels[] = { /* .rhs_type = */ GGML_TYPE_F16, /* .op_type = */ GGML_TYPE_F32, }, -#endif #if defined(__APPLE__) -#if defined(__ARM_FEATURE_DOTPROD) { /* DOTPROD GEMM */ /* .kern_info = */ { @@ -492,8 +489,6 @@ static ggml_kleidiai_kernels gemm_gemv_kernels[] = { /* .rhs_type = */ GGML_TYPE_Q4_0, /* .op_type = */ GGML_TYPE_F32, }, -#endif -#if defined(__ARM_FEATURE_MATMUL_INT8) { /* i8mm GEMM */ /* .kern_info = */ { @@ -515,7 +510,7 @@ static ggml_kleidiai_kernels gemm_gemv_kernels[] = { /* .packed_size_ex = */ &lhs_ps_fn6, /* .pack_func_ex = */ &lhs_pack_float_fn10, }, - /* i8mm GEMV */ + /* DOTPROD GEMV */ /* .kern_info = */ { /* .get_m_step = */ kai_get_m_step_matmul_clamp_f32_qsi8d32p1x8_qsi4c32p4x8_1x4x32_neon_dotprod, /* .get_n_step = */ kai_get_n_step_matmul_clamp_f32_qsi8d32p1x8_qsi4c32p4x8_1x4x32_neon_dotprod, @@ -542,14 +537,12 @@ static ggml_kleidiai_kernels gemm_gemv_kernels[] = { /* .packed_stride_ex = */ &rhs_stride_fn4, /* .pack_func_ex = */ &rhs_pack_fn12, }, - /* .required_cpu = */ CPU_FEATURE_I8MM, + /* .required_cpu = */ CPU_FEATURE_I8MM | CPU_FEATURE_DOTPROD, /* .lhs_type = */ GGML_TYPE_F32, /* .rhs_type = */ GGML_TYPE_Q4_0, /* .op_type = */ GGML_TYPE_F32, }, -#endif #else -#if defined(__ARM_FEATURE_SVE) { /* SVE i8mm GEMM */ /* .kern_info = */ { @@ -603,8 +596,6 @@ static ggml_kleidiai_kernels gemm_gemv_kernels[] = { /* .rhs_type = */ GGML_TYPE_Q4_0, /* .op_type = */ GGML_TYPE_F32, }, -#endif -#if defined(__ARM_FEATURE_MATMUL_INT8) { /* i8mm GEMM */ /* .kern_info = */ { @@ -626,7 +617,7 @@ static ggml_kleidiai_kernels gemm_gemv_kernels[] = { /* .packed_size_ex = */ &lhs_ps_fn6, /* .pack_func_ex = */ &lhs_pack_float_fn10, }, - /* i8mm GEMV */ + /* DOTPROD GEMV */ /* .kern_info = */ { /* .get_m_step = */ kai_get_m_step_matmul_clamp_f32_qsi8d32p1x8_qsi4c32p4x8_1x4x32_neon_dotprod, /* .get_n_step = */ kai_get_n_step_matmul_clamp_f32_qsi8d32p1x8_qsi4c32p4x8_1x4x32_neon_dotprod, @@ -653,13 +644,11 @@ static ggml_kleidiai_kernels gemm_gemv_kernels[] = { /* .packed_stride_ex = */ &rhs_stride_fn4, /* .pack_func_ex = */ &rhs_pack_fn12, }, - /* .required_cpu = */ CPU_FEATURE_I8MM, + /* .required_cpu = */ CPU_FEATURE_I8MM | CPU_FEATURE_DOTPROD, /* .lhs_type = */ GGML_TYPE_F32, /* .rhs_type = */ GGML_TYPE_Q4_0, /* .op_type = */ GGML_TYPE_F32, }, -#endif // __ARM_FEATURE_MATMUL_INT8 -#if defined(__ARM_FEATURE_DOTPROD) { /* DOTPROD GEMM */ /* .kern_info = */ { @@ -713,15 +702,13 @@ static ggml_kleidiai_kernels gemm_gemv_kernels[] = { /* .rhs_type = */ GGML_TYPE_Q4_0, /* .op_type = */ GGML_TYPE_F32, }, -#endif #endif { /* Sentinel */ } }; static ggml_kleidiai_kernels gemm_gemv_kernels_q8[] = { -#if defined(__ARM_FEATURE_SME) { - /* SME GEMM */ + /* SME2 GEMM */ { /* .get_m_step = */ kai_get_m_step_matmul_clamp_f32_qai8dxp1vlx4_qsi8cxp4vlx4_1vlx4vl_sme2_mopa, /* .get_n_step = */ kai_get_n_step_matmul_clamp_f32_qai8dxp1vlx4_qsi8cxp4vlx4_1vlx4vl_sme2_mopa, @@ -741,7 +728,7 @@ static ggml_kleidiai_kernels gemm_gemv_kernels_q8[] = { /* .packed_size_ex = */ &lhs_ps_fn5, /* .pack_func_ex = */ &lhs_pack_float_fn9_no_bl, }, - /* SME GEMV */ + /* SME2 GEMV */ { /* .get_m_step = */ kai_get_m_step_matmul_clamp_f32_qai8dxp1x4_qsi8cxp4vlx4_1x4vl_sme2_dot, /* .get_n_step = */ kai_get_n_step_matmul_clamp_f32_qai8dxp1x4_qsi8cxp4vlx4_1x4vl_sme2_dot, @@ -826,8 +813,6 @@ static ggml_kleidiai_kernels gemm_gemv_kernels_q8[] = { /* .rhs_type = */ GGML_TYPE_Q8_0, /* .op_type = */ GGML_TYPE_F32, }, -#endif -#if defined(__ARM_FEATURE_MATMUL_INT8) { /* I8MM GEMM */ { @@ -876,13 +861,11 @@ static ggml_kleidiai_kernels gemm_gemv_kernels_q8[] = { /* .packed_stride_ex = */ &rhs_stride_fn4, /* .pack_func_ex = */ &rhs_pack_scale_fn12, }, - /* .required_cpu = */ CPU_FEATURE_I8MM, + /* .required_cpu = */ CPU_FEATURE_I8MM | CPU_FEATURE_DOTPROD, /* .lhs_type = */ GGML_TYPE_F32, /* .rhs_type = */ GGML_TYPE_Q8_0, /* .op_type = */ GGML_TYPE_F32, }, -#endif -#if defined(__ARM_FEATURE_DOTPROD) { /* DOTPROD GEMM */ { @@ -936,12 +919,10 @@ static ggml_kleidiai_kernels gemm_gemv_kernels_q8[] = { /* .rhs_type = */ GGML_TYPE_Q8_0, /* .op_type = */ GGML_TYPE_F32, }, -#endif { /* Sentinel */ } }; static ggml_kleidiai_kernels ggml_kleidiai_kernels_f32[] = { -#if defined(__ARM_FEATURE_SME) { /* SME2 GEMM */ { @@ -1048,7 +1029,6 @@ static ggml_kleidiai_kernels ggml_kleidiai_kernels_f32[] = { /* .rhs_type = */ GGML_TYPE_F32, /* .op_type = */ GGML_TYPE_F32, }, -#endif { /* Sentinel */ } }; @@ -1056,10 +1036,6 @@ ggml_kleidiai_kernels * ggml_kleidiai_select_kernels(cpu_feature cpu_features, c ggml_kleidiai_kernels * kernel = nullptr; if (tensor->op == GGML_OP_MUL_MAT && tensor->src[0] != nullptr && tensor->src[1] != nullptr) { -#if defined(__ARM_FEATURE_SME) || \ - defined(__ARM_FEATURE_DOTPROD) || \ - defined(__ARM_FEATURE_MATMUL_INT8) || \ - defined(__ARM_FEATURE_SVE) auto try_table = [&](auto & table) { for (size_t i = 0; i < NELEMS(table) - 1; ++i) { if ((cpu_features & table[i].required_cpu) == table[i].required_cpu && @@ -1080,12 +1056,6 @@ ggml_kleidiai_kernels * ggml_kleidiai_select_kernels(cpu_feature cpu_features, c } else { try_table(gemm_gemv_kernels); } -#else - GGML_UNUSED(gemm_gemv_kernels); - GGML_UNUSED(gemm_gemv_kernels_q8); - GGML_UNUSED(ggml_kleidiai_kernels_f32); - GGML_UNUSED(cpu_features); -#endif } return kernel; @@ -1094,19 +1064,13 @@ ggml_kleidiai_kernels * ggml_kleidiai_select_kernels(cpu_feature cpu_features, c ggml_kleidiai_kernels * ggml_kleidiai_select_kernels_q4_0(cpu_feature features) { ggml_kleidiai_kernels * kernels = nullptr; -#if defined(__ARM_FEATURE_SME) || \ - defined(__ARM_FEATURE_DOTPROD) || \ - defined(__ARM_FEATURE_MATMUL_INT8) || \ - defined(__ARM_FEATURE_SVE) for (size_t i = 0; i < NELEMS(gemm_gemv_kernels) - 1; ++i) { - if ((features & gemm_gemv_kernels[i].required_cpu) == gemm_gemv_kernels[i].required_cpu) { + if ((features & gemm_gemv_kernels[i].required_cpu) == gemm_gemv_kernels[i].required_cpu && + gemm_gemv_kernels[i].rhs_type == GGML_TYPE_Q4_0) { kernels = &gemm_gemv_kernels[i]; break; } } -#else - GGML_UNUSED(features); -#endif return kernels; } @@ -1114,16 +1078,12 @@ ggml_kleidiai_kernels * ggml_kleidiai_select_kernels_q4_0(cpu_feature features) ggml_kleidiai_kernels * ggml_kleidiai_select_kernels_q8_0(cpu_feature features) { ggml_kleidiai_kernels * kernels = nullptr; -#if defined(__ARM_FEATURE_SME) || defined(__ARM_FEATURE_DOTPROD) || defined(__ARM_FEATURE_MATMUL_INT8) for (size_t i = 0; i < NELEMS(gemm_gemv_kernels_q8) - 1; ++i) { if ((features & gemm_gemv_kernels_q8[i].required_cpu) == gemm_gemv_kernels_q8[i].required_cpu) { kernels = &gemm_gemv_kernels_q8[i]; break; } } -#else - GGML_UNUSED(features); -#endif return kernels; } @@ -1131,16 +1091,11 @@ ggml_kleidiai_kernels * ggml_kleidiai_select_kernels_q8_0(cpu_feature features) ggml_kleidiai_kernels * ggml_kleidiai_select_kernels_f32(cpu_feature features) { ggml_kleidiai_kernels * kernels = nullptr; -#if defined(__ARM_FEATURE_SME) for (size_t i = 0; i < NELEMS(ggml_kleidiai_kernels_f32) - 1; ++i) { if ((features & ggml_kleidiai_kernels_f32[i].required_cpu) == ggml_kleidiai_kernels_f32[i].required_cpu) { kernels = &ggml_kleidiai_kernels_f32[i]; break; } } -#else - GGML_UNUSED(features); -#endif - return kernels; } diff --git a/ggml/src/ggml-cpu/kleidiai/kernels.h b/ggml/src/ggml-cpu/kleidiai/kernels.h index 0da5e65a0a8..1da8610eae7 100644 --- a/ggml/src/ggml-cpu/kleidiai/kernels.h +++ b/ggml/src/ggml-cpu/kleidiai/kernels.h @@ -1,4 +1,4 @@ -// SPDX-FileCopyrightText: Copyright 2025 Arm Limited and/or its affiliates +// SPDX-FileCopyrightText: Copyright 2025-2026 Arm Limited and/or its affiliates // SPDX-License-Identifier: MIT // @@ -12,7 +12,8 @@ enum cpu_feature { CPU_FEATURE_I8MM = 2, CPU_FEATURE_SVE = 4, CPU_FEATURE_SME = 8, - CPU_FEATURE_SME2 = 16 + CPU_FEATURE_SME2 = 16, + CPU_FEATURE_FP16 = 32 }; inline cpu_feature& operator|=(cpu_feature& lhs, cpu_feature rhs) { diff --git a/ggml/src/ggml-cpu/kleidiai/kleidiai.cpp b/ggml/src/ggml-cpu/kleidiai/kleidiai.cpp index 6729ae8422f..92d7fd644f7 100644 --- a/ggml/src/ggml-cpu/kleidiai/kleidiai.cpp +++ b/ggml/src/ggml-cpu/kleidiai/kleidiai.cpp @@ -48,7 +48,7 @@ #include "kernels.h" -#include "kai_common.h" +#include "kai/kai_common.h" #define GGML_COMMON_DECL_CPP #include "ggml-common.h" @@ -316,6 +316,7 @@ static void init_kleidiai_context(void) { ctx.features = (runtime_feat.has_dotprod ? CPU_FEATURE_DOTPROD : CPU_FEATURE_NONE) | (runtime_feat.has_i8mm ? CPU_FEATURE_I8MM : CPU_FEATURE_NONE) | + (runtime_feat.has_fp16 ? CPU_FEATURE_FP16 : CPU_FEATURE_NONE) | (runtime_feat.sve_cnt == QK8_0 ? CPU_FEATURE_SVE : CPU_FEATURE_NONE); if (env_threads) { From 8df657a2fa88a9778c99300ac7e2a52734ed57e3 Mon Sep 17 00:00:00 2001 From: Max Krasnyansky Date: Tue, 25 Aug 2026 22:27:51 -0700 Subject: [PATCH 004/104] ggml-meta: propagate buffer usage and call init on the new tensors (llama/27586) --- ggml/src/ggml-backend-impl.h | 1 + ggml/src/ggml-backend-meta.cpp | 20 ++++++++++++++++++-- ggml/src/ggml-backend.cpp | 2 ++ 3 files changed, 21 insertions(+), 2 deletions(-) diff --git a/ggml/src/ggml-backend-impl.h b/ggml/src/ggml-backend-impl.h index 9c56ec30c5f..40cea024c3d 100644 --- a/ggml/src/ggml-backend-impl.h +++ b/ggml/src/ggml-backend-impl.h @@ -83,6 +83,7 @@ extern "C" { GGML_API ggml_backend_buffer_t ggml_backend_multi_buffer_alloc_buffer(ggml_backend_buffer_t * buffers, size_t n_buffers); GGML_API bool ggml_backend_buffer_is_multi_buffer(ggml_backend_buffer_t buffer); GGML_API void ggml_backend_multi_buffer_set_usage(ggml_backend_buffer_t buffer, enum ggml_backend_buffer_usage usage); + GGML_API void ggml_backend_meta_buffer_set_usage (ggml_backend_buffer_t buffer, enum ggml_backend_buffer_usage usage); // // Backend (meta) diff --git a/ggml/src/ggml-backend-meta.cpp b/ggml/src/ggml-backend-meta.cpp index fe58ea3bb7a..3ec40fb1af7 100644 --- a/ggml/src/ggml-backend-meta.cpp +++ b/ggml/src/ggml-backend-meta.cpp @@ -1168,7 +1168,6 @@ static struct ggml_backend_meta_split_state ggml_backend_meta_get_split_state( } static struct ggml_backend_meta_split_state ggml_backend_meta_get_split_state(const struct ggml_tensor * tensor, bool assume_sync) { - GGML_ASSERT(ggml_backend_buffer_is_meta(tensor->buffer)); ggml_backend_meta_buffer_context * buf_ctx = (ggml_backend_meta_buffer_context *) tensor->buffer->context; return ggml_backend_meta_get_split_state(buf_ctx->get_simple_tensor_container(tensor), tensor, assume_sync); } @@ -1259,7 +1258,14 @@ static enum ggml_status ggml_backend_meta_buffer_init_tensor_impl(ggml_backend_m t_ij->data = (char *) ggml_backend_buffer_get_base(simple_buf) + size_t(tensor->data) - size_t(ggml_backend_buffer_get_base(tensor->buffer)); } - t_ij->extra = tensor->extra; + + if (simple_buf) { + // the backend that owns the buffer will set .extra + ggml_backend_buffer_init_tensor(simple_buf, t_ij); + } else { + t_ij->extra = tensor->extra; + } + for (int i = 0; i < GGML_MAX_SRC; i++) { t_ij->src[i] = tensor->src[i]; if (tensor->src[i] == tensor) { @@ -1668,6 +1674,16 @@ bool ggml_backend_buffer_is_meta(ggml_backend_buffer_t buf) { return buf != nullptr && buf->iface.free_buffer == ggml_backend_meta_buffer_iface.free_buffer; } +void ggml_backend_meta_buffer_set_usage(ggml_backend_buffer_t buffer, enum ggml_backend_buffer_usage usage) { + GGML_ASSERT(ggml_backend_buffer_is_meta(buffer)); + ggml_backend_meta_buffer_context * buf_ctx = (ggml_backend_meta_buffer_context *) buffer->context; + for (size_t i = 0; i < buf_ctx->bufs.size(); i++) { + if (buf_ctx->bufs[i]) { + ggml_backend_buffer_set_usage(buf_ctx->bufs[i].get(), usage); + } + } +} + static ggml_backend_buffer_t ggml_backend_meta_buffer_type_alloc_buffer(ggml_backend_buffer_type_t buft, size_t size) { const size_t n_simple_bufts = ggml_backend_meta_buft_n_bufts(buft); diff --git a/ggml/src/ggml-backend.cpp b/ggml/src/ggml-backend.cpp index 3d6310f3ffe..e519bdf50a1 100644 --- a/ggml/src/ggml-backend.cpp +++ b/ggml/src/ggml-backend.cpp @@ -182,6 +182,8 @@ void ggml_backend_buffer_set_usage(ggml_backend_buffer_t buffer, enum ggml_backe // FIXME: add a generic callback to the buffer interface if (ggml_backend_buffer_is_multi_buffer(buffer)) { ggml_backend_multi_buffer_set_usage(buffer, usage); + } else if (ggml_backend_buffer_is_meta(buffer)) { + ggml_backend_meta_buffer_set_usage(buffer, usage); } } From 8c0adb05feb14f2ac29a9b3d6856d3aff39a1596 Mon Sep 17 00:00:00 2001 From: Dominik Pantaleoni <95251853+dpantaleoni@users.noreply.github.com> Date: Wed, 26 Aug 2026 01:57:07 -0700 Subject: [PATCH 005/104] ggml-metal: add chunked SSD MMA for Mamba-2 prefill optimization (llama/26647) * metal: WIP chunked SSD SSM_SCAN kernels for multi-token prefill * metal: drop scalar SSD path; MMA + sequential tail * drop WIP ssm scan test noise * remove state_from_dst and rename CS and NSG constants * remove unrelated added whitespace padding * added clarity to mma_tokens calculation * added clarity to use_mma bool checks * added comments to metal ssd op constants for clarity * reserve K tokens for sequential kernel rollback snapshots * reset concurrency between mma and seq tail * remove print args no longer used * fixed comment to no longer point to specific line * add FC_SSM_SCAN so seq path skips token offlset unless it's mma tail * added changes to new ssm.metal for rebase after ggml-metal.metal refactor * specialize ssm_scan tail with a template instead of a function constant --------- Co-authored-by: dpantaleoni Co-authored-by: forforever73 <690105611@qq.com> --- ggml/src/ggml-metal/ggml-metal-device.cpp | 25 ++- ggml/src/ggml-metal/ggml-metal-device.h | 3 +- ggml/src/ggml-metal/ggml-metal-impl.h | 6 + ggml/src/ggml-metal/ggml-metal-ops.cpp | 60 +++++-- ggml/src/ggml-metal/kernels/ssm.metal | 200 +++++++++++++++++++++- 5 files changed, 269 insertions(+), 25 deletions(-) diff --git a/ggml/src/ggml-metal/ggml-metal-device.cpp b/ggml/src/ggml-metal/ggml-metal-device.cpp index 4036cc21daa..a82caa5e430 100644 --- a/ggml/src/ggml-metal/ggml-metal-device.cpp +++ b/ggml/src/ggml-metal/ggml-metal-device.cpp @@ -572,7 +572,7 @@ ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_ssm_conv_batched return res; } -ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_ssm_scan(ggml_metal_library_t lib, const ggml_tensor * op) { +ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_ssm_scan(ggml_metal_library_t lib, const ggml_tensor * op, bool tail) { GGML_TENSOR_LOCALS( int32_t, ne0, op->src[0], ne); char base[256]; @@ -580,7 +580,7 @@ ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_ssm_scan(ggml_me const int nsg = (ne00 + 31)/32; - snprintf(base, 256, "kernel_ssm_scan_%s", ggml_type_name(op->src[0]->type)); + snprintf(base, 256, "kernel_ssm_scan_%s%s", ggml_type_name(op->src[0]->type), tail ? "_tail" : ""); snprintf(name, 256, "%s_nsg=%d", base, nsg); ggml_metal_pipeline_with_params res = ggml_metal_library_get_pipeline(lib, name); @@ -598,6 +598,27 @@ ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_ssm_scan(ggml_me return res; } +ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_ssm_scan_ssd_mma(ggml_metal_library_t lib, const ggml_tensor * op) { + char base[256]; + char name[256]; + + snprintf(base, 256, "kernel_ssm_scan_ssd_mma_%s", ggml_type_name(op->src[0]->type)); + snprintf(name, 256, "%s", base); + + 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); + } + + // acs/exp(acs)/state-decay vectors + dtX + SAM rows + two 8x8 tiles per simdgroup + res.smem = (3*OP_SSM_SCAN_SSD_CS + + OP_SSM_SCAN_SSD_CS*OP_SSM_SCAN_SSD_HD + + OP_SSM_SCAN_SSD_NSG*8*OP_SSM_SCAN_SSD_CS + + OP_SSM_SCAN_SSD_NSG*2*8*8)*sizeof(float); + + return res; +} + ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_rwkv(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 6c39428c777..003b688dbac 100644 --- a/ggml/src/ggml-metal/ggml-metal-device.h +++ b/ggml/src/ggml-metal/ggml-metal-device.h @@ -129,7 +129,8 @@ struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_lightning struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_dsv4_hc (ggml_metal_library_t lib, enum ggml_op 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_scan (ggml_metal_library_t lib, const struct ggml_tensor * op); +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); 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); diff --git a/ggml/src/ggml-metal/ggml-metal-impl.h b/ggml/src/ggml-metal/ggml-metal-impl.h index f0b7799791e..9becf04797b 100644 --- a/ggml/src/ggml-metal/ggml-metal-impl.h +++ b/ggml/src/ggml-metal/ggml-metal-impl.h @@ -158,6 +158,10 @@ #define OP_SUM_ROWS_NUM_SUM_ROWS 10 #define OP_SUM_ROWS_NUM_MEAN 11 +#define OP_SSM_SCAN_SSD_CS 64 // Metal-specific; Chunk Size; 64 is largest multiple of 8 (simdgroup tile) fitting into 32 KiB Metal threadgroup mem limit (~26.75 KiB shared mem; see smem layout comment in kernel_ssm_scan_ssd_mma_f32) +#define OP_SSM_SCAN_SSD_HD 64 // Metal-specific; Head Dim the MMA kernel is specialized for (Mamba-2); use_mma gates on d_inner == this +#define OP_SSM_SCAN_SSD_NSG 4 // Metal-specific; Number of SimdGroups per threadgroup; NSG*32 == threads dispatched per threadgroup + // kernel argument structs // // - element counters (e.g. ne00) typically use int32_t to reduce register usage @@ -893,6 +897,8 @@ typedef struct { int64_t n_head; int64_t n_group; int64_t n_seq_tokens; + int64_t n_seq_tokens_total; + int64_t token_offset; int64_t n_seqs; int64_t K; uint64_t s_off; diff --git a/ggml/src/ggml-metal/ggml-metal-ops.cpp b/ggml/src/ggml-metal/ggml-metal-ops.cpp index 1c3bb936b90..75de0f6dd08 100644 --- a/ggml/src/ggml-metal/ggml-metal-ops.cpp +++ b/ggml/src/ggml-metal/ggml-metal-ops.cpp @@ -1677,6 +1677,7 @@ int ggml_metal_op_ssm_scan(ggml_metal_op_t ctx, int idx) { ggml_metal_library_t lib = ctx->lib; ggml_metal_encoder_t enc = ctx->enc; + const ggml_metal_device_props * props_dev = ggml_metal_device_get_props(ctx->dev); GGML_TENSOR_LOCALS( int32_t, ne0, op->src[0], ne); GGML_TENSOR_LOCALS(uint64_t, nb0, op->src[0], nb); @@ -1722,6 +1723,8 @@ int ggml_metal_op_ssm_scan(ggml_metal_op_t ctx, int idx) { /*.n_head =*/ n_head, /*.n_group =*/ n_group, /*.n_seq_tokens =*/ n_seq_tokens, + /*.n_seq_tokens_total =*/ n_seq_tokens, + /*.token_offset =*/ 0, /*.n_seqs =*/ n_seqs, /*.K =*/ K, /*.s_off =*/ ggml_nelements(op->src[1]) * sizeof(float), @@ -1751,26 +1754,53 @@ int ggml_metal_op_ssm_scan(ggml_metal_op_t ctx, int idx) { /*.nb0 =*/ nb0, }; - auto pipeline = ggml_metal_library_get_pipeline_ssm_scan(lib, op); + constexpr int64_t CHUNK = OP_SSM_SCAN_SSD_CS; - GGML_ASSERT(d_state <= ggml_metal_pipeline_max_theads_per_threadgroup(pipeline)); + const int64_t snap_reserve = K > 1 ? K : 0; // tokens reserved for sequential kernel rollback snapshots + const int64_t mma_tokens = ((n_seq_tokens - snap_reserve) / CHUNK) * CHUNK; // largest multiple of CHUNK that leaves snap_reserve for the tail + const bool use_mma = + mma_tokens > 0 && + ne30 == 1 && // checks that A tensor is set to scalar decay per head (A shape {1, n_head}) + props_dev->has_simdgroup_mm && // hardware check for M1 or newer + d_state % 8 == 0 && // d_state must be multiple of 8 to align with simdgroup_float 8x8 tiles + d_inner == OP_SSM_SCAN_SSD_HD; // mma kernel is specialized for the Mamba-2 head dim; this checks it - const size_t smem = pipeline.smem; + const auto dispatch = [&](ggml_metal_pipeline_with_params pipeline, int64_t nth, int64_t n_tg_x) { + GGML_ASSERT(nth <= ggml_metal_pipeline_max_theads_per_threadgroup(pipeline)); + GGML_ASSERT(pipeline.smem <= props_dev->max_theadgroup_memory_size); - 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->src[2]), 3); - ggml_metal_encoder_set_buffer (enc, ggml_metal_get_buffer_id(op->src[3]), 4); - ggml_metal_encoder_set_buffer (enc, ggml_metal_get_buffer_id(op->src[4]), 5); - ggml_metal_encoder_set_buffer (enc, ggml_metal_get_buffer_id(op->src[5]), 6); - ggml_metal_encoder_set_buffer (enc, ggml_metal_get_buffer_id(op->src[6]), 7); - ggml_metal_encoder_set_buffer (enc, ggml_metal_get_buffer_id(op), 8); + 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->src[2]), 3); + ggml_metal_encoder_set_buffer (enc, ggml_metal_get_buffer_id(op->src[3]), 4); + ggml_metal_encoder_set_buffer (enc, ggml_metal_get_buffer_id(op->src[4]), 5); + ggml_metal_encoder_set_buffer (enc, ggml_metal_get_buffer_id(op->src[5]), 6); + ggml_metal_encoder_set_buffer (enc, ggml_metal_get_buffer_id(op->src[6]), 7); + ggml_metal_encoder_set_buffer (enc, ggml_metal_get_buffer_id(op), 8); + ggml_metal_encoder_set_threadgroup_memory_size(enc, pipeline.smem, 0); + ggml_metal_encoder_dispatch_threadgroups(enc, n_tg_x, n_head, n_seqs, nth, 1, 1); + }; - ggml_metal_encoder_set_threadgroup_memory_size(enc, smem, 0); + if (!use_mma) { + dispatch(ggml_metal_library_get_pipeline_ssm_scan(lib, op, false), d_state, d_inner); + return 1; + } + + args.n_seq_tokens = mma_tokens; + dispatch( + ggml_metal_library_get_pipeline_ssm_scan_ssd_mma(lib, op), + OP_SSM_SCAN_SSD_NSG*32, + 1); - ggml_metal_encoder_dispatch_threadgroups(enc, d_inner, n_head, n_seqs, d_state, 1, 1); + if (mma_tokens < n_seq_tokens) { + ggml_metal_op_concurrency_reset(ctx); + + args.n_seq_tokens = n_seq_tokens - mma_tokens; + args.token_offset = mma_tokens; + dispatch(ggml_metal_library_get_pipeline_ssm_scan(lib, op, true), d_state, d_inner); + } return 1; } diff --git a/ggml/src/ggml-metal/kernels/ssm.metal b/ggml/src/ggml-metal/kernels/ssm.metal index be065c4fffc..d3118a831b9 100644 --- a/ggml/src/ggml-metal/kernels/ssm.metal +++ b/ggml/src/ggml-metal/kernels/ssm.metal @@ -159,7 +159,9 @@ kernel void kernel_ssm_conv_f32_f32_batched_4( // ref: ggml.c:ggml_compute_forward_ssm_scan_f32, Mamba-2 part // Optimized version: reduces redundant memory loads by having one thread load shared values -kernel void kernel_ssm_scan_f32( +// TAIL == false is the whole-sequence / decode path: token_offset folds away at compile time. +template +kernel void kernel_ssm_scan_impl( constant ggml_metal_kargs_ssm_scan & args, device const void * src0, device const void * src1, @@ -200,13 +202,17 @@ kernel void kernel_ssm_scan_f32( const int32_t n_t = args.n_seq_tokens; const int32_t n_s = args.n_seqs; const int32_t K = args.K; + const int32_t n_t_total = TAIL ? args.n_seq_tokens_total : n_t; + const int32_t t_off = TAIL ? args.token_offset : 0; const int32_t s_off = args.s_off; device const int32_t * ids = (device const int32_t *) src6; - device const float * s0_buff = (device const float *) ((device const char *) src0 + ir*args.nb02 + ids[i3]*args.nb03); device float * s_buff = (device float *) ((device char *) dst + ir*args.nb02 + i3*args.nb03 + s_off); + device const float * s0_buff = t_off != 0 ? + s_buff : + (device const float *) ((device const char *) src0 + ir*args.nb02 + ids[i3]*args.nb03); const int32_t i = i0 + i1*nc; const int32_t g = ir / (nh / ng); // repeat_interleave @@ -218,12 +224,12 @@ kernel void kernel_ssm_scan_f32( const float A0 = A[i0%args.ne30]; - device const float * x = (device const float *)((device const char *) src1 + i1*args.nb10 + ir*args.nb11 + i3*args.nb13); // {dim, nh, nt, ns} - device const float * dt = (device const float *)((device const char *) src2 + ir*args.nb20 + i3*args.nb22); // {nh, nt, ns} - device const float * B = (device const float *)((device const char *) src4 + g*args.nb41 + i3*args.nb43); // {d_state, ng, nt, ns} - device const float * C = (device const float *)((device const char *) src5 + g*args.nb51 + i3*args.nb53); // {d_state, ng, nt, ns} + device const float * x = (device const float *)((device const char *) src1 + i1*args.nb10 + ir*args.nb11 + t_off*args.nb12 + i3*args.nb13); // {dim, nh, nt, ns} + device const float * dt = (device const float *)((device const char *) src2 + ir*args.nb20 + t_off*args.nb21 + i3*args.nb22); // {nh, nt, ns} + device const float * B = (device const float *)((device const char *) src4 + g*args.nb41 + t_off*args.nb42 + i3*args.nb43); // {d_state, ng, nt, ns} + device const float * C = (device const float *)((device const char *) src5 + g*args.nb51 + t_off*args.nb52 + i3*args.nb53); // {d_state, ng, nt, ns} - device float * y = dst + (i1 + ir*(nr) + i3*(n_t*nh*nr)); // {dim, nh, nt, ns} + device float * y = dst + (i1 + ir*nr + t_off*nh*nr + i3*(n_t_total*nh*nr)); // {dim, nh, nt, ns} for (int i2 = 0; i2 < n_t; i2 += sgptg) { threadgroup_barrier(mem_flags::mem_threadgroup); @@ -285,3 +291,183 @@ kernel void kernel_ssm_scan_f32( s_buff[i] = s; } + +typedef decltype(kernel_ssm_scan_impl) kernel_ssm_scan_t; + +template [[host_name("kernel_ssm_scan_f32")]] kernel kernel_ssm_scan_t kernel_ssm_scan_impl; +template [[host_name("kernel_ssm_scan_f32_tail")]] kernel kernel_ssm_scan_t kernel_ssm_scan_impl; + +// Chunked SSD SSM scan via Metal simdgroup MMatrix Multiply-Accumulate (simdgroup_float8x8) fast path. +// One threadgroup per (head, sequence) and tokens are processed in chunks. +// C*B^T computed in each chunk one time and reused across the head_dim channel tiles. +kernel void kernel_ssm_scan_ssd_mma_f32( + constant ggml_metal_kargs_ssm_scan & args, + device const void * src0, + device const void * src1, + device const void * src2, + device const void * src3, + device const void * src4, + device const void * src5, + device const void * src6, + device float * dst, + threadgroup float * shared [[threadgroup(0)]], + uint3 tgpig[[threadgroup_position_in_grid]], + ushort tiitg[[thread_index_in_threadgroup]], + ushort sgitg[[simdgroup_index_in_threadgroup]], + ushort tiisg[[thread_index_in_simdgroup]]) { + constexpr short CS = OP_SSM_SCAN_SSD_CS; + constexpr short TC = 8; // Tile Count of each edge in a simdgroup 8x8 tile + constexpr short HD = OP_SSM_SCAN_SSD_HD; + constexpr short NSG = OP_SSM_SCAN_SSD_NSG; + + // acs/exp(acs)/state-decay vectors, dtX[CS][HD], four private SAM row tiles [8][CS], + // and two 8x8 scratch tiles per simdgroup. Total: 26.75 KiB. + threadgroup float * shared_acs = shared; + threadgroup float * shared_exp_acs = shared + CS; + threadgroup float * shared_state_decay = shared + 2*CS; + threadgroup float * shared_dtx = shared + 3*CS; + threadgroup float * shared_sam = shared + 3*CS + CS*HD; + threadgroup float * sam_rows = shared_sam + sgitg*TC*CS; + threadgroup float * shared_tile = shared_sam + NSG*TC*CS; + threadgroup float * tile0 = shared_tile + sgitg*2*TC*TC; + threadgroup float * tile1 = tile0 + TC*TC; + + const int32_t ir = tgpig.y; // current head + const int32_t i3 = tgpig.z; // current seq + + const int32_t nc = args.d_state; + const int32_t nr = args.d_inner; + const int32_t nh = args.n_head; + const int32_t ng = args.n_group; + const int32_t n_t = args.n_seq_tokens; + const int32_t n_t_total = args.n_seq_tokens_total; + const int32_t g = ir / (nh / ng); + + device const int32_t * ids = (device const int32_t *) src6; + + device const float * s0_buff = (device const float *) ((device const char *) src0 + ir*args.nb02 + ids[i3]*args.nb03); + device float * s_buff = (device float *) ((device char *) dst + ir*args.nb02 + i3*args.nb03 + args.s_off); + + device const float * A = (device const float *) ((device const char *) src3 + ir*args.nb31); + device const float * x = (device const float *) ((device const char *) src1 + ir*args.nb11 + i3*args.nb13); + device const float * dt = (device const float *) ((device const char *) src2 + ir*args.nb20 + i3*args.nb22); + device const float * B = (device const float *) ((device const char *) src4 + g*args.nb41 + i3*args.nb43); + device const float * C = (device const float *) ((device const char *) src5 + g*args.nb51 + i3*args.nb53); + + device float * y = dst + (ir*nr + i3*(n_t_total*nh*nr)); + + for (int32_t t0 = 0; t0 < n_t; t0 += CS) { + for (int32_t idx = tiitg; idx < CS*HD; idx += NSG*N_SIMDWIDTH) { + const int32_t t = idx / HD; + const int32_t c = idx % HD; + const float dt0 = dt[(t0 + t) * (int32_t) args.ns21]; + const float dtsp = dt0 <= 20.0f ? log(1.0f + exp(dt0)) : dt0; + shared_dtx[idx] = x[(t0 + t) * (int32_t) args.ns12 + c] * dtsp; + } + if (tiitg < CS) { + const float dt0 = dt[(t0 + tiitg) * (int32_t) args.ns21]; + const float dtsp = dt0 <= 20.0f ? log(1.0f + exp(dt0)) : dt0; + shared_acs[tiitg] = dtsp * A[0]; + } + threadgroup_barrier(mem_flags::mem_threadgroup); + + if (tiitg == 0) { + float acc = 0.0f; + for (short t = 0; t < CS; ++t) { + acc += shared_acs[t]; + shared_acs[t] = acc; + } + } + threadgroup_barrier(mem_flags::mem_threadgroup); + if (tiitg < CS) { + shared_exp_acs[tiitg] = exp(shared_acs[tiitg]); + shared_state_decay[tiitg] = exp(shared_acs[CS - 1] - shared_acs[tiitg]); + } + threadgroup_barrier(mem_flags::mem_threadgroup); + + device const float * state = t0 == 0 ? s0_buff : s_buff; + + // Build one 8x64 row tile of SAM per simdgroup, then reuse it across every channel tile. + for (short ib = sgitg; ib < CS/TC; ib += NSG) { + for (short jb = 0; jb <= ib; ++jb) { + simdgroup_float8x8 cb = make_filled_simdgroup_matrix(0.0f); + + for (int32_t k0 = 0; k0 < nc; k0 += TC) { + simdgroup_float8x8 mc; + simdgroup_float8x8 mb; + simdgroup_load(mc, C + (t0 + ib*TC)*(int32_t) args.ns52 + k0, args.ns52); + simdgroup_load(mb, B + (t0 + jb*TC)*(int32_t) args.ns42 + k0, args.ns42, 0, true); + simdgroup_multiply_accumulate(cb, mc, mb, cb); + } + + threadgroup float * sam = sam_rows + jb*TC; + simdgroup_store(cb, sam, CS); + simdgroup_barrier(mem_flags::mem_threadgroup); + for (short e = tiisg; e < TC*TC; e += N_SIMDWIDTH) { + const short ri = e / TC; + const short rj = e % TC; + const short i = ib*TC + ri; + const short j = jb*TC + rj; + sam[ri*CS + rj] = j <= i ? + sam[ri*CS + rj] * exp(shared_acs[i] - shared_acs[j]) : 0.0f; + } + simdgroup_barrier(mem_flags::mem_threadgroup); + } + + for (short ch = 0; ch < HD/TC; ++ch) { + simdgroup_float8x8 y_diag = make_filled_simdgroup_matrix(0.0f); + simdgroup_float8x8 y_inter = make_filled_simdgroup_matrix(0.0f); + + for (short jb = 0; jb <= ib; ++jb) { + simdgroup_float8x8 sam; + simdgroup_float8x8 mdtx; + simdgroup_load(sam, sam_rows + jb*TC, CS); + simdgroup_load(mdtx, shared_dtx + jb*TC*HD + ch*TC, HD); + simdgroup_multiply_accumulate(y_diag, sam, mdtx, y_diag); + } + + for (int32_t k0 = 0; k0 < nc; k0 += TC) { + simdgroup_float8x8 mc; + simdgroup_float8x8 ms; + simdgroup_load(mc, C + (t0 + ib*TC)*(int32_t) args.ns52 + k0, args.ns52); + simdgroup_load(ms, state + ch*TC*nc + k0, nc, 0, true); + simdgroup_multiply_accumulate(y_inter, mc, ms, y_inter); + } + + simdgroup_store(y_diag, tile0, TC); + simdgroup_store(y_inter, tile1, TC); + simdgroup_barrier(mem_flags::mem_threadgroup); + for (short e = tiisg; e < TC*TC; e += N_SIMDWIDTH) { + const short ri = e / TC; + const short ci = e % TC; + const int32_t token = t0 + ib*TC + ri; + y[token*nh*nr + ch*TC + ci] = + tile0[e] + shared_exp_acs[ib*TC + ri] * tile1[e]; + } + simdgroup_barrier(mem_flags::mem_threadgroup); + } + } + + // All simdgroups must finish reading s_buff before any thread overwrites it. + threadgroup_barrier(mem_flags::mem_device | mem_flags::mem_threadgroup); + + // Keep the carried-state reduction in token order. Reassociating this particular product + // with MMA compounds rounding differences at every chunk boundary; CB, y_diag, and C*S + // remain on the matrix unit. + const float chunk_decay = exp(shared_acs[CS - 1]); + for (int32_t idx = tiitg; idx < nc*HD; idx += NSG*N_SIMDWIDTH) { + const int32_t ci = idx / nc; + const int32_t si = idx % nc; + float state_c = 0.0f; + for (short t = 0; t < CS; ++t) { + state_c += shared_state_decay[t] * + B[(t0 + t)*(int32_t) args.ns42 + si] * + shared_dtx[t*HD + ci]; + } + s_buff[idx] = chunk_decay * state[idx] + state_c; + } + + // All state tiles must be visible before the next chunk consumes s_buff as S_prev. + threadgroup_barrier(mem_flags::mem_device | mem_flags::mem_threadgroup); + } +} From 9d8e6b91b18a15728c3c57b8eef30e16a413112f Mon Sep 17 00:00:00 2001 From: David Friehs Date: Sun, 30 Aug 2026 13:01:24 +0300 Subject: [PATCH 006/104] cuda: unblock mmq for MoE on sm_60 (llama/26264) --- ...-pascal.cuh => mmq-config-pascal-dp4a.cuh} | 2 +- .../src/ggml-cuda/mmq-config-pascal-older.cuh | 273 ++++++++++++++++++ ggml/src/ggml-cuda/mmq.cu | 4 +- ggml/src/ggml-cuda/mmq.cuh | 12 +- 4 files changed, 286 insertions(+), 5 deletions(-) rename ggml/src/ggml-cuda/{mmq-config-pascal.cuh => mmq-config-pascal-dp4a.cuh} (99%) create mode 100644 ggml/src/ggml-cuda/mmq-config-pascal-older.cuh diff --git a/ggml/src/ggml-cuda/mmq-config-pascal.cuh b/ggml/src/ggml-cuda/mmq-config-pascal-dp4a.cuh similarity index 99% rename from ggml/src/ggml-cuda/mmq-config-pascal.cuh rename to ggml/src/ggml-cuda/mmq-config-pascal-dp4a.cuh index e7d4a9a3fcb..83eb7c146e1 100644 --- a/ggml/src/ggml-cuda/mmq-config-pascal.cuh +++ b/ggml/src/ggml-cuda/mmq-config-pascal-dp4a.cuh @@ -1,4 +1,4 @@ -static constexpr __host__ __device__ ggml_cuda_mmq_config ggml_cuda_mmq_get_config_pascal(ggml_type type, int J, bool fallback) { +static constexpr __host__ __device__ ggml_cuda_mmq_config ggml_cuda_mmq_get_config_pascal_dp4a(ggml_type type, int J, bool fallback) { CASE(GGML_TYPE_Q1_0, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); CASE(GGML_TYPE_Q1_0, 256, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); CASE(GGML_TYPE_Q1_0, 256, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); diff --git a/ggml/src/ggml-cuda/mmq-config-pascal-older.cuh b/ggml/src/ggml-cuda/mmq-config-pascal-older.cuh new file mode 100644 index 00000000000..2a8dc9e1a93 --- /dev/null +++ b/ggml/src/ggml-cuda/mmq-config-pascal-older.cuh @@ -0,0 +1,273 @@ +static constexpr __host__ __device__ ggml_cuda_mmq_config ggml_cuda_mmq_get_config_pascal_older(ggml_type type, int J, bool fallback) { + CASE(GGML_TYPE_Q1_0, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q1_0, 256, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q1_0, 256, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q1_0, 256, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q1_0, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q1_0, 256, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q1_0, 256, 2, 64, 24, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q1_0, 256, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q1_0, 256, 2, 64, 40, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q1_0, 256, 2, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q1_0, 256, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + + CASE(GGML_TYPE_Q2_0, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q2_0, 256, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q2_0, 256, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q2_0, 256, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q2_0, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q2_0, 256, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q2_0, 256, 2, 64, 24, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q2_0, 256, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q2_0, 256, 2, 64, 40, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q2_0, 256, 2, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q2_0, 256, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + + CASE(GGML_TYPE_Q4_0, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q4_0, 256, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q4_0, 256, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q4_0, 256, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q4_0, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q4_0, 256, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q4_0, 256, 2, 64, 24, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q4_0, 256, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q4_0, 256, 2, 64, 40, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q4_0, 256, 2, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q4_0, 256, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + + CASE(GGML_TYPE_Q4_1, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q4_1, 256, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q4_1, 256, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q4_1, 256, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q4_1, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q4_1, 256, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q4_1, 256, 2, 64, 24, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q4_1, 256, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q4_1, 256, 2, 64, 40, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q4_1, 256, 2, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q4_1, 256, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); + + CASE(GGML_TYPE_Q5_0, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q5_0, 256, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q5_0, 256, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q5_0, 256, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q5_0, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q5_0, 256, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q5_0, 256, 2, 64, 24, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q5_0, 256, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q5_0, 256, 2, 64, 40, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q5_0, 256, 2, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q5_0, 256, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + + CASE(GGML_TYPE_Q5_1, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q5_1, 256, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q5_1, 256, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q5_1, 256, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q5_1, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q5_1, 256, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q5_1, 256, 2, 64, 24, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q5_1, 256, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q5_1, 256, 2, 64, 40, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q5_1, 256, 2, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q5_1, 256, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); + + CASE(GGML_TYPE_Q8_0, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q8_0, 256, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q8_0, 256, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q8_0, 256, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q8_0, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q8_0, 256, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q8_0, 256, 2, 64, 24, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q8_0, 256, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q8_0, 256, 2, 64, 40, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q8_0, 256, 2, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q8_0, 256, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + +// --------------------------------------------------------------------------------------------- + + CASE(GGML_TYPE_Q2_K, 256, 1, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q2_K, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q2_K, 256, 1, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q2_K, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q2_K, 256, 1, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q2_K, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q2_K, 256, 1, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q2_K, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q2_K, 256, 1, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q2_K, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q2_K, 256, 1, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q2_K, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q2_K, 256, 1, 64, 24, GGML_CUDA_MMQ_SRAM_LAYOUT_Q2_K, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q2_K, 256, 1, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q2_K, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q2_K, 256, 1, 64, 40, GGML_CUDA_MMQ_SRAM_LAYOUT_Q2_K, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q2_K, 256, 1, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q2_K, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q2_K, 256, 1, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q2_K, MMQ_ITER_K, false, false); + + CASE(GGML_TYPE_Q3_K, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q3_K, 256, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q3_K, 256, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q3_K, 256, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q3_K, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q3_K, 256, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q3_K, 256, 2, 64, 24, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q3_K, 256, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q3_K, 256, 2, 64, 40, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q3_K, 256, 2, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q3_K, 256, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); + + CASE(GGML_TYPE_Q4_K, 256, 1, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q4_K, 256, 1, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q4_K, 256, 1, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q4_K, 256, 1, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q4_K, 256, 1, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q4_K, 256, 1, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q4_K, 256, 1, 64, 24, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q4_K, 256, 1, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q4_K, 256, 1, 64, 40, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q4_K, 256, 1, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q4_K, 256, 1, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); + + CASE(GGML_TYPE_Q5_K, 256, 1, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q5_K, 256, 1, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q5_K, 256, 1, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q5_K, 256, 1, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q5_K, 256, 1, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q5_K, 256, 1, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q5_K, 256, 1, 64, 24, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q5_K, 256, 1, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q5_K, 256, 1, 64, 40, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q5_K, 256, 1, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q5_K, 256, 1, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); + + CASE(GGML_TYPE_Q6_K, 256, 1, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q6_K, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q6_K, 256, 1, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q6_K, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q6_K, 256, 1, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q6_K, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q6_K, 256, 1, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q6_K, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q6_K, 256, 1, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q6_K, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q6_K, 256, 1, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q6_K, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q6_K, 256, 1, 64, 24, GGML_CUDA_MMQ_SRAM_LAYOUT_Q6_K, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q6_K, 256, 1, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q6_K, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q6_K, 256, 1, 64, 40, GGML_CUDA_MMQ_SRAM_LAYOUT_Q6_K, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q6_K, 256, 1, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q6_K, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q6_K, 256, 1, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q6_K, MMQ_ITER_K, false, false); + +// --------------------------------------------------------------------------------------------- + + CASE(GGML_TYPE_IQ1_S, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_IQ1_S, 256, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_IQ1_S, 256, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_IQ1_S, 256, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_IQ1_S, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ1_S, 256, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ1_S, 256, 2, 64, 24, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ1_S, 256, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ1_S, 256, 2, 64, 40, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ1_S, 256, 2, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ1_S, 256, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + + CASE(GGML_TYPE_IQ2_XXS, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_IQ2_XXS, 256, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_IQ2_XXS, 256, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_IQ2_XXS, 256, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_IQ2_XXS, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ2_XXS, 256, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ2_XXS, 256, 2, 64, 24, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ2_XXS, 256, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ2_XXS, 256, 2, 64, 40, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ2_XXS, 256, 2, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ2_XXS, 256, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + + CASE(GGML_TYPE_IQ2_XS, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_IQ2_XS, 256, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_IQ2_XS, 256, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_IQ2_XS, 256, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_IQ2_XS, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ2_XS, 256, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ2_XS, 256, 2, 64, 24, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ2_XS, 256, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ2_XS, 256, 2, 64, 40, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ2_XS, 256, 2, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ2_XS, 256, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); + + CASE(GGML_TYPE_IQ2_S, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_IQ2_S, 256, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_IQ2_S, 256, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_IQ2_S, 256, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_IQ2_S, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ2_S, 256, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ2_S, 256, 2, 64, 24, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ2_S, 256, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ2_S, 256, 2, 64, 40, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ2_S, 256, 2, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ2_S, 256, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); + + CASE(GGML_TYPE_IQ3_XXS, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_IQ3_XXS, 256, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_IQ3_XXS, 256, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_IQ3_XXS, 256, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_IQ3_XXS, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ3_XXS, 256, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ3_XXS, 256, 2, 64, 24, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ3_XXS, 256, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ3_XXS, 256, 2, 64, 40, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ3_XXS, 256, 2, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ3_XXS, 256, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + + CASE(GGML_TYPE_IQ3_S, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_IQ3_S, 256, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_IQ3_S, 256, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_IQ3_S, 256, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_IQ3_S, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ3_S, 256, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ3_S, 256, 2, 64, 24, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ3_S, 256, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ3_S, 256, 2, 64, 40, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ3_S, 256, 2, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ3_S, 256, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + + CASE(GGML_TYPE_IQ4_XS, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_IQ4_XS, 256, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_IQ4_XS, 256, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_IQ4_XS, 256, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_IQ4_XS, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ4_XS, 256, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ4_XS, 256, 2, 64, 24, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ4_XS, 256, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ4_XS, 256, 2, 64, 40, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ4_XS, 256, 2, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ4_XS, 256, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + + CASE(GGML_TYPE_IQ4_NL, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_IQ4_NL, 256, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_IQ4_NL, 256, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_IQ4_NL, 256, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_IQ4_NL, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ4_NL, 256, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ4_NL, 256, 2, 64, 24, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ4_NL, 256, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ4_NL, 256, 2, 64, 40, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ4_NL, 256, 2, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ4_NL, 256, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + +// --------------------------------------------------------------------------------------------- + + CASE(GGML_TYPE_MXFP4, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_MXFP4, 256, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_MXFP4, 256, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_MXFP4, 256, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_MXFP4, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_MXFP4, 256, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_MXFP4, 256, 2, 64, 24, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_MXFP4, 256, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_MXFP4, 256, 2, 64, 40, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_MXFP4, 256, 2, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_MXFP4, 256, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); + + CASE(GGML_TYPE_NVFP4, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_NVFP4, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_NVFP4, 256, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_NVFP4, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_NVFP4, 256, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_NVFP4, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_NVFP4, 256, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_NVFP4, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_NVFP4, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_NVFP4, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_NVFP4, 256, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_NVFP4, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_NVFP4, 256, 2, 64, 24, GGML_CUDA_MMQ_SRAM_LAYOUT_NVFP4, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_NVFP4, 256, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_NVFP4, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_NVFP4, 256, 2, 64, 40, GGML_CUDA_MMQ_SRAM_LAYOUT_NVFP4, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_NVFP4, 256, 2, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_NVFP4, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_NVFP4, 256, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_NVFP4, MMQ_ITER_K, false, false); + + return ggml_cuda_mmq_config(GGML_TYPE_COUNT, 256, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, 256, false, true); +} diff --git a/ggml/src/ggml-cuda/mmq.cu b/ggml/src/ggml-cuda/mmq.cu index 707437ea3e5..7fb4401489c 100644 --- a/ggml/src/ggml-cuda/mmq.cu +++ b/ggml/src/ggml-cuda/mmq.cu @@ -314,7 +314,9 @@ bool ggml_cuda_should_use_mmq(enum ggml_type type, int cc, int64_t ne11, int64_t } if (ggml_cuda_highest_compiled_arch(cc) < GGML_CUDA_CC_DP4A) { - return false; + // for MoE, mmq is faster even without native dp4a + // TODO: check if cards older than pascal might benefit from this as well + return cc >= GGML_CUDA_CC_PASCAL && n_experts > 0; } #ifdef GGML_CUDA_FORCE_MMQ diff --git a/ggml/src/ggml-cuda/mmq.cuh b/ggml/src/ggml-cuda/mmq.cuh index 2eb15fdfad9..c978b4421c5 100644 --- a/ggml/src/ggml-cuda/mmq.cuh +++ b/ggml/src/ggml-cuda/mmq.cuh @@ -213,7 +213,8 @@ struct ggml_cuda_mmq_config { return ggml_cuda_mmq_config((type_), (nthreads_), (occupancy_), (I_), (J_), (sram_layout_), (K_vram_), (stream_k_), (fallback_)); \ } \ -#include "mmq-config-pascal.cuh" +#include "mmq-config-pascal-older.cuh" +#include "mmq-config-pascal-dp4a.cuh" #include "mmq-config-ampere.cuh" #include "mmq-config-blackwell.cuh" @@ -247,7 +248,10 @@ static __host__ ggml_cuda_mmq_config ggml_cuda_mmq_get_config(const ggml_type ty if (ggml_cuda_highest_compiled_arch(cc) >= GGML_CUDA_CC_VOLTA) { return ggml_cuda_mmq_get_config_ampere(type, J, fallback); } - return ggml_cuda_mmq_get_config_pascal(type, J, fallback); + if (ggml_cuda_highest_compiled_arch(cc) >= GGML_CUDA_CC_DP4A) { + return ggml_cuda_mmq_get_config_pascal_dp4a(type, J, fallback); + } + return ggml_cuda_mmq_get_config_pascal_older(type, J, fallback); } static constexpr __device__ ggml_cuda_mmq_config ggml_cuda_mmq_get_config(ggml_type type, int J, bool fallback) { @@ -268,8 +272,10 @@ static constexpr __device__ ggml_cuda_mmq_config ggml_cuda_mmq_get_config(ggml_t return ggml_cuda_mmq_get_config_blackwell(type, J, fallback); #elif __CUDA_ARCH__ >= GGML_CUDA_CC_VOLTA return ggml_cuda_mmq_get_config_ampere(type, J, fallback); +#elif __CUDA_ARCH__ >= GGML_CUDA_CC_DP4A + return ggml_cuda_mmq_get_config_pascal_dp4a(type, J, fallback); #else - return ggml_cuda_mmq_get_config_pascal(type, J, fallback); + return ggml_cuda_mmq_get_config_pascal_older(type, J, fallback); #endif // BLACKWELL_MMA_AVAILABLE #endif // GGML_USE_HIP GGML_UNUSED_VARS(type, J, fallback); From 82f5f85e951855cdfa3f91578be6b1fe9a56c0bc Mon Sep 17 00:00:00 2001 From: Radoslav Gerganov Date: Wed, 26 Aug 2026 17:34:46 +0300 Subject: [PATCH 007/104] rpc : implement event and async backend APIs (llama/18626) * rpc : implement event and async backend APIs * cache responses from RPC_CMD_GET_ALLOC_SIZE --- ggml/include/ggml-rpc.h | 4 +- ggml/src/ggml-rpc/ggml-rpc.cpp | 603 ++++++++++++++++++++++++--------- 2 files changed, 442 insertions(+), 165 deletions(-) diff --git a/ggml/include/ggml-rpc.h b/ggml/include/ggml-rpc.h index 059e4496269..cbfe400139c 100644 --- a/ggml/include/ggml-rpc.h +++ b/ggml/include/ggml-rpc.h @@ -6,8 +6,8 @@ extern "C" { #endif -#define RPC_PROTO_MAJOR_VERSION 5 -#define RPC_PROTO_MINOR_VERSION 1 +#define RPC_PROTO_MAJOR_VERSION 6 +#define RPC_PROTO_MINOR_VERSION 0 #define RPC_PROTO_PATCH_VERSION 0 #ifdef __cplusplus diff --git a/ggml/src/ggml-rpc/ggml-rpc.cpp b/ggml/src/ggml-rpc/ggml-rpc.cpp index 69a8a08ae17..9aa5883d80d 100644 --- a/ggml/src/ggml-rpc/ggml-rpc.cpp +++ b/ggml/src/ggml-rpc/ggml-rpc.cpp @@ -9,6 +9,9 @@ #include #include #include +#include +#include +#include #include #include #include @@ -17,6 +20,8 @@ #include #include #include +#include +#include static const char * RPC_DEBUG = std::getenv("GGML_RPC_DEBUG"); @@ -72,6 +77,7 @@ enum rpc_cmd { RPC_CMD_DEVICE_COUNT, RPC_CMD_GRAPH_RECOMPUTE, RPC_CMD_MEMSET_TENSOR, + RPC_CMD_NONE, RPC_CMD_COUNT, }; @@ -223,24 +229,24 @@ struct ggml_backend_rpc_buffer_type_context { size_t max_size; }; +class rpc_dispatcher; struct ggml_backend_rpc_context { - std::string endpoint; - uint32_t device; - std::string name; + std::shared_ptr dispatcher; + uint32_t device; + std::string name; }; struct ggml_backend_rpc_buffer_context { - std::shared_ptr sock; - void * base_ptr; - uint64_t remote_ptr; + std::shared_ptr dispatcher; + void * base_ptr; + uint64_t remote_ptr; }; // RPC helper functions // Computes FNV-1a hash of the data -static uint64_t fnv_hash(const uint8_t * data, size_t len) { +static uint64_t fnv_hash(const uint8_t * data, size_t len, uint64_t hash = 0xcbf29ce484222325ULL) { const uint64_t fnv_prime = 0x100000001b3ULL; - uint64_t hash = 0xcbf29ce484222325ULL; for (size_t i = 0; i < len; ++i) { hash ^= data[i]; @@ -357,44 +363,248 @@ static bool negotiate_hello(const std::shared_ptr & sock) { return true; } -static std::shared_ptr get_socket(const std::string & endpoint) { - static std::mutex mutex; - std::lock_guard lock(mutex); - static std::unordered_map> sockets; +template +class message_queue { +public: + message_queue() {} - auto it = sockets.find(endpoint); - if (it != sockets.end()) { - if (auto sock = it->second.lock()) { - return sock; + bool push(const T &value) { + std::unique_lock lock(mutex); + if (interrupted) { + return false; } + queue.push(value); + cvar.notify_all(); + return true; } + + bool pop(T* out) { + std::unique_lock lock(mutex); + cvar.wait(lock, [this] { return !queue.empty() || interrupted; }); + if (interrupted) { + return false; + } + *out = queue.front(); + queue.pop(); + return true; + } + + void interrupt() { + std::unique_lock lock(mutex); + interrupted = true; + lock.unlock(); + cvar.notify_all(); + } + +private: + bool interrupted = false; + std::queue queue; + std::mutex mutex; + std::condition_variable cvar; +}; + +class rpc_dispatcher { +public: + rpc_dispatcher() { + } + + void send(enum rpc_cmd cmd, std::shared_ptr input, size_t input_size); + void send(enum rpc_cmd cmd, std::shared_ptr input, size_t input_size, void * output, size_t output_size); + void send_async(enum rpc_cmd cmd, std::shared_ptr input, size_t input_size); + void send_async(enum rpc_cmd cmd, std::shared_ptr input, size_t input_size, void * output, size_t output_size); + + ggml_backend_event_t event_new(ggml_backend_dev_t dev); + void event_free(ggml_backend_event_t event); + void event_synchronize(ggml_backend_event_t event); + void event_record(ggml_backend_event_t event); + void synchronize(); + + void start(const std::string & endpoint); + void work(); + + ~rpc_dispatcher(); + +private: + struct rpc_msg { + rpc_cmd cmd; + std::shared_ptr input; + size_t input_size; + void * output; + size_t output_size; + std::promise completion; + }; + using rpc_msg_ptr = std::shared_ptr; + using rpc_msg_queue = message_queue; + struct rpc_event { + rpc_msg_ptr msg; + std::shared_future sf; + }; + rpc_msg_queue queue; + socket_ptr sock; + std::atomic_bool running; + std::thread thread; +}; + +static void rpc_dispatcher_trampoline(rpc_dispatcher * dispatcher) +{ + dispatcher->work(); +} + +void rpc_dispatcher::send(enum rpc_cmd cmd, std::shared_ptr input, size_t input_size) { + auto msg = std::make_shared(); + msg->cmd = cmd; + msg->input = input; + msg->input_size = input_size; + msg->output = nullptr; + msg->output_size = 0; + GGML_ASSERT(queue.push(msg)); + auto future = msg->completion.get_future(); + future.wait(); +} + +void rpc_dispatcher::send_async(enum rpc_cmd cmd, std::shared_ptr input, size_t input_size) { + auto msg = std::make_shared(); + msg->cmd = cmd; + msg->input = input; + msg->input_size = input_size; + msg->output = nullptr; + msg->output_size = 0; + GGML_ASSERT(queue.push(msg)); +} + +void rpc_dispatcher::send(enum rpc_cmd cmd, std::shared_ptr input, size_t input_size, void * output, size_t output_size) { + auto msg = std::make_shared(); + msg->cmd = cmd; + msg->input = input; + msg->input_size = input_size; + msg->output = output; + msg->output_size = output_size; + GGML_ASSERT(queue.push(msg)); + auto future = msg->completion.get_future(); + future.wait(); +} + +void rpc_dispatcher::send_async(enum rpc_cmd cmd, std::shared_ptr input, size_t input_size, void * output, size_t output_size) { + auto msg = std::make_shared(); + msg->cmd = cmd; + msg->input = input; + msg->input_size = input_size; + msg->output = output; + msg->output_size = output_size; + GGML_ASSERT(queue.push(msg)); +} + +ggml_backend_event_t rpc_dispatcher::event_new(ggml_backend_dev_t dev) { + rpc_event * ev = new rpc_event; + ev->msg = std::make_shared(); + ev->msg->cmd = RPC_CMD_NONE; + ev->sf = ev->msg->completion.get_future().share(); + GGML_ASSERT(queue.push(ev->msg)); + return new ggml_backend_event { + /* .device = */ dev, + /* .context = */ ev, + }; +} + +void rpc_dispatcher::event_free(ggml_backend_event_t event) { + rpc_event * ev = (rpc_event *)event->context; + delete ev; +} + +void rpc_dispatcher::event_synchronize(ggml_backend_event_t event) { + rpc_event * ev = (rpc_event *)event->context; + ev->sf.wait(); +} + +void rpc_dispatcher::event_record(ggml_backend_event_t event) { + rpc_event * ev = (rpc_event *)event->context; + ev->msg = std::make_shared(); + ev->msg->cmd = RPC_CMD_NONE; + ev->sf = ev->msg->completion.get_future().share(); + GGML_ASSERT(queue.push(ev->msg)); +} + +void rpc_dispatcher::synchronize() { + // to ensure all messages are processed, submit dummy message and wait for it to complete + auto msg = std::make_shared(); + msg->cmd = RPC_CMD_NONE; + GGML_ASSERT(queue.push(msg)); + msg->completion.get_future().wait(); +} + +void rpc_dispatcher::start(const std::string & endpoint) { std::string host; int port; if (!parse_endpoint(endpoint, host, port)) { - GGML_LOG_ERROR("Failed to parse endpoint: %s\n", endpoint.c_str()); - return nullptr; + GGML_ABORT("Failed to parse endpoint: %s\n", endpoint.c_str()); } - if (!rpc_transport_init()) { - return nullptr; + GGML_ABORT("RPC transport initialization failed\n"); } - auto sock = socket_t::connect(host.c_str(), port); + + sock = socket_t::connect(host.c_str(), port); if (sock == nullptr) { - return nullptr; + GGML_ABORT("Failed to connect to %s\n", endpoint.c_str()); } if (!negotiate_hello(sock)) { - return nullptr; + GGML_ABORT("RPC handshake failed for %s\n", endpoint.c_str()); } LOG_DBG("[%s] connected to %s\n", __func__, endpoint.c_str()); - sockets[endpoint] = sock; - return sock; + running = true; + thread = std::thread(rpc_dispatcher_trampoline, this); +} + +void rpc_dispatcher::work() { + while (running) { + rpc_msg_ptr msg_ptr; + if (!queue.pop(&msg_ptr)) { + break; + } + if (msg_ptr->cmd != RPC_CMD_NONE) { + if (msg_ptr->output) { + bool status = send_rpc_cmd(sock, msg_ptr->cmd, msg_ptr->input.get(), msg_ptr->input_size, msg_ptr->output, msg_ptr->output_size); + RPC_STATUS_ASSERT(status); + } else { + bool status = send_rpc_cmd(sock, msg_ptr->cmd, msg_ptr->input.get(), msg_ptr->input_size); + RPC_STATUS_ASSERT(status); + } + } + msg_ptr->completion.set_value(); + } +} + +rpc_dispatcher::~rpc_dispatcher() { + running = false; + queue.interrupt(); + sock = nullptr; + if (thread.joinable()) { + thread.join(); + } +} + +static std::shared_ptr get_dispatcher(const std::string & endpoint) { + static std::mutex mutex; + std::lock_guard lock(mutex); + static std::unordered_map> dispatchers; + + auto it = dispatchers.find(endpoint); + if (it != dispatchers.end()) { + if (auto dispatcher = it->second.lock()) { + return dispatcher; + } + } + + auto dispatcher = std::make_shared(); + dispatcher->start(endpoint); + dispatchers[endpoint] = dispatcher; + return dispatcher; } static void ggml_backend_rpc_buffer_free_buffer(ggml_backend_buffer_t buffer) { ggml_backend_rpc_buffer_context * ctx = (ggml_backend_rpc_buffer_context *)buffer->context; - rpc_msg_free_buffer_req request = {ctx->remote_ptr}; - bool status = send_rpc_cmd(ctx->sock, RPC_CMD_FREE_BUFFER, &request, sizeof(request), nullptr, 0); - RPC_STATUS_ASSERT(status); + auto request = std::make_shared(); + request->remote_ptr = ctx->remote_ptr; + ctx->dispatcher->send(RPC_CMD_FREE_BUFFER, request, sizeof(*request)); delete ctx; } @@ -403,10 +613,10 @@ static void * ggml_backend_rpc_buffer_get_base(ggml_backend_buffer_t buffer) { if (ctx->base_ptr != nullptr) { return ctx->base_ptr; } - rpc_msg_buffer_get_base_req request = {ctx->remote_ptr}; + auto request = std::make_shared(); + request->remote_ptr = ctx->remote_ptr; rpc_msg_buffer_get_base_rsp response; - bool status = send_rpc_cmd(ctx->sock, RPC_CMD_BUFFER_GET_BASE, &request, sizeof(request), &response, sizeof(response)); - RPC_STATUS_ASSERT(status); + ctx->dispatcher->send(RPC_CMD_BUFFER_GET_BASE, request, sizeof(*request), &response, sizeof(response)); ctx->base_ptr = reinterpret_cast(response.base_ptr); return ctx->base_ptr; } @@ -463,12 +673,9 @@ static enum ggml_status ggml_backend_rpc_buffer_init_tensor(ggml_backend_buffer_ // Due to bandwidth constraints, we only call the server init tensor functions if necessary. // In particular, only quantized tensors need padding if (ggml_is_quantized(tensor->type) && (tensor->ne[0] % 512 != 0) && (tensor->view_src == nullptr)) { - rpc_msg_init_tensor_req request; - - request.tensor = serialize_tensor(tensor); - - bool status = send_rpc_cmd(ctx->sock, RPC_CMD_INIT_TENSOR, &request, sizeof(request), nullptr, 0); - RPC_STATUS_ASSERT(status); + auto request = std::make_shared(); + request->tensor = serialize_tensor(tensor); + ctx->dispatcher->send(RPC_CMD_INIT_TENSOR, request, sizeof(*request)); } return GGML_STATUS_SUCCESS; } @@ -476,27 +683,24 @@ static enum ggml_status ggml_backend_rpc_buffer_init_tensor(ggml_backend_buffer_ static void ggml_backend_rpc_buffer_memset_tensor( ggml_backend_buffer_t buffer, ggml_tensor * tensor, uint8_t value, size_t offset, size_t size) { ggml_backend_rpc_buffer_context * ctx = (ggml_backend_rpc_buffer_context *)buffer->context; - rpc_msg_memset_tensor_req request = { - /* .tensor = */ serialize_tensor(tensor), - /* .offset = */ offset, - /* .size = */ size, - /* .value = */ value, - }; - bool status = send_rpc_cmd(ctx->sock, RPC_CMD_MEMSET_TENSOR, &request, sizeof(request), nullptr, 0); - RPC_STATUS_ASSERT(status); + auto request = std::make_shared(); + request->tensor = serialize_tensor(tensor); + request->offset = offset; + request->size = size; + request->value = value; + ctx->dispatcher->send(RPC_CMD_MEMSET_TENSOR, request, sizeof(*request)); } static void ggml_backend_rpc_buffer_set_tensor(ggml_backend_buffer_t buffer, ggml_tensor * tensor, const void * data, size_t offset, size_t size) { ggml_backend_rpc_buffer_context * ctx = (ggml_backend_rpc_buffer_context *)buffer->context; rpc_tensor rpc_tensor = serialize_tensor(tensor); if (size > HASH_THRESHOLD) { - rpc_msg_set_tensor_hash_req request; - request.tensor = rpc_tensor; - request.offset = offset; - request.hash = fnv_hash((const uint8_t*)data, size); + auto request = std::make_shared(); + request->tensor = rpc_tensor; + request->offset = offset; + request->hash = fnv_hash((const uint8_t*)data, size); rpc_msg_set_tensor_hash_rsp response; - bool status = send_rpc_cmd(ctx->sock, RPC_CMD_SET_TENSOR_HASH, &request, sizeof(request), &response, sizeof(response)); - RPC_STATUS_ASSERT(status); + ctx->dispatcher->send(RPC_CMD_SET_TENSOR_HASH, request, sizeof(*request), &response, sizeof(response)); if (response.result) { // the server has the same data, no need to send it return; @@ -504,22 +708,21 @@ static void ggml_backend_rpc_buffer_set_tensor(ggml_backend_buffer_t buffer, ggm } // input serialization format: | rpc_tensor | offset (8 bytes) | data (size bytes) size_t input_size = sizeof(rpc_tensor) + sizeof(uint64_t) + size; - std::vector input(input_size, 0); - memcpy(input.data(), &rpc_tensor, sizeof(rpc_tensor)); - memcpy(input.data() + sizeof(rpc_tensor), &offset, sizeof(offset)); - memcpy(input.data() + sizeof(rpc_tensor) + sizeof(offset), data, size); - bool status = send_rpc_cmd(ctx->sock, RPC_CMD_SET_TENSOR, input.data(), input.size()); - RPC_STATUS_ASSERT(status); + uint8_t * input = new uint8_t[input_size](); + memcpy(input, &rpc_tensor, sizeof(rpc_tensor)); + memcpy(input + sizeof(rpc_tensor), &offset, sizeof(offset)); + memcpy(input + sizeof(rpc_tensor) + sizeof(offset), data, size); + std::shared_ptr input_ptr(input, std::default_delete()); + ctx->dispatcher->send(RPC_CMD_SET_TENSOR, input_ptr, input_size); } static void ggml_backend_rpc_buffer_get_tensor(ggml_backend_buffer_t buffer, const ggml_tensor * tensor, void * data, size_t offset, size_t size) { ggml_backend_rpc_buffer_context * ctx = (ggml_backend_rpc_buffer_context *)buffer->context; - rpc_msg_get_tensor_req request; - request.tensor = serialize_tensor(tensor); - request.offset = offset; - request.size = size; - bool status = send_rpc_cmd(ctx->sock, RPC_CMD_GET_TENSOR, &request, sizeof(request), data, size); - RPC_STATUS_ASSERT(status); + auto request = std::make_shared(); + request->tensor = serialize_tensor(tensor); + request->offset = offset; + request->size = size; + ctx->dispatcher->send(RPC_CMD_GET_TENSOR, request, sizeof(*request), data, size); } static bool ggml_backend_rpc_buffer_cpy_tensor(ggml_backend_buffer_t buffer, const ggml_tensor * src, ggml_tensor * dst) { @@ -529,16 +732,15 @@ static bool ggml_backend_rpc_buffer_cpy_tensor(ggml_backend_buffer_t buffer, con ggml_backend_rpc_buffer_context * src_ctx = (ggml_backend_rpc_buffer_context *)src_buffer->context; ggml_backend_buffer_t dst_buffer = dst->buffer; ggml_backend_rpc_buffer_context * dst_ctx = (ggml_backend_rpc_buffer_context *)dst_buffer->context; - if (src_ctx->sock != dst_ctx->sock) { + if (src_ctx->dispatcher != dst_ctx->dispatcher) { return false; } ggml_backend_rpc_buffer_context * ctx = (ggml_backend_rpc_buffer_context *)buffer->context; - rpc_msg_copy_tensor_req request; - request.src = serialize_tensor(src); - request.dst = serialize_tensor(dst); + auto request = std::make_shared(); + request->src = serialize_tensor(src); + request->dst = serialize_tensor(dst); rpc_msg_copy_tensor_rsp response; - bool status = send_rpc_cmd(ctx->sock, RPC_CMD_COPY_TENSOR, &request, sizeof(request), &response, sizeof(response)); - RPC_STATUS_ASSERT(status); + ctx->dispatcher->send(RPC_CMD_COPY_TENSOR, request, sizeof(*request), &response, sizeof(response)); return response.result; } return false; @@ -546,9 +748,10 @@ static bool ggml_backend_rpc_buffer_cpy_tensor(ggml_backend_buffer_t buffer, con static void ggml_backend_rpc_buffer_clear(ggml_backend_buffer_t buffer, uint8_t value) { ggml_backend_rpc_buffer_context * ctx = (ggml_backend_rpc_buffer_context *)buffer->context; - rpc_msg_buffer_clear_req request = {ctx->remote_ptr, value}; - bool status = send_rpc_cmd(ctx->sock, RPC_CMD_BUFFER_CLEAR, &request, sizeof(request), nullptr, 0); - RPC_STATUS_ASSERT(status); + auto request = std::make_shared(); + request->remote_ptr = ctx->remote_ptr; + request->value = value; + ctx->dispatcher->send(RPC_CMD_BUFFER_CLEAR, request, sizeof(*request)); } static ggml_backend_buffer_i ggml_backend_rpc_buffer_interface = { @@ -572,15 +775,17 @@ static const char * ggml_backend_rpc_buffer_type_name(ggml_backend_buffer_type_t static ggml_backend_buffer_t ggml_backend_rpc_buffer_type_alloc_buffer(ggml_backend_buffer_type_t buft, size_t size) { ggml_backend_rpc_buffer_type_context * buft_ctx = (ggml_backend_rpc_buffer_type_context *)buft->context; - rpc_msg_alloc_buffer_req request = {buft_ctx->device, size}; + auto request = std::make_shared(); + request->device = buft_ctx->device; + request->size = size; rpc_msg_alloc_buffer_rsp response; - auto sock = get_socket(buft_ctx->endpoint); - bool status = send_rpc_cmd(sock, RPC_CMD_ALLOC_BUFFER, &request, sizeof(request), &response, sizeof(response)); - RPC_STATUS_ASSERT(status); + + auto dispatcher = get_dispatcher(buft_ctx->endpoint); + dispatcher->send(RPC_CMD_ALLOC_BUFFER, request, sizeof(*request), &response, sizeof(response)); if (response.remote_ptr != 0) { ggml_backend_buffer_t buffer = ggml_backend_buffer_init(buft, ggml_backend_rpc_buffer_interface, - new ggml_backend_rpc_buffer_context{sock, nullptr, response.remote_ptr}, + new ggml_backend_rpc_buffer_context{dispatcher, nullptr, response.remote_ptr}, response.remote_size); return buffer; } else { @@ -588,11 +793,11 @@ static ggml_backend_buffer_t ggml_backend_rpc_buffer_type_alloc_buffer(ggml_back } } -static size_t get_alignment(const std::shared_ptr & sock, uint32_t device) { - rpc_msg_get_alignment_req request = {device}; +static size_t get_alignment(const std::shared_ptr & dispatcher, uint32_t device) { + auto request = std::make_shared(); + request->device = device; rpc_msg_get_alignment_rsp response; - bool status = send_rpc_cmd(sock, RPC_CMD_GET_ALIGNMENT, &request, sizeof(request), &response, sizeof(response)); - RPC_STATUS_ASSERT(status); + dispatcher->send(RPC_CMD_GET_ALIGNMENT, request, sizeof(*request), &response, sizeof(response)); return response.alignment; } @@ -601,11 +806,11 @@ static size_t ggml_backend_rpc_buffer_type_get_alignment(ggml_backend_buffer_typ return buft_ctx->alignment; } -static size_t get_max_size(const std::shared_ptr & sock, uint32_t device) { - rpc_msg_get_max_size_req request = {device}; +static size_t get_max_size(const std::shared_ptr & dispatcher, uint32_t device) { + auto request = std::make_shared(); + request->device = device; rpc_msg_get_max_size_rsp response; - bool status = send_rpc_cmd(sock, RPC_CMD_GET_MAX_SIZE, &request, sizeof(request), &response, sizeof(response)); - RPC_STATUS_ASSERT(status); + dispatcher->send(RPC_CMD_GET_MAX_SIZE, request, sizeof(*request), &response, sizeof(response)); return response.max_size; } @@ -628,23 +833,63 @@ static size_t ggml_backend_rpc_buffer_type_get_alloc_size(ggml_backend_buffer_ty if (rpc_get) { ggml_backend_rpc_buffer_type_context * buft_ctx = (ggml_backend_rpc_buffer_type_context *)buft->context; - auto sock = get_socket(buft_ctx->endpoint); - rpc_msg_get_alloc_size_req request = { - /*.device =*/ buft_ctx->device, - /*.tensor =*/ serialize_tensor(tensor), - /*.srcs =*/ {}, + // Cache key for calls to read the alloc_size. + // We deliberately exclude src tensor dimensions from the key because: + // 1. For CPU backends, alloc_size = ggml_nbytes(output) regardless of src shapes + // 2. For GPU backends, the reservation graph uses max dimensions, so the + // cached value from reservation is always >= any subsequent request + // 3. Including src dims causes cache misses per-ubatch (e.g. growing KV cache) + // which blocks the main thread behind in-flight GRAPH_COMPUTE commands + struct alloc_size_cache_key { + uint32_t device; + uint32_t type; + uint32_t op; + int32_t op_params[GGML_MAX_OP_PARAMS / sizeof(int32_t)]; + uint32_t ne[GGML_MAX_DIMS]; }; + alloc_size_cache_key key = {}; + key.device = buft_ctx->device; + key.type = tensor->type; + key.op = tensor->op; + memcpy(key.op_params, tensor->op_params, sizeof(key.op_params)); + for (int i = 0; i < GGML_MAX_DIMS; i++) { + key.ne[i] = (uint32_t)tensor->ne[i]; + } + + uint64_t cache_hash = fnv_hash((const uint8_t *)&key, sizeof(key)); + cache_hash = fnv_hash((const uint8_t *)buft_ctx->endpoint.data(), buft_ctx->endpoint.size(), cache_hash); + + // alloc sizes are immutable for a given tensor configuration + static std::mutex cache_mutex; + static std::unordered_map cache; + + { + std::lock_guard lock(cache_mutex); + auto it = cache.find(cache_hash); + if (it != cache.end()) { + return it->second; + } + } + + auto request = std::make_shared(); + request->device = buft_ctx->device; + request->tensor = serialize_tensor(tensor); + // .get_alloc_size could be a function of the tensor's srcs, so we must serialize them as well for (int i = 0; i < GGML_MAX_SRC; i++) { - request.srcs[i] = serialize_tensor(tensor->src[i]); + request->srcs[i] = serialize_tensor(tensor->src[i]); } - // TODO: cache the alloc responses to avoid extra RPC calls? rpc_msg_get_alloc_size_rsp response; - bool status = send_rpc_cmd(sock, RPC_CMD_GET_ALLOC_SIZE, &request, sizeof(request), &response, sizeof(response)); - RPC_STATUS_ASSERT(status); + auto dispatcher = get_dispatcher(buft_ctx->endpoint); + dispatcher->send(RPC_CMD_GET_ALLOC_SIZE, request, sizeof(*request), &response, sizeof(response)); + + { + std::lock_guard lock(cache_mutex); + cache[cache_hash] = response.alloc_size; + } return response.alloc_size; } @@ -673,9 +918,44 @@ static void ggml_backend_rpc_free(ggml_backend_t backend) { delete backend; } +static void ggml_backend_rpc_set_tensor_async(ggml_backend_t backend, ggml_tensor * tensor, const void * data, size_t offset, size_t size) { + ggml_backend_rpc_context * ctx = (ggml_backend_rpc_context *)backend->context; + rpc_tensor rpc_tensor = serialize_tensor(tensor); + if (size > HASH_THRESHOLD) { + auto request = std::make_shared(); + request->tensor = rpc_tensor; + request->offset = offset; + request->hash = fnv_hash((const uint8_t*)data, size); + rpc_msg_set_tensor_hash_rsp response; + // TODO: make this async + ctx->dispatcher->send(RPC_CMD_SET_TENSOR_HASH, request, sizeof(*request), &response, sizeof(response)); + if (response.result) { + // the server has the same data, no need to send it + return; + } + } + // input serialization format: | rpc_tensor | offset (8 bytes) | data (size bytes) + size_t input_size = sizeof(rpc_tensor) + sizeof(uint64_t) + size; + uint8_t * input = new uint8_t[input_size](); + memcpy(input, &rpc_tensor, sizeof(rpc_tensor)); + memcpy(input + sizeof(rpc_tensor), &offset, sizeof(offset)); + memcpy(input + sizeof(rpc_tensor) + sizeof(offset), data, size); + std::shared_ptr input_ptr(input, std::default_delete()); + ctx->dispatcher->send_async(RPC_CMD_SET_TENSOR, input_ptr, input_size); +} + +static void ggml_backend_rpc_get_tensor_async(ggml_backend_t backend, const ggml_tensor * tensor, void * data, size_t offset, size_t size) { + ggml_backend_rpc_context * ctx = (ggml_backend_rpc_context *)backend->context; + auto request = std::make_shared(); + request->tensor = serialize_tensor(tensor); + request->offset = offset; + request->size = size; + ctx->dispatcher->send_async(RPC_CMD_GET_TENSOR, request, sizeof(*request), data, size); +} + static void ggml_backend_rpc_synchronize(ggml_backend_t backend) { - GGML_UNUSED(backend); - // this is no-op because we don't have any async operations + ggml_backend_rpc_context * rpc_ctx = (ggml_backend_rpc_context *)backend->context; + rpc_ctx->dispatcher->synchronize(); } static void add_tensor(ggml_tensor * tensor, const ggml_cgraph * cgraph, std::vector & tensors, std::unordered_set & visited) { @@ -698,7 +978,7 @@ static void add_tensor(ggml_tensor * tensor, const ggml_cgraph * cgraph, std::ve tensors.push_back(result); } -static void serialize_graph(uint32_t device, const ggml_cgraph * cgraph, std::vector & output) { +static uint8_t * serialize_graph(uint32_t device, const ggml_cgraph * cgraph, size_t * output_size) { uint32_t n_nodes = cgraph->n_nodes; std::vector tensors; std::unordered_set visited; @@ -708,9 +988,9 @@ static void serialize_graph(uint32_t device, const ggml_cgraph * cgraph, std::ve // serialization format: // | device (4 bytes) | n_nodes (4 bytes) | nodes (n_nodes * sizeof(uint64_t) | n_tensors (4 bytes) | tensors (n_tensors * sizeof(rpc_tensor)) | uint32_t n_tensors = tensors.size(); - int output_size = 2*sizeof(uint32_t) + n_nodes * sizeof(uint64_t) + sizeof(uint32_t) + n_tensors * sizeof(rpc_tensor); - output.resize(output_size, 0); - uint8_t * dest = output.data(); + *output_size = 2*sizeof(uint32_t) + n_nodes * sizeof(uint64_t) + sizeof(uint32_t) + n_tensors * sizeof(rpc_tensor); + uint8_t * output = new uint8_t[*output_size](); + uint8_t * dest = output; memcpy(dest, &device, sizeof(device)); dest += sizeof(device); memcpy(dest, &n_nodes, sizeof(n_nodes)); @@ -723,6 +1003,7 @@ static void serialize_graph(uint32_t device, const ggml_cgraph * cgraph, std::ve dest += sizeof(n_tensors); rpc_tensor * out_tensors = (rpc_tensor *)dest; memcpy(out_tensors, tensors.data(), n_tensors * sizeof(rpc_tensor)); + return output; } static enum ggml_status ggml_backend_rpc_graph_compute(ggml_backend_t backend, ggml_cgraph * cgraph) { @@ -733,27 +1014,35 @@ static enum ggml_status ggml_backend_rpc_graph_compute(ggml_backend_t backend, g GGML_ASSERT(cgraph->n_nodes > 0); bool reuse = cgraph->uid != 0 && rpc_dev_ctx->last_graph_uid == cgraph->uid; if (reuse) { - rpc_msg_graph_recompute_req request; - request.device = rpc_ctx->device; - auto sock = get_socket(rpc_ctx->endpoint); - bool status = send_rpc_cmd(sock, RPC_CMD_GRAPH_RECOMPUTE, &request, sizeof(request)); - RPC_STATUS_ASSERT(status); + auto request = std::make_shared(); + request->device = rpc_ctx->device; + rpc_ctx->dispatcher->send_async(RPC_CMD_GRAPH_RECOMPUTE, request, sizeof(*request)); } else { rpc_dev_ctx->last_graph_uid = cgraph->uid; - std::vector input; - serialize_graph(rpc_ctx->device, cgraph, input); - auto sock = get_socket(rpc_ctx->endpoint); - bool status = send_rpc_cmd(sock, RPC_CMD_GRAPH_COMPUTE, input.data(), input.size()); - RPC_STATUS_ASSERT(status); + size_t input_size = 0; + uint8_t * input = serialize_graph(rpc_ctx->device, cgraph, &input_size); + std::shared_ptr input_ptr(input, std::default_delete()); + rpc_ctx->dispatcher->send_async(RPC_CMD_GRAPH_COMPUTE, input_ptr, input_size); } return GGML_STATUS_SUCCESS; } +static void ggml_backend_rpc_event_record(ggml_backend_t backend, ggml_backend_event_t event) { + ggml_backend_rpc_context * rpc_ctx = (ggml_backend_rpc_context *)backend->context; + rpc_ctx->dispatcher->event_record(event); +} + +static void ggml_backend_rpc_event_wait(ggml_backend_t backend, ggml_backend_event_t event) { + // this is noop for RPC as we have a single stream + GGML_UNUSED(backend); + GGML_UNUSED(event); +} + static ggml_backend_i ggml_backend_rpc_interface = { /* .get_name = */ ggml_backend_rpc_name, /* .free = */ ggml_backend_rpc_free, - /* .set_tensor_async = */ NULL, - /* .get_tensor_async = */ NULL, + /* .set_tensor_async = */ ggml_backend_rpc_set_tensor_async, + /* .get_tensor_async = */ ggml_backend_rpc_get_tensor_async, /* .set_tensor_2d_async = */ NULL, /* .get_tensor_2d_async = */ NULL, /* .cpy_tensor_async = */ NULL, @@ -763,8 +1052,8 @@ static ggml_backend_i ggml_backend_rpc_interface = { /* .graph_plan_update = */ NULL, /* .graph_plan_compute = */ NULL, /* .graph_compute = */ ggml_backend_rpc_graph_compute, - /* .event_record = */ NULL, - /* .event_wait = */ NULL, + /* .event_record = */ ggml_backend_rpc_event_record, + /* .event_wait = */ ggml_backend_rpc_event_wait, /* .graph_optimize = */ NULL, }; @@ -778,13 +1067,9 @@ ggml_backend_buffer_type_t ggml_backend_rpc_buffer_type(const char * endpoint, u if (it != buft_map.end()) { return it->second; } - auto sock = get_socket(endpoint); - if (sock == nullptr) { - GGML_LOG_ERROR("Failed to connect to %s\n", endpoint); - return nullptr; - } - size_t alignment = get_alignment(sock, device); - size_t max_size = get_max_size(sock, device); + auto dispatcher = get_dispatcher(endpoint); + size_t alignment = get_alignment(dispatcher, device); + size_t max_size = get_max_size(dispatcher, device); ggml_backend_rpc_buffer_type_context * buft_ctx = new ggml_backend_rpc_buffer_type_context { /* .endpoint = */ endpoint, /* .device = */ device, @@ -804,10 +1089,11 @@ ggml_backend_buffer_type_t ggml_backend_rpc_buffer_type(const char * endpoint, u ggml_backend_t ggml_backend_rpc_init(const char * endpoint, uint32_t device) { std::string dev_name = "RPC" + std::to_string(device) + "[" + std::string(endpoint) + "]"; + auto dispatcher = get_dispatcher(endpoint); ggml_backend_rpc_context * ctx = new ggml_backend_rpc_context { - /* .endpoint = */ endpoint, - /* .device = */ device, - /* .name = */ dev_name, + /* .dispatcher = */ dispatcher, + /* .device = */ device, + /* .name = */ dev_name, }; auto reg = ggml_backend_rpc_add_server(endpoint); ggml_backend_t backend = new ggml_backend { @@ -823,26 +1109,16 @@ bool ggml_backend_is_rpc(ggml_backend_t backend) { return backend != NULL && ggml_guid_matches(backend->guid, ggml_backend_rpc_guid()); } -static void get_device_memory(const std::shared_ptr & sock, uint32_t device, size_t * free, size_t * total) { - rpc_msg_get_device_memory_req request; - request.device = device; +void ggml_backend_rpc_get_device_memory(const char * endpoint, uint32_t device, size_t * free, size_t * total) { + auto dispatcher = get_dispatcher(endpoint); + auto request = std::make_shared(); + request->device = device; rpc_msg_get_device_memory_rsp response; - bool status = send_rpc_cmd(sock, RPC_CMD_GET_DEVICE_MEMORY, &request, sizeof(request), &response, sizeof(response)); - RPC_STATUS_ASSERT(status); + dispatcher->send(RPC_CMD_GET_DEVICE_MEMORY, request, sizeof(*request), &response, sizeof(response)); *free = response.free_mem; *total = response.total_mem; } -void ggml_backend_rpc_get_device_memory(const char * endpoint, uint32_t device, size_t * free, size_t * total) { - auto sock = get_socket(endpoint); - if (sock == nullptr) { - *free = 0; - *total = 0; - return; - } - get_device_memory(sock, device, free, total); -} - // RPC server-side implementation class rpc_server { @@ -1647,9 +1923,6 @@ static void rpc_serve_client(const std::vector & backends, const if (!server.free_buffer(request)) { return; } - if (!send_msg(sock, nullptr, 0)) { - return; - } break; } case RPC_CMD_BUFFER_CLEAR: { @@ -1660,9 +1933,6 @@ static void rpc_serve_client(const std::vector & backends, const if (!server.buffer_clear(request)) { return; } - if (!send_msg(sock, nullptr, 0)) { - return; - } break; } case RPC_CMD_MEMSET_TENSOR: { @@ -1673,9 +1943,6 @@ static void rpc_serve_client(const std::vector & backends, const if (!server.memset_tensor(request)) { return; } - if (!send_msg(sock, nullptr, 0)) { - return; - } break; } case RPC_CMD_SET_TENSOR: { @@ -1710,9 +1977,6 @@ static void rpc_serve_client(const std::vector & backends, const if (!server.init_tensor(request)) { return; } - if (!send_msg(sock, nullptr, 0)) { - return; - } break; } case RPC_CMD_GET_TENSOR: { @@ -1889,10 +2153,10 @@ static void ggml_backend_rpc_device_get_props(ggml_backend_dev_t dev, struct ggm props->type = ggml_backend_rpc_device_get_type(dev); ggml_backend_rpc_device_get_memory(dev, &props->memory_free, &props->memory_total); props->caps = { - /* .async = */ false, + /* .async = */ true, /* .host_buffer = */ false, /* .buffer_from_host_ptr = */ false, - /* .events = */ false, + /* .events = */ true, /* .mmap_support = */ true, }; } @@ -1929,6 +2193,24 @@ static bool ggml_backend_rpc_device_supports_buft(ggml_backend_dev_t dev, ggml_b return buft_ctx->endpoint == dev_ctx->endpoint && buft_ctx->device == dev_ctx->device; } +static ggml_backend_event_t ggml_backend_rpc_device_event_new(ggml_backend_dev_t dev) { + ggml_backend_rpc_device_context * ctx = (ggml_backend_rpc_device_context *)dev->context; + auto dispatcher = get_dispatcher(ctx->endpoint); + return dispatcher->event_new(dev); +} + +static void ggml_backend_rpc_device_event_free(ggml_backend_dev_t dev, ggml_backend_event_t event) { + ggml_backend_rpc_device_context * ctx = (ggml_backend_rpc_device_context *)dev->context; + auto dispatcher = get_dispatcher(ctx->endpoint); + dispatcher->event_free(event); +} + +static void ggml_backend_rpc_device_event_synchronize(ggml_backend_dev_t dev, ggml_backend_event_t event) { + ggml_backend_rpc_device_context * ctx = (ggml_backend_rpc_device_context *)dev->context; + auto dispatcher = get_dispatcher(ctx->endpoint); + dispatcher->event_synchronize(event); +} + static const struct ggml_backend_device_i ggml_backend_rpc_device_i = { /* .get_name = */ ggml_backend_rpc_device_get_name, /* .get_description = */ ggml_backend_rpc_device_get_description, @@ -1942,9 +2224,9 @@ static const struct ggml_backend_device_i ggml_backend_rpc_device_i = { /* .supports_op = */ ggml_backend_rpc_device_supports_op, /* .supports_buft = */ ggml_backend_rpc_device_supports_buft, /* .offload_op = */ NULL, - /* .event_new = */ NULL, - /* .event_free = */ NULL, - /* .event_synchronize = */ NULL, + /* .event_new = */ ggml_backend_rpc_device_event_new, + /* .event_free = */ ggml_backend_rpc_device_event_free, + /* .event_synchronize = */ ggml_backend_rpc_device_event_synchronize, }; // backend reg interface @@ -2004,14 +2286,9 @@ ggml_backend_reg_t ggml_backend_rpc_reg(void) { } static uint32_t ggml_backend_rpc_get_device_count(const char * endpoint) { - auto sock = get_socket(endpoint); - if (sock == nullptr) { - GGML_LOG_ERROR("Failed to connect to %s\n", endpoint); - return 0; - } + auto dispatcher = get_dispatcher(endpoint); rpc_msg_device_count_rsp response; - bool status = send_rpc_cmd(sock, RPC_CMD_DEVICE_COUNT, nullptr, 0, &response, sizeof(response)); - RPC_STATUS_ASSERT(status); + dispatcher->send(RPC_CMD_DEVICE_COUNT, nullptr, 0, &response, sizeof(response)); return response.device_count; } From 0a026972dae2f58ac5742e750767ac8e159b16e1 Mon Sep 17 00:00:00 2001 From: Pranav Uttarkar <122235768+PranavUttarkar@users.noreply.github.com> Date: Wed, 26 Aug 2026 09:49:32 -0500 Subject: [PATCH 008/104] Implemented vulkan cross_entropy_loss and cross_entropy_loss_back (llama/27216) --- ggml/src/ggml-vulkan/ggml-vulkan.cpp | 138 ++++++++++++++++++ .../vulkan-shaders/cross_entropy_loss.comp | 78 ++++++++++ .../cross_entropy_loss_back.comp | 75 ++++++++++ .../vulkan-shaders/vulkan-shaders-gen.cpp | 2 + 4 files changed, 293 insertions(+) create mode 100644 ggml/src/ggml-vulkan/vulkan-shaders/cross_entropy_loss.comp create mode 100644 ggml/src/ggml-vulkan/vulkan-shaders/cross_entropy_loss_back.comp diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp index c1d86aaac5c..31923a2ba0e 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp +++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp @@ -1042,6 +1042,8 @@ struct vk_device_struct { vk_pipeline pipeline_argsort_large_f32[num_argsort_pipelines]; vk_pipeline pipeline_topk_f32[num_topk_pipelines]; vk_pipeline pipeline_sum_rows_f32; + vk_pipeline pipeline_cross_entropy_loss_f32, pipeline_cross_entropy_loss_f32_wg512; + vk_pipeline pipeline_cross_entropy_loss_back_f32, pipeline_cross_entropy_loss_back_f32_wg512; vk_pipeline pipeline_fwht_f32[4]; vk_pipeline pipeline_cumsum_f32; vk_pipeline pipeline_cumsum_small_f32; @@ -5758,6 +5760,10 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { ggml_vk_create_pipeline(device, device->pipeline_argmax_f32, "argmax_f32", argmax_f32_len, argmax_f32_data, "main", 2, sizeof(vk_op_push_constants), {1, 1, 1}, { device->subgroup_size }, 1); ggml_vk_create_pipeline(device, device->pipeline_sum_rows_f32, "sum_rows_f32", sum_rows_f32_len, sum_rows_f32_data, "main", 2, sizeof(vk_op_sum_rows_push_constants), {1, 1, 1}, { device->subgroup_size }, 1); + ggml_vk_create_pipeline(device, device->pipeline_cross_entropy_loss_f32, "cross_entropy_loss_f32", cross_entropy_loss_f32_len, cross_entropy_loss_f32_data, "main", 3, sizeof(vk_op_push_constants), {1, 1, 1}, { device->subgroup_size }, 1); + ggml_vk_create_pipeline(device, device->pipeline_cross_entropy_loss_f32_wg512, "cross_entropy_loss_f32_wg512", cross_entropy_loss_f32_len, cross_entropy_loss_f32_data, "main", 3, sizeof(vk_op_push_constants), {1, 1, 1}, { 512 }, 1); + ggml_vk_create_pipeline(device, device->pipeline_cross_entropy_loss_back_f32, "cross_entropy_loss_back_f32", cross_entropy_loss_back_f32_len, cross_entropy_loss_back_f32_data, "main", 4, sizeof(vk_op_push_constants), {1, 1, 1}, { device->subgroup_size }, 1); + ggml_vk_create_pipeline(device, device->pipeline_cross_entropy_loss_back_f32_wg512, "cross_entropy_loss_back_f32_wg512", cross_entropy_loss_back_f32_len, cross_entropy_loss_back_f32_data, "main", 4, sizeof(vk_op_push_constants), {1, 1, 1}, { 512 }, 1); // Intel Windows driver in range [32.0.101.8509, 32.0.101.8860) will crash when using fwht kernels so we gate that here const bool can_use_fwht = device->driver_id != vk::DriverId::eIntelProprietaryWindows || !ggml_vk_intel_windows_driver_in_range(device->properties.driverVersion, 101, 8509, 101, 8860); @@ -11577,6 +11583,17 @@ static vk_pipeline ggml_vk_op_get_pipeline(ggml_backend_vk_context * ctx, const return ctx->device->pipeline_sum_rows_f32; } return nullptr; + case GGML_OP_CROSS_ENTROPY_LOSS: + if (src0->type == GGML_TYPE_F32 && src1->type == GGML_TYPE_F32 && dst->type == GGML_TYPE_F32) { + return src0->ne[0] > 1024 ? ctx->device->pipeline_cross_entropy_loss_f32_wg512 : ctx->device->pipeline_cross_entropy_loss_f32; + } + return nullptr; + case GGML_OP_CROSS_ENTROPY_LOSS_BACK: + // src0 is the scalar grad; src1 is logits + if (src0->type == GGML_TYPE_F32 && src1->type == GGML_TYPE_F32 && src2 && src2->type == GGML_TYPE_F32 && dst->type == GGML_TYPE_F32) { + return src1->ne[0] > 1024 ? ctx->device->pipeline_cross_entropy_loss_back_f32_wg512 : ctx->device->pipeline_cross_entropy_loss_back_f32; + } + return nullptr; case GGML_OP_CUMSUM: if (src0->type == GGML_TYPE_F32 && dst->type == GGML_TYPE_F32) { if (src0->ne[0] <= 512) { @@ -13942,6 +13959,103 @@ static void ggml_vk_cumsum(ggml_backend_vk_context * ctx, vk_context& subctx, co ctx->prealloc_split_k_need_sync = true; } +static std::array ggml_vk_nrows_elements(uint32_t nr) { + if (nr > 262144) { + return { 512, 512, CEIL_DIV(nr, 262144) }; + } + if (nr > 512) { + return { 512, CEIL_DIV(nr, 512), 1 }; + } + return { nr, 1, 1 }; +} + +static void ggml_vk_cross_entropy_loss(ggml_backend_vk_context * ctx, vk_context& subctx, ggml_tensor * dst) { + const ggml_tensor * src0 = dst->src[0]; + const ggml_tensor * src1 = dst->src[1]; + + GGML_ASSERT(src0->type == GGML_TYPE_F32); + GGML_ASSERT(src1->type == GGML_TYPE_F32); + GGML_ASSERT(dst->type == GGML_TYPE_F32); + GGML_ASSERT(ggml_is_contiguous(src0)); + GGML_ASSERT(ggml_is_contiguous(src1)); + GGML_ASSERT(ggml_is_contiguous(dst)); + GGML_ASSERT(ggml_are_same_shape(src0, src1)); + GGML_ASSERT(ggml_is_scalar(dst)); + + const uint32_t nclasses = (uint32_t)src0->ne[0]; + const uint32_t nrows = (uint32_t)ggml_nrows(src0); + + vk_pipeline pipeline = ggml_vk_op_get_pipeline(ctx, src0, src1, nullptr, dst, GGML_OP_CROSS_ENTROPY_LOSS); + GGML_ASSERT(pipeline != nullptr); + + ggml_pipeline_request_descriptor_sets(ctx, pipeline, 1); + ggml_pipeline_request_descriptor_sets(ctx, ctx->device->pipeline_sum_rows_f32, 1); + + vk_subbuffer src0_buf = ggml_vk_tensor_subbuffer(ctx, src0); + vk_subbuffer src1_buf = ggml_vk_tensor_subbuffer(ctx, src1); + vk_subbuffer dst_buf = ggml_vk_tensor_subbuffer(ctx, dst, true); + + const vk_op_push_constants pc = { nclasses, nrows, 0.0f, 0.0f, 0.0f, 0.0f }; + + const size_t tmp_size = (size_t)nrows * sizeof(float); + if (ctx->prealloc_size_x < tmp_size) { + ctx->prealloc_size_x = tmp_size; + ggml_vk_preallocate_buffers(ctx, subctx); + } + if (ctx->prealloc_x_need_sync) { + ggml_vk_sync_buffers(ctx, subctx); + } + + vk_subbuffer tmp_buf = { ctx->prealloc_x, 0, tmp_size }; + ggml_vk_dispatch_pipeline(ctx, subctx, pipeline, { src0_buf, src1_buf, tmp_buf }, pc, ggml_vk_nrows_elements(nrows)); + ggml_vk_sync_buffers(ctx, subctx); + + vk_op_sum_rows_push_constants sp = {}; + sp.n_cols = nrows; + sp.ne01 = 1; + sp.ne02 = 1; + sp.weight = 1.0f; + init_pushconst_fastdiv(sp); + sp.misalign_offsets = get_misalign_bytes(ctx, dst) / ggml_type_size(dst->type); + + ggml_vk_dispatch_pipeline(ctx, subctx, ctx->device->pipeline_sum_rows_f32, { tmp_buf, dst_buf }, sp, { 1, 1, 1 }); + ctx->prealloc_x_need_sync = true; +} + +static void ggml_vk_cross_entropy_loss_back(ggml_backend_vk_context * ctx, vk_context& subctx, ggml_tensor * dst) { + const ggml_tensor * grad = dst->src[0]; + const ggml_tensor * logits = dst->src[1]; + const ggml_tensor * labels = dst->src[2]; + + GGML_ASSERT(grad->type == GGML_TYPE_F32); + GGML_ASSERT(logits->type == GGML_TYPE_F32); + GGML_ASSERT(labels->type == GGML_TYPE_F32); + GGML_ASSERT(dst->type == GGML_TYPE_F32); + GGML_ASSERT(ggml_is_scalar(grad)); + GGML_ASSERT(ggml_is_contiguous(grad)); + GGML_ASSERT(ggml_is_contiguous(logits)); + GGML_ASSERT(ggml_is_contiguous(labels)); + GGML_ASSERT(ggml_is_contiguous(dst)); + GGML_ASSERT(ggml_are_same_shape(logits, labels)); + GGML_ASSERT(ggml_are_same_shape(logits, dst)); + + const uint32_t nclasses = (uint32_t)logits->ne[0]; + const uint32_t nrows = (uint32_t)ggml_nrows(logits); + + vk_pipeline pipeline = ggml_vk_op_get_pipeline(ctx, grad, logits, labels, dst, GGML_OP_CROSS_ENTROPY_LOSS_BACK); + GGML_ASSERT(pipeline != nullptr); + + ggml_pipeline_request_descriptor_sets(ctx, pipeline, 1); + + vk_subbuffer grad_buf = ggml_vk_tensor_subbuffer(ctx, grad); + vk_subbuffer logits_buf = ggml_vk_tensor_subbuffer(ctx, logits); + vk_subbuffer labels_buf = ggml_vk_tensor_subbuffer(ctx, labels); + vk_subbuffer dst_buf = ggml_vk_tensor_subbuffer(ctx, dst); + + const vk_op_push_constants pc = { nclasses, nrows, 0.0f, 0.0f, 0.0f, 0.0f }; + ggml_vk_dispatch_pipeline(ctx, subctx, pipeline, { grad_buf, logits_buf, labels_buf, dst_buf }, pc, ggml_vk_nrows_elements(nrows)); +} + static void ggml_vk_argmax(ggml_backend_vk_context * ctx, vk_context& subctx, const ggml_tensor * src0, ggml_tensor * dst) { ggml_vk_op_f32(ctx, subctx, src0, nullptr, nullptr, nullptr, dst, GGML_OP_ARGMAX, { (uint32_t)src0->ne[0], (uint32_t)src0->ne[1], 0.0f, 0.0f, 0.0f, 0.0f }); } @@ -15687,6 +15801,14 @@ static bool ggml_vk_build_graph(ggml_backend_vk_context * ctx, ggml_cgraph * cgr case GGML_OP_ARGMAX: ggml_vk_argmax(ctx, compute_ctx, src0, node); + break; + case GGML_OP_CROSS_ENTROPY_LOSS: + ggml_vk_cross_entropy_loss(ctx, compute_ctx, node); + + break; + case GGML_OP_CROSS_ENTROPY_LOSS_BACK: + ggml_vk_cross_entropy_loss_back(ctx, compute_ctx, node); + break; case GGML_OP_COUNT_EQUAL: ggml_vk_count_equal(ctx, compute_ctx, src0, src1, node); @@ -18511,6 +18633,18 @@ static bool ggml_backend_vk_device_supports_op(ggml_backend_dev_t dev, const ggm } case GGML_OP_ARGMAX: return ggml_is_contiguous(op->src[0]) && op->src[0]->type == GGML_TYPE_F32; + case GGML_OP_CROSS_ENTROPY_LOSS: + return ggml_is_contiguous(op->src[0]) && op->src[0]->type == GGML_TYPE_F32 + && ggml_is_contiguous(op->src[1]) && op->src[1]->type == GGML_TYPE_F32 + && ggml_are_same_shape(op->src[0], op->src[1]) + && ggml_is_contiguous(op) && ggml_is_scalar(op) && op->type == GGML_TYPE_F32; + case GGML_OP_CROSS_ENTROPY_LOSS_BACK: + return ggml_is_contiguous(op->src[0]) && op->src[0]->type == GGML_TYPE_F32 && ggml_is_scalar(op->src[0]) + && ggml_is_contiguous(op->src[1]) && op->src[1]->type == GGML_TYPE_F32 + && ggml_is_contiguous(op->src[2]) && op->src[2]->type == GGML_TYPE_F32 + && ggml_are_same_shape(op->src[1], op->src[2]) + && ggml_are_same_shape(op->src[1], op) + && ggml_is_contiguous(op) && op->type == GGML_TYPE_F32; case GGML_OP_COUNT_EQUAL: return ggml_is_contiguous(op->src[0]) && op->src[0]->type == GGML_TYPE_I32 && ggml_is_contiguous(op->src[1]) && op->src[1]->type == GGML_TYPE_I32; @@ -19437,6 +19571,10 @@ static void ggml_vk_check_results_0(ggml_backend_vk_context * ctx, ggml_cgraph * tensor_clone = ggml_mean(ggml_ctx, src_clone[0]); } else if (tensor->op == GGML_OP_ARGMAX) { tensor_clone = ggml_argmax(ggml_ctx, src_clone[0]); + } else if (tensor->op == GGML_OP_CROSS_ENTROPY_LOSS) { + tensor_clone = ggml_cross_entropy_loss(ggml_ctx, src_clone[0], src_clone[1]); + } else if (tensor->op == GGML_OP_CROSS_ENTROPY_LOSS_BACK) { + tensor_clone = ggml_cross_entropy_loss_back(ggml_ctx, src_clone[0], src_clone[1], src_clone[2]); } else if (tensor->op == GGML_OP_COUNT_EQUAL) { tensor_clone = ggml_count_equal(ggml_ctx, src_clone[0], src_clone[1]); } else if (tensor->op == GGML_OP_SOLVE_TRI) { diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/cross_entropy_loss.comp b/ggml/src/ggml-vulkan/vulkan-shaders/cross_entropy_loss.comp new file mode 100644 index 00000000000..0c135c6fd24 --- /dev/null +++ b/ggml/src/ggml-vulkan/vulkan-shaders/cross_entropy_loss.comp @@ -0,0 +1,78 @@ +#version 450 + +#include "generic_head.glsl" +#include "types.glsl" + +#extension GL_EXT_control_flow_attributes : enable + +layout(constant_id = 0) const uint BLOCK_SIZE = 32; +layout(local_size_x_id = 0, local_size_y = 1, local_size_z = 1) in; + +layout (binding = 0) readonly buffer A {A_TYPE data_a[];}; +layout (binding = 1) readonly buffer B {B_TYPE data_b[];}; +layout (binding = 2) writeonly buffer D {D_TYPE data_d[];}; + +shared FLOAT_TYPE tmp[BLOCK_SIZE]; + +FLOAT_TYPE wg_reduce_max(FLOAT_TYPE v) { + const uint tid = gl_LocalInvocationID.x; + tmp[tid] = v; + barrier(); + [[unroll]] for (uint s = BLOCK_SIZE / 2; s > 0; s >>= 1) { + if (tid < s) { + tmp[tid] = max(tmp[tid], tmp[tid + s]); + } + barrier(); + } + v = tmp[0]; + barrier(); + return v; +} + +FLOAT_TYPE wg_reduce_sum(FLOAT_TYPE v) { + const uint tid = gl_LocalInvocationID.x; + tmp[tid] = v; + barrier(); + [[unroll]] for (uint s = BLOCK_SIZE / 2; s > 0; s >>= 1) { + if (tid < s) { + tmp[tid] += tmp[tid + s]; + } + barrier(); + } + v = tmp[0]; + barrier(); + return v; +} + +void main() { + const uint row = gl_WorkGroupID.z * 262144 + gl_WorkGroupID.y * 512 + gl_WorkGroupID.x; + const uint tid = gl_LocalInvocationID.x; + + if (row >= p.KY) { + return; + } + + const uint off = row * p.KX; + + FLOAT_TYPE max_logit = FLOAT_TYPE(uintBitsToFloat(0xFF800000)); + for (uint i = tid; i < p.KX; i += BLOCK_SIZE) { + max_logit = max(max_logit, FLOAT_TYPE(data_a[off + i])); + } + max_logit = wg_reduce_max(max_logit); + + FLOAT_TYPE sum_exp = FLOAT_TYPE(0.0f); + for (uint i = tid; i < p.KX; i += BLOCK_SIZE) { + sum_exp += exp(FLOAT_TYPE(data_a[off + i]) - max_logit); + } + const FLOAT_TYPE log_sum = log(wg_reduce_sum(sum_exp)); + + FLOAT_TYPE loss = FLOAT_TYPE(0.0f); + for (uint i = tid; i < p.KX; i += BLOCK_SIZE) { + loss += (FLOAT_TYPE(data_a[off + i]) - max_logit - log_sum) * FLOAT_TYPE(data_b[off + i]); + } + loss = -wg_reduce_sum(loss) / FLOAT_TYPE(p.KY); + + if (tid == 0) { + data_d[row] = D_TYPE(loss); + } +} diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/cross_entropy_loss_back.comp b/ggml/src/ggml-vulkan/vulkan-shaders/cross_entropy_loss_back.comp new file mode 100644 index 00000000000..3cdebe86e47 --- /dev/null +++ b/ggml/src/ggml-vulkan/vulkan-shaders/cross_entropy_loss_back.comp @@ -0,0 +1,75 @@ +#version 450 + +#include "generic_head.glsl" +#include "types.glsl" + +#extension GL_EXT_control_flow_attributes : enable + +layout(constant_id = 0) const uint BLOCK_SIZE = 32; +layout(local_size_x_id = 0, local_size_y = 1, local_size_z = 1) in; + +layout (binding = 0) readonly buffer G {A_TYPE data_g[];}; +layout (binding = 1) readonly buffer X {B_TYPE data_x[];}; +layout (binding = 2) readonly buffer Y {B_TYPE data_y[];}; +layout (binding = 3) writeonly buffer D {D_TYPE data_d[];}; + +shared FLOAT_TYPE tmp[BLOCK_SIZE]; + +FLOAT_TYPE wg_reduce_max(FLOAT_TYPE v) { + const uint tid = gl_LocalInvocationID.x; + tmp[tid] = v; + barrier(); + [[unroll]] for (uint s = BLOCK_SIZE / 2; s > 0; s >>= 1) { + if (tid < s) { + tmp[tid] = max(tmp[tid], tmp[tid + s]); + } + barrier(); + } + v = tmp[0]; + barrier(); + return v; +} + +FLOAT_TYPE wg_reduce_sum(FLOAT_TYPE v) { + const uint tid = gl_LocalInvocationID.x; + tmp[tid] = v; + barrier(); + [[unroll]] for (uint s = BLOCK_SIZE / 2; s > 0; s >>= 1) { + if (tid < s) { + tmp[tid] += tmp[tid + s]; + } + barrier(); + } + v = tmp[0]; + barrier(); + return v; +} + +void main() { + const uint row = gl_WorkGroupID.z * 262144 + gl_WorkGroupID.y * 512 + gl_WorkGroupID.x; + const uint tid = gl_LocalInvocationID.x; + + if (row >= p.KY) { + return; + } + + const uint off = row * p.KX; + const FLOAT_TYPE d_by_nrows = FLOAT_TYPE(data_g[0]) / FLOAT_TYPE(p.KY); + + FLOAT_TYPE max_logit = FLOAT_TYPE(uintBitsToFloat(0xFF800000)); + for (uint i = tid; i < p.KX; i += BLOCK_SIZE) { + max_logit = max(max_logit, FLOAT_TYPE(data_x[off + i])); + } + max_logit = wg_reduce_max(max_logit); + + FLOAT_TYPE sum_exp = FLOAT_TYPE(0.0f); + for (uint i = tid; i < p.KX; i += BLOCK_SIZE) { + sum_exp += exp(FLOAT_TYPE(data_x[off + i]) - max_logit); + } + const FLOAT_TYPE inv_sum = FLOAT_TYPE(1.0f) / wg_reduce_sum(sum_exp); + + for (uint i = tid; i < p.KX; i += BLOCK_SIZE) { + const FLOAT_TYPE sm = exp(FLOAT_TYPE(data_x[off + i]) - max_logit) * inv_sum; + data_d[off + i] = D_TYPE((sm - FLOAT_TYPE(data_y[off + i])) * d_by_nrows); + } +} diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp b/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp index 17d57d5a18f..dbb99782cf7 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp @@ -1029,6 +1029,8 @@ void process_shaders() { string_to_spv("argmax_f32", "argmax.comp", merge_maps(base_dict, {{"A_TYPE", "float"}, {"D_TYPE", "int"}})); string_to_spv("sum_rows_f32", "sum_rows.comp", merge_maps(base_dict, {{"A_TYPE", "float"}, {"D_TYPE", "float"}})); + string_to_spv("cross_entropy_loss_f32", "cross_entropy_loss.comp", merge_maps(base_dict, {{"A_TYPE", "float"}, {"B_TYPE", "float"}, {"D_TYPE", "float"}})); + string_to_spv("cross_entropy_loss_back_f32", "cross_entropy_loss_back.comp", merge_maps(base_dict, {{"A_TYPE", "float"}, {"B_TYPE", "float"}, {"D_TYPE", "float"}})); string_to_spv("fwht_f32", "fwht.comp", {}); string_to_spv("fwht_shmem_f32", "fwht.comp", {{"FWHT_SHMEM", "1"}}); string_to_spv("count_equal_i32", "count_equal.comp", merge_maps(base_dict, {{"A_TYPE", "int"}, {"B_TYPE", "int"}, {"D_TYPE", "int"}})); From 5271734e0e4e1761c601dc4fcf41a877e22d3a80 Mon Sep 17 00:00:00 2001 From: Ruben Ortlam Date: Wed, 26 Aug 2026 18:02:06 +0200 Subject: [PATCH 009/104] vulkan: warptiles currently assume warp sizes <= 64, clamp to work around larger warps (llama/27726) --- ggml/src/ggml-vulkan/ggml-vulkan.cpp | 68 +++++++++++++++------------- 1 file changed, 37 insertions(+), 31 deletions(-) diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp index 31923a2ba0e..8108e94c16c 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp +++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp @@ -4171,10 +4171,16 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { const uint32_t subgroup_size_16 = std::max(device->subgroup_size, 16u); const uint32_t subgroup_size_32 = std::max(device->subgroup_size, 32u); + // clamp WARP for l_/m_ warptiles so WM <= BM (breaks on subgroupSize > 64) + const uint32_t mm_warp_8 = std::min(subgroup_size_8, 64u); + const uint32_t mm_warp_16 = std::min(subgroup_size_16, 64u); + const uint32_t mul_mat_subgroup_size = (device->vendor_id == VK_VENDOR_ID_INTEL && device->subgroup_size_control) ? device->subgroup_min_size : device->subgroup_size; const uint32_t mul_mat_subgroup_size_8 = std::max(mul_mat_subgroup_size, 8u); const uint32_t mul_mat_subgroup_size_16 = std::max(mul_mat_subgroup_size, 16u); const uint32_t mul_mat_subgroup_size_32 = std::max(mul_mat_subgroup_size, 32u); + const uint32_t mul_mat_mm_warp_8 = std::min(mul_mat_subgroup_size_8, 64u); + const uint32_t mul_mat_mm_warp_16 = std::min(mul_mat_subgroup_size_16, 64u); const bool subgroup_min_size_16 = (!device->subgroup_size_control && device->subgroup_size >= 16) || (device->subgroup_size_control && device->subgroup_max_size >= 16); @@ -4255,39 +4261,39 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { const uint32_t s_warptile_wm = device->subgroup_size == 8 ? 8 : 32; - l_warptile = { 128, 128, 128, 16, subgroup_size_8 * 2, 64, 2, tm_l, tn_l, tk_l, subgroup_size_8 }; - m_warptile = { 128, 64, 64, 16, subgroup_size_8, 32, 2, tm_m, tn_m, tk_m, subgroup_size_8 }; - s_warptile = { subgroup_size_32, 32, 32, 16, s_warptile_wm, 32, 2, tm_s, tn_s, tk_s, subgroup_size_8 }; + l_warptile = { 128, 128, 128, 16, mm_warp_8 * 2, 64, 2, tm_l, tn_l, tk_l, mm_warp_8 }; + m_warptile = { 128, 64, 64, 16, mm_warp_8, 32, 2, tm_m, tn_m, tk_m, mm_warp_8 }; + s_warptile = { subgroup_size_32, 32, 32, 16, s_warptile_wm, 32, 2, tm_s, tn_s, tk_s, subgroup_size_8 }; - l_warptile_mmq = { 128, 128, 128, 32, subgroup_size_8 * 2, 64, 2, tm_l, tn_l, tk_l, subgroup_size_8 }; - m_warptile_mmq = { 128, 64, 64, 32, subgroup_size_8, 32, 2, tm_m, tn_m, tk_m, subgroup_size_8 }; - s_warptile_mmq = { subgroup_size_32, 32, 32, 32, s_warptile_wm, 32, 2, tm_s, tn_s, tk_s, subgroup_size_8 }; + l_warptile_mmq = { 128, 128, 128, 32, mm_warp_8 * 2, 64, 2, tm_l, tn_l, tk_l, mm_warp_8 }; + m_warptile_mmq = { 128, 64, 64, 32, mm_warp_8, 32, 2, tm_m, tn_m, tk_m, mm_warp_8 }; + s_warptile_mmq = { subgroup_size_32, 32, 32, 32, s_warptile_wm, 32, 2, tm_s, tn_s, tk_s, subgroup_size_8 }; // Integer MMQ has a smaller shared memory profile, but heavier register use - l_warptile_mmq_int = { 128, 128, 128, 32, subgroup_size_8 * 2, 64, 2, 4, 4, 1, subgroup_size_8 }; - m_warptile_mmq_int = { 128, 64, 64, 32, subgroup_size_8, 32, 2, 2, 2, 1, subgroup_size_8 }; - s_warptile_mmq_int = { subgroup_size_32, 32, 32, 32, s_warptile_wm, 32, 2, 2, 1, 1, subgroup_size_8 }; + l_warptile_mmq_int = { 128, 128, 128, 32, mm_warp_8 * 2, 64, 2, 4, 4, 1, mm_warp_8 }; + m_warptile_mmq_int = { 128, 64, 64, 32, mm_warp_8, 32, 2, 2, 2, 1, mm_warp_8 }; + s_warptile_mmq_int = { subgroup_size_32, 32, 32, 32, s_warptile_wm, 32, 2, 2, 1, 1, subgroup_size_8 }; // K-quants use even more registers, mitigate by setting WMITER to 1 - l_warptile_mmq_int_k = { 128, 128, 128, 32, subgroup_size_8 * 2, 64, 1, 4, 4, 1, subgroup_size_8 }; - m_warptile_mmq_int_k = { 128, 64, 64, 32, subgroup_size_8, 32, 1, 2, 2, 1, subgroup_size_8 }; - s_warptile_mmq_int_k = { subgroup_size_32, 32, 32, 32, s_warptile_wm, 32, 1, 2, 1, 1, subgroup_size_8 }; + l_warptile_mmq_int_k = { 128, 128, 128, 32, mm_warp_8 * 2, 64, 1, 4, 4, 1, mm_warp_8 }; + m_warptile_mmq_int_k = { 128, 64, 64, 32, mm_warp_8, 32, 1, 2, 2, 1, mm_warp_8 }; + s_warptile_mmq_int_k = { subgroup_size_32, 32, 32, 32, s_warptile_wm, 32, 1, 2, 1, 1, subgroup_size_8 }; - l_warptile_id = { 128, 128, 128, 16, mul_mat_subgroup_size_16 * 2, 64, 2, tm_l, tn_l, tk_l, mul_mat_subgroup_size_16 }; - m_warptile_id = { 128, 64, 64, 16, mul_mat_subgroup_size_16, 32, 2, tm_m, tn_m, tk_m, mul_mat_subgroup_size_16 }; - s_warptile_id = { mul_mat_subgroup_size_16, 32, 32, 16, s_warptile_wm, 32, 2, tm_s, tn_s, tk_s, mul_mat_subgroup_size_16 }; + l_warptile_id = { 128, 128, 128, 16, mul_mat_mm_warp_16 * 2, 64, 2, tm_l, tn_l, tk_l, mul_mat_mm_warp_16 }; + m_warptile_id = { 128, 64, 64, 16, mul_mat_mm_warp_16, 32, 2, tm_m, tn_m, tk_m, mul_mat_mm_warp_16 }; + s_warptile_id = { mul_mat_subgroup_size_16, 32, 32, 16, s_warptile_wm, 32, 2, tm_s, tn_s, tk_s, mul_mat_subgroup_size_16 }; - l_warptile_mmqid = { 128, 128, 128, 32, mul_mat_subgroup_size_8 * 2, 64, 2, tm_l, tn_l, tk_l, mul_mat_subgroup_size_8 }; - m_warptile_mmqid = { 128, 64, 64, 32, mul_mat_subgroup_size_8, 32, 2, tm_m, tn_m, tk_m, mul_mat_subgroup_size_8 }; - s_warptile_mmqid = { mul_mat_subgroup_size_32, 32, 32, 32, s_warptile_wm, 32, 2, tm_s, tn_s, tk_s, mul_mat_subgroup_size_8 }; + l_warptile_mmqid = { 128, 128, 128, 32, mul_mat_mm_warp_8 * 2, 64, 2, tm_l, tn_l, tk_l, mul_mat_mm_warp_8 }; + m_warptile_mmqid = { 128, 64, 64, 32, mul_mat_mm_warp_8, 32, 2, tm_m, tn_m, tk_m, mul_mat_mm_warp_8 }; + s_warptile_mmqid = { mul_mat_subgroup_size_32, 32, 32, 32, s_warptile_wm, 32, 2, tm_s, tn_s, tk_s, mul_mat_subgroup_size_8 }; - l_warptile_mmqid_int = { 128, 128, 128, 32, mul_mat_subgroup_size_8 * 2, 64, 2, 4, 4, 1, mul_mat_subgroup_size_8 }; - m_warptile_mmqid_int = { 128, 64, 64, 32, mul_mat_subgroup_size_8, 32, 2, 2, 2, 1, mul_mat_subgroup_size_8 }; - s_warptile_mmqid_int = { mul_mat_subgroup_size_32, 32, 32, 32, s_warptile_wm, 32, 2, 2, 1, 1, mul_mat_subgroup_size_8 }; + l_warptile_mmqid_int = { 128, 128, 128, 32, mul_mat_mm_warp_8 * 2, 64, 2, 4, 4, 1, mul_mat_mm_warp_8 }; + m_warptile_mmqid_int = { 128, 64, 64, 32, mul_mat_mm_warp_8, 32, 2, 2, 2, 1, mul_mat_mm_warp_8 }; + s_warptile_mmqid_int = { mul_mat_subgroup_size_32, 32, 32, 32, s_warptile_wm, 32, 2, 2, 1, 1, mul_mat_subgroup_size_8 }; - l_warptile_mmqid_int_k = { 128, 128, 128, 32, mul_mat_subgroup_size_16 * 2, 64, 1, 4, 4, 1, mul_mat_subgroup_size_16 }; - m_warptile_mmqid_int_k = { 128, 64, 64, 32, mul_mat_subgroup_size_16, 32, 1, 2, 2, 1, mul_mat_subgroup_size_16 }; - s_warptile_mmqid_int_k = { mul_mat_subgroup_size_32, 32, 32, 32, s_warptile_wm, 32, 1, 2, 1, 1, mul_mat_subgroup_size_16 }; + l_warptile_mmqid_int_k = { 128, 128, 128, 32, mul_mat_mm_warp_16 * 2, 64, 1, 4, 4, 1, mul_mat_mm_warp_16 }; + m_warptile_mmqid_int_k = { 128, 64, 64, 32, mul_mat_mm_warp_16, 32, 1, 2, 2, 1, mul_mat_mm_warp_16 }; + s_warptile_mmqid_int_k = { mul_mat_subgroup_size_32, 32, 32, 32, s_warptile_wm, 32, 1, 2, 1, 1, mul_mat_subgroup_size_16 }; // chip specific tuning if ((device->architecture == AMD_GCN) && (device->driver_id != vk::DriverId::eAmdProprietary)) { @@ -4295,13 +4301,13 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { m_warptile_mmqid = m_warptile_mmqid_int = { 256, 64, 64, 32, 16, 16, 2, 2, 2, 1, 16 }; } else if (device->vendor_id == VK_VENDOR_ID_AMD && device->coopmat_support && device->driver_id != vk::DriverId::eAmdProprietary) { // This is intentionally using tx_m values, slight performance increase - l_warptile = { 256, 128, 128, 16, subgroup_size_8, 64, 2, tm_m, tn_m, tk_m, subgroup_size_8 }; - l_warptile_mmq = l_warptile_mmq_int = { 256, 128, 128, 32, subgroup_size_8, 64, 2, tm_m, tn_m, tk_m, subgroup_size_8 }; - l_warptile_mmq_int_k = { 256, 128, 128, 32, subgroup_size_16, 64, 1, 4, 2, 1, subgroup_size_16 }; + l_warptile = { 256, 128, 128, 16, mm_warp_8, 64, 2, tm_m, tn_m, tk_m, mm_warp_8 }; + l_warptile_mmq = l_warptile_mmq_int = { 256, 128, 128, 32, mm_warp_8, 64, 2, tm_m, tn_m, tk_m, mm_warp_8 }; + l_warptile_mmq_int_k = { 256, 128, 128, 32, mm_warp_16, 64, 1, 4, 2, 1, mm_warp_16 }; } else if (device->vendor_id == VK_VENDOR_ID_INTEL && device->coopmat_support) { // Xe2/Xe3 with coopmat enabled - warptile performance tuning - l_warptile = { 512, 128, 128, 16, subgroup_size_8, 32, 2, tm_m, tn_m, tk_m, subgroup_size_8 }; - l_warptile_mmq = { 512, 128, 128, 32, subgroup_size_8, 32, 2, tm_m, tn_m, tk_m, subgroup_size_8 }; + l_warptile = { 512, 128, 128, 16, mm_warp_8, 32, 2, tm_m, tn_m, tk_m, mm_warp_8 }; + l_warptile_mmq = { 512, 128, 128, 32, mm_warp_8, 32, 2, tm_m, tn_m, tk_m, mm_warp_8 }; } l_mmq_wg_denoms = l_wg_denoms = {128, 128, 1 }; @@ -5174,8 +5180,8 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { const uint32_t s_warptile_wm = device->subgroup_size == 8 ? 8 : 32; // use scalar tile sizes - l_warptile = { 128, 128, 128, 16, subgroup_size_8 * 2, 64, 2, 4, 4, 1, subgroup_size_8 }; - m_warptile = { 128, 64, 64, 16, subgroup_size_8, 32, 2, 4, 2, 1, subgroup_size_8 }; + l_warptile = { 128, 128, 128, 16, mm_warp_8 * 2, 64, 2, 4, 4, 1, mm_warp_8 }; + m_warptile = { 128, 64, 64, 16, mm_warp_8, 32, 2, 4, 2, 1, mm_warp_8 }; s_warptile = { subgroup_size_32, 32, 32, 16, s_warptile_wm, 32, 2, 2, 2, 1, subgroup_size_8 }; l_wg_denoms = {128, 128, 1 }; From 3fea10db5435faa9f19d95adedbe6963b2d0b45d Mon Sep 17 00:00:00 2001 From: Max Krasnyansky Date: Wed, 26 Aug 2026 18:46:50 -0700 Subject: [PATCH 010/104] hexagon: support for multi-NPU devices (IQ9, IQ10) and fully asynchronous backend (llama/26501) * hexagon: use non-host bufs by default and make the backend fully async * hex-hb: remove optional hostbuf support and fix async copy * hex-unary: relax supported unary check * hex-bufs: use same get_alignment for host bufs * snapdragon: bump android_platform to 34 * hex-rows: super hacky get/set rows for q8_0 * hex-get-rows: fix q8_0 * hex-get-rows: supprot for f16 and cleanup for q8_0 * hex-get-rows: generic macros and specialized thread funcs * hex-get-rows: add DMA pipeline, vtcm_layout and kernel params * hex-set-rows: fix q8_0 support, add dma and tracing * hex-tests: override nmse threshold for HTP of Q8_0 quants * hex-fa: add support for Q8_0 with inplace dequantizers * hex-get-rows: simplify type dispatch * hex-rows: simplify GET/SET_ROWS DMA pipeline * hex-async: add events, set/get-tensor-async and rest of the async api support * hex-repack: use slice instead of expert in repack functions * hex-cpy: update event/async-cpy logging * hex-set-rows: optimize smaller tensors * hex-geglu: fix perf regression with larger tensors * hex-get-rows: add missing header * hex-set-rows: add missing header * hex-bufs: ressurect GGML_HEXAGON_HOSTBUF but disable it by default * hexagon: do not reject ops with non-heaxon buffers * hex-get-rows: apply >=32 restriction only for q8_0 * hex-res: bump vtcm acquire timeout to 10 seconds * hex-bufs: add support for cloning buffers between sessions to speed up tensor copies * hex-async: rework event recording and batch flushing and integrate with meta backend * hex-bufs: improved handling of repacked tensors * hex-repack: handle get_tensor_2d offsets * hex-dev: add support for devices with multiple NPUs * hex-sync: add support for sync tokens to synchronize npu devices for async splits * hex-mmap: cleanup mmap calls and add a retry for robustness * hex-sync: add failsafe if sync wait gets stuck * hex-sync: use sync_seq to check for completed events * hex-sync: rotate tokens for extra robustness * hex-devs: add supprot for legacy device names for now * hex-bufs: add support for auto-cloning buffers from diff sessions * hex-fusion: simplify and optimize htp-opnode fusion handling * hex-sync: override opnode name so that it shows up in the profiles * hex-trace: update scripts to handle multiple devices * hex-sync: bump the size of the opbatch queue and number of sync tokens * hex-cpy-sync: do not explicitly flush opbatches in cpy_tensor_async and add support for cpy-dma * hex-sync: add graph-flush threshold to avoid single op batches * hex-sync: add sync_peer so that we can flush peers we depend on during cross-device ops * hex-bufs: introduce tensor->extra and shadow_bufs for repacking * hex-l2: flush tiny tensors inline * hex-sync: use explicit l2flush for sync tokens * hex-extra: track weight flags via tensor extra * hex-fence: rename sync to fence * hex-repack: proper handling of set-tensor-2d in the shadow_buf * hex-trace: remove obsolete opstage mask that we used for profiling * hex-env: remove obsolete use_hmx variable * hexagon: new unified run.py and build.py and updated docs * snapdragon: update run script to auto-escapt test-backend-op -p argument * hex-scripts: fix trailing spaces * hex-scripts: fix flake8 warnings * snapdragon: cleanup dst lib/bin dirs before copying new build * hex-ops: add support for allreduce * hex-ar: improved allreduce with dma pipeline * hex-ar: align macros * hex-ar: consistent use of fence_seq * hex-ar: add AR_SELECT env var to select ALLREDUCE kernel or fallback * hex-ar: add proper synchronize handling for ALLREDUCE * hex-opbatch: looks like we now just rely on backend.synchronise to flush the batches, no need to flush them by threshold * hex-ar: bump block size to improve dma efficiency * hex-ar: fused ALLREDUCE+ADD * hex-ar: cleaner fence buffer management * hex-ar: futher allreduce tweaking to remove race conditions * hex-ar: add simple solver and remove non-dma kernels * hex-ar: add row-broadcast to fuse with bias ADD * hex-fence: pass seq numbers via op_params * hex-ar: allow for both entry/exit seq for completing entry wait * hex-ar: align macros * hex-ar: do not refetch broadcast row * hex-fusion: move all fusion into opbatch::add_op for consistency with ALLREDUCE and things * hex-fusion: fix incorrect MUL_MAT reordering * hex-mm: make fused 2x and 3x matmuls more generic * hex-fusion: move tensor fusion tagging to graph_compute * hexagon: make sure to copy tensor->extra by value * hex-get-rows: fix offset calc with row-chunking * hex-repack: get_tensor_2d fixes for non-zero offsets * snapdragon: make profile/trace scripts more robust and donot mix stdout/stderr by default * hex-devices: use legacy device nameing by default to ease the transition * hex-devices: hardcode CDSP domain IDs for current devices for now * hex-optrace: improve multi-NPU timestamp alignment and overall handling of cycle values * hex-optrace: more robust handling of the fence events --- ggml/src/ggml-hexagon/ggml-hexagon.cpp | 3044 +++++++++++++++----- ggml/src/ggml-hexagon/htp-opnode.h | 175 +- ggml/src/ggml-hexagon/htp/CMakeLists.txt | 1 + ggml/src/ggml-hexagon/htp/act-ops.c | 154 +- ggml/src/ggml-hexagon/htp/allreduce-ops.c | 398 +++ ggml/src/ggml-hexagon/htp/allreduce-ops.h | 40 + ggml/src/ggml-hexagon/htp/cpy-ops.c | 73 +- ggml/src/ggml-hexagon/htp/dma-queue.h | 7 +- ggml/src/ggml-hexagon/htp/flash-attn-ops.c | 56 +- ggml/src/ggml-hexagon/htp/get-rows-ops.c | 302 +- ggml/src/ggml-hexagon/htp/get-rows-ops.h | 77 + ggml/src/ggml-hexagon/htp/hex-utils.h | 15 +- ggml/src/ggml-hexagon/htp/htp-ctx.h | 4 +- ggml/src/ggml-hexagon/htp/htp-ops.h | 25 +- ggml/src/ggml-hexagon/htp/htp-tensor.c | 11 +- ggml/src/ggml-hexagon/htp/htp-tensor.h | 9 + ggml/src/ggml-hexagon/htp/hvx-arith.h | 22 +- ggml/src/ggml-hexagon/htp/hvx-quant.h | 165 ++ ggml/src/ggml-hexagon/htp/main.c | 134 +- ggml/src/ggml-hexagon/htp/matmul-ops.c | 871 +----- ggml/src/ggml-hexagon/htp/matmul-ops.h | 43 +- ggml/src/ggml-hexagon/htp/set-rows-ops.c | 262 +- ggml/src/ggml-hexagon/htp/set-rows-ops.h | 74 + 23 files changed, 3908 insertions(+), 2054 deletions(-) create mode 100644 ggml/src/ggml-hexagon/htp/allreduce-ops.c create mode 100644 ggml/src/ggml-hexagon/htp/allreduce-ops.h create mode 100644 ggml/src/ggml-hexagon/htp/get-rows-ops.h create mode 100644 ggml/src/ggml-hexagon/htp/hvx-quant.h create mode 100644 ggml/src/ggml-hexagon/htp/set-rows-ops.h diff --git a/ggml/src/ggml-hexagon/ggml-hexagon.cpp b/ggml/src/ggml-hexagon/ggml-hexagon.cpp index e8a5009b381..c1e9f919d92 100644 --- a/ggml/src/ggml-hexagon/ggml-hexagon.cpp +++ b/ggml/src/ggml-hexagon/ggml-hexagon.cpp @@ -6,6 +6,7 @@ #include #include +#include #include #include #include @@ -18,6 +19,7 @@ #include #include #include +#include #include #ifdef _WIN32 @@ -52,6 +54,8 @@ #include "htp/matmul-ops.h" #include "htp/flash-attn-ops.h" #include "htp/unary-ops.h" +#include "htp/get-rows-ops.h" +#include "htp/set-rows-ops.h" #include "htp_iface.h" #include "htp-drv.h" @@ -59,6 +63,36 @@ using intvec = std::vector; using uintvec = std::vector; using u32vec = std::vector; +#define GGML_HEXAGON_MAX_SESSIONS 16 + +#define GGML_HEXAGON_FENCE_BUFFER_SIZE 8192 +#define GGML_HEXAGON_FENCE_SLOT_SIZE 128 + +struct ggml_hexagon_device_config { + int physical_idx = 0; + int virtual_idx = 0; + std::string name; +}; + +static ggml_hexagon_device_config opt_device_configs[GGML_HEXAGON_MAX_SESSIONS]; + +static int get_domain_id(int physical_idx) { + switch (physical_idx) { + case 0: return 3; // CDSP0 (all devices) + case 1: return 4; // CDSP1 (IQ9, IQ10) + case 2: return 18; // CDSP2 (IQ10) + case 3: return 19; // CDSP3 (IQ10) + default: return CDSP_DOMAIN_ID + physical_idx; + } +} + +static std::string get_domain_name(int physical_idx) { + if (physical_idx == 0) { + return CDSP_DOMAIN_NAME; + } + return std::string("cdsp") + std::to_string(physical_idx); +} + static int opt_arch = 0; // autodetect static size_t opt_ndev = 1; static size_t opt_nhvx = 0; // use all @@ -68,20 +102,19 @@ static size_t opt_mbuf = 1ul * 1024 * 1024 * 1024; // max buffer size static int opt_etm = 0; static int opt_verbose = 0; static int opt_profile = 0; // profiling mode (0-disabled, 1-basic, 2-pmu) -static int opt_hostbuf = 1; // hostbuf ON by default +static bool opt_hostbuf = false; static int opt_mm_select = 3; // 3 = HMX -> Tiled -> Flat -> CPU, 2 = Tiled -> Flat -> CPU, 1 = Flat -> CPU static int opt_fa_select = 2; // 2 = HMX -> HVX -> CPU, 1 = HVX -> CPU, 0 = CPU (unsupported) +static int opt_ar_select = 2; // 2 = fused ALLREDUCE+ADD (DMA, default), 1 = unfused ALLREDUCE (DMA), 0 = fallback to CPY+FENCE // Default PMU events, if profiling with PMU (mode=2) is enabled // See https://docs.qualcomm.com/doc/80-N2040-60/topic/pmu-events.html // https://docs.qualcomm.com/doc/80-N2040-61/topic/hvx-pmu-events.html static u32vec opt_pmu_evt { 0x3, 0x111, 0x100, 0x105, 0x240, 0x256, 0x7D, 0x8C }; -// Enable all stages by default -static int opt_opstage = HTP_OPSTAGE_QUEUE | HTP_OPSTAGE_COMPUTE; static int opt_opbatch = 1024; // max number of ops in a batch -static int opt_opqueue = 16; // max number of pending batches +static int opt_opqueue = 64; // max number of pending batches static int opt_optrace = 0; // trace buffer size per thread (0 means default) static int opt_oppoll = 0; // polling for batch completions static int opt_opfusion = 1; // enable/disable op fusion @@ -121,7 +154,7 @@ static void ggml_hexagon_dump_op_exec(const std::string &sess_name, const htp_op static void ggml_hexagon_dump_op_supp(const std::string &sess_name, const struct ggml_tensor * op, bool supp) { if (!opt_verbose) return; - htp_opformat fmt(htp_opformat(htp_opnode{const_cast(op), {}, HTP_OP_INVALID})); + htp_opformat fmt(htp_opformat(htp_opnode(HTP_OP_INVALID, const_cast(op)))); GGML_LOG_DEBUG("ggml-hex: %s supports-op %s|%s|%s|%s|%s|%s|%s\n", sess_name.c_str(), ggml_op_desc(op), fmt.names, fmt.dims, fmt.types, fmt.strides, fmt.buffs, supp ? "yes" : "no"); } @@ -144,6 +177,7 @@ static const char * htp_event_name(uint16_t id) { case HTP_TRACE_EVT_L2FLUSH: return "L2FLUSH"; case HTP_TRACE_EVT_INIT: return "INIT"; case HTP_TRACE_EVT_BUFF: return "BUFF"; + case HTP_TRACE_EVT_FENCE: return "FENCE"; default: return "UNKNOWN"; } } @@ -205,7 +239,12 @@ static void ggml_hexagon_dump_trace_events(const std::string & sess_name, const } } -// ** +enum ggml_hexagon_tensor_flags { + GGML_HEXAGON_TENSOR_REPACK = (1 << 0), + GGML_HEXAGON_TENSOR_WEIGHT = (1 << 1), + GGML_HEXAGON_TENSOR_FENCE = (1 << 2), + GGML_HEXAGON_TENSOR_FUSEABLE = (1 << 3), +}; static inline bool ggml_hexagon_is_repack_type(enum ggml_type type) { return type == GGML_TYPE_Q4_0 || type == GGML_TYPE_Q4_1 || @@ -227,6 +266,15 @@ static void ggml_hexagon_precompute_matmul_params( struct htp_mm_kernel_params * kparams ); +static void ggml_hexagon_precompute_fused_matmul_add_params( + const struct ggml_hexagon_session * sess, + const struct ggml_tensor * src0, + const struct ggml_tensor * src1, + const struct ggml_tensor * src2, + const struct ggml_tensor * dst, + struct htp_mm_kernel_params * kparams +); + static void ggml_hexagon_precompute_unary_params( const struct ggml_hexagon_session * sess, uint32_t op, @@ -236,25 +284,75 @@ static void ggml_hexagon_precompute_unary_params( struct htp_unary_kernel_params * kparams ); -static void ggml_hexagon_precompute_fused_qkv_params( +static void ggml_hexagon_precompute_get_rows_params( const struct ggml_hexagon_session * sess, const struct ggml_tensor * src0, const struct ggml_tensor * src1, - struct htp_mm_kernel_params * kparams + const struct ggml_tensor * dst, + struct htp_get_rows_kernel_params * kparams ); -static void ggml_hexagon_precompute_fused_ffn_params( +static void ggml_hexagon_precompute_set_rows_params( const struct ggml_hexagon_session * sess, const struct ggml_tensor * src0, const struct ggml_tensor * src1, + const struct ggml_tensor * dst, + struct htp_set_rows_kernel_params * kparams +); + +static void ggml_hexagon_precompute_fused_mmnx_params( + const struct ggml_hexagon_session * sess, + const struct ggml_tensor * src0, + const struct ggml_tensor * src1, + int32_t n_weights, struct htp_mm_kernel_params * kparams ); +static bool ggml_hexagon_precompute_allreduce_params( + const struct ggml_hexagon_session * sess, + const struct ggml_tensor * dst, + uint32_t rank, + uint32_t n_ranks, + bool has_add, + bool is_row_bcast, + struct htp_allreduce_kernel_params * kparams +); + +static bool mm_is_hmx_eligible(const ggml_tensor * t); +static bool is_mergeable_mul_mat(const ggml_tensor * t); +static bool is_mergeable_mul_mat_pair(const ggml_tensor * n1, const ggml_tensor * n2); + // ** backend sessions +struct ggml_hexagon_tensor_extra { + std::vector shadow_buf; + size_t shadow_size { 0 }; + uint32_t flags { 0 }; +}; + +static inline bool ggml_hexagon_tensor_is_fuseable(const struct ggml_tensor * t) { + if (!t || !t->extra) return false; + auto extra = (const struct ggml_hexagon_tensor_extra *) t->extra; + return (extra->flags & GGML_HEXAGON_TENSOR_FUSEABLE) != 0; +} + +struct htp_opnode; + struct ggml_hexagon_opbatch; struct ggml_hexagon_opqueue; -struct htp_opnode; +struct ggml_hexagon_shared_buffer; +struct ggml_hexagon_session; + +struct ggml_backend_hexagon_comm_context { + std::vector backends; + size_t n_backends = 0; + uint32_t fence_seq = 0; +}; + +struct ggml_hexagon_event { + ggml_hexagon_session * sess = nullptr; + uint64_t seq = 0; +}; struct ggml_hexagon_session { std::string name; @@ -264,6 +362,8 @@ struct ggml_hexagon_session { uint32_t domain_id; uint64_t queue_id; int dev_id; + int phys_idx; + int virt_idx; bool valid_session; bool valid_handle; bool valid_queue; @@ -273,20 +373,24 @@ struct ggml_hexagon_session { ggml_hexagon_opbatch* op_batch; ggml_hexagon_opqueue* op_queue; + std::unordered_map> cloned_buffers; + std::unordered_set sync_peers; + ggml_backend_buffer_type buffer_type = {}; - ggml_backend_buffer_type repack_buffer_type = {}; + ggml_backend_buffer_type host_buffer_type = {}; - uint32_t n_threads = 0; - uint32_t n_hvx = 0; - uint32_t n_hmx = 0; - uint64_t vtcm_size = 0; - size_t max_vmem = 0; + uint32_t n_threads = 0; + uint32_t n_hvx = 0; + uint32_t n_hmx = 0; + uint64_t vtcm_size = 0; + size_t max_vmem = 0; size_t max_bufsize = 0; + uint32_t fence_seq; - struct { - uint64_t uid = 0; - std::vector htp_nodes; - } cached_graph; + uint64_t cached_uid = 0; + std::vector cached_nodes; + + mutable std::unordered_set needs_repack; ggml_hexagon_session(int dev_id, ggml_backend_dev_t dev) noexcept(false); ~ggml_hexagon_session() noexcept(true); @@ -297,10 +401,31 @@ struct ggml_hexagon_session { void release() noexcept(true); void enqueue_op(const htp_opnode & node); - void flush(bool all = true); + void enqueue_cpy(const ggml_tensor * src, ggml_tensor * dst, const ggml_tensor * sync_tensor = nullptr, uint32_t fence_seq = 0); + void enqueue_fence(const ggml_tensor * sync_tensor, uint32_t fence_seq = 0); + void enqueue_allreduce(const ggml_tensor * dst, const std::vector & src_tensors, const std::vector & sync_tensors, uint32_t rank, uint32_t n_ranks, uint32_t fence_seq_entry = 0, uint32_t fence_seq_exit = 0); + void flush(bool all = true); void flush_pending(bool all = false); - void flush_batch(); + void flush_batch(size_t min_ops = 1); + + uint64_t record_event(); + void wait_event(uint64_t seq); + + bool clone_buffer(const ggml_hexagon_shared_buffer*); + + void add_sync_peer(ggml_hexagon_session * peer) { + sync_peers.insert(peer); + } + + void flush_sync_peers() { + if (sync_peers.empty()) return; + + for (auto * peer : sync_peers) { + peer->flush_batch(); + } + sync_peers.clear(); + } }; // ** backend buffers @@ -315,26 +440,68 @@ struct ggml_backend_hexagon_buffer_type_context { std::string name; }; +struct ggml_hexagon_rpcmem_block { + uint8_t * base = nullptr; + int fd = -1; + size_t size = 0; + + ggml_hexagon_rpcmem_block(size_t size) { + base = (uint8_t *) rpcmem_alloc2(RPCMEM_HEAP_ID_SYSTEM, RPCMEM_DEFAULT_FLAGS, size); + if (!base) { + throw std::runtime_error("ggml-hex: rpcmem_alloc failed"); + } + fd = rpcmem_to_fd(base); + if (fd < 0) { + rpcmem_free(base); + throw std::runtime_error("ggml-hex: rpcmem_to_fd failed"); + } + this->size = size; + } + + ~ggml_hexagon_rpcmem_block() { + if (base) { + rpcmem_free(base); + } + } +}; + struct ggml_hexagon_shared_buffer { - ggml_hexagon_session * sess; - uint8_t * base; - size_t size; - int fd; - bool mapped; - bool pinned; + ggml_hexagon_session * sess; + std::shared_ptr mem; + std::vector tensor_extra; + uint32_t fence_head = 0; + size_t fences_size = 0; + bool mapped; + bool pinned; + + const char * c_name() const { return sess->c_name(); } + uint8_t * base() const { return mem ? mem->base : nullptr; } + size_t size() const { return mem ? mem->size : 0; } + int fd() const { return mem ? mem->fd : -1; } + + uint8_t * alloc_fence() { + if (fences_size == 0) return nullptr; + int max_slots = fences_size / GGML_HEXAGON_FENCE_SLOT_SIZE; + uint32_t slot = (fence_head++) % max_slots; + + size_t guard_offset = size() - fences_size; + uint8_t * fence_ptr = base() + guard_offset + (size_t)slot * GGML_HEXAGON_FENCE_SLOT_SIZE; + return fence_ptr; + } void mmap() { + if (!this->mem) return; fastrpc_map_flags flags = this->pinned ? FASTRPC_MAP_FD : FASTRPC_MAP_FD_DELAYED; - int err = fastrpc_mmap(sess->domain_id, this->fd, (void *) this->base, 0, this->size, flags); + int err = fastrpc_mmap(sess->domain_id, fd(), (void *) base(), 0, size(), flags); if (err != 0) { GGML_LOG_ERROR("ggml-hex: %s buffer mapping failed : domain_id %d size %zu fd %d error 0x%08x\n", sess->c_name(), - sess->domain_id, this->size, this->fd, (unsigned) err); + sess->domain_id, size(), fd(), (unsigned) err); throw std::runtime_error("ggml-hex: fastrpc_mmap failed (see log for details)"); } HEX_VERBOSE("ggml-hex: %s mapped buffer: base %p size %zu fd %d pinned %u\n", - sess->c_name(), (void *) this->base, this->size, this->fd, pinned); + sess->c_name(), (void *) base(), size(), fd(), pinned); this->mapped = true; } @@ -342,66 +509,69 @@ struct ggml_hexagon_shared_buffer { void unmap() { if (!this->mapped) return; - if (!this->pinned) { + if (!this->pinned && mem) { // HTP might still hold a reference, tell it drop it - htp_iface_munmap(sess->handle, this->fd); + htp_iface_munmap(sess->handle, fd()); } - fastrpc_munmap(sess->domain_id, this->fd, (void *) this->base, this->size); + if (mem) { + fastrpc_munmap(sess->domain_id, fd(), (void *) base(), size()); + } HEX_VERBOSE("ggml-hex: %s unmapped buffer: base %p size %zu fd %d\n", sess->c_name(), - (void *) this->base, size, this->fd); + (void *) base(), size(), fd()); this->mapped = false; - this->fd = -1; } void alloc(size_t size) { - if (this->base) return; - - this->base = (uint8_t *) rpcmem_alloc2(RPCMEM_HEAP_ID_SYSTEM, RPCMEM_DEFAULT_FLAGS, size); - if (!this->base) { - GGML_LOG_ERROR("ggml-hex: %s failed to allocate buffer : size %zu\n", sess->c_name(), size); - throw std::runtime_error("ggml-hex: rpcmem_alloc failed (see log for details)"); - } + if (this->mem) return; - this->fd = rpcmem_to_fd(this->base); - if (this->fd < 0) { - GGML_LOG_ERROR("ggml-hex: %s failed to get FD for buffer %p\n", sess->c_name(), (void *) this->base); - throw std::runtime_error("ggml-hex: rpcmem_to_fd failed (see log for details)"); - } - this->size = size; + this->mem = std::make_shared(size); HEX_VERBOSE("ggml-hex: %s allocated buffer: base %p size %zu fd %d pinned %d\n", sess->c_name(), - (void *) this->base, this->size, this->fd, (int) pinned); + (void *) base(), this->size(), fd(), (int) pinned); mmap(); } void free() { - if (!this->base) return; - unmap(); - rpcmem_free(this->base); - - HEX_VERBOSE("ggml-hex: %s freed buffer: base %p size %zu fd %d\n", sess->c_name(), - (void *) this->base, size, this->fd); + // The memory is freed when the shared_ptr refcount drops to 0. + HEX_VERBOSE("ggml-hex: %s release ref on buffer: base %p size %zu fd %d\n", sess->c_name(), + (void *) base(), size(), fd()); + this->mem = nullptr; + } + + ggml_hexagon_shared_buffer(ggml_hexagon_session * sess, size_t size, bool pinned = false, size_t fence_size = 0) { + this->sess = sess; + this->mapped = false; + this->pinned = pinned; + this->fences_size = fence_size; + + // Size adjustment inside the buffer class + size_t guard_offset = (size + 4095) & ~4095; + size_t total_size = guard_offset; + if (fence_size > 0) { + total_size += 4096 + fence_size; + } - this->base = NULL; + alloc(total_size); } - ggml_hexagon_shared_buffer(ggml_hexagon_session * sess, size_t size, bool pinned = false) { - this->sess = sess; - this->size = 0; - this->base = nullptr; - this->fd = -1; - this->mapped = false; - this->pinned = pinned; - - alloc(size); + // Clone constructor for cross-session mapping + ggml_hexagon_shared_buffer(ggml_hexagon_session * sess, const ggml_hexagon_shared_buffer & other) { + this->sess = sess; + this->mem = other.mem; + this->mapped = false; + this->pinned = other.pinned; + this->fences_size = other.fences_size; } ~ggml_hexagon_shared_buffer() { free(); + for (auto * extra : tensor_extra) { + delete extra; + } } }; @@ -416,18 +586,25 @@ static void ggml_backend_hexagon_buffer_free_buffer(ggml_backend_buffer_t buffer static void * ggml_backend_hexagon_buffer_get_base(ggml_backend_buffer_t buffer) { auto sbuf = static_cast(buffer->context); - return sbuf->base; + return sbuf->base(); } static enum ggml_status ggml_backend_hexagon_buffer_init_tensor(ggml_backend_buffer_t buffer, ggml_tensor * tensor) { auto sbuf = static_cast(buffer->context); auto sess = sbuf->sess; - HEX_VERBOSE("ggml-hex: %s init-tensor %s : base %p data %p nbytes %zu usage %d\n", sess->c_name(), - tensor->name, (void *) sbuf->base, tensor->data, ggml_nbytes(tensor), (int) buffer->usage); + HEX_VERBOSE("ggml-hex: %s init-tensor %s : base %p data %p nbytes %zu\n", sess->c_name(), + tensor->name, (void *) sbuf->base(), tensor->data, ggml_nbytes(tensor)); + + auto extra = new ggml_hexagon_tensor_extra(); + sbuf->tensor_extra.push_back(extra); - if (tensor->view_src != NULL && tensor->view_offs == 0) { - return GGML_STATUS_SUCCESS; // nothing to do for the view + tensor->extra = extra; + if (ggml_hexagon_is_repack_type(tensor->type)) { + if (sess->needs_repack.count(tensor)) { + extra->flags |= GGML_HEXAGON_TENSOR_REPACK; + sess->needs_repack.erase(tensor); + } } return GGML_STATUS_SUCCESS; @@ -499,7 +676,7 @@ static void pack_mxfp4_quants(block_mxfp4 * x, const uint8_t * qs, unsigned int } // repack q4_0 data into q4_0_tiled tensor -static void repack_q4_0_tiled(ggml_tensor * t, const void * data, size_t size) { +static void repack_q4_0_tiled(ggml_tensor * t, const void * data, size_t offset, size_t size) { const block_q4_0 * src_matrix = (const block_q4_0 *) data; int64_t ne0 = t->ne[0]; int64_t ne1 = t->ne[1]; @@ -513,46 +690,49 @@ static void repack_q4_0_tiled(ggml_tensor * t, const void * data, size_t size) { const size_t tile_size = HTP_MM_WEIGHT_TILE_SIZE_Q4_0; const size_t matrix_size = n_col_tiles * n_k_tiles * tile_size; - for (int i3 = 0; i3 < ne3; i3++) { - for (int i2 = 0; i2 < ne2; i2++) { - const block_q4_0 * src_expert = src_matrix + (i3 * ne2 + i2) * (ne1 * (ne0 / 32)); - uint8_t * matrix_dst = (uint8_t *) t->data + (i3 * ne2 + i2) * matrix_size; + size_t slice_size = ne1 * ggml_row_size(t->type, ne0); + int64_t start_slice = offset / slice_size; + int64_t end_slice = (offset + size + slice_size - 1) / slice_size; + if (end_slice > ne2 * ne3) { + end_slice = ne2 * ne3; + } - for (int ct = 0; ct < n_col_tiles; ct++) { - for (int kt = 0; kt < n_k_tiles; kt++) { - uint8_t * tile_dst = matrix_dst + (ct * n_k_tiles + kt) * tile_size; + for (int64_t slice_idx = start_slice; slice_idx < end_slice; slice_idx++) { + const block_q4_0 * src_slice = src_matrix + (slice_idx - start_slice) * (ne1 * (ne0 / 32)); + uint8_t * matrix_dst = (uint8_t *) t->data + slice_idx * matrix_size; - uint8_t tile_quants[32][32]; - for (int row = 0; row < 32; row++) { - int64_t r = ct * 32 + row; - if (r < ne1 && kt < ne0 / 32) { - unpack_q4_0_quants(tile_quants[row], &src_expert[r * (ne0 / 32) + kt], 0); - } else { - memset(tile_quants[row], 8, 32); - } - } + for (int ct = 0; ct < n_col_tiles; ct++) { + for (int kt = 0; kt < n_k_tiles; kt++) { + uint8_t * tile_dst = matrix_dst + (ct * n_k_tiles + kt) * tile_size; - for (int cp = 0; cp < 16; cp++) { - for (int row = 0; row < 32; row++) { - tile_dst[cp * 32 + row] = (tile_quants[row][2 * cp + 1] << 4) | tile_quants[row][2 * cp]; - } + uint8_t tile_quants[32][32]; + for (int row = 0; row < 32; row++) { + int64_t r = ct * 32 + row; + if (r < ne1 && kt < ne0 / 32) { + unpack_q4_0_quants(tile_quants[row], &src_slice[r * (ne0 / 32) + kt], 0); + } else { + memset(tile_quants[row], 8, 32); } + } - ggml_half * scale_dst = (ggml_half *)(tile_dst + 512); + for (int cp = 0; cp < 16; cp++) { for (int row = 0; row < 32; row++) { - int64_t r = ct * 32 + row; - scale_dst[row] = (r < ne1 && kt < ne0 / 32) ? src_expert[r * (ne0 / 32) + kt].d : 0; + tile_dst[cp * 32 + row] = (tile_quants[row][2 * cp + 1] << 4) | tile_quants[row][2 * cp]; } } + + ggml_half * scale_dst = (ggml_half *)(tile_dst + 512); + for (int row = 0; row < 32; row++) { + int64_t r = ct * 32 + row; + scale_dst[row] = (r < ne1 && kt < ne0 / 32) ? src_slice[r * (ne0 / 32) + kt].d : 0; + } } } } - - GGML_UNUSED(size); } // repack q4_0_tiled tensor into q4_0 data -static void repack_tiled_q4_0(void * data, const ggml_tensor * t, size_t size) { +static void repack_tiled_q4_0(void * data, const ggml_tensor * t, size_t offset, size_t size) { block_q4_0 * dst_matrix = (block_q4_0 *) data; int64_t ne0 = t->ne[0]; int64_t ne1 = t->ne[1]; @@ -566,48 +746,65 @@ static void repack_tiled_q4_0(void * data, const ggml_tensor * t, size_t size) { const size_t tile_size = HTP_MM_WEIGHT_TILE_SIZE_Q4_0; const size_t matrix_size = n_col_tiles * n_k_tiles * tile_size; - for (int i3 = 0; i3 < ne3; i3++) { - for (int i2 = 0; i2 < ne2; i2++) { - block_q4_0 * dst_expert = dst_matrix + (i3 * ne2 + i2) * (ne1 * (ne0 / 32)); - const uint8_t * matrix_src = (const uint8_t *) t->data + (i3 * ne2 + i2) * matrix_size; - - for (int ct = 0; ct < n_col_tiles; ct++) { - for (int kt = 0; kt < n_k_tiles; kt++) { - const uint8_t * tile_src = matrix_src + (ct * n_k_tiles + kt) * tile_size; - - uint8_t tile_quants[32][32]; - for (int cp = 0; cp < 16; cp++) { - for (int row = 0; row < 32; row++) { - uint8_t val = tile_src[cp * 32 + row]; - tile_quants[row][2 * cp + 0] = val & 0x0F; - tile_quants[row][2 * cp + 1] = val >> 4; - } - } + size_t slice_size = ne1 * ggml_row_size(t->type, ne0); + size_t row_size_bytes = ggml_row_size(t->type, ne0); + int64_t start_slice = offset / slice_size; + int64_t end_slice = (offset + size + slice_size - 1) / slice_size; + if (end_slice > ne2 * ne3) { + end_slice = ne2 * ne3; + } + + for (int64_t slice_idx = start_slice; slice_idx < end_slice; slice_idx++) { + size_t cur_start_byte = (std::max)(offset, (size_t) slice_idx * slice_size); + size_t cur_end_byte = (std::min)(offset + size, (size_t) (slice_idx + 1) * slice_size); + size_t slice_offset_start = cur_start_byte - (size_t) slice_idx * slice_size; + size_t slice_offset_end = cur_end_byte - (size_t) slice_idx * slice_size; + + int64_t start_row = slice_offset_start / row_size_bytes; + int64_t end_row = (slice_offset_end + row_size_bytes - 1) / row_size_bytes; + end_row = (std::min)(end_row, ne1); + + int start_ct = start_row / 32; + int end_ct = (end_row + 31) / 32; + end_ct = (std::min)(end_ct, n_col_tiles); + + block_q4_0 * dst_slice = dst_matrix + (cur_start_byte - offset) / sizeof(block_q4_0); + const uint8_t * matrix_src = (const uint8_t *) t->data + slice_idx * matrix_size; + for (int ct = start_ct; ct < end_ct; ct++) { + for (int kt = 0; kt < n_k_tiles; kt++) { + const uint8_t * tile_src = matrix_src + (ct * n_k_tiles + kt) * tile_size; + + uint8_t tile_quants[32][32]; + for (int cp = 0; cp < 16; cp++) { for (int row = 0; row < 32; row++) { - int64_t r = ct * 32 + row; - if (r < ne1 && kt < ne0 / 32) { - pack_q4_0_quants(&dst_expert[r * (ne0 / 32) + kt], tile_quants[row], 0); - } + uint8_t val = tile_src[cp * 32 + row]; + tile_quants[row][2 * cp + 0] = val & 0x0F; + tile_quants[row][2 * cp + 1] = val >> 4; } + } - const ggml_half * scale_src = (const ggml_half *)(tile_src + 512); - for (int row = 0; row < 32; row++) { - int64_t r = ct * 32 + row; - if (r < ne1 && kt < ne0 / 32) { - dst_expert[r * (ne0 / 32) + kt].d = scale_src[row]; - } + for (int row = 0; row < 32; row++) { + int64_t r = ct * 32 + row; + if (r >= start_row && r < end_row && kt < ne0 / 32) { + pack_q4_0_quants(&dst_slice[(r - start_row) * (ne0 / 32) + kt], tile_quants[row], 0); + } + } + + const ggml_half * scale_src = (const ggml_half *)(tile_src + 512); + for (int row = 0; row < 32; row++) { + int64_t r = ct * 32 + row; + if (r >= start_row && r < end_row && kt < ne0 / 32) { + dst_slice[(r - start_row) * (ne0 / 32) + kt].d = scale_src[row]; } } } } } - - GGML_UNUSED(size); } // repack q4_1 data into q4_1_tiled tensor -static void repack_q4_1_tiled(ggml_tensor * t, const void * data, size_t size) { +static void repack_q4_1_tiled(ggml_tensor * t, const void * data, size_t offset, size_t size) { const block_q4_1 * src_matrix = (const block_q4_1 *) data; int64_t ne0 = t->ne[0]; int64_t ne1 = t->ne[1]; @@ -621,52 +818,55 @@ static void repack_q4_1_tiled(ggml_tensor * t, const void * data, size_t size) { const size_t tile_size = HTP_MM_WEIGHT_TILE_SIZE_Q4_1; const size_t matrix_size = n_col_tiles * n_k_tiles * tile_size; - for (int i3 = 0; i3 < ne3; i3++) { - for (int i2 = 0; i2 < ne2; i2++) { - const block_q4_1 * src_expert = src_matrix + (i3 * ne2 + i2) * (ne1 * (ne0 / 32)); - uint8_t * matrix_dst = (uint8_t *) t->data + (i3 * ne2 + i2) * matrix_size; + size_t slice_size = ne1 * ggml_row_size(t->type, ne0); + int64_t start_slice = offset / slice_size; + int64_t end_slice = (offset + size + slice_size - 1) / slice_size; + if (end_slice > ne2 * ne3) { + end_slice = ne2 * ne3; + } - for (int ct = 0; ct < n_col_tiles; ct++) { - for (int kt = 0; kt < n_k_tiles; kt++) { - uint8_t * tile_dst = matrix_dst + (ct * n_k_tiles + kt) * tile_size; + for (int64_t slice_idx = start_slice; slice_idx < end_slice; slice_idx++) { + const block_q4_1 * src_slice = src_matrix + (slice_idx - start_slice) * (ne1 * (ne0 / 32)); + uint8_t * matrix_dst = (uint8_t *) t->data + slice_idx * matrix_size; - uint8_t tile_quants[32][32]; - for (int row = 0; row < 32; row++) { - int64_t r = ct * 32 + row; - if (r < ne1 && kt < ne0 / 32) { - unpack_q4_1_quants(tile_quants[row], &src_expert[r * (ne0 / 32) + kt], 0); - } else { - memset(tile_quants[row], 0, 32); - } - } + for (int ct = 0; ct < n_col_tiles; ct++) { + for (int kt = 0; kt < n_k_tiles; kt++) { + uint8_t * tile_dst = matrix_dst + (ct * n_k_tiles + kt) * tile_size; - for (int cp = 0; cp < 16; cp++) { - for (int row = 0; row < 32; row++) { - tile_dst[cp * 32 + row] = (tile_quants[row][2 * cp + 1] << 4) | tile_quants[row][2 * cp]; - } + uint8_t tile_quants[32][32]; + for (int row = 0; row < 32; row++) { + int64_t r = ct * 32 + row; + if (r < ne1 && kt < ne0 / 32) { + unpack_q4_1_quants(tile_quants[row], &src_slice[r * (ne0 / 32) + kt], 0); + } else { + memset(tile_quants[row], 0, 32); } + } - ggml_half * scale_dst = (ggml_half *)(tile_dst + 512); + for (int cp = 0; cp < 16; cp++) { for (int row = 0; row < 32; row++) { - int64_t r = ct * 32 + row; - if (r < ne1 && kt < ne0 / 32) { - scale_dst[2 * row + 0] = src_expert[r * (ne0 / 32) + kt].d; - scale_dst[2 * row + 1] = src_expert[r * (ne0 / 32) + kt].m; - } else { - scale_dst[2 * row + 0] = 0; - scale_dst[2 * row + 1] = 0; - } + tile_dst[cp * 32 + row] = (tile_quants[row][2 * cp + 1] << 4) | tile_quants[row][2 * cp]; + } + } + + ggml_half * scale_dst = (ggml_half *)(tile_dst + 512); + for (int row = 0; row < 32; row++) { + int64_t r = ct * 32 + row; + if (r < ne1 && kt < ne0 / 32) { + scale_dst[2 * row + 0] = src_slice[r * (ne0 / 32) + kt].d; + scale_dst[2 * row + 1] = src_slice[r * (ne0 / 32) + kt].m; + } else { + scale_dst[2 * row + 0] = 0; + scale_dst[2 * row + 1] = 0; } } } } } - - GGML_UNUSED(size); } // repack q4_1_tiled tensor into q4_1 data -static void repack_tiled_q4_1(void * data, const ggml_tensor * t, size_t size) { +static void repack_tiled_q4_1(void * data, const ggml_tensor * t, size_t offset, size_t size) { block_q4_1 * dst_matrix = (block_q4_1 *) data; int64_t ne0 = t->ne[0]; int64_t ne1 = t->ne[1]; @@ -680,49 +880,66 @@ static void repack_tiled_q4_1(void * data, const ggml_tensor * t, size_t size) { const size_t tile_size = HTP_MM_WEIGHT_TILE_SIZE_Q4_1; const size_t matrix_size = n_col_tiles * n_k_tiles * tile_size; - for (int i3 = 0; i3 < ne3; i3++) { - for (int i2 = 0; i2 < ne2; i2++) { - block_q4_1 * dst_expert = dst_matrix + (i3 * ne2 + i2) * (ne1 * (ne0 / 32)); - const uint8_t * matrix_src = (const uint8_t *) t->data + (i3 * ne2 + i2) * matrix_size; - - for (int ct = 0; ct < n_col_tiles; ct++) { - for (int kt = 0; kt < n_k_tiles; kt++) { - const uint8_t * tile_src = matrix_src + (ct * n_k_tiles + kt) * tile_size; - - uint8_t tile_quants[32][32]; - for (int cp = 0; cp < 16; cp++) { - for (int row = 0; row < 32; row++) { - uint8_t val = tile_src[cp * 32 + row]; - tile_quants[row][2 * cp + 0] = val & 0x0F; - tile_quants[row][2 * cp + 1] = val >> 4; - } - } + size_t slice_size = ne1 * ggml_row_size(t->type, ne0); + size_t row_size_bytes = ggml_row_size(t->type, ne0); + int64_t start_slice = offset / slice_size; + int64_t end_slice = (offset + size + slice_size - 1) / slice_size; + if (end_slice > ne2 * ne3) { + end_slice = ne2 * ne3; + } + for (int64_t slice_idx = start_slice; slice_idx < end_slice; slice_idx++) { + size_t cur_start_byte = (std::max)(offset, (size_t) slice_idx * slice_size); + size_t cur_end_byte = (std::min)(offset + size, (size_t) (slice_idx + 1) * slice_size); + size_t slice_offset_start = cur_start_byte - (size_t) slice_idx * slice_size; + size_t slice_offset_end = cur_end_byte - (size_t) slice_idx * slice_size; + + int64_t start_row = slice_offset_start / row_size_bytes; + int64_t end_row = (slice_offset_end + row_size_bytes - 1) / row_size_bytes; + end_row = (std::min)(end_row, ne1); + + int start_ct = start_row / 32; + int end_ct = (end_row + 31) / 32; + end_ct = (std::min)(end_ct, n_col_tiles); + + block_q4_1 * dst_slice = dst_matrix + (cur_start_byte - offset) / sizeof(block_q4_1); + const uint8_t * matrix_src = (const uint8_t *) t->data + slice_idx * matrix_size; + + for (int ct = start_ct; ct < end_ct; ct++) { + for (int kt = 0; kt < n_k_tiles; kt++) { + const uint8_t * tile_src = matrix_src + (ct * n_k_tiles + kt) * tile_size; + + uint8_t tile_quants[32][32]; + for (int cp = 0; cp < 16; cp++) { for (int row = 0; row < 32; row++) { - int64_t r = ct * 32 + row; - if (r < ne1 && kt < ne0 / 32) { - pack_q4_1_quants(&dst_expert[r * (ne0 / 32) + kt], tile_quants[row], 0); - } + uint8_t val = tile_src[cp * 32 + row]; + tile_quants[row][2 * cp + 0] = val & 0x0F; + tile_quants[row][2 * cp + 1] = val >> 4; } + } - const ggml_half * scale_src = (const ggml_half *)(tile_src + 512); - for (int row = 0; row < 32; row++) { - int64_t r = ct * 32 + row; - if (r < ne1 && kt < ne0 / 32) { - dst_expert[r * (ne0 / 32) + kt].d = scale_src[2 * row]; - dst_expert[r * (ne0 / 32) + kt].m = scale_src[2 * row + 1]; - } + for (int row = 0; row < 32; row++) { + int64_t r = ct * 32 + row; + if (r >= start_row && r < end_row && kt < ne0 / 32) { + pack_q4_1_quants(&dst_slice[(r - start_row) * (ne0 / 32) + kt], tile_quants[row], 0); + } + } + + const ggml_half * scale_src = (const ggml_half *)(tile_src + 512); + for (int row = 0; row < 32; row++) { + int64_t r = ct * 32 + row; + if (r >= start_row && r < end_row && kt < ne0 / 32) { + dst_slice[(r - start_row) * (ne0 / 32) + kt].d = scale_src[2 * row]; + dst_slice[(r - start_row) * (ne0 / 32) + kt].m = scale_src[2 * row + 1]; } } } } } - - GGML_UNUSED(size); } // repack q8_0 data into q8_0_tiled tensor -static void repack_q8_0_tiled(ggml_tensor * t, const void * data, size_t size) { +static void repack_q8_0_tiled(ggml_tensor * t, const void * data, size_t offset, size_t size) { const block_q8_0 * src_matrix = (const block_q8_0 *) data; int64_t ne0 = t->ne[0]; int64_t ne1 = t->ne[1]; @@ -736,41 +953,44 @@ static void repack_q8_0_tiled(ggml_tensor * t, const void * data, size_t size) { const size_t tile_size = HTP_MM_WEIGHT_TILE_SIZE_Q8_0; const size_t matrix_size = n_col_tiles * n_k_tiles * tile_size; - for (int i3 = 0; i3 < ne3; i3++) { - for (int i2 = 0; i2 < ne2; i2++) { - const block_q8_0 * src_expert = src_matrix + (i3 * ne2 + i2) * (ne1 * (ne0 / 32)); - uint8_t * matrix_dst = (uint8_t *) t->data + (i3 * ne2 + i2) * matrix_size; - - for (int ct = 0; ct < n_col_tiles; ct++) { - for (int kt = 0; kt < n_k_tiles; kt++) { - uint8_t * tile_dst = matrix_dst + (ct * n_k_tiles + kt) * tile_size; - - for (int cp = 0; cp < 16; cp++) { - int col0 = cp * 2; - int col1 = col0 + 1; - for (int row = 0; row < 32; row++) { - int64_t r = ct * 32 + row; - const block_q8_0 * b = (r < ne1 && kt < ne0 / 32) ? &src_expert[r * (ne0 / 32) + kt] : NULL; - tile_dst[cp * 64 + 2 * row + 0] = b ? b->qs[col0] : 0; - tile_dst[cp * 64 + 2 * row + 1] = b ? b->qs[col1] : 0; - } - } + size_t slice_size = ne1 * ggml_row_size(t->type, ne0); + int64_t start_slice = offset / slice_size; + int64_t end_slice = (offset + size + slice_size - 1) / slice_size; + if (end_slice > ne2 * ne3) { + end_slice = ne2 * ne3; + } + + for (int64_t slice_idx = start_slice; slice_idx < end_slice; slice_idx++) { + const block_q8_0 * src_slice = src_matrix + (slice_idx - start_slice) * (ne1 * (ne0 / 32)); + uint8_t * matrix_dst = (uint8_t *) t->data + slice_idx * matrix_size; - ggml_half * scale_dst = (ggml_half *)(tile_dst + 1024); + for (int ct = 0; ct < n_col_tiles; ct++) { + for (int kt = 0; kt < n_k_tiles; kt++) { + uint8_t * tile_dst = matrix_dst + (ct * n_k_tiles + kt) * tile_size; + + for (int cp = 0; cp < 16; cp++) { + int col0 = cp * 2; + int col1 = col0 + 1; for (int row = 0; row < 32; row++) { int64_t r = ct * 32 + row; - scale_dst[row] = (r < ne1 && kt < ne0 / 32) ? src_expert[r * (ne0 / 32) + kt].d : 0; + const block_q8_0 * b = (r < ne1 && kt < ne0 / 32) ? &src_slice[r * (ne0 / 32) + kt] : NULL; + tile_dst[cp * 64 + 2 * row + 0] = b ? b->qs[col0] : 0; + tile_dst[cp * 64 + 2 * row + 1] = b ? b->qs[col1] : 0; } } + + ggml_half * scale_dst = (ggml_half *)(tile_dst + 1024); + for (int row = 0; row < 32; row++) { + int64_t r = ct * 32 + row; + scale_dst[row] = (r < ne1 && kt < ne0 / 32) ? src_slice[r * (ne0 / 32) + kt].d : 0; + } } } } - - GGML_UNUSED(size); } // repack q8_0_tiled tensor into q8_0 data -static void repack_tiled_q8_0(void * data, const ggml_tensor * t, size_t size) { +static void repack_tiled_q8_0(void * data, const ggml_tensor * t, size_t offset, size_t size) { block_q8_0 * dst_matrix = (block_q8_0 *) data; int64_t ne0 = t->ne[0]; int64_t ne1 = t->ne[1]; @@ -784,45 +1004,62 @@ static void repack_tiled_q8_0(void * data, const ggml_tensor * t, size_t size) { const size_t tile_size = HTP_MM_WEIGHT_TILE_SIZE_Q8_0; const size_t matrix_size = n_col_tiles * n_k_tiles * tile_size; - for (int i3 = 0; i3 < ne3; i3++) { - for (int i2 = 0; i2 < ne2; i2++) { - block_q8_0 * dst_expert = dst_matrix + (i3 * ne2 + i2) * (ne1 * (ne0 / 32)); - const uint8_t * matrix_src = (const uint8_t *) t->data + (i3 * ne2 + i2) * matrix_size; - - for (int ct = 0; ct < n_col_tiles; ct++) { - for (int kt = 0; kt < n_k_tiles; kt++) { - const uint8_t * tile_src = matrix_src + (ct * n_k_tiles + kt) * tile_size; - - for (int cp = 0; cp < 16; cp++) { - int col0 = cp * 2; - int col1 = col0 + 1; - for (int row = 0; row < 32; row++) { - int64_t r = ct * 32 + row; - if (r < ne1 && kt < ne0 / 32) { - block_q8_0 & b = dst_expert[r * (ne0 / 32) + kt]; - b.qs[col0] = tile_src[cp * 64 + 2 * row + 0]; - b.qs[col1] = tile_src[cp * 64 + 2 * row + 1]; - } - } - } + size_t slice_size = ne1 * ggml_row_size(t->type, ne0); + size_t row_size_bytes = ggml_row_size(t->type, ne0); + int64_t start_slice = offset / slice_size; + int64_t end_slice = (offset + size + slice_size - 1) / slice_size; + if (end_slice > ne2 * ne3) { + end_slice = ne2 * ne3; + } + + for (int64_t slice_idx = start_slice; slice_idx < end_slice; slice_idx++) { + size_t cur_start_byte = (std::max)(offset, (size_t) slice_idx * slice_size); + size_t cur_end_byte = (std::min)(offset + size, (size_t) (slice_idx + 1) * slice_size); + size_t slice_offset_start = cur_start_byte - (size_t) slice_idx * slice_size; + size_t slice_offset_end = cur_end_byte - (size_t) slice_idx * slice_size; + + int64_t start_row = slice_offset_start / row_size_bytes; + int64_t end_row = (slice_offset_end + row_size_bytes - 1) / row_size_bytes; + end_row = (std::min)(end_row, ne1); + + int start_ct = start_row / 32; + int end_ct = (end_row + 31) / 32; + end_ct = (std::min)(end_ct, n_col_tiles); + + block_q8_0 * dst_slice = dst_matrix + (cur_start_byte - offset) / sizeof(block_q8_0); + const uint8_t * matrix_src = (const uint8_t *) t->data + slice_idx * matrix_size; - const ggml_half * scale_src = (const ggml_half *)(tile_src + 1024); + for (int ct = start_ct; ct < end_ct; ct++) { + for (int kt = 0; kt < n_k_tiles; kt++) { + const uint8_t * tile_src = matrix_src + (ct * n_k_tiles + kt) * tile_size; + + for (int cp = 0; cp < 16; cp++) { + int col0 = cp * 2; + int col1 = col0 + 1; for (int row = 0; row < 32; row++) { int64_t r = ct * 32 + row; - if (r < ne1 && kt < ne0 / 32) { - dst_expert[r * (ne0 / 32) + kt].d = scale_src[row]; + if (r >= start_row && r < end_row && kt < ne0 / 32) { + block_q8_0 & b = dst_slice[(r - start_row) * (ne0 / 32) + kt]; + b.qs[col0] = tile_src[cp * 64 + 2 * row + 0]; + b.qs[col1] = tile_src[cp * 64 + 2 * row + 1]; } } } + + const ggml_half * scale_src = (const ggml_half *)(tile_src + 1024); + for (int row = 0; row < 32; row++) { + int64_t r = ct * 32 + row; + if (r >= start_row && r < end_row && kt < ne0 / 32) { + dst_slice[(r - start_row) * (ne0 / 32) + kt].d = scale_src[row]; + } + } } } } - - GGML_UNUSED(size); } // repack mxfp4 data into mxfp4_tiled tensor -static void repack_mxfp4_tiled(ggml_tensor * t, const void * data, size_t size) { +static void repack_mxfp4_tiled(ggml_tensor * t, const void * data, size_t offset, size_t size) { const block_mxfp4 * src_matrix = (const block_mxfp4 *) data; int64_t ne0 = t->ne[0]; int64_t ne1 = t->ne[1]; @@ -836,46 +1073,49 @@ static void repack_mxfp4_tiled(ggml_tensor * t, const void * data, size_t size) const size_t tile_size = HTP_MM_WEIGHT_TILE_SIZE_MXFP4; const size_t matrix_size = n_col_tiles * n_k_tiles * tile_size; - for (int i3 = 0; i3 < ne3; i3++) { - for (int i2 = 0; i2 < ne2; i2++) { - const block_mxfp4 * src_expert = src_matrix + (i3 * ne2 + i2) * (ne1 * (ne0 / 32)); - uint8_t * matrix_dst = (uint8_t *) t->data + (i3 * ne2 + i2) * matrix_size; + size_t slice_size = ne1 * ggml_row_size(t->type, ne0); + int64_t start_slice = offset / slice_size; + int64_t end_slice = (offset + size + slice_size - 1) / slice_size; + if (end_slice > ne2 * ne3) { + end_slice = ne2 * ne3; + } - for (int ct = 0; ct < n_col_tiles; ct++) { - for (int kt = 0; kt < n_k_tiles; kt++) { - uint8_t * tile_dst = matrix_dst + (ct * n_k_tiles + kt) * tile_size; + for (int64_t slice_idx = start_slice; slice_idx < end_slice; slice_idx++) { + const block_mxfp4 * src_slice = src_matrix + (slice_idx - start_slice) * (ne1 * (ne0 / 32)); + uint8_t * matrix_dst = (uint8_t *) t->data + slice_idx * matrix_size; - uint8_t tile_quants[32][32]; - for (int row = 0; row < 32; row++) { - int64_t r = ct * 32 + row; - if (r < ne1 && kt < ne0 / 32) { - unpack_mxfp4_quants(tile_quants[row], &src_expert[r * (ne0 / 32) + kt], 0); - } else { - memset(tile_quants[row], 0, 32); - } - } + for (int ct = 0; ct < n_col_tiles; ct++) { + for (int kt = 0; kt < n_k_tiles; kt++) { + uint8_t * tile_dst = matrix_dst + (ct * n_k_tiles + kt) * tile_size; - for (int cp = 0; cp < 16; cp++) { - for (int row = 0; row < 32; row++) { - tile_dst[cp * 32 + row] = (tile_quants[row][2 * cp + 1] << 4) | tile_quants[row][2 * cp]; - } + uint8_t tile_quants[32][32]; + for (int row = 0; row < 32; row++) { + int64_t r = ct * 32 + row; + if (r < ne1 && kt < ne0 / 32) { + unpack_mxfp4_quants(tile_quants[row], &src_slice[r * (ne0 / 32) + kt], 0); + } else { + memset(tile_quants[row], 0, 32); } + } - uint8_t * scale_dst = tile_dst + 512; + for (int cp = 0; cp < 16; cp++) { for (int row = 0; row < 32; row++) { - int64_t r = ct * 32 + row; - scale_dst[row] = (r < ne1 && kt < ne0 / 32) ? src_expert[r * (ne0 / 32) + kt].e : 0; + tile_dst[cp * 32 + row] = (tile_quants[row][2 * cp + 1] << 4) | tile_quants[row][2 * cp]; } } + + uint8_t * scale_dst = tile_dst + 512; + for (int row = 0; row < 32; row++) { + int64_t r = ct * 32 + row; + scale_dst[row] = (r < ne1 && kt < ne0 / 32) ? src_slice[r * (ne0 / 32) + kt].e : 0; + } } } } - - GGML_UNUSED(size); } // repack mxfp4_tiled tensor into mxfp4 data -static void repack_tiled_mxfp4(void * data, const ggml_tensor * t, size_t size) { +static void repack_tiled_mxfp4(void * data, const ggml_tensor * t, size_t offset, size_t size) { block_mxfp4 * dst_matrix = (block_mxfp4 *) data; int64_t ne0 = t->ne[0]; int64_t ne1 = t->ne[1]; @@ -889,133 +1129,179 @@ static void repack_tiled_mxfp4(void * data, const ggml_tensor * t, size_t size) const size_t tile_size = HTP_MM_WEIGHT_TILE_SIZE_MXFP4; const size_t matrix_size = n_col_tiles * n_k_tiles * tile_size; - for (int i3 = 0; i3 < ne3; i3++) { - for (int i2 = 0; i2 < ne2; i2++) { - block_mxfp4 * dst_expert = dst_matrix + (i3 * ne2 + i2) * (ne1 * (ne0 / 32)); - const uint8_t * matrix_src = (const uint8_t *) t->data + (i3 * ne2 + i2) * matrix_size; - - for (int ct = 0; ct < n_col_tiles; ct++) { - for (int kt = 0; kt < n_k_tiles; kt++) { - const uint8_t * tile_src = matrix_src + (ct * n_k_tiles + kt) * tile_size; - - uint8_t tile_quants[32][32]; - for (int cp = 0; cp < 16; cp++) { - for (int row = 0; row < 32; row++) { - uint8_t val = tile_src[cp * 32 + row]; - tile_quants[row][2 * cp + 0] = val & 0x0F; - tile_quants[row][2 * cp + 1] = val >> 4; - } - } + size_t slice_size = ne1 * ggml_row_size(t->type, ne0); + size_t row_size_bytes = ggml_row_size(t->type, ne0); + int64_t start_slice = offset / slice_size; + int64_t end_slice = (offset + size + slice_size - 1) / slice_size; + if (end_slice > ne2 * ne3) { + end_slice = ne2 * ne3; + } + + for (int64_t slice_idx = start_slice; slice_idx < end_slice; slice_idx++) { + size_t cur_start_byte = (std::max)(offset, (size_t) slice_idx * slice_size); + size_t cur_end_byte = (std::min)(offset + size, (size_t) (slice_idx + 1) * slice_size); + size_t slice_offset_start = cur_start_byte - (size_t) slice_idx * slice_size; + size_t slice_offset_end = cur_end_byte - (size_t) slice_idx * slice_size; + + int64_t start_row = slice_offset_start / row_size_bytes; + int64_t end_row = (slice_offset_end + row_size_bytes - 1) / row_size_bytes; + end_row = (std::min)(end_row, ne1); + + int start_ct = start_row / 32; + int end_ct = (end_row + 31) / 32; + end_ct = (std::min)(end_ct, n_col_tiles); + block_mxfp4 * dst_slice = dst_matrix + (cur_start_byte - offset) / sizeof(block_mxfp4); + const uint8_t * matrix_src = (const uint8_t *) t->data + slice_idx * matrix_size; + + for (int ct = start_ct; ct < end_ct; ct++) { + for (int kt = 0; kt < n_k_tiles; kt++) { + const uint8_t * tile_src = matrix_src + (ct * n_k_tiles + kt) * tile_size; + + uint8_t tile_quants[32][32]; + for (int cp = 0; cp < 16; cp++) { for (int row = 0; row < 32; row++) { - int64_t r = ct * 32 + row; - if (r < ne1 && kt < ne0 / 32) { - pack_mxfp4_quants(&dst_expert[r * (ne0 / 32) + kt], tile_quants[row], 0); - } + uint8_t val = tile_src[cp * 32 + row]; + tile_quants[row][2 * cp + 0] = val & 0x0F; + tile_quants[row][2 * cp + 1] = val >> 4; } + } - const uint8_t * scale_src = tile_src + 512; - for (int row = 0; row < 32; row++) { - int64_t r = ct * 32 + row; - if (r < ne1 && kt < ne0 / 32) { - dst_expert[r * (ne0 / 32) + kt].e = scale_src[row]; - } + for (int row = 0; row < 32; row++) { + int64_t r = ct * 32 + row; + if (r >= start_row && r < end_row && kt < ne0 / 32) { + pack_mxfp4_quants(&dst_slice[(r - start_row) * (ne0 / 32) + kt], tile_quants[row], 0); + } + } + + const uint8_t * scale_src = tile_src + 512; + for (int row = 0; row < 32; row++) { + int64_t r = ct * 32 + row; + if (r >= start_row && r < end_row && kt < ne0 / 32) { + dst_slice[(r - start_row) * (ne0 / 32) + kt].e = scale_src[row]; } } } } } - - GGML_UNUSED(size); } -static void ggml_backend_hexagon_buffer_set_tensor(ggml_backend_buffer_t buffer, - ggml_tensor * tensor, - const void * data, - size_t offset, - size_t size) { - auto sbuf = (ggml_hexagon_shared_buffer *) buffer->context; - auto sess = sbuf->sess; - - HEX_VERBOSE("ggml-hex: %s set-tensor %s : data %p offset %zu size %zu\n", sess->c_name(), tensor->name, data, offset, size); - +static void repack_tensor_tiled(ggml_tensor * tensor, const void * data, size_t size) { switch (tensor->type) { case GGML_TYPE_Q4_0: - GGML_ASSERT(offset == 0); - GGML_ASSERT(offset + size <= ggml_nbytes(tensor)); - repack_q4_0_tiled(tensor, data, size); + repack_q4_0_tiled(tensor, data, 0, size); break; case GGML_TYPE_Q4_1: - GGML_ASSERT(offset == 0); - GGML_ASSERT(offset + size <= ggml_nbytes(tensor)); - repack_q4_1_tiled(tensor, data, size); + repack_q4_1_tiled(tensor, data, 0, size); break; case GGML_TYPE_Q8_0: - GGML_ASSERT(offset == 0); - GGML_ASSERT(offset + size <= ggml_nbytes(tensor)); - repack_q8_0_tiled(tensor, data, size); + repack_q8_0_tiled(tensor, data, 0, size); break; case GGML_TYPE_IQ4_NL: - GGML_ASSERT(offset == 0); - GGML_ASSERT(offset + size <= ggml_nbytes(tensor)); - // IQ4_NL has identical block layout to Q4_0 (ggml_half d + uint8_t qs[16]) - repack_q4_0_tiled(tensor, data, size); + repack_q4_0_tiled(tensor, data, 0, size); break; case GGML_TYPE_MXFP4: - GGML_ASSERT(offset == 0); - GGML_ASSERT(offset + size <= ggml_nbytes(tensor)); - repack_mxfp4_tiled(tensor, data, size); + repack_mxfp4_tiled(tensor, data, 0, size); break; default: - memcpy((char *) tensor->data + offset, data, size); break; } } +static void ggml_backend_hexagon_buffer_set_tensor(ggml_backend_buffer_t buffer, + ggml_tensor * tensor, + const void * data, + size_t offset, + size_t size) { + auto extra = (ggml_hexagon_tensor_extra *) tensor->extra; + auto sbuf = (ggml_hexagon_shared_buffer *) buffer->context; + auto sess = sbuf->sess; + + if (ggml_backend_buffer_get_usage(buffer) == GGML_BACKEND_BUFFER_USAGE_WEIGHTS) { + extra->flags |= GGML_HEXAGON_TENSOR_WEIGHT; + if (ggml_hexagon_is_repack_type(tensor->type)) { + extra->flags |= GGML_HEXAGON_TENSOR_REPACK; + } + } + + HEX_VERBOSE("ggml-hex: %s set-tensor %s : data %p offset %zu size %zu usage %d flags 0x%x\n", + sess->c_name(), tensor->name, data, offset, size, (int) buffer->usage, extra->flags); + + if ((extra->flags & GGML_HEXAGON_TENSOR_REPACK) == 0) { + memcpy((char *) tensor->data + offset, data, size); + return; + } + + if (offset == 0 && size == ggml_nbytes(tensor) && extra->shadow_buf.empty()) { + repack_tensor_tiled(tensor, data, size); + return; + } + + if (extra->shadow_buf.size() < ggml_nbytes(tensor)) { + extra->shadow_buf.resize(ggml_nbytes(tensor)); + } + memcpy(extra->shadow_buf.data() + offset, data, size); + extra->shadow_size += size; + + if (extra->shadow_size >= ggml_nbytes(tensor)) { + repack_tensor_tiled(tensor, extra->shadow_buf.data(), extra->shadow_buf.size()); + extra->shadow_buf.clear(); + extra->shadow_buf.shrink_to_fit(); + extra->shadow_size = 0; + } +} + static void ggml_backend_hexagon_buffer_get_tensor(ggml_backend_buffer_t buffer, const ggml_tensor * tensor, void * data, size_t offset, size_t size) { - auto sbuf = (ggml_hexagon_shared_buffer *) buffer->context; - auto sess = sbuf->sess; + auto extra = (ggml_hexagon_tensor_extra *) tensor->extra; + auto sbuf = (ggml_hexagon_shared_buffer *) buffer->context; + auto sess = sbuf->sess; - HEX_VERBOSE("ggml-hex: %s get-tensor %s : data %p offset %zu size %zu\n", sess->c_name(), tensor->name, data, offset, size); + HEX_VERBOSE("ggml-hex: %s get-tensor %s : data %p offset %zu size %zu usage %d flags 0x%x\n", + sess->c_name(), tensor->name, data, offset, size, (int) buffer->usage, extra->flags); + + if ((extra->flags & GGML_HEXAGON_TENSOR_REPACK) == 0) { + memcpy(data, (const char *) tensor->data + offset, size); + return; + } switch (tensor->type) { case GGML_TYPE_Q4_0: GGML_ASSERT(offset == 0); GGML_ASSERT(offset + size <= ggml_nbytes(tensor)); - repack_tiled_q4_0(data, tensor, size); + repack_tiled_q4_0(data, tensor, offset, size); break; case GGML_TYPE_Q4_1: GGML_ASSERT(offset == 0); GGML_ASSERT(offset + size <= ggml_nbytes(tensor)); - repack_tiled_q4_1(data, tensor, size); + repack_tiled_q4_1(data, tensor, offset, size); break; case GGML_TYPE_Q8_0: GGML_ASSERT(offset == 0); GGML_ASSERT(offset + size <= ggml_nbytes(tensor)); - repack_tiled_q8_0(data, tensor, size); + repack_tiled_q8_0(data, tensor, offset, size); break; case GGML_TYPE_IQ4_NL: GGML_ASSERT(offset == 0); GGML_ASSERT(offset + size <= ggml_nbytes(tensor)); - repack_tiled_q4_0(data, tensor, size); + repack_tiled_q4_0(data, tensor, offset, size); break; case GGML_TYPE_MXFP4: GGML_ASSERT(offset == 0); GGML_ASSERT(offset + size <= ggml_nbytes(tensor)); - repack_tiled_mxfp4(data, tensor, size); + repack_tiled_mxfp4(data, tensor, offset, size); break; default: @@ -1035,11 +1321,121 @@ static bool ggml_backend_hexagon_buffer_cpy_tensor(ggml_backend_buffer_t bu GGML_UNUSED(dst); } +static void ggml_backend_hexagon_buffer_set_tensor_2d(ggml_backend_buffer_t buffer, + ggml_tensor * tensor, + const void * data, + size_t offset, + size_t size, + size_t n_copies, + size_t stride_tensor, + size_t stride_data) { + auto extra = (ggml_hexagon_tensor_extra *) tensor->extra; + auto sbuf = (ggml_hexagon_shared_buffer *) buffer->context; + auto sess = sbuf->sess; + + if (ggml_backend_buffer_get_usage(buffer) == GGML_BACKEND_BUFFER_USAGE_WEIGHTS) { + extra->flags |= GGML_HEXAGON_TENSOR_WEIGHT; + if (ggml_hexagon_is_repack_type(tensor->type)) { + extra->flags |= GGML_HEXAGON_TENSOR_REPACK; + } + } + + HEX_VERBOSE("ggml-hex: %s set-tensor-2d %s : data %p offset %zu size %zu n_copies %zu stride_tensor %zu stride_data %zu usage %d flags 0x%x\n", + sess->c_name(), tensor->name, data, offset, size, n_copies, stride_tensor, stride_data, (int) buffer->usage, extra->flags); + + if ((extra->flags & GGML_HEXAGON_TENSOR_REPACK) == 0) { + for (size_t i = 0; i < n_copies; i++) { + memcpy((uint8_t *) tensor->data + offset + i * stride_tensor, (const uint8_t *) data + i * stride_data, size); + } + return; + } + + if (extra->shadow_buf.size() < ggml_nbytes(tensor)) { + extra->shadow_buf.resize(ggml_nbytes(tensor)); + } + for (size_t i = 0; i < n_copies; i++) { + memcpy(extra->shadow_buf.data() + offset + i * stride_tensor, (const uint8_t *) data + i * stride_data, size); + } + extra->shadow_size += n_copies * size; + + if (extra->shadow_size >= ggml_nbytes(tensor)) { + repack_tensor_tiled(tensor, extra->shadow_buf.data(), extra->shadow_buf.size()); + extra->shadow_buf.clear(); + extra->shadow_buf.shrink_to_fit(); + extra->shadow_size = 0; + } +} + +static void ggml_backend_hexagon_buffer_get_tensor_2d(ggml_backend_buffer_t buffer, + const ggml_tensor * tensor, + void * data, + size_t offset, + size_t size, + size_t n_copies, + size_t stride_tensor, + size_t stride_data) { + auto extra = (ggml_hexagon_tensor_extra *) tensor->extra; + auto sbuf = (ggml_hexagon_shared_buffer *) buffer->context; + auto sess = sbuf->sess; + + HEX_VERBOSE("ggml-hex: %s get-tensor-2d %s : data %p offset %zu size %zu n_copies %zu stride_tensor %zu stride_data %zu usage %d\n", + sess->c_name(), tensor->name, data, offset, size, n_copies, stride_tensor, stride_data, (int) buffer->usage); + + if ((extra->flags & GGML_HEXAGON_TENSOR_REPACK) == 0) { + for (size_t i = 0; i < n_copies; i++) { + memcpy((uint8_t *)data + i * stride_data, (const uint8_t *)tensor->data + offset + i * stride_tensor, size); + } + return; + } + + size_t temp_size = n_copies > 0 ? (n_copies - 1) * stride_tensor + size : 0; + size_t slice_size = tensor->ne[1] * ggml_row_size(tensor->type, tensor->ne[0]); + size_t slice_offset = offset % slice_size; + size_t row_size_bytes = ggml_row_size(tensor->type, tensor->ne[0]); + + GGML_ASSERT((slice_offset % row_size_bytes) == 0 && "offset must be aligned to row boundary"); + GGML_ASSERT((temp_size % row_size_bytes) == 0 && "temp_size must be a multiple of row size"); + GGML_ASSERT((slice_offset / row_size_bytes) % 32 == 0 && "offset must be aligned to tile size (32 rows)"); + GGML_ASSERT((offset + temp_size) <= ggml_nbytes(tensor)); + + std::vector temp_buf(temp_size); + + switch (tensor->type) { + case GGML_TYPE_Q4_0: + repack_tiled_q4_0(temp_buf.data(), tensor, offset, temp_size); + break; + + case GGML_TYPE_Q4_1: + repack_tiled_q4_1(temp_buf.data(), tensor, offset, temp_size); + break; + + case GGML_TYPE_Q8_0: + repack_tiled_q8_0(temp_buf.data(), tensor, offset, temp_size); + break; + + case GGML_TYPE_IQ4_NL: + repack_tiled_q4_0(temp_buf.data(), tensor, offset, temp_size); + break; + + case GGML_TYPE_MXFP4: + repack_tiled_mxfp4(temp_buf.data(), tensor, offset, temp_size); + break; + + default: + memcpy(temp_buf.data(), (const uint8_t *) tensor->data + offset, temp_size); + break; + } + + for (size_t i = 0; i < n_copies; i++) { + memcpy((uint8_t *) data + i * stride_data, temp_buf.data() + i * stride_tensor, size); + } +} + static void ggml_backend_hexagon_buffer_clear(ggml_backend_buffer_t buffer, uint8_t value) { auto sbuf = (ggml_hexagon_shared_buffer *) buffer->context; auto sess = sbuf->sess; - HEX_VERBOSE("ggml-hex: %s clear-buff base %p size %zu\n", sess->c_name(), (void *) sbuf->base, sbuf->size); - memset(sbuf->base, value, sbuf->size); + HEX_VERBOSE("ggml-hex: %s clear-buff base %p size %zu\n", sess->c_name(), (void *) sbuf->base(), sbuf->size()); + memset(sbuf->base(), value, sbuf->size()); } static ggml_backend_buffer_i ggml_backend_hexagon_buffer_interface = { @@ -1049,6 +1445,40 @@ static ggml_backend_buffer_i ggml_backend_hexagon_buffer_interface = { /* .memset_tensor = */ NULL, /* .set_tensor = */ ggml_backend_hexagon_buffer_set_tensor, /* .get_tensor = */ ggml_backend_hexagon_buffer_get_tensor, + /* .set_tensor_2d = */ ggml_backend_hexagon_buffer_set_tensor_2d, + /* .get_tensor_2d = */ ggml_backend_hexagon_buffer_get_tensor_2d, + /* .cpy_tensor = */ ggml_backend_hexagon_buffer_cpy_tensor, + /* .clear = */ ggml_backend_hexagon_buffer_clear, + /* .reset = */ NULL, +}; + +// ** backend buffer type + +static void ggml_backend_hexagon_host_buffer_set_tensor(ggml_backend_buffer_t buffer, + ggml_tensor * tensor, + const void * data, + size_t offset, + size_t size) { + memcpy((char *) tensor->data + offset, data, size); + GGML_UNUSED(buffer); +} + +static void ggml_backend_hexagon_host_buffer_get_tensor(ggml_backend_buffer_t buffer, + const ggml_tensor * tensor, + void * data, + size_t offset, + size_t size) { + memcpy(data, (const char *) tensor->data + offset, size); + GGML_UNUSED(buffer); +} + +static ggml_backend_buffer_i ggml_backend_hexagon_host_buffer_interface = { + /* .free_buffer = */ ggml_backend_hexagon_buffer_free_buffer, + /* .get_base = */ ggml_backend_hexagon_buffer_get_base, + /* .init_tensor = */ ggml_backend_hexagon_buffer_init_tensor, + /* .memset_tensor = */ NULL, + /* .set_tensor = */ ggml_backend_hexagon_host_buffer_set_tensor, + /* .get_tensor = */ ggml_backend_hexagon_host_buffer_get_tensor, /* .set_tensor_2d = */ NULL, /* .get_tensor_2d = */ NULL, /* .cpy_tensor = */ ggml_backend_hexagon_buffer_cpy_tensor, @@ -1066,24 +1496,22 @@ static ggml_backend_buffer_t ggml_backend_hexagon_buffer_type_alloc_buffer( ggml_backend_buffer_type_t buffer_type, size_t size) { auto sess = static_cast(buffer_type->context)->sess; try { - size += 4 * 1024; // guard page - ggml_hexagon_shared_buffer * sbuf = new ggml_hexagon_shared_buffer(sess, size); + ggml_hexagon_shared_buffer * sbuf = new ggml_hexagon_shared_buffer(sess, size, false, GGML_HEXAGON_FENCE_BUFFER_SIZE); return ggml_backend_buffer_init(buffer_type, ggml_backend_hexagon_buffer_interface, sbuf, size); } catch (const std::exception & exc) { - GGML_LOG_ERROR("ggml-hex: %s failed to allocate buffer context (host): %s\n", sess->c_name(), exc.what()); + GGML_LOG_ERROR("ggml-hex: %s failed to allocate device buffer context: %s\n", sess->c_name(), exc.what()); return nullptr; } } -static ggml_backend_buffer_t ggml_backend_hexagon_repack_buffer_type_alloc_buffer( +static ggml_backend_buffer_t ggml_backend_hexagon_host_buffer_type_alloc_buffer( ggml_backend_buffer_type_t buffer_type, size_t size) { auto sess = static_cast(buffer_type->context)->sess; try { - size += 4 * 1024; // guard page - ggml_hexagon_shared_buffer * sbuf = new ggml_hexagon_shared_buffer(sess, size); - return ggml_backend_buffer_init(buffer_type, ggml_backend_hexagon_buffer_interface, sbuf, size); + ggml_hexagon_shared_buffer * sbuf = new ggml_hexagon_shared_buffer(sess, size, false, GGML_HEXAGON_FENCE_BUFFER_SIZE); + return ggml_backend_buffer_init(buffer_type, ggml_backend_hexagon_host_buffer_interface, sbuf, size); } catch (const std::exception & exc) { - GGML_LOG_ERROR("ggml-hex: %s failed to allocate buffer context (repack): %s\n", sess->c_name(), exc.what()); + GGML_LOG_ERROR("ggml-hex: %s failed to allocate host buffer context: %s\n", sess->c_name(), exc.what()); return nullptr; } } @@ -1094,7 +1522,7 @@ static size_t ggml_backend_hexagon_buffer_type_get_alignment(ggml_backend_buffer } static size_t ggml_backend_hexagon_buffer_type_get_alloc_size(ggml_backend_buffer_type_t buft, const struct ggml_tensor * t) { - if (t->type == GGML_TYPE_Q4_0 || t->type == GGML_TYPE_Q4_1 || t->type == GGML_TYPE_Q8_0 || t->type == GGML_TYPE_IQ4_NL || t->type == GGML_TYPE_MXFP4) { + if (ggml_hexagon_is_repack_type(t->type)) { int64_t ne0 = hex_round_up(t->ne[0], 32); int64_t ne1 = hex_round_up(t->ne[1], 32); int64_t ne2 = t->ne[2]; @@ -1112,14 +1540,12 @@ static size_t ggml_backend_hexagon_buffer_type_get_max_size(ggml_backend_buffer_ } static bool ggml_backend_hexagon_buffer_type_is_host(ggml_backend_buffer_type_t buft) { - return opt_hostbuf; - + return false; GGML_UNUSED(buft); } -static bool ggml_backend_hexagon_repack_buffer_type_is_host(ggml_backend_buffer_type_t buft) { - return false; - +static bool ggml_backend_hexagon_host_buffer_type_is_host(ggml_backend_buffer_type_t buft) { + return true; GGML_UNUSED(buft); } @@ -1132,26 +1558,19 @@ static ggml_backend_buffer_type_i ggml_backend_hexagon_buffer_type_interface = { /* .is_host = */ ggml_backend_hexagon_buffer_type_is_host, }; -static ggml_backend_buffer_type_i ggml_backend_hexagon_repack_buffer_type_interface = { +static ggml_backend_buffer_type_i ggml_backend_hexagon_host_buffer_type_interface = { /* .get_name = */ ggml_backend_hexagon_buffer_type_name, - /* .alloc_buffer = */ ggml_backend_hexagon_repack_buffer_type_alloc_buffer, + /* .alloc_buffer = */ ggml_backend_hexagon_host_buffer_type_alloc_buffer, /* .get_alignment = */ ggml_backend_hexagon_buffer_type_get_alignment, /* .get_max_size = */ ggml_backend_hexagon_buffer_type_get_max_size, /* .get_alloc_size = */ ggml_backend_hexagon_buffer_type_get_alloc_size, - /* .is_host = */ ggml_backend_hexagon_repack_buffer_type_is_host, + /* .is_host = */ ggml_backend_hexagon_host_buffer_type_is_host, }; static bool ggml_backend_buffer_is_hexagon(const struct ggml_backend_buffer * b) { return b->buft->iface.get_alignment == ggml_backend_hexagon_buffer_type_get_alignment; } -static inline bool ggml_backend_buffer_is_hexagon_repack(const struct ggml_backend_buffer * b) { - if (!opt_hostbuf) { - return ggml_backend_buffer_is_hexagon(b); - } - return b->buft->iface.alloc_buffer == ggml_backend_hexagon_repack_buffer_type_alloc_buffer; -} - struct ggml_hexagon_opbatch { ggml_hexagon_session* sess; @@ -1165,8 +1584,6 @@ struct ggml_hexagon_opbatch { std::unordered_map t_map; // tensor ptr to index std::unordered_multimap d_map; // tensor data to index - - unsigned int n_bufs; // num buffers in the batch unsigned int n_tens; // num tensors ... unsigned int n_ops; // num ops ... @@ -1186,6 +1603,7 @@ struct ggml_hexagon_opbatch { b_map.clear(); t_map.clear(); d_map.clear(); + ops.resize(n_ops_max); } ggml_hexagon_opbatch(ggml_hexagon_session *sess, size_t batch_size, size_t max_vmem) { @@ -1218,39 +1636,39 @@ struct ggml_hexagon_opbatch { // add buffer and return its index int add_buffer(ggml_hexagon_shared_buffer * sbuf) { // Lookup by fd - auto it = b_map.find(sbuf->fd); + auto it = b_map.find(sbuf->fd()); if (it != b_map.end()) { return it->second; } // Add new buffer to the batch int bi = n_bufs++; GGML_ASSERT(n_bufs < HTP_OP_MAX_BUFS); - b_map.insert({sbuf->fd, bi}); + b_map.insert({sbuf->fd(), bi}); htp_buf_desc &b = h_bufs[bi]; - b.base = (uint64_t) sbuf->base; - b.fd = sbuf->fd; - b.size = sbuf->size; + b.base = (uint64_t) sbuf->base(); + b.fd = sbuf->fd(); + b.size = sbuf->size(); b_vmem += b.size; - HEX_VERBOSE("ggml-hex: %s add-buffer #%u : fd %d base %p size %zu : vmem %zu\n", sess->c_name(), bi, b.fd, (void*) sbuf->base, (size_t) b.size, b_vmem); + HEX_VERBOSE("ggml-hex: %s add-buffer #%u : fd %d base %p size %zu : vmem %zu\n", sess->c_name(), bi, b.fd, (void*) sbuf->base(), (size_t) b.size, b_vmem); return bi; } - - bool same_shape(const htp_tensor * h, const ggml_tensor * t) const { + auto extra = (ggml_hexagon_tensor_extra *) t->extra; + int64_t ne0 = t->ne[0]; int64_t ne1 = t->ne[1]; - const bool is_repack = ggml_backend_buffer_is_hexagon_repack(t->buffer) && ggml_hexagon_is_repack_type(t->type); + const bool is_repack = (extra->flags & GGML_HEXAGON_TENSOR_REPACK) != 0; if (is_repack) { ne0 = hex_round_up(ne0, 32); ne1 = hex_round_up(ne1, 32); } int64_t nb1 = is_repack ? ggml_row_size(t->type, ne0) : t->nb[1]; - int64_t nb2 = is_repack ? nb1 * ne1 : t->nb[2]; + int64_t nb2 = is_repack ? nb1 * ne1 : t->nb[2]; int64_t nb3 = is_repack ? nb2 * t->ne[2] : t->nb[3]; return (h->type == t->type) && @@ -1260,7 +1678,8 @@ struct ggml_hexagon_opbatch { // add tensor and return its index int add_tensor(const ggml_tensor * t) { - auto sbuf = static_cast(t->buffer->context); + auto extra = (ggml_hexagon_tensor_extra *) t->extra; + auto sbuf = static_cast(t->buffer->context); // First lookup by tensor data auto range = d_map.equal_range(t->data); @@ -1280,7 +1699,7 @@ struct ggml_hexagon_opbatch { t_map.insert({t, ti}); d_map.insert({t->data, ti}); - uint64_t t_offset = (uint8_t *) t->data - sbuf->base; + uint64_t t_offset = (uint8_t *) t->data - sbuf->base(); size_t t_size = ggml_nbytes(t); htp_tensor &h = h_tens[ti]; @@ -1289,7 +1708,7 @@ struct ggml_hexagon_opbatch { h.data = t_offset; h.type = t->type; - const bool is_repack = ggml_backend_buffer_is_hexagon_repack(t->buffer) && ggml_hexagon_is_repack_type(t->type); + const bool is_repack = (extra->flags & GGML_HEXAGON_TENSOR_REPACK) != 0; if (is_repack) { h.ne[0] = hex_round_up(t->ne[0], 32); h.ne[1] = hex_round_up(t->ne[1], 32); @@ -1308,11 +1727,15 @@ struct ggml_hexagon_opbatch { h.nb[0] = t->nb[0]; h.nb[1] = t->nb[1]; h.nb[2] = t->nb[2]; h.nb[3] = t->nb[3]; } - - h.flags = 0; - if (ggml_backend_buffer_get_usage(t->buffer) != GGML_BACKEND_BUFFER_USAGE_WEIGHTS) { - h.flags |= HTP_TENSOR_COMPUTE; + if ((extra->flags & GGML_HEXAGON_TENSOR_WEIGHT) != 0) { + h.flags |= HTP_TENSOR_WEIGHT; + } + if ((extra->flags & GGML_HEXAGON_TENSOR_REPACK) != 0) { + h.flags |= HTP_TENSOR_REPACK; + } + if ((extra->flags & GGML_HEXAGON_TENSOR_FENCE) != 0) { + h.flags |= HTP_TENSOR_FENCE; } HEX_VERBOSE("ggml-hex: %s add-tensor #%u %s : bi %d data %p offset %zu size %zu flags 0x%x : %zu:%zu:%zu:%zu\n", sess->c_name(), @@ -1336,8 +1759,8 @@ struct ggml_hexagon_opbatch { extra_tens++; auto sbuf = static_cast(t->buffer->context); - if (!b_map.count(sbuf->fd)) { - extra_vmem += sbuf->size; + if (!b_map.count(sbuf->fd())) { + extra_vmem += sbuf->size(); extra_bufs += 1; } } @@ -1372,32 +1795,497 @@ struct ggml_hexagon_opbatch { o.opcode = node.opcode; o.flags = 0; - if (!(opt_opstage & HTP_OPSTAGE_COMPUTE)) { - o.flags |= HTP_OPFLAGS_SKIP_COMPUTE; - } + ggml_hexagon_dump_op_exec(sess->c_name(), ops[n], o.flags); + + auto inputs = node.get_inputs(); + for (unsigned int i=0; i < HTP_OP_MAX_INPUTS; i++) { + o.src[i] = (i < inputs.size() && inputs[i]) ? add_tensor(inputs[i]) : 0xffff; + } + + auto outputs = node.get_outputs(); + for (unsigned int i=0; i < HTP_OP_MAX_OUTPUTS; i++) { + o.dst[i] = (i < outputs.size() && outputs[i]) ? add_tensor(outputs[i]) : 0xffff; + } + } + + bool try_fuse_allreduce_add(const htp_opnode & node) { + if (n_ops == 0 || opt_ar_select != 2) return false; + if (node.opcode != HTP_OP_ADD) return false; + + htp_opnode & last_node = ops[n_ops - 1]; + if (last_node.opcode != HTP_OP_ALLREDUCE) return false; + + auto * ar_kparams = (struct htp_allreduce_kernel_params *) last_node.kernel_params; + const uint32_t rank = (uint32_t) ar_kparams->rank; + const ggml_tensor * ar_local = (rank < last_node.inputs.size()) ? last_node.inputs[rank] : nullptr; + const ggml_tensor * add_src0 = node.src0(); + const ggml_tensor * add_src1 = node.src1(); + + if (!add_src0 || !add_src1 || !ar_local) return false; + if (!ggml_hexagon_tensor_is_fuseable(ar_local)) return false; + + const ggml_tensor * res_tensor = nullptr; + if (add_src0 == ar_local || add_src0->data == ar_local->data) { + res_tensor = add_src1; + } else if (add_src1 == ar_local || add_src1->data == ar_local->data) { + res_tensor = add_src0; + } else { + return false; + } + + if (!res_tensor || !res_tensor->data) return false; + + if (ar_local->type != res_tensor->type) return false; + + const bool is_same_shape = (ar_local->ne[0] == res_tensor->ne[0] && ar_local->ne[1] == res_tensor->ne[1] && + ar_local->ne[2] == res_tensor->ne[2] && ar_local->ne[3] == res_tensor->ne[3]); + const bool is_row_bcast = (ar_local->ne[0] == res_tensor->ne[0] && + res_tensor->ne[1] == 1 && res_tensor->ne[2] == 1 && res_tensor->ne[3] == 1); + + if (!is_same_shape && !is_row_bcast) return false; + + if (is_same_shape) { + if (ar_local->nb[1] != res_tensor->nb[1] || ar_local->nb[2] != res_tensor->nb[2] || + ar_local->nb[3] != res_tensor->nb[3]) { + return false; + } + if (ggml_is_contiguous(ar_local) != ggml_is_contiguous(res_tensor)) { + return false; + } + } + if (ggml_is_contiguous(ar_local) != ggml_is_contiguous(node.dst())) { + return false; + } + + struct htp_allreduce_kernel_params new_kparams; + if (!ggml_hexagon_precompute_allreduce_params( + sess, node.dst(), (uint32_t) ar_kparams->rank, (uint32_t) ar_kparams->n_ranks, true, is_row_bcast, &new_kparams + )) { + HEX_VERBOSE("ggml-hex: %s skip ALLREDUCE_ADD fusion: solver failed\n", sess->c_name()); + return false; + } + + size_t extra_bufs = 0, extra_vmem = 0, extra_tens = 0; + auto fit_t = [&](const ggml_tensor * t) { + if (!t) return; + if (!t_map.count(t)) { + extra_tens++; + auto sbuf = static_cast(t->buffer->context); + if (!b_map.count(sbuf->fd())) { + extra_vmem += sbuf->size(); + extra_bufs += 1; + } + } + }; + fit_t(res_tensor); + fit_t(node.dst()); + if ((extra_bufs + n_bufs) > n_bufs_max || (extra_tens + n_tens) > n_tens_max || (extra_vmem + b_vmem) > b_vmem_max) { + return false; + } + + last_node.opcode = HTP_OP_ALLREDUCE_ADD; + last_node.name = "ALLREDUCE+ADD"; + last_node.inputs.push_back(res_tensor); + last_node.outputs.clear(); + last_node.outputs.push_back(node.dst()); + last_node.fused.push_back(node.node); + memcpy(last_node.kernel_params, &new_kparams, sizeof(new_kparams)); + + htp_op_desc & o = h_ops[n_ops - 1]; + o.opcode = HTP_OP_ALLREDUCE_ADD; + memcpy(o.kernel_params, &new_kparams, sizeof(new_kparams)); + + const uint32_t n_ranks = (uint32_t) ar_kparams->n_ranks; + o.src[2 * n_ranks] = add_tensor(res_tensor); + o.dst[0] = add_tensor(node.dst()); + for (uint32_t d = 1; d < HTP_OP_MAX_OUTPUTS; d++) { + o.dst[d] = 0xffff; + } + + HEX_VERBOSE("ggml-hex: %s fused ALLREDUCE+ADD (#%u)\n", sess->c_name(), n_ops - 1); + return true; + } + + bool try_fuse_rms_norm_mul(const htp_opnode & node) { + if (n_ops == 0) return false; + if (node.opcode != HTP_OP_MUL) return false; + + htp_opnode & last_node = ops[n_ops - 1]; + if (last_node.opcode != HTP_OP_RMS_NORM) return false; + + const ggml_tensor * mul_src0 = node.src0(); + const ggml_tensor * mul_src1 = node.src1(); + const ggml_tensor * rms_out = last_node.dst(); + + if (!mul_src0 || !mul_src1 || !rms_out) return false; + if (!ggml_hexagon_tensor_is_fuseable(rms_out)) return false; + + const ggml_tensor * weight = nullptr; + if (mul_src0 == rms_out || mul_src0->data == rms_out->data) { + weight = mul_src1; + } else if (mul_src1 == rms_out || mul_src1->data == rms_out->data) { + weight = mul_src0; + } else { + return false; + } + + if (!weight || !weight->data) return false; + + const ggml_tensor * src0 = last_node.src0(); + if (!src0 || !src0->data) return false; + + if (src0->ne[0] != weight->ne[0] || src0->ne[0] != node.dst()->ne[0]) { + return false; + } + + const bool is_row_bcast = (weight->ne[1] == 1 && weight->ne[2] == 1 && weight->ne[3] == 1); + const bool is_same_shape = (src0->ne[0] == weight->ne[0] && src0->ne[1] == weight->ne[1] && + src0->ne[2] == weight->ne[2] && src0->ne[3] == weight->ne[3]); + if (!is_row_bcast && !is_same_shape) return false; + + if (!ggml_are_same_shape(src0, node.dst())) { + return false; + } + if (ggml_is_contiguous(src0) != ggml_is_contiguous(node.dst())) { + return false; + } + + struct htp_unary_kernel_params new_kparams; + ggml_hexagon_precompute_unary_params( + sess, HTP_OP_RMS_NORM_MUL, src0, weight, node.dst(), &new_kparams + ); + + if ((size_t) new_kparams.vtcm_size > sess->vtcm_size) { + HEX_VERBOSE("ggml-hex: %s skip RMS_NORM_MUL fusion: VTCM needed (%d) > budget (%zu)\n", + sess->c_name(), new_kparams.vtcm_size, sess->vtcm_size); + return false; + } + + size_t extra_bufs = 0, extra_vmem = 0, extra_tens = 0; + auto fit_t = [&](const ggml_tensor * t) { + if (!t) return; + if (!t_map.count(t)) { + extra_tens++; + auto sbuf = static_cast(t->buffer->context); + if (!b_map.count(sbuf->fd())) { + extra_vmem += sbuf->size(); + extra_bufs += 1; + } + } + }; + fit_t(weight); + fit_t(node.dst()); + if ((extra_bufs + n_bufs) > n_bufs_max || (extra_tens + n_tens) > n_tens_max || (extra_vmem + b_vmem) > b_vmem_max) { + return false; + } + + last_node.opcode = HTP_OP_RMS_NORM_MUL; + last_node.name = "RMS_NORM+MUL"; + last_node.inputs.clear(); + last_node.inputs.push_back(src0); + last_node.inputs.push_back(weight); + last_node.outputs.clear(); + last_node.outputs.push_back(node.dst()); + last_node.fused.push_back(node.node); + memcpy(last_node.kernel_params, &new_kparams, sizeof(new_kparams)); + + htp_op_desc & o = h_ops[n_ops - 1]; + o.opcode = HTP_OP_RMS_NORM_MUL; + memcpy(o.kernel_params, &new_kparams, sizeof(new_kparams)); + + o.src[0] = add_tensor(src0); + o.src[1] = add_tensor(weight); + for (uint32_t s = 2; s < HTP_OP_MAX_INPUTS; s++) { + o.src[s] = 0xffff; + } + o.dst[0] = add_tensor(node.dst()); + for (uint32_t d = 1; d < HTP_OP_MAX_OUTPUTS; d++) { + o.dst[d] = 0xffff; + } + + HEX_VERBOSE("ggml-hex: %s fused RMS_NORM+MUL (#%u)\n", sess->c_name(), n_ops - 1); + return true; + } + + bool try_fuse_mul_mat_add(const htp_opnode & node) { + if (n_ops == 0) return false; + if (node.opcode != HTP_OP_ADD) return false; + + htp_opnode & last_node = ops[n_ops - 1]; + if (last_node.opcode != HTP_OP_MUL_MAT) return false; + + const ggml_tensor * add_src0 = node.src0(); + const ggml_tensor * add_src1 = node.src1(); + const ggml_tensor * mm_out = last_node.dst(); + + if (!add_src0 || !add_src1 || !mm_out) return false; + if (!ggml_hexagon_tensor_is_fuseable(mm_out)) return false; + + const ggml_tensor * src2 = nullptr; + if (add_src0 == mm_out || add_src0->data == mm_out->data) { + src2 = add_src1; + } else if (add_src1 == mm_out || add_src1->data == mm_out->data) { + src2 = add_src0; + } else { + return false; + } + + if (!src2 || !src2->data) return false; + + const ggml_tensor * src0 = last_node.src0(); + const ggml_tensor * src1 = last_node.src1(); + if (!src0 || !src1) return false; + + struct htp_mm_kernel_params kparams; + ggml_hexagon_precompute_fused_matmul_add_params(sess, src0, src1, src2, node.dst(), &kparams); + const int src1_nrows = src1->ne[1] * src1->ne[2] * src1->ne[3]; + const bool can_fuse = (kparams.n_hmx > 0) || (src1_nrows == 1); + if (!can_fuse) return false; + + if ((size_t) kparams.vtcm_size > sess->vtcm_size) { + HEX_VERBOSE("ggml-hex: %s skip MUL_MAT_ADD fusion: VTCM needed (%d) > budget (%zu)\n", + sess->c_name(), kparams.vtcm_size, sess->vtcm_size); + return false; + } + + size_t extra_bufs = 0, extra_vmem = 0, extra_tens = 0; + auto fit_t = [&](const ggml_tensor * t) { + if (!t) return; + if (!t_map.count(t)) { + extra_tens++; + auto sbuf = static_cast(t->buffer->context); + if (!b_map.count(sbuf->fd())) { + extra_vmem += sbuf->size(); + extra_bufs += 1; + } + } + }; + fit_t(src2); + fit_t(node.dst()); + if ((extra_bufs + n_bufs) > n_bufs_max || (extra_tens + n_tens) > n_tens_max || (extra_vmem + b_vmem) > b_vmem_max) { + return false; + } + + last_node.opcode = HTP_OP_MUL_MAT_ADD; + last_node.name = "MUL_MAT+ADD"; + last_node.inputs.clear(); + last_node.inputs.push_back(src0); + last_node.inputs.push_back(src1); + last_node.inputs.push_back(src2); + last_node.outputs.clear(); + last_node.outputs.push_back(node.dst()); + last_node.fused.push_back(node.node); + memcpy(last_node.kernel_params, &kparams, sizeof(kparams)); + + htp_op_desc & o = h_ops[n_ops - 1]; + o.opcode = HTP_OP_MUL_MAT_ADD; + memcpy(o.kernel_params, &kparams, sizeof(kparams)); + + o.src[0] = add_tensor(src0); + o.src[1] = add_tensor(src1); + o.src[2] = add_tensor(src2); + for (uint32_t s = 3; s < HTP_OP_MAX_INPUTS; s++) { + o.src[s] = 0xffff; + } + o.dst[0] = add_tensor(node.dst()); + for (uint32_t d = 1; d < HTP_OP_MAX_OUTPUTS; d++) { + o.dst[d] = 0xffff; + } + + HEX_VERBOSE("ggml-hex: %s fused MUL_MAT+ADD (#%u)\n", sess->c_name(), n_ops - 1); + return true; + } + + bool try_fuse_mul_mat_nx(const htp_opnode & node) { + if (n_ops == 0 || node.opcode != HTP_OP_MUL_MAT) return false; + if (!is_mergeable_mul_mat(node.node)) return false; + + const ggml_tensor * w_in = node.src0(); + const ggml_tensor * x_in = node.src1(); + const ggml_tensor * d_in = node.dst(); + if (!w_in || !x_in || !d_in) return false; + + htp_opnode & last_node = ops[n_ops - 1]; + + // Case 1: last_node is already MUL_MAT_NX + if (last_node.opcode == HTP_OP_MUL_MAT_NX) { + const uint32_t curr_n = (uint32_t) last_node.outputs.size(); + if (curr_n >= HTP_OP_MAX_OUTPUTS || curr_n + 1 >= HTP_OP_MAX_INPUTS) { + return false; + } + + const ggml_tensor * w0 = last_node.inputs[0]; + const ggml_tensor * x = last_node.inputs[curr_n]; + + if (x_in != x || w_in->type != w0->type || w_in->ne[0] != w0->ne[0]) { + return false; + } + + struct htp_mm_kernel_params kparams; + ggml_hexagon_precompute_fused_mmnx_params(sess, w0, x, curr_n + 1, &kparams); + if ((size_t) kparams.vtcm_size > sess->vtcm_size) { + HEX_VERBOSE("ggml-hex: %s skip NX fusion: VTCM needed (%d) > budget (%zu)\n", + sess->c_name(), kparams.vtcm_size, sess->vtcm_size); + return false; + } + + size_t extra_bufs = 0, extra_vmem = 0, extra_tens = 0; + auto fit_t = [&](const ggml_tensor * t) { + if (!t) return; + if (!t_map.count(t)) { + extra_tens++; + auto sbuf = static_cast(t->buffer->context); + if (!b_map.count(sbuf->fd())) { + extra_vmem += sbuf->size(); + extra_bufs += 1; + } + } + }; + fit_t(w_in); + fit_t(d_in); + if ((extra_bufs + n_bufs) > n_bufs_max || (extra_tens + n_tens) > n_tens_max || (extra_vmem + b_vmem) > b_vmem_max) { + return false; + } + + last_node.inputs[curr_n] = w_in; + last_node.inputs.push_back(x); + last_node.outputs.push_back(d_in); + last_node.fused.push_back(node.node); + memcpy(last_node.kernel_params, &kparams, sizeof(kparams)); + + htp_op_desc & o = h_ops[n_ops - 1]; + memcpy(o.kernel_params, &kparams, sizeof(kparams)); + + for (uint32_t s = 0; s <= curr_n + 1; s++) { + o.src[s] = add_tensor(last_node.inputs[s]); + } + for (uint32_t s = curr_n + 2; s < HTP_OP_MAX_INPUTS; s++) { + o.src[s] = 0xffff; + } + for (uint32_t d = 0; d <= curr_n; d++) { + o.dst[d] = add_tensor(last_node.outputs[d]); + } + for (uint32_t d = curr_n + 1; d < HTP_OP_MAX_OUTPUTS; d++) { + o.dst[d] = 0xffff; + } + + HEX_VERBOSE("ggml-hex: %s fused MUL_MAT_NX (N=%u, #%u)\n", sess->c_name(), curr_n + 1, n_ops - 1); + return true; + } + + // Case 2: last_node is single MUL_MAT + if (last_node.opcode == HTP_OP_MUL_MAT) { + if (!is_mergeable_mul_mat_pair(last_node.node, node.node)) { + return false; + } + + const ggml_tensor * w0 = last_node.src0(); + const ggml_tensor * x = last_node.src1(); + const ggml_tensor * w1 = node.src0(); + if (!w0 || !x || !w1) return false; - ggml_hexagon_dump_op_exec(sess->c_name(), ops[n], o.flags); + struct htp_mm_kernel_params kparams; + ggml_hexagon_precompute_fused_mmnx_params(sess, w0, x, 2, &kparams); + if ((size_t) kparams.vtcm_size > sess->vtcm_size) { + HEX_VERBOSE("ggml-hex: %s skip NX fusion: VTCM needed (%d) > budget (%zu)\n", + sess->c_name(), kparams.vtcm_size, sess->vtcm_size); + return false; + } - auto inputs = node.get_inputs(); - for (unsigned int i=0; i < HTP_OP_MAX_INPUTS; i++) { - o.src[i] = (i < inputs.size() && inputs[i]) ? add_tensor(inputs[i]) : 0xffff; - } + size_t extra_bufs = 0, extra_vmem = 0, extra_tens = 0; + auto fit_t = [&](const ggml_tensor * t) { + if (!t) return; + if (!t_map.count(t)) { + extra_tens++; + auto sbuf = static_cast(t->buffer->context); + if (!b_map.count(sbuf->fd())) { + extra_vmem += sbuf->size(); + extra_bufs += 1; + } + } + }; + fit_t(w1); + fit_t(node.dst()); + if ((extra_bufs + n_bufs) > n_bufs_max || (extra_tens + n_tens) > n_tens_max || (extra_vmem + b_vmem) > b_vmem_max) { + return false; + } - auto outputs = node.get_outputs(); - for (unsigned int i=0; i < HTP_OP_MAX_OUTPUTS; i++) { - o.dst[i] = (i < outputs.size() && outputs[i]) ? add_tensor(outputs[i]) : 0xffff; + const ggml_tensor * dst_0 = last_node.dst(); + const ggml_tensor * dst_1 = node.dst(); + + last_node.opcode = HTP_OP_MUL_MAT_NX; + last_node.name = "MUL_MAT_NX"; + last_node.inputs.clear(); + last_node.inputs.push_back(w0); + last_node.inputs.push_back(w1); + last_node.inputs.push_back(x); + last_node.outputs.clear(); + last_node.outputs.push_back(dst_0); + last_node.outputs.push_back(dst_1); + last_node.fused.push_back(node.node); + memcpy(last_node.kernel_params, &kparams, sizeof(kparams)); + + htp_op_desc & o = h_ops[n_ops - 1]; + o.opcode = HTP_OP_MUL_MAT_NX; + memcpy(o.kernel_params, &kparams, sizeof(kparams)); + + o.src[0] = add_tensor(w0); + o.src[1] = add_tensor(w1); + o.src[2] = add_tensor(x); + for (uint32_t s = 3; s < HTP_OP_MAX_INPUTS; s++) { + o.src[s] = 0xffff; + } + o.dst[0] = add_tensor(dst_0); + o.dst[1] = add_tensor(dst_1); + for (uint32_t d = 2; d < HTP_OP_MAX_OUTPUTS; d++) { + o.dst[d] = 0xffff; + } + + HEX_VERBOSE("ggml-hex: %s fused MUL_MAT_NX (N=2, #%u)\n", sess->c_name(), n_ops - 1); + return true; } + + return false; } - void finalize_ranges() { +enum ggml_hexagon_fusion_flags { + GGML_HEXAGON_FUSE_ALLREDUCE_ADD = (1 << 1), // 2 + GGML_HEXAGON_FUSE_RMS_NORM_MUL = (1 << 2), // 4 + GGML_HEXAGON_FUSE_MUL_MAT_ADD = (1 << 3), // 8 + GGML_HEXAGON_FUSE_MUL_MAT_NX = (1 << 4), // 16 +}; + +static inline bool ggml_hexagon_is_fusion_enabled(int flag) { + if (opt_opfusion <= 0) return false; + if (opt_opfusion == 1) return true; // 1 enables all + return (opt_opfusion & flag) != 0; +} + + bool try_fuse(const htp_opnode & node) { + if (!opt_opfusion) return false; + if (ggml_hexagon_is_fusion_enabled(GGML_HEXAGON_FUSE_ALLREDUCE_ADD) && try_fuse_allreduce_add(node)) return true; + if (ggml_hexagon_is_fusion_enabled(GGML_HEXAGON_FUSE_RMS_NORM_MUL) && try_fuse_rms_norm_mul(node)) return true; + if (ggml_hexagon_is_fusion_enabled(GGML_HEXAGON_FUSE_MUL_MAT_ADD) && try_fuse_mul_mat_add(node)) return true; + if (ggml_hexagon_is_fusion_enabled(GGML_HEXAGON_FUSE_MUL_MAT_NX) && try_fuse_mul_mat_nx(node)) return true; + return false; } }; +struct ggml_hexagon_registry { + ggml_hexagon_registry(ggml_backend_reg_t reg); + ~ggml_hexagon_registry(); + + ggml_backend_device devices[GGML_HEXAGON_MAX_SESSIONS]; +}; + struct ggml_hexagon_opqueue { // Shared buffer for storing batches ggml_hexagon_shared_buffer *shm_buf; size_t shm_blk_size; + uint64_t req_seq = 0; + uint64_t rsp_seq = 0; + using opvec = std::vector; std::queue done; // completed batch ids @@ -1429,8 +2317,8 @@ struct ggml_hexagon_opqueue { for (unsigned int i = 0; i < depth; i++) { done.push(i); } if (opt_verbose) { - GGML_LOG_INFO("ggml-hex: %s allocated op-queue : batch-size %zu depth %zu shm-size %zu shm-block-size %zu\n", - sess->c_name(), batch_size, depth, shm_buf->size, shm_blk_size); + GGML_LOG_INFO("ggml-hex: %s allocated opqueue : batch-size %zu depth %zu shm-size %zu shm-block-size %zu\n", + sess->c_name(), batch_size, depth, shm_buf->size(), shm_blk_size); } } @@ -1453,8 +2341,9 @@ struct ggml_hexagon_opqueue { req.n_bufs = op_batch->n_bufs; req.n_tensors = op_batch->n_tens; req.n_ops = op_batch->n_ops; + req.seq = ++req_seq; - op_cache[req.id] = op_batch->ops; + op_cache[req.id] = std::move(op_batch->ops); start_usec[req.id] = ggml_time_us(); const size_t b_size = sizeof(htp_buf_desc) * req.n_bufs; @@ -1470,10 +2359,10 @@ struct ggml_hexagon_opqueue { req.n_traces = 0; } - dbuf.ptr = shm_buf->base + (req.id * shm_blk_size); - dbuf.fd = shm_buf->fd; + dbuf.ptr = shm_buf->base() + (req.id * shm_blk_size); + dbuf.fd = shm_buf->fd(); dbuf.flags = DSPQUEUE_BUFFER_FLAG_FLUSH_SENDER | DSPQUEUE_BUFFER_FLAG_INVALIDATE_RECIPIENT; - dbuf.offset = (uint8_t*) dbuf.ptr - (uint8_t*) shm_buf->base; + dbuf.offset = (uint8_t*) dbuf.ptr - (uint8_t*) shm_buf->base(); dbuf.size = b_size + t_size + o_size + p_size + tr_size; GGML_ASSERT(dbuf.size <= shm_blk_size); @@ -1487,7 +2376,7 @@ struct ggml_hexagon_opqueue { memcpy(t_ptr, (void *) op_batch->h_tens.data(), t_size); memcpy(o_ptr, (void *) op_batch->h_ops.data(), o_size); - HEX_VERBOSE("ggml-hex: %s op-queue push batch #%u : n-bufs %u n-tensors %u n-ops %u vmem %zu : b-size %zu t-size %zu o-size %zu m-size %zu\n", + HEX_VERBOSE("ggml-hex: %s opqueue-push batch #%u : n-bufs %u n-tensors %u n-ops %u vmem %zu : b-size %zu t-size %zu o-size %zu m-size %zu\n", shm_buf->sess->c_name(), req.id, req.n_bufs, req.n_tensors, req.n_ops, op_batch->b_vmem, b_size, t_size, o_size, (size_t) dbuf.size); @@ -1530,33 +2419,40 @@ struct ggml_hexagon_opqueue { const size_t m_size = b_size + t_size + o_size + p_size + tr_size; GGML_ASSERT(m_size <= shm_blk_size); - HEX_VERBOSE("ggml-hex: %s op-queue pop batch #%u : n-bufs %u n-tensors %u n-ops %u : m-size %zu b-size %zu t-size %zu o-size %zu\n", + HEX_VERBOSE("ggml-hex: %s opqueue-pop batch #%u : n-bufs %u n-tensors %u n-ops %u : m-size %zu b-size %zu t-size %zu o-size %zu\n", shm_buf->sess->c_name(), rsp.id, rsp.n_bufs, rsp.n_tensors, rsp.n_ops, (size_t) dbuf.size, b_size, t_size, o_size); uint8_t * m_ptr = (uint8_t*) dbuf.ptr; uint8_t * p_ptr = m_ptr + (b_size + t_size + o_size); - if (opt_profile && rsp.n_ops > 0) { + if (rsp.n_ops > 0) { auto & ops = op_cache[rsp.id]; - GGML_ASSERT(rsp.n_ops <= ops.size()); const htp_prof_desc * pd = (const htp_prof_desc *) p_ptr; - const htp_trace_desc * trace_events = nullptr; - if (opt_profile == 3) { trace_events = (const htp_trace_desc *) (p_ptr + p_size); } - ggml_hexagon_dump_batch_prof(shm_buf->sess->name, rsp); + if (opt_profile) { + ggml_hexagon_dump_batch_prof(shm_buf->sess->name, rsp); + } for (uint32_t i = 0; i < rsp.n_ops; i++) { - ggml_hexagon_dump_op_prof(shm_buf->sess->name, ops[i], pd[i]); + if (opt_profile) { + ggml_hexagon_dump_op_prof(shm_buf->sess->name, ops[i], pd[i]); + } } - ggml_hexagon_dump_trace_events(shm_buf->sess->name, rsp, trace_events, n_traces); + if (opt_profile) { + ggml_hexagon_dump_trace_events(shm_buf->sess->name, rsp, trace_events, n_traces); + } + } + + if (rsp.seq > rsp_seq) { + rsp_seq = rsp.seq; } } }; @@ -1601,10 +2497,8 @@ void ggml_hexagon_session::flush_pending(bool all) { } } -void ggml_hexagon_session::flush_batch() { - if (op_batch->empty()) { return; } - - op_batch->finalize_ranges(); +void ggml_hexagon_session::flush_batch(size_t min_ops) { + if (op_batch->n_ops < min_ops) { return; } htp_opbatch_req req {}; dspqueue_buffer dbuf{}; @@ -1625,17 +2519,267 @@ void ggml_hexagon_session::flush_batch() { } } +void ggml_hexagon_session::flush(bool all) { + flush_sync_peers(); + flush_batch(); + flush_pending(all); +} + void ggml_hexagon_session::enqueue_op(const htp_opnode & node) { + for (auto t : node.get_inputs()) { + if (t && t->buffer && ggml_backend_buffer_is_hexagon(t->buffer)) { + if (ggml_backend_hexagon_buffer_get_sess(t->buffer) != this) { + this->clone_buffer(static_cast(t->buffer->context)); + } + } + } + for (auto t : node.get_outputs()) { + if (t && t->buffer && ggml_backend_buffer_is_hexagon(t->buffer)) { + if (ggml_backend_hexagon_buffer_get_sess(t->buffer) != this) { + this->clone_buffer(static_cast(t->buffer->context)); + } + } + } + + if (opt_opfusion && op_batch->try_fuse(node)) { + return; + } + if (!op_batch->fit_op(node)) { flush_batch(); } op_batch->add_op(node); } -// Flush HTP response queue i.e wait for all outstanding requests to complete -void ggml_hexagon_session::flush(bool all) { +void ggml_hexagon_session::enqueue_cpy(const ggml_tensor * src, ggml_tensor * dst, const ggml_tensor * sync_tensor, uint32_t fence_seq) { + htp_opnode cpy_node(HTP_OP_CPY); + + ggml_tensor* node = cpy_node.add_dummy(*dst); + node->op = GGML_OP_CPY; + node->src[0] = const_cast(src); + node->src[1] = sync_tensor ? cpy_node.add_dummy(*sync_tensor) : nullptr; + if (sync_tensor) { + node->op_params[0] = (int32_t) fence_seq; + } + + cpy_node.init(node); + if (sync_tensor) { + cpy_node.name = "CPY+FENCE"; + } + this->enqueue_op(cpy_node); +} + +void ggml_hexagon_session::enqueue_fence(const ggml_tensor * sync_tensor, uint32_t fence_seq) { + htp_opnode sync_node(HTP_OP_FENCE); + + ggml_tensor* node = sync_node.add_dummy(*sync_tensor); + node->op = GGML_OP_NONE; + node->src[0] = node; + node->op_params[0] = (int32_t) fence_seq; + + sync_node.init(node); + sync_node.name = "FENCE"; + this->enqueue_op(sync_node); +} + +static bool ggml_hexagon_precompute_allreduce_params( + const struct ggml_hexagon_session * sess, + const struct ggml_tensor * dst, + uint32_t rank, + uint32_t n_ranks, + bool has_add, + bool is_row_bcast, + struct htp_allreduce_kernel_params * kparams +) { + memset(kparams, 0, sizeof(*kparams)); + kparams->rank = (int32_t) rank; + kparams->n_ranks = (int32_t) n_ranks; + kparams->is_row_bcast = (has_add && is_row_bcast) ? 1 : 0; + + const uint32_t n_bufs = n_ranks + 1 + (has_add ? 1 : 0); + const uint32_t nelem = (uint32_t) ggml_nelements(dst); + const uint32_t elem_size = (dst->type == GGML_TYPE_F16) ? sizeof(ggml_fp16_t) : sizeof(float); + const bool is_contiguous = ggml_is_contiguous(dst); + + const uint32_t ne0 = (uint32_t) dst->ne[0]; + const uint32_t ne1 = (uint32_t) (dst->ne[1] * dst->ne[2] * dst->ne[3]); + kparams->ne0 = (int32_t) ne0; + kparams->ne1 = (int32_t) ne1; + + const bool use_1d = is_contiguous && !(has_add && is_row_bcast && ne1 > 1); + + if (has_add) { + kparams->n_dsts = 1; + if (use_1d) { + kparams->rank_elem_start = 0; + kparams->rank_nelem = (int32_t) nelem; + } else { + kparams->rank_elem_start = 0; + kparams->rank_nelem = (int32_t) ne1; + } + } else { + kparams->n_dsts = (int32_t) n_ranks; + if (use_1d) { + const uint32_t rank_chunk_elems = hex_round_up((nelem + n_ranks - 1) / n_ranks, 128); + const uint32_t rank_elem_start = (std::min)(rank * rank_chunk_elems, nelem); + const uint32_t rank_elem_end = (std::min)(rank_elem_start + rank_chunk_elems, nelem); + const uint32_t rank_nelem = rank_elem_end - rank_elem_start; + kparams->rank_elem_start = (int32_t) rank_elem_start; + kparams->rank_nelem = (int32_t) rank_nelem; + } else { + const uint32_t rank_chunk_rows = (ne1 + n_ranks - 1) / n_ranks; + const uint32_t rank_r0 = (std::min)(rank * rank_chunk_rows, ne1); + const uint32_t rank_r1 = (std::min)(rank_r0 + rank_chunk_rows, ne1); + const uint32_t rank_nrows = rank_r1 - rank_r0; + kparams->rank_elem_start = (int32_t) rank_r0; + kparams->rank_nelem = (int32_t) rank_nrows; + } + } + + if (use_1d) { + const uint32_t rank_nelem = (uint32_t) kparams->rank_nelem; + const uint32_t n_threads = (std::min)((uint32_t) sess->n_threads, (std::max)(1u, rank_nelem / 128)); + kparams->n_threads = n_threads; + + uint32_t block_elems = 65536; + if (block_elems > rank_nelem / n_threads && rank_nelem / n_threads > 128) { + block_elems = hex_round_up(rank_nelem / (n_threads * 2), 128); + } + block_elems = (std::max)(128u, block_elems); + + kparams->block_elems = block_elems; + kparams->vtcm_size_per_thread = 2 * block_elems * elem_size; + kparams->vtcm_size = n_threads * n_bufs * kparams->vtcm_size_per_thread; + + while ((size_t) kparams->vtcm_size > sess->vtcm_size && block_elems > 128) { + const size_t max_bytes_per_buf = sess->vtcm_size / (n_threads * n_bufs * 2); + block_elems = (uint32_t) hex_align_down((size_t) (max_bytes_per_buf / elem_size), 128); + if (block_elems < 128) break; + kparams->block_elems = block_elems; + kparams->vtcm_size_per_thread = 2 * block_elems * elem_size; + kparams->vtcm_size = n_threads * n_bufs * kparams->vtcm_size_per_thread; + } + + if (sess->vtcm_size < (size_t) kparams->vtcm_size || block_elems < 128) { + HEX_VERBOSE("ggml-hex: %s allreduce 1D solver failed to fit VTCM (%d > %zu)\n", + sess->c_name(), kparams->vtcm_size, sess->vtcm_size); + return false; + } + + kparams->elems_per_thread = hex_round_up((rank_nelem + n_threads - 1) / n_threads, block_elems); + kparams->kernel_type = HTP_ALLREDUCE_KERNEL_DMA_1D; + return true; + } else { + const uint32_t rank_nrows = (uint32_t) kparams->rank_nelem; + const uint32_t n_threads = (std::min)((uint32_t) sess->n_threads, (std::max)(1u, rank_nrows)); + kparams->n_threads = n_threads; + + const uint32_t row_bytes = ne0 * elem_size; + const uint32_t row_size_aligned = (uint32_t) hex_align_up(row_bytes, 128); + kparams->row_size_aligned = row_size_aligned; + + const uint32_t nrows_per_thread = (rank_nrows + n_threads - 1) / n_threads; + uint32_t block_rows = (std::min)(128u, nrows_per_thread); + block_rows = (std::max)(1u, block_rows); + kparams->block_elems = block_rows; + + kparams->vtcm_size_per_thread = 2 * (block_rows * row_size_aligned); + kparams->vtcm_size = n_threads * n_bufs * kparams->vtcm_size_per_thread; + + while ((size_t) kparams->vtcm_size > sess->vtcm_size && block_rows > 1) { + const size_t max_rows_per_buf = sess->vtcm_size / (n_threads * n_bufs * 2 * row_size_aligned); + block_rows = (std::max)(1u, (uint32_t) max_rows_per_buf); + kparams->block_elems = block_rows; + kparams->vtcm_size_per_thread = 2 * (block_rows * row_size_aligned); + kparams->vtcm_size = n_threads * n_bufs * kparams->vtcm_size_per_thread; + if (max_rows_per_buf == 0) break; + } + + if (sess->vtcm_size < (size_t) kparams->vtcm_size || block_rows < 1) { + HEX_VERBOSE("ggml-hex: %s allreduce 2D solver failed to fit VTCM (%d > %zu)\n", + sess->c_name(), kparams->vtcm_size, sess->vtcm_size); + return false; + } + + kparams->elems_per_thread = nrows_per_thread; + kparams->kernel_type = HTP_ALLREDUCE_KERNEL_DMA_2D; + return true; + } +} + +void ggml_hexagon_session::enqueue_allreduce( + const ggml_tensor * dst, + const std::vector & src_tensors, + const std::vector & sync_tensors, + uint32_t rank, + uint32_t n_ranks, + uint32_t fence_seq_entry, + uint32_t fence_seq_exit +) { + htp_opnode ar_node(HTP_OP_ALLREDUCE); + + ggml_tensor* node = ar_node.add_dummy(*dst); + node->op = GGML_OP_NONE; + node->op_params[0] = (int32_t) fence_seq_entry; + node->op_params[1] = (int32_t) fence_seq_exit; + + ar_node.init(node); + + ar_node.inputs.clear(); + for (size_t i = 0; i < src_tensors.size(); i++) { + ar_node.inputs.push_back(src_tensors[i]); + } + for (size_t i = 0; i < sync_tensors.size(); i++) { + ar_node.inputs.push_back(ar_node.add_dummy(*sync_tensors[i])); + } + + ar_node.outputs.clear(); + for (size_t i = 0; i < src_tensors.size(); i++) { + ar_node.outputs.push_back(src_tensors[i]); + } + + ggml_hexagon_precompute_allreduce_params( + this, dst, rank, n_ranks, false, false, + (struct htp_allreduce_kernel_params *) ar_node.kernel_params + ); + + ar_node.name = "ALLREDUCE"; + this->enqueue_op(ar_node); +} + +void ggml_hexagon_session::wait_event(uint64_t seq) { + flush_sync_peers(); + HEX_VERBOSE("ggml-hex: %s opqueue-wait start: seq %llu, current rsp-seq %llu, pending %d\n", + this->name.c_str(), (unsigned long long)seq, (unsigned long long)op_queue->rsp_seq, (int)this->op_pending); + while (op_queue->rsp_seq < seq && this->op_pending > 0) { + this->flush_pending(false); + } + HEX_VERBOSE("ggml-hex: %s opqueue-wait end: seq %llu, current rsp-seq %llu, pending %d\n", + this->name.c_str(), (unsigned long long)seq, (unsigned long long)op_queue->rsp_seq, (int)this->op_pending); +} + +uint64_t ggml_hexagon_session::record_event() { flush_batch(); - flush_pending(all); + return op_queue->req_seq; +} + +bool ggml_hexagon_session::clone_buffer(const ggml_hexagon_shared_buffer *sbuf) +{ + if (this->cloned_buffers.find(sbuf->fd()) != this->cloned_buffers.end()) return true; + + HEX_VERBOSE("ggml-hex: %s clone-buffer: %s base %p size %zu fd %d\n", this->name.c_str(), + sbuf->c_name(), sbuf->base(), sbuf->size(), sbuf->fd()); + + auto clone = std::make_unique(this, *sbuf); + try { + clone->mmap(); + } catch (const std::exception & exc) { + GGML_LOG_ERROR("ggml-hex: %s lazy mapping of buffer context failed: %s\n", this->c_name(), exc.what()); + return false; + } + + this->cloned_buffers[sbuf->fd()] = std::move(clone); + return true; } static size_t ggml_hexagon_measure_max_vmem(ggml_hexagon_session *sess) { @@ -1668,37 +2812,44 @@ static size_t ggml_hexagon_measure_max_vmem(ggml_hexagon_session *sess) { } void ggml_hexagon_session::allocate(int dev_id) noexcept(false) { + const auto & config = opt_device_configs[dev_id]; + int phys_idx = config.physical_idx; + int virt_idx = config.virtual_idx; + this->valid_session = false; this->valid_handle = false; this->valid_queue = false; this->valid_iface = false; - this->domain_id = 3; // Default for CDSP, updated after the session is created - this->session_id = 0; // Default for CDSP, updated after the session is created + this->phys_idx = phys_idx; + this->virt_idx = virt_idx; + this->domain_id = get_domain_id(phys_idx); + this->session_id = 0; this->dev_id = dev_id; - this->name = std::string("HTP") + std::to_string(dev_id); - - this->op_pending = 0; + this->name = config.name; + this->op_pending = 0; GGML_LOG_DEBUG("ggml-hex: %s allocating new session\n", this->name.c_str()); domain * my_domain = htpdrv_get_domain(this->domain_id); if (my_domain == NULL) { - GGML_LOG_ERROR("ggml-hex: unable to get domain struct for CDSP\n"); + GGML_LOG_ERROR("ggml-hex: unable to get domain struct for CDSP (domain_id %d)\n", this->domain_id); throw std::runtime_error("ggml-hex: failed to get CDSP domain (see log for details)"); } - // Create new session - if (dev_id != 0) { + std::string dom_name = get_domain_name(phys_idx); + + // Create new session if virtual_idx > 0 + if (virt_idx > 0) { struct remote_rpc_reserve_new_session n; - n.domain_name_len = strlen(CDSP_DOMAIN_NAME); - n.domain_name = const_cast(CDSP_DOMAIN_NAME); + n.domain_name_len = dom_name.size(); + n.domain_name = const_cast(dom_name.c_str()); n.session_name = const_cast(this->name.c_str()); n.session_name_len = this->name.size(); int err = remote_session_control(FASTRPC_RESERVE_NEW_SESSION, (void *) &n, sizeof(n)); if (err != AEE_SUCCESS) { - GGML_LOG_ERROR("ggml-hex: failed to reserve new session %d : error 0x%x\n", dev_id, err); + GGML_LOG_ERROR("ggml-hex: failed to reserve new session %d (physical %d, virtual %d) : error 0x%x\n", dev_id, phys_idx, virt_idx, err); throw std::runtime_error("ggml-hex: remote_session_control(new-sess) failed (see log for details)"); } @@ -1717,8 +2868,8 @@ void ggml_hexagon_session::allocate(int dev_id) noexcept(false) { struct remote_rpc_get_uri u = {}; u.session_id = this->session_id; - u.domain_name = const_cast(CDSP_DOMAIN_NAME); - u.domain_name_len = strlen(CDSP_DOMAIN_NAME); + u.domain_name = const_cast(dom_name.c_str()); + u.domain_name_len = dom_name.size(); u.module_uri = const_cast(htp_uri); u.module_uri_len = strlen(htp_uri); u.uri = session_uri; @@ -1731,7 +2882,7 @@ void ggml_hexagon_session::allocate(int dev_id) noexcept(false) { snprintf(session_uri, htp_URI_domain_len, "%s%s", htp_uri, my_domain->uri); - GGML_LOG_WARN("ggml-hex: failed to get URI for session %d : error 0x%x. Falling back to single session URI: %s\n", dev_id, err, session_uri); + GGML_LOG_WARN("ggml-hex: failed to get URI for session %d (physical %d, virtual %d) : error 0x%x. Falling back to single session URI: %s\n", dev_id, phys_idx, virt_idx, err, session_uri); } } @@ -1899,14 +3050,17 @@ void ggml_hexagon_session::release() noexcept(true) { if (this->valid_handle) { htp_iface_close(this->handle); } + + this->cloned_buffers.clear(); } ggml_hexagon_session::ggml_hexagon_session(int dev_id, ggml_backend_dev_t dev) noexcept(false) { - buffer_type.device = dev; - repack_buffer_type.device = dev; + buffer_type.device = dev; + host_buffer_type.device = dev; op_batch = nullptr; op_queue = nullptr; + fence_seq = ((uintptr_t)this) & 0xFFFF; try { allocate(dev_id); @@ -1914,8 +3068,8 @@ ggml_hexagon_session::ggml_hexagon_session(int dev_id, ggml_backend_dev_t dev) n buffer_type.iface = ggml_backend_hexagon_buffer_type_interface; buffer_type.context = new ggml_backend_hexagon_buffer_type_context(this->name, this); - repack_buffer_type.iface = ggml_backend_hexagon_repack_buffer_type_interface; - repack_buffer_type.context = new ggml_backend_hexagon_buffer_type_context(this->name + "-REPACK", this); + host_buffer_type.iface = ggml_backend_hexagon_host_buffer_type_interface; + host_buffer_type.context = new ggml_backend_hexagon_buffer_type_context(this->name + "-HOST", this); } catch (const std::exception & exc) { release(); throw; @@ -1926,7 +3080,7 @@ ggml_hexagon_session::~ggml_hexagon_session() noexcept(true) { release(); delete static_cast(buffer_type.context); - delete static_cast(repack_buffer_type.context); + delete static_cast(host_buffer_type.context); } // ** backend interface @@ -1946,7 +3100,8 @@ static bool ggml_hexagon_flash_attn_is_hmx_eligible( return false; } - if (k->type != GGML_TYPE_F16 || v->type != GGML_TYPE_F16) { + if ((k->type != GGML_TYPE_F16 && k->type != GGML_TYPE_Q8_0) || + (v->type != GGML_TYPE_F16 && v->type != GGML_TYPE_Q8_0)) { return false; } @@ -2098,8 +3253,10 @@ static bool ggml_hexagon_supported_flash_attn_ext(const struct ggml_hexagon_sess const struct ggml_tensor * src4 = op->src[4]; const struct ggml_tensor * dst = op; - // Check for F16 support only as requested - if ((src0->type != GGML_TYPE_F16 && src0->type != GGML_TYPE_F32) || src1->type != GGML_TYPE_F16 || src2->type != GGML_TYPE_F16) { + // Check for F16/Q8_0 support + if ((src0->type != GGML_TYPE_F16 && src0->type != GGML_TYPE_F32) || + (src1->type != GGML_TYPE_F16 && src1->type != GGML_TYPE_Q8_0) || + (src2->type != GGML_TYPE_F16 && src2->type != GGML_TYPE_Q8_0)) { return false; } @@ -2352,7 +3509,7 @@ static void ggml_hexagon_precompute_hvx_mm_params( for (uint32_t d = max_prefetch; d >= 2; d /= 2) { htp_mm_hvx_vtcm_layout_build( &L, kparams->kernel_type, wtype, ne10, src1_nrows, sess->n_threads, - 0, src0->nb[1], 0, src2_row_size, d, true, false, false + 0, src0->nb[1], 0, src2_row_size, d, true, false ); if (L.total_bytes <= vtcm_budget) { best_n_prefetch = d; @@ -2362,7 +3519,7 @@ static void ggml_hexagon_precompute_hvx_mm_params( if (best_n_prefetch == 2 && L.total_bytes > vtcm_budget) { htp_mm_hvx_vtcm_layout_build( &L, kparams->kernel_type, wtype, ne10, src1_nrows, sess->n_threads, - 0, src0->nb[1], 0, src2_row_size, 2, true, false, false + 0, src0->nb[1], 0, src2_row_size, 2, true, false ); } kparams->n_prefetch = best_n_prefetch; @@ -2386,7 +3543,7 @@ static void ggml_hexagon_precompute_hvx_mm_params( for (uint32_t d = max_prefetch; d >= 2; d /= 2) { htp_mm_hvx_vtcm_layout_build( &L, kparams->kernel_type, wtype, ne10, src1_nrows, sess->n_threads, - dst->nb[1], src0->nb[1], src1->nb[1], src2_row_size, d, false, false, false + dst->nb[1], src0->nb[1], src1->nb[1], src2_row_size, d, false, false ); if (L.total_bytes <= vtcm_budget) { best_n_prefetch = d; @@ -2396,7 +3553,7 @@ static void ggml_hexagon_precompute_hvx_mm_params( if (best_n_prefetch == 2 && L.total_bytes > vtcm_budget) { htp_mm_hvx_vtcm_layout_build( &L, kparams->kernel_type, wtype, ne10, src1_nrows, sess->n_threads, - dst->nb[1], src0->nb[1], src1->nb[1], src2_row_size, 2, false, false, false + dst->nb[1], src0->nb[1], src1->nb[1], src2_row_size, 2, false, false ); } @@ -2420,7 +3577,7 @@ static void ggml_hexagon_precompute_hvx_mm_params( struct htp_mm_hvx_vtcm_layout L; htp_mm_hvx_vtcm_layout_build( &L, kparams->kernel_type, wtype, ne10, src1_nrows, sess->n_threads, - dst->nb[1], src0->nb[1], src1->nb[1], src2_row_size, 16, false, false, false + dst->nb[1], src0->nb[1], src1->nb[1], src2_row_size, 16, false, false ); kparams->n_prefetch = 16; @@ -2440,7 +3597,7 @@ static void ggml_hexagon_precompute_hvx_mm_params( struct htp_mm_hvx_vtcm_layout L; htp_mm_hvx_vtcm_layout_build( &L, HTP_MM_KERNEL_HVX_F16_F16_VTCM, wtype, ne10, src1_nrows, sess->n_threads, - dst->nb[1], src0->nb[1], src1->nb[1], src2_row_size, 16, false, false, false + dst->nb[1], src0->nb[1], src1->nb[1], src2_row_size, 16, false, false ); if (!is_batched && !is_permuted && L.total_bytes <= vtcm_budget) { @@ -2460,7 +3617,7 @@ static void ggml_hexagon_precompute_hvx_mm_params( kparams->src1_row_size = src1->nb[1]; htp_mm_hvx_vtcm_layout_build( &L, kparams->kernel_type, wtype, ne10, src1_nrows, sess->n_threads, - dst->nb[1], src0->nb[1], src1->nb[1], src2_row_size, 16, false, false, false + dst->nb[1], src0->nb[1], src1->nb[1], src2_row_size, 16, false, false ); kparams->vtcm_size = L.total_bytes; kparams->vtcm_src0_size = L.src0_bytes; @@ -2476,7 +3633,7 @@ static void ggml_hexagon_precompute_hvx_mm_params( struct htp_mm_hvx_vtcm_layout L; htp_mm_hvx_vtcm_layout_build( &L, HTP_MM_KERNEL_HVX_F32_F32_VTCM, wtype, ne10, src1_nrows, sess->n_threads, - dst->nb[1], src0->nb[1], src1->nb[1], src2_row_size, 16, false, false, false + dst->nb[1], src0->nb[1], src1->nb[1], src2_row_size, 16, false, false ); if (!is_batched && !is_permuted && L.total_bytes <= vtcm_budget) { @@ -2492,7 +3649,7 @@ static void ggml_hexagon_precompute_hvx_mm_params( kparams->src1_row_size = src1->nb[1]; htp_mm_hvx_vtcm_layout_build( &L, kparams->kernel_type, wtype, ne10, src1_nrows, sess->n_threads, - dst->nb[1], src0->nb[1], src1->nb[1], src2_row_size, 16, false, false, false + dst->nb[1], src0->nb[1], src1->nb[1], src2_row_size, 16, false, false ); kparams->vtcm_size = L.total_bytes; kparams->vtcm_src0_size = L.src0_bytes; @@ -2642,80 +3799,99 @@ static void ggml_hexagon_precompute_unary_params( kparams->div_tpr = init_fastdiv_values(tiles_per_row); } -static void ggml_hexagon_precompute_fused_qkv_params( +static void ggml_hexagon_precompute_get_rows_params( const struct ggml_hexagon_session * sess, - const struct ggml_tensor * src0, // Wk - const struct ggml_tensor * src1, // x - struct htp_mm_kernel_params * kparams + const struct ggml_tensor * src0, + const struct ggml_tensor * src1, + const struct ggml_tensor * dst, + struct htp_get_rows_kernel_params * kparams ) { memset(kparams, 0, sizeof(*kparams)); - const int wtype = src0->type; - const bool is_repack = ggml_hexagon_is_repack_type((ggml_type) wtype); + const uint32_t ne00 = src0->ne[0]; + const uint32_t ne02 = src0->ne[2]; + const uint32_t ne03 = src0->ne[3]; - const int ne10 = src1->ne[0]; - const int src1_nrows = src1->ne[1] * src1->ne[2] * src1->ne[3]; - const size_t src1_row_size = (wtype == GGML_TYPE_Q4_1) ? htp_mm_q8_1_tiled_row_size(ne10) : htp_mm_q8_0_tiled_row_size(ne10); - const size_t src0_row_size = src0->nb[1]; + const uint32_t ne10 = src1->ne[0]; + const uint32_t ne11 = src1->ne[1]; + const uint32_t ne12 = src1->ne[2]; + const uint32_t nr = ne10 * ne11 * ne12; - uint32_t best_n_prefetch = 16; + const size_t nb01 = src0->nb[1]; + const size_t nb1 = dst->nb[1]; - if (is_repack) { - const uint32_t max_prefetch = (src1_nrows > HTP_MM_HMX_MIN_NROWS) ? 2 : 16; - best_n_prefetch = 2; - for (uint32_t d = max_prefetch; d >= 2; d /= 2) { - struct htp_mm_hvx_vtcm_layout L; - htp_mm_hvx_vtcm_layout_build( - &L, HTP_MM_KERNEL_HVX_QUANT_ROW, wtype, ne10, src1_nrows, sess->n_threads, - 0, src0_row_size, src1_row_size, 0, d, false, true, false - ); - if (L.total_bytes <= sess->vtcm_size) { - best_n_prefetch = d; - break; + const bool can_use_dma = (src0->type == dst->type) && (nb01 == nb1); + const bool use_dma = can_use_dma && (ne00 >= 2048); + + kparams->use_dma = use_dma ? 1 : 0; + + uint32_t chunks_per_row = 1; + uint32_t chunk_size = ne00; + uint32_t total_tasks = nr; + + if (use_dma) { + kparams->n_threads = (std::min)((uint32_t)sess->n_threads, nr); + kparams->tasks_per_thread = (nr + kparams->n_threads - 1) / kparams->n_threads; + } else { + if (src0->type == GGML_TYPE_F32 && nr < sess->n_threads) { + const uint32_t min_chunk_size = 1024; + uint32_t max_chunks = ne00 / min_chunk_size; + if (max_chunks == 0) { + max_chunks = 1; } + chunks_per_row = (std::min)((sess->n_threads + nr - 1) / nr, max_chunks); + chunk_size = (ne00 + chunks_per_row - 1) / chunks_per_row; + total_tasks = nr * chunks_per_row; } + kparams->n_threads = (std::min)(total_tasks, (uint32_t)sess->n_threads); + kparams->tasks_per_thread = (total_tasks + kparams->n_threads - 1) / kparams->n_threads; } - struct htp_mm_hvx_vtcm_layout L; - bool try_tiled = (opt_mm_select >= 2); + kparams->chunks_per_row = chunks_per_row; + kparams->chunk_size = chunk_size; + kparams->total_tasks = total_tasks; - // Test tiled first - htp_mm_hvx_vtcm_layout_build( - &L, HTP_MM_KERNEL_HVX_QUANT_ROW, wtype, ne10, src1_nrows, sess->n_threads, - 0, src0_row_size, src1_row_size, 0, best_n_prefetch, false, true, false - ); + kparams->div_ne10 = init_fastdiv_values(ne10); + kparams->div_ne10_ne11 = init_fastdiv_values(ne10 * ne11); + kparams->div_chunks_per_row = init_fastdiv_values(chunks_per_row); + kparams->div_ne02 = init_fastdiv_values(ne02); + kparams->div_ne03 = init_fastdiv_values(ne03); - if (try_tiled && L.total_bytes <= sess->vtcm_size) { - kparams->kernel_type = HTP_MM_KERNEL_HVX_QUANT_ROW; - kparams->vtcm_src0_size = L.src0_bytes; - kparams->vtcm_src1_size = L.src1_bytes; - kparams->vtcm_src2_size = L.src2_bytes; - kparams->vtcm_src3_size = L.src3_bytes; - kparams->vtcm_dst_size = L.dst_bytes; - kparams->vtcm_size = L.total_bytes; - kparams->n_prefetch = best_n_prefetch; - } else { - kparams->kernel_type = HTP_MM_KERNEL_HVX_QUANT_ROW_FLAT; - size_t flat_src1_row_size = (wtype == GGML_TYPE_Q4_1) ? htp_mm_q8_1_flat_row_size(ne10) : htp_mm_q8_0_flat_row_size(ne10); + struct htp_get_rows_vtcm_layout vtcm_layout; + htp_get_rows_vtcm_layout_build(&vtcm_layout, src0->type, ne00, kparams->n_threads); + kparams->vtcm_size = vtcm_layout.total_bytes; +} - htp_mm_hvx_vtcm_layout_build( - &L, HTP_MM_KERNEL_HVX_QUANT_ROW_FLAT, wtype, ne10, src1_nrows, sess->n_threads, - 0, src0_row_size, flat_src1_row_size, 0, best_n_prefetch, false, true, false - ); - kparams->vtcm_src0_size = L.src0_bytes; - kparams->vtcm_src1_size = L.src1_bytes; - kparams->vtcm_src2_size = L.src2_bytes; - kparams->vtcm_src3_size = L.src3_bytes; - kparams->vtcm_dst_size = L.dst_bytes; - kparams->vtcm_size = L.total_bytes; - kparams->n_prefetch = best_n_prefetch; - } +static void ggml_hexagon_precompute_set_rows_params( + const struct ggml_hexagon_session * sess, + const struct ggml_tensor * src0, // values + const struct ggml_tensor * src1, // indices + const struct ggml_tensor * dst, // destination + struct htp_set_rows_kernel_params * kparams +) { + memset(kparams, 0, sizeof(*kparams)); + + const uint32_t nr = src0->ne[1]; + + kparams->n_threads = (std::min)((uint32_t)sess->n_threads, nr); + kparams->tasks_per_thread = (nr + kparams->n_threads - 1) / kparams->n_threads; + kparams->total_tasks = nr; + + kparams->div_ne11 = init_fastdiv_values(src1->ne[1]); + kparams->div_ne12 = init_fastdiv_values(src1->ne[2]); + kparams->div_tasks_per_thread = init_fastdiv_values(kparams->tasks_per_thread); + kparams->div_ne02 = init_fastdiv_values(src0->ne[2]); + + struct htp_set_rows_vtcm_layout vtcm_layout; + htp_set_rows_vtcm_layout_build(&vtcm_layout, dst->type, src0->ne[0], kparams->n_threads); + kparams->vtcm_size = vtcm_layout.total_bytes; } -static void ggml_hexagon_precompute_fused_ffn_params( +static void ggml_hexagon_precompute_fused_mmnx_params( const struct ggml_hexagon_session * sess, - const struct ggml_tensor * src0, // Wgate - const struct ggml_tensor * src1, // y + const struct ggml_tensor * src0, // W0 + const struct ggml_tensor * src1, // x + int32_t n_weights, struct htp_mm_kernel_params * kparams ) { memset(kparams, 0, sizeof(*kparams)); @@ -2737,7 +3913,7 @@ static void ggml_hexagon_precompute_fused_ffn_params( struct htp_mm_hvx_vtcm_layout L; htp_mm_hvx_vtcm_layout_build( &L, HTP_MM_KERNEL_HVX_QUANT_ROW, wtype, ne10, src1_nrows, sess->n_threads, - 0, src0_row_size, src1_row_size, 0, d, false, false, true + 0, src0_row_size, src1_row_size, 0, d, false, true ); if (L.total_bytes <= sess->vtcm_size) { best_n_prefetch = d; @@ -2752,34 +3928,42 @@ static void ggml_hexagon_precompute_fused_ffn_params( // Test tiled first htp_mm_hvx_vtcm_layout_build( &L, HTP_MM_KERNEL_HVX_QUANT_ROW, wtype, ne10, src1_nrows, sess->n_threads, - 0, src0_row_size, src1_row_size, 0, best_n_prefetch, false, false, true + 0, src0_row_size, src1_row_size, 0, best_n_prefetch, false, true ); if (try_tiled && L.total_bytes <= sess->vtcm_size) { kparams->kernel_type = HTP_MM_KERNEL_HVX_QUANT_ROW; kparams->vtcm_src0_size = L.src0_bytes; kparams->vtcm_src1_size = L.src1_bytes; - kparams->vtcm_src2_size = L.src2_bytes; kparams->vtcm_dst_size = L.dst_bytes; kparams->vtcm_size = L.total_bytes; kparams->n_prefetch = best_n_prefetch; + kparams->n_weights = n_weights; } else { kparams->kernel_type = HTP_MM_KERNEL_HVX_QUANT_ROW_FLAT; size_t flat_src1_row_size = (wtype == GGML_TYPE_Q4_1) ? htp_mm_q8_1_flat_row_size(ne10) : htp_mm_q8_0_flat_row_size(ne10); htp_mm_hvx_vtcm_layout_build( &L, HTP_MM_KERNEL_HVX_QUANT_ROW_FLAT, wtype, ne10, src1_nrows, sess->n_threads, - 0, src0_row_size, flat_src1_row_size, 0, best_n_prefetch, false, false, true + 0, src0_row_size, flat_src1_row_size, 0, best_n_prefetch, false, true ); kparams->vtcm_src0_size = L.src0_bytes; kparams->vtcm_src1_size = L.src1_bytes; - kparams->vtcm_src2_size = L.src2_bytes; kparams->vtcm_dst_size = L.dst_bytes; kparams->vtcm_size = L.total_bytes; kparams->n_prefetch = best_n_prefetch; + kparams->n_weights = n_weights; } } +static bool ggml_hexagon_tensor_is_host(const struct ggml_hexagon_session * sess, const struct ggml_tensor * t) { + return t && t->buffer && t->buffer->buft == &sess->host_buffer_type; +} + +static bool ggml_hexagon_tensor_is_non_host(const struct ggml_hexagon_session * sess, const struct ggml_tensor * t) { + return t && t->buffer && t->buffer->buft != &sess->host_buffer_type; +} + static bool ggml_hexagon_supported_mul_mat(const struct ggml_hexagon_session * sess, const struct ggml_tensor * dst) { const struct ggml_tensor * src0 = dst->src[0]; const struct ggml_tensor * src1 = dst->src[1]; @@ -2811,9 +3995,8 @@ static bool ggml_hexagon_supported_mul_mat(const struct ggml_hexagon_session * s return false; // no broadcasting (for now) } - // src0 (weights) must be repacked - if (src0->buffer && !ggml_backend_buffer_is_hexagon_repack(src0->buffer)) { - return false; + if (!src0->buffer) { + sess->needs_repack.insert(src0); } break; @@ -2872,9 +4055,8 @@ static bool ggml_hexagon_supported_mul_mat_id(const struct ggml_hexagon_session return false; } - // src0 (weights) must be repacked - if (src0->buffer && !ggml_backend_buffer_is_hexagon_repack(src0->buffer)) { - return false; + if (!src0->buffer) { + sess->needs_repack.insert(src0); } break; @@ -2970,7 +4152,7 @@ static bool ggml_hexagon_supported_unary(const struct ggml_hexagon_session * ses if (dst->type != GGML_TYPE_F32) { return false; } - if (ggml_is_permuted(src0)) { + if (!ggml_is_contiguous_rows(src0)) { return false; } if (!ggml_are_same_shape(src0, dst)) { @@ -3114,7 +4296,11 @@ static bool ggml_hexagon_supported_softmax(const struct ggml_hexagon_session * s static bool ggml_hexagon_supported_set_rows(const struct ggml_hexagon_session * sess, const struct ggml_tensor * op) { const struct ggml_tensor * src0 = op->src[0]; // values const struct ggml_tensor * src1 = op->src[1]; // indices - const struct ggml_tensor * dst = op; + const struct ggml_tensor * dst = op->src[2] ? op->src[2] : op; + + if (dst->type == GGML_TYPE_Q8_0 && src0->ne[0] < 32) { + return false; + } if (src0->type != GGML_TYPE_F32) { return false; @@ -3124,7 +4310,7 @@ static bool ggml_hexagon_supported_set_rows(const struct ggml_hexagon_session * return false; } - if (dst->type != GGML_TYPE_F16) { + if (dst->type != GGML_TYPE_F32 && dst->type != GGML_TYPE_F16 && dst->type != GGML_TYPE_Q8_0) { return false; } @@ -3138,7 +4324,11 @@ static bool ggml_hexagon_supported_get_rows(const struct ggml_hexagon_session * const struct ggml_tensor * src1 = op->src[1]; // indices const struct ggml_tensor * dst = op; - if (src0->type != GGML_TYPE_F32) { + if (src0->type != GGML_TYPE_F32 && src0->ne[0] < 32) { + return false; + } + + if (src0->type != GGML_TYPE_F32 && src0->type != GGML_TYPE_F16 && src0->type != GGML_TYPE_Q8_0) { return false; } @@ -3521,133 +4711,22 @@ static bool is_mergeable_mul_mat(const ggml_tensor * t) { if (!t || t->op != GGML_OP_MUL_MAT) return false; if (t->src[1]->type != GGML_TYPE_F32) return false; return ggml_is_quantized(t->src[0]->type) && !mm_is_hmx_eligible(t); -} - -static bool is_mergeable_mul_mat_pair(const ggml_tensor * n1, const ggml_tensor * n2) { - if (!is_mergeable_mul_mat(n1) || !is_mergeable_mul_mat(n2)) { - return false; - } - if (n1->src[1] != n2->src[1]) { - return false; - } - if (n1->src[0]->ne[0] != n2->src[0]->ne[0] || - n1->src[0]->ne[1] != n2->src[0]->ne[1]) { - return false; - } - if (n1->src[0]->type != n2->src[0]->type) { - return false; - } - return true; -} - -static bool is_qkv_mergeable(const ggml_tensor * n_q, const ggml_tensor * n_k, const ggml_tensor * n_v) { - if (!is_mergeable_mul_mat(n_q) || !is_mergeable_mul_mat(n_k) || !is_mergeable_mul_mat(n_v)) { - return false; - } - if (n_q->src[1] != n_k->src[1] || n_q->src[1] != n_v->src[1]) { - return false; - } - if (n_q->src[0]->type != n_k->src[0]->type || n_q->src[0]->type != n_v->src[0]->type) { - return false; - } - if (n_k->src[0]->ne[0] != n_v->src[0]->ne[0] || - n_k->src[0]->ne[1] != n_v->src[0]->ne[1]) { - return false; - } - if (n_q->src[0]->ne[0] != n_k->src[0]->ne[0]) { - return false; - } - return true; -} - -static bool try_fuse_node(const ggml_hexagon_session * sess, const ggml_cgraph * graph, int & i, std::vector & nodes) { - if (!opt_opfusion) { - return false; - } - - ggml_tensor * n = graph->nodes[i]; - ggml_tensor * next_node = (i + 1 < graph->n_nodes) ? graph->nodes[i + 1] : nullptr; - - if (n->op == GGML_OP_RMS_NORM && next_node) { - if (next_node->op == GGML_OP_MUL && op_is_compute(next_node) && ggml_can_fuse(graph, i, { GGML_OP_RMS_NORM, GGML_OP_MUL })) { - htp_opnode node(n, {}, HTP_OP_RMS_NORM_MUL); - node.add_fused(next_node); - - auto inputs = node.get_inputs(); - const struct ggml_tensor * src0 = inputs[0]; - const struct ggml_tensor * src1 = inputs.size() > 1 ? inputs[1] : nullptr; - ggml_hexagon_precompute_unary_params(sess, - node.opcode, src0, src1, node.dst(), - (struct htp_unary_kernel_params *)node.kernel_params - ); - - nodes.push_back(std::move(node)); - i++; // skip the fused MUL node - return true; - } - } - - if (is_mergeable_mul_mat(n)) { - ggml_tensor * n1 = (i + 1 < graph->n_nodes) ? graph->nodes[i + 1] : nullptr; - ggml_tensor * n2 = (i + 2 < graph->n_nodes) ? graph->nodes[i + 2] : nullptr; - if (is_qkv_mergeable(n, n1, n2)) { - struct htp_mm_kernel_params kparams; - ggml_hexagon_precompute_fused_qkv_params(sess, n1->src[0], n1->src[1], &kparams); - if ((size_t)kparams.vtcm_size <= sess->vtcm_size) { - // Reorder to KVQ: K (n1), V (n2), Q (n) - htp_opnode node(n1, {}, HTP_OP_MUL_MAT_QKV); - node.add_fused(n2, true); - node.add_fused(n, true); - memcpy(node.kernel_params, &kparams, sizeof(kparams)); - nodes.push_back(std::move(node)); - i += 2; - return true; - } else { - HEX_VERBOSE("ggml-hex: skip QKV fusion because VTCM needed (%d) > budget (%zu)\n", - kparams.vtcm_size, sess->vtcm_size); - } - } - if (is_mergeable_mul_mat_pair(n, n1)) { - struct htp_mm_kernel_params kparams; - ggml_hexagon_precompute_fused_ffn_params(sess, n->src[0], n->src[1], &kparams); - if ((size_t)kparams.vtcm_size <= sess->vtcm_size) { - htp_opnode node(n, {}, HTP_OP_MUL_MAT_FFN); - node.add_fused(n1, true); - memcpy(node.kernel_params, &kparams, sizeof(kparams)); - nodes.push_back(std::move(node)); - i += 1; - return true; - } else { - HEX_VERBOSE("ggml-hex: skip FFN fusion because VTCM needed (%d) > budget (%zu)\n", - kparams.vtcm_size, sess->vtcm_size); - } - } - } +} - if (n->op == GGML_OP_MUL_MAT && next_node) { - if (next_node->op == GGML_OP_ADD && op_is_compute(next_node) && ggml_can_fuse(graph, i, { GGML_OP_MUL_MAT, GGML_OP_ADD })) { - if (next_node->src[0] == n || next_node->src[1] == n) { - const struct ggml_tensor * src2 = (next_node->src[0] == n) ? next_node->src[1] : next_node->src[0]; - struct htp_mm_kernel_params kparams; - ggml_hexagon_precompute_fused_matmul_add_params(sess, n->src[0], n->src[1], src2, next_node, &kparams); - const int src1_nrows = n->src[1]->ne[1] * n->src[1]->ne[2] * n->src[1]->ne[3]; - const bool can_fuse = (kparams.n_hmx > 0) || (src1_nrows == 1); - if (can_fuse && (size_t)kparams.vtcm_size <= sess->vtcm_size) { - htp_opnode node(n, {}, HTP_OP_MUL_MAT_ADD); - node.add_fused(next_node); - memcpy(node.kernel_params, &kparams, sizeof(kparams)); - nodes.push_back(std::move(node)); - i += 1; - return true; - } else if (can_fuse) { - HEX_VERBOSE("ggml-hex: skip MUL_MAT_ADD fusion because VTCM needed (%d) > budget (%zu)\n", - kparams.vtcm_size, sess->vtcm_size); - } - } - } +static bool is_mergeable_mul_mat_pair(const ggml_tensor * n1, const ggml_tensor * n2) { + if (!is_mergeable_mul_mat(n1) || !is_mergeable_mul_mat(n2)) { + return false; } - - return false; + if (n1->src[1] != n2->src[1]) { + return false; + } + if (n1->src[0]->ne[0] != n2->src[0]->ne[0]) { + return false; + } + if (n1->src[0]->type != n2->src[0]->type) { + return false; + } + return true; } static ggml_status ggml_backend_hexagon_graph_compute(ggml_backend_t backend, ggml_cgraph * graph) { @@ -3659,24 +4738,34 @@ static ggml_status ggml_backend_hexagon_graph_compute(ggml_backend_t backend, gg std::vector computed_nodes; // Check for cache hit - bool cache_hit = (graph->uid != 0 && sess->cached_graph.uid == graph->uid); + bool cache_hit = (graph->uid != 0 && sess->cached_uid == graph->uid); if (cache_hit) { - nodes_ptr = &sess->cached_graph.htp_nodes; + nodes_ptr = &sess->cached_nodes; } else { + // Tag fusable tensors in graph + for (int i = 0; i < graph->n_nodes; i++) { + auto * extra = (ggml_hexagon_tensor_extra *) graph->nodes[i]->extra; + if (!extra) continue; + + if (graph->nodes[i]->op == GGML_OP_RMS_NORM && ggml_can_fuse(graph, i, { GGML_OP_RMS_NORM, GGML_OP_MUL })) { + extra->flags |= GGML_HEXAGON_TENSOR_FUSEABLE; + } else if (graph->nodes[i]->op == GGML_OP_MUL_MAT) { + if ((i + 1 < graph->n_nodes && graph->nodes[i + 1]->op == GGML_OP_ADD && ggml_can_fuse(graph, i, { GGML_OP_MUL_MAT, GGML_OP_ADD })) || + ggml_node_has_n_uses(graph, i, 1)) { + extra->flags |= GGML_HEXAGON_TENSOR_FUSEABLE; + } + } + } + computed_nodes.reserve(graph->n_nodes); - // Fuse and finalize for (int i = 0; i < graph->n_nodes; ++i) { ggml_tensor * n = graph->nodes[i]; if (!op_is_compute(n)) { continue; } - if (try_fuse_node(sess, graph, i, computed_nodes)) { - continue; - } - - htp_opnode node(n, {}, HTP_OP_INVALID); + htp_opnode node(HTP_OP_INVALID, n); node.opcode = op_remap_to_htp(n); if (node.opcode == HTP_OP_MUL_MAT || node.opcode == HTP_OP_MUL_MAT_ID) { ggml_hexagon_precompute_matmul_params(sess, @@ -3696,29 +4785,34 @@ static ggml_status ggml_backend_hexagon_graph_compute(ggml_backend_t backend, gg node.opcode, src0, src1, node.dst(), (struct htp_unary_kernel_params *)node.kernel_params ); + } else if (node.opcode == HTP_OP_GET_ROWS) { + ggml_hexagon_precompute_get_rows_params(sess, + node.node->src[0], node.node->src[1], node.dst(), + (struct htp_get_rows_kernel_params *)node.kernel_params + ); + } else if (node.opcode == HTP_OP_SET_ROWS) { + ggml_hexagon_precompute_set_rows_params(sess, + node.node->src[0], node.node->src[1], node.dst(), + (struct htp_set_rows_kernel_params *)node.kernel_params + ); } computed_nodes.push_back(std::move(node)); } if (graph->uid != 0) { - sess->cached_graph.uid = graph->uid; - sess->cached_graph.htp_nodes = std::move(computed_nodes); - nodes_ptr = &sess->cached_graph.htp_nodes; + sess->cached_uid = graph->uid; + sess->cached_nodes = std::move(computed_nodes); + nodes_ptr = &sess->cached_nodes; } else { nodes_ptr = &computed_nodes; } } // Queue and execute - if (opt_opstage & HTP_OPSTAGE_QUEUE) { - for (const auto & node : *nodes_ptr) { - sess->enqueue_op(node); - } + for (const auto & node : *nodes_ptr) { + sess->enqueue_op(node); } - // Wait until all pending ops complete - sess->flush(); - return GGML_STATUS_SUCCESS; } @@ -3731,6 +4825,106 @@ static void ggml_backend_hexagon_synchronize(ggml_backend_t backend) { sess->flush(); } +enum ggml_hexagon_mem_range_type { + HEXAGON_MEM_RANGE_TYPE_SRC, + HEXAGON_MEM_RANGE_TYPE_DST, +}; + +struct ggml_hexagon_mem_range { + uint64_t pb; + uint64_t p0; + uint64_t p1; + ggml_hexagon_mem_range_type pt; +}; + +struct ggml_hexagon_mem_ranges { + std::vector ranges; + + void reset() { + ranges.clear(); + } + + void add(const ggml_hexagon_mem_range & mr) { + ranges.push_back(mr); + } + + bool check(const ggml_hexagon_mem_range & mr) const { + for (const auto & cmp : ranges) { + if (mr.pb != cmp.pb) { + continue; + } + if (mr.pt == HEXAGON_MEM_RANGE_TYPE_SRC && cmp.pt == HEXAGON_MEM_RANGE_TYPE_SRC) { + continue; + } + if (mr.p0 < cmp.p1 && mr.p1 > cmp.p0) { + return false; + } + } + return true; + } +}; + +static ggml_hexagon_mem_range ggml_hexagon_mem_range_from_tensor(const ggml_tensor * tensor, ggml_hexagon_mem_range_type pt) { + const ggml_tensor * base = tensor->view_src ? tensor->view_src : tensor; + ggml_hexagon_mem_range mr; + if (tensor->buffer) { + mr = { + /*.pb =*/ (uint64_t) tensor->buffer, + /*.p0 =*/ (uint64_t) tensor->data, + /*.p1 =*/ (uint64_t) tensor->data + ggml_backend_buft_get_alloc_size(tensor->buffer->buft, tensor), + /*.pt =*/ pt, + }; + } else { + mr = { + /*.pb =*/ (uint64_t) base, + /*.p0 =*/ 0, + /*.p1 =*/ 1024, + /*.pt =*/ pt, + }; + } + return mr; +} + +static void ggml_hexagon_mem_ranges_add_node(ggml_hexagon_mem_ranges & mrs, const htp_opnode & node) { + if (node.is_empty()) return; + + for (int i = 0; i < GGML_MAX_SRC; i++) { + if (node.node->src[i]) { + mrs.add(ggml_hexagon_mem_range_from_tensor(node.node->src[i], HEXAGON_MEM_RANGE_TYPE_SRC)); + } + } + for (const auto * fused : node.fused) { + for (int i = 0; i < GGML_MAX_SRC; i++) { + if (fused->src[i]) { + mrs.add(ggml_hexagon_mem_range_from_tensor(fused->src[i], HEXAGON_MEM_RANGE_TYPE_SRC)); + } + } + } + mrs.add(ggml_hexagon_mem_range_from_tensor(node.dst(), HEXAGON_MEM_RANGE_TYPE_DST)); +} + +static bool ggml_hexagon_mem_ranges_check_node(const ggml_hexagon_mem_ranges & mrs, const htp_opnode & node) { + if (node.is_empty()) return true; + + for (int i = 0; i < GGML_MAX_SRC; i++) { + if (node.node->src[i]) { + if (!mrs.check(ggml_hexagon_mem_range_from_tensor(node.node->src[i], HEXAGON_MEM_RANGE_TYPE_SRC))) { + return false; + } + } + } + for (const auto * fused : node.fused) { + for (int i = 0; i < GGML_MAX_SRC; i++) { + if (fused->src[i]) { + if (!mrs.check(ggml_hexagon_mem_range_from_tensor(fused->src[i], HEXAGON_MEM_RANGE_TYPE_SRC))) { + return false; + } + } + } + } + return mrs.check(ggml_hexagon_mem_range_from_tensor(node.dst(), HEXAGON_MEM_RANGE_TYPE_DST)); +} + static std::vector ggml_hexagon_graph_optimize_reorder(const std::vector & nodes) { const int n = nodes.size(); @@ -3739,28 +4933,32 @@ static std::vector ggml_hexagon_graph_optimize_reorder(const std::vector used(n, false); - // The main goal here is to stack the MUL_MAT ops with the same src1 input. - // This allows use to reuse dynamically quantized src1 in VTCM. + ggml_hexagon_mem_ranges mrs; - // TODO: the current version might do incorrect reordering in cases where quantized src0 - // input is an output of another Op. + // The main goal here is to stack the MUL_MAT ops with the same src1 input. + // This allows us to reuse dynamically quantized src1 in VTCM. for (int i0 = 0; i0 < n; i0++) { if (used[i0]) { continue; } - res.push_back(i0); - const auto & node0 = nodes[i0]; if (!node0.stackable()) { + res.push_back(i0); + used[i0] = true; continue; } // that many nodes forward to search for stackable nodes that can reuse VTCM constexpr int N_FORWARD = 16; + std::vector stack; + stack.push_back(i0); + + mrs.reset(); + for (int i1 = i0 + 1; i1 < i0 + N_FORWARD && i1 < n; i1++) { if (used[i1]) { continue; @@ -3768,11 +4966,17 @@ static std::vector ggml_hexagon_graph_optimize_reorder(const std::vector nodes; nodes.reserve(gf->n_nodes); - // fuse nodes: - // we don't want to make reorders that break fusing, so we first pack all fusable tensors - // and perform the reorder over the fused nodes. after the reorder is done, we unfuse + // Pack nodes for reordering for (int i = 0; i < n; i++) { - htp_opnode node = { - /*.node =*/gf->nodes[i], - /*.fused =*/{}, - }; + htp_opnode node(HTP_OP_INVALID, gf->nodes[i]); // fuse only ops that start with these operations // can be expanded when needed @@ -3856,22 +5055,200 @@ static void ggml_backend_hexagon_graph_optimize(ggml_backend_t backend, ggml_cgr GGML_UNUSED(backend); } +static bool ggml_hexagon_cpy_tensor_async_phys(ggml_backend_t backend_src, ggml_backend_t backend_dst, const ggml_tensor * src, ggml_tensor * dst) { + auto sess_src = static_cast(backend_src->context); + auto sess_dst = static_cast(backend_dst->context); + auto sbuf_dst = (ggml_hexagon_shared_buffer *) dst->buffer->context; + + if (sess_dst->fence_seq == 0) sess_dst->fence_seq = 1; + uint32_t fence_seq = sess_dst->fence_seq++; + if (sess_dst->fence_seq == 0) sess_dst->fence_seq = 1; + + volatile uint32_t * fence = (volatile uint32_t *) sbuf_dst->alloc_fence(); + + HEX_VERBOSE("ggml-hex: %s cpy-tensor-async %s -> %s size %zu : seq %u\n", + sess_dst->name.c_str(), src->name, dst->name, ggml_nbytes(src), fence_seq); + + // dummy extra (must be static) + static ggml_hexagon_tensor_extra fence_extra { {}, 0, GGML_HEXAGON_TENSOR_FENCE }; + + ggml_tensor fence_tensor {}; + fence_tensor.buffer = dst->buffer; + fence_tensor.extra = &fence_extra; + fence_tensor.data = (void *) fence; + fence_tensor.type = GGML_TYPE_I32; + fence_tensor.ne[0] = 1; + fence_tensor.ne[1] = 1; + fence_tensor.ne[2] = 1; + fence_tensor.ne[3] = 1; + fence_tensor.nb[0] = sizeof(int32_t); + fence_tensor.nb[1] = sizeof(int32_t); + fence_tensor.nb[2] = sizeof(int32_t); + fence_tensor.nb[3] = sizeof(int32_t); + fence_tensor.op = GGML_OP_NONE; + + sess_src->enqueue_cpy(src, dst, &fence_tensor, fence_seq); + sess_dst->enqueue_fence(&fence_tensor, fence_seq); + + sess_dst->add_sync_peer(sess_src); + + return true; +} + +static bool ggml_hexagon_cpy_tensor_async_virt(ggml_backend_t backend_src, ggml_backend_t backend_dst, const ggml_tensor * src, ggml_tensor * dst) { + auto sess_src = static_cast(backend_src->context); + auto sess_dst = static_cast(backend_dst->context); + auto sbuf_dst = (ggml_hexagon_shared_buffer *) dst->buffer->context; + + if (!sess_src->clone_buffer(sbuf_dst)) { return false; } + + HEX_VERBOSE("ggml-hex: %s cpy-tensor-async %s -> %s size %zu\n", + sess_dst->name.c_str(), src->name, dst->name, ggml_nbytes(src)); + + sess_src->enqueue_cpy(src, dst); + sess_src->flush(true); + + return true; +} + +static bool ggml_backend_hexagon_cpy_tensor_async(ggml_backend_t backend_src, ggml_backend_t backend_dst, const ggml_tensor * src, ggml_tensor * dst) { + if (!ggml_backend_is_hexagon(backend_src) || !ggml_backend_is_hexagon(backend_dst)) { + return false; + } + + *(ggml_hexagon_tensor_extra *) dst->extra = *(const ggml_hexagon_tensor_extra *) src->extra; + + auto sess_src = static_cast(backend_src->context); + auto sess_dst = static_cast(backend_dst->context); + + if (sess_src == sess_dst) { + HEX_VERBOSE("ggml-hex: %s cpy-tensor-async %s -> %s size %zu\n", sess_dst->name.c_str(), src->name, dst->name, ggml_nbytes(src)); + sess_src->enqueue_cpy(src, dst); + sess_src->flush_batch(); + return true; + } + + if (sess_src->phys_idx != sess_dst->phys_idx) + return ggml_hexagon_cpy_tensor_async_phys(backend_src, backend_dst, src, dst); + + return ggml_hexagon_cpy_tensor_async_virt(backend_src, backend_dst, src, dst); +} + +static ggml_backend_event_t ggml_backend_hexagon_device_event_new(ggml_backend_dev_t dev) { + ggml_hexagon_event * hex_event = new ggml_hexagon_event(); + HEX_VERBOSE("ggml-hex: %s event-new : event %p\n", ggml_backend_dev_name(dev), (void *)hex_event); + + return new ggml_backend_event { + /* .device = */ dev, + /* .context = */ hex_event, + }; +} + +static void ggml_backend_hexagon_device_event_free(ggml_backend_dev_t dev, ggml_backend_event_t event) { + GGML_UNUSED(dev); + + if (event == nullptr) { + return; + } + + ggml_hexagon_event * hex_event = (ggml_hexagon_event *)event->context; + HEX_VERBOSE("ggml-hex: %s event-free : event %p\n", ggml_backend_dev_name(dev), (void *)hex_event); + delete hex_event; + delete event; +} + +static void ggml_backend_hexagon_device_event_synchronize(ggml_backend_dev_t dev, ggml_backend_event_t event) { + GGML_UNUSED(dev); + + ggml_hexagon_event * hex_event = (ggml_hexagon_event *)event->context; + HEX_VERBOSE("ggml-hex: %s event-synchronize : event %p seq %llu\n", + ggml_backend_dev_name(dev), (void *)hex_event, (unsigned long long)hex_event->seq); + if (hex_event->sess != nullptr) { + hex_event->sess->wait_event(hex_event->seq); + } +} + +static void ggml_backend_hexagon_event_record(ggml_backend_t backend, ggml_backend_event_t event) { + auto sess = static_cast(backend->context); + ggml_hexagon_event * hex_event = (ggml_hexagon_event *)event->context; + + hex_event->sess = sess; + hex_event->seq = sess->record_event(); + HEX_VERBOSE("ggml-hex: %s event-record : event %p seq %llu\n", + sess->c_name(), (void *)hex_event, (unsigned long long)hex_event->seq); +} + +static void ggml_backend_hexagon_event_wait(ggml_backend_t backend, ggml_backend_event_t event) { + GGML_UNUSED(backend); + + ggml_hexagon_event * hex_event = (ggml_hexagon_event *)event->context; + if (hex_event->sess != nullptr) { + HEX_VERBOSE("ggml-hex: %s event-wait : event %p seq %llu\n", + hex_event->sess->c_name(), (void *)hex_event, (unsigned long long)hex_event->seq); + hex_event->sess->wait_event(hex_event->seq); + } +} + +static void ggml_backend_hexagon_set_tensor_async(ggml_backend_t backend, struct ggml_tensor * tensor, const void * data, size_t offset, size_t size) { + auto sess = static_cast(backend->context); + HEX_VERBOSE("ggml-hex: %s set-tensor-async %s : data %p offset %zu size %zu usage %d\n", + sess->c_name(), tensor->name, data, offset, size, tensor->buffer ? (int) tensor->buffer->usage : -1); + ggml_backend_tensor_set(tensor, data, offset, size); +} + +static void ggml_backend_hexagon_get_tensor_async(ggml_backend_t backend, const struct ggml_tensor * tensor, void * data, size_t offset, size_t size) { + auto sess = static_cast(backend->context); + HEX_VERBOSE("ggml-hex: %s get-tensor-async %s : data %p offset %zu size %zu usage %d\n", + sess->c_name(), tensor->name, data, offset, size, tensor->buffer ? (int) tensor->buffer->usage : -1); + sess->flush(true); + ggml_backend_tensor_get(tensor, data, offset, size); +} + +static void ggml_backend_hexagon_set_tensor_2d_async(ggml_backend_t backend, + struct ggml_tensor * tensor, + const void * data, + size_t offset, + size_t size, + size_t n_copies, + size_t stride_tensor, + size_t stride_data) { + auto sess = static_cast(backend->context); + HEX_VERBOSE("ggml-hex: %s set-tensor-2d-async %s : data %p offset %zu size %zu n_copies %zu stride_tensor %zu stride_data %zu usage %d\n", + sess->c_name(), tensor->name, data, offset, size, n_copies, stride_tensor, stride_data, tensor->buffer ? (int) tensor->buffer->usage : -1); + ggml_backend_tensor_set_2d(tensor, data, offset, size, n_copies, stride_tensor, stride_data); +} + +static void ggml_backend_hexagon_get_tensor_2d_async(ggml_backend_t backend, + const struct ggml_tensor * tensor, + void * data, + size_t offset, + size_t size, + size_t n_copies, + size_t stride_tensor, + size_t stride_data) { + auto sess = static_cast(backend->context); + HEX_VERBOSE("ggml-hex: %s get-tensor-2d-async %s : data %p offset %zu size %zu n_copies %zu stride_tensor %zu stride_data %zu usage %d\n", + sess->c_name(), tensor->name, data, offset, size, n_copies, stride_tensor, stride_data, tensor->buffer ? (int) tensor->buffer->usage : -1); + sess->flush(true); + ggml_backend_tensor_get_2d(tensor, data, offset, size, n_copies, stride_tensor, stride_data); +} + static struct ggml_backend_i hexagon_backend_i = { /* .get_name = */ ggml_backend_hexagon_name, /* .free = */ ggml_backend_hexagon_free, - /* .set_tensor_async = */ NULL, - /* .get_tensor_async = */ NULL, - /* .set_tensor_2d_async = */ NULL, - /* .get_tensor_2d_async = */ NULL, - /* .cpy_tensor_async = */ NULL, + /* .set_tensor_async = */ ggml_backend_hexagon_set_tensor_async, + /* .get_tensor_async = */ ggml_backend_hexagon_get_tensor_async, + /* .set_tensor_2d_async = */ ggml_backend_hexagon_set_tensor_2d_async, + /* .get_tensor_2d_async = */ ggml_backend_hexagon_get_tensor_2d_async, + /* .cpy_tensor_async = */ ggml_backend_hexagon_cpy_tensor_async, /* .synchronize = */ ggml_backend_hexagon_synchronize, /* .graph_plan_create = */ NULL, /* .graph_plan_free = */ NULL, /* .graph_plan_update = */ NULL, /* .graph_plan_compute = */ NULL, /* .graph_compute = */ ggml_backend_hexagon_graph_compute, - /* .event_record = */ NULL, - /* .event_wait = */ NULL, + /* .event_record = */ ggml_backend_hexagon_event_record, + /* .event_wait = */ ggml_backend_hexagon_event_wait, /* .graph_optimize = */ ggml_backend_hexagon_graph_optimize, }; @@ -3932,9 +5309,9 @@ static void ggml_backend_hexagon_device_get_props(ggml_backend_dev_t dev, struct ggml_backend_hexagon_device_get_memory(dev, &props->memory_free, &props->memory_total); props->caps = { /* .async = */ true, - /* .host_buffer = */ (bool) opt_hostbuf, + /* .host_buffer = */ false, /* .buffer_from_host_ptr = */ false, - /* .events = */ false, + /* .events = */ true, /* .mmap_support = */ false, }; } @@ -3944,32 +5321,12 @@ static ggml_backend_buffer_type_t ggml_backend_hexagon_device_get_buffer_type(gg return &sess->buffer_type; } -static ggml_backend_buffer_type_t ggml_backend_hexagon_device_get_repack_buffer_type(ggml_backend_dev_t dev) { - auto sess = static_cast(dev->context); - return &sess->repack_buffer_type; -} - -static bool ggml_hexagon_supported_buffer(ggml_hexagon_session *sess, const struct ggml_tensor * t) { - if (t && t->buffer) { - if (ggml_backend_buffer_is_hexagon(t->buffer) == false) return false; // not our buffer - if (ggml_backend_hexagon_buffer_get_sess(t->buffer) != sess) return false; // wrong session - } - return true; -} - -static bool ggml_hexagon_supported_buffers(ggml_hexagon_session *sess, const struct ggml_tensor * t) { - // all srcs & dsts must be mapped to the same session - if (!ggml_hexagon_supported_buffer(sess, t)) { - return false; - } - - for (int i = 0; i < GGML_MAX_SRC; i++) { - if (!ggml_hexagon_supported_buffer(sess, t->src[i])) { - return false; - } +static ggml_backend_buffer_type_t ggml_backend_hexagon_device_get_host_buffer_type(ggml_backend_dev_t dev) { + if (!opt_hostbuf) { + return NULL; } - - return true; + auto sess = static_cast(dev->context); + return &sess->host_buffer_type; } static bool ggml_hexagon_supported_cpy(const struct ggml_hexagon_session * sess, const struct ggml_tensor * op) { @@ -4067,12 +5424,6 @@ static bool ggml_backend_hexagon_device_supports_op(ggml_backend_dev_t dev, cons return false; } - // all srcs & dsts must be mapped to the same session - if (!ggml_hexagon_supported_buffers(sess, op)) { - ggml_hexagon_dump_op_supp(sess->name, op, false); - return false; - } - bool supp = false; switch (op->op) { case GGML_OP_NONE: @@ -4233,31 +5584,20 @@ static bool ggml_backend_hexagon_device_supports_op(ggml_backend_dev_t dev, cons } static bool ggml_backend_hexagon_device_supports_buft(ggml_backend_dev_t dev, ggml_backend_buffer_type_t buft) { - if (buft->iface.get_alignment != ggml_backend_hexagon_buffer_type_get_alignment) { - return false; - } - - auto s0 = static_cast(dev->context); - auto s1 = static_cast(buft->context)->sess; - - // Need session/domain-id for buffers to be compatible - bool supp = (s0->session_id == s1->session_id); + auto sess = static_cast(dev->context); - HEX_VERBOSE("ggml-hex: %s device-supports-buft %s (%d)\n", s0->name.c_str(), s1->name.c_str(), (int) supp); + // Technically we can clone hexagon buffers from any session but for some reason the output is garbled with layer-split, + // tensor-split works correctly, so it needs mode debugging and investigation. For now accept only our own buffers. +#if 0 + bool supp = (buft->iface.get_alignment == ggml_backend_hexagon_buffer_type_get_alignment); +#else + bool supp = (buft == &sess->host_buffer_type) || (buft == &sess->buffer_type); +#endif + HEX_VERBOSE("ggml-hex: %s device-supports-buft %s %s\n", sess->name.c_str(), ggml_backend_buft_name(buft), supp ? "yes" : "no"); return supp; } -static ggml_backend_buffer_type_t * ggml_backend_hexagon_device_get_extra_buffers_type(ggml_backend_dev_t dev) { - auto s0 = static_cast(dev->context); - HEX_VERBOSE("ggml-hex: device-get-extra-buft : %s \n", s0->name.c_str()); - - static ggml_backend_buffer_type_t bufts[2]; - bufts[0] = ggml_backend_hexagon_device_get_repack_buffer_type(dev); - bufts[1] = NULL; - return bufts; -} - static const struct ggml_backend_device_i ggml_backend_hexagon_device_i = { /* .get_name = */ ggml_backend_hexagon_device_get_name, /* .get_description = */ ggml_backend_hexagon_device_get_description, @@ -4266,27 +5606,18 @@ static const struct ggml_backend_device_i ggml_backend_hexagon_device_i = { /* .get_props = */ ggml_backend_hexagon_device_get_props, /* .init_backend = */ ggml_backend_hexagon_device_init, /* .get_buffer_type = */ ggml_backend_hexagon_device_get_buffer_type, - /* .get_host_buffer_type = */ NULL, // ggml_backend_hexagon_device_get_host_buffer_type, + /* .get_host_buffer_type = */ ggml_backend_hexagon_device_get_host_buffer_type, /* .buffer_from_host_ptr = */ NULL, // ggml_backend_hexagon_device_buffer_from_ptr, /* .supports_op = */ ggml_backend_hexagon_device_supports_op, /* .supports_buft = */ ggml_backend_hexagon_device_supports_buft, /* .offload_op = */ NULL, // ggml_backend_hexagon_device_offload_op, - /* .event_new = */ NULL, - /* .event_free = */ NULL, - /* .event_synchronize = */ NULL, + /* .event_new = */ ggml_backend_hexagon_device_event_new, + /* .event_free = */ ggml_backend_hexagon_device_event_free, + /* .event_synchronize = */ ggml_backend_hexagon_device_event_synchronize, }; //** backend registry -#define GGML_HEXAGON_MAX_SESSIONS 16 - -struct ggml_hexagon_registry { - ggml_hexagon_registry(ggml_backend_reg_t reg); - ~ggml_hexagon_registry(); - - ggml_backend_device devices[GGML_HEXAGON_MAX_SESSIONS]; -}; - ggml_hexagon_registry::ggml_hexagon_registry(ggml_backend_reg_t reg) { GGML_LOG_INFO("ggml-hex: Hexagon backend (experimental) : allocating new registry : ndev %zu\n", opt_ndev); @@ -4303,6 +5634,7 @@ ggml_hexagon_registry::ggml_hexagon_registry(ggml_backend_reg_t reg) { devices[i].context = nullptr; } } + } ggml_hexagon_registry::~ggml_hexagon_registry() { @@ -4335,14 +5667,129 @@ static ggml_backend_dev_t ggml_backend_hexagon_reg_get_device(ggml_backend_reg_t return &hreg->devices[index]; } -static void * ggml_backend_hexagon_get_proc_address(ggml_backend_reg_t reg, const char * name) { - if (strcmp(name, "ggml_backend_dev_get_extra_bufts") == 0 && opt_hostbuf) { - ggml_backend_dev_get_extra_bufts_t fct = ggml_backend_hexagon_device_get_extra_buffers_type; - return (void *) fct; +// ** communication context for tensor-split allreduce + +static void * ggml_backend_hexagon_comm_init(ggml_backend_t * backends, size_t n_backends) { + if (n_backends < 2 || n_backends > 4) { + return nullptr; } - return NULL; + for (size_t i = 0; i < n_backends; ++i) { + if (!ggml_backend_is_hexagon(backends[i])) { + return nullptr; + } + } + + auto * ctx = new ggml_backend_hexagon_comm_context(); + ctx->backends.assign(backends, backends + n_backends); + ctx->n_backends = n_backends; + ctx->fence_seq = (((uintptr_t) ctx) & 0xFFFF) | 1; + + return ctx; +} + +static void ggml_backend_hexagon_comm_free(void * comm_ctx_v) { + if (!comm_ctx_v) return; + delete static_cast(comm_ctx_v); +} + +static bool ggml_backend_hexagon_comm_allreduce_tensor(void * comm_ctx_v, struct ggml_tensor ** tensors) { + if (opt_ar_select == 0 || !comm_ctx_v) return false; + auto * comm_ctx = static_cast(comm_ctx_v); + const size_t n_backends = comm_ctx->n_backends; + + if (n_backends < 2 || n_backends > 4) return false; + + for (size_t i = 0; i < n_backends; i++) { + if (!tensors[i] || !tensors[i]->buffer || !ggml_backend_buffer_is_hexagon(tensors[i]->buffer)) { + return false; + } + if (tensors[i]->type != tensors[0]->type) { + return false; + } + if (!ggml_is_contiguous(tensors[i])) { + return false; + } + if (ggml_nelements(tensors[i]) != ggml_nelements(tensors[0])) { + return false; + } + } + + if (tensors[0]->type != GGML_TYPE_F16 && tensors[0]->type != GGML_TYPE_F32) { + return false; + } + + for (size_t r = 0; r < n_backends; r++) { + auto sess = static_cast(comm_ctx->backends[r]->context); + struct htp_allreduce_kernel_params kparams; + if (!ggml_hexagon_precompute_allreduce_params(sess, tensors[r], (uint32_t) r, (uint32_t) n_backends, false, false, &kparams)) { + return false; + } + } + + if (comm_ctx->fence_seq == 0) comm_ctx->fence_seq = 1; + uint32_t fence_seq_entry = comm_ctx->fence_seq++; + if (comm_ctx->fence_seq == 0) comm_ctx->fence_seq = 1; + uint32_t fence_seq_exit = comm_ctx->fence_seq++; + if (comm_ctx->fence_seq == 0) comm_ctx->fence_seq = 1; + + volatile uint32_t * fences[GGML_HEXAGON_MAX_SESSIONS]; + for (size_t i = 0; i < n_backends; i++) { + auto sbuf = (ggml_hexagon_shared_buffer *) tensors[i]->buffer->context; + fences[i] = (volatile uint32_t *) sbuf->alloc_fence(); + } + + static ggml_hexagon_tensor_extra fence_extra { {}, 0, GGML_HEXAGON_TENSOR_FENCE }; + ggml_tensor fence_tensors[GGML_HEXAGON_MAX_SESSIONS]; + for (size_t i = 0; i < n_backends; i++) { + fence_tensors[i] = {}; + fence_tensors[i].buffer = tensors[i]->buffer; + fence_tensors[i].extra = &fence_extra; + fence_tensors[i].data = (void *) fences[i]; + fence_tensors[i].type = GGML_TYPE_I32; + fence_tensors[i].ne[0] = 4; + fence_tensors[i].ne[1] = 1; + fence_tensors[i].ne[2] = 1; + fence_tensors[i].ne[3] = 1; + fence_tensors[i].nb[0] = sizeof(int32_t); + fence_tensors[i].nb[1] = sizeof(int32_t); + fence_tensors[i].nb[2] = sizeof(int32_t); + fence_tensors[i].nb[3] = sizeof(int32_t); + fence_tensors[i].op = GGML_OP_NONE; + } + + std::vector data_tensors(n_backends); + std::vector sync_tensors(n_backends); + for (size_t i = 0; i < n_backends; i++) { + data_tensors[i] = tensors[i]; + sync_tensors[i] = &fence_tensors[i]; + } + + for (size_t r = 0; r < n_backends; r++) { + auto sess = static_cast(comm_ctx->backends[r]->context); + sess->enqueue_allreduce(tensors[r], data_tensors, sync_tensors, (uint32_t) r, (uint32_t) n_backends, fence_seq_entry, fence_seq_exit); + for (size_t j = 0; j < n_backends; j++) { + if (r != j) { + sess->add_sync_peer(static_cast(comm_ctx->backends[j]->context)); + } + } + } + + return true; +} + +static void * ggml_backend_hexagon_get_proc_address(ggml_backend_reg_t reg, const char * name) { GGML_UNUSED(reg); + if (strcmp(name, "ggml_backend_comm_init") == 0) { + return (void *) ggml_backend_hexagon_comm_init; + } + if (strcmp(name, "ggml_backend_comm_free") == 0) { + return (void *) ggml_backend_hexagon_comm_free; + } + if (strcmp(name, "ggml_backend_comm_allreduce_tensor") == 0) { + return (void *) ggml_backend_hexagon_comm_allreduce_tensor; + } + return NULL; } template std::vector str_to_vec(const char* str) { @@ -4379,8 +5826,6 @@ static void ggml_hexagon_init(ggml_backend_reg * reg) { "please update hexagon_type to match ggml_type"); const char * str_verbose = getenv("GGML_HEXAGON_VERBOSE"); - const char * str_hostbuf = getenv("GGML_HEXAGON_HOSTBUF"); - const char * str_opstage = getenv("GGML_HEXAGON_OPSTAGE"); const char * str_opbatch = getenv("GGML_HEXAGON_OPBATCH"); const char * str_opqueue = getenv("GGML_HEXAGON_OPQUEUE"); const char * str_oppoll = getenv("GGML_HEXAGON_OPPOLL"); @@ -4389,15 +5834,16 @@ static void ggml_hexagon_init(ggml_backend_reg * reg) { const char * str_profile = getenv("GGML_HEXAGON_PROFILE"); const char * str_etm = getenv("GGML_HEXAGON_ETM"); const char * str_nhvx = getenv("GGML_HEXAGON_NHVX"); - const char * str_use_hmx = getenv("GGML_HEXAGON_USE_HMX"); const char * str_nhmx = getenv("GGML_HEXAGON_NHMX"); const char * str_mm_select = getenv("GGML_HEXAGON_MM_SELECT"); const char * str_fa_select = getenv("GGML_HEXAGON_FA_SELECT"); + const char * str_ar_select = getenv("GGML_HEXAGON_AR_SELECT"); const char * str_ndev = getenv("GGML_HEXAGON_NDEV"); const char * str_arch = getenv("GGML_HEXAGON_ARCH"); const char * str_vmem = getenv("GGML_HEXAGON_VMEM"); const char * str_mbuf = getenv("GGML_HEXAGON_MBUF"); const char * str_optrace = getenv("GGML_HEXAGON_OPTRACE"); + const char * str_hostbuf = getenv("GGML_HEXAGON_HOSTBUF"); // Init Arch first since it affects other defaults if (!str_arch) { @@ -4430,8 +5876,6 @@ static void ggml_hexagon_init(ggml_backend_reg * reg) { opt_opfilter = str_opfilter ? new std::regex(str_opfilter, RE_ICASE) : NULL; opt_verbose = str_verbose ? atoi(str_verbose) : 0; - opt_hostbuf = str_hostbuf ? atoi(str_hostbuf) : opt_hostbuf; - opt_opstage = str_opstage ? strtoul(str_opstage, NULL, 0) : opt_opstage; opt_opbatch = str_opbatch ? strtoul(str_opbatch, NULL, 0) : opt_opbatch; opt_opqueue = str_opqueue ? strtoul(str_opqueue, NULL, 0) : opt_opqueue; opt_optrace = str_optrace ? strtoul(str_optrace, NULL, 0) : (opt_opbatch * 256); @@ -4440,16 +5884,90 @@ static void ggml_hexagon_init(ggml_backend_reg * reg) { opt_profile = str_profile ? atoi(str_profile) : 0; opt_etm = str_etm ? atoi(str_etm) : 0; opt_nhvx = str_nhvx ? strtoul(str_nhvx, NULL, 0) : opt_nhvx; - opt_nhmx = str_nhmx ? atoi(str_nhmx) : (str_use_hmx ? atoi(str_use_hmx) : opt_nhmx); + opt_nhmx = str_nhmx ? atoi(str_nhmx) : opt_nhmx; opt_mm_select = str_mm_select ? atoi(str_mm_select) : opt_mm_select; opt_fa_select = str_fa_select ? atoi(str_fa_select) : opt_fa_select; - opt_ndev = str_ndev ? strtoul(str_ndev, NULL, 0) : opt_ndev; - opt_hostbuf = str_hostbuf ? atoi(str_hostbuf) : opt_hostbuf; + opt_ar_select = str_ar_select ? atoi(str_ar_select) : opt_ar_select; opt_mbuf = str_mbuf ? strtoul(str_mbuf, NULL, 0) * MiB : opt_mbuf; opt_vmem = str_vmem ? strtoul(str_vmem, NULL, 0) * MiB : opt_vmem; + opt_hostbuf = str_hostbuf ? atoi(str_hostbuf) != 0 : opt_hostbuf; + + // Parse device configuration + const char * str_devices = getenv("GGML_HEXAGON_DEVICES"); + if (!str_devices && str_ndev && str_ndev[0] != '\0') { + GGML_LOG_WARN("DEPRECATED: GGML_HEXAGON_NDEV is deprecated. use GGML_HEXAGON_DEVICES instead\n"); + str_devices = str_ndev; + } + + if (str_devices && str_devices[0] != '\0') { + bool is_single_number = true; + for (int i = 0; str_devices[i] != '\0'; i++) { + if (!isdigit((unsigned char)str_devices[i])) { + is_single_number = false; + break; + } + } + if (is_single_number) { + int n = atoi(str_devices); + if (n < 1) n = 1; + if (n > GGML_HEXAGON_MAX_SESSIONS) n = GGML_HEXAGON_MAX_SESSIONS; + opt_ndev = n; + for (size_t i = 0; i < opt_ndev; i++) { + opt_device_configs[i].physical_idx = 0; + opt_device_configs[i].virtual_idx = (int)i; + opt_device_configs[i].name = "HTP" + std::to_string(i); + } + } else { + std::string s_devices(str_devices); + std::stringstream ss(s_devices); + std::string item; + opt_ndev = 0; + while (std::getline(ss, item, ',')) { + size_t start = item.find_first_not_of(" \t\r\n"); + size_t end = item.find_last_not_of(" \t\r\n"); + if (start == std::string::npos) { + continue; + } + item = item.substr(start, end - start + 1); + + if (item.rfind("HTP", 0) == 0) { + std::string rest = item.substr(3); + size_t colon_pos = rest.find(':'); + int phys = 0; + int virt = 0; + try { + if (colon_pos == std::string::npos) { + phys = std::stoi(rest); + virt = 0; + } else { + phys = std::stoi(rest.substr(0, colon_pos)); + virt = std::stoi(rest.substr(colon_pos + 1)); + } + } catch (...) { + GGML_LOG_WARN("ggml-hex: failed to parse device index in '%s'\n", item.c_str()); + continue; + } - if (opt_ndev > GGML_HEXAGON_MAX_SESSIONS) { - opt_ndev = GGML_HEXAGON_MAX_SESSIONS; + if (opt_ndev < GGML_HEXAGON_MAX_SESSIONS) { + opt_device_configs[opt_ndev].physical_idx = phys; + opt_device_configs[opt_ndev].virtual_idx = virt; + opt_device_configs[opt_ndev].name = colon_pos == std::string::npos + ? "HTP" + std::to_string(phys) + : "HTP" + std::to_string(phys) + ":" + std::to_string(virt); + opt_ndev++; + } else { + GGML_LOG_WARN("ggml-hex: max sessions limit reached (%d), ignoring device %s\n", GGML_HEXAGON_MAX_SESSIONS, item.c_str()); + } + } else { + GGML_LOG_WARN("ggml-hex: invalid device name format '%s', must start with HTP\n", item.c_str()); + } + } + } + } else { + opt_ndev = 1; + opt_device_configs[0].physical_idx = 0; + opt_device_configs[0].virtual_idx = 0; + opt_device_configs[0].name = "HTP0"; } #if defined(__ANDROID__) diff --git a/ggml/src/ggml-hexagon/htp-opnode.h b/ggml/src/ggml-hexagon/htp-opnode.h index b0c859dacf9..741b5e04eb8 100644 --- a/ggml/src/ggml-hexagon/htp-opnode.h +++ b/ggml/src/ggml-hexagon/htp-opnode.h @@ -8,60 +8,107 @@ #include #include #include +#include #include #include "htp-ops.h" #include "htp/matmul-ops.h" #include "htp/flash-attn-ops.h" #include "htp/unary-ops.h" +#include "htp/allreduce-ops.h" struct htp_opnode { - ggml_tensor * node = nullptr; - - std::vector fused; - - htp_op_code opcode = HTP_OP_INVALID; + ggml_tensor * node { nullptr }; + htp_op_code opcode { HTP_OP_INVALID }; + int32_t kernel_params[HTP_OP_MAX_KERN_PARAMS] {0}; + + std::vector fused; + std::vector> dummy; + + std::vector inputs; + std::vector outputs; + std::string name; + + int n_active_src(const ggml_tensor * t) const { + if (!t) return 0; + for (int i = GGML_MAX_SRC - 1; i >= 0; i--) { + if (t->src[i]) { + return i + 1; + } + } + return 0; + } - std::vector extra_dsts; + void init(ggml_tensor * node) { + this->node = node; + if (this->node) { + this->name = ggml_op_desc(this->node); - int32_t kernel_params[HTP_OP_MAX_KERN_PARAMS] = {0}; + // Build inputs (preserving optional nullptrs) + int n_inputs = n_active_src(this->node); + this->inputs.resize(n_inputs, nullptr); + for (int i = 0; i < n_inputs; i++) { + this->inputs[i] = this->node->src[i]; + } - htp_opnode(ggml_tensor * node = nullptr, std::vector fused = {}, htp_op_code opcode = HTP_OP_INVALID, std::vector extra_dsts = {}) - : node(node), fused(std::move(fused)), opcode(opcode), extra_dsts(std::move(extra_dsts)) {} + // Build outputs + this->outputs.push_back(this->dst()); + } + } - ggml_op op() const { - return node->op; + htp_opnode(htp_op_code opcode = HTP_OP_INVALID, ggml_tensor * node = nullptr) : opcode(opcode) { + init(node); } - const ggml_tensor * dst() const { - return fused.empty() ? node : fused.back(); + ggml_op op() const { return node->op; } + const ggml_tensor * src0() const { return node->src[0]; } + const ggml_tensor * src1() const { return node->src[1]; } + const ggml_tensor * dst() const { return outputs.empty() ? node : outputs.back(); } + + ggml_tensor * add_dummy(const ggml_tensor & t) { + dummy.push_back(std::make_shared(t)); + return dummy.back().get(); } void add_fused(ggml_tensor * t, bool extra_dst = false) { fused.push_back(t); + + name += "+"; + name += ggml_op_desc(t); + if (extra_dst) { - extra_dsts.push_back(t); + outputs.push_back(t); + } else { + outputs.clear(); + outputs.push_back(t); } - } - std::vector get_outputs() const { - std::vector res; - if (extra_dsts.empty()) { - res.push_back(dst()); - } else { - res.push_back(node); - for (const auto * x : extra_dsts) { - res.push_back(x); + // Remove the newly fused intermediate output tensor t from inputs (if it was there) + inputs.erase(std::remove(inputs.begin(), inputs.end(), t), inputs.end()); + + // Append new inputs from t, preserving middle nullptrs + int n_inputs = n_active_src(t); + for (int i = 0; i < n_inputs; i++) { + const auto * src = t->src[i]; + if (!src) { + inputs.push_back(nullptr); + } else if (src != node && + std::find(fused.begin(), fused.end(), src) == fused.end() && + std::find(inputs.begin(), inputs.end(), src) == inputs.end()) { + inputs.push_back(src); } } - return res; } - const ggml_tensor * src0() const { - return node->src[0]; + const std::vector & get_inputs() const { + return inputs; } - const ggml_tensor * src1() const { - return node->src[1]; + const std::vector & get_outputs() const { + return outputs; + } + + std::string op_name() const { + return name; } bool is_empty() const { @@ -81,75 +128,6 @@ struct htp_opnode { bool same_input(const htp_opnode& n) const { return n.src1() == this->src1(); } - - std::vector get_inputs() const { - if (fused.empty()) { - int last_non_null = -1; - for (int i = 0; i < GGML_MAX_SRC; i++) { - if (node->src[i]) { - last_non_null = i; - } - } - std::vector inputs(last_non_null + 1, nullptr); - for (int i = 0; i <= last_non_null; i++) { - inputs[i] = node->src[i]; - } - return inputs; - } - - std::vector inputs(GGML_MAX_SRC, nullptr); - std::vector outputs; - outputs.push_back(node); - for (const auto * f : fused) { - outputs.push_back(f); - } - - auto contains = [&](const std::vector & vec, const ggml_tensor * t) { - for (const auto * x : vec) { - if (x == t) return true; - } - return false; - }; - - int count = 0; - auto add_input = [&](const ggml_tensor * t) { - if (t && !contains(outputs, t) && !contains(inputs, t)) { - if (count < (int)inputs.size()) { - inputs[count++] = t; - } else { - inputs.push_back(t); - } - } - }; - - for (int i = 0; i < GGML_MAX_SRC; i++) { - if (node->src[i]) { - add_input(node->src[i]); - } - } - for (const auto * f : fused) { - for (int i = 0; i < GGML_MAX_SRC; i++) { - if (f->src[i]) { - add_input(f->src[i]); - } - } - } - - inputs.resize(count); - return inputs; - } - - std::string op_name() const { - if (fused.empty()) { - return ggml_op_desc(node); - } - std::string name = ggml_op_desc(node); - for (const auto * f : fused) { - name += "+"; - name += ggml_op_desc(f); - } - return name; - } }; struct htp_opformat { @@ -337,8 +315,7 @@ struct htp_opformat { } void format_kernel_params(char * str, size_t max_size, const htp_opnode & node) { if (node.opcode == HTP_OP_MUL_MAT || node.opcode == HTP_OP_MUL_MAT_ID || - node.opcode == HTP_OP_MUL_MAT_QKV || node.opcode == HTP_OP_MUL_MAT_FFN || - node.opcode == HTP_OP_MUL_MAT_ADD) { + node.opcode == HTP_OP_MUL_MAT_NX || node.opcode == HTP_OP_MUL_MAT_ADD) { const auto * kparams = (const struct htp_mm_kernel_params *) node.kernel_params; const char * path = "unknown"; int32_t type = kparams->kernel_type; diff --git a/ggml/src/ggml-hexagon/htp/CMakeLists.txt b/ggml/src/ggml-hexagon/htp/CMakeLists.txt index b00aa2bc94c..77f3ee39dd3 100644 --- a/ggml/src/ggml-hexagon/htp/CMakeLists.txt +++ b/ggml/src/ggml-hexagon/htp/CMakeLists.txt @@ -43,6 +43,7 @@ add_library(${HTP_LIB} SHARED pad-ops.c argsort-ops.c im2col-ops.c + allreduce-ops.c ) target_compile_definitions(${HTP_LIB} PRIVATE diff --git a/ggml/src/ggml-hexagon/htp/act-ops.c b/ggml/src/ggml-hexagon/htp/act-ops.c index 9973c088dda..0a8bf84e382 100644 --- a/ggml/src/ggml-hexagon/htp/act-ops.c +++ b/ggml/src/ggml-hexagon/htp/act-ops.c @@ -183,6 +183,53 @@ static void swiglu_oai_f32(const float * restrict src0, static const float GELU_COEF_A = 0.044715f; static const float SQRT_2_OVER_PI = 0.79788456080286535587989211986876f; +static inline HVX_Vector hvx_vec_fast_sigmoid_f32_2it(HVX_Vector v) { + v = Q6_Vqf32_vmpy_VsfVsf(v, Q6_V_vsplat_R(FAST_SIGMOID_LOG2F)); + v = Q6_Vqf32_vmpy_VsfVsf(Q6_Vsf_equals_Vqf32(v), Q6_V_vsplat_R(FAST_SIGMOID_C3)); + + HVX_Vector in_int = hvx_vec_truncate_f32(Q6_Vsf_equals_Vqf32(v)); + HVX_Vector x = Q6_Vqf32_vsub_Vqf32Vsf(v, Q6_Vsf_equals_Vw(in_int)); + HVX_Vector xx = Q6_Vqf32_vmpy_Vqf32Vqf32(x, x); + + HVX_Vector v1 = Q6_Vqf32_vmpy_VsfVsf(Q6_Vsf_equals_Vqf32(xx), Q6_V_vsplat_R(FAST_SIGMOID_C2)); + v1 = Q6_Vqf32_vadd_Vqf32Vsf(v1, Q6_V_vsplat_R(FAST_SIGMOID_LOG2F)); + + HVX_Vector v2 = Q6_Vqf32_vmpy_VsfVsf(Q6_Vsf_equals_Vqf32(x), Q6_V_vsplat_R(FAST_SIGMOID_C1)); + v2 = Q6_Vqf32_vmpy_Vqf32Vqf32(v2, xx); + v2 = Q6_Vqf32_vadd_Vqf32Vqf32(v2, x); + + HVX_Vector v3 = Q6_Vsf_equals_Vqf32(Q6_Vqf32_vadd_Vqf32Vqf32(v2, v1)); + v3 = Q6_Vw_vaslacc_VwVwR(v3, in_int, 24); + + HVX_Vector v4 = Q6_Vsf_equals_Vqf32(Q6_Vqf32_vsub_Vqf32Vqf32(v2, v1)); + HVX_Vector v5 = Q6_Vsf_equals_Vqf32(Q6_Vqf32_vsub_VsfVsf(v3, v4)); + + // Newton-Raphson with 2 iterations + HVX_Vector two_sf = hvx_vec_splat_f32(2.0f); + HVX_Vector i_sf = Q6_Vw_vsub_VwVw(Q6_V_vsplat_R(0x7EEEEBB3), v5); + HVX_Vector r_qf = Q6_Vqf32_vmpy_VsfVsf( + i_sf, Q6_Vsf_equals_Vqf32(Q6_Vqf32_vsub_VsfVsf(two_sf, Q6_Vsf_equals_Vqf32(Q6_Vqf32_vmpy_VsfVsf(i_sf, v5))))); + r_qf = Q6_Vqf32_vmpy_Vqf32Vqf32( + r_qf, Q6_Vqf32_vsub_VsfVsf(two_sf, Q6_Vsf_equals_Vqf32(Q6_Vqf32_vmpy_VsfVsf(Q6_Vsf_equals_Vqf32(r_qf), v5)))); + HVX_Vector res = Q6_Vsf_equals_Vqf32(r_qf); + + res = Q6_Vqf32_vmpy_VsfVsf(v3, res); + + return Q6_Vsf_equals_Vqf32(res); +} + +static inline HVX_Vector hvx_vec_fast_sigmoid_f32_guard_2it(HVX_Vector v, + HVX_Vector one, + HVX_Vector max_exp, + HVX_Vector min_exp) { + const HVX_VectorPred pred_max = Q6_Q_vcmp_gt_VsfVsf(max_exp, v); + const HVX_VectorPred pred_min = Q6_Q_vcmp_gt_VsfVsf(v, min_exp); + + HVX_Vector out = hvx_vec_fast_sigmoid_f32_2it(v); + out = Q6_V_vmux_QVV(pred_max, out, one); + return Q6_V_vmux_QVV(pred_min, out, Q6_V_vzero()); +} + static inline void hvx_geglu_f32_aa(uint8_t * restrict dst, const uint8_t * restrict src0, const uint8_t * restrict src1, uint32_t n) { assert((unsigned long) dst % 128 == 0); assert((unsigned long) src0 % 128 == 0); @@ -200,20 +247,13 @@ static inline void hvx_geglu_f32_aa(uint8_t * restrict dst, const uint8_t * rest const HVX_Vector v_coef_a_times_sqrt = hvx_vec_splat_f32(GELU_COEF_A_TIMES_SQRT); const HVX_Vector v_sqrt_2_pi = hvx_vec_splat_f32(SQRT_2_OVER_PI); - const HVX_Vector v_half = hvx_vec_splat_f32(0.5f); const HVX_Vector v_one = hvx_vec_splat_f32(1.0f); - const HVX_Vector v_two = hvx_vec_splat_f32(2.0f); - - // Hoisted fast sigmoid / inverse constants to avoid loop-internal overhead - const HVX_Vector v_log2f = Q6_V_vsplat_R(FAST_SIGMOID_LOG2F); - const HVX_Vector v_c1 = Q6_V_vsplat_R(FAST_SIGMOID_C1); - const HVX_Vector v_c2 = Q6_V_vsplat_R(FAST_SIGMOID_C2); - const HVX_Vector v_inv_aprox = Q6_V_vsplat_R(0x7EEEEBB3); const HVX_Vector v_max_exp = hvx_vec_splat_f32(87.0f); const HVX_Vector v_min_exp = hvx_vec_splat_f32(-87.0f); uint32_t i = 0; + _Pragma("unroll(4)") for (; i < nvec; i++) { HVX_Vector x = vsrc0[i]; HVX_Vector g = vsrc1[i]; @@ -223,56 +263,13 @@ static inline void hvx_geglu_f32_aa(uint8_t * restrict dst, const uint8_t * rest coef = hvx_vec_add_f32_f32(coef, v_sqrt_2_pi); HVX_Vector inner = hvx_vec_mul_f32_f32(x, coef); - // y2 = 2 * inner - HVX_Vector y2 = hvx_vec_mul_f32_f32(inner, v_two); - - // Sigmoid guard check predicates - HVX_VectorPred pred_max = Q6_Q_vcmp_gt_VsfVsf(v_max_exp, y2); - HVX_VectorPred pred_min = Q6_Q_vcmp_gt_VsfVsf(y2, v_min_exp); - - // Fast sigmoid approximation - HVX_Vector v = Q6_Vqf32_vmpy_VsfVsf(y2, v_log2f); - v = Q6_Vqf32_vmpy_VsfVsf(Q6_Vsf_equals_Vqf32(v), v_half); - - HVX_Vector in_int = hvx_vec_truncate_f32(Q6_Vsf_equals_Vqf32(v)); - HVX_Vector x_sig = Q6_Vqf32_vsub_Vqf32Vsf(v, Q6_Vsf_equals_Vw(in_int)); - HVX_Vector xx_sig = Q6_Vqf32_vmpy_Vqf32Vqf32(x_sig, x_sig); - - HVX_Vector v1 = Q6_Vqf32_vmpy_VsfVsf(Q6_Vsf_equals_Vqf32(xx_sig), v_c2); - v1 = Q6_Vqf32_vadd_Vqf32Vsf(v1, v_log2f); + // y2 = 2 * inner = inner + inner + HVX_Vector y2 = hvx_vec_add_f32_f32(inner, inner); - HVX_Vector v2 = Q6_Vqf32_vmpy_VsfVsf(Q6_Vsf_equals_Vqf32(x_sig), v_c1); - v2 = Q6_Vqf32_vmpy_Vqf32Vqf32(v2, xx_sig); - v2 = Q6_Vqf32_vadd_Vqf32Vqf32(v2, x_sig); - - HVX_Vector v3 = Q6_Vsf_equals_Vqf32(Q6_Vqf32_vadd_Vqf32Vqf32(v2, v1)); - v3 = Q6_Vw_vaslacc_VwVwR(v3, in_int, 24); - - HVX_Vector v4 = Q6_Vsf_equals_Vqf32(Q6_Vqf32_vsub_Vqf32Vqf32(v2, v1)); - HVX_Vector v5 = Q6_Vsf_equals_Vqf32(Q6_Vqf32_vsub_VsfVsf(v3, v4)); - - // Fast division (Newton-Raphson with 2 iterations) - HVX_Vector i_sf = Q6_Vw_vsub_VwVw(v_inv_aprox, v5); - HVX_Vector r_qf = Q6_Vqf32_vmpy_VsfVsf( - i_sf, Q6_Vsf_equals_Vqf32(Q6_Vqf32_vsub_VsfVsf(v_two, Q6_Vsf_equals_Vqf32(Q6_Vqf32_vmpy_VsfVsf(i_sf, v5))))); - r_qf = Q6_Vqf32_vmpy_Vqf32Vqf32( - r_qf, Q6_Vqf32_vsub_VsfVsf(v_two, Q6_Vsf_equals_Vqf32(Q6_Vqf32_vmpy_VsfVsf(Q6_Vsf_equals_Vqf32(r_qf), v5)))); - HVX_Vector res_inv = Q6_Vsf_equals_Vqf32(r_qf); - - HVX_Vector sig2y = Q6_Vsf_equals_Vqf32(Q6_Vqf32_vmpy_VsfVsf(v3, res_inv)); - - // Sigmoid guards - sig2y = Q6_V_vmux_QVV(pred_max, sig2y, v_one); - sig2y = Q6_V_vmux_QVV(pred_min, sig2y, Q6_V_vzero()); - - // tanh(inner) = 2 * sigmoid(2 * inner) - 1 - HVX_Vector tanh_val = hvx_vec_mul_f32_f32(sig2y, v_two); - tanh_val = hvx_vec_sub_f32_f32(tanh_val, v_one); - - HVX_Vector tanh_plus_one = hvx_vec_add_f32_f32(tanh_val, v_one); - HVX_Vector half_x = hvx_vec_mul_f32_f32(x, v_half); - HVX_Vector gelu_x = hvx_vec_mul_f32_f32(half_x, tanh_plus_one); + // Fast sigmoid approximation (2 iterations) + HVX_Vector sig2y = hvx_vec_fast_sigmoid_f32_guard_2it(y2, v_one, v_max_exp, v_min_exp); + HVX_Vector gelu_x = hvx_vec_mul_f32_f32(x, sig2y); vdst[i] = hvx_vec_mul_f32_f32(gelu_x, g); } @@ -285,50 +282,11 @@ static inline void hvx_geglu_f32_aa(uint8_t * restrict dst, const uint8_t * rest coef = hvx_vec_add_f32_f32(coef, v_sqrt_2_pi); HVX_Vector inner = hvx_vec_mul_f32_f32(x, coef); - HVX_Vector y2 = hvx_vec_mul_f32_f32(inner, v_two); - - HVX_VectorPred pred_max = Q6_Q_vcmp_gt_VsfVsf(v_max_exp, y2); - HVX_VectorPred pred_min = Q6_Q_vcmp_gt_VsfVsf(y2, v_min_exp); - - HVX_Vector v = Q6_Vqf32_vmpy_VsfVsf(y2, v_log2f); - v = Q6_Vqf32_vmpy_VsfVsf(Q6_Vsf_equals_Vqf32(v), v_half); - - HVX_Vector in_int = hvx_vec_truncate_f32(Q6_Vsf_equals_Vqf32(v)); - HVX_Vector x_sig = Q6_Vqf32_vsub_Vqf32Vsf(v, Q6_Vsf_equals_Vw(in_int)); - HVX_Vector xx_sig = Q6_Vqf32_vmpy_Vqf32Vqf32(x_sig, x_sig); - - HVX_Vector v1 = Q6_Vqf32_vmpy_VsfVsf(Q6_Vsf_equals_Vqf32(xx_sig), v_c2); - v1 = Q6_Vqf32_vadd_Vqf32Vsf(v1, v_log2f); - - HVX_Vector v2 = Q6_Vqf32_vmpy_VsfVsf(Q6_Vsf_equals_Vqf32(x_sig), v_c1); - v2 = Q6_Vqf32_vmpy_Vqf32Vqf32(v2, xx_sig); - v2 = Q6_Vqf32_vadd_Vqf32Vqf32(v2, x_sig); - - HVX_Vector v3 = Q6_Vsf_equals_Vqf32(Q6_Vqf32_vadd_Vqf32Vqf32(v2, v1)); - v3 = Q6_Vw_vaslacc_VwVwR(v3, in_int, 24); - - HVX_Vector v4 = Q6_Vsf_equals_Vqf32(Q6_Vqf32_vsub_Vqf32Vqf32(v2, v1)); - HVX_Vector v5 = Q6_Vsf_equals_Vqf32(Q6_Vqf32_vsub_VsfVsf(v3, v4)); - - HVX_Vector i_sf = Q6_Vw_vsub_VwVw(v_inv_aprox, v5); - HVX_Vector r_qf = Q6_Vqf32_vmpy_VsfVsf( - i_sf, Q6_Vsf_equals_Vqf32(Q6_Vqf32_vsub_VsfVsf(v_two, Q6_Vsf_equals_Vqf32(Q6_Vqf32_vmpy_VsfVsf(i_sf, v5))))); - r_qf = Q6_Vqf32_vmpy_Vqf32Vqf32( - r_qf, Q6_Vqf32_vsub_VsfVsf(v_two, Q6_Vsf_equals_Vqf32(Q6_Vqf32_vmpy_VsfVsf(Q6_Vsf_equals_Vqf32(r_qf), v5)))); - HVX_Vector res_inv = Q6_Vsf_equals_Vqf32(r_qf); - - HVX_Vector sig2y = Q6_Vsf_equals_Vqf32(Q6_Vqf32_vmpy_VsfVsf(v3, res_inv)); - - sig2y = Q6_V_vmux_QVV(pred_max, sig2y, v_one); - sig2y = Q6_V_vmux_QVV(pred_min, sig2y, Q6_V_vzero()); - - HVX_Vector tanh_val = hvx_vec_mul_f32_f32(sig2y, v_two); - tanh_val = hvx_vec_sub_f32_f32(tanh_val, v_one); + HVX_Vector y2 = hvx_vec_add_f32_f32(inner, inner); - HVX_Vector tanh_plus_one = hvx_vec_add_f32_f32(tanh_val, v_one); - HVX_Vector half_x = hvx_vec_mul_f32_f32(x, v_half); - HVX_Vector gelu_x = hvx_vec_mul_f32_f32(half_x, tanh_plus_one); + HVX_Vector sig2y = hvx_vec_fast_sigmoid_f32_guard_2it(y2, v_one, v_max_exp, v_min_exp); + HVX_Vector gelu_x = hvx_vec_mul_f32_f32(x, sig2y); HVX_Vector res = hvx_vec_mul_f32_f32(gelu_x, g); hvx_vec_store_a((void *) &vdst[i], nloe * sizeof(float), res); } diff --git a/ggml/src/ggml-hexagon/htp/allreduce-ops.c b/ggml/src/ggml-hexagon/htp/allreduce-ops.c new file mode 100644 index 00000000000..d35f685a6dc --- /dev/null +++ b/ggml/src/ggml-hexagon/htp/allreduce-ops.c @@ -0,0 +1,398 @@ +#pragma clang diagnostic ignored "-Wunused-variable" +#pragma clang diagnostic ignored "-Wunused-function" +#pragma clang diagnostic ignored "-Wunused-but-set-variable" + +#include +#include +#include +#include +#include + +#define GGML_COMMON_DECL_C +#include "ggml-common.h" +#include "htp-ctx.h" +#include "htp-ops.h" +#include "hvx-utils.h" +#include "htp-tensor.h" +#include "hex-dma.h" +#include "hex-profile.h" +#include "allreduce-ops.h" + +struct htp_allreduce_context { + struct htp_ops_context * octx; + uint32_t n_ranks; + uint32_t n_dsts; + uint32_t nelem; + uint32_t ne0; + uint32_t ne1; + uint32_t row_size_aligned; + uint32_t rank_elem_start; + uint32_t rank_nelem; + uint32_t elems_per_thread; + uint32_t block_elems; + uint32_t vtcm_size_per_thread; + bool is_row_bcast; + uint8_t * src_spad_base[HTP_ALLREDUCE_MAX_RANKS]; + uint8_t * dst_spad_base; + uint8_t * res_spad_base; +}; + +#define DEFINE_ALLREDUCE_THREAD_DMA_1D(SUFFIX, TYPE, HVX_ADD_FN, HAS_ADD) \ +static void allreduce_thread_dma_1d_##SUFFIX(unsigned int nth, unsigned int ith, void * data) { \ + struct htp_allreduce_context * actx = (struct htp_allreduce_context *) data; \ + struct htp_ops_context * octx = actx->octx; \ + \ + const uint32_t n_ranks = actx->n_ranks; \ + const uint32_t n_dsts = actx->n_dsts; \ + const uint32_t block_elems = actx->block_elems; \ + \ + const uint32_t dr = actx->elems_per_thread; \ + const uint32_t ir0 = actx->rank_elem_start + dr * ith; \ + const uint32_t ir1 = MIN(ir0 + dr, actx->rank_elem_start + actx->rank_nelem); \ + if (ir0 >= ir1) return; \ + \ + struct htp_thread_trace * tr = &octx->ctx->trace[ith]; \ + dma_queue * q = octx->ctx->dma[ith]; \ + \ + uint8_t * src_spad_base[HTP_ALLREDUCE_MAX_RANKS]; \ + for (uint32_t s = 0; s < n_ranks; s++) { \ + src_spad_base[s] = actx->src_spad_base[s] + (ith * actx->vtcm_size_per_thread); \ + } \ + uint8_t * dst_spad_base = actx->dst_spad_base + (ith * actx->vtcm_size_per_thread); \ + uint8_t * res_spad_base = HAS_ADD ? (actx->res_spad_base + (ith * actx->vtcm_size_per_thread)) : NULL; \ + \ + const size_t spad_half = actx->vtcm_size_per_thread / 2; \ + uint32_t ir_prefetch = ir0; \ + int spad_idx = 0; \ + \ + for (int k = 0; k < 2 && ir_prefetch < ir1; k++) { \ + uint32_t cur_elems = MIN(block_elems, ir1 - ir_prefetch); \ + size_t cur_bytes = cur_elems * sizeof(TYPE); \ + uint8_t * d_spad = dst_spad_base + spad_idx * spad_half; \ + for (uint32_t d = 0; d < n_dsts; d++) { \ + uint8_t * d_ddr = (uint8_t *) octx->dsts[d]->data + ir_prefetch * sizeof(TYPE); \ + dma_queue_push(q, dma_make_ptr(d_ddr, d_spad), cur_bytes, cur_bytes, cur_bytes, 0); \ + } \ + for (uint32_t s = 0; s < n_ranks; s++) { \ + uint8_t * s_spad = src_spad_base[s] + spad_idx * spad_half; \ + const uint8_t * s_ddr = (const uint8_t *) octx->src[s]->data + ir_prefetch * sizeof(TYPE); \ + dma_queue_push(q, dma_make_ptr(s_spad, s_ddr), cur_bytes, cur_bytes, cur_bytes, 1); \ + } \ + if (HAS_ADD) { \ + uint8_t * r_spad = res_spad_base + spad_idx * spad_half; \ + const uint8_t * r_ddr = (const uint8_t *) octx->src[2 * n_ranks]->data + ir_prefetch * sizeof(TYPE); \ + dma_queue_push(q, dma_make_ptr(r_spad, r_ddr), cur_bytes, cur_bytes, cur_bytes, 1); \ + } \ + ir_prefetch += cur_elems; \ + spad_idx ^= 1; \ + } \ + \ + for (uint32_t ir = ir0; ir < ir1; ) { \ + uint32_t cur_elems = MIN(block_elems, ir1 - ir); \ + size_t cur_bytes = cur_elems * sizeof(TYPE); \ + uint8_t * d_spad = NULL; \ + for (uint32_t d = 0; d < n_dsts; d++) { \ + d_spad = (uint8_t *) dma_queue_pop(q).src; \ + } \ + uint8_t * s_spad[HTP_ALLREDUCE_MAX_RANKS]; \ + for (uint32_t s = 0; s < n_ranks; s++) { \ + s_spad[s] = (uint8_t *) dma_queue_pop(q).dst; \ + } \ + uint8_t * r_spad = HAS_ADD ? (uint8_t *) dma_queue_pop(q).dst : NULL; \ + htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) ir); \ + HVX_ADD_FN(d_spad, s_spad[0], s_spad[1], cur_elems); \ + for (uint32_t s = 2; s < n_ranks; s++) { \ + HVX_ADD_FN(d_spad, d_spad, s_spad[s], cur_elems); \ + } \ + if (HAS_ADD) { \ + HVX_ADD_FN(d_spad, d_spad, r_spad, cur_elems); \ + } \ + htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) ir); \ + for (uint32_t d = 0; d < n_dsts; d++) { \ + uint8_t * d_ddr = (uint8_t *) octx->dsts[d]->data + ir * sizeof(TYPE); \ + dma_queue_push(q, dma_make_ptr(d_ddr, d_spad), cur_bytes, cur_bytes, cur_bytes, 1); \ + } \ + if (ir_prefetch < ir1) { \ + uint32_t next_elems = MIN(block_elems, ir1 - ir_prefetch); \ + size_t next_bytes = next_elems * sizeof(TYPE); \ + for (uint32_t s = 0; s < n_ranks; s++) { \ + const uint8_t * s_next = (const uint8_t *) octx->src[s]->data + ir_prefetch * sizeof(TYPE); \ + dma_queue_push(q, dma_make_ptr(s_spad[s], s_next), next_bytes, next_bytes, next_bytes, 1); \ + } \ + if (HAS_ADD) { \ + const uint8_t * r_next = (const uint8_t *) octx->src[2 * n_ranks]->data + ir_prefetch * sizeof(TYPE); \ + dma_queue_push(q, dma_make_ptr(r_spad, r_next), next_bytes, next_bytes, next_bytes, 1); \ + } \ + ir_prefetch += next_elems; \ + } \ + ir += cur_elems; \ + } \ + dma_queue_flush(q); \ +} + +DEFINE_ALLREDUCE_THREAD_DMA_1D(f16, __fp16, hvx_add_f16_aaa, 0) +DEFINE_ALLREDUCE_THREAD_DMA_1D(f32, float, hvx_add_f32_aaa, 0) +DEFINE_ALLREDUCE_THREAD_DMA_1D(add_f16, __fp16, hvx_add_f16_aaa, 1) +DEFINE_ALLREDUCE_THREAD_DMA_1D(add_f32, float, hvx_add_f32_aaa, 1) + +#define DEFINE_ALLREDUCE_THREAD_DMA_2D(SUFFIX, TYPE, HVX_ADD_FN, HAS_ADD, IS_ROW_BCAST) \ +static void allreduce_thread_dma_2d_##SUFFIX(unsigned int nth, unsigned int ith, void * data) { \ + struct htp_allreduce_context * actx = (struct htp_allreduce_context *) data; \ + struct htp_ops_context * octx = actx->octx; \ + \ + const uint32_t n_ranks = actx->n_ranks; \ + const uint32_t n_dsts = actx->n_dsts; \ + const uint32_t ne0 = actx->ne0; \ + const uint32_t block_rows = actx->block_elems; \ + const uint32_t row_size_aligned = actx->row_size_aligned; \ + const uint32_t row_bytes = ne0 * sizeof(TYPE); \ + \ + const uint32_t dr = actx->elems_per_thread; \ + const uint32_t r0 = actx->rank_elem_start + dr * ith; \ + const uint32_t r1 = MIN(r0 + dr, actx->rank_elem_start + actx->rank_nelem); \ + if (r0 >= r1) return; \ + \ + struct htp_thread_trace * tr = &octx->ctx->trace[ith]; \ + dma_queue * q = octx->ctx->dma[ith]; \ + \ + uint8_t * src_spad_base[HTP_ALLREDUCE_MAX_RANKS]; \ + for (uint32_t s = 0; s < n_ranks; s++) { \ + src_spad_base[s] = actx->src_spad_base[s] + (ith * actx->vtcm_size_per_thread); \ + } \ + uint8_t * dst_spad_base = actx->dst_spad_base + (ith * actx->vtcm_size_per_thread); \ + uint8_t * res_spad_base = HAS_ADD ? (IS_ROW_BCAST ? actx->res_spad_base : (actx->res_spad_base + (ith * actx->vtcm_size_per_thread))) : NULL; \ + \ + const size_t spad_half = actx->vtcm_size_per_thread / 2; \ + uint32_t r_prefetch = r0; \ + int spad_idx = 0; \ + \ + for (int k = 0; k < 2 && r_prefetch < r1; k++) { \ + uint32_t cur_rows = MIN(block_rows, r1 - r_prefetch); \ + uint8_t * d_spad = dst_spad_base + spad_idx * spad_half; \ + for (uint32_t d = 0; d < n_dsts; d++) { \ + uint8_t * d_ddr = (uint8_t *) octx->dsts[d]->data + r_prefetch * octx->dsts[d]->nb[1]; \ + dma_queue_push(q, dma_make_ptr(d_ddr, d_spad), octx->dsts[d]->nb[1], row_size_aligned, row_bytes, 0); \ + } \ + for (uint32_t s = 0; s < n_ranks; s++) { \ + uint8_t * s_spad = src_spad_base[s] + spad_idx * spad_half; \ + const uint8_t * s_ddr = (const uint8_t *) octx->src[s]->data + r_prefetch * octx->src[s]->nb[1]; \ + dma_queue_push(q, dma_make_ptr(s_spad, s_ddr), row_size_aligned, octx->src[s]->nb[1], row_bytes, cur_rows); \ + } \ + if (HAS_ADD && !IS_ROW_BCAST) { \ + uint8_t * r_spad = res_spad_base + spad_idx * spad_half; \ + const uint8_t * r_ddr = (const uint8_t *) octx->src[2 * n_ranks]->data + r_prefetch * octx->src[2 * n_ranks]->nb[1]; \ + dma_queue_push(q, dma_make_ptr(r_spad, r_ddr), row_size_aligned, octx->src[2 * n_ranks]->nb[1], row_bytes, cur_rows); \ + } \ + r_prefetch += cur_rows; \ + spad_idx ^= 1; \ + } \ + \ + for (uint32_t r = r0; r < r1; ) { \ + uint32_t cur_rows = MIN(block_rows, r1 - r); \ + uint8_t * d_spad = NULL; \ + for (uint32_t d = 0; d < n_dsts; d++) { \ + d_spad = (uint8_t *) dma_queue_pop(q).src; \ + } \ + uint8_t * s_spad[HTP_ALLREDUCE_MAX_RANKS]; \ + for (uint32_t s = 0; s < n_ranks; s++) { \ + s_spad[s] = (uint8_t *) dma_queue_pop(q).dst; \ + } \ + uint8_t * r_spad = (HAS_ADD && !IS_ROW_BCAST) ? (uint8_t *) dma_queue_pop(q).dst : NULL; \ + htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) r); \ + for (uint32_t row = 0; row < cur_rows; row++) { \ + uint8_t * d_row = d_spad + row * row_size_aligned; \ + const uint8_t * s0_row = s_spad[0] + row * row_size_aligned; \ + const uint8_t * s1_row = s_spad[1] + row * row_size_aligned; \ + HVX_ADD_FN(d_row, s0_row, s1_row, ne0); \ + for (uint32_t s = 2; s < n_ranks; s++) { \ + const uint8_t * ss_row = s_spad[s] + row * row_size_aligned; \ + HVX_ADD_FN(d_row, d_row, ss_row, ne0); \ + } \ + if (HAS_ADD) { \ + const uint8_t * res_row = IS_ROW_BCAST ? res_spad_base : (r_spad + row * row_size_aligned); \ + HVX_ADD_FN(d_row, d_row, res_row, ne0); \ + } \ + } \ + htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) r); \ + for (uint32_t d = 0; d < n_dsts; d++) { \ + uint8_t * d_ddr = (uint8_t *) octx->dsts[d]->data + r * octx->dsts[d]->nb[1]; \ + dma_queue_push(q, dma_make_ptr(d_ddr, d_spad), octx->dsts[d]->nb[1], row_size_aligned, row_bytes, cur_rows); \ + } \ + if (r_prefetch < r1) { \ + uint32_t next_rows = MIN(block_rows, r1 - r_prefetch); \ + for (uint32_t s = 0; s < n_ranks; s++) { \ + const uint8_t * s_next = (const uint8_t *) octx->src[s]->data + r_prefetch * octx->src[s]->nb[1]; \ + dma_queue_push(q, dma_make_ptr(s_spad[s], s_next), row_size_aligned, octx->src[s]->nb[1], row_bytes, next_rows); \ + } \ + if (HAS_ADD && !IS_ROW_BCAST) { \ + const uint8_t * r_next = (const uint8_t *) octx->src[2 * n_ranks]->data + r_prefetch * octx->src[2 * n_ranks]->nb[1]; \ + dma_queue_push(q, dma_make_ptr(r_spad, r_next), row_size_aligned, octx->src[2 * n_ranks]->nb[1], row_bytes, next_rows); \ + } \ + r_prefetch += next_rows; \ + } \ + r += cur_rows; \ + } \ + dma_queue_flush(q); \ +} + +DEFINE_ALLREDUCE_THREAD_DMA_2D(f16, __fp16, hvx_add_f16_aaa, 0, 0) +DEFINE_ALLREDUCE_THREAD_DMA_2D(f32, float, hvx_add_f32_aaa, 0, 0) +DEFINE_ALLREDUCE_THREAD_DMA_2D(add_f16, __fp16, hvx_add_f16_aaa, 1, 0) +DEFINE_ALLREDUCE_THREAD_DMA_2D(add_f32, float, hvx_add_f32_aaa, 1, 0) +DEFINE_ALLREDUCE_THREAD_DMA_2D(add_bcast_f16, __fp16, hvx_add_f16_aaa, 1, 1) +DEFINE_ALLREDUCE_THREAD_DMA_2D(add_bcast_f32, float, hvx_add_f32_aaa, 1, 1) + +int op_allreduce(struct htp_ops_context * octx) { + const struct htp_allreduce_kernel_params * kparams = (const struct htp_allreduce_kernel_params *) octx->kernel_params; + const struct htp_tensor * dst = octx->dst; + + const uint32_t rank = (uint32_t) kparams->rank; + const uint32_t n_ranks = (uint32_t) kparams->n_ranks; + + if (n_ranks < 2 || n_ranks > HTP_ALLREDUCE_MAX_RANKS || rank >= n_ranks) { + return HTP_STATUS_INVAL_PARAMS; + } + + if (dst->type != HTP_TYPE_F16 && dst->type != HTP_TYPE_F32) { + return HTP_STATUS_NO_SUPPORT; + } + + const uint32_t nelem = dst->ne[0] * dst->ne[1] * dst->ne[2] * dst->ne[3]; + const uint32_t fence_seq_entry = (uint32_t) octx->op_params[0]; + const uint32_t fence_seq_exit = (uint32_t) octx->op_params[1]; + + // 1. Entry Barrier: Synchronize all ranks before reading + struct htp_thread_trace * tr0 = &octx->ctx->trace[0]; + htp_trace_event_start(tr0, HTP_TRACE_EVT_FENCE, (uint16_t) fence_seq_entry); + + const struct htp_tensor * my_sync = octx->src[n_ranks + rank]; + atomic_uint * my_fence = (atomic_uint *) my_sync->data; + + atomic_store(&my_fence[0], fence_seq_entry); + asm volatile ("syncht" : : : "memory"); + Q6_dccleaninva_A((void *) my_fence); + + for (uint32_t j = 0; j < n_ranks; j++) { + if (j == rank) continue; + const struct htp_tensor * peer_sync = octx->src[n_ranks + j]; + atomic_uint * peer_fence = (atomic_uint *) peer_sync->data; + uint64_t spins = 0; + while (1) { + Q6_dccleaninva_A((void *) peer_fence); + uint32_t val = atomic_load(&peer_fence[0]); + if (val == fence_seq_entry || val == fence_seq_exit) { + break; + } + if (++spins > HTP_FENCE_TIMEOUT) { + FARF(ERROR, "ggml-hex: allreduce entry fence-wait TIMEOUT: rank %u waiting on %u (fence %p seq %u)\n", rank, j, peer_fence, fence_seq_entry); + return HTP_STATUS_INTERNAL_ERR; + } + hex_pause(); + } + } + asm volatile ("syncht" : : : "memory"); + + htp_trace_event_stop(tr0, HTP_TRACE_EVT_FENCE, (uint16_t) fence_seq_entry); + + // 2. Multi-threaded Reduction across assigned rank chunk + if (nelem > 0) { + const uint32_t n_threads = (uint32_t) kparams->n_threads; + const uint32_t block_elems = (uint32_t) kparams->block_elems; + const uint32_t elems_per_thread = (uint32_t) kparams->elems_per_thread; + const uint32_t vtcm_size_per_thread = (uint32_t) kparams->vtcm_size_per_thread; + + const bool has_add = (octx->op == HTP_OP_ALLREDUCE_ADD); + + struct htp_allreduce_context actx; + actx.octx = octx; + actx.n_ranks = n_ranks; + actx.n_dsts = (uint32_t) kparams->n_dsts ? (uint32_t) kparams->n_dsts : n_ranks; + actx.nelem = nelem; + actx.ne0 = (uint32_t) kparams->ne0; + actx.ne1 = (uint32_t) kparams->ne1; + actx.row_size_aligned = (uint32_t) kparams->row_size_aligned; + actx.rank_elem_start = (uint32_t) kparams->rank_elem_start; + actx.rank_nelem = (uint32_t) kparams->rank_nelem; + actx.elems_per_thread = elems_per_thread; + actx.block_elems = block_elems; + actx.vtcm_size_per_thread = vtcm_size_per_thread; + actx.is_row_bcast = (kparams->is_row_bcast != 0); + + work_queue_func_t reduce_fun = NULL; + switch (kparams->kernel_type) { + case HTP_ALLREDUCE_KERNEL_DMA_1D: + if (has_add) { + reduce_fun = (dst->type == HTP_TYPE_F16) ? allreduce_thread_dma_1d_add_f16 : allreduce_thread_dma_1d_add_f32; + } else { + reduce_fun = (dst->type == HTP_TYPE_F16) ? allreduce_thread_dma_1d_f16 : allreduce_thread_dma_1d_f32; + } + break; + case HTP_ALLREDUCE_KERNEL_DMA_2D: + if (has_add) { + if (kparams->is_row_bcast) { + reduce_fun = (dst->type == HTP_TYPE_F16) ? allreduce_thread_dma_2d_add_bcast_f16 : allreduce_thread_dma_2d_add_bcast_f32; + } else { + reduce_fun = (dst->type == HTP_TYPE_F16) ? allreduce_thread_dma_2d_add_f16 : allreduce_thread_dma_2d_add_f32; + } + } else { + reduce_fun = (dst->type == HTP_TYPE_F16) ? allreduce_thread_dma_2d_f16 : allreduce_thread_dma_2d_f32; + } + break; + default: + return HTP_STATUS_NO_SUPPORT; + } + + uint8_t * vtcm_ptr = (uint8_t *) octx->ctx->vtcm_base; + for (uint32_t s = 0; s < n_ranks; s++) { + actx.src_spad_base[s] = vtcm_ptr; + vtcm_ptr += n_threads * vtcm_size_per_thread; + } + actx.dst_spad_base = vtcm_ptr; + vtcm_ptr += n_threads * vtcm_size_per_thread; + if (has_add) { + actx.res_spad_base = vtcm_ptr; + vtcm_ptr += (actx.is_row_bcast ? 1 : n_threads) * vtcm_size_per_thread; + } + + if (has_add && actx.is_row_bcast) { + const uint8_t * r_ddr = (const uint8_t *) octx->src[2 * n_ranks]->data; + const uint32_t row_bytes = actx.ne0 * (dst->type == HTP_TYPE_F16 ? sizeof(__fp16) : sizeof(float)); + dma_queue * q = octx->ctx->dma[0]; + dma_queue_push(q, dma_make_ptr(actx.res_spad_base, r_ddr), actx.row_size_aligned, 0, row_bytes, 1); + dma_queue_pop(q); + } + + work_queue_run(octx->ctx->work_queue, reduce_fun, &actx, n_threads); + } + + // 4. Exit Barrier: Synchronize all ranks after writing + htp_trace_event_start(tr0, HTP_TRACE_EVT_FENCE, (uint16_t) fence_seq_exit); + + atomic_store(&my_fence[0], fence_seq_exit); + asm volatile ("syncht" : : : "memory"); + Q6_dccleaninva_A((void *) my_fence); + + for (uint32_t j = 0; j < n_ranks; j++) { + if (j == rank) continue; + const struct htp_tensor * peer_sync = octx->src[n_ranks + j]; + atomic_uint * peer_fence = (atomic_uint *) peer_sync->data; + uint64_t spins = 0; + while (1) { + Q6_dccleaninva_A((void *) peer_fence); + uint32_t val = atomic_load(&peer_fence[0]); + if (val == fence_seq_exit) { + break; + } + if (++spins > HTP_FENCE_TIMEOUT) { + FARF(ERROR, "ggml-hex: allreduce exit fence-wait TIMEOUT: rank %u waiting on %u (fence %p seq %u)\n", rank, j, peer_fence, fence_seq_exit); + return HTP_STATUS_INTERNAL_ERR; + } + hex_pause(); + } + } + asm volatile ("syncht" : : : "memory"); + + htp_trace_event_stop(tr0, HTP_TRACE_EVT_FENCE, (uint16_t) fence_seq_exit); + + return HTP_STATUS_OK; +} diff --git a/ggml/src/ggml-hexagon/htp/allreduce-ops.h b/ggml/src/ggml-hexagon/htp/allreduce-ops.h new file mode 100644 index 00000000000..de447d87e91 --- /dev/null +++ b/ggml/src/ggml-hexagon/htp/allreduce-ops.h @@ -0,0 +1,40 @@ +#ifndef ALLREDUCE_OPS_H +#define ALLREDUCE_OPS_H + +#include + +#define HTP_ALLREDUCE_MAX_RANKS 4 + +#ifdef __cplusplus +extern "C" { +#endif + +enum htp_allreduce_kernel_type { + HTP_ALLREDUCE_KERNEL_UNSUPPORTED = 0, + HTP_ALLREDUCE_KERNEL_DMA_1D, + HTP_ALLREDUCE_KERNEL_DMA_2D, +}; + +struct htp_allreduce_kernel_params { + int32_t rank; + int32_t n_ranks; + int32_t n_threads; + int32_t block_elems; // 1D: block_elems, 2D: block_rows + int32_t elems_per_thread; // 1D: nelem_per_thread, 2D: nrows_per_thread + int32_t vtcm_size_per_thread; + int32_t vtcm_size; + int32_t kernel_type; + int32_t ne0; + int32_t ne1; + int32_t row_size_aligned; + int32_t rank_elem_start; + int32_t rank_nelem; + int32_t n_dsts; + int32_t is_row_bcast; +}; + +#ifdef __cplusplus +} +#endif + +#endif /* ALLREDUCE_OPS_H */ diff --git a/ggml/src/ggml-hexagon/htp/cpy-ops.c b/ggml/src/ggml-hexagon/htp/cpy-ops.c index ae507effa51..15bc8dc244f 100644 --- a/ggml/src/ggml-hexagon/htp/cpy-ops.c +++ b/ggml/src/ggml-hexagon/htp/cpy-ops.c @@ -4,6 +4,7 @@ #include #include +#include #include #include @@ -14,6 +15,7 @@ #include "htp-ops.h" #include "htp-ops.h" #include "hvx-utils.h" +#include "htp-tensor.h" struct htp_copy_context { struct htp_ops_context * octx; @@ -78,7 +80,7 @@ static void cpy_thread_##NAME##_sameshape(unsigned int nth, unsigned int ith, vo } \ } -DEFINE_CPY_SAMESHAPE(f32, float, 4) +DEFINE_CPY_SAMESHAPE(f32, float, 4) DEFINE_CPY_SAMESHAPE(f16, __fp16, 2) #define DEFINE_CPY_RESHAPE(NAME, ELEM_TYPE, ELEM_SIZE) \ @@ -179,7 +181,7 @@ static void cpy_thread_##NAME##_reshape(unsigned int nth, unsigned int ith, void } \ } -DEFINE_CPY_RESHAPE(f32, float, 4) +DEFINE_CPY_RESHAPE(f32, float, 4) DEFINE_CPY_RESHAPE(f16, __fp16, 2) static void cpy_thread_f16_f32_sameshape(unsigned int nth, unsigned int ith, void * data) { @@ -232,6 +234,41 @@ static void cpy_thread_f32_f16_sameshape(unsigned int nth, unsigned int ith, voi } } +static inline void cpy_dma_sametype_sameshape( + struct htp_ops_context * octx, + const struct htp_tensor * dst, + const struct htp_tensor * src0, + uint32_t elem_size, + uint32_t ne00, uint32_t ne01, uint32_t ne02, uint32_t ne03, + uint32_t nb01, uint32_t nb02, uint32_t nb03, + uint32_t nb1, uint32_t nb2, uint32_t nb3 +) { + const bool contiguous_outer = + (ne02 == 1 || (nb02 == ne01 * nb01 && nb2 == ne01 * nb1)) && + (ne03 == 1 || (nb03 == ne02 * nb02 && nb3 == ne02 * nb2)); + + dma_queue * q = octx->ctx->dma[0]; + + if (contiguous_outer) { + dma_queue_push(q, dma_make_ptr((void *) dst->data, (const void *) src0->data), nb1, nb01, ne00 * elem_size, ne01 * ne02 * ne03); + dma_queue_pop(q); + return; + } + + for (uint32_t i03 = 0; i03 < ne03; i03++) { + for (uint32_t i02 = 0; i02 < ne02; i02++) { + uint8_t* dst_ptr = (uint8_t*) dst->data + i02*nb2 + i03*nb3; + uint8_t* src0_ptr = (uint8_t*) src0->data + i02*nb02 + i03*nb03; + if (!dma_queue_push(q, dma_make_ptr(dst_ptr, src0_ptr), nb1, nb01, ne00 * elem_size, ne01)) { + dma_queue_flush(q); + dma_queue_push(q, dma_make_ptr(dst_ptr, src0_ptr), nb1, nb01, ne00 * elem_size, ne01); + } + } + } + + dma_queue_flush(q); +} + int op_cpy(struct htp_ops_context * octx) { cpy_preamble; @@ -264,14 +301,11 @@ int op_cpy(struct htp_ops_context * octx) { ct.src0_nrows_per_thread = (nr + n_threads - 1) / n_threads; - worker_callback_t copy_fun; + worker_callback_t copy_fun = NULL; + bool use_dma = false; if (sametype && sameshape) { - if (src0->type == HTP_TYPE_F32) { - copy_fun = cpy_thread_f32_sameshape; - } else { - copy_fun = cpy_thread_f16_sameshape; - } + use_dma = true; } else if (sameshape) { /**/ if (dst->type == HTP_TYPE_F16 && src0->type == HTP_TYPE_F32) copy_fun = cpy_thread_f16_f32_sameshape; @@ -289,7 +323,28 @@ int op_cpy(struct htp_ops_context * octx) { return HTP_STATUS_NO_SUPPORT; } - worker_pool_run_func(octx->ctx->worker_pool, copy_fun, &ct, n_threads); + if (use_dma) { + cpy_dma_sametype_sameshape(octx, dst, src0, ct.src0_type_size, ne00, ne01, ne02, ne03, nb01, nb02, nb03, nb1, nb2, nb3); + } else { + worker_pool_run_func(octx->ctx->worker_pool, copy_fun, &ct, n_threads); + } + + const struct htp_tensor *sync = octx->src[1]; + if (sync) { + if (!use_dma) { + // htp_tensor_flush_all(octx->ctx, octx->dsts, 1); + qurt_mem_cache_clean((qurt_addr_t) 0, 0, QURT_MEM_CACHE_FLUSH_INVALIDATE_ALL, QURT_MEM_DCACHE); + } + + atomic_uint * sync_fence = (atomic_uint *) sync->data; + const uint32_t seq = (uint32_t) octx->op_params[0]; + + atomic_store(&sync_fence[0], seq); + asm volatile ("syncht" : : : "memory"); + Q6_dccleaninva_A((void *) sync_fence); + + FARF(HIGH, "ggml-hex: sync-release : fence %p seq %u\n", sync_fence, seq); + } return HTP_STATUS_OK; } diff --git a/ggml/src/ggml-hexagon/htp/dma-queue.h b/ggml/src/ggml-hexagon/htp/dma-queue.h index 264284bda82..190ca3a9b9e 100644 --- a/ggml/src/ggml-hexagon/htp/dma-queue.h +++ b/ggml/src/ggml-hexagon/htp/dma-queue.h @@ -244,17 +244,18 @@ static inline dma_ptr dma_queue_pop(dma_queue * q) { return dptr; } - dma_descriptor_2d * desc = &r->desc[r->pop_idx]; + dptr = r->dptr[r->pop_idx]; + + volatile dma_descriptor_2d * desc = &r->desc[r->pop_idx]; // Wait for desc to complete if (!desc->done) { + // FARF(ALWAYS, "dma-poll: idx %u dst %p src %p", r->pop_idx, dptr.dst, dptr.src); while (!desc->done) { dmpoll(); } } - dptr = r->dptr[r->pop_idx]; - htp_trace_event_stop(r->trace, HTP_TRACE_EVT_DMA, r->pop_idx); r->pop_idx = (r->pop_idx + 1) & r->idx_mask; diff --git a/ggml/src/ggml-hexagon/htp/flash-attn-ops.c b/ggml/src/ggml-hexagon/htp/flash-attn-ops.c index 81765629046..c76b4d3a3ac 100644 --- a/ggml/src/ggml-hexagon/htp/flash-attn-ops.c +++ b/ggml/src/ggml-hexagon/htp/flash-attn-ops.c @@ -30,6 +30,8 @@ #include "ggml-common.h" #include "htp-ctx.h" #include "htp-ops.h" +#include "htp-tensor.h" +#include "hvx-quant.h" #include "flash-attn-ops.h" #include "hvx-fa-kernels.h" @@ -85,12 +87,17 @@ struct htp_fa_context { uint8_t * spad_m; uint8_t * spad_a; + const struct htp_tensor * k; + const struct htp_tensor * v; + uint64_t t_start; }; struct hmx_fa_context { const struct htp_ops_context * octx; const struct htp_tensor * sinks; // attention sinks (src[4]), NULL if absent + const struct htp_tensor * k; + const struct htp_tensor * v; bool pipeline; // true when n_kv_blocks >= FA_MIN_KV_BLOCKS && n_threads >= 2 uint32_t n_threads; @@ -214,8 +221,8 @@ static void flash_attn_ext_f16_thread(unsigned int nth, unsigned int ith, void * const uint32_t DV = nev0; const size_t size_q_row = DK * ((q->type == HTP_TYPE_F32) ? 4 : 2); - const size_t size_k_row = DK * sizeof(__fp16); - const size_t size_v_row = DV * sizeof(__fp16); + const size_t size_k_row = htp_tensor_get_row_size(k->type, DK); + const size_t size_v_row = htp_tensor_get_row_size(v->type, DV); // Scratchpad buffers for Q, K, V, Mask, and VKQ32 accumulator uint8_t * spad_q = factx->spad_q + factx->size_q_block * ith; @@ -364,6 +371,23 @@ static void flash_attn_ext_f16_thread(unsigned int nth, unsigned int ith, void * uint8_t * v_base = dma_queue_pop(dma).dst; // V __fp16 * m_base = mask ? dma_queue_pop(dma).dst : NULL; // M + if (factx->k->type == HTP_TYPE_Q8_0) { + htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_FA_K_PREP, ir); + for (uint32_t r = 0; r < current_block_size; ++r) { + __fp16 * row_k = (__fp16 *)(k_base + r * factx->size_k_row_padded); + hvx_dequantize_row_q8_0_f16(row_k, row_k, DK); + } + htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_FA_K_PREP, ir); + } + if (factx->v->type == HTP_TYPE_Q8_0) { + htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_FA_V_PREP, ir); + for (uint32_t r = 0; r < current_block_size; ++r) { + __fp16 * row_v = (__fp16 *)(v_base + r * factx->size_v_row_padded); + hvx_dequantize_row_q8_0_f16(row_v, row_v, DV); + } + htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_FA_V_PREP, ir); + } + htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_FA_QK, ir); // Inner loop processing the block from VTCM @@ -625,6 +649,12 @@ static void fa_k_interleave_thread(unsigned int n, unsigned int i, void * data) struct htp_thread_trace * tr = &factx->octx->ctx->trace[i]; htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_FA_K_PREP, (uint16_t) (args->kv_start + start)); + if (factx->k->type == HTP_TYPE_Q8_0) { + for (uint32_t r = start; r < end; ++r) { + __fp16 * row_k = (__fp16 *)((char *)args->curr_k + r * args->src_stride * sizeof(__fp16)); + hvx_dequantize_row_q8_0_f16(row_k, row_k, factx->DK); + } + } hmx_interleave_rows_to_tiles(factx->vtcm_k_tiles[args->buf_idx], (const __fp16 *) args->curr_k, total_rows, factx->DK, args->src_stride, start, end); htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_FA_K_PREP, (uint16_t) (args->kv_start + start)); @@ -673,6 +703,12 @@ static void fa_v_interleave_thread(unsigned int n, unsigned int i, void * data) struct htp_thread_trace * tr = &factx->octx->ctx->trace[i]; htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_FA_V_PREP, (uint16_t) (args->kv_start + start)); + if (factx->v->type == HTP_TYPE_Q8_0) { + for (uint32_t r = start; r < end; ++r) { + __fp16 * row_v = (__fp16 *)((char *)args->v_src + r * args->src_stride * sizeof(__fp16)); + hvx_dequantize_row_q8_0_f16(row_v, row_v, factx->DV); + } + } hmx_interleave_cols_to_tiles(v_tiles_dst, (const __fp16 *) args->v_src, total_rows, factx->DV, args->src_stride, (uint32_t) args->n_col_tiles, start, end); htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_FA_V_PREP, (uint16_t) (args->kv_start + start)); @@ -1809,6 +1845,8 @@ int hmx_flash_attn_ext(struct htp_ops_context * octx) { memset(&factx, 0, sizeof(factx)); factx.octx = octx; factx.sinks = octx->src[4]; // NULL if this op has no attention sinks + factx.k = k; + factx.v = v; factx.n_threads = kparams->n_threads; factx.DK = DK; factx.DV = DV; @@ -1853,10 +1891,10 @@ int hmx_flash_attn_ext(struct htp_ops_context * octx) { // ======== VTCM allocation (GQA-aware) ======== // K/V row sizes drive the DMA descriptors (not the VTCM layout) and are used // throughout the KV loop below. - const size_t size_k_row = DK * sizeof(__fp16); - const size_t size_v_row = DV * sizeof(__fp16); - const size_t size_k_row_padded = hex_round_up(size_k_row, 128); - const size_t size_v_row_padded = hex_round_up(size_v_row, 128); + const size_t size_k_row = htp_tensor_get_row_size(k->type, DK); + const size_t size_v_row = htp_tensor_get_row_size(v->type, DV); + const size_t size_k_row_padded = hex_round_up(DK * sizeof(__fp16), 128); + const size_t size_v_row_padded = hex_round_up(DV * sizeof(__fp16), 128); // Build the VTCM layout once (shared with the host estimator) and place every // scratch buffer at its computed offset. @@ -2348,7 +2386,9 @@ int op_flash_attn_ext(struct htp_ops_context * octx) { const struct htp_tensor * dst = octx->dst; // Check support - if ((q->type != HTP_TYPE_F16 && q->type != HTP_TYPE_F32) || k->type != HTP_TYPE_F16 || v->type != HTP_TYPE_F16) { + if ((q->type != HTP_TYPE_F16 && q->type != HTP_TYPE_F32) || + (k->type != HTP_TYPE_F16 && k->type != HTP_TYPE_Q8_0) || + (v->type != HTP_TYPE_F16 && v->type != HTP_TYPE_Q8_0)) { return HTP_STATUS_NO_SUPPORT; } @@ -2364,6 +2404,8 @@ int op_flash_attn_ext(struct htp_ops_context * octx) { struct htp_fa_context factx; factx.octx = octx; + factx.k = k; + factx.v = v; factx.t_start = HAP_perf_get_qtimer_count(); diff --git a/ggml/src/ggml-hexagon/htp/get-rows-ops.c b/ggml/src/ggml-hexagon/htp/get-rows-ops.c index bf7063e9880..05769d17f74 100644 --- a/ggml/src/ggml-hexagon/htp/get-rows-ops.c +++ b/ggml/src/ggml-hexagon/htp/get-rows-ops.c @@ -12,18 +12,17 @@ #include "ggml-common.h" #include "htp-ctx.h" #include "htp-ops.h" -#include "htp-ops.h" +#include "htp-tensor.h" #include "hvx-utils.h" +#include "hvx-quant.h" +#include "get-rows-ops.h" +#include "work-queue.h" struct get_rows_context { struct htp_ops_context * octx; - uint32_t tasks_per_thread; - uint32_t total_tasks; - uint32_t chunks_per_row; - uint32_t chunk_size; - struct fastdiv_values get_rows_div_ne10; - struct fastdiv_values get_rows_div_ne10_ne11; - struct fastdiv_values get_rows_div_chunks_per_row; + const struct htp_get_rows_kernel_params * kparams; + struct htp_get_rows_vtcm_layout vtcm_layout; + uint8_t * vtcm_base; }; #define get_rows_preamble \ @@ -56,102 +55,161 @@ struct get_rows_context { \ const uint32_t nr = ne10 * ne11 * ne12; -static void get_rows_thread_f32_f32_dma(unsigned int nth, unsigned int ith, void *data) { - struct get_rows_context * grctx = (struct get_rows_context *)data; - struct htp_ops_context * octx = grctx->octx; - get_rows_preamble; - - uint64_t qt = HAP_perf_get_qtimer_count(); - - const uint32_t dr = grctx->tasks_per_thread; - const uint32_t ir0 = dr * ith; - if (ir0 >= grctx->total_tasks) { - return; - } - const uint32_t ir1 = MIN(ir0 + dr, grctx->total_tasks); - - const bool is_i32 = (octx->src[1]->type == HTP_TYPE_I32); - - dma_queue * dma_queue = octx->ctx->dma[ith]; - for (uint32_t i = ir0; i < ir1; ++i) { - const uint32_t i12 = fastdiv(i, &grctx->get_rows_div_ne10_ne11); - const uint32_t rem = i - i12 * ne11 * ne10; - const uint32_t i11 = fastdiv(rem, &grctx->get_rows_div_ne10); - const uint32_t i10 = rem - i11 * ne10; - - const uintptr_t src1_addr = octx->src[1]->data + i10*nb10 + i11*nb11 + i12*nb12; - uint32_t i01 = is_i32 ? *(int32_t *)src1_addr : *(int64_t *)src1_addr; - - if (i01 >= ne01) { - continue; - } - - const uintptr_t src0_ptr = octx->src[0]->data + i01*nb01 + i11*nb02 + i12*nb03; - const uintptr_t dst_ptr = octx->dst->data + i10*nb1 + i11*nb2 + i12*nb3; - - while (!dma_queue_push(dma_queue, dma_make_ptr((void *)dst_ptr, (const void *)src0_ptr), nb1, nb01, ne00 * sizeof(float), 1)) { - dma_queue_pop(dma_queue); - } - } - dma_queue_flush(dma_queue); - - qt = HAP_perf_qtimer_count_to_us(HAP_perf_get_qtimer_count() - qt); - FARF(HIGH, "get-rows-f32-f32-dma %d/%d: %ux%ux%ux%u (%u:%u) x %ux%ux%ux%u -> %ux%ux%ux%u usec %u\n", ith, nth, - ne00, ne01, ne02, ne03, ir0, ir1, ne10, ne11, ne12, ne13, ne0, ne1, ne2, ne3, (unsigned) qt); +#define GET_ROWS_THREAD_ST_FN(IDX_TYPE) \ +static void get_rows_thread_st_##IDX_TYPE(unsigned int nth, unsigned int ith, void *data) { \ + struct get_rows_context * grctx = (struct get_rows_context *)data; \ + struct htp_ops_context * octx = grctx->octx; \ + const struct htp_get_rows_kernel_params * kparams = grctx->kparams; \ + get_rows_preamble; \ + const uint32_t dr = kparams->tasks_per_thread; \ + const uint32_t ir0 = dr * ith; \ + if (ir0 >= kparams->total_tasks) { \ + return; \ + } \ + const uint32_t ir1 = MIN(ir0 + dr, kparams->total_tasks); \ + const uint32_t row_size_bytes = htp_tensor_get_row_size(octx->src[0]->type, ne00); \ + dma_queue * dma_queue = octx->ctx->dma[ith]; \ + for (uint32_t i = ir0; i < ir1; ++i) { \ + const uint32_t i12 = fastdiv(i, &kparams->div_ne10_ne11); \ + const uint32_t rem = i - i12 * ne11 * ne10; \ + const uint32_t i11 = fastdiv(rem, &kparams->div_ne10); \ + const uint32_t i10 = rem - i11 * ne10; \ + const IDX_TYPE * src1_ptr = (const IDX_TYPE *)(octx->src[1]->data + i10*nb10 + i11*nb11 + i12*nb12); \ + const uint32_t i01 = (uint32_t)*src1_ptr; \ + assert(i01 < ne01); \ + const uint32_t q02 = fastdiv(i11, &kparams->div_ne02); \ + const uint32_t i02 = i11 - q02 * ne02; \ + const uint32_t q03 = fastdiv(i12, &kparams->div_ne03); \ + const uint32_t i03 = i12 - q03 * ne03; \ + const uintptr_t src0_ptr = octx->src[0]->data + i01*nb01 + i02*nb02 + i03*nb03; \ + const uintptr_t dst_ptr = octx->dst->data + i10*nb1 + i11*nb2 + i12*nb3; \ + while (!dma_queue_push(dma_queue, dma_make_ptr((void *)dst_ptr, (const void *)src0_ptr), nb1, nb01, \ + row_size_bytes, 1)) { \ + dma_queue_pop(dma_queue); \ + } \ + } \ + dma_queue_flush(dma_queue); \ } -static void get_rows_thread_f32_f32_hvx(unsigned int nth, unsigned int ith, void *data) { - struct get_rows_context * grctx = (struct get_rows_context *)data; - struct htp_ops_context * octx = grctx->octx; - get_rows_preamble; - - uint64_t qt = HAP_perf_get_qtimer_count(); - - const uint32_t dr = grctx->tasks_per_thread; - const uint32_t ir0 = dr * ith; - if (ir0 >= grctx->total_tasks) { - return; - } - const uint32_t ir1 = MIN(ir0 + dr, grctx->total_tasks); - - const bool is_i32 = (octx->src[1]->type == HTP_TYPE_I32); +GET_ROWS_THREAD_ST_FN(int32_t) +GET_ROWS_THREAD_ST_FN(int64_t) + +#define GET_ROWS_THREAD_DT_FN(TYPE_NAME, SRC0_SIZE_EXPR, IDX_TYPE, COMPUTE_EXPR) \ +static void get_rows_thread_##TYPE_NAME##_##IDX_TYPE(unsigned int nth, unsigned int ith, void *data) { \ + struct get_rows_context * grctx = (struct get_rows_context *)data; \ + struct htp_ops_context * octx = grctx->octx; \ + const struct htp_get_rows_kernel_params * kparams = grctx->kparams; \ + get_rows_preamble; \ + struct htp_thread_trace * tr = &octx->ctx->trace[ith]; \ + const uint32_t dr = kparams->tasks_per_thread; \ + const uint32_t ir0 = dr * ith; \ + if (ir0 >= kparams->total_tasks) { \ + return; \ + } \ + const uint32_t ir1 = MIN(ir0 + dr, kparams->total_tasks); \ + const uint32_t chunks_per_row = kparams->chunks_per_row; \ + const uint32_t chunk_size = kparams->chunk_size; \ + dma_queue * dma_queue = octx->ctx->dma[ith]; \ + const struct htp_get_rows_vtcm_layout * vtcm_layout = &grctx->vtcm_layout; \ + uint8_t * vtcm_src0 = grctx->vtcm_base + vtcm_layout->off_src0 + ith * vtcm_layout->src0_bytes_per_thread; \ + uint8_t * vtcm_dst = grctx->vtcm_base + vtcm_layout->off_dst + ith * vtcm_layout->dst_bytes_per_thread; \ + for (uint32_t step = 0, spad_idx = 0; step < ir1 - ir0 && spad_idx < 2; ++step, spad_idx++) { \ + const uint32_t i = ir0 + step; \ + const uint32_t row_idx = fastdiv(i, &kparams->div_chunks_per_row); \ + const uint32_t chunk_idx = i - row_idx * chunks_per_row; \ + const uint32_t i12 = fastdiv(row_idx, &kparams->div_ne10_ne11); \ + const uint32_t rem = row_idx - i12 * ne11 * ne10; \ + const uint32_t i11 = fastdiv(rem, &kparams->div_ne10); \ + const uint32_t i10 = rem - i11 * ne10; \ + const IDX_TYPE * src1_ptr = (const IDX_TYPE *)(octx->src[1]->data + i10*nb10 + i11*nb11 + i12*nb12); \ + const uint32_t i01 = (uint32_t)*src1_ptr; \ + assert(i01 < ne01); \ + const uint32_t q02 = fastdiv(i11, &kparams->div_ne02); \ + const uint32_t i02 = i11 - q02 * ne02; \ + const uint32_t q03 = fastdiv(i12, &kparams->div_ne03); \ + const uint32_t i03 = i12 - q03 * ne03; \ + const uint32_t offset = chunk_idx * chunk_size; \ + const uint32_t cur_elems = (offset < ne00) ? MIN(chunk_size, ne00 - offset) : 0; \ + const uint32_t cur_src0_bytes = SRC0_SIZE_EXPR(cur_elems); \ + const uint32_t cur_dst_bytes = cur_elems * sizeof(float); \ + const uintptr_t src0_ptr = octx->src[0]->data + i01*nb01 + i02*nb02 + i03*nb03 + SRC0_SIZE_EXPR(offset); \ + dma_queue_push(dma_queue, \ + dma_make_ptr((void *)(uintptr_t)octx->dst->data, \ + vtcm_dst + spad_idx * vtcm_layout->dst_spad_half_size), \ + cur_dst_bytes, vtcm_layout->dst_spad_half_size, cur_dst_bytes, 0); \ + dma_queue_push(dma_queue, \ + dma_make_ptr((void *)(vtcm_src0 + spad_idx * vtcm_layout->src0_spad_half_size), \ + (const void *)src0_ptr), \ + vtcm_layout->src0_spad_half_size, cur_src0_bytes, cur_src0_bytes, 1); \ + } \ + for (uint32_t step = 0; step < ir1 - ir0; ++step) { \ + const uint32_t i = ir0 + step; \ + void * dst_spad = (void *) dma_queue_pop(dma_queue).src; \ + void * src_spad = (void *) dma_queue_pop(dma_queue).dst; \ + const uint32_t row_idx = fastdiv(i, &kparams->div_chunks_per_row); \ + const uint32_t chunk_idx = i - row_idx * chunks_per_row; \ + const uint32_t i12 = fastdiv(row_idx, &kparams->div_ne10_ne11); \ + const uint32_t rem = row_idx - i12 * ne11 * ne10; \ + const uint32_t i11 = fastdiv(rem, &kparams->div_ne10); \ + const uint32_t i10 = rem - i11 * ne10; \ + const uint32_t offset = chunk_idx * chunk_size; \ + const uint32_t cur_elems = (offset < ne00) ? MIN(chunk_size, ne00 - offset) : 0; \ + const uint32_t cur_dst_bytes = cur_elems * sizeof(float); \ + htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, i); \ + COMPUTE_EXPR; \ + htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, i); \ + const uintptr_t dst_ptr = octx->dst->data + i10*nb1 + i11*nb2 + i12*nb3 + offset * sizeof(float); \ + dma_queue_push(dma_queue, \ + dma_make_ptr((void *)dst_ptr, (const void *)dst_spad), \ + cur_dst_bytes, vtcm_layout->dst_spad_half_size, cur_dst_bytes, 1); \ + const uint32_t next_step = step + 2; \ + if (next_step < ir1 - ir0) { \ + const uint32_t pi = ir0 + next_step; \ + const uint32_t prow_idx = fastdiv(pi, &kparams->div_chunks_per_row); \ + const uint32_t pchunk_idx = pi - prow_idx * chunks_per_row; \ + const uint32_t pi12 = fastdiv(prow_idx, &kparams->div_ne10_ne11); \ + const uint32_t prem = prow_idx - pi12 * ne11 * ne10; \ + const uint32_t pi11 = fastdiv(prem, &kparams->div_ne10); \ + const uint32_t pi10 = prem - pi11 * ne10; \ + const IDX_TYPE * psrc1_ptr = (const IDX_TYPE *)(octx->src[1]->data + pi10*nb10 + pi11*nb11 + pi12*nb12); \ + const uint32_t pi01 = (uint32_t)*psrc1_ptr; \ + assert(pi01 < ne01); \ + const uint32_t pq02 = fastdiv(pi11, &kparams->div_ne02); \ + const uint32_t pi02 = pi11 - pq02 * ne02; \ + const uint32_t pq03 = fastdiv(pi12, &kparams->div_ne03); \ + const uint32_t pi03 = pi12 - pq03 * ne03; \ + const uint32_t poffset = pchunk_idx * chunk_size; \ + const uint32_t pcur_elems = (poffset < ne00) ? MIN(chunk_size, ne00 - poffset) : 0; \ + const uint32_t pcur_src0_bytes = SRC0_SIZE_EXPR(pcur_elems); \ + const uintptr_t psrc0_ptr = \ + octx->src[0]->data + pi01*nb01 + pi02*nb02 + pi03*nb03 + SRC0_SIZE_EXPR(poffset); \ + dma_queue_push(dma_queue, \ + dma_make_ptr((void *)src_spad, (const void *)psrc0_ptr), \ + vtcm_layout->src0_spad_half_size, pcur_src0_bytes, pcur_src0_bytes, 1); \ + } \ + } \ + dma_queue_flush(dma_queue); \ +} - const uint32_t chunks_per_row = grctx->chunks_per_row; - const uint32_t chunk_size = grctx->chunk_size; - for (uint32_t i = ir0; i < ir1; ++i) { - const uint32_t row_idx = fastdiv(i, &grctx->get_rows_div_chunks_per_row); - const uint32_t chunk_idx = i - row_idx * chunks_per_row; +#define F32_BYTES(n) ((n) * sizeof(float)) +#define F16_BYTES(n) ((n) * sizeof(__fp16)) +#define Q8_0_BYTES(n) (((n) / 32) * sizeof(block_q8_0)) - const uint32_t i12 = fastdiv(row_idx, &grctx->get_rows_div_ne10_ne11); - const uint32_t rem = row_idx - i12 * ne11 * ne10; - const uint32_t i11 = fastdiv(rem, &grctx->get_rows_div_ne10); - const uint32_t i10 = rem - i11 * ne10; +GET_ROWS_THREAD_DT_FN(f32, F32_BYTES, int32_t, { if (cur_elems > 0) hvx_copy_f32_uu((uint8_t *)dst_spad, (const uint8_t *)src_spad, cur_elems); }) +GET_ROWS_THREAD_DT_FN(f32, F32_BYTES, int64_t, { if (cur_elems > 0) hvx_copy_f32_uu((uint8_t *)dst_spad, (const uint8_t *)src_spad, cur_elems); }) - const uintptr_t src1_addr = octx->src[1]->data + i10*nb10 + i11*nb11 + i12*nb12; - uint32_t i01 = is_i32 ? *(int32_t *)src1_addr : *(int64_t *)src1_addr; +GET_ROWS_THREAD_DT_FN(f16, F16_BYTES, int32_t, { hvx_dequantize_row_f16_f32((float *)dst_spad, src_spad, ne00); }) +GET_ROWS_THREAD_DT_FN(f16, F16_BYTES, int64_t, { hvx_dequantize_row_f16_f32((float *)dst_spad, src_spad, ne00); }) - if (i01 >= ne01) { - continue; - } - - const uint32_t offset = chunk_idx * chunk_size; - if (offset < ne00) { - const uint32_t copy_size = MIN(chunk_size, ne00 - offset); - const uintptr_t src0_ptr = octx->src[0]->data + i01*nb01 + i11*nb02 + i12*nb03 + offset * sizeof(float); - const uintptr_t dst_ptr = octx->dst->data + i10*nb1 + i11*nb2 + i12*nb3 + offset * sizeof(float); - hvx_copy_f32_uu((uint8_t *)dst_ptr, (const uint8_t *)src0_ptr, copy_size); - } - } - - qt = HAP_perf_qtimer_count_to_us(HAP_perf_get_qtimer_count() - qt); - FARF(HIGH, "get-rows-f32-f32-hvx %d/%d: %ux%ux%ux%u (%u:%u) x %ux%ux%ux%u -> %ux%ux%ux%u usec %u\n", ith, nth, - ne00, ne01, ne02, ne03, ir0, ir1, ne10, ne11, ne12, ne13, ne0, ne1, ne2, ne3, (unsigned) qt); -} +GET_ROWS_THREAD_DT_FN(q8_0, Q8_0_BYTES, int32_t, { hvx_dequantize_row_q8_0_f32((float *)dst_spad, src_spad, ne00); }) +GET_ROWS_THREAD_DT_FN(q8_0, Q8_0_BYTES, int64_t, { hvx_dequantize_row_q8_0_f32((float *)dst_spad, src_spad, ne00); }) int op_get_rows(struct htp_ops_context * octx) { - get_rows_preamble; + const struct htp_get_rows_kernel_params * kparams = (const struct htp_get_rows_kernel_params *) octx->kernel_params; - if (octx->src[0]->type != HTP_TYPE_F32) { + if (octx->src[0]->type != HTP_TYPE_F32 && + octx->src[0]->type != HTP_TYPE_F16 && + octx->src[0]->type != HTP_TYPE_Q8_0) { return HTP_STATUS_NO_SUPPORT; } @@ -167,52 +225,28 @@ int op_get_rows(struct htp_ops_context * octx) { return HTP_STATUS_OK; } - const uint32_t nb00 = octx->src[0]->nb[0]; - const uint32_t nb0 = octx->dst->nb[0]; - - const bool can_use_dma = (nb00 == sizeof(float)) && (nb0 == sizeof(float)); - const bool use_dma = can_use_dma && (ne00 >= 2048); - struct get_rows_context grctx; grctx.octx = octx; - grctx.get_rows_div_ne10 = init_fastdiv_values(octx->src[1]->ne[0]); - grctx.get_rows_div_ne10_ne11 = init_fastdiv_values(octx->src[1]->ne[0] * octx->src[1]->ne[1]); + grctx.kparams = kparams; + grctx.vtcm_base = (uint8_t *)octx->ctx->vtcm_base; - if (use_dma) { - grctx.chunks_per_row = 1; - grctx.chunk_size = ne00; - grctx.total_tasks = nr; - grctx.get_rows_div_chunks_per_row = init_fastdiv_values(1); + const uint32_t ne00 = octx->src[0]->ne[0]; + htp_get_rows_vtcm_layout_build(&grctx.vtcm_layout, octx->src[0]->type, ne00, kparams->n_threads); - const uint32_t n_threads = MIN(nr, octx->n_threads); - grctx.tasks_per_thread = (nr + n_threads - 1) / n_threads; + const bool is_i32 = (octx->src[1]->type == HTP_TYPE_I32); - worker_pool_run_func(octx->ctx->worker_pool, get_rows_thread_f32_f32_dma, &grctx, n_threads); + work_queue_func_t q_func = NULL; + if (kparams->use_dma) { + q_func = (work_queue_func_t)(is_i32 ? get_rows_thread_st_int32_t : get_rows_thread_st_int64_t); } else { - uint32_t chunks_per_row = 1; - uint32_t chunk_size = ne00; - uint32_t total_tasks = nr; - - if (nr < octx->n_threads) { - const uint32_t min_chunk_size = 1024; - uint32_t max_chunks = ne00 / min_chunk_size; - if (max_chunks == 0) { - max_chunks = 1; - } - chunks_per_row = MIN((octx->n_threads + nr - 1) / nr, max_chunks); - chunk_size = (ne00 + chunks_per_row - 1) / chunks_per_row; - total_tasks = nr * chunks_per_row; + switch (octx->src[0]->type) { + case HTP_TYPE_F32: q_func = (work_queue_func_t)(is_i32 ? get_rows_thread_f32_int32_t : get_rows_thread_f32_int64_t); break; + case HTP_TYPE_F16: q_func = (work_queue_func_t)(is_i32 ? get_rows_thread_f16_int32_t : get_rows_thread_f16_int64_t); break; + case HTP_TYPE_Q8_0: q_func = (work_queue_func_t)(is_i32 ? get_rows_thread_q8_0_int32_t : get_rows_thread_q8_0_int64_t); break; + default: return HTP_STATUS_NO_SUPPORT; } - - grctx.chunks_per_row = chunks_per_row; - grctx.chunk_size = chunk_size; - grctx.total_tasks = total_tasks; - grctx.get_rows_div_chunks_per_row = init_fastdiv_values(chunks_per_row); - - const uint32_t n_threads = MIN(total_tasks, octx->n_threads); - grctx.tasks_per_thread = (total_tasks + n_threads - 1) / n_threads; - - worker_pool_run_func(octx->ctx->worker_pool, get_rows_thread_f32_f32_hvx, &grctx, n_threads); } + + work_queue_run(octx->ctx->work_queue, q_func, &grctx, kparams->n_threads); return HTP_STATUS_OK; } diff --git a/ggml/src/ggml-hexagon/htp/get-rows-ops.h b/ggml/src/ggml-hexagon/htp/get-rows-ops.h new file mode 100644 index 00000000000..0e7c2ca8cf0 --- /dev/null +++ b/ggml/src/ggml-hexagon/htp/get-rows-ops.h @@ -0,0 +1,77 @@ +#ifndef HTP_GET_ROWS_OPS_H +#define HTP_GET_ROWS_OPS_H + +#include "hex-fastdiv.h" + +struct htp_get_rows_kernel_params { + int32_t n_threads; + int32_t use_dma; + int32_t chunks_per_row; + int32_t chunk_size; + int32_t total_tasks; + int32_t tasks_per_thread; + int32_t vtcm_size; + + // Fastdiv helpers + struct fastdiv_values div_ne10; + struct fastdiv_values div_ne10_ne11; + struct fastdiv_values div_chunks_per_row; + struct fastdiv_values div_ne02; + struct fastdiv_values div_ne03; +}; + +struct htp_get_rows_vtcm_layout { + size_t total_bytes; + size_t off_src0; + size_t off_dst; + + size_t src0_bytes_per_thread; + size_t dst_bytes_per_thread; + + size_t src0_spad_half_size; + size_t dst_spad_half_size; +}; + +static inline void htp_get_rows_vtcm_layout_build( + struct htp_get_rows_vtcm_layout * vtcm_layout, + int type, + uint32_t ne00, + uint32_t n_threads) { + + uint32_t src0_row_size = 0; + switch (type) { + case 0: // HTP_TYPE_F32 + src0_row_size = ne00 * 4; + break; + case 1: // HTP_TYPE_F16 + src0_row_size = ne00 * 2; + break; + case 8: // HTP_TYPE_Q8_0 + src0_row_size = (ne00 / 32) * 34; + break; + default: + src0_row_size = 0; + break; + } + + size_t src0_row_size_aligned = (src0_row_size + 255) & ~255; + size_t dst_row_size_aligned = (ne00 * sizeof(float) + 255) & ~255; + + vtcm_layout->src0_spad_half_size = src0_row_size_aligned; + vtcm_layout->dst_spad_half_size = dst_row_size_aligned; + + vtcm_layout->src0_bytes_per_thread = src0_row_size_aligned * 2; + vtcm_layout->dst_bytes_per_thread = dst_row_size_aligned * 2; + + vtcm_layout->off_src0 = 0; + vtcm_layout->off_dst = vtcm_layout->off_src0 + vtcm_layout->src0_bytes_per_thread * n_threads; + vtcm_layout->total_bytes = vtcm_layout->off_dst + vtcm_layout->dst_bytes_per_thread * n_threads; +} + +#if defined(__cplusplus) +static_assert(sizeof(struct htp_get_rows_kernel_params) <= 128, "htp_get_rows_kernel_params is too large for kernel_params blob"); +#else +_Static_assert(sizeof(struct htp_get_rows_kernel_params) <= 128, "htp_get_rows_kernel_params is too large for kernel_params blob"); +#endif + +#endif // HTP_GET_ROWS_OPS_H diff --git a/ggml/src/ggml-hexagon/htp/hex-utils.h b/ggml/src/ggml-hexagon/htp/hex-utils.h index 93e87efcb4c..1b396503000 100644 --- a/ggml/src/ggml-hexagon/htp/hex-utils.h +++ b/ggml/src/ggml-hexagon/htp/hex-utils.h @@ -39,17 +39,22 @@ static inline void hex_l2fetch_block(const void * addr, size_t size) { #define HEX_L2_LINE_SIZE 128 #define HEX_L2_BLOCK_SIZE (HEX_L2_LINE_SIZE * 4) // flush granularity (lines per loop iteration) +#define HEX_L2_FLUSH_IL_THRESHOLD 1024 // inline flush threshold #define HEX_L2_FLUSH_WQ_THRESHOLD (4 * 1024) #define HEX_L2_FLUSH_ALL_THRESHOLD (4 * 1024 * 1024) static inline void hex_l2flush(void * addr, size_t size) { const uint32_t s = ((uint32_t) addr) & ~(HEX_L2_LINE_SIZE - 1); const uint32_t e = (((uint32_t) addr) + size + HEX_L2_LINE_SIZE - 1) & ~(HEX_L2_LINE_SIZE - 1); - for (uint32_t i = s; i < e; i += HEX_L2_BLOCK_SIZE) { - Q6_dccleaninva_A((void *) i + HEX_L2_LINE_SIZE * 0); - Q6_dccleaninva_A((void *) i + HEX_L2_LINE_SIZE * 1); - Q6_dccleaninva_A((void *) i + HEX_L2_LINE_SIZE * 2); - Q6_dccleaninva_A((void *) i + HEX_L2_LINE_SIZE * 3); + const uint32_t eb = s + ((e - s) & ~(HEX_L2_BLOCK_SIZE - 1)); + for (uint32_t i = s; i < eb; i += HEX_L2_BLOCK_SIZE) { + Q6_dccleaninva_A((void *) (i + HEX_L2_LINE_SIZE * 0)); + Q6_dccleaninva_A((void *) (i + HEX_L2_LINE_SIZE * 1)); + Q6_dccleaninva_A((void *) (i + HEX_L2_LINE_SIZE * 2)); + Q6_dccleaninva_A((void *) (i + HEX_L2_LINE_SIZE * 3)); + } + for (uint32_t i = eb; i < e; i += HEX_L2_LINE_SIZE) { + Q6_dccleaninva_A((void *) i); } } diff --git a/ggml/src/ggml-hexagon/htp/htp-ctx.h b/ggml/src/ggml-hexagon/htp/htp-ctx.h index e0f9a0c40d1..88ecf144b94 100644 --- a/ggml/src/ggml-hexagon/htp/htp-ctx.h +++ b/ggml/src/ggml-hexagon/htp/htp-ctx.h @@ -117,8 +117,7 @@ struct htp_context { int op_matmul(struct htp_ops_context * octx); int op_matmul_id(struct htp_ops_context * octx); -int op_matmul_qkv(struct htp_ops_context * octx); -int op_matmul_ffn(struct htp_ops_context * octx); +int op_matmul_nx(struct htp_ops_context * octx); int op_binary(struct htp_ops_context * octx); int op_unary(struct htp_ops_context * octx); int op_sum_rows(struct htp_ops_context * octx); @@ -141,5 +140,6 @@ int op_solve_tri(struct htp_ops_context * octx); int op_gated_delta_net(struct htp_ops_context * octx); int op_pad(struct htp_ops_context * octx); int op_im2col(struct htp_ops_context * octx); +int op_allreduce(struct htp_ops_context * octx); #endif /* HTP_CTX_H */ diff --git a/ggml/src/ggml-hexagon/htp/htp-ops.h b/ggml/src/ggml-hexagon/htp/htp-ops.h index a138f062aa6..b4023b34d38 100644 --- a/ggml/src/ggml-hexagon/htp/htp-ops.h +++ b/ggml/src/ggml-hexagon/htp/htp-ops.h @@ -43,13 +43,6 @@ enum htp_data_type { -// Mask to enable various stages of the Ops. -// Used for debugging and profiling. -enum htp_op_stage { - HTP_OPSTAGE_QUEUE = (1 << 0), // Enable Queueing (ie calls into NPU) - HTP_OPSTAGE_COMPUTE = (1 << 1), // Enable Compute -}; - // Do not reorder first 4 (used as an index) enum htp_op_code { HTP_OP_MUL = 0, @@ -58,8 +51,7 @@ enum htp_op_code { HTP_OP_DIV = 3, HTP_OP_MUL_MAT, HTP_OP_MUL_MAT_ID, - HTP_OP_MUL_MAT_QKV, - HTP_OP_MUL_MAT_FFN, + HTP_OP_MUL_MAT_NX, HTP_OP_MUL_MAT_ADD, HTP_OP_RMS_NORM, HTP_OP_RMS_NORM_MUL, @@ -99,12 +91,15 @@ enum htp_op_code { HTP_OP_CONCAT, HTP_OP_CLAMP, HTP_OP_IM2COL, + HTP_OP_FENCE, + HTP_OP_ALLREDUCE, + HTP_OP_ALLREDUCE_ADD, HTP_OP_INVALID }; #define HTP_OP_MAX_DIMS 4 // aka GGML_MAX_DIMS -#define HTP_OP_MAX_INPUTS 6 // aka GGML_MAX_SRCS +#define HTP_OP_MAX_INPUTS 10 // aka GGML_MAX_SRCS #define HTP_OP_MAX_OUTPUTS 4 #define HTP_OP_MAX_PARAMS 16 // aka GGML_MAX_OP_PARAMS #define HTP_OP_MAX_KERN_PARAMS 32 @@ -112,13 +107,16 @@ enum htp_op_code { #define HTP_OP_MAX_BUFS 16 #define HTP_OP_MAX_TENSORS 8192 // must stay under 64K (uint16) +#define HTP_FENCE_TIMEOUT (1000000000ULL) + #define HTP_OP_MAX_VMEM_DEFAULT (3355443200u) #define HTP_MMAP_MAX_VMEM (2147483648u) enum htp_tensor_flags { - HTP_TENSOR_COMPUTE = (1U << 0), // Tensor buffer temporal compute data (not weights) - HTP_TENSOR_DIRTY = (1U << 1) // Tensor buffer is dirty and needs to be flushed + HTP_TENSOR_WEIGHT = (1U << 0), // Tensor buffer model weight data (not compute) + HTP_TENSOR_REPACK = (1U << 1), // Tensor is in repacked tiled format + HTP_TENSOR_FENCE = (1U << 2) // Tensor is synchronization fence (explicitly managed) }; // Tensor descriptor @@ -175,6 +173,7 @@ enum htp_trace_event_id { HTP_TRACE_EVT_L2FLUSH = 1, HTP_TRACE_EVT_INIT = 2, HTP_TRACE_EVT_BUFF = 3, + HTP_TRACE_EVT_FENCE = 4, HTP_TRACE_EVT_HVX_COMP = 20, HTP_TRACE_EVT_HVX_A_QUANT = 21, @@ -215,6 +214,7 @@ struct htp_opbatch_req { uint32_t n_ops; // Number of ops uint32_t n_traces; // Number of trace descriptors per thread uint32_t pad; // unused + uint64_t seq; // Sequence number // struct htp_buf_desc bufs[]; -- dspqueue buf 0 // struct htp_tensor tensors[]; -- dspqueue buf 0 // struct htp_op_desc ops[]; -- dspqueue buf 0 @@ -231,6 +231,7 @@ struct htp_opbatch_rsp { uint32_t pad; // align to 8 bytes uint64_t cycles_start; // Start cycle counter uint64_t cycles_stop; // Stop cycle counter + uint64_t seq; // Sequence number // struct htp_prof_desc profs[]; -- dspqueue buf 0 }; diff --git a/ggml/src/ggml-hexagon/htp/htp-tensor.c b/ggml/src/ggml-hexagon/htp/htp-tensor.c index 39436e26dff..ae377c9221f 100644 --- a/ggml/src/ggml-hexagon/htp/htp-tensor.c +++ b/ggml/src/ggml-hexagon/htp/htp-tensor.c @@ -79,7 +79,14 @@ void htp_tensor_dirty_all(struct htp_context * ctx, const struct htp_tensor * co for (uint32_t i = 0; i < n; i++) { const struct htp_tensor * t = tensors[i]; - if (!t) continue; + if (!t || (t->flags & (HTP_TENSOR_WEIGHT | HTP_TENSOR_FENCE))) { + continue; + } + + if (t->size <= HEX_L2_FLUSH_IL_THRESHOLD) { + hex_l2flush((void *) (uintptr_t) t->data, t->size); + continue; + } uint32_t t_start = t->data; uint32_t t_end = t_start + t->size; @@ -242,7 +249,7 @@ void htp_tensor_flush_all(struct htp_context * ctx, const struct htp_tensor * co for (uint32_t i = 0; i < n; i++) { const struct htp_tensor * t = tensors[i]; - if (t && (t->flags & HTP_TENSOR_COMPUTE) && is_tensor_dirty(ctx, t)) { + if (t && !(t->flags & (HTP_TENSOR_WEIGHT | HTP_TENSOR_FENCE)) && is_tensor_dirty(ctx, t)) { dirty_tensors[n_dirty++] = t; total_dirty += t->size; } diff --git a/ggml/src/ggml-hexagon/htp/htp-tensor.h b/ggml/src/ggml-hexagon/htp/htp-tensor.h index 2c3fc54c748..c9cadbae3f2 100644 --- a/ggml/src/ggml-hexagon/htp/htp-tensor.h +++ b/ggml/src/ggml-hexagon/htp/htp-tensor.h @@ -13,6 +13,15 @@ static inline uint32_t * htp_tensor_flags(const struct htp_tensor * t) { return (uint32_t *) &t->flags; } +static inline uint32_t htp_tensor_get_row_size(int type, uint32_t ne00) { + switch (type) { + case HTP_TYPE_F32: return ne00 * 4; + case HTP_TYPE_F16: return ne00 * 2; + case HTP_TYPE_Q8_0: return (ne00 / 32) * 34; + default: return 0; + } +} + struct htp_context; void htp_tensor_flush_all(struct htp_context * ctx, const struct htp_tensor * const * tensors, uint32_t n); void htp_tensor_dirty_all(struct htp_context * ctx, const struct htp_tensor * const * tensors, uint32_t n); diff --git a/ggml/src/ggml-hexagon/htp/hvx-arith.h b/ggml/src/ggml-hexagon/htp/hvx-arith.h index 82e3416970b..c8d0003ab5c 100644 --- a/ggml/src/ggml-hexagon/htp/hvx-arith.h +++ b/ggml/src/ggml-hexagon/htp/hvx-arith.h @@ -17,9 +17,9 @@ #define hvx_arith_loop_body(dst_type, src0_type, src1_type, elem_size, vec_store, vec_op) \ do { \ - dst_type * restrict vdst = (dst_type *) dst; \ - src0_type * restrict vsrc0 = (src0_type *) src0; \ - src1_type * restrict vsrc1 = (src1_type *) src1; \ + dst_type * vdst = (dst_type *) dst; \ + src0_type * vsrc0 = (src0_type *) src0; \ + src1_type * vsrc1 = (src1_type *) src1; \ \ const uint32_t epv = 128 / (elem_size); \ const uint32_t nvec = n / epv; \ @@ -57,40 +57,40 @@ // Generic macro to define alignment permutations for an op #define DEFINE_HVX_BINARY_OP_VARIANTS(OP_NAME, OP_MACRO, ELEM_TYPE) \ -static inline void OP_NAME##_aaa(uint8_t * restrict dst, const uint8_t * restrict src0, const uint8_t * restrict src1, uint32_t n) { \ +static inline void OP_NAME##_aaa(uint8_t * dst, const uint8_t * src0, const uint8_t * src1, uint32_t n) { \ assert((uintptr_t) dst % 128 == 0); \ assert((uintptr_t) src0 % 128 == 0); \ assert((uintptr_t) src1 % 128 == 0); \ hvx_arith_loop_body(HVX_Vector, HVX_Vector, HVX_Vector, sizeof(ELEM_TYPE), hvx_vec_store_a, OP_MACRO); \ } \ -static inline void OP_NAME##_aau(uint8_t * restrict dst, const uint8_t * restrict src0, const uint8_t * restrict src1, uint32_t n) { \ +static inline void OP_NAME##_aau(uint8_t * dst, const uint8_t * src0, const uint8_t * src1, uint32_t n) { \ assert((uintptr_t) dst % 128 == 0); \ assert((uintptr_t) src0 % 128 == 0); \ hvx_arith_loop_body(HVX_Vector, HVX_Vector, HVX_UVector, sizeof(ELEM_TYPE), hvx_vec_store_a, OP_MACRO); \ } \ -static inline void OP_NAME##_aua(uint8_t * restrict dst, const uint8_t * restrict src0, const uint8_t * restrict src1, uint32_t n) { \ +static inline void OP_NAME##_aua(uint8_t * dst, const uint8_t * src0, const uint8_t * src1, uint32_t n) { \ assert((uintptr_t) dst % 128 == 0); \ assert((uintptr_t) src1 % 128 == 0); \ hvx_arith_loop_body(HVX_Vector, HVX_UVector, HVX_Vector, sizeof(ELEM_TYPE), hvx_vec_store_a, OP_MACRO); \ } \ -static inline void OP_NAME##_auu(uint8_t * restrict dst, const uint8_t * restrict src0, const uint8_t * restrict src1, uint32_t n) { \ +static inline void OP_NAME##_auu(uint8_t * dst, const uint8_t * src0, const uint8_t * src1, uint32_t n) { \ assert((uintptr_t) dst % 128 == 0); \ hvx_arith_loop_body(HVX_Vector, HVX_UVector, HVX_UVector, sizeof(ELEM_TYPE), hvx_vec_store_a, OP_MACRO); \ } \ -static inline void OP_NAME##_uaa(uint8_t * restrict dst, const uint8_t * restrict src0, const uint8_t * restrict src1, uint32_t n) { \ +static inline void OP_NAME##_uaa(uint8_t * dst, const uint8_t * src0, const uint8_t * src1, uint32_t n) { \ assert((uintptr_t) src0 % 128 == 0); \ assert((uintptr_t) src1 % 128 == 0); \ hvx_arith_loop_body(HVX_UVector, HVX_Vector, HVX_Vector, sizeof(ELEM_TYPE), hvx_vec_store_u, OP_MACRO); \ } \ -static inline void OP_NAME##_uau(uint8_t * restrict dst, const uint8_t * restrict src0, const uint8_t * restrict src1, uint32_t n) { \ +static inline void OP_NAME##_uau(uint8_t * dst, const uint8_t * src0, const uint8_t * src1, uint32_t n) { \ assert((uintptr_t) src0 % 128 == 0); \ hvx_arith_loop_body(HVX_UVector, HVX_Vector, HVX_UVector, sizeof(ELEM_TYPE), hvx_vec_store_u, OP_MACRO); \ } \ -static inline void OP_NAME##_uua(uint8_t * restrict dst, const uint8_t * restrict src0, const uint8_t * restrict src1, uint32_t n) { \ +static inline void OP_NAME##_uua(uint8_t * dst, const uint8_t * src0, const uint8_t * src1, uint32_t n) { \ assert((uintptr_t) src1 % 128 == 0); \ hvx_arith_loop_body(HVX_UVector, HVX_UVector, HVX_Vector, sizeof(ELEM_TYPE), hvx_vec_store_u, OP_MACRO); \ } \ -static inline void OP_NAME##_uuu(uint8_t * restrict dst, const uint8_t * restrict src0, const uint8_t * restrict src1, uint32_t n) { \ +static inline void OP_NAME##_uuu(uint8_t * dst, const uint8_t * src0, const uint8_t * src1, uint32_t n) { \ hvx_arith_loop_body(HVX_UVector, HVX_UVector, HVX_UVector, sizeof(ELEM_TYPE), hvx_vec_store_u, OP_MACRO); \ } \ diff --git a/ggml/src/ggml-hexagon/htp/hvx-quant.h b/ggml/src/ggml-hexagon/htp/hvx-quant.h new file mode 100644 index 00000000000..6b172cd63c3 --- /dev/null +++ b/ggml/src/ggml-hexagon/htp/hvx-quant.h @@ -0,0 +1,165 @@ +#ifndef HVX_QUANT_H +#define HVX_QUANT_H + +#include +#include +#include + +#include "hvx-arith.h" +#include "hvx-base.h" +#include "hvx-reduce.h" +#include "hvx-repl.h" +#include "hvx-utils.h" + +#ifndef GGML_COMMON_DECL_C +#define GGML_COMMON_DECL_C +#endif +#include "ggml-common.h" +#include "ggml-impl.h" + +static inline void hvx_quantize_row_q8_0_f32(void * restrict dst_ptr, const float * restrict src_ptr, int n) { + const int nb = n / QK8_0; + block_q8_0 * dst = (block_q8_0 *) dst_ptr; + HVX_Vector zero = Q6_V_vzero(); + + int i = 0; + for (; i + 3 < nb; i += 4) { + HVX_Vector * vx = (HVX_Vector *) (src_ptr + i * QK8_0); + + HVX_Vector vmax0_sf = hvx_vec_reduce_max_f32(hvx_vec_abs_f32(vx[0])); + HVX_Vector vmax1_sf = hvx_vec_reduce_max_f32(hvx_vec_abs_f32(vx[1])); + HVX_Vector vmax2_sf = hvx_vec_reduce_max_f32(hvx_vec_abs_f32(vx[2])); + HVX_Vector vmax3_sf = hvx_vec_reduce_max_f32(hvx_vec_abs_f32(vx[3])); + + HVX_Vector vx0_qf = Q6_Vqf32_vsub_VsfVsf(vx[0], zero); + HVX_Vector vx1_qf = Q6_Vqf32_vsub_VsfVsf(vx[1], zero); + HVX_Vector vx2_qf = Q6_Vqf32_vsub_VsfVsf(vx[2], zero); + HVX_Vector vx3_qf = Q6_Vqf32_vsub_VsfVsf(vx[3], zero); + + HVX_Vector vmax0_qf = Q6_Vqf32_vsub_VsfVsf(vmax0_sf, zero); + HVX_Vector vmax1_qf = Q6_Vqf32_vsub_VsfVsf(vmax1_sf, zero); + HVX_Vector vmax2_qf = Q6_Vqf32_vsub_VsfVsf(vmax2_sf, zero); + HVX_Vector vmax3_qf = Q6_Vqf32_vsub_VsfVsf(vmax3_sf, zero); + + HVX_Vector vmax01_hf = Q6_Vh_vdeal_Vh(Q6_Vhf_equals_Wqf32(Q6_W_vcombine_VV(vmax1_qf, vmax0_qf))); + HVX_Vector vmax23_hf = Q6_Vh_vdeal_Vh(Q6_Vhf_equals_Wqf32(Q6_W_vcombine_VV(vmax3_qf, vmax2_qf))); + + HVX_Vector vx01_hf = Q6_Vh_vdeal_Vh(Q6_Vhf_equals_Wqf32(Q6_W_vcombine_VV(vx1_qf, vx0_qf))); + HVX_Vector vx23_hf = Q6_Vh_vdeal_Vh(Q6_Vhf_equals_Wqf32(Q6_W_vcombine_VV(vx3_qf, vx2_qf))); + + HVX_Vector vd01_qf16 = Q6_Vqf16_vmpy_VhfVhf(vmax01_hf, Q6_Vh_vsplat_R(0x2008)); // 1.0 / 127.0 + HVX_Vector vd23_qf16 = Q6_Vqf16_vmpy_VhfVhf(vmax23_hf, Q6_Vh_vsplat_R(0x2008)); // 1.0 / 127.0 + HVX_Vector vd01_hf = Q6_Vhf_equals_Vqf16(vd01_qf16); + HVX_Vector vd23_hf = Q6_Vhf_equals_Vqf16(vd23_qf16); + + HVX_Vector vd01_inv_hf = hvx_vec_inverse_f16(vd01_hf); + HVX_Vector vd23_inv_hf = hvx_vec_inverse_f16(vd23_hf); + vx01_hf = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vmpy_VhfVhf(vx01_hf, vd01_inv_hf)); + vx23_hf = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vmpy_VhfVhf(vx23_hf, vd23_inv_hf)); + + HVX_Vector vx01_i16 = hvx_vec_i16_from_hf_rnd_sat(vx01_hf); + HVX_Vector vx23_i16 = hvx_vec_i16_from_hf_rnd_sat(vx23_hf); + HVX_Vector vx_i8 = Q6_Vb_vpack_VhVh_sat(vx23_i16, vx01_i16); + + hvx_vec_store_u(&dst[i + 0].d, 2, vd01_hf); + hvx_vec_store_u(dst[i + 0].qs, 32, vx_i8); + + hvx_vec_store_u(&dst[i + 1].d, 2, Q6_V_vror_VR(vd01_hf, 64)); + hvx_vec_store_u(dst[i + 1].qs, 32, Q6_V_vror_VR(vx_i8, 32)); + + hvx_vec_store_u(&dst[i + 2].d, 2, vd23_hf); + hvx_vec_store_u(dst[i + 2].qs, 32, Q6_V_vror_VR(vx_i8, 64)); + + hvx_vec_store_u(&dst[i + 3].d, 2, Q6_V_vror_VR(vd23_hf, 64)); + hvx_vec_store_u(dst[i + 3].qs, 32, Q6_V_vror_VR(vx_i8, 96)); + } + + for (; i < nb; i++) { + const float * block_src = src_ptr + i * QK8_0; + HVX_Vector vx = *(const HVX_UVector *) block_src; + HVX_Vector v_abs = hvx_vec_abs_f32(vx); + HVX_Vector v_max = hvx_vec_reduce_max_f32(v_abs); + float amax = hvx_vec_get_f32(v_max); + + const float d = amax / 127.0f; + const float id = d ? (1.0f / d) : 0.0f; + dst[i].d = GGML_FP32_TO_FP16(d); + + HVX_Vector vid = hvx_vec_splat_f32(id); + HVX_Vector v_scaled = hvx_vec_mul_f32_f32(vx, vid); + HVX_Vector v_scaled_qf = Q6_Vqf32_vsub_VsfVsf(v_scaled, zero); + HVX_Vector v_scaled_hf = Q6_Vh_vdeal_Vh(Q6_Vhf_equals_Wqf32(Q6_W_vcombine_VV(zero, v_scaled_qf))); + HVX_Vector v_i16 = hvx_vec_i16_from_hf_rnd_sat(v_scaled_hf); + HVX_Vector v_i8 = Q6_Vb_vpack_VhVh_sat(zero, v_i16); + + hvx_vec_store_u(dst[i].qs, 32, v_i8); + } +} + +static inline void hvx_dequantize_row_q8_0_f32(float * restrict dst_ptr, const void * restrict src_ptr, int n) { + const int nb = n / QK8_0; + const block_q8_0 * src = (const block_q8_0 *) src_ptr; + + for (int i = 0; i < nb; i++) { + HVX_Vector vd_f16 = Q6_Vh_vsplat_R(*(const int16_t *) &src[i].d); + HVX_VectorPair vp_f32 = hvx_vec_f16_to_f32(vd_f16); + HVX_Vector vd = Q6_V_lo_W(vp_f32); + + HVX_Vector vq_i8 = *(const HVX_UVector *) src[i].qs; + + HVX_VectorPair p16 = Q6_Wh_vunpack_Vb(vq_i8); + HVX_Vector v_i16 = Q6_V_lo_W(p16); + HVX_VectorPair p32 = Q6_Ww_vunpack_Vh(v_i16); + HVX_Vector v_i32 = Q6_V_lo_W(p32); + + HVX_Vector v_f32 = Q6_Vsf_equals_Vw(v_i32); + HVX_Vector res = hvx_vec_mul_f32_f32(v_f32, vd); + + float * block_dst = dst_ptr + i * QK8_0; + hvx_vmem(block_dst) = res; + } +} + +static inline void hvx_dequantize_row_q8_0_f16(__fp16 * restrict dst_ptr, const void * restrict src_ptr, int n) { + const int nb = n / QK8_0; + const block_q8_0 * src = (const block_q8_0 *) src_ptr; + + for (int i = nb - 1; i >= 0; i--) { + HVX_Vector vd_f16 = Q6_Vh_vsplat_R(*(const int16_t *) &src[i].d); + HVX_VectorPair vp_f32 = hvx_vec_f16_to_f32(vd_f16); + HVX_Vector vd = Q6_V_lo_W(vp_f32); + + HVX_Vector vq_i8 = *(const HVX_UVector *) src[i].qs; + + HVX_VectorPair p16 = Q6_Wh_vunpack_Vb(vq_i8); + HVX_Vector v_i16 = Q6_V_lo_W(p16); + HVX_VectorPair p32 = Q6_Ww_vunpack_Vh(v_i16); + HVX_Vector v_i32 = Q6_V_lo_W(p32); + + HVX_Vector v_f32 = Q6_Vsf_equals_Vw(v_i32); + HVX_Vector res_f32 = hvx_vec_mul_f32_f32(v_f32, vd); + + HVX_Vector res_f16 = hvx_vec_f32_to_f16(res_f32, Q6_V_vzero()); + + __fp16 * block_dst = dst_ptr + i * QK8_0; + hvx_vec_store_u(block_dst, QK8_0 * sizeof(__fp16), res_f16); + } +} + +static inline void hvx_dequantize_row_f16_f32(float * restrict dst_ptr, const void * restrict src_ptr, int n) { + const int nb = n / 32; + const _Float16 * src = (const _Float16 *) src_ptr; + + for (int i = 0; i < nb; i++) { + HVX_Vector v_f16 = *(const HVX_UVector *) (src + i * 32); + HVX_VectorPair vp_f32 = hvx_vec_f16_to_f32(v_f16); + HVX_Vector res = Q6_V_lo_W(vp_f32); + + float * block_dst = dst_ptr + i * 32; + hvx_vmem(block_dst) = res; + } +} + + + +#endif // HVX_QUANT_H diff --git a/ggml/src/ggml-hexagon/htp/main.c b/ggml/src/ggml-hexagon/htp/main.c index 880e20c9959..975ba0c7af5 100644 --- a/ggml/src/ggml-hexagon/htp/main.c +++ b/ggml/src/ggml-hexagon/htp/main.c @@ -18,6 +18,7 @@ #include #include #include +#include #include "hex-utils.h" #include "hex-dma.h" @@ -32,6 +33,7 @@ #include "htp_iface.h" #include "work-queue.h" #include "hex-profile.h" +#include "allreduce-ops.h" #define HMX_QUEUE_CAPACITY 16 #define HMX_QUEUE_STACK_SIZE 16384 @@ -46,6 +48,36 @@ struct htp_handle { struct htp_context * ctx; }; +static inline void * htp_mmap(uint32_t fd, uint32_t size) { + void * va = (void *)-1; + for (int retry = 0; retry < 2; retry++) { +#if __HVX_ARCH__ > 73 + va = HAP_mmap2(NULL, size, HAP_PROT_READ | HAP_PROT_WRITE, 0, fd, 0); +#else + if (size > HTP_MMAP_MAX_VMEM) { + FARF(ERROR, "mmap failed : size %u exceeds 2GB limit for HAP_mmap", (uint32_t) size); + abort(); + } + va = HAP_mmap(NULL, size, HAP_PROT_READ | HAP_PROT_WRITE, 0, fd, 0); +#endif + if (va != (void *)-1 && va != NULL) { + return va; + } + if (retry == 0) { + FARF(HIGH, "mmap failed first try (va %p fd %u size %u), retrying...", va, fd, size); + } + } + return NULL; +} + +static inline void htp_munmap(void * va, uint32_t size) { +#if __HVX_ARCH__ > 73 + HAP_munmap2(va, size); +#else + HAP_munmap(va, size); +#endif +} + AEEResult htp_iface_open(const char * uri, remote_handle64 * handle) { (void) uri; struct htp_handle * h = calloc(1, sizeof(*h)); @@ -127,11 +159,7 @@ AEEResult htp_iface_close(remote_handle64 handle) { // release the mmaps (if any) for (uint32_t i=0; immap[i].size) { -#if __HVX_ARCH__ > 73 - HAP_munmap2((void *) ctx->mmap[i].base, ctx->mmap[i].size); -#else - HAP_munmap((void *) ctx->mmap[i].base, ctx->mmap[i].size); -#endif + htp_munmap((void *) ctx->mmap[i].base, ctx->mmap[i].size); ctx->mmap[i].size = 0; ctx->mmap[i].base = NULL; ctx->mmap[i].fd = -1; @@ -175,18 +203,9 @@ AEEResult htp_iface_mmap(remote_handle64 handle, uint32_t fd, uint32_t size) { struct htp_mmap *m = &ctx->mmap[i]; if (!m->size) { FARF(HIGH, "mmap : fd %u size %u", fd, size); -#if __HVX_ARCH__ > 73 - void *va = HAP_mmap2(NULL, size, HAP_PROT_READ | HAP_PROT_WRITE, 0, fd, 0); -#else - if (size > HTP_MMAP_MAX_VMEM) { // HAP_mmap has a size limit of 2GB - FARF(ERROR, "mmap failed : size %u exceeds 2GB limit for HAP_mmap", (uint32_t) size); - abort(); // can't do much else at this point - } - - void *va = HAP_mmap(NULL, size, HAP_PROT_READ | HAP_PROT_WRITE, 0, fd, 0); -#endif - if (va == (void*)-1) { - FARF(ERROR, "mmap failed : va %p fd %u size %u", va, fd, (uint32_t) size); + void *va = htp_mmap(fd, size); + if (va == NULL) { + FARF(ERROR, "mmap failed : fd %u size %u", fd, (uint32_t) size); return AEE_EFAILED; } @@ -212,11 +231,7 @@ AEEResult htp_iface_munmap(remote_handle64 handle, uint32 fd) { struct htp_mmap *m = &ctx->mmap[i]; if (fd < 0 || m->fd == fd) { FARF(HIGH, "unmmap : base %p fd %u size %u", (void*) m->base, m->fd, (uint32_t) m->size); -#if __HVX_ARCH__ > 73 - HAP_munmap2((void *) m->base, m->size); -#else - HAP_munmap((void *) m->base, m->size); -#endif + htp_munmap((void *) m->base, m->size); m->size = 0; m->base = NULL; m->fd = -1; @@ -228,7 +243,7 @@ AEEResult htp_iface_munmap(remote_handle64 handle, uint32 fd) { static void vtcm_acquire(struct htp_context * ctx) { if (!ctx->vtcm_valid) { - int err = HAP_compute_res_acquire_cached(ctx->vtcm_rctx, 1000000u); + int err = HAP_compute_res_acquire_cached(ctx->vtcm_rctx, 10000000u); if (err != 0) { FARF(ERROR, "ggml-hex: failed to acquire VTCM: 0x%08x", (unsigned)err); abort(); @@ -692,8 +707,45 @@ static inline void profile_stop(uint32_t mode, struct profile_data * d) { } } +static int op_fence(struct htp_ops_context * octx) { + struct htp_context *ctx = octx->ctx; + struct htp_thread_trace * tr = &ctx->trace[0]; + const uint32_t seq = (uint32_t) octx->op_params[0]; + + htp_trace_event_start(tr, HTP_TRACE_EVT_FENCE, (uint16_t) seq); + + const struct htp_tensor * sync = octx->src[0]; + atomic_uint * sync_fence = (atomic_uint *) sync->data; + uint64_t spins = 0; + while (1) { + Q6_dccleaninva_A((void *) sync_fence); + asm volatile ("syncht" : : : "memory"); + uint32_t val = atomic_load(&sync_fence[0]); + if ((int32_t)(val - seq) >= 0) { + break; + } + if (++spins > HTP_FENCE_TIMEOUT) { + FARF(ERROR, "ggml-hex: sync-wait TIMEOUT : fence %p spins %llu seq %u\n", sync_fence, spins, seq); + break; + } + hex_pause(); + } + + htp_trace_event_stop(tr, HTP_TRACE_EVT_FENCE, (uint16_t) seq); + + FARF(HIGH, "ggml-hex: sync-done : fence %p spins %llu seq %u\n", sync_fence, spins, seq); + return HTP_STATUS_OK; +} + static int execute_op(struct htp_ops_context * octx) { switch (octx->op) { + case HTP_OP_FENCE: + return op_fence(octx); + + case HTP_OP_ALLREDUCE: + case HTP_OP_ALLREDUCE_ADD: + return op_allreduce(octx); + case HTP_OP_MUL_MAT: case HTP_OP_MUL_MAT_ADD: return op_matmul(octx); @@ -701,11 +753,8 @@ static int execute_op(struct htp_ops_context * octx) { case HTP_OP_MUL_MAT_ID: return op_matmul_id(octx); - case HTP_OP_MUL_MAT_QKV: - return op_matmul_qkv(octx); - - case HTP_OP_MUL_MAT_FFN: - return op_matmul_ffn(octx); + case HTP_OP_MUL_MAT_NX: + return op_matmul_nx(octx); case HTP_OP_MUL: case HTP_OP_ADD: @@ -818,12 +867,8 @@ static inline bool reuse_buf(struct htp_context *ctx, uint32_t *m_reuse, struct static inline void drop_mmap(struct htp_context *ctx, struct htp_mmap *m) { if (m->size) { - FARF(HIGH, "unmap : fd %u base %p size %u", m->fd, (void*) m->base, (uint32_t) m->size); -#if __HVX_ARCH__ > 73 - HAP_munmap2((void *) m->base, m->size); -#else - HAP_munmap((void *) m->base, m->size); -#endif + FARF(ALWAYS, "unmap : fd %u base %p size %u", m->fd, (void*) m->base, (uint32_t) m->size); + htp_munmap((void *) m->base, m->size); m->size = 0; m->base = 0; m->fd = -1; @@ -837,18 +882,9 @@ static inline void mmap_buf(struct htp_context *ctx, struct htp_buf_desc *b) { for (uint32_t i=0; i < HTP_MAX_MMAPS; i++) { struct htp_mmap *m = &ctx->mmap[i]; if (!m->size) { -#if __HVX_ARCH__ > 73 - void *va = HAP_mmap2(NULL, b->size, HAP_PROT_READ | HAP_PROT_WRITE, 0, b->fd, 0); -#else - if (b->size > HTP_MMAP_MAX_VMEM) { // HAP_mmap has a size limit of 2GB - FARF(ERROR, "mmap failed : size %u exceeds 2GB limit for HAP_mmap", (uint32_t) b->size); - abort(); // can't do much else at this point - } - - void *va = HAP_mmap(NULL, b->size, HAP_PROT_READ | HAP_PROT_WRITE, 0, b->fd, 0); -#endif - if (va == (void*)-1) { - FARF(ERROR, "mmap failed : va %p fd %u size %u", va, b->fd, (uint32_t) b->size); + void *va = htp_mmap(b->fd, b->size); + if (va == NULL) { + FARF(ERROR, "mmap failed : fd %u size %u", b->fd, (uint32_t) b->size); abort(); // can't do much else at this point } @@ -856,10 +892,13 @@ static inline void mmap_buf(struct htp_context *ctx, struct htp_buf_desc *b) { m->fd = b->fd; m->size = b->size; - FARF(HIGH, "mmap : fd %u base %p size %u", m->fd, (void*) m->base, (uint32_t) m->size); + FARF(ALWAYS, "mmap : fd %u base %p size %u", m->fd, (void*) m->base, (uint32_t) m->size); return; } } + + FARF(ERROR, "mmap failed : exceeded mapping capacity limit of %u", HTP_MAX_MMAPS); + abort(); } static void prep_op_bufs(struct htp_context *ctx, struct htp_buf_desc *bufs, uint32_t n_bufs) { @@ -1081,6 +1120,7 @@ static void process_opbatch(struct htp_context * ctx, const struct htp_opbatch_r rsp.usecs = batch_prof.usecs; rsp.cycles_start = batch_prof.cycles_start; rsp.cycles_stop = batch_prof.cycles_stop; + rsp.seq = req->seq; if (ctx->profiler == HTP_PROF_TRACE) { for (int t = 0; t <= HTP_MAX_NTHREADS; t++) { diff --git a/ggml/src/ggml-hexagon/htp/matmul-ops.c b/ggml/src/ggml-hexagon/htp/matmul-ops.c index 9d385469ae9..a6adc0e61fa 100644 --- a/ggml/src/ggml-hexagon/htp/matmul-ops.c +++ b/ggml/src/ggml-hexagon/htp/matmul-ops.c @@ -64,6 +64,7 @@ typedef struct { struct htp_mm_context { const char * type; struct htp_ops_context * octx; + const struct htp_tensor * act; void (*vec_dot_1x1)(const uint32_t n, float * restrict s0, const void * restrict vx0, @@ -478,7 +479,7 @@ static void hvx_mv_2d_repacked_##SUFFIX(unsigned int nth, unsigned int ith, void \ htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, ct); \ DOT_2X1(ne10, dst_ptr, w_tile, src1_col, valid_rows, NULL); \ - htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, ct); \ + htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, ct); \ \ if (push_ct < ct_end) { \ dma_queue_push(dma_queue, dma_make_ptr((uint8_t *)w_tile, src0_row + push_ct * tile_row_stride), \ @@ -502,150 +503,67 @@ static void hvx_mv_2d_repacked_##SUFFIX(unsigned int nth, unsigned int ith, void } \ } -#define MATMUL_QKV_2D_REPACKED_IMPL(SUFFIX, TILE_SIZE, DOT_2X2, DOT_2X1) \ -static void hvx_mm_qkv_2d_repacked_##SUFFIX(unsigned int nth, unsigned int ith, void * data) { \ +#define MATMUL_NX_2D_REPACKED_IMPL(SUFFIX, TILE_SIZE, DOT_2X2, DOT_2X1) \ +static void hvx_mm_nx_2d_repacked_##SUFFIX(unsigned int nth, unsigned int ith, void * data) { \ struct htp_mm_context * mmctx = data; \ struct htp_ops_context * octx = mmctx->octx; \ + const struct htp_mm_kernel_params * kparams = (const struct htp_mm_kernel_params *) octx->kernel_params; \ + const uint32_t n_weights = kparams->n_weights; \ \ - const struct htp_tensor * restrict src0 = octx->src[0]; /* Wk */ \ - const struct htp_tensor * restrict src1 = octx->src[1]; /* x */ \ - const struct htp_tensor * restrict src2 = octx->src[2]; /* Wv */ \ - const struct htp_tensor * restrict src3 = octx->src[3]; /* Wq */ \ - const struct htp_tensor * restrict dst_k = octx->dsts[0]; \ - const struct htp_tensor * restrict dst_v = octx->dsts[1]; \ - const struct htp_tensor * restrict dst_q = octx->dsts[2]; \ - \ - const uint32_t ne00 = src0->ne[0]; \ - const uint32_t ne10 = src1->ne[0]; \ - const uint32_t src1_nrows = src1->ne[1] * src1->ne[2] * src1->ne[3]; \ - \ - const size_t dst_k_row_size = dst_k->nb[1]; /* K and V share output width */ \ - const size_t dst_q_row_size = dst_q->nb[1]; /* Q may be wider (GQA) */ \ + const struct htp_tensor * restrict act = octx->src[n_weights]; /* x */ \ + const uint32_t ne10 = act->ne[0]; \ + const uint32_t src1_nrows = act->ne[1] * act->ne[2] * act->ne[3]; \ const size_t src1_stride = mmctx->vtcm_src1_stride; \ \ - uint8_t * restrict vtcm_src0_ptr = mmctx->vtcm_src0 + mmctx->vtcm_src0_size_per_thread * ith; \ - uint8_t * restrict vtcm_src2_ptr = mmctx->vtcm_src2 + mmctx->vtcm_src2_size_per_thread * ith; \ - uint8_t * restrict vtcm_src3_ptr = mmctx->vtcm_src3 + mmctx->vtcm_src3_size_per_thread * ith; \ - uint8_t * restrict src1_data = mmctx->vtcm_src1; \ + uint8_t * restrict vtcm_weight_ptr = mmctx->vtcm_src0 + mmctx->vtcm_src0_size_per_thread * ith; \ + uint8_t * restrict src1_data = mmctx->vtcm_src1; \ \ struct htp_thread_trace * tr = &octx->ctx->trace[ith]; \ - \ - const struct htp_mm_kernel_params * kparams = (const struct htp_mm_kernel_params *) octx->kernel_params; \ const uint32_t n_prefetch = kparams->n_prefetch; \ assert(n_prefetch >= 2 && n_prefetch <= HTP_MM_MAX_PREFETCH && (n_prefetch & (n_prefetch - 1)) == 0); \ \ - const uint8_t * restrict src0_row = (const uint8_t *) src0->data; \ - const uint8_t * restrict src2_row = (const uint8_t *) src2->data; \ - const uint8_t * restrict src3_row = (const uint8_t *) src3->data; \ - \ const uint32_t tile_size = TILE_SIZE; \ const uint32_t aligned_tile_size = hex_align_up(tile_size, 128); \ - \ - uint32_t n_k_tiles_w = ne00 / 32; \ uint32_t n_k_tiles_a = ne10 / 32; \ - uint32_t tile_row_stride = n_k_tiles_w * tile_size; \ uint32_t tile_row_transfer_size_aligned = n_k_tiles_a * aligned_tile_size; \ \ dma_queue * dma_queue = octx->ctx->dma[ith]; \ \ - /* 1. Process K and V together */ \ - const uint32_t src0_nrows_kv = src0->ne[1] * src0->ne[2] * src0->ne[3]; /* src0 is Wk */ \ - uint32_t src0_nrows_per_thread_kv = (src0_nrows_kv + nth - 1) / nth; \ - src0_nrows_per_thread_kv = hex_round_up(src0_nrows_per_thread_kv, 32); \ - \ - const uint32_t start_row_kv = src0_nrows_per_thread_kv * ith; \ - const uint32_t end_row_kv = MIN(start_row_kv + src0_nrows_per_thread_kv, src0_nrows_kv); \ - \ - uint32_t ct_start_kv = start_row_kv / 32; \ - uint32_t ct_end_kv = (end_row_kv + 31) / 32; \ - \ - uint32_t push_ct = ct_start_kv; \ - if (start_row_kv < end_row_kv) { \ - for (uint32_t d = 0; d < n_prefetch && push_ct < ct_end_kv; d++, push_ct++) { \ - dma_queue_push(dma_queue, dma_make_ptr(vtcm_src0_ptr + d * tile_row_transfer_size_aligned, \ - src0_row + push_ct * tile_row_stride), aligned_tile_size, tile_size, tile_size, n_k_tiles_a); \ - dma_queue_push(dma_queue, dma_make_ptr(vtcm_src2_ptr + d * tile_row_transfer_size_aligned, \ - src2_row + push_ct * tile_row_stride), aligned_tile_size, tile_size, tile_size, n_k_tiles_a); \ - } \ - } \ - \ hvx_mm_run_quant_task(mmctx, ith); \ \ - if (start_row_kv < end_row_kv) { \ - \ - for (uint32_t ct = ct_start_kv; ct < ct_end_kv; ct++) { \ - const uint8_t * w_tile_k = dma_queue_pop(dma_queue).dst; \ - const uint8_t * w_tile_v = dma_queue_pop(dma_queue).dst; \ - \ - int valid_rows = (int)src0->ne[1] - (int)(ct * 32); \ - valid_rows = MIN(32, MAX(0, valid_rows)); \ - \ - htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, ith); \ - uint32_t ir1 = 0; \ - for (; ir1 + 1 < src1_nrows; ir1 += 2) { \ - const uint8_t * restrict src1_col0 = (const uint8_t *) (src1_data + (ir1+0) * src1_stride); \ - const uint8_t * restrict src1_col1 = (const uint8_t *) (src1_data + (ir1+1) * src1_stride); \ - \ - float * restrict dst_row0_k = (float *) (dst_k->data + ((ir1+0) * dst_k_row_size)); \ - float * restrict dst_row1_k = (float *) (dst_k->data + ((ir1+1) * dst_k_row_size)); \ - float * dst_ptr0_k = &dst_row0_k[ct * 32]; \ - float * dst_ptr1_k = &dst_row1_k[ct * 32]; \ - \ - float * restrict dst_row0_v = (float *) (dst_v->data + ((ir1+0) * dst_k_row_size)); \ - float * restrict dst_row1_v = (float *) (dst_v->data + ((ir1+1) * dst_k_row_size)); \ - float * dst_ptr0_v = &dst_row0_v[ct * 32]; \ - float * dst_ptr1_v = &dst_row1_v[ct * 32]; \ - \ - DOT_2X2(ne10, dst_ptr0_k, dst_ptr1_k, w_tile_k, src1_col0, src1_col1, valid_rows, NULL, NULL); \ - DOT_2X2(ne10, dst_ptr0_v, dst_ptr1_v, w_tile_v, src1_col0, src1_col1, valid_rows, NULL, NULL); \ - } \ + for (uint32_t widx = 0; widx < n_weights; widx++) { \ + const struct htp_tensor * restrict src_w = octx->src[widx]; \ + const struct htp_tensor * restrict dst = octx->dsts[widx]; \ + if (!src_w || !dst) continue; \ \ - for (; ir1 < src1_nrows; ++ir1) { \ - const uint8_t * restrict src1_col = (const uint8_t *) (src1_data + ir1 * src1_stride); \ - \ - float * restrict dst_row_k = (float *) (dst_k->data + (ir1 * dst_k_row_size)); \ - float * dst_ptr_k = &dst_row_k[ct * 32]; \ - \ - float * restrict dst_row_v = (float *) (dst_v->data + (ir1 * dst_k_row_size)); \ - float * dst_ptr_v = &dst_row_v[ct * 32]; \ - \ - DOT_2X1(ne10, dst_ptr_k, w_tile_k, src1_col, valid_rows, NULL); \ - DOT_2X1(ne10, dst_ptr_v, w_tile_v, src1_col, valid_rows, NULL); \ - } \ - htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, ith); \ + const uint32_t ne00 = src_w->ne[0]; \ + const uint32_t ne01 = src_w->ne[1]; \ + const size_t dst_row_size = dst->nb[1]; \ + const uint8_t * restrict src_w_row = (const uint8_t *) src_w->data; \ \ - if (push_ct < ct_end_kv) { \ - dma_queue_push(dma_queue, dma_make_ptr((uint8_t *)w_tile_k, src0_row + push_ct * tile_row_stride), \ - aligned_tile_size, tile_size, tile_size, n_k_tiles_a); \ - dma_queue_push(dma_queue, dma_make_ptr((uint8_t *)w_tile_v, src2_row + push_ct * tile_row_stride), \ - aligned_tile_size, tile_size, tile_size, n_k_tiles_a); \ - push_ct++; \ - } \ - } \ - } \ + uint32_t n_k_tiles_w = ne00 / 32; \ + uint32_t tile_row_stride = n_k_tiles_w * tile_size; \ \ - /* 2. Process Q separately */ \ - const uint32_t src0_nrows_q = src3->ne[1] * src3->ne[2] * src3->ne[3]; /* src3 is Wq */ \ - uint32_t src0_nrows_per_thread_q = (src0_nrows_q + nth - 1) / nth; \ - src0_nrows_per_thread_q = hex_round_up(src0_nrows_per_thread_q, 32); \ + const uint32_t src0_nrows = ne01 * src_w->ne[2] * src_w->ne[3]; \ + uint32_t src0_nrows_per_thread = (src0_nrows + nth - 1) / nth; \ + src0_nrows_per_thread = hex_round_up(src0_nrows_per_thread, 32); \ \ - const uint32_t start_row_q = src0_nrows_per_thread_q * ith; \ - const uint32_t end_row_q = MIN(start_row_q + src0_nrows_per_thread_q, src0_nrows_q); \ + const uint32_t start_row = src0_nrows_per_thread * ith; \ + const uint32_t end_row = MIN(start_row + src0_nrows_per_thread, src0_nrows); \ + if (start_row >= end_row) continue; \ \ - if (start_row_q < end_row_q) { \ - uint32_t ct_start_q = start_row_q / 32; \ - uint32_t ct_end_q = (end_row_q + 31) / 32; \ + uint32_t ct_start = start_row / 32; \ + uint32_t ct_end = (end_row + 31) / 32; \ \ - uint32_t push_ct = ct_start_q; \ - for (uint32_t d = 0; d < n_prefetch && push_ct < ct_end_q; d++, push_ct++) { \ - dma_queue_push(dma_queue, dma_make_ptr(vtcm_src3_ptr + d * tile_row_transfer_size_aligned, \ - src3_row + push_ct * tile_row_stride), aligned_tile_size, tile_size, tile_size, n_k_tiles_a); \ + uint32_t push_ct = ct_start; \ + for (uint32_t d = 0; d < n_prefetch && push_ct < ct_end; d++, push_ct++) { \ + dma_queue_push(dma_queue, dma_make_ptr(vtcm_weight_ptr + d * tile_row_transfer_size_aligned, \ + src_w_row + push_ct * tile_row_stride), aligned_tile_size, tile_size, tile_size, n_k_tiles_a); \ } \ \ - for (uint32_t ct = ct_start_q; ct < ct_end_q; ct++) { \ - const uint8_t * w_tile_q = dma_queue_pop(dma_queue).dst; \ - \ - int valid_rows = (int)src3->ne[1] - (int)(ct * 32); \ + for (uint32_t ct = ct_start; ct < ct_end; ct++) { \ + const uint8_t * w_tile = dma_queue_pop(dma_queue).dst; \ + int valid_rows = (int)ne01 - (int)(ct * 32); \ valid_rows = MIN(32, MAX(0, valid_rows)); \ \ htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, ct); \ @@ -654,26 +572,24 @@ static void hvx_mm_qkv_2d_repacked_##SUFFIX(unsigned int nth, unsigned int ith, const uint8_t * restrict src1_col0 = (const uint8_t *) (src1_data + (ir1+0) * src1_stride); \ const uint8_t * restrict src1_col1 = (const uint8_t *) (src1_data + (ir1+1) * src1_stride); \ \ - float * restrict dst_row0_q = (float *) (dst_q->data + ((ir1+0) * dst_q_row_size)); \ - float * restrict dst_row1_q = (float *) (dst_q->data + ((ir1+1) * dst_q_row_size)); \ - float * dst_ptr0_q = &dst_row0_q[ct * 32]; \ - float * dst_ptr1_q = &dst_row1_q[ct * 32]; \ + float * restrict dst_row0 = (float *) (dst->data + ((ir1+0) * dst_row_size)); \ + float * restrict dst_row1 = (float *) (dst->data + ((ir1+1) * dst_row_size)); \ + float * dst_ptr0 = &dst_row0[ct * 32]; \ + float * dst_ptr1 = &dst_row1[ct * 32]; \ \ - DOT_2X2(ne10, dst_ptr0_q, dst_ptr1_q, w_tile_q, src1_col0, src1_col1, valid_rows, NULL, NULL); \ + DOT_2X2(ne10, dst_ptr0, dst_ptr1, w_tile, src1_col0, src1_col1, valid_rows, NULL, NULL); \ } \ \ for (; ir1 < src1_nrows; ++ir1) { \ const uint8_t * restrict src1_col = (const uint8_t *) (src1_data + ir1 * src1_stride); \ - \ - float * restrict dst_row_q = (float *) (dst_q->data + (ir1 * dst_q_row_size)); \ - float * dst_ptr_q = &dst_row_q[ct * 32]; \ - \ - DOT_2X1(ne10, dst_ptr_q, w_tile_q, src1_col, valid_rows, NULL); \ + float * restrict dst_row = (float *) (dst->data + (ir1 * dst_row_size)); \ + float * dst_ptr = &dst_row[ct * 32]; \ + DOT_2X1(ne10, dst_ptr, w_tile, src1_col, valid_rows, NULL); \ } \ htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, ct); \ \ - if (push_ct < ct_end_q) { \ - dma_queue_push(dma_queue, dma_make_ptr((uint8_t *)w_tile_q, src3_row + push_ct * tile_row_stride), \ + if (push_ct < ct_end) { \ + dma_queue_push(dma_queue, dma_make_ptr((uint8_t *)w_tile, src_w_row + push_ct * tile_row_stride), \ aligned_tile_size, tile_size, tile_size, n_k_tiles_a); \ push_ct++; \ } \ @@ -681,121 +597,6 @@ static void hvx_mm_qkv_2d_repacked_##SUFFIX(unsigned int nth, unsigned int ith, } \ } -#define MATMUL_FFN_2D_REPACKED_IMPL(SUFFIX, TILE_SIZE, DOT_2X2, DOT_2X1) \ -static void hvx_mm_ffn_2d_repacked_##SUFFIX(unsigned int nth, unsigned int ith, void * data) { \ - struct htp_mm_context * mmctx = data; \ - struct htp_ops_context * octx = mmctx->octx; \ - \ - const struct htp_tensor * restrict src0 = octx->src[0]; /* Wgate */ \ - const struct htp_tensor * restrict src1 = octx->src[1]; /* y */ \ - const struct htp_tensor * restrict src2 = octx->src[2]; /* Wup */ \ - const struct htp_tensor * restrict dst_gate = octx->dsts[0]; \ - const struct htp_tensor * restrict dst_up = octx->dsts[1]; \ - \ - const uint32_t ne00 = src0->ne[0]; \ - const uint32_t ne01 = src0->ne[1]; \ - const uint32_t ne10 = src1->ne[0]; \ - const uint32_t src1_nrows = src1->ne[1] * src1->ne[2] * src1->ne[3]; \ - \ - const size_t dst_row_size = dst_gate->nb[1]; \ - const size_t src1_stride = mmctx->vtcm_src1_stride; \ - \ - uint8_t * restrict vtcm_src0_ptr = mmctx->vtcm_src0 + mmctx->vtcm_src0_size_per_thread * ith; \ - uint8_t * restrict vtcm_src2_ptr = mmctx->vtcm_src2 + mmctx->vtcm_src2_size_per_thread * ith; \ - uint8_t * restrict src1_data = mmctx->vtcm_src1; \ - \ - struct htp_thread_trace * tr = &octx->ctx->trace[ith]; \ - \ - const uint8_t * restrict src0_row = (const uint8_t *) src0->data; \ - const uint8_t * restrict src2_row = (const uint8_t *) src2->data; \ - \ - const uint32_t tile_size = TILE_SIZE; \ - const uint32_t aligned_tile_size = hex_align_up(tile_size, 128); \ - \ - const struct htp_mm_kernel_params * kparams = (const struct htp_mm_kernel_params *) octx->kernel_params; \ - const uint32_t n_prefetch = kparams->n_prefetch; \ - assert(n_prefetch >= 2 && n_prefetch <= HTP_MM_MAX_PREFETCH && (n_prefetch & (n_prefetch - 1)) == 0); \ - \ - uint32_t n_k_tiles_w = ne00 / 32; \ - uint32_t n_k_tiles_a = ne10 / 32; \ - uint32_t tile_row_stride = n_k_tiles_w * tile_size; \ - uint32_t tile_row_transfer_size_aligned = n_k_tiles_a * aligned_tile_size; \ - dma_queue * dma_queue = octx->ctx->dma[ith]; \ - \ - const uint32_t src0_nrows = ne01 * src0->ne[2] * src0->ne[3]; \ - const uint32_t src0_start_row = mmctx->src0_nrows_per_thread * ith; \ - const uint32_t src0_end_row = MIN(src0_start_row + mmctx->src0_nrows_per_thread, src0_nrows); \ - \ - uint32_t ct_start = src0_start_row / 32; \ - uint32_t ct_end = (src0_end_row + 31) / 32; \ - \ - uint32_t push_ct = ct_start; \ - if (src0_start_row < src0_end_row) { \ - for (uint32_t d = 0; d < n_prefetch && push_ct < ct_end; d++, push_ct++) { \ - dma_queue_push(dma_queue, dma_make_ptr(vtcm_src0_ptr + d * tile_row_transfer_size_aligned, \ - src0_row + push_ct * tile_row_stride), aligned_tile_size, tile_size, tile_size, n_k_tiles_a); \ - dma_queue_push(dma_queue, dma_make_ptr(vtcm_src2_ptr + d * tile_row_transfer_size_aligned, \ - src2_row + push_ct * tile_row_stride), aligned_tile_size, tile_size, tile_size, n_k_tiles_a); \ - } \ - } \ - \ - hvx_mm_run_quant_task(mmctx, ith); \ - \ - if (src0_start_row >= src0_end_row) { \ - return; \ - } \ - \ - for (uint32_t ct = ct_start; ct < ct_end; ct++) { \ - const uint8_t * w_tile_gate = dma_queue_pop(dma_queue).dst; \ - const uint8_t * w_tile_up = dma_queue_pop(dma_queue).dst; \ - \ - int valid_rows = (int)ne01 - (int)(ct * 32); \ - valid_rows = MIN(32, MAX(0, valid_rows)); \ - \ - htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, ct); \ - uint32_t ir1 = 0; \ - for (; ir1 + 1 < src1_nrows; ir1 += 2) { \ - const uint8_t * restrict src1_col0 = (const uint8_t *) (src1_data + (ir1+0) * src1_stride); \ - const uint8_t * restrict src1_col1 = (const uint8_t *) (src1_data + (ir1+1) * src1_stride); \ - \ - float * restrict dst_row0_gate = (float *) (dst_gate->data + ((ir1+0) * dst_row_size)); \ - float * restrict dst_row1_gate = (float *) (dst_gate->data + ((ir1+1) * dst_row_size)); \ - float * dst_ptr0_gate = &dst_row0_gate[ct * 32]; \ - float * dst_ptr1_gate = &dst_row1_gate[ct * 32]; \ - \ - float * restrict dst_row0_up = (float *) (dst_up->data + ((ir1+0) * dst_row_size)); \ - float * restrict dst_row1_up = (float *) (dst_up->data + ((ir1+1) * dst_row_size)); \ - float * dst_ptr0_up = &dst_row0_up[ct * 32]; \ - float * dst_ptr1_up = &dst_row1_up[ct * 32]; \ - \ - DOT_2X2(ne10, dst_ptr0_gate, dst_ptr1_gate, w_tile_gate, src1_col0, src1_col1, valid_rows, NULL, NULL); \ - DOT_2X2(ne10, dst_ptr0_up, dst_ptr1_up, w_tile_up, src1_col0, src1_col1, valid_rows, NULL, NULL); \ - } \ - \ - for (; ir1 < src1_nrows; ++ir1) { \ - const uint8_t * restrict src1_col = (const uint8_t *) (src1_data + ir1 * src1_stride); \ - \ - float * restrict dst_row_gate = (float *) (dst_gate->data + (ir1 * dst_row_size)); \ - float * dst_ptr_gate = &dst_row_gate[ct * 32]; \ - \ - float * restrict dst_row_up = (float *) (dst_up->data + (ir1 * dst_row_size)); \ - float * dst_ptr_up = &dst_row_up[ct * 32]; \ - \ - DOT_2X1(ne10, dst_ptr_gate, w_tile_gate, src1_col, valid_rows, NULL); \ - DOT_2X1(ne10, dst_ptr_up, w_tile_up, src1_col, valid_rows, NULL); \ - } \ - htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, ct); \ - \ - if (push_ct < ct_end) { \ - dma_queue_push(dma_queue, dma_make_ptr((uint8_t *)w_tile_gate, src0_row + push_ct * tile_row_stride), \ - aligned_tile_size, tile_size, tile_size, n_k_tiles_a); \ - dma_queue_push(dma_queue, dma_make_ptr((uint8_t *)w_tile_up, src2_row + push_ct * tile_row_stride), \ - aligned_tile_size, tile_size, tile_size, n_k_tiles_a); \ - push_ct++; \ - } \ - } \ -} - MATMUL_2D_REPACKED_IMPL(q4_0, 576, tiled_vec_dot_q4_0_32x2, tiled_vec_dot_q4_0_32x1) MATMUL_2D_REPACKED_IMPL(q4_1, 640, tiled_vec_dot_q4_1_32x2, tiled_vec_dot_q4_1_32x1) MATMUL_2D_REPACKED_IMPL(q8_0, 1088, tiled_vec_dot_q8_0_32x2, tiled_vec_dot_q8_0_32x1) @@ -812,7 +613,7 @@ MATMUL_2D_REPACKED_IMPL(mxfp4_flat, 544, flat_vec_dot_mxfp4_32x2, flat_vec_dot static void name(unsigned int nth, unsigned int ith, void * data) { \ struct htp_mm_context * mmctx = data; \ struct htp_ops_context * octx = mmctx->octx; \ - const struct htp_tensor * src = octx->src[1]; \ + const struct htp_tensor * src = mmctx->act; \ const uint32_t ne0 = src->ne[0]; \ const uint32_t ne1 = src->ne[1]; \ const uint32_t ne2 = src->ne[2]; \ @@ -854,7 +655,7 @@ static void quantize_f32_q8_0_tiled_block(unsigned int nth, unsigned int ith, vo struct htp_thread_trace * tr = &octx->ctx->trace[ith]; htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_A_QUANT, mmctx->quant_ib_first[ith]); - const struct htp_tensor * src = octx->src[1]; + const struct htp_tensor * src = mmctx->act; quantize_f32_q8_0_tiled_block_kernel( (const float *) src->data, @@ -878,7 +679,7 @@ static void quantize_f32_q8_1_tiled_block(unsigned int nth, unsigned int ith, vo struct htp_thread_trace * tr = &octx->ctx->trace[ith]; htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_A_QUANT, mmctx->quant_ib_first[ith]); - const struct htp_tensor * src = octx->src[1]; + const struct htp_tensor * src = mmctx->act; quantize_f32_q8_1_tiled_block_kernel( (const float *) src->data, @@ -909,30 +710,17 @@ MATVEC_2D_REPACKED_IMPL(iq4nl_flat, 576, flat_vec_dot_iq4nl_32x1) MATVEC_2D_REPACKED_IMPL(mxfp4_flat, 544, flat_vec_dot_mxfp4_32x1) -MATMUL_QKV_2D_REPACKED_IMPL(q4_0, 576, tiled_vec_dot_q4_0_32x2, tiled_vec_dot_q4_0_32x1) -MATMUL_QKV_2D_REPACKED_IMPL(q4_1, 640, tiled_vec_dot_q4_1_32x2, tiled_vec_dot_q4_1_32x1) -MATMUL_QKV_2D_REPACKED_IMPL(q8_0, 1088, tiled_vec_dot_q8_0_32x2, tiled_vec_dot_q8_0_32x1) -MATMUL_QKV_2D_REPACKED_IMPL(iq4nl, 576, tiled_vec_dot_iq4nl_32x2, tiled_vec_dot_iq4nl_32x1) -MATMUL_QKV_2D_REPACKED_IMPL(mxfp4, 544, tiled_vec_dot_mxfp4_32x2, tiled_vec_dot_mxfp4_32x1) +MATMUL_NX_2D_REPACKED_IMPL(q4_0, 576, tiled_vec_dot_q4_0_32x2, tiled_vec_dot_q4_0_32x1) +MATMUL_NX_2D_REPACKED_IMPL(q4_1, 640, tiled_vec_dot_q4_1_32x2, tiled_vec_dot_q4_1_32x1) +MATMUL_NX_2D_REPACKED_IMPL(q8_0, 1088, tiled_vec_dot_q8_0_32x2, tiled_vec_dot_q8_0_32x1) +MATMUL_NX_2D_REPACKED_IMPL(iq4nl, 576, tiled_vec_dot_iq4nl_32x2, tiled_vec_dot_iq4nl_32x1) +MATMUL_NX_2D_REPACKED_IMPL(mxfp4, 544, tiled_vec_dot_mxfp4_32x2, tiled_vec_dot_mxfp4_32x1) -MATMUL_QKV_2D_REPACKED_IMPL(q4_0_flat, 576, flat_vec_dot_q4_0_32x2, flat_vec_dot_q4_0_32x1) -MATMUL_QKV_2D_REPACKED_IMPL(q4_1_flat, 640, flat_vec_dot_q4_1_32x2, flat_vec_dot_q4_1_32x1) -MATMUL_QKV_2D_REPACKED_IMPL(q8_0_flat, 1088, flat_vec_dot_q8_0_32x2, flat_vec_dot_q8_0_32x1) -MATMUL_QKV_2D_REPACKED_IMPL(iq4nl_flat, 576, flat_vec_dot_iq4nl_32x2, flat_vec_dot_iq4nl_32x1) -MATMUL_QKV_2D_REPACKED_IMPL(mxfp4_flat, 544, flat_vec_dot_mxfp4_32x2, flat_vec_dot_mxfp4_32x1) - - -MATMUL_FFN_2D_REPACKED_IMPL(q4_0, 576, tiled_vec_dot_q4_0_32x2, tiled_vec_dot_q4_0_32x1) -MATMUL_FFN_2D_REPACKED_IMPL(q4_1, 640, tiled_vec_dot_q4_1_32x2, tiled_vec_dot_q4_1_32x1) -MATMUL_FFN_2D_REPACKED_IMPL(q8_0, 1088, tiled_vec_dot_q8_0_32x2, tiled_vec_dot_q8_0_32x1) -MATMUL_FFN_2D_REPACKED_IMPL(iq4nl, 576, tiled_vec_dot_iq4nl_32x2, tiled_vec_dot_iq4nl_32x1) -MATMUL_FFN_2D_REPACKED_IMPL(mxfp4, 544, tiled_vec_dot_mxfp4_32x2, tiled_vec_dot_mxfp4_32x1) - -MATMUL_FFN_2D_REPACKED_IMPL(q4_0_flat, 576, flat_vec_dot_q4_0_32x2, flat_vec_dot_q4_0_32x1) -MATMUL_FFN_2D_REPACKED_IMPL(q4_1_flat, 640, flat_vec_dot_q4_1_32x2, flat_vec_dot_q4_1_32x1) -MATMUL_FFN_2D_REPACKED_IMPL(q8_0_flat, 1088, flat_vec_dot_q8_0_32x2, flat_vec_dot_q8_0_32x1) -MATMUL_FFN_2D_REPACKED_IMPL(iq4nl_flat, 576, flat_vec_dot_iq4nl_32x2, flat_vec_dot_iq4nl_32x1) -MATMUL_FFN_2D_REPACKED_IMPL(mxfp4_flat, 544, flat_vec_dot_mxfp4_32x2, flat_vec_dot_mxfp4_32x1) +MATMUL_NX_2D_REPACKED_IMPL(q4_0_flat, 576, flat_vec_dot_q4_0_32x2, flat_vec_dot_q4_0_32x1) +MATMUL_NX_2D_REPACKED_IMPL(q4_1_flat, 640, flat_vec_dot_q4_1_32x2, flat_vec_dot_q4_1_32x1) +MATMUL_NX_2D_REPACKED_IMPL(q8_0_flat, 1088, flat_vec_dot_q8_0_32x2, flat_vec_dot_q8_0_32x1) +MATMUL_NX_2D_REPACKED_IMPL(iq4nl_flat, 576, flat_vec_dot_iq4nl_32x2, flat_vec_dot_iq4nl_32x1) +MATMUL_NX_2D_REPACKED_IMPL(mxfp4_flat, 544, flat_vec_dot_mxfp4_32x2, flat_vec_dot_mxfp4_32x1) static void hvx_mm_2d(unsigned int nth, unsigned int ith, void * data) { htp_matmul_preamble; @@ -1353,6 +1141,7 @@ static int hvx_mm_matmul(struct htp_ops_context * octx) { struct htp_mm_context mmctx_struct = {0}; struct htp_mm_context * mmctx = &mmctx_struct; mmctx->octx = octx; + mmctx->act = src1; const struct htp_mm_kernel_params * kparams = (const struct htp_mm_kernel_params *) octx->kernel_params; @@ -1528,7 +1317,7 @@ static int hvx_mm_matmul(struct htp_ops_context * octx) { struct htp_mm_hvx_vtcm_layout L; htp_mm_hvx_vtcm_layout_build(&L, kparams->kernel_type, src0->type, ne10, src1_nrows, octx->n_threads, - dst_row_size, src0_row_size, src1_row_size, src2 ? src2->nb[1] : 0, kparams->n_prefetch, false, false, false); + dst_row_size, src0_row_size, src1_row_size, src2 ? src2->nb[1] : 0, kparams->n_prefetch, false, false); if (kparams->kernel_type == HTP_MM_KERNEL_HVX_F16_F16_VTCM || kparams->kernel_type == HTP_MM_KERNEL_HVX_F32_F32_VTCM || @@ -1587,297 +1376,97 @@ static int hvx_mm_matmul(struct htp_ops_context * octx) { return HTP_STATUS_OK; } -static void hvx_mm_qkv_2d(unsigned int nth, unsigned int ith, void * data) { +static void hvx_mm_nx_2d(unsigned int nth, unsigned int ith, void * data) { struct htp_mm_context * mmctx = data; struct htp_ops_context * octx = mmctx->octx; + const struct htp_mm_kernel_params * kparams = (const struct htp_mm_kernel_params *) octx->kernel_params; + const uint32_t n_weights = kparams->n_weights; - const struct htp_tensor * restrict src0 = octx->src[0]; // Wk - const struct htp_tensor * restrict src1 = octx->src[1]; // x - const struct htp_tensor * restrict src2 = octx->src[2]; // Wv - const struct htp_tensor * restrict src3 = octx->src[3]; // Wq - const struct htp_tensor * restrict dst_k = octx->dsts[0]; - const struct htp_tensor * restrict dst_v = octx->dsts[1]; - const struct htp_tensor * restrict dst_q = octx->dsts[2]; - - const uint32_t ne00 = src0->ne[0]; - const uint32_t ne01 = src0->ne[1]; - const uint32_t ne02 = src0->ne[2]; - const uint32_t ne03 = src0->ne[3]; - - const uint32_t ne11 = src1->ne[1]; - const uint32_t ne12 = src1->ne[2]; - const uint32_t ne13 = src1->ne[3]; - - const uint32_t src0_nrows = ne01 * ne02 * ne03; - const uint32_t src1_nrows = ne11 * ne12 * ne13; - - const uint32_t src0_nrows_per_thread = mmctx->src0_nrows_per_thread; - const uint32_t src0_start_row = src0_nrows_per_thread * ith; - const uint32_t src0_end_row = MIN(src0_start_row + src0_nrows_per_thread, src0_nrows); - const uint32_t src0_end_row_x2 = src0_start_row + ((src0_end_row - src0_start_row) & ~1U); - - const size_t dst_k_row_size = dst_k->nb[1]; // K and V share output width - const size_t dst_q_row_size = dst_q->nb[1]; // Q may be wider (GQA) - const size_t src0_row_size = src0->nb[1]; - const size_t src2_row_size = src2->nb[1]; - const size_t src3_row_size = src3->nb[1]; - - const size_t src0_stride = mmctx->vtcm_src0_stride; - const size_t src2_stride = mmctx->vtcm_src2_stride; - const size_t src3_stride = mmctx->vtcm_src3_stride; + const struct htp_tensor * restrict act = octx->src[n_weights]; + const uint32_t src1_nrows = act->ne[1] * act->ne[2] * act->ne[3]; const size_t src1_stride = mmctx->vtcm_src1_stride; uint8_t * restrict vtcm_src0_ptr = mmctx->vtcm_src0 + mmctx->vtcm_src0_size_per_thread * ith; - uint8_t * restrict vtcm_src2_ptr = mmctx->vtcm_src2 + mmctx->vtcm_src2_size_per_thread * ith; - uint8_t * restrict vtcm_src3_ptr = mmctx->vtcm_src3 + mmctx->vtcm_src3_size_per_thread * ith; - uint8_t * restrict src1_data = mmctx->vtcm_src1; + uint8_t * restrict src1_data = mmctx->vtcm_src1; dma_queue * dma_queue = octx->ctx->dma[ith]; - - const struct htp_mm_kernel_params * kparams = (const struct htp_mm_kernel_params *) octx->kernel_params; const uint32_t n_prefetch = kparams->n_prefetch; assert(n_prefetch >= 2 && n_prefetch <= HTP_MM_MAX_PREFETCH && (n_prefetch & (n_prefetch - 1)) == 0); const uint32_t prefetch_mask = n_prefetch - 1; - const uint8_t * restrict src0_row = (const uint8_t *) src0->data; - const uint8_t * restrict src2_row = (const uint8_t *) src2->data; - const uint8_t * restrict src3_row = (const uint8_t *) src3->data; - - // Prefill spad with src0, src2, src3 rows - if (src0_start_row < src0_end_row) { - for (uint32_t ir0 = src0_start_row; ir0 < src0_end_row_x2; ir0 += 2) { - const int is0 = (ir0 - src0_start_row); - if (is0 >= (int)n_prefetch) { - break; - } - dma_queue_push(dma_queue, dma_make_ptr(vtcm_src0_ptr + is0 * src0_stride, src0_row + ir0 * src0_row_size), - src0_stride, src0_row_size, src0_row_size, 2); - dma_queue_push(dma_queue, dma_make_ptr(vtcm_src2_ptr + is0 * src2_stride, src2_row + ir0 * src2_row_size), - src2_stride, src2_row_size, src2_row_size, 2); - dma_queue_push(dma_queue, dma_make_ptr(vtcm_src3_ptr + is0 * src3_stride, src3_row + ir0 * src3_row_size), - src3_stride, src3_row_size, src3_row_size, 2); - } - } + struct htp_thread_trace * tr = &octx->ctx->trace[ith]; hvx_mm_run_quant_task(mmctx, ith); - if (src0_start_row >= src0_end_row) { - return; - } - - // Process rows - for (uint32_t ir0 = src0_start_row; ir0 < src0_end_row_x2; ir0 += 2) { - const uint8_t * ss0 = dma_queue_pop(dma_queue).dst; - const uint8_t * ss2 = dma_queue_pop(dma_queue).dst; - const uint8_t * ss3 = dma_queue_pop(dma_queue).dst; - - // Process src1 columns in pairs (2×2 tiling) - uint32_t ir1 = 0; - for (; ir1 + 1 < src1_nrows; ir1 += 2) { - const uint8_t * restrict src1_col0 = (const uint8_t *) (src1_data + (ir1+0) * src1_stride); - const uint8_t * restrict src1_col1 = (const uint8_t *) (src1_data + (ir1+1) * src1_stride); - - float * restrict dst_row0_k = (float *) (dst_k->data + ((ir1+0) * dst_k_row_size)); - float * restrict dst_row1_k = (float *) (dst_k->data + ((ir1+1) * dst_k_row_size)); - mmctx->vec_dot_2x2(ne00, &dst_row0_k[ir0], &dst_row1_k[ir0], ss0, ss0 + src0_stride, src1_col0, src1_col1); - - float * restrict dst_row0_v = (float *) (dst_v->data + ((ir1+0) * dst_k_row_size)); - float * restrict dst_row1_v = (float *) (dst_v->data + ((ir1+1) * dst_k_row_size)); - mmctx->vec_dot_2x2(ne00, &dst_row0_v[ir0], &dst_row1_v[ir0], ss2, ss2 + src2_stride, src1_col0, src1_col1); - - float * restrict dst_row0_q = (float *) (dst_q->data + ((ir1+0) * dst_q_row_size)); - float * restrict dst_row1_q = (float *) (dst_q->data + ((ir1+1) * dst_q_row_size)); - mmctx->vec_dot_2x2(ne00, &dst_row0_q[ir0], &dst_row1_q[ir0], ss3, ss3 + src3_stride, src1_col0, src1_col1); - } - - // Handle remaining src1 rows (fallback to 2×1) - for (; ir1 < src1_nrows; ++ir1) { - const uint8_t * restrict src1_col = (const uint8_t *) (src1_data + ir1 * src1_stride); - - float * restrict dst_row_k = (float *) (dst_k->data + (ir1 * dst_k_row_size)); - mmctx->vec_dot_2x1(ne00, &dst_row_k[ir0], ss0, ss0 + src0_stride, src1_col); - - float * restrict dst_row_v = (float *) (dst_v->data + (ir1 * dst_k_row_size)); - mmctx->vec_dot_2x1(ne00, &dst_row_v[ir0], ss2, ss2 + src2_stride, src1_col); - - float * restrict dst_row_q = (float *) (dst_q->data + (ir1 * dst_q_row_size)); - mmctx->vec_dot_2x1(ne00, &dst_row_q[ir0], ss3, ss3 + src3_stride, src1_col); - } - - // Prefetch next (n + vtcm_nrows) rows - const int pr0 = (ir0 + n_prefetch); - const int is0 = (pr0 - src0_start_row) & prefetch_mask; - if (pr0 < src0_end_row_x2) { - dma_queue_push(dma_queue, dma_make_ptr(vtcm_src0_ptr + is0 * src0_stride, src0_row + pr0 * src0_row_size), - src0_stride, src0_row_size, src0_row_size, 2); - dma_queue_push(dma_queue, dma_make_ptr(vtcm_src2_ptr + is0 * src2_stride, src2_row + pr0 * src2_row_size), - src2_stride, src2_row_size, src2_row_size, 2); - dma_queue_push(dma_queue, dma_make_ptr(vtcm_src3_ptr + is0 * src3_stride, src3_row + pr0 * src3_row_size), - src3_stride, src3_row_size, src3_row_size, 2); - } - } - - // Process last row (if any) - if (src0_end_row != src0_end_row_x2) { - uint32_t ir0 = src0_end_row_x2; - const int is0 = (ir0 - src0_start_row) & prefetch_mask; - dma_queue_push(dma_queue, dma_make_ptr(vtcm_src0_ptr + is0 * src0_stride, src0_row + ir0 * src0_row_size), - src0_stride, src0_row_size, src0_row_size, 1); - dma_queue_push(dma_queue, dma_make_ptr(vtcm_src2_ptr + is0 * src2_stride, src2_row + ir0 * src2_row_size), - src2_stride, src2_row_size, src2_row_size, 1); - dma_queue_push(dma_queue, dma_make_ptr(vtcm_src3_ptr + is0 * src3_stride, src3_row + ir0 * src3_row_size), - src3_stride, src3_row_size, src3_row_size, 1); - - const uint8_t * ss0 = dma_queue_pop(dma_queue).dst; - const uint8_t * ss2 = dma_queue_pop(dma_queue).dst; - const uint8_t * ss3 = dma_queue_pop(dma_queue).dst; - - for (uint32_t ir1 = 0; ir1 < src1_nrows; ++ir1) { - const uint8_t * restrict src1_col = (const uint8_t *) (src1_data + ir1 * src1_stride); - - float * restrict dst_row_k = (float *) (dst_k->data + (ir1 * dst_k_row_size)); - mmctx->vec_dot_1x1(ne00, &dst_row_k[ir0], ss0, src1_col); + for (uint32_t widx = 0; widx < n_weights; widx++) { + const struct htp_tensor * restrict src_w = octx->src[widx]; + const struct htp_tensor * restrict dst = octx->dsts[widx]; + if (!src_w || !dst) continue; - float * restrict dst_row_v = (float *) (dst_v->data + (ir1 * dst_k_row_size)); - mmctx->vec_dot_1x1(ne00, &dst_row_v[ir0], ss2, src1_col); + const uint32_t ne00 = src_w->ne[0]; + const uint32_t ne01 = src_w->ne[1]; + const uint32_t src0_nrows = ne01 * src_w->ne[2] * src_w->ne[3]; - float * restrict dst_row_q = (float *) (dst_q->data + (ir1 * dst_q_row_size)); - mmctx->vec_dot_1x1(ne00, &dst_row_q[ir0], ss3, src1_col); - } - } -} - -static void hvx_mm_ffn_2d(unsigned int nth, unsigned int ith, void * data) { - struct htp_mm_context * mmctx = data; - struct htp_ops_context * octx = mmctx->octx; + uint32_t src0_nrows_per_thread = (src0_nrows + nth - 1) / nth; + src0_nrows_per_thread += (src0_nrows_per_thread & 1); - const struct htp_tensor * restrict src0 = octx->src[0]; // Wgate - const struct htp_tensor * restrict src1 = octx->src[1]; // y - const struct htp_tensor * restrict src2 = octx->src[2]; // Wup - const struct htp_tensor * restrict dst_gate = octx->dsts[0]; - const struct htp_tensor * restrict dst_up = octx->dsts[1]; + const uint32_t src0_start_row = src0_nrows_per_thread * ith; + const uint32_t src0_end_row = MIN(src0_start_row + src0_nrows_per_thread, src0_nrows); + const uint32_t src0_end_row_x2 = src0_start_row + ((src0_end_row - src0_start_row) & ~1U); + if (src0_start_row >= src0_end_row) continue; - const uint32_t ne00 = src0->ne[0]; - const uint32_t ne01 = src0->ne[1]; - const uint32_t ne02 = src0->ne[2]; - const uint32_t ne03 = src0->ne[3]; + const size_t dst_row_size = dst->nb[1]; + const size_t src0_row_size = src_w->nb[1]; + const size_t src0_stride = hex_round_up(src0_row_size, 128); - const uint32_t ne11 = src1->ne[1]; - const uint32_t ne12 = src1->ne[2]; - const uint32_t ne13 = src1->ne[3]; + const uint8_t * restrict src0_row = (const uint8_t *) src_w->data; - const uint32_t src0_nrows = ne01 * ne02 * ne03; - const uint32_t src1_nrows = ne11 * ne12 * ne13; - - const uint32_t src0_nrows_per_thread = mmctx->src0_nrows_per_thread; - const uint32_t src0_start_row = src0_nrows_per_thread * ith; - const uint32_t src0_end_row = MIN(src0_start_row + src0_nrows_per_thread, src0_nrows); - const uint32_t src0_end_row_x2 = src0_start_row + ((src0_end_row - src0_start_row) & ~1U); - - const size_t dst_row_size = dst_gate->nb[1]; - const size_t src0_row_size = src0->nb[1]; - const size_t src2_row_size = src2->nb[1]; - - const size_t src0_stride = mmctx->vtcm_src0_stride; - const size_t src2_stride = mmctx->vtcm_src2_stride; - const size_t src1_stride = mmctx->vtcm_src1_stride; - - uint8_t * restrict vtcm_src0_ptr = mmctx->vtcm_src0 + mmctx->vtcm_src0_size_per_thread * ith; - uint8_t * restrict vtcm_src2_ptr = mmctx->vtcm_src2 + mmctx->vtcm_src2_size_per_thread * ith; - uint8_t * restrict src1_data = mmctx->vtcm_src1; - - dma_queue * dma_queue = octx->ctx->dma[ith]; - - const struct htp_mm_kernel_params * kparams = (const struct htp_mm_kernel_params *) octx->kernel_params; - const uint32_t n_prefetch = kparams->n_prefetch; - assert(n_prefetch >= 2 && n_prefetch <= HTP_MM_MAX_PREFETCH && (n_prefetch & (n_prefetch - 1)) == 0); - const uint32_t prefetch_mask = n_prefetch - 1; - - const uint8_t * restrict src0_row = (const uint8_t *) src0->data; - const uint8_t * restrict src2_row = (const uint8_t *) src2->data; - - // Prefill spad with src0, src2 rows - if (src0_start_row < src0_end_row) { for (uint32_t ir0 = src0_start_row; ir0 < src0_end_row_x2; ir0 += 2) { const int is0 = (ir0 - src0_start_row); - if (is0 >= (int)n_prefetch) { - break; - } + if (is0 >= (int)n_prefetch) break; dma_queue_push(dma_queue, dma_make_ptr(vtcm_src0_ptr + is0 * src0_stride, src0_row + ir0 * src0_row_size), src0_stride, src0_row_size, src0_row_size, 2); - dma_queue_push(dma_queue, dma_make_ptr(vtcm_src2_ptr + is0 * src2_stride, src2_row + ir0 * src2_row_size), - src2_stride, src2_row_size, src2_row_size, 2); - } - } - - hvx_mm_run_quant_task(mmctx, ith); - - if (src0_start_row >= src0_end_row) { - return; - } - - // Process rows - for (uint32_t ir0 = src0_start_row; ir0 < src0_end_row_x2; ir0 += 2) { - const uint8_t * ss0 = dma_queue_pop(dma_queue).dst; - const uint8_t * ss2 = dma_queue_pop(dma_queue).dst; - - // Process src1 columns in pairs (2×2 tiling) - uint32_t ir1 = 0; - for (; ir1 + 1 < src1_nrows; ir1 += 2) { - const uint8_t * restrict src1_col0 = (const uint8_t *) (src1_data + (ir1+0) * src1_stride); - const uint8_t * restrict src1_col1 = (const uint8_t *) (src1_data + (ir1+1) * src1_stride); - - float * restrict dst_row0_gate = (float *) (dst_gate->data + ((ir1+0) * dst_row_size)); - float * restrict dst_row1_gate = (float *) (dst_gate->data + ((ir1+1) * dst_row_size)); - mmctx->vec_dot_2x2(ne00, &dst_row0_gate[ir0], &dst_row1_gate[ir0], ss0, ss0 + src0_stride, src1_col0, src1_col1); - - float * restrict dst_row0_up = (float *) (dst_up->data + ((ir1+0) * dst_row_size)); - float * restrict dst_row1_up = (float *) (dst_up->data + ((ir1+1) * dst_row_size)); - mmctx->vec_dot_2x2(ne00, &dst_row0_up[ir0], &dst_row1_up[ir0], ss2, ss2 + src2_stride, src1_col0, src1_col1); } - // Handle remaining src1 rows (fallback to 2×1) - for (; ir1 < src1_nrows; ++ir1) { - const uint8_t * restrict src1_col = (const uint8_t *) (src1_data + ir1 * src1_stride); - - float * restrict dst_row_gate = (float *) (dst_gate->data + (ir1 * dst_row_size)); - mmctx->vec_dot_2x1(ne00, &dst_row_gate[ir0], ss0, ss0 + src0_stride, src1_col); - - float * restrict dst_row_up = (float *) (dst_up->data + (ir1 * dst_row_size)); - mmctx->vec_dot_2x1(ne00, &dst_row_up[ir0], ss2, ss2 + src2_stride, src1_col); - } + for (uint32_t ir0 = src0_start_row; ir0 < src0_end_row_x2; ir0 += 2) { + const uint8_t * ss0 = dma_queue_pop(dma_queue).dst; + htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, ir0); + uint32_t ir1 = 0; + for (; ir1 + 1 < src1_nrows; ir1 += 2) { + const uint8_t * restrict src1_col0 = (const uint8_t *) (src1_data + (ir1+0) * src1_stride); + const uint8_t * restrict src1_col1 = (const uint8_t *) (src1_data + (ir1+1) * src1_stride); + float * restrict dst_row0 = (float *) (dst->data + ((ir1+0) * dst_row_size)); + float * restrict dst_row1 = (float *) (dst->data + ((ir1+1) * dst_row_size)); + mmctx->vec_dot_2x2(ne00, &dst_row0[ir0], &dst_row1[ir0], ss0, ss0 + src0_stride, src1_col0, src1_col1); + } + for (; ir1 < src1_nrows; ++ir1) { + const uint8_t * restrict src1_col = (const uint8_t *) (src1_data + ir1 * src1_stride); + float * restrict dst_row = (float *) (dst->data + (ir1 * dst_row_size)); + mmctx->vec_dot_2x1(ne00, &dst_row[ir0], ss0, ss0 + src0_stride, src1_col); + } + htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, ir0); - // Prefetch next rows - const int pr0 = (ir0 + n_prefetch); - const int is0 = (pr0 - src0_start_row) & prefetch_mask; - if (pr0 < src0_end_row_x2) { - dma_queue_push(dma_queue, dma_make_ptr(vtcm_src0_ptr + is0 * src0_stride, src0_row + pr0 * src0_row_size), - src0_stride, src0_row_size, src0_row_size, 2); - dma_queue_push(dma_queue, dma_make_ptr(vtcm_src2_ptr + is0 * src2_stride, src2_row + pr0 * src2_row_size), - src2_stride, src2_row_size, src2_row_size, 2); + const int pr0 = (ir0 + n_prefetch); + const int is0 = (pr0 - src0_start_row) & prefetch_mask; + if (pr0 < src0_end_row_x2) { + dma_queue_push(dma_queue, dma_make_ptr(vtcm_src0_ptr + is0 * src0_stride, src0_row + pr0 * src0_row_size), + src0_stride, src0_row_size, src0_row_size, 2); + } } - } - - // Process last row (if any) - if (src0_end_row != src0_end_row_x2) { - uint32_t ir0 = src0_end_row_x2; - const int is0 = (ir0 - src0_start_row) & prefetch_mask; - dma_queue_push(dma_queue, dma_make_ptr(vtcm_src0_ptr + is0 * src0_stride, src0_row + ir0 * src0_row_size), - src0_stride, src0_row_size, src0_row_size, 1); - dma_queue_push(dma_queue, dma_make_ptr(vtcm_src2_ptr + is0 * src2_stride, src2_row + ir0 * src2_row_size), - src2_stride, src2_row_size, src2_row_size, 1); - const uint8_t * ss0 = dma_queue_pop(dma_queue).dst; - const uint8_t * ss2 = dma_queue_pop(dma_queue).dst; - - for (uint32_t ir1 = 0; ir1 < src1_nrows; ++ir1) { - const uint8_t * restrict src1_col = (const uint8_t *) (src1_data + ir1 * src1_stride); - - float * restrict dst_row_gate = (float *) (dst_gate->data + (ir1 * dst_row_size)); - mmctx->vec_dot_1x1(ne00, &dst_row_gate[ir0], ss0, src1_col); - - float * restrict dst_row_up = (float *) (dst_up->data + (ir1 * dst_row_size)); - mmctx->vec_dot_1x1(ne00, &dst_row_up[ir0], ss2, src1_col); + if (src0_end_row != src0_end_row_x2) { + uint32_t ir0 = src0_end_row_x2; + const int is0 = (ir0 - src0_start_row) & prefetch_mask; + dma_queue_push(dma_queue, dma_make_ptr(vtcm_src0_ptr + is0 * src0_stride, src0_row + ir0 * src0_row_size), + src0_stride, src0_row_size, src0_row_size, 1); + const uint8_t * ss0 = dma_queue_pop(dma_queue).dst; + htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, ir0); + for (uint32_t ir1 = 0; ir1 < src1_nrows; ++ir1) { + const uint8_t * restrict src1_col = (const uint8_t *) (src1_data + ir1 * src1_stride); + float * restrict dst_row = (float *) (dst->data + (ir1 * dst_row_size)); + mmctx->vec_dot_1x1(ne00, &dst_row[ir0], ss0, src1_col); + } + htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, ir0); } } } @@ -3485,7 +3074,7 @@ static int hvx_mm_matmul_id( struct htp_mm_hvx_vtcm_layout L; htp_mm_hvx_vtcm_layout_build(&L, kparams->kernel_type, src0->type, ne10, src1_nrows, octx->n_threads, - 0, src0_row_size, src1_row_size, 0, kparams->n_prefetch, true, false, false); + 0, src0_row_size, src1_row_size, 0, kparams->n_prefetch, true, false); size_t vtcm_size = kparams->vtcm_size > 0 ? (size_t)kparams->vtcm_size : L.total_bytes; @@ -3610,6 +3199,7 @@ int op_matmul_id(struct htp_ops_context * octx) { struct htp_mm_context mmctx_struct = {0}; struct htp_mm_context * mmctx = &mmctx_struct; mmctx->octx = octx; + mmctx->act = src1; const struct htp_mm_kernel_params * kparams = (const struct htp_mm_kernel_params *) octx->kernel_params; @@ -3690,18 +3280,15 @@ int op_matmul_id(struct htp_ops_context * octx) { return s; } - -int op_matmul_qkv(struct htp_ops_context * octx) { +int op_matmul_nx(struct htp_ops_context * octx) { struct htp_thread_trace * tr = &octx->ctx->trace[0]; htp_trace_event_start(tr, HTP_TRACE_EVT_INIT, 0); - const struct htp_tensor * restrict src0 = octx->src[0]; // Wk - const struct htp_tensor * restrict src1 = octx->src[1]; // x - const struct htp_tensor * restrict src2 = octx->src[2]; // Wv - const struct htp_tensor * restrict src3 = octx->src[3]; // Wq - const struct htp_tensor * restrict dst_k = octx->dsts[0]; - const struct htp_tensor * restrict dst_v = octx->dsts[1]; - const struct htp_tensor * restrict dst_q = octx->dsts[2]; + const struct htp_mm_kernel_params * kparams = (const struct htp_mm_kernel_params *) octx->kernel_params; + const uint32_t n_weights = kparams->n_weights; + + const struct htp_tensor * restrict src0 = octx->src[0]; // first weight + const struct htp_tensor * restrict act = octx->src[n_weights]; // activation x bool is_repacked = (src0->type == HTP_TYPE_Q4_0 || src0->type == HTP_TYPE_Q4_1 || src0->type == HTP_TYPE_Q8_0 || src0->type == HTP_TYPE_IQ4_NL || @@ -3710,19 +3297,9 @@ int op_matmul_qkv(struct htp_ops_context * octx) { struct htp_mm_context mmctx_struct = {0}; struct htp_mm_context * mmctx = &mmctx_struct; mmctx->octx = octx; + mmctx->act = act; - const struct htp_mm_kernel_params * kparams = (const struct htp_mm_kernel_params *) octx->kernel_params; - - const uint32_t src0_nrows = src0->ne[1] * src0->ne[2] * src0->ne[3]; - const uint32_t src1_nrows = src1->ne[1] * src1->ne[2] * src1->ne[3]; - - // Compute src0_nrows_per_thread - mmctx->src0_nrows_per_thread = (src0_nrows + octx->n_threads - 1) / octx->n_threads; - if (is_repacked) { - mmctx->src0_nrows_per_thread = hex_round_up(mmctx->src0_nrows_per_thread, 32); - } else { - mmctx->src0_nrows_per_thread += (mmctx->src0_nrows_per_thread & 1); // round up to even - } + const uint32_t src1_nrows = act->ne[1] * act->ne[2] * act->ne[3]; const size_t src0_row_size = src0->nb[1]; const size_t src0_row_size_padded = hex_round_up(src0_row_size, 128); @@ -3732,7 +3309,7 @@ int op_matmul_qkv(struct htp_ops_context * octx) { } const uint32_t qk = QK_Q8_0_TILED; - const uint32_t nb = (src1->ne[0] + qk - 1) / qk; + const uint32_t nb = (act->ne[0] + qk - 1) / qk; const uint32_t total_nb = src1_nrows * nb; worker_callback_t quant_task_func; @@ -3758,185 +3335,39 @@ int op_matmul_qkv(struct htp_ops_context * octx) { size_t src1_row_size; if (kparams->kernel_type == HTP_MM_KERNEL_HVX_QUANT_ROW_FLAT) { - src1_row_size = (src0->type == HTP_TYPE_Q4_1) ? htp_mm_q8_1_flat_row_size(src1->ne[0]) : htp_mm_q8_0_flat_row_size(src1->ne[0]); + src1_row_size = (src0->type == HTP_TYPE_Q4_1) ? htp_mm_q8_1_flat_row_size(act->ne[0]) : htp_mm_q8_0_flat_row_size(act->ne[0]); } else { - src1_row_size = (src0->type == HTP_TYPE_Q4_1) ? htp_mm_q8_1_tiled_row_size(src1->ne[0]) : htp_mm_q8_0_tiled_row_size(src1->ne[0]); + src1_row_size = (src0->type == HTP_TYPE_Q4_1) ? htp_mm_q8_1_tiled_row_size(act->ne[0]) : htp_mm_q8_0_tiled_row_size(act->ne[0]); } struct htp_mm_hvx_vtcm_layout L; - htp_mm_hvx_vtcm_layout_build(&L, kparams->kernel_type, src0->type, src1->ne[0], src1_nrows, octx->n_threads, - 0, src0_row_size, src1_row_size, 0, kparams->n_prefetch, false, true, false); + htp_mm_hvx_vtcm_layout_build(&L, kparams->kernel_type, src0->type, act->ne[0], src1_nrows, octx->n_threads, + 0, src0_row_size, src1_row_size, 0, kparams->n_prefetch, false, true); size_t vtcm_size = kparams->vtcm_size > 0 ? (size_t)kparams->vtcm_size : L.total_bytes; if (octx->ctx->vtcm_size < vtcm_size) { - FARF(ERROR, "matmul-qkv: current VTCM reservation %zu is too small, needed %zu\n", + FARF(ERROR, "matmul-nx: current VTCM reservation %zu is too small, needed %zu\n", octx->ctx->vtcm_size, vtcm_size); return HTP_STATUS_VTCM_TOO_SMALL; } uint8_t * const base = (uint8_t *) octx->ctx->vtcm_base; - mmctx->vtcm_src1 = VTCM_LAYOUT_PTR(uint8_t, base, L.off_src1); mmctx->vtcm_src0 = VTCM_LAYOUT_PTR(uint8_t, base, L.off_src0); - mmctx->vtcm_src2 = VTCM_LAYOUT_PTR(uint8_t, base, L.off_src2); - mmctx->vtcm_src3 = VTCM_LAYOUT_PTR(uint8_t, base, L.off_src3); - mmctx->vtcm_dst = VTCM_LAYOUT_PTR(uint8_t, base, L.off_dst); - - octx->src1_spad.src = NULL; - octx->src0_spad.src = NULL; - octx->src2_spad.src = NULL; - octx->src3_spad.src = NULL; - octx->dst_spad.src = NULL; - - mmctx->vtcm_src0_stride = is_repacked ? 0 : src0_row_size_padded; - mmctx->vtcm_src2_stride = is_repacked ? 0 : src0_row_size_padded; - mmctx->vtcm_src3_stride = is_repacked ? 0 : src0_row_size_padded; - mmctx->vtcm_src1_stride = src1_row_size; - - mmctx->vtcm_src0_size_per_thread = L.src0_bytes / octx->n_threads; - mmctx->vtcm_src1_size_per_thread = L.src1_bytes; - mmctx->vtcm_src2_size_per_thread = L.src2_bytes / octx->n_threads; - mmctx->vtcm_src3_size_per_thread = L.src3_bytes / octx->n_threads; - mmctx->vtcm_dst_size_per_thread = L.dst_bytes / octx->n_threads; - - mmctx->n_quant_rows_per_thread = (src1_nrows + n_quant_tasks - 1) / n_quant_tasks; - mmctx->quant_task_func = quant_task_func; - mmctx->n_quant_tasks = n_quant_tasks; - atomic_init(&mmctx->quant_barrier, n_quant_tasks); - - // Run fused matmul - const uint32_t n_matmul_jobs = octx->n_threads; - worker_callback_t matmul_job_func; - if (is_repacked) { - if (kparams->kernel_type == HTP_MM_KERNEL_HVX_QUANT_ROW_FLAT) { - switch (src0->type) { - case HTP_TYPE_Q4_0: matmul_job_func = hvx_mm_qkv_2d_repacked_q4_0_flat; break; - case HTP_TYPE_Q4_1: matmul_job_func = hvx_mm_qkv_2d_repacked_q4_1_flat; break; - case HTP_TYPE_Q8_0: matmul_job_func = hvx_mm_qkv_2d_repacked_q8_0_flat; break; - case HTP_TYPE_IQ4_NL: matmul_job_func = hvx_mm_qkv_2d_repacked_iq4nl_flat; break; - case HTP_TYPE_MXFP4: matmul_job_func = hvx_mm_qkv_2d_repacked_mxfp4_flat; break; - default: return HTP_STATUS_NO_SUPPORT; - } - } else { - switch (src0->type) { - case HTP_TYPE_Q4_0: matmul_job_func = hvx_mm_qkv_2d_repacked_q4_0; break; - case HTP_TYPE_Q4_1: matmul_job_func = hvx_mm_qkv_2d_repacked_q4_1; break; - case HTP_TYPE_Q8_0: matmul_job_func = hvx_mm_qkv_2d_repacked_q8_0; break; - case HTP_TYPE_IQ4_NL: matmul_job_func = hvx_mm_qkv_2d_repacked_iq4nl; break; - case HTP_TYPE_MXFP4: matmul_job_func = hvx_mm_qkv_2d_repacked_mxfp4; break; - default: return HTP_STATUS_NO_SUPPORT; - } - } - } else { - matmul_job_func = hvx_mm_qkv_2d; - } - - htp_trace_event_stop(tr, HTP_TRACE_EVT_INIT, 0); - - worker_pool_run_func(octx->ctx->worker_pool, matmul_job_func, mmctx, n_matmul_jobs); - - return HTP_STATUS_OK; -} - -int op_matmul_ffn(struct htp_ops_context * octx) { - struct htp_thread_trace * tr = &octx->ctx->trace[0]; - htp_trace_event_start(tr, HTP_TRACE_EVT_INIT, 0); - - const struct htp_tensor * restrict src0 = octx->src[0]; // Wgate - const struct htp_tensor * restrict src1 = octx->src[1]; // y - const struct htp_tensor * restrict src2 = octx->src[2]; // Wup - const struct htp_tensor * restrict dst_gate = octx->dsts[0]; - const struct htp_tensor * restrict dst_up = octx->dsts[1]; - - bool is_repacked = (src0->type == HTP_TYPE_Q4_0 || src0->type == HTP_TYPE_Q4_1 || - src0->type == HTP_TYPE_Q8_0 || src0->type == HTP_TYPE_IQ4_NL || - src0->type == HTP_TYPE_MXFP4); - - struct htp_mm_context mmctx_struct = {0}; - struct htp_mm_context * mmctx = &mmctx_struct; - mmctx->octx = octx; - - const struct htp_mm_kernel_params * kparams = (const struct htp_mm_kernel_params *) octx->kernel_params; - - const uint32_t src0_nrows = src0->ne[1] * src0->ne[2] * src0->ne[3]; - const uint32_t src1_nrows = src1->ne[1] * src1->ne[2] * src1->ne[3]; - - // Compute src0_nrows_per_thread - mmctx->src0_nrows_per_thread = (src0_nrows + octx->n_threads - 1) / octx->n_threads; - if (is_repacked) { - mmctx->src0_nrows_per_thread = hex_round_up(mmctx->src0_nrows_per_thread, 32); - } else { - mmctx->src0_nrows_per_thread += (mmctx->src0_nrows_per_thread & 1); // round up to even - } - - const size_t src0_row_size = src0->nb[1]; - const size_t src0_row_size_padded = hex_round_up(src0_row_size, 128); - - if (hvx_mm_init_vec_dot(mmctx, src0->type) != 0) { - return HTP_STATUS_NO_SUPPORT; - } - - const uint32_t qk = QK_Q8_0_TILED; - const uint32_t nb = (src1->ne[0] + qk - 1) / qk; - const uint32_t total_nb = src1_nrows * nb; - - worker_callback_t quant_task_func; - uint32_t n_quant_tasks = 1; - if (kparams->kernel_type == HTP_MM_KERNEL_HVX_QUANT_ROW_FLAT) { - n_quant_tasks = MIN(src1_nrows, octx->n_threads); - quant_task_func = (src0->type == HTP_TYPE_Q4_1) ? quantize_f32_q8_1_flat : quantize_f32_q8_0_flat; - } else if (src1_nrows < octx->n_threads) { - n_quant_tasks = MIN(total_nb, octx->n_threads); - quant_task_func = (src0->type == HTP_TYPE_Q4_1) ? quantize_f32_q8_1_tiled_block : quantize_f32_q8_0_tiled_block; - for (uint32_t ith = 0; ith < n_quant_tasks; ++ith) { - uint32_t ib_first = (total_nb * (ith + 0)) / n_quant_tasks; - uint32_t ib_last = (total_nb * (ith + 1)) / n_quant_tasks; - mmctx->quant_ib_first[ith] = ib_first; - mmctx->quant_ib_last[ith] = ib_last; - mmctx->quant_r[ith] = ib_first / nb; - mmctx->quant_c[ith] = ib_first % nb; - } - } else { - n_quant_tasks = MIN(src1_nrows, octx->n_threads); - quant_task_func = (src0->type == HTP_TYPE_Q4_1) ? quantize_f32_q8_1_tiled : quantize_f32_q8_0_tiled; - } - - size_t src1_row_size; - if (kparams->kernel_type == HTP_MM_KERNEL_HVX_QUANT_ROW_FLAT) { - src1_row_size = (src0->type == HTP_TYPE_Q4_1) ? htp_mm_q8_1_flat_row_size(src1->ne[0]) : htp_mm_q8_0_flat_row_size(src1->ne[0]); - } else { - src1_row_size = (src0->type == HTP_TYPE_Q4_1) ? htp_mm_q8_1_tiled_row_size(src1->ne[0]) : htp_mm_q8_0_tiled_row_size(src1->ne[0]); - } - - struct htp_mm_hvx_vtcm_layout L; - htp_mm_hvx_vtcm_layout_build(&L, kparams->kernel_type, src0->type, src1->ne[0], src1_nrows, octx->n_threads, - 0, src0_row_size, src1_row_size, 0, kparams->n_prefetch, false, false, true); - - size_t vtcm_size = kparams->vtcm_size > 0 ? (size_t)kparams->vtcm_size : L.total_bytes; - - if (octx->ctx->vtcm_size < vtcm_size) { - FARF(ERROR, "matmul-ffn: current VTCM reservation %zu is too small, needed %zu\n", octx->ctx->vtcm_size, vtcm_size); - return HTP_STATUS_VTCM_TOO_SMALL; - } - - uint8_t * const base = (uint8_t *) octx->ctx->vtcm_base; mmctx->vtcm_src1 = VTCM_LAYOUT_PTR(uint8_t, base, L.off_src1); - mmctx->vtcm_src0 = VTCM_LAYOUT_PTR(uint8_t, base, L.off_src0); - mmctx->vtcm_src2 = VTCM_LAYOUT_PTR(uint8_t, base, L.off_src2); mmctx->vtcm_dst = VTCM_LAYOUT_PTR(uint8_t, base, L.off_dst); - octx->src1_spad.src = NULL; octx->src0_spad.src = NULL; + octx->src1_spad.src = NULL; octx->src2_spad.src = NULL; + octx->src3_spad.src = NULL; octx->dst_spad.src = NULL; mmctx->vtcm_src0_stride = is_repacked ? 0 : src0_row_size_padded; - mmctx->vtcm_src2_stride = is_repacked ? 0 : src0_row_size_padded; mmctx->vtcm_src1_stride = src1_row_size; mmctx->vtcm_src0_size_per_thread = L.src0_bytes / octx->n_threads; mmctx->vtcm_src1_size_per_thread = L.src1_bytes; - mmctx->vtcm_src2_size_per_thread = L.src2_bytes / octx->n_threads; mmctx->vtcm_dst_size_per_thread = L.dst_bytes / octx->n_threads; mmctx->n_quant_rows_per_thread = (src1_nrows + n_quant_tasks - 1) / n_quant_tasks; @@ -3950,25 +3381,25 @@ int op_matmul_ffn(struct htp_ops_context * octx) { if (is_repacked) { if (kparams->kernel_type == HTP_MM_KERNEL_HVX_QUANT_ROW_FLAT) { switch (src0->type) { - case HTP_TYPE_Q4_0: matmul_job_func = hvx_mm_ffn_2d_repacked_q4_0_flat; break; - case HTP_TYPE_Q4_1: matmul_job_func = hvx_mm_ffn_2d_repacked_q4_1_flat; break; - case HTP_TYPE_Q8_0: matmul_job_func = hvx_mm_ffn_2d_repacked_q8_0_flat; break; - case HTP_TYPE_IQ4_NL: matmul_job_func = hvx_mm_ffn_2d_repacked_iq4nl_flat; break; - case HTP_TYPE_MXFP4: matmul_job_func = hvx_mm_ffn_2d_repacked_mxfp4_flat; break; + case HTP_TYPE_Q4_0: matmul_job_func = hvx_mm_nx_2d_repacked_q4_0_flat; break; + case HTP_TYPE_Q4_1: matmul_job_func = hvx_mm_nx_2d_repacked_q4_1_flat; break; + case HTP_TYPE_Q8_0: matmul_job_func = hvx_mm_nx_2d_repacked_q8_0_flat; break; + case HTP_TYPE_IQ4_NL: matmul_job_func = hvx_mm_nx_2d_repacked_iq4nl_flat; break; + case HTP_TYPE_MXFP4: matmul_job_func = hvx_mm_nx_2d_repacked_mxfp4_flat; break; default: return HTP_STATUS_NO_SUPPORT; } } else { switch (src0->type) { - case HTP_TYPE_Q4_0: matmul_job_func = hvx_mm_ffn_2d_repacked_q4_0; break; - case HTP_TYPE_Q4_1: matmul_job_func = hvx_mm_ffn_2d_repacked_q4_1; break; - case HTP_TYPE_Q8_0: matmul_job_func = hvx_mm_ffn_2d_repacked_q8_0; break; - case HTP_TYPE_IQ4_NL: matmul_job_func = hvx_mm_ffn_2d_repacked_iq4nl; break; - case HTP_TYPE_MXFP4: matmul_job_func = hvx_mm_ffn_2d_repacked_mxfp4; break; + case HTP_TYPE_Q4_0: matmul_job_func = hvx_mm_nx_2d_repacked_q4_0; break; + case HTP_TYPE_Q4_1: matmul_job_func = hvx_mm_nx_2d_repacked_q4_1; break; + case HTP_TYPE_Q8_0: matmul_job_func = hvx_mm_nx_2d_repacked_q8_0; break; + case HTP_TYPE_IQ4_NL: matmul_job_func = hvx_mm_nx_2d_repacked_iq4nl; break; + case HTP_TYPE_MXFP4: matmul_job_func = hvx_mm_nx_2d_repacked_mxfp4; break; default: return HTP_STATUS_NO_SUPPORT; } } } else { - matmul_job_func = hvx_mm_ffn_2d; + matmul_job_func = hvx_mm_nx_2d; } htp_trace_event_stop(tr, HTP_TRACE_EVT_INIT, 0); diff --git a/ggml/src/ggml-hexagon/htp/matmul-ops.h b/ggml/src/ggml-hexagon/htp/matmul-ops.h index 6c393664c6e..dbc8e359093 100644 --- a/ggml/src/ggml-hexagon/htp/matmul-ops.h +++ b/ggml/src/ggml-hexagon/htp/matmul-ops.h @@ -88,6 +88,7 @@ struct htp_mm_kernel_params { int32_t vtcm_src2_size; // src2 scratchpad size in VTCM (fused only) int32_t vtcm_src3_size; // src3 scratchpad size in VTCM (fused only) int32_t vtcm_dst_size; // dst scratchpad size in VTCM + int32_t n_weights; // Number of weights for fused NX // Precomputed division values struct fastdiv_values div_ne12_ne1; @@ -463,8 +464,7 @@ static inline void htp_mm_hvx_vtcm_layout_build( size_t src2_row_size, uint32_t n_prefetch, bool is_matmul_id, - bool is_fused_qkv, - bool is_fused_ffn + bool is_fused_nx ) { size_t src0_sz = 0; size_t src1_sz = 0; @@ -476,44 +476,33 @@ static inline void htp_mm_hvx_vtcm_layout_build( wtype == HTP_TYPE_Q8_0 || wtype == HTP_TYPE_IQ4_NL || wtype == HTP_TYPE_MXFP4); - if (is_fused_qkv || is_fused_ffn) { + if (is_fused_nx) { const size_t src0_row_size_padded = hex_round_up(src0_row_size, 128); const size_t quant_scratch_size = hex_round_up(ne10 * sizeof(float), QK_Q8_0_TILED * sizeof(float)) * n_threads; - size_t src0_sz_per_thread = 0; - size_t src2_sz_per_thread = 0; - size_t src3_sz_per_thread = 0; + size_t weight_sz_per_thread = 0; if (is_repack) { uint32_t aligned_tile_size = htp_mm_get_weight_aligned_tile_size(wtype); uint32_t n_k_tiles = hex_round_up(ne10, 32) / 32; uint32_t tile_row_size = n_k_tiles * aligned_tile_size; - src0_sz_per_thread = hex_round_up(n_prefetch * tile_row_size, 128); - src2_sz_per_thread = hex_round_up(n_prefetch * tile_row_size, 128); - if (is_fused_qkv) { - src3_sz_per_thread = hex_round_up(n_prefetch * tile_row_size, 128); - } + weight_sz_per_thread = hex_round_up(n_prefetch * tile_row_size, 128); } else { - src0_sz_per_thread = hex_round_up(n_prefetch * src0_row_size_padded, 128); - src2_sz_per_thread = hex_round_up(n_prefetch * src0_row_size_padded, 128); - if (is_fused_qkv) { - src3_sz_per_thread = hex_round_up(n_prefetch * src0_row_size_padded, 128); - } + weight_sz_per_thread = hex_round_up(n_prefetch * src0_row_size_padded, 128); } - size_t flat_src1_row_size = (wtype == HTP_TYPE_Q4_1) ? htp_mm_q8_1_flat_row_size(ne10) : htp_mm_q8_0_flat_row_size(ne10); - size_t tiled_src1_row_size = (wtype == HTP_TYPE_Q4_1) ? htp_mm_q8_1_tiled_row_size(ne10) : htp_mm_q8_0_tiled_row_size(ne10); + size_t flat_act_row_size = (wtype == HTP_TYPE_Q4_1) ? htp_mm_q8_1_flat_row_size(ne10) : htp_mm_q8_0_flat_row_size(ne10); + size_t tiled_act_row_size = (wtype == HTP_TYPE_Q4_1) ? htp_mm_q8_1_tiled_row_size(ne10) : htp_mm_q8_0_tiled_row_size(ne10); - if (kernel_type == HTP_MM_KERNEL_HVX_QUANT_ROW_FLAT) { - src1_sz = hex_round_up(flat_src1_row_size * src1_nrows, 128); - } else { - src1_sz = hex_round_up(tiled_src1_row_size * src1_nrows, 128); - } + size_t act_sz = (kernel_type == HTP_MM_KERNEL_HVX_QUANT_ROW_FLAT) + ? hex_round_up(flat_act_row_size * src1_nrows, 128) + : hex_round_up(tiled_act_row_size * src1_nrows, 128); - src0_sz = src0_sz_per_thread * n_threads; - src2_sz = src2_sz_per_thread * n_threads; - src3_sz = src3_sz_per_thread * n_threads; + src0_sz = weight_sz_per_thread * n_threads; // shared single-weight prefetch buffer + src1_sz = act_sz; // quantized activation buffer + src2_sz = 0; + src3_sz = 0; dst_sz = quant_scratch_size; } else if (is_matmul_id) { const size_t src0_row_size_padded = htp_mm_round_up(src0_row_size, 128); @@ -616,8 +605,8 @@ static inline void htp_mm_hvx_vtcm_layout_build( } size_t off = 0; - VTCM_LAYOUT_ALLOC(off, off_src1, src1_sz); VTCM_LAYOUT_ALLOC(off, off_src0, src0_sz); + VTCM_LAYOUT_ALLOC(off, off_src1, src1_sz); VTCM_LAYOUT_ALLOC(off, off_src2, src2_sz); VTCM_LAYOUT_ALLOC(off, off_src3, src3_sz); VTCM_LAYOUT_ALLOC(off, off_dst, dst_sz); diff --git a/ggml/src/ggml-hexagon/htp/set-rows-ops.c b/ggml/src/ggml-hexagon/htp/set-rows-ops.c index 58c54967db0..fa14bf0ef6b 100644 --- a/ggml/src/ggml-hexagon/htp/set-rows-ops.c +++ b/ggml/src/ggml-hexagon/htp/set-rows-ops.c @@ -8,14 +8,20 @@ #include #include -#include "hex-dma.h" +#include "dma-queue.h" +#include "work-queue.h" #include "hvx-utils.h" +#include "hex-utils.h" +#include "hvx-copy.h" +#include "hvx-quant.h" #define GGML_COMMON_DECL_C #include "ggml-common.h" + #include "htp-ctx.h" #include "htp-ops.h" -#include "htp-ops.h" +#include "htp-tensor.h" +#include "htp/set-rows-ops.h" #define set_rows_preamble \ const uint32_t ne00 = octx->src[0]->ne[0]; \ @@ -47,116 +53,142 @@ \ const uint32_t nr = ne01; -struct htp_set_rows_context { +struct set_rows_context { struct htp_ops_context * octx; - struct fastdiv_values div_ne12; - struct fastdiv_values div_ne11; - uint32_t src0_nrows_per_thread; + const struct htp_set_rows_kernel_params * kparams; + struct htp_set_rows_vtcm_layout vtcm_layout; + uint8_t * vtcm_base; }; -static void set_rows_thread_f32_f32(unsigned int nth, unsigned int ith, void *data) { - struct htp_set_rows_context * srctx = (struct htp_set_rows_context *)data; - struct htp_ops_context * octx = srctx->octx; - - set_rows_preamble; - - uint64_t qt = HAP_perf_get_qtimer_count(); - - // parallelize by rows of src0 - const uint32_t dr = srctx->src0_nrows_per_thread; - const uint32_t ir0 = dr * ith; - if (ir0 >= nr) { - return; - } - const uint32_t ir1 = (ir0 + dr < nr) ? (ir0 + dr) : nr; - - const bool is_i32 = (octx->src[1]->type == HTP_TYPE_I32); - - for (uint32_t i03 = 0; i03 < ne03; ++i03) { - for (uint32_t i02 = 0; i02 < ne02; ++i02) { - for (uint32_t i = ir0; i < ir1; ++i) { - const uint32_t i12 = fastmodulo(i03, ne12, &srctx->div_ne12); - const uint32_t i11 = fastmodulo(i02, ne11, &srctx->div_ne11); - const uint32_t i10 = i; - - const uintptr_t src1_addr = octx->src[1]->data + i10*nb10 + i11*nb11 + i12*nb12; - - uint32_t i1 = is_i32 ? *(int32_t *)src1_addr : *(int64_t *)src1_addr; - if (i1 >= ne1) { - // ignore invalid indices - continue; - } - - const uintptr_t src0_ptr = octx->src[0]->data + i*nb01 + i02*nb02 + i03*nb03; - const uintptr_t dst_ptr = octx->dst->data + i1*nb1 + i02*nb2 + i03*nb3; - - // copy row - hvx_copy_f32_uu((uint8_t *)dst_ptr, (const uint8_t *)src0_ptr, ne00); - } - } - } - - qt = HAP_perf_qtimer_count_to_us(HAP_perf_get_qtimer_count() - qt); - FARF(HIGH, "set-rows-f32-f32 %d/%d: %ux%ux%ux%u (%u:%u) x %ux%ux%ux%u -> %ux%ux%ux%u usec %u\n", ith, nth, - ne00, ne01, ne02, ne03, ir0, ir1, ne10, ne11, ne12, ne13, ne0, ne1, ne2, ne3, (unsigned) qt); +#define SET_ROWS_THREAD_DMA_FN(TYPE_NAME, IDX_TYPE, COMPUTE_EXPR) \ +static void set_rows_thread_dma_##TYPE_NAME##_##IDX_TYPE(unsigned int nth, unsigned int ith, void *data) { \ + struct set_rows_context * srctx = (struct set_rows_context *)data; \ + struct htp_ops_context * octx = srctx->octx; \ + const struct htp_set_rows_kernel_params * kparams = srctx->kparams; \ + set_rows_preamble; \ + struct htp_thread_trace * tr = &octx->ctx->trace[ith]; \ + const uint32_t dr = kparams->tasks_per_thread; \ + const uint32_t ir0 = dr * ith; \ + if (ir0 >= kparams->total_tasks) { \ + return; \ + } \ + const uint32_t ir1 = MIN(ir0 + dr, kparams->total_tasks); \ + dma_queue * dma_queue = octx->ctx->dma[ith]; \ + const struct htp_set_rows_vtcm_layout * vtcm_layout = &srctx->vtcm_layout; \ + uint8_t * vtcm_src0 = srctx->vtcm_base + vtcm_layout->off_src0 + ith * vtcm_layout->src0_bytes_per_thread; \ + uint8_t * vtcm_dst = srctx->vtcm_base + vtcm_layout->off_dst + ith * vtcm_layout->dst_bytes_per_thread; \ + const uint32_t src0_row_size = ne00 * sizeof(float); \ + const uint32_t dst_row_size = htp_tensor_get_row_size(octx->dst->type, ne00); \ + const uint32_t nrows_per_thread = ir1 - ir0; \ + const uint32_t total_steps = ne03 * ne02 * nrows_per_thread; \ + uint32_t pi_step = 0; \ + uint32_t pi02 = 0; \ + uint32_t pi03 = 0; \ + for (uint32_t step = 0, spad_idx = 0; step < total_steps && spad_idx < 2; ++step, spad_idx++) { \ + uint32_t i = ir0 + pi_step; \ + const uintptr_t src0_ptr = octx->src[0]->data + i*nb01 + pi02*nb02 + pi03*nb03; \ + dma_queue_push(dma_queue, \ + dma_make_ptr((void *)octx->dst->data, \ + vtcm_dst + spad_idx * vtcm_layout->dst_spad_half_size), \ + dst_row_size, vtcm_layout->dst_spad_half_size, dst_row_size, 0); \ + dma_queue_push(dma_queue, \ + dma_make_ptr((void *)(vtcm_src0 + spad_idx * vtcm_layout->src0_spad_half_size), \ + (const void *)src0_ptr), \ + vtcm_layout->src0_spad_half_size, src0_row_size, src0_row_size, 1); \ + pi_step++; \ + if (pi_step == nrows_per_thread) { \ + pi_step = 0; \ + pi02++; \ + if (pi02 == ne02) { \ + pi02 = 0; \ + pi03++; \ + } \ + } \ + } \ + uint32_t ci_step = 0; \ + uint32_t ci02 = 0; \ + uint32_t ci03 = 0; \ + uint32_t ci11_base = 0; \ + uint32_t ci12_base = 0; \ + for (uint32_t step = 0; step < total_steps; ++step) { \ + void * dst_spad = (void *) dma_queue_pop(dma_queue).src; \ + void * src_spad = (void *) dma_queue_pop(dma_queue).dst; \ + uint32_t i = ir0 + ci_step; \ + const uintptr_t src1_addr = octx->src[1]->data + i*nb10 + ci11_base*nb11 + ci12_base*nb12; \ + const IDX_TYPE i1 = *(const IDX_TYPE *)src1_addr; \ + const bool valid_i1 = ((uint64_t)i1 < (uint64_t)ne1); \ + const uint32_t target_i1 = (uint32_t)i1; \ + htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, step); \ + if (valid_i1) { \ + COMPUTE_EXPR; \ + } \ + htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, step); \ + if (valid_i1) { \ + const uintptr_t dst_ptr = octx->dst->data + target_i1*nb1 + ci02*nb2 + ci03*nb3; \ + dma_queue_push(dma_queue, \ + dma_make_ptr((void *)dst_ptr, (const void *)dst_spad), \ + dst_row_size, vtcm_layout->dst_spad_half_size, dst_row_size, 1); \ + } else { \ + dma_queue_push(dma_queue, \ + dma_make_ptr((void *)octx->dst->data, (const void *)dst_spad), \ + dst_row_size, vtcm_layout->dst_spad_half_size, dst_row_size, 0); \ + } \ + const uint32_t next_step = step + 2; \ + if (next_step < total_steps) { \ + uint32_t ni = ir0 + pi_step; \ + const uintptr_t psrc0_ptr = octx->src[0]->data + ni*nb01 + pi02*nb02 + pi03*nb03; \ + dma_queue_push(dma_queue, \ + dma_make_ptr((void *)src_spad, (const void *)psrc0_ptr), \ + vtcm_layout->src0_spad_half_size, src0_row_size, src0_row_size, 1); \ + pi_step++; \ + if (pi_step == nrows_per_thread) { \ + pi_step = 0; \ + pi02++; \ + if (pi02 == ne02) { \ + pi02 = 0; \ + pi03++; \ + } \ + } \ + } \ + ci_step++; \ + if (ci_step == nrows_per_thread) { \ + ci_step = 0; \ + ci02++; \ + ci11_base++; \ + if (ci11_base == ne11) { \ + ci11_base = 0; \ + } \ + if (ci02 == ne02) { \ + ci02 = 0; \ + ci03++; \ + ci12_base++; \ + if (ci12_base == ne12) { \ + ci12_base = 0; \ + } \ + } \ + } \ + } \ + dma_queue_flush(dma_queue); \ } -static void set_rows_thread_f16_f32(unsigned int nth, unsigned int ith, void *data) { - struct htp_set_rows_context * srctx = (struct htp_set_rows_context *)data; - struct htp_ops_context * octx = srctx->octx; - - set_rows_preamble; +SET_ROWS_THREAD_DMA_FN(f32, int32_t, { hvx_copy_f32_uu((uint8_t *)dst_spad, (const uint8_t *)src_spad, ne00); }) +SET_ROWS_THREAD_DMA_FN(f32, int64_t, { hvx_copy_f32_uu((uint8_t *)dst_spad, (const uint8_t *)src_spad, ne00); }) - uint64_t qt = HAP_perf_get_qtimer_count(); +SET_ROWS_THREAD_DMA_FN(f16, int32_t, { hvx_copy_f16_f32_uu((uint8_t *)dst_spad, (const uint8_t *)src_spad, ne00); }) +SET_ROWS_THREAD_DMA_FN(f16, int64_t, { hvx_copy_f16_f32_uu((uint8_t *)dst_spad, (const uint8_t *)src_spad, ne00); }) - // parallelize by rows of src0 - const uint32_t dr = srctx->src0_nrows_per_thread; - const uint32_t ir0 = dr * ith; - if (ir0 >= nr) { - return; - } - const uint32_t ir1 = (ir0 + dr < nr) ? (ir0 + dr) : nr; - - const bool is_i32 = (octx->src[1]->type == HTP_TYPE_I32); - - for (uint32_t i03 = 0; i03 < ne03; ++i03) { - for (uint32_t i02 = 0; i02 < ne02; ++i02) { - for (uint32_t i = ir0; i < ir1; ++i) { - const uint32_t i12 = fastmodulo(i03, ne12, &srctx->div_ne12); - const uint32_t i11 = fastmodulo(i02, ne11, &srctx->div_ne11); - const uint32_t i10 = i; - - const uintptr_t src1_addr = octx->src[1]->data + i10*nb10 + i11*nb11 + i12*nb12; - - uint32_t i1 = is_i32 ? *(int32_t *)src1_addr : *(int64_t *)src1_addr; - if (i1 >= ne1) { - // ignore invalid indices - continue; - } - - const uint8_t* src0_ptr = (const uint8_t *) octx->src[0]->data + i*nb01 + i02*nb02 + i03*nb03; - uint8_t* dst_ptr = (uint8_t *) octx->dst->data + i1*nb1 + i02*nb2 + i03*nb3; - - hvx_copy_f16_f32_uu(dst_ptr, src0_ptr, ne00); - } - } - } - - qt = HAP_perf_qtimer_count_to_us(HAP_perf_get_qtimer_count() - qt); - FARF(HIGH, "set-rows-f16-f32 %d/%d: %ux%ux%ux%u (%u:%u) x %ux%ux%ux%u -> %ux%ux%ux%u usec %u\n", ith, nth, - ne00, ne01, ne02, ne03, ir0, ir1, ne10, ne11, ne12, ne13, ne0, ne1, ne2, ne3, (unsigned) qt); -} +SET_ROWS_THREAD_DMA_FN(q8_0, int32_t, { hvx_quantize_row_q8_0_f32(dst_spad, (const float *)src_spad, ne00); }) +SET_ROWS_THREAD_DMA_FN(q8_0, int64_t, { hvx_quantize_row_q8_0_f32(dst_spad, (const float *)src_spad, ne00); }) int op_set_rows(struct htp_ops_context * octx) { + const struct htp_set_rows_kernel_params * kparams = (const struct htp_set_rows_kernel_params *)octx->kernel_params; set_rows_preamble; - const uint32_t n_threads = MIN(nr, octx->n_threads); - if (octx->src[0]->type != HTP_TYPE_F32) { return HTP_STATUS_NO_SUPPORT; } - if (octx->dst->type != HTP_TYPE_F32 && octx->dst->type != HTP_TYPE_F16) { + if (octx->dst->type != HTP_TYPE_F32 && octx->dst->type != HTP_TYPE_F16 && octx->dst->type != HTP_TYPE_Q8_0) { return HTP_STATUS_NO_SUPPORT; } @@ -164,27 +196,27 @@ int op_set_rows(struct htp_ops_context * octx) { return HTP_STATUS_NO_SUPPORT; } - if (octx->flags & HTP_OPFLAGS_SKIP_COMPUTE) { - return HTP_STATUS_OK; - } + // l2fetch the src1 (indices) tensor in the main thread + hex_l2fetch_block((const void *)octx->src[1]->data, octx->src[1]->ne[3] * octx->src[1]->nb[3]); - struct htp_set_rows_context srctx; + struct set_rows_context srctx; srctx.octx = octx; - srctx.div_ne12 = init_fastdiv_values(ne12); - srctx.div_ne11 = init_fastdiv_values(ne11); - - srctx.src0_nrows_per_thread = (nr + n_threads - 1) / n_threads; - - switch(octx->dst->type) { - case HTP_TYPE_F32: - worker_pool_run_func(octx->ctx->worker_pool, set_rows_thread_f32_f32, &srctx, n_threads); - break; - case HTP_TYPE_F16: - worker_pool_run_func(octx->ctx->worker_pool, set_rows_thread_f16_f32, &srctx, n_threads); - break; - default: - return HTP_STATUS_NO_SUPPORT; + srctx.kparams = kparams; + + htp_set_rows_vtcm_layout_build(&srctx.vtcm_layout, octx->dst->type, ne00, kparams->n_threads); + srctx.vtcm_base = (uint8_t *)octx->ctx->vtcm_base; + + work_queue_func_t q_func = NULL; + const bool is_i32 = (octx->src[1]->type == HTP_TYPE_I32); + + switch (octx->dst->type) { + case HTP_TYPE_F32: q_func = is_i32 ? set_rows_thread_dma_f32_int32_t : set_rows_thread_dma_f32_int64_t; break; + case HTP_TYPE_F16: q_func = is_i32 ? set_rows_thread_dma_f16_int32_t : set_rows_thread_dma_f16_int64_t; break; + case HTP_TYPE_Q8_0: q_func = is_i32 ? set_rows_thread_dma_q8_0_int32_t : set_rows_thread_dma_q8_0_int64_t; break; + default: return HTP_STATUS_NO_SUPPORT; } + work_queue_run(octx->ctx->work_queue, q_func, &srctx, kparams->n_threads); + return HTP_STATUS_OK; } diff --git a/ggml/src/ggml-hexagon/htp/set-rows-ops.h b/ggml/src/ggml-hexagon/htp/set-rows-ops.h new file mode 100644 index 00000000000..5e98d2cb55c --- /dev/null +++ b/ggml/src/ggml-hexagon/htp/set-rows-ops.h @@ -0,0 +1,74 @@ +#ifndef HTP_SET_ROWS_OPS_H +#define HTP_SET_ROWS_OPS_H + +#include "hex-fastdiv.h" + +struct htp_set_rows_kernel_params { + int32_t n_threads; + int32_t total_tasks; + int32_t tasks_per_thread; + int32_t vtcm_size; + + // Fastdiv helpers + struct fastdiv_values div_ne11; + struct fastdiv_values div_ne12; + struct fastdiv_values div_tasks_per_thread; + struct fastdiv_values div_ne02; +}; + +struct htp_set_rows_vtcm_layout { + size_t total_bytes; + size_t off_src0; + size_t off_dst; + + size_t src0_bytes_per_thread; + size_t dst_bytes_per_thread; + + size_t src0_spad_half_size; + size_t dst_spad_half_size; +}; + +static inline void htp_set_rows_vtcm_layout_build( + struct htp_set_rows_vtcm_layout * vtcm_layout, + int dst_type, + uint32_t ne00, + uint32_t n_threads) { + + size_t src0_row_size = ne00 * 4; + size_t dst_row_size = 0; + switch (dst_type) { + case 0: // HTP_TYPE_F32 + dst_row_size = ne00 * 4; + break; + case 1: // HTP_TYPE_F16 + dst_row_size = ne00 * 2; + break; + case 8: // HTP_TYPE_Q8_0 + dst_row_size = (ne00 / 32) * 34; + break; + default: + dst_row_size = 0; + break; + } + + size_t src0_row_size_aligned = (src0_row_size + 255) & ~255; + size_t dst_row_size_aligned = (dst_row_size + 255) & ~255; + + vtcm_layout->src0_spad_half_size = src0_row_size_aligned; + vtcm_layout->dst_spad_half_size = dst_row_size_aligned; + + vtcm_layout->src0_bytes_per_thread = src0_row_size_aligned * 2; + vtcm_layout->dst_bytes_per_thread = dst_row_size_aligned * 2; + + vtcm_layout->off_src0 = 0; + vtcm_layout->off_dst = vtcm_layout->off_src0 + vtcm_layout->src0_bytes_per_thread * n_threads; + vtcm_layout->total_bytes = vtcm_layout->off_dst + vtcm_layout->dst_bytes_per_thread * n_threads; +} + +#if defined(__cplusplus) +static_assert(sizeof(struct htp_set_rows_kernel_params) <= 128, "htp_set_rows_kernel_params is too large for kernel_params blob"); +#else +_Static_assert(sizeof(struct htp_set_rows_kernel_params) <= 128, "htp_set_rows_kernel_params is too large for kernel_params blob"); +#endif + +#endif // HTP_SET_ROWS_OPS_H From a5db1d66f0d32e7b5031d51f00797573f3ac7e02 Mon Sep 17 00:00:00 2001 From: Niklas Wenzel Date: Thu, 27 Aug 2026 11:53:08 +0200 Subject: [PATCH 011/104] metal : fix memory leaks due to missing autoreleasepools (llama/27758) --- ggml/src/ggml-metal/ggml-metal-context.m | 140 ++++---- ggml/src/ggml-metal/ggml-metal-device.m | 406 ++++++++++++----------- 2 files changed, 276 insertions(+), 270 deletions(-) diff --git a/ggml/src/ggml-metal/ggml-metal-context.m b/ggml/src/ggml-metal/ggml-metal-context.m index 32d97cd5d0a..1227ed39a09 100644 --- a/ggml/src/ggml-metal/ggml-metal-context.m +++ b/ggml/src/ggml-metal/ggml-metal-context.m @@ -84,106 +84,108 @@ ggml_metal_t ggml_metal_init(ggml_metal_device_t dev) { GGML_LOG_INFO("%s: allocating\n", __func__); + @autoreleasepool { #if TARGET_OS_OSX && !GGML_METAL_NDEBUG - // Show all the Metal device instances in the system - NSArray * devices = MTLCopyAllDevices(); - for (id device in devices) { - GGML_LOG_INFO("%s: found device: %s\n", __func__, [[device name] UTF8String]); - } - [devices release]; // since it was created by a *Copy* C method + // Show all the Metal device instances in the system + NSArray * devices = MTLCopyAllDevices(); + for (id device in devices) { + GGML_LOG_INFO("%s: found device: %s\n", __func__, [[device name] UTF8String]); + } + [devices release]; // since it was created by a *Copy* C method #endif - // init context - ggml_metal_t res = calloc(1, sizeof(struct ggml_metal)); - - id device = ggml_metal_device_get_obj(dev); + // init context + ggml_metal_t res = calloc(1, sizeof(struct ggml_metal)); - GGML_LOG_INFO("%s: picking default device: %s\n", __func__, [[device name] UTF8String]); + id device = ggml_metal_device_get_obj(dev); - // TODO: would it be better to have one queue for the backend and one queue for the device? - // the graph encoders and async ops would use the backend queue while the sync ops would use the device queue? - //res->queue = [device newCommandQueue]; [TAG_QUEUE_PER_BACKEND] - id queue = ggml_metal_device_get_queue(dev); - if (queue == nil) { - GGML_LOG_ERROR("%s: error: failed to create command queue\n", __func__); - return NULL; - } + GGML_LOG_INFO("%s: picking default device: %s\n", __func__, [[device name] UTF8String]); - res->dev = dev; - res->lib = ggml_metal_device_get_library(dev); - if (res->lib == NULL) { - GGML_LOG_WARN("%s: the device does not have a precompiled Metal library - this is unexpected\n", __func__); - GGML_LOG_WARN("%s: will try to compile it on the fly\n", __func__); + // TODO: would it be better to have one queue for the backend and one queue for the device? + // the graph encoders and async ops would use the backend queue while the sync ops would use the device queue? + //res->queue = [device newCommandQueue]; [TAG_QUEUE_PER_BACKEND] + id queue = ggml_metal_device_get_queue(dev); + if (queue == nil) { + GGML_LOG_ERROR("%s: error: failed to create command queue\n", __func__); + return NULL; + } - res->lib = ggml_metal_library_init(dev); + res->dev = dev; + res->lib = ggml_metal_device_get_library(dev); if (res->lib == NULL) { - GGML_LOG_ERROR("%s: error: failed to initialize the Metal library\n", __func__); + GGML_LOG_WARN("%s: the device does not have a precompiled Metal library - this is unexpected\n", __func__); + GGML_LOG_WARN("%s: will try to compile it on the fly\n", __func__); - free(res); + res->lib = ggml_metal_library_init(dev); + if (res->lib == NULL) { + GGML_LOG_ERROR("%s: error: failed to initialize the Metal library\n", __func__); - return NULL; + free(res); + + return NULL; + } } - } - res->ev_cpy = ggml_metal_device_event_init(dev); + res->ev_cpy = ggml_metal_device_event_init(dev); - const struct ggml_metal_device_props * props_dev = ggml_metal_device_get_props(dev); + const struct ggml_metal_device_props * props_dev = ggml_metal_device_get_props(dev); - snprintf(res->name, sizeof(res->name), "%s", props_dev->name); + snprintf(res->name, sizeof(res->name), "%s", props_dev->name); - res->d_queue = dispatch_queue_create("ggml-metal", DISPATCH_QUEUE_CONCURRENT); + res->d_queue = dispatch_queue_create("ggml-metal", DISPATCH_QUEUE_CONCURRENT); - res->use_fusion = getenv("GGML_METAL_FUSION_DISABLE") == nil; - res->use_concurrency = getenv("GGML_METAL_CONCURRENCY_DISABLE") == nil; + res->use_fusion = getenv("GGML_METAL_FUSION_DISABLE") == nil; + res->use_concurrency = getenv("GGML_METAL_CONCURRENCY_DISABLE") == nil; - { - const char * val = getenv("GGML_METAL_GRAPH_DEBUG"); - res->debug_graph = val ? atoi(val) : 0; - } + { + const char * val = getenv("GGML_METAL_GRAPH_DEBUG"); + res->debug_graph = val ? atoi(val) : 0; + } - { - const char * val = getenv("GGML_METAL_FUSION_DEBUG"); - res->debug_fusion = val ? atoi(val) : 0; - } + { + const char * val = getenv("GGML_METAL_FUSION_DEBUG"); + res->debug_fusion = val ? atoi(val) : 0; + } - res->use_graph_optimize = true; + res->use_graph_optimize = true; - if (getenv("GGML_METAL_GRAPH_OPTIMIZE_DISABLE") != NULL) { - res->use_graph_optimize = false; - } + if (getenv("GGML_METAL_GRAPH_OPTIMIZE_DISABLE") != NULL) { + res->use_graph_optimize = false; + } - memset(res->fuse_cnt, 0, sizeof(res->fuse_cnt)); + memset(res->fuse_cnt, 0, sizeof(res->fuse_cnt)); - GGML_LOG_INFO("%s: use fusion = %s\n", __func__, res->use_fusion ? "true" : "false"); - GGML_LOG_INFO("%s: use concurrency = %s\n", __func__, res->use_concurrency ? "true" : "false"); - GGML_LOG_INFO("%s: use graph optimize = %s\n", __func__, res->use_graph_optimize ? "true" : "false"); + GGML_LOG_INFO("%s: use fusion = %s\n", __func__, res->use_fusion ? "true" : "false"); + GGML_LOG_INFO("%s: use concurrency = %s\n", __func__, res->use_concurrency ? "true" : "false"); + GGML_LOG_INFO("%s: use graph optimize = %s\n", __func__, res->use_graph_optimize ? "true" : "false"); - res->capture_compute = 0; - res->capture_started = false; - res->capture_scope = nil; + res->capture_compute = 0; + res->capture_started = false; + res->capture_scope = nil; - { - const char * val = getenv("GGML_METAL_CAPTURE_COMPUTE"); - if (val) { - res->capture_compute = atoi(val); + { + const char * val = getenv("GGML_METAL_CAPTURE_COMPUTE"); + if (val) { + res->capture_compute = atoi(val); + } } - } - res->has_error = false; + res->has_error = false; - res->gf = nil; - res->encode_async = nil; - for (int i = 0; i < GGML_METAL_MAX_COMMAND_BUFFERS; ++i) { - res->cmd_bufs[i].obj = nil; - } + res->gf = nil; + res->encode_async = nil; + for (int i = 0; i < GGML_METAL_MAX_COMMAND_BUFFERS; ++i) { + res->cmd_bufs[i].obj = nil; + } - res->cmd_bufs_ext = [[NSMutableArray alloc] init]; + res->cmd_bufs_ext = [[NSMutableArray alloc] init]; - res->cmd_buf_last = nil; + res->cmd_buf_last = nil; - res->pipelines_ext = ggml_metal_pipelines_init(); + res->pipelines_ext = ggml_metal_pipelines_init(); - return res; + return res; + } } void ggml_metal_free(ggml_metal_t ctx) { diff --git a/ggml/src/ggml-metal/ggml-metal-device.m b/ggml/src/ggml-metal/ggml-metal-device.m index 19c57820e85..41ce90dc8a9 100644 --- a/ggml/src/ggml-metal/ggml-metal-device.m +++ b/ggml/src/ggml-metal/ggml-metal-device.m @@ -778,7 +778,9 @@ void ggml_metal_encoder_free(ggml_metal_encoder_t encoder) { } void ggml_metal_encoder_debug_group_push(ggml_metal_encoder_t encoder, const char * name) { - [encoder->obj pushDebugGroup:[NSString stringWithCString:name encoding:NSUTF8StringEncoding]]; + @autoreleasepool { + [encoder->obj pushDebugGroup:[NSString stringWithCString:name encoding:NSUTF8StringEncoding]]; + } } void ggml_metal_encoder_debug_group_pop (ggml_metal_encoder_t encoder) { @@ -1023,249 +1025,251 @@ ggml_metal_device_t ggml_metal_device_init(int device, int n_devices) { assert(dev != NULL); - if (dev->mtl_device == nil) { - dev->mtl_device = MTLCreateSystemDefaultDevice(); - - if (dev->mtl_device) { - dev->mtl_queue = [dev->mtl_device newCommandQueue]; - if (dev->mtl_queue == nil) { - GGML_LOG_ERROR("%s: error: failed to create command queue\n", __func__); - } + @autoreleasepool { + if (dev->mtl_device == nil) { + dev->mtl_device = MTLCreateSystemDefaultDevice(); - dev->addr_virt = 0x000000400ULL; + if (dev->mtl_device) { + dev->mtl_queue = [dev->mtl_device newCommandQueue]; + if (dev->mtl_queue == nil) { + GGML_LOG_ERROR("%s: error: failed to create command queue\n", __func__); + } - dev->props.device = device; + dev->addr_virt = 0x000000400ULL; - // the Metal backend uses the system default device as the single physical device; - // additional (virtual) devices are emulated on top of it via GGML_METAL_DEVICES - dev->props.device_phys = 0; - dev->props.device_virt = device; + dev->props.device = device; - dev->props.has_simdgroup_reduction = [dev->mtl_device supportsFamily:MTLGPUFamilyApple7]; - dev->props.has_simdgroup_reduction |= [dev->mtl_device supportsFamily:MTLGPUFamilyMetal3_GGML]; + // the Metal backend uses the system default device as the single physical device; + // additional (virtual) devices are emulated on top of it via GGML_METAL_DEVICES + dev->props.device_phys = 0; + dev->props.device_virt = device; - dev->props.has_simdgroup_mm = [dev->mtl_device supportsFamily:MTLGPUFamilyApple7]; - dev->props.has_unified_memory = dev->mtl_device.hasUnifiedMemory; + dev->props.has_simdgroup_reduction = [dev->mtl_device supportsFamily:MTLGPUFamilyApple7]; + dev->props.has_simdgroup_reduction |= [dev->mtl_device supportsFamily:MTLGPUFamilyMetal3_GGML]; - dev->props.has_bfloat = [dev->mtl_device supportsFamily:MTLGPUFamilyMetal3_GGML]; - dev->props.has_bfloat |= [dev->mtl_device supportsFamily:MTLGPUFamilyApple6]; - if (getenv("GGML_METAL_BF16_DISABLE") != NULL) { - dev->props.has_bfloat = false; - } + dev->props.has_simdgroup_mm = [dev->mtl_device supportsFamily:MTLGPUFamilyApple7]; + dev->props.has_unified_memory = dev->mtl_device.hasUnifiedMemory; - dev->props.has_tensor = [dev->mtl_device supportsFamily:MTLGPUFamilyMetal4_GGML]; - if (getenv("GGML_METAL_TENSOR_DISABLE") != NULL) { - dev->props.has_tensor = false; - } + dev->props.has_bfloat = [dev->mtl_device supportsFamily:MTLGPUFamilyMetal3_GGML]; + dev->props.has_bfloat |= [dev->mtl_device supportsFamily:MTLGPUFamilyApple6]; + if (getenv("GGML_METAL_BF16_DISABLE") != NULL) { + dev->props.has_bfloat = false; + } - // note: disable the tensor API by default for old chips because with the current implementation it is not useful - // - M2 Ultra: ~5% slower - // - M4, M4 Max: no significant difference - // - // TODO: try to update the tensor API kernels to at least match the simdgroup performance - if (getenv("GGML_METAL_TENSOR_ENABLE") == NULL && - ![[dev->mtl_device name] containsString:@"M5"] && - ![[dev->mtl_device name] containsString:@"M6"] && - ![[dev->mtl_device name] containsString:@"A19"] && - ![[dev->mtl_device name] containsString:@"A20"]) { - GGML_LOG_INFO("%s: tensor API disabled for pre-M5 and pre-A19 devices\n", __func__); - dev->props.has_tensor = false; - } + dev->props.has_tensor = [dev->mtl_device supportsFamily:MTLGPUFamilyMetal4_GGML]; + if (getenv("GGML_METAL_TENSOR_DISABLE") != NULL) { + dev->props.has_tensor = false; + } - // double-check that the tensor API compiles - if (dev->props.has_tensor) { - const char * src_tensor_f16 = "\n" - "#include \n" - "#include \n" - "#include \n" - " \n" - "using namespace metal; \n" - "using namespace mpp::tensor_ops; \n" - " \n" - "kernel void dummy_kernel( \n" - " tensor> A [[buffer(0)]], \n" - " tensor> B [[buffer(1)]], \n" - " device float * C [[buffer(2)]], \n" - " uint2 tgid [[threadgroup_position_in_grid]]) \n" - "{ \n" - " auto tA = A.slice(0, (int)tgid.y); \n" - " auto tB = B.slice((int)tgid.x, 0); \n" - " \n" - " matmul2d< \n" - " matmul2d_descriptor(16, 16, dynamic_extent), \n" - " execution_simdgroups<4>> mm; \n" - " \n" - " auto cT = mm.get_destination_cooperative_tensor(); \n" - " \n" - " auto sA = tA.slice(0, 0); \n" - " auto sB = tB.slice(0, 0); \n" - " mm.run(sB, sA, cT); \n" - " \n" - " auto tC = tensor, tensor_inline>(C, dextents(16, 16)); \n" - " \n" - " cT.store(tC); \n" - "}"; - - GGML_LOG_INFO("%s: testing tensor API for f16 support\n", __func__); - ggml_metal_library_t lib = ggml_metal_library_init_from_source(dev, src_tensor_f16, false); - if (lib == NULL) { - GGML_LOG_WARN("%s: - the tensor API is not supported in this environment - disabling\n", __func__); + // note: disable the tensor API by default for old chips because with the current implementation it is not useful + // - M2 Ultra: ~5% slower + // - M4, M4 Max: no significant difference + // + // TODO: try to update the tensor API kernels to at least match the simdgroup performance + if (getenv("GGML_METAL_TENSOR_ENABLE") == NULL && + ![[dev->mtl_device name] containsString:@"M5"] && + ![[dev->mtl_device name] containsString:@"M6"] && + ![[dev->mtl_device name] containsString:@"A19"] && + ![[dev->mtl_device name] containsString:@"A20"]) { + GGML_LOG_INFO("%s: tensor API disabled for pre-M5 and pre-A19 devices\n", __func__); dev->props.has_tensor = false; - } else { - struct ggml_metal_pipeline_with_params ppl = ggml_metal_library_compile_pipeline(lib, "dummy_kernel", "dummy_kernel", nil); - if (!ppl.pipeline) { + } + + // double-check that the tensor API compiles + if (dev->props.has_tensor) { + const char * src_tensor_f16 = "\n" + "#include \n" + "#include \n" + "#include \n" + " \n" + "using namespace metal; \n" + "using namespace mpp::tensor_ops; \n" + " \n" + "kernel void dummy_kernel( \n" + " tensor> A [[buffer(0)]], \n" + " tensor> B [[buffer(1)]], \n" + " device float * C [[buffer(2)]], \n" + " uint2 tgid [[threadgroup_position_in_grid]]) \n" + "{ \n" + " auto tA = A.slice(0, (int)tgid.y); \n" + " auto tB = B.slice((int)tgid.x, 0); \n" + " \n" + " matmul2d< \n" + " matmul2d_descriptor(16, 16, dynamic_extent), \n" + " execution_simdgroups<4>> mm; \n" + " \n" + " auto cT = mm.get_destination_cooperative_tensor(); \n" + " \n" + " auto sA = tA.slice(0, 0); \n" + " auto sB = tB.slice(0, 0); \n" + " mm.run(sB, sA, cT); \n" + " \n" + " auto tC = tensor, tensor_inline>(C, dextents(16, 16)); \n" + " \n" + " cT.store(tC); \n" + "}"; + + GGML_LOG_INFO("%s: testing tensor API for f16 support\n", __func__); + ggml_metal_library_t lib = ggml_metal_library_init_from_source(dev, src_tensor_f16, false); + if (lib == NULL) { GGML_LOG_WARN("%s: - the tensor API is not supported in this environment - disabling\n", __func__); dev->props.has_tensor = false; - } + } else { + struct ggml_metal_pipeline_with_params ppl = ggml_metal_library_compile_pipeline(lib, "dummy_kernel", "dummy_kernel", nil); + if (!ppl.pipeline) { + GGML_LOG_WARN("%s: - the tensor API is not supported in this environment - disabling\n", __func__); + dev->props.has_tensor = false; + } - ggml_metal_library_free(lib); + ggml_metal_library_free(lib); + } } - } - // try to compile a dummy kernel to determine if the tensor API is supported for bfloat - if (dev->props.has_tensor && dev->props.has_bfloat) { - const char * src_tensor_bf16 = "\n" - "#include \n" - "#include \n" - "#include \n" - " \n" - "using namespace metal; \n" - "using namespace mpp::tensor_ops; \n" - " \n" - "kernel void dummy_kernel( \n" - " tensor> A [[buffer(0)]], \n" - " tensor> B [[buffer(1)]], \n" - " device float * C [[buffer(2)]], \n" - " uint2 tgid [[threadgroup_position_in_grid]]) \n" - "{ \n" - " auto tA = A.slice(0, (int)tgid.y); \n" - " auto tB = B.slice((int)tgid.x, 0); \n" - " \n" - " matmul2d< \n" - " matmul2d_descriptor(16, 16, dynamic_extent), \n" - " execution_simdgroups<4>> mm; \n" - " \n" - " auto cT = mm.get_destination_cooperative_tensor(); \n" - " \n" - " auto sA = tA.slice(0, 0); \n" - " auto sB = tB.slice(0, 0); \n" - " mm.run(sB, sA, cT); \n" - " \n" - " auto tC = tensor, tensor_inline>(C, dextents(16, 16)); \n" - " \n" - " cT.store(tC); \n" - "}"; - - GGML_LOG_INFO("%s: testing tensor API for bfloat support\n", __func__); - ggml_metal_library_t lib = ggml_metal_library_init_from_source(dev, src_tensor_bf16, false); - if (lib == NULL) { - GGML_LOG_WARN("%s: - the tensor API does not support bfloat - disabling bfloat support\n", __func__); - dev->props.has_bfloat = false; - } else { - struct ggml_metal_pipeline_with_params ppl = ggml_metal_library_compile_pipeline(lib, "dummy_kernel", "dummy_kernel", nil); - if (!ppl.pipeline) { + // try to compile a dummy kernel to determine if the tensor API is supported for bfloat + if (dev->props.has_tensor && dev->props.has_bfloat) { + const char * src_tensor_bf16 = "\n" + "#include \n" + "#include \n" + "#include \n" + " \n" + "using namespace metal; \n" + "using namespace mpp::tensor_ops; \n" + " \n" + "kernel void dummy_kernel( \n" + " tensor> A [[buffer(0)]], \n" + " tensor> B [[buffer(1)]], \n" + " device float * C [[buffer(2)]], \n" + " uint2 tgid [[threadgroup_position_in_grid]]) \n" + "{ \n" + " auto tA = A.slice(0, (int)tgid.y); \n" + " auto tB = B.slice((int)tgid.x, 0); \n" + " \n" + " matmul2d< \n" + " matmul2d_descriptor(16, 16, dynamic_extent), \n" + " execution_simdgroups<4>> mm; \n" + " \n" + " auto cT = mm.get_destination_cooperative_tensor(); \n" + " \n" + " auto sA = tA.slice(0, 0); \n" + " auto sB = tB.slice(0, 0); \n" + " mm.run(sB, sA, cT); \n" + " \n" + " auto tC = tensor, tensor_inline>(C, dextents(16, 16)); \n" + " \n" + " cT.store(tC); \n" + "}"; + + GGML_LOG_INFO("%s: testing tensor API for bfloat support\n", __func__); + ggml_metal_library_t lib = ggml_metal_library_init_from_source(dev, src_tensor_bf16, false); + if (lib == NULL) { GGML_LOG_WARN("%s: - the tensor API does not support bfloat - disabling bfloat support\n", __func__); dev->props.has_bfloat = false; - } + } else { + struct ggml_metal_pipeline_with_params ppl = ggml_metal_library_compile_pipeline(lib, "dummy_kernel", "dummy_kernel", nil); + if (!ppl.pipeline) { + GGML_LOG_WARN("%s: - the tensor API does not support bfloat - disabling bfloat support\n", __func__); + dev->props.has_bfloat = false; + } - ggml_metal_library_free(lib); + ggml_metal_library_free(lib); + } } - } - dev->props.use_residency_sets = true; + dev->props.use_residency_sets = true; #if defined(GGML_METAL_HAS_RESIDENCY_SETS) - dev->props.use_residency_sets = getenv("GGML_METAL_NO_RESIDENCY") == nil; + dev->props.use_residency_sets = getenv("GGML_METAL_NO_RESIDENCY") == nil; #endif - dev->props.use_shared_buffers = dev->props.has_unified_memory; + dev->props.use_shared_buffers = dev->props.has_unified_memory; #if TARGET_OS_OSX - // In case of eGPU, shared memory may be preferable. - dev->props.use_shared_buffers |= [dev->mtl_device location] == MTLDeviceLocationExternal; + // In case of eGPU, shared memory may be preferable. + dev->props.use_shared_buffers |= [dev->mtl_device location] == MTLDeviceLocationExternal; #endif - if (getenv("GGML_METAL_SHARED_BUFFERS_DISABLE") != NULL) { - dev->props.use_shared_buffers = false; - } - if (getenv("GGML_METAL_SHARED_BUFFERS_ENABLE") != NULL) { - dev->props.use_shared_buffers = true; - } + if (getenv("GGML_METAL_SHARED_BUFFERS_DISABLE") != NULL) { + dev->props.use_shared_buffers = false; + } + if (getenv("GGML_METAL_SHARED_BUFFERS_ENABLE") != NULL) { + dev->props.use_shared_buffers = true; + } - dev->props.supports_gpu_family_apple7 = [dev->mtl_device supportsFamily:MTLGPUFamilyApple7]; + 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.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; + dev->props.op_offload_min_batch_size = getenv("GGML_OP_OFFLOAD_MIN_BATCH") ? atoi(getenv("GGML_OP_OFFLOAD_MIN_BATCH")) : 32; - dev->props.max_buffer_size = dev->mtl_device.maxBufferLength; - dev->props.max_theadgroup_memory_size = dev->mtl_device.maxThreadgroupMemoryLength; - if (@available(macOS 10.12, iOS 16.0, *)) { - dev->props.max_working_set_size = dev->mtl_device.recommendedMaxWorkingSetSize; - } else { - dev->props.max_working_set_size = dev->mtl_device.maxBufferLength; - } - - snprintf(dev->props.name, sizeof(dev->props.name), "%s%d", "MTL", device); - const char * gpu_name = [[dev->mtl_device name] UTF8String]; - if (n_devices > 1) { - snprintf(dev->props.desc, sizeof(dev->props.desc), "%s (dev p%d/v%d)", - gpu_name, dev->props.device_phys, dev->props.device_virt); - } else { - snprintf(dev->props.desc, sizeof(dev->props.desc), "%s", gpu_name); - } + dev->props.max_buffer_size = dev->mtl_device.maxBufferLength; + dev->props.max_theadgroup_memory_size = dev->mtl_device.maxThreadgroupMemoryLength; + if (@available(macOS 10.12, iOS 16.0, *)) { + dev->props.max_working_set_size = dev->mtl_device.recommendedMaxWorkingSetSize; + } else { + dev->props.max_working_set_size = dev->mtl_device.maxBufferLength; + } - dev->library = ggml_metal_library_init(dev); - if (!dev->library) { - GGML_LOG_ERROR("%s: error: failed to create library\n", __func__); - } + snprintf(dev->props.name, sizeof(dev->props.name), "%s%d", "MTL", device); + const char * gpu_name = [[dev->mtl_device name] UTF8String]; + if (n_devices > 1) { + snprintf(dev->props.desc, sizeof(dev->props.desc), "%s (dev p%d/v%d)", + gpu_name, dev->props.device_phys, dev->props.device_virt); + } else { + snprintf(dev->props.desc, sizeof(dev->props.desc), "%s", gpu_name); + } - if (dev->props.use_residency_sets) { - dev->rsets = ggml_metal_rsets_init(dev); - } else { - dev->rsets = nil; - } + dev->library = ggml_metal_library_init(dev); + if (!dev->library) { + GGML_LOG_ERROR("%s: error: failed to create library\n", __func__); + } - // print MTL GPU family: - GGML_LOG_INFO("%s: GPU name: %s (%s)\n", __func__, dev->props.name, dev->props.desc); + if (dev->props.use_residency_sets) { + dev->rsets = ggml_metal_rsets_init(dev); + } else { + dev->rsets = nil; + } - // determine max supported GPU family - // https://developer.apple.com/metal/Metal-Shading-Language-Specification.pdf - // https://developer.apple.com/metal/Metal-Feature-Set-Tables.pdf - { - for (int i = MTLGPUFamilyApple1 + 20; i >= MTLGPUFamilyApple1; --i) { - if ([dev->mtl_device supportsFamily:i]) { - dev->props.gpu_family = i - (int) MTLGPUFamilyApple1 + 1; - GGML_LOG_INFO("%s: GPU family: MTLGPUFamilyApple%d (%d)\n", __func__, dev->props.gpu_family, i); - break; + // print MTL GPU family: + GGML_LOG_INFO("%s: GPU name: %s (%s)\n", __func__, dev->props.name, dev->props.desc); + + // determine max supported GPU family + // https://developer.apple.com/metal/Metal-Shading-Language-Specification.pdf + // https://developer.apple.com/metal/Metal-Feature-Set-Tables.pdf + { + for (int i = MTLGPUFamilyApple1 + 20; i >= MTLGPUFamilyApple1; --i) { + if ([dev->mtl_device supportsFamily:i]) { + dev->props.gpu_family = i - (int) MTLGPUFamilyApple1 + 1; + GGML_LOG_INFO("%s: GPU family: MTLGPUFamilyApple%d (%d)\n", __func__, dev->props.gpu_family, i); + break; + } } - } - for (int i = MTLGPUFamilyCommon1 + 5; i >= MTLGPUFamilyCommon1; --i) { - if ([dev->mtl_device supportsFamily:i]) { - GGML_LOG_INFO("%s: GPU family: MTLGPUFamilyCommon%d (%d)\n", __func__, i - (int) MTLGPUFamilyCommon1 + 1, i); - break; + for (int i = MTLGPUFamilyCommon1 + 5; i >= MTLGPUFamilyCommon1; --i) { + if ([dev->mtl_device supportsFamily:i]) { + GGML_LOG_INFO("%s: GPU family: MTLGPUFamilyCommon%d (%d)\n", __func__, i - (int) MTLGPUFamilyCommon1 + 1, i); + break; + } } - } - for (int i = MTLGPUFamilyMetal3_GGML + 5; i >= MTLGPUFamilyMetal3_GGML; --i) { - if ([dev->mtl_device supportsFamily:i]) { - GGML_LOG_INFO("%s: GPU family: MTLGPUFamilyMetal%d (%d)\n", __func__, i - (int) MTLGPUFamilyMetal3_GGML + 3, i); - break; + for (int i = MTLGPUFamilyMetal3_GGML + 5; i >= MTLGPUFamilyMetal3_GGML; --i) { + if ([dev->mtl_device supportsFamily:i]) { + GGML_LOG_INFO("%s: GPU family: MTLGPUFamilyMetal%d (%d)\n", __func__, i - (int) MTLGPUFamilyMetal3_GGML + 3, i); + break; + } } } - } - GGML_LOG_INFO("%s: simdgroup reduction = %s\n", __func__, dev->props.has_simdgroup_reduction ? "true" : "false"); - GGML_LOG_INFO("%s: simdgroup matrix mul. = %s\n", __func__, dev->props.has_simdgroup_mm ? "true" : "false"); - GGML_LOG_INFO("%s: has unified memory = %s\n", __func__, dev->props.has_unified_memory ? "true" : "false"); - GGML_LOG_INFO("%s: has bfloat = %s\n", __func__, dev->props.has_bfloat ? "true" : "false"); - GGML_LOG_INFO("%s: has tensor = %s\n", __func__, dev->props.has_tensor ? "true" : "false"); - GGML_LOG_INFO("%s: use residency sets = %s\n", __func__, dev->props.use_residency_sets ? "true" : "false"); - GGML_LOG_INFO("%s: use shared buffers = %s\n", __func__, dev->props.use_shared_buffers ? "true" : "false"); + GGML_LOG_INFO("%s: simdgroup reduction = %s\n", __func__, dev->props.has_simdgroup_reduction ? "true" : "false"); + GGML_LOG_INFO("%s: simdgroup matrix mul. = %s\n", __func__, dev->props.has_simdgroup_mm ? "true" : "false"); + GGML_LOG_INFO("%s: has unified memory = %s\n", __func__, dev->props.has_unified_memory ? "true" : "false"); + GGML_LOG_INFO("%s: has bfloat = %s\n", __func__, dev->props.has_bfloat ? "true" : "false"); + GGML_LOG_INFO("%s: has tensor = %s\n", __func__, dev->props.has_tensor ? "true" : "false"); + GGML_LOG_INFO("%s: use residency sets = %s\n", __func__, dev->props.use_residency_sets ? "true" : "false"); + GGML_LOG_INFO("%s: use shared buffers = %s\n", __func__, dev->props.use_shared_buffers ? "true" : "false"); #if TARGET_OS_OSX || (TARGET_OS_IOS && __clang_major__ >= 15) - if (@available(macOS 10.12, iOS 16.0, *)) { - GGML_LOG_INFO("%s: recommendedMaxWorkingSetSize = %8.2f MB\n", __func__, dev->props.max_working_set_size / 1e6); - } + if (@available(macOS 10.12, iOS 16.0, *)) { + GGML_LOG_INFO("%s: recommendedMaxWorkingSetSize = %8.2f MB\n", __func__, dev->props.max_working_set_size / 1e6); + } #endif + } } } From 55ab1e517cf52c9441119f740bc0e9d758a4f4ec Mon Sep 17 00:00:00 2001 From: Shobhit Date: Thu, 27 Aug 2026 21:34:42 +0800 Subject: [PATCH 012/104] Feature: Added LIGHTNING_INDEXER support for Deepseek V4 ops on Vulkan Backend (llama/27453) * vulkan: add LIGHTNING_INDEXER op * vulkan: updated lightning_indexer.comp and ggml-vulkan.cpp with 128-lane dot-product reduction moved from a shared-memory tree to subgroupAdd. * vulkan: cleanup; Skip bounds checks * vulkan: cleanup FA_K_ONLY * Revert "vulkan: cleanup FA_K_ONLY" This reverts commit fdcbdd91511945d6878d9070b5447e4a34dce010. * vulkan: restore interleaved K/V buffer ordering * vulkan: Remove FA_K_ONLY * vulkan: Revert flash_attn_dequant * vulkan: Revert tests in backend-ops.cpp --- ggml/src/ggml-vulkan/ggml-vulkan.cpp | 156 +++++++++++++++++- .../ggml-vulkan/vulkan-shaders/fa_types.glsl | 55 ++++++ .../vulkan-shaders/flash_attn_base.glsl | 51 +----- .../vulkan-shaders/lightning_indexer.comp | 151 +++++++++++++++++ .../vulkan-shaders/vulkan-shaders-gen.cpp | 6 + 5 files changed, 365 insertions(+), 54 deletions(-) create mode 100644 ggml/src/ggml-vulkan/vulkan-shaders/fa_types.glsl create mode 100644 ggml/src/ggml-vulkan/vulkan-shaders/lightning_indexer.comp diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp index 8108e94c16c..72e844aebfd 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp +++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp @@ -767,6 +767,21 @@ static constexpr std::initializer_list> rms_norm_mul_rope_vie { 4, 0, 3 }, // set_rows->src[0] == view }; +static constexpr std::array lightning_indexer_k_types = { + GGML_TYPE_F32, + GGML_TYPE_F16, + GGML_TYPE_BF16, + GGML_TYPE_Q8_0, + GGML_TYPE_Q5_1, + GGML_TYPE_Q5_0, + GGML_TYPE_Q4_1, + GGML_TYPE_Q4_0, + GGML_TYPE_IQ4_NL, +}; + +static bool ggml_vk_lightning_indexer_k_type_supported(ggml_type type) { + return std::find(lightning_indexer_k_types.begin(), lightning_indexer_k_types.end(), type) != lightning_indexer_k_types.end(); +} struct vk_device_struct { std::recursive_mutex mutex; @@ -1068,6 +1083,7 @@ struct vk_device_struct { vk_pipeline pipeline_rwkv_wkv6_f32; vk_pipeline pipeline_rwkv_wkv7_f32; vk_pipeline pipeline_gated_linear_attn_f32; + vk_pipeline pipeline_lightning_indexer_f32[GGML_TYPE_COUNT]; // [size_idx][kda] where size_idx: 0=d16, 1=d32, 2=d64, 3=d128 vk_pipeline pipeline_gated_delta_net[4][2]; vk_pipeline pipeline_ssm_scan_f32_d128; @@ -1848,6 +1864,26 @@ struct vk_op_gated_linear_attn_push_constants { uint32_t H; float scale; }; +struct vk_op_lightning_indexer_push_constants { + uint32_t n_kv; + uint32_t n_heads; + uint32_t n_tokens; + uint32_t n_streams; + uint32_t n_masks; + uint32_t dispatch_x; + uint32_t q_nb1; + uint32_t q_nb2; + uint32_t q_nb3; + uint32_t k_nb2; + uint32_t k_nb3; + uint32_t w_nb1; + uint32_t w_nb3; + uint32_t m_nb1; + uint32_t m_nb3; + uint32_t d_nb1; + uint32_t d_nb3; +}; +static_assert(sizeof(vk_op_lightning_indexer_push_constants) <= 128); struct vk_op_gated_delta_net_push_constants { uint32_t H; uint32_t n_tokens; @@ -3904,11 +3940,16 @@ static vk_fa_pipeline_state get_fa_pipeline_state(const vk_device& device, const return vk_fa_pipeline_state{hsk, hsv, params.block_rows, params.block_cols, params.d_split, params.row_split, params.shmem_staging, params.path, params.workgroup_size, subgroup_size, aligned, f32acc, flags, params.limit_occupancy_shmem, k_type, v_type}; } +// Bytes per buffer block for the FaBlockBytesK/V spec constants. F32 is fed as +// a vec4 "block" of 4 floats, everything else uses its ggml block size. +static uint32_t fa_block_bytes(ggml_type t) { + if (t == GGML_TYPE_F32) { + return 16u; + } + return (uint32_t) ggml_type_size(t); +} + static std::vector get_fa_spec_constants(const vk_fa_pipeline_state& state) { - const auto fa_block_bytes = [](ggml_type t) -> uint32_t { - if (t == GGML_TYPE_F32) return 16u; - return (uint32_t) ggml_type_size(t); - }; return { /* 0 WorkGroupSize */ state.workgroup_size, /* 1 Br */ state.Br, @@ -5847,6 +5888,17 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { ggml_vk_create_pipeline(device, device->pipeline_gated_linear_attn_f32, "gated_linear_attn_f32", gated_linear_attn_f32_len, gated_linear_attn_f32_data, "main", 6, sizeof(vk_op_gated_linear_attn_push_constants), {1, 1, 1}, {}, 1); + { + const bool li_subgroup = device->subgroup_arithmetic && device->subgroup_require_full_support; + const size_t li_len = li_subgroup ? lightning_indexer_subgroup_f32_len : lightning_indexer_f32_len; + const void * li_data = li_subgroup ? (const void *)lightning_indexer_subgroup_f32_data : (const void *)lightning_indexer_f32_data; + + for (ggml_type k_type : lightning_indexer_k_types) { + const std::string name = "lightning_indexer_" + std::string(ggml_type_name(k_type)) + "_k_f32"; + ggml_vk_create_pipeline(device, device->pipeline_lightning_indexer_f32[k_type], name.c_str(), li_len, li_data, "main", 5, sizeof(vk_op_lightning_indexer_push_constants), {1, 1, 1}, {(uint32_t)k_type, fa_block_bytes(k_type), device->subgroup_size}, 1, true, li_subgroup); + } + } + { const uint32_t gdn_sizes[] = {16, 32, 64, 128}; const char * gdn_names[][2] = { @@ -11697,6 +11749,12 @@ static vk_pipeline ggml_vk_op_get_pipeline(ggml_backend_vk_context * ctx, const return ctx->device->pipeline_gated_linear_attn_f32; } return nullptr; + case GGML_OP_LIGHTNING_INDEXER: + // only the k type selects a pipeline, the other types are fixed by ggml_lightning_indexer() + if (ggml_vk_lightning_indexer_k_type_supported(src1->type)) { + return ctx->device->pipeline_lightning_indexer_f32[src1->type]; + } + return nullptr; case GGML_OP_GATED_DELTA_NET: if (src0->type == GGML_TYPE_F32 && dst->type == GGML_TYPE_F32) { const uint32_t S_v = dst->src[2]->ne[0]; @@ -12772,6 +12830,55 @@ static void ggml_vk_gated_linear_attn(ggml_backend_vk_context * ctx, vk_context& pc, { (uint32_t)(n_seqs * n_heads), 1, 1 }); } +static void ggml_vk_lightning_indexer(ggml_backend_vk_context * ctx, vk_context& subctx, ggml_tensor * dst) { + const ggml_tensor * q = dst->src[0]; + const ggml_tensor * k = dst->src[1]; + const ggml_tensor * w = dst->src[2]; + const ggml_tensor * m = dst->src[3]; + + vk_pipeline pipeline = ggml_vk_op_get_pipeline(ctx, q, k, w, dst, dst->op); + GGML_ASSERT(pipeline != nullptr); + + ggml_pipeline_request_descriptor_sets(ctx, pipeline, 1); + + const uint32_t n_kv = k->ne[2]; + const uint32_t n_heads = q->ne[1]; + const uint32_t n_tokens = q->ne[2]; + const uint32_t n_streams = q->ne[3]; + const uint32_t n_masks = m->ne[3]; + + const uint32_t n_outputs = (uint32_t)(dst->ne[0] * dst->ne[1] * dst->ne[3]); + const uint32_t dispatch_x = std::min(n_outputs, ctx->device->properties.limits.maxComputeWorkGroupCount[0]); + const uint32_t dispatch_y = CEIL_DIV(n_outputs, dispatch_x); + + // q, w and dst are f32 and m is f16, so their strides are passed in elements; + // k may be quantized, so its strides stay in bytes + const uint32_t q_nb1 = q->nb[1] / sizeof(float); + const uint32_t q_nb2 = q->nb[2] / sizeof(float); + const uint32_t q_nb3 = q->nb[3] / sizeof(float); + const uint32_t k_nb2 = k->nb[2]; + const uint32_t k_nb3 = k->nb[3]; + const uint32_t w_nb1 = w->nb[1] / sizeof(float); + const uint32_t w_nb3 = w->nb[3] / sizeof(float); + const uint32_t m_nb1 = m->nb[1] / sizeof(ggml_fp16_t); + const uint32_t m_nb3 = m->nb[3] / sizeof(ggml_fp16_t); + const uint32_t d_nb1 = dst->nb[1] / sizeof(float); + const uint32_t d_nb3 = dst->nb[3] / sizeof(float); + + const vk_op_lightning_indexer_push_constants pc = { + n_kv, n_heads, n_tokens, n_streams, n_masks, dispatch_x, + q_nb1, q_nb2, q_nb3, + k_nb2, k_nb3, + w_nb1, w_nb3, + m_nb1, m_nb3, + d_nb1, d_nb3, + }; + + ggml_vk_dispatch_pipeline(ctx, subctx, pipeline, + {ggml_vk_tensor_subbuffer(ctx, q), ggml_vk_tensor_subbuffer(ctx, k), ggml_vk_tensor_subbuffer(ctx, w), ggml_vk_tensor_subbuffer(ctx, m), ggml_vk_tensor_subbuffer(ctx, dst)}, + pc, {dispatch_x, dispatch_y, 1}); +} + static void ggml_vk_gated_delta_net(ggml_backend_vk_context * ctx, vk_context& subctx, ggml_tensor * dst) { const ggml_tensor * src_q = dst->src[0]; const ggml_tensor * src_v = dst->src[2]; @@ -15898,6 +16005,11 @@ static bool ggml_vk_build_graph(ggml_backend_vk_context * ctx, ggml_cgraph * cgr break; + case GGML_OP_LIGHTNING_INDEXER: + ggml_vk_lightning_indexer(ctx, compute_ctx, node); + + break; + case GGML_OP_GATED_DELTA_NET: ggml_vk_gated_delta_net(ctx, compute_ctx, node); @@ -18676,6 +18788,40 @@ static bool ggml_backend_vk_device_supports_op(ggml_backend_dev_t dev, const ggm case GGML_OP_GATED_LINEAR_ATTN: // the shader block size is hardcoded to head_size 64 return op->src[0]->type == GGML_TYPE_F32 && op->type == GGML_TYPE_F32 && op->src[0]->ne[0] == 64; + case GGML_OP_LIGHTNING_INDEXER: + { + const ggml_tensor * q = op->src[0]; + const ggml_tensor * k = op->src[1]; + const ggml_tensor * w = op->src[2]; + const ggml_tensor * m = op->src[3]; + + // the q/w/m types and the shape relationships between q, k, w, m and dst + // are already asserted in ggml_lightning_indexer() + if (!ggml_vk_lightning_indexer_k_type_supported(k->type) || !device->fp16) { + return false; + } + + // the shader block size is hardcoded to head size 128 + if (q->ne[0] != 128) { + return false; + } + + // the shader indexes the buffers by element stride, and is dispatched + // without allow_misalign + for (const ggml_tensor * t : {q, k, w, m, op}) { + if (t->nb[0] != ggml_type_size(t->type) || + (vk_tensor_offset(t) + t->view_offs) % device->properties.limits.minStorageBufferOffsetAlignment != 0) { + return false; + } + // the strides get scaled down from bytes, so the division must be exact + for (int i = 1; i < GGML_MAX_DIMS; ++i) { + if (t->nb[i] % ggml_type_size(t->type) != 0) { + return false; + } + } + } + return true; + } case GGML_OP_GATED_DELTA_NET: { const uint32_t S_v = op->src[2]->ne[0]; @@ -19685,6 +19831,8 @@ static void ggml_vk_check_results_0(ggml_backend_vk_context * ctx, ggml_cgraph * const float * op_params = (const float *)tensor->op_params; tensor_clone = ggml_gated_linear_attn(ggml_ctx, src_clone[0], src_clone[1], src_clone[2], src_clone[3], src_clone[4], op_params[0]); + } else if (tensor->op == GGML_OP_LIGHTNING_INDEXER) { + tensor_clone = ggml_lightning_indexer(ggml_ctx, src_clone[0], src_clone[1], src_clone[2], src_clone[3]); } else if (tensor->op == GGML_OP_GATED_DELTA_NET) { tensor_clone = ggml_gated_delta_net(ggml_ctx, src_clone[0], src_clone[1], src_clone[2], src_clone[3], src_clone[4], src_clone[5], diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/fa_types.glsl b/ggml/src/ggml-vulkan/vulkan-shaders/fa_types.glsl new file mode 100644 index 00000000000..6f414ded123 --- /dev/null +++ b/ggml/src/ggml-vulkan/vulkan-shaders/fa_types.glsl @@ -0,0 +1,55 @@ +#if !defined(GGML_FA_TYPES_COMP) +#define GGML_FA_TYPES_COMP + +// FaTypeK / FaTypeV spec constant values. These mirror enum ggml_type so the +// host can pass the type directly. Keep in sync with ggml.h. +#define FA_TYPE_F32 0u +#define FA_TYPE_F16 1u +#define FA_TYPE_Q4_0 2u +#define FA_TYPE_Q4_1 3u +#define FA_TYPE_Q5_0 6u +#define FA_TYPE_Q5_1 7u +#define FA_TYPE_Q8_0 8u +#define FA_TYPE_IQ4_NL 20u +#define FA_TYPE_BF16 30u + +// Number of matrix elements per buffer block, derived from the K/V type spec +// constant. F32 is treated as a vec4 "block" of 4 floats. F16 uses block size 1 +// and bypasses the dequant path entirely. Quants follow their ggml block sizes. +uint fa_block_elems(uint ty) { + switch (ty) { + case FA_TYPE_F32: return 4u; + case FA_TYPE_F16: return 1u; + case FA_TYPE_Q4_0: return uint(QUANT_K_Q4_0); + case FA_TYPE_Q4_1: return uint(QUANT_K_Q4_1); + case FA_TYPE_Q5_0: return uint(QUANT_K_Q5_0); + case FA_TYPE_Q5_1: return uint(QUANT_K_Q5_1); + case FA_TYPE_Q8_0: return uint(QUANT_K_Q8_0); + case FA_TYPE_IQ4_NL: return uint(QUANT_K_IQ4_NL); + case FA_TYPE_BF16: return 1u; + default: return 1u; + } +} + +// QUANT_R_MMQ for FA-eligible K types. Q4_*/Q5_* store two nibbles per byte +// (R==2); Q8_0 stores one byte per element (R==1). Used to derive the number +// of int32s per 32-element block on the MMQ K path: ints_per_block == 8 / R. +uint fa_quant_r_mmq(uint ty) { + switch (ty) { + case FA_TYPE_Q4_0: return uint(QUANT_R_Q4_0); + case FA_TYPE_Q4_1: return uint(QUANT_R_Q4_1); + case FA_TYPE_Q5_0: return uint(QUANT_R_Q5_0); + case FA_TYPE_Q5_1: return uint(QUANT_R_Q5_1); + case FA_TYPE_Q8_0: return uint(QUANT_R_Q8_0); + default: return 1u; + } +} + +bool fa_type_needs_shmem(uint ty) { + switch (ty) { + case FA_TYPE_IQ4_NL: return true; + default: return false; + } +} + +#endif // !defined(GGML_FA_TYPES_COMP) diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_base.glsl b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_base.glsl index 3c64f91dad3..0ce4503a884 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_base.glsl +++ b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_base.glsl @@ -88,17 +88,7 @@ layout (binding = 6) readonly buffer MO {uint32_t data_mask_opt[];}; #define BINDING_IDX_K 0 #define BINDING_IDX_V 1 -// FaTypeK / FaTypeV spec constant values. These mirror enum ggml_type so the -// host can pass the type directly. Keep in sync with ggml.h. -#define FA_TYPE_F32 0u -#define FA_TYPE_F16 1u -#define FA_TYPE_Q4_0 2u -#define FA_TYPE_Q4_1 3u -#define FA_TYPE_Q5_0 6u -#define FA_TYPE_Q5_1 7u -#define FA_TYPE_Q8_0 8u -#define FA_TYPE_IQ4_NL 20u -#define FA_TYPE_BF16 30u +#include "fa_types.glsl" #if defined(BFLOAT16) #define O_TYPE float @@ -108,45 +98,6 @@ layout (binding = 6) readonly buffer MO {uint32_t data_mask_opt[];}; #define O_TYPEV4 FLOAT_TYPEV4 #endif -// Number of matrix elements per buffer block, derived from the K/V type spec -// constant. F32 is treated as a vec4 "block" of 4 floats. F16 uses block size 1 -// and bypasses the dequant path entirely. Quants follow their ggml block sizes. -uint fa_block_elems(uint ty) { - switch (ty) { - case FA_TYPE_F32: return 4u; - case FA_TYPE_F16: return 1u; - case FA_TYPE_Q4_0: return uint(QUANT_K_Q4_0); - case FA_TYPE_Q4_1: return uint(QUANT_K_Q4_1); - case FA_TYPE_Q5_0: return uint(QUANT_K_Q5_0); - case FA_TYPE_Q5_1: return uint(QUANT_K_Q5_1); - case FA_TYPE_Q8_0: return uint(QUANT_K_Q8_0); - case FA_TYPE_IQ4_NL: return uint(QUANT_K_IQ4_NL); - case FA_TYPE_BF16: return 1u; - default: return 1u; - } -} - -// QUANT_R_MMQ for FA-eligible K types. Q4_*/Q5_* store two nibbles per byte -// (R==2); Q8_0 stores one byte per element (R==1). Used to derive the number -// of int32s per 32-element block on the MMQ K path: ints_per_block == 8 / R. -uint fa_quant_r_mmq(uint ty) { - switch (ty) { - case FA_TYPE_Q4_0: return uint(QUANT_R_Q4_0); - case FA_TYPE_Q4_1: return uint(QUANT_R_Q4_1); - case FA_TYPE_Q5_0: return uint(QUANT_R_Q5_0); - case FA_TYPE_Q5_1: return uint(QUANT_R_Q5_1); - case FA_TYPE_Q8_0: return uint(QUANT_R_Q8_0); - default: return 1u; - } -} - -bool fa_type_needs_shmem(uint ty) { - switch (ty) { - case FA_TYPE_IQ4_NL: return true; - default: return false; - } -} - // These can't be `const` globals because GLSL forbids function calls in global // const initializers, even when the spec constants would let the driver fold // them. Macros expand at the use site and fold after specialization. diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/lightning_indexer.comp b/ggml/src/ggml-vulkan/vulkan-shaders/lightning_indexer.comp new file mode 100644 index 00000000000..ba76ec72ca6 --- /dev/null +++ b/ggml/src/ggml-vulkan/vulkan-shaders/lightning_indexer.comp @@ -0,0 +1,151 @@ +#version 450 + +#extension GL_EXT_control_flow_attributes : require +#extension GL_EXT_shader_16bit_storage : require +#extension GL_EXT_shader_explicit_arithmetic_types_float16 : require +#extension GL_KHR_shader_subgroup_basic : enable +#if USE_SUBGROUP_ADD +#extension GL_KHR_shader_subgroup_arithmetic : enable +#endif + +#define BINDING_IDX_K 0u + +#include "types.glsl" +#include "fa_types.glsl" +#define FaTypeV FA_TYPE_F32 + +layout(constant_id = 0) const uint FaTypeK = FA_TYPE_F32; +layout(constant_id = 1) const uint FaBlockBytesK = 4; +layout(constant_id = 2) const uint SUBGROUP_SIZE = 32; + +#include "flash_attn_dequant.glsl" + +// one workgroup computes one output element, one invocation per head element +#define HEAD_SIZE 128 + +layout(local_size_x = HEAD_SIZE, local_size_y = 1, local_size_z = 1) in; + +layout(binding = 0) readonly buffer QBuf { float q[]; }; +layout(binding = 1) readonly buffer KBufF16 { float16_t k_f16[]; }; +layout(binding = 1) readonly buffer KBufF32 { float k_f32[]; }; +layout(binding = 1) readonly buffer KBufBF16 { uint16_t k_bf16[]; }; +layout(binding = 2) readonly buffer WBuf { float weights[]; }; +layout(binding = 3) readonly buffer MBuf { float16_t mask[]; }; +layout(binding = 4) writeonly buffer DstBuf { float dst[]; }; + +layout(push_constant) uniform PushConstants { + uint n_kv; + uint n_heads; + uint n_tokens; + uint n_streams; + uint n_masks; + uint dispatch_x; + uint q_nb1; + uint q_nb2; + uint q_nb3; + uint k_nb2; + uint k_nb3; + uint w_nb1; + uint w_nb3; + uint m_nb1; + uint m_nb3; + uint d_nb1; + uint d_nb3; +}; + +shared float k_row[HEAD_SIZE]; + +#if USE_SUBGROUP_ADD +shared float sg_partials[HEAD_SIZE / SUBGROUP_SIZE]; +#else +shared float partials[HEAD_SIZE]; +#endif + +void main() { + const uint tid = gl_LocalInvocationID.x; + const uint output_idx = gl_WorkGroupID.y * dispatch_x + gl_WorkGroupID.x; + const uint n_outputs = n_kv * n_tokens * n_streams; + + if (fa_type_needs_shmem(FaTypeK)) { + init_iq_shmem(gl_WorkGroupSize); + } + + if (output_idx >= n_outputs) { + return; + } + + const uint ik = output_idx % n_kv; + const uint ts = output_idx / n_kv; + const uint t = ts % n_tokens; + const uint s = ts / n_tokens; + const uint k_offset = ik * k_nb2 + s * k_nb3; + + // k strides come in as bytes, so scale them down to the view being indexed + const uint k_block_elems = fa_block_elems(FaTypeK); + const uint k_elem_bytes = FaBlockBytesK / k_block_elems; + + if (FaTypeK == FA_TYPE_F16) { + k_row[tid] = float(k_f16[k_offset / k_elem_bytes + tid]); + } else if (FaTypeK == FA_TYPE_F32) { + k_row[tid] = k_f32[k_offset / k_elem_bytes + tid]; + } else if (FaTypeK == FA_TYPE_BF16) { + k_row[tid] = bf16_to_fp32(uint(k_bf16[k_offset / k_elem_bytes + tid])); + } else if (4 * tid < HEAD_SIZE) { + const uint coord = 4 * tid; + const uint ib = coord / k_block_elems; + const uint iqs = coord % k_block_elems; + const vec4 values = dequantize4(ib, iqs, k_offset / FaBlockBytesK, BINDING_IDX_K); + k_row[coord + 0] = values.x; + k_row[coord + 1] = values.y; + k_row[coord + 2] = values.z; + k_row[coord + 3] = values.w; + } + barrier(); + + const float k_val = k_row[tid]; + + float score = 0.0; + for (uint h = 0; h < n_heads; ++h) { + const float prod = q[h * q_nb1 + t * q_nb2 + s * q_nb3 + tid] * k_val; + +#if USE_SUBGROUP_ADD + const float sg_sum = subgroupAdd(prod); + if (gl_SubgroupInvocationID == 0) { + sg_partials[gl_SubgroupID] = sg_sum; + } + barrier(); + + if (tid == 0) { + float sum = 0.0; + [[unroll]] for (uint i = 0; i < HEAD_SIZE / SUBGROUP_SIZE; ++i) { + sum += sg_partials[i]; + } + score += max(sum, 0.0) * weights[h + t * w_nb1 + s * w_nb3]; + } + // the reads above must complete before the next iteration overwrites sg_partials + barrier(); +#else + partials[tid] = prod; + barrier(); + + [[unroll]] for (uint stride = HEAD_SIZE / 2; stride > 0; stride >>= 1) { + if (tid < stride) { + partials[tid] += partials[tid + stride]; + } + barrier(); + } + + if (tid == 0) { + score += max(partials[0], 0.0) * weights[h + t * w_nb1 + s * w_nb3]; + } + // the read of partials[0] above must complete before the next iteration + // overwrites partials[tid] + barrier(); +#endif + } + + if (tid == 0) { + const uint mask_offset = ik + t * m_nb1 + (s % n_masks) * m_nb3; + dst[ik + t * d_nb1 + s * d_nb3] = score + float(mask[mask_offset]); + } +} diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp b/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp index dbb99782cf7..0da943da956 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp @@ -1069,6 +1069,12 @@ void process_shaders() { string_to_spv("gated_linear_attn_f32", "gla.comp", merge_maps(base_dict, {{"A_TYPE", "float"}})); + // Compile IQ4_NL support in so its shared LUT is available when K uses it. + // K quant type is selected at runtime via the FaTypeK spec constant. + std::map li_dict = {{"FLOAT_TYPE", "float"}, {"FLOAT_TYPEV4", "vec4"}, {"DATA_A_IQ4_NL", "1"}}; + string_to_spv("lightning_indexer_f32", "lightning_indexer.comp", li_dict); + string_to_spv("lightning_indexer_subgroup_f32", "lightning_indexer.comp", merge_maps(li_dict, {{"USE_SUBGROUP_ADD", "1"}})); + string_to_spv("rwkv_wkv7_f32", "wkv7.comp", merge_maps(base_dict, {{"A_TYPE", "float"}})); string_to_spv("gated_delta_net_f32", "gated_delta_net.comp", merge_maps(base_dict, {{"FLOAT_TYPE", "float"}, {"USE_SUBGROUP_ADD", "1"}, {"USE_SUBGROUP_CLUSTERED", "1"}})); From 852997152e4ea554922aac6540c12e112deeb4e4 Mon Sep 17 00:00:00 2001 From: Shawn Gu Date: Thu, 27 Aug 2026 09:44:05 -0700 Subject: [PATCH 013/104] opencl: add bin kernels `kernel_gemm_moe_q4_0_q8_1_dp4a_bin`, `kernel_gemm_moe_mxfp4_q8_1_dp4a_bin` (llama/27768) --- ggml/src/ggml-opencl/ggml-opencl.cpp | 56 ++++++++++++++++++++++++++-- 1 file changed, 53 insertions(+), 3 deletions(-) diff --git a/ggml/src/ggml-opencl/ggml-opencl.cpp b/ggml/src/ggml-opencl/ggml-opencl.cpp index 64f3325b2a5..6ae83449b08 100644 --- a/ggml/src/ggml-opencl/ggml-opencl.cpp +++ b/ggml/src/ggml-opencl/ggml-opencl.cpp @@ -903,6 +903,8 @@ struct ggml_backend_opencl_context { cl_kernel kernel_gemv_moe_mxfp4_f32_ns_wimg = nullptr; // weight-as-texture MoE decode GEMV cl_kernel kernel_gemm_moe_mxfp4_q8_1_dp4a = nullptr; // dp4a (int8) mxfp4 MoE prefill GEMM cl_kernel kernel_gemm_moe_q4_0_q8_1_dp4a = nullptr; // dp4a (int8) q4_0 MoE prefill GEMM + cl_kernel kernel_gemm_moe_mxfp4_q8_1_dp4a_bin = nullptr; // binary dp4a (int8) mxfp4 MoE prefill GEMM + cl_kernel kernel_gemm_moe_q4_0_q8_1_dp4a_bin = nullptr; // binary dp4a (int8) q4_0 MoE prefill GEMM cl_kernel kernel_moe_reorder_b; cl_kernel kernel_moe_histogram, kernel_moe_scan, kernel_moe_fill, kernel_moe_scatter; cl_kernel kernel_moe_scatter_stable = nullptr; // deterministic slot assignment @@ -4248,6 +4250,24 @@ static void load_cl_kernels(ggml_backend_opencl_context *backend_ctx) { GGML_LOG_CONT("."); } + // gemm_moe_mxfp4_q8_1_dp4a_bin (dp4a prefill GEMM) + if (backend_ctx->has_integer_dot) { + size_t bin_size = 0; + backend_ctx->kernel_gemm_moe_mxfp4_q8_1_dp4a_bin = nullptr; + + if (use_adreno_bin_kernels(backend_ctx)) { + const char * kernel_bin = (const char *)backend_ctx->get_adreno_bin_kernel("gemm_moe_mxfp4_q8_1_dp4a_ila", &bin_size); + if (kernel_bin && bin_size > 0) { + cl_program prog = + build_program_from_binary(backend_ctx->context, backend_ctx->device, kernel_bin, CL_moe_compile_opts, bin_size); + + CL_CHECK((backend_ctx->kernel_gemm_moe_mxfp4_q8_1_dp4a_bin = clCreateKernel(prog, "kernel_gemm_moe_mxfp4_q8_1_dp4a_ila", &err), err)); + CL_CHECK(clReleaseProgram(prog)); + GGML_LOG_CONT("."); + } + } + } + // gemm_moe_q4_0_q8_1_dp4a (dp4a prefill GEMM) if (backend_ctx->has_integer_dot) { #ifdef GGML_OPENCL_EMBED_KERNELS @@ -4265,6 +4285,24 @@ static void load_cl_kernels(ggml_backend_opencl_context *backend_ctx) { GGML_LOG_CONT("."); } + // gemm_moe_q4_0_q8_1_dp4a_bin (dp4a prefill GEMM) + if (backend_ctx->has_integer_dot) { + size_t bin_size = 0; + backend_ctx->kernel_gemm_moe_q4_0_q8_1_dp4a_bin = nullptr; + + if (use_adreno_bin_kernels(backend_ctx)) { + const char * kernel_bin = (const char *)backend_ctx->get_adreno_bin_kernel("gemm_moe_q4_0_q8_1_dp4a_ila", &bin_size); + if (kernel_bin && bin_size > 0) { + cl_program prog = + build_program_from_binary(backend_ctx->context, backend_ctx->device, kernel_bin, CL_moe_compile_opts, bin_size); + + CL_CHECK((backend_ctx->kernel_gemm_moe_q4_0_q8_1_dp4a_bin = clCreateKernel(prog, "kernel_gemm_moe_q4_0_q8_1_dp4a_ila", &err), err)); + CL_CHECK(clReleaseProgram(prog)); + GGML_LOG_CONT("."); + } + } + } + // gemm_moe_q8_1_dp4a (generic dp4a MoE GEMM; MOE_QT=80 -> q8_0 expert variant) if (backend_ctx->has_integer_dot) { #ifdef GGML_OPENCL_EMBED_KERNELS @@ -21519,7 +21557,9 @@ static void ggml_cl_mul_mat_id(ggml_backend_t backend, const ggml_tensor * src0, // dot prod has to be available use_moe_dp4a = backend_ctx->has_integer_dot && use_moe_dp4a; // bin kernel takes precedence - use_moe_dp4a = use_moe_dp4a && backend_ctx->kernel_gemm_moe_q4_0_f32_ns_bin == nullptr; + if (backend_ctx->kernel_gemm_moe_q4_0_q8_1_dp4a_bin == nullptr) { + use_moe_dp4a = use_moe_dp4a && backend_ctx->kernel_gemm_moe_q4_0_f32_ns_bin == nullptr; + } cl_buffer_region region; region.origin = 0; @@ -21625,6 +21665,10 @@ static void ggml_cl_mul_mat_id(ggml_backend_t backend, const ggml_tensor * src0, // dp4a GEMM cl_kernel dk = backend_ctx->kernel_gemm_moe_q4_0_q8_1_dp4a; + if (backend_ctx->kernel_gemm_moe_q4_0_q8_1_dp4a_bin) { + dk = backend_ctx->kernel_gemm_moe_q4_0_q8_1_dp4a_bin; + } + int aidx = 0; CL_CHECK(clSetKernelArg(dk, aidx++, sizeof(cl_mem), &extra0_q4_0->q_img)); CL_CHECK(clSetKernelArg(dk, aidx++, sizeof(cl_mem), &extra0_q4_0->d)); @@ -23463,8 +23507,10 @@ static void ggml_cl_mul_mat_id(ggml_backend_t backend, const ggml_tensor * src0, : (backend_ctx->adreno_gen == ADRENO_GPU_GEN::X2E); // dot prod has to be available use_moe_dp4a = backend_ctx->has_integer_dot && use_moe_dp4a; - // bin kernel takes precedence - use_moe_dp4a = use_moe_dp4a && backend_ctx->kernel_gemm_moe_mxfp4_f32_ns_bin == nullptr; + // bin kernel takes precedence, dp4a bin kernel has higher priority than normal bin kernel + if (backend_ctx->kernel_gemm_moe_mxfp4_q8_1_dp4a_bin == nullptr) { + use_moe_dp4a = use_moe_dp4a && backend_ctx->kernel_gemm_moe_mxfp4_f32_ns_bin == nullptr; + } cl_buffer_region region; region.origin = 0; @@ -23573,6 +23619,10 @@ static void ggml_cl_mul_mat_id(ggml_backend_t backend, const ggml_tensor * src0, // dp4a GEMM cl_kernel dk = backend_ctx->kernel_gemm_moe_mxfp4_q8_1_dp4a; + if (backend_ctx->kernel_gemm_moe_mxfp4_q8_1_dp4a_bin) { + dk = backend_ctx->kernel_gemm_moe_mxfp4_q8_1_dp4a_bin; + } + int aidx = 0; CL_CHECK(clSetKernelArg(dk, aidx++, sizeof(cl_mem), &extra0_mxfp4->q_img)); CL_CHECK(clSetKernelArg(dk, aidx++, sizeof(cl_mem), &extra0_mxfp4->e)); From a0614d9ea446915afde48e82e0d6697773c92a18 Mon Sep 17 00:00:00 2001 From: Aparna M P Date: Fri, 28 Aug 2026 03:08:02 +0530 Subject: [PATCH 014/104] hex-unary: fix RMS_NORM_MUL weight-offset bugs for grouped/broadcast norms (llama/27798) --- ggml/src/ggml-hexagon/htp/unary-ops.c | 32 +++++++++++++++++---------- 1 file changed, 20 insertions(+), 12 deletions(-) diff --git a/ggml/src/ggml-hexagon/htp/unary-ops.c b/ggml/src/ggml-hexagon/htp/unary-ops.c index b21415a67d6..971f3b53777 100644 --- a/ggml/src/ggml-hexagon/htp/unary-ops.c +++ b/ggml/src/ggml-hexagon/htp/unary-ops.c @@ -478,6 +478,9 @@ static void unary_task_f32_##NAME(unsigned int nth, unsigned int ith, void * dat const uint32_t nb11 = src1 ? src1->nb[1] : 0; \ const uint32_t nb12 = src1 ? src1->nb[2] : 0; \ const uint32_t nb13 = src1 ? src1->nb[3] : 0; \ + const uint32_t nb11_bc = (src1 && src1->ne[1] > 1) ? nb11 : 0; \ + const uint32_t nb12_bc = (src1 && src1->ne[2] > 1) ? nb12 : 0; \ + const uint32_t nb13_bc = (src1 && src1->ne[3] > 1) ? nb13 : 0; \ const bool src1_contig = src1 ? ((nb12 == (size_t)ne01 * nb11) && (nb13 == (size_t)ne02 * nb12)) : false; \ \ uint8_t * src0_vtcm_data = uctx->vtcm_src0 + (ith * uctx->vtcm_src0_size_per_thread); \ @@ -497,8 +500,12 @@ static void unary_task_f32_##NAME(unsigned int nth, unsigned int ith, void * dat const struct fastdiv_values * div_ne02 = &uctx->kparams->div_ne02; \ const struct fastdiv_values * div_ne012 = &uctx->kparams->div_ne012; \ \ - const uint32_t src0_max_block = src0_contig ? uctx->block : MIN((uint32_t)uctx->block, ne01); \ - const uint32_t dst_max_block = dst_contig ? uctx->block : MIN((uint32_t)uctx->block, ne1); \ + const bool src1_needs_row_clip = (IS_RMS_NORM_MUL) && !uctx->broadcast_weight && !src1_contig; \ + const bool block_src0_contig = src0_contig && !src1_needs_row_clip; \ + const bool block_dst_contig = dst_contig && !src1_needs_row_clip; \ + \ + const uint32_t src0_max_block = block_src0_contig ? uctx->block : MIN((uint32_t)uctx->block, ne01); \ + const uint32_t dst_max_block = block_dst_contig ? uctx->block : MIN((uint32_t)uctx->block, ne1); \ const uint32_t BLOCK = MIN(src0_max_block, dst_max_block); \ if (BLOCK == 0) { \ FARF(ERROR, "unary-f32 : current VTCM reservation %zu is too small, needed at least %zu\n", \ @@ -515,8 +522,8 @@ static void unary_task_f32_##NAME(unsigned int nth, unsigned int ith, void * dat } \ \ for (uint32_t ir = src0_start_row, vtcm_idx = 0; ir < src0_end_row && vtcm_idx < 2; vtcm_idx++) { \ - const uint32_t block_size = unary_block_size(ir, src0_end_row, BLOCK, src0_contig, dst_contig, ne01, \ - div_ne01); \ + const uint32_t block_size = unary_block_size(ir, src0_end_row, BLOCK, block_src0_contig, block_dst_contig, \ + ne01, div_ne01); \ \ dma_queue_push(dma_queue, \ dma_make_ptr(data_dst, dst_vtcm_data + (vtcm_idx * dst_vtcm_half_size)), \ @@ -530,7 +537,7 @@ static void unary_task_f32_##NAME(unsigned int nth, unsigned int ith, void * dat \ if ((IS_RMS_NORM_MUL) && !uctx->broadcast_weight) { \ const size_t src1_off = src1_contig ? (ir * nb11) : \ - unary_row_offset(ir, ne01, ne02, div_ne01, div_ne02, div_ne012, nb11, nb12, nb13); \ + unary_row_offset(ir, ne01, ne02, div_ne01, div_ne02, div_ne012, nb11_bc, nb12_bc, nb13_bc); \ dma_queue_push(dma_queue, \ dma_make_ptr(src1_vtcm_data + (vtcm_idx * src1_vtcm_half_size), data_src1 + src1_off), \ uctx->src1_row_size_aligned, nb11, uctx->src1_data_row_size, block_size); \ @@ -540,8 +547,8 @@ static void unary_task_f32_##NAME(unsigned int nth, unsigned int ith, void * dat } \ \ for (uint32_t ir = src0_start_row; ir < src0_end_row; ) { \ - const uint32_t block_size = unary_block_size(ir, src0_end_row, BLOCK, src0_contig, dst_contig, ne01, \ - div_ne01); \ + const uint32_t block_size = unary_block_size(ir, src0_end_row, BLOCK, block_src0_contig, block_dst_contig, \ + ne01, div_ne01); \ \ float * dst_vtcm = (float *) dma_queue_pop(dma_queue).src; \ float * src0_vtcm = (float *) dma_queue_pop(dma_queue).dst; \ @@ -562,12 +569,12 @@ static void unary_task_f32_##NAME(unsigned int nth, unsigned int ith, void * dat \ const uint32_t next_ir = ir + block_size; \ if (next_ir < src0_end_row) { \ - const uint32_t next_block_size = unary_block_size(next_ir, src0_end_row, BLOCK, src0_contig, dst_contig,\ - ne01, div_ne01); \ + const uint32_t next_block_size = unary_block_size(next_ir, src0_end_row, BLOCK, block_src0_contig, \ + block_dst_contig, ne01, div_ne01); \ const uint32_t pref_ir = next_ir + next_block_size; \ if (pref_ir < src0_end_row) { \ - const uint32_t pref_block_size = unary_block_size(pref_ir, src0_end_row, BLOCK, src0_contig, \ - dst_contig, ne01, div_ne01); \ + const uint32_t pref_block_size = unary_block_size(pref_ir, src0_end_row, BLOCK, block_src0_contig, \ + block_dst_contig, ne01, div_ne01); \ const size_t src0_pref_off = src0_contig ? (pref_ir * nb01) : \ unary_row_offset(pref_ir, ne01, ne02, div_ne01, div_ne02, div_ne012, nb01, nb02, nb03); \ dma_queue_push(dma_queue, \ @@ -576,7 +583,8 @@ static void unary_task_f32_##NAME(unsigned int nth, unsigned int ith, void * dat \ if ((IS_RMS_NORM_MUL) && !uctx->broadcast_weight) { \ const size_t src1_pref_off = src1_contig ? (pref_ir * nb11) : \ - unary_row_offset(pref_ir, ne01, ne02, div_ne01, div_ne02, div_ne012, nb11, nb12, nb13); \ + unary_row_offset(pref_ir, ne01, ne02, div_ne01, div_ne02, div_ne012, nb11_bc, nb12_bc, \ + nb13_bc); \ dma_queue_push(dma_queue, \ dma_make_ptr(src1_vtcm, data_src1 + src1_pref_off), \ uctx->src1_row_size_aligned, nb11, uctx->src1_data_row_size, pref_block_size); \ From b6571e4a55279381fd342c4c9f9ab715f57bd217 Mon Sep 17 00:00:00 2001 From: cqderek Date: Fri, 28 Aug 2026 06:05:57 +0800 Subject: [PATCH 015/104] ggml-hexagon: add HTP unary ops for ABS and LOG (llama/27786) Add HVX-accelerated implementations for GGML_OP_LOG and GGML_UNARY_OP_ABS on the HTP backend. - Register HTP_OP_UNARY_ABS and HTP_OP_UNARY_LOG in op_remap_to_htp() - Add ABS and LOG to ggml_backend_hexagon_device_supports_op() - Implement hvx_abs_f32_aa() in hvx-arith.h using hvx_vec_abs_f32() - Implement hvx_log_f32_aa() in hvx-log.h using hvx_vec_log_f32() - Add abs_f32() and log_f32() row-wise dispatch in unary-ops.c - Define tiled and non-tiled task functions via DEFINE_UNARY_TASK and DEFINE_UNARY_TILED_TASK macros - Route HTP_OP_UNARY_ABS and HTP_OP_UNARY_LOG through execute_op() in main.c --- ggml/src/ggml-hexagon/ggml-hexagon.cpp | 4 +++ ggml/src/ggml-hexagon/htp/htp-ops.h | 2 ++ ggml/src/ggml-hexagon/htp/hvx-arith.h | 28 +++++++++++++++++++ ggml/src/ggml-hexagon/htp/hvx-log.h | 24 ++++++++++++++++ ggml/src/ggml-hexagon/htp/main.c | 2 ++ ggml/src/ggml-hexagon/htp/unary-ops.c | 38 ++++++++++++++++++++++++++ ggml/src/ggml-hexagon/htp/unary-ops.h | 2 ++ 7 files changed, 100 insertions(+) diff --git a/ggml/src/ggml-hexagon/ggml-hexagon.cpp b/ggml/src/ggml-hexagon/ggml-hexagon.cpp index c1e9f919d92..53e86075591 100644 --- a/ggml/src/ggml-hexagon/ggml-hexagon.cpp +++ b/ggml/src/ggml-hexagon/ggml-hexagon.cpp @@ -4643,6 +4643,7 @@ static htp_op_code op_remap_to_htp(const ggml_tensor * t) { case GGML_OP_CLAMP: return HTP_OP_CLAMP; case GGML_OP_SQR: return HTP_OP_SQR; case GGML_OP_SQRT: return HTP_OP_SQRT; + case GGML_OP_LOG: return HTP_OP_UNARY_LOG; case GGML_OP_SOFT_MAX: return HTP_OP_SOFTMAX; case GGML_OP_SSM_CONV: return HTP_OP_SSM_CONV; case GGML_OP_GATED_DELTA_NET: return HTP_OP_GATED_DELTA_NET; @@ -4666,6 +4667,7 @@ static htp_op_code op_remap_to_htp(const ggml_tensor * t) { case GGML_UNARY_OP_EXP: return HTP_OP_UNARY_EXP; case GGML_UNARY_OP_SOFTPLUS: return HTP_OP_UNARY_SOFTPLUS; case GGML_UNARY_OP_TANH: return HTP_OP_UNARY_TANH; + case GGML_UNARY_OP_ABS: return HTP_OP_UNARY_ABS; default: break; } @@ -5463,6 +5465,7 @@ static bool ggml_backend_hexagon_device_supports_op(ggml_backend_dev_t dev, cons case GGML_OP_SQR: case GGML_OP_SQRT: + case GGML_OP_LOG: supp = ggml_hexagon_supported_unary(sess, op); break; @@ -5481,6 +5484,7 @@ static bool ggml_backend_hexagon_device_supports_op(ggml_backend_dev_t dev, cons case GGML_UNARY_OP_SIGMOID: case GGML_UNARY_OP_SOFTPLUS: case GGML_UNARY_OP_TANH: + case GGML_UNARY_OP_ABS: case GGML_UNARY_OP_SILU: case GGML_UNARY_OP_GELU: case GGML_UNARY_OP_GELU_QUICK: diff --git a/ggml/src/ggml-hexagon/htp/htp-ops.h b/ggml/src/ggml-hexagon/htp/htp-ops.h index b4023b34d38..e804844d599 100644 --- a/ggml/src/ggml-hexagon/htp/htp-ops.h +++ b/ggml/src/ggml-hexagon/htp/htp-ops.h @@ -62,6 +62,8 @@ enum htp_op_code { HTP_OP_UNARY_NEG, HTP_OP_UNARY_SOFTPLUS, HTP_OP_UNARY_TANH, + HTP_OP_UNARY_ABS, + HTP_OP_UNARY_LOG, HTP_OP_GLU_SWIGLU, HTP_OP_GLU_SWIGLU_OAI, HTP_OP_GLU_GEGLU, diff --git a/ggml/src/ggml-hexagon/htp/hvx-arith.h b/ggml/src/ggml-hexagon/htp/hvx-arith.h index c8d0003ab5c..765c3577668 100644 --- a/ggml/src/ggml-hexagon/htp/hvx-arith.h +++ b/ggml/src/ggml-hexagon/htp/hvx-arith.h @@ -358,6 +358,34 @@ static inline void hvx_clamp_scalar_f32(uint8_t * restrict dst, const uint8_t * } } +// +// Abs +// + +static inline void hvx_abs_f32_aa(uint8_t * restrict dst, const uint8_t * restrict src, uint32_t n) { + assert((unsigned long) dst % 128 == 0); + assert((unsigned long) src % 128 == 0); + + HVX_Vector * restrict vdst = (HVX_Vector *) dst; + HVX_Vector * restrict vsrc = (HVX_Vector *) src; + + const uint32_t elem_size = sizeof(float); + const uint32_t epv = 128 / elem_size; + const uint32_t nvec = n / epv; + const uint32_t nloe = n % epv; + + uint32_t i = 0; + + _Pragma("unroll(4)") + for (; i < nvec; i++) { + vdst[i] = hvx_vec_abs_f32(vsrc[i]); + } + if (nloe) { + HVX_Vector v = hvx_vec_abs_f32(vsrc[i]); + hvx_vec_store_a((void *) &vdst[i], nloe * elem_size, v); + } +} + // // Square // diff --git a/ggml/src/ggml-hexagon/htp/hvx-log.h b/ggml/src/ggml-hexagon/htp/hvx-log.h index 7013dae785a..a209f88d555 100644 --- a/ggml/src/ggml-hexagon/htp/hvx-log.h +++ b/ggml/src/ggml-hexagon/htp/hvx-log.h @@ -62,4 +62,28 @@ static inline HVX_Vector hvx_vec_log_f32(HVX_Vector x) { return hvx_vec_add_f32_f32(term_e, res); } +static inline void hvx_log_f32_aa(uint8_t * restrict dst, const uint8_t * restrict src, uint32_t n) { + assert((unsigned long) dst % 128 == 0); + assert((unsigned long) src % 128 == 0); + + HVX_Vector * restrict vdst = (HVX_Vector *) dst; + HVX_Vector * restrict vsrc = (HVX_Vector *) src; + + const uint32_t elem_size = sizeof(float); + const uint32_t epv = 128 / elem_size; + const uint32_t nvec = n / epv; + const uint32_t nloe = n % epv; + + uint32_t i = 0; + + _Pragma("unroll(4)") + for (; i < nvec; i++) { + vdst[i] = hvx_vec_log_f32(vsrc[i]); + } + if (nloe) { + HVX_Vector v = hvx_vec_log_f32(vsrc[i]); + hvx_vec_store_a((void *) &vdst[i], nloe * elem_size, v); + } +} + #endif /* HVX_LOG_H */ diff --git a/ggml/src/ggml-hexagon/htp/main.c b/ggml/src/ggml-hexagon/htp/main.c index 975ba0c7af5..fe7d093a81c 100644 --- a/ggml/src/ggml-hexagon/htp/main.c +++ b/ggml/src/ggml-hexagon/htp/main.c @@ -777,6 +777,8 @@ static int execute_op(struct htp_ops_context * octx) { case HTP_OP_UNARY_NEG: case HTP_OP_UNARY_EXP: case HTP_OP_UNARY_TANH: + case HTP_OP_UNARY_ABS: + case HTP_OP_UNARY_LOG: case HTP_OP_L2_NORM: return op_unary(octx); diff --git a/ggml/src/ggml-hexagon/htp/unary-ops.c b/ggml/src/ggml-hexagon/htp/unary-ops.c index 971f3b53777..1a632bf5631 100644 --- a/ggml/src/ggml-hexagon/htp/unary-ops.c +++ b/ggml/src/ggml-hexagon/htp/unary-ops.c @@ -443,6 +443,34 @@ static void tanh_f32(const float * restrict src, } } +static void abs_f32(const float * restrict src, + float * restrict dst, + const uint32_t num_rows, + const struct htp_unary_context * uctx) { + htp_unary_op_preamble; + + for (uint32_t ir = 0; ir < num_rows; ir++) { + const uint8_t * restrict src_local = (const uint8_t *)src + (ir * src0_row_size_aligned); + uint8_t * restrict dst_local = (uint8_t *)dst + (ir * dst_row_size_aligned); + + hvx_abs_f32_aa(dst_local, src_local, ne0); + } +} + +static void log_f32(const float * restrict src, + float * restrict dst, + const uint32_t num_rows, + const struct htp_unary_context * uctx) { + htp_unary_op_preamble; + + for (uint32_t ir = 0; ir < num_rows; ir++) { + const uint8_t * restrict src_local = (const uint8_t *)src + (ir * src0_row_size_aligned); + uint8_t * restrict dst_local = (uint8_t *)dst + (ir * dst_row_size_aligned); + + hvx_log_f32_aa(dst_local, src_local, ne0); + } +} + #define DEFINE_UNARY_TASK(NAME, IS_RMS_NORM_MUL, IS_TRI, CORE_EXPR) \ static void unary_task_f32_##NAME(unsigned int nth, unsigned int ith, void * data) { \ const struct htp_unary_context * uctx = (const struct htp_unary_context *) data; \ @@ -611,6 +639,8 @@ DEFINE_UNARY_TASK(unary_silu, false, false, silu_f32(src0_vtcm, dst_vtcm, bl DEFINE_UNARY_TASK(unary_gelu, false, false, gelu_f32(src0_vtcm, dst_vtcm, block_size, uctx)) DEFINE_UNARY_TASK(unary_softplus, false, false, softplus_f32(src0_vtcm, dst_vtcm, block_size, uctx)) DEFINE_UNARY_TASK(unary_tanh, false, false, tanh_f32(src0_vtcm, dst_vtcm, block_size, uctx)) +DEFINE_UNARY_TASK(unary_abs, false, false, abs_f32(src0_vtcm, dst_vtcm, block_size, uctx)) +DEFINE_UNARY_TASK(unary_log, false, false, log_f32(src0_vtcm, dst_vtcm, block_size, uctx)) DEFINE_UNARY_TASK(l2_norm, false, false, l2_norm_f32(src0_vtcm, dst_vtcm, block_size, uctx)) DEFINE_UNARY_TASK(tri, false, true, tri_f32(src0_vtcm, dst_vtcm, block_size, ir, uctx)) @@ -858,6 +888,8 @@ DEFINE_UNARY_TILED_TASK(unary_silu, false, tile_silu_f32(dst_vtcm, src_vtcm, DEFINE_UNARY_TILED_TASK(unary_gelu, false, tile_gelu_f32(dst_vtcm, src_vtcm, tw)) DEFINE_UNARY_TILED_TASK(unary_softplus, false, tile_unary_softplus_f32(dst_vtcm, src_vtcm, tw)) DEFINE_UNARY_TILED_TASK(unary_tanh, false, hvx_tanh_f32_aa(dst_vtcm, src_vtcm, tw)) +DEFINE_UNARY_TILED_TASK(unary_abs, false, hvx_abs_f32_aa(dst_vtcm, src_vtcm, tw)) +DEFINE_UNARY_TILED_TASK(unary_log, false, hvx_log_f32_aa(dst_vtcm, src_vtcm, tw)) DEFINE_UNARY_TILED_TASK(tri, true, tri_apply_tile_f32(src_vtcm, dst_vtcm, tw, col, i01, ne0, tri_ttype)) static int execute_op_unary_f32(struct htp_ops_context * octx) { @@ -883,6 +915,8 @@ static int execute_op_unary_f32(struct htp_ops_context * octx) { case HTP_OP_UNARY_GELU: op_type = "gelu-f32"; break; case HTP_OP_UNARY_SOFTPLUS: op_type = "softplus-f32"; break; case HTP_OP_UNARY_TANH: op_type = "tanh-f32"; break; + case HTP_OP_UNARY_ABS: op_type = "abs-f32"; break; + case HTP_OP_UNARY_LOG: op_type = "log-f32"; break; case HTP_OP_L2_NORM: op_type = "l2norm-f32"; break; case HTP_OP_TRI: op_type = "tri-f32"; break; @@ -981,6 +1015,8 @@ static int execute_op_unary_f32(struct htp_ops_context * octx) { case HTP_OP_UNARY_GELU: task_func = unary_task_f32_tiled_unary_gelu; break; case HTP_OP_UNARY_SOFTPLUS: task_func = unary_task_f32_tiled_unary_softplus; break; case HTP_OP_UNARY_TANH: task_func = unary_task_f32_tiled_unary_tanh; break; + case HTP_OP_UNARY_ABS: task_func = unary_task_f32_tiled_unary_abs; break; + case HTP_OP_UNARY_LOG: task_func = unary_task_f32_tiled_unary_log; break; case HTP_OP_TRI: task_func = unary_task_f32_tiled_tri; break; default: break; } @@ -1000,6 +1036,8 @@ static int execute_op_unary_f32(struct htp_ops_context * octx) { case HTP_OP_UNARY_GELU: task_func = unary_task_f32_unary_gelu; break; case HTP_OP_UNARY_SOFTPLUS: task_func = unary_task_f32_unary_softplus; break; case HTP_OP_UNARY_TANH: task_func = unary_task_f32_unary_tanh; break; + case HTP_OP_UNARY_ABS: task_func = unary_task_f32_unary_abs; break; + case HTP_OP_UNARY_LOG: task_func = unary_task_f32_unary_log; break; case HTP_OP_L2_NORM: task_func = unary_task_f32_l2_norm; break; case HTP_OP_TRI: task_func = unary_task_f32_tri; break; default: break; diff --git a/ggml/src/ggml-hexagon/htp/unary-ops.h b/ggml/src/ggml-hexagon/htp/unary-ops.h index 1f4c3a5c4d9..458218ff443 100644 --- a/ggml/src/ggml-hexagon/htp/unary-ops.h +++ b/ggml/src/ggml-hexagon/htp/unary-ops.h @@ -55,6 +55,8 @@ static inline bool htp_op_is_unary(uint32_t opcode) { case HTP_OP_UNARY_GELU: case HTP_OP_UNARY_SOFTPLUS: case HTP_OP_UNARY_TANH: + case HTP_OP_UNARY_ABS: + case HTP_OP_UNARY_LOG: case HTP_OP_L2_NORM: case HTP_OP_TRI: return true; From ff38b98e5af6f28329e442ebbd4810a64c780867 Mon Sep 17 00:00:00 2001 From: Brad Smith <1472326+infinitewarp@users.noreply.github.com> Date: Fri, 28 Aug 2026 04:37:43 -0400 Subject: [PATCH 016/104] metal : add fa-vec tunings for M4 Pro (llama/27824) This is a followup contribution to efeda76b948f59ee52ea20db640bc4cf3dfe8ac1 as requested in https://github.com/ggml-org/llama.cpp/discussions/27668 to add support for additional Apple GPUs. I generated this output using the provided instructions: ```sh git clone https://github.com/ggml-org/llama.cpp cd llama.cpp cmake -B build -DGGML_METAL=ON cmake --build build --target ggml-metal-tuning -j ./build/bin/ggml-metal-tuning fa-vec --dtype f16,q8_0 > fa_vec_rows.txt 2> fa_vec_sweep.log ``` This ran on a MacBook Pro (14-inch, Nov 2024) with Apple M4 Pro. The `ggml-metal-tuning` command completed successfully in 1h 13m 1s with no other notable load on the system. --- ggml/src/ggml-metal/ggml-metal-tuning.cpp | 61 ++++++++++++++++++++++- 1 file changed, 60 insertions(+), 1 deletion(-) diff --git a/ggml/src/ggml-metal/ggml-metal-tuning.cpp b/ggml/src/ggml-metal/ggml-metal-tuning.cpp index 6d8c18e6a6a..7e99c1dd661 100644 --- a/ggml/src/ggml-metal/ggml-metal-tuning.cpp +++ b/ggml/src/ggml-metal/ggml-metal-tuning.cpp @@ -448,7 +448,66 @@ constexpr fa_vec_entry_t fa_vec_tuned_table[] = { { { GGML_METAL_DEVICE_M2_ULTRA, GGML_TYPE_Q8_0, 512, 512, -1, 1 }, { 1, 4 } }, { { GGML_METAL_DEVICE_M2_ULTRA, GGML_TYPE_Q8_0, 576, 512, -1, 0 }, { 1, 4 } }, { { GGML_METAL_DEVICE_M2_ULTRA, GGML_TYPE_Q8_0, 576, 512, -1, 1 }, { 1, 4 } }, - + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_F16, 32, 32, 1, 3 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_F16, 32, 32, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_F16, 32, 32, 2, 3 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_F16, 32, 32, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_F16, 32, 32, 3, 3 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_F16, 64, 64, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_F16, 64, 64, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_F16, 64, 64, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_F16, 96, 96, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_F16, 96, 96, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_F16, 96, 96, 1, 4 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_F16, 96, 96, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_F16, 96, 96, 2, 4 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_F16, 128, 128, -1, 1 }, { 2, 2 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_F16, 128, 128, 1, 2 }, { 1, 1 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_F16, 128, 128, 1, 4 }, { 1, 1 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_F16, 192, 192, -1, 1 }, { 2, 2 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_F16, 192, 128, -1, 1 }, { 2, 2 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_F16, 192, 128, 1, 4 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_F16, 256, 256, -1, 0 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_F16, 256, 256, 3, 0 }, { 2, 2 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_F16, 256, 256, -1, 1 }, { 2, 2 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_F16, 320, 256, 3, 0 }, { 2, 2 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_F16, 320, 256, -1, 1 }, { 2, 2 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_F16, 512, 512, 3, 0 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q8_0, 32, 32, -1, 1 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q8_0, 32, 32, 1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q8_0, 32, 32, 1, 4 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q8_0, 32, 32, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q8_0, 32, 32, 2, 4 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q8_0, 32, 32, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q8_0, 64, 64, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q8_0, 64, 64, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q8_0, 64, 64, 3, 2 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q8_0, 96, 96, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q8_0, 96, 96, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q8_0, 128, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q8_0, 128, 128, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q8_0, 128, 128, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q8_0, 128, 128, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q8_0, 128, 128, 3, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q8_0, 192, 192, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q8_0, 192, 192, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q8_0, 192, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q8_0, 192, 128, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q8_0, 192, 128, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q8_0, 192, 128, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q8_0, 192, 128, 3, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q8_0, 256, 256, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q8_0, 256, 256, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q8_0, 320, 256, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q8_0, 320, 256, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q8_0, 512, 512, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q8_0, 512, 512, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q8_0, 576, 512, 2, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q8_0, 576, 512, 3, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q8_0, 576, 512, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q8_0, 576, 512, 1, 1 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q8_0, 576, 512, 1, 3 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q8_0, 576, 512, 1, 4 }, { 1, 2 } }, { { GGML_METAL_DEVICE_M4_MAX, GGML_TYPE_F16, 32, 32, 2, 1 }, { 2, 4 } }, { { GGML_METAL_DEVICE_M4_MAX, GGML_TYPE_F16, 32, 32, 2, 3 }, { 2, 4 } }, { { GGML_METAL_DEVICE_M4_MAX, GGML_TYPE_F16, 32, 32, 3, 1 }, { 2, 4 } }, From 530e3f4834477a267889796e3edd3bdde228ed75 Mon Sep 17 00:00:00 2001 From: Georgi Gerganov Date: Fri, 28 Aug 2026 11:52:03 +0300 Subject: [PATCH 017/104] metal : add fa-vec tunings for M3 Max, M5 and M5 Pro (llama/27863) * metal : add fa-vec tunings for M5 This is a followup contribution to efeda76b948f59ee52ea20db640bc4cf3dfe8ac1 as requested in https://github.com/ggml-org/llama.cpp/discussions/27668 to add support for additional Apple GPUs. I generated this output using the provided instructions: ```sh git clone https://github.com/ggml-org/llama.cpp cd llama.cpp cmake -B build -DGGML_METAL=ON cmake --build build --target ggml-metal-tuning -j ./build/bin/ggml-metal-tuning fa-vec --dtype f16,q8_0 > fa_vec_rows.txt 2> fa_vec_sweep.log ``` This ran on a machine with Apple M5. Assisted-by: pi:llama.cpp/Qwen3.8-27B * metal : add fa-vec tunings for M5 Pro This adds fa_vec_tuned_table records for Apple M5 Pro to ggml-metal-tuning.cpp. Contributed by SerayaEryn in https://github.com/ggml-org/llama.cpp/discussions/27668#discussioncomment-18157544 (F16, Q4_0, Q8_0; M5 Pro, 20 GPU cores). Assisted-by: pi:llama.cpp/Qwen3.8-27B * metal : add fa-vec tunings for M3 Max This adds fa_vec_tuned_table records for Apple M3 Max to ggml-metal-tuning.cpp. Contributed by TeeAaTeeUu in https://github.com/ggml-org/llama.cpp/discussions/27668#discussioncomment-18175220 (F16, Q8_0; M3 Max, MacBook Pro 64GB, low power mode). Assisted-by: pi:llama.cpp/Qwen3.8-27B * cont : whitespaces --- ggml/src/ggml-metal/ggml-metal-tuning.cpp | 305 +++++++++++++++++++++- 1 file changed, 304 insertions(+), 1 deletion(-) diff --git a/ggml/src/ggml-metal/ggml-metal-tuning.cpp b/ggml/src/ggml-metal/ggml-metal-tuning.cpp index 7e99c1dd661..c2139fe20b0 100644 --- a/ggml/src/ggml-metal/ggml-metal-tuning.cpp +++ b/ggml/src/ggml-metal/ggml-metal-tuning.cpp @@ -66,6 +66,7 @@ fa_vec_cfg_t fa_vec_baseline_cfg(int dk, int dv) { // One row per kept bucket, plus per-(dtype,dk,dv) ne11-collapsed domain defaults // (ne11_b = FA_VEC_NE11_DEFAULT, ne01_b = domain). To retune or add a device, re-run the // sweep and paste its output. See ggml-metal-tuning.h for the row/lookup semantics. +// ref: https://github.com/ggml-org/llama.cpp/pull/27824 constexpr fa_vec_entry_t fa_vec_tuned_table[] = { { { GGML_METAL_DEVICE_M1_PRO, GGML_TYPE_F16, 32, 32, 3, 3 }, { 2, 4 } }, { { GGML_METAL_DEVICE_M1_PRO, GGML_TYPE_F16, 64, 64, -1, 0 }, { 1, 4 } }, @@ -448,6 +449,99 @@ constexpr fa_vec_entry_t fa_vec_tuned_table[] = { { { GGML_METAL_DEVICE_M2_ULTRA, GGML_TYPE_Q8_0, 512, 512, -1, 1 }, { 1, 4 } }, { { GGML_METAL_DEVICE_M2_ULTRA, GGML_TYPE_Q8_0, 576, 512, -1, 0 }, { 1, 4 } }, { { GGML_METAL_DEVICE_M2_ULTRA, GGML_TYPE_Q8_0, 576, 512, -1, 1 }, { 1, 4 } }, + + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_F16, 32, 32, 1, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_F16, 32, 32, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_F16, 32, 32, 2, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_F16, 32, 32, 2, 4 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_F16, 32, 32, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_F16, 32, 32, 3, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_F16, 64, 64, 2, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_F16, 64, 64, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_F16, 64, 64, 1, 2 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_F16, 64, 64, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_F16, 64, 64, 3, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_F16, 64, 64, 3, 2 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_F16, 96, 96, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_F16, 96, 96, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_F16, 96, 96, 2, 4 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_F16, 96, 96, 3, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_F16, 96, 96, 3, 3 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_F16, 96, 96, 3, 4 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_F16, 128, 128, -1, 1 }, { 2, 2 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_F16, 128, 128, 1, 4 }, { 1, 1 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_F16, 128, 128, 2, 4 }, { 1, 1 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_F16, 192, 192, 1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_F16, 192, 192, -1, 1 }, { 2, 2 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_F16, 192, 192, 1, 2 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_F16, 192, 128, -1, 1 }, { 2, 2 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_F16, 192, 128, 1, 3 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_F16, 192, 128, 1, 4 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_F16, 192, 128, 3, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_F16, 256, 256, -1, 0 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_F16, 256, 256, -1, 1 }, { 2, 2 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_F16, 256, 256, 2, 3 }, { 1, 1 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_F16, 320, 256, 1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_F16, 320, 256, 3, 0 }, { 2, 2 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_F16, 320, 256, -1, 1 }, { 2, 2 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_F16, 512, 512, 2, 0 }, { 4, 2 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_F16, 512, 512, 3, 0 }, { 4, 1 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_F16, 512, 512, 1, 3 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_F16, 512, 512, 2, 1 }, { 4, 2 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_F16, 512, 512, 2, 3 }, { 4, 2 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_F16, 512, 512, 3, 1 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_F16, 512, 512, 3, 2 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_F16, 576, 512, 2, 0 }, { 4, 2 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_F16, 576, 512, 2, 2 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_F16, 576, 512, 2, 3 }, { 4, 2 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_F16, 576, 512, 3, 1 }, { 4, 2 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q8_0, 32, 32, -1, 1 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q8_0, 32, 32, 1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q8_0, 32, 32, 1, 4 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q8_0, 32, 32, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q8_0, 32, 32, 2, 4 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q8_0, 32, 32, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q8_0, 64, 64, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q8_0, 64, 64, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q8_0, 64, 64, 2, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q8_0, 64, 64, 3, 2 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q8_0, 96, 96, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q8_0, 96, 96, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q8_0, 96, 96, 1, 4 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q8_0, 96, 96, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q8_0, 96, 96, 2, 4 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q8_0, 96, 96, 3, 4 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q8_0, 128, 128, 1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q8_0, 128, 128, 3, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q8_0, 128, 128, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q8_0, 128, 128, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q8_0, 128, 128, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q8_0, 128, 128, 3, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q8_0, 192, 192, 3, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q8_0, 192, 192, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q8_0, 192, 192, 1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q8_0, 192, 192, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q8_0, 192, 192, 1, 4 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q8_0, 192, 192, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q8_0, 192, 192, 3, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q8_0, 192, 128, 1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q8_0, 192, 128, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q8_0, 192, 128, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q8_0, 192, 128, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q8_0, 192, 128, 2, 4 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q8_0, 192, 128, 3, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q8_0, 256, 256, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q8_0, 256, 256, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q8_0, 320, 256, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q8_0, 320, 256, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q8_0, 512, 512, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q8_0, 512, 512, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q8_0, 576, 512, 2, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q8_0, 576, 512, 3, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q8_0, 576, 512, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q8_0, 576, 512, 1, 1 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q8_0, 576, 512, 1, 4 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_F16, 32, 32, 1, 3 }, { 2, 4 } }, { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_F16, 32, 32, 2, 1 }, { 2, 4 } }, { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_F16, 32, 32, 2, 3 }, { 2, 4 } }, @@ -508,6 +602,7 @@ constexpr fa_vec_entry_t fa_vec_tuned_table[] = { { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q8_0, 576, 512, 1, 1 }, { 1, 2 } }, { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q8_0, 576, 512, 1, 3 }, { 1, 2 } }, { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q8_0, 576, 512, 1, 4 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M4_MAX, GGML_TYPE_F16, 32, 32, 2, 1 }, { 2, 4 } }, { { GGML_METAL_DEVICE_M4_MAX, GGML_TYPE_F16, 32, 32, 2, 3 }, { 2, 4 } }, { { GGML_METAL_DEVICE_M4_MAX, GGML_TYPE_F16, 32, 32, 3, 1 }, { 2, 4 } }, @@ -699,7 +794,215 @@ constexpr fa_vec_entry_t fa_vec_tuned_table[] = { { { GGML_METAL_DEVICE_M4_MAX, GGML_TYPE_Q8_0, 576, 512, -1, 0 }, { 1, 4 } }, { { GGML_METAL_DEVICE_M4_MAX, GGML_TYPE_Q8_0, 576, 512, -1, 1 }, { 1, 4 } }, - { { GGML_METAL_DEVICE_M5_MAX, GGML_TYPE_F16, 32, 32, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M5, GGML_TYPE_F16, 32, 32, 1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M5, GGML_TYPE_F16, 32, 32, 1, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M5, GGML_TYPE_F16, 32, 32, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M5, GGML_TYPE_F16, 32, 32, 2, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M5, GGML_TYPE_F16, 32, 32, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M5, GGML_TYPE_F16, 32, 32, 3, 2 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M5, GGML_TYPE_F16, 32, 32, 3, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M5, GGML_TYPE_F16, 32, 32, 3, 4 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M5, GGML_TYPE_F16, 64, 64, 1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M5, GGML_TYPE_F16, 64, 64, 2, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M5, GGML_TYPE_F16, 64, 64, -1, 1 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M5, GGML_TYPE_F16, 64, 64, 1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M5, GGML_TYPE_F16, 64, 64, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M5, GGML_TYPE_F16, 64, 64, 1, 4 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M5, GGML_TYPE_F16, 64, 64, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M5, GGML_TYPE_F16, 64, 64, 2, 4 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M5, GGML_TYPE_F16, 64, 64, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M5, GGML_TYPE_F16, 96, 96, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M5, GGML_TYPE_F16, 96, 96, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M5, GGML_TYPE_F16, 96, 96, 1, 4 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M5, GGML_TYPE_F16, 96, 96, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M5, GGML_TYPE_F16, 96, 96, 3, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M5, GGML_TYPE_F16, 128, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M5, GGML_TYPE_F16, 128, 128, 3, 0 }, { 2, 2 } }, + { { GGML_METAL_DEVICE_M5, GGML_TYPE_F16, 128, 128, -1, 1 }, { 2, 2 } }, + { { GGML_METAL_DEVICE_M5, GGML_TYPE_F16, 128, 128, 1, 2 }, { 1, 1 } }, + { { GGML_METAL_DEVICE_M5, GGML_TYPE_F16, 128, 128, 1, 4 }, { 1, 1 } }, + { { GGML_METAL_DEVICE_M5, GGML_TYPE_F16, 128, 128, 2, 4 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M5, GGML_TYPE_F16, 128, 128, 3, 3 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M5, GGML_TYPE_F16, 128, 128, 3, 4 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M5, GGML_TYPE_F16, 192, 192, -1, 1 }, { 2, 2 } }, + { { GGML_METAL_DEVICE_M5, GGML_TYPE_F16, 192, 192, 1, 2 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M5, GGML_TYPE_F16, 192, 192, 1, 4 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M5, GGML_TYPE_F16, 192, 128, 1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M5, GGML_TYPE_F16, 192, 128, 2, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M5, GGML_TYPE_F16, 192, 128, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M5, GGML_TYPE_F16, 192, 128, 1, 2 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M5, GGML_TYPE_F16, 192, 128, 1, 4 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M5, GGML_TYPE_F16, 192, 128, 3, 2 }, { 2, 2 } }, + { { GGML_METAL_DEVICE_M5, GGML_TYPE_F16, 192, 128, 3, 3 }, { 2, 2 } }, + { { GGML_METAL_DEVICE_M5, GGML_TYPE_F16, 256, 256, -1, 0 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M5, GGML_TYPE_F16, 256, 256, 1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M5, GGML_TYPE_F16, 256, 256, -1, 1 }, { 2, 2 } }, + { { GGML_METAL_DEVICE_M5, GGML_TYPE_F16, 256, 256, 1, 4 }, { 1, 1 } }, + { { GGML_METAL_DEVICE_M5, GGML_TYPE_F16, 320, 256, 2, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M5, GGML_TYPE_F16, 320, 256, 3, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M5, GGML_TYPE_F16, 320, 256, -1, 1 }, { 2, 2 } }, + { { GGML_METAL_DEVICE_M5, GGML_TYPE_F16, 320, 256, 1, 2 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M5, GGML_TYPE_F16, 320, 256, 1, 4 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M5, GGML_TYPE_F16, 512, 512, 2, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M5, GGML_TYPE_F16, 512, 512, 1, 1 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M5, GGML_TYPE_F16, 512, 512, 2, 1 }, { 2, 2 } }, + { { GGML_METAL_DEVICE_M5, GGML_TYPE_F16, 512, 512, 2, 3 }, { 2, 2 } }, + { { GGML_METAL_DEVICE_M5, GGML_TYPE_F16, 512, 512, 3, 1 }, { 2, 2 } }, + { { GGML_METAL_DEVICE_M5, GGML_TYPE_F16, 512, 512, 3, 2 }, { 4, 1 } }, + { { GGML_METAL_DEVICE_M5, GGML_TYPE_F16, 512, 512, 3, 3 }, { 4, 2 } }, + { { GGML_METAL_DEVICE_M5, GGML_TYPE_F16, 512, 512, 3, 4 }, { 2, 2 } }, + { { GGML_METAL_DEVICE_M5, GGML_TYPE_F16, 576, 512, 2, 0 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M5, GGML_TYPE_F16, 576, 512, -1, 1 }, { 2, 2 } }, + { { GGML_METAL_DEVICE_M5, GGML_TYPE_F16, 576, 512, 1, 2 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M5, GGML_TYPE_F16, 576, 512, 1, 4 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M5, GGML_TYPE_F16, 576, 512, 2, 2 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M5, GGML_TYPE_Q8_0, 32, 32, -1, 1 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M5, GGML_TYPE_Q8_0, 32, 32, 1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M5, GGML_TYPE_Q8_0, 32, 32, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M5, GGML_TYPE_Q8_0, 32, 32, 2, 4 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M5, GGML_TYPE_Q8_0, 32, 32, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M5, GGML_TYPE_Q8_0, 32, 32, 3, 4 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M5, GGML_TYPE_Q8_0, 64, 64, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M5, GGML_TYPE_Q8_0, 64, 64, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M5, GGML_TYPE_Q8_0, 64, 64, 1, 2 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M5, GGML_TYPE_Q8_0, 64, 64, 1, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M5, GGML_TYPE_Q8_0, 64, 64, 2, 2 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M5, GGML_TYPE_Q8_0, 64, 64, 3, 2 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M5, GGML_TYPE_Q8_0, 64, 64, 3, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M5, GGML_TYPE_Q8_0, 96, 96, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M5, GGML_TYPE_Q8_0, 96, 96, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M5, GGML_TYPE_Q8_0, 96, 96, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M5, GGML_TYPE_Q8_0, 96, 96, 2, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M5, GGML_TYPE_Q8_0, 96, 96, 3, 4 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M5, GGML_TYPE_Q8_0, 128, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M5, GGML_TYPE_Q8_0, 128, 128, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M5, GGML_TYPE_Q8_0, 128, 128, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M5, GGML_TYPE_Q8_0, 192, 192, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M5, GGML_TYPE_Q8_0, 192, 192, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M5, GGML_TYPE_Q8_0, 192, 192, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M5, GGML_TYPE_Q8_0, 192, 192, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M5, GGML_TYPE_Q8_0, 192, 192, 3, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M5, GGML_TYPE_Q8_0, 192, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M5, GGML_TYPE_Q8_0, 192, 128, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M5, GGML_TYPE_Q8_0, 192, 128, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M5, GGML_TYPE_Q8_0, 192, 128, 2, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M5, GGML_TYPE_Q8_0, 192, 128, 3, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M5, GGML_TYPE_Q8_0, 256, 256, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M5, GGML_TYPE_Q8_0, 256, 256, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M5, GGML_TYPE_Q8_0, 320, 256, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M5, GGML_TYPE_Q8_0, 320, 256, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M5, GGML_TYPE_Q8_0, 320, 256, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M5, GGML_TYPE_Q8_0, 512, 512, 3, 0 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M5, GGML_TYPE_Q8_0, 512, 512, 2, 1 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M5, GGML_TYPE_Q8_0, 512, 512, 3, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M5, GGML_TYPE_Q8_0, 576, 512, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M5, GGML_TYPE_Q8_0, 576, 512, -1, 1 }, { 1, 4 } }, + + { { GGML_METAL_DEVICE_M5_PRO, GGML_TYPE_F16, 32, 32, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M5_PRO, GGML_TYPE_F16, 32, 32, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M5_PRO, GGML_TYPE_F16, 32, 32, 1, 4 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M5_PRO, GGML_TYPE_F16, 32, 32, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M5_PRO, GGML_TYPE_F16, 32, 32, 2, 4 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M5_PRO, GGML_TYPE_F16, 32, 32, 3, 2 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M5_PRO, GGML_TYPE_F16, 32, 32, 3, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M5_PRO, GGML_TYPE_F16, 64, 64, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M5_PRO, GGML_TYPE_F16, 64, 64, 1, 2 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M5_PRO, GGML_TYPE_F16, 64, 64, 3, 2 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M5_PRO, GGML_TYPE_F16, 64, 64, 3, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M5_PRO, GGML_TYPE_F16, 64, 64, 3, 4 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M5_PRO, GGML_TYPE_F16, 96, 96, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M5_PRO, GGML_TYPE_F16, 128, 128, -1, 1 }, { 2, 2 } }, + { { GGML_METAL_DEVICE_M5_PRO, GGML_TYPE_F16, 128, 128, 1, 4 }, { 1, 1 } }, + { { GGML_METAL_DEVICE_M5_PRO, GGML_TYPE_F16, 192, 192, -1, 1 }, { 2, 2 } }, + { { GGML_METAL_DEVICE_M5_PRO, GGML_TYPE_F16, 192, 192, 1, 2 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M5_PRO, GGML_TYPE_F16, 192, 128, 2, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M5_PRO, GGML_TYPE_F16, 192, 128, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M5_PRO, GGML_TYPE_F16, 192, 128, 3, 1 }, { 2, 2 } }, + { { GGML_METAL_DEVICE_M5_PRO, GGML_TYPE_F16, 256, 256, -1, 0 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M5_PRO, GGML_TYPE_F16, 256, 256, -1, 1 }, { 2, 2 } }, + { { GGML_METAL_DEVICE_M5_PRO, GGML_TYPE_F16, 256, 256, 1, 4 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M5_PRO, GGML_TYPE_F16, 320, 256, 3, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M5_PRO, GGML_TYPE_F16, 320, 256, -1, 1 }, { 2, 2 } }, + { { GGML_METAL_DEVICE_M5_PRO, GGML_TYPE_F16, 320, 256, 1, 2 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M5_PRO, GGML_TYPE_F16, 512, 512, 2, 1 }, { 2, 2 } }, + { { GGML_METAL_DEVICE_M5_PRO, GGML_TYPE_F16, 512, 512, 2, 3 }, { 2, 2 } }, + { { GGML_METAL_DEVICE_M5_PRO, GGML_TYPE_F16, 512, 512, 3, 3 }, { 2, 2 } }, + { { GGML_METAL_DEVICE_M5_PRO, GGML_TYPE_F16, 576, 512, -1, 1 }, { 2, 2 } }, + { { GGML_METAL_DEVICE_M5_PRO, GGML_TYPE_F16, 576, 512, 1, 1 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M5_PRO, GGML_TYPE_F16, 576, 512, 1, 2 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M5_PRO, GGML_TYPE_F16, 576, 512, 1, 3 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M5_PRO, GGML_TYPE_F16, 576, 512, 1, 4 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M5_PRO, GGML_TYPE_F16, 576, 512, 2, 2 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M5_PRO, GGML_TYPE_Q4_0, 32, 32, -1, 1 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M5_PRO, GGML_TYPE_Q4_0, 32, 32, 1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M5_PRO, GGML_TYPE_Q4_0, 32, 32, 1, 4 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M5_PRO, GGML_TYPE_Q4_0, 32, 32, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M5_PRO, GGML_TYPE_Q4_0, 32, 32, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M5_PRO, GGML_TYPE_Q4_0, 64, 64, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M5_PRO, GGML_TYPE_Q4_0, 64, 64, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M5_PRO, GGML_TYPE_Q4_0, 64, 64, 1, 2 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M5_PRO, GGML_TYPE_Q4_0, 96, 96, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M5_PRO, GGML_TYPE_Q4_0, 96, 96, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M5_PRO, GGML_TYPE_Q4_0, 96, 96, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M5_PRO, GGML_TYPE_Q4_0, 96, 96, 3, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M5_PRO, GGML_TYPE_Q4_0, 128, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M5_PRO, GGML_TYPE_Q4_0, 128, 128, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M5_PRO, GGML_TYPE_Q4_0, 128, 128, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M5_PRO, GGML_TYPE_Q4_0, 192, 192, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M5_PRO, GGML_TYPE_Q4_0, 192, 192, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M5_PRO, GGML_TYPE_Q4_0, 192, 192, 1, 3 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M5_PRO, GGML_TYPE_Q4_0, 192, 192, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M5_PRO, GGML_TYPE_Q4_0, 192, 192, 2, 3 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M5_PRO, GGML_TYPE_Q4_0, 192, 192, 2, 4 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M5_PRO, GGML_TYPE_Q4_0, 192, 192, 3, 4 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M5_PRO, GGML_TYPE_Q4_0, 192, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M5_PRO, GGML_TYPE_Q4_0, 192, 128, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M5_PRO, GGML_TYPE_Q4_0, 192, 128, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M5_PRO, GGML_TYPE_Q4_0, 192, 128, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M5_PRO, GGML_TYPE_Q4_0, 256, 256, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M5_PRO, GGML_TYPE_Q4_0, 256, 256, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M5_PRO, GGML_TYPE_Q4_0, 320, 256, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M5_PRO, GGML_TYPE_Q4_0, 320, 256, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M5_PRO, GGML_TYPE_Q4_0, 512, 512, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M5_PRO, GGML_TYPE_Q4_0, 512, 512, 1, 0 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M5_PRO, GGML_TYPE_Q4_0, 512, 512, -1, 1 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M5_PRO, GGML_TYPE_Q4_0, 512, 512, 2, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M5_PRO, GGML_TYPE_Q4_0, 512, 512, 2, 3 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M5_PRO, GGML_TYPE_Q8_0, 32, 32, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M5_PRO, GGML_TYPE_Q8_0, 32, 32, 1, 2 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M5_PRO, GGML_TYPE_Q8_0, 32, 32, 1, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M5_PRO, GGML_TYPE_Q8_0, 32, 32, 2, 2 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M5_PRO, GGML_TYPE_Q8_0, 32, 32, 2, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M5_PRO, GGML_TYPE_Q8_0, 32, 32, 3, 2 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M5_PRO, GGML_TYPE_Q8_0, 32, 32, 3, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M5_PRO, GGML_TYPE_Q8_0, 64, 64, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M5_PRO, GGML_TYPE_Q8_0, 64, 64, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M5_PRO, GGML_TYPE_Q8_0, 64, 64, 3, 2 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M5_PRO, GGML_TYPE_Q8_0, 64, 64, 3, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M5_PRO, GGML_TYPE_Q8_0, 96, 96, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M5_PRO, GGML_TYPE_Q8_0, 96, 96, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M5_PRO, GGML_TYPE_Q8_0, 96, 96, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M5_PRO, GGML_TYPE_Q8_0, 128, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M5_PRO, GGML_TYPE_Q8_0, 128, 128, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M5_PRO, GGML_TYPE_Q8_0, 128, 128, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M5_PRO, GGML_TYPE_Q8_0, 128, 128, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M5_PRO, GGML_TYPE_Q8_0, 192, 192, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M5_PRO, GGML_TYPE_Q8_0, 192, 192, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M5_PRO, GGML_TYPE_Q8_0, 192, 192, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M5_PRO, GGML_TYPE_Q8_0, 192, 192, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M5_PRO, GGML_TYPE_Q8_0, 192, 192, 3, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M5_PRO, GGML_TYPE_Q8_0, 192, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M5_PRO, GGML_TYPE_Q8_0, 192, 128, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M5_PRO, GGML_TYPE_Q8_0, 192, 128, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M5_PRO, GGML_TYPE_Q8_0, 192, 128, 3, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M5_PRO, GGML_TYPE_Q8_0, 256, 256, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M5_PRO, GGML_TYPE_Q8_0, 256, 256, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M5_PRO, GGML_TYPE_Q8_0, 320, 256, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M5_PRO, GGML_TYPE_Q8_0, 320, 256, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M5_PRO, GGML_TYPE_Q8_0, 576, 512, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M5_PRO, GGML_TYPE_Q8_0, 576, 512, -1, 1 }, { 1, 4 } }, + + { { GGML_METAL_DEVICE_M5_MAX, GGML_TYPE_F16, 32, 32, -1, 1 }, { 2, 4 } }, { { GGML_METAL_DEVICE_M5_MAX, GGML_TYPE_F16, 32, 32, 1, 2 }, { 1, 4 } }, { { GGML_METAL_DEVICE_M5_MAX, GGML_TYPE_F16, 32, 32, 1, 4 }, { 1, 4 } }, { { GGML_METAL_DEVICE_M5_MAX, GGML_TYPE_F16, 32, 32, 2, 2 }, { 4, 4 } }, From 97d0da26a2e90861a7616a2e12ca4f1e2ed3cbde Mon Sep 17 00:00:00 2001 From: Titaniumtown Date: Fri, 28 Aug 2026 01:53:31 -0700 Subject: [PATCH 018/104] sycl: bind the f16 KV cache in place for the oneDNN SDPA path (llama/27468) Measured at a live KV length of 34816 (32768 depth plus one 2048 ubatch), on Qwen3.8 27B Q4_K_S: per tensor 4 * 34816 * 256 * 2 B = 71.3 MB staged per call K and V, so 2x = 142.6 MB traffic per call read once, write once = 285.2 MB traffic per ubatch 285.2 MB * 16 calls = 4.56 GB One ubatch is one ggml_cgraph submission (llama_context::process_ubatch -> graph_compute), so that 4.56 GB is the cost of a single 2048-token prefill chunk, and it scales with the live KV length: the first ubatch of the same run, at seq = 2048, moves 0.27 GB. Reproduce the two measured inputs with: GGML_SCHED_DEBUG=2 llama-bench -m MODEL -p 8 -n 0 -r 1 -ngl 0 \ -fa on -ctk f16 -ctv f16 -v > nd.txt 2>&1 grep -E 'n_layer|n_head_kv|n_embd_head_k' nd.txt awk '/node # 0 /{g++} g==1 && /\(FLASH_ATTN\)/{n++} END{print n+0}' nd.txt --- ggml/src/ggml-sycl/fattn-onednn.cpp | 62 +++++++++++++++++++++++------ 1 file changed, 49 insertions(+), 13 deletions(-) diff --git a/ggml/src/ggml-sycl/fattn-onednn.cpp b/ggml/src/ggml-sycl/fattn-onednn.cpp index a501295192f..d41c2ddce34 100644 --- a/ggml/src/ggml-sycl/fattn-onednn.cpp +++ b/ggml/src/ggml-sycl/fattn-onednn.cpp @@ -1,3 +1,4 @@ +#include #include #include #include @@ -150,7 +151,8 @@ struct sdpa_partition { // Build + compile the contiguous-input GQA SDPA graph (MatMul->Divide->Add->SoftMax->MatMul), f32 out. // Mirrors the hardware-verified scratch/onednn_sdpa_probe.cpp build_gqa (partitions=1, sdp_primitive_kernel_t). -static sdpa_partition build_sdpa(const engine & eng, int H, int Hkv, int q, int seq, int d) { +static sdpa_partition build_sdpa(const engine & eng, int H, int Hkv, int q, int seq, int d, + const std::array & k_str, const std::array & v_str) try { using ltype = logical_tensor::layout_type; using dt = logical_tensor::data_type; using ldims = logical_tensor::dims; @@ -158,11 +160,12 @@ static sdpa_partition build_sdpa(const engine & eng, int H, int Hkv, int q, int const int rep = H / Hkv; const ldims q_sz = {1, Hkv, rep, q, d}, kv_sz = {1, Hkv, 1, seq, d}, s_sz = {1, Hkv, rep, q, seq}, sc = {1, 1, 1, 1, 1}, msk = {1, 1, 1, q, seq}, o_sz = {1, Hkv, rep, q, d}; + const ldims k_st(k_str.begin(), k_str.end()), v_st(v_str.begin(), v_str.end()); int64_t id = 0; sdpa_partition E; auto query = logical_tensor(id++, t, q_sz, ltype::strided); - auto key = logical_tensor(id++, t, kv_sz, ltype::strided); + auto key = logical_tensor(id++, t, kv_sz, k_st); auto score = logical_tensor(id++, fi, s_sz, ltype::strided); auto bmm1 = op(id++, op::kind::MatMul, "bmm1"); bmm1.set_attr(op::attr::transpose_b, true); // key is [.., seq, d] @@ -184,7 +187,7 @@ static sdpa_partition build_sdpa(const engine & eng, int H, int Hkv, int q, int smax.set_attr(op::attr::mode, "inf_as_zero"); smax.add_inputs({masked}); smax.add_outputs({probs}); - auto value = logical_tensor(id++, t, kv_sz, ltype::strided); + auto value = logical_tensor(id++, t, kv_sz, v_st); // f16 output is REQUIRED to hit sdp_primitive_kernel_t (the systolic micro-kernel); an f32 output // falls to larger_partition_kernel_t which materializes N^2 (confirmed: scratch/onednn_sdpa_kernel_probe.cpp). // converted to the f32 ggml dst in the permute below. @@ -198,6 +201,7 @@ static sdpa_partition build_sdpa(const engine & eng, int H, int Hkv, int q, int auto parts = g.get_partitions(); if (parts.size() != 1 || !parts[0].is_supported()) { + GGML_LOG_WARN("%s: oneDNN did not fuse the SDPA graph; falling back to TILE kernel\n", __func__); return E; // ok stays false -> caller falls back to TILE } E.ins = parts[0].get_input_ports(); @@ -209,6 +213,12 @@ static sdpa_partition build_sdpa(const engine & eng, int H, int Hkv, int q, int E.ok = true; return E; } +catch (const std::exception & e) { + // compile() can reject a stride set the partitioner never inspects; memoise the failure so the + // fallback costs one build rather than one per call. + GGML_LOG_WARN("%s: oneDNN SDPA partition build failed (%s); falling back to TILE kernel\n", __func__, e.what()); + return {}; +} void ggml_sycl_flash_attn_ext_onednn(ggml_backend_sycl_context & ctx, ggml_tensor * dst) try { const ggml_tensor * Q = dst->src[0]; @@ -234,13 +244,34 @@ void ggml_sycl_flash_attn_ext_onednn(ggml_backend_sycl_context & ctx, ggml_tenso ggml_sycl_pool_alloc Qf(ctx.pool(), (size_t) H * q * d); cont_to_f16_sycl((const char *) Q->data, Qf.get(), d, q, H, mb, Q->nb[1], Q->nb[2], Q->nb[3], stream); - // K/V: use pool-alloc for both F16 and dequant paths. + // K/V: bind the f16 cache in place. llama.cpp permutes it to [token][head][dim], so its head + // plane is strided rather than dense, which is what an explicit stride vector expresses. + // Quantized and f32 KV still stage a dense copy -- the layout the k_str/v_str defaults describe. sycl::half * K_ptr = nullptr; sycl::half * V_ptr = nullptr; + std::array k_str{ Hkv * seq * d, seq * d, seq * d, d, 1 }; + std::array v_str = k_str; std::optional> Kf_pool; std::optional> Vf_pool; - if (K->type == GGML_TYPE_F16 && V->type == GGML_TYPE_F16) { + auto bindable = [](const ggml_tensor * t) { + return t->nb[0] == sizeof(sycl::half) && t->nb[1] % sizeof(sycl::half) == 0 && + t->nb[2] % sizeof(sycl::half) == 0 && t->nb[3] % sizeof(sycl::half) == 0; + }; + auto elem_strides = [](const ggml_tensor * t) { + const int64_t s1 = (int64_t) (t->nb[1] / t->nb[0]); + const int64_t s2 = (int64_t) (t->nb[2] / t->nb[0]); + const int64_t s3 = (int64_t) (t->nb[3] / t->nb[0]); + // dims are {mb=1, Hkv, rep=1, seq, d}; the size-1 dims at 0 and 2 never advance an address. + return std::array{ s3, s2, s2, s1, 1 }; + }; + + if (K->type == GGML_TYPE_F16 && V->type == GGML_TYPE_F16 && bindable(K) && bindable(V)) { + K_ptr = (sycl::half *) K->data; + V_ptr = (sycl::half *) V->data; + k_str = elem_strides(K); + v_str = elem_strides(V); + } else if (K->type == GGML_TYPE_F16 && V->type == GGML_TYPE_F16) { Kf_pool.emplace(ctx.pool(), (size_t) Hkv * seq * d); Vf_pool.emplace(ctx.pool(), (size_t) Hkv * seq * d); cont_to_f16_sycl((const char *) K->data, Kf_pool->get(), d, seq, Hkv, mb, K->nb[1], K->nb[2], K->nb[3], stream); @@ -341,19 +372,24 @@ void ggml_sycl_flash_attn_ext_onednn(ggml_backend_sycl_context & ctx, ggml_tenso ggml_sycl_pool_alloc outf(ctx.pool(), (size_t) H * q * d); // f16 contiguous SDPA out [mb,H,q,d] - // compile once per (device, shape), reuse across layers/calls. + // compile once per (device, shape, KV strides), reuse across layers/calls. Stride 2 always + // repeats stride 1 and stride 4 is always 1, so the key covers every entry that can differ. static std::unordered_map cache; - char keyb[96]; - snprintf(keyb, sizeof(keyb), "%d:%lld:%lld:%lld:%lld:%lld", ggml_sycl_get_device(), - (long long) H, (long long) Hkv, (long long) q, (long long) seq, (long long) d); + char keyb[256]; + snprintf(keyb, sizeof(keyb), "%d:%lld:%lld:%lld:%lld:%lld:%lld:%lld:%lld:%lld:%lld:%lld", ggml_sycl_get_device(), + (long long) H, (long long) Hkv, (long long) q, (long long) seq, (long long) d, + (long long) k_str[0], (long long) k_str[1], (long long) k_str[3], + (long long) v_str[0], (long long) v_str[1], (long long) v_str[3]); auto it = cache.find(keyb); if (it == cache.end()) { - it = cache.emplace(keyb, build_sdpa(eng, (int) H, (int) Hkv, (int) q, (int) seq, (int) d)).first; + it = cache.emplace(keyb, build_sdpa(eng, (int) H, (int) Hkv, (int) q, (int) seq, (int) d, k_str, v_str)).first; } sdpa_partition & E = it->second; - // _supported() is authoritative: if it accepted this op the partition must build. - // A failure here is a gap in _supported() -- surface it, don't mask it with a fallback. - GGML_ASSERT(E.ok && "oneDNN SDPA partition failed to build for a _supported() shape"); + if (!E.ok) { + // oneDNN can decline a shape or a stride set that _supported() never sees; build_sdpa warns per key. + ggml_sycl_flash_attn_ext_tile(ctx, dst); + return; + } auto id2ptr = [&](size_t r) -> void * { if (r == E.id_q) return Qf.get(); From fa4d244c939ab9f14e0ce31a08ca85a5bb16353b Mon Sep 17 00:00:00 2001 From: Ozymandias_EBON <112784549+johnkarlhill@users.noreply.github.com> Date: Fri, 28 Aug 2026 03:58:58 -0500 Subject: [PATCH 019/104] sycl: use TILE for quantized KV decode on BMG (llama/26689) Route quantized KV decode to TILE on Xe2 (BMG) only, keep VEC on other archs until validated there. --- ggml/src/ggml-sycl/fattn.cpp | 6 +++++- 1 file changed, 5 insertions(+), 1 deletion(-) diff --git a/ggml/src/ggml-sycl/fattn.cpp b/ggml/src/ggml-sycl/fattn.cpp index a85eb721f6a..a85bca7cb3c 100644 --- a/ggml/src/ggml-sycl/fattn.cpp +++ b/ggml/src/ggml-sycl/fattn.cpp @@ -104,7 +104,6 @@ enum best_fattn_kernel { static best_fattn_kernel ggml_sycl_get_best_fattn_kernel(const int device, const ggml_tensor * dst) { - GGML_UNUSED(device); #ifndef SYCL_FLASH_ATTN GGML_UNUSED(dst); return BEST_FATTN_KERNEL_NONE; @@ -263,6 +262,11 @@ static best_fattn_kernel ggml_sycl_get_best_fattn_kernel(const int device, const } } else { if (Q->ne[1] <= 2) { + // TILE is faster for quantized KV decode on Xe2 (BMG); keep VEC on untested archs + const gpu_arch arch = ggml_sycl_info().devices[device].hw_info.arch; + if (arch == gpu_arch::intel_gpu_bmg_g21 || arch == gpu_arch::intel_gpu_bmg_g31) { + return BEST_FATTN_KERNEL_TILE; + } return BEST_FATTN_KERNEL_VEC; } } From 7f78e1b4694bdae54e8c9101b96be14f89b76f02 Mon Sep 17 00:00:00 2001 From: Zijun Yu Date: Fri, 28 Aug 2026 19:42:07 +0800 Subject: [PATCH 020/104] OpenVINO: Update OV to 2026.3.1, whisper.cpp support, Qwen3.5 on NPU, and new ops (llama/27843) * OpenVINO Backend: Fuse IM2COL + MatMul convolution into OpenVINO convolution * ci:ggml-ov: Skip recurrent state rollback tests * ci:ggml-ov: Skip recurrent state rollback tests * Update OPENVINO.md * ggml-openvino : add env-var gated op support debugging * Fix ggml_rope_set_offset case * OpenVINO backend: Support Whisper.cpp * Fix code style * openvino : enable qwen35 on NPU Static shapes: - get_graph_input_shape() left the s_copy / s_copy-leaf inputs dynamic ([1,1,1,-1]) even in static mode, which propagated a dynamic slot dim through GET_ROWS into the conv/GDN state, the state reshapes and the GDN output. - With -np 1 the s_copy defrag remainder gathers zero rows; short-circuit that CPY to the untouched cache instead of emitting a degenerate Slice/Concat, and skip binding its zero-byte ggml tensor as an output (the dynamic path already did the latter, the static path wrote the full cache over a 0-byte buffer). Token-count independence: - In static mode the compiled model's token count is the prefill chunk size or 1, not the captured cgraph's. Offsets derived from the captured count were therefore wrong. Anchor the GDN state slice at the end of the packed [attn | state] output and drop the rs_src_begin runtime inputs, and make VIEWs over the GDN output / conv_input pass through so the consumer does the slicing. - CONT could not identify its token axis when the graph was captured with a single token (every trailing dim has the same stride and size 1) and baked the captured shape into the prefill model. Chunked prefill: - The last chunk is padded with fabricated tokens. Attention masks them, but the recurrent path folded them into cache_r/cache_s permanently. Add a chunk_valid_len runtime input, use it to zero g and beta for padded steps (making the recurrence an exact identity) and to end the conv snapshot window at the last valid token, and disable the recurrent-cache reset after the first chunk so earlier chunks are not wiped. - get_is_prefill() and the chunk loop bound read inp_pos->ne[0] directly, but IMROPE stacks 4 position planes, so every decode step was run through the padded prefill model and the loop ran extra out-of-bounds chunks. cache_rs_reset_idx/len now stay runtime Parameters in static mode, since can_reuse_statically() does not invalidate the cached model on ComputeParams changes. Add GGML_OPENVINO_FORCE_STATIC to exercise the static path on CPU. * Update to OpenVINO 2026.3.1 * ggml-openvino: forward NPU compilation mode parameters Add GGML_OPENVINO_NPU_COMPILE_CONFIG to the backend's cached environment so callers can configure the NPU compiler without using the generic property escape hatch. When the value is non-empty, pass it to OpenVINO as NPU_COMPILATION_MODE_PARAMS. This enables settings such as optimization-level=3 for NPU compilation while preserving the existing behavior when the variable is unset and leaving CPU and GPU configuration unchanged. Document the variable, its NPU-only scope, and the optimization-level=3 example in the OpenVINO backend runtime configuration table. * ggml-openvino : support RELU, POOL_2D, QUICK_GEGLU, and ROLL ops * reorder op table * exclude GPU/NPU failing POOL_2D case * move op type detection to compute_op_case * Relax rope supported cases * Fix pool case * Update openvino doc, gpu driver in ov docker * openvino: remove unused static remote context branch * openvino: parallelize static model build * Apply editorconfig --------- Co-authored-by: Mostafa Faheem Co-authored-by: Ravi Panchumarthy Co-authored-by: zhaixuejun1993 --- ggml/src/ggml-openvino/CMakeLists.txt | 2 + ggml/src/ggml-openvino/ggml-decoder.cpp | 153 ++++++++++-- ggml/src/ggml-openvino/ggml-decoder.h | 12 +- .../src/ggml-openvino/ggml-openvino-extra.cpp | 12 +- ggml/src/ggml-openvino/ggml-openvino.cpp | 231 ++++++++++-------- ggml/src/ggml-openvino/openvino/op/cpy.cpp | 140 +++++++++-- .../openvino/op/flash_attn_ext.cpp | 104 ++++++-- .../openvino/op/gated_delta_net.cpp | 25 ++ .../openvino/op/glu_geglu_quick.cpp | 64 +++++ .../src/ggml-openvino/openvino/op/pool_2d.cpp | 53 ++++ ggml/src/ggml-openvino/openvino/op/roll.cpp | 36 +++ ggml/src/ggml-openvino/openvino/op/view.cpp | 7 + ggml/src/ggml-openvino/openvino/op_table.cpp | 5 + ggml/src/ggml-openvino/openvino/op_table.h | 3 + .../openvino/pass/fuse_to_conv.cpp | 212 ++++++++++++++++ .../openvino/pass/fuse_to_conv.h | 17 ++ .../openvino/translate_session.cpp | 6 +- ggml/src/ggml-openvino/openvino/utils.cpp | 1 + ggml/src/ggml-openvino/utils.cpp | 159 ++++++++---- ggml/src/ggml-openvino/utils.h | 4 +- 20 files changed, 1048 insertions(+), 198 deletions(-) create mode 100644 ggml/src/ggml-openvino/openvino/op/glu_geglu_quick.cpp create mode 100644 ggml/src/ggml-openvino/openvino/op/pool_2d.cpp create mode 100644 ggml/src/ggml-openvino/openvino/op/roll.cpp create mode 100644 ggml/src/ggml-openvino/openvino/pass/fuse_to_conv.cpp create mode 100644 ggml/src/ggml-openvino/openvino/pass/fuse_to_conv.h diff --git a/ggml/src/ggml-openvino/CMakeLists.txt b/ggml/src/ggml-openvino/CMakeLists.txt index cc089b721fc..af3e0758ca2 100644 --- a/ggml/src/ggml-openvino/CMakeLists.txt +++ b/ggml/src/ggml-openvino/CMakeLists.txt @@ -1,6 +1,8 @@ find_package(OpenVINO REQUIRED COMPONENTS Runtime Threading) find_package(OpenCL REQUIRED) +message(STATUS "Found OpenVINO: ${OpenVINO_DIR} (found version \"${OpenVINO_VERSION}\")") + file(GLOB_RECURSE GGML_HEADERS_OPENVINO "*.h" "*.hpp") file(GLOB_RECURSE GGML_SOURCES_OPENVINO "*.cpp") diff --git a/ggml/src/ggml-openvino/ggml-decoder.cpp b/ggml/src/ggml-openvino/ggml-decoder.cpp index 599f41aebbd..006e005cb7a 100644 --- a/ggml/src/ggml-openvino/ggml-decoder.cpp +++ b/ggml/src/ggml-openvino/ggml-decoder.cpp @@ -357,6 +357,18 @@ int GgmlOvDecoder::compute_op_case(const ggml_tensor * node) const { break; } case GGML_OP_VIEW: { + if (m_is_static && node->src[0] != nullptr && + (node->src[0]->op == GGML_OP_GATED_DELTA_NET || node->src[0]->op == GGML_OP_CONCAT)) { + // VIEW slicing a GATED_DELTA_NET combined [attn|state] output, or the conv_input + // CONCAT. The consuming CPY/RMS_NORM op recovers the true window at runtime via + // ssm_state_size / the fixed conv kernel width, so this VIEW must stay an identity + // pass-through of the full source here too (it already is on the dynamic path); + // otherwise the generic static-mode Slice below would bake in the *captured* + // cgraph's token count, which is wrong once the compiled static model runs with a + // different token count (prefill chunk size or 1). + op_case = 1; + break; + } if (node->src[0]->op == GGML_OP_VIEW) { auto * src = node->src[0]; if (ggml_nelements(node) != ggml_nelements(src)) { @@ -408,6 +420,23 @@ int GgmlOvDecoder::compute_op_case(const ggml_tensor * node) const { } break; } + case GGML_OP_POOL_2D: { + const ggml_op_pool pool_mode = static_cast(node->op_params[0]); + switch (pool_mode) { + case GGML_OP_POOL_MAX: { + op_case = 1; + break; + } + case GGML_OP_POOL_AVG: { + op_case = 2; + break; + } + default: + op_case = 0; + break; + } + break; + } case GGML_OP_CPY: { if (node->src[0]->op == GGML_OP_VIEW) { if (node->src[0]->src[0]->op == GGML_OP_GATED_DELTA_NET) { @@ -425,6 +454,31 @@ int GgmlOvDecoder::compute_op_case(const ggml_tensor * node) const { is_kvcache(node->src[1]->view_src, nullptr)) { // s_copy defrag remainder writeback: gathered extra state rows copied back into the cache op_case = 3; + } else if (node->src[1] != nullptr && node->src[1]->op == GGML_OP_VIEW && node->src[1]->view_src != nullptr) { + // op_case 5: KV write for decoder self-attention (dynamic write offset) + // op_case 6: KV write for encoder self-attn or cross-attn (static offset) + const ggml_tensor * kv_buf = node->src[1]->view_src; + if (kv_buf->ne[1] == 1 && kv_buf->ne[2] == 1 && kv_buf->ne[3] == 1) { + op_case = 6; + // Forward-scan the graph for a FLASH_ATTN_EXT that reads from + // the same buffer. Having a mask (src[3] != nullptr) implies + // decoder self-attention and the write offset is dynamic. + for (int i = 0; i < m_cgraph->n_nodes; i++) { + const ggml_tensor * n = m_cgraph->nodes[i]; + if (n->op != GGML_OP_FLASH_ATTN_EXT) { + continue; + } + // K (src[1]) and V (src[2]) are 3-D views whose view_src is + // the flat KV buffer we are writing to. + if ((n->src[1] != nullptr && n->src[1]->view_src == kv_buf) || + (n->src[2] != nullptr && n->src[2]->view_src == kv_buf)) { + if (n->src[3] != nullptr) { + op_case = 5; // decoder self-attention: mask present + } + break; + } + } + } } break; } @@ -448,6 +502,15 @@ int GgmlOvDecoder::compute_op_case(const ggml_tensor * node) const { } break; } + case GGML_OP_FLASH_ATTN_EXT: { + if (node->src[1] != nullptr && node->src[1]->op == GGML_OP_VIEW && node->src[1]->view_src != nullptr) { + const ggml_tensor * kv_buf = node->src[1]->view_src; + if (kv_buf->ne[1] == 1 && kv_buf->ne[2] == 1 && kv_buf->ne[3] == 1) { + op_case = (node->src[3] != nullptr) ? 1 : 2; + } + } + break; + } default: break; } @@ -479,23 +542,35 @@ std::pair GgmlOvDecoder::compute_llm_params(ggml_cgr switch (node->op) { case GGML_OP_FLASH_ATTN_EXT: - if (node->src[0] == nullptr || node->src[1] == nullptr || node->src[3] == nullptr) { + if (node->src[0] == nullptr || node->src[1] == nullptr) { return -1; } switch (node->src[1]->op) { case GGML_OP_PERMUTE: - // case 0: node op is FLASH_ATTN_EXT, src 1 not null & op is PERMUTE & the permuted tensor src is the view of cache k - if (node->src[1]->src[0] != nullptr && node->src[1]->src[0]->op == GGML_OP_VIEW) { + // case 0: src[1] is PERMUTE of a cache VIEW, mask required + if (node->src[3] != nullptr && node->src[1]->src[0] != nullptr && + node->src[1]->src[0]->op == GGML_OP_VIEW) { return 0; } break; case GGML_OP_CPY: - // case 1: node op is FLASH_ATTN_EXT, src 1 not null & op is CPY & the copied tensor src is PERMUTE & the permuted tensor src is the view of cache k - if (node->src[1]->src[0] != nullptr && node->src[1]->src[0]->op == GGML_OP_PERMUTE && - node->src[1]->src[0]->src[0] != nullptr && node->src[1]->src[0]->src[0]->op == GGML_OP_VIEW) { + // case 1: src[1] is CPY of a PERMUTE(VIEW), mask required + if (node->src[3] != nullptr && node->src[1]->src[0] != nullptr && + node->src[1]->src[0]->op == GGML_OP_PERMUTE && node->src[1]->src[0]->src[0] != nullptr && + node->src[1]->src[0]->src[0]->op == GGML_OP_VIEW) { return 1; } break; + case GGML_OP_VIEW: + // cases 4/5/6: whisper - K is a direct non-contiguous VIEW_3D of a KV cache + if (node->src[1]->view_src != nullptr) { + if (node->src[3] != nullptr) { + return 4; // decoder self-attention + } else { + return 5; // cross-attention or encoder self-attention + }; + } + break; default: break; } @@ -548,6 +623,18 @@ std::pair GgmlOvDecoder::compute_llm_params(ggml_cgr cache_k_permute = node->src[0]->src[0]->src[0]; mask = node->src[1]; break; + case 4: + case 5: { + // whisper: K is a direct VIEW_3D of the KV buffer, no PERMUTE node + auto * cache_k_view = node->src[1]; // VIEW_3D of kv_self.k or kv_cross.k` + compute_params.token_len_per_seq = node->src[0]->ne[1]; + if (attention_pattern_case == 4) { + compute_params.attention_size = cache_k_view->ne[1]; + } else { + compute_params.attention_size_static = cache_k_view->ne[1]; + } + continue; + } default: break; } @@ -654,10 +741,8 @@ std::pair GgmlOvDecoder::compute_llm_params(ggml_cgr ComputeParams::RsWriteback writeback; writeback.slot_begin = (int) (dest_view->view_offs / row_bytes); if (is_conv) { - // conv_input column the copied window starts at writeback.src_begin = (int) (node->src[0]->view_offs / node->src[0]->view_src->nb[0]); } else if (is_gdn) { - // first row of the state part of the gated-delta-net output writeback.src_begin = (int) (node->src[0]->view_offs / node->src[0]->view_src->nb[1]); } compute_params.rs_writebacks[get_tensor_ov_name(cgraph, node)] = writeback; @@ -718,11 +803,15 @@ ov::PartialShape GgmlOvDecoder::get_graph_input_shape(const ggml_tensor * op, } else if (is_kvcache(input, op)) { // kvcache input_shape = ov::PartialShape{get_shape(input)}; - if (!m_is_static) { + // Whisper.cpp uses a fixed size 1D KV buffer [N, 1, 1, 1] (GGML) or [1, 1, 1, N] (OV). + // the token fill level is handled by token_len_per_seq + dynamic mask input. + // skip dynamic dim and stateful reshape for this layout. + const bool is_flat_kv = (input->ne[1] == 1 && input->ne[2] == 1 && input->ne[3] == 1); + if (!m_is_static && !is_flat_kv) { // do not fix ctx size to make llama-bench work across test params input_shape[2] = -1; } - if (is_stateful()) { + if (is_stateful() && !is_flat_kv) { // Convert stateless KV cache layout [1, 1, seq, n_heads_kv * head_size] // to stateful layout [1, seq, n_heads_kv, head_size]. assert(input_shape.size() == 4 && input_shape[0] == 1 && input_shape[1] == 1 && @@ -738,7 +827,9 @@ ov::PartialShape GgmlOvDecoder::get_graph_input_shape(const ggml_tensor * op, input_shape = ov::PartialShape{1, 1, 1, len}; } else if (is_inp_s_copy(input, op) || is_s_copy_leaf(input)) { - input_shape = ov::PartialShape{1, 1, 1, -1}; + // On NPU the total slot count (n_seq_max) is fixed at translation time, so the s_copy + // index list has a static length; on CPU/GPU it may change across compiles (defrag). + input_shape = m_is_static ? ov::PartialShape{get_shape(input)} : ov::PartialShape{1, 1, 1, -1}; } else { input_shape = ov::PartialShape{get_shape(input)}; @@ -790,13 +881,16 @@ void GgmlOvDecoder::add_extra_inputs() { // see llama_kv_cache_unified::get_n_kv and llama_kv_cache_unified::get_padding. // 2. `n_seq_active` and `seq_active_start`, used in FLASH_ATTN_EXT to indicate the active sequences in the batch - auto create_1d_input = [this](const std::string & name, int64_t value) { - m_model_extra_inputs[name] = {ov::element::i64, ov::Shape{1}, value, !m_is_static}; + auto create_1d_input = [this](const std::string & name, int64_t value, bool force_parameter = false) { + m_model_extra_inputs[name] = {ov::element::i64, ov::Shape{1}, value, force_parameter || !m_is_static}; }; if (m_compute_params.attention_size != -1) { create_1d_input("attention_size", m_compute_params.attention_size); } + if (m_compute_params.attention_size_static != -1) { + create_1d_input("attention_size_static", m_compute_params.attention_size_static); + } if (m_compute_params.attention_size_swa != -1) { create_1d_input("attention_size_swa", m_compute_params.attention_size_swa); } @@ -809,17 +903,32 @@ void GgmlOvDecoder::add_extra_inputs() { // create_1d_input("token_len", m_compute_params.token_len_per_seq * m_compute_params.n_seq_active); if (m_compute_params.cache_rs_reset_idx != -1) { - create_1d_input("cache_rs_reset_idx", m_compute_params.cache_rs_reset_idx); - create_1d_input("cache_rs_reset_len", m_compute_params.cache_rs_reset_len); + // Whether/which cache slot to reset varies per compute call (e.g. a new sequence starting + // vs. continued decoding). can_reuse_statically() does not invalidate the cached static + // model on ComputeParams changes, so these must stay runtime Parameters even when static + // (scale.cpp op_case 1 only uses them in value comparisons, never as Slice bounds, so this + // does not reintroduce dynamic shapes). + create_1d_input("cache_rs_reset_idx", m_compute_params.cache_rs_reset_idx, /*force_parameter=*/true); + create_1d_input("cache_rs_reset_len", m_compute_params.cache_rs_reset_len, /*force_parameter=*/true); } if (m_compute_params.s_copy_active_slot_len != -1) { create_1d_input("s_copy_active_slot_len", m_compute_params.s_copy_active_slot_len); + if (m_is_static) { + // Number of real tokens in the current prefill chunk. The last chunk is padded with + // fabricated token ids; attention masks them out, but the recurrent (GDN/conv) path + // would otherwise fold them into cache_r/cache_s permanently. Varies per chunk, so it + // must stay a runtime Parameter; it is only compared against a Range or used as Gather + // indices, so it does not make any shape dynamic. + create_1d_input("chunk_valid_len", get_static_n_tokens(), /*force_parameter=*/true); + } } for (const auto & [node_name, writeback] : m_compute_params.rs_writebacks) { create_1d_input("rs_slot_begin_" + node_name, writeback.slot_begin); - create_1d_input("rs_src_begin_" + node_name, writeback.src_begin); + if (!m_is_static) { + create_1d_input("rs_src_begin_" + node_name, writeback.src_begin); + } } } @@ -1785,13 +1894,23 @@ void GgmlOvDecoder::compute_node_dynamic_dims() { auto dynamic_dim_stride = src_logical_nb[dynamic_dim_idx] / ggml_type_size(node->src[0]->type) * ggml_type_size(node->type); int matched_dim_count = 0; + int first_matched_dim = -1; for (int i = 0; i < GGML_MAX_DIMS; i++) { if (node->nb[i] == dynamic_dim_stride && node->ne[i] == node->src[0]->ne[dynamic_dim_idx]) { + if (first_matched_dim == -1) { + first_matched_dim = i; + } m_node_dynamic_dims[node] = i; matched_dim_count++; } } - if (matched_dim_count != 1) { + if (matched_dim_count > 1 && node->src[0]->ne[dynamic_dim_idx] == 1) { + // Single-token capture: every trailing dim is size 1 with the same stride, so + // the match is ambiguous. The lowest index is the real axis; the rest are + // ggml's size-1 padding. Bailing out here would bake the captured token count + // into the static prefill model, which then runs with a different one. + m_node_dynamic_dims[node] = first_matched_dim; + } else if (matched_dim_count != 1) { m_node_dynamic_dims[node] = -1; GGML_LOG_WARN("ggml-openvino: cannot determine dynamic dim for CONT node '%s', src[0]: '%s'\n", node->name, node->src[0]->name); diff --git a/ggml/src/ggml-openvino/ggml-decoder.h b/ggml/src/ggml-openvino/ggml-decoder.h index 8e39a26c8b7..74cb7385029 100644 --- a/ggml/src/ggml-openvino/ggml-decoder.h +++ b/ggml/src/ggml-openvino/ggml-decoder.h @@ -47,6 +47,7 @@ struct ComputeParams { int seq_active_start = 0; int attention_size = -1; int attention_size_swa = -1; + int attention_size_static = -1; // encoder/cross-attn KV fill level (whisper) int input_len = -1; int token_len_per_seq = -1; int past_kv_len = -1; @@ -84,14 +85,15 @@ struct ComputeParams { struct RsWriteback { int slot_begin = 0; // first cache slot written by the CPY - int src_begin = 0; // where the copied data starts in the source tensor (in rows of it) + int src_begin = 0; // first source row or column copied by the CPY }; std::map rs_writebacks; - // Offsets of the state cache writeback CPY nodes, keyed by node name. They change with the - // batch (kv head, active sequence count, token count) and, with rollback enabled - // (cparams.n_rs_seq > 0), the conv state is written back once per snapshot slot, each snapshot - // taking a different conv_input window. Passed to the cached model as runtime inputs. + // Destination slot offset of each state cache writeback CPY node, keyed by node name. It + // changes with the batch (kv head, active sequence count) and, with rollback enabled + // (cparams.n_rs_seq > 0), the conv state is written back once per snapshot slot. Passed to the + // cached model as a runtime input. Dynamic models also receive the source-side offset; static + // models use a fixed end-anchored offset in the translator. }; class GgmlOvDecoder : public ov::frontend::ggml::GgmlDecoder { diff --git a/ggml/src/ggml-openvino/ggml-openvino-extra.cpp b/ggml/src/ggml-openvino/ggml-openvino-extra.cpp index 36c749244f8..36dfa4d9471 100644 --- a/ggml/src/ggml-openvino/ggml-openvino-extra.cpp +++ b/ggml/src/ggml-openvino/ggml-openvino-extra.cpp @@ -32,6 +32,8 @@ void ggml_openvino_device_config::init() { "GGML_OPENVINO_DEVICE", "GGML_OPENVINO_CACHE_DIR", "GGML_OPENVINO_DEBUG_NODE", + "GGML_OPENVINO_COMPILED_MODEL_CACHE_DIR", + "GGML_OPENVINO_NPU_COMPILE_CONFIG", // Integer values (use ggml_openvino_getenv_int) "GGML_OPENVINO_PREFILL_CHUNK_SIZE", // Boolean toggles (treated as int flags via ggml_openvino_getenv_int) @@ -41,6 +43,9 @@ void ggml_openvino_device_config::init() { "GGML_OPENVINO_DUMP_IR", "GGML_OPENVINO_DEBUG_INPUT", "GGML_OPENVINO_DEBUG_OUTPUT", + // Force the static (NPU-shape) compute path on any device, e.g. GGML_OPENVINO_DEVICE=CPU, + // to test the static-shape translation without NPUW/real NPU hardware in the loop. + "GGML_OPENVINO_FORCE_STATIC", "GGML_OPENVINO_PRINT_CGRAPH_TENSOR_ADDRESS", "GGML_OPENVINO_ENABLE_CACHE", "GGML_OPENVINO_DISABLE_CACHE", @@ -50,7 +55,7 @@ void ggml_openvino_device_config::init() { "GGML_OPENVINO_MEMORY_OPTIMIZE", "GGML_OPENVINO_RELEASE_WEIGHTS", "GGML_OPENVINO_REDUCE_COMPILE_MEM", - "GGML_OPENVINO_COMPILED_MODEL_CACHE_DIR", + "GGML_OPENVINO_LOG_UNSUPPORTED_OPS", }; for (const char * const & env_var : env_var_names) { @@ -85,6 +90,11 @@ void ggml_openvino_device_config::init() { compile_config["NPUW_CACHE_DIR"] = cache_dir; compile_config.insert(ov::cache_mode(ov::CacheMode::OPTIMIZE_SIZE)); } + const char * compilation_mode_params = + ggml_openvino_getenv_str("GGML_OPENVINO_NPU_COMPILE_CONFIG"); + if (compilation_mode_params && strlen(compilation_mode_params) > 0) { + compile_config["NPU_COMPILATION_MODE_PARAMS"] = compilation_mode_params; + } } else if (cache_dir && strlen(cache_dir) > 0) { compile_config.insert(ov::cache_dir(cache_dir)); compile_config.insert(ov::cache_mode(ov::CacheMode::OPTIMIZE_SIZE)); diff --git a/ggml/src/ggml-openvino/ggml-openvino.cpp b/ggml/src/ggml-openvino/ggml-openvino.cpp index e299e16c778..4b1789713d1 100644 --- a/ggml/src/ggml-openvino/ggml-openvino.cpp +++ b/ggml/src/ggml-openvino/ggml-openvino.cpp @@ -908,11 +908,27 @@ static bool has_non_contiguous_view_input(const ggml_tensor * op) { } static bool is_supported_flash_attn_pattern(const ggml_tensor * op) { - // pattern of q,k,v should be q->op==PERMUTE, q->src[0]->op==VIEW, q->src[0]->src[0]->view_src==nullptr + // Each Q/K/V input must follow one of: + // PERMUTE -> VIEW -> base (view_src==nullptr) (llama KV-cache path) + // PERMUTE -> RESHAPE -> base (view_src==nullptr) (whisper Q) + // VIEW -> base (view_src==nullptr) (whisper K/V from kv_pad) for (int i = 0; i < 3; i++) { const ggml_tensor * src = op->src[i]; - if (src->op != GGML_OP_PERMUTE || src->src[0] == nullptr || src->src[0]->op != GGML_OP_VIEW || - src->src[0]->src[0] == nullptr || src->src[0]->src[0]->view_src != nullptr) { + if (src->op == GGML_OP_PERMUTE) { + if (src->src[0] == nullptr) { + return false; + } + if (src->src[0]->op != GGML_OP_VIEW && src->src[0]->op != GGML_OP_RESHAPE) { + return false; + } + if (src->src[0]->src[0] == nullptr || src->src[0]->src[0]->view_src != nullptr) { + return false; + } + } else if (src->op == GGML_OP_VIEW) { + if (src->src[0] == nullptr || src->src[0]->view_src != nullptr) { + return false; + } + } else { return false; } } @@ -1030,18 +1046,29 @@ static bool is_msa_block_mask_expansion(const ggml_tensor * op) { return tensor_name_starts_with(src, "msa_block_mask"); } -static bool is_op_unsupported_case(const ggml_tensor * op) { +namespace { +struct ggml_openvino_op_support { + bool is_supported = true; + std::string reason; + + operator bool() const { + return is_supported; + } +}; +} // namespace + +static ggml_openvino_op_support is_op_supported_case(const ggml_tensor * op) { if (is_msa_block_mask_expansion(op)) { - return true; + return {false, "MSA block mask expansion is not supported"}; } switch (op->op) { case GGML_OP_CONCAT: { if (op->type == GGML_TYPE_I64) { - return true; + return {false, "CONCAT with I64 type is not supported"}; } if (ggml_openvino_get_device_name() == "GPU" && op->type == GGML_TYPE_BF16 && has_view_op_input(op)) { - return true; + return {false, "CONCAT with BF16 type and VIEW input is not supported on GPU"}; } break; } @@ -1052,24 +1079,21 @@ static bool is_op_unsupported_case(const ggml_tensor * op) { // OpenVINO SET translation currently supports dst layouts that match src0 strides. if (op->src[0] == nullptr || nb1 != op->src[0]->nb[1] || nb2 != op->src[0]->nb[2] || nb3 != op->src[0]->nb[3]) { - // std::cout << "Unsupported SET op with dst nb1=" << nb1 << ", nb2=" << nb2 << ", nb3=" << nb3 - // << " that does not match src0 strides nb[1]=" - // << (op->src[0] != nullptr ? std::to_string(op->src[0]->nb[1]) : "null") - // << ", nb[2]=" << (op->src[0] != nullptr ? std::to_string(op->src[0]->nb[2]) : "null") - // << ", nb[3]=" << (op->src[0] != nullptr ? std::to_string(op->src[0]->nb[3]) : "null") - // << std::endl; - return true; + return {false, "SET op with dst nb1=" + std::to_string(nb1) + ", nb2=" + std::to_string(nb2) + ", nb3=" + std::to_string(nb3) + + " that does not match src0 strides nb[1]=" + (op->src[0] != nullptr ? std::to_string(op->src[0]->nb[1]) : "null") + + ", nb[2]=" + (op->src[0] != nullptr ? std::to_string(op->src[0]->nb[2]) : "null") + + ", nb[3]=" + (op->src[0] != nullptr ? std::to_string(op->src[0]->nb[3]) : "null")}; } break; } case GGML_OP_GET_ROWS: case GGML_OP_SET_ROWS: { if (op->ne[3] != 1) { - return true; + return {false, "GET_ROWS/SET_ROWS with ne[3] != 1 (ne[3]=" + std::to_string(op->ne[3]) + ") is not supported"}; } if (op->op == GGML_OP_GET_ROWS && ggml_openvino_get_device_name() == "GPU" && op->src[0]->type == GGML_TYPE_BF16) { - return true; + return {false, "GET_ROWS with BF16 src0 is not supported on GPU"}; } if (op->ne[0] == 256 && (op->src[0]->type == GGML_TYPE_Q4_K || op->src[0]->type == GGML_TYPE_Q5_K || op->src[0]->type == GGML_TYPE_Q4_1 || op->src[0]->type == GGML_TYPE_Q5_1)) { @@ -1078,14 +1102,14 @@ static bool is_op_unsupported_case(const ggml_tensor * op) { // make_int8_weights/make_int4_weights: dequant is done in f16, not f32, to keep the // Convert/Subtract/Multiply chain fusable into GatherMatmulCompressed/FullyConnectedCompressed // for the shared non-test code paths). - return true; + return {false, "GET_ROWS/SET_ROWS with ne[0] == 256 and type " + std::string(ggml_type_name(op->src[0]->type)) + + " rejected due to f16-arithmetic dequant rounding errors that intermittently exceed 1e-7 NMSE threshold"}; } - break; } case GGML_OP_RESHAPE: { if (strncmp(op->name, "ffn_norm_exps", sizeof("ffn_norm_exps") - 1) == 0) { - return true; + return {false, "RESHAPE for ffn_norm_exps is not supported"}; } break; } @@ -1093,11 +1117,13 @@ static bool is_op_unsupported_case(const ggml_tensor * op) { case GGML_OP_MUL: case GGML_OP_SUB: { if (op->src[1]->op == GGML_OP_PERMUTE) { - return true; + return {false, "ADD/MUL/SUB with PERMUTE src1 is not supported"}; } for (int i = 0; i < 4; i++) { if (op->src[0]->ne[i] != op->src[1]->ne[i] && (op->src[0]->ne[i] != 1 && op->src[1]->ne[i] != 1)) { - return true; + return {false, "ADD/MUL/SUB with incompatible broadcast shapes: src0->ne[" + std::to_string(i) + "]=" + + std::to_string(op->src[0]->ne[i]) + ", src1->ne[" + std::to_string(i) + "]=" + + std::to_string(op->src[1]->ne[i])}; } } break; @@ -1106,7 +1132,7 @@ static bool is_op_unsupported_case(const ggml_tensor * op) { // Keep support aligned with the CPU backend implementation, which only handles f32 inputs/output and i32 ids. if (op->type != GGML_TYPE_F32 || op->src[0]->type != GGML_TYPE_F32 || op->src[1]->type != GGML_TYPE_F32 || op->src[2]->type != GGML_TYPE_I32) { - return true; + return {false, "ADD_ID only supports F32 inputs/output and I32 ids"}; } break; } @@ -1116,14 +1142,27 @@ static bool is_op_unsupported_case(const ggml_tensor * op) { // until the fused GPU kernel is reliable. (falied case llama-arch-test mpt) if (ggml_openvino_get_device_name() == "GPU" && op->src[1]->ne[0] == op->ne[0] && op->src[1]->ne[1] == 1 && op->src[1]->ne[2] == 1 && op->src[1]->ne[3] == 1) { - return true; + return {false, "DIV per-channel scale broadcast is not supported on GPU"}; + } + break; + } + case GGML_OP_POOL_2D: { + const auto& name = ggml_openvino_get_device_name(); + if (name == "GPU") { + const int32_t * params = op->op_params; + const int k0 = params[1]; + const int k1 = params[2]; + const int p0 = params[5]; + const int p1 = params[6]; + if ((p0 > 0 || p1 > 0) && (k0 < 3 || k1 < 3)) { + return {false, "POOL_2D with padding and kernel size < 3 is not supported on " + name}; + } } break; } case GGML_OP_SUM_ROWS: { - // if the input is PERMUTE skip if (op->src[0]->op == GGML_OP_PERMUTE) { - return true; + return {false, "SUM_ROWS with PERMUTE input is not supported"}; } break; } @@ -1140,54 +1179,51 @@ static bool is_op_unsupported_case(const ggml_tensor * op) { // accuracy drift in the OpenVINO path. Restrict by scale=1.0 to avoid // affecting non-gemma3n models such as Llama-3.2. if (fabsf(scale - 1.0f) < 1e-6f && is_gemma3n_flash_attn_pattern(op)) { - return true; + return {false, "FLASH_ATTN_EXT gemma3n pattern on GPU is not supported"}; } if (op->src[4] != nullptr) { - // GGML_LOG_WARN("OpenVINO backend does not support FLASH_ATTN_EXT with sinks\n"); - return true; + return {false, "FLASH_ATTN_EXT with sinks is not supported"}; } if (!is_supported_flash_attn_pattern(op)) { - return true; + return {false, "FLASH_ATTN_EXT unsupported attention pattern"}; } if (max_bias > 0) { - // GGML_LOG_WARN("OpenVINO backend does not support FLASH_ATTN_EXT with max_bias > 0\n"); - return true; + return {false, "FLASH_ATTN_EXT with max_bias > 0 (max_bias=" + std::to_string(max_bias) + ") is not supported"}; } if (logit_softcap != 0) { - // GGML_LOG_WARN("OpenVINO backend does not support FLASH_ATTN_EXT with logit_softcap != 0\n"); - return true; + return {false, "FLASH_ATTN_EXT with logit_softcap != 0 (logit_softcap=" + std::to_string(logit_softcap) + ") is not supported"}; } break; } case GGML_OP_PERMUTE: { - if (op->type == GGML_TYPE_BF16) { - // err msg: [GPU] Could not find a suitable kernel for transpose - // GGML_LOG_WARN("OpenVINO backend does not support PERMUTE with BF16 type\n"); - return true; + if (op->type == GGML_TYPE_BF16 && ggml_openvino_get_device_name() == "GPU") { + return {false, "PERMUTE with BF16 type is not supported on GPU"}; } break; } case GGML_OP_CPY: { if (op->src[0]->type == GGML_TYPE_BF16 || op->src[1]->type == GGML_TYPE_BF16) { - // GGML_LOG_WARN("OpenVINO backend does not support CPY with non-contiguous data or bf16 types\n"); - return true; + return {false, "CPY with BF16 src type is not supported"}; } // CPY to a quantized destination (e.g. f32 -> q4_0) is numerically unstable with OpenVINO backend. if (ggml_is_quantized(op->type)) { - return true; + return {false, "CPY to quantized destination (e.g. f32 -> q4_0) is numerically unstable"}; } if (ggml_nelements(op->src[0]) != ggml_nelements(op->src[1])) { - return true; + return {false, "CPY with mismatched element counts is not supported: src0=" + std::to_string(ggml_nelements(op->src[0])) + + " != src1=" + std::to_string(ggml_nelements(op->src[1]))}; } // op test case with non-contiguous src or dst if ((op->ne[0] == 3 && op->ne[1] == 4 && op->ne[2] == 3 && op->ne[3] == 2) || (op->ne[0] == 1 && op->ne[1] == 4 && op->ne[2] == 3 && op->ne[3] == 2) || (op->ne[0] == 2 && op->ne[1] == 4 && op->ne[2] == 3 && op->ne[3] == 2)) { - return true; + return {false, "CPY with non-contiguous shape [" + std::to_string(op->ne[0]) + ", " + + std::to_string(op->ne[1]) + ", " + std::to_string(op->ne[2]) + ", " + + std::to_string(op->ne[3]) + "] is not supported"}; } if (!cpy_output_view_is_supported(op)) { - return true; + return {false, "CPY with non-contiguous output view is not supported"}; } break; } @@ -1196,13 +1232,14 @@ static bool is_op_unsupported_case(const ggml_tensor * op) { ggml_is_quantized(op->src[0]->type) && strcmp(op->src[0]->name, "a") == 0 && strcmp(op->src[1]->name, "b") == 0 && op->src[0]->ne[1] == 1 && op->src[1]->ne[1] == 64 && op->src[0]->ne[0] == 256 && op->src[1]->ne[0] == 256) { - return true; + return {false, "MUL_MAT quantized benchmark test case on GPU is not supported"}; } if (op->src[0]->ne[3] != op->src[1]->ne[3] && op->src[0]->ne[3] != 1 && op->src[1]->ne[3] != 1) { - return true; + return {false, "MUL_MAT with incompatible broadcast on ne[3]: src0->ne[3]=" + std::to_string(op->src[0]->ne[3]) + + ", src1->ne[3]=" + std::to_string(op->src[1]->ne[3])}; } if (op->src[0]->op == GGML_OP_VIEW && op->src[1]->op == GGML_OP_VIEW) { - return true; + return {false, "MUL_MAT with both inputs as VIEW is not supported"}; } break; } @@ -1210,16 +1247,17 @@ static bool is_op_unsupported_case(const ggml_tensor * op) { // Single-expert (or empty) MUL_MAT_ID is a degenerate shape that stresses GatherMatmul edge // cases and never occurs in real MoE; let it fall back to CPU. if (op->src[0] != nullptr && op->src[0]->ne[2] <= 1) { - return true; + return {false, "MUL_MAT_ID with single-expert or empty ne[2] <= 1 (ne[2]=" + + std::to_string(op->src[0]->ne[2]) + ") is not supported"}; } if (ggml_openvino_get_device_name() == "GPU" && op->src[0] != nullptr && op->src[0]->type == GGML_TYPE_BF16) { - return true; + return {false, "MUL_MAT_ID with BF16 weights on GPU is not supported"}; } // GPU MUL_MAT_ID uses a Gather+MatMul fallback because the GPU plugin rejects internal // GatherMatmul for these test shapes. Skip cases that would materialize a large selected // expert-weight temporary. if (ggml_openvino_get_device_name() == "GPU" && mul_mat_id_requires_large_tmp(op)) { - return true; + return {false, "MUL_MAT_ID requires large temporary on GPU"}; } break; } @@ -1229,51 +1267,46 @@ static bool is_op_unsupported_case(const ggml_tensor * op) { const int mode = op_params[2]; if (op_params[15] != 0) { // FIXME: support ggml_rope_set_offset - return true; + return {false, "ggml_rope_set_offset is not supported"}; } if (mode != GGML_ROPE_TYPE_NORMAL && mode != GGML_ROPE_TYPE_NEOX && mode != GGML_ROPE_TYPE_IMROPE) { - // GGML_LOG_WARN("OpenVINO backend does not support ROPE with mode %d\n", mode); - return true; + return {false, "ROPE with mode " + std::to_string(mode) + " is not supported"}; } const int64_t head_dim = op->src[0]->ne[0]; const int64_t rope_dims = n_dims == 0 ? head_dim : n_dims; if (rope_dims <= 0 || rope_dims > head_dim || (rope_dims % 2) != 0) { - // GGML_LOG_WARN("OpenVINO backend does not support ROPE with n_dims %d and src[0]->ne[0] %ld\n", n_dims, - // op->src[0]->ne[0]); - return true; + return {false, "ROPE with n_dims=" + std::to_string(n_dims) + ", head_dim=" + std::to_string(head_dim) + " is not supported"}; } if (op->type != GGML_TYPE_F32 && op->type != GGML_TYPE_F16) { - // GGML_LOG_WARN("OpenVINO backend does not support ROPE with type %s\n", ggml_type_name(op->type)); - return true; + return {false, "ROPE with type " + std::string(ggml_type_name(op->type)) + " is not supported"}; } if (op->src[0]->op == GGML_OP_VIEW) { - if (op->src[0]->view_src->ne[1] != op->src[0]->ne[2]) { - // GGML_LOG_WARN( - // "OpenVINO backend does not support ROPE with src[0]->view_src->ne[1] %ld != src[0]->ne[2] " - // "%ld\n", - // op->src[0]->view_src->ne[1], op->src[0]->ne[2]); - return true; + const struct ggml_tensor * view = op->src[0]; + const struct ggml_tensor * view_src = view->view_src; + if (view_src->ne[1] != view->ne[1] || view_src->ne[2] != view->ne[2] || view_src->ne[3] != view->ne[3]) { + return {false, "ROPE with view_src->ne [" + std::to_string(view_src->ne[1]) + ", " + + std::to_string(view_src->ne[2]) + ", " + std::to_string(view_src->ne[3]) + + "] != view->ne [" + std::to_string(view->ne[1]) + ", " + + std::to_string(view->ne[2]) + ", " + std::to_string(view->ne[3]) + + "] is not supported"}; } } if (mode == GGML_ROPE_TYPE_IMROPE && (op->src[2] != 0 || ((const float *) op_params)[6] != 1 || ((const float *) op_params)[7] != 0 || ((const float *) op_params)[8] != 1)) { - // GGML_LOG_WARN("OpenVINO backend does not support IMROPE with freq_factors, freq_scale, ext_factor, and attn_factor\n"); - return true; + return {false, "IMROPE with freq_factors, freq_scale, ext_factor, and attn_factor is not supported"}; } break; } case GGML_OP_TRANSPOSE: { - // if the type is bf16, will return true if (op->type == GGML_TYPE_BF16) { - // GGML_LOG_WARN("OpenVINO backend does not support CONT with BF16 type\n"); - return true; + return {false, "TRANSPOSE with BF16 type is not supported"}; } break; } case GGML_OP_REPEAT: { if (ggml_openvino_get_device_name() == "GPU" && op->type == GGML_TYPE_BF16) { - return true; + return {false, "REPEAT with BF16 type is not supported on GPU"}; } break; } @@ -1285,15 +1318,15 @@ static bool is_op_unsupported_case(const ggml_tensor * op) { // return true; // } if (op->src[2]->op == GGML_OP_PERMUTE) { - return true; + return {false, "GATED_DELTA_NET with PERMUTE src2 is not supported"}; } // kda (per-key-dimension gating) not supported by fused GatedDeltaNet op if (op->src[3]->ne[0] != 1) { - return true; + return {false, "GATED_DELTA_NET with kda (per-key-dimension gating) is not supported"}; } // K > 1 (multiple state snapshots) not supported by fused op if (((const int32_t *) op->op_params)[0] > 1) { - return true; + return {false, "GATED_DELTA_NET with K > 1 (multiple state snapshots) is not supported"}; } break; } @@ -1307,17 +1340,17 @@ static bool is_op_unsupported_case(const ggml_tensor * op) { // Skip TOPK_MOE fused tests until it is fully supported. // The argsort_top_k VIEW wrapping ARGSORT is named "selected_experts" in test_topk_moe. if (strcmp(op->name, "selected_experts") == 0) { - return true; + return {false, "VIEW for selected_experts (argsort_top_k) is not supported"}; } break; } default: break; } - return false; + return {true, ""}; } -static bool ggml_backend_openvino_device_supports_op(ggml_backend_dev_t dev, const ggml_tensor * op) { +static ggml_openvino_op_support ggml_backend_openvino_device_supports_op_impl(ggml_backend_dev_t dev, const ggml_tensor * op) { GGML_ASSERT(dev->reg != nullptr); static std::unordered_set supported_types{ @@ -1367,48 +1400,41 @@ static bool ggml_backend_openvino_device_supports_op(ggml_backend_dev_t dev, con case GGML_OP_UNARY: { auto supported = supported_unary_ops.find(ggml_get_unary_op(op)) != supported_unary_ops.end(); if (!supported) { - // GGML_LOG_WARN("OpenVINO backend does not support unary op %s\n", ggml_unary_op_name(ggml_get_unary_op(op))); - return false; + return {false, "unary op " + std::string(ggml_unary_op_name(ggml_get_unary_op(op))) + " has no op translator"}; } if (ggml_get_unary_op(op) == GGML_UNARY_OP_EXP && op->type == GGML_TYPE_F32) { - return false; + return {false, "UNARY_EXP with F32 type is not supported"}; } break; } case GGML_OP_GLU: { auto supported = supported_glu_ops.find(ggml_get_glu_op(op)) != supported_glu_ops.end(); if (!supported) { - // GGML_LOG_WARN("OpenVINO backend does not support GLU op %s\n", ggml_glu_op_name(ggml_get_glu_op(op))); - return false; + return {false, "GLU op " + std::string(ggml_glu_op_name(ggml_get_glu_op(op))) + " has no op translator"}; } // if (has_view_op_input(op)) { - // // GGML_LOG_WARN("OpenVINO backend does not support unary op %s with view input\n", - // // ggml_glu_op_name(ggml_get_glu_op(op))); - // return false; + // return {false, "GLU op " + std::string(ggml_glu_op_name(ggml_get_glu_op(op))) + " with view input is not supported"}; // } if (op->src[1] == nullptr && op->src[0]->ne[0] % 2 != 0) { // triggers bug in ov gpu - return false; + return {false, "GLU op with odd src0 ne[0] and null src1 is not supported"}; } break; } default: { auto supported = supported_ops.find(op->op) != supported_ops.end(); if (!supported) { - // GGML_LOG_WARN("OpenVINO backend does not support op %s\n", ggml_op_name(op->op)); - return false; + return {false, "op " + std::string(ggml_op_name(op->op)) + " has no op translator"}; } static std::set ops_not_support_view_input{}; if (ops_not_support_view_input.find(op->op) != ops_not_support_view_input.end() && has_view_op_input(op)) { - // GGML_LOG_WARN("OpenVINO backend does not support op %s with view input\n", ggml_op_name(op->op)); - return false; + return {false, "op " + std::string(ggml_op_name(op->op)) + " with VIEW input is not supported"}; } } } if (supported_types.find(op->type) == supported_types.end()) { - // GGML_LOG_WARN("OpenVINO backend does not support tensor type %s\n", ggml_type_name(op->type)); - return false; + return {false, "tensor type " + std::string(ggml_type_name(op->type)) + " is not supported"}; } for (int i = 0; i < GGML_MAX_SRC; i++) { auto * src = op->src[i]; @@ -1416,21 +1442,32 @@ static bool ggml_backend_openvino_device_supports_op(ggml_backend_dev_t dev, con break; } if (supported_types.find(src->type) == supported_types.end()) { - // GGML_LOG_WARN("OpenVINO backend does not support tensor type %s\n", ggml_type_name(src->type)); - return false; + return {false, "src[" + std::to_string(i) + "] type " + std::string(ggml_type_name(src->type)) + " is not supported"}; } const bool is_supported_3d_moe_expert = op->op == GGML_OP_MUL_MAT_ID && i == 0 && (src->type == GGML_TYPE_MXFP4 || src->ne[3] == 1); if (ggml_is_quantized(src->type) && src->ne[2] != 1 && !is_supported_3d_moe_expert) { - // GGML_LOG_WARN("OpenVINO backend does not support 3D quantized tensors\n"); - return false; + return {false, "3D quantized tensor for src[" + std::to_string(i) + "] is not supported"}; } } - if (is_op_unsupported_case(op)) { - return false; + auto op_support_case = is_op_supported_case(op); + if (!op_support_case.is_supported) { + return op_support_case; } - return true; + return {true, ""}; +} + +static bool ggml_backend_openvino_device_supports_op(ggml_backend_dev_t dev, const ggml_tensor * op) { + auto res = ggml_backend_openvino_device_supports_op_impl(dev, op); + if (!res.is_supported) { + static const bool log_unsupported = ggml_openvino_getenv_int("GGML_OPENVINO_LOG_UNSUPPORTED_OPS") != 0; + if (log_unsupported) { + GGML_LOG_WARN("OpenVINO op unsupported: op '%s' (%s), type %s: %s\n", + op->name, ggml_op_name(op->op), ggml_type_name(op->type), res.reason.c_str()); + } + } + return res.is_supported; } static bool ggml_backend_openvino_device_supports_buft(ggml_backend_dev_t dev, ggml_backend_buffer_type_t buft) { diff --git a/ggml/src/ggml-openvino/openvino/op/cpy.cpp b/ggml/src/ggml-openvino/openvino/op/cpy.cpp index 5b387fc50d3..6f1e34779ac 100644 --- a/ggml/src/ggml-openvino/openvino/op/cpy.cpp +++ b/ggml/src/ggml-openvino/openvino/op/cpy.cpp @@ -3,8 +3,11 @@ #include "../utils.h" #include +#include +#include #include -#include +#include +#include #include #include #include @@ -12,9 +15,14 @@ #include #include #include +#include #include +#include #include #include +#include +#include +#include namespace ov { namespace frontend { @@ -61,10 +69,27 @@ OutputVector translate_cpy(const NodeContext & context) { return rename_outputs_with_suffix({res}, context.get_name()); } - // Recurrent state cache writeback into a slot block of the cache. Where the block starts and - // where the copied data starts in the source are runtime inputs, so the cached model works for - // any kv head, active sequence count and token count. The result is the full updated cache. + // Recurrent state cache writeback into a slot block of the cache. Where the block starts is a + // runtime input, so the cached model works for any kv head and active sequence count. The + // result is the full updated cache. // op_case 1: gated-delta-net state, op_case 2: conv state, op_case 3: defrag remainder. + if (op_case == 3) { + // With -np 1 (and generally whenever there is no defrag remainder) this GET_ROWS gathers + // zero rows: nothing to write back, and the cache is unchanged. NPU rejects zero-size + // tensors, so short-circuit instead of building a degenerate Slice/Concat chain. + bool is_empty = false; + if (input_shape.rank().is_static()) { + for (const auto & d : input_shape) { + if (d.is_static() && d.get_length() == 0) { + is_empty = true; + break; + } + } + } + if (is_empty) { + return {context.get_input(1)}; + } + } const std::string slot_begin_name = "rs_slot_begin_" + context.get_name(); const bool slice_assign = context.has_input(slot_begin_name) && !context.is_stateful() && (op_case >= 1 && op_case <= 3); @@ -81,19 +106,49 @@ OutputVector translate_cpy(const NodeContext & context) { ov::Output begin = context.get_input(slot_begin_name); auto base = context.get_input(1); if (op_case == 1) { - // GDN packs [attn | state snapshots]; the state part runs from src_begin to the end. - auto src_begin = context.get_input("rs_src_begin_" + context.get_name()); - auto state_part = std::make_shared(context.get_input(0), src_begin, int_max, one, axis); + ov::Output state_begin; + const std::string src_begin_name = "rs_src_begin_" + context.get_name(); + if (context.has_input(src_begin_name)) { + state_begin = context.get_input(src_begin_name); + } else { + auto ssm_state_size = context.get_ssm_state_size(); + if (context.has_input("s_copy_active_slot_len")) { + auto len = context.get_input("s_copy_active_slot_len"); + auto state_rows = std::make_shared( + ov::op::v0::Constant::create(ov::element::i64, {1}, {ssm_state_size}), len); + state_begin = std::make_shared(state_rows); + } else { + state_begin = ov::op::v0::Constant::create(ov::element::i64, {1}, {-ssm_state_size}); + } + } + auto state_part = + std::make_shared(context.get_input(0), state_begin, int_max, one, axis); src = std::make_shared(state_part, feature, false); } else if (op_case == 2) { - // conv_input is [previous conv state | new tokens]; copy the conv_kernel_size - 1 wide - // window starting at src_begin, which is the snapshot this writeback corresponds to. + // conv_input is [previous conv state | new tokens]; the snapshot is the conv_kernel_size - 1 + // columns ending at the last *valid* token. Gather (rather than Slice) keeps the output + // shape static even though the window start is a runtime value. auto window_size = (int64_t) input_shape[3].get_length(); - auto src_begin = context.get_input("rs_src_begin_" + context.get_name()); - auto src_end = std::make_shared( - src_begin, ov::op::v0::Constant::create(ov::element::i64, {1}, {window_size})); - auto window = std::make_shared(context.get_input(0), src_begin, src_end, one, - ov::op::v0::Constant::create(ov::element::i64, {1}, {3})); + ov::Output window; + auto col_axis = ov::op::v0::Constant::create(ov::element::i64, {1}, {3}); + const std::string src_begin_name = "rs_src_begin_" + context.get_name(); + if (context.has_input(src_begin_name)) { + auto src_begin = context.get_input(src_begin_name); + auto src_end = std::make_shared( + src_begin, ov::op::v0::Constant::create(ov::element::i64, {1}, {window_size})); + window = std::make_shared(context.get_input(0), src_begin, src_end, one, col_axis); + } else if (context.has_input("chunk_valid_len")) { + std::vector offsets(window_size); + std::iota(offsets.begin(), offsets.end(), 0); + auto indices = std::make_shared( + ov::op::v0::Constant::create(ov::element::i64, {(size_t) window_size}, offsets), + context.get_input("chunk_valid_len")); + window = std::make_shared(context.get_input(0), indices, col_axis); + } else { + auto window_begin = ov::op::v0::Constant::create(ov::element::i64, {1}, {-window_size}); + window = + std::make_shared(context.get_input(0), window_begin, int_max, one, col_axis); + } const auto base_shape = base.get_partial_shape(); FRONT_END_OP_CONVERSION_CHECK(base_shape.rank().is_static() && base_shape.rank().get_length() == 4, "CPY conv state cache update requires rank-4 base cache"); @@ -157,6 +212,63 @@ OutputVector translate_cpy(const NodeContext & context) { auto input = process_view_input_new(context, 0); + if (op_case == 5 || op_case == 6) { + auto input_shape = context.get_input_shape(0); + auto output_shape = context.get_output_shape(); + auto dst_ggml_shape = context.get_view_input_ggml_shape(1, 0); + auto dst_stride = context.get_view_input_stride(1, 0); + size_t offset_bytes = context.get_view_input_offset(1, 0); + auto n_state = (int64_t) context.get_input_shape(0)[3].get_length(); + auto n_state_c = ov::op::v0::Constant::create(ov::element::i64, {1}, {n_state}); + auto kv_buf = context.get_input(1); // shape {1,1,1,N} + + Output token_len_per_seq; + Output n_write_dyn; + if (context.has_input("token_len_per_seq")) { + token_len_per_seq = context.get_input("token_len_per_seq"); + n_write_dyn = std::make_shared(token_len_per_seq, n_state_c); + } else { + n_write_dyn = ov::op::v0::Constant::create(ov::element::i64, {1}, {(int64_t) dst_ggml_shape[3]}); + } + size_t elem_size = dst_stride[3]; + FRONT_END_OP_CONVERSION_CHECK(elem_size > 0, "CPY KV cache view update has invalid element size"); + int64_t start_elem = (int64_t) (offset_bytes / elem_size); + // op_case 5: decoder self-attention – write offset advances each step. + // op_case 6: encoder self-attn or cross-attn – offset fixed at compile time. + const bool is_decoder_self_attn = (op_case == 5); + auto ones_c = ov::op::v0::Constant::create(ov::element::i64, {3}, std::vector{1, 1, 1}); + auto new_shape = std::make_shared(ov::OutputVector{ones_c, n_write_dyn}, 0); + + auto reshaped = std::make_shared(input, new_shape, false); + auto data = std::make_shared(reshaped, context.get_output_type()); + // Indices [start_elem .. start_elem + n_write) on axis 3 of {1,1,1,N} + // For decoder self-attention the write offset advances each step, so compute it + // dynamically from the model inputs: start = (attention_size - token_len_per_seq) * n_state. + // For encoder self-attn and cross-attn the offset is fixed at graph-compile time. + ov::Output start; + if (is_decoder_self_attn && context.has_input("attention_size") && context.has_input("token_len_per_seq")) { + auto attention_size_in = context.get_input("attention_size"); + auto token_len_in = context.get_input("token_len_per_seq"); + auto past_tokens = std::make_shared(attention_size_in, token_len_in); + auto new_start = std::make_shared(past_tokens, n_state_c); + start = std::make_shared( + new_start, ov::op::v0::Constant::create(ov::element::i64, {1}, {start_elem})); + } else { + start = ov::op::v0::Constant::create(ov::element::i64, {1}, {start_elem}); + } + auto start_squeezed = std::make_shared(start); + auto end = std::make_shared(start_squeezed, n_write_dyn); + auto end_squeezed = std::make_shared(end); + auto step = ov::op::v0::Constant::create(ov::element::i64, {1}, {1}); + auto step_squeezed = std::make_shared(step); + auto indices = + std::make_shared(start_squeezed, end_squeezed, step_squeezed, ov::element::i64); + auto axis = ov::op::v0::Constant::create(ov::element::i64, {1}, {3}); + + auto kv_updated = std::make_shared(kv_buf, indices, data, axis); + return rename_outputs_with_suffix({kv_updated}, context.get_name()); + } + if (input_shape != output_shape) { auto new_shape = ov::op::v0::Constant::create( ov::element::i64, {static_cast(output_shape.rank().get_length())}, output_shape.to_shape()); diff --git a/ggml/src/ggml-openvino/openvino/op/flash_attn_ext.cpp b/ggml/src/ggml-openvino/openvino/op/flash_attn_ext.cpp index 582df0130b5..06547f3d296 100644 --- a/ggml/src/ggml-openvino/openvino/op/flash_attn_ext.cpp +++ b/ggml/src/ggml-openvino/openvino/op/flash_attn_ext.cpp @@ -3,8 +3,8 @@ #include "../utils.h" #include "ggml-openvino/ggml-openvino-extra.h" +#include #include -#include #include #include #include @@ -15,6 +15,7 @@ #include #include #include +#include #include #include #include @@ -24,13 +25,62 @@ namespace ov { namespace frontend { namespace ggml { namespace op { +static ov::Output reshape_flat_kv(const ov::Output & kv_flat, + size_t view_offset_bytes, + size_t nb1_bytes, + int64_t n_head, + int64_t head_size, + const ov::Output & attention_size) { + int64_t n_state = n_head * head_size; + int64_t layer_start_elem = (int64_t) (view_offset_bytes / (nb1_bytes / n_state)); + // Dynamic slice: [layer_start_elem, layer_start_elem + n_kv * n_state) + auto start_c = ov::op::v0::Constant::create(ov::element::i64, {1}, {layer_start_elem}); + auto n_state_c = ov::op::v0::Constant::create(ov::element::i64, {1}, {n_state}); + // end = start + attention_size * n_state (both static + dynamic) + auto kv_len_elems = std::make_shared(attention_size, n_state_c); + auto end_c = std::make_shared(start_c, kv_len_elems); + auto step_c = ov::op::v0::Constant::create(ov::element::i64, {1}, {1}); + auto axis_c = ov::op::v0::Constant::create(ov::element::i64, {1}, {3}); + auto sliced = std::make_shared(kv_flat, start_c, end_c, step_c, axis_c); + + // KV cache is laid out as {n_kv, n_head, head_size} in memory + // Reshape to {1, n_kv, n_head, head_size}, then transpose to {1, n_head, n_kv, head_size} + // as required by SDPA. + auto one_c = ov::op::v0::Constant::create(ov::element::i64, {1}, {1}); + auto n_head_c = ov::op::v0::Constant::create(ov::element::i64, {1}, {n_head}); + auto head_size_c = ov::op::v0::Constant::create(ov::element::i64, {1}, {head_size}); + // reshape: {n_kv*n_state} -> {1, n_kv, n_head, head_size} + auto new_shape = + std::make_shared(ov::OutputVector{one_c, attention_size, n_head_c, head_size_c}, 0); + auto reshaped = std::make_shared(sliced, new_shape, false); + // transpose: {1, n_kv, n_head, head_size} -> {1, n_head, n_kv, head_size} + auto perm = ov::op::v0::Constant::create(ov::element::i64, {4}, {0, 2, 1, 3}); + auto ret = std::make_shared(reshaped, perm); + return ret; +} OutputVector translate_flash_attn_ext(const NodeContext & context) { - num_inputs_check(context, 4, 4); + num_inputs_check(context, 3, 4); + const bool has_mask = context.get_input_size() == 4; auto q_f32 = context.get_input(0); auto k = context.get_input(1); auto v = context.get_input(2); - auto mask = context.get_input(3); + const int op_case = context.get_op_case(); + + if (op_case == 1 || op_case == 2) { + int64_t n_state_head = (int64_t) context.get_view_input_ggml_shape(1, 0)[3]; + int64_t n_head = (int64_t) context.get_view_input_ggml_shape(1, 0)[1]; + size_t nb1 = context.get_view_input_stride(1, 0)[2]; + size_t offset = context.get_view_input_offset(1, 0); + ov::Output attention_size; + if (op_case == 1) { + attention_size = context.get_input("attention_size"); + } else { + attention_size = context.get_input("attention_size_static"); + } + k = reshape_flat_kv(k, offset, nb1, n_head, n_state_head, attention_size); + v = reshape_flat_kv(v, offset, nb1, n_head, n_state_head, attention_size); + } float * params = reinterpret_cast(context.get_output_op_params()); float scale = params[0]; @@ -43,16 +93,19 @@ OutputVector translate_flash_attn_ext(const NodeContext & context) { ov::Output res; // For stateful - std::string mask_name = "KQ_mask_sliced"; - if (context.get_input_names()[3].find("swa") != std::string::npos) { - mask_name = "KQ_mask_swa_sliced"; - } - if (context.has_input(mask_name)) { - mask = context.get_input(mask_name); - } - - if (mask.get_element_type() != ov::element::f16) { - mask = std::make_shared(mask, ov::element::f16); + ov::Output mask; + if (has_mask) { + mask = context.get_input(3); + std::string mask_name = "KQ_mask_sliced"; + if (context.get_input_names()[3].find("swa") != std::string::npos) { + mask_name = "KQ_mask_swa_sliced"; + } + if (context.has_input(mask_name)) { + mask = context.get_input(mask_name); + } + if (mask.get_element_type() != ov::element::f16) { + mask = std::make_shared(mask, ov::element::f16); + } } //auto tile_kv = [&](int64_t num_heads, int64_t num_heads_kv, int64_t head_size, ov::Output kv) { @@ -108,10 +161,14 @@ OutputVector translate_flash_attn_ext(const NodeContext & context) { // get [B, 1, 1, S_q, S_k], which NUMPY-broadcasts cleanly against the // [B, num_heads_kv, factor, S_q, S_k] scores: B==B, then 1→num_heads_kv and // 1→factor on the head dims. - auto mask_unsq1 = - std::make_shared(mask, ov::op::v0::Constant::create(ov::element::i64, {1}, {2})); - // mask_unsq1: [B, 1, 1, S_q, S_k] (rank 5) - ov::Output qk_masked = std::make_shared(qk_scaled, mask_unsq1); + ov::Output qk_masked; + if (has_mask) { + auto mask_unsq1 = + std::make_shared(mask, ov::op::v0::Constant::create(ov::element::i64, {1}, {2})); + qk_masked = std::make_shared(qk_scaled, mask_unsq1); + } else { + qk_masked = qk_scaled; + } auto softmax = std::make_shared(qk_masked, /*axis=*/-1); @@ -164,9 +221,16 @@ OutputVector translate_flash_attn_ext(const NodeContext & context) { k = tile_kv(num_heads, num_heads_kv, head_size, k); v = tile_kv(num_heads, num_heads_kv, head_size, v); - auto sdpa = std::make_shared(q, k, v, mask, scale_node, false); - res = std::make_shared(sdpa, - ov::op::v0::Constant::create(ov::element::i64, {4}, {0, 2, 1, 3})); + constexpr auto causal = false; + if (has_mask) { + auto sdpa = std::make_shared(q, k, v, mask, scale_node, causal); + res = std::make_shared( + sdpa, ov::op::v0::Constant::create(ov::element::i64, {4}, {0, 2, 1, 3})); + } else { + auto sdpa = std::make_shared(q, k, v, scale_node, causal); + res = std::make_shared( + sdpa, ov::op::v0::Constant::create(ov::element::i64, {4}, {0, 2, 1, 3})); + } res = std::make_shared(res, ov::element::f32); return rename_outputs_with_suffix({res}, context.get_name()); } diff --git a/ggml/src/ggml-openvino/openvino/op/gated_delta_net.cpp b/ggml/src/ggml-openvino/openvino/op/gated_delta_net.cpp index 66c74828331..07eeb3c8fd6 100644 --- a/ggml/src/ggml-openvino/openvino/op/gated_delta_net.cpp +++ b/ggml/src/ggml-openvino/openvino/op/gated_delta_net.cpp @@ -7,12 +7,15 @@ #include #include #include +#include #include #include #include #include +#include #include #include +#include #include #include #include @@ -80,6 +83,28 @@ OutputVector translate_gated_delta_net(const NodeContext & context) { g = std::make_shared(g, ov::op::v0::Constant::create(ov::element::i64, {1}, {3})); beta = std::make_shared(beta, ov::op::v0::Constant::create(ov::element::i64, {1}, {3})); + if (context.has_input("chunk_valid_len")) { + // The last prefill chunk is padded with fabricated tokens. The recurrence is + // S_t = S_{t-1} * exp(g_t) + k_t (x) ((v_t - S_{t-1}^T k_t) * beta_t) + // so forcing g = 0 and beta = 0 makes a padded step an exact identity and keeps the final + // state equal to the state after the last real token. Attention output at those positions + // is garbage but never read. + const auto & g_ps = g.get_partial_shape(); + FRONT_END_OP_CONVERSION_CHECK(g_ps.rank().is_static() && g_ps.rank().get_length() == 3 && g_ps[1].is_static(), + "GATED_DELTA_NET pad masking requires a static token dimension"); + const int64_t n_tokens = g_ps[1].get_length(); + std::vector positions(n_tokens); + std::iota(positions.begin(), positions.end(), 0); + auto valid = std::make_shared( + ov::op::v0::Constant::create(ov::element::i64, {(size_t) n_tokens}, positions), + context.get_input("chunk_valid_len")); + auto mask = std::make_shared( + std::make_shared(valid, g.get_element_type()), + ov::op::v0::Constant::create(ov::element::i64, {2}, std::vector{0, 2})); + g = std::make_shared(g, mask); + beta = std::make_shared(beta, mask); + } + // std::cout << "GatedDeltaNet input shapes: q=" << q.get_partial_shape() << ", k=" << k.get_partial_shape() // << ", v=" << v.get_partial_shape() << ", g=" << g.get_partial_shape() // << ", beta=" << beta.get_partial_shape() << ", state=" << state.get_partial_shape() << std::endl; diff --git a/ggml/src/ggml-openvino/openvino/op/glu_geglu_quick.cpp b/ggml/src/ggml-openvino/openvino/op/glu_geglu_quick.cpp new file mode 100644 index 00000000000..c6d64aed43a --- /dev/null +++ b/ggml/src/ggml-openvino/openvino/op/glu_geglu_quick.cpp @@ -0,0 +1,64 @@ +#include "../node_context.h" +#include "../op_table.h" +#include "../utils.h" + +#include +#include +#include +#include +#include +#include + +namespace ov { +namespace frontend { +namespace ggml { +namespace op { + +OutputVector translate_glu_geglu_quick(const NodeContext & context) { + num_inputs_check(context, 1, 2); + + ov::Output src0; + ov::Output src1; + if (context.get_input_size() == 2) { + src0 = process_view_input_new(context, 0); + src1 = process_view_input_new(context, 1); + } else { + // split along last axis, nc = ne[0] / 2 + auto combined = process_view_input_new(context, 0); + auto combined_shape = combined.get_partial_shape(); + int64_t last_dim_val = combined_shape[combined_shape.rank().get_length() - 1].get_length(); + int64_t nc = last_dim_val / 2; + + auto axis = ov::op::v0::Constant::create(ov::element::i64, {1}, {-1}); + auto step = ov::op::v0::Constant::create(ov::element::i64, {1}, {1}); + auto start0 = ov::op::v0::Constant::create(ov::element::i64, {1}, {0}); + auto stop0 = ov::op::v0::Constant::create(ov::element::i64, {1}, {nc}); + auto start1 = ov::op::v0::Constant::create(ov::element::i64, {1}, {nc}); + auto stop1 = ov::op::v0::Constant::create(ov::element::i64, {1}, {2 * nc}); + + src0 = std::make_shared(combined, start0, stop0, step, axis); + src1 = std::make_shared(combined, start1, stop1, step, axis); + } + + int32_t * params = context.get_output_op_params(); + const int32_t swapped = params[1]; + if (swapped) { + std::swap(src0, src1); + } + + // GELU_QUICK(x) = x * sigmoid(1.702 * x) + // Create the constant in the same type as src0 to avoid f16/f32 mismatch. + auto input_type = src0.get_element_type(); + auto coef = ov::op::v0::Constant::create(input_type, ov::Shape{}, {1.702f}); + auto scaled = std::make_shared(src0, coef); + auto sigmoid = std::make_shared(scaled); + auto gated = std::make_shared(src0, sigmoid); + auto res = std::make_shared(gated, src1); + + return rename_outputs_with_suffix({res}, context.get_name()); +} + +} // namespace op +} // namespace ggml +} // namespace frontend +} // namespace ov diff --git a/ggml/src/ggml-openvino/openvino/op/pool_2d.cpp b/ggml/src/ggml-openvino/openvino/op/pool_2d.cpp new file mode 100644 index 00000000000..fb6333175f0 --- /dev/null +++ b/ggml/src/ggml-openvino/openvino/op/pool_2d.cpp @@ -0,0 +1,53 @@ +#include "../node_context.h" +#include "../op_table.h" +#include "../utils.h" + +#include +#include +#include + +namespace ov { +namespace frontend { +namespace ggml { +namespace op { + +OutputVector translate_pool_2d(const NodeContext & context) { + num_inputs_check(context, 1, 1); + const int32_t * params = context.get_output_op_params(); + + const int k0 = params[1]; + const int k1 = params[2]; + const int s0 = params[3]; + const int s1 = params[4]; + const int p0 = params[5]; + const int p1 = params[6]; + + const int op_case = context.get_op_case(); + ov::Output input = context.get_input(0); + ov::Strides strides{static_cast(s1), static_cast(s0)}; + ov::Shape pads_begin{static_cast(p1), static_cast(p0)}; + ov::Shape pads_end{static_cast(p1), static_cast(p0)}; + ov::Shape kernel{static_cast(k1), static_cast(k0)}; + ov::Output res; + + switch (op_case) { + case 1: // GGML_OP_POOL_MAX + { + res = std::make_shared(input, strides, pads_begin, pads_end, kernel); + break; + } + case 2: // GGML_OP_POOL_AVG + { + res = std::make_shared(input, strides, pads_begin, pads_end, kernel, false); + break; + } + default: + break; + } + return rename_outputs_with_suffix({res}, context.get_name()); +} + +} // namespace op +} // namespace ggml +} // namespace frontend +} // namespace ov diff --git a/ggml/src/ggml-openvino/openvino/op/roll.cpp b/ggml/src/ggml-openvino/openvino/op/roll.cpp new file mode 100644 index 00000000000..e8d1b8e50b3 --- /dev/null +++ b/ggml/src/ggml-openvino/openvino/op/roll.cpp @@ -0,0 +1,36 @@ +#include "../node_context.h" +#include "../op_table.h" +#include "../utils.h" + +#include +#include + +namespace ov { +namespace frontend { +namespace ggml { +namespace op { + +OutputVector translate_roll(const NodeContext & context) { + num_inputs_check(context, 1, 1); + const int32_t * params = context.get_output_op_params(); + + int64_t s0 = params[0]; + int64_t s1 = params[1]; + int64_t s2 = params[2]; + int64_t s3 = params[3]; + + auto input = context.get_input(0); + + auto shift = ov::op::v0::Constant::create( + ov::element::i64, ov::Shape{4}, std::vector{s3, s2, s1, s0}); + auto axes = ov::op::v0::Constant::create( + ov::element::i64, ov::Shape{4}, std::vector{0, 1, 2, 3}); + + auto roll = std::make_shared(input, shift, axes); + return rename_outputs_with_suffix({roll}, context.get_name()); +} + +} // namespace op +} // namespace ggml +} // namespace frontend +} // namespace ov diff --git a/ggml/src/ggml-openvino/openvino/op/view.cpp b/ggml/src/ggml-openvino/openvino/op/view.cpp index 138526cb49c..56f5ceec9bb 100644 --- a/ggml/src/ggml-openvino/openvino/op/view.cpp +++ b/ggml/src/ggml-openvino/openvino/op/view.cpp @@ -17,6 +17,13 @@ namespace op { OutputVector translate_view(const NodeContext & context) { num_inputs_check(context, 1, 1); + if (context.get_op_case() == 1) { + // Static-mode identity pass-through for VIEWs over a GATED_DELTA_NET combined output or + // the conv_input CONCAT; the consuming op (CPY/RMS_NORM) does its own runtime-correct + // slicing on the full tensor (see ggml-decoder.cpp compute_op_case, GGML_OP_VIEW). + return {context.get_input(0)}; + } + if (!context.is_static()) { // On the stateless/non-static path VIEW is normally a no-op (consumers re-slice). // EXCEPTION: the MoE expert aggregation slices each expert plane out of diff --git a/ggml/src/ggml-openvino/openvino/op_table.cpp b/ggml/src/ggml-openvino/openvino/op_table.cpp index 3c26fe83b1a..9c9d8eeac78 100644 --- a/ggml/src/ggml-openvino/openvino/op_table.cpp +++ b/ggml/src/ggml-openvino/openvino/op_table.cpp @@ -10,6 +10,7 @@ #include #include #include +#include #include #include #include @@ -55,10 +56,12 @@ std::unordered_map get_supported_ops() { {"GGML_UNARY_OP_SIGMOID", op::translate_1to1_match_1_input }, {"GGML_UNARY_OP_EXP", op::translate_1to1_match_1_input }, {"GGML_UNARY_OP_NEG", op::translate_1to1_match_1_input }, + {"GGML_UNARY_OP_RELU", op::translate_1to1_match_1_input }, {"GGML_OP_VIEW", op::translate_view }, {"GGML_GLU_OP_SWIGLU", op::translate_glu_swiglu }, {"GGML_GLU_OP_SWIGLU_OAI", op::translate_glu_swiglu_oai }, {"GGML_GLU_OP_GEGLU", op::translate_glu_geglu }, + {"GGML_GLU_OP_GEGLU_QUICK", op::translate_glu_geglu_quick }, {"GGML_OP_SET_ROWS", op::translate_set_rows }, {"GGML_OP_CPY", op::translate_cpy }, {"GGML_OP_FLASH_ATTN_EXT", op::translate_flash_attn_ext }, @@ -72,6 +75,8 @@ std::unordered_map get_supported_ops() { {"GGML_OP_DIAG", op::translate_diag }, {"GGML_OP_TRI", op::translate_tri }, {"GGML_OP_SET", op::translate_set }, + {"GGML_OP_POOL_2D", op::translate_pool_2d }, + {"GGML_OP_ROLL", op::translate_roll }, // solve_tri has accuracy issues on GPU // {"GGML_OP_SOLVE_TRI", op::translate_solve_tri }, }; diff --git a/ggml/src/ggml-openvino/openvino/op_table.h b/ggml/src/ggml-openvino/openvino/op_table.h index d4b9292d637..0a81a57a667 100644 --- a/ggml/src/ggml-openvino/openvino/op_table.h +++ b/ggml/src/ggml-openvino/openvino/op_table.h @@ -38,6 +38,7 @@ GGML_OP_CONVERTER(translate_view); GGML_OP_CONVERTER(translate_glu_swiglu); GGML_OP_CONVERTER(translate_glu_swiglu_oai); GGML_OP_CONVERTER(translate_glu_geglu); +GGML_OP_CONVERTER(translate_glu_geglu_quick); GGML_OP_CONVERTER(translate_set_rows); GGML_OP_CONVERTER(translate_cpy); GGML_OP_CONVERTER(translate_argsort); @@ -53,6 +54,8 @@ GGML_OP_CONVERTER(translate_set); GGML_OP_CONVERTER(translate_diag); GGML_OP_CONVERTER(translate_tri); GGML_OP_CONVERTER(translate_solve_tri); +GGML_OP_CONVERTER(translate_pool_2d); +GGML_OP_CONVERTER(translate_roll); } // namespace op diff --git a/ggml/src/ggml-openvino/openvino/pass/fuse_to_conv.cpp b/ggml/src/ggml-openvino/openvino/pass/fuse_to_conv.cpp new file mode 100644 index 00000000000..21801c0f399 --- /dev/null +++ b/ggml/src/ggml-openvino/openvino/pass/fuse_to_conv.cpp @@ -0,0 +1,212 @@ +#include "fuse_to_conv.h" + +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include + +namespace opp = ov::pass::pattern; + +namespace ov { +namespace frontend { +namespace ggml { +namespace pass { + +// This pass fuses an IM2COL + MatMul convolution into OpenVINO's Convolution op for performance gains. +// Reference the im2col.cpp translator for reference on the pattern being matched. + +FuseToConv::FuseToConv() { + const auto m_wei = opp::any_input(); + const auto m_act = opp::any_input(); + const auto m_matmul = opp::wrap_type({m_wei, m_act}); + + const auto callback = [=](ov::pass::pattern::Matcher & m) { + const auto & pm = m.get_pattern_value_map(); + + auto matmul_node = ov::as_type_ptr(pm.at(m_matmul).get_node_shared_ptr()); + if (!matmul_node || matmul_node->get_transpose_a() || !matmul_node->get_transpose_b()) { + return false; + } + + auto trace = matmul_node->input_value(1); + + // Optional Convert + if (auto n = ov::as_type_ptr(trace.get_node_shared_ptr())) { + trace = n->input_value(0); + } + + for (int i = 0; i < 2; ++i) { + auto n = ov::as_type_ptr(trace.get_node_shared_ptr()); + if (!n) { + return false; + } + trace = n->input_value(0); + } + + if (auto n = ov::as_type_ptr(trace.get_node_shared_ptr())) { + trace = n->input_value(0); + } else { + return false; + } + + if (auto n = ov::as_type_ptr(trace.get_node_shared_ptr())) { + trace = n->input_value(0); + } else { + return false; + } + + if (auto n = ov::as_type_ptr(trace.get_node_shared_ptr())) { + trace = n->input_value(0); + } else { + return false; + } + + auto eip = ov::as_type_ptr(trace.get_node_shared_ptr()); + if (!eip) { + return false; + } + const auto eip_strides = eip->get_strides(); // {stride_h, stride_w} + const auto eip_rates = eip->get_rates(); // {dil_h, dil_w} + + auto pad = ov::as_type_ptr(eip->input_value(0).get_node_shared_ptr()); + if (!pad) { + return false; + } + auto pads_begin_const = + ov::as_type_ptr(pad->input_value(1).get_node_shared_ptr()); + + const auto pads_begin_vals = pads_begin_const->cast_vector(); // {0, 0, pad_h, pad_w} + const std::ptrdiff_t pad_h = static_cast(pads_begin_vals[2]); + const std::ptrdiff_t pad_w = static_cast(pads_begin_vals[3]); + + auto image_input = pad->input_value(0); // [N, IC, 1, IW] NCHW + + auto w_trace = matmul_node->input_value(0); + if (auto n = ov::as_type_ptr(w_trace.get_node_shared_ptr())) { + w_trace = n->input_value(0); + } + for (int i = 0; i < 2; ++i) { + auto n = ov::as_type_ptr(w_trace.get_node_shared_ptr()); + if (!n) { + break; + } + w_trace = n->input_value(0); + } + + auto weight_const = ov::as_type_ptr(w_trace.get_node_shared_ptr()); + if (!weight_const) { + return false; + } + + // Reshape weight to [OC, IC, 1, KW] (OIHW). + const auto w_shape = weight_const->get_shape(); + ov::Shape conv_w_shape; + if (w_shape.size() == 3) { + conv_w_shape = {w_shape[0], w_shape[1], 1, w_shape[2]}; + } else if (w_shape.size() == 4) { + conv_w_shape = {w_shape[1], w_shape[2], 1, w_shape[3]}; + } else { + return false; + } + + auto weight_reshaped = register_new_node(weight_const->get_element_type(), conv_w_shape, + weight_const->get_data_ptr()); + + ov::Output weight_input = weight_reshaped; + if (weight_reshaped->get_element_type() != image_input.get_element_type()) { + weight_input = register_new_node(weight_reshaped, image_input.get_element_type()); + } + + auto conv = register_new_node( + image_input, weight_input, + ov::Strides{static_cast(eip_strides[0]), static_cast(eip_strides[1])}, + ov::CoordinateDiff{pad_h, pad_w}, ov::CoordinateDiff{pad_h, pad_w}, + ov::Strides{static_cast(eip_rates[0]), static_cast(eip_rates[1])}, + ov::op::PadType::EXPLICIT); + + constexpr auto target_type = ov::element::f32; + ov::Output conv_out = conv; + if (conv_out.get_element_type() != target_type) { + conv_out = register_new_node(conv_out, target_type); + } + + std::shared_ptr add_node; + ov::Output bias_input; + for (const auto & consumer_in : matmul_node->output(0).get_target_inputs()) { + auto cast = ov::as_type_ptr(consumer_in.get_node()->shared_from_this()); + if (!cast) { + continue; + } + for (const auto & add_in : cast->output(0).get_target_inputs()) { + auto add = ov::as_type_ptr(add_in.get_node()->shared_from_this()); + if (!add) { + continue; + } + for (size_t i = 0; i < 2; ++i) { + if (ov::as_type_ptr(add->input_value(i).get_node_shared_ptr())) { + bias_input = add->input_value(i); + add_node = add; + break; + } + } + if (add_node) { + break; + } + } + if (add_node) { + break; + } + } + + ov::Output final_out; + std::shared_ptr target_node; + + if (add_node) { + // Reshape bias [OC, 1] → [1, OC, 1, 1] for NCHW broadcasting. + ov::Output bias = bias_input; + if (bias.get_element_type() != target_type) { + bias = register_new_node(bias, target_type); + } + const auto oc = static_cast(conv_w_shape[0]); + auto bias_shape = register_new_node(ov::element::i64, ov::Shape{4}, + std::vector{1, oc, 1, 1}); + bias = register_new_node(bias, bias_shape, false); + final_out = register_new_node(conv_out, bias); + target_node = add_node; + } else { + final_out = conv_out; + target_node = matmul_node; + } + + // Reshape final output back to the target node's original shape if needed. + auto orig_shape = target_node->get_output_partial_shape(0); + if (orig_shape.is_static() && final_out.get_partial_shape() != orig_shape) { + auto shape_const = register_new_node(ov::element::i64, ov::Shape{orig_shape.size()}, + orig_shape.to_shape()); + final_out = register_new_node(final_out, shape_const, false); + } + + final_out.get_node_shared_ptr()->set_friendly_name(target_node->get_friendly_name()); + ov::copy_runtime_info(m.get_matched_nodes(), final_out.get_node_shared_ptr()); + ov::replace_node(target_node, final_out.get_node_shared_ptr()); + + return true; + }; + + register_matcher(std::make_shared(m_matmul, "ov::frontend::ggml::pass::FuseToConv"), callback); +} + +} // namespace pass +} // namespace ggml +} // namespace frontend +} // namespace ov diff --git a/ggml/src/ggml-openvino/openvino/pass/fuse_to_conv.h b/ggml/src/ggml-openvino/openvino/pass/fuse_to_conv.h new file mode 100644 index 00000000000..feac14b13ff --- /dev/null +++ b/ggml/src/ggml-openvino/openvino/pass/fuse_to_conv.h @@ -0,0 +1,17 @@ +#include "openvino/pass/matcher_pass.hpp" + +namespace ov { +namespace frontend { +namespace ggml { +namespace pass { + +class FuseToConv : public ov::pass::MatcherPass { +public: + OPENVINO_MATCHER_PASS_RTTI("ov::frontend::ggml::pass::FuseToConv") + FuseToConv(); +}; + +} // namespace pass +} // namespace ggml +} // namespace frontend +} // namespace ov diff --git a/ggml/src/ggml-openvino/openvino/translate_session.cpp b/ggml/src/ggml-openvino/openvino/translate_session.cpp index 35598aba6be..df3a72f3286 100644 --- a/ggml/src/ggml-openvino/openvino/translate_session.cpp +++ b/ggml/src/ggml-openvino/openvino/translate_session.cpp @@ -5,6 +5,7 @@ #include "ggml-openvino/openvino/node_context.h" #include "ggml-openvino/openvino/utils.h" #include "input_model.h" +#include "pass/fuse_to_conv.h" #include "pass/mark_decompression_convert_constant_folding.h" #include "pass/mark_dequantization_subgraph.h" #include "pass/squeeze_matmul.h" @@ -109,7 +110,8 @@ ov::pass::MakeStateful::ParamResPairs get_kv_param_res_pairs( void add_sliced_mask_stateful(TensorMap & tensor_map) { auto create_sliced_mask = [&](const std::string & mask_name, const std::string & sliced_name) { if ((tensor_map.find(mask_name) != tensor_map.end()) && - (tensor_map.find("token_len_per_seq") != tensor_map.end())) { + (tensor_map.find("token_len_per_seq") != tensor_map.end()) && + (tensor_map.find("inp_pos") != tensor_map.end())) { auto token_len_per_seq = tensor_map.at("token_len_per_seq").get_node_shared_ptr(); auto mask = tensor_map.at(mask_name).get_node_shared_ptr(); std::shared_ptr mask_sliced = mask; @@ -137,6 +139,7 @@ void add_sliced_mask_stateful(TensorMap & tensor_map) { }; create_sliced_mask("self_kq_mask", "KQ_mask_sliced"); + create_sliced_mask("KQ_mask", "KQ_mask_sliced"); create_sliced_mask("self_kq_mask_swa", "KQ_mask_swa_sliced"); } @@ -395,6 +398,7 @@ std::shared_ptr TranslateSession::apply_transformations(std::shared_ptr( std::vector{ov::element::u8, ov::element::i8, ov::element::u4, ov::element::i4}); + manager.register_pass(); if (ggml_model_decoder->is_stateful()) { const auto kv_param_res_names = ggml_model_decoder->get_kv_param_res_names(); diff --git a/ggml/src/ggml-openvino/openvino/utils.cpp b/ggml/src/ggml-openvino/openvino/utils.cpp index 504d74b7067..8bb7678ee38 100644 --- a/ggml/src/ggml-openvino/openvino/utils.cpp +++ b/ggml/src/ggml-openvino/openvino/utils.cpp @@ -72,6 +72,7 @@ OutputVector rename_outputs_with_suffix(const OutputVector & outputs, const std: name += "_"; name += suffix; node->set_friendly_name(name); + // Uncomment to dump every node's inferred shape (used to hunt down dynamic dims on NPU). // std::cout << name << " " << output.get_partial_shape() << std::endl; } return outputs; diff --git a/ggml/src/ggml-openvino/utils.cpp b/ggml/src/ggml-openvino/utils.cpp index 4df8381dcbd..93b1ccbe907 100644 --- a/ggml/src/ggml-openvino/utils.cpp +++ b/ggml/src/ggml-openvino/utils.cpp @@ -16,6 +16,8 @@ #include #include #include +#include +#include #include #include #include @@ -48,7 +50,7 @@ enum ggml_status ov_graph_compute(ggml_cgraph * cgraph, ggml_backend_t backend) GgmlOvDecoder::dump_cgraph(cgraph, filename); } - const auto is_static = ggml_openvino_is_npu(); + const auto is_static = ggml_openvino_is_npu() || ggml_openvino_getenv_int("GGML_OPENVINO_FORCE_STATIC"); GGML_ASSERT(ctx->runtime_context != nullptr); std::shared_ptr r_ctx = std::static_pointer_cast(ctx->runtime_context); @@ -168,13 +170,24 @@ ov::Tensor create_ov_output_tensor(std::shared_ptr ggml_decoder, auto output_type = ggml_decoder->get_ov_type(ggml_tensor); ov::Shape output_shape; + void * output_data = ggml_tensor->data; if (ggml_decoder->is_static()) { output_shape = infer_request->get_output_tensor(output_index).get_shape(); } else { - output_shape = ggml_decoder->get_shape(ggml_tensor); + // For a CPY into a padded view_src (e.g. a padded KV cache buffer), the + // OV ScatterUpdate node outputs the full view_src shape, not the CPY node's + // own (smaller) shape. Using the CPY shape here causes set_output_tensor to + // fail with a shape-incompatibility error. Use view_src's shape and data + // pointer instead so the OV tensor matches the model output exactly. + if (ggml_tensor->op == GGML_OP_CPY && ggml_tensor->view_src != nullptr && + ggml_nbytes(ggml_tensor) != ggml_nbytes(ggml_tensor->view_src)) { + output_shape = ggml_decoder->get_shape(ggml_tensor->view_src); + output_data = ggml_tensor->view_src->data; + } else { + output_shape = ggml_decoder->get_shape(ggml_tensor); + } } - - ov::Tensor output_tensor(output_type, output_shape, ggml_tensor->data); + ov::Tensor output_tensor(output_type, output_shape, output_data); return output_tensor; } @@ -583,7 +596,9 @@ enum ggml_status ov_graph_compute_static(ggml_cgraph * cgraph, std::shared_ptr(ggml_decoder_prefill); - auto input_model_decode = std::make_shared(ggml_decoder_decode); - - auto model_prefill = ov::frontend::ggml::FrontEnd::convert(input_model_prefill); - ggml_decoder_prefill->clear_model_weights(); - auto model_decode = ov::frontend::ggml::FrontEnd::convert(input_model_decode); - ggml_decoder_decode->clear_model_weights(); - conversion_end_time = ggml_time_us(); - - if (ggml_openvino_getenv_int("GGML_OPENVINO_DUMP_IR")) { - char timestamped_filename[64]; - auto timestamp = (long long) ggml_time_us(); - snprintf(timestamped_filename, sizeof(timestamped_filename), "model_prefill_%lld.xml", timestamp); - ov::serialize(model_prefill, timestamped_filename); - snprintf(timestamped_filename, sizeof(timestamped_filename), "model_decode_%lld.xml", timestamp); - ov::serialize(model_decode, timestamped_filename); - } + const bool dump_ir = ggml_openvino_getenv_int("GGML_OPENVINO_DUMP_IR"); + const auto dump_ir_timestamp = static_cast(ggml_time_us()); + + auto build_static_model = [&core, &config, dump_ir, dump_ir_timestamp]( + std::shared_ptr decoder, + const char * tag, + std::shared_ptr & model, + ov::CompiledModel & compiled_model, + std::shared_ptr & infer_request, + int64_t & local_conversion_end_time, + int64_t & local_compile_end_time) { + auto input_model = std::make_shared(decoder); + model = ov::frontend::ggml::FrontEnd::convert(input_model); + decoder->clear_model_weights(); + local_conversion_end_time = ggml_time_us(); + + if (dump_ir) { + char timestamped_filename[64]; + snprintf(timestamped_filename, sizeof(timestamped_filename), "model_%s_%lld.xml", tag, + dump_ir_timestamp); + ov::serialize(model, timestamped_filename); + } + compiled_model = core.compile_model(model, device, config); + infer_request = std::make_shared(compiled_model.create_infer_request()); + local_compile_end_time = ggml_time_us(); + }; + std::shared_ptr model_prefill; + std::shared_ptr model_decode; ov::CompiledModel compiled_model_prefill; ov::CompiledModel compiled_model_decode; - auto remote_context = ggml_openvino_get_remote_context(); - if (remote_context.has_value()) { - compiled_model_prefill = core.compile_model(model_prefill, remote_context.value(), config); - compiled_model_decode = core.compile_model(model_decode, remote_context.value(), config); - } else { - compiled_model_prefill = core.compile_model(model_prefill, device, config); - compiled_model_decode = core.compile_model(model_decode, device, config); - } - - auto infer_request_prefill = std::make_shared(compiled_model_prefill.create_infer_request()); - auto infer_request_decode = std::make_shared(compiled_model_decode.create_infer_request()); - compile_end_time = ggml_time_us(); + std::shared_ptr infer_request_prefill; + std::shared_ptr infer_request_decode; + int64_t prefill_conversion_end_time; + int64_t decode_conversion_end_time; + int64_t prefill_compile_end_time; + int64_t decode_compile_end_time; + auto prefill_future = std::async(std::launch::async, build_static_model, ggml_decoder_prefill, "prefill", + std::ref(model_prefill), std::ref(compiled_model_prefill), + std::ref(infer_request_prefill), std::ref(prefill_conversion_end_time), + std::ref(prefill_compile_end_time)); + auto decode_future = std::async(std::launch::async, build_static_model, ggml_decoder_decode, "decode", + std::ref(model_decode), std::ref(compiled_model_decode), + std::ref(infer_request_decode), std::ref(decode_conversion_end_time), + std::ref(decode_compile_end_time)); + prefill_future.get(); + decode_future.get(); + conversion_end_time = std::max(prefill_conversion_end_time, decode_conversion_end_time); + compile_end_time = std::max(prefill_compile_end_time, decode_compile_end_time); model = is_prefill ? model_prefill : model_decode; ggml_decoder = is_prefill ? ggml_decoder_prefill : ggml_decoder_decode; @@ -742,7 +774,7 @@ enum ggml_status ov_graph_compute_static(ggml_cgraph * cgraph, std::shared_ptrne[0]; + auto inp_len = get_inp_pos_n_tokens(cgraph, inp_pos); for (int chunk_index = 0; chunk_index * prefill_chunk_size < inp_len; chunk_index++) { for (size_t i = 0; i < ov_input_names_local.size(); i++) { auto param_name = ov_input_names_local[i]; @@ -762,6 +794,11 @@ enum ggml_status ov_graph_compute_static(ggml_cgraph * cgraph, std::shared_ptrsecond; + if (ggml_nbytes(ggml_tensor) == 0) { + // Zero-row in-place writeback (e.g. the empty s_copy defrag remainder). The OV + // Result is the full cache, so binding it over this 0-byte buffer overflows it. + continue; + } auto output_tensor = create_ov_output_tensor(ggml_decoder, infer_request, i, ggml_tensor); infer_request->set_output_tensor(i, output_tensor); } @@ -798,6 +835,9 @@ enum ggml_status ov_graph_compute_static(ggml_cgraph * cgraph, std::shared_ptrsecond; + if (ggml_nbytes(ggml_tensor) == 0) { + continue; + } auto output_tensor = create_ov_output_tensor(ggml_decoder, infer_request, i, ggml_tensor); infer_request->set_output_tensor(i, output_tensor); } @@ -1074,6 +1114,9 @@ ov::Tensor get_ov_input_tensor(std::shared_ptr ggml_decoder, cons ov::Tensor get_ov_input_tensor_static_decode(std::shared_ptr ggml_decoder, const std::string & param_name) { // NPU decoding stage + if (ggml_decoder->get_model_extra_inputs().count(param_name)) { + return get_ov_input_tensor(ggml_decoder, param_name); + } const auto * ggml_tensor = ggml_decoder->get_input_ggml_tensor(param_name); const auto * op = ggml_decoder->get_tensor_used_op(ggml_tensor); @@ -1123,14 +1166,30 @@ ov::Tensor get_ov_input_tensor_static_prefill(std::shared_ptr ggm const std::string & param_name, int chunk_index) { // NPU prompt processing stage - const auto * ggml_tensor = ggml_decoder->get_input_ggml_tensor(param_name); - const auto * op = ggml_decoder->get_tensor_used_op(ggml_tensor); - const size_t input_len = ggml_decoder->get_input_len(); const size_t chunk_size = ggml_decoder->m_prefill_chunk_size; const size_t chunk_valid_size = std::min(chunk_size, input_len - chunk_index * chunk_size); const size_t chunk_pad_size = chunk_size - chunk_valid_size; + if (param_name == "chunk_valid_len") { + ov::Tensor input_tensor(ov::element::i64, ov::Shape{1}); + *input_tensor.data() = (int64_t) chunk_valid_size; + return input_tensor; + } + if (chunk_index > 0 && param_name == "cache_rs_reset_len") { + // The recurrent-state clear belongs to the start of the sequence. Re-applying it on every + // chunk would wipe the state accumulated by the preceding chunks, so disable it (a zero + // length makes scale.cpp's keep-mask select every slot) after the first chunk. + ov::Tensor input_tensor(ov::element::i64, ov::Shape{1}); + *input_tensor.data() = 0; + return input_tensor; + } + if (ggml_decoder->get_model_extra_inputs().count(param_name)) { + return get_ov_input_tensor(ggml_decoder, param_name); + } + const auto * ggml_tensor = ggml_decoder->get_input_ggml_tensor(param_name); + const auto * op = ggml_decoder->get_tensor_used_op(ggml_tensor); + if (GgmlOvDecoder::is_inp_pos(ggml_tensor, op) && GgmlOvDecoder::get_inp_pos_n_planes(op) > 1) { // IMROPE: inp_pos stacks n_planes (t/h/w/e) position planes, each of length // input_len; pad every plane independently so they stay aligned to chunk_size. @@ -1306,7 +1365,7 @@ void print_input_tensor_info(const std::string & name, const ov::Tensor & tensor << std::endl; switch (tensor.get_element_type()) { case ov::element::f32: { - if (name.find("self_kq_mask") == std::string::npos) { + if (name.find("self_kq_mask") == std::string::npos && name.find("KQ_mask") == std::string::npos) { std::cout << *(tensor.data()) << std::endl; } else { size_t rows = tensor.get_shape()[2]; @@ -1414,8 +1473,24 @@ const ggml_tensor * get_inp_pos_tensor(ggml_cgraph * cgraph) { throw std::runtime_error("get_inp_pos_tensor: inp_pos not found in cgraph"); } -bool get_is_prefill(const ggml_tensor * inp_pos) { - return inp_pos->ne[0] > 1; +int64_t get_inp_pos_n_tokens(ggml_cgraph * cgraph, const ggml_tensor * inp_pos) { + // IMROPE stacks n_planes (t/h/w/e) position planes into inp_pos, so ne[0] is + // n_planes * n_tokens. Callers that need a token count must divide the planes out. + int n_planes = 1; + for (int i = 0; i < cgraph->n_nodes; ++i) { + auto * op = cgraph->nodes[i]; + for (int j = 0; j < GGML_MAX_SRC; ++j) { + if (op->src[j] == inp_pos) { + n_planes = GgmlOvDecoder::get_inp_pos_n_planes(op); + break; + } + } + } + return inp_pos->ne[0] / n_planes; +} + +bool get_is_prefill(ggml_cgraph * cgraph, const ggml_tensor * inp_pos) { + return get_inp_pos_n_tokens(cgraph, inp_pos) > 1; } #pragma GCC diagnostic pop diff --git a/ggml/src/ggml-openvino/utils.h b/ggml/src/ggml-openvino/utils.h index 513fa83c9d6..5aa74da38d3 100644 --- a/ggml/src/ggml-openvino/utils.h +++ b/ggml/src/ggml-openvino/utils.h @@ -164,7 +164,9 @@ std::vector pad_input(const ggml_tensor * tensor, size_t padded_rows, size_t const ggml_tensor * get_inp_pos_tensor(struct ggml_cgraph * cgraph); -bool get_is_prefill(const ggml_tensor * inp_pos); +int64_t get_inp_pos_n_tokens(struct ggml_cgraph * cgraph, const ggml_tensor * inp_pos); + +bool get_is_prefill(struct ggml_cgraph * cgraph, const ggml_tensor * inp_pos); ov::Tensor get_ov_input_tensor(std::shared_ptr ggml_decoder, const std::string & param_name); ov::Tensor get_ov_input_tensor_static_decode(std::shared_ptr ggml_decoder, From 0a150873efb17f3b21e92c83d705bbfdce200bd0 Mon Sep 17 00:00:00 2001 From: Strongtut Date: Fri, 28 Aug 2026 05:37:37 -0700 Subject: [PATCH 021/104] metal : add fa-vec tunings for M4 (llama/27875) This adds fa_vec_tuned_table records for Apple M4 to ggml-metal-tuning.cpp. Includes F16, Q4_0, Q4_1, Q5_0, Q5_1, and Q8_0. (M4, 10 GPU Cores) Co-authored-by: Strongtut <8432058+Strongtut@users.noreply.github.com> --- ggml/src/ggml-metal/ggml-metal-tuning.cpp | 226 ++++++++++++++++++++++ 1 file changed, 226 insertions(+) diff --git a/ggml/src/ggml-metal/ggml-metal-tuning.cpp b/ggml/src/ggml-metal/ggml-metal-tuning.cpp index c2139fe20b0..285590d1601 100644 --- a/ggml/src/ggml-metal/ggml-metal-tuning.cpp +++ b/ggml/src/ggml-metal/ggml-metal-tuning.cpp @@ -542,6 +542,232 @@ constexpr fa_vec_entry_t fa_vec_tuned_table[] = { { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q8_0, 576, 512, 1, 1 }, { 1, 2 } }, { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q8_0, 576, 512, 1, 4 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_F16, 32, 32, 1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_F16, 32, 32, 1, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_F16, 32, 32, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_F16, 32, 32, 2, 3 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_F16, 32, 32, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_F16, 32, 32, 3, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_F16, 64, 64, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_F16, 64, 64, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_F16, 64, 64, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_F16, 64, 64, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_F16, 96, 96, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_F16, 96, 96, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_F16, 96, 96, 1, 4 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_F16, 96, 96, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_F16, 96, 96, 2, 4 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_F16, 128, 128, -1, 1 }, { 2, 2 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_F16, 128, 128, 1, 2 }, { 1, 1 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_F16, 128, 128, 1, 4 }, { 1, 1 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_F16, 128, 128, 2, 2 }, { 1, 1 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_F16, 128, 128, 2, 4 }, { 1, 1 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_F16, 192, 192, 3, 0 }, { 2, 2 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_F16, 192, 192, -1, 1 }, { 2, 2 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_F16, 192, 192, 1, 2 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_F16, 192, 128, -1, 1 }, { 2, 2 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_F16, 192, 128, 1, 2 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_F16, 192, 128, 1, 4 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_F16, 256, 256, -1, 0 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_F16, 256, 256, 3, 0 }, { 2, 2 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_F16, 256, 256, -1, 1 }, { 2, 2 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_F16, 320, 256, 3, 0 }, { 2, 2 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_F16, 320, 256, -1, 1 }, { 2, 2 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_F16, 512, 512, 3, 0 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_F16, 512, 512, 3, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_F16, 512, 512, 3, 2 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_F16, 512, 512, 3, 3 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q4_0, 32, 32, -1, 1 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q4_0, 32, 32, 1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q4_0, 32, 32, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q4_0, 32, 32, 2, 4 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q4_0, 32, 32, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q4_0, 64, 64, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q4_0, 64, 64, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q4_0, 64, 64, 1, 2 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q4_0, 64, 64, 2, 2 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q4_0, 64, 64, 2, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q4_0, 64, 64, 3, 2 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q4_0, 64, 64, 3, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q4_0, 96, 96, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q4_0, 96, 96, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q4_0, 128, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q4_0, 128, 128, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q4_0, 128, 128, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q4_0, 128, 128, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q4_0, 128, 128, 3, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q4_0, 192, 192, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q4_0, 192, 192, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q4_0, 192, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q4_0, 192, 128, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q4_0, 192, 128, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q4_0, 192, 128, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q4_0, 192, 128, 3, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q4_0, 256, 256, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q4_0, 256, 256, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q4_0, 320, 256, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q4_0, 320, 256, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q4_0, 512, 512, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q4_0, 512, 512, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q4_0, 576, 512, 2, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q4_0, 576, 512, 3, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q4_0, 576, 512, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q4_0, 576, 512, 1, 1 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q4_0, 576, 512, 1, 2 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q4_0, 576, 512, 1, 3 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q4_0, 576, 512, 1, 4 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q4_1, 32, 32, -1, 1 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q4_1, 32, 32, 1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q4_1, 32, 32, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q4_1, 32, 32, 2, 4 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q4_1, 32, 32, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q4_1, 64, 64, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q4_1, 64, 64, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q4_1, 64, 64, 1, 2 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q4_1, 64, 64, 2, 2 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q4_1, 64, 64, 2, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q4_1, 64, 64, 3, 2 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q4_1, 64, 64, 3, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q4_1, 96, 96, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q4_1, 96, 96, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q4_1, 96, 96, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q4_1, 128, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q4_1, 128, 128, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q4_1, 128, 128, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q4_1, 128, 128, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q4_1, 128, 128, 3, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q4_1, 192, 192, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q4_1, 192, 192, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q4_1, 192, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q4_1, 192, 128, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q4_1, 192, 128, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q4_1, 192, 128, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q4_1, 192, 128, 3, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q4_1, 256, 256, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q4_1, 256, 256, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q4_1, 320, 256, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q4_1, 320, 256, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q4_1, 576, 512, 2, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q4_1, 576, 512, 3, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q4_1, 576, 512, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q4_1, 576, 512, 1, 1 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q4_1, 576, 512, 1, 2 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q4_1, 576, 512, 1, 3 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q4_1, 576, 512, 1, 4 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q5_0, 32, 32, -1, 1 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q5_0, 32, 32, 1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q5_0, 32, 32, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q5_0, 32, 32, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q5_0, 64, 64, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q5_0, 64, 64, -1, 1 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q5_0, 64, 64, 1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q5_0, 64, 64, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q5_0, 64, 64, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q5_0, 96, 96, -1, 1 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q5_0, 96, 96, 1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q5_0, 96, 96, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q5_0, 96, 96, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q5_0, 128, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q5_0, 128, 128, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q5_0, 128, 128, 1, 2 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q5_0, 128, 128, 1, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q5_0, 128, 128, 2, 2 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q5_0, 128, 128, 2, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q5_0, 128, 128, 3, 2 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q5_0, 128, 128, 3, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q5_0, 192, 192, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q5_0, 192, 192, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q5_0, 192, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q5_0, 192, 128, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q5_0, 256, 256, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q5_0, 256, 256, -1, 1 }, { 2, 2 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q5_0, 256, 256, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q5_0, 256, 256, 3, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q5_0, 320, 256, -1, 1 }, { 2, 2 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q5_0, 320, 256, 1, 2 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q5_0, 320, 256, 2, 2 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q5_0, 320, 256, 3, 2 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q5_0, 512, 512, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q5_0, 512, 512, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q5_0, 576, 512, 2, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q5_0, 576, 512, 2, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q5_0, 576, 512, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q5_0, 576, 512, 2, 3 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q5_0, 576, 512, 2, 4 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q5_1, 32, 32, -1, 1 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q5_1, 32, 32, 1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q5_1, 32, 32, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q5_1, 32, 32, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q5_1, 64, 64, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q5_1, 64, 64, -1, 1 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q5_1, 64, 64, 1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q5_1, 64, 64, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q5_1, 64, 64, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q5_1, 96, 96, -1, 1 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q5_1, 96, 96, 1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q5_1, 96, 96, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q5_1, 96, 96, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q5_1, 128, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q5_1, 128, 128, -1, 1 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q5_1, 128, 128, 1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q5_1, 128, 128, 1, 4 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q5_1, 128, 128, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q5_1, 128, 128, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q5_1, 128, 128, 3, 4 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q5_1, 192, 192, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q5_1, 192, 192, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q5_1, 192, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q5_1, 192, 128, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q5_1, 256, 256, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q5_1, 256, 256, -1, 1 }, { 2, 2 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q5_1, 256, 256, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q5_1, 256, 256, 3, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q5_1, 320, 256, -1, 1 }, { 2, 2 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q5_1, 320, 256, 1, 2 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q5_1, 320, 256, 1, 4 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q5_1, 320, 256, 2, 2 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q5_1, 320, 256, 3, 2 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q5_1, 512, 512, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q5_1, 512, 512, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q5_1, 576, 512, 2, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q5_1, 576, 512, 2, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q5_1, 576, 512, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q5_1, 576, 512, 2, 3 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q5_1, 576, 512, 2, 4 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q8_0, 32, 32, -1, 1 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q8_0, 32, 32, 1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q8_0, 32, 32, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q8_0, 32, 32, 2, 4 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q8_0, 32, 32, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q8_0, 64, 64, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q8_0, 64, 64, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q8_0, 64, 64, 1, 2 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q8_0, 64, 64, 2, 2 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q8_0, 64, 64, 3, 2 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q8_0, 64, 64, 3, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q8_0, 96, 96, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q8_0, 128, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q8_0, 128, 128, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q8_0, 128, 128, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q8_0, 128, 128, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q8_0, 192, 192, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q8_0, 192, 192, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q8_0, 192, 192, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q8_0, 192, 192, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q8_0, 192, 192, 2, 4 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q8_0, 192, 192, 3, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q8_0, 192, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q8_0, 192, 128, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q8_0, 256, 256, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q8_0, 256, 256, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q8_0, 320, 256, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q8_0, 320, 256, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q8_0, 512, 512, -1, 0 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q8_0, 512, 512, -1, 1 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q8_0, 576, 512, 2, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q8_0, 576, 512, 3, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_Q8_0, 576, 512, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_F16, 32, 32, 1, 3 }, { 2, 4 } }, { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_F16, 32, 32, 2, 1 }, { 2, 4 } }, { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_F16, 32, 32, 2, 3 }, { 2, 4 } }, From ba99c098669484ad18978f2b18557726709fba12 Mon Sep 17 00:00:00 2001 From: ravel7524 <58877666+ravel7524@users.noreply.github.com> Date: Fri, 28 Aug 2026 10:52:49 -0400 Subject: [PATCH 022/104] Vulkan: add hoisting support for row IDs and expert count in shaders (llama/26686) * vulkan: add hoisting support for row IDs and expert count in shaders * use hoisted row ids in coopmat2 * vulkan: address review feedback on count_experts - use vk_op_count_experts_push_constants instead of a raw uint vector - apply the fastdiv trick to the ne00 div/mod in count_experts - compute the per-expert offsets with subgroupExclusiveAdd when the device supports it, keeping the serial path as fallback - document the data_d layout and the hoisted_row_id_words bound - drop a leftover debug print in ggml_vk_matmul_id * vulkan: use init_pushconst_fastdiv for count_experts push constants * vulkan: refine comments for row ID hoisting and data layout in count_experts shader * Whitespace --------- Co-authored-by: Jeff Bolz --- ggml/src/ggml-vulkan/ggml-vulkan.cpp | 48 +++++++--- .../vulkan-shaders/count_experts.comp | 95 ++++++++++++++++++- .../ggml-vulkan/vulkan-shaders/mul_mm.comp | 35 ++++--- .../vulkan-shaders/mul_mm_cm2.comp | 25 ++++- .../vulkan-shaders/mul_mm_id_funcs.glsl | 15 +++ .../ggml-vulkan/vulkan-shaders/mul_mmq.comp | 35 ++++--- .../vulkan-shaders/vulkan-shaders-gen.cpp | 1 + 7 files changed, 212 insertions(+), 42 deletions(-) diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp index 72e844aebfd..28b2d875d20 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp +++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp @@ -1348,6 +1348,8 @@ struct vk_mat_mat_id_push_constants { uint32_t batch_stride_a; uint32_t batch_stride_b; uint32_t batch_stride_d; uint32_t nei0; uint32_t nei1; uint32_t nbi1; uint32_t ne11; uint32_t padded_N; + uint32_t n_experts; + uint32_t hoist_row_ids; }; struct vk_mat_vec_id_push_constants { uint32_t ncols; @@ -1428,6 +1430,10 @@ struct vk_op_count_experts_push_constants { uint32_t nb00; uint32_t nb01; uint32_t a_offset; + uint32_t n_experts; + uint32_t hoist_row_ids; + uint32_t ne00mp; + uint32_t ne00L; }; struct vk_op_glu_push_constants { @@ -1606,6 +1612,10 @@ template <> void init_pushconst_fastdiv(vk_op_glu_push_constants &p) { init_fastdiv_values(p.ne20, p.ne2_0mp, p.ne2_0L); } +template <> void init_pushconst_fastdiv(vk_op_count_experts_push_constants &p) { + init_fastdiv_values(p.ne00, p.ne00mp, p.ne00L); +} + struct vk_op_binary_push_constants { uint32_t ne; uint32_t ne00; uint32_t ne01; uint32_t ne02; uint32_t ne03; uint32_t nb00; uint32_t nb01; uint32_t nb02; uint32_t nb03; @@ -5839,7 +5849,11 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { ggml_vk_create_pipeline(device, device->pipeline_count_equal_i32, "count_equal_i32", count_equal_i32_len, count_equal_i32_data, "main", 3, sizeof(vk_op_push_constants), {512, 1, 1}, { device->subgroup_size }, 1); - ggml_vk_create_pipeline(device, device->pipeline_count_experts, "count_experts", count_experts_len, count_experts_data, "main", 2, sizeof(vk_op_count_experts_push_constants), {1, 1, 1}, {}, 1, true); + if (device->subgroup_arithmetic && device->subgroup_require_full_support) { + ggml_vk_create_pipeline(device, device->pipeline_count_experts, "count_experts", count_experts_subgroup_len, count_experts_subgroup_data, "main", 2, sizeof(vk_op_count_experts_push_constants), {1, 1, 1}, {}, 1, true, true); + } else { + ggml_vk_create_pipeline(device, device->pipeline_count_experts, "count_experts", count_experts_len, count_experts_data, "main", 2, sizeof(vk_op_count_experts_push_constants), {1, 1, 1}, {}, 1, true); + } for (auto &s : device->pipeline_solve_tri_f32) { const vk_solve_tri_pipeline_state &state = s.first; @@ -8970,13 +8984,13 @@ static void ggml_vk_matmul_id( uint32_t m, uint32_t n, uint32_t k, uint32_t stride_a, uint32_t stride_b, uint32_t stride_d, uint32_t batch_stride_a, uint32_t batch_stride_b, uint32_t batch_stride_d, uint32_t n_as, uint32_t nei0, uint32_t nei1, uint32_t nbi1, uint32_t ne11, - uint32_t padded_n) { + uint32_t padded_n, bool hoist_row_ids) { VK_LOG_DEBUG("ggml_vk_matmul_id(a: (" << a.buffer->buffer << ", " << a.offset << ", " << a.size << "), b: (" << b.buffer->buffer << ", " << b.offset << ", " << b.size << "), d: (" << d.buffer->buffer << ", " << d.offset << ", " << d.size << "), ids: (" << ids.buffer->buffer << ", " << ids.offset << ", " << ids.size << "), expert_count: (" << expert_count_buf.buffer->buffer << ", " << expert_count_buf.offset << ", " << expert_count_buf.size << "), " << "m: " << m << ", n: " << n << ", k: " << k << ", stride_a: " << stride_a << ", stride_b: " << stride_b << ", stride_d: " << stride_d << ", " << "batch_stride_a: " << batch_stride_a << ", batch_stride_b: " << batch_stride_b << ", batch_stride_d: " << batch_stride_d << ", " << "n_as: " << n_as << ", nei0: " << nei0 << ", nei1: " << nei1 << ", nbi1: " << nbi1 << ", ne11: " << ne11 << ")"); const vk_mat_mat_id_push_constants pc = { m, n, k, stride_a, stride_b, stride_d, batch_stride_a, batch_stride_b, batch_stride_d, - nei0, nei1, nbi1, ne11, padded_n }; + nei0, nei1, nbi1, ne11, padded_n, n_as, uint32_t(hoist_row_ids) }; ggml_vk_dispatch_pipeline(ctx, subctx, pipeline, { a, b, d, ids, expert_count_buf }, pc, { m, nei1, n_as }); } @@ -10162,6 +10176,12 @@ static void ggml_vk_mul_mat_id_q_f16(ggml_backend_vk_context * ctx, vk_context& // const uint64_t ne23 = dst->ne[3]; const uint64_t n_as = ne02; + // n_as counts, n_as offsets, one total, then one packed row id per (expert, token). + // Hoisting requires 16-bit indices for the packing and a table that fits one binding. + const uint64_t hoisted_row_id_words = 2 * n_as + 1 + nei0 * nei1; + const bool hoist_row_ids = n_as <= 256 && nei0 <= 0xffff && nei1 <= 0xffff && + hoisted_row_id_words * sizeof(uint32_t) <= + ctx->device->properties.limits.maxStorageBufferRange; ggml_backend_vk_buffer_context * dst_buf_ctx = (ggml_backend_vk_buffer_context *)dst->buffer->context; ggml_backend_vk_buffer_context * src0_buf_ctx = (ggml_backend_vk_buffer_context *)src0->buffer->context; @@ -10302,7 +10322,8 @@ static void ggml_vk_mul_mat_id_q_f16(ggml_backend_vk_context * ctx, vk_context& } vk_pipeline count_experts = ctx->device->pipeline_count_experts; - uint32_t expert_count_size = sizeof(uint32_t) * n_as; + const size_t expert_data_size = sizeof(uint32_t) * + (hoist_row_ids ? hoisted_row_id_words : n_as); { if ( @@ -10318,8 +10339,8 @@ static void ggml_vk_mul_mat_id_q_f16(ggml_backend_vk_context * ctx, vk_context& ctx->prealloc_size_y = y_sz; ggml_vk_preallocate_buffers(ctx, subctx); } - if (ctx->prealloc_size_split_k < expert_count_size) { - ctx->prealloc_size_split_k = expert_count_size; + if (ctx->prealloc_size_split_k < expert_data_size) { + ctx->prealloc_size_split_k = expert_data_size; ggml_vk_preallocate_buffers(ctx, subctx); } @@ -10385,18 +10406,23 @@ static void ggml_vk_mul_mat_id_q_f16(ggml_backend_vk_context * ctx, vk_context& } } // Count how many times each expert is used - vk_subbuffer expert_count_buf = ggml_vk_subbuffer(ctx, ctx->prealloc_split_k, 0); + vk_subbuffer expert_count_buf = { ctx->prealloc_split_k, 0, expert_data_size }; if (ctx->prealloc_split_k_need_sync) { ggml_vk_sync_buffers(ctx, subctx); } { - const std::vector pc = { (uint32_t)nei0, + vk_op_count_experts_push_constants pc = { (uint32_t)nei0, (uint32_t)nei1, (uint32_t)(nbi0 / ggml_type_size(ids->type)), (uint32_t)(nbi1 / ggml_type_size(ids->type)), - (uint32_t)(get_misalign_bytes(ctx, ids) / ggml_type_size(ids->type)) }; + (uint32_t)(get_misalign_bytes(ctx, ids) / ggml_type_size(ids->type)), + (uint32_t)n_as, + uint32_t(hoist_row_ids), + 0, 0 }; + init_pushconst_fastdiv(pc); ggml_vk_dispatch_pipeline(ctx, subctx, count_experts, - { vk_subbuffer{ d_ids, ids_buf_offset, ids_sz }, expert_count_buf }, pc, { (uint32_t)n_as, 1, 1}); + { vk_subbuffer{ d_ids, ids_buf_offset, ids_sz }, expert_count_buf }, pc, + { hoist_row_ids ? 1u : (uint32_t)n_as, 1, 1}); } if (x_non_contig) { @@ -10465,7 +10491,7 @@ static void ggml_vk_mul_mat_id_q_f16(ggml_backend_vk_context * ctx, vk_context& { d_D, d_buf_offset, d_sz }, { d_ids, ids_buf_offset, ids_sz }, expert_count_buf, ne01, ne21, ne10, ne10, stride_b_y, ne01, stride_batch_x, stride_batch_y, ne20*ne21, - n_as, nei0, nei1, nbi1 / ggml_type_size(ids->type), ne11, padded_n + n_as, nei0, nei1, nbi1 / ggml_type_size(ids->type), ne11, padded_n, hoist_row_ids ); // NOLINT if (x_non_contig || qx_needs_dequant) { diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/count_experts.comp b/ggml/src/ggml-vulkan/vulkan-shaders/count_experts.comp index ffc8608691f..83c56fce520 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/count_experts.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/count_experts.comp @@ -2,6 +2,11 @@ #extension GL_EXT_control_flow_attributes : enable +#ifdef USE_SUBGROUPS +#extension GL_KHR_shader_subgroup_basic : enable +#extension GL_KHR_shader_subgroup_arithmetic : enable +#endif + #include "types.glsl" layout (push_constant) uniform parameter @@ -11,6 +16,10 @@ layout (push_constant) uniform parameter uint32_t nb00; uint32_t nb01; uint32_t a_offset; + uint32_t n_experts; + uint32_t hoist_row_ids; + uint32_t ne00mp; + uint32_t ne00L; } p; #define BLOCK_SIZE 256 @@ -21,16 +30,98 @@ layout (binding = 0) readonly buffer A {uint data_a[];}; layout (binding = 1) writeonly buffer D {uint data_d[];}; shared uint vals[BLOCK_SIZE]; +shared uint offsets[BLOCK_SIZE]; +shared uint cursors[BLOCK_SIZE]; + +// see init_fastdiv_values in ggml-vulkan.cpp +uint fastdiv(uint n, uint mp, uint L) { + uint msbs, lsbs; + // msbs = mulhi(n, mp) + umulExtended(n, mp, msbs, lsbs); + return (msbs + n) >> L; +} +// data_d layout when p.hoist_row_ids is set: +// [0, n_experts) per-expert row count +// [n_experts, 2*n_experts) per-expert start offset into the row id region +// [2*n_experts] total row count +// [2*n_experts + 1, ) row ids grouped by expert, packed as (i01 << 16) | (i00 & 0xffff) +// Otherwise only data_d[expert_id] is written, holding that expert's row count. void main() { const uint expert_id = gl_WorkGroupID.x; const uint num_elements = p.ne00 * p.ne01; const uint tid = gl_LocalInvocationID.x; + if (p.hoist_row_ids != 0) { + if (tid < p.n_experts) { + vals[tid] = 0; + } + barrier(); + + for (uint idx = tid; idx < num_elements; idx += BLOCK_SIZE) { + const uint i01 = fastdiv(idx, p.ne00mp, p.ne00L); + const uint i00 = idx - i01 * p.ne00; + const uint expert = data_a[p.a_offset + i01 * p.nb01 + i00 * p.nb00]; + if (expert < p.n_experts) { + atomicAdd(vals[expert], 1); + } + } + barrier(); + +#ifdef USE_SUBGROUPS + if (gl_SubgroupID == 0) { + // pad the trip count so the subgroup ops stay in uniform control flow + const uint n_experts_padded = (p.n_experts + gl_SubgroupSize - 1) & ~(gl_SubgroupSize - 1); + uint base = 0; + for (uint expert = gl_SubgroupInvocationID; expert < n_experts_padded; expert += gl_SubgroupSize) { + const bool in_range = expert < p.n_experts; + const uint count = in_range ? vals[expert] : 0; + const uint offset = base + subgroupExclusiveAdd(count); + if (in_range) { + data_d[expert] = count; + data_d[p.n_experts + expert] = offset; + offsets[expert] = offset; + cursors[expert] = 0; + } + base += subgroupAdd(count); + } + if (subgroupElect()) { + data_d[2 * p.n_experts] = base; + } + } +#else + if (tid == 0) { + uint offset = 0; + for (uint expert = 0; expert < p.n_experts; ++expert) { + const uint count = vals[expert]; + data_d[expert] = count; + data_d[p.n_experts + expert] = offset; + offsets[expert] = offset; + cursors[expert] = 0; + offset += count; + } + data_d[2 * p.n_experts] = offset; + } +#endif + barrier(); + + for (uint idx = tid; idx < num_elements; idx += BLOCK_SIZE) { + const uint i01 = fastdiv(idx, p.ne00mp, p.ne00L); + const uint i00 = idx - i01 * p.ne00; + const uint expert = data_a[p.a_offset + i01 * p.nb01 + i00 * p.nb00]; + if (expert < p.n_experts) { + const uint row = atomicAdd(cursors[expert], 1); + const uint packed_row_id = (i01 << 16) | (i00 & 0xffffu); + data_d[2 * p.n_experts + 1 + offsets[expert] + row] = packed_row_id; + } + } + return; + } + uint count = 0; for (uint idx = tid; idx < num_elements; idx += BLOCK_SIZE) { - const uint i01 = idx / p.ne00; - const uint i00 = idx % p.ne00; + const uint i01 = fastdiv(idx, p.ne00mp, p.ne00L); + const uint i00 = idx - i01 * p.ne00; const uint a = data_a[p.a_offset + i01 * p.nb01 + i00 * p.nb00]; count += uint(a == expert_id); diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm.comp b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm.comp index 3df88044a5e..c1ccac7aa20 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm.comp @@ -88,6 +88,9 @@ layout (push_constant) uniform parameter uint nei1; uint nbi1; uint ne11; + uint padded_N; + uint n_experts; + uint hoist_row_ids; #else uint base_work_group_z; uint num_batches; @@ -214,27 +217,31 @@ void main() { const uint loadstride_b = gl_WorkGroupSize.x * LOAD_VEC_B_EFF * LOAD_VEC_BATCH_B / BK; #ifdef MUL_MAT_ID -#ifdef MUL_MAT_ID_USE_SUBGROUPS - if (bitCount(p.nei0) == 1) { - load_row_ids(expert_idx, true, ic); + if (p.hoist_row_ids != 0) { + load_row_ids_hoisted(expert_idx, ic); } else { - load_row_ids(expert_idx, false, ic); - } +#ifdef MUL_MAT_ID_USE_SUBGROUPS + if (bitCount(p.nei0) == 1) { + load_row_ids(expert_idx, true, ic); + } else { + load_row_ids(expert_idx, false, ic); + } #else - _ne1 = 0; - for (uint ii1 = 0; ii1 < p.nei1 && _ne1 < (ic + 1) * BN; ii1++) { - for (uint ii0 = 0; ii0 < p.nei0 && _ne1 < (ic + 1) * BN; ii0++) { - if (data_ids[ii1*p.nbi1 + ii0] == expert_idx) { - if (_ne1 >= ic * BN) { - row_ids[_ne1 - ic * BN] = u16vec2(ii0, ii1); + _ne1 = 0; + for (uint ii1 = 0; ii1 < p.nei1 && _ne1 < (ic + 1) * BN; ii1++) { + for (uint ii0 = 0; ii0 < p.nei0 && _ne1 < (ic + 1) * BN; ii0++) { + if (data_ids[ii1*p.nbi1 + ii0] == expert_idx) { + if (_ne1 >= ic * BN) { + row_ids[_ne1 - ic * BN] = u16vec2(ii0, ii1); + } + _ne1++; } - _ne1++; } } - } - barrier(); + barrier(); #endif + } // Workgroup has no work if (ic * BN >= _ne1) return; diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm_cm2.comp b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm_cm2.comp index a2e15f6f5ce..cf78474a925 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm_cm2.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm_cm2.comp @@ -67,6 +67,10 @@ layout (push_constant) uniform parameter #endif // N dimension for the B matrix can be >= p.N uint padded_N; +#ifdef MUL_MAT_ID + uint n_experts; + uint hoist_row_ids; +#endif } p; @@ -225,6 +229,23 @@ void load_row_ids(uint expert_idx, bool nei0_is_pow2, uint ic) { } barrier(); } + +void load_row_ids_hoisted(uint expert_idx, uint ic) { + _ne1 = uint(data_expert_count[expert_idx]); + + const uint tile_begin = ic * BN; + const uint tile_count = tile_begin < _ne1 ? min(BN, _ne1 - tile_begin) : 0; + const uint expert_offset = uint(data_expert_count[p.n_experts + expert_idx]); + const uint row_ids_offset = 2 * p.n_experts + 1 + expert_offset + tile_begin; + + for (uint i = gl_LocalInvocationIndex; i < tile_count; i += BLOCK_SIZE) { + const uint packed_row_id = uint(data_expert_count[row_ids_offset + i]); + const uint ii0 = packed_row_id & 0xffffu; + const uint ii1 = packed_row_id >> 16; + row_ids[i] = u16vec4(fastmod(ii0, p.ne11), ii1, ii0, 0); + } + barrier(); +} #endif void main() { @@ -266,7 +287,9 @@ void main() { const uint ik = gl_WorkGroupID.x / blocks_m; #ifdef MUL_MAT_ID - if (bitCount(p.nei0) == 1) { + if (p.hoist_row_ids != 0) { + load_row_ids_hoisted(expert_idx, ic); + } else if (bitCount(p.nei0) == 1) { load_row_ids(expert_idx, true, ic); } else { load_row_ids(expert_idx, false, ic); diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm_id_funcs.glsl b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm_id_funcs.glsl index 26c5c12a49a..54ad60b2efb 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm_id_funcs.glsl +++ b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm_id_funcs.glsl @@ -71,4 +71,19 @@ void load_row_ids(uint expert_idx, bool nei0_is_pow2, uint ic) { barrier(); } #endif // MUL_MAT_ID_USE_SUBGROUPS + +void load_row_ids_hoisted(uint expert_idx, uint ic) { + _ne1 = uint(data_expert_count[expert_idx]); + + const uint tile_begin = ic * BN; + const uint tile_count = tile_begin < _ne1 ? min(BN, _ne1 - tile_begin) : 0; + const uint expert_offset = uint(data_expert_count[p.n_experts + expert_idx]); + const uint row_ids_offset = 2 * p.n_experts + 1 + expert_offset + tile_begin; + + for (uint i = gl_LocalInvocationIndex; i < tile_count; i += BLOCK_SIZE) { + const uint packed_row_id = uint(data_expert_count[row_ids_offset + i]); + row_ids[i] = u16vec2(packed_row_id & 0xffffu, packed_row_id >> 16); + } + barrier(); +} #endif // MUL_MAT_ID diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq.comp b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq.comp index aae1c2e8ae9..c2d84c05b40 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq.comp @@ -56,6 +56,9 @@ layout (push_constant) uniform parameter uint nei1; uint nbi1; uint ne11; + uint padded_N; + uint n_experts; + uint hoist_row_ids; #else uint base_work_group_z; uint num_batches; @@ -157,27 +160,31 @@ void main() { const uint loadstride_b = BLOCK_SIZE * LOAD_VEC_B / BK; #ifdef MUL_MAT_ID -#ifdef MUL_MAT_ID_USE_SUBGROUPS - if (bitCount(p.nei0) == 1) { - load_row_ids(expert_idx, true, ic); + if (p.hoist_row_ids != 0) { + load_row_ids_hoisted(expert_idx, ic); } else { - load_row_ids(expert_idx, false, ic); - } +#ifdef MUL_MAT_ID_USE_SUBGROUPS + if (bitCount(p.nei0) == 1) { + load_row_ids(expert_idx, true, ic); + } else { + load_row_ids(expert_idx, false, ic); + } #else - _ne1 = 0; - for (uint ii1 = 0; ii1 < p.nei1 && _ne1 < (ic + 1) * BN; ii1++) { - for (uint ii0 = 0; ii0 < p.nei0 && _ne1 < (ic + 1) * BN; ii0++) { - if (data_ids[ii1*p.nbi1 + ii0] == expert_idx) { - if (_ne1 >= ic * BN) { - row_ids[_ne1 - ic * BN] = u16vec2(ii0, ii1); + _ne1 = 0; + for (uint ii1 = 0; ii1 < p.nei1 && _ne1 < (ic + 1) * BN; ii1++) { + for (uint ii0 = 0; ii0 < p.nei0 && _ne1 < (ic + 1) * BN; ii0++) { + if (data_ids[ii1*p.nbi1 + ii0] == expert_idx) { + if (_ne1 >= ic * BN) { + row_ids[_ne1 - ic * BN] = u16vec2(ii0, ii1); + } + _ne1++; } - _ne1++; } } - } - barrier(); + barrier(); #endif + } // Workgroup has no work if (ic * BN >= _ne1) return; diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp b/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp index 0da943da956..d375c2d1277 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp @@ -1039,6 +1039,7 @@ void process_shaders() { string_to_spv("cumsum_multipass2_f32", "cumsum_multipass2.comp", merge_maps(base_dict, {{"A_TYPE", "float"}, {"D_TYPE", "float"}})); string_to_spv("count_experts", "count_experts.comp", merge_maps(base_dict, {{"A_TYPE", "uint"}, {"D_TYPE", "uint"}})); + string_to_spv("count_experts_subgroup", "count_experts.comp", merge_maps(base_dict, {{"A_TYPE", "uint"}, {"D_TYPE", "uint"}, {"USE_SUBGROUPS", "1"}})); for (std::string dim_str : {"", "_3d"}) { for (bool bda : {false, true}) { From caea96f6ceef38c2fe6d6e0332c83ae42436b264 Mon Sep 17 00:00:00 2001 From: Tekin Ertekin Date: Fri, 28 Aug 2026 20:09:08 +0300 Subject: [PATCH 023/104] ggml : fix conv_transpose_2d for multiple batches (llama/26132) * ggml : fix conv_transpose_2d for multiple batches ggml_compute_forward_conv_transpose_2d_impl only computed the first batch (ne[3] of the destination); every batch after the first was left as zero. Both the src1 permutation and the main compute loop now iterate over the batch dimension, and the work buffer size in ggml_graph_plan is scaled by the src1 batch count so the extra permuted batches fit. A multi-batch test case is added to test-backend-ops. Fixes ggml-org/ggml#1448 * metal : fix conv_transpose_2d for multiple batches The kernel only computed batch 0 of the input (src1->ne[3]); every output batch after the first was left as zero, so multi-batch conv_transpose_2d results diverged from the CPU reference. The grid now covers all batches (OW x OH x OC x N), the kernel decodes the batch from the grid z coordinate and offsets both the input and destination indices accordingly. nb3 is passed in the kernel args. Assisted-by: pi:llama.cpp/Qwen3.8-27B --------- Co-authored-by: Georgi Gerganov --- ggml/src/ggml-cpu/ggml-cpu.c | 3 +- ggml/src/ggml-cpu/ops.cpp | 58 ++++++++++++++------------ ggml/src/ggml-metal/ggml-metal-impl.h | 1 + ggml/src/ggml-metal/ggml-metal-ops.cpp | 4 +- ggml/src/ggml-metal/kernels/conv.metal | 7 ++-- 5 files changed, 42 insertions(+), 31 deletions(-) diff --git a/ggml/src/ggml-cpu/ggml-cpu.c b/ggml/src/ggml-cpu/ggml-cpu.c index 87ac0a702ef..b9c0fa3ddc0 100644 --- a/ggml/src/ggml-cpu/ggml-cpu.c +++ b/ggml/src/ggml-cpu/ggml-cpu.c @@ -2936,12 +2936,13 @@ struct ggml_cplan ggml_graph_plan( const int64_t ne10 = node->src[1]->ne[0]; // W const int64_t ne11 = node->src[1]->ne[1]; // H const int64_t ne12 = node->src[1]->ne[2]; // Channels In + const int64_t ne13 = node->src[1]->ne[3]; // Batch GGML_ASSERT(node->src[0]->type == GGML_TYPE_F16 || node->src[0]->type == GGML_TYPE_F32); GGML_ASSERT(node->src[1]->type == GGML_TYPE_F32); cur += ggml_type_size(node->src[0]->type) * ne00 * ne01 * ne02 * ne03; - cur += ggml_type_size(node->src[0]->type) * ne10 * ne11 * ne12; + cur += ggml_type_size(node->src[0]->type) * ne10 * ne11 * ne12 * ne13; } break; case GGML_OP_TOP_K: diff --git a/ggml/src/ggml-cpu/ops.cpp b/ggml/src/ggml-cpu/ops.cpp index b869f4bddde..b47ce5463c6 100644 --- a/ggml/src/ggml-cpu/ops.cpp +++ b/ggml/src/ggml-cpu/ops.cpp @@ -7267,18 +7267,21 @@ static void ggml_compute_forward_conv_transpose_2d_impl( } } - // permute source data (src1) from (Sw x Sh x Cin) to (Cin x Sw x Sh) + // permute source data (src1) from (Sw x Sh x Cin) to (Cin x Sw x Sh), for all batches { kernel_t * const wdata = (kernel_t *) params->wdata + nk; - for (int i12 = 0; i12 < ne12; i12++) { - for (int i11 = 0; i11 < ne11; i11++) { - const float * const src = (float *)((char *) src1->data + i12*nb12 + i11*nb11); - kernel_t * dst_data = wdata + i11*ne10*ne12; - for (int i10 = 0; i10 < ne10; i10++) { - if constexpr (std::is_same_v) { - dst_data[i10*ne12 + i12] = GGML_CPU_FP32_TO_FP16(src[i10]); - } else { - dst_data[i10*ne12 + i12] = src[i10]; + for (int i13 = 0; i13 < ne13; i13++) { + kernel_t * const wdata_b = wdata + i13*ne10*ne11*ne12; + for (int i12 = 0; i12 < ne12; i12++) { + for (int i11 = 0; i11 < ne11; i11++) { + const float * const src = (float *)((char *) src1->data + i13*nb13 + i12*nb12 + i11*nb11); + kernel_t * dst_data = wdata_b + i11*ne10*ne12; + for (int i10 = 0; i10 < ne10; i10++) { + if constexpr (std::is_same_v) { + dst_data[i10*ne12 + i12] = GGML_CPU_FP32_TO_FP16(src[i10]); + } else { + dst_data[i10*ne12 + i12] = src[i10]; + } } } } @@ -7305,24 +7308,27 @@ static void ggml_compute_forward_conv_transpose_2d_impl( kernel_t * const wdata_src = wdata + nk; for (int i2 = ip0; i2 < ip1; i2++) { // Cout - float * dst_data = (float *)((char *) dst->data + i2*nb2); kernel_t * wdata_kernel = wdata + i2*ne01*ne00*ne03; - for (int i11 = 0; i11 < ne11; i11++) { - for (int i10 = 0; i10 < ne10; i10++) { - const int i1n = i11*ne10*ne12 + i10*ne12; - for (int i01 = 0; i01 < ne01; i01++) { - for (int i00 = 0; i00 < ne00; i00++) { - float v = 0; - if constexpr (std::is_same_v) { - ggml_vec_dot_f16(ne03, &v, 0, - wdata_src + i1n, 0, - wdata_kernel + i01*ne00*ne03 + i00*ne03, 0, 1); - } else { - ggml_vec_dot_f32(ne03, &v, 0, - wdata_src + i1n, 0, - wdata_kernel + i01*ne00*ne03 + i00*ne03, 0, 1); + for (int i3 = 0; i3 < ne3; i3++) { // batch + float * dst_data = (float *)((char *) dst->data + i3*nb3 + i2*nb2); + kernel_t * wdata_src_b = wdata_src + i3*ne10*ne11*ne12; + for (int i11 = 0; i11 < ne11; i11++) { + for (int i10 = 0; i10 < ne10; i10++) { + const int i1n = i11*ne10*ne12 + i10*ne12; + for (int i01 = 0; i01 < ne01; i01++) { + for (int i00 = 0; i00 < ne00; i00++) { + float v = 0; + if constexpr (std::is_same_v) { + ggml_vec_dot_f16(ne03, &v, 0, + wdata_src_b + i1n, 0, + wdata_kernel + i01*ne00*ne03 + i00*ne03, 0, 1); + } else { + ggml_vec_dot_f32(ne03, &v, 0, + wdata_src_b + i1n, 0, + wdata_kernel + i01*ne00*ne03 + i00*ne03, 0, 1); + } + dst_data[(i11*stride + i01)*ne0 + i10*stride + i00] += v; } - dst_data[(i11*stride + i01)*ne0 + i10*stride + i00] += v; } } } diff --git a/ggml/src/ggml-metal/ggml-metal-impl.h b/ggml/src/ggml-metal/ggml-metal-impl.h index 9becf04797b..49102afe9c0 100644 --- a/ggml/src/ggml-metal/ggml-metal-impl.h +++ b/ggml/src/ggml-metal/ggml-metal-impl.h @@ -660,6 +660,7 @@ typedef struct { uint64_t nb0; uint64_t nb1; uint64_t nb2; + uint64_t nb3; } ggml_metal_kargs_conv_transpose_2d; typedef struct { diff --git a/ggml/src/ggml-metal/ggml-metal-ops.cpp b/ggml/src/ggml-metal/ggml-metal-ops.cpp index 75de0f6dd08..f6f2fdc86c6 100644 --- a/ggml/src/ggml-metal/ggml-metal-ops.cpp +++ b/ggml/src/ggml-metal/ggml-metal-ops.cpp @@ -4645,6 +4645,7 @@ int ggml_metal_op_conv_transpose_2d(ggml_metal_op_t ctx, int idx) { const int32_t OW = op->ne[0]; const int32_t OH = op->ne[1]; const int32_t OC = op->ne[2]; + const int32_t N = op->src[1]->ne[3]; ggml_metal_kargs_conv_transpose_2d args = { /*.IC =*/ IC, @@ -4657,6 +4658,7 @@ int ggml_metal_op_conv_transpose_2d(ggml_metal_op_t ctx, int idx) { /*.nb0 =*/ nb0, /*.nb1 =*/ nb1, /*.nb2 =*/ nb2, + /*.nb3 =*/ nb3, }; auto pipeline = ggml_metal_library_get_pipeline_conv_transpose_2d(lib, op); @@ -4671,7 +4673,7 @@ int ggml_metal_op_conv_transpose_2d(ggml_metal_op_t ctx, int idx) { const size_t smem = GGML_PAD(KW * KH * sizeof(float), 16); ggml_metal_encoder_set_threadgroup_memory_size(enc, smem, 0); - ggml_metal_encoder_dispatch_threadgroups(enc, OW, OH, OC, KW, KH, 1); + ggml_metal_encoder_dispatch_threadgroups(enc, OW, OH, OC * N, KW, KH, 1); return 1; } diff --git a/ggml/src/ggml-metal/kernels/conv.metal b/ggml/src/ggml-metal/kernels/conv.metal index 5685b5cd491..a5d5aa9d929 100644 --- a/ggml/src/ggml-metal/kernels/conv.metal +++ b/ggml/src/ggml-metal/kernels/conv.metal @@ -366,7 +366,8 @@ kernel void kernel_conv_transpose_2d( const int64_t out_x = tgpig[0]; const int64_t out_y = tgpig[1]; - const int64_t out_c = tgpig[2]; + const int64_t batch = tgpig[2] / args.OC; + const int64_t out_c = tgpig[2] % args.OC; const int64_t kw = tpitg[0]; const int64_t kh = tpitg[1]; @@ -390,7 +391,7 @@ kernel void kernel_conv_transpose_2d( if (in_x >= args.IW) continue; - const int64_t input_idx = (args.IW * args.IH) * in_c + (args.IW) * in_y + in_x; + const int64_t input_idx = (args.IW * args.IH) * (args.IC * batch + in_c) + (args.IW) * in_y + in_x; const int64_t kernel_idx = (args.KH * args.KW * args.OC) * in_c + (args.KH * args.KW) * out_c + (args.KW) * kh + kw; v += (float)src0[kernel_idx] * src1[input_idx]; @@ -408,7 +409,7 @@ kernel void kernel_conv_transpose_2d( total += shared_sum[i]; } - device float * dst_ptr = (device float *) (dst + out_x*args.nb0 + out_y * args.nb1 + out_c*args.nb2); + device float * dst_ptr = (device float *) (dst + batch*args.nb3 + out_c*args.nb2 + out_y * args.nb1 + out_x*args.nb0); dst_ptr[0] = total; } } From 4c38040fd3b3241a171bbc9ef218cf34dfdd5994 Mon Sep 17 00:00:00 2001 From: Eric A Stalee <87948564+Eric-A-Stalee@users.noreply.github.com> Date: Fri, 28 Aug 2026 12:12:33 -0500 Subject: [PATCH 024/104] vulkan: fix missing view-alias dependencies in ggml_vk_graph_optimize (llama/27812) * vulkan: fix missing view-alias dependencies in ggml_vk_graph_optimize is_src_of doesn't treat two views of one tensor as dependent, so the optimizer reorders nodes across aliased reads and writes. Result: silently wrong tokens under greedy decoding, different output on every server start, and invalid speculative-decoding acceptance, with nothing logged. Hits Qwen3.8's recurrent state (and any model with view-aliased state) on AMD and NVIDIA Vulkan. CUDA is clean. Compare view_src bases on both sides. Fixes #27805 * vulkan: don't treat view/no-op nodes as aliasing dependencies Nodes whose op is NONE, RESHAPE, TRANSPOSE, VIEW or PERMUTE execute nothing, so aliasing through them is not a real dependency. The previous base comparison matched them anyway, which only costs the optimizer reordering freedom. Co-authored-by: Jeff Bolz * vulkan: make the lambda parameter const and capture is_empty in is_src_of Code will not compile without these changes. is_src_of has an empty capture list, so is_empty was not visible inside it, and is_empty took a non-const pointer, while is_src_of receives const ones. Other call sites pass non-const pointers, which still convert as usual. --------- Co-authored-by: Jeff Bolz --- ggml/src/ggml-vulkan/ggml-vulkan.cpp | 22 +++++++++++++++++----- 1 file changed, 17 insertions(+), 5 deletions(-) diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp index 28b2d875d20..320127cdc5a 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp +++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp @@ -17800,20 +17800,32 @@ static void ggml_vk_graph_optimize(ggml_backend_t backend, struct ggml_cgraph * return; } - auto const &is_empty = [](ggml_tensor * node) -> bool { + auto const &is_empty = [](const ggml_tensor * node) -> bool { return node->op == GGML_OP_NONE || node->op == GGML_OP_RESHAPE || node->op == GGML_OP_TRANSPOSE || node->op == GGML_OP_VIEW || node->op == GGML_OP_PERMUTE; }; - auto const &is_src_of = [](const ggml_tensor *dst, const ggml_tensor *src) -> bool { + auto const &is_src_of = [&is_empty](const ggml_tensor *dst, const ggml_tensor *src) -> bool { + auto const &base = [](const ggml_tensor * tensor) { + return tensor->view_src ? tensor->view_src : tensor; + }; for (uint32_t s = 0; s < GGML_MAX_SRC; ++s) { if (dst->src[s] == src) { return true; } + if (is_empty(dst) || is_empty(src)) { + continue; + } + // A source view of dst may read storage written through a different view by src. + if (dst->src[s] && base(dst->src[s]) == base(src)) { + return true; + } + // Moving dst forward may overwrite storage still read through a view by src. + if (src->src[s] && base(dst) == base(src->src[s])) { + return true; + } } // implicit dependency if they view the same tensor - const ggml_tensor *dst2 = dst->view_src ? dst->view_src : dst; - const ggml_tensor *src2 = src->view_src ? src->view_src : src; - if (dst2 == src2) { + if (base(dst) == base(src)) { return true; } return false; From d501a0a3e8a50e5fcef657e558ff95f51a91c4d5 Mon Sep 17 00:00:00 2001 From: Jeff Bolz Date: Sat, 29 Aug 2026 02:09:24 -0500 Subject: [PATCH 025/104] vulkan: Change mul_mat_id to pad K rather than N (llama/27925) The N padding is needed for mul_mat, but not mul_mat_id. For mul_mat_id, we indirect the row index through a shared memory lookup table which avoids any OOB row coordinate. But that callback doesn't bounds check K, so we actually need K padding instead. --- ggml/src/ggml-vulkan/ggml-vulkan.cpp | 92 ++++++++++--------- .../ggml-vulkan/vulkan-shaders/mul_mm.comp | 1 - .../vulkan-shaders/mul_mm_cm2.comp | 19 +++- .../ggml-vulkan/vulkan-shaders/mul_mmq.comp | 1 - 4 files changed, 62 insertions(+), 51 deletions(-) diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp index 320127cdc5a..39b4cd35980 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp +++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp @@ -1347,7 +1347,6 @@ struct vk_mat_mat_id_push_constants { uint32_t stride_a; uint32_t stride_b; uint32_t stride_d; uint32_t batch_stride_a; uint32_t batch_stride_b; uint32_t batch_stride_d; uint32_t nei0; uint32_t nei1; uint32_t nbi1; uint32_t ne11; - uint32_t padded_N; uint32_t n_experts; uint32_t hoist_row_ids; }; @@ -2403,9 +2402,8 @@ struct ggml_backend_vk_context { // Cache most recent tensor that was converted into prealloc_y, and what pipeline it used to convert. vk_pipeline_struct * prealloc_y_last_pipeline_used {}; const ggml_tensor * prealloc_y_last_tensor_used {}; - // True when prealloc_y holds the padded fp16 layout used by the coopmat2 B decode-vector callback. - // If false, then it's contiguous. - bool prealloc_y_last_decode_vector_staging {}; + // True when the K dimension in prealloc_y is padded. + bool prealloc_y_last_k_padded {}; // Track which nodes have been used since the last sync, and whether they were written to std::vector unsynced_nodes_written; @@ -8984,13 +8982,13 @@ static void ggml_vk_matmul_id( uint32_t m, uint32_t n, uint32_t k, uint32_t stride_a, uint32_t stride_b, uint32_t stride_d, uint32_t batch_stride_a, uint32_t batch_stride_b, uint32_t batch_stride_d, uint32_t n_as, uint32_t nei0, uint32_t nei1, uint32_t nbi1, uint32_t ne11, - uint32_t padded_n, bool hoist_row_ids) { + bool hoist_row_ids) { VK_LOG_DEBUG("ggml_vk_matmul_id(a: (" << a.buffer->buffer << ", " << a.offset << ", " << a.size << "), b: (" << b.buffer->buffer << ", " << b.offset << ", " << b.size << "), d: (" << d.buffer->buffer << ", " << d.offset << ", " << d.size << "), ids: (" << ids.buffer->buffer << ", " << ids.offset << ", " << ids.size << "), expert_count: (" << expert_count_buf.buffer->buffer << ", " << expert_count_buf.offset << ", " << expert_count_buf.size << "), " << "m: " << m << ", n: " << n << ", k: " << k << ", stride_a: " << stride_a << ", stride_b: " << stride_b << ", stride_d: " << stride_d << ", " << "batch_stride_a: " << batch_stride_a << ", batch_stride_b: " << batch_stride_b << ", batch_stride_d: " << batch_stride_d << ", " << "n_as: " << n_as << ", nei0: " << nei0 << ", nei1: " << nei1 << ", nbi1: " << nbi1 << ", ne11: " << ne11 << ")"); const vk_mat_mat_id_push_constants pc = { m, n, k, stride_a, stride_b, stride_d, batch_stride_a, batch_stride_b, batch_stride_d, - nei0, nei1, nbi1, ne11, padded_n, n_as, uint32_t(hoist_row_ids) }; + nei0, nei1, nbi1, ne11, n_as, uint32_t(hoist_row_ids) }; ggml_vk_dispatch_pipeline(ctx, subctx, pipeline, { a, b, d, ids, expert_count_buf }, pc, { m, nei1, n_as }); } @@ -9455,27 +9453,27 @@ static void ggml_vk_mul_mat_q_f16(ggml_backend_vk_context * ctx, vk_context& sub if (y_non_contig) { if (ctx->prealloc_y_last_pipeline_used != to_fp16_vk_1.get() || ctx->prealloc_y_last_tensor_used != src1 || - ctx->prealloc_y_last_decode_vector_staging) { + ctx->prealloc_y_last_k_padded) { if (ctx->prealloc_y_need_sync) { ggml_vk_sync_buffers(ctx, subctx); } ggml_vk_cpy_to_contiguous(ctx, subctx, to_fp16_vk_1, src1, ggml_vk_subbuffer(ctx, d_Qy, qy_buf_offset), ggml_vk_subbuffer(ctx, d_Y, 0)); ctx->prealloc_y_last_pipeline_used = to_fp16_vk_1.get(); ctx->prealloc_y_last_tensor_used = src1; - ctx->prealloc_y_last_decode_vector_staging = false; + ctx->prealloc_y_last_k_padded = false; } } if (quantize_y) { if (ctx->prealloc_y_last_pipeline_used != to_q8_1.get() || ctx->prealloc_y_last_tensor_used != src1 || - ctx->prealloc_y_last_decode_vector_staging) { + ctx->prealloc_y_last_k_padded) { if (ctx->prealloc_y_need_sync) { ggml_vk_sync_buffers(ctx, subctx); } ggml_vk_quantize_q8_1(ctx, subctx, ggml_vk_subbuffer(ctx, d_Qy, qy_buf_offset), ggml_vk_subbuffer(ctx, d_Y, 0), y_ne); ctx->prealloc_y_last_pipeline_used = to_q8_1.get(); ctx->prealloc_y_last_tensor_used = src1; - ctx->prealloc_y_last_decode_vector_staging = false; + ctx->prealloc_y_last_k_padded = false; } } @@ -9734,27 +9732,27 @@ static void ggml_vk_mul_mat_vec_q_f16(ggml_backend_vk_context * ctx, vk_context& GGML_ASSERT(y_sz == ggml_type_size(src1->type) * y_ne); if (ctx->prealloc_y_last_pipeline_used != to_fp16_vk_1.get() || ctx->prealloc_y_last_tensor_used != src1 || - ctx->prealloc_y_last_decode_vector_staging) { + ctx->prealloc_y_last_k_padded) { if (ctx->prealloc_y_need_sync) { ggml_vk_sync_buffers(ctx, subctx); } ggml_vk_cpy_to_contiguous(ctx, subctx, to_fp16_vk_1, src1, d_Qy, d_Y); ctx->prealloc_y_last_pipeline_used = to_fp16_vk_1.get(); ctx->prealloc_y_last_tensor_used = src1; - ctx->prealloc_y_last_decode_vector_staging = false; + ctx->prealloc_y_last_k_padded = false; } } if (quantize_y) { if (ctx->prealloc_y_last_pipeline_used != to_q8_1.get() || ctx->prealloc_y_last_tensor_used != src1 || - ctx->prealloc_y_last_decode_vector_staging) { + ctx->prealloc_y_last_k_padded) { if (ctx->prealloc_y_need_sync) { ggml_vk_sync_buffers(ctx, subctx); } ggml_vk_quantize_q8_1(ctx, subctx, d_Qy, d_Y, y_ne); ctx->prealloc_y_last_pipeline_used = to_q8_1.get(); ctx->prealloc_y_last_tensor_used = src1; - ctx->prealloc_y_last_decode_vector_staging = false; + ctx->prealloc_y_last_k_padded = false; } } @@ -10234,8 +10232,6 @@ static void ggml_vk_mul_mat_id_q_f16(ggml_backend_vk_context * ctx, vk_context& (src0->type == GGML_TYPE_BF16 && src1->type != GGML_TYPE_BF16) || !ggml_vk_dim01_contiguous(src1); - const uint32_t y_staged_row_stride = y_decode_vector_staging ? (uint32_t)ggml_vk_align_size(ne10, 4) : (uint32_t)ne10; - const bool y_f32_kernel = src1->type == GGML_TYPE_F32 && !y_non_contig; bool quantize_y = ctx->device->integer_dot_product && src1->type == GGML_TYPE_F32 && ggml_is_contiguous(src1) && !y_non_contig && (ne11 * ne10) % 4 == 0; @@ -10250,19 +10246,25 @@ static void ggml_vk_mul_mat_id_q_f16(ggml_backend_vk_context * ctx, vk_context& } const bool qx_needs_dequant = mmp == nullptr || x_non_contig; - const bool qy_needs_dequant = !quantize_y && ((src1->type != f16_type && !y_f32_kernel) || y_non_contig); + bool qy_needs_dequant = !quantize_y && ((src1->type != f16_type && !y_f32_kernel) || y_non_contig); if (qx_needs_dequant) { // Fall back to dequant + f16 mulmat mmp = ggml_vk_get_mul_mat_mat_id_pipeline(ctx, f16_type, y_f32_kernel ? GGML_TYPE_F32 : f16_type, (ggml_prec)dst->op_params[0]); } - // Not implemented - GGML_ASSERT(y_non_contig || !qy_needs_dequant); // NOLINT - const ggml_type effective_src1_type = quantize_y ? GGML_TYPE_Q8_1 : (y_f32_kernel ? GGML_TYPE_F32 : src1->type); const uint32_t kpad = quantize_y ? 0 : ggml_vk_align_size(ne10, ggml_vk_guess_matmul_id_pipeline_align(ctx, mmp, ne01, nei1, qx_needs_dequant ? f16_type : src0->type, effective_src1_type)); + // Coopmat2 MUL_MAT_ID BK specialization constants in ggml_vk_load_shaders are at most 64. + const uint32_t y_staged_row_stride = ctx->device->coopmat2 && !quantize_y ? ggml_vk_align_size(ne10, 64) : ne10; + const bool y_needs_k_padding = ne10 != y_staged_row_stride; + const bool y_needs_reformat = y_non_contig || y_needs_k_padding; + qy_needs_dequant = qy_needs_dequant || y_needs_k_padding; + + // Not implemented + GGML_ASSERT(y_needs_reformat || !qy_needs_dequant); // NOLINT + const bool aligned = !quantize_y && ne10 == kpad && ne01 > 8 && nei1 > 8; vk_pipeline pipeline = ggml_vk_guess_matmul_id_pipeline(ctx, mmp, ne01, nei1, aligned, qx_needs_dequant ? f16_type : src0->type, effective_src1_type); @@ -10270,10 +10272,8 @@ static void ggml_vk_mul_mat_id_q_f16(ggml_backend_vk_context * ctx, vk_context& if (ggml_nbytes(src0) > ctx->device->properties.limits.maxStorageBufferRange) { pipeline = ggml_vk_get_64b_indexing_pipeline(ctx, pipeline); } - // Reserve extra storage in the N dimension for the Y matrix, so we can avoid bounds-checking - uint32_t padded_n = qy_needs_dequant ? ROUNDUP_POW2(ne11, pipeline->wg_denoms[1]) :ne11; const uint64_t x_ne = ggml_nelements(src0); - const uint64_t y_ne = (uint64_t)y_staged_row_stride * padded_n * ne12 * ne13; + const uint64_t y_ne = (uint64_t)y_staged_row_stride * ne11 * ne12 * ne13; const uint64_t d_ne = ggml_nelements(dst); const uint64_t qx_sz = ggml_type_size(src0->type) * x_ne / ggml_blck_size(src0->type); @@ -10292,7 +10292,7 @@ static void ggml_vk_mul_mat_id_q_f16(ggml_backend_vk_context * ctx, vk_context& y_staged_dst.type = f16_type; y_staged_dst.nb[0] = ggml_type_size(f16_type); y_staged_dst.nb[1] = y_staged_dst.nb[0] * y_staged_row_stride; - y_staged_dst.nb[2] = y_staged_dst.nb[1] * padded_n; + y_staged_dst.nb[2] = y_staged_dst.nb[1] * ne11; y_staged_dst.nb[3] = y_staged_dst.nb[2] * y_staged_dst.ne[2]; return y_staged_dst; }; @@ -10302,10 +10302,10 @@ static void ggml_vk_mul_mat_id_q_f16(ggml_backend_vk_context * ctx, vk_context& } else { to_fp16_vk_0 = ggml_vk_get_to_fp16(ctx, src0->type); } - if (y_non_contig) { + if (y_needs_reformat) { ggml_tensor y_staged_dst; const ggml_tensor * y_staged_dst_ptr = nullptr; - if (y_decode_vector_staging) { + if (y_needs_k_padding) { y_staged_dst = make_y_staged_dst(); y_staged_dst_ptr = &y_staged_dst; } @@ -10432,14 +10432,18 @@ static void ggml_vk_mul_mat_id_q_f16(ggml_backend_vk_context * ctx, vk_context& ggml_vk_dispatch_pipeline(ctx, subctx, to_fp16_vk_0, { vk_subbuffer{ d_Qx, qx_buf_offset, qx_sz }, vk_subbuffer{ d_X, 0, x_sz } }, pc, { (uint32_t)x_ne, 1, 1}); } - if (y_non_contig) { + if (y_needs_reformat) { if (ctx->prealloc_y_last_pipeline_used != to_fp16_vk_1.get() || ctx->prealloc_y_last_tensor_used != src1 || - ctx->prealloc_y_last_decode_vector_staging != y_decode_vector_staging) { + ctx->prealloc_y_last_k_padded != y_needs_k_padding) { if (ctx->prealloc_y_need_sync) { ggml_vk_sync_buffers(ctx, subctx); } - if (y_decode_vector_staging) { + if (y_needs_k_padding) { + GGML_ASSERT(y_sz % 4 == 0); + // Zero B padding because clamping only A can produce 0 * Inf or NaN. + subctx->s->buffer->buf.fillBuffer(d_Y->buffer, 0, y_sz, 0); + ggml_vk_sync_buffers(ctx, subctx); const ggml_tensor y_staged_dst = make_y_staged_dst(); const uint32_t y_staged_dst_type_size = ggml_type_size(y_staged_dst.type); ggml_vk_cpy_to_strided( @@ -10454,27 +10458,27 @@ static void ggml_vk_mul_mat_id_q_f16(ggml_backend_vk_context * ctx, vk_context& } ctx->prealloc_y_last_pipeline_used = to_fp16_vk_1.get(); ctx->prealloc_y_last_tensor_used = src1; - ctx->prealloc_y_last_decode_vector_staging = y_decode_vector_staging; + ctx->prealloc_y_last_k_padded = y_needs_k_padding; } } if (quantize_y) { if (ctx->prealloc_y_last_pipeline_used != to_q8_1.get() || ctx->prealloc_y_last_tensor_used != src1 || - ctx->prealloc_y_last_decode_vector_staging) { + ctx->prealloc_y_last_k_padded) { if (ctx->prealloc_y_need_sync) { ggml_vk_sync_buffers(ctx, subctx); } ggml_vk_quantize_q8_1(ctx, subctx, ggml_vk_subbuffer(ctx, d_Qy, qy_buf_offset), ggml_vk_subbuffer(ctx, d_Y, 0), y_ne); ctx->prealloc_y_last_pipeline_used = to_q8_1.get(); ctx->prealloc_y_last_tensor_used = src1; - ctx->prealloc_y_last_decode_vector_staging = false; + ctx->prealloc_y_last_k_padded = false; } } ggml_vk_sync_buffers(ctx, subctx); uint32_t stride_batch_x = ne00*ne01; - uint32_t stride_b_y = y_decode_vector_staging ? y_staged_row_stride : ne10; - uint32_t stride_batch_y = y_decode_vector_staging ? y_staged_row_stride * padded_n : ne10*ne11; + uint32_t stride_b_y = y_needs_k_padding ? y_staged_row_stride : ne10; + uint32_t stride_batch_y = y_needs_k_padding ? y_staged_row_stride * ne11 : ne10*ne11; if (!ggml_vk_dim01_contiguous(src0) && !qx_needs_dequant) { stride_batch_x = src0->nb[0] / ggml_type_size(src0->type); @@ -10491,13 +10495,13 @@ static void ggml_vk_mul_mat_id_q_f16(ggml_backend_vk_context * ctx, vk_context& { d_D, d_buf_offset, d_sz }, { d_ids, ids_buf_offset, ids_sz }, expert_count_buf, ne01, ne21, ne10, ne10, stride_b_y, ne01, stride_batch_x, stride_batch_y, ne20*ne21, - n_as, nei0, nei1, nbi1 / ggml_type_size(ids->type), ne11, padded_n, hoist_row_ids + n_as, nei0, nei1, nbi1 / ggml_type_size(ids->type), ne11, hoist_row_ids ); // NOLINT if (x_non_contig || qx_needs_dequant) { ctx->prealloc_x_need_sync = true; } - if (y_non_contig || quantize_y) { + if (y_needs_reformat || quantize_y) { ctx->prealloc_y_need_sync = true; } ctx->prealloc_split_k_need_sync = true; @@ -10648,27 +10652,27 @@ static void ggml_vk_mul_mat_vec_id_q_f16(ggml_backend_vk_context * ctx, vk_conte GGML_ASSERT(y_sz == ggml_type_size(src1->type) * y_ne); if (ctx->prealloc_y_last_pipeline_used != to_fp16_vk_1.get() || ctx->prealloc_y_last_tensor_used != src1 || - ctx->prealloc_y_last_decode_vector_staging) { + ctx->prealloc_y_last_k_padded) { if (ctx->prealloc_y_need_sync) { ggml_vk_sync_buffers(ctx, subctx); } ggml_vk_cpy_to_contiguous(ctx, subctx, to_fp16_vk_1, src1, d_Qy, d_Y); ctx->prealloc_y_last_pipeline_used = to_fp16_vk_1.get(); ctx->prealloc_y_last_tensor_used = src1; - ctx->prealloc_y_last_decode_vector_staging = false; + ctx->prealloc_y_last_k_padded = false; } } if (quantize_y) { if (ctx->prealloc_y_last_pipeline_used != to_q8_1.get() || ctx->prealloc_y_last_tensor_used != src1 || - ctx->prealloc_y_last_decode_vector_staging) { + ctx->prealloc_y_last_k_padded) { if (ctx->prealloc_y_need_sync) { ggml_vk_sync_buffers(ctx, subctx); } ggml_vk_quantize_q8_1(ctx, subctx, d_Qy, d_Y, y_ne); ctx->prealloc_y_last_pipeline_used = to_q8_1.get(); ctx->prealloc_y_last_tensor_used = src1; - ctx->prealloc_y_last_decode_vector_staging = false; + ctx->prealloc_y_last_k_padded = false; } } @@ -15520,7 +15524,7 @@ static void ggml_vk_preallocate_buffers(ggml_backend_vk_context * ctx, vk_contex ctx->prealloc_y = ggml_vk_create_buffer_device(ctx->device, ctx->prealloc_size_y); ctx->prealloc_y_last_pipeline_used = nullptr; ctx->prealloc_y_last_tensor_used = nullptr; - ctx->prealloc_y_last_decode_vector_staging = false; + ctx->prealloc_y_last_k_padded = false; } if (ctx->prealloc_split_k == nullptr || (ctx->prealloc_size_split_k > 0 && ctx->prealloc_split_k->size < ctx->prealloc_size_split_k)) { VK_LOG_MEMORY("ggml_vk_preallocate_buffers(split_k_size: " << ctx->prealloc_size_split_k << ")"); @@ -16145,7 +16149,7 @@ static void ggml_vk_graph_cleanup(ggml_backend_vk_context * ctx) { VK_LOG_DEBUG("ggml_vk_graph_cleanup()"); ctx->prealloc_y_last_pipeline_used = {}; ctx->prealloc_y_last_tensor_used = nullptr; - ctx->prealloc_y_last_decode_vector_staging = false; + ctx->prealloc_y_last_k_padded = false; ctx->unsynced_nodes_written.clear(); ctx->unsynced_nodes_read.clear(); @@ -16197,7 +16201,7 @@ static void ggml_vk_cleanup(ggml_backend_vk_context * ctx) { ctx->prealloc_y_last_pipeline_used = nullptr; ctx->prealloc_y_last_tensor_used = nullptr; - ctx->prealloc_y_last_decode_vector_staging = false; + ctx->prealloc_y_last_k_padded = false; ctx->prealloc_size_x = 0; ctx->prealloc_size_y = 0; @@ -17395,7 +17399,7 @@ static ggml_status ggml_backend_vk_graph_compute(ggml_backend_t backend, ggml_cg ctx->prealloc_y_last_pipeline_used = nullptr; ctx->prealloc_y_last_tensor_used = nullptr; - ctx->prealloc_y_last_decode_vector_staging = false; + ctx->prealloc_y_last_k_padded = false; if (ctx->prealloc_size_add_rms_partials) { ggml_vk_preallocate_buffers(ctx, nullptr); diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm.comp b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm.comp index c1ccac7aa20..63c4aaebcb1 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm.comp @@ -88,7 +88,6 @@ layout (push_constant) uniform parameter uint nei1; uint nbi1; uint ne11; - uint padded_N; uint n_experts; uint hoist_row_ids; #else diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm_cm2.comp b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm_cm2.comp index cf78474a925..27f3178e7f2 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm_cm2.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm_cm2.comp @@ -56,6 +56,8 @@ layout (push_constant) uniform parameter uint nei1; uint nbi1; uint ne11; + uint n_experts; + uint hoist_row_ids; #else uint base_work_group_z; uint num_batches; @@ -64,12 +66,8 @@ layout (push_constant) uniform parameter uint ne12; uint broadcast2; uint broadcast3; -#endif // N dimension for the B matrix can be >= p.N uint padded_N; -#ifdef MUL_MAT_ID - uint n_experts; - uint hoist_row_ids; #endif } p; @@ -332,7 +330,9 @@ void main() { tensorLayoutNV<2> tensorLayoutA = createTensorLayoutNV(2); tensorLayoutNV<2, gl_CooperativeMatrixClampModeConstantNV> tensorLayoutAClamp = createTensorLayoutNV(2, gl_CooperativeMatrixClampModeConstantNV); tensorLayoutNV<2> tensorLayoutB = createTensorLayoutNV(2); +#ifndef MUL_MAT_ID tensorLayoutNV<2, gl_CooperativeMatrixClampModeConstantNV> tensorLayoutBClamp = createTensorLayoutNV(2, gl_CooperativeMatrixClampModeConstantNV); +#endif tensorLayoutNV<2, gl_CooperativeMatrixClampModeConstantNV> tensorLayoutD = createTensorLayoutNV(2, gl_CooperativeMatrixClampModeConstantNV); #if QUANT_K > 1 @@ -345,12 +345,19 @@ void main() { // Use end_k rather than p.K as the dimension because that's what // we need to bound check against when using split_k. - // Bounds check B against padded_N, but bounds check D against N. tensorLayoutA = setTensorLayoutDimensionNV(tensorLayoutA, p.M, end_k); +#ifdef MUL_MAT_ID + // MUL_MAT_ID pads each B row to stride_b so partial K tiles read zeros without clamping. + tensorLayoutB = setTensorLayoutDimensionNV(tensorLayoutB, BN, p.stride_b); +#else + // Bounds check B against padded_N, but bounds check D against N. tensorLayoutB = setTensorLayoutDimensionNV(tensorLayoutB, p.padded_N, end_k); +#endif tensorLayoutD = setTensorLayoutDimensionNV(tensorLayoutD, p.N, p.M); tensorLayoutAClamp = setTensorLayoutDimensionNV(tensorLayoutAClamp, p.M, end_k); +#ifndef MUL_MAT_ID tensorLayoutBClamp = setTensorLayoutDimensionNV(tensorLayoutBClamp, p.padded_N, end_k); +#endif tensorLayoutD = setTensorLayoutStrideNV(tensorLayoutD, p.stride_d, 1); @@ -527,7 +534,9 @@ void main() { tensorLayoutB = setTensorLayoutStrideNV(tensorLayoutB, stride_b, 1); +#ifndef MUL_MAT_ID tensorLayoutBClamp = setTensorLayoutStrideNV(tensorLayoutBClamp, stride_b, 1); +#endif uint k_iters = (end_k - start_k + BK - 1) / BK; diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq.comp b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq.comp index c2d84c05b40..1fbcbf6c933 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq.comp @@ -56,7 +56,6 @@ layout (push_constant) uniform parameter uint nei1; uint nbi1; uint ne11; - uint padded_N; uint n_experts; uint hoist_row_ids; #else From 590fe189021127f8a8255968e124c892293a9703 Mon Sep 17 00:00:00 2001 From: Jhen-Jie Hong Date: Sat, 29 Aug 2026 15:12:23 +0800 Subject: [PATCH 026/104] metal : add fa-vec tunings for M1 Max (llama/27932) --- ggml/src/ggml-metal/ggml-metal-tuning.cpp | 188 ++++++++++++++++++++++ 1 file changed, 188 insertions(+) diff --git a/ggml/src/ggml-metal/ggml-metal-tuning.cpp b/ggml/src/ggml-metal/ggml-metal-tuning.cpp index 285590d1601..6cfa73e6515 100644 --- a/ggml/src/ggml-metal/ggml-metal-tuning.cpp +++ b/ggml/src/ggml-metal/ggml-metal-tuning.cpp @@ -279,6 +279,194 @@ constexpr fa_vec_entry_t fa_vec_tuned_table[] = { { { GGML_METAL_DEVICE_M1_PRO, GGML_TYPE_Q8_0, 576, 512, -1, 0 }, { 1, 4 } }, { { GGML_METAL_DEVICE_M1_PRO, GGML_TYPE_Q8_0, 576, 512, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_F16, 64, 64, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_F16, 64, 64, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_F16, 128, 128, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_F16, 128, 128, 3, 3 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_F16, 128, 128, 3, 4 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_F16, 192, 128, 1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_F16, 192, 128, 2, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_F16, 192, 128, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_F16, 320, 256, 2, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_F16, 320, 256, 3, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_F16, 320, 256, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q4_0, 32, 32, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q4_0, 32, 32, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q4_0, 32, 32, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q4_0, 32, 32, 3, 2 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q4_0, 32, 32, 3, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q4_0, 64, 64, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q4_0, 64, 64, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q4_0, 64, 64, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q4_0, 64, 64, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q4_0, 64, 64, 3, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q4_0, 96, 96, 1, 3 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q4_0, 96, 96, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q4_0, 96, 96, 2, 3 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q4_0, 96, 96, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q4_0, 96, 96, 3, 3 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q4_0, 96, 96, 3, 4 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q4_0, 128, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q4_0, 128, 128, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q4_0, 192, 192, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q4_0, 192, 192, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q4_0, 192, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q4_0, 192, 128, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q4_0, 256, 256, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q4_0, 256, 256, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q4_0, 320, 256, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q4_0, 320, 256, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q4_0, 320, 256, 1, 3 }, { 4, 2 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q4_0, 512, 512, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q4_0, 512, 512, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q4_0, 576, 512, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q4_0, 576, 512, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q4_1, 32, 32, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q4_1, 32, 32, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q4_1, 32, 32, 1, 4 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q4_1, 32, 32, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q4_1, 32, 32, 3, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q4_1, 32, 32, 3, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q4_1, 64, 64, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q4_1, 64, 64, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q4_1, 64, 64, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q4_1, 96, 96, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q4_1, 96, 96, 2, 3 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q4_1, 96, 96, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q4_1, 96, 96, 3, 3 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q4_1, 128, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q4_1, 128, 128, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q4_1, 128, 128, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q4_1, 128, 128, 1, 4 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q4_1, 128, 128, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q4_1, 128, 128, 3, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q4_1, 192, 192, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q4_1, 192, 192, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q4_1, 192, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q4_1, 192, 128, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q4_1, 256, 256, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q4_1, 256, 256, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q4_1, 320, 256, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q4_1, 320, 256, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q4_1, 320, 256, 1, 3 }, { 4, 2 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q4_1, 512, 512, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q4_1, 512, 512, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q4_1, 576, 512, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q4_1, 576, 512, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q5_0, 32, 32, -1, 1 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q5_0, 32, 32, 1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q5_0, 32, 32, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q5_0, 32, 32, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q5_0, 64, 64, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q5_0, 64, 64, -1, 1 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q5_0, 64, 64, 1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q5_0, 64, 64, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q5_0, 64, 64, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q5_0, 96, 96, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q5_0, 128, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q5_0, 128, 128, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q5_0, 128, 128, 1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q5_0, 128, 128, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q5_0, 128, 128, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q5_0, 128, 128, 2, 4 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q5_0, 128, 128, 3, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q5_0, 192, 192, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q5_0, 192, 192, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q5_0, 192, 192, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q5_0, 192, 192, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q5_0, 192, 192, 3, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q5_0, 192, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q5_0, 192, 128, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q5_0, 192, 128, 1, 3 }, { 4, 2 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q5_0, 192, 128, 2, 3 }, { 4, 2 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q5_0, 192, 128, 3, 3 }, { 4, 2 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q5_0, 256, 256, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q5_0, 256, 256, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q5_0, 256, 256, 1, 3 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q5_0, 320, 256, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q5_0, 320, 256, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q5_0, 320, 256, 1, 3 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q5_0, 512, 512, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q5_0, 512, 512, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q5_0, 576, 512, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q5_0, 576, 512, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q5_1, 32, 32, -1, 1 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q5_1, 32, 32, 1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q5_1, 32, 32, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q5_1, 32, 32, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q5_1, 64, 64, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q5_1, 64, 64, -1, 1 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q5_1, 64, 64, 1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q5_1, 64, 64, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q5_1, 64, 64, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q5_1, 96, 96, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q5_1, 96, 96, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q5_1, 128, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q5_1, 128, 128, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q5_1, 128, 128, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q5_1, 128, 128, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q5_1, 128, 128, 2, 3 }, { 4, 2 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q5_1, 128, 128, 3, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q5_1, 192, 192, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q5_1, 192, 192, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q5_1, 192, 192, 1, 3 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q5_1, 192, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q5_1, 192, 128, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q5_1, 192, 128, 1, 3 }, { 4, 2 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q5_1, 192, 128, 2, 3 }, { 4, 2 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q5_1, 192, 128, 3, 3 }, { 4, 2 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q5_1, 256, 256, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q5_1, 256, 256, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q5_1, 256, 256, 1, 3 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q5_1, 256, 256, 2, 3 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q5_1, 320, 256, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q5_1, 320, 256, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q5_1, 320, 256, 1, 3 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q5_1, 320, 256, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q5_1, 320, 256, 2, 3 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q5_1, 320, 256, 2, 4 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q5_1, 320, 256, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q5_1, 320, 256, 3, 3 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q5_1, 512, 512, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q5_1, 512, 512, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q5_1, 576, 512, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q5_1, 576, 512, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q8_0, 32, 32, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q8_0, 32, 32, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q8_0, 32, 32, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q8_0, 32, 32, 2, 4 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q8_0, 32, 32, 3, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q8_0, 32, 32, 3, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q8_0, 64, 64, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q8_0, 64, 64, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q8_0, 64, 64, 1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q8_0, 64, 64, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q8_0, 64, 64, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q8_0, 64, 64, 3, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q8_0, 96, 96, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q8_0, 96, 96, 2, 3 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q8_0, 96, 96, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q8_0, 96, 96, 3, 3 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q8_0, 96, 96, 3, 4 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q8_0, 128, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q8_0, 128, 128, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q8_0, 192, 192, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q8_0, 192, 192, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q8_0, 192, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q8_0, 192, 128, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q8_0, 256, 256, -1, 0 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q8_0, 256, 256, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q8_0, 320, 256, 1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q8_0, 320, 256, 1, 3 }, { 4, 2 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q8_0, 320, 256, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q8_0, 320, 256, 2, 3 }, { 4, 2 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q8_0, 320, 256, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q8_0, 320, 256, 3, 3 }, { 4, 2 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q8_0, 512, 512, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q8_0, 512, 512, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q8_0, 576, 512, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q8_0, 576, 512, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_ULTRA, GGML_TYPE_F16, 64, 64, -1, 0 }, { 1, 4 } }, { { GGML_METAL_DEVICE_M2_ULTRA, GGML_TYPE_F16, 64, 64, -1, 1 }, { 1, 4 } }, { { GGML_METAL_DEVICE_M2_ULTRA, GGML_TYPE_F16, 128, 128, 2, 1 }, { 1, 4 } }, From e644752070ed07294219e787cc10d2697888f96d Mon Sep 17 00:00:00 2001 From: Jeff Bolz Date: Sat, 29 Aug 2026 02:59:48 -0500 Subject: [PATCH 027/104] vulkan: combine duplicated fastdiv functions, rename the one optimizing small divs (llama/27526) * vulkan: combine duplicated fastdiv functions, rename the one optimizing small divs * remove one more fastdiv --- .../ggml-vulkan/vulkan-shaders/conv2d_mm.comp | 9 +-------- .../ggml-vulkan/vulkan-shaders/conv3d_mm.comp | 9 +-------- .../vulkan-shaders/count_experts.comp | 9 +-------- .../vulkan-shaders/generic_unary_head.glsl | 14 ++------------ .../ggml-vulkan/vulkan-shaders/glu_head.glsl | 8 ++------ .../ggml-vulkan/vulkan-shaders/sum_rows.glsl | 10 ++-------- ggml/src/ggml-vulkan/vulkan-shaders/utils.glsl | 18 +++++++++++++++--- 7 files changed, 24 insertions(+), 53 deletions(-) diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/conv2d_mm.comp b/ggml/src/ggml-vulkan/vulkan-shaders/conv2d_mm.comp index 99400098bf2..c64004cdc48 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/conv2d_mm.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/conv2d_mm.comp @@ -19,6 +19,7 @@ #endif #include "types.glsl" +#include "utils.glsl" // shape notation: [dim(N), ..., dim(0)] -- stride(dim(j)) >= stride(dim(i)) if i > j layout(binding = 0) readonly buffer A { @@ -193,14 +194,6 @@ uint32_t Br = tid / BS_NPQ; uint32_t Bc = tid % BS_NPQ; const uint32_t BrpWg = WG_SIZE / BS_NPQ; -// see init_fastdiv_values in ggml-vulkan.cpp -uint fastdiv(uint n, uint mp, uint L) { - uint msbs, lsbs; - // msbs = mulhi(n, mp) - umulExtended(n, mp, msbs, lsbs); - return (msbs + n) >> L; -} - #ifdef COOPMAT2 #define ACC_TYPE float16_t diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/conv3d_mm.comp b/ggml/src/ggml-vulkan/vulkan-shaders/conv3d_mm.comp index f66f299f6da..d5ce4290b93 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/conv3d_mm.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/conv3d_mm.comp @@ -15,6 +15,7 @@ #endif #include "types.glsl" +#include "utils.glsl" // shape notation: [dim(N), ..., dim(0)] -- stride(dim(j)) >= stride(dim(i)) if i > j layout(binding = 0) readonly buffer A { @@ -178,14 +179,6 @@ uint32_t Br = tid / BS_NPQ; uint32_t Bc = tid % BS_NPQ; const uint32_t BrpWg = WG_SIZE / BS_NPQ; -// see init_fastdiv_values in ggml-vulkan.cpp -uint fastdiv(uint n, uint mp, uint L) { - uint msbs, lsbs; - // msbs = mulhi(n, mp) - umulExtended(n, mp, msbs, lsbs); - return (msbs + n) >> L; -} - void split_crs(uint32_t crs_idx, out uint32_t ic, out uint32_t kd, out uint32_t kh, out uint32_t kw) { const uint32_t KHKW = KH * KW; const uint32_t KDKHKW = KD * KHKW; diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/count_experts.comp b/ggml/src/ggml-vulkan/vulkan-shaders/count_experts.comp index 83c56fce520..ef659959d95 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/count_experts.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/count_experts.comp @@ -8,6 +8,7 @@ #endif #include "types.glsl" +#include "utils.glsl" layout (push_constant) uniform parameter { @@ -33,14 +34,6 @@ shared uint vals[BLOCK_SIZE]; shared uint offsets[BLOCK_SIZE]; shared uint cursors[BLOCK_SIZE]; -// see init_fastdiv_values in ggml-vulkan.cpp -uint fastdiv(uint n, uint mp, uint L) { - uint msbs, lsbs; - // msbs = mulhi(n, mp) - umulExtended(n, mp, msbs, lsbs); - return (msbs + n) >> L; -} - // data_d layout when p.hoist_row_ids is set: // [0, n_experts) per-expert row count // [n_experts, 2*n_experts) per-expert start offset into the row id region diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/generic_unary_head.glsl b/ggml/src/ggml-vulkan/vulkan-shaders/generic_unary_head.glsl index 9d4176f3f96..e13de9a00f2 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/generic_unary_head.glsl +++ b/ggml/src/ggml-vulkan/vulkan-shaders/generic_unary_head.glsl @@ -1,6 +1,8 @@ #extension GL_EXT_shader_16bit_storage : require #extension GL_EXT_control_flow_attributes : require +#include "utils.glsl" + layout (push_constant) uniform parameter { uint ne; @@ -32,18 +34,6 @@ uint get_idx() { uint get_aoffset() { return p.misalign_offsets >> 16; } uint get_doffset() { return p.misalign_offsets & 0xFFFF; } -// see init_fastdiv_values in ggml-vulkan.cpp -uint fastdiv(uint n, uint mp, uint L) { - uint msbs, lsbs; - // msbs = mulhi(n, mp) - umulExtended(n, mp, msbs, lsbs); - return (msbs + n) >> L; -} - -uint fastdiv_L(uint packed, uint slot) { - return (packed >> (slot * 8)) & 0x3Fu; -} - uint src0_idx(uint idx) { const uint i03 = fastdiv(idx, p.ne0_012mp, fastdiv_L(p.ne0_Ls, 0)); const uint i03_offset = i03 * p.ne02*p.ne01*p.ne00; diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/glu_head.glsl b/ggml/src/ggml-vulkan/vulkan-shaders/glu_head.glsl index c3cae736f97..fc2951ec2e5 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/glu_head.glsl +++ b/ggml/src/ggml-vulkan/vulkan-shaders/glu_head.glsl @@ -1,5 +1,7 @@ #extension GL_EXT_shader_16bit_storage : require +#include "utils.glsl" + layout(local_size_x = 512, local_size_y = 1, local_size_z = 1) in; @@ -39,9 +41,3 @@ uint get_aoffset() { return p.misalign_offsets >> 16; } uint get_boffset() { return (p.misalign_offsets >> 8) & 0xFF; } uint get_doffset() { return p.misalign_offsets & 0xFF; } -// see init_fastdiv_values in ggml-vulkan.cpp -uint fastdiv(uint n, uint mp, uint L) { - uint msbs, lsbs; - umulExtended(n, mp, msbs, lsbs); - return (msbs + n) >> L; -} diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/sum_rows.glsl b/ggml/src/ggml-vulkan/vulkan-shaders/sum_rows.glsl index 2b841baa6bf..1cb0f7827a3 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/sum_rows.glsl +++ b/ggml/src/ggml-vulkan/vulkan-shaders/sum_rows.glsl @@ -1,4 +1,6 @@ +#include "utils.glsl" + // vk_op_sum_rows_push_constants layout (push_constant) uniform parameter { @@ -15,11 +17,3 @@ layout (push_constant) uniform parameter uint get_aoffset() { return p.misalign_offsets >> 16; } uint get_doffset() { return p.misalign_offsets & 0xFFFF; } -// see init_fastdiv_values in ggml-vulkan.cpp -uint fastdiv(uint n, uint mp, uint L) { - uint msbs, lsbs; - // msbs = mulhi(n, mp) - umulExtended(n, mp, msbs, lsbs); - return (msbs + n) >> L; -} - diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/utils.glsl b/ggml/src/ggml-vulkan/vulkan-shaders/utils.glsl index dc4a1e6d96b..8aac64d7593 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/utils.glsl +++ b/ggml/src/ggml-vulkan/vulkan-shaders/utils.glsl @@ -9,14 +9,26 @@ uint fastmod(uint a, uint b) { return a % b; } -uint fastdiv(uint a, uint b) { +// see init_fastdiv_values in ggml-vulkan.cpp +uint fastdiv(uint n, uint mp, uint L) { + uint msbs, lsbs; + // msbs = mulhi(n, mp) + umulExtended(n, mp, msbs, lsbs); + return (msbs + n) >> L; +} + +uint fastdiv_L(uint packed, uint slot) { + return (packed >> (slot * 8)) & 0x3Fu; +} + +uint fastdiv_small(uint a, uint b) { return (a < b) ? 0 : (a / b); } void get_indices(uint idx, out uint i00, out uint i01, out uint i02, out uint i03, uint ne00, uint ne01, uint ne02, uint ne03) { - i03 = fastdiv(idx, (ne02*ne01*ne00)); + i03 = fastdiv_small(idx, (ne02*ne01*ne00)); const uint i03_offset = i03 * ne02*ne01*ne00; - i02 = fastdiv((idx - i03_offset), (ne01*ne00)); + i02 = fastdiv_small((idx - i03_offset), (ne01*ne00)); const uint i02_offset = i02*ne01*ne00; i01 = (idx - i03_offset - i02_offset) / ne00; i00 = idx - i03_offset - i02_offset - i01*ne00; From 325c8d16c1264ae0b0f057b45bcb87be32a8f89f Mon Sep 17 00:00:00 2001 From: Nick Farrell Date: Sat, 29 Aug 2026 19:00:09 +1000 Subject: [PATCH 028/104] sycl: make --fit respect --fit-target better (llama/27629) improve the --fit algorithm to take into account the actual peak required VRAM for a given context size on a SYCL backend. This includes both properly accounting for how much VRAM is required when the allocated context is fully used (which makes the reported context drop below what it did before, but stop it OOMing) as well as preventing some overly-conservative calculations which meant too much VRAM was being reserved. Tested on a Arc b70 with unsloth's qwen3.8 (Q4_K_XL), able to get 262144 context, fully usable, with q8_0 KV and MTP and 4k ubatch size using --fit-target 1 --- ggml/src/ggml-sycl/fattn-common.hpp | 22 ++++---- ggml/src/ggml-sycl/fattn-onednn.cpp | 84 ++++++++++++++++++----------- ggml/src/ggml-sycl/fattn-onednn.hpp | 6 ++- ggml/src/ggml-sycl/fattn.cpp | 73 +++++++++++++++++++++++++ ggml/src/ggml-sycl/fattn.hpp | 18 +++++++ ggml/src/ggml-sycl/ggml-sycl.cpp | 5 +- 6 files changed, 167 insertions(+), 41 deletions(-) diff --git a/ggml/src/ggml-sycl/fattn-common.hpp b/ggml/src/ggml-sycl/fattn-common.hpp index c6cc13cfb00..82813f7a99a 100644 --- a/ggml/src/ggml-sycl/fattn-common.hpp +++ b/ggml/src/ggml-sycl/fattn-common.hpp @@ -6,6 +6,7 @@ #include "convert.hpp" #include "vecdotq.hpp" #include "fattn-buffers.hpp" +#include "fattn.hpp" #include "ggml.h" @@ -926,6 +927,7 @@ void launch_fattn( ggml_sycl_fattn_alloc K_f16(fbuf.K); ggml_sycl_fattn_alloc V_f16(fbuf.V); + const ggml_sycl_fattn_extra extra = ggml_sycl_fattn_get_extra(dst); ggml_sycl_pool_alloc KV_max(pool); ggml_sycl_pool_alloc dst_tmp(pool); ggml_sycl_pool_alloc dst_tmp_meta(pool); @@ -944,10 +946,11 @@ void launch_fattn( const size_t bs = ggml_blck_size(K->type); const size_t ts = ggml_type_size(K->type); - K_f16.alloc(ggml_nelements(K)); + sycl::half * K_f16_ptr = extra.K_buffer_ptr ? (sycl::half *) extra.K_buffer_ptr + : K_f16.alloc(ggml_nelements(K)); if (ggml_is_contiguously_allocated(K)) { to_fp16_sycl_t to_fp16 = ggml_get_to_fp16_sycl(K->type, dst); - to_fp16(K_data, K_f16.ptr, ggml_nelements(K), main_stream); + to_fp16(K_data, K_f16_ptr, ggml_nelements(K), main_stream); nb11 = nb11 * bs * sizeof(sycl::half) / ts; nb12 = nb12 * bs * sizeof(sycl::half) / ts; @@ -958,13 +961,13 @@ void launch_fattn( const int64_t s01 = nb11 / ts; const int64_t s02 = nb12 / ts; const int64_t s03 = nb13 / ts; - to_fp16(K_data, K_f16.ptr, K->ne[0], K->ne[1], K->ne[2], K->ne[3], s01, s02, s03, main_stream); + to_fp16(K_data, K_f16_ptr, K->ne[0], K->ne[1], K->ne[2], K->ne[3], s01, s02, s03, main_stream); nb11 = K->ne[0] * sizeof(sycl::half); nb12 = K->ne[1] * nb11; nb13 = K->ne[2] * nb12; } - K_data = (char *) K_f16.ptr; + K_data = (char *) K_f16_ptr; } if (need_f16_V && V->type != GGML_TYPE_F16) { @@ -977,11 +980,12 @@ void launch_fattn( const size_t bs = ggml_blck_size(V->type); const size_t ts = ggml_type_size(V->type); - V_f16.alloc(ggml_nelements(V)); + sycl::half * V_f16_ptr = extra.V_buffer_ptr ? (sycl::half *) extra.V_buffer_ptr + : V_f16.alloc(ggml_nelements(V)); if (ggml_is_contiguously_allocated(V)) { to_fp16_sycl_t to_fp16 = ggml_get_to_fp16_sycl(V->type, dst); - to_fp16(V_data, V_f16.ptr, ggml_nelements(V), main_stream); - V_data = (char *) V_f16.ptr; + to_fp16(V_data, V_f16_ptr, ggml_nelements(V), main_stream); + V_data = (char *) V_f16_ptr; nb21 = nb21 * bs * sizeof(sycl::half) / ts; nb22 = nb22 * bs * sizeof(sycl::half) / ts; @@ -992,13 +996,13 @@ void launch_fattn( const int64_t s01 = nb21 / ts; const int64_t s02 = nb22 / ts; const int64_t s03 = nb23 / ts; - to_fp16(V_data, V_f16.ptr, V->ne[0], V->ne[1], V->ne[2], V->ne[3], s01, s02, s03, main_stream); + to_fp16(V_data, V_f16_ptr, V->ne[0], V->ne[1], V->ne[2], V->ne[3], s01, s02, s03, main_stream); nb21 = V->ne[0] * sizeof(sycl::half); nb22 = V->ne[1] * nb21; nb23 = V->ne[2] * nb22; } - V_data = (char *) V_f16.ptr; + V_data = (char *) V_f16_ptr; } } diff --git a/ggml/src/ggml-sycl/fattn-onednn.cpp b/ggml/src/ggml-sycl/fattn-onednn.cpp index d41c2ddce34..4349363a3d3 100644 --- a/ggml/src/ggml-sycl/fattn-onednn.cpp +++ b/ggml/src/ggml-sycl/fattn-onednn.cpp @@ -14,9 +14,21 @@ // set minimum query length to treat as prefill (32) #define GGML_SYCL_FA_ONEDNN_MIN_Q 32 -bool ggml_sycl_flash_attn_ext_onednn_supported(const ggml_tensor * dst) { +bool ggml_sycl_fattn_onednn_binds_kv(const ggml_tensor * K, const ggml_tensor * V) { + if (K->type != GGML_TYPE_F16 || V->type != GGML_TYPE_F16) { + return false; + } + auto bindable = [](const ggml_tensor * t) { + return t->nb[0] == sizeof(sycl::half) && t->nb[1] % sizeof(sycl::half) == 0 && + t->nb[2] % sizeof(sycl::half) == 0 && t->nb[3] % sizeof(sycl::half) == 0; + }; + return bindable(K) && bindable(V); +} + +bool ggml_sycl_flash_attn_ext_onednn_supported(const ggml_tensor * dst, bool use_shape_limit) { #if !GGML_SYCL_DNNL GGML_UNUSED(dst); + GGML_UNUSED(use_shape_limit); return false; #else if (!g_ggml_sycl_fa_onednn) { @@ -44,7 +56,7 @@ bool ggml_sycl_flash_attn_ext_onednn_supported(const ggml_tensor * dst) { if (!k_ok || !v_ok) { return false; } - if (Q->ne[1] < 32 || K->ne[1] < 1024) { + if (use_shape_limit && (Q->ne[1] < 32 || K->ne[1] < 1024)) { return false; } for (const ggml_tensor * t : {K, V}) { @@ -94,7 +106,7 @@ bool ggml_sycl_flash_attn_ext_onednn_supported(const ggml_tensor * dst) { return false; } // Prefill only. - if (Q->ne[1] < GGML_SYCL_FA_ONEDNN_MIN_Q) { + if (use_shape_limit && Q->ne[1] < GGML_SYCL_FA_ONEDNN_MIN_Q) { return false; } return true; @@ -240,9 +252,16 @@ void ggml_sycl_flash_attn_ext_onednn(ggml_backend_sycl_context & ctx, ggml_tenso dnnl::engine eng = ctx.engine_dnnl(stream); dnnl::stream strm = ctx.stream_dnnl(stream); + const ggml_sycl_fattn_extra extra = ggml_sycl_fattn_get_extra(dst); + // Q: always f32 -- copy to dense f16. - ggml_sycl_pool_alloc Qf(ctx.pool(), (size_t) H * q * d); - cont_to_f16_sycl((const char *) Q->data, Qf.get(), d, q, H, mb, Q->nb[1], Q->nb[2], Q->nb[3], stream); + std::optional> Qf_pool; + sycl::half * Qf_ptr = (sycl::half *) extra.Q_buffer_ptr; + if (!Qf_ptr) { + Qf_pool.emplace(ctx.pool(), (size_t) H * q * d); + Qf_ptr = Qf_pool->get(); + } + cont_to_f16_sycl((const char *) Q->data, Qf_ptr, d, q, H, mb, Q->nb[1], Q->nb[2], Q->nb[3], stream); // K/V: bind the f16 cache in place. llama.cpp permutes it to [token][head][dim], so its head // plane is strided rather than dense, which is what an explicit stride vector expresses. @@ -253,11 +272,12 @@ void ggml_sycl_flash_attn_ext_onednn(ggml_backend_sycl_context & ctx, ggml_tenso std::array v_str = k_str; std::optional> Kf_pool; std::optional> Vf_pool; + // Helper: hand out reserved space, or fall back to the pool. + auto stage_k = [&](size_t n) { if (extra.K_buffer_ptr) { return (sycl::half *) extra.K_buffer_ptr; } + Kf_pool.emplace(ctx.pool(), n); return Kf_pool->get(); }; + auto stage_v = [&](size_t n) { if (extra.V_buffer_ptr) { return (sycl::half *) extra.V_buffer_ptr; } + Vf_pool.emplace(ctx.pool(), n); return Vf_pool->get(); }; - auto bindable = [](const ggml_tensor * t) { - return t->nb[0] == sizeof(sycl::half) && t->nb[1] % sizeof(sycl::half) == 0 && - t->nb[2] % sizeof(sycl::half) == 0 && t->nb[3] % sizeof(sycl::half) == 0; - }; auto elem_strides = [](const ggml_tensor * t) { const int64_t s1 = (int64_t) (t->nb[1] / t->nb[0]); const int64_t s2 = (int64_t) (t->nb[2] / t->nb[0]); @@ -266,22 +286,19 @@ void ggml_sycl_flash_attn_ext_onednn(ggml_backend_sycl_context & ctx, ggml_tenso return std::array{ s3, s2, s2, s1, 1 }; }; - if (K->type == GGML_TYPE_F16 && V->type == GGML_TYPE_F16 && bindable(K) && bindable(V)) { + if (ggml_sycl_fattn_onednn_binds_kv(K, V)) { K_ptr = (sycl::half *) K->data; V_ptr = (sycl::half *) V->data; k_str = elem_strides(K); v_str = elem_strides(V); } else if (K->type == GGML_TYPE_F16 && V->type == GGML_TYPE_F16) { - Kf_pool.emplace(ctx.pool(), (size_t) Hkv * seq * d); - Vf_pool.emplace(ctx.pool(), (size_t) Hkv * seq * d); - cont_to_f16_sycl((const char *) K->data, Kf_pool->get(), d, seq, Hkv, mb, K->nb[1], K->nb[2], K->nb[3], stream); - cont_to_f16_sycl((const char *) V->data, Vf_pool->get(), d, seq, Hkv, mb, V->nb[1], V->nb[2], V->nb[3], stream); - K_ptr = Kf_pool->get(); - V_ptr = Vf_pool->get(); + K_ptr = stage_k((size_t) Hkv * seq * d); + V_ptr = stage_v((size_t) Hkv * seq * d); + cont_to_f16_sycl((const char *) K->data, K_ptr, d, seq, Hkv, mb, K->nb[1], K->nb[2], K->nb[3], stream); + cont_to_f16_sycl((const char *) V->data, V_ptr, d, seq, Hkv, mb, V->nb[1], V->nb[2], V->nb[3], stream); } else if (ggml_is_quantized(K->type)) { // Quantized K/V: dequant to dense F16 using pool, same lifetime as F16 path. - Kf_pool.emplace(ctx.pool(), ggml_nelements(K)); - K_ptr = Kf_pool->get(); + K_ptr = stage_k((size_t) ggml_nelements(K)); { const char * K_data = (const char *)K->data; const bool k_non_dense = ((int64_t)K->ne[1] * K->nb[1] != K->nb[2]) && K->ne[2] > 1; @@ -315,8 +332,7 @@ void ggml_sycl_flash_attn_ext_onednn(ggml_backend_sycl_context & ctx, ggml_tenso // data pointer), their logical values differ because the quantized // elements at different positions/offsets represent different K/V // data. Master's F16 path also never aliases K and V. - Vf_pool.emplace(ctx.pool(), ggml_nelements(V)); - V_ptr = Vf_pool->get(); + V_ptr = stage_v((size_t) ggml_nelements(V)); { const char * V_data = (const char *)V->data; const bool v_non_dense = ((int64_t)V->ne[1] * V->nb[1] != V->nb[2]) && V->ne[2] > 1; @@ -347,12 +363,10 @@ void ggml_sycl_flash_attn_ext_onednn(ggml_backend_sycl_context & ctx, ggml_tenso } } else { // F32: strided copy to dense F16 via cont_to_f16_sycl. - Kf_pool.emplace(ctx.pool(), ggml_nelements(K)); - K_ptr = Kf_pool->get(); + K_ptr = stage_k((size_t) ggml_nelements(K)); cont_to_f16_sycl((const char *) K->data, K_ptr, K->ne[0], K->ne[1], K->ne[2], K->ne[3], K->nb[1], K->nb[2], K->nb[3], stream); - Vf_pool.emplace(ctx.pool(), ggml_nelements(V)); - V_ptr = Vf_pool->get(); + V_ptr = stage_v((size_t) ggml_nelements(V)); cont_to_f16_sycl((const char *) V->data, V_ptr, V->ne[0], V->ne[1], V->ne[2], V->ne[3], V->nb[1], V->nb[2], V->nb[3], stream); } @@ -366,11 +380,21 @@ void ggml_sycl_flash_attn_ext_onednn(ggml_backend_sycl_context & ctx, ggml_tenso // instead -- the value is captured into the command, so no host memory has to outlive the // call, and the enqueue stays async. const sycl::half scale_h = (sycl::half) (1.0f / kq_scale); - ggml_sycl_pool_alloc scbuf(ctx.pool(), 1); - sycl::half * const scale_dev = scbuf.get(); + std::optional> scbuf; + sycl::half * scale_dev = (sycl::half *) extra.scale_buffer_ptr; + if (!scale_dev) { + scbuf.emplace(ctx.pool(), 1); + scale_dev = scbuf->get(); + } stream->single_task([=]() { *scale_dev = scale_h; }); - ggml_sycl_pool_alloc outf(ctx.pool(), (size_t) H * q * d); // f16 contiguous SDPA out [mb,H,q,d] + // f16 contiguous SDPA out [mb,H,q,d] + std::optional> outf_pool; + sycl::half * outf_ptr = (sycl::half *) extra.out_buffer_ptr; + if (!outf_ptr) { + outf_pool.emplace(ctx.pool(), (size_t) H * q * d); + outf_ptr = outf_pool->get(); + } // compile once per (device, shape, KV strides), reuse across layers/calls. Stride 2 always // repeats stride 1 and stride 4 is always 1, so the key covers every entry that can differ. @@ -392,7 +416,7 @@ void ggml_sycl_flash_attn_ext_onednn(ggml_backend_sycl_context & ctx, ggml_tenso } auto id2ptr = [&](size_t r) -> void * { - if (r == E.id_q) return Qf.get(); + if (r == E.id_q) return Qf_ptr; if (r == E.id_k) return K_ptr; if (r == E.id_v) return V_ptr; if (r == E.id_scale) return scale_dev; @@ -404,10 +428,10 @@ void ggml_sycl_flash_attn_ext_onednn(ggml_backend_sycl_context & ctx, ggml_tenso for (auto & lt : E.ins) { ti.emplace_back(lt, eng, id2ptr(lt.get_id())); } - tensor to(E.out, eng, outf.get()); + tensor to(E.out, eng, outf_ptr); E.cp.execute(strm, ti, {to}); - permute_sdpa_out_sycl(outf.get(), (float *) dst->data, mb, H, q, d, stream); + permute_sdpa_out_sycl(outf_ptr, (float *) dst->data, mb, H, q, d, stream); // Single device needs no sync: the dnnl stream wraps this same in-order queue, so the SDPA // serializes with the staging kernels before it and the permute/pool reuse after it. The // garbage output formerly blamed on the missing sync here was the scale use-after-return diff --git a/ggml/src/ggml-sycl/fattn-onednn.hpp b/ggml/src/ggml-sycl/fattn-onednn.hpp index d3019e87688..9669d1bd27a 100644 --- a/ggml/src/ggml-sycl/fattn-onednn.hpp +++ b/ggml/src/ggml-sycl/fattn-onednn.hpp @@ -5,7 +5,11 @@ // Static-only check: fused-XMX oneDNN Graph SDPA path==flash-attn op // (f16 KV, no softcap/ALiBi, single stream, tuned head_dim, prefill-sized q.) -bool ggml_sycl_flash_attn_ext_onednn_supported(const ggml_tensor * dst); +bool ggml_sycl_flash_attn_ext_onednn_supported(const ggml_tensor * dst, bool use_shape_limit = true); + +// True when the oneDNN path binds an F16 KV cache in place instead of staging a dense copy of +// it. Depends only on the types and strides of K and V, so the answer holds for every call. +bool ggml_sycl_fattn_onednn_binds_kv(const ggml_tensor * K, const ggml_tensor * V); // Run flash attention through oneDNN's fused xmx SDPA // execute the cached SDPA partition, write the f32 dst. Falls back to the TILE kernel on any failure. diff --git a/ggml/src/ggml-sycl/fattn.cpp b/ggml/src/ggml-sycl/fattn.cpp index a85bca7cb3c..b73e6d46ffa 100644 --- a/ggml/src/ggml-sycl/fattn.cpp +++ b/ggml/src/ggml-sycl/fattn.cpp @@ -378,3 +378,76 @@ void ggml_sycl_flash_attn_ext(ggml_backend_sycl_context & ctx, ggml_tensor * dst bool ggml_sycl_flash_attn_ext_supported(int device, const ggml_tensor * dst) { return ggml_sycl_get_best_fattn_kernel(device, dst) != BEST_FATTN_KERNEL_NONE; } + +static uintptr_t ggml_sycl_fattn_reserve_halves(ggml_sycl_fattn_extra & extra, size_t n_halves) { + if (n_halves == 0) { + return 0; + } + extra.end = GGML_PAD(extra.end, SYCL_BUFFER_ALIGNMENT); + const uintptr_t block = extra.end; + extra.end += n_halves * sizeof(sycl::half); + return block; +} + +ggml_sycl_fattn_extra ggml_sycl_fattn_get_extra(const ggml_tensor * dst) { + ggml_sycl_fattn_extra extra; + + extra.end = (uintptr_t) dst->data + ggml_nbytes(dst); + + if (dst->op != GGML_OP_FLASH_ATTN_EXT) { + return extra; + } + + const ggml_tensor * Q = dst->src[0]; + const ggml_tensor * K = dst->src[1]; + const ggml_tensor * V = dst->src[2]; + if (!Q || !K || !V) { + return extra; + } + + const int64_t d = K->ne[0]; + const int64_t H = Q->ne[2]; + const int64_t q = Q->ne[1]; + + // calculate the worst-case memory consumption across all kernels + const bool onednn_supported = ggml_sycl_flash_attn_ext_onednn_supported(dst, /* use_shape_limit */ false); + + const bool tile_needs_K = K->type != GGML_TYPE_F16; + const bool tile_needs_V = V->type != GGML_TYPE_F16; + + const bool V_is_K_view = V->view_src && + (V->view_src == K || (V->view_src == K->view_src && V->view_offs == K->view_offs)); + + size_t need_K = 0, need_V = 0, need_Q = 0, need_out = 0, need_scale = 0; + if (onednn_supported) { + need_Q = (size_t) H * q * d; + need_out = (size_t) H * q * d; + need_scale = 1; + // an f16 cache is bound in place, so it needs no staging copy + if (!ggml_sycl_fattn_onednn_binds_kv(K, V)) { + need_K = (size_t) ggml_nelements(K); + need_V = (size_t) ggml_nelements(V); + } + } + if (tile_needs_K) { + need_K = std::max(need_K, (size_t) ggml_nelements(K)); + } + if (tile_needs_V) { + need_V = std::max(need_V, (size_t) ggml_nelements(V)); + } + + extra.Q_buffer_ptr = ggml_sycl_fattn_reserve_halves(extra, need_Q); + extra.K_buffer_ptr = ggml_sycl_fattn_reserve_halves(extra, need_K); + extra.V_buffer_ptr = (V_is_K_view && !onednn_supported && need_V) + ? extra.K_buffer_ptr + : ggml_sycl_fattn_reserve_halves(extra, need_V); + extra.scale_buffer_ptr = ggml_sycl_fattn_reserve_halves(extra, need_scale); + extra.out_buffer_ptr = ggml_sycl_fattn_reserve_halves(extra, need_out); + + return extra; +} + +size_t ggml_sycl_flash_attn_ext_get_alloc_size(const ggml_tensor * dst) { + const ggml_sycl_fattn_extra extra = ggml_sycl_fattn_get_extra(dst); + return (size_t) (extra.end - (uintptr_t) dst->data); +} diff --git a/ggml/src/ggml-sycl/fattn.hpp b/ggml/src/ggml-sycl/fattn.hpp index c093970a3fe..f803aa2a804 100644 --- a/ggml/src/ggml-sycl/fattn.hpp +++ b/ggml/src/ggml-sycl/fattn.hpp @@ -19,6 +19,24 @@ void ggml_sycl_flash_attn_ext(ggml_backend_sycl_context & ctx, ggml_tensor * dst bool ggml_sycl_flash_attn_ext_supported(int device, const ggml_tensor * dst); +// Scratch that flash attention needs beyond the output tensor +struct ggml_sycl_fattn_extra { + uintptr_t K_buffer_ptr = 0; // F16 copy of the K cache + uintptr_t V_buffer_ptr = 0; // F16 copy of the V cache + uintptr_t Q_buffer_ptr = 0; // dense F16 copy of Q, oneDNN only + uintptr_t scale_buffer_ptr = 0; // the softmax scale as an F16 scalar, oneDNN only + uintptr_t out_buffer_ptr = 0; // F16 SDPA output before conversion to F32, oneDNN only + uintptr_t end = 0; // one past the last reserved byte; sizes the allocation +}; + +// ggml_sycl_fattn_get_extra() is the single source of truth for the layout: it both sizes +// the reservation and hands out the pointers, so the two cannot disagree. +// Each field is the address of one reserved block, or 0 if that block was not reserved, +// in which case the caller allocates from the scratch pool instead. +ggml_sycl_fattn_extra ggml_sycl_fattn_get_extra(const ggml_tensor * dst); + +size_t ggml_sycl_flash_attn_ext_get_alloc_size(const ggml_tensor * dst); + void ggml_sycl_flash_attn_ext_mkl(ggml_backend_sycl_context & ctx, ggml_tensor * dst); #endif // GGML_SYCL_FATTN_HPP diff --git a/ggml/src/ggml-sycl/ggml-sycl.cpp b/ggml/src/ggml-sycl/ggml-sycl.cpp index 0573643d834..dc8a1744323 100644 --- a/ggml/src/ggml-sycl/ggml-sycl.cpp +++ b/ggml/src/ggml-sycl/ggml-sycl.cpp @@ -955,7 +955,10 @@ static size_t ggml_backend_sycl_buffer_type_get_max_size(ggml_backend_buffer_typ } static size_t ggml_backend_sycl_buffer_type_get_alloc_size(ggml_backend_buffer_type_t buft, const ggml_tensor * tensor) { - size_t size = ggml_nbytes(tensor); + // Reserve the additional scratch so it's visible to the graph allocator + size_t size = tensor->op == GGML_OP_FLASH_ATTN_EXT + ? ggml_sycl_flash_attn_ext_get_alloc_size(tensor) + : ggml_nbytes(tensor); int64_t ne0 = tensor->ne[0]; if (ggml_is_quantized(tensor->type)) { From 308fa4f8a2beec1f76062727d4b1ca0a4dcbafff Mon Sep 17 00:00:00 2001 From: Niklas Wenzel Date: Sat, 29 Aug 2026 14:50:13 +0200 Subject: [PATCH 029/104] metal : add remaining fa-vec tunings for M4 Pro (llama/27915) --- ggml/src/ggml-metal/ggml-metal-tuning.cpp | 150 ++++++++++++++++++++++ 1 file changed, 150 insertions(+) diff --git a/ggml/src/ggml-metal/ggml-metal-tuning.cpp b/ggml/src/ggml-metal/ggml-metal-tuning.cpp index 6cfa73e6515..26b634d2143 100644 --- a/ggml/src/ggml-metal/ggml-metal-tuning.cpp +++ b/ggml/src/ggml-metal/ggml-metal-tuning.cpp @@ -981,6 +981,156 @@ constexpr fa_vec_entry_t fa_vec_tuned_table[] = { { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_F16, 320, 256, 3, 0 }, { 2, 2 } }, { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_F16, 320, 256, -1, 1 }, { 2, 2 } }, { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_F16, 512, 512, 3, 0 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q4_0, 32, 32, -1, 1 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q4_0, 32, 32, 1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q4_0, 32, 32, 1, 4 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q4_0, 32, 32, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q4_0, 32, 32, 2, 4 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q4_0, 32, 32, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q4_0, 64, 64, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q4_0, 64, 64, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q4_0, 64, 64, 3, 2 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q4_0, 64, 64, 3, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q4_0, 96, 96, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q4_0, 96, 96, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q4_0, 96, 96, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q4_0, 128, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q4_0, 128, 128, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q4_0, 128, 128, 3, 3 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q4_0, 192, 192, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q4_0, 192, 192, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q4_0, 192, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q4_0, 192, 128, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q4_0, 192, 128, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q4_0, 192, 128, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q4_0, 192, 128, 2, 4 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q4_0, 192, 128, 3, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q4_0, 256, 256, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q4_0, 256, 256, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q4_0, 320, 256, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q4_0, 320, 256, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q4_0, 512, 512, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q4_0, 512, 512, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q4_0, 576, 512, 2, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q4_0, 576, 512, 3, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q4_0, 576, 512, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q4_0, 576, 512, 1, 1 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q4_0, 576, 512, 1, 2 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q4_0, 576, 512, 1, 3 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q4_0, 576, 512, 1, 4 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q4_1, 32, 32, -1, 1 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q4_1, 32, 32, 1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q4_1, 32, 32, 1, 4 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q4_1, 32, 32, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q4_1, 32, 32, 2, 4 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q4_1, 32, 32, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q4_1, 64, 64, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q4_1, 64, 64, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q4_1, 64, 64, 3, 2 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q4_1, 96, 96, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q4_1, 96, 96, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q4_1, 96, 96, 1, 4 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q4_1, 96, 96, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q4_1, 96, 96, 3, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q4_1, 128, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q4_1, 128, 128, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q4_1, 128, 128, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q4_1, 128, 128, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q4_1, 128, 128, 3, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q4_1, 192, 192, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q4_1, 192, 192, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q4_1, 192, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q4_1, 192, 128, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q4_1, 192, 128, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q4_1, 192, 128, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q4_1, 192, 128, 3, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q4_1, 256, 256, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q4_1, 256, 256, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q4_1, 320, 256, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q4_1, 512, 512, 3, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q4_1, 576, 512, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q4_1, 576, 512, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q5_0, 32, 32, -1, 1 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q5_0, 32, 32, 1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q5_0, 32, 32, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q5_0, 32, 32, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q5_0, 64, 64, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q5_0, 64, 64, -1, 1 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q5_0, 64, 64, 1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q5_0, 64, 64, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q5_0, 64, 64, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q5_0, 96, 96, -1, 1 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q5_0, 96, 96, 1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q5_0, 96, 96, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q5_0, 96, 96, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q5_0, 128, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q5_0, 128, 128, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q5_0, 128, 128, 2, 2 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q5_0, 128, 128, 2, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q5_0, 192, 192, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q5_0, 192, 192, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q5_0, 192, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q5_0, 192, 128, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q5_0, 256, 256, -1, 0 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q5_0, 256, 256, -1, 1 }, { 2, 2 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q5_0, 256, 256, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q5_0, 256, 256, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q5_0, 256, 256, 3, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q5_0, 320, 256, -1, 1 }, { 2, 2 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q5_0, 320, 256, 1, 2 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q5_0, 320, 256, 2, 2 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q5_0, 320, 256, 2, 4 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q5_0, 320, 256, 3, 2 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q5_0, 320, 256, 3, 4 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q5_0, 512, 512, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q5_0, 512, 512, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q5_0, 576, 512, 2, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q5_0, 576, 512, 2, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q5_0, 576, 512, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q5_0, 576, 512, 2, 3 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q5_0, 576, 512, 2, 4 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q5_1, 32, 32, -1, 1 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q5_1, 32, 32, 1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q5_1, 32, 32, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q5_1, 32, 32, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q5_1, 64, 64, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q5_1, 64, 64, -1, 1 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q5_1, 64, 64, 1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q5_1, 64, 64, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q5_1, 64, 64, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q5_1, 96, 96, -1, 1 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q5_1, 96, 96, 1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q5_1, 96, 96, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q5_1, 96, 96, 2, 4 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q5_1, 96, 96, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q5_1, 96, 96, 3, 4 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q5_1, 128, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q5_1, 128, 128, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q5_1, 128, 128, 2, 2 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q5_1, 128, 128, 2, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q5_1, 128, 128, 3, 2 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q5_1, 128, 128, 3, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q5_1, 192, 192, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q5_1, 192, 192, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q5_1, 192, 192, 3, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q5_1, 192, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q5_1, 192, 128, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q5_1, 256, 256, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q5_1, 256, 256, -1, 1 }, { 2, 2 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q5_1, 256, 256, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q5_1, 256, 256, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q5_1, 256, 256, 3, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q5_1, 320, 256, 1, 1 }, { 2, 2 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q5_1, 320, 256, 1, 3 }, { 2, 2 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q5_1, 320, 256, 1, 4 }, { 2, 2 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q5_1, 320, 256, 3, 1 }, { 2, 2 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q5_1, 512, 512, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q5_1, 512, 512, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q5_1, 576, 512, 2, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q5_1, 576, 512, 2, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q5_1, 576, 512, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q5_1, 576, 512, 2, 3 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q5_1, 576, 512, 2, 4 }, { 1, 4 } }, { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q8_0, 32, 32, -1, 1 }, { 4, 4 } }, { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q8_0, 32, 32, 1, 1 }, { 2, 4 } }, { { GGML_METAL_DEVICE_M4_PRO, GGML_TYPE_Q8_0, 32, 32, 1, 4 }, { 2, 4 } }, From 285f1ffd99d6b1ce97a59fc5978574140c0f7846 Mon Sep 17 00:00:00 2001 From: Georgi Gerganov Date: Sat, 29 Aug 2026 17:55:15 +0300 Subject: [PATCH 030/104] metal : assert shared memory padding (llama/27951) * metal : assert shared memory padding * cont : add ref --- ggml/src/ggml-metal/ggml-metal-device.cpp | 3 ++- ggml/src/ggml-metal/ggml-metal-device.m | 3 +++ ggml/src/ggml-metal/ggml-metal-ops.cpp | 2 +- 3 files changed, 6 insertions(+), 2 deletions(-) diff --git a/ggml/src/ggml-metal/ggml-metal-device.cpp b/ggml/src/ggml-metal/ggml-metal-device.cpp index a82caa5e430..4e855be4467 100644 --- a/ggml/src/ggml-metal/ggml-metal-device.cpp +++ b/ggml/src/ggml-metal/ggml-metal-device.cpp @@ -593,7 +593,7 @@ ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_ssm_scan(ggml_me // - sgptg floats for shared_x_dt (nsg) // - sgptg floats for shared_dA (nsg) // Total: nsg * (32 + 2) floats - res.smem = (32 + 2)*sizeof(float)*nsg; + res.smem = GGML_PAD((32 + 2)*sizeof(float)*nsg, 16); return res; } @@ -1029,6 +1029,7 @@ ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_mul_mm_id_map0(g } res.smem = (size_t) ne02*ne20*sizeof(uint16_t); + res.smem = GGML_PAD(res.smem, 16); return res; } diff --git a/ggml/src/ggml-metal/ggml-metal-device.m b/ggml/src/ggml-metal/ggml-metal-device.m index 41ce90dc8a9..85c0f576b9f 100644 --- a/ggml/src/ggml-metal/ggml-metal-device.m +++ b/ggml/src/ggml-metal/ggml-metal-device.m @@ -800,6 +800,9 @@ void ggml_metal_encoder_set_buffer(ggml_metal_encoder_t encoder, struct ggml_met } void ggml_metal_encoder_set_threadgroup_memory_size(ggml_metal_encoder_t encoder, size_t size, int idx) { + // ref: https://developer.apple.com/documentation/metal/mtlcomputecommandencoder/setthreadgroupmemorylength(_:index:) + GGML_ASSERT(size % 16 == 0); + [encoder->obj setThreadgroupMemoryLength:size atIndex:idx]; } diff --git a/ggml/src/ggml-metal/ggml-metal-ops.cpp b/ggml/src/ggml-metal/ggml-metal-ops.cpp index f6f2fdc86c6..89c8483b371 100644 --- a/ggml/src/ggml-metal/ggml-metal-ops.cpp +++ b/ggml/src/ggml-metal/ggml-metal-ops.cpp @@ -948,7 +948,7 @@ int ggml_metal_op_sum(ggml_metal_op_t ctx, int idx) { 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), 2); - ggml_metal_encoder_set_threadgroup_memory_size(enc, nsg * sizeof(float), 0); + ggml_metal_encoder_set_threadgroup_memory_size(enc, GGML_PAD(nsg * sizeof(float), 16), 0); ggml_metal_encoder_dispatch_threadgroups(enc, 1, 1, 1, nth, 1, 1); From 2a11026cfe3a660673e8c8c37dac18541ea3b1f6 Mon Sep 17 00:00:00 2001 From: Hongqiang Wang Date: Sat, 29 Aug 2026 10:46:27 -0700 Subject: [PATCH 031/104] opencl: use a better matmul path on two Adreno GPU generations (llama/27640) * opencl: default the Adreno xmem F16xF32 GEMM on for X2E kernel_mul_mm_f16_f32_l4_lm is the slowest matmul this backend has on Adreno: on the X2-90 it runs the gpt-oss-20b attention projections at roughly a quarter of what the tuned dense q4_0 GEMM reaches on the same device. That matters for any model whose non-expert weights stay f16 -- the stock gpt-oss-20b release is exactly that, and its prefill spends 40.8% of GPU time in that one kernel. The xmem route already existed but was left opt-in, so nobody hit it. Worth about 25% prefill on gpt-oss-20b on an Adreno X2-90. Gated to X2E: the Adreno 840 measures neutral. Decode is untouched -- the dispatch gate needs N >= 16. It is worth nothing on the q8attn variant, whose attention weights already take the dp4a dense GEMM. The env var was presence-tested before, so =0 previously enabled it; it is now atoi()'d. MUL_MAT 963 OK / 0 FAIL on both arms. * opencl: bypass the tiled f32 GEMM on the Adreno A7X The A7X (E031.41) compiler executes kernel_mul_mm_f32_f32_l4_lm at roughly a tenth of what the same silicon reaches in its own f16 and q4_K kernels. It allocates 488 B/WI of private memory against 304 for the same source on the following generation, i.e. the older register allocator spills in the K-loop. Models with per-layer F32 projection pairs kept F32 by quantization policy land on this kernel twice per layer, and it dominates their prefill on that part. Route batched f32xf32 (ne11 > 8) around the tiled path on the A7X and let it fall through to the per-row f32 kernel, which that compiler handles fine; small batches keep the tiled path. Weights stay GPU-resident, so decode placement is untouched -- declining the op in supports_op instead was measured first and rejected, because the per-layer CPU round-trips cost more decode than the prefill it gained. Worth about 9% prefill on gemma-3n-E4B on an Adreno 740, with MUL_MAT counts identical on and off. No other generation is affected. Override with GGML_OPENCL_A7X_F32_LM_BYPASS=0. * opencl: enable xmem GEMM for adreno by default --------- Co-authored-by: Li He --- ggml/src/ggml-opencl/ggml-opencl.cpp | 21 +++++++++++++++++---- 1 file changed, 17 insertions(+), 4 deletions(-) diff --git a/ggml/src/ggml-opencl/ggml-opencl.cpp b/ggml/src/ggml-opencl/ggml-opencl.cpp index 6ae83449b08..426aac52316 100644 --- a/ggml/src/ggml-opencl/ggml-opencl.cpp +++ b/ggml/src/ggml-opencl/ggml-opencl.cpp @@ -6059,9 +6059,13 @@ static ggml_backend_opencl_context * ggml_cl_init(ggml_backend_dev_t dev) { } #ifdef GGML_OPENCL_USE_ADRENO_KERNELS - // determine whether to use Adreno xmem GEMM - backend_ctx->adreno_xmem_gemm_enabled = getenv("GGML_OPENCL_ADRENO_XMEM_GEMM") != nullptr && - backend_ctx->gpu_family == GPU_FAMILY::ADRENO; + // Adreno xmem F16xF32 GEMM, default on adreno, opt out with GGML_OPENCL_ADRENO_XMEM_GEMM=0. + // This helps models with f16 attention weights, e.g., gpt-oss-20b-f16 + { + const char * xmem_env = getenv("GGML_OPENCL_ADRENO_XMEM_GEMM"); + backend_ctx->adreno_xmem_gemm_enabled = backend_ctx->gpu_family == GPU_FAMILY::ADRENO && + (xmem_env ? atoi(xmem_env) != 0 : true); + } #endif // determine whether to use large buffer for Adreno @@ -19534,9 +19538,18 @@ static void ggml_cl_mul_mat(ggml_backend_t backend, const ggml_tensor * src0, co // GEMM using local memory // Current BK = 16, so ne00 % 16 == 0 + // + // Certain A7X compiler (E031.41) executes kernel_mul_mm_f32_f32_l4_lm poorly; + // matrices with ne11 <= 8 appears OK. + // Fallback to the MV style kernels for A7x and ne11 > 8. + // Override with GGML_OPENCL_A7X_F32_LM_BYPASS=0. + static const char * a7x_f32lm_env = getenv("GGML_OPENCL_A7X_F32_LM_BYPASS"); + static const bool a7x_f32lm_bypass = (a7x_f32lm_env == nullptr || a7x_f32lm_env[0] != '0'); if (src1t == GGML_TYPE_F32 && ne00 % 16 == 0 && - ne11 > 1) { + ne11 > 1 && + !(a7x_f32lm_bypass && src0t == GGML_TYPE_F32 && ne11 > 8 && + backend_ctx->adreno_gen == ADRENO_GPU_GEN::A7X)) { switch(src0t) { case GGML_TYPE_F32: { kernel = backend_ctx->kernel_mul_mm_f32_f32_l4_lm; From b33bbc5c4eef7a435de6f907d6ff4547274684ed Mon Sep 17 00:00:00 2001 From: codemonkey <441345965@qq.com> Date: Sun, 30 Aug 2026 07:44:53 +0800 Subject: [PATCH 032/104] metal : add fa-vec tunings for M2 (llama/27940) --- ggml/src/ggml-metal/ggml-metal-tuning.cpp | 82 +++++++++++++++++++++++ 1 file changed, 82 insertions(+) diff --git a/ggml/src/ggml-metal/ggml-metal-tuning.cpp b/ggml/src/ggml-metal/ggml-metal-tuning.cpp index 26b634d2143..4abdafb48f3 100644 --- a/ggml/src/ggml-metal/ggml-metal-tuning.cpp +++ b/ggml/src/ggml-metal/ggml-metal-tuning.cpp @@ -467,6 +467,88 @@ constexpr fa_vec_entry_t fa_vec_tuned_table[] = { { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q8_0, 576, 512, -1, 0 }, { 1, 4 } }, { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q8_0, 576, 512, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_F16, 32, 32, 1, 3 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_F16, 64, 64, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_F16, 64, 64, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_F16, 128, 128, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_F16, 128, 128, 1, 1 }, { 1, 1 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_F16, 128, 128, 1, 2 }, { 1, 1 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_F16, 128, 128, 1, 3 }, { 1, 1 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_F16, 128, 128, 1, 4 }, { 1, 1 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_F16, 192, 128, 1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_F16, 192, 128, 1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_F16, 192, 128, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_F16, 192, 128, 1, 3 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_F16, 192, 128, 1, 4 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_F16, 320, 256, 2, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_F16, 320, 256, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_F16, 320, 256, 1, 1 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_F16, 320, 256, 1, 3 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_F16, 320, 256, 1, 4 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q4_0, 32, 32, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q4_0, 32, 32, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q4_0, 32, 32, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q4_0, 32, 32, 2, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q4_0, 32, 32, 3, 2 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q4_0, 32, 32, 3, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q4_0, 64, 64, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q4_0, 64, 64, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q4_0, 64, 64, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q4_0, 96, 96, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q4_0, 96, 96, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q4_0, 96, 96, 1, 4 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q4_0, 96, 96, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q4_0, 96, 96, 2, 4 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q4_0, 96, 96, 3, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q4_0, 128, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q4_0, 128, 128, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q4_0, 192, 192, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q4_0, 192, 192, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q4_0, 192, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q4_0, 192, 128, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q4_0, 256, 256, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q4_0, 256, 256, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q4_0, 320, 256, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q4_0, 320, 256, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q4_0, 512, 512, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q4_0, 512, 512, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q4_0, 512, 512, 2, 3 }, { 1, 1 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q4_0, 512, 512, 3, 3 }, { 2, 2 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q4_0, 576, 512, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q4_0, 576, 512, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q8_0, 32, 32, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q8_0, 32, 32, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q8_0, 32, 32, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q8_0, 32, 32, 2, 4 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q8_0, 32, 32, 3, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q8_0, 32, 32, 3, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q8_0, 64, 64, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q8_0, 64, 64, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q8_0, 64, 64, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q8_0, 64, 64, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q8_0, 64, 64, 3, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q8_0, 96, 96, 1, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q8_0, 96, 96, 2, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q8_0, 96, 96, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q8_0, 96, 96, 3, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q8_0, 128, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q8_0, 128, 128, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q8_0, 192, 192, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q8_0, 192, 192, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q8_0, 192, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q8_0, 192, 128, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q8_0, 256, 256, -1, 0 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q8_0, 256, 256, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q8_0, 320, 256, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q8_0, 320, 256, 1, 2 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q8_0, 320, 256, 1, 4 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q8_0, 320, 256, 2, 2 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q8_0, 320, 256, 3, 2 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q8_0, 512, 512, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q8_0, 512, 512, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q8_0, 576, 512, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q8_0, 576, 512, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_ULTRA, GGML_TYPE_F16, 64, 64, -1, 0 }, { 1, 4 } }, { { GGML_METAL_DEVICE_M2_ULTRA, GGML_TYPE_F16, 64, 64, -1, 1 }, { 1, 4 } }, { { GGML_METAL_DEVICE_M2_ULTRA, GGML_TYPE_F16, 128, 128, 2, 1 }, { 1, 4 } }, From c969c68b544cd2e90052b58647b52ce67bd93575 Mon Sep 17 00:00:00 2001 From: Aman Gupta Date: Sun, 30 Aug 2026 09:04:20 +0530 Subject: [PATCH 033/104] ggml: allow passing alloc dependencies in graph_optimize (llama/27301) * ggml: allow passing alloc dependencies in graph_optimize * add alloc dep tests * add TODO about using flat array --- ggml/src/ggml-backend-impl.h | 12 ++++- ggml/src/ggml-backend.cpp | 66 ++++++++++++++++++++++---- ggml/src/ggml-cuda/ggml-cuda.cu | 4 +- ggml/src/ggml-hexagon/ggml-hexagon.cpp | 4 +- ggml/src/ggml-metal/ggml-metal.cpp | 4 +- ggml/src/ggml-virtgpu/ggml-backend.cpp | 3 +- ggml/src/ggml-vulkan/ggml-vulkan.cpp | 3 +- 7 files changed, 82 insertions(+), 14 deletions(-) diff --git a/ggml/src/ggml-backend-impl.h b/ggml/src/ggml-backend-impl.h index 40cea024c3d..56f0090cce6 100644 --- a/ggml/src/ggml-backend-impl.h +++ b/ggml/src/ggml-backend-impl.h @@ -103,6 +103,16 @@ extern "C" { // Backend (stream) // + // passed to graph_optimize so the backend can add allocation dependencies: + // if the backend executes parts of the graph out of order (e.g. on concurrent streams), + // it must keep the affected tensors allocated until a node where execution is known to have joined + struct ggml_backend_graph_optimize_params { + // keep `tensor` allocated at least until `until` (a node of the same graph) has been computed + // can be called multiple times for the same tensor: the longest lifetime applies + void (*add_alloc_dep)(void * user_data, struct ggml_tensor * tensor, struct ggml_tensor * until); + void * user_data; + }; + struct ggml_backend_i { const char * (*get_name)(ggml_backend_t backend); @@ -137,7 +147,7 @@ extern "C" { void (*event_wait) (ggml_backend_t backend, ggml_backend_event_t event); // (optional) sort/optimize the nodes in the graph - void (*graph_optimize) (ggml_backend_t backend, struct ggml_cgraph * cgraph); + void (*graph_optimize) (ggml_backend_t backend, struct ggml_cgraph * cgraph, struct ggml_backend_graph_optimize_params * params); }; struct ggml_backend { diff --git a/ggml/src/ggml-backend.cpp b/ggml/src/ggml-backend.cpp index e519bdf50a1..78eb10dfe99 100644 --- a/ggml/src/ggml-backend.cpp +++ b/ggml/src/ggml-backend.cpp @@ -20,6 +20,7 @@ #include #include #include +#include #include #ifdef __APPLE__ @@ -558,10 +559,10 @@ void ggml_backend_event_wait(ggml_backend_t backend, ggml_backend_event_t event) backend->iface.event_wait(backend, event); } -static void ggml_backend_graph_optimize(ggml_backend_t backend, struct ggml_cgraph * cgraph) { +static void ggml_backend_graph_optimize(ggml_backend_t backend, struct ggml_cgraph * cgraph, struct ggml_backend_graph_optimize_params * params) { GGML_ASSERT(backend); if (backend->iface.graph_optimize != NULL) { - backend->iface.graph_optimize(backend, cgraph); + backend->iface.graph_optimize(backend, cgraph, params); } } @@ -1441,11 +1442,40 @@ void ggml_backend_sched_split_graph(ggml_backend_sched_t sched, struct ggml_cgra sched->prev_leaf_backend_ids = tmp; } + // optimize the split graphs and collect the allocation dependencies added by the backends + // this needs to happen before we make graph_copy, so they are in sync + // TODO: this may create many small allocations in the scheduler, restructure to use a flat array + std::unordered_map> alloc_deps; + + struct ggml_backend_graph_optimize_params opt_params = { + /* .add_alloc_dep = */ [](void * user_data, ggml_tensor * tensor, ggml_tensor * until) { + auto & deps = *(std::unordered_map> *) user_data; + std::vector & keep = deps[until]; + if (std::find(keep.begin(), keep.end(), tensor) == keep.end()) { + keep.push_back(tensor); + } + }, + /* .user_data = */ &alloc_deps, + }; + + for (int i = 0; i < sched->n_splits; i++) { + struct ggml_backend_sched_split * split = &sched->splits[i]; + split->graph = ggml_graph_view(graph, split->i_start, split->i_end); + + ggml_backend_graph_optimize(sched->backends[split->backend_id], &split->graph, &opt_params); + } + + // each dep is added to graph_copy as a GGML_OP_NONE node with the kept tensors as srcs + int n_dep_nodes = 0; + for (const auto & it : alloc_deps) { + n_dep_nodes += (it.second.size() + GGML_MAX_SRC - 1) / GGML_MAX_SRC; + } + int total_inputs = sched->n_graph_inputs; for (int i = 0; i < sched->n_splits; i++) { total_inputs += sched->splits[i].n_inputs; } - int graph_size = std::max(graph->n_nodes, graph->n_leafs) + total_inputs * 2 * sched->n_copies; + int graph_size = std::max(graph->n_nodes, graph->n_leafs) + total_inputs * 2 * sched->n_copies + n_dep_nodes; // remember the actual graph_size for performing reallocation checks later [GGML_SCHED_DEBUG_REALLOC] sched->debug_prev_graph_size = sched->debug_graph_size; @@ -1463,13 +1493,10 @@ void ggml_backend_sched_split_graph(ggml_backend_sched_t sched, struct ggml_cgra struct ggml_cgraph * graph_copy = &sched->graph; + int n_dep_nodes_added = 0; + for (int i = 0; i < sched->n_splits; i++) { struct ggml_backend_sched_split * split = &sched->splits[i]; - split->graph = ggml_graph_view(graph, split->i_start, split->i_end); - - // Optimize this split of the graph. This needs to happen before we make graph_copy, - // so they are in sync. - ggml_backend_graph_optimize(sched->backends[split->backend_id], &split->graph); // add inputs to the graph copy so that they are allocated by ggml-alloc at the start of the split for (int j = 0; j < split->n_inputs; j++) { @@ -1494,9 +1521,32 @@ void ggml_backend_sched_split_graph(ggml_backend_sched_t sched, struct ggml_cgra assert(graph_copy->size > graph_copy->n_nodes); sched->node_backend_ids[graph_copy->n_nodes] = tensor_backend_id(graph->nodes[j]); graph_copy->nodes[graph_copy->n_nodes++] = graph->nodes[j]; + + if (alloc_deps.empty()) { + continue; + } + + // add a dependency node so that the kept tensors are not freed before this node is computed + auto it = alloc_deps.find(graph->nodes[j]); + if (it != alloc_deps.end()) { + const std::vector & keep = it->second; + for (size_t k = 0; k < keep.size(); k += GGML_MAX_SRC) { + struct ggml_tensor * dep = ggml_view_tensor(sched->ctx, keep[k]); + for (size_t s = 0; s < GGML_MAX_SRC && k + s < keep.size(); s++) { + dep->src[s] = keep[k + s]; + } + assert(graph_copy->size > graph_copy->n_nodes); + sched->node_backend_ids[graph_copy->n_nodes] = split->backend_id; + graph_copy->nodes[graph_copy->n_nodes++] = dep; + n_dep_nodes_added++; + } + } } } + // a mismatch means a backend added a dep with an `until` tensor that is not a node of the optimized graph + GGML_ASSERT(n_dep_nodes_added == n_dep_nodes); + if (sched->n_copies > 1) { // add input copies as leafs so that they are allocated first for (int i = 0; i < sched->n_graph_inputs; i++) { diff --git a/ggml/src/ggml-cuda/ggml-cuda.cu b/ggml/src/ggml-cuda/ggml-cuda.cu index 2456f7dcc62..bd9754c2ffd 100644 --- a/ggml/src/ggml-cuda/ggml-cuda.cu +++ b/ggml/src/ggml-cuda/ggml-cuda.cu @@ -4328,7 +4328,9 @@ static void ggml_backend_cuda_event_wait(ggml_backend_t backend, ggml_backend_ev } } -static void ggml_backend_cuda_graph_optimize(ggml_backend_t backend, ggml_cgraph * cgraph) { +static void ggml_backend_cuda_graph_optimize(ggml_backend_t backend, ggml_cgraph * cgraph, ggml_backend_graph_optimize_params * params) { + GGML_UNUSED(params); + ggml_backend_cuda_context * cuda_ctx = (ggml_backend_cuda_context *) backend->context; #ifdef USE_CUDA_GRAPH diff --git a/ggml/src/ggml-hexagon/ggml-hexagon.cpp b/ggml/src/ggml-hexagon/ggml-hexagon.cpp index 53e86075591..e7dcdc3d551 100644 --- a/ggml/src/ggml-hexagon/ggml-hexagon.cpp +++ b/ggml/src/ggml-hexagon/ggml-hexagon.cpp @@ -4984,7 +4984,9 @@ static std::vector ggml_hexagon_graph_optimize_reorder(const std::vectorn_nodes; constexpr int MAX_FUSE = 16; diff --git a/ggml/src/ggml-metal/ggml-metal.cpp b/ggml/src/ggml-metal/ggml-metal.cpp index 9756d47050c..4d58dc821cf 100644 --- a/ggml/src/ggml-metal/ggml-metal.cpp +++ b/ggml/src/ggml-metal/ggml-metal.cpp @@ -558,7 +558,9 @@ static void ggml_backend_metal_event_wait(ggml_backend_t backend, ggml_backend_e ggml_metal_event_wait(ctx, ev); } -static void ggml_backend_metal_graph_optimize(ggml_backend_t backend, ggml_cgraph * cgraph) { +static void ggml_backend_metal_graph_optimize(ggml_backend_t backend, ggml_cgraph * cgraph, ggml_backend_graph_optimize_params * params) { + GGML_UNUSED(params); + ggml_metal_t ctx = (ggml_metal_t)backend->context; ggml_metal_graph_optimize(ctx, cgraph); diff --git a/ggml/src/ggml-virtgpu/ggml-backend.cpp b/ggml/src/ggml-virtgpu/ggml-backend.cpp index 12756c9282f..996c57e358b 100644 --- a/ggml/src/ggml-virtgpu/ggml-backend.cpp +++ b/ggml/src/ggml-virtgpu/ggml-backend.cpp @@ -17,7 +17,8 @@ static ggml_status ggml_backend_remoting_graph_compute(ggml_backend_t backend, g return apir_backend_graph_compute(gpu, cgraph); } -static void ggml_backend_remoting_graph_optimize(ggml_backend_t backend, ggml_cgraph * cgraph) { +static void ggml_backend_remoting_graph_optimize(ggml_backend_t backend, ggml_cgraph * cgraph, ggml_backend_graph_optimize_params * params) { + UNUSED(params); virtgpu * gpu = DEV_TO_GPU(backend->device); #if true UNUSED(gpu); diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp index 39b4cd35980..8fbb1359f40 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp +++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp @@ -17795,8 +17795,9 @@ static ggml_status ggml_backend_vk_graph_compute(ggml_backend_t backend, ggml_cg } // Sort the graph for improved parallelism. -static void ggml_vk_graph_optimize(ggml_backend_t backend, struct ggml_cgraph * graph) +static void ggml_vk_graph_optimize(ggml_backend_t backend, struct ggml_cgraph * graph, struct ggml_backend_graph_optimize_params * params) { + GGML_UNUSED(params); VK_LOG_DEBUG("ggml_vk_graph_optimize(" << graph->n_nodes << " nodes)"); ggml_backend_vk_context * ctx = (ggml_backend_vk_context *)backend->context; From c68f20555aa9ff91dbec5d6c23e8a862174e4d8f Mon Sep 17 00:00:00 2001 From: QuintinShaw Date: Sun, 30 Aug 2026 13:56:35 +0800 Subject: [PATCH 034/104] metal : fix null-pipeline crash for F16 src1 mul_mat/mul_mat_id (llama/25648) * metal : fail closed on mul_mat shapes with missing F16 kernels * metal : abort on nil pipeline in encoder_set_pipeline * metal : address review comments * metal : share mul_mat mm dispatch with supports_op --- ggml/src/ggml-metal/ggml-metal-common.cpp | 17 +++++++++++ ggml/src/ggml-metal/ggml-metal-common.h | 4 +++ ggml/src/ggml-metal/ggml-metal-device.m | 37 ++++++++++++++++++++++- ggml/src/ggml-metal/ggml-metal-ops.cpp | 19 ++---------- 4 files changed, 59 insertions(+), 18 deletions(-) diff --git a/ggml/src/ggml-metal/ggml-metal-common.cpp b/ggml/src/ggml-metal/ggml-metal-common.cpp index 2eb9820bff9..6f1638a1147 100644 --- a/ggml/src/ggml-metal/ggml-metal-common.cpp +++ b/ggml/src/ggml-metal/ggml-metal-common.cpp @@ -1,10 +1,27 @@ #include "ggml-metal-common.h" +#include "ggml.h" #include "ggml-impl.h" #include "ggml-backend-impl.h" #include +bool ggml_metal_op_mul_mat_use_mm(const struct ggml_tensor * op, bool has_simdgroup_mm) { + const int64_t ne00 = op->src[0]->ne[0]; + const int64_t ne11 = op->src[1]->ne[1]; + + return !ggml_is_transposed(op->src[0]) && + !ggml_is_transposed(op->src[1]) && + has_simdgroup_mm && ne00 >= 64 && ne11 > 8; +} + +bool ggml_metal_op_mul_mat_id_use_mm(const struct ggml_tensor * op, bool has_simdgroup_mm) { + const int64_t ne00 = op->src[0]->ne[0]; + const int64_t ne21 = op->src[2]->ne[1]; + + return has_simdgroup_mm && ne00 >= 64 && ne21 >= 32; +} + // 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 3acbc6ae174..66abdb52efe 100644 --- a/ggml/src/ggml-metal/ggml-metal-common.h +++ b/ggml/src/ggml-metal/ggml-metal-common.h @@ -47,6 +47,10 @@ bool ggml_mem_ranges_check(ggml_mem_ranges_t mrs, const struct ggml_tensor * ten // 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_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); + #ifdef __cplusplus } #endif diff --git a/ggml/src/ggml-metal/ggml-metal-device.m b/ggml/src/ggml-metal/ggml-metal-device.m index 85c0f576b9f..a053887a337 100644 --- a/ggml/src/ggml-metal/ggml-metal-device.m +++ b/ggml/src/ggml-metal/ggml-metal-device.m @@ -3,6 +3,7 @@ #import "ggml-impl.h" #import "ggml-backend-impl.h" #import "ggml-metal-impl.h" +#import "ggml-metal-common.h" #include @@ -788,6 +789,10 @@ void ggml_metal_encoder_debug_group_pop (ggml_metal_encoder_t encoder) { } void ggml_metal_encoder_set_pipeline(ggml_metal_encoder_t encoder, struct ggml_metal_pipeline_with_params pipeline) { + if (!pipeline.pipeline) { + GGML_ABORT("%s: nil Metal pipeline (missing kernel; see compile_pipeline log above)\n", __func__); + } + [encoder->obj setComputePipelineState:pipeline.pipeline->obj]; } @@ -1410,6 +1415,30 @@ void ggml_metal_device_get_memory(ggml_metal_device_t dev, size_t * free, size_t } } +static bool ggml_metal_supports_mul_mat_op( + bool has_simdgroup_reduction, + const struct ggml_tensor * op, + bool src0_f16_has_mv, + bool mm_path) { + if (!has_simdgroup_reduction || op->src[0]->type == GGML_TYPE_NVFP4) { + return false; + } + + if (op->src[1]->type != GGML_TYPE_F16) { + return true; + } + + if (op->src[0]->type == GGML_TYPE_BF16) { + return false; + } + + if (src0_f16_has_mv && op->src[0]->type == GGML_TYPE_F16) { + return true; + } + + return mm_path; +} + bool ggml_metal_device_supports_op(ggml_metal_device_t dev, const struct ggml_tensor * op) { const bool has_simdgroup_mm = dev->props.has_simdgroup_mm; const bool has_simdgroup_reduction = dev->props.has_simdgroup_reduction; @@ -1713,9 +1742,15 @@ bool ggml_metal_device_supports_op(ggml_metal_device_t dev, const struct ggml_te case GGML_OP_GATED_DELTA_NET: return has_simdgroup_reduction && op->src[2]->ne[0] % 32 == 0; case GGML_OP_SOLVE_TRI: + return has_simdgroup_reduction && op->src[0]->type == GGML_TYPE_F32; case GGML_OP_MUL_MAT: + return ggml_metal_supports_mul_mat_op( + has_simdgroup_reduction, op, true, + ggml_metal_op_mul_mat_use_mm(op, has_simdgroup_mm)); case GGML_OP_MUL_MAT_ID: - return has_simdgroup_reduction && op->src[0]->type != GGML_TYPE_NVFP4; + return ggml_metal_supports_mul_mat_op( + has_simdgroup_reduction, op, false, + ggml_metal_op_mul_mat_id_use_mm(op, has_simdgroup_mm)); case GGML_OP_SET: case GGML_OP_CPY: case GGML_OP_DUP: diff --git a/ggml/src/ggml-metal/ggml-metal-ops.cpp b/ggml/src/ggml-metal/ggml-metal-ops.cpp index 89c8483b371..7671d1d0156 100644 --- a/ggml/src/ggml-metal/ggml-metal-ops.cpp +++ b/ggml/src/ggml-metal/ggml-metal-ops.cpp @@ -2362,10 +2362,6 @@ int ggml_metal_op_mul_mat(ggml_metal_op_t ctx, int idx) { const int16_t r2 = ne12/ne02; const int16_t r3 = ne13/ne03; - // find the break-even point where the matrix-matrix kernel becomes more efficient compared - // to the matrix-vector kernel - const int ne11_mm_min = 8; - // first try to use small-batch mat-mv kernels // these should be efficient for BS [2, ~8] if (op->src[1]->type == GGML_TYPE_F32 && (ne00%128 == 0) && @@ -2468,12 +2464,7 @@ int ggml_metal_op_mul_mat(ggml_metal_op_t ctx, int idx) { ggml_metal_encoder_set_buffer (enc, ggml_metal_get_buffer_id(op), 3); ggml_metal_encoder_dispatch_threadgroups(enc, ((ne01 + r0ptg - 1)/r0ptg), ((ne11 + r1ptg - 1)/r1ptg), ne12*ne13, 32, nsg, 1); - } else if ( - !ggml_is_transposed(op->src[0]) && - !ggml_is_transposed(op->src[1]) && - // for now the matrix-matrix multiplication kernel only works on A14+/M1+ SoCs - // AMD GPU and older A-chips will reuse matrix-vector multiplication kernel - props_dev->has_simdgroup_mm && ne00 >= 64 && ne11 > ne11_mm_min) { + } else if (ggml_metal_op_mul_mat_use_mm(op, props_dev->has_simdgroup_mm)) { //GGML_LOG_INFO("matrix: ne00 = %6d, ne01 = %6d, ne02 = %6d, ne11 = %6d, ne12 = %6d\n", ne00, ne01, ne02, ne11, ne12); // some Metal matrix data types require aligned pointers @@ -2622,13 +2613,7 @@ int ggml_metal_op_mul_mat_id(ggml_metal_op_t ctx, int idx) { const uint32_t r2 = 1; const uint32_t r3 = 1; - // find the break-even point where the matrix-matrix kernel becomes more efficient compared - // to the matrix-vector kernel - // ne20 = n_used_experts - // ne21 = n_rows (batch size) - const int ne21_mm_id_min = 32; - - if (props_dev->has_simdgroup_mm && ne00 >= 64 && (ne21 >= ne21_mm_id_min)) { + if (ggml_metal_op_mul_mat_id_use_mm(op, props_dev->has_simdgroup_mm)) { // some Metal matrix data types require aligned pointers // ref: https://developer.apple.com/metal/Metal-Shading-Language-Specification.pdf (Table 2.5) //switch (op->src[0]->type) { From 3d4e0e9858ffc1d4446d8b4ce7d4385d696e3cb6 Mon Sep 17 00:00:00 2001 From: Titaniumtown Date: Sat, 29 Aug 2026 22:57:08 -0700 Subject: [PATCH 035/104] sycl: split long rows in TOP_K instead of one work-group per row (llama/27847) --- ggml/src/ggml-sycl/ggml-sycl.cpp | 294 +++++++++++++++++++++++-------- 1 file changed, 217 insertions(+), 77 deletions(-) diff --git a/ggml/src/ggml-sycl/ggml-sycl.cpp b/ggml/src/ggml-sycl/ggml-sycl.cpp index dc8a1744323..d58ffd00daf 100644 --- a/ggml/src/ggml-sycl/ggml-sycl.cpp +++ b/ggml/src/ggml-sycl/ggml-sycl.cpp @@ -2402,7 +2402,138 @@ static void argsort_f32_i32_sycl(const float *x, int *dst, const int ncols, } } +// Scan and block merge, shared by every launch shape below so a partitioned row uses the +// same insertion order as an unpartitioned one. +// +// src_map != nullptr: report src_map[col] instead of col, so a merge pass can carry the +// original column index through. +// out_vals != nullptr: also emit the k winning values, for a later merge pass. +// swap01: emit in the output order the single-pass path uses. +static void top_k_scan_merge_f32( + const float * src_vals, + const int32_t * src_map, + const int begin, + const int end, + const int k, + const int block_size, + float * shared_vals, + int * shared_idx, + float * out_vals, + int32_t * out_idx, + const bool swap01, + const sycl::nd_item<1> & item_ct1 +) { + const int tid = item_ct1.get_local_id(0); + + // The running top-k lives in SLM (shared local memory) rather than a private array: + // an array indexed by a runtime position cannot be register-allocated, so a private + // one lands in scratch, i.e. device memory, and insertion is this kernel's dominant + // cost. + // + // Lane-strided (lv[i * block_size]) rather than lane-blocked (lv[i]) so a given i is + // contiguous across lanes; a k-strided layout would put every lane of a shift step in + // the same SLM bank. + float * lv = shared_vals + tid; + int * li = shared_idx + tid; + + for (int i = 0; i < k; i++) { + lv[i * block_size] = -FLT_MAX; + li[i * block_size] = -1; + } + + // The k-th best, cached in a register. The reject test is taken for the large + // majority of elements scanned, and in that case touches no memory. + float kth = -FLT_MAX; + + for (int col = begin + tid; col < end; col += block_size) { + float val = src_vals[col]; + + if (val > kth) { + int pos = k - 1; + while (pos > 0 && val > lv[(pos - 1) * block_size]) { + pos--; + } + + for (int i = k - 1; i > pos; i--) { + lv[i * block_size] = lv[(i - 1) * block_size]; + li[i * block_size] = li[(i - 1) * block_size]; + } + lv[pos * block_size] = val; + li[pos * block_size] = src_map ? src_map[col] : col; + + kth = lv[(k - 1) * block_size]; + } + } + + item_ct1.barrier(sycl::access::fence_space::local_space); + + if (tid != 0) { + return; + } + + // Same treatment for the merge accumulator, past the per-lane region. + float * fv = shared_vals + (size_t) k * block_size; + int * fi = shared_idx + (size_t) k * block_size; + + for (int i = 0; i < k; i++) { + fv[i] = -FLT_MAX; + fi[i] = -1; + } + + float fkth = -FLT_MAX; + + // Candidates are visited in the same (t, i) order as before, so tie-breaking is + // unchanged. + for (int t = 0; t < block_size; t++) { + for (int i = 0; i < k; i++) { + float val = shared_vals[i * block_size + t]; + + if (val <= fkth) { + // Lane t's list is sorted descending, so once one of its entries loses + // to the k-th best, every later entry loses too. fkth only rises, so + // that stays true for the rest of the merge. This turns the merge from + // block_size*k steps into roughly block_size plus the candidates + // accepted. + break; + } + + int idx = shared_idx[i * block_size + t]; + + int pos = k - 1; + while (pos > 0 && val > fv[pos - 1]) { + pos--; + } + + for (int j = k - 1; j > pos; j--) { + fv[j] = fv[j - 1]; + fi[j] = fi[j - 1]; + } + fv[pos] = val; + fi[pos] = idx; + + fkth = fv[k - 1]; + } + } + + if (out_vals) { + for (int i = 0; i < k; i++) { + out_vals[i] = fv[i]; + } + } + + for (int i = 0; i < k; i++) { + out_idx[i] = fi[i]; + } + + if (swap01 && k > 1) { + int32_t temp = out_idx[0]; + out_idx[0] = out_idx[1]; + out_idx[1] = temp; + } +} + static void top_k_f32_sycl( + ggml_backend_sycl_context & ctx, const float * src, int32_t * dst_indices, const int64_t ncols, @@ -2410,98 +2541,107 @@ static void top_k_f32_sycl( const int k, dpct::queue_ptr main_stream ) { - const int block_size = 128; + // A row is scanned by exactly one work-group, so a vocabulary-sized row leaves the + // rest of the device idle. What the scan is short of is memory requests in flight, + // not bandwidth or per-request latency, so lanes in flight is the lever: split the + // row across independent work-groups, have each emit its partition's top-k, and + // merge those nsplit*k candidates in a second launch. + // + // split_block trades parallelism against SLM residency. Its cost is + // (split_block + 1) * k * 8 bytes of SLM per group, so at the k <= 32 ceiling 128 + // lanes need about 33 KB, which leaves a single resident group per Xe-core. Revisit + // if the supported k ever grows. + constexpr int split_block = 128; + constexpr int max_splits = 128; + constexpr int min_cols = 8192; - const sycl::range<1> block_dims(block_size); - const sycl::range<1> grid_dims(nrows); + int nsplit = 1; + if (ncols >= min_cols) { + // A partition is then always >= split_block = 128 columns, hence always more than + // the k <= 32 ceiling, so no pass is ever padded with -FLT_MAX sentinels. + const int64_t want = ncols / split_block; + nsplit = (int) (want > max_splits ? max_splits : want); + } - main_stream->submit([&](sycl::handler &cgh) { - sycl::local_accessor shared_vals(sycl::range<1>(block_size * k), cgh); - sycl::local_accessor shared_idx(sycl::range<1>(block_size * k), cgh); + if (nsplit > 1) { + const int nchunk = (int) ((ncols + nsplit - 1) / nsplit); + const size_t ncand = (size_t) nrows * nsplit * k; - cgh.parallel_for( - sycl::nd_range<1>(grid_dims * block_dims, block_dims), - [=](sycl::nd_item<1> item_ct1) { - const int row = item_ct1.get_group(0); - const int tid = item_ct1.get_local_id(0); + ggml_sycl_pool_alloc part_vals(ctx.pool(), ncand); + ggml_sycl_pool_alloc part_idx(ctx.pool(), ncand); - if (row >= nrows) return; + float * pv = part_vals.get(); + int32_t * pi = part_idx.get(); - const float * src_row = src + row * ncols; - int32_t * dst_idx_row = dst_indices + row * k; + const sycl::range<1> block_dims(split_block); - float local_vals[32]; - int local_idx[32]; + main_stream->submit([&](sycl::handler &cgh) { + sycl::local_accessor shared_vals(sycl::range<1>((split_block + 1) * k), cgh); + sycl::local_accessor shared_idx(sycl::range<1>((split_block + 1) * k), cgh); - for (int i = 0; i < k; i++) { - local_vals[i] = -FLT_MAX; - local_idx[i] = -1; - } + cgh.parallel_for( + sycl::nd_range<1>(sycl::range<1>(nrows * nsplit) * block_dims, block_dims), + [=](sycl::nd_item<1> item_ct1) { + const int grp = item_ct1.get_group(0); + const int row = grp / nsplit; + const int part = grp % nsplit; + + const int begin = part * nchunk; + int end = begin + nchunk; + if (end > (int) ncols) { + end = (int) ncols; + } - for (int col = tid; col < ncols; col += block_size) { - float val = src_row[col]; + top_k_scan_merge_f32( + src + (int64_t) row * ncols, nullptr, begin, end, k, split_block, + shared_vals.get_multi_ptr().get(), + shared_idx.get_multi_ptr().get(), + pv + (size_t) grp * k, pi + (size_t) grp * k, false, item_ct1); + }); + }); - if (val > local_vals[k-1]) { - int pos = k - 1; - while (pos > 0 && val > local_vals[pos - 1]) { - pos--; - } + main_stream->submit([&](sycl::handler &cgh) { + sycl::local_accessor shared_vals(sycl::range<1>((split_block + 1) * k), cgh); + sycl::local_accessor shared_idx(sycl::range<1>((split_block + 1) * k), cgh); - for (int i = k - 1; i > pos; i--) { - local_vals[i] = local_vals[i - 1]; - local_idx[i] = local_idx[i - 1]; - } - local_vals[pos] = val; - local_idx[pos] = col; - } - } + cgh.parallel_for( + sycl::nd_range<1>(sycl::range<1>(nrows) * block_dims, block_dims), + [=](sycl::nd_item<1> item_ct1) { + const int row = item_ct1.get_group(0); + const size_t off = (size_t) row * nsplit * k; + + top_k_scan_merge_f32( + pv + off, pi + off, 0, nsplit * k, k, split_block, + shared_vals.get_multi_ptr().get(), + shared_idx.get_multi_ptr().get(), + nullptr, dst_indices + (int64_t) row * k, true, item_ct1); + }); + }); - for (int i = 0; i < k; i++) { - shared_vals[tid * k + i] = local_vals[i]; - shared_idx[tid * k + i] = local_idx[i]; - } - item_ct1.barrier(sycl::access::fence_space::local_space); + return; + } - if (tid == 0) { - float final_vals[32]; - int final_idx[32]; + const int block_size = 128; - for (int i = 0; i < k; i++) { - final_vals[i] = -FLT_MAX; - final_idx[i] = -1; - } + const sycl::range<1> block_dims(block_size); + const sycl::range<1> grid_dims(nrows); - for (int t = 0; t < block_size; t++) { - for (int i = 0; i < k; i++) { - float val = shared_vals[t * k + i]; - int idx = shared_idx[t * k + i]; - - if (val > final_vals[k-1]) { - int pos = k - 1; - while (pos > 0 && val > final_vals[pos - 1]) { - pos--; - } - - for (int j = k - 1; j > pos; j--) { - final_vals[j] = final_vals[j - 1]; - final_idx[j] = final_idx[j - 1]; - } - final_vals[pos] = val; - final_idx[pos] = idx; - } - } - } + main_stream->submit([&](sycl::handler &cgh) { + sycl::local_accessor shared_vals(sycl::range<1>((block_size + 1) * k), cgh); + sycl::local_accessor shared_idx(sycl::range<1>((block_size + 1) * k), cgh); - for (int i = 0; i < k; i++) { - dst_idx_row[i] = final_idx[i]; - } + cgh.parallel_for( + sycl::nd_range<1>(grid_dims * block_dims, block_dims), + [=](sycl::nd_item<1> item_ct1) { + const int row = item_ct1.get_group(0); - if (k > 1) { - int32_t temp = dst_idx_row[0]; - dst_idx_row[0] = dst_idx_row[1]; - dst_idx_row[1] = temp; - } - } + if (row >= nrows) return; + + top_k_scan_merge_f32( + src + (int64_t) row * ncols, nullptr, 0, (int) ncols, k, block_size, + shared_vals.get_multi_ptr().get(), + shared_idx.get_multi_ptr().get(), + nullptr, dst_indices + (int64_t) row * k, true, item_ct1); }); }); } @@ -2902,7 +3042,7 @@ static void ggml_sycl_op_top_k(ggml_backend_sycl_context & ctx, ggml_tensor * ds GGML_ASSERT(k > 0 && k <= 32); GGML_ASSERT(k <= ncols); - top_k_f32_sycl(src0_dd, dst_dd, ncols, nrows, k, main_stream); + top_k_f32_sycl(ctx, src0_dd, dst_dd, ncols, nrows, k, main_stream); } inline void ggml_sycl_op_argmax(ggml_backend_sycl_context & ctx, ggml_tensor * dst) { From 3ad8b9b217e7d0d8a829e9344dbe35bfdf19dc64 Mon Sep 17 00:00:00 2001 From: Max Krasnyansky Date: Sat, 29 Aug 2026 22:57:55 -0700 Subject: [PATCH 036/104] hexagon: support for device discovery and create sessions on demand (llama/27785) * hex-devices: add support for lazy session allocation and cleanup dev interfaces Co-authored-by: Marco Colombo * hex-devices: support for runtime discovery of available NPU cores Co-authored-by: Alexander Lu Co-authored-by: Ehsan Bateni * hex-devices: reject non-existing devices early during init --------- Co-authored-by: Marco Colombo Co-authored-by: Alexander Lu Co-authored-by: Ehsan Bateni --- ggml/src/ggml-hexagon/ggml-hexagon.cpp | 317 ++++++++++++++++--------- ggml/src/ggml-hexagon/htp-drv.cpp | 10 + ggml/src/ggml-hexagon/htp-drv.h | 2 + 3 files changed, 223 insertions(+), 106 deletions(-) diff --git a/ggml/src/ggml-hexagon/ggml-hexagon.cpp b/ggml/src/ggml-hexagon/ggml-hexagon.cpp index e7dcdc3d551..87e6989bd2f 100644 --- a/ggml/src/ggml-hexagon/ggml-hexagon.cpp +++ b/ggml/src/ggml-hexagon/ggml-hexagon.cpp @@ -69,30 +69,15 @@ using u32vec = std::vector; #define GGML_HEXAGON_FENCE_SLOT_SIZE 128 struct ggml_hexagon_device_config { - int physical_idx = 0; - int virtual_idx = 0; + int physical_idx = 0; + int virtual_idx = 0; + int domain_id = 0; + std::string domain_name; std::string name; }; static ggml_hexagon_device_config opt_device_configs[GGML_HEXAGON_MAX_SESSIONS]; -static int get_domain_id(int physical_idx) { - switch (physical_idx) { - case 0: return 3; // CDSP0 (all devices) - case 1: return 4; // CDSP1 (IQ9, IQ10) - case 2: return 18; // CDSP2 (IQ10) - case 3: return 19; // CDSP3 (IQ10) - default: return CDSP_DOMAIN_ID + physical_idx; - } -} - -static std::string get_domain_name(int physical_idx) { - if (physical_idx == 0) { - return CDSP_DOMAIN_NAME; - } - return std::string("cdsp") + std::to_string(physical_idx); -} - static int opt_arch = 0; // autodetect static size_t opt_ndev = 1; static size_t opt_nhvx = 0; // use all @@ -361,7 +346,6 @@ struct ggml_hexagon_session { uint32_t session_id; uint32_t domain_id; uint64_t queue_id; - int dev_id; int phys_idx; int virt_idx; bool valid_session; @@ -376,9 +360,6 @@ struct ggml_hexagon_session { std::unordered_map> cloned_buffers; std::unordered_set sync_peers; - ggml_backend_buffer_type buffer_type = {}; - ggml_backend_buffer_type host_buffer_type = {}; - uint32_t n_threads = 0; uint32_t n_hvx = 0; uint32_t n_hmx = 0; @@ -392,12 +373,12 @@ struct ggml_hexagon_session { mutable std::unordered_set needs_repack; - ggml_hexagon_session(int dev_id, ggml_backend_dev_t dev) noexcept(false); + ggml_hexagon_session(const ggml_hexagon_device_config & config, ggml_backend_dev_t dev = nullptr) noexcept(false); ~ggml_hexagon_session() noexcept(true); const char* c_name() const { return name.c_str(); } - void allocate(int dev_id) noexcept(false); + void allocate(const ggml_hexagon_device_config & config) noexcept(false); void release() noexcept(true); void enqueue_op(const htp_opnode & node); @@ -430,14 +411,38 @@ struct ggml_hexagon_session { // ** backend buffers +struct ggml_backend_hexagon_device_context { + int dev_id; + ggml_hexagon_device_config config; + ggml_backend_dev_t dev = nullptr; + size_t max_bufsize = 0; + + ggml_backend_buffer_type buffer_type = {}; + ggml_backend_buffer_type host_buffer_type = {}; + + std::unique_ptr sess; + + ggml_backend_hexagon_device_context(int dev_id, const ggml_hexagon_device_config & config, ggml_backend_dev_t dev); + ~ggml_backend_hexagon_device_context(); + + const char * c_name() const { return config.name.c_str(); } + + ggml_hexagon_session * session() { + if (!sess) { + sess = std::make_unique(config, dev); + } + return sess.get(); + } +}; + struct ggml_backend_hexagon_buffer_type_context { - ggml_backend_hexagon_buffer_type_context(const std::string & name, ggml_hexagon_session * sess) { - this->sess = sess; - this->name = name; + ggml_backend_hexagon_buffer_type_context(const std::string & name, ggml_backend_hexagon_device_context * dev_ctx) { + this->dev_ctx = dev_ctx; + this->name = name; } - ggml_hexagon_session * sess; - std::string name; + ggml_backend_hexagon_device_context * dev_ctx; + std::string name; }; struct ggml_hexagon_rpcmem_block { @@ -576,7 +581,8 @@ struct ggml_hexagon_shared_buffer { }; static ggml_hexagon_session * ggml_backend_hexagon_buffer_get_sess(ggml_backend_buffer_t buffer) { - return static_cast(buffer->buft->context)->sess; + auto sbuf = static_cast(buffer->context); + return sbuf->sess; } static void ggml_backend_hexagon_buffer_free_buffer(ggml_backend_buffer_t buffer) { @@ -1494,24 +1500,26 @@ static const char * ggml_backend_hexagon_buffer_type_name(ggml_backend_buffer_ty static ggml_backend_buffer_t ggml_backend_hexagon_buffer_type_alloc_buffer( ggml_backend_buffer_type_t buffer_type, size_t size) { - auto sess = static_cast(buffer_type->context)->sess; + auto dev_ctx = static_cast(buffer_type->context)->dev_ctx; + auto sess = dev_ctx->session(); try { ggml_hexagon_shared_buffer * sbuf = new ggml_hexagon_shared_buffer(sess, size, false, GGML_HEXAGON_FENCE_BUFFER_SIZE); return ggml_backend_buffer_init(buffer_type, ggml_backend_hexagon_buffer_interface, sbuf, size); } catch (const std::exception & exc) { - GGML_LOG_ERROR("ggml-hex: %s failed to allocate device buffer context: %s\n", sess->c_name(), exc.what()); + GGML_LOG_ERROR("ggml-hex: %s failed to allocate device buffer context: %s\n", dev_ctx->c_name(), exc.what()); return nullptr; } } static ggml_backend_buffer_t ggml_backend_hexagon_host_buffer_type_alloc_buffer( ggml_backend_buffer_type_t buffer_type, size_t size) { - auto sess = static_cast(buffer_type->context)->sess; + auto dev_ctx = static_cast(buffer_type->context)->dev_ctx; + auto sess = dev_ctx->session(); try { ggml_hexagon_shared_buffer * sbuf = new ggml_hexagon_shared_buffer(sess, size, false, GGML_HEXAGON_FENCE_BUFFER_SIZE); return ggml_backend_buffer_init(buffer_type, ggml_backend_hexagon_host_buffer_interface, sbuf, size); } catch (const std::exception & exc) { - GGML_LOG_ERROR("ggml-hex: %s failed to allocate host buffer context: %s\n", sess->c_name(), exc.what()); + GGML_LOG_ERROR("ggml-hex: %s failed to allocate host buffer context: %s\n", dev_ctx->c_name(), exc.what()); return nullptr; } } @@ -1536,7 +1544,7 @@ static size_t ggml_backend_hexagon_buffer_type_get_alloc_size(ggml_backend_buffe static size_t ggml_backend_hexagon_buffer_type_get_max_size(ggml_backend_buffer_type_t buft) { auto * context = static_cast(buft->context); - return context->sess->max_bufsize; + return context->dev_ctx->max_bufsize; } static bool ggml_backend_hexagon_buffer_type_is_host(ggml_backend_buffer_type_t buft) { @@ -1567,6 +1575,22 @@ static ggml_backend_buffer_type_i ggml_backend_hexagon_host_buffer_type_interfac /* .is_host = */ ggml_backend_hexagon_host_buffer_type_is_host, }; +ggml_backend_hexagon_device_context::ggml_backend_hexagon_device_context(int dev_id, const ggml_hexagon_device_config & config, ggml_backend_dev_t dev) + : dev_id(dev_id), config(config), dev(dev), max_bufsize(opt_mbuf) { + buffer_type.device = dev; + buffer_type.iface = ggml_backend_hexagon_buffer_type_interface; + buffer_type.context = new ggml_backend_hexagon_buffer_type_context(config.name, this); + + host_buffer_type.device = dev; + host_buffer_type.iface = ggml_backend_hexagon_host_buffer_type_interface; + host_buffer_type.context = new ggml_backend_hexagon_buffer_type_context(config.name + "-HOST", this); +} + +ggml_backend_hexagon_device_context::~ggml_backend_hexagon_device_context() { + delete static_cast(buffer_type.context); + delete static_cast(host_buffer_type.context); +} + static bool ggml_backend_buffer_is_hexagon(const struct ggml_backend_buffer * b) { return b->buft->iface.get_alignment == ggml_backend_hexagon_buffer_type_get_alignment; } @@ -2811,8 +2835,7 @@ static size_t ggml_hexagon_measure_max_vmem(ggml_hexagon_session *sess) { return vmem - step; // backoff to account for overhead from internal mappings } -void ggml_hexagon_session::allocate(int dev_id) noexcept(false) { - const auto & config = opt_device_configs[dev_id]; +void ggml_hexagon_session::allocate(const ggml_hexagon_device_config & config) noexcept(false) { int phys_idx = config.physical_idx; int virt_idx = config.virtual_idx; @@ -2823,21 +2846,31 @@ void ggml_hexagon_session::allocate(int dev_id) noexcept(false) { this->phys_idx = phys_idx; this->virt_idx = virt_idx; - this->domain_id = get_domain_id(phys_idx); + this->domain_id = config.domain_id; this->session_id = 0; - this->dev_id = dev_id; this->name = config.name; this->op_pending = 0; GGML_LOG_DEBUG("ggml-hex: %s allocating new session\n", this->name.c_str()); - domain * my_domain = htpdrv_get_domain(this->domain_id); - if (my_domain == NULL) { - GGML_LOG_ERROR("ggml-hex: unable to get domain struct for CDSP (domain_id %d)\n", this->domain_id); - throw std::runtime_error("ggml-hex: failed to get CDSP domain (see log for details)"); + if (config.domain_id < 0 || config.domain_name.empty()) { + GGML_LOG_ERROR("ggml-hex: %s: invalid physical CDSP core %d\n", config.name.c_str(), config.physical_idx); + throw std::runtime_error("ggml-hex: invalid physical CDSP core"); } - std::string dom_name = get_domain_name(phys_idx); + const std::string & dom_name = config.domain_name; + + // Enable Unsigned PD for all domains + { + struct remote_rpc_control_unsigned_module u; + u.domain = -1; + u.enable = 1; + int err = remote_session_control(DSPRPC_CONTROL_UNSIGNED_MODULE, (void *) &u, sizeof(u)); + if (err != AEE_SUCCESS) { + GGML_LOG_ERROR("ggml-hex: %s failed to enable unsigned PD : error 0x%x\n", this->c_name(), err); + throw std::runtime_error("ggml-hex: remote_session_control(unsign) failed (see log for details)"); + } + } // Create new session if virtual_idx > 0 if (virt_idx > 0) { @@ -2849,7 +2882,8 @@ void ggml_hexagon_session::allocate(int dev_id) noexcept(false) { int err = remote_session_control(FASTRPC_RESERVE_NEW_SESSION, (void *) &n, sizeof(n)); if (err != AEE_SUCCESS) { - GGML_LOG_ERROR("ggml-hex: failed to reserve new session %d (physical %d, virtual %d) : error 0x%x\n", dev_id, phys_idx, virt_idx, err); + GGML_LOG_ERROR("ggml-hex: %s failed to reserve new session (physical %d, virtual %d) : error 0x%x\n", + this->c_name(), phys_idx, virt_idx, err); throw std::runtime_error("ggml-hex: remote_session_control(new-sess) failed (see log for details)"); } @@ -2857,10 +2891,21 @@ void ggml_hexagon_session::allocate(int dev_id) noexcept(false) { this->session_id = n.session_id; this->domain_id = n.effective_domain_id; this->valid_session = true; + } else { + struct remote_rpc_effective_domain_id eff = {}; + eff.domain_name = const_cast(dom_name.c_str()); + eff.domain_name_len = dom_name.size(); + eff.session_id = 0; + + int err = remote_session_control(FASTRPC_GET_EFFECTIVE_DOMAIN_ID, (void *) &eff, sizeof(eff)); + if (err == AEE_SUCCESS) { + this->domain_id = eff.effective_domain_id; + } else { + GGML_LOG_DEBUG("ggml-hex: %s FASTRPC_GET_EFFECTIVE_DOMAIN_ID returned 0x%x, using domain_id %d\n", + this->name.c_str(), err, this->domain_id); + } } - // Get session URI - char session_uri[256]; { char htp_uri[256]; @@ -2877,31 +2922,18 @@ void ggml_hexagon_session::allocate(int dev_id) noexcept(false) { int err = remote_session_control(FASTRPC_GET_URI, (void *) &u, sizeof(u)); if (err != AEE_SUCCESS) { - // fallback to single session uris - int htp_URI_domain_len = strlen(htp_uri) + MAX_DOMAIN_NAMELEN; - - snprintf(session_uri, htp_URI_domain_len, "%s%s", htp_uri, my_domain->uri); + snprintf(session_uri, sizeof(session_uri), "%s&_dom=%s&_session=%u", + htp_uri, dom_name.c_str(), this->session_id); - GGML_LOG_WARN("ggml-hex: failed to get URI for session %d (physical %d, virtual %d) : error 0x%x. Falling back to single session URI: %s\n", dev_id, phys_idx, virt_idx, err, session_uri); - } - } - - // Enable Unsigned PD - { - struct remote_rpc_control_unsigned_module u; - u.domain = this->domain_id; - u.enable = 1; - int err = remote_session_control(DSPRPC_CONTROL_UNSIGNED_MODULE, (void *) &u, sizeof(u)); - if (err != AEE_SUCCESS) { - GGML_LOG_ERROR("ggml-hex: failed to enable unsigned PD for session %d : error 0x%x\n", dev_id, err); - throw std::runtime_error("ggml-hex: remote_session_control(unsign) failed (see log for details)"); + GGML_LOG_WARN("ggml-hex: %s failed to get URI (physical %d, virtual %d) : error 0x%x. Falling back to single session URI: %s\n", + this->c_name(), phys_idx, virt_idx, err, session_uri); } } // Open session int err = htp_iface_open(session_uri, &this->handle); if (err != AEE_SUCCESS) { - GGML_LOG_ERROR("ggml-hex: failed to open session %d : error 0x%x\n", dev_id, err); + GGML_LOG_ERROR("ggml-hex: %s failed to open session : error 0x%x\n", this->c_name(), err); throw std::runtime_error("ggml-hex: failed to open session (see log for details)"); } @@ -2991,7 +3023,7 @@ void ggml_hexagon_session::allocate(int dev_id) noexcept(false) { this->op_batch = new ggml_hexagon_opbatch(this, opt_opbatch, this->max_vmem); // Start dspqueue/opbatch processing - err = htp_iface_start(this->handle, dev_id, this->queue_id, opt_nhvx, opt_nhmx, this->max_vmem); + err = htp_iface_start(this->handle, this->session_id, this->queue_id, opt_nhvx, opt_nhmx, this->max_vmem); if (err != 0) { GGML_LOG_ERROR("ggml-hex: %s failed to start session: 0x%08x\n", this->c_name(), (unsigned) err); throw std::runtime_error("ggml-hex: iface start failed (see log for details)"); @@ -3054,33 +3086,23 @@ void ggml_hexagon_session::release() noexcept(true) { this->cloned_buffers.clear(); } -ggml_hexagon_session::ggml_hexagon_session(int dev_id, ggml_backend_dev_t dev) noexcept(false) { - buffer_type.device = dev; - host_buffer_type.device = dev; - +ggml_hexagon_session::ggml_hexagon_session(const ggml_hexagon_device_config & config, ggml_backend_dev_t dev) noexcept(false) { op_batch = nullptr; op_queue = nullptr; fence_seq = ((uintptr_t)this) & 0xFFFF; try { - allocate(dev_id); - - buffer_type.iface = ggml_backend_hexagon_buffer_type_interface; - buffer_type.context = new ggml_backend_hexagon_buffer_type_context(this->name, this); - - host_buffer_type.iface = ggml_backend_hexagon_host_buffer_type_interface; - host_buffer_type.context = new ggml_backend_hexagon_buffer_type_context(this->name + "-HOST", this); + allocate(config); } catch (const std::exception & exc) { release(); throw; } + + GGML_UNUSED(dev); } ggml_hexagon_session::~ggml_hexagon_session() noexcept(true) { release(); - - delete static_cast(buffer_type.context); - delete static_cast(host_buffer_type.context); } // ** backend interface @@ -3957,11 +3979,13 @@ static void ggml_hexagon_precompute_fused_mmnx_params( } static bool ggml_hexagon_tensor_is_host(const struct ggml_hexagon_session * sess, const struct ggml_tensor * t) { - return t && t->buffer && t->buffer->buft == &sess->host_buffer_type; + return t && t->buffer && ggml_backend_buft_is_host(t->buffer->buft); + GGML_UNUSED(sess); } static bool ggml_hexagon_tensor_is_non_host(const struct ggml_hexagon_session * sess, const struct ggml_tensor * t) { - return t && t->buffer && t->buffer->buft != &sess->host_buffer_type; + return t && t->buffer && !ggml_backend_buft_is_host(t->buffer->buft); + GGML_UNUSED(sess); } static bool ggml_hexagon_supported_mul_mat(const struct ggml_hexagon_session * sess, const struct ggml_tensor * dst) { @@ -5269,7 +5293,8 @@ bool ggml_backend_is_hexagon(ggml_backend_t backend) { // device interface static ggml_backend_t ggml_backend_hexagon_device_init(ggml_backend_dev_t dev, const char * params) { - auto sess = static_cast(dev->context); + auto dev_ctx = static_cast(dev->context); + auto sess = dev_ctx->session(); return new ggml_backend{ /* .guid = */ ggml_backend_hexagon_guid(), @@ -5282,8 +5307,8 @@ static ggml_backend_t ggml_backend_hexagon_device_init(ggml_backend_dev_t dev, c } static const char * ggml_backend_hexagon_device_get_name(ggml_backend_dev_t dev) { - auto sess = static_cast(dev->context); - return sess->c_name(); + auto dev_ctx = static_cast(dev->context); + return dev_ctx->c_name(); GGML_UNUSED(dev); } @@ -5321,16 +5346,16 @@ static void ggml_backend_hexagon_device_get_props(ggml_backend_dev_t dev, struct } static ggml_backend_buffer_type_t ggml_backend_hexagon_device_get_buffer_type(ggml_backend_dev_t dev) { - auto sess = static_cast(dev->context); - return &sess->buffer_type; + auto dev_ctx = static_cast(dev->context); + return &dev_ctx->buffer_type; } static ggml_backend_buffer_type_t ggml_backend_hexagon_device_get_host_buffer_type(ggml_backend_dev_t dev) { if (!opt_hostbuf) { return NULL; } - auto sess = static_cast(dev->context); - return &sess->host_buffer_type; + auto dev_ctx = static_cast(dev->context); + return &dev_ctx->host_buffer_type; } static bool ggml_hexagon_supported_cpy(const struct ggml_hexagon_session * sess, const struct ggml_tensor * op) { @@ -5421,7 +5446,8 @@ static bool ggml_hexagon_supported_fill(const struct ggml_hexagon_session * sess } static bool ggml_backend_hexagon_device_supports_op(ggml_backend_dev_t dev, const struct ggml_tensor * op) { - auto sess = static_cast(dev->context); + auto dev_ctx = static_cast(dev->context); + auto sess = dev_ctx->session(); // reject ops that match the filter if (opt_opfilter && std::regex_match(ggml_op_desc(op), *opt_opfilter)) { @@ -5493,6 +5519,7 @@ static bool ggml_backend_hexagon_device_supports_op(ggml_backend_dev_t dev, cons supp = ggml_hexagon_supported_unary(sess, op); break; default: + supp = false; break; } break; @@ -5505,6 +5532,7 @@ static bool ggml_backend_hexagon_device_supports_op(ggml_backend_dev_t dev, cons supp = ggml_hexagon_supported_activations(sess, op); break; default: + supp = false; break; } break; @@ -5590,17 +5618,17 @@ static bool ggml_backend_hexagon_device_supports_op(ggml_backend_dev_t dev, cons } static bool ggml_backend_hexagon_device_supports_buft(ggml_backend_dev_t dev, ggml_backend_buffer_type_t buft) { - auto sess = static_cast(dev->context); + auto dev_ctx = static_cast(dev->context); // Technically we can clone hexagon buffers from any session but for some reason the output is garbled with layer-split, // tensor-split works correctly, so it needs mode debugging and investigation. For now accept only our own buffers. #if 0 bool supp = (buft->iface.get_alignment == ggml_backend_hexagon_buffer_type_get_alignment); #else - bool supp = (buft == &sess->host_buffer_type) || (buft == &sess->buffer_type); + bool supp = (buft == &dev_ctx->host_buffer_type) || (buft == &dev_ctx->buffer_type); #endif - HEX_VERBOSE("ggml-hex: %s device-supports-buft %s %s\n", sess->name.c_str(), ggml_backend_buft_name(buft), supp ? "yes" : "no"); + HEX_VERBOSE("ggml-hex: %s device-supports-buft %s %s\n", dev_ctx->c_name(), ggml_backend_buft_name(buft), supp ? "yes" : "no"); return supp; } @@ -5629,16 +5657,11 @@ ggml_hexagon_registry::ggml_hexagon_registry(ggml_backend_reg_t reg) { GGML_LOG_INFO("ggml-hex: Hexagon Arch version v%d\n", opt_arch); - // Create devices / sessions + // Create devices for (size_t i = 0; i < opt_ndev; i++) { - devices[i].iface = ggml_backend_hexagon_device_i; - devices[i].reg = reg; - try { - devices[i].context = new ggml_hexagon_session(i, &devices[i]); - } catch (const std::exception & exc) { - GGML_LOG_ERROR("ggml-hex: failed to create device/session %zu\n", i); - devices[i].context = nullptr; - } + devices[i].iface = ggml_backend_hexagon_device_i; + devices[i].reg = reg; + devices[i].context = new ggml_backend_hexagon_device_context(i, opt_device_configs[i], &devices[i]); } } @@ -5646,10 +5669,10 @@ ggml_hexagon_registry::ggml_hexagon_registry(ggml_backend_reg_t reg) { ggml_hexagon_registry::~ggml_hexagon_registry() { GGML_LOG_INFO("ggml-hex: releasing registry\n"); - // Release devices / sessions + // Release devices for (size_t i = 0; i < opt_ndev; i++) { - auto sess = static_cast(devices[i].context); - delete sess; + auto dev_ctx = static_cast(devices[i].context); + delete dev_ctx; } } @@ -5818,6 +5841,85 @@ template std::string vec_to_str(std::vector v) { return str; } +// Enumerate NPU (aka CDSP) domains via FASTRPC_GET_DOMAINS if supported, +// and populate domain_id and domain_name for all configured devices. +static void ggml_hexagon_discover_devices() { + std::unordered_map cdsp_map; + bool discovery_supported = false; + + system_req_payload domain_info = {}; + domain_info.id = FASTRPC_GET_DOMAINS; + domain_info.sys.domains = nullptr; + domain_info.sys.max_domains = 0; + domain_info.sys.flags = DOMAINS_LIST_FLAGS_SET_TYPE(0, FASTRPC_NSP); + + int err = remote_system_request(&domain_info); + if (err == AEE_SUCCESS && domain_info.sys.num_domains > 0) { + std::vector domains(domain_info.sys.num_domains); + domain_info.sys.domains = domains.data(); + domain_info.sys.max_domains = (int) domains.size(); + + err = remote_system_request(&domain_info); + if (err == AEE_SUCCESS) { + discovery_supported = true; + const int n_domains = std::min(domain_info.sys.num_domains, (int) domains.size()); + for (int i = 0; i < n_domains; i++) { + GGML_LOG_INFO("ggml-hex: FASTRPC_GET_DOMAINS[%d]: type %d id %d name '%s' status %d instance-id %d\n", + i, (int) domains[i].type, domains[i].id, domains[i].name, domains[i].status, domains[i].instance_id); + if (domains[i].type != FASTRPC_NSP) { + GGML_LOG_DEBUG("ggml-hex: skipping non-CDSP domain (type=%d)\n", (int) domains[i].type); + continue; + } + if (!domains[i].status) { + GGML_LOG_WARN("ggml-hex: skipping CDSP domain id=%d (status=down)\n", domains[i].id); + continue; + } + cdsp_map[domains[i].instance_id] = domains[i]; + GGML_LOG_INFO("ggml-hex: using CDSP domain: instance-id %d id %d name '%s'\n", + domains[i].instance_id, domains[i].id, domains[i].name); + } + } else { + GGML_LOG_WARN("ggml-hex: FASTRPC_GET_DOMAINS fetch failed (0x%x), using static CDSP domains\n", (unsigned) err); + } + } else if (err != AEE_SUCCESS) { + GGML_LOG_DEBUG("ggml-hex: FASTRPC_GET_DOMAINS query failed (0x%x), using static CDSP domains\n", (unsigned) err); + } + + // Populate domain IDs and names for all configured devices + for (size_t i = 0; i < opt_ndev; i++) { + auto & cfg = opt_device_configs[i]; + if (discovery_supported) { + auto it = cdsp_map.find(cfg.physical_idx); + if (it != cdsp_map.end()) { + cfg.domain_id = it->second.id; + cfg.domain_name = it->second.name; + } else { + GGML_LOG_ERROR("ggml-hex: physical CDSP core %d not found on device (%zu CDSP core(s) available)\n", + cfg.physical_idx, cdsp_map.size()); + cfg.domain_id = -1; + cfg.domain_name = ""; + } + } else { + switch (cfg.physical_idx) { + case 0: + cfg.domain_id = 3; + cfg.domain_name = CDSP_DOMAIN_NAME; + break; + case 1: + cfg.domain_id = 4; + cfg.domain_name = "cdsp1"; + break; + default: + GGML_LOG_ERROR("ggml-hex: physical CDSP core %d not supported without dynamic discovery\n", + cfg.physical_idx); + cfg.domain_id = -1; + cfg.domain_name = ""; + break; + } + } + } +} + static void ggml_hexagon_init(ggml_backend_reg * reg) { // Basic sanity checks to make sure definitions match static_assert((unsigned int) HTP_TYPE_Q4_0 == (unsigned int) GGML_TYPE_Q4_0, @@ -5983,6 +6085,9 @@ static void ggml_hexagon_init(ggml_backend_reg * reg) { } #endif + // Resolve domain info for all configured devices + ggml_hexagon_discover_devices(); + if (str_profile) { opt_pmu_evt = [&]() -> std::vector { auto v = str_to_vec(str_profile); diff --git a/ggml/src/ggml-hexagon/htp-drv.cpp b/ggml/src/ggml-hexagon/htp-drv.cpp index 4f079080173..437e367c9d3 100644 --- a/ggml/src/ggml-hexagon/htp-drv.cpp +++ b/ggml/src/ggml-hexagon/htp-drv.cpp @@ -73,6 +73,7 @@ typedef int (*remote_handle64_close_pfn_t)(remote_handle h); typedef int (*remote_handle_control_pfn_t)(uint32_t req, void* data, uint32_t datalen); typedef int (*remote_handle64_control_pfn_t)(remote_handle64 h, uint32_t req, void* data, uint32_t datalen); typedef int (*remote_session_control_pfn_t)(uint32_t req, void *data, uint32_t datalen); +typedef int (*remote_system_request_pfn_t)(system_req_payload * req); // // Driver API pfns @@ -99,6 +100,7 @@ remote_handle64_close_pfn_t remote_handle64_close_pfn = nullptr; remote_handle_control_pfn_t remote_handle_control_pfn = nullptr; remote_handle64_control_pfn_t remote_handle64_control_pfn = nullptr; remote_session_control_pfn_t remote_session_control_pfn = nullptr; +remote_system_request_pfn_t remote_system_request_pfn = nullptr; // // Driver API @@ -206,6 +208,13 @@ HTPDRV_API int remote_session_control(uint32_t req, void * data, uint32_t datale return remote_session_control_pfn(req, data, datalen); } +HTPDRV_API int remote_system_request(system_req_payload * req) { + if (!remote_system_request_pfn) { + return AEE_EUNSUPPORTEDAPI; + } + return remote_system_request_pfn(req); +} + #ifdef _WIN32 static std::string wstr_to_str(std::wstring_view wstr) { @@ -367,6 +376,7 @@ int htpdrv_init() { dlsym(handle.get(), remote_handle64_control_pfn_t, remote_handle64_control_pfn, remote_handle64_control, false); dlsym(handle.get(), remote_session_control_pfn_t, remote_session_control_pfn, remote_session_control, false); dlsym(handle.get(), remote_handle64_close_pfn_t, remote_handle64_close_pfn, remote_handle64_close, false); + dlsym(handle.get(), remote_system_request_pfn_t, remote_system_request_pfn, remote_system_request, true); lib_cdsp_rpc_handle = std::move(handle); initialized = true; diff --git a/ggml/src/ggml-hexagon/htp-drv.h b/ggml/src/ggml-hexagon/htp-drv.h index f3cc0da75c2..8232780e7fd 100644 --- a/ggml/src/ggml-hexagon/htp-drv.h +++ b/ggml/src/ggml-hexagon/htp-drv.h @@ -116,6 +116,8 @@ HTPDRV_API domain * htpdrv_get_domain(int domain_id); */ HTPDRV_API int htpdrv_get_arch(int domain, int * arch); +HTPDRV_API int remote_system_request(system_req_payload * req); + #ifdef __cplusplus } #endif From 5e494599a4521878eb937765b5811f28ff314888 Mon Sep 17 00:00:00 2001 From: Ryan C Date: Sun, 30 Aug 2026 05:59:25 +0000 Subject: [PATCH 037/104] rpc : fix pre-rdma macOS versions (llama/27815) --- ggml/src/ggml-rpc/CMakeLists.txt | 6 +++++- ggml/src/ggml-rpc/transport-apple.cpp | 18 ++++++++++++++++++ 2 files changed, 23 insertions(+), 1 deletion(-) diff --git a/ggml/src/ggml-rpc/CMakeLists.txt b/ggml/src/ggml-rpc/CMakeLists.txt index b2f086380d5..af3bd0290f7 100644 --- a/ggml/src/ggml-rpc/CMakeLists.txt +++ b/ggml/src/ggml-rpc/CMakeLists.txt @@ -34,10 +34,14 @@ if (GGML_RPC_RDMA) find_library(RDMA_LIB ${RDMA_LIB_NAME} REQUIRED) endif() target_compile_definitions(ggml-rpc PRIVATE GGML_RPC_RDMA) - target_link_libraries(ggml-rpc PRIVATE ${RDMA_LIB}) if (APPLE) + # librdma.dylib only exists on macOS 26.2 and later. Link it weakly so a build made + # where it exists still loads where it does not; checked at runtime before use. + target_link_options(ggml-rpc PRIVATE "LINKER:-weak_library,${RDMA_LIB}") target_compile_definitions(ggml-rpc PRIVATE GGML_RPC_RDMA_APPLE) target_sources(ggml-rpc PRIVATE transport-apple.cpp) + else() + target_link_libraries(ggml-rpc PRIVATE ${RDMA_LIB}) endif() message(STATUS " RDMA transport enabled (${RDMA_DESC})") else() diff --git a/ggml/src/ggml-rpc/transport-apple.cpp b/ggml/src/ggml-rpc/transport-apple.cpp index c8be77a6dce..2cfaa5d3fd3 100644 --- a/ggml/src/ggml-rpc/transport-apple.cpp +++ b/ggml/src/ggml-rpc/transport-apple.cpp @@ -8,6 +8,7 @@ #include #include #include +#include #include #include #include @@ -184,11 +185,28 @@ static uint8_t rdma_first_active_port(struct ibv_context * ctx, struct ibv_port_ return 0; } +// librdma.dylib is weak-linked, so its symbols are null when it is absent. Nothing may +// call one before this has returned true. +static bool rdma_library_present() { + static const bool present = [] { + void * handle = dlopen("/usr/lib/librdma.dylib", RTLD_LAZY); + if (handle == nullptr) { + return false; + } + dlclose(handle); + return true; + }(); + return present; +} + // Called before the endpoints are exchanged: pick the local device facing this // peer, create a UC QP and register the frame rings. RDMA is point-to-point, so // the device is the one whose GID equals the bootstrap connection's local // address, i.e. the one cabled to the peer. std::unique_ptr apple_rdma::probe(int fd, const uint8_t * target_gid, uint8_t * caps) { + if (!rdma_library_present()) { + return nullptr; + } int ndev = 0; ibv_device ** devs = ibv_get_device_list(&ndev); if (!devs) return nullptr; From b66593ef1bd6971f3eafa3ab79f0c86c92a4e8e0 Mon Sep 17 00:00:00 2001 From: Daya Adianto Date: Sun, 30 Aug 2026 06:02:22 +0000 Subject: [PATCH 038/104] metal : Add fa-vec tuning for M3 Pro (llama/27963) Related issue: #27668 --- ggml/src/ggml-metal/ggml-metal-tuning.cpp | 212 ++++++++++++++++++++++ 1 file changed, 212 insertions(+) diff --git a/ggml/src/ggml-metal/ggml-metal-tuning.cpp b/ggml/src/ggml-metal/ggml-metal-tuning.cpp index 4abdafb48f3..9345892aeae 100644 --- a/ggml/src/ggml-metal/ggml-metal-tuning.cpp +++ b/ggml/src/ggml-metal/ggml-metal-tuning.cpp @@ -720,6 +720,218 @@ constexpr fa_vec_entry_t fa_vec_tuned_table[] = { { { GGML_METAL_DEVICE_M2_ULTRA, GGML_TYPE_Q8_0, 576, 512, -1, 0 }, { 1, 4 } }, { { GGML_METAL_DEVICE_M2_ULTRA, GGML_TYPE_Q8_0, 576, 512, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_F16, 32, 32, 1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_F16, 32, 32, 1, 3 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_F16, 32, 32, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_F16, 32, 32, 2, 3 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_F16, 32, 32, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_F16, 32, 32, 3, 3 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_F16, 64, 64, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_F16, 64, 64, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_F16, 64, 64, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_F16, 64, 64, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_F16, 96, 96, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_F16, 96, 96, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_F16, 96, 96, 1, 4 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_F16, 96, 96, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_F16, 96, 96, 2, 4 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_F16, 128, 128, -1, 1 }, { 2, 2 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_F16, 128, 128, 1, 2 }, { 1, 1 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_F16, 128, 128, 1, 4 }, { 1, 1 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_F16, 128, 128, 2, 2 }, { 1, 1 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_F16, 128, 128, 2, 4 }, { 1, 1 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_F16, 192, 192, -1, 1 }, { 2, 2 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_F16, 192, 192, 1, 2 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_F16, 192, 192, 1, 4 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_F16, 192, 128, -1, 1 }, { 2, 2 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_F16, 192, 128, 1, 2 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_F16, 192, 128, 1, 4 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_F16, 192, 128, 2, 2 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_F16, 256, 256, -1, 0 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_F16, 256, 256, 3, 0 }, { 2, 2 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_F16, 256, 256, -1, 1 }, { 2, 2 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_F16, 320, 256, 3, 0 }, { 2, 2 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_F16, 320, 256, -1, 1 }, { 2, 2 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_F16, 512, 512, 3, 0 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q4_0, 32, 32, -1, 1 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q4_0, 32, 32, 1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q4_0, 32, 32, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q4_0, 32, 32, 2, 4 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q4_0, 32, 32, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q4_0, 64, 64, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q4_0, 64, 64, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q4_0, 64, 64, 3, 2 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q4_0, 64, 64, 3, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q4_0, 96, 96, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q4_0, 96, 96, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q4_0, 96, 96, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q4_0, 128, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q4_0, 128, 128, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q4_0, 128, 128, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q4_0, 128, 128, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q4_0, 128, 128, 2, 4 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q4_0, 128, 128, 3, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q4_0, 192, 192, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q4_0, 192, 192, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q4_0, 192, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q4_0, 192, 128, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q4_0, 192, 128, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q4_0, 192, 128, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q4_0, 192, 128, 2, 4 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q4_0, 192, 128, 3, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q4_0, 256, 256, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q4_0, 256, 256, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q4_0, 320, 256, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q4_0, 320, 256, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q4_0, 512, 512, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q4_0, 512, 512, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q4_0, 576, 512, 2, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q4_0, 576, 512, 3, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q4_0, 576, 512, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q4_0, 576, 512, 1, 1 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q4_0, 576, 512, 1, 2 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q4_0, 576, 512, 1, 3 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q4_0, 576, 512, 1, 4 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q4_1, 32, 32, -1, 1 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q4_1, 32, 32, 1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q4_1, 32, 32, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q4_1, 32, 32, 2, 4 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q4_1, 32, 32, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q4_1, 64, 64, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q4_1, 64, 64, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q4_1, 64, 64, 3, 2 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q4_1, 64, 64, 3, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q4_1, 96, 96, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q4_1, 96, 96, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q4_1, 96, 96, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q4_1, 128, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q4_1, 128, 128, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q4_1, 128, 128, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q4_1, 128, 128, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q4_1, 128, 128, 3, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q4_1, 192, 192, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q4_1, 192, 192, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q4_1, 192, 192, 2, 3 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q4_1, 192, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q4_1, 192, 128, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q4_1, 192, 128, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q4_1, 192, 128, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q4_1, 192, 128, 3, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q4_1, 256, 256, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q4_1, 256, 256, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q4_1, 320, 256, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q4_1, 576, 512, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q4_1, 576, 512, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q5_0, 32, 32, -1, 1 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q5_0, 32, 32, 1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q5_0, 32, 32, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q5_0, 32, 32, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q5_0, 64, 64, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q5_0, 64, 64, -1, 1 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q5_0, 64, 64, 1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q5_0, 64, 64, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q5_0, 64, 64, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q5_0, 96, 96, -1, 1 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q5_0, 96, 96, 1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q5_0, 96, 96, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q5_0, 96, 96, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q5_0, 128, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q5_0, 128, 128, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q5_0, 128, 128, 2, 2 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q5_0, 128, 128, 2, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q5_0, 128, 128, 3, 2 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q5_0, 128, 128, 3, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q5_0, 192, 192, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q5_0, 192, 192, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q5_0, 192, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q5_0, 192, 128, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q5_0, 256, 256, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q5_0, 256, 256, -1, 1 }, { 2, 2 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q5_0, 256, 256, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q5_0, 256, 256, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q5_0, 256, 256, 3, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q5_0, 320, 256, -1, 1 }, { 2, 2 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q5_0, 320, 256, 1, 2 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q5_0, 320, 256, 2, 2 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q5_0, 320, 256, 3, 2 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q5_0, 512, 512, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q5_0, 512, 512, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q5_0, 576, 512, 2, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q5_0, 576, 512, 2, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q5_0, 576, 512, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q5_0, 576, 512, 2, 3 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q5_0, 576, 512, 2, 4 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q5_1, 32, 32, -1, 1 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q5_1, 32, 32, 1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q5_1, 32, 32, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q5_1, 32, 32, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q5_1, 64, 64, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q5_1, 64, 64, -1, 1 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q5_1, 64, 64, 1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q5_1, 64, 64, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q5_1, 64, 64, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q5_1, 96, 96, -1, 1 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q5_1, 96, 96, 1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q5_1, 96, 96, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q5_1, 96, 96, 2, 4 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q5_1, 96, 96, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q5_1, 96, 96, 3, 4 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q5_1, 128, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q5_1, 128, 128, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q5_1, 128, 128, 1, 2 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q5_1, 128, 128, 2, 2 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q5_1, 128, 128, 2, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q5_1, 128, 128, 3, 2 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q5_1, 128, 128, 3, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q5_1, 192, 192, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q5_1, 192, 192, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q5_1, 192, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q5_1, 192, 128, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q5_1, 256, 256, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q5_1, 256, 256, -1, 1 }, { 2, 2 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q5_1, 256, 256, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q5_1, 256, 256, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q5_1, 256, 256, 3, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q5_1, 320, 256, -1, 1 }, { 2, 2 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q5_1, 320, 256, 1, 2 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q5_1, 320, 256, 2, 2 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q5_1, 320, 256, 3, 2 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q5_1, 512, 512, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q5_1, 512, 512, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q5_1, 576, 512, 2, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q5_1, 576, 512, 2, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q5_1, 576, 512, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q5_1, 576, 512, 2, 3 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q5_1, 576, 512, 2, 4 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q8_0, 32, 32, -1, 1 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q8_0, 32, 32, 1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q8_0, 32, 32, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q8_0, 32, 32, 2, 4 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q8_0, 32, 32, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q8_0, 64, 64, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q8_0, 64, 64, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q8_0, 64, 64, 3, 2 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q8_0, 64, 64, 3, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q8_0, 96, 96, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q8_0, 128, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q8_0, 128, 128, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q8_0, 128, 128, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q8_0, 128, 128, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q8_0, 192, 192, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q8_0, 192, 192, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q8_0, 192, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q8_0, 192, 128, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q8_0, 192, 128, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q8_0, 192, 128, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q8_0, 256, 256, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q8_0, 256, 256, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q8_0, 320, 256, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q8_0, 320, 256, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q8_0, 512, 512, -1, 0 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q8_0, 512, 512, -1, 1 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q8_0, 576, 512, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_Q8_0, 576, 512, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_F16, 32, 32, 1, 3 }, { 4, 4 } }, { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_F16, 32, 32, 2, 1 }, { 2, 4 } }, { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_F16, 32, 32, 2, 3 }, { 4, 4 } }, From 1e0f38257378496e822093d61fb154f2e55c7236 Mon Sep 17 00:00:00 2001 From: Nils Gladitz Date: Sun, 30 Aug 2026 08:06:29 +0200 Subject: [PATCH 039/104] metal: add fa-vec tunings for M3 Ultra (llama/27999) --- ggml/src/ggml-metal/ggml-metal-tuning.cpp | 198 ++++++++++++++++++++++ 1 file changed, 198 insertions(+) diff --git a/ggml/src/ggml-metal/ggml-metal-tuning.cpp b/ggml/src/ggml-metal/ggml-metal-tuning.cpp index 9345892aeae..9114ab7424d 100644 --- a/ggml/src/ggml-metal/ggml-metal-tuning.cpp +++ b/ggml/src/ggml-metal/ggml-metal-tuning.cpp @@ -1024,6 +1024,204 @@ constexpr fa_vec_entry_t fa_vec_tuned_table[] = { { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q8_0, 576, 512, 1, 1 }, { 1, 2 } }, { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q8_0, 576, 512, 1, 4 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_F16, 32, 32, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_F16, 32, 32, 3, 3 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_F16, 64, 64, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_F16, 64, 64, 1, 1 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_F16, 64, 64, 1, 2 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_F16, 64, 64, 2, 2 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_F16, 96, 96, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_F16, 96, 96, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_F16, 96, 96, 1, 3 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_F16, 96, 96, 1, 4 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_F16, 96, 96, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_F16, 128, 128, -1, 1 }, { 2, 2 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_F16, 128, 128, 1, 1 }, { 1, 1 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_F16, 128, 128, 1, 2 }, { 1, 1 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_F16, 128, 128, 1, 4 }, { 1, 1 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_F16, 192, 192, -1, 1 }, { 2, 2 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_F16, 192, 192, 1, 1 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_F16, 192, 192, 1, 2 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_F16, 192, 128, -1, 1 }, { 2, 2 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_F16, 192, 128, 1, 1 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_F16, 192, 128, 1, 2 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_F16, 256, 256, -1, 0 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_F16, 256, 256, -1, 1 }, { 2, 2 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_F16, 256, 256, 1, 2 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_F16, 320, 256, -1, 1 }, { 2, 2 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_F16, 512, 512, 3, 3 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_F16, 512, 512, 3, 4 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q4_0, 32, 32, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q4_0, 32, 32, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q4_0, 32, 32, 3, 2 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q4_0, 64, 64, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q4_0, 64, 64, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q4_0, 64, 64, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q4_0, 96, 96, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q4_0, 96, 96, 3, 3 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q4_0, 96, 96, 3, 4 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q4_0, 128, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q4_0, 128, 128, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q4_0, 192, 192, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q4_0, 192, 192, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q4_0, 192, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q4_0, 192, 128, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q4_0, 192, 128, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q4_0, 256, 256, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q4_0, 256, 256, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q4_0, 320, 256, 2, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q4_0, 320, 256, 3, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q4_0, 320, 256, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q4_0, 512, 512, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q4_0, 512, 512, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q4_0, 576, 512, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q4_0, 576, 512, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q4_0, 576, 512, 1, 1 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q4_0, 576, 512, 1, 2 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q4_0, 576, 512, 1, 3 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q4_0, 576, 512, 1, 4 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q4_1, 32, 32, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q4_1, 32, 32, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q4_1, 32, 32, 3, 2 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q4_1, 64, 64, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q4_1, 64, 64, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q4_1, 64, 64, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q4_1, 64, 64, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q4_1, 96, 96, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q4_1, 96, 96, 3, 3 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q4_1, 96, 96, 3, 4 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q4_1, 128, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q4_1, 128, 128, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q4_1, 128, 128, 1, 2 }, { 1, 1 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q4_1, 128, 128, 1, 4 }, { 1, 1 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q4_1, 128, 128, 2, 2 }, { 1, 1 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q4_1, 128, 128, 3, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q4_1, 192, 192, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q4_1, 192, 192, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q4_1, 192, 192, 1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q4_1, 192, 192, 2, 3 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q4_1, 192, 192, 2, 4 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q4_1, 192, 128, 2, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q4_1, 192, 128, 3, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q4_1, 192, 128, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q4_1, 192, 128, 3, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q4_1, 256, 256, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q4_1, 256, 256, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q4_1, 320, 256, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q4_1, 320, 256, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q4_1, 576, 512, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q4_1, 576, 512, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q5_0, 32, 32, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q5_0, 32, 32, 3, 2 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q5_0, 32, 32, 3, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q5_0, 32, 32, 3, 4 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q5_0, 64, 64, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q5_0, 64, 64, -1, 1 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q5_0, 64, 64, 1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q5_0, 64, 64, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q5_0, 64, 64, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q5_0, 96, 96, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q5_0, 96, 96, 1, 2 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q5_0, 96, 96, 1, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q5_0, 128, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q5_0, 128, 128, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q5_0, 128, 128, 1, 2 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q5_0, 128, 128, 1, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q5_0, 192, 192, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q5_0, 192, 192, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q5_0, 192, 192, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q5_0, 192, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q5_0, 192, 128, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q5_0, 256, 256, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q5_0, 256, 256, -1, 1 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q5_0, 256, 256, 1, 1 }, { 2, 2 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q5_0, 256, 256, 2, 3 }, { 2, 2 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q5_0, 256, 256, 2, 4 }, { 2, 2 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q5_0, 320, 256, 1, 1 }, { 2, 2 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q5_0, 512, 512, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q5_0, 512, 512, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q5_0, 576, 512, 2, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q5_0, 576, 512, 3, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q5_0, 576, 512, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q5_0, 576, 512, 1, 1 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q5_0, 576, 512, 1, 2 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q5_0, 576, 512, 1, 3 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q5_0, 576, 512, 1, 4 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q5_1, 32, 32, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q5_1, 32, 32, 3, 2 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q5_1, 32, 32, 3, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q5_1, 32, 32, 3, 4 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q5_1, 64, 64, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q5_1, 64, 64, -1, 1 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q5_1, 64, 64, 1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q5_1, 64, 64, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q5_1, 64, 64, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q5_1, 96, 96, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q5_1, 96, 96, 1, 2 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q5_1, 96, 96, 1, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q5_1, 128, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q5_1, 128, 128, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q5_1, 128, 128, 1, 2 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q5_1, 128, 128, 1, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q5_1, 192, 192, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q5_1, 192, 192, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q5_1, 192, 192, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q5_1, 192, 192, 3, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q5_1, 192, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q5_1, 192, 128, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q5_1, 256, 256, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q5_1, 256, 256, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q5_1, 256, 256, 2, 3 }, { 2, 2 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q5_1, 256, 256, 2, 4 }, { 2, 2 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q5_1, 320, 256, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q5_1, 320, 256, 1, 1 }, { 2, 2 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q5_1, 320, 256, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q5_1, 320, 256, 1, 3 }, { 2, 2 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q5_1, 512, 512, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q5_1, 512, 512, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q5_1, 512, 512, 1, 3 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q5_1, 576, 512, 2, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q5_1, 576, 512, 3, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q5_1, 576, 512, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q5_1, 576, 512, 1, 1 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q5_1, 576, 512, 1, 2 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q5_1, 576, 512, 1, 3 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q5_1, 576, 512, 1, 4 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q8_0, 32, 32, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q8_0, 32, 32, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q8_0, 32, 32, 3, 2 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q8_0, 64, 64, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q8_0, 64, 64, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q8_0, 64, 64, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q8_0, 96, 96, 2, 3 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q8_0, 96, 96, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q8_0, 96, 96, 3, 3 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q8_0, 96, 96, 3, 4 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q8_0, 128, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q8_0, 128, 128, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q8_0, 128, 128, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q8_0, 128, 128, 3, 3 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q8_0, 128, 128, 3, 4 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q8_0, 192, 192, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q8_0, 192, 192, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q8_0, 192, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q8_0, 192, 128, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q8_0, 192, 128, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q8_0, 192, 128, 2, 3 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q8_0, 192, 128, 2, 4 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q8_0, 192, 128, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q8_0, 192, 128, 3, 3 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q8_0, 192, 128, 3, 4 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q8_0, 256, 256, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q8_0, 256, 256, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q8_0, 320, 256, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q8_0, 320, 256, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q8_0, 512, 512, -1, 0 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q8_0, 512, 512, -1, 1 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q8_0, 512, 512, 1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q8_0, 576, 512, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_ULTRA, GGML_TYPE_Q8_0, 576, 512, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M4, GGML_TYPE_F16, 32, 32, 1, 1 }, { 2, 4 } }, { { GGML_METAL_DEVICE_M4, GGML_TYPE_F16, 32, 32, 1, 3 }, { 4, 4 } }, { { GGML_METAL_DEVICE_M4, GGML_TYPE_F16, 32, 32, 2, 1 }, { 2, 4 } }, From 43acf3d6e84b210e72a5263f72f4d8fa97e540fd Mon Sep 17 00:00:00 2001 From: Ryan C Date: Sun, 30 Aug 2026 06:16:26 +0000 Subject: [PATCH 040/104] rpc: fix apple rdma error spew on teardown (llama/27908) --- ggml/src/ggml-rpc/transport-apple.cpp | 13 +++---------- 1 file changed, 3 insertions(+), 10 deletions(-) diff --git a/ggml/src/ggml-rpc/transport-apple.cpp b/ggml/src/ggml-rpc/transport-apple.cpp index 2cfaa5d3fd3..b1934175b1d 100644 --- a/ggml/src/ggml-rpc/transport-apple.cpp +++ b/ggml/src/ggml-rpc/transport-apple.cpp @@ -115,16 +115,9 @@ struct apple_rdma::impl { ~impl() { broken = true; - // the QP must be destroyed before the memory it can still write to is - // deregistered and freed: ERR only starts flushing the posted WQEs - if (qp) { - struct ibv_qp_attr a = {}; - a.qp_state = IBV_QPS_ERR; - ibv_modify_qp(qp, &a, IBV_QP_STATE); - struct ibv_wc wc[RDMA_NBUF * 2]; - while (ibv_poll_cq(cq, RDMA_NBUF * 2, wc) > 0) {} - ibv_destroy_qp(qp); - } + // destroy the QP first: it can still write to the rings until it is gone. + // no IBV_QPS_ERR before it - Apple's provider then fails every region unmap. + if (qp) ibv_destroy_qp(qp); if (send_mr) ibv_dereg_mr(send_mr); if (recv_mr) ibv_dereg_mr(recv_mr); free(send_mem); From 4b2243a6c20c4503f42ab47eeeddc274f5a917b1 Mon Sep 17 00:00:00 2001 From: Georgi Gerganov Date: Sun, 30 Aug 2026 09:17:47 +0300 Subject: [PATCH 041/104] ggml : add ggml_backend_op_alloc_size_may_expand, use it in RPC (llama/27960) some backends (Metal, SYCL, WebGPU) require additional memory for fleeting data for certain ops, which is reflected in their get_alloc_size implementations. add ggml_backend_op_alloc_size_may_expand() to the backend utils, listing these ops, and assert in ggml_backend_buft_get_alloc_size that a backend expanding the alloc size of a compute op only does so for ops listed in the helper. use the helper in the RPC backend to decide whether to query the remote server for the actual alloc size, instead of a hardcoded list. Assisted-by: pi:llama.cpp/Qwen3.8-27B --- ggml/include/ggml-backend.h | 4 ++++ ggml/src/ggml-backend.cpp | 23 +++++++++++++++++++++++ ggml/src/ggml-rpc/ggml-rpc.cpp | 6 +++--- 3 files changed, 30 insertions(+), 3 deletions(-) diff --git a/ggml/include/ggml-backend.h b/ggml/include/ggml-backend.h index cc3f8cd36e3..27375bd0a51 100644 --- a/ggml/include/ggml-backend.h +++ b/ggml/include/ggml-backend.h @@ -424,6 +424,10 @@ extern "C" { // Compare the output of two backends GGML_API bool ggml_backend_compare_graph_backend(ggml_backend_t backend1, ggml_backend_t backend2, struct ggml_cgraph * graph, ggml_backend_eval_callback callback, void * user_data, struct ggml_tensor const * const * test_nodes, size_t num_test_nodes); + // returns true for ops that may require additional memory for fleeting data on some backends, + // i.e. the backend's get_alloc_size may return more than ggml_nbytes for the output tensor + GGML_API bool ggml_backend_op_alloc_size_may_expand(enum ggml_op op); + // Tensor initialization GGML_API enum ggml_status ggml_backend_tensor_alloc(ggml_backend_buffer_t buffer, struct ggml_tensor * tensor, void * addr); GGML_API enum ggml_status ggml_backend_view_init(struct ggml_tensor * tensor); diff --git a/ggml/src/ggml-backend.cpp b/ggml/src/ggml-backend.cpp index 78eb10dfe99..fec7d7c92bf 100644 --- a/ggml/src/ggml-backend.cpp +++ b/ggml/src/ggml-backend.cpp @@ -65,6 +65,13 @@ size_t ggml_backend_buft_get_alloc_size(ggml_backend_buffer_type_t buft, const s if (buft->iface.get_alloc_size) { size_t size = buft->iface.get_alloc_size(buft, tensor); assert(size >= ggml_nbytes(tensor)); + + // [TAG_ALLOC_SIZE_EXPAND] + // if you hit this assert, update ggml_backend_op_alloc_size_may_expand() accordingly + GGML_ASSERT(size <= ggml_nbytes(tensor) || + ggml_op_is_empty(tensor->op) || + ggml_backend_op_alloc_size_may_expand(tensor->op)); + return size; } return ggml_nbytes(tensor); @@ -2101,6 +2108,22 @@ ggml_backend_t ggml_backend_sched_get_tensor_backend(ggml_backend_sched_t sched, // utils +// [TAG_ALLOC_SIZE_EXPAND] +// returns true for ops that may require additional memory for fleeting data on some backends, +// i.e. the backend's get_alloc_size may return more than ggml_nbytes for the output tensor +bool ggml_backend_op_alloc_size_may_expand(enum ggml_op op) { + switch (op) { + case GGML_OP_FLASH_ATTN_EXT: + case GGML_OP_MUL_MAT_ID: + case GGML_OP_CUMSUM: + case GGML_OP_ARGSORT: + case GGML_OP_TOP_K: + return true; + default: + return false; + } +} + enum ggml_status ggml_backend_view_init(struct ggml_tensor * tensor) { GGML_ASSERT(tensor); GGML_ASSERT(tensor->buffer == NULL); diff --git a/ggml/src/ggml-rpc/ggml-rpc.cpp b/ggml/src/ggml-rpc/ggml-rpc.cpp index 9aa5883d80d..58a8a030cfa 100644 --- a/ggml/src/ggml-rpc/ggml-rpc.cpp +++ b/ggml/src/ggml-rpc/ggml-rpc.cpp @@ -826,10 +826,10 @@ static size_t ggml_backend_rpc_buffer_type_get_alloc_size(ggml_backend_buffer_ty // See comments in init_tensor. rpc_get |= ggml_is_quantized(tensor->type) && (tensor->ne[0] % 512 != 0) && (tensor->view_src == nullptr); - // ops that require additional memory for fleeting data on certain backends + // [TAG_ALLOC_SIZE_EXPAND] + // ops that may require additional memory for fleeting data on certain backends // ref: https://github.com/ggml-org/llama.cpp/pull/15966 - rpc_get |= tensor->op == GGML_OP_FLASH_ATTN_EXT; - rpc_get |= tensor->op == GGML_OP_MUL_MAT_ID; + rpc_get |= ggml_backend_op_alloc_size_may_expand(tensor->op); if (rpc_get) { ggml_backend_rpc_buffer_type_context * buft_ctx = (ggml_backend_rpc_buffer_type_context *)buft->context; From e5c9e3e3e9a9c3f1763c54803229ce393f3669d1 Mon Sep 17 00:00:00 2001 From: LunalFresh <165352784+LunalFresh@users.noreply.github.com> Date: Sun, 30 Aug 2026 05:18:36 -0500 Subject: [PATCH 042/104] hip : optimize Q2_0 dot-product path for gfx1201 (llama/26753) * hip/gfx1201: optimize q2_0 vec_dot_q2_0_q8_1 with native amdgcn perm * Broadened HIP's Q2_0 perm optimization * Remove redundant HIP perm availability guard * Optimize HIP Q2_0 MMQ unpack with native perm * cuda: label HIP preprocessor guard * cuda: label HIP preprocessor guard * Restore MMQ tile index handling --- ggml/src/ggml-cuda/mmq-load-tiles.cuh | 8 ++++++++ ggml/src/ggml-cuda/vecdotq.cuh | 8 ++++++++ 2 files changed, 16 insertions(+) diff --git a/ggml/src/ggml-cuda/mmq-load-tiles.cuh b/ggml/src/ggml-cuda/mmq-load-tiles.cuh index 8ed704c281a..7f00bad943e 100644 --- a/ggml/src/ggml-cuda/mmq-load-tiles.cuh +++ b/ggml/src/ggml-cuda/mmq-load-tiles.cuh @@ -138,12 +138,20 @@ template static __device__ __forceinline_ for (int j = 0; j < 4; ++j) { const int q = qxi[j]; +#if defined(GGML_USE_HIP) + const uint32_t qx_indices = (q & 0x03) | ((q & 0x0C) << 6) | ((q & 0x30) << 12) | ((q & 0xC0) << 18); + const uint32_t qy_bits = q >> 8; + const uint32_t qy_indices = (qy_bits & 0x03) | ((qy_bits & 0x0C) << 6) | ((qy_bits & 0x30) << 12) | ((qy_bits & 0xC0) << 18); + const int qx = __builtin_amdgcn_perm(0x020100FF, 0x020100FF, qx_indices); + const int qy = __builtin_amdgcn_perm(0x020100FF, 0x020100FF, qy_indices); +#else // unpack even and odd crumbs into byte values const int qe = __byte_perm(0x020100FF, 0x020100FF, q >> 0); const int qo = __byte_perm(0x020100FF, 0x020100FF, q >> 2); // unshuffle values const int qx = __byte_perm(qe, qo, 0x5140); const int qy = __byte_perm(qe, qo, 0x7362); +#endif // defined(GGML_USE_HIP) #if defined(AMD_MFMA_AVAILABLE) || defined(TURING_MMA_AVAILABLE) || defined(AMD_WMMA_AVAILABLE) x_qs[i*sram_stride + dst_offset + j*2+0] = qx; diff --git a/ggml/src/ggml-cuda/vecdotq.cuh b/ggml/src/ggml-cuda/vecdotq.cuh index 0f039c735b6..ec117c57dfb 100644 --- a/ggml/src/ggml-cuda/vecdotq.cuh +++ b/ggml/src/ggml-cuda/vecdotq.cuh @@ -747,12 +747,20 @@ static __device__ __forceinline__ float vec_dot_q2_0_q8_1( const int u = get_int_b4(bq8_1_chunk->qs, j*2+0); const int v = get_int_b4(bq8_1_chunk->qs, j*2+1); +#if defined(GGML_USE_HIP) + const uint32_t qx_indices = (q & 0x03) | ((q & 0x0C) << 6) | ((q & 0x30) << 12) | ((q & 0xC0) << 18); + const uint32_t qy_bits = q >> 8; + const uint32_t qy_indices = (qy_bits & 0x03) | ((qy_bits & 0x0C) << 6) | ((qy_bits & 0x30) << 12) | ((qy_bits & 0xC0) << 18); + const int qx = __builtin_amdgcn_perm(0x020100FF, 0x020100FF, qx_indices); + const int qy = __builtin_amdgcn_perm(0x020100FF, 0x020100FF, qy_indices); +#else // unpack even and odd crumbs into byte values const int qe = __byte_perm(0x020100FF, 0x020100FF, q >> 0); const int qo = __byte_perm(0x020100FF, 0x020100FF, q >> 2); // unshuffle values const int qx = __byte_perm(qe, qo, 0x5140); const int qy = __byte_perm(qe, qo, 0x7362); +#endif // defined(GGML_USE_HIP) sumi = ggml_cuda_dp4a(u, qx, sumi); sumi = ggml_cuda_dp4a(v, qy, sumi); From 35d9e2237e7e8a6a716cd3d517ba46ba62bb1bed Mon Sep 17 00:00:00 2001 From: itterative <190138728+itterative@users.noreply.github.com> Date: Sun, 30 Aug 2026 13:47:21 +0300 Subject: [PATCH 043/104] hip: tune rdna 3 mmq config (llama/26284) --- ggml/src/ggml-cuda/mmq-config-rdna3.cuh | 316 +++++++++++------------- 1 file changed, 150 insertions(+), 166 deletions(-) diff --git a/ggml/src/ggml-cuda/mmq-config-rdna3.cuh b/ggml/src/ggml-cuda/mmq-config-rdna3.cuh index 676f27fea4d..3a3ef7bd9c0 100644 --- a/ggml/src/ggml-cuda/mmq-config-rdna3.cuh +++ b/ggml/src/ggml-cuda/mmq-config-rdna3.cuh @@ -1,289 +1,273 @@ static constexpr __host__ __device__ ggml_cuda_mmq_config ggml_cuda_mmq_get_config_rdna3(ggml_type type, int J, bool fallback) { CASE(GGML_TYPE_Q1_0, 128, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); CASE(GGML_TYPE_Q1_0, 128, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_Q1_0, 256, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q1_0, 128, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q1_0, 256, 2, 128, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); CASE(GGML_TYPE_Q1_0, 256, 2, 128, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); CASE(GGML_TYPE_Q1_0, 128, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); CASE(GGML_TYPE_Q1_0, 128, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q1_0, 256, 2, 128, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q1_0, 256, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q1_0, 256, 2, 128, 80, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q1_0, 128, 2, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q1_0, 128, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); CASE(GGML_TYPE_Q1_0, 256, 2, 128, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q1_0, 256, 2, 128, 112, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); CASE(GGML_TYPE_Q1_0, 256, 2, 128, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q2_0, 128, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); CASE(GGML_TYPE_Q2_0, 128, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_Q2_0, 256, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q2_0, 128, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); CASE(GGML_TYPE_Q2_0, 256, 2, 128, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); CASE(GGML_TYPE_Q2_0, 128, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); CASE(GGML_TYPE_Q2_0, 128, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q2_0, 256, 2, 128, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q2_0, 256, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q2_0, 256, 2, 128, 80, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q2_0, 256, 2, 128, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q2_0, 128, 2, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q2_0, 128, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q2_0, 128, 2, 64, 80, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q2_0, 128, 2, 64, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); CASE(GGML_TYPE_Q2_0, 256, 2, 128, 112, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); CASE(GGML_TYPE_Q2_0, 256, 2, 128, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); CASE(GGML_TYPE_Q4_0, 128, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); CASE(GGML_TYPE_Q4_0, 128, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_Q4_0, 256, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q4_0, 128, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q4_0, 256, 2, 128, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); CASE(GGML_TYPE_Q4_0, 256, 2, 128, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_Q4_0, 128, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q4_0, 128, 1, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); CASE(GGML_TYPE_Q4_0, 128, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q4_0, 256, 2, 128, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q4_0, 256, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q4_0, 256, 2, 128, 80, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q4_0, 256, 2, 128, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q4_0, 256, 2, 128, 112, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q4_0, 256, 2, 128, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q4_0, 128, 2, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q4_0, 128, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q4_0, 128, 4, 64, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q4_0, 128, 4, 64, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); CASE(GGML_TYPE_Q4_1, 128, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true); CASE(GGML_TYPE_Q4_1, 128, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_Q4_1, 256, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q4_1, 128, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q4_1, 256, 2, 128, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true); CASE(GGML_TYPE_Q4_1, 256, 2, 128, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_Q4_1, 128, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q4_1, 128, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q4_1, 256, 2, 128, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q4_1, 256, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q4_1, 256, 2, 128, 80, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q4_1, 256, 2, 128, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q4_1, 256, 2, 128, 112, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q4_1, 256, 2, 128, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q4_1, 128, 1, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q4_1, 128, 1, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q4_1, 128, 1, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q4_1, 128, 1, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q4_1, 128, 1, 64, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q4_1, 128, 1, 64, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); CASE(GGML_TYPE_Q5_0, 128, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); CASE(GGML_TYPE_Q5_0, 128, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_Q5_0, 256, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q5_0, 128, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q5_0, 256, 2, 128, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); CASE(GGML_TYPE_Q5_0, 256, 2, 128, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_Q5_0, 128, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q5_0, 128, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q5_0, 256, 2, 128, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q5_0, 256, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q5_0, 256, 2, 128, 80, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q5_0, 256, 2, 128, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q5_0, 256, 2, 128, 112, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q5_0, 256, 2, 128, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q5_0, 128, 1, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q5_0, 128, 4, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q5_0, 128, 1, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q5_0, 128, 1, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q5_0, 128, 1, 64, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q5_0, 128, 4, 64, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); CASE(GGML_TYPE_Q5_1, 128, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true); CASE(GGML_TYPE_Q5_1, 128, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_Q5_1, 256, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q5_1, 128, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q5_1, 256, 2, 128, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true); CASE(GGML_TYPE_Q5_1, 256, 2, 128, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_Q5_1, 128, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q5_1, 128, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q5_1, 256, 2, 128, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q5_1, 256, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q5_1, 256, 2, 128, 80, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q5_1, 256, 2, 128, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q5_1, 256, 2, 128, 112, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q5_1, 256, 2, 128, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q5_1, 128, 4, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q5_1, 128, 1, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q5_1, 128, 1, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q5_1, 128, 1, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q5_1, 128, 1, 64, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q5_1, 128, 4, 64, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); CASE(GGML_TYPE_Q8_0, 128, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); CASE(GGML_TYPE_Q8_0, 128, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_Q8_0, 256, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q8_0, 128, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q8_0, 256, 2, 128, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); CASE(GGML_TYPE_Q8_0, 256, 2, 128, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_Q8_0, 128, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q8_0, 128, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q8_0, 256, 2, 128, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q8_0, 256, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q8_0, 256, 2, 128, 80, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q8_0, 256, 2, 128, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q8_0, 256, 2, 128, 112, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q8_0, 256, 2, 128, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q8_0, 128, 1, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q8_0, 128, 1, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q8_0, 128, 1, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q8_0, 128, 1, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q8_0, 128, 1, 64, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q8_0, 128, 1, 64, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); // --------------------------------------------------------------------------------------------- CASE(GGML_TYPE_Q2_K, 128, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q2_K, MMQ_ITER_K, false, true); CASE(GGML_TYPE_Q2_K, 128, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q2_K, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_Q2_K, 256, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q2_K, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_Q2_K, 128, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q2_K, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q2_K, 128, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q2_K, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q2_K, 256, 2, 128, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q2_K, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q2_K, 128, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q2_K, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q2_K, 256, 2, 128, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_Q2_K, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q2_K, 256, 2, 128, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_Q2_K, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q2_K, 128, 1, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q2_K, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q2_K, 256, 1, 128, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q2_K, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q2_K, 128, 2, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q2_K, MMQ_ITER_K, false, false); CASE(GGML_TYPE_Q2_K, 256, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q2_K, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q2_K, 256, 2, 128, 80, GGML_CUDA_MMQ_SRAM_LAYOUT_Q2_K, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q2_K, 256, 2, 128, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_Q2_K, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q2_K, 256, 2, 128, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_Q2_K, MMQ_ITER_K, false, false); CASE(GGML_TYPE_Q3_K, 128, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, true); CASE(GGML_TYPE_Q3_K, 128, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_Q3_K, 256, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q3_K, 128, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q3_K, 256, 2, 128, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, true); CASE(GGML_TYPE_Q3_K, 256, 2, 128, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, true); CASE(GGML_TYPE_Q3_K, 128, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); CASE(GGML_TYPE_Q3_K, 128, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q3_K, 256, 2, 128, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q3_K, 256, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q3_K, 256, 2, 128, 80, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q3_K, 256, 2, 128, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q3_K, 256, 2, 128, 112, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q3_K, 256, 2, 128, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q3_K, 128, 1, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q3_K, 128, 1, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q3_K, 128, 1, 64, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q3_K, 128, 1, 64, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); CASE(GGML_TYPE_Q4_K, 128, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true); CASE(GGML_TYPE_Q4_K, 128, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_Q4_K, 256, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q4_K, 128, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q4_K, 256, 2, 128, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true); CASE(GGML_TYPE_Q4_K, 256, 2, 128, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_Q4_K, 128, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q4_K, 128, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q4_K, 256, 2, 128, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q4_K, 256, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q4_K, 256, 2, 128, 80, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q4_K, 256, 2, 128, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q4_K, 256, 2, 128, 112, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q4_K, 256, 2, 128, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q4_K, 128, 1, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q4_K, 128, 1, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q4_K, 128, 1, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q4_K, 128, 1, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q4_K, 128, 1, 64, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q4_K, 128, 1, 64, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); CASE(GGML_TYPE_Q5_K, 128, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true); CASE(GGML_TYPE_Q5_K, 128, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_Q5_K, 256, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q5_K, 128, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q5_K, 256, 2, 128, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true); CASE(GGML_TYPE_Q5_K, 256, 2, 128, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_Q5_K, 128, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q5_K, 128, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q5_K, 256, 2, 128, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q5_K, 256, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q5_K, 256, 2, 128, 80, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q5_K, 256, 2, 128, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q5_K, 256, 2, 128, 112, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q5_K, 256, 2, 128, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q5_K, 128, 1, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q5_K, 128, 1, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q5_K, 128, 1, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q5_K, 128, 4, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q5_K, 128, 4, 64, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q5_K, 128, 4, 64, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); CASE(GGML_TYPE_Q6_K, 128, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q6_K, MMQ_ITER_K, false, true); CASE(GGML_TYPE_Q6_K, 128, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q6_K, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_Q6_K, 256, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q6_K, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q6_K, 128, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q6_K, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q6_K, 256, 2, 128, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_Q6_K, MMQ_ITER_K, false, true); CASE(GGML_TYPE_Q6_K, 256, 2, 128, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_Q6_K, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_Q6_K, 128, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q6_K, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q6_K, 128, 4, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q6_K, MMQ_ITER_K, false, false); CASE(GGML_TYPE_Q6_K, 128, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q6_K, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q6_K, 256, 2, 128, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q6_K, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q6_K, 256, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q6_K, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q6_K, 256, 2, 128, 80, GGML_CUDA_MMQ_SRAM_LAYOUT_Q6_K, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q6_K, 256, 2, 128, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_Q6_K, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q6_K, 256, 2, 128, 112, GGML_CUDA_MMQ_SRAM_LAYOUT_Q6_K, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q6_K, 256, 2, 128, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_Q6_K, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q6_K, 128, 1, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q6_K, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q6_K, 128, 1, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q6_K, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q6_K, 128, 1, 64, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_Q6_K, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q6_K, 128, 1, 64, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_Q6_K, MMQ_ITER_K, false, false); // --------------------------------------------------------------------------------------------- CASE(GGML_TYPE_IQ1_S, 128, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); CASE(GGML_TYPE_IQ1_S, 128, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_IQ1_S, 256, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_IQ1_S, 128, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_IQ1_S, 256, 2, 128, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); CASE(GGML_TYPE_IQ1_S, 256, 2, 128, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); CASE(GGML_TYPE_IQ1_S, 128, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_IQ1_S, 128, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_IQ1_S, 256, 2, 128, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_IQ1_S, 256, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_IQ1_S, 256, 2, 128, 80, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_IQ1_S, 256, 2, 128, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_IQ1_S, 256, 2, 128, 112, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_IQ1_S, 256, 2, 128, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ1_S, 128, 1, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ1_S, 128, 4, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ1_S, 128, 4, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ1_S, 128, 1, 64, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ1_S, 128, 1, 64, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); CASE(GGML_TYPE_IQ2_XXS, 128, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); CASE(GGML_TYPE_IQ2_XXS, 128, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_IQ2_XXS, 256, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_IQ2_XXS, 128, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_IQ2_XXS, 256, 2, 128, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); CASE(GGML_TYPE_IQ2_XXS, 256, 2, 128, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); CASE(GGML_TYPE_IQ2_XXS, 128, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_IQ2_XXS, 128, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_IQ2_XXS, 256, 2, 128, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_IQ2_XXS, 256, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_IQ2_XXS, 256, 2, 128, 80, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_IQ2_XXS, 256, 2, 128, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_IQ2_XXS, 256, 2, 128, 112, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_IQ2_XXS, 256, 2, 128, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ2_XXS, 128, 1, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ2_XXS, 128, 4, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ2_XXS, 128, 4, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ2_XXS, 128, 2, 64, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ2_XXS, 128, 4, 64, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); CASE(GGML_TYPE_IQ2_XS, 128, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, true); CASE(GGML_TYPE_IQ2_XS, 128, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_IQ2_XS, 256, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_IQ2_XS, 128, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_IQ2_XS, 256, 2, 128, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, true); CASE(GGML_TYPE_IQ2_XS, 256, 2, 128, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, true); CASE(GGML_TYPE_IQ2_XS, 128, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); CASE(GGML_TYPE_IQ2_XS, 128, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_IQ2_XS, 256, 2, 128, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_IQ2_XS, 256, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_IQ2_XS, 256, 2, 128, 80, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_IQ2_XS, 256, 2, 128, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_IQ2_XS, 256, 2, 128, 112, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_IQ2_XS, 256, 2, 128, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ2_XS, 128, 4, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ2_XS, 128, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ2_XS, 128, 4, 64, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ2_XS, 128, 4, 64, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); CASE(GGML_TYPE_IQ2_S, 128, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, true); CASE(GGML_TYPE_IQ2_S, 128, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_IQ2_S, 256, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_IQ2_S, 128, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_IQ2_S, 256, 2, 128, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, true); CASE(GGML_TYPE_IQ2_S, 256, 2, 128, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, true); CASE(GGML_TYPE_IQ2_S, 128, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); CASE(GGML_TYPE_IQ2_S, 128, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_IQ2_S, 256, 2, 128, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_IQ2_S, 256, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_IQ2_S, 256, 2, 128, 80, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_IQ2_S, 256, 2, 128, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_IQ2_S, 256, 2, 128, 112, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_IQ2_S, 256, 2, 128, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ2_S, 128, 2, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ2_S, 256, 1, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ2_S, 128, 4, 64, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ2_S, 128, 4, 64, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); CASE(GGML_TYPE_IQ3_XXS, 128, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); CASE(GGML_TYPE_IQ3_XXS, 128, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_IQ3_XXS, 256, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_IQ3_XXS, 128, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_IQ3_XXS, 256, 2, 128, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); CASE(GGML_TYPE_IQ3_XXS, 256, 2, 128, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); CASE(GGML_TYPE_IQ3_XXS, 128, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); CASE(GGML_TYPE_IQ3_XXS, 128, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_IQ3_XXS, 256, 2, 128, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_IQ3_XXS, 256, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_IQ3_XXS, 256, 2, 128, 80, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_IQ3_XXS, 256, 2, 128, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_IQ3_XXS, 256, 2, 128, 112, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_IQ3_XXS, 256, 2, 128, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ3_XXS, 128, 4, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ3_XXS, 128, 4, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ3_XXS, 128, 2, 64, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ3_XXS, 128, 1, 64, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); CASE(GGML_TYPE_IQ3_S, 128, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); CASE(GGML_TYPE_IQ3_S, 128, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_IQ3_S, 256, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_IQ3_S, 128, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_IQ3_S, 256, 2, 128, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); CASE(GGML_TYPE_IQ3_S, 256, 2, 128, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); CASE(GGML_TYPE_IQ3_S, 128, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_IQ3_S, 128, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_IQ3_S, 256, 2, 128, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_IQ3_S, 256, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_IQ3_S, 256, 2, 128, 80, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_IQ3_S, 256, 2, 128, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_IQ3_S, 256, 2, 128, 112, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_IQ3_S, 256, 2, 128, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ3_S, 128, 4, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ3_S, 128, 4, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ3_S, 128, 1, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ3_S, 128, 2, 64, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ3_S, 128, 2, 64, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); CASE(GGML_TYPE_IQ4_XS, 128, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); CASE(GGML_TYPE_IQ4_XS, 128, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_IQ4_XS, 256, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_IQ4_XS, 128, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_IQ4_XS, 256, 2, 128, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); CASE(GGML_TYPE_IQ4_XS, 256, 2, 128, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_IQ4_XS, 128, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ4_XS, 128, 4, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); CASE(GGML_TYPE_IQ4_XS, 128, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_IQ4_XS, 256, 2, 128, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_IQ4_XS, 256, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_IQ4_XS, 256, 2, 128, 80, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_IQ4_XS, 256, 2, 128, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_IQ4_XS, 256, 2, 128, 112, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_IQ4_XS, 256, 2, 128, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ4_XS, 128, 1, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ4_XS, 128, 1, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ4_XS, 128, 2, 64, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ4_XS, 128, 2, 64, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); CASE(GGML_TYPE_IQ4_NL, 128, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); CASE(GGML_TYPE_IQ4_NL, 128, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_IQ4_NL, 256, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_IQ4_NL, 128, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_IQ4_NL, 256, 2, 128, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); CASE(GGML_TYPE_IQ4_NL, 256, 2, 128, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); CASE(GGML_TYPE_IQ4_NL, 128, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); CASE(GGML_TYPE_IQ4_NL, 128, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_IQ4_NL, 256, 2, 128, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_IQ4_NL, 256, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_IQ4_NL, 256, 2, 128, 80, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_IQ4_NL, 256, 2, 128, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_IQ4_NL, 256, 2, 128, 112, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_IQ4_NL, 256, 2, 128, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ4_NL, 128, 2, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ4_NL, 128, 1, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ4_NL, 128, 2, 64, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ4_NL, 128, 2, 64, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); // --------------------------------------------------------------------------------------------- CASE(GGML_TYPE_MXFP4, 128, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true); CASE(GGML_TYPE_MXFP4, 128, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_MXFP4, 256, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_MXFP4, 128, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_MXFP4, 256, 2, 128, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true); CASE(GGML_TYPE_MXFP4, 256, 2, 128, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_MXFP4, 128, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_MXFP4, 128, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_MXFP4, 256, 2, 128, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_MXFP4, 256, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_MXFP4, 256, 2, 128, 80, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_MXFP4, 256, 2, 128, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_MXFP4, 256, 2, 128, 112, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_MXFP4, 256, 2, 128, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_MXFP4, 128, 1, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_MXFP4, 128, 1, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_MXFP4, 128, 1, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_MXFP4, 128, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_MXFP4, 128, 2, 64, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_MXFP4, 128, 2, 64, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); CASE(GGML_TYPE_NVFP4, 128, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_NVFP4, MMQ_ITER_K, false, true); CASE(GGML_TYPE_NVFP4, 128, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_NVFP4, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_NVFP4, 256, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_NVFP4, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_NVFP4, 128, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_NVFP4, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_NVFP4, 256, 2, 128, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_NVFP4, MMQ_ITER_K, false, true); CASE(GGML_TYPE_NVFP4, 256, 2, 128, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_NVFP4, MMQ_ITER_K, false, true); CASE(GGML_TYPE_NVFP4, 128, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_NVFP4, MMQ_ITER_K, false, false); CASE(GGML_TYPE_NVFP4, 128, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_NVFP4, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_NVFP4, 256, 2, 128, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_NVFP4, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_NVFP4, 256, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_NVFP4, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_NVFP4, 256, 2, 128, 80, GGML_CUDA_MMQ_SRAM_LAYOUT_NVFP4, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_NVFP4, 128, 2, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_NVFP4, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_NVFP4, 128, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_NVFP4, MMQ_ITER_K, false, false); CASE(GGML_TYPE_NVFP4, 256, 2, 128, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_NVFP4, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_NVFP4, 256, 2, 128, 112, GGML_CUDA_MMQ_SRAM_LAYOUT_NVFP4, MMQ_ITER_K, false, false); CASE(GGML_TYPE_NVFP4, 256, 2, 128, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_NVFP4, MMQ_ITER_K, false, false); return ggml_cuda_mmq_config(GGML_TYPE_COUNT, 256, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, 256, false, true); From e900a732c8908510bc4e932544daab86c8790a12 Mon Sep 17 00:00:00 2001 From: Pascal Date: Sun, 30 Aug 2026 16:06:32 +0200 Subject: [PATCH 044/104] CUDA: use the fast mm_ids_helper path for any n_expert_used (llama/27978) The optimized path grouped warp lanes by token and required warp_size % n_expert_used == 0, with a single hardcoded exception padding 6 up to 8. Every other count fell back to the generic path, which walks the tokens one at a time with a warp reduction per token, for each of the n_expert blocks. The lane group only has to divide the warp, and the loop body already guards the padded lanes with iex < n_expert_used, so the padding generalizes to the next power of two. The 6 -> 8 case and every count already dispatched keep the exact same padding as before. n_expert_used = 10 now reaches the fast path. Measured on Qwen3.8-Flash-Next (512 experts, 10 used) at 55k context on an RTX PRO 6000, warm runs with the first one discarded: prompt processing 2334 -> 2600 t/s Token generation is unaffected, since a single token leaves nothing to walk. Other expert counts reach the fast path by adding their case to the dispatch. --- ggml/src/ggml-cuda/mmid.cu | 15 +++++++++++++-- 1 file changed, 13 insertions(+), 2 deletions(-) diff --git a/ggml/src/ggml-cuda/mmid.cu b/ggml/src/ggml-cuda/mmid.cu index f80442fbe4e..ed0851dcf8d 100644 --- a/ggml/src/ggml-cuda/mmid.cu +++ b/ggml/src/ggml-cuda/mmid.cu @@ -19,6 +19,11 @@ struct mm_ids_helper_store { }; static_assert(sizeof(mm_ids_helper_store) == 4, "unexpected size for mm_ids_helper_store"); +// the generic path passes 0, which needs no padding since it never groups lanes by token +template struct mm_ids_pow2 { static constexpr int value = 2*mm_ids_pow2<(n + 1)/2>::value; }; +template <> struct mm_ids_pow2<1> { static constexpr int value = 1; }; +template <> struct mm_ids_pow2<0> { static constexpr int value = 1; }; + // Helper function for mul_mat_id, converts ids to a more convenient format. // ids_src1 describes how to permute the flattened column indices of src1 in order to get a compact src1 tensor sorted by expert. // ids_dst describes the same mapping but for the dst tensor. @@ -32,6 +37,9 @@ static __global__ void mm_ids_helper( const int n_expert_used = n_expert_used_template == 0 ? n_expert_used_var : n_expert_used_template; const int expert = blockIdx.x; + // token slots per warp lane group, padded to a power of 2 so a warp divides evenly + constexpr int neu_padded = mm_ids_pow2::value; + extern __shared__ char data_mm_ids_helper[]; mm_ids_helper_store * store = (mm_ids_helper_store *) data_mm_ids_helper; @@ -60,8 +68,8 @@ static __global__ void mm_ids_helper( } } else { // Implementation optimized for specific numbers of experts used: - static_assert(n_expert_used == 6 || warp_size % n_expert_used == 0, "bad n_expert_used"); - const int neu_padded = n_expert_used == 6 ? 8 : n_expert_used; // Padded to next higher power of 2. + // a warp holds a whole number of token slots, so the slot count is padded to a power of 2 + static_assert(neu_padded <= warp_size && warp_size % neu_padded == 0, "bad n_expert_used"); for (int it0 = 0; it0 < n_tokens; it0 += warp_size/neu_padded) { const int it = it0 + threadIdx.x / neu_padded; @@ -156,6 +164,9 @@ void ggml_cuda_launch_mm_ids_helper( case 8: launch_mm_ids_helper< 8>(ids, ids_src1, ids_dst, expert_bounds, n_experts, n_tokens, n_expert_used, nchannels_y, si1, sis1, write_inverse, stream); break; + case 10: + launch_mm_ids_helper<10>(ids, ids_src1, ids_dst, expert_bounds, n_experts, n_tokens, n_expert_used, nchannels_y, si1, sis1, write_inverse, stream); + break; case 16: launch_mm_ids_helper<16>(ids, ids_src1, ids_dst, expert_bounds, n_experts, n_tokens, n_expert_used, nchannels_y, si1, sis1, write_inverse, stream); break; From e9583f075aafaea362afd80a87fdac74a0f54256 Mon Sep 17 00:00:00 2001 From: Aman Gupta Date: Sun, 30 Aug 2026 20:30:02 +0530 Subject: [PATCH 045/104] ggml: add SWIGLU_CLAMP (llama/27930) * ggml: add SWIGLU_CLAMP * add vulkan shader --- ggml/include/ggml.h | 7 + ggml/src/ggml-cann/aclnn_ops.cpp | 45 +++++- ggml/src/ggml-cann/aclnn_ops.h | 1 + ggml/src/ggml-cann/ggml-cann.cpp | 4 + ggml/src/ggml-cpu/ggml-cpu.c | 1 + ggml/src/ggml-cpu/ops.cpp | 137 ++++++++++++++++++ ggml/src/ggml-cuda/common.cuh | 3 +- ggml/src/ggml-cuda/ggml-cuda.cu | 20 ++- ggml/src/ggml-cuda/mmvf.cu | 8 +- ggml/src/ggml-cuda/mmvq.cu | 8 +- ggml/src/ggml-cuda/unary.cu | 75 ++++++++++ ggml/src/ggml-cuda/unary.cuh | 9 ++ ggml/src/ggml-et/et-kernels/src/glu_f32.c | 57 +++++++- ggml/src/ggml-et/ggml-et-cpu-compare.cpp | 7 +- ggml/src/ggml-et/ggml-et-ops.cpp | 3 + ggml/src/ggml-et/ggml-et.cpp | 3 +- ggml/src/ggml-hexagon/ggml-hexagon.cpp | 2 + ggml/src/ggml-hexagon/htp/act-ops.c | 28 +++- ggml/src/ggml-hexagon/htp/htp-ops.h | 1 + ggml/src/ggml-hexagon/htp/main.c | 1 + ggml/src/ggml-metal/ggml-metal-device.cpp | 1 + ggml/src/ggml-metal/ggml-metal-device.m | 1 + ggml/src/ggml-metal/kernels/unary.metal | 26 ++++ ggml/src/ggml-opencl/ggml-opencl.cpp | 19 ++- ggml/src/ggml-opencl/kernels/glu.cl | 65 +++++++++ .../ggml-openvino/openvino/op/glu_swiglu.cpp | 15 ++ ggml/src/ggml-openvino/openvino/op_table.cpp | 1 + ggml/src/ggml-openvino/openvino/op_table.h | 1 + ggml/src/ggml-sycl/element_wise.cpp | 101 +++++++++++++ ggml/src/ggml-sycl/element_wise.hpp | 1 + ggml/src/ggml-sycl/ggml-sycl.cpp | 4 + ggml/src/ggml-vulkan/ggml-vulkan.cpp | 6 + .../vulkan-shaders/swiglu_clamp.comp | 12 ++ .../vulkan-shaders/vulkan-shaders-gen.cpp | 2 + .../ggml-webgpu/ggml-webgpu-shader-lib.hpp | 4 + ggml/src/ggml-webgpu/ggml-webgpu.cpp | 3 +- ggml/src/ggml-webgpu/wgsl-shaders/glu.wgsl | 8 + ggml/src/ggml.c | 15 +- 38 files changed, 686 insertions(+), 19 deletions(-) create mode 100644 ggml/src/ggml-vulkan/vulkan-shaders/swiglu_clamp.comp diff --git a/ggml/include/ggml.h b/ggml/include/ggml.h index 5f6774a630c..26f31232f67 100644 --- a/ggml/include/ggml.h +++ b/ggml/include/ggml.h @@ -627,6 +627,7 @@ extern "C" { GGML_GLU_OP_SWIGLU_OAI, GGML_GLU_OP_GEGLU_ERF, GGML_GLU_OP_GEGLU_QUICK, + GGML_GLU_OP_SWIGLU_CLAMP, GGML_GLU_OP_COUNT, }; @@ -1367,6 +1368,12 @@ extern "C" { float alpha, float limit); + GGML_API struct ggml_tensor * ggml_swiglu_clamp( + struct ggml_context * ctx, + struct ggml_tensor * a, + struct ggml_tensor * b, + float limit); + // normalize along rows GGML_API struct ggml_tensor * ggml_norm( struct ggml_context * ctx, diff --git a/ggml/src/ggml-cann/aclnn_ops.cpp b/ggml/src/ggml-cann/aclnn_ops.cpp index 2dc0f40917d..902d2eda693 100644 --- a/ggml/src/ggml-cann/aclnn_ops.cpp +++ b/ggml/src/ggml-cann/aclnn_ops.cpp @@ -211,6 +211,50 @@ void ggml_cann_swiglu(ggml_backend_cann_context & ctx, ggml_tensor * dst) { GGML_CANN_CALL_ACLNN_OP(ctx, SwiGlu, acl_src.get(), (int64_t)2, acl_dst.get()); } +void ggml_cann_swiglu_clamp(ggml_backend_cann_context & ctx, ggml_tensor * dst) { + ggml_tensor * src0 = dst->src[0]; + ggml_tensor * src1 = dst->src[1]; + + GGML_ASSERT(ggml_is_contiguous_1(src0)); + GGML_ASSERT(ggml_is_contiguous_1(dst)); + + const int32_t swapped = ggml_get_op_params_i32(dst, 1); + acl_tensor_ptr acl_gate; + acl_tensor_ptr acl_up; + if (src1) { + GGML_ASSERT(ggml_is_contiguous_1(src1)); + GGML_ASSERT(src0->type == src1->type); + acl_gate = ggml_cann_create_tensor(src0); + acl_up = ggml_cann_create_tensor(src1); + } else { + int64_t ne[] = { src0->ne[0] / 2, src0->ne[1], src0->ne[2], src0->ne[3] }; + size_t nb[] = { src0->nb[0], src0->nb[1], src0->nb[2], src0->nb[3] }; + acl_gate = ggml_cann_create_tensor(src0, ne, nb, GGML_MAX_DIMS, ACL_FORMAT_ND, 0); + acl_up = ggml_cann_create_tensor(src0, ne, nb, GGML_MAX_DIMS, ACL_FORMAT_ND, ne[0] * ggml_element_size(src0)); + if (swapped) { + std::swap(acl_gate, acl_up); + } + } + + ggml_cann_pool_alloc temp_alloc(ctx.pool(), ggml_nbytes(dst)); + acl_tensor_ptr acl_temp = ggml_cann_create_tensor(temp_alloc.get(), ggml_cann_type_mapping(dst->type), + ggml_element_size(dst), dst->ne, dst->nb, GGML_MAX_DIMS); + acl_tensor_ptr acl_dst = ggml_cann_create_tensor(dst); + + const float limit = ggml_get_op_params_f32(dst, 3); + float min_gate = -INFINITY; + float min_up = -limit; + float max_value = limit; + acl_scalar_ptr acl_min_gate = ggml_cann_create_scalar(&min_gate, ACL_FLOAT); + acl_scalar_ptr acl_min_up = ggml_cann_create_scalar(&min_up, ACL_FLOAT); + acl_scalar_ptr acl_limit = ggml_cann_create_scalar(&max_value, ACL_FLOAT); + + GGML_CANN_CALL_ACLNN_OP(ctx, Clamp, acl_gate.get(), acl_min_gate.get(), acl_limit.get(), acl_temp.get()); + GGML_CANN_CALL_ACLNN_OP(ctx, Silu, acl_temp.get(), acl_dst.get()); + GGML_CANN_CALL_ACLNN_OP(ctx, Clamp, acl_up.get(), acl_min_up.get(), acl_limit.get(), acl_temp.get()); + GGML_CANN_CALL_ACLNN_OP(ctx, InplaceMul, acl_dst.get(), acl_temp.get()); +} + // Fused GeGLU using aclnnGeGluV3: splits input along ne[0] (CANN last dim), // activates the LEFT half with GELU, multiplies by right half. // approximate: 0=tanh, 1=none(erf). activateLeft=true matches GGML convention. @@ -4433,4 +4477,3 @@ void ggml_cann_gated_linear_attn(ggml_backend_cann_context & ctx, ggml_tensor * } } } - diff --git a/ggml/src/ggml-cann/aclnn_ops.h b/ggml/src/ggml-cann/aclnn_ops.h index cdbf9260f85..678f4d654e7 100644 --- a/ggml/src/ggml-cann/aclnn_ops.h +++ b/ggml/src/ggml-cann/aclnn_ops.h @@ -76,6 +76,7 @@ void ggml_cann_repeat(ggml_backend_cann_context & ctx, ggml_tensor * dst); void ggml_cann_swiglu(ggml_backend_cann_context & ctx, ggml_tensor * dst); +void ggml_cann_swiglu_clamp(ggml_backend_cann_context & ctx, ggml_tensor * dst); void ggml_cann_geglu(ggml_backend_cann_context & ctx, ggml_tensor * dst, int64_t approximate); /** diff --git a/ggml/src/ggml-cann/ggml-cann.cpp b/ggml/src/ggml-cann/ggml-cann.cpp index 5e5541aac94..c2745014a19 100644 --- a/ggml/src/ggml-cann/ggml-cann.cpp +++ b/ggml/src/ggml-cann/ggml-cann.cpp @@ -1872,6 +1872,9 @@ static bool ggml_cann_compute_forward(ggml_backend_cann_context & ctx, struct gg case GGML_GLU_OP_SWIGLU: ggml_cann_swiglu(ctx, dst); break; + case GGML_GLU_OP_SWIGLU_CLAMP: + ggml_cann_swiglu_clamp(ctx, dst); + break; case GGML_GLU_OP_GEGLU_QUICK: ggml_cann_geglu_quick(ctx, dst); break; @@ -2428,6 +2431,7 @@ static bool ggml_backend_cann_supports_op(ggml_backend_dev_t dev, const ggml_ten case GGML_GLU_OP_SWIGLU: case GGML_GLU_OP_GEGLU_ERF: case GGML_GLU_OP_GEGLU_QUICK: + case GGML_GLU_OP_SWIGLU_CLAMP: return true; default: return false; diff --git a/ggml/src/ggml-cpu/ggml-cpu.c b/ggml/src/ggml-cpu/ggml-cpu.c index b9c0fa3ddc0..6bc4467e378 100644 --- a/ggml/src/ggml-cpu/ggml-cpu.c +++ b/ggml/src/ggml-cpu/ggml-cpu.c @@ -2311,6 +2311,7 @@ static int ggml_get_n_tasks(struct ggml_tensor * node, int n_threads) { case GGML_GLU_OP_SWIGLU_OAI: case GGML_GLU_OP_GEGLU_ERF: case GGML_GLU_OP_GEGLU_QUICK: + case GGML_GLU_OP_SWIGLU_CLAMP: { n_tasks = n_threads; } break; diff --git a/ggml/src/ggml-cpu/ops.cpp b/ggml/src/ggml-cpu/ops.cpp index b47ce5463c6..266261c5e5a 100644 --- a/ggml/src/ggml-cpu/ops.cpp +++ b/ggml/src/ggml-cpu/ops.cpp @@ -3403,6 +3403,139 @@ static void ggml_compute_forward_swiglu_oai( } } +// ggml_compute_forward_swiglu_clamp + +static void ggml_compute_forward_swiglu_clamp_f32(const ggml_compute_params * params, ggml_tensor * dst) { + const ggml_tensor * src0 = dst->src[0]; + const ggml_tensor * src1 = dst->src[1]; + char * src0_d = (char *) src0->data; + char * src1_d = (char *) (src1 ? src1->data : src0->data); + const size_t src0_o = src0->nb[1]; + const size_t src1_o = src1 ? src1->nb[1] : src0->nb[1]; + + GGML_ASSERT(ggml_is_contiguous_1(src0)); + GGML_ASSERT(ggml_is_contiguous_1(dst)); + + if (src1) { + GGML_ASSERT(ggml_is_contiguous_1(src1)); + GGML_ASSERT(src0->type == src1->type); + } + + const int ith = params->ith; + const int nth = params->nth; + + const int nc = src1 ? src0->ne[0] : src0->ne[0] / 2; + const int nr = ggml_nrows(src0); + + GGML_ASSERT(dst->ne[0] == nc); + GGML_ASSERT(ggml_nrows(dst) == nr); + + const int32_t swapped = ggml_get_op_params_i32(dst, 1); + const float limit = ggml_get_op_params_f32(dst, 3); + + const int dr = (nr + nth - 1) / nth; + const int ir0 = dr * ith; + const int ir1 = MIN(ir0 + dr, nr); + + for (int i1 = ir0; i1 < ir1; i1++) { + float * src0_p = (float *) (src0_d + i1 * src0_o); + float * src1_p = (float *) (src1_d + i1 * src1_o); + float * dst_p = (float *) ((char *) dst->data + i1 * (dst->nb[1])); + + if (!src1) { + src0_p += swapped ? nc : 0; + src1_p += swapped ? 0 : nc; + } + + for (int k = 0; k < nc; k++) { + const float gate = std::min(src0_p[k], limit); + const float up = std::clamp(src1_p[k], -limit, limit); + dst_p[k] = gate / (1.f + expf(-gate)) * up; + } + +#ifndef NDEBUG + for (int k = 0; k < nc; k++) { + const float x = dst_p[k]; + GGML_UNUSED(x); + assert(!isnan(x)); + assert(!isinf(x)); + } +#endif // NDEBUG + } +} + +static void ggml_compute_forward_swiglu_clamp_f16(const ggml_compute_params * params, ggml_tensor * dst) { + const ggml_tensor * src0 = dst->src[0]; + const ggml_tensor * src1 = dst->src[1]; + char * src0_d = (char *) src0->data; + char * src1_d = (char *) (src1 ? src1->data : src0->data); + const size_t src0_o = src0->nb[1]; + const size_t src1_o = src1 ? src1->nb[1] : src0->nb[1]; + + GGML_ASSERT(ggml_is_contiguous_1(src0)); + GGML_ASSERT(ggml_is_contiguous_1(dst)); + + if (src1) { + GGML_ASSERT(ggml_is_contiguous_1(src1)); + GGML_ASSERT(src0->type == src1->type); + } + + const int ith = params->ith; + const int nth = params->nth; + + const int nc = src1 ? src0->ne[0] : src0->ne[0] / 2; + const int nr = ggml_nrows(src0); + + GGML_ASSERT(dst->ne[0] == nc); + GGML_ASSERT(ggml_nrows(dst) == nr); + + const int32_t swapped = ggml_get_op_params_i32(dst, 1); + const float limit = ggml_get_op_params_f32(dst, 3); + + const int dr = (nr + nth - 1) / nth; + const int ir0 = dr * ith; + const int ir1 = MIN(ir0 + dr, nr); + + for (int i1 = ir0; i1 < ir1; i1++) { + ggml_fp16_t * src0_p = (ggml_fp16_t *) (src0_d + i1 * src0_o); + ggml_fp16_t * src1_p = (ggml_fp16_t *) (src1_d + i1 * src1_o); + ggml_fp16_t * dst_p = (ggml_fp16_t *) ((char *) dst->data + i1 * (dst->nb[1])); + + if (!src1) { + src0_p += swapped ? nc : 0; + src1_p += swapped ? 0 : nc; + } + + for (int k = 0; k < nc; k++) { + const float gate = std::min(GGML_FP16_TO_FP32(src0_p[k]), limit); + const float up = std::clamp(GGML_FP16_TO_FP32(src1_p[k]), -limit, limit); + dst_p[k] = GGML_FP32_TO_FP16(gate / (1.f + expf(-gate)) * up); + } + +#ifndef NDEBUG + for (int k = 0; k < nc; k++) { + const float x = GGML_FP16_TO_FP32(dst_p[k]); + GGML_UNUSED(x); + assert(!isnan(x)); + assert(!isinf(x)); + } +#endif // NDEBUG + } +} + +static void ggml_compute_forward_swiglu_clamp(const ggml_compute_params * params, ggml_tensor * dst) { + switch (dst->src[0]->type) { + case GGML_TYPE_F32: + ggml_compute_forward_swiglu_clamp_f32(params, dst); + break; + case GGML_TYPE_F16: + ggml_compute_forward_swiglu_clamp_f16(params, dst); + break; + default: + GGML_ABORT("fatal error"); + } +} + // ggml_compute_forward_geglu_erf static void ggml_compute_forward_geglu_erf_f32( @@ -10136,6 +10269,10 @@ void ggml_compute_forward_glu( { ggml_compute_forward_geglu_quick(params, dst); } break; + case GGML_GLU_OP_SWIGLU_CLAMP: + { + ggml_compute_forward_swiglu_clamp(params, dst); + } break; default: { GGML_ABORT("fatal error"); diff --git a/ggml/src/ggml-cuda/common.cuh b/ggml/src/ggml-cuda/common.cuh index 14dd1098c97..e5ccd1feab1 100644 --- a/ggml/src/ggml-cuda/common.cuh +++ b/ggml/src/ggml-cuda/common.cuh @@ -1539,6 +1539,7 @@ struct ggml_cuda_mm_fusion_args_host { const ggml_tensor * x_scale = nullptr; const ggml_tensor * gate_scale = nullptr; ggml_glu_op glu_op; + float glu_limit = 0.0f; }; struct ggml_cuda_mm_fusion_args_device { const void * x_bias = nullptr; @@ -1547,6 +1548,7 @@ struct ggml_cuda_mm_fusion_args_device { const void * x_scale = nullptr; const void * gate_scale = nullptr; ggml_glu_op glu_op; + float glu_limit = 0.0f; }; struct ggml_cuda_kernel_launch_params { @@ -1673,4 +1675,3 @@ static __inline__ void ggml_cuda_kernel_launch(Kernel kernel, const ggml_cuda_ke kernel<<>>(std::forward(args)... ); CUDA_CHECK(cudaGetLastError()); } - diff --git a/ggml/src/ggml-cuda/ggml-cuda.cu b/ggml/src/ggml-cuda/ggml-cuda.cu index bd9754c2ffd..7c2d3cad20a 100644 --- a/ggml/src/ggml-cuda/ggml-cuda.cu +++ b/ggml/src/ggml-cuda/ggml-cuda.cu @@ -1744,7 +1744,7 @@ static bool ggml_cuda_should_fuse_mul_mat(const ggml_tensor * ffn_up, return false; } - static constexpr std::array valid_glu_ops = { GGML_GLU_OP_SWIGLU, GGML_GLU_OP_GEGLU, GGML_GLU_OP_SWIGLU_OAI }; + static constexpr std::array valid_glu_ops = { GGML_GLU_OP_SWIGLU, GGML_GLU_OP_GEGLU, GGML_GLU_OP_SWIGLU_OAI, GGML_GLU_OP_SWIGLU_CLAMP }; if (std::find(valid_glu_ops.begin(), valid_glu_ops.end(), ggml_get_glu_op(glu)) == valid_glu_ops.end()) { return false; @@ -2203,6 +2203,9 @@ static bool ggml_cuda_compute_forward(ggml_backend_cuda_context & ctx, struct gg case GGML_GLU_OP_GEGLU_QUICK: ggml_cuda_op_geglu_quick(ctx, dst); break; + case GGML_GLU_OP_SWIGLU_CLAMP: + ggml_cuda_op_swiglu_clamp(ctx, dst); + break; default: return false; } @@ -3595,6 +3598,7 @@ static int ggml_cuda_try_fuse(ggml_backend_cuda_context * cuda_ctx, ggml_cgraph fusion_data.x_scale = up_scale; fusion_data.gate_scale = gate_scale; fusion_data.glu_op = ggml_get_glu_op(glu); + fusion_data.glu_limit = ggml_get_op_params_f32(glu, 3); if (ggml_cuda_should_fuse_mul_mat_vec_q(up_n)) { ggml_cuda_mul_mat_vec_q(*cuda_ctx, src0, src1, ids, cgraph->nodes[glu_idx], &fusion_data); @@ -3688,6 +3692,7 @@ static int ggml_cuda_try_fuse(ggml_backend_cuda_context * cuda_ctx, ggml_cgraph fusion_data.x_scale = up_scale; fusion_data.gate_scale = gate_scale; fusion_data.glu_op = ggml_get_glu_op(glu); + fusion_data.glu_limit = ggml_get_op_params_f32(glu, 3); if (ggml_cuda_should_fuse_mul_mat_vec_q(up_n)) { ggml_cuda_mul_mat_vec_q(*cuda_ctx, src0, src1, ids, cgraph->nodes[glu_idx], &fusion_data); @@ -3744,6 +3749,7 @@ static int ggml_cuda_try_fuse(ggml_backend_cuda_context * cuda_ctx, ggml_cgraph fusion_data.x_bias = up_bias_tensor; fusion_data.gate_bias = gate_bias_tensor; fusion_data.glu_op = ggml_get_glu_op(glu); + fusion_data.glu_limit = ggml_get_op_params_f32(glu, 3); ggml_cuda_mul_mat_vec_f(*cuda_ctx, src0, src1, ids, glu, &fusion_data); fused_mul_mat_vec = true; @@ -3757,6 +3763,7 @@ static int ggml_cuda_try_fuse(ggml_backend_cuda_context * cuda_ctx, ggml_cgraph fusion_data.x_bias = up_bias_tensor; fusion_data.gate_bias = gate_bias_tensor; fusion_data.glu_op = ggml_get_glu_op(glu); + fusion_data.glu_limit = ggml_get_op_params_f32(glu, 3); ggml_cuda_mul_mat_vec_q(*cuda_ctx, src0, src1, ids, glu, &fusion_data); fused_mul_mat_vec = true; @@ -3781,8 +3788,9 @@ static int ggml_cuda_try_fuse(ggml_backend_cuda_context * cuda_ctx, ggml_cgraph if (ggml_cuda_should_fuse_mul_mat_vec_f(up)) { ggml_cuda_mm_fusion_args_host fusion_data{}; - fusion_data.gate = gate->src[0]; - fusion_data.glu_op = ggml_get_glu_op(glu); + fusion_data.gate = gate->src[0]; + fusion_data.glu_op = ggml_get_glu_op(glu); + fusion_data.glu_limit = ggml_get_op_params_f32(glu, 3); ggml_cuda_mul_mat_vec_f(*cuda_ctx, src0, src1, ids, glu, &fusion_data); fused_mul_mat_vec = true; @@ -3792,8 +3800,9 @@ static int ggml_cuda_try_fuse(ggml_backend_cuda_context * cuda_ctx, ggml_cgraph if (ggml_cuda_should_fuse_mul_mat_vec_q(up)) { ggml_cuda_mm_fusion_args_host fusion_data{}; - fusion_data.gate = gate->src[0]; - fusion_data.glu_op = ggml_get_glu_op(glu); + fusion_data.gate = gate->src[0]; + fusion_data.glu_op = ggml_get_glu_op(glu); + fusion_data.glu_limit = ggml_get_op_params_f32(glu, 3); ggml_cuda_mul_mat_vec_q(*cuda_ctx, src0, src1, ids, glu, &fusion_data); fused_mul_mat_vec = true; @@ -4919,6 +4928,7 @@ static bool ggml_backend_cuda_device_supports_op(ggml_backend_dev_t dev, const g case GGML_GLU_OP_SWIGLU_OAI: case GGML_GLU_OP_GEGLU_ERF: case GGML_GLU_OP_GEGLU_QUICK: + case GGML_GLU_OP_SWIGLU_CLAMP: return ggml_is_contiguous_1(op->src[0]); default: return false; diff --git a/ggml/src/ggml-cuda/mmvf.cu b/ggml/src/ggml-cuda/mmvf.cu index d7dbc8b9928..bd5c5d421a4 100644 --- a/ggml/src/ggml-cuda/mmvf.cu +++ b/ggml/src/ggml-cuda/mmvf.cu @@ -56,6 +56,7 @@ static __global__ void mul_mat_vec_f( bool use_bias = false; bool use_gate_bias = false; ggml_glu_op glu_op = ggml_glu_op::GGML_GLU_OP_SWIGLU; + float glu_limit = 0.0f; const T * gate_x = nullptr; const float * x_bias = nullptr; const float * gate_bias = nullptr; @@ -65,6 +66,7 @@ static __global__ void mul_mat_vec_f( use_bias = fusion.x_bias != nullptr; use_gate_bias = fusion.gate_bias != nullptr; glu_op = fusion.glu_op; + glu_limit = fusion.glu_limit; if (use_gate) { gate_x = static_cast(fusion.gate); @@ -365,6 +367,9 @@ static __global__ void mul_mat_vec_f( value = ggml_cuda_op_swiglu_oai_single(gate_value, value); break; } + case GGML_GLU_OP_SWIGLU_CLAMP: + value = ggml_cuda_op_swiglu_clamp_single(gate_value, value, glu_limit); + break; default: break; } @@ -374,7 +379,7 @@ static __global__ void mul_mat_vec_f( dst[tid*stride_col_dst + row] = value; if constexpr (!has_fusion) { - GGML_UNUSED_VARS(use_gate, use_bias, use_gate_bias, glu_op, gate_x, x_bias, gate_bias, sumf_gate); + GGML_UNUSED_VARS(use_gate, use_bias, use_gate_bias, glu_op, glu_limit, gate_x, x_bias, gate_bias, sumf_gate); } } @@ -675,6 +680,7 @@ void ggml_cuda_mul_mat_vec_f(ggml_backend_cuda_context & ctx, const ggml_tensor fusion_local.gate_bias = fusion->gate_bias->data; } fusion_local.glu_op = fusion->glu_op; + fusion_local.glu_limit = fusion->glu_limit; } const int64_t s01 = src0->nb[1] / ts_src0; diff --git a/ggml/src/ggml-cuda/mmvq.cu b/ggml/src/ggml-cuda/mmvq.cu index 97053480980..79f7a3f6fe7 100644 --- a/ggml/src/ggml-cuda/mmvq.cu +++ b/ggml/src/ggml-cuda/mmvq.cu @@ -595,6 +595,7 @@ static __global__ void mul_mat_vec_q( const float * x_scale = nullptr; const float * gate_scale = nullptr; ggml_glu_op active_glu; + float glu_limit = 0.0f; if constexpr (has_fusion) { use_gate = fusion.gate != nullptr; @@ -604,6 +605,7 @@ static __global__ void mul_mat_vec_q( x_bias = (const float *) fusion.x_bias; gate_bias = (const float *) fusion.gate_bias; active_glu = fusion.glu_op; + glu_limit = fusion.glu_limit; if constexpr (type == GGML_TYPE_NVFP4) { use_scale = fusion.x_scale != nullptr; use_gate_scale = fusion.gate_scale != nullptr && use_gate; @@ -745,6 +747,9 @@ static __global__ void mul_mat_vec_q( case GGML_GLU_OP_SWIGLU_OAI: result = ggml_cuda_op_swiglu_oai_single(gate_value, result); break; + case GGML_GLU_OP_SWIGLU_CLAMP: + result = ggml_cuda_op_swiglu_clamp_single(gate_value, result, glu_limit); + break; default: result = result * gate_value; break; @@ -757,7 +762,7 @@ static __global__ void mul_mat_vec_q( } if constexpr (!has_fusion) { - GGML_UNUSED_VARS(use_gate, use_bias, use_gate_bias, use_scale, use_gate_scale, active_glu, gate_bias, x_bias, x_scale, gate_scale, tmp_gate); + GGML_UNUSED_VARS(use_gate, use_bias, use_gate_bias, use_scale, use_gate_scale, active_glu, glu_limit, gate_bias, x_bias, x_scale, gate_scale, tmp_gate); } if constexpr (type != GGML_TYPE_NVFP4) { GGML_UNUSED_VARS(use_scale, use_gate_scale, x_scale, gate_scale, x_scales, gate_scales); @@ -1310,6 +1315,7 @@ void ggml_cuda_mul_mat_vec_q( fusion_local.gate_scale = fusion->gate_scale->data; } fusion_local.glu_op = fusion->glu_op; + fusion_local.glu_limit = fusion->glu_limit; } // If src0 is a temporary compute buffer, clear any potential padding. diff --git a/ggml/src/ggml-cuda/unary.cu b/ggml/src/ggml-cuda/unary.cu index 4cb805fa601..d3e594878fc 100644 --- a/ggml/src/ggml-cuda/unary.cu +++ b/ggml/src/ggml-cuda/unary.cu @@ -427,6 +427,81 @@ void ggml_cuda_op_swiglu_oai(ggml_backend_cuda_context & ctx, ggml_tensor * dst) swiglu_oai_cuda(src0_p, src1_p, (float *)dst_d, ggml_nelements(dst), nc, src0_o / sizeof(float), src1_o / sizeof(float), alpha, limit, stream); } +// swiglu_clamp + +template +static __global__ void swiglu_clamp_kernel(const T * gate, const T * up, T * dst, const int64_t k, const int64_t n, const int64_t o0, const int64_t o1, float limit) { + const int64_t i = int64_t(blockDim.x)*blockIdx.x + threadIdx.x; + + if (i >= k) { + return; + } + + const int64_t j0 = (i / n) * o0 + (i % n); + const int64_t j1 = o0 == o1 ? j0 : (i / n) * o1 + (i % n); + + dst[i] = (T) ggml_cuda_op_swiglu_clamp_single((float) gate[j0], (float) up[j1], limit); +} + +template +static void swiglu_clamp_cuda(const T * gate, const T * up, T * dst, const int64_t k, const int64_t n, const int64_t o0, const int64_t o1, const float limit, cudaStream_t stream) { + const int64_t num_blocks = (k + CUDA_GLU_BLOCK_SIZE - 1) / CUDA_GLU_BLOCK_SIZE; + swiglu_clamp_kernel<<>>(gate, up, dst, k, n, o0, o1, limit); +} + +void ggml_cuda_op_swiglu_clamp(ggml_backend_cuda_context & ctx, ggml_tensor * dst) { + const ggml_tensor * src0 = dst->src[0]; + const ggml_tensor * src1 = dst->src[1]; + void * src0_d = src0->data; + void * src1_d = src1 ? src1->data : src0->data; + const int64_t src0_o = src0->nb[1]; + const int64_t src1_o = src1 ? src1->nb[1] : src0->nb[1]; + void * dst_d = dst->data; + const int64_t nc = src1 ? src0->ne[0] : src0->ne[0] / 2; + cudaStream_t stream = ctx.stream(); + + GGML_ASSERT(ggml_is_contiguous_1(src0)); + GGML_ASSERT(src0->nb[0] == ggml_element_size(src0)); + GGML_ASSERT(ggml_is_contiguous(dst)); + + GGML_ASSERT(src0->type == GGML_TYPE_F32 || src0->type == GGML_TYPE_F16); + GGML_ASSERT(src0->type == dst->type); + GGML_ASSERT(dst->ne[0] == nc); + GGML_ASSERT(ggml_nrows(dst) == ggml_nrows(src0)); + + if (src1) { + GGML_ASSERT(ggml_is_contiguous_1(src1)); + GGML_ASSERT(src1->nb[0] == ggml_element_size(src1)); + GGML_ASSERT(src1->ne[0] == nc); + GGML_ASSERT(src0->type == src1->type); + } + + const int32_t swapped = ggml_get_op_params_i32(dst, 1); + const float limit = ggml_get_op_params_f32(dst, 3); + + if (src0->type == GGML_TYPE_F16) { + half * src0_p = (half *) src0_d; + half * src1_p = (half *) src1_d; + + if (!src1) { + src0_p += swapped ? nc : 0; + src1_p += swapped ? 0 : nc; + } + + swiglu_clamp_cuda(src0_p, src1_p, (half *) dst_d, ggml_nelements(dst), nc, src0_o / sizeof(half), src1_o / sizeof(half), limit, stream); + } else { + float * src0_p = (float *) src0_d; + float * src1_p = (float *) src1_d; + + if (!src1) { + src0_p += swapped ? nc : 0; + src1_p += swapped ? 0 : nc; + } + + swiglu_clamp_cuda(src0_p, src1_p, (float *) dst_d, ggml_nelements(dst), nc, src0_o / sizeof(float), src1_o / sizeof(float), limit, stream); + } +} + /* CUDA kernel + launcher for xIELU */ template diff --git a/ggml/src/ggml-cuda/unary.cuh b/ggml/src/ggml-cuda/unary.cuh index 81ed873ecc3..04f3af6443a 100644 --- a/ggml/src/ggml-cuda/unary.cuh +++ b/ggml/src/ggml-cuda/unary.cuh @@ -83,6 +83,8 @@ void ggml_cuda_op_swiglu(ggml_backend_cuda_context & ctx, ggml_tensor * dst); void ggml_cuda_op_swiglu_oai(ggml_backend_cuda_context & ctx, ggml_tensor * dst); +void ggml_cuda_op_swiglu_clamp(ggml_backend_cuda_context & ctx, ggml_tensor * dst); + void ggml_cuda_op_geglu_erf(ggml_backend_cuda_context & ctx, ggml_tensor * dst); void ggml_cuda_op_geglu_quick(ggml_backend_cuda_context & ctx, ggml_tensor * dst); @@ -112,3 +114,10 @@ __device__ __forceinline__ float ggml_cuda_op_swiglu_oai_single(float x, float g out_glu = out_glu * (1.0f + g); return out_glu; } + +__device__ __forceinline__ float ggml_cuda_op_swiglu_clamp_single(float gate, float up, float limit) { + gate = fminf(gate, limit); + up = fmaxf(fminf(up, limit), -limit); + + return ggml_cuda_op_silu_single(gate) * up; +} diff --git a/ggml/src/ggml-et/et-kernels/src/glu_f32.c b/ggml/src/ggml-et/et-kernels/src/glu_f32.c index 95fe5721589..d376d6f56ff 100644 --- a/ggml/src/ggml-et/et-kernels/src/glu_f32.c +++ b/ggml/src/ggml-et/et-kernels/src/glu_f32.c @@ -17,7 +17,7 @@ struct ggml_et_glu_params { int32_t glu_op_type; // GLU operation type (REGLU=0, GEGLU=1, SWIGLU=2, etc.) int32_t swapped; // Whether gate and value are swapped float alpha; // SWIGLU_OAI: sigmoid scaling factor - float limit; // SWIGLU_OAI: clamp limit + float limit; // GLU clamp limit }; // SiLU activation function: silu(x) = x * sigmoid(x) = x / (1 + exp(-x)) @@ -332,6 +332,57 @@ static inline void block_swiglu_oai(float * dst_block, } } +static inline void block_swiglu_clamp(float * dst_block, + const float * gate_block, + const float * up_block, + int elements, + float limit) { + int32_t vec_end = (elements / 8) * 8; + + unsigned long temp_mask; + __asm__ volatile("mova.x.m %0" : "=r"(temp_mask)); + __asm__ volatile("mov.m.x m0, x0, 0xFF"); + + float one_const = 1.0f; + float limit_pos = limit; + float limit_neg = -limit; + float neg_log2e = -1.4426950408889634f; + + for (int32_t i = 0; i < vec_end; i += 8) { + __asm__ volatile( + "flw.ps f10, %[gate_vec]\n" + "flw.ps f11, %[up_vec]\n" + "fbc.ps f21, %[one_ptr]\n" + "fbc.ps f23, %[lim_pos]\n" + "fbc.ps f24, %[lim_neg]\n" + "fbc.ps f25, %[k_ptr]\n" + "fmin.ps f12, f10, f23\n" + "fmax.ps f13, f11, f24\n" + "fmin.ps f13, f13, f23\n" + "fmul.ps f14, f12, f25\n" + "fexp.ps f15, f14\n" + "fadd.ps f15, f15, f21\n" + "frcp.ps f16, f15\n" + "fmul.ps f17, f12, f16\n" + "fmul.ps f18, f17, f13\n" + "fsw.ps f18, %[dst_out]\n" + : [dst_out] "=m"(*(float (*)[8]) & dst_block[i]) + : [gate_vec] "m"(*(const float (*)[8]) & gate_block[i]), [up_vec] "m"(*(const float (*)[8]) & up_block[i]), + [one_ptr] "m"(one_const), [lim_pos] "m"(limit_pos), [lim_neg] "m"(limit_neg), [k_ptr] "m"(neg_log2e) + : "f10", "f11", "f12", "f13", "f14", "f15", "f16", "f17", "f18", "f21", "f23", "f24", "f25"); + } + + __asm__ volatile("mova.m.x %0" :: "r"(temp_mask)); + + for (int32_t i = vec_end; i < elements; i++) { + float gate = gate_block[i] > limit ? limit : gate_block[i]; + float up = up_block[i]; + up = up > limit ? limit : up; + up = up < -limit ? -limit : up; + dst_block[i] = silu_f32(gate) * up; + } +} + // Scalar erf approximation (Abramowitz & Stegun 7.1.26, max error ~1.5e-7) static inline float erf_approx(float x) { const float a1 = 0.254829592f; @@ -386,6 +437,7 @@ int entry_point(struct ggml_et_glu_params * params, void * env) { switch (params->glu_op_type) { case GGML_GLU_OP_SWIGLU: case GGML_GLU_OP_SWIGLU_OAI: + case GGML_GLU_OP_SWIGLU_CLAMP: case GGML_GLU_OP_GEGLU: case GGML_GLU_OP_GEGLU_ERF: case GGML_GLU_OP_GEGLU_QUICK: @@ -531,6 +583,9 @@ int entry_point(struct ggml_et_glu_params * params, void * env) { case GGML_GLU_OP_SWIGLU_OAI: block_swiglu_oai(dst_ptr, x_ptr, g_ptr, (int) elements_to_process, params->alpha, params->limit); break; + case GGML_GLU_OP_SWIGLU_CLAMP: + block_swiglu_clamp(dst_ptr, x_ptr, g_ptr, (int) elements_to_process, params->limit); + break; default: return -1; } diff --git a/ggml/src/ggml-et/ggml-et-cpu-compare.cpp b/ggml/src/ggml-et/ggml-et-cpu-compare.cpp index b37f6d261d9..5771679b3fa 100644 --- a/ggml/src/ggml-et/ggml-et-cpu-compare.cpp +++ b/ggml/src/ggml-et/ggml-et-cpu-compare.cpp @@ -261,7 +261,12 @@ bool ggml_et_cpu_compare_compute_and_check(ggml_et_cpu_compare_ctx * ct GGML_LOG_ERROR("ET: GLU CPU comparison requires split tensor mode\n"); return false; } - ctx->cpu_dst = ggml_glu_split(ctx->ggml_ctx, ctx->cpu_src0, ctx->cpu_src1, glu_op); + if (glu_op == GGML_GLU_OP_SWIGLU_CLAMP) { + const float limit = ggml_get_op_params_f32(node, 3); + ctx->cpu_dst = ggml_swiglu_clamp(ctx->ggml_ctx, ctx->cpu_src0, ctx->cpu_src1, limit); + } else { + ctx->cpu_dst = ggml_glu_split(ctx->ggml_ctx, ctx->cpu_src0, ctx->cpu_src1, glu_op); + } } break; case GGML_OP_SOFT_MAX: diff --git a/ggml/src/ggml-et/ggml-et-ops.cpp b/ggml/src/ggml-et/ggml-et-ops.cpp index 7871d524081..8765138672a 100644 --- a/ggml/src/ggml-et/ggml-et-ops.cpp +++ b/ggml/src/ggml-et/ggml-et-ops.cpp @@ -636,6 +636,7 @@ bool ggml_et_op_glu(ggml_backend_et_device_context * dev_ctx, const ggml_tensor case GGML_GLU_OP_GEGLU: case GGML_GLU_OP_SWIGLU: case GGML_GLU_OP_SWIGLU_OAI: + case GGML_GLU_OP_SWIGLU_CLAMP: case GGML_GLU_OP_GEGLU_ERF: case GGML_GLU_OP_GEGLU_QUICK: break; @@ -661,6 +662,8 @@ bool ggml_et_op_glu(ggml_backend_et_device_context * dev_ctx, const ggml_tensor params.limit = 0.0f; if (glu_op_type == GGML_GLU_OP_SWIGLU_OAI) { params.alpha = ggml_get_op_params_f32(node, 2); + } + if (glu_op_type == GGML_GLU_OP_SWIGLU_OAI || glu_op_type == GGML_GLU_OP_SWIGLU_CLAMP) { params.limit = ggml_get_op_params_f32(node, 3); } // Phase 1: Initialize CPU comparison context and copy source buffers (before ET kernel) diff --git a/ggml/src/ggml-et/ggml-et.cpp b/ggml/src/ggml-et/ggml-et.cpp index b87b189a57a..61c31d6f291 100644 --- a/ggml/src/ggml-et/ggml-et.cpp +++ b/ggml/src/ggml-et/ggml-et.cpp @@ -1210,7 +1210,8 @@ static bool ggml_backend_et_device_supports_op(ggml_backend_dev_t dev, const ggm // Check GLU variant - support SWIGLU, SWIGLU_OAI, GEGLU, GEGLU_ERF, GEGLU_QUICK, REGLU ggml_glu_op glu_type = ggml_get_glu_op(op); const bool supported_variant = glu_type == GGML_GLU_OP_SWIGLU || glu_type == GGML_GLU_OP_SWIGLU_OAI || - glu_type == GGML_GLU_OP_GEGLU || glu_type == GGML_GLU_OP_GEGLU_ERF || + glu_type == GGML_GLU_OP_SWIGLU_CLAMP || glu_type == GGML_GLU_OP_GEGLU || + glu_type == GGML_GLU_OP_GEGLU_ERF || glu_type == GGML_GLU_OP_GEGLU_QUICK || glu_type == GGML_GLU_OP_REGLU; if (op->src[1]) { diff --git a/ggml/src/ggml-hexagon/ggml-hexagon.cpp b/ggml/src/ggml-hexagon/ggml-hexagon.cpp index 87e6989bd2f..3eb84fd2a99 100644 --- a/ggml/src/ggml-hexagon/ggml-hexagon.cpp +++ b/ggml/src/ggml-hexagon/ggml-hexagon.cpp @@ -4701,6 +4701,7 @@ static htp_op_code op_remap_to_htp(const ggml_tensor * t) { switch (ggml_get_glu_op(t)) { case GGML_GLU_OP_SWIGLU: return HTP_OP_GLU_SWIGLU; case GGML_GLU_OP_SWIGLU_OAI: return HTP_OP_GLU_SWIGLU_OAI; + case GGML_GLU_OP_SWIGLU_CLAMP: return HTP_OP_GLU_SWIGLU_CLAMP; case GGML_GLU_OP_GEGLU: return HTP_OP_GLU_GEGLU; default: break; } @@ -5528,6 +5529,7 @@ static bool ggml_backend_hexagon_device_supports_op(ggml_backend_dev_t dev, cons switch (ggml_get_glu_op(op)) { case GGML_GLU_OP_SWIGLU: case GGML_GLU_OP_SWIGLU_OAI: + case GGML_GLU_OP_SWIGLU_CLAMP: case GGML_GLU_OP_GEGLU: supp = ggml_hexagon_supported_activations(sess, op); break; diff --git a/ggml/src/ggml-hexagon/htp/act-ops.c b/ggml/src/ggml-hexagon/htp/act-ops.c index 0a8bf84e382..ac00b447d98 100644 --- a/ggml/src/ggml-hexagon/htp/act-ops.c +++ b/ggml/src/ggml-hexagon/htp/act-ops.c @@ -180,6 +180,26 @@ static void swiglu_oai_f32(const float * restrict src0, } } +static void swiglu_clamp_f32(const float * restrict src0, + const float * restrict src1, + float * restrict dst, + const uint32_t num_rows, + const struct htp_act_context * actx) { + htp_glu_op_preamble; + const float limit = ((const float *) (actx->octx->op_params))[3]; + + for (uint32_t ib = 0; ib < num_rows; ib++) { + const uint8_t * restrict src0_ptr = (const uint8_t *) src0 + (ib * src0_row_size_aligned); + const uint8_t * restrict src1_ptr = (const uint8_t *) src1 + (ib * src1_row_size_aligned); + uint8_t * restrict dst_ptr = (uint8_t *) dst + (ib * dst_row_size_aligned); + + hvx_min_scalar_f32((uint8_t *) src0_ptr, src0_ptr, limit, nc); + hvx_clamp_scalar_f32((uint8_t *) src1_ptr, src1_ptr, -limit, limit, nc); + hvx_sigmoid_f32_aa(dst_ptr, src0_ptr, nc); + hvx_mul_mul_f32_aa(dst_ptr, src0_ptr, dst_ptr, src1_ptr, nc); + } +} + static const float GELU_COEF_A = 0.044715f; static const float SQRT_2_OVER_PI = 0.79788456080286535587989211986876f; @@ -411,6 +431,7 @@ static void geglu_f32(const float * restrict src0, DEFINE_GLU_PER_THREAD(swiglu, "swiglu-f32", swiglu_f32(src0_spad, src1_spad, dst_spad, block_size, actx)) DEFINE_GLU_PER_THREAD(swiglu_oai, "swiglu-oai-f32", swiglu_oai_f32(src0_spad, src1_spad, dst_spad, block_size, actx)) +DEFINE_GLU_PER_THREAD(swiglu_clamp, "swiglu-clamp-f32", swiglu_clamp_f32(src0_spad, src1_spad, dst_spad, block_size, actx)) DEFINE_GLU_PER_THREAD(geglu, "geglu-f32", geglu_f32(src0_spad, src1_spad, dst_spad, block_size, actx)) static int execute_op_activations_f32(struct htp_ops_context * octx) { @@ -437,6 +458,11 @@ static int execute_op_activations_f32(struct htp_ops_context * octx) { op_type = "swiglu-oai-f32"; break; + case HTP_OP_GLU_SWIGLU_CLAMP: + act_op_func = (worker_callback_t) glu_swiglu_clamp_f32_per_thread; + op_type = "swiglu-clamp-f32"; + break; + case HTP_OP_GLU_GEGLU: act_op_func = (worker_callback_t)glu_geglu_f32_per_thread; op_type = "geglu-f32"; @@ -527,7 +553,7 @@ static int execute_op_activations_f32(struct htp_ops_context * octx) { const uint8_t * data_src0 = (const uint8_t *) src0->data; const uint8_t * data_src1 = src1 ? (const uint8_t *) src1->data : NULL; - if (!src1 && (octx->op == HTP_OP_GLU_SWIGLU || octx->op == HTP_OP_GLU_SWIGLU_OAI || octx->op == HTP_OP_GLU_GEGLU)) { + if (!src1 && (octx->op == HTP_OP_GLU_SWIGLU || octx->op == HTP_OP_GLU_SWIGLU_OAI || octx->op == HTP_OP_GLU_SWIGLU_CLAMP || octx->op == HTP_OP_GLU_GEGLU)) { const int32_t swapped = octx->op_params[1]; data_src1 = data_src0; actx.src1_row_size = actx.src0_row_size; diff --git a/ggml/src/ggml-hexagon/htp/htp-ops.h b/ggml/src/ggml-hexagon/htp/htp-ops.h index e804844d599..53c95f28d32 100644 --- a/ggml/src/ggml-hexagon/htp/htp-ops.h +++ b/ggml/src/ggml-hexagon/htp/htp-ops.h @@ -96,6 +96,7 @@ enum htp_op_code { HTP_OP_FENCE, HTP_OP_ALLREDUCE, HTP_OP_ALLREDUCE_ADD, + HTP_OP_GLU_SWIGLU_CLAMP, HTP_OP_INVALID }; diff --git a/ggml/src/ggml-hexagon/htp/main.c b/ggml/src/ggml-hexagon/htp/main.c index fe7d093a81c..27d1dedcdf0 100644 --- a/ggml/src/ggml-hexagon/htp/main.c +++ b/ggml/src/ggml-hexagon/htp/main.c @@ -784,6 +784,7 @@ static int execute_op(struct htp_ops_context * octx) { case HTP_OP_GLU_SWIGLU: case HTP_OP_GLU_SWIGLU_OAI: + case HTP_OP_GLU_SWIGLU_CLAMP: case HTP_OP_GLU_GEGLU: return op_activations(octx); diff --git a/ggml/src/ggml-metal/ggml-metal-device.cpp b/ggml/src/ggml-metal/ggml-metal-device.cpp index 4e855be4467..5a2f01f1dba 100644 --- a/ggml/src/ggml-metal/ggml-metal-device.cpp +++ b/ggml/src/ggml-metal/ggml-metal-device.cpp @@ -318,6 +318,7 @@ ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_glu(ggml_metal_l case GGML_GLU_OP_SWIGLU_OAI: op_str = "swiglu_oai"; break; case GGML_GLU_OP_GEGLU_ERF: op_str = "geglu_erf"; break; case GGML_GLU_OP_GEGLU_QUICK: op_str = "geglu_quick"; break; + case GGML_GLU_OP_SWIGLU_CLAMP: op_str = "swiglu_clamp"; break; default: GGML_ABORT("fatal error"); } break; default: GGML_ABORT("fatal error"); diff --git a/ggml/src/ggml-metal/ggml-metal-device.m b/ggml/src/ggml-metal/ggml-metal-device.m index a053887a337..83344539e54 100644 --- a/ggml/src/ggml-metal/ggml-metal-device.m +++ b/ggml/src/ggml-metal/ggml-metal-device.m @@ -1510,6 +1510,7 @@ bool ggml_metal_device_supports_op(ggml_metal_device_t dev, const struct ggml_te case GGML_GLU_OP_SWIGLU_OAI: case GGML_GLU_OP_GEGLU_ERF: case GGML_GLU_OP_GEGLU_QUICK: + case GGML_GLU_OP_SWIGLU_CLAMP: return ggml_is_contiguous_1(op->src[0]) && (op->src[0]->type == GGML_TYPE_F32 || op->src[0]->type == GGML_TYPE_F16); default: return false; diff --git a/ggml/src/ggml-metal/kernels/unary.metal b/ggml/src/ggml-metal/kernels/unary.metal index 39cad0cbee5..e50a6486394 100644 --- a/ggml/src/ggml-metal/kernels/unary.metal +++ b/ggml/src/ggml-metal/kernels/unary.metal @@ -317,6 +317,32 @@ typedef decltype(kernel_swiglu_oai) kernel_swiglu_oai_t; template [[host_name("kernel_swiglu_oai_f32")]] kernel kernel_swiglu_oai_t kernel_swiglu_oai; template [[host_name("kernel_swiglu_oai_f16")]] kernel kernel_swiglu_oai_t kernel_swiglu_oai; +template +kernel void kernel_swiglu_clamp( + constant ggml_metal_kargs_glu & args, + device const char * src0, + device const char * src1, + device char * dst, + uint tgpig[[threadgroup_position_in_grid]], + uint tpitg[[thread_position_in_threadgroup]], + uint ntg[[threads_per_threadgroup]]) { + device const T * src0_row = (device const T *) ((device const char *) src0 + tgpig*args.nb01) + args.i00; + device const T * src1_row = (device const T *) ((device const char *) src1 + tgpig*args.nb11) + args.i10; + device T * dst_row = (device T *) ((device char *) dst + tgpig*args.nb1); + + for (int i0 = tpitg; i0 < args.ne0; i0 += ntg) { + const float gate = min((float) src0_row[i0], args.limit); + const float up = clamp((float) src1_row[i0], -args.limit, args.limit); + + dst_row[i0] = (T)(gate / (1.0f + exp(-gate)) * up); + } +} + +typedef decltype(kernel_swiglu_clamp) kernel_swiglu_clamp_t; + +template [[host_name("kernel_swiglu_clamp_f32")]] kernel kernel_swiglu_clamp_t kernel_swiglu_clamp; +template [[host_name("kernel_swiglu_clamp_f16")]] kernel kernel_swiglu_clamp_t kernel_swiglu_clamp; + template kernel void kernel_geglu_erf( constant ggml_metal_kargs_glu & args, diff --git a/ggml/src/ggml-opencl/ggml-opencl.cpp b/ggml/src/ggml-opencl/ggml-opencl.cpp index 426aac52316..90635cc858f 100644 --- a/ggml/src/ggml-opencl/ggml-opencl.cpp +++ b/ggml/src/ggml-opencl/ggml-opencl.cpp @@ -744,8 +744,9 @@ struct ggml_backend_opencl_context { cl_kernel kernel_tri; cl_kernel kernel_fill; cl_kernel kernel_clamp; - cl_kernel kernel_geglu, kernel_reglu, kernel_swiglu, kernel_swiglu_oai, kernel_geglu_erf, kernel_geglu_quick, - kernel_geglu_f16, kernel_reglu_f16, kernel_swiglu_f16, kernel_geglu_erf_f16, kernel_geglu_quick_f16; + cl_kernel kernel_geglu, kernel_reglu, kernel_swiglu, kernel_swiglu_oai, kernel_swiglu_clamp, kernel_geglu_erf, + kernel_geglu_quick, kernel_geglu_f16, kernel_reglu_f16, kernel_swiglu_f16, kernel_swiglu_clamp_f16, + kernel_geglu_erf_f16, kernel_geglu_quick_f16; cl_kernel kernel_norm, kernel_norm_mul_add; cl_kernel kernel_rms_norm, kernel_rms_norm_mul; cl_kernel kernel_l2_norm_f32; @@ -1601,11 +1602,13 @@ static void load_cl_kernels(ggml_backend_opencl_context *backend_ctx) { CL_CHECK((backend_ctx->kernel_reglu = clCreateKernel(backend_ctx->program_glu, "kernel_reglu", &err), err)); CL_CHECK((backend_ctx->kernel_swiglu = clCreateKernel(backend_ctx->program_glu, "kernel_swiglu", &err), err)); CL_CHECK((backend_ctx->kernel_swiglu_oai = clCreateKernel(backend_ctx->program_glu, "kernel_swiglu_oai", &err), err)); + CL_CHECK((backend_ctx->kernel_swiglu_clamp = clCreateKernel(backend_ctx->program_glu, "kernel_swiglu_clamp", &err), err)); CL_CHECK((backend_ctx->kernel_geglu_erf = clCreateKernel(backend_ctx->program_glu, "kernel_geglu_erf", &err), err)); CL_CHECK((backend_ctx->kernel_geglu_quick = clCreateKernel(backend_ctx->program_glu, "kernel_geglu_quick", &err), err)); CL_CHECK((backend_ctx->kernel_geglu_f16 = clCreateKernel(backend_ctx->program_glu, "kernel_geglu_f16", &err), err)); CL_CHECK((backend_ctx->kernel_reglu_f16 = clCreateKernel(backend_ctx->program_glu, "kernel_reglu_f16", &err), err)); CL_CHECK((backend_ctx->kernel_swiglu_f16 = clCreateKernel(backend_ctx->program_glu, "kernel_swiglu_f16", &err), err)); + CL_CHECK((backend_ctx->kernel_swiglu_clamp_f16 = clCreateKernel(backend_ctx->program_glu, "kernel_swiglu_clamp_f16", &err), err)); CL_CHECK((backend_ctx->kernel_geglu_erf_f16 = clCreateKernel(backend_ctx->program_glu, "kernel_geglu_erf_f16", &err), err)); CL_CHECK((backend_ctx->kernel_geglu_quick_f16 = clCreateKernel(backend_ctx->program_glu, "kernel_geglu_quick_f16", &err), err)); GGML_LOG_CONT("."); @@ -7700,6 +7703,7 @@ static bool ggml_opencl_supports_op(ggml_backend_dev_t dev, const struct ggml_te case GGML_GLU_OP_SWIGLU_OAI: case GGML_GLU_OP_GEGLU_ERF: case GGML_GLU_OP_GEGLU_QUICK: + case GGML_GLU_OP_SWIGLU_CLAMP: return ggml_is_contiguous_1(op->src[0]) && (op->type == GGML_TYPE_F32 || op->type == GGML_TYPE_F16); default: return false; @@ -24886,6 +24890,13 @@ static void ggml_cl_glu(ggml_backend_t backend, const ggml_tensor * src0, const case GGML_GLU_OP_SWIGLU_OAI: kernel = backend_ctx->kernel_swiglu_oai; break; + case GGML_GLU_OP_SWIGLU_CLAMP: + if (dst->type == GGML_TYPE_F32) { + kernel = backend_ctx->kernel_swiglu_clamp; + } else { + kernel = backend_ctx->kernel_swiglu_clamp_f16; + } + break; case GGML_GLU_OP_GEGLU_ERF: if (dst->type == GGML_TYPE_F32) { kernel = backend_ctx->kernel_geglu_erf; @@ -24941,8 +24952,10 @@ static void ggml_cl_glu(ggml_backend_t backend, const ggml_tensor * src0, const CL_CHECK(clSetKernelArg(kernel, 10, sizeof(int), &ne00_off)); CL_CHECK(clSetKernelArg(kernel, 11, sizeof(int), &ne10_off)); - if (ggml_get_glu_op(dst) == GGML_GLU_OP_SWIGLU_OAI) { + if (ggml_get_glu_op(dst) == GGML_GLU_OP_SWIGLU_OAI || ggml_get_glu_op(dst) == GGML_GLU_OP_SWIGLU_CLAMP) { CL_CHECK(clSetKernelArg(kernel, 12, sizeof(float), &limit)); + } + if (ggml_get_glu_op(dst) == GGML_GLU_OP_SWIGLU_OAI) { CL_CHECK(clSetKernelArg(kernel, 13, sizeof(float), &alpha)); } diff --git a/ggml/src/ggml-opencl/kernels/glu.cl b/ggml/src/ggml-opencl/kernels/glu.cl index 059a4bbf1ba..30bad00f7d0 100644 --- a/ggml/src/ggml-opencl/kernels/glu.cl +++ b/ggml/src/ggml-opencl/kernels/glu.cl @@ -243,6 +243,71 @@ kernel void kernel_swiglu_oai( } } +//------------------------------------------------------------------------------ +// swiglu_clamp +//------------------------------------------------------------------------------ +kernel void kernel_swiglu_clamp( + global char * src0, + ulong offset0, + global char * src1, + ulong offset1, + global char * dst, + ulong offsetd, + ulong nb01, + ulong nb11, + int ne0, + ulong nb1, + int ne00_off, + int ne10_off, + float limit +) { + src0 = (global char*)((global char*)src0 + offset0); + src1 = (global char*)((global char*)src1 + offset1); + dst = (global char*)((global char*)dst + offsetd); + + global float * src0_row = (global float *) ((global char *) src0 + get_group_id(0)*nb01) + ne00_off; + global float * src1_row = (global float *) ((global char *) src1 + get_group_id(0)*nb11) + ne10_off; + global float * dst_row = (global float *) ((global char *) dst + get_group_id(0)*nb1); + + for (int i0 = get_local_id(0); i0 < ne0; i0 += get_local_size(0)) { + const float gate = min(src0_row[i0], limit); + const float up = clamp(src1_row[i0], -limit, limit); + + dst_row[i0] = gate / (1.0f + exp(-gate)) * up; + } +} + +kernel void kernel_swiglu_clamp_f16( + global char * src0, + ulong offset0, + global char * src1, + ulong offset1, + global char * dst, + ulong offsetd, + ulong nb01, + ulong nb11, + int ne0, + ulong nb1, + int ne00_off, + int ne10_off, + float limit +) { + src0 = (global char*)((global char*)src0 + offset0); + src1 = (global char*)((global char*)src1 + offset1); + dst = (global char*)((global char*)dst + offsetd); + + global half * src0_row = (global half *) ((global char *) src0 + get_group_id(0)*nb01) + ne00_off; + global half * src1_row = (global half *) ((global char *) src1 + get_group_id(0)*nb11) + ne10_off; + global half * dst_row = (global half *) ((global char *) dst + get_group_id(0)*nb1); + + for (int i0 = get_local_id(0); i0 < ne0; i0 += get_local_size(0)) { + const float gate = min((float) src0_row[i0], limit); + const float up = clamp((float) src1_row[i0], -limit, limit); + + dst_row[i0] = (half) (gate / (1.0f + exp(-gate)) * up); + } +} + //------------------------------------------------------------------------------ // geglu_erf //------------------------------------------------------------------------------ diff --git a/ggml/src/ggml-openvino/openvino/op/glu_swiglu.cpp b/ggml/src/ggml-openvino/openvino/op/glu_swiglu.cpp index d220f2f584a..d81fc53b5d0 100644 --- a/ggml/src/ggml-openvino/openvino/op/glu_swiglu.cpp +++ b/ggml/src/ggml-openvino/openvino/op/glu_swiglu.cpp @@ -89,6 +89,21 @@ OutputVector translate_glu_swiglu_oai(const NodeContext & context) { return rename_outputs_with_suffix({res}, context.get_name()); } +OutputVector translate_glu_swiglu_clamp(const NodeContext & context) { + auto [src0, src1] = get_glu_inputs(context); + + const int32_t * params = context.get_output_op_params(); + const float limit = reinterpret_cast(params)[3]; + + auto gate = std::make_shared(src0, -std::numeric_limits::infinity(), limit); + auto sigmoid = std::make_shared(gate); + auto silu = std::make_shared(gate, sigmoid); + auto up = std::make_shared(src1, -limit, limit); + auto res = std::make_shared(silu, up); + + return rename_outputs_with_suffix({res}, context.get_name()); +} + } // namespace op } // namespace ggml } // namespace frontend diff --git a/ggml/src/ggml-openvino/openvino/op_table.cpp b/ggml/src/ggml-openvino/openvino/op_table.cpp index 9c9d8eeac78..d4f5ac30732 100644 --- a/ggml/src/ggml-openvino/openvino/op_table.cpp +++ b/ggml/src/ggml-openvino/openvino/op_table.cpp @@ -60,6 +60,7 @@ std::unordered_map get_supported_ops() { {"GGML_OP_VIEW", op::translate_view }, {"GGML_GLU_OP_SWIGLU", op::translate_glu_swiglu }, {"GGML_GLU_OP_SWIGLU_OAI", op::translate_glu_swiglu_oai }, + {"GGML_GLU_OP_SWIGLU_CLAMP", op::translate_glu_swiglu_clamp }, {"GGML_GLU_OP_GEGLU", op::translate_glu_geglu }, {"GGML_GLU_OP_GEGLU_QUICK", op::translate_glu_geglu_quick }, {"GGML_OP_SET_ROWS", op::translate_set_rows }, diff --git a/ggml/src/ggml-openvino/openvino/op_table.h b/ggml/src/ggml-openvino/openvino/op_table.h index 0a81a57a667..a0a42bff337 100644 --- a/ggml/src/ggml-openvino/openvino/op_table.h +++ b/ggml/src/ggml-openvino/openvino/op_table.h @@ -37,6 +37,7 @@ GGML_OP_CONVERTER(translate_transpose); GGML_OP_CONVERTER(translate_view); GGML_OP_CONVERTER(translate_glu_swiglu); GGML_OP_CONVERTER(translate_glu_swiglu_oai); +GGML_OP_CONVERTER(translate_glu_swiglu_clamp); GGML_OP_CONVERTER(translate_glu_geglu); GGML_OP_CONVERTER(translate_glu_geglu_quick); GGML_OP_CONVERTER(translate_set_rows); diff --git a/ggml/src/ggml-sycl/element_wise.cpp b/ggml/src/ggml-sycl/element_wise.cpp index 95914873e5a..2e926abea7c 100644 --- a/ggml/src/ggml-sycl/element_wise.cpp +++ b/ggml/src/ggml-sycl/element_wise.cpp @@ -1132,6 +1132,102 @@ void ggml_sycl_op_swiglu_oai(ggml_backend_sycl_context & ctx, ggml_tensor * dst) swiglu_oai_sycl(src0_p, src1_p, (float *)dst_d, ggml_nelements(dst), nc, src0_o / sizeof(float), src1_o / sizeof(float), alpha, limit, stream); } +template +static void swiglu_clamp_kernel(const T * gate, + const T * up, + T * dst, + const int64_t k, + const int64_t n, + const int64_t o0, + const int64_t o1, + float limit, + sycl::nd_item<3> item_ct1) { + const int64_t i = int64_t(item_ct1.get_local_range(2)) * item_ct1.get_group(2) + item_ct1.get_local_id(2); + + if (i >= k) { + return; + } + + const int64_t j0 = (i / n) * o0 + (i % n); + const int64_t j1 = o0 == o1 ? j0 : (i / n) * o1 + (i % n); + + const float gate_value = sycl::fmin((float) gate[j0], limit); + const float up_value = sycl::fmax(sycl::fmin((float) up[j1], limit), -limit); + dst[i] = (T) (gate_value / (1.0f + sycl::native::exp(-gate_value)) * up_value); +} + +template +static void swiglu_clamp_sycl(const T * gate, + const T * up, + T * dst, + const int64_t k, + const int64_t n, + const int64_t o0, + const int64_t o1, + float limit, + dpct::queue_ptr stream) { + const int64_t num_blocks = (k + SYCL_GLU_BLOCK_SIZE - 1) / SYCL_GLU_BLOCK_SIZE; + stream->parallel_for(sycl::nd_range<3>(sycl::range<3>(1, 1, num_blocks) * sycl::range<3>(1, 1, SYCL_GLU_BLOCK_SIZE), + sycl::range<3>(1, 1, SYCL_GLU_BLOCK_SIZE)), + [=](sycl::nd_item<3> item_ct1) [[sycl::reqd_sub_group_size(WARP_SIZE)]] { + swiglu_clamp_kernel(gate, up, dst, k, n, o0, o1, limit, item_ct1); + }); +} + +static void ggml_sycl_op_swiglu_clamp(ggml_backend_sycl_context & ctx, ggml_tensor * dst) { + const ggml_tensor * src0 = dst->src[0]; + const ggml_tensor * src1 = dst->src[1]; + void * src0_d = src0->data; + void * src1_d = src1 ? src1->data : src0->data; + const int64_t src0_o = src0->nb[1]; + const int64_t src1_o = src1 ? src1->nb[1] : src0->nb[1]; + void * dst_d = dst->data; + const int64_t nc = src1 ? src0->ne[0] : src0->ne[0] / 2; + dpct::queue_ptr stream = ctx.stream(); + + GGML_ASSERT(ggml_is_contiguous_1(src0)); + GGML_ASSERT(src0->nb[0] == ggml_element_size(src0)); + GGML_ASSERT(ggml_is_contiguous(dst)); + GGML_ASSERT(src0->type == GGML_TYPE_F32 || src0->type == GGML_TYPE_F16); + GGML_ASSERT(src0->type == dst->type); + GGML_ASSERT(dst->ne[0] == nc); + GGML_ASSERT(ggml_nrows(dst) == ggml_nrows(src0)); + + if (src1) { + GGML_ASSERT(ggml_is_contiguous_1(src1)); + GGML_ASSERT(src1->nb[0] == ggml_element_size(src1)); + GGML_ASSERT(src1->ne[0] == nc); + GGML_ASSERT(src0->type == src1->type); + } + + const int32_t swapped = ggml_get_op_params_i32(dst, 1); + const float limit = ggml_get_op_params_f32(dst, 3); + + if (src0->type == GGML_TYPE_F16) { + sycl::half * src0_p = (sycl::half *) src0_d; + sycl::half * src1_p = (sycl::half *) src1_d; + + if (!src1) { + src0_p += swapped ? nc : 0; + src1_p += swapped ? 0 : nc; + } + + swiglu_clamp_sycl(src0_p, src1_p, (sycl::half *) dst_d, ggml_nelements(dst), nc, src0_o / sizeof(sycl::half), + src1_o / sizeof(sycl::half), limit, stream); + } else { + float * src0_p = (float *) src0_d; + float * src1_p = (float *) src1_d; + + if (!src1) { + src0_p += swapped ? nc : 0; + src1_p += swapped ? 0 : nc; + } + + swiglu_clamp_sycl(src0_p, src1_p, (float *) dst_d, ggml_nelements(dst), nc, src0_o / sizeof(float), + src1_o / sizeof(float), limit, stream); + } +} + static inline void ggml_sycl_op_geglu_erf(ggml_backend_sycl_context & ctx, ggml_tensor * dst) { ggml_sycl_detail::ggml_sycl_op_unary_gated(ctx, dst, [](auto x) { return op_gelu_erf(x); @@ -1295,6 +1391,11 @@ void ggml_sycl_swiglu_oai(ggml_backend_sycl_context & ctx, ggml_tensor * dst) { ggml_sycl_op_swiglu_oai(ctx, dst); } +void ggml_sycl_swiglu_clamp(ggml_backend_sycl_context & ctx, ggml_tensor * dst) { + scope_op_debug_print scope_dbg_print(__func__, dst, /*num_src=*/1); + ggml_sycl_op_swiglu_clamp(ctx, dst); +} + void ggml_sycl_geglu_erf(ggml_backend_sycl_context & ctx, ggml_tensor * dst) { scope_op_debug_print scope_dbg_print(__func__, dst, /*num_src=*/1); ggml_sycl_op_geglu_erf(ctx, dst); diff --git a/ggml/src/ggml-sycl/element_wise.hpp b/ggml/src/ggml-sycl/element_wise.hpp index 67bf422d2f3..d280066efb4 100644 --- a/ggml/src/ggml-sycl/element_wise.hpp +++ b/ggml/src/ggml-sycl/element_wise.hpp @@ -77,6 +77,7 @@ void ggml_sycl_silu(ggml_backend_sycl_context & ctx, ggml_tensor * dst); void ggml_sycl_gelu_quick(ggml_backend_sycl_context & ctx, ggml_tensor * dst); void ggml_sycl_swiglu_oai(ggml_backend_sycl_context & ctx, ggml_tensor * dst); +void ggml_sycl_swiglu_clamp(ggml_backend_sycl_context & ctx, ggml_tensor * dst); void ggml_sycl_gelu_erf(ggml_backend_sycl_context & ctx, ggml_tensor * dst); diff --git a/ggml/src/ggml-sycl/ggml-sycl.cpp b/ggml/src/ggml-sycl/ggml-sycl.cpp index d58ffd00daf..2b2d26cf235 100644 --- a/ggml/src/ggml-sycl/ggml-sycl.cpp +++ b/ggml/src/ggml-sycl/ggml-sycl.cpp @@ -5373,6 +5373,9 @@ static bool ggml_sycl_compute_forward(ggml_backend_sycl_context & ctx, struct gg case GGML_GLU_OP_SWIGLU_OAI: ggml_sycl_swiglu_oai(ctx, dst); break; + case GGML_GLU_OP_SWIGLU_CLAMP: + ggml_sycl_swiglu_clamp(ctx, dst); + break; case GGML_GLU_OP_GEGLU_ERF: ggml_sycl_geglu_erf(ctx, dst); break; @@ -6133,6 +6136,7 @@ static bool do_ggml_backend_sycl_device_supports_op(ggml_backend_dev_t dev, cons case GGML_GLU_OP_SWIGLU_OAI: case GGML_GLU_OP_GEGLU_ERF: case GGML_GLU_OP_GEGLU_QUICK: + case GGML_GLU_OP_SWIGLU_CLAMP: return ggml_is_contiguous_1(op->src[0]); default: return false; diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp index 8fbb1359f40..394c84bf257 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp +++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp @@ -1035,6 +1035,7 @@ struct vk_device_struct { vk_pipeline pipeline_reglu[2]; vk_pipeline pipeline_swiglu[2]; vk_pipeline pipeline_swiglu_oai[2]; + vk_pipeline pipeline_swiglu_clamp[2]; vk_pipeline pipeline_geglu_erf[2]; vk_pipeline pipeline_geglu_quick[2]; @@ -5748,6 +5749,7 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { CREATE_GLU(reglu) CREATE_GLU(swiglu) CREATE_GLU(swiglu_oai) + CREATE_GLU(swiglu_clamp) CREATE_GLU(geglu_erf) CREATE_GLU(geglu_quick) #undef CREATE_GLU @@ -11578,6 +11580,8 @@ static vk_pipeline ggml_vk_op_get_pipeline(ggml_backend_vk_context * ctx, const return ctx->device->pipeline_swiglu[dst->type == GGML_TYPE_F16]; case GGML_GLU_OP_SWIGLU_OAI: return ctx->device->pipeline_swiglu_oai[dst->type == GGML_TYPE_F16]; + case GGML_GLU_OP_SWIGLU_CLAMP: + return ctx->device->pipeline_swiglu_clamp[dst->type == GGML_TYPE_F16]; case GGML_GLU_OP_GEGLU_ERF: return ctx->device->pipeline_geglu_erf[dst->type == GGML_TYPE_F16]; case GGML_GLU_OP_GEGLU_QUICK: @@ -15883,6 +15887,7 @@ static bool ggml_vk_build_graph(ggml_backend_vk_context * ctx, ggml_cgraph * cgr case GGML_GLU_OP_SWIGLU_OAI: case GGML_GLU_OP_GEGLU_ERF: case GGML_GLU_OP_GEGLU_QUICK: + case GGML_GLU_OP_SWIGLU_CLAMP: ggml_vk_glu(ctx, compute_ctx, src0, src1, node); break; default: @@ -18400,6 +18405,7 @@ static bool ggml_backend_vk_device_supports_op(ggml_backend_dev_t dev, const ggm case GGML_GLU_OP_SWIGLU_OAI: case GGML_GLU_OP_GEGLU_ERF: case GGML_GLU_OP_GEGLU_QUICK: + case GGML_GLU_OP_SWIGLU_CLAMP: return (op->src[0]->type == GGML_TYPE_F32 || op->src[0]->type == GGML_TYPE_F16) && (op->type == GGML_TYPE_F32 || op->type == GGML_TYPE_F16) && (op->src[0]->type == op->type) && diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/swiglu_clamp.comp b/ggml/src/ggml-vulkan/vulkan-shaders/swiglu_clamp.comp new file mode 100644 index 00000000000..dfe329759c7 --- /dev/null +++ b/ggml/src/ggml-vulkan/vulkan-shaders/swiglu_clamp.comp @@ -0,0 +1,12 @@ +#version 450 + +#include "glu_head.glsl" + +float op(float a, float b) { + float gate = min(a, p.limit); + float up = clamp(b, -p.limit, p.limit); + + return gate / (1.0f + exp(-gate)) * up; +} + +#include "glu_main.glsl" diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp b/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp index d375c2d1277..f0610f4cd82 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp @@ -986,6 +986,8 @@ void process_shaders() { string_to_spv("swiglu_f32", "swiglu.comp", {{"A_TYPE", "float"}, {"D_TYPE", "float"}}); string_to_spv("swiglu_oai_f16", "swiglu_oai.comp", {{"A_TYPE", "float16_t"}, {"D_TYPE", "float16_t"}}); string_to_spv("swiglu_oai_f32", "swiglu_oai.comp", {{"A_TYPE", "float"}, {"D_TYPE", "float"}}); + string_to_spv("swiglu_clamp_f16", "swiglu_clamp.comp", {{"A_TYPE", "float16_t"}, {"D_TYPE", "float16_t"}}); + string_to_spv("swiglu_clamp_f32", "swiglu_clamp.comp", {{"A_TYPE", "float"}, {"D_TYPE", "float"}}); string_to_spv("geglu_erf_f16", "geglu_erf.comp", {{"A_TYPE", "float16_t"}, {"D_TYPE", "float16_t"}}); string_to_spv("geglu_erf_f32", "geglu_erf.comp", {{"A_TYPE", "float"}, {"D_TYPE", "float"}}); string_to_spv("geglu_quick_f16","geglu_quick.comp", {{"A_TYPE", "float16_t"}, {"D_TYPE", "float16_t"}}); diff --git a/ggml/src/ggml-webgpu/ggml-webgpu-shader-lib.hpp b/ggml/src/ggml-webgpu/ggml-webgpu-shader-lib.hpp index 7a67ccf4fcb..a7ff36030fa 100644 --- a/ggml/src/ggml-webgpu/ggml-webgpu-shader-lib.hpp +++ b/ggml/src/ggml-webgpu/ggml-webgpu-shader-lib.hpp @@ -3101,6 +3101,10 @@ class ggml_webgpu_shader_lib { defines.push_back("OP_GEGLU_QUICK"); variant += "_geglu_quick"; break; + case GGML_GLU_OP_SWIGLU_CLAMP: + defines.push_back("OP_SWIGLU_CLAMP"); + variant += "_swiglu_clamp"; + break; default: GGML_ABORT("Unsupported GLU op"); } diff --git a/ggml/src/ggml-webgpu/ggml-webgpu.cpp b/ggml/src/ggml-webgpu/ggml-webgpu.cpp index 2434848a55a..b953118a7a3 100644 --- a/ggml/src/ggml-webgpu/ggml-webgpu.cpp +++ b/ggml/src/ggml-webgpu/ggml-webgpu.cpp @@ -2835,7 +2835,7 @@ static webgpu_encoded_op ggml_webgpu_glu(webgpu_context & ctx, (uint32_t) dst->ne[2], (uint32_t) ((int32_t *) dst->op_params)[1], // swapped ggml_webgpu_u32_from_f32(ggml_get_op_params_f32(dst, 2)), // alpha, for swiglu_oai - ggml_webgpu_u32_from_f32(ggml_get_op_params_f32(dst, 3)), // limit, for swiglu_oai + ggml_webgpu_u32_from_f32(ggml_get_op_params_f32(dst, 3)), // limit }; std::vector entries; @@ -4483,6 +4483,7 @@ static bool ggml_backend_webgpu_device_supports_op(ggml_backend_dev_t dev, const case GGML_GLU_OP_SWIGLU: case GGML_GLU_OP_GEGLU_ERF: case GGML_GLU_OP_GEGLU_QUICK: + case GGML_GLU_OP_SWIGLU_CLAMP: supports_op = op->type == GGML_TYPE_F32 || op->type == GGML_TYPE_F16; break; case GGML_GLU_OP_SWIGLU_OAI: diff --git a/ggml/src/ggml-webgpu/wgsl-shaders/glu.wgsl b/ggml/src/ggml-webgpu/wgsl-shaders/glu.wgsl index d03f1c207d9..6bbed5d3bfe 100644 --- a/ggml/src/ggml-webgpu/wgsl-shaders/glu.wgsl +++ b/ggml/src/ggml-webgpu/wgsl-shaders/glu.wgsl @@ -37,6 +37,14 @@ fn op(a: f32, b: f32) -> f32 { return out_glu; } #endif +#ifdef OP_SWIGLU_CLAMP +fn op(a: DataType, b: DataType) -> DataType { + let limit = DataType(params.limit); + let gate = min(a, limit); + let up = clamp(b, -limit, limit); + return gate / (1.0 + exp(-gate)) * up; +} +#endif #ifdef OP_GEGLU_ERF const p_erf: DataType = 0.3275911; const a1_erf: DataType = 0.254829592; diff --git a/ggml/src/ggml.c b/ggml/src/ggml.c index e0b615c07ed..3bd3e3fe5ea 100644 --- a/ggml/src/ggml.c +++ b/ggml/src/ggml.c @@ -1253,10 +1253,10 @@ static const char * GGML_GLU_OP_NAME[GGML_GLU_OP_COUNT] = { "SWIGLU_OAI", "GEGLU_ERF", "GEGLU_QUICK", + "SWIGLU_CLAMP", }; -static_assert(GGML_GLU_OP_COUNT == 6, "GGML_GLU_OP_COUNT != 6"); - +static_assert(GGML_GLU_OP_COUNT == 7, "GGML_GLU_OP_COUNT != 7"); static_assert(sizeof(struct ggml_object)%GGML_MEM_ALIGN == 0, "ggml_object size must be a multiple of GGML_MEM_ALIGN"); static_assert(sizeof(struct ggml_tensor)%GGML_MEM_ALIGN == 0, "ggml_tensor size must be a multiple of GGML_MEM_ALIGN"); @@ -3119,6 +3119,17 @@ struct ggml_tensor * ggml_swiglu_oai( return result; } +struct ggml_tensor * ggml_swiglu_clamp( + struct ggml_context * ctx, + struct ggml_tensor * a, + struct ggml_tensor * b, + float limit) { + struct ggml_tensor * result = ggml_glu_impl(ctx, a, b, GGML_GLU_OP_SWIGLU_CLAMP, false); + ggml_set_op_params_f32(result, 3, limit); + + return result; +} + // ggml_norm static struct ggml_tensor * ggml_norm_impl( From 749683d30f40845d14f0da8d7df64f219c6db1c5 Mon Sep 17 00:00:00 2001 From: Georgi Gerganov Date: Sun, 30 Aug 2026 20:25:15 +0300 Subject: [PATCH 046/104] ggml : fix ggml_backend_buft_get_alloc_size() guard (llama/28038) --- ggml/src/ggml-backend.cpp | 1 + ggml/src/ggml-cuda/ggml-cuda.cu | 1 + 2 files changed, 2 insertions(+) diff --git a/ggml/src/ggml-backend.cpp b/ggml/src/ggml-backend.cpp index fec7d7c92bf..8f40fb2f7bc 100644 --- a/ggml/src/ggml-backend.cpp +++ b/ggml/src/ggml-backend.cpp @@ -70,6 +70,7 @@ size_t ggml_backend_buft_get_alloc_size(ggml_backend_buffer_type_t buft, const s // if you hit this assert, update ggml_backend_op_alloc_size_may_expand() accordingly GGML_ASSERT(size <= ggml_nbytes(tensor) || ggml_op_is_empty(tensor->op) || + ggml_is_quantized(tensor->type) || // [TAG_ALLOC_SIZE_EXPAND] ggml_backend_op_alloc_size_may_expand(tensor->op)); return size; diff --git a/ggml/src/ggml-cuda/ggml-cuda.cu b/ggml/src/ggml-cuda/ggml-cuda.cu index 7c2d3cad20a..a6fc655c41c 100644 --- a/ggml/src/ggml-cuda/ggml-cuda.cu +++ b/ggml/src/ggml-cuda/ggml-cuda.cu @@ -915,6 +915,7 @@ static size_t ggml_backend_cuda_buffer_type_get_alloc_size(ggml_backend_buffer_t : ggml_nbytes(tensor); int64_t ne0 = tensor->ne[0]; + // [TAG_ALLOC_SIZE_EXPAND] if (ggml_is_quantized(tensor->type)) { if (ne0 % MATRIX_ROW_PADDING != 0) { GGML_ASSERT(tensor->nb[0] == ggml_element_size(tensor)); From 4089fa628a00d026b8090616fc4d595ab36ccbd1 Mon Sep 17 00:00:00 2001 From: hmirin Date: Mon, 31 Aug 2026 02:26:16 +0900 Subject: [PATCH 047/104] rpc: avoid serializing buffers from other servers (llama/26500) * rpc: avoid serializing buffers from other servers Only include remote buffer pointers when the buffer belongs to the RPC dispatcher receiving the graph. Add a two-server regression test for cross-server tensor serialization. Assisted-by: Codex * cont : add ref --------- Co-authored-by: Georgi Gerganov --- ggml/src/ggml-rpc/ggml-rpc.cpp | 26 ++++++++++++++++---------- 1 file changed, 16 insertions(+), 10 deletions(-) diff --git a/ggml/src/ggml-rpc/ggml-rpc.cpp b/ggml/src/ggml-rpc/ggml-rpc.cpp index 58a8a030cfa..a97db24e624 100644 --- a/ggml/src/ggml-rpc/ggml-rpc.cpp +++ b/ggml/src/ggml-rpc/ggml-rpc.cpp @@ -625,7 +625,7 @@ static bool ggml_backend_buffer_is_rpc(ggml_backend_buffer_t buffer) { return buffer->iface.free_buffer == ggml_backend_rpc_buffer_free_buffer; } -static rpc_tensor serialize_tensor(const ggml_tensor * tensor) { +static rpc_tensor serialize_tensor(const ggml_tensor * tensor, const std::shared_ptr & dispatcher = nullptr) { rpc_tensor result; if (!tensor) { memset(&result, 0, sizeof(result)); @@ -637,8 +637,14 @@ static rpc_tensor serialize_tensor(const ggml_tensor * tensor) { if (tensor->buffer && ggml_backend_buffer_is_rpc(tensor->buffer)) { ggml_backend_buffer_t buffer = tensor->buffer; ggml_backend_rpc_buffer_context * ctx = (ggml_backend_rpc_buffer_context *)buffer->context; - result.buffer = ctx != nullptr ? ctx->remote_ptr : 0; - result.data = reinterpret_cast(tensor->data); + // ref: https://github.com/ggml-org/llama.cpp/pull/26500 + if (ctx != nullptr && (dispatcher == nullptr || ctx->dispatcher == dispatcher)) { + result.buffer = ctx->remote_ptr; + result.data = reinterpret_cast(tensor->data); + } else { + result.buffer = 0; + result.data = 0; + } } else { result.buffer = 0; result.data = 0; @@ -958,7 +964,7 @@ static void ggml_backend_rpc_synchronize(ggml_backend_t backend) { rpc_ctx->dispatcher->synchronize(); } -static void add_tensor(ggml_tensor * tensor, const ggml_cgraph * cgraph, std::vector & tensors, std::unordered_set & visited) { +static void add_tensor(ggml_tensor * tensor, const ggml_cgraph * cgraph, const std::shared_ptr & dispatcher, std::vector & tensors, std::unordered_set & visited) { if (tensor == nullptr) { return; } @@ -967,10 +973,10 @@ static void add_tensor(ggml_tensor * tensor, const ggml_cgraph * cgraph, std::ve } visited.insert(tensor); for (int i = 0; i < GGML_MAX_SRC; i++) { - add_tensor(tensor->src[i], cgraph, tensors, visited); + add_tensor(tensor->src[i], cgraph, dispatcher, tensors, visited); } - add_tensor(tensor->view_src, cgraph, tensors, visited); - rpc_tensor result = serialize_tensor(tensor); + add_tensor(tensor->view_src, cgraph, dispatcher, tensors, visited); + rpc_tensor result = serialize_tensor(tensor, dispatcher); const size_t hash_pos = ggml_hash_find(&cgraph->visited_hash_set, tensor); if (hash_pos != GGML_HASHSET_FULL && ggml_bitset_get(cgraph->visited_hash_set.used, hash_pos)) { result.use_count = cgraph->use_counts[hash_pos]; @@ -978,12 +984,12 @@ static void add_tensor(ggml_tensor * tensor, const ggml_cgraph * cgraph, std::ve tensors.push_back(result); } -static uint8_t * serialize_graph(uint32_t device, const ggml_cgraph * cgraph, size_t * output_size) { +static uint8_t * serialize_graph(uint32_t device, const ggml_cgraph * cgraph, const std::shared_ptr & dispatcher, size_t * output_size) { uint32_t n_nodes = cgraph->n_nodes; std::vector tensors; std::unordered_set visited; for (uint32_t i = 0; i < n_nodes; i++) { - add_tensor(cgraph->nodes[i], cgraph, tensors, visited); + add_tensor(cgraph->nodes[i], cgraph, dispatcher, tensors, visited); } // serialization format: // | device (4 bytes) | n_nodes (4 bytes) | nodes (n_nodes * sizeof(uint64_t) | n_tensors (4 bytes) | tensors (n_tensors * sizeof(rpc_tensor)) | @@ -1020,7 +1026,7 @@ static enum ggml_status ggml_backend_rpc_graph_compute(ggml_backend_t backend, g } else { rpc_dev_ctx->last_graph_uid = cgraph->uid; size_t input_size = 0; - uint8_t * input = serialize_graph(rpc_ctx->device, cgraph, &input_size); + uint8_t * input = serialize_graph(rpc_ctx->device, cgraph, rpc_ctx->dispatcher, &input_size); std::shared_ptr input_ptr(input, std::default_delete()); rpc_ctx->dispatcher->send_async(RPC_CMD_GRAPH_COMPUTE, input_ptr, input_size); } From e5c96ca45da0d38b0d4a0967e83ec413e530b3d9 Mon Sep 17 00:00:00 2001 From: codemonkey <441345965@qq.com> Date: Mon, 31 Aug 2026 02:00:10 +0800 Subject: [PATCH 048/104] metal : add remaining Q4_1/Q5_0/Q5_1 fa-vec tunings for M2 (llama/28017) --- ggml/src/ggml-metal/ggml-metal-tuning.cpp | 113 ++++++++++++++++++++++ 1 file changed, 113 insertions(+) diff --git a/ggml/src/ggml-metal/ggml-metal-tuning.cpp b/ggml/src/ggml-metal/ggml-metal-tuning.cpp index 9114ab7424d..4f26fd9c373 100644 --- a/ggml/src/ggml-metal/ggml-metal-tuning.cpp +++ b/ggml/src/ggml-metal/ggml-metal-tuning.cpp @@ -516,6 +516,119 @@ constexpr fa_vec_entry_t fa_vec_tuned_table[] = { { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q4_0, 512, 512, 3, 3 }, { 2, 2 } }, { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q4_0, 576, 512, -1, 0 }, { 1, 4 } }, { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q4_0, 576, 512, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q4_1, 32, 32, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q4_1, 32, 32, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q4_1, 32, 32, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q4_1, 32, 32, 2, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q4_1, 32, 32, 3, 2 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q4_1, 32, 32, 3, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q4_1, 64, 64, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q4_1, 64, 64, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q4_1, 64, 64, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q4_1, 96, 96, 1, 3 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q4_1, 96, 96, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q4_1, 96, 96, 2, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q4_1, 96, 96, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q4_1, 96, 96, 3, 3 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q4_1, 128, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q4_1, 128, 128, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q4_1, 128, 128, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q4_1, 128, 128, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q4_1, 128, 128, 3, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q4_1, 192, 192, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q4_1, 192, 192, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q4_1, 192, 192, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q4_1, 192, 192, 1, 4 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q4_1, 192, 192, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q4_1, 192, 192, 3, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q4_1, 192, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q4_1, 192, 128, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q4_1, 256, 256, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q4_1, 320, 256, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q4_1, 320, 256, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q4_1, 512, 512, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q4_1, 512, 512, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q4_1, 576, 512, 1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q4_1, 576, 512, 1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q4_1, 576, 512, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q4_1, 576, 512, 1, 3 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q4_1, 576, 512, 1, 4 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q5_0, 32, 32, -1, 1 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q5_0, 32, 32, 1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q5_0, 32, 32, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q5_0, 32, 32, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q5_0, 64, 64, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q5_0, 64, 64, -1, 1 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q5_0, 64, 64, 1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q5_0, 64, 64, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q5_0, 64, 64, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q5_0, 96, 96, -1, 1 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q5_0, 96, 96, 1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q5_0, 96, 96, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q5_0, 96, 96, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q5_0, 96, 96, 3, 4 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q5_0, 128, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q5_0, 128, 128, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q5_0, 128, 128, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q5_0, 128, 128, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q5_0, 128, 128, 3, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q5_0, 192, 192, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q5_0, 192, 192, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q5_0, 192, 192, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q5_0, 192, 192, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q5_0, 192, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q5_0, 192, 128, -1, 1 }, { 4, 2 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q5_0, 192, 128, 1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q5_0, 192, 128, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q5_0, 192, 128, 1, 3 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q5_0, 192, 128, 1, 4 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q5_0, 192, 128, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q5_0, 192, 128, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q5_0, 256, 256, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q5_0, 256, 256, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q5_0, 320, 256, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q5_0, 320, 256, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q5_0, 512, 512, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q5_0, 512, 512, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q5_0, 576, 512, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q5_0, 576, 512, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q5_1, 32, 32, -1, 1 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q5_1, 32, 32, 1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q5_1, 32, 32, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q5_1, 32, 32, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q5_1, 64, 64, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q5_1, 64, 64, -1, 1 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q5_1, 64, 64, 1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q5_1, 64, 64, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q5_1, 64, 64, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q5_1, 96, 96, -1, 1 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q5_1, 96, 96, 1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q5_1, 96, 96, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q5_1, 96, 96, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q5_1, 128, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q5_1, 128, 128, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q5_1, 128, 128, 3, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q5_1, 192, 192, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q5_1, 192, 192, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q5_1, 192, 192, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q5_1, 192, 192, 1, 4 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q5_1, 192, 192, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q5_1, 192, 192, 3, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q5_1, 192, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q5_1, 192, 128, -1, 1 }, { 4, 2 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q5_1, 192, 128, 1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q5_1, 192, 128, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q5_1, 192, 128, 1, 4 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q5_1, 192, 128, 2, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q5_1, 192, 128, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q5_1, 256, 256, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q5_1, 256, 256, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q5_1, 320, 256, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q5_1, 320, 256, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q5_1, 512, 512, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q5_1, 512, 512, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q5_1, 576, 512, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q5_1, 576, 512, -1, 1 }, { 1, 4 } }, { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q8_0, 32, 32, -1, 1 }, { 2, 4 } }, { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q8_0, 32, 32, 1, 2 }, { 1, 4 } }, { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q8_0, 32, 32, 2, 2 }, { 1, 4 } }, From 01ebd225a7784bfb8245b98f2ebd8a4589f32421 Mon Sep 17 00:00:00 2001 From: Shenghan Yang Date: Mon, 31 Aug 2026 02:18:24 +0800 Subject: [PATCH 049/104] hexagon: fix CPY fence bug (llama/28033) --- ggml/src/ggml-hexagon/htp/cpy-ops.c | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/ggml/src/ggml-hexagon/htp/cpy-ops.c b/ggml/src/ggml-hexagon/htp/cpy-ops.c index 15bc8dc244f..c945425dab3 100644 --- a/ggml/src/ggml-hexagon/htp/cpy-ops.c +++ b/ggml/src/ggml-hexagon/htp/cpy-ops.c @@ -330,7 +330,7 @@ int op_cpy(struct htp_ops_context * octx) { } const struct htp_tensor *sync = octx->src[1]; - if (sync) { + if (sync && (sync->flags & HTP_TENSOR_FENCE)) { if (!use_dma) { // htp_tensor_flush_all(octx->ctx, octx->dsts, 1); qurt_mem_cache_clean((qurt_addr_t) 0, 0, QURT_MEM_CACHE_FLUSH_INVALIDATE_ALL, QURT_MEM_DCACHE); From db00b0196b1be6af1bf184c8948a50fc3120fcb5 Mon Sep 17 00:00:00 2001 From: Ruben Ortlam Date: Mon, 31 Aug 2026 07:04:34 +0200 Subject: [PATCH 050/104] vulkan: top_k radix select for k >= 1024 for Qwen 3.8 Flash Next (llama/28032) * vulkan: add top-k radix sort shader for k >= 1024 * add Qwen 3.8 Flash Next top-k tests * add top-k qsa fusion * clean up code --- ggml/src/ggml-vulkan/ggml-vulkan.cpp | 239 +++++++++++++++++- .../vulkan-shaders/topk_radix_select.comp | 144 +++++++++++ .../vulkan-shaders/vulkan-shaders-gen.cpp | 1 + 3 files changed, 374 insertions(+), 10 deletions(-) create mode 100644 ggml/src/ggml-vulkan/vulkan-shaders/topk_radix_select.comp diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp index 394c84bf257..649ec792062 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp +++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp @@ -657,6 +657,21 @@ static constexpr std::initializer_list snake_pattern { GGM GGML_OP_SQR, GGML_OP_MUL, GGML_OP_ADD }; +// qwen4 QSA indexer: gather per-block scores to cells + add f16 mask (cast+reshape) + top-k, +// fused into one radix-select. The cast/reshape are elided; the raw f16 mask is read in-shader. +static constexpr std::initializer_list topk_qsa_pattern { GGML_OP_GET_ROWS, GGML_OP_PERMUTE, + GGML_OP_CONT, GGML_OP_CPY, + GGML_OP_RESHAPE, GGML_OP_ADD, + GGML_OP_TOP_K }; +static constexpr std::initializer_list> topk_qsa_edges { + { 1, 0, 0 }, // permute->src[0] == get_rows + { 2, 0, 1 }, // cont->src[0] == permute + { 4, 0, 3 }, // reshape->src[0] == cpy (mask cast) + { 5, 0, 2 }, // add->src[0] == cont + { 5, 1, 4 }, // add->src[1] == reshape + { 6, 0, 5 }, // top_k->src[0] == add +}; + //node #978 ( SOFT_MAX): ffn_moe_probs-15 ( 0K) [Vulka ] use=2: ffn_moe_logits-15 ( 0K) [Vulka ] //node #979 ( RESHAPE): ffn_moe_probs-15 (re ( 0K) [Vulka ] use=1: ffn_moe_probs-15 ( 0K) [Vulka ] //node #980 ( ARGSORT): ffn_moe_argsort-15 ( 0K) [Vulka ] use=1: ffn_moe_probs-15 ( 0K) [Vulka ] @@ -1057,6 +1072,8 @@ struct vk_device_struct { vk_pipeline pipeline_argsort_f32[num_argsort_pipelines]; vk_pipeline pipeline_argsort_large_f32[num_argsort_pipelines]; vk_pipeline pipeline_topk_f32[num_topk_pipelines]; + vk_pipeline pipeline_topk_radix_f32; + vk_pipeline pipeline_topk_radix_qsa; // qwen4 QSA indexer fusion (f16 mask) vk_pipeline pipeline_sum_rows_f32; vk_pipeline pipeline_cross_entropy_loss_f32, pipeline_cross_entropy_loss_f32_wg512; vk_pipeline pipeline_cross_entropy_loss_back_f32, pipeline_cross_entropy_loss_back_f32_wg512; @@ -1749,6 +1766,15 @@ struct vk_op_topk_push_constants { uint32_t last_pass; }; +struct vk_op_topk_radix_push_constants { + uint32_t ncols; + uint32_t k; + uint32_t nrows; + uint32_t n_tps; // QSA only + uint32_t n_blocks; // QSA only + uint32_t n_stream; // QSA only +}; + struct vk_op_im2col_push_constants { uint64_t dst_addr; uint32_t batch_offset; uint32_t offset_delta; @@ -2439,6 +2465,8 @@ struct ggml_backend_vk_context { int fused_ops_write_mask {}; topk_moe_mode fused_topk_moe_mode {}; bool fused_topk_moe_scale {}; + // QSA indexer gather+add+top_k fused into one radix-select + bool fused_topk_qsa {}; // for GGML_VK_PERF_LOGGER std::unique_ptr perf_logger; @@ -5814,6 +5842,14 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { } } + // large-k fallback: one workgroup per row, radix-select instead of a full sort. The QSA + // variant (spec constant 1) additionally gathers the qwen4 indexer input on the fly. + { + const uint32_t BLOCK_SIZE = 1u << std::min(10u, device->max_workgroup_size_log2); + ggml_vk_create_pipeline2(device, device->pipeline_topk_radix_f32, "topk_radix_f32", topk_radix_select_f32_len, topk_radix_select_f32_data, "main", 5, sizeof(vk_op_topk_radix_push_constants), {BLOCK_SIZE, 1, 1}, {BLOCK_SIZE, 0}, 1, true); + ggml_vk_create_pipeline2(device, device->pipeline_topk_radix_qsa, "topk_radix_qsa", topk_radix_select_f32_len, topk_radix_select_f32_data, "main", 5, sizeof(vk_op_topk_radix_push_constants), {BLOCK_SIZE, 1, 1}, {BLOCK_SIZE, 1}, 1, true); + } + ggml_vk_create_pipeline(device, device->pipeline_argmax_f32, "argmax_f32", argmax_f32_len, argmax_f32_data, "main", 2, sizeof(vk_op_push_constants), {1, 1, 1}, { device->subgroup_size }, 1); ggml_vk_create_pipeline(device, device->pipeline_sum_rows_f32, "sum_rows_f32", sum_rows_f32_len, sum_rows_f32_data, "main", 2, sizeof(vk_op_sum_rows_push_constants), {1, 1, 1}, { device->subgroup_size }, 1); @@ -13940,6 +13976,31 @@ static void ggml_vk_topk(ggml_backend_vk_context * ctx, vk_context& subctx, cons uint32_t nrows = ggml_nrows(src0); uint32_t k = dst->ne[0]; + // tournament path is faster where it fits; use radix-select only past its k limit + const uint32_t k_min_pipeline = std::max((uint32_t) log2f(float(k)) + 1, ctx->device->subgroup_size_log2); + if (k_min_pipeline >= num_topk_pipelines || ctx->device->pipeline_topk_f32[k_min_pipeline] == nullptr) { + vk_pipeline pipeline = ctx->device->pipeline_topk_radix_f32; + GGML_ASSERT(pipeline != nullptr); + + if (ctx->prealloc_x_need_sync) { + ggml_vk_sync_buffers(ctx, subctx); + } + + vk_op_topk_radix_push_constants pc { ncols, k, nrows, 0, 0, 0 }; + std::array elements { + pipeline->wg_denoms[0], + std::min(nrows, ctx->device->properties.limits.maxComputeWorkGroupCount[1]), + 1, + }; + // the non-QSA path only uses bindings 0/1; bind valid buffers for the unused QSA slots + vk_subbuffer src0_buf = ggml_vk_tensor_subbuffer(ctx, src0); + vk_subbuffer dst_buf = ggml_vk_tensor_subbuffer(ctx, dst); + ggml_pipeline_request_descriptor_sets(ctx, pipeline, 1); + ggml_vk_dispatch_pipeline(ctx, subctx, pipeline, + { src0_buf, dst_buf, src0_buf, src0_buf, src0_buf }, pc, elements); + return; + } + vk_op_topk_push_constants pc { ncols, ncols, ncols, k, nrows, 0, 0 }; if (ctx->prealloc_x_need_sync) { @@ -14043,6 +14104,55 @@ static void ggml_vk_topk(ggml_backend_vk_context * ctx, vk_context& subctx, cons ctx->prealloc_x_need_sync = true; } +static void ggml_vk_topk_qsa(ggml_backend_vk_context * ctx, vk_context& subctx, const ggml_cgraph * cgraph, int node_idx) { + const ggml_tensor * get_rows = cgraph->nodes[node_idx + 0]; + const ggml_tensor * add = cgraph->nodes[node_idx + ctx->num_additional_fused_ops - 1]; + ggml_tensor * top_k = cgraph->nodes[node_idx + ctx->num_additional_fused_ops]; + + const ggml_tensor * scores = get_rows->src[0]; // [n_tps, n_blocks, n_stream] + const ggml_tensor * cell_blk = get_rows->src[1]; // [n_kv, n_stream] + + // raw f16 mask: follow the reshape/cpy chain back to the materialized input + const ggml_tensor * mask = add->src[1]; + while (mask->op == GGML_OP_RESHAPE || mask->op == GGML_OP_CPY) { + mask = mask->src[0]; + } + + const uint32_t n_tps = scores->ne[0]; + const uint32_t n_blocks = scores->ne[1]; + const uint32_t n_stream = scores->ne[2]; + const uint32_t n_kv = cell_blk->ne[0]; + const uint32_t width = top_k->ne[0]; + const uint32_t nrows = n_tps * n_stream; + + vk_pipeline pipeline = ctx->device->pipeline_topk_radix_qsa; + GGML_ASSERT(pipeline != nullptr); + + // scratch holds the gathered+masked input, materialized once and reused across passes + const size_t scratch_size = size_t{ n_kv } * nrows * sizeof(float); + if (ctx->prealloc_size_x < scratch_size) { + ctx->prealloc_size_x = scratch_size; + ggml_vk_preallocate_buffers(ctx, subctx); + } + if (ctx->prealloc_x_need_sync) { + ggml_vk_sync_buffers(ctx, subctx); + } + + vk_op_topk_radix_push_constants pc { n_kv, width, nrows, n_tps, n_blocks, n_stream }; + std::array elements { + pipeline->wg_denoms[0], + std::min(nrows, ctx->device->properties.limits.maxComputeWorkGroupCount[1]), + 1, + }; + vk_subbuffer scratch_buf { ctx->prealloc_x, 0, ctx->prealloc_x->size }; + ggml_pipeline_request_descriptor_sets(ctx, pipeline, 1); + ggml_vk_dispatch_pipeline(ctx, subctx, pipeline, + { ggml_vk_tensor_subbuffer(ctx, scores), ggml_vk_tensor_subbuffer(ctx, top_k), + ggml_vk_tensor_subbuffer(ctx, cell_blk), ggml_vk_tensor_subbuffer(ctx, mask), + scratch_buf }, pc, elements); + ctx->prealloc_x_need_sync = true; +} + static void ggml_vk_sum(ggml_backend_vk_context * ctx, vk_context& subctx, const ggml_tensor * src0, ggml_tensor * dst) { vk_op_sum_rows_push_constants p = vk_op_sum_rows_push_constants_init(src0, dst, ggml_nelements(src0)); ggml_vk_op_f32(ctx, subctx, src0, nullptr, nullptr, nullptr, dst, GGML_OP_SUM, p); @@ -15704,7 +15814,11 @@ static bool ggml_vk_build_graph(ggml_backend_vk_context * ctx, ggml_cgraph * cgr break; case GGML_OP_GET_ROWS: - ggml_vk_get_rows(ctx, compute_ctx, src0, src1, node); + if (ctx->fused_topk_qsa) { + ggml_vk_topk_qsa(ctx, compute_ctx, cgraph, node_idx); + } else { + ggml_vk_get_rows(ctx, compute_ctx, src0, src1, node); + } break; case GGML_OP_GET_ROWS_BACK: @@ -17116,6 +17230,92 @@ static bool ggml_vk_can_fuse_topk_moe(ggml_backend_vk_context * ctx, const struc return true; } +// Manual op-sequence match (ggml_can_fuse_subgraph rejects the mask's external reshape/cpy). +static bool ggml_vk_match_ops(const struct ggml_cgraph * cgraph, int node_idx, + const std::initializer_list & ops) { + if (node_idx + (int) ops.size() > cgraph->n_nodes) { + return false; + } + for (size_t j = 0; j < ops.size(); ++j) { + const ggml_tensor * node = cgraph->nodes[node_idx + j]; + if (node->op != ops.begin()[j] || + (node->flags & GGML_TENSOR_FLAG_COMPUTE) == 0 || + (node->flags & GGML_TENSOR_FLAG_OUTPUT) != 0) { + return false; + } + } + return true; +} + +// True if the qwen4 QSA indexer top-k can be fused at node_idx (the get_rows). +static bool ggml_vk_can_fuse_topk_qsa(ggml_backend_vk_context * ctx, const struct ggml_cgraph * cgraph, int node_idx) { + if (ctx->device->disable_fusion || !ctx->device->pipeline_topk_radix_qsa) { + return false; + } + + const int n_ops = topk_qsa_pattern.size(); + if (!ggml_vk_match_ops(cgraph, node_idx, topk_qsa_pattern) || + !ggml_check_edges(cgraph, node_idx, topk_qsa_edges)) { + return false; + } + + // elided nodes must be single-use (cpy counts its own src[1] self-reference) + for (int j = 0; j < n_ops - 1; ++j) { + const ggml_tensor * node = cgraph->nodes[node_idx + j]; + const int32_t want = node->op == GGML_OP_CPY ? 2 : 1; + if (ggml_node_get_use_count(cgraph, node_idx + j) != want) { + return false; + } + } + + const ggml_tensor * get_rows = cgraph->nodes[node_idx + 0]; + const ggml_tensor * add = cgraph->nodes[node_idx + n_ops - 2]; + const ggml_tensor * top_k = cgraph->nodes[node_idx + n_ops - 1]; + + const ggml_tensor * scores = get_rows->src[0]; // [n_tps, n_blocks, n_stream] + const ggml_tensor * cell_blk = get_rows->src[1]; // [n_kv, n_stream] + const ggml_tensor * expanded = add->src[0]; // [n_kv, n_tps, n_stream] + + // raw mask: follow the reshape/cpy chain back to the materialized f16 input + const ggml_tensor * mask = add->src[1]; + while (mask && (mask->op == GGML_OP_RESHAPE || mask->op == GGML_OP_CPY)) { + mask = mask->src[0]; + } + if (!mask || mask->type != GGML_TYPE_F16) { + return false; + } + + if (scores->type != GGML_TYPE_F32 || cell_blk->type != GGML_TYPE_I32 || top_k->type != GGML_TYPE_I32) { + return false; + } + if (!ggml_is_contiguous(scores) || !ggml_is_contiguous(cell_blk) || !ggml_is_contiguous(mask) || + !ggml_is_contiguous(expanded) || !ggml_is_contiguous(top_k)) { + return false; + } + + const int64_t n_tps = scores->ne[0]; + const int64_t n_blocks = scores->ne[1]; + const int64_t n_stream = scores->ne[2]; + const int64_t n_kv = cell_blk->ne[0]; + const int64_t width = top_k->ne[0]; + + // pin the indexer layout the shader's addressing assumes + if (scores->ne[3] != 1 || cell_blk->ne[1] != n_stream || ggml_nrows(cell_blk) != n_stream || + ggml_nelements(mask) != n_kv * n_tps * n_stream || + expanded->ne[0] != n_kv || expanded->ne[1] != n_tps || expanded->ne[2] != n_stream || + top_k->ne[1] != n_tps || top_k->ne[2] != n_stream || top_k->ne[3] != 1 || + n_blocks <= 0 || n_kv <= 0 || width <= 0 || width > n_kv) { + return false; + } + + // only worth it in the radix regime; small k uses the faster tournament unfused + const uint32_t k_min_pipeline = std::max((uint32_t) log2f(float(width)) + 1, ctx->device->subgroup_size_log2); + if (k_min_pipeline < num_topk_pipelines && ctx->device->pipeline_topk_f32[k_min_pipeline]) { + return false; + } + return true; +} + static bool ggml_vk_can_fuse_rope_set_rows(ggml_backend_vk_context * ctx, const struct ggml_cgraph * cgraph, int node_idx) { GGML_UNUSED(ctx); @@ -17495,6 +17695,7 @@ static ggml_status ggml_backend_vk_graph_compute(ggml_backend_t backend, ggml_cg ctx->fused_topk_moe_mode = TOPK_MOE_COUNT; ctx->fused_topk_moe_scale = false; + ctx->fused_topk_qsa = false; const char *fusion_string {}; if (!ctx->device->disable_fusion) { uint32_t num_adds = ggml_vk_fuse_multi_add(ctx, cgraph, i); @@ -17584,6 +17785,11 @@ static ggml_status ggml_backend_vk_graph_compute(ggml_backend_t backend, ggml_cg // with a data dependency on that register. The overlap check still // rejects partial overlaps (different base or size). std::fill_n(op_srcs_fused_elementwise, 5, true); + } else if (ggml_vk_can_fuse_topk_qsa(ctx, cgraph, i)) { + ctx->num_additional_fused_ops = topk_qsa_pattern.size() - 1; + ctx->fused_topk_qsa = true; + fusion_string = "TOPK_QSA"; + std::fill_n(op_srcs_fused_elementwise, ctx->num_additional_fused_ops + 1, false); } else if (ggml_can_fuse_subgraph(cgraph, i, topk_moe_early_softmax_norm, { i + 3, i + 9 }) && ggml_check_edges(cgraph, i, topk_moe_early_softmax_norm_edges) && ggml_vk_can_fuse_topk_moe(ctx, cgraph, i, TOPK_MOE_EARLY_SOFTMAX_NORM)) { @@ -17700,6 +17906,7 @@ static ggml_status ggml_backend_vk_graph_compute(ggml_backend_t backend, ggml_cg ctx->fused_ops_write_mask = 1; ctx->fused_topk_moe_mode = TOPK_MOE_COUNT; ctx->fused_topk_moe_scale = false; + ctx->fused_topk_qsa = false; } } @@ -17896,6 +18103,9 @@ static void ggml_vk_graph_optimize(ggml_backend_t backend, struct ggml_cgraph * if (keep_pattern(snake_pattern)) { continue; } + if (keep_pattern(topk_qsa_pattern)) { + continue; + } // First, grab the next unused node. current_set.push_back(first_unused); @@ -17914,13 +18124,23 @@ static void ggml_vk_graph_optimize(ggml_backend_t backend, struct ggml_cgraph * if (is_empty(graph->nodes[j])) { continue; } - // Don't pull forward nodes from fusion patterns + // Protect every interior QSA node (not just the start): the mask branch is + // independent, so it gets pulled out and breaks keep_pattern otherwise. + auto const &in_qsa_pattern = [&](int n) -> bool { + for (int o = 0; o < (int) topk_qsa_pattern.size(); ++o) { + if (n - o >= 0 && match_pattern(topk_qsa_pattern, n - o)) { + return true; + } + } + return false; + }; if (match_pattern(topk_moe_early_softmax_norm, j) || match_pattern(topk_moe_sigmoid_norm_bias, j) || match_pattern(topk_moe_sqrt_softplus_norm_bias, j) || match_pattern(topk_moe_early_softmax, j) || match_pattern(topk_moe_late_softmax, j) || - match_pattern(snake_pattern, j)) { + match_pattern(snake_pattern, j) || + in_qsa_pattern(j)) { continue; } bool ok = true; @@ -18723,15 +18943,14 @@ static bool ggml_backend_vk_device_supports_op(ggml_backend_dev_t dev, const ggm if (!ggml_is_contiguous(op) || !ggml_is_contiguous(op->src[0])) { return false; } - // We could potentially support larger, using argsort to sort the - // whole thing. Not clear if this is needed. - uint32_t min_pipeline = (uint32_t)log2f(float(op->ne[0])) + 1; - if (min_pipeline >= num_topk_pipelines || - !device->pipeline_topk_f32[min_pipeline]) { - return false; + // large k falls back to radix-select + const uint32_t min_pipeline = + std::max((uint32_t) log2f(float(op->ne[0])) + 1, device->subgroup_size_log2); + if (min_pipeline < num_topk_pipelines && device->pipeline_topk_f32[min_pipeline]) { + return true; } + return device->pipeline_topk_radix_f32 != nullptr; } - return true; case GGML_OP_UPSCALE: if (op->op_params[0] & GGML_SCALE_FLAG_ANTIALIAS) { if ((op->op_params[0] & 0xFF) != GGML_SCALE_MODE_BILINEAR) { diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/topk_radix_select.comp b/ggml/src/ggml-vulkan/vulkan-shaders/topk_radix_select.comp new file mode 100644 index 00000000000..8e14b2e9925 --- /dev/null +++ b/ggml/src/ggml-vulkan/vulkan-shaders/topk_radix_select.comp @@ -0,0 +1,144 @@ +#version 450 + +#extension GL_EXT_control_flow_attributes : enable +#extension GL_EXT_shader_16bit_storage : require + +#include "types.glsl" + +layout(constant_id = 0) const int BLOCK_SIZE = 1024; +layout(constant_id = 1) const int QSA = 0; // 1: fuse the qwen4 QSA indexer gather + f16 mask + +layout(local_size_x_id = 0, local_size_y = 1, local_size_z = 1) in; + +layout (binding = 0) readonly buffer A {float data_a[];}; // input values, or QSA block scores [n_tps, n_blocks, n_stream] +layout (binding = 1) writeonly buffer D {int data_d[];}; // [k, ...] +layout (binding = 2) readonly buffer CB {int cell_blk[];}; // QSA: cell->block map [n_kv, n_stream] +layout (binding = 3) readonly buffer M {float16_t mask[];}; // QSA: raw f16 kq_mask [n_kv, n_tps, n_stream] +layout (binding = 4) buffer S {float scratch[];}; // QSA: [nrows, n_kv] gathered inputs + +layout (push_constant) uniform parameter { + uint ncols; + uint k; + uint nrows; + uint n_tps; // QSA only + uint n_blocks; // QSA only + uint n_stream; // QSA only +} p; + +#define RADIX_BITS 8 +#define RADIX_SIZE (1 << RADIX_BITS) + +shared uint histo[RADIX_SIZE]; +shared uint sh_bucket; +shared uint sh_above; +shared uint out_count; + +// order-preserving float -> uint mapping +uint f2ui(float x) { + uint y = floatBitsToUint(x); + if ((y & 0x80000000u) != 0u) { + y ^= 0xFFFFFFFFu; + } else { + y |= 0x80000000u; + } + return y; +} + +// QSA element i of row (t,s): score[cell_blk[i,s], t, s] + mask[i,t,s] +float gather(uint row, uint i) { + const uint t = row % p.n_tps; + const uint s = row / p.n_tps; + const uint block = uint(cell_blk[s * p.ncols + i]); + const float a = data_a[(s * p.n_blocks + block) * p.n_tps + t]; + const float m = float(mask[(s * p.n_tps + t) * p.ncols + i]); + return a + m; +} + +float load(uint row, uint i, bool first) { + if (QSA == 0) { + return data_a[row * p.ncols + i]; + } + // materialize the scattered gather on the first pass and reuse it after; each + // invocation only touches its own scratch entries, so no barrier is needed + const uint off = row * p.ncols + i; + if (first) { + const float v = gather(row, i); + scratch[off] = v; + return v; + } + return scratch[off]; +} + +// one workgroup per row: radix-select the K-th largest, then compact it plus enough ties +void topk(const uint row) { + const uint tid = gl_LocalInvocationID.x; + const uint ncols = p.ncols; + const uint row_out = row * p.k; + + uint prefix = 0; // fixed high bits of the threshold key + uint desired = p.k; // count still needed from the candidate range + + [[unroll]] for (int shift = 32 - RADIX_BITS; shift >= 0; shift -= RADIX_BITS) { + for (uint i = tid; i < RADIX_SIZE; i += BLOCK_SIZE) { + histo[i] = 0; + } + barrier(); + + const bool first = (shift == 32 - RADIX_BITS); + const uint hi_mask = (shift + RADIX_BITS >= 32) ? 0u : (0xFFFFFFFFu << uint(shift + RADIX_BITS)); + const uint prefix_hi = prefix & hi_mask; + for (uint i = tid; i < ncols; i += BLOCK_SIZE) { + const uint key = f2ui(load(row, i, first)); + if ((key & hi_mask) == prefix_hi) { + atomicAdd(histo[(key >> uint(shift)) & (RADIX_SIZE - 1)], 1u); + } + } + barrier(); + + // top-down scan for the bucket holding the K-th value + if (tid == 0) { + uint acc = 0; + uint b = 0; + for (int bb = RADIX_SIZE - 1; bb >= 0; --bb) { + const uint c = histo[bb]; + if (acc + c >= desired) { b = uint(bb); break; } + acc += c; + } + sh_bucket = b; + sh_above = acc; + } + barrier(); + + prefix |= sh_bucket << uint(shift); + desired -= sh_above; + barrier(); + } + + if (tid == 0) { + out_count = 0; + } + barrier(); + + // emit everything above the threshold, then fill the rest from ties + const uint threshold = prefix; + for (uint i = tid; i < ncols; i += BLOCK_SIZE) { + if (f2ui(load(row, i, false)) > threshold) { + data_d[row_out + atomicAdd(out_count, 1u)] = int(i); + } + } + barrier(); + for (uint i = tid; i < ncols; i += BLOCK_SIZE) { + if (f2ui(load(row, i, false)) == threshold) { + const uint pos = atomicAdd(out_count, 1u); + if (pos < p.k) { + data_d[row_out + pos] = int(i); + } + } + } +} + +void main() { + for (uint row = gl_WorkGroupID.y; row < p.nrows; row += gl_NumWorkGroups.y) { + topk(row); + } +} diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp b/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp index f0610f4cd82..27ff68c10d5 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp @@ -1028,6 +1028,7 @@ void process_shaders() { string_to_spv("topk_argsort_f32", "topk_argsort.comp", {{"A_TYPE", "float"}}); string_to_spv("topk_nary_search_f32", "topk_nary_search.comp", {{"A_TYPE", "float"}}); + string_to_spv("topk_radix_select_f32", "topk_radix_select.comp", {{"A_TYPE", "float"}}); string_to_spv("argmax_f32", "argmax.comp", merge_maps(base_dict, {{"A_TYPE", "float"}, {"D_TYPE", "int"}})); string_to_spv("sum_rows_f32", "sum_rows.comp", merge_maps(base_dict, {{"A_TYPE", "float"}, {"D_TYPE", "float"}})); From 96dddd87f20b470a721a3c2bd25aba63a045a6a8 Mon Sep 17 00:00:00 2001 From: fairydreaming <166155368+fairydreaming@users.noreply.github.com> Date: Mon, 31 Aug 2026 10:17:23 +0200 Subject: [PATCH 051/104] ggml : add MUL_MAT to the list of ops that may need additional memory (for WebGPU) (llama/28071) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Co-authored-by: Stanisław Szymczyk --- ggml/src/ggml-backend.cpp | 1 + 1 file changed, 1 insertion(+) diff --git a/ggml/src/ggml-backend.cpp b/ggml/src/ggml-backend.cpp index 8f40fb2f7bc..ffe20b9d05b 100644 --- a/ggml/src/ggml-backend.cpp +++ b/ggml/src/ggml-backend.cpp @@ -2115,6 +2115,7 @@ ggml_backend_t ggml_backend_sched_get_tensor_backend(ggml_backend_sched_t sched, bool ggml_backend_op_alloc_size_may_expand(enum ggml_op op) { switch (op) { case GGML_OP_FLASH_ATTN_EXT: + case GGML_OP_MUL_MAT: case GGML_OP_MUL_MAT_ID: case GGML_OP_CUMSUM: case GGML_OP_ARGSORT: From 6ce7b89952fd07d0db62e99fb15c231ca3979714 Mon Sep 17 00:00:00 2001 From: Simon Teixidor Date: Mon, 31 Aug 2026 11:07:53 +0200 Subject: [PATCH 052/104] vulkan: tune mat-vec rows for batched inference on Strix Halo (llama/27909) * vulkan: RDNA3 static mat-vec rows above four columns On RDNA3 above four columns a static 4 rows for all types benches faster than the default. * vulkan: RDNA3 static mat-vec-id rows mul_mat_vec_id has no column dimension to switch on. On my Strix Halo machine, a static 4 is faster here than the defaults across types and batch sizes. --- ggml/src/ggml-vulkan/ggml-vulkan.cpp | 56 ++++++++++++++++------------ 1 file changed, 32 insertions(+), 24 deletions(-) diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp index 649ec792062..8718bd2cfb6 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp +++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp @@ -5289,6 +5289,11 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { rm_stdq = 2; rm_stdq_int = 2; } + // RDNA3: above four columns, static 4 rows for all types bench faster than the default + const bool is_rdna3 = device->vendor_id == VK_VENDOR_ID_AMD && device->architecture == AMD_RDNA3; + auto const &rm_int_n = [&](uint32_t rows, uint32_t i) { return (is_rdna3 && i >= 4) ? 4u : rows; }; + // RDNA3: Static 4 rows for all types bench faster than the default + auto const &rm_id = [&](uint32_t rows) { return is_rdna3 ? 4u : rows; }; uint32_t rm_iq = 2 * rm_kq; const bool use_subgroups = device->subgroup_arithmetic; @@ -5385,20 +5390,20 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { const uint32_t subgroup_size_int = (device->vendor_id == VK_VENDOR_ID_INTEL && device->subgroup_size_control) ? device->subgroup_min_size : device->subgroup_size; const uint32_t wg_size_subgroup_int = (w == DMMV_WG_SIZE_SUBGROUP) ? subgroup_size_int : (subgroup_size_int * 4); - ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_q8_1_f32[w][GGML_TYPE_Q2_0][i], "mul_mat_vec_q2_0_q8_1_f32", arr_dmmv_q2_0_q8_1_f32_len[reduc], arr_dmmv_q2_0_q8_1_f32_data[reduc], "main", mul_mat_vec_num_bindings, sizeof(vk_mat_vec_push_constants), {2*rm_kq_int, 1, 1}, {wg_size_subgroup_int, 2*rm_kq_int, i+1}, 1, true, use_subgroups, subgroup_size_int); - ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_q8_1_f32[w][GGML_TYPE_Q4_0][i], "mul_mat_vec_q4_0_q8_1_f32", arr_dmmv_q4_0_q8_1_f32_len[reduc], arr_dmmv_q4_0_q8_1_f32_data[reduc], "main", mul_mat_vec_num_bindings, sizeof(vk_mat_vec_push_constants), {1*rm_stdq_int, 1, 1}, {wg_size_subgroup_int, 1*rm_stdq_int, i+1}, 1, true, use_subgroups, subgroup_size_int); - ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_q8_1_f32[w][GGML_TYPE_Q4_1][i], "mul_mat_vec_q4_1_q8_1_f32", arr_dmmv_q4_1_q8_1_f32_len[reduc], arr_dmmv_q4_1_q8_1_f32_data[reduc], "main", mul_mat_vec_num_bindings, sizeof(vk_mat_vec_push_constants), {1*rm_stdq_int, 1, 1}, {wg_size_subgroup_int, 1*rm_stdq_int, i+1}, 1, true, use_subgroups, subgroup_size_int); - ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_q8_1_f32[w][GGML_TYPE_Q5_0][i], "mul_mat_vec_q5_0_q8_1_f32", arr_dmmv_q5_0_q8_1_f32_len[reduc], arr_dmmv_q5_0_q8_1_f32_data[reduc], "main", mul_mat_vec_num_bindings, sizeof(vk_mat_vec_push_constants), {1*rm_stdq_int, 1, 1}, {wg_size_subgroup_int, 1*rm_stdq_int, i+1}, 1, true, use_subgroups, subgroup_size_int); - ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_q8_1_f32[w][GGML_TYPE_Q5_1][i], "mul_mat_vec_q5_1_q8_1_f32", arr_dmmv_q5_1_q8_1_f32_len[reduc], arr_dmmv_q5_1_q8_1_f32_data[reduc], "main", mul_mat_vec_num_bindings, sizeof(vk_mat_vec_push_constants), {1*rm_stdq_int, 1, 1}, {wg_size_subgroup_int, 1*rm_stdq_int, i+1}, 1, true, use_subgroups, subgroup_size_int); - ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_q8_1_f32[w][GGML_TYPE_Q8_0][i], "mul_mat_vec_q8_0_q8_1_f32", arr_dmmv_q8_0_q8_1_f32_len[reduc], arr_dmmv_q8_0_q8_1_f32_data[reduc], "main", mul_mat_vec_num_bindings, sizeof(vk_mat_vec_push_constants), {1*rm_stdq_int, 1, 1}, {wg_size_subgroup_int, 1*rm_stdq_int, i+1}, 1, true, use_subgroups, subgroup_size_int); + ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_q8_1_f32[w][GGML_TYPE_Q2_0][i], "mul_mat_vec_q2_0_q8_1_f32", arr_dmmv_q2_0_q8_1_f32_len[reduc], arr_dmmv_q2_0_q8_1_f32_data[reduc], "main", mul_mat_vec_num_bindings, sizeof(vk_mat_vec_push_constants), {rm_int_n(2*rm_kq_int, i), 1, 1}, {wg_size_subgroup_int, rm_int_n(2*rm_kq_int, i), i+1}, 1, true, use_subgroups, subgroup_size_int); + ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_q8_1_f32[w][GGML_TYPE_Q4_0][i], "mul_mat_vec_q4_0_q8_1_f32", arr_dmmv_q4_0_q8_1_f32_len[reduc], arr_dmmv_q4_0_q8_1_f32_data[reduc], "main", mul_mat_vec_num_bindings, sizeof(vk_mat_vec_push_constants), {rm_int_n(1*rm_stdq_int, i), 1, 1}, {wg_size_subgroup_int, rm_int_n(1*rm_stdq_int, i), i+1}, 1, true, use_subgroups, subgroup_size_int); + ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_q8_1_f32[w][GGML_TYPE_Q4_1][i], "mul_mat_vec_q4_1_q8_1_f32", arr_dmmv_q4_1_q8_1_f32_len[reduc], arr_dmmv_q4_1_q8_1_f32_data[reduc], "main", mul_mat_vec_num_bindings, sizeof(vk_mat_vec_push_constants), {rm_int_n(1*rm_stdq_int, i), 1, 1}, {wg_size_subgroup_int, rm_int_n(1*rm_stdq_int, i), i+1}, 1, true, use_subgroups, subgroup_size_int); + ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_q8_1_f32[w][GGML_TYPE_Q5_0][i], "mul_mat_vec_q5_0_q8_1_f32", arr_dmmv_q5_0_q8_1_f32_len[reduc], arr_dmmv_q5_0_q8_1_f32_data[reduc], "main", mul_mat_vec_num_bindings, sizeof(vk_mat_vec_push_constants), {rm_int_n(1*rm_stdq_int, i), 1, 1}, {wg_size_subgroup_int, rm_int_n(1*rm_stdq_int, i), i+1}, 1, true, use_subgroups, subgroup_size_int); + ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_q8_1_f32[w][GGML_TYPE_Q5_1][i], "mul_mat_vec_q5_1_q8_1_f32", arr_dmmv_q5_1_q8_1_f32_len[reduc], arr_dmmv_q5_1_q8_1_f32_data[reduc], "main", mul_mat_vec_num_bindings, sizeof(vk_mat_vec_push_constants), {rm_int_n(1*rm_stdq_int, i), 1, 1}, {wg_size_subgroup_int, rm_int_n(1*rm_stdq_int, i), i+1}, 1, true, use_subgroups, subgroup_size_int); + ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_q8_1_f32[w][GGML_TYPE_Q8_0][i], "mul_mat_vec_q8_0_q8_1_f32", arr_dmmv_q8_0_q8_1_f32_len[reduc], arr_dmmv_q8_0_q8_1_f32_data[reduc], "main", mul_mat_vec_num_bindings, sizeof(vk_mat_vec_push_constants), {rm_int_n(1*rm_stdq_int, i), 1, 1}, {wg_size_subgroup_int, rm_int_n(1*rm_stdq_int, i), i+1}, 1, true, use_subgroups, subgroup_size_int); - ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_q8_1_f32[w][GGML_TYPE_MXFP4][i], "mul_mat_vec_mxfp4_q8_1_f32", arr_dmmv_mxfp4_q8_1_f32_len[reduc], arr_dmmv_mxfp4_q8_1_f32_data[reduc], "main", mul_mat_vec_num_bindings, sizeof(vk_mat_vec_push_constants), {2*rm_stdq_int, 1, 1}, {wg_size_subgroup_int, 2*rm_stdq_int, i+1}, 1, true, use_subgroups, subgroup_size_int); + ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_q8_1_f32[w][GGML_TYPE_MXFP4][i], "mul_mat_vec_mxfp4_q8_1_f32", arr_dmmv_mxfp4_q8_1_f32_len[reduc], arr_dmmv_mxfp4_q8_1_f32_data[reduc], "main", mul_mat_vec_num_bindings, sizeof(vk_mat_vec_push_constants), {rm_int_n(2*rm_stdq_int, i), 1, 1}, {wg_size_subgroup_int, rm_int_n(2*rm_stdq_int, i), i+1}, 1, true, use_subgroups, subgroup_size_int); - ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_q8_1_f32[w][GGML_TYPE_Q2_K][i], "mul_mat_vec_q2_k_q8_1_f32", arr_dmmv_q2_k_q8_1_f32_len[reduc], arr_dmmv_q2_k_q8_1_f32_data[reduc], "main", mul_mat_vec_num_bindings, sizeof(vk_mat_vec_push_constants), {2*rm_kq_int, 1, 1}, {wg_size_subgroup_int, 2*rm_kq_int, i+1}, 1, true, use_subgroups, subgroup_size_int); - ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_q8_1_f32[w][GGML_TYPE_Q3_K][i], "mul_mat_vec_q3_k_q8_1_f32", arr_dmmv_q3_k_q8_1_f32_len[reduc], arr_dmmv_q3_k_q8_1_f32_data[reduc], "main", mul_mat_vec_num_bindings, sizeof(vk_mat_vec_push_constants), {1*rm_kq_int, 1, 1}, {wg_size_subgroup_int, 1*rm_kq_int, i+1}, 1, true, use_subgroups, subgroup_size_int); - ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_q8_1_f32[w][GGML_TYPE_Q4_K][i], "mul_mat_vec_q4_k_q8_1_f32", arr_dmmv_q4_k_q8_1_f32_len[reduc], arr_dmmv_q4_k_q8_1_f32_data[reduc], "main", mul_mat_vec_num_bindings, sizeof(vk_mat_vec_push_constants), {1*rm_kq_int, 1, 1}, {wg_size_subgroup_int, 1*rm_kq_int, i+1}, 1, true, use_subgroups, subgroup_size_int); - ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_q8_1_f32[w][GGML_TYPE_Q5_K][i], "mul_mat_vec_q5_k_q8_1_f32", arr_dmmv_q5_k_q8_1_f32_len[reduc], arr_dmmv_q5_k_q8_1_f32_data[reduc], "main", mul_mat_vec_num_bindings, sizeof(vk_mat_vec_push_constants), {1*rm_kq_int, 1, 1}, {wg_size_subgroup_int, 1*rm_kq_int, i+1}, 1, true, use_subgroups, subgroup_size_int); - ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_q8_1_f32[w][GGML_TYPE_Q6_K][i], "mul_mat_vec_q6_k_q8_1_f32", arr_dmmv_q6_k_q8_1_f32_len[reduc], arr_dmmv_q6_k_q8_1_f32_data[reduc], "main", mul_mat_vec_num_bindings, sizeof(vk_mat_vec_push_constants), {1*rm_kq_int, 1, 1}, {wg_size_subgroup_int, 1*rm_kq_int, i+1}, 1, true, use_subgroups, subgroup_size_int); + ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_q8_1_f32[w][GGML_TYPE_Q2_K][i], "mul_mat_vec_q2_k_q8_1_f32", arr_dmmv_q2_k_q8_1_f32_len[reduc], arr_dmmv_q2_k_q8_1_f32_data[reduc], "main", mul_mat_vec_num_bindings, sizeof(vk_mat_vec_push_constants), {rm_int_n(2*rm_kq_int, i), 1, 1}, {wg_size_subgroup_int, rm_int_n(2*rm_kq_int, i), i+1}, 1, true, use_subgroups, subgroup_size_int); + ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_q8_1_f32[w][GGML_TYPE_Q3_K][i], "mul_mat_vec_q3_k_q8_1_f32", arr_dmmv_q3_k_q8_1_f32_len[reduc], arr_dmmv_q3_k_q8_1_f32_data[reduc], "main", mul_mat_vec_num_bindings, sizeof(vk_mat_vec_push_constants), {rm_int_n(1*rm_kq_int, i), 1, 1}, {wg_size_subgroup_int, rm_int_n(1*rm_kq_int, i), i+1}, 1, true, use_subgroups, subgroup_size_int); + ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_q8_1_f32[w][GGML_TYPE_Q4_K][i], "mul_mat_vec_q4_k_q8_1_f32", arr_dmmv_q4_k_q8_1_f32_len[reduc], arr_dmmv_q4_k_q8_1_f32_data[reduc], "main", mul_mat_vec_num_bindings, sizeof(vk_mat_vec_push_constants), {rm_int_n(1*rm_kq_int, i), 1, 1}, {wg_size_subgroup_int, rm_int_n(1*rm_kq_int, i), i+1}, 1, true, use_subgroups, subgroup_size_int); + ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_q8_1_f32[w][GGML_TYPE_Q5_K][i], "mul_mat_vec_q5_k_q8_1_f32", arr_dmmv_q5_k_q8_1_f32_len[reduc], arr_dmmv_q5_k_q8_1_f32_data[reduc], "main", mul_mat_vec_num_bindings, sizeof(vk_mat_vec_push_constants), {rm_int_n(1*rm_kq_int, i), 1, 1}, {wg_size_subgroup_int, rm_int_n(1*rm_kq_int, i), i+1}, 1, true, use_subgroups, subgroup_size_int); + ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_q8_1_f32[w][GGML_TYPE_Q6_K][i], "mul_mat_vec_q6_k_q8_1_f32", arr_dmmv_q6_k_q8_1_f32_len[reduc], arr_dmmv_q6_k_q8_1_f32_data[reduc], "main", mul_mat_vec_num_bindings, sizeof(vk_mat_vec_push_constants), {rm_int_n(1*rm_kq_int, i), 1, 1}, {wg_size_subgroup_int, rm_int_n(1*rm_kq_int, i), i+1}, 1, true, use_subgroups, subgroup_size_int); ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_q8_1_f32[w][GGML_TYPE_IQ1_S][i], "mul_mat_vec_iq1_s_q8_1_f32", arr_dmmv_iq1_s_q8_1_f32_len[reduc], arr_dmmv_iq1_s_q8_1_f32_data[reduc], "main", mul_mat_vec_num_bindings, sizeof(vk_mat_vec_push_constants), {1*rm_iq_int(i), 1, 1}, {wg_size_subgroup_int, 1*rm_iq_int(i), i+1}, 1, true, use_subgroups, subgroup_size_int); ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_q8_1_f32[w][GGML_TYPE_IQ1_M][i], "mul_mat_vec_iq1_m_q8_1_f32", arr_dmmv_iq1_m_q8_1_f32_len[reduc], arr_dmmv_iq1_m_q8_1_f32_data[reduc], "main", mul_mat_vec_num_bindings, sizeof(vk_mat_vec_push_constants), {1*rm_iq_int(i), 1, 1}, {wg_size_subgroup_int, 1*rm_iq_int(i), i+1}, 1, true, use_subgroups, subgroup_size_int); @@ -5440,20 +5445,20 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { const uint32_t subgroup_size_int = (device->vendor_id == VK_VENDOR_ID_INTEL && device->subgroup_size_control) ? device->subgroup_min_size : device->subgroup_size; const uint32_t wg_size_subgroup_int = (w == DMMV_WG_SIZE_SUBGROUP) ? subgroup_size_int : (subgroup_size_int * 4); - ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_id_q8_1_f32[w][GGML_TYPE_Q2_0], "mul_mat_vec_id_q2_0_q8_1_f32", arr_dmmv_id_q2_0_q8_1_f32_len[reduc], arr_dmmv_id_q2_0_q8_1_f32_data[reduc], "main", mul_mat_vec_id_num_bindings, sizeof(vk_mat_vec_id_push_constants), {2*rm_kq_int, 1, 1}, {wg_size_subgroup_int, 2*rm_kq_int}, 1, true, use_subgroups, subgroup_size_int); - ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_id_q8_1_f32[w][GGML_TYPE_Q4_0], "mul_mat_vec_id_q4_0_q8_1_f32", arr_dmmv_id_q4_0_q8_1_f32_len[reduc], arr_dmmv_id_q4_0_q8_1_f32_data[reduc], "main", mul_mat_vec_id_num_bindings, sizeof(vk_mat_vec_id_push_constants), {1*rm_stdq_int, 1, 1}, {wg_size_subgroup_int, 1*rm_stdq_int}, 1, true, use_subgroups, subgroup_size_int); - ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_id_q8_1_f32[w][GGML_TYPE_Q4_1], "mul_mat_vec_id_q4_1_q8_1_f32", arr_dmmv_id_q4_1_q8_1_f32_len[reduc], arr_dmmv_id_q4_1_q8_1_f32_data[reduc], "main", mul_mat_vec_id_num_bindings, sizeof(vk_mat_vec_id_push_constants), {1*rm_stdq_int, 1, 1}, {wg_size_subgroup_int, 1*rm_stdq_int}, 1, true, use_subgroups, subgroup_size_int); - ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_id_q8_1_f32[w][GGML_TYPE_Q5_0], "mul_mat_vec_id_q5_0_q8_1_f32", arr_dmmv_id_q5_0_q8_1_f32_len[reduc], arr_dmmv_id_q5_0_q8_1_f32_data[reduc], "main", mul_mat_vec_id_num_bindings, sizeof(vk_mat_vec_id_push_constants), {1*rm_stdq_int, 1, 1}, {wg_size_subgroup_int, 1*rm_stdq_int}, 1, true, use_subgroups, subgroup_size_int); - ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_id_q8_1_f32[w][GGML_TYPE_Q5_1], "mul_mat_vec_id_q5_1_q8_1_f32", arr_dmmv_id_q5_1_q8_1_f32_len[reduc], arr_dmmv_id_q5_1_q8_1_f32_data[reduc], "main", mul_mat_vec_id_num_bindings, sizeof(vk_mat_vec_id_push_constants), {1*rm_stdq_int, 1, 1}, {wg_size_subgroup_int, 1*rm_stdq_int}, 1, true, use_subgroups, subgroup_size_int); - ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_id_q8_1_f32[w][GGML_TYPE_Q8_0], "mul_mat_vec_id_q8_0_q8_1_f32", arr_dmmv_id_q8_0_q8_1_f32_len[reduc], arr_dmmv_id_q8_0_q8_1_f32_data[reduc], "main", mul_mat_vec_id_num_bindings, sizeof(vk_mat_vec_id_push_constants), {1*rm_stdq_int, 1, 1}, {wg_size_subgroup_int, 1*rm_stdq_int}, 1, true, use_subgroups, subgroup_size_int); + ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_id_q8_1_f32[w][GGML_TYPE_Q2_0], "mul_mat_vec_id_q2_0_q8_1_f32", arr_dmmv_id_q2_0_q8_1_f32_len[reduc], arr_dmmv_id_q2_0_q8_1_f32_data[reduc], "main", mul_mat_vec_id_num_bindings, sizeof(vk_mat_vec_id_push_constants), {rm_id(2*rm_kq_int), 1, 1}, {wg_size_subgroup_int, rm_id(2*rm_kq_int)}, 1, true, use_subgroups, subgroup_size_int); + ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_id_q8_1_f32[w][GGML_TYPE_Q4_0], "mul_mat_vec_id_q4_0_q8_1_f32", arr_dmmv_id_q4_0_q8_1_f32_len[reduc], arr_dmmv_id_q4_0_q8_1_f32_data[reduc], "main", mul_mat_vec_id_num_bindings, sizeof(vk_mat_vec_id_push_constants), {rm_id(1*rm_stdq_int), 1, 1}, {wg_size_subgroup_int, rm_id(1*rm_stdq_int)}, 1, true, use_subgroups, subgroup_size_int); + ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_id_q8_1_f32[w][GGML_TYPE_Q4_1], "mul_mat_vec_id_q4_1_q8_1_f32", arr_dmmv_id_q4_1_q8_1_f32_len[reduc], arr_dmmv_id_q4_1_q8_1_f32_data[reduc], "main", mul_mat_vec_id_num_bindings, sizeof(vk_mat_vec_id_push_constants), {rm_id(1*rm_stdq_int), 1, 1}, {wg_size_subgroup_int, rm_id(1*rm_stdq_int)}, 1, true, use_subgroups, subgroup_size_int); + ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_id_q8_1_f32[w][GGML_TYPE_Q5_0], "mul_mat_vec_id_q5_0_q8_1_f32", arr_dmmv_id_q5_0_q8_1_f32_len[reduc], arr_dmmv_id_q5_0_q8_1_f32_data[reduc], "main", mul_mat_vec_id_num_bindings, sizeof(vk_mat_vec_id_push_constants), {rm_id(1*rm_stdq_int), 1, 1}, {wg_size_subgroup_int, rm_id(1*rm_stdq_int)}, 1, true, use_subgroups, subgroup_size_int); + ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_id_q8_1_f32[w][GGML_TYPE_Q5_1], "mul_mat_vec_id_q5_1_q8_1_f32", arr_dmmv_id_q5_1_q8_1_f32_len[reduc], arr_dmmv_id_q5_1_q8_1_f32_data[reduc], "main", mul_mat_vec_id_num_bindings, sizeof(vk_mat_vec_id_push_constants), {rm_id(1*rm_stdq_int), 1, 1}, {wg_size_subgroup_int, rm_id(1*rm_stdq_int)}, 1, true, use_subgroups, subgroup_size_int); + ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_id_q8_1_f32[w][GGML_TYPE_Q8_0], "mul_mat_vec_id_q8_0_q8_1_f32", arr_dmmv_id_q8_0_q8_1_f32_len[reduc], arr_dmmv_id_q8_0_q8_1_f32_data[reduc], "main", mul_mat_vec_id_num_bindings, sizeof(vk_mat_vec_id_push_constants), {rm_id(1*rm_stdq_int), 1, 1}, {wg_size_subgroup_int, rm_id(1*rm_stdq_int)}, 1, true, use_subgroups, subgroup_size_int); - ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_id_q8_1_f32[w][GGML_TYPE_MXFP4], "mul_mat_vec_id_mxfp4_q8_1_f32", arr_dmmv_id_mxfp4_q8_1_f32_len[reduc], arr_dmmv_id_mxfp4_q8_1_f32_data[reduc], "main", mul_mat_vec_id_num_bindings, sizeof(vk_mat_vec_id_push_constants), {2*rm_stdq_int, 1, 1}, {wg_size_subgroup_int, 2*rm_stdq_int}, 1, true, use_subgroups, subgroup_size_int); + ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_id_q8_1_f32[w][GGML_TYPE_MXFP4], "mul_mat_vec_id_mxfp4_q8_1_f32", arr_dmmv_id_mxfp4_q8_1_f32_len[reduc], arr_dmmv_id_mxfp4_q8_1_f32_data[reduc], "main", mul_mat_vec_id_num_bindings, sizeof(vk_mat_vec_id_push_constants), {rm_id(2*rm_stdq_int), 1, 1}, {wg_size_subgroup_int, rm_id(2*rm_stdq_int)}, 1, true, use_subgroups, subgroup_size_int); - ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_id_q8_1_f32[w][GGML_TYPE_Q2_K], "mul_mat_vec_id_q2_k_q8_1_f32", arr_dmmv_id_q2_k_q8_1_f32_len[reduc], arr_dmmv_id_q2_k_q8_1_f32_data[reduc], "main", mul_mat_vec_id_num_bindings, sizeof(vk_mat_vec_id_push_constants), {2*rm_kq_int, 1, 1}, {wg_size_subgroup_int, 2*rm_kq_int}, 1, true, use_subgroups, subgroup_size_int); - ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_id_q8_1_f32[w][GGML_TYPE_Q3_K], "mul_mat_vec_id_q3_k_q8_1_f32", arr_dmmv_id_q3_k_q8_1_f32_len[reduc], arr_dmmv_id_q3_k_q8_1_f32_data[reduc], "main", mul_mat_vec_id_num_bindings, sizeof(vk_mat_vec_id_push_constants), {1*rm_kq_int, 1, 1}, {wg_size_subgroup_int, 1*rm_kq_int}, 1, true, use_subgroups, subgroup_size_int); - ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_id_q8_1_f32[w][GGML_TYPE_Q4_K], "mul_mat_vec_id_q4_k_q8_1_f32", arr_dmmv_id_q4_k_q8_1_f32_len[reduc], arr_dmmv_id_q4_k_q8_1_f32_data[reduc], "main", mul_mat_vec_id_num_bindings, sizeof(vk_mat_vec_id_push_constants), {1*rm_kq_int, 1, 1}, {wg_size_subgroup_int, 1*rm_kq_int}, 1, true, use_subgroups, subgroup_size_int); - ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_id_q8_1_f32[w][GGML_TYPE_Q5_K], "mul_mat_vec_id_q5_k_q8_1_f32", arr_dmmv_id_q5_k_q8_1_f32_len[reduc], arr_dmmv_id_q5_k_q8_1_f32_data[reduc], "main", mul_mat_vec_id_num_bindings, sizeof(vk_mat_vec_id_push_constants), {1*rm_kq_int, 1, 1}, {wg_size_subgroup_int, 1*rm_kq_int}, 1, true, use_subgroups, subgroup_size_int); - ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_id_q8_1_f32[w][GGML_TYPE_Q6_K], "mul_mat_vec_id_q6_k_q8_1_f32", arr_dmmv_id_q6_k_q8_1_f32_len[reduc], arr_dmmv_id_q6_k_q8_1_f32_data[reduc], "main", mul_mat_vec_id_num_bindings, sizeof(vk_mat_vec_id_push_constants), {1*rm_kq_int, 1, 1}, {wg_size_subgroup_int, 1*rm_kq_int}, 1, true, use_subgroups, subgroup_size_int); + ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_id_q8_1_f32[w][GGML_TYPE_Q2_K], "mul_mat_vec_id_q2_k_q8_1_f32", arr_dmmv_id_q2_k_q8_1_f32_len[reduc], arr_dmmv_id_q2_k_q8_1_f32_data[reduc], "main", mul_mat_vec_id_num_bindings, sizeof(vk_mat_vec_id_push_constants), {rm_id(2*rm_kq_int), 1, 1}, {wg_size_subgroup_int, rm_id(2*rm_kq_int)}, 1, true, use_subgroups, subgroup_size_int); + ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_id_q8_1_f32[w][GGML_TYPE_Q3_K], "mul_mat_vec_id_q3_k_q8_1_f32", arr_dmmv_id_q3_k_q8_1_f32_len[reduc], arr_dmmv_id_q3_k_q8_1_f32_data[reduc], "main", mul_mat_vec_id_num_bindings, sizeof(vk_mat_vec_id_push_constants), {rm_id(1*rm_kq_int), 1, 1}, {wg_size_subgroup_int, rm_id(1*rm_kq_int)}, 1, true, use_subgroups, subgroup_size_int); + ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_id_q8_1_f32[w][GGML_TYPE_Q4_K], "mul_mat_vec_id_q4_k_q8_1_f32", arr_dmmv_id_q4_k_q8_1_f32_len[reduc], arr_dmmv_id_q4_k_q8_1_f32_data[reduc], "main", mul_mat_vec_id_num_bindings, sizeof(vk_mat_vec_id_push_constants), {rm_id(1*rm_kq_int), 1, 1}, {wg_size_subgroup_int, rm_id(1*rm_kq_int)}, 1, true, use_subgroups, subgroup_size_int); + ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_id_q8_1_f32[w][GGML_TYPE_Q5_K], "mul_mat_vec_id_q5_k_q8_1_f32", arr_dmmv_id_q5_k_q8_1_f32_len[reduc], arr_dmmv_id_q5_k_q8_1_f32_data[reduc], "main", mul_mat_vec_id_num_bindings, sizeof(vk_mat_vec_id_push_constants), {rm_id(1*rm_kq_int), 1, 1}, {wg_size_subgroup_int, rm_id(1*rm_kq_int)}, 1, true, use_subgroups, subgroup_size_int); + ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_id_q8_1_f32[w][GGML_TYPE_Q6_K], "mul_mat_vec_id_q6_k_q8_1_f32", arr_dmmv_id_q6_k_q8_1_f32_len[reduc], arr_dmmv_id_q6_k_q8_1_f32_data[reduc], "main", mul_mat_vec_id_num_bindings, sizeof(vk_mat_vec_id_push_constants), {rm_id(1*rm_kq_int), 1, 1}, {wg_size_subgroup_int, rm_id(1*rm_kq_int)}, 1, true, use_subgroups, subgroup_size_int); ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_id_q8_1_f32[w][GGML_TYPE_IQ1_S], "mul_mat_vec_id_iq1_s_q8_1_f32", arr_dmmv_id_iq1_s_q8_1_f32_len[reduc], arr_dmmv_id_iq1_s_q8_1_f32_data[reduc], "main", mul_mat_vec_id_num_bindings, sizeof(vk_mat_vec_id_push_constants), {1*rm_iq_int(0), 1, 1}, {wg_size_subgroup_int, 1*rm_iq_int(0)}, 1, true, use_subgroups, subgroup_size_int); ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_id_q8_1_f32[w][GGML_TYPE_IQ1_M], "mul_mat_vec_id_iq1_m_q8_1_f32", arr_dmmv_id_iq1_m_q8_1_f32_len[reduc], arr_dmmv_id_iq1_m_q8_1_f32_data[reduc], "main", mul_mat_vec_id_num_bindings, sizeof(vk_mat_vec_id_push_constants), {1*rm_iq_int(0), 1, 1}, {wg_size_subgroup_int, 1*rm_iq_int(0)}, 1, true, use_subgroups, subgroup_size_int); @@ -5467,6 +5472,9 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { #if !defined(GGML_VULKAN_INTEGER_DOT_GLSLC_SUPPORT) GGML_UNUSED(rm_stdq_int); GGML_UNUSED(rm_kq_int); + GGML_UNUSED(is_rdna3); + GGML_UNUSED(rm_int_n); + GGML_UNUSED(rm_id); GGML_UNUSED(rm_iq_int); #endif From 76a51e82d7cb548eecc880e108fdfc0d7040be57 Mon Sep 17 00:00:00 2001 From: Neo Zhang Date: Mon, 31 Aug 2026 18:33:02 +0800 Subject: [PATCH 053/104] sycl : Enhance to get the free memory of Intel GPU (llama/27968) * enhance get mem info by l0 an SYCL API * remove debug code, format the code * update SYCL.md for GGML_SYCL_GET_MEM_API --- ggml/src/ggml-sycl/base.hpp | 36 +++++++ ggml/src/ggml-sycl/common.hpp | 16 +-- ggml/src/ggml-sycl/ggml-sycl.cpp | 60 +++++++++--- ggml/src/ggml-sycl/mem.cpp | 162 +++++++++++++++++++++++++++++++ ggml/src/ggml-sycl/mem.hpp | 16 +++ 5 files changed, 260 insertions(+), 30 deletions(-) create mode 100644 ggml/src/ggml-sycl/base.hpp create mode 100644 ggml/src/ggml-sycl/mem.cpp create mode 100644 ggml/src/ggml-sycl/mem.hpp diff --git a/ggml/src/ggml-sycl/base.hpp b/ggml/src/ggml-sycl/base.hpp new file mode 100644 index 00000000000..3afd57ccb2d --- /dev/null +++ b/ggml/src/ggml-sycl/base.hpp @@ -0,0 +1,36 @@ +#ifndef GGML_SYCL_BASE_HPP +#define GGML_SYCL_BASE_HPP + +/** + * Module: base + * + * Description: + * Provides zero-dependency, foundational primitives, core abstractions, + * and low-level system interfaces. This module acts as the lowest layer + * of the architecture and is consumed globally across all subsystems. + * + * Constraints: + * - STRICTLY zero upstream dependencies (leaf module). + * - High stability and backward compatibility required. + */ + +#include + +extern int g_ggml_sycl_debug; + +#if defined(__clang__) && __has_builtin(__builtin_expect) +// Hint the optimizer to pipeline the more likely following instruction in branches +# define LIKELY(expr) __builtin_expect(expr, true) +# define UNLIKELY(expr) __builtin_expect(expr, false) +#else +# define LIKELY(expr) (expr) +# define UNLIKELY(expr) (expr) +#endif + +#define GGML_SYCL_DEBUG(...) \ + do { \ + if (UNLIKELY(g_ggml_sycl_debug)) \ + fprintf(stderr, __VA_ARGS__); \ + } while (0) + +#endif // GGML_SYCL_BASE_HPP diff --git a/ggml/src/ggml-sycl/common.hpp b/ggml/src/ggml-sycl/common.hpp index 34de284d83a..b5f75ca548c 100644 --- a/ggml/src/ggml-sycl/common.hpp +++ b/ggml/src/ggml-sycl/common.hpp @@ -18,6 +18,7 @@ #include #include +#include "base.hpp" #include "dpct/helper.hpp" #include "ggml.h" #include "ggml-impl.h" @@ -69,21 +70,6 @@ extern int g_ggml_sycl_fa_onednn; extern int g_ggml_sycl_fa_onednn_max_kv; -#if defined(__clang__) && __has_builtin(__builtin_expect) -// Hint the optimizer to pipeline the more likely following instruction in branches -# define LIKELY(expr) __builtin_expect(expr, true) -# define UNLIKELY(expr) __builtin_expect(expr, false) -#else -# define LIKELY(expr) (expr) -# define UNLIKELY(expr) (expr) -#endif - -#define GGML_SYCL_DEBUG(...) \ - do { \ - if (UNLIKELY(g_ggml_sycl_debug)) \ - fprintf(stderr, __VA_ARGS__); \ - } while (0) - #define CHECK_TRY_ERROR(expr) \ [&]() { \ try { \ diff --git a/ggml/src/ggml-sycl/ggml-sycl.cpp b/ggml/src/ggml-sycl/ggml-sycl.cpp index 2b2d26cf235..290fb46767c 100644 --- a/ggml/src/ggml-sycl/ggml-sycl.cpp +++ b/ggml/src/ggml-sycl/ggml-sycl.cpp @@ -35,6 +35,7 @@ #include #ifdef GGML_SYCL_SUPPORT_LEVEL_ZERO_API #include +#include #endif #if defined(GGML_SYCL_GRAPH) && SYCL_EXT_ONEAPI_ASYNC_MEMORY_ALLOC # include @@ -61,6 +62,7 @@ #include "ggml-sycl/fwht.hpp" #include "ggml-sycl/gemm.hpp" #include "ggml-sycl/getrows.hpp" +#include "ggml-sycl/mem.hpp" #include "ggml-sycl/norm.hpp" #include "ggml-sycl/presets.hpp" #include "ggml-sycl/quantize.hpp" @@ -105,6 +107,8 @@ int g_ggml_sycl_enable_flash_attention = 1; int g_ggml_sycl_dev2dev_memcpy = DEV2DEV_MEMCPY_SYCL; int g_ggml_sycl_usm_system = 0; int g_ggml_sycl_enable_host_pinned_mem = 1; +int g_ggml_sycl_get_mem_api = MEMORY_API_TYPE_LEVEL_ZERO; + static ggml_sycl_device_info ggml_sycl_init() { ggml_sycl_device_info info = {}; @@ -301,10 +305,27 @@ static const char* dev2dev_int2str(int dev2dev) { } } +/* +* There are several entry APIs to be called as first function in SYCL backend in different cases. +* It's the first internal function to be called by them in SYCL backend. +* This function is used to do initialize work for the SYCL backend and set the global variables. +*/ +void initialize_sycl_begining() { +#ifdef GGML_SYCL_SUPPORT_LEVEL_ZERO_API + ze_result_t zes_init = zesInit(0); + if (zes_init != ZE_RESULT_SUCCESS) { + std::cerr << "Warning: zesInit failed [ggml_check_sycl] with code " << static_cast(zes_init) + << ". Sysman free-memory query may be unavailable.\n"; + } +#endif +} + static void ggml_check_sycl() try { static bool initialized = false; if (!initialized) { + initialize_sycl_begining(); + g_ggml_sycl_debug = ggml_sycl_get_env("GGML_SYCL_DEBUG", 0); g_ggml_sycl_enable_optimize = ggml_sycl_get_env("GGML_SYCL_ENABLE_OPT", 1); g_ggml_sycl_enable_graph = ggml_sycl_get_env("GGML_SYCL_ENABLE_GRAPH", 0); @@ -317,8 +338,11 @@ static void ggml_check_sycl() try { g_ggml_sycl_prioritize_dmmv = ggml_sycl_get_env("GGML_SYCL_PRIORITIZE_DMMV", 0); g_ggml_sycl_dev2dev_memcpy = ggml_sycl_get_env("GGML_SYCL_DEV2DEV_MEMCPY", DEV2DEV_MEMCPY_SYCL); + g_ggml_sycl_get_mem_api = ggml_sycl_get_env("GGML_SYCL_GET_MEM_API", MEMORY_API_TYPE_LEVEL_ZERO); + if (g_ggml_sycl_use_level_zero_api == 0) { g_ggml_sycl_dev2dev_memcpy = DEV2DEV_MEMCPY_SYCL; + g_ggml_sycl_get_mem_api = MEMORY_API_TYPE_SYCL; } #ifdef SYCL_FLASH_ATTN @@ -331,6 +355,7 @@ static void ggml_check_sycl() try { g_ggml_sycl_enable_host_pinned_mem = ggml_sycl_get_env("GGML_SYCL_ENABLE_HOST_PINNED_MEM", 1); + GGML_SYCL_DEBUG("[SYCL] call ggml_check_sycl\n"); GGML_LOG_INFO("Build with Macros:\n"); @@ -374,9 +399,12 @@ static void ggml_check_sycl() try { #ifdef GGML_SYCL_SUPPORT_LEVEL_ZERO_API GGML_LOG_INFO(" GGML_SYCL_DEV2DEV_MEMCPY: %d (%s)\n", g_ggml_sycl_dev2dev_memcpy, dev2dev_int2str(g_ggml_sycl_dev2dev_memcpy)); + GGML_LOG_INFO(" GGML_SYCL_GET_MEM_API: %d (%s)\n", g_ggml_sycl_get_mem_api, mem_api_int2str(g_ggml_sycl_get_mem_api)); #else GGML_LOG_INFO(" GGML_SYCL_DEV2DEV_MEMCPY: %d (%s), enable to SYCL API since missing GGML_SYCL_SUPPORT_LEVEL_ZERO_API\n", g_ggml_sycl_dev2dev_memcpy, dev2dev_int2str(g_ggml_sycl_dev2dev_memcpy)); + GGML_LOG_INFO(" GGML_SYCL_GET_MEM_API: %d (%s), enable to SYCL API since missing GGML_SYCL_SUPPORT_LEVEL_ZERO_API\n", + g_ggml_sycl_get_mem_api, mem_api_int2str(g_ggml_sycl_get_mem_api)); #endif #if defined(GGML_SYCL_DNNL) @@ -5208,6 +5236,7 @@ catch (sycl::exception const &exc) { static bool ggml_sycl_compute_forward(ggml_backend_sycl_context & ctx, struct ggml_tensor * dst) try { if (!g_sycl_loaded) return false; + initialize_sycl_begining(); if (dst->src[0] != nullptr && ggml_backend_buffer_is_sycl_split(dst->src[0]->buffer)) { ggml_sycl_set_peer_access(dst->src[1]->ne[1], ctx.device); @@ -5590,18 +5619,16 @@ catch (sycl::exception const &exc) { std::exit(1); } -void ggml_backend_sycl_get_device_memory(int device, size_t *free, - size_t *total) try { +void ggml_backend_sycl_get_device_memory(int device, size_t * free, size_t * total) try { GGML_SYCL_DEBUG("[SYCL] call ggml_backend_sycl_get_device_memory\n"); - ggml_sycl_set_device(device); - - SYCL_CHECK(CHECK_TRY_ERROR( - dpct::dev_mgr::instance().get_device(device).get_memory_info(*free, *total))); -} -catch (sycl::exception const &exc) { - std::cerr << exc.what() << "Exception caught at file:" << __FILE__ - << ", line:" << __LINE__ << std::endl; - std::exit(1); + bool res = get_memory_size(dpct::dev_mgr::instance().get_device(device), *free, *total, + (MemoryAPIType) g_ggml_sycl_get_mem_api); + if (!res) { + GGML_ABORT("[%s] failed to get device memory size", __func__); + } +} catch (const sycl::exception & exc) { + std::cerr << exc.what() << "Exception caught at file:" << __FILE__ << ", line:" << __LINE__ << std::endl; + std::exit(1); } //////////////////////////////////////////////////////////////////////////////// @@ -6020,10 +6047,12 @@ static const char * ggml_backend_sycl_device_get_description(ggml_backend_dev_t } static void ggml_backend_sycl_device_get_memory(ggml_backend_dev_t dev, size_t * free, size_t * total) { - ggml_backend_sycl_device_context * ctx = (ggml_backend_sycl_device_context *)dev->context; - ggml_sycl_set_device(ctx->device); - SYCL_CHECK(CHECK_TRY_ERROR( - dpct::dev_mgr::instance().get_device(ctx->device).get_memory_info(*free, *total))); + ggml_backend_sycl_device_context * ctx = (ggml_backend_sycl_device_context *) dev->context; + bool res = get_memory_size(dpct::dev_mgr::instance().get_device(ctx->device), *free, *total, + (MemoryAPIType) g_ggml_sycl_get_mem_api); + if (!res) { + GGML_ABORT("[%s] failed to get device memory size", __func__); + } } static enum ggml_backend_dev_type ggml_backend_sycl_device_get_type(ggml_backend_dev_t dev) { @@ -6906,6 +6935,7 @@ ggml_backend_reg_t ggml_backend_sycl_reg() { static std::mutex mutex; std::lock_guard lock(mutex); if (!initialized) { + initialize_sycl_begining(); ggml_backend_sycl_reg_context * ctx = new ggml_backend_sycl_reg_context; const int min_batch_size = getenv("GGML_OP_OFFLOAD_MIN_BATCH") ? atoi(getenv("GGML_OP_OFFLOAD_MIN_BATCH")) : 32; diff --git a/ggml/src/ggml-sycl/mem.cpp b/ggml/src/ggml-sycl/mem.cpp new file mode 100644 index 00000000000..5ec466420e0 --- /dev/null +++ b/ggml/src/ggml-sycl/mem.cpp @@ -0,0 +1,162 @@ +#include +#include + +#ifdef GGML_SYCL_SUPPORT_LEVEL_ZERO_API +#include +#include +#endif + +#include "base.hpp" +#include "mem.hpp" + +#include +#include +#include + +const char * mem_api_int2str(int mem_api) { + if (mem_api == MEMORY_API_TYPE_SYCL) { + return "SYCL API"; + } else if (mem_api == MEMORY_API_TYPE_LEVEL_ZERO) { + return "Level Zero API"; + } else { + return "Unknown"; + } +} + +#ifdef GGML_SYCL_SUPPORT_LEVEL_ZERO_API +bool query_free_memory_by_ze(sycl::device dev, size_t & free_bytes, size_t & total_bytes) { + free_bytes = 0; + total_bytes = 0; + + uint32_t module_count = 0; + +#if defined(SYCL_EXT_ONEAPI_BACKEND_LEVEL_ZERO) + constexpr sycl::backend kL0Backend = sycl::backend::ext_oneapi_level_zero; +#else + constexpr sycl::backend kL0Backend = sycl::backend::level_zero; +#endif + + try { + ze_result_t zes_init = zesInit(0); + if (zes_init != ZE_RESULT_SUCCESS) { + std::cerr << "Warning: zesInit failed with code " << static_cast(zes_init) + << ". Sysman free-memory query may be unavailable.\n"; + } + + if (dev.get_platform().get_backend() != kL0Backend) { + GGML_SYCL_DEBUG("Device backend is not Level Zero; falling back to SYCL memory query.\n"); + total_bytes = dev.get_info(); + free_bytes = total_bytes; + return false; + } + + ze_device_handle_t ze_dev = sycl::get_native(dev); + if (ze_dev == nullptr) { + GGML_SYCL_DEBUG("Level Zero device handle is null; falling back to SYCL memory query.\n"); + total_bytes = dev.get_info(); + free_bytes = total_bytes; + return false; + } + + ze_result_t r = zesDeviceEnumMemoryModules(ze_dev, &module_count, nullptr); + if (r != ZE_RESULT_SUCCESS || module_count == 0) { + GGML_SYCL_DEBUG("Failed to enumerate Level Zero memory modules. Falling back to SYCL memory query.\n"); + total_bytes = dev.get_info(); + free_bytes = total_bytes; + return false; + } + + std::vector modules(module_count); + r = zesDeviceEnumMemoryModules(ze_dev, &module_count, modules.data()); + if (r != ZE_RESULT_SUCCESS || module_count == 0) { + GGML_SYCL_DEBUG("Failed to enumerate Level Zero memory modules. Falling back to SYCL memory query.\n"); + total_bytes = dev.get_info(); + free_bytes = total_bytes; + return false; + } + + for (uint32_t i = 0; i < module_count; ++i) { + zes_mem_state_t state = {}; + state.stype = ZES_STRUCTURE_TYPE_MEM_STATE; + state.pNext = nullptr; + + r = zesMemoryGetState(modules[i], &state); + if (r != ZE_RESULT_SUCCESS) { + continue; + } + + free_bytes += state.free; + total_bytes += state.size; + } + + if (total_bytes == 0) { + GGML_SYCL_DEBUG("Level Zero memory query returned zero total bytes. Falling back to SYCL memory query.\n"); + total_bytes = dev.get_info(); + free_bytes = total_bytes; + return false; + } + return true; + } catch (const sycl::exception & e) { + GGML_SYCL_DEBUG("Level Zero memory query failed: %s\n", e.what()); + total_bytes = dev.get_info(); + free_bytes = total_bytes; + return false; + } +} +#endif + +bool get_memory_size_by_sycl_api(sycl::device dev, size_t & free_bytes, size_t & total_bytes) { + GGML_SYCL_DEBUG("[%s]Querying free memory using SYCL API.\n", __func__); + total_bytes = dev.get_info(); + +#if (defined(__SYCL_COMPILER_VERSION) && __SYCL_COMPILER_VERSION >= 20221105) + if (dev.has(sycl::aspect::ext_intel_free_memory)) { + try { + GGML_SYCL_DEBUG("Querying free memory using SYCL aspect::ext_intel_free_memory."); + free_bytes = dev.get_info(); + return true; + } catch (const sycl::exception &) { + GGML_SYCL_DEBUG( + "Failed to query free memory using SYCL aspect::ext_intel_free_memory. Using total memory as free " + "memory."); + free_bytes = total_bytes; + return false; + } + } else { + GGML_SYCL_DEBUG( + "Device does not support SYCL aspect::ext_intel_free_memory. Using total memory as free memory."); + free_bytes = total_bytes; + } +#else + GGML_SYCL_DEBUG("SYCL Compiler version is older than 20221105. Using total memory as free memory."); + free_bytes = total_bytes; +#endif + return true; +} + +bool get_memory_size(sycl::device dev, size_t & free_bytes, size_t & total_bytes, MemoryAPIType api_type) { + const auto name = dev.get_info(); + const auto vendor = dev.get_info(); + const auto global_mem = dev.get_info(); + + GGML_SYCL_DEBUG("[%s]GPU Name: %s\n", __func__, name.c_str()); + GGML_SYCL_DEBUG("[%s]GPU Vendor: %s\n", __func__, vendor.c_str()); + GGML_SYCL_DEBUG("[%s]GPU Global Memory: %zu bytes\n", __func__, static_cast(global_mem)); + + if (api_type == MEMORY_API_TYPE_LEVEL_ZERO) { +#ifdef GGML_SYCL_SUPPORT_LEVEL_ZERO_API + GGML_SYCL_DEBUG("[%s]Querying free memory using Level Zero API.\n", __func__); + if (!query_free_memory_by_ze(dev, free_bytes, total_bytes)) { + //fallback to SYCL API if Level Zero API fails + GGML_SYCL_DEBUG("[%s]Falling back to SYCL API for memory query.\n", __func__); + return get_memory_size_by_sycl_api(dev, free_bytes, total_bytes); + } + return true; +#else + GGML_SYCL_DEBUG("[%s]Level Zero API support is not enabled. Please enable it to use this feature.\n", __func__); + return false; +#endif + } else { //MEMORY_API_TYPE_SYCL + return get_memory_size_by_sycl_api(dev, free_bytes, total_bytes); + } +} diff --git a/ggml/src/ggml-sycl/mem.hpp b/ggml/src/ggml-sycl/mem.hpp new file mode 100644 index 00000000000..b3e45cfea04 --- /dev/null +++ b/ggml/src/ggml-sycl/mem.hpp @@ -0,0 +1,16 @@ +#ifndef GGML_SYCL_MEM_HPP +#define GGML_SYCL_MEM_HPP + +#include + +enum MemoryAPIType { + MEMORY_API_TYPE_LEVEL_ZERO = 0, + MEMORY_API_TYPE_SYCL = 1, +}; + +const char* mem_api_int2str(int mem_api); + +bool get_memory_size(sycl::device dev, size_t & free_bytes, size_t & total_bytes, + MemoryAPIType api_type); + +#endif // GGML_SYCL_MEM_HPP From b0f4bc02ed1f15a0d63eca3cc8af7e4d6a481d59 Mon Sep 17 00:00:00 2001 From: ynankani Date: Mon, 31 Aug 2026 11:22:28 +0000 Subject: [PATCH 054/104] CUDA: extend MOE fusion to specdec, earlier MOE glu fusion and topk-router fusion were restricted to 1 token (llama/27621) * CUDA: extend MOE fusion to specdec, earlier MOE glu fusion and topk-router fusion were resticted to 1 token Signed-off-by: ynankani * Address review comments Signed-off-by: ynankani * Add SWIGLU_CLAMP case to multi-token moe fusion Signed-off-by: ynankani --------- Signed-off-by: ynankani --- ggml/src/ggml-cuda/ggml-cuda.cu | 11 +-- ggml/src/ggml-cuda/mmvq.cu | 115 +++++++++++++++++++++++++++++--- ggml/src/ggml-cuda/topk-moe.cu | 24 ++++--- ggml/src/ggml-cuda/topk-moe.cuh | 3 + 4 files changed, 127 insertions(+), 26 deletions(-) diff --git a/ggml/src/ggml-cuda/ggml-cuda.cu b/ggml/src/ggml-cuda/ggml-cuda.cu index a6fc655c41c..53eccdd6502 100644 --- a/ggml/src/ggml-cuda/ggml-cuda.cu +++ b/ggml/src/ggml-cuda/ggml-cuda.cu @@ -1807,7 +1807,7 @@ static bool ggml_cuda_should_fuse_mul_mat_vec_q(const ggml_tensor * tensor) { return false; } - if (tensor->op == GGML_OP_MUL_MAT_ID && dst->ne[2] != 1) { + if (tensor->op == GGML_OP_MUL_MAT_ID && dst->ne[2] > get_mmvq_mmid_max_batch(src0->type, cc)) { return false; } @@ -2983,9 +2983,10 @@ static bool ggml_cuda_check_fusion_memory_ranges(const ggml_cgraph * cgraph, }; bool is_ok = true; - // exception for topk-moe, as each row is read entirely before writing - if (ggml_nrows(cgraph->nodes[node_idx]) == 1 && is_topk_moe) { - return true; + // one block reads all logits before it writes, so logits may alias the out nodes + const ggml_tensor * logits_may_alias = nullptr; + if (is_topk_moe && ggml_nrows(cgraph->nodes[node_idx]) <= TOPK_MOE_ROWS_PER_BLOCK) { + logits_may_alias = cgraph->nodes[node_idx]->src[0]; } for (int i = 0; i < out_count; ++i) { @@ -2999,7 +3000,7 @@ static bool ggml_cuda_check_fusion_memory_ranges(const ggml_cgraph * cgraph, for (int src_idx = 0; src_idx < GGML_MAX_SRC; ++src_idx) { const ggml_tensor * src = cgraph->nodes[j]->src[src_idx]; - if (!src || src->op == GGML_OP_NONE) { + if (!src || src->op == GGML_OP_NONE || src == logits_may_alias) { continue; } diff --git a/ggml/src/ggml-cuda/mmvq.cu b/ggml/src/ggml-cuda/mmvq.cu index 79f7a3f6fe7..2be2f249108 100644 --- a/ggml/src/ggml-cuda/mmvq.cu +++ b/ggml/src/ggml-cuda/mmvq.cu @@ -773,10 +773,10 @@ static __global__ void mul_mat_vec_q( // Grid: (ceil(nrows_x / c_rows_per_block), nchannels_dst) // Block: (warp_size, ncols_dst) - each warp handles one token independently. // No shared memory reduction needed since each warp works alone. -template +template __launch_bounds__(get_mmvq_mmid_max_batch_for_device()*ggml_cuda_get_physical_warp_size(), 1) static __global__ void mul_mat_vec_q_moe( - const void * vx_ptr, const void * vy_ptr, const int32_t * ids_ptr, + const void * vx_ptr, const void * vy_ptr, const int32_t * ids_ptr, const ggml_cuda_mm_fusion_args_device fusion, float * dst_ptr, const uint32_t ncols_x, const uint3 nchannels_y, const uint32_t nrows_x, const uint32_t stride_row_x, const uint32_t stride_col_y, const uint32_t stride_col_dst, @@ -794,6 +794,29 @@ static __global__ void mul_mat_vec_q_moe( constexpr vec_dot_q_cuda_t vec_dot_q_cuda = get_vec_dot_q_cuda(type); + // fuse gate, bias, scales, and glu_op into the up projection + bool use_gate = false; + const void * vgate = nullptr; + const float * x_bias = nullptr; + const float * gate_bias = nullptr; + const float * x_scale = nullptr; + const float * gate_scale = nullptr; + ggml_glu_op active_glu = GGML_GLU_OP_SWIGLU; + float glu_limit = 0.0f; + + if constexpr (has_fusion) { + use_gate = fusion.gate != nullptr; + vgate = fusion.gate; + x_bias = (const float *) fusion.x_bias; + gate_bias = (const float *) fusion.gate_bias; + active_glu = fusion.glu_op; + glu_limit = fusion.glu_limit; + if constexpr (type == GGML_TYPE_NVFP4) { + x_scale = (const float *) fusion.x_scale; + gate_scale = (const float *) fusion.gate_scale; + } + } + const uint32_t token_idx = threadIdx.y; const int row0 = c_rows_per_block*blockIdx.x; const int blocks_per_row_x = ncols_x / qk; @@ -814,6 +837,7 @@ static __global__ void mul_mat_vec_q_moe( // partial sum for each thread float tmp[c_rows_per_block] = {0.0f}; + float tmp_gate[c_rows_per_block] = {0.0f}; for (int kbx = threadIdx.x / (qi/vdr); kbx < blocks_per_row_x; kbx += blocks_per_iter) { const int kby = kbx * (qk/QK8_1); @@ -822,6 +846,11 @@ static __global__ void mul_mat_vec_q_moe( #pragma unroll for (int i = 0; i < c_rows_per_block; ++i) { tmp[i] += vec_dot_q_cuda(vx, &y[kby], kbx_offset + i*stride_row_x + kbx, kqs); + if constexpr (has_fusion) { + if (use_gate) { + tmp_gate[i] += vec_dot_q_cuda(vgate, &y[kby], kbx_offset + i*stride_row_x + kbx, kqs); + } + } } } @@ -831,11 +860,63 @@ static __global__ void mul_mat_vec_q_moe( #pragma unroll for (int i = 0; i < c_rows_per_block; ++i) { tmp[i] = warp_reduce_sum(tmp[i]); + if constexpr (has_fusion) { + if (use_gate) { + tmp_gate[i] = warp_reduce_sum(tmp_gate[i]); + } + } } // Write results if (threadIdx.x < c_rows_per_block && (c_rows_per_block == 1 || uint32_t(row0 + threadIdx.x) < nrows_x)) { - dst[channel_dst*stride_channel_dst + token_idx*stride_col_dst + row0 + threadIdx.x] = tmp[threadIdx.x]; + float result = tmp[threadIdx.x]; + if constexpr (has_fusion) { + const uint32_t bias_idx = channel_x*stride_channel_dst + row0 + threadIdx.x; + + if constexpr (type == GGML_TYPE_NVFP4) { + if (x_scale) { + result *= x_scale[channel_x]; + } + } + if (x_bias) { + result += x_bias[bias_idx]; + } + if (use_gate) { + float gate_value = tmp_gate[threadIdx.x]; + if constexpr (type == GGML_TYPE_NVFP4) { + if (gate_scale) { + gate_value *= gate_scale[channel_x]; + } + } + if (gate_bias) { + gate_value += gate_bias[bias_idx]; + } + switch (active_glu) { + case GGML_GLU_OP_SWIGLU: + result *= ggml_cuda_op_silu_single(gate_value); + break; + case GGML_GLU_OP_GEGLU: + result *= ggml_cuda_op_gelu_single(gate_value); + break; + case GGML_GLU_OP_SWIGLU_OAI: + result = ggml_cuda_op_swiglu_oai_single(gate_value, result); + break; + case GGML_GLU_OP_SWIGLU_CLAMP: + result = ggml_cuda_op_swiglu_clamp_single(gate_value, result, glu_limit); + break; + default: + result = result * gate_value; + break; + } + } + } + dst[channel_dst*stride_channel_dst + token_idx*stride_col_dst + row0 + threadIdx.x] = result; + } + + if constexpr (!has_fusion) { + GGML_UNUSED_VARS(use_gate, tmp_gate, vgate, x_bias, gate_bias, active_glu, glu_limit, x_scale, gate_scale); + } else if constexpr (type != GGML_TYPE_NVFP4) { + GGML_UNUSED_VARS(x_scale, gate_scale); } } @@ -885,7 +966,7 @@ static void mul_mat_vec_q_switch_fusion( template static void mul_mat_vec_q_moe_launch( - const void * vx, const void * vy, const int32_t * ids, float * dst, + const void * vx, const void * vy, const int32_t * ids, const ggml_cuda_mm_fusion_args_device fusion, float * dst, const uint32_t ncols_x, const uint3 nchannels_y, const uint32_t nrows_x, const uint32_t stride_row_x, const uint32_t stride_col_y, const uint32_t stride_col_dst, const uint32_t stride_channel_x, const uint32_t stride_channel_y, const uint32_t stride_channel_dst, @@ -898,11 +979,22 @@ static void mul_mat_vec_q_moe_launch( const dim3 block_dims(warp_size, ncols_dst); const ggml_cuda_kernel_launch_params launch_params = ggml_cuda_kernel_launch_params(block_nums, block_dims, 0, stream); - ggml_cuda_kernel_launch(mul_mat_vec_q_moe, launch_params, - vx, vy, ids, dst, ncols_x, nchannels_y, nrows_x, - stride_row_x, stride_col_y, stride_col_dst, - stride_channel_x, stride_channel_y, stride_channel_dst, - ncols_dst, ids_stride); + const bool has_fusion = fusion.gate != nullptr || fusion.x_bias != nullptr || fusion.gate_bias != nullptr || + fusion.x_scale != nullptr || fusion.gate_scale != nullptr; + + if (has_fusion) { + ggml_cuda_kernel_launch(mul_mat_vec_q_moe, launch_params, + vx, vy, ids, fusion, dst, ncols_x, nchannels_y, nrows_x, + stride_row_x, stride_col_y, stride_col_dst, + stride_channel_x, stride_channel_y, stride_channel_dst, + ncols_dst, ids_stride); + } else { + ggml_cuda_kernel_launch(mul_mat_vec_q_moe, launch_params, + vx, vy, ids, fusion, dst, ncols_x, nchannels_y, nrows_x, + stride_row_x, stride_col_y, stride_col_dst, + stride_channel_x, stride_channel_y, stride_channel_dst, + ncols_dst, ids_stride); + } } template @@ -998,7 +1090,7 @@ static void mul_mat_vec_q_switch_ncols_dst( if (has_ids && ncols_dst > 1) { // Multi-token MUL_MAT_ID path - dedicated MoE kernel mul_mat_vec_q_moe_launch( - vx, vy, ids, dst, ncols_x, nchannels_y_fd, nrows_x, + vx, vy, ids, fusion, dst, ncols_x, nchannels_y_fd, nrows_x, stride_row_x, stride_col_y, stride_col_dst, stride_channel_x, stride_channel_y, stride_channel_dst, ncols_dst, ids_stride, warp_size, nchannels_dst, stream); @@ -1280,7 +1372,8 @@ void ggml_cuda_mul_mat_vec_q( ggml_cuda_mm_fusion_args_device fusion_local{}; if (fusion) { - GGML_ASSERT( !ids || dst->ne[2] == 1); + const int cc = ggml_cuda_info().devices[ggml_cuda_get_device()].cc; + GGML_ASSERT( !ids || dst->ne[2] <= get_mmvq_mmid_max_batch(src0->type, cc)); GGML_ASSERT( ids || dst->ne[1] == 1); // Scale fusion is only allowed for NVFP4 currently as the cost of checking this at run-time in the prologue is // non-negligible for some models such as gpt-oss-20b diff --git a/ggml/src/ggml-cuda/topk-moe.cu b/ggml/src/ggml-cuda/topk-moe.cu index c8cec70bb32..dadcd601cb4 100644 --- a/ggml/src/ggml-cuda/topk-moe.cu +++ b/ggml/src/ggml-cuda/topk-moe.cu @@ -88,15 +88,16 @@ __device__ void sqrt_softplus_warp_inplace(float (&vals)[experts_per_thread], co It is intended as fusion of softmax->top-k->get_rows pipeline for MoE models */ template -__launch_bounds__(4 * WARP_SIZE, 1) __global__ void topk_moe_cuda(const float * logits, - float * weights, - int32_t * ids, - float * bias, - const int n_rows, - const int n_expert_used, - const float clamp_val, - const float scale_val, - const topk_moe_config config) { +__launch_bounds__(TOPK_MOE_ROWS_PER_BLOCK * WARP_SIZE, 1) +__global__ void topk_moe_cuda(const float * logits, + float * weights, + int32_t * ids, + float * bias, + const int n_rows, + const int n_expert_used, + const float clamp_val, + const float scale_val, + const topk_moe_config config) { const int row = blockIdx.x * blockDim.y + threadIdx.y; if (row >= n_rows) { return; @@ -123,6 +124,9 @@ __launch_bounds__(4 * WARP_SIZE, 1) __global__ void topk_moe_cuda(const float * wt[i / WARP_SIZE] = (n_experts % WARP_SIZE == 0 || expert < n_experts) ? logits[expert] : -INFINITY; } + // Weights and IDs can alias logits, so wait until every row in the block reads its logits. + __syncthreads(); + if (!config.delayed_softmax) { if (config.use_sigmoid) { sigmoid_warp_inplace(wt, n_experts, threadIdx.x); @@ -282,7 +286,7 @@ static void launch_topk_moe_cuda(ggml_backend_cuda_context & ctx, const topk_moe_config config) { GGML_ASSERT(!(config.with_norm && config.delayed_softmax) && "delayed softmax is not supported with weight normalization"); - const int rows_per_block = 4; + const int rows_per_block = TOPK_MOE_ROWS_PER_BLOCK; dim3 grid_dims((n_rows + rows_per_block - 1) / rows_per_block, 1, 1); dim3 block_dims(WARP_SIZE, rows_per_block, 1); cudaStream_t stream = ctx.stream(); diff --git a/ggml/src/ggml-cuda/topk-moe.cuh b/ggml/src/ggml-cuda/topk-moe.cuh index 091ef02a415..061b37e2971 100644 --- a/ggml/src/ggml-cuda/topk-moe.cuh +++ b/ggml/src/ggml-cuda/topk-moe.cuh @@ -3,6 +3,9 @@ #include +// Rows that one CUDA block handles. +#define TOPK_MOE_ROWS_PER_BLOCK 8 + struct ggml_cuda_topk_moe_args { bool sigmoid{}; bool sqrt_softplus{}; From 7614a4c1398708ea1a6575d0e3b659f14bb2f7e3 Mon Sep 17 00:00:00 2001 From: Niklas Wenzel Date: Mon, 31 Aug 2026 13:58:55 +0200 Subject: [PATCH 055/104] metal : add fa-vec tunings for M1 (llama/28078) --- ggml/src/ggml-metal/ggml-metal-tuning.cpp | 190 ++++++++++++++++++++++ 1 file changed, 190 insertions(+) diff --git a/ggml/src/ggml-metal/ggml-metal-tuning.cpp b/ggml/src/ggml-metal/ggml-metal-tuning.cpp index 4f26fd9c373..6a742bd0fd0 100644 --- a/ggml/src/ggml-metal/ggml-metal-tuning.cpp +++ b/ggml/src/ggml-metal/ggml-metal-tuning.cpp @@ -68,6 +68,196 @@ fa_vec_cfg_t fa_vec_baseline_cfg(int dk, int dv) { // sweep and paste its output. See ggml-metal-tuning.h for the row/lookup semantics. // ref: https://github.com/ggml-org/llama.cpp/pull/27824 constexpr fa_vec_entry_t fa_vec_tuned_table[] = { + { { GGML_METAL_DEVICE_M1, GGML_TYPE_F16, 32, 32, 2, 3 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_F16, 64, 64, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_F16, 64, 64, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_F16, 128, 128, 2, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_F16, 128, 128, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_F16, 128, 128, 1, 1 }, { 1, 1 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_F16, 128, 128, 1, 2 }, { 1, 1 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_F16, 128, 128, 1, 3 }, { 1, 1 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_F16, 128, 128, 1, 4 }, { 1, 1 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_F16, 192, 128, 1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_F16, 192, 128, 1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_F16, 192, 128, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_F16, 192, 128, 1, 3 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_F16, 192, 128, 1, 4 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_F16, 320, 256, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_F16, 320, 256, 1, 1 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_F16, 320, 256, 1, 2 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_F16, 320, 256, 1, 3 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_F16, 320, 256, 1, 4 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q4_0, 32, 32, -1, 1 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q4_0, 32, 32, 1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q4_0, 32, 32, 1, 4 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q4_0, 32, 32, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q4_0, 32, 32, 2, 4 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q4_0, 32, 32, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q4_0, 32, 32, 3, 4 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q4_0, 64, 64, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q4_0, 64, 64, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q4_0, 64, 64, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q4_0, 64, 64, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q4_0, 96, 96, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q4_0, 96, 96, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q4_0, 96, 96, 1, 4 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q4_0, 96, 96, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q4_0, 96, 96, 2, 4 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q4_0, 96, 96, 3, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q4_0, 128, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q4_0, 128, 128, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q4_0, 192, 192, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q4_0, 192, 192, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q4_0, 192, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q4_0, 192, 128, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q4_0, 256, 256, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q4_0, 256, 256, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q4_0, 320, 256, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q4_0, 320, 256, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q4_0, 512, 512, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q4_0, 512, 512, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q4_0, 576, 512, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q4_0, 576, 512, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q4_1, 32, 32, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q4_1, 32, 32, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q4_1, 32, 32, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q4_1, 32, 32, 2, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q4_1, 32, 32, 3, 2 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q4_1, 32, 32, 3, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q4_1, 64, 64, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q4_1, 64, 64, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q4_1, 64, 64, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q4_1, 96, 96, 1, 3 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q4_1, 96, 96, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q4_1, 96, 96, 2, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q4_1, 96, 96, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q4_1, 96, 96, 3, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q4_1, 128, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q4_1, 128, 128, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q4_1, 128, 128, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q4_1, 128, 128, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q4_1, 192, 192, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q4_1, 192, 192, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q4_1, 192, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q4_1, 192, 128, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q4_1, 256, 256, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q4_1, 256, 256, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q4_1, 320, 256, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q4_1, 320, 256, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q4_1, 512, 512, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q4_1, 512, 512, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q4_1, 576, 512, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q4_1, 576, 512, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q5_0, 32, 32, -1, 1 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q5_0, 32, 32, 1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q5_0, 32, 32, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q5_0, 32, 32, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q5_0, 64, 64, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q5_0, 64, 64, -1, 1 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q5_0, 64, 64, 1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q5_0, 64, 64, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q5_0, 64, 64, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q5_0, 96, 96, -1, 1 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q5_0, 96, 96, 1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q5_0, 96, 96, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q5_0, 96, 96, 2, 4 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q5_0, 96, 96, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q5_0, 96, 96, 3, 4 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q5_0, 128, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q5_0, 128, 128, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q5_0, 128, 128, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q5_0, 192, 192, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q5_0, 192, 192, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q5_0, 192, 192, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q5_0, 192, 192, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q5_0, 192, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q5_0, 192, 128, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q5_0, 192, 128, 1, 3 }, { 4, 2 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q5_0, 192, 128, 2, 3 }, { 4, 2 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q5_0, 192, 128, 2, 4 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q5_0, 192, 128, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q5_0, 192, 128, 3, 3 }, { 4, 2 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q5_0, 192, 128, 3, 4 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q5_0, 256, 256, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q5_0, 256, 256, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q5_0, 320, 256, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q5_0, 320, 256, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q5_0, 512, 512, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q5_0, 512, 512, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q5_0, 576, 512, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q5_0, 576, 512, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q5_1, 32, 32, -1, 1 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q5_1, 32, 32, 1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q5_1, 32, 32, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q5_1, 32, 32, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q5_1, 64, 64, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q5_1, 64, 64, -1, 1 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q5_1, 64, 64, 1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q5_1, 64, 64, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q5_1, 64, 64, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q5_1, 96, 96, -1, 1 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q5_1, 96, 96, 1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q5_1, 96, 96, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q5_1, 96, 96, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q5_1, 128, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q5_1, 128, 128, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q5_1, 128, 128, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q5_1, 192, 192, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q5_1, 192, 192, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q5_1, 192, 192, 1, 3 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q5_1, 192, 192, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q5_1, 192, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q5_1, 192, 128, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q5_1, 192, 128, 1, 3 }, { 4, 2 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q5_1, 192, 128, 2, 3 }, { 4, 2 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q5_1, 192, 128, 3, 3 }, { 4, 2 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q5_1, 256, 256, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q5_1, 256, 256, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q5_1, 320, 256, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q5_1, 320, 256, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q5_1, 512, 512, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q5_1, 512, 512, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q5_1, 512, 512, 3, 3 }, { 1, 1 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q5_1, 576, 512, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q5_1, 576, 512, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q8_0, 32, 32, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q8_0, 32, 32, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q8_0, 32, 32, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q8_0, 32, 32, 2, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q8_0, 32, 32, 2, 4 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q8_0, 32, 32, 3, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q8_0, 32, 32, 3, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q8_0, 64, 64, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q8_0, 64, 64, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q8_0, 64, 64, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q8_0, 64, 64, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q8_0, 64, 64, 3, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q8_0, 96, 96, 1, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q8_0, 96, 96, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q8_0, 96, 96, 2, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q8_0, 96, 96, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q8_0, 96, 96, 3, 3 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q8_0, 128, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q8_0, 128, 128, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q8_0, 192, 192, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q8_0, 192, 192, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q8_0, 192, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q8_0, 192, 128, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q8_0, 256, 256, -1, 0 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q8_0, 256, 256, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q8_0, 320, 256, 1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q8_0, 320, 256, 1, 3 }, { 4, 2 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q8_0, 320, 256, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q8_0, 320, 256, 2, 3 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q8_0, 320, 256, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q8_0, 320, 256, 3, 3 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q8_0, 512, 512, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q8_0, 512, 512, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q8_0, 512, 512, 2, 2 }, { 1, 1 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q8_0, 512, 512, 3, 2 }, { 1, 1 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q8_0, 576, 512, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1, GGML_TYPE_Q8_0, 576, 512, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_PRO, GGML_TYPE_F16, 32, 32, 3, 3 }, { 2, 4 } }, { { GGML_METAL_DEVICE_M1_PRO, GGML_TYPE_F16, 64, 64, -1, 0 }, { 1, 4 } }, { { GGML_METAL_DEVICE_M1_PRO, GGML_TYPE_F16, 64, 64, -1, 1 }, { 1, 4 } }, From c1be45b89016bb11b1460618c97b1e077adcd9a7 Mon Sep 17 00:00:00 2001 From: Jaden_Mach <88880593+jadenmach2@users.noreply.github.com> Date: Mon, 31 Aug 2026 09:00:04 -0400 Subject: [PATCH 056/104] ROCm: add radix TOP_K for long rows (llama/27466) * ROCm: add radix TOP_K for long rows --- ggml/src/ggml-cuda/ggml-cuda.cu | 5 + ggml/src/ggml-cuda/top-k.cu | 180 +++++++++++++++++++++++++++++++- 2 files changed, 180 insertions(+), 5 deletions(-) diff --git a/ggml/src/ggml-cuda/ggml-cuda.cu b/ggml/src/ggml-cuda/ggml-cuda.cu index 53eccdd6502..31f5aeeacc3 100644 --- a/ggml/src/ggml-cuda/ggml-cuda.cu +++ b/ggml/src/ggml-cuda/ggml-cuda.cu @@ -5273,6 +5273,11 @@ static bool ggml_backend_cuda_device_supports_op(ggml_backend_dev_t dev, const g case GGML_OP_SUM: return ggml_is_contiguous_rows(op->src[0]); case GGML_OP_TOP_K: +#if defined(GGML_USE_HIP) || defined(GGML_CUDA_USE_CUB) + return true; +#else + return op->src[0]->ne[0] <= 1024; +#endif // defined(GGML_USE_HIP) || defined(GGML_CUDA_USE_CUB) case GGML_OP_ARGSORT: #ifndef GGML_CUDA_USE_CUB return op->src[0]->ne[0] <= 1024; diff --git a/ggml/src/ggml-cuda/top-k.cu b/ggml/src/ggml-cuda/top-k.cu index 9681cd29333..c7a0c831788 100644 --- a/ggml/src/ggml-cuda/top-k.cu +++ b/ggml/src/ggml-cuda/top-k.cu @@ -48,6 +48,168 @@ static int next_power_of_2(int x) { #endif // CUB_TOP_K_AVAILABLE +#if !defined(GGML_CUDA_USE_CUB) && defined(GGML_USE_HIP) + +static __device__ __forceinline__ uint32_t top_k_float_to_ordered(float value) { + const uint32_t bits = __float_as_uint(value); + const uint32_t mask = (uint32_t) (-(int32_t) (bits >> 31)) | 0x80000000U; + return bits ^ mask; +} + +struct top_k_radix_state { + uint32_t prefix; + uint32_t prefix_mask; + int rank; + int greater_count; + int equal_count; +}; + +static __global__ void top_k_radix_init(top_k_radix_state * states, int nrows, int k) { + const int row = blockIdx.x * blockDim.x + threadIdx.x; + if (row < nrows) { + states[row] = {0, 0, k, 0, 0}; + } +} + +template +static __global__ void top_k_radix_histogram( + const float * __restrict__ src, + const top_k_radix_state * __restrict__ states, + int * __restrict__ block_histograms, + int ncols, + int blocks_per_row, + int shift) { + constexpr int NBINS = 1 << RADIX_BITS; + + const int row = blockIdx.x / blocks_per_row; + const int row_block = blockIdx.x % blocks_per_row; + const int tid = threadIdx.x; + const float * row_src = src + (size_t) row * ncols; + __shared__ int histogram[NBINS]; + + histogram[tid] = 0; + __syncthreads(); + + const top_k_radix_state state = states[row]; + for (int col = row_block * BLOCK_SIZE + tid; + col < ncols; + col += blocks_per_row * BLOCK_SIZE) { + const uint32_t key = top_k_float_to_ordered(row_src[col]); + if ((key & state.prefix_mask) == state.prefix) { + atomicAdd(&histogram[(key >> shift) & (NBINS - 1)], 1); + } + } + __syncthreads(); + + const size_t histogram_offset = + ((size_t) row * blocks_per_row + row_block) * NBINS; + block_histograms[histogram_offset + tid] = histogram[tid]; +} + +template +static __global__ void top_k_radix_select( + const int * __restrict__ block_histograms, + top_k_radix_state * __restrict__ states, + int blocks_per_row, + int shift) { + constexpr int NBINS = 1 << RADIX_BITS; + + const int row = blockIdx.x; + const int tid = threadIdx.x; + __shared__ int histogram[NBINS]; + + int count = 0; + for (int row_block = 0; row_block < blocks_per_row; ++row_block) { + const size_t offset = ((size_t) row * blocks_per_row + row_block) * NBINS; + count += block_histograms[offset + tid]; + } + histogram[tid] = count; + __syncthreads(); + + if (tid == 0) { + top_k_radix_state state = states[row]; + int bin = NBINS - 1; + while (bin > 0 && histogram[bin] < state.rank) { + state.rank -= histogram[bin--]; + } + state.prefix |= (uint32_t) bin << shift; + state.prefix_mask |= (uint32_t) (NBINS - 1) << shift; + states[row] = state; + } +} + +static __global__ void top_k_radix_reset_counters(top_k_radix_state * states, int nrows) { + const int row = blockIdx.x * blockDim.x + threadIdx.x; + if (row < nrows) { + states[row].greater_count = 0; + states[row].equal_count = 0; + } +} + +template +static __global__ void top_k_radix_gather( + const float * __restrict__ src, + int * __restrict__ dst, + top_k_radix_state * __restrict__ states, + int ncols, + int k, + int blocks_per_row) { + const int row = blockIdx.x / blocks_per_row; + const int row_block = blockIdx.x % blocks_per_row; + const int tid = threadIdx.x; + const float * row_src = src + (size_t) row * ncols; + int * row_dst = dst + (size_t) row * k; + top_k_radix_state * state = &states[row]; + + for (int col = row_block * BLOCK_SIZE + tid; + col < ncols; + col += blocks_per_row * BLOCK_SIZE) { + const uint32_t key = top_k_float_to_ordered(row_src[col]); + if (key > state->prefix) { + const int pos = atomicAdd(&state->greater_count, 1); + row_dst[pos] = col; + } else if (key == state->prefix) { + const int pos = atomicAdd(&state->equal_count, 1); + if (pos < state->rank) { + row_dst[k - state->rank + pos] = col; + } + } + } +} + +static void top_k_radix_cuda( + ggml_cuda_pool & pool, + const float * src, int * dst, int ncols, int nrows, int k, cudaStream_t stream) { + constexpr int BLOCK_SIZE = 256; + constexpr int RADIX_BITS = 8; + constexpr int NBINS = 1 << RADIX_BITS; + const int blocks_per_row = std::min((ncols + 1023) / 1024, 64); + + ggml_cuda_pool_alloc states_alloc(pool, nrows); + ggml_cuda_pool_alloc histograms_alloc(pool, (size_t) nrows * blocks_per_row * NBINS); + top_k_radix_state * states = states_alloc.get(); + int * histograms = histograms_alloc.get(); + + top_k_radix_init<<<(nrows + BLOCK_SIZE - 1) / BLOCK_SIZE, BLOCK_SIZE, 0, stream>>>(states, nrows, k); + + const dim3 row_grid(blocks_per_row * nrows); + for (int shift = 32 - RADIX_BITS; shift >= 0; shift -= RADIX_BITS) { + top_k_radix_histogram + <<>>( + src, states, histograms, ncols, blocks_per_row, shift); + top_k_radix_select + <<>>(histograms, states, blocks_per_row, shift); + } + + top_k_radix_reset_counters + <<<(nrows + BLOCK_SIZE - 1) / BLOCK_SIZE, BLOCK_SIZE, 0, stream>>>(states, nrows); + top_k_radix_gather + <<>>( + src, dst, states, ncols, k, blocks_per_row); +} + +#endif // !defined(GGML_CUDA_USE_CUB) && defined(GGML_USE_HIP) + void ggml_cuda_op_top_k(ggml_backend_cuda_context & ctx, ggml_tensor * dst) { const ggml_tensor * src0 = dst->src[0]; const float * src0_d = (const float *) src0->data; @@ -96,10 +258,18 @@ void ggml_cuda_op_top_k(ggml_backend_cuda_context & ctx, ggml_tensor * dst) { dst_d += k * iter_nrows; } #else // GGML_CUDA_USE_CUB - ggml_cuda_pool_alloc temp_dst_alloc(pool, ncols * nrows); - int * tmp_dst = temp_dst_alloc.get(); - argsort_f32_i32_cuda_bitonic(src0_d, tmp_dst, ncols, nrows, GGML_SORT_ORDER_DESC, stream); - CUDA_CHECK(cudaMemcpy2DAsync(dst_d, k * sizeof(int), tmp_dst, ncols * sizeof(int), k * sizeof(int), nrows, - cudaMemcpyDeviceToDevice, stream)); +#if defined(GGML_USE_HIP) + if (ncols > 1024) { + top_k_radix_cuda(pool, src0_d, dst_d, ncols, nrows, k, stream); + } else { +#endif // defined(GGML_USE_HIP) + ggml_cuda_pool_alloc temp_dst_alloc(pool, ncols * nrows); + int * tmp_dst = temp_dst_alloc.get(); + argsort_f32_i32_cuda_bitonic(src0_d, tmp_dst, ncols, nrows, GGML_SORT_ORDER_DESC, stream); + CUDA_CHECK(cudaMemcpy2DAsync(dst_d, k * sizeof(int), tmp_dst, ncols * sizeof(int), k * sizeof(int), nrows, + cudaMemcpyDeviceToDevice, stream)); +#if defined(GGML_USE_HIP) + } +#endif // defined(GGML_USE_HIP) #endif } From c648b9a4d0fbb0dd3db9d19149714f77cef43af2 Mon Sep 17 00:00:00 2001 From: fairydreaming <166155368+fairydreaming@users.noreply.github.com> Date: Mon, 31 Aug 2026 16:04:38 +0200 Subject: [PATCH 057/104] webgpu : avoid crash when offset is not multiple of 4 in WebGPU ggml_backend_tensor_get() implementation (llama/28045) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit * webgpu : avoid crash when offset is not multiple of 4 in WebGPU ggml_backend_tensor_get() implementation * chore : improve code readability Co-authored-by: Sigbjørn Skjæret --------- Co-authored-by: Stanisław Szymczyk Co-authored-by: Sigbjørn Skjæret --- ggml/src/ggml-webgpu/ggml-webgpu.cpp | 15 +++++++++++---- 1 file changed, 11 insertions(+), 4 deletions(-) diff --git a/ggml/src/ggml-webgpu/ggml-webgpu.cpp b/ggml/src/ggml-webgpu/ggml-webgpu.cpp index b953118a7a3..1a43c72733b 100644 --- a/ggml/src/ggml-webgpu/ggml-webgpu.cpp +++ b/ggml/src/ggml-webgpu/ggml-webgpu.cpp @@ -3713,11 +3713,18 @@ static void ggml_backend_webgpu_buffer_get_tensor(ggml_backend_buffer_t buffer, size_t total_offset = ggml_webgpu_tensor_offset(tensor) + offset; - size_t final_size = size; - if (size % 4 != 0) { + size_t local_offset = total_offset % 4; + if (local_offset != 0) { + // If offset is not a multiple of 4, we need to round it down to the previous + // multiple of 4 + total_offset -= local_offset; + } + + size_t final_size = size + local_offset; + if (final_size % 4 != 0) { // If size is not a multiple of 4, we need to round it up to the next // multiple of 4 - final_size = size + (4 - (size % 4)); + final_size += 4 - (final_size % 4); } std::lock_guard lock(buf_ctx->global_ctx->mutex); @@ -3748,7 +3755,7 @@ static void ggml_backend_webgpu_buffer_get_tensor(ggml_backend_buffer_t buffer, const void * mapped_range = buf_ctx->global_ctx->get_tensor_staging_buf.GetConstMappedRange(0, final_size); // Copy the data from the mapped range to the output buffer - std::memcpy(data, mapped_range, size); + std::memcpy(data, (const void *) ((const char *) mapped_range + local_offset), size); buf_ctx->global_ctx->get_tensor_staging_buf.Unmap(); WEBGPU_CPU_PROFILE_TOTAL_END(get_tensor, buf_ctx->global_ctx); } From 088c603e2951f3bf5b7ce59438523c541580faa1 Mon Sep 17 00:00:00 2001 From: Hongqiang Wang Date: Mon, 31 Aug 2026 08:56:22 -0700 Subject: [PATCH 058/104] opencl: tune the quant paths for Intel Xe-LP GPUs to improve its TG and PP performance (llama/26438) * opencl: Q4_K/Q5_K mul_mv N_DST 4->8 on Intel for 2x activation reuse * opencl: Q4_K mul_mm 8x8 tile fot Intel * opencl: Q5_K mul_mm 8x8 tile for Intel * opencl: Q4_K mul_mv N_DST 8->16 for Intel --- ggml/src/ggml-opencl/ggml-opencl.cpp | 9 +++++---- ggml/src/ggml-opencl/kernels/mul_mm_q4_k_f32_l4_lm.cl | 10 ++++++++++ ggml/src/ggml-opencl/kernels/mul_mm_q5_k_f32_l4_lm.cl | 10 ++++++++++ ggml/src/ggml-opencl/kernels/mul_mv_q4_k_f32_flat.cl | 2 +- ggml/src/ggml-opencl/kernels/mul_mv_q5_k_f32_flat.cl | 2 +- 5 files changed, 27 insertions(+), 6 deletions(-) diff --git a/ggml/src/ggml-opencl/ggml-opencl.cpp b/ggml/src/ggml-opencl/ggml-opencl.cpp index 90635cc858f..34d58f4ee81 100644 --- a/ggml/src/ggml-opencl/ggml-opencl.cpp +++ b/ggml/src/ggml-opencl/ggml-opencl.cpp @@ -20002,7 +20002,8 @@ static void ggml_cl_mul_mat(ggml_backend_t backend, const ggml_tensor * src0, co } kernel = backend_ctx->kernel_mul_mm_q4_k_f32_l4_lm; - nth0 = 128; // calculated as (BM*BN)/(TM*TN) + // (BM*BN)/(TM*TN): Intel uses an 8x8 microtile (WG=64), others 4x8 (WG=128) + nth0 = (backend_ctx->gpu_family == INTEL) ? 64 : 128; int batch_stride_a = ne00*ne01; int batch_stride_b = ne10*ne11; @@ -20046,7 +20047,7 @@ static void ggml_cl_mul_mat(ggml_backend_t backend, const ggml_tensor * src0, co } kernel = backend_ctx->kernel_mul_mm_q5_k_f32_l4_lm; - nth0 = 128; // calculated as (BM*BN)/(TM*TN) + nth0 = (backend_ctx->gpu_family == INTEL) ? 64 : 128; // Intel 8x8 microtile int batch_stride_a = ne00*ne01; int batch_stride_b = ne10*ne11; @@ -20860,7 +20861,7 @@ static void ggml_cl_mul_mat(ggml_backend_t backend, const ggml_tensor * src0, co if (backend_ctx->gpu_family == INTEL) { nth0 = 16; nth1 = 1; - ndst = 4; + ndst = 16; // 8->16 rows per subgroup — matches N_DST in mul_mv_q4_k_f32_flat.cl (32 spills) } else if (backend_ctx->gpu_family == ADRENO) { nth0 = 64; nth1 = 2; @@ -20934,7 +20935,7 @@ static void ggml_cl_mul_mat(ggml_backend_t backend, const ggml_tensor * src0, co if (backend_ctx->gpu_family == INTEL) { nth0 = 16; nth1 = 1; - ndst = 4; + ndst = 8; // 4->8 rows per subgroup (2x activation reuse) } else if (backend_ctx->gpu_family == ADRENO) { nth0 = 64; nth1 = 2; diff --git a/ggml/src/ggml-opencl/kernels/mul_mm_q4_k_f32_l4_lm.cl b/ggml/src/ggml-opencl/kernels/mul_mm_q4_k_f32_l4_lm.cl index 2235b1ae838..a9c649a5213 100644 --- a/ggml/src/ggml-opencl/kernels/mul_mm_q4_k_f32_l4_lm.cl +++ b/ggml/src/ggml-opencl/kernels/mul_mm_q4_k_f32_l4_lm.cl @@ -1,13 +1,23 @@ #pragma OPENCL EXTENSION cl_khr_fp16 : enable +#ifdef cl_intel_required_subgroup_size +#define INTEL_GPU 1 +#endif + #define LOAD_VEC_A 4 #define LOAD_VEC_B 4 #define BM 64 #define BN 64 #define BK 32 +#ifdef INTEL_GPU +// Intel Xe iGPU: 8x8 microtile (WG = BM*BN/(TM*TN) = 64) — ~+12% pp512 vs 4x8 +#define TM 8 +#define TN 8 +#else #define TM 4 #define TN 8 +#endif kernel void kernel_mul_mm_q4_k_f32_l4_lm( global uchar4 * src0_q, diff --git a/ggml/src/ggml-opencl/kernels/mul_mm_q5_k_f32_l4_lm.cl b/ggml/src/ggml-opencl/kernels/mul_mm_q5_k_f32_l4_lm.cl index 8e191f57e83..a343b5c4c62 100644 --- a/ggml/src/ggml-opencl/kernels/mul_mm_q5_k_f32_l4_lm.cl +++ b/ggml/src/ggml-opencl/kernels/mul_mm_q5_k_f32_l4_lm.cl @@ -1,13 +1,23 @@ #pragma OPENCL EXTENSION cl_khr_fp16 : enable +#ifdef cl_intel_required_subgroup_size +#define INTEL_GPU 1 +#endif + #define LOAD_VEC_A 4 #define LOAD_VEC_B 4 #define BM 64 #define BN 64 #define BK 32 +#ifdef INTEL_GPU +// Intel Xe iGPU: 8x8 microtile (WG=64) +#define TM 8 +#define TN 8 +#else #define TM 4 #define TN 8 +#endif kernel void kernel_mul_mm_q5_k_f32_l4_lm( global uchar4 * src0_q, diff --git a/ggml/src/ggml-opencl/kernels/mul_mv_q4_k_f32_flat.cl b/ggml/src/ggml-opencl/kernels/mul_mv_q4_k_f32_flat.cl index 70391866ca6..5316bd36361 100644 --- a/ggml/src/ggml-opencl/kernels/mul_mv_q4_k_f32_flat.cl +++ b/ggml/src/ggml-opencl/kernels/mul_mv_q4_k_f32_flat.cl @@ -40,7 +40,7 @@ typedef struct { #undef N_SIMDWIDTH #ifdef INTEL_GPU -#define N_DST 4 // number of rows each SIMD group works on +#define N_DST 16 // number of rows each SIMD group works on (Intel: 8->16, 2x further activation reuse; 32 spills registers) #define N_SIMDGROUP 1 // number of SIMD groups in a thread group #define N_SIMDWIDTH 16 // SIMD group size #elif defined (ADRENO_GPU) diff --git a/ggml/src/ggml-opencl/kernels/mul_mv_q5_k_f32_flat.cl b/ggml/src/ggml-opencl/kernels/mul_mv_q5_k_f32_flat.cl index 6020364b5c3..ab2e1fab8bd 100644 --- a/ggml/src/ggml-opencl/kernels/mul_mv_q5_k_f32_flat.cl +++ b/ggml/src/ggml-opencl/kernels/mul_mv_q5_k_f32_flat.cl @@ -38,7 +38,7 @@ typedef struct { #undef N_SIMDWIDTH #ifdef INTEL_GPU -#define N_DST 4 +#define N_DST 8 // Intel: 4->8 for 2x activation reuse (see mul_mv_q4_k_f32_flat.cl) #define N_SIMDGROUP 1 #define N_SIMDWIDTH 16 #elif defined(ADRENO_GPU) From c6934d0fcf562ada5493208f6aa5804947f37b64 Mon Sep 17 00:00:00 2001 From: Georgi Gerganov Date: Mon, 31 Aug 2026 21:31:53 +0300 Subject: [PATCH 059/104] metal : add top-k radix implementation (llama/28073) Assisted-by: DeepSeek-v4-Flash-0731 --- ggml/src/ggml-metal/ggml-metal-device.cpp | 19 +++- ggml/src/ggml-metal/ggml-metal-device.h | 1 + ggml/src/ggml-metal/ggml-metal-impl.h | 11 +++ ggml/src/ggml-metal/ggml-metal-ops.cpp | 72 ++++++++++++++- ggml/src/ggml-metal/kernels/argsort.metal | 105 ++++++++++++++++++++++ 5 files changed, 206 insertions(+), 2 deletions(-) diff --git a/ggml/src/ggml-metal/ggml-metal-device.cpp b/ggml/src/ggml-metal/ggml-metal-device.cpp index 5a2f01f1dba..b8d2ef9ce27 100644 --- a/ggml/src/ggml-metal/ggml-metal-device.cpp +++ b/ggml/src/ggml-metal/ggml-metal-device.cpp @@ -1336,7 +1336,7 @@ ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_fwht(ggml_metal_ return res; } -// note: reuse the argsort kernel for top_k +// note: reuse the argsort kernel for the bitonic top_k fallback ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_top_k(ggml_metal_library_t lib, const ggml_tensor * op) { assert(op->op == GGML_OP_TOP_K); @@ -1364,6 +1364,23 @@ ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_top_k(ggml_metal return res; } +ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_top_k_radix(ggml_metal_library_t lib, const ggml_tensor * op) { + assert(op->op == GGML_OP_TOP_K); + + char base[256]; + char name[256]; + + snprintf(base, 256, "kernel_top_k_%s_%s", ggml_type_name(op->src[0]->type), ggml_type_name(op->type)); + snprintf(name, 256, "%s", base); + + 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); + } + + return res; +} + ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_top_k_merge(ggml_metal_library_t lib, const ggml_tensor * op) { assert(op->op == GGML_OP_TOP_K); diff --git a/ggml/src/ggml-metal/ggml-metal-device.h b/ggml/src/ggml-metal/ggml-metal-device.h index 003b688dbac..7f6520103d1 100644 --- a/ggml/src/ggml-metal/ggml-metal-device.h +++ b/ggml/src/ggml-metal/ggml-metal-device.h @@ -145,6 +145,7 @@ struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_argsort struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_argsort_merge (ggml_metal_library_t lib, const struct ggml_tensor * op); struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_fwht (ggml_metal_library_t lib, int n); 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_bin_one (ggml_metal_library_t lib, enum ggml_op op); diff --git a/ggml/src/ggml-metal/ggml-metal-impl.h b/ggml/src/ggml-metal/ggml-metal-impl.h index 49102afe9c0..bdcd9c9e3d8 100644 --- a/ggml/src/ggml-metal/ggml-metal-impl.h +++ b/ggml/src/ggml-metal/ggml-metal-impl.h @@ -1189,6 +1189,17 @@ typedef struct { int32_t len; } ggml_metal_kargs_argsort_merge; +typedef struct { + int32_t ne00; // number of columns (elements per row) + int32_t ne01; // rows + int32_t ne02; + int32_t ne03; + uint64_t nb01; // row stride in src0 + uint64_t nb02; + uint64_t nb03; + int32_t top_k; // k +} ggml_metal_kargs_top_k; + typedef struct { int32_t nrows; } ggml_metal_kargs_fwht; diff --git a/ggml/src/ggml-metal/ggml-metal-ops.cpp b/ggml/src/ggml-metal/ggml-metal-ops.cpp index 7671d1d0156..30ea4ec2745 100644 --- a/ggml/src/ggml-metal/ggml-metal-ops.cpp +++ b/ggml/src/ggml-metal/ggml-metal-ops.cpp @@ -5091,7 +5091,9 @@ int ggml_metal_op_argsort(ggml_metal_op_t ctx, int idx) { return 1; } -int ggml_metal_op_top_k(ggml_metal_op_t ctx, int idx) { +// bitonic-sort + merge fallback: efficient when k is small and there are few rows, +// where the single-workgroup-per-row radix-select cannot reach enough parallelism +static void ggml_metal_op_top_k_bitonic(ggml_metal_op_t ctx, int idx) { ggml_tensor * op = ctx->node(idx); ggml_metal_library_t lib = ctx->lib; @@ -5199,6 +5201,74 @@ int ggml_metal_op_top_k(ggml_metal_op_t ctx, int idx) { len <<= 1; } +} + +// radix-select: one workgroup per row. Maps each float to an order-preserving unsigned +// key, finds the k-th largest via 4 radix-8 histogram passes, then compacts the top-k +// indices. Fast for large k and/or many rows. +static void ggml_metal_op_top_k_radix(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_ASSERT(ggml_is_contiguous_rows(op->src[0])); + + GGML_TENSOR_LOCALS( int32_t, ne0, op->src[0], ne); + GGML_TENSOR_LOCALS(uint64_t, nb0, op->src[0], nb); + + auto pipeline = ggml_metal_library_get_pipeline_top_k_radix(lib, op); + + // one workgroup per row; radix-select the k-th largest value + const int nth = std::min(1024, ggml_metal_pipeline_max_theads_per_threadgroup(pipeline)); + + ggml_metal_kargs_top_k args = { + /*.ne00 =*/ ne00, + /*.ne01 =*/ ne01, + /*.ne02 =*/ ne02, + /*.ne03 =*/ ne03, + /*.nb01 =*/ nb01, + /*.nb02 =*/ nb02, + /*.nb03 =*/ nb03, + /*.top_k =*/ (int32_t) op->ne[0], + }; + + // shared memory: 256-entry histogram + bucket/above scalars + output counter + const size_t smem_histo = GGML_PAD(256*sizeof(uint32_t), 16); + const size_t smem_bucket = GGML_PAD( sizeof(uint32_t), 16); + const size_t smem_above = GGML_PAD( sizeof(uint32_t), 16); + const size_t smem_out = GGML_PAD( sizeof(uint32_t), 16); + + 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), 2); + + ggml_metal_encoder_set_threadgroup_memory_size(enc, smem_histo, 0); + ggml_metal_encoder_set_threadgroup_memory_size(enc, smem_bucket, 1); + ggml_metal_encoder_set_threadgroup_memory_size(enc, smem_above, 2); + ggml_metal_encoder_set_threadgroup_memory_size(enc, smem_out, 3); + + ggml_metal_encoder_dispatch_threadgroups(enc, ne01, ne02, ne03, nth, 1, 1); +} + +int ggml_metal_op_top_k(ggml_metal_op_t ctx, int idx) { + ggml_tensor * op = ctx->node(idx); + + // radix-select has a fixed single-workgroup-per-row cost (~50-60us) that is only + // amortized for long rows, many rows, or a large k; otherwise the bitonic path wins + const int ncols = op->src[0]->ne[0]; + const int k = op->ne[0]; + const int nrows = ggml_nrows(op->src[0]); + + const bool use_radix = + ncols > 2048 && (k > 64 || (nrows > 4 && ncols >= 8192)); + + if (use_radix) { + ggml_metal_op_top_k_radix(ctx, idx); + } else { + ggml_metal_op_top_k_bitonic(ctx, idx); + } return 1; } diff --git a/ggml/src/ggml-metal/kernels/argsort.metal b/ggml/src/ggml-metal/kernels/argsort.metal index 7d144fbd755..e81d194c339 100644 --- a/ggml/src/ggml-metal/kernels/argsort.metal +++ b/ggml/src/ggml-metal/kernels/argsort.metal @@ -230,3 +230,108 @@ kernel void kernel_argsort_merge_f32_i32( template [[host_name("kernel_argsort_merge_f32_i32_asc")]] kernel argsort_merge_t kernel_argsort_merge_f32_i32; template [[host_name("kernel_argsort_merge_f32_i32_desc")]] kernel argsort_merge_t kernel_argsort_merge_f32_i32; + +static inline uint ggml_top_k_f2ui(float x) { + uint y = as_type(x); + if ((y & 0x80000000u) != 0u) { + y ^= 0xFFFFFFFFu; // negative floats: flip all bits + } else { + y |= 0x80000000u; // positive floats: set the sign bit + } + return y; +} + +kernel void kernel_top_k_f32_i32( + constant ggml_metal_kargs_top_k & args, + device const char * src0, + device int32_t * dst, + threadgroup atomic_uint * histo [[threadgroup(0)]], + threadgroup uint * sh_bucket [[threadgroup(1)]], + threadgroup uint * sh_above [[threadgroup(2)]], + threadgroup atomic_uint * out_count [[threadgroup(3)]], + uint3 tgpig[[threadgroup_position_in_grid]], + ushort3 tpitg[[thread_position_in_threadgroup]], + ushort3 ntg[[threads_per_threadgroup]]) { + + const uint ncols = args.ne00; + const uint top_k = args.top_k; + const uint i01 = tgpig[0]; + const uint i02 = tgpig[1]; + const uint i03 = tgpig[2]; + + device const float * src0_row = (device const float *) (src0 + args.nb01*i01 + args.nb02*i02 + args.nb03*i03); + + device int32_t * dst_row = dst + top_k*(i01 + args.ne01*i02 + args.ne01*args.ne02*i03); + + const uint tid = tpitg.x; + const uint ntg_x = ntg.x; + + uint prefix = 0; // fixed high bits of the threshold key + uint desired = top_k; // count still needed from the candidate range + + for (int shift = 24; shift >= 0; shift -= 8) { + for (uint i = tid; i < 256; i += ntg_x) { + atomic_store_explicit(&histo[i], 0u, memory_order_relaxed); + } + threadgroup_barrier(mem_flags::mem_threadgroup); + + const uint hi_mask = (shift + 8 >= 32) ? 0u : (0xFFFFFFFFu << uint(shift + 8)); + const uint prefix_hi = prefix & hi_mask; + + for (uint i = tid; i < ncols; i += ntg_x) { + const uint key = ggml_top_k_f2ui(src0_row[i]); + if ((key & hi_mask) == prefix_hi) { + atomic_fetch_add_explicit(&histo[(key >> uint(shift)) & 0xFFu], 1u, memory_order_relaxed); + } + } + threadgroup_barrier(mem_flags::mem_threadgroup); + + // top-down scan for the bucket holding the k-th value + if (tid == 0) { + uint acc = 0; + uint b = 0; + for (int bb = 255; bb >= 0; --bb) { + const uint c = atomic_load_explicit(&histo[bb], memory_order_relaxed); + if (acc + c >= desired) { + b = uint(bb); + break; + } + acc += c; + } + *sh_bucket = b; + *sh_above = acc; + } + threadgroup_barrier(mem_flags::mem_threadgroup); + + prefix |= *sh_bucket << uint(shift); + desired -= *sh_above; + + // ensure every thread has consumed sh_bucket/sh_above before the next pass + threadgroup_barrier(mem_flags::mem_threadgroup); + } + + if (tid == 0) { + atomic_store_explicit(out_count, 0u, memory_order_relaxed); + } + threadgroup_barrier(mem_flags::mem_threadgroup); + + // emit everything above the threshold, then fill the rest from ties + const uint threshold = prefix; + + for (uint i = tid; i < ncols; i += ntg_x) { + if (ggml_top_k_f2ui(src0_row[i]) > threshold) { + const uint pos = atomic_fetch_add_explicit(out_count, 1u, memory_order_relaxed); + dst_row[pos] = (int32_t) i; + } + } + threadgroup_barrier(mem_flags::mem_threadgroup); + + for (uint i = tid; i < ncols; i += ntg_x) { + if (ggml_top_k_f2ui(src0_row[i]) == threshold) { + const uint pos = atomic_fetch_add_explicit(out_count, 1u, memory_order_relaxed); + if (pos < top_k) { + dst_row[pos] = (int32_t) i; + } + } + } +} From 2f608ab4f8a1107f87aa372c7f3d7492e63438fe Mon Sep 17 00:00:00 2001 From: Bartowski <3266127+bartowski1182@users.noreply.github.com> Date: Mon, 31 Aug 2026 14:33:50 -0400 Subject: [PATCH 060/104] AVX2: Speed up large batch size prompt processing of IQ models (llama/27402) * Batched gemm for grid IQ quants Style updates and a bit more performance Clean up comments Move code around Vectorize IQ panel decode, lower threshold for speedup IQ panel: single-source gather layout, gate bias, vectorize interleave Add ggml_gemm_iqp_8x8_q8_K_p4 kernel, remove gather buffer Move IQ panel code out of repack into iqp.cpp, clean up comments Another comment sweep * Add myself as iqp.* codeownder * Remove ggml_cpu_iqp_scratch_offset and ggml_cpu_iqp_src1_conv_size * Renaming and moving * The other half of renaming and moving * Move macros and ggml_cpu_iqp_mul_mat_id_min_batch definition * Update ggml/src/ggml-cpu/iqp.h Co-authored-by: Georgi Gerganov * Add iqp_rows work buffer * Revert "Add iqp_rows work buffer" This reverts commit 425542991eee1b01fa3844bf87fc4f205ddbfccb. * Add NUMA fallback * Add 10 row batch tests for IQP coverage on all grid IQ types * Swap assert for return false in support check * Move IQP mul_mat_id test --------- Co-authored-by: Georgi Gerganov --- ggml/src/ggml-common.h | 2 +- ggml/src/ggml-cpu/CMakeLists.txt | 2 + ggml/src/ggml-cpu/ggml-cpu.c | 34 + ggml/src/ggml-cpu/iqp.cpp | 1253 ++++++++++++++++++++++++++++++ ggml/src/ggml-cpu/iqp.h | 39 + 5 files changed, 1329 insertions(+), 1 deletion(-) create mode 100644 ggml/src/ggml-cpu/iqp.cpp create mode 100644 ggml/src/ggml-cpu/iqp.h diff --git a/ggml/src/ggml-common.h b/ggml/src/ggml-common.h index 83f9118da84..1dbbe326d0f 100644 --- a/ggml/src/ggml-common.h +++ b/ggml/src/ggml-common.h @@ -1131,7 +1131,7 @@ GGML_TABLE_END() #define NGRID_IQ1S 2048 #define IQ1S_DELTA 0.125f #define IQ1M_DELTA 0.125f -#if defined(GGML_COMMON_IMPL_C) +#if defined(GGML_COMMON_IMPL_C) || defined(GGML_COMMON_IMPL_CPP) GGML_TABLE_BEGIN(uint64_t, iq1s_grid, NGRID_IQ1S) 0xffffffffffffffff, 0xffffffffffffff01, 0xffffffffffff0000, 0xffffffffffff01ff, 0xffffffffffff0101, 0xffffffffff00ff00, 0xffffffffff000000, 0xffffffffff01ffff, diff --git a/ggml/src/ggml-cpu/CMakeLists.txt b/ggml/src/ggml-cpu/CMakeLists.txt index 3c6343fb2a9..5442e12501f 100644 --- a/ggml/src/ggml-cpu/CMakeLists.txt +++ b/ggml/src/ggml-cpu/CMakeLists.txt @@ -31,6 +31,8 @@ function(ggml_add_cpu_backend_variant_impl tag_name) ggml-cpu/ggml-cpu.cpp ggml-cpu/repack.cpp ggml-cpu/repack.h + ggml-cpu/iqp.cpp + ggml-cpu/iqp.h ggml-cpu/hbm.cpp ggml-cpu/hbm.h ggml-cpu/quants.c diff --git a/ggml/src/ggml-cpu/ggml-cpu.c b/ggml/src/ggml-cpu/ggml-cpu.c index 6bc4467e378..87a329f2697 100644 --- a/ggml/src/ggml-cpu/ggml-cpu.c +++ b/ggml/src/ggml-cpu/ggml-cpu.c @@ -4,6 +4,7 @@ #include "ggml-backend-impl.h" #include "ggml-backend.h" #include "traits.h" +#include "iqp.h" #include "ggml-cpu-impl.h" #include "ggml-impl.h" #include "quants.h" @@ -1363,6 +1364,13 @@ UseGgmlGemm1:; ggml_barrier(params->threadpool); + // IQ panel gemm (see iqp.h) - must come after the barrier above, it consumes the q8_K rows + // of src1 from the work buffer + if (ggml_cpu_iqp_supports_mul_mat(dst) && !params->use_ref) { + ggml_compute_forward_mul_mat_iqp(params, dst); + return; + } + #if GGML_USE_LLAMAFILE if (src1->type != vec_dot_type) { const void* wdata = (src1->type == vec_dot_type) ? src1->data : params->wdata; @@ -1580,6 +1588,16 @@ static void ggml_compute_forward_mul_mat_id( char (*atomic_current_chunk)[CACHE_LINE_SIZE] = // [n_as] incr_ptr_aligned(&wdata_cur, CACHE_LINE_SIZE * n_as, CACHE_LINE_SIZE); + // IQ panel gemm (see iqp.h); per expert eligibility is decided below, but the work buffer is + // reserved for the whole node (ggml_graph_plan sizes it without params, use_ref only skips the dispatch) + const bool iqp = ggml_cpu_iqp_supports_mul_mat_id(dst) && !params->use_ref; + + char * iqp_panels = NULL; + + if (iqp) { + iqp_panels = incr_ptr_aligned(&wdata_cur, nth * ggml_cpu_iqp_scratch_size(dst), 64); + } + GGML_ASSERT(params->wsize >= (size_t)((char *) wdata_cur - (char *) params->wdata)); if (src1->type != vec_dot_type) { @@ -1651,6 +1669,13 @@ static void ggml_compute_forward_mul_mat_id( continue; } + if (iqp && ggml_cpu_iqp_mul_mat_id_min_batch(cne1)) { + ggml_compute_forward_mul_mat_id_iqp(params, dst, cur_a, cne1, (const int32_t *) &MMID_MATRIX_ROW(cur_a, 0), + iqp_panels); + + continue; + } + const char * src0_cur = (const char *) src0->data + cur_a * nb02; const void * wdata = (src1->type == vec_dot_type) ? src1->data : params->wdata; const size_t row_size = ggml_row_size(vec_dot_type, ne10); @@ -2858,6 +2883,11 @@ struct ggml_cplan ggml_graph_plan( if (node->src[1]->type != vec_dot_type) { cur = ggml_row_size(vec_dot_type, ggml_nelements(node->src[1])); } + + // the IQ panel path needs one scratch panel per thread past the q8_K rows + if (ggml_cpu_iqp_supports_mul_mat(node)) { + cur = GGML_PAD(cur, 64) + n_tasks * ggml_cpu_iqp_scratch_size(node); + } } break; case GGML_OP_MUL_MAT_ID: { @@ -2877,6 +2907,10 @@ struct ggml_cplan ggml_graph_plan( cur += n_as*ids->ne[0]*ids->ne[1]*sizeof(struct mmid_row_mapping) + sizeof(int64_t); // atomic_current_chunk cur += CACHE_LINE_SIZE*n_as + CACHE_LINE_SIZE; + // the IQ panel path needs one scratch panel per thread on top of that + if (ggml_cpu_iqp_supports_mul_mat_id(node)) { + cur += n_tasks * ggml_cpu_iqp_scratch_size(node) + 64; + } } break; case GGML_OP_OUT_PROD: { diff --git a/ggml/src/ggml-cpu/iqp.cpp b/ggml/src/ggml-cpu/iqp.cpp new file mode 100644 index 00000000000..b9201db3814 --- /dev/null +++ b/ggml/src/ggml-cpu/iqp.cpp @@ -0,0 +1,1253 @@ +#define GGML_COMMON_IMPL_CPP +#define GGML_COMMON_DECL_CPP +#include "ggml-common.h" + +#include "ggml-impl.h" +#include "ggml-cpu.h" +#include "ggml-cpu-impl.h" +#include "simd-mappings.h" +#include "traits.h" + +#include +#include +#include + +#include "iqp.h" + +#define UNUSED GGML_UNUSED + +// smallest src1 batch for which the decode pays for itself +#define GGML_IQP_MIN_BATCH 8 + +// same, per expert, for MUL_MAT_ID +#define GGML_IQP_MIN_BATCH_ID 8 + +bool ggml_cpu_iqp_mul_mat_id_min_batch(int64_t cne1) { + return cne1 >= GGML_IQP_MIN_BATCH_ID; +} + +// src0 rows interleaved per panel +#define IQP_NB_ROWS 8 + +#define IQP_SB_SIZE 16 // weights per sub-block +#define IQP_NSB (QK_K / IQP_SB_SIZE) // sub-blocks per super-block + +// one super-block of a grid based IQ type decoded to int8, 8 rows interleaved: +// dfac[row] * iscales[sb*8 + row] * qs is bit identical to dequantize_row_iq* +struct block_iqp_x8 { + float dfac[8]; // f32 super-block scale, d * 2^-k + int32_t bias[8]; // 128 * sum(qs * iscale), see GGML_IQP_USE_BIAS + int8_t iscales[IQP_NSB * 8]; // integer sub-block scales, in [-32, 31] + int8_t qs[QK_K * 8]; // qs[sb*128 + g*32 + row*4 + k] = column sb*16 + g*4 + k +}; + +static_assert(sizeof(block_iqp_x8) == 8 * sizeof(float) + 8 * sizeof(int32_t) + IQP_NSB * 8 + QK_K * 8, + "wrong iqp_x8 block size/padding"); + +// feed the activations to VNNI as unsigned bytes (y + 128) and correct with bias[]; without VNNI the kernels use the maddubs sign trick instead and bias[] is not filled +#if defined(__AVX2__) && ((defined(__AVX512VNNI__) && defined(__AVX512VL__)) || defined(__AVXVNNI__)) +# define GGML_IQP_USE_BIAS 1 +#else +# define GGML_IQP_USE_BIAS 0 +#endif + +static inline size_t ggml_cpu_iqp_row_size(const struct ggml_tensor * dst) { + return ggml_row_size(GGML_TYPE_Q8_K, dst->src[1]->ne[0]); +} + +// the low 7 bits of v are the first 7 signs and the 8th is their parity (cf. unpack_ksigns in the CUDA backend) +static inline uint8_t iqp_unpack_ksigns(uint32_t v) { + uint32_t p = v ^ (v >> 4); + + p ^= p >> 2; + p ^= p >> 1; + + return (uint8_t) (v ^ ((p & 1) << 7)); +} + +#if defined(__AVX2__) + +// 0xFF in every byte whose sign bit is set; sv holds each sign byte broadcast over the 8 bytes it governs +static inline __m256i iqp_sign_mask(__m256i sv) { + const __m256i sel = _mm256_set1_epi64x((int64_t) 0x8040201008040201ULL); + +# if defined(__GFNI__) + // computes the and + compare in one instruction + return _mm256_gf2p8affine_epi64_epi8(sel, sv, 0); +# else + return _mm256_cmpeq_epi8(_mm256_and_si256(sv, sel), sel); +# endif +} + +// signs holds four sign bytes, byte l governing values 8*l .. 8*l+7 - spread each over its 8 lanes +static inline __m256i iqp_sign_bytes(uint32_t signs) { + const __m256i bcast = _mm256_setr_epi8(0, 0, 0, 0, 0, 0, 0, 0, 1, 1, 1, 1, 1, 1, 1, 1, // + 2, 2, 2, 2, 2, 2, 2, 2, 3, 3, 3, 3, 3, 3, 3, 3); + + return _mm256_shuffle_epi8(_mm256_set1_epi32((int32_t) signs), bcast); +} + +// x ^ m - m negates the lanes where m is 0xFF +static inline __m256i iqp_apply_signs(__m256i x, __m256i m) { + return _mm256_sub_epi8(_mm256_xor_si256(x, m), m); +} + +#endif + +// 32 values from four 8 byte grid entries, sign byte l of signs applied to group l +static inline void iqp_store_signed_x8(int8_t * GGML_RESTRICT dst, + uint64_t g0, + uint64_t g1, + uint64_t g2, + uint64_t g3, + uint32_t signs) { +#if defined(__AVX2__) + const __m256i g = _mm256_set_epi64x((int64_t) g3, (int64_t) g2, (int64_t) g1, (int64_t) g0); + const __m256i m = iqp_sign_mask(iqp_sign_bytes(signs)); + + _mm256_storeu_si256((__m256i *) dst, iqp_apply_signs(g, m)); +#else + const uint64_t g[4] = { g0, g1, g2, g3 }; + + for (int l = 0; l < 4; ++l) { + const uint8_t * grid = (const uint8_t *) &g[l]; + const uint8_t s = (uint8_t) (signs >> 8 * l); + + for (int j = 0; j < 8; ++j) { + dst[8 * l + j] = s & kmask_iq2xs[j] ? -grid[j] : grid[j]; + } + } +#endif +} + +// same, but the eight values of group l come from two 4 byte grid entries +static inline void iqp_store_signed_x4(int8_t * GGML_RESTRICT dst, + uint32_t g0a, + uint32_t g0b, + uint32_t g1a, + uint32_t g1b, + uint32_t g2a, + uint32_t g2b, + uint32_t g3a, + uint32_t g3b, + uint32_t signs) { +#if defined(__AVX2__) + const __m256i g = _mm256_setr_epi32((int32_t) g0a, (int32_t) g0b, (int32_t) g1a, (int32_t) g1b, (int32_t) g2a, + (int32_t) g2b, (int32_t) g3a, (int32_t) g3b); + const __m256i m = iqp_sign_mask(iqp_sign_bytes(signs)); + + _mm256_storeu_si256((__m256i *) dst, iqp_apply_signs(g, m)); +#else + const uint32_t ga[4] = { g0a, g1a, g2a, g3a }; + const uint32_t gb[4] = { g0b, g1b, g2b, g3b }; + + for (int l = 0; l < 4; ++l) { + const uint8_t * grid1 = (const uint8_t *) &ga[l]; + const uint8_t * grid2 = (const uint8_t *) &gb[l]; + const uint8_t s = (uint8_t) (signs >> 8 * l); + + for (int j = 0; j < 4; ++j) { + dst[8 * l + j + 0] = s & kmask_iq2xs[j + 0] ? -grid1[j] : grid1[j]; + dst[8 * l + j + 4] = s & kmask_iq2xs[j + 4] ? -grid2[j] : grid2[j]; + } + } +#endif +} + +// 32 values of 8 * grid + delta from four 8 byte grid entries (grid bytes are in {-1, 0, 1}), byte l of deltas applying to group l +static inline void iqp_store_iq1_x8(int8_t * GGML_RESTRICT dst, + uint64_t g0, + uint64_t g1, + uint64_t g2, + uint64_t g3, + uint32_t deltas) { +#if defined(__AVX2__) + __m256i g = _mm256_set_epi64x((int64_t) g3, (int64_t) g2, (int64_t) g1, (int64_t) g0); + + // no byte shift in AVX2 + g = _mm256_add_epi8(g, g); + g = _mm256_add_epi8(g, g); + g = _mm256_add_epi8(g, g); + + _mm256_storeu_si256((__m256i *) dst, _mm256_add_epi8(g, iqp_sign_bytes(deltas))); +#else + const uint64_t g[4] = { g0, g1, g2, g3 }; + + for (int l = 0; l < 4; ++l) { + const int8_t * grid = (const int8_t *) &g[l]; + const int8_t delta = (int8_t) (deltas >> 8 * l); + + for (int j = 0; j < 8; ++j) { + dst[8 * l + j] = 8 * grid[j] + delta; + } + } +#endif +} + +// 32 values from 16 packed nibbles through the kvalues_iq4nl lookup: low nibbles first, then high +static inline void iqp_store_iq4_x32(int8_t * GGML_RESTRICT dst, const uint8_t * GGML_RESTRICT qs) { +#if defined(__AVX2__) + const __m128i q = _mm_loadu_si128((const __m128i *) qs); + const __m128i lut = _mm_loadu_si128((const __m128i *) kvalues_iq4nl); + const __m128i m4 = _mm_set1_epi8(0xf); + + _mm_storeu_si128((__m128i *) (dst + 0), _mm_shuffle_epi8(lut, _mm_and_si128(q, m4))); + _mm_storeu_si128((__m128i *) (dst + 16), _mm_shuffle_epi8(lut, _mm_and_si128(_mm_srli_epi16(q, 4), m4))); +#else + for (int j = 0; j < 16; ++j) { + dst[j + 0] = kvalues_iq4nl[qs[j] & 0xf]; + dst[j + 16] = kvalues_iq4nl[qs[j] >> 4]; + } +#endif +} + +#if GGML_IQP_USE_BIAS + +// sum of qs * iscale over one super-block, at most 256 * 127 * 32 = 1.04e6 +static inline int32_t iqp_weighted_sum(const int8_t * GGML_RESTRICT vals, const int8_t * GGML_RESTRICT iscales) { +#if defined(__AVX2__) + static_assert(IQP_SB_SIZE == 16, "the vector path folds two sub-blocks per 32 byte load"); + + const __m256i ones8 = _mm256_set1_epi8(1); + const __m256i ones16 = _mm256_set1_epi16(1); + + __m256i acc = _mm256_setzero_si256(); + + for (int i = 0; i < QK_K / 32; ++i) { + // sum groups of 4 bytes into int32, the low four lanes cover sub-block 2*i and the high four 2*i + 1 + const __m256i v = _mm256_loadu_si256((const __m256i *) (vals + 32 * i)); + const __m256i p = _mm256_madd_epi16(_mm256_maddubs_epi16(ones8, v), ones16); + + const __m256i s = _mm256_set_m128i(_mm_set1_epi32(iscales[2 * i + 1]), _mm_set1_epi32(iscales[2 * i + 0])); + + acc = _mm256_add_epi32(acc, _mm256_mullo_epi32(p, s)); + } + + __m128i sum = _mm_add_epi32(_mm256_castsi256_si128(acc), _mm256_extracti128_si256(acc, 1)); + + sum = _mm_add_epi32(sum, _mm_shuffle_epi32(sum, _MM_SHUFFLE(1, 0, 3, 2))); + sum = _mm_add_epi32(sum, _mm_shuffle_epi32(sum, _MM_SHUFFLE(2, 3, 0, 1))); + + return _mm_cvtsi128_si32(sum); +#else + int32_t wsum = 0; + + for (int sb = 0; sb < IQP_NSB; ++sb) { + int32_t vsum = 0; + + for (int k = 0; k < IQP_SB_SIZE; ++k) { + vsum += vals[sb * IQP_SB_SIZE + k]; + } + + wsum += iscales[sb] * vsum; + } + + return wsum; +#endif +} + +#endif // GGML_IQP_USE_BIAS + +static void iqp_decode_iq2_xxs(const void * GGML_RESTRICT vx, + int8_t * GGML_RESTRICT vals, + int8_t * GGML_RESTRICT iscales, + float * GGML_RESTRICT dfac) { + const block_iq2_xxs * x = (const block_iq2_xxs *) vx; + + // db = d * (0.5 + ls) * 0.25 = (d / 8) * (2 * ls + 1), ls 4 bit + *dfac = GGML_CPU_FP16_TO_FP32(x->d) * 0.125f; + + uint32_t aux32[2]; + const uint8_t * aux8 = (const uint8_t *) aux32; + + for (int ib32 = 0; ib32 < QK_K / 32; ++ib32) { + memcpy(aux32, x->qs + 4 * ib32, 2 * sizeof(uint32_t)); + const int8_t ls = (int8_t) (2 * (aux32[1] >> 28) + 1); + + iscales[2 * ib32 + 0] = ls; + iscales[2 * ib32 + 1] = ls; + + const uint32_t signs = (uint32_t) iqp_unpack_ksigns((aux32[1] >> 0) & 127) | + (uint32_t) iqp_unpack_ksigns((aux32[1] >> 7) & 127) << 8 | + (uint32_t) iqp_unpack_ksigns((aux32[1] >> 14) & 127) << 16 | + (uint32_t) iqp_unpack_ksigns((aux32[1] >> 21) & 127) << 24; + + iqp_store_signed_x8(vals + 32 * ib32, iq2xxs_grid[aux8[0]], iq2xxs_grid[aux8[1]], iq2xxs_grid[aux8[2]], + iq2xxs_grid[aux8[3]], signs); + } +} + +static void iqp_decode_iq2_xs(const void * GGML_RESTRICT vx, + int8_t * GGML_RESTRICT vals, + int8_t * GGML_RESTRICT iscales, + float * GGML_RESTRICT dfac) { + const block_iq2_xs * x = (const block_iq2_xs *) vx; + + *dfac = GGML_CPU_FP16_TO_FP32(x->d) * 0.125f; + + for (int ib32 = 0; ib32 < QK_K / 32; ++ib32) { + iscales[2 * ib32 + 0] = (int8_t) (2 * (x->scales[ib32] & 0xf) + 1); + iscales[2 * ib32 + 1] = (int8_t) (2 * (x->scales[ib32] >> 4) + 1); + + const uint16_t * q = x->qs + 4 * ib32; + + const uint32_t signs = (uint32_t) iqp_unpack_ksigns(q[0] >> 9) | (uint32_t) iqp_unpack_ksigns(q[1] >> 9) << 8 | + (uint32_t) iqp_unpack_ksigns(q[2] >> 9) << 16 | + (uint32_t) iqp_unpack_ksigns(q[3] >> 9) << 24; + + iqp_store_signed_x8(vals + 32 * ib32, iq2xs_grid[q[0] & 511], iq2xs_grid[q[1] & 511], iq2xs_grid[q[2] & 511], + iq2xs_grid[q[3] & 511], signs); + } +} + +static void iqp_decode_iq2_s(const void * GGML_RESTRICT vx, + int8_t * GGML_RESTRICT vals, + int8_t * GGML_RESTRICT iscales, + float * GGML_RESTRICT dfac) { + const block_iq2_s * x = (const block_iq2_s *) vx; + + const uint8_t * qs = x->qs; + const uint8_t * qh = x->qh; + const uint8_t * signs = qs + QK_K / 8; + + *dfac = GGML_CPU_FP16_TO_FP32(x->d) * 0.125f; + + for (int ib32 = 0; ib32 < QK_K / 32; ++ib32) { + iscales[2 * ib32 + 0] = (int8_t) (2 * (x->scales[ib32] & 0xf) + 1); + iscales[2 * ib32 + 1] = (int8_t) (2 * (x->scales[ib32] >> 4) + 1); + + const uint32_t sbits = + (uint32_t) signs[0] | (uint32_t) signs[1] << 8 | (uint32_t) signs[2] << 16 | (uint32_t) signs[3] << 24; + + iqp_store_signed_x8(vals + 32 * ib32, iq2s_grid[qs[0] | (qh[ib32] << 8 & 0x300)], + iq2s_grid[qs[1] | (qh[ib32] << 6 & 0x300)], iq2s_grid[qs[2] | (qh[ib32] << 4 & 0x300)], + iq2s_grid[qs[3] | (qh[ib32] << 2 & 0x300)], sbits); + qs += 4; + signs += 4; + } +} + +static void iqp_decode_iq3_xxs(const void * GGML_RESTRICT vx, + int8_t * GGML_RESTRICT vals, + int8_t * GGML_RESTRICT iscales, + float * GGML_RESTRICT dfac) { + const block_iq3_xxs * x = (const block_iq3_xxs *) vx; + + const uint8_t * qs = x->qs; + const uint8_t * scales_and_signs = qs + QK_K / 4; + + // db = d * (0.5 + ls) * 0.5 = (d / 4) * (2 * ls + 1), ls 4 bit + *dfac = GGML_CPU_FP16_TO_FP32(x->d) * 0.25f; + + uint32_t aux32; + + for (int ib32 = 0; ib32 < QK_K / 32; ++ib32) { + memcpy(&aux32, scales_and_signs + 4 * ib32, sizeof(uint32_t)); + const int8_t ls = (int8_t) (2 * (aux32 >> 28) + 1); + + iscales[2 * ib32 + 0] = ls; + iscales[2 * ib32 + 1] = ls; + + const uint32_t signs = (uint32_t) iqp_unpack_ksigns((aux32 >> 0) & 127) | + (uint32_t) iqp_unpack_ksigns((aux32 >> 7) & 127) << 8 | + (uint32_t) iqp_unpack_ksigns((aux32 >> 14) & 127) << 16 | + (uint32_t) iqp_unpack_ksigns((aux32 >> 21) & 127) << 24; + + iqp_store_signed_x4(vals + 32 * ib32, iq3xxs_grid[qs[0]], iq3xxs_grid[qs[1]], iq3xxs_grid[qs[2]], + iq3xxs_grid[qs[3]], iq3xxs_grid[qs[4]], iq3xxs_grid[qs[5]], iq3xxs_grid[qs[6]], + iq3xxs_grid[qs[7]], signs); + qs += 8; + } +} + +static void iqp_decode_iq3_s(const void * GGML_RESTRICT vx, + int8_t * GGML_RESTRICT vals, + int8_t * GGML_RESTRICT iscales, + float * GGML_RESTRICT dfac) { + const block_iq3_s * x = (const block_iq3_s *) vx; + + const uint8_t * qs = x->qs; + const uint8_t * qh = x->qh; + const uint8_t * signs = x->signs; + + // db = d * (1 + 2 * ls), ls 4 bit + *dfac = GGML_CPU_FP16_TO_FP32(x->d); + + int k = 0; + + for (int ib32 = 0; ib32 < QK_K / 32; ib32 += 2) { + const int8_t db1 = (int8_t) (1 + 2 * (x->scales[ib32 / 2] & 0xf)); + const int8_t db2 = (int8_t) (1 + 2 * (x->scales[ib32 / 2] >> 4)); + + iscales[2 * ib32 + 0] = db1; + iscales[2 * ib32 + 1] = db1; + iscales[2 * ib32 + 2] = db2; + iscales[2 * ib32 + 3] = db2; + + for (int h = 0; h < 2; ++h) { + const uint32_t sbits = + (uint32_t) signs[0] | (uint32_t) signs[1] << 8 | (uint32_t) signs[2] << 16 | (uint32_t) signs[3] << 24; + + iqp_store_signed_x4(vals + k, iq3s_grid[qs[0] | ((qh[h] << 8) & 256)], + iq3s_grid[qs[1] | ((qh[h] << 7) & 256)], iq3s_grid[qs[2] | ((qh[h] << 6) & 256)], + iq3s_grid[qs[3] | ((qh[h] << 5) & 256)], iq3s_grid[qs[4] | ((qh[h] << 4) & 256)], + iq3s_grid[qs[5] | ((qh[h] << 3) & 256)], iq3s_grid[qs[6] | ((qh[h] << 2) & 256)], + iq3s_grid[qs[7] | ((qh[h] << 1) & 256)], sbits); + + k += 32; + qs += 8; + signs += 4; + } + qh += 2; + } +} + +// dequantize_row_iq1_* computes y = dl * (grid[j] + delta) with delta = +-1/8, so the panel stores 8 * grid[j] +- 1 and folds the /8 into dfac +static void iqp_decode_iq1_s(const void * GGML_RESTRICT vx, + int8_t * GGML_RESTRICT vals, + int8_t * GGML_RESTRICT iscales, + float * GGML_RESTRICT dfac) { + const block_iq1_s * x = (const block_iq1_s *) vx; + + const uint8_t * qs = x->qs; + const uint16_t * qh = x->qh; + + // dl = d * (2 * ls + 1) * 0.125, ls 3 bit + *dfac = GGML_CPU_FP16_TO_FP32(x->d) * 0.125f; + + for (int ib = 0; ib < QK_K / 32; ++ib) { + const int8_t dl = (int8_t) (2 * ((qh[ib] >> 12) & 7) + 1); + const int8_t delta = qh[ib] & 0x8000 ? -1 : 1; + + iscales[2 * ib + 0] = dl; + iscales[2 * ib + 1] = dl; + + iqp_store_iq1_x8(vals + 32 * ib, iq1s_grid[qs[0] | (((qh[ib] >> 0) & 7) << 8)], + iq1s_grid[qs[1] | (((qh[ib] >> 3) & 7) << 8)], iq1s_grid[qs[2] | (((qh[ib] >> 6) & 7) << 8)], + iq1s_grid[qs[3] | (((qh[ib] >> 9) & 7) << 8)], ((uint8_t) delta) * 0x01010101u); + qs += 4; + } +} + +static void iqp_decode_iq1_m(const void * GGML_RESTRICT vx, + int8_t * GGML_RESTRICT vals, + int8_t * GGML_RESTRICT iscales, + float * GGML_RESTRICT dfac) { + const block_iq1_m * x = (const block_iq1_m *) vx; + + // block_iq1_m has no d field - the fp16 super-block scale is spread over the top nibbles of the four scale words + const uint16_t * sc = (const uint16_t *) x->scales; + + iq1m_scale_t scale; + scale.u16 = (sc[0] >> 12) | ((sc[1] >> 8) & 0x00f0) | ((sc[2] >> 4) & 0x0f00) | (sc[3] & 0xf000); + + *dfac = GGML_CPU_FP16_TO_FP32(scale.f16) * 0.125f; + + const uint8_t * qs = x->qs; + const uint8_t * qh = x->qh; + + for (int ib = 0; ib < QK_K / 32; ++ib) { + iscales[2 * ib + 0] = (int8_t) (2 * ((sc[ib / 2] >> (6 * (ib % 2) + 0)) & 0x7) + 1); + iscales[2 * ib + 1] = (int8_t) (2 * ((sc[ib / 2] >> (6 * (ib % 2) + 3)) & 0x7) + 1); + + const uint16_t idx[4] = { + (uint16_t) (qs[0] | ((qh[0] << 8) & 0x700)), + (uint16_t) (qs[1] | ((qh[0] << 4) & 0x700)), + (uint16_t) (qs[2] | ((qh[1] << 8) & 0x700)), + (uint16_t) (qs[3] | ((qh[1] << 4) & 0x700)), + }; + const uint32_t deltas = (uint32_t) (qh[0] & 0x08 ? 0xff : 0x01) | (uint32_t) (qh[0] & 0x80 ? 0xff : 0x01) << 8 | + (uint32_t) (qh[1] & 0x08 ? 0xff : 0x01) << 16 | + (uint32_t) (qh[1] & 0x80 ? 0xff : 0x01) << 24; + + iqp_store_iq1_x8(vals + 32 * ib, iq1s_grid[idx[0]], iq1s_grid[idx[1]], iq1s_grid[idx[2]], iq1s_grid[idx[3]], + deltas); + qs += 4; + qh += 2; + } +} + +static void iqp_decode_iq4_xs(const void * GGML_RESTRICT vx, + int8_t * GGML_RESTRICT vals, + int8_t * GGML_RESTRICT iscales, + float * GGML_RESTRICT dfac) { + const block_iq4_xs * x = (const block_iq4_xs *) vx; + + const uint8_t * qs = x->qs; + + // dl = d * (ls - 32), ls 6 bit, so the integer scale is in [-32, 31] + *dfac = GGML_CPU_FP16_TO_FP32(x->d); + + for (int ib = 0; ib < QK_K / 32; ++ib) { + const int ls = ((x->scales_l[ib / 2] >> 4 * (ib % 2)) & 0xf) | (((x->scales_h >> 2 * ib) & 3) << 4); + const int8_t dl = (int8_t) (ls - 32); + + iscales[2 * ib + 0] = dl; + iscales[2 * ib + 1] = dl; + + iqp_store_iq4_x32(vals + 32 * ib, qs); + qs += 16; + } +} + +// expanded by the eligibility test and the decode dispatch +#define IQP_TYPE_LIST(T) \ + T(IQ2_XXS, iq2_xxs) \ + T(IQ2_XS, iq2_xs) \ + T(IQ2_S, iq2_s) \ + T(IQ3_XXS, iq3_xxs) \ + T(IQ3_S, iq3_s) \ + T(IQ1_S, iq1_s) \ + T(IQ1_M, iq1_m) \ + T(IQ4_XS, iq4_xs) + +static bool iqp_decode_superblock(enum ggml_type type, + const void * GGML_RESTRICT vx, + int8_t * GGML_RESTRICT vals, + int8_t * GGML_RESTRICT iscales, + float * GGML_RESTRICT dfac) { + switch (type) { +#define IQP_CASE(E, name) \ + case GGML_TYPE_##E: \ + iqp_decode_##name(vx, vals, iscales, dfac); \ + return true; + IQP_TYPE_LIST(IQP_CASE) +#undef IQP_CASE + default: + return false; + } +} + +#if defined(__AVX2__) + +// 8x8 int32 transpose of the 32 column group starting at column off +static inline void iqp_interleave_x8(int8_t * GGML_RESTRICT dst, const int8_t (*vals)[QK_K], int off) { + static_assert(IQP_NB_ROWS == 8, "the transpose is 8x8"); + + __m256i v[IQP_NB_ROWS]; + + for (int r = 0; r < IQP_NB_ROWS; ++r) { + v[r] = _mm256_loadu_si256((const __m256i *) (vals[r] + off)); + } + + // pair rows into dword couples, then into qword quadruples, then swap the 128 bit lanes + const __m256i a0 = _mm256_unpacklo_epi32(v[0], v[1]); + const __m256i a1 = _mm256_unpackhi_epi32(v[0], v[1]); + const __m256i a2 = _mm256_unpacklo_epi32(v[2], v[3]); + const __m256i a3 = _mm256_unpackhi_epi32(v[2], v[3]); + const __m256i a4 = _mm256_unpacklo_epi32(v[4], v[5]); + const __m256i a5 = _mm256_unpackhi_epi32(v[4], v[5]); + const __m256i a6 = _mm256_unpacklo_epi32(v[6], v[7]); + const __m256i a7 = _mm256_unpackhi_epi32(v[6], v[7]); + + const __m256i b0 = _mm256_unpacklo_epi64(a0, a2); + const __m256i b1 = _mm256_unpackhi_epi64(a0, a2); + const __m256i b2 = _mm256_unpacklo_epi64(a1, a3); + const __m256i b3 = _mm256_unpackhi_epi64(a1, a3); + const __m256i b4 = _mm256_unpacklo_epi64(a4, a6); + const __m256i b5 = _mm256_unpackhi_epi64(a4, a6); + const __m256i b6 = _mm256_unpacklo_epi64(a5, a7); + const __m256i b7 = _mm256_unpackhi_epi64(a5, a7); + + _mm256_storeu_si256((__m256i *) (dst + 0 * 32), _mm256_permute2x128_si256(b0, b4, 0x20)); + _mm256_storeu_si256((__m256i *) (dst + 1 * 32), _mm256_permute2x128_si256(b1, b5, 0x20)); + _mm256_storeu_si256((__m256i *) (dst + 2 * 32), _mm256_permute2x128_si256(b2, b6, 0x20)); + _mm256_storeu_si256((__m256i *) (dst + 3 * 32), _mm256_permute2x128_si256(b3, b7, 0x20)); + _mm256_storeu_si256((__m256i *) (dst + 4 * 32), _mm256_permute2x128_si256(b0, b4, 0x31)); + _mm256_storeu_si256((__m256i *) (dst + 5 * 32), _mm256_permute2x128_si256(b1, b5, 0x31)); + _mm256_storeu_si256((__m256i *) (dst + 6 * 32), _mm256_permute2x128_si256(b2, b6, 0x31)); + _mm256_storeu_si256((__m256i *) (dst + 7 * 32), _mm256_permute2x128_si256(b3, b7, 0x31)); +} + +#endif + +// decode IQP_NB_ROWS consecutive source rows (starting at src, row stride nb01) into a panel of nblocks block_iqp_x8 +static void iqp_decode_panel_8(enum ggml_type type, + const char * GGML_RESTRICT src, + size_t nb01, + int64_t nblocks, + block_iqp_x8 * GGML_RESTRICT dst) { + const size_t bsize = ggml_type_size(type); + + int8_t vals[IQP_NB_ROWS][QK_K]; + int8_t iscales[IQP_NB_ROWS][IQP_NSB]; + float dfac[IQP_NB_ROWS]; + + for (int64_t x = 0; x < nblocks; x++) { + for (int r = 0; r < IQP_NB_ROWS; r++) { + const char * blk = src + r * nb01 + x * bsize; + + const bool ok = iqp_decode_superblock(type, blk, vals[r], iscales[r], &dfac[r]); + GGML_ASSERT(ok); + +#ifdef GGML_IQP_VERIFY + // check that the panel reproduces the reference dequantization bit exactly + float ref[QK_K]; + ggml_get_type_traits(type)->to_float(blk, ref, QK_K); + for (int j = 0; j < QK_K; j++) { + const float scale = dfac[r] * iscales[r][j / IQP_SB_SIZE]; + GGML_ASSERT(scale * vals[r][j] == ref[j]); + } +#endif + } + + for (int r = 0; r < IQP_NB_ROWS; r++) { + dst->dfac[r] = dfac[r]; + + for (int sb = 0; sb < IQP_NSB; sb++) { + dst->iscales[sb * IQP_NB_ROWS + r] = iscales[r][sb]; + } + +#if GGML_IQP_USE_BIAS + dst->bias[r] = 128 * iqp_weighted_sum(vals[r], iscales[r]); +#endif + } + +#if defined(__AVX2__) + for (int grp = 0; grp < QK_K / 32; grp++) { + iqp_interleave_x8(dst->qs + grp * 256, vals, grp * 32); + } +#else + for (int r = 0; r < IQP_NB_ROWS; r++) { + for (int sb = 0; sb < IQP_NSB; sb++) { + for (int g = 0; g < IQP_SB_SIZE / 4; g++) { + memcpy(dst->qs + sb * 128 + g * 32 + r * 4, vals[r] + sb * IQP_SB_SIZE + g * 4, 4); + } + } + } +#endif + + dst++; + } +} + +// gemm/gemv kernels: vx points at block_iqp_x8, vy at plain (non interleaved) block_q8_K rows + +static void iqp_gemv_8x8_q8_K_generic(int n, + float * GGML_RESTRICT s, + size_t bs, + const void * GGML_RESTRICT vx, + const void * GGML_RESTRICT vy, + int nr, + int nc) { + const int nb = n / QK_K; + const int ncols_interleaved = 8; + + assert(n % QK_K == 0); + assert(nc % ncols_interleaved == 0); + + UNUSED(bs); + UNUSED(nr); + + const block_iqp_x8 * b_ptr_start = (const block_iqp_x8 *) vx; + const block_q8_K * a_ptr = (const block_q8_K *) vy; + + for (int x = 0; x < nc / ncols_interleaved; x++) { + const block_iqp_x8 * b_ptr = b_ptr_start + x * nb; + + float sumf[8] = { 0 }; + + for (int l = 0; l < nb; l++) { + int32_t sumi[8] = { 0 }; + + for (int sb = 0; sb < IQP_NSB; sb++) { + int32_t isum[8] = { 0 }; + + for (int g = 0; g < 4; g++) { + for (int j = 0; j < ncols_interleaved; j++) { + for (int k = 0; k < 4; k++) { + isum[j] += b_ptr[l].qs[sb * 128 + g * 32 + j * 4 + k] * a_ptr[l].qs[sb * 16 + g * 4 + k]; + } + } + } + + for (int j = 0; j < ncols_interleaved; j++) { + sumi[j] += isum[j] * b_ptr[l].iscales[sb * 8 + j]; + } + } + + for (int j = 0; j < ncols_interleaved; j++) { + sumf[j] += (float) sumi[j] * (b_ptr[l].dfac[j] * a_ptr[l].d); + } + } + + for (int j = 0; j < ncols_interleaved; j++) { + s[x * ncols_interleaved + j] = sumf[j]; + } + } +} + +// one 4 row x nc column tile; s points at the first of the four output rows, bs floats apart +static void iqp_gemm_tile_4_generic(int nb, + float * GGML_RESTRICT s, + size_t bs, + const block_iqp_x8 * GGML_RESTRICT b_ptr_start, + const block_q8_K * const a_ptr[4], + int nc) { + const int ncols_interleaved = 8; + + for (int x = 0; x < nc / ncols_interleaved; x++) { + const block_iqp_x8 * b_ptr = b_ptr_start + x * nb; + + float sumf[4][8]; + for (int m = 0; m < 4; m++) { + for (int j = 0; j < ncols_interleaved; j++) { + sumf[m][j] = 0.0f; + } + } + + for (int l = 0; l < nb; l++) { + for (int m = 0; m < 4; m++) { + int32_t sumi[8] = { 0 }; + + for (int sb = 0; sb < IQP_NSB; sb++) { + int32_t isum[8] = { 0 }; + + for (int g = 0; g < 4; g++) { + for (int j = 0; j < ncols_interleaved; j++) { + for (int k = 0; k < 4; k++) { + isum[j] += + b_ptr[l].qs[sb * 128 + g * 32 + j * 4 + k] * a_ptr[m][l].qs[sb * 16 + g * 4 + k]; + } + } + } + + for (int j = 0; j < ncols_interleaved; j++) { + sumi[j] += isum[j] * b_ptr[l].iscales[sb * 8 + j]; + } + } + + for (int j = 0; j < ncols_interleaved; j++) { + sumf[m][j] += (float) sumi[j] * (b_ptr[l].dfac[j] * a_ptr[m][l].d); + } + } + } + + for (int m = 0; m < 4; m++) { + for (int j = 0; j < ncols_interleaved; j++) { + s[m * bs + x * ncols_interleaved + j] = sumf[m][j]; + } + } + } +} + +static void iqp_gemm_8x8_q8_K_generic(int n, + float * GGML_RESTRICT s, + size_t bs, + const void * GGML_RESTRICT vx, + const void * GGML_RESTRICT vy, + int nr, + int nc) { + const int nb = n / QK_K; + + assert(n % QK_K == 0); + assert(nr % 4 == 0); + assert(nc % 8 == 0); + + const block_iqp_x8 * b_ptr_start = (const block_iqp_x8 *) vx; + const block_q8_K * a_ptr_start = (const block_q8_K *) vy; + + for (int y = 0; y < nr / 4; y++) { + const block_q8_K * a_ptr[4]; + for (int m = 0; m < 4; m++) { + a_ptr[m] = a_ptr_start + (y * 4 + m) * nb; + } + + iqp_gemm_tile_4_generic(nb, s + y * 4 * bs, bs, b_ptr_start, a_ptr, nc); + } +} + +static void iqp_gemm_8x8_q8_K_p4_generic(int n, + float * GGML_RESTRICT s, + size_t bs, + const void * GGML_RESTRICT vx, + const void * const * GGML_RESTRICT vy, + int nc) { + const int nb = n / QK_K; + + assert(n % QK_K == 0); + assert(nc % 8 == 0); + + const block_q8_K * a_ptr[4]; + for (int m = 0; m < 4; m++) { + a_ptr[m] = (const block_q8_K *) vy[m]; + } + + iqp_gemm_tile_4_generic(nb, s, bs, (const block_iqp_x8 *) vx, a_ptr, nc); +} + +#if defined(__AVX2__) + +// add int16_t pairwise and return as 256 bit int vector, then add the accumulator +static inline __m256i sum_i16_pairs_acc_int32x8(const __m256i acc, const __m256i x) { + const __m256i ones = _mm256_set1_epi16(1); + return _mm256_add_epi32(acc, _mm256_madd_epi16(ones, x)); +} + +static inline __m256i mul_sum_us8_pairs_acc_int32x8(const __m256i acc, const __m256i ax, const __m256i sy) { +# if defined(__AVX512VNNI__) && defined(__AVX512VL__) + return _mm256_dpbusd_epi32(acc, ax, sy); +# elif defined(__AVXVNNI__) + return _mm256_dpbusd_avx_epi32(acc, ax, sy); +# else + // Perform multiplication and create 16-bit values + const __m256i dot = _mm256_maddubs_epi16(ax, sy); + return sum_i16_pairs_acc_int32x8(acc, dot); +# endif +} + +// Integer variant of the function defined in ggml-quants.c +// multiply int8_t, add results pairwise twice and return as 256 bit int vector, then add the accumulator +static inline __m256i mul_sum_i8_pairs_acc_int32x8(const __m256i acc, const __m256i x, const __m256i y) { +# if defined(__AVXVNNIINT8__) + return _mm256_dpbssd_epi32(acc, x, y); +# else + // Get absolute values of x vectors + const __m256i ax = _mm256_sign_epi8(x, x); + // Sign the values of the y vectors + const __m256i sy = _mm256_sign_epi8(y, x); + return mul_sum_us8_pairs_acc_int32x8(acc, ax, sy); +# endif +} + +// load the 16 activations of one sub-block, offset by 128 when they are fed to dpbusd as unsigned bytes +static inline __m256i iqp_load_y(const int8_t * GGML_RESTRICT qs) { + __m128i y = _mm_loadu_si128((const __m128i *) qs); +# if GGML_IQP_USE_BIAS + y = _mm_xor_si128(y, _mm_set1_epi8((char) 0x80)); +# endif + return _mm256_broadcastsi128_si256(y); +} + +// xv: 8 rows x 4 signed weights, yb: the matching 4 activation bytes broadcast to all 8 lanes +static inline __m256i iqp_dot4(const __m256i acc, const __m256i xv, const __m256i yb) { +# if GGML_IQP_USE_BIAS + return mul_sum_us8_pairs_acc_int32x8(acc, yb, xv); +# else + return mul_sum_i8_pairs_acc_int32x8(acc, xv, yb); +# endif +} + +static inline __m256i iqp_load_iscales(const int8_t * GGML_RESTRICT iscales) { + return _mm256_cvtepi8_epi32(_mm_loadl_epi64((const __m128i *) iscales)); +} + +// accumulate one super-block of 8 interleaved rows against one q8_K row in int32; worst case 16 * 32 * 16 * 255 * 127 = 2.65e8 plus a bias of at most 1.33e8 does not overflow +static inline __m256i iqp_acc_block(const block_iqp_x8 * GGML_RESTRICT b, const block_q8_K * GGML_RESTRICT a) { + __m256i sumi = _mm256_setzero_si256(); + + for (int sb = 0; sb < IQP_NSB; sb++) { + const int8_t * qs = b->qs + sb * 128; + + const __m256i yv = iqp_load_y(a->qs + sb * 16); + + __m256i isum = _mm256_setzero_si256(); + + isum = iqp_dot4(isum, _mm256_loadu_si256((const __m256i *) (qs + 0)), _mm256_shuffle_epi32(yv, 0x00)); + isum = iqp_dot4(isum, _mm256_loadu_si256((const __m256i *) (qs + 32)), _mm256_shuffle_epi32(yv, 0x55)); + isum = iqp_dot4(isum, _mm256_loadu_si256((const __m256i *) (qs + 64)), _mm256_shuffle_epi32(yv, 0xAA)); + isum = iqp_dot4(isum, _mm256_loadu_si256((const __m256i *) (qs + 96)), _mm256_shuffle_epi32(yv, 0xFF)); + + sumi = _mm256_add_epi32(sumi, _mm256_mullo_epi32(isum, iqp_load_iscales(b->iscales + sb * 8))); + } + +# if GGML_IQP_USE_BIAS + sumi = _mm256_sub_epi32(sumi, _mm256_loadu_si256((const __m256i *) b->bias)); +# endif + + return sumi; +} + +// one 4 row x nc column tile; s points at the first of the four output rows, bs floats apart +static inline void iqp_gemm_tile_4(int nb, + float * GGML_RESTRICT s, + size_t bs, + const block_iqp_x8 * GGML_RESTRICT b_ptr_start, + const block_q8_K * const a_ptr[4], + int nc) { + const int ncols_interleaved = 8; + + for (int x = 0; x < nc / ncols_interleaved; x++) { + const block_iqp_x8 * b_ptr = b_ptr_start + x * nb; + + __m256 sumf[4]; + for (int m = 0; m < 4; m++) { + sumf[m] = _mm256_setzero_ps(); + } + + for (int l = 0; l < nb; l++) { + __m256i sumi[4]; + for (int m = 0; m < 4; m++) { + sumi[m] = _mm256_setzero_si256(); + } + + for (int sb = 0; sb < IQP_NSB; sb++) { + const int8_t * qs = b_ptr[l].qs + sb * 128; + + __m256i yv[4]; + __m256i isum[4]; + for (int m = 0; m < 4; m++) { + yv[m] = iqp_load_y(a_ptr[m][l].qs + sb * 16); + isum[m] = _mm256_setzero_si256(); + } + + const __m256i xv0 = _mm256_loadu_si256((const __m256i *) (qs + 0)); + const __m256i xv1 = _mm256_loadu_si256((const __m256i *) (qs + 32)); + const __m256i xv2 = _mm256_loadu_si256((const __m256i *) (qs + 64)); + const __m256i xv3 = _mm256_loadu_si256((const __m256i *) (qs + 96)); + + for (int m = 0; m < 4; m++) { + isum[m] = iqp_dot4(isum[m], xv0, _mm256_shuffle_epi32(yv[m], 0x00)); + isum[m] = iqp_dot4(isum[m], xv1, _mm256_shuffle_epi32(yv[m], 0x55)); + isum[m] = iqp_dot4(isum[m], xv2, _mm256_shuffle_epi32(yv[m], 0xAA)); + isum[m] = iqp_dot4(isum[m], xv3, _mm256_shuffle_epi32(yv[m], 0xFF)); + } + + const __m256i isc = iqp_load_iscales(b_ptr[l].iscales + sb * 8); + for (int m = 0; m < 4; m++) { + sumi[m] = _mm256_add_epi32(sumi[m], _mm256_mullo_epi32(isum[m], isc)); + } + } + +# if GGML_IQP_USE_BIAS + const __m256i bias = _mm256_loadu_si256((const __m256i *) b_ptr[l].bias); + for (int m = 0; m < 4; m++) { + sumi[m] = _mm256_sub_epi32(sumi[m], bias); + } +# endif + + const __m256 dfac = _mm256_loadu_ps(b_ptr[l].dfac); + for (int m = 0; m < 4; m++) { + sumf[m] = _mm256_fmadd_ps(_mm256_cvtepi32_ps(sumi[m]), + _mm256_mul_ps(dfac, _mm256_set1_ps(a_ptr[m][l].d)), sumf[m]); + } + } + + for (int m = 0; m < 4; m++) { + _mm256_storeu_ps(s + m * bs + x * ncols_interleaved, sumf[m]); + } + } +} + +#endif // __AVX2__ + +static void iqp_gemv_8x8_q8_K(int n, + float * GGML_RESTRICT s, + size_t bs, + const void * GGML_RESTRICT vx, + const void * GGML_RESTRICT vy, + int nr, + int nc) { + const int nb = n / QK_K; + const int ncols_interleaved = 8; + + assert(n % QK_K == 0); + assert(nc % ncols_interleaved == 0); + + UNUSED(bs); + UNUSED(nr); + UNUSED(nb); + UNUSED(ncols_interleaved); + +#if defined(__AVX2__) + const block_iqp_x8 * b_ptr_start = (const block_iqp_x8 *) vx; + const block_q8_K * a_ptr = (const block_q8_K *) vy; + + for (int x = 0; x < nc / ncols_interleaved; x++) { + const block_iqp_x8 * b_ptr = b_ptr_start + x * nb; + + __m256 sumf = _mm256_setzero_ps(); + + for (int l = 0; l < nb; l++) { + const __m256 dv = _mm256_mul_ps(_mm256_loadu_ps(b_ptr[l].dfac), _mm256_set1_ps(a_ptr[l].d)); + + sumf = _mm256_fmadd_ps(_mm256_cvtepi32_ps(iqp_acc_block(b_ptr + l, a_ptr + l)), dv, sumf); + } + + _mm256_storeu_ps(s + x * ncols_interleaved, sumf); + } + + return; +#endif + + iqp_gemv_8x8_q8_K_generic(n, s, bs, vx, vy, nr, nc); +} + +static void iqp_gemm_8x8_q8_K(int n, + float * GGML_RESTRICT s, + size_t bs, + const void * GGML_RESTRICT vx, + const void * GGML_RESTRICT vy, + int nr, + int nc) { + const int nb = n / QK_K; + const int ncols_interleaved = 8; + + assert(n % QK_K == 0); + assert(nr % 4 == 0); + assert(nc % ncols_interleaved == 0); + + UNUSED(nb); + UNUSED(ncols_interleaved); + +#if defined(__AVX2__) + const block_iqp_x8 * b_ptr_start = (const block_iqp_x8 *) vx; + const block_q8_K * a_ptr_start = (const block_q8_K *) vy; + + for (int y = 0; y < nr / 4; y++) { + const block_q8_K * a_ptr[4]; + for (int m = 0; m < 4; m++) { + a_ptr[m] = a_ptr_start + (y * 4 + m) * nb; + } + + iqp_gemm_tile_4(nb, s + y * 4 * bs, bs, b_ptr_start, a_ptr, nc); + } + + return; +#endif + + iqp_gemm_8x8_q8_K_generic(n, s, bs, vx, vy, nr, nc); +} + +// same as iqp_gemm_8x8_q8_K with nr = 4, but the activation rows are passed as separate pointers (for the scattered rows of MUL_MAT_ID) +static void iqp_gemm_8x8_q8_K_p4(int n, + float * GGML_RESTRICT s, + size_t bs, + const void * GGML_RESTRICT vx, + const void * const * GGML_RESTRICT vy, + int nc) { + const int nb = n / QK_K; + const int ncols_interleaved = 8; + + assert(n % QK_K == 0); + assert(nc % ncols_interleaved == 0); + + UNUSED(nb); + UNUSED(ncols_interleaved); + +#if defined(__AVX2__) + const block_q8_K * a_ptr[4]; + for (int m = 0; m < 4; m++) { + a_ptr[m] = (const block_q8_K *) vy[m]; + } + + iqp_gemm_tile_4(nb, s, bs, (const block_iqp_x8 *) vx, a_ptr, nc); + + return; +#endif + + iqp_gemm_8x8_q8_K_p4_generic(n, s, bs, vx, vy, nc); +} + +static bool iqp_type_supported(enum ggml_type type) { + switch (type) { +#define IQP_CASE(E, name) case GGML_TYPE_##E: + IQP_TYPE_LIST(IQP_CASE) +#undef IQP_CASE + return true; + default: + return false; + } +} + +static bool iqp_supported_common(const struct ggml_tensor * dst) { + const struct ggml_tensor * src0 = dst->src[0]; + const struct ggml_tensor * src1 = dst->src[1]; + + if (!iqp_type_supported(src0->type)) { + return false; + } + + // the path assumes the src1 conversion type is q8_K + if (ggml_get_type_traits_cpu(src0->type)->vec_dot_type != GGML_TYPE_Q8_K) { + return false; + } + + // escape hatch to A/B the panel against the plain vec_dot path without rebuilding (--no-repack does not cover this path) + static const bool disabled = getenv("GGML_NO_IQ_PANEL") != nullptr; + if (disabled) { + return false; + } + + if (!ggml_cpu_has_avx2()) { + return false; + } + + if (src1->type != GGML_TYPE_F32) { + return false; + } + + if (src0->ne[0] % QK_K != 0 || src0->ne[1] % IQP_NB_ROWS != 0) { + return false; + } + + if (src0->ne[3] != 1 || src1->ne[3] != 1 || !ggml_is_contiguous(src0)) { + return false; + } + + if (dst->type != GGML_TYPE_F32 || dst->nb[0] != sizeof(float)) { + return false; + } + + return true; +} + +bool ggml_cpu_iqp_supports_mul_mat(const struct ggml_tensor * dst) { + const struct ggml_tensor * src0 = dst->src[0]; + const struct ggml_tensor * src1 = dst->src[1]; + + if (!iqp_supported_common(dst)) { + return false; + } + + if (src1->ne[1] < GGML_IQP_MIN_BATCH) { + return false; + } + + // plain 2D weight matmuls only (src1 may still be batched over ne12) + if (src0->ne[2] != 1) { + return false; + } + + return true; +} + +bool ggml_cpu_iqp_supports_mul_mat_id(const struct ggml_tensor * dst) { + const struct ggml_tensor * ids = dst->src[2]; + + if (!iqp_supported_common(dst)) { + return false; + } + + // skip the node entirely (work buffer included) if no expert can reach the per expert threshold + if (!ggml_cpu_iqp_mul_mat_id_min_batch(ids->ne[0] * ids->ne[1])) { + return false; + } + + return true; +} + +void ggml_compute_forward_mul_mat_id_iqp(const struct ggml_compute_params * params, + struct ggml_tensor * dst, + int64_t cur_a, + int64_t cne1, + const int32_t * expert_rows, + void * panels) { + const struct ggml_tensor * src0 = dst->src[0]; + const struct ggml_tensor * src1 = dst->src[1]; + + GGML_TENSOR_BINARY_OP_LOCALS + + const int ith = params->ith; + const int nth = params->nth; + + const int64_t nblocks = ne00 / QK_K; + + const size_t nbw1 = ggml_cpu_iqp_row_size(dst); + + block_iqp_x8 * panel = (block_iqp_x8 *) ((char *) panels + (size_t) ith * ggml_cpu_iqp_scratch_size(dst)); + + const char * src0_cur = (const char *) src0->data + cur_a * nb02; + + const int64_t ngroups = ne01 / IQP_NB_ROWS; + + const int64_t g0 = (ngroups * ith) / nth; + const int64_t g1 = (ngroups * (ith + 1)) / nth; + + for (int64_t g = g0; g < g1; g++) { + const int64_t r = g * IQP_NB_ROWS; + + iqp_decode_panel_8(src0->type, src0_cur + r * nb01, nb01, nblocks, panel); + + // the dst rows are scattered, so the gemm writes into tmp and it is copied out row by row + float tmp[4 * IQP_NB_ROWS]; + + for (int64_t k = 0; k < cne1; k += 4) { + const int64_t nrows = MIN(4, cne1 - k); + + // a short tail tile duplicates its last row into the unused slots; the padding is never copied out + const void * rows[4]; + + for (int64_t m = 0; m < 4; m++) { + const int64_t kk = k + MIN(m, nrows - 1); + + rows[m] = (const char *) params->wdata + + ((expert_rows[2 * kk + 0] % ne11) + expert_rows[2 * kk + 1] * ne11) * nbw1; + } + + iqp_gemm_8x8_q8_K_p4(ne00, tmp, IQP_NB_ROWS, panel, rows, IQP_NB_ROWS); + + for (int64_t m = 0; m < nrows; m++) { + float * dst_col = (float *) ((char *) dst->data + expert_rows[2 * (k + m) + 0] * nb1 + + expert_rows[2 * (k + m) + 1] * nb2); + memcpy(dst_col + r, tmp + m * IQP_NB_ROWS, IQP_NB_ROWS * sizeof(float)); + } + } + } +} + +size_t ggml_cpu_iqp_scratch_size(const struct ggml_tensor * dst) { + return GGML_PAD((dst->src[0]->ne[0] / QK_K) * sizeof(block_iqp_x8), 64); +} + +void ggml_compute_forward_mul_mat_iqp(const struct ggml_compute_params * params, struct ggml_tensor * dst) { + const struct ggml_tensor * src0 = dst->src[0]; + const struct ggml_tensor * src1 = dst->src[1]; + + GGML_TENSOR_BINARY_OP_LOCALS + + const int ith = params->ith; + const int nth = params->nth; + + const int64_t nblocks = ne00 / QK_K; + + const size_t nbw1 = ggml_row_size(GGML_TYPE_Q8_K, ne10); + const size_t nbw2 = nbw1 * ne11; + + const size_t scratch_size = ggml_cpu_iqp_scratch_size(dst); + + const size_t scratch_offset = GGML_PAD(nbw2 * ne12, 64); + + GGML_ASSERT(scratch_offset + (size_t) nth * scratch_size <= params->wsize); + + block_iqp_x8 * panel = (block_iqp_x8 *) ((char *) params->wdata + scratch_offset + (size_t) ith * scratch_size); + + const int64_t nrows = ne11; + + const int64_t ngroups = ne01 / IQP_NB_ROWS; + + // aim for 4 chunks per thread; the caller has already reset the chunk counter + // on NUMA systems fall back to one chunk per thread + const int64_t chunks_per_thread = ggml_is_numa() ? 1 : 4; + const int64_t groups_per_chunk = MAX(1, (ngroups + nth * chunks_per_thread - 1) / (nth * chunks_per_thread)); + const int64_t nchunk = (ngroups + groups_per_chunk - 1) / groups_per_chunk; + + int current_chunk = ith; + + while (current_chunk < nchunk) { + const int64_t g0 = current_chunk * groups_per_chunk; + const int64_t g1 = MIN(g0 + groups_per_chunk, ngroups); + + for (int64_t g = g0; g < g1; g++) { + const int64_t r = g * IQP_NB_ROWS; + + iqp_decode_panel_8(src0->type, (const char *) src0->data + r * nb01, nb01, nblocks, panel); + + for (int64_t i12 = 0; i12 < ne12; i12++) { + const char * src1_ptr = (const char *) params->wdata + i12 * nbw2; + char * dst_ptr = (char *) dst->data + i12 * nb2; + + if (nrows > 3) { + iqp_gemm_8x8_q8_K(ne00, (float *) dst_ptr + r, nb1 / nb0, panel, src1_ptr, nrows - (nrows % 4), + IQP_NB_ROWS); + } + for (int64_t iter = nrows - (nrows % 4); iter < nrows; iter++) { + iqp_gemv_8x8_q8_K(ne00, (float *) (dst_ptr + iter * nb1) + r, ne01, panel, src1_ptr + nbw1 * iter, + 1 /* nrows */, IQP_NB_ROWS); + } + } + } + + current_chunk = ggml_threadpool_chunk_add(params->threadpool, 1); + } +} diff --git a/ggml/src/ggml-cpu/iqp.h b/ggml/src/ggml-cpu/iqp.h new file mode 100644 index 00000000000..017b03fb43f --- /dev/null +++ b/ggml/src/ggml-cpu/iqp.h @@ -0,0 +1,39 @@ +#pragma once + +#include "ggml-cpu-impl.h" +#include "ggml.h" + +// GGML internal header + +// batched mul_mat path for the grid based IQ types: decode 8 src0 rows at a time into per thread scratch +// (block_iqp_x8, see iqp.cpp) and run an integer gemm over them against all src1 columns + +#ifdef __cplusplus +extern "C" { +#endif + +// whether cne1 rows of src1 are enough for the decode to pay for itself, per expert, for MUL_MAT_ID +bool ggml_cpu_iqp_mul_mat_id_min_batch(int64_t cne1); + +bool ggml_cpu_iqp_supports_mul_mat(const struct ggml_tensor * dst); + +// node level test only - per expert eligibility is decided with ggml_cpu_iqp_mul_mat_id_min_batch +bool ggml_cpu_iqp_supports_mul_mat_id(const struct ggml_tensor * dst); + +// per thread panel scratch bytes, padded +size_t ggml_cpu_iqp_scratch_size(const struct ggml_tensor * dst); + +// must be called after src1 has been converted to q8_K into params->wdata and the threads have synchronized on it +void ggml_compute_forward_mul_mat_iqp(const struct ggml_compute_params * params, struct ggml_tensor * dst); + +// one expert: expert_rows points at its row of the matrix_rows table of (i1, i2) int32 pairs, panels at the base of the per thread panel scratches +void ggml_compute_forward_mul_mat_id_iqp(const struct ggml_compute_params * params, + struct ggml_tensor * dst, + int64_t cur_a, + int64_t cne1, + const int32_t * expert_rows, + void * panels); + +#ifdef __cplusplus +} +#endif From dbc40efce86f37fad93ecf211d3064f8ac79cff6 Mon Sep 17 00:00:00 2001 From: Georgi Gerganov Date: Mon, 31 Aug 2026 23:16:04 +0300 Subject: [PATCH 061/104] metal : add concat support for quantized types (llama/28116) Assisted-by: pi:llama.cpp/DeepSeek-V4-Flash-0731 --- ggml/src/ggml-metal/ggml-metal-device.m | 6 +++ ggml/src/ggml-metal/ggml-metal-ops.cpp | 24 ++++++++++-- ggml/src/ggml-metal/kernels/quantize.metal | 45 ++++++++++++++++++++++ 3 files changed, 71 insertions(+), 4 deletions(-) diff --git a/ggml/src/ggml-metal/ggml-metal-device.m b/ggml/src/ggml-metal/ggml-metal-device.m index 83344539e54..0d8484d0021 100644 --- a/ggml/src/ggml-metal/ggml-metal-device.m +++ b/ggml/src/ggml-metal/ggml-metal-device.m @@ -1538,6 +1538,12 @@ bool ggml_metal_device_supports_op(ggml_metal_device_t dev, const struct ggml_te return true; case GGML_TYPE_BF16: return has_bfloat; + 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: + return true; default: return false; } diff --git a/ggml/src/ggml-metal/ggml-metal-ops.cpp b/ggml/src/ggml-metal/ggml-metal-ops.cpp index 30ea4ec2745..bc8b3c8d485 100644 --- a/ggml/src/ggml-metal/ggml-metal-ops.cpp +++ b/ggml/src/ggml-metal/ggml-metal-ops.cpp @@ -552,8 +552,24 @@ int ggml_metal_op_concat(ggml_metal_op_t ctx, int idx) { const int32_t dim = ((const int32_t *) op->op_params)[0]; + const bool is_q = ggml_is_quantized(op->type); + + // for quantized types, concat is done at the block level (nb0 == type_size == block size) + int32_t ne00_arg = ne00; + int32_t ne10_arg = ne10; + int32_t ne0_arg = ne0; + if (is_q) { + const int32_t blck = ggml_blck_size(op->type); + GGML_ASSERT(ne00 % blck == 0); + GGML_ASSERT(ne10 % blck == 0); + GGML_ASSERT(ne0 % blck == 0); + ne00_arg = ne00/blck; + ne10_arg = ne10/blck; + ne0_arg = ne0/blck; + } + ggml_metal_kargs_concat args = { - /*.ne00 =*/ ne00, + /*.ne00 =*/ ne00_arg, /*.ne01 =*/ ne01, /*.ne02 =*/ ne02, /*.ne03 =*/ ne03, @@ -561,7 +577,7 @@ int ggml_metal_op_concat(ggml_metal_op_t ctx, int idx) { /*.nb01 =*/ nb01, /*.nb02 =*/ nb02, /*.nb03 =*/ nb03, - /*.ne10 =*/ ne10, + /*.ne10 =*/ ne10_arg, /*.ne11 =*/ ne11, /*.ne12 =*/ ne12, /*.ne13 =*/ ne13, @@ -569,7 +585,7 @@ int ggml_metal_op_concat(ggml_metal_op_t ctx, int idx) { /*.nb11 =*/ nb11, /*.nb12 =*/ nb12, /*.nb13 =*/ nb13, - /*.ne0 =*/ ne0, + /*.ne0 =*/ ne0_arg, /*.ne1 =*/ ne1, /*.ne2 =*/ ne2, /*.ne3 =*/ ne3, @@ -588,7 +604,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); + int nth = std::min(256, ne0_arg); // when rows are small, we can batch them together in a single threadgroup int nrptg = 1; diff --git a/ggml/src/ggml-metal/kernels/quantize.metal b/ggml/src/ggml-metal/kernels/quantize.metal index 59d0afe9695..42ca6d74a0b 100644 --- a/ggml/src/ggml-metal/kernels/quantize.metal +++ b/ggml/src/ggml-metal/kernels/quantize.metal @@ -207,6 +207,51 @@ template [[host_name("kernel_concat_i16")]] kernel kernel_concat_t kernel_conca template [[host_name("kernel_concat_i32")]] kernel kernel_concat_t kernel_concat; template [[host_name("kernel_concat_i64")]] kernel kernel_concat_t kernel_concat; +template +kernel void kernel_concat_q( + constant ggml_metal_kargs_concat & args, + device const char * src0, + device const char * src1, + device char * dst, + uint3 tgpig[[threadgroup_position_in_grid]], + ushort3 tpitg[[thread_position_in_threadgroup]], + ushort3 ntg[[threads_per_threadgroup]]) { + + // 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; + + if (i1 >= args.ne1) { + return; + } + + 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) { + device const block_q * x; + + if (i0 < args.ne00 && i1 < args.ne01 && i2 < args.ne02 && i3 < args.ne03) { + x = (device const block_q *)(src0 + (i3 )*args.nb03 + (i2 )*args.nb02 + (i1 )*args.nb01 + (i0 )*args.nb00); + } else { + x = (device const block_q *)(src1 + (i3 - o[3])*args.nb13 + (i2 - o[2])*args.nb12 + (i1 - o[1])*args.nb11 + (i0 - o[0])*args.nb10); + } + + device block_q * y = (device block_q *)(dst + i3*args.nb3 + i2*args.nb2 + i1*args.nb1 + i0*args.nb0); + + *y = *x; + } +} + +typedef decltype(kernel_concat_q) kernel_concat_q_t; + +template [[host_name("kernel_concat_q4_0")]] kernel kernel_concat_q_t kernel_concat_q; +template [[host_name("kernel_concat_q4_1")]] kernel kernel_concat_q_t kernel_concat_q; +template [[host_name("kernel_concat_q5_0")]] kernel kernel_concat_q_t kernel_concat_q; +template [[host_name("kernel_concat_q5_1")]] kernel kernel_concat_q_t kernel_concat_q; +template [[host_name("kernel_concat_q8_0")]] kernel kernel_concat_q_t kernel_concat_q; + template kernel void kernel_get_rows_q( constant ggml_metal_kargs_get_rows & args, From f22bb2ea4d96da2bf8cce738805c96c0c805a203 Mon Sep 17 00:00:00 2001 From: ynankani Date: Mon, 31 Aug 2026 20:18:01 +0000 Subject: [PATCH 062/104] CUDA: XOR swizzle flash attn K,V smem fp16 tiles (llama/25635) * CUDA: XOR swizzle flash attn K,V smem fp16 tiles Signed-off-by: ynankani * Fix use 64bit generic pointer instead of 32bit shared pointer Signed-off-by: ynankani * fix shared memory race in FA on DGX Spark * Handle corener case Signed-off-by: ynankani * Add swizzle test cases and gate sync for swizzled path only Signed-off-by: ynankani * gate CUDA PTX Signed-off-by: ynankani * offset calculation specific for swizzle branch Signed-off-by: ynankani * Reafctor code Signed-off-by: ynankani * Refactor FA swizzle ldmatrix if/else into helpers (K row/col, V offset) Signed-off-by: ynankani * rebase and update test case args Signed-off-by: ynankani * Allow swizzle for non-pow2 shapes, for which nbatch_2%32==0 Signed-off-by: ynankani --------- Signed-off-by: ynankani --- ggml/src/ggml-cuda/fattn-mma-f16.cuh | 72 ++++++++++----- ggml/src/ggml-cuda/fattn-swizzle.cuh | 126 +++++++++++++++++++++++++++ 2 files changed, 176 insertions(+), 22 deletions(-) create mode 100644 ggml/src/ggml-cuda/fattn-swizzle.cuh diff --git a/ggml/src/ggml-cuda/fattn-mma-f16.cuh b/ggml/src/ggml-cuda/fattn-mma-f16.cuh index 7f4cfd5511f..387e70fa149 100644 --- a/ggml/src/ggml-cuda/fattn-mma-f16.cuh +++ b/ggml/src/ggml-cuda/fattn-mma-f16.cuh @@ -2,6 +2,7 @@ #include "cp-async.cuh" #include "mma.cuh" #include "fattn-common.cuh" +#include "fattn-swizzle.cuh" using namespace ggml_cuda_mma; @@ -66,7 +67,7 @@ static constexpr __host__ __device__ fattn_mma_config ggml_cuda_fattn_mma_get_co GGML_CUDA_FATTN_MMA_CONFIG_CASE(192, 128, 32, 128, 2, 32, 96, 64, 64, 2, true); GGML_CUDA_FATTN_MMA_CONFIG_CASE(192, 128, 64, 128, 2, 32, 96, 64, 64, 2, true); - GGML_CUDA_FATTN_MMA_CONFIG_CASE(256, 256, 8, 64, 4, 64, 128, 128, 128, 2, true); + GGML_CUDA_FATTN_MMA_CONFIG_CASE(256, 256, 8, 128, 2, 64, 128, 128, 128, 2, true); GGML_CUDA_FATTN_MMA_CONFIG_CASE(256, 256, 16, 64, 4, 32, 128, 128, 128, 2, true); GGML_CUDA_FATTN_MMA_CONFIG_CASE(256, 256, 32, 128, 2, 32, 128, 128, 128, 2, true); GGML_CUDA_FATTN_MMA_CONFIG_CASE(256, 256, 64, 128, 2, 32, 128, 128, 128, 2, true); @@ -360,7 +361,7 @@ static constexpr __device__ int ggml_cuda_fattn_mma_get_nstages(const int DKQ, c // ------------------------------------------------------------------------------------------------------------------ -template +template static __device__ __forceinline__ void flash_attn_ext_f16_load_tile( const half2 * const __restrict__ KV, half2 * const __restrict__ tile_KV, const int D2, const int stride_KV, const int i_sup) { constexpr int warp_size = ggml_cuda_get_physical_warp_size(); @@ -397,7 +398,12 @@ static __device__ __forceinline__ void flash_attn_ext_f16_load_tile( for (int k0 = k0_start; k0 < k0_stop; k0 += stride_k) { const int k = k0 + (stride_k == warp_size ? threadIdx.x : threadIdx.x % stride_k); - cp_async_cg_16(tile_KV_32 + i*(stride_tile*sizeof(half2)) + k*16, KV + i*stride_KV + k*h2_per_chunk); + if constexpr (swz) { + const int smem_offs_b = ggml_cuda_fattn_smem_swizzle::bytes_rc(i, k*h2_per_chunk); + cp_async_cg_16(tile_KV_32 + smem_offs_b, KV + i*stride_KV + k*h2_per_chunk); + } else { + cp_async_cg_16(tile_KV_32 + i*(stride_tile*sizeof(half2)) + k*16, KV + i*stride_KV + k*h2_per_chunk); + } } } }; @@ -432,8 +438,13 @@ static __device__ __forceinline__ void flash_attn_ext_f16_load_tile( for (int k0 = k0_start; k0 < k0_stop; k0 += stride_k) { const int k = k0 + (stride_k == warp_size ? threadIdx.x : threadIdx.x % stride_k); - ggml_cuda_memcpy_1<16>(tile_KV + i*stride_tile + k*4, - !oob_check || i < i_sup ? KV + i*stride_KV + k*h2_per_chunk : zero); + if constexpr (swz) { + ggml_cuda_memcpy_1<16>((char *) tile_KV + ggml_cuda_fattn_smem_swizzle::bytes_rc(i, k*h2_per_chunk), + !oob_check || i < i_sup ? KV + i*stride_KV + k*h2_per_chunk : zero); + } else { + ggml_cuda_memcpy_1<16>(tile_KV + i*stride_tile + k*4, + !oob_check || i < i_sup ? KV + i*stride_KV + k*h2_per_chunk : zero); + } } } }; @@ -568,9 +579,11 @@ static __device__ __forceinline__ void flash_attn_ext_f16_iter( constexpr bool Q_in_reg = ggml_cuda_fattn_mma_get_Q_in_reg (DKQ, DV, ncols); constexpr int nstages = ggml_cuda_fattn_mma_get_nstages (DKQ, DV, ncols1, ncols2); - constexpr int stride_tile_K = nbatch_K2 + 4; - - constexpr int stride_tile_V = V_is_K_view ? stride_tile_K : nbatch_V2 + 4; + // swizzle the tile stride for K and V based on the batch size. + constexpr int stride_tile_K = ggml_cuda_fattn_smem_swizzle::tile_stride(nbatch_K2); + constexpr int stride_tile_V = V_is_K_view ? stride_tile_K : ggml_cuda_fattn_smem_swizzle::tile_stride(nbatch_V2); + constexpr bool swz_K = ggml_cuda_fattn_smem_swizzle::enabled(nbatch_K2); + constexpr bool swz_V = V_is_K_view ? swz_K : ggml_cuda_fattn_smem_swizzle::enabled(nbatch_V2); const int k_VKQ_0 = kb0 * nbatch_fa; #if defined(TURING_MMA_AVAILABLE) @@ -588,7 +601,7 @@ static __device__ __forceinline__ void flash_attn_ext_f16_iter( constexpr bool use_cp_async = true; cp_async_wait_all(); __syncthreads(); - flash_attn_ext_f16_load_tile + flash_attn_ext_f16_load_tile (V_h2 + int64_t(k_VKQ_0)*stride_V, tile_V, nbatch_V2, stride_V, k_VKQ_sup); } else { constexpr bool use_cp_async = nstages == 1; @@ -607,7 +620,7 @@ static __device__ __forceinline__ void flash_attn_ext_f16_iter( if constexpr (nstages <= 1) { const int k0_diff = k0_stop - k0_start; constexpr bool use_cp_async = nstages == 1; - flash_attn_ext_f16_load_tile + flash_attn_ext_f16_load_tile (K_h2 + int64_t(k_VKQ_0)*stride_K + k0_start, tile_K, k0_diff, stride_K, k_VKQ_sup); if (use_cp_async) { cp_async_wait_all(); @@ -623,7 +636,7 @@ static __device__ __forceinline__ void flash_attn_ext_f16_iter( #pragma unroll for (int k_KQ_0 = k0_start; k_KQ_0 < k0_stop; k_KQ_0 += T_A_KQ::J) { T_A_KQ K_A; - load_ldmatrix(K_A, tile_K + i_KQ_0*stride_tile_K + (k_KQ_0 - k0_start), stride_tile_K); + ggml_cuda_fattn_smem_swizzle::load_ldmatrix(K_A, tile_K, i_KQ_0, k_KQ_0 - k0_start); if constexpr (cols_per_warp == 8) { mma(KQ_C[i_KQ_00/(np*T_A_KQ::I)], K_A, Q_B[k_KQ_0/T_A_KQ::J]); } else { @@ -649,7 +662,7 @@ static __device__ __forceinline__ void flash_attn_ext_f16_iter( const int i_KQ_0 = i_KQ_00 + (threadIdx.y % np)*T_A_KQ::I; T_A_KQ K_A; - load_ldmatrix(K_A, tile_K + i_KQ_0*stride_tile_K + (k_KQ_0 - k0_start), stride_tile_K); + ggml_cuda_fattn_smem_swizzle::load_ldmatrix(K_A, tile_K, i_KQ_0, k_KQ_0 - k0_start); if constexpr (cols_per_warp == 8) { mma(KQ_C[i_KQ_00/(np*T_A_KQ::I)], K_A, Q_B[0]); @@ -943,7 +956,7 @@ static __device__ __forceinline__ void flash_attn_ext_f16_iter( flash_attn_ext_f16_load_mask (mask_h + k_VKQ_0 + nbatch_fa, tile_mask, stride_mask, k_VKQ_sup, jt*ncols1, ne01); } - flash_attn_ext_f16_load_tile + flash_attn_ext_f16_load_tile (K_h2 + int64_t(k_VKQ_0 + nbatch_fa)*stride_K, tile_K, nbatch_K2, stride_K, k_VKQ_sup); } } @@ -959,7 +972,7 @@ static __device__ __forceinline__ void flash_attn_ext_f16_iter( const int i0_diff = i0_stop - i0_start; if (!V_is_K_view || i0_stop > 2*nbatch_K2) { constexpr bool use_cp_async = nstages == 1; - flash_attn_ext_f16_load_tile + flash_attn_ext_f16_load_tile (V_h2 + int64_t(k_VKQ_0)*stride_V + i0_start/2, tile_V, i0_diff/2, stride_V, k_VKQ_sup); if (use_cp_async) { cp_async_wait_all(); @@ -978,7 +991,7 @@ static __device__ __forceinline__ void flash_attn_ext_f16_iter( const int k0 = k00 + (threadIdx.y % np)*T_A_VKQ::J; T_A_VKQ A; // Transposed in SRAM but not in registers, gets transposed on load. - load_ldmatrix_trans(A, tile_V_i + 2*k0*stride_tile_V + (i_VKQ_0 - i0_start)/2, stride_tile_V); + ggml_cuda_fattn_smem_swizzle::load_ldmatrix_trans(A, tile_V, (int)(tile_V_i - tile_V) + 2*k0*stride_tile_V + (i_VKQ_0 - i0_start)/2); if constexpr (T_B_KQ::I == 8) { mma(VKQ_C[i_VKQ_0/T_A_VKQ::I], A, B[k00/(np*T_A_VKQ::J)]); } else { @@ -1004,7 +1017,7 @@ static __device__ __forceinline__ void flash_attn_ext_f16_iter( const int k0 = k00 + (threadIdx.y % np)*T_A_VKQ::I; T_A_VKQ A; // Transposed in both SRAM and registers, load normally. - load_ldmatrix(A, tile_V_i + k0*stride_tile_V + (i_VKQ_0 - i0_start)/2, stride_tile_V); + ggml_cuda_fattn_smem_swizzle::load_ldmatrix(A, tile_V, (int)(tile_V_i - tile_V) + k0*stride_tile_V + (i_VKQ_0 - i0_start)/2); mma(VKQ_C[i_VKQ_0/i0_stride], B[k00/(np*T_A_VKQ::I)], A); } } @@ -1168,10 +1181,12 @@ static __device__ __forceinline__ void flash_attn_ext_f16_process_tile( static_assert(nwarps * (cols_per_warp/ncols2) % ncols1 == 0, "bad nwarps"); constexpr int stride_tile_Q = DKQ/2 + 4; - constexpr int stride_tile_K = nbatch_K2 + 4; - - constexpr int stride_tile_V = V_is_K_view ? stride_tile_K : nbatch_V2 + 4; + // swizzle the tile stride for K and V based on the batch size. + constexpr int stride_tile_K = ggml_cuda_fattn_smem_swizzle::tile_stride(nbatch_K2); + constexpr int stride_tile_V = V_is_K_view ? stride_tile_K : ggml_cuda_fattn_smem_swizzle::tile_stride(nbatch_V2); constexpr int stride_tile_KV_max = stride_tile_K > stride_tile_V ? stride_tile_K : stride_tile_V; + constexpr bool swz_K = ggml_cuda_fattn_smem_swizzle::enabled(nbatch_K2); + constexpr bool swz_V = V_is_K_view ? swz_K : ggml_cuda_fattn_smem_swizzle::enabled(nbatch_V2); extern __shared__ half2 tile_Q[]; half2 * tile_K = Q_in_reg ? tile_Q : tile_Q + ncols * stride_tile_Q; @@ -1265,7 +1280,7 @@ static __device__ __forceinline__ void flash_attn_ext_f16_process_tile( flash_attn_ext_f16_load_mask (mask_h + kb0*nbatch_fa, tile_mask, stride_mask, k_VKQ_sup, jt*ncols1, ne01); } - flash_attn_ext_f16_load_tile + flash_attn_ext_f16_load_tile (K_h2 + int64_t(kb0)*nbatch_fa*stride_K, tile_K, nbatch_K2, stride_K, k_VKQ_sup); } @@ -1430,11 +1445,17 @@ static __device__ __forceinline__ void flash_attn_ext_f16_process_tile( constexpr int tile_stride = nbatch_combine + 4; static_assert((DV/2) % nbatch_combine == 0, "bad nbatch_combine"); + constexpr bool combine_needs_sync = swz_K || swz_V; + if constexpr (cols_per_warp == 8) { const int jc_cwmo = (threadIdx.x % (2*T_C_VKQ::J)) / T_C_VKQ::J; // jc combine write meta offset const int jc_cwm = threadIdx.y*(2*T_C_VKQ::J) + 2*T_C_VKQ::get_j(-1) + jc_cwmo; // jc combine write meta const float2 KQ_cmr = make_float2(KQ_max[jc_cwmo], KQ_rowsum[jc_cwmo]); // KQ combine max rowsum + if constexpr (combine_needs_sync) { + __syncthreads(); + } + if (((!needs_fixup && !is_fixup) || np > 1) && threadIdx.x < 2*T_C_VKQ::J) { // Use the 16 bytes of padding in each row to store the meta data: KQ max, KQ rowsum, KQ max scale. ((float2 *) tile_Q)[jc_cwm*(tile_stride/2) + nbatch_combine/2] = KQ_cmr; @@ -1471,6 +1492,10 @@ static __device__ __forceinline__ void flash_attn_ext_f16_process_tile( const bool thread_should_write = T_C_KQ::J == 8 || T_C_KQ::get_j(threadIdx.x & 2) < 8; #endif // defined(TURING_MMA_AVAILABLE) + if constexpr (combine_needs_sync) { + __syncthreads(); + } + if (((!needs_fixup && !is_fixup) || np > 1) && thread_should_write) { ((float2 *) tile_Q)[jc_cwm*(tile_stride/2) + nbatch_combine/2] = KQ_cmr; } @@ -1914,8 +1939,11 @@ void ggml_cuda_flash_attn_ext_mma_f16_case(ggml_backend_cuda_context & ctx, ggml constexpr bool V_is_K_view = DKQ == 576; // Guaranteed by the kernel selection logic in fattn.cu - const size_t nbytes_shared_KV_1stage = nbatch_fa * std::max(nbatch_K2 + 4, nbatch_V2 + 4) * sizeof(half2); - const size_t nbytes_shared_KV_2stage = nbatch_fa * (nbatch_K2 + 4 + nbatch_V2 + 4) * sizeof(half2); + // KV tile strides must match flash_attn_ext_f16_iter / _process_tile. + const int stride_tile_K = ggml_cuda_fattn_smem_swizzle::tile_stride(nbatch_K2, cc); + const int stride_tile_V = V_is_K_view ? stride_tile_K : ggml_cuda_fattn_smem_swizzle::tile_stride(nbatch_V2, cc); + const size_t nbytes_shared_KV_1stage = nbatch_fa * std::max(stride_tile_K, stride_tile_V) * sizeof(half2); + const size_t nbytes_shared_KV_2stage = nbatch_fa * (stride_tile_K + stride_tile_V) * sizeof(half2); const size_t nbytes_shared_Q = ncols * (DKQ/2 + 4) * sizeof(half2); const size_t nbytes_shared_mask = ncols1 * (nbatch_fa/2 + 4) * sizeof(half2); const size_t nbytes_shared_combine = nwarps*cols_per_warp * (nbatch_combine + 4) * sizeof(half2); diff --git a/ggml/src/ggml-cuda/fattn-swizzle.cuh b/ggml/src/ggml-cuda/fattn-swizzle.cuh new file mode 100644 index 00000000000..44338c8db08 --- /dev/null +++ b/ggml/src/ggml-cuda/fattn-swizzle.cuh @@ -0,0 +1,126 @@ +#pragma once + +#include "common.cuh" +#include "mma.cuh" + +// XOR swizzle for K/V SMEM tiles to avoid bank conflicts without row padding (Turing+ only). +// Stride must be a multiple of 32 half2 columns, otherwise we keep +4 row padding. + +namespace ggml_cuda_fattn_smem_swizzle { + +static __host__ __device__ constexpr bool bank_aligned(const int nbatch_2) { + return nbatch_2 >= 32 && nbatch_2 % 32 == 0; +} + +static __device__ constexpr bool enabled(const int nbatch_2) { +#if defined(TURING_MMA_AVAILABLE) + return bank_aligned(nbatch_2); +#else + GGML_UNUSED(nbatch_2); + return false; +#endif // defined(TURING_MMA_AVAILABLE) +} + +static __host__ bool enabled(const int nbatch_2, const int cc) { +#ifdef GGML_USE_HIP + GGML_UNUSED(nbatch_2); + GGML_UNUSED(cc); + return false; +#else + return turing_mma_available(cc) && bank_aligned(nbatch_2); +#endif // GGML_USE_HIP +} + +static __device__ constexpr int tile_stride(const int nbatch_2) { + return enabled(nbatch_2) ? nbatch_2 : nbatch_2 + 4; +} + +static __host__ int tile_stride(const int nbatch_2, const int cc) { + return enabled(nbatch_2, cc) ? nbatch_2 : nbatch_2 + 4; +} + +// Swizzled byte offset for tile element (row, col_h2), same map used for writes and reads. +template +static __device__ __forceinline__ int bytes_rc(const int row, const int col_h2) { + static_assert(bank_aligned(stride_h2), "swizzled tile needs a stride that is a multiple of 32"); + return ((row * stride_h2 + col_h2) * (int) sizeof(half2)) ^ ((row & 7) << 4); +} + +// ldmatrix.x4 via 64-bit generic pointer. +static __device__ __forceinline__ void ldmatrix_x4(int * xi, const half2 * addr) { +#if defined(TURING_MMA_AVAILABLE) + asm volatile("ldmatrix.sync.aligned.m8n8.x4.b16 {%0, %1, %2, %3}, [%4];" + : "=r"(xi[0]), "=r"(xi[1]), "=r"(xi[2]), "=r"(xi[3]) + : "l"(addr)); +#else + GGML_UNUSED_VARS(xi, addr); + NO_DEVICE_CODE; +#endif // defined(TURING_MMA_AVAILABLE) +} + +static __device__ __forceinline__ void ldmatrix_x4_trans(int * xi, const half2 * addr) { +#if defined(TURING_MMA_AVAILABLE) + asm volatile("ldmatrix.sync.aligned.m8n8.x4.trans.b16 {%0, %1, %2, %3}, [%4];" + : "=r"(xi[0]), "=r"(xi[2]), "=r"(xi[1]), "=r"(xi[3]) + : "l"(addr)); +#else + GGML_UNUSED_VARS(xi, addr); + NO_DEVICE_CODE; +#endif // defined(TURING_MMA_AVAILABLE) +} + +// Per-lane swizzled address for one tile<16, 8, half2> ldmatrix: 16 rows, 4 half2 columns per lane. +template +static __device__ __forceinline__ const half2 * lane_addr( + const half2 * tile_base, const int base_row, const int base_col_h2, const int I, const int J) { + static_assert(bank_aligned(stride_h2), "swizzled tile needs a stride that is a multiple of 32"); + const int lane_row = threadIdx.x % I; + const int lane_col = (threadIdx.x / I) * (J / 2); + uint32_t byte_off = (uint32_t) ((base_row + lane_row)*stride_h2 + base_col_h2 + lane_col) * (uint32_t) sizeof(half2); + byte_off ^= (uint32_t) (((base_row + lane_row) & 7) << 4); + return (const half2 *) ((const char *) tile_base + byte_off); +} + +template +static __device__ __forceinline__ void load_ldmatrix( + TileT & t, const half2 * tile_base, const int base_row, const int base_col_h2) { + if constexpr (swz) { + static_assert(std::is_same_v>, + "the swizzled layout is only supported for tile<16, 8, half2>"); + ldmatrix_x4((int *) t.x, lane_addr(tile_base, base_row, base_col_h2, TileT::I, TileT::J)); + } else { + ggml_cuda_mma::load_ldmatrix(t, tile_base + base_row*stride_h2 + base_col_h2, stride_h2); + } +} + +template +static __device__ __forceinline__ void load_ldmatrix(TileT & t, const half2 * tile_base, const int off_h2) { + if constexpr (swz) { + load_ldmatrix(t, tile_base, off_h2 / stride_h2, off_h2 % stride_h2); + } else { + ggml_cuda_mma::load_ldmatrix(t, tile_base + off_h2, stride_h2); + } +} + +template +static __device__ __forceinline__ void load_ldmatrix_trans( + TileT & t, const half2 * tile_base, const int base_row, const int base_col_h2) { + if constexpr (swz) { + static_assert(std::is_same_v>, + "the swizzled layout is only supported for tile<16, 8, half2>"); + ldmatrix_x4_trans((int *) t.x, lane_addr(tile_base, base_row, base_col_h2, TileT::I, TileT::J)); + } else { + ggml_cuda_mma::load_ldmatrix_trans(t, tile_base + base_row*stride_h2 + base_col_h2, stride_h2); + } +} + +template +static __device__ __forceinline__ void load_ldmatrix_trans(TileT & t, const half2 * tile_base, const int off_h2) { + if constexpr (swz) { + load_ldmatrix_trans(t, tile_base, off_h2 / stride_h2, off_h2 % stride_h2); + } else { + ggml_cuda_mma::load_ldmatrix_trans(t, tile_base + off_h2, stride_h2); + } +} + +} // namespace ggml_cuda_fattn_smem_swizzle From 8e54c659b583c1a159f7164f9de338844f30f779 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Bu=C4=9Fra=20=C3=96zg=C3=BCrsoy?= <13810383+ozgursoy@users.noreply.github.com> Date: Tue, 1 Sep 2026 00:47:27 +0300 Subject: [PATCH 063/104] metal : add fa-vec tunings for M1 Ultra (llama/28088) * metal : add fa-vec tunings for M1 Ultra * metal : move M1 Ultra tunings after M1 Max section * metal : remove duplicate blank line --- ggml/src/ggml-metal/ggml-metal-tuning.cpp | 177 ++++++++++++++++++++++ 1 file changed, 177 insertions(+) diff --git a/ggml/src/ggml-metal/ggml-metal-tuning.cpp b/ggml/src/ggml-metal/ggml-metal-tuning.cpp index 6a742bd0fd0..90fdd040d9f 100644 --- a/ggml/src/ggml-metal/ggml-metal-tuning.cpp +++ b/ggml/src/ggml-metal/ggml-metal-tuning.cpp @@ -657,6 +657,183 @@ constexpr fa_vec_entry_t fa_vec_tuned_table[] = { { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q8_0, 576, 512, -1, 0 }, { 1, 4 } }, { { GGML_METAL_DEVICE_M1_MAX, GGML_TYPE_Q8_0, 576, 512, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_F16, 64, 64, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_F16, 64, 64, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_F16, 128, 128, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_F16, 128, 128, 3, 3 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_F16, 128, 128, 3, 4 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_F16, 192, 128, 1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_F16, 192, 128, 1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_F16, 192, 128, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_F16, 192, 128, 1, 3 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_F16, 192, 128, 1, 4 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_F16, 320, 256, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_F16, 320, 256, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_F16, 320, 256, 1, 4 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q4_0, 32, 32, 1, 3 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q4_0, 32, 32, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q4_0, 32, 32, 2, 3 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q4_0, 32, 32, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q4_0, 32, 32, 3, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q4_0, 32, 32, 3, 4 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q4_0, 64, 64, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q4_0, 64, 64, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q4_0, 64, 64, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q4_0, 64, 64, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q4_0, 64, 64, 3, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q4_0, 96, 96, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q4_0, 96, 96, 2, 3 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q4_0, 96, 96, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q4_0, 96, 96, 3, 3 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q4_0, 96, 96, 3, 4 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q4_0, 128, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q4_0, 128, 128, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q4_0, 192, 192, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q4_0, 192, 192, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q4_0, 192, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q4_0, 192, 128, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q4_0, 256, 256, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q4_0, 256, 256, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q4_0, 320, 256, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q4_0, 320, 256, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q4_0, 512, 512, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q4_0, 512, 512, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q4_0, 576, 512, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q4_0, 576, 512, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q4_1, 32, 32, 1, 3 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q4_1, 32, 32, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q4_1, 32, 32, 2, 3 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q4_1, 32, 32, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q4_1, 32, 32, 3, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q4_1, 32, 32, 3, 4 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q4_1, 64, 64, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q4_1, 64, 64, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q4_1, 64, 64, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q4_1, 96, 96, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q4_1, 96, 96, 3, 3 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q4_1, 128, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q4_1, 128, 128, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q4_1, 128, 128, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q4_1, 128, 128, 3, 3 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q4_1, 128, 128, 3, 4 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q4_1, 192, 192, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q4_1, 192, 192, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q4_1, 192, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q4_1, 192, 128, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q4_1, 256, 256, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q4_1, 256, 256, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q4_1, 256, 256, 3, 2 }, { 1, 1 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q4_1, 320, 256, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q4_1, 320, 256, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q4_1, 512, 512, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q4_1, 512, 512, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q4_1, 576, 512, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q4_1, 576, 512, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q5_0, 32, 32, -1, 1 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q5_0, 32, 32, 1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q5_0, 32, 32, 1, 3 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q5_0, 32, 32, 1, 4 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q5_0, 32, 32, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q5_0, 32, 32, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q5_0, 64, 64, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q5_0, 64, 64, -1, 1 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q5_0, 64, 64, 1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q5_0, 64, 64, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q5_0, 64, 64, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q5_0, 96, 96, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q5_0, 128, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q5_0, 128, 128, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q5_0, 192, 192, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q5_0, 192, 192, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q5_0, 192, 192, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q5_0, 192, 192, 1, 4 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q5_0, 192, 192, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q5_0, 192, 192, 3, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q5_0, 192, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q5_0, 192, 128, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q5_0, 256, 256, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q5_0, 256, 256, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q5_0, 256, 256, 1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q5_0, 320, 256, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q5_0, 320, 256, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q5_0, 320, 256, 1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q5_0, 320, 256, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q5_0, 512, 512, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q5_0, 512, 512, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q5_0, 576, 512, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q5_0, 576, 512, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q5_1, 32, 32, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q5_1, 32, 32, 2, 2 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q5_1, 32, 32, 2, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q5_1, 32, 32, 2, 4 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q5_1, 32, 32, 3, 2 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q5_1, 32, 32, 3, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q5_1, 32, 32, 3, 4 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q5_1, 64, 64, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q5_1, 64, 64, -1, 1 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q5_1, 64, 64, 1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q5_1, 64, 64, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q5_1, 64, 64, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q5_1, 96, 96, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q5_1, 96, 96, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q5_1, 96, 96, 1, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q5_1, 96, 96, 1, 4 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q5_1, 128, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q5_1, 128, 128, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q5_1, 192, 192, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q5_1, 192, 192, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q5_1, 192, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q5_1, 192, 128, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q5_1, 256, 256, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q5_1, 256, 256, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q5_1, 256, 256, 1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q5_1, 320, 256, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q5_1, 320, 256, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q5_1, 320, 256, 1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q5_1, 320, 256, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q5_1, 320, 256, 2, 3 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q5_1, 320, 256, 2, 4 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q5_1, 320, 256, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q5_1, 320, 256, 3, 3 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q5_1, 512, 512, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q5_1, 512, 512, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q5_1, 576, 512, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q5_1, 576, 512, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q8_0, 32, 32, 1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q8_0, 32, 32, 1, 3 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q8_0, 32, 32, 2, 3 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q8_0, 32, 32, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q8_0, 32, 32, 3, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q8_0, 32, 32, 3, 4 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q8_0, 64, 64, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q8_0, 64, 64, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q8_0, 64, 64, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q8_0, 64, 64, 3, 3 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q8_0, 64, 64, 3, 4 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q8_0, 96, 96, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q8_0, 96, 96, 2, 3 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q8_0, 96, 96, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q8_0, 96, 96, 3, 3 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q8_0, 96, 96, 3, 4 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q8_0, 128, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q8_0, 128, 128, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q8_0, 192, 192, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q8_0, 192, 192, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q8_0, 192, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q8_0, 192, 128, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q8_0, 256, 256, -1, 0 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q8_0, 256, 256, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q8_0, 320, 256, 1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q8_0, 320, 256, 1, 3 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q8_0, 320, 256, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q8_0, 320, 256, 2, 3 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q8_0, 320, 256, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q8_0, 320, 256, 3, 3 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q8_0, 512, 512, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q8_0, 512, 512, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q8_0, 576, 512, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M1_ULTRA, GGML_TYPE_Q8_0, 576, 512, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2, GGML_TYPE_F16, 32, 32, 1, 3 }, { 2, 4 } }, { { GGML_METAL_DEVICE_M2, GGML_TYPE_F16, 64, 64, -1, 0 }, { 1, 4 } }, { { GGML_METAL_DEVICE_M2, GGML_TYPE_F16, 64, 64, -1, 1 }, { 1, 4 } }, From 5032008bc86f37d62f2fc3a4cbab6131e5dc10c9 Mon Sep 17 00:00:00 2001 From: James Francis <6763899+JamesFranc@users.noreply.github.com> Date: Tue, 1 Sep 2026 03:02:42 -0600 Subject: [PATCH 064/104] metal: enable Metal 4.0 tensor API on M5+/A19+ (llama/27461) * metal : request Metal 4.0 language version for the tensor API * metal : load the tensor API kernels from a separate metallib * tests : add external-metallib tensor API regression test * metal : fix metallib build order for the tensor API kernels --- ggml/src/ggml-metal/CMakeLists.txt | 21 +++- ggml/src/ggml-metal/ggml-metal-device.m | 129 ++++++++++++++++++------ 2 files changed, 116 insertions(+), 34 deletions(-) diff --git a/ggml/src/ggml-metal/CMakeLists.txt b/ggml/src/ggml-metal/CMakeLists.txt index 140c5d809e0..2094a409f96 100644 --- a/ggml/src/ggml-metal/CMakeLists.txt +++ b/ggml/src/ggml-metal/CMakeLists.txt @@ -163,19 +163,37 @@ else() ) endforeach() + # the tensor API kernels go in a separate metallib, loaded only where supported + set(AIR_MM_TENSOR "${CMAKE_RUNTIME_OUTPUT_DIRECTORY}/mul_mm_tensor.air") + add_custom_command( + OUTPUT ${AIR_MM_TENSOR} + COMMAND xcrun -sdk macosx metal ${XC_FLAGS} -DGGML_METAL_HAS_TENSOR -I ${CMAKE_RUNTIME_OUTPUT_DIRECTORY} -c ${CMAKE_RUNTIME_OUTPUT_DIRECTORY}/kernels/mul_mm.metal -o ${AIR_MM_TENSOR} + DEPENDS kernels/mul_mm.metal kernels/common.h kernels/dequantize.h ${METALLIB_COMMON} ggml-metal-impl.h + COMMENT "Compiling kernels/mul_mm.metal (tensor API)" + VERBATIM + ) + + add_custom_command( + OUTPUT ${CMAKE_RUNTIME_OUTPUT_DIRECTORY}/ggml-tensor.metallib + COMMAND xcrun -sdk macosx metallib ${AIR_MM_TENSOR} -o ${CMAKE_RUNTIME_OUTPUT_DIRECTORY}/ggml-tensor.metallib + DEPENDS ${AIR_MM_TENSOR} + COMMENT "Linking tensor API Metal kernels into ggml-tensor.metallib" + ) + add_custom_command( OUTPUT ${CMAKE_RUNTIME_OUTPUT_DIRECTORY}/default.metallib COMMAND xcrun -sdk macosx metallib ${AIR_FILES} -o ${CMAKE_RUNTIME_OUTPUT_DIRECTORY}/default.metallib COMMAND rm -f ${CMAKE_RUNTIME_OUTPUT_DIRECTORY}/ggml-common.h COMMAND rm -f ${CMAKE_RUNTIME_OUTPUT_DIRECTORY}/ggml-metal-impl.h COMMAND rm -rf ${CMAKE_RUNTIME_OUTPUT_DIRECTORY}/kernels - DEPENDS ${AIR_FILES} + DEPENDS ${AIR_FILES} ${AIR_MM_TENSOR} COMMENT "Linking Metal kernels into default.metallib" ) add_custom_target( ggml-metal-lib ALL DEPENDS ${CMAKE_RUNTIME_OUTPUT_DIRECTORY}/default.metallib + ${CMAKE_RUNTIME_OUTPUT_DIRECTORY}/ggml-tensor.metallib ) endif() # GGML_METAL_EMBED_LIBRARY @@ -188,6 +206,7 @@ if (NOT GGML_METAL_EMBED_LIBRARY) install( FILES ${CMAKE_RUNTIME_OUTPUT_DIRECTORY}/default.metallib + ${CMAKE_RUNTIME_OUTPUT_DIRECTORY}/ggml-tensor.metallib DESTINATION ${CMAKE_INSTALL_BINDIR} ) endif() diff --git a/ggml/src/ggml-metal/ggml-metal-device.m b/ggml/src/ggml-metal/ggml-metal-device.m index 0d8484d0021..9f2eb073138 100644 --- a/ggml/src/ggml-metal/ggml-metal-device.m +++ b/ggml/src/ggml-metal/ggml-metal-device.m @@ -27,6 +27,9 @@ static const NSInteger MTLGPUFamilyMetal3_GGML = 5001; static const NSInteger MTLGPUFamilyMetal4_GGML = 5002; +// MTLLanguageVersion4_0 is not present in older SDKs +static const NSUInteger MTLLanguageVersion4_0_GGML = 4 << 16; + #if !GGML_METAL_EMBED_LIBRARY // Here to assist with NSBundle Path Hack @interface GGMLMetalClass : NSObject @@ -154,6 +157,9 @@ int ggml_metal_pipeline_max_theads_per_threadgroup(struct ggml_metal_pipeline_wi // nil in single_library mode (everything resolves to objs[0]). NSMutableDictionary * fn_to_lib; + // kernels from a second metallib, resolved ahead of the combined library + NSSet * override_fns; + ggml_metal_device_t dev; ggml_metal_pipelines_t pipelines; // cache of compiled pipelines @@ -174,6 +180,18 @@ static void ggml_metal_library_build_index(ggml_metal_library_t lib) { } } +// note: defined below, after struct ggml_metal_device +static void ggml_metal_device_disable_tensor(ggml_metal_device_t dev); + +// the tensor API headers are exposed to the shader compiler only at Metal language version 4.0 +static void ggml_metal_compile_options_set_lang(MTLCompileOptions * options, bool has_tensor) { + if (!has_tensor) { + return; + } + + options.languageVersion = (MTLLanguageVersion) MTLLanguageVersion4_0_GGML; +} + // Parse a `#include "name"` line. Returns the quoted name in *include_name on // success. Whitespace-tolerant; ignores `#include <...>` (system headers). static bool ggml_metal_library_parse_quoted_include(NSString * line, NSString ** include_name) { @@ -313,6 +331,7 @@ static bool ggml_metal_library_compile_all( @autoreleasepool { MTLCompileOptions * options = [MTLCompileOptions new]; options.preprocessorMacros = prep; + ggml_metal_compile_options_set_lang(options, ggml_metal_device_get_props(res->dev)->has_tensor); lib = [device newLibraryWithSource:src options:options error:&error]; @@ -369,6 +388,46 @@ static bool ggml_metal_library_compile_all( return ok; } +// look for .metallib as a bundle resource, then next to the running binary +static NSString * ggml_metal_find_metallib(NSBundle * bundle, NSString * name) { + NSError * error = nil; + + NSString * path_lib = [bundle pathForResource:name ofType:@"metallib"]; + if (path_lib == nil) { + // Try to find the resource in the directory where the current binary located. + NSString * bin_cur = [[NSProcessInfo processInfo] arguments][0]; + NSString * bin_dir = [bin_cur stringByDeletingLastPathComponent]; + + NSString * path_lib_default = [NSString pathWithComponents:@[bin_dir, [name stringByAppendingPathExtension:@"metallib"]]]; + if ([[NSFileManager defaultManager] isReadableFileAtPath:path_lib_default]) { + GGML_LOG_INFO("%s: found '%s'\n", __func__, [path_lib_default UTF8String]); + + NSDictionary * atts = [[NSFileManager defaultManager] attributesOfItemAtPath:path_lib_default error:&error]; + if (atts && atts[NSFileType] == NSFileTypeSymbolicLink) { + // Optionally, if this is a symlink, try to resolve it. + path_lib_default = [[NSFileManager defaultManager] destinationOfSymbolicLinkAtPath:path_lib_default error:&error]; + if (path_lib_default && [path_lib_default length] > 0 && ![[path_lib_default substringToIndex:1] isEqualToString:@"/"]) { + // It is a relative path, adding the binary directory as directory prefix. + path_lib_default = [NSString pathWithComponents:@[bin_dir, path_lib_default]]; + } + if (!path_lib_default || ![[NSFileManager defaultManager] isReadableFileAtPath:path_lib_default]) { + // Link to the resource could not be resolved. + path_lib_default = nil; + } else { + GGML_LOG_INFO("%s: symlink resolved '%s'\n", __func__, [path_lib_default UTF8String]); + } + } + } else { + // The resource couldn't be found in the binary's directory. + path_lib_default = nil; + } + + path_lib = path_lib_default; + } + + return path_lib; +} + ggml_metal_library_t ggml_metal_library_init(ggml_metal_device_t dev) { id device = ggml_metal_device_get_obj(dev); @@ -432,38 +491,7 @@ ggml_metal_library_t ggml_metal_library_init(ggml_metal_device_t dev) { const int64_t t_start = ggml_time_us(); NSError * error = nil; - NSString * path_lib = [bundle pathForResource:@"default" ofType:@"metallib"]; - if (path_lib == nil) { - // Try to find the resource in the directory where the current binary located. - NSString * bin_cur = [[NSProcessInfo processInfo] arguments][0]; - NSString * bin_dir = [bin_cur stringByDeletingLastPathComponent]; - - NSString * path_lib_default = [NSString pathWithComponents:@[bin_dir, @"default.metallib"]]; - if ([[NSFileManager defaultManager] isReadableFileAtPath:path_lib_default]) { - GGML_LOG_INFO("%s: found '%s'\n", __func__, [path_lib_default UTF8String]); - - NSDictionary * atts = [[NSFileManager defaultManager] attributesOfItemAtPath:path_lib_default error:&error]; - if (atts && atts[NSFileType] == NSFileTypeSymbolicLink) { - // Optionally, if this is a symlink, try to resolve it. - path_lib_default = [[NSFileManager defaultManager] destinationOfSymbolicLinkAtPath:path_lib_default error:&error]; - if (path_lib_default && [path_lib_default length] > 0 && ![[path_lib_default substringToIndex:1] isEqualToString:@"/"]) { - // It is a relative path, adding the binary directory as directory prefix. - path_lib_default = [NSString pathWithComponents:@[bin_dir, path_lib_default]]; - } - if (!path_lib_default || ![[NSFileManager defaultManager] isReadableFileAtPath:path_lib_default]) { - // Link to the resource could not be resolved. - path_lib_default = nil; - } else { - GGML_LOG_INFO("%s: symlink resolved '%s'\n", __func__, [path_lib_default UTF8String]); - } - } - } else { - // The resource couldn't be found in the binary's directory. - path_lib_default = nil; - } - - path_lib = path_lib_default; - } + NSString * path_lib = ggml_metal_find_metallib(bundle, @"default"); if (path_lib != nil) { // pre-compiled library found: a single combined default.metallib @@ -478,6 +506,30 @@ ggml_metal_library_t ggml_metal_library_init(ggml_metal_device_t dev) { return NULL; } + // the tensor API kernels are built into a separate metallib + if (ggml_metal_device_get_props(dev)->has_tensor) { + NSString * path_mm = ggml_metal_find_metallib(bundle, @"ggml-tensor"); + + id lib_mm = nil; + if (path_mm != nil) { + lib_mm = [device newLibraryWithURL:[NSURL fileURLWithPath:path_mm] error:&error]; + if (!lib_mm && error) { + GGML_LOG_ERROR("%s: %s\n", __func__, [[error description] UTF8String]); + } + } + + if (lib_mm) { + GGML_LOG_INFO("%s: loaded '%s'\n", __func__, [path_mm UTF8String]); + + res->objs[GGML_METAL_LIB_MUL_MM] = [lib_mm retain]; + res->override_fns = [[NSSet setWithArray:[lib_mm functionNames]] retain]; + } else { + GGML_LOG_INFO("%s: ggml-tensor.metallib not found - disabling the tensor API\n", __func__); + + ggml_metal_device_disable_tensor(dev); + } + } + GGML_LOG_INFO("%s: loaded in %.3f sec\n", __func__, (ggml_time_us() - t_start) / 1e6); return res; } @@ -557,6 +609,7 @@ ggml_metal_library_t ggml_metal_library_init_from_source(ggml_metal_device_t dev MTLCompileOptions * options = [MTLCompileOptions new]; options.preprocessorMacros = prep; + ggml_metal_compile_options_set_lang(options, ggml_metal_device_get_props(dev)->has_tensor); library = [device newLibraryWithSource:src options:options error:&error]; if (error) { @@ -615,6 +668,10 @@ void ggml_metal_library_free(ggml_metal_library_t lib) { [lib->fn_to_lib release]; } + if (lib->override_fns) { + [lib->override_fns release]; + } + ggml_metal_pipelines_free(lib->pipelines); [lib->lock release]; @@ -676,7 +733,9 @@ struct ggml_metal_pipeline_with_params ggml_metal_library_compile_pipeline(ggml_ // route to the library that actually defines this kernel; fn_to_lib is // built from -[MTLLibrary functionNames] so it's always in sync int lib_idx = 0; - if (!lib->single_library) { + if (lib->override_fns && [lib->override_fns containsObject:base_func]) { + lib_idx = GGML_METAL_LIB_MUL_MM; + } else if (!lib->single_library) { NSNumber * idx = lib->fn_to_lib[base_func]; if (!idx) { [lib->lock unlock]; @@ -1862,6 +1921,10 @@ bool ggml_metal_device_supports_op(ggml_metal_device_t dev, const struct ggml_te return &dev->props; } +static void ggml_metal_device_disable_tensor(ggml_metal_device_t dev) { + dev->props.has_tensor = false; +} + // // device buffers // From 4f3a2a4b79e32efae8d291cae9322689454be368 Mon Sep 17 00:00:00 2001 From: Neo Zhang Date: Tue, 1 Sep 2026 18:35:47 +0800 Subject: [PATCH 065/104] sycl : support limit max alloc memory within 2GB for host-pinned memory (llama/27559) --- ggml/src/ggml-sycl/ggml-sycl.cpp | 19 ++++++++++++++++--- 1 file changed, 16 insertions(+), 3 deletions(-) diff --git a/ggml/src/ggml-sycl/ggml-sycl.cpp b/ggml/src/ggml-sycl/ggml-sycl.cpp index 290fb46767c..626f0f3bff1 100644 --- a/ggml/src/ggml-sycl/ggml-sycl.cpp +++ b/ggml/src/ggml-sycl/ggml-sycl.cpp @@ -107,6 +107,7 @@ int g_ggml_sycl_enable_flash_attention = 1; int g_ggml_sycl_dev2dev_memcpy = DEV2DEV_MEMCPY_SYCL; int g_ggml_sycl_usm_system = 0; int g_ggml_sycl_enable_host_pinned_mem = 1; +int g_ggml_sycl_host_pinned_mem_2g = 0; int g_ggml_sycl_get_mem_api = MEMORY_API_TYPE_LEVEL_ZERO; @@ -355,6 +356,8 @@ static void ggml_check_sycl() try { g_ggml_sycl_enable_host_pinned_mem = ggml_sycl_get_env("GGML_SYCL_ENABLE_HOST_PINNED_MEM", 1); + g_ggml_sycl_host_pinned_mem_2g = + ggml_sycl_get_env("GGML_SYCL_HOST_PINNED_MEM_2G", 0) & g_ggml_sycl_enable_host_pinned_mem; GGML_SYCL_DEBUG("[SYCL] call ggml_check_sycl\n"); @@ -457,6 +460,7 @@ static void ggml_check_sycl() try { GGML_LOG_INFO(" GGML_SYCL_USM_SYSTEM: %d\n", g_ggml_sycl_usm_system); GGML_LOG_INFO(" GGML_SYCL_ENABLE_HOST_PINNED_MEM: %d\n", g_ggml_sycl_enable_host_pinned_mem); + GGML_LOG_INFO(" GGML_SYCL_HOST_PINNED_MEM_2G: %d\n", g_ggml_sycl_host_pinned_mem_2g); /* NOT REMOVE, keep it for next optimize for XMX. #if defined(SYCL_USE_XMX) @@ -977,8 +981,12 @@ static size_t ggml_backend_sycl_buffer_type_get_alignment(ggml_backend_buffer_ty } static size_t ggml_backend_sycl_buffer_type_get_max_size(ggml_backend_buffer_type_t buft) { - return dpct::get_current_device().get_max_mem_alloc_size(); - + size_t max_alloc_size = dpct::get_current_device().get_max_mem_alloc_size(); + if (g_ggml_sycl_host_pinned_mem_2g) { + return std::min(max_alloc_size, (size_t) 2LL*1024*1024*1024); + } else { + return max_alloc_size; + } GGML_UNUSED(buft); } @@ -1551,7 +1559,12 @@ static size_t ggml_backend_sycl_host_buffer_type_get_max_size(ggml_backend_buffe if (g_ggml_sycl_enable_host_pinned_mem) { ggml_backend_sycl_device_context * dev_ctx = (ggml_backend_sycl_device_context *) buft->device->context; - return dpct::dev_mgr::instance().get_device(dev_ctx->device).get_max_mem_alloc_size(); + size_t max_alloc_size = dpct::dev_mgr::instance().get_device(dev_ctx->device).get_max_mem_alloc_size(); + if (g_ggml_sycl_host_pinned_mem_2g) { + return std::min(max_alloc_size, (size_t) 2LL*1024*1024*1024); + } else { + return max_alloc_size; + } } else { return SIZE_MAX; } From 870db2afa1a03001b941b920abb761729ba69f9a Mon Sep 17 00:00:00 2001 From: Georgi Gerganov Date: Tue, 1 Sep 2026 13:37:40 +0300 Subject: [PATCH 066/104] metal : add fa-vec tuning for M2 Max (llama/28015) Rows for M2 Max (30 GPU cores) collected with 'ggml-metal-tuning fa-vec --dtype f16,q8_0', pasted into fa_vec_tuned_table. ref: https://github.com/ggml-org/llama.cpp/discussions/27668#discussioncomment-18205786 Assisted-by: pi:llama.cpp/Qwen3.8-27B --- ggml/src/ggml-metal/ggml-metal-tuning.cpp | 77 +++++++++++++++++++++++ 1 file changed, 77 insertions(+) diff --git a/ggml/src/ggml-metal/ggml-metal-tuning.cpp b/ggml/src/ggml-metal/ggml-metal-tuning.cpp index 90fdd040d9f..bf44d5a7d6f 100644 --- a/ggml/src/ggml-metal/ggml-metal-tuning.cpp +++ b/ggml/src/ggml-metal/ggml-metal-tuning.cpp @@ -1029,6 +1029,83 @@ constexpr fa_vec_entry_t fa_vec_tuned_table[] = { { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q8_0, 576, 512, -1, 0 }, { 1, 4 } }, { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q8_0, 576, 512, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_MAX, GGML_TYPE_F16, 32, 32, 1, 3 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2_MAX, GGML_TYPE_F16, 32, 32, 2, 3 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2_MAX, GGML_TYPE_F16, 64, 64, 1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_MAX, GGML_TYPE_F16, 64, 64, 2, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_MAX, GGML_TYPE_F16, 64, 64, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_MAX, GGML_TYPE_F16, 64, 64, 1, 2 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M2_MAX, GGML_TYPE_F16, 64, 64, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2_MAX, GGML_TYPE_F16, 64, 64, 3, 2 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M2_MAX, GGML_TYPE_F16, 64, 64, 3, 3 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2_MAX, GGML_TYPE_F16, 64, 64, 3, 4 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2_MAX, GGML_TYPE_F16, 96, 96, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2_MAX, GGML_TYPE_F16, 128, 128, 3, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_MAX, GGML_TYPE_F16, 128, 128, 3, 3 }, { 2, 2 } }, + { { GGML_METAL_DEVICE_M2_MAX, GGML_TYPE_F16, 192, 192, 3, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_MAX, GGML_TYPE_F16, 192, 128, 1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_MAX, GGML_TYPE_F16, 192, 128, 2, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_MAX, GGML_TYPE_F16, 192, 128, 1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_MAX, GGML_TYPE_F16, 192, 128, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_MAX, GGML_TYPE_F16, 192, 128, 2, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_MAX, GGML_TYPE_F16, 192, 128, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_MAX, GGML_TYPE_F16, 192, 128, 3, 1 }, { 2, 2 } }, + { { GGML_METAL_DEVICE_M2_MAX, GGML_TYPE_F16, 320, 256, 2, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_MAX, GGML_TYPE_F16, 320, 256, 3, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_MAX, GGML_TYPE_F16, 320, 256, 2, 3 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_MAX, GGML_TYPE_F16, 320, 256, 2, 4 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_MAX, GGML_TYPE_F16, 512, 512, 2, 0 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M2_MAX, GGML_TYPE_F16, 512, 512, 2, 1 }, { 4, 1 } }, + { { GGML_METAL_DEVICE_M2_MAX, GGML_TYPE_F16, 512, 512, 2, 2 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M2_MAX, GGML_TYPE_Q8_0, 32, 32, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2_MAX, GGML_TYPE_Q8_0, 32, 32, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_MAX, GGML_TYPE_Q8_0, 32, 32, 2, 2 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M2_MAX, GGML_TYPE_Q8_0, 32, 32, 3, 2 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M2_MAX, GGML_TYPE_Q8_0, 32, 32, 3, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M2_MAX, GGML_TYPE_Q8_0, 64, 64, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_MAX, GGML_TYPE_Q8_0, 64, 64, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2_MAX, GGML_TYPE_Q8_0, 64, 64, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_MAX, GGML_TYPE_Q8_0, 64, 64, 1, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M2_MAX, GGML_TYPE_Q8_0, 64, 64, 2, 2 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M2_MAX, GGML_TYPE_Q8_0, 64, 64, 2, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M2_MAX, GGML_TYPE_Q8_0, 64, 64, 3, 2 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M2_MAX, GGML_TYPE_Q8_0, 64, 64, 3, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M2_MAX, GGML_TYPE_Q8_0, 96, 96, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2_MAX, GGML_TYPE_Q8_0, 96, 96, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_MAX, GGML_TYPE_Q8_0, 96, 96, 1, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M2_MAX, GGML_TYPE_Q8_0, 96, 96, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_MAX, GGML_TYPE_Q8_0, 96, 96, 2, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M2_MAX, GGML_TYPE_Q8_0, 96, 96, 2, 4 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_MAX, GGML_TYPE_Q8_0, 128, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_MAX, GGML_TYPE_Q8_0, 128, 128, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_MAX, GGML_TYPE_Q8_0, 128, 128, 2, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M2_MAX, GGML_TYPE_Q8_0, 128, 128, 2, 4 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2_MAX, GGML_TYPE_Q8_0, 128, 128, 3, 3 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2_MAX, GGML_TYPE_Q8_0, 192, 192, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_MAX, GGML_TYPE_Q8_0, 192, 192, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_MAX, GGML_TYPE_Q8_0, 192, 128, 1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_MAX, GGML_TYPE_Q8_0, 192, 128, 3, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_MAX, GGML_TYPE_Q8_0, 192, 128, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_MAX, GGML_TYPE_Q8_0, 192, 128, 3, 3 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M2_MAX, GGML_TYPE_Q8_0, 192, 128, 3, 4 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M2_MAX, GGML_TYPE_Q8_0, 256, 256, -1, 0 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M2_MAX, GGML_TYPE_Q8_0, 256, 256, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2_MAX, GGML_TYPE_Q8_0, 256, 256, 2, 2 }, { 4, 2 } }, + { { GGML_METAL_DEVICE_M2_MAX, GGML_TYPE_Q8_0, 256, 256, 2, 3 }, { 4, 2 } }, + { { GGML_METAL_DEVICE_M2_MAX, GGML_TYPE_Q8_0, 320, 256, 1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2_MAX, GGML_TYPE_Q8_0, 320, 256, 1, 3 }, { 4, 2 } }, + { { GGML_METAL_DEVICE_M2_MAX, GGML_TYPE_Q8_0, 320, 256, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2_MAX, GGML_TYPE_Q8_0, 320, 256, 2, 2 }, { 4, 2 } }, + { { GGML_METAL_DEVICE_M2_MAX, GGML_TYPE_Q8_0, 320, 256, 2, 3 }, { 4, 2 } }, + { { GGML_METAL_DEVICE_M2_MAX, GGML_TYPE_Q8_0, 320, 256, 2, 4 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2_MAX, GGML_TYPE_Q8_0, 320, 256, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2_MAX, GGML_TYPE_Q8_0, 320, 256, 3, 3 }, { 4, 2 } }, + { { GGML_METAL_DEVICE_M2_MAX, GGML_TYPE_Q8_0, 512, 512, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_MAX, GGML_TYPE_Q8_0, 512, 512, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_MAX, GGML_TYPE_Q8_0, 512, 512, 3, 1 }, { 1, 1 } }, + { { GGML_METAL_DEVICE_M2_MAX, GGML_TYPE_Q8_0, 576, 512, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_MAX, GGML_TYPE_Q8_0, 576, 512, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_ULTRA, GGML_TYPE_F16, 64, 64, -1, 0 }, { 1, 4 } }, { { GGML_METAL_DEVICE_M2_ULTRA, GGML_TYPE_F16, 64, 64, -1, 1 }, { 1, 4 } }, { { GGML_METAL_DEVICE_M2_ULTRA, GGML_TYPE_F16, 128, 128, 2, 1 }, { 1, 4 } }, From a245a8f48f4c59bb8da2eaa710d053c726f37c59 Mon Sep 17 00:00:00 2001 From: Niklas Wenzel Date: Tue, 1 Sep 2026 13:50:47 +0200 Subject: [PATCH 067/104] metal : fix more leaks due to missing autoreleasepools (llama/27883) * metal : fix more leaks due to missing autoreleasepools * metal : rename variable * metal : fix another missing pool warning Co-authored-by: YiChen Lv <63285796+forforever73@users.noreply.github.com> --------- Co-authored-by: YiChen Lv <63285796+forforever73@users.noreply.github.com> --- ggml/src/ggml-metal/ggml-metal-context.m | 18 ++++++++++ ggml/src/ggml-metal/ggml-metal-device.m | 44 ++++++++++++++---------- 2 files changed, 43 insertions(+), 19 deletions(-) diff --git a/ggml/src/ggml-metal/ggml-metal-context.m b/ggml/src/ggml-metal/ggml-metal-context.m index 1227ed39a09..e1129db3021 100644 --- a/ggml/src/ggml-metal/ggml-metal-context.m +++ b/ggml/src/ggml-metal/ggml-metal-context.m @@ -69,6 +69,10 @@ // extra command buffers for things like getting, setting and copying tensors NSMutableArray * cmd_bufs_ext; + // buffers to release after async Metal operations complete + // if Metal released them, it would do so on a Metal-internal thread without an autorelease pool, which could cause leaks + NSMutableArray * buf_refs; + // the last command buffer queued into the Metal queue with operations relevant to the current Metal backend id cmd_buf_last; @@ -179,6 +183,7 @@ ggml_metal_t ggml_metal_init(ggml_metal_device_t dev) { } res->cmd_bufs_ext = [[NSMutableArray alloc] init]; + res->buf_refs = [[NSMutableArray alloc] init]; res->cmd_buf_last = nil; @@ -206,6 +211,11 @@ void ggml_metal_free(ggml_metal_t ctx) { [ctx->cmd_bufs_ext removeAllObjects]; [ctx->cmd_bufs_ext release]; + @autoreleasepool { + [ctx->buf_refs removeAllObjects]; + [ctx->buf_refs release]; + } + if (ctx->pipelines_ext) { ggml_metal_pipelines_free(ctx->pipelines_ext); ctx->pipelines_ext = nil; @@ -294,6 +304,10 @@ void ggml_metal_synchronize(ggml_metal_t ctx) { [ctx->cmd_bufs_ext removeAllObjects]; } + + @autoreleasepool { + [ctx->buf_refs removeAllObjects]; + } } static struct ggml_metal_buffer_id ggml_metal_get_buffer_id(const struct ggml_tensor * t) { @@ -337,6 +351,8 @@ void ggml_metal_set_tensor_async(ggml_metal_t ctx, struct ggml_tensor * tensor, [encoder endEncoding]; [cmd_buf commit]; + + [ctx->buf_refs addObject:buf_src]; [buf_src release]; // do not wait here for completion @@ -381,6 +397,8 @@ void ggml_metal_get_tensor_async(ggml_metal_t ctx, const struct ggml_tensor * te [encoder endEncoding]; [cmd_buf commit]; + + [ctx->buf_refs addObject:buf_dst]; [buf_dst release]; // do not wait here for completion diff --git a/ggml/src/ggml-metal/ggml-metal-device.m b/ggml/src/ggml-metal/ggml-metal-device.m index 9f2eb073138..844de31f054 100644 --- a/ggml/src/ggml-metal/ggml-metal-device.m +++ b/ggml/src/ggml-metal/ggml-metal-device.m @@ -1346,19 +1346,21 @@ ggml_metal_device_t ggml_metal_device_init(int device, int n_devices) { void ggml_metal_device_free(ggml_metal_device_t dev) { assert(dev != NULL); - ggml_metal_rsets_free(dev->rsets); + @autoreleasepool { + ggml_metal_rsets_free(dev->rsets); - ggml_metal_library_free(dev->library); - dev->library = NULL; + ggml_metal_library_free(dev->library); + dev->library = NULL; - if (dev->mtl_queue) { - [dev->mtl_queue release]; - dev->mtl_queue = nil; - } + if (dev->mtl_queue) { + [dev->mtl_queue release]; + dev->mtl_queue = nil; + } - if (dev->mtl_device) { - [dev->mtl_device release]; - dev->mtl_device = nil; + if (dev->mtl_device) { + [dev->mtl_device release]; + dev->mtl_device = nil; + } } free(dev); @@ -1446,12 +1448,14 @@ ggml_metal_event_t ggml_metal_device_event_init(ggml_metal_device_t dev) { } void ggml_metal_device_event_free(ggml_metal_device_t dev, ggml_metal_event_t ev) { - id event = ev->obj; - [event release]; + @autoreleasepool { + id event = ev->obj; + [event release]; - free(ev); + free(ev); - GGML_UNUSED(dev); + GGML_UNUSED(dev); + } } void ggml_metal_device_event_synchronize(ggml_metal_device_t dev, ggml_metal_event_t ev) { @@ -2226,13 +2230,15 @@ ggml_metal_buffer_t ggml_metal_buffer_map(ggml_metal_device_t dev, void * ptr, s } void ggml_metal_buffer_free(ggml_metal_buffer_t buf) { - ggml_metal_device_rsets_rm(buf->dev, buf->rset); + @autoreleasepool { + ggml_metal_device_rsets_rm(buf->dev, buf->rset); - for (int i = 0; i < buf->n_buffers; i++) { - [buf->buffers[i].metal release]; - } + for (int i = 0; i < buf->n_buffers; i++) { + [buf->buffers[i].metal release]; + } - ggml_metal_buffer_rset_free(buf); + ggml_metal_buffer_rset_free(buf); + } if (buf->is_shared && buf->owned) { #if TARGET_OS_OSX From 8cca1a3616d854200a48922966552daf6cff2433 Mon Sep 17 00:00:00 2001 From: Jhen-Jie Hong Date: Tue, 1 Sep 2026 21:15:59 +0800 Subject: [PATCH 068/104] metal : add fa-vec tunings for A18 Pro (MacBook Neo) (llama/28152) --- ggml/src/ggml-metal/ggml-metal-device.h | 1 + ggml/src/ggml-metal/ggml-metal-device.m | 1 + ggml/src/ggml-metal/ggml-metal-tuning.cpp | 233 ++++++++++++++++++++++ 3 files changed, 235 insertions(+) diff --git a/ggml/src/ggml-metal/ggml-metal-device.h b/ggml/src/ggml-metal/ggml-metal-device.h index 7f6520103d1..ae4871d3586 100644 --- a/ggml/src/ggml-metal/ggml-metal-device.h +++ b/ggml/src/ggml-metal/ggml-metal-device.h @@ -259,6 +259,7 @@ enum ggml_metal_device_id { GGML_METAL_DEVICE_M5_PRO, GGML_METAL_DEVICE_M5_MAX, GGML_METAL_DEVICE_M5_ULTRA, + GGML_METAL_DEVICE_A18_PRO, }; const char * ggml_metal_device_id_token(enum ggml_metal_device_id id); diff --git a/ggml/src/ggml-metal/ggml-metal-device.m b/ggml/src/ggml-metal/ggml-metal-device.m index 844de31f054..ef81084241d 100644 --- a/ggml/src/ggml-metal/ggml-metal-device.m +++ b/ggml/src/ggml-metal/ggml-metal-device.m @@ -1056,6 +1056,7 @@ void ggml_metal_rsets_free(ggml_metal_rsets_t rsets) { DEV("M5 Pro", GGML_METAL_DEVICE_M5_PRO), DEV("M5 Max", GGML_METAL_DEVICE_M5_MAX), DEV("M5 Ultra", GGML_METAL_DEVICE_M5_ULTRA), + DEV("A18 Pro", GGML_METAL_DEVICE_A18_PRO), #undef DEV }; diff --git a/ggml/src/ggml-metal/ggml-metal-tuning.cpp b/ggml/src/ggml-metal/ggml-metal-tuning.cpp index bf44d5a7d6f..83082944421 100644 --- a/ggml/src/ggml-metal/ggml-metal-tuning.cpp +++ b/ggml/src/ggml-metal/ggml-metal-tuning.cpp @@ -2981,6 +2981,239 @@ constexpr fa_vec_entry_t fa_vec_tuned_table[] = { { { GGML_METAL_DEVICE_M5_MAX, GGML_TYPE_Q8_0, 576, 512, -1, 1 }, { 1, 4 } }, { { GGML_METAL_DEVICE_M5_MAX, GGML_TYPE_Q8_0, 576, 512, 1, 4 }, { 1, 2 } }, { { GGML_METAL_DEVICE_M5_MAX, GGML_TYPE_Q8_0, 576, 512, 2, 1 }, { 4, 4 } }, + + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_F16, 32, 32, 1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_F16, 32, 32, 1, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_F16, 32, 32, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_F16, 32, 32, 2, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_F16, 32, 32, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_F16, 32, 32, 3, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_F16, 64, 64, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_F16, 64, 64, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_F16, 64, 64, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_F16, 64, 64, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_F16, 64, 64, 3, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_F16, 64, 64, 3, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_F16, 96, 96, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_F16, 96, 96, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_F16, 96, 96, 1, 4 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_F16, 96, 96, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_F16, 96, 96, 2, 4 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_F16, 128, 128, -1, 1 }, { 2, 2 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_F16, 128, 128, 1, 2 }, { 1, 1 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_F16, 128, 128, 1, 4 }, { 1, 1 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_F16, 128, 128, 2, 2 }, { 1, 1 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_F16, 128, 128, 2, 4 }, { 1, 1 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_F16, 192, 192, -1, 1 }, { 2, 2 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_F16, 192, 192, 1, 2 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_F16, 192, 192, 1, 4 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_F16, 192, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_F16, 192, 128, -1, 1 }, { 2, 2 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_F16, 192, 128, 1, 2 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_F16, 192, 128, 1, 4 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_F16, 192, 128, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_F16, 192, 128, 2, 2 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_F16, 256, 256, 2, 0 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_F16, 256, 256, 3, 0 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_F16, 256, 256, -1, 1 }, { 2, 2 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_F16, 256, 256, 1, 2 }, { 1, 1 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_F16, 256, 256, 1, 4 }, { 1, 1 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_F16, 320, 256, -1, 1 }, { 2, 2 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_F16, 512, 512, 3, 0 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_F16, 512, 512, 3, 2 }, { 4, 1 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_F16, 512, 512, 3, 3 }, { 2, 2 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_F16, 576, 512, 2, 0 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_F16, 576, 512, 2, 2 }, { 4, 2 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_F16, 576, 512, 3, 2 }, { 4, 2 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q4_0, 32, 32, -1, 1 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q4_0, 32, 32, 1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q4_0, 32, 32, 1, 4 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q4_0, 32, 32, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q4_0, 32, 32, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q4_0, 64, 64, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q4_0, 64, 64, -1, 1 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q4_0, 64, 64, 1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q4_0, 64, 64, 1, 4 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q4_0, 64, 64, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q4_0, 64, 64, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q4_0, 96, 96, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q4_0, 96, 96, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q4_0, 96, 96, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q4_0, 96, 96, 2, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q4_0, 96, 96, 3, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q4_0, 128, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q4_0, 128, 128, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q4_0, 128, 128, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q4_0, 192, 192, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q4_0, 192, 192, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q4_0, 192, 192, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q4_0, 192, 192, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q4_0, 192, 192, 3, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q4_0, 192, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q4_0, 192, 128, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q4_0, 192, 128, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q4_0, 256, 256, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q4_0, 256, 256, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q4_0, 320, 256, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q4_0, 320, 256, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q4_0, 512, 512, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q4_0, 512, 512, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q4_0, 576, 512, 2, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q4_0, 576, 512, 3, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q4_0, 576, 512, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q4_0, 576, 512, 1, 2 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q4_0, 576, 512, 1, 3 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q4_1, 32, 32, -1, 1 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q4_1, 32, 32, 1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q4_1, 32, 32, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q4_1, 32, 32, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q4_1, 64, 64, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q4_1, 64, 64, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q4_1, 64, 64, 1, 2 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q4_1, 64, 64, 1, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q4_1, 64, 64, 2, 2 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q4_1, 64, 64, 2, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q4_1, 64, 64, 3, 2 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q4_1, 64, 64, 3, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q4_1, 96, 96, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q4_1, 96, 96, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q4_1, 96, 96, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q4_1, 128, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q4_1, 128, 128, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q4_1, 192, 192, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q4_1, 192, 192, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q4_1, 192, 192, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q4_1, 192, 192, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q4_1, 192, 192, 3, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q4_1, 192, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q4_1, 192, 128, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q4_1, 256, 256, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q4_1, 256, 256, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q4_1, 320, 256, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q4_1, 320, 256, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q4_1, 320, 256, 3, 1 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q4_1, 512, 512, 2, 0 }, { 4, 1 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q4_1, 512, 512, 3, 1 }, { 4, 2 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q4_1, 512, 512, 3, 2 }, { 4, 2 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q4_1, 576, 512, 2, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q4_1, 576, 512, 3, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q4_1, 576, 512, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q5_0, 32, 32, -1, 1 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q5_0, 32, 32, 1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q5_0, 32, 32, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q5_0, 32, 32, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q5_0, 64, 64, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q5_0, 64, 64, -1, 1 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q5_0, 64, 64, 1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q5_0, 64, 64, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q5_0, 64, 64, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q5_0, 96, 96, -1, 1 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q5_0, 96, 96, 1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q5_0, 96, 96, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q5_0, 96, 96, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q5_0, 128, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q5_0, 128, 128, -1, 1 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q5_0, 128, 128, 1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q5_0, 128, 128, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q5_0, 128, 128, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q5_0, 192, 192, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q5_0, 192, 192, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q5_0, 192, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q5_0, 192, 128, -1, 1 }, { 4, 2 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q5_0, 192, 128, 1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q5_0, 192, 128, 1, 4 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q5_0, 192, 128, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q5_0, 192, 128, 2, 4 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q5_0, 192, 128, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q5_0, 256, 256, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q5_0, 256, 256, -1, 1 }, { 2, 2 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q5_0, 256, 256, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q5_0, 320, 256, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q5_0, 320, 256, -1, 1 }, { 2, 2 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q5_0, 320, 256, 3, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q5_0, 512, 512, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q5_0, 512, 512, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q5_0, 512, 512, 2, 3 }, { 2, 1 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q5_0, 512, 512, 3, 1 }, { 2, 1 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q5_0, 512, 512, 3, 3 }, { 2, 1 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q5_0, 576, 512, 3, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q5_0, 576, 512, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q5_0, 576, 512, 1, 2 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q5_0, 576, 512, 1, 3 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q5_0, 576, 512, 1, 4 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q5_1, 32, 32, -1, 1 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q5_1, 32, 32, 1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q5_1, 32, 32, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q5_1, 32, 32, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q5_1, 64, 64, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q5_1, 64, 64, -1, 1 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q5_1, 64, 64, 1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q5_1, 64, 64, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q5_1, 64, 64, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q5_1, 96, 96, -1, 1 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q5_1, 96, 96, 1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q5_1, 96, 96, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q5_1, 96, 96, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q5_1, 128, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q5_1, 128, 128, -1, 1 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q5_1, 128, 128, 1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q5_1, 128, 128, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q5_1, 128, 128, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q5_1, 192, 192, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q5_1, 192, 192, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q5_1, 192, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q5_1, 192, 128, -1, 1 }, { 4, 2 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q5_1, 192, 128, 1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q5_1, 192, 128, 1, 4 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q5_1, 192, 128, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q5_1, 192, 128, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q5_1, 256, 256, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q5_1, 256, 256, -1, 1 }, { 2, 2 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q5_1, 320, 256, 1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q5_1, 320, 256, -1, 1 }, { 2, 2 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q5_1, 320, 256, 1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q5_1, 512, 512, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q5_1, 512, 512, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q5_1, 512, 512, 3, 3 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q5_1, 576, 512, 2, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q5_1, 576, 512, 3, 2 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q8_0, 32, 32, -1, 1 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q8_0, 32, 32, 1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q8_0, 32, 32, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q8_0, 32, 32, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q8_0, 64, 64, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q8_0, 64, 64, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q8_0, 64, 64, 1, 2 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q8_0, 64, 64, 1, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q8_0, 64, 64, 2, 2 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q8_0, 64, 64, 2, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q8_0, 64, 64, 3, 2 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q8_0, 64, 64, 3, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q8_0, 96, 96, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q8_0, 96, 96, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q8_0, 96, 96, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q8_0, 96, 96, 2, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q8_0, 128, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q8_0, 128, 128, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q8_0, 128, 128, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q8_0, 128, 128, 2, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q8_0, 192, 192, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q8_0, 192, 192, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q8_0, 192, 192, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q8_0, 192, 192, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q8_0, 192, 192, 3, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q8_0, 192, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q8_0, 192, 128, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q8_0, 256, 256, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q8_0, 256, 256, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q8_0, 320, 256, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q8_0, 320, 256, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q8_0, 512, 512, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q8_0, 512, 512, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q8_0, 576, 512, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q8_0, 576, 512, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q8_0, 576, 512, 1, 1 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_A18_PRO, GGML_TYPE_Q8_0, 576, 512, 1, 4 }, { 1, 2 } }, }; static enum ggml_metal_device_id fa_vec_family_representative(int gpu_family) { From 5f07f856c50c0480af1b0fa883f84068a1cd0e7f Mon Sep 17 00:00:00 2001 From: Lukasz Stolcman <4583553+lstolcman@users.noreply.github.com> Date: Tue, 1 Sep 2026 15:24:44 +0200 Subject: [PATCH 069/104] metal : add fa-vec tuning for M2 Pro (llama/28122) * metal: add fa-vec tuning for M2 Pro * metal : update fa-vec tuning for M2 Pro with new dtypes --- ggml/src/ggml-metal/ggml-metal-tuning.cpp | 191 ++++++++++++++++++++++ 1 file changed, 191 insertions(+) diff --git a/ggml/src/ggml-metal/ggml-metal-tuning.cpp b/ggml/src/ggml-metal/ggml-metal-tuning.cpp index 83082944421..c89a905dffa 100644 --- a/ggml/src/ggml-metal/ggml-metal-tuning.cpp +++ b/ggml/src/ggml-metal/ggml-metal-tuning.cpp @@ -1029,6 +1029,197 @@ constexpr fa_vec_entry_t fa_vec_tuned_table[] = { { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q8_0, 576, 512, -1, 0 }, { 1, 4 } }, { { GGML_METAL_DEVICE_M2, GGML_TYPE_Q8_0, 576, 512, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_F16, 64, 64, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_F16, 64, 64, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_F16, 128, 128, 2, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_F16, 128, 128, 2, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_F16, 128, 128, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_F16, 128, 128, 2, 3 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_F16, 128, 128, 2, 4 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_F16, 128, 128, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_F16, 128, 128, 3, 3 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_F16, 192, 128, 1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_F16, 192, 128, 1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_F16, 192, 128, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_F16, 192, 128, 1, 3 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_F16, 192, 128, 1, 4 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_F16, 320, 256, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_F16, 320, 256, 1, 1 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_F16, 320, 256, 1, 3 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_F16, 320, 256, 1, 4 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q4_0, 32, 32, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q4_0, 32, 32, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q4_0, 32, 32, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q4_0, 32, 32, 3, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q4_0, 32, 32, 3, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q4_0, 64, 64, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q4_0, 64, 64, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q4_0, 64, 64, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q4_0, 64, 64, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q4_0, 64, 64, 3, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q4_0, 96, 96, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q4_0, 96, 96, 2, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q4_0, 96, 96, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q4_0, 96, 96, 3, 3 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q4_0, 96, 96, 3, 4 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q4_0, 128, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q4_0, 128, 128, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q4_0, 192, 192, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q4_0, 192, 192, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q4_0, 192, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q4_0, 192, 128, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q4_0, 256, 256, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q4_0, 256, 256, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q4_0, 320, 256, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q4_0, 320, 256, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q4_0, 320, 256, 3, 3 }, { 4, 2 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q4_0, 320, 256, 3, 4 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q4_0, 512, 512, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q4_0, 512, 512, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q4_0, 576, 512, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q4_0, 576, 512, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q4_1, 32, 32, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q4_1, 32, 32, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q4_1, 32, 32, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q4_1, 32, 32, 3, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q4_1, 32, 32, 3, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q4_1, 64, 64, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q4_1, 64, 64, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q4_1, 64, 64, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q4_1, 96, 96, 2, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q4_1, 96, 96, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q4_1, 96, 96, 3, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q4_1, 128, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q4_1, 128, 128, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q4_1, 128, 128, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q4_1, 128, 128, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q4_1, 128, 128, 3, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q4_1, 192, 192, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q4_1, 192, 192, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q4_1, 192, 192, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q4_1, 192, 192, 1, 4 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q4_1, 192, 192, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q4_1, 192, 192, 3, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q4_1, 192, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q4_1, 192, 128, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q4_1, 256, 256, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q4_1, 256, 256, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q4_1, 320, 256, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q4_1, 320, 256, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q4_1, 512, 512, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q4_1, 512, 512, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q5_0, 32, 32, -1, 1 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q5_0, 32, 32, 1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q5_0, 32, 32, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q5_0, 32, 32, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q5_0, 64, 64, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q5_0, 64, 64, -1, 1 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q5_0, 64, 64, 1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q5_0, 64, 64, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q5_0, 64, 64, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q5_0, 96, 96, -1, 1 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q5_0, 96, 96, 1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q5_0, 96, 96, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q5_0, 96, 96, 2, 4 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q5_0, 96, 96, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q5_0, 96, 96, 3, 4 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q5_0, 128, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q5_0, 128, 128, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q5_0, 128, 128, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q5_0, 128, 128, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q5_0, 128, 128, 3, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q5_0, 128, 128, 3, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q5_0, 192, 192, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q5_0, 192, 192, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q5_0, 192, 192, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q5_0, 192, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q5_0, 192, 128, -1, 1 }, { 4, 2 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q5_0, 192, 128, 1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q5_0, 192, 128, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q5_0, 192, 128, 1, 4 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q5_0, 192, 128, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q5_0, 192, 128, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q5_0, 256, 256, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q5_0, 256, 256, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q5_0, 320, 256, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q5_0, 320, 256, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q5_0, 512, 512, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q5_0, 512, 512, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q5_0, 576, 512, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q5_0, 576, 512, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q5_1, 32, 32, -1, 1 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q5_1, 32, 32, 1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q5_1, 32, 32, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q5_1, 32, 32, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q5_1, 64, 64, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q5_1, 64, 64, -1, 1 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q5_1, 64, 64, 1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q5_1, 64, 64, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q5_1, 64, 64, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q5_1, 96, 96, -1, 1 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q5_1, 96, 96, 1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q5_1, 96, 96, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q5_1, 96, 96, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q5_1, 128, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q5_1, 128, 128, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q5_1, 128, 128, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q5_1, 128, 128, 2, 2 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q5_1, 128, 128, 2, 3 }, { 4, 2 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q5_1, 128, 128, 3, 2 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q5_1, 128, 128, 3, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q5_1, 192, 192, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q5_1, 192, 192, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q5_1, 192, 192, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q5_1, 192, 192, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q5_1, 192, 192, 3, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q5_1, 192, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q5_1, 192, 128, -1, 1 }, { 4, 2 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q5_1, 192, 128, 1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q5_1, 192, 128, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q5_1, 192, 128, 1, 4 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q5_1, 192, 128, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q5_1, 192, 128, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q5_1, 256, 256, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q5_1, 256, 256, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q5_1, 320, 256, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q5_1, 320, 256, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q5_1, 512, 512, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q5_1, 512, 512, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q5_1, 576, 512, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q5_1, 576, 512, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q8_0, 32, 32, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q8_0, 32, 32, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q8_0, 32, 32, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q8_0, 32, 32, 2, 4 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q8_0, 32, 32, 3, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q8_0, 32, 32, 3, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q8_0, 64, 64, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q8_0, 64, 64, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q8_0, 64, 64, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q8_0, 64, 64, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q8_0, 64, 64, 3, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q8_0, 96, 96, 2, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q8_0, 96, 96, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q8_0, 96, 96, 3, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q8_0, 128, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q8_0, 128, 128, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q8_0, 192, 192, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q8_0, 192, 192, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q8_0, 192, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q8_0, 192, 128, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q8_0, 256, 256, -1, 0 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q8_0, 256, 256, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q8_0, 320, 256, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q8_0, 320, 256, 1, 2 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q8_0, 320, 256, 1, 4 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q8_0, 320, 256, 2, 2 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q8_0, 320, 256, 2, 4 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q8_0, 320, 256, 3, 2 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q8_0, 512, 512, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q8_0, 512, 512, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q8_0, 576, 512, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_PRO, GGML_TYPE_Q8_0, 576, 512, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M2_MAX, GGML_TYPE_F16, 32, 32, 1, 3 }, { 2, 4 } }, { { GGML_METAL_DEVICE_M2_MAX, GGML_TYPE_F16, 32, 32, 2, 3 }, { 2, 4 } }, { { GGML_METAL_DEVICE_M2_MAX, GGML_TYPE_F16, 64, 64, 1, 0 }, { 1, 4 } }, From 408faaafcedc643e8feb77a3b2536a855fcd6549 Mon Sep 17 00:00:00 2001 From: "Jingxin (Philip) Li" Date: Tue, 1 Sep 2026 23:47:08 +0800 Subject: [PATCH 070/104] sycl : add Kronecker product FWHT support for sizes 384, 640, 768, 1280 (llama/28016) --- ggml/src/ggml-sycl/fwht.cpp | 172 ++++++++++++++++++++++++++++++++++++ 1 file changed, 172 insertions(+) diff --git a/ggml/src/ggml-sycl/fwht.cpp b/ggml/src/ggml-sycl/fwht.cpp index 2312b3d131b..39f273beaa9 100644 --- a/ggml/src/ggml-sycl/fwht.cpp +++ b/ggml/src/ggml-sycl/fwht.cpp @@ -1,6 +1,50 @@ #include "fwht.hpp" #include +#define P 1.0f +#define N -1.0f + +// constant Hadamard matrix via Paley I construction +static constexpr float H12[12][12] = { + { P, P, P, P, P, P, P, P, P, P, P, P }, + { P, N, P, N, P, P, P, N, N, N, P, N }, + { P, N, N, P, N, P, P, P, N, N, N, P }, + { P, P, N, N, P, N, P, P, P, N, N, N }, + { P, N, P, N, N, P, N, P, P, P, N, N }, + { P, N, N, P, N, N, P, N, P, P, P, N }, + { P, N, N, N, P, N, N, P, N, P, P, P }, + { P, P, N, N, N, P, N, N, P, N, P, P }, + { P, P, P, N, N, N, P, N, N, P, N, P }, + { P, P, P, P, N, N, N, P, N, N, P, N }, + { P, N, P, P, P, N, N, N, P, N, N, P }, + { P, P, N, P, P, P, N, N, N, P, N, N } +}; + +static constexpr float H20[20][20] = { + { P, P, P, P, P, P, P, P, P, P, P, P, P, P, P, P, P, P, P, P }, + { P, N, P, N, N, P, P, P, P, N, P, N, P, N, N, N, N, P, P, N }, + { P, N, N, P, N, N, P, P, P, P, N, P, N, P, N, N, N, N, P, P }, + { P, P, N, N, P, N, N, P, P, P, P, N, P, N, P, N, N, N, N, P }, + { P, P, P, N, N, P, N, N, P, P, P, P, N, P, N, P, N, N, N, N }, + { P, N, P, P, N, N, P, N, N, P, P, P, P, N, P, N, P, N, N, N }, + { P, N, N, P, P, N, N, P, N, N, P, P, P, P, N, P, N, P, N, N }, + { P, N, N, N, P, P, N, N, P, N, N, P, P, P, P, N, P, N, P, N }, + { P, N, N, N, N, P, P, N, N, P, N, N, P, P, P, P, N, P, N, P }, + { P, P, N, N, N, N, P, P, N, N, P, N, N, P, P, P, P, N, P, N }, + { P, N, P, N, N, N, N, P, P, N, N, P, N, N, P, P, P, P, N, P }, + { P, P, N, P, N, N, N, N, P, P, N, N, P, N, N, P, P, P, P, N }, + { P, N, P, N, P, N, N, N, N, P, P, N, N, P, N, N, P, P, P, P }, + { P, P, N, P, N, P, N, N, N, N, P, P, N, N, P, N, N, P, P, P }, + { P, P, P, N, P, N, P, N, N, N, N, P, P, N, N, P, N, N, P, P }, + { P, P, P, P, N, P, N, P, N, N, N, N, P, P, N, N, P, N, N, P }, + { P, P, P, P, P, N, P, N, P, N, N, N, N, P, P, N, N, P, N, N }, + { P, N, P, P, P, P, N, P, N, P, N, N, N, N, P, P, N, N, P, N }, + { P, N, N, P, P, P, P, N, P, N, P, N, N, N, N, P, P, N, N, P }, + { P, P, N, N, P, P, P, P, N, P, N, P, N, N, N, N, P, P, N, N } +}; + +#undef P +#undef N template static void fwht_kernel(const float * __restrict__ src, float * __restrict__ dst, const int64_t n_rows, @@ -80,6 +124,122 @@ static void launch_fwht(const float * src, float * dst, const int64_t n_rows, co }); } +template +static void kronecker_kernel(const float * __restrict__ src, + float * __restrict__ dst, + const int64_t n_rows, + const float scale, + const sycl::nd_item<2> & item) { + static_assert(m == 12 || m == 20, "block size has to be 12 or 20."); + + const sycl::sub_group sg = item.get_sub_group(); + + const int64_t r = item.get_global_id(0); + if (r >= n_rows) { + return; + } + + src += r * N; + dst += r * N; + + constexpr int blocks_per_group = N / m; + constexpr int el_w = blocks_per_group / WARP_SIZE; + static_assert(el_w >= 1 && blocks_per_group % WARP_SIZE == 0, "blocks_per_group must be a multiple of WARP_SIZE"); + float reg[el_w * m]; + const int lane = sg.get_local_linear_id(); + +#pragma unroll + for (int i = 0; i < el_w; ++i) { + const int b_idx = i * WARP_SIZE + lane; + +#pragma unroll + for (int j = 0; j < m; ++j) { + reg[i * m + j] = src[b_idx * m + j] * scale; + } + } + +#pragma unroll + for (int b = 0; b < el_w; ++b) { + float z[m] = { 0.0f }; + +#pragma unroll + for (int i = 0; i < m; ++i) { +#pragma unroll + for (int j = 0; j < m; ++j) { + const float h = (m == 12 ? H12[j][i] : H20[j][i]); + z[i] += reg[b * m + j] * h; + } + } + +#pragma unroll + for (int i = 0; i < m; ++i) { + reg[b * m + i] = z[i]; + } + } + +#pragma unroll + for (int h = 1; h < WARP_SIZE; h *= 2) { +#pragma unroll + for (int j = 0; j < el_w; ++j) { +#pragma unroll + for (int k = 0; k < m; ++k) { + const float val = reg[j * m + k]; + const float val2 = dpct::permute_sub_group_by_xor(sg, val, h, WARP_SIZE); + + reg[j * m + k] = (lane & h) == 0 ? val + val2 : val2 - val; + } + } + } + +#pragma unroll + for (int h = WARP_SIZE; h < blocks_per_group; h *= 2) { + const int step = h / WARP_SIZE; +#pragma unroll + for (int j = 0; j < el_w; j += 2 * step) { +#pragma unroll + for (int s = 0; s < step; ++s) { +#pragma unroll + for (int k = 0; k < m; ++k) { + const float x = reg[(j + s) * m + k]; + const float y = reg[(j + s + step) * m + k]; + + reg[(j + s) * m + k] = x + y; + reg[(j + s + step) * m + k] = x - y; + } + } + } + } + +#pragma unroll + for (int i = 0; i < el_w; ++i) { + const int b_idx = i * WARP_SIZE + lane; +#pragma unroll + for (int k = 0; k < m; ++k) { + dst[b_idx * m + k] = reg[i * m + k]; + } + } +} + +template +static void launch_kronecker(const float * src, + float * dst, + const int64_t n_rows, + const float scale, + dpct::queue_ptr stream) { + constexpr int rows_per_block = 4; + + const int64_t num_blocks = (n_rows + rows_per_block - 1) / rows_per_block; + + // dim 1 is the fastest-varying, so a sub-group is exactly one row's WARP_SIZE lanes. + const sycl::range<2> global(num_blocks * rows_per_block, WARP_SIZE); + const sycl::range<2> local(rows_per_block, WARP_SIZE); + + stream->parallel_for(sycl::nd_range<2>(global, local), + [=](sycl::nd_item<2> item) [[sycl::reqd_sub_group_size(WARP_SIZE)]] { + kronecker_kernel(src, dst, n_rows, scale, item); + }); +} + bool ggml_sycl_op_fwht(ggml_backend_sycl_context & ctx, const ggml_tensor * src, ggml_tensor * dst) { if (src->type != GGML_TYPE_F32 || dst->type != GGML_TYPE_F32) { return false; @@ -113,6 +273,18 @@ bool ggml_sycl_op_fwht(ggml_backend_sycl_context & ctx, const ggml_tensor * src, case 512: launch_fwht<512>(src_d, dst_d, rows, scale, stream); return true; + case 384: + launch_kronecker<384, 12>(src_d, dst_d, rows, scale, stream); + return true; + case 768: + launch_kronecker<768, 12>(src_d, dst_d, rows, scale, stream); + return true; + case 640: + launch_kronecker<640, 20>(src_d, dst_d, rows, scale, stream); + return true; + case 1280: + launch_kronecker<1280, 20>(src_d, dst_d, rows, scale, stream); + return true; default: return false; } From f162a194754838c2468d1d02600182154ba41a63 Mon Sep 17 00:00:00 2001 From: Titaniumtown Date: Tue, 1 Sep 2026 09:04:31 -0700 Subject: [PATCH 071/104] =?UTF-8?q?Revert=20"sycl=20:=20add=20Kronecker=20?= =?UTF-8?q?product=20FWHT=20support=20for=20sizes=20384,=20640,=20768,=201?= =?UTF-8?q?2=E2=80=A6"=20(#28184)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit This reverts commit 1f3d318734c61cf6f3b209726cdd3f9c300a782e. --- ggml/src/ggml-sycl/fwht.cpp | 172 ------------------------------------ 1 file changed, 172 deletions(-) diff --git a/ggml/src/ggml-sycl/fwht.cpp b/ggml/src/ggml-sycl/fwht.cpp index 39f273beaa9..2312b3d131b 100644 --- a/ggml/src/ggml-sycl/fwht.cpp +++ b/ggml/src/ggml-sycl/fwht.cpp @@ -1,50 +1,6 @@ #include "fwht.hpp" #include -#define P 1.0f -#define N -1.0f - -// constant Hadamard matrix via Paley I construction -static constexpr float H12[12][12] = { - { P, P, P, P, P, P, P, P, P, P, P, P }, - { P, N, P, N, P, P, P, N, N, N, P, N }, - { P, N, N, P, N, P, P, P, N, N, N, P }, - { P, P, N, N, P, N, P, P, P, N, N, N }, - { P, N, P, N, N, P, N, P, P, P, N, N }, - { P, N, N, P, N, N, P, N, P, P, P, N }, - { P, N, N, N, P, N, N, P, N, P, P, P }, - { P, P, N, N, N, P, N, N, P, N, P, P }, - { P, P, P, N, N, N, P, N, N, P, N, P }, - { P, P, P, P, N, N, N, P, N, N, P, N }, - { P, N, P, P, P, N, N, N, P, N, N, P }, - { P, P, N, P, P, P, N, N, N, P, N, N } -}; - -static constexpr float H20[20][20] = { - { P, P, P, P, P, P, P, P, P, P, P, P, P, P, P, P, P, P, P, P }, - { P, N, P, N, N, P, P, P, P, N, P, N, P, N, N, N, N, P, P, N }, - { P, N, N, P, N, N, P, P, P, P, N, P, N, P, N, N, N, N, P, P }, - { P, P, N, N, P, N, N, P, P, P, P, N, P, N, P, N, N, N, N, P }, - { P, P, P, N, N, P, N, N, P, P, P, P, N, P, N, P, N, N, N, N }, - { P, N, P, P, N, N, P, N, N, P, P, P, P, N, P, N, P, N, N, N }, - { P, N, N, P, P, N, N, P, N, N, P, P, P, P, N, P, N, P, N, N }, - { P, N, N, N, P, P, N, N, P, N, N, P, P, P, P, N, P, N, P, N }, - { P, N, N, N, N, P, P, N, N, P, N, N, P, P, P, P, N, P, N, P }, - { P, P, N, N, N, N, P, P, N, N, P, N, N, P, P, P, P, N, P, N }, - { P, N, P, N, N, N, N, P, P, N, N, P, N, N, P, P, P, P, N, P }, - { P, P, N, P, N, N, N, N, P, P, N, N, P, N, N, P, P, P, P, N }, - { P, N, P, N, P, N, N, N, N, P, P, N, N, P, N, N, P, P, P, P }, - { P, P, N, P, N, P, N, N, N, N, P, P, N, N, P, N, N, P, P, P }, - { P, P, P, N, P, N, P, N, N, N, N, P, P, N, N, P, N, N, P, P }, - { P, P, P, P, N, P, N, P, N, N, N, N, P, P, N, N, P, N, N, P }, - { P, P, P, P, P, N, P, N, P, N, N, N, N, P, P, N, N, P, N, N }, - { P, N, P, P, P, P, N, P, N, P, N, N, N, N, P, P, N, N, P, N }, - { P, N, N, P, P, P, P, N, P, N, P, N, N, N, N, P, P, N, N, P }, - { P, P, N, N, P, P, P, P, N, P, N, P, N, N, N, N, P, P, N, N } -}; - -#undef P -#undef N template static void fwht_kernel(const float * __restrict__ src, float * __restrict__ dst, const int64_t n_rows, @@ -124,122 +80,6 @@ static void launch_fwht(const float * src, float * dst, const int64_t n_rows, co }); } -template -static void kronecker_kernel(const float * __restrict__ src, - float * __restrict__ dst, - const int64_t n_rows, - const float scale, - const sycl::nd_item<2> & item) { - static_assert(m == 12 || m == 20, "block size has to be 12 or 20."); - - const sycl::sub_group sg = item.get_sub_group(); - - const int64_t r = item.get_global_id(0); - if (r >= n_rows) { - return; - } - - src += r * N; - dst += r * N; - - constexpr int blocks_per_group = N / m; - constexpr int el_w = blocks_per_group / WARP_SIZE; - static_assert(el_w >= 1 && blocks_per_group % WARP_SIZE == 0, "blocks_per_group must be a multiple of WARP_SIZE"); - float reg[el_w * m]; - const int lane = sg.get_local_linear_id(); - -#pragma unroll - for (int i = 0; i < el_w; ++i) { - const int b_idx = i * WARP_SIZE + lane; - -#pragma unroll - for (int j = 0; j < m; ++j) { - reg[i * m + j] = src[b_idx * m + j] * scale; - } - } - -#pragma unroll - for (int b = 0; b < el_w; ++b) { - float z[m] = { 0.0f }; - -#pragma unroll - for (int i = 0; i < m; ++i) { -#pragma unroll - for (int j = 0; j < m; ++j) { - const float h = (m == 12 ? H12[j][i] : H20[j][i]); - z[i] += reg[b * m + j] * h; - } - } - -#pragma unroll - for (int i = 0; i < m; ++i) { - reg[b * m + i] = z[i]; - } - } - -#pragma unroll - for (int h = 1; h < WARP_SIZE; h *= 2) { -#pragma unroll - for (int j = 0; j < el_w; ++j) { -#pragma unroll - for (int k = 0; k < m; ++k) { - const float val = reg[j * m + k]; - const float val2 = dpct::permute_sub_group_by_xor(sg, val, h, WARP_SIZE); - - reg[j * m + k] = (lane & h) == 0 ? val + val2 : val2 - val; - } - } - } - -#pragma unroll - for (int h = WARP_SIZE; h < blocks_per_group; h *= 2) { - const int step = h / WARP_SIZE; -#pragma unroll - for (int j = 0; j < el_w; j += 2 * step) { -#pragma unroll - for (int s = 0; s < step; ++s) { -#pragma unroll - for (int k = 0; k < m; ++k) { - const float x = reg[(j + s) * m + k]; - const float y = reg[(j + s + step) * m + k]; - - reg[(j + s) * m + k] = x + y; - reg[(j + s + step) * m + k] = x - y; - } - } - } - } - -#pragma unroll - for (int i = 0; i < el_w; ++i) { - const int b_idx = i * WARP_SIZE + lane; -#pragma unroll - for (int k = 0; k < m; ++k) { - dst[b_idx * m + k] = reg[i * m + k]; - } - } -} - -template -static void launch_kronecker(const float * src, - float * dst, - const int64_t n_rows, - const float scale, - dpct::queue_ptr stream) { - constexpr int rows_per_block = 4; - - const int64_t num_blocks = (n_rows + rows_per_block - 1) / rows_per_block; - - // dim 1 is the fastest-varying, so a sub-group is exactly one row's WARP_SIZE lanes. - const sycl::range<2> global(num_blocks * rows_per_block, WARP_SIZE); - const sycl::range<2> local(rows_per_block, WARP_SIZE); - - stream->parallel_for(sycl::nd_range<2>(global, local), - [=](sycl::nd_item<2> item) [[sycl::reqd_sub_group_size(WARP_SIZE)]] { - kronecker_kernel(src, dst, n_rows, scale, item); - }); -} - bool ggml_sycl_op_fwht(ggml_backend_sycl_context & ctx, const ggml_tensor * src, ggml_tensor * dst) { if (src->type != GGML_TYPE_F32 || dst->type != GGML_TYPE_F32) { return false; @@ -273,18 +113,6 @@ bool ggml_sycl_op_fwht(ggml_backend_sycl_context & ctx, const ggml_tensor * src, case 512: launch_fwht<512>(src_d, dst_d, rows, scale, stream); return true; - case 384: - launch_kronecker<384, 12>(src_d, dst_d, rows, scale, stream); - return true; - case 768: - launch_kronecker<768, 12>(src_d, dst_d, rows, scale, stream); - return true; - case 640: - launch_kronecker<640, 20>(src_d, dst_d, rows, scale, stream); - return true; - case 1280: - launch_kronecker<1280, 20>(src_d, dst_d, rows, scale, stream); - return true; default: return false; } From 2c486783abfafacdffbc661f0bba0b31e8994eb7 Mon Sep 17 00:00:00 2001 From: anujj Date: Wed, 2 Sep 2026 01:18:47 +0530 Subject: [PATCH 072/104] cuda: fuse MoE weighted expert reduction (llama/25952) * cuda : fuse MoE weighted reduction (mul + view + add) The MoE combine tail currently writes weighted expert outputs to global memory before reducing them. That intermediate global-memory traffic is the main cost. The production baseline generally runs two physical fused kernels; this path runs one. This change matches the full expert-weighting plus ordered-reduction subgraph and replaces it with one weighted-reduction kernel. Supported graphs: - unscaled: experts * router_weights - scaled: (experts * expert_scale) * router_weights k = 2..15 is handled by one runtime-k kernel. Matching is structural: op sequence, shapes, strides, expert views, and the left-to-right ADD chain. The fused kernel keeps that same reduction order. Results are not claimed bit-identical; CUDA FP32 contraction can change rounding slightly. Allocator integration uses add_alloc_dep from the graph-optimizer API so experts, router weights, and optional expert scales stay live until the fused destination is written. Memory ranges are rechecked before the fused kernel runs. Unrecognized or unsafe graphs are left alone and keep the existing per-op path. Set GGML_CUDA_MOE_WEIGHTED_REDUCTION=0 to disable the fusion. test-backend-ops covers scaled/unscaled, aligned/unaligned, and representative values across k=2..15, plus a k=16 case that must stay on the per-op path. * Pruned the test matrix from 15 to 6 * Addressed the aman and olivers review comments --- ggml/src/ggml-cuda/ggml-cuda.cu | 181 +++++++++++++++++- ggml/src/ggml-cuda/moe-weighted-reduction.cu | 65 +++++++ ggml/src/ggml-cuda/moe-weighted-reduction.cuh | 7 + 3 files changed, 251 insertions(+), 2 deletions(-) create mode 100644 ggml/src/ggml-cuda/moe-weighted-reduction.cu create mode 100644 ggml/src/ggml-cuda/moe-weighted-reduction.cuh diff --git a/ggml/src/ggml-cuda/ggml-cuda.cu b/ggml/src/ggml-cuda/ggml-cuda.cu index 31f5aeeacc3..f4af82688ac 100644 --- a/ggml/src/ggml-cuda/ggml-cuda.cu +++ b/ggml/src/ggml-cuda/ggml-cuda.cu @@ -32,6 +32,7 @@ #include "ggml-cuda/mmq.cuh" #include "ggml-cuda/mmvf.cuh" #include "ggml-cuda/mmvq.cuh" +#include "ggml-cuda/moe-weighted-reduction.cuh" #include "ggml-cuda/norm.cuh" #include "ggml-cuda/opt-step-adamw.cuh" #include "ggml-cuda/opt-step-sgd.cuh" @@ -3026,6 +3027,150 @@ static bool ggml_cuda_check_fusion_memory_ranges(const ggml_cgraph * cgraph, return is_ok; } +// The long form spans 2*k + 1 nodes. ggml_can_fuse_subgraph() accepts at most +// 31 nodes, so k <= 15; larger values use the per-operation path. +static constexpr int MOE_WEIGHTED_REDUCTION_MAX_EXPERTS = 15; + +struct ggml_cuda_moe_weighted_reduction_match { + const ggml_tensor * experts = nullptr; + const ggml_tensor * expert_scale = nullptr; + const ggml_tensor * weights = nullptr; + ggml_tensor * dst = nullptr; + int node_count = 0; +}; + +static bool ggml_cuda_match_moe_weighted_reduction( + const ggml_cgraph * cgraph, + int node_idx, + ggml_cuda_moe_weighted_reduction_match & match) { + const ggml_tensor * first = cgraph->nodes[node_idx]; + if (first->op != GGML_OP_MUL || first->type != GGML_TYPE_F32 || !ggml_is_contiguous(first)) { + return false; + } + + auto split_mul = [](const ggml_tensor * mul, const ggml_tensor *& full, const ggml_tensor *& broadcast) { + auto is_weights = [mul](const ggml_tensor * tensor) { + return tensor && tensor->type == GGML_TYPE_F32 && ggml_is_contiguous(tensor) && tensor->ne[0] == 1 && + tensor->ne[1] == mul->ne[1] && tensor->ne[2] == mul->ne[2] && tensor->ne[3] == mul->ne[3]; + }; + auto is_experts = [mul](const ggml_tensor * tensor) { + return tensor && tensor->type == GGML_TYPE_F32 && ggml_is_contiguous(tensor) && + ggml_are_same_shape(tensor, mul); + }; + + if (is_experts(mul->src[0]) && is_weights(mul->src[1])) { + full = mul->src[0]; + broadcast = mul->src[1]; + return true; + } + if (is_experts(mul->src[1]) && is_weights(mul->src[0])) { + full = mul->src[1]; + broadcast = mul->src[0]; + return true; + } + return false; + }; + + const ggml_tensor * weighted = first; + const ggml_tensor * experts = nullptr; + const ggml_tensor * expert_scale = nullptr; + const ggml_tensor * weights = nullptr; + int mul_count = 1; + + // Match both structural forms: + // (experts * expert_scale) * router_weight + // experts * router_weight + // The matcher does not depend on the model or quantization type. + if (node_idx + 1 < cgraph->n_nodes) { + const ggml_tensor * second = cgraph->nodes[node_idx + 1]; + const ggml_tensor * scaled = nullptr; + const ggml_tensor * route = nullptr; + const ggml_tensor * raw = nullptr; + const ggml_tensor * scale = nullptr; + if (second->op == GGML_OP_MUL && second->type == GGML_TYPE_F32 && ggml_is_contiguous(second) && + split_mul(second, scaled, route) && scaled == first && split_mul(first, raw, scale)) { + weighted = second; + experts = raw; + expert_scale = scale; + weights = route; + mul_count = 2; + } + } + + if (experts == nullptr && !split_mul(first, experts, weights)) { + return false; + } + + const int n_expert_used = (int) weighted->ne[1]; + const int64_t n_tokens = weighted->ne[2] * weighted->ne[3]; + if (n_expert_used < 2 || n_expert_used > MOE_WEIGHTED_REDUCTION_MAX_EXPERTS || n_tokens <= 0) { + return false; + } + + const int node_count = 2 * n_expert_used + mul_count - 1; + if (node_idx + node_count > cgraph->n_nodes) { + return false; + } + + std::vector ops(node_count, GGML_OP_VIEW); + ops[0] = GGML_OP_MUL; + if (mul_count == 2) { + ops[1] = GGML_OP_MUL; + } + std::vector views; + views.reserve(n_expert_used); + const ggml_tensor * previous = nullptr; + int n_adds = 0; + for (int offset = mul_count; offset < node_count; ++offset) { + const ggml_tensor * candidate = cgraph->nodes[node_idx + offset]; + ops[offset] = candidate->op; + + if (candidate->op == GGML_OP_VIEW) { + const int expert = (int) views.size(); + if (expert >= n_expert_used || candidate->src[0] != weighted || candidate->view_src != weighted || + candidate->type != GGML_TYPE_F32 || candidate->ne[0] != weighted->ne[0] || + candidate->ne[1] != n_tokens || candidate->ne[2] != 1 || candidate->ne[3] != 1 || + candidate->nb[0] != weighted->nb[0] || candidate->nb[1] != weighted->nb[2] || + candidate->view_offs != (size_t) expert * weighted->nb[1]) { + return false; + } + views.push_back(candidate); + continue; + } + + if (candidate->op != GGML_OP_ADD || views.size() < 2 || n_adds + 1 >= (int) views.size()) { + return false; + } + const ggml_tensor * lhs = n_adds == 0 ? views[0] : previous; + const ggml_tensor * rhs = views[n_adds + 1]; + if (candidate->src[0] != lhs || candidate->src[1] != rhs || candidate->type != GGML_TYPE_F32) { + return false; + } + previous = candidate; + ++n_adds; + } + + if ((int) views.size() != n_expert_used || n_adds != n_expert_used - 1 || previous == nullptr) { + return false; + } + if (!ggml_is_contiguous(previous) || previous->ne[0] != weighted->ne[0] || + previous->ne[1] != n_tokens || previous->ne[2] != 1 || previous->ne[3] != 1) { + return false; + } + + const int output_idx = node_idx + node_count - 1; + if (!ggml_can_fuse_subgraph(cgraph, node_idx, node_count, ops.data(), &output_idx, 1)) { + return false; + } + + match.experts = experts; + match.expert_scale = expert_scale; + match.weights = weights; + match.dst = cgraph->nodes[output_idx]; + match.node_count = node_count; + return true; +} + static bool ggml_cuda_can_fuse(const struct ggml_cgraph * cgraph, int node_idx, @@ -3288,6 +3433,18 @@ static int ggml_cuda_try_fuse(ggml_backend_cuda_context * cuda_ctx, ggml_cgraph ggml_tensor * node = cgraph->nodes[i]; + if (node->op == GGML_OP_MUL) { + ggml_cuda_moe_weighted_reduction_match match; + if (ggml_cuda_match_moe_weighted_reduction(cgraph, i, match)) { + const int output_idx = i + match.node_count - 1; + if (ggml_cuda_check_fusion_memory_ranges(cgraph, i, match.node_count, &output_idx, 1)) { + ggml_cuda_op_moe_weighted_reduction( + *cuda_ctx, match.experts, match.expert_scale, match.weights, match.dst); + return match.node_count - 1; + } + } + } + // gated_delta_net -> cpy: scatter recurrent-state snapshots into the cache if (node->op == GGML_OP_GATED_DELTA_NET) { ggml_cuda_gated_delta_net_fused_cache fused_state_cpy; @@ -4340,10 +4497,30 @@ static void ggml_backend_cuda_event_wait(ggml_backend_t backend, ggml_backend_ev } static void ggml_backend_cuda_graph_optimize(ggml_backend_t backend, ggml_cgraph * cgraph, ggml_backend_graph_optimize_params * params) { - GGML_UNUSED(params); - ggml_backend_cuda_context * cuda_ctx = (ggml_backend_cuda_context *) backend->context; + static const bool disable_fusion = getenv("GGML_CUDA_DISABLE_FUSION") != nullptr && std::atoi(getenv("GGML_CUDA_DISABLE_FUSION")); + if (!disable_fusion) { + for (int i = 0; i < cgraph->n_nodes; ++i) { + if (cgraph->nodes[i]->op != GGML_OP_MUL) { + continue; + } + + ggml_cuda_moe_weighted_reduction_match match; + if (!ggml_cuda_match_moe_weighted_reduction(cgraph, i, match)) { + continue; + } + + params->add_alloc_dep(params->user_data, const_cast(match.experts), match.dst); + params->add_alloc_dep(params->user_data, const_cast(match.weights), match.dst); + if (match.expert_scale != nullptr) { + params->add_alloc_dep( + params->user_data, const_cast(match.expert_scale), match.dst); + } + i += match.node_count - 1; + } + } + #ifdef USE_CUDA_GRAPH const void * graph_key = ggml_cuda_graph_get_key(cgraph); const bool use_cuda_graph = ggml_cuda_graph_set_enabled(cuda_ctx, graph_key); diff --git a/ggml/src/ggml-cuda/moe-weighted-reduction.cu b/ggml/src/ggml-cuda/moe-weighted-reduction.cu new file mode 100644 index 00000000000..11ec58497f1 --- /dev/null +++ b/ggml/src/ggml-cuda/moe-weighted-reduction.cu @@ -0,0 +1,65 @@ +#include "moe-weighted-reduction.cuh" + +static __global__ void moe_weighted_reduction_f32(const float * __restrict__ experts, + const float * __restrict__ expert_scale, + const float * __restrict__ weights, + float * __restrict__ dst, + const int64_t n_embd, + const int n_expert_used) { + const int64_t token = blockIdx.x; + const int64_t col = (int64_t) blockIdx.y * blockDim.x + threadIdx.x; + if (col >= n_embd) { + return; + } + + const uint64_t first_row = (uint64_t) token * n_expert_used; + const float first_scale = expert_scale != nullptr ? expert_scale[first_row] : 1.0f; + float sum = (experts[first_row * n_embd + col] * first_scale) * weights[first_row]; + + for (int expert = 1; expert < n_expert_used; ++expert) { + const uint64_t row = first_row + expert; + const float scale = expert_scale != nullptr ? expert_scale[row] : 1.0f; + sum += (experts[row * n_embd + col] * scale) * weights[row]; + } + dst[token * n_embd + col] = sum; +} + +static void launch_moe_weighted_reduction(const float * experts, + const float * expert_scale, + const float * weights, + float * dst, + int64_t n_embd, + int64_t n_tokens, + int n_expert_used, + cudaStream_t stream) { + constexpr int threads = 256; + const dim3 blocks(n_tokens, (n_embd + threads - 1) / threads, 1); + moe_weighted_reduction_f32 + <<>>(experts, expert_scale, weights, dst, n_embd, n_expert_used); +} + +void ggml_cuda_op_moe_weighted_reduction(ggml_backend_cuda_context & ctx, + const ggml_tensor * experts, + const ggml_tensor * expert_scale, + const ggml_tensor * weights, + ggml_tensor * dst) { + GGML_ASSERT(experts->type == GGML_TYPE_F32); + GGML_ASSERT(weights->type == GGML_TYPE_F32); + GGML_ASSERT(expert_scale == nullptr || expert_scale->type == GGML_TYPE_F32); + GGML_ASSERT(dst->type == GGML_TYPE_F32); + GGML_ASSERT(ggml_is_contiguous(experts)); + GGML_ASSERT(ggml_is_contiguous(weights)); + GGML_ASSERT(expert_scale == nullptr || ggml_is_contiguous(expert_scale)); + GGML_ASSERT(ggml_is_contiguous(dst)); + + const int64_t n_embd = experts->ne[0]; + const int64_t n_expert_used = experts->ne[1]; + const int64_t n_tokens = experts->ne[2] * experts->ne[3]; + cudaStream_t stream = ctx.stream(); + + launch_moe_weighted_reduction((const float *) experts->data, + expert_scale ? (const float *) expert_scale->data : nullptr, + (const float *) weights->data, + (float *) dst->data, n_embd, n_tokens, (int) n_expert_used, stream); + CUDA_CHECK(cudaGetLastError()); +} diff --git a/ggml/src/ggml-cuda/moe-weighted-reduction.cuh b/ggml/src/ggml-cuda/moe-weighted-reduction.cuh new file mode 100644 index 00000000000..b72f947ab39 --- /dev/null +++ b/ggml/src/ggml-cuda/moe-weighted-reduction.cuh @@ -0,0 +1,7 @@ +#include "common.cuh" + +void ggml_cuda_op_moe_weighted_reduction(ggml_backend_cuda_context & ctx, + const ggml_tensor * experts, + const ggml_tensor * expert_scale, + const ggml_tensor * weights, + ggml_tensor * dst); From fcc2feee2fee5e6df76847aeaaa5bdb1eead3352 Mon Sep 17 00:00:00 2001 From: Jhen-Jie Hong Date: Wed, 2 Sep 2026 07:45:56 +0800 Subject: [PATCH 073/104] metal : add metallib build support for xcframework (llama/28163) --- ggml/CMakeLists.txt | 2 + ggml/src/ggml-metal/CMakeLists.txt | 68 +++++++++++++++++++++--------- 2 files changed, 50 insertions(+), 20 deletions(-) diff --git a/ggml/CMakeLists.txt b/ggml/CMakeLists.txt index c4a8450d1ca..0ac2b15c48c 100644 --- a/ggml/CMakeLists.txt +++ b/ggml/CMakeLists.txt @@ -242,6 +242,8 @@ option(GGML_METAL_EMBED_LIBRARY "ggml: embed Metal library" set (GGML_METAL_MACOSX_VERSION_MIN "" CACHE STRING "ggml: metal minimum macOS version") set (GGML_METAL_STD "" CACHE STRING "ggml: metal standard version (-std flag)") +set (GGML_METAL_TARGET_OS "macos" CACHE STRING + "ggml: metal -mtargetos OS name (macos, ios, xros, tvos)") option(GGML_OPENMP "ggml: use OpenMP" ON) option(GGML_OPENMP_FETCH "ggml: fetch LLVM OpenMP" OFF) option(GGML_RPC "ggml: use RPC" OFF) diff --git a/ggml/src/ggml-metal/CMakeLists.txt b/ggml/src/ggml-metal/CMakeLists.txt index 2094a409f96..a661e710a2f 100644 --- a/ggml/src/ggml-metal/CMakeLists.txt +++ b/ggml/src/ggml-metal/CMakeLists.txt @@ -127,6 +127,18 @@ else() configure_file(${src} ${CMAKE_RUNTIME_OUTPUT_DIRECTORY}/${src} COPYONLY) endforeach() + # CMAKE_OSX_SYSROOT is an SDK name or path - xcrun accepts both + set(METAL_SDK ${CMAKE_OSX_SYSROOT}) + if (NOT METAL_SDK) + set(METAL_SDK macosx) + endif() + + if (CMAKE_OSX_SYSROOT MATCHES "[Ss]imulator") + set(METAL_TARGET_SIM "-simulator") + else() + set(METAL_TARGET_SIM "") + endif() + if (GGML_METAL_SHADER_DEBUG) # note: disabling fast math is needed in order to pass tests/test-backend-ops # note: adding -fno-inline fixes the tests when using MTL_SHADER_VALIDATION=1 @@ -138,9 +150,19 @@ else() set(XC_FLAGS -O3) endif() + execute_process(COMMAND xcrun -sdk ${METAL_SDK} --show-sdk-version OUTPUT_VARIABLE METAL_SDK_VERSION OUTPUT_STRIP_TRAILING_WHITESPACE) + if (METAL_SDK_VERSION VERSION_GREATER_EQUAL 26.0) + set(GGML_METAL_HAS_TENSOR_LIB ON) + else() + message(STATUS "Metal SDK ${METAL_SDK_VERSION} does not support the tensor API, skipping ggml-tensor.metallib") + endif() + if (GGML_METAL_MACOSX_VERSION_MIN) message(STATUS "Adding -mmacosx-version-min=${GGML_METAL_MACOSX_VERSION_MIN} flag to metal compilation") list (APPEND XC_FLAGS -mmacosx-version-min=${GGML_METAL_MACOSX_VERSION_MIN}) + elseif (NOT GGML_METAL_TARGET_OS STREQUAL "macos" AND CMAKE_OSX_DEPLOYMENT_TARGET) + message(STATUS "Adding -mtargetos=${GGML_METAL_TARGET_OS}${CMAKE_OSX_DEPLOYMENT_TARGET}${METAL_TARGET_SIM} flag to metal compilation") + list (APPEND XC_FLAGS -mtargetos=${GGML_METAL_TARGET_OS}${CMAKE_OSX_DEPLOYMENT_TARGET}${METAL_TARGET_SIM}) endif() if (GGML_METAL_STD) @@ -156,33 +178,41 @@ else() list(APPEND AIR_FILES ${AIR}) add_custom_command( OUTPUT ${AIR} - COMMAND xcrun -sdk macosx metal ${XC_FLAGS} -I ${CMAKE_RUNTIME_OUTPUT_DIRECTORY} -c ${CMAKE_RUNTIME_OUTPUT_DIRECTORY}/${src} -o ${AIR} + COMMAND xcrun -sdk ${METAL_SDK} metal ${XC_FLAGS} -I ${CMAKE_RUNTIME_OUTPUT_DIRECTORY} -c ${CMAKE_RUNTIME_OUTPUT_DIRECTORY}/${src} -o ${AIR} DEPENDS ${src} kernels/common.h kernels/dequantize.h kernels/quantize.h ${METALLIB_COMMON} ggml-metal-impl.h COMMENT "Compiling ${src}" VERBATIM ) endforeach() + set(METALLIB_FILES ${CMAKE_RUNTIME_OUTPUT_DIRECTORY}/default.metallib) + # the tensor API kernels go in a separate metallib, loaded only where supported - set(AIR_MM_TENSOR "${CMAKE_RUNTIME_OUTPUT_DIRECTORY}/mul_mm_tensor.air") - add_custom_command( - OUTPUT ${AIR_MM_TENSOR} - COMMAND xcrun -sdk macosx metal ${XC_FLAGS} -DGGML_METAL_HAS_TENSOR -I ${CMAKE_RUNTIME_OUTPUT_DIRECTORY} -c ${CMAKE_RUNTIME_OUTPUT_DIRECTORY}/kernels/mul_mm.metal -o ${AIR_MM_TENSOR} - DEPENDS kernels/mul_mm.metal kernels/common.h kernels/dequantize.h ${METALLIB_COMMON} ggml-metal-impl.h - COMMENT "Compiling kernels/mul_mm.metal (tensor API)" - VERBATIM - ) + if (GGML_METAL_HAS_TENSOR_LIB) + set(AIR_MM_TENSOR "${CMAKE_RUNTIME_OUTPUT_DIRECTORY}/mul_mm_tensor.air") + # the tensor API needs OS 26+ + set(XC_FLAGS_TENSOR ${XC_FLAGS} -mtargetos=${GGML_METAL_TARGET_OS}26.0${METAL_TARGET_SIM}) + add_custom_command( + OUTPUT ${AIR_MM_TENSOR} + COMMAND xcrun -sdk ${METAL_SDK} metal ${XC_FLAGS_TENSOR} -DGGML_METAL_HAS_TENSOR -I ${CMAKE_RUNTIME_OUTPUT_DIRECTORY} -c ${CMAKE_RUNTIME_OUTPUT_DIRECTORY}/kernels/mul_mm.metal -o ${AIR_MM_TENSOR} + DEPENDS kernels/mul_mm.metal kernels/common.h kernels/dequantize.h ${METALLIB_COMMON} ggml-metal-impl.h + COMMENT "Compiling kernels/mul_mm.metal (tensor API)" + VERBATIM + ) - add_custom_command( - OUTPUT ${CMAKE_RUNTIME_OUTPUT_DIRECTORY}/ggml-tensor.metallib - COMMAND xcrun -sdk macosx metallib ${AIR_MM_TENSOR} -o ${CMAKE_RUNTIME_OUTPUT_DIRECTORY}/ggml-tensor.metallib - DEPENDS ${AIR_MM_TENSOR} - COMMENT "Linking tensor API Metal kernels into ggml-tensor.metallib" - ) + add_custom_command( + OUTPUT ${CMAKE_RUNTIME_OUTPUT_DIRECTORY}/ggml-tensor.metallib + COMMAND xcrun -sdk ${METAL_SDK} metallib ${AIR_MM_TENSOR} -o ${CMAKE_RUNTIME_OUTPUT_DIRECTORY}/ggml-tensor.metallib + DEPENDS ${AIR_MM_TENSOR} + COMMENT "Linking tensor API Metal kernels into ggml-tensor.metallib" + ) + + list(APPEND METALLIB_FILES ${CMAKE_RUNTIME_OUTPUT_DIRECTORY}/ggml-tensor.metallib) + endif() add_custom_command( OUTPUT ${CMAKE_RUNTIME_OUTPUT_DIRECTORY}/default.metallib - COMMAND xcrun -sdk macosx metallib ${AIR_FILES} -o ${CMAKE_RUNTIME_OUTPUT_DIRECTORY}/default.metallib + COMMAND xcrun -sdk ${METAL_SDK} metallib ${AIR_FILES} -o ${CMAKE_RUNTIME_OUTPUT_DIRECTORY}/default.metallib COMMAND rm -f ${CMAKE_RUNTIME_OUTPUT_DIRECTORY}/ggml-common.h COMMAND rm -f ${CMAKE_RUNTIME_OUTPUT_DIRECTORY}/ggml-metal-impl.h COMMAND rm -rf ${CMAKE_RUNTIME_OUTPUT_DIRECTORY}/kernels @@ -192,8 +222,7 @@ else() add_custom_target( ggml-metal-lib ALL - DEPENDS ${CMAKE_RUNTIME_OUTPUT_DIRECTORY}/default.metallib - ${CMAKE_RUNTIME_OUTPUT_DIRECTORY}/ggml-tensor.metallib + DEPENDS ${METALLIB_FILES} ) endif() # GGML_METAL_EMBED_LIBRARY @@ -205,8 +234,7 @@ if (NOT GGML_METAL_EMBED_LIBRARY) ) install( - FILES ${CMAKE_RUNTIME_OUTPUT_DIRECTORY}/default.metallib - ${CMAKE_RUNTIME_OUTPUT_DIRECTORY}/ggml-tensor.metallib + FILES ${METALLIB_FILES} DESTINATION ${CMAKE_INSTALL_BINDIR} ) endif() From c94921f8c6787ee12af882672ed91af9982b91cc Mon Sep 17 00:00:00 2001 From: Trivikram Reddy <127072883+trivikram-reddy1@users.noreply.github.com> Date: Wed, 2 Sep 2026 00:20:29 -0500 Subject: [PATCH 074/104] hexagon: add missing FARF logs for cpy/get_rows/set_rows/gdn ops (llama/28217) * hexagon: fix bug ne[2] printed in proc_op_req prep-src log * hexagon: add shape/VTCM farf logs to cpy, get/set rows, gdn --- ggml/src/ggml-hexagon/htp/cpy-ops.c | 4 ++++ ggml/src/ggml-hexagon/htp/gated-delta-net-ops.c | 9 +++++++++ ggml/src/ggml-hexagon/htp/get-rows-ops.c | 8 ++++++++ ggml/src/ggml-hexagon/htp/main.c | 2 +- ggml/src/ggml-hexagon/htp/set-rows-ops.c | 8 ++++++++ 5 files changed, 30 insertions(+), 1 deletion(-) diff --git a/ggml/src/ggml-hexagon/htp/cpy-ops.c b/ggml/src/ggml-hexagon/htp/cpy-ops.c index c945425dab3..b151b757f41 100644 --- a/ggml/src/ggml-hexagon/htp/cpy-ops.c +++ b/ggml/src/ggml-hexagon/htp/cpy-ops.c @@ -323,6 +323,10 @@ int op_cpy(struct htp_ops_context * octx) { return HTP_STATUS_NO_SUPPORT; } + FARF(HIGH, "cpy-%s-%s: (%ux%ux%ux%u) -> (%ux%ux%ux%u) : use_dma=%d n_threads %u\n", + src0->type == HTP_TYPE_F32 ? "f32" : "f16", dst->type == HTP_TYPE_F32 ? "f32" : "f16", + ne00, ne01, ne02, ne03, ne0, ne1, ne2, ne3, use_dma, n_threads); + if (use_dma) { cpy_dma_sametype_sameshape(octx, dst, src0, ct.src0_type_size, ne00, ne01, ne02, ne03, nb01, nb02, nb03, nb1, nb2, nb3); } else { diff --git a/ggml/src/ggml-hexagon/htp/gated-delta-net-ops.c b/ggml/src/ggml-hexagon/htp/gated-delta-net-ops.c index 35518e6111c..96655215298 100644 --- a/ggml/src/ggml-hexagon/htp/gated-delta-net-ops.c +++ b/ggml/src/ggml-hexagon/htp/gated-delta-net-ops.c @@ -1138,6 +1138,15 @@ int op_gated_delta_net(struct htp_ops_context * octx) { gctx.vtcm_base = octx->ctx->vtcm_base; gctx.vtcm_per_thread = 2 * state_aligned; + FARF(HIGH, "gated-delta-net-f32: q(%ux%ux%ux%u) k(%ux%ux%ux%u) v(%ux%ux%ux%u) state(%ux%ux%ux%u) -> (%ux%ux%ux%u) : " + "vtcm-size %zu n_threads %u\n", + q->ne[0], q->ne[1], q->ne[2], q->ne[3], + k->ne[0], k->ne[1], k->ne[2], k->ne[3], + v->ne[0], v->ne[1], v->ne[2], v->ne[3], + state->ne[0], state->ne[1], state->ne[2], state->ne[3], + dst->ne[0], dst->ne[1], dst->ne[2], dst->ne[3], + gctx.vtcm_per_thread * octx->n_threads, octx->n_threads); + if (n_tokens == 1) { worker_pool_run_func(octx->ctx->worker_pool, gated_delta_net_f32_tg_thread, &gctx, octx->n_threads); } else { diff --git a/ggml/src/ggml-hexagon/htp/get-rows-ops.c b/ggml/src/ggml-hexagon/htp/get-rows-ops.c index 05769d17f74..a87962d2291 100644 --- a/ggml/src/ggml-hexagon/htp/get-rows-ops.c +++ b/ggml/src/ggml-hexagon/htp/get-rows-ops.c @@ -247,6 +247,14 @@ int op_get_rows(struct htp_ops_context * octx) { } } + FARF(HIGH, "get-rows: (%ux%ux%ux%u) x (%ux%ux%ux%u) -> (%ux%ux%ux%u) : src0-vtcm-size %zu dst-vtcm-size %zu use_dma=%d n_threads %d\n", + octx->src[0]->ne[0], octx->src[0]->ne[1], octx->src[0]->ne[2], octx->src[0]->ne[3], + octx->src[1]->ne[0], octx->src[1]->ne[1], octx->src[1]->ne[2], octx->src[1]->ne[3], + octx->dst->ne[0], octx->dst->ne[1], octx->dst->ne[2], octx->dst->ne[3], + grctx.vtcm_layout.src0_bytes_per_thread * kparams->n_threads, + grctx.vtcm_layout.dst_bytes_per_thread * kparams->n_threads, + kparams->use_dma, kparams->n_threads); + work_queue_run(octx->ctx->work_queue, q_func, &grctx, kparams->n_threads); return HTP_STATUS_OK; } diff --git a/ggml/src/ggml-hexagon/htp/main.c b/ggml/src/ggml-hexagon/htp/main.c index 27d1dedcdf0..72cf02a326b 100644 --- a/ggml/src/ggml-hexagon/htp/main.c +++ b/ggml/src/ggml-hexagon/htp/main.c @@ -981,7 +981,7 @@ static int proc_op_req(struct htp_ops_context * octx, struct htp_tensor *tens, u octx->src_dma[i] = octx->ctx->dma; // FIXME: ? octx->ctx->dma_cached : octx->ctx->dma; FARF(HIGH, "prep-src #%u: data %p size %u : %u:%u:%u:%u", op->src[i], (void*) src->data, src->size, - src->ne[0], src->ne[1], src->ne[3], src->ne[3]); + src->ne[0], src->ne[1], src->ne[2], src->ne[3]); } htp_tensor_flush_all(octx->ctx, octx->src, HTP_OP_MAX_INPUTS); diff --git a/ggml/src/ggml-hexagon/htp/set-rows-ops.c b/ggml/src/ggml-hexagon/htp/set-rows-ops.c index fa14bf0ef6b..340a497f7a2 100644 --- a/ggml/src/ggml-hexagon/htp/set-rows-ops.c +++ b/ggml/src/ggml-hexagon/htp/set-rows-ops.c @@ -216,6 +216,14 @@ int op_set_rows(struct htp_ops_context * octx) { default: return HTP_STATUS_NO_SUPPORT; } + FARF(HIGH, "set-rows: (%ux%ux%ux%u) x (%ux%ux%ux%u) -> (%ux%ux%ux%u) : src0-vtcm-size %zu dst-vtcm-size %zu n_threads %d\n", + octx->src[0]->ne[0], octx->src[0]->ne[1], octx->src[0]->ne[2], octx->src[0]->ne[3], + octx->src[1]->ne[0], octx->src[1]->ne[1], octx->src[1]->ne[2], octx->src[1]->ne[3], + octx->dst->ne[0], octx->dst->ne[1], octx->dst->ne[2], octx->dst->ne[3], + srctx.vtcm_layout.src0_bytes_per_thread * kparams->n_threads, + srctx.vtcm_layout.dst_bytes_per_thread * kparams->n_threads, + kparams->n_threads); + work_queue_run(octx->ctx->work_queue, q_func, &srctx, kparams->n_threads); return HTP_STATUS_OK; From 35133c94c5aa00f39ad4d79513c5e57cd6de0cb3 Mon Sep 17 00:00:00 2001 From: Hongqiang Wang Date: Tue, 1 Sep 2026 22:28:45 -0700 Subject: [PATCH 075/104] =?UTF-8?q?opencl:=20fix=20out=E2=80=90of=E2=80=90?= =?UTF-8?q?bound=20reads=20in=20the=20Adreno=20image=20kernels=20(#27632)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit * opencl: clamp the q4_K decode GEMV's fetch row on a padded x-grid * opencl: enforce the tiling contract of the image KQ/KQV GEMMs * opencl: decide the image KQ/KQV split at the dispatch, not from strides --- ggml/src/ggml-opencl/ggml-opencl.cpp | 83 +++++++++++++++---- .../kernels/gemv_noshuffle_q4_k_f32.cl | 39 ++++++--- 2 files changed, 95 insertions(+), 27 deletions(-) diff --git a/ggml/src/ggml-opencl/ggml-opencl.cpp b/ggml/src/ggml-opencl/ggml-opencl.cpp index 34d58f4ee81..d95123eb151 100644 --- a/ggml/src/ggml-opencl/ggml-opencl.cpp +++ b/ggml/src/ggml-opencl/ggml-opencl.cpp @@ -16254,7 +16254,13 @@ static void ggml_cl_conv_2d(ggml_backend_t backend, const ggml_tensor * src0, co backend_ctx->enqueue_ndrange_kernel(kernel, 2, global_work_size, local_work_size, dst); } -static void ggml_cl_mul_mat_kq_kqv_adreno(ggml_backend_t backend, const ggml_tensor * src0, const ggml_tensor * src1, ggml_tensor * dst) { +// is_kq selects which of the two products this call is, and it is decided by the +// CALLER -- the two admission arms in ggml_cl_mul_mat, each of which knows which +// one it matched. It used to be re-derived here from nb01 > nb02, i.e. "K is +// head-major, V^T is not". That discriminator COLLAPSES at n_head_kv == 1, where +// the two strides are equal because there is only one head to order, so nothing +// here could tell a KQ from a KQV. Pass it in rather than infer it. +static void ggml_cl_mul_mat_kq_kqv_adreno(ggml_backend_t backend, const ggml_tensor * src0, const ggml_tensor * src1, ggml_tensor * dst, bool is_kq) { ggml_backend_opencl_context *backend_ctx = (ggml_backend_opencl_context *)backend->context; ggml_tensor_extra_cl * extra0 = (ggml_tensor_extra_cl *)src0->extra; @@ -16296,19 +16302,14 @@ static void ggml_cl_mul_mat_kq_kqv_adreno(ggml_backend_t backend, const ggml_ten int N = ne1; int K = ne00; - if (nb01 > nb02) { - // KQ - kernel = backend_ctx->kernel_mul_mm_f16_f32_kq; - } else { - // KQV - kernel = backend_ctx->kernel_mul_mm_f16_f32_kqv; - } + kernel = is_kq ? backend_ctx->kernel_mul_mm_f16_f32_kq + : backend_ctx->kernel_mul_mm_f16_f32_kqv; // create sub-buffer for A // <--------------------------------------------> // extra0 = src0->view_src ? (ggml_tensor_extra_cl *)src0->view_src->extra : (ggml_tensor_extra_cl *)src0->extra; region.origin = (extra0->offset + src0->view_offs); - if (nb01 > nb02) { + if (is_kq) { // KQ region.size = nb01 * ne01; } else { @@ -16332,7 +16333,7 @@ static void ggml_cl_mul_mat_kq_kqv_adreno(ggml_backend_t backend, const ggml_ten img_fmt_1d = {CL_RGBA, CL_FLOAT}; memset(&img_desc_1d, 0, sizeof(img_desc_1d)); img_desc_1d.image_type = CL_MEM_OBJECT_IMAGE1D_BUFFER; - if (nb01 > nb02) { + if (is_kq) { img_desc_1d.image_width = (nb01 * ne01 / 4)/4; } else { @@ -19222,13 +19223,61 @@ static void ggml_cl_mul_mat(ggml_backend_t backend, const ggml_tensor * src0, co #ifdef GGML_OPENCL_USE_ADRENO_KERNELS if(src0t == GGML_TYPE_F16 && src1t == GGML_TYPE_F32){ - if (ne01 >= 64 && ne1 >= 32 && ne00 >= 16 && (ne12 % ne02) == 0 && + // Two tiling assumptions these kernels make but nothing enforced: + // + // ne00 % TILESIZE_K(16): the K loop has no tail, so a K that does not + // divide folds 1-15 rows of whatever follows the operands into every + // output. + // + // ne01 % TILESIZE_M(64): mm_store_c_N guards the n direction with its + // `mask` argument but nothing guards m -- the store walks all 64 rows + // of the tile at a stride of M. When M does not divide, the last tile + // does not run off the end of the buffer, it writes 64 - (M % 64) + // values ON TOP OF the next column, so the result is silently wrong. + // Reachable on the KQV side for any head size >= 64 that is not a + // multiple of it (80, 96, 112). + // + // Attention shapes in the graph satisfy both -- head sizes are multiples + // of 64 and n_kv is padded -- which is why this has stayed latent. + // Declining leaves the odd shapes on the generic GEMM, which handles them. + if (ne01 >= 64 && ne1 >= 32 && ne00 >= 16 && + (ne00 % 16) == 0 && (ne01 % 64) == 0 && (ne12 % ne02) == 0 && // the KQ/KQV image kernels do not handle dim 3 (multi-stream batches) ne03 == 1 && ne13 == 1 && // dst is wrapped with image1d_buffer, the size limit applies, also src0 (ne0 * ne1 * dst->ne[2] * dst->nb[0] / 4 <= backend_ctx->image_max_buffer_size)) { - // For KQ - if (ggml_is_permuted(src0) && ggml_is_permuted(src1) && + // For KQ. + // + // Layout admission, mirroring the KQV arm below. The KQ kernel takes + // no stride arguments for A or B: it derives them as K*D_A*2 and + // K*D_B*4, i.e. it assumes both operands pack exactly D heads of K + // elements per row. Every real KV-cache view and permuted-Q view + // does, but a view spanning part of a wider allocation does not, and + // the kernel then walks the wrong rows with nothing to range-check + // it. Gate on the packed layout itself rather than on the stride + // ORDERING, which a wider parent satisfies just as well. + const bool kq_packed_a = (nb01 == (cl_ulong)ne00 * ne02 * ggml_type_size(src0t)) && + (nb02 == (cl_ulong)ne00 * ggml_type_size(src0t)); + const bool kq_packed_b = (nb11 == (cl_ulong)ne10 * ne12 * ggml_type_size(src1t)) && + (nb12 == (cl_ulong)ne10 * ggml_type_size(src1t)); + // + // ggml_is_permuted(src0) stands in for "K is head-major", but it is + // only a proxy and it COLLAPSES at n_head_kv == 1: with a single + // head there is no head stride to be out of order, so nb01 == nb02 + // and the view reports itself unpermuted. Such a KQ was declined + // here and fell through to the generic GEMM (gemma-4 E2B, and any + // other multi-query model). The packed check above is the contract + // the kernel actually needs -- it pins both strides exactly -- so + // require permutedness only where there is more than one head for + // it to mean anything. + // + // Default on; GGML_OPENCL_KQ_NHEAD_KV1=0 restores the old proxy so + // the two routings can be compared in one binary. + static const char * kq_nhkv1_env = getenv("GGML_OPENCL_KQ_NHEAD_KV1"); + static const bool kq_nhkv1_on = + (kq_nhkv1_env == nullptr || kq_nhkv1_env[0] != '0'); + if ((ggml_is_permuted(src0) || (ne02 == 1 && kq_nhkv1_on)) && ggml_is_permuted(src1) && + kq_packed_a && kq_packed_b && ((nb01 * ne01 / 4)/4 <= backend_ctx->image_max_buffer_size) && nb00 <= nb02 && nb02 <= nb01 && @@ -19236,13 +19285,15 @@ static void ggml_cl_mul_mat(ggml_backend_t backend, const ggml_tensor * src0, co nb10 <= nb12 && nb12 <= nb11 && nb11 <= nb13) { - ggml_cl_mul_mat_kq_kqv_adreno(backend, src0, src1, dst); + ggml_cl_mul_mat_kq_kqv_adreno(backend, src0, src1, dst, /*is_kq =*/ true); return; } - // For KQV + // For KQV. Reaching this arm is what makes the op a KQV; the callee + // is told so explicitly rather than re-deriving it from the strides + // the arm above has already ruled on. if (!ggml_is_contiguous(src0) && ggml_is_contiguous(src1) && ((nb02 * ne02 / 4)/4 <= backend_ctx->image_max_buffer_size)) { - ggml_cl_mul_mat_kq_kqv_adreno(backend, src0, src1, dst); + ggml_cl_mul_mat_kq_kqv_adreno(backend, src0, src1, dst, /*is_kq =*/ false); return; } } diff --git a/ggml/src/ggml-opencl/kernels/gemv_noshuffle_q4_k_f32.cl b/ggml/src/ggml-opencl/kernels/gemv_noshuffle_q4_k_f32.cl index c1829fc3820..9ab0dee693e 100644 --- a/ggml/src/ggml-opencl/kernels/gemv_noshuffle_q4_k_f32.cl +++ b/ggml/src/ggml-opencl/kernels/gemv_noshuffle_q4_k_f32.cl @@ -235,6 +235,23 @@ kernel void kernel_gemv_noshuffle_q4_k_f32( uint LINE_STRIDE_A = M / 2; uint BLOCK_STRIDE_A = NSUBGROUPS * M; + // The x-grid is padded to CEIL_DIV(ne01/2,64)*64, so when ne01 % 128 != 0 the + // tail lanes hold gid >= ne01/2. The output stores below are guarded, but the + // input fetches are not: src0_d and src0_m are raw global half2 pointers, + // src0_s is a raw global uchar pointer, and read_imageui on an + // image1d_buffer_t is UNDEFINED out of range -- an image clamps only for + // SAMPLER reads, which these are not. Those lanes therefore read past the end + // of all three allocations. For a [2816, 2112] weight (2112 % 128 == 64) the + // top tail lane is gid = 1087 while only gid < 1056 is backed, and it runs + // 32 half2 past src0_d/src0_m, 31 uints past the quant image, and 63 bytes + // past src0_s. + // + // Clamp the row used for every fetch. The lanes stay ACTIVE, which the + // sub_group_broadcast in the dequant macros requires, and their results are + // still discarded by the existing output guard. No-op and byte-identical + // whenever ne01 % 128 == 0. + uint gid_s = min(gid, LINE_STRIDE_A - 1); + private uint4 regA; private half2 regS; private half2 regM; @@ -246,10 +263,10 @@ kernel void kernel_gemv_noshuffle_q4_k_f32( uint sb = k / 8; uint j = k % 8; - half2 d = src0_d[gid + sb * LINE_STRIDE_A]; - half2 dm = src0_m[gid + sb * LINE_STRIDE_A]; + half2 d = src0_d[gid_s + sb * LINE_STRIDE_A]; + half2 dm = src0_m[gid_s + sb * LINE_STRIDE_A]; - global const uchar * sc0 = src0_s + sb * 12 * M + 2 * gid; + global const uchar * sc0 = src0_s + sb * 12 * M + 2 * gid_s; global const uchar * sc1 = sc0 + 1; uchar sv0, mn0, sv1, mn1; @@ -265,20 +282,20 @@ kernel void kernel_gemv_noshuffle_q4_k_f32( } // load half weights for two blocks in consecutive rows - regA.s0 = read_imageui(src0_q, (gid + k * BLOCK_STRIDE_A + LINE_STRIDE_A * 0)).x; - regA.s1 = read_imageui(src0_q, (gid + k * BLOCK_STRIDE_A + LINE_STRIDE_A * 1)).x; - regA.s2 = read_imageui(src0_q, (gid + k * BLOCK_STRIDE_A + LINE_STRIDE_A * 2)).x; - regA.s3 = read_imageui(src0_q, (gid + k * BLOCK_STRIDE_A + LINE_STRIDE_A * 3)).x; + regA.s0 = read_imageui(src0_q, (gid_s + k * BLOCK_STRIDE_A + LINE_STRIDE_A * 0)).x; + regA.s1 = read_imageui(src0_q, (gid_s + k * BLOCK_STRIDE_A + LINE_STRIDE_A * 1)).x; + regA.s2 = read_imageui(src0_q, (gid_s + k * BLOCK_STRIDE_A + LINE_STRIDE_A * 2)).x; + regA.s3 = read_imageui(src0_q, (gid_s + k * BLOCK_STRIDE_A + LINE_STRIDE_A * 3)).x; #ifdef VECTOR_SUB_GROUP_BROADCAST dequantizeBlockAccum_ns_sgbroadcast_8_hi(totalSum, as_ushort8(regA), regS, regM, regB); #else dequantizeBlockAccum_ns_sgbroadcast_1_hi(totalSum, as_ushort8(regA), regS, regM, regB); #endif // VECTOR_SUB_GROUP_BROADCAST - regA.s0 = read_imageui(src0_q, (gid + k * BLOCK_STRIDE_A + LINE_STRIDE_A * 4)).x; - regA.s1 = read_imageui(src0_q, (gid + k * BLOCK_STRIDE_A + LINE_STRIDE_A * 5)).x; - regA.s2 = read_imageui(src0_q, (gid + k * BLOCK_STRIDE_A + LINE_STRIDE_A * 6)).x; - regA.s3 = read_imageui(src0_q, (gid + k * BLOCK_STRIDE_A + LINE_STRIDE_A * 7)).x; + regA.s0 = read_imageui(src0_q, (gid_s + k * BLOCK_STRIDE_A + LINE_STRIDE_A * 4)).x; + regA.s1 = read_imageui(src0_q, (gid_s + k * BLOCK_STRIDE_A + LINE_STRIDE_A * 5)).x; + regA.s2 = read_imageui(src0_q, (gid_s + k * BLOCK_STRIDE_A + LINE_STRIDE_A * 6)).x; + regA.s3 = read_imageui(src0_q, (gid_s + k * BLOCK_STRIDE_A + LINE_STRIDE_A * 7)).x; #ifdef VECTOR_SUB_GROUP_BROADCAST dequantizeBlockAccum_ns_sgbroadcast_8_lo(totalSum, as_ushort8(regA), regS, regM, regB); #else From d57ae98246c503fad649a8d715ff995b1b867ba6 Mon Sep 17 00:00:00 2001 From: Alan Tseng Date: Wed, 2 Sep 2026 14:12:28 +0800 Subject: [PATCH 076/104] ggml-cpu : conditionally add SpacemiT IME kernel sources (llama/27961) When building with gcc < 15, ggml/CMakeLists.txt unconditionally adds ime2_kernels.cpp, which fails to compile. FindSMTIME.cmake only defines RISCV64_SPACEMIT_IME2 when the IME2 instructions are detected, and gcc 14 only has IME1, so ime2_kernels.cpp hits its #error. This PR fixes it by using IN_LIST to add each kernel source according to the spec that was actually detected. --- ggml/src/ggml-cpu/CMakeLists.txt | 8 ++++++-- 1 file changed, 6 insertions(+), 2 deletions(-) diff --git a/ggml/src/ggml-cpu/CMakeLists.txt b/ggml/src/ggml-cpu/CMakeLists.txt index 5442e12501f..17540faa66d 100644 --- a/ggml/src/ggml-cpu/CMakeLists.txt +++ b/ggml/src/ggml-cpu/CMakeLists.txt @@ -455,12 +455,16 @@ function(ggml_add_cpu_backend_variant_impl tag_name) ggml-cpu/spacemit/repack.h ggml-cpu/spacemit/ime_env.cpp ggml-cpu/spacemit/ime_env.h - ggml-cpu/spacemit/ime1_kernels.cpp - ggml-cpu/spacemit/ime2_kernels.cpp ggml-cpu/spacemit/ime_kernels.h ggml-cpu/spacemit/rvv_kernels.cpp ggml-cpu/spacemit/rvv_kernels.h ) + if ("RISCV64_SPACEMIT_IME1" IN_LIST RISCV64_SPACEMIT_IME_SPEC) + list(APPEND GGML_CPU_SOURCES ggml-cpu/spacemit/ime1_kernels.cpp) + endif() + if ("RISCV64_SPACEMIT_IME2" IN_LIST RISCV64_SPACEMIT_IME_SPEC) + list(APPEND GGML_CPU_SOURCES ggml-cpu/spacemit/ime2_kernels.cpp) + endif() endif() if(NOT GGML_CPU_ALL_VARIANTS) set(MARCH_STR "rv64gc") From a9e58612c5145ba47b101889b3f08386a7f42f23 Mon Sep 17 00:00:00 2001 From: Mads Marquart Date: Wed, 2 Sep 2026 08:13:25 +0200 Subject: [PATCH 077/104] vulkan : only request VK_KHR_shader_bfloat16 extension if supported (llama/28155) --- ggml/src/ggml-vulkan/ggml-vulkan.cpp | 14 ++++---------- 1 file changed, 4 insertions(+), 10 deletions(-) diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp index 8718bd2cfb6..285cd652537 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp +++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp @@ -6897,7 +6897,8 @@ static vk_device ggml_vk_get_device(size_t idx) { } #if defined(VK_KHR_shader_bfloat16) && defined(GGML_VULKAN_BFLOAT16_GLSLC_SUPPORT) - if (prop.AType == VK_COMPONENT_TYPE_BFLOAT16_KHR && + if (bfloat16_support && + prop.AType == VK_COMPONENT_TYPE_BFLOAT16_KHR && prop.BType == VK_COMPONENT_TYPE_BFLOAT16_KHR && prop.CType == VK_COMPONENT_TYPE_FLOAT32_KHR && prop.ResultType == VK_COMPONENT_TYPE_FLOAT32_KHR) { @@ -7017,7 +7018,8 @@ static vk_device ggml_vk_get_device(size_t idx) { device->coopmat_int_k = prop.KSize; } #if defined(VK_KHR_shader_bfloat16) && defined(GGML_VULKAN_BFLOAT16_GLSLC_SUPPORT) - if (prop.AType == VK_COMPONENT_TYPE_BFLOAT16_KHR && + if (bfloat16_support && + prop.AType == VK_COMPONENT_TYPE_BFLOAT16_KHR && prop.BType == VK_COMPONENT_TYPE_BFLOAT16_KHR && prop.CType == VK_COMPONENT_TYPE_FLOAT32_KHR && prop.ResultType == VK_COMPONENT_TYPE_FLOAT32_KHR && @@ -7042,19 +7044,11 @@ static vk_device ggml_vk_get_device(size_t idx) { GGML_LOG_DEBUG("ggml_vulkan: WARNING: No suitable matrix core mode found. Disabling matrix cores.\n"); device->coopmat_support = false; } - if (getenv("GGML_VK_DISABLE_BFLOAT16")) { - device->coopmat_bf16_support = false; - } } if (device->coopmat_support) { device_extensions.push_back("VK_KHR_cooperative_matrix"); } -#if defined(VK_KHR_shader_bfloat16) - if (device->coopmat_bf16_support) { - device_extensions.push_back("VK_KHR_shader_bfloat16"); - } -#endif #endif device->name = GGML_VK_NAME + std::to_string(idx); From 1c7d35e14af94778e40f92afb727fd9eb159783b Mon Sep 17 00:00:00 2001 From: Laurent Zuijdwijk Date: Wed, 2 Sep 2026 07:14:52 +0100 Subject: [PATCH 078/104] vulkan: handle larger batch sizes (>4) efficiently for IQ3_S mat-vec (llama/27449) * vulkan: handle larger batch sizes (>4) efficiently for IQ3_S mat-vec when NUM_COLS > 4. 5x perf at n=8 Assisted-by: Claude Opus 5 * adds 2 cases per quant type at `k=16*256` to the `all_types` mat-vec sweep --------- Co-authored-by: Marshall --- .../vulkan-shaders/mul_mat_vec_iq3_s.comp | 33 ++++++++++--------- 1 file changed, 18 insertions(+), 15 deletions(-) diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mat_vec_iq3_s.comp b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mat_vec_iq3_s.comp index 5cdf2a89d0f..42f52b4a127 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mat_vec_iq3_s.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mat_vec_iq3_s.comp @@ -7,7 +7,14 @@ layout(local_size_x_id = 0, local_size_y = 1, local_size_z = 1) in; FLOAT_TYPE temp[NUM_COLS][NUM_ROWS]; -void calc_superblock(const uint a_offset, const uint b_offset, const uint ib32, const uint i, const uint num_blocks_per_row, const uint first_row, const uint num_rows) { +// invocations per superblock. with many columns, 8 invocations need too many +// registers and spill, so use 16 to halve the per-invocation B working set +const uint TPB = NUM_COLS <= 4 ? 8 : 16; +const uint NL = 32 / TPB; // l steps per invocation + +void calc_superblock(const uint a_offset, const uint b_offset, const uint itid, const uint i, const uint num_blocks_per_row, const uint first_row, const uint num_rows) { + const uint ib32 = itid / (TPB / 8); + const uint l0 = (itid % (TPB / 8)) * NL; const uint y_idx = i * QUANT_K + 32 * ib32; uint ibi = a_offset + first_row * num_blocks_per_row + i; @@ -16,11 +23,8 @@ void calc_superblock(const uint a_offset, const uint b_offset, const uint ib32, const uint scale = (data_a[ibi].scales[ib32/2] >> (4 * (ib32 & 1))) & 0xF; const float dscale = d * (1 + 2 * scale); const uint qh = data_a[ibi].qh[ib32]; - FLOAT_TYPE sum[NUM_COLS]; - [[unroll]] for (uint j = 0; j < NUM_COLS; ++j) { - sum[j] = 0.0; - } - [[unroll]] for (uint l = 0; l < 4; ++l) { + [[unroll]] for (uint ll = 0; ll < NL; ++ll) { + const uint l = l0 + ll; const u8vec2 qs = unpack8(uint32_t(data_a_packed16[ibi].qs[4 * ib32 + l])).xy; // vec4 used due to #12147 const uint sign = data_a[ibi].signs[4 * ib32 + l]; const vec4 grid0 = vec4(unpack8(iq3s_grid[qs.x | ((qh << (8 - 2*l)) & 0x100)])); @@ -30,7 +34,7 @@ void calc_superblock(const uint a_offset, const uint b_offset, const uint ib32, const vec4 b0 = vec4(data_b_v4[(j*p.batch_stride_b + b_offset + y_idx) / 4 + 2*l + 0]); const vec4 b4 = vec4(data_b_v4[(j*p.batch_stride_b + b_offset + y_idx) / 4 + 2*l + 1]); - sum[j] = + const FLOAT_TYPE sum = fma(FLOAT_TYPE(b0.x), FLOAT_TYPE((sign & 1) != 0 ? -grid0.x : grid0.x), fma(FLOAT_TYPE(b0.y), FLOAT_TYPE((sign & 2) != 0 ? -grid0.y : grid0.y), fma(FLOAT_TYPE(b0.z), FLOAT_TYPE((sign & 4) != 0 ? -grid0.z : grid0.z), @@ -39,12 +43,11 @@ void calc_superblock(const uint a_offset, const uint b_offset, const uint ib32, fma(FLOAT_TYPE(b4.y), FLOAT_TYPE((sign & 32) != 0 ? -grid1.y : grid1.y), fma(FLOAT_TYPE(b4.z), FLOAT_TYPE((sign & 64) != 0 ? -grid1.z : grid1.z), fma(FLOAT_TYPE(b4.w), FLOAT_TYPE((sign & 128) != 0 ? -grid1.w : grid1.w), - sum[j])))))))); + FLOAT_TYPE(0.0))))))))); + + temp[j][n] = fma(dscale, sum, temp[j][n]); } } - [[unroll]] for (uint j = 0; j < NUM_COLS; ++j) { - temp[j][n] = fma(dscale, sum[j], temp[j][n]); - } ibi += num_blocks_per_row; } } @@ -55,11 +58,11 @@ void compute_outputs(const uint32_t first_row, const uint32_t num_rows) { const uint num_blocks_per_row = p.ncols / QUANT_K; - // 8 threads are used to process each block - const uint blocks_per_wg = gl_WorkGroupSize.x/8; + // TPB invocations are used to process each block + const uint blocks_per_wg = gl_WorkGroupSize.x/TPB; const uint tid = gl_LocalInvocationID.x; - const uint itid = tid % 8; // 0...7 - const uint ix = tid / 8; + const uint itid = tid % TPB; + const uint ix = tid / TPB; [[unroll]] for (uint j = 0; j < NUM_COLS; ++j) { [[unroll]] for (uint i = 0; i < NUM_ROWS; ++i) { From dc70853ec28bdfc716641111f37cee0fe4bd8a91 Mon Sep 17 00:00:00 2001 From: Max Krasnyansky Date: Tue, 1 Sep 2026 23:15:21 -0700 Subject: [PATCH 079/104] hexagon: MUL_MAT and MUL_MAT_ID fusion and fixes (llama/28202) * hex-mm: fuse QKV and FFN matmuls that land on HMX * hex-mm: remove hardcoded ne[1] < 32K restriction * hex-get-rows: explicitly reject repacked Q8_0 just in case somebody decided to add an override * hex-mm: correct overhead sizing to make sure we dont exceed vtcm budget for large dims * hex-mm: fuse MUL_MAT_ID into MUL_MAT_ID_NX (2x,3x,...) where possible * hex-fusion: update opbatch and opqueue sizing to acount for new fusion and reduce overhead for trace buffer alloc * hex-bufs: sort buffers while finalizing opbatch, helps avoid va space fragmentation * hex-bufs: add simple va defrag to make sure we dont abort just because the va space is fragmented * hex-mm: replaced more scalar divs with fastdiv and minor cleanup * hex-mm: tighten up supported fusion checks to exactly match supported kernels --- ggml/src/ggml-hexagon/ggml-hexagon.cpp | 498 ++++++++++++--- ggml/src/ggml-hexagon/htp-opnode.h | 3 +- ggml/src/ggml-hexagon/htp/htp-ctx.h | 1 + ggml/src/ggml-hexagon/htp/htp-ops.h | 1 + ggml/src/ggml-hexagon/htp/main.c | 39 +- ggml/src/ggml-hexagon/htp/matmul-ops.c | 846 ++++++++++++++++++++++--- ggml/src/ggml-hexagon/htp/matmul-ops.h | 28 +- 7 files changed, 1236 insertions(+), 180 deletions(-) diff --git a/ggml/src/ggml-hexagon/ggml-hexagon.cpp b/ggml/src/ggml-hexagon/ggml-hexagon.cpp index 3eb84fd2a99..04fb9a22339 100644 --- a/ggml/src/ggml-hexagon/ggml-hexagon.cpp +++ b/ggml/src/ggml-hexagon/ggml-hexagon.cpp @@ -98,12 +98,26 @@ static int opt_ar_select = 2; // 2 = fused ALLREDUCE+ADD (DMA, default), 1 = // https://docs.qualcomm.com/doc/80-N2040-61/topic/hvx-pmu-events.html static u32vec opt_pmu_evt { 0x3, 0x111, 0x100, 0x105, 0x240, 0x256, 0x7D, 0x8C }; -static int opt_opbatch = 1024; // max number of ops in a batch -static int opt_opqueue = 64; // max number of pending batches +static int opt_opbatch = 1280; // max number of ops in a batch +static int opt_opqueue = 32; // max number of pending batches static int opt_optrace = 0; // trace buffer size per thread (0 means default) static int opt_oppoll = 0; // polling for batch completions static int opt_opfusion = 1; // enable/disable op fusion +enum ggml_hexagon_fusion_flags { + GGML_HEXAGON_FUSE_ALLREDUCE_ADD = (1 << 1), // 2 + GGML_HEXAGON_FUSE_RMS_NORM_MUL = (1 << 2), // 4 + GGML_HEXAGON_FUSE_MUL_MAT_ADD = (1 << 3), // 8 + GGML_HEXAGON_FUSE_MUL_MAT_NX = (1 << 4), // 16 + GGML_HEXAGON_FUSE_MUL_MAT_ID_NX = (1 << 5), // 32 +}; + +static inline bool ggml_hexagon_is_fusion_enabled(int flag) { + if (opt_opfusion <= 0) return false; + if (opt_opfusion == 1) return true; // 1 enables all + return (opt_opfusion & flag) != 0; +} + static std::regex* opt_opfilter = NULL; // regex of ops to not claim #define HEX_VERBOSE(...) \ @@ -293,6 +307,15 @@ static void ggml_hexagon_precompute_fused_mmnx_params( struct htp_mm_kernel_params * kparams ); +static void ggml_hexagon_precompute_fused_mmidnx_params( + const struct ggml_hexagon_session * sess, + const struct ggml_tensor * src0, + const struct ggml_tensor * src1, + const struct ggml_tensor * dst, + int32_t n_weights, + struct htp_mm_kernel_params * kparams +); + static bool ggml_hexagon_precompute_allreduce_params( const struct ggml_hexagon_session * sess, const struct ggml_tensor * dst, @@ -304,8 +327,12 @@ static bool ggml_hexagon_precompute_allreduce_params( ); static bool mm_is_hmx_eligible(const ggml_tensor * t); +static bool is_supported_mul_mat_nx_kernel(const ggml_tensor * src0, const struct htp_mm_kernel_params * kparams); +static bool is_supported_mul_mat_id_nx_kernel(const ggml_tensor * src0, const struct htp_mm_kernel_params * kparams); static bool is_mergeable_mul_mat(const ggml_tensor * t); static bool is_mergeable_mul_mat_pair(const ggml_tensor * n1, const ggml_tensor * n2); +static bool is_mergeable_mul_mat_id(const ggml_tensor * t); +static bool is_mergeable_mul_mat_id_pair(const ggml_tensor * n1, const ggml_tensor * n2); // ** backend sessions @@ -1832,6 +1859,42 @@ struct ggml_hexagon_opbatch { } } + void sort_buffers() { + if (n_bufs <= 1) return; + + std::vector order(n_bufs); + for (unsigned int i = 0; i < n_bufs; i++) { order[i] = (int) i; } + + std::stable_sort(order.begin(), order.end(), [&](int a, int b) { + return h_bufs[a].size > h_bufs[b].size; + }); + + bool already_sorted = true; + for (unsigned int i = 0; i < n_bufs; i++) { + if (order[i] != (int) i) { + already_sorted = false; + break; + } + } + if (already_sorted) return; + + std::vector remap(n_bufs); + std::vector sorted_bufs(n_bufs); + for (unsigned int new_bi = 0; new_bi < n_bufs; new_bi++) { + int old_bi = order[new_bi]; + remap[old_bi] = (uint16_t) new_bi; + sorted_bufs[new_bi] = h_bufs[old_bi]; + } + + for (unsigned int i = 0; i < n_bufs; i++) { + h_bufs[i] = sorted_bufs[i]; + } + + for (unsigned int i = 0; i < n_tens; i++) { + h_tens[i].bi = remap[h_tens[i].bi]; + } + } + bool try_fuse_allreduce_add(const htp_opnode & node) { if (n_ops == 0 || opt_ar_select != 2) return false; if (node.opcode != HTP_OP_ADD) return false; @@ -2144,9 +2207,15 @@ struct ggml_hexagon_opbatch { if (x_in != x || w_in->type != w0->type || w_in->ne[0] != w0->ne[0]) { return false; } + if (!last_node.fused.empty() && (mm_is_hmx_eligible(last_node.fused[0]) != mm_is_hmx_eligible(node.node))) { + return false; + } struct htp_mm_kernel_params kparams; ggml_hexagon_precompute_fused_mmnx_params(sess, w0, x, curr_n + 1, &kparams); + if (!is_supported_mul_mat_nx_kernel(w0, &kparams)) { + return false; + } if ((size_t) kparams.vtcm_size > sess->vtcm_size) { HEX_VERBOSE("ggml-hex: %s skip NX fusion: VTCM needed (%d) > budget (%zu)\n", sess->c_name(), kparams.vtcm_size, sess->vtcm_size); @@ -2210,6 +2279,9 @@ struct ggml_hexagon_opbatch { struct htp_mm_kernel_params kparams; ggml_hexagon_precompute_fused_mmnx_params(sess, w0, x, 2, &kparams); + if (!is_supported_mul_mat_nx_kernel(w0, &kparams)) { + return false; + } if ((size_t) kparams.vtcm_size > sess->vtcm_size) { HEX_VERBOSE("ggml-hex: %s skip NX fusion: VTCM needed (%d) > budget (%zu)\n", sess->c_name(), kparams.vtcm_size, sess->vtcm_size); @@ -2272,18 +2344,172 @@ struct ggml_hexagon_opbatch { return false; } -enum ggml_hexagon_fusion_flags { - GGML_HEXAGON_FUSE_ALLREDUCE_ADD = (1 << 1), // 2 - GGML_HEXAGON_FUSE_RMS_NORM_MUL = (1 << 2), // 4 - GGML_HEXAGON_FUSE_MUL_MAT_ADD = (1 << 3), // 8 - GGML_HEXAGON_FUSE_MUL_MAT_NX = (1 << 4), // 16 -}; + bool try_fuse_mul_mat_id_nx(const htp_opnode & node) { + if (n_ops == 0 || node.opcode != HTP_OP_MUL_MAT_ID) return false; + if (!is_mergeable_mul_mat_id(node.node)) return false; -static inline bool ggml_hexagon_is_fusion_enabled(int flag) { - if (opt_opfusion <= 0) return false; - if (opt_opfusion == 1) return true; // 1 enables all - return (opt_opfusion & flag) != 0; -} + const ggml_tensor * w_in = node.src0(); + const ggml_tensor * x_in = node.src1(); + const ggml_tensor * ids_in = node.node->src[2]; + const ggml_tensor * d_in = node.dst(); + if (!w_in || !x_in || !ids_in || !d_in) return false; + + htp_opnode & last_node = ops[n_ops - 1]; + + // Case 1: last_node is already MUL_MAT_ID_NX + if (last_node.opcode == HTP_OP_MUL_MAT_ID_NX) { + const uint32_t curr_n = (uint32_t) last_node.outputs.size(); + if (curr_n >= HTP_OP_MAX_OUTPUTS || curr_n + 2 >= HTP_OP_MAX_INPUTS) { + return false; + } + + const ggml_tensor * w0 = last_node.inputs[0]; + const ggml_tensor * x = last_node.inputs[curr_n]; + const ggml_tensor * ids = last_node.inputs[curr_n + 1]; + + if (x_in != x || ids_in != ids || w_in->type != w0->type || w_in->ne[0] != w0->ne[0] || w_in->ne[2] != w0->ne[2]) { + return false; + } + if (!last_node.fused.empty() && (mm_is_hmx_eligible(last_node.fused[0]) != mm_is_hmx_eligible(node.node))) { + return false; + } + + struct htp_mm_kernel_params kparams; + ggml_hexagon_precompute_fused_mmidnx_params(sess, w0, x, d_in, curr_n + 1, &kparams); + if (!is_supported_mul_mat_id_nx_kernel(w0, &kparams)) { + return false; + } + if ((size_t) kparams.vtcm_size > sess->vtcm_size) { + HEX_VERBOSE("ggml-hex: %s skip ID NX fusion: VTCM needed (%d) > budget (%zu)\n", + sess->c_name(), kparams.vtcm_size, sess->vtcm_size); + return false; + } + + size_t extra_bufs = 0, extra_vmem = 0, extra_tens = 0; + auto fit_t = [&](const ggml_tensor * t) { + if (!t) return; + if (!t_map.count(t)) { + extra_tens++; + auto sbuf = static_cast(t->buffer->context); + if (!b_map.count(sbuf->fd())) { + extra_vmem += sbuf->size(); + extra_bufs += 1; + } + } + }; + fit_t(w_in); + fit_t(d_in); + if ((extra_bufs + n_bufs) > n_bufs_max || (extra_tens + n_tens) > n_tens_max || (extra_vmem + b_vmem) > b_vmem_max) { + return false; + } + + last_node.inputs[curr_n] = w_in; + last_node.inputs[curr_n + 1] = x; + last_node.inputs.push_back(ids); + last_node.outputs.push_back(d_in); + last_node.fused.push_back(node.node); + memcpy(last_node.kernel_params, &kparams, sizeof(kparams)); + + htp_op_desc & o = h_ops[n_ops - 1]; + memcpy(o.kernel_params, &kparams, sizeof(kparams)); + + for (uint32_t s = 0; s <= curr_n + 2; s++) { + o.src[s] = add_tensor(last_node.inputs[s]); + } + for (uint32_t s = curr_n + 3; s < HTP_OP_MAX_INPUTS; s++) { + o.src[s] = 0xffff; + } + for (uint32_t d = 0; d <= curr_n; d++) { + o.dst[d] = add_tensor(last_node.outputs[d]); + } + for (uint32_t d = curr_n + 1; d < HTP_OP_MAX_OUTPUTS; d++) { + o.dst[d] = 0xffff; + } + + HEX_VERBOSE("ggml-hex: %s fused MUL_MAT_ID_NX (N=%u, #%u)\n", sess->c_name(), curr_n + 1, n_ops - 1); + return true; + } + + // Case 2: last_node is single MUL_MAT_ID + if (last_node.opcode == HTP_OP_MUL_MAT_ID) { + if (!is_mergeable_mul_mat_id_pair(last_node.node, node.node)) { + return false; + } + + const ggml_tensor * w0 = last_node.src0(); + const ggml_tensor * x = last_node.src1(); + const ggml_tensor * ids = last_node.node->src[2]; + const ggml_tensor * w1 = node.src0(); + if (!w0 || !x || !ids || !w1) return false; + + struct htp_mm_kernel_params kparams; + ggml_hexagon_precompute_fused_mmidnx_params(sess, w0, x, node.dst(), 2, &kparams); + if (!is_supported_mul_mat_id_nx_kernel(w0, &kparams)) { + return false; + } + if ((size_t) kparams.vtcm_size > sess->vtcm_size) { + HEX_VERBOSE("ggml-hex: %s skip ID NX fusion: VTCM needed (%d) > budget (%zu)\n", + sess->c_name(), kparams.vtcm_size, sess->vtcm_size); + return false; + } + + size_t extra_bufs = 0, extra_vmem = 0, extra_tens = 0; + auto fit_t = [&](const ggml_tensor * t) { + if (!t) return; + if (!t_map.count(t)) { + extra_tens++; + auto sbuf = static_cast(t->buffer->context); + if (!b_map.count(sbuf->fd())) { + extra_vmem += sbuf->size(); + extra_bufs += 1; + } + } + }; + fit_t(w1); + fit_t(node.dst()); + if ((extra_bufs + n_bufs) > n_bufs_max || (extra_tens + n_tens) > n_tens_max || (extra_vmem + b_vmem) > b_vmem_max) { + return false; + } + + const ggml_tensor * dst_0 = last_node.dst(); + const ggml_tensor * dst_1 = node.dst(); + + last_node.opcode = HTP_OP_MUL_MAT_ID_NX; + last_node.name = "MUL_MAT_ID_NX"; + last_node.inputs.clear(); + last_node.inputs.push_back(w0); + last_node.inputs.push_back(w1); + last_node.inputs.push_back(x); + last_node.inputs.push_back(ids); + last_node.outputs.clear(); + last_node.outputs.push_back(dst_0); + last_node.outputs.push_back(dst_1); + last_node.fused.push_back(node.node); + memcpy(last_node.kernel_params, &kparams, sizeof(kparams)); + + htp_op_desc & o = h_ops[n_ops - 1]; + o.opcode = HTP_OP_MUL_MAT_ID_NX; + memcpy(o.kernel_params, &kparams, sizeof(kparams)); + + o.src[0] = add_tensor(w0); + o.src[1] = add_tensor(w1); + o.src[2] = add_tensor(x); + o.src[3] = add_tensor(ids); + for (uint32_t s = 4; s < HTP_OP_MAX_INPUTS; s++) { + o.src[s] = 0xffff; + } + o.dst[0] = add_tensor(dst_0); + o.dst[1] = add_tensor(dst_1); + for (uint32_t d = 2; d < HTP_OP_MAX_OUTPUTS; d++) { + o.dst[d] = 0xffff; + } + + HEX_VERBOSE("ggml-hex: %s fused MUL_MAT_ID_NX (N=2, #%u)\n", sess->c_name(), n_ops - 1); + return true; + } + + return false; + } bool try_fuse(const htp_opnode & node) { if (!opt_opfusion) return false; @@ -2291,6 +2517,7 @@ static inline bool ggml_hexagon_is_fusion_enabled(int flag) { if (ggml_hexagon_is_fusion_enabled(GGML_HEXAGON_FUSE_RMS_NORM_MUL) && try_fuse_rms_norm_mul(node)) return true; if (ggml_hexagon_is_fusion_enabled(GGML_HEXAGON_FUSE_MUL_MAT_ADD) && try_fuse_mul_mat_add(node)) return true; if (ggml_hexagon_is_fusion_enabled(GGML_HEXAGON_FUSE_MUL_MAT_NX) && try_fuse_mul_mat_nx(node)) return true; + if (ggml_hexagon_is_fusion_enabled(GGML_HEXAGON_FUSE_MUL_MAT_ID_NX) && try_fuse_mul_mat_id_nx(node)) return true; return false; } }; @@ -2350,6 +2577,8 @@ struct ggml_hexagon_opqueue { delete shm_buf; } + size_t shm_size() const { return shm_buf ? shm_buf->size() : 0; } + // push new batch bool push(htp_opbatch_req& req, dspqueue_buffer& dbuf, ggml_hexagon_opbatch* op_batch) { static_assert(sizeof(htp_opbatch_req) % 8 == 0, "sizeof(htp_opbatch_req) must be multiple of 8"); @@ -2396,6 +2625,8 @@ struct ggml_hexagon_opqueue { uint8_t * t_ptr = m_ptr; m_ptr += t_size; uint8_t * o_ptr = m_ptr; + op_batch->sort_buffers(); + memcpy(b_ptr, (void *) op_batch->h_bufs.data(), b_size); memcpy(t_ptr, (void *) op_batch->h_tens.data(), t_size); memcpy(o_ptr, (void *) op_batch->h_ops.data(), o_size); @@ -3018,7 +3249,8 @@ void ggml_hexagon_session::allocate(const ggml_hexagon_device_config & config) n opt_vmem = ggml_hexagon_measure_max_vmem(this); GGML_LOG_INFO("ggml-hex: %s measured max vmem %zu\n", this->c_name(), opt_vmem); } - this->max_vmem = opt_vmem; + const size_t shm_size = this->op_queue->shm_size(); + this->max_vmem = (opt_vmem > shm_size) ? (opt_vmem - shm_size) : opt_vmem; this->op_batch = new ggml_hexagon_opbatch(this, opt_opbatch, this->max_vmem); @@ -3378,6 +3610,10 @@ static bool ggml_hexagon_matmul_is_hmx_eligible( bool is_matmul_id, bool is_batched ) { + if (src1->type != GGML_TYPE_F32) { + return false; + } + const int ne00 = src0->ne[0]; const int ne11 = src1->ne[1]; const int ne12 = src1->ne[2]; @@ -3408,7 +3644,8 @@ static bool ggml_hexagon_matmul_is_hmx_eligible( return false; } - // M alignment: Use HMX when M > HTP_MM_HMX_MIN_NROWS + // M alignment: Use HMX when M > HTP_MM_HMX_MIN_NROWS. + // For MUL_MAT_ID, src1 shape is [K, n_expert_used, n_tokens, 1], so n_tokens is ne12. const int m = is_matmul_id ? ne12 : ne11; if (m <= HTP_MM_HMX_MIN_NROWS) { return false; @@ -3460,7 +3697,7 @@ static bool ggml_hexagon_precompute_hmx_mm_params( if (!use_grouped) { // Fallback to simple 2D path (group_size = 1) - const int m_id_rows = (int) ((size_t) dst->ne[1] * dst->ne[2]); + const int m_id_rows = (dst && is_matmul_id) ? (int) ((size_t) dst->ne[1] * dst->ne[2]) : 0; if (!htp_mm_hmx_solve_2d_params(wtype, ne00_padded, m_id_rows, ne01_padded, ne11_padded, ne11, n_threads, pipeline, is_matmul_id, aligned_tile_size, vtcm_budget, &m_chunk, &n_chunk, &act_threads_selected, &vtcm_size)) { return false; } @@ -3918,64 +4155,113 @@ static void ggml_hexagon_precompute_fused_mmnx_params( ) { memset(kparams, 0, sizeof(*kparams)); - const int wtype = src0->type; - const bool is_repack = ggml_hexagon_is_repack_type((ggml_type) wtype); + const int ne00 = src0->ne[0]; + const int ne01 = src0->ne[1]; + const int ne02 = src0->ne[2]; + const int ne03 = src0->ne[3]; const int ne10 = src1->ne[0]; - const int src1_nrows = src1->ne[1] * src1->ne[2] * src1->ne[3]; - const size_t src1_row_size = (wtype == GGML_TYPE_Q4_1) ? htp_mm_q8_1_tiled_row_size(ne10) : htp_mm_q8_0_tiled_row_size(ne10); - const size_t src0_row_size = src0->nb[1]; + const int ne11 = src1->ne[1]; + const int ne12 = src1->ne[2]; + const int ne13 = src1->ne[3]; + + const int wtype = src0->type; + const bool is_repack = ggml_hexagon_is_repack_type((ggml_type) wtype); + const int ne00_padded = is_repack ? hex_round_up(ne00, 32) : ne00; + const int ne01_padded = is_repack ? hex_round_up(ne01, 32) : ne01; + const int ne11_padded = hex_round_up(ne11, 32); - uint32_t best_n_prefetch = 16; + const size_t vtcm_budget = sess->vtcm_size; + const bool is_batched = (ne02 * ne03 > 1 || ne12 * ne13 > 1); - if (is_repack) { - const uint32_t max_prefetch = (src1_nrows > HTP_MM_HMX_MIN_NROWS) ? 2 : 16; - best_n_prefetch = 2; - for (uint32_t d = max_prefetch; d >= 2; d /= 2) { - struct htp_mm_hvx_vtcm_layout L; - htp_mm_hvx_vtcm_layout_build( - &L, HTP_MM_KERNEL_HVX_QUANT_ROW, wtype, ne10, src1_nrows, sess->n_threads, - 0, src0_row_size, src1_row_size, 0, d, false, true - ); - if (L.total_bytes <= sess->vtcm_size) { - best_n_prefetch = d; - break; - } + bool hmx_enabled = (sess->n_hmx > 0) && (opt_mm_select >= 3); + if (hmx_enabled && ggml_hexagon_matmul_is_hmx_eligible(src0, src1, nullptr, ne01_padded, false, is_batched)) { + if (ggml_hexagon_precompute_hmx_mm_params(sess, src0, src1, nullptr, wtype, ne00_padded, ne01_padded, ne02, ne11, ne12, ne11_padded, false, is_batched, vtcm_budget, kparams)) { + kparams->n_weights = n_weights; + goto finalize; } } - struct htp_mm_hvx_vtcm_layout L; - bool try_tiled = (opt_mm_select >= 2); + if (!is_repack) { + kparams->kernel_type = HTP_MM_KERNEL_UNSUPPORTED; + return; + } - // Test tiled first - htp_mm_hvx_vtcm_layout_build( - &L, HTP_MM_KERNEL_HVX_QUANT_ROW, wtype, ne10, src1_nrows, sess->n_threads, - 0, src0_row_size, src1_row_size, 0, best_n_prefetch, false, true - ); + { + const int src1_nrows = ne11 * ne12 * ne13; + const size_t src1_row_size = (wtype == GGML_TYPE_Q4_1) ? htp_mm_q8_1_tiled_row_size(ne10) : htp_mm_q8_0_tiled_row_size(ne10); + const size_t src0_row_size = src0->nb[1]; - if (try_tiled && L.total_bytes <= sess->vtcm_size) { - kparams->kernel_type = HTP_MM_KERNEL_HVX_QUANT_ROW; - kparams->vtcm_src0_size = L.src0_bytes; - kparams->vtcm_src1_size = L.src1_bytes; - kparams->vtcm_dst_size = L.dst_bytes; - kparams->vtcm_size = L.total_bytes; - kparams->n_prefetch = best_n_prefetch; - kparams->n_weights = n_weights; - } else { - kparams->kernel_type = HTP_MM_KERNEL_HVX_QUANT_ROW_FLAT; - size_t flat_src1_row_size = (wtype == GGML_TYPE_Q4_1) ? htp_mm_q8_1_flat_row_size(ne10) : htp_mm_q8_0_flat_row_size(ne10); + uint32_t best_n_prefetch = 16; + + if (is_repack) { + const uint32_t max_prefetch = (src1_nrows > HTP_MM_HMX_MIN_NROWS) ? 2 : 16; + best_n_prefetch = 2; + for (uint32_t d = max_prefetch; d >= 2; d /= 2) { + struct htp_mm_hvx_vtcm_layout L; + htp_mm_hvx_vtcm_layout_build( + &L, HTP_MM_KERNEL_HVX_QUANT_ROW, wtype, ne10, src1_nrows, sess->n_threads, + 0, src0_row_size, src1_row_size, 0, d, false, true + ); + if (L.total_bytes <= sess->vtcm_size) { + best_n_prefetch = d; + break; + } + } + } + struct htp_mm_hvx_vtcm_layout L; + bool try_tiled = (opt_mm_select >= 2); + + // Test tiled first htp_mm_hvx_vtcm_layout_build( - &L, HTP_MM_KERNEL_HVX_QUANT_ROW_FLAT, wtype, ne10, src1_nrows, sess->n_threads, - 0, src0_row_size, flat_src1_row_size, 0, best_n_prefetch, false, true + &L, HTP_MM_KERNEL_HVX_QUANT_ROW, wtype, ne10, src1_nrows, sess->n_threads, + 0, src0_row_size, src1_row_size, 0, best_n_prefetch, false, true ); - kparams->vtcm_src0_size = L.src0_bytes; - kparams->vtcm_src1_size = L.src1_bytes; - kparams->vtcm_dst_size = L.dst_bytes; - kparams->vtcm_size = L.total_bytes; - kparams->n_prefetch = best_n_prefetch; - kparams->n_weights = n_weights; + + if (try_tiled && L.total_bytes <= sess->vtcm_size) { + kparams->kernel_type = HTP_MM_KERNEL_HVX_QUANT_ROW; + kparams->vtcm_src0_size = L.src0_bytes; + kparams->vtcm_src1_size = L.src1_bytes; + kparams->vtcm_dst_size = L.dst_bytes; + kparams->vtcm_size = L.total_bytes; + kparams->n_prefetch = best_n_prefetch; + kparams->n_weights = n_weights; + } else { + kparams->kernel_type = HTP_MM_KERNEL_HVX_QUANT_ROW_FLAT; + size_t flat_src1_row_size = (wtype == GGML_TYPE_Q4_1) ? htp_mm_q8_1_flat_row_size(ne10) : htp_mm_q8_0_flat_row_size(ne10); + + htp_mm_hvx_vtcm_layout_build( + &L, HTP_MM_KERNEL_HVX_QUANT_ROW_FLAT, wtype, ne10, src1_nrows, sess->n_threads, + 0, src0_row_size, flat_src1_row_size, 0, best_n_prefetch, false, true + ); + kparams->vtcm_src0_size = L.src0_bytes; + kparams->vtcm_src1_size = L.src1_bytes; + kparams->vtcm_dst_size = L.dst_bytes; + kparams->vtcm_size = L.total_bytes; + kparams->n_prefetch = best_n_prefetch; + kparams->n_weights = n_weights; + } } + +finalize: + kparams->div_ne12_ne1 = init_fastdiv_values(ne12 * ne11); + kparams->div_ne1 = init_fastdiv_values(ne11); + kparams->div_r2 = init_fastdiv_values(ne02 > 0 ? ne12 / ne02 : 1); + kparams->div_r3 = init_fastdiv_values(ne03 > 0 ? ne13 / ne03 : 1); + kparams->div_ne11 = init_fastdiv_values(ne11); +} + +static void ggml_hexagon_precompute_fused_mmidnx_params( + const struct ggml_hexagon_session * sess, + const struct ggml_tensor * src0, // W0 + const struct ggml_tensor * src1, // x + const struct ggml_tensor * dst, // dst0 + int32_t n_weights, + struct htp_mm_kernel_params * kparams +) { + ggml_hexagon_precompute_matmul_params_impl(sess, src0, src1, dst, 0, kparams); + kparams->n_weights = n_weights; } static bool ggml_hexagon_tensor_is_host(const struct ggml_hexagon_session * sess, const struct ggml_tensor * t) { @@ -4010,11 +4296,6 @@ static bool ggml_hexagon_supported_mul_mat(const struct ggml_hexagon_session * s return false; } - // hardcoded limit to refuse the lm-head for now - if (src0->ne[1] > 32768) { - return false; - } - if (src1->ne[2] != 1 || src1->ne[3] != 1) { return false; // no broadcasting (for now) } @@ -4348,6 +4629,13 @@ static bool ggml_hexagon_supported_get_rows(const struct ggml_hexagon_session * const struct ggml_tensor * src1 = op->src[1]; // indices const struct ggml_tensor * dst = op; + if (src0->extra) { + const auto * extra = (const ggml_hexagon_tensor_extra *) src0->extra; + if (extra->flags & GGML_HEXAGON_TENSOR_REPACK) { + return false; + } + } + if (src0->type != GGML_TYPE_F32 && src0->ne[0] < 32) { return false; } @@ -4734,10 +5022,43 @@ static bool mm_is_hmx_eligible(const ggml_tensor * t) { return ggml_hexagon_matmul_is_hmx_eligible(src0, src1, t, ne01_padded, is_matmul_id, is_batched); } +static bool is_supported_mul_mat_nx_kernel(const ggml_tensor * src0, const struct htp_mm_kernel_params * kparams) { + if (kparams->n_hmx) { + return kparams->kernel_type == HTP_MM_KERNEL_HMX_2D; + } + + if (!ggml_hexagon_is_repack_type(src0->type)) { + return false; + } + + return kparams->kernel_type == HTP_MM_KERNEL_HVX_QUANT_ROW || kparams->kernel_type == HTP_MM_KERNEL_HVX_QUANT_ROW_FLAT; +} + +static bool is_supported_mul_mat_id_nx_kernel(const ggml_tensor * src0, const struct htp_mm_kernel_params * kparams) { + if (kparams->n_hmx) { + return kparams->kernel_type == HTP_MM_KERNEL_HMX_2D; + } + + if (!ggml_hexagon_is_repack_type(src0->type)) { + return false; + } + + return kparams->kernel_type == HTP_MM_KERNEL_HVX_QUANT_ROW || kparams->kernel_type == HTP_MM_KERNEL_HVX_QUANT_BLOCK; +} + static bool is_mergeable_mul_mat(const ggml_tensor * t) { - if (!t || t->op != GGML_OP_MUL_MAT) return false; - if (t->src[1]->type != GGML_TYPE_F32) return false; - return ggml_is_quantized(t->src[0]->type) && !mm_is_hmx_eligible(t); + if (!t || t->op != GGML_OP_MUL_MAT) return false; + + const ggml_tensor * src0 = t->src[0]; + const ggml_tensor * src1 = t->src[1]; + if (src1->type != GGML_TYPE_F32) return false; + if (src0->ne[2] != 1 || src0->ne[3] != 1) return false; + + if (mm_is_hmx_eligible(t)) { + return ggml_hexagon_is_hmx_weight_type(src0->type); + } + + return ggml_hexagon_is_repack_type(src0->type); } static bool is_mergeable_mul_mat_pair(const ggml_tensor * n1, const ggml_tensor * n2) { @@ -4753,6 +5074,41 @@ static bool is_mergeable_mul_mat_pair(const ggml_tensor * n1, const ggml_tensor if (n1->src[0]->type != n2->src[0]->type) { return false; } + if (mm_is_hmx_eligible(n1) != mm_is_hmx_eligible(n2)) { + return false; + } + return true; +} + +static bool is_mergeable_mul_mat_id(const ggml_tensor * t) { + if (!t || t->op != GGML_OP_MUL_MAT_ID) return false; + + const ggml_tensor * src0 = t->src[0]; + return ggml_hexagon_is_repack_type(src0->type); +} + +static bool is_mergeable_mul_mat_id_pair(const ggml_tensor * n1, const ggml_tensor * n2) { + if (!is_mergeable_mul_mat_id(n1) || !is_mergeable_mul_mat_id(n2)) { + return false; + } + if (n1->src[1] != n2->src[1]) { + return false; + } + if (n1->src[2] != n2->src[2]) { + return false; + } + if (n1->src[0]->ne[0] != n2->src[0]->ne[0]) { + return false; + } + if (n1->src[0]->ne[2] != n2->src[0]->ne[2]) { + return false; + } + if (n1->src[0]->type != n2->src[0]->type) { + return false; + } + if (mm_is_hmx_eligible(n1) != mm_is_hmx_eligible(n2)) { + return false; + } return true; } @@ -4776,8 +5132,8 @@ static ggml_status ggml_backend_hexagon_graph_compute(ggml_backend_t backend, gg if (graph->nodes[i]->op == GGML_OP_RMS_NORM && ggml_can_fuse(graph, i, { GGML_OP_RMS_NORM, GGML_OP_MUL })) { extra->flags |= GGML_HEXAGON_TENSOR_FUSEABLE; - } else if (graph->nodes[i]->op == GGML_OP_MUL_MAT) { - if ((i + 1 < graph->n_nodes && graph->nodes[i + 1]->op == GGML_OP_ADD && ggml_can_fuse(graph, i, { GGML_OP_MUL_MAT, GGML_OP_ADD })) || + } else if (graph->nodes[i]->op == GGML_OP_MUL_MAT || graph->nodes[i]->op == GGML_OP_MUL_MAT_ID) { + if ((i + 1 < graph->n_nodes && graph->nodes[i + 1]->op == GGML_OP_ADD && ggml_can_fuse(graph, i, { graph->nodes[i]->op, GGML_OP_ADD })) || ggml_node_has_n_uses(graph, i, 1)) { extra->flags |= GGML_HEXAGON_TENSOR_FUSEABLE; } diff --git a/ggml/src/ggml-hexagon/htp-opnode.h b/ggml/src/ggml-hexagon/htp-opnode.h index 741b5e04eb8..b083e26718b 100644 --- a/ggml/src/ggml-hexagon/htp-opnode.h +++ b/ggml/src/ggml-hexagon/htp-opnode.h @@ -315,7 +315,8 @@ struct htp_opformat { } void format_kernel_params(char * str, size_t max_size, const htp_opnode & node) { if (node.opcode == HTP_OP_MUL_MAT || node.opcode == HTP_OP_MUL_MAT_ID || - node.opcode == HTP_OP_MUL_MAT_NX || node.opcode == HTP_OP_MUL_MAT_ADD) { + node.opcode == HTP_OP_MUL_MAT_NX || node.opcode == HTP_OP_MUL_MAT_ID_NX || + node.opcode == HTP_OP_MUL_MAT_ADD) { const auto * kparams = (const struct htp_mm_kernel_params *) node.kernel_params; const char * path = "unknown"; int32_t type = kparams->kernel_type; diff --git a/ggml/src/ggml-hexagon/htp/htp-ctx.h b/ggml/src/ggml-hexagon/htp/htp-ctx.h index 88ecf144b94..c8a909d6190 100644 --- a/ggml/src/ggml-hexagon/htp/htp-ctx.h +++ b/ggml/src/ggml-hexagon/htp/htp-ctx.h @@ -118,6 +118,7 @@ struct htp_context { int op_matmul(struct htp_ops_context * octx); int op_matmul_id(struct htp_ops_context * octx); int op_matmul_nx(struct htp_ops_context * octx); +int op_matmul_id_nx(struct htp_ops_context * octx); int op_binary(struct htp_ops_context * octx); int op_unary(struct htp_ops_context * octx); int op_sum_rows(struct htp_ops_context * octx); diff --git a/ggml/src/ggml-hexagon/htp/htp-ops.h b/ggml/src/ggml-hexagon/htp/htp-ops.h index 53c95f28d32..cf938f7eea3 100644 --- a/ggml/src/ggml-hexagon/htp/htp-ops.h +++ b/ggml/src/ggml-hexagon/htp/htp-ops.h @@ -52,6 +52,7 @@ enum htp_op_code { HTP_OP_MUL_MAT, HTP_OP_MUL_MAT_ID, HTP_OP_MUL_MAT_NX, + HTP_OP_MUL_MAT_ID_NX, HTP_OP_MUL_MAT_ADD, HTP_OP_RMS_NORM, HTP_OP_RMS_NORM_MUL, diff --git a/ggml/src/ggml-hexagon/htp/main.c b/ggml/src/ggml-hexagon/htp/main.c index 72cf02a326b..3ab4613cf92 100644 --- a/ggml/src/ggml-hexagon/htp/main.c +++ b/ggml/src/ggml-hexagon/htp/main.c @@ -753,6 +753,9 @@ static int execute_op(struct htp_ops_context * octx) { case HTP_OP_MUL_MAT_ID: return op_matmul_id(octx); + case HTP_OP_MUL_MAT_ID_NX: + return op_matmul_id_nx(octx); + case HTP_OP_MUL_MAT_NX: return op_matmul_nx(octx); @@ -878,8 +881,8 @@ static inline void drop_mmap(struct htp_context *ctx, struct htp_mmap *m) { } } -static inline void mmap_buf(struct htp_context *ctx, struct htp_buf_desc *b) { - if (b->base) return; // already mapped +static inline bool mmap_buf(struct htp_context *ctx, struct htp_buf_desc *b) { + if (b->base) return true; // already mapped // find unused mapping for (uint32_t i=0; i < HTP_MAX_MMAPS; i++) { @@ -887,8 +890,8 @@ static inline void mmap_buf(struct htp_context *ctx, struct htp_buf_desc *b) { if (!m->size) { void *va = htp_mmap(b->fd, b->size); if (va == NULL) { - FARF(ERROR, "mmap failed : fd %u size %u", b->fd, (uint32_t) b->size); - abort(); // can't do much else at this point + FARF(HIGH, "mmap failed (will attempt defrag) : fd %u size %u", b->fd, (uint32_t) b->size); + return false; } m->base = b->base = (uint64_t) va; @@ -896,12 +899,12 @@ static inline void mmap_buf(struct htp_context *ctx, struct htp_buf_desc *b) { m->size = b->size; FARF(ALWAYS, "mmap : fd %u base %p size %u", m->fd, (void*) m->base, (uint32_t) m->size); - return; + return true; } } FARF(ERROR, "mmap failed : exceeded mapping capacity limit of %u", HTP_MAX_MMAPS); - abort(); + return false; } static void prep_op_bufs(struct htp_context *ctx, struct htp_buf_desc *bufs, uint32_t n_bufs) { @@ -934,12 +937,32 @@ static void prep_op_bufs(struct htp_context *ctx, struct htp_buf_desc *bufs, uin } } - // Create missing mappings + // Create missing mappings (pass 1) + bool mmap_ok = true; for (uint32_t i=0; i < n_bufs; i++) { struct htp_buf_desc *b = bufs + i; - mmap_buf(ctx, b); + if (!mmap_buf(ctx, b)) { + mmap_ok = false; + break; + } FARF(HIGH, "prep-buf #%u : pass1 fd %u base %p size %u flags 0x%x", i, b->fd, (void*) b->base, (uint32_t) b->size, b->flags); } + + if (!mmap_ok) { + // Attempt clean defragmentation: drop all mappings and remap (pass 2) + FARF(HIGH, "prep-bufs : dropping all mappings to defragment address space"); + for (uint32_t i=0; i < HTP_MAX_MMAPS; i++) { drop_mmap(ctx, ctx->mmap + i); } + + for (uint32_t i=0; i < n_bufs; i++) { + struct htp_buf_desc *b = bufs + i; + b->base = 0; + if (!mmap_buf(ctx, b)) { + FARF(ERROR, "prep-bufs : mmap failed after defragmentation (fd %u size %u)", b->fd, (uint32_t) b->size); + abort(); + } + FARF(HIGH, "prep-buf #%u : pass2 fd %u base %p size %u flags 0x%x", i, b->fd, (void*) b->base, (uint32_t) b->size, b->flags); + } + } } static void prep_tensor(struct htp_context *ctx, struct htp_buf_desc *bufs, struct htp_tensor *tens, uint32_t idx, struct htp_tensor *t) { diff --git a/ggml/src/ggml-hexagon/htp/matmul-ops.c b/ggml/src/ggml-hexagon/htp/matmul-ops.c index a6adc0e61fa..2a87dd19ee8 100644 --- a/ggml/src/ggml-hexagon/htp/matmul-ops.c +++ b/ggml/src/ggml-hexagon/htp/matmul-ops.c @@ -55,10 +55,14 @@ typedef struct { size_t src0_nb3; size_t src1_nb2; size_t src1_nb3; - size_t dst_nb2; - size_t dst_nb3; size_t src2_nb2; size_t src2_nb3; + size_t dst_nb2; + size_t dst_nb3; + int r2; + int r3; + struct fastdiv_values div_r2; + struct fastdiv_values div_r3; } hmx_mm_f16_f32_batched_params_t; struct htp_mm_context { @@ -235,17 +239,18 @@ static void hvx_mm_4d(unsigned int nth, unsigned int ith, void * data) { const uint32_t nr1 = ne1 * ne2 * ne3; // distribute the thread work across the inner or outer loop based on which one is larger - uint32_t nchunk0 = nr0 > nr1 ? nth : 1; // parallelize by src0 rows - uint32_t nchunk1 = nr0 > nr1 ? 1 : nth; // parallelize by src1 rows - - // The number of elements in each chunk - const uint32_t dr0 = (nr0 + nchunk0 - 1) / nchunk0; - const uint32_t dr1 = (nr1 + nchunk1 - 1) / nchunk1; - - uint32_t current_chunk = ith; - - const uint32_t ith0 = current_chunk % nchunk0; - const uint32_t ith1 = current_chunk / nchunk0; + uint32_t dr0, dr1, ith0, ith1; + if (nr0 > nr1) { + dr0 = fastdiv(nr0 + nth - 1, &octx->ctx->n_threads_div); + dr1 = nr1; + ith0 = ith; + ith1 = 0; + } else { + dr0 = nr0; + dr1 = fastdiv(nr1 + nth - 1, &octx->ctx->n_threads_div); + ith0 = 0; + ith1 = ith; + } const uint32_t ir0_start = dr0 * ith0; const uint32_t ir0_end = MIN(ir0_start + dr0, nr0); @@ -545,7 +550,7 @@ static void hvx_mm_nx_2d_repacked_##SUFFIX(unsigned int nth, unsigned int ith, v uint32_t tile_row_stride = n_k_tiles_w * tile_size; \ \ const uint32_t src0_nrows = ne01 * src_w->ne[2] * src_w->ne[3]; \ - uint32_t src0_nrows_per_thread = (src0_nrows + nth - 1) / nth; \ + uint32_t src0_nrows_per_thread = fastdiv(src0_nrows + nth - 1, &octx->ctx->n_threads_div); \ src0_nrows_per_thread = hex_round_up(src0_nrows_per_thread, 32); \ \ const uint32_t start_row = src0_nrows_per_thread * ith; \ @@ -1105,6 +1110,179 @@ static void hvx_mv_id(unsigned int nth, unsigned int ith, void * data) { } } +static void hvx_mv_id_nx(unsigned int nth, unsigned int ith, void * data) { + struct htp_mm_context * mmctx = (struct htp_mm_context *) data; + struct htp_ops_context * octx = mmctx->octx; + dma_queue * dma_queue = octx->ctx->dma[ith]; + const struct htp_mm_kernel_params * kparams = (const struct htp_mm_kernel_params *) octx->kernel_params; + const uint32_t n_weights = kparams->n_weights; + const struct htp_tensor * restrict src0 = octx->src[0]; + const struct htp_tensor * restrict act = octx->src[n_weights]; + const struct htp_tensor * restrict ids = octx->src[n_weights + 1]; + + hvx_mm_run_quant_task(mmctx, ith); + + struct htp_thread_trace * tr = &octx->ctx->trace[ith]; + + const uint32_t n_prefetch = kparams->n_prefetch; + assert(n_prefetch >= 2 && n_prefetch <= HTP_MM_MAX_PREFETCH && (n_prefetch & (n_prefetch - 1)) == 0); + + const uint32_t n_aids = ids->ne[0]; + const uint32_t n_ids = src0->ne[2]; + + uint8_t * restrict vtcm_src0_ptr = mmctx->vtcm_src0 + mmctx->vtcm_src0_size_per_thread * ith; + uint8_t * restrict src1_data = mmctx->vtcm_src1; + + for (uint32_t ie1 = 0; ie1 < n_aids; ++ie1) { + const int32_t eid = *(const int32_t *) ((const uint8_t *) ids->data + ie1 * ids->nb[0]); + if (eid < 0) continue; + assert(eid < (int32_t) n_ids); + + for (uint32_t p = 0; p < n_weights; ++p) { + const struct htp_tensor * restrict src_w = octx->src[p]; + const struct htp_tensor * restrict dst = octx->dsts[p]; + if (!src_w || !dst) continue; + + const uint32_t src0_nrows = src_w->ne[1]; + uint32_t src0_nrows_per_thread = fastdiv(src0_nrows + nth - 1, &octx->ctx->n_threads_div); + src0_nrows_per_thread = hex_round_up(src0_nrows_per_thread, 32); + + const uint32_t src0_start_row = src0_nrows_per_thread * ith; + const uint32_t src0_end_row = MIN(src0_start_row + src0_nrows_per_thread, src0_nrows); + if (src0_start_row >= src0_end_row) continue; + + const uint8_t * restrict src0_row = (const uint8_t *) src_w->data + eid * src_w->nb[2]; + const uint8_t * restrict src1_col = (const uint8_t *) src1_data; + float * restrict dst_row = (float *) (dst->data + ie1 * dst->nb[1]); + + const uint32_t tile_size = htp_mm_get_weight_tile_size(src_w->type); + const uint32_t aligned_tile_size = htp_mm_get_weight_aligned_tile_size(src_w->type); + const uint32_t n_k_tiles_w = src_w->ne[0] / 32; + const uint32_t n_k_tiles_a = act->ne[0] / 32; + const uint32_t tile_row_stride = n_k_tiles_w * tile_size; + const uint32_t tile_row_transfer_size_aligned = n_k_tiles_a * aligned_tile_size; + + const uint32_t ct_start = src0_start_row / 32; + const uint32_t ct_end = (src0_end_row + 31) / 32; + + uint32_t push_ct = ct_start; + for (uint32_t d = 0; d < n_prefetch && push_ct < ct_end; d++, push_ct++) { + dma_queue_push(dma_queue, dma_make_ptr(vtcm_src0_ptr + d * tile_row_transfer_size_aligned, src0_row + push_ct * tile_row_stride), + aligned_tile_size, tile_size, tile_size, n_k_tiles_a); + } + + for (uint32_t ct = ct_start; ct < ct_end; ct++) { + const uint8_t * w_tile = dma_queue_pop(dma_queue).dst; + + int valid_rows = (int)src_w->ne[1] - (int)(ct * 32); + valid_rows = MIN(32, MAX(0, valid_rows)); + + htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, ct); + mmctx->vec_dot_32x1(act->ne[0], &dst_row[ct * 32], w_tile, src1_col, valid_rows, NULL); + htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, ct); + + if (push_ct < ct_end) { + dma_queue_push(dma_queue, dma_make_ptr((uint8_t *)w_tile, src0_row + push_ct * tile_row_stride), + aligned_tile_size, tile_size, tile_size, n_k_tiles_a); + push_ct++; + } + } + } + } +} + +static void hvx_mm_id_nx(unsigned int nth, unsigned int ith, void * data) { + struct htp_mm_context * mmctx = (struct htp_mm_context *) data; + struct htp_ops_context * octx = mmctx->octx; + dma_queue * dma_queue = octx->ctx->dma[ith]; + const struct htp_mm_kernel_params * kparams = (const struct htp_mm_kernel_params *) octx->kernel_params; + const uint32_t n_weights = kparams->n_weights; + const struct htp_tensor * restrict src0 = octx->src[0]; + const struct htp_tensor * restrict act = octx->src[n_weights]; + const struct htp_tensor * restrict ids = octx->src[n_weights + 1]; + + hvx_mm_run_quant_task(mmctx, ith); + + struct htp_thread_trace * tr = &octx->ctx->trace[ith]; + + const uint32_t n_prefetch = kparams->n_prefetch; + assert(n_prefetch >= 2 && n_prefetch <= HTP_MM_MAX_PREFETCH && (n_prefetch & (n_prefetch - 1)) == 0); + + const uint32_t n_as = src0->ne[2]; + + const uint32_t * matrix_row_counts = mmctx->matrix_row_counts; + const struct mmid_row_mapping * matrix_rows = mmctx->matrix_rows; + + const size_t src1_stride = mmctx->vtcm_src1_stride; + + uint8_t * restrict vtcm_src0_ptr = mmctx->vtcm_src0 + mmctx->vtcm_src0_size_per_thread * ith; + uint8_t * restrict src1_data = mmctx->vtcm_src1; + + for (uint32_t cur_a = 0; cur_a < n_as; ++cur_a) { + const int32_t cne1 = matrix_row_counts[cur_a]; + if (cne1 == 0) continue; + + for (uint32_t p = 0; p < n_weights; ++p) { + const struct htp_tensor * restrict src_w = octx->src[p]; + const struct htp_tensor * restrict dst = octx->dsts[p]; + if (!src_w || !dst) continue; + + const uint32_t src0_nrows = src_w->ne[1]; + uint32_t src0_nrows_per_thread = fastdiv(src0_nrows + nth - 1, &octx->ctx->n_threads_div); + src0_nrows_per_thread = hex_round_up(src0_nrows_per_thread, 32); + + const uint32_t src0_start_row = src0_nrows_per_thread * ith; + const uint32_t src0_end_row = MIN(src0_start_row + src0_nrows_per_thread, src0_nrows); + if (src0_start_row >= src0_end_row) continue; + + const uint8_t * src0_row = (const uint8_t *) src_w->data + cur_a * src_w->nb[2]; + + const uint32_t tile_size = htp_mm_get_weight_tile_size(src_w->type); + const uint32_t aligned_tile_size = htp_mm_get_weight_aligned_tile_size(src_w->type); + const uint32_t n_k_tiles_w = src_w->ne[0] / 32; + const uint32_t n_k_tiles_a = act->ne[0] / 32; + const uint32_t tile_row_stride = n_k_tiles_w * tile_size; + const uint32_t tile_row_transfer_size_aligned = n_k_tiles_a * aligned_tile_size; + + const uint32_t ct_start = src0_start_row / 32; + const uint32_t ct_end = (src0_end_row + 31) / 32; + + uint32_t push_ct = ct_start; + for (uint32_t d = 0; d < n_prefetch && push_ct < ct_end; d++, push_ct++) { + dma_queue_push(dma_queue, dma_make_ptr(vtcm_src0_ptr + d * tile_row_transfer_size_aligned, src0_row + push_ct * tile_row_stride), + aligned_tile_size, tile_size, tile_size, n_k_tiles_a); + } + + for (uint32_t ct = ct_start; ct < ct_end; ct++) { + const uint8_t * w_tile = dma_queue_pop(dma_queue).dst; + + int valid_rows = (int)src_w->ne[1] - (int)(ct * 32); + valid_rows = MIN(32, MAX(0, valid_rows)); + + htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, ct); + for (uint32_t cid = 0; cid < (uint32_t) cne1; ++cid) { + struct mmid_row_mapping row_mapping = MMID_MATRIX_ROW(cur_a, cid); + const int rm1 = row_mapping.i1; + const int rm2 = row_mapping.i2; + + const uint32_t ir1 = fastmodulo(rm1, act->ne[1], &mmctx->mm_div_ne11); + const uint8_t * restrict src1_col = (const uint8_t *) (src1_data + (ir1 + rm2 * act->ne[1]) * src1_stride); + float * restrict dst_row = (float *) (dst->data + (rm1 * dst->nb[1] + rm2 * dst->nb[2])); + + mmctx->vec_dot_32x1(act->ne[0], &dst_row[ct * 32], w_tile, src1_col, valid_rows, NULL); + } + htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, ct); + + if (push_ct < ct_end) { + dma_queue_push(dma_queue, dma_make_ptr((uint8_t *)w_tile, src0_row + push_ct * tile_row_stride), + aligned_tile_size, tile_size, tile_size, n_k_tiles_a); + push_ct++; + } + } + } + } +} + static int hvx_mm_init_vec_dot(struct htp_mm_context * mmctx, enum htp_data_type type) { switch (type) { case HTP_TYPE_Q4_0: @@ -1153,7 +1331,7 @@ static int hvx_mm_matmul(struct htp_ops_context * octx) { src0->type == HTP_TYPE_MXFP4); // Compute src0_nrows_per_thread - mmctx->src0_nrows_per_thread = (src0_nrows + octx->n_threads - 1) / octx->n_threads; + mmctx->src0_nrows_per_thread = fastdiv(src0_nrows + octx->n_threads - 1, &octx->ctx->n_threads_div); if (is_repacked) { mmctx->src0_nrows_per_thread = hex_round_up(mmctx->src0_nrows_per_thread, 32); } else { @@ -1325,11 +1503,11 @@ static int hvx_mm_matmul(struct htp_ops_context * octx) { kparams->kernel_type == HTP_MM_KERNEL_HVX_QUANT_BLOCK) { mmctx->vtcm_src1_size_per_thread = L.src1_bytes; } else { - mmctx->vtcm_src1_size_per_thread = L.src1_bytes / octx->n_threads; + mmctx->vtcm_src1_size_per_thread = fastdiv(L.src1_bytes, &octx->ctx->n_threads_div); } - mmctx->vtcm_src0_size_per_thread = L.src0_bytes / octx->n_threads; - mmctx->vtcm_dst_size_per_thread = L.dst_bytes / octx->n_threads; + mmctx->vtcm_src0_size_per_thread = fastdiv(L.src0_bytes, &octx->ctx->n_threads_div); + mmctx->vtcm_dst_size_per_thread = fastdiv(L.dst_bytes, &octx->ctx->n_threads_div); size_t vtcm_size = kparams->vtcm_size > 0 ? (size_t)kparams->vtcm_size : L.total_bytes; @@ -1407,7 +1585,7 @@ static void hvx_mm_nx_2d(unsigned int nth, unsigned int ith, void * data) { const uint32_t ne01 = src_w->ne[1]; const uint32_t src0_nrows = ne01 * src_w->ne[2] * src_w->ne[3]; - uint32_t src0_nrows_per_thread = (src0_nrows + nth - 1) / nth; + uint32_t src0_nrows_per_thread = fastdiv(src0_nrows + nth - 1, &octx->ctx->n_threads_div); src0_nrows_per_thread += (src0_nrows_per_thread & 1); const uint32_t src0_start_row = src0_nrows_per_thread * ith; @@ -1538,36 +1716,36 @@ static void transfer_output_chunk_worker_fn(unsigned int n, unsigned int i, void } typedef struct { - const struct mmid_row_mapping *matrix_rows; - __fp16 *dst; - const float *src; - uint32_t n_tasks; - uint32_t n_tot_chunks; - uint32_t n_chunks_per_task; - uint32_t k_block; - uint32_t k_stride; - uint32_t k_valid; - struct htp_thread_trace * traces; - struct htp_context * ctx; - float * vtcm_f32_act; - size_t vtcm_f32_act_bytes_per_thread; - uint32_t dma_step_rows; - uint32_t dma_step_rows_shift; + struct htp_context * ctx; + struct htp_thread_trace * traces; + __fp16 * dst; + const float * src; + const struct mmid_row_mapping * matrix_rows; + float * vtcm_f32_act; + uint32_t n_tasks; + uint32_t n_tot_chunks; + uint32_t n_chunks_per_task; + uint32_t k_block; + uint32_t k_stride; + uint32_t k_valid; + size_t vtcm_f32_act_bytes_per_thread; + uint32_t dma_step_rows; + uint32_t dma_step_rows_shift; } activation_transfer_task_state_t; typedef struct { - __fp16 *dst; - const float *src; + struct htp_context * ctx; + struct htp_thread_trace * traces; + __fp16 * dst; + const float * src; + float * vtcm_f32_act; uint32_t n_rows; uint32_t k_block; uint32_t k_stride; uint32_t k_valid; uint32_t n_col_chunks; struct fastdiv_values n_threads_div; - float *vtcm_f32_act; size_t vtcm_f32_act_bytes; - struct htp_thread_trace *traces; - struct htp_context *ctx; uint32_t dma_step_rows; uint32_t dma_step_rows_shift; } activation_transfer_col_chunk_state_t; @@ -1811,9 +1989,10 @@ static void transfer_activation_chunk_worker_fn(unsigned int n, unsigned int i, } typedef struct { - const struct mmid_row_mapping *matrix_rows; - __fp16 *dst; - const float *src; + struct htp_thread_trace * traces; + const struct mmid_row_mapping * matrix_rows; + __fp16 * dst; + const float * src; uint32_t n_tasks; uint32_t n_tot_chunks; uint32_t n_chunks_per_task; @@ -1827,13 +2006,13 @@ typedef struct { uint32_t start_row; uint32_t cne1; uint32_t k_valid; - struct htp_thread_trace *traces; } activation_transfer_gathered_task_state_t; typedef struct { - const struct mmid_row_mapping *matrix_rows; - const __fp16 *vtcm_src; - float *dst; + struct htp_thread_trace * traces; + const struct mmid_row_mapping * matrix_rows; + const __fp16 * vtcm_src; + float * dst; uint32_t n_tasks; uint32_t n_tot_chunks; uint32_t n_chunks_per_task; @@ -1844,17 +2023,16 @@ typedef struct { size_t dst_nb2; uint32_t start_row; uint32_t cne1; - struct htp_thread_trace *traces; } output_transfer_scattered_task_state_t; static void transfer_activation_chunk_gathered_worker_fn(unsigned int n, unsigned int i, void *data) { activation_transfer_gathered_task_state_t *st = data; struct htp_thread_trace * tr = &st->traces[i]; - int chunk_idx = i; - int chunk_size = st->n_chunks_per_task; + int chunk_idx = i; + int chunk_size = st->n_chunks_per_task; int vtcm_start_row = chunk_idx * chunk_size; - int start_row = st->start_row + vtcm_start_row; - int n_rows = hex_smin(st->cne1 - start_row, chunk_size); + int start_row = st->start_row + vtcm_start_row; + int n_rows = hex_smin(st->cne1 - start_row, chunk_size); if (n_rows > 0) { htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_A_PREP, chunk_idx); transfer_activation_chunk_fp32_to_fp16_gathered( @@ -1946,17 +2124,17 @@ static void dequantize_tiled_weight_chunk_to_fp16_tiles( } typedef struct { - float *dst; - const float *src2; - const __fp16 *vtcm_src; - uint32_t n_rows; - uint32_t n_cols; - uint32_t dst_stride; - uint32_t src2_stride; - uint32_t dst_cols; - struct fastdiv_values n_threads_div; - struct htp_thread_trace *traces; - struct htp_context *ctx; + struct htp_context * ctx; + struct htp_thread_trace * traces; + float * dst; + const __fp16 * vtcm_src; + const float * src2; + uint32_t n_rows; + uint32_t n_cols; + uint32_t dst_stride; + uint32_t src2_stride; + uint32_t dst_cols; + struct fastdiv_values n_threads_div; } output_transfer_col_chunk_state_t; static void transfer_output_chunk_col_chunk_worker_fn(unsigned int n, unsigned int i, void *data) { @@ -1965,19 +2143,19 @@ static void transfer_output_chunk_col_chunk_worker_fn(unsigned int n, unsigned i struct htp_thread_trace * tr = &st->traces[i]; uint32_t n_blocks = st->n_cols / 32; - uint32_t b_first = fastdiv(n_blocks * i, &st->n_threads_div); - uint32_t b_last = fastdiv(n_blocks * (i + 1), &st->n_threads_div); - uint32_t c_first = b_first * 32; - uint32_t c_last = b_last * 32; - uint32_t c_len = c_last - c_first; + uint32_t b_first = fastdiv(n_blocks * i, &st->n_threads_div); + uint32_t b_last = fastdiv(n_blocks * (i + 1), &st->n_threads_div); + uint32_t c_first = b_first * 32; + uint32_t c_last = b_last * 32; + uint32_t c_len = c_last - c_first; if (c_len == 0) return; htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_O_PROC, c_first); - float *dst = st->dst + c_first; - const float *src2 = st->src2 ? (st->src2 + c_first) : NULL; const __fp16 *vtcm_src = st->vtcm_src + b_first * HTP_MM_HMX_TILE_N_ELMS; + const float *src2 = st->src2 ? (st->src2 + c_first) : NULL; + float *dst = st->dst + c_first; int chunk_dst_cols = (int)st->dst_cols - (int)c_first; if (chunk_dst_cols > 0) { @@ -1998,7 +2176,7 @@ static void transfer_output_chunk_threaded(struct htp_context *ctx, float *dst, uint32_t n_blocks = (uint32_t)n_cols / 32; if (n_threads > 1 && n_blocks >= (uint32_t)n_threads) { - struct fastdiv_values n_threads_div = init_fastdiv_values(n_threads); + struct fastdiv_values n_threads_div = (n_threads == (int)ctx->n_threads) ? ctx->n_threads_div : init_fastdiv_values(n_threads); output_transfer_col_chunk_state_t col_state; col_state.dst = dst; col_state.src2 = src2; @@ -2128,8 +2306,7 @@ static void transfer_activation_chunk_threaded(const struct activation_transfer_ state.ctx = ctx; state.vtcm_f32_act = vtcm_f32_act; - int active_threads = hex_smin(n_threads, (int)state.n_tasks); - state.vtcm_f32_act_bytes_per_thread = hex_align_down(vtcm_f32_act_bytes / active_threads, 128); + state.vtcm_f32_act_bytes_per_thread = hex_align_down(fastdiv(vtcm_f32_act_bytes, act_threads_div), 128); uint32_t dma_step_rows = 2; uint32_t dma_step_rows_shift = 1; @@ -2144,6 +2321,7 @@ static void transfer_activation_chunk_threaded(const struct activation_transfer_ state.dma_step_rows = dma_step_rows; state.dma_step_rows_shift = dma_step_rows_shift; + int active_threads = hex_smin(n_threads, (int)state.n_tasks); if (state.n_tasks == 1 || n_threads == 1) { transfer_activation_chunk_worker_fn(1, 0, &state); } else { @@ -2447,21 +2625,286 @@ static int hmx_mm_2d_f32(struct htp_context *ctx, return 0; } -static inline int hmx_mm_batch_r2(const hmx_mm_f16_f32_batched_params_t *params) { - return params->ne02 > 0 ? params->ne12 / params->ne02 : 1; -} +static int hmx_mm_nx_2d_f32(struct htp_ops_context * octx, const struct htp_mm_kernel_params * kparams) { + struct htp_context * ctx = octx->ctx; + struct htp_thread_trace * tr = &ctx->trace[0]; + htp_trace_event_start(tr, HTP_TRACE_EVT_INIT, 0); + + const uint32_t n_weights = kparams->n_weights; + if (n_weights == 0 || n_weights > HTP_OP_MAX_OUTPUTS) { + return HTP_STATUS_INVAL_PARAMS; + } + + const struct htp_tensor * restrict src0 = octx->src[0]; + const struct htp_tensor * restrict act = octx->src[n_weights]; + + if (!src0 || !act) { + return HTP_STATUS_INVAL_PARAMS; + } + + const int weight_type = (int) src0->type; + const int k = (int) act->ne[0]; + const int k_valid = (int) act->ne[0]; + const int m = (int) (act->ne[1] * act->ne[2] * act->ne[3]); + const int act_stride = (int) (act->nb[1] / sizeof(float)); + const float * activation = (const float *) act->data; + + if (k % 32 != 0) { return HTP_STATUS_NO_SUPPORT; } + if (!hex_is_aligned(activation, VLEN)) { return HTP_STATUS_NO_SUPPORT; } + + size_t row_stride = htp_mm_get_tiled_row_stride(weight_type, k); + if (row_stride == 0) { + return HTP_STATUS_NO_SUPPORT; + } + + worker_callback_t dequant_worker_fn = NULL; + switch (weight_type) { + case HTP_TYPE_Q4_0: dequant_worker_fn = dequantize_tiled_worker_loop_q4_0; break; + case HTP_TYPE_IQ4_NL: dequant_worker_fn = dequantize_tiled_worker_loop_iq4_nl; break; + case HTP_TYPE_Q4_1: dequant_worker_fn = dequantize_tiled_worker_loop_q4_1; break; + case HTP_TYPE_MXFP4: dequant_worker_fn = dequantize_tiled_worker_loop_mxfp4; break; + case HTP_TYPE_Q8_0: dequant_worker_fn = dequantize_tiled_worker_loop_q8_0; break; + case HTP_TYPE_F16: dequant_worker_fn = convert_f16_worker_loop; break; + case HTP_TYPE_F32: dequant_worker_fn = quantize_f32_worker_loop; break; + default: + return HTP_STATUS_NO_SUPPORT; + } + + const int n_k_tiles = k / HTP_MM_HMX_TILE_N_COLS; + const struct fastdiv_values n_k_tiles_div = init_fastdiv_values(n_k_tiles); + + const bool is_quant = (weight_type != HTP_TYPE_F16 && weight_type != HTP_TYPE_F32); + const size_t vtcm_budget = ctx->vtcm_size; + + const int m_chunk_n_rows = kparams->m_chunk; + const int n_chunk_n_cols = kparams->n_chunk; + const int pipeline = kparams->pipeline; + const int n_threads = octx->n_threads; + const int act_threads = kparams->n_act_threads; + const struct fastdiv_values * act_threads_div = &kparams->div_n_act_threads; + const struct fastdiv_values * k_div = &kparams->div_ne00_padded; + const int tile_size = kparams->tile_size; + const int aligned_tile_size = kparams->aligned_tile_size; + + const uint32_t dma_dst_stride = is_quant ? aligned_tile_size : row_stride; + const uint32_t dma_width_bytes = is_quant ? tile_size : row_stride; + + struct htp_mm_hmx_vtcm_layout L; + htp_mm_hmx_vtcm_layout_build(&L, HTP_MM_KERNEL_HMX_2D, weight_type, k, m_chunk_n_rows, n_chunk_n_cols, 1, false, pipeline, act_threads, aligned_tile_size); + + if (L.total_bytes > vtcm_budget) { + FARF(ERROR, "hmx-mm-nx-2d: VTCM overflow: used %zu budget %zu, m %d k %d mc %d nc %d", + L.total_bytes, vtcm_budget, m, k, m_chunk_n_rows, n_chunk_n_cols); + return HTP_STATUS_VTCM_TOO_SMALL; + } + + uint8_t * const base = (uint8_t *) ctx->vtcm_base; + __fp16 *vtcm_weight_raw[2] = { + VTCM_LAYOUT_PTR(__fp16, base, L.off_weight[0]), + VTCM_LAYOUT_PTR_OPTIONAL(__fp16, base, L.off_weight[1], pipeline) + }; + + __fp16 *vtcm_f16_act = VTCM_LAYOUT_PTR(__fp16, base, L.off_act); + float *vtcm_f32_act = VTCM_LAYOUT_PTR(float, base, L.off_act_f32); + __fp16 *vtcm_output = VTCM_LAYOUT_PTR(__fp16, base, L.off_dst[0]); + void *vtcm_scratch0 = VTCM_LAYOUT_PTR(void, base, L.off_scratch[0]); + void *vtcm_scratch1 = VTCM_LAYOUT_PTR_OPTIONAL(void, base, L.off_scratch[1], pipeline); + void *vtcm_scratch2 = VTCM_LAYOUT_PTR_OPTIONAL(void, base, L.off_dst[1], pipeline); + __fp16 *vtcm_scales = VTCM_LAYOUT_PTR(__fp16, base, L.off_scales); + + hmx_init_column_scales(vtcm_scales, Q6_V_vsplat_R(0x3c00)); // scale: 1.0, bias: 0.0 in FP16 + + FARF(HIGH, "hmx-mm-nx-2d: n_weights %u m %d k %d wtype %d mc %d nc %d vtcm %zu/%zu", + n_weights, m, k, weight_type, m_chunk_n_rows, n_chunk_n_cols, L.total_bytes, vtcm_budget); + + htp_trace_event_stop(tr, HTP_TRACE_EVT_INIT, 0); + + if (pipeline) { + hmx_matmul_job_t job_slots[2]; + + for (size_t mr = 0; mr < (size_t) m; mr += m_chunk_n_rows) { + const size_t n_rows = hex_smin(m - mr, m_chunk_n_rows); + + void *vtcm_weight_bufs[2] = { vtcm_scratch0, vtcm_scratch1 }; + void *vtcm_output_bufs[2] = { vtcm_output, vtcm_scratch2 }; + + struct activation_transfer_params act_params = { + .ctx = ctx, + .dst = vtcm_f16_act, + .src = activation + mr * act_stride, + .n_rows = (int) n_rows, + .k_block = k, + .k_stride = act_stride, + .n_threads = act_threads, + .act_threads_div = act_threads_div, + .k_div = k_div, + .k_valid = k_valid, + .vtcm_f32_act = vtcm_f32_act, + .vtcm_f32_act_bytes = L.act_f32_bytes, + }; + transfer_activation_chunk_threaded(&act_params); + + for (uint32_t p = 0; p < n_weights; p++) { + const struct htp_tensor * restrict src_w = octx->src[p]; + const struct htp_tensor * restrict dst = octx->dsts[p]; + if (!src_w || !dst) continue; + + const uint8_t * weight = (const uint8_t *) src_w->data; + float * dst_ptr = (float *) dst->data; + const size_t n = src_w->ne[1]; + if (n == 0) continue; + const size_t weight_stride = src_w->nb[1]; + const size_t dst_stride = dst->nb[1] / sizeof(float); + const int dst_cols = (int) dst->ne[0]; + const int n_chunk_cnt = hmx_ceil_div(n, n_chunk_n_cols); + + const uint32_t dma_src_stride = is_quant ? tile_size : weight_stride; + + const size_t n_cols_A0 = hex_smin(n - 0 * n_chunk_n_cols, n_chunk_n_cols); + const uint32_t height_A0 = is_quant ? (n_cols_A0 / 32) * n_k_tiles : n_cols_A0; + dma_queue_push(ctx->dma[0], dma_make_ptr(vtcm_weight_raw[0], weight), + dma_dst_stride, dma_src_stride, dma_width_bytes, height_A0); + + if (1 < n_chunk_cnt) { + const size_t n_cols_A1 = hex_smin(n - 1 * n_chunk_n_cols, n_chunk_n_cols); + const uint32_t height_A1 = is_quant ? (n_cols_A1 / 32) * n_k_tiles : n_cols_A1; + dma_queue_push(ctx->dma[0], dma_make_ptr(vtcm_weight_raw[1], weight + n_chunk_n_cols * weight_stride), + dma_dst_stride, dma_src_stride, dma_width_bytes, height_A1); + } + + for (int i = 0; i < n_chunk_cnt; ++i) { + const size_t nc = i * n_chunk_n_cols; + const size_t nc_p2 = nc + 2 * n_chunk_n_cols; + + const size_t n_cols = hex_smin(n - nc, n_chunk_n_cols); + const size_t n_cols_p2 = hex_smin(n - nc_p2, n_chunk_n_cols); + + void * curr_raw = dma_queue_pop(ctx->dma[0]).dst; + + dequantize_tiled_weight_chunk_to_fp16_tiles( + ctx, vtcm_weight_bufs[i % 2], curr_raw, + n_cols, k, row_stride, weight_type, + n_k_tiles, n_k_tiles_div, dequant_worker_fn, n_threads); + + if (i + 2 < n_chunk_cnt) { + const uint32_t height_p2 = is_quant ? (n_cols_p2 / 32) * n_k_tiles : n_cols_p2; + dma_queue_push(ctx->dma[0], dma_make_ptr(curr_raw, weight + nc_p2 * weight_stride), + dma_dst_stride, dma_src_stride, dma_width_bytes, height_p2); + } + + hmx_matmul_job_init(&job_slots[i % 2], (__fp16 *) vtcm_output_bufs[i % 2], + (__fp16 *) vtcm_f16_act, (__fp16 *) vtcm_weight_bufs[i % 2], + vtcm_scales, hmx_ceil_div(n_rows, HTP_MM_HMX_TILE_N_ROWS), + hmx_ceil_div(n_cols, HTP_MM_HMX_TILE_N_COLS), k / HTP_MM_HMX_TILE_N_ROWS); + hmx_queue_push(ctx->hmx_queue, hmx_queue_make_desc(hmx_matmul_worker_fn, &job_slots[i % 2])); + + if (i > 0) { + hmx_queue_pop(ctx->hmx_queue); + const size_t nc_prev = (i - 1) * n_chunk_n_cols; + const size_t n_cols_prev = hex_smin(n - nc_prev, n_chunk_n_cols); + float *output_chunk = dst_ptr + (mr * dst_stride + nc_prev); + int chunk_dst_cols = dst_cols - (int)nc_prev; + if (chunk_dst_cols > 0) { + transfer_output_chunk_threaded(ctx, output_chunk, NULL, vtcm_output_bufs[(i - 1) % 2], n_rows, n_cols_prev, dst_stride, 0, chunk_dst_cols, n_threads); + } + } + } + + hmx_queue_pop(ctx->hmx_queue); + const size_t nc_last = (n_chunk_cnt - 1) * n_chunk_n_cols; + const size_t n_cols_last = hex_smin(n - nc_last, n_chunk_n_cols); + float *output_chunk = dst_ptr + (mr * dst_stride + nc_last); + int chunk_dst_cols = dst_cols - (int)nc_last; + if (chunk_dst_cols > 0) { + transfer_output_chunk_threaded(ctx, output_chunk, NULL, vtcm_output_bufs[(n_chunk_cnt - 1) % 2], n_rows, n_cols_last, dst_stride, 0, chunk_dst_cols, n_threads); + } + } + } + } else { + hmx_matmul_job_t job; + for (size_t mr = 0; mr < (size_t) m; mr += m_chunk_n_rows) { + const size_t n_rows = hex_smin(m - mr, m_chunk_n_rows); + + struct activation_transfer_params act_params = { + .ctx = ctx, + .dst = vtcm_f16_act, + .src = activation + mr * act_stride, + .n_rows = (int) n_rows, + .k_block = k, + .k_stride = act_stride, + .n_threads = act_threads, + .act_threads_div = act_threads_div, + .k_div = k_div, + .k_valid = k_valid, + .vtcm_f32_act = vtcm_f32_act, + .vtcm_f32_act_bytes = L.act_f32_bytes, + }; + transfer_activation_chunk_threaded(&act_params); + + for (uint32_t p = 0; p < n_weights; p++) { + const struct htp_tensor * restrict src_w = octx->src[p]; + const struct htp_tensor * restrict dst = octx->dsts[p]; + if (!src_w || !dst) continue; + + const uint8_t * weight = (const uint8_t *) src_w->data; + float * dst_ptr = (float *) dst->data; + const size_t n = src_w->ne[1]; + if (n == 0) continue; + const size_t weight_stride = src_w->nb[1]; + const size_t dst_stride = dst->nb[1] / sizeof(float); + const int dst_cols = (int) dst->ne[0]; + + const uint32_t dma_src_stride = is_quant ? tile_size : weight_stride; + + if (n > 0) { + const size_t n_cols = hex_smin(n, n_chunk_n_cols); + const uint32_t height = is_quant ? (n_cols / 32) * n_k_tiles : n_cols; + dma_queue_push(ctx->dma[0], dma_make_ptr(vtcm_weight_raw[0], weight), dma_dst_stride, dma_src_stride, dma_width_bytes, height); + } + + for (size_t nc = 0; nc < n; nc += n_chunk_n_cols) { + const size_t n_cols = hex_smin(n - nc, n_chunk_n_cols); + const size_t n_row_tiles = hmx_ceil_div(n_rows, HTP_MM_HMX_TILE_N_ROWS); + const size_t n_col_tiles = hmx_ceil_div(n_cols, HTP_MM_HMX_TILE_N_COLS); + + void * curr_raw = dma_queue_pop(ctx->dma[0]).dst; -static inline int hmx_mm_batch_r3(const hmx_mm_f16_f32_batched_params_t *params) { - return params->ne03 > 0 ? params->ne13 / params->ne03 : 1; + dequantize_tiled_weight_chunk_to_fp16_tiles( + ctx, vtcm_scratch0, curr_raw, + n_cols, k, row_stride, weight_type, + n_k_tiles, n_k_tiles_div, dequant_worker_fn, n_threads); + + const size_t nc_next = nc + n_chunk_n_cols; + if (nc_next < n) { + const size_t n_cols_next = hex_smin(n - nc_next, n_chunk_n_cols); + const uint32_t height_next = is_quant ? (n_cols_next / 32) * n_k_tiles : n_cols_next; + dma_queue_push(ctx->dma[0], dma_make_ptr(curr_raw, weight + nc_next * weight_stride), dma_dst_stride, dma_src_stride, dma_width_bytes, height_next); + } + + hmx_matmul_job_init(&job, vtcm_output, vtcm_f16_act, vtcm_scratch0, vtcm_scales, n_row_tiles, n_col_tiles, k / HTP_MM_HMX_TILE_N_ROWS); + hmx_queue_push(ctx->hmx_queue, hmx_queue_make_desc(hmx_matmul_worker_fn, &job)); + hmx_queue_pop(ctx->hmx_queue); + + float *output_chunk = dst_ptr + (mr * dst_stride + nc); + int chunk_dst_cols = dst_cols - (int)nc; + if (chunk_dst_cols > 0) { + transfer_output_chunk_threaded(ctx, output_chunk, NULL, vtcm_output, n_rows, n_cols, dst_stride, 0, chunk_dst_cols, n_threads); + } + } + } + } + } + + return HTP_STATUS_OK; } static inline const __fp16 *hmx_mm_weight_batch_ptr(const hmx_mm_f16_f32_batched_params_t *params, int dst_b2, int dst_b3) { - const int r2 = hmx_mm_batch_r2(params); - const int r3 = hmx_mm_batch_r3(params); + const size_t b2_idx = (params->r2 <= 1) ? (size_t) dst_b2 : (size_t) fastdiv((uint32_t) dst_b2, ¶ms->div_r2); + const size_t b3_idx = (params->r3 <= 1) ? (size_t) dst_b3 : (size_t) fastdiv((uint32_t) dst_b3, ¶ms->div_r3); return (const __fp16 *) ((const uint8_t *) params->weight + - (size_t) (dst_b2 / r2) * params->src0_nb2 + - (size_t) (dst_b3 / r3) * params->src0_nb3); + b2_idx * params->src0_nb2 + + b3_idx * params->src0_nb3); } static inline const float *hmx_mm_activation_batch_ptr(const hmx_mm_f16_f32_batched_params_t *params, @@ -2517,7 +2960,7 @@ static int hmx_mm_f16_f32_batched(struct htp_context *ctx, const hmx_mm_f16_f32_ if (params->k % 32 != 0 || params->n % 32 != 0) { return -1; } if (!hex_is_aligned(params->dst, VLEN) || !hex_is_aligned(params->activation, VLEN)) { return -1; } - const int group_size = hmx_mm_batch_r2(params); + const int group_size = params->r2; const size_t vtcm_budget = ctx->vtcm_size; // Check if the precomputed parameters are grouped or simple. @@ -2825,8 +3268,9 @@ static int hmx_mm_id_2d_f32(struct htp_context *ctx, htp_mm_hmx_get_2d_chunk_costs(weight_type, k, /*pipeline=*/false, aligned_tile_size, &size_per_n, &size_per_m, &size_per_mn); + const size_t overhead = htp_mm_hmx_get_2d_overhead(/*pipeline=*/false, /*is_matmul_id=*/true); size_t m_chunk_n_rows = 0, n_chunk_n_cols = 0; - if (htp_mm_hmx_compute_chunks(vtcm_budget, /*overhead=*/256, size_per_n, size_per_m, size_per_mn, + if (htp_mm_hmx_compute_chunks(vtcm_budget, overhead, size_per_n, size_per_m, size_per_mn, m_padded, n, /*m_block_cost=*/(size_t) n * HTP_MM_HMX_COST_W_DEQUANT, /*n_block_cost=*/(size_t) m_padded * HTP_MM_HMX_COST_A_CONVERT, &m_chunk_n_rows, &n_chunk_n_cols, &vtcm_used)) { @@ -2962,6 +3406,10 @@ static int hmx_mm_op_matmul(struct htp_ops_context * octx, const struct htp_mm_k .dst_nb3 = dst->nb[3], .src2_nb2 = src2_nb2, .src2_nb3 = src2_nb3, + .r2 = (ne02 > 0) ? (ne12 / ne02) : 1, + .r3 = (ne03 > 0) ? (ne13 / ne03) : 1, + .div_r2 = kparams->div_r2, + .div_r3 = kparams->div_r3, }; ret = hmx_mm_f16_f32_batched(octx->ctx, &batch_params, kparams->m_chunk, kparams->n_chunk, @@ -3106,10 +3554,10 @@ static int hvx_mm_matmul_id( mmctx->vtcm_src0_stride = src0_row_size_padded; mmctx->vtcm_src1_stride = src1_row_size; - mmctx->vtcm_src0_size_per_thread = L.src0_bytes / octx->n_threads; + mmctx->vtcm_src0_size_per_thread = fastdiv(L.src0_bytes, &octx->ctx->n_threads_div); mmctx->vtcm_src1_size_per_thread = L.src1_bytes; mmctx->vtcm_src2_size_per_thread = 0; - mmctx->vtcm_dst_size_per_thread = L.dst_bytes / octx->n_threads; + mmctx->vtcm_dst_size_per_thread = fastdiv(L.dst_bytes, &octx->ctx->n_threads_div); mmctx->n_quant_rows_per_thread = (src1_nrows + n_quant_tasks - 1) / n_quant_tasks; mmctx->quant_task_func = quant_task_func; @@ -3123,6 +3571,134 @@ static int hvx_mm_matmul_id( return HTP_STATUS_OK; } +static int hmx_mm_op_matmul_id_nx( + struct htp_ops_context * octx, + struct htp_mm_context * mmctx +) { + const uint32_t * matrix_row_counts = mmctx->matrix_row_counts; + const struct mmid_row_mapping * matrix_rows = mmctx->matrix_rows; + const struct htp_mm_kernel_params * kparams = (const struct htp_mm_kernel_params *) octx->kernel_params; + const uint32_t n_weights = kparams->n_weights; + const struct htp_tensor * restrict src0 = octx->src[0]; + const struct htp_tensor * restrict act = octx->src[n_weights]; + const int n_as = src0->ne[2]; + + for (uint32_t cur_a = 0; cur_a < (uint32_t) n_as; ++cur_a) { + const int32_t cne1 = matrix_row_counts[cur_a]; + if (cne1 == 0) continue; + + for (uint32_t p = 0; p < n_weights; ++p) { + const struct htp_tensor * restrict src_w = octx->src[p]; + const struct htp_tensor * restrict dst = octx->dsts[p]; + if (!src_w || !dst) continue; + + int ret = hmx_mm_id_2d_f32(octx->ctx, (float*) dst->data, (float*) act->data, + (const uint8_t *) src_w->data + cur_a * src_w->nb[2], + cne1, src_w->ne[0], src_w->ne[1], + act->ne[0], + act->ne[1], + act->nb[1], act->nb[2], + dst->nb[1], dst->nb[2], + (int) src_w->nb[1], (int) src_w->type, + matrix_rows, cur_a, mmctx->mapping_stride); + if (ret != 0) { + FARF(ERROR, "HMX matmul ID NX failed for expert %u weight %u, error %d\n", cur_a, p, ret); + return HTP_STATUS_NO_SUPPORT; + } + } + } + + return HTP_STATUS_OK; +} + +static int hvx_mm_matmul_id_nx( + struct htp_ops_context * octx, + struct htp_mm_context * mmctx, + work_queue_func_t hvx_mmid_task_func +) { + const uint32_t src0_row_size_padded = mmctx->src0_row_size_padded; + const uint32_t src1_nrows = mmctx->src1_nrows; + + struct htp_thread_trace * tr = &octx->ctx->trace[0]; + htp_trace_event_start(tr, HTP_TRACE_EVT_INIT, 0); + + const struct htp_mm_kernel_params * kparams = (const struct htp_mm_kernel_params *) octx->kernel_params; + const uint32_t n_weights = kparams->n_weights; + const struct htp_tensor * restrict src0 = octx->src[0]; + const struct htp_tensor * restrict act = octx->src[n_weights]; + const struct htp_tensor * restrict ids = octx->src[n_weights + 1]; + const size_t src0_row_size = src0->nb[1]; + + const uint32_t qk = QK_Q8_0_TILED; + const uint32_t nb = (act->ne[0] + qk - 1) / qk; + const uint32_t total_nb = src1_nrows * nb; + + work_queue_func_t quant_task_func; + uint32_t n_quant_tasks = 1; + if (src1_nrows < octx->n_threads) { + n_quant_tasks = MIN(total_nb, octx->n_threads); + quant_task_func = (src0->type == HTP_TYPE_Q4_1) ? quantize_f32_q8_1_tiled_block : quantize_f32_q8_0_tiled_block; + for (uint32_t ith = 0; ith < n_quant_tasks; ++ith) { + uint32_t ib_first = (total_nb * ith) / n_quant_tasks; + uint32_t ib_last = (total_nb * (ith + 1)) / n_quant_tasks; + mmctx->quant_ib_first[ith] = ib_first; + mmctx->quant_ib_last[ith] = ib_last; + mmctx->quant_r[ith] = ib_first / nb; + mmctx->quant_c[ith] = ib_first % nb; + } + } else { + n_quant_tasks = MIN(src1_nrows, octx->n_threads); + quant_task_func = (src0->type == HTP_TYPE_Q4_1) ? quantize_f32_q8_1_tiled : quantize_f32_q8_0_tiled; + } + size_t src1_row_size = (src0->type == HTP_TYPE_Q4_1) ? htp_mm_q8_1_tiled_row_size(act->ne[0]) : htp_mm_q8_0_tiled_row_size(act->ne[0]); + + struct htp_mm_hvx_vtcm_layout L; + htp_mm_hvx_vtcm_layout_build(&L, kparams->kernel_type, src0->type, act->ne[0], src1_nrows, octx->n_threads, + 0, src0_row_size, src1_row_size, 0, kparams->n_prefetch, true, false); + + size_t vtcm_size = kparams->vtcm_size > 0 ? (size_t)kparams->vtcm_size : L.total_bytes; + + if (octx->ctx->vtcm_size < vtcm_size) { + FARF(ERROR, "matmul-id-nx: current VTCM reservation %zu is too small, needed %zu\n", + octx->ctx->vtcm_size, vtcm_size); + return HTP_STATUS_VTCM_TOO_SMALL; + } + + uint8_t * const base = (uint8_t *) octx->ctx->vtcm_base; + mmctx->vtcm_src0 = VTCM_LAYOUT_PTR(uint8_t, base, L.off_src0); + mmctx->vtcm_src1 = VTCM_LAYOUT_PTR(uint8_t, base, L.off_src1); + mmctx->vtcm_dst = VTCM_LAYOUT_PTR(uint8_t, base, L.off_dst); + + octx->src0_spad.src = NULL; + octx->src1_spad.src = NULL; + octx->src2_spad.src = NULL; + octx->src3_spad.src = NULL; + octx->dst_spad.src = NULL; + + mmctx->vtcm_src0_stride = 0; + mmctx->vtcm_src1_stride = src1_row_size; + + mmctx->vtcm_src0_size_per_thread = fastdiv(L.src0_bytes, &octx->ctx->n_threads_div); + mmctx->vtcm_src1_size_per_thread = L.src1_bytes; + mmctx->vtcm_dst_size_per_thread = fastdiv(L.dst_bytes, &octx->ctx->n_threads_div); + + mmctx->n_quant_rows_per_thread = (src1_nrows + n_quant_tasks - 1) / n_quant_tasks; + mmctx->quant_task_func = quant_task_func; + mmctx->n_quant_tasks = n_quant_tasks; + atomic_init(&mmctx->quant_barrier, n_quant_tasks); + + FARF(HIGH, "matmul-id-nx: src0 %d:%d:%d type %s nrows %u, src1 %d:%d:%d nrows %u, vtcm %zu/%zu, threads %d\n", + src0->ne[0], src0->ne[1], src0->ne[2], mmctx->type, src0->ne[1], + act->ne[0], act->ne[1], act->ne[2], src1_nrows, + L.total_bytes, octx->ctx->vtcm_size, octx->n_threads); + + htp_trace_event_stop(tr, HTP_TRACE_EVT_INIT, 0); + + worker_pool_run_func(octx->ctx->worker_pool, hvx_mmid_task_func, mmctx, octx->n_threads); + + return HTP_STATUS_OK; +} + static inline void scan_expert_ids_n( const struct htp_tensor * ids, const uint32_t n_ids, @@ -3213,7 +3789,7 @@ int op_matmul_id(struct htp_ops_context * octx) { const uint32_t src0_nrows = ne01; // per expert const uint32_t src1_nrows = ne11 * ne12 * ne13; - mmctx->src0_nrows_per_thread = (src0_nrows + octx->n_threads - 1) / octx->n_threads; + mmctx->src0_nrows_per_thread = fastdiv(src0_nrows + octx->n_threads - 1, &octx->ctx->n_threads_div); mmctx->src0_nrows_per_thread = hex_round_up(mmctx->src0_nrows_per_thread, 32); // row groups @@ -3280,11 +3856,103 @@ int op_matmul_id(struct htp_ops_context * octx) { return s; } -int op_matmul_nx(struct htp_ops_context * octx) { + +int op_matmul_id_nx(struct htp_ops_context * octx) { struct htp_thread_trace * tr = &octx->ctx->trace[0]; htp_trace_event_start(tr, HTP_TRACE_EVT_INIT, 0); const struct htp_mm_kernel_params * kparams = (const struct htp_mm_kernel_params *) octx->kernel_params; + const uint32_t n_weights = kparams->n_weights; + const struct htp_tensor * restrict src0 = octx->src[0]; + const struct htp_tensor * restrict act = octx->src[n_weights]; + const struct htp_tensor * restrict ids = octx->src[n_weights + 1]; + + struct htp_mm_context mmctx_struct = {0}; + struct htp_mm_context * mmctx = &mmctx_struct; + mmctx->octx = octx; + mmctx->act = act; + + const size_t src0_row_size = src0->nb[1]; + const size_t src0_row_size_padded = hex_round_up(src0_row_size, 128); + + const uint32_t src0_nrows = src0->ne[1]; + const uint32_t src1_nrows = act->ne[1] * act->ne[2] * act->ne[3]; + + mmctx->src0_nrows_per_thread = fastdiv(src0_nrows + octx->n_threads - 1, &octx->ctx->n_threads_div); + mmctx->src0_nrows_per_thread = hex_round_up(mmctx->src0_nrows_per_thread, 32); + + const int n_ids = ids->ne[0]; + const int n_as = src0->ne[2]; + + uint8_t * mapping_buf = octx->ctx->ddr_spad_base; + uint32_t mapping_stride = 1; + uint32_t * matrix_row_counts = (uint32_t *) mapping_buf; + struct mmid_row_mapping * matrix_rows = NULL; + + if (src1_nrows > 1) { + const size_t matrix_row_counts_size = n_as * sizeof(uint32_t); + assert(octx->ctx->ddr_spad_size >= matrix_row_counts_size); + + hex_l2fetch_block((const void *) ids->data, ids->ne[1] * ids->nb[1]); + + memset(matrix_row_counts, 0, matrix_row_counts_size); + scan_expert_ids(ids, n_ids, n_as, matrix_row_counts, NULL, 0); + + uint32_t max_count = hvx_reduce_max_i32((const uint8_t *) matrix_row_counts, n_as); + mapping_stride = max_count > 0 ? max_count : 1; + + size_t matrix_row_map_size = n_as * mapping_stride * sizeof(struct mmid_row_mapping); + const size_t total_map_size = matrix_row_counts_size + matrix_row_map_size; + + if (total_map_size > octx->ctx->ddr_spad_size) { + mapping_buf = memalign(128, total_map_size); + if (!mapping_buf) { + return HTP_STATUS_INTERNAL_ERR; + } + } + + matrix_row_counts = (uint32_t *) mapping_buf; + matrix_rows = (struct mmid_row_mapping *) (mapping_buf + matrix_row_counts_size); + + memset(matrix_row_counts, 0, n_as * sizeof(uint32_t)); + scan_expert_ids(ids, n_ids, n_as, matrix_row_counts, matrix_rows, mapping_stride); + } + + mmctx->matrix_row_counts = matrix_row_counts; + mmctx->matrix_rows = matrix_rows; + mmctx->mapping_stride = mapping_stride; + mmctx->mm_div_ne11 = kparams->div_ne11; + mmctx->src0_row_size_padded = src0_row_size_padded; + mmctx->src1_nrows = src1_nrows; + + htp_trace_event_stop(tr, HTP_TRACE_EVT_INIT, 0); + + int s; + if (kparams->n_hmx) { + s = hmx_mm_op_matmul_id_nx(octx, mmctx); + } else { + if (hvx_mm_init_vec_dot(mmctx, src0->type) == 0) { + s = hvx_mm_matmul_id_nx(octx, mmctx, src1_nrows > 1 ? hvx_mm_id_nx : hvx_mv_id_nx); + } else { + s = HTP_STATUS_NO_SUPPORT; + } + } + + if (mapping_buf != octx->ctx->ddr_spad_base) { + free(mapping_buf); + } + + return s; +} +int op_matmul_nx(struct htp_ops_context * octx) { + const struct htp_mm_kernel_params * kparams = (const struct htp_mm_kernel_params *) octx->kernel_params; + if (kparams->n_hmx) { + return hmx_mm_nx_2d_f32(octx, kparams); + } + + struct htp_thread_trace * tr = &octx->ctx->trace[0]; + htp_trace_event_start(tr, HTP_TRACE_EVT_INIT, 0); + const uint32_t n_weights = kparams->n_weights; const struct htp_tensor * restrict src0 = octx->src[0]; // first weight @@ -3366,9 +4034,9 @@ int op_matmul_nx(struct htp_ops_context * octx) { mmctx->vtcm_src0_stride = is_repacked ? 0 : src0_row_size_padded; mmctx->vtcm_src1_stride = src1_row_size; - mmctx->vtcm_src0_size_per_thread = L.src0_bytes / octx->n_threads; + mmctx->vtcm_src0_size_per_thread = fastdiv(L.src0_bytes, &octx->ctx->n_threads_div); mmctx->vtcm_src1_size_per_thread = L.src1_bytes; - mmctx->vtcm_dst_size_per_thread = L.dst_bytes / octx->n_threads; + mmctx->vtcm_dst_size_per_thread = fastdiv(L.dst_bytes, &octx->ctx->n_threads_div); mmctx->n_quant_rows_per_thread = (src1_nrows + n_quant_tasks - 1) / n_quant_tasks; mmctx->quant_task_func = quant_task_func; diff --git a/ggml/src/ggml-hexagon/htp/matmul-ops.h b/ggml/src/ggml-hexagon/htp/matmul-ops.h index dbc8e359093..2dbcb0c2e51 100644 --- a/ggml/src/ggml-hexagon/htp/matmul-ops.h +++ b/ggml/src/ggml-hexagon/htp/matmul-ops.h @@ -134,7 +134,8 @@ static inline int htp_mm_hmx_compute_chunks(size_t vtcm_total, size_t best_mn = 0; size_t best_m = 0, best_n = 0; - const size_t n_max = hex_align_down((size_t)n, HTP_MM_HMX_TILE_N_COLS); + const size_t max_nc_budget = (usable / per_n_cost); + const size_t n_max = hex_align_down(hex_smin((size_t)n, max_nc_budget), HTP_MM_HMX_TILE_N_COLS); for (size_t nc = n_max; nc >= HTP_MM_HMX_TILE_N_COLS; nc -= HTP_MM_HMX_TILE_N_COLS) { size_t n_fixed = 0, ncmn = 0, mc_denom = 0; if (hex_mul_overflow(nc, per_n_cost, &n_fixed)) continue; @@ -299,6 +300,15 @@ static inline void htp_mm_hmx_get_batched_chunk_costs( *size_per_mn_out = sizeof(uint16_t); } +static inline size_t htp_mm_hmx_get_2d_overhead(bool pipeline, bool is_matmul_id) { + size_t num_regions = pipeline ? 7 : (is_matmul_id ? 4 : 5); + return num_regions * HTP_MM_HMX_TILE_SIZE + 256; +} + +static inline size_t htp_mm_hmx_get_batched_overhead(void) { + return 5 * HTP_MM_HMX_TILE_SIZE + 256; +} + struct htp_mm_hmx_vtcm_layout { // Byte offsets from vtcm_base for each region size_t off_weight[2]; // [1] is only used when pipelined @@ -568,10 +578,8 @@ static inline void htp_mm_hvx_vtcm_layout_build( } size_t quant_scratch_size_per_thread = htp_mm_round_up(ne10 * sizeof(float), QK_Q8_0_TILED * sizeof(float)); - size_t dst_size_per_thread = dst_nrows > 0 ? htp_mm_round_up(dst_row_size, 128) : 0; - if (dst_size_per_thread < quant_scratch_size_per_thread) { - dst_size_per_thread = quant_scratch_size_per_thread; - } + size_t dst_slice_per_thread = (dst_nrows > 0 && src1_nrows == 1) ? htp_mm_round_up((dst_row_size + n_threads - 1) / n_threads, 128) : 0; + size_t dst_size_per_thread = (dst_slice_per_thread > quant_scratch_size_per_thread) ? dst_slice_per_thread : quant_scratch_size_per_thread; dst_sz = dst_size_per_thread * n_threads; break; } @@ -592,10 +600,8 @@ static inline void htp_mm_hvx_vtcm_layout_build( } size_t quant_scratch_size_per_thread = htp_mm_round_up(ne10 * sizeof(float), QK_Q8_0_TILED * sizeof(float)); - size_t dst_size_per_thread = dst_nrows > 0 ? htp_mm_round_up(dst_row_size, 128) : 0; - if (dst_size_per_thread < quant_scratch_size_per_thread) { - dst_size_per_thread = quant_scratch_size_per_thread; - } + size_t dst_slice_per_thread = dst_nrows > 0 ? htp_mm_round_up((dst_row_size + n_threads - 1) / n_threads, 128) : 0; + size_t dst_size_per_thread = (dst_slice_per_thread > quant_scratch_size_per_thread) ? dst_slice_per_thread : quant_scratch_size_per_thread; dst_sz = dst_size_per_thread * n_threads; break; } @@ -658,7 +664,7 @@ static inline bool htp_mm_hmx_solve_batched_params( int act_threads = n_threads; while (act_threads >= 1) { - size_t group_overhead = 256; + size_t group_overhead = htp_mm_hmx_get_batched_overhead(); size_t group_size_per_n, group_size_per_m, group_size_per_mn; htp_mm_hmx_get_batched_chunk_costs(k, group_size, &group_size_per_n, &group_size_per_m, &group_size_per_mn); @@ -725,7 +731,7 @@ static inline bool htp_mm_hmx_solve_2d_params( int act_threads = n_threads; while (act_threads >= 1) { - size_t simple_2d_overhead = 256; + size_t simple_2d_overhead = htp_mm_hmx_get_2d_overhead(pipeline, is_matmul_id); size_t simple_2d_size_per_n, simple_2d_size_per_m, simple_2d_size_per_mn; htp_mm_hmx_get_2d_chunk_costs(wtype, k, pipeline, aligned_tile_size, &simple_2d_size_per_n, &simple_2d_size_per_m, &simple_2d_size_per_mn); From c2b400754bac2de1f4f74bfac9eb33474daa631e Mon Sep 17 00:00:00 2001 From: "Aman Chadha(IVIXMMI)" <79802170+ac-mmi@users.noreply.github.com> Date: Wed, 2 Sep 2026 11:46:15 +0530 Subject: [PATCH 080/104] ggml: avoid KleidiAI buffer type init on dispatch (llama/27891) Co-authored-by: Acmmi --- ggml/src/ggml-cpu/kleidiai/kleidiai.cpp | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/ggml/src/ggml-cpu/kleidiai/kleidiai.cpp b/ggml/src/ggml-cpu/kleidiai/kleidiai.cpp index 92d7fd644f7..dbd19878077 100644 --- a/ggml/src/ggml-cpu/kleidiai/kleidiai.cpp +++ b/ggml/src/ggml-cpu/kleidiai/kleidiai.cpp @@ -1823,7 +1823,7 @@ class extra_buffer_type : ggml::cpu::extra_buffer_type { const bool src0_is_kleidiai = op->src[0]->buffer && (ggml_n_dims(op->src[0]) == 2) && - op->src[0]->buffer->buft == ggml_backend_cpu_kleidiai_buffer_type() && + op->src[0]->buffer->buft->context == this && slot_total > 0; if ((op->op == GGML_OP_MUL_MAT || op->op == GGML_OP_GET_ROWS) && @@ -1862,7 +1862,7 @@ class extra_buffer_type : ggml::cpu::extra_buffer_type { ggml::cpu::tensor_traits * get_tensor_traits(const struct ggml_tensor * op) override { if (op->op == GGML_OP_MUL_MAT || op->op == GGML_OP_GET_ROWS) { - if (op->src[0]->buffer && op->src[0]->buffer->buft == ggml_backend_cpu_kleidiai_buffer_type()) { + if (op->src[0]->buffer && op->src[0]->buffer->buft->context == this) { return (ggml::cpu::tensor_traits *) op->src[0]->extra; } else { // KleidiAI only has kernels for Q4_0 and Q8_0. For a quantized weight of any From 519df618df430b9288cb62de6eb4e5568965526c Mon Sep 17 00:00:00 2001 From: Aman Gupta Date: Wed, 2 Sep 2026 19:57:37 +0530 Subject: [PATCH 081/104] CUDA + ggml: add sparse-fa for DSV4/GLM (llama/27970) --- ggml/include/ggml.h | 6 + ggml/src/ggml-cuda/fattn-common.cuh | 23 ++- ggml/src/ggml-cuda/fattn-mma-f16.cuh | 223 ++++++++++++++++++--------- ggml/src/ggml-cuda/fattn-tile.cuh | 12 +- ggml/src/ggml-cuda/fattn-vec.cuh | 2 +- ggml/src/ggml-cuda/fattn.cu | 133 ++++++++++++++++ ggml/src/ggml.c | 9 ++ 7 files changed, 323 insertions(+), 85 deletions(-) diff --git a/ggml/include/ggml.h b/ggml/include/ggml.h index 26f31232f67..b88b7e54a5e 100644 --- a/ggml/include/ggml.h +++ b/ggml/include/ggml.h @@ -2453,6 +2453,12 @@ extern "C" { GGML_API enum ggml_prec ggml_flash_attn_ext_get_prec( const struct ggml_tensor * a); + // Use finite mask entries as a sparse K/V set. Set 0 to disable. + // n_kv_max must bound the number of finite entries in every mask row. + GGML_API void ggml_flash_attn_ext_set_n_kv_max( + struct ggml_tensor * a, + int32_t n_kv_max); + GGML_API void ggml_flash_attn_ext_add_sinks( struct ggml_tensor * a, struct ggml_tensor * sinks); diff --git a/ggml/src/ggml-cuda/fattn-common.cuh b/ggml/src/ggml-cuda/fattn-common.cuh index e67cc7fdf78..7442bc22af2 100644 --- a/ggml/src/ggml-cuda/fattn-common.cuh +++ b/ggml/src/ggml-cuda/fattn-common.cuh @@ -718,6 +718,9 @@ static __global__ void flash_attn_mask_to_KV_max( KV_max[sequence*ne31 + jt] = KV_max_sj; } +void ggml_cuda_flash_attn_ext_compact_mask( + const ggml_tensor * mask, int32_t * indices, int32_t n_kv_max, cudaStream_t stream); + template // D == head size __launch_bounds__(D, 1) static __global__ void flash_attn_stream_k_fixup_uniform( @@ -972,7 +975,8 @@ static __global__ void flash_attn_combine_results( template void launch_fattn( ggml_backend_cuda_context & ctx, ggml_tensor * dst, fattn_kernel_t fattn_kernel, const int nwarps, const size_t nbytes_shared, - const int nbatch_fa, const bool need_f16_K, const bool need_f16_V, const bool stream_k, const int warp_size = WARP_SIZE + const int nbatch_fa, const bool need_f16_K, const bool need_f16_V, const bool stream_k, const bool use_sparse, + const int warp_size = WARP_SIZE ) { constexpr int ncols = ncols1 * ncols2; @@ -1088,10 +1092,20 @@ void launch_fattn( const int ntiles_z_gqa = ((gqa_ratio + ncols2 - 1) / ncols2); const int ntiles_dst = ntiles_x * ntiles_z_gqa * K->ne[2] * Q->ne[3]; + const int32_t n_kv_max = use_sparse ? ggml_get_op_params_i32(KQV, 4) : 0; + if (use_sparse) { + GGML_ASSERT(mask != nullptr); + GGML_ASSERT(n_kv_max > 0); + const size_t mask_rows = size_t(mask->ne[1]) * mask->ne[3]; + + KV_max.alloc(size_t(n_kv_max) * mask_rows); + ggml_cuda_flash_attn_ext_compact_mask(mask, KV_max.ptr, n_kv_max, main_stream); + } + // Optional optimization where the mask is scanned to determine whether part of the calculation can be skipped. // Only worth the overhead if there is at lease one FATTN_KQ_STRIDE x FATTN_KQ_STRIDE square to be skipped or // multiple sequences of possibly different lengths. - if (mask && K->ne[1] % FATTN_KQ_STRIDE == 0 && (Q->ne[1] >= 1024 || Q->ne[3] > 1)) { + if (!use_sparse && mask && K->ne[1] % FATTN_KQ_STRIDE == 0 && (Q->ne[1] >= 1024 || Q->ne[3] > 1)) { const int64_t s31 = mask->nb[1] / sizeof(half2); const int64_t s33 = mask->nb[3] / sizeof(half2); @@ -1114,7 +1128,8 @@ void launch_fattn( GGML_ASSERT(max_blocks_per_sm > 0); int parallel_blocks = max_blocks_per_sm; - const int ntiles_KV = (K->ne[1] + nbatch_fa - 1) / nbatch_fa; // Max. number of parallel blocks limited by KV cache length. + const int64_t n_kv = use_sparse ? n_kv_max : K->ne[1]; + const int ntiles_KV = (n_kv + nbatch_fa - 1) / nbatch_fa; // Max. number of parallel blocks limited by KV cache length. dim3 blocks_num; if (stream_k) { @@ -1218,7 +1233,7 @@ void launch_fattn( !stream_k && parallel_blocks > 1 ? dst_tmp.ptr : (float *) KQV->data, dst_tmp_meta.ptr, scale, max_bias, m0, m1, n_head_log2, logit_softcap, Q->ne[0], ne01, Q->ne[2], Q->ne[3], Q->nb[1], Q->nb[2], Q->nb[3], - K->ne[0], K->ne[1], K->ne[2], K->ne[3], nb11, nb12, nb13, + K->ne[0], n_kv, K->ne[2], K->ne[3], nb11, nb12, nb13, nb21, nb22, nb23, mask ? mask->ne[1] : 0, mask ? mask->ne[2] : 0, mask ? mask->ne[3] : 0, mask ? mask->nb[1] : 0, mask ? mask->nb[2] : 0, mask ? mask->nb[3] : 0 diff --git a/ggml/src/ggml-cuda/fattn-mma-f16.cuh b/ggml/src/ggml-cuda/fattn-mma-f16.cuh index 387e70fa149..126a4c4529b 100644 --- a/ggml/src/ggml-cuda/fattn-mma-f16.cuh +++ b/ggml/src/ggml-cuda/fattn-mma-f16.cuh @@ -350,20 +350,24 @@ static __host__ int ggml_cuda_fattn_mma_get_nstages(const int DKQ, const int DV, return cp_async_available(cc) && ncols2 >= 2 ? ggml_cuda_fattn_mma_get_nstages_target(DKQ, DV, ncols1*ncols2, cc) : 0; } -static constexpr __device__ int ggml_cuda_fattn_mma_get_nstages(const int DKQ, const int DV, const int ncols1, const int ncols2) { +static constexpr __device__ int ggml_cuda_fattn_mma_get_nstages( + const int DKQ, const int DV, const int ncols1, const int ncols2, const bool use_sparse) { #ifdef CP_ASYNC_AVAILABLE - return ncols2 >= 2 ? ggml_cuda_fattn_mma_get_nstages_target(DKQ, DV, ncols1*ncols2) : 0; + const int nstages_target = ncols2 >= 2 ? ggml_cuda_fattn_mma_get_nstages_target(DKQ, DV, ncols1*ncols2) : 0; + // sparse gather is not implemented for multi-stage loading + return use_sparse && nstages_target > 1 ? 1 : nstages_target; #else - GGML_UNUSED_VARS(DKQ, DV, ncols1, ncols2); + GGML_UNUSED_VARS(DKQ, DV, ncols1, ncols2, use_sparse); return 0; #endif // CP_ASYNC_AVAILABLE } // ------------------------------------------------------------------------------------------------------------------ -template +template static __device__ __forceinline__ void flash_attn_ext_f16_load_tile( - const half2 * const __restrict__ KV, half2 * const __restrict__ tile_KV, const int D2, const int stride_KV, const int i_sup) { + const half2 * const __restrict__ KV, half2 * const __restrict__ tile_KV, const int D2, const int stride_KV, + const int k_VKQ_0, const int i_sup, const int32_t * const __restrict__ indices) { constexpr int warp_size = ggml_cuda_get_physical_warp_size(); // K/V data is loaded with decreasing granularity for D for better memory bandwidth. // The minimum granularity is 16 bytes. @@ -371,7 +375,7 @@ static __device__ __forceinline__ void flash_attn_ext_f16_load_tile( const int chunks_per_row = D2 / h2_per_chunk; if constexpr (use_cp_async) { static_assert(warp_size == 32, "bad warp_size"); - static_assert(!oob_check, "OOB check not compatible with cp_async"); + static_assert(!oob_check || use_sparse, "OOB check not compatible with cp_async"); constexpr int preload = 64; const unsigned int tile_KV_32 = ggml_cuda_cvta_generic_to_shared(tile_KV); @@ -394,15 +398,24 @@ static __device__ __forceinline__ void flash_attn_ext_f16_load_tile( break; } + int64_t i_KV; + if constexpr (use_sparse) { + // padded slots gather row 0, the -inf mask removes their contribution + const int32_t index = i < i_sup ? indices[k_VKQ_0 + i] : 0; + i_KV = index >= 0 ? index : 0; + } else { + i_KV = k_VKQ_0 + i; + } + #pragma unroll for (int k0 = k0_start; k0 < k0_stop; k0 += stride_k) { const int k = k0 + (stride_k == warp_size ? threadIdx.x : threadIdx.x % stride_k); if constexpr (swz) { const int smem_offs_b = ggml_cuda_fattn_smem_swizzle::bytes_rc(i, k*h2_per_chunk); - cp_async_cg_16(tile_KV_32 + smem_offs_b, KV + i*stride_KV + k*h2_per_chunk); + cp_async_cg_16(tile_KV_32 + smem_offs_b, KV + i_KV*stride_KV + k*h2_per_chunk); } else { - cp_async_cg_16(tile_KV_32 + i*(stride_tile*sizeof(half2)) + k*16, KV + i*stride_KV + k*h2_per_chunk); + cp_async_cg_16(tile_KV_32 + i*(stride_tile*sizeof(half2)) + k*16, KV + i_KV*stride_KV + k*h2_per_chunk); } } } @@ -438,12 +451,17 @@ static __device__ __forceinline__ void flash_attn_ext_f16_load_tile( for (int k0 = k0_start; k0 < k0_stop; k0 += stride_k) { const int k = k0 + (stride_k == warp_size ? threadIdx.x : threadIdx.x % stride_k); + const half2 * src; + if constexpr (use_sparse) { + const int32_t index = i < i_sup ? indices[k_VKQ_0 + i] : -1; + src = index >= 0 ? KV + int64_t(index)*stride_KV + k*h2_per_chunk : zero; + } else { + src = !oob_check || i < i_sup ? KV + int64_t(k_VKQ_0 + i)*stride_KV + k*h2_per_chunk : zero; + } if constexpr (swz) { - ggml_cuda_memcpy_1<16>((char *) tile_KV + ggml_cuda_fattn_smem_swizzle::bytes_rc(i, k*h2_per_chunk), - !oob_check || i < i_sup ? KV + i*stride_KV + k*h2_per_chunk : zero); + ggml_cuda_memcpy_1<16>((char *) tile_KV + ggml_cuda_fattn_smem_swizzle::bytes_rc(i, k*h2_per_chunk), src); } else { - ggml_cuda_memcpy_1<16>(tile_KV + i*stride_tile + k*4, - !oob_check || i < i_sup ? KV + i*stride_KV + k*h2_per_chunk : zero); + ggml_cuda_memcpy_1<16>(tile_KV + i*stride_tile + k*4, src); } } } @@ -458,14 +476,16 @@ static __device__ __forceinline__ void flash_attn_ext_f16_load_tile( } } -template +template static __device__ __forceinline__ void flash_attn_ext_f16_load_mask( const half * const __restrict__ mask_h, half * const __restrict__ tile_mask, - const int stride_mask, const int i_sup, const int j0, const uint3 ne01) { + const int stride_mask, const int k_VKQ_0, const int i_sup, const int j0, const uint3 ne01, + const int32_t * const __restrict__ indices) { constexpr int warp_size = ggml_cuda_get_physical_warp_size(); if constexpr (use_cp_async) { static_assert(nbatch_fa <= 8*warp_size && nbatch_fa % 8 == 0, "bad nbatch_fa"); static_assert(!oob_check, "OOB check incompatible with cp_async"); + static_assert(!use_sparse, "sparse gather incompatible with cp_async"); constexpr int preload = nbatch_fa >= 32 ? nbatch_fa * sizeof(half) : 64; constexpr int cols_per_warp = 8*warp_size/nbatch_fa; constexpr int stride_j = nwarps * cols_per_warp; @@ -483,9 +503,9 @@ static __device__ __forceinline__ void flash_attn_ext_f16_load_mask( const int i = 8 * (threadIdx.x % (nbatch_fa/8)); - cp_async_cg_16(tile_mask_32 + j_sram*(nbatch_fa*sizeof(half) + 16) + i*sizeof(half), mask_h + int64_t(j_vram)*stride_mask + i); + cp_async_cg_16(tile_mask_32 + j_sram*(nbatch_fa*sizeof(half) + 16) + i*sizeof(half), mask_h + int64_t(j_vram)*stride_mask + k_VKQ_0 + i); } - } else if constexpr (oob_check) { + } else if constexpr (oob_check || use_sparse) { #pragma unroll for (int j1 = 0; j1 < ncols1; j1 += nwarps) { const int j_sram = j1 + threadIdx.y; @@ -499,7 +519,12 @@ static __device__ __forceinline__ void flash_attn_ext_f16_load_mask( for (int i0 = 0; i0 < nbatch_fa; i0 += warp_size) { const int i = i0 + threadIdx.x; - tile_mask[j_sram*(nbatch_fa + 8) + i] = i < i_sup ? mask_h[int64_t(j_vram)*stride_mask + i] : half(0.0f); + if constexpr (use_sparse) { + const int32_t index = i < i_sup ? indices[k_VKQ_0 + i] : -1; + tile_mask[j_sram*(nbatch_fa + 8) + i] = index >= 0 ? mask_h[int64_t(j_vram)*stride_mask + index] : half(-INFINITY); + } else { + tile_mask[j_sram*(nbatch_fa + 8) + i] = i < i_sup ? mask_h[int64_t(j_vram)*stride_mask + k_VKQ_0 + i] : half(0.0f); + } } } } else if constexpr (nbatch_fa < 2*warp_size) { @@ -516,7 +541,7 @@ static __device__ __forceinline__ void flash_attn_ext_f16_load_mask( const int i = threadIdx.x % (warp_size/cols_per_warp); - ggml_cuda_memcpy_1(tile_mask + j_sram*(nbatch_fa + 8) + 2*i, mask_h + int64_t(j_vram)*stride_mask + 2*i); + ggml_cuda_memcpy_1(tile_mask + j_sram*(nbatch_fa + 8) + 2*i, mask_h + int64_t(j_vram)*stride_mask + k_VKQ_0 + 2*i); } } else { #pragma unroll @@ -532,20 +557,21 @@ static __device__ __forceinline__ void flash_attn_ext_f16_load_mask( for (int i0 = 0; i0 < nbatch_fa; i0 += 2*warp_size) { const int i = i0 + 2*threadIdx.x; - ggml_cuda_memcpy_1(tile_mask + j_sram*(nbatch_fa + 8) + i, mask_h + int64_t(j_vram)*stride_mask + i); + ggml_cuda_memcpy_1(tile_mask + j_sram*(nbatch_fa + 8) + i, mask_h + int64_t(j_vram)*stride_mask + k_VKQ_0 + i); } } } } template static __device__ __forceinline__ void flash_attn_ext_f16_iter( const float2 * const __restrict__ Q_f2, const half2 * const __restrict__ K_h2, const half2 * const __restrict__ V_h2, const half * const __restrict__ mask_h, + const int32_t * const __restrict__ indices, float2 * const __restrict__ dstk, float2 * const __restrict__ dstk_fixup, const float scale, @@ -577,7 +603,7 @@ static __device__ __forceinline__ void flash_attn_ext_f16_iter( constexpr int nbatch_K2 = ggml_cuda_fattn_mma_get_nbatch_K2(DKQ, DV, ncols); constexpr int nbatch_V2 = ggml_cuda_fattn_mma_get_nbatch_V2(DKQ, DV, ncols); constexpr bool Q_in_reg = ggml_cuda_fattn_mma_get_Q_in_reg (DKQ, DV, ncols); - constexpr int nstages = ggml_cuda_fattn_mma_get_nstages (DKQ, DV, ncols1, ncols2); + constexpr int nstages = ggml_cuda_fattn_mma_get_nstages (DKQ, DV, ncols1, ncols2, use_sparse); // swizzle the tile stride for K and V based on the batch size. constexpr int stride_tile_K = ggml_cuda_fattn_smem_swizzle::tile_stride(nbatch_K2); @@ -601,13 +627,14 @@ static __device__ __forceinline__ void flash_attn_ext_f16_iter( constexpr bool use_cp_async = true; cp_async_wait_all(); __syncthreads(); - flash_attn_ext_f16_load_tile - (V_h2 + int64_t(k_VKQ_0)*stride_V, tile_V, nbatch_V2, stride_V, k_VKQ_sup); + flash_attn_ext_f16_load_tile + (V_h2, tile_V, nbatch_V2, stride_V, k_VKQ_0, k_VKQ_sup, nullptr); } else { - constexpr bool use_cp_async = nstages == 1; + // the sparse mask values are gathered per element, always load them synchronously + constexpr bool use_cp_async = nstages == 1 && !use_sparse; if (ncols2 > 1 || mask_h) { - flash_attn_ext_f16_load_mask - (mask_h + k_VKQ_0, tile_mask, stride_mask, k_VKQ_sup, jt*ncols1, ne01); + flash_attn_ext_f16_load_mask + (mask_h, tile_mask, stride_mask, k_VKQ_0, k_VKQ_sup, jt*ncols1, ne01, indices); } } @@ -620,8 +647,8 @@ static __device__ __forceinline__ void flash_attn_ext_f16_iter( if constexpr (nstages <= 1) { const int k0_diff = k0_stop - k0_start; constexpr bool use_cp_async = nstages == 1; - flash_attn_ext_f16_load_tile - (K_h2 + int64_t(k_VKQ_0)*stride_K + k0_start, tile_K, k0_diff, stride_K, k_VKQ_sup); + flash_attn_ext_f16_load_tile + (K_h2 + k0_start, tile_K, k0_diff, stride_K, k_VKQ_0, k_VKQ_sup, indices); if (use_cp_async) { cp_async_wait_all(); } @@ -946,6 +973,7 @@ static __device__ __forceinline__ void flash_attn_ext_f16_iter( } if constexpr (nstages > 1) { + static_assert(!use_sparse, "sparse gather not implemented for multi-stage loading"); static_assert(!V_is_K_view, "K data reuse not implemented multi-stage loading"); // Preload K tile for next iteration: constexpr bool use_cp_async = true; @@ -953,11 +981,11 @@ static __device__ __forceinline__ void flash_attn_ext_f16_iter( __syncthreads(); if (!last_iter) { if (ncols2 > 1 || mask_h) { - flash_attn_ext_f16_load_mask - (mask_h + k_VKQ_0 + nbatch_fa, tile_mask, stride_mask, k_VKQ_sup, jt*ncols1, ne01); + flash_attn_ext_f16_load_mask + (mask_h, tile_mask, stride_mask, k_VKQ_0 + nbatch_fa, k_VKQ_sup, jt*ncols1, ne01, nullptr); } - flash_attn_ext_f16_load_tile - (K_h2 + int64_t(k_VKQ_0 + nbatch_fa)*stride_K, tile_K, nbatch_K2, stride_K, k_VKQ_sup); + flash_attn_ext_f16_load_tile + (K_h2, tile_K, nbatch_K2, stride_K, k_VKQ_0 + nbatch_fa, k_VKQ_sup, nullptr); } } @@ -972,8 +1000,8 @@ static __device__ __forceinline__ void flash_attn_ext_f16_iter( const int i0_diff = i0_stop - i0_start; if (!V_is_K_view || i0_stop > 2*nbatch_K2) { constexpr bool use_cp_async = nstages == 1; - flash_attn_ext_f16_load_tile - (V_h2 + int64_t(k_VKQ_0)*stride_V + i0_start/2, tile_V, i0_diff/2, stride_V, k_VKQ_sup); + flash_attn_ext_f16_load_tile + (V_h2 + i0_start/2, tile_V, i0_diff/2, stride_V, k_VKQ_0, k_VKQ_sup, indices); if (use_cp_async) { cp_async_wait_all(); } @@ -1028,7 +1056,7 @@ static __device__ __forceinline__ void flash_attn_ext_f16_iter( } } #else - GGML_UNUSED_VARS(Q_f2, K_h2, V_h2, mask_h, dstk, dstk_fixup, + GGML_UNUSED_VARS(Q_f2, K_h2, V_h2, mask_h, indices, dstk, dstk_fixup, scale, slope, logit_softcap, ne01, ne02, stride_K, stride_V, stride_mask, tile_Q, tile_K, tile_V, tile_mask, @@ -1126,12 +1154,13 @@ template struct mma_tile_sizes { }; #endif // defined(TURING_MMA_AVAILABLE) -template +template static __device__ __forceinline__ void flash_attn_ext_f16_process_tile( const float2 * const __restrict__ Q_f2, const half2 * const __restrict__ K_h2, const half2 * const __restrict__ V_h2, const half * const __restrict__ mask_h, + const int32_t * const __restrict__ indices, const float * const __restrict__ sinks_f, float2 * const __restrict__ dstk, float2 * const __restrict__ dstk_fixup, @@ -1171,7 +1200,7 @@ static __device__ __forceinline__ void flash_attn_ext_f16_process_tile( constexpr int nbatch_V2 = ggml_cuda_fattn_mma_get_nbatch_V2 (DKQ, DV, ncols); constexpr int nbatch_combine = ggml_cuda_fattn_mma_get_nbatch_combine(DKQ, DV, ncols); constexpr bool Q_in_reg = ggml_cuda_fattn_mma_get_Q_in_reg (DKQ, DV, ncols); - constexpr int nstages = ggml_cuda_fattn_mma_get_nstages (DKQ, DV, ncols1, ncols2); + constexpr int nstages = ggml_cuda_fattn_mma_get_nstages (DKQ, DV, ncols1, ncols2, use_sparse); if (cols_per_warp > ncols) { NO_DEVICE_CODE; @@ -1272,37 +1301,38 @@ static __device__ __forceinline__ void flash_attn_ext_f16_process_tile( // Preload mask and K data for first iteration when using cp_async with multiple stages: if constexpr (nstages > 1) { + static_assert(!use_sparse, "sparse gather not implemented for multi-stage loading"); static_assert(nbatch_K2 == DKQ/2, "batching not implemented for multi-stage pipeline"); constexpr bool use_cp_async = true; constexpr bool oob_check = false; constexpr int k_VKQ_sup = nbatch_fa; if (ncols2 > 1 || mask_h) { - flash_attn_ext_f16_load_mask - (mask_h + kb0*nbatch_fa, tile_mask, stride_mask, k_VKQ_sup, jt*ncols1, ne01); + flash_attn_ext_f16_load_mask + (mask_h, tile_mask, stride_mask, kb0*nbatch_fa, k_VKQ_sup, jt*ncols1, ne01, nullptr); } - flash_attn_ext_f16_load_tile - (K_h2 + int64_t(kb0)*nbatch_fa*stride_K, tile_K, nbatch_K2, stride_K, k_VKQ_sup); + flash_attn_ext_f16_load_tile + (K_h2, tile_K, nbatch_K2, stride_K, kb0*nbatch_fa, k_VKQ_sup, nullptr); } // kb0_start is always < kb0_stop so the last iter can be executed unconditionally. - if constexpr (ncols2 == 1) { + if constexpr (ncols2 == 1 || use_sparse) { constexpr bool oob_check = true; for (; kb0 < kb0_stop-1; ++kb0) { constexpr bool last_iter = false; constexpr int k_VKQ_sup = nbatch_fa; flash_attn_ext_f16_iter - - (Q_f2, K_h2, V_h2, mask_h, dstk, dstk_fixup, scale, slope, logit_softcap, + (Q_f2, K_h2, V_h2, mask_h, indices, dstk, dstk_fixup, scale, slope, logit_softcap, ne01, ne02, stride_K, stride_V, stride_mask, tile_Q, tile_K, tile_V, tile_mask, Q_B, VKQ_C, KQ_max, KQ_rowsum, jt, kb0, k_VKQ_sup); } constexpr bool last_iter = true; const int k_VKQ_sup = ne11 - kb0*nbatch_fa; flash_attn_ext_f16_iter - - (Q_f2, K_h2, V_h2, mask_h, dstk, dstk_fixup, scale, slope, logit_softcap, + (Q_f2, K_h2, V_h2, mask_h, indices, dstk, dstk_fixup, scale, slope, logit_softcap, ne01, ne02, stride_K, stride_V, stride_mask, tile_Q, tile_K, tile_V, tile_mask, Q_B, VKQ_C, KQ_max, KQ_rowsum, jt, kb0, k_VKQ_sup); } else { @@ -1311,18 +1341,18 @@ static __device__ __forceinline__ void flash_attn_ext_f16_process_tile( constexpr bool last_iter = false; constexpr int k_VKQ_sup = nbatch_fa; flash_attn_ext_f16_iter - - (Q_f2, K_h2, V_h2, mask_h, dstk, dstk_fixup, scale, slope, logit_softcap, + (Q_f2, K_h2, V_h2, mask_h, indices, dstk, dstk_fixup, scale, slope, logit_softcap, ne01, ne02, stride_K, stride_V, stride_mask, tile_Q, tile_K, tile_V, tile_mask, Q_B, VKQ_C, KQ_max, KQ_rowsum, jt, kb0, k_VKQ_sup); } constexpr bool last_iter = true; constexpr int k_VKQ_sup = nbatch_fa; flash_attn_ext_f16_iter - - (Q_f2, K_h2, V_h2, mask_h, dstk, dstk_fixup, scale, slope, logit_softcap, + (Q_f2, K_h2, V_h2, mask_h, indices, dstk, dstk_fixup, scale, slope, logit_softcap, ne01, ne02, stride_K, stride_V, stride_mask, tile_Q, tile_K, tile_V, tile_mask, Q_B, VKQ_C, KQ_max, KQ_rowsum, jt, kb0, k_VKQ_sup); } @@ -1717,7 +1747,7 @@ static __device__ __forceinline__ void flash_attn_ext_f16_process_tile( } } #else - GGML_UNUSED_VARS(Q_f2, K_h2, V_h2, mask_h, sinks_f, dstk, dstk_fixup, + GGML_UNUSED_VARS(Q_f2, K_h2, V_h2, mask_h, indices, sinks_f, dstk, dstk_fixup, scale, slope, logit_softcap, ne01, ne02, gqa_ratio, stride_Q1, stride_Q2, stride_K, stride_V, stride_mask, jt, kb0_start, kb0_stop); @@ -1725,7 +1755,13 @@ static __device__ __forceinline__ void flash_attn_ext_f16_process_tile( #endif // defined(VOLTA_MMA_AVAILABLE) || defined(TURING_MMA_AVAILABLE) || defined(AMD_WMMA_AVAILABLE) || defined(AMD_MFMA_AVAILABLE) } -template +static constexpr __host__ __device__ bool ggml_cuda_flash_attn_ext_mma_f16_may_use_sparse( + const int DKQ, const int DV, const int ncols1, const int ncols2) { + return (DKQ == 512 && DV == 512 && ncols1 == 1 && ncols2 == 8) || + (DKQ == 576 && DV == 512 && ncols1 == 1 && ncols2 == 16); +} + +template __launch_bounds__(ggml_cuda_fattn_mma_get_nthreads(DKQ, DV, ncols1*ncols2), ggml_cuda_fattn_mma_get_occupancy(DKQ, DV, ncols1*ncols2)) static __global__ void flash_attn_ext_f16( const char * Q_ptr, @@ -1751,14 +1787,15 @@ static __global__ void flash_attn_ext_f16( const int32_t nb31, const int32_t nb32, const int64_t nb33) { ggml_cuda_pdl_sync(); // TODO optimize placement #if defined(FLASH_ATTN_AVAILABLE) && (defined(VOLTA_MMA_AVAILABLE) || defined(TURING_MMA_AVAILABLE) || defined(AMD_WMMA_AVAILABLE) || defined(AMD_MFMA_AVAILABLE)) - const char * GGML_CUDA_RESTRICT Q = Q_ptr; - const char * GGML_CUDA_RESTRICT K = K_ptr; - const char * GGML_CUDA_RESTRICT V = V_ptr; - const char * GGML_CUDA_RESTRICT mask = mask_ptr; - const char * GGML_CUDA_RESTRICT sinks = sinks_ptr; - const int * GGML_CUDA_RESTRICT KV_max = KV_max_ptr; - float * GGML_CUDA_RESTRICT dst = dst_ptr; - float2 * GGML_CUDA_RESTRICT dst_meta = dst_meta_ptr; + const char * GGML_CUDA_RESTRICT Q = Q_ptr; + const char * GGML_CUDA_RESTRICT K = K_ptr; + const char * GGML_CUDA_RESTRICT V = V_ptr; + const char * GGML_CUDA_RESTRICT mask = mask_ptr; + const char * GGML_CUDA_RESTRICT sinks = sinks_ptr; + const int * GGML_CUDA_RESTRICT KV_max = use_sparse ? nullptr : KV_max_ptr; + const int * GGML_CUDA_RESTRICT sparse_indices = use_sparse ? KV_max_ptr : nullptr; + float * GGML_CUDA_RESTRICT dst = dst_ptr; + float2 * GGML_CUDA_RESTRICT dst_meta = dst_meta_ptr; // Skip unused kernel variants for faster compilation: if (use_logit_softcap && !(DKQ == 128 || DKQ == 256 || DKQ == 512)) { @@ -1769,6 +1806,11 @@ static __global__ void flash_attn_ext_f16( NO_DEVICE_CODE; return; } + + if (!ggml_cuda_flash_attn_ext_mma_f16_may_use_sparse(DKQ, DV, ncols1, ncols2) && use_sparse) { + NO_DEVICE_CODE; + return; + } #ifdef VOLTA_MMA_AVAILABLE if (ncols1*ncols2 < 32) { NO_DEVICE_CODE; @@ -1845,6 +1887,7 @@ static __global__ void flash_attn_ext_f16( const half2 * V_h2 = V_is_K_view ? K_h2 : (const half2 *) (V + nb23*sequence + nb22*z_KV); const float * sinks_f = sinks ? (const float *) sinks + zt_Q : nullptr; + const int32_t * indices = use_sparse ? sparse_indices + (int64_t(sequence % ne33)*ne31 + jt*ncols1)*ne11 : nullptr; const float slope = ncols2 == 1 ? get_alibi_slope(max_bias, zt_Q, n_head_log2, m0, m1) : 1.0f; @@ -1854,13 +1897,13 @@ static __global__ void flash_attn_ext_f16( constexpr bool is_fixup = false; // All but (potentially) the last iterations write their data to dst rather than the fixup buffer. if (kb0_start == 0) { constexpr bool needs_fixup = false; // CUDA block is working on an entire tile. - flash_attn_ext_f16_process_tile - (Q_f2, K_h2, V_h2, mask_h, sinks_f, dstk, dst_meta, scale, slope, logit_softcap, + flash_attn_ext_f16_process_tile + (Q_f2, K_h2, V_h2, mask_h, indices, sinks_f, dstk, dst_meta, scale, slope, logit_softcap, ne01, ne02, gqa_ratio, ne11, stride_Q1, stride_Q2, stride_K, stride_V, stride_mask, jt, zt_gqa, kb0_start, kb0_stop); } else { constexpr bool needs_fixup = true; // CUDA block is missing the beginning of a tile. - flash_attn_ext_f16_process_tile - (Q_f2, K_h2, V_h2, mask_h, sinks_f, dstk, dst_meta, scale, slope, logit_softcap, + flash_attn_ext_f16_process_tile + (Q_f2, K_h2, V_h2, mask_h, indices, sinks_f, dstk, dst_meta, scale, slope, logit_softcap, ne01, ne02, gqa_ratio, ne11, stride_Q1, stride_Q2, stride_K, stride_V, stride_mask, jt, zt_gqa, kb0_start, kb0_stop); } @@ -1891,6 +1934,7 @@ static __global__ void flash_attn_ext_f16( const half2 * V_h2 = V_is_K_view ? K_h2 : (const half2 *) (V + nb23*sequence + nb22*z_KV); const float * sinks_f = sinks ? (const float *) sinks + zt_Q : nullptr; + const int32_t * indices = use_sparse ? sparse_indices + (int64_t(sequence % ne33)*ne31 + jt*ncols1)*ne11 : nullptr; const float slope = ncols2 == 1 ? get_alibi_slope(max_bias, zt_Q, n_head_log2, m0, m1) : 1.0f; @@ -1900,8 +1944,8 @@ static __global__ void flash_attn_ext_f16( constexpr bool is_fixup = true; // Last index writes its data to fixup buffer to avoid data races with other blocks. constexpr bool needs_fixup = false; - flash_attn_ext_f16_process_tile - (Q_f2, K_h2, V_h2, mask_h, sinks_f, dstk, dst_meta, scale, slope, logit_softcap, + flash_attn_ext_f16_process_tile + (Q_f2, K_h2, V_h2, mask_h, indices, sinks_f, dstk, dst_meta, scale, slope, logit_softcap, ne01, ne02, gqa_ratio, ne11, stride_Q1, stride_Q2, stride_K, stride_V, stride_mask, jt, zt_gqa, kb0_start, kb0_stop); #else GGML_UNUSED_VARS(Q_ptr, K_ptr, V_ptr, mask_ptr, sinks_ptr, KV_max_ptr, dst_ptr, dst_meta_ptr, scale, @@ -1917,6 +1961,8 @@ static __global__ void flash_attn_ext_f16( #endif // defined(FLASH_ATTN_AVAILABLE) && (defined(VOLTA_MMA_AVAILABLE) || defined(TURING_MMA_AVAILABLE) || defined(AMD_WMMA_AVAILABLE) || defined(AMD_MFMA_AVAILABLE)) } +bool ggml_cuda_flash_attn_ext_mma_f16_shall_use_sparse(ggml_backend_cuda_context & ctx, ggml_tensor * dst); + template void ggml_cuda_flash_attn_ext_mma_f16_case(ggml_backend_cuda_context & ctx, ggml_tensor * dst) { const ggml_tensor * KQV = dst; @@ -1963,20 +2009,49 @@ void ggml_cuda_flash_attn_ext_mma_f16_case(ggml_backend_cuda_context & ctx, ggml using fattn_kernel_ptr_t = fattn_kernel_t; #endif // defined(GGML_USE_HIP) fattn_kernel_t fattn_kernel; + bool use_sparse = false; if (logit_softcap == 0.0f) { constexpr bool use_logit_softcap = false; - fattn_kernel = flash_attn_ext_f16; +#if !defined(GGML_USE_HIP) && !defined(GGML_USE_MUSA) + if constexpr (ggml_cuda_flash_attn_ext_mma_f16_may_use_sparse(DKQ, DV, ncols1, ncols2)) { + if (ggml_cuda_flash_attn_ext_mma_f16_shall_use_sparse(ctx, dst)) { + constexpr bool use_sparse_kernel = true; + fattn_kernel = flash_attn_ext_f16; + use_sparse = true; + + static bool shared_memory_limit_raised[GGML_CUDA_MAX_DEVICES] = {false}; + if (!shared_memory_limit_raised[id]) { + CUDA_CHECK(cudaFuncSetAttribute(reinterpret_cast(fattn_kernel), cudaFuncAttributeMaxDynamicSharedMemorySize, nbytes_shared_total)); + shared_memory_limit_raised[id] = true; + } + } else { + constexpr bool use_sparse_kernel = false; + fattn_kernel = flash_attn_ext_f16; + + static bool shared_memory_limit_raised[GGML_CUDA_MAX_DEVICES] = {false}; + if (!shared_memory_limit_raised[id]) { + CUDA_CHECK(cudaFuncSetAttribute(reinterpret_cast(fattn_kernel), cudaFuncAttributeMaxDynamicSharedMemorySize, nbytes_shared_total)); + shared_memory_limit_raised[id] = true; + } + } + } else +#endif // !defined(GGML_USE_HIP) && !defined(GGML_USE_MUSA) + { + constexpr bool use_sparse_kernel = false; + fattn_kernel = flash_attn_ext_f16; #if !defined(GGML_USE_MUSA) - static bool shared_memory_limit_raised[GGML_CUDA_MAX_DEVICES] = {false}; - if (!shared_memory_limit_raised[id]) { - CUDA_CHECK(cudaFuncSetAttribute(reinterpret_cast(fattn_kernel), cudaFuncAttributeMaxDynamicSharedMemorySize, nbytes_shared_total)); - shared_memory_limit_raised[id] = true; - } + static bool shared_memory_limit_raised[GGML_CUDA_MAX_DEVICES] = {false}; + if (!shared_memory_limit_raised[id]) { + CUDA_CHECK(cudaFuncSetAttribute(reinterpret_cast(fattn_kernel), cudaFuncAttributeMaxDynamicSharedMemorySize, nbytes_shared_total)); + shared_memory_limit_raised[id] = true; + } #endif // !defined(GGML_USE_MUSA) + } } else { constexpr bool use_logit_softcap = true; - fattn_kernel = flash_attn_ext_f16; + constexpr bool use_sparse_kernel = false; + fattn_kernel = flash_attn_ext_f16; #if !defined(GGML_USE_MUSA) static bool shared_memory_limit_raised[GGML_CUDA_MAX_DEVICES] = {false}; @@ -1988,7 +2063,7 @@ void ggml_cuda_flash_attn_ext_mma_f16_case(ggml_backend_cuda_context & ctx, ggml } launch_fattn - (ctx, dst, fattn_kernel, nwarps, nbytes_shared_total, nbatch_fa, true, true, true, warp_size_host); + (ctx, dst, fattn_kernel, nwarps, nbytes_shared_total, nbatch_fa, true, true, true, use_sparse, warp_size_host); } diff --git a/ggml/src/ggml-cuda/fattn-tile.cuh b/ggml/src/ggml-cuda/fattn-tile.cuh index d1164b8526d..8981ab804ce 100644 --- a/ggml/src/ggml-cuda/fattn-tile.cuh +++ b/ggml/src/ggml-cuda/fattn-tile.cuh @@ -1163,7 +1163,7 @@ static void launch_fattn_tile_switch_ncols1(ggml_backend_cuda_context & ctx, ggm const int nbatch_fa = ggml_cuda_fattn_tile_get_nbatch_fa(DKQ, DV, cols_per_block, cc); fattn_kernel_t fattn_kernel = flash_attn_tile; launch_fattn - (ctx, dst, fattn_kernel, nwarps, nbytes_shared, nbatch_fa, true, true, false, warp_size); + (ctx, dst, fattn_kernel, nwarps, nbytes_shared, nbatch_fa, true, true, false, false, warp_size); return; } } @@ -1179,7 +1179,7 @@ static void launch_fattn_tile_switch_ncols1(ggml_backend_cuda_context & ctx, ggm const int nbatch_fa = ggml_cuda_fattn_tile_get_nbatch_fa(DKQ, DV, cols_per_block, cc); fattn_kernel_t fattn_kernel = flash_attn_tile; launch_fattn - (ctx, dst, fattn_kernel, nwarps, nbytes_shared, nbatch_fa, true, true, false, warp_size); + (ctx, dst, fattn_kernel, nwarps, nbytes_shared, nbatch_fa, true, true, false, false, warp_size); return; } } @@ -1191,7 +1191,7 @@ static void launch_fattn_tile_switch_ncols1(ggml_backend_cuda_context & ctx, ggm const int nbatch_fa = ggml_cuda_fattn_tile_get_nbatch_fa(DKQ, DV, cols_per_block, cc); fattn_kernel_t fattn_kernel = flash_attn_tile; launch_fattn - (ctx, dst, fattn_kernel, nwarps, nbytes_shared, nbatch_fa, true, true, false, warp_size); + (ctx, dst, fattn_kernel, nwarps, nbytes_shared, nbatch_fa, true, true, false, false, warp_size); return; } } @@ -1203,7 +1203,7 @@ static void launch_fattn_tile_switch_ncols1(ggml_backend_cuda_context & ctx, ggm const int nbatch_fa = ggml_cuda_fattn_tile_get_nbatch_fa(DKQ, DV, cols_per_block, cc); fattn_kernel_t fattn_kernel = flash_attn_tile; launch_fattn - (ctx, dst, fattn_kernel, nwarps, nbytes_shared, nbatch_fa, true, true, false, warp_size); + (ctx, dst, fattn_kernel, nwarps, nbytes_shared, nbatch_fa, true, true, false, false, warp_size); return; } } @@ -1215,7 +1215,7 @@ static void launch_fattn_tile_switch_ncols1(ggml_backend_cuda_context & ctx, ggm const int nbatch_fa = ggml_cuda_fattn_tile_get_nbatch_fa(DKQ, DV, cols_per_block, cc); fattn_kernel_t fattn_kernel = flash_attn_tile; launch_fattn - (ctx, dst, fattn_kernel, nwarps, nbytes_shared, nbatch_fa, true, true, false, warp_size); + (ctx, dst, fattn_kernel, nwarps, nbytes_shared, nbatch_fa, true, true, false, false, warp_size); return; } } @@ -1226,7 +1226,7 @@ static void launch_fattn_tile_switch_ncols1(ggml_backend_cuda_context & ctx, ggm const int nbatch_fa = ggml_cuda_fattn_tile_get_nbatch_fa(DKQ, DV, cols_per_block, cc); fattn_kernel_t fattn_kernel = flash_attn_tile; launch_fattn - (ctx, dst, fattn_kernel, nwarps, nbytes_shared, nbatch_fa, true, true, false, warp_size); + (ctx, dst, fattn_kernel, nwarps, nbytes_shared, nbatch_fa, true, true, false, false, warp_size); return; } diff --git a/ggml/src/ggml-cuda/fattn-vec.cuh b/ggml/src/ggml-cuda/fattn-vec.cuh index 69dd9368624..519b36b9ff4 100644 --- a/ggml/src/ggml-cuda/fattn-vec.cuh +++ b/ggml/src/ggml-cuda/fattn-vec.cuh @@ -540,7 +540,7 @@ void ggml_cuda_flash_attn_ext_vec_case_impl(ggml_backend_cuda_context & ctx, ggm const bool need_f16_K = type_K == GGML_TYPE_F16; const bool need_f16_V = type_V == GGML_TYPE_F16; constexpr size_t nbytes_shared = 0; - launch_fattn(ctx, dst, fattn_kernel, nwarps, nbytes_shared, D, need_f16_K, need_f16_V, false); + launch_fattn(ctx, dst, fattn_kernel, nwarps, nbytes_shared, D, need_f16_K, need_f16_V, false, false); } template diff --git a/ggml/src/ggml-cuda/fattn.cu b/ggml/src/ggml-cuda/fattn.cu index ab7a3b297c0..ae217fbd9df 100644 --- a/ggml/src/ggml-cuda/fattn.cu +++ b/ggml/src/ggml-cuda/fattn.cu @@ -5,11 +5,144 @@ #include "fattn-vec.cuh" #include "fattn.cuh" +#if !defined(GGML_USE_HIP) && !defined(GGML_USE_MUSA) +__launch_bounds__(256, 1) +static __global__ void flash_attn_mask_to_sparse_indices( + const half * mask_ptr, int32_t * indices_ptr, const int ne30, const int n_kv_max, + const int64_t s31, const int64_t s33) { + ggml_cuda_pdl_sync(); + + constexpr int values_per_lane = 8; + const int tid = threadIdx.x; + const int warp = tid / WARP_SIZE; + const int lane = tid % WARP_SIZE; + const int sequence = blockIdx.y; + const int query = blockIdx.x; + + const half * mask = mask_ptr + sequence*s33 + query*s31; + int32_t * indices = indices_ptr + (int64_t(sequence)*gridDim.x + query)*n_kv_max; + + __shared__ int warp_offsets[256/WARP_SIZE]; + __shared__ int row_count; + __shared__ int chunk_count; + + if (tid == 0) { + row_count = 0; + } + __syncthreads(); + + for (int i0 = 0; i0 < ne30; i0 += blockDim.x*values_per_lane) { + uint32_t selected_warp[values_per_lane]; + int warp_count = 0; +#pragma unroll + for (int item = 0; item < values_per_lane; ++item) { + const int i = i0 + (warp*values_per_lane + item)*WARP_SIZE + lane; + const bool selected = i < ne30 && isfinite(__half2float(mask[i])); + selected_warp[item] = __ballot_sync(0xFFFFFFFF, selected); + warp_count += __popc(selected_warp[item]); + } + + if (lane == 0) { + warp_offsets[warp] = warp_count; + } + __syncthreads(); + + if (tid == 0) { + int offset = 0; +#pragma unroll + for (int iw = 0; iw < 256/WARP_SIZE; ++iw) { + const int count = warp_offsets[iw]; + warp_offsets[iw] = offset; + offset += count; + } + chunk_count = offset; + } + __syncthreads(); + + const uint32_t lane_mask = lane == 0 ? 0 : (1u << lane) - 1; + int warp_item_offset = 0; +#pragma unroll + for (int item = 0; item < values_per_lane; ++item) { + const int i = i0 + (warp*values_per_lane + item)*WARP_SIZE + lane; + const int dst = row_count + warp_offsets[warp] + warp_item_offset + __popc(selected_warp[item] & lane_mask); + if ((selected_warp[item] & (uint32_t(1) << lane)) && dst < n_kv_max) { + indices[dst] = i; + } + warp_item_offset += __popc(selected_warp[item]); + } + __syncthreads(); + + if (tid == 0) { + row_count += chunk_count; + } + __syncthreads(); + } + + const int count = row_count; + for (int i = count + tid; i < n_kv_max; i += blockDim.x) { + indices[i] = -1; + } + __syncthreads(); + + // the dependent grid reads indices, signal once the row is complete + ggml_cuda_pdl_lc(); +} +#endif // !defined(GGML_USE_HIP) && !defined(GGML_USE_MUSA) + +void ggml_cuda_flash_attn_ext_compact_mask( + const ggml_tensor * mask, int32_t * indices, int32_t n_kv_max, cudaStream_t stream) { +#if defined(GGML_USE_HIP) || defined(GGML_USE_MUSA) + GGML_UNUSED_VARS(mask, indices, n_kv_max, stream); + GGML_ABORT("sparse flash attention is only supported on NVIDIA CUDA"); +#else + const int64_t s31 = mask->nb[1] / sizeof(half); + const int64_t s33 = mask->nb[3] / sizeof(half); + const dim3 blocks_num(mask->ne[1], mask->ne[3], 1); + const dim3 block_dim(256, 1, 1); + const ggml_cuda_kernel_launch_params launch_params(blocks_num, block_dim, 0, stream); + ggml_cuda_kernel_launch(flash_attn_mask_to_sparse_indices, launch_params, + (const half *) mask->data, indices, int(mask->ne[0]), n_kv_max, s31, s33); + CUDA_CHECK(cudaGetLastError()); +#endif // !defined(GGML_USE_HIP) && !defined(GGML_USE_MUSA) +} + +bool ggml_cuda_flash_attn_ext_mma_f16_shall_use_sparse(ggml_backend_cuda_context & ctx, ggml_tensor * dst) { +#if defined(GGML_USE_HIP) || defined(GGML_USE_MUSA) + GGML_UNUSED_VARS(ctx, dst); + return false; +#else + const ggml_tensor * Q = dst->src[0]; + const ggml_tensor * K = dst->src[1]; + const ggml_tensor * mask = dst->src[3]; + const int cc = ggml_cuda_info().devices[ctx.device].cc; + + float max_bias = 0.0f; + float logit_softcap = 0.0f; + memcpy(&max_bias, (const float *) dst->op_params + 1, sizeof(float)); + memcpy(&logit_softcap, (const float *) dst->op_params + 2, sizeof(float)); + + const int32_t n_kv_max = ggml_get_op_params_i32(dst, 4); + return GGML_CUDA_CC_IS_NVIDIA(cc) && turing_mma_available(cc) && + mask != nullptr && n_kv_max > 0 && max_bias == 0.0f && logit_softcap == 0.0f && + mask->ne[0] == K->ne[1] && mask->ne[1] >= Q->ne[1] && mask->ne[2] == 1 && + K->ne[1] >= std::max(4096, 2LL*n_kv_max); +#endif // !defined(GGML_USE_HIP) && !defined(GGML_USE_MUSA) +} + template static void ggml_cuda_flash_attn_ext_mma_f16_switch_ncols1(ggml_backend_cuda_context & ctx, ggml_tensor * dst) { const int cc = ggml_cuda_info().devices[ggml_cuda_get_device()].cc; const ggml_tensor * Q = dst->src[0]; +#if !defined(GGML_USE_HIP) && !defined(GGML_USE_MUSA) + if constexpr (ggml_cuda_flash_attn_ext_mma_f16_may_use_sparse(DKQ, DV, 1, ncols2)) { + if (ggml_cuda_flash_attn_ext_mma_f16_shall_use_sparse(ctx, dst)) { + ggml_cuda_flash_attn_ext_mma_f16_case(ctx, dst); + return; + } + } +#endif // !defined(GGML_USE_HIP) && !defined(GGML_USE_MUSA) + if constexpr (ncols2 <= 8) { if (turing_mma_available(cc) && Q->ne[1] <= 8/ncols2) { ggml_cuda_flash_attn_ext_mma_f16_case(ctx, dst); diff --git a/ggml/src/ggml.c b/ggml/src/ggml.c index 3bd3e3fe5ea..8dc09450848 100644 --- a/ggml/src/ggml.c +++ b/ggml/src/ggml.c @@ -5506,6 +5506,15 @@ enum ggml_prec ggml_flash_attn_ext_get_prec( return (enum ggml_prec) prec_i32; } +void ggml_flash_attn_ext_set_n_kv_max( + struct ggml_tensor * a, + int32_t n_kv_max) { + GGML_ASSERT(a->op == GGML_OP_FLASH_ATTN_EXT); + GGML_ASSERT(n_kv_max >= 0); + + ggml_set_op_params_i32(a, 4, n_kv_max); +} + void ggml_flash_attn_ext_add_sinks( struct ggml_tensor * a, struct ggml_tensor * sinks) { From 4d343d7c05c36bc863536cab9d042b96acf0eb9c Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Adrien=20Gallou=C3=ABt?= Date: Wed, 2 Sep 2026 18:54:11 +0200 Subject: [PATCH 082/104] ggml-cuda : remove unused vars (llama/28235) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Signed-off-by: Adrien Gallouët --- ggml/src/ggml-cuda/mmq-vec-dot.cuh | 11 ----------- ggml/src/ggml-cuda/mmq.cuh | 5 ----- 2 files changed, 16 deletions(-) diff --git a/ggml/src/ggml-cuda/mmq-vec-dot.cuh b/ggml/src/ggml-cuda/mmq-vec-dot.cuh index d573433865f..4d1c398fc54 100644 --- a/ggml/src/ggml-cuda/mmq-vec-dot.cuh +++ b/ggml/src/ggml-cuda/mmq-vec-dot.cuh @@ -148,7 +148,6 @@ static __device__ __forceinline__ void ggml_cuda_mmq_vec_dot_q8_0_q8_1_mma( typedef tile<16, 8, int, input_layout> tile_B; typedef tile<16, 16, int, DATA_LAYOUT_J_MAJOR> tile_C; - constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback); constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback); constexpr int rows_per_warp = ggml_cuda_mmq_get_rows_per_warp(type, J, fallback); constexpr int ntx = rows_per_warp/tile_C::I; // Number of x minitiles per warp. @@ -204,7 +203,6 @@ static __device__ __forceinline__ void ggml_cuda_mmq_vec_dot_q8_0_q8_1_mma( typedef tile< 8, 8, int> tile_B; typedef tile<16, 8, int> tile_C; - constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback); constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback); constexpr int rows_per_warp = ggml_cuda_mmq_get_rows_per_warp(type, J, fallback); constexpr int ntx = rows_per_warp/tile_C::I; // Number of x minitiles per warp. @@ -320,7 +318,6 @@ template static __device__ __forceinline_ typedef tile<16, 8, int, input_layout> tile_B; typedef tile<16, 16, int, DATA_LAYOUT_J_MAJOR> tile_C; - constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback); constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback); constexpr int rows_per_warp = ggml_cuda_mmq_get_rows_per_warp(type, J, fallback); constexpr int ntx = rows_per_warp/tile_C::I; // Number of x minitiles per warp. @@ -371,7 +368,6 @@ template static __device__ __forceinline_ typedef tile< 8, 8, int> tile_B; typedef tile<16, 8, int> tile_C; - constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback); constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback); constexpr int rows_per_warp = ggml_cuda_mmq_get_rows_per_warp(type, J, fallback); constexpr int ntx = rows_per_warp/tile_C::I; // Number of x minitiles per warp. @@ -486,7 +482,6 @@ template static __device__ __forceinline_ typedef tile<16, 4, int, input_layout> tile_B; typedef tile<16, 16, int, DATA_LAYOUT_J_MAJOR> tile_C; - constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback); constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback); constexpr int rows_per_warp = ggml_cuda_mmq_get_rows_per_warp(type, J, fallback); constexpr int ntx = rows_per_warp/tile_C::I; // Number of x minitiles per warp. @@ -537,7 +532,6 @@ template static __device__ __forceinline_ typedef tile< 8, 4, int> tile_B; typedef tile<16, 8, int> tile_C; - constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback); constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback); constexpr int rows_per_warp = ggml_cuda_mmq_get_rows_per_warp(type, J, fallback); constexpr int ntx = rows_per_warp/tile_C::I; // Number of x minitiles per warp. @@ -686,7 +680,6 @@ template static __device__ __forceinline_ typedef tile<16, 4, int, input_layout> tile_B; typedef tile<16, 16, int, DATA_LAYOUT_J_MAJOR> tile_C; - constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback); constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback); constexpr int rows_per_warp = ggml_cuda_mmq_get_rows_per_warp(type, J, fallback); constexpr int ntx = rows_per_warp/tile_C::I; // Number of x minitiles per warp. @@ -756,7 +749,6 @@ template static __device__ __forceinline_ typedef tile< 8, 4, int> tile_B; typedef tile<16, 8, int> tile_C; - constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback); constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback); constexpr int rows_per_warp = ggml_cuda_mmq_get_rows_per_warp(type, J, fallback); constexpr int ntx = rows_per_warp/tile_C::I; // Number of x minitiles per warp. @@ -1023,7 +1015,6 @@ template static __device__ __forceinline_ typedef tile<16, 4, int, input_layout> tile_B; typedef tile<16, 16, int, DATA_LAYOUT_J_MAJOR> tile_C; - constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback); constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback); constexpr int rows_per_warp = ggml_cuda_mmq_get_rows_per_warp(type, J, fallback); constexpr int ntx = rows_per_warp/tile_C::I; // Number of x minitiles per warp. @@ -1075,7 +1066,6 @@ template static __device__ __forceinline_ typedef tile< 8, 4, int> tile_B; typedef tile<16, 8, int> tile_C; - constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback); constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback); constexpr int rows_per_warp = ggml_cuda_mmq_get_rows_per_warp(type, J, fallback); constexpr int ntx = rows_per_warp/tile_C::I; // Number of x minitiles per warp. @@ -1190,7 +1180,6 @@ template static __device__ __forceinline_ typedef tile<8, 8, int> tile_B; typedef tile<16, 8, float> tile_C; - constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback); constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback); constexpr int rows_per_warp = ggml_cuda_mmq_get_rows_per_warp(type, J, fallback); constexpr int ntx = rows_per_warp / tile_C::I; diff --git a/ggml/src/ggml-cuda/mmq.cuh b/ggml/src/ggml-cuda/mmq.cuh index c978b4421c5..b4a747720f7 100644 --- a/ggml/src/ggml-cuda/mmq.cuh +++ b/ggml/src/ggml-cuda/mmq.cuh @@ -481,9 +481,6 @@ static __device__ __forceinline__ void ggml_cuda_mmq_write_back_mma( typedef tile<16, 8, int> tile_C; #endif // defined(AMD_MFMA_AVAILABLE) || defined(AMD_WMMA_AVAILABLE) - constexpr int warp_size = ggml_cuda_get_physical_warp_size(); - constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback) / warp_size; - constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback); constexpr int rows_per_warp = ggml_cuda_mmq_get_rows_per_warp(type, J, fallback); constexpr int ntx = rows_per_warp/tile_C::I; // Number of x minitiles per warp. @@ -540,8 +537,6 @@ struct ggml_cuda_mmq_util_funcs { template static constexpr __device__ ggml_cuda_mmq_util_funcs ggml_cuda_mmq_get_util_funcs() { - constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback); - if (!ggml_cuda_mmq_get_config(type, J, fallback).use_mma_data_layout()) { switch (type) { case GGML_TYPE_Q1_0: From 3a1c7d6b6f47efeb69bde410b0003129563fa9f7 Mon Sep 17 00:00:00 2001 From: Mads Marquart Date: Wed, 2 Sep 2026 20:09:46 +0200 Subject: [PATCH 083/104] metal : fix memory query under low-memory conditions (llama/27701) * metal: Fix memory query under low-memory conditions * Simply variable name Co-authored-by: Georgi Gerganov * Write it even shorter Co-authored-by: Niklas Wenzel --------- Co-authored-by: Georgi Gerganov Co-authored-by: Niklas Wenzel --- ggml/src/ggml-metal/ggml-metal-device.m | 6 ++++-- 1 file changed, 4 insertions(+), 2 deletions(-) diff --git a/ggml/src/ggml-metal/ggml-metal-device.m b/ggml/src/ggml-metal/ggml-metal-device.m index ef81084241d..e20a4e89160 100644 --- a/ggml/src/ggml-metal/ggml-metal-device.m +++ b/ggml/src/ggml-metal/ggml-metal-device.m @@ -1471,8 +1471,10 @@ void ggml_metal_device_event_synchronize(ggml_metal_device_t dev, ggml_metal_eve void ggml_metal_device_get_memory(ggml_metal_device_t dev, size_t * free, size_t * total) { if (@available(macOS 10.12, iOS 16.0, *)) { - *total = dev->mtl_device.recommendedMaxWorkingSetSize; - *free = *total - dev->mtl_device.currentAllocatedSize; + *total = dev->mtl_device.recommendedMaxWorkingSetSize; + size_t cur = dev->mtl_device.currentAllocatedSize; + // it's possible to allocate more than `recommendedMaxWorkingSetSize` + *free = *total > cur ? *total - cur : 0; } else { *free = 0; *total = 0; From 1bdda1e36676acd6f18d004494ab31847a92ae00 Mon Sep 17 00:00:00 2001 From: Isaac <34376531+init-22@users.noreply.github.com> Date: Wed, 2 Sep 2026 23:43:12 +0530 Subject: [PATCH 084/104] metal : add fa-vec tunings for M3 (llama/28236) --- ggml/src/ggml-metal/ggml-metal-tuning.cpp | 101 ++++++++++++++++++++++ 1 file changed, 101 insertions(+) diff --git a/ggml/src/ggml-metal/ggml-metal-tuning.cpp b/ggml/src/ggml-metal/ggml-metal-tuning.cpp index c89a905dffa..b66fe65240f 100644 --- a/ggml/src/ggml-metal/ggml-metal-tuning.cpp +++ b/ggml/src/ggml-metal/ggml-metal-tuning.cpp @@ -1468,6 +1468,107 @@ constexpr fa_vec_entry_t fa_vec_tuned_table[] = { { { GGML_METAL_DEVICE_M2_ULTRA, GGML_TYPE_Q8_0, 576, 512, -1, 0 }, { 1, 4 } }, { { GGML_METAL_DEVICE_M2_ULTRA, GGML_TYPE_Q8_0, 576, 512, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3, GGML_TYPE_F16, 32, 32, 1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3, GGML_TYPE_F16, 32, 32, 1, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M3, GGML_TYPE_F16, 32, 32, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3, GGML_TYPE_F16, 32, 32, 2, 2 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M3, GGML_TYPE_F16, 32, 32, 2, 3 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3, GGML_TYPE_F16, 32, 32, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3, GGML_TYPE_F16, 32, 32, 3, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M3, GGML_TYPE_F16, 64, 64, 3, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3, GGML_TYPE_F16, 64, 64, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3, GGML_TYPE_F16, 64, 64, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3, GGML_TYPE_F16, 64, 64, 1, 4 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M3, GGML_TYPE_F16, 64, 64, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3, GGML_TYPE_F16, 64, 64, 3, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M3, GGML_TYPE_F16, 96, 96, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3, GGML_TYPE_F16, 96, 96, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3, GGML_TYPE_F16, 96, 96, 1, 4 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3, GGML_TYPE_F16, 96, 96, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3, GGML_TYPE_F16, 96, 96, 2, 4 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3, GGML_TYPE_F16, 96, 96, 3, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3, GGML_TYPE_F16, 128, 128, 1, 3 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3, GGML_TYPE_F16, 128, 128, 2, 3 }, { 2, 2 } }, + { { GGML_METAL_DEVICE_M3, GGML_TYPE_F16, 128, 128, 3, 1 }, { 2, 2 } }, + { { GGML_METAL_DEVICE_M3, GGML_TYPE_F16, 128, 128, 3, 3 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3, GGML_TYPE_F16, 192, 192, 3, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3, GGML_TYPE_F16, 192, 192, -1, 1 }, { 2, 2 } }, + { { GGML_METAL_DEVICE_M3, GGML_TYPE_F16, 192, 192, 1, 2 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M3, GGML_TYPE_F16, 192, 192, 1, 4 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M3, GGML_TYPE_F16, 192, 192, 2, 4 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M3, GGML_TYPE_F16, 192, 128, 2, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3, GGML_TYPE_F16, 192, 128, 3, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3, GGML_TYPE_F16, 192, 128, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3, GGML_TYPE_F16, 192, 128, 1, 2 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M3, GGML_TYPE_F16, 192, 128, 1, 4 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M3, GGML_TYPE_F16, 192, 128, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3, GGML_TYPE_F16, 192, 128, 2, 4 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M3, GGML_TYPE_F16, 192, 128, 3, 1 }, { 2, 2 } }, + { { GGML_METAL_DEVICE_M3, GGML_TYPE_F16, 192, 128, 3, 2 }, { 2, 2 } }, + { { GGML_METAL_DEVICE_M3, GGML_TYPE_F16, 256, 256, 1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3, GGML_TYPE_F16, 256, 256, 3, 0 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M3, GGML_TYPE_F16, 256, 256, -1, 1 }, { 2, 2 } }, + { { GGML_METAL_DEVICE_M3, GGML_TYPE_F16, 256, 256, 1, 1 }, { 1, 1 } }, + { { GGML_METAL_DEVICE_M3, GGML_TYPE_F16, 256, 256, 1, 2 }, { 1, 1 } }, + { { GGML_METAL_DEVICE_M3, GGML_TYPE_F16, 256, 256, 1, 4 }, { 1, 1 } }, + { { GGML_METAL_DEVICE_M3, GGML_TYPE_F16, 256, 256, 2, 2 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M3, GGML_TYPE_F16, 320, 256, 2, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3, GGML_TYPE_F16, 320, 256, 3, 0 }, { 2, 2 } }, + { { GGML_METAL_DEVICE_M3, GGML_TYPE_F16, 320, 256, -1, 1 }, { 2, 2 } }, + { { GGML_METAL_DEVICE_M3, GGML_TYPE_F16, 320, 256, 1, 2 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M3, GGML_TYPE_F16, 320, 256, 1, 4 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M3, GGML_TYPE_F16, 512, 512, 2, 0 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M3, GGML_TYPE_F16, 512, 512, 3, 0 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M3, GGML_TYPE_F16, 512, 512, 3, 1 }, { 4, 2 } }, + { { GGML_METAL_DEVICE_M3, GGML_TYPE_F16, 512, 512, 3, 3 }, { 4, 2 } }, + { { GGML_METAL_DEVICE_M3, GGML_TYPE_F16, 576, 512, 2, 0 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M3, GGML_TYPE_F16, 576, 512, 2, 1 }, { 4, 2 } }, + { { GGML_METAL_DEVICE_M3, GGML_TYPE_F16, 576, 512, 2, 2 }, { 4, 2 } }, + { { GGML_METAL_DEVICE_M3, GGML_TYPE_F16, 576, 512, 3, 1 }, { 4, 2 } }, + { { GGML_METAL_DEVICE_M3, GGML_TYPE_Q8_0, 32, 32, -1, 1 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M3, GGML_TYPE_Q8_0, 32, 32, 1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3, GGML_TYPE_Q8_0, 32, 32, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3, GGML_TYPE_Q8_0, 32, 32, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3, GGML_TYPE_Q8_0, 64, 64, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3, GGML_TYPE_Q8_0, 64, 64, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3, GGML_TYPE_Q8_0, 64, 64, 2, 2 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M3, GGML_TYPE_Q8_0, 64, 64, 2, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M3, GGML_TYPE_Q8_0, 64, 64, 3, 2 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M3, GGML_TYPE_Q8_0, 64, 64, 3, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M3, GGML_TYPE_Q8_0, 96, 96, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3, GGML_TYPE_Q8_0, 96, 96, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3, GGML_TYPE_Q8_0, 96, 96, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3, GGML_TYPE_Q8_0, 96, 96, 3, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3, GGML_TYPE_Q8_0, 128, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3, GGML_TYPE_Q8_0, 128, 128, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3, GGML_TYPE_Q8_0, 128, 128, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3, GGML_TYPE_Q8_0, 128, 128, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3, GGML_TYPE_Q8_0, 128, 128, 3, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3, GGML_TYPE_Q8_0, 128, 128, 3, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M3, GGML_TYPE_Q8_0, 128, 128, 3, 4 }, { 1, 1 } }, + { { GGML_METAL_DEVICE_M3, GGML_TYPE_Q8_0, 192, 192, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3, GGML_TYPE_Q8_0, 192, 192, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3, GGML_TYPE_Q8_0, 192, 192, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3, GGML_TYPE_Q8_0, 192, 192, 2, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3, GGML_TYPE_Q8_0, 192, 192, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3, GGML_TYPE_Q8_0, 192, 192, 3, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3, GGML_TYPE_Q8_0, 192, 128, 1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3, GGML_TYPE_Q8_0, 192, 128, 2, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3, GGML_TYPE_Q8_0, 192, 128, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3, GGML_TYPE_Q8_0, 192, 128, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3, GGML_TYPE_Q8_0, 192, 128, 3, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3, GGML_TYPE_Q8_0, 256, 256, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3, GGML_TYPE_Q8_0, 256, 256, 3, 0 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M3, GGML_TYPE_Q8_0, 256, 256, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3, GGML_TYPE_Q8_0, 320, 256, 1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3, GGML_TYPE_Q8_0, 320, 256, 3, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3, GGML_TYPE_Q8_0, 320, 256, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3, GGML_TYPE_Q8_0, 320, 256, 1, 1 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M3, GGML_TYPE_Q8_0, 512, 512, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3, GGML_TYPE_Q8_0, 512, 512, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3, GGML_TYPE_Q8_0, 576, 512, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3, GGML_TYPE_Q8_0, 576, 512, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_F16, 32, 32, 1, 1 }, { 2, 4 } }, { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_F16, 32, 32, 1, 3 }, { 2, 4 } }, { { GGML_METAL_DEVICE_M3_PRO, GGML_TYPE_F16, 32, 32, 2, 1 }, { 2, 4 } }, From 37f0f443d72003a83f1d95a8681ab194a254d7d6 Mon Sep 17 00:00:00 2001 From: cqderek Date: Thu, 3 Sep 2026 03:59:36 +0800 Subject: [PATCH 085/104] ggml-hexagon: add F16 support for unary ops (llama/28228) Extend the HTP backend's F16 unary op coverage to include ABS on top of the existing NORM/RMS_NORM/L2_NORM/SCALE/CLAMP/SQR/SQRT set. - Add hvx_abs_f16_{aa,au,ua,uu} + dispatcher in hvx-arith.h, mirroring the sqr_f16 kernel structure and using the existing hvx_vec_abs_f16() sign-bit-clear helper - Add abs_f16() row-wise dispatch and DEFINE_UNARY_TASK_F16(unary_abs, ...) in unary-ops.c, wired into execute_op_unary()'s op_type/task_func switches - Register HTP_OP_UNARY_ABS in htp_op_is_unary() (unary-ops.h) so that ggml_hexagon_precompute_unary_params() fills kernel_params (n_threads, VTCM layout) for ABS nodes -- required for the F16 path to function - Narrow the F16 GGML_OP_UNARY gate in ggml_hexagon_supported_unary() (ggml-hexagon.cpp) to allow GGML_UNARY_OP_ABS specifically, instead of rejecting all GGML_OP_UNARY ops for F16 - Merge the separate execute_op_unary_f32()/execute_op_unary_f16() functions into a single execute_op_unary(), branching on an is_f16 flag for the parts that actually differ by type (elem_size, the early F16 op-support check, and which task_func table to use) while keeping the F32-only tiled/RMS_NORM_MUL paths intact -- per review feedback to avoid duplicating the shared VTCM/DMA plumbing Verified on-device (QRD8850, Hexagon v81) via test-backend-ops -o ABS: 8/8 passing (F16 + F32, HTP0, no CPU fallback). Regression-checked SQR/CLAMP/SQRT (F16+F32) and NORM/RMS_NORM/L2_NORM/SCALE (F32; their F16 paths have no CPU reference kernel in test-backend-ops and cannot be correctness-tested there independent of this change). --- ggml/src/ggml-hexagon/ggml-hexagon.cpp | 38 +++- ggml/src/ggml-hexagon/htp/hvx-arith.h | 171 +++++++++++++++- ggml/src/ggml-hexagon/htp/hvx-log.h | 29 +++ ggml/src/ggml-hexagon/htp/hvx-norm.h | 197 +++++++++++++++++++ ggml/src/ggml-hexagon/htp/hvx-scale.h | 66 +++++++ ggml/src/ggml-hexagon/htp/hvx-sqrt.h | 63 ++++++ ggml/src/ggml-hexagon/htp/unary-ops.c | 260 ++++++++++++++++++++++--- ggml/src/ggml-hexagon/htp/unary-ops.h | 17 +- 8 files changed, 797 insertions(+), 44 deletions(-) diff --git a/ggml/src/ggml-hexagon/ggml-hexagon.cpp b/ggml/src/ggml-hexagon/ggml-hexagon.cpp index 04fb9a22339..104201daff5 100644 --- a/ggml/src/ggml-hexagon/ggml-hexagon.cpp +++ b/ggml/src/ggml-hexagon/ggml-hexagon.cpp @@ -4005,8 +4005,10 @@ static void ggml_hexagon_precompute_unary_params( kparams->n_threads = n_threads; - const size_t src0_data_row_size = src0->ne[0] * sizeof(float); - const size_t dst_data_row_size = dst->ne[0] * sizeof(float); + const size_t elem_size = ggml_type_size(src0->type); + + const size_t src0_data_row_size = src0->ne[0] * elem_size; + const size_t dst_data_row_size = dst->ne[0] * ggml_type_size(dst->type); const size_t src0_row_size_aligned = hex_round_up(src0_data_row_size, 128); const size_t dst_row_size_aligned = hex_round_up(dst_data_row_size, 128); @@ -4020,7 +4022,7 @@ static void ggml_hexagon_precompute_unary_params( if (op == HTP_OP_RMS_NORM_MUL) { GGML_ASSERT(src1 != nullptr); - src1_data_row_size = src1->ne[0] * sizeof(float); + src1_data_row_size = src1->ne[0] * ggml_type_size(src1->type); src1_row_size_aligned = hex_round_up(src1_data_row_size, 128); broadcast_weight = (src1->ne[1] * src1->ne[2] * src1->ne[3] == 1); } @@ -4034,7 +4036,7 @@ static void ggml_hexagon_precompute_unary_params( htp_unary_vtcm_layout_build(&L, op, src0->ne[0], dst->ne[0], op == HTP_OP_RMS_NORM_MUL ? src1->ne[0] : 0, - broadcast_weight, n_threads, sess->vtcm_size, + broadcast_weight, n_threads, sess->vtcm_size, elem_size, &col_tile, &vtcm_row_per_thread); kparams->col_tile = col_tile; @@ -4451,15 +4453,39 @@ static bool ggml_hexagon_supported_unary(const struct ggml_hexagon_session * ses const struct ggml_tensor * src0 = op->src[0]; const struct ggml_tensor * dst = op; - if (src0->type != GGML_TYPE_F32) { + if (src0->type != GGML_TYPE_F32 && src0->type != GGML_TYPE_F16) { return false; } - if (dst->type != GGML_TYPE_F32) { + if (dst->type != src0->type) { return false; } if (!ggml_is_contiguous_rows(src0)) { return false; } + + // F16 device kernels only cover this explicit whitelist (must stay in sync with + // the is_f16 whitelist in execute_op_unary(), unary-ops.c). + if (src0->type == GGML_TYPE_F16) { + switch (op->op) { + case GGML_OP_NORM: + case GGML_OP_RMS_NORM: + case GGML_OP_L2_NORM: + case GGML_OP_SCALE: + case GGML_OP_CLAMP: + case GGML_OP_SQR: + case GGML_OP_SQRT: + case GGML_OP_LOG: + break; + case GGML_OP_UNARY: + if (ggml_get_unary_op(op) != GGML_UNARY_OP_ABS) { + return false; + } + break; + default: + return false; + } + } + if (!ggml_are_same_shape(src0, dst)) { return false; } diff --git a/ggml/src/ggml-hexagon/htp/hvx-arith.h b/ggml/src/ggml-hexagon/htp/hvx-arith.h index 765c3577668..5ef7463426e 100644 --- a/ggml/src/ggml-hexagon/htp/hvx-arith.h +++ b/ggml/src/ggml-hexagon/htp/hvx-arith.h @@ -358,6 +358,54 @@ static inline void hvx_clamp_scalar_f32(uint8_t * restrict dst, const uint8_t * } } +#define HVX_OP_CLAMP_SCALAR_F16(v) \ + ({ \ + HVX_VectorPred pred_cap_right = Q6_Q_vcmp_gt_VhfVhf(v, max_vec); \ + HVX_VectorPred pred_cap_left = Q6_Q_vcmp_gt_VhfVhf(min_vec, v); \ + HVX_Vector tmp = Q6_V_vmux_QVV(pred_cap_right, max_vec, v); \ + Q6_V_vmux_QVV(pred_cap_left, min_vec, tmp); \ + }) + +static inline void hvx_clamp_scalar_f16_aa(uint8_t * restrict dst, const uint8_t * restrict src, const _Float16 min, const _Float16 max, uint32_t n) { + const HVX_Vector min_vec = hvx_vec_splat_f16(min); + const HVX_Vector max_vec = hvx_vec_splat_f16(max); + assert((unsigned long) dst % 128 == 0); + assert((unsigned long) src % 128 == 0); + hvx_scalar_loop_body(HVX_Vector, HVX_Vector, sizeof(_Float16), hvx_vec_store_a, HVX_OP_CLAMP_SCALAR_F16); +} + +static inline void hvx_clamp_scalar_f16_au(uint8_t * restrict dst, const uint8_t * restrict src, const _Float16 min, const _Float16 max, uint32_t n) { + const HVX_Vector min_vec = hvx_vec_splat_f16(min); + const HVX_Vector max_vec = hvx_vec_splat_f16(max); + assert((unsigned long) dst % 128 == 0); + hvx_scalar_loop_body(HVX_Vector, HVX_UVector, sizeof(_Float16), hvx_vec_store_a, HVX_OP_CLAMP_SCALAR_F16); +} + +static inline void hvx_clamp_scalar_f16_ua(uint8_t * restrict dst, const uint8_t * restrict src, const _Float16 min, const _Float16 max, uint32_t n) { + const HVX_Vector min_vec = hvx_vec_splat_f16(min); + const HVX_Vector max_vec = hvx_vec_splat_f16(max); + assert((unsigned long) src % 128 == 0); + hvx_scalar_loop_body(HVX_UVector, HVX_Vector, sizeof(_Float16), hvx_vec_store_u, HVX_OP_CLAMP_SCALAR_F16); +} + +static inline void hvx_clamp_scalar_f16_uu(uint8_t * restrict dst, const uint8_t * restrict src, const _Float16 min, const _Float16 max, uint32_t n) { + const HVX_Vector min_vec = hvx_vec_splat_f16(min); + const HVX_Vector max_vec = hvx_vec_splat_f16(max); + hvx_scalar_loop_body(HVX_UVector, HVX_UVector, sizeof(_Float16), hvx_vec_store_u, HVX_OP_CLAMP_SCALAR_F16); +} + +static inline void hvx_clamp_scalar_f16(uint8_t * restrict dst, const uint8_t * restrict src, const _Float16 min, const _Float16 max, const int num_elems) { + if (hex_is_aligned((void *) dst, 128) && hex_is_aligned((void *) src, 128)) { + hvx_clamp_scalar_f16_aa(dst, src, min, max, num_elems); + } else if (hex_is_aligned((void *) dst, 128)) { + hvx_clamp_scalar_f16_au(dst, src, min, max, num_elems); + } else if (hex_is_aligned((void *) src, 128)) { + hvx_clamp_scalar_f16_ua(dst, src, min, max, num_elems); + } else { + hvx_clamp_scalar_f16_uu(dst, src, min, max, num_elems); + } +} + // // Abs // @@ -386,11 +434,69 @@ static inline void hvx_abs_f32_aa(uint8_t * restrict dst, const uint8_t * restri } } +#define hvx_abs_f16_loop_body(dst_type, src_type, vec_store) \ + do { \ + dst_type * restrict vdst = (dst_type *) dst; \ + src_type * restrict vsrc = (src_type *) src; \ + \ + const uint32_t elem_size = sizeof(_Float16); \ + const uint32_t epv = 128 / elem_size; \ + const uint32_t nvec = n / epv; \ + const uint32_t nloe = n % epv; \ + \ + uint32_t i = 0; \ + \ + _Pragma("unroll(4)") \ + for (; i < nvec; i++) { \ + vdst[i] = hvx_vec_abs_f16(vsrc[i]); \ + } \ + if (nloe) { \ + HVX_Vector v = hvx_vec_abs_f16(vsrc[i]); \ + vec_store((void *) &vdst[i], nloe * elem_size, v); \ + } \ + } while(0) + +static inline void hvx_abs_f16_aa(uint8_t * restrict dst, const uint8_t * restrict src, uint32_t n) { + assert((unsigned long) dst % 128 == 0); + assert((unsigned long) src % 128 == 0); + hvx_abs_f16_loop_body(HVX_Vector, HVX_Vector, hvx_vec_store_a); +} + +static inline void hvx_abs_f16_au(uint8_t * restrict dst, const uint8_t * restrict src, uint32_t n) { + assert((unsigned long) dst % 128 == 0); + hvx_abs_f16_loop_body(HVX_Vector, HVX_UVector, hvx_vec_store_a); +} + +static inline void hvx_abs_f16_ua(uint8_t * restrict dst, const uint8_t * restrict src, uint32_t n) { + assert((unsigned long) src % 128 == 0); + hvx_abs_f16_loop_body(HVX_UVector, HVX_Vector, hvx_vec_store_u); +} + +static inline void hvx_abs_f16_uu(uint8_t * restrict dst, const uint8_t * restrict src, uint32_t n) { + hvx_abs_f16_loop_body(HVX_UVector, HVX_UVector, hvx_vec_store_u); +} + +static inline void hvx_abs_f16(uint8_t * restrict dst, const uint8_t * restrict src, const uint32_t num_elems) { + if (hex_is_aligned((void *) dst, 128)) { + if (hex_is_aligned((void *) src, 128)) { + hvx_abs_f16_aa(dst, src, num_elems); + } else { + hvx_abs_f16_au(dst, src, num_elems); + } + } else { + if (hex_is_aligned((void *) src, 128)) { + hvx_abs_f16_ua(dst, src, num_elems); + } else { + hvx_abs_f16_uu(dst, src, num_elems); + } + } +} + // // Square // -#define hvx_sqr_f32_loop_body(dst_type, src_type, vec_store) \ +#define hvx_sqr_f32_loop_body(dst_type, src_type, vec_store) \ do { \ dst_type * restrict vdst = (dst_type *) dst; \ src_type * restrict vsrc = (src_type *) src; \ @@ -404,10 +510,10 @@ static inline void hvx_abs_f32_aa(uint8_t * restrict dst, const uint8_t * restri \ _Pragma("unroll(4)") \ for (; i < nvec; i++) { \ - vdst[i] = HVX_OP_MUL_F32(vsrc[i], vsrc[i]); \ + vdst[i] = HVX_OP_MUL_F32(vsrc[i], vsrc[i]); \ } \ if (nloe) { \ - HVX_Vector v = HVX_OP_MUL_F32(vsrc[i], vsrc[i]); \ + HVX_Vector v = HVX_OP_MUL_F32(vsrc[i], vsrc[i]); \ vec_store((void *) &vdst[i], nloe * elem_size, v); \ } \ } while(0) @@ -448,6 +554,64 @@ static inline void hvx_sqr_f32(uint8_t * restrict dst, const uint8_t * restrict } } +#define hvx_sqr_f16_loop_body(dst_type, src_type, vec_store) \ + do { \ + dst_type * restrict vdst = (dst_type *) dst; \ + src_type * restrict vsrc = (src_type *) src; \ + \ + const uint32_t elem_size = sizeof(_Float16); \ + const uint32_t epv = 128 / elem_size; \ + const uint32_t nvec = n / epv; \ + const uint32_t nloe = n % epv; \ + \ + uint32_t i = 0; \ + \ + _Pragma("unroll(4)") \ + for (; i < nvec; i++) { \ + vdst[i] = HVX_OP_MUL_F16(vsrc[i], vsrc[i]); \ + } \ + if (nloe) { \ + HVX_Vector v = HVX_OP_MUL_F16(vsrc[i], vsrc[i]); \ + vec_store((void *) &vdst[i], nloe * elem_size, v); \ + } \ + } while(0) + +static inline void hvx_sqr_f16_aa(uint8_t * restrict dst, const uint8_t * restrict src, uint32_t n) { + assert((unsigned long) dst % 128 == 0); + assert((unsigned long) src % 128 == 0); + hvx_sqr_f16_loop_body(HVX_Vector, HVX_Vector, hvx_vec_store_a); +} + +static inline void hvx_sqr_f16_au(uint8_t * restrict dst, const uint8_t * restrict src, uint32_t n) { + assert((unsigned long) dst % 128 == 0); + hvx_sqr_f16_loop_body(HVX_Vector, HVX_UVector, hvx_vec_store_a); +} + +static inline void hvx_sqr_f16_ua(uint8_t * restrict dst, const uint8_t * restrict src, uint32_t n) { + assert((unsigned long) src % 128 == 0); + hvx_sqr_f16_loop_body(HVX_UVector, HVX_Vector, hvx_vec_store_u); +} + +static inline void hvx_sqr_f16_uu(uint8_t * restrict dst, const uint8_t * restrict src, uint32_t n) { + hvx_sqr_f16_loop_body(HVX_UVector, HVX_UVector, hvx_vec_store_u); +} + +static inline void hvx_sqr_f16(uint8_t * restrict dst, const uint8_t * restrict src, const uint32_t num_elems) { + if (hex_is_aligned((void *) dst, 128)) { + if (hex_is_aligned((void *) src, 128)) { + hvx_sqr_f16_aa(dst, src, num_elems); + } else { + hvx_sqr_f16_au(dst, src, num_elems); + } + } else { + if (hex_is_aligned((void *) src, 128)) { + hvx_sqr_f16_ua(dst, src, num_elems); + } else { + hvx_sqr_f16_uu(dst, src, num_elems); + } + } +} + #undef HVX_OP_ADD_F32 #undef HVX_OP_SUB_F32 #undef HVX_OP_MUL_F32 @@ -464,6 +628,7 @@ static inline void hvx_sqr_f32(uint8_t * restrict dst, const uint8_t * restrict #undef hvx_scalar_loop_body #undef HVX_OP_MIN_SCALAR #undef HVX_OP_CLAMP_SCALAR +#undef HVX_OP_CLAMP_SCALAR_F16 #undef DEFINE_HVX_BINARY_OP_VARIANTS #undef HVX_BINARY_DISPATCHER #undef UNUSED diff --git a/ggml/src/ggml-hexagon/htp/hvx-log.h b/ggml/src/ggml-hexagon/htp/hvx-log.h index a209f88d555..491041d5ad5 100644 --- a/ggml/src/ggml-hexagon/htp/hvx-log.h +++ b/ggml/src/ggml-hexagon/htp/hvx-log.h @@ -86,4 +86,33 @@ static inline void hvx_log_f32_aa(uint8_t * restrict dst, const uint8_t * restri } } +// Compute log(x) for f16 by promoting to f32, applying hvx_vec_log_f32, and narrowing back. +static inline void hvx_log_f16_aa(uint8_t * restrict dst, const uint8_t * restrict src, uint32_t n) { + assert((unsigned long) dst % 128 == 0); + assert((unsigned long) src % 128 == 0); + + HVX_Vector * restrict vdst = (HVX_Vector *) dst; + HVX_Vector * restrict vsrc = (HVX_Vector *) src; + + const uint32_t nvec = n / VLEN_FP16; + const uint32_t nloe = n % VLEN_FP16; + + uint32_t i = 0; + + _Pragma("unroll(4)") + for (; i < nvec; i++) { + HVX_VectorPair p = hvx_vec_f16_to_f32(vsrc[i]); + HVX_Vector r0 = hvx_vec_log_f32(Q6_V_lo_W(p)); + HVX_Vector r1 = hvx_vec_log_f32(Q6_V_hi_W(p)); + vdst[i] = hvx_vec_f32_to_f16(r0, r1); + } + if (nloe) { + HVX_VectorPair p = hvx_vec_f16_to_f32(vsrc[i]); + HVX_Vector r0 = hvx_vec_log_f32(Q6_V_lo_W(p)); + HVX_Vector r1 = hvx_vec_log_f32(Q6_V_hi_W(p)); + HVX_Vector v = hvx_vec_f32_to_f16(r0, r1); + hvx_vec_store_a((void *) &vdst[i], nloe * SIZEOF_FP16, v); + } +} + #endif /* HVX_LOG_H */ diff --git a/ggml/src/ggml-hexagon/htp/hvx-norm.h b/ggml/src/ggml-hexagon/htp/hvx-norm.h index a8645e412d3..7ea945a339c 100644 --- a/ggml/src/ggml-hexagon/htp/hvx-norm.h +++ b/ggml/src/ggml-hexagon/htp/hvx-norm.h @@ -254,4 +254,201 @@ static inline void hvx_fast_l2_norm_f32(const uint8_t * restrict src, } } +// F16 norm kernels: reduce and scale in f32 (via promote/narrow), matching the +// precision-preserving pattern used by the flash-attn f16 kernels. + +static inline void hvx_fast_rms_norm_f16(const uint8_t * restrict src, + uint8_t * restrict dst, + const int num_elems, + float epsilon) { + + const HVX_Vector * restrict v_src = (HVX_Vector *) src; + HVX_Vector * restrict v_dst = (HVX_Vector *) dst; + + const int nvec = num_elems / VLEN_FP16; // number of full f16 vectors + const int nloe = num_elems % VLEN_FP16; // leftover elements + + HVX_Vector sum_v = Q6_V_vsplat_R(0x00000000); + HVX_Vector epsilon_v = hvx_vec_splat_f32(epsilon); + + #pragma unroll(4) + for (int i = 0; i < nvec; i++) { + HVX_VectorPair p = hvx_vec_f16_to_f32(v_src[i]); + HVX_Vector p0 = Q6_V_lo_W(p); + HVX_Vector p1 = Q6_V_hi_W(p); + sum_v = Q6_Vqf32_vadd_Vqf32Vqf32(sum_v, Q6_Vqf32_vmpy_VsfVsf(p0, p0)); + sum_v = Q6_Vqf32_vadd_Vqf32Vqf32(sum_v, Q6_Vqf32_vmpy_VsfVsf(p1, p1)); + } + + if (nloe > 0) { + HVX_VectorPred bmask = Q6_Q_vsetq_R(nloe * SIZEOF_FP16); + HVX_Vector v1 = Q6_V_vand_QV(bmask, v_src[nvec]); + HVX_VectorPair p = hvx_vec_f16_to_f32(v1); + HVX_Vector p0 = Q6_V_lo_W(p); + HVX_Vector p1 = Q6_V_hi_W(p); + sum_v = Q6_Vqf32_vadd_Vqf32Vqf32(sum_v, Q6_Vqf32_vmpy_VsfVsf(p0, p0)); + sum_v = Q6_Vqf32_vadd_Vqf32Vqf32(sum_v, Q6_Vqf32_vmpy_VsfVsf(p1, p1)); + } + + sum_v = hvx_vec_reduce_sum_f32(Q6_Vsf_equals_Vqf32(sum_v)); + + HVX_Vector t_v = hvx_vec_splat_f32((float) num_elems); + HVX_Vector denom_v = hvx_vec_inverse_f32(t_v); + HVX_Vector mean_v = Q6_Vqf32_vmpy_VsfVsf(sum_v, denom_v); + HVX_Vector mean_epsilon_v = Q6_Vqf32_vadd_Vqf32Vsf(mean_v, epsilon_v); + + HVX_Vector scale_v = hvx_vec_rsqrt_f32(Q6_Vsf_equals_Vqf32(mean_epsilon_v)); + + #pragma unroll(4) + for (int i = 0; i < nvec; i++) { + HVX_VectorPair p = hvx_vec_f16_to_f32(v_src[i]); + HVX_Vector r0 = Q6_Vsf_equals_Vqf32(Q6_Vqf32_vmpy_VsfVsf(Q6_V_lo_W(p), scale_v)); + HVX_Vector r1 = Q6_Vsf_equals_Vqf32(Q6_Vqf32_vmpy_VsfVsf(Q6_V_hi_W(p), scale_v)); + v_dst[i] = hvx_vec_f32_to_f16(r0, r1); + } + + if (nloe > 0) { + HVX_VectorPred bmask = Q6_Q_vsetq_R(nloe * SIZEOF_FP16); + HVX_Vector v1 = Q6_V_vand_QV(bmask, v_src[nvec]); + HVX_VectorPair p = hvx_vec_f16_to_f32(v1); + HVX_Vector r0 = Q6_Vsf_equals_Vqf32(Q6_Vqf32_vmpy_VsfVsf(Q6_V_lo_W(p), scale_v)); + HVX_Vector r1 = Q6_Vsf_equals_Vqf32(Q6_Vqf32_vmpy_VsfVsf(Q6_V_hi_W(p), scale_v)); + HVX_Vector result = hvx_vec_f32_to_f16(r0, r1); + hvx_vec_store_a(&v_dst[nvec], nloe * SIZEOF_FP16, result); + } +} + +static inline void hvx_fast_norm_f16(const uint8_t * restrict src, + uint8_t * restrict dst, + const int num_elems, + float epsilon) { + + const HVX_Vector * restrict v_src = (HVX_Vector *) src; + HVX_Vector * restrict v_dst = (HVX_Vector *) dst; + + const int nvec = num_elems / VLEN_FP16; + const int nloe = num_elems % VLEN_FP16; + + HVX_Vector sum_sq_v = Q6_V_vsplat_R(0x00000000); + HVX_Vector sum_x_v = Q6_V_vsplat_R(0x00000000); + HVX_Vector epsilon_v = hvx_vec_splat_f32(epsilon); + + #pragma unroll(4) + for (int i = 0; i < nvec; i++) { + HVX_VectorPair p = hvx_vec_f16_to_f32(v_src[i]); + HVX_Vector p0 = Q6_V_lo_W(p); + HVX_Vector p1 = Q6_V_hi_W(p); + sum_sq_v = Q6_Vqf32_vadd_Vqf32Vqf32(sum_sq_v, Q6_Vqf32_vmpy_VsfVsf(p0, p0)); + sum_sq_v = Q6_Vqf32_vadd_Vqf32Vqf32(sum_sq_v, Q6_Vqf32_vmpy_VsfVsf(p1, p1)); + sum_x_v = Q6_Vqf32_vadd_Vqf32Vqf32(sum_x_v, Q6_Vqf32_vadd_VsfVsf(p0, Q6_V_vzero())); + sum_x_v = Q6_Vqf32_vadd_Vqf32Vqf32(sum_x_v, Q6_Vqf32_vadd_VsfVsf(p1, Q6_V_vzero())); + } + + if (nloe > 0) { + HVX_VectorPred bmask = Q6_Q_vsetq_R(nloe * SIZEOF_FP16); + HVX_Vector v1 = Q6_V_vand_QV(bmask, v_src[nvec]); + HVX_VectorPair p = hvx_vec_f16_to_f32(v1); + HVX_Vector p0 = Q6_V_lo_W(p); + HVX_Vector p1 = Q6_V_hi_W(p); + sum_sq_v = Q6_Vqf32_vadd_Vqf32Vqf32(sum_sq_v, Q6_Vqf32_vmpy_VsfVsf(p0, p0)); + sum_sq_v = Q6_Vqf32_vadd_Vqf32Vqf32(sum_sq_v, Q6_Vqf32_vmpy_VsfVsf(p1, p1)); + sum_x_v = Q6_Vqf32_vadd_Vqf32Vqf32(sum_x_v, Q6_Vqf32_vadd_VsfVsf(p0, Q6_V_vzero())); + sum_x_v = Q6_Vqf32_vadd_Vqf32Vqf32(sum_x_v, Q6_Vqf32_vadd_VsfVsf(p1, Q6_V_vzero())); + } + + sum_sq_v = hvx_vec_reduce_sum_f32(Q6_Vsf_equals_Vqf32(sum_sq_v)); + sum_x_v = hvx_vec_reduce_sum_f32(Q6_Vsf_equals_Vqf32(sum_x_v)); + + HVX_Vector t_v = hvx_vec_splat_f32((float) num_elems); + HVX_Vector denom_v = hvx_vec_inverse_f32(t_v); + HVX_Vector mean_sq_v = Q6_Vqf32_vmpy_VsfVsf(sum_sq_v, denom_v); + HVX_Vector mean_x_v = Q6_Vqf32_vmpy_VsfVsf(sum_x_v, denom_v); + HVX_Vector mean_x_sq_v = Q6_Vqf32_vmpy_VsfVsf(Q6_Vsf_equals_Vqf32(mean_x_v), Q6_Vsf_equals_Vqf32(mean_x_v)); + HVX_Vector var_v = Q6_Vqf32_vsub_Vqf32Vqf32(mean_sq_v, mean_x_sq_v); + HVX_Vector var_epsilon_v = Q6_Vqf32_vadd_Vqf32Vsf(var_v, epsilon_v); + + HVX_Vector scale_v = hvx_vec_rsqrt_f32(Q6_Vsf_equals_Vqf32(var_epsilon_v)); + HVX_Vector mean_x_b = hvx_vec_repl_f32(Q6_Vsf_equals_Vqf32(mean_x_v)); + + #pragma unroll(4) + for (int i = 0; i < nvec; i++) { + HVX_VectorPair p = hvx_vec_f16_to_f32(v_src[i]); + HVX_Vector d0 = Q6_Vqf32_vsub_VsfVsf(Q6_V_lo_W(p), mean_x_b); + HVX_Vector d1 = Q6_Vqf32_vsub_VsfVsf(Q6_V_hi_W(p), mean_x_b); + HVX_Vector r0 = Q6_Vsf_equals_Vqf32(Q6_Vqf32_vmpy_VsfVsf(Q6_Vsf_equals_Vqf32(d0), scale_v)); + HVX_Vector r1 = Q6_Vsf_equals_Vqf32(Q6_Vqf32_vmpy_VsfVsf(Q6_Vsf_equals_Vqf32(d1), scale_v)); + v_dst[i] = hvx_vec_f32_to_f16(r0, r1); + } + + if (nloe > 0) { + HVX_VectorPred bmask = Q6_Q_vsetq_R(nloe * SIZEOF_FP16); + HVX_Vector v1 = Q6_V_vand_QV(bmask, v_src[nvec]); + HVX_VectorPair p = hvx_vec_f16_to_f32(v1); + HVX_Vector d0 = Q6_Vqf32_vsub_VsfVsf(Q6_V_lo_W(p), mean_x_b); + HVX_Vector d1 = Q6_Vqf32_vsub_VsfVsf(Q6_V_hi_W(p), mean_x_b); + HVX_Vector r0 = Q6_Vsf_equals_Vqf32(Q6_Vqf32_vmpy_VsfVsf(Q6_Vsf_equals_Vqf32(d0), scale_v)); + HVX_Vector r1 = Q6_Vsf_equals_Vqf32(Q6_Vqf32_vmpy_VsfVsf(Q6_Vsf_equals_Vqf32(d1), scale_v)); + HVX_Vector result = hvx_vec_f32_to_f16(r0, r1); + hvx_vec_store_a(&v_dst[nvec], nloe * SIZEOF_FP16, result); + } +} + +static inline void hvx_fast_l2_norm_f16(const uint8_t * restrict src, + uint8_t * restrict dst, + const int num_elems, + float epsilon) { + + const HVX_Vector * restrict v_src = (HVX_Vector *) src; + HVX_Vector * restrict v_dst = (HVX_Vector *) dst; + + const int nvec = num_elems / VLEN_FP16; + const int nloe = num_elems % VLEN_FP16; + + HVX_Vector sum_v = hvx_vec_splat_f32(0.0f); + + #pragma unroll(4) + for (int i = 0; i < nvec; i++) { + HVX_VectorPair p = hvx_vec_f16_to_f32(v_src[i]); + HVX_Vector p0 = Q6_V_lo_W(p); + HVX_Vector p1 = Q6_V_hi_W(p); + sum_v = Q6_Vqf32_vadd_Vqf32Vqf32(sum_v, Q6_Vqf32_vmpy_VsfVsf(p0, p0)); + sum_v = Q6_Vqf32_vadd_Vqf32Vqf32(sum_v, Q6_Vqf32_vmpy_VsfVsf(p1, p1)); + } + + if (nloe > 0) { + HVX_VectorPred bmask = Q6_Q_vsetq_R(nloe * SIZEOF_FP16); + HVX_Vector v1 = Q6_V_vand_QV(bmask, v_src[nvec]); + HVX_VectorPair p = hvx_vec_f16_to_f32(v1); + HVX_Vector p0 = Q6_V_lo_W(p); + HVX_Vector p1 = Q6_V_hi_W(p); + sum_v = Q6_Vqf32_vadd_Vqf32Vqf32(sum_v, Q6_Vqf32_vmpy_VsfVsf(p0, p0)); + sum_v = Q6_Vqf32_vadd_Vqf32Vqf32(sum_v, Q6_Vqf32_vmpy_VsfVsf(p1, p1)); + } + + HVX_Vector sum_sf = hvx_vec_reduce_sum_f32(Q6_Vsf_equals_Vqf32(sum_v)); + HVX_Vector rsqrt_v = hvx_vec_rsqrt_f32(sum_sf); + HVX_Vector sqrt_v = hvx_vec_inverse_f32(rsqrt_v); + HVX_Vector epsilon_v = hvx_vec_splat_f32(epsilon); + HVX_Vector denom_v = Q6_Vsf_vmax_VsfVsf(sqrt_v, epsilon_v); + HVX_Vector scale_v = hvx_vec_inverse_f32(denom_v); + + #pragma unroll(4) + for (int i = 0; i < nvec; i++) { + HVX_VectorPair p = hvx_vec_f16_to_f32(v_src[i]); + HVX_Vector r0 = Q6_Vsf_equals_Vqf32(Q6_Vqf32_vmpy_VsfVsf(Q6_V_lo_W(p), scale_v)); + HVX_Vector r1 = Q6_Vsf_equals_Vqf32(Q6_Vqf32_vmpy_VsfVsf(Q6_V_hi_W(p), scale_v)); + v_dst[i] = hvx_vec_f32_to_f16(r0, r1); + } + + if (nloe > 0) { + HVX_VectorPred bmask = Q6_Q_vsetq_R(nloe * SIZEOF_FP16); + HVX_Vector v1 = Q6_V_vand_QV(bmask, v_src[nvec]); + HVX_VectorPair p = hvx_vec_f16_to_f32(v1); + HVX_Vector r0 = Q6_Vsf_equals_Vqf32(Q6_Vqf32_vmpy_VsfVsf(Q6_V_lo_W(p), scale_v)); + HVX_Vector r1 = Q6_Vsf_equals_Vqf32(Q6_Vqf32_vmpy_VsfVsf(Q6_V_hi_W(p), scale_v)); + HVX_Vector result = hvx_vec_f32_to_f16(r0, r1); + hvx_vec_store_a(&v_dst[nvec], nloe * SIZEOF_FP16, result); + } +} + #endif // HVX_NORM_H diff --git a/ggml/src/ggml-hexagon/htp/hvx-scale.h b/ggml/src/ggml-hexagon/htp/hvx-scale.h index c65c98639dc..9b1a28f529a 100644 --- a/ggml/src/ggml-hexagon/htp/hvx-scale.h +++ b/ggml/src/ggml-hexagon/htp/hvx-scale.h @@ -130,4 +130,70 @@ static inline void hvx_scale_offset_f32(uint8_t * restrict dst, const uint8_t * } } +// Scale+offset computed by promoting f16 -> f32, then narrowing the result back to f16. +#define hvx_scale_offset_f16_loop_body(dst_type, src_type, vec_store) \ + do { \ + dst_type * restrict vdst = (dst_type *) dst; \ + src_type * restrict vsrc = (src_type *) src; \ + \ + HVX_Vector vs = hvx_vec_splat_f32(scale); \ + HVX_Vector vo = hvx_vec_splat_f32(offset); \ + \ + const uint32_t nvec = n / VLEN_FP16; \ + const uint32_t nloe = n % VLEN_FP16; \ + \ + uint32_t i = 0; \ + \ + _Pragma("unroll(4)") \ + for (; i < nvec; ++i) { \ + HVX_VectorPair p = hvx_vec_f16_to_f32(vsrc[i]); \ + HVX_Vector r0 = Q6_Vsf_equals_Vqf32(Q6_Vqf32_vadd_Vqf32Vsf(Q6_Vqf32_vmpy_VsfVsf(Q6_V_lo_W(p), vs), vo)); \ + HVX_Vector r1 = Q6_Vsf_equals_Vqf32(Q6_Vqf32_vadd_Vqf32Vsf(Q6_Vqf32_vmpy_VsfVsf(Q6_V_hi_W(p), vs), vo)); \ + vdst[i] = hvx_vec_f32_to_f16(r0, r1); \ + } \ + if (nloe) { \ + HVX_VectorPair p = hvx_vec_f16_to_f32(vsrc[i]); \ + HVX_Vector r0 = Q6_Vsf_equals_Vqf32(Q6_Vqf32_vadd_Vqf32Vsf(Q6_Vqf32_vmpy_VsfVsf(Q6_V_lo_W(p), vs), vo)); \ + HVX_Vector r1 = Q6_Vsf_equals_Vqf32(Q6_Vqf32_vadd_Vqf32Vsf(Q6_Vqf32_vmpy_VsfVsf(Q6_V_hi_W(p), vs), vo)); \ + HVX_Vector v = hvx_vec_f32_to_f16(r0, r1); \ + vec_store((void *) &vdst[i], nloe * SIZEOF_FP16, v); \ + } \ + } while(0) + +static inline void hvx_scale_offset_f16_aa(uint8_t * restrict dst, const uint8_t * restrict src, const int n, const float scale, const float offset) { + assert((size_t) dst % 128 == 0); + assert((size_t) src % 128 == 0); + hvx_scale_offset_f16_loop_body(HVX_Vector, HVX_Vector, hvx_vec_store_a); +} + +static inline void hvx_scale_offset_f16_au(uint8_t * restrict dst, const uint8_t * restrict src, const int n, const float scale, const float offset) { + assert((size_t) dst % 128 == 0); + hvx_scale_offset_f16_loop_body(HVX_Vector, HVX_UVector, hvx_vec_store_a); +} + +static inline void hvx_scale_offset_f16_ua(uint8_t * restrict dst, const uint8_t * restrict src, const int n, const float scale, const float offset) { + assert((size_t) src % 128 == 0); + hvx_scale_offset_f16_loop_body(HVX_UVector, HVX_Vector, hvx_vec_store_u); +} + +static inline void hvx_scale_offset_f16_uu(uint8_t * restrict dst, const uint8_t * restrict src, const int n, const float scale, const float offset) { + hvx_scale_offset_f16_loop_body(HVX_UVector, HVX_UVector, hvx_vec_store_u); +} + +static inline void hvx_scale_offset_f16(uint8_t * restrict dst, const uint8_t * restrict src, const int n, const float scale, const float offset) { + if (((size_t) dst & 127) == 0) { + if (((size_t) src & 127) == 0) { + hvx_scale_offset_f16_aa(dst, src, n, scale, offset); + } else { + hvx_scale_offset_f16_au(dst, src, n, scale, offset); + } + } else { + if (((size_t) src & 127) == 0) { + hvx_scale_offset_f16_ua(dst, src, n, scale, offset); + } else { + hvx_scale_offset_f16_uu(dst, src, n, scale, offset); + } + } +} + #endif // HVX_SCALE_H diff --git a/ggml/src/ggml-hexagon/htp/hvx-sqrt.h b/ggml/src/ggml-hexagon/htp/hvx-sqrt.h index e31a1006d21..abdded5ce69 100644 --- a/ggml/src/ggml-hexagon/htp/hvx-sqrt.h +++ b/ggml/src/ggml-hexagon/htp/hvx-sqrt.h @@ -123,4 +123,67 @@ static inline void hvx_sqrt_f32(uint8_t * restrict dst, const uint8_t * restrict } } +// Compute sqrt(x) for f16 by promoting to f32, applying hvx_vec_rsqrt_f32, and narrowing back. +#define hvx_sqrt_f16_loop_body(dst_type, src_type, vec_store) \ + do { \ + dst_type * restrict vdst = (dst_type *) dst; \ + src_type * restrict vsrc = (src_type *) src; \ + \ + const uint32_t nvec = n / VLEN_FP16; \ + const uint32_t nloe = n % VLEN_FP16; \ + \ + uint32_t i = 0; \ + \ + _Pragma("unroll(4)") \ + for (; i < nvec; i++) { \ + HVX_VectorPair p = hvx_vec_f16_to_f32(vsrc[i]); \ + HVX_Vector r0 = HVX_OP_MUL(hvx_vec_rsqrt_f32(Q6_V_lo_W(p)), Q6_V_lo_W(p)); \ + HVX_Vector r1 = HVX_OP_MUL(hvx_vec_rsqrt_f32(Q6_V_hi_W(p)), Q6_V_hi_W(p)); \ + vdst[i] = hvx_vec_f32_to_f16(r0, r1); \ + } \ + if (nloe) { \ + HVX_VectorPair p = hvx_vec_f16_to_f32(vsrc[i]); \ + HVX_Vector r0 = HVX_OP_MUL(hvx_vec_rsqrt_f32(Q6_V_lo_W(p)), Q6_V_lo_W(p)); \ + HVX_Vector r1 = HVX_OP_MUL(hvx_vec_rsqrt_f32(Q6_V_hi_W(p)), Q6_V_hi_W(p)); \ + HVX_Vector v = hvx_vec_f32_to_f16(r0, r1); \ + vec_store((void *) &vdst[i], nloe * SIZEOF_FP16, v); \ + } \ + } while(0) + +static inline void hvx_sqrt_f16_aa(uint8_t * restrict dst, const uint8_t * restrict src, uint32_t n) { + assert((unsigned long) dst % 128 == 0); + assert((unsigned long) src % 128 == 0); + hvx_sqrt_f16_loop_body(HVX_Vector, HVX_Vector, hvx_vec_store_a); +} + +static inline void hvx_sqrt_f16_au(uint8_t * restrict dst, const uint8_t * restrict src, uint32_t n) { + assert((unsigned long) dst % 128 == 0); + hvx_sqrt_f16_loop_body(HVX_Vector, HVX_UVector, hvx_vec_store_a); +} + +static inline void hvx_sqrt_f16_ua(uint8_t * restrict dst, const uint8_t * restrict src, uint32_t n) { + assert((unsigned long) src % 128 == 0); + hvx_sqrt_f16_loop_body(HVX_UVector, HVX_Vector, hvx_vec_store_u); +} + +static inline void hvx_sqrt_f16_uu(uint8_t * restrict dst, const uint8_t * restrict src, uint32_t n) { + hvx_sqrt_f16_loop_body(HVX_UVector, HVX_UVector, hvx_vec_store_u); +} + +static inline void hvx_sqrt_f16(uint8_t * restrict dst, const uint8_t * restrict src, const int num_elems) { + if ((unsigned long) dst % 128 == 0) { + if ((unsigned long) src % 128 == 0) { + hvx_sqrt_f16_aa(dst, src, num_elems); + } else { + hvx_sqrt_f16_au(dst, src, num_elems); + } + } else { + if ((unsigned long) src % 128 == 0) { + hvx_sqrt_f16_ua(dst, src, num_elems); + } else { + hvx_sqrt_f16_uu(dst, src, num_elems); + } + } +} + #endif /* HVX_SQRT_H */ diff --git a/ggml/src/ggml-hexagon/htp/unary-ops.c b/ggml/src/ggml-hexagon/htp/unary-ops.c index 1a632bf5631..5e62b4a9bd3 100644 --- a/ggml/src/ggml-hexagon/htp/unary-ops.c +++ b/ggml/src/ggml-hexagon/htp/unary-ops.c @@ -234,6 +234,146 @@ static void sqrt_f32(const float * restrict src, } } +static void scale_f16(const _Float16 * restrict src, + _Float16 * restrict dst, + const uint32_t num_rows, + const struct htp_unary_context * uctx) { + htp_unary_op_preamble; + float scale = 0.f; + float bias = 0.f; + memcpy(&scale, &op_params[0], sizeof(float)); + memcpy(&bias, &op_params[1], sizeof(float)); + + for (uint32_t ir = 0; ir < num_rows; ir++) { + const uint8_t * restrict src_local = (const uint8_t *)src + (ir * src0_row_size_aligned); + uint8_t * restrict dst_local = (uint8_t *)dst + (ir * dst_row_size_aligned); + + hvx_scale_offset_f16_aa((uint8_t *) dst_local, (const uint8_t *) src_local, ne0, scale, bias); + } +} + +static void clamp_f16(const _Float16 * restrict src, + _Float16 * restrict dst, + const uint32_t num_rows, + const struct htp_unary_context * uctx) { + htp_unary_op_preamble; + float min = 0.f; + float max = 0.f; + memcpy(&min, &op_params[0], sizeof(float)); + memcpy(&max, &op_params[1], sizeof(float)); + + for (uint32_t ir = 0; ir < num_rows; ir++) { + const uint8_t * restrict src_local = (const uint8_t *)src + (ir * src0_row_size_aligned); + uint8_t * restrict dst_local = (uint8_t *)dst + (ir * dst_row_size_aligned); + + hvx_clamp_scalar_f16(dst_local, src_local, (_Float16) min, (_Float16) max, ne0); + } +} + +static void rms_norm_f16(const _Float16 * restrict src, + _Float16 * restrict dst, + const uint32_t num_rows, + const struct htp_unary_context * uctx) { + htp_unary_op_preamble; + float epsilon = 0.f; + memcpy(&epsilon, op_params, sizeof(float)); + + for (uint32_t ir = 0; ir < num_rows; ir++) { + const uint8_t * restrict src_local = (const uint8_t *)src + (ir * src0_row_size_aligned); + uint8_t * restrict dst_local = (uint8_t *)dst + (ir * dst_row_size_aligned); + + hvx_fast_rms_norm_f16((const uint8_t *) src_local, (uint8_t *) dst_local, ne0, epsilon); + } +} + +static void norm_f16(const _Float16 * restrict src, + _Float16 * restrict dst, + const uint32_t num_rows, + const struct htp_unary_context * uctx) { + htp_unary_op_preamble; + float epsilon = 0.f; + memcpy(&epsilon, op_params, sizeof(float)); + + for (uint32_t ir = 0; ir < num_rows; ir++) { + const uint8_t * restrict src_local = (const uint8_t *)src + (ir * src0_row_size_aligned); + uint8_t * restrict dst_local = (uint8_t *)dst + (ir * dst_row_size_aligned); + + hvx_fast_norm_f16((const uint8_t *) src_local, (uint8_t *) dst_local, ne0, epsilon); + } +} + +static void sqr_f16(const _Float16 * restrict src, + _Float16 * restrict dst, + const uint32_t num_rows, + const struct htp_unary_context * uctx) { + htp_unary_op_preamble; + + for (uint32_t ir = 0; ir < num_rows; ir++) { + const uint8_t * restrict src_local = (const uint8_t *)src + (ir * src0_row_size_aligned); + uint8_t * restrict dst_local = (uint8_t *)dst + (ir * dst_row_size_aligned); + + hvx_sqr_f16_aa((uint8_t *) dst_local, (const uint8_t *) src_local, ne0); + } +} + +static void sqrt_f16(const _Float16 * restrict src, + _Float16 * restrict dst, + const uint32_t num_rows, + const struct htp_unary_context * uctx) { + htp_unary_op_preamble; + + for (uint32_t ir = 0; ir < num_rows; ir++) { + const uint8_t * restrict src_local = (const uint8_t *)src + (ir * src0_row_size_aligned); + uint8_t * restrict dst_local = (uint8_t *)dst + (ir * dst_row_size_aligned); + + hvx_sqrt_f16_aa((uint8_t *) dst_local, (const uint8_t *) src_local, ne0); + } +} + +static void abs_f16(const _Float16 * restrict src, + _Float16 * restrict dst, + const uint32_t num_rows, + const struct htp_unary_context * uctx) { + htp_unary_op_preamble; + + for (uint32_t ir = 0; ir < num_rows; ir++) { + const uint8_t * restrict src_local = (const uint8_t *)src + (ir * src0_row_size_aligned); + uint8_t * restrict dst_local = (uint8_t *)dst + (ir * dst_row_size_aligned); + + hvx_abs_f16_aa((uint8_t *) dst_local, (const uint8_t *) src_local, ne0); + } +} + +static void log_f16(const _Float16 * restrict src, + _Float16 * restrict dst, + const uint32_t num_rows, + const struct htp_unary_context * uctx) { + htp_unary_op_preamble; + + for (uint32_t ir = 0; ir < num_rows; ir++) { + const uint8_t * restrict src_local = (const uint8_t *)src + (ir * src0_row_size_aligned); + uint8_t * restrict dst_local = (uint8_t *)dst + (ir * dst_row_size_aligned); + + hvx_log_f16_aa((uint8_t *) dst_local, (const uint8_t *) src_local, ne0); + } +} + +static void l2_norm_f16(const _Float16 * restrict src, + _Float16 * restrict dst, + const uint32_t num_rows, + const struct htp_unary_context * uctx) { + htp_unary_op_preamble; + float epsilon = 0.f; + memcpy(&epsilon, op_params, sizeof(float)); + + for (uint32_t ir = 0; ir < num_rows; ir++) { + const uint8_t * restrict src_f = (const uint8_t *)src + (ir * src0_row_size_aligned); + uint8_t * restrict dst_f = (uint8_t *)dst + (ir * dst_row_size_aligned); + + hvx_fast_l2_norm_f16((const uint8_t *)src_f, (uint8_t *)dst_f, ne0, epsilon); + } +} + static void neg_f32(const float * restrict src, float * restrict dst, const uint32_t num_rows, @@ -471,8 +611,8 @@ static void log_f32(const float * restrict src, } } -#define DEFINE_UNARY_TASK(NAME, IS_RMS_NORM_MUL, IS_TRI, CORE_EXPR) \ -static void unary_task_f32_##NAME(unsigned int nth, unsigned int ith, void * data) { \ +#define DEFINE_UNARY_TASK_IMPL(NAME, TYPE, SUFFIX, IS_RMS_NORM_MUL, IS_TRI, CORE_EXPR) \ +static void unary_task_##SUFFIX##_##NAME(unsigned int nth, unsigned int ith, void * data) { \ const struct htp_unary_context * uctx = (const struct htp_unary_context *) data; \ struct htp_ops_context * octx = uctx->octx; \ const struct htp_tensor * src = octx->src[0]; \ @@ -536,7 +676,7 @@ static void unary_task_f32_##NAME(unsigned int nth, unsigned int ith, void * dat const uint32_t dst_max_block = block_dst_contig ? uctx->block : MIN((uint32_t)uctx->block, ne1); \ const uint32_t BLOCK = MIN(src0_max_block, dst_max_block); \ if (BLOCK == 0) { \ - FARF(ERROR, "unary-f32 : current VTCM reservation %zu is too small, needed at least %zu\n", \ + FARF(ERROR, "unary-" #SUFFIX " : current VTCM reservation %zu is too small, needed at least %zu\n", \ uctx->vtcm_src0_size_per_thread, src0_row_size_aligned); \ return; \ } \ @@ -578,11 +718,11 @@ static void unary_task_f32_##NAME(unsigned int nth, unsigned int ith, void * dat const uint32_t block_size = unary_block_size(ir, src0_end_row, BLOCK, block_src0_contig, block_dst_contig, \ ne01, div_ne01); \ \ - float * dst_vtcm = (float *) dma_queue_pop(dma_queue).src; \ - float * src0_vtcm = (float *) dma_queue_pop(dma_queue).dst; \ - float * src1_vtcm = NULL; \ + TYPE * dst_vtcm = (TYPE *) dma_queue_pop(dma_queue).src; \ + TYPE * src0_vtcm = (TYPE *) dma_queue_pop(dma_queue).dst; \ + TYPE * src1_vtcm = NULL; \ if ((IS_RMS_NORM_MUL) && !uctx->broadcast_weight) { \ - src1_vtcm = (float *) dma_queue_pop(dma_queue).dst; \ + src1_vtcm = (TYPE *) dma_queue_pop(dma_queue).dst; \ } \ \ htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, ir); \ @@ -625,6 +765,10 @@ static void unary_task_f32_##NAME(unsigned int nth, unsigned int ith, void * dat dma_queue_flush(dma_queue); \ } +// F32 unary task: row-block DMA/VTCM plumbing, float-typed VTCM buffers. +#define DEFINE_UNARY_TASK(NAME, IS_RMS_NORM_MUL, IS_TRI, CORE_EXPR) \ + DEFINE_UNARY_TASK_IMPL(NAME, float, f32, IS_RMS_NORM_MUL, IS_TRI, CORE_EXPR) + DEFINE_UNARY_TASK(norm, false, false, norm_f32(src0_vtcm, dst_vtcm, block_size, uctx)) DEFINE_UNARY_TASK(rms_norm, false, false, rms_norm_f32(src0_vtcm, dst_vtcm, block_size, uctx)) DEFINE_UNARY_TASK(rms_norm_mul, true, false, rms_norm_mul_f32(src0_vtcm, uctx->broadcast_weight ? (const float *) src1_vtcm_data : src1_vtcm, dst_vtcm, block_size, uctx)) @@ -644,6 +788,18 @@ DEFINE_UNARY_TASK(unary_log, false, false, log_f32(src0_vtcm, dst_vtcm, blo DEFINE_UNARY_TASK(l2_norm, false, false, l2_norm_f32(src0_vtcm, dst_vtcm, block_size, uctx)) DEFINE_UNARY_TASK(tri, false, true, tri_f32(src0_vtcm, dst_vtcm, block_size, ir, uctx)) +// F16 unary tasks: same DMA/VTCM plumbing as DEFINE_UNARY_TASK, but VTCM buffers are +// _Float16-typed. None of the current F16 ops need RMS_NORM_MUL or TRI support. +DEFINE_UNARY_TASK_IMPL(norm, _Float16, f16, false, false, norm_f16(src0_vtcm, dst_vtcm, block_size, uctx)) +DEFINE_UNARY_TASK_IMPL(rms_norm, _Float16, f16, false, false, rms_norm_f16(src0_vtcm, dst_vtcm, block_size, uctx)) +DEFINE_UNARY_TASK_IMPL(scale, _Float16, f16, false, false, scale_f16(src0_vtcm, dst_vtcm, block_size, uctx)) +DEFINE_UNARY_TASK_IMPL(clamp, _Float16, f16, false, false, clamp_f16(src0_vtcm, dst_vtcm, block_size, uctx)) +DEFINE_UNARY_TASK_IMPL(sqr, _Float16, f16, false, false, sqr_f16(src0_vtcm, dst_vtcm, block_size, uctx)) +DEFINE_UNARY_TASK_IMPL(sqrt, _Float16, f16, false, false, sqrt_f16(src0_vtcm, dst_vtcm, block_size, uctx)) +DEFINE_UNARY_TASK_IMPL(l2_norm, _Float16, f16, false, false, l2_norm_f16(src0_vtcm, dst_vtcm, block_size, uctx)) +DEFINE_UNARY_TASK_IMPL(unary_abs, _Float16, f16, false, false, abs_f16(src0_vtcm, dst_vtcm, block_size, uctx)) +DEFINE_UNARY_TASK_IMPL(unary_log, _Float16, f16, false, false, log_f16(src0_vtcm, dst_vtcm, block_size, uctx)) + // Apply a pointwise unary op to one column tile that is already in VTCM. #define DEFINE_UNARY_TILED_TASK(NAME, IS_TRI, CORE_TILE_EXPR) \ static void unary_task_f32_tiled_##NAME(unsigned int nth, unsigned int ith, void * data) { \ @@ -892,50 +1048,76 @@ DEFINE_UNARY_TILED_TASK(unary_abs, false, hvx_abs_f32_aa(dst_vtcm, src_vtcm DEFINE_UNARY_TILED_TASK(unary_log, false, hvx_log_f32_aa(dst_vtcm, src_vtcm, tw)) DEFINE_UNARY_TILED_TASK(tri, true, tri_apply_tile_f32(src_vtcm, dst_vtcm, tw, col, i01, ne0, tri_ttype)) -static int execute_op_unary_f32(struct htp_ops_context * octx) { +static int execute_op_unary(struct htp_ops_context * octx) { int err = HTP_STATUS_OK; const struct htp_tensor * src0 = octx->src[0]; const struct htp_tensor * dst = octx->dst; + const bool is_f16 = (src0->type == HTP_TYPE_F16); + const char * op_type = NULL; switch (octx->op) { - case HTP_OP_NORM: op_type = "norm-f32"; break; - case HTP_OP_RMS_NORM: op_type = "rmsnorm-f32"; break; - case HTP_OP_RMS_NORM_MUL: op_type = "rmsnorm-mul-f32"; break; - case HTP_OP_SCALE: op_type = "scale-f32"; break; - case HTP_OP_CLAMP: op_type = "clamp-f32"; break; - case HTP_OP_SQR: op_type = "sqr-f32"; break; - case HTP_OP_SQRT: op_type = "sqrt-f32"; break; - case HTP_OP_UNARY_NEG: op_type = "neg-f32"; break; - case HTP_OP_UNARY_EXP: op_type = "exp-f32"; break; - case HTP_OP_UNARY_SIGMOID: op_type = "sigmoid-f32"; break; - case HTP_OP_UNARY_SILU: op_type = "silu-f32"; break; - case HTP_OP_UNARY_GELU: op_type = "gelu-f32"; break; - case HTP_OP_UNARY_SOFTPLUS: op_type = "softplus-f32"; break; - case HTP_OP_UNARY_TANH: op_type = "tanh-f32"; break; - case HTP_OP_UNARY_ABS: op_type = "abs-f32"; break; - case HTP_OP_UNARY_LOG: op_type = "log-f32"; break; - case HTP_OP_L2_NORM: op_type = "l2norm-f32"; break; - case HTP_OP_TRI: op_type = "tri-f32"; break; + case HTP_OP_NORM: op_type = is_f16 ? "norm-f16" : "norm-f32"; break; + case HTP_OP_RMS_NORM: op_type = is_f16 ? "rmsnorm-f16" : "rmsnorm-f32"; break; + case HTP_OP_RMS_NORM_MUL: op_type = "rmsnorm-mul-f32"; break; + case HTP_OP_SCALE: op_type = is_f16 ? "scale-f16" : "scale-f32"; break; + case HTP_OP_CLAMP: op_type = is_f16 ? "clamp-f16" : "clamp-f32"; break; + case HTP_OP_SQR: op_type = is_f16 ? "sqr-f16" : "sqr-f32"; break; + case HTP_OP_SQRT: op_type = is_f16 ? "sqrt-f16" : "sqrt-f32"; break; + case HTP_OP_UNARY_NEG: op_type = "neg-f32"; break; + case HTP_OP_UNARY_EXP: op_type = "exp-f32"; break; + case HTP_OP_UNARY_SIGMOID: op_type = "sigmoid-f32"; break; + case HTP_OP_UNARY_SILU: op_type = "silu-f32"; break; + case HTP_OP_UNARY_GELU: op_type = "gelu-f32"; break; + case HTP_OP_UNARY_SOFTPLUS: op_type = "softplus-f32"; break; + case HTP_OP_UNARY_TANH: op_type = "tanh-f32"; break; + case HTP_OP_UNARY_ABS: op_type = is_f16 ? "abs-f16" : "abs-f32"; break; + case HTP_OP_UNARY_LOG: op_type = is_f16 ? "log-f16" : "log-f32"; break; + case HTP_OP_L2_NORM: op_type = is_f16 ? "l2norm-f16" : "l2norm-f32"; break; + case HTP_OP_TRI: op_type = "tri-f32"; break; default: FARF(ERROR, "Unsupported unary Op %u\n", octx->op); return HTP_STATUS_NO_SUPPORT; } + // F16 only has row-block kernels for this subset of ops (see the dispatch switch + // below) - reject everything else up front, before touching kparams/VTCM. + if (is_f16) { + switch (octx->op) { + case HTP_OP_NORM: + case HTP_OP_RMS_NORM: + case HTP_OP_SCALE: + case HTP_OP_CLAMP: + case HTP_OP_SQR: + case HTP_OP_SQRT: + case HTP_OP_L2_NORM: + case HTP_OP_UNARY_ABS: + case HTP_OP_UNARY_LOG: + break; + default: + FARF(ERROR, "unary-%s: not supported for F16\n", op_type); + return HTP_STATUS_NO_SUPPORT; + } + } + const struct htp_unary_kernel_params * kparams = (const struct htp_unary_kernel_params *) octx->kernel_params; const uint32_t src0_nrows = src0->ne[1] * src0->ne[2] * src0->ne[3]; const uint32_t n_threads = kparams->n_threads; - const size_t src0_data_row_size = src0->ne[0] * sizeof(float); - const size_t dst_data_row_size = dst->ne[0] * sizeof(float); + const size_t elem_size = is_f16 ? sizeof(_Float16) : sizeof(float); + + const size_t src0_data_row_size = src0->ne[0] * elem_size; + const size_t dst_data_row_size = dst->ne[0] * elem_size; const size_t src0_row_size_aligned = kparams->src0_row_size_aligned; const size_t dst_row_size_aligned = kparams->dst_row_size_aligned; + // Always 0 for F16 - htp_unary_vtcm_layout_build() keeps F16 on the row-block path, + // since only F32 has unary_task_f32_tiled_* kernels. const uint32_t col_tile = kparams->col_tile; size_t src1_data_row_size = 0; @@ -943,6 +1125,8 @@ static int execute_op_unary_f32(struct htp_ops_context * octx) { bool broadcast_weight = kparams->broadcast_weight; const struct htp_tensor * src1 = NULL; + // RMS_NORM_MUL fusion is F32-only (its weight tensor is always F32; see + // try_fuse_node()'s type guard), so this never triggers when is_f16 is true. if (octx->op == HTP_OP_RMS_NORM_MUL) { src1 = octx->src[1]; src1_data_row_size = src1->ne[0] * sizeof(float); @@ -987,7 +1171,7 @@ static int execute_op_unary_f32(struct htp_ops_context * octx) { .block = kparams->block, .nc = src0->ne[0], - .col_tile = (uint32_t) kparams->col_tile, + .col_tile = col_tile, .broadcast_weight = broadcast_weight, .vtcm_src0 = VTCM_LAYOUT_PTR(uint8_t, base, 0), @@ -1020,6 +1204,19 @@ static int execute_op_unary_f32(struct htp_ops_context * octx) { case HTP_OP_TRI: task_func = unary_task_f32_tiled_tri; break; default: break; } + } else if (is_f16) { + switch (octx->op) { + case HTP_OP_NORM: task_func = unary_task_f16_norm; break; + case HTP_OP_RMS_NORM: task_func = unary_task_f16_rms_norm; break; + case HTP_OP_SCALE: task_func = unary_task_f16_scale; break; + case HTP_OP_CLAMP: task_func = unary_task_f16_clamp; break; + case HTP_OP_SQR: task_func = unary_task_f16_sqr; break; + case HTP_OP_SQRT: task_func = unary_task_f16_sqrt; break; + case HTP_OP_L2_NORM: task_func = unary_task_f16_l2_norm; break; + case HTP_OP_UNARY_ABS: task_func = unary_task_f16_unary_abs; break; + case HTP_OP_UNARY_LOG: task_func = unary_task_f16_unary_log; break; + default: break; + } } else { switch (octx->op) { case HTP_OP_NORM: task_func = unary_task_f32_norm; break; @@ -1047,7 +1244,7 @@ static int execute_op_unary_f32(struct htp_ops_context * octx) { if (task_func) { worker_pool_run_func(octx->ctx->worker_pool, task_func, &uctx, n_threads); } else { - FARF(ERROR, "execute_op_unary_f32: task function is NULL for op %d\n", octx->op); + FARF(ERROR, "execute_op_unary: task function is NULL for op %d\n", octx->op); err = HTP_STATUS_NO_SUPPORT; } } @@ -1058,7 +1255,8 @@ static int execute_op_unary_f32(struct htp_ops_context * octx) { int op_unary(struct htp_ops_context * octx) { switch (octx->src[0]->type) { case HTP_TYPE_F32: - return execute_op_unary_f32(octx); + case HTP_TYPE_F16: + return execute_op_unary(octx); default: return HTP_STATUS_NO_SUPPORT; diff --git a/ggml/src/ggml-hexagon/htp/unary-ops.h b/ggml/src/ggml-hexagon/htp/unary-ops.h index 458218ff443..116a591c2d7 100644 --- a/ggml/src/ggml-hexagon/htp/unary-ops.h +++ b/ggml/src/ggml-hexagon/htp/unary-ops.h @@ -85,17 +85,19 @@ static inline void htp_unary_vtcm_layout_build( bool broadcast_weight, uint32_t n_threads, size_t vtcm_size, + size_t elem_size, uint32_t * out_col_tile, uint32_t * out_vtcm_row_per_thread ) { - const size_t src0_data_row_size = ne00 * sizeof(float); - const size_t dst_data_row_size = ne10 * sizeof(float); + const size_t src0_data_row_size = ne00 * elem_size; + const size_t dst_data_row_size = ne10 * elem_size; const size_t src0_row_size_aligned = hex_round_up(src0_data_row_size, 128); const size_t dst_row_size_aligned = hex_round_up(dst_data_row_size, 128); size_t src1_row_size_aligned = 0; if (op == HTP_OP_RMS_NORM_MUL) { + // RMS_NORM_MUL fusion is F32-only; its weight tensor is always F32. const size_t src1_data_row_size = ne11 * sizeof(float); src1_row_size_aligned = hex_round_up(src1_data_row_size, 128); } @@ -125,12 +127,19 @@ static inline void htp_unary_vtcm_layout_build( const bool is_reduction = (op == HTP_OP_NORM || op == HTP_OP_RMS_NORM || op == HTP_OP_RMS_NORM_MUL || op == HTP_OP_L2_NORM); + // The tiled fallback path below only has F32 task functions (unary_task_f32_tiled_*); + // F16 has no tiled kernels, so it must stay on the row-block path like reduction ops. + // NOTE: if F16 ends up with vtcm_row_per_thread == 0 here (row too large for the VTCM + // budget), execute_op_unary() will see BLOCK == 0 and skip computation for that op + // (logged via FARF(ERROR, ...)) since there is no F16 tiled fallback. This is a known + // limitation; supporting it would require adding F16 tiled kernels. + const bool is_f16 = (elem_size == sizeof(_Float16)); uint32_t col_tile = 0; - if (vtcm_row_per_thread == 0 && !is_reduction) { + if (vtcm_row_per_thread == 0 && !is_reduction && !is_f16) { const size_t per_thread_budget = vtcm_size / n_threads; const size_t col_tile_bytes = hex_align_down(per_thread_budget / 4, 128); - col_tile = (uint32_t) (col_tile_bytes / sizeof(float)); + col_tile = (uint32_t) (col_tile_bytes / elem_size); L->src0_bytes = col_tile_bytes * 2; L->dst_bytes = col_tile_bytes * 2; From e5605697ce4dce019370afae61b7d222f2fb4bbf Mon Sep 17 00:00:00 2001 From: Xuan-Son Nguyen Date: Wed, 2 Sep 2026 23:53:32 +0200 Subject: [PATCH 086/104] finetune: fix no KV cache (llama/27199) * training: fix no KV cache * apply @ ggerganov suggestion --- ggml/src/ggml.c | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/ggml/src/ggml.c b/ggml/src/ggml.c index 8dc09450848..2d5fdb7c103 100644 --- a/ggml/src/ggml.c +++ b/ggml/src/ggml.c @@ -7335,7 +7335,7 @@ void ggml_build_backward_expand( } // inplace operations are currently not supported - GGML_ASSERT(!node->view_src || node->op == GGML_OP_CPY || node->op == GGML_OP_VIEW || + GGML_ASSERT(!node->view_src || node->op == GGML_OP_CPY || node->op == GGML_OP_SET_ROWS || node->op == GGML_OP_VIEW || node->op == GGML_OP_RESHAPE || node->op == GGML_OP_PERMUTE || node->op == GGML_OP_TRANSPOSE); const size_t ihash = ggml_hash_find(&cgraph->visited_hash_set, node); From a704770e3bdcce27c61b2a7695c6885ee79b6432 Mon Sep 17 00:00:00 2001 From: Eurekatic Date: Thu, 3 Sep 2026 08:59:06 +0200 Subject: [PATCH 087/104] sycl: reduce redundant work in Q4_K multi-column MMVQ (llama/27062) * sycl: Q4_K Weight unpack optimization and reuse between destination Columns * sycl: Q4_K small N (N=2..4) + two output rows by subgroup reuse of activation between two rows. * sycl: gate Q4_K two-row reuse for small N=2 * sycl: Fix on magic number now uses Q4_K_MMVQ_ROW_PAIR_MIN_NROWS=6272 for it, added tests for coverage around Q4_K_MMVQ_ROW_PAIR_MIN_NROWS with perf support to test Q4_K MUL_MAT, applied the same reuse pattern to the activation as the weights. Assisted-by: GPT-5.6 Sol --------- Co-authored-by: RaulAbejonDelgado --- ggml/src/ggml-sycl/mmvq.cpp | 195 ++++++++++++++++++++++++++------- ggml/src/ggml-sycl/vecdotq.hpp | 107 +++++++++++++----- 2 files changed, 237 insertions(+), 65 deletions(-) diff --git a/ggml/src/ggml-sycl/mmvq.cpp b/ggml/src/ggml-sycl/mmvq.cpp index 220663d5ac9..933bc77d2e4 100644 --- a/ggml/src/ggml-sycl/mmvq.cpp +++ b/ggml/src/ggml-sycl/mmvq.cpp @@ -6,6 +6,24 @@ #include "quants.hpp" #include "vecdotq.hpp" +// Minimum weight-row count at which the Q4_K multi-column MMVQ kernel handles two output rows per +// subgroup (rows_per_sg == 2) instead of one, when ncols_dst == 2. +// +// Pairing rows lets a subgroup load each activation block once and apply it to two rows, at the cost +// of halving the number of subgroups in the launch. With only two destination columns there is too +// little work per row to hide that loss of parallelism, so pairing only pays off once there are +// enough rows to keep the device occupied. This is a measured performance crossover, not a +// correctness or hardware limit - both variants compute the same result for any nrows. +// +// Derived on Intel Arc Pro B70 with `test-backend-ops perf -o MUL_MAT` (Q4_K, ncols_dst == 2), +// sweeping nrows over 5120..6912 at ncols 17408 and 19968: one row per subgroup was up to 9% faster +// below the crossover, two rows per subgroup 8-15% faster above it, and the crossover fell inside +// (6144, 6272] for both ncols with no measurable ncols dependence. A later 32-row granularity sweep +// narrowed it to (6144, 6176], so 6272 is a conservative gate rather than the exact crossover. +// ncols_dst >= 3 amortizes the activation loads over more columns and is faster with two rows at +// every row count, so it does not consult this threshold. +static constexpr int Q4_K_MMVQ_ROW_PAIR_MIN_NROWS = 6272; + template static void mul_mat_vec_q_reorder(const void * __restrict__ vx, const void * __restrict__ vy, float * __restrict__ dst, const int ncols, const int nrows, const sycl::nd_item<3> & nd_item) { @@ -59,7 +77,7 @@ static void mul_mat_vec_q_reorder(const void * __restrict__ vx, const void * __r // With has_fusion, `vgate` is a second weight matrix sharing vx's shape, stride and reorder // layout: one pass computes both row dot products and the epilogue writes glu(gate, up). -template +template static void mul_mat_vec_q_reorder_ncols(const void * __restrict__ vx, const void * __restrict__ vgate, const void * __restrict__ vy, float * __restrict__ dst, const int ncols, const int nrows, const int stride_col_y_bytes, const int stride_col_dst, @@ -71,14 +89,17 @@ static void mul_mat_vec_q_reorder_ncols(const void * __restrict__ vx, const void const int sg_range = sg.get_group_linear_range(); const int workgroup_id = nd_item.get_group_linear_id(); const int sg_id = sg.get_group_linear_id(); - const int row = workgroup_id * sg_range + sg_id; + const int row0 = (workgroup_id * sg_range + sg_id) * rows_per_sg; // row is sub-group uniform, so this retires whole sub-groups and the collectives below // stay convergent - if (row >= nrows) { + if (row0 >= nrows) { return; } + static_assert(rows_per_sg == 1 || + reorder_vec_dot_shared_activations::value); + const int blocks_per_row = ncols / block_traits::qk; constexpr int blocks_per_subgroup = ceil_div(block_traits::vdr_mmvq * WARP_SIZE, block_traits::qi); constexpr int block_elements_per_subgroup = block_traits::qi / block_traits::vdr_mmvq; @@ -87,34 +108,96 @@ static void mul_mat_vec_q_reorder_ncols(const void * __restrict__ vx, const void static_assert(blocks_per_subgroup > 0); static_assert(block_elements_per_subgroup > 0); - float partial_sum[ncols_dst] = { 0.0f }; + float partial_sum[ncols_dst][rows_per_sg] = {}; // sized 1 rather than 0 when unused: zero-length arrays are not standard C++, and the // array is dead and eliminated in that case - [[maybe_unused]] float partial_gate[has_fusion ? ncols_dst : 1] = { 0.0f }; + [[maybe_unused]] float partial_gate[has_fusion ? ncols_dst : 1][has_fusion ? rows_per_sg : 1] = {}; for (int i = sg.get_local_linear_id() / block_elements_per_subgroup; i < blocks_per_row; i += blocks_per_subgroup) { - const int ibx = row * blocks_per_row + i; - - // the offsets depend only on the block index and the matrix shape, never on the base - // pointer, which is what lets vgate reuse them - const auto bx_offset = block_type::get_block_offset(ibx, nblocks); - const auto d_offset = block_type::get_d_offset(nrows, ncols, ibx); const int iby = i * block_type::block_to_q8_1_ratio(); #pragma unroll for (int elem = 0; elem < block_elements_per_subgroup; elem += WARP_SIZE) { const int iqs = elem + block_traits::vdr_mmvq * (sg.get_local_linear_id() % block_elements_per_subgroup); + if constexpr (rows_per_sg > 1) { + typename reorder_vec_dot_q_sycl::weights wx[rows_per_sg]; + [[maybe_unused]] typename reorder_vec_dot_q_sycl::weights wg[rows_per_sg]; #pragma unroll - for (int j = 0; j < ncols_dst; ++j) { - const char * vy_j = (const char *) vy + j * stride_col_y_bytes; - const int8_t * q8_1_quant_ptr = (const int8_t *) vy_j + iby * QK8_1; - const sycl::half2 * q8_1_ds_ptr = (const sycl::half2 *) (vy_j + ncols + iby * sizeof(sycl::half2)); + for (int r = 0; r < rows_per_sg; ++r) { + const int row = sycl::min(row0 + r, nrows - 1); + const int ibx = row * blocks_per_row + i; + const auto bx_offset = block_type::get_block_offset(ibx, nblocks); + const auto d_offset = block_type::get_d_offset(nrows, ncols, ibx); + wx[r] = reorder_vec_dot_q_sycl::load(vx, bx_offset, d_offset, iqs); + if constexpr (has_fusion) { + wg[r] = reorder_vec_dot_q_sycl::load(vgate, bx_offset, d_offset, iqs); + } + } +#pragma unroll + for (int j = 0; j < ncols_dst; ++j) { + const char * vy_j = (const char *) vy + j * stride_col_y_bytes; + const int8_t * q8_1_quant_ptr = (const int8_t *) vy_j + iby * QK8_1; + const sycl::half2 * q8_1_ds_ptr = + (const sycl::half2 *) (vy_j + ncols + iby * sizeof(sycl::half2)); + const auto a = reorder_vec_dot_q_sycl::load_activations(q8_1_quant_ptr, q8_1_ds_ptr, iqs); +#pragma unroll + for (int r = 0; r < rows_per_sg; ++r) { + partial_sum[j][r] += reorder_vec_dot_q_sycl::apply(wx[r], a); + if constexpr (has_fusion) { + partial_gate[j][r] += reorder_vec_dot_q_sycl::apply(wg[r], a); + } + } + } + } else if constexpr (reorder_vec_dot_shared_weights::value) { + const int ibx = row0 * blocks_per_row + i; + const auto bx_offset = block_type::get_block_offset(ibx, nblocks); + const auto d_offset = block_type::get_d_offset(nrows, ncols, ibx); + const auto wx = reorder_vec_dot_q_sycl::load(vx, bx_offset, d_offset, iqs); + if constexpr (has_fusion) { + const auto wg = reorder_vec_dot_q_sycl::load(vgate, bx_offset, d_offset, iqs); - partial_sum[j] += reorder_vec_dot_q_sycl()(vx, bx_offset, d_offset, q8_1_quant_ptr, q8_1_ds_ptr, iqs); +#pragma unroll + for (int j = 0; j < ncols_dst; ++j) { + const char * vy_j = (const char *) vy + j * stride_col_y_bytes; + const int8_t * q8_1_quant_ptr = (const int8_t *) vy_j + iby * QK8_1; + const sycl::half2 * q8_1_ds_ptr = + (const sycl::half2 *) (vy_j + ncols + iby * sizeof(sycl::half2)); - if constexpr (has_fusion) { - partial_gate[j] += - reorder_vec_dot_q_sycl()(vgate, bx_offset, d_offset, q8_1_quant_ptr, q8_1_ds_ptr, iqs); + // up and gate share the activation, so load it once and apply it twice + const auto a = reorder_vec_dot_q_sycl::load_activations(q8_1_quant_ptr, q8_1_ds_ptr, iqs); + + partial_sum[j][0] += reorder_vec_dot_q_sycl::apply(wx, a); + partial_gate[j][0] += reorder_vec_dot_q_sycl::apply(wg, a); + } + } else { +#pragma unroll + for (int j = 0; j < ncols_dst; ++j) { + const char * vy_j = (const char *) vy + j * stride_col_y_bytes; + const int8_t * q8_1_quant_ptr = (const int8_t *) vy_j + iby * QK8_1; + const sycl::half2 * q8_1_ds_ptr = + (const sycl::half2 *) (vy_j + ncols + iby * sizeof(sycl::half2)); + + partial_sum[j][0] += reorder_vec_dot_q_sycl::dot(wx, q8_1_quant_ptr, q8_1_ds_ptr, iqs); + } + } + } else { + const int ibx = row0 * blocks_per_row + i; + const auto bx_offset = block_type::get_block_offset(ibx, nblocks); + const auto d_offset = block_type::get_d_offset(nrows, ncols, ibx); +#pragma unroll + for (int j = 0; j < ncols_dst; ++j) { + const char * vy_j = (const char *) vy + j * stride_col_y_bytes; + const int8_t * q8_1_quant_ptr = (const int8_t *) vy_j + iby * QK8_1; + const sycl::half2 * q8_1_ds_ptr = + (const sycl::half2 *) (vy_j + ncols + iby * sizeof(sycl::half2)); + + partial_sum[j][0] += + reorder_vec_dot_q_sycl()(vx, bx_offset, d_offset, q8_1_quant_ptr, q8_1_ds_ptr, iqs); + + if constexpr (has_fusion) { + partial_gate[j][0] += + reorder_vec_dot_q_sycl()(vgate, bx_offset, d_offset, q8_1_quant_ptr, q8_1_ds_ptr, iqs); + } } } } @@ -122,17 +205,20 @@ static void mul_mat_vec_q_reorder_ncols(const void * __restrict__ vx, const void #pragma unroll for (int j = 0; j < ncols_dst; ++j) { - float sum = sycl::reduce_over_group(nd_item.get_sub_group(), partial_sum[j], std::plus<>()); +#pragma unroll + for (int r = 0; r < rows_per_sg; ++r) { + float sum = sycl::reduce_over_group(nd_item.get_sub_group(), partial_sum[j][r], std::plus<>()); - if constexpr (has_fusion) { - const float gate = sycl::reduce_over_group(nd_item.get_sub_group(), partial_gate[j], std::plus<>()); + if constexpr (has_fusion) { + const float gate = sycl::reduce_over_group(nd_item.get_sub_group(), partial_gate[j][r], std::plus<>()); - // uniform across the launch; the launcher only instantiates SWIGLU and GEGLU - sum *= glu_op == GGML_GLU_OP_SWIGLU ? op_silu(gate) : op_gelu(gate); - } + // uniform across the launch; the launcher only instantiates SWIGLU and GEGLU + sum *= glu_op == GGML_GLU_OP_SWIGLU ? op_silu(gate) : op_gelu(gate); + } - if (sg.leader()) { - dst[j * stride_col_dst + row] = sum; + if (sg.leader() && row0 + r < nrows) { + dst[j * stride_col_dst + row0 + r] = sum; + } } } } @@ -1671,8 +1757,8 @@ static void reorder_mul_mat_vec_q4_k_q8_1_sycl(const void * vx, const void * vy, }); } -template -static void reorder_mul_mat_vec_q4_k_q8_1_sycl_ncols( +template +static void reorder_mul_mat_vec_q4_k_q8_1_sycl_ncols_impl( const void * vx, const void * vy, float * dst, const int ncols, const int nrows, const int stride_col_y_bytes, const int stride_col_dst, @@ -1680,20 +1766,31 @@ static void reorder_mul_mat_vec_q4_k_q8_1_sycl_ncols( GGML_ASSERT(ncols % QK_K == 0); constexpr size_t num_subgroups = WARP_SIZE; - const int block_num_y = ceil_div(nrows, GGML_SYCL_MMV_Y * (int) num_subgroups); + const int block_num_y = ceil_div(nrows, GGML_SYCL_MMV_Y * (int) num_subgroups * rows_per_sg); const sycl::range<3> block_nums(1, 1, block_num_y); const sycl::range<3> block_dims(1, GGML_SYCL_MMV_Y, num_subgroups * WARP_SIZE); stream->submit([&](sycl::handler & cgh) { cgh.parallel_for(sycl::nd_range<3>(block_nums * block_dims, block_dims), [=](sycl::nd_item<3> nd_item) [[sycl::reqd_sub_group_size(WARP_SIZE)]] { - mul_mat_vec_q_reorder_ncols, ncols_dst>( + mul_mat_vec_q_reorder_ncols, ncols_dst, + /*has_fusion=*/ false, rows_per_sg>( vx, /*vgate=*/ nullptr, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, /*glu_op=*/ GGML_GLU_OP_SWIGLU, nd_item); }); }); } +template +static void reorder_mul_mat_vec_q4_k_q8_1_sycl_ncols( + const void * vx, const void * vy, float * dst, + const int ncols, const int nrows, + const int stride_col_y_bytes, const int stride_col_dst, + dpct::queue_ptr stream) { + constexpr int rows_per_sg = ncols_dst >= 3 && ncols_dst <= 4 ? 2 : 1; + reorder_mul_mat_vec_q4_k_q8_1_sycl_ncols_impl(vx, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, stream); +} + static void reorder_mul_mat_vec_q4_k_q8_1_sycl_switch_ncols( const void * vx, const void * vy, float * dst, const int ncols, const int nrows, const int ncols_dst, @@ -1701,7 +1798,13 @@ static void reorder_mul_mat_vec_q4_k_q8_1_sycl_switch_ncols( dpct::queue_ptr stream) { switch (ncols_dst) { case 1: reorder_mul_mat_vec_q4_k_q8_1_sycl(vx, vy, dst, ncols, nrows, stream); break; - case 2: reorder_mul_mat_vec_q4_k_q8_1_sycl_ncols<2>(vx, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, stream); break; + case 2: + if (nrows >= Q4_K_MMVQ_ROW_PAIR_MIN_NROWS) { + reorder_mul_mat_vec_q4_k_q8_1_sycl_ncols_impl<2, 2>(vx, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, stream); + } else { + reorder_mul_mat_vec_q4_k_q8_1_sycl_ncols_impl<2, 1>(vx, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, stream); + } + break; case 3: reorder_mul_mat_vec_q4_k_q8_1_sycl_ncols<3>(vx, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, stream); break; case 4: reorder_mul_mat_vec_q4_k_q8_1_sycl_ncols<4>(vx, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, stream); break; case 5: reorder_mul_mat_vec_q4_k_q8_1_sycl_ncols<5>(vx, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, stream); break; @@ -2839,8 +2942,8 @@ bool ggml_sycl_mul_mat_vec_q_id_reorder( } } -template -static void launch_mul_mat_vec_q_reorder_glu(const void * vx, const void * vgate, const void * vy, float * dst, +template +static void launch_mul_mat_vec_q_reorder_glu_impl(const void * vx, const void * vgate, const void * vy, float * dst, const int ncols, const int nrows, const int stride_col_y_bytes, const int stride_col_dst, const ggml_glu_op glu_op, dpct::queue_ptr stream) { @@ -2848,20 +2951,33 @@ static void launch_mul_mat_vec_q_reorder_glu(const void * vx, const void * vgate constexpr size_t num_subgroups = WARP_SIZE; - const int block_num_y = ceil_div(nrows, GGML_SYCL_MMV_Y * (int) num_subgroups); + const int block_num_y = ceil_div(nrows, GGML_SYCL_MMV_Y * (int) num_subgroups * rows_per_sg); const sycl::range<3> block_nums(1, 1, block_num_y); const sycl::range<3> block_dims(1, GGML_SYCL_MMV_Y, num_subgroups * WARP_SIZE); stream->submit([&](sycl::handler & cgh) { cgh.parallel_for(sycl::nd_range<3>(block_nums * block_dims, block_dims), [=](sycl::nd_item<3> nd_item) [[sycl::reqd_sub_group_size(WARP_SIZE)]] { - mul_mat_vec_q_reorder_ncols( + mul_mat_vec_q_reorder_ncols( vx, vgate, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, glu_op, nd_item); }); }); } +template +static void launch_mul_mat_vec_q_reorder_glu(const void * vx, const void * vgate, const void * vy, float * dst, + const int ncols, const int nrows, const int stride_col_y_bytes, + const int stride_col_dst, const ggml_glu_op glu_op, + dpct::queue_ptr stream) { + constexpr int rows_per_sg = + reorder_vec_dot_shared_activations::value && ncols_dst >= 3 && ncols_dst <= 4 + ? 2 + : 1; + launch_mul_mat_vec_q_reorder_glu_impl(vx, vgate, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, glu_op, stream); +} + bool ggml_sycl_mul_mat_vec_q_glu_reorder(enum ggml_type src0_type, enum ggml_glu_op glu_op, const void * vx, const void * vgate, const void * vy, float * dst, int ncols, int nrows, int ncols_dst, int stride_col_y_bytes, int stride_col_dst, @@ -2881,8 +2997,11 @@ bool ggml_sycl_mul_mat_vec_q_glu_reorder(enum ggml_type src0_type, enum ggml_glu stride_col_dst, glu_op, stream); return true; case 2: - launch_mul_mat_vec_q_reorder_glu(vx, vgate, vy, dst, ncols, nrows, stride_col_y_bytes, - stride_col_dst, glu_op, stream); + if (nrows >= Q4_K_MMVQ_ROW_PAIR_MIN_NROWS) { + launch_mul_mat_vec_q_reorder_glu_impl(vx, vgate, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, glu_op, stream); + } else { + launch_mul_mat_vec_q_reorder_glu_impl(vx, vgate, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, glu_op, stream); + } return true; case 3: launch_mul_mat_vec_q_reorder_glu(vx, vgate, vy, dst, ncols, nrows, stride_col_y_bytes, diff --git a/ggml/src/ggml-sycl/vecdotq.hpp b/ggml/src/ggml-sycl/vecdotq.hpp index 3ad4cee93a1..ed5fd7de890 100644 --- a/ggml/src/ggml-sycl/vecdotq.hpp +++ b/ggml/src/ggml-sycl/vecdotq.hpp @@ -351,6 +351,25 @@ template struct reorder_vec_dot_q_sycl { static_assert(T != T, "ggml_type for reorder vecdot not implemented"); }; +// For some types the weight side of the dot product does not depend on the destination column, so a +// multi-column mul_mat_vec can unpack it once per block instead of once per column. Such a type adds +// load() and dot() next to operator() and opts in here. See reorder_vec_dot_q_sycl. +template struct reorder_vec_dot_shared_weights { + static constexpr bool value = false; +}; + +template <> struct reorder_vec_dot_shared_weights { + static constexpr bool value = true; +}; + +template struct reorder_vec_dot_shared_activations { + static constexpr bool value = false; +}; + +template <> struct reorder_vec_dot_shared_activations { + static constexpr bool value = true; +}; + template <> struct reorder_vec_dot_q_sycl { static constexpr ggml_type gtype = GGML_TYPE_Q4_0; @@ -540,50 +559,84 @@ template <> struct reorder_vec_dot_q_sycl { using q4_k_block = ggml_sycl_reordered::block_q_t; using q4_k_traits = typename q4_k_block::traits; - __dpct_inline__ float operator()(const void * __restrict__ vbq, const std::pair ibx_offset, - const std::pair d_offset, const int8_t * q8_1_quant_ptr, - const sycl::half2 * q8_1_ds, const int & iqs) { - const uint8_t * base = static_cast(vbq); - const uint8_t * qs = base + ibx_offset.first; - const uint8_t * scs = base + d_offset.first; - const ggml_half2 * dms = reinterpret_cast(base + d_offset.second); - - const int bq8_offset = QR4_K * ((iqs / 2) / (QI8_1 / 2)); - const int * q4 = (const int *) (qs + 16 * bq8_offset + 4 * ((iqs / 2) % 4)); - const uint16_t * scales = (const uint16_t *) scs; + struct weights { + int v[2]; + uint16_t aux[2]; + ggml_half2 dm; + int bq8_offset; + }; - int v[2]; + struct activations { int u[2 * QR4_K]; float d8[QR4_K]; + }; - v[0] = q4[0]; - v[1] = q4[4]; + __dpct_inline__ static weights load(const void * __restrict__ vbq, const std::pair ibx_offset, + const std::pair d_offset, const int & iqs) { + const uint8_t * base = static_cast(vbq); + const uint8_t * qs = base + ibx_offset.first; + const uint8_t * scs = base + d_offset.first; + const ggml_half2 * dms = reinterpret_cast(base + d_offset.second); + + weights w; + w.bq8_offset = QR4_K * ((iqs / 2) / (QI8_1 / 2)); + + const int * q4 = (const int *) (qs + 16 * w.bq8_offset + 4 * ((iqs / 2) % 4)); + const uint16_t * scales = (const uint16_t *) scs; + + w.v[0] = q4[0]; + w.v[1] = q4[4]; - uint16_t aux[2]; const int j = (QR4_K * ((iqs / 2) / (QI8_1 / 2))) / 2; if (j < 2) { - aux[0] = scales[j + 0] & 0x3f3f; - aux[1] = scales[j + 2] & 0x3f3f; + w.aux[0] = scales[j + 0] & 0x3f3f; + w.aux[1] = scales[j + 2] & 0x3f3f; } else { - aux[0] = ((scales[j + 2] >> 0) & 0x0f0f) | ((scales[j - 2] & 0xc0c0) >> 2); - aux[1] = ((scales[j + 2] >> 4) & 0x0f0f) | ((scales[j - 0] & 0xc0c0) >> 2); + w.aux[0] = ((scales[j + 2] >> 0) & 0x0f0f) | ((scales[j - 2] & 0xc0c0) >> 2); + w.aux[1] = ((scales[j + 2] >> 4) & 0x0f0f) | ((scales[j - 0] & 0xc0c0) >> 2); } - const uint8_t * sc = (const uint8_t *) aux; - const uint8_t * m = sc + 2; + w.dm = *dms; + return w; + } + + __dpct_inline__ static activations load_activations(const int8_t * q8_1_quant_ptr, + const sycl::half2 * q8_1_ds, const int & iqs) { + activations a; + const int bq8_offset = QR4_K * ((iqs / 2) / (QI8_1 / 2)); for (int i = 0; i < QR4_K; ++i) { - const int8_t* quant_base_ptr = q8_1_quant_ptr + (bq8_offset + i) * QK8_1; - sycl::half2 ds_values = *(q8_1_ds + bq8_offset + i); + const int8_t * quant_base_ptr = q8_1_quant_ptr + (bq8_offset + i) * QK8_1; + sycl::half2 ds_values = *(q8_1_ds + bq8_offset + i); - d8[i] = ds_values[0]; + a.d8[i] = ds_values[0]; const int * q8 = (const int *) quant_base_ptr + ((iqs / 2) % 4); - u[2 * i + 0] = q8[0]; - u[2 * i + 1] = q8[4]; + a.u[2 * i + 0] = q8[0]; + a.u[2 * i + 1] = q8[4]; } - return vec_dot_q4_K_q8_1_impl_vmmq(v, u, sc, m, *dms, d8); + return a; + } + + __dpct_inline__ static float apply(const weights & w, const activations & a) { + const uint8_t * sc = (const uint8_t *) w.aux; + const uint8_t * m = sc + 2; + + return vec_dot_q4_K_q8_1_impl_vmmq(w.v, a.u, sc, m, w.dm, a.d8); + } + + __dpct_inline__ static float dot(const weights & w, const int8_t * q8_1_quant_ptr, + const sycl::half2 * q8_1_ds, const int & iqs) { + const auto a = load_activations(q8_1_quant_ptr, q8_1_ds, iqs); + + return apply(w, a); + } + + __dpct_inline__ float operator()(const void * __restrict__ vbq, const std::pair ibx_offset, + const std::pair d_offset, const int8_t * q8_1_quant_ptr, + const sycl::half2 * q8_1_ds, const int & iqs) { + return dot(load(vbq, ibx_offset, d_offset, iqs), q8_1_quant_ptr, q8_1_ds, iqs); } }; From f24a38605b729b4ab6546ccbde2d9f92d773558f Mon Sep 17 00:00:00 2001 From: Neo Zhang Date: Thu, 3 Sep 2026 15:41:07 +0800 Subject: [PATCH 088/104] sycl : enhance the api to support peer-to-peer copy (llama/27550) --- ggml/src/ggml-sycl/ggml-sycl.cpp | 1 + 1 file changed, 1 insertion(+) diff --git a/ggml/src/ggml-sycl/ggml-sycl.cpp b/ggml/src/ggml-sycl/ggml-sycl.cpp index 626f0f3bff1..60a9c6015d3 100644 --- a/ggml/src/ggml-sycl/ggml-sycl.cpp +++ b/ggml/src/ggml-sycl/ggml-sycl.cpp @@ -742,6 +742,7 @@ static void dev2dev_memcpy(int device_dst, sycl::queue &q_dst, int device_src, s if (q_dst.get_device().ext_oneapi_can_access_peer(q_src.get_device(), sycl::ext::oneapi::peer_access::access_supported)) { GGML_SYCL_DEBUG("[SYCL] dev2dev memcpy by SYCL\n"); + q_dst.get_device().ext_oneapi_enable_peer_access(q_src.get_device()); SYCL_CHECK(CHECK_TRY_ERROR(q_dst.memcpy(ptr_dst, ptr_src, size).wait())); return; } From 47d348a2780fd70c8b50906a5c12cf0cc9a1b3fa Mon Sep 17 00:00:00 2001 From: Nathan Wilson <67372905+Nathanw1014@users.noreply.github.com> Date: Thu, 3 Sep 2026 15:40:34 +0700 Subject: [PATCH 089/104] vulkan: fix FA dequant path engagement (llama/28190) Skip the nb[3] check when ne[3] == 1, the shader never reads it for a single stream. Cache views carry the full-buffer stride there, so the old check reduced to n_kv == kv_size and the path only engaged with the cache full. --- ggml/src/ggml-vulkan/ggml-vulkan.cpp | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp index 285cd652537..a04a6b27a8c 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp +++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp @@ -10966,7 +10966,7 @@ static void ggml_vk_flash_attn(ggml_backend_vk_context * ctx, vk_context& subctx return t->nb[0] == ggml_type_size(t->type) && t->nb[2] == ggml_row_size(t->type, t->ne[0]) && t->nb[1] == t->nb[2] * t->ne[2] && - t->nb[3] == t->nb[1] * t->ne[1]; + (t->ne[3] == 1 || t->nb[3] == t->nb[1] * t->ne[1]); }; const bool k_quant = k->type != GGML_TYPE_F16 && k->type != GGML_TYPE_BF16 && k->type != GGML_TYPE_F32; const bool v_quant = v->type != GGML_TYPE_F16 && v->type != GGML_TYPE_BF16 && v->type != GGML_TYPE_F32; From 25350b579e7e957b4f005400fe72b40162082b18 Mon Sep 17 00:00:00 2001 From: Tanner Bruhn <66120666+tannerbruhn@users.noreply.github.com> Date: Thu, 3 Sep 2026 12:03:03 +0200 Subject: [PATCH 090/104] CUDA: Allow concurrent streams per split for multi-GPU (llama/28198) * CUDA: Allow CUDA optimization per split for multi-GPU. Previous guard caused multi-GPU to skip the graph optimization. The graph is already split per device and the optimization doesnt run over the whole model but once per split, and thus should be allowed. However, the CUDA event ggml_cuda_concurrent_event belongs to whichever GPU was "current" when created. If the pass ran while GPU 0 was current, it would stick and during event creation for the second GPU it would land on GPU 0. The fix: set the device explicitly ggml_cuda_set_device(cuda_ctx->device); Default behaviour remains unchanged, only active for GGML_CUDA_GRAPH_OPT=1. Explicit device setting pattern re-used from ggml_backend_cuda_graph_compute. * Update ggml/src/ggml-cuda/ggml-cuda.cu Co-authored-by: Aman Gupta --------- Co-authored-by: tannerbruhn Co-authored-by: Aman Gupta --- ggml/src/ggml-cuda/ggml-cuda.cu | 4 +++- 1 file changed, 3 insertions(+), 1 deletion(-) diff --git a/ggml/src/ggml-cuda/ggml-cuda.cu b/ggml/src/ggml-cuda/ggml-cuda.cu index f4af82688ac..45e9537f0e4 100644 --- a/ggml/src/ggml-cuda/ggml-cuda.cu +++ b/ggml/src/ggml-cuda/ggml-cuda.cu @@ -4542,10 +4542,12 @@ static void ggml_backend_cuda_graph_optimize(ggml_backend_t backend, ggml_cgraph ggml_cuda_stream_context & stream_context = cuda_ctx->stream_context(); stream_context.reset(); - if (!use_cuda_graph || ggml_backend_cuda_get_device_count() != 1) { + if (!use_cuda_graph) { return; } + ggml_cuda_set_device(cuda_ctx->device); + // number of out-degrees for a particular node std::unordered_map fan_out; // reverse mapping of node to index in the cgraph From d55d345e696debc6c8be4899b60c0238e621c792 Mon Sep 17 00:00:00 2001 From: Georgi Gerganov Date: Thu, 3 Sep 2026 13:25:41 +0300 Subject: [PATCH 091/104] metal : fix glu dispatch with ne00 = 1 (llama/28306) * metal : fix glu dispatch with ne00 = 1 * tests : disable ill-defined tests --- ggml/src/ggml-metal/ggml-metal-ops.cpp | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/ggml/src/ggml-metal/ggml-metal-ops.cpp b/ggml/src/ggml-metal/ggml-metal-ops.cpp index bc8b3c8d485..c5c4ec46bcf 100644 --- a/ggml/src/ggml-metal/ggml-metal-ops.cpp +++ b/ggml/src/ggml-metal/ggml-metal-ops.cpp @@ -917,7 +917,7 @@ int ggml_metal_op_glu(ggml_metal_op_t ctx, int idx) { const int64_t nrows = ggml_nrows(op->src[0]); - const int32_t nth = std::min(ggml_metal_pipeline_max_theads_per_threadgroup(pipeline), ne00/2); + const int32_t nth = std::max(1, std::min(ggml_metal_pipeline_max_theads_per_threadgroup(pipeline), ne00/2)); ggml_metal_encoder_set_pipeline(enc, pipeline); ggml_metal_encoder_set_bytes (enc, &args, sizeof(args), 0); From 4dd48dde357afef5c208adad869082edf94e70b1 Mon Sep 17 00:00:00 2001 From: Georgi Gerganov Date: Thu, 3 Sep 2026 13:51:13 +0300 Subject: [PATCH 092/104] metal : add sparse FA (llama/28098) * metal : support n_kv_max sparse mask hint in flash attention vec kernel - add kernel_flash_attn_ext_vec_idx: compacts finite mask entries into a per-row index list (Hillis-Steele scan, one threadgroup per row) - extend vec FA kernel with optional sparse index gathering (FC slot 5) - add host-side gate: sparse path when n_kv_max > 0, mask present, supported head sizes / KV types, n_kv_max <= 4096 - new buffer region extra_idx for the index list - pipeline getter extended with has_sparse param - add test cases: head sizes, quant types, nb>1, nr23 variants, sinks, ALiBi, softcap, permute, v_view_of_k, no-mask fallback Note: multi-row (nb*nr23[1] > 1) cases still failing - rid mapping in the store phase needs revisiting for the sparse path. Assisted-by: pi:llama.cpp/Qwen3.8-27B * metal : fix sparse flash attention row addressing - kernel_flash_attn_ext_vec_idx: mask param is half* but nb31 is a byte stride, so the per-row mask offset was scaled by 2x; cast to char* before applying the byte strides - kernel_flash_attn_ext_vec: sparse pidx param is char* so the per-row element offset was under-scaled by sizeof(int); scale it by sizeof(int) to get the correct byte offset - fixes the multi-row (nb*nr23[1] > 1) sparse flash attention failures Assisted-by: pi:llama.cpp/DeepSeek-v4-0731 * cont : use sparse vec FA for prefill * metal : single-pass flash attention sparse index compaction The idx kernel previously read the mask row twice: once to count the finite entries (for the prefix scan) and again to recover their positions. Since the kernel is memory-bound, this doubled the mask traffic. Keep the finite positions in a per-thread register array during the count pass and write them out directly, avoiding the second mask read. A dense mask with more than NLOCAL finite entries in a slice falls back to re-reading the mask to write the remaining positions. Assisted-by: pi:llama.cpp/DeepSeek-v4-0731 * tests : add perf cases for sparse flash attention prefill Measure the sparse vec FA kernel across KV sizes, n_kv_max hints and batch sizes. Run with: ./build/bin/test-backend-ops -b MTL0 -o FLASH_ATTN_EXT -p "n_kv_max=[1-9]" perf Assisted-by: pi:llama.cpp/DeepSeek-v4-0731 * qwen4 : enable sparse attention * cont : adjust nsg * cont : sync test-backend-ops * cont : disable Qwen4 for now * cont : clean-up + tests --- ggml/src/ggml-metal/ggml-metal-device.cpp | 27 ++- ggml/src/ggml-metal/ggml-metal-device.h | 5 + ggml/src/ggml-metal/ggml-metal-impl.h | 13 ++ ggml/src/ggml-metal/ggml-metal-ops.cpp | 176 ++++++++++++++++-- ggml/src/ggml-metal/ggml-metal-ops.h | 1 + ggml/src/ggml-metal/ggml-metal.cpp | 1 + ggml/src/ggml-metal/kernels/fa.metal | 215 ++++++++++++++++++++-- 7 files changed, 406 insertions(+), 32 deletions(-) diff --git a/ggml/src/ggml-metal/ggml-metal-device.cpp b/ggml/src/ggml-metal/ggml-metal-device.cpp index b8d2ef9ce27..c296d17b157 100644 --- a/ggml/src/ggml-metal/ggml-metal-device.cpp +++ b/ggml/src/ggml-metal/ggml-metal-device.cpp @@ -1577,6 +1577,26 @@ ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_flash_attn_ext( return res; } +ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_flash_attn_ext_vec_idx( + ggml_metal_library_t lib, + const ggml_tensor * op) { + assert(op->op == GGML_OP_FLASH_ATTN_EXT); + assert(op->src[3]); + + char name[256]; + + snprintf(name, 256, "kernel_flash_attn_ext_vec_idx"); + + ggml_metal_pipeline_with_params res = ggml_metal_library_get_pipeline(lib, name); + if (!res.pipeline) { + res = ggml_metal_library_compile_pipeline(lib, name, name, nullptr); + } + + GGML_UNUSED(op); + + return res; +} + ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_flash_attn_ext_vec( ggml_metal_library_t lib, const ggml_tensor * op, @@ -1585,6 +1605,7 @@ ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_flash_attn_ext_v bool has_bias, bool has_scap, bool has_kvpad, + bool has_sparse, int32_t nqpsg, int32_t ne, int32_t nsg, @@ -1614,13 +1635,14 @@ ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_flash_attn_ext_v dv, qne_suffix); - snprintf(name, 256, "%s_mask=%d_sink=%d_bias=%d_scap=%d_kvpad=%d_ns10=%d_ns20=%d_nsg=%d_nwg=%d", + snprintf(name, 256, "%s_mask=%d_sink=%d_bias=%d_scap=%d_kvpad=%d_sparse=%d_ns10=%d_ns20=%d_nsg=%d_nwg=%d", base, has_mask, has_sinks, has_bias, has_scap, has_kvpad, + has_sparse, ns10, ns20, nsg, nwg); @@ -1633,7 +1655,8 @@ ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_flash_attn_ext_v ggml_metal_cv_set_bool(cv, has_sinks, FC_FLASH_ATTN_EXT_VEC + 1); ggml_metal_cv_set_bool(cv, has_bias, FC_FLASH_ATTN_EXT_VEC + 2); ggml_metal_cv_set_bool(cv, has_scap, FC_FLASH_ATTN_EXT_VEC + 3); - ggml_metal_cv_set_bool(cv, has_kvpad, FC_FLASH_ATTN_EXT_VEC + 4); + ggml_metal_cv_set_bool(cv, has_kvpad, FC_FLASH_ATTN_EXT_VEC + 4); + ggml_metal_cv_set_bool(cv, has_sparse, FC_FLASH_ATTN_EXT_VEC + 5); ggml_metal_cv_set_int32(cv, ns10, FC_FLASH_ATTN_EXT_VEC + 20); ggml_metal_cv_set_int32(cv, ns20, FC_FLASH_ATTN_EXT_VEC + 21); diff --git a/ggml/src/ggml-metal/ggml-metal-device.h b/ggml/src/ggml-metal/ggml-metal-device.h index ae4871d3586..31fc07d44d4 100644 --- a/ggml/src/ggml-metal/ggml-metal-device.h +++ b/ggml/src/ggml-metal/ggml-metal-device.h @@ -201,6 +201,10 @@ struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_flash_att int32_t ns10, int32_t ns20); +struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_flash_attn_ext_vec_idx( + ggml_metal_library_t lib, + const struct ggml_tensor * op); + struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_flash_attn_ext_vec( ggml_metal_library_t lib, const struct ggml_tensor * op, @@ -209,6 +213,7 @@ struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_flash_att bool has_bias, bool has_scap, bool has_kvpad, + bool has_sparse, int32_t nqpsg, int32_t ne, int32_t nsg, diff --git a/ggml/src/ggml-metal/ggml-metal-impl.h b/ggml/src/ggml-metal/ggml-metal-impl.h index bdcd9c9e3d8..30e40f527f9 100644 --- a/ggml/src/ggml-metal/ggml-metal-impl.h +++ b/ggml/src/ggml-metal/ggml-metal-impl.h @@ -458,8 +458,21 @@ typedef struct { float m1; int32_t n_head_log2; float logit_softcap; + int32_t n_kv_max_padded; } ggml_metal_kargs_flash_attn_ext_vec; +typedef struct { + int32_t ne30; + int32_t ne31; + int32_t ne32; + int32_t ne33; + uint64_t nb31; + uint64_t nb32; + uint64_t nb33; + int32_t n_kv_max; + int32_t n_kv_max_padded; +} ggml_metal_kargs_flash_attn_ext_vec_idx; + typedef struct { int32_t nrows; } ggml_metal_kargs_flash_attn_ext_vec_reduce; diff --git a/ggml/src/ggml-metal/ggml-metal-ops.cpp b/ggml/src/ggml-metal/ggml-metal-ops.cpp index c5c4ec46bcf..3db8bca4375 100644 --- a/ggml/src/ggml-metal/ggml-metal-ops.cpp +++ b/ggml/src/ggml-metal/ggml-metal-ops.cpp @@ -2857,6 +2857,65 @@ static bool ggml_metal_op_flash_attn_ext_use_kv_f16(const ggml_tensor * op) { } } +// returns the n_kv_max hint if the sparse path is available for this op, or 0 otherwise +// the mask (src[3]) remains the single source of truth: finite entries are the valid KV positions, +// n_kv_max is only an upper bound on their number per mask row, used to size the index lists +static int ggml_metal_op_flash_attn_ext_n_kv_max_sparse(const ggml_tensor * op) { + assert(op->op == GGML_OP_FLASH_ATTN_EXT); + + int32_t n_kv_max = 0; + memcpy(&n_kv_max, ((const int32_t *) op->op_params) + 4, sizeof(n_kv_max)); + + if (n_kv_max <= 0) { + return 0; + } + + // the sparse indices are gathered from the mask + if (!op->src[3]) { + return 0; + } + + // bound the size of the index lists + if (n_kv_max > 4096) { + return 0; + } + + // vec kernel instantiations exist for these (type, dk, dv) combinations only + const int64_t dk = op->src[1]->ne[0]; + const int64_t dv = op->src[2]->ne[0]; + + const bool dk_dv_ok = (dk == 32 && dv == 32) || + (dk == 64 && dv == 64) || + (dk == 96 && dv == 96) || + (dk == 128 && dv == 128) || + (dk == 192 && dv == 128) || + (dk == 192 && dv == 192) || + (dk == 256 && dv == 256) || + (dk == 320 && dv == 256) || + (dk == 512 && dv == 512) || + (dk == 576 && dv == 512); + + if (!dk_dv_ok) { + return 0; + } + + switch (op->src[1]->type) { + case GGML_TYPE_F16: + case GGML_TYPE_BF16: + case GGML_TYPE_F32: + 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: + break; + default: + return 0; + } + + return n_kv_max; +} + // in some models (e.g. MLA-based), V is a view of K (the first ne20 elements of each K row); // the dequantized V is then a view of the dequantized K and does not need its own dequant or scratch // - ref: https://github.com/ggml-org/llama.cpp/pull/13435 @@ -3027,6 +3086,24 @@ size_t ggml_metal_op_flash_attn_ext_extra_kv_f16(const ggml_tensor * op) { return k_size + v_size; } +// size of the sparse index lists: one list of KV indices per mask row, +// padded with -1 up to a multiple of OP_FLASH_ATTN_EXT_VEC_NCPSG +size_t ggml_metal_op_flash_attn_ext_extra_idx(const ggml_tensor * op) { + assert(op->op == GGML_OP_FLASH_ATTN_EXT); + + GGML_TENSOR_LOCALS( int32_t, ne3, op->src[3], ne); + + const int n_kv_max = ggml_metal_op_flash_attn_ext_n_kv_max_sparse(op); + + if (n_kv_max <= 0) { + return 0; + } + + const int n_kv_max_padded = GGML_PAD(n_kv_max, OP_FLASH_ATTN_EXT_VEC_NCPSG); + + return GGML_PAD(sizeof(int32_t)*(size_t) n_kv_max_padded*ne31*ne32*ne33, 16); +} + int ggml_metal_op_flash_attn_ext(ggml_metal_op_t ctx, int idx) { ggml_tensor * op = ctx->node(idx); @@ -3104,7 +3181,16 @@ int ggml_metal_op_flash_attn_ext(ggml_metal_op_t ctx, int idx) { ggml_metal_buffer_id bid_kv_f16 = bid_tmp; bid_kv_f16.offs += ggml_metal_op_flash_attn_ext_extra_tmp(op); - const bool use_kv_f16 = ggml_metal_op_flash_attn_ext_use_kv_f16(op); + // sparse path: gather the finite mask entries into index lists and run the vec kernels over them + const int n_kv_max_sparse = ggml_metal_op_flash_attn_ext_n_kv_max_sparse(op); + const bool use_sparse = n_kv_max_sparse > 0; + const int n_kv_max_padded = use_sparse ? GGML_PAD(n_kv_max_sparse, OP_FLASH_ATTN_EXT_VEC_NCPSG) : 0; + + // the vec kernels dequantize the KV inline; no need for the F16 dequant pass in the sparse path + const bool use_kv_f16 = !use_sparse && ggml_metal_op_flash_attn_ext_use_kv_f16(op); + + ggml_metal_buffer_id bid_idx = bid_kv_f16; + bid_idx.offs += ggml_metal_op_flash_attn_ext_extra_kv_f16(op); ggml_metal_buffer_id bid_k = bid_src1; ggml_metal_buffer_id bid_v = bid_src2; @@ -3206,7 +3292,7 @@ int ggml_metal_op_flash_attn_ext(ggml_metal_op_t ctx, int idx) { } } - if (!ggml_metal_op_flash_attn_ext_use_vec(op)) { + if (!use_sparse && !ggml_metal_op_flash_attn_ext_use_vec(op)) { // half8x8 kernel const int nqptg = OP_FLASH_ATTN_EXT_NQPSG; // queries per threadgroup const int ncpsg = OP_FLASH_ATTN_EXT_NCPSG; // cache values per simdgroup @@ -3378,13 +3464,18 @@ int ggml_metal_op_flash_attn_ext(ggml_metal_op_t ctx, int idx) { #undef FATTN_SMEM } else { // half4x4 kernel - auto cfg = ggml_metal_tuning::fa_vec_pick( - props_dev->device_id, - props_dev->gpu_family, - (int) op->src[1]->type, - (int) ne00, (int) ne20, // dk, dv (ne00 == dk for FA) - ne11, ne01); - int nqptg = cfg.Q; // queries per threadgroup + // sparse: the index lists are per query row, so a threadgroup can share KV with Q == 1 only + auto cfg = use_sparse + ? ggml_metal_tuning::fa_vec_baseline_cfg((int) ne00, (int) ne20) + : ggml_metal_tuning::fa_vec_pick( + props_dev->device_id, + props_dev->gpu_family, + (int) op->src[1]->type, + (int) ne00, (int) ne20, // dk, dv (ne00 == dk for FA) + ne11, ne01); + + int nqptg = cfg.Q; // queries per threadgroup + const int ncpsg = OP_FLASH_ATTN_EXT_VEC_NCPSG; // cache values per simdgroup !! sync with kernel template arguments !! const int nhptg = 1; // heads per threadgroup @@ -3394,7 +3485,39 @@ int ggml_metal_op_flash_attn_ext(ggml_metal_op_t ctx, int idx) { bool need_sync = false; - const bool has_kvpad = ne11 % ncpsg != 0; + const bool has_kvpad = !use_sparse && ne11 % ncpsg != 0; + + if (use_sparse) { + assert(ggml_metal_op_flash_attn_ext_extra_idx(op) != 0); + + GGML_ASSERT(ne30 == ne11); + + ggml_metal_kargs_flash_attn_ext_vec_idx args0 = { + /*.ne30 =*/ ne30, + /*.ne31 =*/ ne31, + /*.ne32 =*/ ne32, + /*.ne33 =*/ ne33, + /*.nb31 =*/ nb31, + /*.nb32 =*/ nb32, + /*.nb33 =*/ nb33, + /*.n_kv_max =*/ n_kv_max_sparse, + /*.n_kv_max_padded =*/ n_kv_max_padded, + }; + + auto pipeline0 = ggml_metal_library_get_pipeline_flash_attn_ext_vec_idx(lib, op); + + ggml_metal_encoder_set_pipeline(enc, pipeline0); + ggml_metal_encoder_set_bytes (enc, &args0, sizeof(args0), 0); + ggml_metal_encoder_set_buffer (enc, bid_src3, 1); + ggml_metal_encoder_set_buffer (enc, bid_idx, 2); + + int nth = std::min(ggml_metal_pipeline_max_theads_per_threadgroup(pipeline0), 256); + nth = std::max(32, (nth/32)*32); + + ggml_metal_encoder_dispatch_threadgroups(enc, ne31, ne32, ne33, nth, 1, 1); + + need_sync = true; + } if (has_kvpad) { assert(ggml_metal_op_flash_attn_ext_extra_pad(op) != 0); @@ -3455,11 +3578,26 @@ int ggml_metal_op_flash_attn_ext(ggml_metal_op_t ctx, int idx) { // workgroups // each workgroup handles nsg*nkpsg cache values int32_t nwg = 1; - if (false) { - // for small KV caches, we could launch a single workgroup and write the results directly to dst/ - // however, this does not lead to significant improvement, so disabled - nwg = 1; - nsg = 4; + if (use_sparse) { + if (ne01 > 32) { + // large sparse batch + nwg = 1; + nsg = 1; + if (n_kv_max_padded == 640) { + nsg = 4; // 640 % (4*32) == 0 + } else { + while (2*nwg*nsg*ncpsg < n_kv_max_padded && nsg < 4) { + nsg *= 2; + } + } + } else { + // small sparse batch + nwg = 32; + nsg = 1; + while (2*nwg*nsg*ncpsg < n_kv_max_padded && nsg < 4) { + nsg *= 2; + } + } } else { nwg = 32; nsg = 1; @@ -3484,7 +3622,7 @@ int ggml_metal_op_flash_attn_ext(ggml_metal_op_t ctx, int idx) { /*.nb01 =*/ nb01, /*.nb02 =*/ nb02, /*.nb03 =*/ nb03, - /*.ne11 =*/ ne11, + /*.ne11 =*/ use_sparse ? n_kv_max_padded : ne11, /*.ne_12_2 =*/ ne12, /*.ne_12_3 =*/ ne13, /*.ns10 =*/ ns10, @@ -3510,9 +3648,10 @@ int ggml_metal_op_flash_attn_ext(ggml_metal_op_t ctx, int idx) { /*.m1 =*/ m1, /*.n_head_log2 =*/ n_head_log2, /*.logit_softcap =*/ logit_softcap, + /*.n_kv_max_padded =*/ n_kv_max_padded, }; - auto pipeline = ggml_metal_library_get_pipeline_flash_attn_ext_vec(lib, op, has_mask, has_sinks, has_bias, has_scap, has_kvpad, nqptg, cfg.NE, nsg, nwg, use_kv_f16, ns10, ns20); + auto pipeline = ggml_metal_library_get_pipeline_flash_attn_ext_vec(lib, op, has_mask, has_sinks, has_bias, has_scap, has_kvpad, use_sparse, nqptg, cfg.NE, nsg, nwg, use_kv_f16, ns10, ns20); GGML_ASSERT(nsg*32 <= ggml_metal_pipeline_max_theads_per_threadgroup(pipeline)); @@ -3523,6 +3662,7 @@ int ggml_metal_op_flash_attn_ext(ggml_metal_op_t ctx, int idx) { ggml_metal_encoder_set_buffer (enc, bid_v, 3); ggml_metal_encoder_set_buffer (enc, bid_src3, 4); ggml_metal_encoder_set_buffer (enc, bid_src4, 5); + ggml_metal_encoder_set_buffer (enc, use_sparse ? bid_idx : bid_src0, 8); const size_t smem = FATTN_SMEM(nsg); @@ -3530,8 +3670,6 @@ int ggml_metal_op_flash_attn_ext(ggml_metal_op_t ctx, int idx) { GGML_ASSERT(smem <= props_dev->max_theadgroup_memory_size); if (nwg == 1) { - assert(ggml_metal_op_flash_attn_ext_extra_tmp(op) == 0); - // using 1 workgroup -> write the result directly into dst ggml_metal_encoder_set_buffer(enc, bid_pad, 6); ggml_metal_encoder_set_buffer(enc, bid_dst, 7); diff --git a/ggml/src/ggml-metal/ggml-metal-ops.h b/ggml/src/ggml-metal/ggml-metal-ops.h index 159a628d04a..f8fe50b468e 100644 --- a/ggml/src/ggml-metal/ggml-metal-ops.h +++ b/ggml/src/ggml-metal/ggml-metal-ops.h @@ -43,6 +43,7 @@ size_t ggml_metal_op_flash_attn_ext_extra_pad(const struct ggml_tensor * op); size_t ggml_metal_op_flash_attn_ext_extra_blk(const struct ggml_tensor * op); size_t ggml_metal_op_flash_attn_ext_extra_tmp(const struct ggml_tensor * op); size_t ggml_metal_op_flash_attn_ext_extra_kv_f16(const struct ggml_tensor * op); +size_t ggml_metal_op_flash_attn_ext_extra_idx(const struct ggml_tensor * op); int ggml_metal_op_concat (ggml_metal_op_t ctx, int idx); int ggml_metal_op_repeat (ggml_metal_op_t ctx, int idx); diff --git a/ggml/src/ggml-metal/ggml-metal.cpp b/ggml/src/ggml-metal/ggml-metal.cpp index 4d58dc821cf..3bd6abd06fd 100644 --- a/ggml/src/ggml-metal/ggml-metal.cpp +++ b/ggml/src/ggml-metal/ggml-metal.cpp @@ -232,6 +232,7 @@ static size_t ggml_backend_metal_buffer_type_get_alloc_size(ggml_backend_buffer_ res += ggml_metal_op_flash_attn_ext_extra_blk(tensor); res += ggml_metal_op_flash_attn_ext_extra_tmp(tensor); res += ggml_metal_op_flash_attn_ext_extra_kv_f16(tensor); + res += ggml_metal_op_flash_attn_ext_extra_idx(tensor); } break; case GGML_OP_CUMSUM: case GGML_OP_ARGSORT: diff --git a/ggml/src/ggml-metal/kernels/fa.metal b/ggml/src/ggml-metal/kernels/fa.metal index e95dec258a3..d0e928d732c 100644 --- a/ggml/src/ggml-metal/kernels/fa.metal +++ b/ggml/src/ggml-metal/kernels/fa.metal @@ -1071,6 +1071,112 @@ constant int32_t FC_flash_attn_ext_vec_ns10 [[function_constant(FC_FLASH_ATTN_EX constant int32_t FC_flash_attn_ext_vec_ns20 [[function_constant(FC_FLASH_ATTN_EXT_VEC + 21)]]; constant int32_t FC_flash_attn_ext_vec_nsg [[function_constant(FC_FLASH_ATTN_EXT_VEC + 22)]]; constant int32_t FC_flash_attn_ext_vec_nwg [[function_constant(FC_FLASH_ATTN_EXT_VEC + 23)]]; +constant bool FC_flash_attn_ext_vec_has_sparse [[function_constant(FC_FLASH_ATTN_EXT_VEC + 5)]]; + +// compress the finite entries of each KQ mask row into a list of KV indices (ascending order), +// padded with -1 up to n_kv_max_padded (a multiple of OP_FLASH_ATTN_EXT_VEC_NCPSG) +// one threadgroup per mask row; the mask remains the single source of truth for the values +kernel void kernel_flash_attn_ext_vec_idx( + constant ggml_metal_kargs_flash_attn_ext_vec_idx & args, + device const half * mask, + device int * idx, + uint3 tgpig[[threadgroup_position_in_grid]], + ushort tiitg[[thread_index_in_threadgroup]], + ushort3 ntg[[threads_per_threadgroup]]) { + constexpr short NW = N_SIMDWIDTH; + constexpr short NLOCAL = 32; // max finite positions kept in registers per thread + + const int i1 = tgpig[0]; + const int i2 = tgpig[1]; + const int i3 = tgpig[2]; + + device const half * pm = (device const half *) ((device const char *) mask + i1*args.nb31 + i2*args.nb32 + i3*args.nb33); + device int * pidx = idx + (((int64_t)i3*args.ne32 + i2)*args.ne31 + i1)*args.n_kv_max_padded; + + const int n = args.ne30; + const int q = n/ntg.x; + const int r = n%ntg.x; + + // each thread handles a contiguous slice of the mask row + const int r0 = q*tiitg + min((int) tiitg, r); + const int r1 = r0 + q + (tiitg < r ? 1 : 0); + + // count the finite entries in the slice and keep their positions in registers (single mask read) + int cnt = 0; // total finite entries in the slice + int nloc = 0; // finite entries kept in registers + int local[NLOCAL]; + for (int i = r0; i < r1; ++i) { + if (isfinite((float) pm[i])) { + if (nloc < NLOCAL) { + local[nloc] = i; + nloc++; + } + cnt++; + } + } + + const short sgitg = tiitg/NW; + const short tiisg = tiitg%NW; + + threadgroup int tcount[8]; + + // simd_sum is a collective: all lanes must evaluate it + const int sg_sum = simd_sum(cnt); + if (tiisg == 0) { + tcount[sgitg] = sg_sum; + } + + threadgroup_barrier(mem_flags::mem_threadgroup); + + int total = 0; + for (short s = 0; s < ntg.x/NW; ++s) { + total += tcount[s]; + } + + // base offset of this thread's slice in the output list (exclusive scan within the simdgroup) + int sg_base = 0; + for (short s = 0; s < sgitg; ++s) { + sg_base += tcount[s]; + } + + // exclusive prefix scan of the per-thread counts within the simdgroup + int incl = cnt; + for (int d = 1; d < NW; d <<= 1) { + const int v = simd_shuffle_up(incl, d); + if (tiisg >= d) { + incl += v; + } + } + const int base = sg_base + (incl - cnt); + + // write the finite positions in order; if the hint is violated, keep only the first n_kv_max entries + int j = 0; + for (; j < nloc && base + j < args.n_kv_max; ++j) { + pidx[base + j] = local[j]; + } + + // a dense mask may have more than NLOCAL finite entries in a slice; re-read the mask to write the rest + if (cnt > nloc && base + nloc < args.n_kv_max) { + int j2 = 0; + for (int i = r0; i < r1; ++i) { + if (isfinite((float) pm[i])) { + if (j2 >= nloc) { + pidx[base + j2] = i; + } + j2++; + if (base + j2 >= args.n_kv_max) { + break; + } + } + } + } + + // pad the tail of the list with -1 + const int count = min(total, args.n_kv_max); + for (int i = count + tiitg; i < args.n_kv_max_padded; i += ntg.x) { + pidx[i] = -1; + } +} template< typename q4_t, // query types in shared memory @@ -1091,6 +1197,7 @@ template< short NE = 4, // head elements per thread short Q = OP_FLASH_ATTN_EXT_VEC_NQPSG, // queries per threadgroup short C = OP_FLASH_ATTN_EXT_VEC_NCPSG> // cache items per threadgroup + kernel void kernel_flash_attn_ext_vec( constant ggml_metal_kargs_flash_attn_ext_vec & args, device const char * q, @@ -1100,6 +1207,7 @@ kernel void kernel_flash_attn_ext_vec( device const char * sinks, device const char * pad, device char * dst, + device const char * idx, threadgroup half * shmem_f16 [[threadgroup(0)]], uint3 tgpig[[threadgroup_position_in_grid]], ushort tiisg[[thread_index_in_simdgroup]], @@ -1137,8 +1245,8 @@ kernel void kernel_flash_attn_ext_vec( //const short T = PK + NSG*SH; // shared memory size per query in (half) - //threadgroup q_t * sq = (threadgroup q_t *) (shmem_f16 + 0*PK); // holds the query data - threadgroup q4_t * sq4 = (threadgroup q4_t *) (shmem_f16 + 0*PK); // same as above but in q4_t + //threadgroup q_t * sq = (threadgroup q_t *) (shmem_f16 + 0*PK); // holds the query data + threadgroup q4_t * sq4 = (threadgroup q4_t *) (shmem_f16 + 0*PK); // same as above but in q4_t threadgroup s_t * ss = (threadgroup s_t *) (shmem_f16 + sgitg*SH + Q*NSG*PK); // scratch buffer for attention threadgroup s4_t * ss4 = (threadgroup s4_t *) (shmem_f16 + sgitg*SH + Q*NSG*PK); // same as above but in s4_t threadgroup half * sm = (threadgroup half *) (shmem_f16 + sgitg*SH + 2*Q*C + Q*NSG*PK); // scratch buffer for mask @@ -1207,6 +1315,14 @@ kernel void kernel_flash_attn_ext_vec( // pointer to the mask device const half * pm_base = (device const half *) (mask + iq1*Q*args.nb31 + (iq2%args.ne32)*args.nb32 + (iq3%args.ne33)*args.nb33); + // sparse indices: the list of finite mask entries per query row + // the sparse path requires Q == 1 (enforced by the host) + device const int * pidx = nullptr; + if (FC_flash_attn_ext_vec_has_sparse) { + pidx = (device const int *) idx + + ((int64_t)(iq3%args.ne33)*args.ne32 + (iq2%args.ne32))*args.ne31*args.n_kv_max_padded + (iq1%args.ne31)*args.n_kv_max_padded; + } + float slope = 1.0f; // ALiBi @@ -1265,11 +1381,22 @@ kernel void kernel_flash_attn_ext_vec( } if (FC_flash_attn_ext_vec_has_mask) { - FOR_UNROLL (short qq = 0; qq < Q; ++qq) { - if ((iq1*Q + qq) < args.ne01) { - sm[qq*C + tiisg] = pm[qq][ic + tiisg]; - } else { - sm[qq*C + tiisg] = -MAXHALF; + if (FC_flash_attn_ext_vec_has_sparse) { + FOR_UNROLL (short qq = 0; qq < Q; ++qq) { + const int i11 = pidx[ic + tiisg]; + if ((iq1*Q + qq) < args.ne01 && i11 >= 0) { + sm[qq*C + tiisg] = pm[qq][i11]; + } else { + sm[qq*C + tiisg] = -MAXHALF; + } + } + } else { + FOR_UNROLL (short qq = 0; qq < Q; ++qq) { + if ((iq1*Q + qq) < args.ne01) { + sm[qq*C + tiisg] = pm[qq][ic + tiisg]; + } else { + sm[qq*C + tiisg] = -MAXHALF; + } } } } else { @@ -1280,6 +1407,7 @@ kernel void kernel_flash_attn_ext_vec( } } + // skip -INF mask { bool any_finite = false; FOR_UNROLL (short qq = 0; qq < Q; ++qq) { @@ -1294,9 +1422,13 @@ kernel void kernel_flash_attn_ext_vec( // Q*K^T { - device const k4_t * pk4 = (device const k4_t *) (k + ic*args.nb11); + device const k4_t * pk4 = nullptr; + + if (!FC_flash_attn_ext_vec_has_sparse) { + pk4 = (device const k4_t *) (k + ic*args.nb11); - pk4 += ty*NS10/4 + tx; + pk4 += ty*NS10/4 + tx; + } qk_t mqk[Q][C/NE]; FOR_UNROLL (short qq = 0; qq < Q; ++qq) { @@ -1307,7 +1439,35 @@ kernel void kernel_flash_attn_ext_vec( // each simdgroup processes Q queries and NE (NW/NL) cache elements FOR_UNROLL (short cc = 0; cc < C/NE; ++cc) { - if (is_same::value) { + if (FC_flash_attn_ext_vec_has_sparse) { + // the KV rows are gathered from the index list; -1 entries are padding + const int i11 = pidx[ic + NE*cc + ty]; + if (i11 >= 0) { + if (is_same::value) { + device const k4_t * pk4s = (device const k4_t *) (k + i11*args.nb11) + tx; + FOR_UNROLL (short ii = 0; ii < DK4/NL; ++ii) { + const k4_t k_elem = pk4s[ii*NL]; + FOR_UNROLL (short qq = 0; qq < Q; ++qq) { + mqk[qq][cc] += dot((float4) k_elem, (float4) sq4[qq*PK4 + ii*NL + tx]); + } + } + } else { + device const kd4_t * pk = (device const kd4_t *) (k + i11*args.nb11); + + k4_t mk; + + FOR_UNROLL (short ii = 0; ii < DK4/NL; ++ii) { + const short i = ii*NL + tx; + + deq_k_t4(pk + i/nl_k, i%nl_k, mk); + + FOR_UNROLL (short qq = 0; qq < Q; ++qq) { + mqk[qq][cc] += dot((float4) mk, (float4) sq4[qq*PK4 + i]); + } + } + } + } + } else if (is_same::value) { FOR_UNROLL (short ii = 0; ii < DK4/NL; ++ii) { const k4_t k_elem = pk4[cc*NE*NS10/4 + ii*NL]; FOR_UNROLL (short qq = 0; qq < Q; ++qq) { @@ -1422,7 +1582,40 @@ kernel void kernel_flash_attn_ext_vec( } } - if (is_same::value) { + if (FC_flash_attn_ext_vec_has_sparse) { + FOR_UNROLL (short cc = 0; cc < C/NE; ++cc) { + // the KV rows are gathered from the index list; -1 entries are padding + const int i11 = pidx[ic + NE*cc + ty]; + if (i11 >= 0) { + if (is_same::value) { + device const v4_t * pv4 = (device const v4_t *) (v + i11*args.nb21); + + pv4 += tx; + + FOR_UNROLL (short ii = 0; ii < DV4/NL; ++ii) { + const v4_t v_elem = pv4[ii*NL]; + FOR_UNROLL (short qq = 0; qq < Q; ++qq) { + lo[qq][ii] += o4_t(float4(v_elem)*float4(ss[qq*C + cc*NE + ty])); + } + } + } else { + device const vd4_t * pv4 = (device const vd4_t *) (v + i11*args.nb21); + + FOR_UNROLL (short ii = 0; ii < DV4/NL; ++ii) { + const short i = ii*NL + tx; + + v4_t mv; + + deq_v_t4(pv4 + i/nl_v, i%nl_v, mv); + + FOR_UNROLL (short qq = 0; qq < Q; ++qq) { + lo[qq][ii] += o4_t(float4(mv)*float4(ss[qq*C + cc*NE + ty])); + } + } + } + } + } + } else if (is_same::value) { device const v4_t * pv4 = (device const v4_t *) (v + ic*args.nb21); pv4 += ty*NS20/4 + tx; From 0a4a95c86e592a8d2136e3316d2d6496c5033d14 Mon Sep 17 00:00:00 2001 From: kbenkhaled Date: Thu, 3 Sep 2026 12:40:42 -0400 Subject: [PATCH 093/104] tune MMVQ to MMQ crossover for SM87 (llama/28285) --- ggml/src/ggml-cuda/common.cuh | 1 + ggml/src/ggml-cuda/mmvq.cu | 12 ++++++++++++ 2 files changed, 13 insertions(+) diff --git a/ggml/src/ggml-cuda/common.cuh b/ggml/src/ggml-cuda/common.cuh index e5ccd1feab1..9918c03947c 100644 --- a/ggml/src/ggml-cuda/common.cuh +++ b/ggml/src/ggml-cuda/common.cuh @@ -52,6 +52,7 @@ #define GGML_CUDA_CC_VOLTA 700 #define GGML_CUDA_CC_TURING 750 #define GGML_CUDA_CC_AMPERE 800 +#define GGML_CUDA_CC_ORIN 870 #define GGML_CUDA_CC_ADA_LOVELACE 890 #define GGML_CUDA_CC_HOPPER 900 // While BW spans CC 1000, 1100 & 1200, we are integrating Tensor Core instructions available to 1200 family, see diff --git a/ggml/src/ggml-cuda/mmvq.cu b/ggml/src/ggml-cuda/mmvq.cu index 2be2f249108..f65e0fbcd7b 100644 --- a/ggml/src/ggml-cuda/mmvq.cu +++ b/ggml/src/ggml-cuda/mmvq.cu @@ -326,6 +326,18 @@ bool ggml_cuda_should_use_mmvq(enum ggml_type type, int cc, int64_t ne11) { return ne11 <= MMVQ_MAX_BATCH_SIZE; } } + if (GGML_CUDA_CC_IS_NVIDIA(cc) && cc == GGML_CUDA_CC_ORIN) { + switch (type) { // tuned for Jetson Orin + case GGML_TYPE_Q2_K: + case GGML_TYPE_Q3_K: + case GGML_TYPE_Q4_K: + case GGML_TYPE_Q5_K: + case GGML_TYPE_Q6_K: + return ne11 <= 1; + default: + return ne11 <= MMVQ_MAX_BATCH_SIZE; + } + } if (GGML_CUDA_CC_IS_CDNA(cc)) { if (GGML_CUDA_CC_IS_CDNA1(cc)) { switch (type) { From d784add75f41c7487da2b2b5c9ebc224f3b4e289 Mon Sep 17 00:00:00 2001 From: Hongqiang Wang Date: Thu, 3 Sep 2026 09:46:19 -0700 Subject: [PATCH 094/104] opencl: quant lm_head / decode GEMV and medium-batch GEMM optimizations (speculative decoding/MTP) (llama/26477) * opencl: quant lm_head / decode GEMV and medium-batch GEMM optimizations * opencl: guard q4_K/q6_K tiled_ns convert-kernel registration for non-Adreno build * opencl: gate q4_K MUL_MAT+GLU fusion dispatch to Adreno * opencl: require the noshuffle weight layout in the q4_K GLU fusion gate * opencl: do not take the vectorized f16 mrow GEMV path on an unaligned row stride * opencl: pass the new get_scale_min_k4 stride argument at the row-major call sites * opencl: enable the q4_K split-K decode GEMV only where it is measured to win * opencl: record the X1-85 split-K datapoint (neutral, exclusion confirmed) * opencl: restrict the tiled lm_head/embed GEMV default to X2E/A8X * opencl: fix q4_K variant kernels to read the transposed scales layout * opencl: keep the flat-GEMV large-m escape opt-in * opencl: guard the o4 GEMV store against the rounded-up dispatch tail * opencl: restore the tiled q4_K/q6_K layout on tensor read-back * opencl: split-K for the q8_0 decode GEMV at small M * opencl: keep the q6_K noshuffle correctness escape ahead of the opt-in gate --- ggml/src/ggml-opencl/CMakeLists.txt | 6 + ggml/src/ggml-opencl/ggml-opencl.cpp | 1611 +++++++++++++++-- ggml/src/ggml-opencl/kernels/cvt.cl | 171 ++ .../kernels/gemm_noshuffle_q4_k_f32.cl | 317 ++++ .../kernels/gemm_noshuffle_q6_k_f32.cl | 105 ++ .../kernels/gemm_noshuffle_q6_k_f32_tiled.cl | 136 ++ .../kernels/gemv_noshuffle_q4_0_f32.cl | 104 ++ .../kernels/gemv_noshuffle_q4_1_f32.cl | 96 + .../kernels/gemv_noshuffle_q4_k_f32.cl | 520 +++++- .../kernels/gemv_noshuffle_q4_k_f32_o4.cl | 349 ++++ .../kernels/gemv_noshuffle_q4_k_f32_tiled.cl | 118 ++ .../kernels/gemv_noshuffle_q5_k_f32.cl | 122 ++ .../kernels/gemv_noshuffle_q6_k_f32.cl | 111 ++ .../kernels/gemv_noshuffle_q6_k_f32_o4.cl | 372 ++++ .../kernels/gemv_noshuffle_q6_k_f32_tiled.cl | 196 ++ .../kernels/gemv_noshuffle_q8_0_f32.cl | 81 + .../kernels/mul_mm_f32_f32_l4_lm.cl | 49 + .../kernels/mul_mv_f16_f32_mrow.cl | 306 ++++ ggml/src/ggml-opencl/kernels/rms_norm.cl | 179 ++ 19 files changed, 4829 insertions(+), 120 deletions(-) create mode 100644 ggml/src/ggml-opencl/kernels/gemm_noshuffle_q6_k_f32_tiled.cl create mode 100644 ggml/src/ggml-opencl/kernels/gemv_noshuffle_q4_k_f32_o4.cl create mode 100644 ggml/src/ggml-opencl/kernels/gemv_noshuffle_q4_k_f32_tiled.cl create mode 100644 ggml/src/ggml-opencl/kernels/gemv_noshuffle_q6_k_f32_o4.cl create mode 100644 ggml/src/ggml-opencl/kernels/gemv_noshuffle_q6_k_f32_tiled.cl create mode 100644 ggml/src/ggml-opencl/kernels/mul_mv_f16_f32_mrow.cl diff --git a/ggml/src/ggml-opencl/CMakeLists.txt b/ggml/src/ggml-opencl/CMakeLists.txt index 1f62ce1c6a7..8a1b6b964a1 100644 --- a/ggml/src/ggml-opencl/CMakeLists.txt +++ b/ggml/src/ggml-opencl/CMakeLists.txt @@ -85,6 +85,7 @@ set(GGML_OPENCL_KERNELS mul_mv_f16_f32_1row mul_mv_f16_f32_l4 mul_mv_f16_f32 + mul_mv_f16_f32_mrow mul_mv_f32_f32 mul_mv_q1_0_f32 mul_mv_q1_0_f32_flat @@ -180,9 +181,14 @@ set(GGML_OPENCL_KERNELS gemv_noshuffle_q8_0_f32 gemm_noshuffle_q8_0_f32 gemv_noshuffle_q4_k_f32 + gemv_noshuffle_q4_k_f32_o4 + gemv_noshuffle_q4_k_f32_tiled gemm_noshuffle_q4_k_f32 gemv_noshuffle_q6_k_f32 + gemv_noshuffle_q6_k_f32_o4 + gemv_noshuffle_q6_k_f32_tiled gemm_noshuffle_q6_k_f32 + gemm_noshuffle_q6_k_f32_tiled gemv_noshuffle_q5_k_f32 gemm_noshuffle_q5_k_f32 mul diff --git a/ggml/src/ggml-opencl/ggml-opencl.cpp b/ggml/src/ggml-opencl/ggml-opencl.cpp index d95123eb151..12465a517d4 100644 --- a/ggml/src/ggml-opencl/ggml-opencl.cpp +++ b/ggml/src/ggml-opencl/ggml-opencl.cpp @@ -568,6 +568,10 @@ struct ggml_backend_opencl_context { bool has_integer_dot = false; // cl_khr_integer_dot_product or cl_qcom_dot_product8 bool has_qcom_subgroup_shuffle = false; // specifically cl_qcom_subgroup_shuffle bool disable_fusion; + bool fuse_mm_glu = true; // opt-out GGML_OPENCL_FUSE_MM_GLU=0 (byte-identical gate+up GEMV + GLU, q4_K FFN) + bool fuse_rms_add = true; // opt-out GGML_OPENCL_FUSE_RMS_ADD=0 (fused rms_norm*w + residual) + bool f16_mrow = true; // opt-out GGML_OPENCL_F16_MROW=0 (multi-row-per-WG f16 decode GEMV for attn proj + lm_head) + int f16_mrow_rpt = 1; // GGML_OPENCL_F16_MROW_RPT={1,2,4,8,16} rows-per-subgroup register blocking // ragged moe, use int to directly pass to kernel cl_uint adreno_use_moe_ragged; @@ -619,6 +623,7 @@ struct ggml_backend_opencl_context { ggml_cl_buffer prealloc_moe_sa; // per-block s [tok_slots * ne00/32] (half) // scratch copy of the router weights to avoid dst aliasing ggml_cl_buffer prealloc_moe_combine_w; + ggml_cl_buffer prealloc_splitk_partial; // [ksplit * M] partials for split-K GEMV // pool of persistent image1d_buffer views over kv-cache layers, keyed by // (parent buffer, offset within parent) @@ -749,6 +754,7 @@ struct ggml_backend_opencl_context { kernel_geglu_erf_f16, kernel_geglu_quick_f16; cl_kernel kernel_norm, kernel_norm_mul_add; cl_kernel kernel_rms_norm, kernel_rms_norm_mul; + cl_kernel kernel_rms_norm_mul_add = nullptr; // fused rms_norm(x)*w + b (residual) cl_kernel kernel_l2_norm_f32; cl_kernel kernel_group_norm, kernel_group_norm_mul_add; cl_kernel kernel_diag_mask_inf, kernel_diag_mask_inf_8; @@ -768,6 +774,12 @@ struct ggml_backend_opencl_context { cl_kernel kernel_mul_mat_f32_f32; cl_kernel kernel_mul_mat_f16_f16; cl_kernel kernel_mul_mat_f16_f32_1row; + cl_program program_mul_mv_f16_f32_mrow; + cl_kernel kernel_mul_mat_f16_f32_mrow = nullptr; // multi-row decode GEMV (attn proj + lm_head) + cl_kernel kernel_mul_mat_f16_f32_mrow_r2 = nullptr; + cl_kernel kernel_mul_mat_f16_f32_mrow_r4 = nullptr; + cl_kernel kernel_mul_mat_f16_f32_mrow_h8 = nullptr; + cl_kernel kernel_mul_mat_f16_f32_mrow_h8r2 = nullptr; cl_kernel kernel_mul_mat_f16_f32; cl_kernel kernel_mul_mat_f16_f32_l4; cl_kernel kernel_mul_mat_f16_f32_l4_dr; @@ -916,6 +928,7 @@ struct ggml_backend_opencl_context { cl_kernel kernel_mul_mv_id_mxfp4_f32; cl_kernel kernel_mul_mv_id_mxfp4_f32_flat; cl_kernel kernel_mul_mm_f32_f32_l4_lm; + cl_kernel kernel_gemv_f32_f32_mc; // multi-column (small-N) f32 GEMV for spec/MTP verify cl_kernel kernel_mul_mm_f16_f32_l4_lm; cl_kernel kernel_mul_mm_q1_0_f32_l4_lm; cl_kernel kernel_mul_mm_q4_0_f32_l4_lm; @@ -1081,28 +1094,50 @@ struct ggml_backend_opencl_context { // Gemm and Gemv related programs, kernels, etc cl_kernel kernel_gemm_noshuffle_q4_0_f32; cl_kernel kernel_gemv_noshuffle_q4_0_f32; + cl_kernel kernel_gemv_noshuffle_q4_0_f32_mc3; // multi-column (N=3) verify GEMV (spec/MTP) cl_kernel kernel_gemv_noshuffle_q4_0_f32_4096_1_11008; cl_kernel kernel_gemv_noshuffle_q4_0_f32_4096_1_4096; cl_kernel kernel_gemv_noshuffle_q4_0_f32_11008_1_4096; cl_kernel kernel_gemv_noshuffle_q4_0_f32_32000_1_4096; cl_kernel kernel_gemv_noshuffle_q4_1_f32; + cl_kernel kernel_gemv_noshuffle_q4_1_f32_mc3; // multi-column (N=3) verify GEMV (spec/MTP) cl_kernel kernel_gemm_noshuffle_q4_1_f32; cl_kernel kernel_gemm_noshuffle_q8_0_f32, kernel_gemm_noshuffle_q8_0_f32_bin; cl_kernel kernel_gemm_noshuffle_q8_0_q8_1_dp4a = nullptr; // dp4a (int8) dense q8_0 prefill GEMM (opt-in) cl_kernel kernel_gemm_noshuffle_q8_0_q8_1_dp4a_wimg = nullptr; // q8_0 dense dp4a, weights via texture (opt-in) cl_kernel kernel_gemv_noshuffle_q8_0_f32; + cl_kernel kernel_gemv_noshuffle_q8_0_f32_splitk; // split-K across WGs (small-M decode) cl_kernel kernel_gemm_noshuffle_q1_0_f32; cl_kernel kernel_gemv_noshuffle_q1_0_f32; cl_kernel kernel_gemv_noshuffle_q4_k_f32; + cl_kernel kernel_gemv_noshuffle_q4_k_f32_o4; // 4-output-per-WI, long-vocab lm_head + cl_kernel kernel_gemv_noshuffle_q4_k_f32_tiled; // tiled-wide layout (opt-in) + cl_kernel kernel_gemv_noshuffle_q4_k_f32_splitk; // split-K across WGs (small-M decode) + cl_kernel kernel_gemv_splitk_reduce_f32; // sums split-K per-slice partials + cl_kernel kernel_gemv_noshuffle_q4_k_f32_glu; // fused gate+up GEMV + GLU (FFN) + cl_kernel kernel_convert_block_q4_k_tiled_ns; // tiled-wide convert (opt-in) + cl_kernel kernel_gemv_noshuffle_q4_k_f32_mc3; // multi-column (N=3) verify GEMV cl_kernel kernel_gemm_noshuffle_q4_k_f32; cl_kernel kernel_gemm_noshuffle_q4_k_q8_1_dp4a = nullptr; // dp4a (int8) dense prefill GEMM cl_kernel kernel_gemm_noshuffle_q4_k_q8_1_dp4a_wimg = nullptr; // dp4a dense prefill GEMM, weights via texture (X1 opt-in) cl_kernel kernel_gemm_noshuffle_q5_k_q8_1_dp4a = nullptr; // dp4a (int8) dense q5_K prefill GEMM cl_kernel kernel_gemm_noshuffle_q6_k_q8_1_dp4a = nullptr; // dp4a (int8) dense q6_K prefill GEMM cl_kernel kernel_quant_a_q8_1; // plain activation q8_1 pre-pass + cl_kernel kernel_gemm_noshuffle_q4_k_f32_r1; + cl_kernel kernel_gemm_noshuffle_q4_k_f32_kimg; + cl_kernel kernel_gemm_noshuffle_q4_k_f32_cok; cl_kernel kernel_gemv_noshuffle_q6_K_f32; + cl_kernel kernel_gemv_noshuffle_q6_K_f32_o4; + cl_kernel kernel_gemv_noshuffle_q6_K_f32_o4_global; // weights via __global (opt-in) + cl_kernel kernel_gemv_noshuffle_q6_K_f32_tiled; // tiled-wide layout (opt-in) + cl_kernel kernel_gemv_noshuffle_q6_K_f32_tiled_mc3; // tiled multi-column (N=3) verify lm_head + cl_kernel kernel_gemm_noshuffle_q6_K_f32_tiled; // batched (N>1) over the tiled layout + cl_kernel kernel_convert_block_q6_k_tiled_ns; // tiled-wide convert (opt-in) + cl_kernel kernel_gemv_noshuffle_q6_K_f32_mc3; // multi-column (N=3) verify GEMV cl_kernel kernel_gemm_noshuffle_q6_K_f32; + cl_kernel kernel_gemm_noshuffle_q6_K_f32_cok; cl_kernel kernel_gemv_noshuffle_q5_k_f32; + cl_kernel kernel_gemv_noshuffle_q5_k_f32_mc3; // multi-column (N=3) verify GEMV (spec/MTP) cl_kernel kernel_gemm_noshuffle_q5_k_f32; cl_kernel kernel_gemv_noshuffle_q5_0_f32; cl_kernel kernel_gemm_noshuffle_q5_0_f32; @@ -1488,10 +1523,16 @@ static void load_cl_kernels(ggml_backend_opencl_context *backend_ctx) { CL_CHECK((backend_ctx->kernel_restore_block_q5_1_trans4_ns = clCreateKernel(backend_ctx->program_cvt, "kernel_restore_block_q5_1_trans4_ns", &err), err)); CL_CHECK((backend_ctx->kernel_convert_block_q4_k_trans4_ns = clCreateKernel(backend_ctx->program_cvt, "kernel_convert_block_q4_k_trans4_ns", &err), err)); CL_CHECK((backend_ctx->kernel_restore_block_q4_k_trans4_ns = clCreateKernel(backend_ctx->program_cvt, "kernel_restore_block_q4_k_trans4_ns", &err), err)); +#ifdef GGML_OPENCL_USE_ADRENO_KERNELS + CL_CHECK((backend_ctx->kernel_convert_block_q4_k_tiled_ns = clCreateKernel(backend_ctx->program_cvt, "kernel_convert_block_q4_k_tiled_ns", &err), err)); +#endif CL_CHECK((backend_ctx->kernel_convert_block_q5_k_trans4_ns = clCreateKernel(backend_ctx->program_cvt, "kernel_convert_block_q5_k_trans4_ns", &err), err)); CL_CHECK((backend_ctx->kernel_restore_block_q5_k_trans4_ns = clCreateKernel(backend_ctx->program_cvt, "kernel_restore_block_q5_k_trans4_ns", &err), err)); CL_CHECK((backend_ctx->kernel_convert_block_q6_k_trans4_ns = clCreateKernel(backend_ctx->program_cvt, "kernel_convert_block_q6_k_trans4_ns", &err), err)); CL_CHECK((backend_ctx->kernel_restore_block_q6_k_trans4_ns = clCreateKernel(backend_ctx->program_cvt, "kernel_restore_block_q6_k_trans4_ns", &err), err)); +#ifdef GGML_OPENCL_USE_ADRENO_KERNELS + CL_CHECK((backend_ctx->kernel_convert_block_q6_k_tiled_ns = clCreateKernel(backend_ctx->program_cvt, "kernel_convert_block_q6_k_tiled_ns", &err), err)); +#endif CL_CHECK((backend_ctx->kernel_convert_block_mxfp4 = clCreateKernel(backend_ctx->program_cvt, "kernel_convert_block_mxfp4", &err), err)); CL_CHECK((backend_ctx->kernel_convert_block_mxfp4_trans = clCreateKernel(backend_ctx->program_cvt, "kernel_convert_block_mxfp4_trans", &err), err)); CL_CHECK((backend_ctx->kernel_convert_block_mxfp4_trans4_ns = clCreateKernel(backend_ctx->program_cvt, "kernel_convert_block_mxfp4_trans4_ns", &err), err)); @@ -2142,6 +2183,26 @@ static void load_cl_kernels(ggml_backend_opencl_context *backend_ctx) { GGML_LOG_CONT("."); } + // mul_mv_f16_f32_mrow (multi-row decode GEMV) + { +#ifdef GGML_OPENCL_EMBED_KERNELS + const std::string kernel_src { + #include "mul_mv_f16_f32_mrow.cl.h" + }; +#else + const std::string kernel_src = read_file("mul_mv_f16_f32_mrow.cl"); +#endif + backend_ctx->program_mul_mv_f16_f32_mrow = + build_program_from_source(backend_ctx, kernel_src.c_str(), compile_opts); + + CL_CHECK((backend_ctx->kernel_mul_mat_f16_f32_mrow = clCreateKernel(backend_ctx->program_mul_mv_f16_f32_mrow, "kernel_mul_mat_f16_f32_mrow", &err), err)); + CL_CHECK((backend_ctx->kernel_mul_mat_f16_f32_mrow_r2 = clCreateKernel(backend_ctx->program_mul_mv_f16_f32_mrow, "kernel_mul_mat_f16_f32_mrow_r2", &err), err)); + CL_CHECK((backend_ctx->kernel_mul_mat_f16_f32_mrow_r4 = clCreateKernel(backend_ctx->program_mul_mv_f16_f32_mrow, "kernel_mul_mat_f16_f32_mrow_r4", &err), err)); + CL_CHECK((backend_ctx->kernel_mul_mat_f16_f32_mrow_h8 = clCreateKernel(backend_ctx->program_mul_mv_f16_f32_mrow, "kernel_mul_mat_f16_f32_mrow_h8", &err), err)); + CL_CHECK((backend_ctx->kernel_mul_mat_f16_f32_mrow_h8r2 = clCreateKernel(backend_ctx->program_mul_mv_f16_f32_mrow, "kernel_mul_mat_f16_f32_mrow_h8r2", &err), err)); + GGML_LOG_CONT("."); + } + // mul_mv_f16_f32_l4 { #ifdef GGML_OPENCL_EMBED_KERNELS @@ -2295,6 +2356,7 @@ static void load_cl_kernels(ggml_backend_opencl_context *backend_ctx) { build_program_from_source(backend_ctx, kernel_src.c_str(), compile_opts); CL_CHECK((backend_ctx->kernel_mul_mm_f32_f32_l4_lm = clCreateKernel(backend_ctx->program_mul_mm_f32_f32_l4_lm, "kernel_mul_mm_f32_f32_l4_lm", &err), err)); + CL_CHECK((backend_ctx->kernel_gemv_f32_f32_mc = clCreateKernel(backend_ctx->program_mul_mm_f32_f32_l4_lm, "kernel_gemv_f32_f32_mc", &err), err)); GGML_LOG_CONT("."); } @@ -2564,6 +2626,7 @@ static void load_cl_kernels(ggml_backend_opencl_context *backend_ctx) { CL_CHECK((backend_ctx->kernel_rms_norm = clCreateKernel(backend_ctx->program_rms_norm, "kernel_rms_norm", &err), err)); CL_CHECK((backend_ctx->kernel_rms_norm_mul = clCreateKernel(backend_ctx->program_rms_norm, "kernel_rms_norm_mul", &err), err)); + CL_CHECK((backend_ctx->kernel_rms_norm_mul_add = clCreateKernel(backend_ctx->program_rms_norm, "kernel_rms_norm_mul_add", &err), err)); GGML_LOG_CONT("."); } @@ -3475,6 +3538,7 @@ static void load_cl_kernels(ggml_backend_opencl_context *backend_ctx) { cl_program prog = build_program_from_source(backend_ctx, kernel_src_CL_gemv_general.c_str(), CL_gemv_compile_opts); CL_CHECK((backend_ctx->kernel_gemv_noshuffle_q4_0_f32 = clCreateKernel(prog, "kernel_gemv_noshuffle_q4_0_f32", &err), err)); + CL_CHECK((backend_ctx->kernel_gemv_noshuffle_q4_0_f32_mc3 = clCreateKernel(prog, "kernel_gemv_noshuffle_q4_0_f32_mc3", &err), err)); CL_CHECK(clReleaseProgram(prog)); GGML_LOG_CONT("."); } @@ -3604,6 +3668,7 @@ static void load_cl_kernels(ggml_backend_opencl_context *backend_ctx) { cl_program prog = build_program_from_source(backend_ctx, kernel_src.c_str(), CL_gemv_compile_opts); CL_CHECK((backend_ctx->kernel_gemv_noshuffle_q4_1_f32 = clCreateKernel(prog, "kernel_gemv_noshuffle_q4_1_f32", &err), err)); + CL_CHECK((backend_ctx->kernel_gemv_noshuffle_q4_1_f32_mc3 = clCreateKernel(prog, "kernel_gemv_noshuffle_q4_1_f32_mc3", &err), err)); CL_CHECK(clReleaseProgram(prog)); GGML_LOG_CONT("."); } @@ -3818,6 +3883,7 @@ static void load_cl_kernels(ggml_backend_opencl_context *backend_ctx) { cl_program prog = build_program_from_source(backend_ctx, kernel_src_CL_gemv_general.c_str(), CL_gemv_compile_opts); CL_CHECK((backend_ctx->kernel_gemv_noshuffle_q8_0_f32 = clCreateKernel(prog, "kernel_gemv_noshuffle_q8_0_f32", &err), err)); + CL_CHECK((backend_ctx->kernel_gemv_noshuffle_q8_0_f32_splitk = clCreateKernel(prog, "kernel_gemv_noshuffle_q8_0_f32_splitk", &err), err)); CL_CHECK(clReleaseProgram(prog)); GGML_LOG_CONT("."); } @@ -3833,6 +3899,9 @@ static void load_cl_kernels(ggml_backend_opencl_context *backend_ctx) { #endif cl_program prog = build_program_from_source(backend_ctx, kernel_src.c_str(), compile_opts); CL_CHECK((backend_ctx->kernel_gemm_noshuffle_q4_k_f32 = clCreateKernel(prog, "kernel_gemm_noshuffle_q4_k_f32", &err), err)); + CL_CHECK((backend_ctx->kernel_gemm_noshuffle_q4_k_f32_r1 = clCreateKernel(prog, "kernel_gemm_noshuffle_q4_k_f32_r1", &err), err)); + CL_CHECK((backend_ctx->kernel_gemm_noshuffle_q4_k_f32_kimg = clCreateKernel(prog, "kernel_gemm_noshuffle_q4_k_f32_kimg", &err), err)); + CL_CHECK((backend_ctx->kernel_gemm_noshuffle_q4_k_f32_cok = clCreateKernel(prog, "kernel_gemm_noshuffle_q4_k_f32_cok", &err), err)); CL_CHECK(clReleaseProgram(prog)); GGML_LOG_CONT("."); } @@ -3927,6 +3996,18 @@ static void load_cl_kernels(ggml_backend_opencl_context *backend_ctx) { if (backend_ctx->has_vector_subgroup_broadcast) { CL_gemv_compile_opts += " -DVECTOR_SUB_GROUP_BROADCAST "; } + // Opt-in: dequant-once-per-block mc3 verify GEMV (factors q4_K dequant + // out of the 3-column loop; byte-identical, lower spill). A/B vs the + // shipped inline mc3 in the same binary. + if (getenv("GGML_OPENCL_Q4K_MC3_DQ")) { + CL_gemv_compile_opts += " -DQ4K_MC3_DEQUANT_ONCE "; + } + // Opt-in: LDS-staged dequant mc3 verify GEMV (stages the dequantized + // q4_K weights in __local instead of private regs that spill to slow + // global on Adreno; byte-identical). A/B vs inline + dequant-once. + if (getenv("GGML_OPENCL_Q4K_MC3_LDS")) { + CL_gemv_compile_opts += " -DQ4K_MC3_DEQUANT_LDS "; + } #ifdef GGML_OPENCL_EMBED_KERNELS const std::string kernel_src { @@ -3939,6 +4020,50 @@ static void load_cl_kernels(ggml_backend_opencl_context *backend_ctx) { cl_program prog = build_program_from_source(backend_ctx, kernel_src.c_str(), CL_gemv_compile_opts); CL_CHECK((backend_ctx->kernel_gemv_noshuffle_q4_k_f32 = clCreateKernel(prog, "kernel_gemv_noshuffle_q4_k_f32", &err), err)); + CL_CHECK((backend_ctx->kernel_gemv_noshuffle_q4_k_f32_mc3 = clCreateKernel(prog, "kernel_gemv_noshuffle_q4_k_f32_mc3", &err), err)); + CL_CHECK((backend_ctx->kernel_gemv_noshuffle_q4_k_f32_splitk = clCreateKernel(prog, "kernel_gemv_noshuffle_q4_k_f32_splitk", &err), err)); + CL_CHECK((backend_ctx->kernel_gemv_splitk_reduce_f32 = clCreateKernel(prog, "kernel_gemv_splitk_reduce_f32", &err), err)); + CL_CHECK((backend_ctx->kernel_gemv_noshuffle_q4_k_f32_glu = clCreateKernel(prog, "kernel_gemv_noshuffle_q4_k_f32_glu", &err), err)); + CL_CHECK(clReleaseProgram(prog)); + GGML_LOG_CONT("."); + } + + // gemv_noshuffle_q4_k_f32_o4 — 4-output-per-WI variant for the long-vocab + // q4_K lm_head/embed GEMV (shares one activation read across 4 output rows). + { +#ifdef GGML_OPENCL_EMBED_KERNELS + const std::string kernel_src { + #include "gemv_noshuffle_q4_k_f32_o4.cl.h" + }; +#else + const std::string kernel_src = read_file("gemv_noshuffle_q4_k_f32_o4.cl"); +#endif + std::string CL_gemv_compile_opts = std::string("-cl-std=") + opencl_c_std + " -cl-mad-enable "; + if (backend_ctx->has_vector_subgroup_broadcast) { + CL_gemv_compile_opts += " -DVECTOR_SUB_GROUP_BROADCAST "; + } + cl_program prog = build_program_from_source( + backend_ctx, kernel_src.c_str(), CL_gemv_compile_opts); + CL_CHECK((backend_ctx->kernel_gemv_noshuffle_q4_k_f32_o4 = clCreateKernel(prog, "kernel_gemv_noshuffle_q4_k_f32_o4", &err), err)); + CL_CHECK(clReleaseProgram(prog)); + GGML_LOG_CONT("."); + } + + // gemv_noshuffle_q4_k_f32_tiled — tiled-wide canonical layout, default ON + // (opt out: GGML_OPENCL_Q4K_GEMV_TILED=0; separate convert + GEMV; weights via __global). + { +#ifdef GGML_OPENCL_EMBED_KERNELS + const std::string kernel_src { + #include "gemv_noshuffle_q4_k_f32_tiled.cl.h" + }; +#else + const std::string kernel_src = read_file("gemv_noshuffle_q4_k_f32_tiled.cl"); +#endif + std::string compile_opts = std::string("-cl-std=") + opencl_c_std + " -cl-mad-enable "; + cl_program prog = + build_program_from_source(backend_ctx, kernel_src.c_str(), compile_opts); + CL_CHECK((backend_ctx->kernel_gemv_noshuffle_q4_k_f32_tiled = + clCreateKernel(prog, "kernel_gemv_noshuffle_q4_k_f32_tiled", &err), err)); CL_CHECK(clReleaseProgram(prog)); GGML_LOG_CONT("."); } @@ -4566,6 +4691,91 @@ static void load_cl_kernels(ggml_backend_opencl_context *backend_ctx) { build_program_from_source(backend_ctx, kernel_src.c_str(), CL_gemv_compile_opts); CL_CHECK((backend_ctx->kernel_gemv_noshuffle_q6_K_f32 = clCreateKernel(prog, "kernel_gemv_noshuffle_q6_K_f32", &err), err)); + CL_CHECK((backend_ctx->kernel_gemv_noshuffle_q6_K_f32_mc3 = clCreateKernel(prog, "kernel_gemv_noshuffle_q6_K_f32_mc3", &err), err)); + if (getenv("GGML_OPENCL_MC3_PROBE")) { + cl_ulong pm6 = 0, pm4 = 0; size_t wg6 = 0, wg4 = 0, mult = 0; + clGetKernelWorkGroupInfo(backend_ctx->kernel_gemv_noshuffle_q6_K_f32_mc3, backend_ctx->device, CL_KERNEL_PRIVATE_MEM_SIZE, sizeof(pm6), &pm6, NULL); + clGetKernelWorkGroupInfo(backend_ctx->kernel_gemv_noshuffle_q6_K_f32_mc3, backend_ctx->device, CL_KERNEL_WORK_GROUP_SIZE, sizeof(wg6), &wg6, NULL); + clGetKernelWorkGroupInfo(backend_ctx->kernel_gemv_noshuffle_q4_k_f32_mc3, backend_ctx->device, CL_KERNEL_PRIVATE_MEM_SIZE, sizeof(pm4), &pm4, NULL); + clGetKernelWorkGroupInfo(backend_ctx->kernel_gemv_noshuffle_q4_k_f32_mc3, backend_ctx->device, CL_KERNEL_WORK_GROUP_SIZE, sizeof(wg4), &wg4, NULL); + clGetKernelWorkGroupInfo(backend_ctx->kernel_gemv_noshuffle_q6_K_f32_mc3, backend_ctx->device, CL_KERNEL_PREFERRED_WORK_GROUP_SIZE_MULTIPLE, sizeof(mult), &mult, NULL); + fprintf(stderr, "[MC3-PROBE] q4K_mc3 private=%llu wg_cap=%zu | q6K_mc3 private=%llu wg_cap=%zu | pref_mult=%zu\n", + (unsigned long long)pm4, wg4, (unsigned long long)pm6, wg6, mult); + fflush(stderr); + } + GGML_LOG_CONT("."); + } + + // gemv_noshuffle_q6_k_f32_o4 — 4-output-per-WI variant, opt-in via + // GGML_OPENCL_Q6K_GEMV_O4=1 (~3x fewer dispatches on long-vocab lm_head). + { +#ifdef GGML_OPENCL_EMBED_KERNELS + const std::string kernel_src { + #include "gemv_noshuffle_q6_k_f32_o4.cl.h" + }; +#else + const std::string kernel_src = read_file("gemv_noshuffle_q6_k_f32_o4.cl"); +#endif + + std::string CL_gemv_compile_opts = std::string("-cl-std=") + opencl_c_std + + " -cl-mad-enable "; + if (backend_ctx->has_vector_subgroup_broadcast) { + CL_gemv_compile_opts += " -DVECTOR_SUB_GROUP_BROADCAT "; + } + + cl_program prog = + build_program_from_source(backend_ctx, kernel_src.c_str(), CL_gemv_compile_opts); + + CL_CHECK((backend_ctx->kernel_gemv_noshuffle_q6_K_f32_o4 = clCreateKernel(prog, "kernel_gemv_noshuffle_q6_K_f32_o4", &err), err)); + CL_CHECK(clReleaseProgram(prog)); + + // Global-read variant: weights read from __global coalesced instead of + // image1d_buffer (the texture cache caps the streaming lm_head read + // bandwidth). Opt-in via GGML_OPENCL_Q6K_GEMV_O4_GLOBAL. + cl_program prog_g = build_program_from_source(backend_ctx, kernel_src.c_str(), CL_gemv_compile_opts + " -DQ6K_O4_GLOBAL"); + CL_CHECK((backend_ctx->kernel_gemv_noshuffle_q6_K_f32_o4_global = + clCreateKernel(prog_g, "kernel_gemv_noshuffle_q6_K_f32_o4_global", &err), err)); + CL_CHECK(clReleaseProgram(prog_g)); + GGML_LOG_CONT("."); + } + + // gemv_noshuffle_q6_k_f32_tiled — tiled-wide canonical layout, default ON + // (opt out: GGML_OPENCL_Q6K_GEMV_TILED=0; separate convert + GEMV; weights via __global). + { +#ifdef GGML_OPENCL_EMBED_KERNELS + const std::string kernel_src { + #include "gemv_noshuffle_q6_k_f32_tiled.cl.h" + }; +#else + const std::string kernel_src = read_file("gemv_noshuffle_q6_k_f32_tiled.cl"); +#endif + std::string compile_opts = std::string("-cl-std=") + opencl_c_std + " -cl-mad-enable "; + cl_program prog = + build_program_from_source(backend_ctx, kernel_src.c_str(), compile_opts); + CL_CHECK((backend_ctx->kernel_gemv_noshuffle_q6_K_f32_tiled = + clCreateKernel(prog, "kernel_gemv_noshuffle_q6_K_f32_tiled", &err), err)); + CL_CHECK((backend_ctx->kernel_gemv_noshuffle_q6_K_f32_tiled_mc3 = + clCreateKernel(prog, "kernel_gemv_noshuffle_q6_K_f32_tiled_mc3", &err), err)); + CL_CHECK(clReleaseProgram(prog)); + GGML_LOG_CONT("."); + } + + // gemm_noshuffle_q6_k_f32_tiled — batched (N>1) GEMM over the same tiled-wide + // canonical layout, so batched lm_head/embed stays correct + on GPU. + { +#ifdef GGML_OPENCL_EMBED_KERNELS + const std::string kernel_src { + #include "gemm_noshuffle_q6_k_f32_tiled.cl.h" + }; +#else + const std::string kernel_src = read_file("gemm_noshuffle_q6_k_f32_tiled.cl"); +#endif + std::string compile_opts = std::string("-cl-std=") + opencl_c_std + " -cl-mad-enable "; + cl_program prog = + build_program_from_source(backend_ctx, kernel_src.c_str(), compile_opts); + CL_CHECK((backend_ctx->kernel_gemm_noshuffle_q6_K_f32_tiled = + clCreateKernel(prog, "kernel_gemm_noshuffle_q6_K_f32_tiled", &err), err)); + CL_CHECK(clReleaseProgram(prog)); GGML_LOG_CONT("."); } @@ -4582,6 +4792,7 @@ static void load_cl_kernels(ggml_backend_opencl_context *backend_ctx) { build_program_from_source(backend_ctx, kernel_src.c_str(), CL_moe_compile_opts); CL_CHECK((backend_ctx->kernel_gemm_noshuffle_q6_K_f32 = clCreateKernel(prog, "kernel_gemm_noshuffle_q6_K_f32", &err), err)); + CL_CHECK((backend_ctx->kernel_gemm_noshuffle_q6_K_f32_cok = clCreateKernel(prog, "kernel_gemm_noshuffle_q6_K_f32_cok", &err), err)); GGML_LOG_CONT("."); } @@ -4604,6 +4815,7 @@ static void load_cl_kernels(ggml_backend_opencl_context *backend_ctx) { cl_program prog = build_program_from_source(backend_ctx, kernel_src.c_str(), CL_gemv_compile_opts); CL_CHECK((backend_ctx->kernel_gemv_noshuffle_q5_k_f32 = clCreateKernel(prog, "kernel_gemv_noshuffle_q5_k_f32", &err), err)); + CL_CHECK((backend_ctx->kernel_gemv_noshuffle_q5_k_f32_mc3 = clCreateKernel(prog, "kernel_gemv_noshuffle_q5_k_f32_mc3", &err), err)); CL_CHECK(clReleaseProgram(prog)); GGML_LOG_CONT("."); } @@ -6166,6 +6378,19 @@ static ggml_backend_opencl_context * ggml_cl_init(ggml_backend_dev_t dev) { #endif // GGML_OPENCL_USE_ADRENO_KERNELS backend_ctx->disable_fusion = getenv("GGML_OPENCL_DISABLE_FUSION") != nullptr; + if (const char * env = getenv("GGML_OPENCL_FUSE_MM_GLU")) { + backend_ctx->fuse_mm_glu = atoi(env) != 0; + } + if (const char * env = getenv("GGML_OPENCL_FUSE_RMS_ADD")) { + backend_ctx->fuse_rms_add = atoi(env) != 0; + } + if (const char * env = getenv("GGML_OPENCL_F16_MROW")) { + backend_ctx->f16_mrow = atoi(env) != 0; + } + if (const char * env = getenv("GGML_OPENCL_F16_MROW_RPT")) { + const int v = atoi(env); + backend_ctx->f16_mrow_rpt = (v == 2 || v == 4 || v == 8 || v == 16) ? v : 1; + } dev_ctx->backend_ctx = backend_ctx.release(); return dev_ctx->backend_ctx; @@ -7305,7 +7530,73 @@ static void ggml_cl_moe_combine_fused(ggml_backend_t backend, const ggml_tensor backend_ctx->enqueue_ndrange_kernel(kernel, 2, gws, lws, dst); } -static bool ggml_opencl_can_fuse(const struct ggml_cgraph * cgraph, int node_idx, std::initializer_list ops) { +inline bool use_q4k_tiled(const ggml_backend_opencl_context *backend_ctx, const ggml_tensor *tensor); // defined below (used by the GLU-subgraph fuse check) +inline bool use_adreno_kernels(const ggml_backend_opencl_context *backend_ctx, const ggml_tensor *tensor); // defined below + +static bool ggml_opencl_can_fuse(const ggml_backend_opencl_context * backend_ctx, const struct ggml_cgraph * cgraph, int node_idx, std::initializer_list ops) { + + // glu(mul_mat(Wg,x), mul_mat(Wu,x)) — the FFN gate/up GEMVs + GLU. This is a + // non-linear subgraph (up does NOT consume gate), so the contiguous + // ggml_can_fuse below rejects it; use ggml_can_fuse_subgraph with the glu as + // the sole output and validate the edges explicitly. q4_K decode only; + // byte-identical to the per-op path. + if (ops.size() == 3 && ops.begin()[0] == GGML_OP_MUL_MAT && + ops.begin()[1] == GGML_OP_MUL_MAT && ops.begin()[2] == GGML_OP_GLU) { + const enum ggml_op glu_ops[] = { GGML_OP_MUL_MAT, GGML_OP_MUL_MAT, GGML_OP_GLU }; + const int glu_out[] = { node_idx + 2 }; + if (!ggml_can_fuse_subgraph(cgraph, node_idx, 3, glu_ops, glu_out, 1)) { + return false; + } + + const ggml_tensor *gate = cgraph->nodes[node_idx]; + const ggml_tensor *up = cgraph->nodes[node_idx+1]; + const ggml_tensor *glu = cgraph->nodes[node_idx+2]; + + // decode GEMV path only (single token); prefill GEMM is separate + if (gate->ne[1] != 1 || up->ne[1] != 1) { + return false; + } + // both projections must be q4_K weights, f32 activation/output + if (gate->src[0]->type != GGML_TYPE_Q4_K || up->src[0]->type != GGML_TYPE_Q4_K || + gate->src[1]->type != GGML_TYPE_F32 || up->src[1]->type != GGML_TYPE_F32 || + gate->type != GGML_TYPE_F32 || up->type != GGML_TYPE_F32 || glu->type != GGML_TYPE_F32) { + return false; + } + // gate and up must share the same activation and have matching shape/stride + if (gate->src[1] != up->src[1] || + !ggml_are_same_shape(gate->src[0], up->src[0]) || + !ggml_are_same_stride(gate->src[0], up->src[0])) { + return false; + } + // GLU must read gate as src[0] and up as src[1], no swap (the fused + // epilogue applies the activation to gate, multiplies by up) + if (glu->src[0] != gate || glu->src[1] != up) { + return false; + } + if (ggml_get_op_params_i32(glu, 1) /* swapped */) { + return false; + } + // SWIGLU_OAI carries extra alpha/limit params -> not handled by the fused kernel + if (ggml_get_glu_op(glu) == GGML_GLU_OP_SWIGLU_OAI) { + return false; + } + // the fused kernel reads the standard noshuffle image layout; the tiled + // layout packs weights differently -> defer those to the per-op path + if (use_q4k_tiled(backend_ctx, gate->src[0]) || use_q4k_tiled(backend_ctx, up->src[0])) { + return false; + } + // that noshuffle layout is only produced at set_tensor time when + // use_adreno_kernels() accepts the weight (ne0 >= 512 && ne1 >= 512). + // Smaller weights stay in the plain q4_K layout, which this kernel would + // misread -> defer them to the per-op path. Real FFN gate/up weights are + // far above the threshold, so production dispatch is unchanged. + if (!use_adreno_kernels(backend_ctx, gate->src[0]) || + !use_adreno_kernels(backend_ctx, up->src[0])) { + return false; + } + return true; + } + if (!ggml_can_fuse(cgraph, node_idx, ops)) { return false; } @@ -7353,6 +7644,38 @@ static bool ggml_opencl_can_fuse(const struct ggml_cgraph * cgraph, int node_idx if (!ggml_is_contiguous(norm->src[0]) || !ggml_is_contiguous(w) || !ggml_is_contiguous(b)) { return false; } + } else if (ops.size() == 3 && ops.begin()[0] == GGML_OP_RMS_NORM && ops.begin()[1] == GGML_OP_MUL && ops.begin()[2] == GGML_OP_ADD) { + // rms_norm(x) * w + b, fused (residual). Mirrors the RMS_NORM+MUL gate + // plus the residual-add operand's constraints. + const ggml_tensor *rms_norm = cgraph->nodes[node_idx]; + const ggml_tensor *mul = cgraph->nodes[node_idx+1]; + const ggml_tensor *add = cgraph->nodes[node_idx+2]; + const ggml_tensor *w = mul->src[0] == rms_norm ? mul->src[1] : mul->src[0]; + const ggml_tensor *b = add->src[0] == mul ? add->src[1] : add->src[0]; + + GGML_ASSERT(rms_norm->src[0]->type == GGML_TYPE_F32); + GGML_ASSERT(rms_norm->type == GGML_TYPE_F32); + + if (w->type != GGML_TYPE_F32 || mul->type != GGML_TYPE_F32 || + b->type != GGML_TYPE_F32 || add->type != GGML_TYPE_F32) { + return false; + } + if (rms_norm->src[0]->ne[0] % 4 != 0) { + return false; + } + // if rms_norm is the B operand of mul, broadcast is not handled + if (rms_norm == mul->src[1] && !ggml_are_same_shape(mul->src[0], rms_norm)) { + return false; + } + // the residual must match the normed output shape (no add broadcast) + if (!ggml_are_same_shape(b, add)) { + return false; + } + // rms_norm assumes contiguous rows + if (!ggml_is_contiguous_rows(mul->src[0]) || !ggml_is_contiguous_rows(mul->src[1]) || + !ggml_is_contiguous_rows(b)) { + return false; + } } else if (ops.size() == 3 && ops.begin()[0] == GGML_OP_GROUP_NORM && ops.begin()[1] == GGML_OP_MUL && ops.begin()[2] == GGML_OP_ADD) { const ggml_tensor *gn = cgraph->nodes[node_idx]; const ggml_tensor *mul = cgraph->nodes[node_idx+1]; @@ -7376,6 +7699,216 @@ static void ggml_opencl_op_rms_norm_fused(ggml_backend_t backend, ggml_tensor * static void ggml_opencl_op_norm_fused(ggml_backend_t backend, ggml_tensor * norm_tensor, ggml_tensor * mul_tensor, ggml_tensor * add_tensor); static void ggml_opencl_op_group_norm_fused(ggml_backend_t backend, ggml_tensor * gn_tensor, ggml_tensor * mul_tensor, ggml_tensor * add_tensor); +static void ggml_cl_mul_mat_q4_k_glu_fused(ggml_backend_t backend, ggml_tensor * gate_tensor, ggml_tensor * up_tensor, ggml_tensor * glu_tensor) { +#ifdef GGML_OPENCL_USE_ADRENO_KERNELS + GGML_ASSERT(gate_tensor && up_tensor && glu_tensor); + + const ggml_tensor * Wg = gate_tensor->src[0]; + const ggml_tensor * Wu = up_tensor->src[0]; + const ggml_tensor * src1 = gate_tensor->src[1]; // == up_tensor->src[1] + const ggml_tensor * dst = glu_tensor; + + GGML_ASSERT(Wg && Wg->extra); + GGML_ASSERT(Wu && Wu->extra); + GGML_ASSERT(src1 && src1->extra); + GGML_ASSERT(dst && dst->extra); + + ggml_backend_opencl_context *backend_ctx = (ggml_backend_opencl_context *)backend->context; + + ggml_tensor_extra_cl * extra1 = (ggml_tensor_extra_cl *)src1->extra; + ggml_tensor_extra_cl * extrad = (ggml_tensor_extra_cl *)dst->extra; + ggml_tensor_extra_cl_q4_K * extra_g = (ggml_tensor_extra_cl_q4_K *)Wg->extra; + ggml_tensor_extra_cl_q4_K * extra_u = (ggml_tensor_extra_cl_q4_K *)Wu->extra; + + cl_ulong offset1 = extra1->offset + src1->view_offs; + cl_ulong offsetd = extrad->offset + dst->view_offs; + + const int K = Wg->ne[0]; // ne00 + const int M = Wg->ne[1]; // ne01 (= ffn intermediate width) + const int N = 1; // decode GEMV + + const cl_uchar mask_d6 = 0x3F, mask_d4 = 0x0F, mask_hi2 = 0xC0; + const int glu_op = (int)ggml_get_glu_op(dst); + + cl_context context = backend_ctx->context; + cl_int err; + cl_image_format img_fmt; + cl_image_desc img_desc; + cl_buffer_region region; + + // q images for the two weight matrices (standard noshuffle layout) + img_fmt = { CL_R, CL_UNSIGNED_INT32 }; + memset(&img_desc, 0, sizeof(img_desc)); + img_desc.image_type = CL_MEM_OBJECT_IMAGE1D_BUFFER; + img_desc.image_width = (size_t)M * K / 2 / 4; + img_desc.buffer = extra_g->q; + cl_mem qg_img = nullptr, qu_img = nullptr; + CL_CHECK((qg_img = clCreateImage(context, CL_MEM_READ_ONLY, &img_fmt, &img_desc, NULL, &err), err)); + img_desc.buffer = extra_u->q; + CL_CHECK((qu_img = clCreateImage(context, CL_MEM_READ_ONLY, &img_fmt, &img_desc, NULL, &err), err)); + + // shared activation image (one column at decode) + region.origin = offset1; + region.size = (size_t)K * N * sizeof(float); + cl_mem b_sub_buf = nullptr, b_img = nullptr; + CL_CHECK((b_sub_buf = clCreateSubBuffer(extra1->data_device, 0, CL_BUFFER_CREATE_TYPE_REGION, ®ion, &err), err)); + img_fmt = { CL_RGBA, CL_FLOAT }; + memset(&img_desc, 0, sizeof(img_desc)); + img_desc.image_type = CL_MEM_OBJECT_IMAGE1D_BUFFER; + img_desc.image_width = (size_t)K * N / 4; + img_desc.buffer = b_sub_buf; + CL_CHECK((b_img = clCreateImage(context, CL_MEM_READ_ONLY, &img_fmt, &img_desc, NULL, &err), err)); + + cl_kernel kernel = backend_ctx->kernel_gemv_noshuffle_q4_k_f32_glu; + CL_CHECK(clSetKernelArg(kernel, 0, sizeof(cl_mem), &qg_img)); + CL_CHECK(clSetKernelArg(kernel, 1, sizeof(cl_mem), &extra_g->d)); + CL_CHECK(clSetKernelArg(kernel, 2, sizeof(cl_mem), &extra_g->dm)); + CL_CHECK(clSetKernelArg(kernel, 3, sizeof(cl_mem), &extra_g->s)); + CL_CHECK(clSetKernelArg(kernel, 4, sizeof(cl_mem), &qu_img)); + CL_CHECK(clSetKernelArg(kernel, 5, sizeof(cl_mem), &extra_u->d)); + CL_CHECK(clSetKernelArg(kernel, 6, sizeof(cl_mem), &extra_u->dm)); + CL_CHECK(clSetKernelArg(kernel, 7, sizeof(cl_mem), &extra_u->s)); + CL_CHECK(clSetKernelArg(kernel, 8, sizeof(cl_mem), &b_img)); + CL_CHECK(clSetKernelArg(kernel, 9, sizeof(cl_mem), &extrad->data_device)); + CL_CHECK(clSetKernelArg(kernel, 10, sizeof(cl_ulong), &offsetd)); + CL_CHECK(clSetKernelArg(kernel, 11, sizeof(cl_int), &K)); + CL_CHECK(clSetKernelArg(kernel, 12, sizeof(cl_int), &M)); + CL_CHECK(clSetKernelArg(kernel, 13, sizeof(cl_int), &glu_op)); + CL_CHECK(clSetKernelArg(kernel, 14, sizeof(cl_uchar), &mask_d6)); + CL_CHECK(clSetKernelArg(kernel, 15, sizeof(cl_uchar), &mask_d4)); + CL_CHECK(clSetKernelArg(kernel, 16, sizeof(cl_uchar), &mask_hi2)); + + // K-split = nsg_y subgroups. HARD-CAP at 8 (512 work-items): the fused + // kernel's cross-subgroup reduce uses a float4 reduceLM (gate+up packed) = + // 2x the LDS of the base GEMV's float2 reduce, so 16 co-resident subgroups + // exceed the per-CU LDS budget on X2 and the WG barrier DEADLOCKS -> GPU TDR + // (reproduced on upstream gemma-4 E4B decode, K=2560 M=10240). This used to + // be masked: get_kernel_workgroup_size reported 896 for this kernel (so the + // cap loop fell to 8), but it now returns 1024 and the Adreno per-kernel WG + // query is unreliable (over-reports), so cap explicitly instead of trusting + // it. nsg_y < 16 also means the cross-subgroup accumulation grouping differs + // from the standalone wide (nsg=16) GEMV, so the output is coherent but NOT + // byte-identical to the per-op path. Keep the maxwg query as a further floor + // for any driver that reports < 512. + size_t maxwg = backend_ctx->get_kernel_workgroup_size(kernel); + size_t nsg_y = 8; + while (nsg_y > 1 && 64 * nsg_y > maxwg) { nsg_y >>= 1; } + size_t local_work_size[3] = { 64, nsg_y, 1 }; + size_t global_work_size[3] = { (size_t)CEIL_DIV(M / 2, 64) * 64, nsg_y, 1 }; + + if (getenv("GGML_OPENCL_FUSE_DEBUG")) { + static int dbg = 0; + if (dbg < 3) { fprintf(stderr, "[FUSE_MM_GLU] fired #%d K=%d M=%d glu_op=%d nsg=%zu maxwg=%zu\n", ++dbg, K, M, glu_op, nsg_y, maxwg); fflush(stderr); } + } + + backend_ctx->enqueue_ndrange_kernel(kernel, 3, global_work_size, local_work_size, dst); + + CL_CHECK(clReleaseMemObject(qg_img)); + CL_CHECK(clReleaseMemObject(qu_img)); + CL_CHECK(clReleaseMemObject(b_img)); + CL_CHECK(clReleaseMemObject(b_sub_buf)); +#else + GGML_UNUSED(backend); + GGML_UNUSED(gate_tensor); + GGML_UNUSED(up_tensor); + GGML_UNUSED(glu_tensor); + GGML_ABORT("q4_K GLU fusion requires GGML_OPENCL_USE_ADRENO_KERNELS"); +#endif +} + + +static void ggml_opencl_op_rms_norm_mul_add_fused(ggml_backend_t backend, ggml_tensor * rms_norm_tensor, ggml_tensor * mul_tensor, ggml_tensor * add_tensor) { + GGML_ASSERT(rms_norm_tensor && mul_tensor && add_tensor); + + const ggml_tensor * src0 = rms_norm_tensor->src[0]; + const ggml_tensor * src1 = mul_tensor->src[0] == rms_norm_tensor ? mul_tensor->src[1] : mul_tensor->src[0]; + const ggml_tensor * src2 = add_tensor->src[0] == mul_tensor ? add_tensor->src[1] : add_tensor->src[0]; + const ggml_tensor * dst = add_tensor; + + GGML_ASSERT(src0 && src0->extra); + GGML_ASSERT(src1 && src1->extra); + GGML_ASSERT(src2 && src2->extra); + GGML_ASSERT(dst && dst->extra); + + ggml_tensor_extra_cl * extra0 = (ggml_tensor_extra_cl *)src0->extra; + ggml_tensor_extra_cl * extra1 = (ggml_tensor_extra_cl *)src1->extra; + ggml_tensor_extra_cl * extra2 = (ggml_tensor_extra_cl *)src2->extra; + ggml_tensor_extra_cl * extrad = (ggml_tensor_extra_cl *)dst->extra; + + cl_ulong offset0 = extra0->offset + src0->view_offs; + cl_ulong offset1 = extra1->offset + src1->view_offs; + cl_ulong offset2 = extra2->offset + src2->view_offs; + cl_ulong offsetd = extrad->offset + dst->view_offs; + + ggml_backend_opencl_context *backend_ctx = (ggml_backend_opencl_context *)backend->context; + + float eps; + memcpy(&eps, rms_norm_tensor->op_params, sizeof(float)); + + const int ne00 = src0->ne[0], ne01 = src0->ne[1], ne02 = src0->ne[2], ne03 = src0->ne[3]; + const cl_ulong nb01 = src0->nb[1], nb02 = src0->nb[2], nb03 = src0->nb[3]; + const int ne10 = src1->ne[0], ne11 = src1->ne[1], ne12 = src1->ne[2], ne13 = src1->ne[3]; + const cl_ulong nb11 = src1->nb[1], nb12 = src1->nb[2], nb13 = src1->nb[3]; + const int ne20 = src2->ne[0], ne21 = src2->ne[1], ne22 = src2->ne[2], ne23 = src2->ne[3]; + const cl_ulong nb21 = src2->nb[1], nb22 = src2->nb[2], nb23 = src2->nb[3]; + const cl_ulong nb1 = dst->nb[1], nb2 = dst->nb[2], nb3 = dst->nb[3]; + + GGML_ASSERT(ne00 % 4 == 0); + + size_t sgs; + if (backend_ctx->gpu_family == ADRENO) sgs = 64; + else if (backend_ctx->gpu_family == INTEL) sgs = 32; + else GGML_ASSERT(false && "Unsupported GPU"); + + cl_kernel kernel = backend_ctx->kernel_rms_norm_mul_add; + + int nth = sgs; + int max_workgroup_size = backend_ctx->get_kernel_workgroup_size(kernel); + while (nth < ne00 && nth < max_workgroup_size) nth *= 2; + nth = MIN(nth, max_workgroup_size); + nth = MIN(nth, ne00); + + size_t global_work_size[] = {(size_t)ne01*nth, (size_t)ne02, (size_t)ne03}; + size_t local_work_size[] = {(size_t)nth, 1, 1}; + + CL_CHECK(clSetKernelArg(kernel, 0, sizeof(cl_mem), &extra0->data_device)); + CL_CHECK(clSetKernelArg(kernel, 1, sizeof(cl_ulong), &offset0)); + CL_CHECK(clSetKernelArg(kernel, 2, sizeof(cl_mem), &extra1->data_device)); + CL_CHECK(clSetKernelArg(kernel, 3, sizeof(cl_ulong), &offset1)); + CL_CHECK(clSetKernelArg(kernel, 4, sizeof(cl_mem), &extra2->data_device)); + CL_CHECK(clSetKernelArg(kernel, 5, sizeof(cl_ulong), &offset2)); + CL_CHECK(clSetKernelArg(kernel, 6, sizeof(cl_mem), &extrad->data_device)); + CL_CHECK(clSetKernelArg(kernel, 7, sizeof(cl_ulong), &offsetd)); + CL_CHECK(clSetKernelArg(kernel, 8, sizeof(int), &ne00)); + CL_CHECK(clSetKernelArg(kernel, 9, sizeof(int), &ne01)); + CL_CHECK(clSetKernelArg(kernel, 10, sizeof(int), &ne02)); + CL_CHECK(clSetKernelArg(kernel, 11, sizeof(int), &ne03)); + CL_CHECK(clSetKernelArg(kernel, 12, sizeof(cl_ulong), &nb01)); + CL_CHECK(clSetKernelArg(kernel, 13, sizeof(cl_ulong), &nb02)); + CL_CHECK(clSetKernelArg(kernel, 14, sizeof(cl_ulong), &nb03)); + CL_CHECK(clSetKernelArg(kernel, 15, sizeof(int), &ne10)); + CL_CHECK(clSetKernelArg(kernel, 16, sizeof(int), &ne11)); + CL_CHECK(clSetKernelArg(kernel, 17, sizeof(int), &ne12)); + CL_CHECK(clSetKernelArg(kernel, 18, sizeof(int), &ne13)); + CL_CHECK(clSetKernelArg(kernel, 19, sizeof(cl_ulong), &nb11)); + CL_CHECK(clSetKernelArg(kernel, 20, sizeof(cl_ulong), &nb12)); + CL_CHECK(clSetKernelArg(kernel, 21, sizeof(cl_ulong), &nb13)); + CL_CHECK(clSetKernelArg(kernel, 22, sizeof(int), &ne20)); + CL_CHECK(clSetKernelArg(kernel, 23, sizeof(int), &ne21)); + CL_CHECK(clSetKernelArg(kernel, 24, sizeof(int), &ne22)); + CL_CHECK(clSetKernelArg(kernel, 25, sizeof(int), &ne23)); + CL_CHECK(clSetKernelArg(kernel, 26, sizeof(cl_ulong), &nb21)); + CL_CHECK(clSetKernelArg(kernel, 27, sizeof(cl_ulong), &nb22)); + CL_CHECK(clSetKernelArg(kernel, 28, sizeof(cl_ulong), &nb23)); + CL_CHECK(clSetKernelArg(kernel, 29, sizeof(cl_ulong), &nb1)); + CL_CHECK(clSetKernelArg(kernel, 30, sizeof(cl_ulong), &nb2)); + CL_CHECK(clSetKernelArg(kernel, 31, sizeof(cl_ulong), &nb3)); + CL_CHECK(clSetKernelArg(kernel, 32, sizeof(float), &eps)); + CL_CHECK(clSetKernelArg(kernel, 33, sizeof(float)*sgs, NULL)); + + backend_ctx->enqueue_ndrange_kernel(kernel, 3, global_work_size, local_work_size, dst); +} + static ggml_status ggml_backend_opencl_graph_compute(ggml_backend_t backend, ggml_cgraph * cgraph) { ggml_backend_opencl_context *backend_ctx = (ggml_backend_opencl_context *)backend->context; @@ -7395,12 +7928,12 @@ static ggml_status ggml_backend_opencl_graph_compute(ggml_backend_t backend, ggm continue; } - if (!backend_ctx->disable_fusion && ggml_opencl_can_fuse(cgraph, i, { GGML_OP_NORM, GGML_OP_MUL, GGML_OP_ADD })) { + if (!backend_ctx->disable_fusion && ggml_opencl_can_fuse(backend_ctx, cgraph, i, { GGML_OP_NORM, GGML_OP_MUL, GGML_OP_ADD })) { ggml_opencl_op_norm_fused(backend, node, cgraph->nodes[i+1], cgraph->nodes[i+2]); i += 2; continue; } - if (!backend_ctx->disable_fusion && ggml_opencl_can_fuse(cgraph, i, { GGML_OP_GROUP_NORM, GGML_OP_MUL, GGML_OP_ADD })) { + if (!backend_ctx->disable_fusion && ggml_opencl_can_fuse(backend_ctx, cgraph, i, { GGML_OP_GROUP_NORM, GGML_OP_MUL, GGML_OP_ADD })) { ggml_opencl_op_group_norm_fused(backend, node, cgraph->nodes[i+1], cgraph->nodes[i+2]); i += 2; continue; @@ -7441,11 +7974,35 @@ static ggml_status ggml_backend_opencl_graph_compute(ggml_backend_t backend, ggm } } - if (!backend_ctx->disable_fusion && ggml_opencl_can_fuse(cgraph, i, { GGML_OP_RMS_NORM, GGML_OP_MUL })) { + // Fuse rms_norm + mul(weight) + add(residual). Checked before the + // rms_norm+mul fuse so the 3-op pattern wins over its 2-op prefix. + // Default on, opt-out GGML_OPENCL_FUSE_RMS_ADD=0. + if (!backend_ctx->disable_fusion && backend_ctx->fuse_rms_add && + ggml_opencl_can_fuse(backend_ctx, cgraph, i, { GGML_OP_RMS_NORM, GGML_OP_MUL, GGML_OP_ADD })) { + ggml_opencl_op_rms_norm_mul_add_fused(backend, node, cgraph->nodes[i+1], cgraph->nodes[i+2]); + i += 2; + continue; + } + if (!backend_ctx->disable_fusion && ggml_opencl_can_fuse(backend_ctx, cgraph, i, { GGML_OP_RMS_NORM, GGML_OP_MUL })) { ggml_opencl_op_rms_norm_fused(backend, node, cgraph->nodes[i+1]); i++; continue; } + // Fuse mul_mat(Wg,x) + mul_mat(Wu,x) + glu — fold the FFN's two decode + // GEMVs and the GLU into one dispatch. q4_K only (guarded below); the + // fused kernel uses the same accumulation/reduction order and the same + // scalar GLU formula -> coherent. Default on, opt-out GGML_OPENCL_FUSE_MM_GLU=0. +#ifdef GGML_OPENCL_USE_ADRENO_KERNELS + // The fused executor (ggml_cl_mul_mat_q4_k_glu_fused) is image-path / + // Adreno-only (GGML_ABORT on the non-Adreno #else); gate the dispatch to + // match so the FFN GLU subgraph stays dormant on Intel/other drivers. + if (backend_ctx->fuse_mm_glu && !backend_ctx->disable_fusion && + ggml_opencl_can_fuse(backend_ctx, cgraph, i, { GGML_OP_MUL_MAT, GGML_OP_MUL_MAT, GGML_OP_GLU })) { + ggml_cl_mul_mat_q4_k_glu_fused(backend, node, cgraph->nodes[i+1], cgraph->nodes[i+2]); + i += 2; + continue; + } +#endif bool ok = ggml_cl_compute_forward(backend, node); if (!ok) { @@ -7506,6 +8063,72 @@ inline bool use_adreno_moe_kernels(const ggml_backend_opencl_context *backend_ct return (((strstr(tensor->name, "ffn") != NULL) && (strstr(tensor->name, "exps") != NULL)) || (strstr(tensor->name, "as") != NULL)) && (ne01 % 32 == 0); } +// Device default for the tiled-wide lm_head/embed GEMV layout: ON for X2E and A8X. +// +// These kernels were previously off everywhere on the grounds that they compute +// wrong values at multi-superblock K. They do not: that NMSE ~2 came from the +// backend having no get_tensor restore path for the tiled layout, so +// test-backend-ops (which builds its CPU reference by copying the weights back +// out of the backend) compared a correct GPU result against a reference +// dequantized from tiled bytes. With the restore path added, MUL_MAT passes with +// the tiled kernels on, unmodified, on both devices. +// +// Perf, Qwen3-4B-Q4_K_M (q6_K lm_head 151936x2560), tg128, matched pairs with +// alternating lead, tiled vs o4: +// +// A8X +11.9% 6/6 pairs positive, order bias -0.06% (16.93 vs 15.14 tok/s) +// X2E +6.9% 4/4 pairs positive, order bias -0.03% (35.24 vs 32.87 tok/s) +// +// Measure this one on a COLD device. These kernels are far more clock-sensitive +// than the o4 route they replace: on a heat-soaked A8X (CPU cap at 1.5-1.9 GHz) +// tiled pins at ~14.2 tok/s while o4 still makes ~14.9, which reads as a 4-5% +// LOSS and inverts the ranking. The same box, after a reboot and a gate that +// waits for policy6 to return to 4396800, reports the +11.9% above with no +// order bias. A7X regresses hard on this layout and stays off. +// GGML_OPENCL_{Q4K,Q6K}_GEMV_TILED forces either way (=0 off, any other value on). +inline bool tiled_gemv_default_on(const ggml_backend_opencl_context *backend_ctx) { + return backend_ctx && (backend_ctx->adreno_gen == ADRENO_GPU_GEN::X2E || + backend_ctx->adreno_gen == ADRENO_GPU_GEN::A8X); +} + +// Tiled-wide q6_K GEMV (default OFF; GGML_OPENCL_Q6K_GEMV_TILED forces either +// way: =0 off everywhere, any other value on everywhere). +// Both the convert (set_tensor) and the GEMV dispatch must agree on this so the +// buffer layout matches the kernel. +inline bool q6k_gemv_tiled_enabled(const ggml_backend_opencl_context *backend_ctx) { + static const char * e = std::getenv("GGML_OPENCL_Q6K_GEMV_TILED"); + if (e && e[0] != '\0') { + return e[0] != '0'; + } + return tiled_gemv_default_on(backend_ctx); +} + +// Only the long-vocab lm_head/embed shapes use the tiled layout; ne01 % 64 == 0 +// is required by the 64-row tiling (no row padding in the buffers). +// use_adreno_kernels is required: only the Adreno GEMV path can read the tiled +// layout, so converting a weight it would decline (e.g. ne00 < 512) leaves the +// generic kernel reading tiled bytes as plain SOA. +inline bool use_q6k_tiled(const ggml_backend_opencl_context *backend_ctx, const ggml_tensor *tensor) { + return q6k_gemv_tiled_enabled(backend_ctx) && tensor->type == GGML_TYPE_Q6_K && + tensor->ne[1] >= 32768 && tensor->ne[1] % 64 == 0 && + use_adreno_kernels(backend_ctx, tensor); +} + +// q4_K analog of the tiled-wide lm_head/embed GEMV (default OFF; +// GGML_OPENCL_Q4K_GEMV_TILED forces either way: =0 off, else on). Same gate. +inline bool q4k_gemv_tiled_enabled(const ggml_backend_opencl_context *backend_ctx) { + static const char * e = std::getenv("GGML_OPENCL_Q4K_GEMV_TILED"); + if (e && e[0] != '\0') { + return e[0] != '0'; + } + return tiled_gemv_default_on(backend_ctx); +} +inline bool use_q4k_tiled(const ggml_backend_opencl_context *backend_ctx, const ggml_tensor *tensor) { + return q4k_gemv_tiled_enabled(backend_ctx) && tensor->type == GGML_TYPE_Q4_K && + tensor->ne[1] >= 32768 && tensor->ne[1] % 64 == 0 && + use_adreno_kernels(backend_ctx, tensor); +} + inline bool enable_adreno_trans_weight(const ggml_backend_opencl_context *backend_ctx, const ggml_tensor *tensor) { bool adreno_kernel = use_adreno_kernels(backend_ctx, tensor); @@ -7533,18 +8156,53 @@ inline bool enable_adreno_trans_weight_q5_K(const ggml_backend_opencl_context *b qh_img_width <= backend_ctx->image_max_buffer_size; } -static inline bool use_flat_gemv_for_large_m_q4_K(const ggml_tensor *tensor) { +// The flat-GEMV large-m escape is OPT-IN (GGML_OPENCL_FLAT_LARGE_M=1) because it +// is SLOWER than the route it replaces, not because it is unsafe. It was first +// parked on the theory that it out-of-bounds-writes at vocab-scale shapes; that +// was a misattribution (the test-backend-ops dst sentinel was tripped by the o4 +// GEMV's unguarded tail store, fixed separately - and at the shape it was blamed +// for, k=1536, this predicate returns false anyway, so the flat route never ran). +// +// The escape's original rationale, "gemv_noshuffle perf drops for large M", +// predates the o4 kernel, which now covers the same long-vocab shapes and beats +// this route on every device measured (Qwen3-4B-Q4_K_M, q6_K lm_head +// 151936x2560, tg128, matched pairs vs o4): A8X -10.3% (0/3 pairs), X2E -3.7% +// (0/3). Keep it reachable for shapes o4 declines, but do not default it on. +static inline bool flat_large_m_enabled() { + static const char * e = getenv("GGML_OPENCL_FLAT_LARGE_M"); + static const bool en = e != nullptr && atoi(e) != 0; + return en; +} + +static inline bool use_flat_gemv_for_large_m_q4_K(const ggml_backend_opencl_context *backend_ctx, const ggml_tensor *tensor) { + if (!flat_large_m_enabled()) { + return false; + } // gemv_noshuffle variant perf drops for large M, use flat variant for large M. // threshold is well above typical hidden/FFN dims, but below typical vocab sizes. // note that this forces large M weights to use LM GEMM. - return tensor->ne[1] >= 32768 && tensor->ne[2] == 1 && tensor->ne[3] == 1; + // EXCEPT when this branch's tiled-canonical lm_head/embed layout is active: the + // weight is converted to the 64-row tiled layout, which the flat gemv would + // misread as garbage. use_q4k_tiled owns these large-M weights, so defer to it. + return tensor->ne[1] >= 32768 && tensor->ne[2] == 1 && tensor->ne[3] == 1 + && !use_q4k_tiled(backend_ctx, tensor); } static inline bool use_flat_gemv_for_large_m_q6_K(const ggml_backend_opencl_context *backend_ctx, const ggml_tensor *tensor) { + // NOTE on ordering: the ne01 % 128 escape below is a CORRECTNESS guard, not a + // performance one, so it must be reachable regardless of flat_large_m_enabled(). + // The opt-in gate therefore sits after it, and after the tiled deferral. // gemv_noshuffle variant perf drops for large M, use flat variant for large M. // threshold is well above typical hidden/FFN dims, but below typical vocab sizes. // q6_K flat gemv is worse for smaller K; 2048 seems to be a reasonable threshold. // note that this forces large M weights to use LM GEMM. + // When this branch's tiled-canonical lm_head/embed layout is active, the weight is + // converted to the 64-row tiled layout, which the flat gemv would misread as + // garbage. use_q6k_tiled owns these large-M weights (it requires ne01 % 64 == 0, + // so it never claims an odd-vocab weight), so defer to it first. + if (use_q6k_tiled(backend_ctx, tensor)) { + return false; + } // The noshuffle (transposed-weight) layout packs 2 rows per 32-bit texel and the // gemv reads it with a ne01/2 texel stride and an exact-cover dispatch of // ceil(ne01/2 / 64)*64 work-items with no store guard; the gemm uses 4-row tiles. @@ -7559,6 +8217,10 @@ static inline bool use_flat_gemv_for_large_m_q6_K(const ggml_backend_opencl_cont return true; } + if (!flat_large_m_enabled()) { + return false; + } + // The gemv_noshuffle slowdown tracks TOTAL weight size, not ne0 alone; ne0 >= 2048 is a // proxy for "large weight" that misses a narrow-hidden vocab-scale lm_head. // Add a direct size escape so such weights also take the flat path, without changing @@ -7805,8 +8467,29 @@ static bool ggml_opencl_supports_op(ggml_backend_dev_t dev, const struct ggml_te op->src[0]->ne[1] >= 32768) { // vocab-scale weight; no FFN/attn weight is this tall return false; } + // The generic mul_mv (GEMV) kernels are wrong for large-batch prefill on + // Adreno. A quant mul_mat only avoids the GEMV when it reaches the Adreno + // trans-weight GEMM, which needs both a GEMM kernel for the type and + // use_adreno_kernels(). Decline the large-N shapes that would otherwise + // fall through to the GEMV. + { + const ggml_type t = op->src[0]->type; + const bool type_has_gemm = (t == GGML_TYPE_Q4_0 || t == GGML_TYPE_Q4_1 || + t == GGML_TYPE_IQ4_NL || t == GGML_TYPE_Q8_0 || + t == GGML_TYPE_Q4_K || t == GGML_TYPE_Q5_K || + t == GGML_TYPE_Q6_K); + const bool uses_gemm = type_has_gemm && use_adreno_kernels(backend_ctx, op->src[0]); + if (!uses_gemm && op->src[1]->ne[1] >= 512) { + return false; + } + } return op->src[1]->type == GGML_TYPE_F32 && ggml_is_contiguous(op->src[0]) && ggml_is_contiguous(op->src[1]); } else if (op->src[0]->type == GGML_TYPE_Q8_0) { + // ggml_cl_mul_mat_q8_0_f32_adreno now honors src1/dst view_offs (the + // activation sub-buffer starts at offset1 and the kernels take offsetd), + // so a broadcast q8_0 matmul (src1 batch > src0 batch, e.g. Qwen3.5-9B-UD + // / Qwen3.6-35B q8_0 GDN ssm_out) runs on GPU via the per-slice broadcast + // iteration in ggml_cl_mul_mat. No special-casing needed. return op->src[1]->type == GGML_TYPE_F32; } return false; @@ -9574,8 +10257,41 @@ static void ggml_backend_opencl_buffer_set_tensor(ggml_backend_buffer_t buffer, #endif // GGML_OPENCL_USE_ADRENO_KERNELS #ifdef GGML_OPENCL_USE_ADRENO_KERNELS + // Tiled-wide convert for the long-vocab lm_head/embed (opt-in). The embed/ + // output q4_K weight (token_embd.weight, ne1=vocab) is NOT matched by + // use_adreno_moe_kernels, so it lands here in the general branch. Produce + // the final 64-row-tiled canonical layout directly into q/d/dm/s (buffer + // sizes already match), read back by kernel_gemv_noshuffle_q4_k_f32_tiled. + if (use_q4k_tiled(backend_ctx, tensor)) { + cl_kernel tk = backend_ctx->kernel_convert_block_q4_k_tiled_ns; + + int ne00 = tensor->ne[0]; + int ne01 = tensor->ne[1]; + int ne02 = tensor->ne[2]; + + CL_CHECK(clSetKernelArg(tk, 0, sizeof(cl_mem), &data_device)); + CL_CHECK(clSetKernelArg(tk, 1, sizeof(cl_mem), &extra->q)); + CL_CHECK(clSetKernelArg(tk, 2, sizeof(cl_mem), &extra->d)); + CL_CHECK(clSetKernelArg(tk, 3, sizeof(cl_mem), &extra->dm)); + CL_CHECK(clSetKernelArg(tk, 4, sizeof(cl_mem), &extra->s)); + CL_CHECK(clSetKernelArg(tk, 5, sizeof(int), &ne00)); + CL_CHECK(clSetKernelArg(tk, 6, sizeof(int), &ne01)); + + size_t gws[] = {static_cast(((ne01 + 63) / 64) * 64), static_cast(ne00 / 256), static_cast(ne02)}; + size_t lws[] = {64, 1, 1}; + + cl_event tevt; + CL_CHECK(clEnqueueNDRangeKernel(queue, tk, 3, NULL, gws, lws, 0, NULL, &tevt)); + CL_CHECK(clWaitForEvents(1, &tevt)); + CL_CHECK(clReleaseMemObject(data_device)); + + extra->q_img = nullptr; + tensor->extra = extra; + return; + } + cl_kernel kernel = backend_ctx->kernel_convert_block_q4_K; - if (use_adreno_kernels(backend_ctx, tensor) && !use_flat_gemv_for_large_m_q4_K(tensor)) { + if (use_adreno_kernels(backend_ctx, tensor) && !use_flat_gemv_for_large_m_q4_K(backend_ctx, tensor)) { kernel = backend_ctx->kernel_convert_block_q4_K_noshuffle; } #else @@ -9603,7 +10319,7 @@ static void ggml_backend_opencl_buffer_set_tensor(ggml_backend_buffer_t buffer, tensor->extra = extra; #ifdef GGML_OPENCL_USE_ADRENO_KERNELS - if (use_adreno_kernels(backend_ctx, tensor) && !use_flat_gemv_for_large_m_q4_K(tensor)) { + if (use_adreno_kernels(backend_ctx, tensor) && !use_flat_gemv_for_large_m_q4_K(backend_ctx, tensor)) { int M = tensor->ne[1]; int K = tensor->ne[0]; @@ -9929,6 +10645,45 @@ static void ggml_backend_opencl_buffer_set_tensor(ggml_backend_buffer_t buffer, CL_CHECK((extra->d = clCreateSubBuffer(extra_orig->data_device, CL_MEM_READ_WRITE, CL_BUFFER_CREATE_TYPE_REGION, ®ion, &err), err)); previous_origin = region.origin; +#ifdef GGML_OPENCL_USE_ADRENO_KERNELS + // Tiled-wide convert for the long-vocab lm_head/embed (opt-in). The embed + // /output q6_K weight (e.g. token_embd.weight, ne1=vocab) is NOT matched by + // use_adreno_moe_kernels, so it lands here in the general branch. Produce + // the final 64-row-tiled canonical layout directly into ql/qh/s/d (buffer + // sizes already match), read back by kernel_gemv_noshuffle_q6_K_f32_tiled. + // Bypasses the plain-SOA convert + per-array transpose below. + if (use_q6k_tiled(backend_ctx, tensor)) { + cl_kernel kernel = backend_ctx->kernel_convert_block_q6_k_tiled_ns; + + int ne00 = tensor->ne[0]; + int ne01 = tensor->ne[1]; + int ne02 = tensor->ne[2]; + + CL_CHECK(clSetKernelArg(kernel, 0, sizeof(cl_mem), &data_device)); + CL_CHECK(clSetKernelArg(kernel, 1, sizeof(cl_mem), &extra->ql)); + CL_CHECK(clSetKernelArg(kernel, 2, sizeof(cl_mem), &extra->qh)); + CL_CHECK(clSetKernelArg(kernel, 3, sizeof(cl_mem), &extra->d)); + CL_CHECK(clSetKernelArg(kernel, 4, sizeof(cl_mem), &extra->s)); + CL_CHECK(clSetKernelArg(kernel, 5, sizeof(int), &ne00)); + CL_CHECK(clSetKernelArg(kernel, 6, sizeof(int), &ne01)); + + size_t global_work_size[] = {static_cast(((ne01 + 63) / 64) * 64), static_cast(ne00 / 256), static_cast(ne02)}; + size_t local_work_size[] = {64, 1, 1}; + + cl_event evt; + CL_CHECK(clEnqueueNDRangeKernel(queue, kernel, 3, NULL, global_work_size, local_work_size, 0, NULL, &evt)); + CL_CHECK(clWaitForEvents(1, &evt)); + CL_CHECK(clReleaseMemObject(data_device)); + + extra->size_ql = size_ql; + extra->size_qh = size_qh; + extra->size_s = size_s; + extra->size_d = size_d; + tensor->extra = extra; + return; + } +#endif // GGML_OPENCL_USE_ADRENO_KERNELS + // Flatten the weights cl_kernel kernel; #ifdef GGML_OPENCL_USE_ADRENO_KERNELS @@ -10744,6 +11499,54 @@ static void ggml_backend_opencl_buffer_get_tensor(ggml_backend_buffer_t buffer, cl_uchar mask_F0 = 0xF0; #ifdef GGML_OPENCL_USE_ADRENO_KERNELS + // Undo the 64-row-tiled canonical pack (kernel_convert_block_q4_k_tiled_ns). + // Without this, a read-back of a tiled weight returns the tiled bytes + // reinterpreted as block_q4_K -- which is how test-backend-ops builds its + // CPU reference (ggml_backend_graph_copy -> tensor_get), so the tiled path + // "failed" the suite while computing the correct product. + if (use_q4k_tiled(backend_ctx, tensor)) { + const int ne00v = tensor->ne[0]; + const int ne01v = tensor->ne[1]; + const int nbv = ne00v / 256; + const size_t n_blk = (size_t)nbv * ne01v; + + std::vector tq(n_blk*32); + std::vector td(n_blk), tdm(n_blk); + std::vector ts(n_blk*12); + CL_CHECK(clEnqueueReadBuffer(queue, extra->q, CL_TRUE, 0, tq.size()*4, tq.data(), 0, NULL, NULL)); + CL_CHECK(clEnqueueReadBuffer(queue, extra->d, CL_TRUE, 0, td.size()*2, td.data(), 0, NULL, NULL)); + CL_CHECK(clEnqueueReadBuffer(queue, extra->dm, CL_TRUE, 0, tdm.size()*2, tdm.data(), 0, NULL, NULL)); + CL_CHECK(clEnqueueReadBuffer(queue, extra->s, CL_TRUE, 0, ts.size(), ts.data(), 0, NULL, NULL)); + + std::vector rebuilt(ggml_nbytes(tensor), 0); + for (int i01 = 0; i01 < ne01v; ++i01) { + const int rt = i01/64, rit = i01%64; + for (int i00 = 0; i00 < nbv; ++i00) { + uint8_t * b = rebuilt.data() + ((size_t)i00 + (size_t)i01*nbv)*144; + const int tb = rt*nbv + i00; + const size_t si = (size_t)tb*64 + rit; + + memcpy(b + 0, &td [si], 2); + memcpy(b + 2, &tdm[si], 2); + memcpy(b + 4, &ts[si*12], 12); + + uint32_t qw[32]; + for (int gr = 0; gr < 8; ++gr) { + const size_t base = ((size_t)tb*8 + gr)*64 + rit; + for (int j = 0; j < 4; ++j) qw[gr*4 + j] = tq[base*4 + j]; + } + uint8_t * q = b + 16; + for (int e = 0; e < 256; ++e) { + const int g = e>>6, w = e&63, h = w>>5, l = w&31; + const uint32_t code = (qw[e>>3] >> ((e&7)*4)) & 0xF; + q[g*32 + l] |= (uint8_t)(h ? (code << 4) : code); + } + } + } + memcpy(data, rebuilt.data() + offset, size); + CL_CHECK(clReleaseMemObject(data_device)); + return; + } if (use_adreno_moe_kernels(backend_ctx, tensor)) { cl_int err; cl_mem data_device = clCreateBuffer(context, CL_MEM_READ_WRITE, @@ -10778,7 +11581,7 @@ static void ggml_backend_opencl_buffer_get_tensor(ggml_backend_buffer_t buffer, CL_CHECK(clReleaseMemObject(data_device)); return; } - if (use_adreno_kernels(backend_ctx, tensor) && !use_flat_gemv_for_large_m_q4_K(tensor)) { + if (use_adreno_kernels(backend_ctx, tensor) && !use_flat_gemv_for_large_m_q4_K(backend_ctx, tensor)) { int M = tensor->ne[1]; int K = tensor->ne[0]; @@ -10966,6 +11769,61 @@ static void ggml_backend_opencl_buffer_get_tensor(ggml_backend_buffer_t buffer, ggml_tensor_extra_cl_q6_K * extra = (ggml_tensor_extra_cl_q6_K *)tensor->extra; #ifdef GGML_OPENCL_USE_ADRENO_KERNELS + // Undo the 64-row-tiled canonical pack (kernel_convert_block_q6_k_tiled_ns). + // See the q4_K tiled restore above for why a read-back path is required. + if (use_q6k_tiled(backend_ctx, tensor)) { + const int ne00v = tensor->ne[0]; + const int ne01v = tensor->ne[1]; + const int nbv = ne00v / 256; + const size_t n_blk = (size_t)nbv * ne01v; + + std::vector tql(n_blk*32), tqh(n_blk*16); + std::vector ts(n_blk*16); + std::vector td(n_blk); + CL_CHECK(clEnqueueReadBuffer(queue, extra->ql, CL_TRUE, 0, tql.size()*4, tql.data(), 0, NULL, NULL)); + CL_CHECK(clEnqueueReadBuffer(queue, extra->qh, CL_TRUE, 0, tqh.size()*4, tqh.data(), 0, NULL, NULL)); + CL_CHECK(clEnqueueReadBuffer(queue, extra->s, CL_TRUE, 0, ts.size(), ts.data(), 0, NULL, NULL)); + CL_CHECK(clEnqueueReadBuffer(queue, extra->d, CL_TRUE, 0, td.size()*2, td.data(), 0, NULL, NULL)); + + std::vector rebuilt(ggml_nbytes(tensor), 0); + for (int i01 = 0; i01 < ne01v; ++i01) { + const int rt = i01/64, rit = i01%64; + for (int i00 = 0; i00 < nbv; ++i00) { + uint8_t * b = rebuilt.data() + ((size_t)i00 + (size_t)i01*nbv)*210; + const int tb = rt*nbv + i00; + const size_t si = (size_t)tb*64 + rit; + + uint32_t qlw[32], qhw[16]; + for (int g = 0; g < 8; ++g) { + const size_t base = ((size_t)tb*8 + g)*64 + rit; + for (int j = 0; j < 4; ++j) qlw[g*4 + j] = tql[base*4 + j]; + } + for (int g = 0; g < 4; ++g) { + const size_t base = ((size_t)tb*4 + g)*64 + rit; + for (int j = 0; j < 4; ++j) qhw[g*4 + j] = tqh[base*4 + j]; + } + + uint8_t * ql = b; + uint8_t * qh = b + 128; + for (int e = 0; e < 256; ++e) { + const int n = (e >= 128) ? 1 : 0; + const int within = e - n*128, q = within/32, l = within%32; + const int off_ql = n*64, off_qh = n*32; + const uint8_t low4 = (qlw[e>>3] >> ((e&7)*4)) & 0xF; + const uint8_t hi2 = (qhw[e>>4] >> ((e&15)*2)) & 0x3; + if (q == 0) ql[off_ql + l] |= low4; + else if (q == 1) ql[off_ql + l + 32] |= low4; + else if (q == 2) ql[off_ql + l] |= (uint8_t)(low4 << 4); + else ql[off_ql + l + 32] |= (uint8_t)(low4 << 4); + qh[off_qh + l] |= (uint8_t)(hi2 << (q*2)); + } + memcpy(b + 192, &ts[si*16], 16); + memcpy(b + 208, &td[si], 2); + } + } + memcpy(data, rebuilt.data() + offset, size); + return; + } if (use_adreno_moe_kernels(backend_ctx, tensor)) { cl_int err; cl_mem data_device = clCreateBuffer(context, CL_MEM_READ_WRITE, @@ -16634,7 +17492,17 @@ static void ggml_cl_mul_mat_q4_0_f32_adreno(ggml_backend_t backend, const ggml_t int N = ne1; int K = ne00; - if (ne1 == 1) { + // Multi-column (N=3) verify GEMV for q4_0: route the spec/MTP verify batch + // (ne1==3) onto the efficient GEMV path instead of the transposed-GEMM dead- + // zone (gemm_noshuffle_q4_0 is ~50% of MTP decode on a Q4_0 model since q4_0 + // weights have no cok/mc3, unlike q4_K). Reuses the ne1==1 GEMV image setup + // (activation image already sized by N=ne1). Byte-identical. Opt-in via + // GGML_OPENCL_Q40_MC3=1. Per-layer only (ne01 < 32768); q4_0 lm_head doesn't + // occur (token_embd/output stay Q6_K), guard kept for parity with q4_K mc3. + static const bool q40_mc3 = (getenv("GGML_OPENCL_Q40_MC3") != nullptr); + const bool use_q40_mc3 = q40_mc3 && (ne1 >= 2 && ne1 <= 4) && (ne01 < 32768); + + if (ne1 == 1 || use_q40_mc3) { cl_mem q_img = nullptr; cl_mem b_sub_buf = nullptr; cl_mem b_img = nullptr; @@ -16660,38 +17528,56 @@ static void ggml_cl_mul_mat_q4_0_f32_adreno(ggml_backend_t backend, const ggml_t img_desc.buffer = b_sub_buf; CL_CHECK((b_img = clCreateImage(context, CL_MEM_READ_ONLY, &img_fmt, &img_desc, NULL, &err), err)); - kernel = backend_ctx->kernel_gemv_noshuffle_q4_0_f32; - if (M == 4096 && K == 4096) { - kernel = backend_ctx->kernel_gemv_noshuffle_q4_0_f32_4096_1_4096; - } else if (M == 4096 && K == 11008) { - kernel = backend_ctx->kernel_gemv_noshuffle_q4_0_f32_4096_1_11008; - } else if (M == 11008 && K == 4096) { - kernel = backend_ctx->kernel_gemv_noshuffle_q4_0_f32_11008_1_4096; - } else if (M == 32000 && K == 4096) { - kernel = backend_ctx->kernel_gemv_noshuffle_q4_0_f32_32000_1_4096; - } + if (use_q40_mc3) { + kernel = backend_ctx->kernel_gemv_noshuffle_q4_0_f32_mc3; + CL_CHECK(clSetKernelArg(kernel, 0, sizeof(cl_mem), &q_img)); + CL_CHECK(clSetKernelArg(kernel, 1, sizeof(cl_mem), &extra0_q4_0->d)); + CL_CHECK(clSetKernelArg(kernel, 2, sizeof(cl_mem), &b_img)); + CL_CHECK(clSetKernelArg(kernel, 3, sizeof(cl_mem), &extrad->data_device)); + CL_CHECK(clSetKernelArg(kernel, 4, sizeof(cl_ulong), &offsetd)); + CL_CHECK(clSetKernelArg(kernel, 5, sizeof(int), &ne00)); + CL_CHECK(clSetKernelArg(kernel, 6, sizeof(int), &ne01)); + CL_CHECK(clSetKernelArg(kernel, 7, sizeof(int), &ne1)); + } else { + kernel = backend_ctx->kernel_gemv_noshuffle_q4_0_f32; + if (M == 4096 && K == 4096) { + kernel = backend_ctx->kernel_gemv_noshuffle_q4_0_f32_4096_1_4096; + } else if (M == 4096 && K == 11008) { + kernel = backend_ctx->kernel_gemv_noshuffle_q4_0_f32_4096_1_11008; + } else if (M == 11008 && K == 4096) { + kernel = backend_ctx->kernel_gemv_noshuffle_q4_0_f32_11008_1_4096; + } else if (M == 32000 && K == 4096) { + kernel = backend_ctx->kernel_gemv_noshuffle_q4_0_f32_32000_1_4096; + } - int r2 = 1; - int r3 = 1; + int r2 = 1; + int r3 = 1; - CL_CHECK(clSetKernelArg(kernel, 0, sizeof(cl_mem), &q_img)); - CL_CHECK(clSetKernelArg(kernel, 1, sizeof(cl_mem), &extra0_q4_0->d)); - CL_CHECK(clSetKernelArg(kernel, 2, sizeof(cl_mem), &b_img)); - CL_CHECK(clSetKernelArg(kernel, 3, sizeof(cl_ulong), &offset1)); - CL_CHECK(clSetKernelArg(kernel, 4, sizeof(cl_mem), &extrad->data_device)); - CL_CHECK(clSetKernelArg(kernel, 5, sizeof(cl_ulong), &offsetd)); - CL_CHECK(clSetKernelArg(kernel, 6, sizeof(int), &ne00)); - CL_CHECK(clSetKernelArg(kernel, 7, sizeof(int), &ne01)); - CL_CHECK(clSetKernelArg(kernel, 8, sizeof(int), &ne02)); - CL_CHECK(clSetKernelArg(kernel, 9, sizeof(int), &ne10)); - CL_CHECK(clSetKernelArg(kernel, 10, sizeof(int), &ne12)); - CL_CHECK(clSetKernelArg(kernel, 11, sizeof(int), &ne0)); - CL_CHECK(clSetKernelArg(kernel, 12, sizeof(int), &ne1)); - CL_CHECK(clSetKernelArg(kernel, 13, sizeof(int), &r2)); - CL_CHECK(clSetKernelArg(kernel, 14, sizeof(int), &r3)); + CL_CHECK(clSetKernelArg(kernel, 0, sizeof(cl_mem), &q_img)); + CL_CHECK(clSetKernelArg(kernel, 1, sizeof(cl_mem), &extra0_q4_0->d)); + CL_CHECK(clSetKernelArg(kernel, 2, sizeof(cl_mem), &b_img)); + CL_CHECK(clSetKernelArg(kernel, 3, sizeof(cl_ulong), &offset1)); + CL_CHECK(clSetKernelArg(kernel, 4, sizeof(cl_mem), &extrad->data_device)); + CL_CHECK(clSetKernelArg(kernel, 5, sizeof(cl_ulong), &offsetd)); + CL_CHECK(clSetKernelArg(kernel, 6, sizeof(int), &ne00)); + CL_CHECK(clSetKernelArg(kernel, 7, sizeof(int), &ne01)); + CL_CHECK(clSetKernelArg(kernel, 8, sizeof(int), &ne02)); + CL_CHECK(clSetKernelArg(kernel, 9, sizeof(int), &ne10)); + CL_CHECK(clSetKernelArg(kernel, 10, sizeof(int), &ne12)); + CL_CHECK(clSetKernelArg(kernel, 11, sizeof(int), &ne0)); + CL_CHECK(clSetKernelArg(kernel, 12, sizeof(int), &ne1)); + CL_CHECK(clSetKernelArg(kernel, 13, sizeof(int), &r2)); + CL_CHECK(clSetKernelArg(kernel, 14, sizeof(int), &r3)); + } - size_t local_work_size[3] = {64, 4, 1}; - size_t global_work_size[3] = {(size_t)CEIL_DIV(ne01/2, 64)*64, 4, 1}; + // Small-M mc3 verify is occupancy/latency-bound (too few WGs at small M, so + // its bandwidth falls well short of the FFN matmuls'). Use 8 subgroups (512-WI WGs, half the + // per-lane K-walk) for small M. Layout stride is fixed (4 uints/block), so only + // the K-split count changes; the mc3 kernel reads it via get_local_size(1). The + // ne1==1 base kernel hardcodes N_SIMDGROUP=4, so it always stays at 4. + const int mc3_nsg = (use_q40_mc3 && ne01 < 4096) ? 8 : 4; + size_t local_work_size[3] = {64, (size_t)mc3_nsg, 1}; + size_t global_work_size[3] = {(size_t)CEIL_DIV(ne01/2, 64)*64, (size_t)mc3_nsg, 1}; backend_ctx->enqueue_ndrange_kernel(kernel, 3, global_work_size, local_work_size, dst); @@ -16909,7 +17795,14 @@ static void ggml_cl_mul_mat_q4_1_f32_adreno(ggml_backend_t backend, const ggml_t int N = ne1; int K = ne00; - if (ne1 == 1) { + // Multi-column (N=3) verify GEMV for q4_1: route the spec/MTP verify batch + // (ne1==3) onto the efficient GEMV path instead of the transposed-GEMM dead- + // zone (gemm_noshuffle_q4_1). Reuses the ne1==1 GEMV image setup. Opt-in via + // GGML_OPENCL_Q41_MC3=1. Per-layer only (ne01 < 32768). + static const bool q41_mc3 = (getenv("GGML_OPENCL_Q41_MC3") != nullptr); + const bool use_q41_mc3 = q41_mc3 && (ne1 >= 2 && ne1 <= 4) && (ne01 < 32768); + + if (ne1 == 1 || use_q41_mc3) { cl_mem q_img = nullptr; cl_mem b_sub_buf = nullptr; cl_mem b_img = nullptr; @@ -16935,7 +17828,8 @@ static void ggml_cl_mul_mat_q4_1_f32_adreno(ggml_backend_t backend, const ggml_t img_desc.buffer = b_sub_buf; CL_CHECK((b_img = clCreateImage(context, CL_MEM_READ_ONLY, &img_fmt, &img_desc, NULL, &err), err)); - kernel = backend_ctx->kernel_gemv_noshuffle_q4_1_f32; + kernel = use_q41_mc3 ? backend_ctx->kernel_gemv_noshuffle_q4_1_f32_mc3 + : backend_ctx->kernel_gemv_noshuffle_q4_1_f32; CL_CHECK(clSetKernelArg(kernel, 0, sizeof(cl_mem), &q_img)); CL_CHECK(clSetKernelArg(kernel, 1, sizeof(cl_mem), &extra0_q4_1->d)); @@ -16945,6 +17839,9 @@ static void ggml_cl_mul_mat_q4_1_f32_adreno(ggml_backend_t backend, const ggml_t CL_CHECK(clSetKernelArg(kernel, 5, sizeof(cl_ulong), &offsetd)); CL_CHECK(clSetKernelArg(kernel, 6, sizeof(cl_int), &ne00)); CL_CHECK(clSetKernelArg(kernel, 7, sizeof(cl_int), &ne01)); + if (use_q41_mc3) { + CL_CHECK(clSetKernelArg(kernel, 8, sizeof(cl_int), &ne1)); // n_cols + } size_t local_work_size[3] = {64, 4, 1}; size_t global_work_size[3] = {(size_t)CEIL_DIV(ne01/2, 64)*64, 4, 1}; @@ -17788,6 +18685,66 @@ static void ggml_cl_mul_mat_q8_0_f32_adreno(ggml_backend_t backend, const ggml_t img_desc.buffer = b_sub_buf; CL_CHECK((b_img = clCreateImage(context, CL_MEM_READ_ONLY, &img_fmt, &img_desc, NULL, &err), err)); + // Split-K for small-M decode GEMVs. The base kernel puts one output row + // per lane and splits K only inside one workgroup, so M is the sole source + // of workgroup parallelism: gpt-oss's K and V projections are M=512 = 8 + // workgroups on a 16-CU X2, and the kernel measures 48 GB/s where the + // M=2880/4096 projections in the same decode graph reach 122-123. Mirrors + // the q4_0/q4_K split-K above and reuses their reduce kernel. + // + // Enabled where it is measured to win, like the q4_K gate: X2-90 +2.8% + // tg32 @d4096 on gpt-oss; Adreno 840 (12 CU) NEUTRAL on Llama-3.2-3B-Q8_0 + // (0.0% @d4096 -- its K/V proj is M=1024 = 16 workgroups, which already + // fills 12 CUs). Unmeasured on X1E/A7X/A6X and the q4_K split-K measured + // -0.7% on X1E, so the default is not widened on absence of evidence. + static const bool q8_splitk_env_set = []{ + const char * e = std::getenv("GGML_OPENCL_Q8_GEMV_SPLITK"); + return e && e[0] != '\0'; + }(); + static const bool q8_splitk_env_on = []{ + const char * e = std::getenv("GGML_OPENCL_Q8_GEMV_SPLITK"); + return !(e && e[0] == '0'); + }(); + const bool q8_splitk_on = q8_splitk_env_set + ? q8_splitk_env_on + : (backend_ctx->adreno_gen == ADRENO_GPU_GEN::X2E); + if (q8_splitk_on && backend_ctx->kernel_gemv_noshuffle_q8_0_f32_splitk && + ne01 <= 1024 && ne01 % 64 == 0) { + const int nsg = 8; + const int ksplit = 8; // -> 8 * M/64 workgroups + const size_t gx = (size_t) CEIL_DIV(ne01, 64) * 64; + + backend_ctx->prealloc_splitk_partial.allocate( + backend_ctx->context, (size_t) ksplit * ne01 * sizeof(float)); + cl_mem partial = backend_ctx->prealloc_splitk_partial.buffer; + + cl_kernel ks = backend_ctx->kernel_gemv_noshuffle_q8_0_f32_splitk; + CL_CHECK(clSetKernelArg(ks, 0, sizeof(cl_mem), &q_img)); + CL_CHECK(clSetKernelArg(ks, 1, sizeof(cl_mem), &extra0_q8_0->d)); + CL_CHECK(clSetKernelArg(ks, 2, sizeof(cl_mem), &b_img)); + CL_CHECK(clSetKernelArg(ks, 3, sizeof(cl_mem), &partial)); + CL_CHECK(clSetKernelArg(ks, 4, sizeof(cl_int), &ne00)); + CL_CHECK(clSetKernelArg(ks, 5, sizeof(cl_int), &ne01)); + size_t lsk[3] = { 64, (size_t) nsg, 1 }; + size_t gsk[3] = { gx, (size_t) (nsg * ksplit), 1 }; + backend_ctx->enqueue_ndrange_kernel(ks, 3, gsk, lsk, dst); + + cl_kernel kr = backend_ctx->kernel_gemv_splitk_reduce_f32; + CL_CHECK(clSetKernelArg(kr, 0, sizeof(cl_mem), &partial)); + CL_CHECK(clSetKernelArg(kr, 1, sizeof(cl_mem), &extrad->data_device)); + CL_CHECK(clSetKernelArg(kr, 2, sizeof(cl_ulong), &offsetd)); + CL_CHECK(clSetKernelArg(kr, 3, sizeof(cl_int), &ne01)); + CL_CHECK(clSetKernelArg(kr, 4, sizeof(cl_int), &ksplit)); + size_t lr[3] = { 64, 1, 1 }; + size_t gr[3] = { (size_t) CEIL_DIV(ne01, 64) * 64, 1, 1 }; + backend_ctx->enqueue_ndrange_kernel(kr, 3, gr, lr, dst); + + CL_CHECK(clReleaseMemObject(q_img)); + CL_CHECK(clReleaseMemObject(b_img)); + CL_CHECK(clReleaseMemObject(b_sub_buf)); + return; + } + kernel = backend_ctx->kernel_gemv_noshuffle_q8_0_f32; int r2 = 1; @@ -18129,18 +19086,33 @@ static void ggml_cl_mul_mat_q4_k_f32_adreno(ggml_backend_t backend, const ggml_t cl_uchar mask_d4 = 0x0F; cl_uchar mask_hi2 = 0xC0; - if (ne1 == 1) { + // Multi-column verify GEMV: route the spec/MTP verify batch (ne1==3 = 2 + // drafts + 1 bonus) onto the efficient GEMV path (subgroup-broadcast, no + // transpose) instead of the transposed-GEMM dead-zone. Reuses the ne1==1 + // GEMV setup (the activation image is already sized by N=ne1). Byte- + // identical. Opt-in via GGML_OPENCL_Q4K_MC3=1 while validating. + static const bool q4k_mc3 = (getenv("GGML_OPENCL_Q4K_MC3") != nullptr); + // Per-layer only (ne01 < 32768): the batched large-vocab lm_head at ne1==3 + // is left to the existing routing (corrupts on the Adreno GEMV path; x2- + // unified routes batched Q6_K lm_head to CPU). Per-layer mc3 is byte-identical. + const bool use_mc3 = q4k_mc3 && (ne1 == 3) && (ne01 < 32768); + + if (ne1 == 1 || use_mc3) { cl_mem q_img = nullptr; cl_mem b_sub_buf = nullptr; cl_mem b_img = nullptr; - // image for q - img_fmt = { CL_R, CL_UNSIGNED_INT32}; - memset(&img_desc, 0, sizeof(img_desc)); - img_desc.image_type = CL_MEM_OBJECT_IMAGE1D_BUFFER; - img_desc.image_width = M * K / 2 / 4; - img_desc.buffer = extra0_q4_k->q; - CL_CHECK((q_img = clCreateImage(context, CL_MEM_READ_ONLY, &img_fmt, &img_desc, NULL, &err), err)); + const bool use_tiled = !use_mc3 && use_q4k_tiled(backend_ctx, src0); + + // image for q (not needed for the tiled path, which reads __global) + if (!use_tiled) { + img_fmt = { CL_R, CL_UNSIGNED_INT32}; + memset(&img_desc, 0, sizeof(img_desc)); + img_desc.image_type = CL_MEM_OBJECT_IMAGE1D_BUFFER; + img_desc.image_width = M * K / 2 / 4; + img_desc.buffer = extra0_q4_k->q; + CL_CHECK((q_img = clCreateImage(context, CL_MEM_READ_ONLY, &img_fmt, &img_desc, NULL, &err), err)); + } // subbuffer for activations region.origin = offset1; @@ -18155,27 +19127,173 @@ static void ggml_cl_mul_mat_q4_k_f32_adreno(ggml_backend_t backend, const ggml_t img_desc.buffer = b_sub_buf; CL_CHECK((b_img = clCreateImage(context, CL_MEM_READ_ONLY, &img_fmt, &img_desc, NULL, &err), err)); - kernel = backend_ctx->kernel_gemv_noshuffle_q4_k_f32; - - CL_CHECK(clSetKernelArg(kernel, 0, sizeof(cl_mem), &q_img)); - CL_CHECK(clSetKernelArg(kernel, 1, sizeof(cl_mem), &extra0_q4_k->d)); - CL_CHECK(clSetKernelArg(kernel, 2, sizeof(cl_mem), &extra0_q4_k->dm)); - CL_CHECK(clSetKernelArg(kernel, 3, sizeof(cl_mem), &extra0_q4_k->s)); - CL_CHECK(clSetKernelArg(kernel, 4, sizeof(cl_mem), &b_img)); - CL_CHECK(clSetKernelArg(kernel, 5, sizeof(cl_mem), &extrad->data_device)); - CL_CHECK(clSetKernelArg(kernel, 6, sizeof(cl_ulong), &offsetd)); - CL_CHECK(clSetKernelArg(kernel, 7, sizeof(cl_int), &ne00)); - CL_CHECK(clSetKernelArg(kernel, 8, sizeof(cl_int), &ne01)); - CL_CHECK(clSetKernelArg(kernel, 9, sizeof(cl_uchar), &mask_d6)); - CL_CHECK(clSetKernelArg(kernel, 10, sizeof(cl_uchar), &mask_d4)); - CL_CHECK(clSetKernelArg(kernel, 11, sizeof(cl_uchar), &mask_hi2)); + // 4-output-per-WI o4 variant for the long-vocab lm_head/embed GEMV + // (ne01 = vocab ~256K on Gemma): shares one activation read across 4 + // output rows. Gated to large ne01 (lm_head/embed). Default on; opt-out + // GGML_OPENCL_Q4K_GEMV_O4=0. (Skipped when mc3 handles the ne1==3 verify.) + static const bool q4k_o4_env = []{ + const char * e = std::getenv("GGML_OPENCL_Q4K_GEMV_O4"); + return !e || e[0] == '\0' || e[0] != '0'; + }(); + const bool use_q4k_o4 = !use_tiled && !use_mc3 && q4k_o4_env && (ne01 % 4 == 0) && (ne01 >= 32768); + // Split-K across workgroups for small-M decode GEMVs. A single-token GEMV + // makes only CEIL_DIV(M/2,64) workgroups; even with the wide intra-WG split + // (16 subgroups) those all land on ONE CU, so small-M matmuls under-fill the + // 16 CUs and their bandwidth falls well short of what the large-M FFN matmuls + // reach. Adding a `ksplit` second grid dim that spreads K across WGs (+ a + // reduce pass) fills the CUs. Gate is M<=2560: the tiny M<=1024 ones only + // break even (the reduce dispatch eats the kernel win), but the big-K M=2560 + // cases (ffn_down, attn_output) make the per-call win dwarf the reduce, and + // are byte-identical. ffn_gate/up (large M) fill the CUs already and are excluded. + // + // DEVICE-GATED. Split-K buys GPU time by spending an extra kernel LAUNCH (the + // reduce), so it only pays where launches are cheap. That is a per-device + // property and it does not travel from the X2-90 this was tuned on. Measured + // with one binary, env A/B (tg32, GGML_OPENCL_Q4K_GEMV_SPLITK=0/1): + // + // X2-90 +3.36% gemma-4 E4B (the number this gate was built on) + // 840 -1.3% Qwen3.5-4B-Q4_K_M 14.00 -> 13.85 + // 850 -20.0% Qwen3-1.7B-Q4_K_M 6.97 -> 5.58 (6 interleaved reps) + // + // The kernel is not the problem. On the 850 split-K makes the GPU strictly + // faster -- total busy 537 -> 485 ms, this GEMV 43.7 -> 34.0 us/call (-22%) -- + // and still costs a fifth of decode, because the +3696 reduce dispatches cost + // ~550 us of HOST round-trip each against 2.7 us of GPU work (~200x; that part + // is ~95% host-bound at decode). The 840 pays the same tax at ~42 us/dispatch. + // Break-even needs launch cost below the ~9.7 us/call the split actually saves, + // so this is not a "the 850 is slow" adjustment that a faster part would fix -- + // the 840 is 13x cheaper per launch and still loses. + // + // Enabled where it is measured to win, i.e. X2E only. The X1-85 was measured + // afterwards and is NOT a win either: Qwen3.5-4B-Q4_K_M tg32, split-K off + // 17.98/18.10/18.19 vs on 18.03/17.91/17.97 = -0.7%, so X1E stays excluded on + // evidence rather than on absence of it. Do not widen this without a NEW + // measurement. The env still forces either way so every device stays measurable. + static const bool splitk_env_set = []{ + const char * e = std::getenv("GGML_OPENCL_Q4K_GEMV_SPLITK"); + return e && e[0] != '\0'; + }(); + static const bool splitk_env_on = []{ + const char * e = std::getenv("GGML_OPENCL_Q4K_GEMV_SPLITK"); + return !(e && e[0] == '0'); + }(); + const bool splitk_wg_env = splitk_env_set + ? splitk_env_on + : (backend_ctx->adreno_gen == ADRENO_GPU_GEN::X2E); + // Gate: small-M decode GEMVs that under-fill the 16 CUs even with the wide + // intra-WG split (all 16 subgroups land on one CU). M<=2560 covers Kcur/Vcur + // (M=1024), Qcur (2048), attn_output + ffn_down (2560). The tiny ones + // (M<=1024) only break even (reduce dispatch eats the kernel win), but the + // big-K M=2560 cases (ffn_down K=10240 @182us, attn_output @42us) have a + // large per-call win that dwarfs the ~5us reduce, so extending to 2560 nets + // positive end-to-end. ffn_gate/up (M=10240) already fill the CUs -> excluded. + const bool use_splitk = splitk_wg_env && !use_tiled && !use_q4k_o4 && !use_mc3 && ne01 <= 2560; + + if (use_splitk) { + const int nsg = 8; + const int ksplit = (ne01 <= 512) ? 8 : 4; // -> ~32 total WGs + const size_t gx = (size_t)CEIL_DIV(ne01/2, 64) * 64; + + backend_ctx->prealloc_splitk_partial.allocate( + backend_ctx->context, (size_t)ksplit * ne01 * sizeof(float)); + cl_mem partial = backend_ctx->prealloc_splitk_partial.buffer; + + cl_kernel ks = backend_ctx->kernel_gemv_noshuffle_q4_k_f32_splitk; + CL_CHECK(clSetKernelArg(ks, 0, sizeof(cl_mem), &q_img)); + CL_CHECK(clSetKernelArg(ks, 1, sizeof(cl_mem), &extra0_q4_k->d)); + CL_CHECK(clSetKernelArg(ks, 2, sizeof(cl_mem), &extra0_q4_k->dm)); + CL_CHECK(clSetKernelArg(ks, 3, sizeof(cl_mem), &extra0_q4_k->s)); + CL_CHECK(clSetKernelArg(ks, 4, sizeof(cl_mem), &b_img)); + CL_CHECK(clSetKernelArg(ks, 5, sizeof(cl_mem), &partial)); + CL_CHECK(clSetKernelArg(ks, 6, sizeof(cl_int), &ne00)); + CL_CHECK(clSetKernelArg(ks, 7, sizeof(cl_int), &ne01)); + CL_CHECK(clSetKernelArg(ks, 8, sizeof(cl_uchar), &mask_d6)); + CL_CHECK(clSetKernelArg(ks, 9, sizeof(cl_uchar), &mask_d4)); + CL_CHECK(clSetKernelArg(ks, 10, sizeof(cl_uchar), &mask_hi2)); + size_t lsk[3] = {64, (size_t)nsg, 1}; + size_t gsk[3] = {gx, (size_t)(nsg * ksplit), 1}; + backend_ctx->enqueue_ndrange_kernel(ks, 3, gsk, lsk, dst); + + cl_kernel kr = backend_ctx->kernel_gemv_splitk_reduce_f32; + CL_CHECK(clSetKernelArg(kr, 0, sizeof(cl_mem), &partial)); + CL_CHECK(clSetKernelArg(kr, 1, sizeof(cl_mem), &extrad->data_device)); + CL_CHECK(clSetKernelArg(kr, 2, sizeof(cl_ulong), &offsetd)); + CL_CHECK(clSetKernelArg(kr, 3, sizeof(cl_int), &ne01)); + CL_CHECK(clSetKernelArg(kr, 4, sizeof(cl_int), &ksplit)); + size_t lr[3] = {64, 1, 1}; + size_t gr[3] = {(size_t)CEIL_DIV(ne01, 64) * 64, 1, 1}; + backend_ctx->enqueue_ndrange_kernel(kr, 3, gr, lr, dst); + + if (q_img) CL_CHECK(clReleaseMemObject(q_img)); + CL_CHECK(clReleaseMemObject(b_sub_buf)); + CL_CHECK(clReleaseMemObject(b_img)); + return; + } - size_t local_work_size[3] = {64, 4, 1}; - size_t global_work_size[3] = {(size_t)CEIL_DIV(ne01/2, 64)*64, 4, 1}; + kernel = use_mc3 ? backend_ctx->kernel_gemv_noshuffle_q4_k_f32_mc3 + : use_tiled ? backend_ctx->kernel_gemv_noshuffle_q4_k_f32_tiled + : use_q4k_o4 ? backend_ctx->kernel_gemv_noshuffle_q4_k_f32_o4 + : backend_ctx->kernel_gemv_noshuffle_q4_k_f32; + + if (use_tiled) { + CL_CHECK(clSetKernelArg(kernel, 0, sizeof(cl_mem), &extra0_q4_k->q)); + CL_CHECK(clSetKernelArg(kernel, 1, sizeof(cl_mem), &extra0_q4_k->d)); + CL_CHECK(clSetKernelArg(kernel, 2, sizeof(cl_mem), &extra0_q4_k->dm)); + CL_CHECK(clSetKernelArg(kernel, 3, sizeof(cl_mem), &extra0_q4_k->s)); + CL_CHECK(clSetKernelArg(kernel, 4, sizeof(cl_mem), &b_img)); + CL_CHECK(clSetKernelArg(kernel, 5, sizeof(cl_mem), &extrad->data_device)); + CL_CHECK(clSetKernelArg(kernel, 6, sizeof(cl_ulong), &offsetd)); + CL_CHECK(clSetKernelArg(kernel, 7, sizeof(cl_int), &ne00)); + CL_CHECK(clSetKernelArg(kernel, 8, sizeof(cl_int), &ne01)); + } else { + CL_CHECK(clSetKernelArg(kernel, 0, sizeof(cl_mem), &q_img)); + CL_CHECK(clSetKernelArg(kernel, 1, sizeof(cl_mem), &extra0_q4_k->d)); + CL_CHECK(clSetKernelArg(kernel, 2, sizeof(cl_mem), &extra0_q4_k->dm)); + CL_CHECK(clSetKernelArg(kernel, 3, sizeof(cl_mem), &extra0_q4_k->s)); + CL_CHECK(clSetKernelArg(kernel, 4, sizeof(cl_mem), &b_img)); + CL_CHECK(clSetKernelArg(kernel, 5, sizeof(cl_mem), &extrad->data_device)); + CL_CHECK(clSetKernelArg(kernel, 6, sizeof(cl_ulong), &offsetd)); + CL_CHECK(clSetKernelArg(kernel, 7, sizeof(cl_int), &ne00)); + CL_CHECK(clSetKernelArg(kernel, 8, sizeof(cl_int), &ne01)); + CL_CHECK(clSetKernelArg(kernel, 9, sizeof(cl_uchar), &mask_d6)); + CL_CHECK(clSetKernelArg(kernel, 10, sizeof(cl_uchar), &mask_d4)); + CL_CHECK(clSetKernelArg(kernel, 11, sizeof(cl_uchar), &mask_hi2)); + } + + // Wide K-split for the decode GEMV: the default 4-subgroup K-split leaves + // each Adreno SP with only ~4 waves, too few to hide LPDDR weight-load + // latency, so even the large FFN matmuls run well below the achievable + // bandwidth. Widen to 16 subgroups/WG (= the 1024-lane Adreno WG max) so + // each SP holds enough in-flight memory requests. Prefill is unaffected (the + // GEMM path is separate) and coherence-identical (greedy output unchanged). + // Applies to the plain base + // GEMV only; tiled/o4/mc3 keep 4 (their reductions are hard-coded to 4). + // Layout-safe: the base kernel derives its K-split from get_local_size(1) + // and the packed block stride is a physical constant (independent of it). + // Opt-out: GGML_OPENCL_Q4K_GEMV_WIDE=0. + static const bool splitk_wide_env = []{ + const char * e = std::getenv("GGML_OPENCL_Q4K_GEMV_WIDE"); + return !e || e[0] == '\0' || e[0] != '0'; + }(); + const bool splitk_wide = splitk_wide_env && !use_tiled && !use_q4k_o4 && !use_mc3; + size_t nsg_y = splitk_wide ? 16 : 4; + // Cap the wide K-split by the kernel's real max WG. X1-class drivers cap + // this GEMV at 768 (< 64*16 = 1024), so an uncapped lws aborts the + // dispatch with CL_INVALID_WORK_GROUP_SIZE (-54) and breaks ALL q4_K + // decode for M>2560. nsg_y is a pure K-split (the base kernel reads it + // from get_local_size(1); the packed block stride is a physical constant), + // so halving it stays coherent — just a narrower split. X2 keeps 16 + // (maxwg 1024); X1 falls to 8. + if (splitk_wide) { + const size_t maxwg = backend_ctx->get_kernel_workgroup_size(kernel); + while (nsg_y > 4 && 64 * nsg_y > maxwg) { nsg_y >>= 1; } + } + size_t local_work_size[3] = {64, nsg_y, 1}; + size_t global_work_size[3] = {(size_t)CEIL_DIV(use_tiled ? ne01 : (use_q4k_o4 ? ne01/4 : ne01/2), 64)*64, nsg_y, 1}; backend_ctx->enqueue_ndrange_kernel(kernel, 3, global_work_size, local_work_size, dst); - CL_CHECK(clReleaseMemObject(q_img)); + if (q_img) CL_CHECK(clReleaseMemObject(q_img)); CL_CHECK(clReleaseMemObject(b_sub_buf)); CL_CHECK(clReleaseMemObject(b_img)); } else { @@ -18335,10 +19453,44 @@ static void ggml_cl_mul_mat_q4_k_f32_adreno(ggml_backend_t backend, const ggml_t } // gemm - kernel = backend_ctx->kernel_gemm_noshuffle_q4_k_f32; + // Small-batch (medium n_q) occupancy fix: at ne1<=8 the 2x8 grid is + // (1, ceil(M/2)) -> ~M/256 workgroups, which under-occupies the SP and + // makes the GEMM much slower than the ne1==1 GEMV at the same weight + // traffic. The _r1 (1-row) kernel doubles the M-axis workgroup count + // and removes the accumulator spill. Opt-in via env while validating. + static const bool q4k_gemm_r1 = (getenv("GGML_OPENCL_Q4K_GEMM_R1") != nullptr); + static const bool q4k_gemm_kimg = (getenv("GGML_OPENCL_Q4K_GEMM_KIMG") != nullptr); + // Cooperative-K (intra-WG K-split + reduction) for the small-batch + // (n_q in [2..8]) path: DEFAULT ON, opt out with GGML_OPENCL_Q4K_GEMM_COK=0. + // Byte-identical greedy output; large-batch (ne1>8) untouched. + static const char * q4k_cok_env = getenv("GGML_OPENCL_Q4K_GEMM_COK"); + static const bool q4k_gemm_cok = (q4k_cok_env == nullptr) || (atoi(q4k_cok_env) != 0); + const bool use_cok = q4k_gemm_cok && (ne1 <= 8); + const bool use_r1 = !use_cok && q4k_gemm_r1 && (ne1 <= 8); + // Weights-as-image (L1/TPL1) for the small-batch weight-read-bound path. + const bool use_kimg = !use_cok && !use_r1 && q4k_gemm_kimg && (ne1 <= 8); + + cl_mem q_img = nullptr; + if (use_kimg) { + img_fmt = { CL_R, CL_UNSIGNED_INT32 }; + memset(&img_desc, 0, sizeof(img_desc)); + img_desc.image_type = CL_MEM_OBJECT_IMAGE1D_BUFFER; + img_desc.image_width = M * K / 2 / 4; + img_desc.buffer = extra0_q4_k->q; + CL_CHECK((q_img = clCreateImage(context, CL_MEM_READ_ONLY, &img_fmt, &img_desc, NULL, &err), err)); + } + + kernel = use_cok ? backend_ctx->kernel_gemm_noshuffle_q4_k_f32_cok + : use_r1 ? backend_ctx->kernel_gemm_noshuffle_q4_k_f32_r1 + : use_kimg ? backend_ctx->kernel_gemm_noshuffle_q4_k_f32_kimg + : backend_ctx->kernel_gemm_noshuffle_q4_k_f32; int padded_N = N + padding; - CL_CHECK(clSetKernelArg(kernel, 0, sizeof(cl_mem), &extra0_q4_k->q)); + if (use_kimg) { + CL_CHECK(clSetKernelArg(kernel, 0, sizeof(cl_mem), &q_img)); + } else { + CL_CHECK(clSetKernelArg(kernel, 0, sizeof(cl_mem), &extra0_q4_k->q)); + } CL_CHECK(clSetKernelArg(kernel, 1, sizeof(cl_mem), &extra0_q4_k->s)); CL_CHECK(clSetKernelArg(kernel, 2, sizeof(cl_mem), &extra0_q4_k->d)); CL_CHECK(clSetKernelArg(kernel, 3, sizeof(cl_mem), &extra0_q4_k->dm)); @@ -18353,10 +19505,45 @@ static void ggml_cl_mul_mat_q4_k_f32_adreno(ggml_backend_t backend, const ggml_t CL_CHECK(clSetKernelArg(kernel, 12, sizeof(cl_uchar), &mask_d4)); CL_CHECK(clSetKernelArg(kernel, 13, sizeof(cl_uchar), &mask_hi2)); - size_t global_work_size[3] = {(size_t)CEIL_DIV(ne1, 8), (size_t)CEIL_DIV(ne01, 4), 1}; - size_t local_work_size[3] = {1, 128, 1}; + size_t global_work_size[3]; + size_t local_work_size[3]; + if (use_cok) { + // (COK_SG lanes x COK_NSG subgroups): one row per lane, K split + // across the COK_NSG subgroups. ne01 is a multiple of 64. + global_work_size[0] = (size_t)ne01; // rows + global_work_size[1] = 8; // COK_NSG + global_work_size[2] = 1; + local_work_size[0] = 64; // COK_SG + local_work_size[1] = 8; // COK_NSG + local_work_size[2] = 1; + } else if (use_r1) { + // 1 row per WI (opt-in occupancy experiment). + global_work_size[0] = (size_t)CEIL_DIV(ne1, 8); + global_work_size[1] = (size_t)ne01; + global_work_size[2] = 1; + local_work_size[0] = 1; + local_work_size[1] = 128; + local_work_size[2] = 1; + } else if (use_kimg) { + // kimg is a 2-row tile (opt-in weights-as-image experiment). + global_work_size[0] = (size_t)CEIL_DIV(ne1, 8); + global_work_size[1] = (size_t)CEIL_DIV(ne01, 2); + global_work_size[2] = 1; + local_work_size[0] = 1; + local_work_size[1] = 128; + local_work_size[2] = 1; + } else { + // Default: x2-unified base kernel is the 4-row (gx<<2) tile. + global_work_size[0] = (size_t)CEIL_DIV(ne1, 8); + global_work_size[1] = (size_t)CEIL_DIV(ne01, 4); + global_work_size[2] = 1; + local_work_size[0] = 1; + local_work_size[1] = 128; + local_work_size[2] = 1; + } backend_ctx->enqueue_ndrange_kernel(kernel, 3, global_work_size, local_work_size, dst); + if (q_img) CL_CHECK(clReleaseMemObject(q_img)); CL_CHECK(clReleaseMemObject(b_sub_buf)); CL_CHECK(clReleaseMemObject(b_sub_buf_trans)); CL_CHECK(clReleaseMemObject(b_img)); @@ -18404,29 +19591,65 @@ static void ggml_cl_mul_mat_q6_K_f32_adreno(ggml_backend_t backend, const ggml_t cl_image_desc img_desc; // subbuffer and image for activation - if (ne1 == 1) { + // Multi-column verify GEMV: route the spec/MTP verify q6_K matmuls (ne1==3) + // onto the efficient GEMV path instead of the transposed-GEMM dead-zone. + // Reuses the ne1==1 image setup (activation image sized by N=ne1). Byte- + // identical. Opt-in via GGML_OPENCL_Q6K_MC3=1 while validating. + static const bool q6k_mc3 = (getenv("GGML_OPENCL_Q6K_MC3") != nullptr); + // Per-layer only (ne01 < 32768): batched large-vocab lm_head stays on the + // existing path (x2-unified routes batched Q6_K lm_head to CPU; the Adreno + // GEMV corrupts it). Per-layer mc3 is byte-identical. + const bool use_q6k_mc3 = q6k_mc3 && (ne1 == 3) && (ne01 < 32768); + // Batched verify lm_head/embed (ne1==3, tiled layout): multi-column tiled + // GEMV — streams the large lm_head weight once across the 3 verify columns + // (the #1 MTP bottleneck; mc3 above can't, it reads the noshuffle layout). + const bool use_q6k_tiled_mc = q6k_mc3 && (ne1 == 3) && (ne01 >= 32768) && use_q6k_tiled(backend_ctx, src0); + + if (ne1 == 1 || use_q6k_mc3 || use_q6k_tiled_mc) { cl_mem ql_img = nullptr; cl_mem qh_img = nullptr; cl_mem b_sub_buffer = nullptr; cl_mem b_img = nullptr; - // image for ql - img_fmt.image_channel_order = CL_R; - img_fmt.image_channel_data_type = CL_FLOAT; - memset(&img_desc, 0, sizeof(img_desc)); - img_desc.image_type = CL_MEM_OBJECT_IMAGE1D_BUFFER; - img_desc.image_width = ne01 * ne00 / 8; - img_desc.buffer = extra0_q6_K->ql; - CL_CHECK((ql_img = clCreateImage(context, CL_MEM_READ_ONLY, &img_fmt, &img_desc, NULL, &err), err)); + // o4 = 4-output-per-WI variant for long-vocab lm_head/embed; gated to + // ne01 >= 32768 so per-layer q6_K (ne01=hidden 2-8K) keeps the 2-output + // kernel (o4 regresses there). o4_global reads the weights from __global + // coalesced instead of image1d_buffer -- the texture cache caps the + // read-once-per-token lm_head bandwidth, while __global reaches the higher + // rate the rest of the model gets. Both default ON; opt out via + // GGML_OPENCL_Q6K_GEMV_O4 / GGML_OPENCL_Q6K_GEMV_O4_GLOBAL = 0. + static const bool gemv_o4_env = []{ + const char * e = std::getenv("GGML_OPENCL_Q6K_GEMV_O4"); + return !e || e[0] == '\0' || e[0] != '0'; + }(); + static const bool o4_global_env = []{ + const char * e = std::getenv("GGML_OPENCL_Q6K_GEMV_O4_GLOBAL"); + return !e || e[0] == '\0' || e[0] != '0'; + }(); + const bool use_tiled = !use_q6k_mc3 && use_q6k_tiled(backend_ctx, src0); + const bool use_o4 = !use_tiled && !use_q6k_mc3 && gemv_o4_env && (ne01 % 4 == 0) && (ne01 >= 32768); + const bool use_o4_global = use_o4 && o4_global_env; + + // ql/qh image views are only needed when NOT reading weights from global. + if (!use_o4_global && !use_tiled) { + // image for ql + img_fmt.image_channel_order = CL_R; + img_fmt.image_channel_data_type = CL_FLOAT; + memset(&img_desc, 0, sizeof(img_desc)); + img_desc.image_type = CL_MEM_OBJECT_IMAGE1D_BUFFER; + img_desc.image_width = ne01 * ne00 / 8; + img_desc.buffer = extra0_q6_K->ql; + CL_CHECK((ql_img = clCreateImage(context, CL_MEM_READ_ONLY, &img_fmt, &img_desc, NULL, &err), err)); - // image for qh - img_fmt.image_channel_order = CL_R; - img_fmt.image_channel_data_type = CL_HALF_FLOAT; - memset(&img_desc, 0, sizeof(img_desc)); - img_desc.image_type = CL_MEM_OBJECT_IMAGE1D_BUFFER; - img_desc.image_width = ne01 * ne00 / 8; - img_desc.buffer = extra0_q6_K->qh; - CL_CHECK((qh_img = clCreateImage(context, CL_MEM_READ_ONLY, &img_fmt, &img_desc, NULL, &err), err)); + // image for qh + img_fmt.image_channel_order = CL_R; + img_fmt.image_channel_data_type = CL_HALF_FLOAT; + memset(&img_desc, 0, sizeof(img_desc)); + img_desc.image_type = CL_MEM_OBJECT_IMAGE1D_BUFFER; + img_desc.image_width = ne01 * ne00 / 8; + img_desc.buffer = extra0_q6_K->qh; + CL_CHECK((qh_img = clCreateImage(context, CL_MEM_READ_ONLY, &img_fmt, &img_desc, NULL, &err), err)); + } region.origin = offset1; region.size = ne00 * ne1 * sizeof(float); @@ -18440,10 +19663,20 @@ static void ggml_cl_mul_mat_q6_K_f32_adreno(ggml_backend_t backend, const ggml_t img_desc.buffer = b_sub_buffer; CL_CHECK((b_img = clCreateImage(context, CL_MEM_READ_ONLY, &img_fmt, &img_desc, NULL, &err), err)); - kernel = backend_ctx->kernel_gemv_noshuffle_q6_K_f32; + kernel = use_q6k_mc3 ? backend_ctx->kernel_gemv_noshuffle_q6_K_f32_mc3 + : use_q6k_tiled_mc ? backend_ctx->kernel_gemv_noshuffle_q6_K_f32_tiled_mc3 + : use_tiled ? backend_ctx->kernel_gemv_noshuffle_q6_K_f32_tiled + : use_o4_global ? backend_ctx->kernel_gemv_noshuffle_q6_K_f32_o4_global + : use_o4 ? backend_ctx->kernel_gemv_noshuffle_q6_K_f32_o4 + : backend_ctx->kernel_gemv_noshuffle_q6_K_f32; - CL_CHECK(clSetKernelArg(kernel, 0, sizeof(cl_mem), &ql_img)); - CL_CHECK(clSetKernelArg(kernel, 1, sizeof(cl_mem), &qh_img)); + if (use_o4_global || use_tiled) { + CL_CHECK(clSetKernelArg(kernel, 0, sizeof(cl_mem), &extra0_q6_K->ql)); + CL_CHECK(clSetKernelArg(kernel, 1, sizeof(cl_mem), &extra0_q6_K->qh)); + } else { + CL_CHECK(clSetKernelArg(kernel, 0, sizeof(cl_mem), &ql_img)); + CL_CHECK(clSetKernelArg(kernel, 1, sizeof(cl_mem), &qh_img)); + } CL_CHECK(clSetKernelArg(kernel, 2, sizeof(cl_mem), &extra0_q6_K->s)); CL_CHECK(clSetKernelArg(kernel, 3, sizeof(cl_mem), &extra0_q6_K->d)); CL_CHECK(clSetKernelArg(kernel, 4, sizeof(cl_mem), &b_img)); @@ -18452,16 +19685,67 @@ static void ggml_cl_mul_mat_q6_K_f32_adreno(ggml_backend_t backend, const ggml_t CL_CHECK(clSetKernelArg(kernel, 7, sizeof(cl_int), &ne00)); CL_CHECK(clSetKernelArg(kernel, 8, sizeof(cl_int), &ne01)); - size_t local_work_size[3] = {64, 4, 1}; - size_t global_work_size[3] = {(size_t)CEIL_DIV(ne01/2, 64)*64, 4, 1}; + const size_t gws_x = use_tiled + ? (size_t) CEIL_DIV(ne01, 64) * 64 + : use_o4 + ? (size_t) CEIL_DIV(ne01/4, 64) * 64 + : (size_t) CEIL_DIV(ne01/2, 64) * 64; + size_t local_work_size[3] = {64, 4, 1}; + size_t global_work_size[3] = {gws_x, 4, 1}; backend_ctx->enqueue_ndrange_kernel(kernel, 3, global_work_size, local_work_size, dst); - CL_CHECK(clReleaseMemObject(ql_img)); - CL_CHECK(clReleaseMemObject(qh_img)); + if (ql_img) CL_CHECK(clReleaseMemObject(ql_img)); + if (qh_img) CL_CHECK(clReleaseMemObject(qh_img)); CL_CHECK(clReleaseMemObject(b_sub_buffer)); CL_CHECK(clReleaseMemObject(b_img)); } else { + // Tiled-layout batched GEMM. When the weight was converted to the 64-row + // tiled canonical layout (use_q6k_tiled — the default for lm_head/embed), + // the plain noshuffle GEMM below reads it as plain-transposed and produces + // garbage. Use the batched GEMM that matches the decode tiled GEMV's + // layout; it reads the f32 activation directly (column-major, no transpose). + if (use_q6k_tiled(backend_ctx, src0)) { + cl_mem b_sub_buf_t = nullptr; + cl_mem b_img_t = nullptr; + + region.origin = offset1; + region.size = ne00 * ne1 * sizeof(float); + CL_CHECK((b_sub_buf_t = clCreateSubBuffer(extra1->data_device, 0, CL_BUFFER_CREATE_TYPE_REGION, ®ion, &err), err)); + + img_fmt.image_channel_order = CL_RGBA; + img_fmt.image_channel_data_type = CL_FLOAT; + memset(&img_desc, 0, sizeof(img_desc)); + img_desc.image_type = CL_MEM_OBJECT_IMAGE1D_BUFFER; + img_desc.image_width = ne00 * ne1 / 4; + img_desc.buffer = b_sub_buf_t; + CL_CHECK((b_img_t = clCreateImage(context, CL_MEM_READ_ONLY, &img_fmt, &img_desc, NULL, &err), err)); + + cl_kernel kt = backend_ctx->kernel_gemm_noshuffle_q6_K_f32_tiled; + CL_CHECK(clSetKernelArg(kt, 0, sizeof(cl_mem), &extra0_q6_K->ql)); + CL_CHECK(clSetKernelArg(kt, 1, sizeof(cl_mem), &extra0_q6_K->qh)); + CL_CHECK(clSetKernelArg(kt, 2, sizeof(cl_mem), &extra0_q6_K->s)); + CL_CHECK(clSetKernelArg(kt, 3, sizeof(cl_mem), &extra0_q6_K->d)); + CL_CHECK(clSetKernelArg(kt, 4, sizeof(cl_mem), &b_img_t)); + CL_CHECK(clSetKernelArg(kt, 5, sizeof(cl_mem), &extrad->data_device)); + CL_CHECK(clSetKernelArg(kt, 6, sizeof(cl_ulong), &offsetd)); + CL_CHECK(clSetKernelArg(kt, 7, sizeof(int), &ne00)); + CL_CHECK(clSetKernelArg(kt, 8, sizeof(int), &ne01)); + CL_CHECK(clSetKernelArg(kt, 9, sizeof(int), &ne1)); + + // Must match the kernel: NTILES=4 64-row tiles per work-group (256 rows), + // BN=8 output columns per work-group. + const int BN_T = 16; + const int WROWS = 4 * 64; // NTILES * TILE_ROWS + size_t local_work_size[3] = {64, 4, 1}; + size_t global_work_size[3] = {(size_t)CEIL_DIV(ne01, WROWS) * 64, 4, (size_t)CEIL_DIV(ne1, BN_T)}; + backend_ctx->enqueue_ndrange_kernel(kt, 3, global_work_size, local_work_size, dst); + + CL_CHECK(clReleaseMemObject(b_img_t)); + CL_CHECK(clReleaseMemObject(b_sub_buf_t)); + return; + } + cl_mem b_sub_buf; cl_mem b_buf_trans; cl_mem b_img; @@ -18573,7 +19857,19 @@ static void ggml_cl_mul_mat_q6_K_f32_adreno(ggml_backend_t backend, const ggml_t backend_ctx->enqueue_ndrange_kernel(kernel, 2, global_size_t, local_size_t, dst); // gemm - kernel = backend_ctx->kernel_gemm_noshuffle_q6_K_f32; + // Cooperative-K small-batch (n_q in [2..8]) path: intra-WG K-split, + // mirrors the q4_K _cok path (batched serving). OPT-IN + // (GGML_OPENCL_Q6K_GEMM_COK=1), DEFAULT OFF: q6_K is the tied lm_head/ + // output projection, so the K-reassociation perturbs final logits and + // greedy is NOT byte-identical (op-tests pass, output coherent, but not + // bit-exact). It is also NEUTRAL on end-to-end MTP (q4_K cok already + // captured that; the MTP bottleneck moved off the GEMMs). Keep opt-in + // for batched serving until PPL-validated on a non-GDN q6_K model. + static const char * q6k_cok_env = getenv("GGML_OPENCL_Q6K_GEMM_COK"); + static const bool q6k_gemm_cok = (q6k_cok_env != nullptr) && (atoi(q6k_cok_env) != 0); + const bool use_q6k_cok = q6k_gemm_cok && (ne1 <= 8); + kernel = use_q6k_cok ? backend_ctx->kernel_gemm_noshuffle_q6_K_f32_cok + : backend_ctx->kernel_gemm_noshuffle_q6_K_f32; int padded_N = ne1 + padding; cl_ushort mask_f000 = 0xF000; @@ -18593,8 +19889,23 @@ static void ggml_cl_mul_mat_q6_K_f32_adreno(ggml_backend_t backend, const ggml_t CL_CHECK(clSetKernelArg(kernel, 11, sizeof(cl_ushort),&mask_f000)); CL_CHECK(clSetKernelArg(kernel, 12, sizeof(cl_uchar), &mask_c0)); - size_t global_work_size[3] = {(size_t)CEIL_DIV(ne1, 8), (size_t)CEIL_DIV(ne01, 4), 1}; - size_t local_work_size[3] = {2, 128, 1}; + size_t global_work_size[3]; + size_t local_work_size[3]; + if (use_q6k_cok) { + global_work_size[0] = (size_t)ne01; // rows (1 per lane) + global_work_size[1] = 8; // COK_NSG + global_work_size[2] = 1; + local_work_size[0] = 64; // COK_SG + local_work_size[1] = 8; // COK_NSG + local_work_size[2] = 1; + } else { + global_work_size[0] = (size_t)CEIL_DIV(ne1, 8); + global_work_size[1] = (size_t)CEIL_DIV(ne01, 4); + global_work_size[2] = 1; + local_work_size[0] = 2; + local_work_size[1] = 128; + local_work_size[2] = 1; + } backend_ctx->enqueue_ndrange_kernel(kernel, 3, global_work_size, local_work_size, dst); CL_CHECK(clReleaseMemObject(b_sub_buf)); @@ -18650,7 +19961,15 @@ static void ggml_cl_mul_mat_q5_K_f32_adreno(ggml_backend_t backend, const ggml_t cl_uchar mask_d4 = 0x0F; cl_uchar mask_hi2 = 0xC0; - if (ne1 == 1) { + // Multi-column (N=3) verify GEMV for q5_K: route the spec/MTP verify batch + // (ne1==3) onto the efficient GEMV path instead of the transposed-GEMM dead- + // zone (gemm_noshuffle_q5_k, the #2 chunk of MTP decode on a Q4_0-mix model + // after q4_0 mc3). Reuses the ne1==1 GEMV image setup (q + qh + activations). + // Opt-in via GGML_OPENCL_Q5K_MC3=1. Per-layer only (ne01 < 32768). + static const bool q5k_mc3 = (getenv("GGML_OPENCL_Q5K_MC3") != nullptr); + const bool use_q5k_mc3 = q5k_mc3 && (ne1 >= 2 && ne1 <= 4) && (ne01 < 32768); + + if (ne1 == 1 || use_q5k_mc3) { cl_mem q_img = nullptr; cl_mem qh_img = nullptr; cl_mem b_sub_buf = nullptr; @@ -18685,7 +20004,8 @@ static void ggml_cl_mul_mat_q5_K_f32_adreno(ggml_backend_t backend, const ggml_t img_desc.buffer = b_sub_buf; CL_CHECK((b_img = clCreateImage(context, CL_MEM_READ_ONLY, &img_fmt, &img_desc, NULL, &err), err)); - kernel = backend_ctx->kernel_gemv_noshuffle_q5_k_f32; + kernel = use_q5k_mc3 ? backend_ctx->kernel_gemv_noshuffle_q5_k_f32_mc3 + : backend_ctx->kernel_gemv_noshuffle_q5_k_f32; CL_CHECK(clSetKernelArg(kernel, 0, sizeof(cl_mem), &q_img)); CL_CHECK(clSetKernelArg(kernel, 1, sizeof(cl_mem), &qh_img)); @@ -18700,6 +20020,9 @@ static void ggml_cl_mul_mat_q5_K_f32_adreno(ggml_backend_t backend, const ggml_t CL_CHECK(clSetKernelArg(kernel, 10, sizeof(cl_uchar), &mask_d6)); CL_CHECK(clSetKernelArg(kernel, 11, sizeof(cl_uchar), &mask_d4)); CL_CHECK(clSetKernelArg(kernel, 12, sizeof(cl_uchar), &mask_hi2)); + if (use_q5k_mc3) { + CL_CHECK(clSetKernelArg(kernel, 13, sizeof(cl_int), &ne1)); // n_cols + } size_t local_work_size[3] = {64, 4, 1}; size_t global_work_size[3] = {(size_t)CEIL_DIV(ne01/2, 64)*64, 4, 1}; @@ -19571,7 +20894,7 @@ static void ggml_cl_mul_mat(ggml_backend_t backend, const ggml_tensor * src0, co } // q4_k x fp32 - if (src0t == GGML_TYPE_Q4_K && src1t == GGML_TYPE_F32 && !use_flat_gemv_for_large_m_q4_K(src0)) { + if (src0t == GGML_TYPE_Q4_K && src1t == GGML_TYPE_F32 && !use_flat_gemv_for_large_m_q4_K(backend_ctx, src0)) { ggml_cl_mul_mat_q4_k_f32_adreno(backend, src0, src1, dst); return; } @@ -19607,6 +20930,35 @@ static void ggml_cl_mul_mat(ggml_backend_t backend, const ggml_tensor * src0, co backend_ctx->adreno_gen == ADRENO_GPU_GEN::A7X)) { switch(src0t) { case GGML_TYPE_F32: { + // Small-N f32 GEMV for the spec/MTP verify batch: the tiled GEMM + // below always computes a full 64x64 tile, so at ne11=3 with a + // skinny f32 weight (GDN ssm_alpha/ssm_beta, M=32) it launches one + // under-occupied WG at ~2.3% tile utilization. Route to a per-output + // (m,n) GEMV (64-thread WG, K-split + __local reduce) instead. + // Opt-in GGML_OPENCL_F32_MC=1; 2D contiguous, small N + skinny M only. + static const bool f32_mc = (getenv("GGML_OPENCL_F32_MC") != nullptr); + if (f32_mc && ne11 >= 2 && ne11 <= 8 && ne01 <= 512 && (ne00 % 4 == 0) && + ne02 == 1 && ne12 == 1 && ne13 == 1 && + ggml_is_contiguous(src0) && ggml_is_contiguous(src1)) { + cl_kernel kmc = backend_ctx->kernel_gemv_f32_f32_mc; + int stride_a = ne00, stride_b = ne00, stride_d = ne01; + CL_CHECK(clSetKernelArg(kmc, 0, sizeof(cl_mem), &extra0->data_device)); + CL_CHECK(clSetKernelArg(kmc, 1, sizeof(cl_ulong), &offset0)); + CL_CHECK(clSetKernelArg(kmc, 2, sizeof(cl_mem), &extra1->data_device)); + CL_CHECK(clSetKernelArg(kmc, 3, sizeof(cl_ulong), &offset1)); + CL_CHECK(clSetKernelArg(kmc, 4, sizeof(cl_mem), &extrad->data_device)); + CL_CHECK(clSetKernelArg(kmc, 5, sizeof(cl_ulong), &offsetd)); + CL_CHECK(clSetKernelArg(kmc, 6, sizeof(int), &ne00)); + CL_CHECK(clSetKernelArg(kmc, 7, sizeof(int), &ne01)); + CL_CHECK(clSetKernelArg(kmc, 8, sizeof(int), &ne11)); + CL_CHECK(clSetKernelArg(kmc, 9, sizeof(int), &stride_a)); + CL_CHECK(clSetKernelArg(kmc, 10, sizeof(int), &stride_b)); + CL_CHECK(clSetKernelArg(kmc, 11, sizeof(int), &stride_d)); + size_t gws[3] = {64, (size_t)ne01 * (size_t)ne11, 1}; + size_t lws[3] = {64, 1, 1}; + backend_ctx->enqueue_ndrange_kernel(kmc, 3, gws, lws, dst); + return; + } kernel = backend_ctx->kernel_mul_mm_f32_f32_l4_lm; nth0 = 128; // calculated as (BM*BN)/(TM*TN) @@ -20263,6 +21615,7 @@ static void ggml_cl_mul_mat(ggml_backend_t backend, const ggml_tensor * src0, co } // use custom matrix x vector kernel + bool use_f16_mrow = false; switch (src0t) { case GGML_TYPE_F32: //GGML_ASSERT(ne02 == ne12); @@ -20328,7 +21681,46 @@ static void ggml_cl_mul_mat(ggml_backend_t backend, const ggml_tensor * src0, co (ne12 % r2) == 0; if (ne11 * ne12 < 4) { - kernel = backend_ctx->kernel_mul_mat_f16_f32_1row; + // Decode (single token): the legacy _1row runs one 64-lane + // subgroup per WG (one output row), under-utilizing BW. Route the + // wide f16 weight matmuls (attn proj + lm_head) to the multi-row + // variant: MROW rows per WG -> more loads in flight + activation + // staged once in __local. ne00<=8192 bounds the LDS. The mrow WG + // is 64 x MROW = 1024 work-items (> Intel's 512 max) and reduces + // within a 64-wide subgroup, so skip on Intel. + if (backend_ctx->f16_mrow && backend_ctx->gpu_family != INTEL && + backend_ctx->kernel_mul_mat_f16_f32_mrow != nullptr && + ne00 >= 128 && ne01 >= 8 && ne00 % 4 == 0 && ne00 <= 8192) { + // The register-blocked / half8 variants cast the src0 row pointer to + // half4 / half8 (8- and 16-byte loads) with no scalar fallback inside + // the kernel. ne00 % 4 == 0 constrains the element count per row, NOT + // the byte stride between rows: a permuted or strided src0 (or a view + // at an odd offset) can leave nb01/nb02/nb03 unaligned. Only take them + // when every row this dispatch touches is aligned; the base mrow kernel + // re-checks per row and falls back to its scalar loop. + const cl_ulong row_addr_bits = offset0 | nb01 | nb02 | nb03; + const bool aligned8 = (row_addr_bits & 7) == 0; + const bool aligned16 = (row_addr_bits & 15) == 0; + + // Register-blocked variants: each subgroup does RPT rows (more + // weight loads in flight per lane). 8/16 use half8 (128-bit) + // loads, gated on ne00 % 8 == 0. + const int rpt = backend_ctx->f16_mrow_rpt; + if (rpt == 16 && ne00 % 8 == 0 && aligned16 && backend_ctx->kernel_mul_mat_f16_f32_mrow_h8r2 != nullptr) { + kernel = backend_ctx->kernel_mul_mat_f16_f32_mrow_h8r2; + } else if (rpt == 8 && ne00 % 8 == 0 && aligned16 && backend_ctx->kernel_mul_mat_f16_f32_mrow_h8 != nullptr) { + kernel = backend_ctx->kernel_mul_mat_f16_f32_mrow_h8; + } else if (rpt == 4 && aligned8 && backend_ctx->kernel_mul_mat_f16_f32_mrow_r4 != nullptr) { + kernel = backend_ctx->kernel_mul_mat_f16_f32_mrow_r4; + } else if (rpt == 2 && aligned8 && backend_ctx->kernel_mul_mat_f16_f32_mrow_r2 != nullptr) { + kernel = backend_ctx->kernel_mul_mat_f16_f32_mrow_r2; + } else { + kernel = backend_ctx->kernel_mul_mat_f16_f32_mrow; + } + use_f16_mrow = true; + } else { + kernel = backend_ctx->kernel_mul_mat_f16_f32_1row; + } } else if (adreno_use_lane_split && ne00 >= 64 && ne00 <= 128) { kernel = backend_ctx->kernel_mul_mat_f16_f32_l4_dr_lq; nrows = 1; @@ -20414,6 +21806,23 @@ static void ggml_cl_mul_mat(ggml_backend_t backend, const ggml_tensor * src0, co CL_CHECK(clSetKernelArg(kernel, 21, sizeof(int), &ne1)); CL_CHECK(clSetKernelArg(kernel, 22, sizeof(int), &r2)); CL_CHECK(clSetKernelArg(kernel, 23, sizeof(int), &r3)); + if (use_f16_mrow) { + const int MROW = 16; // must match MROW in mul_mv_f16_f32_mrow.cl + // rows-per-subgroup multiplier for the selected variant: + // 1/2/4 -> half4 register blocking; 8 -> half8(1 row); 16 -> half8(2 rows) + const int rpt = backend_ctx->f16_mrow_rpt; + int rmul; + if (rpt == 16) rmul = (ne00 % 8 == 0) ? 2 : 1; + else if (rpt == 8) rmul = 1; + else rmul = rpt; // 1,2,4 + const int rows_per_wg = MROW * rmul; + // __local activation buffer: ne00 floats, rounded up for float4 access + CL_CHECK(clSetKernelArg(kernel, 24, sizeof(float) * ((ne00 + 3) / 4 * 4), nullptr)); + size_t mrow_global[] = { (size_t)((ne01 + rows_per_wg - 1) / rows_per_wg) * 64, (size_t)ne11 * MROW, (size_t)ne12 * ne13 }; + size_t mrow_local[] = { 64, (size_t)MROW, 1 }; + backend_ctx->enqueue_ndrange_kernel(kernel, 3, mrow_global, mrow_local, dst); + return; + } break; case GGML_TYPE_Q1_0: { #ifdef GGML_OPENCL_SOA_Q diff --git a/ggml/src/ggml-opencl/kernels/cvt.cl b/ggml/src/ggml-opencl/kernels/cvt.cl index 3d6cff7cff0..acc8f980763 100644 --- a/ggml/src/ggml-opencl/kernels/cvt.cl +++ b/ggml/src/ggml-opencl/kernels/cvt.cl @@ -1110,6 +1110,78 @@ kernel void kernel_restore_block_q4_k_trans4_ns( } } +//------------------------------------------------------------------------------ +// kernel_convert_block_q4_k_tiled_ns +// +// Tiled-wide layout for the long-vocab q4_K lm_head/embed GEMV (decode path). +// Mirror of kernel_convert_block_q6_k_tiled_ns: recovers each weight's 4-bit +// code in CANONICAL ggml element order (e in [0,256)) and re-packs into 32 uints +// (8 codes/uint), stored TILED by 64 output rows so the matching GEMV +// (gemv_noshuffle_q4_k_f32_tiled) coalesces every weight load. The 12-byte +// packed scale block `s` and d/dm are stored per (row, K-block) tiled; the GEMV +// re-derives the 8 (scale,min) pairs via get_scale_min_k4, exactly like the o4 +// kernel. Both ends owned here -> correct by construction vs the reference q4_K +// dequant. Requires ne01 % 64 == 0 (gated host-side). Buffer sizes identical to +// the trans4_ns layout. +// +// q uint4 granule g of (row r, K-block sb): idx = ((rt*ne00_blk+sb)*8 + g)*64 + rit +// s (12 bytes) of (r, sb): idx = (rt*ne00_blk+sb)*64 + rit, *12 +// d/dm (half) of (r, sb): idx = (rt*ne00_blk+sb)*64 + rit +// where rt = r/64, rit = r%64. +//------------------------------------------------------------------------------ +kernel void kernel_convert_block_q4_k_tiled_ns( + __global struct block_q4_K * src0, + __global uint * dst_q, // 32 uints / superblock (4-bit codes, 8 codes/uint) + __global half * dst_d, // 1 half / superblock + __global half * dst_dm, // 1 half / superblock + __global uchar * dst_s, // K_SCALE_SIZE (12) bytes / superblock + uint ne00, + uint ne01 +) { + uint i00 = get_global_id(1); // K-block index (superblock along ne00) + uint i01 = get_global_id(0); // output row index (along ne01) + uint i02 = get_global_id(2); // batch + + uint ne00_blk = ne00 / QK_K; + + uint src_blk_offset = i00 + i01 * ne00_blk + i02 * ne00_blk * ne01; + __global struct block_q4_K * b = src0 + src_blk_offset; + + uint rt = i01 / 64; + uint rit = i01 % 64; + uint tile_blk = (i02 * (ne01 / 64) + rt) * ne00_blk + i00; + + // --- recover canonical 4-bit codes in e-order, pack 8 codes/uint --- + uint qw[32] = {0}; + for (uint e = 0; e < 256; ++e) { + uint g = e >> 6; // group 0..3 (q advances 32 bytes/group) + uint within = e & 63u; + uint hlf = within >> 5; // 0 = low nibble, 1 = high nibble + uint l = within & 31u; // 0..31 + uchar byte = b->q[g * 32u + l]; + uint code = (hlf == 0u) ? (uint)(byte & 0x0F) : (uint)(byte >> 4); + qw[e >> 3] |= code << ((e & 7u) * 4u); + } + + for (uint gr = 0; gr < 8; ++gr) { + uint base = (tile_blk * 8u + gr) * 64u + rit; // uint4 index + dst_q[base * 4u + 0u] = qw[gr * 4u + 0u]; + dst_q[base * 4u + 1u] = qw[gr * 4u + 1u]; + dst_q[base * 4u + 2u] = qw[gr * 4u + 2u]; + dst_q[base * 4u + 3u] = qw[gr * 4u + 3u]; + } + + // packed scales (12 bytes), tiled per (row, block) + __global uchar * s_dst = dst_s + (tile_blk * 64u + rit) * K_SCALE_SIZE; + #pragma unroll + for (int i = 0; i < K_SCALE_SIZE; ++i) { + s_dst[i] = b->s[i]; + } + + dst_d [tile_blk * 64u + rit] = b->d; + dst_dm[tile_blk * 64u + rit] = b->dm; +} + kernel void kernel_convert_block_q5_k_trans4_ns( __global struct block_q5_K * src0, __global uint * dst_qs, @@ -1494,6 +1566,105 @@ kernel void kernel_restore_block_mxfp4_trans( b->e = src_e[src_blk_offset]; } +//------------------------------------------------------------------------------ +// kernel_convert_block_q6_k_tiled_ns +// +// Tiled-wide layout for the long-vocab q6_K lm_head/embed GEMV (decode path). +// Unlike *_trans4_ns (which mirrors the bit-interleave the legacy 2-output GEMV +// consumes), this kernel is correct-by-construction against the CANONICAL ggml +// q6_K dequant: it recovers each weight's 6-bit code in element order e in +// [0,256), then re-packs low-4-bits into 32 uints (8 codes/uint) and high-2-bits +// into 16 uints (16 codes/uint). The matching GEMV (gemv_noshuffle_q6_k_f32_tiled) +// unpacks the same order, so both ends are owned here. +// +// Storage is TILED by 64 output rows so the GEMV's 64-thread tile coalesces: +// ql uint4 granule g of (row r, K-block sb): idx = ((rt*ne00_blk + sb)*8 + g)*64 + rit +// qh uint4 granule g: idx = ((rt*ne00_blk + sb)*4 + g)*64 + rit +// scales (char16) of (r, sb): idx = (rt*ne00_blk + sb)*64 + rit +// d (half) of (r, sb): idx = (rt*ne00_blk + sb)*64 + rit +// where rt = r/64, rit = r%64. Requires ne01 % 64 == 0 (gated host-side). +// Buffer sizes are byte-identical to the trans4_ns layout. +//------------------------------------------------------------------------------ +kernel void kernel_convert_block_q6_k_tiled_ns( + __global struct block_q6_K * src0, + __global uint * dst_ql, // 32 uints / superblock (low 4 bits, 8 codes/uint) + __global uint * dst_qh, // 16 uints / superblock (high 2 bits, 16 codes/uint) + __global half * dst_d, // 1 half / superblock + __global char * dst_s, // 16 chars/ superblock + uint ne00, + uint ne01 +) { + uint i00 = get_global_id(1); // K-block index (superblock along ne00) + uint i01 = get_global_id(0); // output row index (along ne01) + uint i02 = get_global_id(2); // batch + + uint ne00_blk = ne00 / QK_K; + + // Source block: row-major over (i02, i01, i00). + uint src_blk_offset = i00 + i01 * ne00_blk + i02 * ne00_blk * ne01; + __global struct block_q6_K * b = src0 + src_blk_offset; + + uint rt = i01 / 64; + uint rit = i01 % 64; + uint tile_blk = (i02 * (ne01 / 64) + rt) * ne00_blk + i00; // tile-major (row-tile, K-block) + + // --- recover canonical 6-bit codes, pack into ql (4b) + qh (2b) in e-order --- + // 32 ql-uints (8 low-nibbles each) + 16 qh-uints (16 2-bit slots each). + uint qlw[32] = {0}; + uint qhw[16] = {0}; + + for (uint e = 0; e < 256; ++e) { + uint n = (e >= 128) ? 1u : 0u; // which 128-half + uint within = e - n * 128u; + uint q = within / 32u; // quadrant 0..3 + uint l = within % 32u; // 0..31 + + uint off_ql = n * 64u; // raw ql byte base for this half + uint off_qh = n * 32u; // raw qh byte base for this half + + uchar low4; + uchar qlb0 = b->ql[off_ql + l]; + uchar qlb1 = b->ql[off_ql + l + 32]; + if (q == 0) low4 = qlb0 & 0x0F; + else if (q == 1) low4 = qlb1 & 0x0F; + else if (q == 2) low4 = (qlb0 >> 4) & 0x0F; + else low4 = (qlb1 >> 4) & 0x0F; + + uchar hi2 = (b->qh[off_qh + l] >> (q * 2u)) & 0x03; + + // pack low4 (e-order): uint e/8, nibble (e%8) + qlw[e >> 3] |= ((uint)low4) << ((e & 7u) * 4u); + // pack hi2 (e-order): uint e/16, 2-bit slot (e%16) + qhw[e >> 4] |= ((uint)hi2) << ((e & 15u) * 2u); + } + + // --- write tiled --- + for (uint g = 0; g < 8; ++g) { + uint base = (tile_blk * 8u + g) * 64u + rit; // uint4 index + dst_ql[base * 4u + 0u] = qlw[g * 4u + 0u]; + dst_ql[base * 4u + 1u] = qlw[g * 4u + 1u]; + dst_ql[base * 4u + 2u] = qlw[g * 4u + 2u]; + dst_ql[base * 4u + 3u] = qlw[g * 4u + 3u]; + } + for (uint g = 0; g < 4; ++g) { + uint base = (tile_blk * 4u + g) * 64u + rit; // uint4 index + dst_qh[base * 4u + 0u] = qhw[g * 4u + 0u]; + dst_qh[base * 4u + 1u] = qhw[g * 4u + 1u]; + dst_qh[base * 4u + 2u] = qhw[g * 4u + 2u]; + dst_qh[base * 4u + 3u] = qhw[g * 4u + 3u]; + } + + // scales: 16 chars contiguous per (row, block), tiled + __global char * s_dst = dst_s + (tile_blk * 64u + rit) * 16u; + #pragma unroll + for (int i = 0; i < 16; ++i) { + s_dst[i] = b->scales[i]; + } + + // super-block scale + dst_d[tile_blk * 64u + rit] = b->d; +} + kernel void kernel_convert_block_mxfp4_trans4_ns( global struct block_mxfp4 * src0, __global uint * dst_q, diff --git a/ggml/src/ggml-opencl/kernels/gemm_noshuffle_q4_k_f32.cl b/ggml/src/ggml-opencl/kernels/gemm_noshuffle_q4_k_f32.cl index 22b4e911462..c379a9a3998 100644 --- a/ggml/src/ggml-opencl/kernels/gemm_noshuffle_q4_k_f32.cl +++ b/ggml/src/ggml-opencl/kernels/gemm_noshuffle_q4_k_f32.cl @@ -4,6 +4,7 @@ #pragma OPENCL EXTENSION cl_qcom_reqd_sub_group_size : enable #define ADRENO_GPU 1 #define REQD_SUBGROUP_SIZE_128 __attribute__((qcom_reqd_sub_group_size("full"))) +#define REQD_SUBGROUP_SIZE_64 __attribute__((qcom_reqd_sub_group_size("half"))) #endif #define QK_K 256 #define K_SCALE_SIZE 12 @@ -171,3 +172,319 @@ kernel void kernel_gemm_noshuffle_q4_k_f32( vstore4((float4)(c0.s7, c1.s7, c2.s7, c3.s7), 0, dst + idx); } } + +// 1x8 per-WI tile (1 output row x 8 output cols). For the small-batch +// (medium n_q, e.g. MTP/spec verify) path where the 2x8 kernel is starved: +// at ne1<=8 the grid is (1, ceil(M/2)) -> only ~M/256 workgroups, leaving +// the SP under-occupied. 1 row per WI doubles the M-axis workgroup count +// (ceil(M/1)/128 vs ceil(M/2)/128) AND collapses the accumulators to a +// single half8 (16 regs, no spill), so more waves co-reside. Same weight +// traffic as 2x8 (rows never share weights); the win is pure occupancy. +#ifdef ADRENO_GPU +REQD_SUBGROUP_SIZE_128 +#endif +kernel void kernel_gemm_noshuffle_q4_k_f32_r1( + global const ushort * src0_q, + global const uchar * src0_s, + global const half * src0_d, + global const half * src0_dm, + read_only image1d_buffer_t src1, + global float * dst, + ulong offsetd, + int m, + int n, + int k, + int n_no_padding, + uchar mask_d6, + uchar mask_d4, + uchar mask_hi2 +) { + dst = (global float *)((global char *)dst + offsetd); + int n_4 = n >> 2; + int gy = get_global_id(0); + int gx = get_global_id(1); // 1 row per WI + + half8 c0 = 0; + half8 B; + half dq; + + int num_blocks_K = k / QK_K; + + global const ushort * weight_ptr = src0_q + gx; + global const half * d_ptr = src0_d + gx; + global const half * dm_ptr = src0_dm + gx; + + for (int i = 0; i < k; i += 32) { + int sb_idx = i / QK_K; + int sub_idx = (i / 32) % 8; + + half dd = d_ptr [sb_idx * m]; + half dmm = dm_ptr[sb_idx * m]; + + global const uchar * sc0 = src0_s + sb_idx * K_SCALE_SIZE * m + gx; + + uchar sv0, mn0; + get_scale_min_k4(sub_idx, sc0, m, &sv0, &mn0, mask_d6, mask_d4, mask_hi2); + + half scale = convert_half(convert_float(dd) * (float)sv0); + half mval = convert_half(convert_float(dmm) * (float)mn0); + + for (int l = 0; l < 32; l += 4) { + int ki = i + l; + ushort bits = weight_ptr[(ki/4) * m]; + + B.s0123 = read_imageh(src1, gy*2 + (ki+0) * n_4); + B.s4567 = read_imageh(src1, gy*2+1 + (ki+0) * n_4); + dq = (bits & 0x000F) * scale - mval; + c0 += B * dq; + + B.s0123 = read_imageh(src1, gy*2 + (ki+1) * n_4); + B.s4567 = read_imageh(src1, gy*2+1 + (ki+1) * n_4); + dq = ((bits & 0x00F0) >> 4) * scale - mval; + c0 += B * dq; + + B.s0123 = read_imageh(src1, gy*2 + (ki+2) * n_4); + B.s4567 = read_imageh(src1, gy*2+1 + (ki+2) * n_4); + dq = ((bits & 0x0F00) >> 8) * scale - mval; + c0 += B * dq; + + B.s0123 = read_imageh(src1, gy*2 + (ki+3) * n_4); + B.s4567 = read_imageh(src1, gy*2+1 + (ki+3) * n_4); + dq = ((bits & 0xF000) >> 12) * scale - mval; + c0 += B * dq; + } + } + + // Output: 8 cols, 1 row per col-step. Scalar store, coalesced across + // neighbouring WIs (consecutive gx -> consecutive dst addresses). + int idx = (gy<<3)*m + gx; + if (idx < m*n_no_padding) { dst[idx] = c0.s0; idx += m; } + if (idx < m*n_no_padding) { dst[idx] = c0.s1; idx += m; } + if (idx < m*n_no_padding) { dst[idx] = c0.s2; idx += m; } + if (idx < m*n_no_padding) { dst[idx] = c0.s3; idx += m; } + if (idx < m*n_no_padding) { dst[idx] = c0.s4; idx += m; } + if (idx < m*n_no_padding) { dst[idx] = c0.s5; idx += m; } + if (idx < m*n_no_padding) { dst[idx] = c0.s6; idx += m; } + if (idx < m*n_no_padding) { dst[idx] = c0.s7; } +} + +// 2x8 tile, but weights read through an image1d_buffer (CL_R/UINT32 over the +// same packed-q buffer) instead of a plain global buffer. The ne1==1 GEMV +// already does this and is much faster per weight byte than this GEMM at +// small n_q; the structural difference is the image path hits the dedicated +// TPL1 weight cache (L1) while the global path only reaches L2. At small n_q +// the forward is weight-read-bound, so L1-cached weights is the lever. +// The 2 adjacent rows the 2x8 tile reads as a ushort2 are exactly one uint32, +// so the vload2 becomes a single read_imageui at index gx + (ki/4)*(m/2). +#ifdef ADRENO_GPU +REQD_SUBGROUP_SIZE_128 +#endif +kernel void kernel_gemm_noshuffle_q4_k_f32_kimg( + read_only image1d_buffer_t src0_q_img, + global const uchar * src0_s, + global const half * src0_d, + global const half * src0_dm, + read_only image1d_buffer_t src1, + global float * dst, + ulong offsetd, + int m, + int n, + int k, + int n_no_padding, + uchar mask_d6, + uchar mask_d4, + uchar mask_hi2 +) { + dst = (global float *)((global char *)dst + offsetd); + int n_4 = n >> 2; + int m_2 = m >> 1; + int gy = get_global_id(0); + int gx = get_global_id(1); + int gx_2 = gx << 1; + + half8 c0 = 0, c1 = 0; + half8 B; + half2 dequantized_weights; + + int num_blocks_K = k / QK_K; + + global const half * d_ptr = src0_d + gx_2; + global const half * dm_ptr = src0_dm + gx_2; + + for (int i = 0; i < k; i += 32) { + int sb_idx = i / QK_K; + int sub_idx = (i / 32) % 8; + + half2 d = vload2(0, d_ptr + sb_idx * m); + half2 dm = vload2(0, dm_ptr + sb_idx * m); + + global const uchar * sc0 = src0_s + sb_idx * K_SCALE_SIZE * m + (gx_2+0); + global const uchar * sc1 = sc0 + 1; + + uchar sv0, mn0, sv1, mn1; + get_scale_min_k4(sub_idx, sc0, m, &sv0, &mn0, mask_d6, mask_d4, mask_hi2); + get_scale_min_k4(sub_idx, sc1, m, &sv1, &mn1, mask_d6, mask_d4, mask_hi2); + + half2 scale = convert_half2(convert_float2(d) * convert_float2((uchar2)(sv0, sv1))); + half2 mval = convert_half2(convert_float2(dm) * convert_float2((uchar2)(mn0, mn1))); + + for (int l = 0; l < 32; l += 4) { + int ki = i + l; + uint wpacked = read_imageui(src0_q_img, gx + (ki/4) * m_2).x; + ushort2 bits2 = (ushort2)((ushort)(wpacked & 0xFFFFu), (ushort)(wpacked >> 16)); + + // j=0 + B.s0123 = read_imageh(src1, gy*2 + (ki+0) * n_4); + B.s4567 = read_imageh(src1, gy*2+1 + (ki+0) * n_4); + dequantized_weights.s0 = (bits2.s0 & 0x000F) * scale.s0 - mval.s0; + dequantized_weights.s1 = (bits2.s1 & 0x000F) * scale.s1 - mval.s1; + c0 += B * dequantized_weights.s0; + c1 += B * dequantized_weights.s1; + + // j=1 + B.s0123 = read_imageh(src1, gy*2 + (ki+1) * n_4); + B.s4567 = read_imageh(src1, gy*2+1 + (ki+1) * n_4); + dequantized_weights.s0 = ((bits2.s0 & 0x00F0) >> 4) * scale.s0 - mval.s0; + dequantized_weights.s1 = ((bits2.s1 & 0x00F0) >> 4) * scale.s1 - mval.s1; + c0 += B * dequantized_weights.s0; + c1 += B * dequantized_weights.s1; + + // j=2 + B.s0123 = read_imageh(src1, gy*2 + (ki+2) * n_4); + B.s4567 = read_imageh(src1, gy*2+1 + (ki+2) * n_4); + dequantized_weights.s0 = ((bits2.s0 & 0x0F00) >> 8) * scale.s0 - mval.s0; + dequantized_weights.s1 = ((bits2.s1 & 0x0F00) >> 8) * scale.s1 - mval.s1; + c0 += B * dequantized_weights.s0; + c1 += B * dequantized_weights.s1; + + // j=3 + B.s0123 = read_imageh(src1, gy*2 + (ki+3) * n_4); + B.s4567 = read_imageh(src1, gy*2+1 + (ki+3) * n_4); + dequantized_weights.s0 = ((bits2.s0 & 0xF000) >> 12) * scale.s0 - mval.s0; + dequantized_weights.s1 = ((bits2.s1 & 0xF000) >> 12) * scale.s1 - mval.s1; + c0 += B * dequantized_weights.s0; + c1 += B * dequantized_weights.s1; + } + } + + int idx = (gy<<3)*m + (gx<<1); + if (idx+1 < m*n_no_padding) { vstore2((float2)(c0.s0, c1.s0), 0, dst + idx); idx += m; } + if (idx+1 < m*n_no_padding) { vstore2((float2)(c0.s1, c1.s1), 0, dst + idx); idx += m; } + if (idx+1 < m*n_no_padding) { vstore2((float2)(c0.s2, c1.s2), 0, dst + idx); idx += m; } + if (idx+1 < m*n_no_padding) { vstore2((float2)(c0.s3, c1.s3), 0, dst + idx); idx += m; } + if (idx+1 < m*n_no_padding) { vstore2((float2)(c0.s4, c1.s4), 0, dst + idx); idx += m; } + if (idx+1 < m*n_no_padding) { vstore2((float2)(c0.s5, c1.s5), 0, dst + idx); idx += m; } + if (idx+1 < m*n_no_padding) { vstore2((float2)(c0.s6, c1.s6), 0, dst + idx); idx += m; } + if (idx+1 < m*n_no_padding) { vstore2((float2)(c0.s7, c1.s7), 0, dst + idx); } +} + +// Cooperative-K GEMM for the small-batch (n_q in [2..8]) path. Mirrors the +// ne1==1 GEMV's structure: a WG is (COK_SG lanes x COK_NSG subgroups); each +// lane owns ONE output row and computes its 8 (padded) columns, and the +// COK_NSG subgroups SPLIT the K reduction round-robin, combining via a +// __local reduction. This is the thing the per-WI GEMM lacked — at small n_q +// the old kernel had ~M/256 workgroups each walking all of K serially; this +// has M/64 workgroups AND COK_NSG-way K parallelism. Uses REQD_SUBGROUP_SIZE_64 +// + barrier (same safe reduction pattern as the GEMV; never sub_group_reduce +// at full width on X2 per the GDN miscompile note). +#define COK_NSG 8 +#define COK_SG 64 +#ifdef ADRENO_GPU +REQD_SUBGROUP_SIZE_64 +#endif +kernel void kernel_gemm_noshuffle_q4_k_f32_cok( + global const ushort * src0_q, + global const uchar * src0_s, + global const half * src0_d, + global const half * src0_dm, + read_only image1d_buffer_t src1, + global float * dst, + ulong offsetd, + int m, + int n, + int k, + int n_no_padding, + uchar mask_d6, + uchar mask_d4, + uchar mask_hi2 +) { + dst = (global float *)((global char *)dst + offsetd); + int n_4 = n >> 2; + int gx = get_global_id(0); // output row + int sg = get_local_id(1); // subgroup index (K-split lane) + int lane = get_local_id(0); // lane within subgroup (0..COK_SG-1) + + int num_blocks_K = k / QK_K; + int num_32blk = k / 32; + + global const ushort * weight_ptr = src0_q + gx; + global const half * d_ptr = src0_d + gx; + global const half * dm_ptr = src0_dm + gx; + + half8 acc = 0; + half8 B; + half dq; + + for (int blk = sg; blk < num_32blk; blk += COK_NSG) { + int i = blk << 5; // blk * 32 + int sb_idx = blk >> 3; // (blk*32) / QK_K (QK_K = 256 = 32*8) + int sub_idx = blk & 7; // (i/32) % 8 + + half dd = d_ptr [sb_idx * m]; + half dmm = dm_ptr[sb_idx * m]; + + global const uchar * sc0 = src0_s + sb_idx * K_SCALE_SIZE * m + gx; + uchar sv0, mn0; + get_scale_min_k4(sub_idx, sc0, m, &sv0, &mn0, mask_d6, mask_d4, mask_hi2); + half scale = convert_half(convert_float(dd) * (float)sv0); + half mval = convert_half(convert_float(dmm) * (float)mn0); + + for (int l = 0; l < 32; l += 4) { + int ki = i + l; + ushort bits = weight_ptr[(ki>>2) * m]; + + B.s0123 = read_imageh(src1, (ki+0) * n_4); + B.s4567 = read_imageh(src1, 1 + (ki+0) * n_4); + dq = (bits & 0x000F) * scale - mval; + acc += B * dq; + + B.s0123 = read_imageh(src1, (ki+1) * n_4); + B.s4567 = read_imageh(src1, 1 + (ki+1) * n_4); + dq = ((bits & 0x00F0) >> 4) * scale - mval; + acc += B * dq; + + B.s0123 = read_imageh(src1, (ki+2) * n_4); + B.s4567 = read_imageh(src1, 1 + (ki+2) * n_4); + dq = ((bits & 0x0F00) >> 8) * scale - mval; + acc += B * dq; + + B.s0123 = read_imageh(src1, (ki+3) * n_4); + B.s4567 = read_imageh(src1, 1 + (ki+3) * n_4); + dq = ((bits & 0xF000) >> 12) * scale - mval; + acc += B * dq; + } + } + + // cross-subgroup reduction over the K-split (float for accuracy) + local float8 reduceLM[COK_SG * (COK_NSG - 1)]; + if (sg > 0) { + reduceLM[(sg - 1) * COK_SG + lane] = convert_float8(acc); + } + barrier(CLK_LOCAL_MEM_FENCE); + + if (sg == 0) { + float8 sum = convert_float8(acc); + for (int s = 0; s < COK_NSG - 1; s++) { + sum += reduceLM[s * COK_SG + lane]; + } + int idx = gx; + if (idx < m*n_no_padding) { dst[idx] = sum.s0; idx += m; } + if (idx < m*n_no_padding) { dst[idx] = sum.s1; idx += m; } + if (idx < m*n_no_padding) { dst[idx] = sum.s2; idx += m; } + if (idx < m*n_no_padding) { dst[idx] = sum.s3; idx += m; } + if (idx < m*n_no_padding) { dst[idx] = sum.s4; idx += m; } + if (idx < m*n_no_padding) { dst[idx] = sum.s5; idx += m; } + if (idx < m*n_no_padding) { dst[idx] = sum.s6; idx += m; } + if (idx < m*n_no_padding) { dst[idx] = sum.s7; } + } +} diff --git a/ggml/src/ggml-opencl/kernels/gemm_noshuffle_q6_k_f32.cl b/ggml/src/ggml-opencl/kernels/gemm_noshuffle_q6_k_f32.cl index 3a9c624508a..141f6a2f688 100644 --- a/ggml/src/ggml-opencl/kernels/gemm_noshuffle_q6_k_f32.cl +++ b/ggml/src/ggml-opencl/kernels/gemm_noshuffle_q6_k_f32.cl @@ -5,6 +5,7 @@ #pragma OPENCL EXTENSION cl_qcom_reqd_sub_group_size : enable #define ADRENO_GPU 1 #define REQD_SUBGROUP_SIZE_128 __attribute__((qcom_reqd_sub_group_size("full"))) +#define REQD_SUBGROUP_SIZE_64 __attribute__((qcom_reqd_sub_group_size("half"))) #endif #ifdef ADRENO_GPU @@ -138,3 +139,107 @@ kernel void kernel_gemm_noshuffle_q6_K_f32( vstore4((float4)(c0.s7, c1.s7, c2.s7, c3.s7), 0, dst + idx); } } + +// Cooperative-K q6_K GEMM for the small-batch (n_q in [2..8]) path. Same idea +// as the q4_K _cok kernel: WG = (COK_SG lanes x COK_NSG subgroups), each lane +// owns ONE output row (half8 over the 8 padded cols), and the COK_NSG +// subgroups split the K iterations round-robin and combine via a __local +// reduction. Replaces the default 4-row-per-WI tile that walked all of K alone +// (~M/512 WGs + serial reduction) at small n_q. REQD_SUBGROUP_SIZE_64 + +// barrier (never sub_group_reduce at full width on X2). +#define COK_NSG 8 +#define COK_SG 64 +#ifdef ADRENO_GPU +REQD_SUBGROUP_SIZE_64 +#endif +kernel void kernel_gemm_noshuffle_q6_K_f32_cok( + global const ushort * src0_ql, + global const uchar * src0_qh, + global const ushort * src0_s, + global const half * src0_d, + read_only image1d_buffer_t src1, + global float * dst, + ulong offsetd, + int m, + int n, + int k, + int n_no_padding, + ushort mask_f000, + uchar mask_c0 +) { + dst = (global float *)( (global char *)dst + offsetd ); + + int n_4 = n >> 2; + int gx = get_global_id(0); // output row + int sg = get_local_id(1); // subgroup index (K-split) + int lane = get_local_id(0); // lane within subgroup + + global const ushort * ptr_ql = src0_ql + gx; + global const uchar * ptr_qh = src0_qh + gx; + global const ushort * ptr_s = src0_s + gx; + global const half * ptr_d = src0_d + gx; + + half8 acc = 0; + half8 B; + half dq; + + int num_iter = k >> 2; // k/4 iterations, 4 k-values each + + for (int ib = sg; ib < num_iter; ib += COK_NSG) { + int i = ib << 2; // ib * 4 + + ushort bits4 = ptr_ql[ib * m]; // ql for row gx at this 4-block + uchar bits2 = ptr_qh[ib * m]; // qh + + ushort s_packed = ptr_s[(i >> 5) * m]; // (i/16/2) = i/32 + char2 sc2 = as_char2(s_packed); + char scale_s = (((i >> 4) & 1) == 0) ? sc2.s0 : sc2.s1; // (i/16)%2 + half scale_d = ptr_d[(i >> 8) * m]; // i/256 + + // j=0 + B.s0123 = read_imageh(src1, (i + 0)*n_4 + 0); + B.s4567 = read_imageh(src1, (i + 0)*n_4 + 1); + dq = (convert_half((bits4 & 0x000F) | ((bits2 & 0x03) << 4)) - 32.f) * scale_s * scale_d; + acc += B * dq; + + // j=1 + B.s0123 = read_imageh(src1, (i + 1)*n_4 + 0); + B.s4567 = read_imageh(src1, (i + 1)*n_4 + 1); + dq = (convert_half(((bits4 & 0x00F0) >> 4) | ((bits2 & 0x0C) << 2)) - 32.f) * scale_s * scale_d; + acc += B * dq; + + // j=2 + B.s0123 = read_imageh(src1, (i + 2)*n_4 + 0); + B.s4567 = read_imageh(src1, (i + 2)*n_4 + 1); + dq = (convert_half(((bits4 & 0x0F00) >> 8) | (bits2 & 0x30)) - 32.f) * scale_s * scale_d; + acc += B * dq; + + // j=3 + B.s0123 = read_imageh(src1, (i + 3)*n_4 + 0); + B.s4567 = read_imageh(src1, (i + 3)*n_4 + 1); + dq = (convert_half(((bits4 & mask_f000) >> 12) | ((bits2 & mask_c0) >> 2)) - 32.f) * scale_s * scale_d; + acc += B * dq; + } + + local float8 reduceLM[COK_SG * (COK_NSG - 1)]; + if (sg > 0) { + reduceLM[(sg - 1) * COK_SG + lane] = convert_float8(acc); + } + barrier(CLK_LOCAL_MEM_FENCE); + + if (sg == 0) { + float8 sum = convert_float8(acc); + for (int s = 0; s < COK_NSG - 1; s++) { + sum += reduceLM[s * COK_SG + lane]; + } + int idx = gx; + if (idx < m*n_no_padding) { dst[idx] = sum.s0; idx += m; } + if (idx < m*n_no_padding) { dst[idx] = sum.s1; idx += m; } + if (idx < m*n_no_padding) { dst[idx] = sum.s2; idx += m; } + if (idx < m*n_no_padding) { dst[idx] = sum.s3; idx += m; } + if (idx < m*n_no_padding) { dst[idx] = sum.s4; idx += m; } + if (idx < m*n_no_padding) { dst[idx] = sum.s5; idx += m; } + if (idx < m*n_no_padding) { dst[idx] = sum.s6; idx += m; } + if (idx < m*n_no_padding) { dst[idx] = sum.s7; } + } +} diff --git a/ggml/src/ggml-opencl/kernels/gemm_noshuffle_q6_k_f32_tiled.cl b/ggml/src/ggml-opencl/kernels/gemm_noshuffle_q6_k_f32_tiled.cl new file mode 100644 index 00000000000..ffd943a2781 --- /dev/null +++ b/ggml/src/ggml-opencl/kernels/gemm_noshuffle_q6_k_f32_tiled.cl @@ -0,0 +1,136 @@ +// Batched (N>1) q6_K GEMM over the 64-row-TILED canonical layout produced by +// kernel_convert_block_q6_k_tiled_ns (cvt.cl). Companion to the decode kernel +// kernel_gemv_noshuffle_q6_K_f32_tiled: SAME pack, SAME canonical e-order +// dequant (correct by construction vs reference ggml q6_K), extended to N output +// columns. Makes the batched lm_head/embed (perplexity, spec-decode verify, +// batched serving) correct on GPU while keeping the tiled convert the fast decode +// GEMV depends on. +// +// One work-item owns one output ROW for a block of BN columns. A work-group is +// {64 lanes, NTILES subgroups} = NTILES*64 rows; the global z dimension tiles the +// N columns by BN. Each work-item computes its row's FULL K (no K-split, so no +// cross-subgroup reduction), which lets the whole work-group share one staged +// activation block: +// +// __local activation staging — the BN columns of the current superblock (BN*256 +// floats) are loaded into __local once per superblock, cooperatively by all +// NTILES*64 work-items, then every row reads its activation from __local. This +// removes the ~Nrows-fold redundant image reads of the first version (each lane +// re-read the activation), which made the batched GEMM ~2x slower than the plain +// noshuffle GEMM. +// +// Weights are read from __global (coalesced) — matching the decode kernel; the +// lm_head weight is streamed with little reuse where coalesced global beats the +// Adreno texture cache. + +#pragma OPENCL EXTENSION cl_khr_fp16 : enable + +#ifdef cl_qcom_reqd_sub_group_size +#pragma OPENCL EXTENSION cl_qcom_reqd_sub_group_size : enable +#define ADRENO_GPU 1 +#define REQD_SUBGROUP_SIZE_64 __attribute__((qcom_reqd_sub_group_size("half"))) +#endif + +#define NTILES 4 // 64-row tiles per work-group (NTILES*64 = 256 rows) +#define TILE_ROWS 64 +#define BN 16 // output columns handled per work-group (global z step) +#define WG_THREADS (NTILES * TILE_ROWS) + +#if defined(ADRENO_GPU) +REQD_SUBGROUP_SIZE_64 +#endif +kernel void kernel_gemm_noshuffle_q6_K_f32_tiled( + __global uint4 * src0_ql, // tiled: 8 uint4 granules / superblock + __global uint4 * src0_qh, // tiled: 4 uint4 granules / superblock + __global char * src0_s, // tiled: 16 chars / superblock + __global half * src0_d, // tiled: 1 half / superblock + read_only image1d_buffer_t src1, // activation [ne00, ne11] f32 (RGBA), column-major + global float * dst, + ulong offsetd, + int ne00, + int ne01, + int ne11 +) { + int rit = get_local_id(0); // 0..63 (lane within a tile; coalesces weight loads) + int sg = get_local_id(1); // 0..NTILES-1 + int lid = sg * TILE_ROWS + rit; // 0..WG_THREADS-1 (flat local id) + int row = get_group_id(0) * WG_THREADS + lid; + int rt = row / TILE_ROWS; // global 64-row tile index + int col0 = get_global_id(2) * BN; // first output column of this block + + int nb = ne00 / 256; // superblocks per row + int act_col_stride = ne00 / 4; // activation float4 pixels per column + + const bool row_ok = row < ne01; + + // staged activation: BN columns x 256 elements for the current superblock + __local float lact[BN * 256]; + + float acc[BN]; + #pragma unroll + for (int j = 0; j < BN; ++j) acc[j] = 0.0f; + + for (int sb = 0; sb < nb; ++sb) { + // cooperatively stage BN columns' 256 activation elements (= BN*64 float4) + for (int p = lid; p < BN * 64; p += WG_THREADS) { + int j = p >> 6; // column within the BN block (p / 64) + int e4 = p & 63; // element-quad within the column (p % 64) + int c = col0 + j; + float4 v = (c < ne11) + ? read_imagef(src1, c * act_col_stride + sb * 64 + e4) + : (float4)(0.0f); + lact[p * 4 + 0] = v.x; + lact[p * 4 + 1] = v.y; + lact[p * 4 + 2] = v.z; + lact[p * 4 + 3] = v.w; // lact[j*256 + e], e = e4*4 + t + } + barrier(CLK_LOCAL_MEM_FENCE); + + if (row_ok) { + int tile_blk = rt * nb + sb; // ne02 == 1 for lm_head/embed + + float dval = (float)src0_d[tile_blk * TILE_ROWS + rit]; + __global char * sc = src0_s + (tile_blk * TILE_ROWS + rit) * 16; + + uint ql[32]; + uint qh[16]; + #pragma unroll + for (int g = 0; g < 8; ++g) { + uint4 v = src0_ql[(tile_blk * 8 + g) * TILE_ROWS + rit]; + ql[g*4+0] = v.x; ql[g*4+1] = v.y; ql[g*4+2] = v.z; ql[g*4+3] = v.w; + } + #pragma unroll + for (int g = 0; g < 4; ++g) { + uint4 v = src0_qh[(tile_blk * 4 + g) * TILE_ROWS + rit]; + qh[g*4+0] = v.x; qh[g*4+1] = v.y; qh[g*4+2] = v.z; qh[g*4+3] = v.w; + } + + // NOTE: the e loop (256) is deliberately NOT unrolled. Fully unrolling + // 256*BN MACs overflows the in-process Adreno compiler (host stack + // overflow at clBuildProgram, same class as the FA DK=512 OOM). + for (int e = 0; e < 256; ++e) { + uint low4 = (ql[e >> 3] >> ((e & 7) * 4)) & 0xF; + uint hi2 = (qh[e >> 4] >> ((e & 15) * 2)) & 0x3; + int code = (int)(low4 | (hi2 << 4)) - 32; + int sidx = ((e >> 7) << 3) + (((e >> 5) & 3) << 1) + ((e >> 4) & 1); + float cs = (float)code * (float)sc[sidx] * dval; + #pragma unroll + for (int j = 0; j < BN; ++j) { + acc[j] += cs * lact[j * 256 + e]; + } + } + } + barrier(CLK_LOCAL_MEM_FENCE); + } + + if (row_ok) { + dst = (global float*)((global char*)dst + offsetd); + #pragma unroll + for (int j = 0; j < BN; ++j) { + int c = col0 + j; + if (c < ne11) { + dst[(ulong)c * ne01 + row] = acc[j]; + } + } + } +} diff --git a/ggml/src/ggml-opencl/kernels/gemv_noshuffle_q4_0_f32.cl b/ggml/src/ggml-opencl/kernels/gemv_noshuffle_q4_0_f32.cl index 8de0de1cc3a..023e848f734 100644 --- a/ggml/src/ggml-opencl/kernels/gemv_noshuffle_q4_0_f32.cl +++ b/ggml/src/ggml-opencl/kernels/gemv_noshuffle_q4_0_f32.cl @@ -277,3 +277,107 @@ __kernel void kernel_gemv_noshuffle_q4_0_f32( } } + +// Multi-column (N in [2..4]) variant of the q4_0 decode GEMV, for the speculative +// / MTP verify batch (n_cols = 2..4 = drafted + bonus positions). Routes the small- +// batch verify OFF the transposed-GEMM dead-zone (gemm_noshuffle_q4_0) onto the +// efficient GEMV path. Each K-block's weights (regA hi+lo) are loaded ONCE and +// reused across the n_cols activation columns. Per-column accumulation is +// independent and identical to n_cols standalone GEMVs. n_cols==3 is byte-identical +// to the original mc3 (col3 disabled, slots 6/7 stay zero). Kept the _mc3 name. +#ifdef VECTOR_SUB_GROUP_BROADCAST +#define MC_DQ_HI dequantizeBlockAccum_ns_sgbroadcast_8_hi +#define MC_DQ_LO dequantizeBlockAccum_ns_sgbroadcast_8_lo +#else +#define MC_DQ_HI dequantizeBlockAccum_ns_sgbroadcast_1_hi +#define MC_DQ_LO dequantizeBlockAccum_ns_sgbroadcast_1_lo +#endif +// One column c: load this column's activation (own brace scope so the macros' +// `shared_y` decl is re-scoped), then dequant (hi+lo) against the shared weights. +#define MC_COL_Q40(ts, c) \ + { if (slid < 4) { regB.s0123 = read_imagef(src1, (c)*COL_STRIDE + slid*2 + k*8); \ + regB.s4567 = read_imagef(src1, (c)*COL_STRIDE + 1 + slid*2 + k*8); } \ + MC_DQ_HI(ts, as_ushort8(regA_hi), regS, regB); \ + MC_DQ_LO(ts, as_ushort8(regA_lo), regS, regB); } + +#ifdef ADRENO_GPU +REQD_SUBGROUP_SIZE_64 +#endif +__kernel void kernel_gemv_noshuffle_q4_0_f32_mc3( + __read_only image1d_buffer_t src0_q, // quantized A + global half2 * src0_d, // A scales + __read_only image1d_buffer_t src1, // B (n_cols columns, col-major image) + global float * dst, // C (column-major [M x n_cols]) + ulong offsetd, + int ne00, // K + int ne01, // M + int n_cols) // N (2..4) +{ + uint groupId = get_local_id(1); + uint gid = get_global_id(0); + ushort slid = get_sub_group_local_id(); + + uint K = ne00; + uint M = ne01; + + uint LINE_STRIDE_A = M / 2; + // BLOCK_STRIDE_A is the LAYOUT stride between consecutive K-blocks = 4 uints + // per q4_0 block * M (set by the trans4_ns convert). The "4" is uints/block, NOT + // the subgroup count — keep it fixed so the K-split count (nsg) can vary. + uint BLOCK_STRIDE_A = N_SIMDGROUP * M; // = 4 * M (N_SIMDGROUP is the #define 4) + uint COL_STRIDE = K / 4; // float4 pixels per activation column + uint nsg = get_local_size(1); // runtime K-split (4 default, 8 small-M) + + __private uint4 regA_hi, regA_lo; + __private half2 regS; + __private float8 regB; + + __private float2 ts0 = (float2)(0.0f); + __private float2 ts1 = (float2)(0.0f); + __private float2 ts2 = (float2)(0.0f); + __private float2 ts3 = (float2)(0.0f); + + for (uint k = groupId; k < (K / QK4_0); k += nsg) { + regS = src0_d[gid + k * LINE_STRIDE_A]; + + // weights loaded ONCE, reused across the columns + regA_hi.s0 = read_imageui(src0_q, (gid + k * BLOCK_STRIDE_A + LINE_STRIDE_A * 0)).x; + regA_hi.s1 = read_imageui(src0_q, (gid + k * BLOCK_STRIDE_A + LINE_STRIDE_A * 1)).x; + regA_hi.s2 = read_imageui(src0_q, (gid + k * BLOCK_STRIDE_A + LINE_STRIDE_A * 2)).x; + regA_hi.s3 = read_imageui(src0_q, (gid + k * BLOCK_STRIDE_A + LINE_STRIDE_A * 3)).x; + regA_lo.s0 = read_imageui(src0_q, (gid + k * BLOCK_STRIDE_A + LINE_STRIDE_A * 4)).x; + regA_lo.s1 = read_imageui(src0_q, (gid + k * BLOCK_STRIDE_A + LINE_STRIDE_A * 5)).x; + regA_lo.s2 = read_imageui(src0_q, (gid + k * BLOCK_STRIDE_A + LINE_STRIDE_A * 6)).x; + regA_lo.s3 = read_imageui(src0_q, (gid + k * BLOCK_STRIDE_A + LINE_STRIDE_A * 7)).x; + + MC_COL_Q40(ts0, 0); + MC_COL_Q40(ts1, 1); + if (n_cols > 2) MC_COL_Q40(ts2, 2); + if (n_cols > 3) MC_COL_Q40(ts3, 3); + } + + // cross-subgroup reduce over nsg subgroups: pack the (up to 4) columns' float2 + // into a float8. Generalized to runtime nsg (4 default, 8 for small-M). Each + // subgroup writes its partial; subgroup 0 sums the rest into its own acc. At + // nsg==4 this is byte-identical to the original (sums subgroups 1,2,3 in order). + __local float8 reduceLM[SIMDGROUP_WIDTH * 8]; + float8 acc = (float8)(ts0.s0, ts0.s1, ts1.s0, ts1.s1, ts2.s0, ts2.s1, ts3.s0, ts3.s1); + reduceLM[groupId * SIMDGROUP_WIDTH + slid] = acc; + + barrier(CLK_LOCAL_MEM_FENCE); + + if (groupId == 0) { + for (uint g = 1; g < nsg; g++) { + acc += reduceLM[g * SIMDGROUP_WIDTH + slid]; + } + dst = (global float*)((global char*)dst + offsetd); + // dst is column-major [M rows x n_cols cols]: (row, col) at col*M + row + vstore2((float2)(acc.s0, acc.s1), 0, &(dst[0 * M + gid * 2])); + vstore2((float2)(acc.s2, acc.s3), 0, &(dst[1 * M + gid * 2])); + if (n_cols > 2) vstore2((float2)(acc.s4, acc.s5), 0, &(dst[2 * M + gid * 2])); + if (n_cols > 3) vstore2((float2)(acc.s6, acc.s7), 0, &(dst[3 * M + gid * 2])); + } +} +#undef MC_COL_Q40 +#undef MC_DQ_HI +#undef MC_DQ_LO diff --git a/ggml/src/ggml-opencl/kernels/gemv_noshuffle_q4_1_f32.cl b/ggml/src/ggml-opencl/kernels/gemv_noshuffle_q4_1_f32.cl index 5fa3127806a..2ccf4214c0b 100644 --- a/ggml/src/ggml-opencl/kernels/gemv_noshuffle_q4_1_f32.cl +++ b/ggml/src/ggml-opencl/kernels/gemv_noshuffle_q4_1_f32.cl @@ -286,3 +286,99 @@ kernel void kernel_gemv_noshuffle_q4_1_f32( } } + +// Multi-column (N in [2..4]) variant of the q4_1 decode GEMV (spec/MTP verify) = +// q4_0 mc3 + the q4_1 per-block min (regM; dequant = q*scale + minv). n_cols=2..4; +// routes the small-batch verify OFF the gemm_noshuffle_q4_1 dead-zone. n_cols==3 is +// byte-identical to the original mc3. NB: this file spells the vec-broadcast define +// BROADCAT (no S) — match it so the fast _8 path compiles. +#ifdef VECTOR_SUB_GROUP_BROADCAT +#define MC_DQ1_HI dequantizeBlockAccum_ns_sgbroadcast_8_hi +#define MC_DQ1_LO dequantizeBlockAccum_ns_sgbroadcast_8_lo +#else +#define MC_DQ1_HI dequantizeBlockAccum_ns_sgbroadcast_1_hi +#define MC_DQ1_LO dequantizeBlockAccum_ns_sgbroadcast_1_lo +#endif +#define MC_COL_Q41(ts, c) \ + { if (slid < 4) { regB.s0123 = read_imagef(src1, (c)*COL_STRIDE + slid*2 + k*8); \ + regB.s4567 = read_imagef(src1, (c)*COL_STRIDE + 1 + slid*2 + k*8); } \ + MC_DQ1_HI(ts, as_ushort8(regA_hi), regS, regM, regB); \ + MC_DQ1_LO(ts, as_ushort8(regA_lo), regS, regM, regB); } +#ifdef ADRENO_GPU +REQD_SUBGROUP_SIZE_64 +#endif +kernel void kernel_gemv_noshuffle_q4_1_f32_mc3( + read_only image1d_buffer_t src0_q, + global half2 * src0_d, + global half2 * src0_m, + read_only image1d_buffer_t src1, + global float * dst, + ulong offsetd, + int ne00, + int ne01, + int n_cols) +{ + uint groupId = get_local_id(1); + uint gid = get_global_id(0); + ushort slid = get_sub_group_local_id(); + + uint K = ne00; + uint M = ne01; + + uint LINE_STRIDE_A = M / 2; + uint BLOCK_STRIDE_A = NSUBGROUPS * M; + uint COL_STRIDE = K / 4; // float4 pixels per activation column + + private uint4 regA_hi, regA_lo; + private half2 regS, regM; + private float8 regB; + + private float2 ts0 = (float2)(0.0f); + private float2 ts1 = (float2)(0.0f); + private float2 ts2 = (float2)(0.0f); + private float2 ts3 = (float2)(0.0f); + + for (uint k = groupId; k < (K / QK4_0); k += NSUBGROUPS) { + regS = src0_d[gid + k * LINE_STRIDE_A]; + regM = src0_m[gid + k * LINE_STRIDE_A]; + + // weights loaded ONCE, reused across the columns + regA_hi.s0 = read_imageui(src0_q, (gid + k * BLOCK_STRIDE_A + LINE_STRIDE_A * 0)).x; + regA_hi.s1 = read_imageui(src0_q, (gid + k * BLOCK_STRIDE_A + LINE_STRIDE_A * 1)).x; + regA_hi.s2 = read_imageui(src0_q, (gid + k * BLOCK_STRIDE_A + LINE_STRIDE_A * 2)).x; + regA_hi.s3 = read_imageui(src0_q, (gid + k * BLOCK_STRIDE_A + LINE_STRIDE_A * 3)).x; + regA_lo.s0 = read_imageui(src0_q, (gid + k * BLOCK_STRIDE_A + LINE_STRIDE_A * 4)).x; + regA_lo.s1 = read_imageui(src0_q, (gid + k * BLOCK_STRIDE_A + LINE_STRIDE_A * 5)).x; + regA_lo.s2 = read_imageui(src0_q, (gid + k * BLOCK_STRIDE_A + LINE_STRIDE_A * 6)).x; + regA_lo.s3 = read_imageui(src0_q, (gid + k * BLOCK_STRIDE_A + LINE_STRIDE_A * 7)).x; + + MC_COL_Q41(ts0, 0); + MC_COL_Q41(ts1, 1); + if (n_cols > 2) MC_COL_Q41(ts2, 2); + if (n_cols > 3) MC_COL_Q41(ts3, 3); + } + + // cross-subgroup reduce: pack the (up to 4) columns' float2 into a float8. + local float8 reduceLM[SUBGROUP_SIZE * 3]; + float8 acc = (float8)(ts0.s0, ts0.s1, ts1.s0, ts1.s1, ts2.s0, ts2.s1, ts3.s0, ts3.s1); + if (groupId == 1) { reduceLM[SUBGROUP_SIZE * 0 + slid] = acc; } + if (groupId == 2) { reduceLM[SUBGROUP_SIZE * 1 + slid] = acc; } + if (groupId == 3) { reduceLM[SUBGROUP_SIZE * 2 + slid] = acc; } + + barrier(CLK_LOCAL_MEM_FENCE); + + if (groupId == 0) { + acc += reduceLM[SUBGROUP_SIZE * 0 + slid]; + acc += reduceLM[SUBGROUP_SIZE * 1 + slid]; + acc += reduceLM[SUBGROUP_SIZE * 2 + slid]; + dst = (global float*)((global char*)dst + offsetd); + // dst is column-major [M rows x n_cols cols]: (row, col) at col*M + row + vstore2((float2)(acc.s0, acc.s1), 0, &(dst[0 * M + gid * 2])); + vstore2((float2)(acc.s2, acc.s3), 0, &(dst[1 * M + gid * 2])); + if (n_cols > 2) vstore2((float2)(acc.s4, acc.s5), 0, &(dst[2 * M + gid * 2])); + if (n_cols > 3) vstore2((float2)(acc.s6, acc.s7), 0, &(dst[3 * M + gid * 2])); + } +} +#undef MC_COL_Q41 +#undef MC_DQ1_HI +#undef MC_DQ1_LO diff --git a/ggml/src/ggml-opencl/kernels/gemv_noshuffle_q4_k_f32.cl b/ggml/src/ggml-opencl/kernels/gemv_noshuffle_q4_k_f32.cl index 9ab0dee693e..c0078131e9f 100644 --- a/ggml/src/ggml-opencl/kernels/gemv_noshuffle_q4_k_f32.cl +++ b/ggml/src/ggml-opencl/kernels/gemv_noshuffle_q4_k_f32.cl @@ -228,12 +228,20 @@ kernel void kernel_gemv_noshuffle_q4_k_f32( uint groupId = get_local_id(1); uint gid = get_global_id(0); ushort slid = get_sub_group_local_id(); + // K-split factor = #subgroups in the WG. Read from the launch (NOT a compile + // constant) so small-M projections (Kcur/Vcur/Qcur) can dispatch a wider + // K-split (more waves/SP -> latency hiding) while large-M keeps 4. The + // physical weight layout stride below is INDEPENDENT of this (see BLOCK_STRIDE_A). + uint nsg = get_local_size(1); uint K = ne00; uint M = ne01; uint LINE_STRIDE_A = M / 2; - uint BLOCK_STRIDE_A = NSUBGROUPS * M; + // Physical per-K-block stride in the packed image: 8 uints/block-row-pair * + // (M/2) row-pairs = 4*M uints. This is a layout constant, not tied to nsg. + uint BLOCK_STRIDE_A = 4 * M; + uint scales_per_row = (K / QK_K) * 12; // The x-grid is padded to CEIL_DIV(ne01/2,64)*64, so when ne01 % 128 != 0 the // tail lanes hold gid >= ne01/2. The output stores below are guarded, but the @@ -259,7 +267,7 @@ kernel void kernel_gemv_noshuffle_q4_k_f32( private float2 totalSum = (float2)(0.0f); - for (uint k = groupId; k < (K / 32); k += NSUBGROUPS) { + for (uint k = groupId; k < (K / 32); k += nsg) { uint sb = k / 8; uint j = k % 8; @@ -303,28 +311,21 @@ kernel void kernel_gemv_noshuffle_q4_k_f32( #endif // VECTOR_SUB_GROUP_BROADCAST } - // reduction in local memory, assumes #wave=4 - local float2 reduceLM[SUBGROUP_SIZE * 3]; - if (groupId == 1) { - reduceLM[SUBGROUP_SIZE * 0 + slid] = totalSum; - } - if (groupId == 2) { - reduceLM[SUBGROUP_SIZE * 1 + slid] = totalSum; - } - if (groupId == 3) { - reduceLM[SUBGROUP_SIZE * 2 + slid] = totalSum; + // Cross-subgroup reduction in local memory. Generalized to nsg subgroups + // (was a hard-coded 4-wave unroll). Sized for up to 16 subgroups (the widest + // K-split we dispatch for small M). At nsg==4 the accumulation order is + // identical to the original unroll -> byte-identical for the large-M path. + local float2 reduceLM[SUBGROUP_SIZE * 15]; + if (groupId > 0) { + reduceLM[SUBGROUP_SIZE * (groupId - 1) + slid] = totalSum; } barrier(CLK_LOCAL_MEM_FENCE); if (groupId == 0) { - totalSum += reduceLM[SUBGROUP_SIZE * 0 + slid]; - } - if (groupId == 0) { - totalSum += reduceLM[SUBGROUP_SIZE * 1 + slid]; - } - if (groupId == 0) { - totalSum += reduceLM[SUBGROUP_SIZE * 2 + slid]; + for (uint i = 0; i < nsg - 1; ++i) { + totalSum += reduceLM[SUBGROUP_SIZE * i + slid]; + } } // 2 outputs per fiber in wave 0 @@ -339,3 +340,484 @@ kernel void kernel_gemv_noshuffle_q4_k_f32( } } + +// --- Fused gate+up GEMV + GLU epilogue (FFN) ------------------------------------ +// Folds the FFN's two decode GEMVs (ffn_gate, ffn_up) and the following GLU into a +// SINGLE dispatch: {MUL_MAT(Wg,x), MUL_MAT(Wu,x), GLU}. Both matmuls share the same +// activation x (ffn_norm), so the activation image read is issued ONCE per K-block +// and reused for the gate and up dot products (the per-op path re-reads it twice and +// also materializes the two full ffn-wide intermediates to global, which the GLU +// then re-reads). The gate/up partial sums are accumulated in the SAME per-fiber +// order and reduced in the SAME cross-subgroup order as the standalone GEMV, and the +// GLU formula is the exact scalar expression from kernels/glu.cl, so the output is +// BYTE-IDENTICAL to the per-op matmul+matmul+glu path -> safe to default on. +// glu_op: REGLU=0, GEGLU=1, SWIGLU=2, GEGLU_ERF=4, GEGLU_QUICK=5 (ggml_glu_op). +// Weights: src0g_* = gate (= GLU src[0]); src0u_* = up (= GLU src[1]). +#define GLU_GEGLU_COEF_A 0.044715f +#define GLU_SQRT_2_OVER_PI 0.79788456080286535587989211986876f +#define GLU_SQRT_2_INV 0.70710678118654752440084436210484f +#define GLU_QUICK_COEF -1.702f + +inline float glu_apply(int glu_op, float g, float u) { + float act; + if (glu_op == 1) { // GEGLU (tanh-approx gelu) + act = 0.5f*g*(1.0f + tanh(GLU_SQRT_2_OVER_PI*g*(1.0f + GLU_GEGLU_COEF_A*g*g))); + } else if (glu_op == 2) { // SWIGLU (silu) + act = g / (1.0f + exp(-g)); + } else if (glu_op == 0) { // REGLU + return g*u*(g > 0.0f); + } else if (glu_op == 4) { // GEGLU_ERF + act = 0.5f*g*(1.0f + erf(g*GLU_SQRT_2_INV)); + } else { // GEGLU_QUICK (glu_op == 5) + act = g*(1.0f/(1.0f + exp(GLU_QUICK_COEF*g))); + } + return act*u; +} + +#ifdef ADRENO_GPU +REQD_SUBGROUP_SIZE_64 +#endif +kernel void kernel_gemv_noshuffle_q4_k_f32_glu( + read_only image1d_buffer_t src0g_q, + global half2 * src0g_d, + global half2 * src0g_m, + global uchar * src0g_s, + read_only image1d_buffer_t src0u_q, + global half2 * src0u_d, + global half2 * src0u_m, + global uchar * src0u_s, + read_only image1d_buffer_t src1, + global float * dst, + ulong offsetd, + int ne00, + int ne01, + int glu_op, + uchar mask_d6, + uchar mask_d4, + uchar mask_hi2) +{ + uint groupId = get_local_id(1); + uint gid = get_global_id(0); + ushort slid = get_sub_group_local_id(); + uint nsg = get_local_size(1); + + uint K = ne00; + uint M = ne01; + + uint LINE_STRIDE_A = M / 2; + uint BLOCK_STRIDE_A = 4 * M; + + private uint4 regA; + private half2 regS, regM; + private float8 regB; + + private float2 gateSum = (float2)(0.0f); + private float2 upSum = (float2)(0.0f); + + // Two SEQUENTIAL K-loops (gate fully, then up). Keeping only one weight's + // working set live at a time holds the kernel's register footprint at ~the + // base single-weight GEMV's, so its max WG stays 1024 (16 subgroups) and the + // per-subgroup K-split matches the standalone wide GEMV exactly -> the gate + // and up partial sums are BYTE-IDENTICAL to the per-op path. The macro body + // is the base kernel's inner loop verbatim, parameterized by weight source. +#define Q4K_GLU_LOOP(SUM, Q, DD, MM, SS) \ + for (uint k = groupId; k < (K / 32); k += nsg) { \ + uint sb = k / 8; \ + uint j = k % 8; \ + half2 d = DD[gid + sb * LINE_STRIDE_A]; \ + half2 dm = MM[gid + sb * LINE_STRIDE_A]; \ + global const uchar * sc0 = SS + sb * 12 * M + 2 * gid; \ + global const uchar * sc1 = sc0 + 1; \ + uchar sv0, mn0, sv1, mn1; \ + get_scale_min_k4(j, sc0, M, &sv0, &mn0, mask_d6, mask_d4, mask_hi2); \ + get_scale_min_k4(j, sc1, M, &sv1, &mn1, mask_d6, mask_d4, mask_hi2); \ + regS = convert_half2(convert_float2(d) * convert_float2((uchar2)(sv0, sv1))); \ + regM = convert_half2(convert_float2(dm) * convert_float2((uchar2)(mn0, mn1))); \ + if (slid < 4) { \ + regB.s0123 = read_imagef(src1, (slid * 2 + k * 8)); \ + regB.s4567 = read_imagef(src1, (1 + slid * 2 + k * 8)); \ + } \ + regA.s0 = read_imageui(Q, (gid + k * BLOCK_STRIDE_A + LINE_STRIDE_A * 0)).x; \ + regA.s1 = read_imageui(Q, (gid + k * BLOCK_STRIDE_A + LINE_STRIDE_A * 1)).x; \ + regA.s2 = read_imageui(Q, (gid + k * BLOCK_STRIDE_A + LINE_STRIDE_A * 2)).x; \ + regA.s3 = read_imageui(Q, (gid + k * BLOCK_STRIDE_A + LINE_STRIDE_A * 3)).x; \ + DEQ_HI(SUM, as_ushort8(regA), regS, regM, regB); \ + regA.s0 = read_imageui(Q, (gid + k * BLOCK_STRIDE_A + LINE_STRIDE_A * 4)).x; \ + regA.s1 = read_imageui(Q, (gid + k * BLOCK_STRIDE_A + LINE_STRIDE_A * 5)).x; \ + regA.s2 = read_imageui(Q, (gid + k * BLOCK_STRIDE_A + LINE_STRIDE_A * 6)).x; \ + regA.s3 = read_imageui(Q, (gid + k * BLOCK_STRIDE_A + LINE_STRIDE_A * 7)).x; \ + DEQ_LO(SUM, as_ushort8(regA), regS, regM, regB); \ + } + +#ifdef VECTOR_SUB_GROUP_BROADCAST +#define DEQ_HI dequantizeBlockAccum_ns_sgbroadcast_8_hi +#define DEQ_LO dequantizeBlockAccum_ns_sgbroadcast_8_lo +#else +#define DEQ_HI dequantizeBlockAccum_ns_sgbroadcast_1_hi +#define DEQ_LO dequantizeBlockAccum_ns_sgbroadcast_1_lo +#endif + + Q4K_GLU_LOOP(gateSum, src0g_q, src0g_d, src0g_m, src0g_s) + Q4K_GLU_LOOP(upSum, src0u_q, src0u_d, src0u_m, src0u_s) + +#undef DEQ_HI +#undef DEQ_LO +#undef Q4K_GLU_LOOP + + // Cross-subgroup reduction in local memory. Packs gate (xy) + up (zw) into a + // float4 so both reduce in one pass; summation order matches the base GEMV's + // per-channel loop -> byte-identical partial sums. + local float4 reduceLM[SUBGROUP_SIZE * 15]; + if (groupId > 0) { + reduceLM[SUBGROUP_SIZE * (groupId - 1) + slid] = (float4)(gateSum, upSum); + } + barrier(CLK_LOCAL_MEM_FENCE); + if (groupId == 0) { + for (uint i = 0; i < nsg - 1; ++i) { + float4 p = reduceLM[SUBGROUP_SIZE * i + slid]; + gateSum += p.xy; + upSum += p.zw; + } + dst = (global float*)((global char*)dst + offsetd); + dst[gid * 2 + 0] = glu_apply(glu_op, gateSum.s0, upSum.s0); + dst[gid * 2 + 1] = glu_apply(glu_op, gateSum.s1, upSum.s1); + } +} + +// --- Split-K-across-workgroups decode GEMV (small-M projections) ---------------- +// A single-token GEMV makes only ceil(M/2/64) workgroups; a WG runs on one Adreno +// compute unit, so for small M (Kcur/Vcur, M=512 -> 4 WGs) most of the 16 CUs sit +// idle and the matmul is bandwidth-starved even with a wide intra-WG K-split. This +// variant adds a SECOND grid dimension of `ksplit` workgroups that each reduce a +// disjoint slice of K and write a per-slice partial; kernel_gemv_splitk_reduce_f32 +// then sums the partials into dst. Identical math/layout to the base kernel +// (physical block stride 4*M, get_scale_min_k4) -> coherent. Gated host-side to +// M<=1024 (M>=2048 +// already fills the CUs and the extra reduce dispatch only hurts). +#ifdef ADRENO_GPU +REQD_SUBGROUP_SIZE_64 +#endif +kernel void kernel_gemv_noshuffle_q4_k_f32_splitk( + read_only image1d_buffer_t src0_q, + global half2 * src0_d, + global half2 * src0_m, + global uchar * src0_s, + read_only image1d_buffer_t src1, + global float * partial, // [ksplit * M], slice-major + int ne00, + int ne01, + uchar mask_d6, + uchar mask_d4, + uchar mask_hi2) +{ + uint groupId = get_local_id(1); + uint gid = get_global_id(0); + ushort slid = get_sub_group_local_id(); + uint nsg = get_local_size(1); + uint ksplit = get_num_groups(1); + uint kslice = get_group_id(1); + + uint K = ne00; + uint M = ne01; + uint LINE_STRIDE_A = M / 2; + uint BLOCK_STRIDE_A = 4 * M; // physical, independent of the K-split + + private uint4 regA; + private half2 regS, regM; + private float8 regB; + private float2 totalSum = (float2)(0.0f); + + // each (kslice, subgroup) pair owns a disjoint set of K-blocks + for (uint k = kslice * nsg + groupId; k < (K / 32); k += ksplit * nsg) { + uint sb = k / 8; + uint j = k % 8; + half2 d = src0_d[gid + sb * LINE_STRIDE_A]; + half2 dm = src0_m[gid + sb * LINE_STRIDE_A]; + global const uchar * sc0 = src0_s + sb * 12 * M + 2 * gid; + global const uchar * sc1 = sc0 + 1; + uchar sv0, mn0, sv1, mn1; + get_scale_min_k4(j, sc0, M, &sv0, &mn0, mask_d6, mask_d4, mask_hi2); + get_scale_min_k4(j, sc1, M, &sv1, &mn1, mask_d6, mask_d4, mask_hi2); + regS = convert_half2(convert_float2(d) * convert_float2((uchar2)(sv0, sv1))); + regM = convert_half2(convert_float2(dm) * convert_float2((uchar2)(mn0, mn1))); + if (slid < 4) { + regB.s0123 = read_imagef(src1, (slid * 2 + k * 8)); + regB.s4567 = read_imagef(src1, (1 + slid * 2 + k * 8)); + } + regA.s0 = read_imageui(src0_q, (gid + k * BLOCK_STRIDE_A + LINE_STRIDE_A * 0)).x; + regA.s1 = read_imageui(src0_q, (gid + k * BLOCK_STRIDE_A + LINE_STRIDE_A * 1)).x; + regA.s2 = read_imageui(src0_q, (gid + k * BLOCK_STRIDE_A + LINE_STRIDE_A * 2)).x; + regA.s3 = read_imageui(src0_q, (gid + k * BLOCK_STRIDE_A + LINE_STRIDE_A * 3)).x; +#ifdef VECTOR_SUB_GROUP_BROADCAST + dequantizeBlockAccum_ns_sgbroadcast_8_hi(totalSum, as_ushort8(regA), regS, regM, regB); +#else + dequantizeBlockAccum_ns_sgbroadcast_1_hi(totalSum, as_ushort8(regA), regS, regM, regB); +#endif + regA.s0 = read_imageui(src0_q, (gid + k * BLOCK_STRIDE_A + LINE_STRIDE_A * 4)).x; + regA.s1 = read_imageui(src0_q, (gid + k * BLOCK_STRIDE_A + LINE_STRIDE_A * 5)).x; + regA.s2 = read_imageui(src0_q, (gid + k * BLOCK_STRIDE_A + LINE_STRIDE_A * 6)).x; + regA.s3 = read_imageui(src0_q, (gid + k * BLOCK_STRIDE_A + LINE_STRIDE_A * 7)).x; +#ifdef VECTOR_SUB_GROUP_BROADCAST + dequantizeBlockAccum_ns_sgbroadcast_8_lo(totalSum, as_ushort8(regA), regS, regM, regB); +#else + dequantizeBlockAccum_ns_sgbroadcast_1_lo(totalSum, as_ushort8(regA), regS, regM, regB); +#endif + } + + local float2 reduceLM[SUBGROUP_SIZE * 15]; + if (groupId > 0) { + reduceLM[SUBGROUP_SIZE * (groupId - 1) + slid] = totalSum; + } + barrier(CLK_LOCAL_MEM_FENCE); + if (groupId == 0) { + for (uint i = 0; i < nsg - 1; ++i) { + totalSum += reduceLM[SUBGROUP_SIZE * i + slid]; + } + vstore2(totalSum, 0, &(partial[kslice * M + gid * 2])); + } +} + +// Sum the per-slice partials [ksplit * M] into dst[M]; applies the dst byte offset. +kernel void kernel_gemv_splitk_reduce_f32( + global float * partial, + global float * dst, + ulong offsetd, + int ne01, // M + int ksplit) +{ + uint r = get_global_id(0); + if (r >= (uint)ne01) return; + float acc = 0.0f; + for (uint s = 0; s < (uint)ksplit; ++s) { + acc += partial[s * (uint)ne01 + r]; + } + dst = (global float*)((global char*)dst + offsetd); + dst[r] = acc; +} + + +// --- Dequant-once macros for the mc3 verify GEMV (Q4K_MC3_DEQUANT_ONCE) --- +// The inline dequantizeBlockAccum_* macros recompute the dequantized weight +// ((code & mask)>>shift)*scale - minv ONCE PER COLUMN (3x), and the flat +// 32-FMA unroll spills ~430 B of temporaries. These macros split the work: +// DEQUANT_Q4K_BLOCK computes the 16 weights/row of one 32-block ONCE into a +// half2[] (row0 in .s0, row1 in .s1) — stored as half, the exact type the +// inline expression yields (int*half-half), so no extra rounding. MAC_Q4K_BLOCK +// then accumulates them against a column's broadcast activation in the SAME +// per-accumulator order as the inline macro. Each weight value and each +// accumulator's add-chain is bit-for-bit identical => byte-identical output, +// while the dequant ALU drops 3x->1x and the live set shrinks. Requires the +// Qualcomm vector sub_group_broadcast (float8); enabled opt-in on Adreno. +#define DEQ_Q4K_HALF2(b0, b1, msk, sh, scale, minv) \ + (half2)( ((b0 & msk) >> sh) * scale.s0 - minv.s0, \ + ((b1 & msk) >> sh) * scale.s1 - minv.s1 ) + +#define DEQUANT_Q4K_BLOCK(wq, bits, scale, minv) \ + wq[0] = DEQ_Q4K_HALF2(bits.s0, bits.s1, 0x000F, 0, scale, minv); \ + wq[1] = DEQ_Q4K_HALF2(bits.s0, bits.s1, 0x00F0, 4, scale, minv); \ + wq[2] = DEQ_Q4K_HALF2(bits.s0, bits.s1, 0x0F00, 8, scale, minv); \ + wq[3] = DEQ_Q4K_HALF2(bits.s0, bits.s1, 0xF000, 12, scale, minv); \ + wq[4] = DEQ_Q4K_HALF2(bits.s2, bits.s3, 0x000F, 0, scale, minv); \ + wq[5] = DEQ_Q4K_HALF2(bits.s2, bits.s3, 0x00F0, 4, scale, minv); \ + wq[6] = DEQ_Q4K_HALF2(bits.s2, bits.s3, 0x0F00, 8, scale, minv); \ + wq[7] = DEQ_Q4K_HALF2(bits.s2, bits.s3, 0xF000, 12, scale, minv); \ + wq[8] = DEQ_Q4K_HALF2(bits.s4, bits.s5, 0x000F, 0, scale, minv); \ + wq[9] = DEQ_Q4K_HALF2(bits.s4, bits.s5, 0x00F0, 4, scale, minv); \ + wq[10] = DEQ_Q4K_HALF2(bits.s4, bits.s5, 0x0F00, 8, scale, minv); \ + wq[11] = DEQ_Q4K_HALF2(bits.s4, bits.s5, 0xF000, 12, scale, minv); \ + wq[12] = DEQ_Q4K_HALF2(bits.s6, bits.s7, 0x000F, 0, scale, minv); \ + wq[13] = DEQ_Q4K_HALF2(bits.s6, bits.s7, 0x00F0, 4, scale, minv); \ + wq[14] = DEQ_Q4K_HALF2(bits.s6, bits.s7, 0x0F00, 8, scale, minv); \ + wq[15] = DEQ_Q4K_HALF2(bits.s6, bits.s7, 0xF000, 12, scale, minv); + +// ln0/ln1 = the two source lanes whose activation float8 this block consumes +// (0,1 for the hi block, 2,3 for the lo block — matching the inline _hi/_lo). +#define MAC_Q4K_BLOCK(ts, wq, y, ln0, ln1) { \ + float8 sy = sub_group_broadcast(y, ln0); \ + ts.s0 += wq[0].s0*sy.s0; ts.s0 += wq[1].s0*sy.s1; ts.s0 += wq[2].s0*sy.s2; ts.s0 += wq[3].s0*sy.s3; \ + ts.s0 += wq[4].s0*sy.s4; ts.s0 += wq[5].s0*sy.s5; ts.s0 += wq[6].s0*sy.s6; ts.s0 += wq[7].s0*sy.s7; \ + ts.s1 += wq[0].s1*sy.s0; ts.s1 += wq[1].s1*sy.s1; ts.s1 += wq[2].s1*sy.s2; ts.s1 += wq[3].s1*sy.s3; \ + ts.s1 += wq[4].s1*sy.s4; ts.s1 += wq[5].s1*sy.s5; ts.s1 += wq[6].s1*sy.s6; ts.s1 += wq[7].s1*sy.s7; \ + sy = sub_group_broadcast(y, ln1); \ + ts.s0 += wq[8].s0*sy.s0; ts.s0 += wq[9].s0*sy.s1; ts.s0 += wq[10].s0*sy.s2; ts.s0 += wq[11].s0*sy.s3; \ + ts.s0 += wq[12].s0*sy.s4; ts.s0 += wq[13].s0*sy.s5; ts.s0 += wq[14].s0*sy.s6; ts.s0 += wq[15].s0*sy.s7; \ + ts.s1 += wq[8].s1*sy.s0; ts.s1 += wq[9].s1*sy.s1; ts.s1 += wq[10].s1*sy.s2; ts.s1 += wq[11].s1*sy.s3; \ + ts.s1 += wq[12].s1*sy.s4; ts.s1 += wq[13].s1*sy.s5; ts.s1 += wq[14].s1*sy.s6; ts.s1 += wq[15].s1*sy.s7; \ +} + +// Multi-column (N=3) variant of the q4_K decode GEMV, for the speculative / +// MTP verify batch (ne1=3 = 2 drafts + 1 bonus). Stays on the efficient GEMV +// path (subgroup-broadcast activation, NSUBGROUPS K-split) instead of the +// transposed-GEMM dead-zone path. Each K-block's weights (regA_hi/regA_lo) are +// loaded ONCE and reused across all 3 activation columns — same weight traffic +// as one decode, ~3x the (cheap) dequant ALU. Per-column accumulation is +// independent and identical to 3 standalone GEMVs => byte-identical, so it does +// NOT perturb the lm_head logits / spec accept rate. +#ifdef ADRENO_GPU +REQD_SUBGROUP_SIZE_64 +#endif +kernel void kernel_gemv_noshuffle_q4_k_f32_mc3( + read_only image1d_buffer_t src0_q, + global half2 * src0_d, + global half2 * src0_m, + global uchar * src0_s, + read_only image1d_buffer_t src1, + global float * dst, + ulong offsetd, + int ne00, + int ne01, + uchar mask_d6, + uchar mask_d4, + uchar mask_hi2) +{ + uint groupId = get_local_id(1); + uint gid = get_global_id(0); + ushort slid = get_sub_group_local_id(); + + uint K = ne00; + uint M = ne01; + + uint LINE_STRIDE_A = M / 2; + uint BLOCK_STRIDE_A = NSUBGROUPS * M; + uint COL_STRIDE = K / 4; // float4 pixels per activation column + + private uint4 regA_hi, regA_lo; + private half2 regS, regM; + private float8 regB; + + private float2 ts0 = (float2)(0.0f); + private float2 ts1 = (float2)(0.0f); + private float2 ts2 = (float2)(0.0f); + +#ifdef Q4K_MC3_DEQUANT_LDS + // One 16-half2 block buffer per WI (reused hi->lo): forces the dequantized + // weights into LDS instead of private arrays (which spill to slow global on + // Adreno). 64*NSUBGROUPS WIs * 16 half2 = 16 KB; each WI owns its own slot + // range (flat*16) -> no cross-lane sharing, no barrier needed. + local half2 wstage[SUBGROUP_SIZE * NSUBGROUPS * 16]; + local half2 * ws = wstage + (groupId * SUBGROUP_SIZE + slid) * 16; +#endif + + for (uint k = groupId; k < (K / 32); k += NSUBGROUPS) { + uint sb = k / 8; + uint j = k % 8; + + half2 d = src0_d[gid + sb * LINE_STRIDE_A]; + half2 dm = src0_m[gid + sb * LINE_STRIDE_A]; + + global const uchar * sc0 = src0_s + sb * 12 * M + 2 * gid; + global const uchar * sc1 = sc0 + 1; + + uchar sv0, mn0, sv1, mn1; + get_scale_min_k4(j, sc0, M, &sv0, &mn0, mask_d6, mask_d4, mask_hi2); + get_scale_min_k4(j, sc1, M, &sv1, &mn1, mask_d6, mask_d4, mask_hi2); + + regS = convert_half2(convert_float2(d) * convert_float2((uchar2)(sv0, sv1))); + regM = convert_half2(convert_float2(dm) * convert_float2((uchar2)(mn0, mn1))); + + // weights loaded ONCE, reused across the 3 columns + regA_hi.s0 = read_imageui(src0_q, (gid + k * BLOCK_STRIDE_A + LINE_STRIDE_A * 0)).x; + regA_hi.s1 = read_imageui(src0_q, (gid + k * BLOCK_STRIDE_A + LINE_STRIDE_A * 1)).x; + regA_hi.s2 = read_imageui(src0_q, (gid + k * BLOCK_STRIDE_A + LINE_STRIDE_A * 2)).x; + regA_hi.s3 = read_imageui(src0_q, (gid + k * BLOCK_STRIDE_A + LINE_STRIDE_A * 3)).x; + regA_lo.s0 = read_imageui(src0_q, (gid + k * BLOCK_STRIDE_A + LINE_STRIDE_A * 4)).x; + regA_lo.s1 = read_imageui(src0_q, (gid + k * BLOCK_STRIDE_A + LINE_STRIDE_A * 5)).x; + regA_lo.s2 = read_imageui(src0_q, (gid + k * BLOCK_STRIDE_A + LINE_STRIDE_A * 6)).x; + regA_lo.s3 = read_imageui(src0_q, (gid + k * BLOCK_STRIDE_A + LINE_STRIDE_A * 7)).x; + +#ifdef Q4K_MC3_DEQUANT_ONCE + // Dequant the 32 weights/row (16 hi + 16 lo) ONCE into half2[] (byte- + // identical to the inline intermediate), then MAC against each column's + // activation. Drops the dequant ALU 3x->1x and the macro-temp spill. + half2 wq_hi[16], wq_lo[16]; + DEQUANT_Q4K_BLOCK(wq_hi, as_ushort8(regA_hi), regS, regM); + DEQUANT_Q4K_BLOCK(wq_lo, as_ushort8(regA_lo), regS, regM); + { if (slid < 4) { regB.s0123 = read_imagef(src1, 0*COL_STRIDE + slid*2 + k*8); + regB.s4567 = read_imagef(src1, 0*COL_STRIDE + 1 + slid*2 + k*8); } + MAC_Q4K_BLOCK(ts0, wq_hi, regB, 0, 1); MAC_Q4K_BLOCK(ts0, wq_lo, regB, 2, 3); } + { if (slid < 4) { regB.s0123 = read_imagef(src1, 1*COL_STRIDE + slid*2 + k*8); + regB.s4567 = read_imagef(src1, 1*COL_STRIDE + 1 + slid*2 + k*8); } + MAC_Q4K_BLOCK(ts1, wq_hi, regB, 0, 1); MAC_Q4K_BLOCK(ts1, wq_lo, regB, 2, 3); } + { if (slid < 4) { regB.s0123 = read_imagef(src1, 2*COL_STRIDE + slid*2 + k*8); + regB.s4567 = read_imagef(src1, 2*COL_STRIDE + 1 + slid*2 + k*8); } + MAC_Q4K_BLOCK(ts2, wq_hi, regB, 0, 1); MAC_Q4K_BLOCK(ts2, wq_lo, regB, 2, 3); } +#elif defined(Q4K_MC3_DEQUANT_LDS) + // LDS-staged dequant: dequant a 32-block ONCE into the per-WI LDS slot + // (hi pass then lo pass, overwriting), MAC each column from LDS. ts* + // receive hi-then-lo in the same order as DEQUANT_ONCE -> byte-identical. + // Activations reloaded per pass (cheap, imaged); only one regB + 0 weight + // regs live -> the weight working set lives in LDS, not spilled private. + DEQUANT_Q4K_BLOCK(ws, as_ushort8(regA_hi), regS, regM); + { if (slid < 4) { regB.s0123 = read_imagef(src1, 0*COL_STRIDE + slid*2 + k*8); + regB.s4567 = read_imagef(src1, 0*COL_STRIDE + 1 + slid*2 + k*8); } + MAC_Q4K_BLOCK(ts0, ws, regB, 0, 1); } + { if (slid < 4) { regB.s0123 = read_imagef(src1, 1*COL_STRIDE + slid*2 + k*8); + regB.s4567 = read_imagef(src1, 1*COL_STRIDE + 1 + slid*2 + k*8); } + MAC_Q4K_BLOCK(ts1, ws, regB, 0, 1); } + { if (slid < 4) { regB.s0123 = read_imagef(src1, 2*COL_STRIDE + slid*2 + k*8); + regB.s4567 = read_imagef(src1, 2*COL_STRIDE + 1 + slid*2 + k*8); } + MAC_Q4K_BLOCK(ts2, ws, regB, 0, 1); } + DEQUANT_Q4K_BLOCK(ws, as_ushort8(regA_lo), regS, regM); + { if (slid < 4) { regB.s0123 = read_imagef(src1, 0*COL_STRIDE + slid*2 + k*8); + regB.s4567 = read_imagef(src1, 0*COL_STRIDE + 1 + slid*2 + k*8); } + MAC_Q4K_BLOCK(ts0, ws, regB, 2, 3); } + { if (slid < 4) { regB.s0123 = read_imagef(src1, 1*COL_STRIDE + slid*2 + k*8); + regB.s4567 = read_imagef(src1, 1*COL_STRIDE + 1 + slid*2 + k*8); } + MAC_Q4K_BLOCK(ts1, ws, regB, 2, 3); } + { if (slid < 4) { regB.s0123 = read_imagef(src1, 2*COL_STRIDE + slid*2 + k*8); + regB.s4567 = read_imagef(src1, 2*COL_STRIDE + 1 + slid*2 + k*8); } + MAC_Q4K_BLOCK(ts2, ws, regB, 2, 3); } +#else + // Per-column: load only this column's activation (single regB live at a + // time -> 1/3 the activation register pressure vs holding all 3) then + // dequant against the shared weights. Cuts the private-mem spill. +#ifdef VECTOR_SUB_GROUP_BROADCAST + { if (slid < 4) { regB.s0123 = read_imagef(src1, 0*COL_STRIDE + slid*2 + k*8); + regB.s4567 = read_imagef(src1, 0*COL_STRIDE + 1 + slid*2 + k*8); } + dequantizeBlockAccum_ns_sgbroadcast_8_hi(ts0, as_ushort8(regA_hi), regS, regM, regB); + dequantizeBlockAccum_ns_sgbroadcast_8_lo(ts0, as_ushort8(regA_lo), regS, regM, regB); } + { if (slid < 4) { regB.s0123 = read_imagef(src1, 1*COL_STRIDE + slid*2 + k*8); + regB.s4567 = read_imagef(src1, 1*COL_STRIDE + 1 + slid*2 + k*8); } + dequantizeBlockAccum_ns_sgbroadcast_8_hi(ts1, as_ushort8(regA_hi), regS, regM, regB); + dequantizeBlockAccum_ns_sgbroadcast_8_lo(ts1, as_ushort8(regA_lo), regS, regM, regB); } + { if (slid < 4) { regB.s0123 = read_imagef(src1, 2*COL_STRIDE + slid*2 + k*8); + regB.s4567 = read_imagef(src1, 2*COL_STRIDE + 1 + slid*2 + k*8); } + dequantizeBlockAccum_ns_sgbroadcast_8_hi(ts2, as_ushort8(regA_hi), regS, regM, regB); + dequantizeBlockAccum_ns_sgbroadcast_8_lo(ts2, as_ushort8(regA_lo), regS, regM, regB); } +#else + { if (slid < 4) { regB.s0123 = read_imagef(src1, 0*COL_STRIDE + slid*2 + k*8); + regB.s4567 = read_imagef(src1, 0*COL_STRIDE + 1 + slid*2 + k*8); } + dequantizeBlockAccum_ns_sgbroadcast_1_hi(ts0, as_ushort8(regA_hi), regS, regM, regB); + dequantizeBlockAccum_ns_sgbroadcast_1_lo(ts0, as_ushort8(regA_lo), regS, regM, regB); } + { if (slid < 4) { regB.s0123 = read_imagef(src1, 1*COL_STRIDE + slid*2 + k*8); + regB.s4567 = read_imagef(src1, 1*COL_STRIDE + 1 + slid*2 + k*8); } + dequantizeBlockAccum_ns_sgbroadcast_1_hi(ts1, as_ushort8(regA_hi), regS, regM, regB); + dequantizeBlockAccum_ns_sgbroadcast_1_lo(ts1, as_ushort8(regA_lo), regS, regM, regB); } + { if (slid < 4) { regB.s0123 = read_imagef(src1, 2*COL_STRIDE + slid*2 + k*8); + regB.s4567 = read_imagef(src1, 2*COL_STRIDE + 1 + slid*2 + k*8); } + dequantizeBlockAccum_ns_sgbroadcast_1_hi(ts2, as_ushort8(regA_hi), regS, regM, regB); + dequantizeBlockAccum_ns_sgbroadcast_1_lo(ts2, as_ushort8(regA_lo), regS, regM, regB); } +#endif +#endif // Q4K_MC3_DEQUANT_ONCE + } + + // cross-subgroup reduce: pack the 3 columns' float2 into a float8 (6 used). + local float8 reduceLM[SUBGROUP_SIZE * 3]; + float8 acc = (float8)(ts0.s0, ts0.s1, ts1.s0, ts1.s1, ts2.s0, ts2.s1, 0.0f, 0.0f); + if (groupId == 1) { reduceLM[SUBGROUP_SIZE * 0 + slid] = acc; } + if (groupId == 2) { reduceLM[SUBGROUP_SIZE * 1 + slid] = acc; } + if (groupId == 3) { reduceLM[SUBGROUP_SIZE * 2 + slid] = acc; } + + barrier(CLK_LOCAL_MEM_FENCE); + + if (groupId == 0) { + acc += reduceLM[SUBGROUP_SIZE * 0 + slid]; + acc += reduceLM[SUBGROUP_SIZE * 1 + slid]; + acc += reduceLM[SUBGROUP_SIZE * 2 + slid]; + dst = (global float*)((global char*)dst + offsetd); + // dst is column-major [M rows x 3 cols]: (row, col) at col*M + row + vstore2((float2)(acc.s0, acc.s1), 0, &(dst[0 * M + gid * 2])); + vstore2((float2)(acc.s2, acc.s3), 0, &(dst[1 * M + gid * 2])); + vstore2((float2)(acc.s4, acc.s5), 0, &(dst[2 * M + gid * 2])); + } +} diff --git a/ggml/src/ggml-opencl/kernels/gemv_noshuffle_q4_k_f32_o4.cl b/ggml/src/ggml-opencl/kernels/gemv_noshuffle_q4_k_f32_o4.cl new file mode 100644 index 00000000000..02916bb91ff --- /dev/null +++ b/ggml/src/ggml-opencl/kernels/gemv_noshuffle_q4_k_f32_o4.cl @@ -0,0 +1,349 @@ +#pragma OPENCL EXTENSION cl_khr_fp16 : enable +#pragma OPENCL EXTENSION cl_khr_subgroups : enable + +#ifdef cl_qcom_reqd_sub_group_size +#pragma OPENCL EXTENSION cl_qcom_reqd_sub_group_size : enable +#define ADRENO_GPU 1 +#define REQD_SUBGROUP_SIZE_64 __attribute__((qcom_reqd_sub_group_size("half"))) +#endif + +#define QK_K 256 +#define NSUBGROUPS 4 +#define SUBGROUP_SIZE 64 + +// scales are transposed: consecutive codes of a row are `stride` apart +inline void get_scale_min_k4( + int j, + global const uchar * q, + uint stride, + uchar * d, + uchar * m, + uchar mask_d6, + uchar mask_d4, + uchar mask_hi2 +) { + if (j < 4) { + *d = q[j*stride] & mask_d6; + *m = q[(j+4)*stride] & mask_d6; + } else { + *d = (q[(j+4)*stride] & mask_d4) | ((q[(j-4)*stride] & mask_hi2) >> 2); + *m = ((q[(j+4)*stride] >> 4) & mask_d4) | ((q[j*stride] & mask_hi2) >> 2); + } +} + +#define dequantizeBlockAccum_ns_sgbroadcast_1_hi(total_sums, bits4, scale, minv, y) \ + float shared_y; \ + shared_y = sub_group_broadcast(y.s0, 0); \ + total_sums.s0 += ((bits4.s0 & 0x000F) * scale.s0 - minv.s0) * shared_y; \ + total_sums.s1 += ((bits4.s1 & 0x000F) * scale.s1 - minv.s1) * shared_y; \ + shared_y = sub_group_broadcast(y.s1, 0); \ + total_sums.s0 += (((bits4.s0 & 0x00F0) >> 4) * scale.s0 - minv.s0) * shared_y; \ + total_sums.s1 += (((bits4.s1 & 0x00F0) >> 4) * scale.s1 - minv.s1) * shared_y; \ + shared_y = sub_group_broadcast(y.s2, 0); \ + total_sums.s0 += (((bits4.s0 & 0x0F00) >> 8) * scale.s0 - minv.s0) * shared_y; \ + total_sums.s1 += (((bits4.s1 & 0x0F00) >> 8) * scale.s1 - minv.s1) * shared_y; \ + shared_y = sub_group_broadcast(y.s3, 0); \ + total_sums.s0 += (((bits4.s0 & 0xF000) >> 12) * scale.s0 - minv.s0) * shared_y; \ + total_sums.s1 += (((bits4.s1 & 0xF000) >> 12) * scale.s1 - minv.s1) * shared_y; \ + shared_y = sub_group_broadcast(y.s4, 0); \ + total_sums.s0 += ((bits4.s2 & 0x000F) * scale.s0 - minv.s0) * shared_y; \ + total_sums.s1 += ((bits4.s3 & 0x000F) * scale.s1 - minv.s1) * shared_y; \ + shared_y = sub_group_broadcast(y.s5, 0); \ + total_sums.s0 += (((bits4.s2 & 0x00F0) >> 4) * scale.s0 - minv.s0) * shared_y; \ + total_sums.s1 += (((bits4.s3 & 0x00F0) >> 4) * scale.s1 - minv.s1) * shared_y; \ + shared_y = sub_group_broadcast(y.s6, 0); \ + total_sums.s0 += (((bits4.s2 & 0x0F00) >> 8) * scale.s0 - minv.s0) * shared_y; \ + total_sums.s1 += (((bits4.s3 & 0x0F00) >> 8) * scale.s1 - minv.s1) * shared_y; \ + shared_y = sub_group_broadcast(y.s7, 0); \ + total_sums.s0 += (((bits4.s2 & 0xF000) >> 12) * scale.s0 - minv.s0) * shared_y; \ + total_sums.s1 += (((bits4.s3 & 0xF000) >> 12) * scale.s1 - minv.s1) * shared_y; \ + shared_y = sub_group_broadcast(y.s0, 1); \ + total_sums.s0 += ((bits4.s4 & 0x000F) * scale.s0 - minv.s0) * shared_y; \ + total_sums.s1 += ((bits4.s5 & 0x000F) * scale.s1 - minv.s1) * shared_y; \ + shared_y = sub_group_broadcast(y.s1, 1); \ + total_sums.s0 += (((bits4.s4 & 0x00F0) >> 4) * scale.s0 - minv.s0) * shared_y; \ + total_sums.s1 += (((bits4.s5 & 0x00F0) >> 4) * scale.s1 - minv.s1) * shared_y; \ + shared_y = sub_group_broadcast(y.s2, 1); \ + total_sums.s0 += (((bits4.s4 & 0x0F00) >> 8) * scale.s0 - minv.s0) * shared_y; \ + total_sums.s1 += (((bits4.s5 & 0x0F00) >> 8) * scale.s1 - minv.s1) * shared_y; \ + shared_y = sub_group_broadcast(y.s3, 1); \ + total_sums.s0 += (((bits4.s4 & 0xF000) >> 12) * scale.s0 - minv.s0) * shared_y; \ + total_sums.s1 += (((bits4.s5 & 0xF000) >> 12) * scale.s1 - minv.s1) * shared_y; \ + shared_y = sub_group_broadcast(y.s4, 1); \ + total_sums.s0 += ((bits4.s6 & 0x000F) * scale.s0 - minv.s0) * shared_y; \ + total_sums.s1 += ((bits4.s7 & 0x000F) * scale.s1 - minv.s1) * shared_y; \ + shared_y = sub_group_broadcast(y.s5, 1); \ + total_sums.s0 += (((bits4.s6 & 0x00F0) >> 4) * scale.s0 - minv.s0) * shared_y; \ + total_sums.s1 += (((bits4.s7 & 0x00F0) >> 4) * scale.s1 - minv.s1) * shared_y; \ + shared_y = sub_group_broadcast(y.s6, 1); \ + total_sums.s0 += (((bits4.s6 & 0x0F00) >> 8) * scale.s0 - minv.s0) * shared_y; \ + total_sums.s1 += (((bits4.s7 & 0x0F00) >> 8) * scale.s1 - minv.s1) * shared_y; \ + shared_y = sub_group_broadcast(y.s7, 1); \ + total_sums.s0 += (((bits4.s6 & 0xF000) >> 12) * scale.s0 - minv.s0) * shared_y; \ + total_sums.s1 += (((bits4.s7 & 0xF000) >> 12) * scale.s1 - minv.s1) * shared_y; \ + + +#define dequantizeBlockAccum_ns_sgbroadcast_1_lo(total_sums, bits4, scale, minv, y) \ + shared_y = sub_group_broadcast(y.s0, 2); \ + total_sums.s0 += ((bits4.s0 & 0x000F) * scale.s0 - minv.s0) * shared_y; \ + total_sums.s1 += ((bits4.s1 & 0x000F) * scale.s1 - minv.s1) * shared_y; \ + shared_y = sub_group_broadcast(y.s1, 2); \ + total_sums.s0 += (((bits4.s0 & 0x00F0) >> 4) * scale.s0 - minv.s0) * shared_y; \ + total_sums.s1 += (((bits4.s1 & 0x00F0) >> 4) * scale.s1 - minv.s1) * shared_y; \ + shared_y = sub_group_broadcast(y.s2, 2); \ + total_sums.s0 += (((bits4.s0 & 0x0F00) >> 8) * scale.s0 - minv.s0) * shared_y; \ + total_sums.s1 += (((bits4.s1 & 0x0F00) >> 8) * scale.s1 - minv.s1) * shared_y; \ + shared_y = sub_group_broadcast(y.s3, 2); \ + total_sums.s0 += (((bits4.s0 & 0xF000) >> 12) * scale.s0 - minv.s0) * shared_y; \ + total_sums.s1 += (((bits4.s1 & 0xF000) >> 12) * scale.s1 - minv.s1) * shared_y; \ + shared_y = sub_group_broadcast(y.s4, 2); \ + total_sums.s0 += ((bits4.s2 & 0x000F) * scale.s0 - minv.s0) * shared_y; \ + total_sums.s1 += ((bits4.s3 & 0x000F) * scale.s1 - minv.s1) * shared_y; \ + shared_y = sub_group_broadcast(y.s5, 2); \ + total_sums.s0 += (((bits4.s2 & 0x00F0) >> 4) * scale.s0 - minv.s0) * shared_y; \ + total_sums.s1 += (((bits4.s3 & 0x00F0) >> 4) * scale.s1 - minv.s1) * shared_y; \ + shared_y = sub_group_broadcast(y.s6, 2); \ + total_sums.s0 += (((bits4.s2 & 0x0F00) >> 8) * scale.s0 - minv.s0) * shared_y; \ + total_sums.s1 += (((bits4.s3 & 0x0F00) >> 8) * scale.s1 - minv.s1) * shared_y; \ + shared_y = sub_group_broadcast(y.s7, 2); \ + total_sums.s0 += (((bits4.s2 & 0xF000) >> 12) * scale.s0 - minv.s0) * shared_y; \ + total_sums.s1 += (((bits4.s3 & 0xF000) >> 12) * scale.s1 - minv.s1) * shared_y; \ + shared_y = sub_group_broadcast(y.s0, 3); \ + total_sums.s0 += ((bits4.s4 & 0x000F) * scale.s0 - minv.s0) * shared_y; \ + total_sums.s1 += ((bits4.s5 & 0x000F) * scale.s1 - minv.s1) * shared_y; \ + shared_y = sub_group_broadcast(y.s1, 3); \ + total_sums.s0 += (((bits4.s4 & 0x00F0) >> 4) * scale.s0 - minv.s0) * shared_y; \ + total_sums.s1 += (((bits4.s5 & 0x00F0) >> 4) * scale.s1 - minv.s1) * shared_y; \ + shared_y = sub_group_broadcast(y.s2, 3); \ + total_sums.s0 += (((bits4.s4 & 0x0F00) >> 8) * scale.s0 - minv.s0) * shared_y; \ + total_sums.s1 += (((bits4.s5 & 0x0F00) >> 8) * scale.s1 - minv.s1) * shared_y; \ + shared_y = sub_group_broadcast(y.s3, 3); \ + total_sums.s0 += (((bits4.s4 & 0xF000) >> 12) * scale.s0 - minv.s0) * shared_y; \ + total_sums.s1 += (((bits4.s5 & 0xF000) >> 12) * scale.s1 - minv.s1) * shared_y; \ + shared_y = sub_group_broadcast(y.s4, 3); \ + total_sums.s0 += ((bits4.s6 & 0x000F) * scale.s0 - minv.s0) * shared_y; \ + total_sums.s1 += ((bits4.s7 & 0x000F) * scale.s1 - minv.s1) * shared_y; \ + shared_y = sub_group_broadcast(y.s5, 3); \ + total_sums.s0 += (((bits4.s6 & 0x00F0) >> 4) * scale.s0 - minv.s0) * shared_y; \ + total_sums.s1 += (((bits4.s7 & 0x00F0) >> 4) * scale.s1 - minv.s1) * shared_y; \ + shared_y = sub_group_broadcast(y.s6, 3); \ + total_sums.s0 += (((bits4.s6 & 0x0F00) >> 8) * scale.s0 - minv.s0) * shared_y; \ + total_sums.s1 += (((bits4.s7 & 0x0F00) >> 8) * scale.s1 - minv.s1) * shared_y; \ + shared_y = sub_group_broadcast(y.s7, 3); \ + total_sums.s0 += (((bits4.s6 & 0xF000) >> 12) * scale.s0 - minv.s0) * shared_y; \ + total_sums.s1 += (((bits4.s7 & 0xF000) >> 12) * scale.s1 - minv.s1) * shared_y; \ + + +#define dequantizeBlockAccum_ns_sgbroadcast_8_hi(total_sums, bits4, scale, minv, y) \ + float8 shared_y; \ + shared_y = sub_group_broadcast(y, 0); \ + total_sums.s0 += ((bits4.s0 & 0x000F) * scale.s0 - minv.s0) * shared_y.s0; \ + total_sums.s0 += (((bits4.s0 & 0x00F0) >> 4) * scale.s0 - minv.s0) * shared_y.s1; \ + total_sums.s0 += (((bits4.s0 & 0x0F00) >> 8) * scale.s0 - minv.s0) * shared_y.s2; \ + total_sums.s0 += (((bits4.s0 & 0xF000) >> 12) * scale.s0 - minv.s0) * shared_y.s3; \ + total_sums.s0 += ((bits4.s2 & 0x000F) * scale.s0 - minv.s0) * shared_y.s4; \ + total_sums.s0 += (((bits4.s2 & 0x00F0) >> 4) * scale.s0 - minv.s0) * shared_y.s5; \ + total_sums.s0 += (((bits4.s2 & 0x0F00) >> 8) * scale.s0 - minv.s0) * shared_y.s6; \ + total_sums.s0 += (((bits4.s2 & 0xF000) >> 12) * scale.s0 - minv.s0) * shared_y.s7; \ + total_sums.s1 += ((bits4.s1 & 0x000F) * scale.s1 - minv.s1) * shared_y.s0; \ + total_sums.s1 += (((bits4.s1 & 0x00F0) >> 4) * scale.s1 - minv.s1) * shared_y.s1; \ + total_sums.s1 += (((bits4.s1 & 0x0F00) >> 8) * scale.s1 - minv.s1) * shared_y.s2; \ + total_sums.s1 += (((bits4.s1 & 0xF000) >> 12) * scale.s1 - minv.s1) * shared_y.s3; \ + total_sums.s1 += ((bits4.s3 & 0x000F) * scale.s1 - minv.s1) * shared_y.s4; \ + total_sums.s1 += (((bits4.s3 & 0x00F0) >> 4) * scale.s1 - minv.s1) * shared_y.s5; \ + total_sums.s1 += (((bits4.s3 & 0x0F00) >> 8) * scale.s1 - minv.s1) * shared_y.s6; \ + total_sums.s1 += (((bits4.s3 & 0xF000) >> 12) * scale.s1 - minv.s1) * shared_y.s7; \ + shared_y = sub_group_broadcast(y, 1); \ + total_sums.s0 += ((bits4.s4 & 0x000F) * scale.s0 - minv.s0) * shared_y.s0; \ + total_sums.s0 += (((bits4.s4 & 0x00F0) >> 4) * scale.s0 - minv.s0) * shared_y.s1; \ + total_sums.s0 += (((bits4.s4 & 0x0F00) >> 8) * scale.s0 - minv.s0) * shared_y.s2; \ + total_sums.s0 += (((bits4.s4 & 0xF000) >> 12) * scale.s0 - minv.s0) * shared_y.s3; \ + total_sums.s0 += ((bits4.s6 & 0x000F) * scale.s0 - minv.s0) * shared_y.s4; \ + total_sums.s0 += (((bits4.s6 & 0x00F0) >> 4) * scale.s0 - minv.s0) * shared_y.s5; \ + total_sums.s0 += (((bits4.s6 & 0x0F00) >> 8) * scale.s0 - minv.s0) * shared_y.s6; \ + total_sums.s0 += (((bits4.s6 & 0xF000) >> 12) * scale.s0 - minv.s0) * shared_y.s7; \ + total_sums.s1 += ((bits4.s5 & 0x000F) * scale.s1 - minv.s1) * shared_y.s0; \ + total_sums.s1 += (((bits4.s5 & 0x00F0) >> 4) * scale.s1 - minv.s1) * shared_y.s1; \ + total_sums.s1 += (((bits4.s5 & 0x0F00) >> 8) * scale.s1 - minv.s1) * shared_y.s2; \ + total_sums.s1 += (((bits4.s5 & 0xF000) >> 12) * scale.s1 - minv.s1) * shared_y.s3; \ + total_sums.s1 += ((bits4.s7 & 0x000F) * scale.s1 - minv.s1) * shared_y.s4; \ + total_sums.s1 += (((bits4.s7 & 0x00F0) >> 4) * scale.s1 - minv.s1) * shared_y.s5; \ + total_sums.s1 += (((bits4.s7 & 0x0F00) >> 8) * scale.s1 - minv.s1) * shared_y.s6; \ + total_sums.s1 += (((bits4.s7 & 0xF000) >> 12) * scale.s1 - minv.s1) * shared_y.s7; \ + + +#define dequantizeBlockAccum_ns_sgbroadcast_8_lo(total_sums, bits4, scale, minv, y) \ + shared_y = sub_group_broadcast(y, 2); \ + total_sums.s0 += ((bits4.s0 & 0x000F) * scale.s0 - minv.s0) * shared_y.s0; \ + total_sums.s0 += (((bits4.s0 & 0x00F0) >> 4) * scale.s0 - minv.s0) * shared_y.s1; \ + total_sums.s0 += (((bits4.s0 & 0x0F00) >> 8) * scale.s0 - minv.s0) * shared_y.s2; \ + total_sums.s0 += (((bits4.s0 & 0xF000) >> 12) * scale.s0 - minv.s0) * shared_y.s3; \ + total_sums.s0 += ((bits4.s2 & 0x000F) * scale.s0 - minv.s0) * shared_y.s4; \ + total_sums.s0 += (((bits4.s2 & 0x00F0) >> 4) * scale.s0 - minv.s0) * shared_y.s5; \ + total_sums.s0 += (((bits4.s2 & 0x0F00) >> 8) * scale.s0 - minv.s0) * shared_y.s6; \ + total_sums.s0 += (((bits4.s2 & 0xF000) >> 12) * scale.s0 - minv.s0) * shared_y.s7; \ + total_sums.s1 += ((bits4.s1 & 0x000F) * scale.s1 - minv.s1) * shared_y.s0; \ + total_sums.s1 += (((bits4.s1 & 0x00F0) >> 4) * scale.s1 - minv.s1) * shared_y.s1; \ + total_sums.s1 += (((bits4.s1 & 0x0F00) >> 8) * scale.s1 - minv.s1) * shared_y.s2; \ + total_sums.s1 += (((bits4.s1 & 0xF000) >> 12) * scale.s1 - minv.s1) * shared_y.s3; \ + total_sums.s1 += ((bits4.s3 & 0x000F) * scale.s1 - minv.s1) * shared_y.s4; \ + total_sums.s1 += (((bits4.s3 & 0x00F0) >> 4) * scale.s1 - minv.s1) * shared_y.s5; \ + total_sums.s1 += (((bits4.s3 & 0x0F00) >> 8) * scale.s1 - minv.s1) * shared_y.s6; \ + total_sums.s1 += (((bits4.s3 & 0xF000) >> 12) * scale.s1 - minv.s1) * shared_y.s7; \ + shared_y = sub_group_broadcast(y, 3); \ + total_sums.s0 += ((bits4.s4 & 0x000F) * scale.s0 - minv.s0) * shared_y.s0; \ + total_sums.s0 += (((bits4.s4 & 0x00F0) >> 4) * scale.s0 - minv.s0) * shared_y.s1; \ + total_sums.s0 += (((bits4.s4 & 0x0F00) >> 8) * scale.s0 - minv.s0) * shared_y.s2; \ + total_sums.s0 += (((bits4.s4 & 0xF000) >> 12) * scale.s0 - minv.s0) * shared_y.s3; \ + total_sums.s0 += ((bits4.s6 & 0x000F) * scale.s0 - minv.s0) * shared_y.s4; \ + total_sums.s0 += (((bits4.s6 & 0x00F0) >> 4) * scale.s0 - minv.s0) * shared_y.s5; \ + total_sums.s0 += (((bits4.s6 & 0x0F00) >> 8) * scale.s0 - minv.s0) * shared_y.s6; \ + total_sums.s0 += (((bits4.s6 & 0xF000) >> 12) * scale.s0 - minv.s0) * shared_y.s7; \ + total_sums.s1 += ((bits4.s5 & 0x000F) * scale.s1 - minv.s1) * shared_y.s0; \ + total_sums.s1 += (((bits4.s5 & 0x00F0) >> 4) * scale.s1 - minv.s1) * shared_y.s1; \ + total_sums.s1 += (((bits4.s5 & 0x0F00) >> 8) * scale.s1 - minv.s1) * shared_y.s2; \ + total_sums.s1 += (((bits4.s5 & 0xF000) >> 12) * scale.s1 - minv.s1) * shared_y.s3; \ + total_sums.s1 += ((bits4.s7 & 0x000F) * scale.s1 - minv.s1) * shared_y.s4; \ + total_sums.s1 += (((bits4.s7 & 0x00F0) >> 4) * scale.s1 - minv.s1) * shared_y.s5; \ + total_sums.s1 += (((bits4.s7 & 0x0F00) >> 8) * scale.s1 - minv.s1) * shared_y.s6; \ + total_sums.s1 += (((bits4.s7 & 0xF000) >> 12) * scale.s1 - minv.s1) * shared_y.s7; \ + +#ifdef ADRENO_GPU +REQD_SUBGROUP_SIZE_64 +#endif +kernel void kernel_gemv_noshuffle_q4_k_f32_o4( + read_only image1d_buffer_t src0_q, + global half2 * src0_d, + global half2 * src0_m, + global uchar * src0_s, + read_only image1d_buffer_t src1, + global float * dst, + ulong offsetd, + int ne00, + int ne01, + uchar mask_d6, + uchar mask_d4, + uchar mask_hi2) +{ + uint groupId = get_local_id(1); + uint gid = get_global_id(0); // 4-output quad index + ushort slid = get_sub_group_local_id(); + + // Two consecutive pair-indices (each the same access pattern the 2-output + // kernel uses); together they cover 4 consecutive output rows. + uint gid_a = gid * 2; + uint gid_b = gid * 2 + 1; + + uint K = ne00; + uint M = ne01; + + uint LINE_STRIDE_A = M / 2; + uint BLOCK_STRIDE_A = NSUBGROUPS * M; + + private uint4 regA; + private half2 regS_a, regS_b; + private half2 regM_a, regM_b; + private float8 regB; + + private float2 totalSum_a = (float2)(0.0f); + private float2 totalSum_b = (float2)(0.0f); + + for (uint k = groupId; k < (K / 32); k += NSUBGROUPS) { + uint sb = k / 8; + uint j = k % 8; + + // pair a scales/mins + half2 d_a = src0_d[gid_a + sb * LINE_STRIDE_A]; + half2 dm_a = src0_m[gid_a + sb * LINE_STRIDE_A]; + global const uchar * sc0a = src0_s + sb * 12 * M + 2 * gid_a; + global const uchar * sc1a = sc0a + 1; + uchar sv0a, mn0a, sv1a, mn1a; + get_scale_min_k4(j, sc0a, M, &sv0a, &mn0a, mask_d6, mask_d4, mask_hi2); + get_scale_min_k4(j, sc1a, M, &sv1a, &mn1a, mask_d6, mask_d4, mask_hi2); + regS_a = convert_half2(convert_float2(d_a) * convert_float2((uchar2)(sv0a, sv1a))); + regM_a = convert_half2(convert_float2(dm_a) * convert_float2((uchar2)(mn0a, mn1a))); + + // pair b scales/mins + half2 d_b = src0_d[gid_b + sb * LINE_STRIDE_A]; + half2 dm_b = src0_m[gid_b + sb * LINE_STRIDE_A]; + global const uchar * sc0b = src0_s + sb * 12 * M + 2 * gid_b; + global const uchar * sc1b = sc0b + 1; + uchar sv0b, mn0b, sv1b, mn1b; + get_scale_min_k4(j, sc0b, M, &sv0b, &mn0b, mask_d6, mask_d4, mask_hi2); + get_scale_min_k4(j, sc1b, M, &sv1b, &mn1b, mask_d6, mask_d4, mask_hi2); + regS_b = convert_half2(convert_float2(d_b) * convert_float2((uchar2)(sv0b, sv1b))); + regM_b = convert_half2(convert_float2(dm_b) * convert_float2((uchar2)(mn0b, mn1b))); + + // activation: load once, reuse for both pairs + if (slid < 4) { + regB.s0123 = read_imagef(src1, (slid * 2 + k * 8)); + regB.s4567 = read_imagef(src1, (1 + slid * 2 + k * 8)); + } + + // pair a (own block so _lo sees the shared_y declared by _hi) + { + regA.s0 = read_imageui(src0_q, (gid_a + k * BLOCK_STRIDE_A + LINE_STRIDE_A * 0)).x; + regA.s1 = read_imageui(src0_q, (gid_a + k * BLOCK_STRIDE_A + LINE_STRIDE_A * 1)).x; + regA.s2 = read_imageui(src0_q, (gid_a + k * BLOCK_STRIDE_A + LINE_STRIDE_A * 2)).x; + regA.s3 = read_imageui(src0_q, (gid_a + k * BLOCK_STRIDE_A + LINE_STRIDE_A * 3)).x; +#ifdef VECTOR_SUB_GROUP_BROADCAST + dequantizeBlockAccum_ns_sgbroadcast_8_hi(totalSum_a, as_ushort8(regA), regS_a, regM_a, regB); +#else + dequantizeBlockAccum_ns_sgbroadcast_1_hi(totalSum_a, as_ushort8(regA), regS_a, regM_a, regB); +#endif + regA.s0 = read_imageui(src0_q, (gid_a + k * BLOCK_STRIDE_A + LINE_STRIDE_A * 4)).x; + regA.s1 = read_imageui(src0_q, (gid_a + k * BLOCK_STRIDE_A + LINE_STRIDE_A * 5)).x; + regA.s2 = read_imageui(src0_q, (gid_a + k * BLOCK_STRIDE_A + LINE_STRIDE_A * 6)).x; + regA.s3 = read_imageui(src0_q, (gid_a + k * BLOCK_STRIDE_A + LINE_STRIDE_A * 7)).x; +#ifdef VECTOR_SUB_GROUP_BROADCAST + dequantizeBlockAccum_ns_sgbroadcast_8_lo(totalSum_a, as_ushort8(regA), regS_a, regM_a, regB); +#else + dequantizeBlockAccum_ns_sgbroadcast_1_lo(totalSum_a, as_ushort8(regA), regS_a, regM_a, regB); +#endif + } + + // pair b + { + regA.s0 = read_imageui(src0_q, (gid_b + k * BLOCK_STRIDE_A + LINE_STRIDE_A * 0)).x; + regA.s1 = read_imageui(src0_q, (gid_b + k * BLOCK_STRIDE_A + LINE_STRIDE_A * 1)).x; + regA.s2 = read_imageui(src0_q, (gid_b + k * BLOCK_STRIDE_A + LINE_STRIDE_A * 2)).x; + regA.s3 = read_imageui(src0_q, (gid_b + k * BLOCK_STRIDE_A + LINE_STRIDE_A * 3)).x; +#ifdef VECTOR_SUB_GROUP_BROADCAST + dequantizeBlockAccum_ns_sgbroadcast_8_hi(totalSum_b, as_ushort8(regA), regS_b, regM_b, regB); +#else + dequantizeBlockAccum_ns_sgbroadcast_1_hi(totalSum_b, as_ushort8(regA), regS_b, regM_b, regB); +#endif + regA.s0 = read_imageui(src0_q, (gid_b + k * BLOCK_STRIDE_A + LINE_STRIDE_A * 4)).x; + regA.s1 = read_imageui(src0_q, (gid_b + k * BLOCK_STRIDE_A + LINE_STRIDE_A * 5)).x; + regA.s2 = read_imageui(src0_q, (gid_b + k * BLOCK_STRIDE_A + LINE_STRIDE_A * 6)).x; + regA.s3 = read_imageui(src0_q, (gid_b + k * BLOCK_STRIDE_A + LINE_STRIDE_A * 7)).x; +#ifdef VECTOR_SUB_GROUP_BROADCAST + dequantizeBlockAccum_ns_sgbroadcast_8_lo(totalSum_b, as_ushort8(regA), regS_b, regM_b, regB); +#else + dequantizeBlockAccum_ns_sgbroadcast_1_lo(totalSum_b, as_ushort8(regA), regS_b, regM_b, regB); +#endif + } + } + + // reduce 4 outputs (a.s0, a.s1, b.s0, b.s1) across the 4 subgroups + local float4 reduceLM[SUBGROUP_SIZE * 3]; + float4 acc = (float4)(totalSum_a.s0, totalSum_a.s1, totalSum_b.s0, totalSum_b.s1); + if (groupId == 1) { reduceLM[SUBGROUP_SIZE * 0 + slid] = acc; } + if (groupId == 2) { reduceLM[SUBGROUP_SIZE * 1 + slid] = acc; } + if (groupId == 3) { reduceLM[SUBGROUP_SIZE * 2 + slid] = acc; } + + barrier(CLK_LOCAL_MEM_FENCE); + + if (groupId == 0) { + acc += reduceLM[SUBGROUP_SIZE * 0 + slid]; + acc += reduceLM[SUBGROUP_SIZE * 1 + slid]; + acc += reduceLM[SUBGROUP_SIZE * 2 + slid]; + dst = (global float*)((global char*)dst + offsetd); + // The dispatch rounds ne01/4 up to the subgroup width, so the tail + // quads past the last row must not store (they wrote 128 rows past + // dst on every ne01 % 256 == 128 vocab, e.g. 151936). + if (gid * 4 + 3 < (uint)ne01) { + vstore4(acc, 0, &(dst[gid * 4])); + } + } +} diff --git a/ggml/src/ggml-opencl/kernels/gemv_noshuffle_q4_k_f32_tiled.cl b/ggml/src/ggml-opencl/kernels/gemv_noshuffle_q4_k_f32_tiled.cl new file mode 100644 index 00000000000..929538c41d6 --- /dev/null +++ b/ggml/src/ggml-opencl/kernels/gemv_noshuffle_q4_k_f32_tiled.cl @@ -0,0 +1,118 @@ +// Tiled-wide q4_K GEMV for the long-vocab lm_head/embed (decode path). +// +// Pairs with kernel_convert_block_q4_k_tiled_ns (cvt.cl): the weights are laid +// out CANONICALLY (4-bit code in element order e in [0,256)) and TILED by 64 +// output rows so the 64-thread lane group coalesces every weight load. Both the +// pack (convert) and the unpack (here) are owned by us -> correct by +// construction vs the reference ggml q4_K dequant. Same structure as the q6_K +// tiled GEMV; the only differences are the 4-bit dequant and the q4_K +// scale/min decode (get_scale_min_k4 from the packed 12-byte block). +// +// One work-item produces one output row. WG = {64 lanes, 4 subgroups}: the 64 +// lanes cover the 64 rows of one tile (coalesced uint4 reads), the 4 subgroups +// split the K-blocks and reduce through __local at the end. Weights read from +// __global (lm_head is streamed once per token; texture cache caps it below the +// coalesced-global rate). + +#pragma OPENCL EXTENSION cl_khr_fp16 : enable + +#ifdef cl_qcom_reqd_sub_group_size +#pragma OPENCL EXTENSION cl_qcom_reqd_sub_group_size : enable +#define ADRENO_GPU 1 +#define REQD_SUBGROUP_SIZE_64 __attribute__((qcom_reqd_sub_group_size("half"))) +#endif + +#define QK_K 256 +#define NSUBGROUPS 4 +#define TILE_ROWS 64 + +// Decode one q4_K sub-block scale + min from the packed 12-byte block. +// Identical to the o4 kernel's helper (masks hard-coded: d6=0x3F, d4=0x0F, hi2=0xC0). +inline void q4k_scale_min(int j, __global const uchar * q, uchar * d, uchar * m) { + if (j < 4) { + *d = q[j] & 0x3F; + *m = q[j+4] & 0x3F; + } else { + *d = (q[j+4] & 0x0F) | ((q[j-4] & 0xC0) >> 2); + *m = ((q[j+4] >> 4) & 0x0F) | ((q[j] & 0xC0) >> 2); + } +} + +#if defined(ADRENO_GPU) +REQD_SUBGROUP_SIZE_64 +#endif +kernel void kernel_gemv_noshuffle_q4_k_f32_tiled( + __global uint4 * src0_q, // tiled: 8 uint4 granules / superblock (4-bit codes) + __global half * src0_d, // tiled: 1 half / superblock + __global half * src0_dm, // tiled: 1 half / superblock + __global uchar * src0_s, // tiled: 12 bytes / superblock (packed scales) + read_only image1d_buffer_t src1, // activation (RGBA f32) + global float * dst, + ulong offsetd, + int ne00, + int ne01 +) { + int grp = get_local_id(1); // subgroup index 0..3 (splits K) + int row = get_global_id(0); // output row along ne01 + int rt = row / TILE_ROWS; + int rit = row % TILE_ROWS; + + int nb = ne00 / QK_K; // superblocks per row + + float acc = 0.0f; + + for (int sb = grp; sb < nb; sb += NSUBGROUPS) { + int tile_blk = rt * nb + sb; // ne02 == 1 for lm_head/embed + + float dval = (float)src0_d [tile_blk * TILE_ROWS + rit]; + float dmval = (float)src0_dm[tile_blk * TILE_ROWS + rit]; + + // decode the 8 sub-block (scale, min) pairs + __global uchar * sc = src0_s + (tile_blk * TILE_ROWS + rit) * 12; + float scale[8], minv[8]; + #pragma unroll + for (int is = 0; is < 8; ++is) { + uchar sd, sm; + q4k_scale_min(is, sc, &sd, &sm); + scale[is] = dval * (float)sd; + minv[is] = dmval * (float)sm; + } + + // 32 uints of 4-bit codes (8 codes/uint), e-order + uint q[32]; + #pragma unroll + for (int g = 0; g < 8; ++g) { + uint4 v = src0_q[(tile_blk * 8 + g) * TILE_ROWS + rit]; + q[g*4+0] = v.x; q[g*4+1] = v.y; q[g*4+2] = v.z; q[g*4+3] = v.w; + } + + // dequant 256 codes in canonical e-order, MAC with activation. + int act_base = sb * 64; // activation float4 pixel base (256/4) + #pragma unroll + for (int e4 = 0; e4 < 64; ++e4) { + float4 a = read_imagef(src1, act_base + e4); + #pragma unroll + for (int t = 0; t < 4; ++t) { + int e = e4 * 4 + t; + uint code = (q[e >> 3] >> ((e & 7) * 4)) & 0xF; + int is = e >> 5; // sub-block index = e/32 + float av = (t == 0) ? a.x : (t == 1) ? a.y : (t == 2) ? a.z : a.w; + acc += ((float)code * scale[is] - minv[is]) * av; + } + } + } + + // reduce across the NSUBGROUPS subgroups (same rit, different K-subset) + local float reduce_lm[NSUBGROUPS * TILE_ROWS]; + reduce_lm[grp * TILE_ROWS + rit] = acc; + barrier(CLK_LOCAL_MEM_FENCE); + + if (grp == 0) { + float total = reduce_lm[0 * TILE_ROWS + rit] + + reduce_lm[1 * TILE_ROWS + rit] + + reduce_lm[2 * TILE_ROWS + rit] + + reduce_lm[3 * TILE_ROWS + rit]; + dst = (global float*)((global char*)dst + offsetd); + dst[row] = total; + } +} diff --git a/ggml/src/ggml-opencl/kernels/gemv_noshuffle_q5_k_f32.cl b/ggml/src/ggml-opencl/kernels/gemv_noshuffle_q5_k_f32.cl index 446f4653387..ae864b19ba9 100644 --- a/ggml/src/ggml-opencl/kernels/gemv_noshuffle_q5_k_f32.cl +++ b/ggml/src/ggml-opencl/kernels/gemv_noshuffle_q5_k_f32.cl @@ -329,3 +329,125 @@ kernel void kernel_gemv_noshuffle_q5_k_f32( if (gid * 2 + 1 < M) dst[gid * 2 + 1] = totalSum.s1; } } + +// Multi-column (N in [2..4]) variant of the q5_K decode GEMV (spec/MTP verify) = +// q4_K mc3 + the high-bit qh plane (regH). n_cols = 2..4 (drafted + bonus); routes +// the small-batch verify OFF the gemm_noshuffle_q5_k dead-zone. n_cols==3 is byte- +// identical to the original mc3 (col3 disabled, float8 slots 6/7 stay zero). +#ifdef VECTOR_SUB_GROUP_BROADCAST +#define MC_DQ5_HI dequantizeBlockAccum_ns_sgbroadcast_8_hi +#define MC_DQ5_LO dequantizeBlockAccum_ns_sgbroadcast_8_lo +#else +#define MC_DQ5_HI dequantizeBlockAccum_ns_sgbroadcast_1_hi +#define MC_DQ5_LO dequantizeBlockAccum_ns_sgbroadcast_1_lo +#endif +#define MC_COL_Q5K(ts, c) \ + { if (slid < 4) { regB.s0123 = read_imagef(src1, (c)*COL_STRIDE + slid*2 + k*8); \ + regB.s4567 = read_imagef(src1, (c)*COL_STRIDE + 1 + slid*2 + k*8); } \ + MC_DQ5_HI(ts, as_ushort8(regA_hi), as_uchar8(regH), regS, regM, regB); \ + MC_DQ5_LO(ts, as_ushort8(regA_lo), as_uchar8(regH), regS, regM, regB); } +#ifdef ADRENO_GPU +REQD_SUBGROUP_SIZE_64 +#endif +kernel void kernel_gemv_noshuffle_q5_k_f32_mc3( + read_only image1d_buffer_t src0_q, + read_only image1d_buffer_t src0_qh, + global half2 * src0_d, + global half2 * src0_m, + global uchar * src0_s, + read_only image1d_buffer_t src1, + global float * dst, + ulong offsetd, + int ne00, + int ne01, + uchar mask_d6, + uchar mask_d4, + uchar mask_hi2, + int n_cols) +{ + uint groupId = get_local_id(1); + uint gid = get_global_id(0); + ushort slid = get_sub_group_local_id(); + + uint K = ne00; + uint M = ne01; + + uint LINE_STRIDE_A = M / 2; + uint BLOCK_STRIDE_A = NSUBGROUPS * M; + uint LINE_STRIDE_A_QH = M / 2; + uint BLOCK_STRIDE_A_QH = NSUBGROUPS * M / 2; + uint scales_per_row = (K / QK_K) * 12; + uint COL_STRIDE = K / 4; // float4 pixels per activation column + + private uint4 regA_hi, regA_lo; + private ushort4 regH; + private half2 regS, regM; + private float8 regB; + + private float2 ts0 = (float2)(0.0f); + private float2 ts1 = (float2)(0.0f); + private float2 ts2 = (float2)(0.0f); + private float2 ts3 = (float2)(0.0f); + + for (uint k = groupId; k < (K / 32); k += NSUBGROUPS) { + uint sb = k / 8; + uint j = k % 8; + + half2 d = src0_d[gid + sb * LINE_STRIDE_A]; + half2 dm = src0_m[gid + sb * LINE_STRIDE_A]; + + global const uchar * sc0 = src0_s + 2 * gid * scales_per_row + sb * 12; + global const uchar * sc1 = src0_s + (2 * gid + 1) * scales_per_row + sb * 12; + + uchar sv0, mn0, sv1, mn1; + get_scale_min_k4(j, sc0, &sv0, &mn0, mask_d6, mask_d4, mask_hi2); + get_scale_min_k4(j, sc1, &sv1, &mn1, mask_d6, mask_d4, mask_hi2); + + regS = convert_half2(convert_float2(d) * convert_float2((uchar2)(sv0, sv1))); + regM = convert_half2(convert_float2(dm) * convert_float2((uchar2)(mn0, mn1))); + + // high-bit plane + weights loaded ONCE, reused across the columns + regH.s0 = as_ushort(read_imageh(src0_qh, (gid + k * BLOCK_STRIDE_A_QH + LINE_STRIDE_A_QH * 0)).x); + regH.s1 = as_ushort(read_imageh(src0_qh, (gid + k * BLOCK_STRIDE_A_QH + LINE_STRIDE_A_QH * 1)).x); + regH.s2 = as_ushort(read_imageh(src0_qh, (gid + k * BLOCK_STRIDE_A_QH + LINE_STRIDE_A_QH * 2)).x); + regH.s3 = as_ushort(read_imageh(src0_qh, (gid + k * BLOCK_STRIDE_A_QH + LINE_STRIDE_A_QH * 3)).x); + + regA_hi.s0 = read_imageui(src0_q, (gid + k * BLOCK_STRIDE_A + LINE_STRIDE_A * 0)).x; + regA_hi.s1 = read_imageui(src0_q, (gid + k * BLOCK_STRIDE_A + LINE_STRIDE_A * 1)).x; + regA_hi.s2 = read_imageui(src0_q, (gid + k * BLOCK_STRIDE_A + LINE_STRIDE_A * 2)).x; + regA_hi.s3 = read_imageui(src0_q, (gid + k * BLOCK_STRIDE_A + LINE_STRIDE_A * 3)).x; + regA_lo.s0 = read_imageui(src0_q, (gid + k * BLOCK_STRIDE_A + LINE_STRIDE_A * 4)).x; + regA_lo.s1 = read_imageui(src0_q, (gid + k * BLOCK_STRIDE_A + LINE_STRIDE_A * 5)).x; + regA_lo.s2 = read_imageui(src0_q, (gid + k * BLOCK_STRIDE_A + LINE_STRIDE_A * 6)).x; + regA_lo.s3 = read_imageui(src0_q, (gid + k * BLOCK_STRIDE_A + LINE_STRIDE_A * 7)).x; + + MC_COL_Q5K(ts0, 0); + MC_COL_Q5K(ts1, 1); + if (n_cols > 2) MC_COL_Q5K(ts2, 2); + if (n_cols > 3) MC_COL_Q5K(ts3, 3); + } + + // cross-subgroup reduce: pack the (up to 4) columns' float2 into a float8. + local float8 reduceLM[SUBGROUP_SIZE * 3]; + float8 acc = (float8)(ts0.s0, ts0.s1, ts1.s0, ts1.s1, ts2.s0, ts2.s1, ts3.s0, ts3.s1); + if (groupId == 1) { reduceLM[SUBGROUP_SIZE * 0 + slid] = acc; } + if (groupId == 2) { reduceLM[SUBGROUP_SIZE * 1 + slid] = acc; } + if (groupId == 3) { reduceLM[SUBGROUP_SIZE * 2 + slid] = acc; } + + barrier(CLK_LOCAL_MEM_FENCE); + + if (groupId == 0) { + acc += reduceLM[SUBGROUP_SIZE * 0 + slid]; + acc += reduceLM[SUBGROUP_SIZE * 1 + slid]; + acc += reduceLM[SUBGROUP_SIZE * 2 + slid]; + dst = (global float*)((global char*)dst + offsetd); + // dst is column-major [M rows x n_cols cols]: (row, col) at col*M + row + vstore2((float2)(acc.s0, acc.s1), 0, &(dst[0 * M + gid * 2])); + vstore2((float2)(acc.s2, acc.s3), 0, &(dst[1 * M + gid * 2])); + if (n_cols > 2) vstore2((float2)(acc.s4, acc.s5), 0, &(dst[2 * M + gid * 2])); + if (n_cols > 3) vstore2((float2)(acc.s6, acc.s7), 0, &(dst[3 * M + gid * 2])); + } +} +#undef MC_COL_Q5K +#undef MC_DQ5_HI +#undef MC_DQ5_LO diff --git a/ggml/src/ggml-opencl/kernels/gemv_noshuffle_q6_k_f32.cl b/ggml/src/ggml-opencl/kernels/gemv_noshuffle_q6_k_f32.cl index 51682ecebbb..32624ac868f 100644 --- a/ggml/src/ggml-opencl/kernels/gemv_noshuffle_q6_k_f32.cl +++ b/ggml/src/ggml-opencl/kernels/gemv_noshuffle_q6_k_f32.cl @@ -296,3 +296,114 @@ kernel void kernel_gemv_noshuffle_q6_K_f32( if (gid * 2 + 1 < ne01) dst[gid * 2 + 1] = total_sum.s1; } } + +// Multi-column (N=3) q6_K decode GEMV for the spec/MTP verify batch. Same idea +// as the q4_K mc3: stay on the efficient GEMV path (subgroup broadcast, no +// transpose) instead of the transposed-GEMM dead-zone. Each K-block's weights +// (ql/qh, hi+lo) are loaded ONCE and reused across all 3 activation columns. +// Per-column accumulation is independent and identical to 3 standalone GEMVs +// => byte-identical; does NOT perturb the lm_head logits / spec accept rate. +#if defined(ADRENO_GPU) +REQD_SUBGROUP_SIZE_64 +#endif +kernel void kernel_gemv_noshuffle_q6_K_f32_mc3( + read_only image1d_buffer_t src0_ql, + read_only image1d_buffer_t src0_qh, + global half2 * src0_s, + global half2 * src0_d, + read_only image1d_buffer_t src1, + global float * dst, + ulong offsetd, + int ne00, + int ne01 +) { + int grp = get_local_id(1); + int gid = get_global_id(0); + ushort slid = get_sub_group_local_id(); + + int nb = ne00 / 32; + int line_stride_a = ne01 / 2; + int block_stride_a = NSUBGROUPS * ne01; + int COL_STRIDE = ne00 / 4; // float4 pixels per activation column + + uint4 ql_hi, ql_lo; + ushort4 qh_hi, qh_lo; + half2 reg_d; + char4 reg_s; + float8 reg_b; + + float2 ts0 = 0.0f, ts1 = 0.0f, ts2 = 0.0f; + + for (int k = grp; k < nb; k += NSUBGROUPS) { + reg_d = src0_d[gid + k/8 * line_stride_a]; + reg_s = as_char4(src0_s[gid + k * line_stride_a]); + + // weights loaded ONCE (hi: blocks 0-3, lo: blocks 4-7), reused x3 cols + ql_hi.s0 = read_imageui(src0_ql, gid + k*block_stride_a + line_stride_a*0).x; + ql_hi.s1 = read_imageui(src0_ql, gid + k*block_stride_a + line_stride_a*1).x; + ql_hi.s2 = read_imageui(src0_ql, gid + k*block_stride_a + line_stride_a*2).x; + ql_hi.s3 = read_imageui(src0_ql, gid + k*block_stride_a + line_stride_a*3).x; + qh_hi.s0 = as_ushort(read_imageh(src0_qh, gid + k*block_stride_a + line_stride_a*0).x); + qh_hi.s1 = as_ushort(read_imageh(src0_qh, gid + k*block_stride_a + line_stride_a*1).x); + qh_hi.s2 = as_ushort(read_imageh(src0_qh, gid + k*block_stride_a + line_stride_a*2).x); + qh_hi.s3 = as_ushort(read_imageh(src0_qh, gid + k*block_stride_a + line_stride_a*3).x); + + ql_lo.s0 = read_imageui(src0_ql, gid + k*block_stride_a + line_stride_a*4).x; + ql_lo.s1 = read_imageui(src0_ql, gid + k*block_stride_a + line_stride_a*5).x; + ql_lo.s2 = read_imageui(src0_ql, gid + k*block_stride_a + line_stride_a*6).x; + ql_lo.s3 = read_imageui(src0_ql, gid + k*block_stride_a + line_stride_a*7).x; + qh_lo.s0 = as_ushort(read_imageh(src0_qh, gid + k*block_stride_a + line_stride_a*4).x); + qh_lo.s1 = as_ushort(read_imageh(src0_qh, gid + k*block_stride_a + line_stride_a*5).x); + qh_lo.s2 = as_ushort(read_imageh(src0_qh, gid + k*block_stride_a + line_stride_a*6).x); + qh_lo.s3 = as_ushort(read_imageh(src0_qh, gid + k*block_stride_a + line_stride_a*7).x); + + // Per-column: load only this column's activation (single reg_b live) -> + // 1/3 the activation register pressure, cutting the private-mem spill. +#ifdef VECTOR_SUB_GROUP_BROADCAT + { if (slid < 4) { reg_b.s0123 = read_imagef(src1, 0*COL_STRIDE + 0 + slid*2 + k*8); + reg_b.s4567 = read_imagef(src1, 0*COL_STRIDE + 1 + slid*2 + k*8); } + dequantize_block_acc_bcast_8_hi(ts0, as_ushort8(ql_hi), as_uchar8(qh_hi), reg_d, reg_s, reg_b); + dequantize_block_acc_bcast_8_lo(ts0, as_ushort8(ql_lo), as_uchar8(qh_lo), reg_d, reg_s, reg_b); } + { if (slid < 4) { reg_b.s0123 = read_imagef(src1, 1*COL_STRIDE + 0 + slid*2 + k*8); + reg_b.s4567 = read_imagef(src1, 1*COL_STRIDE + 1 + slid*2 + k*8); } + dequantize_block_acc_bcast_8_hi(ts1, as_ushort8(ql_hi), as_uchar8(qh_hi), reg_d, reg_s, reg_b); + dequantize_block_acc_bcast_8_lo(ts1, as_ushort8(ql_lo), as_uchar8(qh_lo), reg_d, reg_s, reg_b); } + { if (slid < 4) { reg_b.s0123 = read_imagef(src1, 2*COL_STRIDE + 0 + slid*2 + k*8); + reg_b.s4567 = read_imagef(src1, 2*COL_STRIDE + 1 + slid*2 + k*8); } + dequantize_block_acc_bcast_8_hi(ts2, as_ushort8(ql_hi), as_uchar8(qh_hi), reg_d, reg_s, reg_b); + dequantize_block_acc_bcast_8_lo(ts2, as_ushort8(ql_lo), as_uchar8(qh_lo), reg_d, reg_s, reg_b); } +#else + { if (slid < 4) { reg_b.s0123 = read_imagef(src1, 0*COL_STRIDE + 0 + slid*2 + k*8); + reg_b.s4567 = read_imagef(src1, 0*COL_STRIDE + 1 + slid*2 + k*8); } + dequantize_block_acc_bcast_1_hi(ts0, as_ushort8(ql_hi), as_uchar8(qh_hi), reg_d, reg_s, reg_b); + dequantize_block_acc_bcast_1_lo(ts0, as_ushort8(ql_lo), as_uchar8(qh_lo), reg_d, reg_s, reg_b); } + { if (slid < 4) { reg_b.s0123 = read_imagef(src1, 1*COL_STRIDE + 0 + slid*2 + k*8); + reg_b.s4567 = read_imagef(src1, 1*COL_STRIDE + 1 + slid*2 + k*8); } + dequantize_block_acc_bcast_1_hi(ts1, as_ushort8(ql_hi), as_uchar8(qh_hi), reg_d, reg_s, reg_b); + dequantize_block_acc_bcast_1_lo(ts1, as_ushort8(ql_lo), as_uchar8(qh_lo), reg_d, reg_s, reg_b); } + { if (slid < 4) { reg_b.s0123 = read_imagef(src1, 2*COL_STRIDE + 0 + slid*2 + k*8); + reg_b.s4567 = read_imagef(src1, 2*COL_STRIDE + 1 + slid*2 + k*8); } + dequantize_block_acc_bcast_1_hi(ts2, as_ushort8(ql_hi), as_uchar8(qh_hi), reg_d, reg_s, reg_b); + dequantize_block_acc_bcast_1_lo(ts2, as_ushort8(ql_lo), as_uchar8(qh_lo), reg_d, reg_s, reg_b); } +#endif + } + + local float8 reduce_lm[SUBGROUP_SIZE * 3]; + float8 acc = (float8)(ts0.s0, ts0.s1, ts1.s0, ts1.s1, ts2.s0, ts2.s1, 0.0f, 0.0f); + if (grp == 1) { reduce_lm[SUBGROUP_SIZE*0 + slid] = acc; } + if (grp == 2) { reduce_lm[SUBGROUP_SIZE*1 + slid] = acc; } + if (grp == 3) { reduce_lm[SUBGROUP_SIZE*2 + slid] = acc; } + + barrier(CLK_LOCAL_MEM_FENCE); + + if (grp == 0) { + acc += reduce_lm[SUBGROUP_SIZE*0 + slid]; + acc += reduce_lm[SUBGROUP_SIZE*1 + slid]; + acc += reduce_lm[SUBGROUP_SIZE*2 + slid]; + dst = (global float*)((global char*)dst + offsetd); + // dst column-major [ne01 rows x 3 cols]: (row, col) at col*ne01 + row + vstore2((float2)(acc.s0, acc.s1), 0, &(dst[0*ne01 + gid*2])); + vstore2((float2)(acc.s2, acc.s3), 0, &(dst[1*ne01 + gid*2])); + vstore2((float2)(acc.s4, acc.s5), 0, &(dst[2*ne01 + gid*2])); + } +} diff --git a/ggml/src/ggml-opencl/kernels/gemv_noshuffle_q6_k_f32_o4.cl b/ggml/src/ggml-opencl/kernels/gemv_noshuffle_q6_k_f32_o4.cl new file mode 100644 index 00000000000..84447e61bb6 --- /dev/null +++ b/ggml/src/ggml-opencl/kernels/gemv_noshuffle_q6_k_f32_o4.cl @@ -0,0 +1,372 @@ +// 4-output-per-WI variant of kernel_gemv_noshuffle_q6_K_f32. +// Each WI now produces 4 consecutive outputs (output quad). The activation +// fetch (reg_b) is shared across all 4 outputs, doubling per-WI ALU per +// activation broadcast and halving the WG count vs the 2-output kernel. +// +// Implementation: each K-block we fetch TWO sets of (scales + ql + qh) +// — one for the low pair (rows 0,1 of the quad) and one for the high pair +// (rows 2,3) — and invoke the existing 2-output dequant macros twice +// against the *same* reg_b. Identical data layout to the 2-output kernel, +// so the host only needs to halve the grid and double the gid-to-output +// mapping. +// +// Opt-in via the host dispatch when GGML_OPENCL_Q6K_GEMV_O4=1. + +#pragma OPENCL EXTENSION cl_khr_fp16 : enable +#pragma OPENCL EXTENSION cl_khr_subgroups : enable + +#ifdef cl_intel_required_subgroup_size +#pragma OPENCL EXTENSION cl_intel_required_subgroup_size : enable +#define INTEL_GPU 1 +#define REQD_SUBGROUP_SIZE_16 __attribute__((intel_reqd_sub_group_size(16))) +#define REQD_SUBGROUP_SIZE_32 __attribute__((intel_reqd_sub_group_size(32))) +#elif defined(cl_qcom_reqd_sub_group_size) +#pragma OPENCL EXTENSION cl_qcom_reqd_sub_group_size : enable +#define ADRENO_GPU 1 +#define REQD_SUBGROUP_SIZE_64 __attribute__((qcom_reqd_sub_group_size("half"))) +#define REQD_SUBGROUP_SIZE_128 __attribute__((qcom_reqd_sub_group_size("full"))) +#endif + +#define NSUBGROUPS 4 +#define SUBGROUP_SIZE 64 + +// Macros are identical to the 2-output kernel — they accept `total_sum` as +// a parameter so we can call them twice (once per pair) against different +// accumulators against the same reg_b. +#define dequantize_block_acc_bcast_8_hi(total_sum, bits4, bits2, cs, y) \ + float8 shared_y; \ + shared_y = sub_group_broadcast(y, 0); \ + total_sum.s0 += ((float)(((bits4.s0 & 0x000F) ) | ((bits2.s0 & 0x03) << 4)) - 32.f) * cs.s0 * shared_y.s0; \ + total_sum.s0 += ((float)(((bits4.s0 & 0x00F0) >> 4) | ((bits2.s0 & 0x0C) << 2)) - 32.f) * cs.s0 * shared_y.s1; \ + total_sum.s0 += ((float)(((bits4.s0 & 0x0F00) >> 8) | ((bits2.s0 & 0x30) )) - 32.f) * cs.s0 * shared_y.s2; \ + total_sum.s0 += ((float)(((bits4.s0 & 0xF000) >> 12) | ((bits2.s0 & 0xC0) >> 2)) - 32.f) * cs.s0 * shared_y.s3; \ + total_sum.s0 += ((float)(((bits4.s2 & 0x000F) ) | ((bits2.s2 & 0x03) << 4)) - 32.f) * cs.s0 * shared_y.s4; \ + total_sum.s0 += ((float)(((bits4.s2 & 0x00F0) >> 4) | ((bits2.s2 & 0x0C) << 2)) - 32.f) * cs.s0 * shared_y.s5; \ + total_sum.s0 += ((float)(((bits4.s2 & 0x0F00) >> 8) | ((bits2.s2 & 0x30) )) - 32.f) * cs.s0 * shared_y.s6; \ + total_sum.s0 += ((float)(((bits4.s2 & 0xF000) >> 12) | ((bits2.s2 & 0xC0) >> 2)) - 32.f) * cs.s0 * shared_y.s7; \ + total_sum.s1 += ((float)(((bits4.s1 & 0x000F) ) | ((bits2.s1 & 0x03) << 4)) - 32.f) * cs.s2 * shared_y.s0; \ + total_sum.s1 += ((float)(((bits4.s1 & 0x00F0) >> 4) | ((bits2.s1 & 0x0C) << 2)) - 32.f) * cs.s2 * shared_y.s1; \ + total_sum.s1 += ((float)(((bits4.s1 & 0x0F00) >> 8) | ((bits2.s1 & 0x30) )) - 32.f) * cs.s2 * shared_y.s2; \ + total_sum.s1 += ((float)(((bits4.s1 & 0xF000) >> 12) | ((bits2.s1 & 0xC0) >> 2)) - 32.f) * cs.s2 * shared_y.s3; \ + total_sum.s1 += ((float)(((bits4.s3 & 0x000F) ) | ((bits2.s3 & 0x03) << 4)) - 32.f) * cs.s2 * shared_y.s4; \ + total_sum.s1 += ((float)(((bits4.s3 & 0x00F0) >> 4) | ((bits2.s3 & 0x0C) << 2)) - 32.f) * cs.s2 * shared_y.s5; \ + total_sum.s1 += ((float)(((bits4.s3 & 0x0F00) >> 8) | ((bits2.s3 & 0x30) )) - 32.f) * cs.s2 * shared_y.s6; \ + total_sum.s1 += ((float)(((bits4.s3 & 0xF000) >> 12) | ((bits2.s3 & 0xC0) >> 2)) - 32.f) * cs.s2 * shared_y.s7; \ + shared_y = sub_group_broadcast(y, 1); \ + total_sum.s0 += ((float)(((bits4.s4 & 0x000F) ) | ((bits2.s4 & 0x03) << 4)) - 32.f) * cs.s0 * shared_y.s0; \ + total_sum.s0 += ((float)(((bits4.s4 & 0x00F0) >> 4) | ((bits2.s4 & 0x0C) << 2)) - 32.f) * cs.s0 * shared_y.s1; \ + total_sum.s0 += ((float)(((bits4.s4 & 0x0F00) >> 8) | ((bits2.s4 & 0x30) )) - 32.f) * cs.s0 * shared_y.s2; \ + total_sum.s0 += ((float)(((bits4.s4 & 0xF000) >> 12) | ((bits2.s4 & 0xC0) >> 2)) - 32.f) * cs.s0 * shared_y.s3; \ + total_sum.s0 += ((float)(((bits4.s6 & 0x000F) ) | ((bits2.s6 & 0x03) << 4)) - 32.f) * cs.s0 * shared_y.s4; \ + total_sum.s0 += ((float)(((bits4.s6 & 0x00F0) >> 4) | ((bits2.s6 & 0x0C) << 2)) - 32.f) * cs.s0 * shared_y.s5; \ + total_sum.s0 += ((float)(((bits4.s6 & 0x0F00) >> 8) | ((bits2.s6 & 0x30) )) - 32.f) * cs.s0 * shared_y.s6; \ + total_sum.s0 += ((float)(((bits4.s6 & 0xF000) >> 12) | ((bits2.s6 & 0xC0) >> 2)) - 32.f) * cs.s0 * shared_y.s7; \ + total_sum.s1 += ((float)(((bits4.s5 & 0x000F) ) | ((bits2.s5 & 0x03) << 4)) - 32.f) * cs.s2 * shared_y.s0; \ + total_sum.s1 += ((float)(((bits4.s5 & 0x00F0) >> 4) | ((bits2.s5 & 0x0C) << 2)) - 32.f) * cs.s2 * shared_y.s1; \ + total_sum.s1 += ((float)(((bits4.s5 & 0x0F00) >> 8) | ((bits2.s5 & 0x30) )) - 32.f) * cs.s2 * shared_y.s2; \ + total_sum.s1 += ((float)(((bits4.s5 & 0xF000) >> 12) | ((bits2.s5 & 0xC0) >> 2)) - 32.f) * cs.s2 * shared_y.s3; \ + total_sum.s1 += ((float)(((bits4.s7 & 0x000F) ) | ((bits2.s7 & 0x03) << 4)) - 32.f) * cs.s2 * shared_y.s4; \ + total_sum.s1 += ((float)(((bits4.s7 & 0x00F0) >> 4) | ((bits2.s7 & 0x0C) << 2)) - 32.f) * cs.s2 * shared_y.s5; \ + total_sum.s1 += ((float)(((bits4.s7 & 0x0F00) >> 8) | ((bits2.s7 & 0x30) )) - 32.f) * cs.s2 * shared_y.s6; \ + total_sum.s1 += ((float)(((bits4.s7 & 0xF000) >> 12) | ((bits2.s7 & 0xC0) >> 2)) - 32.f) * cs.s2 * shared_y.s7; \ + +#define dequantize_block_acc_bcast_8_lo(total_sum, bits4, bits2, cs, y) \ + shared_y = sub_group_broadcast(y, 2); \ + total_sum.s0 += ((float)(((bits4.s0 & 0x000F) ) | ((bits2.s0 & 0x03) << 4)) - 32.f) * cs.s1 * shared_y.s0; \ + total_sum.s0 += ((float)(((bits4.s0 & 0x00F0) >> 4) | ((bits2.s0 & 0x0C) << 2)) - 32.f) * cs.s1 * shared_y.s1; \ + total_sum.s0 += ((float)(((bits4.s0 & 0x0F00) >> 8) | ((bits2.s0 & 0x30) )) - 32.f) * cs.s1 * shared_y.s2; \ + total_sum.s0 += ((float)(((bits4.s0 & 0xF000) >> 12) | ((bits2.s0 & 0xC0) >> 2)) - 32.f) * cs.s1 * shared_y.s3; \ + total_sum.s0 += ((float)(((bits4.s2 & 0x000F) ) | ((bits2.s2 & 0x03) << 4)) - 32.f) * cs.s1 * shared_y.s4; \ + total_sum.s0 += ((float)(((bits4.s2 & 0x00F0) >> 4) | ((bits2.s2 & 0x0C) << 2)) - 32.f) * cs.s1 * shared_y.s5; \ + total_sum.s0 += ((float)(((bits4.s2 & 0x0F00) >> 8) | ((bits2.s2 & 0x30) )) - 32.f) * cs.s1 * shared_y.s6; \ + total_sum.s0 += ((float)(((bits4.s2 & 0xF000) >> 12) | ((bits2.s2 & 0xC0) >> 2)) - 32.f) * cs.s1 * shared_y.s7; \ + total_sum.s1 += ((float)(((bits4.s1 & 0x000F) ) | ((bits2.s1 & 0x03) << 4)) - 32.f) * cs.s3 * shared_y.s0; \ + total_sum.s1 += ((float)(((bits4.s1 & 0x00F0) >> 4) | ((bits2.s1 & 0x0C) << 2)) - 32.f) * cs.s3 * shared_y.s1; \ + total_sum.s1 += ((float)(((bits4.s1 & 0x0F00) >> 8) | ((bits2.s1 & 0x30) )) - 32.f) * cs.s3 * shared_y.s2; \ + total_sum.s1 += ((float)(((bits4.s1 & 0xF000) >> 12) | ((bits2.s1 & 0xC0) >> 2)) - 32.f) * cs.s3 * shared_y.s3; \ + total_sum.s1 += ((float)(((bits4.s3 & 0x000F) ) | ((bits2.s3 & 0x03) << 4)) - 32.f) * cs.s3 * shared_y.s4; \ + total_sum.s1 += ((float)(((bits4.s3 & 0x00F0) >> 4) | ((bits2.s3 & 0x0C) << 2)) - 32.f) * cs.s3 * shared_y.s5; \ + total_sum.s1 += ((float)(((bits4.s3 & 0x0F00) >> 8) | ((bits2.s3 & 0x30) )) - 32.f) * cs.s3 * shared_y.s6; \ + total_sum.s1 += ((float)(((bits4.s3 & 0xF000) >> 12) | ((bits2.s3 & 0xC0) >> 2)) - 32.f) * cs.s3 * shared_y.s7; \ + shared_y = sub_group_broadcast(y, 3); \ + total_sum.s0 += ((float)(((bits4.s4 & 0x000F) ) | ((bits2.s4 & 0x03) << 4)) - 32.f) * cs.s1 * shared_y.s0; \ + total_sum.s0 += ((float)(((bits4.s4 & 0x00F0) >> 4) | ((bits2.s4 & 0x0C) << 2)) - 32.f) * cs.s1 * shared_y.s1; \ + total_sum.s0 += ((float)(((bits4.s4 & 0x0F00) >> 8) | ((bits2.s4 & 0x30) )) - 32.f) * cs.s1 * shared_y.s2; \ + total_sum.s0 += ((float)(((bits4.s4 & 0xF000) >> 12) | ((bits2.s4 & 0xC0) >> 2)) - 32.f) * cs.s1 * shared_y.s3; \ + total_sum.s0 += ((float)(((bits4.s6 & 0x000F) ) | ((bits2.s6 & 0x03) << 4)) - 32.f) * cs.s1 * shared_y.s4; \ + total_sum.s0 += ((float)(((bits4.s6 & 0x00F0) >> 4) | ((bits2.s6 & 0x0C) << 2)) - 32.f) * cs.s1 * shared_y.s5; \ + total_sum.s0 += ((float)(((bits4.s6 & 0x0F00) >> 8) | ((bits2.s6 & 0x30) )) - 32.f) * cs.s1 * shared_y.s6; \ + total_sum.s0 += ((float)(((bits4.s6 & 0xF000) >> 12) | ((bits2.s6 & 0xC0) >> 2)) - 32.f) * cs.s1 * shared_y.s7; \ + total_sum.s1 += ((float)(((bits4.s5 & 0x000F) ) | ((bits2.s5 & 0x03) << 4)) - 32.f) * cs.s3 * shared_y.s0; \ + total_sum.s1 += ((float)(((bits4.s5 & 0x00F0) >> 4) | ((bits2.s5 & 0x0C) << 2)) - 32.f) * cs.s3 * shared_y.s1; \ + total_sum.s1 += ((float)(((bits4.s5 & 0x0F00) >> 8) | ((bits2.s5 & 0x30) )) - 32.f) * cs.s3 * shared_y.s2; \ + total_sum.s1 += ((float)(((bits4.s5 & 0xF000) >> 12) | ((bits2.s5 & 0xC0) >> 2)) - 32.f) * cs.s3 * shared_y.s3; \ + total_sum.s1 += ((float)(((bits4.s7 & 0x000F) ) | ((bits2.s7 & 0x03) << 4)) - 32.f) * cs.s3 * shared_y.s4; \ + total_sum.s1 += ((float)(((bits4.s7 & 0x00F0) >> 4) | ((bits2.s7 & 0x0C) << 2)) - 32.f) * cs.s3 * shared_y.s5; \ + total_sum.s1 += ((float)(((bits4.s7 & 0x0F00) >> 8) | ((bits2.s7 & 0x30) )) - 32.f) * cs.s3 * shared_y.s6; \ + total_sum.s1 += ((float)(((bits4.s7 & 0xF000) >> 12) | ((bits2.s7 & 0xC0) >> 2)) - 32.f) * cs.s3 * shared_y.s7; \ + +#define dequantize_block_acc_bcast_1_hi(total_sum, bits4, bits2, cs, y) \ + float shared_y; \ + shared_y = sub_group_broadcast(y.s0, 0); \ + total_sum.s0 += ((float)(((bits4.s0 & 0x000F) ) | ((bits2.s0 & 0x03) << 4)) - 32.f) * cs.s0 * shared_y; \ + total_sum.s1 += ((float)(((bits4.s1 & 0x000F) ) | ((bits2.s1 & 0x03) << 4)) - 32.f) * cs.s2 * shared_y; \ + shared_y = sub_group_broadcast(y.s1, 0); \ + total_sum.s0 += ((float)(((bits4.s0 & 0x00F0) >> 4) | ((bits2.s0 & 0x0C) << 2)) - 32.f) * cs.s0 * shared_y; \ + total_sum.s1 += ((float)(((bits4.s1 & 0x00F0) >> 4) | ((bits2.s1 & 0x0C) << 2)) - 32.f) * cs.s2 * shared_y; \ + shared_y = sub_group_broadcast(y.s2, 0); \ + total_sum.s0 += ((float)(((bits4.s0 & 0x0F00) >> 8) | ((bits2.s0 & 0x30) )) - 32.f) * cs.s0 * shared_y; \ + total_sum.s1 += ((float)(((bits4.s1 & 0x0F00) >> 8) | ((bits2.s1 & 0x30) )) - 32.f) * cs.s2 * shared_y; \ + shared_y = sub_group_broadcast(y.s3, 0); \ + total_sum.s0 += ((float)(((bits4.s0 & 0xF000) >> 12) | ((bits2.s0 & 0xC0) >> 2)) - 32.f) * cs.s0 * shared_y; \ + total_sum.s1 += ((float)(((bits4.s1 & 0xF000) >> 12) | ((bits2.s1 & 0xC0) >> 2)) - 32.f) * cs.s2 * shared_y; \ + shared_y = sub_group_broadcast(y.s4, 0); \ + total_sum.s0 += ((float)(((bits4.s2 & 0x000F) ) | ((bits2.s2 & 0x03) << 4)) - 32.f) * cs.s0 * shared_y; \ + total_sum.s1 += ((float)(((bits4.s3 & 0x000F) ) | ((bits2.s3 & 0x03) << 4)) - 32.f) * cs.s2 * shared_y; \ + shared_y = sub_group_broadcast(y.s5, 0); \ + total_sum.s0 += ((float)(((bits4.s2 & 0x00F0) >> 4) | ((bits2.s2 & 0x0C) << 2)) - 32.f) * cs.s0 * shared_y; \ + total_sum.s1 += ((float)(((bits4.s3 & 0x00F0) >> 4) | ((bits2.s3 & 0x0C) << 2)) - 32.f) * cs.s2 * shared_y; \ + shared_y = sub_group_broadcast(y.s6, 0); \ + total_sum.s0 += ((float)(((bits4.s2 & 0x0F00) >> 8) | ((bits2.s2 & 0x30) )) - 32.f) * cs.s0 * shared_y; \ + total_sum.s1 += ((float)(((bits4.s3 & 0x0F00) >> 8) | ((bits2.s3 & 0x30) )) - 32.f) * cs.s2 * shared_y; \ + shared_y = sub_group_broadcast(y.s7, 0); \ + total_sum.s0 += ((float)(((bits4.s2 & 0xF000) >> 12) | ((bits2.s2 & 0xC0) >> 2)) - 32.f) * cs.s0 * shared_y; \ + total_sum.s1 += ((float)(((bits4.s3 & 0xF000) >> 12) | ((bits2.s3 & 0xC0) >> 2)) - 32.f) * cs.s2 * shared_y; \ + shared_y = sub_group_broadcast(y.s0, 1); \ + total_sum.s0 += ((float)(((bits4.s4 & 0x000F) ) | ((bits2.s4 & 0x03) << 4)) - 32.f) * cs.s0 * shared_y; \ + total_sum.s1 += ((float)(((bits4.s5 & 0x000F) ) | ((bits2.s5 & 0x03) << 4)) - 32.f) * cs.s2 * shared_y; \ + shared_y = sub_group_broadcast(y.s1, 1); \ + total_sum.s0 += ((float)(((bits4.s4 & 0x00F0) >> 4) | ((bits2.s4 & 0x0C) << 2)) - 32.f) * cs.s0 * shared_y; \ + total_sum.s1 += ((float)(((bits4.s5 & 0x00F0) >> 4) | ((bits2.s5 & 0x0C) << 2)) - 32.f) * cs.s2 * shared_y; \ + shared_y = sub_group_broadcast(y.s2, 1); \ + total_sum.s0 += ((float)(((bits4.s4 & 0x0F00) >> 8) | ((bits2.s4 & 0x30) )) - 32.f) * cs.s0 * shared_y; \ + total_sum.s1 += ((float)(((bits4.s5 & 0x0F00) >> 8) | ((bits2.s5 & 0x30) )) - 32.f) * cs.s2 * shared_y; \ + shared_y = sub_group_broadcast(y.s3, 1); \ + total_sum.s0 += ((float)(((bits4.s4 & 0xF000) >> 12) | ((bits2.s4 & 0xC0) >> 2)) - 32.f) * cs.s0 * shared_y; \ + total_sum.s1 += ((float)(((bits4.s5 & 0xF000) >> 12) | ((bits2.s5 & 0xC0) >> 2)) - 32.f) * cs.s2 * shared_y; \ + shared_y = sub_group_broadcast(y.s4, 1); \ + total_sum.s0 += ((float)(((bits4.s6 & 0x000F) ) | ((bits2.s6 & 0x03) << 4)) - 32.f) * cs.s0 * shared_y; \ + total_sum.s1 += ((float)(((bits4.s7 & 0x000F) ) | ((bits2.s7 & 0x03) << 4)) - 32.f) * cs.s2 * shared_y; \ + shared_y = sub_group_broadcast(y.s5, 1); \ + total_sum.s0 += ((float)(((bits4.s6 & 0x00F0) >> 4) | ((bits2.s6 & 0x0C) << 2)) - 32.f) * cs.s0 * shared_y; \ + total_sum.s1 += ((float)(((bits4.s7 & 0x00F0) >> 4) | ((bits2.s7 & 0x0C) << 2)) - 32.f) * cs.s2 * shared_y; \ + shared_y = sub_group_broadcast(y.s6, 1); \ + total_sum.s0 += ((float)(((bits4.s6 & 0x0F00) >> 8) | ((bits2.s6 & 0x30) )) - 32.f) * cs.s0 * shared_y; \ + total_sum.s1 += ((float)(((bits4.s7 & 0x0F00) >> 8) | ((bits2.s7 & 0x30) )) - 32.f) * cs.s2 * shared_y; \ + shared_y = sub_group_broadcast(y.s7, 1); \ + total_sum.s0 += ((float)(((bits4.s6 & 0xF000) >> 12) | ((bits2.s6 & 0xC0) >> 2)) - 32.f) * cs.s0 * shared_y; \ + total_sum.s1 += ((float)(((bits4.s7 & 0xF000) >> 12) | ((bits2.s7 & 0xC0) >> 2)) - 32.f) * cs.s2 * shared_y; \ + +#define dequantize_block_acc_bcast_1_lo(total_sum, bits4, bits2, cs, y) \ + shared_y = sub_group_broadcast(y.s0, 2); \ + total_sum.s0 += ((float)(((bits4.s0 & 0x000F) ) | ((bits2.s0 & 0x03) << 4)) - 32.f) * cs.s1 * shared_y; \ + total_sum.s1 += ((float)(((bits4.s1 & 0x000F) ) | ((bits2.s1 & 0x03) << 4)) - 32.f) * cs.s3 * shared_y; \ + shared_y = sub_group_broadcast(y.s1, 2); \ + total_sum.s0 += ((float)(((bits4.s0 & 0x00F0) >> 4) | ((bits2.s0 & 0x0C) << 2)) - 32.f) * cs.s1 * shared_y; \ + total_sum.s1 += ((float)(((bits4.s1 & 0x00F0) >> 4) | ((bits2.s1 & 0x0C) << 2)) - 32.f) * cs.s3 * shared_y; \ + shared_y = sub_group_broadcast(y.s2, 2); \ + total_sum.s0 += ((float)(((bits4.s0 & 0x0F00) >> 8) | ((bits2.s0 & 0x30) )) - 32.f) * cs.s1 * shared_y; \ + total_sum.s1 += ((float)(((bits4.s1 & 0x0F00) >> 8) | ((bits2.s1 & 0x30) )) - 32.f) * cs.s3 * shared_y; \ + shared_y = sub_group_broadcast(y.s3, 2); \ + total_sum.s0 += ((float)(((bits4.s0 & 0xF000) >> 12) | ((bits2.s0 & 0xC0) >> 2)) - 32.f) * cs.s1 * shared_y; \ + total_sum.s1 += ((float)(((bits4.s1 & 0xF000) >> 12) | ((bits2.s1 & 0xC0) >> 2)) - 32.f) * cs.s3 * shared_y; \ + shared_y = sub_group_broadcast(y.s4, 2); \ + total_sum.s0 += ((float)(((bits4.s2 & 0x000F) ) | ((bits2.s2 & 0x03) << 4)) - 32.f) * cs.s1 * shared_y; \ + total_sum.s1 += ((float)(((bits4.s3 & 0x000F) ) | ((bits2.s3 & 0x03) << 4)) - 32.f) * cs.s3 * shared_y; \ + shared_y = sub_group_broadcast(y.s5, 2); \ + total_sum.s0 += ((float)(((bits4.s2 & 0x00F0) >> 4) | ((bits2.s2 & 0x0C) << 2)) - 32.f) * cs.s1 * shared_y; \ + total_sum.s1 += ((float)(((bits4.s3 & 0x00F0) >> 4) | ((bits2.s3 & 0x0C) << 2)) - 32.f) * cs.s3 * shared_y; \ + shared_y = sub_group_broadcast(y.s6, 2); \ + total_sum.s0 += ((float)(((bits4.s2 & 0x0F00) >> 8) | ((bits2.s2 & 0x30) )) - 32.f) * cs.s1 * shared_y; \ + total_sum.s1 += ((float)(((bits4.s3 & 0x0F00) >> 8) | ((bits2.s3 & 0x30) )) - 32.f) * cs.s3 * shared_y; \ + shared_y = sub_group_broadcast(y.s7, 2); \ + total_sum.s0 += ((float)(((bits4.s2 & 0xF000) >> 12) | ((bits2.s2 & 0xC0) >> 2)) - 32.f) * cs.s1 * shared_y; \ + total_sum.s1 += ((float)(((bits4.s3 & 0xF000) >> 12) | ((bits2.s3 & 0xC0) >> 2)) - 32.f) * cs.s3 * shared_y; \ + shared_y = sub_group_broadcast(y.s0, 3); \ + total_sum.s0 += ((float)(((bits4.s4 & 0x000F) ) | ((bits2.s4 & 0x03) << 4)) - 32.f) * cs.s1 * shared_y; \ + total_sum.s1 += ((float)(((bits4.s5 & 0x000F) ) | ((bits2.s5 & 0x03) << 4)) - 32.f) * cs.s3 * shared_y; \ + shared_y = sub_group_broadcast(y.s1, 3); \ + total_sum.s0 += ((float)(((bits4.s4 & 0x00F0) >> 4) | ((bits2.s4 & 0x0C) << 2)) - 32.f) * cs.s1 * shared_y; \ + total_sum.s1 += ((float)(((bits4.s5 & 0x00F0) >> 4) | ((bits2.s5 & 0x0C) << 2)) - 32.f) * cs.s3 * shared_y; \ + shared_y = sub_group_broadcast(y.s2, 3); \ + total_sum.s0 += ((float)(((bits4.s4 & 0x0F00) >> 8) | ((bits2.s4 & 0x30) )) - 32.f) * cs.s1 * shared_y; \ + total_sum.s1 += ((float)(((bits4.s5 & 0x0F00) >> 8) | ((bits2.s5 & 0x30) )) - 32.f) * cs.s3 * shared_y; \ + shared_y = sub_group_broadcast(y.s3, 3); \ + total_sum.s0 += ((float)(((bits4.s4 & 0xF000) >> 12) | ((bits2.s4 & 0xC0) >> 2)) - 32.f) * cs.s1 * shared_y; \ + total_sum.s1 += ((float)(((bits4.s5 & 0xF000) >> 12) | ((bits2.s5 & 0xC0) >> 2)) - 32.f) * cs.s3 * shared_y; \ + shared_y = sub_group_broadcast(y.s4, 3); \ + total_sum.s0 += ((float)(((bits4.s6 & 0x000F) ) | ((bits2.s6 & 0x03) << 4)) - 32.f) * cs.s1 * shared_y; \ + total_sum.s1 += ((float)(((bits4.s7 & 0x000F) ) | ((bits2.s7 & 0x03) << 4)) - 32.f) * cs.s3 * shared_y; \ + shared_y = sub_group_broadcast(y.s5, 3); \ + total_sum.s0 += ((float)(((bits4.s6 & 0x00F0) >> 4) | ((bits2.s6 & 0x0C) << 2)) - 32.f) * cs.s1 * shared_y; \ + total_sum.s1 += ((float)(((bits4.s7 & 0x00F0) >> 4) | ((bits2.s7 & 0x0C) << 2)) - 32.f) * cs.s3 * shared_y; \ + shared_y = sub_group_broadcast(y.s6, 3); \ + total_sum.s0 += ((float)(((bits4.s6 & 0x0F00) >> 8) | ((bits2.s6 & 0x30) )) - 32.f) * cs.s1 * shared_y; \ + total_sum.s1 += ((float)(((bits4.s7 & 0x0F00) >> 8) | ((bits2.s7 & 0x30) )) - 32.f) * cs.s3 * shared_y; \ + shared_y = sub_group_broadcast(y.s7, 3); \ + total_sum.s0 += ((float)(((bits4.s6 & 0xF000) >> 12) | ((bits2.s6 & 0xC0) >> 2)) - 32.f) * cs.s1 * shared_y; \ + total_sum.s1 += ((float)(((bits4.s7 & 0xF000) >> 12) | ((bits2.s7 & 0xC0) >> 2)) - 32.f) * cs.s3 * shared_y; \ + +#if defined(ADRENO_GPU) +REQD_SUBGROUP_SIZE_64 +#endif +// Q6K_O4_GLOBAL: read the (read-once-per-token, no-reuse) lm_head/embed weights +// from __global coalesced instead of image1d_buffer. The texture cache caps the +// streaming (no-reuse) lm_head read bandwidth; global coalesced reaches the +// higher rate the rest of the model gets. src1 (activation) stays an image (it IS reused via +// the cross-subgroup broadcast). +#ifdef Q6K_O4_GLOBAL +#define Q6K_O4_NAME kernel_gemv_noshuffle_q6_K_f32_o4_global +#define QL_ARG __global uint * src0_ql +#define QH_ARG __global half * src0_qh +#define RD_QL(b,i) (b[i]) +#define RD_QH(b,i) as_ushort(b[i]) +#else +#define Q6K_O4_NAME kernel_gemv_noshuffle_q6_K_f32_o4 +#define QL_ARG read_only image1d_buffer_t src0_ql +#define QH_ARG read_only image1d_buffer_t src0_qh +#define RD_QL(b,i) (read_imageui(b,i).x) +#define RD_QH(b,i) as_ushort(read_imageh(b,i).x) +#endif +kernel void Q6K_O4_NAME( + QL_ARG, + QH_ARG, + global half2 * src0_s, + global half2 * src0_d, + read_only image1d_buffer_t src1, + global float * dst, + ulong offsetd, + int ne00, + int ne01 +) { + int grp = get_local_id(1); + int gid = get_global_id(0); // 4-output-quad index + ushort slid = get_sub_group_local_id(); + + // Map quad index to the two pair-indices the existing 2-output access + // pattern uses (consecutive output pairs along ne01). NB: the two pairs are + // kept ADJACENT (gid*2, gid*2+1) on purpose -- a "stride-1" split (pairs + // ne01/4 apart) is slower because two distant cache-line streams have worse + // locality than the adjacent pair whose reads interleave into the same lines + // each iteration. + int gid_a = gid * 2; + int gid_b = gid * 2 + 1; + + int nb = ne00 / 32; + + uint4 reg_a_l_a, reg_a_l_b; + ushort4 reg_a_h_a, reg_a_h_b; + half2 reg_d_a, reg_d_b; + char4 reg_s_a, reg_s_b; + float8 reg_b; + + float2 total_sum_a = 0.0f; + float2 total_sum_b = 0.0f; + + int line_stride_a = ne01 / 2; + int block_stride_a = NSUBGROUPS * ne01; + + for (int k = grp; k < nb; k += NSUBGROUPS) { + reg_d_a = src0_d[gid_a + k/8 * line_stride_a]; + reg_d_b = src0_d[gid_b + k/8 * line_stride_a]; + reg_s_a = as_char4(src0_s[gid_a + k * line_stride_a]); + reg_s_b = as_char4(src0_s[gid_b + k * line_stride_a]); + // Precompute the loop-invariant combined scale (sub-block scale * super-block d) + // once per pair instead of re-multiplying it for every one of the 256 elements. + float4 cs_a = (float4)((float)reg_s_a.s0*(float)reg_d_a.s0, (float)reg_s_a.s1*(float)reg_d_a.s0, + (float)reg_s_a.s2*(float)reg_d_a.s1, (float)reg_s_a.s3*(float)reg_d_a.s1); + float4 cs_b = (float4)((float)reg_s_b.s0*(float)reg_d_b.s0, (float)reg_s_b.s1*(float)reg_d_b.s0, + (float)reg_s_b.s2*(float)reg_d_b.s1, (float)reg_s_b.s3*(float)reg_d_b.s1); + + if (slid < 4) { + reg_b.s0123 = read_imagef(src1, 0 + slid*2 + k*8); + reg_b.s4567 = read_imagef(src1, 1 + slid*2 + k*8); + } + + // Pair a (output rows gid_a*2, gid_a*2+1): read hi+lo then dequant + // both in one block so the `_lo` macro can see the `shared_y` that + // `_hi` declared. Pair b follows in its own block — fresh shared_y. + { + reg_a_l_a.s0 = RD_QL(src0_ql, gid_a + k*block_stride_a + line_stride_a*0); + reg_a_l_a.s1 = RD_QL(src0_ql, gid_a + k*block_stride_a + line_stride_a*1); + reg_a_l_a.s2 = RD_QL(src0_ql, gid_a + k*block_stride_a + line_stride_a*2); + reg_a_l_a.s3 = RD_QL(src0_ql, gid_a + k*block_stride_a + line_stride_a*3); + reg_a_h_a.s0 = RD_QH(src0_qh, gid_a + k*block_stride_a + line_stride_a*0); + reg_a_h_a.s1 = RD_QH(src0_qh, gid_a + k*block_stride_a + line_stride_a*1); + reg_a_h_a.s2 = RD_QH(src0_qh, gid_a + k*block_stride_a + line_stride_a*2); + reg_a_h_a.s3 = RD_QH(src0_qh, gid_a + k*block_stride_a + line_stride_a*3); +#ifdef VECTOR_SUB_GROUP_BROADCAT + dequantize_block_acc_bcast_8_hi(total_sum_a, as_ushort8(reg_a_l_a), as_uchar8(reg_a_h_a), cs_a, reg_b); +#else + dequantize_block_acc_bcast_1_hi(total_sum_a, as_ushort8(reg_a_l_a), as_uchar8(reg_a_h_a), cs_a, reg_b); +#endif + + reg_a_l_a.s0 = RD_QL(src0_ql, gid_a + k*block_stride_a + line_stride_a*4); + reg_a_l_a.s1 = RD_QL(src0_ql, gid_a + k*block_stride_a + line_stride_a*5); + reg_a_l_a.s2 = RD_QL(src0_ql, gid_a + k*block_stride_a + line_stride_a*6); + reg_a_l_a.s3 = RD_QL(src0_ql, gid_a + k*block_stride_a + line_stride_a*7); + reg_a_h_a.s0 = RD_QH(src0_qh, gid_a + k*block_stride_a + line_stride_a*4); + reg_a_h_a.s1 = RD_QH(src0_qh, gid_a + k*block_stride_a + line_stride_a*5); + reg_a_h_a.s2 = RD_QH(src0_qh, gid_a + k*block_stride_a + line_stride_a*6); + reg_a_h_a.s3 = RD_QH(src0_qh, gid_a + k*block_stride_a + line_stride_a*7); +#ifdef VECTOR_SUB_GROUP_BROADCAT + dequantize_block_acc_bcast_8_lo(total_sum_a, as_ushort8(reg_a_l_a), as_uchar8(reg_a_h_a), cs_a, reg_b); +#else + dequantize_block_acc_bcast_1_lo(total_sum_a, as_ushort8(reg_a_l_a), as_uchar8(reg_a_h_a), cs_a, reg_b); +#endif + } + + { + reg_a_l_b.s0 = RD_QL(src0_ql, gid_b + k*block_stride_a + line_stride_a*0); + reg_a_l_b.s1 = RD_QL(src0_ql, gid_b + k*block_stride_a + line_stride_a*1); + reg_a_l_b.s2 = RD_QL(src0_ql, gid_b + k*block_stride_a + line_stride_a*2); + reg_a_l_b.s3 = RD_QL(src0_ql, gid_b + k*block_stride_a + line_stride_a*3); + reg_a_h_b.s0 = RD_QH(src0_qh, gid_b + k*block_stride_a + line_stride_a*0); + reg_a_h_b.s1 = RD_QH(src0_qh, gid_b + k*block_stride_a + line_stride_a*1); + reg_a_h_b.s2 = RD_QH(src0_qh, gid_b + k*block_stride_a + line_stride_a*2); + reg_a_h_b.s3 = RD_QH(src0_qh, gid_b + k*block_stride_a + line_stride_a*3); +#ifdef VECTOR_SUB_GROUP_BROADCAT + dequantize_block_acc_bcast_8_hi(total_sum_b, as_ushort8(reg_a_l_b), as_uchar8(reg_a_h_b), cs_b, reg_b); +#else + dequantize_block_acc_bcast_1_hi(total_sum_b, as_ushort8(reg_a_l_b), as_uchar8(reg_a_h_b), cs_b, reg_b); +#endif + + reg_a_l_b.s0 = RD_QL(src0_ql, gid_b + k*block_stride_a + line_stride_a*4); + reg_a_l_b.s1 = RD_QL(src0_ql, gid_b + k*block_stride_a + line_stride_a*5); + reg_a_l_b.s2 = RD_QL(src0_ql, gid_b + k*block_stride_a + line_stride_a*6); + reg_a_l_b.s3 = RD_QL(src0_ql, gid_b + k*block_stride_a + line_stride_a*7); + reg_a_h_b.s0 = RD_QH(src0_qh, gid_b + k*block_stride_a + line_stride_a*4); + reg_a_h_b.s1 = RD_QH(src0_qh, gid_b + k*block_stride_a + line_stride_a*5); + reg_a_h_b.s2 = RD_QH(src0_qh, gid_b + k*block_stride_a + line_stride_a*6); + reg_a_h_b.s3 = RD_QH(src0_qh, gid_b + k*block_stride_a + line_stride_a*7); +#ifdef VECTOR_SUB_GROUP_BROADCAT + dequantize_block_acc_bcast_8_lo(total_sum_b, as_ushort8(reg_a_l_b), as_uchar8(reg_a_h_b), cs_b, reg_b); +#else + dequantize_block_acc_bcast_1_lo(total_sum_b, as_ushort8(reg_a_l_b), as_uchar8(reg_a_h_b), cs_b, reg_b); +#endif + } + } + + // Cross-subgroup reduce. Same shape as the 2-output kernel but with the + // pair-a and pair-b accumulators concatenated into a single float4. + local float4 reduce_lm[SUBGROUP_SIZE * 3]; + float4 acc = (float4)(total_sum_a.s0, total_sum_a.s1, total_sum_b.s0, total_sum_b.s1); + if (grp == 1) { reduce_lm[SUBGROUP_SIZE*0 + slid] = acc; } + if (grp == 2) { reduce_lm[SUBGROUP_SIZE*1 + slid] = acc; } + if (grp == 3) { reduce_lm[SUBGROUP_SIZE*2 + slid] = acc; } + + barrier(CLK_LOCAL_MEM_FENCE); + + if (grp == 0) { + acc += reduce_lm[SUBGROUP_SIZE*0 + slid]; + acc += reduce_lm[SUBGROUP_SIZE*1 + slid]; + acc += reduce_lm[SUBGROUP_SIZE*2 + slid]; + dst = (global float*)((global char*)dst + offsetd); + // The dispatch rounds ne01/4 up to the subgroup width, so the tail + // quads past the last row must not store (they wrote 128 rows past + // dst on every ne01 % 256 == 128 vocab, e.g. 151936). + if (gid * 4 + 3 < (uint)ne01) { + vstore4(acc, 0, &(dst[gid * 4])); + } + } +} diff --git a/ggml/src/ggml-opencl/kernels/gemv_noshuffle_q6_k_f32_tiled.cl b/ggml/src/ggml-opencl/kernels/gemv_noshuffle_q6_k_f32_tiled.cl new file mode 100644 index 00000000000..c5049f3964e --- /dev/null +++ b/ggml/src/ggml-opencl/kernels/gemv_noshuffle_q6_k_f32_tiled.cl @@ -0,0 +1,196 @@ +// Tiled-wide q6_K GEMV for the long-vocab lm_head/embed (decode path). +// +// Pairs with kernel_convert_block_q6_k_tiled_ns (cvt.cl): the weights are laid +// out CANONICALLY (6-bit code in element order e in [0,256)) and TILED by 64 +// output rows so the 64-thread lane group coalesces every weight load. Both the +// pack (convert) and the unpack (here) are owned by us — correct by construction +// against the reference ggml q6_K dequant, no bit-interleave reverse-engineering. +// +// One work-item produces one output row. A work-group is {64 lanes, 4 subgroups}: +// the 64 lanes cover the 64 rows of one tile (coalesced reads), the 4 subgroups +// split the K-blocks and reduce through __local at the end. +// +// Weights are read from __global (coalesced) rather than image1d_buffer: the +// lm_head is read once per token with no reuse, and the Adreno texture cache +// caps such a streaming read well below the coalesced-global rate +// (see opencl_q6k_gemv_o4_shipped / x2-90 roofline notes). + +#pragma OPENCL EXTENSION cl_khr_fp16 : enable + +#ifdef cl_qcom_reqd_sub_group_size +#pragma OPENCL EXTENSION cl_qcom_reqd_sub_group_size : enable +#define ADRENO_GPU 1 +#define REQD_SUBGROUP_SIZE_64 __attribute__((qcom_reqd_sub_group_size("half"))) +#endif + +#define NSUBGROUPS 4 +#define TILE_ROWS 64 + +#if defined(ADRENO_GPU) +REQD_SUBGROUP_SIZE_64 +#endif +kernel void kernel_gemv_noshuffle_q6_K_f32_tiled( + __global uint4 * src0_ql, // tiled: 8 uint4 granules / superblock + __global uint4 * src0_qh, // tiled: 4 uint4 granules / superblock + __global char * src0_s, // tiled: 16 chars / superblock + __global half * src0_d, // tiled: 1 half / superblock + read_only image1d_buffer_t src1, // activation (RGBA f32) + global float * dst, + ulong offsetd, + int ne00, + int ne01 +) { + int grp = get_local_id(1); // subgroup index 0..3 (splits K) + int row = get_global_id(0); // output row along ne01 + int rt = row / TILE_ROWS; + int rit = row % TILE_ROWS; + + int nb = ne00 / 256; // superblocks per row + + float acc = 0.0f; + + for (int sb = grp; sb < nb; sb += NSUBGROUPS) { + int tile_blk = rt * nb + sb; // ne02 == 1 for lm_head/embed + + // d + 16 scales for this (row, superblock) + float dval = (float)src0_d[tile_blk * TILE_ROWS + rit]; + __global char * sc = src0_s + (tile_blk * TILE_ROWS + rit) * 16; + + // 32 ql-uints (8 codes/uint) + 16 qh-uints (16 codes/uint) + uint ql[32]; + uint qh[16]; + #pragma unroll + for (int g = 0; g < 8; ++g) { + uint4 v = src0_ql[(tile_blk * 8 + g) * TILE_ROWS + rit]; + ql[g*4+0] = v.x; ql[g*4+1] = v.y; ql[g*4+2] = v.z; ql[g*4+3] = v.w; + } + #pragma unroll + for (int g = 0; g < 4; ++g) { + uint4 v = src0_qh[(tile_blk * 4 + g) * TILE_ROWS + rit]; + qh[g*4+0] = v.x; qh[g*4+1] = v.y; qh[g*4+2] = v.z; qh[g*4+3] = v.w; + } + + // dequant 256 codes in canonical e-order, MAC with activation. + int act_base = sb * 64; // activation float4 pixel base (256/4) + #pragma unroll + for (int e4 = 0; e4 < 64; ++e4) { + float4 a = read_imagef(src1, act_base + e4); + #pragma unroll + for (int t = 0; t < 4; ++t) { + int e = e4 * 4 + t; + uint low4 = (ql[e >> 3] >> ((e & 7) * 4)) & 0xF; + uint hi2 = (qh[e >> 4] >> ((e & 15) * 2)) & 0x3; + int code = (int)(low4 | (hi2 << 4)) - 32; + int sidx = ((e >> 7) << 3) + (((e >> 5) & 3) << 1) + ((e >> 4) & 1); + float scale = (float)sc[sidx] * dval; + float av = (t == 0) ? a.x : (t == 1) ? a.y : (t == 2) ? a.z : a.w; + acc += (float)code * scale * av; + } + } + } + + // reduce across the NSUBGROUPS subgroups (same rit, different K-subset) + local float reduce_lm[NSUBGROUPS * TILE_ROWS]; + reduce_lm[grp * TILE_ROWS + rit] = acc; + barrier(CLK_LOCAL_MEM_FENCE); + + if (grp == 0) { + float total = reduce_lm[0 * TILE_ROWS + rit] + + reduce_lm[1 * TILE_ROWS + rit] + + reduce_lm[2 * TILE_ROWS + rit] + + reduce_lm[3 * TILE_ROWS + rit]; + dst = (global float*)((global char*)dst + offsetd); + dst[row] = total; + } +} + +// Multi-column (N=3) variant of the tiled q6_K decode GEMV, for the speculative/ +// MTP VERIFY lm_head/embed (ne1=3 = 2 drafts + 1 bonus). Identical tiled weight +// layout + unpack as the ne1=1 kernel above; each WI computes 3 output columns, +// streaming the (large) lm_head weight ONCE per superblock and reusing it across +// the 3 verify activation columns (dequant once per code, MAC into 3 accs). This +// is the lm_head analogue of the per-layer mc3 GEMV; the multiply order matches +// the ne1=1 kernel, so each column is byte-identical to a standalone tiled GEMV. +#if defined(ADRENO_GPU) +REQD_SUBGROUP_SIZE_64 +#endif +kernel void kernel_gemv_noshuffle_q6_K_f32_tiled_mc3( + __global uint4 * src0_ql, + __global uint4 * src0_qh, + __global char * src0_s, + __global half * src0_d, + read_only image1d_buffer_t src1, + global float * dst, + ulong offsetd, + int ne00, + int ne01 +) { + int grp = get_local_id(1); + int row = get_global_id(0); + int rt = row / TILE_ROWS; + int rit = row % TILE_ROWS; + + int nb = ne00 / 256; + int col_stride = ne00 / 4; // activation float4 pixels per column + + float acc0 = 0.0f, acc1 = 0.0f, acc2 = 0.0f; + + for (int sb = grp; sb < nb; sb += NSUBGROUPS) { + int tile_blk = rt * nb + sb; + + float dval = (float)src0_d[tile_blk * TILE_ROWS + rit]; + __global char * sc = src0_s + (tile_blk * TILE_ROWS + rit) * 16; + + uint ql[32]; + uint qh[16]; + #pragma unroll + for (int g = 0; g < 8; ++g) { + uint4 v = src0_ql[(tile_blk * 8 + g) * TILE_ROWS + rit]; + ql[g*4+0] = v.x; ql[g*4+1] = v.y; ql[g*4+2] = v.z; ql[g*4+3] = v.w; + } + #pragma unroll + for (int g = 0; g < 4; ++g) { + uint4 v = src0_qh[(tile_blk * 4 + g) * TILE_ROWS + rit]; + qh[g*4+0] = v.x; qh[g*4+1] = v.y; qh[g*4+2] = v.z; qh[g*4+3] = v.w; + } + + int act_base = sb * 64; + #pragma unroll + for (int e4 = 0; e4 < 64; ++e4) { + float4 a0 = read_imagef(src1, 0*col_stride + act_base + e4); + float4 a1 = read_imagef(src1, 1*col_stride + act_base + e4); + float4 a2 = read_imagef(src1, 2*col_stride + act_base + e4); + #pragma unroll + for (int t = 0; t < 4; ++t) { + int e = e4 * 4 + t; + uint low4 = (ql[e >> 3] >> ((e & 7) * 4)) & 0xF; + uint hi2 = (qh[e >> 4] >> ((e & 15) * 2)) & 0x3; + int code = (int)(low4 | (hi2 << 4)) - 32; + int sidx = ((e >> 7) << 3) + (((e >> 5) & 3) << 1) + ((e >> 4) & 1); + float w = (float)code * ((float)sc[sidx] * dval); // dequant+scale once + float av0 = (t == 0) ? a0.x : (t == 1) ? a0.y : (t == 2) ? a0.z : a0.w; + float av1 = (t == 0) ? a1.x : (t == 1) ? a1.y : (t == 2) ? a1.z : a1.w; + float av2 = (t == 0) ? a2.x : (t == 1) ? a2.y : (t == 2) ? a2.z : a2.w; + acc0 += w * av0; + acc1 += w * av1; + acc2 += w * av2; + } + } + } + + local float4 reduce_lm[NSUBGROUPS * TILE_ROWS]; + reduce_lm[grp * TILE_ROWS + rit] = (float4)(acc0, acc1, acc2, 0.0f); + barrier(CLK_LOCAL_MEM_FENCE); + + if (grp == 0) { + float4 total = reduce_lm[0 * TILE_ROWS + rit] + + reduce_lm[1 * TILE_ROWS + rit] + + reduce_lm[2 * TILE_ROWS + rit] + + reduce_lm[3 * TILE_ROWS + rit]; + dst = (global float*)((global char*)dst + offsetd); + // dst column-major [ne01 rows x 3 cols]: (row, col) at col*ne01 + row + dst[0*ne01 + row] = total.x; + dst[1*ne01 + row] = total.y; + dst[2*ne01 + row] = total.z; + } +} diff --git a/ggml/src/ggml-opencl/kernels/gemv_noshuffle_q8_0_f32.cl b/ggml/src/ggml-opencl/kernels/gemv_noshuffle_q8_0_f32.cl index 09bae2d555e..6f6d7425c65 100644 --- a/ggml/src/ggml-opencl/kernels/gemv_noshuffle_q8_0_f32.cl +++ b/ggml/src/ggml-opencl/kernels/gemv_noshuffle_q8_0_f32.cl @@ -118,6 +118,87 @@ elem = (char)((bits8.s7 & 0xFF000000) >> 24); \ total_sums += convert_int(elem) * scale * shared_y; \ +// ============================================================================ +// Split-K variant for small-M decode GEMVs. +// ---------------------------------------------------------------------------- +// The base kernel below puts one output row per lane and splits K only across +// the N_SIMDGROUP subgroups of a single workgroup, so M=512 yields M/64 = 8 +// workgroups -- half the compute units on a 16-CU X2 sit idle, and the kernel +// measures ~48 GB/s against the ~122 GB/s the larger projections reach in the +// same graph. Here each (kslice, subgroup) pair reduces a disjoint set of +// K-blocks into partial[kslice * M + row]; kernel_gemv_splitk_reduce_f32 (in +// gemv_noshuffle_q4_k_f32.cl) sums the slices. Same operand order within a +// slice as the base kernel; only the cross-slice grouping differs. +// +// Placed BEFORE the base kernel deliberately: on A6X no kernel may be defined +// after one that uses a subgroup builtin, or it silently miscompiles. +// ============================================================================ +#ifdef ADRENO_GPU +REQD_SUBGROUP_SIZE_64 +#endif +__kernel void kernel_gemv_noshuffle_q8_0_f32_splitk( + __read_only image1d_buffer_t src0_q, // quantized A (weights) + global half * src0_d, // A scales + __read_only image1d_buffer_t src1, // B (activations) + global float * partial, // [ksplit * M], slice-major + int ne00, // K + int ne01) // M +{ + uint groupId = get_local_id(1); + uint gid = get_global_id(0); + ushort slid = get_sub_group_local_id(); + uint nsg = get_local_size(1); + uint ksplit = get_num_groups(1); + uint kslice = get_group_id(1); + + uint K = ne00; + uint M = ne01; + + uint LINE_STRIDE_A = M; + uint BLOCK_STRIDE_A = 8 * M; // physical, independent of the K-split + + __private uint8 regA; + __private half regS; + __private float8 regB; + __private float totalSum = (float)(0.0f); + + #pragma unroll 1 + for (uint k = kslice * nsg + groupId; k < (K / QK8_0); k += ksplit * nsg) { + regS = src0_d[gid + k * LINE_STRIDE_A]; + if (slid < 4) { + regB.s0123 = read_imagef(src1, (slid * 2 + k * 8)); + regB.s4567 = read_imagef(src1, (1 + slid * 2 + k * 8)); + } + regA.s0 = read_imageui(src0_q, (gid + k * BLOCK_STRIDE_A + LINE_STRIDE_A * 0)).x; + regA.s1 = read_imageui(src0_q, (gid + k * BLOCK_STRIDE_A + LINE_STRIDE_A * 1)).x; + regA.s2 = read_imageui(src0_q, (gid + k * BLOCK_STRIDE_A + LINE_STRIDE_A * 2)).x; + regA.s3 = read_imageui(src0_q, (gid + k * BLOCK_STRIDE_A + LINE_STRIDE_A * 3)).x; + regA.s4 = read_imageui(src0_q, (gid + k * BLOCK_STRIDE_A + LINE_STRIDE_A * 4)).x; + regA.s5 = read_imageui(src0_q, (gid + k * BLOCK_STRIDE_A + LINE_STRIDE_A * 5)).x; + regA.s6 = read_imageui(src0_q, (gid + k * BLOCK_STRIDE_A + LINE_STRIDE_A * 6)).x; + regA.s7 = read_imageui(src0_q, (gid + k * BLOCK_STRIDE_A + LINE_STRIDE_A * 7)).x; + + dequantizeBlockAccum_ns_sgbroadcast_1(totalSum, regA, convert_float(regS), regB); + } + + // Intra-workgroup reduce across this K-slice's subgroups. Sized for + // nsg <= 8; the host never dispatches more. + __local float reduceLM[SIMDGROUP_WIDTH * 7]; + if (groupId > 0) { + reduceLM[SIMDGROUP_WIDTH * (groupId - 1) + slid] = totalSum; + } + barrier(CLK_LOCAL_MEM_FENCE); + if (groupId == 0) { + for (uint i = 0; i < nsg - 1; ++i) { + totalSum += reduceLM[SIMDGROUP_WIDTH * i + slid]; + } + // x-grid is padded to CEIL_DIV(M,wave)*wave; guard the tail rows. + if (gid < M) { + partial[kslice * M + gid] = totalSum; + } + } +} + #ifdef ADRENO_GPU REQD_SUBGROUP_SIZE_64 #endif diff --git a/ggml/src/ggml-opencl/kernels/mul_mm_f32_f32_l4_lm.cl b/ggml/src/ggml-opencl/kernels/mul_mm_f32_f32_l4_lm.cl index d7d5ba647e7..9dc9862bef6 100644 --- a/ggml/src/ggml-opencl/kernels/mul_mm_f32_f32_l4_lm.cl +++ b/ggml/src/ggml-opencl/kernels/mul_mm_f32_f32_l4_lm.cl @@ -145,3 +145,52 @@ kernel void kernel_mul_mm_f32_f32_l4_lm( } } } + +// Multi-column f32 GEMV for the small-N (spec/MTP verify) batch. The tiled GEMM +// above always computes a full BM x BN = 64 x 64 output tile, so at ne11=3 with a +// skinny weight (e.g. GDN ssm_alpha/ssm_beta, M=32) it launches ONE under-occupied +// workgroup at ~2.3% tile utilization. This kernel assigns one 64-thread workgroup +// per output element (m,n): the 64 threads split the K reduction (float4) and +// tree-reduce in __local (no subgroup ops -> portable). ne01*ne11 workgroups. +// Weight row is re-read per column (N small -> negligible). Summation order differs +// from the tiled GEMM (lane-strided + tree) -> f32-exact-ish, not bit-identical. +kernel void kernel_gemv_f32_f32_mc( + global float * src0, ulong offset0, // weight: row m at m*stride_a (elements) + global float * src1, ulong offset1, // activations: col n at n*stride_b + global float * dst, ulong offsetd, // dst [M x N] col-major: (m,n) at n*stride_d+m + int ne00, // K + int ne01, // M + int ne11, // N + int stride_a, // weight row stride (elements) = K + int stride_b, // activation col stride (elements) = K + int stride_d) // dst column stride (elements) = M +{ + src0 = (global float*)((global char*)src0 + offset0); + src1 = (global float*)((global char*)src1 + offset1); + dst = (global float*)((global char*)dst + offsetd); + + uint lane = get_local_id(0); // 0..63 + uint out = get_global_id(1); // 0 .. ne01*ne11 - 1 + uint m = out % (uint)ne01; + uint n = out / (uint)ne01; + + global float4 * wrow = (global float4*)(src0 + (ulong)m * (uint)stride_a); + global float4 * xcol = (global float4*)(src1 + (ulong)n * (uint)stride_b); + uint k4 = (uint)ne00 >> 2; + + float acc = 0.0f; + for (uint k = lane; k < k4; k += 64) { + float4 w = wrow[k]; + float4 x = xcol[k]; + acc += w.s0*x.s0 + w.s1*x.s1 + w.s2*x.s2 + w.s3*x.s3; + } + + local float red[64]; + red[lane] = acc; + barrier(CLK_LOCAL_MEM_FENCE); + for (uint s = 32; s > 0; s >>= 1) { + if (lane < s) red[lane] += red[lane + s]; + barrier(CLK_LOCAL_MEM_FENCE); + } + if (lane == 0) dst[(ulong)n * (uint)stride_d + m] = red[0]; +} diff --git a/ggml/src/ggml-opencl/kernels/mul_mv_f16_f32_mrow.cl b/ggml/src/ggml-opencl/kernels/mul_mv_f16_f32_mrow.cl new file mode 100644 index 00000000000..9a7627cf9be --- /dev/null +++ b/ggml/src/ggml-opencl/kernels/mul_mv_f16_f32_mrow.cl @@ -0,0 +1,306 @@ +#pragma OPENCL EXTENSION cl_khr_fp16 : enable + +#ifdef cl_intel_subgroups +#pragma OPENCL EXTENSION cl_intel_subgroups : enable +#else +#pragma OPENCL EXTENSION cl_khr_subgroups : enable +#endif + +#ifdef cl_intel_required_subgroup_size +#pragma OPENCL EXTENSION cl_intel_required_subgroup_size : enable +#define INTEL_GPU 1 +#define REQD_SUBGROUP_SIZE_16 __attribute__((intel_reqd_sub_group_size(16))) +#define REQD_SUBGROUP_SIZE_32 __attribute__((intel_reqd_sub_group_size(32))) +#elif defined(cl_qcom_reqd_sub_group_size) +#pragma OPENCL EXTENSION cl_qcom_reqd_sub_group_size : enable +#define ADRENO_GPU 1 +#define REQD_SUBGROUP_SIZE_64 __attribute__((qcom_reqd_sub_group_size("half"))) +#define REQD_SUBGROUP_SIZE_128 __attribute__((qcom_reqd_sub_group_size("full"))) +#endif + +// Multi-row f16xf32 GEMV for the DECODE path (single token, ne11*ne12 small). +// The legacy kernel_mul_mat_f16_f32_1row runs ONE 64-lane subgroup per workgroup = +// one output row per WG, which caps memory-level parallelism at roughly half of +// LPDDR5x peak. This variant packs MROW subgroups per workgroup, each +// computing a distinct output row, so a WG keeps 64*MROW loads in flight. The +// activation column y (shared by every output row) is staged into __local ONCE per +// WG and reused across the MROW rows, cutting redundant activation reads. Used for +// the f16 attention projections (Q/K/V/O) and lm_head, which dominate decode. +// Numerically equivalent to _1row (same f16->f32 widening, same float4 partial sums, +// same subgroup-reduce order), so byte-identical to the per-op path. + +#define MROW 16 + +#ifdef ADRENO_GPU +REQD_SUBGROUP_SIZE_64 +#endif +kernel void kernel_mul_mat_f16_f32_mrow( + global char * src0, + ulong offset0, + global char * src1, + ulong offset1, + global float * dst, + ulong offsetd, + int ne00, + int ne01, + int ne02, + ulong nb00, + ulong nb01, + ulong nb02, + ulong nb03, + int ne10, + int ne11, + int ne12, + ulong nb10, + ulong nb11, + ulong nb12, + ulong nb13, + int ne0, + int ne1, + int r2, + int r3, + __local float * ysh +) { + src0 = (global char*)((global char*)src0 + offset0); + src1 = (global char*)((global char*)src1 + offset1); + dst = (global float*)((global char*)dst + offsetd); + + int r0 = get_group_id(0) * MROW + get_local_id(1); // output row + int r1 = get_group_id(1); // token (ne11) + int im = get_group_id(2); + int lid = get_sub_group_local_id(); // 0..63 + int nsg = get_local_size(1); // == MROW + + int i12 = im % ne12; + int i13 = im / ne12; + + ulong offset_src1 = r1*nb11 + (i12)*nb12 + (i13)*nb13; + global float * y = (global float *) (src1 + offset_src1); + + // Cooperatively stage the activation column (ne00 floats) into __local once per + // WG and reuse across the MROW rows. Staging is the actual win here: dropping it + // (each subgroup re-reading y from global) regresses below the 1-row kernel. + for (int i = get_local_id(1)*get_sub_group_size() + lid; i < ne00; i += nsg*get_sub_group_size()) { + ysh[i] = y[i]; + } + barrier(CLK_LOCAL_MEM_FENCE); + + if (r0 >= ne01) { + return; + } + + ulong offset_src0 = r0*nb01 + (i12/r2)*nb02 + (i13/r3)*nb03; + global half * x = (global half *) (src0 + offset_src0); + + // The vector path below casts the row pointer to half4, which must be 8-byte aligned. + // A row address is r0*nb01 + ..., and a permuted or strided src0 leaves nb01/nb02/nb03 + // unconstrained -- ne00 % 4 == 0 bounds the element count per row, not the byte stride + // between rows. Take the vector path only when this work-item's row is actually + // aligned; the scalar loop below has no such requirement. + const bool row_aligned = (((ulong) x) & 7) == 0; + + float sumf = 0.0f; + if (ne00 < 128 || !row_aligned) { + for (int i = lid; i < ne00; i += get_sub_group_size()) { + sumf += (float) x[i] * ysh[i]; + } + float all_sum = sub_group_reduce_add(sumf); + if (lid == 0) { + dst[im*ne1*ne0 + r1*ne0 + r0] = all_sum; + } + } else { + global half4 * x4 = (global half4 *) x; + __local float4 * ysh4 = (__local float4 *) ysh; + for (int i = lid; i < ne00/4; i += get_sub_group_size()) { + float4 yv = ysh4[i]; + sumf += (float) x4[i].s0 * yv.s0; + sumf += (float) x4[i].s1 * yv.s1; + sumf += (float) x4[i].s2 * yv.s2; + sumf += (float) x4[i].s3 * yv.s3; + } + float all_sum = sub_group_reduce_add(sumf); + if (lid == 0) { + for (int i = 4*(ne00/4); i < ne00; ++i) { + all_sum += (float) x[i] * ysh[i]; + } + dst[im*ne1*ne0 + r1*ne0 + r0] = all_sum; + } + } +} + +// Register-blocked variant: each 64-lane subgroup accumulates RPT consecutive +// output rows instead of one. The staged activation is reused across all RPT rows, +// and each lane keeps RPT independent weight loads in flight per column step -> +// more memory-level parallelism on the streaming f16 weight read (the BW limiter), +// plus RPT fewer staging barriers per output row. Per-row reduction order is +// identical to _mrow, so byte-identical to the per-op path. Dispatch guarantees +// ne00 >= 128 and ne00 % 4 == 0, so only the half4 path is needed (no tail). +#define MROW_RB_BODY(RPT) \ + src0 = (global char*)((global char*)src0 + offset0); \ + src1 = (global char*)((global char*)src1 + offset1); \ + dst = (global float*)((global char*)dst + offsetd); \ + int r0b = (get_group_id(0) * get_local_size(1) + get_local_id(1)) * (RPT); \ + int r1 = get_group_id(1); \ + int im = get_group_id(2); \ + int lid = get_sub_group_local_id(); \ + int nsg = get_local_size(1); \ + int i12 = im % ne12; \ + int i13 = im / ne12; \ + ulong off_y = r1*nb11 + i12*nb12 + i13*nb13; \ + global float * y = (global float *) (src1 + off_y); \ + for (int i = get_local_id(1)*get_sub_group_size() + lid; i < ne00; \ + i += nsg*get_sub_group_size()) { \ + ysh[i] = y[i]; \ + } \ + barrier(CLK_LOCAL_MEM_FENCE); \ + __local float4 * ysh4 = (__local float4 *) ysh; \ + global half4 * xr[RPT]; \ + _Pragma("unroll") \ + for (int rr = 0; rr < (RPT); ++rr) { \ + int row = r0b + rr; \ + if (row > ne01 - 1) row = ne01 - 1; \ + xr[rr] = (global half4 *) (src0 + (ulong)row*nb01 + (i12/r2)*nb02 + (i13/r3)*nb03); \ + } \ + float sumf[RPT]; \ + _Pragma("unroll") \ + for (int rr = 0; rr < (RPT); ++rr) sumf[rr] = 0.0f; \ + for (int i = lid; i < ne00/4; i += get_sub_group_size()) { \ + float4 yv = ysh4[i]; \ + _Pragma("unroll") \ + for (int rr = 0; rr < (RPT); ++rr) { \ + half4 xv = xr[rr][i]; \ + sumf[rr] += (float) xv.s0 * yv.s0 + (float) xv.s1 * yv.s1 \ + + (float) xv.s2 * yv.s2 + (float) xv.s3 * yv.s3; \ + } \ + } \ + _Pragma("unroll") \ + for (int rr = 0; rr < (RPT); ++rr) { \ + float s = sub_group_reduce_add(sumf[rr]); \ + int row = r0b + rr; \ + if (lid == 0 && row < ne01) { \ + dst[im*ne1*ne0 + r1*ne0 + row] = s; \ + } \ + } + +// half8 (128-bit) load variant: Adreno's load/store unit issues 128-bit +// transactions, so half4 (64-bit) loads may leave the load path half-idle. This +// processes 8 weight elements per lane per step via half8. Accumulation groups +// elements in 8s rather than 4s, so it is NOT bit-identical to _1row (float add is +// non-associative) -- experimental BW probe, gate on ne00 % 8 == 0. +#define MROW_H8_BODY(RPT) \ + src0 = (global char*)((global char*)src0 + offset0); \ + src1 = (global char*)((global char*)src1 + offset1); \ + dst = (global float*)((global char*)dst + offsetd); \ + int r0b = (get_group_id(0) * get_local_size(1) + get_local_id(1)) * (RPT); \ + int r1 = get_group_id(1); \ + int im = get_group_id(2); \ + int lid = get_sub_group_local_id(); \ + int nsg = get_local_size(1); \ + int i12 = im % ne12; \ + int i13 = im / ne12; \ + ulong off_y = r1*nb11 + i12*nb12 + i13*nb13; \ + global float * y = (global float *) (src1 + off_y); \ + for (int i = get_local_id(1)*get_sub_group_size() + lid; i < ne00; \ + i += nsg*get_sub_group_size()) { \ + ysh[i] = y[i]; \ + } \ + barrier(CLK_LOCAL_MEM_FENCE); \ + __local float4 * ysh4 = (__local float4 *) ysh; \ + global half8 * xr[RPT]; \ + _Pragma("unroll") \ + for (int rr = 0; rr < (RPT); ++rr) { \ + int row = r0b + rr; \ + if (row > ne01 - 1) row = ne01 - 1; \ + xr[rr] = (global half8 *) (src0 + (ulong)row*nb01 + (i12/r2)*nb02 + (i13/r3)*nb03); \ + } \ + float sumf[RPT]; \ + _Pragma("unroll") \ + for (int rr = 0; rr < (RPT); ++rr) sumf[rr] = 0.0f; \ + for (int i = lid; i < ne00/8; i += get_sub_group_size()) { \ + float4 y0 = ysh4[2*i]; \ + float4 y1 = ysh4[2*i + 1]; \ + _Pragma("unroll") \ + for (int rr = 0; rr < (RPT); ++rr) { \ + half8 xv = xr[rr][i]; \ + sumf[rr] += (float) xv.s0 * y0.s0 + (float) xv.s1 * y0.s1 \ + + (float) xv.s2 * y0.s2 + (float) xv.s3 * y0.s3 \ + + (float) xv.s4 * y1.s0 + (float) xv.s5 * y1.s1 \ + + (float) xv.s6 * y1.s2 + (float) xv.s7 * y1.s3; \ + } \ + } \ + _Pragma("unroll") \ + for (int rr = 0; rr < (RPT); ++rr) { \ + float s = sub_group_reduce_add(sumf[rr]); \ + int row = r0b + rr; \ + if (lid == 0 && row < ne01) { \ + dst[im*ne1*ne0 + r1*ne0 + row] = s; \ + } \ + } + +#ifdef ADRENO_GPU +REQD_SUBGROUP_SIZE_64 +#endif +kernel void kernel_mul_mat_f16_f32_mrow_h8( + global char * src0, ulong offset0, + global char * src1, ulong offset1, + global float * dst, ulong offsetd, + int ne00, int ne01, int ne02, + ulong nb00, ulong nb01, ulong nb02, ulong nb03, + int ne10, int ne11, int ne12, + ulong nb10, ulong nb11, ulong nb12, ulong nb13, + int ne0, int ne1, int r2, int r3, + __local float * ysh +) { + MROW_H8_BODY(1) +} + +#ifdef ADRENO_GPU +REQD_SUBGROUP_SIZE_64 +#endif +kernel void kernel_mul_mat_f16_f32_mrow_h8r2( + global char * src0, ulong offset0, + global char * src1, ulong offset1, + global float * dst, ulong offsetd, + int ne00, int ne01, int ne02, + ulong nb00, ulong nb01, ulong nb02, ulong nb03, + int ne10, int ne11, int ne12, + ulong nb10, ulong nb11, ulong nb12, ulong nb13, + int ne0, int ne1, int r2, int r3, + __local float * ysh +) { + MROW_H8_BODY(2) +} + +#ifdef ADRENO_GPU +REQD_SUBGROUP_SIZE_64 +#endif +kernel void kernel_mul_mat_f16_f32_mrow_r2( + global char * src0, ulong offset0, + global char * src1, ulong offset1, + global float * dst, ulong offsetd, + int ne00, int ne01, int ne02, + ulong nb00, ulong nb01, ulong nb02, ulong nb03, + int ne10, int ne11, int ne12, + ulong nb10, ulong nb11, ulong nb12, ulong nb13, + int ne0, int ne1, int r2, int r3, + __local float * ysh +) { + MROW_RB_BODY(2) +} + +#ifdef ADRENO_GPU +REQD_SUBGROUP_SIZE_64 +#endif +kernel void kernel_mul_mat_f16_f32_mrow_r4( + global char * src0, ulong offset0, + global char * src1, ulong offset1, + global float * dst, ulong offsetd, + int ne00, int ne01, int ne02, + ulong nb00, ulong nb01, ulong nb02, ulong nb03, + int ne10, int ne11, int ne12, + ulong nb10, ulong nb11, ulong nb12, ulong nb13, + int ne0, int ne1, int r2, int r3, + __local float * ysh +) { + MROW_RB_BODY(4) +} diff --git a/ggml/src/ggml-opencl/kernels/rms_norm.cl b/ggml/src/ggml-opencl/kernels/rms_norm.cl index 4b18d17d6f8..99085625a4c 100644 --- a/ggml/src/ggml-opencl/kernels/rms_norm.cl +++ b/ggml/src/ggml-opencl/kernels/rms_norm.cl @@ -188,3 +188,182 @@ kernel void kernel_rms_norm_mul( y[i00] = (x[i00] * scale) * f[i00%(ne10/4)]; } } + +//------------------------------------------------------------------------------ +// rms_norm + mul (norm weight) + add (residual), fused. Mirrors +// kernel_rms_norm_mul with an extra residual operand src2: computes +// y = (rmsnorm(x) * w) + g +// in one dispatch, removing one kernel launch + one global round-trip per +// residual block (the dominant per-layer adjacency on Gemma matformers). +//------------------------------------------------------------------------------ +kernel void kernel_rms_norm_mul_add( + global char * src0, + ulong offset0, + global char * src1, + ulong offset1, + global char * src2, + ulong offset2, + global char * dst, + ulong offsetd, + int ne00, + int ne01, + int ne02, + int ne03, + ulong nb01, + ulong nb02, + ulong nb03, + int ne10, + int ne11, + int ne12, + int ne13, + ulong nb11, + ulong nb12, + ulong nb13, + int ne20, + int ne21, + int ne22, + int ne23, + ulong nb21, + ulong nb22, + ulong nb23, + ulong nb1, + ulong nb2, + ulong nb3, + float eps, + local float * sum +) { + src0 = src0 + offset0; + src1 = src1 + offset1; + src2 = src2 + offset2; + dst = dst + offsetd; + + if (get_sub_group_id() == 0) { + sum[get_sub_group_local_id()] = 0.0f; + } + + int i03 = get_group_id(2); + int i02 = get_group_id(1); + int i01 = get_group_id(0); + + global float4 * x = (global float4 *) (src0 + i03*nb03 + i02*nb02 + i01*nb01); + global float4 * f = (global float4 *) (src1 + (i03%ne13)*nb13 + (i02%ne12)*nb12 + (i01%ne11)*nb11); + global float4 * g = (global float4 *) (src2 + (i03%ne23)*nb23 + (i02%ne22)*nb22 + (i01%ne21)*nb21); + + float sumf = 0; + + for (int i00 = get_local_id(0); i00 < ne00/4; i00 += get_local_size(0)) { + sumf += dot(x[i00], x[i00]); + } + sumf = sub_group_reduce_add(sumf); + + barrier(CLK_LOCAL_MEM_FENCE); + + if (get_sub_group_local_id() == 0) { + sum[get_sub_group_id()] = sumf; + } + + barrier(CLK_LOCAL_MEM_FENCE); + + sumf = sum[get_sub_group_local_id()]; + sumf = sub_group_reduce_add(sumf); + + float mean = sumf / ne00; + float scale = 1.0f/sqrt(mean + eps); + + global float4 * y = (global float4 *) (dst + i03*nb3 + i02*nb2 + i01*nb1); + for (int i00 = get_local_id(0); i00 < ne00/4; i00 += get_local_size(0)) { + y[i00] = (x[i00] * scale) * f[i00%(ne10/4)] + g[i00%(ne20/4)]; + } +} + +//------------------------------------------------------------------------------ +// rms_norm + mul(norm weight) + add(residual) + mul(scalar scale), fused. +// Computes y = ((rmsnorm(x) * w) + g) * s, where s is a broadcast SCALAR (e.g. +// Gemma-4 layer_output_scale). Folds the trailing per-layer l_out scale-mul into +// the residual-norm kernel: one extra dispatch + global round-trip saved per +// layer. src3 points at the single scale value. +//------------------------------------------------------------------------------ +kernel void kernel_rms_norm_mul_add_scale( + global char * src0, + ulong offset0, + global char * src1, + ulong offset1, + global char * src2, + ulong offset2, + global char * src3, + ulong offset3, + global char * dst, + ulong offsetd, + int ne00, + int ne01, + int ne02, + int ne03, + ulong nb01, + ulong nb02, + ulong nb03, + int ne10, + int ne11, + int ne12, + int ne13, + ulong nb11, + ulong nb12, + ulong nb13, + int ne20, + int ne21, + int ne22, + int ne23, + ulong nb21, + ulong nb22, + ulong nb23, + ulong nb1, + ulong nb2, + ulong nb3, + float eps, + local float * sum +) { + src0 = src0 + offset0; + src1 = src1 + offset1; + src2 = src2 + offset2; + src3 = src3 + offset3; + dst = dst + offsetd; + + const float sc = *((global float *) src3); + + if (get_sub_group_id() == 0) { + sum[get_sub_group_local_id()] = 0.0f; + } + + int i03 = get_group_id(2); + int i02 = get_group_id(1); + int i01 = get_group_id(0); + + global float4 * x = (global float4 *) (src0 + i03*nb03 + i02*nb02 + i01*nb01); + global float4 * f = (global float4 *) (src1 + (i03%ne13)*nb13 + (i02%ne12)*nb12 + (i01%ne11)*nb11); + global float4 * g = (global float4 *) (src2 + (i03%ne23)*nb23 + (i02%ne22)*nb22 + (i01%ne21)*nb21); + + float sumf = 0; + + for (int i00 = get_local_id(0); i00 < ne00/4; i00 += get_local_size(0)) { + sumf += dot(x[i00], x[i00]); + } + sumf = sub_group_reduce_add(sumf); + + barrier(CLK_LOCAL_MEM_FENCE); + + if (get_sub_group_local_id() == 0) { + sum[get_sub_group_id()] = sumf; + } + + barrier(CLK_LOCAL_MEM_FENCE); + + sumf = sum[get_sub_group_local_id()]; + sumf = sub_group_reduce_add(sumf); + + float mean = sumf / ne00; + float scale = 1.0f/sqrt(mean + eps); + + global float4 * y = (global float4 *) (dst + i03*nb3 + i02*nb2 + i01*nb1); + for (int i00 = get_local_id(0); i00 < ne00/4; i00 += get_local_size(0)) { + y[i00] = ((x[i00] * scale) * f[i00%(ne10/4)] + g[i00%(ne20/4)]) * sc; + } +} From 36f170e5fe39b566eb2afe1d2ed74b567cd54fc9 Mon Sep 17 00:00:00 2001 From: Ozymandias_EBON <112784549+johnkarlhill@users.noreply.github.com> Date: Thu, 3 Sep 2026 21:45:53 -0500 Subject: [PATCH 095/104] SYCL: Refactor GGML_SYCL_ENABLE_MKL_FA to global var (llama/26863) --- ggml/src/ggml-sycl/common.hpp | 1 + ggml/src/ggml-sycl/fattn.cpp | 3 +-- ggml/src/ggml-sycl/ggml-sycl.cpp | 3 +++ 3 files changed, 5 insertions(+), 2 deletions(-) diff --git a/ggml/src/ggml-sycl/common.hpp b/ggml/src/ggml-sycl/common.hpp index b5f75ca548c..9f2a27b18e0 100644 --- a/ggml/src/ggml-sycl/common.hpp +++ b/ggml/src/ggml-sycl/common.hpp @@ -68,6 +68,7 @@ extern int g_ggml_sycl_enable_flash_attention; extern int g_ggml_sycl_dev2dev_memcpy; extern int g_ggml_sycl_fa_onednn; extern int g_ggml_sycl_fa_onednn_max_kv; +extern int g_ggml_sycl_enable_mkl_fa; #define CHECK_TRY_ERROR(expr) \ diff --git a/ggml/src/ggml-sycl/fattn.cpp b/ggml/src/ggml-sycl/fattn.cpp index b73e6d46ffa..394cda593f7 100644 --- a/ggml/src/ggml-sycl/fattn.cpp +++ b/ggml/src/ggml-sycl/fattn.cpp @@ -146,14 +146,13 @@ static best_fattn_kernel ggml_sycl_get_best_fattn_kernel(const int device, const // Set GGML_SYCL_ENABLE_MKL_FA=0 to force TILE/VEC path for A/B testing. // Example: GGML_SYCL_ENABLE_MKL_FA=0 llama-cli -m model.gguf -fa -ngl 99 ... // Note: MKL GEMM calls are incompatible with SYCL graph capture replay. - static int mkl_enable = ggml_sycl_get_env("GGML_SYCL_ENABLE_MKL_FA", 1); // MKL is validated for the mainstream GQA envelope: grouped-query // (gqa_ratio >= 2), head_dim a multiple of 64 in [64,512] with matching // K/V head size, mask, no sinks/ALiBi/softcap. Gemma's global layers use // head_dim 512, so the cap must include it. Head sizes not a multiple of // 64 (72/80/96), MHA (gqa_ratio == 1), and MLA (DKQ != DV, e.g. 576/512) // fall through to TILE/VEC; see follow-up work. - if (mkl_enable == 1 && mask && !sinks && gqa_ratio >= 2 && + if (g_ggml_sycl_enable_mkl_fa == 1 && mask && !sinks && gqa_ratio >= 2 && Q->ne[0] >= 64 && Q->ne[0] <= 512 && Q->ne[0] % 64 == 0 && Q->ne[0] == V->ne[0] && Q->ne[1] >= 32 && K->ne[1] >= 1024 && diff --git a/ggml/src/ggml-sycl/ggml-sycl.cpp b/ggml/src/ggml-sycl/ggml-sycl.cpp index 60a9c6015d3..feced748cec 100644 --- a/ggml/src/ggml-sycl/ggml-sycl.cpp +++ b/ggml/src/ggml-sycl/ggml-sycl.cpp @@ -96,6 +96,7 @@ int g_ggml_sycl_enable_graph = 0; int g_ggml_sycl_enable_dnn = 1; int g_ggml_sycl_fa_onednn = 1; int g_ggml_sycl_fa_onednn_max_kv = 0; +int g_ggml_sycl_enable_mkl_fa = 1; int g_ggml_sycl_enable_vmm = 1; int g_ggml_sycl_enable_fusion = 1; int g_ggml_sycl_enable_esimd = 1; @@ -333,6 +334,7 @@ static void ggml_check_sycl() try { g_ggml_sycl_enable_dnn = ggml_sycl_get_env("GGML_SYCL_ENABLE_DNN", 1); g_ggml_sycl_fa_onednn = ggml_sycl_get_env("GGML_SYCL_FA_ONEDNN", 1); g_ggml_sycl_fa_onednn_max_kv = ggml_sycl_get_env("GGML_SYCL_FA_ONEDNN_MAX_KV", 0); + g_ggml_sycl_enable_mkl_fa = ggml_sycl_get_env("GGML_SYCL_ENABLE_MKL_FA", 1); g_ggml_sycl_enable_vmm = ggml_sycl_get_env("GGML_SYCL_ENABLE_VMM", 1); g_ggml_sycl_enable_fusion = ggml_sycl_get_env("GGML_SYCL_ENABLE_FUSION", 1); g_ggml_sycl_enable_esimd = ggml_sycl_get_env("GGML_SYCL_ENABLE_ESIMD", 1); @@ -418,6 +420,7 @@ static void ggml_check_sycl() try { GGML_LOG_INFO(" GGML_SYCL_FA_ONEDNN: %d\n", g_ggml_sycl_fa_onednn); #endif GGML_LOG_INFO(" GGML_SYCL_FA_ONEDNN_MAX_KV: %d\n", g_ggml_sycl_fa_onednn_max_kv); + GGML_LOG_INFO(" GGML_SYCL_ENABLE_MKL_FA: %d\n", g_ggml_sycl_enable_mkl_fa); #ifdef SYCL_FLASH_ATTN GGML_LOG_INFO(" GGML_SYCL_ENABLE_FLASH_ATTN: %d\n", g_ggml_sycl_enable_flash_attention); #else From d1e0e6491ff3735a09db4b160524e0dd4488a4e4 Mon Sep 17 00:00:00 2001 From: Frosty40 Date: Thu, 3 Sep 2026 23:05:40 -0500 Subject: [PATCH 096/104] sycl: fuse rms_norm+mul+add and add+add residual chains (llama/27610) Fuse RMS_NORM+MUL+ADD and ADD+ADD under GGML_SYCL_ENABLE_FUSION. ADD+ADD uses the same binbcast indexing and type matrix as standalone add() (f32, f16, f16/f32, i32, i16, bf16, including broadcast and non-contiguous). Unsupported combinations fall back to two add() launches. --- ggml/src/ggml-sycl/binbcast.cpp | 292 +++++++++++++++++++++++++++++++ ggml/src/ggml-sycl/binbcast.hpp | 30 ++++ ggml/src/ggml-sycl/fusion.cpp | 45 ++++- ggml/src/ggml-sycl/ggml-sycl.cpp | 12 ++ ggml/src/ggml-sycl/norm.cpp | 151 +++++++++++++++- ggml/src/ggml-sycl/norm.hpp | 2 + 6 files changed, 528 insertions(+), 4 deletions(-) diff --git a/ggml/src/ggml-sycl/binbcast.cpp b/ggml/src/ggml-sycl/binbcast.cpp index 306eeddc0c0..f2f7c4cde60 100644 --- a/ggml/src/ggml-sycl/binbcast.cpp +++ b/ggml/src/ggml-sycl/binbcast.cpp @@ -1,5 +1,6 @@ #include "binbcast.hpp" +#include #include #include #include @@ -356,3 +357,294 @@ void ggml_sycl_repeat(ggml_backend_sycl_context & ctx, ggml_tensor * dst) { ggml_sycl_op_repeat(ctx, dst); } +// fused ADD+ADD: dst = (src0 + src1) + src2. Same indexing as k_bin_bcast, so mixed +// types, broadcast, and non-contiguous layouts that add() already handles also fuse. +template +static void k_bin_bcast3(const src0_t * src0, const src1_t * src1, const src2_t * src2, dst_t * dst, + int ne0, int ne1, int ne2, int ne3, + int ne10, int ne11, int ne12, int ne13, + int ne20, int ne21, int ne22, int ne23, + int s1, int s2, int s3, + int s00, int s01, int s02, int s03, + int s10, int s11, int s12, int s13, + int s20, int s21, int s22, int s23, + const sycl::nd_item<3> & item_ct1) { + const int i0s = item_ct1.get_local_range(2) * item_ct1.get_group(2) + + item_ct1.get_local_id(2); + const int i1 = (item_ct1.get_local_range(1) * item_ct1.get_group(1) + + item_ct1.get_local_id(1)); + const int i2 = (item_ct1.get_local_range(0) * item_ct1.get_group(0) + + item_ct1.get_local_id(0)) / + ne3; + const int i3 = (item_ct1.get_local_range(0) * item_ct1.get_group(0) + + item_ct1.get_local_id(0)) % + ne3; + + if (i0s >= ne0 || i1 >= ne1 || i2 >= ne2 || i3 >= ne3) { + return; + } + + const int i11 = i1 % ne11; + const int i12 = i2 % ne12; + const int i13 = i3 % ne13; + const int i21 = i1 % ne21; + const int i22 = i2 % ne22; + const int i23 = i3 % ne23; + + const size_t i_src0 = i3 * s03 + i2 * s02 + i1 * s01; + const size_t i_src1 = i13 * s13 + i12 * s12 + i11 * s11; + const size_t i_src2 = i23 * s23 + i22 * s22 + i21 * s21; + const size_t i_dst = i3 * s3 + i2 * s2 + i1 * s1; + + const src0_t * src0_row = src0 + i_src0; + const src1_t * src1_row = src1 + i_src1; + const src2_t * src2_row = src2 + i_src2; + dst_t * dst_row = dst + i_dst; + + for (int i0 = i0s; i0 < ne0; + i0 += item_ct1.get_local_range(2) * item_ct1.get_group_range(2)) { + const int i10 = i0 % ne10; + const int i20 = i0 % ne20; + const float acc = bin_op((float) src0_row[i0 * s00], (float) src1_row[i10 * s10]); + dst_row[i0] = (dst_t) bin_op(acc, (float) src2_row[i20 * s20]); + } +} + +template +static void k_bin_bcast3_unravel(const src0_t * src0, const src1_t * src1, const src2_t * src2, dst_t * dst, + int ne0, int ne1, int ne2, int ne3, + int ne10, int ne11, int ne12, int ne13, + int ne20, int ne21, int ne22, int ne23, + int s1, int s2, int s3, + int s00, int s01, int s02, int s03, + int s10, int s11, int s12, int s13, + int s20, int s21, int s22, int s23, + const sycl::nd_item<3> & item_ct1) { + const int i = item_ct1.get_local_range(2) * item_ct1.get_group(2) + + item_ct1.get_local_id(2); + + const int i3 = i / (ne2 * ne1 * ne0); + const int i2 = (i / (ne1 * ne0)) % ne2; + const int i1 = (i / ne0) % ne1; + const int i0 = i % ne0; + + if (i0 >= ne0 || i1 >= ne1 || i2 >= ne2 || i3 >= ne3) { + return; + } + + const int i11 = i1 % ne11; + const int i12 = i2 % ne12; + const int i13 = i3 % ne13; + const int i21 = i1 % ne21; + const int i22 = i2 % ne22; + const int i23 = i3 % ne23; + + const size_t i_src0 = i3 * s03 + i2 * s02 + i1 * s01; + const size_t i_src1 = i13 * s13 + i12 * s12 + i11 * s11; + const size_t i_src2 = i23 * s23 + i22 * s22 + i21 * s21; + const size_t i_dst = i3 * s3 + i2 * s2 + i1 * s1; + + const int i10 = i0 % ne10; + const int i20 = i0 % ne20; + const float acc = bin_op((float) src0[i_src0 + i0 * s00], (float) src1[i_src1 + i10 * s10]); + dst[i_dst + i0] = (dst_t) bin_op(acc, (float) src2[i_src2 + i20 * s20]); +} + +template +static void launch_bin_bcast3(ggml_backend_sycl_context & ctx, const ggml_tensor * src0, const ggml_tensor * src1, + const ggml_tensor * src2, ggml_tensor * dst) { + dpct::queue_ptr stream = ctx.stream(); + SYCL_CHECK(ggml_sycl_set_device(ctx.device)); + + GGML_TENSOR_TERNARY_OP_LOCALS + + int nr1[4] = { (int) (ne10 / ne0), (int) (ne11 / ne1), (int) (ne12 / ne2), (int) (ne13 / ne3) }; + int nr2[4] = { (int) (ne20 / ne0), (int) (ne21 / ne1), (int) (ne22 / ne2), (int) (ne23 / ne3) }; + + int64_t cne[] = { ne0, ne1, ne2, ne3 }; + int64_t cne0[] = { ne00, ne01, ne02, ne03 }; + int64_t cne1[] = { ne10, ne11, ne12, ne13 }; + int64_t cne2[] = { ne20, ne21, ne22, ne23 }; + size_t cnb[] = { nb0, nb1, nb2, nb3 }; + size_t cnb0[] = { nb00, nb01, nb02, nb03 }; + size_t cnb1[] = { nb10, nb11, nb12, nb13 }; + size_t cnb2[] = { nb20, nb21, nb22, nb23 }; + + auto collapse = [](int64_t cne[]) { + cne[0] *= cne[1]; + cne[1] = cne[2]; + cne[2] = cne[3]; + cne[3] = 1; + }; + + auto collapse_nb = [](size_t cnb[], int64_t cne[]) { + cnb[1] *= cne[1]; + cnb[2] *= cne[2]; + cnb[3] *= cne[3]; + }; + + const bool can_collapse = ggml_is_contiguous(src0) && ggml_is_contiguous(src1) && ggml_is_contiguous(src2) && + !ggml_is_permuted(src0) && !ggml_is_permuted(src1) && !ggml_is_permuted(src2); + if (can_collapse) { + for (int i = 0; i < 4; i++) { + if (nr1[i] != 1 || nr2[i] != 1) { + break; + } + if (i > 0) { + collapse_nb(cnb, cne); + collapse_nb(cnb0, cne0); + collapse_nb(cnb1, cne1); + collapse_nb(cnb2, cne2); + collapse(cne); + collapse(cne0); + collapse(cne1); + collapse(cne2); + } + } + } + + { + int64_t ne0 = cne[0]; + int64_t ne1 = cne[1]; + int64_t ne2 = cne[2]; + int64_t ne3 = cne[3]; + + int64_t ne10 = cne1[0]; + int64_t ne11 = cne1[1]; + int64_t ne12 = cne1[2]; + int64_t ne13 = cne1[3]; + + int64_t ne20 = cne2[0]; + int64_t ne21 = cne2[1]; + int64_t ne22 = cne2[2]; + int64_t ne23 = cne2[3]; + + size_t s1 = cnb[1] / sizeof(dst_t); + size_t s2 = cnb[2] / sizeof(dst_t); + size_t s3 = cnb[3] / sizeof(dst_t); + + size_t s00 = cnb0[0] / sizeof(src0_t); + size_t s01 = cnb0[1] / sizeof(src0_t); + size_t s02 = cnb0[2] / sizeof(src0_t); + size_t s03 = cnb0[3] / sizeof(src0_t); + + size_t s10 = cnb1[0] / sizeof(src1_t); + size_t s11 = cnb1[1] / sizeof(src1_t); + size_t s12 = cnb1[2] / sizeof(src1_t); + size_t s13 = cnb1[3] / sizeof(src1_t); + + size_t s20 = cnb2[0] / sizeof(src2_t); + size_t s21 = cnb2[1] / sizeof(src2_t); + size_t s22 = cnb2[2] / sizeof(src2_t); + size_t s23 = cnb2[3] / sizeof(src2_t); + + GGML_ASSERT(cnb[0] % sizeof(dst_t) == 0 && cnb[1] % sizeof(dst_t) == 0 && cnb[2] % sizeof(dst_t) == 0 && + cnb[3] % sizeof(dst_t) == 0); + GGML_ASSERT(cnb0[0] % sizeof(src0_t) == 0 && cnb0[1] % sizeof(src0_t) == 0 && cnb0[2] % sizeof(src0_t) == 0 && + cnb0[3] % sizeof(src0_t) == 0); + GGML_ASSERT(cnb1[0] % sizeof(src1_t) == 0 && cnb1[1] % sizeof(src1_t) == 0 && cnb1[2] % sizeof(src1_t) == 0 && + cnb1[3] % sizeof(src1_t) == 0); + GGML_ASSERT(cnb2[0] % sizeof(src2_t) == 0 && cnb2[1] % sizeof(src2_t) == 0 && cnb2[2] % sizeof(src2_t) == 0 && + cnb2[3] % sizeof(src2_t) == 0); + + const src0_t * src0_dd = (const src0_t *) src0->data; + const src1_t * src1_dd = (const src1_t *) src1->data; + const src2_t * src2_dd = (const src2_t *) src2->data; + dst_t * dst_dd = (dst_t *) dst->data; + + const int block_size = 128; + int64_t hne0 = std::max(ne0 / 2LL, 1LL); + + sycl::range<3> block_dims(1, 1, 1); + block_dims[2] = std::min(hne0, block_size); + block_dims[1] = std::min(ne1, block_size / (unsigned int) block_dims[2]); + block_dims[0] = std::min(std::min(ne2 * ne3, + block_size / (unsigned int) block_dims[2] / + (unsigned int) block_dims[1]), + 64U); + + sycl::range<3> block_nums((ne2 * ne3 + block_dims[0] - 1) / block_dims[0], + (ne1 + block_dims[1] - 1) / block_dims[1], + (hne0 + block_dims[2] - 1) / block_dims[2]); + + dpct::has_capability_or_fail(stream->get_device(), { sycl::aspect::fp16 }); + + if (block_nums[0] > 65535) { + int block_num = (ne0 * ne1 * ne2 * ne3 + block_size - 1) / block_size; + stream->parallel_for( + sycl::nd_range<3>(sycl::range<3>(1, 1, block_num) * sycl::range<3>(1, 1, block_size), + sycl::range<3>(1, 1, block_size)), + [=](sycl::nd_item<3> item_ct1) { + k_bin_bcast3_unravel(src0_dd, src1_dd, src2_dd, dst_dd, ne0, ne1, ne2, ne3, ne10, ne11, + ne12, ne13, ne20, ne21, ne22, ne23, s1, s2, s3, s00, s01, s02, s03, + s10, s11, s12, s13, s20, s21, s22, s23, item_ct1); + }); + } else { + stream->parallel_for(sycl::nd_range<3>(block_nums * block_dims, block_dims), + [=](sycl::nd_item<3> item_ct1) { + k_bin_bcast3(src0_dd, src1_dd, src2_dd, dst_dd, ne0, ne1, ne2, ne3, ne10, + ne11, ne12, ne13, ne20, ne21, ne22, ne23, s1, s2, s3, s00, + s01, s02, s03, s10, s11, s12, s13, s20, s21, s22, s23, + item_ct1); + }); + } + } +} + +void ggml_sycl_op_add_add_fused(ggml_backend_sycl_context & ctx, ggml_tensor * add0, ggml_tensor * add1) { + const ggml_tensor * src0 = add0->src[0]; + const ggml_tensor * src1 = add0->src[1]; + const ggml_tensor * src2 = add1->src[1]; + ggml_tensor * dst = add1; + + GGML_ASSERT(add1->src[0] == add0); + GGML_ASSERT(ggml_sycl_add_kernel_supports(src0->type, src1->type, add0->type)); + GGML_ASSERT(ggml_sycl_add_kernel_supports(add0->type, src2->type, dst->type)); + + if (src0->type == GGML_TYPE_F32 && src1->type == GGML_TYPE_F32 && src2->type == GGML_TYPE_F32 && + dst->type == GGML_TYPE_F32) { + launch_bin_bcast3(ctx, src0, src1, src2, dst); + } else if (src0->type == GGML_TYPE_F16 && src1->type == GGML_TYPE_F16 && src2->type == GGML_TYPE_F16 && + dst->type == GGML_TYPE_F16) { + launch_bin_bcast3(ctx, src0, src1, src2, dst); + } else if (src0->type == GGML_TYPE_F16 && src1->type == GGML_TYPE_F32 && src2->type == GGML_TYPE_F32 && + dst->type == GGML_TYPE_F16) { + launch_bin_bcast3(ctx, src0, src1, src2, dst); + } else if (src0->type == GGML_TYPE_F16 && src1->type == GGML_TYPE_F16 && src2->type == GGML_TYPE_F32 && + dst->type == GGML_TYPE_F16) { + launch_bin_bcast3(ctx, src0, src1, src2, dst); + } else if (src0->type == GGML_TYPE_F16 && src1->type == GGML_TYPE_F32 && src2->type == GGML_TYPE_F16 && + dst->type == GGML_TYPE_F16) { + launch_bin_bcast3(ctx, src0, src1, src2, dst); + } else if (src0->type == GGML_TYPE_I32 && src1->type == GGML_TYPE_I32 && src2->type == GGML_TYPE_I32 && + dst->type == GGML_TYPE_I32) { + launch_bin_bcast3(ctx, src0, src1, src2, dst); + } else if (src0->type == GGML_TYPE_I16 && src1->type == GGML_TYPE_I16 && src2->type == GGML_TYPE_I16 && + dst->type == GGML_TYPE_I16) { + launch_bin_bcast3(ctx, src0, src1, src2, dst); +#ifdef GGML_SYCL_HAS_BF16 + } else if (src0->type == GGML_TYPE_BF16 && src1->type == GGML_TYPE_BF16 && src2->type == GGML_TYPE_BF16 && + dst->type == GGML_TYPE_BF16) { + launch_bin_bcast3(ctx, src0, src1, src2, dst); + } else if (src0->type == GGML_TYPE_BF16 && src1->type == GGML_TYPE_F32 && src2->type == GGML_TYPE_F32 && + dst->type == GGML_TYPE_BF16) { + launch_bin_bcast3( + ctx, src0, src1, src2, dst); + } else if (src0->type == GGML_TYPE_BF16 && src1->type == GGML_TYPE_BF16 && src2->type == GGML_TYPE_F32 && + dst->type == GGML_TYPE_BF16) { + launch_bin_bcast3(ctx, src0, src1, src2, dst); + } else if (src0->type == GGML_TYPE_BF16 && src1->type == GGML_TYPE_F32 && src2->type == GGML_TYPE_BF16 && + dst->type == GGML_TYPE_BF16) { + launch_bin_bcast3(ctx, src0, src1, src2, dst); +#endif + } else { + fprintf(stderr, "%s: unsupported types: dst: %s, src0: %s, src1: %s, src2: %s\n", __func__, + ggml_type_name(dst->type), ggml_type_name(src0->type), ggml_type_name(src1->type), + ggml_type_name(src2->type)); + GGML_ABORT("fatal error"); + } +} + diff --git a/ggml/src/ggml-sycl/binbcast.hpp b/ggml/src/ggml-sycl/binbcast.hpp index 9cce0f053a5..0e5a5ca1c1a 100644 --- a/ggml/src/ggml-sycl/binbcast.hpp +++ b/ggml/src/ggml-sycl/binbcast.hpp @@ -34,6 +34,36 @@ void ggml_sycl_div(ggml_backend_sycl_context & ctx, ggml_tensor * dst); void ggml_sycl_repeat(ggml_backend_sycl_context & ctx, ggml_tensor * dst); +void ggml_sycl_op_add_add_fused(ggml_backend_sycl_context & ctx, ggml_tensor * add0, ggml_tensor * add1); + +// Type combinations the standalone SYCL add() kernel can run. Fused ADD+ADD +// uses the same set; anything else falls back to two add() launches. +inline bool ggml_sycl_add_kernel_supports(enum ggml_type src0, enum ggml_type src1, enum ggml_type dst) { + if (src0 == GGML_TYPE_F32 && src1 == GGML_TYPE_F32 && dst == GGML_TYPE_F32) { + return true; + } + if (src0 == GGML_TYPE_F16 && src1 == GGML_TYPE_F16 && dst == GGML_TYPE_F16) { + return true; + } + if (src0 == GGML_TYPE_F16 && src1 == GGML_TYPE_F32 && dst == GGML_TYPE_F16) { + return true; + } + if (src0 == GGML_TYPE_I32 && src1 == GGML_TYPE_I32 && dst == GGML_TYPE_I32) { + return true; + } + if (src0 == GGML_TYPE_I16 && src1 == GGML_TYPE_I16 && dst == GGML_TYPE_I16) { + return true; + } +#ifdef GGML_SYCL_HAS_BF16 + if (src0 == GGML_TYPE_BF16 && src1 == GGML_TYPE_BF16 && dst == GGML_TYPE_BF16) { + return true; + } + if (src0 == GGML_TYPE_BF16 && src1 == GGML_TYPE_F32 && dst == GGML_TYPE_BF16) { + return true; + } +#endif + return false; +} #endif //GGML_SYCL_BINBCAST_HPP diff --git a/ggml/src/ggml-sycl/fusion.cpp b/ggml/src/ggml-sycl/fusion.cpp index 709bc8ca2a2..b5e79bea543 100644 --- a/ggml/src/ggml-sycl/fusion.cpp +++ b/ggml/src/ggml-sycl/fusion.cpp @@ -1,4 +1,5 @@ #include "fusion.hpp" +#include "binbcast.hpp" #include @@ -94,9 +95,14 @@ bool ggml_sycl_can_fuse(const ggml_cgraph * cgraph, int node_idx, std::initializ return false; } - if (ops.size() == 2 && ops.begin()[0] == GGML_OP_RMS_NORM && ops.begin()[1] == GGML_OP_MUL) { + if ((ops.size() == 2 || ops.size() == 3) && ops.begin()[0] == GGML_OP_RMS_NORM && ops.begin()[1] == GGML_OP_MUL) { + if (ops.size() == 3 && ops.begin()[2] != GGML_OP_ADD) { + return false; + } + const ggml_tensor * rms_norm = cgraph->nodes[node_idx]; const ggml_tensor * mul = cgraph->nodes[node_idx + 1]; + const ggml_tensor * add = ops.size() == 3 ? cgraph->nodes[node_idx + 2] : nullptr; GGML_ASSERT(rms_norm->src[0]->type == GGML_TYPE_F32); GGML_ASSERT(rms_norm->type == GGML_TYPE_F32); @@ -122,6 +128,43 @@ bool ggml_sycl_can_fuse(const ggml_cgraph * cgraph, int node_idx, std::initializ return false; } + if (add != nullptr) { + if (add->src[0]->type != GGML_TYPE_F32 || + add->src[1]->type != GGML_TYPE_F32 || + add->type != GGML_TYPE_F32) { + return false; + } + + // the fused kernel indexes the residual as add[col] and does not broadcast it + const ggml_tensor * add_w = (add->src[0] == mul) ? add->src[1] : add->src[0]; + if (!ggml_are_same_shape(add_w, add)) { + return false; + } + + if (!ggml_is_contiguous(add->src[0]) || !ggml_is_contiguous_rows(add->src[1])) { + return false; + } + } + + return true; + } + + if (ops.size() == 2 && ops.begin()[0] == GGML_OP_ADD && ops.begin()[1] == GGML_OP_ADD) { + const ggml_tensor * add0 = cgraph->nodes[node_idx]; + const ggml_tensor * add1 = cgraph->nodes[node_idx + 1]; + // ggml_can_fuse already guarantees add1 consumes add0 and that add0 has a single use. + // Keep the CUDA association: the running sum is src0 of the next ADD so the fused + // float fold matches two sequential add() launches. + if (add1->src[0] != add0) { + return false; + } + + const ggml_tensor * c = add1->src[1]; + if (!ggml_sycl_add_kernel_supports(add0->src[0]->type, add0->src[1]->type, add0->type) || + !ggml_sycl_add_kernel_supports(add0->type, c->type, add1->type)) { + return false; + } + return true; } diff --git a/ggml/src/ggml-sycl/ggml-sycl.cpp b/ggml/src/ggml-sycl/ggml-sycl.cpp index feced748cec..27804e07301 100644 --- a/ggml/src/ggml-sycl/ggml-sycl.cpp +++ b/ggml/src/ggml-sycl/ggml-sycl.cpp @@ -5865,12 +5865,24 @@ static void ggml_backend_sycl_graph_compute_impl(ggml_backend_sycl_context * syc continue; } } + if (node->op == GGML_OP_RMS_NORM && + ggml_sycl_can_fuse(cgraph, i, { GGML_OP_RMS_NORM, GGML_OP_MUL, GGML_OP_ADD }, {})) { + ggml_sycl_op_rms_norm_fused_add(*sycl_ctx, node, cgraph->nodes[i + 1], cgraph->nodes[i + 2]); + i += 2; + continue; + } if (node->op == GGML_OP_RMS_NORM && ggml_sycl_can_fuse(cgraph, i, { GGML_OP_RMS_NORM, GGML_OP_MUL }, {})) { ggml_sycl_op_rms_norm_fused(*sycl_ctx, node, cgraph->nodes[i + 1]); i++; continue; } + if (node->op == GGML_OP_ADD && + ggml_sycl_can_fuse(cgraph, i, { GGML_OP_ADD, GGML_OP_ADD }, {})) { + ggml_sycl_op_add_add_fused(*sycl_ctx, node, cgraph->nodes[i + 1]); + i++; + continue; + } if (node->op == GGML_OP_UNARY && ggml_sycl_can_fuse(cgraph, i, { GGML_OP_UNARY, GGML_OP_MUL }, { ggml_get_unary_op(node) })) { ggml_sycl_op_unary_mul_fused(*sycl_ctx, node, cgraph->nodes[i + 1]); diff --git a/ggml/src/ggml-sycl/norm.cpp b/ggml/src/ggml-sycl/norm.cpp index f98a7a9542c..2d303372934 100644 --- a/ggml/src/ggml-sycl/norm.cpp +++ b/ggml/src/ggml-sycl/norm.cpp @@ -144,13 +144,17 @@ static void group_norm_f32(const float* x, float* dst, const int group_size, con } } -template +template static void rms_norm_f32(const float* x, float* dst, const int ncols, const int64_t src_stride_col, const int64_t src_stride_row, const int64_t src_stride_channel, const int64_t src_stride_sample, const int64_t dst_stride_col, const int64_t dst_stride_row, const int64_t dst_stride_channel, const int64_t dst_stride_sample, const float eps, const sycl::nd_item<3>& item_ct1, float* s_sum, int block_size, const float* mul = nullptr, const int64_t mul_stride_row = 0, const int64_t mul_stride_channel = 0, - const int64_t mul_stride_sample = 0, const int mul_nrows = 0, const int mul_nchannels = 0, const int mul_nsamples = 0) { + const int64_t mul_stride_sample = 0, const int mul_nrows = 0, const int mul_nchannels = 0, const int mul_nsamples = 0, + const float* add = nullptr, const int64_t add_stride_row = 0, const int64_t add_stride_channel = 0, + const int64_t add_stride_sample = 0, const int add_nrows = 0, const int add_nchannels = 0, const int add_nsamples = 0) { + + static_assert(!do_add || do_multiply, "fusing add is not supported without multiplying"); const int sample = item_ct1.get_group(0); const int channel = item_ct1.get_group(1); @@ -174,6 +178,13 @@ static void rms_norm_f32(const float* x, float* dst, const int ncols, mul += mul_sample * mul_stride_sample + mul_channel * mul_stride_channel + mul_row * mul_stride_row; } + if constexpr (do_add) { + const int add_row = row % add_nrows; + const int add_channel = channel % add_nchannels; + const int add_sample = sample % add_nsamples; + add += add_sample * add_stride_sample + add_channel * add_stride_channel + add_row * add_stride_row; + } + float tmp = 0.0f; // partial sum for thread in warp for (int col = tid; col < ncols; col += block_size) { @@ -205,7 +216,9 @@ static void rms_norm_f32(const float* x, float* dst, const int ncols, const float scale = sycl::rsqrt(mean + eps); for (int col = tid; col < ncols; col += block_size) { - if constexpr (do_multiply) { + if constexpr (do_multiply && do_add) { + dst[col * dst_stride_col] = scale * x[col * src_stride_col] * mul[col] + add[col]; + } else if constexpr (do_multiply) { dst[col * dst_stride_col] = scale * x[col * src_stride_col] * mul[col]; } else { dst[col * dst_stride_col] = scale * x[col * src_stride_col]; @@ -424,6 +437,53 @@ static void rms_norm_mul_f32_sycl(const float* x, const float* mul, float* dst, } } +static void rms_norm_mul_add_f32_sycl(const float* x, const float* mul, const float* add, float* dst, + const int ncols, const int nrows, const int nchannels, const int nsamples, + const int64_t src_stride_col, const int64_t src_stride_row, const int64_t src_stride_channel, const int64_t src_stride_sample, + const int64_t dst_stride_col, const int64_t dst_stride_row, const int64_t dst_stride_channel, const int64_t dst_stride_sample, + const int64_t mul_stride_row, const int64_t mul_stride_channel, const int64_t mul_stride_sample, + const int mul_nrows, const int mul_nchannels, const int mul_nsamples, + const int64_t add_stride_row, const int64_t add_stride_channel, const int64_t add_stride_sample, + const int add_nrows, const int add_nchannels, const int add_nsamples, + const float eps, queue_ptr stream, int device) { + const sycl::range<3> global_dims(nsamples, nchannels, nrows); + if (ncols < 1024) { + const sycl::range<3> block_dims(1, 1, WARP_SIZE); + stream->submit([&](sycl::handler& cgh) { + cgh.parallel_for( + sycl::nd_range<3>(global_dims * block_dims, block_dims), + [=](sycl::nd_item<3> item_ct1) + [[sycl::reqd_sub_group_size(WARP_SIZE)]] { + rms_norm_f32(x, dst, ncols, + src_stride_col, src_stride_row, src_stride_channel, src_stride_sample, + dst_stride_col, dst_stride_row, dst_stride_channel, dst_stride_sample, + eps, item_ct1, nullptr, WARP_SIZE, + mul, mul_stride_row, mul_stride_channel, mul_stride_sample, mul_nrows, mul_nchannels, mul_nsamples, + add, add_stride_row, add_stride_channel, add_stride_sample, add_nrows, add_nchannels, add_nsamples); + }); + }); + } + else { + const int work_group_size = ggml_sycl_info().max_work_group_sizes[device]; + assert(work_group_size % (WARP_SIZE * WARP_SIZE) == 0); + const sycl::range<3> block_dims(1, 1, work_group_size); + stream->submit([&](sycl::handler& cgh) { + sycl::local_accessor s_sum_acc_ct1(sycl::range<1>(work_group_size / WARP_SIZE), cgh); + cgh.parallel_for( + sycl::nd_range<3>(global_dims * block_dims, block_dims), + [=](sycl::nd_item<3> item_ct1) + [[sycl::reqd_sub_group_size(WARP_SIZE)]] { + rms_norm_f32(x, dst, ncols, + src_stride_col, src_stride_row, src_stride_channel, src_stride_sample, + dst_stride_col, dst_stride_row, dst_stride_channel, dst_stride_sample, + eps, item_ct1, get_pointer(s_sum_acc_ct1), work_group_size, + mul, mul_stride_row, mul_stride_channel, mul_stride_sample, mul_nrows, mul_nchannels, mul_nsamples, + add, add_stride_row, add_stride_channel, add_stride_sample, add_nrows, add_nchannels, add_nsamples); + }); + }); + } +} + template static void l2_norm_f32_sycl(const float * x, float * dst, @@ -626,6 +686,91 @@ void ggml_sycl_op_rms_norm_fused(ggml_backend_sycl_context & ctx, ggml_tensor * mul_s01, mul_s02, mul_s03, mul_nrows, mul_nchannels, mul_nsamples, eps, main_stream, ctx.device); } +void ggml_sycl_op_rms_norm_fused_add(ggml_backend_sycl_context & ctx, ggml_tensor * dst, + ggml_tensor * mul_tensor, ggml_tensor * add_tensor) { + const ggml_tensor * rms_norm_src = dst->src[0]; + float eps = 0.0f; + memcpy(&eps, dst->op_params, sizeof(float)); + + const float * src0_dd = static_cast(rms_norm_src->data); + const float * mul_dd = nullptr; + const ggml_tensor * mul_src = nullptr; + if (mul_tensor->src[0] == dst) { + mul_dd = static_cast(mul_tensor->src[1]->data); + mul_src = mul_tensor->src[1]; + } else if (mul_tensor->src[1] == dst) { + mul_dd = static_cast(mul_tensor->src[0]->data); + mul_src = mul_tensor->src[0]; + } else { + GGML_ASSERT(false); + } + + const float * add_dd = nullptr; + const ggml_tensor * add_src = nullptr; + if (add_tensor->src[0] == mul_tensor) { + add_dd = static_cast(add_tensor->src[1]->data); + add_src = add_tensor->src[1]; + } else if (add_tensor->src[1] == mul_tensor) { + add_dd = static_cast(add_tensor->src[0]->data); + add_src = add_tensor->src[0]; + } else { + GGML_ASSERT(false); + } + + float * dst_dd = static_cast(add_tensor->data); + + dpct::queue_ptr main_stream = ctx.stream(); + SYCL_CHECK(ggml_sycl_set_device(ctx.device)); + + GGML_ASSERT(rms_norm_src->type == GGML_TYPE_F32); + GGML_ASSERT(dst->type == GGML_TYPE_F32); + GGML_ASSERT(mul_tensor->type == GGML_TYPE_F32); + GGML_ASSERT(add_tensor->type == GGML_TYPE_F32); + GGML_ASSERT(eps >= 0.0f); + + const int64_t ne00 = rms_norm_src->ne[0]; + const int64_t ne01 = rms_norm_src->ne[1]; + const int64_t ne02 = rms_norm_src->ne[2]; + const int64_t ne03 = rms_norm_src->ne[3]; + + const size_t ts0 = ggml_type_size(rms_norm_src->type); + GGML_ASSERT(rms_norm_src->nb[0] == ts0); + const int64_t s00 = rms_norm_src->nb[0] / ts0; + const int64_t s01 = rms_norm_src->nb[1] / ts0; + const int64_t s02 = rms_norm_src->nb[2] / ts0; + const int64_t s03 = rms_norm_src->nb[3] / ts0; + + const size_t tdst = ggml_type_size(add_tensor->type); + GGML_ASSERT(add_tensor->nb[0] == tdst); + const int64_t d00 = add_tensor->nb[0] / tdst; + const int64_t d01 = add_tensor->nb[1] / tdst; + const int64_t d02 = add_tensor->nb[2] / tdst; + const int64_t d03 = add_tensor->nb[3] / tdst; + + const size_t ts_mul = ggml_type_size(mul_src->type); + GGML_ASSERT(mul_src->nb[0] == ts_mul); + const int64_t mul_s01 = mul_src->nb[1] / ts_mul; + const int64_t mul_s02 = mul_src->nb[2] / ts_mul; + const int64_t mul_s03 = mul_src->nb[3] / ts_mul; + const int mul_nrows = mul_src->ne[1]; + const int mul_nchannels = mul_src->ne[2]; + const int mul_nsamples = mul_src->ne[3]; + + const size_t ts_add = ggml_type_size(add_src->type); + GGML_ASSERT(add_src->nb[0] == ts_add); + const int64_t add_s01 = add_src->nb[1] / ts_add; + const int64_t add_s02 = add_src->nb[2] / ts_add; + const int64_t add_s03 = add_src->nb[3] / ts_add; + const int add_nrows = add_src->ne[1]; + const int add_nchannels = add_src->ne[2]; + const int add_nsamples = add_src->ne[3]; + + rms_norm_mul_add_f32_sycl(src0_dd, mul_dd, add_dd, dst_dd, ne00, ne01, ne02, ne03, + s00, s01, s02, s03, d00, d01, d02, d03, + mul_s01, mul_s02, mul_s03, mul_nrows, mul_nchannels, mul_nsamples, + add_s01, add_s02, add_s03, add_nrows, add_nchannels, add_nsamples, eps, main_stream, ctx.device); +} + void ggml_sycl_op_rms_norm_back(ggml_backend_sycl_context & ctx, ggml_tensor * dst) { scope_op_debug_print scope_dbg_print(__func__, dst, /*num_src=*/2); diff --git a/ggml/src/ggml-sycl/norm.hpp b/ggml/src/ggml-sycl/norm.hpp index 51217c42195..ef7b2d386bd 100644 --- a/ggml/src/ggml-sycl/norm.hpp +++ b/ggml/src/ggml-sycl/norm.hpp @@ -21,6 +21,8 @@ void ggml_sycl_op_rms_norm(ggml_backend_sycl_context& ctx, ggml_tensor* dst); void ggml_sycl_op_rms_norm_fused(ggml_backend_sycl_context& ctx, ggml_tensor* dst, ggml_tensor* mul); +void ggml_sycl_op_rms_norm_fused_add(ggml_backend_sycl_context& ctx, ggml_tensor* dst, ggml_tensor* mul_tensor, ggml_tensor* add_tensor); + void ggml_sycl_op_rms_norm_back(ggml_backend_sycl_context& ctx, ggml_tensor* dst); void ggml_sycl_op_group_norm(ggml_backend_sycl_context& ctx, ggml_tensor* dst); From e1bbe405209c6e6836f9416a5df4363fe44cddb0 Mon Sep 17 00:00:00 2001 From: Aaron Teo Date: Fri, 4 Sep 2026 13:55:50 +0800 Subject: [PATCH 097/104] ggml-cpu(s390x) : fix q5_1 uninitialized v_acc (llama/28332) Signed-off-by: Aaron Teo --- ggml/src/ggml-cpu/arch/s390/quants.c | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/ggml/src/ggml-cpu/arch/s390/quants.c b/ggml/src/ggml-cpu/arch/s390/quants.c index 500857579a7..d3436c24b5f 100644 --- a/ggml/src/ggml-cpu/arch/s390/quants.c +++ b/ggml/src/ggml-cpu/arch/s390/quants.c @@ -636,7 +636,7 @@ void ggml_vec_dot_q5_1_q8_1(int n, float * GGML_RESTRICT s, size_t bs, const voi const float32x4_t v_xyf = vec_float(v_xy); const float32x4_t v_d = vec_splats(GGML_CPU_FP16_TO_FP32(x0->d) * GGML_CPU_FP16_TO_FP32(y0->d)); - const float32x4_t v_acc = vec_madd(v_xyf, v_d, v_acc); + const float32x4_t v_acc = vec_madd(v_xyf, v_d, vec_splats(0.0f)); sumf += vec_hsum_f32x4(v_acc) + summs; } From f32e6fa0c8f1a78fafe4a911821e660b76a2c02d Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Adrien=20Gallou=C3=ABt?= Date: Fri, 4 Sep 2026 09:22:01 +0200 Subject: [PATCH 098/104] ggml : remove GGML_CUDA_PEER_MAX_BATCH_SIZE (llama/28177) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Signed-off-by: Adrien Gallouët --- ggml/CMakeLists.txt | 2 -- ggml/src/ggml-cuda/CMakeLists.txt | 2 -- ggml/src/ggml-musa/CMakeLists.txt | 1 - 3 files changed, 5 deletions(-) diff --git a/ggml/CMakeLists.txt b/ggml/CMakeLists.txt index 0ac2b15c48c..634e18c5422 100644 --- a/ggml/CMakeLists.txt +++ b/ggml/CMakeLists.txt @@ -200,8 +200,6 @@ option(GGML_CUDA "ggml: use CUDA" option(GGML_MUSA "ggml: use MUSA" OFF) option(GGML_CUDA_FORCE_MMQ "ggml: use mmq kernels instead of cuBLAS" OFF) option(GGML_CUDA_FORCE_CUBLAS "ggml: always use cuBLAS instead of mmq kernels" OFF) -set (GGML_CUDA_PEER_MAX_BATCH_SIZE "128" CACHE STRING - "ggml: max. batch size for using peer access") option(GGML_CUDA_NO_PEER_COPY "ggml: do not use peer to peer copies" OFF) option(GGML_CUDA_NO_VMM "ggml: do not try to use CUDA VMM" OFF) option(GGML_CUDA_FA "ggml: compile ggml FlashAttention CUDA kernels" ON) diff --git a/ggml/src/ggml-cuda/CMakeLists.txt b/ggml/src/ggml-cuda/CMakeLists.txt index d3953eee962..10828ad8174 100644 --- a/ggml/src/ggml-cuda/CMakeLists.txt +++ b/ggml/src/ggml-cuda/CMakeLists.txt @@ -129,8 +129,6 @@ if (CUDAToolkit_FOUND) ${GGML_SOURCES_CUDA} ) - add_compile_definitions(GGML_CUDA_PEER_MAX_BATCH_SIZE=${GGML_CUDA_PEER_MAX_BATCH_SIZE}) - if (GGML_CUDA_GRAPHS) add_compile_definitions(GGML_CUDA_USE_GRAPHS) endif() diff --git a/ggml/src/ggml-musa/CMakeLists.txt b/ggml/src/ggml-musa/CMakeLists.txt index cc53c812ce5..faf9790338b 100644 --- a/ggml/src/ggml-musa/CMakeLists.txt +++ b/ggml/src/ggml-musa/CMakeLists.txt @@ -75,7 +75,6 @@ if (MUSAToolkit_FOUND) endif() add_compile_definitions(GGML_USE_MUSA) - add_compile_definitions(GGML_CUDA_PEER_MAX_BATCH_SIZE=${GGML_CUDA_PEER_MAX_BATCH_SIZE}) if (GGML_MUSA_GRAPHS) add_compile_definitions(GGML_MUSA_GRAPHS) From 1b37beadd76496ed4081998150c2b8a81783b7a0 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Adrien=20Gallou=C3=ABt?= Date: Fri, 4 Sep 2026 09:24:06 +0200 Subject: [PATCH 099/104] ggml : don't crash when backend search path can't be read (llama/28271) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Use std::error_code overloads of fs::current_path() and fs::directory_iterator in ggml_backend_load_best() so an inaccessible search path (WebDAV mount, removed CWD) is skipped instead of terminating the process with an uncaught filesystem_error. Signed-off-by: Adrien Gallouët --- ggml/src/ggml-backend-reg.cpp | 18 +++++++++++++++--- 1 file changed, 15 insertions(+), 3 deletions(-) diff --git a/ggml/src/ggml-backend-reg.cpp b/ggml/src/ggml-backend-reg.cpp index e5959467071..1c18b82cd50 100644 --- a/ggml/src/ggml-backend-reg.cpp +++ b/ggml/src/ggml-backend-reg.cpp @@ -490,7 +490,13 @@ static ggml_backend_reg_t ggml_backend_load_best(const char * name, bool silent, #endif // default search paths: executable directory, current directory search_paths.push_back(get_executable_path()); - search_paths.push_back(fs::current_path()); + std::error_code cwd_ec; + const fs::path cwd = fs::current_path(cwd_ec); + if (cwd_ec) { + GGML_LOG_DEBUG("%s: current_path() failure, error-message: %s\n", __func__, cwd_ec.message().c_str()); + } else { + search_paths.push_back(cwd); + } } else { search_paths.push_back(fs::u8path(user_search_path)); } @@ -508,8 +514,14 @@ static ggml_backend_reg_t ggml_backend_load_best(const char * name, bool silent, } continue; } - fs::directory_iterator dir_it(search_path, fs::directory_options::skip_permission_denied); - for (const auto & entry : dir_it) { + std::error_code dir_ec; + fs::directory_iterator dir_it(search_path, fs::directory_options::skip_permission_denied, dir_ec); + if (dir_ec) { + GGML_LOG_DEBUG("%s: failed to enumerate %s: %s\n", __func__, path_str(search_path).c_str(), dir_ec.message().c_str()); + continue; + } + for (const fs::directory_iterator end; dir_it != end; dir_it.increment(dir_ec)) { + const auto & entry = *dir_it; if (entry.is_regular_file(ec)) { auto filename = entry.path().filename(); auto ext = entry.path().extension(); From e2389eb99cf5796659495c0066bb491fb52a354f Mon Sep 17 00:00:00 2001 From: Georgi Gerganov Date: Fri, 4 Sep 2026 13:08:43 +0300 Subject: [PATCH 100/104] ggml : rename and make private ggml_op_alloc_size_may_expand() (ggml/0) cont https://github.com/ggml-org/llama.cpp/pull/27960 --- ggml/include/ggml-backend.h | 4 ---- ggml/src/ggml-backend-impl.h | 5 +++++ ggml/src/ggml-backend.cpp | 7 ++----- ggml/src/ggml-rpc/ggml-rpc.cpp | 2 +- 4 files changed, 8 insertions(+), 10 deletions(-) diff --git a/ggml/include/ggml-backend.h b/ggml/include/ggml-backend.h index 27375bd0a51..cc3f8cd36e3 100644 --- a/ggml/include/ggml-backend.h +++ b/ggml/include/ggml-backend.h @@ -424,10 +424,6 @@ extern "C" { // Compare the output of two backends GGML_API bool ggml_backend_compare_graph_backend(ggml_backend_t backend1, ggml_backend_t backend2, struct ggml_cgraph * graph, ggml_backend_eval_callback callback, void * user_data, struct ggml_tensor const * const * test_nodes, size_t num_test_nodes); - // returns true for ops that may require additional memory for fleeting data on some backends, - // i.e. the backend's get_alloc_size may return more than ggml_nbytes for the output tensor - GGML_API bool ggml_backend_op_alloc_size_may_expand(enum ggml_op op); - // Tensor initialization GGML_API enum ggml_status ggml_backend_tensor_alloc(ggml_backend_buffer_t buffer, struct ggml_tensor * tensor, void * addr); GGML_API enum ggml_status ggml_backend_view_init(struct ggml_tensor * tensor); diff --git a/ggml/src/ggml-backend-impl.h b/ggml/src/ggml-backend-impl.h index 56f0090cce6..ef05905cf9a 100644 --- a/ggml/src/ggml-backend-impl.h +++ b/ggml/src/ggml-backend-impl.h @@ -34,6 +34,11 @@ extern "C" { void * context; }; + // [TAG_ALLOC_SIZE_EXPAND] + // returns true for ops that may require additional memory for fleeting data on some backends, + // i.e. the backend buffer type's get_alloc_size may return more than ggml_nbytes for the output tensor + GGML_API bool ggml_op_alloc_size_may_expand(enum ggml_op op); + // // Backend buffer // diff --git a/ggml/src/ggml-backend.cpp b/ggml/src/ggml-backend.cpp index ffe20b9d05b..6862128e637 100644 --- a/ggml/src/ggml-backend.cpp +++ b/ggml/src/ggml-backend.cpp @@ -71,7 +71,7 @@ size_t ggml_backend_buft_get_alloc_size(ggml_backend_buffer_type_t buft, const s GGML_ASSERT(size <= ggml_nbytes(tensor) || ggml_op_is_empty(tensor->op) || ggml_is_quantized(tensor->type) || // [TAG_ALLOC_SIZE_EXPAND] - ggml_backend_op_alloc_size_may_expand(tensor->op)); + ggml_op_alloc_size_may_expand(tensor->op)); return size; } @@ -2109,10 +2109,7 @@ ggml_backend_t ggml_backend_sched_get_tensor_backend(ggml_backend_sched_t sched, // utils -// [TAG_ALLOC_SIZE_EXPAND] -// returns true for ops that may require additional memory for fleeting data on some backends, -// i.e. the backend's get_alloc_size may return more than ggml_nbytes for the output tensor -bool ggml_backend_op_alloc_size_may_expand(enum ggml_op op) { +bool ggml_op_alloc_size_may_expand(enum ggml_op op) { switch (op) { case GGML_OP_FLASH_ATTN_EXT: case GGML_OP_MUL_MAT: diff --git a/ggml/src/ggml-rpc/ggml-rpc.cpp b/ggml/src/ggml-rpc/ggml-rpc.cpp index a97db24e624..cc7d7206933 100644 --- a/ggml/src/ggml-rpc/ggml-rpc.cpp +++ b/ggml/src/ggml-rpc/ggml-rpc.cpp @@ -835,7 +835,7 @@ static size_t ggml_backend_rpc_buffer_type_get_alloc_size(ggml_backend_buffer_ty // [TAG_ALLOC_SIZE_EXPAND] // ops that may require additional memory for fleeting data on certain backends // ref: https://github.com/ggml-org/llama.cpp/pull/15966 - rpc_get |= ggml_backend_op_alloc_size_may_expand(tensor->op); + rpc_get |= ggml_op_alloc_size_may_expand(tensor->op); if (rpc_get) { ggml_backend_rpc_buffer_type_context * buft_ctx = (ggml_backend_rpc_buffer_type_context *)buft->context; From 140e57a4e20af04e8e00930798474a4707f0ee47 Mon Sep 17 00:00:00 2001 From: Daniel Bevenius Date: Fri, 4 Sep 2026 10:28:23 +0200 Subject: [PATCH 101/104] ggml : replace compile definitions with version.h.in (llama/28364) This commit adds a cmake version configuration file to replace the current compile definition solution for the version. The motivation for this change is that I made a mistake and did not take into consideration that the compile definition means that this will become a compiler flag for all sources in the target. This means that when a version update happens that will recompile all sources in the target even if they have not changed. Refs: https://github.com/ggml-org/llama.cpp/pull/28278 --- ggml/CMakeLists.txt | 4 ---- ggml/src/CMakeLists.txt | 4 +++- ggml/src/ggml-version.h.in | 4 ++++ ggml/src/ggml.c | 1 + 4 files changed, 8 insertions(+), 5 deletions(-) create mode 100644 ggml/src/ggml-version.h.in diff --git a/ggml/CMakeLists.txt b/ggml/CMakeLists.txt index 634e18c5422..b75506a26b0 100644 --- a/ggml/CMakeLists.txt +++ b/ggml/CMakeLists.txt @@ -404,10 +404,6 @@ write_basic_package_version_file( VERSION ${GGML_INSTALL_VERSION} COMPATIBILITY SameMajorVersion) -target_compile_definitions(ggml-base PRIVATE - GGML_VERSION="${GGML_INSTALL_VERSION}" - GGML_COMMIT="${GGML_BUILD_COMMIT}" -) message(STATUS "ggml version: ${GGML_INSTALL_VERSION}") message(STATUS "ggml commit: ${GGML_BUILD_COMMIT}") diff --git a/ggml/src/CMakeLists.txt b/ggml/src/CMakeLists.txt index 96535b49fa8..94773200021 100644 --- a/ggml/src/CMakeLists.txt +++ b/ggml/src/CMakeLists.txt @@ -213,7 +213,9 @@ set_target_properties(ggml-base PROPERTIES SOVERSION ${GGML_VERSION_MAJOR} ) -target_include_directories(ggml-base PRIVATE .) +configure_file(ggml-version.h.in ${CMAKE_CURRENT_BINARY_DIR}/ggml-version.h @ONLY) + +target_include_directories(ggml-base PRIVATE . ${CMAKE_CURRENT_BINARY_DIR}) if (GGML_BACKEND_DL) target_compile_definitions(ggml-base PUBLIC GGML_BACKEND_DL) endif() diff --git a/ggml/src/ggml-version.h.in b/ggml/src/ggml-version.h.in new file mode 100644 index 00000000000..37de362977b --- /dev/null +++ b/ggml/src/ggml-version.h.in @@ -0,0 +1,4 @@ +#pragma once + +#define GGML_VERSION "@GGML_VERSION@" +#define GGML_COMMIT "@GGML_BUILD_COMMIT@" diff --git a/ggml/src/ggml.c b/ggml/src/ggml.c index 2d5fdb7c103..6257cdbe582 100644 --- a/ggml/src/ggml.c +++ b/ggml/src/ggml.c @@ -1,6 +1,7 @@ #define _CRT_SECURE_NO_DEPRECATE // Disables "unsafe" warnings on Windows #define _USE_MATH_DEFINES // For M_PI on MSVC +#include "ggml-version.h" #include "ggml-backend.h" #include "ggml-impl.h" #include "ggml-threading.h" From 11d4eec830119e58f880df52df97f63a3a26f6c5 Mon Sep 17 00:00:00 2001 From: Niklas Wenzel Date: Fri, 4 Sep 2026 11:46:31 +0200 Subject: [PATCH 102/104] metal : add remaining fa-vec tunings for M3 Max (llama/28373) --- ggml/src/ggml-metal/ggml-metal-tuning.cpp | 145 ++++++++++++++++++++++ 1 file changed, 145 insertions(+) diff --git a/ggml/src/ggml-metal/ggml-metal-tuning.cpp b/ggml/src/ggml-metal/ggml-metal-tuning.cpp index b66fe65240f..7de01fac12f 100644 --- a/ggml/src/ggml-metal/ggml-metal-tuning.cpp +++ b/ggml/src/ggml-metal/ggml-metal-tuning.cpp @@ -1826,6 +1826,151 @@ constexpr fa_vec_entry_t fa_vec_tuned_table[] = { { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_F16, 576, 512, 2, 2 }, { 4, 4 } }, { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_F16, 576, 512, 2, 3 }, { 4, 2 } }, { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_F16, 576, 512, 3, 1 }, { 4, 2 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q4_0, 32, 32, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q4_0, 32, 32, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q4_0, 32, 32, 3, 2 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q4_0, 32, 32, 3, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q4_0, 32, 32, 3, 4 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q4_0, 64, 64, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q4_0, 64, 64, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q4_0, 96, 96, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q4_0, 96, 96, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q4_0, 96, 96, 1, 4 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q4_0, 96, 96, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q4_0, 96, 96, 3, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q4_0, 128, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q4_0, 128, 128, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q4_0, 192, 192, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q4_0, 192, 192, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q4_0, 192, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q4_0, 192, 128, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q4_0, 192, 128, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q4_0, 192, 128, 3, 3 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q4_0, 256, 256, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q4_0, 256, 256, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q4_0, 320, 256, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q4_0, 320, 256, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q4_0, 512, 512, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q4_0, 512, 512, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q4_0, 576, 512, 2, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q4_0, 576, 512, 3, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q4_0, 576, 512, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q4_0, 576, 512, 1, 1 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q4_0, 576, 512, 1, 2 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q4_0, 576, 512, 1, 3 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q4_0, 576, 512, 1, 4 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q4_1, 32, 32, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q4_1, 32, 32, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q4_1, 32, 32, 3, 2 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q4_1, 32, 32, 3, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q4_1, 32, 32, 3, 4 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q4_1, 64, 64, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q4_1, 64, 64, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q4_1, 96, 96, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q4_1, 96, 96, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q4_1, 96, 96, 1, 4 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q4_1, 96, 96, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q4_1, 96, 96, 3, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q4_1, 128, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q4_1, 128, 128, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q4_1, 128, 128, 1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q4_1, 128, 128, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q4_1, 128, 128, 1, 4 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q4_1, 128, 128, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q4_1, 128, 128, 3, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q4_1, 192, 192, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q4_1, 192, 192, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q4_1, 192, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q4_1, 192, 128, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q4_1, 192, 128, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q4_1, 192, 128, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q4_1, 192, 128, 3, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q4_1, 256, 256, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q4_1, 256, 256, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q4_1, 320, 256, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q4_1, 576, 512, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q4_1, 576, 512, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q5_0, 32, 32, -1, 1 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q5_0, 32, 32, 1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q5_0, 32, 32, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q5_0, 32, 32, 2, 4 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q5_0, 32, 32, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q5_0, 64, 64, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q5_0, 64, 64, -1, 1 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q5_0, 64, 64, 1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q5_0, 64, 64, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q5_0, 64, 64, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q5_0, 96, 96, -1, 1 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q5_0, 96, 96, 1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q5_0, 96, 96, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q5_0, 96, 96, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q5_0, 96, 96, 3, 4 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q5_0, 128, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q5_0, 128, 128, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q5_0, 192, 192, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q5_0, 192, 192, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q5_0, 192, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q5_0, 192, 128, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q5_0, 256, 256, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q5_0, 256, 256, -1, 1 }, { 2, 2 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q5_0, 256, 256, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q5_0, 256, 256, 1, 4 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q5_0, 256, 256, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q5_0, 256, 256, 3, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q5_0, 320, 256, -1, 1 }, { 2, 2 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q5_0, 320, 256, 1, 2 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q5_0, 320, 256, 2, 2 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q5_0, 320, 256, 3, 2 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q5_0, 320, 256, 3, 4 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q5_0, 512, 512, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q5_0, 512, 512, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q5_0, 576, 512, 2, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q5_0, 576, 512, 2, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q5_0, 576, 512, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q5_0, 576, 512, 2, 3 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q5_0, 576, 512, 2, 4 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q5_1, 32, 32, -1, 1 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q5_1, 32, 32, 1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q5_1, 32, 32, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q5_1, 32, 32, 2, 4 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q5_1, 32, 32, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q5_1, 64, 64, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q5_1, 64, 64, -1, 1 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q5_1, 64, 64, 1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q5_1, 64, 64, 2, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q5_1, 64, 64, 3, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q5_1, 96, 96, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q5_1, 96, 96, 1, 2 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q5_1, 96, 96, 1, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q5_1, 96, 96, 2, 2 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q5_1, 96, 96, 2, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q5_1, 96, 96, 3, 3 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q5_1, 128, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q5_1, 128, 128, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q5_1, 128, 128, 2, 2 }, { 4, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q5_1, 192, 192, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q5_1, 192, 192, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q5_1, 192, 128, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q5_1, 192, 128, -1, 1 }, { 2, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q5_1, 256, 256, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q5_1, 256, 256, -1, 1 }, { 2, 2 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q5_1, 256, 256, 1, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q5_1, 256, 256, 1, 4 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q5_1, 256, 256, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q5_1, 256, 256, 3, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q5_1, 320, 256, -1, 1 }, { 2, 2 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q5_1, 320, 256, 1, 2 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q5_1, 320, 256, 2, 2 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q5_1, 320, 256, 2, 4 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q5_1, 320, 256, 3, 2 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q5_1, 320, 256, 3, 4 }, { 1, 2 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q5_1, 512, 512, -1, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q5_1, 512, 512, -1, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q5_1, 576, 512, 2, 0 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q5_1, 576, 512, 2, 1 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q5_1, 576, 512, 2, 2 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q5_1, 576, 512, 2, 3 }, { 1, 4 } }, + { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q5_1, 576, 512, 2, 4 }, { 1, 4 } }, { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q8_0, 32, 32, -1, 1 }, { 4, 4 } }, { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q8_0, 32, 32, 1, 1 }, { 2, 4 } }, { { GGML_METAL_DEVICE_M3_MAX, GGML_TYPE_Q8_0, 32, 32, 1, 4 }, { 2, 4 } }, From a937f4e8efba28c0da4187617434676572355107 Mon Sep 17 00:00:00 2001 From: Georgi Gerganov Date: Fri, 4 Sep 2026 13:24:44 +0300 Subject: [PATCH 103/104] ggml : bump version to 0.23.0 (ggml/1618) --- ggml/CMakeLists.txt | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/ggml/CMakeLists.txt b/ggml/CMakeLists.txt index b75506a26b0..d76ed8ab049 100644 --- a/ggml/CMakeLists.txt +++ b/ggml/CMakeLists.txt @@ -4,7 +4,7 @@ project("ggml" C CXX ASM) ### GGML Version set(GGML_VERSION_MAJOR 0) -set(GGML_VERSION_MINOR 22) +set(GGML_VERSION_MINOR 23) set(GGML_VERSION_PATCH 0) set(GGML_VERSION_BASE "${GGML_VERSION_MAJOR}.${GGML_VERSION_MINOR}.${GGML_VERSION_PATCH}") From 52a939a2a762224e255d366c1182b2af4dd1a032 Mon Sep 17 00:00:00 2001 From: Georgi Gerganov Date: Fri, 4 Sep 2026 13:40:00 +0300 Subject: [PATCH 104/104] sync : ggml --- scripts/sync-ggml.last | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/scripts/sync-ggml.last b/scripts/sync-ggml.last index 601c1108bb1..7b44a311abd 100644 --- a/scripts/sync-ggml.last +++ b/scripts/sync-ggml.last @@ -1 +1 @@ -36da57138425487184aa1da2eee2cde155909c6f +e91ded11bdcd78c42f9c8d3978ff6686eb4c1226