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);