#ifndef GGML_SYCL_ESIMD_HPP
#define GGML_SYCL_ESIMD_HPP
#include <sycl/ext/intel/esimd.hpp>
#include "common.hpp"
namespace ggml_sycl_esimd {
constexpr int GGML_SYCL_DMMV_ESIMD_WG_SIZE = 4;
template <ggml_type T> struct esimd_reorder_q_traits;
static ESIMD_INLINE sycl::ext::intel::esimd::simd<float, 32> splat_lo_hi(float lo, float hi) {
using namespace sycl::ext::intel::esimd;
simd<float, 32> v;
v.select<16, 1>(0) = lo;
v.select<16, 1>(16) = hi;
return v;
}
static ESIMD_INLINE void unpack_scale_min_k4(
sycl::ext::intel::esimd::simd<uint8_t, 12> scales, float dall, float dmin,
sycl::ext::intel::esimd::simd<float, 8> & scale_f,
sycl::ext::intel::esimd::simd<float, 8> & min_f) {
using namespace sycl::ext::intel::esimd;
simd<uint8_t, 8> sc = 0;
simd<uint8_t, 8> m = 0;
simd<uint8_t, 4> scale_lo = scales.select<4, 1>(0);
simd<uint8_t, 4> min_lo = scales.select<4, 1>(4);
simd<uint8_t, 4> hi_bits = scales.select<4, 1>(8);
sc.select<4, 1>(0) = scale_lo & simd<uint8_t, 4>(0x3F);
sc.select<4, 1>(4) = (hi_bits & simd<uint8_t, 4>(0x0F)) |
((scale_lo >> simd<uint8_t, 4>(6)) << simd<uint8_t, 4>(4));
m.select<4, 1>(0) = min_lo & simd<uint8_t, 4>(0x3F);
m.select<4, 1>(4) = (hi_bits >> simd<uint8_t, 4>(4)) |
((min_lo >> simd<uint8_t, 4>(6)) << simd<uint8_t, 4>(4));
scale_f = convert<float>(sc) * dall;
min_f = convert<float>(m) * (-dmin);
}
template <> struct esimd_reorder_q_traits<GGML_TYPE_Q3_K> {
struct ptrs {
const uint8_t * qs;
const uint8_t * hmask;
const uint8_t * scales;
const sycl::half * d;
};
static ESIMD_INLINE ptrs make_ptrs(const void * vx, size_t nb) {
const uint8_t * qs = (const uint8_t *) vx;
const uint8_t * hmask = qs + nb * (QK_K / 4);
const uint8_t * scales = hmask + nb * (QK_K / 8);
const sycl::half * d = (const sycl::half *) (scales + nb * 12);
return { qs, hmask, scales, d };
}
static ESIMD_INLINE sycl::ext::intel::esimd::simd<float, 16> unpack_scales(
sycl::ext::intel::esimd::simd<uint8_t, 12> in, float d) {
using namespace sycl::ext::intel::esimd;
simd<uint8_t, 8> lo8 = in.select<8, 1>(0);
simd<uint8_t, 16> code;
code.select<8, 1>(0) = lo8 & simd<uint8_t, 8>(0x0F);
code.select<8, 1>(8) = lo8 >> simd<uint8_t, 8>(4);
simd<uint8_t, 16> hib;
hib.select<4, 1>(0) = in.select<4, 1>(8);
hib.select<4, 1>(4) = in.select<4, 1>(8);
hib.select<4, 1>(8) = in.select<4, 1>(8);
hib.select<4, 1>(12) = in.select<4, 1>(8);
simd<uint8_t, 16> hshift;
hshift.select<4, 1>(0) = 0;
hshift.select<4, 1>(4) = 2;
hshift.select<4, 1>(8) = 4;
hshift.select<4, 1>(12) = 6;
hib = (hib >> hshift) & simd<uint8_t, 16>(0x03);
code = code | (hib << simd<uint8_t, 16>(4));
return (convert<float>(code) - 32.0f) * d;
}
static ESIMD_INLINE void mac_pair(
const ptrs & pa, size_t bia,
const ptrs & pb, size_t bib, bool has_b,
sycl::ext::intel::esimd::simd<float, 256> & y_vec,
sycl::ext::intel::esimd::simd<float, 32> & acc_a,
sycl::ext::intel::esimd::simd<float, 32> & acc_b) {
using namespace sycl::ext::intel::esimd;
simd<uint8_t, 64> qs_a = block_load<uint8_t, 64>(pa.qs + bia * (QK_K / 4));
simd<uint8_t, 64> qs_b = 0;
simd<uint8_t, 32> hmask_a = block_load<uint8_t, 32>(pa.hmask + bia * (QK_K / 8));
simd<uint8_t, 32> hmask_b = 0;
simd<uint8_t, 12> scales_a = block_load<uint8_t, 12>(pa.scales + bia * 12);
simd<uint8_t, 12> scales_b = 0;
const float d_a = (float) pa.d[bia];
float d_b = 0.0f;
if (has_b) {
qs_b = block_load<uint8_t, 64>(pb.qs + bib * (QK_K / 4));
hmask_b = block_load<uint8_t, 32>(pb.hmask + bib * (QK_K / 8));
scales_b = block_load<uint8_t, 12>(pb.scales + bib * 12);
d_b = (float) pb.d[bib];
}
simd<float, 16> scale_f_a = unpack_scales(scales_a, d_a);
simd<float, 16> scale_f_b = unpack_scales(scales_b, d_b);
#pragma unroll
for (int s = 0; s < 8; ++s) {
const int byte_base = 32 * (s / 4);
const uint8_t shift = (uint8_t) (2 * (s % 4));
simd<float, 32> y_s = y_vec.select<32, 1>(s * 32);
simd<uint16_t, 32> q3_a = convert<uint16_t>(
(qs_a.select<32, 1>(byte_base) >> shift) & simd<uint8_t, 32>(3));
q3_a |= convert<uint16_t>(
((hmask_a >> simd<uint8_t, 32>((uint8_t) s)) & simd<uint8_t, 32>(1)) << simd<uint8_t, 32>(2));
simd<uint16_t, 32> q3_b = convert<uint16_t>(
(qs_b.select<32, 1>(byte_base) >> shift) & simd<uint8_t, 32>(3));
q3_b |= convert<uint16_t>(
((hmask_b >> simd<uint8_t, 32>((uint8_t) s)) & simd<uint8_t, 32>(1)) << simd<uint8_t, 32>(2));
simd<float, 32> qf_a = convert<float>(q3_a) - 4.0f;
simd<float, 32> qf_b = convert<float>(q3_b) - 4.0f;
const float scale_a_lo = scale_f_a[2 * s + 0];
const float scale_a_hi = scale_f_a[2 * s + 1];
const float scale_b_lo = scale_f_b[2 * s + 0];
const float scale_b_hi = scale_f_b[2 * s + 1];
simd<float, 32> scale_vec_a = splat_lo_hi(scale_a_lo, scale_a_hi);
simd<float, 32> scale_vec_b = splat_lo_hi(scale_b_lo, scale_b_hi);
simd<float, 32> deq_a = qf_a * scale_vec_a;
simd<float, 32> deq_b = qf_b * scale_vec_b;
acc_a += y_s * deq_a;
acc_b += y_s * deq_b;
}
}
};
template <> struct esimd_reorder_q_traits<GGML_TYPE_Q4_K> {
struct ptrs {
const uint8_t * qs;
const uint8_t * scales;
const sycl::half * dm;
};
static ESIMD_INLINE ptrs make_ptrs(const void * vx, size_t nb) {
const uint8_t * qs = (const uint8_t *) vx;
const uint8_t * scales = qs + nb * (QK_K / 2);
const sycl::half * dm = (const sycl::half *) (scales + nb * K_SCALE_SIZE);
return { qs, scales, dm };
}
static ESIMD_INLINE void mac_pair(
const ptrs & pa, size_t bia,
const ptrs & pb, size_t bib, bool has_b,
sycl::ext::intel::esimd::simd<float, 256> & y_vec,
sycl::ext::intel::esimd::simd<float, 32> & acc_a,
sycl::ext::intel::esimd::simd<float, 32> & acc_b) {
using namespace sycl::ext::intel::esimd;
simd<uint8_t, 128> qs_a = block_load<uint8_t, 128>(pa.qs + bia * (QK_K / 2));
simd<uint8_t, 128> qs_b = 0;
simd<uint8_t, 12> scales_a = block_load<uint8_t, 12>(pa.scales + bia * K_SCALE_SIZE);
simd<uint8_t, 12> scales_b = 0;
const float dall_a = (float) pa.dm[bia * 2 + 0];
const float dmin_a = (float) pa.dm[bia * 2 + 1];
float dall_b = 0.0f;
float dmin_b = 0.0f;
if (has_b) {
qs_b = block_load<uint8_t, 128>(pb.qs + bib * (QK_K / 2));
scales_b = block_load<uint8_t, 12>(pb.scales + bib * K_SCALE_SIZE);
dall_b = (float) pb.dm[bib * 2 + 0];
dmin_b = (float) pb.dm[bib * 2 + 1];
}
simd<float, 8> scale_f_a, min_f_a, scale_f_b, min_f_b;
unpack_scale_min_k4(scales_a, dall_a, dmin_a, scale_f_a, min_f_a);
unpack_scale_min_k4(scales_b, dall_b, dmin_b, scale_f_b, min_f_b);
simd<uint8_t, 128> qs_lo_a = qs_a & simd<uint8_t, 128>(0x0F);
simd<uint8_t, 128> qs_hi_a = qs_a >> simd<uint8_t, 128>(4);
simd<uint8_t, 128> qs_lo_b = qs_b & simd<uint8_t, 128>(0x0F);
simd<uint8_t, 128> qs_hi_b = qs_b >> simd<uint8_t, 128>(4);
#pragma unroll
for (int sb = 0; sb < 8; sb += 2) {
const int q_offset = sb * 16;
simd<float, 32> y_lo = y_vec.select<32, 1>(sb * 32);
simd<float, 32> y_hi = y_vec.select<32, 1>((sb + 1) * 32);
const float scale_a_lo = scale_f_a[sb];
const float scale_a_hi = scale_f_a[sb + 1];
const float min_a_lo = min_f_a[sb];
const float min_a_hi = min_f_a[sb + 1];
const float scale_b_lo = scale_f_b[sb];
const float scale_b_hi = scale_f_b[sb + 1];
const float min_b_lo = min_f_b[sb];
const float min_b_hi = min_f_b[sb + 1];
simd<uint8_t, 32> qa_lo = qs_lo_a.select<32, 1>(q_offset);
simd<uint8_t, 32> qa_hi = qs_hi_a.select<32, 1>(q_offset);
simd<uint8_t, 32> qb_lo = qs_lo_b.select<32, 1>(q_offset);
simd<uint8_t, 32> qb_hi = qs_hi_b.select<32, 1>(q_offset);
simd<float, 32> deq_a_lo = convert<float>(qa_lo) * scale_a_lo + min_a_lo;
simd<float, 32> deq_a_hi = convert<float>(qa_hi) * scale_a_hi + min_a_hi;
simd<float, 32> deq_b_lo = convert<float>(qb_lo) * scale_b_lo + min_b_lo;
simd<float, 32> deq_b_hi = convert<float>(qb_hi) * scale_b_hi + min_b_hi;
acc_a += y_lo * deq_a_lo;
acc_b += y_lo * deq_b_lo;
acc_a += y_hi * deq_a_hi;
acc_b += y_hi * deq_b_hi;
}
}
};
template <> struct esimd_reorder_q_traits<GGML_TYPE_Q6_K> {
struct ptrs {
const uint8_t * ql;
const uint8_t * qh;
const int8_t * scales;
const sycl::half * d;
};
static ESIMD_INLINE ptrs make_ptrs(const void * vx, size_t nb) {
const uint8_t * ql = (const uint8_t *) vx;
const uint8_t * qh = ql + nb * (QK_K / 2);
const int8_t * scales = (const int8_t *) (qh + nb * (QK_K / 4));
const sycl::half * d = (const sycl::half *) (scales + nb * (QK_K / 16));
return { ql, qh, scales, d };
}
static ESIMD_INLINE void mac_pair(
const ptrs & pa, size_t bia,
const ptrs & pb, size_t bib, bool has_b,
sycl::ext::intel::esimd::simd<float, 256> & y_vec,
sycl::ext::intel::esimd::simd<float, 32> & acc_a,
sycl::ext::intel::esimd::simd<float, 32> & acc_b) {
using namespace sycl::ext::intel::esimd;
simd<uint8_t, 128> ql_a = block_load<uint8_t, 128>(pa.ql + bia * (QK_K / 2));
simd<uint8_t, 128> ql_b = 0;
simd<uint8_t, 64> qh_a = block_load<uint8_t, 64>(pa.qh + bia * (QK_K / 4));
simd<uint8_t, 64> qh_b = 0;
simd<int8_t, 16> scales_a = block_load<int8_t, 16>(pa.scales + bia * (QK_K / 16));
simd<int8_t, 16> scales_b = 0;
const float d_a = (float) pa.d[bia];
float d_b = 0.0f;
if (has_b) {
ql_b = block_load<uint8_t, 128>(pb.ql + bib * (QK_K / 2));
qh_b = block_load<uint8_t, 64>(pb.qh + bib * (QK_K / 4));
scales_b = block_load<int8_t, 16>(pb.scales + bib * (QK_K / 16));
d_b = (float) pb.d[bib];
}
simd<float, 16> sc_a = convert<float>(scales_a);
simd<float, 16> sc_b = convert<float>(scales_b);
#pragma unroll
for (int im = 0; im < 2; ++im) {
simd<uint8_t, 32> ql_lo_a = ql_a.select<32, 1>(64 * im);
simd<uint8_t, 32> ql_hi_a = ql_a.select<32, 1>(64 * im + 32);
simd<uint8_t, 32> qh_bits_a = qh_a.select<32, 1>(32 * im);
simd<uint8_t, 32> ql_lo_b = ql_b.select<32, 1>(64 * im);
simd<uint8_t, 32> ql_hi_b = ql_b.select<32, 1>(64 * im + 32);
simd<uint8_t, 32> qh_bits_b = qh_b.select<32, 1>(32 * im);
#pragma unroll
for (int g = 0; g < 4; ++g) {
simd<float, 32> y_g = y_vec.select<32, 1>(32 * (4 * im + g));
const float scale_a_lo = sc_a[8 * im + 2 * g + 0] * d_a;
const float scale_a_hi = sc_a[8 * im + 2 * g + 1] * d_a;
const float scale_b_lo = sc_b[8 * im + 2 * g + 0] * d_b;
const float scale_b_hi = sc_b[8 * im + 2 * g + 1] * d_b;
simd<float, 32> scale_vec_a = splat_lo_hi(scale_a_lo, scale_a_hi);
simd<float, 32> scale_vec_b = splat_lo_hi(scale_b_lo, scale_b_hi);
simd<uint8_t, 32> qa;
simd<uint8_t, 32> qb;
switch (g) {
case 0:
qa = (ql_lo_a & simd<uint8_t, 32>(0x0F)) | ((qh_bits_a & simd<uint8_t, 32>(0x03)) << simd<uint8_t, 32>(4));
qb = (ql_lo_b & simd<uint8_t, 32>(0x0F)) | ((qh_bits_b & simd<uint8_t, 32>(0x03)) << simd<uint8_t, 32>(4));
break;
case 1:
qa = (ql_hi_a & simd<uint8_t, 32>(0x0F)) | ((qh_bits_a & simd<uint8_t, 32>(0x0C)) << simd<uint8_t, 32>(2));
qb = (ql_hi_b & simd<uint8_t, 32>(0x0F)) | ((qh_bits_b & simd<uint8_t, 32>(0x0C)) << simd<uint8_t, 32>(2));
break;
case 2:
qa = (ql_lo_a >> simd<uint8_t, 32>(4)) | (qh_bits_a & simd<uint8_t, 32>(0x30));
qb = (ql_lo_b >> simd<uint8_t, 32>(4)) | (qh_bits_b & simd<uint8_t, 32>(0x30));
break;
default:
qa = (ql_hi_a >> simd<uint8_t, 32>(4)) | ((qh_bits_a & simd<uint8_t, 32>(0xC0)) >> simd<uint8_t, 32>(2));
qb = (ql_hi_b >> simd<uint8_t, 32>(4)) | ((qh_bits_b & simd<uint8_t, 32>(0xC0)) >> simd<uint8_t, 32>(2));
break;
}
simd<float, 32> deq_a = (convert<float>(qa) - 32.0f) * scale_vec_a;
simd<float, 32> deq_b = (convert<float>(qb) - 32.0f) * scale_vec_b;
acc_a += y_g * deq_a;
acc_b += y_g * deq_b;
}
}
}
};
}
#endif