From 92d0bd21857d65544a85ac339f7ceaabb5513263 Mon Sep 17 00:00:00 2001 From: Ruben Ortlam Date: Mon, 24 Aug 2026 15:29:39 +0200 Subject: [PATCH] add q4_1, q5_0, q5_1 support --- ggml/src/ggml-vulkan/ggml-vulkan.cpp | 3 ++ .../vulkan-shaders/mul_mmq_cm1.comp | 44 +++++++++++++++++ .../vulkan-shaders/mul_mmq_cm1_funcs.glsl | 49 ++++++++++++++++++- .../vulkan-shaders/vulkan-shaders-gen.cpp | 2 +- 4 files changed, 95 insertions(+), 3 deletions(-) diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp index 203c6275f5..d2c7dc158a 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp +++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp @@ -4817,6 +4817,9 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { if (device->coopmat_int_support) { CREATE_MMQ2(GGML_TYPE_Q4_0, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_Q4_0], matmul_q4_0_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_push_constants, 3, ); + CREATE_MMQ2(GGML_TYPE_Q4_1, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_Q4_1], matmul_q4_1_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_push_constants, 3, ); + CREATE_MMQ2(GGML_TYPE_Q5_0, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_Q5_0], matmul_q5_0_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_push_constants, 3, ); + CREATE_MMQ2(GGML_TYPE_Q5_1, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_Q5_1], matmul_q5_1_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_push_constants, 3, ); CREATE_MMQ2(GGML_TYPE_Q8_0, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_Q8_0], matmul_q8_0_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_push_constants, 3, ); } diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1.comp b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1.comp index 9a4b0cc5be..4673dbce13 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1.comp @@ -90,6 +90,17 @@ shared float buf_a_d[BM * BK_STEP]; 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) +shared float buf_a_m[BM * BK_STEP]; +shared float buf_b_s[BN * BK_STEP]; +#endif + +#if defined(DATA_A_Q4_1) +#define QUANT_OFFSET 8.0 +#elif defined(DATA_A_Q5_1) +#define QUANT_OFFSET 16.0 +#endif + #define LOAD_VEC_A (4 * QUANT_R) #define LOAD_VEC_B 16 @@ -245,6 +256,14 @@ void main() { ivec4 pre_b_qs[B_LOADS * BK_STEP]; float16_t pre_b_d [B_LOADS * BK_STEP]; +#if defined(DATA_A_Q5_0) || defined(DATA_A_Q5_1) + uint32_t pre_a_qh[A_LOADS * BK_STEP]; +#endif +#if defined(DATA_A_Q4_1) || defined(DATA_A_Q5_1) + float16_t pre_a_m [A_LOADS * BK_STEP]; + float16_t pre_b_s [B_LOADS * BK_STEP]; +#endif + #include "mul_mmq_cm1_funcs.glsl" // Prefetch first block @@ -308,6 +327,22 @@ void main() { } } +#if defined(DATA_A_Q4_1) || defined(DATA_A_Q5_1) + float corr_a[cms_per_row * CM_ELEMS]; + float sum_b[cms_per_col * CM_ELEMS]; + [[unroll]] for (uint r = 0; r < cms_per_row; r++) { + [[unroll]] for (uint e = 0; e < CM_ELEMS; e++) { + corr_a[r * CM_ELEMS + e] = float(QUANT_OFFSET) * scale_a[r * CM_ELEMS + e] + + buf_a_m[ks * BM + a_row0 + r * TM + elem_row[e]]; + } + } + [[unroll]] for (uint c = 0; c < cms_per_col; c++) { + [[unroll]] for (uint e = 0; e < CM_ELEMS; e++) { + sum_b[c * CM_ELEMS + e] = buf_b_s[ks * BN + b_col0 + c * TN + elem_col[e]]; + } + } +#endif + coopmat accs[cms_per_row * cms_per_col]; [[unroll]] for (uint r = 0; r < cms_per_row; r++) { @@ -338,6 +373,10 @@ void main() { sums[tile_idx * CM_ELEMS + e] += ACC_TYPE(float(accs[tile_idx][e]) * scale_a[r * CM_ELEMS + e] * scale_b[c * CM_ELEMS + e]); } +#if defined(DATA_A_Q4_1) || defined(DATA_A_Q5_1) + sums[tile_idx * CM_ELEMS + e] += ACC_TYPE( + corr_a[r * CM_ELEMS + e] * sum_b[c * CM_ELEMS + e]); +#endif } } } @@ -351,6 +390,11 @@ void main() { #undef STORE_BLOCK_TO_LDS #undef B_IB_CALC #undef PREFETCH_BLOCK +#undef PREFETCH_A_QH +#undef PREFETCH_A_M +#undef PREFETCH_B_S +#undef STORE_A_M +#undef STORE_B_S const uint dr = ir * BM + a_row0; const uint dc = ic * BN + b_col0; 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 index fcb914a443..453190fe47 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1_funcs.glsl +++ b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1_funcs.glsl @@ -7,11 +7,51 @@ hi4 = ((hi4 | 0x80808080) - 0x08080808) ^ 0x80808080; \ buf_a_qs[(buf_ib) * QPITCH + (ks) * (BK / 4) + loadr_a ] = lo4; \ buf_a_qs[(buf_ib) * QPITCH + (ks) * (BK / 4) + loadr_a + 4] = hi4; +#elif defined(DATA_A_Q5_0) || defined(DATA_A_Q5_1) +#define STORE_A_QS(buf_ib, ks, idx) \ + uint32_t lo4 = pre_a_qs[idx] & 0x0F0F0F0F; \ + uint32_t hi4 = (pre_a_qs[idx] >> 4) & 0x0F0F0F0F; \ + const uint32_t qh = pre_a_qh[idx]; \ + lo4 |= ((qh >> (4u * loadr_a )) & 0xFu) * 0x02040810u & 0x10101010u; \ + hi4 |= ((qh >> (4u * loadr_a + 16u )) & 0xFu) * 0x02040810u & 0x10101010u; \ + lo4 = ((lo4 | 0x80808080) - 0x10101010) ^ 0x80808080; \ + hi4 = ((hi4 | 0x80808080) - 0x10101010) ^ 0x80808080; \ + buf_a_qs[(buf_ib) * QPITCH + (ks) * (BK / 4) + loadr_a ] = lo4; \ + buf_a_qs[(buf_ib) * QPITCH + (ks) * (BK / 4) + loadr_a + 4] = hi4; #elif defined(DATA_A_Q8_0) #define STORE_A_QS(buf_ib, ks, idx) \ buf_a_qs[(buf_ib) * QPITCH + (ks) * (BK / 4) + loadr_a] = pre_a_qs[idx]; #endif +// Quant-specific extra prefetch helpers (no-ops for types that don't need them) +#if defined(DATA_A_Q5_0) +#define PREFETCH_A_QH(li, ks, ib) \ + pre_a_qh[(li) * BK_STEP + (ks)] = \ + pack32(u16vec2(data_a_packed16[(ib) + (ks)].qh[0], \ + data_a_packed16[(ib) + (ks)].qh[1])); +#elif defined(DATA_A_Q5_1) +#define PREFETCH_A_QH(li, ks, ib) \ + pre_a_qh[(li) * BK_STEP + (ks)] = data_a_packed16[(ib) + (ks)].qh; +#else +#define PREFETCH_A_QH(li, ks, ib) +#endif + +#if defined(DATA_A_Q4_1) || defined(DATA_A_Q5_1) +#define PREFETCH_A_M(li, ks, ib) \ + pre_a_m[(li) * BK_STEP + (ks)] = data_a_packed16[(ib) + (ks)].m; +#define PREFETCH_B_S(li, ks, ib_outer, ib_inner) \ + pre_b_s[(li) * BK_STEP + (ks)] = data_b[(ib_outer)].ds[(ib_inner)].y; +#define STORE_A_M(buf_ib, ks, idx) \ + buf_a_m[(ks) * BM + (buf_ib)] = float(pre_a_m[idx]); +#define STORE_B_S(in_bounds, buf_ib, ks, idx) \ + buf_b_s[(ks) * BN + (buf_ib)] = (in_bounds) ? float(pre_b_s[idx]) : 0.0f; +#else +#define PREFETCH_A_M(li, ks, ib) +#define PREFETCH_B_S(li, ks, ib_outer, ib_inner) +#define STORE_A_M(buf_ib, ks, idx) +#define STORE_B_S(in_bounds, buf_ib, ks, idx) +#endif + #ifdef MUL_MAT_ID #define B_IB_CALC \ const u16vec2 row_idx = row_ids[buf_ib]; \ @@ -33,6 +73,8 @@ pack32(u16vec2(data_a_packed16[ib + ks].qs[loadr_a * 2], \ data_a_packed16[ib + ks].qs[loadr_a * 2 + 1])); \ pre_a_d[li * BK_STEP + ks] = data_a_packed16[ib + ks].d; \ + PREFETCH_A_QH(li, ks, ib) \ + PREFETCH_A_M(li, ks, ib) \ } \ } \ } \ @@ -46,6 +88,7 @@ const uint ib_inner = ib_k % 4; \ pre_b_qs[li * BK_STEP + ks] = data_b[ib_outer].qs[ib_inner * 2 + loadr_b]; \ pre_b_d[li * BK_STEP + ks] = data_b[ib_outer].ds[ib_inner].x; \ + PREFETCH_B_S(li, ks, ib_outer, ib_inner) \ } \ } \ } @@ -59,7 +102,8 @@ const uint idx = li * BK_STEP + ks; \ STORE_A_QS(buf_ib, ks, idx) \ if (loadr_a == 0) { \ - buf_a_d[ks * BM + buf_ib] = float(pre_a_d[idx]); \ + buf_a_d[ks * BM + buf_ib] = float(pre_a_d[idx]); \ + STORE_A_M(buf_ib, ks, idx) \ } \ } \ } \ @@ -77,7 +121,8 @@ buf_b_qs[base + 2] = v.z; \ buf_b_qs[base + 3] = v.w; \ if (loadr_b == 0) { \ - buf_b_d[ks * BN + buf_ib] = in_bounds ? float(pre_b_d[idx]) : 0.0f; \ + buf_b_d[ks * BN + buf_ib] = in_bounds ? float(pre_b_d[idx]) : 0.0f; \ + STORE_B_S(in_bounds, buf_ib, ks, idx) \ } \ } \ } \ 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 d96d2a916d..cb50b5538a 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp @@ -629,7 +629,7 @@ void matmul_shaders(bool fp16, MatMulIdType matmul_id_type, bool coopmat, bool c } #endif - if (coopmat && (tname == "q4_0" || tname == "q8_0")) { + if (coopmat && (tname == "q4_0" || tname == "q4_1" || tname == "q5_0" || tname == "q5_1" || tname == "q8_0")) { 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); } }