From a2be61dc87947e069dcdfddcf276abd7d629e924 Mon Sep 17 00:00:00 2001 From: Neo Zhang Date: Fri, 31 Jul 2026 14:19:41 +0800 Subject: [PATCH] [SYCL] Support q2 mul_mat (#26231) * support q2_0 in mul_mat * support more q2_0 case --- ggml/src/ggml-sycl/convert.cpp | 4 ++ ggml/src/ggml-sycl/dequantize.hpp | 22 +++++++++ ggml/src/ggml-sycl/mmvq.cpp | 79 +++++++++++++++++++++++++++++++ ggml/src/ggml-sycl/vecdotq.hpp | 69 +++++++++++++++++++++++++++ 4 files changed, 174 insertions(+) diff --git a/ggml/src/ggml-sycl/convert.cpp b/ggml/src/ggml-sycl/convert.cpp index 060d0aca2d..9ec9276952 100644 --- a/ggml/src/ggml-sycl/convert.cpp +++ b/ggml/src/ggml-sycl/convert.cpp @@ -644,6 +644,8 @@ to_fp16_sycl_t ggml_get_to_fp16_sycl(ggml_type type, ggml_tensor * dst) { switch (type) { case GGML_TYPE_Q1_0: return dequantize_block_sycl; + case GGML_TYPE_Q2_0: + return dequantize_block_sycl; case GGML_TYPE_Q4_0: if (dst->src[0]->extra && ((ggml_tensor_extra_gpu*)dst->src[0]->extra)->optimized_feature.reorder) { @@ -728,6 +730,8 @@ to_fp32_sycl_t ggml_get_to_fp32_sycl(ggml_type type, ggml_tensor *dst) { switch (type) { case GGML_TYPE_Q1_0: return dequantize_block_sycl; + case GGML_TYPE_Q2_0: + return dequantize_block_sycl; case GGML_TYPE_Q4_0: if (dst->src[0]->extra && ((ggml_tensor_extra_gpu*)dst->src[0]->extra)->optimized_feature.reorder) { diff --git a/ggml/src/ggml-sycl/dequantize.hpp b/ggml/src/ggml-sycl/dequantize.hpp index 3db55319fe..876ba1b444 100644 --- a/ggml/src/ggml-sycl/dequantize.hpp +++ b/ggml/src/ggml-sycl/dequantize.hpp @@ -25,6 +25,28 @@ typedef void (*dequantize_kernel_f32_t)(const void * vx, const int64_t ib, const static inline void get_scale_min_k4(int j, const uint8_t * q, uint8_t & d, uint8_t & m); #endif +static __dpct_inline__ void dequantize_q2_0(const void *vx, const int64_t ib, + const int iqs, dfloat2 &v) { + const block_q2_0 * x = (const block_q2_0 *) vx; + + const dfloat d = x[ib].d; + + const int byte_idx = iqs / 4; + const int shift = (iqs % 4) * 2; + const uint8_t vui = x[ib].qs[byte_idx]; + + v.x() = (vui >> shift) & 3; + v.y() = (vui >> (shift + 2)) & 3; + +#ifdef GGML_SYCL_F16 + v.s0() = ((dfloat)v.s0() - 1.0f) * d; + v.s1() = ((dfloat)v.s1() - 1.0f) * d; +#else + v.x() = ((dfloat)v.x() - 1.0f) * d; + v.y() = ((dfloat)v.y() - 1.0f) * d; +#endif // GGML_SYCL_F16 +} + static __dpct_inline__ void dequantize_q4_0(const void *vx, const int64_t ib, const int iqs, dfloat2 &v) { const block_q4_0 * x = (const block_q4_0 *) vx; diff --git a/ggml/src/ggml-sycl/mmvq.cpp b/ggml/src/ggml-sycl/mmvq.cpp index 7b1b3d467f..863d34eabb 100644 --- a/ggml/src/ggml-sycl/mmvq.cpp +++ b/ggml/src/ggml-sycl/mmvq.cpp @@ -1254,6 +1254,66 @@ static void mul_mat_vec_q1_0_q8_1_sycl_switch_ncols( } } +static void mul_mat_vec_q2_0_q8_1_sycl(const void * vx, const void * vy, + float * dst, const int ncols, + const int nrows, + dpct::queue_ptr stream) { + GGML_ASSERT(ncols % QK2_0 == 0); + const int block_num_y = (nrows + GGML_SYCL_MMV_Y - 1) / GGML_SYCL_MMV_Y; + const sycl::range<3> block_nums(1, 1, block_num_y); + const sycl::range<3> block_dims(1, GGML_SYCL_MMV_Y, WARP_SIZE); + + stream->submit([&](sycl::handler & cgh) { + cgh.parallel_for( + sycl::nd_range<3>(block_nums * block_dims, block_dims), + [=](sycl::nd_item<3> item_ct1) [[sycl::reqd_sub_group_size(WARP_SIZE)]] { + mul_mat_vec_q( + vx, vy, dst, ncols, nrows, item_ct1); + }); + }); +} + +template +static void mul_mat_vec_q2_0_q8_1_sycl_ncols( + const void * vx, const void * vy, float * dst, + const int ncols, const int nrows, + const int stride_col_y, const int stride_col_dst, + dpct::queue_ptr stream) { + GGML_ASSERT(ncols % QK2_0 == 0); + const int block_num_y = (nrows + GGML_SYCL_MMV_Y - 1) / GGML_SYCL_MMV_Y; + const sycl::range<3> block_nums(1, 1, block_num_y); + const sycl::range<3> block_dims(1, GGML_SYCL_MMV_Y, WARP_SIZE); + + stream->submit([&](sycl::handler & cgh) { + cgh.parallel_for( + sycl::nd_range<3>(block_nums * block_dims, block_dims), + [=](sycl::nd_item<3> item_ct1) [[sycl::reqd_sub_group_size(WARP_SIZE)]] { + mul_mat_vec_q_ncols( + vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, item_ct1); + }); + }); +} + +static void mul_mat_vec_q2_0_q8_1_sycl_switch_ncols( + const void * vx, const void * vy, float * dst, + const int ncols, const int nrows, const int ncols_dst, + const int stride_col_y, const int stride_col_dst, + dpct::queue_ptr stream) { + switch (ncols_dst) { + case 1: mul_mat_vec_q2_0_q8_1_sycl(vx, vy, dst, ncols, nrows, stream); break; + case 2: mul_mat_vec_q2_0_q8_1_sycl_ncols<2>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break; + case 3: mul_mat_vec_q2_0_q8_1_sycl_ncols<3>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break; + case 4: mul_mat_vec_q2_0_q8_1_sycl_ncols<4>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break; + case 5: mul_mat_vec_q2_0_q8_1_sycl_ncols<5>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break; + case 6: mul_mat_vec_q2_0_q8_1_sycl_ncols<6>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break; + case 7: mul_mat_vec_q2_0_q8_1_sycl_ncols<7>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break; + case 8: mul_mat_vec_q2_0_q8_1_sycl_ncols<8>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break; + default: GGML_ABORT("unsupported ncols_dst=%d for Q2_0 multi-col MMVQ", ncols_dst); + } +} + static void mul_mat_vec_q2_K_q8_1_sycl(const void *vx, const void *vy, float *dst, const int ncols, const int nrows, @@ -2194,6 +2254,20 @@ void ggml_sycl_op_mul_mat_vec_q(ggml_backend_sycl_context & ctx, const ggml_tens mul_mat_vec_q1_0_q8_1_sycl(src0_dd_i, src1_ddq_i_bs, dst_dd_i_bs, ne00, row_diff, stream); } break; + case GGML_TYPE_Q2_0: + if (i == 0 && src1_ncols > 1 && src1_ncols <= 8) { + const int stride_col_y = src1_padded_col_size / QK8_1; + const int stride_col_dst = dst->ne[0]; + GGML_SYCL_DEBUG("Calling mul_mat_vec_q2_0_q8_1_sycl_switch_ncols ncols=%d\n", (int)src1_ncols); + mul_mat_vec_q2_0_q8_1_sycl_switch_ncols( + src0_dd_i, src1_ddq_i, dst_dd_i, ne00, row_diff, + src1_ncols, stride_col_y, stride_col_dst, stream); + return; + } else if (i == 0 || src1_ncols == 1) { + GGML_SYCL_DEBUG("Calling mul_mat_vec_q2_0_q8_1_sycl\n"); + mul_mat_vec_q2_0_q8_1_sycl(src0_dd_i, src1_ddq_i_bs, dst_dd_i_bs, ne00, row_diff, stream); + } + break; case GGML_TYPE_Q2_K: if (i == 0 && src1_ncols > 1 && src1_ncols <= 8) { const int stride_col_y = src1_padded_col_size / QK8_1; @@ -2503,6 +2577,11 @@ bool ggml_sycl_mul_mat_vec_q_id( vx_base, vy, ids_dev, dst_base, ncols, nrows, n_experts_used, expert_weight_stride, dst_row_stride, src1_row_stride, stream); return true; + case GGML_TYPE_Q2_0: + launch_mul_mat_vec_q_moe( + vx_base, vy, ids_dev, dst_base, ncols, nrows, n_experts_used, + expert_weight_stride, dst_row_stride, src1_row_stride, stream); + return true; case GGML_TYPE_Q2_K: launch_mul_mat_vec_q_moe( vx_base, vy, ids_dev, dst_base, ncols, nrows, n_experts_used, diff --git a/ggml/src/ggml-sycl/vecdotq.hpp b/ggml/src/ggml-sycl/vecdotq.hpp index 765fb7f159..c11a6e8f9c 100644 --- a/ggml/src/ggml-sycl/vecdotq.hpp +++ b/ggml/src/ggml-sycl/vecdotq.hpp @@ -658,6 +658,40 @@ template <> struct reorder_vec_dot_q_sycl { #define VDR_Q4_0_Q8_1_MMVQ 2 #define VDR_Q4_0_Q8_1_MMQ 4 +#define VDR_Q2_0_Q8_1_MMVQ 1 + +template +static __dpct_inline__ float vec_dot_q2_0_q8_1_impl( + const int * v, + const int * u, + const float & d2, + const sycl::half2 & ds8) { + int sumi = 0; + +#pragma unroll + for (int i = 0; i < vdr; ++i) { +#pragma unroll + for (int j = 0; j < 4; ++j) { + const uint8_t q = (uint8_t) ((uint32_t) v[i] >> (8 * j)); + + // unpack 2-bit values to byte lanes (0..3), then apply zero-point + // correction with ds8f.y() below, mirroring the q4_0 style. + int vi = 0; + vi |= (((q >> 0) & 0x3) & 0xFF) << 0; + vi |= (((q >> 2) & 0x3) & 0xFF) << 8; + vi |= (((q >> 4) & 0x3) & 0xFF) << 16; + vi |= (((q >> 6) & 0x3) & 0xFF) << 24; + + sumi = dpct::dp4a(vi, u[4 * i + j], sumi); + } + } + + const sycl::float2 ds8f = ds8.convert(); + // q2_0 has zero-point 1. Scale ds8f.y() by processed-lane ratio, + // consistent with q4_0's explicit zero-point subtraction style. + return d2 * (sumi * ds8f.x() - ((float) vdr / (float) QI2_0) * ds8f.y()); +} + template static __dpct_inline__ float vec_dot_q4_0_q8_1_impl(const int * v, const int * u, const float & d4, const sycl::half2 & ds8) { @@ -882,6 +916,41 @@ vec_dot_q4_0_q8_1(const void *__restrict__ vbq, return vec_dot_q4_0_q8_1_impl(v, u, bq4_0->d, bq8_1->ds); } +static __dpct_inline__ float +vec_dot_q2_0_q8_1(const void *__restrict__ vbq, + const block_q8_1 *__restrict__ bq8_1, const int &iqs) { + + const block_q2_0 * bq2_0 = (const block_q2_0 *) vbq; + + int v[2 * VDR_Q2_0_Q8_1_MMVQ]; + int u[8 * VDR_Q2_0_Q8_1_MMVQ]; + +#pragma unroll + for (int i = 0; i < VDR_Q2_0_Q8_1_MMVQ; ++i) { + const int base = 4 * (iqs + i); + + // Q2_0 has QK2_0 = 64 and uses 2 x QK8_1 blocks on the RHS. + v[2 * i + 0] = get_int_from_uint8(bq2_0->qs, iqs + i); + v[2 * i + 1] = get_int_from_uint8(bq2_0->qs, iqs + i + QI2_0); + + u[8 * i + 0] = get_int_from_int8_aligned(bq8_1[0].qs, base + 0); + u[8 * i + 1] = get_int_from_int8_aligned(bq8_1[0].qs, base + 1); + u[8 * i + 2] = get_int_from_int8_aligned(bq8_1[0].qs, base + 2); + u[8 * i + 3] = get_int_from_int8_aligned(bq8_1[0].qs, base + 3); + + u[8 * i + 4] = get_int_from_int8_aligned(bq8_1[1].qs, base + 0); + u[8 * i + 5] = get_int_from_int8_aligned(bq8_1[1].qs, base + 1); + u[8 * i + 6] = get_int_from_int8_aligned(bq8_1[1].qs, base + 2); + u[8 * i + 7] = get_int_from_int8_aligned(bq8_1[1].qs, base + 3); + } + + const float sum0 = vec_dot_q2_0_q8_1_impl( + v + 0, u + 0, bq2_0->d, bq8_1[0].ds); + const float sum1 = vec_dot_q2_0_q8_1_impl( + v + VDR_Q2_0_Q8_1_MMVQ, u + 4 * VDR_Q2_0_Q8_1_MMVQ, bq2_0->d, bq8_1[1].ds); + return sum0 + sum1; +} + static __dpct_inline__ float vec_dot_q4_1_q8_1(const void *__restrict__ vbq, const block_q8_1 *__restrict__ bq8_1, const int &iqs) {