Commit 8216c8462 for llama.cpp

commit 8216c84623cf5b22b29319344ca5e685003892f5
Author: Masashi Yoshimura <yoshimura.masashi.frbs@gmail.com>
Date:   Mon Oct 5 15:56:45 2026 +0900

    webgpu: add MMVQ support for Q1_0/Q5_0/Q5_1/Q3_K/Q5_K/Q6_K/MXFP4 (#29483)

    * add supports for q1/q5/q3_k/q5_k/q6_k/mxfp4 of mmvq path

    * Add K_QUANTS_HANDLING macro to q1_0 of mmvq path

diff --git a/ggml/src/ggml-webgpu/ggml-webgpu-shader-lib.hpp b/ggml/src/ggml-webgpu/ggml-webgpu-shader-lib.hpp
index d4cc0258c..65556f352 100644
--- a/ggml/src/ggml-webgpu/ggml-webgpu-shader-lib.hpp
+++ b/ggml/src/ggml-webgpu/ggml-webgpu-shader-lib.hpp
@@ -1181,11 +1181,18 @@ inline bool ggml_webgpu_can_use_mmvq(const ggml_tensor * src0,
             switch (src1->type) {
                 case GGML_TYPE_F32:
                     switch (src0->type) {
+                        case GGML_TYPE_Q1_0:
                         case GGML_TYPE_Q4_0:
                         case GGML_TYPE_Q4_1:
+                        case GGML_TYPE_Q5_0:
+                        case GGML_TYPE_Q5_1:
                         case GGML_TYPE_Q8_0:
+                        case GGML_TYPE_MXFP4:
                         case GGML_TYPE_Q2_K:
+                        case GGML_TYPE_Q3_K:
                         case GGML_TYPE_Q4_K:
+                        case GGML_TYPE_Q5_K:
+                        case GGML_TYPE_Q6_K:
                             return src0->ne[0] % 4 == 0;
                         default:
                             break;
@@ -2036,17 +2043,23 @@ class ggml_webgpu_shader_lib {
                     defines.push_back("U32_DEQUANT_HELPERS");
                     defines.push_back("SRC0_INNER_TYPE=u32");
                     switch (context.src0->type) {
-                        case GGML_TYPE_Q8_0:
                         case GGML_TYPE_Q4_0:
                         case GGML_TYPE_Q4_1:
+                        case GGML_TYPE_Q5_0:
+                        case GGML_TYPE_Q5_1:
+                        case GGML_TYPE_Q8_0:
                             if (key.use_mmvq) {
-                                defines.push_back("LEGACY_QUANTS");
+                                defines.push_back("LEGACY_QUANTS_HANDLING");
                             }
                             break;
+                        case GGML_TYPE_Q1_0:
                         case GGML_TYPE_Q2_K:
+                        case GGML_TYPE_Q3_K:
                         case GGML_TYPE_Q4_K:
+                        case GGML_TYPE_Q5_K:
+                        case GGML_TYPE_Q6_K:
                             if (key.use_mmvq) {
-                                defines.push_back("K_QUANTS");
+                                defines.push_back("K_QUANTS_HANDLING");
                             }
                             break;
                         case GGML_TYPE_IQ1_S:
@@ -2064,6 +2077,11 @@ class ggml_webgpu_shader_lib {
                             defines.push_back(type_upper + "_TABLES");
                             break;
                         case GGML_TYPE_MXFP4:
+                            defines.push_back(type_upper + "_LUT");
+                            if (key.use_mmvq) {
+                                defines.push_back("LEGACY_QUANTS_HANDLING");
+                            }
+                            break;
                         case GGML_TYPE_NVFP4:
                             defines.push_back(type_upper + "_LUT");
                             break;
diff --git a/ggml/src/ggml-webgpu/wgsl-shaders/mul_mat_vec_q_acc.tmpl b/ggml/src/ggml-webgpu/wgsl-shaders/mul_mat_vec_q_acc.tmpl
index 6ccaf61a6..3dcf7faee 100644
--- a/ggml/src/ggml-webgpu/wgsl-shaders/mul_mat_vec_q_acc.tmpl
+++ b/ggml/src/ggml-webgpu/wgsl-shaders/mul_mat_vec_q_acc.tmpl
@@ -14,10 +14,13 @@ fn sbyte_of(v: u32, b: u32) -> i32 {
 #define SRC0_TYPE SRC0_INNER_TYPE
 #define SRC1_TYPE SRC1_INNER_TYPE

-#ifdef LEGACY_QUANTS
+#ifdef LEGACY_QUANTS_HANDLING
 #define BLOCK_SIZE 32
 #define THREADS_PER_BLOCK 4
-#elif K_QUANTS
+#elif defined(MUL_ACC_Q1_0)
+#define BLOCK_SIZE 128
+#define THREADS_PER_BLOCK 8
+#elif defined(K_QUANTS_HANDLING)
 #define BLOCK_SIZE 256
 #define THREADS_PER_BLOCK 16
 #endif
@@ -25,6 +28,22 @@ fn sbyte_of(v: u32, b: u32) -> i32 {
 #define ELEMS_PER_THREAD (BLOCK_SIZE/THREADS_PER_BLOCK)
 #define Q8_BLOCK_SIZE 32

+#if (defined(LEGACY_QUANTS_HANDLING) || defined(MUL_ACC_MXFP4)) && !defined(MUL_ACC_Q8_0)
+fn repack_b_qs(block:u32, inner_id: u32) -> vec2<u32> {
+    return vec2<u32>(
+            src1q[block].qs[inner_id],
+            src1q[block].qs[inner_id + 4u],
+        );
+}
+#endif
+
+#if defined(MUL_ACC_Q5_0) || defined(MUL_ACC_Q5_1)
+fn qh_bits(qh: u32, shift: u32) -> u32 {
+    // multiply by 0x00204081 moves bit b to bit 8*b
+    return ((((qh >> shift) & 0xFu) * 0x00204081u) & 0x01010101u) << 4u;
+}
+#endif
+
 #ifdef MUL_ACC_Q4_0
 #define BLOCK_SIZE_BYTES 18
 #define B_DS_TYPE vec2<f32>
@@ -36,12 +55,6 @@ fn repack_a(block_byte_base: u32, inner_id: u32) -> vec2<u32> {
         (qs_packed >> 4u) & 0x0F0F0F0Fu
     );
 }
-fn repack_b_qs(block:u32, inner_id: u32) -> vec2<u32> {
-    return vec2<u32>(
-            src1q[block].qs[inner_id],
-            src1q[block].qs[inner_id + 4u],
-        );
-}
 fn repack_b_dm(block: u32) -> B_DS_TYPE {
     return B_DS_TYPE(
         f32(src1q[block].d),
@@ -64,11 +77,54 @@ fn repack_a(block_byte_base: u32, inner_id: u32) -> vec2<u32> {
         (qs_packed >> 4u) & 0x0F0F0F0Fu
     );
 }
-fn repack_b_qs(block:u32, inner_id: u32) -> vec2<u32> {
+fn repack_b_dm(block: u32) -> B_DS_TYPE {
+    return B_DS_TYPE(
+        f32(src1q[block].d),
+        f32(src1q[block].s)
+    );
+}
+fn get_dm(block_byte_base: u32) -> vec2<f32> {
+    return vec2<f32>(
+        f32(load_f16_at_src0(block_byte_base)),
+        f32(load_f16_at_src0(block_byte_base + 2u))
+    );
+}
+#endif // MUL_ACC_Q4_1
+
+#ifdef MUL_ACC_Q5_0
+#define BLOCK_SIZE_BYTES 22
+#define B_DS_TYPE vec2<f32>
+fn repack_a(block_byte_base: u32, inner_id: u32) -> vec2<u32> {
+    let qh        = load_u32_at_src0(block_byte_base + 2u);
+    let qs_packed = load_u32_at_src0(block_byte_base + 6u + 4u * inner_id);
+
     return vec2<u32>(
-            src1q[block].qs[inner_id],
-            src1q[block].qs[inner_id + 4u],
-        );
+        (qs_packed & 0x0F0F0F0Fu) | qh_bits(qh, 4u * inner_id),
+        ((qs_packed >> 4u) & 0x0F0F0F0Fu) | qh_bits(qh, 16u + 4u * inner_id)
+    );
+}
+fn repack_b_dm(block: u32) -> B_DS_TYPE {
+    return B_DS_TYPE(
+        f32(src1q[block].d),
+        f32(src1q[block].s)
+    );
+}
+fn get_dm(block_byte_base: u32) -> f32 {
+    return f32(load_f16_at_src0(block_byte_base));
+}
+#endif // MUL_ACC_Q5_0
+
+#ifdef MUL_ACC_Q5_1
+#define BLOCK_SIZE_BYTES 24
+#define B_DS_TYPE vec2<f32>
+fn repack_a(block_byte_base: u32, inner_id: u32) -> vec2<u32> {
+    let qh        = load_u32_at_src0(block_byte_base + 4u);
+    let qs_packed = load_u32_at_src0(block_byte_base + 8u + 4u * inner_id);
+
+    return vec2<u32>(
+        (qs_packed & 0x0F0F0F0Fu) | qh_bits(qh, 4u * inner_id),
+        ((qs_packed >> 4u) & 0x0F0F0F0Fu) | qh_bits(qh, 16u + 4u * inner_id)
+    );
 }
 fn repack_b_dm(block: u32) -> B_DS_TYPE {
     return B_DS_TYPE(
@@ -82,7 +138,30 @@ fn get_dm(block_byte_base: u32) -> vec2<f32> {
         f32(load_f16_at_src0(block_byte_base + 2u))
     );
 }
-#endif // MUL_ACC_Q4_1
+#endif // MUL_ACC_Q5_1
+
+#ifdef MUL_ACC_MXFP4
+#define BLOCK_SIZE_BYTES 17
+#define B_DS_TYPE f32
+fn repack_a(block_byte_base: u32, inner_id: u32) -> vec2<u32> {
+    let qs_packed = load_u32_at_src0(block_byte_base + 1u + 4u * inner_id);
+
+    var lo = 0u;
+    var hi = 0u;
+    for (var b = 0u; b < 4u; b++) {
+        let q_byte = byte_of(qs_packed, b);
+        lo |= (bitcast<u32>(kvalues_mxfp4[q_byte & 0xFu]) & 0xFFu) << (8u * b);
+        hi |= (bitcast<u32>(kvalues_mxfp4[q_byte >> 4u]) & 0xFFu) << (8u * b);
+    }
+    return vec2<u32>(lo, hi);
+}
+fn repack_b_dm(block: u32) -> B_DS_TYPE {
+    return B_DS_TYPE(src1q[block].d);
+}
+fn get_dm(block_byte_base: u32) -> f32 {
+    return ldexp(1.0, i32(byte_of(load_u32_at_src0(block_byte_base), 0u)) - 128);
+}
+#endif // MUL_ACC_MXFP4

 #ifdef MUL_ACC_Q8_0
 #define BLOCK_SIZE_BYTES 34
@@ -107,7 +186,7 @@ fn get_dm(block_byte_base: u32) -> f32 {
 }
 #endif // MUL_ACC_Q8_0

-#if defined(LEGACY_QUANTS)
+#if defined(LEGACY_QUANTS_HANDLING)
 fn accumulate_vec_q_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1q_idx_base: u32) -> array<array<f32, OUTPUTS_PER_WG>, NUM_COLS> {
     var acc: array<array<f32, OUTPUTS_PER_WG>, NUM_COLS>;

@@ -115,6 +194,13 @@ fn accumulate_vec_q_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, s

     for (var block = thread_id / THREADS_PER_BLOCK; block < num_blocks; block += WG_SIZE / THREADS_PER_BLOCK) {
         let inner_id = thread_id % THREADS_PER_BLOCK;
+        var b_qs_cols: array<vec2<u32>, NUM_COLS>;
+        var b_ds_cols: array<B_DS_TYPE, NUM_COLS>;
+        for (var col = 0u;col < NUM_COLS;col += 1) {
+            let src1q_idx = src1q_idx_base + col * (params.k / Q8_BLOCK_SIZE) + block;
+            b_qs_cols[col] = repack_b_qs(src1q_idx, inner_id);
+            b_ds_cols[col] = repack_b_dm(src1q_idx);
+        }
         for (var row = 0u; row < OUTPUTS_PER_WG; row++) {
             let output_row = row_base + row;
             if (output_row < params.m) {
@@ -122,9 +208,8 @@ fn accumulate_vec_q_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, s
                 let a_repacked = repack_a(block_byte_base, inner_id);
                 let da = get_dm(block_byte_base);
                 for (var col = 0u;col < NUM_COLS;col += 1) {
-                    let src1q_idx = src1q_idx_base + col * (params.k / Q8_BLOCK_SIZE) + block;
-                    let b_repacked = repack_b_qs(src1q_idx, inner_id);
-                    let b_ds = repack_b_dm(src1q_idx);
+                    let b_repacked = b_qs_cols[col];
+                    let b_ds = b_ds_cols[col];

                     let row_sum = dot4I8Packed(a_repacked[0], b_repacked[0]) + dot4I8Packed(a_repacked[1], b_repacked[1]);

@@ -132,13 +217,17 @@ fn accumulate_vec_q_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, s
                     acc[col][row] += f32(row_sum) * (da * b_ds.x) - 8.0 * da * b_ds.y / THREADS_PER_BLOCK;
 #endif // MUL_ACC_Q4_0

-#if defined(MUL_ACC_Q4_1)
+#if defined(MUL_ACC_Q5_0)
+                    acc[col][row] += f32(row_sum) * (da * b_ds.x) - 16.0 * da * b_ds.y / THREADS_PER_BLOCK;
+#endif // MUL_ACC_Q5_0
+
+#if defined(MUL_ACC_Q4_1) || defined(MUL_ACC_Q5_1)
                     acc[col][row] += f32(row_sum) * (da.x * b_ds.x) + da.y * b_ds.y / THREADS_PER_BLOCK;
-#endif // MUL_ACC_Q4_1
+#endif // MUL_ACC_Q4_1 || MUL_ACC_Q5_1

-#if defined(MUL_ACC_Q8_0)
+#if defined(MUL_ACC_Q8_0) || defined(MUL_ACC_MXFP4)
                     acc[col][row] += f32(row_sum) * (da * b_ds);
-#endif // MUL_ACC_Q8_0
+#endif // MUL_ACC_Q8_0 || MUL_ACC_MXFP4
                 }
             }
         }
@@ -146,7 +235,49 @@ fn accumulate_vec_q_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, s

     return acc;
 }
-#endif // LEGACY_QUANTS
+#endif // LEGACY_QUANTS_HANDLING
+
+// every k-quant thread covers 16 elements
+#if defined(K_QUANTS_HANDLING)
+fn repack_b_qs(q8_block_idx: u32, tid: u32) -> vec4<u32> {
+    let phase = tid % 2u;
+    return vec4<u32>(
+        src1q[q8_block_idx].qs[4u * phase],
+        src1q[q8_block_idx].qs[4u * phase + 1u],
+        src1q[q8_block_idx].qs[4u * phase + 2u],
+        src1q[q8_block_idx].qs[4u * phase + 3u],
+    );
+}
+#endif
+
+#if defined(MUL_ACC_Q1_0) || defined(MUL_ACC_Q3_K) || defined(MUL_ACC_Q6_K)
+// subtract c from every byte of v, giving 4 packed i8; each byte of v must be below 128
+fn sub_packed_bytes(v: u32, c: u32) -> u32 {
+    return (v + (0x80u - c) * 0x01010101u) ^ 0x80808080u;
+}
+#endif
+
+#ifdef MUL_ACC_Q1_0
+#define BLOCK_SIZE_BYTES 18
+#define B_DS_TYPE f32
+fn repack_a(block_byte_base: u32, tid: u32) -> vec4<u32> {
+    let bits = load_u16_at_src0(block_byte_base + 2u + 2u * tid);
+
+    var res: vec4<u32>;
+    for (var i = 0u; i < 4u; i++) {
+        // spread 4 bits to bit 1 of each byte
+        let twice_bits = ((((bits >> (4u * i)) & 0xFu) * 0x00204081u) & 0x01010101u) << 1u;
+        res[i] = sub_packed_bytes(twice_bits, 1u);
+    }
+    return res;
+}
+fn repack_b_dm(q8_block_idx: u32) -> B_DS_TYPE {
+    return B_DS_TYPE(src1q[q8_block_idx].d);
+}
+fn get_dm(block_byte_base: u32) -> f32 {
+    return f32(load_f16_at_src0(block_byte_base));
+}
+#endif // MUL_ACC_Q1_0

 #ifdef MUL_ACC_Q2_K
 #define BLOCK_SIZE_BYTES 84
@@ -164,15 +295,6 @@ fn repack_a(block_byte_base: u32, tid: u32) -> vec4<u32> {
         (load_u32_at_src0_aligned(qs_byte_base + 12u) >> qs_shift) & 0x03030303u,
     );
 }
-fn repack_b_qs(q8_block_idx: u32, tid: u32) -> vec4<u32> {
-    let phase = tid % 2u;
-    return vec4<u32>(
-        src1q[q8_block_idx].qs[4u * phase],
-        src1q[q8_block_idx].qs[4u * phase + 1u],
-        src1q[q8_block_idx].qs[4u * phase + 2u],
-        src1q[q8_block_idx].qs[4u * phase + 3u],
-    );
-}
 fn repack_b_dm(q8_block_idx: u32) -> B_DS_TYPE {
     return B_DS_TYPE(src1q[q8_block_idx].d);
 }
@@ -189,31 +311,51 @@ fn get_scale_min(block_byte_base: u32, tid: u32) -> vec2<f32> {
 }
 #endif // MUL_ACC_Q2_K

-#ifdef MUL_ACC_Q4_K
-#define BLOCK_SIZE_BYTES 144
-#define B_DS_TYPE vec2<f32>
+#ifdef MUL_ACC_Q3_K
+#define BLOCK_SIZE_BYTES 110
+#define B_DS_TYPE f32
 fn repack_a(block_byte_base: u32, tid: u32) -> vec4<u32> {
-    let iq4 = tid / 4u;
-    let phase = tid % 2u;
-    let nibble = (tid >> 1u) % 2u;
-    let q_qs_byte_base = block_byte_base + 16u + 32u * iq4 + 16u * phase;
-    let qs_shift = 4u * nibble;
-    return vec4<u32>(
-        (load_u32_at_src0_aligned(q_qs_byte_base) >> qs_shift) & 0x0F0F0F0Fu,
-        (load_u32_at_src0_aligned(q_qs_byte_base + 4u) >> qs_shift) & 0x0F0F0F0Fu,
-        (load_u32_at_src0_aligned(q_qs_byte_base + 8u) >> qs_shift) & 0x0F0F0F0Fu,
-        (load_u32_at_src0_aligned(q_qs_byte_base + 12u) >> qs_shift) & 0x0F0F0F0Fu,
-    );
+    let half_blk = tid / 8u;
+    let sub      = (tid % 8u) / 2u;
+    let phase    = tid % 2u;
+    let qs_byte_base = block_byte_base + 32u + 32u * half_blk + 16u * phase;
+    let hm_byte_base = block_byte_base + 16u * phase;
+    let qs_shift = 2u * sub;
+    let hm_shift = 4u * half_blk + sub;
+
+    var res: vec4<u32>;
+    for (var i = 0u; i < 4u; i++) {
+        let qs = (load_u32_at_src0(qs_byte_base + 4u * i) >> qs_shift) & 0x03030303u;
+        let hm = (load_u32_at_src0(hm_byte_base + 4u * i) >> hm_shift) & 0x01010101u;
+        // the high bit is stored inverted: a clear bit means the value is 4 lower
+        res[i] = sub_packed_bytes(qs | (hm << 2u), 4u);
+    }
+    return res;
 }
-fn repack_b_qs(q8_block_idx: u32, tid: u32) -> vec4<u32> {
-    let phase = tid % 2u;
-    return vec4<u32>(
-        src1q[q8_block_idx].qs[4u * phase],
-        src1q[q8_block_idx].qs[4u * phase + 1u],
-        src1q[q8_block_idx].qs[4u * phase + 2u],
-        src1q[q8_block_idx].qs[4u * phase + 3u],
-    );
+fn repack_b_dm(q8_block_idx: u32) -> B_DS_TYPE {
+    return B_DS_TYPE(src1q[q8_block_idx].d);
+}
+fn get_dm(block_byte_base: u32) -> f32 {
+    return f32(load_f16_at_src0(block_byte_base + 108u));
 }
+fn get_scale_min(block_byte_base: u32, tid: u32) -> f32 {
+    let byte_idx = tid & 3u;
+    let group    = tid / 4u;
+
+    let scales_lo = load_u32_at_src0(block_byte_base + 96u + 4u * (group & 1u));
+    let scales_hi = load_u32_at_src0(block_byte_base + 104u);
+
+    let lo_byte = byte_of(scales_lo, byte_idx);
+    let lo      = select(lo_byte >> 4u, lo_byte & 0x0Fu, group < 2u);
+    let hi      = (byte_of(scales_hi, byte_idx) >> (2u * group)) & 3u;
+
+    return f32(i32(lo | (hi << 4u)) - 32);
+}
+#endif // MUL_ACC_Q3_K
+
+// Q4_K and Q5_K share the scale/min layout
+#if defined(MUL_ACC_Q4_K) || defined(MUL_ACC_Q5_K)
+#define B_DS_TYPE vec2<f32>
 fn repack_b_dm(q8_block_idx: u32) -> B_DS_TYPE {
     return B_DS_TYPE(
         f32(src1q[q8_block_idx].d),
@@ -246,26 +388,104 @@ fn get_scale_min(block_byte_base: u32, tid: u32) -> vec2<f32> {

     return vec2<f32>(scale, min_val);
 }
+#endif // MUL_ACC_Q4_K || MUL_ACC_Q5_K
+
+#ifdef MUL_ACC_Q4_K
+#define BLOCK_SIZE_BYTES 144
+fn repack_a(block_byte_base: u32, tid: u32) -> vec4<u32> {
+    let iq4 = tid / 4u;
+    let phase = tid % 2u;
+    let nibble = (tid >> 1u) % 2u;
+    let q_qs_byte_base = block_byte_base + 16u + 32u * iq4 + 16u * phase;
+    let qs_shift = 4u * nibble;
+    return vec4<u32>(
+        (load_u32_at_src0_aligned(q_qs_byte_base) >> qs_shift) & 0x0F0F0F0Fu,
+        (load_u32_at_src0_aligned(q_qs_byte_base + 4u) >> qs_shift) & 0x0F0F0F0Fu,
+        (load_u32_at_src0_aligned(q_qs_byte_base + 8u) >> qs_shift) & 0x0F0F0F0Fu,
+        (load_u32_at_src0_aligned(q_qs_byte_base + 12u) >> qs_shift) & 0x0F0F0F0Fu,
+    );
+}
 #endif // MUL_ACC_Q4_K

-#ifdef K_QUANTS
+#ifdef MUL_ACC_Q5_K
+#define BLOCK_SIZE_BYTES 176
+fn repack_a(block_byte_base: u32, tid: u32) -> vec4<u32> {
+    let iq4 = tid / 4u;
+    let phase = tid % 2u;
+    let nibble = (tid >> 1u) % 2u;
+    let ql_byte_base = block_byte_base + 48u + 32u * iq4 + 16u * phase;
+    let qh_byte_base = block_byte_base + 16u + 16u * phase;
+    let ql_shift = 4u * nibble;
+    let qh_shift = 2u * iq4 + nibble;
+
+    var res: vec4<u32>;
+    for (var i = 0u; i < 4u; i++) {
+        let ql = (load_u32_at_src0_aligned(ql_byte_base + 4u * i) >> ql_shift) & 0x0F0F0F0Fu;
+        let qh = (load_u32_at_src0_aligned(qh_byte_base + 4u * i) >> qh_shift) & 0x01010101u;
+        res[i] = ql | (qh << 4u);
+    }
+    return res;
+}
+#endif // MUL_ACC_Q5_K
+
+#ifdef MUL_ACC_Q6_K
+#define BLOCK_SIZE_BYTES 210
+#define B_DS_TYPE f32
+fn repack_a(block_byte_base: u32, tid: u32) -> vec4<u32> {
+    let half_blk = tid / 8u;
+    let sub      = (tid % 8u) / 2u;
+    let phase    = tid % 2u;
+    let ql_byte_base = block_byte_base + 64u * half_blk + 32u * (sub & 1u) + 16u * phase;
+    let qh_byte_base = block_byte_base + 128u + 32u * half_blk + 16u * phase;
+    let ql_shift = 4u * (sub >> 1u);
+    let qh_shift = 2u * sub;
+
+    var res: vec4<u32>;
+    for (var i = 0u; i < 4u; i++) {
+        let ql = (load_u32_at_src0(ql_byte_base + 4u * i) >> ql_shift) & 0x0F0F0F0Fu;
+        let qh = (load_u32_at_src0(qh_byte_base + 4u * i) >> qh_shift) & 0x03030303u;
+        res[i] = sub_packed_bytes(ql | (qh << 4u), 32u);
+    }
+    return res;
+}
+fn repack_b_dm(q8_block_idx: u32) -> B_DS_TYPE {
+    return B_DS_TYPE(src1q[q8_block_idx].d);
+}
+fn get_dm(block_byte_base: u32) -> f32 {
+    return f32(load_f16_at_src0(block_byte_base + 208u));
+}
+fn get_scale_min(block_byte_base: u32, tid: u32) -> f32 {
+    let scale_byte = block_byte_base + 192u + tid;
+    return f32(sbyte_of(load_u32_at_src0_aligned(scale_byte), scale_byte & 3u));
+}
+#endif // MUL_ACC_Q6_K
+
+#if defined(K_QUANTS_HANDLING)
 fn accumulate_vec_q_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1q_idx_base: u32) -> array<array<f32, OUTPUTS_PER_WG>, NUM_COLS> {
     var acc: array<array<f32, OUTPUTS_PER_WG>, NUM_COLS>;

     let tid = thread_id % THREADS_PER_BLOCK;

     for (var block = thread_id / THREADS_PER_BLOCK; block < params.k / BLOCK_SIZE; block += WG_SIZE / THREADS_PER_BLOCK) {
+        var b_qs_cols: array<vec4<u32>, NUM_COLS>;
+        var b_ds_cols: array<B_DS_TYPE, NUM_COLS>;
+        for (var col = 0u;col < NUM_COLS;col += 1) {
+            let src1q_idx = src1q_idx_base + col * (params.k / Q8_BLOCK_SIZE) + (block * BLOCK_SIZE + ELEMS_PER_THREAD * tid) / Q8_BLOCK_SIZE;
+            b_qs_cols[col] = repack_b_qs(src1q_idx, tid);
+            b_ds_cols[col] = repack_b_dm(src1q_idx);
+        }
         for (var row = 0u; row < OUTPUTS_PER_WG; row++) {
             let output_row = row_base + row;
             if (output_row < params.m) {
                 let block_byte_base = (src0_batch_offset + output_row * params.stride_01 + block) * BLOCK_SIZE_BYTES;
                 let a_repacked = repack_a(block_byte_base, tid);
                 let dm = get_dm(block_byte_base);
+#ifndef MUL_ACC_Q1_0
                 let scale_min = get_scale_min(block_byte_base, tid);
+#endif
                 for (var col = 0u;col < NUM_COLS;col += 1) {
-                    let src1q_idx = src1q_idx_base + col * (params.k / Q8_BLOCK_SIZE) + (block * BLOCK_SIZE + ELEMS_PER_THREAD * tid) / Q8_BLOCK_SIZE;
-                    let b_repacked = repack_b_qs(src1q_idx, tid);
-                    let b_ds = repack_b_dm(src1q_idx);
+                    let b_repacked = b_qs_cols[col];
+                    let b_ds = b_ds_cols[col];

 #if defined(MUL_ACC_Q2_K)
                     let scale_q = i32(scale_min.x);
@@ -279,13 +499,27 @@ fn accumulate_vec_q_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, s
                     acc[col][row] += b_ds * (dm.x * f32(row_sum_d) - dm.y * f32(row_sum_m));
 #endif // MUL_ACC_Q2_K

-#if defined(MUL_ACC_Q4_K)
+#if defined(MUL_ACC_Q4_K) || defined(MUL_ACC_Q5_K)
                     let row_sum = dot4I8Packed(a_repacked[0], b_repacked[0]) + dot4I8Packed(a_repacked[1], b_repacked[1])
                                     + dot4I8Packed(a_repacked[2], b_repacked[2]) + dot4I8Packed(a_repacked[3], b_repacked[3]);

                     // Each thread covers half of the Q8_1 block, so add only b_ds.y/2.
                     acc[col][row] += b_ds.x * dm.x * scale_min.x * f32(row_sum) - dm.y * scale_min.y * (b_ds.y / (Q8_BLOCK_SIZE / ELEMS_PER_THREAD));
-#endif // MUL_ACC_Q4_K
+#endif // MUL_ACC_Q4_K || MUL_ACC_Q5_K
+
+#if defined(MUL_ACC_Q3_K) || defined(MUL_ACC_Q6_K)
+                    let row_sum = dot4I8Packed(a_repacked[0], b_repacked[0]) + dot4I8Packed(a_repacked[1], b_repacked[1])
+                                    + dot4I8Packed(a_repacked[2], b_repacked[2]) + dot4I8Packed(a_repacked[3], b_repacked[3]);
+
+                    acc[col][row] += b_ds * dm * scale_min * f32(row_sum);
+#endif // MUL_ACC_Q3_K || MUL_ACC_Q6_K
+
+#if defined(MUL_ACC_Q1_0)
+                    let row_sum = dot4I8Packed(a_repacked[0], b_repacked[0]) + dot4I8Packed(a_repacked[1], b_repacked[1])
+                                    + dot4I8Packed(a_repacked[2], b_repacked[2]) + dot4I8Packed(a_repacked[3], b_repacked[3]);
+
+                    acc[col][row] += b_ds * dm * f32(row_sum);
+#endif // MUL_ACC_Q1_0

                 }
             }
@@ -294,4 +528,4 @@ fn accumulate_vec_q_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, s

     return acc;
 }
-#endif // K_QUANTS
+#endif // K_QUANTS_HANDLING
diff --git a/ggml/src/ggml-webgpu/wgsl-shaders/quantize_q8.wgsl b/ggml/src/ggml-webgpu/wgsl-shaders/quantize_q8.wgsl
index 847b27ffa..db8fafd41 100644
--- a/ggml/src/ggml-webgpu/wgsl-shaders/quantize_q8.wgsl
+++ b/ggml/src/ggml-webgpu/wgsl-shaders/quantize_q8.wgsl
@@ -34,7 +34,7 @@ fn cluster_max_8(v: f32) -> f32 {
     return r;
 }

-#if defined(MUL_ACC_Q4_0) || defined(MUL_ACC_Q4_1) || defined(MUL_ACC_Q4_K)
+#if defined(MUL_ACC_Q4_0) || defined(MUL_ACC_Q4_1) || defined(MUL_ACC_Q5_0) || defined(MUL_ACC_Q5_1) || defined(MUL_ACC_Q4_K) || defined(MUL_ACC_Q5_K)
 fn cluster_add_i4x8(v: i32) -> i32 {
     var r= v;
     r += subgroupShuffleXor(r, 1u);
@@ -113,7 +113,7 @@ fn main(
         src1q[src1q_idx].qs[qs_idx] = q4_quants;
     }

-#if defined(MUL_ACC_Q4_0) || defined(MUL_ACC_Q4_1) || defined(MUL_ACC_Q4_K)
+#if defined(MUL_ACC_Q4_0) || defined(MUL_ACC_Q4_1) || defined(MUL_ACC_Q5_0) || defined(MUL_ACC_Q5_1) || defined(MUL_ACC_Q4_K) || defined(MUL_ACC_Q5_K)
     let q4_quants_sum = dot4I8Packed(q4_quants, 0x01010101u);
     let s = f16(d * f32(cluster_add_i4x8(q4_quants_sum)));

@@ -158,7 +158,7 @@ fn main(
         }
     }

-#if defined(MUL_ACC_Q4_0) || defined(MUL_ACC_Q4_1) || defined(MUL_ACC_Q4_K)
+#if defined(MUL_ACC_Q4_0) || defined(MUL_ACC_Q4_1) || defined(MUL_ACC_Q5_0) || defined(MUL_ACC_Q5_1) || defined(MUL_ACC_Q4_K) || defined(MUL_ACC_Q5_K)

     partial_sums[cluster_id][qs_idx] = dot4I8Packed(q4_quants, 0x01010101u);