Commit 847f447c3 for llama.cpp

commit 847f447c310a3784d623a953e6f0bb07f04c0f27
Author: cwriter <silvan.niederer@bluewin.ch>
Date:   Thu Oct 8 07:41:13 2026 +0200

    sycl: add grouped MoE XMX GEMM (#29245)

    Co-authored-by: cwriter <cwriter@localhost>

diff --git a/docs/backend/SYCL.md b/docs/backend/SYCL.md
index 70f5f2eff..c2ae0c1ba 100644
--- a/docs/backend/SYCL.md
+++ b/docs/backend/SYCL.md
@@ -816,6 +816,10 @@ User can use the device management in [docs/multi-gpu.md](https://github.com/ggm
 | GGML_SYCL_MKL_FA_DIAG | 0 (default) or 1 | Enable output fingerprinting for MKL flash attention. Dumps the first 64 float output values for the first 6 FA calls with n_kv ≥ 1024, labeled with kernel type (MKL/TILE/VEC) for cross-kernel comparison. |
 | GGML_SYCL_ENABLE_FUSION | 0 or 1 (default) | Enable fused-kernel dispatch in graph compute. Unsupported types and layouts fall back to the standalone op kernels. See `ggml_sycl_can_fuse()`. |
 | GGML_SYCL_ENABLE_ESIMD | 0 or 1 (default)| Enable ESIMD kernels when available. |
+| GGML_SYCL_XMX_GATHER_TYPES | decimal bitmask, all bits set (default) | Weight formats that may use the XMX dequant-GEMM paths, which dequantize weights straight into the XMX tiles. This speeds up prompt processing of MoE models on GPUs with XMX units (Arc A- and B-series, Arc Pro, Data Center GPU Max), for example pp512 of Qwen3-30B-A3B UD-IQ3_XXS by about 50% on an Arc Pro B60. Bits:<br>* 1: IQ4_NL, 2: IQ3_S, 4: IQ4_XS, 8: IQ3_XXS, 16: IQ2_XXS, 32: IQ2_XS, 64: IQ2_S, 128: IQ1_S, 256: IQ1_M<br>* 512: Q8_0, 1024: Q4_K, 2048: Q5_K, 4096: Q6_K (MoE `MUL_MAT_ID` only)<br>Add values to combine them, for example `3` for IQ4_NL and IQ3_S; `0` disables the paths. A set bit does not force the path: batches of more than 64 tokens per expert or row lengths that are not a multiple of 256 (32 for IQ4_NL and Q8_0) use the library GEMM. |
+| GGML_SYCL_XMX_GATHER_SHAPES | decimal bitmask, 255 (default) | XMX `joint_matrix` combinations the paths of `GGML_SYCL_XMX_GATHER_TYPES` may use; the operand type comes from `GGML_SYCL_DYNAMIC_PRECISION` and the best supported combination is picked automatically (logged as `fg_pick_combo`). Bits:<br>* Xe2, Xe3, Xe-HPC: 1: f16 8x16x16, 2: f16 16x16x16, 4: f16 32x64x16, 8: f16 32x64x32, 32: tf32 8x16x8, 64: bf16 8x16x16<br>* Xe-HPG (Arc A770, ARL-H): 16: f16 8x8x16, 128: bf16 8x8x16<br>Clear a bit to exclude a combination, or set a single bit to force one for testing. |
+| GGML_SYCL_DYNAMIC_PRECISION | `F16` (default with `GGML_SYCL_F16=ON`), `BF16`, `TF32` or `F32` (default otherwise) | Operand type of the XMX dequant-GEMM paths (`GGML_SYCL_XMX_GATHER_TYPES`); accumulation is always f32. `F16` is the fastest, but activations above 65504 overflow. `BF16` keeps the f32 range at a 7-bit mantissa, `TF32` keeps the range and the f16 mantissa but is about 30% slower and needs Xe2, Xe3 or Xe-HPC, and `F32` turns the XMX paths off. Ops that request a higher src1 precision ([TAG_GGML_PREC]) get it regardless of this setting. |
+| GGML_SYCL_DYNAMIC_REQUIRED_PRECISION | `F32` (default), `TF32`, `BF16` or `F16` | Lowest type the XMX paths may use for an op that requests an F32 src1, such as Mistral 4 `ffn_down_exps`. The default runs such ops on the library f32 GEMM; `TF32` or `BF16` trade mantissa for speed while keeping the f32 range. `F16` ignores the request and can overflow; it is meant for testing only. |
 | GGML_SYCL_MMVQ_WIDE | 0 or 1 (default) | Use the wide-load variant of the reordered Q8_0 mat-vec kernel, which reads four contiguous dwords per operand instead of one value at a time. Set to 0 to fall back to the per-value loads. Only affects Q8_0 weights in the reordered layout. |
 | GGML_SYCL_SPARSE_FA | 0 (default) or 1 | Enable Sparse Flash-attention.|
 | GGML_SYCL_SPARSE_FA_DEBUG | 0 (default) or 1 | Enable to debug for Sparse Flash-attention.|
diff --git a/ggml/src/ggml-sycl/CMakeLists.txt b/ggml/src/ggml-sycl/CMakeLists.txt
index d2196f74d..1697ecd95 100644
--- a/ggml/src/ggml-sycl/CMakeLists.txt
+++ b/ggml/src/ggml-sycl/CMakeLists.txt
@@ -221,4 +221,27 @@ if (GGML_SYCL_DEVICE_ARCH)
         "SHELL:-Xsycl-target-backend=spir64_gen \"-device ${GGML_SYCL_DEVICE_ARCH}\""
         -fsycl-max-parallel-link-jobs=${GGML_SYCL_MAX_PARALLEL_LINK_JOBS}
     )
+
+    # The XMX dequant-GEMM tiles need the sub-group size of the target: 8 on Xe-HPG (DG2, ARL-H),
+    # 16 on Xe-HPC and Xe2 or newer. ocloc fails on the other size, so build only the one that fits.
+    # 0 (unknown name, mixed list, or no XMX) builds no XMX tile and the path stays off.
+    set(_ggml_sycl_xmx_sg "")
+    string(TOLOWER "${GGML_SYCL_DEVICE_ARCH}" _ggml_sycl_archs)
+    string(REPLACE "," ";" _ggml_sycl_archs "${_ggml_sycl_archs}")
+    foreach(_arch IN LISTS _ggml_sycl_archs)
+        if (_arch MATCHES "^(dg2|acm|ats-m|arl-h|xe-hpg|12\\.5[567]\\.|12\\.74\\.)")
+            set(_sg 8)
+        elseif (_arch MATCHES "^(pvc|bmg|lnl|ptl|wcl|nvl|cri|xe2|xe3|xe-hpc|12\\.60\\.|20\\.|30\\.)")
+            set(_sg 16)
+        else()
+            set(_sg 0)
+        endif()
+        if (_ggml_sycl_xmx_sg STREQUAL "" OR _ggml_sycl_xmx_sg EQUAL _sg)
+            set(_ggml_sycl_xmx_sg ${_sg})
+        else()
+            set(_ggml_sycl_xmx_sg 0)
+        endif()
+    endforeach()
+    message(STATUS "GGML_SYCL_DEVICE_ARCH: XMX dequant-GEMM sub-group size ${_ggml_sycl_xmx_sg} (0 = off)")
+    target_compile_definitions(ggml-sycl PRIVATE GGML_SYCL_XMX_AOT_SG=${_ggml_sycl_xmx_sg})
 endif()
diff --git a/ggml/src/ggml-sycl/common.hpp b/ggml/src/ggml-sycl/common.hpp
index 5e904bd2a..475509372 100644
--- a/ggml/src/ggml-sycl/common.hpp
+++ b/ggml/src/ggml-sycl/common.hpp
@@ -65,6 +65,61 @@ extern int g_ggml_sycl_enable_fusion;
 extern int g_ggml_sycl_enable_esimd;
 extern int g_ggml_sycl_mmvq_wide;
 extern int g_ggml_sycl_prioritize_dmmv;
+
+// Which quantized weight formats may take the XMX dequant-GEMM paths. A bitmask rather than one
+// flag per path, so a format can be enabled or measured on its own and adding a format is one bit.
+enum ggml_sycl_xmx_gather_type {
+    GGML_SYCL_XMX_GATHER_IQ4_NL   = 1 << 0,
+    GGML_SYCL_XMX_GATHER_IQ3_S    = 1 << 1,
+    GGML_SYCL_XMX_GATHER_IQ4_XS   = 1 << 2,
+    GGML_SYCL_XMX_GATHER_IQ3_XXS  = 1 << 3,
+    GGML_SYCL_XMX_GATHER_IQ2_XXS  = 1 << 4,
+    GGML_SYCL_XMX_GATHER_IQ2_XS   = 1 << 5,
+    GGML_SYCL_XMX_GATHER_IQ2_S    = 1 << 6,
+    GGML_SYCL_XMX_GATHER_IQ1_S    = 1 << 7,
+    GGML_SYCL_XMX_GATHER_IQ1_M    = 1 << 8,
+    GGML_SYCL_XMX_GATHER_Q8_0     = 1 << 9,
+    GGML_SYCL_XMX_GATHER_Q4_K     = 1 << 10,
+    GGML_SYCL_XMX_GATHER_Q5_K     = 1 << 11,
+    GGML_SYCL_XMX_GATHER_Q6_K     = 1 << 12,
+};
+static constexpr int GGML_SYCL_XMX_GATHER_TYPES_DEFAULT = ~0;
+extern int g_ggml_sycl_xmx_gather_types;
+// Which joint_matrix combinations the XMX dequant-GEMM paths may use, one bit each (see fused-gemm.cpp).
+// GGML_SYCL_DYNAMIC_PRECISION picks the operand type, this mask the combinations of that type.
+static constexpr int GGML_SYCL_XMX_GATHER_SHAPES_DEFAULT = 0xff;
+extern int g_ggml_sycl_xmx_gather_shapes;
+
+// GGML_SYCL_DYNAMIC_PRECISION: operand type of the XMX dequant-GEMM paths. F32 turns them off and
+// keeps the library GEMM in f32. A src1 precision request of an op [TAG_GGML_PREC] is always met.
+enum ggml_sycl_dynamic_precision {
+    GGML_SYCL_DYNAMIC_PRECISION_F16,
+    GGML_SYCL_DYNAMIC_PRECISION_BF16,
+    GGML_SYCL_DYNAMIC_PRECISION_TF32,
+    GGML_SYCL_DYNAMIC_PRECISION_F32,
+};
+#ifdef GGML_SYCL_F16
+static constexpr int GGML_SYCL_DYNAMIC_PRECISION_DEFAULT = GGML_SYCL_DYNAMIC_PRECISION_F16;
+#else
+static constexpr int GGML_SYCL_DYNAMIC_PRECISION_DEFAULT = GGML_SYCL_DYNAMIC_PRECISION_F32;
+#endif
+extern int g_ggml_sycl_dynamic_precision;
+// GGML_SYCL_DYNAMIC_REQUIRED_PRECISION: the XMX type an F32 src1 request may run on instead of f32
+// (TF32, or BF16 which also allows tf32). F32 (default): none. F16: src1 requests are ignored.
+extern int g_ggml_sycl_dynamic_required_precision;
+
+// [TAG_GGML_PREC] src1 precision request of the MUL_MAT/MUL_MAT_ID op dst
+static inline int32_t ggml_sycl_src1_prec(const ggml_tensor * dst) {
+    return g_ggml_sycl_dynamic_required_precision == GGML_SYCL_DYNAMIC_PRECISION_F16 ? GGML_PREC_UNDEFINED :
+                                                                                       dst->op_params[3];
+}
+
+// [TAG_GGML_PREC] the library GEMM and dmmv may convert src1 of the MUL_MAT/MUL_MAT_ID op dst to f16
+static inline bool ggml_sycl_src1_f16_ok(const ggml_tensor * dst) {
+    const int32_t src1_prec = ggml_sycl_src1_prec(dst);
+    return g_ggml_sycl_dynamic_precision != GGML_SYCL_DYNAMIC_PRECISION_F32 &&
+           (src1_prec == GGML_PREC_UNDEFINED || src1_prec >= GGML_PREC_F16);
+}
 extern int g_ggml_sycl_enable_flash_attention;
 extern int g_ggml_sycl_dev2dev_memcpy;
 extern int g_ggml_sycl_fa_onednn;
@@ -333,6 +388,12 @@ struct mmid_row_mapping {
     int32_t i2;
 };

+struct ggml_sycl_gg_tile {
+    int32_t expert;
+    int32_t n0;
+    int32_t n1;
+};
+
 namespace sycl_ex = sycl::ext::oneapi::experimental;
 struct ggml_backend_sycl_context {
     int device;
@@ -410,6 +471,7 @@ struct ggml_backend_sycl_context {
     std::unique_ptr<ggml_sycl_pool> host_pools[GGML_SYCL_MAX_DEVICES];

     std::vector<mmid_row_mapping> mmid_row_mapping_host;
+    std::vector<ggml_sycl_gg_tile> mmid_tile_schedule_host;

     static std::unique_ptr<ggml_sycl_pool> new_pool_for_device(queue_ptr qptr, int device);

diff --git a/ggml/src/ggml-sycl/fused-gemm.cpp b/ggml/src/ggml-sycl/fused-gemm.cpp
new file mode 100644
index 000000000..38c69bc6b
--- /dev/null
+++ b/ggml/src/ggml-sycl/fused-gemm.cpp
@@ -0,0 +1,1093 @@
+#include "fused-gemm.hpp"
+
+#include <sycl/ext/oneapi/matrix/matrix.hpp>
+
+#include <algorithm>
+#include <string>
+#include <tuple>
+#include <mutex>
+#include <set>
+#include <unordered_map>
+
+namespace mx = sycl::ext::oneapi::experimental::matrix;
+
+// FG_ / fg_ is short for fused GEMM: the weights are dequantized inside the GEMM, into the XMX tiles.
+
+// A k step is one 32-value weight sub-block; iq3_s and the other superblock formats split their
+// superblock into steps of this width. The sub-groups of a work-group each walk their own K range
+// and are summed at the end.
+static constexpr int FG_BK     = QK4_NL;
+static constexpr int FG_KSPLIT = 4;
+
+// Element traits of one joint_matrix operand type. The A stage and the B pack compute in f32 and
+// convert once, in registers, when they write the element, so any type costs the same one pass.
+//   store: storage in SLM (A) and in the packed B buffer
+//   mtype: matrix_type in matrix_combinations
+//   mode:  GGML_SYCL_DYNAMIC_PRECISION value that selects this type
+//   src:   ggml type that needs no conversion into this type (GGML_TYPE_COUNT: none)
+//   slow:  XMX throughput class, 0 is fastest. f16 and bf16 share the DPAS rate; tf32 does half the
+//          K per instruction. B60, Qwen3-30B-A3B pp512: f16 1108, bf16 1000, tf32 751 t/s
+template <typename T> struct fg_elem;
+
+template <> struct fg_elem<sycl::half> {
+    using store = sycl::half;
+    using pair  = sycl::half2;
+    static constexpr mx::matrix_type mtype = mx::matrix_type::fp16;
+    static constexpr int             mode  = GGML_SYCL_DYNAMIC_PRECISION_F16;
+    static constexpr ggml_type       src   = GGML_TYPE_F16;
+    static constexpr int             mant  = 10;
+    static constexpr int             slow  = 0;
+    static store cvt(float x) { return (store) x; }
+    static pair make(float x, float y) { return pair((store) x, (store) y); }
+};
+
+// tf32 rounds to nearest even with plain bit ops: round_to_tf32 needs a SPIR-V extension that the
+// DG2 AOT target rejects
+static inline uint32_t fg_round_bits(float x, int drop) {
+    const uint32_t u = sycl::bit_cast<uint32_t>(x);
+    if ((u & 0x7f800000u) == 0x7f800000u) {
+        return (u & 0x7fffffu) ? u | (1u << drop) : u; // nan stays nan
+    }
+    return u + ((1u << (drop - 1)) - 1) + ((u >> drop) & 1);
+}
+
+struct alignas(4) fg_bf16x2 {
+    sycl::ext::oneapi::bfloat16 x, y;
+};
+
+template <> struct fg_elem<sycl::ext::oneapi::bfloat16> {
+    using store = sycl::ext::oneapi::bfloat16;
+    using pair  = fg_bf16x2;
+    static constexpr mx::matrix_type mtype = mx::matrix_type::bf16;
+    static constexpr int             mode  = GGML_SYCL_DYNAMIC_PRECISION_BF16;
+    static constexpr ggml_type       src   = GGML_TYPE_BF16;
+    static constexpr int             mant  = 7;
+    static constexpr int             slow  = 0;
+    static store cvt(float x) { return store(x); }
+    static pair make(float x, float y) { return { cvt(x), cvt(y) }; }
+};
+
+// tf32 keeps f32 range and f16 mantissa, in f32 storage
+template <> struct fg_elem<mx::precision::tf32> {
+    using store = float;
+    using pair  = sycl::float2;
+    static constexpr mx::matrix_type mtype = mx::matrix_type::tf32;
+    static constexpr int             mode  = GGML_SYCL_DYNAMIC_PRECISION_TF32;
+    static constexpr ggml_type       src   = GGML_TYPE_COUNT;
+    static constexpr int             mant  = 10;
+    static constexpr int             slow  = 1;
+    static store cvt(float x) { return sycl::bit_cast<float>(fg_round_bits(x, 13) & ~0x1fffu); }
+    static pair make(float x, float y) { return pair(cvt(x), cvt(y)); }
+};
+
+// One joint_matrix combination (A type, B type, TM x TN x TK, sub-group size; C and D are f32) and
+// the tiling built on it. A sub-group owns SG_ROWS rows of A (at least 16) and BN (at least 32)
+// columns of B. A and B may differ: the device lists the pairs it supports.
+template <typename TA, typename TB, int TM_, int TN_, int TK_, int SG_> struct fg_combo {
+    using ta  = TA;
+    using tb  = TB;
+    using EA  = fg_elem<TA>;
+    using EB  = fg_elem<TB>;
+    using tsa = typename EA::store;
+    using tsb = typename EB::store;
+    static constexpr int TM = TM_;
+    static constexpr int TN = TN_;
+    static constexpr int TK = TK_;
+    static constexpr int SG = SG_;
+    static constexpr int VNNI    = 4 / sizeof(tsb);  // K rows of B packed in one 32-bit word
+    static constexpr int SG_ROWS = TM > 16 ? TM : 16;
+    static constexpr int RPL     = SG_ROWS / SG;     // A rows one lane decodes per k step
+    static constexpr int MT      = SG_ROWS / TM;
+    static constexpr int BN      = TN > 32 ? TN : 32;
+    static constexpr int NT      = BN / TN;
+    static constexpr int WG_SIZE = FG_KSPLIT * SG;
+    static constexpr mx::layout b_layout = VNNI == 1 ? mx::layout::row_major : mx::layout::ext_intel_packed;
+    // a 64-wide N is mostly padding here and a 32x64 f32 accumulator needs 128 registers per lane,
+    // so it spills: 13x slower on B60
+    static constexpr bool efficient = TN <= 32;
+    static_assert(SG_ROWS % SG == 0 && SG_ROWS % TM == 0 && BN % TN == 0 && FG_BK % TK == 0, "bad tile");
+    static_assert(BN <= GGML_SYCL_FG_MAX_N, "header gate must cover the tile width");
+};
+
+using fg_half = sycl::half;
+using fg_bf16 = sycl::ext::oneapi::bfloat16;
+using fg_tf32 = mx::precision::tf32;
+
+// One bit of GGML_SYCL_XMX_GATHER_SHAPES per combination. Only combinations some device lists in
+// matrix_combinations are built (appendix of sycl_ext_oneapi_matrix and the runtime's own list).
+template <typename F> static void fg_visit_combo(int idx, F && f);
+static constexpr int FG_N_COMBOS = 8;
+
+// A spir64_gen AOT build (GGML_SYCL_XMX_AOT_SG) drops the combinations of the other sub-group size
+// entirely: ocloc rejects even an empty kernel that asks for a sub-group size it lacks.
+template <int SG> static constexpr bool fg_listed() {
+#if defined(GGML_SYCL_XMX_AOT_SG)
+    return SG == GGML_SYCL_XMX_AOT_SG;
+#else
+    return true;
+#endif
+}
+
+template <typename S, typename F> static void fg_call_combo(F && f) {
+    if constexpr (fg_listed<S::SG>()) {
+        f(S{});
+    }
+}
+
+template <typename F> static void fg_visit_combo(int idx, F && f) {
+    switch (idx) {
+        case 0: fg_call_combo<fg_combo<fg_half, fg_half, 8, 16, 16, 16>>(f);  break; // Xe2, Xe3, Xe-HPC
+        case 1: fg_call_combo<fg_combo<fg_half, fg_half, 16, 16, 16, 16>>(f); break; // Xe2, Xe3, Xe-HPC
+        case 2: fg_call_combo<fg_combo<fg_half, fg_half, 32, 64, 16, 16>>(f); break; // Xe2, Xe3, Xe-HPC
+        case 3: fg_call_combo<fg_combo<fg_half, fg_half, 32, 64, 32, 16>>(f); break; // Xe2, Xe3, Xe-HPC
+        case 4: fg_call_combo<fg_combo<fg_half, fg_half, 8, 8, 16, 8>>(f);    break; // Xe-HPG (Arc A), ARL-H
+        case 5: fg_call_combo<fg_combo<fg_tf32, fg_tf32, 8, 16, 8, 16>>(f);   break; // Xe2, Xe3, Xe-HPC
+        case 6: fg_call_combo<fg_combo<fg_bf16, fg_bf16, 8, 16, 16, 16>>(f);  break; // Xe2, Xe3, Xe-HPC
+        case 7: fg_call_combo<fg_combo<fg_bf16, fg_bf16, 8, 8, 16, 8>>(f);    break; // Xe-HPG (Arc A), ARL-H
+        default: GGML_ABORT("bad XMX combination %d", idx);
+    }
+}
+
+// AOT with -fsycl-targets=intel_gpu_*: compile each tile body only for targets with its sub-group
+// size, since IGC fails on the other ones. A JIT build keeps them all, but each combination lands in
+// its own device image (joint_matrix is an optional kernel feature) and only a combination the
+// device reports is launched, so the runtime never asks IGC for the others.
+#if defined(__SYCL_DEVICE_ONLY__)
+#    if __SYCL_TARGET_INTEL_GPU_ACM_G10__ || __SYCL_TARGET_INTEL_GPU_ACM_G11__ || __SYCL_TARGET_INTEL_GPU_ACM_G12__ || \
+        __SYCL_TARGET_INTEL_GPU_ARL_H__
+#        define FG_AOT_SG 8
+#    elif __SYCL_TARGET_INTEL_GPU_PVC__ || __SYCL_TARGET_INTEL_GPU_PVC_VG__ || __SYCL_TARGET_INTEL_GPU_BMG_G21__ || \
+        __SYCL_TARGET_INTEL_GPU_BMG_G31__ || __SYCL_TARGET_INTEL_GPU_LNL_M__ || __SYCL_TARGET_INTEL_GPU_PTL_H__ ||   \
+        __SYCL_TARGET_INTEL_GPU_PTL_U__ || __SYCL_TARGET_INTEL_GPU_WCL__ || __SYCL_TARGET_INTEL_GPU_NVL_S__ ||       \
+        __SYCL_TARGET_INTEL_GPU_NVL_U__ || __SYCL_TARGET_INTEL_GPU_NVL_P__
+#        define FG_AOT_SG 16
+#    elif __SYCL_TARGET_INTEL_GPU_TGLLP__ || __SYCL_TARGET_INTEL_GPU_RKL__ || __SYCL_TARGET_INTEL_GPU_ADL_S__ || \
+        __SYCL_TARGET_INTEL_GPU_ADL_P__ || __SYCL_TARGET_INTEL_GPU_ADL_N__ || __SYCL_TARGET_INTEL_GPU_DG1__ ||   \
+        __SYCL_TARGET_INTEL_GPU_MTL_U__ || __SYCL_TARGET_INTEL_GPU_MTL_H__
+#        define FG_AOT_SG 0 // no XMX
+#    endif
+#endif
+
+template <int SG> static constexpr bool fg_built() {
+#if defined(FG_AOT_SG)
+    return SG == FG_AOT_SG;
+#else
+    return true;
+#endif
+}
+
+// Upper bound on the tile count when total_rows rows are routed to n_as experts: the worst case gives
+// each expert one row and fills whole tiles with the rest. The bound depends only on the shape, not
+// on the routing, so the pool reuses one buffer every ubatch instead of keeping one per size seen.
+static constexpr int64_t grouped_gemm_max_tiles(int64_t total_rows, int64_t n_as, int64_t BN) {
+    return total_rows <= n_as ? total_rows : n_as + (total_rows - n_as) / BN;
+}
+// Tiles do not cross experts, so the bound is not ceil(total_rows / BN): 34 rows over 2 experts with
+// BN = 16 split 17 + 17 need 2 + 2 tiles, where the ceil gives 3.
+static_assert(grouped_gemm_max_tiles(34, 2, 16) == 4);
+
+// the device lists S with an f32 accumulator and output
+template <typename S> static bool fg_device_has_combo(const std::vector<mx::combination> & combinations) {
+    for (const auto & c : combinations) {
+        if (c.atype == S::EA::mtype && c.btype == S::EB::mtype && c.ctype == mx::matrix_type::fp32 &&
+            c.dtype == mx::matrix_type::fp32 &&
+            (c.max_msize >= (size_t) S::TM || c.msize == (size_t) S::TM) &&
+            (c.max_nsize >= (size_t) S::TN || c.nsize == (size_t) S::TN) &&
+            (c.max_ksize >= (size_t) S::TK || c.ksize == (size_t) S::TK)) {
+            return true;
+        }
+    }
+    return false;
+}
+
+template <typename T> static const char * fg_type_name() {
+    return std::is_same_v<T, fg_half> ? "f16" : std::is_same_v<T, fg_bf16> ? "bf16" : "tf32";
+}
+
+static std::string fg_combo_name(int idx) {
+    std::string name;
+    fg_visit_combo(idx, [&](auto s) {
+        using S = decltype(s);
+        name = std::string(fg_type_name<typename S::ta>()) + "x" + fg_type_name<typename S::tb>() + " " +
+               std::to_string(S::TM) + "x" + std::to_string(S::TN) + "x" + std::to_string(S::TK) + " sg" +
+               std::to_string(S::SG);
+    });
+    return name;
+}
+
+// Combinations this build has kernels for and the device lists, one bit each. Cached per device:
+// on a mixed box the first caller's verdict is not the others'.
+static int fg_device_combos(const sycl::device & dev) {
+    static std::mutex                            mtx;
+    static std::unordered_map<sycl::device, int> known;
+    std::lock_guard<std::mutex>                  lock(mtx);
+    const auto                                   it = known.find(dev);
+    if (it != known.end()) {
+        return it->second;
+    }
+    int available = 0;
+    try {
+        const auto combinations = dev.get_info<sycl::ext::oneapi::experimental::info::device::matrix_combinations>();
+        const auto sg_sizes     = dev.get_info<sycl::info::device::sub_group_sizes>();
+        for (int idx = 0; idx < FG_N_COMBOS; ++idx) {
+            fg_visit_combo(idx, [&](auto s) {
+                using S = decltype(s);
+                const bool sg_ok = std::find(sg_sizes.begin(), sg_sizes.end(), (size_t) S::SG) != sg_sizes.end();
+                if (sg_ok && fg_device_has_combo<S>(combinations)) {
+                    available |= 1 << idx;
+                }
+            });
+        }
+    } catch (const sycl::exception &) {
+        available = 0;
+    }
+    GGML_LOG_INFO("%s: %s: XMX dequant-GEMM combinations available 0x%x, allowed 0x%x\n", __func__,
+                  dev.get_info<sycl::info::device::name>().c_str(), available, g_ggml_sycl_xmx_gather_shapes);
+    known.emplace(dev, available);
+    return available;
+}
+
+// Rank of combination S for a src1 of type src1_type, lower is better. Order:
+//  1. throughput: a tile that does not spill, then the fastest type class of A and B
+//  2. B type equal to the src1 type, so the pack is a plain copy
+//  3. B at least as precise as f16
+//  4. the device's native DPAS tile (8 x SG x 32 bytes of K), then the largest M x K
+// A costs nothing to convert: the A stage emits any type at the same cost.
+template <typename S> static int64_t fg_rank(ggml_type src1_type) {
+    const int64_t spills  = !S::efficient;
+    const int64_t slow    = std::max(S::EA::slow, S::EB::slow);
+    const int64_t convert = S::EB::src != src1_type;
+    const int64_t lossy   = S::EB::mant < 10;
+    const int64_t foreign = !(S::TM == 8 && S::TN == S::SG);
+    const int64_t mk      = 1024 - S::TM * S::TK;
+    return ((((spills * 2 + slow) * 2 + convert) * 2 + lossy) * 2 + foreign) * 2048 + mk;
+}
+
+// whether XMX operands of type mode (a GGML_SYCL_DYNAMIC_PRECISION value) meet the src1 request
+// [TAG_GGML_PREC]. f16 lacks the f32 range that BF16 and F32 ask for; an F32 request goes only as far
+// down as GGML_SYCL_DYNAMIC_REQUIRED_PRECISION allows.
+static bool fg_mode_meets(int mode, int32_t src1_prec) {
+    if (src1_prec == GGML_PREC_UNDEFINED || src1_prec >= GGML_PREC_F16) {
+        return true;
+    }
+    if (src1_prec >= GGML_PREC_BF16) {
+        return mode != GGML_SYCL_DYNAMIC_PRECISION_F16;
+    }
+    switch (g_ggml_sycl_dynamic_required_precision) {
+        case GGML_SYCL_DYNAMIC_PRECISION_TF32: return mode == GGML_SYCL_DYNAMIC_PRECISION_TF32;
+        case GGML_SYCL_DYNAMIC_PRECISION_BF16: return mode != GGML_SYCL_DYNAMIC_PRECISION_F16;
+        default:                               return false;
+    }
+}
+
+// Best allowed combination for this call, or -1 if none. The type is GGML_SYCL_DYNAMIC_PRECISION if it
+// meets the src1 request; if not, bf16 then tf32 for a BF16 request (fastest first) and tf32 then bf16
+// for an F32 request (most mantissa first).
+static int fg_pick_combo(dpct::queue_ptr stream, ggml_type src1_type, int32_t src1_prec) {
+    const sycl::device dev     = stream->get_device();
+    const int          allowed = fg_device_combos(dev) & g_ggml_sycl_xmx_gather_shapes;
+    const bool         f32_req = src1_prec != GGML_PREC_UNDEFINED && src1_prec < GGML_PREC_BF16;
+    const int          modes[] = {
+        g_ggml_sycl_dynamic_precision,
+        f32_req ? GGML_SYCL_DYNAMIC_PRECISION_TF32 : GGML_SYCL_DYNAMIC_PRECISION_BF16,
+        f32_req ? GGML_SYCL_DYNAMIC_PRECISION_BF16 : GGML_SYCL_DYNAMIC_PRECISION_TF32,
+    };
+    int     best      = -1;
+    int64_t best_rank = 0;
+    for (int i = 0; i < 3 && best < 0; ++i) {
+        const int mode = modes[i];
+        if (!fg_mode_meets(mode, src1_prec)) {
+            continue;
+        }
+        for (int idx = 0; idx < FG_N_COMBOS; ++idx) {
+            if (!(allowed & (1 << idx))) {
+                continue;
+            }
+            fg_visit_combo(idx, [&](auto s) {
+                using S = decltype(s);
+                if (S::EA::mode != mode || S::EB::mode != mode) {
+                    return;
+                }
+                const int64_t rank = fg_rank<S>(src1_type);
+                if (best < 0 || rank < best_rank) {
+                    best      = idx;
+                    best_rank = rank;
+                }
+            });
+        }
+    }
+    // log each distinct decision once
+    static std::mutex                                      mtx;
+    static std::set<std::tuple<size_t, int, int32_t, int>> seen;
+    std::lock_guard<std::mutex>                            lock(mtx);
+    if (seen.emplace(std::hash<sycl::device>{}(dev), (int) src1_type, src1_prec, best).second) {
+        GGML_LOG_INFO("%s: src1 %s, src1 prec %d -> %s\n", __func__, ggml_type_name(src1_type), src1_prec,
+                      best >= 0 ? fg_combo_name(best).c_str() : "none (library GEMM)");
+    }
+    return best;
+}
+
+// src1 [N][K] -> packed [K/V][Npad][V] so B tiles load straight from global memory
+template <typename E, typename T_src>
+static void fused_gemm_pack_b(const T_src * y, typename E::store * packed, int N, int Npad, int K,
+                              dpct::queue_ptr stream) {
+    constexpr int V   = 4 / sizeof(typename E::store);
+    const int     kqs = K / V;
+    stream->parallel_for(sycl::range<1>((size_t) Npad * kqs), [=](sycl::id<1> id) {
+        const int idx = id[0];
+        const int n   = idx / kqs;
+        const int kq  = idx - n * kqs;
+        typename E::store vals[V] = {};
+        if (n < N) {
+            const T_src * src = y + (size_t) n * K + V * kq;
+#pragma unroll
+            for (int v = 0; v < V; ++v) {
+                vals[v] = E::cvt((float) src[v]);
+            }
+        }
+        typename E::store * out = packed + ((size_t) kq * Npad + n) * V;
+#pragma unroll
+        for (int v = 0; v < V; ++v) {
+            out[v] = vals[v];
+        }
+    });
+}
+
+// A stage: one lane owns one row and decodes FG_BK values of it per k step, with every scale
+// folded into the value so the mad below sees plain A elements. One overload per weight format.
+template <typename E>
+static __dpct_inline__ void fg_stage_a(const block_iq4_nl * __restrict__ xrow, const int kb, typename E::pair * a) {
+    const block_iq4_nl blk = xrow[kb];
+    const float        d   = (float) blk.d;
+#pragma unroll
+    for (int j = 0; j < QK4_NL / 2; j += 2) {
+        const uint8_t q0 = blk.qs[j];
+        const uint8_t q1 = blk.qs[j + 1];
+        a[j / 2]     = E::make(d * kvalues_iq4nl[q0 & 0xf], d * kvalues_iq4nl[q1 & 0xf]);
+        a[j / 2 + 8] = E::make(d * kvalues_iq4nl[q0 >> 4], d * kvalues_iq4nl[q1 >> 4]);
+    }
+}
+
+// iq3_s: k step kb is sub-block kb % 8 of superblock kb / 8. The superblock is 110 bytes, so read
+// only the fields of that sub-block instead of copying the block. Same decode as
+// dequantize_block_iq3_s: grid entries are taken as dwords and the sign bit is a plain shift.
+template <typename E>
+static __dpct_inline__ void fg_stage_a(const block_iq3_s * __restrict__ xrow, const int kb, typename E::pair * a) {
+    static_assert(QK_K == 256, "the iq3_s A stage assumes 8 sub-blocks per superblock");
+    const block_iq3_s * blk = xrow + kb / (QK_K / 32);
+    const int           ib8 = kb % (QK_K / 32);
+    const uint8_t *     qs  = blk->qs + 8 * ib8;
+    const int           qh  = blk->qh[ib8];
+    const float         d   = (float) blk->d * (1 + 2 * ((blk->scales[ib8 / 2] >> (4 * (ib8 % 2))) & 0xf));
+#pragma unroll
+    for (int il = 0; il < 4; ++il) {
+        const uint32_t grid1 = iq3s_grid[qs[2 * il + 0] | ((qh << (8 - 2 * il)) & 256)];
+        const uint32_t grid2 = iq3s_grid[qs[2 * il + 1] | ((qh << (7 - 2 * il)) & 256)];
+        const int      signs = blk->signs[4 * ib8 + il];
+#pragma unroll
+        for (int j = 0; j < 2; ++j) {
+            const float g1a = (float) ((grid1 >> (16 * j + 0)) & 0xff);
+            const float g1b = (float) ((grid1 >> (16 * j + 8)) & 0xff);
+            const float g2a = (float) ((grid2 >> (16 * j + 0)) & 0xff);
+            const float g2b = (float) ((grid2 >> (16 * j + 8)) & 0xff);
+            const int   s   = 2 * j;
+            a[4 * il + j]     = E::make(d * ((signs & (1 << (s + 0))) ? -g1a : g1a),
+                                        d * ((signs & (1 << (s + 1))) ? -g1b : g1b));
+            a[4 * il + j + 2] = E::make(d * ((signs & (1 << (s + 4))) ? -g2a : g2a),
+                                        d * ((signs & (1 << (s + 5))) ? -g2b : g2b));
+        }
+    }
+}
+
+// values per stored block, so a row of K values is K/qk blocks
+template <typename block_q_t> struct fg_block_traits;
+template <> struct fg_block_traits<block_iq4_nl> { static constexpr int qk = QK4_NL; };
+template <> struct fg_block_traits<block_iq3_s>  { static constexpr int qk = QK_K; };
+template <> struct fg_block_traits<block_iq3_xxs> { static constexpr int qk = QK_K; };
+template <> struct fg_block_traits<block_iq4_xs>  { static constexpr int qk = QK_K; };
+template <> struct fg_block_traits<block_iq2_xxs> { static constexpr int qk = QK_K; };
+template <> struct fg_block_traits<block_iq2_xs>  { static constexpr int qk = QK_K; };
+template <> struct fg_block_traits<block_iq2_s>   { static constexpr int qk = QK_K; };
+template <> struct fg_block_traits<block_iq1_s>   { static constexpr int qk = QK_K; };
+template <> struct fg_block_traits<block_iq1_m>   { static constexpr int qk = QK_K; };
+template <> struct fg_block_traits<block_q8_0>    { static constexpr int qk = QK8_0; };
+template <> struct fg_block_traits<block_q4_K>    { static constexpr int qk = QK_K; };
+template <> struct fg_block_traits<block_q5_K>    { static constexpr int qk = QK_K; };
+template <> struct fg_block_traits<block_q6_K>    { static constexpr int qk = QK_K; };
+
+// The A stages below are the dequantize_block_iq* kernels rewritten for one k step. There a
+// work-item handled one quarter (il) of one 32-wide sub-block (ib); here one lane produces the
+// whole step, so il becomes a loop and ib is kb inside the superblock. Each quarter yields 8
+// consecutive values, i.e. 4 pairs at a[4*il], so nothing larger than 8 floats is ever live.
+#define FG_SUPERBLOCK(T)                                                         \
+    static_assert(QK_K == 256, "the " #T " A stage assumes 8 sub-blocks per superblock"); \
+    const T * blk = xrow + kb / (QK_K / 32);                                     \
+    const int ib  = kb % (QK_K / 32)
+
+template <typename E>
+static __dpct_inline__ void fg_pack_quarter(const float * __restrict__ t, typename E::pair * a, int il) {
+#pragma unroll
+    for (int j = 0; j < 4; ++j) {
+        a[4 * il + j] = E::make(t[2 * j], t[2 * j + 1]);
+    }
+}
+
+template <typename E>
+static __dpct_inline__ void fg_stage_a(const block_iq4_xs * __restrict__ xrow, const int kb, typename E::pair * a) {
+    FG_SUPERBLOCK(block_iq4_xs);
+    // low nibbles fill the first half of the step, high nibbles the second, so the two halves
+    // land at a[0..7] and a[8..15] and no quarter loop is needed
+    const float d = (float) blk->d *
+        ((((blk->scales_l[ib / 2] >> (4 * (ib % 2))) & 0xf) | (((blk->scales_h >> (2 * ib)) & 3) << 4)) - 32);
+    const uint8_t * q4 = blk->qs + 16 * ib;
+#pragma unroll
+    for (int j = 0; j < 8; ++j) {
+        a[j]     = E::make(d * kvalues_iq4nl[q4[2 * j] & 0xf], d * kvalues_iq4nl[q4[2 * j + 1] & 0xf]);
+        a[8 + j] = E::make(d * kvalues_iq4nl[q4[2 * j] >> 4],  d * kvalues_iq4nl[q4[2 * j + 1] >> 4]);
+    }
+}
+
+template <typename E>
+static __dpct_inline__ void fg_stage_a(const block_iq3_xxs * __restrict__ xrow, const int kb, typename E::pair * a) {
+    FG_SUPERBLOCK(block_iq3_xxs);
+    const uint8_t *  q3    = blk->qs + 8 * ib;
+    const uint16_t * gas   = (const uint16_t *) (blk->qs + QK_K / 4) + 2 * ib;
+    const uint32_t   aux32 = gas[0] | (gas[1] << 16);
+    const float      d     = (float) blk->d * (0.5f + (aux32 >> 28)) * 0.5f;
+#pragma unroll
+    for (int il = 0; il < 4; ++il) {
+        const uint8_t * grid1 = (const uint8_t *) (iq3xxs_grid + q3[2 * il + 0]);
+        const uint8_t * grid2 = (const uint8_t *) (iq3xxs_grid + q3[2 * il + 1]);
+        const uint8_t   signs = ksigns_iq2xs[(aux32 >> (7 * il)) & 127];
+        float t[8];
+#pragma unroll
+        for (int j = 0; j < 4; ++j) {
+            t[j + 0] = d * grid1[j] * (signs & kmask_iq2xs[j + 0] ? -1.f : 1.f);
+            t[j + 4] = d * grid2[j] * (signs & kmask_iq2xs[j + 4] ? -1.f : 1.f);
+        }
+        fg_pack_quarter<E>(t, a, il);
+    }
+}
+
+template <typename E>
+static __dpct_inline__ void fg_stage_a(const block_iq2_xxs * __restrict__ xrow, const int kb, typename E::pair * a) {
+    FG_SUPERBLOCK(block_iq2_xxs);
+    const uint16_t * q2    = blk->qs + 4 * ib;
+    const uint8_t *  aux8  = (const uint8_t *) q2;
+    const uint32_t   aux32 = q2[2] | (q2[3] << 16);
+    const float      d     = (float) blk->d * (0.5f + (aux32 >> 28)) * 0.25f;
+#pragma unroll
+    for (int il = 0; il < 4; ++il) {
+        const uint8_t * grid  = (const uint8_t *) (iq2xxs_grid + aux8[il]);
+        const uint8_t   signs = ksigns_iq2xs[(aux32 >> (7 * il)) & 127];
+        float t[8];
+#pragma unroll
+        for (int j = 0; j < 8; ++j) {
+            t[j] = d * grid[j] * (signs & kmask_iq2xs[j] ? -1.f : 1.f);
+        }
+        fg_pack_quarter<E>(t, a, il);
+    }
+}
+
+template <typename E>
+static __dpct_inline__ void fg_stage_a(const block_iq2_xs * __restrict__ xrow, const int kb, typename E::pair * a) {
+    FG_SUPERBLOCK(block_iq2_xs);
+    const uint16_t * q2 = blk->qs + 4 * ib;
+#pragma unroll
+    for (int il = 0; il < 4; ++il) {
+        const uint8_t * grid  = (const uint8_t *) (iq2xs_grid + (q2[il] & 511));
+        const float     d     = (float) blk->d * (0.5f + ((blk->scales[ib] >> (4 * (il / 2))) & 0xf)) * 0.25f;
+        const uint8_t   signs = ksigns_iq2xs[q2[il] >> 9];
+        float t[8];
+#pragma unroll
+        for (int j = 0; j < 8; ++j) {
+            t[j] = d * grid[j] * (signs & kmask_iq2xs[j] ? -1.f : 1.f);
+        }
+        fg_pack_quarter<E>(t, a, il);
+    }
+}
+
+template <typename E>
+static __dpct_inline__ void fg_stage_a(const block_iq2_s * __restrict__ xrow, const int kb, typename E::pair * a) {
+    FG_SUPERBLOCK(block_iq2_s);
+#pragma unroll
+    for (int il = 0; il < 4; ++il) {
+        const uint8_t * grid =
+            (const uint8_t *) (iq2s_grid + (blk->qs[4 * ib + il] | ((blk->qh[ib] << (8 - 2 * il)) & 0x300)));
+        const float   d     = (float) blk->d * (0.5f + ((blk->scales[ib] >> (4 * (il / 2))) & 0xf)) * 0.25f;
+        const uint8_t signs = blk->qs[QK_K / 8 + 4 * ib + il];
+        float t[8];
+#pragma unroll
+        for (int j = 0; j < 8; ++j) {
+            t[j] = d * grid[j] * (signs & kmask_iq2xs[j] ? -1.f : 1.f);
+        }
+        fg_pack_quarter<E>(t, a, il);
+    }
+}
+
+template <typename E>
+static __dpct_inline__ void fg_stage_a(const block_iq1_s * __restrict__ xrow, const int kb, typename E::pair * a) {
+    FG_SUPERBLOCK(block_iq1_s);
+    const float delta = blk->qh[ib] & 0x8000 ? -1 - IQ1S_DELTA : -1 + IQ1S_DELTA;
+    const float d     = (float) blk->d * (2 * ((blk->qh[ib] >> 12) & 7) + 1);
+#pragma unroll
+    for (int il = 0; il < 4; ++il) {
+        uint32_t       grid32[2];
+        const int8_t * q = (const int8_t *) grid32;
+        grid32[0] = iq1s_grid_gpu[blk->qs[4 * ib + il] | (((blk->qh[ib] >> (3 * il)) & 7) << 8)];
+        grid32[1] = (grid32[0] >> 4) & 0x0f0f0f0f;
+        grid32[0] &= 0x0f0f0f0f;
+        float t[8];
+#pragma unroll
+        for (int j = 0; j < 8; ++j) {
+            t[j] = d * (q[j] + delta);
+        }
+        fg_pack_quarter<E>(t, a, il);
+    }
+}
+
+template <typename E>
+static __dpct_inline__ void fg_stage_a(const block_iq1_m * __restrict__ xrow, const int kb, typename E::pair * a) {
+    FG_SUPERBLOCK(block_iq1_m);
+    const uint16_t * sc = (const uint16_t *) blk->scales;
+    iq1m_scale_t     scale;
+    scale.u16 = (sc[0] >> 12) | ((sc[1] >> 8) & 0x00f0) | ((sc[2] >> 4) & 0x0f00) | (sc[3] & 0xf000);
+#pragma unroll
+    for (int il = 0; il < 4; ++il) {
+        const int   ib16  = 2 * ib + il / 2;
+        const float d     = (float) scale.f16 * (2 * ((sc[ib16 / 4] >> (3 * (ib16 % 4))) & 0x7) + 1);
+        const float delta = blk->qh[2 * ib + il / 2] & (0x08 << (4 * (il % 2))) ? -1 - IQ1M_DELTA : -1 + IQ1M_DELTA;
+        uint32_t       grid32[2];
+        const int8_t * q = (const int8_t *) grid32;
+        grid32[0] = iq1s_grid_gpu[blk->qs[4 * ib + il] |
+                                  (((blk->qh[2 * ib + il / 2] >> (4 * (il % 2))) & 7) << 8)];
+        grid32[1] = (grid32[0] >> 4) & 0x0f0f0f0f;
+        grid32[0] &= 0x0f0f0f0f;
+        float t[8];
+#pragma unroll
+        for (int j = 0; j < 8; ++j) {
+            t[j] = d * (q[j] + delta);
+        }
+        fg_pack_quarter<E>(t, a, il);
+    }
+}
+
+
+// q8_0 and the k-quants. Each decode takes the fields of one 32-value sub-block, so the canonical
+// layout and the reorder (SoA) layout of reorder_qw() share it and differ only in where the fields
+// live, as in dequantize.hpp.
+template <typename E>
+static __dpct_inline__ void fg_decode_q8_0(const int8_t * __restrict__ qs, const float d, typename E::pair * a) {
+#pragma unroll
+    for (int j = 0; j < QK8_0 / 2; ++j) {
+        a[j] = E::make(d * qs[2 * j], d * qs[2 * j + 1]);
+    }
+}
+
+// same unpack as get_scale_min_k4() in dequantize.hpp
+static __dpct_inline__ void fg_scale_min_k4(const int j, const uint8_t * __restrict__ q, uint8_t & d, uint8_t & m) {
+    if (j < 4) {
+        d = q[j] & 63;
+        m = q[j + 4] & 63;
+    } else {
+        d = (q[j + 4] & 0xF) | ((q[j - 4] >> 6) << 4);
+        m = (q[j + 4] >> 4) | ((q[j - 0] >> 6) << 4);
+    }
+}
+
+// q4_K sub-block ib (0..7) uses scale/min pair ib and the low (even ib) or high (odd ib) nibbles of
+// qs[32 * (ib / 2) ...], as in dequantize_row_q4_K
+template <typename E>
+static __dpct_inline__ void fg_decode_q4_K(const uint8_t * __restrict__ qs, const uint8_t * __restrict__ scales,
+                                           const sycl::half2 dm, const int ib, typename E::pair * a) {
+    uint8_t sc, mb;
+    fg_scale_min_k4(ib, scales, sc, mb);
+    const float     d     = (float) dm[0] * sc;
+    const float     m     = (float) dm[1] * mb;
+    const uint8_t * q     = qs + 32 * (ib / 2);
+    const int       shift = 4 * (ib % 2);
+#pragma unroll
+    for (int j = 0; j < 16; ++j) {
+        a[j] = E::make(d * ((q[2 * j] >> shift) & 0xF) - m, d * ((q[2 * j + 1] >> shift) & 0xF) - m);
+    }
+}
+
+// q5_K: q4_K plus one high bit per value, bit ib of qh[l]
+template <typename E>
+static __dpct_inline__ void fg_decode_q5_K(const uint8_t * __restrict__ qs, const uint8_t * __restrict__ qh,
+                                           const uint8_t * __restrict__ scales, const sycl::half2 dm, const int ib,
+                                           typename E::pair * a) {
+    uint8_t sc, mb;
+    fg_scale_min_k4(ib, scales, sc, mb);
+    const float     d     = (float) dm[0] * sc;
+    const float     m     = (float) dm[1] * mb;
+    const uint8_t * q     = qs + 32 * (ib / 2);
+    const int       shift = 4 * (ib % 2);
+#pragma unroll
+    for (int j = 0; j < 16; ++j) {
+        const int l = 2 * j;
+        a[j] = E::make(d * (((q[l] >> shift) & 0xF) | (((qh[l] >> ib) & 1) << 4)) - m,
+                       d * (((q[l + 1] >> shift) & 0xF) | (((qh[l + 1] >> ib) & 1) << 4)) - m);
+    }
+}
+
+// q6_K: sub-block ib is quarter r = ib % 4 of half h = ib / 4. The half selects ql + 64h, qh + 32h
+// and scales + 8h; the quarter selects ql + 32(r & 1), the ql nibble r / 2, the qh bit pair r and
+// scales + 2r, as in dequantize_row_q6_K. The scale changes at value 16 of the sub-block.
+template <typename E>
+static __dpct_inline__ void fg_decode_q6_K(const uint8_t * __restrict__ ql, const uint8_t * __restrict__ qh,
+                                           const int8_t * __restrict__ scales, const float d, const int ib,
+                                           typename E::pair * a) {
+    const int       h  = ib / 4;
+    const int       r  = ib % 4;
+    const uint8_t * q  = ql + 64 * h + 32 * (r & 1);
+    const uint8_t * hb = qh + 32 * h;
+    const int8_t *  sc = scales + 8 * h + 2 * r;
+#pragma unroll
+    for (int j = 0; j < 16; ++j) {
+        const int   l  = 2 * j;
+        const float dl = d * sc[j / 8];
+        const int   q0 = (((q[l] >> (4 * (r / 2))) & 0xF) | (((hb[l] >> (2 * r)) & 3) << 4)) - 32;
+        const int   q1 = (((q[l + 1] >> (4 * (r / 2))) & 0xF) | (((hb[l + 1] >> (2 * r)) & 3) << 4)) - 32;
+        a[j] = E::make(dl * q0, dl * q1);
+    }
+}
+
+template <typename E>
+static __dpct_inline__ void fg_stage_a(const block_q8_0 * __restrict__ xrow, const int kb, typename E::pair * a) {
+    fg_decode_q8_0<E>(xrow[kb].qs, (float) xrow[kb].d, a);
+}
+
+template <typename E>
+static __dpct_inline__ void fg_stage_a(const block_q4_K * __restrict__ xrow, const int kb, typename E::pair * a) {
+    FG_SUPERBLOCK(block_q4_K);
+    fg_decode_q4_K<E>(blk->qs, blk->scales, blk->dm, ib, a);
+}
+
+template <typename E>
+static __dpct_inline__ void fg_stage_a(const block_q5_K * __restrict__ xrow, const int kb, typename E::pair * a) {
+    FG_SUPERBLOCK(block_q5_K);
+    fg_decode_q5_K<E>(blk->qs, blk->qh, blk->scales, blk->dm, ib, a);
+}
+
+template <typename E>
+static __dpct_inline__ void fg_stage_a(const block_q6_K * __restrict__ xrow, const int kb, typename E::pair * a) {
+    FG_SUPERBLOCK(block_q6_K);
+    fg_decode_q6_K<E>(blk->ql, blk->qh, blk->scales, (float) blk->d, ib, a);
+}
+
+// Reorder (SoA) layout: each block field is one stream over the nblocks of the matrix (of the expert
+// slice for MUL_MAT_ID), in the order reorder_qw() writes them. Block ib holds k step kb.
+template <typename block_q_t> struct fg_soa {
+    static constexpr bool supported = false;
+};
+
+template <> struct fg_soa<block_q8_0> {
+    static constexpr bool supported = true;
+    // [qs][d]
+    template <typename E>
+    static __dpct_inline__ void stage(const uint8_t * x, const size_t nblocks, const size_t ib, const int,
+                                      typename E::pair * a) {
+        const float d = (float) ((const sycl::half *) (x + nblocks * QK8_0))[ib];
+        fg_decode_q8_0<E>((const int8_t *) x + ib * QK8_0, d, a);
+    }
+};
+
+template <> struct fg_soa<block_q4_K> {
+    static constexpr bool supported = true;
+    // [qs][scales][dm]
+    template <typename E>
+    static __dpct_inline__ void stage(const uint8_t * x, const size_t nblocks, const size_t ib, const int kb,
+                                      typename E::pair * a) {
+        const uint8_t *   scales = x + nblocks * (QK_K / 2);
+        const sycl::half2 dm     = ((const sycl::half2 *) (scales + nblocks * K_SCALE_SIZE))[ib];
+        fg_decode_q4_K<E>(x + ib * (QK_K / 2), scales + ib * K_SCALE_SIZE, dm, kb % (QK_K / 32), a);
+    }
+};
+
+template <> struct fg_soa<block_q5_K> {
+    static constexpr bool supported = true;
+    // [qs][qh][scales][dm]
+    template <typename E>
+    static __dpct_inline__ void stage(const uint8_t * x, const size_t nblocks, const size_t ib, const int kb,
+                                      typename E::pair * a) {
+        const uint8_t *   qh     = x + nblocks * (QK_K / 2);
+        const uint8_t *   scales = qh + nblocks * (QK_K / 8);
+        const sycl::half2 dm     = ((const sycl::half2 *) (scales + nblocks * K_SCALE_SIZE))[ib];
+        fg_decode_q5_K<E>(x + ib * (QK_K / 2), qh + ib * (QK_K / 8), scales + ib * K_SCALE_SIZE, dm,
+                          kb % (QK_K / 32), a);
+    }
+};
+
+template <> struct fg_soa<block_q6_K> {
+    static constexpr bool supported = true;
+    // [ql][qh][scales][d]
+    template <typename E>
+    static __dpct_inline__ void stage(const uint8_t * x, const size_t nblocks, const size_t ib, const int kb,
+                                      typename E::pair * a) {
+        const uint8_t * qh     = x + nblocks * (QK_K / 2);
+        const uint8_t * scales = qh + nblocks * (QK_K / 4);
+        const float     d      = (float) ((const sycl::half *) (scales + nblocks * (QK_K / 16)))[ib];
+        fg_decode_q6_K<E>(x + ib * (QK_K / 2), qh + ib * (QK_K / 4), (const int8_t *) scales + ib * (QK_K / 16), d,
+                          kb % (QK_K / 32), a);
+    }
+};
+
+
+// one SG_ROWS x BN output tile: B columns [b0, b0 + BN) of packed_b go to dst columns [n0, n1),
+// n1 - n0 <= BN
+template <typename S, typename block_q_t, bool reordered>
+static void fused_dequant_gemm_tile(
+    const block_q_t * __restrict__ x,
+    const typename S::tsb * __restrict__ packed_b,
+    float * __restrict__ dst,
+    const int M, const int Npad, const int K, const int ldd,
+    const int b0, const int n0, const int n1,
+    sycl::local_accessor<typename S::tsa, 1> tile_a,
+    sycl::local_accessor<float, 1> tile_c,
+    const sycl::nd_item<2> & item) {
+    if constexpr (fg_built<S::SG>()) {
+        using EA = typename S::EA;
+        using TA = typename S::ta;
+        using TB = typename S::tb;
+        const auto sg     = item.get_sub_group();
+        const int  sg_id  = sg.get_group_id()[0];
+        const int  lane   = sg.get_local_id()[0];
+        const int  m0     = item.get_group(1) * S::SG_ROWS;
+        const int  nstep  = K / FG_BK;
+        const int  a_base = sg_id * S::SG_ROWS * FG_BK;
+        const int  c_base = sg_id * S::SG_ROWS * S::BN;
+
+        mx::joint_matrix<sycl::sub_group, float, mx::use::accumulator, S::TM, S::TN> acc[S::MT][S::NT];
+#pragma unroll
+        for (int mt = 0; mt < S::MT; ++mt) {
+#pragma unroll
+            for (int nt = 0; nt < S::NT; ++nt) {
+                mx::joint_matrix_fill(sg, acc[mt][nt], 0.0f);
+            }
+        }
+
+        // lane decodes rows lane, lane + SG, ... of the sub-group's SG_ROWS
+        constexpr int         KPB     = fg_block_traits<block_q_t>::qk / FG_BK; // k steps per block
+        const size_t          bpr     = K / fg_block_traits<block_q_t>::qk;
+        const size_t          nblocks = (size_t) M * bpr;
+        size_t                row_blk[S::RPL];
+        bool                  row_ok[S::RPL];
+        typename EA::pair *   a[S::RPL];
+#pragma unroll
+        for (int r = 0; r < S::RPL; ++r) {
+            const int row = m0 + r * S::SG + lane;
+            row_ok[r]  = row < M;
+            row_blk[r] = (size_t) (row_ok[r] ? row : 0) * bpr;
+            a[r]       = (typename EA::pair *) &tile_a[a_base + (r * S::SG + lane) * FG_BK];
+        }
+
+        const auto b_ptr = sycl::address_space_cast<sycl::access::address_space::global_space,
+                                                    sycl::access::decorated::no>(packed_b);
+        const int b_stride = Npad * S::VNNI;
+
+        const int kb_begin = (sg_id * nstep) / FG_KSPLIT;
+        const int kb_end   = ((sg_id + 1) * nstep) / FG_KSPLIT;
+        for (int kb = kb_begin; kb < kb_end; ++kb) {
+#pragma unroll
+            for (int r = 0; r < S::RPL; ++r) {
+                if (row_ok[r]) {
+                    if constexpr (reordered) {
+                        fg_soa<block_q_t>::template stage<EA>((const uint8_t *) x, nblocks, row_blk[r] + kb / KPB, kb,
+                                                              a[r]);
+                    } else {
+                        fg_stage_a<EA>(x + row_blk[r], kb, a[r]);
+                    }
+                } else {
+#pragma unroll
+                    for (int j = 0; j < FG_BK / 2; ++j) {
+                        a[r][j] = EA::make(0.0f, 0.0f);
+                    }
+                }
+            }
+            sycl::group_barrier(sg);
+
+#pragma unroll
+            for (int kt = 0; kt < FG_BK / S::TK; ++kt) {
+                const int kq0 = (kb * FG_BK + kt * S::TK) / S::VNNI;
+                mx::joint_matrix<sycl::sub_group, TB, mx::use::b, S::TK, S::TN, S::b_layout> sub_b[S::NT];
+#pragma unroll
+                for (int nt = 0; nt < S::NT; ++nt) {
+                    mx::joint_matrix_load(sg, sub_b[nt], b_ptr + (size_t) kq0 * b_stride + (b0 + nt * S::TN) * S::VNNI, b_stride);
+                }
+#pragma unroll
+                for (int mt = 0; mt < S::MT; ++mt) {
+                    mx::joint_matrix<sycl::sub_group, TA, mx::use::a, S::TM, S::TK, mx::layout::row_major> sub_a;
+                    mx::joint_matrix_load(sg, sub_a,
+                        tile_a.template get_multi_ptr<sycl::access::decorated::no>() + a_base + (mt * S::TM) * FG_BK + kt * S::TK,
+                        FG_BK);
+#pragma unroll
+                    for (int nt = 0; nt < S::NT; ++nt) {
+                        mx::joint_matrix_mad(sg, acc[mt][nt], sub_a, sub_b[nt], acc[mt][nt]);
+                    }
+                }
+            }
+            // the next step overwrites tile_a
+            sycl::group_barrier(sg);
+        }
+
+#pragma unroll
+        for (int mt = 0; mt < S::MT; ++mt) {
+#pragma unroll
+            for (int nt = 0; nt < S::NT; ++nt) {
+                mx::joint_matrix_store(sg, acc[mt][nt],
+                    tile_c.template get_multi_ptr<sycl::access::decorated::no>() + c_base + (mt * S::TM) * S::BN + nt * S::TN,
+                    S::BN, mx::layout::row_major);
+            }
+        }
+        sycl::group_barrier(item.get_group());
+
+        // sum the K splits; consecutive lanes write consecutive rows of one dst column
+        for (int idx = item.get_local_linear_id(); idx < S::SG_ROWS * S::BN; idx += S::WG_SIZE) {
+            const int r = idx % S::SG_ROWS;
+            const int c = idx / S::SG_ROWS;
+            const int m = m0 + r;
+            const int n = n0 + c;
+            if (m < M && n < n1) {
+                float sum = 0.0f;
+#pragma unroll
+                for (int s = 0; s < FG_KSPLIT; ++s) {
+                    sum += tile_c[s * S::SG_ROWS * S::BN + r * S::BN + c];
+                }
+                dst[(size_t) n * ldd + m] = sum;
+            }
+        }
+    }
+}
+
+template <typename S, typename block_q_t>
+static void fused_dequant_gemm_launch(const void * src0, const typename S::tsb * packed, float * dst, const int M,
+                                      const int N, const int Npad, const int K, const int ldd,
+                                      const int64_t groups_n, const int64_t groups_m, dpct::queue_ptr stream) {
+    stream->submit([&](sycl::handler & cgh) {
+        sycl::local_accessor<typename S::tsa, 1> tile_a(FG_KSPLIT * S::SG_ROWS * FG_BK, cgh);
+        sycl::local_accessor<float, 1>          tile_c(FG_KSPLIT * S::SG_ROWS * S::BN, cgh);
+        cgh.parallel_for(
+            sycl::nd_range<2>(sycl::range<2>(groups_n, groups_m * S::WG_SIZE), sycl::range<2>(1, S::WG_SIZE)),
+            [=](sycl::nd_item<2> item) [[sycl::reqd_sub_group_size(S::SG)]] {
+                const int n0 = item.get_group(0) * S::BN;
+                fused_dequant_gemm_tile<S, block_q_t, false>((const block_q_t *) src0, packed, dst, M, Npad, K, ldd,
+                                                      n0, n0, N, tile_a, tile_c, item);
+            });
+    });
+}
+
+// grouped: work-group (t, mt) is tile t of the schedule; its B columns sit at t * BN
+template <typename S, typename block_q_t, bool reordered>
+static void grouped_dequant_gemm_launch(const char * src0_dd, const size_t expert_stride,
+                                        const ggml_sycl_gg_tile * tiles_ptr, const typename S::tsb * packed, float * dst,
+                                        const int M, const int Npad, const int K, const int64_t n_tiles,
+                                        const int64_t groups_m, dpct::queue_ptr stream) {
+    stream->submit([&](sycl::handler & cgh) {
+        sycl::local_accessor<typename S::tsa, 1> tile_a(FG_KSPLIT * S::SG_ROWS * FG_BK, cgh);
+        sycl::local_accessor<float, 1>          tile_c(FG_KSPLIT * S::SG_ROWS * S::BN, cgh);
+        cgh.parallel_for(
+            sycl::nd_range<2>(sycl::range<2>(n_tiles, groups_m * S::WG_SIZE), sycl::range<2>(1, S::WG_SIZE)),
+            [=](sycl::nd_item<2> item) [[sycl::reqd_sub_group_size(S::SG)]] {
+                const int               t    = item.get_group(0);
+                const ggml_sycl_gg_tile tile = tiles_ptr[t];
+                const block_q_t *       x    = (const block_q_t *) (src0_dd + (size_t) tile.expert * expert_stride);
+                fused_dequant_gemm_tile<S, block_q_t, reordered>(x, packed, dst, M, Npad, K, M, t * S::BN, tile.n0,
+                                                                 tile.n1, tile_a, tile_c, item);
+            });
+    });
+}
+
+// src1 f32 rows -> packed [K/V][n_tiles*BN][V], tile t holds its rows [n0, n1) at columns t*BN..,
+// zero past n1. The column runs fastest so a sub-group writes one contiguous run.
+template <typename S>
+static void grouped_gemm_pack_b(const float * y, typename S::tsb * packed, const ggml_sycl_gg_tile * tiles, int Npad,
+                                int K, dpct::queue_ptr stream) {
+    using E        = typename S::EB;
+    constexpr int V = S::VNNI;
+    const int kqs   = K / V;
+    stream->parallel_for(sycl::range<1>((size_t) Npad * kqs), [=](sycl::id<1> id) {
+        const size_t idx = id[0];
+        const int    kq  = idx / Npad;
+        const int    n   = idx - (size_t) kq * Npad;
+        const ggml_sycl_gg_tile tile = tiles[n / S::BN];
+        const int    row = tile.n0 + n % S::BN;
+        // one guarded load run per work-item, as a per-element select costs ~1.5% prefill
+        typename S::tsb vals[V] = {};
+        if (row < tile.n1) {
+            const float * src = y + (size_t) row * K + V * kq;
+#pragma unroll
+            for (int v = 0; v < V; ++v) {
+                vals[v] = E::cvt(src[v]);
+            }
+        }
+        typename S::tsb * out = packed + ((size_t) kq * Npad + n) * V;
+#pragma unroll
+        for (int v = 0; v < V; ++v) {
+            out[v] = vals[v];
+        }
+    });
+}
+
+// q8_0 and the k-quants take only the grouped path. The plain kernel decodes A again for every BN columns
+// of a dense batch, and for these formats that costs more than the one dequantization of the library GEMM.
+template <typename T> static constexpr bool fg_plain_ok() {
+    return !std::is_same_v<T, block_q8_0> && !std::is_same_v<T, block_q4_K> && !std::is_same_v<T, block_q5_K> &&
+           !std::is_same_v<T, block_q6_K>;
+}
+
+template <typename T, bool R> struct fg_tag {
+    using type                     = T;
+    static constexpr bool reordered = R;
+};
+
+// calls f(fg_tag<block_q_t, reordered>{}) for the weight format and layout; false if it has no A stage
+template <typename T, typename F> static bool fg_visit_layout(bool reordered, F && f) {
+    if (!reordered) {
+        f(fg_tag<T, false>{});
+        return true;
+    }
+    if constexpr (fg_soa<T>::supported) {
+        f(fg_tag<T, true>{});
+        return true;
+    }
+    return false;
+}
+
+template <typename F> static bool fg_visit_type(ggml_type type, bool reordered, F && f) {
+    switch (type) {
+        case GGML_TYPE_IQ4_NL:  return fg_visit_layout<block_iq4_nl>(reordered, f);
+        case GGML_TYPE_IQ3_S:   return fg_visit_layout<block_iq3_s>(reordered, f);
+        case GGML_TYPE_IQ4_XS:  return fg_visit_layout<block_iq4_xs>(reordered, f);
+        case GGML_TYPE_IQ3_XXS: return fg_visit_layout<block_iq3_xxs>(reordered, f);
+        case GGML_TYPE_IQ2_XXS: return fg_visit_layout<block_iq2_xxs>(reordered, f);
+        case GGML_TYPE_IQ2_XS:  return fg_visit_layout<block_iq2_xs>(reordered, f);
+        case GGML_TYPE_IQ2_S:   return fg_visit_layout<block_iq2_s>(reordered, f);
+        case GGML_TYPE_IQ1_S:   return fg_visit_layout<block_iq1_s>(reordered, f);
+        case GGML_TYPE_IQ1_M:   return fg_visit_layout<block_iq1_m>(reordered, f);
+        case GGML_TYPE_Q8_0:    return fg_visit_layout<block_q8_0>(reordered, f);
+        case GGML_TYPE_Q4_K:    return fg_visit_layout<block_q4_K>(reordered, f);
+        case GGML_TYPE_Q5_K:    return fg_visit_layout<block_q5_K>(reordered, f);
+        case GGML_TYPE_Q6_K:    return fg_visit_layout<block_q6_K>(reordered, f);
+        default:                return false;
+    }
+}
+
+template <typename S>
+static bool fg_fused_run(ggml_type src0_type, const void * src0, const void * src1, ggml_type src1_type,
+                         float * dst, int64_t M, int64_t N, int64_t K, int64_t ldd, ggml_sycl_pool & pool,
+                         dpct::queue_ptr stream) {
+    const int64_t groups_n = (N + S::BN - 1) / S::BN;
+    const int64_t groups_m = (M + S::SG_ROWS - 1) / S::SG_ROWS;
+    const int     Npad     = (int) (groups_n * S::BN);
+
+    // src1 is read in its own type: one pass, converted in registers only if B differs
+    ggml_sycl_pool_alloc<typename S::tsb> packed_b(pool, (size_t) K * Npad);
+    if (src1_type == GGML_TYPE_F16) {
+        fused_gemm_pack_b<typename S::EB>((const sycl::half *) src1, packed_b.get(), (int) N, Npad, (int) K, stream);
+    } else if (src1_type == GGML_TYPE_BF16) {
+        fused_gemm_pack_b<typename S::EB>((const fg_bf16 *) src1, packed_b.get(), (int) N, Npad, (int) K, stream);
+    } else {
+        fused_gemm_pack_b<typename S::EB>((const float *) src1, packed_b.get(), (int) N, Npad, (int) K, stream);
+    }
+
+    const typename S::tsb * packed = packed_b.get();
+    return fg_visit_type(src0_type, false, [&](auto tag) {
+        using block_q_t = typename decltype(tag)::type;
+        if constexpr (fg_plain_ok<block_q_t>()) {
+            fused_dequant_gemm_launch<S, block_q_t>(src0, packed, dst, (int) M, (int) N, Npad, (int) K, (int) ldd,
+                                                    groups_n, groups_m, stream);
+        }
+    });
+}
+
+bool ggml_sycl_fused_dequant_gemm(ggml_type src0_type, const void * src0, const void * src1, ggml_type src1_type,
+                                  int32_t src1_prec, float * dst, int64_t M, int64_t N, int64_t K, int64_t ldd,
+                                  ggml_sycl_pool & pool, dpct::queue_ptr stream) {
+    // every BN columns dequantize A again, so wide N is left to the library GEMM
+    bool plain_ok = false;
+    fg_visit_type(src0_type, false, [&](auto tag) { plain_ok = fg_plain_ok<typename decltype(tag)::type>(); });
+    if (g_ggml_sycl_dynamic_precision == GGML_SYCL_DYNAMIC_PRECISION_F32 ||
+        !ggml_sycl_xmx_gather_type_enabled(src0_type) || !plain_ok) {
+        return false;
+    }
+    if (src1_type != GGML_TYPE_F32 && src1_type != GGML_TYPE_F16 && src1_type != GGML_TYPE_BF16) {
+        return false;
+    }
+    if (!ggml_sycl_fused_dequant_gemm_shape_ok(src0_type, M, N, K, ldd)) {
+        return false;
+    }
+    const int combo = fg_pick_combo(stream, src1_type, src1_prec);
+    if (combo < 0) {
+        return false;
+    }
+    bool launched = false;
+    fg_visit_combo(combo, [&](auto s) {
+        launched = fg_fused_run<decltype(s)>(src0_type, src0, src1, src1_type, dst, M, N, K, ldd, pool, stream);
+    });
+    return launched;
+}
+
+template <typename S>
+static bool fg_grouped_run(ggml_type src0_type, bool reordered, const void * src0_base, size_t expert_stride,
+                           const float * src1, float * dst, const int64_t * expert_row_offsets, int64_t n_as, int64_t M,
+                           int64_t K, std::vector<ggml_sycl_gg_tile> & tiles, ggml_sycl_pool & pool,
+                           dpct::queue_ptr stream) {
+    // the host knows every slice, so it lays out the work-groups: no search on the device
+    tiles.clear();
+    for (int64_t e = 0; e < n_as; ++e) {
+        const int64_t end = expert_row_offsets[e + 1];
+        for (int64_t n0 = expert_row_offsets[e]; n0 < end; n0 += S::BN) {
+            tiles.push_back({ (int32_t) e, (int32_t) n0, (int32_t) std::min<int64_t>(n0 + S::BN, end) });
+        }
+    }
+    const int64_t n_tiles  = tiles.size();
+    const int64_t groups_m = (M + S::SG_ROWS - 1) / S::SG_ROWS;
+    const int     Npad     = (int) (n_tiles * S::BN);
+
+    const int64_t max_tiles = grouped_gemm_max_tiles(expert_row_offsets[n_as], n_as, S::BN);
+    GGML_ASSERT(n_tiles <= max_tiles);
+    ggml_sycl_pool_alloc<ggml_sycl_gg_tile> tiles_dev(pool, max_tiles);
+    SYCL_CHECK(CHECK_TRY_ERROR(stream->memcpy(tiles_dev.get(), tiles.data(), n_tiles * sizeof(ggml_sycl_gg_tile))));
+
+    ggml_sycl_pool_alloc<typename S::tsb> packed_b(pool, (size_t) K * max_tiles * S::BN);
+    grouped_gemm_pack_b<S>(src1, packed_b.get(), tiles_dev.get(), Npad, (int) K, stream);
+
+    const typename S::tsb *   packed    = packed_b.get();
+    const ggml_sycl_gg_tile * tiles_ptr = tiles_dev.get();
+    const char *              src0_dd   = (const char *) src0_base;
+    return fg_visit_type(src0_type, reordered, [&](auto tag) {
+        using T = decltype(tag);
+        grouped_dequant_gemm_launch<S, typename T::type, T::reordered>(src0_dd, expert_stride, tiles_ptr, packed, dst,
+                                                                       (int) M, Npad, (int) K, n_tiles, groups_m,
+                                                                       stream);
+    });
+}
+
+bool ggml_sycl_grouped_dequant_gemm(ggml_type src0_type, bool reordered, const void * src0_base, size_t expert_stride,
+                                    const float * src1, int32_t src1_prec, float * dst,
+                                    const int64_t * expert_row_offsets, int64_t n_as, int64_t M, int64_t K,
+                                    int64_t total_rows, std::vector<ggml_sycl_gg_tile> & tiles,
+                                    ggml_sycl_pool & pool, dpct::queue_ptr stream) {
+    int64_t n_active = 0;
+    for (int64_t e = 0; e < n_as; ++e) {
+        n_active += expert_row_offsets[e + 1] > expert_row_offsets[e];
+    }
+    if (g_ggml_sycl_dynamic_precision == GGML_SYCL_DYNAMIC_PRECISION_F32 ||
+        !ggml_sycl_xmx_gather_type_enabled(src0_type) || !fg_visit_type(src0_type, reordered, [](auto) {})) {
+        return false;
+    }
+    if (!ggml_sycl_grouped_dequant_gemm_shape_ok(src0_type, M, K, total_rows, n_active)) {
+        return false;
+    }
+    const int combo = fg_pick_combo(stream, GGML_TYPE_F32, src1_prec);
+    if (combo < 0) {
+        return false;
+    }
+    bool launched = false;
+    fg_visit_combo(combo, [&](auto s) {
+        launched = fg_grouped_run<decltype(s)>(src0_type, reordered, src0_base, expert_stride, src1, dst,
+                                               expert_row_offsets, n_as, M, K, tiles, pool, stream);
+    });
+    return launched;
+}
diff --git a/ggml/src/ggml-sycl/fused-gemm.hpp b/ggml/src/ggml-sycl/fused-gemm.hpp
new file mode 100644
index 000000000..298989204
--- /dev/null
+++ b/ggml/src/ggml-sycl/fused-gemm.hpp
@@ -0,0 +1,88 @@
+#ifndef GGML_SYCL_FUSED_GEMM_HPP
+#define GGML_SYCL_FUSED_GEMM_HPP
+
+#include "common.hpp"
+
+
+// Shape and type gates for the kernels below. Device capability is separate: it needs a queue to ask.
+static constexpr int GGML_SYCL_FG_MAX_N = 64; // widest N taken; each shape covers it in BN-wide tiles
+
+// weight formats the fused A stage decodes; K must cover whole stored blocks
+constexpr bool ggml_sycl_fused_dequant_gemm_type_ok(ggml_type src0_type, int64_t K) {
+    // iq4_nl and q8_0 store 32 values per block; every other format here is a 256-value superblock
+    // that the A stage walks in steps of 32, so K must cover whole superblocks.
+    if (src0_type == GGML_TYPE_IQ4_NL || src0_type == GGML_TYPE_Q8_0) {
+        return K % 32 == 0;
+    }
+    const bool superblock =
+           src0_type == GGML_TYPE_Q4_K ||
+           src0_type == GGML_TYPE_Q5_K ||
+           src0_type == GGML_TYPE_Q6_K ||
+           src0_type == GGML_TYPE_IQ3_S ||
+           src0_type == GGML_TYPE_IQ4_XS ||
+           src0_type == GGML_TYPE_IQ3_XXS ||
+           src0_type == GGML_TYPE_IQ2_XXS ||
+           src0_type == GGML_TYPE_IQ2_XS ||
+           src0_type == GGML_TYPE_IQ2_S ||
+           src0_type == GGML_TYPE_IQ1_S ||
+           src0_type == GGML_TYPE_IQ1_M;
+    return superblock && QK_K == 256 && K % QK_K == 0;
+}
+
+constexpr bool ggml_sycl_fused_dequant_gemm_shape_ok(ggml_type src0_type, int64_t M, int64_t N, int64_t K,
+                                                     int64_t ldd) {
+    return ggml_sycl_fused_dequant_gemm_type_ok(src0_type, K) && M > 0 && N > 0 && K > 0 &&
+           N <= GGML_SYCL_FG_MAX_N &&
+           M <= INT32_MAX && N <= INT32_MAX && K <= INT32_MAX && ldd <= INT32_MAX;
+}
+
+// grouped variant: the per-expert fused kernel is only worth it while each expert is narrow,
+// so wider average slices are left to the per-expert library GEMM loop
+constexpr bool ggml_sycl_grouped_dequant_gemm_shape_ok(ggml_type src0_type, int64_t M, int64_t K,
+                                                       int64_t total_rows, int64_t n_active) {
+    return ggml_sycl_fused_dequant_gemm_shape_ok(src0_type, M, 1, K, M) && total_rows > 0 &&
+           total_rows <= INT32_MAX && total_rows <= n_active * GGML_SYCL_FG_MAX_N;
+}
+
+// Runtime type gate, kept out of the constexpr predicates above so those stay pure.
+inline bool ggml_sycl_xmx_gather_type_enabled(ggml_type src0_type) {
+    switch (src0_type) {
+        case GGML_TYPE_IQ4_NL:  return (g_ggml_sycl_xmx_gather_types & GGML_SYCL_XMX_GATHER_IQ4_NL  ) != 0;
+        case GGML_TYPE_IQ3_S:   return (g_ggml_sycl_xmx_gather_types & GGML_SYCL_XMX_GATHER_IQ3_S   ) != 0;
+        case GGML_TYPE_IQ4_XS:  return (g_ggml_sycl_xmx_gather_types & GGML_SYCL_XMX_GATHER_IQ4_XS  ) != 0;
+        case GGML_TYPE_IQ3_XXS: return (g_ggml_sycl_xmx_gather_types & GGML_SYCL_XMX_GATHER_IQ3_XXS ) != 0;
+        case GGML_TYPE_IQ2_XXS: return (g_ggml_sycl_xmx_gather_types & GGML_SYCL_XMX_GATHER_IQ2_XXS ) != 0;
+        case GGML_TYPE_IQ2_XS:  return (g_ggml_sycl_xmx_gather_types & GGML_SYCL_XMX_GATHER_IQ2_XS  ) != 0;
+        case GGML_TYPE_IQ2_S:   return (g_ggml_sycl_xmx_gather_types & GGML_SYCL_XMX_GATHER_IQ2_S   ) != 0;
+        case GGML_TYPE_IQ1_S:   return (g_ggml_sycl_xmx_gather_types & GGML_SYCL_XMX_GATHER_IQ1_S   ) != 0;
+        case GGML_TYPE_IQ1_M:   return (g_ggml_sycl_xmx_gather_types & GGML_SYCL_XMX_GATHER_IQ1_M   ) != 0;
+        case GGML_TYPE_Q8_0:    return (g_ggml_sycl_xmx_gather_types & GGML_SYCL_XMX_GATHER_Q8_0    ) != 0;
+        case GGML_TYPE_Q4_K:    return (g_ggml_sycl_xmx_gather_types & GGML_SYCL_XMX_GATHER_Q4_K    ) != 0;
+        case GGML_TYPE_Q5_K:    return (g_ggml_sycl_xmx_gather_types & GGML_SYCL_XMX_GATHER_Q5_K    ) != 0;
+        case GGML_TYPE_Q6_K:    return (g_ggml_sycl_xmx_gather_types & GGML_SYCL_XMX_GATHER_Q6_K    ) != 0;
+        default:          return false;
+    }
+}
+
+// dst[n*ldd + m] = sum_k dequant(src0)[m*K + k] * src1[n*K + k], src1 is F32, F16 or BF16.
+// The XMX combination is picked per call from the src1 type and its precision request src1_prec
+// (op_params[3], [TAG_GGML_PREC]); the accumulator is f32, which meets any request.
+// q8_0 and the k-quants are not handled here, only in the grouped path below.
+// Returns false when the case is not handled (type, device, precision, or shape).
+bool ggml_sycl_fused_dequant_gemm(ggml_type src0_type, const void * src0, const void * src1, ggml_type src1_type,
+                                  int32_t src1_prec, float * dst, int64_t M, int64_t N, int64_t K, int64_t ldd,
+                                  ggml_sycl_pool & pool, dpct::queue_ptr stream);
+
+// One launch for every expert of a MUL_MAT_ID: rows of src1/dst are grouped by expert, expert e
+// owns rows [expert_row_offsets[e], expert_row_offsets[e+1]) and reads its weights at
+// src0_base + e*expert_stride. tiles is host scratch that must stay alive until the queue drains.
+// reordered: each expert slice is in the reorder (SoA) layout of reorder_qw().
+// dst[n*M + m] = sum_k dequant(src0_e)[m*K + k] * src1[n*K + k]
+// Returns false when the case is not handled (type, layout, device, precision, or shape).
+bool ggml_sycl_grouped_dequant_gemm(ggml_type src0_type, bool reordered, const void * src0_base, size_t expert_stride,
+                                    const float * src1, int32_t src1_prec, float * dst,
+                                    const int64_t * expert_row_offsets, int64_t n_as, int64_t M, int64_t K,
+                                    int64_t total_rows, std::vector<ggml_sycl_gg_tile> & tiles,
+                                    ggml_sycl_pool & pool, dpct::queue_ptr stream);
+
+#endif // GGML_SYCL_FUSED_GEMM_HPP
diff --git a/ggml/src/ggml-sycl/ggml-sycl.cpp b/ggml/src/ggml-sycl/ggml-sycl.cpp
index b4d200526..ebca92886 100644
--- a/ggml/src/ggml-sycl/ggml-sycl.cpp
+++ b/ggml/src/ggml-sycl/ggml-sycl.cpp
@@ -14,6 +14,7 @@
 #include <array>
 #include <assert.h>
 #include <atomic>
+#include <cctype>
 #include <cinttypes>
 #include <cstddef>
 #include <cstdint>
@@ -60,6 +61,7 @@
 #include "ggml-sycl/common.hpp"
 #include "ggml-sycl/element_wise.hpp"
 #include "ggml-sycl/fwht.hpp"
+#include "ggml-sycl/fused-gemm.hpp"
 #include "ggml-sycl/gemm.hpp"
 #include "ggml-sycl/getrows.hpp"
 #include "ggml-sycl/mem.hpp"
@@ -105,6 +107,30 @@ int g_ggml_sycl_enable_fusion = 1;
 int g_ggml_sycl_enable_esimd = 1;
 int g_ggml_sycl_mmvq_wide = 1;
 int g_ggml_sycl_prioritize_dmmv = 0;
+int g_ggml_sycl_xmx_gather_types = GGML_SYCL_XMX_GATHER_TYPES_DEFAULT;
+int g_ggml_sycl_xmx_gather_shapes = GGML_SYCL_XMX_GATHER_SHAPES_DEFAULT;
+int g_ggml_sycl_dynamic_precision = GGML_SYCL_DYNAMIC_PRECISION_DEFAULT;
+int g_ggml_sycl_dynamic_required_precision = GGML_SYCL_DYNAMIC_PRECISION_F32;
+static const char * ggml_sycl_dynamic_precision_names[] = { "F16", "BF16", "TF32", "F32" };
+
+// value of a GGML_SYCL_DYNAMIC_PRECISION-style variable; def if unset or invalid
+static int ggml_sycl_get_env_precision(const char * name, int def) {
+    const char * env = getenv(name);
+    if (!env) {
+        return def;
+    }
+    std::string mode(env);
+    for (char & c : mode) {
+        c = (char) std::toupper((unsigned char) c);
+    }
+    for (int i = GGML_SYCL_DYNAMIC_PRECISION_F16; i <= GGML_SYCL_DYNAMIC_PRECISION_F32; i++) {
+        if (mode == ggml_sycl_dynamic_precision_names[i]) {
+            return i;
+        }
+    }
+    GGML_LOG_WARN("%s: unknown %s=%s, using %s\n", __func__, name, env, ggml_sycl_dynamic_precision_names[def]);
+    return def;
+}
 int g_ggml_sycl_use_async_mem_op = 0;
 int g_ggml_sycl_use_async_mem_op_requested = 1;
 int g_ggml_sycl_use_level_zero_api = 0;
@@ -401,6 +427,12 @@ static void ggml_check_sycl() try {
         g_ggml_sycl_enable_esimd = ggml_sycl_get_env("GGML_SYCL_ENABLE_ESIMD", 1);
         g_ggml_sycl_mmvq_wide = ggml_sycl_get_env("GGML_SYCL_MMVQ_WIDE", 1);
         g_ggml_sycl_prioritize_dmmv = ggml_sycl_get_env("GGML_SYCL_PRIORITIZE_DMMV", 0);
+        g_ggml_sycl_xmx_gather_types = ggml_sycl_get_env("GGML_SYCL_XMX_GATHER_TYPES", GGML_SYCL_XMX_GATHER_TYPES_DEFAULT);
+        g_ggml_sycl_xmx_gather_shapes = ggml_sycl_get_env("GGML_SYCL_XMX_GATHER_SHAPES", GGML_SYCL_XMX_GATHER_SHAPES_DEFAULT);
+        g_ggml_sycl_dynamic_precision =
+            ggml_sycl_get_env_precision("GGML_SYCL_DYNAMIC_PRECISION", GGML_SYCL_DYNAMIC_PRECISION_DEFAULT);
+        g_ggml_sycl_dynamic_required_precision =
+            ggml_sycl_get_env_precision("GGML_SYCL_DYNAMIC_REQUIRED_PRECISION", GGML_SYCL_DYNAMIC_PRECISION_F32);

 #ifdef GGML_SYCL_SUPPORT_LEVEL_ZERO_API
         g_ggml_sycl_use_level_zero_api = ggml_sycl_get_env("GGML_SYCL_USE_LEVEL_ZERO_API", 1);
@@ -509,6 +541,12 @@ static void ggml_check_sycl() try {
 #endif

         GGML_LOG_INFO("  GGML_SYCL_ENABLE_OPT: %d\n", g_ggml_sycl_enable_optimize);
+        GGML_LOG_INFO("  GGML_SYCL_XMX_GATHER_TYPES: %d\n", g_ggml_sycl_xmx_gather_types);
+        GGML_LOG_INFO("  GGML_SYCL_XMX_GATHER_SHAPES: %d\n", g_ggml_sycl_xmx_gather_shapes);
+        GGML_LOG_INFO("  GGML_SYCL_DYNAMIC_PRECISION: %s\n",
+                      ggml_sycl_dynamic_precision_names[g_ggml_sycl_dynamic_precision]);
+        GGML_LOG_INFO("  GGML_SYCL_DYNAMIC_REQUIRED_PRECISION: %s\n",
+                      ggml_sycl_dynamic_precision_names[g_ggml_sycl_dynamic_required_precision]);

 #if defined(GGML_SYCL_SUPPORT_VMM)
         GGML_LOG_INFO("  GGML_SYCL_ENABLE_VMM: %d\n", g_ggml_sycl_enable_vmm);
@@ -3027,22 +3065,18 @@ inline void ggml_sycl_op_mul_mat_sycl(
     }
 #endif

+    // dequantize inside the GEMM instead of writing the f16 weights out and reading them back; src1
+    // goes in its own type, so there is no separate conversion pass
+    if (ggml_is_quantized(src0->type) && ggml_is_contiguous(src0) && row_diff == src0->ne[1] &&
+        ggml_sycl_fused_dequant_gemm(src0->type, src0_dd_i, src1_ddf_i, src1->type, ggml_sycl_src1_prec(dst), dst_dd_i,
+                                     row_diff, src1_ncols, ne10, ldc, ctx.pool(), stream)) {
+        return;
+    }
+
+    // the f16 route converts src1 to f16 [TAG_GGML_PREC]
+    use_fp16 = use_fp16 && ggml_sycl_src1_f16_ok(dst);
     if ((src0->type == GGML_TYPE_F16 || ggml_is_quantized(src0->type)) && use_fp16 && ggml_is_contiguous(src0) &&
         row_diff == src0->ne[1] && dst->op_params[0] == GGML_PREC_DEFAULT) {
-        ggml_sycl_pool_alloc<sycl::half> src0_as_f16(ctx.pool());
-        if (src0->type != GGML_TYPE_F16) {
-            scope_op_debug_print scope_dbg_print(__func__, "/to_fp16_sycl", dst, /*num_src=*/2,
-                                                 " : converting src0 to fp16");
-            const to_fp16_sycl_t to_fp16_sycl = ggml_get_to_fp16_sycl(src0->type, dst);
-            GGML_ASSERT(to_fp16_sycl != nullptr);
-            size_t ne = row_diff*ne00;
-            src0_as_f16.alloc(ne);
-            to_fp16_sycl(src0_dd_i, src0_as_f16.get(), ne, stream);
-        }
-        const sycl::half *src0_ptr = src0->type == GGML_TYPE_F16
-                                         ? (const sycl::half *)src0_dd_i
-                                         : src0_as_f16.get();
-
         ggml_sycl_pool_alloc<sycl::half> src1_as_f16(ctx.pool());
         if (src1->type != GGML_TYPE_F16) {
             scope_op_debug_print scope_dbg_print(__func__, "/to_fp16_sycl", dst, /*num_src=*/2,
@@ -3057,6 +3091,20 @@ inline void ggml_sycl_op_mul_mat_sycl(
                 ? (const sycl::half *)src1->data + src1_padded_row_size
                                          : src1_as_f16.get();

+        ggml_sycl_pool_alloc<sycl::half> src0_as_f16(ctx.pool());
+        if (src0->type != GGML_TYPE_F16) {
+            scope_op_debug_print scope_dbg_print(__func__, "/to_fp16_sycl", dst, /*num_src=*/2,
+                                                 " : converting src0 to fp16");
+            const to_fp16_sycl_t to_fp16_sycl = ggml_get_to_fp16_sycl(src0->type, dst);
+            GGML_ASSERT(to_fp16_sycl != nullptr);
+            size_t ne = row_diff*ne00;
+            src0_as_f16.alloc(ne);
+            to_fp16_sycl(src0_dd_i, src0_as_f16.get(), ne, stream);
+        }
+        const sycl::half *src0_ptr = src0->type == GGML_TYPE_F16
+                                         ? (const sycl::half *)src0_dd_i
+                                         : src0_as_f16.get();
+
 #if GGML_SYCL_DNNL
         if (g_ggml_sycl_enable_dnn && ggml_sycl_dnnl_has_optimized_gemm(ggml_sycl_get_device())) {
                 DnnlGemmWrapper::row_gemm(ctx,row_diff, src1_ncols , ne10, src0_ptr,
@@ -4859,6 +4907,10 @@ static void ggml_sycl_mul_mat(ggml_backend_sycl_context & ctx, const ggml_tensor

     // check data types and tensor shapes for custom matrix multiplication kernels:
     bool use_dequantize_mul_mat_vec = can_use_dequantize_mul_mat_vec(src0, src1, dst);
+#ifdef GGML_SYCL_F16
+    // dmmv may convert src1 to f16 in this build [TAG_GGML_PREC]
+    use_dequantize_mul_mat_vec = use_dequantize_mul_mat_vec && ggml_sycl_src1_f16_ok(dst);
+#endif

     bool use_mul_mat_vec_q = can_use_mul_mat_vec_q(src0, src1, dst);

@@ -5272,7 +5324,9 @@ static void ggml_sycl_mul_mat_id(ggml_backend_sycl_context & ctx,
     SYCL_CHECK(CHECK_TRY_ERROR(
         stream->memcpy(ids_host.data(), ids_dev, ggml_nbytes(ids))));

-    // also ensures ctx.mmid_row_mapping_host is drained before we use it again
+    // also ensures ctx.mmid_row_mapping_host and ctx.mmid_tile_schedule_host are drained before we
+    // refill them: the grouped GEMM enqueues an async copy out of the tile schedule, so removing
+    // this wait would let the next node overwrite a buffer the device is still reading
     SYCL_CHECK(CHECK_TRY_ERROR(stream->wait()));

     ggml_tensor src0_row = *src0;
@@ -5363,7 +5417,25 @@ static void ggml_sycl_mul_mat_id(ggml_backend_sycl_context & ctx,
             });
         }

-        for (int64_t i02 = 0; i02 < n_as; i02++) {
+        bool grouped = false;
+        if (ggml_is_contiguous(src0) && src1->type == GGML_TYPE_F32 &&
+            dst->type == GGML_TYPE_F32 && nb11 == sizeof(float)*ne10 && nb1 == sizeof(float)*ne0) {
+            // the grouped GEMM reads the reorder (SoA) layout faster, and the first decode step installs it
+            // anyway: install it here already, so prefill does not depend on whether a decode ran before
+            if (g_ggml_sycl_dynamic_precision != GGML_SYCL_DYNAMIC_PRECISION_F32 &&
+                ggml_sycl_xmx_gather_type_enabled(src0->type)) {
+                opt_for_reorder_id(&ctx, src0);
+            }
+            const bool src0_reordered =
+                src0->extra && ((const ggml_tensor_extra_gpu *) src0->extra)->optimized_feature.reorder;
+            grouped = ggml_sycl_grouped_dequant_gemm(src0->type, src0_reordered, src0_original, nb02,
+                                                     (const float *) src1_contiguous.get(), ggml_sycl_src1_prec(dst),
+                                                     (float *) dst_contiguous.get(),
+                                                     expert_row_offsets.data(), n_as, ne01, ne10, n_routed_rows,
+                                                     ctx.mmid_tile_schedule_host, ctx.pool(), stream);
+        }
+
+        for (int64_t i02 = 0; i02 < n_as && !grouped; i02++) {
             const int64_t num_src1_rows = expert_row_counts[i02];

             if (num_src1_rows == 0) {