mirror of
https://github.com/ggml-org/llama.cpp.git
synced 2026-09-18 16:55:05 +02:00
add q4_1, q5_0, q5_1 support
This commit is contained in:
@@ -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);
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user