[SYCL] Support q2 mul_mat (#26231)

* support q2_0 in mul_mat

* support more q2_0 case
This commit is contained in:
Neo Zhang
2026-07-31 14:19:41 +08:00
committed by GitHub
parent 1553725965
commit a2be61dc87
4 changed files with 174 additions and 0 deletions

View File

@@ -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<QK1_0, QR1_0, dequantize_q1_0>;
case GGML_TYPE_Q2_0:
return dequantize_block_sycl<QK2_0, QR2_0, dequantize_q2_0>;
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<QK1_0, QR1_0, dequantize_q1_0>;
case GGML_TYPE_Q2_0:
return dequantize_block_sycl<QK2_0, QR2_0, dequantize_q2_0>;
case GGML_TYPE_Q4_0:
if (dst->src[0]->extra &&
((ggml_tensor_extra_gpu*)dst->src[0]->extra)->optimized_feature.reorder) {

View File

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

View File

@@ -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<QK2_0, QI2_0, block_q2_0,
VDR_Q2_0_Q8_1_MMVQ, vec_dot_q2_0_q8_1>(
vx, vy, dst, ncols, nrows, item_ct1);
});
});
}
template <int ncols_dst>
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<QK2_0, QI2_0, block_q2_0,
VDR_Q2_0_Q8_1_MMVQ, vec_dot_q2_0_q8_1, ncols_dst>(
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<QK2_0, QI2_0, block_q2_0, VDR_Q2_0_Q8_1_MMVQ, vec_dot_q2_0_q8_1>(
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<QK_K, QI2_K, block_q2_K, VDR_Q2_K_Q8_1_MMVQ, vec_dot_q2_K_q8_1>(
vx_base, vy, ids_dev, dst_base, ncols, nrows, n_experts_used,

View File

@@ -658,6 +658,40 @@ template <> struct reorder_vec_dot_q_sycl<GGML_TYPE_Q6_K> {
#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 <int vdr>
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<float, sycl::rounding_mode::automatic>();
// 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 <int vdr>
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<VDR_Q4_0_Q8_1_MMVQ>(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<VDR_Q2_0_Q8_1_MMVQ>(
v + 0, u + 0, bq2_0->d, bq8_1[0].ds);
const float sum1 = vec_dot_q2_0_q8_1_impl<VDR_Q2_0_Q8_1_MMVQ>(
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) {