add q4_1, q5_0, q5_1 support

This commit is contained in:
Ruben Ortlam
2026-08-24 15:29:39 +02:00
parent 8fc44e56be
commit 92d0bd2185
4 changed files with 95 additions and 3 deletions
+3
View File
@@ -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, );
}
@@ -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<int32_t, gl_ScopeSubgroup, TM, TN, gl_MatrixUseAccumulator> 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;
@@ -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) \
} \
} \
} \
@@ -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);
}
}