Commit 70c4e1582 for llama.cpp

commit 70c4e1582e37e4fd94104eb09301711a0f2675bc
Author: Ruben Ortlam <rortlam@redhat.com>
Date:   Thu Sep 24 15:18:47 2026 +0200

    vulkan: int8 coopmat1 matmul implementation for AMD RDNA3 and RDNA4 (#27952)

    * vulkan: add int8 coopmat quantized matmul shader

    * apply scales inline

    * use scalar sums

    * probe and directly access coopmat values instead of going through shmem

    * add q8_0 support

    * add BK_STEP to shader, default to 2

    * use larger workgroups

    * double buffering

    * preload scales

    * coopmat load first, then wmma

    * use float for scales

    * add faster RDNA int->float conversion

    * workgroup scheduling for cache proximity

    * clean up

    * use wave32

    * restructure for vgpr use

    * skip computation for inactive tiles

    * only force subgroup size 32 on AMD RDNA

    * use BK_STEP 4

    * fix compilation

    * move quant-specific prefetch function out of main file

    * add q4_1, q5_0, q5_1 support

    * restructure mmq cm1 functions

    * enable mul_mat_id support

    * fix segfault

    * fix mul_mat_id bug

    * support iq4_nl and mxfp4

    * remove elem row/col fast path, invalid for RDNA4

    * use shmem arrays for LUTs

    * use 4-byte loads where possible

    * add q3_k, q4_k, q5_k, q6_k and nvfp4 support

    * fix l warptile

    * improve performance

    * improve performance

    * improvements

    * dedup b scales

    * merge shmem arrays

    * undo uint8_t, gate to RDNA3/4

    * add RDNA4 architecture, use for hardcoded coopmat elem thread access, set BK_STEP back to 4

    * improve offset application

    * clean up

    * fix iq4_nl and nvfp4 performance

    * rdna4 tuning

    * use BK_STEP 2 on MUL_MAT_ID

    * adapt to upstream changes

    * fix shmem support function, clean up comments

    * fix warptile logic

    Co-authored-by: Piotr Wilkin (ilintar) <piotr.wilkin@syndatis.com>

    * vulkan: add IQ4_XS support to the coopmat1 integer matmul shader (#28440)

    Adds IQ4_XS to mul_mmq_cm1: dedicated block_a_load/block_a_to_shmem that
    expand both nibbles of each packed32 word through cm1_kvalues, LOAD_VEC_A 8
    and an IQ4_XS-sized a_panel_bytes estimate for the L2-friendly scheduling.

    Assisted-by: OpenAI Codex

    Co-authored-by: Claude Opus 5 (1M context) <noreply@anthropic.com>

    * avoid compiling f16 acc shader variants

    ---------

    Co-authored-by: Piotr Wilkin (ilintar) <piotr.wilkin@syndatis.com>
    Co-authored-by: Claude Opus 5 (1M context) <noreply@anthropic.com>

diff --git a/ggml/src/ggml-vulkan/ggml-vulkan-types.h b/ggml/src/ggml-vulkan/ggml-vulkan-types.h
index d12cc0071..252359bf8 100644
--- a/ggml/src/ggml-vulkan/ggml-vulkan-types.h
+++ b/ggml/src/ggml-vulkan/ggml-vulkan-types.h
@@ -386,6 +386,7 @@ enum vk_device_architecture {
     AMD_RDNA1,
     AMD_RDNA2,
     AMD_RDNA3,
+    AMD_RDNA4,
     INTEL_XE1,
     INTEL_XE2,
     NVIDIA_PRE_TURING,
diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp
index ded654197..bd416a4e0 100644
--- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp
+++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp
@@ -14,6 +14,7 @@ static vk_device_architecture get_device_architecture(const vk::PhysicalDevice&
         bool amd_shader_core_properties = false;
         bool integer_dot_product = false;
         bool subgroup_size_control = false;
+        bool shader_float8 = false;

         for (const auto& properties : ext_props) {
             if (strcmp("VK_AMD_shader_core_properties", properties.extensionName) == 0) {
@@ -22,6 +23,8 @@ static vk_device_architecture get_device_architecture(const vk::PhysicalDevice&
                 integer_dot_product = true;
             } else if (strcmp("VK_EXT_subgroup_size_control", properties.extensionName) == 0) {
                 subgroup_size_control = true;
+            } else if (strcmp("VK_EXT_shader_float8", properties.extensionName) == 0) {
+                shader_float8 = true;
             }
         }

@@ -48,6 +51,9 @@ static vk_device_architecture get_device_architecture(const vk::PhysicalDevice&
             if (shader_core_props_amd.wavefrontsPerSimd == 20) {
                 return vk_device_architecture::AMD_RDNA1;
             }
+            if (shader_float8) {
+                return vk_device_architecture::AMD_RDNA4;
+            }
             if (integer_dot_props.integerDotProduct4x8BitPackedMixedSignednessAccelerated) {
                 return vk_device_architecture::AMD_RDNA3;
             }
@@ -1497,6 +1503,8 @@ static bool ggml_vk_matmul_int_shmem_support(const vk_device& device, const std:
         case GGML_TYPE_Q8_0:    block_a_size = std430_size({{32, 4}, {fp_size,  fp_align}});                  break; // qs[8] + dm
         case GGML_TYPE_IQ4_XS:  block_a_size = std430_size({{32, 4}, {fp_size,  fp_align}});                  break; // qs[8] + d
         case GGML_TYPE_MXFP4:   block_a_size = std430_size({{32, 4}, {fp_size,  fp_align}});                  break; // qs[8] + d
+        case GGML_TYPE_IQ4_NL:  block_a_size = std430_size({{32, 4}, {fp_size,  fp_align}});                  break; // qs[8] + d
+        case GGML_TYPE_NVFP4:   block_a_size = std430_size({{32, 4}, {fp2_size, fp2_align}});                 break; // qs[8] + d_scales(vec2)
         case GGML_TYPE_Q2_K:    block_a_size = std430_size({{ 8, 4}, {2, 2}, {fp2_size, fp2_align}});         break; // qs[2] + scales(u8vec2) + dm(vec2)
         case GGML_TYPE_Q3_K:    block_a_size = std430_size({{16, 4}, {fp2_size, fp2_align}});                 break; // qs[4] + d_scales(vec2)
         case GGML_TYPE_Q4_K:    block_a_size = std430_size({{16, 4}, {fp2_size, fp2_align}});                 break; // qs[4] + dm(vec2)
@@ -1534,6 +1542,66 @@ static bool ggml_vk_matmul_int_shmem_support(const vk_device& device, const std:
     return supported;
 }

+static bool ggml_vk_matmul_cm1_int_shmem_support(const vk_device& device, const std::vector<uint32_t>& warptile, bool mul_mat_id, ggml_type src0_type) {
+
+    bool kscales2 = false;    // two scale sets per block
+    bool has_dm   = false;    // d+m as vec2 + b-side sum
+    bool has_kvalues = false;
+    switch (src0_type) {
+        case GGML_TYPE_Q4_0: case GGML_TYPE_Q5_0: case GGML_TYPE_Q8_0:
+            break;
+        case GGML_TYPE_Q4_1: case GGML_TYPE_Q5_1:
+        case GGML_TYPE_Q4_K: case GGML_TYPE_Q5_K:
+            has_dm = true;                          break;
+        case GGML_TYPE_IQ4_NL: case GGML_TYPE_IQ4_XS: case GGML_TYPE_MXFP4:
+            has_kvalues = true;                     break;
+        case GGML_TYPE_Q3_K: case GGML_TYPE_Q6_K:
+            kscales2 = true;                        break;
+        case GGML_TYPE_NVFP4:
+            kscales2 = true; has_kvalues = true;    break;
+        default:
+            return false;
+    }
+
+    const uint32_t BLOCK_SIZE = warptile[0];
+    const uint32_t BM         = warptile[1];
+    const uint32_t BN         = warptile[2];
+    const uint32_t WARP       = warptile[10];
+
+    const uint32_t BK      = 32;
+    const uint32_t BK_STEP = mul_mat_id ? 2u : 4u;
+    const uint32_t QPITCH  = BK_STEP * (BK / 4u) + 4u;
+    const uint32_t KSCALES = kscales2 ? 2u : 1u;
+
+    uint32_t total = 0;
+    total += BM * QPITCH * (uint32_t)sizeof(uint32_t);   // buf_a_qs
+    total += BN * QPITCH * (uint32_t)sizeof(uint32_t);   // buf_b_qs
+    total += has_dm ? (BM * BK_STEP * 2u * (uint32_t)sizeof(float))   // buf_a_dm (vec2)
+                    : (BM * BK_STEP * KSCALES * (uint32_t)sizeof(float)); // buf_a_d
+    total += BN * BK_STEP * (uint32_t)sizeof(float);     // buf_b_d
+    if (has_dm) {
+        total += BN * BK_STEP * (uint32_t)sizeof(float); // buf_b_s
+    }
+    if (has_kvalues) {
+        total += 16u * (uint32_t)sizeof(int8_t);         // cm1_kvalues[16]
+    }
+    if (src0_type == GGML_TYPE_NVFP4 && !device->ocp_fp4) {
+        total += 128u * (uint32_t)sizeof(float);         // ue4m3_fp32_lut[128]
+    }
+    if (mul_mat_id) {
+        total += BN * 2u * (uint32_t)sizeof(uint16_t);   // row_ids[BN] (u16vec2)
+        const uint32_t num_warps = BLOCK_SIZE / std::max(WARP, 1u);
+        total += num_warps * 4u * (uint32_t)sizeof(uint32_t); // ballots_sh[NUM_WARPS] (uvec4)
+    }
+
+    const bool supported = total <= device->properties.limits.maxComputeSharedMemorySize;
+
+    VK_LOG_DEBUG("ggml_vk_matmul_cm1_int_shmem_support(warptile=(" << warptile[0] << "," << warptile[1] << "," << warptile[2] << "), "
+                 "mul_mat_id=" << mul_mat_id << ", src0_type=" << ggml_type_name(src0_type) << ", total=" << total << ", supported=" << supported);
+
+    return supported;
+}
+
 static const std::unordered_map<std::string, uint32_t> rdna1_pipelines = {
     {"soft_max", 64}, {"im2col", 64},
     {"argmax", 64}, {"mul_mat_vec", 64},
@@ -1637,6 +1705,8 @@ void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) {
                           l_warptile_id, m_warptile_id, s_warptile_id,
                           l_warptile_mmq, m_warptile_mmq, s_warptile_mmq,
                           l_warptile_mmq_int, m_warptile_mmq_int, s_warptile_mmq_int,
+                          l_warptile_mmq_cm1_int, m_warptile_mmq_cm1_int, s_warptile_mmq_cm1_int,
+                          l_warptile_mmq_cm1_int_k, m_warptile_mmq_cm1_int_k, s_warptile_mmq_cm1_int_k,
                           l_warptile_mmq_int_k, m_warptile_mmq_int_k, s_warptile_mmq_int_k,
                           l_warptile_mmq_k, m_warptile_mmq_k, s_warptile_mmq_k,
                           l_warptile_mmqid, m_warptile_mmqid, s_warptile_mmqid,
@@ -1645,10 +1715,17 @@ void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) {
     std::array<uint32_t, 3> l_wg_denoms, m_wg_denoms, s_wg_denoms,
                             l_mmq_wg_denoms, m_mmq_wg_denoms, s_mmq_wg_denoms,
                             l_mmq_wg_denoms_k, m_mmq_wg_denoms_k, s_mmq_wg_denoms_k,
+                            l_mmq_cm1_wg_denoms_k, m_mmq_cm1_wg_denoms_k, s_mmq_cm1_wg_denoms_k,
                             l_mmqid_wg_denoms, m_mmqid_wg_denoms, s_mmqid_wg_denoms;

     uint32_t l_align, m_align, s_align;

+    // RDNA3.5 preferred wave32 here
+    const bool cm1_use_wave32 = device->vendor_id == VK_VENDOR_ID_AMD &&
+                                device->subgroup_size_control &&
+                                device->subgroup_min_size <= 32 && device->subgroup_max_size >= 32;
+    const uint32_t cm1_sg = cm1_use_wave32 ? 32 : device->subgroup_size;
+
     vk_pipeline wait_pipeline;
     CompileTask claimed_task {};
     bool has_claimed_task = false;
@@ -1706,6 +1783,10 @@ void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) {
         const uint32_t tk_m = device->coopmat_support ? device->coopmat_k : 1;
         const uint32_t tk_s = device->coopmat_support ? device->coopmat_k : 1;

+        const uint32_t itm = device->coopmat_int_m;
+        const uint32_t itn = device->coopmat_int_n;
+        const uint32_t itk = device->coopmat_int_k;
+
         const uint32_t s_warptile_wm = device->subgroup_size == 8 ? 8 : 32;

         l_warptile = { 128,             128, 128, 16, mm_warp_8 * 2, 64, 2, tm_l, tn_l, tk_l, mm_warp_8 };
@@ -1721,6 +1802,22 @@ void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) {
         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 };

+        const auto cm1_bs = [cm1_sg](uint32_t bm, uint32_t bn) {
+            return cm1_sg * (bm / std::min(cm1_sg, bm)) * (bn / 32);
+        };
+
+        l_warptile_mmq_cm1_int = { cm1_bs(128, 128), 128, 128, 32, std::min(cm1_sg, 128u), 32, 2, itm, itn, itk, cm1_sg, (uint32_t)device->architecture };
+        m_warptile_mmq_cm1_int = { cm1_bs( 64,  64),  64,  64, 32, std::min(cm1_sg,  64u), 32, 2, itm, itn, itk, cm1_sg, (uint32_t)device->architecture };
+        s_warptile_mmq_cm1_int = { cm1_bs( 32,  32),  32,  32, 32, std::min(cm1_sg,  32u), 32, 2, itm, itn, itk, cm1_sg, (uint32_t)device->architecture };
+
+        l_warptile_mmq_cm1_int_k = { cm1_bs( 64, 128),  64, 128, 32, std::min(cm1_sg,  64u), 32, 2, itm, itn, itk, cm1_sg, (uint32_t)device->architecture };
+        m_warptile_mmq_cm1_int_k = { cm1_bs( 64,  64),  64,  64, 32, std::min(cm1_sg,  64u), 32, 2, itm, itn, itk, cm1_sg, (uint32_t)device->architecture };
+        s_warptile_mmq_cm1_int_k = { cm1_bs( 32,  32),  32,  32, 32, std::min(cm1_sg,  32u), 32, 2, itm, itn, itk, cm1_sg, (uint32_t)device->architecture };
+
+        l_mmq_cm1_wg_denoms_k = { l_warptile_mmq_cm1_int_k[1], l_warptile_mmq_cm1_int_k[2], 1 };
+        m_mmq_cm1_wg_denoms_k = { m_warptile_mmq_cm1_int_k[1], m_warptile_mmq_cm1_int_k[2], 1 };
+        s_mmq_cm1_wg_denoms_k = { s_warptile_mmq_cm1_int_k[1], s_warptile_mmq_cm1_int_k[2], 1 };
+
         // K-quants use even more registers, mitigate by setting WMITER to 1
         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 };
@@ -1777,6 +1874,9 @@ void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) {
             }
         }

+        const bool use_cm1_int = device->coopmat_int_support &&
+                                 (device->architecture == AMD_RDNA3 || device->architecture == AMD_RDNA4);
+
         for (uint32_t i = 0; i < GGML_TYPE_COUNT; ++i) {
             ggml_type t = (ggml_type)i;
             // Disable medium and large matrix multiplication if not enough shared memory is available
@@ -1806,35 +1906,50 @@ void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) {

             // The q8_1 mmq path has its own (larger) shmem layout, check it separately.
             // K-quants and IQ3_S use the _int_k warptiles, others use _int.
+            // cm1 splits k-tiles on the KSCALES==2 types and shares tiles between dense/id.
             const bool is_k_quant = (t == GGML_TYPE_Q2_K || t == GGML_TYPE_Q3_K ||
                                      t == GGML_TYPE_Q4_K || t == GGML_TYPE_Q5_K ||
                                      t == GGML_TYPE_Q6_K || t == GGML_TYPE_IQ3_S);
-            const auto & s_int   = is_k_quant ? s_warptile_mmq_int_k   : s_warptile_mmq_int;
-            const auto & m_int   = is_k_quant ? m_warptile_mmq_int_k   : m_warptile_mmq_int;
-            const auto & l_int   = is_k_quant ? l_warptile_mmq_int_k   : l_warptile_mmq_int;
-            const auto & s_intid = is_k_quant ? s_warptile_mmqid_int_k : s_warptile_mmqid_int;
-            const auto & m_intid = is_k_quant ? m_warptile_mmqid_int_k : m_warptile_mmqid_int;
-            const auto & l_intid = is_k_quant ? l_warptile_mmqid_int_k : l_warptile_mmqid_int;
-
-            if (!ggml_vk_matmul_int_shmem_support(device, s_int, false, t)) {
+            const bool cm1_k_tile = (t == GGML_TYPE_Q3_K || t == GGML_TYPE_Q6_K ||
+                                     t == GGML_TYPE_NVFP4);
+
+            const auto & s_int   = use_cm1_int ? (cm1_k_tile ? s_warptile_mmq_cm1_int_k : s_warptile_mmq_cm1_int)
+                                               : (is_k_quant  ? s_warptile_mmq_int_k     : s_warptile_mmq_int);
+            const auto & m_int   = use_cm1_int ? (cm1_k_tile ? m_warptile_mmq_cm1_int_k : m_warptile_mmq_cm1_int)
+                                               : (is_k_quant  ? m_warptile_mmq_int_k     : m_warptile_mmq_int);
+            const auto & l_int   = use_cm1_int ? (cm1_k_tile ? l_warptile_mmq_cm1_int_k : l_warptile_mmq_cm1_int)
+                                               : (is_k_quant  ? l_warptile_mmq_int_k     : l_warptile_mmq_int);
+            const auto & s_intid = use_cm1_int ? (cm1_k_tile ? s_warptile_mmq_cm1_int_k : s_warptile_mmq_cm1_int)
+                                               : (is_k_quant  ? s_warptile_mmqid_int_k   : s_warptile_mmqid_int);
+            const auto & m_intid = use_cm1_int ? (cm1_k_tile ? m_warptile_mmq_cm1_int_k : m_warptile_mmq_cm1_int)
+                                               : (is_k_quant  ? m_warptile_mmqid_int_k   : m_warptile_mmqid_int);
+            const auto & l_intid = use_cm1_int ? (cm1_k_tile ? l_warptile_mmq_cm1_int_k : l_warptile_mmq_cm1_int)
+                                               : (is_k_quant  ? l_warptile_mmqid_int_k   : l_warptile_mmqid_int);
+
+            const auto int_shmem_support = [&](const std::vector<uint32_t>& wt, bool id) {
+                return use_cm1_int ? ggml_vk_matmul_cm1_int_shmem_support(device, wt, id, t)
+                                   : ggml_vk_matmul_int_shmem_support(device, wt, id, t);
+            };
+
+            if (!int_shmem_support(s_int, false)) {
                 device->mul_mat_s_int[i] = false;
                 device->mul_mat_m_int[i] = false;
                 device->mul_mat_l_int[i] = false;
-            } else if (!ggml_vk_matmul_int_shmem_support(device, m_int, false, t)) {
+            } else if (!int_shmem_support(m_int, false)) {
                 device->mul_mat_m_int[i] = false;
                 device->mul_mat_l_int[i] = false;
-            } else if (!ggml_vk_matmul_int_shmem_support(device, l_int, false, t)) {
+            } else if (!int_shmem_support(l_int, false)) {
                 device->mul_mat_l_int[i] = false;
             }

-            if (!ggml_vk_matmul_int_shmem_support(device, s_intid, true, t)) {
+            if (!int_shmem_support(s_intid, true)) {
                 device->mul_mat_id_s_int[i] = false;
                 device->mul_mat_id_m_int[i] = false;
                 device->mul_mat_id_l_int[i] = false;
-            } else if (!ggml_vk_matmul_int_shmem_support(device, m_intid, true, t)) {
+            } else if (!int_shmem_support(m_intid, true)) {
                 device->mul_mat_id_m_int[i] = false;
                 device->mul_mat_id_l_int[i] = false;
-            } else if (!ggml_vk_matmul_int_shmem_support(device, l_intid, true, t)) {
+            } else if (!int_shmem_support(l_intid, true)) {
                 device->mul_mat_id_l_int[i] = false;
             }
         }
@@ -2283,6 +2398,29 @@ void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) {
             auto tc = filter_tc(tc_base, key.type_a, key.mul_mat_id);
             if (!tc.empty()) create_mm_pipelines(key, tc, name, len, data, pc_size, pc, qs, false, true, 0, true, cm1_pin);
         };
+        // int8 MMQ helper: per-type cm1 shader, warptile passed as-is (carries DEVICE_ARCH in
+        // spec constant WARP_SIZE_IDX+1), subgroup size pinned to the warptile WARP element, no aligned variant.
+        auto cm1_create_mmq = [&](vk_matmul_pipeline_key key, const std::vector<vk_tile_config>& tc_base,
+                                  const std::string& name, size_t len, const void* data, uint32_t pc_size, uint32_t pc) {
+            spec_fn_t identity = [](const std::vector<uint32_t>& wt, bool) { return wt; };
+            auto tc = filter_tc(tc_base, key.type_a, key.mul_mat_id, true);
+            if (!tc.empty()) create_mm_pipelines(key, tc, name, len, data, pc_size, pc, identity, false, false, 0, false, true);
+        };
+
+        std::vector<vk_tile_config> tc_mmq_cm1_int = {
+            {s_warptile_mmq_cm1_int, s_mmq_wg_denoms, s_align},
+            {m_warptile_mmq_cm1_int, m_mmq_wg_denoms, m_align},
+            {l_warptile_mmq_cm1_int, l_mmq_wg_denoms, l_align},
+        };
+        std::vector<vk_tile_config> tc_mmq_cm1_int_k = {
+            {s_warptile_mmq_cm1_int_k, s_mmq_cm1_wg_denoms_k, s_align},
+            {m_warptile_mmq_cm1_int_k, m_mmq_cm1_wg_denoms_k, m_align},
+            {l_warptile_mmq_cm1_int_k, l_mmq_cm1_wg_denoms_k, l_align},
+        };
+
+        // Some quants are not performant on RDNA4, those fall back to FP16 matmul
+        const bool rdna3 = device->architecture == AMD_RDNA3;
+        const bool rdna4 = device->architecture == AMD_RDNA4;

         cm1_create({GGML_TYPE_F32, GGML_TYPE_F32, false, false}, tc_mm, "matmul_f32_f32",     matmul_f32_f32_cm1_len,     matmul_f32_f32_cm1_data,     sizeof(vk_mat_mat_push_constants), 3);
         cm1_create({GGML_TYPE_F32, GGML_TYPE_F16, false, false}, tc_mm, "matmul_f32_f16",     matmul_f32_f16_cm1_len,     matmul_f32_f16_cm1_data,     sizeof(vk_mat_mat_push_constants), 3);
@@ -2340,6 +2478,22 @@ void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) {
         }
 #undef X_CM1

+        if (device->coopmat_int_support && (rdna3 || rdna4)) {
+            cm1_create_mmq({GGML_TYPE_Q4_0,   GGML_TYPE_Q8_1, false, false}, tc_mmq_cm1_int,   "matmul_q4_0_q8_1",   matmul_q4_0_q8_1_cm1_len,   matmul_q4_0_q8_1_cm1_data,   sizeof(vk_mat_mat_push_constants), 3);
+            if (!rdna4) { cm1_create_mmq({GGML_TYPE_Q4_1, GGML_TYPE_Q8_1, false, false}, tc_mmq_cm1_int,   "matmul_q4_1_q8_1",   matmul_q4_1_q8_1_cm1_len,   matmul_q4_1_q8_1_cm1_data,   sizeof(vk_mat_mat_push_constants), 3); }
+            cm1_create_mmq({GGML_TYPE_Q5_0,   GGML_TYPE_Q8_1, false, false}, tc_mmq_cm1_int,   "matmul_q5_0_q8_1",   matmul_q5_0_q8_1_cm1_len,   matmul_q5_0_q8_1_cm1_data,   sizeof(vk_mat_mat_push_constants), 3);
+            if (!rdna4) { cm1_create_mmq({GGML_TYPE_Q5_1, GGML_TYPE_Q8_1, false, false}, tc_mmq_cm1_int,   "matmul_q5_1_q8_1",   matmul_q5_1_q8_1_cm1_len,   matmul_q5_1_q8_1_cm1_data,   sizeof(vk_mat_mat_push_constants), 3); }
+            cm1_create_mmq({GGML_TYPE_Q8_0,   GGML_TYPE_Q8_1, false, false}, tc_mmq_cm1_int,   "matmul_q8_0_q8_1",   matmul_q8_0_q8_1_cm1_len,   matmul_q8_0_q8_1_cm1_data,   sizeof(vk_mat_mat_push_constants), 3);
+            cm1_create_mmq({GGML_TYPE_IQ4_NL, GGML_TYPE_Q8_1, false, false}, tc_mmq_cm1_int,   "matmul_iq4_nl_q8_1", matmul_iq4_nl_q8_1_cm1_len, matmul_iq4_nl_q8_1_cm1_data, sizeof(vk_mat_mat_push_constants), 3);
+            cm1_create_mmq({GGML_TYPE_IQ4_XS, GGML_TYPE_Q8_1, false, false}, tc_mmq_cm1_int,   "matmul_iq4_xs_q8_1", matmul_iq4_xs_q8_1_cm1_len, matmul_iq4_xs_q8_1_cm1_data, sizeof(vk_mat_mat_push_constants), 3);
+            cm1_create_mmq({GGML_TYPE_MXFP4,  GGML_TYPE_Q8_1, false, false}, tc_mmq_cm1_int,   "matmul_mxfp4_q8_1",  matmul_mxfp4_q8_1_cm1_len,  matmul_mxfp4_q8_1_cm1_data,  sizeof(vk_mat_mat_push_constants), 3);
+            cm1_create_mmq({GGML_TYPE_Q3_K,   GGML_TYPE_Q8_1, false, false}, tc_mmq_cm1_int_k, "matmul_q3_k_q8_1",   matmul_q3_k_q8_1_cm1_len,   matmul_q3_k_q8_1_cm1_data,   sizeof(vk_mat_mat_push_constants), 3);
+            if (!rdna4) { cm1_create_mmq({GGML_TYPE_Q4_K, GGML_TYPE_Q8_1, false, false}, tc_mmq_cm1_int,   "matmul_q4_k_q8_1",   matmul_q4_k_q8_1_cm1_len,   matmul_q4_k_q8_1_cm1_data,   sizeof(vk_mat_mat_push_constants), 3); }
+            if (!rdna4) { cm1_create_mmq({GGML_TYPE_Q5_K, GGML_TYPE_Q8_1, false, false}, tc_mmq_cm1_int,   "matmul_q5_k_q8_1",   matmul_q5_k_q8_1_cm1_len,   matmul_q5_k_q8_1_cm1_data,   sizeof(vk_mat_mat_push_constants), 3); }
+            cm1_create_mmq({GGML_TYPE_Q6_K,   GGML_TYPE_Q8_1, false, false}, tc_mmq_cm1_int_k, "matmul_q6_k_q8_1",   matmul_q6_k_q8_1_cm1_len,   matmul_q6_k_q8_1_cm1_data,   sizeof(vk_mat_mat_push_constants), 3);
+            if (!rdna4) { cm1_create_mmq({GGML_TYPE_NVFP4, GGML_TYPE_Q8_1, false, false}, tc_mmq_cm1_int_k, "matmul_nvfp4_q8_1",  matmul_nvfp4_q8_1_cm1_len,  matmul_nvfp4_q8_1_cm1_data,  sizeof(vk_mat_mat_push_constants), 3); }
+        }
+
         GGML_ASSERT(device->subgroup_ballot);

         cm1_create({GGML_TYPE_F32, GGML_TYPE_F32, true, false}, tc_mm, "matmul_id_subgroup_f32_f32", matmul_id_subgroup_f32_f32_cm1_len, matmul_id_subgroup_f32_f32_cm1_data, sizeof(vk_mat_mat_id_push_constants), mul_mat_id_param_count);
@@ -2396,6 +2550,22 @@ void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) {
             FOR_EACH_LUT_FP4_TYPE(X_CM1_ID)
         }
 #undef X_CM1_ID
+
+        if (device->coopmat_int_support && (rdna3 || rdna4)) {
+            cm1_create_mmq({GGML_TYPE_Q4_0,   GGML_TYPE_Q8_1, true, false}, tc_mmq_cm1_int,   "matmul_id_subgroup_q4_0_q8_1",   matmul_id_subgroup_q4_0_q8_1_cm1_len,   matmul_id_subgroup_q4_0_q8_1_cm1_data,   sizeof(vk_mat_mat_id_push_constants), mul_mat_id_param_count);
+            cm1_create_mmq({GGML_TYPE_Q4_1,   GGML_TYPE_Q8_1, true, false}, tc_mmq_cm1_int,   "matmul_id_subgroup_q4_1_q8_1",   matmul_id_subgroup_q4_1_q8_1_cm1_len,   matmul_id_subgroup_q4_1_q8_1_cm1_data,   sizeof(vk_mat_mat_id_push_constants), mul_mat_id_param_count);
+            cm1_create_mmq({GGML_TYPE_Q5_0,   GGML_TYPE_Q8_1, true, false}, tc_mmq_cm1_int,   "matmul_id_subgroup_q5_0_q8_1",   matmul_id_subgroup_q5_0_q8_1_cm1_len,   matmul_id_subgroup_q5_0_q8_1_cm1_data,   sizeof(vk_mat_mat_id_push_constants), mul_mat_id_param_count);
+            cm1_create_mmq({GGML_TYPE_Q5_1,   GGML_TYPE_Q8_1, true, false}, tc_mmq_cm1_int,   "matmul_id_subgroup_q5_1_q8_1",   matmul_id_subgroup_q5_1_q8_1_cm1_len,   matmul_id_subgroup_q5_1_q8_1_cm1_data,   sizeof(vk_mat_mat_id_push_constants), mul_mat_id_param_count);
+            cm1_create_mmq({GGML_TYPE_Q8_0,   GGML_TYPE_Q8_1, true, false}, tc_mmq_cm1_int,   "matmul_id_subgroup_q8_0_q8_1",   matmul_id_subgroup_q8_0_q8_1_cm1_len,   matmul_id_subgroup_q8_0_q8_1_cm1_data,   sizeof(vk_mat_mat_id_push_constants), mul_mat_id_param_count);
+            cm1_create_mmq({GGML_TYPE_IQ4_NL, GGML_TYPE_Q8_1, true, false}, tc_mmq_cm1_int,   "matmul_id_subgroup_iq4_nl_q8_1", matmul_id_subgroup_iq4_nl_q8_1_cm1_len, matmul_id_subgroup_iq4_nl_q8_1_cm1_data, sizeof(vk_mat_mat_id_push_constants), mul_mat_id_param_count);
+            cm1_create_mmq({GGML_TYPE_IQ4_XS, GGML_TYPE_Q8_1, true, false}, tc_mmq_cm1_int,   "matmul_id_subgroup_iq4_xs_q8_1", matmul_id_subgroup_iq4_xs_q8_1_cm1_len, matmul_id_subgroup_iq4_xs_q8_1_cm1_data, sizeof(vk_mat_mat_id_push_constants), mul_mat_id_param_count);
+            cm1_create_mmq({GGML_TYPE_MXFP4,  GGML_TYPE_Q8_1, true, false}, tc_mmq_cm1_int,   "matmul_id_subgroup_mxfp4_q8_1",  matmul_id_subgroup_mxfp4_q8_1_cm1_len,  matmul_id_subgroup_mxfp4_q8_1_cm1_data,  sizeof(vk_mat_mat_id_push_constants), mul_mat_id_param_count);
+            cm1_create_mmq({GGML_TYPE_Q3_K,   GGML_TYPE_Q8_1, true, false}, tc_mmq_cm1_int_k, "matmul_id_subgroup_q3_k_q8_1",   matmul_id_subgroup_q3_k_q8_1_cm1_len,   matmul_id_subgroup_q3_k_q8_1_cm1_data,   sizeof(vk_mat_mat_id_push_constants), mul_mat_id_param_count);
+            cm1_create_mmq({GGML_TYPE_Q4_K,   GGML_TYPE_Q8_1, true, false}, tc_mmq_cm1_int,   "matmul_id_subgroup_q4_k_q8_1",   matmul_id_subgroup_q4_k_q8_1_cm1_len,   matmul_id_subgroup_q4_k_q8_1_cm1_data,   sizeof(vk_mat_mat_id_push_constants), mul_mat_id_param_count);
+            cm1_create_mmq({GGML_TYPE_Q5_K,   GGML_TYPE_Q8_1, true, false}, tc_mmq_cm1_int,   "matmul_id_subgroup_q5_k_q8_1",   matmul_id_subgroup_q5_k_q8_1_cm1_len,   matmul_id_subgroup_q5_k_q8_1_cm1_data,   sizeof(vk_mat_mat_id_push_constants), mul_mat_id_param_count);
+            cm1_create_mmq({GGML_TYPE_Q6_K,   GGML_TYPE_Q8_1, true, false}, tc_mmq_cm1_int_k, "matmul_id_subgroup_q6_k_q8_1",   matmul_id_subgroup_q6_k_q8_1_cm1_len,   matmul_id_subgroup_q6_k_q8_1_cm1_data,   sizeof(vk_mat_mat_id_push_constants), mul_mat_id_param_count);
+            if (!rdna4) { cm1_create_mmq({GGML_TYPE_NVFP4, GGML_TYPE_Q8_1, true, false}, tc_mmq_cm1_int_k, "matmul_id_subgroup_nvfp4_q8_1",  matmul_id_subgroup_nvfp4_q8_1_cm1_len,  matmul_id_subgroup_nvfp4_q8_1_cm1_data,  sizeof(vk_mat_mat_id_push_constants), mul_mat_id_param_count); }
+        }
     } else
 #endif  // defined(VK_KHR_cooperative_matrix) && defined(GGML_VULKAN_COOPMAT_GLSLC_SUPPORT)
     {
@@ -6114,27 +6284,34 @@ static void ggml_vk_mul_mat_q_f16(ggml_backend_vk_context * ctx, vk_context& sub
     // Reformat and convert to fp16 if non-contiguous, or for coopmat2 for better perf
     const bool x_non_contig = (ctx->device->coopmat2 && src0->type == GGML_TYPE_F32) ||
                               !ggml_vk_dim01_contiguous(src0);
-    const bool y_non_contig = (ctx->device->coopmat2 && src1->type == GGML_TYPE_F32) ||
-                              // coopmat1: force f32->f16 conversion so the f16 B-type quant pipeline is used.
-                              (ctx->device->coopmat_support && !ctx->device->coopmat2 &&
-                               ggml_is_quantized(src0->type) && src1->type == GGML_TYPE_F32) ||
-                              (src0->type == GGML_TYPE_BF16 && src1->type != GGML_TYPE_BF16) ||
-                              !ggml_vk_dim01_contiguous(src1);
-
     // If src0 is BF16, try to use a BF16 x BF16 multiply
     ggml_type f16_type = src0->type == GGML_TYPE_BF16 ? GGML_TYPE_BF16 : GGML_TYPE_F16;

-    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;
+    // Prefer the int8 MMQ path (quantize src1 to q8_1) whenever a matching pipeline exists.
+    // The pipeline lookup returns nullptr for types without a q8_1 pipeline (e.g. RDNA4-skipped
+    // quants), in which case coopmat1 falls back to the f16 B-type quant matmul below.
+    bool quantize_y = (ctx->device->integer_dot_product || ctx->device->coopmat_int_support) &&
+                      src1->type == GGML_TYPE_F32 && ggml_is_contiguous(src1) && (ne11 * ne10) % 4 == 0;

     // Check for mmq first
     const std::vector<vk_matmul_pipeline_pair>* mmp_map = quantize_y ? ggml_vk_get_mul_mat_mat_pipeline_map(ctx, src0->type, GGML_TYPE_Q8_1, (ggml_prec)dst->op_params[0]) : nullptr;
+    if (mmp_map == nullptr) {
+        quantize_y = false;
+    }
+
+    const bool y_non_contig = (ctx->device->coopmat2 && src1->type == GGML_TYPE_F32) ||
+                              // coopmat1: force f32->f16 conversion so the f16 B-type quant pipeline is
+                              // used, but only when the int8 MMQ path above is not taken.
+                              (ctx->device->coopmat_support && !ctx->device->coopmat2 && !quantize_y &&
+                               ggml_is_quantized(src0->type) && src1->type == GGML_TYPE_F32) ||
+                              (src0->type == GGML_TYPE_BF16 && src1->type != GGML_TYPE_BF16) ||
+                              !ggml_vk_dim01_contiguous(src1);
+
+    const bool y_f32_kernel = src1->type == GGML_TYPE_F32 && !y_non_contig;

     if (mmp_map == nullptr) {
         // Fall back to f16 dequant mul mat
         mmp_map = ggml_vk_get_mul_mat_mat_pipeline_map(ctx, src0->type, y_non_contig ? f16_type : src1->type, (ggml_prec)dst->op_params[0]);
-        quantize_y = false;
     }

     const bool qx_needs_dequant = mmp_map == nullptr || x_non_contig;
@@ -15975,7 +16152,7 @@ bool ggml_vk_khr_cooperative_matrix_support(const vk::PhysicalDeviceProperties&
     case VK_VENDOR_ID_AMD:
         if (driver_props.driverID == vk::DriverId::eAmdProprietary || driver_props.driverID == vk::DriverId::eAmdOpenSource) {
             // Workaround for AMD proprietary driver reporting support on all GPUs
-            return arch == vk_device_architecture::AMD_RDNA3;
+            return arch == vk_device_architecture::AMD_RDNA3 || arch == vk_device_architecture::AMD_RDNA4;
         }
         return true;
     case VK_VENDOR_ID_QUALCOMM:
diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1.comp b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1.comp
new file mode 100644
index 000000000..7cab9a119
--- /dev/null
+++ b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1.comp
@@ -0,0 +1,510 @@
+#version 450
+
+#extension GL_EXT_control_flow_attributes : enable
+#extension GL_EXT_shader_16bit_storage : require
+#extension GL_EXT_shader_explicit_arithmetic_types_int8 : require
+#extension GL_EXT_shader_explicit_arithmetic_types_float16 : require
+
+#extension GL_KHR_shader_subgroup_basic : require
+#extension GL_KHR_cooperative_matrix : require
+#extension GL_KHR_memory_scope_semantics : enable
+
+#if defined(MUL_MAT_ID_USE_SUBGROUPS)
+#extension GL_KHR_shader_subgroup_ballot : enable
+#endif
+
+#ifdef MUL_MAT_ID
+#extension GL_EXT_shader_explicit_arithmetic_types_int16 : require
+#endif
+
+#include "types.glsl"
+
+#if defined(DATA_A_Q3_K) || defined(DATA_A_Q6_K) || defined(DATA_A_NVFP4)
+#define KSCALES 2
+#else
+#define KSCALES 1
+#endif
+
+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[];};
+#if defined(A_TYPE_PACKED16)
+layout (binding = 0) readonly buffer A_PACKED16 {A_TYPE_PACKED16 data_a_packed16[];};
+#endif
+#if defined(A_TYPE_PACKED32)
+layout (binding = 0) readonly buffer A_PACKED32 {A_TYPE_PACKED32 data_a_packed32[];};
+#endif
+layout (binding = 1) readonly buffer B {block_q8_1_x4_packed128 data_b[];};
+layout (binding = 2) writeonly buffer D {D_TYPE data_d[];};
+
+#ifdef MUL_MAT_ID
+layout (binding = 3) readonly buffer IDS {int data_ids[];};
+layout (binding = 4) readonly buffer Counts {int data_expert_count[];};
+#endif
+
+layout (push_constant) uniform parameter
+{
+    uint M;
+    uint N;
+    uint K;
+    uint stride_a;
+    uint stride_b;
+    uint stride_d;
+
+    uint batch_stride_a;
+    uint batch_stride_b;
+    uint batch_stride_d;
+
+#ifdef MUL_MAT_ID
+    uint nei0;
+    uint nei1;
+    uint nbi1;
+    uint ne11;
+    uint n_experts;
+    uint hoist_row_ids;
+#else
+    uint base_work_group_z;
+    uint num_batches;
+    uint k_split;
+    uint ne02;
+    uint ne12;
+    uint broadcast2;
+    uint broadcast3;
+#endif
+} p;
+
+layout (constant_id = 0) const uint BLOCK_SIZE = 256;
+layout (constant_id = 1) const uint BM = 128;
+layout (constant_id = 2) const uint BN = 128;
+// layout (constant_id = 3) const uint BK = 32;
+layout (constant_id = 4) const uint WM = 64;
+layout (constant_id = 5) const uint WN = 32;
+layout (constant_id = 7) const uint TM = 16;
+layout (constant_id = 8) const uint TN = 16;
+layout (constant_id = 9) const uint TK = 16;
+layout (constant_id = 10) const uint WARP = 32;
+layout (constant_id = 11) const uint DEVICE_ARCH = 0; // vk_device_architecture (ggml-vulkan.cpp)
+#define VK_ARCH_AMD_RDNA4 5u
+
+#define BK 32
+#ifdef MUL_MAT_ID
+#define BK_STEP 2
+#else
+#define BK_STEP 4
+#endif
+#define GROUP_A_BUDGET (16u * 1024u * 1024u)
+
+const uint QPITCH = BK_STEP * (BK / 4) + 4;
+
+shared uint32_t buf_a_qs[BM * QPITCH];
+#if defined(DATA_A_Q4_1) || defined(DATA_A_Q5_1) || defined(DATA_A_Q4_K) || defined(DATA_A_Q5_K)
+shared vec2 buf_a_dm[BM * BK_STEP];   // .x = d, .y = m
+#else
+shared float buf_a_d[BM * BK_STEP * KSCALES];
+#endif
+
+shared uint32_t buf_b_qs[BN * QPITCH];
+shared float buf_b_d[BN * BK_STEP];
+
+#if defined(DATA_A_Q4_1) || defined(DATA_A_Q5_1) || defined(DATA_A_Q4_K) || defined(DATA_A_Q5_K)
+shared float buf_b_s[BN * BK_STEP];
+#endif
+
+#if defined(DATA_A_IQ4_NL) || defined(DATA_A_IQ4_XS) || defined(DATA_A_MXFP4) || defined(DATA_A_NVFP4)
+shared int8_t cm1_kvalues[16];
+#endif
+
+#if defined(DATA_A_QUANT_K) || defined(DATA_A_IQ4_XS) || defined(DATA_A_NVFP4)
+#define LOAD_VEC_A 8
+#else
+#define LOAD_VEC_A (4 * QUANT_R)
+#endif
+#define LOAD_VEC_B 16
+
+const uint CM_ELEMS = (TM * TN) / WARP;
+#define ACC_BIAS_BITS 0x4B400000
+#define ACC_BIAS_F    12582912.0f
+const bool USE_MAGIC_BIAS = WARP != 32;
+
+// Accumulator row for element e: RDNA4 blocked, RDNA3/3.5 interleaved.
+uint cm_elem_row(uint e) {
+    const uint row_half = gl_SubgroupInvocationID / TN;
+    return (DEVICE_ARCH == VK_ARCH_AMD_RDNA4) ? (e + row_half * CM_ELEMS) : (row_half + 2u * e);
+}
+
+// min_term = asymmetric-quant min*b_sum correction (0 for symmetric types).
+ACC_TYPE cm1_accumulate(ACC_TYPE prev, int acc_e, float scale_a, float nbias_a, float scale_b, float min_term) {
+    if (USE_MAGIC_BIAS) {
+        const float t = fma(intBitsToFloat(acc_e), scale_a, nbias_a);
+        return ACC_TYPE(fma(t, scale_b, float(prev) + min_term));
+    }
+    return prev + ACC_TYPE(fma(float(acc_e) * scale_a, scale_b, min_term));
+}
+
+#ifdef MUL_MAT_ID
+#define NUM_WARPS (BLOCK_SIZE / WARP)
+#include "mul_mm_id_funcs.glsl"
+#endif
+
+#include "mul_mmq_cm1_funcs.glsl"
+
+void main() {
+#if defined(DATA_A_IQ4_NL) || defined(DATA_A_IQ4_XS)
+    if (gl_LocalInvocationIndex < 16u) {
+        cm1_kvalues[gl_LocalInvocationIndex] = kvalues_iq4nl_const[gl_LocalInvocationIndex];
+    }
+    barrier();
+#elif defined(DATA_A_MXFP4)
+    if (gl_LocalInvocationIndex < 16u) {
+        cm1_kvalues[gl_LocalInvocationIndex] = kvalues_mxfp4_const[gl_LocalInvocationIndex];
+    }
+    barrier();
+#elif defined(DATA_A_NVFP4)
+    if (gl_LocalInvocationIndex < 16u) {
+        cm1_kvalues[gl_LocalInvocationIndex] = kvalues_mxfp4_const[gl_LocalInvocationIndex];
+    }
+#if !defined(USE_OCP_FP4)
+    for (uint i = gl_LocalInvocationIndex; i < 128u; i += BLOCK_SIZE) {
+        ue4m3_fp32_lut[i] = ue4m3_to_fp32_build(i);
+    }
+#endif
+    barrier();
+#endif
+
+    const uint blocks_m = (p.M + BM - 1) / BM;
+    const uint ik = gl_WorkGroupID.x / blocks_m;
+
+#ifdef MUL_MAT_ID
+    const uint ic = gl_WorkGroupID.y;
+    const uint ir = gl_WorkGroupID.x % blocks_m;
+    const uint expert_idx = gl_WorkGroupID.z;
+    if (ic * BN >= data_expert_count[expert_idx]) {
+        return;
+    }
+#else
+    // L2-friendly workgroup scheduling
+    const uint blocks_n = (p.N + BN - 1) / BN;
+#if defined(DATA_A_IQ4_XS)
+    const uint a_panel_bytes = (BM * p.K) / 2 + (BM * p.K) / 32;
+#else
+    const uint a_panel_bytes = BM * p.K + (BM * p.K) / 16;
+#endif
+    const uint group_m = clamp(GROUP_A_BUDGET / max(a_panel_bytes, 1u), 1u, min(blocks_m, 32u));
+    const uint tiles_per_group = group_m * blocks_n;
+    const uint lin = gl_WorkGroupID.y * blocks_m + (gl_WorkGroupID.x % blocks_m);
+    const uint group_id = lin / tiles_per_group;
+    const uint first_m = group_id * group_m;
+    const uint gsize = min(blocks_m - first_m, group_m);
+    const uint in_group = lin - group_id * tiles_per_group;
+    const uint ir = first_m + in_group % gsize;
+    const uint ic = in_group / gsize;
+#endif
+
+#ifndef MUL_MAT_ID
+    const uint batch_idx = gl_WorkGroupID.z + p.base_work_group_z;
+
+    const uint i13 = batch_idx / p.ne12;
+    const uint i12 = batch_idx % p.ne12;
+
+    const uint i03 = i13 / p.broadcast3;
+    const uint i02 = i12 / p.broadcast2;
+
+    const uint batch_idx_a = i03 * p.ne02 + i02;
+#endif
+
+    const uint warp_i = gl_SubgroupID;
+
+    const uint cms_per_row = WM / TM;
+    const uint cms_per_col = WN / TN;
+
+    const uint warp_r = warp_i % (BM / WM);
+    const uint warp_c = warp_i / (BM / WM);
+
+    const uint elem_col0 = gl_SubgroupInvocationID % TN;
+
+    const uint loadr_a = gl_LocalInvocationID.x % (BK / LOAD_VEC_A);
+    const uint loadc_a = gl_LocalInvocationID.x / (BK / LOAD_VEC_A);
+    const uint loadr_b = gl_LocalInvocationID.x % (BK / LOAD_VEC_B);
+    const uint loadc_b = gl_LocalInvocationID.x / (BK / LOAD_VEC_B);
+
+    const uint loadstride_a = BLOCK_SIZE * LOAD_VEC_A / BK;
+    const uint loadstride_b = BLOCK_SIZE * LOAD_VEC_B / BK;
+
+#ifdef MUL_MAT_ID
+    if (p.hoist_row_ids != 0) {
+        load_row_ids_hoisted(expert_idx, ic);
+    } else {
+#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++;
+                }
+            }
+        }
+
+        barrier();
+#endif
+    }
+
+    if (ic * BN >= _ne1) return;
+#endif
+
+#ifdef MUL_MAT_ID
+    const uint start_k = 0;
+    const uint end_k = p.K;
+#else
+    const uint start_k = ik * p.k_split;
+    const uint end_k = min(p.K, (ik + 1) * p.k_split);
+#endif
+
+    uint pos_a_ib =
+#ifdef MUL_MAT_ID
+        expert_idx * (p.batch_stride_a / BK) +
+#else
+        batch_idx_a * (p.batch_stride_a / BK) +
+#endif
+        (ir * BM * p.stride_a + start_k) / BK;
+#ifdef MUL_MAT_ID
+    uint pos_b_ib = 0;
+#else
+    uint pos_b_ib = (batch_idx * p.batch_stride_b + ic * BN * p.stride_b + start_k) / BK;
+#endif
+
+    ACC_TYPE sums[cms_per_row * cms_per_col * CM_ELEMS];
+    [[unroll]] for (uint i = 0; i < cms_per_row * cms_per_col * CM_ELEMS; i++) {
+        sums[i] = ACC_TYPE(0.0);
+    }
+
+    // Double-buffering: prefetch registers
+    const uint A_LOADS = (BM + loadstride_a - 1) / loadstride_a;
+    const uint B_LOADS = (BN + loadstride_b - 1) / loadstride_b;
+
+    block_a_prefetch pre_a[A_LOADS * BK_STEP];
+    block_b_prefetch pre_b[B_LOADS * BK_STEP];
+
+    if (start_k < end_k) {
+        PREFETCH_BLOCK(start_k)
+    }
+
+    const uint a_row0 = warp_r * WM;
+    const uint b_col0 = warp_c * WN;
+#ifdef MUL_MAT_ID
+    const bool active_col_tile = ic * BN + b_col0 < _ne1;
+#else
+    const bool active_col_tile = ic * BN + b_col0 < p.N;
+#endif
+
+    barrier();
+
+    for (uint block = start_k; block < end_k; block += BK * BK_STEP) {
+        STORE_BLOCK_TO_LDS(block)
+
+        barrier();
+
+        pos_a_ib += BK_STEP;
+        pos_b_ib += BK_STEP;
+
+        const uint next_block = block + BK * BK_STEP;
+        if (next_block < end_k) {
+            PREFETCH_BLOCK(next_block)
+        }
+
+        if (active_col_tile) {
+        [[unroll]] for (uint ks = 0; ks < BK_STEP; ks++) {
+            const uint K_SUB = BK / TK;
+
+#if KSCALES == 2
+            [[unroll]] for (uint h = 0; h < K_SUB; h++) {
+                [[unroll]] for (uint r = 0; r < cms_per_row; r++) {
+                    coopmat<int8_t, gl_ScopeSubgroup, TM, TK, gl_MatrixUseA> cache_a;
+                    coopMatLoad(cache_a, buf_a_qs, (a_row0 + r * TM) * QPITCH + ks * (BK / 4) + h * (TK / 4), QPITCH, gl_CooperativeMatrixLayoutRowMajor);
+
+                    float scale_a[CM_ELEMS];
+                    float nbias_a[CM_ELEMS];
+                    [[unroll]] for (uint e = 0; e < CM_ELEMS; e++) {
+                        scale_a[e] = buf_a_d[(ks * KSCALES + h) * BM + a_row0 + r * TM + cm_elem_row(e)];
+                        if (USE_MAGIC_BIAS) {
+                            nbias_a[e] = -ACC_BIAS_F * scale_a[e];
+                        }
+                    }
+
+                    [[unroll]] for (uint c = 0; c < cms_per_col; c++) {
+                        coopmat<int8_t, gl_ScopeSubgroup, TK, TN, gl_MatrixUseB> cache_b;
+                        coopMatLoad(cache_b, buf_b_qs, (b_col0 + c * TN) * QPITCH + ks * (BK / 4) + h * (TK / 4), QPITCH, gl_CooperativeMatrixLayoutColumnMajor);
+
+                        const float scale_b_v = buf_b_d[ks * BN + b_col0 + c * TN + elem_col0];
+
+                        coopmat<int32_t, gl_ScopeSubgroup, TM, TN, gl_MatrixUseAccumulator> acc =
+                            coopmat<int32_t, gl_ScopeSubgroup, TM, TN, gl_MatrixUseAccumulator>(
+                                USE_MAGIC_BIAS ? ACC_BIAS_BITS : 0);
+                        acc = coopMatMulAdd(cache_a, cache_b, acc);
+
+                        const uint tile_idx = r * cms_per_col + c;
+                        [[unroll]] for (uint e = 0; e < CM_ELEMS; e++) {
+                            sums[tile_idx * CM_ELEMS + e] = cm1_accumulate(
+                                sums[tile_idx * CM_ELEMS + e], int(acc[e]),
+                                scale_a[e], nbias_a[e], scale_b_v, 0.0);
+                        }
+                    }
+                }
+            }
+#elif defined(DATA_A_Q4_1) || defined(DATA_A_Q5_1) || defined(DATA_A_Q4_K) || defined(DATA_A_Q5_K)
+            // Preload all A/B fragments up front (ILP).
+            coopmat<int8_t, gl_ScopeSubgroup, TM, TK, gl_MatrixUseA> cache_a[cms_per_row * K_SUB];
+            coopmat<int8_t, gl_ScopeSubgroup, TK, TN, gl_MatrixUseB> cache_b[cms_per_col * K_SUB];
+
+            [[unroll]] for (uint r = 0; r < cms_per_row; r++) {
+                [[unroll]] for (uint h = 0; h < K_SUB; h++) {
+                    coopMatLoad(cache_a[r * K_SUB + h], buf_a_qs, (a_row0 + r * TM) * QPITCH + ks * (BK / 4) + h * (TK / 4), QPITCH, gl_CooperativeMatrixLayoutRowMajor);
+                }
+            }
+            [[unroll]] for (uint c = 0; c < cms_per_col; c++) {
+                [[unroll]] for (uint h = 0; h < K_SUB; h++) {
+                    coopMatLoad(cache_b[c * K_SUB + h], buf_b_qs, (b_col0 + c * TN) * QPITCH + ks * (BK / 4) + h * (TK / 4), QPITCH, gl_CooperativeMatrixLayoutColumnMajor);
+                }
+            }
+
+            float scale_b[cms_per_col];
+            float bs[cms_per_col];
+            [[unroll]] for (uint c = 0; c < cms_per_col; c++) {
+                scale_b[c] = buf_b_d[ks * BN + b_col0 + c * TN + elem_col0];
+                bs[c] = float(buf_b_s[ks * BN + b_col0 + c * TN + elem_col0]);
+            }
+
+            [[unroll]] for (uint r = 0; r < cms_per_row; r++) {
+                float scale_a[CM_ELEMS];
+                float nbias_a[CM_ELEMS];
+                float ma[CM_ELEMS];
+                [[unroll]] for (uint e = 0; e < CM_ELEMS; e++) {
+                    vec2 dm = buf_a_dm[ks * BM + a_row0 + r * TM + cm_elem_row(e)];
+                    scale_a[e] = dm.x;
+                    if (USE_MAGIC_BIAS) {
+                        nbias_a[e] = -ACC_BIAS_F * scale_a[e];
+                    }
+                    ma[e] = dm.y;
+                }
+
+                [[unroll]] for (uint c = 0; c < cms_per_col; c++) {
+                    coopmat<int32_t, gl_ScopeSubgroup, TM, TN, gl_MatrixUseAccumulator> acc =
+                        coopmat<int32_t, gl_ScopeSubgroup, TM, TN, gl_MatrixUseAccumulator>(
+                            USE_MAGIC_BIAS ? ACC_BIAS_BITS : 0);
+
+                    [[unroll]] for (uint h = 0; h < K_SUB; h++) {
+                        acc = coopMatMulAdd(cache_a[r * K_SUB + h], cache_b[c * K_SUB + h], acc);
+                    }
+
+                    const uint tile_idx = r * cms_per_col + c;
+                    [[unroll]] for (uint e = 0; e < CM_ELEMS; e++) {
+                        sums[tile_idx * CM_ELEMS + e] = cm1_accumulate(
+                            sums[tile_idx * CM_ELEMS + e], int(acc[e]),
+                            scale_a[e], nbias_a[e], scale_b[c], ma[e] * bs[c]);
+                    }
+                }
+            }
+#else
+            // Preload all A/B fragments up front (ILP).
+            coopmat<int8_t, gl_ScopeSubgroup, TM, TK, gl_MatrixUseA> cache_a[cms_per_row * K_SUB];
+            coopmat<int8_t, gl_ScopeSubgroup, TK, TN, gl_MatrixUseB> cache_b[cms_per_col * K_SUB];
+
+            [[unroll]] for (uint r = 0; r < cms_per_row; r++) {
+                [[unroll]] for (uint h = 0; h < K_SUB; h++) {
+                    coopMatLoad(cache_a[r * K_SUB + h], buf_a_qs, (a_row0 + r * TM) * QPITCH + ks * (BK / 4) + h * (TK / 4), QPITCH, gl_CooperativeMatrixLayoutRowMajor);
+                }
+            }
+            [[unroll]] for (uint c = 0; c < cms_per_col; c++) {
+                [[unroll]] for (uint h = 0; h < K_SUB; h++) {
+                    coopMatLoad(cache_b[c * K_SUB + h], buf_b_qs, (b_col0 + c * TN) * QPITCH + ks * (BK / 4) + h * (TK / 4), QPITCH, gl_CooperativeMatrixLayoutColumnMajor);
+                }
+            }
+
+            float scale_b[cms_per_col];
+            [[unroll]] for (uint c = 0; c < cms_per_col; c++) {
+                scale_b[c] = buf_b_d[ks * BN + b_col0 + c * TN + elem_col0];
+            }
+
+            [[unroll]] for (uint r = 0; r < cms_per_row; r++) {
+                float scale_a[CM_ELEMS];
+                float nbias_a[CM_ELEMS];
+                [[unroll]] for (uint e = 0; e < CM_ELEMS; e++) {
+                    scale_a[e] = buf_a_d[ks * BM + a_row0 + r * TM + cm_elem_row(e)];
+                    if (USE_MAGIC_BIAS) {
+                        nbias_a[e] = -ACC_BIAS_F * scale_a[e];
+                    }
+                }
+
+                [[unroll]] for (uint c = 0; c < cms_per_col; c++) {
+                    coopmat<int32_t, gl_ScopeSubgroup, TM, TN, gl_MatrixUseAccumulator> acc =
+                        coopmat<int32_t, gl_ScopeSubgroup, TM, TN, gl_MatrixUseAccumulator>(
+                            USE_MAGIC_BIAS ? ACC_BIAS_BITS : 0);
+
+                    [[unroll]] for (uint h = 0; h < K_SUB; h++) {
+                        acc = coopMatMulAdd(cache_a[r * K_SUB + h], cache_b[c * K_SUB + h], acc);
+                    }
+
+                    const uint tile_idx = r * cms_per_col + c;
+                    [[unroll]] for (uint e = 0; e < CM_ELEMS; e++) {
+                        sums[tile_idx * CM_ELEMS + e] = cm1_accumulate(
+                            sums[tile_idx * CM_ELEMS + e], int(acc[e]),
+                            scale_a[e], nbias_a[e], scale_b[c], 0.0);
+                    }
+                }
+            }
+#endif // KSCALES
+        }
+        }
+
+        barrier();
+    }
+
+#undef PREFETCH_BLOCK
+#undef STORE_BLOCK_TO_LDS
+#undef B_IB_CALC
+
+    const uint dr = ir * BM + a_row0;
+    const uint dc = ic * BN + b_col0;
+
+#ifdef MUL_MAT_ID
+    [[unroll]] for (uint r = 0; r < cms_per_row; r++) {
+        [[unroll]] for (uint c = 0; c < cms_per_col; c++) {
+            const uint tile_idx = r * cms_per_col + c;
+            [[unroll]] for (uint e = 0; e < CM_ELEMS; e++) {
+                const uint col_i = dc + c * TN + elem_col0;
+                if (col_i >= _ne1) continue;
+
+                const uint row_g = dr + r * TM + cm_elem_row(e);
+                if (row_g >= p.M) continue;
+
+                const u16vec2 row_idx = row_ids[col_i - ic * BN];
+                const uint store_offset = row_idx.y * p.batch_stride_d + row_idx.x * p.stride_d + row_g;
+                data_d[store_offset] = D_TYPE(sums[tile_idx * CM_ELEMS + e]);
+            }
+        }
+    }
+#else
+    const uint offsets = batch_idx * p.batch_stride_d + ik * p.batch_stride_d * p.num_batches;
+
+    [[unroll]] for (uint r = 0; r < cms_per_row; r++) {
+        [[unroll]] for (uint c = 0; c < cms_per_col; c++) {
+            const uint tile_idx = r * cms_per_col + c;
+            [[unroll]] for (uint e = 0; e < CM_ELEMS; e++) {
+                const uint row_g = dr + r * TM + cm_elem_row(e);
+                const uint col_g = dc + c * TN + elem_col0;
+                if (row_g < p.M && col_g < p.N) {
+                    data_d[offsets + col_g * p.stride_d + row_g] = D_TYPE(sums[tile_idx * CM_ELEMS + e]);
+                }
+            }
+        }
+    }
+#endif // MUL_MAT_ID
+}
diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1_funcs.glsl b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1_funcs.glsl
new file mode 100644
index 000000000..1760c138e
--- /dev/null
+++ b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1_funcs.glsl
@@ -0,0 +1,594 @@
+// Per-quant-type data structures and functions for the cm1 int8 coopmat path.
+// Each quant type defines:
+//   struct block_a_prefetch  — register data for one A-block per thread
+//   block_a_load()           — load from global memory into a block_a_prefetch
+//   block_a_to_shmem()       — unpack and write to shared memory
+
+#if defined(DATA_A_Q4_0)
+
+struct block_a_prefetch {
+    uint32_t qs;
+    float16_t d;
+};
+
+block_a_prefetch block_a_load(uint ib, uint loadr) {
+    block_a_prefetch blk;
+    blk.qs = pack32(u16vec2(data_a_packed16[ib].qs[loadr * 2],
+                             data_a_packed16[ib].qs[loadr * 2 + 1]));
+    blk.d = data_a_packed16[ib].d;
+    return blk;
+}
+
+void block_a_to_shmem(block_a_prefetch blk, uint buf_ib, uint ks, uint loadr) {
+    uint32_t lo4 = blk.qs & 0x0F0F0F0F;
+    uint32_t hi4 = (blk.qs >> 4) & 0x0F0F0F0F;
+    lo4 = ((lo4 | 0x80808080) - 0x08080808) ^ 0x80808080;
+    hi4 = ((hi4 | 0x80808080) - 0x08080808) ^ 0x80808080;
+    buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + loadr    ] = lo4;
+    buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + loadr + 4] = hi4;
+
+    if (loadr == 0) {
+        buf_a_d[ks * BM + buf_ib] = float(blk.d);
+    }
+}
+
+#elif defined(DATA_A_Q4_1)
+
+struct block_a_prefetch {
+    uint32_t qs;
+    f16vec2 dm;
+};
+
+block_a_prefetch block_a_load(uint ib, uint loadr) {
+    block_a_prefetch blk;
+    blk.qs = data_a_packed32[ib].qs[loadr];
+    blk.dm = data_a_packed32[ib].dm;
+    return blk;
+}
+
+void block_a_to_shmem(block_a_prefetch blk, uint buf_ib, uint ks, uint loadr) {
+    // Store raw unsigned nibbles; the -8 offset is absorbed by the min term.
+    uint32_t lo4 = blk.qs & 0x0F0F0F0F;
+    uint32_t hi4 = (blk.qs >> 4) & 0x0F0F0F0F;
+    buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + loadr    ] = lo4;
+    buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + loadr + 4] = hi4;
+
+    if (loadr == 0) {
+        buf_a_dm[ks * BM + buf_ib] = vec2(float(blk.dm.x), float(blk.dm.y));
+    }
+}
+
+#elif defined(DATA_A_Q5_0)
+
+struct block_a_prefetch {
+    uint32_t qs;
+    float16_t d;
+    uint32_t qh;
+};
+
+block_a_prefetch block_a_load(uint ib, uint loadr) {
+    block_a_prefetch blk;
+    blk.qs = pack32(u16vec2(data_a_packed16[ib].qs[loadr * 2],
+                             data_a_packed16[ib].qs[loadr * 2 + 1]));
+    blk.d = data_a_packed16[ib].d;
+    blk.qh = pack32(u16vec2(data_a_packed16[ib].qh[0], data_a_packed16[ib].qh[1]));
+    return blk;
+}
+
+void block_a_to_shmem(block_a_prefetch blk, uint buf_ib, uint ks, uint loadr) {
+    uint32_t lo4 = blk.qs & 0x0F0F0F0F;
+    uint32_t hi4 = (blk.qs >> 4) & 0x0F0F0F0F;
+    lo4 |= ((blk.qh >> (4u * loadr       )) & 0xFu) * 0x02040810u & 0x10101010u;
+    hi4 |= ((blk.qh >> (4u * loadr + 16u )) & 0xFu) * 0x02040810u & 0x10101010u;
+    lo4 = ((lo4 | 0x80808080) - 0x10101010) ^ 0x80808080;
+    hi4 = ((hi4 | 0x80808080) - 0x10101010) ^ 0x80808080;
+    buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + loadr    ] = lo4;
+    buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + loadr + 4] = hi4;
+
+    if (loadr == 0) {
+        buf_a_d[ks * BM + buf_ib] = float(blk.d);
+    }
+}
+
+#elif defined(DATA_A_Q5_1)
+
+struct block_a_prefetch {
+    uint32_t qs;
+    f16vec2 dm;
+    uint32_t qh;
+};
+
+block_a_prefetch block_a_load(uint ib, uint loadr) {
+    block_a_prefetch blk;
+    blk.qs = data_a_packed32[ib].qs[loadr];
+    blk.dm = data_a_packed32[ib].dm;
+    blk.qh = data_a_packed32[ib].qh;
+    return blk;
+}
+
+void block_a_to_shmem(block_a_prefetch blk, uint buf_ib, uint ks, uint loadr) {
+    // Store raw unsigned 5-bit values; the -16 offset is absorbed by the min term.
+    uint32_t lo4 = blk.qs & 0x0F0F0F0F;
+    uint32_t hi4 = (blk.qs >> 4) & 0x0F0F0F0F;
+    lo4 |= ((blk.qh >> (4u * loadr       )) & 0xFu) * 0x02040810u & 0x10101010u;
+    hi4 |= ((blk.qh >> (4u * loadr + 16u )) & 0xFu) * 0x02040810u & 0x10101010u;
+    buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + loadr    ] = lo4;
+    buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + loadr + 4] = hi4;
+
+    if (loadr == 0) {
+        buf_a_dm[ks * BM + buf_ib] = vec2(float(blk.dm.x), float(blk.dm.y));
+    }
+}
+
+#elif defined(DATA_A_Q8_0)
+
+struct block_a_prefetch {
+    uint32_t qs;
+    float16_t d;
+};
+
+block_a_prefetch block_a_load(uint ib, uint loadr) {
+    block_a_prefetch blk;
+    blk.qs = pack32(u16vec2(data_a_packed16[ib].qs[loadr * 2],
+                             data_a_packed16[ib].qs[loadr * 2 + 1]));
+    blk.d = data_a_packed16[ib].d;
+    return blk;
+}
+
+void block_a_to_shmem(block_a_prefetch blk, uint buf_ib, uint ks, uint loadr) {
+    buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + loadr] = blk.qs;
+
+    if (loadr == 0) {
+        buf_a_d[ks * BM + buf_ib] = float(blk.d);
+    }
+}
+
+#elif defined(DATA_A_IQ4_NL)
+
+struct block_a_prefetch {
+    uint32_t qs;
+    float16_t d;
+};
+
+block_a_prefetch block_a_load(uint ib, uint loadr) {
+    block_a_prefetch blk;
+    blk.qs = pack32(u16vec2(data_a_packed16[ib].qs[loadr * 2],
+                             data_a_packed16[ib].qs[loadr * 2 + 1]));
+    blk.d = data_a_packed16[ib].d;
+    return blk;
+}
+
+void block_a_to_shmem(block_a_prefetch blk, uint buf_ib, uint ks, uint loadr) {
+    const u8vec4 lo_idx = unpack8(blk.qs & 0x0F0F0F0F);
+    const u8vec4 hi_idx = unpack8((blk.qs >> 4) & 0x0F0F0F0F);
+    buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + loadr    ] =
+        pack32(i8vec4(cm1_kvalues[lo_idx.x], cm1_kvalues[lo_idx.y],
+                      cm1_kvalues[lo_idx.z], cm1_kvalues[lo_idx.w]));
+    buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + loadr + 4] =
+        pack32(i8vec4(cm1_kvalues[hi_idx.x], cm1_kvalues[hi_idx.y],
+                      cm1_kvalues[hi_idx.z], cm1_kvalues[hi_idx.w]));
+
+    if (loadr == 0) {
+        buf_a_d[ks * BM + buf_ib] = float(blk.d);
+    }
+}
+
+#elif defined(DATA_A_IQ4_XS)
+
+struct block_a_prefetch {
+    uint32_t qs;
+    float d;
+};
+
+block_a_prefetch block_a_load(uint ib, uint loadr) {
+    block_a_prefetch blk;
+    const uint ib_k = ib / 8;
+    const uint ib32 = ib % 8;
+    blk.qs = data_a_packed32[ib_k].qs[4 * ib32 + loadr];
+    blk.d = 0.0;
+    if (loadr == 0) {
+        const uint sl = (data_a_packed32[ib_k].scales_l >> (4 * ib32)) & 0xF;
+        const uint sh = (data_a_packed32[ib_k].scales_h >> (2 * ib32)) & 3;
+        blk.d = float(data_a_packed32[ib_k].d) * float(int(sl | (sh << 4)) - 32);
+    }
+    return blk;
+}
+
+void block_a_to_shmem(block_a_prefetch blk, uint buf_ib, uint ks, uint loadr) {
+    const u8vec4 lo_idx = unpack8(blk.qs & 0x0F0F0F0F);
+    const u8vec4 hi_idx = unpack8((blk.qs >> 4) & 0x0F0F0F0F);
+    buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + loadr    ] =
+        pack32(i8vec4(cm1_kvalues[lo_idx.x], cm1_kvalues[lo_idx.y],
+                      cm1_kvalues[lo_idx.z], cm1_kvalues[lo_idx.w]));
+    buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + loadr + 4] =
+        pack32(i8vec4(cm1_kvalues[hi_idx.x], cm1_kvalues[hi_idx.y],
+                      cm1_kvalues[hi_idx.z], cm1_kvalues[hi_idx.w]));
+
+    if (loadr == 0) {
+        buf_a_d[ks * BM + buf_ib] = blk.d;
+    }
+}
+
+#elif defined(DATA_A_MXFP4)
+
+struct block_a_prefetch {
+    uint32_t qs;
+    uint8_t e;
+};
+
+block_a_prefetch block_a_load(uint ib, uint loadr) {
+    block_a_prefetch blk;
+    blk.qs = pack32(u8vec4(data_a[ib].qs[loadr * 4],
+                            data_a[ib].qs[loadr * 4 + 1],
+                            data_a[ib].qs[loadr * 4 + 2],
+                            data_a[ib].qs[loadr * 4 + 3]));
+    blk.e = data_a[ib].e;
+    return blk;
+}
+
+void block_a_to_shmem(block_a_prefetch blk, uint buf_ib, uint ks, uint loadr) {
+    const u8vec4 lo_idx = unpack8(blk.qs & 0x0F0F0F0F);
+    const u8vec4 hi_idx = unpack8((blk.qs >> 4) & 0x0F0F0F0F);
+    buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + loadr    ] =
+        pack32(i8vec4(cm1_kvalues[lo_idx.x], cm1_kvalues[lo_idx.y],
+                      cm1_kvalues[lo_idx.z], cm1_kvalues[lo_idx.w]));
+    buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + loadr + 4] =
+        pack32(i8vec4(cm1_kvalues[hi_idx.x], cm1_kvalues[hi_idx.y],
+                      cm1_kvalues[hi_idx.z], cm1_kvalues[hi_idx.w]));
+
+    if (loadr == 0) {
+        buf_a_d[ks * BM + buf_ib] = e8m0_to_fp32(blk.e) * 0.5;
+    }
+}
+
+// LOAD_VEC_A=8 for k-quants and NVFP4: loadr has 4 positions, each writes 2 uint32
+
+#elif defined(DATA_A_Q4_K)
+
+struct block_a_prefetch {
+    uint32_t qs0;
+    uint32_t qs1;
+    uint ib;
+};
+
+block_a_prefetch block_a_load(uint ib, uint loadr) {
+    block_a_prefetch blk;
+    const uint ib_k = ib / 8;
+    const uint sub = ib % 8;
+    const uint qs_base = (sub >> 1) * 8;
+
+    uint32_t raw0 = data_a_packed32[ib_k].qs[qs_base + loadr * 2];
+    uint32_t raw1 = data_a_packed32[ib_k].qs[qs_base + loadr * 2 + 1];
+    if ((sub & 1u) != 0u) {
+        blk.qs0 = (raw0 >> 4) & 0x0F0F0F0F;
+        blk.qs1 = (raw1 >> 4) & 0x0F0F0F0F;
+    } else {
+        blk.qs0 = raw0 & 0x0F0F0F0F;
+        blk.qs1 = raw1 & 0x0F0F0F0F;
+    }
+    blk.ib = ib;
+
+    return blk;
+}
+
+void block_a_to_shmem(block_a_prefetch blk, uint buf_ib, uint ks, uint loadr) {
+    // Store raw unsigned nibbles (blk.qs already masked); no -8 recentering needed.
+    buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + loadr * 2    ] = blk.qs0;
+    buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + loadr * 2 + 1] = blk.qs1;
+
+    if (loadr == 0) {
+        const uint ib_k = blk.ib / 8;
+        const uint sub = blk.ib % 8;
+        const uint j = sub & 3u;
+        const uint s_j  = uint(data_a[ib_k].scales[j]);
+        const uint s_j4 = uint(data_a[ib_k].scales[j + 4]);
+        const uint s_j8 = uint(data_a[ib_k].scales[j + 8]);
+        const uint sc_val = (sub < 4) ? (s_j  & 0x3Fu) : ((s_j8 & 0x0Fu) | ((s_j  >> 6) << 4));
+        const uint mn_val = (sub < 4) ? (s_j4 & 0x3Fu) : ((s_j8 >> 4)    | ((s_j4 >> 6) << 4));
+        vec2 dm = vec2(data_a_packed32[ib_k].dm);
+        float d_scaled = dm.x * float(sc_val);
+        buf_a_dm[ks * BM + buf_ib] = vec2(d_scaled, -(dm.y * float(mn_val)));
+    }
+}
+
+#elif defined(DATA_A_Q5_K)
+
+struct block_a_prefetch {
+    uint32_t qs0;
+    uint32_t qs1;
+    uint32_t qh0;
+    uint32_t qh1;
+    uint ib;
+};
+
+block_a_prefetch block_a_load(uint ib, uint loadr) {
+    block_a_prefetch blk;
+    const uint ib_k = ib / 8;
+    const uint sub = ib % 8;
+    const uint qs_base = (sub >> 1) * 8;
+
+    uint32_t raw0 = data_a_packed32[ib_k].qs[qs_base + loadr * 2];
+    uint32_t raw1 = data_a_packed32[ib_k].qs[qs_base + loadr * 2 + 1];
+    if ((sub & 1u) != 0u) {
+        blk.qs0 = (raw0 >> 4) & 0x0F0F0F0F;
+        blk.qs1 = (raw1 >> 4) & 0x0F0F0F0F;
+    } else {
+        blk.qs0 = raw0 & 0x0F0F0F0F;
+        blk.qs1 = raw1 & 0x0F0F0F0F;
+    }
+    blk.qh0 = ((data_a_packed32[ib_k].qh[loadr * 2    ] >> sub) & 0x01010101) << 4;
+    blk.qh1 = ((data_a_packed32[ib_k].qh[loadr * 2 + 1] >> sub) & 0x01010101) << 4;
+    blk.ib = ib;
+
+    return blk;
+}
+
+void block_a_to_shmem(block_a_prefetch blk, uint buf_ib, uint ks, uint loadr) {
+    // Store raw unsigned 5-bit values (qs nibble | qh bit); no -16 recentering needed.
+    uint32_t v0 = blk.qs0 | blk.qh0;
+    uint32_t v1 = blk.qs1 | blk.qh1;
+    buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + loadr * 2    ] = v0;
+    buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + loadr * 2 + 1] = v1;
+
+    if (loadr == 0) {
+        const uint ib_k = blk.ib / 8;
+        const uint sub = blk.ib % 8;
+        const uint j = sub & 3u;
+        const uint s_j  = uint(data_a[ib_k].scales[j]);
+        const uint s_j4 = uint(data_a[ib_k].scales[j + 4]);
+        const uint s_j8 = uint(data_a[ib_k].scales[j + 8]);
+        const uint sc_val = (sub < 4) ? (s_j  & 0x3Fu) : ((s_j8 & 0x0Fu) | ((s_j  >> 6) << 4));
+        const uint mn_val = (sub < 4) ? (s_j4 & 0x3Fu) : ((s_j8 >> 4)    | ((s_j4 >> 6) << 4));
+        vec2 dm = vec2(data_a_packed32[ib_k].dm);
+        float d_scaled = dm.x * float(sc_val);
+        buf_a_dm[ks * BM + buf_ib] = vec2(d_scaled, -(dm.y * float(mn_val)));
+    }
+}
+
+#elif defined(DATA_A_Q6_K)
+
+struct block_a_prefetch {
+    uint32_t qs0;
+    uint32_t qs1;
+    uint ib;
+};
+
+block_a_prefetch block_a_load(uint ib, uint loadr) {
+    block_a_prefetch blk;
+    const uint ib_k = ib / 8;
+    const uint sub = ib % 8;
+    const uint g = sub / 4;
+    const uint j = sub % 4;
+
+    const uint ql_u16 = g * 32 + (j & 1) * 16 + loadr * 4;
+    const uint qh_u16 = g * 16 + loadr * 4;
+    const uint qh_shift = j * 2;
+
+    uint32_t ql0 = pack32(u16vec2(data_a_packed16[ib_k].ql[ql_u16    ],
+                                   data_a_packed16[ib_k].ql[ql_u16 + 1]));
+    uint32_t ql1 = pack32(u16vec2(data_a_packed16[ib_k].ql[ql_u16 + 2],
+                                   data_a_packed16[ib_k].ql[ql_u16 + 3]));
+    if (j >= 2) {
+        ql0 = (ql0 >> 4) & 0x0F0F0F0F;
+        ql1 = (ql1 >> 4) & 0x0F0F0F0F;
+    } else {
+        ql0 = ql0 & 0x0F0F0F0F;
+        ql1 = ql1 & 0x0F0F0F0F;
+    }
+
+    uint32_t qh0 = pack32(u16vec2(data_a_packed16[ib_k].qh[qh_u16    ],
+                                   data_a_packed16[ib_k].qh[qh_u16 + 1]));
+    uint32_t qh1 = pack32(u16vec2(data_a_packed16[ib_k].qh[qh_u16 + 2],
+                                   data_a_packed16[ib_k].qh[qh_u16 + 3]));
+
+    blk.qs0 = ql0 | (((qh0 >> qh_shift) & 0x03030303) << 4);
+    blk.qs1 = ql1 | (((qh1 >> qh_shift) & 0x03030303) << 4);
+    blk.ib = ib;
+
+    return blk;
+}
+
+void block_a_to_shmem(block_a_prefetch blk, uint buf_ib, uint ks, uint loadr) {
+    uint32_t v0 = ((blk.qs0 | 0x80808080) - 0x20202020) ^ 0x80808080;
+    uint32_t v1 = ((blk.qs1 | 0x80808080) - 0x20202020) ^ 0x80808080;
+    buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + loadr * 2    ] = v0;
+    buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + loadr * 2 + 1] = v1;
+
+    if (loadr == 0) {
+        const uint ib_k = blk.ib / 8;
+        const uint sub = blk.ib % 8;
+        i8vec2 sc = unpack8(int32_t(int16_t(data_a_packed16[ib_k].scales[sub]))).xy;
+        buf_a_d[(ks * KSCALES    ) * BM + buf_ib] = float(data_a_packed16[ib_k].d) * float(sc.x);
+        buf_a_d[(ks * KSCALES + 1) * BM + buf_ib] = float(data_a_packed16[ib_k].d) * float(sc.y);
+    }
+}
+
+#elif defined(DATA_A_Q3_K)
+
+struct block_a_prefetch {
+    uint32_t qs0;
+    uint32_t qs1;
+    uint ib;
+};
+
+block_a_prefetch block_a_load(uint ib, uint loadr) {
+    block_a_prefetch blk;
+    const uint ib_k = ib / 8;
+    const uint sub = ib % 8;
+    const uint g = sub / 4;
+    const uint j = sub % 4;
+    const uint qs_shift = j * 2;
+    const uint hm_bit = j + g * 4;
+
+    const uint qs_u16 = g * 16 + loadr * 4;
+    uint32_t qs0 = pack32(u16vec2(data_a_packed16[ib_k].qs[qs_u16    ],
+                                   data_a_packed16[ib_k].qs[qs_u16 + 1]));
+    uint32_t qs1 = pack32(u16vec2(data_a_packed16[ib_k].qs[qs_u16 + 2],
+                                   data_a_packed16[ib_k].qs[qs_u16 + 3]));
+
+    const uint hm_u16 = loadr * 4;
+    uint32_t hm0 = pack32(u16vec2(data_a_packed16[ib_k].hmask[hm_u16    ],
+                                   data_a_packed16[ib_k].hmask[hm_u16 + 1]));
+    uint32_t hm1 = pack32(u16vec2(data_a_packed16[ib_k].hmask[hm_u16 + 2],
+                                   data_a_packed16[ib_k].hmask[hm_u16 + 3]));
+
+    blk.qs0 = ((qs0 >> qs_shift) & 0x03030303) | (((hm0 >> hm_bit) & 0x01010101) << 2);
+    blk.qs1 = ((qs1 >> qs_shift) & 0x03030303) | (((hm1 >> hm_bit) & 0x01010101) << 2);
+    blk.ib = ib;
+
+    return blk;
+}
+
+void block_a_to_shmem(block_a_prefetch blk, uint buf_ib, uint ks, uint loadr) {
+    uint32_t v0 = ((blk.qs0 | 0x80808080) - 0x04040404) ^ 0x80808080;
+    uint32_t v1 = ((blk.qs1 | 0x80808080) - 0x04040404) ^ 0x80808080;
+    buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + loadr * 2    ] = v0;
+    buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + loadr * 2 + 1] = v1;
+
+    if (loadr == 0) {
+        const uint ib_k = blk.ib / 8;
+        const uint sub = blk.ib % 8;
+        const uint is = sub * 2;
+        uint lo = uint(data_a_packed16[ib_k].scales[(is % 8) / 2]);
+        lo = (lo >> (4 * (is / 8))) & 0x0F0Fu;
+        uint hi = uint(data_a_packed16[ib_k].scales[(8 + (is % 4)) / 2]);
+        hi = (hi >> (2 * (is / 4))) & 0x0303u;
+        uint combined = lo | (hi << 4);
+        i8vec2 sc = unpack8(int32_t(combined)).xy;
+        float d = float(data_a_packed16[ib_k].d);
+        buf_a_d[(ks * KSCALES    ) * BM + buf_ib] = d * float(int(sc.x) - 32);
+        buf_a_d[(ks * KSCALES + 1) * BM + buf_ib] = d * float(int(sc.y) - 32);
+    }
+}
+
+#elif defined(DATA_A_NVFP4)
+
+struct block_a_prefetch {
+    uint32_t qs;
+    uint8_t d0;
+    uint8_t d1;
+};
+
+block_a_prefetch block_a_load(uint ib, uint loadr) {
+    block_a_prefetch blk;
+    const uint ib_k = ib / 2;
+    const uint ihalf = ib % 2;
+    const uint sub = ihalf * 2 + (loadr >> 1);
+    const uint byte_group = loadr & 1u;
+
+    blk.qs = pack32(u8vec4(data_a[ib_k].qs[sub * 8 + byte_group * 4],
+                            data_a[ib_k].qs[sub * 8 + byte_group * 4 + 1],
+                            data_a[ib_k].qs[sub * 8 + byte_group * 4 + 2],
+                            data_a[ib_k].qs[sub * 8 + byte_group * 4 + 3]));
+    blk.d0 = data_a[ib_k].d[ihalf * 2];
+    blk.d1 = data_a[ib_k].d[ihalf * 2 + 1];
+
+    return blk;
+}
+
+void block_a_to_shmem(block_a_prefetch blk, uint buf_ib, uint ks, uint loadr) {
+    const u8vec4 lo_idx = unpack8(blk.qs & 0x0F0F0F0F);
+    const u8vec4 hi_idx = unpack8((blk.qs >> 4) & 0x0F0F0F0F);
+    const uint sub_base = (loadr >> 1) * 4;
+    const uint byte_group = loadr & 1u;
+    buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + sub_base + byte_group] =
+        pack32(i8vec4(cm1_kvalues[lo_idx.x], cm1_kvalues[lo_idx.y],
+                      cm1_kvalues[lo_idx.z], cm1_kvalues[lo_idx.w]));
+    buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + sub_base + 2 + byte_group] =
+        pack32(i8vec4(cm1_kvalues[hi_idx.x], cm1_kvalues[hi_idx.y],
+                      cm1_kvalues[hi_idx.z], cm1_kvalues[hi_idx.w]));
+
+    if (loadr == 0) {
+        buf_a_d[(ks * KSCALES    ) * BM + buf_ib] = ue4m3_to_fp32(blk.d0) * 0.5;
+        buf_a_d[(ks * KSCALES + 1) * BM + buf_ib] = ue4m3_to_fp32(blk.d1) * 0.5;
+    }
+}
+
+#endif
+
+// ===== B-side: load and store =====
+
+struct block_b_prefetch {
+    ivec4 qs;
+    float16_t d;
+#if defined(DATA_A_Q4_1) || defined(DATA_A_Q5_1) || defined(DATA_A_Q4_K) || defined(DATA_A_Q5_K)
+    float16_t s;
+#endif
+};
+
+block_b_prefetch block_b_load(uint ib_outer, uint ib_inner, uint loadr) {
+    block_b_prefetch blk;
+    blk.qs = data_b[ib_outer].qs[ib_inner * 2 + loadr];
+    blk.d = data_b[ib_outer].ds[ib_inner].x;
+#if defined(DATA_A_Q4_1) || defined(DATA_A_Q5_1) || defined(DATA_A_Q4_K) || defined(DATA_A_Q5_K)
+    blk.s = data_b[ib_outer].ds[ib_inner].y;
+#endif
+    return blk;
+}
+
+void block_b_to_shmem(block_b_prefetch blk, uint buf_ib, uint ks, uint loadr, bool in_bounds) {
+    const ivec4 v = in_bounds ? blk.qs : ivec4(0);
+    const uint base = buf_ib * QPITCH + ks * (BK / 4) + loadr * 4;
+    buf_b_qs[base    ] = v.x;
+    buf_b_qs[base + 1] = v.y;
+    buf_b_qs[base + 2] = v.z;
+    buf_b_qs[base + 3] = v.w;
+    if (loadr == 0) {
+        buf_b_d[ks * BN + buf_ib] = in_bounds ? float(blk.d) : 0.0f;
+#if defined(DATA_A_Q4_1) || defined(DATA_A_Q5_1) || defined(DATA_A_Q4_K) || defined(DATA_A_Q5_K)
+        buf_b_s[ks * BN + buf_ib] = in_bounds ? float(blk.s) : 0.0f;
+#endif
+    }
+}
+
+// ===== Framework macros =====
+
+#ifdef MUL_MAT_ID
+#define B_IB_CALC                                                                               \
+            const u16vec2 row_idx = row_ids[buf_ib];                                            \
+            const uint ib = pos_b_ib + row_idx.y * p.batch_stride_b / BK                        \
+                          + (row_idx.x % p.ne11) * p.stride_b / BK;
+#else
+#define B_IB_CALC                                                                               \
+            const uint ib = pos_b_ib + buf_ib * p.stride_b / BK;
+#endif
+
+#define PREFETCH_BLOCK(blk)                                                                     \
+    [[unroll]] for (uint li = 0; li < A_LOADS; li++) {                                          \
+        const uint buf_ib = loadc_a + li * loadstride_a;                                        \
+        if (buf_ib < BM) {                                                                      \
+            const uint ib = pos_a_ib + buf_ib * p.stride_a / BK;                                \
+            [[unroll]] for (uint ks = 0; ks < BK_STEP; ks++) {                                  \
+                pre_a[li * BK_STEP + ks] = block_a_load(ib + ks, loadr_a);                      \
+            }                                                                                   \
+        }                                                                                       \
+    }                                                                                           \
+    [[unroll]] for (uint li = 0; li < B_LOADS; li++) {                                          \
+        const uint buf_ib = loadc_b + li * loadstride_b;                                        \
+        if (buf_ib < BN) {                                                                      \
+            B_IB_CALC                                                                           \
+            [[unroll]] for (uint ks = 0; ks < BK_STEP; ks++) {                                  \
+                const uint ib_k = ((blk) + ks * BK < end_k) ? (ib + ks) : ib;                   \
+                pre_b[li * BK_STEP + ks] = block_b_load(ib_k / 4, ib_k % 4, loadr_b);          \
+            }                                                                                   \
+        }                                                                                       \
+    }
+
+#define STORE_BLOCK_TO_LDS(blk)                                                                 \
+    [[unroll]] for (uint li = 0; li < A_LOADS; li++) {                                          \
+        const uint buf_ib = loadc_a + li * loadstride_a;                                        \
+        if (buf_ib < BM) {                                                                      \
+            [[unroll]] for (uint ks = 0; ks < BK_STEP; ks++) {                                  \
+                block_a_to_shmem(pre_a[li * BK_STEP + ks], buf_ib, ks, loadr_a);                \
+            }                                                                                   \
+        }                                                                                       \
+    }                                                                                           \
+    [[unroll]] for (uint li = 0; li < B_LOADS; li++) {                                          \
+        const uint buf_ib = loadc_b + li * loadstride_b;                                        \
+        if (buf_ib < BN) {                                                                      \
+            [[unroll]] for (uint ks = 0; ks < BK_STEP; ks++) {                                  \
+                const bool in_bounds = (blk) + ks * BK < end_k;                                 \
+                block_b_to_shmem(pre_b[li * BK_STEP + ks], buf_ib, ks, loadr_b, in_bounds);     \
+            }                                                                                   \
+        }                                                                                       \
+    }
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 12f9b3f56..cde5d36dc 100644
--- a/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp
+++ b/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp
@@ -480,8 +480,9 @@ void matmul_shaders(bool fp16, MatMulIdType matmul_id_type, bool coopmat, bool c
         base_dict["FLOAT16"] = "1";
     }

-    base_dict["ACC_TYPE"  ] = f16acc ? "float16_t" : "float";
-    base_dict["ACC_TYPEV2"] = f16acc ? "f16vec2"   : "vec2";
+    base_dict["ACC_TYPE"     ] = f16acc ? "float16_t" : "float";
+    base_dict["ACC_TYPEV2"   ] = f16acc ? "f16vec2"   : "vec2";
+    base_dict["ACC_TYPE_VEC4"] = f16acc ? "f16vec4"   : "vec4";
     if (f16acc) {
         base_dict["ACC_TYPE_MAX"] = "float16_t(65504.0)";
     }
@@ -629,6 +630,11 @@ void matmul_shaders(bool fp16, MatMulIdType matmul_id_type, bool coopmat, bool c
         }
 #endif

+        if (!f16acc && coopmat && (tname == "q4_0" || tname == "q4_1" || tname == "q5_0" || tname == "q5_1" || tname == "q8_0" || tname == "iq4_nl" || tname == "iq4_xs" || tname == "mxfp4"
+                     || tname == "q3_k" || tname == "q4_k" || tname == "q5_k" || tname == "q6_k" || tname == "nvfp4")) {
+            string_to_spv(shader_name + "_" + tname + "_q8_1", "mul_mmq_cm1.comp", merge_maps(merge_maps(base_dict, float_type_dict), {{data_a_key, "1"}, {"D_TYPE", "float"}, {"D_TYPE_VEC4", "vec4"}}), fp16, coopmat, coopmat2, f16acc);
+        }
+
         if (is_lut_quant(tname)) {
             std::string lva = lut_load_vec_a(tname);