33static inline float hsum256_ps_q4k(__m256 v) {
34 __m128 sum = _mm_add_ps(_mm256_castps256_ps128(v), _mm256_extractf128_ps(v, 1));
35 sum = _mm_add_ps(sum, _mm_movehl_ps(sum, sum));
36 sum = _mm_add_ss(sum, _mm_movehdup_ps(sum));
37 return _mm_cvtss_f32(sum);
40static inline __m256i q4k_scale_pair_shuffle(
int pair)
42 const int lo = 2 * pair;
43 return _mm256_set1_epi16((
short)(((lo + 1) << 8) | lo));
46static float dot_q4_k_q8_k_avx2(
const block_q4_K *w,
50 const int nb = k /
QK_K;
51 __m256 acc = _mm256_setzero_ps();
52 __m128 acc_min = _mm_setzero_ps();
53 const __m256i nibble_mask = _mm256_set1_epi8(0x0F);
54 const uint32_t mask_6bit = UINT32_C(0x3f3f3f3f);
55 const uint32_t mask_4bit = UINT32_C(0x0f0f0f0f);
56 const uint32_t mask_2bit = UINT32_C(0x03030303);
58 for (
int i = 0; i < nb; ++i) {
60 memcpy(packed, w[i].scales, 12);
61 packed[3] = ((packed[2] >> 4) & mask_4bit) |
62 (((packed[1] >> 6) & mask_2bit) << 4);
63 const uint32_t upper_scales = packed[1] & mask_6bit;
64 packed[1] = (packed[2] & mask_4bit) |
65 (((packed[0] >> 6) & mask_2bit) << 4);
66 packed[2] = upper_scales;
67 packed[0] &= mask_6bit;
69 const __m256i mins_and_scales = _mm256_cvtepu8_epi16(
70 _mm_set_epi32((
int)packed[3], (
int)packed[2],
71 (
int)packed[1], (
int)packed[0]));
72 const __m128i scale_bytes =
73 _mm256_extracti128_si256(mins_and_scales, 0);
74 const __m256i scales = _mm256_set_m128i(scale_bytes, scale_bytes);
79 const __m128i q8_sums = _mm_hadd_epi16(
80 _mm_loadu_si128((
const __m128i *)&x[i].bsums[0]),
81 _mm_loadu_si128((
const __m128i *)&x[i].bsums[8]));
82 const __m128i min_products = _mm_madd_epi16(
83 _mm256_extracti128_si256(mins_and_scales, 1), q8_sums);
84 acc_min = _mm_fmadd_ps(
85 _mm_set1_ps(dmin), _mm_cvtepi32_ps(min_products), acc_min);
87 __m256i block_sum = _mm256_setzero_si256();
88 for (
int group = 0; group <
QK_K / 64; ++group) {
89 const __m256i packed = _mm256_loadu_si256(
90 (
const __m256i *)&w[i].qs[group * 32]);
91 const __m256i q4_lo = _mm256_and_si256(packed, nibble_mask);
92 const __m256i q4_hi = _mm256_and_si256(
93 _mm256_srli_epi16(packed, 4), nibble_mask);
94 const __m256i q8_lo = _mm256_loadu_si256(
95 (
const __m256i *)&x[i].qs[group * 64]);
96 const __m256i q8_hi = _mm256_loadu_si256(
97 (
const __m256i *)&x[i].qs[group * 64 + 32]);
98 __m256i lo = _mm256_maddubs_epi16(q4_lo, q8_lo);
99 __m256i hi = _mm256_maddubs_epi16(q4_hi, q8_hi);
100 const __m256i scale_lo = _mm256_shuffle_epi8(
101 scales, q4k_scale_pair_shuffle(2 * group));
102 const __m256i scale_hi = _mm256_shuffle_epi8(
103 scales, q4k_scale_pair_shuffle(2 * group + 1));
104 lo = _mm256_madd_epi16(scale_lo, lo);
105 hi = _mm256_madd_epi16(scale_hi, hi);
106 block_sum = _mm256_add_epi32(block_sum, _mm256_add_epi32(lo, hi));
108 acc = _mm256_fmadd_ps(
109 _mm256_set1_ps(d), _mm256_cvtepi32_ps(block_sum), acc);
112 acc_min = _mm_add_ps(acc_min, _mm_movehl_ps(acc_min, acc_min));
113 acc_min = _mm_add_ss(acc_min, _mm_movehdup_ps(acc_min));
114 return hsum256_ps_q4k(acc) + _mm_cvtss_f32(acc_min);
124 if (!y || !W || !x_q8 || M <= 0 || K <= 0) {
130 const int blocks_per_row = K /
QK_K;
132 for (
int row = 0; row < M; ++row) {
133 const block_q4_K *w_row = blocks + (size_t)row * (
size_t)blocks_per_row;
134 y[row] = dot_q4_k_q8_k_avx2(w_row, x, K);
Quantization block structures for weight-only quantization.
#define CK_FP16_TO_FP32(x)
void gemv_q4_k_q8_k_avx2(float *y, const void *W, const void *x_q8, int M, int K)
void gemv_q4_k_q8_k_ref(float *y, const void *W, const void *x_q8, int M, int K)