}
}
+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,
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;
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,
#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) {
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) {